retriever.py 5.5 KB

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