| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183 |
- """
- 混合检索器:pgvector 向量检索 + BM25 关键词检索
- """
- import os
- import re
- 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()
- # --- 已知药物过敏检测:XX过敏处理/怎么办/急救/替代 → 用药安全咨询 ---
- # 必须在 symptom_keywords 之前检测,否则会被 "过敏" 误判为症状求药
- known_allergy_drugs = [
- "青霉素", "头孢", "磺胺", "阿莫西林", "布洛芬",
- "阿司匹林", "链霉素", "庆大霉素", "四环素",
- "红霉素", "氯霉素", "万古霉素", "喹诺酮",
- "普鲁卡因", "利多卡因", "碘", "破伤风",
- ]
- if "过敏" in q and any(drug in q for drug in known_allergy_drugs):
- return "usage_guide"
- # --- 已知药物过敏通用模式:对XX过敏 / XX药过敏 / 药物过敏 ---
- known_allergy_patterns = [
- r"对.{1,6}过敏", # 对青霉素过敏
- r".{1,6}药过敏", # 头孢药过敏
- r"药物过敏", # 药物过敏
- r".+过敏(处理|怎么办|急救|替代|注意|救治|抢救|应急)",
- ]
- if any(re.search(pat, q) for pat in known_allergy_patterns):
- return "usage_guide"
- # --- 否定语义检测:先检查是否包含否定模式 ---
- negation_patterns = [
- r"不是", r"没有", r"并非", r"不算", r"不属于",
- r"这不是", r"我没有", r"不包含", r"不涉及",
- ]
- has_negation = any(re.search(pat, q) for pat in negation_patterns)
- # 如果有否定语义,直接返回药品查询(用户可能在排除某些情况)
- if has_negation:
- return "drug_query"
- # --- 关键词定义 ---
- usage_keywords = [
- "怎么吃", "吃多少", "怎么用", "一天几次", "多长时间",
- "能一起吃", "孕妇能用", "儿童用量", "哺乳期",
- "饭前还是饭后", "空腹", "过量", "漏服", "停药",
- "副作用多大", "伤肝吗", "伤肾吗", "安全吗",
- ]
- safety_sections = [
- "副作用", "不良反应", "禁忌", "注意事项",
- "能不能", "可以吗", "会不会",
- ]
- regulation_keywords = [
- "凡例", "通则规定", "制剂通则",
- "一般规定", "通用技术要求", "检验方法通则",
- ]
- exam_keywords = [
- "执业药师考试", "考点", "历年真题", "考试大纲",
- "高频考点", "报名时间",
- ]
- # 症状/疾病求药关键词
- symptom_keywords = [
- "吃了什么药", "吃什么药", "该吃", "推荐用药", "推荐下用药",
- "买什么药", "推荐什么药", "用什么药", "用药建议",
- "发烧", "咳嗽", "感冒", "腹泻", "头疼", "头痛",
- "嗓子疼", "流鼻涕", "鼻塞", "肚子疼", "胃疼",
- "过敏", "皮肤痒", "失眠", "便秘", "牙疼",
- "体温", "多少度", "退烧", "止痛", "止泻",
- ]
- # --- 优先级匹配(usage > symptom > regulation > exam > drug_query) ---
- # 注意:"过敏"从 safety_sections 中移除,只保留在 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
- self._engine = None
- @property
- def engine(self):
- """延迟创建数据库引擎,复用连接池。"""
- if self._engine is None:
- self._engine = create_async_engine(
- DB_URL,
- pool_size=5,
- max_overflow=10,
- pool_pre_ping=True,
- )
- return self._engine
- async def close(self):
- """释放数据库连接池。"""
- if self._engine is not None:
- await self._engine.dispose()
- self._engine = None
- 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) + "]"
- async with self.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
|