kgrag/graph_search/route_label.py
2026-06-30 13:35:52 +08:00

316 lines
13 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
===========================================
路由标签识别模块 - 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)
)
logger.info(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
}