#!/usr/bin/env python3
"""
记忆组件 - 基于 ProfileManager 的记忆管理

实现对话事件监听、事实提取、矛盾检测等功能。
"""

import asyncio
import logging
import time
from typing import Dict, Any, List
from datetime import datetime

from event_bus import EventBus, EventType, EventBusEvent
from components.base import BaseComponent, ComponentConfig, ComponentState, MemoryComponent
from profile_manager import ProfileManager

logger = logging.getLogger(__name__)


class MemoryComponentImpl(MemoryComponent):
    """记忆组件实现"""
    
    def __init__(self, config: Dict, event_bus: EventBus):
        # 提取记忆相关配置
        memory_config = config.get("memory", {})
        llm_config = config.get("llm", {})
        
        super().__init__(
            config=ComponentConfig(
                name="memory_component",
                enabled=memory_config.get("enabled", True),
                priority=5,
                dependencies=["stt", "llm"]
            ),
            event_bus=event_bus
        )
        
        self.memory_config = memory_config
        self.llm_config = llm_config
        
        # 初始化 ProfileManager
        self.pm = ProfileManager()
        self.pm.initialize()
        
        # 统计
        self._facts_extracted = 0
        self._facts_saved = 0
        self._conflicts_detected = 0
    
    async def _do_initialize(self):
        """初始化记忆组件"""
        logger.info("🧠 初始化记忆组件")
        
        # 注册事件处理器
        self.subscribe(EventType.USER_SPEECH, self._on_user_speech)
        self.subscribe(EventType.SYSTEM_ERROR, self._on_system_error)
        
        logger.info("✅ 记忆组件初始化完成")
    
    async def _run_loop(self):
        """主循环 - 定期清理过期记忆"""
        interval = self.memory_config.get("learn_interval", 300)  # 默认5分钟
        logger.info(f"🔄 记忆组件主循环启动，学习间隔: {interval}s")

        while self.state == ComponentState.RUNNING:
            try:
                # 清理过期记忆
                expired = self.pm.cleanup_expired_facts()
                if expired > 0:
                    logger.info(f"🧹 清理了 {expired} 条过期记忆")

                await asyncio.sleep(interval)

            except asyncio.CancelledError:
                break
            except Exception as e:
                logger.error(f"❌ 记忆组件主循环错误: {e}")
                await asyncio.sleep(60)

    def is_ready(self) -> bool:
        """检查是否就绪"""
        return self.pm._conn is not None

    def health_check(self) -> Dict[str, Any]:
        """健康检查"""
        return {
            "status": "healthy" if self.pm._conn is not None else "unhealthy",
            "facts_count": self._facts_saved,
            "db_path": self.pm.db_path
        }
    async def _on_user_speech(self, event: EventBusEvent):
        """处理用户语音事件"""
        text = event.data.get("text", "")
        speaker = event.data.get("speaker", "unknown")
        
        if not text:
            return
        
        logger.info(f"🧠 处理用户语音: {text[:50]}...")
        
        try:
            # 1. 提取事实
            facts = self.extract_facts(text, speaker)
            self._facts_extracted += len(facts)
            
            if facts:
                # 2. 保存事实
                saved = await self.save_facts(speaker, facts)
                self._facts_saved += saved
                
                # 3. 发布提取完成事件
                self.event_bus.publish(EventType.MEMORY_EXTRACTED, {
                    "speaker": speaker,
                    "facts_count": len(facts),
                    "saved_count": saved
                }, source="memory_component")
            
            # 4. 检测矛盾
            conflicts = self.detect_conflicts(speaker)
            if conflicts:
                self._conflicts_detected += len(conflicts)
                self.event_bus.publish(EventType.CONFLICT_DETECTED, {
                    "speaker": speaker,
                    "conflicts": conflicts
                }, source="memory_component")
                
                logger.warning(f"⚠️ 检测到 {len(conflicts)} 个矛盾")
            
        except Exception as e:
            logger.error(f"❌ 处理用户语音失败: {e}")
            self.event_bus.publish(EventType.SYSTEM_ERROR, {
                "component": "memory",
                "error": str(e)
            }, source="memory_component")
    
    async def _on_system_error(self, event: EventBusEvent):
        """处理系统错误事件"""
        error = event.data.get("error", "")
        logger.error(f"⚠️ 记忆组件收到系统错误: {error}")
    
    async def extract_facts(self, text: str, speaker: str) -> List[Dict]:
        """
        从对话文本中提取事实
        
        Args:
            text: 对话文本
            speaker: 说话人ID
        
        Returns:
            List[Dict]: 提取的事实列表
        """
        return self.pm.extract_facts_from_dialogue(text, speaker)
    
    async def save_facts(self, speaker: str, facts: List[Dict]) -> int:
        """
        保存提取的事实
        
        Args:
            speaker: 说话人ID
            facts: 事实列表
        
        Returns:
            int: 成功保存的数量
        """
        saved = 0
        for fact in facts:
            try:
                self.pm.add_fact(
                    subject_id=speaker,
                    predicate=fact.get("predicate", ""),
                    object=fact.get("object", ""),
                    topic=fact.get("topic", "other"),
                    confidence=fact.get("confidence", 0.8),
                    who=speaker,
                    what=f"对话提取: {fact.get('predicate')}",
                    time_desc=datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
                    location="语音对话",
                    reason="自动提取",
                    trigger_method="对话处理"
                )
                saved += 1
            except Exception as e:
                logger.error(f"❌ 保存事实失败: {e}")
        
        return saved
    
    async def get_context(self, identity: str, hours: int = 24) -> Dict:
        """
        获取人物上下文
        
        Args:
            identity: 人物ID
            hours: 时间范围（小时）
        
        Returns:
            Dict: 上下文信息
        """
        return self.pm.get_context(identity, hours)
    
    def detect_conflicts(self, identity: str) -> List[Dict]:
        """
        检测矛盾
        
        Args:
            identity: 人物ID
        
        Returns:
            List[Dict]: 矛盾列表
        """
        return self.pm.detect_conflicts(identity)
    
    async def get_summary(self, identity: str, hours: int = 720) -> Dict:
        """
        获取人物画像
        
        Args:
            identity: 人物ID
            hours: 时间范围
        
        Returns:
            Dict: 人物画像
        """
        return self.pm.get_summary(identity, hours)
    
    def get_stats(self) -> Dict:
        """获取统计信息"""
        return {
            "facts_extracted": self._facts_extracted,
            "facts_saved": self._facts_saved,
            "conflicts_detected": self._conflicts_detected
        }
    
    async def stop(self) -> bool:
        """停止记忆组件"""
        success = await super().stop()
        self.pm.close()
        return success
