""" 混合检索器: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