579 lines
25 KiB
Python
579 lines
25 KiB
Python
"""对话服务:调用 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)
|