#!/usr/bin/env python3
"""
Podcast TTS — custom MCP server (streamable HTTP) for Claude.

Exposes one tool, generate_podcast_audio: takes a (possibly segmented) spoken
script, synthesizes it with Gemini TTS, stitches the audio with real silences
at segment boundaries, encodes an mp3 with ffmpeg, uploads it to a public GCS
bucket, and returns ONLY metadata: {url, duration_seconds, size_bytes}.

The audio itself never travels through the MCP channel — that is the whole
point of this design: tool results stay tiny, storage does the heavy lifting.

Env vars:
  GEMINI_API_KEY    required — aistudio.google.com API key
  GCS_BUCKET        required — public bucket for the mp3s
  SECRET_PATH       recommended — random string; server only answers under /<SECRET_PATH>/mcp
  GEMINI_TTS_MODEL  default "gemini-3.1-flash-tts-preview"
  DEFAULT_VOICE     default "Charon"
"""
import base64
import io
import json
import os
import re
import subprocess
import tempfile
import time

import httpx
from google.cloud import storage
from mcp.server.fastmcp import FastMCP
from mcp.server.transport_security import TransportSecuritySettings

GEMINI_API_KEY = os.environ.get("GEMINI_API_KEY", "")
GCS_BUCKET = os.environ.get("GCS_BUCKET", "")
SECRET_PATH = os.environ.get("SECRET_PATH", "").strip("/")
MODEL = os.environ.get("GEMINI_TTS_MODEL", "gemini-3.1-flash-tts-preview")
DEFAULT_VOICE = os.environ.get("DEFAULT_VOICE", "Achird")

DIALOGUE_VOICE_A = os.environ.get("DIALOGUE_VOICE_A", "Puck")
DIALOGUE_VOICE_B = os.environ.get("DIALOGUE_VOICE_B", "Kore")

SAMPLE_RATE = 24000  # Gemini TTS returns 16-bit mono PCM at 24 kHz
BYTES_PER_SEC = SAMPLE_RATE * 2
MAX_CHUNK_CHARS = 2200  # smaller requests: the preview model returns empty
                        # responses more often as the payload grows

# A dialogue line looks like "BUYER: November Santos, what have you got?"
SPEAKER_RE = re.compile(r"^([A-Z][A-Z0-9 _-]{1,14}):\s*(.+)$")
GEMINI_URL = (
    "https://generativelanguage.googleapis.com/v1beta/models/"
    "{model}:generateContent?key={key}"
)

# DNS-rebinding protection is designed for local servers ("Host must be
# localhost"); on Cloud Run the Host is the public run.app domain, so the
# default check rejects every request with 421. The URL secret path is this
# server's actual access control, so we disable the host check explicitly.
mcp = FastMCP(
    "podcast-tts",
    stateless_http=True,
    transport_security=TransportSecuritySettings(enable_dns_rebinding_protection=False),
)


# ----------------------------------------------------------------- script ----

def parse_segments(script: str):
    """Accepts 'TEXT ||| PAUSE' lines or plain text. Returns [(text, pause)]."""
    segments = []
    for raw in script.splitlines():
        line = raw.strip()
        if not line or line.startswith("#"):
            continue
        if "|||" in line:
            text, _, pause_s = line.rpartition("|||")
            try:
                pause = max(0.0, min(float(pause_s.strip()), 3.0))
            except ValueError:
                pause = 0.4
            text = text.strip()
        else:
            text, pause = line, 0.4
        if text:
            segments.append((text, pause))
    return segments


def is_dialogue(text: str) -> bool:
    return bool(SPEAKER_RE.match(text.strip()))


def group_chunks(segments):
    """Group segments into API-sized chunks, breaking at strong pauses.

    Returns [(chunk_text, trailing_pause, kind)] where kind is "narration" or
    "dialogue". Dialogue lines ("SPEAKER: text") are never mixed with narration
    in the same chunk, so they can be rendered with distinct voices.
    """
    chunks, buf, buf_len, buf_kind = [], [], 0, None

    def flush(pause):
        if buf:
            joined = "\n".join(buf) if buf_kind == "dialogue" else " ".join(buf)
            chunks.append((joined, pause, buf_kind))

    for i, (text, pause) in enumerate(segments):
        kind = "dialogue" if is_dialogue(text) else "narration"
        if buf and kind != buf_kind:
            flush(0.35)  # small beat when switching between narration and dialogue
            buf, buf_len = [], 0
        buf.append(text)
        buf_kind = kind
        buf_len += len(text)
        if pause >= 0.6 or buf_len >= MAX_CHUNK_CHARS or i == len(segments) - 1:
            flush(pause)
            buf, buf_len = [], 0
    return chunks


