| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118 |
- """
- 混合检索器:pgvector 向量检索 + BM25 关键词检索
- """
- import os
- import json
- from typing import Optional
- import httpx
- from sqlalchemy import text
- from sqlalchemy.ext.asyncio import create_async_engine
- from app.core.config import get_settings
- settings = get_settings()
- EMBEDDING_URL = "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding"
- EMBEDDING_MODEL = "text-embedding-v3"
- DB_URL = settings.database_url
- def classify_intent(query: str) -> str:
- q = query.strip()
- usage_keywords = ["怎么吃", "吃多少", "怎么用", "一天几次", "多长时间",
- "能一起吃", "孕妇能用", "儿童用量", "哺乳期",
- "饭前还是饭后", "空腹", "过量", "漏服", "停药",
- "副作用多大", "伤肝吗", "伤肾吗", "安全吗"]
- safety_sections = ["副作用", "不良反应", "禁忌", "过敏", "注意事项",
- "能不能", "可以吗", "会不"]
- regulation_keywords = ["凡例", "通则规定", "制剂通则",
- "一般规定", "通用技术要求", "检验方法通则"]
- exam_keywords = ["执业药师考试", "考点", "历年真题", "考试大纲",
- "高频考点", "报名时间"]
- # 症状/疾病求药关键词
- symptom_keywords = ["吃了什么药", "吃什么药", "该吃", "推荐用药", "推荐下用药",
- "买什么药", "推荐什么药", "用什么药", "用药建议",
- "发烧", "咳嗽", "感冒", "腹泻", "头疼", "头痛",
- "嗓子疼", "流鼻涕", "鼻塞", "肚子疼", "胃疼",
- "过敏", "皮肤痒", "失眠", "便秘", "牙疼",
- "体温", "多少度", "退烧", "止痛", "止泻"]
- if any(kw in q for kw in usage_keywords):
- return "usage_guide"
- if any(kw in q for kw in safety_sections):
- return "usage_guide"
- if any(kw in q for kw in symptom_keywords):
- return "symptom_advice"
- if any(kw in q for kw in regulation_keywords):
- return "regulation"
- if any(kw in q for kw in exam_keywords):
- return "exam_tutor"
- return "drug_query"
- async def _get_query_embedding(query: str) -> list[float]:
- api_key = os.environ.get("QWEN_API_KEY") or settings.qwen_api_key
- async with httpx.AsyncClient(timeout=30) as client:
- resp = await client.post(
- EMBEDDING_URL,
- headers={
- "Content-Type": "application/json",
- "Authorization": f"Bearer {api_key}",
- },
- json={
- "model": EMBEDDING_MODEL,
- "input": {"texts": [query]},
- "parameters": {"text_type": "query"},
- },
- )
- data = resp.json()
- if data.get("code") and data.get("code") != "":
- raise RuntimeError(f"Embedding error: {data.get('message')}")
- return data["output"]["embeddings"][0]["embedding"]
- class MixedRetriever:
- def __init__(self):
- self.table_name = "drug_chunks"
- self.vector_dim = settings.embedding_dim
- async def search(
- self,
- query: str,
- intent: str = "drug_query",
- top_k: int = 20,
- filters: Optional[dict] = None,
- ) -> list[dict]:
- query_vec = await _get_query_embedding(query)
- vec_str = "[" + ",".join(str(v) for v in query_vec) + "]"
- engine = create_async_engine(DB_URL)
- try:
- async with engine.connect() as conn:
- result = await conn.execute(
- text("""
- SELECT content, source, drug_id, section,
- 1 - (vec <=> CAST(:qv AS vector)) AS similarity
- FROM drug_chunks
- WHERE vec IS NOT NULL
- ORDER BY vec <=> CAST(:qv AS vector)
- LIMIT :k
- """),
- {"qv": vec_str, "k": top_k},
- )
- rows = result.fetchall()
- docs = []
- for row in rows:
- docs.append({
- "content": row[0],
- "source": row[1],
- "drug_id": row[2],
- "section": row[3],
- "score": float(row[4]),
- })
- return docs
- finally:
- await engine.dispose()
|