Files
Owner/src/owner_voice_pet/config.py
T

350 lines
14 KiB
Python

from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from .models import ErrorCode, ProviderError
@dataclass(frozen=True, slots=True)
class AppConfig:
wake_word: str = "小杰小杰"
sample_rate: int = 16000
channels: int = 1
llm_base_url: str = "https://token-plan-cn.xiaomimimo.com/v1"
llm_api_key: str | None = None
llm_model: str = "mimo-v2.5"
llm_api_style: str = "chat_completions"
llm_stream: bool = True
realtime_transcript_enabled: bool = True
audio_input_device: str | None = None
audio_output_device: str | None = None
asset_dir: Path = Path("assets/pet")
log_dir: Path = Path("logs")
wake_provider: str = "local_kws"
wake_keywords_file: Path | None = None
wake_kws_threshold: float = 0.15
wake_kws_score: float = 1.0
wake_ack_text: str = "我在"
post_playback_drain_ms: int = 0
pipeline_mode: str = "live_turn_based"
endpoint_mode: str = "primary_speaker"
speaker_profile_ms: int = 600
speaker_profile_min_ms: int = 120
speaker_absent_ms: int = 300
speaker_similarity_threshold: float = 0.70
speaker_min_rms: float = 0.012
vad_provider: str = "hybrid"
vad_threshold: float = 0.5
vad_min_duration_ms: int = 250
vad_end_silence_ms: int = 350
vad_no_speech_timeout_ms: int = 5000
vad_max_recording_ms: int = 12000
speech_provider: str = "cloud"
asr_model: str = "mimo-v2.5-asr"
tts_model: str = "mimo-v2.5-tts"
tts_voice: str = "mimo_default"
speech_models_dir: Path = Path("models")
context_mode: str = "session_memory"
context_max_messages: int = 12
context_max_chars: int = 12000
@classmethod
def from_dotenv(cls, path: str | Path = ".env", prefix: str = "OWNER_") -> "AppConfig":
values = parse_dotenv(Path(path))
def get(name: str, default: str | None = None) -> str | None:
value = values.get(f"{prefix}{name}")
return default if value is None or value == "" else value
return cls(
wake_word=get("WAKE_WORD", "小杰小杰") or "小杰小杰",
sample_rate=int(get("SAMPLE_RATE", "16000") or "16000"),
channels=int(get("CHANNELS", "1") or "1"),
llm_base_url=(get("LLM_BASE_URL", "https://token-plan-cn.xiaomimimo.com/v1") or "").rstrip("/"),
llm_api_key=get("LLM_API_KEY"),
llm_model=get("LLM_MODEL", "mimo-v2.5") or "mimo-v2.5",
llm_api_style=get("LLM_API_STYLE", "chat_completions") or "chat_completions",
llm_stream=(get("LLM_STREAM", "1") or "1").lower() not in {"0", "false", "no"},
realtime_transcript_enabled=(get("REALTIME_TRANSCRIPT_ENABLED", "1") or "1").lower()
not in {"0", "false", "no"},
audio_input_device=get("AUDIO_INPUT_DEVICE"),
audio_output_device=get("AUDIO_OUTPUT_DEVICE"),
asset_dir=Path(get("ASSET_DIR", "assets/pet") or "assets/pet"),
log_dir=Path(get("LOG_DIR", "logs") or "logs"),
wake_provider=(get("WAKE_PROVIDER", "local_kws") or "local_kws").lower(),
wake_keywords_file=Path(value) if (value := get("WAKE_KEYWORDS_FILE")) else None,
wake_kws_threshold=float(get("WAKE_KWS_THRESHOLD", "0.15") or "0.15"),
wake_kws_score=float(get("WAKE_KWS_SCORE", "1.0") or "1.0"),
wake_ack_text=get("WAKE_ACK_TEXT", "我在") or "我在",
post_playback_drain_ms=int(get("POST_PLAYBACK_DRAIN_MS", "0") or "0"),
pipeline_mode=(get("PIPELINE_MODE", "live_turn_based") or "live_turn_based").lower(),
endpoint_mode=(get("ENDPOINT_MODE", "primary_speaker") or "primary_speaker").lower(),
speaker_profile_ms=int(get("SPEAKER_PROFILE_MS", "600") or "600"),
speaker_profile_min_ms=int(get("SPEAKER_PROFILE_MIN_MS", "120") or "120"),
speaker_absent_ms=int(get("SPEAKER_ABSENT_MS", "300") or "300"),
speaker_similarity_threshold=float(
get("SPEAKER_SIMILARITY_THRESHOLD", "0.70") or "0.70"
),
speaker_min_rms=float(get("SPEAKER_MIN_RMS", "0.012") or "0.012"),
vad_provider=(get("VAD_PROVIDER", "hybrid") or "hybrid").lower(),
vad_threshold=float(get("VAD_THRESHOLD", "0.5") or "0.5"),
vad_min_duration_ms=int(get("VAD_MIN_DURATION_MS", "250") or "250"),
vad_end_silence_ms=int(get("VAD_END_SILENCE_MS", "350") or "350"),
vad_no_speech_timeout_ms=int(get("VAD_NO_SPEECH_TIMEOUT_MS", "5000") or "5000"),
vad_max_recording_ms=int(get("VAD_MAX_RECORDING_MS", "12000") or "12000"),
speech_provider=(get("SPEECH_PROVIDER", "cloud") or "cloud").lower(),
asr_model=get("ASR_MODEL", "mimo-v2.5-asr") or "mimo-v2.5-asr",
tts_model=get("TTS_MODEL", "mimo-v2.5-tts") or "mimo-v2.5-tts",
tts_voice=get("TTS_VOICE", "mimo_default") or "mimo_default",
speech_models_dir=Path(get("SPEECH_MODELS_DIR", "models") or "models"),
context_mode=(get("CONTEXT_MODE", "session_memory") or "session_memory").lower(),
context_max_messages=int(get("CONTEXT_MAX_MESSAGES", "12") or "12"),
context_max_chars=int(get("CONTEXT_MAX_CHARS", "12000") or "12000"),
)
@classmethod
def from_env(cls, prefix: str = "OWNER_") -> "AppConfig":
return cls.from_dotenv(".env", prefix=prefix)
def require_llm_credentials(self) -> None:
if not self.llm_api_key:
raise ProviderError(
code=ErrorCode.LLM_API_KEY_MISSING,
message="OWNER_LLM_API_KEY is required for cloud LLM calls",
retryable=False,
provider="openai-compatible",
stage="llm",
)
def validate_basic(self) -> list[ProviderError]:
errors: list[ProviderError] = []
if self.sample_rate <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"sample_rate must be positive",
False,
"config",
"startup",
)
)
if self.channels <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"channels must be positive",
False,
"config",
"startup",
)
)
if self.llm_api_style not in {"chat_completions", "responses"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_LLM_API_STYLE must be chat_completions or responses",
False,
"config",
"startup",
)
)
if self.speech_provider not in {"cloud", "local"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_SPEECH_PROVIDER must be cloud or local",
False,
"config",
"startup",
)
)
if self.wake_provider not in {"local_kws"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_WAKE_PROVIDER must be local_kws",
False,
"config",
"startup",
)
)
if self.wake_kws_threshold <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_WAKE_KWS_THRESHOLD must be positive",
False,
"config",
"startup",
)
)
if self.wake_kws_score <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_WAKE_KWS_SCORE must be positive",
False,
"config",
"startup",
)
)
if self.post_playback_drain_ms < 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_POST_PLAYBACK_DRAIN_MS must be non-negative",
False,
"config",
"startup",
)
)
if self.pipeline_mode not in {"live_turn_based"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_PIPELINE_MODE must be live_turn_based",
False,
"config",
"startup",
)
)
if self.endpoint_mode not in {"primary_speaker", "vad"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_ENDPOINT_MODE must be primary_speaker or vad",
False,
"config",
"startup",
)
)
for name, value in {
"OWNER_SPEAKER_PROFILE_MS": self.speaker_profile_ms,
"OWNER_SPEAKER_PROFILE_MIN_MS": self.speaker_profile_min_ms,
"OWNER_SPEAKER_ABSENT_MS": self.speaker_absent_ms,
}.items():
if value <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
f"{name} must be positive",
False,
"config",
"startup",
)
)
if not 0 < self.speaker_similarity_threshold <= 1:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_SPEAKER_SIMILARITY_THRESHOLD must be in (0, 1]",
False,
"config",
"startup",
)
)
if self.speaker_min_rms <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_SPEAKER_MIN_RMS must be positive",
False,
"config",
"startup",
)
)
if self.vad_provider not in {"hybrid", "local", "energy"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_VAD_PROVIDER must be hybrid, local, or energy",
False,
"config",
"startup",
)
)
if self.vad_threshold <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_VAD_THRESHOLD must be positive",
False,
"config",
"startup",
)
)
for name, value in {
"OWNER_VAD_MIN_DURATION_MS": self.vad_min_duration_ms,
"OWNER_VAD_END_SILENCE_MS": self.vad_end_silence_ms,
"OWNER_VAD_NO_SPEECH_TIMEOUT_MS": self.vad_no_speech_timeout_ms,
"OWNER_VAD_MAX_RECORDING_MS": self.vad_max_recording_ms,
}.items():
if value <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
f"{name} must be positive",
False,
"config",
"startup",
)
)
if not self.llm_base_url.startswith(("http://", "https://")):
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_LLM_BASE_URL must start with http:// or https://",
False,
"config",
"startup",
)
)
if self.context_mode not in {"session_memory"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_CONTEXT_MODE must be session_memory",
False,
"config",
"startup",
)
)
return errors
def api_url(self, path: str) -> str:
normalized = path if path.startswith("/") else f"/{path}"
base = self.llm_base_url.rstrip("/")
if base.endswith("/v1") and normalized.startswith("/v1/"):
return base + normalized[3:]
return base + normalized
def parse_dotenv(path: Path) -> dict[str, str]:
if not path.exists():
return {}
values: dict[str, str] = {}
for line_no, raw_line in enumerate(path.read_text(encoding="utf-8").splitlines(), start=1):
line = raw_line.strip()
if not line or line.startswith("#"):
continue
if line.startswith("export "):
line = line.removeprefix("export ").strip()
if "=" not in line:
raise ValueError(f"invalid .env line {line_no}: missing '='")
key, value = line.split("=", 1)
key = key.strip()
value = _strip_dotenv_value(value.strip())
if not key:
raise ValueError(f"invalid .env line {line_no}: empty key")
values[key] = value
return values
def _strip_dotenv_value(value: str) -> str:
if len(value) >= 2 and value[0] == value[-1] and value[0] in {"'", '"'}:
return value[1:-1]
if " #" in value:
return value.split(" #", 1)[0].rstrip()
return value