83 lines
2.6 KiB
Python
83 lines
2.6 KiB
Python
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass, field
|
|
from typing import Any, Callable
|
|
|
|
from .models import PipelineState, ProviderError
|
|
|
|
|
|
PIPELINE_STARTED = "pipeline_started"
|
|
WAKE_LISTENING = "wake_listening"
|
|
WAKE_DETECTED = "wake_detected"
|
|
ACK_STARTED = "ack_started"
|
|
QUESTION_PROMPT = "question_prompt"
|
|
CAPTURE_STARTED = "capture_started"
|
|
SPEECH_STARTED = "speech_started"
|
|
SPEECH_ENDED = "speech_ended"
|
|
STT_STARTED = "stt_started"
|
|
TRANSCRIPT_PARTIAL = "transcript_partial"
|
|
TRANSCRIPT_FINAL = "transcript_final"
|
|
LLM_STARTED = "llm_started"
|
|
TTS_STARTED = "tts_started"
|
|
PLAYBACK_FINISHED = "playback_finished"
|
|
STANDBY_RESUMED = "standby_resumed"
|
|
STAGE_ERROR = "stage_error"
|
|
RECOVERING = "recovering"
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class PipelineEvent:
|
|
type: str
|
|
turn_id: int | None = None
|
|
state: PipelineState | None = None
|
|
message: str = ""
|
|
payload: dict[str, Any] = field(default_factory=dict)
|
|
|
|
|
|
class PipelineEventBus:
|
|
def __init__(self) -> None:
|
|
self.events: list[PipelineEvent] = []
|
|
self._listeners: list[Callable[[PipelineEvent], None]] = []
|
|
|
|
def subscribe(self, listener: Callable[[PipelineEvent], None]) -> None:
|
|
self._listeners.append(listener)
|
|
|
|
def emit(
|
|
self,
|
|
event_type: str,
|
|
*,
|
|
turn_id: int | None = None,
|
|
state: PipelineState | None = None,
|
|
message: str = "",
|
|
payload: dict[str, Any] | None = None,
|
|
) -> PipelineEvent:
|
|
event = PipelineEvent(event_type, turn_id=turn_id, state=state, message=message, payload=payload or {})
|
|
self.events.append(event)
|
|
for listener in list(self._listeners):
|
|
listener(event)
|
|
return event
|
|
|
|
|
|
def dispatch_pipeline_event(reporter: Any, event: PipelineEvent) -> None:
|
|
if event.type in {TRANSCRIPT_PARTIAL, TRANSCRIPT_FINAL}:
|
|
reporter.transcript(
|
|
str(event.payload.get("text", event.message)),
|
|
final=event.type == TRANSCRIPT_FINAL,
|
|
turn_id=event.turn_id,
|
|
)
|
|
return
|
|
if event.type == STAGE_ERROR:
|
|
error = event.payload.get("error")
|
|
if isinstance(error, ProviderError):
|
|
reporter.error(error.stage, error.code.value, error.message, turn_id=event.turn_id)
|
|
return
|
|
reporter.error(
|
|
str(event.payload.get("stage", "pipeline")),
|
|
str(event.payload.get("code", "STAGE_ERROR")),
|
|
event.message,
|
|
turn_id=event.turn_id,
|
|
)
|
|
return
|
|
if event.message:
|
|
reporter.status((event.state.value if event.state else event.type), event.message, turn_id=event.turn_id)
|