1729 lines
46 KiB
Markdown
1729 lines
46 KiB
Markdown
# 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**
|
||
|
||
```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**
|
||
|
||
```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**
|
||
|
||
```gitignore
|
||
__pycache__/
|
||
*.pyc
|
||
data/
|
||
.env
|
||
*.egg-info/
|
||
.pytest_cache/
|
||
```
|
||
|
||
- [ ] **Step 5: 安装依赖**
|
||
|
||
```bash
|
||
cd D:/Code/doing_exercises/programs/md-vector-db
|
||
uv sync
|
||
```
|
||
|
||
- [ ] **Step 6: Commit**
|
||
|
||
```bash
|
||
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**
|
||
|
||
```python
|
||
"""配置加载模块测试."""
|
||
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: 运行测试验证失败**
|
||
|
||
```bash
|
||
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**
|
||
|
||
```python
|
||
"""应用配置加载模块."""
|
||
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: 运行测试验证通过**
|
||
|
||
```bash
|
||
uv run pytest tests/test_config.py -v
|
||
```
|
||
|
||
Expected: PASS
|
||
|
||
- [ ] **Step 5: Commit**
|
||
|
||
```bash
|
||
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**
|
||
|
||
```python
|
||
"""数据库层测试."""
|
||
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: 运行测试验证失败**
|
||
|
||
```bash
|
||
uv run pytest tests/test_db.py -v
|
||
```
|
||
|
||
Expected: FAIL
|
||
|
||
- [ ] **Step 3: 实现 db.py**
|
||
|
||
```python
|
||
"""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: 运行测试验证通过**
|
||
|
||
```bash
|
||
uv run pytest tests/test_db.py -v
|
||
```
|
||
|
||
Expected: PASS
|
||
|
||
- [ ] **Step 5: Commit**
|
||
|
||
```bash
|
||
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**
|
||
|
||
```python
|
||
"""嵌入模型测试."""
|
||
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: 运行测试验证失败**
|
||
|
||
```bash
|
||
uv run pytest tests/test_embedder.py -v
|
||
```
|
||
|
||
Expected: FAIL
|
||
|
||
- [ ] **Step 3: 实现 embedder.py**
|
||
|
||
```python
|
||
"""嵌入模型抽象层."""
|
||
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: 运行测试验证通过**
|
||
|
||
```bash
|
||
uv run pytest tests/test_embedder.py -v
|
||
```
|
||
|
||
Expected: PASS (首次运行会下载模型)
|
||
|
||
- [ ] **Step 5: Commit**
|
||
|
||
```bash
|
||
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**
|
||
|
||
```python
|
||
"""文档入库测试."""
|
||
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**
|
||
|
||
```python
|
||
"""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: 运行测试验证通过**
|
||
|
||
```bash
|
||
uv run pytest tests/test_ingest.py -v
|
||
```
|
||
|
||
Expected: PASS
|
||
|
||
- [ ] **Step 5: Commit**
|
||
|
||
```bash
|
||
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**
|
||
|
||
```python
|
||
"""检索模块测试."""
|
||
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: 运行测试验证失败**
|
||
|
||
```bash
|
||
uv run pytest tests/test_search.py -v
|
||
```
|
||
|
||
Expected: FAIL
|
||
|
||
- [ ] **Step 3: 实现 search.py**
|
||
|
||
```python
|
||
"""语义检索模块."""
|
||
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: 运行测试验证通过**
|
||
|
||
```bash
|
||
uv run pytest tests/test_search.py -v
|
||
```
|
||
|
||
Expected: PASS
|
||
|
||
- [ ] **Step 5: Commit**
|
||
|
||
```bash
|
||
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**
|
||
|
||
```python
|
||
"""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: 运行测试验证失败**
|
||
|
||
```bash
|
||
uv run pytest tests/test_api.py -v
|
||
```
|
||
|
||
Expected: FAIL
|
||
|
||
- [ ] **Step 3: 实现 server/app.py**
|
||
|
||
```python
|
||
"""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: 运行测试验证通过**
|
||
|
||
```bash
|
||
uv run pytest tests/test_api.py -v
|
||
```
|
||
|
||
Expected: PASS
|
||
|
||
- [ ] **Step 5: Commit**
|
||
|
||
```bash
|
||
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**
|
||
|
||
```python
|
||
"""命令行工具入口."""
|
||
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 可用**
|
||
|
||
```bash
|
||
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**
|
||
|
||
```bash
|
||
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**
|
||
|
||
```python
|
||
"""便捷启动脚本."""
|
||
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]` 段末尾添加:
|
||
|
||
```toml
|
||
[project.scripts]
|
||
md-vector-db = "src.cli.main:app"
|
||
```
|
||
|
||
- [ ] **Step 3: 验证 pip install 后可命令行调用**
|
||
|
||
```bash
|
||
uv pip install -e .
|
||
md-vector-db --help
|
||
```
|
||
|
||
- [ ] **Step 4: Commit**
|
||
|
||
```bash
|
||
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 文件**
|
||
|
||
```bash
|
||
mkdir -p md_docs
|
||
cat > md_docs/test-guide.md << 'EOF'
|
||
# 向量数据库入门指南
|
||
|
||
## 什么是向量数据库
|
||
|
||
向量数据库是一种专门用于存储和检索向量嵌入的数据库。
|
||
|
||
## 为什么需要向量数据库
|
||
|
||
传统的数据库只能做精确匹配,而向量数据库可以做语义相似度搜索。
|
||
|
||
## 常用的向量数据库
|
||
|
||
- ChromaDB:轻量级,适合小项目
|
||
- Qdrant:高性能,适合生产环境
|
||
- Milvus:分布式,适合大规模数据
|
||
|
||
## ChromaDB 快速上手
|
||
|
||
安装 ChromaDB:
|
||
|
||
```bash
|
||
pip install chromadb
|
||
```
|
||
|
||
创建 collection 并添加数据:
|
||
|
||
```python
|
||
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: 运行全部测试验证覆盖率**
|
||
|
||
```bash
|
||
uv run pytest tests/ -v --tb=short
|
||
```
|
||
|
||
Expected: 所有测试 PASS
|
||
|
||
- [ ] **Step 4: 创建 README.md**
|
||
|
||
```markdown
|
||
# 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"
|
||
```
|