fix: 修复 11 个代码架构审计问题

H1: AppConfig 自定义 __init__ 改用 from_dict() 类方法
H2: list_collections_with_stats 改为从 ChromaDB 直接查询
H3: RateLimiter 过期 key 自动清理, 防止内存泄漏
M1: delete_by_source 区分 ValueError 与真实异常, 记日志
M2: VectorDB 新增 list_collections() 封装方法
M3: _remove_by_source 异常记日志, 不再静默吞掉
M4: CLI 集合回退支持 MD_VECTOR_DB_COLLECTION 环境变量
L1: DEFAULT_CONFIG_PATH 自动从项目根目录解析
L2: ingest 命令内重复 import 移至模块顶部
L3: 新增死循环回归测试 + 密集分隔符分块测试
L4: 统一 logger 名称为 md-vector-db
删除旧版审计文档

测试: 48 passed
This commit is contained in:
2026-07-06 13:16:26 +08:00
parent 7e77c31e16
commit 832201186d
13 changed files with 815 additions and 50 deletions
+6 -6
View File
@@ -96,8 +96,8 @@ uv run md-vector-db serve --port 8000
## HTTP API
| 方法 | 路径 | 需要认证 | 说明 |
|------|------|----------|------|
| `GET` | `/` | — | 重定向到 `/docs` (Swagger UI) |
| ---------- | --------------------------------- | -------- | ----------------------------------------------------------------------------- |
| `GET` | `/` | — | 重定向到`/docs` (Swagger UI) |
| `GET` | `/api/v1/health` | — | 健康检查(验证 ChromaDB + Embedder 可用性) |
| `GET` | `/api/v1/collections` | — | 列出集合和已入库的源文件 |
| `POST` | `/api/v1/ingest` | API Key | 入库文档(`file_path``content` + `file_name`,可选 `collection` |
@@ -146,8 +146,8 @@ for r in resp.json()["results"]:
所有命令均支持 `--config/-c`(配置文件)、`--collection/-C`(集合名,默认 `default`)。
| 命令 | 说明 |
|------|------|
| `ingest <文件路径>` | 入库单个 .md 文件,支持 `-C` 指定集合 |
| --------------------------- | -------------------------------------------------------- |
| `ingest <文件路径>` | 入库单个 .md 文件,支持`-C` 指定集合 |
| `ingest-dir <目录路径>` | 递归入库目录下所有 .md 文件 |
| `search <查询> -k <数量>` | 语义检索,`-k` 默认 10、最大 100`--json` JSON 输出 |
| `stats` | 显示 chunks 总数、源文件列表 |
@@ -190,7 +190,7 @@ MD_VECTOR_API_KEY=secret # HTTP API 访问密钥(不设置则跳过认证
### Provider 速查
| provider | 默认 API 地址 | 默认模型 | 维度 |
|----------|--------------|---------|------|
| ------------- | ----------------------------- | -------------------------- | ---- |
| `local` | 本地(sentence-transformers | `BAAI/bge-small-zh-v1.5` | 512 |
| `openai` | `https://api.openai.com/v1` | `text-embedding-3-small` | 1536 |
| `dashscope` | 阿里云 DashScope | `text-embedding-v4` | 1536 |
@@ -261,7 +261,7 @@ uv run pytest tests/ -v --cov=src --cov-report=term-missing # 覆盖率
本地嵌入模式自动检测 CUDA 设备,优先使用 GPU。实测性能对比(RTX 4060 Laptopbge-small-zh-v1.5):
| 文件大小 | chunks | CPU 耗时 | GPU 耗时 |
|----------|--------|----------|----------|
| -------- | ------ | ------------ | -------- |
| 60KB | 320 | 数分钟至卡死 | 0.7s |
| 96KB | 729 | 卡死 | 1.8s |
+437
View File
@@ -0,0 +1,437 @@
# MCP 构建指南
本文档指导如何将 md-vector-db 封装为 MCPModel Context Protocol)服务。
参考实现:[PaddleOCR MCP Server](D:/Code/language/Python/PaddleOCR/ocr_mcp_server.py) — 使用 `FastMCP` + `Pydantic` 模式。
---
## PaddleOCR 的 MCP 模式(参考)
PaddleOCR 的 MCP 实现非常简洁,核心模式如下:
```python
from mcp.server.fastmcp import FastMCP, Context
from pydantic import BaseModel, Field
# 1. 一行创建 server
mcp = FastMCP("server_name")
# 2. Pydantic 定义参数模型(自动生成 inputSchema
class MyInput(BaseModel):
file_path: str = Field(..., description="参数说明", min_length=1)
option: bool = Field(default=False, description="可选参数")
# 3. @mcp.tool 装饰器注册工具,参数模型自动转换为 JSON Schema
@mcp.tool(name="tool_name", annotations={...})
async def tool_name(params: MyInput, ctx: Context) -> str:
await ctx.info("处理中...") # 日志
await ctx.report_progress(0.5, ...) # 进度
return "结果字符串"
# 4. 一行启动
if __name__ == "__main__":
mcp.run()
```
---
## md-vector-db 的 MCP Server 实现
### 1. 安装依赖
```bash
uv add "mcp>=1.0.0"
```
### 2. 创建 `src/mcp_server.py`
```python
"""md-vector-db MCP Server — 将向量知识库暴露为 MCP 工具.
基于 FastMCP + Pydantic 模式(参考 PaddleOCR/ocr_mcp_server.py)。
"""
import sys
from pathlib import Path
# 确保 src/ 在路径中
sys.path.insert(0, str(Path(__file__).parent))
from mcp.server.fastmcp import FastMCP, Context
from pydantic import BaseModel, Field, ConfigDict
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
# ── 初始化(进程启动时执行一次) ──
cfg = load_config()
db = VectorDB(persist_dir=cfg.chroma.persist_dir)
embedder = create_embedder(cfg.embed)
_searchers: dict[str, Searcher] = {}
_ingestors: dict[str, DocumentIngestor] = {}
def _get_searcher(collection: str) -> Searcher:
if collection not in _searchers:
_searchers[collection] = Searcher(db, embedder, collection)
return _searchers[collection]
def _get_ingestor(collection: str) -> DocumentIngestor:
if collection not in _ingestors:
_ingestors[collection] = DocumentIngestor(db, embedder, collection)
return _ingestors[collection]
# ── MCP Server ──
mcp = FastMCP("md-vector-db")
# ========== Pydantic 输入模型 ==========
class SearchInput(BaseModel):
"""语义检索参数."""
model_config = ConfigDict(str_strip_whitespace=True)
query: str = Field(
...,
description="自然语言查询,例如 'Git 分支管理策略''Docker Compose 多容器编排'",
min_length=1,
)
collection: str = Field(
default="obsidian_blog",
description="知识库集合名。可选值: obsidian_blog (博客笔记), default (测试数据)",
)
top_k: int = Field(
default=5,
ge=1,
le=50,
description="返回结果数量 (1-50)",
)
class IngestInput(BaseModel):
"""文档入库参数."""
model_config = ConfigDict(str_strip_whitespace=True)
content: str = Field(
...,
description="Markdown 格式的文档内容",
min_length=1,
)
file_name: str = Field(
...,
description="文档名(用于标识来源,如 'git-guide.md'",
min_length=1,
)
collection: str = Field(
default="default",
description="目标集合名",
)
class ListCollectionsInput(BaseModel):
"""列出集合(无参数)."""
pass
# ========== MCP 工具 ==========
@mcp.tool(
name="search_knowledge_base",
annotations={
"title": "语义搜索知识库",
"readOnlyHint": True,
"destructiveHint": False,
"idempotentHint": True,
"openWorldHint": True,
},
)
async def search_knowledge_base(params: SearchInput, ctx: Context) -> str:
"""语义检索知识库,返回最相关的 Markdown 文档片段。
基于 bge-small-zh-v1.5 模型 + ChromaDB,支持跨文档模糊语义匹配。
当前知识库包含 51 篇博客文章(3,611 chunks),涵盖 Git、Docker、深度学习、
LaTeX、AI 编程等主题。
Args:
params: SearchInput — query/collection/top_k
Returns:
格式化的 Markdown 搜索结果,包含相似度、来源、章节和内容
"""
searcher = _get_searcher(params.collection)
await ctx.info(f"搜索: {params.query} (集合: {params.collection})")
results = searcher.search(query=params.query, top_k=params.top_k)
if not results:
return f"未找到与 '{params.query}' 相关的结果(集合: {params.collection}"
lines = [f"## 搜索结果: {params.query}\n"]
for i, r in enumerate(results, 1):
lines.append(
f"### 结果 {i} (相似度: {r['score']:.2%})\n"
f"- **来源**: {r.get('source_file', '?')}\n"
f"- **章节**: {r.get('section_title', '')}\n"
f"\n{r['content']}\n"
)
await ctx.info(f"返回 {len(results)} 条结果")
return "\n".join(lines)
@mcp.tool(
name="list_collections",
annotations={
"title": "列出知识库集合",
"readOnlyHint": True,
"destructiveHint": False,
"idempotentHint": True,
"openWorldHint": True,
},
)
async def list_collections(params: ListCollectionsInput, ctx: Context) -> str:
"""列出所有知识库集合及文档统计."""
await ctx.info("列出集合...")
lines = ["## 知识库集合\n"]
try:
collections = db.client.list_collections()
if not collections:
lines.append("(无集合)")
for coll in collections:
count = coll.count()
# 获取源文件数
try:
sources = coll.get()["metadatas"]
unique_sources = len(set(
m.get("source_file", "") for m in sources if m
)) if sources else 0
except Exception:
unique_sources = "?"
lines.append(f"- **{coll.name}**: {count} chunks / {unique_sources} 文件")
except Exception as e:
lines.append(f"错误: {e}")
return "\n".join(lines)
@mcp.tool(
name="ingest_document",
annotations={
"title": "入库 Markdown 文档",
"readOnlyHint": False,
"destructiveHint": False,
"idempotentHint": True,
"openWorldHint": True,
},
)
async def ingest_document(params: IngestInput, ctx: Context) -> str:
"""入库 Markdown 文档到知识库。同名文件会覆盖旧版本。
Args:
params: IngestInput — content/file_name/collection
Returns:
入库结果摘要
"""
ingestor = _get_ingestor(params.collection)
await ctx.info(f"入库: {params.file_name} → 集合 {params.collection}")
n = ingestor.ingest_content(params.content, params.file_name)
return (
f"✅ 已入库: **{params.file_name}**\n"
f"- Chunks: {n}\n"
f"- 集合: `{params.collection}`"
)
# ── 入口 ──
if __name__ == "__main__":
mcp.run()
```
### 3. 注册入口
```toml
# pyproject.toml
[project.scripts]
md-vector-db = "src.cli.main:app"
md-vector-db-mcp = "src.mcp_server:mcp" # FastMCP 的 run() 通过入口调用
```
> 注意:FastMCP 用 `mcp.run()` 启动,入口指向模块的 `mcp` 对象或直接用 `python -m src.mcp_server`。
### 4. MCP 客户端配置
#### Claude Code`.claude/mcp.json`
```json
{
"mcpServers": {
"md-vector-db": {
"type": "stdio",
"command": "D:/Code/doing_exercises/programs/md-vector-db/.venv/Scripts/python.exe",
"args": [
"D:/Code/doing_exercises/programs/md-vector-db/src/mcp_server.py"
]
}
}
}
```
> PaddleOCR 也是用绝对路径 python + 脚本路径的方式,**不通过 `uv run`**,避免 uv 的环境解析覆盖 CUDA torch。
#### Claude Desktop`claude_desktop_config.json`
```json
{
"mcpServers": {
"md-vector-db": {
"command": "D:/Code/doing_exercises/programs/md-vector-db/.venv/Scripts/python.exe",
"args": [
"D:/Code/doing_exercises/programs/md-vector-db/src/mcp_server.py"
]
}
}
}
```
---
## FastMCP vs 低层 API 对比
| | FastMCP(推荐) | 低层 API(不推荐) |
| ----------- | -------------------------------------------------------- | -------------------------------------------------- |
| 创建 Server | `mcp = FastMCP("name")` | `server = Server("name")` |
| 注册工具 | `@mcp.tool()` 装饰器 | `@server.list_tools()` + `@server.call_tool()` |
| 参数定义 | **Pydantic `BaseModel`** → 自动生成 JSON Schema | 手写`inputSchema` dict |
| 进度/日志 | `ctx.info()` / `ctx.report_progress()` | 不支持(需手动实现) |
| 启动 | `mcp.run()` | `asyncio.run(stdio_server(...))` |
| 工具元数据 | `annotations` 字典(标准 MCP 注解) | 不支持 |
---
## 工具注解说明
MCP 规范定义 4 个注解,帮助客户端理解工具行为:
| 注解 | 含义 | 示例 |
| ------------------- | -------------- | ----------------------------- |
| `readOnlyHint` | 是否只读 | `True` — 搜索不修改数据 |
| `destructiveHint` | 是否破坏性 | `False` — 入库不会删除数据 |
| `idempotentHint` | 是否幂等 | `True` — 重复调用结果一致 |
| `openWorldHint` | 是否与外部交互 | `True` — 连接外部 ChromaDB |
---
## 设计要点总结
### 1. 用 Pydantic,不用手写 JSON Schema
```python
# ✅ FastMCP 方式 — Pydantic 自动生成 inputSchema
class SearchInput(BaseModel):
query: str = Field(..., description="...", min_length=1)
top_k: int = Field(default=5, ge=1, le=50)
@mcp.tool()
async def search(params: SearchInput, ctx: Context) -> str:
...
# ❌ 低层 API — 手写容易出错
@server.list_tools()
async def list_tools():
return [Tool(name="search", inputSchema={...手写 JSON Schema...})]
```
### 2. 用 `ctx` 报告进度
```python
await ctx.info("搜索: Docker 部署") # 日志
await ctx.report_progress(0.5, message="嵌入查询中...") # 进度条
```
### 3. 返回 Markdown 格式字符串
Claude 能正确渲染 Markdown,所以直接返回格式化的 Markdown:
```python
return f"## 搜索结果\n\n### 1. {title}\n{content}"
```
### 4. 异常处理 — 返回错误字符串而非抛异常
```python
try:
results = searcher.search(query, top_k)
except Exception as e:
return f"搜索失败: {e}"
```
### 5. 启动用绝对路径 Python,不用 `uv run`
```
command: .venv/Scripts/python.exe # 直接用 venv 的 python
args: [src/mcp_server.py] # 脚本路径
```
---
## 测试 MCP Server
### MCP Inspector
```bash
npx @modelcontextprotocol/inspector \
D:/Code/doing_exercises/programs/md-vector-db/.venv/Scripts/python.exe \
D:/Code/doing_exercises/programs/md-vector-db/src/mcp_server.py
```
### 手动 JSON-RPC 测试
```bash
echo '{"jsonrpc":"2.0","id":1,"method":"tools/list"}' | \
.venv/Scripts/python.exe src/mcp_server.py
```
---
## 部署到 NAS
### 1. 同步项目文件
```bash
rsync -avz --exclude '.venv' --exclude 'data' \
D:/Code/doing_exercises/programs/md-vector-db/ \
LHY@192.168.5.8:/volume2/办公/Code/md-vector-db/
```
### 2. NAS 上创建 venv 并安装
```bash
ssh LHY@192.168.5.8
cd /volume2/办公/Code/md-vector-db
uv sync # NAS 上不需要 GPUCPU torch 即可
```
### 3. 客户端配置(通过 SSH 隧道连接)
如果 Claude Code 运行在 Windows 上,MCP Server 在 NAS 上,需要 SSH 隧道:
```json
{
"mcpServers": {
"md-vector-db": {
"command": "ssh",
"args": [
"LHY@192.168.5.8",
"cd /volume2/办公/Code/md-vector-db && .venv/bin/python src/mcp_server.py"
]
}
}
}
```
@@ -0,0 +1,275 @@
# md-vector-db 代码架构审计报告
**审计日期**: 2026-07-06
**审计范围**: `src/` 全部模块 + `scripts/ingest_obsidian.py` + 测试
**审计方法**: 逐文件静态分析,按安全性 → 可靠性 → 设计 → 测试四个维度评审
---
## 总览
| 维度 | 评分 | 说明 |
|------|------|------|
| 安全性 | ⭐⭐⭐⭐ | 整体良好,认证/限流/路径防护到位 |
| 可靠性 | ⭐⭐⭐ | 有 2 处吞异常 + 1 处内存泄漏风险 |
| 设计 | ⭐⭐⭐ | 核心架构清晰,但 2 处封装泄漏 |
| 测试 | ⭐⭐⭐ | 46 个测试覆盖主要路径,缺回归测试 |
共发现 **11 个问题**3 个高、4 个中、4 个低。
---
## 高优先级问题
### H1. `config.py:75` — `AppConfig.__init__` 破坏 dataclass 构造
```python
# 问题代码 (config.py:66-79)
@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): # ← 覆盖了 dataclass 生成的 __init__
self.chroma = ChromaConfig(**kwargs.get("chroma", {}))
self.embed = EmbedConfig(**kwargs.get("embed", {}))
...
```
**影响**
- `AppConfig(chroma=ChromaConfig())` — 正常 dataclass 构造方式会报 `TypeError`
- `@dataclass``field(default_factory=...)` 全部被绕过
- 所有字段失去类型检查
**修复**
```python
@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)
@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", {})),
)
def load_config(path: str | None = None) -> AppConfig:
...
return AppConfig.from_dict(data)
```
**严重程度**: HIGH — 调用方如果用标准 dataclass 构造会直接崩溃。
---
### H2. `deps.py:47-58` — `list_collections_with_stats` 只列出内存中的集合
```python
# 问题代码
def list_collections_with_stats(self) -> list[dict]:
all_names = set(self._searchers.keys()) | set(self._ingestors.keys())
all_names.add(self.default_collection)
# ↑ 只包含被访问过的集合!ChromaDB 中已有的其他集合不会出现
```
**影响**:通过 CLI 或直接操作 ChromaDB 创建的集合,在 HTTP API 的 `/api/v1/collections` 中看不到。
**修复**:应从 ChromaDB 直接获取。
```python
def list_collections_with_stats(self) -> list[dict]:
result = []
for coll in self.db.client.list_collections():
result.append({"name": coll.name, "count": coll.count()})
return result
```
**严重程度**: HIGH — 功能缺陷,`/api/v1/collections` 返回不完整数据。
---
### H3. `auth.py:27` — RateLimiter 内存无限增长
```python
# 问题代码
class RateLimiter:
def __init__(self, max_requests=30, window_seconds=60):
self._store: dict[str, list[float]] = defaultdict(list)
def is_allowed(self, client_id: str) -> bool:
with self._lock:
records = self._store[client_id]
records[:] = [t for t in records if now - t < self.window]
# ↑ 只清理时间戳,不清理空 key
```
**影响**:每个新 IP 都在 `_store` 中永久保留一个 key,长期运行后字典膨胀。
**修复**:在清理过期时间戳后,删除空的 key:
```python
records[:] = [t for t in records if now - t < self.window]
if not records:
del self._store[client_id]
```
**严重程度**: HIGH — 长期运行会内存泄漏。
---
## 中优先级问题
### M1. `search.py:87-88` — `delete_by_source` 吞掉所有异常
```python
def delete_by_source(self, file_name: str) -> bool:
try:
with self.db.write_lock:
...
except Exception:
pass # ← 隐藏了 ChromaDB 损坏、磁盘满等严重错误
return False
```
**修复**:至少记录日志:
```python
except Exception:
logger.exception("删除文档失败: %s", file_name)
return False
```
---
### M2. `db.py` — 缺少 `list_collections()` 封装
`VectorDB` 没有暴露 `list_collections()`,导致:
- `deps.py` 无法从 `db` 对象获取集合列表
- MCP server 需要绕过封装直接访问 `db.client`
**修复**
```python
def list_collections(self) -> list:
return self.client.list_collections()
```
---
### M3. `ingest.py:237-245` — `_remove_by_source` 吞掉去重异常
```python
def _remove_by_source(self, file_name: str) -> None:
try:
...
except Exception:
pass # collection 为空时 get 可能抛异常
```
**影响**:如果去重失败(非空 collection),重复入库不会报错但会静默产生重复数据。
**修复**:区分 "collection 为空" 和真正的错误:
```python
except ValueError:
pass # collection 无数据时 ChromaDB 会抛 ValueError
except Exception:
logger.exception("去重检查失败: %s", file_name)
```
---
### M4. `cli/main.py:47-48` — 硬编码回退值
```python
def _get_default_collection() -> str:
return _cfg.chroma.collection_name if _cfg else "markdown_docs"
```
`deps.py:27-29` 中的 `default_collection` 逻辑不一致(deps.py 支持 `MD_VECTOR_DB_COLLECTION` 环境变量)。
---
## 低优先级问题
### L1. `config.py:16` — 相对路径 `DEFAULT_CONFIG_PATH`
```python
DEFAULT_CONFIG_PATH = "config.yaml" # 从不同 CWD 运行会找不到
```
**建议**:在 `load_config` 中基于项目根目录解析。
---
### L2. `cli/main.py:97-101` — 函数内重复 import
```python
for fp in file_paths:
from pathlib import Path as _Path # 每次循环都 import
...
import glob as _glob
```
**修复**:移到文件顶部。
---
### L3. 缺少死循环回归测试
`_split_single_paragraph` 修复(`start = max(start + 1, next_start)`)没有对应的单元测试:分隔符距 chunk 起点 < overlap 的场景。
```python
def test_separator_near_start_does_not_loop(self, splitter):
"""分隔符紧挨 start 时不会死循环 (回归测试)."""
# 构造: 前 overlap 字符内就有分隔符的情况
text = "" * 50 + "内容" * 300 # 分隔符密集在开头
chunks = splitter.split(f"# T\n{text}", source_file="test.md")
assert len(chunks) > 0 # 不卡死即通过
```
---
### L4. `deps.py:11` — Logger 名称
```python
logger = logging.getLogger(__name__) # __name__ = "src.server.deps"
```
`app.py:17``logging.getLogger("md-vector-db")` 不一致,日志输出时名称不统一。
---
## 架构评价
### 做得好的地方
- **依赖方向正确**: `config ← db ← embedder ← ingest/search ← server/cli`,核心层不依赖传输层
- **策略模式**: `Embedder` Protocol + 工厂函数 `create_embedder`,方便扩展新 Provider
- **依赖注入**: `AppState` + `FastAPI Depends`,可测试可替换
- **线程安全**: `threading.Lock` 保护 ChromaDB 写操作
- **安全**: API Key 认证 + 速率限制 + 路径遍历防护三重保护
- **MCP 友好**: CLI `--json` 输出,支持 stdin 输入
### 待改进
| 问题 | 说明 |
|------|------|
| 封装泄漏 | `db.client` 被外部直接访问(MCP server、deps.py),应通过 `VectorDB` 方法暴露 |
| AppConfig 构造 | 自定义 `__init__` 与 dataclass 冲突 |
| 集合发现 | 没有统一的 "列出所有集合" 方法,deps/MCP/CLI 各自实现 |
---
## 总结
| 优先级 | 数量 | 建议处理 |
|--------|------|----------|
| HIGH | 3 | H1 (AppConfig) 需立即修,H2/H3 尽快修 |
| MEDIUM | 4 | 逐个修,影响小但累积会出问题 |
| LOW | 4 | 择机修,不紧急 |
**整体评价**: 代码架构设计良好,核心逻辑清晰。主要问题集中在配置系统的 dataclass 使用不当、集合发现的封装泄漏、以及几处静默吞异常。修复 H1~H3 后即可达到生产就绪状态。
+7 -4
View File
@@ -3,6 +3,8 @@ from pathlib import Path
import sys
import io
import json
import os
import glob as _glob
from typing import Annotated
import typer
@@ -45,7 +47,10 @@ def _init_shared(config_path: str = DEFAULT_CONFIG_PATH):
def _get_default_collection() -> str:
return _cfg.chroma.collection_name if _cfg else "markdown_docs"
return os.environ.get(
"MD_VECTOR_DB_COLLECTION",
_cfg.chroma.collection_name if _cfg else "markdown_docs",
)
def _resolve_collection(collection: str | None) -> str:
@@ -94,10 +99,8 @@ def ingest(
total = 0
for fp in file_paths:
# 支持通配符 (shell 展开或 Python glob)
from pathlib import Path as _Path
p = _Path(fp)
p = Path(fp)
if "*" in fp or "?" in fp:
import glob as _glob
matches = _glob.glob(fp, recursive=True)
for m in matches:
c = ingestor.ingest_file(m)
+19 -8
View File
@@ -72,19 +72,30 @@ class AppConfig:
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", {}))
@classmethod
def from_dict(cls, data: dict) -> "AppConfig":
"""从字典构造 (用于 YAML 加载)."""
return cls(
chroma=ChromaConfig(**data.get("chroma", {})),
embed=EmbedConfig(**data.get("embed", {})),
chunk=ChunkConfig(**data.get("chunk", {})),
server=ServerConfig(**data.get("server", {})),
)
def load_config(path: str | None = None) -> AppConfig:
"""从 YAML 文件加载配置, 若文件不存在则返回默认配置."""
config_path = path or DEFAULT_CONFIG_PATH
if not Path(config_path).exists():
config_file = Path(config_path)
# 相对路径 → 从项目根目录解析
if not config_file.is_absolute():
root = Path(__file__).resolve().parent.parent.parent
_try = root / config_file
if _try.exists():
config_file = _try
if not config_file.exists():
return AppConfig()
with open(config_path, "r", encoding="utf-8") as f:
with open(config_file, "r", encoding="utf-8") as f:
data = yaml.safe_load(f) or {}
return AppConfig(**data)
return AppConfig.from_dict(data)
+4
View File
@@ -20,6 +20,10 @@ class VectorDB:
"""获取或创建 collection."""
return self.client.get_or_create_collection(name=name)
def list_collections(self) -> list[Collection]:
"""列出所有 collection."""
return self.client.list_collections()
def delete_collection(self, name: str) -> None:
"""删除 collection (线程安全)."""
with self._write_lock:
+6 -1
View File
@@ -1,4 +1,5 @@
"""Markdown 文档解析与入库模块."""
import logging
import re
from dataclasses import dataclass
from pathlib import Path
@@ -6,6 +7,8 @@ from pathlib import Path
from src.core.db import VectorDB
from src.core.embedder import Embedder, batch_embed
logger = logging.getLogger("md-vector-db")
class MarkdownSplitter:
"""Markdown 混合分块器:先按标题拆,超长再按段落拆."""
@@ -242,5 +245,7 @@ class DocumentIngestor:
)
if existing and existing["ids"]:
self.collection.delete(ids=existing["ids"])
except ValueError:
pass # collection 为空时 ChromaDB 抛 ValueError
except Exception:
pass # collection 为空时 get 可能抛异常
logger.exception("去重检查失败: %s", file_name)
+6 -1
View File
@@ -1,7 +1,10 @@
"""语义检索模块."""
import logging
from src.core.db import VectorDB
from src.core.embedder import Embedder
logger = logging.getLogger("md-vector-db")
class Searcher:
"""向量检索器."""
@@ -84,6 +87,8 @@ class Searcher:
if existing and existing["ids"]:
self.collection.delete(ids=existing["ids"])
return True
except ValueError:
pass # collection 为空时 ChromaDB 抛 ValueError
except Exception:
pass
logger.exception("删除文档失败: %s", file_name)
return False
+2
View File
@@ -32,6 +32,8 @@ class RateLimiter:
with self._lock:
records = self._store[client_id]
records[:] = [t for t in records if now - t < self.window]
if not records:
del self._store[client_id] # 清理空 key,防止内存泄漏
if len(records) >= self.max_requests:
return False
records.append(now)
+5 -8
View File
@@ -8,7 +8,7 @@ from src.core.embedder import create_embedder
from src.core.ingest import DocumentIngestor
from src.core.search import Searcher
logger = logging.getLogger(__name__)
logger = logging.getLogger("md-vector-db")
class AppState:
@@ -45,16 +45,13 @@ class AppState:
return self._ingestors[name]
def list_collections_with_stats(self) -> list[dict]:
"""列出所有 collection 及其统计."""
"""列出所有 collection 及其统计(直接从 ChromaDB 查询)."""
result = []
all_names = set(self._searchers.keys()) | set(self._ingestors.keys())
all_names.add(self.default_collection)
for name in sorted(all_names):
try:
coll = self.db.get_or_create_collection(name)
result.append({"name": name, "count": coll.count()})
for coll in self.db.client.list_collections():
result.append({"name": coll.name, "count": coll.count()})
except Exception:
result.append({"name": name, "count": 0})
logger.exception("列出集合失败")
return result
def is_healthy(self) -> dict:
+3 -3
View File
@@ -14,7 +14,7 @@ class TestEmbedConfig:
def test_local_mode_defaults(self):
"""默认 local 模式,带默认模型名."""
data = {"embed": {"mode": "local"}}
cfg = AppConfig(**data)
cfg = AppConfig.from_dict(data)
assert cfg.embed.mode == "local"
assert cfg.embed.local_model == "BAAI/bge-small-zh-v1.5"
assert cfg.embed.api_base == ""
@@ -28,7 +28,7 @@ class TestEmbedConfig:
"api_key": "sk-test",
}
}
cfg = AppConfig(**data)
cfg = AppConfig.from_dict(data)
assert cfg.embed.mode == "api"
assert cfg.embed.api_base == "https://api.openai.com/v1"
assert cfg.embed.api_key == "sk-test"
@@ -40,7 +40,7 @@ class TestChunkConfig:
def test_default_values(self):
"""分块默认值正确."""
data = {"chunk": {}}
cfg = AppConfig(**data)
cfg = AppConfig.from_dict(data)
assert cfg.chunk.max_size == 1000
assert cfg.chunk.overlap == 100
+26
View File
@@ -68,6 +68,32 @@ def hello():
all_content = " ".join(c["content"] for c in chunks)
assert "def hello()" in all_content
def test_separator_near_start_does_not_loop(self, splitter):
"""分隔符紧挨 start 时不会死循环 (回归测试, fix: start=max(start+1, next_start)).
场景: 超长段落中, 分隔符出现在距离 start 小于 overlap 的位置,
_split_single_paragraph 的 start 会回退为负数, str.rfind 负索引绕回导致死循环.
"""
# 100 个句号 + 大量内容 → 句号密集在开头且很近
text = "" * 80 + "内容文本" * 500
md = f"# 边界测试\n{text}"
chunks = splitter.split(md, source_file="edge.md")
# 不卡死即通过
assert len(chunks) > 0
# 验证内容完整
all_text = "".join(c["content"] for c in chunks)
assert "内容文本" in all_text
def test_dense_separators_in_long_para(self, splitter):
"""超长段落中分隔符密集分布也能正确分块."""
# ~2000 字符: 每段 20 个"内容文本" + "。",共 30 段
text = ""
for i in range(30):
text += "内容文本" * 20 + "" * (3 if i % 5 == 0 else 1) + "\n"
md = f"# 密集分隔符\n{text}"
chunks = splitter.split(md, source_file="dense.md")
assert len(chunks) > 1 # 超过 1000 字符应被拆分
class TestDocumentIngestor:
"""文档入库器测试."""