wx-agent/utils/ship_number_search528.py
2026-06-30 13:47:51 +08:00

468 lines
18 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.

"""
舷号搜索策略模块
提供三级搜索策略:
1. 指定舷号搜索
2. 同型号所有舷号搜索
3. 全局搜索
新增功能:
- 智能舷号映射:支持舷号、型号、舰名的智能识别和映射
"""
from typing import List, Dict, Any, Optional, Tuple
import random
import re
import httpx
from utils.function_tracker import get_all_callbacks
from tools.ship_model_db import load_ship_model_mapping
from tools.rag_tools import rag_search
from modelsAPI.model_api import OpenaiAPI
from config import KB_TREE_CONFIG
# ==================== 舷号 -> kb_id 映射 ====================
_kb_name_to_id_cache: Optional[Dict[str, str]] = None
async def _load_kb_tree() -> Dict[str, str]:
"""
从知识库树 API 加载所有知识库,返回 {名称: kb_id} 的映射
递归遍历树结构,带缓存,只请求一次
"""
global _kb_name_to_id_cache
if _kb_name_to_id_cache is not None:
return _kb_name_to_id_cache
url = KB_TREE_CONFIG.get("url", "")
if not url:
print("[ship_number_search] KB_TREE_CONFIG.url 未配置")
return {}
headers = {
"accept": "application/json",
"x-user-id": "1",
"x-user-name": "testuser",
"x-role": "admin",
}
try:
async with httpx.AsyncClient(timeout=30.0, verify=False) as client:
response = await client.get(url, headers=headers)
if response.status_code != 200:
print(f"[ship_number_search] 知识库树 API 请求失败: HTTP {response.status_code}")
return {}
result = response.json()
# API 返回的是列表,每个元素可能有 children
if isinstance(result, list):
top_nodes = result
elif isinstance(result, dict):
kb_data = result.get("data") or result.get("list") or {}
top_nodes = kb_data.get("children", []) if isinstance(kb_data, dict) else []
else:
top_nodes = []
# 递归收集所有知识库
name_to_id = {}
def _collect(nodes):
for node in nodes:
kb_name = node.get("name") or ""
kb_id = str(node.get("id", "")) if node.get("id") is not None else ""
if kb_name and kb_id:
name_to_id[kb_name] = kb_id
children = node.get("children", [])
if children:
_collect(children)
_collect(top_nodes)
_kb_name_to_id_cache = name_to_id
return name_to_id
except Exception as e:
print(f"[ship_number_search] 加载知识库树失败: {str(e)}")
return {}
async def resolve_kb_id(ship_number: str) -> Optional[str]:
"""
根据舷号从知识库树中查找对应的真实 kb_id
匹配策略:在知识库名称中查找包含该舷号的条目
例如 ship_number="163",知识库名称为 "163舰资料",则匹配成功
Args:
ship_number: 舷号(如 "163"
Returns:
匹配到的 kb_id未找到返回 None
"""
if not ship_number:
return None
name_to_id = await _load_kb_tree()
# 优先精确匹配:名称就是舷号本身
if ship_number in name_to_id:
return name_to_id[ship_number]
# 其次:名称包含该舷号(确保是完整数字,避免 "16" 匹配到 "163"
for kb_name, kb_id in name_to_id.items():
# 用正则匹配完整数字,确保舷号作为独立数字出现在名称中
if re.search(rf'(?<!\d){re.escape(ship_number)}(?!\d)', kb_name):
return kb_id
return None
async def get_ship_numbers_by_model(ship_number: str) -> list:
"""
根据舷号获取同型号的所有舷号
Args:
ship_number: 舷号
Returns:
list: 同型号的所有舷号列表,如果未找到则返回只包含输入舷号的列表
"""
ship_number = str(ship_number).strip()
ship_model_mapping = await load_ship_model_mapping()
for model, info in ship_model_mapping.items():
numbers = info.get("numbers", [])
if ship_number in numbers:
return numbers.copy()
return [ship_number]
async def _build_ship_mapping_candidates() -> Dict[str, Dict[str, Any]]:
"""
构建舷号映射候选词库
舰名与舷号按位置一一对应,如 names[0] 对应 numbers[0]
"""
candidates = {}
ship_model_mapping = await load_ship_model_mapping()
for model_name, info in ship_model_mapping.items():
numbers = info.get("numbers", [])
aliases = info.get("aliases", [])
names = info.get("names", [])
mapping_info = {
"model": model_name,
"numbers": numbers.copy()
}
for num in numbers:
candidates[str(num)] = {
"type": "number",
"ship_number": str(num),
**mapping_info
}
for alias in aliases:
if alias:
candidates[alias] = {
"type": "alias",
"ship_number": None,
**mapping_info
}
for idx, name in enumerate(names):
if name:
ship_num = str(numbers[idx]) if idx < len(numbers) else None
candidates[name] = {
"type": "name",
"ship_number": ship_num,
**mapping_info
}
return candidates
_SHIP_MAPPING_CANDIDATES = None
async def get_ship_mapping_candidates() -> Dict[str, Dict[str, Any]]:
"""
获取舷号映射候选词库(带缓存)
"""
global _SHIP_MAPPING_CANDIDATES
if _SHIP_MAPPING_CANDIDATES is None:
_SHIP_MAPPING_CANDIDATES = await _build_ship_mapping_candidates()
return _SHIP_MAPPING_CANDIDATES
async def smart_ship_number_mapping(
user_input: str,
similarity_threshold: float = 0.8
) -> Tuple[Optional[str], Optional[str], Optional[List[str]], str]:
"""
智能舷号映射函数
将用户输入的舷号/型号/舰名映射到准确的舷号或型号
Returns:
Tuple: (matched_ship_number, matched_model, all_numbers, match_type)
- matched_ship_number: 匹配到的具体舷号(舰名也会映射到对应舷号)
- matched_model: 匹配到的型号名称
- all_numbers: 该型号下的所有舷号列表
- match_type: 匹配类型 ("exact_number", "exact_alias", "exact_name", "embedding", "none")
"""
if not user_input or not user_input.strip():
return None, None, None, "none"
user_input = user_input.strip()
candidates = await get_ship_mapping_candidates()
if user_input in candidates:
info = candidates[user_input]
match_type = f"exact_{info['type']}"
model = info["model"]
numbers = info["numbers"]
ship_number = info.get("ship_number")
return ship_number, model, numbers, match_type
try:
candidate_texts = list(candidates.keys())
query_embedding = await OpenaiAPI.get_embeddings_async(user_input)
candidate_embeddings = await OpenaiAPI.get_embeddings_batch_async(candidate_texts)
best_score = -1
best_idx = -1
for idx, emb in enumerate(candidate_embeddings):
score = OpenaiAPI.cosine_similarity(query_embedding, emb)
if score > best_score:
best_score = score
best_idx = idx
if best_score >= similarity_threshold and best_idx >= 0:
matched_text = candidate_texts[best_idx]
info = candidates[matched_text]
model = info["model"]
numbers = info["numbers"]
ship_number = info.get("ship_number")
print(f"[智能映射] 用户输入 '{user_input}' -> 匹配到 '{matched_text}' (相似度: {best_score:.3f}, 类型: {info['type']})")
return ship_number, model, numbers, "embedding"
except Exception as e:
print(f"[智能映射] Embedding 匹配失败: {str(e)}")
return None, None, None, "none"
def send_search_status(title: str, details: str):
"""发送搜索状态事件"""
title_options = [f"🔍 {title}", f"🔎 {title}", f"📚 {title}"]
start_event = {
"type": "function_execution",
"title": random.choice(title_options),
"details": details
}
for callback in get_all_callbacks():
try:
callback(start_event)
except Exception:
pass
def convert_rag_result(rag_result):
"""将 rag_search 的返回结果转换为原来的格式"""
if not rag_result or not rag_result.get("success", False):
return []
source_citation = rag_result.get("sourceCitation", {})
flat_results = []
for resource, items in source_citation.items():
for item in items:
flat_results.append(item)
return flat_results
async def search_with_ship_number_strategy(
base_query: str,
ship_number: str,
top_k: int = 5,
search_type: str = "fault"
) -> Tuple[List[Dict[str, Any]], str, str, str]:
"""
使用舷号三级搜索策略进行RAG检索
返回: (rag_results, matched_kb_name, matched_kb_id, matched_ship_number)
matched_ship_number: 映射后的舷号,未映射到具体舷号时为空字符串
"""
rag_results = []
matched_kb_name = ""
matched_kb_id = ""
matched_ship_number = None
matched_model = None
model_all_numbers = None
match_type = "none"
if ship_number and ship_number.strip():
send_search_status(
"正在智能识别",
f"识别 '{ship_number}' 的含义(舷号/型号/舰名)..."
)
matched_ship_number, matched_model, model_all_numbers, match_type = await smart_ship_number_mapping(ship_number)
if match_type != "none":
print(f"[智能映射结果] 输入='{ship_number}' -> 舷号={matched_ship_number}, 型号={matched_model}, 所有舷号={model_all_numbers}, 匹配类型={match_type}")
else:
print(f"[智能映射结果] 输入='{ship_number}' -> 未找到匹配,将使用原始值进行检索")
if match_type in ["exact_number", "exact_name", "embedding"] and matched_ship_number:
real_kb_id = await resolve_kb_id(matched_ship_number)
if real_kb_id:
send_search_status(
"正在指定舷号检索",
f"使用舷号 {matched_ship_number} 进行检索..."
)
rag_query = f"{base_query}"
try:
rag_result_raw = await rag_search(rag_query, top_k=top_k, kb_id=real_kb_id)
rag_results = convert_rag_result(rag_result_raw)
if rag_results:
print(f"步骤1 - 舷号 {matched_ship_number} 检索完成,返回 {len(rag_results)} 条结果")
for item in rag_results:
if isinstance(item, dict):
found_kb_name = item.get("kb_name", "")
found_kb_id = item.get("kb_id", "")
if found_kb_name and not matched_kb_name:
matched_kb_name = str(found_kb_name).strip()
if found_kb_id and not matched_kb_id:
matched_kb_id = str(found_kb_id).strip()
return rag_results, matched_kb_name, matched_kb_id, matched_ship_number or ""
except Exception as e:
print(f"步骤1 - 舷号 {matched_ship_number} 检索失败: {str(e)}")
else:
print(f"步骤1 - 舷号 {matched_ship_number} 无对应知识库,跳过指定舷号检索")
if match_type in ["exact_number", "exact_alias", "exact_name", "embedding"] and model_all_numbers:
send_search_status(
"正在型号检索",
f"识别为型号 '{matched_model}',使用该型号所有舷号 {', '.join(model_all_numbers)} 进行检索..."
)
for num in model_all_numbers:
rag_query = f"{base_query}"
try:
real_kb_id = await resolve_kb_id(num)
if not real_kb_id:
continue
rag_result_raw = await rag_search(rag_query, top_k=top_k, kb_id=real_kb_id)
current_results = convert_rag_result(rag_result_raw)
if current_results:
rag_results.extend(current_results)
print(f"步骤2 - 舷号 {num} 检索完成,返回 {len(current_results)} 条结果")
if not matched_kb_name:
for item in current_results:
if isinstance(item, dict):
found_kb_name = item.get("kb_name", "")
found_kb_id = item.get("kb_id", "")
if found_kb_name and not matched_kb_name:
matched_kb_name = str(found_kb_name).strip()
if found_kb_id and not matched_kb_id:
matched_kb_id = str(found_kb_id).strip()
except Exception as e:
print(f"步骤2 - 舷号 {num} 检索失败: {str(e)}")
if rag_results:
return rag_results, matched_kb_name, matched_kb_id, matched_ship_number or ""
if ship_number and ship_number.strip() and match_type == "none":
real_kb_id = await resolve_kb_id(ship_number)
if real_kb_id:
send_search_status(
"正在指定舰名检索",
f"使用舰名 {ship_number} 进行检索..."
)
rag_query = f"{base_query}"
try:
rag_result_raw = await rag_search(rag_query, top_k=top_k, kb_id=real_kb_id)
rag_results = convert_rag_result(rag_result_raw)
if rag_results:
print(f"步骤1 - 舷号 {ship_number} 检索完成,返回 {len(rag_results)} 条结果")
for item in rag_results:
if isinstance(item, dict):
found_kb_name = item.get("kb_name", "")
found_kb_id = item.get("kb_id", "")
if found_kb_name and not matched_kb_name:
matched_kb_name = str(found_kb_name).strip()
if found_kb_id and not matched_kb_id:
matched_kb_id = str(found_kb_id).strip()
return rag_results, matched_kb_name, matched_kb_id, matched_ship_number or ""
except Exception as e:
print(f"步骤1 - 舷号 {ship_number} 检索失败: {str(e)}")
else:
print(f"步骤1 - 舷号 {ship_number} 无对应知识库,跳过指定舰名检索")
model_ship_numbers = await get_ship_numbers_by_model(ship_number)
if len(model_ship_numbers) > 1 or (len(model_ship_numbers) == 1 and ship_number not in model_ship_numbers):
send_search_status(
"正在同型号检索",
f"使用同型号舷号 {', '.join(model_ship_numbers)} 进行检索..."
)
for num in model_ship_numbers:
if num == ship_number:
continue
rag_query = f"{base_query}"
try:
real_kb_id = await resolve_kb_id(num)
if not real_kb_id:
continue
rag_result_raw = await rag_search(rag_query, top_k=top_k, kb_id=real_kb_id)
current_results = convert_rag_result(rag_result_raw)
if current_results:
rag_results.extend(current_results)
print(f"步骤2 - 舷号 {num} 检索完成,返回 {len(current_results)} 条结果")
if not matched_kb_name:
for item in current_results:
if isinstance(item, dict):
found_kb_name = item.get("kb_name", "")
found_kb_id = item.get("kb_id", "")
if found_kb_name and not matched_kb_name:
matched_kb_name = str(found_kb_name).strip()
if found_kb_id and not matched_kb_id:
matched_kb_id = str(found_kb_id).strip()
except Exception as e:
print(f"步骤2 - 舷号 {num} 检索失败: {str(e)}")
if rag_results:
return rag_results, matched_kb_name, matched_kb_id, matched_ship_number or ""
send_search_status(
"正在全局检索",
"未找到舷号相关资料,进行全局检索..."
)
try:
rag_result_raw = await rag_search(base_query, top_k=top_k)
rag_results = convert_rag_result(rag_result_raw)
if rag_results:
print(f"步骤3 - 全局检索完成,返回 {len(rag_results)} 条结果")
for item in rag_results:
if isinstance(item, dict):
found_kb_name = item.get("kb_name", "")
found_kb_id = item.get("kb_id", "")
if found_kb_name and not matched_kb_name:
matched_kb_name = str(found_kb_name).strip()
if found_kb_id and not matched_kb_id:
matched_kb_id = str(found_kb_id).strip()
except Exception as e:
print(f"步骤3 - 全局检索失败: {str(e)}")
rag_results = []
return rag_results, matched_kb_name, matched_kb_id, matched_ship_number or ""