# -------------------------------------------------------------------- tts ----

def speech_config(text: str, voice: str, kind: str, voice_a: str, voice_b: str):
    """Single-voice config for narration; two-voice config for dialogue chunks."""
    if kind != "dialogue":
        return {"voiceConfig": {"prebuiltVoiceConfig": {"voiceName": voice}}}, None

    speakers = []
    for line in text.splitlines():
        m = SPEAKER_RE.match(line.strip())
        if m and m.group(1) not in speakers:
            speakers.append(m.group(1))
    if not speakers:  # defensive: fall back to narration
        return {"voiceConfig": {"prebuiltVoiceConfig": {"voiceName": voice}}}, None

    # Gemini supports up to 2 distinct speakers; extras alternate between them.
    voices = [voice_a, voice_b]
    configs = [
        {
            "speaker": name,
            "voiceConfig": {"prebuiltVoiceConfig": {"voiceName": voices[i % 2]}},
        }
        for i, name in enumerate(speakers[:2])
    ]
    return {"multiSpeakerVoiceConfig": {"speakerVoiceConfigs": configs}}, speakers[:2]


def tts_chunk(text: str, voice: str, style: str, kind: str = "narration",
              voice_a: str = "", voice_b: str = "") -> bytes:
    """One Gemini TTS call -> raw PCM bytes. Retries transient failures."""
    cfg, speakers = speech_config(
        text, voice, kind, voice_a or DIALOGUE_VOICE_A, voice_b or DIALOGUE_VOICE_B
    )
    if speakers:
        lead = (
            f"TTS the following short trading-desk exchange between "
            f"{' and '.join(speakers)}. Deliver it naturally, as real market chatter: "
            f"brisk, clipped, matter-of-fact, no theatrical acting.\n\n"
        )
        prompt = lead + text
    else:
        prompt = f"{style.strip()}\n\n{text}" if style.strip() else text
    body = {
        "contents": [{"parts": [{"text": prompt}]}],
        "generationConfig": {
            "responseModalities": ["AUDIO"],
            "speechConfig": cfg,
        },
    }
    url = GEMINI_URL.format(model=MODEL, key=GEMINI_API_KEY)
    last_error = None
    for attempt in range(7):
        try:
            resp = httpx.post(url, json=body, timeout=120)
            if resp.status_code == 429:
                time.sleep(15 * (attempt + 1))
                last_error = "429 rate limited"
                continue
            if resp.status_code >= 400:
                # Surface the API's own message — essential for diagnosis.
                last_error = f"HTTP {resp.status_code}: {resp.text[:700]}"
                # The preview TTS model returns sporadic 400 INVALID_ARGUMENT on
                # perfectly valid input, so retry those too. Only auth errors are
                # genuinely terminal.
                if resp.status_code in (401, 403):
                    break
                time.sleep(4 * (attempt + 1))
                continue
            data = resp.json()
            cands = data.get("candidates") or []
            if not cands:
                # 200 with no candidates: the preview model does this
                # intermittently. Retryable.
                last_error = f"empty response (no candidates): {str(data)[:300]}"
                time.sleep(6 * (attempt + 1))
                continue
            part = cands[0]["content"]["parts"][0]
            return base64.b64decode(part["inlineData"]["data"])
        except Exception as exc:  # noqa: BLE001
            last_error = f"{type(exc).__name__}: {str(exc)[:500]}"
            time.sleep(6 * (attempt + 1))

    if speakers:
        # A dialogue must never cost us the episode: fall back to the narrator
        # reading the exchange, speaker labels stripped.
        plain = "\n".join(
            (m.group(2) if (m := SPEAKER_RE.match(l.strip())) else l)
            for l in text.splitlines()
        )
        return tts_chunk(plain, voice, style, "narration", voice_a, voice_b)

    raise RuntimeError(f"Gemini TTS chunk failed after 7 attempts — {last_error}")


