chat.py 25 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644
  1. import base64
  2. import json
  3. import uuid
  4. from typing import Optional
  5. from fastapi import APIRouter, Depends, HTTPException, Query, UploadFile, File, Form
  6. from fastapi.responses import StreamingResponse
  7. from pydantic import BaseModel, Field
  8. from sqlalchemy import text
  9. from sqlalchemy.ext.asyncio import create_async_engine
  10. from app.core.security import get_current_user, RateLimiter
  11. from app.rag.retriever import MixedRetriever, classify_intent
  12. from app.rag.reranker import Reranker
  13. from app.rag.prompt import build_prompt
  14. from app.core.llm_client import llm_client
  15. from app.core.config import get_settings
  16. settings = get_settings()
  17. router = APIRouter(prefix="/chat", tags=["对话"])
  18. retriever = MixedRetriever()
  19. reranker = Reranker()
  20. class ChatRequest(BaseModel):
  21. message: str = Field(..., min_length=1, max_length=2000)
  22. conversation_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
  23. class ImageChatRequest(BaseModel):
  24. """图片对话请求:base64 图片"""
  25. image_base64: str = Field(..., min_length=1, description="Base64 编码的图片")
  26. mime_type: str = Field(default="image/jpeg", description="图片 MIME 类型")
  27. message: str = Field(default="", max_length=2000, description="可选的附加文字问题")
  28. conversation_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
  29. class MultimodalChatRequest(BaseModel):
  30. """统一多模态对话请求:支持文本 + 图片 + 视频"""
  31. message: str = Field(default="", max_length=2000, description="文字问题(可选)")
  32. # 媒体附件(图片和视频二选一或都不传,纯文本也可以)
  33. media_type: str = Field(default="", description="媒体类型: image / video / 空=纯文本")
  34. media_base64: str = Field(default="", description="Base64 编码的图片或视频")
  35. media_mime: str = Field(default="", description="媒体 MIME 类型,如 image/jpeg, video/mp4")
  36. conversation_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
  37. class ChatResponse(BaseModel):
  38. answer: str
  39. sources: list[dict]
  40. conversation_id: str
  41. intent: str
  42. class FeedbackRequest(BaseModel):
  43. conversation_id: str
  44. message_id: int
  45. feedback: str
  46. # ============================================
  47. # DB helpers
  48. # ============================================
  49. async def _ensure_user(openid: str) -> int:
  50. engine = create_async_engine(settings.database_url)
  51. try:
  52. async with engine.begin() as conn:
  53. result = await conn.execute(
  54. text("SELECT id FROM users WHERE openid = :openid"),
  55. {"openid": openid},
  56. )
  57. row = result.fetchone()
  58. if row:
  59. return row[0]
  60. result = await conn.execute(
  61. text("INSERT INTO users (openid) VALUES (:openid) RETURNING id"),
  62. {"openid": openid},
  63. )
  64. return result.fetchone()[0]
  65. finally:
  66. await engine.dispose()
  67. async def _save_message(conversation_id: str, role: str, content: str,
  68. intent: str = None, sources: list = None):
  69. engine = create_async_engine(settings.database_url)
  70. try:
  71. async with engine.begin() as conn:
  72. # 确保 conversation 存在(不要求 user_id 外键,因为 dev token 的 user 可能不在库中)
  73. await conn.execute(
  74. text("""
  75. INSERT INTO conversations (conversation_id, user_id, title)
  76. VALUES (:cid, 0, :title)
  77. ON CONFLICT (conversation_id) DO NOTHING
  78. """),
  79. {
  80. "cid": conversation_id,
  81. "title": content[:50] if role == "user" else "",
  82. },
  83. )
  84. await conn.execute(
  85. text("""
  86. INSERT INTO messages (conversation_id, role, content, intent, sources)
  87. VALUES (:cid, :role, :content, :intent, :sources)
  88. """),
  89. {
  90. "cid": conversation_id,
  91. "role": role,
  92. "content": content,
  93. "intent": intent,
  94. "sources": json.dumps(sources) if sources else None,
  95. },
  96. )
  97. finally:
  98. await engine.dispose()
  99. # ============================================
  100. # Chat endpoints
  101. # ============================================
  102. @router.post("/ask", response_model=ChatResponse)
  103. async def chat_ask(req: ChatRequest, user: dict = Depends(get_current_user)):
  104. intent = classify_intent(req.message)
  105. documents = await retriever.search(req.message, intent=intent, top_k=20)
  106. documents = reranker.rerank(req.message, documents, top_k=5)
  107. msgs = build_prompt(req.message, documents, intent=intent)
  108. answer = await llm_client.chat(msgs)
  109. sources = [
  110. {"name": d.get("drug_name", d.get("source", "")),
  111. "section": d.get("section", ""), "source": d.get("source", ""),
  112. "score": d.get("score", 0)}
  113. for d in documents
  114. ]
  115. # 保存到 DB
  116. await _save_message(req.conversation_id, "user", req.message, intent)
  117. await _save_message(req.conversation_id, "assistant", answer, intent, sources)
  118. return ChatResponse(answer=answer, sources=sources,
  119. conversation_id=req.conversation_id, intent=intent)
  120. @router.post("/stream")
  121. async def chat_stream(req: ChatRequest, user: dict = Depends(get_current_user)):
  122. async def stream_gen():
  123. intent = classify_intent(req.message)
  124. yield f"event: intent\ndata: {intent}\n\n"
  125. yield "event: status\ndata: 正在检索...\n\n"
  126. documents = await retriever.search(req.message, intent=intent, top_k=20)
  127. documents = reranker.rerank(req.message, documents, top_k=5)
  128. yield f"event: status\ndata: 已匹配 {len(documents)} 条,生成中...\n\n"
  129. msgs = build_prompt(req.message, documents, intent=intent)
  130. sources = [
  131. {"name": d.get("drug_name", d.get("source", "")),
  132. "section": d.get("section", ""), "source": d.get("source", ""),
  133. "score": d.get("score", 0)}
  134. for d in documents
  135. ]
  136. yield "event: content\n"
  137. full_answer = []
  138. async for token in llm_client.chat_stream(msgs):
  139. full_answer.append(token)
  140. yield f"data: {token}\n\n"
  141. yield "data: [DONE]\n\n"
  142. # 元数据追加
  143. import json
  144. yield f"event: meta\ndata: {json.dumps({'intent': intent, 'sources': sources, 'cid': req.conversation_id})}\n\n"
  145. answer_text = "".join(full_answer)
  146. await _save_message(req.conversation_id, "user", req.message, intent)
  147. await _save_message(req.conversation_id, "assistant", answer_text, intent, sources)
  148. return StreamingResponse(stream_gen(), media_type="text/event-stream")
  149. # ============================================
  150. # 图片对话 API(Qwen VL 分析 + OCR → RAG 检索 → 联网搜索)
  151. # ============================================
  152. @router.post("/ask-image", response_model=ChatResponse)
  153. async def chat_ask_image(req: ImageChatRequest, user: dict = Depends(get_current_user)):
  154. """
  155. 图片对话 — 非流式:
  156. 1. Qwen VL 分析图片 + OCR 提取文字
  157. 2. 用提取文字做 RAG 检索
  158. 3. 结合检索结果 + 联网搜索生成回答
  159. """
  160. # Step 1: Qwen VL 分析图片 → 提取文字
  161. ocr_text = await llm_client.analyze_image(
  162. req.image_base64, req.mime_type,
  163. prompt="请分析这张图片,提取其中所有文字信息(OCR),特别是药品名称、成分、用法用量等关键药学信息。简要输出即可。",
  164. )
  165. # Step 2: 拼接用户附加文字 + OCR 结果 → RAG 检索
  166. query = req.message.strip() if req.message else ocr_text
  167. if req.message:
  168. query = f"{req.message}\n\n(图片OCR提取内容:{ocr_text})"
  169. intent = classify_intent(query)
  170. documents = await retriever.search(query, intent=intent, top_k=20)
  171. documents = reranker.rerank(query, documents, top_k=5)
  172. # Step 3: 构建 Prompt(含图片分析结果)+ 联网搜索
  173. image_context = f"\n\n【图片分析结果】\n{ocr_text}\n"
  174. msgs = build_prompt(query, documents, intent=intent)
  175. # 在 system prompt 中追加图片分析上下文
  176. msgs[0]["content"] += image_context
  177. answer = await llm_client.chat(msgs, enable_search=True)
  178. sources = [
  179. {"name": d.get("drug_name", d.get("source", "")),
  180. "section": d.get("section", ""), "source": d.get("source", ""),
  181. "score": d.get("score", 0)}
  182. for d in documents
  183. ]
  184. await _save_message(req.conversation_id, "user",
  185. f"[图片] {req.message}" if req.message else "[图片]",
  186. intent)
  187. await _save_message(req.conversation_id, "assistant", answer, intent, sources)
  188. return ChatResponse(answer=answer, sources=sources,
  189. conversation_id=req.conversation_id, intent=intent)
  190. @router.post("/stream-image")
  191. async def chat_stream_image(req: ImageChatRequest, user: dict = Depends(get_current_user)):
  192. """图片对话 — SSE 流式"""
  193. async def stream_gen():
  194. intent = "drug_query"
  195. yield f"event: intent\ndata: {intent}\n\n"
  196. # Step 1: Qwen VL 流式分析图片 — 实时推给用户
  197. yield "event: status\ndata: 🔍 正在分析图片...\n\n"
  198. yield "event: content\ndata: 【📷 图片分析】\n\n"
  199. ocr_parts = []
  200. async for token in llm_client.analyze_image_stream(
  201. req.image_base64, req.mime_type,
  202. prompt="请分析这张图片,提取其中所有文字信息(OCR),特别是药品名称、成分、用法用量等。简要输出。",
  203. ):
  204. ocr_parts.append(token)
  205. yield f"data: {token}\n\n" # ← 实时流给用户
  206. ocr_text = "".join(ocr_parts)
  207. yield "data: \n\n"
  208. # Step 2: 拼接查询 → RAG
  209. query = req.message.strip() if req.message else ocr_text
  210. if req.message:
  211. query = f"{req.message}\n\n(图片OCR提取内容:{ocr_text})"
  212. intent = classify_intent(query)
  213. yield f"event: intent\ndata: {intent}\n\n"
  214. yield f"event: status\ndata: 📚 检索药典知识库...\n\n"
  215. documents = await retriever.search(query, intent=intent, top_k=20)
  216. documents = reranker.rerank(query, documents, top_k=5)
  217. yield f"event: status\ndata: 已匹配 {len(documents)} 条药典资料,生成回答中(已启用联网搜索)...\n\n"
  218. # Step 3: 构建 Prompt + 联网搜索流式生成
  219. image_context = f"\n\n【图片分析结果】\n{ocr_text}\n"
  220. msgs = build_prompt(query, documents, intent=intent)
  221. msgs[0]["content"] += image_context
  222. sources = [
  223. {"name": d.get("drug_name", d.get("source", "")),
  224. "section": d.get("section", ""), "source": d.get("source", ""),
  225. "score": d.get("score", 0)}
  226. for d in documents
  227. ]
  228. yield "event: content\ndata: \n【📚 药典参考回答】\n\n"
  229. full_answer = []
  230. async for token in llm_client.chat_stream(msgs, enable_search=True):
  231. full_answer.append(token)
  232. yield f"data: {token}\n\n"
  233. yield "data: [DONE]\n\n"
  234. yield f"event: meta\ndata: {json.dumps({'intent': intent, 'sources': sources, 'cid': req.conversation_id, 'ocr_text': ocr_text[:200]})}\n\n"
  235. answer_text = "".join(full_answer)
  236. await _save_message(req.conversation_id, "user",
  237. f"[图片] {req.message}" if req.message else "[图片]",
  238. intent)
  239. await _save_message(req.conversation_id, "assistant", answer_text, intent, sources)
  240. return StreamingResponse(stream_gen(), media_type="text/event-stream")
  241. @router.post("/upload-image")
  242. async def chat_upload_image(
  243. file: UploadFile = File(...),
  244. message: str = Form(default=""),
  245. conversation_id: str = Form(default=""),
  246. user: dict = Depends(get_current_user),
  247. ):
  248. """
  249. 上传图片文件 → 转为 base64 → 走图片对话流程
  250. 支持格式:jpg, jpeg, png, webp, bmp
  251. """
  252. allowed = {"image/jpeg", "image/png", "image/webp", "image/bmp"}
  253. if file.content_type and file.content_type not in allowed:
  254. raise HTTPException(400, f"不支持的图片格式: {file.content_type},支持 jpg/png/webp/bmp")
  255. contents = await file.read()
  256. if len(contents) > 10 * 1024 * 1024:
  257. raise HTTPException(400, "图片大小不能超过 10MB")
  258. image_b64 = base64.b64encode(contents).decode("utf-8")
  259. mime = file.content_type or "image/jpeg"
  260. cid = conversation_id or str(uuid.uuid4())
  261. req = ImageChatRequest(
  262. image_base64=image_b64,
  263. mime_type=mime,
  264. message=message,
  265. conversation_id=cid,
  266. )
  267. return await chat_ask_image(req, user)
  268. # ============================================================
  269. # 统一多模态对话 API(文本 + 图片 + 视频,一个接口全搞定)
  270. # ============================================================
  271. @router.post("/ask-multimodal", response_model=ChatResponse)
  272. async def chat_ask_multimodal(req: MultimodalChatRequest, user: dict = Depends(get_current_user)):
  273. """
  274. 统一多模态对话 — 非流式:
  275. 支持纯文本 / 文本+图片 / 文本+视频 / 纯图片 / 纯视频
  276. 流程:媒体分析(OCR) → RAG 检索 → 联网搜索 → 回答
  277. """
  278. conversation_id = req.conversation_id or str(uuid.uuid4())
  279. ocr_text = ""
  280. # Step 1: 如果有媒体附件,先做视觉分析 + OCR
  281. if req.media_base64 and req.media_type in ("image", "video"):
  282. media_label = "视频" if req.media_type == "video" else "图片"
  283. ocr_text = await llm_client.analyze_media(
  284. req.media_base64,
  285. req.media_type,
  286. mime_type=req.media_mime or ("" if req.media_type != "video" else "video/mp4"),
  287. )
  288. # Step 2: 拼接查询文本
  289. query = req.message.strip()
  290. if query and ocr_text:
  291. query = f"{query}\n\n({media_label}OCR提取内容:{ocr_text})"
  292. elif ocr_text:
  293. query = ocr_text
  294. elif not query:
  295. query = "请介绍一下自己"
  296. # Step 3: RAG 检索
  297. intent = classify_intent(query)
  298. documents = await retriever.search(query, intent=intent, top_k=20)
  299. documents = reranker.rerank(query, documents, top_k=5)
  300. # Step 4: 构建 Prompt + 联网搜索
  301. msgs = build_prompt(query, documents, intent=intent)
  302. if ocr_text:
  303. msgs[0]["content"] += f"\n\n【{media_label}分析结果】\n{ocr_text}\n"
  304. answer = await llm_client.chat(msgs, enable_search=bool(ocr_text) or settings.enable_web_search)
  305. sources = [
  306. {"name": d.get("drug_name", d.get("source", "")),
  307. "section": d.get("section", ""), "source": d.get("source", ""),
  308. "score": d.get("score", 0)}
  309. for d in documents
  310. ]
  311. user_msg = req.message or f"[{media_label}]" if ocr_text else req.message
  312. await _save_message(conversation_id, "user", user_msg, intent)
  313. await _save_message(conversation_id, "assistant", answer, intent, sources)
  314. return ChatResponse(answer=answer, sources=sources,
  315. conversation_id=conversation_id, intent=intent)
  316. @router.post("/stream-multimodal")
  317. async def chat_stream_multimodal(req: MultimodalChatRequest, user: dict = Depends(get_current_user)):
  318. """
  319. 统一多模态对话 — SSE 流式(全链路流式):
  320. OCR 分析 → 实时推送给用户 → 立即 RAG 检索 → 流式生成回答
  321. 用户无需等待,每一步都在实时输出
  322. """
  323. async def stream_gen():
  324. nonlocal req
  325. conversation_id = req.conversation_id or str(uuid.uuid4())
  326. media_label = ""
  327. ocr_text = ""
  328. has_media = req.media_base64 and req.media_type in ("image", "video")
  329. if has_media:
  330. media_label = "视频" if req.media_type == "video" else "图片"
  331. # Step 1: 流式 OCR 分析 — 实时推送给用户
  332. yield "event: status\ndata: 🔍 正在分析...\n\n"
  333. yield f"event: content\ndata: 【📷 {media_label}分析】\n\n"
  334. ocr_parts = []
  335. async for token in llm_client.analyze_media_stream(
  336. req.media_base64, req.media_type,
  337. mime_type=req.media_mime or "",
  338. ):
  339. ocr_parts.append(token)
  340. yield f"data: {token}\n\n" # ← OCR token 实时流给用户
  341. ocr_text = "".join(ocr_parts)
  342. yield "data: \n\n" # 分隔
  343. else:
  344. yield "event: intent\ndata: drug_query\n\n"
  345. # Step 2: 拼接查询 → RAG 检索(此时 OCR 已全部拿到)
  346. query = req.message.strip()
  347. if query and ocr_text:
  348. query = f"{query}\n\n({media_label}OCR提取内容:{ocr_text})"
  349. elif ocr_text:
  350. query = ocr_text
  351. elif not query:
  352. query = "请介绍一下自己"
  353. intent = classify_intent(query)
  354. yield f"event: intent\ndata: {intent}\n\n"
  355. yield f"event: status\ndata: 📚 检索药典知识库...\n\n"
  356. documents = await retriever.search(query, intent=intent, top_k=20)
  357. documents = reranker.rerank(query, documents, top_k=5)
  358. search_hint = "(已启用联网搜索)" if (ocr_text or settings.enable_web_search) else ""
  359. yield f"event: status\ndata: 已匹配 {len(documents)} 条,生成回答中{search_hint}...\n\n"
  360. # Step 3: LLM 流式生成
  361. msgs = build_prompt(query, documents, intent=intent)
  362. if ocr_text:
  363. msgs[0]["content"] += f"\n\n【{media_label}分析结果】\n{ocr_text}\n"
  364. sources = [
  365. {"name": d.get("drug_name", d.get("source", "")),
  366. "section": d.get("section", ""), "source": d.get("source", ""),
  367. "score": d.get("score", 0)}
  368. for d in documents
  369. ]
  370. yield "event: content\ndata: \n【📚 药典参考回答】\n\n"
  371. full_answer = []
  372. async for token in llm_client.chat_stream(msgs, enable_search=bool(ocr_text) or settings.enable_web_search):
  373. full_answer.append(token)
  374. yield f"data: {token}\n\n"
  375. yield "data: [DONE]\n\n"
  376. yield f"event: meta\ndata: {json.dumps({'intent': intent, 'sources': sources, 'cid': conversation_id, 'ocr_text': ocr_text[:200] if ocr_text else ''})}\n\n"
  377. answer_text = "".join(full_answer)
  378. user_msg = req.message or f"[{media_label}]" if ocr_text else req.message
  379. await _save_message(conversation_id, "user", user_msg, intent)
  380. await _save_message(conversation_id, "assistant", answer_text, intent, sources)
  381. return StreamingResponse(stream_gen(), media_type="text/event-stream")
  382. @router.post("/upload-media")
  383. async def chat_upload_media(
  384. file: UploadFile = File(...),
  385. message: str = Form(default=""),
  386. conversation_id: str = Form(default=""),
  387. user: dict = Depends(get_current_user),
  388. ):
  389. """
  390. 上传媒体文件(图片/视频)→ 自动识别类型 → 走多模态对话流程
  391. 支持:jpg, jpeg, png, webp, bmp, mp4, mov, avi, webm
  392. """
  393. mime = file.content_type or ""
  394. media_type = ""
  395. if mime.startswith("image/"):
  396. media_type = "image"
  397. max_size = 10 * 1024 * 1024 # 10MB
  398. elif mime.startswith("video/"):
  399. media_type = "video"
  400. max_size = 50 * 1024 * 1024 # 50MB
  401. else:
  402. raise HTTPException(400, f"不支持的媒体格式: {mime},支持 jpg/png/webp/bmp/mp4/mov/avi/webm")
  403. contents = await file.read()
  404. if len(contents) > max_size:
  405. raise HTTPException(400, f"文件大小不能超过 {max_size // 1024 // 1024}MB")
  406. media_b64 = base64.b64encode(contents).decode("utf-8")
  407. cid = conversation_id or str(uuid.uuid4())
  408. req = MultimodalChatRequest(
  409. message=message,
  410. media_type=media_type,
  411. media_base64=media_b64,
  412. media_mime=mime,
  413. conversation_id=cid,
  414. )
  415. return await chat_ask_multimodal(req, user)
  416. # ============================================
  417. # 对话历史 API
  418. # ============================================
  419. @router.get("/history")
  420. async def get_history(
  421. page: int = Query(1, ge=1),
  422. page_size: int = Query(20, ge=1, le=50),
  423. user: dict = Depends(get_current_user),
  424. ):
  425. engine = create_async_engine(settings.database_url)
  426. try:
  427. async with engine.connect() as conn:
  428. offset = (page - 1) * page_size
  429. result = await conn.execute(
  430. text("""
  431. SELECT c.conversation_id, c.title, c.created_at,
  432. COUNT(m.id) as msg_count
  433. FROM conversations c
  434. LEFT JOIN messages m ON m.conversation_id = c.conversation_id
  435. GROUP BY c.id
  436. ORDER BY c.created_at DESC
  437. LIMIT :limit OFFSET :offset
  438. """),
  439. {"limit": page_size, "offset": offset},
  440. )
  441. items = []
  442. for row in result.fetchall():
  443. items.append({
  444. "conversation_id": row[0],
  445. "title": row[1] or "新的对话",
  446. "created_at": row[2].isoformat() if row[2] else "",
  447. "message_count": row[3],
  448. })
  449. return {"items": items, "page": page, "page_size": page_size}
  450. finally:
  451. await engine.dispose()
  452. @router.get("/history/{conversation_id}")
  453. async def get_conversation_detail(
  454. conversation_id: str,
  455. user: dict = Depends(get_current_user),
  456. ):
  457. engine = create_async_engine(settings.database_url)
  458. try:
  459. async with engine.connect() as conn:
  460. result = await conn.execute(
  461. text("""
  462. SELECT id, role, content, intent, sources, created_at
  463. FROM messages
  464. WHERE conversation_id = :cid
  465. ORDER BY created_at ASC
  466. """),
  467. {"cid": conversation_id},
  468. )
  469. messages = []
  470. for row in result.fetchall():
  471. messages.append({
  472. "id": row[0],
  473. "role": row[1],
  474. "content": row[2],
  475. "intent": row[3],
  476. "sources": row[4] if row[4] else [],
  477. "created_at": row[5].isoformat() if row[5] else "",
  478. })
  479. return {"conversation_id": conversation_id, "messages": messages}
  480. finally:
  481. await engine.dispose()
  482. @router.post("/feedback")
  483. async def submit_feedback(req: FeedbackRequest, user: dict = Depends(get_current_user)):
  484. engine = create_async_engine(settings.database_url)
  485. try:
  486. async with engine.begin() as conn:
  487. await conn.execute(
  488. text("UPDATE messages SET feedback = :fb WHERE id = :mid"),
  489. {"fb": req.feedback, "mid": req.message_id},
  490. )
  491. return {"ok": True, "message_id": req.message_id, "feedback": req.feedback}
  492. finally:
  493. await engine.dispose()
  494. # ============================================
  495. # 管理员:查看全部对话
  496. # ============================================
  497. @router.get("/admin/conversations")
  498. async def admin_list_conversations(
  499. page: int = Query(1, ge=1),
  500. page_size: int = Query(20, ge=1, le=100),
  501. keyword: Optional[str] = Query(None, description="搜索用户提问关键词"),
  502. user: dict = Depends(get_current_user),
  503. ):
  504. engine = create_async_engine(settings.database_url)
  505. try:
  506. async with engine.connect() as conn:
  507. offset = (page - 1) * page_size
  508. where = ""
  509. params = {"limit": page_size, "offset": offset}
  510. if keyword:
  511. where = "WHERE m.content LIKE :kw"
  512. params["kw"] = f"%{keyword}%"
  513. result = await conn.execute(
  514. text(f"""
  515. SELECT DISTINCT ON (c.conversation_id)
  516. c.conversation_id, c.title, c.created_at,
  517. m.content as last_msg, m.role
  518. FROM conversations c
  519. JOIN messages m ON m.conversation_id = c.conversation_id
  520. {where}
  521. ORDER BY c.conversation_id, m.created_at DESC
  522. LIMIT :limit OFFSET :offset
  523. """),
  524. params,
  525. )
  526. items = []
  527. for row in result.fetchall():
  528. items.append({
  529. "conversation_id": row[0],
  530. "title": row[1] or "新的对话",
  531. "created_at": row[2].isoformat() if row[2] else "",
  532. "last_message": (row[3] or "")[:200],
  533. "last_role": row[4],
  534. })
  535. return {"items": items, "page": page, "page_size": page_size}
  536. finally:
  537. await engine.dispose()