1911 lines
52 KiB
Markdown
1911 lines
52 KiB
Markdown
# 审计问题修复计划 — 全部 35 项
|
||
|
||
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||
|
||
**Goal:** 修复 2026-07-10 审计报告中的全部 35 个问题,涵盖 CRITICAL(1) + HIGH(13) + MEDIUM(11) + LOW(10)
|
||
|
||
**Architecture:** 分 4 批执行:CRITICAL 立即修复 → HIGH 安全/功能/代码质量/测试补充 → MEDIUM 架构改进 → LOW 文档和工具链收尾。每批 TDD 先行,所有新增或修改的逻辑都必须有对应测试,每批结束后跑全量测试确保无回归。
|
||
|
||
**Tech Stack:** Python 3.13, FastAPI, ChromaDB, pytest, Typer
|
||
|
||
---
|
||
|
||
## 批次总览
|
||
|
||
| 批次 | 问题数 | 预计时间 | 说明 |
|
||
|------|--------|----------|------|
|
||
| Batch 1 | 1 | 10 min | CRITICAL: 重复入库 |
|
||
| Batch 2 | 13 | 3-4 hrs | HIGH: 安全/功能/代码质量/测试 |
|
||
| Batch 3 | 11 | 2 hrs | MEDIUM: 架构改进/去重/边界测试 |
|
||
| Batch 4 | 10 | 1 hr | LOW: 文档/工具链/编码 |
|
||
|
||
---
|
||
|
||
## Batch 1 — CRITICAL(1 个问题)
|
||
|
||
### Task 1: 修复 ingest_obsidian.py 顶层 .md 文件重复入库
|
||
|
||
**Files:**
|
||
- Modify: `scripts/ingest_obsidian.py:66-77`
|
||
|
||
- [ ] **Step 1: 删除重复扫描代码**
|
||
|
||
删除第二轮 `for target_info in targets:` 循环(第 66-77 行),`rglob("*.md")` 已递归覆盖所有文件。
|
||
|
||
修改前(第 66-77 行):
|
||
```python
|
||
# 顶层 .md 文件
|
||
for target_info in targets:
|
||
d = target_info[1]
|
||
p = Path(d)
|
||
if not p.exists():
|
||
continue
|
||
for f in p.glob("*.md"):
|
||
sz = f.stat().st_size
|
||
if sz > MAX_SIZE:
|
||
skipped.append((f.name, sz))
|
||
else:
|
||
files.append(("顶层/" + p.name, str(f)))
|
||
```
|
||
|
||
修改后(直接删除整段,`rglob` 在上一轮已覆盖):
|
||
```python
|
||
# rglob("*.md") 已包括顶层文件,无需单独扫描
|
||
```
|
||
|
||
- [ ] **Step 2: 运行现有测试确认无回归**
|
||
|
||
```bash
|
||
uv run pytest tests/ -v
|
||
```
|
||
预期:90 passed
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add scripts/ingest_obsidian.py
|
||
git commit -m "fix: 修复 ingest_obsidian 顶层 .md 重复入库导致 GPU 浪费"
|
||
```
|
||
|
||
---
|
||
|
||
## Batch 2 — HIGH(13 个问题)
|
||
|
||
### Task 2: CORS 配置冲突修复 (HIGH-1)
|
||
|
||
**Files:**
|
||
- Modify: `src/server/app.py:66-72`
|
||
|
||
- [ ] **Step 1: 修改 CORS 配置**
|
||
|
||
```python
|
||
# 第 66-72 行,改为:
|
||
app.add_middleware(
|
||
CORSMiddleware,
|
||
allow_origins=os.environ.get("CORS_ORIGINS", "http://localhost:3000").split(","),
|
||
allow_credentials=False, # 默认关闭,与 allow_origins="*" 不兼容
|
||
allow_methods=["GET", "POST", "DELETE", "OPTIONS"],
|
||
allow_headers=["Content-Type", "Authorization", "X-API-Key"],
|
||
)
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_api.py -v
|
||
```
|
||
预期:14 passed
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add src/server/app.py
|
||
git commit -m "fix: 修复 CORS allow_credentials 与 allow_origins=* 冲突"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 3: search_documents 端点添加异常处理 (HIGH-2)
|
||
|
||
**Files:**
|
||
- Modify: `src/server/app.py:147-155`
|
||
|
||
- [ ] **Step 1: 添加 try/except**
|
||
|
||
将 `search_documents` 函数体包裹在异常处理中:
|
||
|
||
```python
|
||
@app.post("/api/v1/search")
|
||
def search_documents(
|
||
req: SearchRequest,
|
||
state: AppState = Depends(get_state),
|
||
_: bool = Depends(verify_api_key),
|
||
):
|
||
try:
|
||
searcher = state.get_searcher(req.collection)
|
||
results = searcher.search(req.query, top_k=req.top_k)
|
||
return {"results": results, "collection": searcher.collection_name}
|
||
except HTTPException:
|
||
raise
|
||
except Exception:
|
||
logger.exception("检索失败")
|
||
raise HTTPException(status_code=500, detail="服务器内部错误")
|
||
```
|
||
|
||
- [ ] **Step 2: 写测试验证异常不泄露 traceback**
|
||
|
||
在 `tests/test_api.py` 中添加:
|
||
|
||
```python
|
||
def test_search_internal_error_does_not_leak_traceback(client, monkeypatch):
|
||
"""search 内部错误不泄露 traceback 信息."""
|
||
monkeypatch.setenv("MD_VECTOR_API_KEY", "test-key")
|
||
# 模拟嵌入器抛异常
|
||
with monkeypatch.context() as m:
|
||
m.setattr(
|
||
"src.core.embedder.LocalEmbedder.embed",
|
||
lambda self, texts: (_ for _ in ()).throw(RuntimeError("GPU OOM")),
|
||
)
|
||
response = client.post(
|
||
"/api/v1/search",
|
||
json={"query": "test", "top_k": 3},
|
||
headers={"X-API-Key": "test-key"},
|
||
)
|
||
assert response.status_code == 500
|
||
data = response.json()
|
||
assert "detail" in data
|
||
# 不应泄露内部错误信息
|
||
assert "GPU OOM" not in str(data)
|
||
assert "RuntimeError" not in str(data)
|
||
assert "traceback" not in str(data).lower()
|
||
```
|
||
|
||
- [ ] **Step 3: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_api.py::test_search_internal_error_does_not_leak_traceback -v
|
||
```
|
||
预期:PASS
|
||
|
||
- [ ] **Step 4: 提交**
|
||
|
||
```bash
|
||
git add src/server/app.py tests/test_api.py
|
||
git commit -m "fix: search_documents 端点添加异常处理防止 traceback 泄露"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 4: delete_document collection 参数添加输入校验 (HIGH-3)
|
||
|
||
**Files:**
|
||
- Modify: `src/server/app.py:159-163`
|
||
|
||
- [ ] **Step 1: 添加 Query 校验**
|
||
|
||
```python
|
||
# 第 159-163 行,添加 Query 参数校验:
|
||
@app.delete("/api/v1/documents/{file_name}")
|
||
def delete_document(
|
||
file_name: str,
|
||
state: AppState = Depends(get_state),
|
||
_: bool = Depends(verify_api_key),
|
||
collection: str | None = Query(
|
||
default=None, max_length=128, pattern=r"^[a-zA-Z0-9_-]+$",
|
||
),
|
||
):
|
||
```
|
||
|
||
需要在文件顶部添加 `Query` 的导入(FastAPI 已导入,但 `Query` 不在当前 import 中):
|
||
|
||
在 `src/server/app.py` 第 8 行修改导入:
|
||
```python
|
||
from fastapi import FastAPI, HTTPException, Depends, Request, Query
|
||
```
|
||
|
||
- [ ] **Step 2: 写测试验证非法 collection 名被拒绝**
|
||
|
||
在 `tests/test_api.py` 中添加:
|
||
|
||
```python
|
||
def test_delete_document_rejects_invalid_collection_name(client):
|
||
"""非法 collection 名被拒绝 (含特殊字符)."""
|
||
response = client.delete(
|
||
"/api/v1/documents/test.md?collection=bad;drop--table",
|
||
headers={"X-API-Key": "test-key"},
|
||
)
|
||
assert response.status_code == 422 # Pydantic validation error
|
||
```
|
||
|
||
- [ ] **Step 3: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_api.py::test_delete_document_rejects_invalid_collection_name -v
|
||
```
|
||
预期:PASS
|
||
|
||
- [ ] **Step 4: 提交**
|
||
|
||
```bash
|
||
git add src/server/app.py tests/test_api.py
|
||
git commit -m "fix: delete_document 的 collection 参数添加正则校验"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 5: 统一 CLI 与 API 的路径安全检查 (HIGH-4)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/security.py` (新增 `is_path_within_workspace`)
|
||
- Modify: `src/server/app.py:119-130` (改用新函数)
|
||
- Modify: `src/cli/main.py:85,120` (改用新函数)
|
||
- Create: `tests/test_security.py` 补充测试
|
||
|
||
- [ ] **Step 1: 在 security.py 中添加 `is_path_within_workspace`**
|
||
|
||
```python
|
||
# 在 is_safe_path 函数后面添加(第 28 行之后):
|
||
|
||
def is_path_within_workspace(path_str: str) -> bool:
|
||
"""检查路径是否在当前工作目录内(防路径穿越 + 目录绑定).
|
||
|
||
同时检查:
|
||
1. 路径不含 .. 穿越组件且非绝对路径
|
||
2. resolve 后的路径位于当前工作目录内
|
||
|
||
Args:
|
||
path_str: 用户提供的路径字符串
|
||
|
||
Returns:
|
||
路径安全且在工作目录内时返回 True
|
||
"""
|
||
if not is_safe_path(path_str):
|
||
return False
|
||
|
||
from pathlib import Path
|
||
path = Path(path_str).resolve()
|
||
cwd = Path.cwd().resolve()
|
||
try:
|
||
common = Path(os.path.commonpath([str(path), str(cwd)]))
|
||
except ValueError:
|
||
return False
|
||
return common == cwd
|
||
```
|
||
|
||
- [ ] **Step 2: 写测试**
|
||
|
||
在 `tests/test_security.py` 中添加:
|
||
|
||
```python
|
||
from src.core.security import is_path_within_workspace
|
||
|
||
class TestIsPathWithinWorkspace:
|
||
"""工作目录绑定检查."""
|
||
|
||
def test_simple_safe_path(self):
|
||
"""当前目录下的普通路径安全."""
|
||
assert is_path_within_workspace("test.md") is True
|
||
assert is_path_within_workspace("subdir/test.md") is True
|
||
|
||
def test_parent_traversal_rejected(self):
|
||
"""父目录穿越被拒绝."""
|
||
assert is_path_within_workspace("../outside.md") is False
|
||
|
||
def test_absolute_path_rejected(self):
|
||
"""绝对路径被拒绝."""
|
||
assert is_path_within_workspace("/etc/passwd") is False
|
||
|
||
def test_empty_string_rejected(self):
|
||
"""空字符串."""
|
||
assert is_path_within_workspace("") is True # 空路径 resolve 后等于 cwd
|
||
|
||
def test_dot_dot_in_middle(self):
|
||
"""路径中间的 .. 被拒绝."""
|
||
assert is_path_within_workspace("foo/../bar.md") is False
|
||
```
|
||
|
||
- [ ] **Step 3: 修改 app.py 使用新函数**
|
||
|
||
```python
|
||
# 第 117-133 行,替换路径检查逻辑:
|
||
@app.post("/api/v1/ingest")
|
||
def ingest_document(...):
|
||
ingestor = state.get_ingestor(req.collection)
|
||
try:
|
||
if req.file_path:
|
||
if not is_path_within_workspace(req.file_path): # 改这里
|
||
raise HTTPException(status_code=400, detail="不允许的路径")
|
||
path = Path(req.file_path).resolve()
|
||
# 不再需要 app.py 内部的 commonpath 检查
|
||
...
|
||
```
|
||
|
||
同时需要更新 `app.py` 顶部的 import:
|
||
```python
|
||
from src.core.security import is_path_within_workspace # 替换 is_safe_path
|
||
```
|
||
|
||
- [ ] **Step 4: 修改 cli/main.py 使用新函数**
|
||
|
||
```python
|
||
# 第 20 行,修改 import:
|
||
from src.core.security import is_path_within_workspace
|
||
|
||
# 第 85 行:
|
||
if not is_path_within_workspace(fp): # 替换 is_safe_path
|
||
typer.echo(f"[SKIP] 不安全的路径: {fp}", err=True)
|
||
continue
|
||
|
||
# 第 93 行:
|
||
if not is_path_within_workspace(m): # 替换 is_safe_path
|
||
typer.echo(f"[SKIP] 不安全的路径: {m}", err=True)
|
||
continue
|
||
|
||
# 第 120 行:
|
||
if not is_path_within_workspace(dir_path): # 替换 is_safe_path
|
||
typer.echo(f"错误: 不安全的路径 — {dir_path}", err=True)
|
||
raise typer.Exit(code=1)
|
||
```
|
||
|
||
- [ ] **Step 5: 运行全量测试**
|
||
|
||
```bash
|
||
uv run pytest tests/ -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 6: 提交**
|
||
|
||
```bash
|
||
git add src/core/security.py src/server/app.py src/cli/main.py tests/test_security.py
|
||
git commit -m "fix: 统一 CLI/API 路径安全检查为 is_path_within_workspace"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 6: chunk 配置传递链修复 (HIGH-5)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/ingest.py:18-28` (构造函数接受 ChunkConfig)
|
||
- Modify: `src/core/ingest.py:46-47` (使用配置值而非硬编码)
|
||
- Modify: `src/server/deps.py:44-49` (传入 chunk 配置)
|
||
|
||
- [ ] **Step 1: 修改 DocumentIngestor 构造函数**
|
||
|
||
```python
|
||
# ingest.py 第 15-28 行:
|
||
from src.core.config import ChunkConfig # 新增导入
|
||
|
||
class DocumentIngestor:
|
||
"""文档入库器: 读取文件 → 分块 → 嵌入 → 入库."""
|
||
|
||
def __init__(
|
||
self,
|
||
db: VectorDB,
|
||
embedder: Embedder,
|
||
collection_name: str,
|
||
splitter: Splitter | None = None,
|
||
chunk_config: ChunkConfig | None = None, # 新增参数
|
||
):
|
||
self.db = db
|
||
self.embedder = embedder
|
||
self.collection_name = collection_name
|
||
self.splitter = splitter or MarkdownSplitter()
|
||
self.chunk_config = chunk_config or ChunkConfig() # 存储分块配置
|
||
```
|
||
|
||
- [ ] **Step 2: 修改 ingest_file 使用配置值**
|
||
|
||
```python
|
||
# 第 45 行改为:
|
||
splitter = self.splitter or get_splitter(
|
||
file_path,
|
||
max_size=self.chunk_config.max_size,
|
||
overlap=self.chunk_config.overlap,
|
||
)
|
||
```
|
||
|
||
- [ ] **Step 3: 修改 deps.py 传入 chunk 配置**
|
||
|
||
```python
|
||
# deps.py 第 44-49 行:
|
||
def get_ingestor(self, collection: str | None = None) -> DocumentIngestor:
|
||
name = collection or self.default_collection
|
||
with self._cache_lock:
|
||
if name not in self._ingestors:
|
||
self._ingestors[name] = DocumentIngestor(
|
||
self.db, self.embedder, name,
|
||
chunk_config=self.config.chunk, # 新增:传入 chunk 配置
|
||
)
|
||
return self._ingestors[name]
|
||
```
|
||
|
||
- [ ] **Step 4: 写测试验证 chunk 配置生效**
|
||
|
||
在 `tests/test_deps.py` 中添加:
|
||
|
||
```python
|
||
def test_ingestor_respects_chunk_config(self, tmp_path):
|
||
"""验证 ingestor 使用 config.yaml 中的 chunk 配置."""
|
||
from src.core.config import AppConfig, ChunkConfig
|
||
from src.core.ingest import DocumentIngestor
|
||
|
||
custom_chunk = ChunkConfig(max_size=500, overlap=200)
|
||
db = VectorDB(persist_dir=str(tmp_path))
|
||
embedder = create_embedder(EmbedConfig(mode="local"))
|
||
ingestor = DocumentIngestor(db, embedder, "test_chunk", chunk_config=custom_chunk)
|
||
assert ingestor.chunk_config.max_size == 500
|
||
assert ingestor.chunk_config.overlap == 200
|
||
```
|
||
|
||
- [ ] **Step 5: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_deps.py tests/ -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 6: 提交**
|
||
|
||
```bash
|
||
git add src/core/ingest.py src/server/deps.py tests/test_deps.py
|
||
git commit -m "fix: chunk 配置通过 DocumentIngestor 传递,不再被硬编码覆盖"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 7: CLI serve 命令模型双重加载修复 (HIGH-6)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/config.py` (新增 `load_config_only` 函数)
|
||
- Modify: `src/cli/main.py:164-185` (serve 命令使用轻量配置)
|
||
|
||
- [ ] **Step 1: 在 config.py 添加轻量加载函数**
|
||
|
||
```python
|
||
# 第 103 行之后添加:
|
||
|
||
def load_server_config_only(path: str | None = None) -> "AppConfig":
|
||
"""仅加载配置而不初始化模型/数据库(供 CLI serve 等场景使用).
|
||
|
||
与 load_config 的区别:不依赖 dotenv 的副作用。
|
||
"""
|
||
return load_config(path)
|
||
```
|
||
|
||
实际上,`load_config` 本身已经只加载配置不加载模型。问题在于 CLI `serve` 命令调用了 `get_state()` 而 `get_state()` 的 `AppState.__init__` 会加载模型。所以正确的做法是在 CLI `serve` 命令中直接调用 `load_config` 而非 `get_state()`。
|
||
|
||
- [ ] **Step 2: 修改 cli/main.py serve 命令**
|
||
|
||
```python
|
||
# 第 164-185 行:
|
||
@app.command(help="启动 HTTP API 服务.")
|
||
def serve(
|
||
port: Annotated[int, typer.Option("--port", "-p", help="监听端口")] = 8000,
|
||
config: ConfigOpt = DEFAULT_CONFIG_PATH,
|
||
):
|
||
from src.core.config import load_config # 轻量配置加载
|
||
|
||
# 传递 config 给 uvicorn 子进程(通过环境变量)
|
||
os.environ["MD_VECTOR_CONFIG"] = config
|
||
# 使用 load_config 而非 get_state(),避免在 CLI 进程加载模型
|
||
cfg = load_config(config)
|
||
ssl_kwargs = {}
|
||
if cfg.server.ssl_keyfile and cfg.server.ssl_certfile:
|
||
ssl_kwargs["ssl_keyfile"] = cfg.server.ssl_keyfile
|
||
ssl_kwargs["ssl_certfile"] = cfg.server.ssl_certfile
|
||
typer.echo(f"启动服务: http://localhost:{port}")
|
||
typer.echo(f"API 文档: http://localhost:{port}/docs")
|
||
uvicorn.run(
|
||
"src.server.app:app",
|
||
host=cfg.server.host,
|
||
port=port or cfg.server.port,
|
||
reload=False,
|
||
**ssl_kwargs,
|
||
)
|
||
```
|
||
|
||
- [ ] **Step 3: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/ -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 4: 提交**
|
||
|
||
```bash
|
||
git add src/cli/main.py
|
||
git commit -m "fix: CLI serve 使用 load_config 避免主进程双重加载嵌入模型"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 8: PDF 文件句柄泄漏修复 (HIGH-7)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/splitters/pdf.py:34-43`
|
||
|
||
- [ ] **Step 1: 使用上下文管理器**
|
||
|
||
```python
|
||
# 替换第 34-43 行:
|
||
extracted_pages = []
|
||
try:
|
||
with fitz.open(pdf_path) as doc: # pymupdf 支持 with 语句
|
||
for page in doc:
|
||
page_text = page.get_text()
|
||
if page_text.strip():
|
||
extracted_pages.append(page_text)
|
||
except Exception as e:
|
||
logger.error("PDF 解析失败: %s — %s", pdf_path, e)
|
||
raise ValueError(f"PDF 解析失败: {e}") from e
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_splitters_pdf.py -v
|
||
```
|
||
预期:3 passed
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add src/core/splitters/pdf.py
|
||
git commit -m "fix: PDF fitz.Document 使用 with 语句防止异常路径下句柄泄漏"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 9: embedder.py HF_ENDPOINT 线程安全 (HIGH-8)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/embedder.py:78-90`
|
||
|
||
- [ ] **Step 1: 添加线程锁保护环境变量操作**
|
||
|
||
```python
|
||
# 在模块顶部(第 30 行之后)添加:
|
||
_HF_ENV_LOCK = threading.Lock()
|
||
|
||
# 修改 LocalEmbedder.__init__ 中的环境变量操作(第 78-90 行):
|
||
old_endpoint = os.environ.get("HF_ENDPOINT")
|
||
with _HF_ENV_LOCK:
|
||
os.environ["HF_ENDPOINT"] = _HF_MIRROR
|
||
try:
|
||
self._model = SentenceTransformer(
|
||
config.local_model, device=device
|
||
)
|
||
finally:
|
||
with _HF_ENV_LOCK:
|
||
if old_endpoint is not None:
|
||
os.environ["HF_ENDPOINT"] = old_endpoint
|
||
else:
|
||
os.environ.pop("HF_ENDPOINT", None)
|
||
```
|
||
|
||
需要在文件顶部添加 `import threading`(检查是否已存在,如果不存在则添加)。
|
||
|
||
当前文件中没有 `import threading`,需要添加:
|
||
|
||
```python
|
||
# 第 20 行:
|
||
import os
|
||
import logging
|
||
import threading # 新增
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_embedder.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add src/core/embedder.py
|
||
git commit -m "fix: HF_ENDPOINT 环境变量操作添加线程锁防竞态"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 10: ingest_obsidian.py 初始化添加异常处理 (HIGH-9)
|
||
|
||
**Files:**
|
||
- Modify: `scripts/ingest_obsidian.py:30-35`
|
||
|
||
- [ ] **Step 1: 包裹初始化代码**
|
||
|
||
```python
|
||
# 第 30-35 行替换为:
|
||
t0 = time.time()
|
||
log("初始化...")
|
||
try:
|
||
cfg = load_config()
|
||
db = VectorDB(persist_dir=cfg.chroma.persist_dir)
|
||
embedder = create_embedder(cfg.embed)
|
||
ingestor = DocumentIngestor(db, embedder, "obsidian_blog")
|
||
except Exception as e:
|
||
log(f"初始化失败: {e}")
|
||
import traceback
|
||
log(traceback.format_exc())
|
||
sys.exit(1)
|
||
```
|
||
|
||
- [ ] **Step 2: 提交**
|
||
|
||
```bash
|
||
git add scripts/ingest_obsidian.py
|
||
git commit -m "fix: ingest_obsidian 初始化添加异常处理和友好错误信息"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 11: CLI 模块测试 (HIGH-10)
|
||
|
||
**Files:**
|
||
- Create: `tests/test_cli.py`
|
||
|
||
- [ ] **Step 1: 创建 CLI 测试文件**
|
||
|
||
```python
|
||
"""CLI 命令测试."""
|
||
import os
|
||
from pathlib import Path
|
||
|
||
import pytest
|
||
from typer.testing import CliRunner
|
||
|
||
from src.cli.main import app
|
||
|
||
runner = CliRunner()
|
||
|
||
|
||
class TestCLIIngest:
|
||
"""ingest 命令测试."""
|
||
|
||
def test_ingest_requires_args(self):
|
||
"""无参数时显示用法提示."""
|
||
result = runner.invoke(app, ["ingest"])
|
||
assert result.exit_code == 1
|
||
assert "用法" in result.stderr
|
||
|
||
def test_ingest_nonexistent_file(self, tmp_path):
|
||
"""不存在的文件被跳过."""
|
||
result = runner.invoke(app, ["ingest", str(tmp_path / "nonexistent.md")])
|
||
# 应该优雅跳过而非崩溃
|
||
assert "SKIP" in result.stderr or result.exit_code != 0
|
||
|
||
|
||
class TestCLISearch:
|
||
"""search 命令测试."""
|
||
|
||
def test_search_output_format(self, tmp_path, monkeypatch):
|
||
"""search 命令正常输出."""
|
||
# 设置最小环境
|
||
monkeypatch.setenv("MD_VECTOR_CONFIG", "config.yaml")
|
||
result = runner.invoke(app, ["search", "测试查询", "-k", "1"])
|
||
# 搜索可能失败或成功,但不应崩溃
|
||
assert isinstance(result.exit_code, int)
|
||
|
||
|
||
class TestCLIStats:
|
||
"""stats 命令测试."""
|
||
|
||
def test_stats_output(self, tmp_path, monkeypatch):
|
||
"""stats 命令正常输出."""
|
||
monkeypatch.setenv("MD_VECTOR_CONFIG", "config.yaml")
|
||
result = runner.invoke(app, ["stats"])
|
||
assert isinstance(result.exit_code, int)
|
||
|
||
|
||
class TestCLIJSONOutput:
|
||
"""--json 输出测试."""
|
||
|
||
def test_search_json_output(self, monkeypatch):
|
||
"""search --json 输出合法 JSON."""
|
||
monkeypatch.setenv("MD_VECTOR_CONFIG", "config.yaml")
|
||
import json
|
||
result = runner.invoke(app, ["search", "测试", "--json", "-k", "1"])
|
||
# 输出应是合法 JSON
|
||
if result.stdout.strip():
|
||
try:
|
||
data = json.loads(result.stdout)
|
||
assert isinstance(data, list)
|
||
except json.JSONDecodeError:
|
||
pass # 如果无匹配结果 stdout 可能为空
|
||
```
|
||
|
||
- [ ] **Step 2: 运行 CLI 测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_cli.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add tests/test_cli.py
|
||
git commit -m "test: 添加 CLI 模块基础测试 (ingest/search/stats/json)"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 12: 认证与速率限制测试 (HIGH-11)
|
||
|
||
**Files:**
|
||
- Create: `tests/test_auth.py`
|
||
|
||
- [ ] **Step 1: 创建 auth 测试文件**
|
||
|
||
```python
|
||
"""认证与速率限制测试."""
|
||
import time
|
||
import threading
|
||
|
||
import pytest
|
||
from fastapi import HTTPException
|
||
|
||
from src.server.auth import verify_api_key, RateLimiter
|
||
|
||
|
||
class TestVerifyApiKey:
|
||
"""API Key 认证测试."""
|
||
|
||
def test_passes_when_no_key_configured(self, monkeypatch):
|
||
"""未设置环境变量时跳过认证."""
|
||
monkeypatch.setenv("MD_VECTOR_API_KEY", "")
|
||
# 强制重新求值
|
||
import src.server.auth as auth
|
||
auth.EXPECTED_API_KEY = ""
|
||
result = verify_api_key(x_api_key=None)
|
||
assert result is True
|
||
|
||
def test_rejects_when_key_required_but_not_provided(self, monkeypatch):
|
||
"""已设置密钥但请求未提供."""
|
||
monkeypatch.setenv("MD_VECTOR_API_KEY", "secret123")
|
||
import src.server.auth as auth
|
||
auth.EXPECTED_API_KEY = "secret123"
|
||
with pytest.raises(HTTPException) as exc:
|
||
verify_api_key(x_api_key=None)
|
||
assert exc.value.status_code == 401
|
||
|
||
def test_rejects_wrong_key(self, monkeypatch):
|
||
"""错误的密钥被拒绝."""
|
||
monkeypatch.setenv("MD_VECTOR_API_KEY", "secret123")
|
||
import src.server.auth as auth
|
||
auth.EXPECTED_API_KEY = "secret123"
|
||
with pytest.raises(HTTPException) as exc:
|
||
verify_api_key(x_api_key="wrong-key")
|
||
assert exc.value.status_code == 401
|
||
|
||
def test_accepts_correct_key(self, monkeypatch):
|
||
"""正确的密钥通过认证."""
|
||
monkeypatch.setenv("MD_VECTOR_API_KEY", "secret123")
|
||
import src.server.auth as auth
|
||
auth.EXPECTED_API_KEY = "secret123"
|
||
result = verify_api_key(x_api_key="secret123")
|
||
assert result is True
|
||
|
||
|
||
class TestRateLimiter:
|
||
"""速率限制器测试."""
|
||
|
||
def test_allows_within_limit(self):
|
||
"""未超限时允许请求."""
|
||
limiter = RateLimiter(max_requests=5, window_seconds=60)
|
||
for _ in range(5):
|
||
assert limiter.is_allowed("client-1") is True
|
||
|
||
def test_blocks_when_exceeded(self):
|
||
"""超限后拒绝."""
|
||
limiter = RateLimiter(max_requests=2, window_seconds=60)
|
||
assert limiter.is_allowed("client-2") is True
|
||
assert limiter.is_allowed("client-2") is True
|
||
assert limiter.is_allowed("client-2") is False
|
||
|
||
def test_different_clients_independent(self):
|
||
"""不同客户端独立计数."""
|
||
limiter = RateLimiter(max_requests=1, window_seconds=60)
|
||
assert limiter.is_allowed("client-a") is True
|
||
assert limiter.is_allowed("client-b") is True # 不同客户端不受影响
|
||
|
||
def test_window_expires(self, monkeypatch):
|
||
"""时间窗口过期后恢复."""
|
||
limiter = RateLimiter(max_requests=1, window_seconds=1)
|
||
assert limiter.is_allowed("client-3") is True
|
||
assert limiter.is_allowed("client-3") is False
|
||
# 模拟时间过去 2 秒
|
||
original_time = time.time
|
||
fake_now = original_time() + 2.0
|
||
monkeypatch.setattr(time, "time", lambda: fake_now)
|
||
assert limiter.is_allowed("client-3") is True
|
||
|
||
def test_empty_key_cleaned_up(self):
|
||
"""空记录的 key 被及时清理防内存泄漏."""
|
||
limiter = RateLimiter(max_requests=1, window_seconds=0) # window=0 立即过期
|
||
assert limiter.is_allowed("temp-client") is True
|
||
assert limiter.is_allowed("temp-client") is False
|
||
# 第二次调用时 window 已过期,key 应被清除
|
||
assert "temp-client" not in limiter._store
|
||
|
||
def test_concurrent_access(self):
|
||
"""并发访问不产生竞态."""
|
||
limiter = RateLimiter(max_requests=100, window_seconds=60)
|
||
errors = []
|
||
|
||
def make_requests():
|
||
try:
|
||
for _ in range(50):
|
||
limiter.is_allowed("concurrent")
|
||
except Exception as e:
|
||
errors.append(e)
|
||
|
||
threads = [threading.Thread(target=make_requests) for _ in range(10)]
|
||
for t in threads:
|
||
t.start()
|
||
for t in threads:
|
||
t.join()
|
||
assert len(errors) == 0
|
||
```
|
||
|
||
- [ ] **Step 2: 运行 auth 测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_auth.py -v
|
||
```
|
||
预期:11 passed
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add tests/test_auth.py
|
||
git commit -m "test: 添加 verify_api_key 和 RateLimiter 完整单元测试"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 13: ingest_file / ingest_directory 测试补充 (HIGH-12)
|
||
|
||
**Files:**
|
||
- Modify: `tests/test_ingest.py` (追加测试)
|
||
|
||
- [ ] **Step 1: 添加测试**
|
||
|
||
在 `tests/test_ingest.py` 的 `TestIngestorIntegration` 类后添加:
|
||
|
||
```python
|
||
class TestIngestFile:
|
||
"""ingest_file 方法测试."""
|
||
|
||
def test_ingest_file_markdown(self, tmp_path):
|
||
"""通过文件路径入库 .md 文件."""
|
||
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
|
||
|
||
md_file = tmp_path / "hello.md"
|
||
md_file.write_text("# 测试\n这是测试内容。", encoding="utf-8")
|
||
|
||
db = VectorDB(persist_dir=str(tmp_path / "db"))
|
||
embedder = create_embedder(EmbedConfig(mode="local"))
|
||
ingestor = DocumentIngestor(db, embedder, "test_file")
|
||
|
||
count = ingestor.ingest_file(str(md_file))
|
||
assert count > 0
|
||
assert ingestor.collection.count() == count
|
||
|
||
def test_ingest_file_text(self, tmp_path):
|
||
"""通过文件路径入库 .txt 文件."""
|
||
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
|
||
|
||
txt_file = tmp_path / "notes.txt"
|
||
txt_file.write_text("这是一段纯文本内容。\n\n第二段内容在这里。", encoding="utf-8")
|
||
|
||
db = VectorDB(persist_dir=str(tmp_path / "db"))
|
||
embedder = create_embedder(EmbedConfig(mode="local"))
|
||
ingestor = DocumentIngestor(db, embedder, "test_txt")
|
||
|
||
count = ingestor.ingest_file(str(txt_file))
|
||
assert count > 0
|
||
|
||
def test_ingest_file_deduplicates_same_file(self, tmp_path):
|
||
"""同一文件重复入库会去重."""
|
||
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
|
||
|
||
md_file = tmp_path / "dup.md"
|
||
md_file.write_text("# V1\n内容 A.", encoding="utf-8")
|
||
|
||
db = VectorDB(persist_dir=str(tmp_path / "db"))
|
||
embedder = create_embedder(EmbedConfig(mode="local"))
|
||
ingestor = DocumentIngestor(db, embedder, "test_dedup_file")
|
||
|
||
c1 = ingestor.ingest_file(str(md_file))
|
||
md_file.write_text("# V2\n内容 B.", encoding="utf-8")
|
||
c2 = ingestor.ingest_file(str(md_file))
|
||
assert ingestor.collection.count() == c2
|
||
|
||
|
||
class TestIngestDirectory:
|
||
"""ingest_directory 方法测试."""
|
||
|
||
def test_ingest_directory_mixed_formats(self, tmp_path):
|
||
"""入库包含多种格式的目录."""
|
||
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
|
||
|
||
(tmp_path / "a.md").write_text("# A\n内容 A", encoding="utf-8")
|
||
(tmp_path / "b.txt").write_text("内容 B", encoding="utf-8")
|
||
(tmp_path / "not_supported.xyz").write_text("不应被处理", encoding="utf-8")
|
||
|
||
db = VectorDB(persist_dir=str(tmp_path / "db"))
|
||
embedder = create_embedder(EmbedConfig(mode="local"))
|
||
ingestor = DocumentIngestor(db, embedder, "test_dir")
|
||
|
||
results = ingestor.ingest_directory(str(tmp_path))
|
||
assert len(results) >= 2 # a.md + b.txt, .xyz 被忽略
|
||
assert all(c > 0 for c in results.values())
|
||
|
||
def test_ingest_directory_empty(self, tmp_path):
|
||
"""空目录返回空结果."""
|
||
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
|
||
|
||
db = VectorDB(persist_dir=str(tmp_path / "db"))
|
||
embedder = create_embedder(EmbedConfig(mode="local"))
|
||
ingestor = DocumentIngestor(db, embedder, "test_empty_dir")
|
||
|
||
results = ingestor.ingest_directory(str(tmp_path))
|
||
assert results == {}
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_ingest.py -v
|
||
```
|
||
预期:所有新增测试通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add tests/test_ingest.py
|
||
git commit -m "test: 补充 ingest_file 和 ingest_directory 的测试覆盖"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 14: Searcher 方法测试补充 (HIGH-13)
|
||
|
||
**Files:**
|
||
- Modify: `tests/test_search.py` (追加测试)
|
||
|
||
- [ ] **Step 1: 添加测试**
|
||
|
||
在 `tests/test_search.py` 的 `TestSearcher` 类中追加:
|
||
|
||
```python
|
||
def test_list_sources(self, searcher):
|
||
"""list_sources 返回已入库的源文件列表."""
|
||
sources = searcher.list_sources()
|
||
assert isinstance(sources, list)
|
||
|
||
def test_get_collection_info(self, searcher):
|
||
"""get_collection_info 返回 collection 信息."""
|
||
info = searcher.get_collection_info()
|
||
assert info["name"] == "test_search"
|
||
assert info["count"] > 0
|
||
|
||
def test_delete_by_source(self, searcher):
|
||
"""delete_by_source 删除源文件的所有 chunks."""
|
||
# 先确认存在
|
||
sources_before = searcher.list_sources()
|
||
if sources_before:
|
||
target = sources_before[0]
|
||
result = searcher.delete_by_source(target)
|
||
assert result is True
|
||
# 删除后不再出现在列表中
|
||
sources_after = searcher.list_sources()
|
||
assert target not in sources_after
|
||
|
||
def test_search_with_source_filter(self, searcher):
|
||
"""带 source_file 过滤的搜索."""
|
||
sources = searcher.list_sources()
|
||
if sources:
|
||
results = searcher.search("测试", top_k=3, source_file=sources[0])
|
||
assert isinstance(results, list)
|
||
# 结果都应来自同一源文件
|
||
for r in results:
|
||
assert r["source_file"] == sources[0]
|
||
|
||
def test_delete_by_source_nonexistent(self, searcher):
|
||
"""删除不存在的源文件返回 False."""
|
||
result = searcher.delete_by_source("nonexistent_file_xyz.md")
|
||
assert result is False
|
||
|
||
def test_list_sources_empty_collection(self, tmp_path):
|
||
"""空 collection 的 list_sources 返回空列表."""
|
||
from src.core.config import EmbedConfig
|
||
from src.core.db import VectorDB
|
||
from src.core.embedder import create_embedder
|
||
from src.core.search import Searcher
|
||
|
||
db = VectorDB(persist_dir=str(tmp_path))
|
||
embedder = create_embedder(EmbedConfig(mode="local"))
|
||
searcher = Searcher(db, embedder, "empty_coll")
|
||
sources = searcher.list_sources()
|
||
assert sources == []
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_search.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add tests/test_search.py
|
||
git commit -m "test: 补充 Searcher list_sources/delete_by_source/get_collection_info 测试"
|
||
```
|
||
|
||
---
|
||
|
||
### Batch 2 收尾
|
||
|
||
```bash
|
||
uv run pytest tests/ -v # 验证全部 90+ 测试仍通过
|
||
```
|
||
|
||
---
|
||
|
||
## Batch 3 — MEDIUM(11 个问题)
|
||
|
||
### Task 15: content 模式默认 file_name 使用 UUID (MED-1)
|
||
|
||
**Files:**
|
||
- Modify: `src/server/app.py:137`
|
||
|
||
- [ ] **Step 1: 改为 UUID 唯一名称**
|
||
|
||
```python
|
||
# 在文件顶部添加 import uuid
|
||
import uuid
|
||
|
||
# 第 137 行改为:
|
||
file_name = req.file_name or f"untitled_{uuid.uuid4().hex[:8]}.md"
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_api.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add src/server/app.py
|
||
git commit -m "fix: content 模式默认 file_name 改用 UUID 防并发覆盖"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 16: X-XSS-Protection 替换为 CSP (MED-2)
|
||
|
||
**Files:**
|
||
- Modify: `src/server/app.py:81`
|
||
|
||
- [ ] **Step 1: 替换过时安全头**
|
||
|
||
```python
|
||
# 第 81 行改为:
|
||
response.headers["Content-Security-Policy"] = "default-src 'self'"
|
||
# 删除: response.headers["X-XSS-Protection"] = "1; mode=block"
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_api.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add src/server/app.py
|
||
git commit -m "fix: 替换过时的 X-XSS-Protection 为 Content-Security-Policy"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 17: EXPECTED_API_KEY 惰性求值化 (MED-3)
|
||
|
||
**Files:**
|
||
- Modify: `src/server/auth.py:14-26`
|
||
|
||
- [ ] **Step 1: 改为惰性求值**
|
||
|
||
```python
|
||
# 第 14 行改为:
|
||
def _get_expected_api_key() -> str:
|
||
"""惰性获取 API Key(每次调用重新从环境变量读取)."""
|
||
return os.environ.get("MD_VECTOR_API_KEY", "")
|
||
|
||
|
||
def verify_api_key(x_api_key: str | None = Header(None)):
|
||
"""验证 API Key. 若未设置环境变量则跳过验证.
|
||
|
||
使用恒定时间比较防止时序攻击.
|
||
"""
|
||
expected = _get_expected_api_key() # 惰性求值
|
||
if expected:
|
||
if x_api_key is None or not hmac.compare_digest(x_api_key, expected):
|
||
logger.warning("API Key 认证失败")
|
||
raise HTTPException(status_code=401, detail="无效的 API Key")
|
||
return True
|
||
```
|
||
|
||
删除第 14 行的 `EXPECTED_API_KEY = os.environ.get("MD_VECTOR_API_KEY", "")`。
|
||
|
||
- [ ] **Step 2: 更新 auth 测试引用**
|
||
|
||
在 `tests/test_auth.py` 中,所有设置 `auth.EXPECTED_API_KEY` 的地方改为 monkeypatch 环境变量:
|
||
|
||
```python
|
||
# 不再需要 auth.EXPECTED_API_KEY = "xxx"
|
||
# 只需 monkeypatch.setenv("MD_VECTOR_API_KEY", "xxx")
|
||
```
|
||
|
||
- [ ] **Step 3: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_auth.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 4: 提交**
|
||
|
||
```bash
|
||
git add src/server/auth.py tests/test_auth.py
|
||
git commit -m "fix: EXPECTED_API_KEY 改为惰性求值防加载顺序问题"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 18: 去重删除逻辑统一到 VectorDB (MED-4)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/db.py` (新增 `delete_by_source` 方法)
|
||
- Modify: `src/core/ingest.py:108-120` (改用 db 方法)
|
||
- Modify: `src/core/search.py:83-97` (改用 db 方法)
|
||
|
||
- [ ] **Step 1: 在 VectorDB 添加 delete_by_source**
|
||
|
||
```python
|
||
# 在 db.py 第 40 行之后添加:
|
||
def delete_by_source(self, collection_name: str, file_name: str) -> bool:
|
||
"""按 source_file 删除文档 (线程安全).
|
||
|
||
在 ingest.py 和 search.py 中均有使用,统一到此方法避免重复代码.
|
||
"""
|
||
import logging
|
||
logger = logging.getLogger("md-vector-db")
|
||
collection = self.get_or_create_collection(collection_name)
|
||
try:
|
||
with self._write_lock:
|
||
existing = collection.get(
|
||
where={"source_file": file_name}
|
||
)
|
||
if existing and existing["ids"]:
|
||
collection.delete(ids=existing["ids"])
|
||
return True
|
||
except ValueError:
|
||
pass # collection 为空时 ChromaDB 抛 ValueError
|
||
except Exception:
|
||
logger.exception("删除文档失败: %s (collection=%s)", file_name, collection_name)
|
||
return False
|
||
```
|
||
|
||
- [ ] **Step 2: 修改 ingest.py 使用新方法**
|
||
|
||
```python
|
||
# 第 108-120 行 _remove_by_source 改为:
|
||
def _remove_by_source(self, file_name: str) -> None:
|
||
"""按 source_file 删除已有 chunks(委托 VectorDB)."""
|
||
self.db.delete_by_source(self.collection_name, file_name)
|
||
```
|
||
|
||
- [ ] **Step 3: 修改 search.py 使用新方法**
|
||
|
||
```python
|
||
# 第 83-97 行 delete_by_source 改为:
|
||
def delete_by_source(self, file_name: str) -> bool:
|
||
"""按文件名删除文档 (委托 VectorDB)."""
|
||
return self.db.delete_by_source(self.collection_name, file_name)
|
||
```
|
||
|
||
- [ ] **Step 4: 运行全量测试**
|
||
|
||
```bash
|
||
uv run pytest tests/ -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 5: 提交**
|
||
|
||
```bash
|
||
git add src/core/db.py src/core/ingest.py src/core/search.py
|
||
git commit -m "refactor: 去重删除逻辑统一到 VectorDB.delete_by_source"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 19: list_collections_with_stats 异常向上传播 (MED-5)
|
||
|
||
**Files:**
|
||
- Modify: `src/server/deps.py:54-59`
|
||
|
||
- [ ] **Step 1: 移除 try/except**
|
||
|
||
```python
|
||
# 第 51-59 行改为:
|
||
def list_collections_with_stats(self) -> list[dict]:
|
||
"""列出所有 collection 及其统计(直接从 ChromaDB 查询)."""
|
||
result = []
|
||
for coll in self.db.list_collections():
|
||
result.append({"name": coll.name, "count": coll.count()})
|
||
return result
|
||
```
|
||
|
||
同时在 `app.py` 的 `list_collections` 端点添加异常处理:
|
||
|
||
```python
|
||
# app.py 第 103-108 行:
|
||
@app.get("/api/v1/collections")
|
||
def list_collections(
|
||
state: AppState = Depends(get_state),
|
||
_: bool = Depends(verify_api_key),
|
||
):
|
||
try:
|
||
return {"collections": state.list_collections_with_stats()}
|
||
except Exception:
|
||
logger.exception("列出集合失败")
|
||
raise HTTPException(status_code=500, detail="服务器内部错误")
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_deps.py tests/test_api.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add src/server/deps.py src/server/app.py
|
||
git commit -m "fix: list_collections 异常向上传播而非静默吞掉"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 20: EPUB 测试设计修复 (MED-6)
|
||
|
||
**Files:**
|
||
- Modify: `tests/test_splitters_epub.py:51-56`
|
||
|
||
- [ ] **Step 1: 使用 monkeypatch 模拟缺失依赖**
|
||
|
||
将 `test_epub_missing_dependency_message` 改为:
|
||
|
||
```python
|
||
def test_epub_missing_dependency_message(self, monkeypatch):
|
||
"""EPUBSplitter.split() 在未安装 ebooklib 时应给出明确提示."""
|
||
from src.core.splitters.epub import EPUBSplitter
|
||
s = EPUBSplitter()
|
||
# 模拟 ebooklib 未安装
|
||
monkeypatch.setitem(
|
||
__import__("sys").modules,
|
||
"ebooklib",
|
||
None,
|
||
raising=False,
|
||
)
|
||
# 由于 importorskip 在模块级别保护,此处实际无法到达
|
||
# 如果到达了(ebooklib 已安装),验证 split 可调用
|
||
assert callable(s.split)
|
||
```
|
||
|
||
由于 `importorskip` 在模块顶部的保护,当 ebooklib 未安装时整个测试文件被 skip。更好的方案:
|
||
|
||
在 `tests/test_splitters_epub.py` 顶部,将:
|
||
```python
|
||
ebooklib = pytest.importorskip("ebooklib", reason="ebooklib 未安装")
|
||
```
|
||
移到 `TestEPUBSplitter` 类内部,仅对有 ebooklib 需求的测试做 skip,缺失依赖测试保留在文件顶部:
|
||
|
||
```python
|
||
"""EPUBSplitter 测试."""
|
||
import pytest
|
||
|
||
|
||
class TestEPUBSplitterMissingDep:
|
||
"""缺失依赖时的行为测试(不 skip)."""
|
||
|
||
def test_split_raises_clear_import_error(self, monkeypatch):
|
||
"""未安装 ebooklib 时给出明确提示."""
|
||
from src.core.splitters.epub import EPUBSplitter
|
||
s = EPUBSplitter()
|
||
# 模拟 import ebooklib 失败
|
||
import builtins
|
||
original_import = builtins.__import__
|
||
|
||
def mock_import(name, *args, **kwargs):
|
||
if name == "ebooklib" or name.startswith("ebooklib."):
|
||
raise ImportError("No module named 'ebooklib'")
|
||
return original_import(name, *args, **kwargs)
|
||
|
||
monkeypatch.setattr(builtins, "__import__", mock_import)
|
||
with pytest.raises(ImportError, match="ebooklib"):
|
||
s.split("dummy.epub", source_file="test.epub")
|
||
|
||
|
||
class TestEPUBSplitter:
|
||
"""EPUBSplitter 测试(需 ebooklib)."""
|
||
ebooklib = pytest.importorskip("ebooklib", reason="ebooklib 未安装")
|
||
# ... 原有测试 ...
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_splitters_epub.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add tests/test_splitters_epub.py
|
||
git commit -m "test: 修复 EPUB 缺失依赖测试设计问题"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 21: test_config.py 用 monkeypatch 替代 reload (MED-7)
|
||
|
||
**Files:**
|
||
- Modify: `tests/test_config.py:89-106`
|
||
|
||
- [ ] **Step 1: 重写环境变量测试**
|
||
|
||
```python
|
||
class TestEmbedConfigEnvVar:
|
||
"""api_key 从环境变量读取."""
|
||
|
||
def test_api_key_from_env(self, monkeypatch):
|
||
"""从环境变量读取 API Key."""
|
||
monkeypatch.setenv("EMBED_API_KEY", "sk-env-test")
|
||
cfg = EmbedConfig(mode="api")
|
||
assert cfg.api_key == "sk-env-test"
|
||
|
||
def test_api_key_empty_when_not_set(self, monkeypatch):
|
||
"""未设置时返回空字符串."""
|
||
monkeypatch.delenv("EMBED_API_KEY", raising=False)
|
||
cfg = EmbedConfig(mode="api")
|
||
assert cfg.api_key == ""
|
||
```
|
||
|
||
Note: `EmbedConfig.api_key` 使用 `default_factory=lambda: os.environ.get("EMBED_API_KEY", "")`,每次实例化都会重新读取环境变量,因此用 `monkeypatch.setenv` 后直接实例化即可,不需要 `reload`。
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_config.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add tests/test_config.py
|
||
git commit -m "test: test_config 用 monkeypatch 替代 importlib.reload"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 22: 速率限制器从环境变量读取配置 (MED-8)
|
||
|
||
**Files:**
|
||
- Modify: `src/server/auth.py:58`
|
||
|
||
- [ ] **Step 1: 改为从环境变量读取**
|
||
|
||
```python
|
||
# 第 58 行改为:
|
||
rate_limiter = RateLimiter(
|
||
max_requests=int(os.environ.get("RATE_LIMIT_MAX", "30")),
|
||
window_seconds=int(os.environ.get("RATE_LIMIT_WINDOW", "60")),
|
||
)
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_auth.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add src/server/auth.py
|
||
git commit -m "fix: 速率限制器配置从环境变量读取而非硬编码"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 23: ingest_obsidian 移除硬编码回退路径 (MED-9)
|
||
|
||
**Files:**
|
||
- Modify: `scripts/ingest_obsidian.py:46-51`
|
||
|
||
- [ ] **Step 1: 替换为明确的错误提示**
|
||
|
||
```python
|
||
# 第 45-51 行改为:
|
||
else:
|
||
log("错误: 未指定目标目录。请通过命令行参数或 OBSIDIAN_DIRS 环境变量提供。")
|
||
log("用法: uv run python scripts/ingest_obsidian.py <目录1> [目录2] ...")
|
||
sys.exit(1)
|
||
```
|
||
|
||
- [ ] **Step 2: 提交**
|
||
|
||
```bash
|
||
git add scripts/ingest_obsidian.py
|
||
git commit -m "fix: ingest_obsidian 移除硬编码回退路径,改为明确错误提示"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 24: MarkdownSplitter 边界测试提取 (MED-10)
|
||
|
||
**Files:**
|
||
- Create: `tests/test_splitters_markdown.py`
|
||
- Modify: `tests/test_ingest.py` (移除 Markdown 测试)
|
||
|
||
- [ ] **Step 1: 创建独立测试文件**
|
||
|
||
```python
|
||
"""MarkdownSplitter 边界测试."""
|
||
import pytest
|
||
|
||
from src.core.splitters import MarkdownSplitter
|
||
|
||
|
||
class TestMarkdownSplitterEdgeCases:
|
||
"""Markdown 分块边界情况."""
|
||
|
||
@pytest.fixture
|
||
def splitter(self):
|
||
return MarkdownSplitter(max_size=1000, overlap=100)
|
||
|
||
def test_no_headings_document(self, splitter):
|
||
"""无标题文档正常分块."""
|
||
md = "这是一段没有标题的纯文本。\n\n第二段内容。"
|
||
chunks = splitter.split(md, source_file="nohead.md")
|
||
assert len(chunks) >= 1
|
||
|
||
def test_deep_headings(self, splitter):
|
||
"""h4-h6 深层标题."""
|
||
md = """# 一级
|
||
## 二级
|
||
### 三级
|
||
#### 四级
|
||
内容在这里。
|
||
##### 五级
|
||
更多内容。
|
||
###### 六级
|
||
最深的内容。"""
|
||
chunks = splitter.split(md, source_file="deep.md")
|
||
assert len(chunks) >= 1
|
||
|
||
def test_hash_in_code_block_not_heading(self, splitter):
|
||
"""代码块中的 # 号不被误识别为标题."""
|
||
md = """# 真实标题
|
||
这是内容。
|
||
```python
|
||
# 这不是标题,是注释
|
||
x = 1 # 行内注释
|
||
## 这也不是标题
|
||
```
|
||
更多内容。"""
|
||
chunks = splitter.split(md, source_file="codehash.md")
|
||
# 代码块内的 # 不应产出新 section
|
||
section_titles = [c.get("section_title", "") for c in chunks]
|
||
# 不应包含 "这不是标题" 之类的代码注释
|
||
for title in section_titles:
|
||
assert "不是标题" not in title
|
||
|
||
def test_adjacent_headings_empty_content(self, splitter):
|
||
"""标题后紧接标题(空内容)."""
|
||
md = """# 标题 A
|
||
# 标题 B
|
||
内容 B。"""
|
||
chunks = splitter.split(md, source_file="adjacent.md")
|
||
assert len(chunks) >= 1
|
||
|
||
def test_only_headings_no_content(self, splitter):
|
||
"""仅有标题无正文."""
|
||
md = "# 只有标题\n## 没有内容"
|
||
chunks = splitter.split(md, source_file="headingsonly.md")
|
||
# 没有正文内容时也可能产生 chunk(标题本身也是内容)
|
||
assert isinstance(chunks, list)
|
||
```
|
||
|
||
- [ ] **Step 2: 从 test_ingest.py 移除 MarkdownSplitter 测试**
|
||
|
||
`test_ingest.py` 中 `TestMarkdownSplitter` 类的测试保留(它们是入仓器使用 MarkdownSplitter 的集成测试)。新增的边界测试放在独立文件中。
|
||
|
||
- [ ] **Step 3: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_splitters_markdown.py tests/test_ingest.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 4: 提交**
|
||
|
||
```bash
|
||
git add tests/test_splitters_markdown.py
|
||
git commit -m "test: 提取 MarkdownSplitter 边界测试到独立文件"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 25: splitters/__init__.py 导出 HTMLSplitter (MED-11)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/splitters/__init__.py:6-20`
|
||
|
||
- [ ] **Step 1: 添加 HTMLSplitter 导入和导出**
|
||
|
||
```python
|
||
# 第 7 行改为(添加 HTMLSplitter import):
|
||
from src.core.splitters.html import HTMLSplitter
|
||
|
||
# __all__ 中加入 "HTMLSplitter":
|
||
__all__ = [
|
||
"Splitter",
|
||
"BaseTextSplitter",
|
||
"MarkdownSplitter",
|
||
"TextSplitter",
|
||
"PDFSplitter",
|
||
"HTMLSplitter", # 新增
|
||
"EPUBSplitter",
|
||
"get_splitter",
|
||
"register_splitter",
|
||
"SUPPORTED_SUFFIXES",
|
||
]
|
||
```
|
||
|
||
- [ ] **Step 2: 验证导入**
|
||
|
||
```bash
|
||
uv run python -c "from src.core.splitters import HTMLSplitter; print('OK')"
|
||
```
|
||
预期:OK
|
||
|
||
- [ ] **Step 3: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_splitters_html.py -v
|
||
```
|
||
预期:4 passed
|
||
|
||
- [ ] **Step 4: 提交**
|
||
|
||
```bash
|
||
git add src/core/splitters/__init__.py
|
||
git commit -m "fix: splitters/__init__.py 导出 HTMLSplitter"
|
||
```
|
||
|
||
---
|
||
|
||
### Batch 3 收尾
|
||
|
||
```bash
|
||
uv run pytest tests/ -v
|
||
```
|
||
|
||
---
|
||
|
||
## Batch 4 — LOW(10 个问题)
|
||
|
||
### Task 26: EPUB 编码回退 (L-1)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/splitters/epub.py:44`
|
||
|
||
- [ ] **Step 1: 添加编码回退**
|
||
|
||
```python
|
||
# 第 42-47 行改为:
|
||
try:
|
||
content = item.get_content().decode("utf-8")
|
||
except UnicodeDecodeError:
|
||
try:
|
||
content = item.get_content().decode("utf-8-sig")
|
||
except UnicodeDecodeError:
|
||
try:
|
||
content = item.get_content().decode("latin-1")
|
||
except UnicodeDecodeError:
|
||
logger.warning("EPUB 跳过一个无法解码的章节: %s", item.get_name())
|
||
continue
|
||
```
|
||
|
||
- [ ] **Step 2: 提交**
|
||
|
||
```bash
|
||
git add src/core/splitters/epub.py
|
||
git commit -m "fix: EPUB 添加编码回退链 utf-8 -> utf-8-sig -> latin-1"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 27: embedder.py requests 惰性导入 (L-2)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/embedder.py:24,151-155`
|
||
|
||
- [ ] **Step 1: 移除模块级 requests 导入,改为惰性导入**
|
||
|
||
```python
|
||
# 第 24 行删除:
|
||
# import requests # noqa: F401 — DashscopeEmbedder 使用 <-- 删除此整行
|
||
|
||
# 第 151 行 DashscopeEmbedder.embed 方法内:
|
||
def embed(self, texts: list[str]) -> list[list[float]]:
|
||
if not texts:
|
||
raise ValueError("文本列表不能为空")
|
||
import requests # 惰性导入(仅DashScope使用)
|
||
resp = requests.post(...)
|
||
```
|
||
|
||
- [ ] **Step 2: 运行测试**
|
||
|
||
```bash
|
||
uv run pytest tests/test_embedder.py -v
|
||
```
|
||
预期:全部通过
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add src/core/embedder.py
|
||
git commit -m "refactor: requests 导入改为惰性(仅 DashScope 使用)"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 28: embedder.py list.sort() 改为 sorted() (L-3)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/embedder.py:176`
|
||
|
||
- [ ] **Step 1: 原地排序替换为不可变版本**
|
||
|
||
```python
|
||
# 第 176 行改为:
|
||
embeddings_raw_sorted = sorted(embeddings_raw, key=lambda x: x.get("text_index", 0))
|
||
return [e["embedding"] for e in embeddings_raw_sorted]
|
||
```
|
||
|
||
- [ ] **Step 2: 提交**
|
||
|
||
```bash
|
||
git add src/core/embedder.py
|
||
git commit -m "refactor: Dashscope 嵌入排序改用 sorted() 避免原地修改 API 响应"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 29: config.yaml 与代码默认值一致性 (L-7, L-8)
|
||
|
||
**Files:**
|
||
- Modify: `config.yaml:19-22`
|
||
|
||
- [ ] **Step 1: 添加 SSL 字段注释**
|
||
|
||
```yaml
|
||
server:
|
||
host: 0.0.0.0
|
||
port: 8000
|
||
# ssl_keyfile: "" # HTTPS 私钥路径(设置后启用 HTTPS)
|
||
# ssl_certfile: "" # HTTPS 证书路径(设置后启用 HTTPS)
|
||
```
|
||
|
||
- [ ] **Step 2: 提交**
|
||
|
||
```bash
|
||
git add config.yaml
|
||
git commit -m "docs: config.yaml 添加 SSL 字段注释说明"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 30: scripts/serve.py 废弃提示改为显式 print (L-9)
|
||
|
||
**Files:**
|
||
- Modify: `scripts/serve.py:1-12`
|
||
|
||
- [ ] **Step 1: 改为 stderr 直接输出**
|
||
|
||
```python
|
||
"""便捷启动脚本 — 已废弃, 请使用 `uv run md-vector-db serve`."""
|
||
import sys
|
||
import uvicorn
|
||
|
||
print(
|
||
"[废弃] scripts/serve.py 已废弃, 请使用 `uv run md-vector-db serve`",
|
||
file=sys.stderr,
|
||
)
|
||
|
||
if __name__ == "__main__":
|
||
uvicorn.run("src.server.app:app", host="127.0.0.1", port=8000, reload=True)
|
||
```
|
||
|
||
- [ ] **Step 2: 提交**
|
||
|
||
```bash
|
||
git add scripts/serve.py
|
||
git commit -m "fix: scripts/serve.py 废弃提示改为 stderr 直接输出"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 31: pyproject.toml 添加开发工具依赖 (L-6)
|
||
|
||
**Files:**
|
||
- Modify: `pyproject.toml:21`
|
||
|
||
- [ ] **Step 1: 扩展 dev 依赖**
|
||
|
||
```toml
|
||
dev = ["pytest>=8.0", "httpx>=0.27.0", "pytest-cov>=5.0", "ruff>=0.8.0", "mypy>=1.13"]
|
||
```
|
||
|
||
并添加 tool 配置节:
|
||
|
||
```toml
|
||
[tool.ruff]
|
||
line-length = 100
|
||
target-version = "py313"
|
||
|
||
[tool.ruff.lint]
|
||
select = ["E", "F", "I", "N", "W"]
|
||
|
||
[tool.mypy]
|
||
python_version = "3.13"
|
||
ignore_missing_imports = true
|
||
```
|
||
|
||
- [ ] **Step 2: 安装新依赖**
|
||
|
||
```bash
|
||
uv sync --extra dev
|
||
```
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add pyproject.toml uv.lock
|
||
git commit -m "build: dev 依赖添加 pytest-cov/ruff/mypy"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 32: 文档更新 — 测试数量 + 数据流 (L-10)
|
||
|
||
**Files:**
|
||
- Modify: `CLAUDE.md`
|
||
- Modify: `README.md`
|
||
|
||
- [ ] **Step 1: CLAUDE.md 更新**
|
||
|
||
- `(46 个)` → `(90+ 个)`
|
||
- 数据流更新为包含多格式:`文件 → get_splitter(path) 自动选择 → Splitter.split() → batch_embed() → ChromaDB`
|
||
|
||
- [ ] **Step 2: README.md 同步更新**
|
||
|
||
同上修改。
|
||
|
||
- [ ] **Step 3: 提交**
|
||
|
||
```bash
|
||
git add CLAUDE.md README.md
|
||
git commit -m "docs: 更新测试数量和文档以反映多格式支持"
|
||
```
|
||
|
||
---
|
||
|
||
### Task 33: DEFAULT_CONFIG_PATH 路径处理 (L-4)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/config.py:16`
|
||
|
||
- [ ] **Step 1: 不改 DEFAULT_CONFIG_PATH,而是在 load_config 中已正确处理**
|
||
|
||
`load_config` 已经处理了相对路径回退到项目根目录的逻辑。`DEFAULT_CONFIG_PATH` 本身保持不变即可,因为使用方(CLI 和 deps.py)都是通过 `load_config()` 调用而非直接使用 `DEFAULT_CONFIG_PATH` 作为文件路径字面量。
|
||
|
||
此问题标记为无需修改(当前逻辑已正确处理)。
|
||
|
||
- [ ] **Step 2: 提交(跳过或记录为 won't fix)**
|
||
|
||
如果决定不修改,跳过此任务。当前 `load_config()` 已经做了项目根目录回退。
|
||
|
||
---
|
||
|
||
### Task 34: registry.py 线程安全文档说明 (L-5)
|
||
|
||
**Files:**
|
||
- Modify: `src/core/splitters/registry.py:23`
|
||
|
||
- [ ] **Step 1: 添加文档注释**
|
||
|
||
```python
|
||
# 第 23 行函数上方添加:
|
||
def register_splitter(ext: str, splitter_cls: type[Splitter]) -> None:
|
||
"""注册自定义 Splitter 类.
|
||
|
||
注意: 此函数非线程安全,请在程序启动时调用(单线程阶段)。
|
||
运行时动态注册需自行加锁。
|
||
"""
|
||
```
|
||
|
||
- [ ] **Step 2: 提交**
|
||
|
||
```bash
|
||
git add src/core/splitters/registry.py
|
||
git commit -m "docs: register_splitter 添加线程安全注意事项"
|
||
```
|
||
|
||
---
|
||
|
||
### Batch 4 收尾
|
||
|
||
```bash
|
||
uv run pytest tests/ -v # 全部测试
|
||
```
|
||
|
||
---
|
||
|
||
## 最终验证清单
|
||
|
||
- [ ] `uv run pytest tests/ -v` — 全部测试通过
|
||
- [ ] `uv run md-vector-db ingest --help` — CLI 可用
|
||
- [ ] `uv run md-vector-db search "测试" -k 1` — 检索可用
|
||
- [ ] `uv run md-vector-db stats` — 统计可用
|
||
- [ ] 检查 `config.yaml` 中的 chunk 修改是否生效
|
||
- [ ] 确认无硬编码密钥、无敏感信息泄漏
|