331 lines
11 KiB
Python
331 lines
11 KiB
Python
"""
|
||
故障记录数据库模块
|
||
用于在故障诊断智能体生成方案后自动记录故障信息到PostgreSQL数据库
|
||
记录内容:舷号、设备、故障、备件、系统
|
||
"""
|
||
import json
|
||
import re
|
||
import asyncio
|
||
from datetime import datetime
|
||
from typing import Dict, Any, Optional, List
|
||
|
||
from config import POSTGRES_CONNECTION_STRING
|
||
|
||
|
||
async def init_fault_records_table():
|
||
"""
|
||
初始化 fault_records 表
|
||
"""
|
||
try:
|
||
from psycopg_pool import AsyncConnectionPool
|
||
except ImportError:
|
||
print("[fault_record_db] psycopg_pool 未安装,无法初始化表")
|
||
return
|
||
|
||
async with AsyncConnectionPool(
|
||
POSTGRES_CONNECTION_STRING,
|
||
kwargs={"autocommit": True}
|
||
) as pool:
|
||
async with pool.connection() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute("""
|
||
CREATE TABLE IF NOT EXISTS fault_records (
|
||
id SERIAL PRIMARY KEY,
|
||
ship_number TEXT NOT NULL DEFAULT '',
|
||
device_name TEXT NOT NULL DEFAULT '',
|
||
fault TEXT NOT NULL DEFAULT '',
|
||
spare_parts JSONB NOT NULL DEFAULT '[]',
|
||
system_name TEXT NOT NULL DEFAULT '',
|
||
created_at TIMESTAMP NOT NULL DEFAULT NOW()
|
||
)
|
||
""")
|
||
await cur.execute("""
|
||
CREATE INDEX IF NOT EXISTS fault_records_ship_number_idx ON fault_records(ship_number)
|
||
""")
|
||
await cur.execute("""
|
||
CREATE INDEX IF NOT EXISTS fault_records_device_name_idx ON fault_records(device_name)
|
||
""")
|
||
await cur.execute("""
|
||
CREATE INDEX IF NOT EXISTS fault_records_system_name_idx ON fault_records(system_name)
|
||
""")
|
||
await cur.execute("""
|
||
CREATE INDEX IF NOT EXISTS fault_records_created_at_idx ON fault_records(created_at)
|
||
""")
|
||
print("[fault_record_db] fault_records 表初始化完成")
|
||
|
||
|
||
|
||
async def extract_spare_parts_from_scheme(scheme_content: str) -> List[str]:
|
||
"""
|
||
从方案内容中用大模型抽取备件名称列表
|
||
|
||
Args:
|
||
scheme_content: 方案内容(Markdown格式)
|
||
|
||
Returns:
|
||
备件名称列表
|
||
"""
|
||
if not scheme_content or len(scheme_content.strip()) < 10:
|
||
return []
|
||
|
||
prompt = f"""从以下维修方案内容中提取所有提到的备品备件名称。
|
||
|
||
方案内容:
|
||
{scheme_content[:3000]}
|
||
|
||
提取规则:
|
||
1. 找到"备品备件"、"备件"等相关章节,提取其中的项目名称
|
||
2. 只提取名称,不要提取型号、规格、数量、材质、用途等
|
||
3. 如果没有专门的章节,也从全文中识别备件名称
|
||
|
||
以JSON数组格式返回:
|
||
{{
|
||
"备品备件": [备件名称1, 备件名称2, ...]
|
||
}}
|
||
|
||
如果没有提到备品备件,请返回空数组 []"""
|
||
|
||
try:
|
||
from modelsAPI.model_api import OpenaiAPI
|
||
result_text = await OpenaiAPI.open_api_chat_without_thinking(
|
||
query=prompt,
|
||
json_output=True
|
||
)
|
||
print("抽取出来的备品备件", result_text)
|
||
|
||
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)
|
||
|
||
try:
|
||
parsed = json.loads(result_text)
|
||
except json.JSONDecodeError:
|
||
json_match = re.search(r'\{.*\}', result_text, re.DOTALL)
|
||
if json_match:
|
||
parsed = json.loads(json_match.group(0))
|
||
else:
|
||
return []
|
||
|
||
if isinstance(parsed, dict):
|
||
if "备品备件" in parsed and isinstance(parsed["备品备件"], list):
|
||
return [str(part).strip() for part in parsed["备品备件"] if part]
|
||
all_parts = []
|
||
for v in parsed.values():
|
||
if isinstance(v, list):
|
||
all_parts.extend([str(part).strip() for part in v if part])
|
||
return all_parts
|
||
elif isinstance(parsed, list):
|
||
return [str(part).strip() for part in parsed if part]
|
||
|
||
return []
|
||
|
||
except Exception as e:
|
||
print(f"[fault_record_db] 从方案提取备件失败: {str(e)}")
|
||
return []
|
||
|
||
|
||
async def save_fault_record(
|
||
ship_number: str,
|
||
device_name: str,
|
||
fault: str,
|
||
spare_parts: List[str],
|
||
system_name: str
|
||
) -> bool:
|
||
"""
|
||
保存故障记录到数据库
|
||
|
||
Args:
|
||
ship_number: 舷号
|
||
device_name: 设备名称
|
||
fault: 故障现象
|
||
spare_parts: 备件列表
|
||
system_name: 系统名称
|
||
|
||
Returns:
|
||
是否保存成功
|
||
"""
|
||
if not ship_number and not device_name and not fault:
|
||
print("[fault_record_db] 舷号、设备、故障均为空,跳过保存")
|
||
return False
|
||
|
||
try:
|
||
from psycopg_pool import AsyncConnectionPool
|
||
except ImportError:
|
||
print("[fault_record_db] psycopg_pool 未安装,无法保存")
|
||
return False
|
||
|
||
try:
|
||
async with AsyncConnectionPool(
|
||
POSTGRES_CONNECTION_STRING,
|
||
kwargs={"autocommit": True}
|
||
) as pool:
|
||
async with pool.connection() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"""
|
||
INSERT INTO fault_records (ship_number, device_name, fault, spare_parts, system_name, created_at)
|
||
VALUES (%s, %s, %s, %s, %s, %s)
|
||
""",
|
||
(
|
||
ship_number or "",
|
||
device_name or "",
|
||
fault or "",
|
||
json.dumps(spare_parts, ensure_ascii=False),
|
||
system_name or "",
|
||
datetime.now()
|
||
)
|
||
)
|
||
print(f"[fault_record_db] 故障记录已保存: 舷号={ship_number}, 设备={device_name}, 故障={fault}, 系统={system_name}, 备件数={len(spare_parts)}")
|
||
return True
|
||
except Exception as e:
|
||
print(f"[fault_record_db] 保存故障记录失败: {str(e)}")
|
||
import traceback
|
||
traceback.print_exc()
|
||
return False
|
||
|
||
|
||
async def record_fault_from_state(state: Dict[str, Any]) -> bool:
|
||
"""
|
||
从工作流状态中提取故障信息并记录到数据库
|
||
在故障诊断智能体生成方案后调用
|
||
|
||
流程:
|
||
1. 从状态中获取舷号、设备、故障
|
||
2. 从方案内容中用大模型抽取备件
|
||
3. 通过设备名称反推系统
|
||
4. 保存到数据库
|
||
|
||
Args:
|
||
state: 工作流状态字典
|
||
|
||
Returns:
|
||
是否记录成功
|
||
"""
|
||
ship_number = state.get("ship_number", "") or ""
|
||
device_name = state.get("device_name", "") or ""
|
||
fault = state.get("fault", "") or ""
|
||
|
||
if not device_name and not fault:
|
||
print("[fault_record_db] 设备和故障均为空,跳过记录")
|
||
return False
|
||
|
||
scheme_content = state.get("last_generated_scheme", "") or state.get("response", "") or ""
|
||
|
||
spare_parts = await extract_spare_parts_from_scheme(scheme_content)
|
||
|
||
system_name = "其他"
|
||
if device_name:
|
||
from tools.fault_statistics import get_device_system_from_neo4j
|
||
system_name = await get_device_system_from_neo4j(device_name)
|
||
|
||
return await save_fault_record(
|
||
ship_number=ship_number,
|
||
device_name=device_name,
|
||
fault=fault,
|
||
spare_parts=spare_parts,
|
||
system_name=system_name
|
||
)
|
||
|
||
|
||
async def query_fault_records(
|
||
ship_number: Optional[str] = None,
|
||
device_name: Optional[str] = None,
|
||
system_name: Optional[str] = None,
|
||
start_date: Optional[str] = None,
|
||
end_date: Optional[str] = None,
|
||
limit: int = 1000
|
||
) -> List[Dict[str, Any]]:
|
||
"""
|
||
从数据库查询故障记录
|
||
|
||
Args:
|
||
ship_number: 舷号过滤
|
||
device_name: 设备名称过滤
|
||
system_name: 系统名称过滤
|
||
start_date: 开始日期(格式:YYYY-MM-DD)
|
||
end_date: 结束日期(格式:YYYY-MM-DD)
|
||
limit: 返回记录数上限
|
||
|
||
Returns:
|
||
故障记录列表
|
||
"""
|
||
try:
|
||
from psycopg_pool import AsyncConnectionPool
|
||
except ImportError:
|
||
print("[fault_record_db] psycopg_pool 未安装,无法查询")
|
||
return []
|
||
|
||
from datetime import timedelta
|
||
|
||
conditions = []
|
||
params = []
|
||
idx = 1
|
||
|
||
if ship_number:
|
||
conditions.append(f"ship_number = %s")
|
||
params.append(ship_number)
|
||
|
||
if device_name:
|
||
conditions.append(f"device_name = %s")
|
||
params.append(device_name)
|
||
|
||
if system_name:
|
||
conditions.append(f"system_name = %s")
|
||
params.append(system_name)
|
||
|
||
if start_date:
|
||
try:
|
||
start_dt = datetime.strptime(start_date, "%Y-%m-%d")
|
||
conditions.append(f"created_at >= %s")
|
||
params.append(start_dt)
|
||
except ValueError:
|
||
pass
|
||
|
||
if end_date:
|
||
try:
|
||
end_dt = datetime.strptime(end_date, "%Y-%m-%d") + timedelta(days=1) - timedelta(seconds=1)
|
||
conditions.append(f"created_at <= %s")
|
||
params.append(end_dt)
|
||
except ValueError:
|
||
pass
|
||
|
||
where_clause = " AND ".join(conditions) if conditions else "1=1"
|
||
|
||
try:
|
||
async with AsyncConnectionPool(
|
||
POSTGRES_CONNECTION_STRING,
|
||
kwargs={"autocommit": True}
|
||
) as pool:
|
||
async with pool.connection() as conn:
|
||
async with conn.cursor() as cur:
|
||
sql = f"""
|
||
SELECT id, ship_number, device_name, fault, spare_parts, system_name, created_at
|
||
FROM fault_records
|
||
WHERE {where_clause}
|
||
ORDER BY created_at DESC
|
||
LIMIT %s
|
||
"""
|
||
params.append(limit)
|
||
await cur.execute(sql, params)
|
||
rows = await cur.fetchall()
|
||
columns = ["id", "ship_number", "device_name", "fault", "spare_parts", "system_name", "created_at"]
|
||
results = []
|
||
for row in rows:
|
||
record = dict(zip(columns, row))
|
||
if isinstance(record.get("spare_parts"), str):
|
||
try:
|
||
record["spare_parts"] = json.loads(record["spare_parts"])
|
||
except Exception:
|
||
record["spare_parts"] = []
|
||
if isinstance(record.get("created_at"), datetime):
|
||
record["created_at"] = record["created_at"].isoformat()
|
||
results.append(record)
|
||
return results
|
||
except Exception as e:
|
||
print(f"[fault_record_db] 查询故障记录失败: {str(e)}")
|
||
import traceback
|
||
traceback.print_exc()
|
||
return []
|
||
|