fix: remove api_key from config.yaml, read from EMBED_API_KEY env var; unify config path constant
This commit is contained in:
+1
-1
@@ -6,7 +6,7 @@ embed:
|
|||||||
mode: local # local | api
|
mode: local # local | api
|
||||||
local_model: BAAI/bge-small-zh-v1.5
|
local_model: BAAI/bge-small-zh-v1.5
|
||||||
api_base: "" # api 模式下填写
|
api_base: "" # api 模式下填写
|
||||||
api_key: "" # api 模式下填写
|
# api_key 请通过环境变量 EMBED_API_KEY 设置
|
||||||
|
|
||||||
chunk:
|
chunk:
|
||||||
max_size: 1000 # 分块最大字符数
|
max_size: 1000 # 分块最大字符数
|
||||||
|
|||||||
@@ -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 "测试"
|
||||||
|
```
|
||||||
+2
-2
@@ -13,7 +13,7 @@ if sys.stdout.encoding != "utf-8":
|
|||||||
# 确保 src 在路径中
|
# 确保 src 在路径中
|
||||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
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.db import VectorDB
|
||||||
from src.core.embedder import create_embedder
|
from src.core.embedder import create_embedder
|
||||||
from src.core.ingest import DocumentIngestor
|
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 文档向量数据库管理工具")
|
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)
|
cfg = load_config(config_path)
|
||||||
db = VectorDB(persist_dir=cfg.chroma.persist_dir)
|
db = VectorDB(persist_dir=cfg.chroma.persist_dir)
|
||||||
|
|||||||
+8
-2
@@ -1,11 +1,15 @@
|
|||||||
"""应用配置加载模块."""
|
"""应用配置加载模块."""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import yaml
|
import yaml
|
||||||
|
|
||||||
|
# 默认配置文件路径
|
||||||
|
DEFAULT_CONFIG_PATH = "config.yaml"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ChromaConfig:
|
class ChromaConfig:
|
||||||
@@ -22,7 +26,9 @@ class EmbedConfig:
|
|||||||
mode: str = "local" # "local" | "api"
|
mode: str = "local" # "local" | "api"
|
||||||
local_model: str = "BAAI/bge-small-zh-v1.5"
|
local_model: str = "BAAI/bge-small-zh-v1.5"
|
||||||
api_base: str = ""
|
api_base: str = ""
|
||||||
api_key: str = ""
|
api_key: str = field(
|
||||||
|
default_factory=lambda: os.environ.get("EMBED_API_KEY", "")
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -59,7 +65,7 @@ class AppConfig:
|
|||||||
|
|
||||||
def load_config(path: str | None = None) -> AppConfig:
|
def load_config(path: str | None = None) -> AppConfig:
|
||||||
"""从 YAML 文件加载配置, 若文件不存在则返回默认配置."""
|
"""从 YAML 文件加载配置, 若文件不存在则返回默认配置."""
|
||||||
config_path = path or "config.yaml"
|
config_path = path or DEFAULT_CONFIG_PATH
|
||||||
if not Path(config_path).exists():
|
if not Path(config_path).exists():
|
||||||
return AppConfig()
|
return AppConfig()
|
||||||
|
|
||||||
|
|||||||
+4
-4
@@ -5,7 +5,7 @@ from pathlib import Path
|
|||||||
from fastapi import FastAPI, HTTPException
|
from fastapi import FastAPI, HTTPException
|
||||||
from pydantic import BaseModel
|
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.db import VectorDB
|
||||||
from src.core.embedder import create_embedder
|
from src.core.embedder import create_embedder
|
||||||
from src.core.ingest import DocumentIngestor
|
from src.core.ingest import DocumentIngestor
|
||||||
@@ -42,9 +42,9 @@ def _get_db():
|
|||||||
def _get_embedder():
|
def _get_embedder():
|
||||||
global _embedder
|
global _embedder
|
||||||
if _embedder is None:
|
if _embedder is None:
|
||||||
# 先尝试从 config.yaml 加载, 失败则用默认 local
|
# 先尝试从配置文件加载, 失败则用默认 local
|
||||||
try:
|
try:
|
||||||
cfg = load_config("config.yaml")
|
cfg = load_config(DEFAULT_CONFIG_PATH)
|
||||||
embed_cfg = cfg.embed
|
embed_cfg = cfg.embed
|
||||||
except Exception:
|
except Exception:
|
||||||
embed_cfg = EmbedConfig(mode="local")
|
embed_cfg = EmbedConfig(mode="local")
|
||||||
@@ -56,7 +56,7 @@ def _init_services():
|
|||||||
global _searcher, _ingestor
|
global _searcher, _ingestor
|
||||||
collection = os.getenv("MD_VECTOR_DB_COLLECTION", "markdown_docs")
|
collection = os.getenv("MD_VECTOR_DB_COLLECTION", "markdown_docs")
|
||||||
try:
|
try:
|
||||||
cfg = load_config("config.yaml")
|
cfg = load_config(DEFAULT_CONFIG_PATH)
|
||||||
collection = cfg.chroma.collection_name
|
collection = cfg.chroma.collection_name
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
"""配置加载模块测试."""
|
"""配置加载模块测试."""
|
||||||
|
import os
|
||||||
import tempfile
|
import tempfile
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
@@ -80,3 +81,26 @@ server:
|
|||||||
cfg = load_config("/nonexistent/config.yaml")
|
cfg = load_config("/nonexistent/config.yaml")
|
||||||
assert cfg.chroma.persist_dir == "./data"
|
assert cfg.chroma.persist_dir == "./data"
|
||||||
assert cfg.embed.mode == "local"
|
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)
|
||||||
|
|||||||
Reference in New Issue
Block a user