fix: PubMed搜索合规 — 34项修复 + query_expansion UnboundLocalError + HomeView precision_mode残留
CI / backend (push) Canceled after 0s
CI / frontend (push) Canceled after 0s

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:
34047007@qq.com
2026-07-27 09:45:17 +08:00
parent 43392438c8
commit 723c4fc5c9
8 changed files with 129 additions and 81 deletions
+9 -3
View File
@@ -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
+24 -13
View File
@@ -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,
)
+27 -3
View File
@@ -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] = []
+18 -16
View File
@@ -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
+6 -5
View File
@@ -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: