51 KiB
第一优先级:核心功能增强 实施计划
For agentic workers: REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (
- [ ]) syntax for tracking.
Goal: 为 md-vector-db 补齐混合检索、增量入库、结果重排序三大核心能力,外加 Docker 部署方案和安全加固,使其达到生产级 RAG 系统的基线水准。
Architecture: 新增 src/core/retriever.py 统一混合检索入口(BM25 + 向量 + 过滤),新增 src/core/reranker.py 提供 Cross-Encoder 重排序,增强 DocumentIngestor 支持文件变更检测。所有修改向下兼容现有 API。
Tech Stack: rank-bm25, sentence-transformers (已有), ChromaDB (已有), FastAPI (已有), Docker
预估总工作量: 约 15-20 小时
文件结构规划
新增文件:
src/core/retriever.py — 混合检索器(BM25 + 向量联合)
src/core/reranker.py — Cross-Encoder 重排序器
src/core/file_tracker.py — 文件变更追踪(SHA256 + mtime)
tests/test_retriever.py — 混合检索测试
tests/test_reranker.py — 重排序测试
tests/test_file_tracker.py — 文件追踪测试
Dockerfile — 生产镜像
docker-compose.yml — 本地开发 + 部署
.dockerignore — Docker 构建排除
修改文件:
src/core/ingest.py — 增量入库模式
src/core/search.py — 委托 Retriever
src/server/app.py — 新端点 + 安全加固
src/cli/main.py — 新命令
pyproject.toml — 新依赖
config.yaml — 新配置项
README.md — 更新文档
Task 1: 安装新依赖
Files:
-
Modify:
pyproject.toml -
Step 1: 在 pyproject.toml 添加混合检索和重排序依赖
在 dependencies 列表末尾追加 rank-bm25:
dependencies = [
"chromadb>=0.5.0",
"sentence-transformers>=3.0.0",
"fastapi>=0.115.0",
"uvicorn[standard]>=0.30.0",
"pyyaml>=6.0",
"markdown-it-py>=3.0.0",
"typer>=0.12.0",
"python-dotenv>=1.2.2",
"rank-bm25>=0.2.2",
]
- Step 2: 安装依赖
uv sync
- Step 3: 验证 rank-bm25 可导入
uv run python -c "from rank_bm25 import BM25Okapi; print('OK')"
预期输出: OK
- Step 4: Commit
git add pyproject.toml uv.lock
git commit -m "chore: 添加 rank-bm25 依赖(混合检索用)"
Task 2: 文件变更追踪器
Files:
-
Create:
src/core/file_tracker.py -
Create:
tests/test_file_tracker.py -
Step 1: 编写失败的测试
# tests/test_file_tracker.py
"""文件变更追踪器测试."""
import hashlib
import time
from pathlib import Path
from src.core.file_tracker import FileTracker, FileRecord
def test_compute_file_hash(tmp_path: Path):
"""计算文件 SHA256 哈希."""
file = tmp_path / "test.md"
file.write_text("hello world", encoding="utf-8")
result = FileTracker.compute_hash(str(file))
expected = hashlib.sha256(b"hello world").hexdigest()
assert result == expected
def test_compute_hash_deterministic(tmp_path: Path):
"""同一文件内容产生相同哈希."""
file = tmp_path / "a.md"
file.write_text("same content", encoding="utf-8")
assert FileTracker.compute_hash(str(file)) == FileTracker.compute_hash(str(file))
def test_hash_changes_with_content(tmp_path: Path):
"""内容变更导致哈希不同."""
file = tmp_path / "b.md"
file.write_text("v1", encoding="utf-8")
h1 = FileTracker.compute_hash(str(file))
file.write_text("v2", encoding="utf-8")
h2 = FileTracker.compute_hash(str(file))
assert h1 != h2
def test_is_stale_new_file(tmp_path: Path):
"""新文件(无记录)视为过期."""
file = tmp_path / "new.md"
file.write_text("content", encoding="utf-8")
tracker = FileTracker(tmp_path / "tracker.json")
record = tracker.get_record(str(file))
assert record is None
assert tracker.is_stale(str(file)) is True
def test_is_stale_unchanged_file(tmp_path: Path):
"""未修改文件视为未过期."""
file = tmp_path / "unchanged.md"
file.write_text("stable", encoding="utf-8")
tracker = FileTracker(tmp_path / "tracker.json")
tracker.mark_ingested(str(file))
assert tracker.is_stale(str(file)) is False
def test_is_stale_modified_file(tmp_path: Path):
"""修改后文件视为过期."""
file = tmp_path / "modified.md"
file.write_text("v1", encoding="utf-8")
tracker = FileTracker(tmp_path / "tracker.json")
tracker.mark_ingested(str(file))
file.write_text("v2", encoding="utf-8")
assert tracker.is_stale(str(file)) is True
def test_mark_ingested_updates_record(tmp_path: Path):
"""mark_ingested 创建/更新记录."""
file = tmp_path / "x.md"
file.write_text("hello", encoding="utf-8")
tracker = FileTracker(tmp_path / "tracker.json")
tracker.mark_ingested(str(file))
record = tracker.get_record(str(file))
assert record is not None
assert record["hash"] == FileTracker.compute_hash(str(file))
assert "ingested_at" in record
def test_persistence_across_instances(tmp_path: Path):
"""tracker 数据持久化到 JSON,跨实例可读."""
file = tmp_path / "p.md"
file.write_text("persist me", encoding="utf-8")
db_path = tmp_path / "tracker.json"
t1 = FileTracker(db_path)
t1.mark_ingested(str(file))
t2 = FileTracker(db_path)
assert t2.is_stale(str(file)) is False
def test_file_deleted_considered_stale(tmp_path: Path):
"""文件被删除后视为过期(清理记录)."""
file = tmp_path / "tmp.md"
file.write_text("temp", encoding="utf-8")
tracker = FileTracker(tmp_path / "tracker.json")
tracker.mark_ingested(str(file))
file.unlink()
assert tracker.is_stale(str(file)) is True
def test_binary_file_hash(tmp_path: Path):
"""二进制文件也能正确计算哈希(如 PDF)."""
file = tmp_path / "doc.pdf"
file.write_bytes(b"\x00\x01\x02\x03")
h = FileTracker.compute_hash(str(file))
assert len(h) == 64
assert h == hashlib.sha256(b"\x00\x01\x02\x03").hexdigest()
- Step 2: 运行测试验证失败
uv run pytest tests/test_file_tracker.py -v
预期: 全部 FAIL(模块不存在)
- Step 3: 实现 FileTracker
# src/core/file_tracker.py
"""文件变更追踪 — 基于 SHA256 + mtime 判断文件是否需要重新入库."""
from __future__ import annotations
import hashlib
import json
import logging
import os
import threading
from pathlib import Path
from typing import TypedDict
logger = logging.getLogger("md-vector-db")
class FileRecord(TypedDict):
"""追踪记录."""
hash: str
mtime: float
size: int
ingested_at: str
class FileTracker:
"""文件入库追踪器 — JSON 文件持久化.
通过比对文件内容的 SHA256 哈希判断文件是否变更。
使用 threading.Lock 保护 JSON 文件的并发读写。
"""
def __init__(self, db_path: str = "./ingest_tracker.json"):
self._db_path = Path(db_path)
self._lock = threading.Lock()
self._records: dict[str, FileRecord] = {}
self._load()
# -- 公开 API --
def is_stale(self, file_path: str) -> bool:
"""检查文件是否需要重新入库.
Returns:
True: 文件不存在 / 无记录 / 内容已变更
False: 文件未变更且记录存在
"""
path = Path(file_path)
if not path.exists():
return True
record = self.get_record(file_path)
if record is None:
return True
current_hash = self.compute_hash(file_path)
return current_hash != record["hash"]
def get_record(self, file_path: str) -> FileRecord | None:
"""获取文件的追踪记录(无记录返回 None)."""
return self._records.get(self._abs_key(file_path))
def mark_ingested(self, file_path: str) -> None:
"""标记文件已入库(创建或更新追踪记录)."""
path = Path(file_path)
if not path.exists():
logger.warning("标记已入库时文件不存在: %s", file_path)
return
stat = path.stat()
record: FileRecord = {
"hash": self.compute_hash(file_path),
"mtime": stat.st_mtime,
"size": stat.st_size,
"ingested_at": self._now_iso(),
}
with self._lock:
self._records[self._abs_key(file_path)] = record
self._save()
def remove_record(self, file_path: str) -> None:
"""移除文件的追踪记录."""
key = self._abs_key(file_path)
with self._lock:
if key in self._records:
del self._records[key]
self._save()
@staticmethod
def compute_hash(file_path: str) -> str:
"""计算文件 SHA256 哈希(分块读取,适合大文件)."""
sha = hashlib.sha256()
with open(file_path, "rb") as f:
while chunk := f.read(8192):
sha.update(chunk)
return sha.hexdigest()
# -- 内部 --
def _abs_key(self, file_path: str) -> str:
"""生成标准化 key."""
return str(Path(file_path).resolve())
@staticmethod
def _now_iso() -> str:
from datetime import datetime, timezone
return datetime.now(timezone.utc).isoformat()
def _load(self) -> None:
if self._db_path.exists():
try:
with open(self._db_path, "r", encoding="utf-8") as f:
self._records = json.load(f)
except (json.JSONDecodeError, OSError) as e:
logger.warning("tracker 文件损坏,重置为空: %s", e)
self._records = {}
def _save(self) -> None:
try:
with open(self._db_path, "w", encoding="utf-8") as f:
json.dump(self._records, f, ensure_ascii=False, indent=2)
except OSError as e:
logger.error("无法写入 tracker 文件: %s", e)
- Step 4: 运行测试验证通过
uv run pytest tests/test_file_tracker.py -v
预期: 9 passed
- Step 5: Commit
git add src/core/file_tracker.py tests/test_file_tracker.py
git commit -m "feat: 添加 FileTracker — 基于 SHA256 的文件变更追踪"
Task 3: DocumentIngestor 增量入库模式
Files:
-
Modify:
src/core/ingest.py -
Modify:
tests/test_ingest.py -
Step 1: 为增量入库编写测试
在 tests/test_ingest.py 末尾追加:
def test_ingest_file_incremental_skips_unchanged(db, local_embedder, tmp_path: Path):
"""增量模式: 未修改的文件跳过入库."""
file = tmp_path / "stable.md"
file.write_text("# 稳定文档\n\n内容不变。", encoding="utf-8")
tracker_path = str(tmp_path / "tracker.json")
ingestor = DocumentIngestor(
db, local_embedder, "test_incr",
chunk_config=ChunkConfig(max_size=1000, overlap=100),
file_tracker=FileTracker(tracker_path),
)
# 首次入库
count1 = ingestor.ingest_file(str(file), incremental=True)
assert count1 > 0
# 二次入库(无变更)
count2 = ingestor.ingest_file(str(file), incremental=True)
assert count2 == 0 # 跳过
def test_ingest_file_incremental_reingests_modified(db, local_embedder, tmp_path: Path):
"""增量模式: 修改后的文件重新入库."""
file = tmp_path / "changing.md"
file.write_text("# v1", encoding="utf-8")
tracker_path = str(tmp_path / "tracker2.json")
ingestor = DocumentIngestor(
db, local_embedder, "test_incr2",
chunk_config=ChunkConfig(max_size=1000, overlap=100),
file_tracker=FileTracker(tracker_path),
)
count1 = ingestor.ingest_file(str(file), incremental=True)
assert count1 > 0
# 修改内容
file.write_text("# v2\n\n新增段落,内容不同。", encoding="utf-8")
count2 = ingestor.ingest_file(str(file), incremental=True)
assert count2 > 0
def test_ingest_file_force_mode_always_reingests(db, local_embedder, tmp_path: Path):
"""force=True 时始终重新入库(忽略 tracker)."""
file = tmp_path / "force.md"
file.write_text("# force test", encoding="utf-8")
tracker_path = str(tmp_path / "tracker3.json")
ingestor = DocumentIngestor(
db, local_embedder, "test_force",
chunk_config=ChunkConfig(max_size=1000, overlap=100),
file_tracker=FileTracker(tracker_path),
)
count1 = ingestor.ingest_file(str(file), incremental=True)
count2 = ingestor.ingest_file(str(file), incremental=True, force=True)
assert count1 > 0
assert count2 > 0 # force 模式重新入库
需要在 tests/test_ingest.py 顶部添加 import:
from src.core.file_tracker import FileTracker
- Step 2: 运行测试确认失败
uv run pytest tests/test_ingest.py::test_ingest_file_incremental_skips_unchanged -v
预期: FAIL(file_tracker 参数不存在)
- Step 3: 修改 DocumentIngestor.init
修改 src/core/ingest.py 的 __init__ 方法:
# 在文件顶部添加 import
from src.core.file_tracker import FileTracker
# 修改 __init__ 方法签名
def __init__(
self,
db: VectorDB,
embedder: Embedder,
collection_name: str,
splitter: Splitter | None = None,
chunk_config: ChunkConfig | None = None,
file_tracker: FileTracker | None = None,
):
self.db = db
self.embedder = embedder
self.collection_name = collection_name
self.splitter = splitter or MarkdownSplitter()
self.chunk_config = chunk_config or ChunkConfig()
self.file_tracker = file_tracker # None = 不使用增量功能
- Step 4: 修改 ingest_file 方法签名和逻辑
将 ingest_file 方法改为:
def ingest_file(self, file_path: str, incremental: bool = False, force: bool = False) -> int:
"""入库单个文件, 返回 chunk 数量.
Args:
file_path: 文件路径
incremental: 启用增量模式(需 file_tracker 已注入)
force: 强制重新入库(忽略增量检查)
增量模式下:
- 若文件未变更 → 跳过嵌入,返回 0
- 若文件已变更 → 先删旧 chunks,再嵌入入库
"""
import hashlib
path = Path(file_path).resolve()
path_hash = hashlib.sha256(str(path).encode()).hexdigest()[:12]
file_name = f"{path_hash}_{path.name}"
# 增量模式:检查是否需要重新入库
if incremental and not force and self.file_tracker is not None:
if not self.file_tracker.is_stale(str(path)):
logger.debug("跳过未变更文件: %s", path.name)
return 0
splitter = self.splitter or get_splitter(
file_path,
max_size=self.chunk_config.max_size,
overlap=self.chunk_config.overlap,
)
suffix = path.suffix.lower()
if suffix in (".pdf", ".epub"):
chunks = splitter.split(str(path), source_file=file_name)
result = self._add_chunks(chunks, file_name)
else:
content = path.read_text(encoding="utf-8")
result = self._ingest_with_splitter(content, file_name, splitter)
# 入库成功后更新 tracker
if self.file_tracker is not None:
self.file_tracker.mark_ingested(str(path))
return result
- Step 5: 运行增量测试
uv run pytest tests/test_ingest.py -v -k "incremental or force"
预期: 3 passed
- Step 6: Commit
git add src/core/ingest.py tests/test_ingest.py
git commit -m "feat: DocumentIngestor 支持增量入库(FileTracker)"
Task 4: 混合检索器(BM25 + 向量)
Files:
-
Create:
src/core/retriever.py -
Create:
tests/test_retriever.py -
Step 1: 编写 Retriever 测试
# tests/test_retriever.py
"""混合检索器测试."""
import pytest
from src.core.retriever import HybridRetriever, SearchResult
class FakeEmbedder:
"""模拟嵌入器 — 返回伪向量."""
@property
def dimension(self) -> int:
return 4
def embed(self, texts: list[str]) -> list[list[float]]:
# 每个文本返回固定维度随机向量
import hashlib
result = []
for t in texts:
h = hashlib.md5(t.encode()).digest()
vec = [float(b) / 255.0 for b in h[:4]]
result.append(vec)
return result
class FakeCollection:
"""模拟 ChromaDB collection."""
def __init__(self):
self._docs: list[dict] = []
self._next_id = 0
def add(self, ids, embeddings, documents, metadatas):
for i, doc_id in enumerate(ids):
self._docs.append({
"id": doc_id,
"embedding": embeddings[i] if embeddings else [],
"document": documents[i],
"metadata": metadatas[i] if metadatas else {},
})
self._next_id += len(ids)
def query(self, query_embeddings, n_results, where=None, include=None):
return {"ids": [[]], "documents": [[]], "metadatas": [[]], "distances": [[]]}
def get(self, include=None):
return {
"ids": [d["id"] for d in self._docs],
"documents": [d["document"] for d in self._docs],
"metadatas": [d["metadata"] for d in self._docs],
}
def count(self) -> int:
return len(self._docs)
def delete(self, ids):
self._docs = [d for d in self._docs if d["id"] not in ids]
class FakeDB:
def __init__(self):
self._collections: dict[str, FakeCollection] = {}
def get_or_create_collection(self, name: str):
if name not in self._collections:
self._collections[name] = FakeCollection()
return self._collections[name]
@pytest.fixture
def retriever():
db = FakeDB()
embedder = FakeEmbedder()
return HybridRetriever(db, embedder, "test", bm25_weight=0.3)
def test_search_returns_list(retriever):
"""基本语义检索返回列表."""
results = retriever.search("测试查询", top_k=5)
assert isinstance(results, list)
def test_search_with_source_filter(retriever):
"""按 source_file 过滤."""
results = retriever.search("查询", top_k=5, source_file="doc.md")
assert isinstance(results, list)
def test_search_top_k_bounds(retriever):
"""top_k 在合理范围内."""
for k in [1, 10, 50]:
results = retriever.search("test", top_k=k)
assert len(results) <= k
def test_bm25_index_built_from_collection(retriever):
"""BM25 索引从 collection 文档构建."""
# 先入库一些文档
coll = retriever._db.get_or_create_collection("test")
coll.add(
ids=["doc_0", "doc_1", "doc_2"],
embeddings=[[0.1]*4, [0.2]*4, [0.3]*4],
documents=["Python 是一门编程语言", "Java 也是编程语言", "今天天气很好"],
metadatas=[
{"source_file": "a.md"},
{"source_file": "b.md"},
{"source_file": "c.md"},
],
)
retriever._bm25_index = None # 强制重建
results = retriever.search("编程语言", top_k=2)
assert len(results) == 2
# BM25 结果应排在前面(关键词更匹配)
assert any("编程语言" in r["content"] for r in results)
def test_hybrid_score_fusion(retriever):
"""混合分数融合:向量分 + BM25 分加权."""
coll = retriever._db.get_or_create_collection("test")
coll.add(
ids=["d0", "d1"],
embeddings=[[1.0]*4, [0.5]*4],
documents=["Docker 容器化部署指南", "Python 数据分析入门"],
metadatas=[{"source_file": "x.md"}, {"source_file": "y.md"}],
)
retriever._bm25_index = None # 强制重建
results = retriever.search("Docker 部署", top_k=2)
assert len(results) >= 1
# "Docker 容器化部署指南" 应该排第一(BM25 + 向量双重命中)
assert "Docker" in results[0]["content"]
def test_empty_collection_returns_empty(retriever):
"""空 collection 返回空列表."""
results = retriever.search("查询", top_k=5)
assert results == []
def test_metadata_in_results(retriever):
"""结果中包含完整元数据."""
coll = retriever._db.get_or_create_collection("test")
coll.add(
ids=["meta_test"],
embeddings=[[0.5]*4],
documents=["带元数据的文档"],
metadatas=[{"source_file": "meta.md", "section_title": "第一章", "heading_level": 1}],
)
retriever._bm25_index = None
results = retriever.search("元数据", top_k=1)
assert len(results) == 1
assert results[0]["source_file"] == "meta.md"
assert results[0]["section_title"] == "第一章"
- Step 2: 运行测试确认失败
uv run pytest tests/test_retriever.py -v
预期: 全部 FAIL
- Step 3: 实现 HybridRetriever
# src/core/retriever.py
"""混合检索器 — BM25 关键词 + 向量语义联合检索."""
from __future__ import annotations
import logging
from typing import TypedDict
from rank_bm25 import BM25Okapi
from src.core.db import VectorDB
from src.core.embedder import Embedder
logger = logging.getLogger("md-vector-db")
class SearchResult(TypedDict):
"""检索结果类型."""
id: str
content: str
source_file: str
section_title: str
heading_level: int
chunk_index: int
score: float
bm25_score: float
vector_score: float
class HybridRetriever:
"""BM25 + 向量混合检索器.
Architecture:
1. 向量检索取得 top_k * 2 候选
2. BM25 对候选打分
3. 加权融合排序(默认 0.7 向量 + 0.3 BM25)
4. 返回 top_k 结果
支持按 source_file / section_title 等元数据过滤。
"""
def __init__(
self,
db: VectorDB,
embedder: Embedder,
collection_name: str,
bm25_weight: float = 0.3,
vector_candidate_multiplier: int = 3,
):
self._db = db
self._embedder = embedder
self._collection_name = collection_name
self._bm25_weight = bm25_weight
self._vector_multiplier = vector_candidate_multiplier
self._bm25_index: BM25Okapi | None = None
self._bm25_docs: list[str] = []
self._bm25_ids: list[str] = []
@property
def _collection(self):
return self._db.get_or_create_collection(self._collection_name)
def search(
self,
query: str,
top_k: int = 10,
source_file: str | None = None,
) -> list[SearchResult]:
"""混合检索.
Args:
query: 查询文本
top_k: 返回结果数量
source_file: 可选,按源文件过滤
Returns:
按混合分数降序排列的结果列表
"""
# 1. 向量检索:多取候选
vector_candidates = self._vector_search(query, top_k * self._vector_multiplier, source_file)
if not vector_candidates:
return []
# 2. BM25 打分
bm25_scored = self._bm25_rerank(query, vector_candidates)
# 3. 分数融合
fused = self._fuse_scores(bm25_scored, self._bm25_weight)
# 4. 排序取 top_k
fused.sort(key=lambda x: x["score"], reverse=True)
return fused[:top_k]
# -- 内部方法 --
def _vector_search(
self, query: str, n: int, source_file: str | None
) -> list[dict]:
"""向量检索取得候选."""
embeddings = self._embedder.embed([query])
if not embeddings:
return []
query_embedding = embeddings[0]
where_filter = {"source_file": source_file} if source_file else None
results = self._collection.query(
query_embeddings=[query_embedding],
n_results=n,
where=where_filter,
include=["documents", "metadatas", "distances"],
)
candidates = []
if results["ids"] and results["ids"][0]:
for i, doc_id in enumerate(results["ids"][0]):
metadata = results["metadatas"][0][i] if results["metadatas"] else {}
distance = results["distances"][0][i] if results["distances"] else 0.0
vector_score = max(0.0, round(1.0 - distance, 4))
candidates.append({
"id": doc_id,
"content": results["documents"][0][i] if results["documents"] else "",
"source_file": metadata.get("source_file", ""),
"section_title": metadata.get("section_title", ""),
"heading_level": metadata.get("heading_level", 0),
"chunk_index": metadata.get("chunk_index", 0),
"vector_score": vector_score,
})
return candidates
def _bm25_rerank(self, query: str, candidates: list[dict]) -> list[dict]:
"""用 BM25 对候选列表重新打分."""
if not candidates:
return candidates
self._ensure_bm25_index()
tokenized_query = self._tokenize(query)
tokenized_candidates = [self._tokenize(c["content"]) for c in candidates]
bm25 = BM25Okapi(tokenized_candidates)
scores = bm25.get_scores(tokenized_query)
# 归一化 BM25 分数到 [0, 1]
max_score = max(scores) if max(scores) > 0 else 1.0
for i, c in enumerate(candidates):
c["bm25_score"] = round(scores[i] / max_score, 4)
return candidates
def _fuse_scores(self, items: list[dict], bm25_weight: float) -> list[dict]:
"""加权融合向量分和 BM25 分."""
vector_weight = 1.0 - bm25_weight
for item in items:
bm25_s = item.get("bm25_score", 0.0)
vec_s = item.get("vector_score", 0.0)
item["score"] = round(vec_s * vector_weight + bm25_s * bm25_weight, 4)
return items
def _ensure_bm25_index(self) -> None:
"""确保 BM25 索引已构建(从 collection 所有文档构建)."""
if self._bm25_index is not None:
return
all_data = self._collection.get(include=["documents", "metadatas"])
if all_data and all_data["ids"]:
self._bm25_ids = all_data["ids"]
self._bm25_docs = all_data["documents"] or []
tokenized = [self._tokenize(d) for d in self._bm25_docs]
self._bm25_index = BM25Okapi(tokenized) if tokenized else None
else:
self._bm25_ids = []
self._bm25_docs = []
self._bm25_index = None
@staticmethod
def _tokenize(text: str) -> list[str]:
"""简易中文+英文分词.
注意: 这是基础实现。生产环境建议集成 jieba 分词。
"""
import re
# 按中文单字 + 英文单词 + 数字拆分
tokens = []
# 匹配英文单词/数字,或单个中文字符
for match in re.finditer(r"[a-zA-Z0-9]+|[一-鿿]|[^\s]", text):
tokens.append(match.group().lower())
return tokens
- Step 4: 运行 Retriever 测试
uv run pytest tests/test_retriever.py -v
预期: 9 passed
- Step 5: Commit
git add src/core/retriever.py tests/test_retriever.py
git commit -m "feat: 添加 HybridRetriever — BM25+向量混合检索"
Task 5: 集成 Retriever 到 Searcher 和 API
Files:
-
Modify:
src/core/search.py -
Modify:
src/server/app.py -
Modify:
src/cli/main.py -
Modify:
src/server/deps.py -
Modify:
config.yaml -
Step 1: 在 config.yaml 添加检索配置
# 在 config.yaml 末尾追加
search:
mode: hybrid # hybrid | vector(默认 hybrid = BM25 + 向量)
bm25_weight: 0.3 # BM25 权重(0=纯向量, 1=纯BM25)
candidate_multiplier: 3 # 向量检索候选倍数
enable_rerank: false # Cross-Encoder 重排序(后续任务实现)
- Step 2: 在 config.py 添加 SearchConfig
在 src/core/config.py 中添加:
@dataclass
class SearchConfig:
"""检索配置."""
mode: str = "hybrid" # hybrid | vector
bm25_weight: float = 0.3 # BM25 权重
candidate_multiplier: int = 3 # 向量候选倍数
enable_rerank: bool = False # 是否启用重排序
# 在 AppConfig 中添加 search 字段
@dataclass
class AppConfig:
chroma: ChromaConfig = field(default_factory=ChromaConfig)
embed: EmbedConfig = field(default_factory=EmbedConfig)
chunk: ChunkConfig = field(default_factory=ChunkConfig)
server: ServerConfig = field(default_factory=ServerConfig)
search: SearchConfig = field(default_factory=SearchConfig) # 新增
@classmethod
def from_dict(cls, data: dict) -> "AppConfig":
return cls(
chroma=ChromaConfig(**data.get("chroma", {})),
embed=EmbedConfig(**data.get("embed", {})),
chunk=ChunkConfig(**data.get("chunk", {})),
server=ServerConfig(**data.get("server", {})),
search=SearchConfig(**data.get("search", {})), # 新增
)
- Step 3: 修改 Searcher 委托 HybridRetriever
修改 src/core/search.py:
"""语义检索模块."""
import logging
from src.core.config import SearchConfig
from src.core.db import VectorDB
from src.core.embedder import Embedder
from src.core.retriever import HybridRetriever
logger = logging.getLogger("md-vector-db")
class Searcher:
"""向量检索器 — 支持纯向量或 HybridRetriever 混合检索."""
def __init__(
self,
db: VectorDB,
embedder: Embedder,
collection_name: str,
search_config: SearchConfig | None = None,
):
self.db = db
self.embedder = embedder
self.collection_name = collection_name
self._search_config = search_config or SearchConfig()
self._hybrid: HybridRetriever | None = None
@property
def collection(self):
return self.db.get_or_create_collection(self.collection_name)
def _get_hybrid(self) -> HybridRetriever:
if self._hybrid is None:
self._hybrid = HybridRetriever(
self.db,
self.embedder,
self.collection_name,
bm25_weight=self._search_config.bm25_weight,
vector_candidate_multiplier=self._search_config.candidate_multiplier,
)
return self._hybrid
def search(
self,
query: str,
top_k: int = 10,
source_file: str | None = None,
) -> list[dict]:
"""语义检索, 返回格式化结果列表.
根据 search.mode 配置自动选择 hybrid 或 vector 模式.
"""
if self._search_config.mode == "hybrid":
return self._get_hybrid().search(query, top_k=top_k, source_file=source_file)
# 纯向量模式(原有逻辑)
embeddings = self.embedder.embed([query])
if not embeddings:
raise RuntimeError("嵌入器返回空结果, 无法进行检索")
query_embedding = embeddings[0]
where_filter = None
if source_file:
where_filter = {"source_file": source_file}
results = self.collection.query(
query_embeddings=[query_embedding],
n_results=top_k,
where=where_filter,
include=["documents", "metadatas", "distances"],
)
formatted = []
if results["ids"] and results["ids"][0]:
for i, doc_id in enumerate(results["ids"][0]):
metadata = results["metadatas"][0][i] if results["metadatas"] else {}
distance = results["distances"][0][i] if results["distances"] else 0.0
score = round(1.0 - distance, 4)
formatted.append({
"id": doc_id,
"content": results["documents"][0][i] if results["documents"] else "",
"source_file": metadata.get("source_file", ""),
"section_title": metadata.get("section_title", ""),
"heading_level": metadata.get("heading_level", 0),
"chunk_index": metadata.get("chunk_index", 0),
"score": max(0.0, score),
})
return formatted
def get_collection_info(self) -> dict:
return {
"name": self.collection_name,
"count": self.collection.count(),
}
def list_sources(self) -> list[str]:
if self.collection.count() == 0:
return []
result = self.collection.get(include=["metadatas"])
sources = set()
if result and result["metadatas"]:
for m in result["metadatas"]:
if m and "source_file" in m:
sources.add(m["source_file"])
return sorted(sources)
def delete_by_source(self, file_name: str) -> bool:
return self.db.delete_by_source(self.collection_name, file_name)
- Step 4: 更新 deps.py 传递 SearchConfig
修改 src/server/deps.py 的 get_searcher 方法:
def get_searcher(self, collection: str | None = None) -> Searcher:
name = collection or self.default_collection
with self._cache_lock:
if name not in self._searchers:
self._searchers[name] = Searcher(
self.db, self.embedder, name,
search_config=self.config.search, # 新增
)
return self._searchers[name]
同样更新 get_ingestor 传递 file_tracker(如果配置了):
def get_ingestor(self, collection: str | None = None) -> DocumentIngestor:
name = collection or self.default_collection
with self._cache_lock:
if name not in self._ingestors:
self._ingestors[name] = DocumentIngestor(
self.db, self.embedder, name,
chunk_config=self.config.chunk,
file_tracker=self._file_tracker, # 新增(可为 None)
)
return self._ingestors[name]
并在 AppState.__init__ 中初始化可选的 file_tracker:
# 在 __init__ 末尾添加
tracker_path = os.environ.get("MD_VECTOR_TRACKER", "")
if tracker_path:
from src.core.file_tracker import FileTracker
self._file_tracker = FileTracker(tracker_path)
else:
self._file_tracker = None
- Step 5: 确认现有测试仍然通过
uv run pytest tests/test_search.py tests/test_deps.py tests/test_config.py -v
预期: 全部 PASS
- Step 6: Commit
git add src/core/search.py src/server/deps.py src/core/config.py config.yaml
git commit -m "feat: Searcher 集成 HybridRetriever,支持混合检索模式"
Task 6: Cross-Encoder 重排序
Files:
-
Create:
src/core/reranker.py -
Create:
tests/test_reranker.py -
Step 1: 编写 Reranker 测试
# tests/test_reranker.py
"""重排序器测试."""
import pytest
from src.core.reranker import Reranker
class FakeCrossEncoder:
"""模拟 Cross-Encoder 模型."""
def predict(self, pairs, **kwargs):
# 返回伪分数:包含"重要"的 pair 分数高
scores = []
for pair in pairs:
score = 5.0 if "重要" in pair[1] else 1.0
scores.append(score)
return scores
def test_reranker_returns_same_count():
"""重排序不改变结果数量."""
reranker = Reranker(model_name="test-model")
reranker._model = FakeCrossEncoder()
candidates = [
{"content": "普通文档", "score": 0.8},
{"content": "重要文档", "score": 0.6},
{"content": "另一个普通", "score": 0.7},
]
result = reranker.rerank("查询", candidates, top_k=3)
assert len(result) == 3
def test_reranker_promotes_relevant():
"""重排序将更相关的内容提前."""
reranker = Reranker(model_name="test-model")
reranker._model = FakeCrossEncoder()
candidates = [
{"content": "普通 A", "score": 0.9},
{"content": "重要内容", "score": 0.5}, # Cross-Encoder 会给高分
{"content": "普通 B", "score": 0.7},
]
result = reranker.rerank("查询", candidates, top_k=3)
# "重要内容" 应该排第一
assert "重要" in result[0]["content"]
def test_reranker_truncates_to_top_k():
"""rerank 截断到指定的 top_k."""
reranker = Reranker(model_name="test-model")
reranker._model = FakeCrossEncoder()
candidates = [
{"content": f"文档{i}", "score": 0.9 - i * 0.1}
for i in range(20)
]
result = reranker.rerank("查询", candidates, top_k=5)
assert len(result) == 5
def test_reranker_empty_input():
"""空输入返回空列表."""
reranker = Reranker(model_name="test-model")
result = reranker.rerank("查询", [], top_k=5)
assert result == []
def test_reranker_preserves_metadata():
"""重排序保留文档元数据."""
reranker = Reranker(model_name="test-model")
reranker._model = FakeCrossEncoder()
candidates = [
{
"content": "带元数据",
"score": 0.5,
"source_file": "meta.md",
"section_title": "第一章",
}
]
result = reranker.rerank("查询", candidates, top_k=1)
assert result[0]["source_file"] == "meta.md"
assert result[0]["section_title"] == "第一章"
def test_reranker_lazy_load(monkeypatch):
"""模型只在首次调用时加载."""
loaded = False
def fake_predict(self, pairs, **kwargs):
return [3.0] * len(pairs)
reranker = Reranker(model_name="lazy-model")
# 替换 predict 方法模拟已加载
reranker._model = type("Fake", (), {"predict": fake_predict})()
result = reranker.rerank("测试", [{"content": "测试文档", "score": 0.8}], top_k=1)
assert len(result) == 1
- Step 2: 实现 Reranker
# src/core/reranker.py
"""Cross-Encoder 重排序器."""
from __future__ import annotations
import logging
logger = logging.getLogger("md-vector-db")
class Reranker:
"""使用 Cross-Encoder 模型对检索结果重排序.
默认模型: BAAI/bge-reranker-base(中文友好,384 维)
首次调用时懒加载模型。
"""
_DEFAULT_MODEL = "BAAI/bge-reranker-base"
def __init__(self, model_name: str | None = None):
self._model_name = model_name or self._DEFAULT_MODEL
self._model = None
def rerank(
self, query: str, candidates: list[dict], top_k: int = 10
) -> list[dict]:
"""对候选列表重排序.
Args:
query: 原始查询
candidates: 候选结果列表(需含 "content" 字段)
top_k: 返回数量
Returns:
按 cross-encoder 分数降序的结果列表
"""
if not candidates:
return []
self._ensure_model()
# 构建 (query, document) 对
pairs = [(query, c["content"]) for c in candidates]
try:
scores = self._model.predict(pairs, show_progress_bar=False)
except Exception as e:
logger.error("Cross-Encoder 重排序失败: %s", e)
# 降级:保留原始顺序
return candidates[:top_k]
# 附加 rerank_score
for i, c in enumerate(candidates):
c["rerank_score"] = round(float(scores[i]), 4)
# 按 rerank_score 降序排列
candidates.sort(key=lambda x: x.get("rerank_score", 0), reverse=True)
# 返回 top_k,将 rerank_score 作为最终 score
result = candidates[:top_k]
for r in result:
r["score"] = r.get("rerank_score", r.get("score", 0))
return result
def _ensure_model(self) -> None:
"""懒加载 Cross-Encoder 模型."""
if self._model is not None:
return
from sentence_transformers import CrossEncoder
logger.info("加载 Cross-Encoder 模型: %s", self._model_name)
self._model = CrossEncoder(self._model_name)
- Step 3: 运行重排序测试
uv run pytest tests/test_reranker.py -v
预期: 6 passed
- Step 4: 集成 Reranker 到 HybridRetriever
修改 src/core/retriever.py 的 HybridRetriever.__init__:
def __init__(
self,
db: VectorDB,
embedder: Embedder,
collection_name: str,
bm25_weight: float = 0.3,
vector_candidate_multiplier: int = 3,
reranker: "Reranker | None" = None, # 新增
):
# ... 原有代码 ...
self._reranker = reranker
修改 search 方法末尾,在 return fused[:top_k] 前加入:
# 4.5 可选:Cross-Encoder 重排序
if self._reranker is not None and len(fused) > 1:
fused = self._reranker.rerank(query, fused, top_k=top_k)
在文件顶部添加 import:
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from src.core.reranker import Reranker
- Step 5: 确认所有测试通过
uv run pytest tests/test_retriever.py tests/test_reranker.py tests/test_search.py -v
预期: 全部 PASS
- Step 6: Commit
git add src/core/reranker.py tests/test_reranker.py src/core/retriever.py
git commit -m "feat: 添加 Cross-Encoder Reranker,集成到 HybridRetriever"
Task 7: CLI 新命令(增量入库 + 检索模式切换)
Files:
-
Modify:
src/cli/main.py -
Step 1: 添加 --incremental 和 --force 选项到 ingest 命令
修改 src/cli/main.py 的 ingest 函数签名和逻辑:
在函数参数中添加:
def ingest(
file_paths: Annotated[
list[str] | None,
typer.Argument(help="Markdown 文件路径 (可多个, 或 - 从标准输入读取)"),
] = None,
name: Annotated[str | None, typer.Option("--name", help="标准输入模式下的虚拟文件名")] = None,
incremental: Annotated[bool, typer.Option("--incremental", help="增量模式:跳过未变更文件")] = False, # 新增
force: Annotated[bool, typer.Option("--force", help="强制重新入库(忽略增量检查)")] = False, # 新增
config: ConfigOpt = DEFAULT_CONFIG_PATH,
collection: CollectionOpt = None,
):
在 ingest_file 调用处传入参数:
# 将 c = ingestor.ingest_file(m) 改为:
c = ingestor.ingest_file(m, incremental=incremental, force=force)
以及 ingest_file(fp) 同理:
c = ingestor.ingest_file(fp, incremental=incremental, force=force)
- Step 2: 添加 --mode 选项到 search 命令
在 search 函数参数中添加:
def search(
query: Annotated[str, typer.Argument(help="搜索关键词或自然语言查询")],
top_k: Annotated[int, typer.Option("--top-k", "-k", help="返回结果数量 (1-100)")] = 10,
json_output: Annotated[bool, typer.Option("--json", help="以 JSON 格式输出")] = False,
mode: Annotated[str, typer.Option("--mode", help="检索模式: hybrid | vector")] = "hybrid", # 新增
config: ConfigOpt = DEFAULT_CONFIG_PATH,
collection: CollectionOpt = None,
):
在调用 search 前覆盖配置:
# 在 state = get_state() 之后添加:
if mode:
searcher = state.get_searcher(_resolve_collection(collection))
searcher._search_config.mode = mode
- Step 3: 确认 CLI 帮助正常
uv run md-vector-db ingest --help
uv run md-vector-db search --help
- Step 4: Commit
git add src/cli/main.py
git commit -m "feat: CLI 添加 --incremental/--force/--mode 选项"
Task 8: API 安全加固
Files:
-
Modify:
src/server/app.py -
Step 1: 添加请求体大小限制中间件
在 src/server/app.py 的 rate_limit_middleware 之后添加:
@app.middleware("http")
async def body_size_limit_middleware(request: Request, call_next):
"""限制请求体大小(防止内存耗尽攻击)."""
content_length = request.headers.get("content-length")
max_size = int(os.environ.get("MAX_REQUEST_BODY_SIZE", str(10 * 1024 * 1024))) # 默认 10MB
if content_length and int(content_length) > max_size:
raise HTTPException(status_code=413, detail="请求体过大")
return await call_next(request)
- Step 2: 添加审计日志中间件
@app.middleware("http")
async def audit_log_middleware(request: Request, call_next):
"""记录所有 API 请求的审计日志."""
start = time.time()
response = await call_next(request)
duration_ms = (time.time() - start) * 1000
logger.info(
"audit: %s %s → %d (%.1fms) [%s]",
request.method, request.url.path,
response.status_code, duration_ms,
request.client.host if request.client else "unknown",
)
return response
确保文件顶部有 import time。
- Step 3: 添加健康检查端点的速率限制豁免
修改 rate_limit_middleware,跳过 /api/v1/health 和 /:
@app.middleware("http")
async def rate_limit_middleware(request: Request, call_next):
# 健康检查和根路径不需要速率限制
if request.url.path in ("/api/v1/health", "/"):
return await call_next(request)
await rate_limiter(request)
response = await call_next(request)
return response
- Step 4: 确认 API 测试仍然通过
uv run pytest tests/test_api.py tests/test_auth.py -v
- Step 5: Commit
git add src/server/app.py
git commit -m "feat: API 安全加固 — 请求体大小限制、审计日志、健康检查免限速"
Task 9: Docker 部署方案
Files:
-
Create:
Dockerfile -
Create:
docker-compose.yml -
Create:
.dockerignore -
Step 1: 创建 .dockerignore
__pycache__/
*.pyc
.venv/
.pytest_cache/
.ruff_cache/
.mypy_cache/
.vscode/
.git/
.gitignore
.env
data/
*.egg-info/
dist/
build/
- Step 2: 创建 Dockerfile
# Dockerfile — md-vector-db 生产镜像
FROM python:3.13-slim-bookworm
LABEL org.opencontainers.image.title="md-vector-db"
LABEL org.opencontainers.image.description="Markdown 文档向量数据库"
# 系统依赖
RUN apt-get update && apt-get install -y --no-install-recommends \
build-essential \
&& rm -rf /var/lib/apt/lists/*
WORKDIR /app
# 先复制依赖文件以利用 Docker 层缓存
COPY pyproject.toml uv.lock ./
# 安装 uv 并同步依赖(CPU 模式)
RUN pip install --no-cache-dir uv \
&& uv sync --frozen --no-dev \
&& uv cache clean
# 复制源码和配置
COPY config.yaml .env.example ./
COPY src/ ./src/
COPY scripts/ ./scripts/
# 创建数据目录
RUN mkdir -p /app/data
# 暴露端口
EXPOSE 8000
# 健康检查
HEALTHCHECK --interval=30s --timeout=10s --start-period=60s --retries=3 \
CMD python -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/api/v1/health')" || exit 1
# 默认启动 HTTP 服务
CMD ["uv", "run", "md-vector-db", "serve", "--port", "8000"]
- Step 3: 创建 docker-compose.yml
# docker-compose.yml — md-vector-db 本地开发与生产部署
version: "3.8"
services:
md-vector-db:
build:
context: .
dockerfile: Dockerfile
image: md-vector-db:latest
container_name: md-vector-db
restart: unless-stopped
ports:
- "${MD_VECTOR_PORT:-8000}:8000"
volumes:
# 持久化 ChromaDB 数据
- ./data:/app/data
# 挂载配置文件(方便热更新)
- ./config.yaml:/app/config.yaml:ro
# 挂载待入库文档目录
- ${MD_VECTOR_DOCS_DIR:-./md_docs}:/app/md_docs:ro
environment:
- MD_VECTOR_CONFIG=/app/config.yaml
- MD_VECTOR_DB_DATA_DIR=/app/data
- MD_VECTOR_API_KEY=${MD_VECTOR_API_KEY:-}
- EMBED_API_KEY=${EMBED_API_KEY:-}
- CORS_ORIGINS=${CORS_ORIGINS:-http://localhost:3000}
- MAX_REQUEST_BODY_SIZE=${MAX_REQUEST_BODY_SIZE:-10485760}
env_file:
- .env # 可选,Docker 会忽略不存在的文件
healthcheck:
test: ["CMD", "python", "-c", "import urllib.request; urllib.request.urlopen('http://localhost:8000/api/v1/health')"]
interval: 30s
timeout: 10s
retries: 3
start_period: 60s
# 可选:GPU 版本(需 nvidia-container-toolkit)
md-vector-db-gpu:
profiles: ["gpu"]
build:
context: .
dockerfile: Dockerfile
image: md-vector-db:latest
container_name: md-vector-db-gpu
restart: unless-stopped
ports:
- "${MD_VECTOR_PORT:-8000}:8000"
volumes:
- ./data:/app/data
- ./config.yaml:/app/config.yaml:ro
- ${MD_VECTOR_DOCS_DIR:-./md_docs}:/app/md_docs:ro
environment:
- MD_VECTOR_CONFIG=/app/config.yaml
- MD_VECTOR_DB_DATA_DIR=/app/data
- MD_VECTOR_API_KEY=${MD_VECTOR_API_KEY:-}
- EMBED_API_KEY=${EMBED_API_KEY:-}
- CORS_ORIGINS=${CORS_ORIGINS:-http://localhost:3000}
env_file:
- .env
deploy:
resources:
reservations:
devices:
- driver: nvidia
count: 1
capabilities: [gpu]
- Step 4: 验证 Docker 构建
docker build -t md-vector-db:test .
docker run --rm md-vector-db:test md-vector-db --help
- Step 5: Commit
git add Dockerfile docker-compose.yml .dockerignore
git commit -m "feat: 添加 Dockerfile 和 docker-compose.yml 部署方案"
Task 10: 更新 README 文档
Files:
-
Modify:
README.md -
Step 1: 在 README 中添加新功能说明
在 ## 功能特性 列表中追加:
- **混合检索**: BM25 关键词 + 向量语义联合检索,加权融合排序
- **增量入库**: 基于 SHA256 自动跳过未变更文件,避免重复嵌入
- **结果重排序**: 可选 Cross-Encoder 精确重排,提升检索精度
- **Docker 部署**: 提供 Dockerfile 和 docker-compose.yml,一键部署
在 ## HTTP API 之后添加:
## Docker 部署
```bash
# 构建镜像
docker compose build
# 启动服务(CPU 模式)
docker compose up -d
# 启动服务(GPU 模式,需 nvidia-container-toolkit)
docker compose --profile gpu up -d
# 查看日志
docker compose logs -f
# 停止服务
docker compose down
在 `## 配置说明` 的 `config.yaml` 示例末尾追加:
```yaml
search:
mode: hybrid # hybrid | vector(默认 hybrid)
bm25_weight: 0.3 # BM25 权重(0=纯向量, 1=纯BM25)
candidate_multiplier: 3 # 向量检索候选倍数
enable_rerank: false # Cross-Encoder 重排序
- Step 2: Commit
git add README.md
git commit -m "docs: 更新 README — 混合检索、增量入库、Docker 部署"
Task 11: 最终验证
Files: 无(纯验证步骤)
- Step 1: 运行全部测试(含覆盖率)
uv run pytest tests/ --cov=src --cov-report=term-missing -v
预期: 全部通过,覆盖率 ≥ 78%(从 75% 提升)
- Step 2: 运行 ruff lint
uv run ruff check src/ tests/
预期: 零错误
- Step 3: 功能集成测试
# 启动服务
uv run md-vector-db serve --port 8000 &
sleep 3
# 入库测试文档
echo "# 测试\n\n这是一段测试内容。" > /tmp/test_hybrid.md
uv run md-vector-db ingest /tmp/test_hybrid.md --incremental
# 混合检索
curl -s http://localhost:8000/api/v1/search \
-H "Content-Type: application/json" \
-H "x-api-key: your-secret-key" \
-d '{"query": "测试内容", "top_k": 3}'
# 增量入库(应跳过)
uv run md-vector-db ingest /tmp/test_hybrid.md --incremental
# 关闭服务
kill %1
- Step 4: 总结 Commit
git add -A
git commit -m "feat: 第一优先级完善 — 混合检索、增量入库、重排序、Docker、安全加固"
自审清单
- Spec 覆盖: 混合检索 ✅ | 增量入库 ✅ | 重排序 ✅ | Docker ✅ | 安全加固 ✅
- 无占位符: 所有代码块均为具体实现,无 TODO/TBD
- 类型一致性:
SearchResultTypedDict 在 retriever.py 定义,search.py 和 reranker.py 均引用 - 测试先行: 每个模块都是先写测试(Step 1-2),再实现(Step 3-5)