#!/usr/bin/env python3 """Run the reusable Urdu S2S service with live reply and speech-text providers.""" from __future__ import annotations import argparse import json from pathlib import Path import sys from typing import Any ROOT = Path(__file__).resolve().parents[1] for path in (ROOT / "src", ROOT / "scripts"): if str(path) not in sys.path: sys.path.insert(0, str(path)) from urdu_s2s.asr_providers import FasterWhisperASRProvider # noqa: E402 from urdu_s2s.live_providers import ( # noqa: E402 OpenAIBridgeProvider, OpenAICompatibleChatClient, OpenAIReplyProvider, ) from urdu_s2s.pipeline import SpeechToSpeechPipeline # noqa: E402 from urdu_s2s.schemas import ( # noqa: E402 ASRResult, BridgeResult, ReplyResult, SpeechToSpeechRequest, SpeechToSpeechResult, TTSResult, ) from urdu_s2s.tts_providers import ChatterboxPraxyTTSProvider # noqa: E402 from urdu_s2s.tracing import to_jsonable # noqa: E402 DEFAULT_PRAXY_ANCHOR = ROOT / "data/processed/voice_anchors/chatterbox_praxy_v1/bench_025.wav" class TranscriptASRProvider: """Temporary ASR adapter for live S2S testing from a known transcript.""" def __init__(self, transcript: str) -> None: self.transcript = transcript def transcribe(self, request: SpeechToSpeechRequest) -> ASRResult: return ASRResult( text=self.transcript, provider="provided_transcript", model="manual_asr_transcript", language="ur", ) class PlaceholderTTSProvider: """Temporary TTS adapter until the live Chatterbox/Praxy provider is wired.""" def __init__(self, audio_path: Path) -> None: self.audio_path = audio_path def synthesize( self, request: SpeechToSpeechRequest, reply: ReplyResult, bridge: BridgeResult, ) -> TTSResult: return TTSResult( audio_path=self.audio_path, provider="placeholder_tts", model="not_synthesized_yet", ) def build_text_live_pipeline( *, audio_path: Path, request_id: str, asr_transcript: str, prompt_roman_urdu: str, tts_audio_path: Path, chat_client: OpenAICompatibleChatClient, ) -> tuple[SpeechToSpeechPipeline, SpeechToSpeechRequest]: return build_live_pipeline( audio_path=audio_path, request_id=request_id, asr_provider_name="transcript", asr_transcript=asr_transcript, prompt_roman_urdu=prompt_roman_urdu, tts_provider_name="placeholder", tts_audio_path=tts_audio_path, chat_client=chat_client, ) def build_live_pipeline( *, audio_path: Path, request_id: str, asr_provider_name: str, asr_transcript: str, prompt_roman_urdu: str, tts_provider_name: str = "placeholder", tts_audio_path: Path, voice_prompt_audio_path: Path = DEFAULT_PRAXY_ANCHOR, chat_client: OpenAICompatibleChatClient, whisper_model_factory: Any | None = None, whisper_model: str = "large-v3", whisper_language: str = "ur", whisper_device: str = "cpu", whisper_compute_type: str = "int8", chatterbox_model_loader: Any | None = None, chatterbox_wav_writer: Any | None = None, chatterbox_duration_reader: Any | None = None, chatterbox_device: str = "cuda", chatterbox_t3_model: str = "v3", ) -> tuple[SpeechToSpeechPipeline, SpeechToSpeechRequest]: if asr_provider_name == "transcript": if not asr_transcript.strip(): raise ValueError("--asr-transcript is required when --asr-provider transcript") asr_provider = TranscriptASRProvider(asr_transcript) elif asr_provider_name == "faster_whisper": asr_provider = FasterWhisperASRProvider( model_name=whisper_model, language=whisper_language, device=whisper_device, compute_type=whisper_compute_type, model_factory=whisper_model_factory, ) else: raise ValueError(f"Unknown ASR provider: {asr_provider_name}") if tts_provider_name == "placeholder": tts_provider = PlaceholderTTSProvider(tts_audio_path) elif tts_provider_name == "chatterbox_praxy": tts_provider = ChatterboxPraxyTTSProvider( output_audio_path=tts_audio_path, voice_prompt_audio_path=voice_prompt_audio_path, device=chatterbox_device, t3_model=chatterbox_t3_model, model_loader=chatterbox_model_loader, wav_writer=chatterbox_wav_writer, duration_reader=chatterbox_duration_reader, ) else: raise ValueError(f"Unknown TTS provider: {tts_provider_name}") pipeline = SpeechToSpeechPipeline( asr_provider=asr_provider, reply_provider=OpenAIReplyProvider(chat_client=chat_client), bridge_provider=OpenAIBridgeProvider(chat_client=chat_client), tts_provider=tts_provider, ) request = SpeechToSpeechRequest( request_id=request_id, audio_path=audio_path, metadata={"prompt_roman_urdu": prompt_roman_urdu}, ) return pipeline, request def result_to_payload(result: SpeechToSpeechResult) -> dict[str, object]: return { "request_id": result.request.request_id, "input_audio_path": str(result.request.audio_path), "asr_transcript": result.asr.text, "assistant_reply_urdu": result.reply.text_urdu, "devanagari_tts_text": result.bridge.text_devanagari, "tts_audio_path": str(result.tts.audio_path), "trace": to_jsonable(result.trace), } def write_json_result(payload: dict[str, object], output_path: Path) -> None: output_path.parent.mkdir(parents=True, exist_ok=True) output_path.write_text( json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8", ) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--audio-path", required=True, type=Path) parser.add_argument( "--asr-provider", choices=["transcript", "faster_whisper"], default="transcript", ) parser.add_argument( "--asr-transcript", default="", help="Temporary transcript input until a live ASR provider is wired.", ) parser.add_argument("--whisper-model", default="large-v3") parser.add_argument("--whisper-language", default="ur") parser.add_argument("--whisper-device", default="cpu") parser.add_argument("--whisper-compute-type", default="int8") parser.add_argument("--prompt-roman-urdu", default="") parser.add_argument("--request-id", default="") parser.add_argument("--model", default="") parser.add_argument("--base-url", default="") parser.add_argument("--output-json", type=Path, default=ROOT / "reports/evals/s2s_live_result.json") parser.add_argument( "--tts-provider", choices=["placeholder", "chatterbox_praxy"], default="placeholder", ) parser.add_argument( "--tts-audio-path", type=Path, default=ROOT / "reports/evals/s2s_live_placeholder_tts.wav", help="Response WAV path for Chatterbox, or placeholder path in placeholder mode.", ) parser.add_argument("--voice-prompt-audio-path", type=Path, default=DEFAULT_PRAXY_ANCHOR) parser.add_argument("--chatterbox-device", default="cuda") parser.add_argument("--chatterbox-t3-model", default="v3") return parser.parse_args() def main() -> int: args = parse_args() request_id = args.request_id or args.audio_path.stem chat_client = OpenAICompatibleChatClient( base_url=args.base_url or None, model=args.model or None, ) pipeline, request = build_live_pipeline( audio_path=args.audio_path, request_id=request_id, asr_provider_name=args.asr_provider, asr_transcript=args.asr_transcript, prompt_roman_urdu=args.prompt_roman_urdu, tts_provider_name=args.tts_provider, tts_audio_path=args.tts_audio_path, voice_prompt_audio_path=args.voice_prompt_audio_path, chat_client=chat_client, whisper_model=args.whisper_model, whisper_language=args.whisper_language, whisper_device=args.whisper_device, whisper_compute_type=args.whisper_compute_type, chatterbox_device=args.chatterbox_device, chatterbox_t3_model=args.chatterbox_t3_model, ) payload = result_to_payload(pipeline.run(request)) write_json_result(payload, args.output_json) print(json.dumps(payload, ensure_ascii=False, indent=2)) print(f"wrote={args.output_json}") return 0 if __name__ == "__main__": raise SystemExit(main())