chat.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297
  1. import json
  2. import uuid
  3. from typing import Optional
  4. from fastapi import APIRouter, Depends, HTTPException, Query
  5. from fastapi.responses import StreamingResponse
  6. from pydantic import BaseModel, Field
  7. from sqlalchemy import text
  8. from sqlalchemy.ext.asyncio import create_async_engine
  9. from app.core.security import get_current_user, RateLimiter
  10. from app.rag.retriever import MixedRetriever, classify_intent
  11. from app.rag.reranker import Reranker
  12. from app.rag.prompt import build_prompt
  13. from app.core.llm_client import llm_client
  14. from app.core.config import get_settings
  15. settings = get_settings()
  16. router = APIRouter(prefix="/chat", tags=["对话"])
  17. retriever = MixedRetriever()
  18. reranker = Reranker()
  19. class ChatRequest(BaseModel):
  20. message: str = Field(..., min_length=1, max_length=2000)
  21. conversation_id: str = Field(default_factory=lambda: str(uuid.uuid4()))
  22. class ChatResponse(BaseModel):
  23. answer: str
  24. sources: list[dict]
  25. conversation_id: str
  26. intent: str
  27. class FeedbackRequest(BaseModel):
  28. conversation_id: str
  29. message_id: int
  30. feedback: str
  31. # ============================================
  32. # DB helpers
  33. # ============================================
  34. async def _ensure_user(openid: str) -> int:
  35. engine = create_async_engine(settings.database_url)
  36. try:
  37. async with engine.begin() as conn:
  38. result = await conn.execute(
  39. text("SELECT id FROM users WHERE openid = :openid"),
  40. {"openid": openid},
  41. )
  42. row = result.fetchone()
  43. if row:
  44. return row[0]
  45. result = await conn.execute(
  46. text("INSERT INTO users (openid) VALUES (:openid) RETURNING id"),
  47. {"openid": openid},
  48. )
  49. return result.fetchone()[0]
  50. finally:
  51. await engine.dispose()
  52. async def _save_message(conversation_id: str, role: str, content: str,
  53. intent: str = None, sources: list = None):
  54. engine = create_async_engine(settings.database_url)
  55. try:
  56. async with engine.begin() as conn:
  57. # 确保 conversation 存在(不要求 user_id 外键,因为 dev token 的 user 可能不在库中)
  58. await conn.execute(
  59. text("""
  60. INSERT INTO conversations (conversation_id, user_id, title)
  61. VALUES (:cid, 0, :title)
  62. ON CONFLICT (conversation_id) DO NOTHING
  63. """),
  64. {
  65. "cid": conversation_id,
  66. "title": content[:50] if role == "user" else "",
  67. },
  68. )
  69. await conn.execute(
  70. text("""
  71. INSERT INTO messages (conversation_id, role, content, intent, sources)
  72. VALUES (:cid, :role, :content, :intent, :sources)
  73. """),
  74. {
  75. "cid": conversation_id,
  76. "role": role,
  77. "content": content,
  78. "intent": intent,
  79. "sources": json.dumps(sources) if sources else None,
  80. },
  81. )
  82. finally:
  83. await engine.dispose()
  84. # ============================================
  85. # Chat endpoints
  86. # ============================================
  87. @router.post("/ask", response_model=ChatResponse)
  88. async def chat_ask(req: ChatRequest, user: dict = Depends(get_current_user)):
  89. intent = classify_intent(req.message)
  90. documents = await retriever.search(req.message, intent=intent, top_k=20)
  91. documents = reranker.rerank(req.message, documents, top_k=5)
  92. msgs = build_prompt(req.message, documents, intent=intent)
  93. answer = await llm_client.chat(msgs)
  94. sources = [
  95. {"name": d.get("drug_name", d.get("source", "")),
  96. "section": d.get("section", ""), "source": d.get("source", ""),
  97. "score": d.get("score", 0)}
  98. for d in documents
  99. ]
  100. # 保存到 DB
  101. await _save_message(req.conversation_id, "user", req.message, intent)
  102. await _save_message(req.conversation_id, "assistant", answer, intent, sources)
  103. return ChatResponse(answer=answer, sources=sources,
  104. conversation_id=req.conversation_id, intent=intent)
  105. @router.post("/stream")
  106. async def chat_stream(req: ChatRequest, user: dict = Depends(get_current_user)):
  107. async def stream_gen():
  108. intent = classify_intent(req.message)
  109. yield f"event: intent\ndata: {intent}\n\n"
  110. yield "event: status\ndata: 正在检索...\n\n"
  111. documents = await retriever.search(req.message, intent=intent, top_k=20)
  112. documents = reranker.rerank(req.message, documents, top_k=5)
  113. yield f"event: status\ndata: 已匹配 {len(documents)} 条,生成中...\n\n"
  114. msgs = build_prompt(req.message, documents, intent=intent)
  115. sources = [
  116. {"name": d.get("drug_name", d.get("source", "")),
  117. "section": d.get("section", ""), "source": d.get("source", ""),
  118. "score": d.get("score", 0)}
  119. for d in documents
  120. ]
  121. yield "event: content\n"
  122. full_answer = []
  123. async for token in llm_client.chat_stream(msgs):
  124. full_answer.append(token)
  125. yield f"data: {token}\n\n"
  126. yield "data: [DONE]\n\n"
  127. # 元数据追加
  128. import json
  129. yield f"event: meta\ndata: {json.dumps({'intent': intent, 'sources': sources, 'cid': req.conversation_id})}\n\n"
  130. answer_text = "".join(full_answer)
  131. await _save_message(req.conversation_id, "user", req.message, intent)
  132. await _save_message(req.conversation_id, "assistant", answer_text, intent, sources)
  133. return StreamingResponse(stream_gen(), media_type="text/event-stream")
  134. # ============================================
  135. # 对话历史 API
  136. # ============================================
  137. @router.get("/history")
  138. async def get_history(
  139. page: int = Query(1, ge=1),
  140. page_size: int = Query(20, ge=1, le=50),
  141. user: dict = Depends(get_current_user),
  142. ):
  143. engine = create_async_engine(settings.database_url)
  144. try:
  145. async with engine.connect() as conn:
  146. offset = (page - 1) * page_size
  147. result = await conn.execute(
  148. text("""
  149. SELECT c.conversation_id, c.title, c.created_at,
  150. COUNT(m.id) as msg_count
  151. FROM conversations c
  152. LEFT JOIN messages m ON m.conversation_id = c.conversation_id
  153. GROUP BY c.id
  154. ORDER BY c.created_at DESC
  155. LIMIT :limit OFFSET :offset
  156. """),
  157. {"limit": page_size, "offset": offset},
  158. )
  159. items = []
  160. for row in result.fetchall():
  161. items.append({
  162. "conversation_id": row[0],
  163. "title": row[1] or "新的对话",
  164. "created_at": row[2].isoformat() if row[2] else "",
  165. "message_count": row[3],
  166. })
  167. return {"items": items, "page": page, "page_size": page_size}
  168. finally:
  169. await engine.dispose()
  170. @router.get("/history/{conversation_id}")
  171. async def get_conversation_detail(
  172. conversation_id: str,
  173. user: dict = Depends(get_current_user),
  174. ):
  175. engine = create_async_engine(settings.database_url)
  176. try:
  177. async with engine.connect() as conn:
  178. result = await conn.execute(
  179. text("""
  180. SELECT id, role, content, intent, sources, created_at
  181. FROM messages
  182. WHERE conversation_id = :cid
  183. ORDER BY created_at ASC
  184. """),
  185. {"cid": conversation_id},
  186. )
  187. messages = []
  188. for row in result.fetchall():
  189. messages.append({
  190. "id": row[0],
  191. "role": row[1],
  192. "content": row[2],
  193. "intent": row[3],
  194. "sources": row[4] if row[4] else [],
  195. "created_at": row[5].isoformat() if row[5] else "",
  196. })
  197. return {"conversation_id": conversation_id, "messages": messages}
  198. finally:
  199. await engine.dispose()
  200. @router.post("/feedback")
  201. async def submit_feedback(req: FeedbackRequest, user: dict = Depends(get_current_user)):
  202. engine = create_async_engine(settings.database_url)
  203. try:
  204. async with engine.begin() as conn:
  205. await conn.execute(
  206. text("UPDATE messages SET feedback = :fb WHERE id = :mid"),
  207. {"fb": req.feedback, "mid": req.message_id},
  208. )
  209. return {"ok": True, "message_id": req.message_id, "feedback": req.feedback}
  210. finally:
  211. await engine.dispose()
  212. # ============================================
  213. # 管理员:查看全部对话
  214. # ============================================
  215. @router.get("/admin/conversations")
  216. async def admin_list_conversations(
  217. page: int = Query(1, ge=1),
  218. page_size: int = Query(20, ge=1, le=100),
  219. keyword: Optional[str] = Query(None, description="搜索用户提问关键词"),
  220. user: dict = Depends(get_current_user),
  221. ):
  222. engine = create_async_engine(settings.database_url)
  223. try:
  224. async with engine.connect() as conn:
  225. offset = (page - 1) * page_size
  226. where = ""
  227. params = {"limit": page_size, "offset": offset}
  228. if keyword:
  229. where = "WHERE m.content LIKE :kw"
  230. params["kw"] = f"%{keyword}%"
  231. result = await conn.execute(
  232. text(f"""
  233. SELECT DISTINCT ON (c.conversation_id)
  234. c.conversation_id, c.title, c.created_at,
  235. m.content as last_msg, m.role
  236. FROM conversations c
  237. JOIN messages m ON m.conversation_id = c.conversation_id
  238. {where}
  239. ORDER BY c.conversation_id, m.created_at DESC
  240. LIMIT :limit OFFSET :offset
  241. """),
  242. params,
  243. )
  244. items = []
  245. for row in result.fetchall():
  246. items.append({
  247. "conversation_id": row[0],
  248. "title": row[1] or "新的对话",
  249. "created_at": row[2].isoformat() if row[2] else "",
  250. "last_message": (row[3] or "")[:200],
  251. "last_role": row[4],
  252. })
  253. return {"items": items, "page": page, "page_size": page_size}
  254. finally:
  255. await engine.dispose()