Files
Owner/src/owner_voice_pet/models.py
T

201 lines
5.9 KiB
Python

from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Mapping
class PipelineState(str, Enum):
IDLE = "idle"
LISTENING = "listening"
WAKE_LISTENING = "wake_listening"
SPEECH_DETECTING = "speech_detecting"
RECORDING = "recording"
TRANSCRIBING = "transcribing"
THINKING = "thinking"
SPEAKING = "speaking"
INTERRUPTED = "interrupted"
TOOL_RUNNING = "tool_running"
RECOVERING = "recovering"
ERROR_RECOVERING = "error_recovering"
class ErrorCode(str, Enum):
AUDIO_INPUT_DEVICE_MISSING = "AUDIO_INPUT_DEVICE_MISSING"
AUDIO_OUTPUT_DEVICE_MISSING = "AUDIO_OUTPUT_DEVICE_MISSING"
AUDIO_PERMISSION_DENIED = "AUDIO_PERMISSION_DENIED"
AUDIO_STREAM_UNDERRUN = "AUDIO_STREAM_UNDERRUN"
AUDIO_FORMAT_UNSUPPORTED = "AUDIO_FORMAT_UNSUPPORTED"
AUDIO_APM_UNAVAILABLE = "AUDIO_APM_UNAVAILABLE"
AUDIO_APM_FORMAT_MISMATCH = "AUDIO_APM_FORMAT_MISMATCH"
AUDIO_APM_PROCESS_FAILED = "AUDIO_APM_PROCESS_FAILED"
AUDIO_BUFFER_OVERRUN = "AUDIO_BUFFER_OVERRUN"
CONFIG_MISSING_VALUE = "CONFIG_MISSING_VALUE"
WAKE_MODEL_MISSING = "WAKE_MODEL_MISSING"
WAKE_MODEL_LOAD_FAILED = "WAKE_MODEL_LOAD_FAILED"
WAKE_AUDIO_FORMAT_INVALID = "WAKE_AUDIO_FORMAT_INVALID"
VAD_MODEL_LOAD_FAILED = "VAD_MODEL_LOAD_FAILED"
VAD_TIMEOUT_NO_SPEECH = "VAD_TIMEOUT_NO_SPEECH"
VAD_MAX_RECORDING_REACHED = "VAD_MAX_RECORDING_REACHED"
STT_MODEL_MISSING = "STT_MODEL_MISSING"
STT_TRANSCRIBE_FAILED = "STT_TRANSCRIBE_FAILED"
STT_EMPTY_TRANSCRIPT = "STT_EMPTY_TRANSCRIPT"
LLM_API_KEY_MISSING = "LLM_API_KEY_MISSING"
LLM_REQUEST_TIMEOUT = "LLM_REQUEST_TIMEOUT"
LLM_RATE_LIMITED = "LLM_RATE_LIMITED"
LLM_NETWORK_ERROR = "LLM_NETWORK_ERROR"
LLM_EMPTY_REPLY = "LLM_EMPTY_REPLY"
TTS_MODEL_MISSING = "TTS_MODEL_MISSING"
TTS_SYNTHESIS_FAILED = "TTS_SYNTHESIS_FAILED"
TTS_EMPTY_AUDIO = "TTS_EMPTY_AUDIO"
NOISE_FILTER_MODEL_MISSING = "NOISE_FILTER_MODEL_MISSING"
NOISE_FILTER_FAILED = "NOISE_FILTER_FAILED"
ASSET_MISSING = "ASSET_MISSING"
VALIDATION_FAILED = "VALIDATION_FAILED"
@dataclass(slots=True)
class ProviderError(Exception):
code: ErrorCode
message: str
retryable: bool
provider: str
stage: str
def __str__(self) -> str:
return f"{self.code.value} [{self.stage}/{self.provider}]: {self.message}"
@dataclass(frozen=True, slots=True)
class AudioFrame:
pcm: bytes
sample_rate: int
channels: int
timestamp_ms: int
frame_id: int
metadata: Mapping[str, Any] = field(default_factory=dict)
def __post_init__(self) -> None:
if self.sample_rate <= 0:
raise ValueError("sample_rate must be positive")
if self.channels <= 0:
raise ValueError("channels must be positive")
if self.timestamp_ms < 0:
raise ValueError("timestamp_ms must be non-negative")
if self.frame_id < 0:
raise ValueError("frame_id must be non-negative")
@property
def duration_ms(self) -> int:
metadata_duration = self.metadata.get("duration_ms")
if isinstance(metadata_duration, int):
return metadata_duration
bytes_per_sample = 2
if self.channels <= 0:
return 0
sample_count = len(self.pcm) // (bytes_per_sample * self.channels)
return round(sample_count * 1000 / self.sample_rate)
def to_fixture(self) -> dict[str, Any]:
return {
"pcm_hex": self.pcm.hex(),
"sample_rate": self.sample_rate,
"channels": self.channels,
"timestamp_ms": self.timestamp_ms,
"frame_id": self.frame_id,
"metadata": dict(self.metadata),
}
@classmethod
def from_fixture(cls, data: Mapping[str, Any]) -> "AudioFrame":
return cls(
pcm=bytes.fromhex(str(data["pcm_hex"])),
sample_rate=int(data["sample_rate"]),
channels=int(data["channels"]),
timestamp_ms=int(data["timestamp_ms"]),
frame_id=int(data["frame_id"]),
metadata=dict(data.get("metadata") or {}),
)
@dataclass(frozen=True, slots=True)
class AudioSegment:
pcm: bytes
sample_rate: int
channels: int
start_time_ms: int
end_time_ms: int
metadata: Mapping[str, Any] = field(default_factory=dict)
def __post_init__(self) -> None:
if self.sample_rate <= 0:
raise ValueError("sample_rate must be positive")
if self.channels <= 0:
raise ValueError("channels must be positive")
if self.start_time_ms < 0:
raise ValueError("start_time_ms must be non-negative")
if self.end_time_ms < self.start_time_ms:
raise ValueError("end_time_ms must be >= start_time_ms")
@property
def duration_ms(self) -> int:
return self.end_time_ms - self.start_time_ms
@dataclass(frozen=True, slots=True)
class WakeEvent:
keyword: str
confidence: float
timestamp_ms: int
@dataclass(frozen=True, slots=True)
class VadResult:
is_speech: bool
confidence: float
speech_ms: int
silence_ms: int
end_reason: str | None = None
@dataclass(frozen=True, slots=True)
class Transcript:
text: str
language: str
confidence: float | None
duration_ms: int
provider: str
raw_metadata: Mapping[str, Any] = field(default_factory=dict)
@property
def normalized_text(self) -> str:
return self.text.strip()
@dataclass(frozen=True, slots=True)
class Message:
role: str
content: str
created_at: float
@dataclass(frozen=True, slots=True)
class ReplyDelta:
text_delta: str
is_sentence_boundary: bool = False
finish_reason: str | None = None
@dataclass(frozen=True, slots=True)
class PlaybackResult:
played: bool
duration_ms: int
error: ProviderError | None = None
@dataclass(frozen=True, slots=True)
class TransportHealth:
input_available: bool
output_available: bool
message: str = ""