retriever.py 6.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183
  1. """
  2. 混合检索器:pgvector 向量检索 + BM25 关键词检索
  3. """
  4. import os
  5. import re
  6. import json
  7. from typing import Optional
  8. import httpx
  9. from sqlalchemy import text
  10. from sqlalchemy.ext.asyncio import create_async_engine
  11. from app.core.config import get_settings
  12. settings = get_settings()
  13. EMBEDDING_URL = "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding"
  14. EMBEDDING_MODEL = "text-embedding-v3"
  15. DB_URL = settings.database_url
  16. def classify_intent(query: str) -> str:
  17. """基于关键词的意图分类,含否定语义处理和置信度兜底。"""
  18. q = query.strip()
  19. # --- 已知药物过敏检测:XX过敏处理/怎么办/急救/替代 → 用药安全咨询 ---
  20. # 必须在 symptom_keywords 之前检测,否则会被 "过敏" 误判为症状求药
  21. known_allergy_drugs = [
  22. "青霉素", "头孢", "磺胺", "阿莫西林", "布洛芬",
  23. "阿司匹林", "链霉素", "庆大霉素", "四环素",
  24. "红霉素", "氯霉素", "万古霉素", "喹诺酮",
  25. "普鲁卡因", "利多卡因", "碘", "破伤风",
  26. ]
  27. if "过敏" in q and any(drug in q for drug in known_allergy_drugs):
  28. return "usage_guide"
  29. # --- 已知药物过敏通用模式:对XX过敏 / XX药过敏 / 药物过敏 ---
  30. known_allergy_patterns = [
  31. r"对.{1,6}过敏", # 对青霉素过敏
  32. r".{1,6}药过敏", # 头孢药过敏
  33. r"药物过敏", # 药物过敏
  34. r".+过敏(处理|怎么办|急救|替代|注意|救治|抢救|应急)",
  35. ]
  36. if any(re.search(pat, q) for pat in known_allergy_patterns):
  37. return "usage_guide"
  38. # --- 否定语义检测:先检查是否包含否定模式 ---
  39. negation_patterns = [
  40. r"不是", r"没有", r"并非", r"不算", r"不属于",
  41. r"这不是", r"我没有", r"不包含", r"不涉及",
  42. ]
  43. has_negation = any(re.search(pat, q) for pat in negation_patterns)
  44. # 如果有否定语义,直接返回药品查询(用户可能在排除某些情况)
  45. if has_negation:
  46. return "drug_query"
  47. # --- 关键词定义 ---
  48. usage_keywords = [
  49. "怎么吃", "吃多少", "怎么用", "一天几次", "多长时间",
  50. "能一起吃", "孕妇能用", "儿童用量", "哺乳期",
  51. "饭前还是饭后", "空腹", "过量", "漏服", "停药",
  52. "副作用多大", "伤肝吗", "伤肾吗", "安全吗",
  53. ]
  54. safety_sections = [
  55. "副作用", "不良反应", "禁忌", "注意事项",
  56. "能不能", "可以吗", "会不会",
  57. ]
  58. regulation_keywords = [
  59. "凡例", "通则规定", "制剂通则",
  60. "一般规定", "通用技术要求", "检验方法通则",
  61. ]
  62. exam_keywords = [
  63. "执业药师考试", "考点", "历年真题", "考试大纲",
  64. "高频考点", "报名时间",
  65. ]
  66. # 症状/疾病求药关键词
  67. symptom_keywords = [
  68. "吃了什么药", "吃什么药", "该吃", "推荐用药", "推荐下用药",
  69. "买什么药", "推荐什么药", "用什么药", "用药建议",
  70. "发烧", "咳嗽", "感冒", "腹泻", "头疼", "头痛",
  71. "嗓子疼", "流鼻涕", "鼻塞", "肚子疼", "胃疼",
  72. "过敏", "皮肤痒", "失眠", "便秘", "牙疼",
  73. "体温", "多少度", "退烧", "止痛", "止泻",
  74. ]
  75. # --- 优先级匹配(usage > symptom > regulation > exam > drug_query) ---
  76. # 注意:"过敏"从 safety_sections 中移除,只保留在 symptom_keywords,
  77. # 避免"过敏"被误判为用法询问而非症状求药。
  78. if any(kw in q for kw in usage_keywords):
  79. return "usage_guide"
  80. if any(kw in q for kw in safety_sections):
  81. return "usage_guide"
  82. if any(kw in q for kw in symptom_keywords):
  83. return "symptom_advice"
  84. if any(kw in q for kw in regulation_keywords):
  85. return "regulation"
  86. if any(kw in q for kw in exam_keywords):
  87. return "exam_tutor"
  88. # 兜底:无明确意图时走药品通用查询(向量检索命中率最高)
  89. return "drug_query"
  90. async def _get_query_embedding(query: str) -> list[float]:
  91. api_key = os.environ.get("QWEN_API_KEY") or settings.qwen_api_key
  92. async with httpx.AsyncClient(timeout=30) as client:
  93. resp = await client.post(
  94. EMBEDDING_URL,
  95. headers={
  96. "Content-Type": "application/json",
  97. "Authorization": f"Bearer {api_key}",
  98. },
  99. json={
  100. "model": EMBEDDING_MODEL,
  101. "input": {"texts": [query]},
  102. "parameters": {"text_type": "query"},
  103. },
  104. )
  105. data = resp.json()
  106. if data.get("code") and data.get("code") != "":
  107. raise RuntimeError(f"Embedding error: {data.get('message')}")
  108. return data["output"]["embeddings"][0]["embedding"]
  109. class MixedRetriever:
  110. def __init__(self):
  111. self.table_name = "drug_chunks"
  112. self.vector_dim = settings.embedding_dim
  113. self._engine = None
  114. @property
  115. def engine(self):
  116. """延迟创建数据库引擎,复用连接池。"""
  117. if self._engine is None:
  118. self._engine = create_async_engine(
  119. DB_URL,
  120. pool_size=5,
  121. max_overflow=10,
  122. pool_pre_ping=True,
  123. )
  124. return self._engine
  125. async def close(self):
  126. """释放数据库连接池。"""
  127. if self._engine is not None:
  128. await self._engine.dispose()
  129. self._engine = None
  130. async def search(
  131. self,
  132. query: str,
  133. intent: str = "drug_query",
  134. top_k: int = 20,
  135. filters: Optional[dict] = None,
  136. ) -> list[dict]:
  137. query_vec = await _get_query_embedding(query)
  138. vec_str = "[" + ",".join(str(v) for v in query_vec) + "]"
  139. async with self.engine.connect() as conn:
  140. result = await conn.execute(
  141. text("""
  142. SELECT content, source, drug_id, section,
  143. 1 - (vec <=> CAST(:qv AS vector)) AS similarity
  144. FROM drug_chunks
  145. WHERE vec IS NOT NULL
  146. ORDER BY vec <=> CAST(:qv AS vector)
  147. LIMIT :k
  148. """),
  149. {"qv": vec_str, "k": top_k},
  150. )
  151. rows = result.fetchall()
  152. docs = []
  153. for row in rows:
  154. docs.append({
  155. "content": row[0],
  156. "source": row[1],
  157. "drug_id": row[2],
  158. "section": row[3],
  159. "score": float(row[4]),
  160. })
  161. return docs