"""火山引擎(豆包)大模型语音合成 —— V3 WebSocket 双向流式二进制协议。

文档:wss://openspeech.bytedance.com/api/v3/tts/bidirection(用户提供)。
对外:`synthesize(text, voice) -> mp3 bytes`(同步,内部跑 asyncio)。引擎线程可直接调。
鉴权(旧版控制台):X-Api-App-Id + X-Api-Access-Key + X-Api-Resource-Id(seed-tts-1.0/2.0)。
帧:4 字节 header + [event int32] + [id_len+id] + payload_len + payload;整数大端。
"""

from __future__ import annotations

import asyncio
import gzip
import json
import os
import struct
import uuid
from pathlib import Path
from typing import Optional

from genesis.obs.logging_setup import get_logger

logger = get_logger("voice.tts")

WS_URL = "wss://openspeech.bytedance.com/api/v3/tts/bidirection"

# 9 个座位的声线(实测 seed-tts-1.0 可用,4 男 4 女交替;座位>8 循环复用)
VOICES = [
    "zh_male_M392_conversation_wvae_bigtts", "zh_female_cancan_mars_bigtts",
    "zh_male_ahu_conversation_wvae_bigtts", "zh_female_shuangkuaisisi_moon_bigtts",
    "zh_male_beijingxiaoye_moon_bigtts", "zh_female_wanwanxiaohe_moon_bigtts",
    "zh_male_jingqiangkanye_moon_bigtts", "zh_female_tianmeixiaoyuan_moon_bigtts",
]


# 法官/主持声线:沉稳叙事腔,刻意选在 8 个座位轮换之外,一耳朵能和玩家区分开
JUDGE_VOICE = "zh_male_changtianyi_mars_bigtts"


def voice_for_seat(seat: int) -> str:
    if int(seat) == 0:                       # 0 号 = 法官
        return JUDGE_VOICE
    return VOICES[(int(seat) - 1) % len(VOICES)]

# Message types(byte1 高 4 位)
MT_FULL_CLIENT = 0b0001
MT_FULL_SERVER = 0b1001
MT_AUDIO_ONLY = 0b1011
MT_ERROR = 0b1111
FLAG_EVENT = 0b0100          # 带 event number
SER_JSON, SER_RAW = 0b0001, 0b0000

# Events
EV_START_CONN, EV_FINISH_CONN = 1, 2
EV_CONN_STARTED, EV_CONN_FAILED = 50, 51
EV_START_SESSION, EV_FINISH_SESSION = 100, 102
EV_SESSION_STARTED, EV_SESSION_FINISHED, EV_SESSION_FAILED = 150, 152, 153
EV_TASK_REQUEST = 200
EV_TTS_SENTENCE_START, EV_TTS_SENTENCE_END, EV_TTS_RESPONSE = 350, 351, 352


def _creds() -> tuple[str, str]:
    """从环境或 .env 读 VOLC_APP_ID / VOLC_ACCESS_TOKEN(不回显)。"""
    appid, token = os.getenv("VOLC_APP_ID"), os.getenv("VOLC_ACCESS_TOKEN")
    if not (appid and token):
        envf = Path(__file__).resolve().parents[3] / ".env"
        if envf.exists():
            for ln in envf.read_text(encoding="utf-8").splitlines():
                ln = ln.strip()
                if ln.startswith("VOLC_APP_ID="):
                    appid = appid or ln.split("=", 1)[1]
                elif ln.startswith("VOLC_ACCESS_TOKEN="):
                    token = token or ln.split("=", 1)[1]
    if not (appid and token):
        raise RuntimeError("缺少 VOLC_APP_ID / VOLC_ACCESS_TOKEN(请配 .env)")
    return appid, token


def _frame(event: int, session_id: Optional[str], payload, *, msg_type=MT_FULL_CLIENT, serial=SER_JSON) -> bytes:
    out = bytes([(0b0001 << 4) | 0b0001, (msg_type << 4) | FLAG_EVENT, (serial << 4) | 0, 0])
    out += struct.pack(">i", event)
    if session_id is not None:
        sid = session_id.encode("utf-8")
        out += struct.pack(">I", len(sid)) + sid
    body = json.dumps(payload, ensure_ascii=False).encode("utf-8") if serial == SER_JSON else payload
    out += struct.pack(">I", len(body)) + body
    return out


