766 lines
34 KiB
Python
766 lines
34 KiB
Python
"""对话服务:调用 BaoDan Chat API,处理 SSE 流式响应。"""
|
||
import json
|
||
import requests
|
||
from flask import current_app
|
||
|
||
|
||
class ChatService:
|
||
"""智能问答业务逻辑。"""
|
||
|
||
# 缓存模型名称,避免重复查询数据库
|
||
_model_name_cache: dict[str, str] = {}
|
||
|
||
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 _parse_model_name(self, provider: str = "", model_id: str = "", model_json=None) -> str:
|
||
"""Extract provider/model from Dify app model config fields."""
|
||
import json
|
||
|
||
if provider and model_id:
|
||
return f"{provider}/{model_id}"
|
||
|
||
if not model_json:
|
||
return ""
|
||
|
||
try:
|
||
model_data = json.loads(model_json) if isinstance(model_json, str) and model_json.startswith("{") else model_json
|
||
except (json.JSONDecodeError, TypeError):
|
||
return ""
|
||
|
||
if not isinstance(model_data, dict):
|
||
return ""
|
||
|
||
parsed_provider = model_data.get("provider", "")
|
||
parsed_model = model_data.get("name", "") or model_data.get("model", "") or model_data.get("model_id", "")
|
||
if parsed_provider and parsed_model and not isinstance(parsed_model, dict):
|
||
return f"{parsed_provider}/{parsed_model}"
|
||
|
||
nested = model_data.get("model", {})
|
||
if isinstance(nested, dict):
|
||
nested_provider = nested.get("provider", "") or parsed_provider
|
||
nested_model = nested.get("name", "") or nested.get("model", "") or nested.get("model_id", "")
|
||
if nested_provider and nested_model:
|
||
return f"{nested_provider}/{nested_model}"
|
||
|
||
return ""
|
||
|
||
def _get_model_name(self, api_token: str) -> str:
|
||
"""从 BaoDan 数据库获取应用使用的模型名称。
|
||
|
||
BaoDan 的 SSE 响应不包含模型信息,需要从 app_model_configs 表获取。
|
||
结果会缓存,避免每次请求都查数据库。
|
||
"""
|
||
import logging
|
||
token = (api_token or "").replace("Bearer ", "").strip()
|
||
cache_key = token or "__default__"
|
||
if cache_key in self._model_name_cache:
|
||
return self._model_name_cache[cache_key]
|
||
|
||
try:
|
||
from insurance.db.compat import db
|
||
from sqlalchemy import text
|
||
|
||
if token:
|
||
result = db.session.execute(text("""
|
||
SELECT am.id, am.provider, am.model_id, am.model
|
||
FROM api_tokens t
|
||
JOIN apps a ON a.id = t.app_id
|
||
JOIN app_model_configs am
|
||
ON am.id = a.app_model_config_id OR am.app_id = a.id
|
||
WHERE t.token = :token AND t.type = 'app'
|
||
ORDER BY
|
||
CASE WHEN am.id = a.app_model_config_id THEN 0 ELSE 1 END,
|
||
am.updated_at DESC
|
||
"""), {"token": token})
|
||
else:
|
||
result = db.session.execute(text("""
|
||
SELECT id, provider, model_id, model FROM app_model_configs
|
||
ORDER BY updated_at DESC
|
||
"""))
|
||
rows = result.fetchall()
|
||
if token and not rows:
|
||
result = db.session.execute(text("""
|
||
SELECT id, provider, model_id, model FROM app_model_configs
|
||
ORDER BY updated_at DESC
|
||
"""))
|
||
rows = result.fetchall()
|
||
logging.info(f"[_get_model_name] 查到 {len(rows)} 条 app_model_configs 记录")
|
||
|
||
for row in rows:
|
||
row_id, provider, model_id, model_json = row[0], row[1], row[2], row[3]
|
||
logging.info(f"[_get_model_name] 记录 {row_id}: provider={provider}, model_id={model_id}, model={str(model_json)[:200] if model_json else None}")
|
||
model_name = self._parse_model_name(provider, model_id, model_json)
|
||
if model_name:
|
||
self._model_name_cache[cache_key] = model_name
|
||
logging.info(f"[_get_model_name] 获取模型名称: {model_name}")
|
||
return model_name
|
||
|
||
except Exception as e:
|
||
logging.warning(f"[_get_model_name] 查询失败: {e}")
|
||
logging.info(f"[_get_model_name] 未找到模型信息")
|
||
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 = "", model_id: str = "",
|
||
message_tokens: int = 0, answer_tokens: int = 0):
|
||
"""保存对话记录到本地数据库。"""
|
||
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,
|
||
message_id=message_id,
|
||
model_id=model_id,
|
||
message_tokens=message_tokens,
|
||
answer_tokens=answer_tokens,
|
||
)
|
||
db.session.add(record)
|
||
try:
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
raise
|
||
logging.info(f"[SAVE] 消息保存成功: role={role}, session_id={session_id[:8]}..., record_id={record.id}, model_id={model_id}, message_tokens={message_tokens}, answer_tokens={answer_tokens}")
|
||
|
||
# 用户发消息时更新最后活跃时间(用于管理后台统计活跃用户)
|
||
if role == "user":
|
||
try:
|
||
from datetime import datetime
|
||
from insurance.models.wecom_user import WeComUserMapping
|
||
# user_id 是数据库自增 ID(JWT 中的 user_id = mapping.id)
|
||
# 访客模式 user_id 格式为 "guest_xxx",无需更新
|
||
if user_id.isdigit():
|
||
WeComUserMapping.query.filter_by(id=int(user_id)).update(
|
||
{"last_active_at": datetime.now()}
|
||
)
|
||
db.session.commit()
|
||
except Exception:
|
||
db.session.rollback()
|
||
|
||
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
|
||
import time
|
||
logging.info(f"调用 BaoDan API: {base_url}/v1/chat-messages, conversation_id={baodan_conversation_id}")
|
||
|
||
# 重试机制:最多重试2次(共3次尝试),仅对连接错误和5xx错误重试
|
||
max_retries = 2
|
||
resp = None
|
||
for attempt in range(max_retries + 1):
|
||
try:
|
||
resp = requests.post(
|
||
f"{base_url}/v1/chat-messages",
|
||
json=payload,
|
||
headers=headers,
|
||
stream=True,
|
||
timeout=120,
|
||
)
|
||
# 5xx 错误且还有重试次数时重试
|
||
if resp.status_code >= 500 and attempt < max_retries:
|
||
logging.warning(f"BaoDan API 返回 {resp.status_code},第 {attempt + 1} 次重试...")
|
||
time.sleep(1 * (attempt + 1)) # 递增延迟
|
||
resp.close()
|
||
continue
|
||
break # 非5xx错误或已用完重试次数,跳出循环
|
||
except (requests.ConnectionError, requests.Timeout) as e:
|
||
if attempt < max_retries:
|
||
logging.warning(f"BaoDan API 连接失败: {e},第 {attempt + 1} 次重试...")
|
||
time.sleep(1 * (attempt + 1))
|
||
continue
|
||
raise # 最后一次尝试仍失败,抛出异常
|
||
|
||
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", "")
|
||
# 提取 token 使用信息
|
||
usage = metadata.get("usage", {})
|
||
# BaoDan/Dify 的 SSE 响应不包含模型信息
|
||
# 优先从 metadata 获取,后备从数据库获取
|
||
model_id = metadata.get("ls_model_name", "") or event_data.get("model", "")
|
||
if not model_id:
|
||
model_id = self._get_model_name(api_key)
|
||
message_tokens = usage.get("prompt_tokens", 0)
|
||
answer_tokens = usage.get("completion_tokens", 0)
|
||
# 记录 token 使用情况(用于调试)
|
||
import logging
|
||
logging.info(f"[TOKEN] usage={usage}, model={model_id}, prompt_tokens={message_tokens}, completion_tokens={answer_tokens}")
|
||
# 优先使用 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,
|
||
model_id=model_id, message_tokens=message_tokens)
|
||
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,
|
||
model_id=model_id, answer_tokens=answer_tokens)
|
||
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
|
||
break # 回答已完成,关闭与 BaoDan 的连接,让前端 reader.read() 收到 EOF
|
||
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,
|
||
"correction": assistant.correction 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,
|
||
"correction": record.correction,
|
||
"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, correction: 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
|
||
if correction:
|
||
record.correction = correction
|
||
db.session.commit()
|
||
except Exception as e:
|
||
logging.warning(f"保存本地反馈失败: {e}")
|
||
try:
|
||
from insurance.db.compat import db
|
||
db.session.rollback()
|
||
except Exception:
|
||
pass
|
||
|
||
return {"code": 0, "message": "success", "data": None}
|
||
|
||
def get_suggestions(self, message: str) -> dict:
|
||
"""基于回答生成推荐追问(调用 LLM)。"""
|
||
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")
|
||
|
||
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,
|
||
)
|
||
if resp.status_code != 200:
|
||
logging.warning(f"推荐追问 API 返回 {resp.status_code}")
|
||
return {"code": 0, "data": {"suggestions": []}}
|
||
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]}}
|
||
except Exception as e:
|
||
logging.warning(f"获取推荐追问失败: {e}")
|
||
return {"code": 0, "data": {"suggestions": []}}
|
||
|
||
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)
|
||
|
||
def get_app_opening(self, api_token: str) -> dict:
|
||
"""获取 Dify 应用的开场配置(开场白 + 推荐问题)。
|
||
|
||
调用 Dify 的 GET /v1/parameters 接口,返回 opening_statement 和
|
||
suggested_questions,供前端在新会话时展示。
|
||
"""
|
||
import logging
|
||
|
||
from flask import current_app
|
||
base_url = current_app.config.get("BAODAN_API_URL", "http://localhost:5001")
|
||
|
||
if not api_token:
|
||
logging.warning("[APP_OPENING] api_token 为空,跳过获取开场配置")
|
||
return {"code": 0, "data": {"opening_statement": "", "suggested_questions": []}}
|
||
|
||
try:
|
||
url = f"{base_url}/v1/parameters"
|
||
logging.info(f"[APP_OPENING] 请求 Dify: {url}")
|
||
resp = requests.get(
|
||
url,
|
||
headers={"Authorization": f"Bearer {api_token}"},
|
||
timeout=10,
|
||
)
|
||
logging.info(f"[APP_OPENING] Dify 响应: status={resp.status_code}")
|
||
|
||
if resp.status_code != 200:
|
||
logging.warning(f"[APP_OPENING] Dify parameters 接口返回 {resp.status_code}: {resp.text[:200]}")
|
||
return {"code": 0, "data": {"opening_statement": "", "suggested_questions": []}}
|
||
|
||
data = resp.json()
|
||
logging.info(f"[APP_OPENING] Dify 原始响应 keys: {list(data.keys())}")
|
||
opening = data.get("opening_statement", "")
|
||
questions = data.get("suggested_questions", [])
|
||
logging.info(f"[APP_OPENING] opening_statement={opening[:50] if opening else '(空)'}, suggested_questions={questions}, count={len(questions)}")
|
||
return {
|
||
"code": 0,
|
||
"data": {
|
||
"opening_statement": opening,
|
||
"suggested_questions": questions,
|
||
},
|
||
}
|
||
except Exception as e:
|
||
logging.exception(f"[APP_OPENING] 获取应用开场配置失败: {e}")
|
||
return {"code": 0, "data": {"opening_statement": "", "suggested_questions": []}}
|