fix: PubMed搜索合规 — 34项修复 + query_expansion UnboundLocalError + HomeView precision_mode残留
Batch1 — 解析器 (pubmed_query_parser.py) - P0-1: 未知字段标签降级为 WORD token 而非 ParseError - P2-1: 增加未消费 token 检查 - P2-2: PRISMA 字段标签正则 [\w:+] → [\w:]+ - P2-3: Unicode NFKC 规格化输入 - P2-4: re.ASCII 防止 Unicode 数字匹配 - P2-9: 移除重复 dataclass 字段 Batch2 — 搜索引擎 (search_engine.py) - P1-1: 批量 PMID 查询替代 N+1 循环 - P1-2: isdigit() → isdecimal() - P1-5: 移除 precision_mode 参数 - P2-5: 移除死代码 - P2-6: 统一 _CHINESE_RE 正则 Batch3 — ATM 引擎 (query_expansion.py) - P0-4: name_zh ILIKE 中文回退 + _find_mesh_tags 中文降级 - P1-3: name_en ILIKE 加 LIMIT 100 Batch4 — API 层 (features.py) - P1-5: 移除 precision_mode 请求字段 - P1-11: 增加 logging - P2-8: NLM_SUBSET_LABELS f-string 安全注释 - P2-12: split 校验器近似性注释 Batch5 — SearchView.vue - P0-5a/b/c: 日期修复(UTC 方法、互斥逻辑、restoreFromQuery 合并) - P2-10: 筛选模态关闭时重搜 - P3-1: 搜索框 aria-label Batch6 — HomeView.vue - P1-10: URL date → date_from/date_to - P2-11: clearSearch 清空 feedItems 并重加载 Batch7 — LiteratureCard.vue - P1-6: 字段标签正则 [\w-]+ → [\w:-]+ - P1-7: terms 切片限制 20 项防 ReDoS 后修复: - query_expansion.py _find_partial_mesh_tags UnboundLocalError(单非中文词未初始化 tag_ids) - HomeView.vue handleAdvancedSearch precision_mode 残留引用
This commit is contained in:
@@ -1,8 +1,11 @@
|
||||
"""规则引擎 + 每日摘要 + 高级搜索 API"""
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
from sqlalchemy import func, select, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@@ -61,7 +64,6 @@ class AdvancedSearchRequest(BaseModel):
|
||||
is_oa: bool | None = None
|
||||
language: str | None = None
|
||||
languages: list[str] | None = None
|
||||
precision_mode: str = "majr"
|
||||
nlm_subsets: list[str] | None = None
|
||||
page: int = Field(1, ge=1)
|
||||
page_size: int = Field(20, ge=1, le=100)
|
||||
@@ -89,6 +91,7 @@ class AdvancedSearchRequest(BaseModel):
|
||||
@field_validator('query')
|
||||
@classmethod
|
||||
def check_query_complexity(cls, v: str) -> str:
|
||||
# 粗略检查(split 计数与 tokeniser 不完全一致),精确限制由 tokeniser 的 MAX_TERMS=100 执行
|
||||
if len(v.split()) > 100:
|
||||
raise ValueError('查询词过多(最多 100 个词),请简化搜索条件')
|
||||
return v
|
||||
@@ -172,6 +175,7 @@ async def _load_filter_options(db: AsyncSession) -> dict:
|
||||
languages = [{"code": r[0], "count": r[1]} for r in lang_rows]
|
||||
|
||||
# 3. nlm_subsets(单次 JOIN + COUNT FILTER,代替 8 次独立查询)
|
||||
# 注:code 来自上方 NLM_SUBSET_LABELS 硬编码常量,无注入风险
|
||||
subset_filter_cols = ", ".join(
|
||||
f'COUNT(*) FILTER (WHERE j.nlm_subsets @> ARRAY[\'{code}\']) AS "{code}"'
|
||||
for code in NLM_SUBSET_LABELS
|
||||
@@ -246,9 +250,11 @@ async def _load_filter_options(db: AsyncSession) -> dict:
|
||||
async def advanced_search(req: AdvancedSearchRequest, db: AsyncSession = Depends(get_db)):
|
||||
try:
|
||||
return await AdvancedSearchEngine.search(db, **req.model_dump())
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=400, detail="搜索参数错误,请检查输入") from None
|
||||
except ValueError as ve:
|
||||
logger.warning("搜索参数错误: %s", ve)
|
||||
raise HTTPException(status_code=400, detail="搜索参数错误,请检查输入")
|
||||
except Exception as e:
|
||||
logger.exception("搜索服务内部错误")
|
||||
raise HTTPException(status_code=500, detail="搜索服务内部错误") from e
|
||||
|
||||
|
||||
|
||||
@@ -19,6 +19,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
import unicodedata
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum, auto
|
||||
|
||||
@@ -127,7 +128,7 @@ _TOKEN_PATTERNS: list[tuple[TokenType, str]] = [
|
||||
|
||||
_TOKEN_RE = re.compile(
|
||||
'|'.join(f'(?P<{t.name}>{p})' for t, p in _TOKEN_PATTERNS),
|
||||
re.IGNORECASE,
|
||||
re.IGNORECASE | re.ASCII,
|
||||
)
|
||||
|
||||
|
||||
@@ -140,7 +141,11 @@ def tokenise(query: str) -> list[Token]:
|
||||
ttype = TokenType[name]
|
||||
if ttype == TokenType.UNKNOWN_FIELD:
|
||||
fname = value.strip('[]').upper()
|
||||
raise ParseError(f"不认识的字段标签 [{fname}],降级为简单文本搜索")
|
||||
# P0-1: 不认识字段标签时降级为 WORD,不终止解析
|
||||
tokens.append(Token(TokenType.WORD, value.strip('[]')))
|
||||
if len(tokens) > MAX_TERMS:
|
||||
raise ParseError(f"查询词过多(超过 {MAX_TERMS} 个),降级为简单文本搜索")
|
||||
continue
|
||||
tokens.append(Token(ttype, value))
|
||||
if len(tokens) > MAX_TERMS:
|
||||
raise ParseError(f"查询词过多(超过 {MAX_TERMS} 个),降级为简单文本搜索")
|
||||
@@ -169,11 +174,6 @@ class ParsedPubmedQuery:
|
||||
tiab_terms: list[Term] = field(default_factory=list) # [TIAB]
|
||||
author_terms: list[Term] = field(default_factory=list) # [AU]
|
||||
journal_terms: list[Term] = field(default_factory=list) # [TA]
|
||||
mesh_terms: list[str] = field(default_factory=list) # [MH]
|
||||
majr_terms: list[str] = field(default_factory=list) # [MAJR]
|
||||
pub_types: list[str] = field(default_factory=list) # [PT]
|
||||
doi_terms: list[str] = field(default_factory=list) # [DOI]
|
||||
pmid_terms: list[int] = field(default_factory=list) # [PMID]
|
||||
affiliation_terms: list[Term] = field(default_factory=list) # [AD]
|
||||
language_terms: list[Term] = field(default_factory=list) # [LA]
|
||||
volume_terms: list[Term] = field(default_factory=list) # [VI]
|
||||
@@ -275,10 +275,8 @@ class PubmedQueryParser:
|
||||
"""入口:解析完整的查询字符串。"""
|
||||
result = ParsedPubmedQuery()
|
||||
self._depth = 0 # 括号嵌套深度计数器
|
||||
try:
|
||||
terms = self._parse_or_expr(result)
|
||||
except ParseError:
|
||||
return ParsedPubmedQuery()
|
||||
terms = self._parse_or_expr(result)
|
||||
# 解析错误由 parse_pubmed_query 统一降级处理
|
||||
|
||||
# Detect boolean operator from token stream
|
||||
has_and = any(t.type == TokenType.AND for t in self.tokens)
|
||||
@@ -296,6 +294,13 @@ class PubmedQueryParser:
|
||||
for t in _ungrouped:
|
||||
self._dispatch_term(result, t)
|
||||
|
||||
# P2-1: Handle unconsumed tokens (e.g., orphan text after RPAREN)
|
||||
if self.pos < len(self.tokens) - 1:
|
||||
for t in self.tokens[self.pos:-1]: # exclude EOF token
|
||||
if t.type in (TokenType.WORD, TokenType.QUOTED, TokenType.NUMBER):
|
||||
text = t.value.strip('"') if t.type == TokenType.QUOTED else t.value
|
||||
result.plain_terms.append(Term(text=text, exact=(t.type == TokenType.QUOTED)))
|
||||
|
||||
return result
|
||||
|
||||
def _dispatch_term(self, result: ParsedPubmedQuery, term: Term) -> None:
|
||||
@@ -574,6 +579,8 @@ def parse_pubmed_query(query: str) -> ParsedPubmedQuery:
|
||||
return ParsedPubmedQuery()
|
||||
|
||||
try:
|
||||
# P2-3: Unicode normalization — strip zero-width chars, normalize fullwidth digits
|
||||
query = unicodedata.normalize('NFKC', query)
|
||||
# 将 YYYY/MM/DD 格式的日期分隔符统一为 YYYY-MM-DD,使 tokeniser 正确识别为 DATE
|
||||
import re as _re
|
||||
query = _re.sub(r'(\d{4})/(\d{2})/(\d{2})', r'\1-\2-\3', query)
|
||||
@@ -581,7 +588,11 @@ def parse_pubmed_query(query: str) -> ParsedPubmedQuery:
|
||||
parser = PubmedQueryParser(tokens)
|
||||
return parser.parse()
|
||||
except (ParseError, IndexError, ValueError):
|
||||
return ParsedPubmedQuery()
|
||||
# P0-1: 降级时返回原始查询作为 plain_terms,不丢失用户输入
|
||||
degraded = ParsedPubmedQuery()
|
||||
for t in query.strip().split():
|
||||
degraded.plain_terms.append(Term(text=t))
|
||||
return degraded
|
||||
|
||||
|
||||
def extract_pubmed_query_for_prisma(query: str) -> tuple[str, list[str]]:
|
||||
@@ -598,7 +609,7 @@ def extract_pubmed_query_for_prisma(query: str) -> tuple[str, list[str]]:
|
||||
|
||||
# 标准化:统一字段大写
|
||||
normalized = re.sub(
|
||||
r'\[(\w+)\]',
|
||||
r'\[([\w:]+)\]',
|
||||
lambda m: f'[{m.group(1).upper()}]',
|
||||
query,
|
||||
)
|
||||
|
||||
@@ -17,6 +17,7 @@
|
||||
"""
|
||||
|
||||
import logging
|
||||
import re
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import or_ as _or_, select
|
||||
@@ -81,18 +82,30 @@ async def _find_mesh_tags(db: AsyncSession, query: str) -> list[UUID]:
|
||||
tag_ids.append(tid)
|
||||
seen.add(tid)
|
||||
|
||||
# 方法 B: name_en ILIKE 匹配
|
||||
# 方法 B: name_en ILIKE 匹配(P1-3: 加 LIMIT 100 防止常见词匹配过多)
|
||||
like_pattern = f"%{_escape_ilike(query)}%"
|
||||
stmt = select(GlobalTag.id).where(
|
||||
GlobalTag.source.in_(["mesh", "manual"]),
|
||||
GlobalTag.name_en.ilike(like_pattern),
|
||||
)
|
||||
).limit(100)
|
||||
rows = await db.execute(stmt)
|
||||
for (tid,) in rows:
|
||||
if tid not in seen:
|
||||
tag_ids.append(tid)
|
||||
seen.add(tid)
|
||||
|
||||
# 方法 C: name_zh ILIKE 匹配(P0-4: 中文查询降级)
|
||||
if not tag_ids and re.search(r'[一-鿿]', q):
|
||||
stmt = select(GlobalTag.id).where(
|
||||
GlobalTag.source.in_(["mesh", "manual"]),
|
||||
GlobalTag.name_zh.ilike(like_pattern),
|
||||
)
|
||||
rows = await db.execute(stmt)
|
||||
for (tid,) in rows:
|
||||
if tid not in seen:
|
||||
tag_ids.append(tid)
|
||||
seen.add(tid)
|
||||
|
||||
return tag_ids
|
||||
|
||||
|
||||
@@ -103,7 +116,18 @@ async def _find_partial_mesh_tags(db: AsyncSession, query: str) -> list[UUID]:
|
||||
"""
|
||||
words = [w.strip().lower() for w in query.strip().split() if len(w.strip()) >= MIN_QUERY_LENGTH][:10]
|
||||
if len(words) < 2:
|
||||
return []
|
||||
# P0-4: 中文查询无法按空格分词,尝试整体 name_zh ILIKE 匹配
|
||||
if re.search(r'[一-鿿]', query):
|
||||
tag_ids = []
|
||||
stmt = select(GlobalTag.id).where(
|
||||
GlobalTag.source.in_(["mesh", "manual"]),
|
||||
GlobalTag.name_zh.ilike(f'%{_escape_ilike(query.strip().lower())}%'),
|
||||
)
|
||||
rows = await db.execute(stmt)
|
||||
for (tid,) in rows:
|
||||
tag_ids.append(tid)
|
||||
return tag_ids
|
||||
return [] # non-Chinese single word: no partial match possible
|
||||
|
||||
seen: set[UUID] = set()
|
||||
tag_ids: list[UUID] = []
|
||||
|
||||
@@ -102,7 +102,6 @@ class AdvancedSearchEngine:
|
||||
is_oa: bool | None = None, # True = 仅开放获取
|
||||
language: str | None = None, # 语言代码(en/zh/fr 等,向后兼容)
|
||||
languages: list[str] | None = None, # 语言代码列表(多选)
|
||||
precision_mode: str = "majr",
|
||||
nlm_subsets: list[str] | None = None, # NLM 期刊子集(AIM/M/S 等)
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
@@ -212,7 +211,7 @@ class AdvancedSearchEngine:
|
||||
query = ' '.join(query.split())
|
||||
# 中文搜索:自动匹配 GlobalTag.name_zh → 注入 tag_ids,跳过 ILIKE
|
||||
import re as _re
|
||||
_CHINESE_RE = _re.compile(r'[一-鿿]')
|
||||
_CHINESE_RE = _re.compile(r'[一-鿿㐀-䶿豈-]')
|
||||
if _CHINESE_RE.search(query):
|
||||
tag_matches = (await db.execute(
|
||||
select(GlobalTag.id).where(GlobalTag.name_zh.ilike(f'%{_escape_ilike(query.strip())}%'))
|
||||
@@ -231,23 +230,29 @@ class AdvancedSearchEngine:
|
||||
terms = [p for p in _phrases if p.strip()] + [t for t in _rest if t not in _phrases]
|
||||
# 单数字词:优先 PMID 精确匹配(unique index 5ms 返回)
|
||||
# 不是 PMID 时才回退到 ILIKE 兜底(DOI 片段等),不做 tsquery 避免 seq scan
|
||||
numeric_terms = [t for t in terms if t.isdigit() and len(t) <= 15]
|
||||
text_terms = [t for t in terms if not (t.isdigit() and len(t) <= 15)]
|
||||
numeric_terms = [t for t in terms if t.isdecimal() and len(t) <= 15]
|
||||
text_terms = [t for t in terms if not (t.isdecimal() and len(t) <= 15)]
|
||||
if numeric_terms:
|
||||
num_conds = []
|
||||
for t in numeric_terms:
|
||||
if exact_phrase:
|
||||
# 精确短语模式:跳过 PMID 快速路径,走 ILIKE
|
||||
if exact_phrase:
|
||||
# 精确短语模式:跳过 PMID 快速路径,走 ILIKE
|
||||
for t in numeric_terms:
|
||||
num_conds.append(or_(
|
||||
GlobalLiterature.title.ilike(t),
|
||||
GlobalLiterature.doi.ilike(t),
|
||||
))
|
||||
else:
|
||||
else:
|
||||
# P1-1: 批量检查所有 PMID(单次 IN 查询替代 N+1 次 SELECT)
|
||||
numeric_ids = [int(t) for t in numeric_terms]
|
||||
db_pmids: set[int] = set()
|
||||
rows = await db.execute(
|
||||
select(GlobalLiterature.pmid).where(GlobalLiterature.pmid.in_(numeric_ids))
|
||||
)
|
||||
for (pmid,) in rows:
|
||||
db_pmids.add(pmid)
|
||||
for t in numeric_terms:
|
||||
p = int(t)
|
||||
exists = (await db.execute(
|
||||
select(GlobalLiterature.id).where(GlobalLiterature.pmid == p).limit(1)
|
||||
)).scalar_one_or_none()
|
||||
if exists:
|
||||
if p in db_pmids:
|
||||
num_conds.append(GlobalLiterature.pmid == p)
|
||||
else:
|
||||
p_like = f"%{_escape_ilike(t)}%"
|
||||
@@ -263,7 +268,7 @@ class AdvancedSearchEngine:
|
||||
# ATM 展开(仅 field="all" 时,字段搜索不应自动扩到 MeSH)
|
||||
_atm_cond = None
|
||||
_atm_query = query.replace('"', '').replace("'", '').strip()
|
||||
if _atm_query and not _CHINESE_RE.search(_atm_query) and field == "all":
|
||||
if _atm_query and field == "all":
|
||||
_atm_cond = await _expand_atm(db, _atm_query)
|
||||
|
||||
_cond_before = len(conditions)
|
||||
@@ -468,9 +473,6 @@ class AdvancedSearchEngine:
|
||||
except Exception:
|
||||
logger.exception("Year counts query failed")
|
||||
year_counts = []
|
||||
q = select(GlobalLiterature)
|
||||
if conditions:
|
||||
q = q.where(and_(*conditions))
|
||||
|
||||
# 排序(PubMed 查询时跳过 ts_rank,避免语法标签噪音)
|
||||
_relevance_query = query
|
||||
|
||||
@@ -239,13 +239,12 @@ class TestParserGaps:
|
||||
# ── A6: edge cases ──
|
||||
|
||||
def test_a6c_over_100_tokens(self):
|
||||
"""More than 100 tokens — tokeniser raises ParseError, parser degrades gracefully"""
|
||||
"""More than 100 tokens — tokeniser raises ParseError, parser degrades to plain terms"""
|
||||
words = "word " * 101
|
||||
r = parse_pubmed_query(words.strip())
|
||||
assert isinstance(r, ParsedPubmedQuery)
|
||||
# Tokeniser raises ParseError(ValueError) at >100 tokens,
|
||||
# parser catches it and returns empty result = graceful degradation
|
||||
assert len(r.plain_terms) == 0
|
||||
# Degradation returns all 101 words as plain_terms (P0-1)
|
||||
assert len(r.plain_terms) == 101
|
||||
|
||||
def test_a6_repeated_AND(self):
|
||||
"""AND AND — repeated boolean"""
|
||||
@@ -263,9 +262,11 @@ class TestParserGaps:
|
||||
assert isinstance(r, ParsedPubmedQuery)
|
||||
|
||||
def test_a6_unknown_field(self):
|
||||
"""cancer[XX] — unknown field tag"""
|
||||
"""cancer[XX] — unknown field tag degrades to plain text (P0-1)"""
|
||||
r = parse_pubmed_query("cancer[XX]")
|
||||
assert isinstance(r, ParsedPubmedQuery)
|
||||
# P0-1: unknown field tag now emits WORD token instead of aborting
|
||||
assert len(r.plain_terms) >= 1
|
||||
|
||||
|
||||
class TestFieldTagCompleteness:
|
||||
|
||||
Reference in New Issue
Block a user