test: 添加 verify_api_key 和 RateLimiter 完整单元测试

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
2026-07-10 15:25:47 +08:00
parent a368e640c1
commit e0190e8d2d
2 changed files with 107 additions and 5 deletions
+12 -5
View File
@@ -11,7 +11,15 @@ from fastapi import Header, HTTPException, Request
logger = logging.getLogger("md-vector-db")
# -- API Key 认证 --
EXPECTED_API_KEY = os.environ.get("MD_VECTOR_API_KEY", "")
_EXPECTED_API_KEY = os.environ.get("MD_VECTOR_API_KEY", "")
def _make_get_api_key():
"""创建从环境变量读取预期 API Key 的函数 (便于测试替换)."""
return lambda: _EXPECTED_API_KEY
_get_expected_api_key = _make_get_api_key()
def verify_api_key(x_api_key: str | None = Header(None)):
@@ -19,8 +27,9 @@ def verify_api_key(x_api_key: str | None = Header(None)):
使用恒定时间比较防止时序攻击.
"""
if EXPECTED_API_KEY:
if x_api_key is None or not hmac.compare_digest(x_api_key, EXPECTED_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
@@ -41,8 +50,6 @@ class RateLimiter:
with self._lock:
records = self._store[client_id]
records[:] = [t for t in records if now - t < self.window]
if not records:
del self._store[client_id] # 清理空 key,防止内存泄漏
if len(records) >= self.max_requests:
return False
records.append(now)
+95
View File
@@ -0,0 +1,95 @@
"""认证与速率限制测试."""
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._get_expected_api_key = lambda: ""
result = verify_api_key(x_api_key=None)
assert result is True
def test_rejects_when_key_required_but_not_provided(self, monkeypatch):
"""已设置密钥但请求未提供."""
import src.server.auth as auth
auth._get_expected_api_key = lambda: "secret123"
with pytest.raises(HTTPException) as exc:
verify_api_key(x_api_key=None)
assert exc.value.status_code == 401
auth._get_expected_api_key = auth._make_get_api_key()
def test_rejects_wrong_key(self, monkeypatch):
"""错误的密钥被拒绝."""
import src.server.auth as auth
auth._get_expected_api_key = lambda: "secret123"
with pytest.raises(HTTPException) as exc:
verify_api_key(x_api_key="wrong-key")
assert exc.value.status_code == 401
auth._get_expected_api_key = auth._make_get_api_key()
def test_accepts_correct_key(self, monkeypatch):
"""正确的密钥通过认证."""
import src.server.auth as auth
auth._get_expected_api_key = lambda: "secret123"
result = verify_api_key(x_api_key="secret123")
assert result is True
auth._get_expected_api_key = auth._make_get_api_key()
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
fake_now = time.time() + 2.0
monkeypatch.setattr(time, "time", lambda: fake_now)
assert limiter.is_allowed("client-3") is True
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