def encode_mp3(pcm: bytes, bitrate: str = "96k") -> bytes:
    """Raw s16le PCM -> mp3 via ffmpeg."""
    with tempfile.NamedTemporaryFile(suffix=".mp3", delete=False) as out:
        out_path = out.name
    try:
        subprocess.run(
            [
                "ffmpeg", "-y", "-loglevel", "error",
                "-f", "s16le", "-ar", str(SAMPLE_RATE), "-ac", "1", "-i", "pipe:0",
                "-b:a", bitrate, out_path,
            ],
            input=pcm,
            check=True,
        )
        return open(out_path, "rb").read()
    finally:
        os.unlink(out_path)


def upload_public(data: bytes, object_name: str, content_type: str = "audio/mpeg",
                  cache_control: str = "public, max-age=3600") -> str:
    client = storage.Client()
    blob = client.bucket(GCS_BUCKET).blob(object_name)
    blob.cache_control = cache_control
    blob.upload_from_string(data, content_type=content_type)
    return f"https://storage.googleapis.com/{GCS_BUCKET}/{object_name}"


def read_meta(show: str, episode: str):
    """Read the sidecar status json for an episode, or None."""
    client = storage.Client()
    blob = client.bucket(GCS_BUCKET).blob(f"{show}/{episode}.json")
    if not blob.exists():
        return None
    return json.loads(blob.download_as_bytes())


def write_meta(show: str, episode: str, meta: dict):
    upload_public(json.dumps(meta).encode(), f"{show}/{episode}.json",
                  content_type="application/json")


# ------------------------------------------------------------------- tool ----

def _generate_impl(script, voice, style, show, episode, voice_a="", voice_b="") -> str:
    if not GEMINI_API_KEY or not GCS_BUCKET:
        return json.dumps({"error": "server misconfigured: missing GEMINI_API_KEY or GCS_BUCKET"})

    segments = parse_segments(script)
    if not segments:
        return json.dumps({"error": "empty script"})
    chunks = group_chunks(segments)

    safe_show = re.sub(r"[^a-z0-9_-]", "-", show.lower()) or "default"
    safe_ep = re.sub(r"[^a-z0-9_.-]", "-", episode.lower()) or "episode"
    voice = voice or DEFAULT_VOICE

    # Status sidecar: lets a client that timed out on this (long) call recover
    # the result later via check_audio_ready.
    n_dialogue = sum(1 for c in chunks if c[2] == "dialogue")
    write_meta(safe_show, safe_ep,
               {"status": "working", "chunks": len(chunks), "dialogue_chunks": n_dialogue})

    try:
        pcm = io.BytesIO()
        failed = []
        for i, (text, trailing_pause, kind) in enumerate(chunks):
            try:
                pcm.write(tts_chunk(text, voice, style, kind, voice_a, voice_b))
            except Exception as exc:  # noqa: BLE001
                # One stubborn chunk must not cost the whole episode: leave a
                # short gap, record it, and keep going.
                failed.append({"chunk": i + 1, "error": str(exc)[:300],
                               "preview": text[:80]})
                pcm.write(b"\x00" * int(BYTES_PER_SEC * 0.4))
            if i < len(chunks) - 1 and trailing_pause > 0:
                pcm.write(b"\x00" * int(BYTES_PER_SEC * trailing_pause))
        if len(failed) == len(chunks):
            raise RuntimeError(f"every chunk failed; first error: {failed[0]['error']}")

        raw = pcm.getvalue()
        duration = int(len(raw) / BYTES_PER_SEC)
        mp3 = encode_mp3(raw)
        url = upload_public(mp3, f"{safe_show}/{safe_ep}.mp3")
    except Exception as exc:  # noqa: BLE001
        write_meta(safe_show, safe_ep, {"status": "error", "error": str(exc)[:500]})
        raise

    meta = {
        "status": "done",
        "url": url,
        "duration_seconds": duration,
        "size_bytes": len(mp3),
        "chunks": len(chunks),
        "dialogue_chunks": n_dialogue,
        "failed_chunks": failed,
        "voice": voice,
        "model": MODEL,
    }
    write_meta(safe_show, safe_ep, meta)
    return json.dumps(meta)


