diff --git a/backend/app/repository.py b/backend/app/repository.py index 3a4582c..ba3918c 100644 --- a/backend/app/repository.py +++ b/backend/app/repository.py @@ -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 [] diff --git a/backend/app/retrieval/engine.py b/backend/app/retrieval/engine.py index 8f28e47..0009b50 100644 --- a/backend/app/retrieval/engine.py +++ b/backend/app/retrieval/engine.py @@ -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,8 +104,11 @@ class RetrievalEngine: ordered = normalize_scores(ordered) - # 5. 分页 - total = len(ordered) + # 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] return SearchResponse( diff --git a/backend/app/services/note_service.py b/backend/app/services/note_service.py index d3e6753..e0feb19 100644 --- a/backend/app/services/note_service.py +++ b/backend/app/services/note_service.py @@ -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 保持一致) - await index_note(parsed) + _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,18 +154,23 @@ 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) - parsed = parse_note( - markdown=new_md, file_path=record.file_path, folder=record.folder, tags=tags, - created_at=record.created_at, updated_at=now, - ) - if title is not None: - parsed.title = title # 显式传入的 title 覆盖正文推导结果 + 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, + ) + if title is not None: + parsed.title = title # 显式传入的 title 覆盖正文推导结果 - await index_note(parsed) + 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) diff --git a/backend/tests/test_retrieval.py b/backend/tests/test_retrieval.py index 7639f2c..d6cb7ec 100644 --- a/backend/tests/test_retrieval.py +++ b/backend/tests/test_retrieval.py @@ -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 # 跨过旧候选池边界仍能取到结果