405303e82c
Batch 1 — CRITICAL (1): - 提取 is_safe_path() 到 src/core/security.py 公共模块 - CLI 和 ingest_obsidian.py 统一添加路径遍历防护 Batch 2 — HIGH (13) + 架构重构: - CLI 复用 deps.py AppState, 消除 30 行重复代码 - AppState/get_state 添加线程安全锁 - serve 命令传递 --config 到 uvicorn (H1) - OpenAIEmbedder 懒创建+复用 HTTP 客户端 (H2) - DashscopeEmbedder import 移到模块顶部 (H3) - 路径检查改用 os.path.commonpath (H4) - embedder.embed() 返回值长度检查 (H5) - 健康检查不泄露内部错误详情 (H7) - /api/v1/collections 添加 API Key 认证 (H8) - API Key 使用 hmac.compare_digest 恒定时间比较 (H9) - 添加 CORS 中间件 (H10) - ServerConfig 支持 SSL 配置 (H11) - HF_ENDPOINT 修改添加详细注释 (H12) Batch 3 — MEDIUM (20) + Splitter Protocol: - 定义 Splitter(Protocol) 接口, DocumentIngestor 接受可选 splitter - DashScope 响应添加结构验证 (M2) - ingest_obsidian.py 支持 CLI 参数和 OBSIDIAN_DIRS 环境变量 (M6) - scripts/serve.py 添加废弃警告 (M7) - content 限制 500KB, collection 正则限制字符集 (M12-M14) - 默认监听地址 127.0.0.1 (M16) - 添加安全响应头中间件 (M17) - verify_api_key 认证失败记录日志 (M19) Batch 4 — LOW (10): - CLI emoji 清理为纯文本标记 (L5) - logging.basicConfig 移到 FastAPI lifespan (L1) - VectorDB 添加 write_guard() 上下文管理器 (L3) - IngestRequest file_path/content 互斥校验 (L10) - ingest_obsidian.py 注释修正 (L6) 测试: 46 → 70 (+24) - tests/test_security.py: 11 个路径安全测试 - tests/test_deps.py: 11 个依赖注入测试 Co-Authored-By: Claude <noreply@anthropic.com>
98 lines
3.3 KiB
Python
98 lines
3.3 KiB
Python
"""AppState 和依赖注入测试."""
|
|
import os
|
|
import tempfile
|
|
|
|
import pytest
|
|
|
|
from src.core.config import load_config
|
|
from src.server.deps import AppState
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def _clean_env(monkeypatch):
|
|
"""清除环境变量防止其他测试污染."""
|
|
monkeypatch.delenv("MD_VECTOR_DB_COLLECTION", raising=False)
|
|
monkeypatch.delenv("MD_VECTOR_DB_DATA_DIR", raising=False)
|
|
|
|
|
|
class TestAppState:
|
|
"""AppState 类测试."""
|
|
|
|
def test_init_with_config(self):
|
|
"""使用默认配置初始化."""
|
|
state = AppState()
|
|
assert state.config is not None
|
|
assert state.db is not None
|
|
assert state.embedder is not None
|
|
assert state.default_collection == state.config.chroma.collection_name
|
|
|
|
def test_get_searcher_returns_cached(self):
|
|
"""同一 collection 多次调用返回同一实例."""
|
|
state = AppState()
|
|
s1 = state.get_searcher("test_coll")
|
|
s2 = state.get_searcher("test_coll")
|
|
assert s1 is s2
|
|
|
|
def test_get_searcher_different_collections(self):
|
|
"""不同 collection 返回不同实例."""
|
|
state = AppState()
|
|
s1 = state.get_searcher("coll_a")
|
|
s2 = state.get_searcher("coll_b")
|
|
assert s1 is not s2
|
|
|
|
def test_get_searcher_uses_default_when_none(self):
|
|
"""collection 为 None 时使用默认值."""
|
|
state = AppState()
|
|
s = state.get_searcher(None)
|
|
assert s.collection_name == state.default_collection
|
|
|
|
def test_get_ingestor_returns_cached(self):
|
|
"""同一 collection 多次调用返回同一实例."""
|
|
state = AppState()
|
|
i1 = state.get_ingestor("test_coll")
|
|
i2 = state.get_ingestor("test_coll")
|
|
assert i1 is i2
|
|
|
|
def test_default_collection_from_config(self):
|
|
"""默认 collection 名从配置读取."""
|
|
state = AppState()
|
|
assert isinstance(state.default_collection, str)
|
|
assert len(state.default_collection) > 0
|
|
|
|
def test_list_collections_with_stats(self):
|
|
"""列出集合统计."""
|
|
state = AppState()
|
|
result = state.list_collections_with_stats()
|
|
assert isinstance(result, list)
|
|
|
|
def test_is_healthy_returns_status(self):
|
|
"""健康检查返回正确结构."""
|
|
state = AppState()
|
|
result = state.is_healthy()
|
|
assert "status" in result
|
|
assert "checks" in result
|
|
assert "chromadb" in result["checks"]
|
|
assert "embedder" in result["checks"]
|
|
|
|
def test_is_healthy_chromadb_ok(self):
|
|
"""健康检查 ChromaDB 正常."""
|
|
state = AppState()
|
|
result = state.is_healthy()
|
|
assert result["checks"]["chromadb"]["status"] == "ok"
|
|
|
|
def test_is_healthy_embedder_ok(self):
|
|
"""健康检查 Embedder 正常."""
|
|
state = AppState()
|
|
result = state.is_healthy()
|
|
assert result["checks"]["embedder"]["status"] == "ok"
|
|
|
|
def test_is_healthy_no_detail_leak(self):
|
|
"""健康检查不泄露内部详情."""
|
|
state = AppState()
|
|
result = state.is_healthy()
|
|
for component in ("chromadb", "embedder"):
|
|
detail = result["checks"][component].get("detail", "")
|
|
if detail:
|
|
# 如果出错,detail 应该是通用消息,不是异常堆栈
|
|
assert "unavailable" in str(detail).lower() or len(detail) < 100
|