#!/usr/bin/env python3
"""
FastRTC 风格接口 - 简化版 ReplyOnPause
借鉴 FastRTC 的易用性，保持我们的性能优势
"""
import asyncio
import logging
from typing import AsyncGenerator, Optional
import numpy as np

logger = logging.getLogger(__name__)


class ReplyOnPause:
    """
    借鉴 FastRTC 的 ReplyOnPause
    
    特点:
    1. 自动 VAD 检测
    2. 自动话轮切换
    3. 支持打断
    4. 底层使用 TEN VAD (更快)
    """
    
    def __init__(self, handler, can_interrupt: bool = True):
        """
        Args:
            handler: 异步生成器函数，接收音频数据，yield 音频块
            can_interrupt: 是否允许打断
        """
        self.handler = handler
        self.can_interrupt = can_interrupt
        self._current_task = None
        self._interrupted = False
    
    async def __call__(self, audio_chunk: bytes) -> Optional[bytes]:
        """
        处理音频帧
        
        Returns:
            输出的音频块，或 None
        """
        # 这里简化实现，实际应集成到主 Agent
        pass
    
    async def start_streaming(self, audio_source):
        """
        开始流式处理
        
        Args:
            audio_source: 音频源，提供 get_audio() 方法
        """
        logger.info("🎙️ ReplyOnPause 启动")
        
        buffer = bytearray()
        vad = TenVADWrapper()
        await vad.initialize()
        
        while True:
            chunk = await audio_source.get_audio()
            if not chunk:
                continue
            
            buffer.extend(chunk)
            
            # VAD 检测
            vad_result = vad.process_frame(chunk)
            
            if vad_result is True:
                # 检测到语音开始
                logger.info("🎤 检测到语音")
                self._interrupted = False
            elif vad_result is False:
                # 检测到语音结束
                if len(buffer) > 3200:
                    logger.info("📝 处理语音...")
                    await self._process_speech(bytes(buffer))
                    buffer.clear()
            
            await asyncio.sleep(0.01)
    
    async def _process_speech(self, audio_data: bytes):
        """处理完整语音"""
        try:
            # 调用用户处理器
            async for output_chunk in self.handler(audio_data):
                if self._interrupted:
                    break
                yield output_chunk
        except Exception as e:
            logger.error(f"处理失败: {e}")
    
    def interrupt(self):
        """打断当前处理"""
        self._interrupted = True
        logger.info("⏹️ 打断请求")


class TenVADWrapper:
    """TEN VAD 包装"""
    
    def __init__(self, threshold: float = 0.3):
        self.threshold = threshold
        self._vad = None
        self._is_speaking = False
        self._speech_start = 0.0
        self._silence_start = 0.0
        
    async def initialize(self):
        try:
            import ten_vad
            self._vad = ten_vad.TenVad(hop_size=256, threshold=self.threshold)
            logger.info(f"✅ TEN VAD 就绪 (threshold={self.threshold})")
        except ImportError:
            logger.warning("⚠️ TEN VAD 未安装")
            self._vad = None
    
    def process_frame(self, audio_data: bytes) -> Optional[bool]:
        if self._vad is None:
            return self._energy_detect(audio_data)
        
        try:
            import numpy as np
            audio_np = np.frombuffer(audio_data, dtype=np.int16)[:256]
            score, is_speech = self._vad.process(audio_np)
            
            now = asyncio.get_event_loop().time() * 1000
            
            if is_speech:
                if not self._is_speaking:
                    self._is_speaking = True
                    self._speech_start = now
                    return True
            else:
                if self._is_speaking:
                    if now - self._speech_start < 300:
                        self._is_speaking = False
                        return False
                    if self._silence_start == 0:
                        self._silence_start = now
                    if now - self._silence_start >= 500:
                        self._is_speaking = False
                        return False
            return None
        except:
            return None
    
    def _energy_detect(self, audio_data: bytes) -> Optional[bool]:
        try:
            arr = np.frombuffer(audio_data, dtype=np.int16)
            return float(np.mean(arr.astype(float) ** 2)) > 100
        except:
            return None


# ========== 使用示例 ==========
async def example_handler(audio: bytes) -> AsyncGenerator[bytes, None]:
    """示例处理器"""
    # 1. STT
    text = await stt.recognize(audio)
    logger.info(f"🎤 识别: {text}")
    
    # 2. LLM
    response_text = await llm.chat(text)
    
    # 3. TTS (流式)
    async for chunk in tts.synthesize_stream(response_text):
        yield chunk


# 使用
# reply_on_pause = ReplyOnPause(example_handler)
# await reply_on_pause.start_streaming(audio_source)
