#!/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"
  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")
DEFAULT_VOICE = os.environ.get("DEFAULT_VOICE", "Charon")

SAMPLE_RATE = 24000  # Gemini TTS returns 16-bit mono PCM at 24 kHz
BYTES_PER_SEC = SAMPLE_RATE * 2
MAX_CHUNK_CHARS = 2500
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) -> str:
    client = storage.Client()
    blob = client.bucket(GCS_BUCKET).blob(object_name)
    blob.cache_control = "public, max-age=31536000"
    blob.upload_from_string(data, content_type="audio/mpeg")
    return f"https://storage.googleapis.com/{GCS_BUCKET}/{object_name}"


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

    voice = voice or DEFAULT_VOICE
    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)

    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(mp3, f"{safe_show}/{safe_ep}.mp3")

    return json.dumps(
        {
            "url": url,
            "duration_seconds": duration,
            "size_bytes": len(mp3),
            "chunks": len(chunks),
            "voice": voice,
            "model": MODEL,
        }
    )


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


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