Files
Owner/src/owner_voice_pet/events.py
T

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)