@mcp.tool()
async def generate_podcast_audio(
    script: str,
    voice: str = "",
    style: str = "Narrate calmly and clearly, like a professional podcast host. Mark a brief pause between ideas.",
    show: str = "default",
    episode: str = "episode",
    dialogue_voice_a: str = "",
    dialogue_voice_b: str = "",
) -> str:
    """Generate podcast audio from a script with Gemini TTS and host it publicly.

    Supports two-voice dialogue: any script line written as "SPEAKER: text"
    (uppercase label, e.g. "BUYER: November Santos, what have you got?") is
    rendered as a two-person exchange with distinct voices, separate from the
    narrator. Consecutive dialogue lines are kept together; up to two distinct
    speakers per exchange.

    Args:
        script: The spoken script. Either plain text, or one segment per line in
            the format "TEXT ||| PAUSE" where PAUSE (seconds) is inserted as real
            silence after the segment (pauses >= 0.6 s become chunk boundaries).
            Dialogue lines use "SPEAKER: text ||| PAUSE".
        voice: Narrator voice — Gemini prebuilt name (Achird, Charon, Kore,
            Puck, Fenrir...). Empty = server default.
        style: Natural-language delivery instruction for narration.
        show: Show slug — becomes the storage folder (one per podcast).
        episode: Episode slug — becomes the file name, e.g. "ep02".
        dialogue_voice_a: Voice for the first speaker in an exchange (default Puck).
        dialogue_voice_b: Voice for the second speaker (default Kore).

    Returns JSON: {"url", "duration_seconds", "size_bytes", "chunks", "model"}.
    The audio itself is uploaded to public storage; only this metadata returns.

    Runs the blocking work (HTTP to Gemini, ffmpeg, GCS upload) in a worker
    thread so concurrent MCP calls are served without blocking the event loop.
    """
    import anyio

    return await anyio.to_thread.run_sync(
        _generate_impl, script, voice, style, show, episode,
        dialogue_voice_a, dialogue_voice_b
    )


def _probe_impl(text: str, voice_a: str, voice_b: str) -> str:
    """Diagnostic: try a multi-speaker call and return the raw API outcome."""
    cfg, speakers = speech_config(text, "Achird", "dialogue",
                                  voice_a or DIALOGUE_VOICE_A, voice_b or DIALOGUE_VOICE_B)
    body = {
        "contents": [{"parts": [{"text": text}]}],
        "generationConfig": {"responseModalities": ["AUDIO"], "speechConfig": cfg},
    }
    url = GEMINI_URL.format(model=MODEL, key=GEMINI_API_KEY)
    try:
        resp = httpx.post(url, json=body, timeout=120)
        ok = resp.status_code == 200
        return json.dumps({
            "speakers": speakers,
            "status_code": resp.status_code,
            "ok": ok,
            "body": "" if ok else resp.text[:1500],
            "config_sent": cfg,
        })
    except Exception as exc:  # noqa: BLE001
        return json.dumps({"error": f"{type(exc).__name__}: {str(exc)[:500]}"})


@mcp.tool()
async def probe_multispeaker(
    text: str = "BUYER: November Santos, what have you got?\nSELLER: I make you plus eighty-five.",
    voice_a: str = "",
    voice_b: str = "",
) -> str:
    """Diagnostic tool: attempt one multi-speaker TTS call and report the raw result.

    Returns the HTTP status, the exact config sent, and the API's error body if
    it failed. Use this to debug two-voice dialogue without generating audio.
    """
    import anyio

    return await anyio.to_thread.run_sync(_probe_impl, text, voice_a, voice_b)


def _check_impl(show: str, episode: str) -> str:
    safe_show = re.sub(r"[^a-z0-9_-]", "-", show.lower()) or "default"
    safe_ep = re.sub(r"[^a-z0-9_.-]", "-", episode.lower()) or "episode"
    meta = read_meta(safe_show, safe_ep)
    if meta:
        return json.dumps(meta)
    return json.dumps({"status": "unknown"})


