import json import uuid from typing import Optional from fastapi import APIRouter, Depends, HTTPException, Query from fastapi.responses import StreamingResponse from pydantic import BaseModel, Field from sqlalchemy import text from sqlalchemy.ext.asyncio import create_async_engine from app.core.security import get_current_user, RateLimiter from app.rag.retriever import MixedRetriever, classify_intent from app.rag.reranker import Reranker from app.rag.prompt import build_prompt from app.core.llm_client import llm_client from app.core.config import get_settings settings = get_settings() router = APIRouter(prefix="/chat", tags=["对话"]) retriever = MixedRetriever() reranker = Reranker() class ChatRequest(BaseModel): message: str = Field(..., min_length=1, max_length=2000) conversation_id: str = Field(default_factory=lambda: str(uuid.uuid4())) class ChatResponse(BaseModel): answer: str sources: list[dict] conversation_id: str intent: str class FeedbackRequest(BaseModel): conversation_id: str message_id: int feedback: str # ============================================ # DB helpers # ============================================ async def _ensure_user(openid: str) -> int: engine = create_async_engine(settings.database_url) try: async with engine.begin() as conn: result = await conn.execute( text("SELECT id FROM users WHERE openid = :openid"), {"openid": openid}, ) row = result.fetchone() if row: return row[0] result = await conn.execute( text("INSERT INTO users (openid) VALUES (:openid) RETURNING id"), {"openid": openid}, ) return result.fetchone()[0] finally: await engine.dispose() async def _save_message(conversation_id: str, role: str, content: str, intent: str = None, sources: list = None): engine = create_async_engine(settings.database_url) try: async with engine.begin() as conn: # 确保 conversation 存在(不要求 user_id 外键,因为 dev token 的 user 可能不在库中) await conn.execute( text(""" INSERT INTO conversations (conversation_id, user_id, title) VALUES (:cid, 0, :title) ON CONFLICT (conversation_id) DO NOTHING """), { "cid": conversation_id, "title": content[:50] if role == "user" else "", }, ) await conn.execute( text(""" INSERT INTO messages (conversation_id, role, content, intent, sources) VALUES (:cid, :role, :content, :intent, :sources) """), { "cid": conversation_id, "role": role, "content": content, "intent": intent, "sources": json.dumps(sources) if sources else None, }, ) finally: await engine.dispose() # ============================================ # Chat endpoints # ============================================ @router.post("/ask", response_model=ChatResponse) async def chat_ask(req: ChatRequest, user: dict = Depends(get_current_user)): intent = classify_intent(req.message) documents = await retriever.search(req.message, intent=intent, top_k=20) documents = reranker.rerank(req.message, documents, top_k=5) msgs = build_prompt(req.message, documents, intent=intent) answer = await llm_client.chat(msgs) sources = [ {"name": d.get("drug_name", d.get("source", "")), "section": d.get("section", ""), "source": d.get("source", ""), "score": d.get("score", 0)} for d in documents ] # 保存到 DB await _save_message(req.conversation_id, "user", req.message, intent) await _save_message(req.conversation_id, "assistant", answer, intent, sources) return ChatResponse(answer=answer, sources=sources, conversation_id=req.conversation_id, intent=intent) @router.post("/stream") async def chat_stream(req: ChatRequest, user: dict = Depends(get_current_user)): async def stream_gen(): intent = classify_intent(req.message) yield f"event: intent\ndata: {intent}\n\n" yield "event: status\ndata: 正在检索...\n\n" documents = await retriever.search(req.message, intent=intent, top_k=20) documents = reranker.rerank(req.message, documents, top_k=5) yield f"event: status\ndata: 已匹配 {len(documents)} 条,生成中...\n\n" msgs = build_prompt(req.message, documents, intent=intent) sources = [ {"name": d.get("drug_name", d.get("source", "")), "section": d.get("section", ""), "source": d.get("source", ""), "score": d.get("score", 0)} for d in documents ] yield "event: content\n" full_answer = [] async for token in llm_client.chat_stream(msgs): full_answer.append(token) yield f"data: {token}\n\n" yield "data: [DONE]\n\n" # 元数据追加 import json yield f"event: meta\ndata: {json.dumps({'intent': intent, 'sources': sources, 'cid': req.conversation_id})}\n\n" answer_text = "".join(full_answer) await _save_message(req.conversation_id, "user", req.message, intent) await _save_message(req.conversation_id, "assistant", answer_text, intent, sources) return StreamingResponse(stream_gen(), media_type="text/event-stream") # ============================================ # 对话历史 API # ============================================ @router.get("/history") async def get_history( page: int = Query(1, ge=1), page_size: int = Query(20, ge=1, le=50), user: dict = Depends(get_current_user), ): engine = create_async_engine(settings.database_url) try: async with engine.connect() as conn: offset = (page - 1) * page_size result = await conn.execute( text(""" SELECT c.conversation_id, c.title, c.created_at, COUNT(m.id) as msg_count FROM conversations c LEFT JOIN messages m ON m.conversation_id = c.conversation_id GROUP BY c.id ORDER BY c.created_at DESC LIMIT :limit OFFSET :offset """), {"limit": page_size, "offset": offset}, ) items = [] for row in result.fetchall(): items.append({ "conversation_id": row[0], "title": row[1] or "新的对话", "created_at": row[2].isoformat() if row[2] else "", "message_count": row[3], }) return {"items": items, "page": page, "page_size": page_size} finally: await engine.dispose() @router.get("/history/{conversation_id}") async def get_conversation_detail( conversation_id: str, user: dict = Depends(get_current_user), ): engine = create_async_engine(settings.database_url) try: async with engine.connect() as conn: result = await conn.execute( text(""" SELECT id, role, content, intent, sources, created_at FROM messages WHERE conversation_id = :cid ORDER BY created_at ASC """), {"cid": conversation_id}, ) messages = [] for row in result.fetchall(): messages.append({ "id": row[0], "role": row[1], "content": row[2], "intent": row[3], "sources": row[4] if row[4] else [], "created_at": row[5].isoformat() if row[5] else "", }) return {"conversation_id": conversation_id, "messages": messages} finally: await engine.dispose() @router.post("/feedback") async def submit_feedback(req: FeedbackRequest, user: dict = Depends(get_current_user)): engine = create_async_engine(settings.database_url) try: async with engine.begin() as conn: await conn.execute( text("UPDATE messages SET feedback = :fb WHERE id = :mid"), {"fb": req.feedback, "mid": req.message_id}, ) return {"ok": True, "message_id": req.message_id, "feedback": req.feedback} finally: await engine.dispose() # ============================================ # 管理员:查看全部对话 # ============================================ @router.get("/admin/conversations") async def admin_list_conversations( page: int = Query(1, ge=1), page_size: int = Query(20, ge=1, le=100), keyword: Optional[str] = Query(None, description="搜索用户提问关键词"), user: dict = Depends(get_current_user), ): engine = create_async_engine(settings.database_url) try: async with engine.connect() as conn: offset = (page - 1) * page_size where = "" params = {"limit": page_size, "offset": offset} if keyword: where = "WHERE m.content LIKE :kw" params["kw"] = f"%{keyword}%" result = await conn.execute( text(f""" SELECT DISTINCT ON (c.conversation_id) c.conversation_id, c.title, c.created_at, m.content as last_msg, m.role FROM conversations c JOIN messages m ON m.conversation_id = c.conversation_id {where} ORDER BY c.conversation_id, m.created_at DESC LIMIT :limit OFFSET :offset """), params, ) items = [] for row in result.fetchall(): items.append({ "conversation_id": row[0], "title": row[1] or "新的对话", "created_at": row[2].isoformat() if row[2] else "", "last_message": (row[3] or "")[:200], "last_role": row[4], }) return {"items": items, "page": page, "page_size": page_size} finally: await engine.dispose()