#!/usr/bin/env python3
"""
远程渲染 HTTP 客户端 — 对接独立部署的 remotion-renderer 服务。

接口（位于 /remotion-renderer/server/index.ts）：
  - POST /render        → { code, data: { taskId } }
        body 可选 coverCompositionId（如 "SpotlightCardCover"）；传了的话视频
        渲染完成后服务端会再串行渲一帧静态封面、一并上传。
  - POST /renderStatus  → { code, data: { status, progress, fileUrl?,
                                          coverUrl?, coverError?, error? } }
        coverUrl  视频成功且封面成功时存在
        coverError 视频成功但封面失败时存在（视频本身仍然 succeeded）

认证头（复用 X-Priv-Token 规范，服务端通过后端 API 校验）：
  X-Priv-Token: <PRIV_TOKEN>
  x-invoke-skill:         render-video
  x-invoke-agent:         <AGENT_NAME>  # optional

Base URL 优先级：
  1. 调用方显式传入 base_url
  2. 环境变量 REMOTION_RENDER_API_URL（指向本地/staging 时在此覆盖）
  3. 默认 https://api-render.remixmate.ai（生产，零配置可用）

本模块只负责 HTTP 层；调用方（render_video.py）负责把结果写回 render_plan。
"""

from __future__ import annotations

import json
import os
import sys
import time
import urllib.error
import urllib.parse
import urllib.request
from typing import Callable, Optional
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "template-registry" / "scripts"))
from log_diagnostics import trace_headers

DEFAULT_API_BASE_URL = "https://api-render.remixmate.ai"
API_BASE_URL = (
    os.environ.get("REMOTION_RENDER_API_URL")
    or os.environ.get("RENDER_API_URL")
    or DEFAULT_API_BASE_URL
).strip()
SKILL_NAME = "render-video"
AGENT_NAME = os.environ.get("AGENT_NAME", "")

STATUS_PENDING = "pending"
STATUS_PROCESSING = "processing"
STATUS_RUNNING = "running"
STATUS_SUCCEEDED = "succeeded"
STATUS_SUCCEED = "succeed"
STATUS_COMPLETED = "completed"
STATUS_FAILED = "failed"
STATUS_CANCELLED = "cancelled"

TERMINAL_SUCCESS = {STATUS_SUCCEEDED, STATUS_SUCCEED, STATUS_COMPLETED}
TERMINAL_FAILURE = {STATUS_FAILED, STATUS_CANCELLED}


class RemoteRenderError(RuntimeError):
    """抛出至调用方，由上层写入 render_plan.errors。"""


class RemoteRenderTimeout(RemoteRenderError):
    """轮询超时 —— 任务**没有失败**，只是我们不等了。

    与普通 RemoteRenderError 分开，是因为上层的处置完全相反：
    真失败 → 可以重新提交一个任务；超时 → 后端还在渲，重新提交等于白烧一遍，
    必须保留 taskId 让下一次调用续上。
    """

    def __init__(self, message: str, task_id: str):
        super().__init__(message)
        self.task_id = task_id


def _build_headers(private_token: str, content_type: str = "application/json", conversation_id: Optional[str] = None) -> dict:
    if not private_token:
        raise RemoteRenderError("PrivToken is not set (PRIV_TOKEN or --priv-token)")
    headers = {
        "X-Priv-Token": private_token,
        "x-invoke-skill": SKILL_NAME,
    }
    if content_type:
        headers["Content-Type"] = content_type
    if AGENT_NAME:
        headers["x-invoke-agent"] = AGENT_NAME
    if conversation_id:
        headers["x-conversation-id"] = conversation_id
    headers.update(trace_headers())
    return headers


def _request(
    method: str,
    path: str,
    *,
    private_token: str,
    payload: Optional[dict] = None,
    query: Optional[dict] = None,
    base_url: Optional[str] = None,
    timeout: float = 120.0,
    conversation_id: Optional[str] = None,
) -> dict:
    base = (base_url or API_BASE_URL).rstrip("/")
    if not base:
        raise RemoteRenderError("renderer service base URL is not set")
    url = f"{base}{path}"
    if query:
        url = f"{url}?{urllib.parse.urlencode(query)}"

    data = None
    if payload is not None:
        data = json.dumps(payload, ensure_ascii=False).encode("utf-8")

    headers = _build_headers(private_token, conversation_id=conversation_id)
    req = urllib.request.Request(url, data=data, headers=headers, method=method)

    with urllib.request.urlopen(req, timeout=timeout) as resp:
        body = resp.read().decode("utf-8")
    try:
        return json.loads(body)
    except json.JSONDecodeError as exc:
        raise RemoteRenderError(f"remote response is not JSON: {body[:200]}") from exc