def _publish_feed_impl(show: str, feed_xml: str) -> str:
    safe_show = re.sub(r"[^a-z0-9_-]", "-", show.lower()) or "default"
    url = upload_public(
        feed_xml.encode("utf-8"),
        f"{safe_show}/feed.xml",
        content_type="application/rss+xml; charset=utf-8",
        cache_control="public, max-age=300",  # 5 min: podcast apps see updates fast
    )
    return json.dumps({"feed_url": url})


@mcp.tool()
async def publish_feed(show: str, feed_xml: str) -> str:
    """Host the podcast RSS feed on public storage with a 5-minute cache.

    Pass the full feed.xml content. Returns {"feed_url": ...} — a stable URL
    podcast apps can subscribe to, refreshed within minutes of each publish
    (unlike CDN-cached copies which can lag by hours).
    """
    import anyio

    return await anyio.to_thread.run_sync(_publish_feed_impl, show, feed_xml)


def _publish_page_impl(show: str, episode: str, page_html: str) -> str:
    safe_show = re.sub(r"[^a-z0-9_-]", "-", show.lower()) or "default"
    safe_ep = re.sub(r"[^a-z0-9_.-]", "-", episode.lower()) or "episode"
    url = upload_public(
        page_html.encode("utf-8"),
        f"{safe_show}/{safe_ep}.html",
        content_type="text/html; charset=utf-8",
        cache_control="public, max-age=300",
    )
    return json.dumps({"page_url": url})


@mcp.tool()
async def publish_page(show: str, episode: str, page_html: str) -> str:
    """Host an episode's HTML page on public storage and return its URL.

    Pass the full HTML document (as produced by the toolkit's build_page.py).
    Returns {"page_url": ...} — a stable link to put in the episode email.
    Re-publishing the same episode replaces the page (5-minute cache).
    """
    import anyio

    return await anyio.to_thread.run_sync(_publish_page_impl, show, episode, page_html)


def _purge_impl(paths) -> str:
    results = {}
    for path in paths:
        try:
            resp = httpx.get(f"https://purge.jsdelivr.net/{path.lstrip('/')}", timeout=30)
            results[path] = resp.status_code
        except Exception as exc:  # noqa: BLE001
            results[path] = str(exc)[:200]
    return json.dumps(results)


@mcp.tool()
async def purge_feed_cache(
    paths: list[str] = ["npm/@sdelsad/commodity-desk-daily@latest/feed.xml"],
) -> str:
    """Purge jsDelivr's CDN cache for the given paths (default: the podcast RSS feed).

    Call this right after publishing a new episode so podcast apps see the
    updated feed immediately instead of waiting out the CDN cache (up to ~12 h).
    Returns per-path HTTP status codes from the purge API.
    """
    import anyio

    return await anyio.to_thread.run_sync(_purge_impl, paths)


@mcp.tool()
async def check_audio_ready(show: str, episode: str) -> str:
    """Check whether a previously requested episode audio is ready.

    Use this when a generate_podcast_audio call timed out client-side: the
    server keeps working and uploads the result anyway. Poll this (it returns
    instantly) until status is "done", then use the returned url,
    duration_seconds and size_bytes exactly as if the original call had
    returned them. Statuses: "working", "done", "error" (with message),
    "unknown" (no such job).
    """
    import anyio

    return await anyio.to_thread.run_sync(_check_impl, show, episode)


# ------------------------------------------------------------------- asgi ----

_inner = mcp.streamable_http_app()


async def app(scope, receive, send):
    """Health endpoint + optional secret path gate around the MCP app."""
    if scope["type"] == "http":
        path = scope.get("path", "")
        if path in ("/", "/healthz"):
            await send({"type": "http.response.start", "status": 200,
                        "headers": [(b"content-type", b"text/plain")]})
            await send({"type": "http.response.body", "body": b"ok"})
            return
        if SECRET_PATH:
            prefix = f"/{SECRET_PATH}"
            if not path.startswith(prefix + "/") and path != prefix:
                await send({"type": "http.response.start", "status": 404,
                            "headers": [(b"content-type", b"text/plain")]})
                await send({"type": "http.response.body", "body": b"not found"})
                return
            scope = dict(scope)
            scope["path"] = path[len(prefix):] or "/"
            raw = scope.get("raw_path")
            if raw:
                scope["raw_path"] = raw[len(prefix.encode()):] or b"/"
    await _inner(scope, receive, send)
