#!/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")

SAMPLE_RATE = 24000  # Gemini TTS returns 16-bit mono PCM at 24 kHz
BYTES_PER_SEC = SAMPLE_RATE * 2
MAX_CHUNK_CHARS = 4000  # fewer, larger chunks = fewer requests = less rate limiting
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 group_chunks(segments):
    """Group segments into API-sized chunks, breaking at strong pauses.

    Returns [(chunk_text, trailing_pause)]. Intra-chunk pauses are left to the
    model's own prosody; boundary pauses become real inserted silence.
    """
    chunks, buf, buf_len = [], [], 0
    for i, (text, pause) in enumerate(segments):
        buf.append(text)
        buf_len += len(text)
        boundary = pause >= 0.6 or buf_len >= MAX_CHUNK_CHARS
        if boundary or i == len(segments) - 1:
            chunks.append((" ".join(buf), pause))
            buf, buf_len = [], 0
    return chunks


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

def tts_chunk(text: str, voice: str, style: str) -> bytes:
    """One Gemini TTS call -> raw PCM bytes. Retries transient failures."""
    prompt = f"{style.strip()}\n\n{text}" if style.strip() else text
    body = {
        "contents": [{"parts": [{"text": prompt}]}],
        "generationConfig": {
            "responseModalities": ["AUDIO"],
            "speechConfig": {
                "voiceConfig": {"prebuiltVoiceConfig": {"voiceName": voice}}
            },
        },
    }
    url = GEMINI_URL.format(model=MODEL, key=GEMINI_API_KEY)
    last_error = None
    for attempt in range(3):
        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
            resp.raise_for_status()
            data = resp.json()
            part = data["candidates"][0]["content"]["parts"][0]
            return base64.b64decode(part["inlineData"]["data"])
        except Exception as exc:  # noqa: BLE001 — surface the last error cleanly
            last_error = str(exc)[:500]
            time.sleep(5 * (attempt + 1))
    raise RuntimeError(f"Gemini TTS failed after 3 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") -> str:
    client = storage.Client()
    blob = client.bucket(GCS_BUCKET).blob(object_name)
    # Short cache so a re-published episode (same name) propagates quickly.
    blob.cache_control = "public, max-age=3600"
    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) -> 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.
    write_meta(safe_show, safe_ep, {"status": "working", "chunks": len(chunks)})

    try:
        pcm = io.BytesIO()
        for i, (text, trailing_pause) in enumerate(chunks):
            pcm.write(tts_chunk(text, voice, style))
            if i < len(chunks) - 1 and trailing_pause > 0:
                pcm.write(b"\x00" * int(BYTES_PER_SEC * trailing_pause))

        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),
        "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",
) -> str:
    """Generate podcast audio from a script with Gemini TTS and host it publicly.

    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).
        voice: Gemini prebuilt voice name (e.g. Charon, Kore, Puck, Fenrir).
            Empty = server default.
        style: Natural-language delivery instruction prepended to each chunk.
        show: Show slug — becomes the storage folder (one per podcast).
        episode: Episode slug — becomes the file name, e.g. "ep02".

    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
    )


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"})


@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)
