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

970 lines
30 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.

"""
故障类型统计工具模块
用于统计故障诊断智能体记录的故障信息
数据来源PostgreSQL fault_records 表
"""
import os
import re
import json
import httpx
from datetime import datetime, timedelta
from typing import Dict, Any, Optional, List
from difflib import SequenceMatcher
import asyncio
from langchain_core.tools import tool
from utils.function_tracker import track_function_calls
from config import REPAIR_FEEDBACK_CONFIG
neo4j_cache = {}
def fuzzy_match_system(user_input: str, system_name: str, threshold: float = 0.6) -> bool:
if not user_input or not system_name:
return False
user_input = user_input.strip().lower()
system_name = system_name.strip().lower()
if user_input == system_name:
return True
if user_input in system_name or system_name in user_input:
return True
similarity = SequenceMatcher(None, user_input, system_name).ratio()
return similarity >= threshold
async def get_device_system_from_neo4j(device_name: str) -> str:
"""
从 Neo4j 图谱中查询设备对应的系统(调用外部接口),带缓存机制
"""
if not device_name or not device_name.strip():
return "其他"
device_name_clean = device_name.strip().lower()
if device_name_clean in neo4j_cache:
return neo4j_cache[device_name_clean]
from config import DEVICE_SYSTEM_CONFIG
endpoint = DEVICE_SYSTEM_CONFIG.get("endpoint")
timeout = DEVICE_SYSTEM_CONFIG.get("timeout", 30.0)
min_hops = DEVICE_SYSTEM_CONFIG.get("default_min_hops", 2)
max_hops = DEVICE_SYSTEM_CONFIG.get("default_max_hops", 8)
top_k = DEVICE_SYSTEM_CONFIG.get("default_top_k", 10)
if not endpoint:
result = "其他"
neo4j_cache[device_name_clean] = result
return result
try:
payload = {
"device_name": device_name.strip(),
"min_hops": min_hops,
"max_hops": max_hops,
"top_k": top_k
}
async with httpx.AsyncClient(timeout=timeout) as client:
response = await client.post(
endpoint,
json=payload,
headers={"Content-Type": "application/json"}
)
if response.status_code != 200:
result = "其他"
neo4j_cache[device_name_clean] = result
return result
result_json = response.json()
if not result_json.get("success", False):
result = "其他"
neo4j_cache[device_name_clean] = result
return result
systems = result_json.get("systems", [])
if not systems or not isinstance(systems, list):
result = "其他"
neo4j_cache[device_name_clean] = result
return result
first_system = systems[0]
if isinstance(first_system, dict):
system_name = first_system.get("name") or first_system.get("系统名称") or first_system.get("名称")
if system_name:
result = system_name.strip()
neo4j_cache[device_name_clean] = result
return result
result = "其他"
neo4j_cache[device_name_clean] = result
return result
except Exception as e:
print(f"查询设备系统失败: {str(e)}")
result = "其他"
neo4j_cache[device_name_clean] = result
return result
def generate_color_palette(n: int) -> List[str]:
predefined_colors = [
"#00C2FF", "#4E7CFE", "#9153FF", "#FFB800", "#00D16B",
"#FF6B6B", "#4ECDC4", "#45B7D1", "#FFA07A", "#98D8C8",
"#F7DC6F", "#BB8FCE", "#85C1E9", "#F8B500", "#00E676",
"#FF4081", "#7C4DFF", "#18FFFF", "#69F0AE", "#FFD740"
]
if n <= len(predefined_colors):
return predefined_colors[:n]
colors = predefined_colors.copy()
import colorsys
for i in range(len(predefined_colors), n):
hue = (i * 0.618033988749895) % 1.0
saturation = 0.7 + (i % 3) * 0.1
value = 0.8 + (i % 2) * 0.1
rgb = colorsys.hsv_to_rgb(hue, saturation, value)
color = "#{:02X}{:02X}{:02X}".format(
int(rgb[0] * 255),
int(rgb[1] * 255),
int(rgb[2] * 255)
)
colors.append(color)
return colors
async def _query_fault_records_from_db(
ship_number: Optional[str] = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
system_name: Optional[str] = None,
limit: int = 5000
) -> List[Dict[str, Any]]:
"""
从数据库查询故障记录(内部共用方法)
"""
from tools.fault_record_db import query_fault_records
return await query_fault_records(
ship_number=ship_number,
system_name=system_name,
start_date=start_date,
end_date=end_date,
limit=limit
)
async def statistics_fault_types(
ship_number_filter: Optional[str] = None
) -> Dict[str, Any]:
"""
统计故障类型(从数据库读取),按系统分类计算百分比
Args:
ship_number_filter: 舷号过滤(可选)
Returns:
包含统计结果的字典
"""
try:
records = await _query_fault_records_from_db(ship_number=ship_number_filter)
if not records:
return {
"success": True,
"faultTypeData": []
}
system_totals = {}
valid_count = 0
for record in records:
system_name = record.get("system_name", "其他")
if not system_name:
system_name = "其他"
valid_count += 1
if system_name not in system_totals:
system_totals[system_name] = 0
system_totals[system_name] += 1
fault_type_data = []
if valid_count > 0:
unique_systems = sorted(system_totals.keys())
colors = generate_color_palette(len(unique_systems))
system_color_map = {system: color for system, color in zip(unique_systems, colors)}
for system_name, count in system_totals.items():
percentage = round((count / valid_count) * 100)
color = system_color_map.get(system_name, "#6E7D91")
fault_type_data.append({
"value": percentage,
"name": system_name,
"itemStyle": {
"color": color
}
})
fault_type_data.sort(key=lambda x: x["value"], reverse=True)
return {
"success": True,
"faultTypeData": fault_type_data
}
except Exception as e:
print(f"[故障类型统计错误] {str(e)}")
import traceback
traceback.print_exc()
return {
"success": False,
"faultTypeData": []
}
@tool
async def fault_type_statistics_tool(
ship_number_filter: Optional[str] = None
) -> Dict[str, Any]:
"""
故障类型统计工具(供工作流调用)
统计故障类型,按系统分类并计算百分比
Args:
ship_number_filter: 舷号过滤(可选)
Returns:
统计结果,包含 success 和 faultTypeData 字段
"""
result = await statistics_fault_types(ship_number_filter=ship_number_filter)
return result
async def merge_similar_faults_with_embedding(
fault_counter: Dict[str, int],
fault_system_map: Dict[str, str],
similarity_threshold: float = 0.85
) -> tuple:
"""
使用embedding合并相似的故障
"""
if not fault_counter:
return {}, {}, {}
fault_keys = list(fault_counter.keys())
if len(fault_keys) <= 1:
return fault_counter.copy(), fault_system_map.copy(), {}
try:
from modelsAPI.model_api import OpenaiAPI
embeddings = await OpenaiAPI.get_embeddings_batch_async(fault_keys)
fault_embedding_map = {key: emb for key, emb in zip(fault_keys, embeddings)}
merged_counter = {}
merged_system_map = {}
merge_mapping = {}
merged_keys = set()
sorted_faults = sorted(fault_counter.items(), key=lambda x: x[1], reverse=True)
for fault_key, count in sorted_faults:
if fault_key in merged_keys:
continue
merged_keys.add(fault_key)
merged_counter[fault_key] = count
merged_system_map[fault_key] = fault_system_map.get(fault_key, "其他")
for other_key, other_count in sorted_faults:
if other_key in merged_keys or other_key == fault_key:
continue
if fault_embedding_map.get(fault_key) and fault_embedding_map.get(other_key):
similarity = OpenaiAPI.cosine_similarity(
fault_embedding_map[fault_key],
fault_embedding_map[other_key]
)
if similarity >= similarity_threshold:
merged_counter[fault_key] += other_count
merged_keys.add(other_key)
merge_mapping[other_key] = fault_key
return merged_counter, merged_system_map, merge_mapping
except Exception as e:
print(f"Embedding合并失败使用原始数据: {str(e)}")
return fault_counter.copy(), fault_system_map.copy(), {}
async def statistics_fault_frequency(
top_n: int = 10,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
system_filter: Optional[str] = None,
use_embedding_merge: bool = True,
similarity_threshold: float = 0.85
) -> Dict[str, Any]:
"""
统计故障发生次数(从数据库读取),按频率排序
统计所有故障的数据但只返回top N展示在表格中
Args:
top_n: 返回前 N 个高频故障(默认为 10
start_date: 开始日期格式YYYY-MM-DD可选
end_date: 结束日期格式YYYY-MM-DD可选
system_filter: 系统名称过滤(可选)
use_embedding_merge: 是否使用embedding合并相似故障默认True
similarity_threshold: 相似度阈值默认0.85
Returns:
包含统计结果的字典
"""
try:
records = await _query_fault_records_from_db(
start_date=start_date,
end_date=end_date
)
if not records:
return {
"success": True,
"totalCount": 0,
"totalTypes": 0,
"highFreqFaults": []
}
fault_counter = {}
fault_system_map = {}
for record in records:
device_name = record.get("device_name", "")
fault = record.get("fault", "")
system_name = record.get("system_name", "其他")
if not device_name or not fault:
continue
if system_filter and not fuzzy_match_system(system_filter, system_name):
continue
fault_key = f"{device_name}{fault}"
if fault_key not in fault_counter:
fault_counter[fault_key] = 0
fault_system_map[fault_key] = system_name
fault_counter[fault_key] += 1
if use_embedding_merge and len(fault_counter) > 1:
fault_counter, fault_system_map, merge_mapping = await merge_similar_faults_with_embedding(
fault_counter,
fault_system_map,
similarity_threshold
)
if merge_mapping:
print(f"合并了 {len(merge_mapping)} 个相似故障")
# 先计算全部数据的统计量
total_count = sum(fault_counter.values())
total_types = len(fault_counter)
# 对全部数据排序
sorted_faults_all = sorted(fault_counter.items(), key=lambda x: x[1], reverse=True)
# 只取top N用于展示表格
sorted_faults = sorted_faults_all[:top_n]
high_freq_faults = []
for rank, (fault_key, count) in enumerate(sorted_faults, start=1):
system_name = fault_system_map.get(fault_key, "其他")
high_freq_faults.append({
"rank": rank,
"name": fault_key,
"system": system_name,
"count": count
})
return {
"success": True,
"totalCount": total_count,
"totalTypes": total_types,
"highFreqFaults": high_freq_faults
}
except Exception as e:
print(f"[故障频率统计错误] {str(e)}")
import traceback
traceback.print_exc()
return {
"success": False,
"highFreqFaults": []
}
@tool
async def fault_frequency_statistics_tool(
top_n: Optional[int] = 6,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
system_filter: Optional[str] = None
) -> Dict[str, Any]:
"""
故障发生次数统计工具(供工作流调用)
统计故障发生次数,按频率排序
Args:
top_n: 返回前 N 个高频故障(默认为 6
start_date: 开始日期格式YYYY-MM-DD可选
end_date: 结束日期格式YYYY-MM-DD可选
system_filter: 系统名称过滤(可选)
Returns:
统计结果,包含 success 和 highFreqFaults 字段
"""
result = await statistics_fault_frequency(
top_n=top_n or 10,
start_date=start_date,
end_date=end_date,
system_filter=system_filter
)
return result
async def statistics_spare_parts(
ship_number_filter: Optional[str] = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
top_n: int = 6
) -> Dict[str, Any]:
"""
统计备件消耗(从数据库读取)
统计所有备件的数据但只返回top N展示在表格中
Args:
ship_number_filter: 舷号过滤(可选)
start_date: 开始日期格式YYYY-MM-DD可选
end_date: 结束日期格式YYYY-MM-DD可选
top_n: 表格中展示的前 N 个备件(默认为 5
Returns:
包含统计结果的字典
"""
try:
records = await _query_fault_records_from_db(
ship_number=ship_number_filter,
start_date=start_date,
end_date=end_date
)
if not records:
return {
"success": True,
"totalCount": 0,
"totalTypes": 0,
"sparePartData": {
"categories": [],
"system": [],
"counter": []
}
}
spare_part_counter = {}
spare_part_system_map = {}
for record in records:
spare_parts = record.get("spare_parts", [])
system_name = record.get("system_name", "其他")
if not spare_parts:
continue
for part in spare_parts:
if not part or not part.strip():
continue
part_key = part.strip()
if part_key not in spare_part_counter:
spare_part_counter[part_key] = 0
spare_part_system_map[part_key] = system_name
spare_part_counter[part_key] += 1
# 先计算全部数据的统计量
total_count = sum(spare_part_counter.values())
total_types = len(spare_part_counter)
# 对全部数据排序
sorted_parts_all = sorted(spare_part_counter.items(), key=lambda x: x[1], reverse=True)
# 只取top N用于展示表格
sorted_parts = sorted_parts_all[:top_n]
categories = []
systems = []
counters = []
for part_key, count in sorted_parts:
categories.append(part_key)
systems.append(spare_part_system_map.get(part_key, "其他"))
counters.append(count)
return {
"success": True,
"totalCount": total_count,
"totalTypes": total_types,
"sparePartData": {
"categories": categories,
"system": systems,
"counter": counters
}
}
except Exception as e:
print(f"[备件统计错误] {str(e)}")
import traceback
traceback.print_exc()
return {
"success": False,
"sparePartData": {
"categories": [],
"system": [],
"counter": []
}
}
@tool
async def spare_parts_statistics_tool(
ship_number: Optional[str] = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
top_n: Optional[int] = 5
) -> Dict[str, Any]:
"""
备件消耗统计工具(供工作流调用)
统计备件消耗情况
Args:
ship_number: 舷号过滤(可选)
start_date: 开始日期格式YYYY-MM-DD可选
end_date: 结束日期格式YYYY-MM-DD可选
top_n: 返回前 N 个备件(默认为 5
Returns:
统计结果,包含 success 和 sparePartData 字段
"""
result = await statistics_spare_parts(
ship_number_filter=ship_number,
start_date=start_date,
end_date=end_date,
top_n=top_n or 5
)
return result
async def extract_device_from_text(text: str) -> Optional[str]:
"""
从文本中提取设备名称(使用大模型)
"""
if not text or not text.strip():
return None
prompt = f"""从以下文本中提取设备名称:
文本内容:
{text}
请识别文本中提到的具体设备名称,如:发动机、压缩机、水泵、发电机、空调机组、液压泵等。
请以JSON格式返回
{{"device": "设备名称"}}
如果文本中没有提到任何设备,请返回 {{"device": null}}
只返回JSON不要有其他内容。"""
try:
from modelsAPI.model_api import OpenaiAPI
result_text = await OpenaiAPI.open_api_chat_without_thinking(
query=prompt,
json_output=True
)
result_text = result_text.strip()
result_text = re.sub(r'^```json\s*', '', result_text, flags=re.IGNORECASE)
result_text = re.sub(r'^```\s*', '', result_text)
result_text = re.sub(r'```$', '', result_text)
json_match = re.search(r'\{.*\}', result_text, re.DOTALL)
if json_match:
result = json.loads(json_match.group(0))
device = result.get("device")
if device and device != "null":
return str(device).strip()
return None
except Exception as e:
print(f"大模型提取设备名称失败: {str(e)}")
return None
async def classify_feedback_content(text: str) -> str:
"""
对反馈内容进行分类(使用大模型)
"""
if not text or not text.strip():
return "其他"
prompt = f"""请对以下反馈内容进行分类。
反馈内容:
{text}
分类类别(单选,选择最匹配的一个):
1. 信息过时 - 反映文档或手册中的信息已经过时,不再适用
2. 步骤错误 - 反映操作步骤或流程存在错误
3. 步骤缺失 - 反映缺少必要的操作步骤或流程说明
4. 备件信息有误 - 反映备件型号、规格、数量等信息有误
5. 安全提示不足 - 反映缺少必要的安全警示或注意事项
6. 图示不清 - 反映图片、示意图不清晰或难以理解
请分析反馈内容判断属于哪个类别以JSON格式返回
{{"category": "类别名称"}}
如果反馈内容不属于以上任何类别,请返回 {{"category": "其他"}}
只返回JSON不要有其他内容。"""
try:
from modelsAPI.model_api import OpenaiAPI
result_text = await OpenaiAPI.open_api_chat_without_thinking(
query=prompt,
json_output=True
)
result_text = result_text.strip()
result_text = re.sub(r'^```json\s*', '', result_text, flags=re.IGNORECASE)
result_text = re.sub(r'^```\s*', '', result_text)
result_text = re.sub(r'```$', '', result_text)
json_match = re.search(r'\{.*\}', result_text, re.DOTALL)
if json_match:
result = json.loads(json_match.group(0))
category = result.get("category")
valid_categories = ["信息过时", "步骤错误", "步骤缺失", "备件信息有误", "安全提示不足", "图示不清", "其他"]
if category in valid_categories:
return category
return "其他"
except Exception as e:
print(f"大模型分类反馈内容失败: {str(e)}")
return "其他"
async def analyze_feedback_and_classify(feedback_text: str) -> Dict[str, Any]:
"""
分析反馈内容并分类(主接口)
"""
if not feedback_text or not feedback_text.strip():
return {
"success": False,
"device": None,
"system": None,
"category": "其他",
"error": "反馈内容为空"
}
try:
device = await extract_device_from_text(feedback_text)
system = "其他"
if device:
system = await get_device_system_from_neo4j(device)
category = await classify_feedback_content(feedback_text)
return {
"success": True,
"system": system,
"category": category
}
except Exception as e:
return {
"success": False,
"system": None,
"category": "其他",
"error": str(e)
}
_feedback_token_cache = {"token": None, "expires_at": 0}
async def _get_feedback_api_token() -> Optional[str]:
"""
获取维修反馈API的认证token带缓存避免每次请求都认证
"""
import time
if _feedback_token_cache["token"] and time.time() < _feedback_token_cache["expires_at"]:
return _feedback_token_cache["token"]
base_url = REPAIR_FEEDBACK_CONFIG["url"]
email = REPAIR_FEEDBACK_CONFIG.get("email", "")
password = REPAIR_FEEDBACK_CONFIG.get("password", "")
if not email or not password:
print("[维修反馈认证错误] 未配置认证凭据")
return None
try:
async with httpx.AsyncClient(timeout=30.0) as client:
signin_resp = await client.post(
f"{base_url}/api/v1/auths/signin",
json={"email": email, "password": password}
)
if signin_resp.status_code != 200:
print(f"[维修反馈认证错误] 登录失败,状态码: {signin_resp.status_code}")
return None
token = signin_resp.json().get("token", "")
if not token:
print("[维修反馈认证错误] 响应中未找到token")
return None
_feedback_token_cache["token"] = token
_feedback_token_cache["expires_at"] = time.time() + 3600
return token
except Exception as e:
print(f"[维修反馈认证错误] {str(e)}")
return None
async def statistics_repair_feedbacks(
source: Optional[str] = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
top_n: int = 10
) -> Dict[str, Any]:
"""
统计维修反馈数据从外部API获取带认证
Args:
source: 来源筛选(好评/差评/反馈),可选
start_date: 开始日期格式YYYY-MM-DD可选
end_date: 结束日期格式YYYY-MM-DD可选
top_n: 返回前 N 条记录(默认为 10
Returns:
包含统计结果的字典
"""
try:
token = await _get_feedback_api_token()
if not token:
return {
"success": False,
"totalCount": 0,
"feedbackData": [],
"error": "API认证失败无法获取token"
}
base_url = REPAIR_FEEDBACK_CONFIG["url"]
endpoint = f"{base_url}/api/v1/repair-feedbacks/"
params = {}
if source and source.strip():
params["source"] = source.strip()
async with httpx.AsyncClient(timeout=30.0) as client:
response = await client.get(
endpoint,
params=params,
headers={
"Accept": "application/json",
"Content-Type": "application/json",
"Authorization": f"Bearer {token}"
}
)
if response.status_code == 401 or response.status_code == 403:
_feedback_token_cache["token"] = None
_feedback_token_cache["expires_at"] = 0
token = await _get_feedback_api_token()
if token:
response = await client.get(
endpoint,
params=params,
headers={
"Accept": "application/json",
"Content-Type": "application/json",
"Authorization": f"Bearer {token}"
}
)
if response.status_code != 200:
print(f"[维修反馈统计错误] API返回状态码: {response.status_code}")
return {
"success": False,
"totalCount": 0,
"feedbackData": [],
"error": f"API请求失败状态码: {response.status_code}"
}
result_json = response.json()
if isinstance(result_json, list):
feedback_list = result_json
elif isinstance(result_json, dict):
feedback_list = result_json.get("results", result_json.get("data", []))
else:
feedback_list = []
if not feedback_list:
return {
"success": True,
"totalCount": 0,
"totalTypes": 0,
"feedbackData": [],
"sourceDistribution": {}
}
start_dt = None
end_dt = None
if start_date:
try:
start_dt = datetime.strptime(start_date, "%Y-%m-%d")
except ValueError:
pass
if end_date:
try:
end_dt = datetime.strptime(end_date, "%Y-%m-%d") + timedelta(days=1) - timedelta(seconds=1)
except ValueError:
pass
if start_dt or end_dt:
filtered_list = []
for item in feedback_list:
if not isinstance(item, dict):
continue
raw_created_at = item.get("created_at", "")
if not raw_created_at:
continue
try:
item_dt = datetime.fromtimestamp(int(raw_created_at))
if start_dt and item_dt < start_dt:
continue
if end_dt and item_dt > end_dt:
continue
filtered_list.append(item)
except (ValueError, TypeError, OSError):
continue
feedback_list = filtered_list
total_count = len(feedback_list)
processed_data = []
for idx, item in enumerate(feedback_list[:top_n], start=1):
if isinstance(item, dict):
raw_title = item.get("title", "")
raw_content = item.get("content", "")
raw_source = item.get("source", source or "")
raw_created_at = item.get("created_at", "")
raw_user_name = item.get("user_name", "")
created_at_str = ""
if raw_created_at:
try:
ts = int(raw_created_at)
created_at_str = datetime.fromtimestamp(ts).strftime("%Y-%m-%d %H:%M")
except (ValueError, TypeError, OSError):
created_at_str = str(raw_created_at)
processed_item = {
"rank": idx,
"title": str(raw_title) if raw_title else "",
"content": str(raw_content)[:200] if raw_content else "",
"source": str(raw_source) if raw_source else "未知",
"created_at": created_at_str,
"user_name": str(raw_user_name) if raw_user_name else ""
}
processed_data.append(processed_item)
source_counter = {}
for item in processed_data:
s = item.get("source", "未知")
source_counter[s] = source_counter.get(s, 0) + 1
return {
"success": True,
"totalCount": total_count,
"totalTypes": len(feedback_list),
"sourceDistribution": source_counter,
"feedbackData": processed_data
}
except Exception as e:
print(f"[维修反馈统计错误] {str(e)}")
import traceback
traceback.print_exc()
return {
"success": False,
"totalCount": 0,
"feedbackData": [],
"error": str(e)
}
@tool
async def repair_feedback_statistics_tool(
source: Optional[str] = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
top_n: Optional[int] = 10
) -> Dict[str, Any]:
"""
维修反馈统计工具(供工作流调用)
统计维修反馈情况,支持按来源和时间筛选
Args:
source: 来源筛选(可选),如:好评、差评、反馈
start_date: 开始日期格式YYYY-MM-DD可选
end_date: 结束日期格式YYYY-MM-DD可选
top_n: 返回前 N 条记录(默认为 10
Returns:
统计结果,包含 success 和 feedbackData 字段
"""
result = await statistics_repair_feedbacks(
source=source,
start_date=start_date,
end_date=end_date,
top_n=top_n or 10
)
return result