| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648 |
- import base64
- import json
- import uuid
- from typing import Optional
- from fastapi import APIRouter, Depends, HTTPException, Query, UploadFile, File, Form
- 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 ImageChatRequest(BaseModel):
- """图片对话请求:base64 图片"""
- model_config = {"populate_by_name": True}
- image_base64: str = Field(..., min_length=1, alias="imageBase64", description="Base64 编码的图片")
- mime_type: str = Field(default="image/jpeg", alias="mimeType", description="图片 MIME 类型")
- message: str = Field(default="", max_length=2000, description="可选的附加文字问题")
- conversation_id: str = Field(default_factory=lambda: str(uuid.uuid4()), alias="conversationId")
- class MultimodalChatRequest(BaseModel):
- """统一多模态对话请求:支持文本 + 图片 + 视频"""
- model_config = {"populate_by_name": True}
- message: str = Field(default="", max_length=2000, description="文字问题(可选)")
- # 媒体附件(图片和视频二选一或都不传,纯文本也可以)
- media_type: str = Field(default="", alias="mediaType", description="媒体类型: image / video / 空=纯文本")
- media_base64: str = Field(default="", alias="mediaBase64", description="Base64 编码的图片或视频")
- media_mime: str = Field(default="", alias="mediaMime", description="媒体 MIME 类型,如 image/jpeg, video/mp4")
- conversation_id: str = Field(default_factory=lambda: str(uuid.uuid4()), alias="conversationId")
- 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(Qwen VL 分析 + OCR → RAG 检索 → 联网搜索)
- # ============================================
- @router.post("/ask-image", response_model=ChatResponse)
- async def chat_ask_image(req: ImageChatRequest, user: dict = Depends(get_current_user)):
- """
- 图片对话 — 非流式:
- 1. Qwen VL 分析图片 + OCR 提取文字
- 2. 用提取文字做 RAG 检索
- 3. 结合检索结果 + 联网搜索生成回答
- """
- # Step 1: Qwen VL 分析图片 → 提取文字
- ocr_text = await llm_client.analyze_image(
- req.image_base64, req.mime_type,
- prompt="请分析这张图片,提取其中所有文字信息(OCR),特别是药品名称、成分、用法用量等关键药学信息。简要输出即可。",
- )
- # Step 2: 拼接用户附加文字 + OCR 结果 → RAG 检索
- query = req.message.strip() if req.message else ocr_text
- if req.message:
- query = f"{req.message}\n\n(图片OCR提取内容:{ocr_text})"
- intent = classify_intent(query)
- documents = await retriever.search(query, intent=intent, top_k=20)
- documents = reranker.rerank(query, documents, top_k=5)
- # Step 3: 构建 Prompt(含图片分析结果)+ 联网搜索
- image_context = f"\n\n【图片分析结果】\n{ocr_text}\n"
- msgs = build_prompt(query, documents, intent=intent)
- # 在 system prompt 中追加图片分析上下文
- msgs[0]["content"] += image_context
- answer = await llm_client.chat(msgs, enable_search=True)
- 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
- ]
- await _save_message(req.conversation_id, "user",
- f"[图片] {req.message}" if req.message else "[图片]",
- 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-image")
- async def chat_stream_image(req: ImageChatRequest, user: dict = Depends(get_current_user)):
- """图片对话 — SSE 流式"""
- async def stream_gen():
- intent = "drug_query"
- yield f"event: intent\ndata: {intent}\n\n"
- # Step 1: Qwen VL 流式分析图片 — 实时推给用户
- yield "event: status\ndata: 🔍 正在分析图片...\n\n"
- yield "event: content\ndata: 【📷 图片分析】\n\n"
- ocr_parts = []
- async for token in llm_client.analyze_image_stream(
- req.image_base64, req.mime_type,
- prompt="请分析这张图片,提取其中所有文字信息(OCR),特别是药品名称、成分、用法用量等。简要输出。",
- ):
- ocr_parts.append(token)
- yield f"data: {token}\n\n" # ← 实时流给用户
- ocr_text = "".join(ocr_parts)
- yield "data: \n\n"
- # Step 2: 拼接查询 → RAG
- query = req.message.strip() if req.message else ocr_text
- if req.message:
- query = f"{req.message}\n\n(图片OCR提取内容:{ocr_text})"
- intent = classify_intent(query)
- yield f"event: intent\ndata: {intent}\n\n"
- yield f"event: status\ndata: 📚 检索药典知识库...\n\n"
- documents = await retriever.search(query, intent=intent, top_k=20)
- documents = reranker.rerank(query, documents, top_k=5)
- yield f"event: status\ndata: 已匹配 {len(documents)} 条药典资料,生成回答中(已启用联网搜索)...\n\n"
- # Step 3: 构建 Prompt + 联网搜索流式生成
- image_context = f"\n\n【图片分析结果】\n{ocr_text}\n"
- msgs = build_prompt(query, documents, intent=intent)
- msgs[0]["content"] += image_context
- 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\ndata: \n【📚 药典参考回答】\n\n"
- full_answer = []
- async for token in llm_client.chat_stream(msgs, enable_search=True):
- full_answer.append(token)
- yield f"data: {token}\n\n"
- yield "data: [DONE]\n\n"
- yield f"event: meta\ndata: {json.dumps({'intent': intent, 'sources': sources, 'cid': req.conversation_id, 'ocr_text': ocr_text[:200]})}\n\n"
- answer_text = "".join(full_answer)
- await _save_message(req.conversation_id, "user",
- f"[图片] {req.message}" if req.message else "[图片]",
- intent)
- await _save_message(req.conversation_id, "assistant", answer_text, intent, sources)
- return StreamingResponse(stream_gen(), media_type="text/event-stream")
- @router.post("/upload-image")
- async def chat_upload_image(
- file: UploadFile = File(...),
- message: str = Form(default=""),
- conversation_id: str = Form(default=""),
- user: dict = Depends(get_current_user),
- ):
- """
- 上传图片文件 → 转为 base64 → 走图片对话流程
- 支持格式:jpg, jpeg, png, webp, bmp
- """
- allowed = {"image/jpeg", "image/png", "image/webp", "image/bmp"}
- if file.content_type and file.content_type not in allowed:
- raise HTTPException(400, f"不支持的图片格式: {file.content_type},支持 jpg/png/webp/bmp")
- contents = await file.read()
- if len(contents) > 10 * 1024 * 1024:
- raise HTTPException(400, "图片大小不能超过 10MB")
- image_b64 = base64.b64encode(contents).decode("utf-8")
- mime = file.content_type or "image/jpeg"
- cid = conversation_id or str(uuid.uuid4())
- req = ImageChatRequest(
- image_base64=image_b64,
- mime_type=mime,
- message=message,
- conversation_id=cid,
- )
- return await chat_ask_image(req, user)
- # ============================================================
- # 统一多模态对话 API(文本 + 图片 + 视频,一个接口全搞定)
- # ============================================================
- @router.post("/ask-multimodal", response_model=ChatResponse)
- async def chat_ask_multimodal(req: MultimodalChatRequest, user: dict = Depends(get_current_user)):
- """
- 统一多模态对话 — 非流式:
- 支持纯文本 / 文本+图片 / 文本+视频 / 纯图片 / 纯视频
- 流程:媒体分析(OCR) → RAG 检索 → 联网搜索 → 回答
- """
- conversation_id = req.conversation_id or str(uuid.uuid4())
- ocr_text = ""
- # Step 1: 如果有媒体附件,先做视觉分析 + OCR
- if req.media_base64 and req.media_type in ("image", "video"):
- media_label = "视频" if req.media_type == "video" else "图片"
- ocr_text = await llm_client.analyze_media(
- req.media_base64,
- req.media_type,
- mime_type=req.media_mime or ("" if req.media_type != "video" else "video/mp4"),
- )
- # Step 2: 拼接查询文本
- query = req.message.strip()
- if query and ocr_text:
- query = f"{query}\n\n({media_label}OCR提取内容:{ocr_text})"
- elif ocr_text:
- query = ocr_text
- elif not query:
- query = "请介绍一下自己"
- # Step 3: RAG 检索
- intent = classify_intent(query)
- documents = await retriever.search(query, intent=intent, top_k=20)
- documents = reranker.rerank(query, documents, top_k=5)
- # Step 4: 构建 Prompt + 联网搜索
- msgs = build_prompt(query, documents, intent=intent)
- if ocr_text:
- msgs[0]["content"] += f"\n\n【{media_label}分析结果】\n{ocr_text}\n"
- answer = await llm_client.chat(msgs, enable_search=bool(ocr_text) or settings.enable_web_search)
- 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
- ]
- user_msg = req.message or f"[{media_label}]" if ocr_text else req.message
- await _save_message(conversation_id, "user", user_msg, intent)
- await _save_message(conversation_id, "assistant", answer, intent, sources)
- return ChatResponse(answer=answer, sources=sources,
- conversation_id=conversation_id, intent=intent)
- @router.post("/stream-multimodal")
- async def chat_stream_multimodal(req: MultimodalChatRequest, user: dict = Depends(get_current_user)):
- """
- 统一多模态对话 — SSE 流式(全链路流式):
- OCR 分析 → 实时推送给用户 → 立即 RAG 检索 → 流式生成回答
- 用户无需等待,每一步都在实时输出
- """
- async def stream_gen():
- nonlocal req
- conversation_id = req.conversation_id or str(uuid.uuid4())
- media_label = ""
- ocr_text = ""
- has_media = req.media_base64 and req.media_type in ("image", "video")
- if has_media:
- media_label = "视频" if req.media_type == "video" else "图片"
- # Step 1: 流式 OCR 分析 — 实时推送给用户
- yield "event: status\ndata: 🔍 正在分析...\n\n"
- yield f"event: content\ndata: 【📷 {media_label}分析】\n\n"
- ocr_parts = []
- async for token in llm_client.analyze_media_stream(
- req.media_base64, req.media_type,
- mime_type=req.media_mime or "",
- ):
- ocr_parts.append(token)
- yield f"data: {token}\n\n" # ← OCR token 实时流给用户
- ocr_text = "".join(ocr_parts)
- yield "data: \n\n" # 分隔
- else:
- yield "event: intent\ndata: drug_query\n\n"
- # Step 2: 拼接查询 → RAG 检索(此时 OCR 已全部拿到)
- query = req.message.strip()
- if query and ocr_text:
- query = f"{query}\n\n({media_label}OCR提取内容:{ocr_text})"
- elif ocr_text:
- query = ocr_text
- elif not query:
- query = "请介绍一下自己"
- intent = classify_intent(query)
- yield f"event: intent\ndata: {intent}\n\n"
- yield f"event: status\ndata: 📚 检索药典知识库...\n\n"
- documents = await retriever.search(query, intent=intent, top_k=20)
- documents = reranker.rerank(query, documents, top_k=5)
- search_hint = "(已启用联网搜索)" if (ocr_text or settings.enable_web_search) else ""
- yield f"event: status\ndata: 已匹配 {len(documents)} 条,生成回答中{search_hint}...\n\n"
- # Step 3: LLM 流式生成
- msgs = build_prompt(query, documents, intent=intent)
- if ocr_text:
- msgs[0]["content"] += f"\n\n【{media_label}分析结果】\n{ocr_text}\n"
- 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\ndata: \n【📚 药典参考回答】\n\n"
- full_answer = []
- async for token in llm_client.chat_stream(msgs, enable_search=bool(ocr_text) or settings.enable_web_search):
- full_answer.append(token)
- yield f"data: {token}\n\n"
- yield "data: [DONE]\n\n"
- 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"
- answer_text = "".join(full_answer)
- user_msg = req.message or f"[{media_label}]" if ocr_text else req.message
- await _save_message(conversation_id, "user", user_msg, intent)
- await _save_message(conversation_id, "assistant", answer_text, intent, sources)
- return StreamingResponse(stream_gen(), media_type="text/event-stream")
- @router.post("/upload-media")
- async def chat_upload_media(
- file: UploadFile = File(...),
- message: str = Form(default=""),
- conversation_id: str = Form(default=""),
- user: dict = Depends(get_current_user),
- ):
- """
- 上传媒体文件(图片/视频)→ 自动识别类型 → 走多模态对话流程
- 支持:jpg, jpeg, png, webp, bmp, mp4, mov, avi, webm
- """
- mime = file.content_type or ""
- media_type = ""
- if mime.startswith("image/"):
- media_type = "image"
- max_size = 10 * 1024 * 1024 # 10MB
- elif mime.startswith("video/"):
- media_type = "video"
- max_size = 50 * 1024 * 1024 # 50MB
- else:
- raise HTTPException(400, f"不支持的媒体格式: {mime},支持 jpg/png/webp/bmp/mp4/mov/avi/webm")
- contents = await file.read()
- if len(contents) > max_size:
- raise HTTPException(400, f"文件大小不能超过 {max_size // 1024 // 1024}MB")
- media_b64 = base64.b64encode(contents).decode("utf-8")
- cid = conversation_id or str(uuid.uuid4())
- req = MultimodalChatRequest(
- message=message,
- media_type=media_type,
- media_base64=media_b64,
- media_mime=mime,
- conversation_id=cid,
- )
- return await chat_ask_multimodal(req, user)
- # ============================================
- # 对话历史 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()
|