retriever.py 4.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118
  1. """
  2. 混合检索器:pgvector 向量检索 + BM25 关键词检索
  3. """
  4. import os
  5. import json
  6. from typing import Optional
  7. import httpx
  8. from sqlalchemy import text
  9. from sqlalchemy.ext.asyncio import create_async_engine
  10. from app.core.config import get_settings
  11. settings = get_settings()
  12. EMBEDDING_URL = "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding"
  13. EMBEDDING_MODEL = "text-embedding-v3"
  14. DB_URL = settings.database_url
  15. def classify_intent(query: str) -> str:
  16. q = query.strip()
  17. usage_keywords = ["怎么吃", "吃多少", "怎么用", "一天几次", "多长时间",
  18. "能一起吃", "孕妇能用", "儿童用量", "哺乳期",
  19. "饭前还是饭后", "空腹", "过量", "漏服", "停药",
  20. "副作用多大", "伤肝吗", "伤肾吗", "安全吗"]
  21. safety_sections = ["副作用", "不良反应", "禁忌", "过敏", "注意事项",
  22. "能不能", "可以吗", "会不"]
  23. regulation_keywords = ["凡例", "通则规定", "制剂通则",
  24. "一般规定", "通用技术要求", "检验方法通则"]
  25. exam_keywords = ["执业药师考试", "考点", "历年真题", "考试大纲",
  26. "高频考点", "报名时间"]
  27. # 症状/疾病求药关键词
  28. symptom_keywords = ["吃了什么药", "吃什么药", "该吃", "推荐用药", "推荐下用药",
  29. "买什么药", "推荐什么药", "用什么药", "用药建议",
  30. "发烧", "咳嗽", "感冒", "腹泻", "头疼", "头痛",
  31. "嗓子疼", "流鼻涕", "鼻塞", "肚子疼", "胃疼",
  32. "过敏", "皮肤痒", "失眠", "便秘", "牙疼",
  33. "体温", "多少度", "退烧", "止痛", "止泻"]
  34. if any(kw in q for kw in usage_keywords):
  35. return "usage_guide"
  36. if any(kw in q for kw in safety_sections):
  37. return "usage_guide"
  38. if any(kw in q for kw in symptom_keywords):
  39. return "symptom_advice"
  40. if any(kw in q for kw in regulation_keywords):
  41. return "regulation"
  42. if any(kw in q for kw in exam_keywords):
  43. return "exam_tutor"
  44. return "drug_query"
  45. async def _get_query_embedding(query: str) -> list[float]:
  46. api_key = os.environ.get("QWEN_API_KEY") or settings.qwen_api_key
  47. async with httpx.AsyncClient(timeout=30) as client:
  48. resp = await client.post(
  49. EMBEDDING_URL,
  50. headers={
  51. "Content-Type": "application/json",
  52. "Authorization": f"Bearer {api_key}",
  53. },
  54. json={
  55. "model": EMBEDDING_MODEL,
  56. "input": {"texts": [query]},
  57. "parameters": {"text_type": "query"},
  58. },
  59. )
  60. data = resp.json()
  61. if data.get("code") and data.get("code") != "":
  62. raise RuntimeError(f"Embedding error: {data.get('message')}")
  63. return data["output"]["embeddings"][0]["embedding"]
  64. class MixedRetriever:
  65. def __init__(self):
  66. self.table_name = "drug_chunks"
  67. self.vector_dim = settings.embedding_dim
  68. async def search(
  69. self,
  70. query: str,
  71. intent: str = "drug_query",
  72. top_k: int = 20,
  73. filters: Optional[dict] = None,
  74. ) -> list[dict]:
  75. query_vec = await _get_query_embedding(query)
  76. vec_str = "[" + ",".join(str(v) for v in query_vec) + "]"
  77. engine = create_async_engine(DB_URL)
  78. try:
  79. async with engine.connect() as conn:
  80. result = await conn.execute(
  81. text("""
  82. SELECT content, source, drug_id, section,
  83. 1 - (vec <=> CAST(:qv AS vector)) AS similarity
  84. FROM drug_chunks
  85. WHERE vec IS NOT NULL
  86. ORDER BY vec <=> CAST(:qv AS vector)
  87. LIMIT :k
  88. """),
  89. {"qv": vec_str, "k": top_k},
  90. )
  91. rows = result.fetchall()
  92. docs = []
  93. for row in rows:
  94. docs.append({
  95. "content": row[0],
  96. "source": row[1],
  97. "drug_id": row[2],
  98. "section": row[3],
  99. "score": float(row[4]),
  100. })
  101. return docs
  102. finally:
  103. await engine.dispose()