baodan/api/insurance/chat/service.py

579 lines
25 KiB
Python
Raw 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.

"""对话服务:调用 BaoDan Chat API处理 SSE 流式响应。"""
import json
import requests
from flask import current_app
class ChatService:
"""智能问答业务逻辑。"""
def _get_baodan_conversation_id(self, user_id: str, session_id: str) -> str:
"""获取调用 BaoDan 时应使用的 conversation_id。
BaoDan 只接受它自己返回过的 conversation_id。
本地新建的 session_id 对 BaoDan 来说不存在,必须传空让它创建新对话。
"""
if not session_id:
return ""
try:
from insurance.models.chat_session import ChatSession
from insurance.models.chat_record import ChatRecord
# 检查本地会话是否存在
session = ChatSession.query.filter_by(
user_id=user_id,
session_id=session_id,
is_deleted=False,
).first()
if not session:
# 本地无此会话记录,传空让 BaoDan 创建新对话
return ""
# 检查该会话是否有聊天记录(说明之前已经和 BaoDan 建立过对话)
has_records = ChatRecord.query.filter_by(
user_id=user_id,
session_id=session_id,
).first() is not None
if not has_records:
# 新会话,没有记录,传空让 BaoDan 创建新对话
return ""
# 有记录,说明之前已经和 BaoDan 对话过
# 优先使用存储的 baodan_conversation_id
if hasattr(session, 'baodan_conversation_id') and session.baodan_conversation_id:
return session.baodan_conversation_id
# 如果没有存储 baodan_conversation_id使用本地 session_id
# (因为 session_id 可能就是 BaoDan 之前返回的 conversation_id
return session_id
except Exception:
return ""
def _finalize_session(self, user_id: str, requested_session_id: str, final_session_id: str, first_message: str) -> str:
"""确保会话落库,并把本地占位会话 ID 绑定为 BaoDan 返回的真实 ID。"""
import logging
persisted_session_id = final_session_id or requested_session_id
if not persisted_session_id:
return ""
session_name = first_message[:20].replace("\n", " ")
if len(first_message) > 20:
session_name += "..."
session_name = session_name or "新会话"
try:
from insurance.db.compat import db
from insurance.models.chat_session import ChatSession
from insurance.models.chat_record import ChatRecord
if requested_session_id:
local_session = ChatSession.query.filter_by(
user_id=user_id,
session_id=requested_session_id,
is_deleted=False,
).first()
else:
local_session = None
final_session = ChatSession.query.filter_by(
user_id=user_id,
session_id=persisted_session_id,
is_deleted=False,
).first()
if local_session and requested_session_id != persisted_session_id:
if final_session:
# 目标会话已存在,合并记录
ChatRecord.query.filter_by(
user_id=user_id,
session_id=requested_session_id,
).update({"session_id": persisted_session_id})
db.session.delete(local_session)
else:
# 目标会话不存在,更新本地会话的 session_id并同步更新聊天记录
local_session.session_id = persisted_session_id
ChatRecord.query.filter_by(
user_id=user_id,
session_id=requested_session_id,
).update({"session_id": persisted_session_id})
if not local_session.name or local_session.name == "新会话":
local_session.name = session_name
elif final_session:
if not final_session.name or final_session.name == "新会话":
final_session.name = session_name
elif not local_session:
db.session.add(ChatSession(
session_id=persisted_session_id,
user_id=user_id,
name=session_name,
))
elif not local_session.name or local_session.name == "新会话":
local_session.name = session_name
# 保存 BaoDan 返回的 conversation_id用于后续对话保持上下文
# 只要 BaoDan 返回了 conversation_id就保存到本地会话
if final_session_id and final_session_id != requested_session_id:
# BaoDan 返回了不同于本地 session_id 的 conversation_id需要保存
target_session = local_session or final_session
if not target_session:
# 可能是新建的会话,重新查询
target_session = ChatSession.query.filter_by(
user_id=user_id,
session_id=persisted_session_id,
is_deleted=False,
).first()
if target_session:
target_session.baodan_conversation_id = final_session_id
logging.info(f"保存会话的 BaoDan conversation_id: session_id={persisted_session_id}, baodan_conversation_id={final_session_id}")
db.session.commit()
return persisted_session_id
except Exception as e:
logging.exception(f"保存会话失败: {e}")
try:
from insurance.db.compat import db
db.session.rollback()
except Exception:
pass
return persisted_session_id
def _save_chat_record(self, user_id: str, session_id: str, role: str, content: str, message_id: str = ""):
"""保存对话记录到本地数据库。"""
import logging
logging.info(f"[SAVE] 开始保存消息: role={role}, session_id={session_id}, user_id={user_id}")
from insurance.db.compat import db
from insurance.models.chat_record import ChatRecord
record = ChatRecord(
user_id=user_id,
session_id=session_id,
role=role,
content=content[:100] + "..." if len(content) > 100 else content,
message_id=message_id,
)
db.session.add(record)
db.session.commit()
logging.info(f"[SAVE] 消息保存成功: role={role}, session_id={session_id[:8]}..., record_id={record.id}")
def send_message_stream(self, user_id: str, message: str, session_id: str, filters: dict,
api_token: str = "", base_url: str = ""):
"""调用 BaoDan Chat API以 SSE 流式返回结果。"""
api_key = api_token or ""
base_url = base_url or "http://localhost:5001"
# 注意:用户消息的保存延迟到 message_end 事件中,确保使用正确的 conversation_id
baodan_conversation_id = self._get_baodan_conversation_id(user_id, session_id)
payload = {
"inputs": {},
"query": message,
"response_mode": "streaming",
"user": f"user_{user_id}",
"conversation_id": baodan_conversation_id,
"files": [],
}
# 如果有筛选条件,通过 inputs 传递给 BaoDan Workflow
if filters:
payload["inputs"] = filters
headers = {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
}
# 收集完整的助手回复
full_answer = ""
final_message_id = ""
final_session_id = session_id
done_yielded = False
try:
import logging
logging.info(f"调用 BaoDan API: {base_url}/v1/chat-messages, conversation_id={baodan_conversation_id}")
resp = requests.post(
f"{base_url}/v1/chat-messages",
json=payload,
headers=headers,
stream=True,
timeout=120,
)
logging.info(f"BaoDan API 响应状态码: {resp.status_code}")
# 检查响应状态码
if resp.status_code != 200:
error_text = resp.text[:500] if resp.text else "无响应内容"
logging.error(f"BaoDan API 返回错误: {resp.status_code} - {error_text}")
yield json.dumps({
"type": "error",
"data": {"message": f"AI 服务返回错误 (HTTP {resp.status_code})"},
})
return
for line in resp.iter_lines():
if not line:
continue
decoded = line.decode("utf-8")
if decoded.startswith("data: "):
event_data = json.loads(decoded[6:])
event_type = event_data.get("event", "")
if event_type == "message":
# 文本增量
answer_chunk = event_data.get("answer", "")
full_answer += answer_chunk
yield json.dumps({
"type": "delta",
"data": answer_chunk,
})
elif event_type == "message_end":
# 回答完成
metadata = event_data.get("metadata", {})
retriever = metadata.get("retriever_resources", [])
final_message_id = event_data.get("message_id", "")
# 优先使用 BaoDan 返回的 conversation_id如果没有则使用原始 session_id
returned_conversation_id = event_data.get("conversation_id", "")
final_session_id = returned_conversation_id if returned_conversation_id else session_id
final_session_id = self._finalize_session(user_id, session_id, final_session_id, message)
# 调试日志:检查 BaoDan 返回的 conversation_id
import logging
logging.info(f"BaoDan message_end: conversation_id={returned_conversation_id}, session_id={session_id}, final_session_id={final_session_id}")
# 保存用户消息到本地数据库(延迟保存,确保使用正确的 conversation_id
logging.info(f"[SAVE] 准备保存用户消息: session_id={final_session_id}, message={message[:50]}...")
self._save_chat_record(user_id, final_session_id, "user", message)
logging.info(f"[SAVE] 用户消息保存完成")
# 保存助手回复到本地数据库
if full_answer:
logging.info(f"[SAVE] 准备保存助手回复: session_id={final_session_id}, answer={full_answer[:50]}...")
self._save_chat_record(user_id, final_session_id, "assistant", full_answer, final_message_id)
logging.info(f"[SAVE] 助手回复保存完成")
# 发送来源引用
for ref in retriever:
yield json.dumps({
"type": "source",
"data": {
"doc_name": ref.get("document_name", ""),
"chunk": ref.get("content", ""),
"score": ref.get("score", 0),
},
})
# 发送完成事件
yield json.dumps({
"type": "done",
"data": {
"message_id": final_message_id,
"conversation_id": final_session_id,
},
})
done_yielded = True
elif event_type == "error":
yield json.dumps({
"type": "error",
"data": {"message": event_data.get("message", "未知错误")},
})
except requests.Timeout:
yield json.dumps({"type": "error", "data": {"message": "BaoDan API 超时"}})
except Exception as e:
import logging
logging.exception(f"对话流式响应异常: {e}")
yield json.dumps({"type": "error", "data": {"message": str(e)}})
finally:
# 确保前端总能收到 done 事件,防止无限加载
if not done_yielded:
import logging
logging.warning("流式响应结束但未收到 done 事件,发送兜底 done")
yield json.dumps({
"type": "done",
"data": {
"message_id": final_message_id,
"conversation_id": final_session_id or session_id,
},
})
def get_sessions(self, user_id: str, page: int, page_size: int) -> dict:
"""获取用户会话列表(从本地数据库)。"""
import logging
try:
from insurance.models.chat_session import ChatSession
query = ChatSession.query.filter_by(
user_id=user_id,
is_deleted=False,
).order_by(ChatSession.created_at.desc())
total = query.count()
sessions = query.offset((page - 1) * page_size).limit(page_size).all()
items = [s.to_dict() for s in sessions]
logging.info(f"获取会话列表成功: user_id={user_id}, count={len(items)}")
return {
"code": 0,
"message": "success",
"data": {
"data": items,
"items": items,
"total": total,
"page": page,
"limit": page_size,
"has_more": (page * page_size) < total,
},
}
except Exception as e:
logging.exception(f"获取会话列表失败: {e}")
return {"code": 5001, "message": f"获取会话列表失败: {str(e)}", "data": {"data": [], "items": []}}
def create_session(self, user_id: str, session_id: str = "", name: str = "新会话") -> dict:
"""创建新会话(保存到本地数据库)。"""
import logging
import uuid
try:
from insurance.db.compat import db
from insurance.models.chat_session import ChatSession
if not session_id:
session_id = str(uuid.uuid4())
existing = ChatSession.query.filter_by(
user_id=user_id,
session_id=session_id,
is_deleted=False,
).first()
if existing:
return {"code": 0, "message": "success", "data": {"session_id": session_id}}
def _persist_session(session_obj: ChatSession) -> None:
db.session.add(session_obj)
db.session.commit()
session = ChatSession(
session_id=session_id,
user_id=user_id,
name=name,
)
try:
_persist_session(session)
logging.info(f"会话已创建: session_id={session_id}, user_id={user_id}")
except Exception as db_error:
db.session.rollback()
logging.warning(f"当前上下文创建会话失败,尝试 app context: {db_error}")
try:
from insurance.app import app as flask_app
with flask_app.app_context():
from insurance.db.compat import db as ctx_db
from insurance.models.chat_session import ChatSession as CtxChatSession
ctx_session = CtxChatSession(
session_id=session_id,
user_id=user_id,
name=name,
)
ctx_db.session.add(ctx_session)
ctx_db.session.commit()
logging.info(f"会话已创建with app context: session_id={session_id}, user_id={user_id}")
except Exception as ctx_error:
logging.exception(f"创建会话失败: {ctx_error}")
return {"code": 5001, "message": f"创建会话失败: {str(ctx_error)}", "data": None}
return {"code": 0, "message": "success", "data": {"session_id": session_id}}
except Exception as e:
logging.exception(f"创建会话失败: {e}")
return {"code": 5001, "message": f"创建会话失败: {str(e)}", "data": None}
def delete_session(self, user_id: str, session_id: str) -> dict:
"""软删除会话(本地数据库)。"""
try:
from insurance.db.compat import db
from insurance.models.chat_session import ChatSession
session = ChatSession.query.filter_by(
session_id=session_id,
user_id=user_id,
).first()
if session:
session.is_deleted = True
db.session.commit()
import logging
logging.info(f"会话已删除: session_id={session_id}")
return {"code": 0, "message": "success", "data": None}
else:
return {"code": 404, "message": "会话不存在", "data": None}
except Exception as e:
import logging
logging.error(f"删除会话失败: {e}")
return {"code": 5001, "message": f"删除会话失败: {str(e)}", "data": None}
def get_messages(self, user_id: str, session_id: str) -> dict:
"""获取会话消息记录(从本地数据库)。"""
import logging
logging.info(f"[QUERY] 查询消息: user_id={user_id}, session_id={session_id}")
try:
from insurance.models.chat_record import ChatRecord
records = ChatRecord.query.filter_by(
user_id=user_id,
session_id=session_id,
).order_by(ChatRecord.created_at.asc(), ChatRecord.id.asc()).all()
logging.info(f"[QUERY] 查询到 {len(records)} 条记录")
# 调试:列出所有该用户的消息
all_records = ChatRecord.query.filter_by(user_id=user_id).order_by(ChatRecord.id.desc()).limit(5).all()
for r in all_records:
logging.info(f"[QUERY] 最近消息: id={r.id}, session_id={r.session_id}, role={r.role}")
messages = [record.to_dict() for record in records]
pairs = []
i = 0
while i < len(records):
record = records[i]
if record.role == "user":
assistant = records[i + 1] if i + 1 < len(records) and records[i + 1].role == "assistant" else None
pairs.append({
"id": assistant.message_id if assistant else record.message_id or str(record.id),
"query": record.content,
"answer": assistant.content if assistant else "",
"feedback": assistant.rating if assistant else None,
"sources": [],
"created_at": str(record.created_at) if record.created_at else None,
})
i += 2 if assistant else 1
else:
pairs.append({
"id": record.message_id or str(record.id),
"query": "",
"answer": record.content,
"feedback": record.rating,
"sources": [],
"created_at": str(record.created_at) if record.created_at else None,
})
i += 1
return {
"code": 0,
"message": "success",
"data": {
"messages": pairs,
"records": messages,
"data": pairs,
},
}
except Exception as e:
import logging
logging.exception(f"获取本地聊天记录失败: {e}")
return {"code": 5001, "message": f"获取聊天记录失败: {str(e)}", "data": {"messages": [], "records": []}}
def submit_feedback(self, user_id: str, message_id: str, rating: str, comment: str) -> dict:
"""提交反馈并保存到本地数据库。"""
import logging
try:
from insurance.db.compat import db
from insurance.models.chat_record import ChatRecord
record = ChatRecord.query.filter_by(
user_id=user_id,
message_id=message_id,
role="assistant",
).first()
if record:
record.rating = rating
db.session.commit()
except Exception as e:
logging.warning(f"保存本地反馈失败: {e}")
try:
from insurance.db.compat import db
db.session.rollback()
except Exception:
pass
try:
from flask import current_app
api_key = current_app.config.get("BAODAN_CHAT_API_KEY", "")
base_url = current_app.config.get("BAODAN_API_URL", "http://localhost:5001")
if api_key and message_id:
requests.post(
f"{base_url}/v1/messages/{message_id}/feedbacks",
json={
"rating": rating,
"user": f"user_{user_id}",
"content": comment,
},
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
timeout=10,
)
except Exception as e:
logging.warning(f"转发反馈到 BaoDan 失败: {e}")
return {"code": 0, "message": "success", "data": None}
def get_suggestions(self, message: str) -> dict:
"""基于回答生成推荐追问(调用 LLM"""
from flask import current_app
api_key = current_app.config.get("BAODAN_CHAT_API_KEY", "")
base_url = current_app.config.get("BAODAN_API_URL", "http://localhost:5001")
prompt = f"基于以下回答,生成 3 个用户可能会追问的相关问题,每行一个:\n\n{message}"
resp = requests.post(
f"{base_url}/v1/chat-messages",
json={
"inputs": {},
"query": prompt,
"response_mode": "blocking",
"user": "system",
},
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
timeout=30,
)
data = resp.json()
answer = data.get("answer", "")
suggestions = [line.strip().lstrip("0123456789.、") for line in answer.split("\n") if line.strip()]
return {"code": 0, "data": {"suggestions": suggestions[:3]}}
def rename_session(self, user_id: str, session_id: str, name: str) -> dict:
"""重命名会话(本地数据库)。"""
try:
from insurance.db.compat import db
from insurance.models.chat_session import ChatSession
session = ChatSession.query.filter_by(
session_id=session_id,
user_id=user_id,
).first()
if session:
session.name = name
db.session.commit()
import logging
logging.info(f"会话已重命名: session_id={session_id}, name={name}")
return {"code": 0, "message": "success", "data": None}
else:
return {"code": 404, "message": "会话不存在", "data": None}
except Exception as e:
import logging
logging.error(f"重命名会话失败: {e}")
return {"code": 5001, "message": f"重命名失败: {str(e)}", "data": None}
def auto_name_session(self, user_id: str, session_id: str, first_message: str) -> dict:
"""根据首条消息自动命名会话(提取关键词)。"""
# 简单提取取前20个字符作为标题
name = first_message[:20].replace("\n", " ")
if len(first_message) > 20:
name += "..."
return self.rename_session(user_id, session_id, name)