def _check_code(result: dict) -> dict:
    code = result.get("code")
    if code != 0:
        msg = result.get("msg") or result.get("message") or "unknown error"
        raise RemoteRenderError(f"API code={code}: {msg}")
    return result.get("data") or {}


def start_render(
    payload: dict,
    *,
    private_token: str,
    base_url: Optional[str] = None,
    timeout: float = 120.0,
    conversation_id: Optional[str] = None,
    path: str = "/render",
) -> str:
    """POST <path>（默认 /render，私有模板动态渲染传 /renderDraft）；返回 taskId。"""
    result = _request(
        "POST",
        path,
        private_token=private_token,
        payload=payload,
        base_url=base_url,
        timeout=timeout,
        conversation_id=conversation_id,
    )
    data = _check_code(result)
    task_id = data.get("taskId")
    if not task_id:
        raise RemoteRenderError(f"{path} did not return a taskId: {result}")
    return str(task_id)


def poll_render(
    task_id: str,
    *,
    private_token: str,
    base_url: Optional[str] = None,
    timeout: float = 1800.0,
    interval: float = 2.0,
    request_timeout: float = 60.0,
    max_consecutive_errors: int = 5,
    on_progress: Optional[Callable[[dict], None]] = None,
    adaptive_interval: bool = False,
    status_path: str = "/renderStatus",
) -> dict:
    """POST <status_path>（默认 /renderStatus，私有模板传 /renderDraftStatus）轮询直到完成/失败/超时。

    成功时返回 data dict（至少包含 fileUrl），否则抛 RemoteRenderError。
    on_progress 回调在每次成功请求后触发，参数为完整 data dict。

    adaptive_interval=True 时根据渲染进度自动调整轮询间隔：
      0-20%: max(interval * 0.5, 1.0)（起步阶段采样加密，短渲染也能看到进度）
      20-80%: interval（正常渲染阶段）
      80-100%: max(interval * 0.6, 1.0)（编码收尾阶段，更频繁检查）
    """
    start = time.monotonic()
    consecutive_errors = 0
    current_progress = 0.0

    while True:
        elapsed = time.monotonic() - start
        if elapsed > timeout:
            raise RemoteRenderTimeout(
                f"remote render polling timed out (waited {elapsed:.0f}s, taskId={task_id}); "
                f"the job may still be running on the backend — query {status_path} later",
                task_id,
            )

        try:
            result = _request(
                "POST",
                status_path,
                private_token=private_token,
                payload={"taskId": task_id},
                base_url=base_url,
                timeout=request_timeout,
            )
        except (urllib.error.URLError, OSError, ConnectionError) as exc:
            consecutive_errors += 1
            print(
                f"   ⚠️  polling network error ({consecutive_errors}/{max_consecutive_errors}): {exc}",
                file=sys.stderr,
                flush=True,
            )
            if consecutive_errors >= max_consecutive_errors:
                raise RemoteRenderError(
                    f"{max_consecutive_errors} consecutive network errors; giving up on polling (taskId={task_id})"
                ) from exc
            time.sleep(interval)
            continue
        except urllib.error.HTTPError as exc:
            body = ""
            try:
                body = exc.read().decode("utf-8")
            except Exception:
                pass
            raise RemoteRenderError(f"HTTP {exc.code} {status_path}: {body[:200]}") from exc

        consecutive_errors = 0

        data = _check_code(result)
        if on_progress is not None:
            try:
                on_progress(data)
            except Exception:  # noqa: BLE001 — 进度回调异常不应中断轮询
                pass

        status = str(data.get("status") or "").lower()
        if status in TERMINAL_SUCCESS:
            return data
        if status in TERMINAL_FAILURE:
            err = data.get("error") or data.get("errorMsg") or data.get("message") or "unknown"
            raise RemoteRenderError(f"remote render failed ({status}): {err}")

        # Adaptive interval: adjust sleep based on current progress
        current_progress = float(data.get("progress") or 0.0)
        if adaptive_interval:
            if current_progress < 0.2:
                # 起步阶段(含服务端很快就渲完的短片)采样加密，避免 0%→完成的跳变
                sleep_time = max(interval * 0.5, 1.0)
            elif current_progress >= 0.8:
                sleep_time = max(interval * 0.6, 1.0)
            else:
                sleep_time = interval
        else:
            sleep_time = interval
        time.sleep(sleep_time)


__all__ = ["start_render", "poll_render", "RemoteRenderError"]
