fix(retrieval): 修复路径逃逸、部分提交、失效向量与分页问题
- 路径逃逸:清洗 folder(拒绝 ..、绝对路径/盘符),_abs_path 增加 Vault 边界校验 - 部分提交:create/update 索引失败时回滚文件 - 失效向量残留:replace_note_metadata 返回旧 id,index_note 清理旧向量 - 分页不完整:fts_count 返回真实命中总数,候选池覆盖 offset+limit Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -71,11 +71,18 @@ def replace_note_metadata(
|
||||
created_at: datetime,
|
||||
updated_at: datetime,
|
||||
blocks: list[NoteBlock],
|
||||
) -> None:
|
||||
"""整体替换一条笔记的元数据、Block 与 FTS5 索引(单事务)。"""
|
||||
) -> list[str]:
|
||||
"""整体替换一条笔记的元数据、Block 与 FTS5 索引(单事务)。
|
||||
|
||||
返回替换前的旧 block_id 列表,供调用方清理 vec_blocks 中已失效的向量。
|
||||
"""
|
||||
conn = connect()
|
||||
try:
|
||||
with transaction(conn):
|
||||
old_block_ids = [
|
||||
row["block_id"]
|
||||
for row in conn.execute("SELECT block_id FROM blocks WHERE note_id = ?", (note_id,))
|
||||
]
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO notes (note_id, title, file_path, folder, tags, created_at, updated_at)
|
||||
@@ -109,6 +116,7 @@ def replace_note_metadata(
|
||||
"INSERT INTO blocks_fts (block_id, note_id, heading_path, content) VALUES (?, ?, ?, ?)",
|
||||
(block.block_id, note_id, segment(" ".join(block.heading_path)), segment(block.content)),
|
||||
)
|
||||
return old_block_ids
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
@@ -208,6 +216,17 @@ def fts_search(match: str, limit: int = 100) -> list[FtsHit]:
|
||||
conn.close()
|
||||
|
||||
|
||||
def fts_count(match: str) -> int:
|
||||
"""返回 FTS5 命中总数,用于分页 total(不受候选池截断影响)。"""
|
||||
conn = connect()
|
||||
try:
|
||||
return conn.execute(
|
||||
"SELECT COUNT(*) FROM blocks_fts WHERE blocks_fts MATCH ?", (match,)
|
||||
).fetchone()[0]
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def get_block_hits(block_ids: list[str]) -> list[BlockHit]:
|
||||
if not block_ids:
|
||||
return []
|
||||
|
||||
@@ -27,6 +27,8 @@ from app.textutils import make_snippet, match_query
|
||||
|
||||
# 每个通道的候选池大小;真实规模上来后按 Retrieval Config 调整
|
||||
CANDIDATE_POOL = 50
|
||||
# 分页窗口上限:候选池至少覆盖 offset+limit,但设上限防止超大 offset 撑爆内存
|
||||
MAX_CANDIDATE_POOL = 200
|
||||
|
||||
|
||||
class RetrievalEngine:
|
||||
@@ -41,23 +43,30 @@ class RetrievalEngine:
|
||||
self.vector_store = vector_store
|
||||
|
||||
async def search(self, request: SearchRequest) -> SearchResponse:
|
||||
# 候选池至少覆盖本次请求的 offset+limit,保证分页能取到目标页;设上限防内存失控
|
||||
window = min(request.offset + request.limit, MAX_CANDIDATE_POOL)
|
||||
pool_size = max(CANDIDATE_POOL, window)
|
||||
|
||||
# 1. 按模式收集候选(FTS 与 Vector 各产出「按相关性降序」的 block_id 列表)
|
||||
fts_ranked: list[str] = []
|
||||
vec_ranked: list[str] = []
|
||||
fts_scores: dict[str, float] = {}
|
||||
vec_scores: dict[str, float] = {}
|
||||
fts_total = 0
|
||||
|
||||
if request.mode in (SearchMode.fts, SearchMode.hybrid):
|
||||
match = match_query(request.query)
|
||||
if match:
|
||||
fts_hits = repository.fts_search(match, CANDIDATE_POOL)
|
||||
fts_hits = repository.fts_search(match, pool_size)
|
||||
fts_ranked = [h.block_id for h in fts_hits]
|
||||
# bm25 越小越相关,取反后统一为「越大越相关」
|
||||
fts_scores = {h.block_id: -h.bm25 for h in fts_hits}
|
||||
if request.mode == SearchMode.fts:
|
||||
fts_total = repository.fts_count(match)
|
||||
|
||||
if request.mode in (SearchMode.vector, SearchMode.hybrid):
|
||||
query_vec = await self.embedding.embed_query(request.query)
|
||||
vec_hits = await self.vector_store.search(query_vec, top_k=CANDIDATE_POOL)
|
||||
vec_hits = await self.vector_store.search(query_vec, top_k=pool_size)
|
||||
vec_ranked = [v.id for v in vec_hits]
|
||||
vec_scores = {v.id: v.score for v in vec_hits}
|
||||
|
||||
@@ -95,7 +104,10 @@ class RetrievalEngine:
|
||||
|
||||
ordered = normalize_scores(ordered)
|
||||
|
||||
# 5. 分页
|
||||
# 5. 分页:fts 用真实命中总数;vector/hybrid 为 KNN 候选集,无全局 total
|
||||
if request.mode == SearchMode.fts:
|
||||
total = fts_total
|
||||
else:
|
||||
total = len(ordered)
|
||||
page = ordered[request.offset : request.offset + request.limit]
|
||||
items = [self._build_result(hits[block_id], request, score) for block_id, score in page]
|
||||
|
||||
@@ -32,16 +32,43 @@ def _safe_name(title: str) -> str:
|
||||
return name or "untitled"
|
||||
|
||||
|
||||
def _rel_path(folder: str | None, title: str) -> str:
|
||||
folder_part = folder.strip().strip("/") if folder else ""
|
||||
def _normalize_folder(folder: str | None) -> str:
|
||||
"""清洗 folder 为安全的相对目录,拒绝 `..`/`.`/绝对路径/盘符/空字节,防路径逃逸。"""
|
||||
if not folder:
|
||||
return ""
|
||||
if "\x00" in folder:
|
||||
raise ApiError(400, "INVALID_PATH", "folder must not contain NUL bytes", {"folder": folder})
|
||||
segments: list[str] = []
|
||||
for part in re.split(r"[\\/]+", folder):
|
||||
if part == "":
|
||||
continue
|
||||
if part in (".", ".."):
|
||||
raise ApiError(400, "INVALID_PATH", "folder must not contain '.' or '..'", {"folder": folder})
|
||||
if ":" in part:
|
||||
raise ApiError(400, "INVALID_PATH", "folder must be a relative path", {"folder": folder})
|
||||
segments.append(part)
|
||||
return "/".join(segments)
|
||||
|
||||
|
||||
def _rel_path(folder: str | None, title: str) -> tuple[str, str]:
|
||||
"""由 folder + title 生成安全的相对路径,返回 (rel_path, 清洗后的 folder)。"""
|
||||
clean_folder = _normalize_folder(folder)
|
||||
name = _safe_name(title)
|
||||
if not name.endswith(".md"):
|
||||
name += ".md"
|
||||
return f"{folder_part}/{name}" if folder_part else name
|
||||
rel = f"{clean_folder}/{name}" if clean_folder else name
|
||||
return rel, clean_folder
|
||||
|
||||
|
||||
def _abs_path(rel_path: str) -> Path:
|
||||
return _vault() / rel_path
|
||||
"""把相对路径解析为 Vault 内的绝对路径;越界即报 400,杜绝路径逃逸。"""
|
||||
if not rel_path or "\x00" in rel_path:
|
||||
raise ApiError(400, "INVALID_PATH", "invalid file path", {"file_path": rel_path})
|
||||
root = _vault().resolve()
|
||||
candidate = (_vault() / rel_path).resolve()
|
||||
if not candidate.is_relative_to(root):
|
||||
raise ApiError(400, "INVALID_PATH", "path escapes vault", {"file_path": rel_path})
|
||||
return candidate
|
||||
|
||||
|
||||
def _read_markdown(rel_path: str) -> str:
|
||||
@@ -62,9 +89,13 @@ def _delete_markdown(rel_path: str) -> None:
|
||||
|
||||
|
||||
async def index_note(parsed: ParsedNote) -> None:
|
||||
"""把解析结果写入元数据 + FTS5 + 向量(三层可重建索引)。"""
|
||||
"""把解析结果写入元数据 + FTS5 + 向量(三层可重建索引)。
|
||||
|
||||
替换元数据时拿到旧 block_id:清理已删除/内容变化的旧向量,只为新增 block 写向量,
|
||||
避免失效向量残留(内容未变的 block 其向量仍有效,无需重复写入)。
|
||||
"""
|
||||
vectors = await embedding.embed_documents([block.content for block in parsed.blocks])
|
||||
repository.replace_note_metadata(
|
||||
old_block_ids = repository.replace_note_metadata(
|
||||
note_id=parsed.note_id,
|
||||
title=parsed.title,
|
||||
file_path=parsed.file_path,
|
||||
@@ -74,24 +105,35 @@ async def index_note(parsed: ParsedNote) -> None:
|
||||
updated_at=parsed.updated_at,
|
||||
blocks=parsed.blocks,
|
||||
)
|
||||
old_ids = set(old_block_ids)
|
||||
new_ids = {block.block_id for block in parsed.blocks}
|
||||
stale_ids = [bid for bid in old_ids if bid not in new_ids]
|
||||
if stale_ids:
|
||||
await vector_store.delete(stale_ids)
|
||||
missing_ids = [bid for bid in new_ids if bid not in old_ids]
|
||||
records = [
|
||||
VectorRecord(id=block.block_id, vector=vector)
|
||||
for block, vector in zip(parsed.blocks, vectors)
|
||||
if block.block_id in missing_ids
|
||||
]
|
||||
await vector_store.upsert(records)
|
||||
repository.set_index_meta({"embedding_model": embedding.model_id, "embedding_dim": str(embedding.dim)})
|
||||
|
||||
|
||||
async def create_note(*, title: str, markdown: str, folder: str | None, tags: list[str]) -> Note:
|
||||
rel_path = _rel_path(folder, title)
|
||||
_write_markdown(rel_path, markdown)
|
||||
rel_path, clean_folder = _rel_path(folder, title)
|
||||
now = datetime.now(timezone.utc)
|
||||
parsed = parse_note(
|
||||
markdown=markdown, file_path=rel_path, folder=folder or "", tags=tags,
|
||||
markdown=markdown, file_path=rel_path, folder=clean_folder, tags=tags,
|
||||
created_at=now, updated_at=now,
|
||||
)
|
||||
parsed.title = title # 显式传入的 title 优先于正文推导(与 update_note 保持一致)
|
||||
_write_markdown(rel_path, markdown)
|
||||
try:
|
||||
await index_note(parsed)
|
||||
except BaseException:
|
||||
_delete_markdown(rel_path) # 索引失败时回滚,避免「文件已写、索引缺失」的部分提交
|
||||
raise
|
||||
return _build_note(parsed.note_id, parsed.title, parsed.file_path, parsed.tags,
|
||||
parsed.created_at, parsed.updated_at, parsed.blocks, markdown)
|
||||
|
||||
@@ -112,10 +154,12 @@ async def update_note(
|
||||
if record is None:
|
||||
raise ApiError(404, "RESOURCE_NOT_FOUND", "note not found", {"note_id": note_id})
|
||||
|
||||
new_md = _read_markdown(record.file_path) if markdown is None else markdown
|
||||
old_md = _read_markdown(record.file_path)
|
||||
new_md = old_md if markdown is None else markdown
|
||||
_write_markdown(record.file_path, new_md)
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
try:
|
||||
parsed = parse_note(
|
||||
markdown=new_md, file_path=record.file_path, folder=record.folder, tags=tags,
|
||||
created_at=record.created_at, updated_at=now,
|
||||
@@ -124,6 +168,9 @@ async def update_note(
|
||||
parsed.title = title # 显式传入的 title 覆盖正文推导结果
|
||||
|
||||
await index_note(parsed)
|
||||
except BaseException:
|
||||
_write_markdown(record.file_path, old_md) # 索引失败时回滚正文,避免部分提交
|
||||
raise
|
||||
return _build_note(parsed.note_id, parsed.title, parsed.file_path, parsed.tags,
|
||||
parsed.created_at, parsed.updated_at, parsed.blocks, new_md)
|
||||
|
||||
|
||||
@@ -244,3 +244,80 @@ def test_note_crud_roundtrip(vault) -> None:
|
||||
|
||||
assert asyncio.run(note_service.delete_note(note.note_id)) is True
|
||||
assert asyncio.run(note_service.get_note(note.note_id)) is None
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 审阅回归:路径逃逸 / 部分提交回滚 / 失效向量 / 搜索分页
|
||||
# --------------------------------------------------------------------------- #
|
||||
@pytest.mark.parametrize("folder", ["../../outside", "..", "..\\..\\etc", "C:\\Windows", "a/../b"])
|
||||
def test_create_note_rejects_path_traversal(vault, folder) -> None:
|
||||
from app.errors import ApiError
|
||||
from app.services import note_service
|
||||
|
||||
with pytest.raises(ApiError) as exc:
|
||||
asyncio.run(
|
||||
note_service.create_note(title="逃逸", markdown="# 逃逸", folder=folder, tags=[])
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
assert exc.value.code == "INVALID_PATH"
|
||||
|
||||
|
||||
def test_update_note_rolls_back_file_on_index_error(vault, monkeypatch) -> None:
|
||||
from app.services import note_service
|
||||
|
||||
note = asyncio.run(
|
||||
note_service.create_note(title="回滚", markdown="# 原文\n\n旧内容。", folder="", tags=[])
|
||||
)
|
||||
path = vault / note.file_path
|
||||
before = path.read_text(encoding="utf-8")
|
||||
|
||||
async def _boom(_contents):
|
||||
raise RuntimeError("embedding down")
|
||||
|
||||
monkeypatch.setattr(note_service.embedding, "embed_documents", _boom)
|
||||
with pytest.raises(RuntimeError):
|
||||
asyncio.run(note_service.update_note(note.note_id, markdown="# 新文\n\n新内容。"))
|
||||
|
||||
assert path.read_text(encoding="utf-8") == before # 文件已回滚,无部分提交
|
||||
|
||||
|
||||
def test_update_removes_stale_vectors(vault) -> None:
|
||||
from app.database.db import connect
|
||||
from app.services import note_service
|
||||
|
||||
def vec_count() -> int:
|
||||
conn = connect()
|
||||
try:
|
||||
return conn.execute("SELECT COUNT(*) FROM vec_blocks").fetchone()[0]
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
note = asyncio.run(
|
||||
note_service.create_note(
|
||||
title="向量清理", markdown="# 标题\n\n段落一。\n\n段落二。", folder="", tags=[]
|
||||
)
|
||||
)
|
||||
assert vec_count() == 3 # 标题 + 段落一 + 段落二
|
||||
|
||||
asyncio.run(note_service.update_note(note.note_id, markdown="# 标题\n\n段落一。"))
|
||||
assert vec_count() == 2 # 段落二的旧向量被清理,不再残留
|
||||
|
||||
|
||||
def test_search_pagination_total_reflects_all_matches(vault) -> None:
|
||||
from app.retrieval.engine import engine
|
||||
from app.services import index_service
|
||||
|
||||
body = "\n\n".join(f"第{i}段 内容。" for i in range(60))
|
||||
_write_vault(vault, {"多段.md": f"# 大量段落\n\n{body}"})
|
||||
asyncio.run(index_service.rebuild(IndexRebuildRequest(scope="all")))
|
||||
|
||||
page1 = asyncio.run(
|
||||
engine.search(SearchRequest(query="段", mode=SearchMode.fts, limit=10, offset=0))
|
||||
)
|
||||
assert page1.page.total >= 60 # total 反映真实命中数,而非候选池上限 50
|
||||
assert len(page1.items) == 10
|
||||
|
||||
page2 = asyncio.run(
|
||||
engine.search(SearchRequest(query="段", mode=SearchMode.fts, limit=10, offset=55))
|
||||
)
|
||||
assert page2.items # 跨过旧候选池边界仍能取到结果
|
||||
|
||||
Reference in New Issue
Block a user