From d771397ac73826412393e0ff3a5842ddc1d0daa0 Mon Sep 17 00:00:00 2001 From: berges99 Date: Wed, 29 Jul 2026 09:50:22 +0200 Subject: [PATCH] feat(voice): typed VoiceConfig with fail-fast validation and a client override allowlist MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit voice_config was an untyped dict threaded through three merge layers: a typo'd key silently fell back to defaults, and merge_client_voice_overrides applied any client key over server config — recording stayed server-only only because build_voice_session happened to read it from defaults. - Add timbal.voice.config with VoiceConfig + RecordingConfig (extra="forbid"). merge_voice_config validates at server boot, so bad keys/values fail fast; dicts, callables, and VoiceConfig instances are all accepted. - Filter client hello keys through an explicit CLIENT_SETTABLE_VOICE_FIELDS allowlist. recording is structurally excluded; model stays client-settable for the playground picker. Client values skip re-validation on purpose — the per-key guards in build_voice_session degrade instead of 500-ing. - Consume attributes instead of merged.get() in build_voice_session, warmup_voice_stack, and both transports; rding env+user merge now validates through RecordingConfig inside the existing non-fatal try. Also includes the previous (uncommitted) round: voice_config added to codegen AGENT_FIELDS, and TIMBAL_VOICE_LANGUAGE unset now means provider auto-detect instead of defaulting to "es". --- python/tests/codegen/test_set_config.py | 11 + python/tests/server/test_voice_config.py | 120 +++++++---- python/tests/server/test_voice_ws.py | 4 +- .../timbal/codegen/transformers/set_config.py | 1 + python/timbal/server/README.md | 6 +- python/timbal/server/rtc.py | 5 +- python/timbal/server/voice.py | 191 +++++++++++------- python/timbal/voice/__init__.py | 6 + python/timbal/voice/config.py | 73 +++++++ 9 files changed, 296 insertions(+), 121 deletions(-) create mode 100644 python/timbal/voice/config.py diff --git a/python/tests/codegen/test_set_config.py b/python/tests/codegen/test_set_config.py index c04209bd..0ddc2677 100644 --- a/python/tests/codegen/test_set_config.py +++ b/python/tests/codegen/test_set_config.py @@ -99,6 +99,17 @@ def test_set_description(self, workspace): ns = _exec_agent(output) assert ns["agent"].description == "A helpful agent" + def test_set_voice_config(self, workspace): + ws = workspace("""\ + from timbal.core import Agent + + agent = Agent(name="a", model="openai/gpt-4o-mini") + """) + config = json.dumps({"voice_config": {"language": "en", "turn_detector": "local", "tts_extra": {"auto_mode": True}}}) + output = _run_dry(ws, "--config", config) + ns = _exec_agent(output) + assert ns["agent"].voice_config == {"language": "en", "turn_detector": "local", "tts_extra": {"auto_mode": True}} + def test_set_max_iter(self, workspace): ws = workspace("""\ from timbal.core import Agent diff --git a/python/tests/server/test_voice_config.py b/python/tests/server/test_voice_config.py index db71e3ed..8050fa82 100644 --- a/python/tests/server/test_voice_config.py +++ b/python/tests/server/test_voice_config.py @@ -12,10 +12,12 @@ import pytest from fastapi import FastAPI from fastapi.testclient import TestClient +from pydantic import ValidationError from timbal import __version__ as timbal_version from timbal.server import voice as voice_routes from timbal.server.http import create_app, lifespan from timbal.utils import ImportSpec +from timbal.voice.config import DEFAULT_VOICE_ID, RecordingConfig, VoiceConfig from .voice_env import VOICE_ENV_KEYS @@ -24,14 +26,14 @@ class TestDefaultVoiceConfigFromEnv: def test_defaults_when_unset(self) -> None: cfg = voice_routes.default_voice_config_from_env() - assert cfg["stt_provider"] == "elevenlabs" - assert cfg["stt_model"] == "scribe_v2_realtime" - assert cfg["tts_model"] == "eleven_flash_v2_5" - assert cfg["voice"] == voice_routes._DEFAULT_VOICE_ID - assert cfg["language"] == "es" - assert cfg["sample_rate"] == 16_000 - assert cfg["stt_extra"]["commit_strategy"] == "vad" - assert cfg["tts_extra"]["auto_mode"] is True + assert cfg.stt_provider == "elevenlabs" + assert cfg.stt_model == "scribe_v2_realtime" + assert cfg.tts_model == "eleven_flash_v2_5" + assert cfg.voice == DEFAULT_VOICE_ID + assert cfg.language is None # provider auto-detect + assert cfg.sample_rate == 16_000 + assert cfg.stt_extra["commit_strategy"] == "vad" + assert cfg.tts_extra["auto_mode"] is True def test_env_overrides(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("TIMBAL_STT_PROVIDER", "deepgram") @@ -39,18 +41,18 @@ def test_env_overrides(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("TIMBAL_TTS_MODEL", "custom_tts") monkeypatch.setenv("TIMBAL_VOICE_LANGUAGE", "en") cfg = voice_routes.default_voice_config_from_env() - assert cfg["stt_provider"] == "deepgram" - assert cfg["stt_model"] == "custom_stt" - assert cfg["tts_model"] == "custom_tts" - assert cfg["language"] == "en" + assert cfg.stt_provider == "deepgram" + assert cfg.stt_model == "custom_stt" + assert cfg.tts_model == "custom_tts" + assert cfg.language == "en" def test_elevenlabs_voice_id_precedence(self, monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv("ELEVENLABS_VOICE_ID", raising=False) monkeypatch.delenv("TIMBAL_VOICE_ID", raising=False) monkeypatch.setenv("TIMBAL_VOICE_ID", "from_timbal") - assert voice_routes.default_voice_config_from_env()["voice"] == "from_timbal" + assert voice_routes.default_voice_config_from_env().voice == "from_timbal" monkeypatch.setenv("ELEVENLABS_VOICE_ID", "from_el") - assert voice_routes.default_voice_config_from_env()["voice"] == "from_el" + assert voice_routes.default_voice_config_from_env().voice == "from_el" @pytest.mark.usefixtures("clear_voice_env") @@ -60,17 +62,17 @@ class R: pass merged = voice_routes.merge_voice_config(R()) - assert merged["language"] == "es" - assert merged["voice"] == voice_routes._DEFAULT_VOICE_ID + assert merged.language is None + assert merged.voice == DEFAULT_VOICE_ID def test_dict_overrides_top_level(self) -> None: class R: voice_config = {"voice": "v1", "language": "pt"} merged = voice_routes.merge_voice_config(R()) - assert merged["voice"] == "v1" - assert merged["language"] == "pt" - assert merged["stt_model"] == "scribe_v2_realtime" + assert merged.voice == "v1" + assert merged.language == "pt" + assert merged.stt_model == "scribe_v2_realtime" def test_callable_voice_config(self) -> None: class R: @@ -79,44 +81,90 @@ def voice_config(): return {"voice": "callable_v"} merged = voice_routes.merge_voice_config(R()) - assert merged["voice"] == "callable_v" + assert merged.voice == "callable_v" + + def test_voice_config_instance(self) -> None: + class R: + voice_config = VoiceConfig(voice="typed_v", language="de") + + merged = voice_routes.merge_voice_config(R()) + assert merged.voice == "typed_v" + assert merged.language == "de" + assert merged.stt_model == "scribe_v2_realtime" def test_stt_extra_deep_merge(self) -> None: class R: voice_config = {"stt_extra": {"vad_threshold": 0.99}} merged = voice_routes.merge_voice_config(R()) - assert merged["stt_extra"]["commit_strategy"] == "vad" - assert merged["stt_extra"]["vad_threshold"] == 0.99 + assert merged.stt_extra["commit_strategy"] == "vad" + assert merged.stt_extra["vad_threshold"] == 0.99 def test_tts_extra_deep_merge(self) -> None: class R: voice_config = {"tts_extra": {"auto_mode": False}} merged = voice_routes.merge_voice_config(R()) - assert merged["tts_extra"]["auto_mode"] is False + assert merged.tts_extra["auto_mode"] is False def test_none_values_in_runnable_dict_skipped(self) -> None: class R: voice_config = {"voice": None, "language": "it"} merged = voice_routes.merge_voice_config(R()) - assert merged["voice"] == voice_routes._DEFAULT_VOICE_ID - assert merged["language"] == "it" + assert merged.voice == DEFAULT_VOICE_ID + assert merged.language == "it" + + def test_unknown_key_fails_fast(self) -> None: + """A typo'd voice_config key must raise (at server boot), not silently no-op.""" + + class R: + voice_config = {"languag": "en"} + + with pytest.raises(ValidationError, match="languag"): + voice_routes.merge_voice_config(R()) + + def test_recording_dict_is_validated(self) -> None: + class R: + voice_config = {"recording": {"dir": "/tmp/rec", "layout": "sideways"}} + + with pytest.raises(ValidationError, match="layout"): + voice_routes.merge_voice_config(R()) + + def test_turn_detector_instance_passes_through(self) -> None: + sentinel = object() + + class R: + voice_config = {"turn_detector": sentinel} + + assert voice_routes.merge_voice_config(R()).turn_detector is sentinel class TestMergeClientVoiceOverrides: def test_overlay(self) -> None: - base = {"voice": "a", "language": "es", "sample_rate": 16_000} + base = VoiceConfig(voice="a", language="es") out = voice_routes.merge_client_voice_overrides(base, {"language": "en", "voice": "b"}) - assert out["language"] == "en" - assert out["voice"] == "b" - assert out["sample_rate"] == 16_000 + assert out.language == "en" + assert out.voice == "b" + assert out.sample_rate == 16_000 def test_none_skipped(self) -> None: - base = {"voice": "a", "language": "es"} + base = VoiceConfig(voice="a", language="es") out = voice_routes.merge_client_voice_overrides(base, {"language": None}) - assert out["language"] == "es" + assert out.language == "es" + + def test_non_allowlisted_keys_ignored(self) -> None: + base = VoiceConfig(recording=RecordingConfig(dir="/srv/rec")) + out = voice_routes.merge_client_voice_overrides( + base, + {"recording": {"dir": "/tmp/evil"}, "turn_detector": "heuristic", "bogus": 1, "voice": "b"}, + ) + assert out.recording is not None and out.recording.dir == "/srv/rec" + assert out.voice == "b" + + def test_allowlist_is_subset_of_model_fields(self) -> None: + assert voice_routes.CLIENT_SETTABLE_VOICE_FIELDS <= set(VoiceConfig.model_fields) + assert "recording" not in voice_routes.CLIENT_SETTABLE_VOICE_FIELDS class TestLifespanVoiceState: @@ -139,9 +187,9 @@ async def test_lifespan_sets_voice_config_from_runnable( app = FastAPI() spec = ImportSpec.from_fqn(os.environ["TIMBAL_RUNNABLE"]) async with lifespan(app, spec): - assert app.state.voice_config["language"] == "nl" - assert app.state.voice_config["voice"] == "nl_voice" - assert "stt_extra" in app.state.voice_config + assert app.state.voice_config.language == "nl" + assert app.state.voice_config.voice == "nl_voice" + assert app.state.voice_config.stt_extra["commit_strategy"] == "vad" @pytest.mark.asyncio async def test_lifespan_merge_with_env( @@ -159,7 +207,7 @@ async def test_lifespan_merge_with_env( app = FastAPI() spec = ImportSpec.from_fqn(os.environ["TIMBAL_RUNNABLE"]) async with lifespan(app, spec): - assert app.state.voice_config["voice"] == "env_only_voice" + assert app.state.voice_config.voice == "env_only_voice" class TestCreateAppVoiceIntegration: @@ -182,7 +230,7 @@ def test_testclient_startup_sets_voice_config( with TestClient(app) as client: r = client.get("/healthcheck") assert r.status_code == 204 - assert app.state.voice_config["language"] == "sv" + assert app.state.voice_config.language == "sv" def test_voice_page_injects_runnable_meta( self, diff --git a/python/tests/server/test_voice_ws.py b/python/tests/server/test_voice_ws.py index bbd424f0..b0c71a8f 100644 --- a/python/tests/server/test_voice_ws.py +++ b/python/tests/server/test_voice_ws.py @@ -491,7 +491,7 @@ async def start(self, config) -> None: shared = _TrackingDetector() app = create_app() with TestClient(app) as client: - app.state.voice_config = {**(app.state.voice_config or {}), "turn_detector": shared} + app.state.voice_config = app.state.voice_config.model_copy(update={"turn_detector": shared}) for _ in range(2): with client.websocket_connect("/voice/ws") as ws: ws.send_json({}) @@ -518,7 +518,7 @@ def _run_session(self, monkeypatch, tmp_path, hello: dict, server_td=None) -> di app = create_app() with TestClient(app) as client: if server_td is not None: - app.state.voice_config = {**(app.state.voice_config or {}), "turn_detector": server_td} + app.state.voice_config = app.state.voice_config.model_copy(update={"turn_detector": server_td}) with client.websocket_connect("/voice/ws") as ws: ws.send_json(hello) messages = _collect_ws_messages(ws) diff --git a/python/timbal/codegen/transformers/set_config.py b/python/timbal/codegen/transformers/set_config.py index 5e5c2352..7d966286 100644 --- a/python/timbal/codegen/transformers/set_config.py +++ b/python/timbal/codegen/transformers/set_config.py @@ -25,6 +25,7 @@ "api_key", "model_params", # Deprecated, kept for backward compatibility "skills_path", + "voice_config", # dict consumed by the voice server (merge_voice_config) } diff --git a/python/timbal/server/README.md b/python/timbal/server/README.md index a8b9d481..eba7fdf5 100644 --- a/python/timbal/server/README.md +++ b/python/timbal/server/README.md @@ -54,9 +54,9 @@ Whether acks were received is reported per turn in `metrics.playback_acks_receiv ### Config overrides -Optional **first** text frame: a JSON object merged on top of `app.state.voice_config`, which is built at startup from environment defaults and optional `runnable.voice_config` on the loaded agent (`http` lifespan). +Optional **first** text frame: a JSON object merged on top of `app.state.voice_config`, which is built at startup from environment defaults and optional `runnable.voice_config` on the loaded agent (`http` lifespan). Server-side `voice_config` (a dict, zero-arg callable, or `timbal.voice.VoiceConfig`) is validated strictly at startup — an unknown key fails server boot instead of being silently ignored. -Only send keys you need; omitted keys keep server defaults. +Only send keys you need; omitted keys keep server defaults. Client keys are allowlist-filtered (`CLIENT_SETTABLE_VOICE_FIELDS`): the table below plus `model` (per-session LLM override, `"provider/model"`), `turn_timeout_secs`, and `turn_timeout_fallback` (`""` disables the spoken apology). Anything else — notably `recording` — is server policy and is ignored with a log line. | Key | Description | |---------------|-------------| @@ -64,7 +64,7 @@ Only send keys you need; omitted keys keep server defaults. | `stt_model` | Speech-to-text model id (ElevenLabs realtime `scribe_*`, Deepgram `flux-general-en`/`flux-general-multi`/`nova-3*`). Model ids that don't belong to the selected provider are ignored (provider default used). | | `tts_model` | Text-to-speech model id. | | `voice` | ElevenLabs voice id string. | -| `language` | e.g. `"es"`. | +| `language` | e.g. `"es"`. Unset → provider auto-detect. | | `sample_rate` | Hz; STT/TTS audio use this unless extended later. | | `encoding` | Default `"pcm_s16le"`. | | `stt_extra` | Object merged with default STT options (e.g. VAD). | diff --git a/python/timbal/server/rtc.py b/python/timbal/server/rtc.py index cec21d30..34f8f232 100644 --- a/python/timbal/server/rtc.py +++ b/python/timbal/server/rtc.py @@ -32,6 +32,7 @@ from fastapi import APIRouter, Request from fastapi.responses import JSONResponse +from ..voice.config import VoiceConfig from .voice import build_voice_session, event_to_payloads, merge_client_voice_overrides logger = structlog.get_logger("timbal.server.rtc") @@ -92,8 +93,8 @@ async def voice_rtc(request: Request) -> JSONResponse: logger.error("voice_rtc_rejected", reason="runnable is not an Agent", type=type(runnable).__name__) return JSONResponse(status_code=400, content={"error": "Voice requires an Agent runnable."}) - defaults: dict = getattr(request.app.state, "voice_config", None) or {} - sample_rate = int(merge_client_voice_overrides(defaults, config).get("sample_rate", 16_000)) + defaults = getattr(request.app.state, "voice_config", None) or VoiceConfig() + sample_rate = int(merge_client_voice_overrides(defaults, config).sample_rate) downlink = PcmQueueTrack(sample_rate=sample_rate) session, meta = build_voice_session( diff --git a/python/timbal/server/voice.py b/python/timbal/server/voice.py index 9ebe706d..7b4c7bda 100644 --- a/python/timbal/server/voice.py +++ b/python/timbal/server/voice.py @@ -2,10 +2,11 @@ Serves ``GET /voice`` (HTML) and ``/voice/ws`` for the same runnable as ``/run``. Defaults come from :func:`default_voice_config_from_env` and optional -``runnable.voice_config`` (dict or callable); the client can override with a JSON -first message on the socket. +``runnable.voice_config`` (dict, callable, or :class:`VoiceConfig`); the client +can override allowlisted keys with a JSON first message on the socket. -Heavy imports (``VoiceSession``, ElevenLabs) load on first WebSocket connection only. +Heavy imports (``VoiceSession``, ElevenLabs) load on first WebSocket connection +only — ``timbal.voice.config`` itself is import-light. """ from __future__ import annotations @@ -24,6 +25,7 @@ from fastapi.responses import HTMLResponse from .. import __version__ as timbal_version +from ..voice.config import RecordingConfig, VoiceConfig logger = structlog.get_logger("timbal.server.voice") @@ -31,31 +33,21 @@ _HTML_PATH = Path(__file__).parent / "voice.html" -# Override with ELEVENLABS_VOICE_ID / TIMBAL_VOICE_ID (cloned/custom voices are account-specific). -_DEFAULT_VOICE_ID = "1SM7GgM6IMuvQlz2BwM3" - - -def default_voice_config_from_env() -> dict[str, Any]: - """STT/TTS defaults for ``/voice/ws`` (ElevenLabs). Override with env or ``runnable.voice_config``.""" - cfg: dict[str, Any] = { - "stt_provider": os.environ.get("TIMBAL_STT_PROVIDER", "elevenlabs"), - "stt_model": os.environ.get("TIMBAL_STT_MODEL", "scribe_v2_realtime"), - "tts_model": os.environ.get("TIMBAL_TTS_MODEL", "eleven_flash_v2_5"), - "voice": (os.environ.get("ELEVENLABS_VOICE_ID") or os.environ.get("TIMBAL_VOICE_ID") or _DEFAULT_VOICE_ID), - "language": os.environ.get("TIMBAL_VOICE_LANGUAGE", "es"), - "sample_rate": 16_000, - "stt_extra": { - "commit_strategy": "vad", - # 100ms is what ElevenLabs' own realtime examples use. 300ms made - # short replies ("work.", "yes.") transcribe as partials but never - # commit — the session then stalls until the user speaks again. - "min_speech_duration_ms": 100, - "vad_silence_threshold_secs": 1.2, - "vad_threshold": 0.4, - }, - "tts_extra": {"auto_mode": True}, - } - return cfg + +def default_voice_config_from_env() -> VoiceConfig: + """:class:`VoiceConfig` defaults, with env overrides where set.""" + kwargs: dict[str, Any] = {} + if v := os.environ.get("TIMBAL_STT_PROVIDER"): + kwargs["stt_provider"] = v + if v := os.environ.get("TIMBAL_STT_MODEL"): + kwargs["stt_model"] = v + if v := os.environ.get("TIMBAL_TTS_MODEL"): + kwargs["tts_model"] = v + if v := os.environ.get("ELEVENLABS_VOICE_ID") or os.environ.get("TIMBAL_VOICE_ID"): + kwargs["voice"] = v + if v := os.environ.get("TIMBAL_VOICE_LANGUAGE"): + kwargs["language"] = v + return VoiceConfig(**kwargs) def _recording_config_from_env() -> dict[str, Any]: @@ -88,29 +80,69 @@ def _recording_config_from_env() -> dict[str, Any]: ) -def merge_voice_config(runnable: Any) -> dict[str, Any]: - """Env defaults, then optional ``runnable.voice_config`` dict or ``lambda -> dict``.""" +def merge_voice_config(runnable: Any) -> VoiceConfig: + """Env defaults, then optional ``runnable.voice_config`` (dict, callable, or ``VoiceConfig``). + + Validated strictly: an unknown key or bad value raises at server boot + instead of silently falling back to defaults on the first call. + """ base = default_voice_config_from_env() vc = getattr(runnable, "voice_config", None) if callable(vc): vc = vc() + if isinstance(vc, VoiceConfig): + vc = vc.model_dump(include=vc.model_fields_set) if not isinstance(vc, dict): return base skip = frozenset({"stt_extra", "tts_extra"}) - merged = { - **base, + data = { + **base.model_dump(), **{k: v for k, v in vc.items() if v is not None and k not in skip}, } if isinstance(vc.get("stt_extra"), dict): - merged["stt_extra"] = {**base.get("stt_extra", {}), **vc["stt_extra"]} + data["stt_extra"] = {**base.stt_extra, **vc["stt_extra"]} if isinstance(vc.get("tts_extra"), dict): - merged["tts_extra"] = {**base.get("tts_extra", {}), **vc["tts_extra"]} - return merged - - -def merge_client_voice_overrides(server_defaults: dict[str, Any], client: dict[str, Any]) -> dict[str, Any]: - """Apply optional first WebSocket JSON message over ``app.state.voice_config``.""" - return {**server_defaults, **{k: v for k, v in client.items() if v is not None}} + data["tts_extra"] = {**base.tts_extra, **vc["tts_extra"]} + return VoiceConfig(**data) + + +# Keys a browser may override via the WS/RTC hello. Everything else — +# ``recording`` above all — is server policy. ``turn_detector`` is negotiated +# separately in :func:`select_turn_detector_spec` (mode names only, never +# callables). ``model`` is deliberately client-settable for the playground +# model picker; trim this set in deployments where the client is untrusted. +CLIENT_SETTABLE_VOICE_FIELDS = frozenset({ + "stt_provider", + "stt_model", + "tts_model", + "voice", + "language", + "sample_rate", + "encoding", + "stt_extra", + "tts_extra", + "vad_endpointing", + "model", + "turn_timeout_secs", + "turn_timeout_fallback", +}) + + +def merge_client_voice_overrides(server_defaults: VoiceConfig, client: dict[str, Any]) -> VoiceConfig: + """Apply the optional first WebSocket JSON message over the server config. + + Allowlist-filtered. Values are deliberately NOT re-validated here — client + input gets per-key guards with fallbacks in :func:`build_voice_session`, + so one bad value degrades that knob instead of closing the socket. + """ + updates = {k: v for k, v in client.items() if k in CLIENT_SETTABLE_VOICE_FIELDS and v is not None} + ignored = sorted( + k for k, v in client.items() + if v is not None and k not in CLIENT_SETTABLE_VOICE_FIELDS and k != "turn_detector" + ) + if ignored: + logger.info("voice_client_config_ignored", keys=ignored) + return server_defaults.model_copy(update=updates) def runnable_meta_for_voice_page(runnable: Any, import_spec: str) -> dict[str, Any]: @@ -147,7 +179,7 @@ def runnable_meta_for_voice_page(runnable: Any, import_spec: str) -> dict[str, A _VOICE_HTML_META_TOKEN = "__TIMBAL_VOICE_RUNNABLE_META_JSON__" -async def warmup_voice_stack(voice_config: dict[str, Any]) -> None: +async def warmup_voice_stack(voice_config: VoiceConfig) -> None: """Background warmup at server boot so the first voice session starts fast. Two tiers, both best-effort: @@ -186,7 +218,7 @@ def _import_stack() -> None: from ..voice.turn_detection import LocalAudioTurnDetector, resolve_turn_detector from ..voice.vad import SileroVad - td = voice_config.get("turn_detector") + td = voice_config.turn_detector detector = None if isinstance(td, LocalAudioTurnDetector): detector = td @@ -282,7 +314,7 @@ def select_turn_detector_spec(server_spec: Any, client_spec: Any, *, stt_is_flux def build_voice_session( runnable: Any, - defaults: dict[str, Any], + defaults: VoiceConfig, client_config: dict[str, Any], *, playback_tracker: Any = None, @@ -311,8 +343,8 @@ def build_voice_session( merged = merge_client_voice_overrides(defaults, client_config) - stt_provider = merged.get("stt_provider") - stt_model_requested = merged.get("stt_model") + stt_provider = merged.stt_provider + stt_model_requested = merged.stt_model try: stt = resolve_stt(stt_provider, model=stt_model_requested) except ValueError as e: @@ -332,24 +364,27 @@ def build_voice_session( stt_model = effective_stt_model(stt, stt_model_requested) tts = ElevenLabsStreamTTS() - stt_extra = dict(merged.get("stt_extra", {})) + # Client extras are unvalidated (model_copy in the merge): tolerate a + # non-dict rather than 500-ing the socket. + stt_extra = dict(merged.stt_extra) if isinstance(merged.stt_extra, dict) else {} + tts_extra = dict(merged.tts_extra) if isinstance(merged.tts_extra, dict) else {} if stt_is_flux: # Scribe-tuned VAD knobs don't apply to Flux's turn machine. for k in ("commit_strategy", "min_speech_duration_ms", "vad_silence_threshold_secs", "vad_threshold"): stt_extra.pop(k, None) audio_in = AudioInputConfig( model=stt_model, - language=merged.get("language"), - sample_rate=merged.get("sample_rate", 16_000), - encoding=merged.get("encoding", "pcm_s16le"), + language=merged.language, + sample_rate=merged.sample_rate, + encoding=merged.encoding, extra=stt_extra, ) audio_out = AudioOutputConfig( - model=merged.get("tts_model"), - voice=merged.get("voice"), - sample_rate=merged.get("sample_rate", 16_000), - encoding=merged.get("encoding", "pcm_s16le"), - extra=merged.get("tts_extra", {}), + model=merged.tts_model, + voice=merged.voice, + sample_rate=merged.sample_rate, + encoding=merged.encoding, + extra=tts_extra, ) # ``runnable.voice_config`` may supply a TurnDetector instance, factory, or @@ -361,7 +396,7 @@ def build_voice_session( # only here is the STT's endpointing behaviour known. turn_detector = None raw_td = select_turn_detector_spec( - defaults.get("turn_detector"), + defaults.turn_detector, client_config.get("turn_detector"), stt_is_flux=stt_is_flux, ) @@ -377,7 +412,7 @@ def build_voice_session( # Optional bool: force the local VAD endpointing fast path on/off. Default # (absent / non-bool) is auto — on when the detector has an audio EOU model # and timbal[voice] is installed. Client hello may override the server value. - vad_endpointing = merged.get("vad_endpointing") + vad_endpointing = merged.vad_endpointing if not isinstance(vad_endpointing, bool): vad_endpointing = None if stt_is_flux: @@ -386,7 +421,7 @@ def build_voice_session( vad_endpointing = False # Playground / client may override the Agent's LLM for this session only. - raw_model = merged.get("model") + raw_model = merged.model model_override = ( raw_model.strip() if isinstance(raw_model, str) and "/" in raw_model.strip() @@ -400,7 +435,7 @@ def build_voice_session( stt=type(stt).__name__, stt_provider=stt_provider, stt_model=stt_model, - stt_model_requested=merged.get("stt_model"), + stt_model_requested=merged.stt_model, model=llm_model, turn_detector=turn_detector_label, vad_endpointing="auto" if vad_endpointing is None else vad_endpointing, @@ -410,34 +445,34 @@ def build_voice_session( # (list of phrases or `(tool_name) -> str | None` callable) and pass it to # `VoiceSession(filler=...)` once tool-call filler speech returns. session_kwargs: dict[str, Any] = {} - if "turn_timeout_secs" in merged: + if merged.turn_timeout_secs is not None: try: - session_kwargs["turn_timeout_secs"] = float(merged["turn_timeout_secs"]) + session_kwargs["turn_timeout_secs"] = float(merged.turn_timeout_secs) except (TypeError, ValueError): - logger.warning("voice_ws_bad_turn_timeout_secs", value=repr(merged.get("turn_timeout_secs"))) - if "turn_timeout_fallback" in merged: - fb = merged["turn_timeout_fallback"] - session_kwargs["turn_timeout_fallback"] = None if fb is None else str(fb) + logger.warning("voice_ws_bad_turn_timeout_secs", value=repr(merged.turn_timeout_secs)) + if merged.turn_timeout_fallback is not None: + # "" disables the spoken apology; None means "VoiceSession default". + session_kwargs["turn_timeout_fallback"] = str(merged.turn_timeout_fallback) if playback_tracker is not None: session_kwargs["playback_tracker"] = playback_tracker # Call recording is read from *server* config only — env (per session, - # CRIU-safe) under ``runnable.voice_config["recording"]`` (user keys win) - # — never from the merged dict: merge_client_voice_overrides applies - # client keys freely, and a browser must not be able to switch recording - # on or off. - user_recording = defaults.get("recording") - recording_cfg = { + # CRIU-safe) under ``runnable.voice_config["recording"]`` (user keys win). + # ``recording`` is not in CLIENT_SETTABLE_VOICE_FIELDS: a browser must not + # be able to switch recording on or off. + user_recording = defaults.recording + recording_data = { **_recording_config_from_env(), - **(user_recording if isinstance(user_recording, dict) else {}), + **(user_recording.model_dump(include=user_recording.model_fields_set) if user_recording else {}), } - if recording_cfg.get("dir"): + if recording_data.get("dir"): try: from uuid_extensions import uuid7 from ..voice.recording import CallRecorder - on_saved = recording_cfg.get("on_saved") + recording_cfg = RecordingConfig(**recording_data) + on_saved = recording_cfg.on_saved if on_saved is None and os.environ.get("TIMBAL_VOICE_RECORDING_UPLOAD") == "platform": from .recording_upload import platform_recording_upload_hook @@ -445,10 +480,10 @@ def build_voice_session( session_id = uuid7(as_type="str").replace("-", "") session_kwargs["session_id"] = session_id session_kwargs["recorder"] = CallRecorder( - Path(recording_cfg["dir"]) / f"{session_id}.mp3", - sample_rate=int(merged.get("sample_rate", 16_000)), - layout=recording_cfg.get("layout", "mixed"), - bitrate_kbps=int(recording_cfg.get("bitrate_kbps", 32)), + Path(recording_cfg.dir) / f"{session_id}.mp3", + sample_rate=int(merged.sample_rate), + layout=recording_cfg.layout, + bitrate_kbps=recording_cfg.bitrate_kbps, on_saved=on_saved, meta={k: v for env_key, k in _RECORDING_IDENTITY_ENV if (v := os.environ.get(env_key))} or None, ) @@ -607,7 +642,7 @@ async def voice_ws(ws: WebSocket) -> None: except Exception as e: logger.warning("voice_ws_first_frame_error", error=str(e)) - defaults: dict = getattr(ws.app.state, "voice_config", None) or {} + defaults: VoiceConfig = getattr(ws.app.state, "voice_config", None) or VoiceConfig() session, meta = build_voice_session(runnable, defaults, config) meta = {"playback_acks": "recommended", "transport": "websocket", **meta} session.recording_meta = meta diff --git a/python/timbal/voice/__init__.py b/python/timbal/voice/__init__.py index eb58a991..e67552ad 100644 --- a/python/timbal/voice/__init__.py +++ b/python/timbal/voice/__init__.py @@ -1,5 +1,9 @@ """timbal.voice — voice pipeline: VoiceSession, STT/TTS ABCs, turn detection, metrics, and provider implementations.""" +from .config import ( + RecordingConfig, + VoiceConfig, +) from .endpointing import ( VadEndpointer, endpointing_delay, @@ -116,6 +120,7 @@ def __getattr__(name: str): "RealtimeEvent", "RealtimeModel", "RealtimeSession", + "RecordingConfig", "SemanticTurnDetector", "SessionEnded", "SileroVad", @@ -136,6 +141,7 @@ def __getattr__(name: str): "TurnMetricsEvent", "TurnState", "VadEndpointer", + "VoiceConfig", "VoiceSession", "VoiceSessionEvent", "endpointing_delay", diff --git a/python/timbal/voice/config.py b/python/timbal/voice/config.py new file mode 100644 index 00000000..c1cd462c --- /dev/null +++ b/python/timbal/voice/config.py @@ -0,0 +1,73 @@ +"""Typed configuration for the voice server. + +``Agent(voice_config=...)`` — a dict, callable, or :class:`VoiceConfig` — is +validated against this model at server boot, so a typo'd key fails fast +instead of silently falling back to defaults on the first call. Defaults are +the ElevenLabs realtime stack. + +Kept import-light on purpose: the server imports this at module load, while +provider SDKs stay behind ``timbal.voice``'s lazy ``__getattr__``. +""" + +from __future__ import annotations + +from typing import Any, Literal + +from pydantic import BaseModel, ConfigDict, Field + +# Override with ELEVENLABS_VOICE_ID / TIMBAL_VOICE_ID (cloned/custom voices +# are account-specific). +DEFAULT_VOICE_ID = "1SM7GgM6IMuvQlz2BwM3" + + +def _default_stt_extra() -> dict[str, Any]: + return { + "commit_strategy": "vad", + # 100ms is what ElevenLabs' own realtime examples use. 300ms made + # short replies ("work.", "yes.") transcribe as partials but never + # commit — the session then stalls until the user speaks again. + "min_speech_duration_ms": 100, + "vad_silence_threshold_secs": 1.2, + "vad_threshold": 0.4, + } + + +class RecordingConfig(BaseModel): + """Call-recording knobs. Server-side only — never client-settable.""" + + model_config = ConfigDict(extra="forbid") + + dir: str | None = None + layout: Literal["mixed", "split"] = "mixed" + bitrate_kbps: int = 32 + on_saved: Any = None + """Async callable invoked with the ``RecordingResult``. Python-only.""" + + +class VoiceConfig(BaseModel): + """Cross-transport voice session configuration (WS and WebRTC).""" + + model_config = ConfigDict(extra="forbid") + + stt_provider: str = "elevenlabs" + stt_model: str = "scribe_v2_realtime" + tts_model: str = "eleven_flash_v2_5" + voice: str = DEFAULT_VOICE_ID + language: str | None = None + """None → provider auto-detect.""" + sample_rate: int = 16_000 + encoding: str = "pcm_s16le" + stt_extra: dict[str, Any] = Field(default_factory=_default_stt_extra) + tts_extra: dict[str, Any] = Field(default_factory=lambda: {"auto_mode": True}) + turn_detector: Any = None + """Mode name, ``TurnDetector`` instance, or zero-arg factory. + Clients may only send mode names (see ``select_turn_detector_spec``).""" + vad_endpointing: bool | None = None + """None → auto: on when the turn detector exposes an audio EOU model.""" + model: str | None = None + """Per-session LLM override ("provider/model").""" + turn_timeout_secs: float | None = None + """None → ``VoiceSession`` default.""" + turn_timeout_fallback: str | None = None + """None → ``VoiceSession`` default; "" → no spoken apology on timeout.""" + recording: RecordingConfig | None = None