| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297 |
- 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()
|