def _parse(data: bytes) -> dict:
    msg_type = (data[1] >> 4) & 0x0F
    flags = data[1] & 0x0F
    compress = data[2] & 0x0F
    off = 4
    event = None
    if flags & FLAG_EVENT:
        event = struct.unpack(">i", data[off:off + 4])[0]; off += 4
    if msg_type == MT_ERROR:
        code = struct.unpack(">I", data[off:off + 4])[0]; off += 4
        plen = struct.unpack(">I", data[off:off + 4])[0]; off += 4
        return {"kind": "error", "event": event, "code": code, "payload": data[off:off + plen]}
    idlen = struct.unpack(">I", data[off:off + 4])[0]; off += 4          # connection_id / session_id
    off += idlen
    plen = struct.unpack(">I", data[off:off + 4])[0]; off += 4
    payload = data[off:off + plen]
    if compress == 0b0001 and payload:
        payload = gzip.decompress(payload)
    return {"kind": "audio" if msg_type == MT_AUDIO_ONLY else "server", "event": event, "payload": payload}


async def _connect(headers):
    """websockets 版本兼容:新版用 additional_headers,旧版用 extra_headers。"""
    import websockets
    try:
        return await websockets.connect(WS_URL, additional_headers=headers, max_size=None)
    except TypeError:
        return await websockets.connect(WS_URL, extra_headers=headers, max_size=None)


async def _synthesize(text: str, voice: str, appid: str, token: str, resource_id: str,
                      fmt: str = "mp3", sample_rate: int = 24000) -> bytes:
    headers = {"X-Api-App-Id": appid, "X-Api-Access-Key": token,
               "X-Api-Resource-Id": resource_id, "X-Api-Connect-Id": str(uuid.uuid4())}
    sid = str(uuid.uuid4())
    ap = {"format": fmt, "sample_rate": sample_rate}
    audio = bytearray()
    ws = await _connect(headers)
    try:
        await ws.send(_frame(EV_START_CONN, None, {}))
        r = _parse(await ws.recv())
        if r["event"] != EV_CONN_STARTED:
            raise RuntimeError(f"建连失败:{r}")
        await ws.send(_frame(EV_START_SESSION, sid, {
            "user": {"uid": "genesis"}, "event": EV_START_SESSION, "namespace": "BidirectionalTTS",
            "req_params": {"speaker": voice, "audio_params": ap}}))
        r = _parse(await ws.recv())
        if r["event"] != EV_SESSION_STARTED:
            raise RuntimeError(f"会话失败:{r.get('payload', b'')[:200]}")
        await ws.send(_frame(EV_TASK_REQUEST, sid, {
            "user": {"uid": "genesis"}, "event": EV_TASK_REQUEST, "namespace": "BidirectionalTTS",
            "req_params": {"text": text, "speaker": voice, "audio_params": ap}}))
        await ws.send(_frame(EV_FINISH_SESSION, sid, {}))
        while True:
            r = _parse(await ws.recv())
            if r["kind"] == "error":
                raise RuntimeError(f"TTS 错误 code={r['code']} {r['payload'][:200]}")
            if r["kind"] == "audio" and r["event"] == EV_TTS_RESPONSE:
                audio += r["payload"]
            if r["event"] == EV_SESSION_FAILED:
                raise RuntimeError(f"会话失败:{r['payload'][:200]}")
            if r["event"] == EV_SESSION_FINISHED:
                break
        try:
            await ws.send(_frame(EV_FINISH_CONN, None, {}))
        except Exception:
            pass
    finally:
        await ws.close()
    return bytes(audio)


def synthesize(text: str, voice: str = "zh_female_cancan_mars_bigtts",
               resource_id: Optional[str] = None, fmt: str = "mp3", sample_rate: int = 24000) -> bytes:
    """文本 → 音频字节(同步)。fmt=mp3/pcm/ogg_opus。resource_id 默认 VOLC_TTS_RESOURCE_ID 或 seed-tts-1.0。"""
    appid, token = _creds()
    rid = resource_id or os.getenv("VOLC_TTS_RESOURCE_ID") or "seed-tts-1.0"
    return asyncio.run(_synthesize(text, voice, appid, token, rid, fmt, sample_rate))
