"""Stream a PCM WAV to Kendr Voice and save the received speech as a WAV.

Python 3.11+, pip install 'websockets>=14,<16'
KENDR_API_KEY must be a server-side key with models:invoke.
python voice_stream.py question.wav --voice tiffany --output reply.wav

This deliberately makes ONE paid session; it never retries or renews itself.
It captures received audio for inspection, rather than playing it live.
"""
from __future__ import annotations

import argparse
import asyncio
from contextlib import suppress
import json
import math
import os
from pathlib import Path
import sys
import wave

FRAME_BYTES = 1024
FRAME_SECONDS = 0.032  # 512 signed 16-bit samples / 16,000 samples per second
STREAM_URL = "wss://api.kendr.org/v1/voice/stream"


def load_pcm(path: Path) -> bytes:
    with wave.open(str(path), "rb") as source:
        if (source.getnchannels(), source.getsampwidth(), source.getframerate(), source.getcomptype()) != (1, 2, 16000, "NONE"):
            raise ValueError("Input must be an uncompressed 16 kHz, mono, 16-bit PCM WAV.")
        if source.getnframes() > 16000 * 400:
            raise ValueError("This single-session example accepts at most 400 seconds of input.")
        pcm = source.readframes(source.getnframes())
    if not pcm:
        raise ValueError("The input WAV is empty.")
    return pcm


def pcm_frames(pcm: bytes):
    if len(pcm) % 2:
        raise ValueError("PCM must contain complete 16-bit samples.")
    for offset in range(0, len(pcm), FRAME_BYTES):
        yield pcm[offset:offset + FRAME_BYTES].ljust(FRAME_BYTES, b"\0")


async def send_audio(ws, pcm: bytes, listen_seconds: float) -> None:
    """Pace frames without catching up in a burst after a slow send."""
    for frame in pcm_frames(pcm):
        await ws.send(frame)
        await asyncio.sleep(FRAME_SECONDS)
    # Silence lets the model detect the end of the spoken turn and answer.
    # `end` ends the whole session; it is not an end-of-utterance marker.
    for _ in range(math.ceil(listen_seconds / FRAME_SECONDS)):
        await ws.send(bytes(FRAME_BYTES))
        await asyncio.sleep(FRAME_SECONDS)
    await ws.send(json.dumps({"type": "end"}))


async def exchange(ws, pcm: bytes, output, *, voice: str, listen_seconds: float, report=print) -> dict:
    """Transport-injected example: receive while the audio sender runs."""
    await ws.send(json.dumps({"type": "start", "mode": "full", "voice_id": voice, "tools_enabled": False}))
    sender = None
    ready = False
    ended = False
    # Covers the configurable start timeout, then the server's advertised duration.
    deadline = asyncio.get_running_loop().time() + 65
    try:
        while True:
            remaining = deadline - asyncio.get_running_loop().time()
            if remaining <= 0:
                raise TimeoutError("Timed out waiting for Kendr Voice; session was not completed.")
            receive = asyncio.create_task(ws.recv())
            # Notice producer failures immediately instead of waiting for the whole session.
            pending = {receive}
            if sender is not None:
                pending.add(sender)
            try:
                done, _ = await asyncio.wait(pending, timeout=remaining, return_when=asyncio.FIRST_COMPLETED)
                if not done:
                    raise TimeoutError("Kendr Voice did not finish before the deadline.")
                if sender is not None and sender in done:
                    sender.result()
                    sender = None
                    deadline = min(deadline, asyncio.get_running_loop().time() + 10)
                if receive not in done:
                    continue
                message = receive.result()
            finally:
                if not receive.done():
                    receive.cancel()
                    with suppress(asyncio.CancelledError):
                        await receive
            if isinstance(message, bytes):
                if not ready or len(message) % 2:
                    raise ValueError("Unexpected PCM frame received from the server.")
                output.writeframesraw(message)
                continue
            event = json.loads(message)
            kind = event.get("type")
            if kind == "session.ready":
                if ready or event.get("protocol") != "kendr.voice.v1":
                    raise ValueError("Unsupported or duplicate session.ready event.")
                for field, rate in (("input_audio", 16000), ("output_audio", 24000)):
                    if event.get(field) != {"encoding": "pcm_s16le", "sample_rate_hz": rate, "channels": 1}:
                        raise ValueError("Unsupported server audio format.")
                duration = event["max_duration_seconds"]
                if math.ceil(len(pcm) / FRAME_BYTES) * FRAME_SECONDS + listen_seconds + 2 >= duration:
                    raise ValueError("Recording plus listening time exceeds the advertised session duration; shorten it.")
                ready = True
                report("Voice session ready: " + event["session_id"])
                deadline = asyncio.get_running_loop().time() + duration + 10
                sender = asyncio.create_task(send_audio(ws, pcm, listen_seconds))
            elif kind == "transcript" and event.get("final"):
                # Final text is the whole utterance, not another delta.
                report(f"{event['role']}: {event['text']}")
            elif kind == "error":
                report(f"Voice error {event.get('code')}: {event.get('message')}")
                if event.get("fatal"):
                    raise RuntimeError("Voice session rejected; no automatic retry was attempted.")
                # Nonfatal errors can precede a renewable session.ended. Keep reading.
            elif kind == "session.ended":
                ended = True
                report("Session ended: " + event["reason"])
                report("Usage (not a settled credit receipt): " + json.dumps(event.get("usage", {})))
                if event.get("renewable"):
                    report("Renewal is available; this example does not start another paid session.")
                return event
            elif kind == "playback.clear":
                # A live player MUST stop and clear queued playback here. This example
                # archives already received audio, so there is no playback queue.
                pass
    finally:
        if sender is not None:
            sender.cancel()
            with suppress(asyncio.CancelledError, Exception):
                await sender
        if not ended:
            with suppress(Exception):
                await asyncio.wait_for(ws.send(json.dumps({"type": "end"})), 2)


async def run(args) -> None:
    from websockets.asyncio.client import connect

    key = os.environ.get("KENDR_API_KEY", "").strip()
    if not key:
        raise ValueError("Set KENDR_API_KEY in the server environment.")
    if args.input.resolve() == args.output.resolve():
        raise ValueError("Input and output must be different files.")
    if not math.isfinite(args.listen_seconds) or not 1 <= args.listen_seconds <= 60:
        raise ValueError("--listen-seconds must be between 1 and 60.")
    pcm = load_pcm(args.input)
    # No Origin is required for a server client. The key never enters the URL.
    async with connect(STREAM_URL, additional_headers={"Authorization": "Bearer " + key},
                       open_timeout=20, close_timeout=5, max_size=128 * 1024) as ws:
        with wave.open(str(args.output), "wb") as output:
            output.setnchannels(1)
            output.setsampwidth(2)
            output.setframerate(24000)
            await exchange(ws, pcm, output, voice=args.voice, listen_seconds=args.listen_seconds)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("input", type=Path)
    parser.add_argument("--voice", default="tiffany", help="ID from GET /v1/voice/voices")
    parser.add_argument("--output", type=Path, default=Path("reply.wav"))
    parser.add_argument("--listen-seconds", type=float, default=10, help="Send silence while receiving a reply, then end")
    args = parser.parse_args()
    try:
        asyncio.run(run(args))
    except KeyboardInterrupt:
        raise SystemExit(130) from None
    except Exception as error:
        # Avoid dumping handshake headers, credentials, or continuation tokens.
        print(f"Voice session failed ({type(error).__name__}); check credentials, scopes, credits, input format, and service availability.", file=sys.stderr)
        raise SystemExit(1) from None


if __name__ == "__main__":
    main()
