"""
===========================================
路由标签识别模块 - RouteLabel
===========================================
功能：使用LLM识别用户查询中涉及的节点类型和实体
     作为后续Neo4j查询的入口节点
===========================================
"""

import json
import logging
import os
import time
import asyncio
from typing import List, Dict, Any, Optional, Tuple
from pydantic import BaseModel, Field
from graph_search.init_prompts import get_extract_entities_prompt, get_rewrite_query_prompt

logger = logging.getLogger(__name__)
#logger = logging.getLogger(__name__)
logger.setLevel(logging.INFO)   # 不依赖继承链


class RouteItem(BaseModel):
    """路由项：节点类型和实体"""
    label: str = Field(..., description="节点类型，比如'Drug'、'Disease'")
    entity: str = Field(..., description="实体文本")


class RouteOutput(BaseModel):
    """路由输出：包含多个路由项"""
    outputs: List[RouteItem]


class RouteLabel:
    """
    路由标签识别类
    
    功能：使用LLM识别用户查询中涉及的节点类型和实体
    """
    
    def __init__(self, llm_api, optional_labels: list[str], schema: str = None):
        """
        初始化路由识别器
        
        Args:
            llm_api: LLM API实例（OpenaiAPI）
            optional_labels: 可选节点标签列表
            schema: Neo4j图谱schema信息（可选，用于帮助模型识别节点类型）
        """
        self.llm_api = llm_api
        self.optional_labels = optional_labels
        self.schema = schema
    
    async def _extract_entities(self, query: str) -> List[Dict[str, str]]:
        """
        任务1：抽取主语实体
        
        Args:
            query: 用户查询
            
        Returns:
            List[Dict]: 路由结果列表，格式如 [{"label": "Drug", "entity": "xxx"}]
        """
        # 使用 init_prompts 中的 prompt 生成函数
        prompt = get_extract_entities_prompt(
            query=query,
            optional_labels=self.optional_labels,
            schema=self.schema
        )
        
        try:
            llm_start = time.time()
            model = os.getenv("OPENAI_MODEL")
            response = await self.llm_api.open_api_chat_async_json(prompt, model, temperature=0)
            llm_elapsed = (time.time() - llm_start) * 1000
            
            # 解析JSON响应
            parse_start = time.time()
            content = response.strip()
            
            # 如果包含markdown代码块，提取JSON
            if "```json" in content:
                json_start = content.find("```json") + 7
                json_end = content.find("```", json_start)
                if json_end > json_start:
                    content = content[json_start:json_end].strip()
            elif "```" in content:
                json_start = content.find("```") + 3
                json_end = content.find("```", json_start)
                if json_end > json_start:
                    content = content[json_start:json_end].strip()
            
            # 尝试找到JSON数组或对象的开始和结束位置
            if not content.startswith('[') and not content.startswith('{'):
                for start_char in ['[', '{']:
                    start_idx = content.find(start_char)
                    if start_idx != -1:
                        content = content[start_idx:]
                        break
            
            # 解析JSON
            try:
                outputs_data = json.loads(content)
            except json.JSONDecodeError as e:
                logger.warning(f"任务1 JSON解析失败，尝试修复: {e}")
                if content.rfind(']') > content.rfind('}'):
                    end_idx = content.rfind(']') + 1
                    content = content[:end_idx]
                elif content.rfind('}') > -1:
                    end_idx = content.rfind('}') + 1
                    content = content[:end_idx]
                try:
                    outputs_data = json.loads(content)
                except json.JSONDecodeError:
                    logger.error(f"任务1 JSON解析最终失败，原始内容: {content[:200]}...")
                    return []
            
            # 提取entities
            if isinstance(outputs_data, dict) and "entities" in outputs_data:
                route_results = [
                    {"label": item.get("label", ""), "entity": item.get("entity", "")}
                    for item in outputs_data["entities"]
                ]
            elif isinstance(outputs_data, list):
                # 兼容旧格式：直接是列表
                route_results = [
                    {"label": item.get("label", ""), "entity": item.get("entity", "")}
                    for item in outputs_data
                ]
            elif isinstance(outputs_data, dict) and "outputs" in outputs_data:
                # 兼容旧格式：包含 outputs 字段
                route_results = [
                    {"label": item.get("label", ""), "entity": item.get("entity", "")}
                    for item in outputs_data["outputs"]
                ]
            else:
                route_results = []
            
            return route_results
            
        except Exception as e:
            logger.error(f"任务1（实体抽取）失败: {e}")
            return []
    
    async def _rewrite_query(self, query: str) -> Dict[str, Any]:
        """
        任务2：重写用户问题并进行分类
        
        Args:
            query: 用户查询
            
        Returns:
            Dict: 包含重写后的查询和分类信息
                - rewritten_query: 重写后的查询字符串
                - query_classification: 分类信息字典
                    - query_type: "aggregate"（统计类）或 "detail"（明细类）
                    - has_time: 是否包含时间信息（bool）
                    - time_type: 时间统计类型（"single"/"dual"/"multi"/None）
        """
        # 使用 init_prompts 中的 prompt 生成函数
        prompt = get_rewrite_query_prompt(
            query=query,
            schema=self.schema
        )

        
        try:
            llm_start = time.time()
            model = os.getenv("OPENAI_MODEL")
            response = await self.llm_api.open_api_chat_async_json(prompt, model, temperature=0)
            
            # 解析JSON响应
            content = response.strip()
            
            # 如果包含markdown代码块，提取JSON
            if "```json" in content:
                json_start = content.find("```json") + 7
                json_end = content.find("```", json_start)
                if json_end > json_start:
                    content = content[json_start:json_end].strip()
            elif "```" in content:
                json_start = content.find("```") + 3
                json_end = content.find("```", json_start)
                if json_end > json_start:
                    content = content[json_start:json_end].strip()
            
            # 尝试找到JSON数组或对象的开始和结束位置
            if not content.startswith('[') and not content.startswith('{'):
                for start_char in ['[', '{']:
                    start_idx = content.find(start_char)
                    if start_idx != -1:
                        content = content[start_idx:]
                        break
            
            # 解析JSON
            try:
                outputs_data = json.loads(content)
            except json.JSONDecodeError as e:
                logger.warning(f"任务2 JSON解析失败，尝试修复: {e}")
                if content.rfind(']') > content.rfind('}'):
                    end_idx = content.rfind(']') + 1
                    content = content[:end_idx]
                elif content.rfind('}') > -1:
                    end_idx = content.rfind('}') + 1
                    content = content[:end_idx]
                try:
                    outputs_data = json.loads(content)
                except json.JSONDecodeError:
                    logger.error(f"任务2 JSON解析最终失败，原始内容: {content[:200]}...")
                    return query
            
            # 提取rewritten_query和分类信息
            if isinstance(outputs_data, dict) and "rewritten_query" in outputs_data:
                rewritten_query = outputs_data.get("rewritten_query", query)
                # 提取分类信息
                query_classification = outputs_data.get("query_classification", {})
                if not isinstance(query_classification, dict):
                    query_classification = {}
                
                # 确保分类信息格式正确
                result = {
                    "rewritten_query": rewritten_query,
                    "query_classification": {
                        "query_type": query_classification.get("query_type", "detail"),
                        "has_time": query_classification.get("has_time", False),
                        "time_type": query_classification.get("time_type")
                    }
                }
            else:
                # 如果解析失败，返回默认值
                result = {
                    "rewritten_query": query,
                    "query_classification": {
                        "query_type": "detail",
                        "has_time": False,
                        "time_type": None
                    }
                }
            
            return result
            
        except Exception as e:
            logger.error(f"任务2（查询重写）失败: {e}")
            # 返回默认值
            return {
                "rewritten_query": query,
                "query_classification": {
                    "query_type": "detail",
                    "has_time": False,
                    "time_type": None
                }
            }
    
    async def route(self, query: str) -> Tuple[List[Dict[str, str]], str, Dict[str, Any]]:
        """
        路由标签识别：识别标签，抽取实体
        
        功能：使用LLM识别用户查询中涉及的节点类型和实体。
        将两个任务（实体抽取和查询重写）拆分为独立的异步调用，并行执行以提高效率。
        
        Args:
            query: 用户查询
            
        Returns:
            tuple: (路由结果列表, 重写后的查询, 分类信息)
                路由结果列表格式如 [{"label": "Drug", "entity": "xxx"}]
                重写后的查询字符串，如果未获取到则返回原始query
                分类信息字典，包含 query_type, has_time, time_type
        """
        try:
            # ========== 并行执行两个任务 ==========
            total_start = time.time()
            logger.info(f"[时间统计] RouteLabel: 开始并行执行两个任务（实体抽取 + 查询重写+分类）...")
            
            # 使用 asyncio.gather 并行执行两个任务
            route_results, rewrite_result = await asyncio.gather(
                self._extract_entities(query),
                self._rewrite_query(query)
            )
            print(route_results) 
            logger.info(111111111111111111111111111111111111111)           
            # 从重写结果中提取信息
            if isinstance(rewrite_result, dict):
                rewritten_query = rewrite_result.get("rewritten_query", query)
                query_classification = rewrite_result.get("query_classification", {
                    "query_type": "detail",
                    "has_time": False,
                    "time_type": None
                })
            else:
                # 兼容旧格式（如果返回的是字符串）
                rewritten_query = rewrite_result if isinstance(rewrite_result, str) else query
                query_classification = {
                    "query_type": "detail",
                    "has_time": False,
                    "time_type": None
                }
            
            total_elapsed = (time.time() - total_start) * 1000
            logger.info(f"[时间统计] RouteLabel: 两个任务并行执行完成，总耗时: {total_elapsed:.2f}ms")
            logger.info(f"最终结果 - 入口节点标签与实体: {route_results}")
            logger.info(f"最终结果 - 原始查询: '{query}' -> 重写后查询: '{rewritten_query}'")
            logger.info(f"最终结果 - 分类信息: {query_classification}")
            
            return route_results, rewritten_query, query_classification
            
        except Exception as e:
            logger.error(f"路由标签识别失败: {e}")
            return [], query, {
                "query_type": "detail",
                "has_time": False,
                "time_type": None
            }

