test: 添加 verify_api_key 和 RateLimiter 完整单元测试
Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
+12
-5
@@ -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)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user