diff --git a/config.yaml b/config.yaml index ac8a627..ec1ada4 100644 --- a/config.yaml +++ b/config.yaml @@ -6,7 +6,7 @@ embed: mode: local # local | api local_model: BAAI/bge-small-zh-v1.5 api_base: "" # api 模式下填写 - api_key: "" # api 模式下填写 + # api_key 请通过环境变量 EMBED_API_KEY 设置 chunk: max_size: 1000 # 分块最大字符数 diff --git a/docs/superpowers/plans/2026-07-05-audit-fixes.md b/docs/superpowers/plans/2026-07-05-audit-fixes.md new file mode 100644 index 0000000..de226b6 --- /dev/null +++ b/docs/superpowers/plans/2026-07-05-audit-fixes.md @@ -0,0 +1,889 @@ +# 审计问题修复计划 + +> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) to implement this plan task-by-task. + +**Goal:** 修复架构审计报告中 Phase 1+2 的全部严重和高风险问题,以及部分中风险问题。 + +**Architecture:** 按影响面从底层到上层依次修复:config → embedder → db → ingest → search → server → cli。每层修复后运行测试验证不引入回归。 + +**Tech Stack:** Python 3.13, FastAPI, ChromaDB, sentence-transformers + +--- + +## 修复范围 + +| 审计编号 | 问题 | 本次修复 | +|----------|------|----------| +| S-01 | API 密钥明文 | ✅ | +| S-02 | 路径遍历 | ✅ | +| S-03 | 无认证 | ✅ | +| S-05 | 错误信息泄露 | ✅ | +| C-01 | 并发不安全 | ✅ | +| C-02 | 全局环境变量污染 | ✅ | +| C-04 | top_k 无上界 | ✅ | +| C-05 | 大文件 OOM | ✅ | +| C-06 | 配置路径硬编码 | ✅ | +| C-07 | 假健康检查 | ✅ | +| D-02 | 全局单例 | ✅ | +| D-03 | 同步嵌入阻塞 | ✅ | +| M-02 | 无日志 | ✅ | +| T-02 | DELETE 端点无测试 | ✅ | + +--- + +### Task 1: 配置安全修复 (S-01, C-06) + +**Files:** +- Modify: `config.yaml` — 移除 `api_key` 字段 +- Modify: `src/core/config.py` — `api_key` 默认从环境变量读取 +- Modify: `src/cli/main.py:14` — 移除重复 sys.path 操作(顺带) +- Add Test: `tests/test_config.py` — 补充 API key 环境变量测试 + +- [ ] **Step 1: 补充测试 — tests/test_config.py 追加** + +```python +class TestEmbedConfigEnvVar: + """api_key 从环境变量读取.""" + + def test_api_key_from_env(self, monkeypatch): + monkeypatch.setenv("EMBED_API_KEY", "sk-env-test") + from src.core.config import EmbedConfig + 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) + from src.core.config import EmbedConfig + cfg = EmbedConfig(mode="api") + assert cfg.api_key == "" +``` + +- [ ] **Step 2: 运行测试确认失败** + +```bash +cd D:/Code/doing_exercises/programs/md-vector-db && uv run pytest tests/test_config.py -v +``` +Expected: 新测试 FAIL + +- [ ] **Step 3: 修改 config.yaml — 删除 api_key 行** + +```yaml +embed: + mode: local # local | api + local_model: BAAI/bge-small-zh-v1.5 + api_base: "" # api 模式下填写 + # api_key 请通过环境变量 EMBED_API_KEY 设置 +``` + +- [ ] **Step 4: 修改 config.py — api_key 默认从环境变量读取** + +将 `EmbedConfig` 的 `api_key` 字段改为: + +```python +@dataclass +class EmbedConfig: + """嵌入模型配置.""" + mode: str = "local" + local_model: str = "BAAI/bge-small-zh-v1.5" + api_base: str = "" + api_key: str = field(default_factory=lambda: os.environ.get("EMBED_API_KEY", "")) +``` + +并在文件顶部加 `import os`。 + +- [ ] **Step 5: 统一配置路径常量** + +在 `config.py` 末尾添加: +```python +DEFAULT_CONFIG_PATH = "config.yaml" +``` + +其他文件中的 `"config.yaml"` 字符串引用改为 `from src.core.config import DEFAULT_CONFIG_PATH`。 + +修改 `src/server/app.py:47,59` 和 `src/cli/main.py:26` 中的硬编码 `"config.yaml"`。 + +- [ ] **Step 6: 运行全部测试** + +```bash +cd D:/Code/doing_exercises/programs/md-vector-db && uv run pytest tests/ -v --tb=short +``` +Expected: 33 passed (31 + 2 new) + +- [ ] **Step 7: Commit** + +```bash +git add -A && git commit -m "fix: remove api_key from config.yaml, read from env var; unify config path" +``` + +--- + +### Task 2: 路径遍历防护 (S-02) + +**Files:** +- Modify: `src/server/app.py` — `ingest_document()` 增加路径白名单 +- Add Test: `tests/test_api.py` — 路径遍历攻击测试 + +- [ ] **Step 1: 补充测试 — tests/test_api.py 追加** + +```python +class TestSecurity: + """安全测试.""" + + def test_ingest_rejects_path_traversal(self, client): + """拒绝路径遍历攻击.""" + response = client.post( + "/api/v1/ingest", + json={"file_path": "../../../etc/passwd", "file_name": "hack.md"}, + ) + assert response.status_code in (400, 403) + + def test_ingest_rejects_absolute_path(self, client): + """拒绝绝对路径.""" + response = client.post( + "/api/v1/ingest", + json={"file_path": "C:\\Windows\\System32\\config\\SAM"}, + ) + assert response.status_code in (400, 403) +``` + +- [ ] **Step 2: 运行测试确认失败** + +```bash +uv run pytest tests/test_api.py::TestSecurity -v +``` +Expected: FAIL (当前无防护) + +- [ ] **Step 3: 修改 app.py — 路径校验函数** + +在 `ingest_document()` 最前面加: + +```python +import os as _os + +def _is_safe_path(path_str: str) -> bool: + """检查路径是否安全: 仅允许相对路径且不含 .. 穿越.""" + normalized = _os.path.normpath(path_str) + if _os.path.isabs(normalized): + return False + if ".." in normalized.split(_os.sep): + return False + return True + +@app.post("/api/v1/ingest") +def ingest_document(req: IngestRequest): + if req.file_path: + if not _is_safe_path(req.file_path): + raise HTTPException(status_code=400, detail="不允许的路径") + path = Path(req.file_path) + ... +``` + +- [ ] **Step 4: 运行测试确认通过** + +```bash +uv run pytest tests/test_api.py -v +``` +Expected: 8 passed (6 original + 2 new security tests) + +- [ ] **Step 5: Commit** + +```bash +git add -A && git commit -m "fix: add path traversal protection to ingest endpoint" +``` + +--- + +### Task 3: API 认证 + 速率限制 + 安全错误消息 (S-03, S-04, S-05) + +**Files:** +- Modify: `src/server/app.py` — 添加 API Key 中间件、速率限制、安全错误消息 +- Create: `src/server/auth.py` — 认证依赖 +- Add Test: `tests/test_api.py` — 认证/速率测试 + +- [ ] **Step 1: 创建 auth.py** + +```python +"""API 认证与安全中间件.""" +import os +from fastapi import Header, HTTPException, Request +from fastapi.responses import JSONResponse +import time +from collections import defaultdict +import threading + + +# -- API Key 认证 -- +EXPECTED_API_KEY = os.environ.get("MD_VECTOR_API_KEY", "") + + +def verify_api_key(x_api_key: str | None = Header(None)): + """验证 API Key. 若未设置环境变量则跳过验证.""" + if EXPECTED_API_KEY and x_api_key != EXPECTED_API_KEY: + raise HTTPException(status_code=401, detail="无效的 API Key") + return True + + +# -- 简易速率限制 -- +class RateLimiter: + """基于内存的简易速率限制器.""" + + def __init__(self, max_requests: int = 30, window_seconds: int = 60): + self.max_requests = max_requests + self.window = window_seconds + self._store: dict[str, list[float]] = defaultdict(list) + self._lock = threading.Lock() + + def is_allowed(self, client_id: str) -> bool: + now = time.time() + with self._lock: + records = self._store[client_id] + # 清理过期记录 + records[:] = [t for t in records if now - t < self.window] + if len(records) >= self.max_requests: + return False + records.append(now) + return True + + async def __call__(self, request: Request): + client_id = request.client.host if request.client else "unknown" + if not self.is_allowed(client_id): + raise HTTPException(status_code=429, detail="请求过于频繁,请稍后再试") + return True + + +rate_limiter = RateLimiter(max_requests=30, window_seconds=60) +``` + +- [ ] **Step 2: 修改 app.py — 集成认证和速率限制** + +在 `app.py` 中: + +```python +from src.server.auth import verify_api_key, rate_limiter +import logging + +logger = logging.getLogger("md-vector-db") + +# 速率限制中间件 +@app.middleware("http") +async def rate_limit_middleware(request: Request, call_next): + await rate_limiter(request) + response = await call_next(request) + return response + + +@app.post("/api/v1/ingest") +def ingest_document(req: IngestRequest, _: bool = Depends(verify_api_key)): + ... + except HTTPException: + raise + except Exception as e: + logger.exception("入库失败: %s", e) + raise HTTPException(status_code=500, detail="服务器内部错误") + + +@app.post("/api/v1/search") +def search_documents(req: SearchRequest, _: bool = Depends(verify_api_key)): + ... + + +@app.delete("/api/v1/documents/{file_name}") +def delete_document(file_name: str, _: bool = Depends(verify_api_key)): + ... +``` + +- [ ] **Step 3: 运行测试确认通过** + +```bash +uv run pytest tests/test_api.py -v +``` +Expected: 全部通过 (未设 API Key 环境变量时认证跳过) + +- [ ] **Step 4: Commit** + +```bash +git add -A && git commit -m "fix: add API key auth, rate limiting, and safe error messages" +``` + +--- + +### Task 4: 请求校验 (C-04, M-03) + +**Files:** +- Modify: `src/server/app.py` — `SearchRequest.top_k` 加 Pydantic Field 校验;`IngestRequest.file_name` 校验 + +- [ ] **Step 1: 修改请求模型** + +```python +from pydantic import BaseModel, Field + +class IngestRequest(BaseModel): + file_path: str | None = None + content: str | None = None + file_name: str | None = Field(default=None, max_length=255, pattern=r"^[^\\/:*?\"<>|]+\.md$") + + +class SearchRequest(BaseModel): + query: str = Field(..., min_length=1, max_length=2000) + top_k: int = Field(default=10, ge=1, le=100) +``` + +- [ ] **Step 2: 运行测试** + +```bash +uv run pytest tests/test_api.py -v +``` +Expected: 全部通过 + +- [ ] **Step 3: Commit** + +```bash +git add -A && git commit -m "fix: add field validation to API request models" +``` + +--- + +### Task 5: ChromaDB 线程安全 (C-01) + +**Files:** +- Modify: `src/core/db.py` — VectorDB 写操作加锁 + +- [ ] **Step 1: 修改 db.py** + +```python +"""ChromaDB 数据库层.""" +import threading +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) + self._write_lock = threading.Lock() + + 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 (线程安全).""" + with self._write_lock: + try: + self.client.delete_collection(name=name) + except ValueError: + pass + + def close(self) -> None: + """释放数据库连接.""" + self.client.close() + + @property + def write_lock(self) -> threading.Lock: + """获取写锁,供外部在 add/delete/update 操作时使用.""" + return self._write_lock +``` + +- [ ] **Step 2: 修改 ingest.py — 使用写锁** + +```python +class DocumentIngestor: + ... + def ingest_content(self, content: str, file_name: str) -> int: + ... + with self.db.write_lock: + self.collection.add( + ids=ids, embeddings=embeddings, + documents=texts, metadatas=metadatas, + ) + return len(chunks) +``` + +- [ ] **Step 3: 修改 search.py — delete 使用写锁** + +```python +def delete_by_source(self, file_name: str) -> bool: + ... + with self.db.write_lock: + existing = self.collection.get(where={"source_file": file_name}) + if existing and existing["ids"]: + self.collection.delete(ids=existing["ids"]) + return True + return False +``` + +- [ ] **Step 4: 运行测试** + +```bash +uv run pytest tests/ -v --tb=short +``` +Expected: 31 passed + +- [ ] **Step 5: Commit** + +```bash +git add -A && git commit -m "fix: add thread-safe write lock to ChromaDB operations" +``` + +--- + +### Task 6: Embedder 重构 (C-02, C-05, D-01) + +**Files:** +- Modify: `src/core/embedder.py` — 拆分为策略模式;移除全局 env var 副作用;加分批嵌入 +- Modify: `tests/test_embedder.py` — 补充 API 模式测试 + +- [ ] **Step 1: 重写 embedder.py** + +```python +"""嵌入模型抽象层 — 策略模式.""" +import os +import logging +from dataclasses import dataclass +from typing import Protocol + +from src.core.config import EmbedConfig + +logger = logging.getLogger(__name__) + +# 模型下载镜像 +HF_MIRROR = "https://hf-mirror.com" + + +class Embedder(Protocol): + """嵌入器接口.""" + @property + def dimension(self) -> int: ... + def embed(self, texts: list[str]) -> list[list[float]]: ... + + +class LocalEmbedder: + """本地 sentence-transformers 模型嵌入器.""" + + def __init__(self, config: EmbedConfig): + from sentence_transformers import SentenceTransformer + self._config = config + try: + self._model = SentenceTransformer( + config.local_model, local_files_only=True + ) + except Exception: + logger.info("模型未缓存, 通过镜像下载 %s", config.local_model) + os.environ["HF_ENDPOINT"] = HF_MIRROR + self._model = SentenceTransformer(config.local_model) + os.environ.pop("HF_ENDPOINT", None) + + @property + def dimension(self) -> int: + return self._model.get_embedding_dimension() + + def embed(self, texts: list[str]) -> list[list[float]]: + if not texts: + raise ValueError("文本列表不能为空") + embeddings = self._model.encode(texts, normalize_embeddings=True) + return embeddings.tolist() + + +class APIEmbedder: + """OpenAI 兼容 API 嵌入器.""" + + def __init__(self, config: EmbedConfig): + self._api_base = config.api_base + self._api_key = config.api_key + + @property + def dimension(self) -> int: + return 1536 + + def embed(self, texts: list[str]) -> list[list[float]]: + if not texts: + raise ValueError("文本列表不能为空") + 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: + """工厂函数.""" + if config.mode == "local": + return LocalEmbedder(config) + elif config.mode == "api": + return APIEmbedder(config) + else: + raise ValueError(f"不支持的嵌入模式: {config.mode}") + + +def batch_embed(embedder: Embedder, texts: list[str], batch_size: int = 32) -> list[list[float]]: + """分批嵌入,避免一次性传入过多文本导致 OOM.""" + all_embeddings = [] + for i in range(0, len(texts), batch_size): + batch = texts[i:i + batch_size] + all_embeddings.extend(embedder.embed(batch)) + return all_embeddings +``` + +- [ ] **Step 2: 修改 ingest.py — 使用 batch_embed** + +将 `ingest_content()` 中的: +```python +embeddings = self.embedder.embed(texts) +``` +改为: +```python +from src.core.embedder import batch_embed +embeddings = batch_embed(self.embedder, texts) +``` + +- [ ] **Step 3: 补充 API 模式测试** + +在 `tests/test_embedder.py` 追加: +```python +class TestAPIEmbedder: + """API 嵌入器测试.""" + + def test_api_embedder_init(self): + from src.core.embedder import APIEmbedder + cfg = EmbedConfig(mode="api", api_base="https://api.test.com", api_key="sk-test") + emb = APIEmbedder(cfg) + assert emb.dimension == 1536 + + def test_api_embedder_empty_raises(self): + from src.core.embedder import APIEmbedder + cfg = EmbedConfig(mode="api") + emb = APIEmbedder(cfg) + with pytest.raises(ValueError): + emb.embed([]) + +class TestBatchEmbed: + """分批嵌入测试.""" + + def test_batch_embed(self): + from src.core.embedder import batch_embed, create_embedder + embedder = create_embedder(EmbedConfig(mode="local")) + texts = ["测试文本"] * 70 # 超过 batch_size 32 + results = batch_embed(embedder, texts, batch_size=32) + assert len(results) == 70 + assert len(results[0]) == embedder.dimension +``` + +同时把原来的 `TestEmbedder` 中的 `create_embedder(EmbedConfig(mode="local"))` 改为检查返回类型是 `LocalEmbedder`。 + +- [ ] **Step 4: 运行测试** + +```bash +uv run pytest tests/test_embedder.py tests/test_ingest.py tests/test_search.py tests/test_api.py -v +``` +Expected: 全部通过 + +- [ ] **Step 5: Commit** + +```bash +git add -A && git commit -m "refactor: strategy-pattern embedder, remove global env vars, add batch embedding" +``` + +--- + +### Task 7: 服务层重构 (D-02, D-03, C-07) + +**Files:** +- Modify: `src/server/app.py` — FastAPI Depends 注入替代全局单例;异步嵌入;真实健康检查 +- Create: `src/server/deps.py` — 依赖注入工厂 + +- [ ] **Step 1: 创建 deps.py** + +```python +"""FastAPI 依赖注入.""" +import os +import functools +import logging + +from src.core.config import load_config, 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 + +logger = logging.getLogger(__name__) + + +class AppState: + """应用级共享状态 (替代模块级全局变量).""" + + def __init__(self): + config_path = os.environ.get("MD_VECTOR_CONFIG", "config.yaml") + self.config = load_config(config_path) + + data_dir = os.environ.get("MD_VECTOR_DB_DATA_DIR", self.config.chroma.persist_dir) + self.db = VectorDB(persist_dir=data_dir) + + self.embedder = create_embedder(self.config.embed) + + collection = os.environ.get("MD_VECTOR_DB_COLLECTION", self.config.chroma.collection_name) + self.searcher = Searcher(self.db, self.embedder, collection) + self.ingestor = DocumentIngestor(self.db, self.embedder, collection) + + def is_healthy(self) -> dict: + """真实健康检查.""" + status = {"status": "ok", "checks": {}} + try: + count = self.searcher.collection.count() + status["checks"]["chromadb"] = {"status": "ok", "count": count} + except Exception as e: + status["checks"]["chromadb"] = {"status": "error", "detail": str(e)} + status["status"] = "degraded" + try: + _ = self.embedder.dimension + status["checks"]["embedder"] = {"status": "ok"} + except Exception as e: + status["checks"]["embedder"] = {"status": "error", "detail": str(e)} + status["status"] = "degraded" + return status + + +_state: AppState | None = None + + +def get_state() -> AppState: + """获取应用状态单例.""" + global _state + if _state is None: + logger.info("初始化应用状态...") + _state = AppState() + return _state +``` + +- [ ] **Step 2: 重写 app.py — 使用 Depends 注入** + +```python +"""FastAPI 服务层.""" +import logging +from pathlib import Path + +from fastapi import FastAPI, HTTPException, Depends, Request +from fastapi.responses import RedirectResponse, JSONResponse +from pydantic import BaseModel, Field + +from src.server.auth import verify_api_key, rate_limiter +from src.server.deps import get_state, AppState + +logger = logging.getLogger(__name__) + +# -- 请求模型 -- +class IngestRequest(BaseModel): + file_path: str | None = None + content: str | None = None + file_name: str | None = Field(default=None, max_length=255, pattern=r"^[^\\/:*?\"<>|]+\.md$") + +class SearchRequest(BaseModel): + query: str = Field(..., min_length=1, max_length=2000) + top_k: int = Field(default=10, ge=1, le=100) + +# -- App -- +app = FastAPI(title="md-vector-db", version="0.1.0") + + +@app.middleware("http") +async def rate_limit_middleware(request: Request, call_next): + await rate_limiter(request) + return await call_next(request) + + +@app.get("/") +def root(): + return RedirectResponse(url="/docs") + + +@app.get("/api/v1/health") +def health(state: AppState = Depends(get_state)): + return state.is_healthy() + + +@app.get("/api/v1/collections") +def list_collections(state: AppState = Depends(get_state)): + info = state.searcher.get_collection_info() + sources = state.searcher.list_sources() + return {"collections": [info], "sources": sources} + + +@app.post("/api/v1/ingest") +def ingest_document(req: IngestRequest, state: AppState = Depends(get_state), + _: bool = Depends(verify_api_key)): + try: + if req.file_path: + path = Path(req.file_path).resolve() + if not str(path).startswith(str(Path.cwd().resolve())): + raise HTTPException(status_code=400, detail="不允许访问该路径") + if not path.exists(): + raise HTTPException(status_code=404, detail=f"文件不存在: {path.name}") + count = state.ingestor.ingest_file(str(path)) + file_name = path.name + elif req.content: + file_name = req.file_name or "untitled.md" + count = state.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: + logger.exception("入库失败") + raise HTTPException(status_code=500, detail="服务器内部错误") + + +@app.post("/api/v1/search") +def search_documents(req: SearchRequest, state: AppState = Depends(get_state), + _: bool = Depends(verify_api_key)): + results = state.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, state: AppState = Depends(get_state), + _: bool = Depends(verify_api_key)): + deleted = state.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 3: 运行全部测试** + +```bash +uv run pytest tests/ -v --tb=short +``` +Expected: 全部通过 + +- [ ] **Step 4: Commit** + +```bash +git add -A && git commit -m "refactor: FastAPI Depends injection, async-safe, real health check" +``` + +--- + +### Task 8: 测试补充与回归 (T-02, D-05, T-03) + +**Files:** +- Modify: `tests/test_api.py` — 补充 DELETE 端点测试、ingest_file 测试 +- Modify: `tests/test_ingest.py` — 补充 DocumentIngestor 集成测试 + +- [ ] **Step 1: 补充 API 测试** + +在 `tests/test_api.py` 追加: +```python +class TestDeleteEndpoint: + """删除端点.""" + + def test_delete_nonexistent(self, client): + response = client.delete("/api/v1/documents/nonexistent.md") + assert response.status_code == 404 + + def test_delete_ingested(self, client): + client.post("/api/v1/ingest", json={"content": "# Test", "file_name": "del.md"}) + response = client.delete("/api/v1/documents/del.md") + assert response.status_code == 200 + # 删除后搜索不应返回结果 + search_resp = client.post("/api/v1/search", json={"query": "Test", "top_k": 3}) + results = search_resp.json()["results"] + sources = [r["source_file"] for r in results] + assert "del.md" not in sources +``` + +- [ ] **Step 2: 补充 ingest 集成测试** + +在 `tests/test_ingest.py` 追加: +```python +class TestIngestorIntegration: + """入库器集成测试 (使用真实 embedder).""" + + def test_ingest_content_real(self, tmp_path): + import tempfile + 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)) + embedder = create_embedder(EmbedConfig(mode="local")) + ingestor = DocumentIngestor(db, embedder, "test_integration") + + count = ingestor.ingest_content("# Hello\nWorld.", "hello.md") + assert count > 0 + assert ingestor.collection.count() == count + + def test_ingest_deduplicates(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)) + embedder = create_embedder(EmbedConfig(mode="local")) + ingestor = DocumentIngestor(db, embedder, "test_dedup") + + c1 = ingestor.ingest_content("# A", "dup.md") + c2 = ingestor.ingest_content("# B", "dup.md") + # 第二次入库应覆盖第一次,总数不翻倍 + assert ingestor.collection.count() == c2 +``` + +- [ ] **Step 3: 运行全部测试** + +```bash +uv run pytest tests/ -v --tb=short +``` +Expected: 35+ passed + +- [ ] **Step 4: 最终 Commit** + +```bash +git add -A && git commit -m "test: add DELETE endpoint tests and ingest integration tests" +``` + +--- + +### Task 9: 日志系统 + README 更新 + +**Files:** +- Modify: `src/server/app.py` — 启动时配置 logging +- Modify: `src/cli/main.py` — 添加 logging +- Modify: `README.md` — 更新 API 文档 + +- [ ] **Step 1: 在 app.py 顶部配置日志** + +```python +import logging +logging.basicConfig( + level=logging.INFO, + format="%(asctime)s [%(levelname)s] %(name)s: %(message)s", + handlers=[logging.StreamHandler()] +) +``` + +- [ ] **Step 2: 运行测试确认无回归** + +```bash +uv run pytest tests/ -v --tb=short +``` + +- [ ] **Step 3: Commit** + +```bash +git add -A && git commit -m "feat: add structured logging and update README" +``` + +--- + +## 最终验证 + +```bash +cd D:/Code/doing_exercises/programs/md-vector-db +uv run pytest tests/ -v +uv run python -m src.cli.main search "测试" +``` diff --git a/src/cli/main.py b/src/cli/main.py index b161211..a1a3e60 100644 --- a/src/cli/main.py +++ b/src/cli/main.py @@ -13,7 +13,7 @@ if sys.stdout.encoding != "utf-8": # 确保 src 在路径中 sys.path.insert(0, str(Path(__file__).parent.parent)) -from src.core.config import load_config +from src.core.config import load_config, DEFAULT_CONFIG_PATH from src.core.db import VectorDB from src.core.embedder import create_embedder from src.core.ingest import DocumentIngestor @@ -23,7 +23,7 @@ from src.core.search import Searcher app = typer.Typer(name="md-vector-db", help="Markdown 文档向量数据库管理工具") -def _get_components(config_path: str = "config.yaml"): +def _get_components(config_path: str = DEFAULT_CONFIG_PATH): """初始化所有组件.""" cfg = load_config(config_path) db = VectorDB(persist_dir=cfg.chroma.persist_dir) diff --git a/src/core/config.py b/src/core/config.py index de0f201..6f46c4d 100644 --- a/src/core/config.py +++ b/src/core/config.py @@ -1,11 +1,15 @@ """应用配置加载模块.""" from __future__ import annotations +import os from dataclasses import dataclass, field from pathlib import Path import yaml +# 默认配置文件路径 +DEFAULT_CONFIG_PATH = "config.yaml" + @dataclass class ChromaConfig: @@ -22,7 +26,9 @@ class EmbedConfig: mode: str = "local" # "local" | "api" local_model: str = "BAAI/bge-small-zh-v1.5" api_base: str = "" - api_key: str = "" + api_key: str = field( + default_factory=lambda: os.environ.get("EMBED_API_KEY", "") + ) @dataclass @@ -59,7 +65,7 @@ class AppConfig: def load_config(path: str | None = None) -> AppConfig: """从 YAML 文件加载配置, 若文件不存在则返回默认配置.""" - config_path = path or "config.yaml" + config_path = path or DEFAULT_CONFIG_PATH if not Path(config_path).exists(): return AppConfig() diff --git a/src/server/app.py b/src/server/app.py index 09ecffc..8fc075b 100644 --- a/src/server/app.py +++ b/src/server/app.py @@ -5,7 +5,7 @@ from pathlib import Path from fastapi import FastAPI, HTTPException from pydantic import BaseModel -from src.core.config import load_config, EmbedConfig +from src.core.config import load_config, EmbedConfig, DEFAULT_CONFIG_PATH from src.core.db import VectorDB from src.core.embedder import create_embedder from src.core.ingest import DocumentIngestor @@ -42,9 +42,9 @@ def _get_db(): def _get_embedder(): global _embedder if _embedder is None: - # 先尝试从 config.yaml 加载, 失败则用默认 local + # 先尝试从配置文件加载, 失败则用默认 local try: - cfg = load_config("config.yaml") + cfg = load_config(DEFAULT_CONFIG_PATH) embed_cfg = cfg.embed except Exception: embed_cfg = EmbedConfig(mode="local") @@ -56,7 +56,7 @@ def _init_services(): global _searcher, _ingestor collection = os.getenv("MD_VECTOR_DB_COLLECTION", "markdown_docs") try: - cfg = load_config("config.yaml") + cfg = load_config(DEFAULT_CONFIG_PATH) collection = cfg.chroma.collection_name except Exception: pass diff --git a/tests/test_config.py b/tests/test_config.py index 8739ac9..e1c30ca 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -1,4 +1,5 @@ """配置加载模块测试.""" +import os import tempfile from pathlib import Path @@ -80,3 +81,26 @@ server: cfg = load_config("/nonexistent/config.yaml") assert cfg.chroma.persist_dir == "./data" assert cfg.embed.mode == "local" + + +class TestEmbedConfigEnvVar: + """api_key 从环境变量读取.""" + + def test_api_key_from_env(self, monkeypatch): + monkeypatch.setenv("EMBED_API_KEY", "sk-env-test") + import importlib + from src.core import config + importlib.reload(config) + cfg = config.EmbedConfig(mode="api") + assert cfg.api_key == "sk-env-test" + # 清理 + importlib.reload(config) + + def test_api_key_empty_when_not_set(self, monkeypatch): + monkeypatch.delenv("EMBED_API_KEY", raising=False) + import importlib + from src.core import config + importlib.reload(config) + cfg = config.EmbedConfig(mode="api") + assert cfg.api_key == "" + importlib.reload(config)