46 KiB
md-vector-db 实现计划
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: 构建一个 Markdown 文档向量数据库,支持文档入库、嵌入、语义检索,并通过 FastAPI HTTP 服务供其他项目调用。
Architecture: 分层设计 — core 层封装 ChromaDB + 嵌入模型逻辑,server 层提供 FastAPI HTTP 接口,cli 层提供命令行工具。core 不依赖 server/cli,两者各自调用 core。
Tech Stack: Python 3.13, ChromaDB, sentence-transformers (BAAI/bge-small-zh-v1.5), FastAPI + uvicorn, Typer, PyYAML, markdown-it-py
文件结构
| 文件 | 职责 |
|---|---|
pyproject.toml |
项目元数据和依赖声明 |
config.yaml |
用户配置文件(嵌入模式、分块参数、服务端口等) |
src/core/config.py |
配置加载,将 YAML 解析为 dataclass 对象 |
src/core/db.py |
ChromaDB PersistentClient 初始化,collection 管理 |
src/core/embedder.py |
嵌入模型抽象(本地/API 双模),EmbedConfig → Embedder |
src/core/ingest.py |
Markdown 解析、标题提取、混合分块、入库、去重 |
src/core/search.py |
语义检索,query 嵌入→ChromaDB 相似度搜索→格式化结果 |
src/server/app.py |
FastAPI 路由:ingest/search/collections/documents/health |
src/cli/main.py |
Typer CLI:ingest/ingest-dir/search/serve/stats |
scripts/serve.py |
一行启动脚本:uvicorn src.server.app:app |
tests/test_config.py |
配置加载测试 |
tests/test_embedder.py |
嵌入模型测试(本地模式核心验证) |
tests/test_ingest.py |
分块逻辑与入库测试 |
tests/test_search.py |
检索接口测试 |
tests/test_api.py |
FastAPI 端点集成测试 |
Task 1: 项目脚手架
Files:
-
Create:
pyproject.toml -
Create:
config.yaml -
Create:
src/__init__.py -
Create:
src/core/__init__.py -
Create:
src/server/__init__.py -
Create:
src/cli/__init__.py -
Create:
tests/__init__.py -
Create:
scripts/__init__.py(空,不创建也行) -
Create:
.gitignore -
Step 1: 创建 pyproject.toml
[project]
name = "md-vector-db"
version = "0.1.0"
description = "Markdown 文档向量数据库,支持语义检索"
requires-python = ">=3.13"
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",
]
[project.optional-dependencies]
dev = [
"pytest>=8.0",
"httpx>=0.27.0",
]
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
[tool.pytest.ini_options]
testpaths = ["tests"]
pythonpath = ["src"]
- Step 2: 创建 config.yaml
chroma:
persist_dir: ./data
collection_name: markdown_docs
embed:
mode: local # local | api
local_model: BAAI/bge-small-zh-v1.5
api_base: "" # api 模式下填写
api_key: "" # api 模式下填写
chunk:
max_size: 1000 # 分块最大字符数
overlap: 100 # 相邻块重叠字符数
server:
host: 0.0.0.0
port: 8000
- Step 3: 创建所有 __init__.py(空文件)
创建以下空文件:
-
src/__init__.py -
src/core/__init__.py -
src/server/__init__.py -
src/cli/__init__.py -
tests/__init__.py -
Step 4: 创建 .gitignore
__pycache__/
*.pyc
data/
.env
*.egg-info/
.pytest_cache/
- Step 5: 安装依赖
cd D:/Code/doing_exercises/programs/md-vector-db
uv sync
- Step 6: Commit
git init
git add -A
git commit -m "chore: scaffold md-vector-db project structure"
Task 2: 配置加载模块 (config.py)
Files:
-
Create:
src/core/config.py -
Create:
tests/test_config.py -
Step 1: 写测试 — tests/test_config.py
"""配置加载模块测试."""
import tempfile
from pathlib import Path
import yaml
import pytest
from src.core.config import AppConfig, EmbedConfig, ChunkConfig, load_config
class TestEmbedConfig:
"""嵌入配置解析."""
def test_local_mode_defaults(self):
"""默认 local 模式,带默认模型名."""
data = {"embed": {"mode": "local"}}
cfg = AppConfig(**data)
assert cfg.embed.mode == "local"
assert cfg.embed.local_model == "BAAI/bge-small-zh-v1.5"
assert cfg.embed.api_base == ""
def test_api_mode_fields(self):
"""api 模式下 api_base 和 api_key 可设置."""
data = {
"embed": {
"mode": "api",
"api_base": "https://api.openai.com/v1",
"api_key": "sk-test",
}
}
cfg = AppConfig(**data)
assert cfg.embed.mode == "api"
assert cfg.embed.api_base == "https://api.openai.com/v1"
assert cfg.embed.api_key == "sk-test"
class TestChunkConfig:
"""分块配置解析."""
def test_default_values(self):
"""分块默认值正确."""
data = {"chunk": {}}
cfg = AppConfig(**data)
assert cfg.chunk.max_size == 1000
assert cfg.chunk.overlap == 100
class TestLoadConfig:
"""load_config 函数测试."""
def test_load_from_yaml_file(self):
"""从 YAML 文件加载配置."""
yaml_content = """
chroma:
persist_dir: /tmp/test_data
collection_name: test_collection
embed:
mode: local
chunk:
max_size: 500
server:
port: 9000
"""
with tempfile.NamedTemporaryFile(
mode="w", suffix=".yaml", delete=False
) as f:
f.write(yaml_content)
tmp_path = f.name
try:
cfg = load_config(tmp_path)
assert cfg.chroma.persist_dir == "/tmp/test_data"
assert cfg.chroma.collection_name == "test_collection"
assert cfg.chunk.max_size == 500
assert cfg.server.port == 9000
finally:
Path(tmp_path).unlink()
def test_load_missing_file_uses_defaults(self):
"""配置文件不存在时使用默认值."""
cfg = load_config("/nonexistent/config.yaml")
assert cfg.chroma.persist_dir == "./data"
assert cfg.embed.mode == "local"
- Step 2: 运行测试验证失败
cd D:/Code/doing_exercises/programs/md-vector-db
uv run pytest tests/test_config.py -v
Expected: FAIL (module not found)
- Step 3: 实现 config.py
"""应用配置加载模块."""
from dataclasses import dataclass, field
from pathlib import Path
import yaml
@dataclass
class ChromaConfig:
"""ChromaDB 配置."""
persist_dir: str = "./data"
collection_name: str = "markdown_docs"
@dataclass
class EmbedConfig:
"""嵌入模型配置."""
mode: str = "local" # "local" | "api"
local_model: str = "BAAI/bge-small-zh-v1.5"
api_base: str = ""
api_key: str = ""
@dataclass
class ChunkConfig:
"""文档分块配置."""
max_size: int = 1000
overlap: int = 100
@dataclass
class ServerConfig:
"""HTTP 服务配置."""
host: str = "0.0.0.0"
port: int = 8000
@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)
def __init__(self, **kwargs):
self.chroma = ChromaConfig(**kwargs.get("chroma", {}))
self.embed = EmbedConfig(**kwargs.get("embed", {}))
self.chunk = ChunkConfig(**kwargs.get("chunk", {}))
self.server = ServerConfig(**kwargs.get("server", {}))
def load_config(path: str | None = None) -> AppConfig:
"""从 YAML 文件加载配置, 若文件不存在则返回默认配置."""
config_path = path or "config.yaml"
if not Path(config_path).exists():
return AppConfig()
with open(config_path, "r", encoding="utf-8") as f:
data = yaml.safe_load(f) or {}
return AppConfig(**data)
- Step 4: 运行测试验证通过
uv run pytest tests/test_config.py -v
Expected: PASS
- Step 5: Commit
git add src/core/config.py tests/test_config.py
git commit -m "feat: add config loading module with YAML support"
Task 3: 数据库层 (db.py)
Files:
-
Create:
src/core/db.py -
Create:
tests/test_db.py -
Step 1: 写测试 — tests/test_db.py
"""数据库层测试."""
import tempfile
from pathlib import Path
import pytest
from src.core.db import VectorDB
class TestVectorDB:
"""VectorDB 测试."""
@pytest.fixture
def temp_dir(self):
with tempfile.TemporaryDirectory() as d:
yield d
def test_init_creates_persist_directory(self, temp_dir):
"""初始化时自动创建持久化目录."""
db = VectorDB(persist_dir=temp_dir + "/subdir")
assert Path(temp_dir + "/subdir").exists()
assert db.client is not None
def test_get_or_create_collection(self, temp_dir):
"""创建或获取 collection."""
db = VectorDB(persist_dir=temp_dir)
col = db.get_or_create_collection("test_col")
assert col.name == "test_col"
# 再次获取应返回同一个
col2 = db.get_or_create_collection("test_col")
assert col2.name == "test_col"
def test_count_returns_zero_for_empty_collection(self, temp_dir):
"""空 collection 文档数为 0."""
db = VectorDB(persist_dir=temp_dir)
col = db.get_or_create_collection("test_col")
assert col.count() == 0
def test_delete_collection(self, temp_dir):
"""删除 collection."""
db = VectorDB(persist_dir=temp_dir)
db.get_or_create_collection("tmp_col")
db.delete_collection("tmp_col")
# 再次获取会创建新的
col = db.get_or_create_collection("tmp_col")
assert col.count() == 0
- Step 2: 运行测试验证失败
uv run pytest tests/test_db.py -v
Expected: FAIL
- Step 3: 实现 db.py
"""ChromaDB 数据库层."""
import chromadb
from chromadb.api.models.Collection import Collection
class VectorDB:
"""向量数据库封装."""
def __init__(self, persist_dir: str = "./data"):
self.client = chromadb.PersistentClient(path=persist_dir)
def get_or_create_collection(self, name: str) -> Collection:
"""获取或创建 collection."""
return self.client.get_or_create_collection(name=name)
def delete_collection(self, name: str) -> None:
"""删除 collection."""
try:
self.client.delete_collection(name=name)
except ValueError:
pass # collection 不存在则忽略
- Step 4: 运行测试验证通过
uv run pytest tests/test_db.py -v
Expected: PASS
- Step 5: Commit
git add src/core/db.py tests/test_db.py
git commit -m "feat: add ChromaDB database layer"
Task 4: 嵌入层 (embedder.py)
Files:
-
Create:
src/core/embedder.py -
Create:
tests/test_embedder.py -
Step 1: 写测试 — tests/test_embedder.py
"""嵌入模型测试."""
import pytest
from src.core.config import EmbedConfig
from src.core.embedder import Embedder, create_embedder
class TestEmbedder:
"""Embedder 单元测试.
注意:本地模型测试需要下载 sentence-transformers 模型(约 100MB),
首次运行耗时较长。API 模式测试使用 mock 避免网络依赖。
"""
@pytest.fixture
def local_config(self):
return EmbedConfig(mode="local")
def test_create_local_embedder(self, local_config):
"""创建本地嵌入器,验证维度正确."""
embedder = create_embedder(local_config)
assert embedder.dimension > 0
assert isinstance(embedder.dimension, int)
def test_embed_single_text(self, local_config):
"""嵌入单条文本返回正确维度向量."""
embedder = create_embedder(local_config)
result = embedder.embed(["你好世界"])
assert len(result) == 1
assert len(result[0]) == embedder.dimension
assert all(isinstance(v, float) for v in result[0])
def test_embed_multiple_texts(self, local_config):
"""嵌入多条文本返回对应数量的向量."""
embedder = create_embedder(local_config)
texts = ["第一段文本", "第二段文本", "第三段文本"]
result = embedder.embed(texts)
assert len(result) == 3
for vec in result:
assert len(vec) == embedder.dimension
def test_embed_empty_list_raises(self, local_config):
"""空列表应抛出异常."""
embedder = create_embedder(local_config)
with pytest.raises(ValueError):
embedder.embed([])
def test_create_embedder_from_factory_function(self, local_config):
"""工厂函数正确创建 Embedder 实例."""
embedder = create_embedder(local_config)
assert isinstance(embedder, Embedder)
- Step 2: 运行测试验证失败
uv run pytest tests/test_embedder.py -v
Expected: FAIL
- Step 3: 实现 embedder.py
"""嵌入模型抽象层."""
from sentence_transformers import SentenceTransformer
from src.core.config import EmbedConfig
class Embedder:
"""文本嵌入器, 支持本地模型和 API 两种模式."""
def __init__(self, config: EmbedConfig):
self._config = config
if config.mode == "local":
self._model = SentenceTransformer(config.local_model)
elif config.mode == "api":
self._model = None # 延迟初始化, 需要 openai 包
self._api_base = config.api_base
self._api_key = config.api_key
else:
raise ValueError(f"不支持的嵌入模式: {config.mode}")
@property
def mode(self) -> str:
"""当前嵌入模式."""
return self._config.mode
@property
def dimension(self) -> int:
"""嵌入向量维度."""
if self.mode == "local":
return self._model.get_sentence_embedding_dimension()
else:
# 默认 OpenAI text-embedding-ada-002 / text-embedding-3-small 维度
return 1536
def embed(self, texts: list[str]) -> list[list[float]]:
"""对文本列表进行嵌入, 返回向量列表."""
if not texts:
raise ValueError("文本列表不能为空")
if self.mode == "local":
embeddings = self._model.encode(texts, normalize_embeddings=True)
return embeddings.tolist()
else:
return self._embed_via_api(texts)
def _embed_via_api(self, texts: list[str]) -> list[list[float]]:
"""通过 OpenAI 兼容 API 嵌入(延迟导入 openai)."""
from openai import OpenAI
client = OpenAI(base_url=self._api_base, api_key=self._api_key)
response = client.embeddings.create(
model="text-embedding-3-small",
input=texts,
)
return [d.embedding for d in response.data]
def create_embedder(config: EmbedConfig) -> Embedder:
"""工厂函数: 根据配置创建嵌入器."""
return Embedder(config)
- Step 4: 运行测试验证通过
uv run pytest tests/test_embedder.py -v
Expected: PASS (首次运行会下载模型)
- Step 5: Commit
git add src/core/embedder.py tests/test_embedder.py
git commit -m "feat: add embedder with local/API dual mode"
Task 5: 文档入库模块 (ingest.py)
Files:
-
Create:
src/core/ingest.py -
Create:
tests/test_ingest.py -
Step 1: 写测试 — tests/test_ingest.py
"""文档入库测试."""
import tempfile
from pathlib import Path
import pytest
from src.core.ingest import MarkdownSplitter, DocumentIngestor
class TestMarkdownSplitter:
"""Markdown 分块器测试."""
@pytest.fixture
def splitter(self):
return MarkdownSplitter(max_size=1000, overlap=100)
def test_split_simple_document(self, splitter):
"""简单文档按标题拆分."""
md = """# 标题一
这是第一段内容。
## 标题二
这是第二段内容。
# 标题三
这是第三段内容。"""
chunks = splitter.split(md, source_file="test.md")
assert len(chunks) >= 3
# 每个 chunk 有元数据
for chunk in chunks:
assert "content" in chunk
assert chunk["source_file"] == "test.md"
def test_chunk_has_heading_metadata(self, splitter):
"""chunk 附带标题元数据."""
md = "# 配置指南\n这里是配置说明。"
chunks = splitter.split(md, source_file="config.md")
assert len(chunks) >= 1
title = chunks[0]["section_title"]
assert "配置指南" in title or title == ""
def test_long_section_is_split(self, splitter):
"""超长章节被进一步拆分."""
# 创建一个超过 max_size 的段落
long_text = "这是很长的文本。" * 300 # ~3000 字符
md = f"# 长章节\n{long_text}"
small_splitter = MarkdownSplitter(max_size=500, overlap=50)
chunks = small_splitter.split(md, source_file="long.md")
assert len(chunks) > 1
def test_empty_document(self, splitter):
"""空文档返回空列表."""
chunks = splitter.split("", source_file="empty.md")
assert chunks == []
def test_code_blocks_preserved(self, splitter):
"""代码块不被拆分."""
md = """# 代码示例
```python
def hello():
print("world")
""" chunks = splitter.split(md, source_file="code.md") assert len(chunks) >= 1 # 代码块内容应在某个 chunk 中 all_content = " ".join(c["content"] for c in chunks) assert "def hello()" in all_content
class TestDocumentIngestor: """文档入库器测试."""
@pytest.fixture
def temp_md_dir(self):
with tempfile.TemporaryDirectory() as d:
# 创建测试 Markdown 文件
md_path = Path(d) / "test.md"
md_path.write_text("# 测试\n这是测试内容。", encoding="utf-8")
yield d
def test_read_markdown_file(self, temp_md_dir):
"""读取 Markdown 文件."""
ingestor = DocumentIngestor.__new__(DocumentIngestor)
content = Path(temp_md_dir + "/test.md").read_text(encoding="utf-8")
assert "测试" in content
assert "这是测试内容" in content
- [ ] **Step 2: 运行测试验证失败**
```bash
uv run pytest tests/test_ingest.py -v
Expected: FAIL
- Step 3: 实现 ingest.py
"""Markdown 文档解析与入库模块."""
import re
from dataclasses import dataclass, field
from pathlib import Path
from src.core.db import VectorDB
from src.core.embedder import Embedder
@dataclass
class Chunk:
"""文档分块."""
content: str
source_file: str
section_title: str = ""
heading_level: int = 0
chunk_index: int = 0
def to_metadata(self) -> dict:
"""转为 ChromaDB 元数据."""
return {
"source_file": self.source_file,
"section_title": self.section_title,
"heading_level": self.heading_level,
"chunk_index": self.chunk_index,
}
class MarkdownSplitter:
"""Markdown 混合分块器:先按标题拆,超长再按段落拆."""
def __init__(self, max_size: int = 1000, overlap: int = 100):
self.max_size = max_size
self.overlap = overlap
def split(self, text: str, source_file: str = "") -> list[dict]:
"""将 Markdown 文本拆分为带元数据的 chunk 列表."""
if not text.strip():
return []
sections = self._split_by_headings(text)
chunks = []
for section in sections:
if len(section["content"]) <= self.max_size:
chunks.append(section)
else:
sub_chunks = self._split_by_paragraphs(
section["content"],
section["section_title"],
section["heading_level"],
)
chunks.extend(sub_chunks)
# 为所有 chunk 补充 source_file 和 chunk_index
for i, chunk in enumerate(chunks):
chunk["source_file"] = source_file or chunk.get("source_file", "")
chunk["chunk_index"] = i
return chunks
def _split_by_headings(self, text: str) -> list[dict]:
"""按 Markdown 标题拆分."""
# 匹配行首的 # 标题
heading_pattern = re.compile(r"^(#{1,6})\s+(.+)$", re.MULTILINE)
matches = list(heading_pattern.finditer(text))
if not matches:
return [{
"content": text.strip(),
"section_title": "",
"heading_level": 0,
}]
sections = []
for i, match in enumerate(matches):
level = len(match.group(1))
title = match.group(2).strip()
start = match.end()
end = matches[i + 1].start() if i + 1 < len(matches) else len(text)
content = text[start:end].strip()
if content:
sections.append({
"content": f"{match.group(0)}\n{content}",
"section_title": title,
"heading_level": level,
})
# 处理第一个标题之前的内容
if matches and matches[0].start() > 0:
preamble = text[:matches[0].start()].strip()
if preamble:
sections.insert(0, {
"content": preamble,
"section_title": "",
"heading_level": 0,
})
return sections
def _split_by_paragraphs(
self, text: str, section_title: str, heading_level: int
) -> list[dict]:
"""按段落边界拆分超长章节."""
paragraphs = re.split(r"\n\n+", text)
chunks = []
current = ""
count = 0
for para in paragraphs:
if len(current) + len(para) > self.max_size and current:
chunks.append({
"content": current.strip(),
"section_title": section_title,
"heading_level": heading_level,
})
# overlap: 保留上一块的末尾部分
if self.overlap > 0 and len(current) > self.overlap:
current = current[-self.overlap:] + "\n\n" + para
else:
current = para
else:
if current:
current += "\n\n" + para
else:
current = para
if current.strip():
chunks.append({
"content": current.strip(),
"section_title": section_title,
"heading_level": heading_level,
})
return chunks
class DocumentIngestor:
"""文档入库器: 读取 MD 文件 → 分块 → 嵌入 → 入库."""
def __init__(self, db: VectorDB, embedder: Embedder, collection_name: str):
self.db = db
self.embedder = embedder
self.collection_name = collection_name
self.splitter = MarkdownSplitter()
@property
def collection(self):
return self.db.get_or_create_collection(self.collection_name)
def ingest_file(self, file_path: str) -> int:
"""入库单个 Markdown 文件, 返回 chunk 数量."""
path = Path(file_path)
content = path.read_text(encoding="utf-8")
file_name = path.name
return self.ingest_content(content, file_name)
def ingest_content(self, content: str, file_name: str) -> int:
"""入库 Markdown 内容(无需实际文件)."""
# 去重:先删旧 chunks
self._remove_by_source(file_name)
# 分块
chunks = self.splitter.split(content, source_file=file_name)
if not chunks:
return 0
# 嵌入
texts = [c["content"] for c in chunks]
embeddings = self.embedder.embed(texts)
# 入库
ids = [f"{file_name}_{i}" for i in range(len(chunks))]
metadatas = [
{
"source_file": c.get("source_file", file_name),
"section_title": c.get("section_title", ""),
"heading_level": c.get("heading_level", 0),
"chunk_index": i,
}
for i, c in enumerate(chunks)
]
self.collection.add(
ids=ids,
embeddings=embeddings,
documents=texts,
metadatas=metadatas,
)
return len(chunks)
def ingest_directory(self, dir_path: str) -> dict[str, int]:
"""入库目录下所有 Markdown 文件."""
results = {}
for md_file in Path(dir_path).rglob("*.md"):
count = self.ingest_file(str(md_file))
results[md_file.name] = count
return results
def _remove_by_source(self, file_name: str) -> None:
"""按 source_file 删除已有 chunks."""
try:
existing = self.collection.get(
where={"source_file": file_name}
)
if existing and existing["ids"]:
self.collection.delete(ids=existing["ids"])
except Exception:
pass # collection 为空时 get 可能抛异常
- Step 4: 运行测试验证通过
uv run pytest tests/test_ingest.py -v
Expected: PASS
- Step 5: Commit
git add src/core/ingest.py tests/test_ingest.py
git commit -m "feat: add Markdown splitter and document ingestor"
Task 6: 检索模块 (search.py)
Files:
-
Create:
src/core/search.py -
Create:
tests/test_search.py -
Step 1: 写测试 — tests/test_search.py
"""检索模块测试."""
import tempfile
from pathlib import Path
import pytest
from src.core.config import EmbedConfig
from src.core.db import VectorDB
from src.core.embedder import create_embedder
from src.core.ingest import DocumentIngestor
from src.core.search import Searcher
@pytest.fixture
def searcher():
"""创建带测试数据的 Searcher."""
tmpdir = tempfile.mkdtemp()
db = VectorDB(persist_dir=tmpdir)
embedder = create_embedder(EmbedConfig(mode="local"))
ingestor = DocumentIngestor(db, embedder, "test_search")
# 入库一些测试文档
content = """# Python 入门
Python 是一种解释型编程语言。
## 安装 Python
从 python.org 下载安装包。
# 向量数据库
ChromaDB 是一个轻量级向量数据库。
## ChromaDB 安装
使用 pip install chromadb 安装。"""
ingestor.ingest_content(content, "guide.md")
return Searcher(db, embedder, "test_search")
class TestSearcher:
"""检索器测试."""
def test_search_returns_results(self, searcher):
"""搜索返回至少一条结果."""
results = searcher.search("Python 编程", top_k=3)
assert len(results) > 0
for r in results:
assert r["content"]
assert r["source_file"]
assert "score" in r
def test_search_scores_are_descending(self, searcher):
"""搜索结果按相似度降序排列."""
results = searcher.search("向量数据库", top_k=5)
scores = [r["score"] for r in results]
assert scores == sorted(scores)
def test_search_respects_top_k(self, searcher):
"""top_k 参数限制返回数量."""
results = searcher.search("安装", top_k=2)
assert len(results) <= 2
def test_search_returns_all_fields(self, searcher):
"""搜索结果包含完整字段."""
results = searcher.search("ChromaDB", top_k=1)
if results:
r = results[0]
assert "content" in r
assert "source_file" in r
assert "section_title" in r
assert "heading_level" in r
assert "chunk_index" in r
assert "score" in r
def test_search_no_results(self, searcher):
"""无语义匹配时不崩溃."""
results = searcher.search("xyzxyz不存在的内容abcabc", top_k=3)
assert isinstance(results, list)
# ChromaDB 总是返回最近的向量, 所以即使不匹配也会返回结果, 只是 score 低
# 这里只验证不抛异常
- Step 2: 运行测试验证失败
uv run pytest tests/test_search.py -v
Expected: FAIL
- Step 3: 实现 search.py
"""语义检索模块."""
from src.core.db import VectorDB
from src.core.embedder import Embedder
class Searcher:
"""向量检索器."""
def __init__(self, db: VectorDB, embedder: Embedder, collection_name: str):
self.db = db
self.embedder = embedder
self.collection_name = collection_name
@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[dict]:
"""语义检索, 返回格式化结果列表."""
query_embedding = self.embedder.embed([query])[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
# ChromaDB 默认用余弦距离, 转为相似度分数 (0~1)
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:
"""获取 collection 信息."""
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:
"""按文件名删除文档."""
try:
existing = self.collection.get(
where={"source_file": file_name}
)
if existing and existing["ids"]:
self.collection.delete(ids=existing["ids"])
return True
except Exception:
pass
return False
- Step 4: 运行测试验证通过
uv run pytest tests/test_search.py -v
Expected: PASS
- Step 5: Commit
git add src/core/search.py tests/test_search.py
git commit -m "feat: add semantic search module"
Task 7: FastAPI 服务层 (server/app.py)
Files:
-
Create:
src/server/app.py -
Create:
tests/test_api.py -
Step 1: 写测试 — tests/test_api.py
"""API 端点集成测试."""
import tempfile
from pathlib import Path
import pytest
from fastapi.testclient import TestClient
@pytest.fixture
def app_client(monkeypatch):
"""创建测试客户端(使用临时目录)."""
tmpdir = tempfile.mkdtemp()
# 在导入 app 前设置环境
monkeypatch.setattr(
"src.server.app._get_config",
lambda: type(
"C",
(),
{
"chroma": type("Ch", (), {"persist_dir": tmpdir, "collection_name": "test_api"})(),
"embed": type("Em", (), {"mode": "local"})(),
"chunk": type("Ck", (), {"max_size": 1000, "overlap": 100})(),
"server": type("Sv", (), {"host": "0.0.0.0", "port": 8000})(),
},
)(),
)
from src.server.app import app
return TestClient(app)
class TestHealthEndpoint:
"""健康检查."""
def test_health_returns_ok(self, app_client):
response = app_client.get("/api/v1/health")
assert response.status_code == 200
data = response.json()
assert data["status"] == "ok"
class TestCollectionEndpoint:
"""Collection 端点."""
def test_list_collections(self, app_client):
response = app_client.get("/api/v1/collections")
assert response.status_code == 200
data = response.json()
assert "collections" in data
class TestSearchEndpoint:
"""搜索端点."""
def test_search_requires_query(self, app_client):
response = app_client.post("/api/v1/search", json={})
assert response.status_code == 422 # 缺少必填字段
def test_search_empty_collection(self, app_client):
response = app_client.post(
"/api/v1/search",
json={"query": "test", "top_k": 5},
)
assert response.status_code == 200
data = response.json()
assert "results" in data
assert data["results"] == []
class TestIngestEndpoint:
"""入库端点."""
def test_ingest_content(self, app_client):
response = app_client.post(
"/api/v1/ingest",
json={
"content": "# Test\nHello world.",
"file_name": "test.md",
},
)
assert response.status_code == 200
data = response.json()
assert data["status"] == "ok"
assert data["chunks"] > 0
def test_ingest_then_search(self, app_client):
"""入库后能检索到."""
# 先入库
app_client.post(
"/api/v1/ingest",
json={"content": "# 配置说明\nChromaDB 配置很简单。", "file_name": "config.md"},
)
# 再搜索
response = app_client.post(
"/api/v1/search",
json={"query": "配置", "top_k": 3},
)
assert response.status_code == 200
data = response.json()
assert len(data["results"]) > 0
- Step 2: 运行测试验证失败
uv run pytest tests/test_api.py -v
Expected: FAIL
- Step 3: 实现 server/app.py
"""FastAPI 服务层."""
import os
from pathlib import Path
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from src.core.config import load_config, AppConfig
from src.core.db import VectorDB
from src.core.embedder import create_embedder
from src.core.ingest import DocumentIngestor
from src.core.search import Searcher
# -- 模型 --
class IngestRequest(BaseModel):
file_path: str | None = None
content: str | None = None
file_name: str | None = None
class SearchRequest(BaseModel):
query: str
top_k: int = 10
# -- 懒加载单例 --
_config: AppConfig | None = None
_db: VectorDB | None = None
_embedder = None
_searcher: Searcher | None = None
_ingestor: DocumentIngestor | None = None
def _get_config():
global _config
if _config is None:
_config = load_config(os.getenv("MD_VECTOR_CONFIG", "config.yaml"))
return _config
def _get_db():
global _db
if _db is None:
cfg = _get_config()
_db = VectorDB(persist_dir=cfg.chroma.persist_dir)
return _db
def _get_embedder():
global _embedder
if _embedder is None:
cfg = _get_config()
_embedder = create_embedder(cfg.embed)
return _embedder
def _init_services():
global _searcher, _ingestor
cfg = _get_config()
db = _get_db()
embedder = _get_embedder()
_searcher = Searcher(db, embedder, cfg.chroma.collection_name)
_ingestor = DocumentIngestor(db, embedder, cfg.chroma.collection_name)
def _get_searcher():
if _searcher is None:
_init_services()
return _searcher
def _get_ingestor():
if _ingestor is None:
_init_services()
return _ingestor
# -- App --
app = FastAPI(
title="md-vector-db",
description="Markdown 文档向量数据库 API",
version="0.1.0",
)
@app.on_event("startup")
def startup():
_init_services()
@app.get("/api/v1/health")
def health():
return {"status": "ok"}
@app.get("/api/v1/collections")
def list_collections():
searcher = _get_searcher()
info = searcher.get_collection_info()
sources = searcher.list_sources()
return {"collections": [info], "sources": sources}
@app.post("/api/v1/ingest")
def ingest_document(req: IngestRequest):
ingestor = _get_ingestor()
try:
if req.file_path:
path = Path(req.file_path)
if not path.exists():
raise HTTPException(status_code=404, detail=f"文件不存在: {req.file_path}")
count = ingestor.ingest_file(str(path))
file_name = path.name
elif req.content:
file_name = req.file_name or "untitled.md"
count = ingestor.ingest_content(req.content, file_name)
else:
raise HTTPException(
status_code=400, detail="需要提供 file_path 或 content"
)
return {"status": "ok", "chunks": count, "file": file_name}
except HTTPException:
raise
except Exception as e:
raise HTTPException(status_code=500, detail=str(e))
@app.post("/api/v1/search")
def search_documents(req: SearchRequest):
searcher = _get_searcher()
results = searcher.search(req.query, top_k=req.top_k)
return {"results": results}
@app.delete("/api/v1/documents/{file_name}")
def delete_document(file_name: str):
searcher = _get_searcher()
deleted = searcher.delete_by_source(file_name)
if not deleted:
raise HTTPException(status_code=404, detail=f"文档不存在: {file_name}")
return {"status": "ok", "file": file_name}
- Step 4: 运行测试验证通过
uv run pytest tests/test_api.py -v
Expected: PASS
- Step 5: Commit
git add src/server/app.py tests/test_api.py
git commit -m "feat: add FastAPI HTTP API layer"
Task 8: CLI 工具 (cli/main.py)
Files:
-
Create:
src/cli/main.py -
Step 1: 实现 CLI
"""命令行工具入口."""
import sys
from pathlib import Path
import typer
import uvicorn
# 添加 src 到路径
sys.path.insert(0, str(Path(__file__).parent.parent))
from src.core.config import load_config
from src.core.db import VectorDB
from src.core.embedder import create_embedder
from src.core.ingest import DocumentIngestor
from src.core.search import Searcher
app = typer.Typer(name="md-vector-db", help="Markdown 文档向量数据库管理工具")
def _get_components(config_path: str = "config.yaml"):
"""初始化所有组件."""
cfg = load_config(config_path)
db = VectorDB(persist_dir=cfg.chroma.persist_dir)
embedder = create_embedder(cfg.embed)
searcher = Searcher(db, embedder, cfg.chroma.collection_name)
ingestor = DocumentIngestor(db, embedder, cfg.chroma.collection_name)
return cfg, searcher, ingestor
@app.command()
def ingest(file_path: str):
"""入库单个 Markdown 文件."""
_, _, ingestor = _get_components()
count = ingestor.ingest_file(file_path)
typer.echo(f"✅ 已入库: {file_path} ({count} 个 chunks)")
@app.command()
def ingest_dir(dir_path: str):
"""入库目录下所有 Markdown 文件."""
_, _, ingestor = _get_components()
results = ingestor.ingest_directory(dir_path)
total = sum(results.values())
for name, count in results.items():
typer.echo(f" 📄 {name}: {count} chunks")
typer.echo(f"✅ 共入库 {len(results)} 个文件, {total} 个 chunks")
@app.command()
def search(query: str, top_k: int = 10):
"""语义检索."""
_, searcher, _ = _get_components()
results = searcher.search(query, top_k=top_k)
if not results:
typer.echo("未找到匹配结果。")
return
for i, r in enumerate(results, 1):
typer.echo(f"\n--- 结果 {i} (相似度: {r['score']:.4f}) ---")
typer.echo(f"📄 来源: {r['source_file']}")
if r["section_title"]:
typer.echo(f"📑 章节: {r['section_title']}")
preview = r["content"][:200] + "..." if len(r["content"]) > 200 else r["content"]
typer.echo(preview)
@app.command()
def serve(port: int = 8000):
"""启动 HTTP 服务."""
typer.echo(f"🚀 启动服务: http://0.0.0.0:{port}")
uvicorn.run("src.server.app:app", host="0.0.0.0", port=port, reload=False)
@app.command()
def stats():
"""查看统计信息."""
_, searcher, _ = _get_components()
info = searcher.get_collection_info()
sources = searcher.list_sources()
typer.echo(f"📊 Collection: {info['name']}")
typer.echo(f"📦 总 chunks: {info['count']}")
typer.echo(f"📄 源文件数: {len(sources)}")
if sources:
typer.echo("\n源文件列表:")
for s in sources:
typer.echo(f" - {s}")
if __name__ == "__main__":
app()
- Step 2: 验证 CLI 可用
cd D:/Code/doing_exercises/programs/md-vector-db
uv run python -m src.cli.main --help
Expected: 显示帮助信息,列出 ingest/ingest-dir/search/serve/stats 命令
- Step 3: Commit
git add src/cli/main.py
git commit -m "feat: add CLI management tool"
Task 9: 启动脚本与入口完善 (scripts/serve.py + pyproject.toml 更新)
Files:
-
Create:
scripts/serve.py -
Modify:
pyproject.toml -
Step 1: 创建 scripts/serve.py
"""便捷启动脚本."""
import uvicorn
if __name__ == "__main__":
uvicorn.run("src.server.app:app", host="0.0.0.0", port=8000, reload=True)
- Step 2: 更新 pyproject.toml 添加 scripts 入口
在 pyproject.toml 的 [project] 段末尾添加:
[project.scripts]
md-vector-db = "src.cli.main:app"
- Step 3: 验证 pip install 后可命令行调用
uv pip install -e .
md-vector-db --help
- Step 4: Commit
git add scripts/serve.py pyproject.toml
git commit -m "feat: add serve script and console entry point"
Task 10: 端到端验证与文档
Files:
-
Create:
README.md -
Step 1: 创建测试用的 Markdown 文件
mkdir -p md_docs
cat > md_docs/test-guide.md << 'EOF'
# 向量数据库入门指南
## 什么是向量数据库
向量数据库是一种专门用于存储和检索向量嵌入的数据库。
## 为什么需要向量数据库
传统的数据库只能做精确匹配,而向量数据库可以做语义相似度搜索。
## 常用的向量数据库
- ChromaDB:轻量级,适合小项目
- Qdrant:高性能,适合生产环境
- Milvus:分布式,适合大规模数据
## ChromaDB 快速上手
安装 ChromaDB:
```bash
pip install chromadb
创建 collection 并添加数据:
import chromadb
client = chromadb.PersistentClient(path="./data")
collection = client.get_or_create_collection("my_docs")
collection.add(
documents=["这是第一篇文档", "这是第二篇文档"],
ids=["doc1", "doc2"],
)
EOF
- [ ] **Step 2: 端到端测试**
```bash
# 入库测试文档
uv run python -m src.cli.main ingest md_docs/test-guide.md
# 命令行搜索
uv run python -m src.cli.main search "如何安装ChromaDB" --top-k 3
# 查看统计
uv run python -m src.cli.main stats
Expected: 搜索返回相关结果,统计显示 1 个源文件
- Step 3: 运行全部测试验证覆盖率
uv run pytest tests/ -v --tb=short
Expected: 所有测试 PASS
- Step 4: 创建 README.md
# md-vector-db
Markdown 文档向量数据库,支持文档入库、语义检索,可通过 HTTP API 供其他项目调用。
## 快速开始
### 安装
\`\`\`bash
cd md-vector-db
uv sync
\`\`\`
### CLI 使用
\`\`\`bash
# 入库文档
uv run python -m src.cli.main ingest path/to/file.md
# 批量入库
uv run python -m src.cli.main ingest-dir ./md_docs/
# 搜索
uv run python -m src.cli.main search "关键词" --top-k 10
# 启动 HTTP 服务
uv run python -m src.cli.main serve --port 8000
\`\`\`
### HTTP API
\`\`\`bash
# 启动服务
uv run python -m src.cli.main serve
# 入库
curl -X POST http://localhost:8000/api/v1/ingest \\
-H "Content-Type: application/json" \\
-d '{"content": "# 标题\\n内容...", "file_name": "doc.md"}'
# 搜索
curl -X POST http://localhost:8000/api/v1/search \\
-H "Content-Type: application/json" \\
-d '{"query": "关键词", "top_k": 5}'
\`\`\`
## 配置
编辑 `config.yaml` 切换嵌入模式、分块参数等。
## 项目结构
- `src/core/` — 核心逻辑(数据库、嵌入、分块、检索)
- `src/server/` — FastAPI HTTP 服务
- `src/cli/` — 命令行工具
- `tests/` — 测试
\`\`\`
- [ ] **Step 5: 最终 Commit**
```bash
git add README.md md_docs/
git commit -m "docs: add README and test data"