mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
feat(rate-limit): 实现分层 RPM 限速,支持系统默认/用户/独立Key三级配置
- 新增用户级 rate_limit 字段,支持系统默认/用户自定义/不限制三种模式 - 独立 Key 的 rate_limit 语义调整:null=跟随系统默认,0=不限制,>0=自定义 - 实现 UserRpmLimiter 基于 Redis sliding window 的 RPM 限速引擎 - Pipeline 请求流程集成用户级 RPM 检查 - 管理后台和用户面板新增 RPM 限速配置与实时状态查看 - 系统设置新增全局默认 RPM 配置项 - 迁移脚本回填现有 API Key 的 rate_limit 默认值 - 新增用户/Key RPM 状态监控 API 和前端展示 Closes #231 Co-authored-by: LewisPen <LewisPen@nyadoo.com>
This commit is contained in:
@@ -1,7 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime, timezone
|
||||
from types import SimpleNamespace
|
||||
from typing import Generator
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
@@ -18,6 +20,7 @@ from src.api.admin.api_keys.routes import router as admin_api_keys_router
|
||||
from src.api.admin.users.routes import (
|
||||
AdminGetUserKeyFullKeyAdapter,
|
||||
AdminToggleUserKeyLockAdapter,
|
||||
AdminUpdateUserKeyAdapter,
|
||||
)
|
||||
from src.api.admin.users.routes import router as admin_users_router
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
@@ -25,6 +28,15 @@ from src.database import get_db
|
||||
from src.models.api import CreateApiKeyRequest
|
||||
|
||||
|
||||
def _patch_get_db_context(monkeypatch: pytest.MonkeyPatch, db: MagicMock) -> None:
|
||||
@contextmanager
|
||||
def _fake_ctx() -> Generator[MagicMock, None, None]:
|
||||
yield db
|
||||
|
||||
monkeypatch.setattr("src.api.admin.users.routes.get_db_context", _fake_ctx)
|
||||
monkeypatch.setattr("src.api.admin.api_keys.routes.get_db_context", _fake_ctx)
|
||||
|
||||
|
||||
def _build_context(db: MagicMock) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
db=db,
|
||||
@@ -46,11 +58,15 @@ def _build_admin_users_app(db: MagicMock, monkeypatch: pytest.MonkeyPatch) -> Te
|
||||
*, adapter: object, http_request: object, db: MagicMock, mode: object
|
||||
) -> object:
|
||||
_ = http_request, mode
|
||||
try:
|
||||
payload = await http_request.json()
|
||||
except Exception:
|
||||
payload = {}
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
user=SimpleNamespace(id="admin-1"),
|
||||
ensure_json_body=lambda: {},
|
||||
ensure_json_body=lambda: payload,
|
||||
add_audit_metadata=lambda **_: None,
|
||||
)
|
||||
return await adapter.handle(context)
|
||||
@@ -82,10 +98,11 @@ def _build_admin_api_keys_app(db: MagicMock, monkeypatch: pytest.MonkeyPatch) ->
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_toggle_user_key_lock_adapter_success() -> None:
|
||||
async def test_toggle_user_key_lock_adapter_success(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
db = MagicMock()
|
||||
api_key = SimpleNamespace(id="key-1", user_id="user-1", is_standalone=False, is_locked=False)
|
||||
_mock_query_first(db, api_key)
|
||||
_patch_get_db_context(monkeypatch, db)
|
||||
|
||||
adapter = AdminToggleUserKeyLockAdapter(user_id="user-1", key_id="key-1")
|
||||
result = await adapter.handle(_build_context(db))
|
||||
@@ -98,9 +115,12 @@ async def test_toggle_user_key_lock_adapter_success() -> None:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_toggle_user_key_lock_adapter_not_found_for_standalone_or_wrong_owner() -> None:
|
||||
async def test_toggle_user_key_lock_adapter_not_found_for_standalone_or_wrong_owner(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
_mock_query_first(db, None)
|
||||
_patch_get_db_context(monkeypatch, db)
|
||||
|
||||
adapter = AdminToggleUserKeyLockAdapter(user_id="user-1", key_id="key-standalone")
|
||||
with pytest.raises(NotFoundException):
|
||||
@@ -170,7 +190,9 @@ async def test_get_user_key_full_key_adapter_returns_500_on_decrypt_error(
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_standalone_toggle_adapters_reject_normal_user_key() -> None:
|
||||
async def test_standalone_toggle_adapters_reject_normal_user_key(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
normal_key = SimpleNamespace(
|
||||
id="key-user",
|
||||
@@ -182,6 +204,7 @@ async def test_standalone_toggle_adapters_reject_normal_user_key() -> None:
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
_mock_query_first(db, normal_key)
|
||||
_patch_get_db_context(monkeypatch, db)
|
||||
context = _build_context(db)
|
||||
|
||||
with pytest.raises(InvalidRequestException):
|
||||
@@ -195,6 +218,7 @@ def test_user_key_lock_route_path_smoke(monkeypatch: pytest.MonkeyPatch) -> None
|
||||
db = MagicMock()
|
||||
api_key = SimpleNamespace(id="key-5", user_id="user-2", is_standalone=False, is_locked=False)
|
||||
_mock_query_first(db, api_key)
|
||||
_patch_get_db_context(monkeypatch, db)
|
||||
client = _build_admin_users_app(db, monkeypatch)
|
||||
|
||||
response = client.patch("/api/admin/users/user-2/api-keys/key-5/lock")
|
||||
@@ -220,6 +244,72 @@ def test_user_key_full_key_route_path_smoke(monkeypatch: pytest.MonkeyPatch) ->
|
||||
assert response.json() == {"key": "sk-user-route-key"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_user_key_adapter_passes_rate_limit_and_name(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def _update_user_key_sync(
|
||||
user_id: str, key_id: str, request: object
|
||||
) -> tuple[dict[str, object], dict[str, object]]:
|
||||
captured["user_id"] = user_id
|
||||
captured["key_id"] = key_id
|
||||
captured["name"] = getattr(request, "name", None)
|
||||
captured["rate_limit"] = getattr(request, "rate_limit", None)
|
||||
return {"id": key_id, "name": captured["name"], "rate_limit": captured["rate_limit"]}, {}
|
||||
|
||||
monkeypatch.setattr("src.api.admin.users.routes._update_user_key_sync", _update_user_key_sync)
|
||||
|
||||
adapter = AdminUpdateUserKeyAdapter(user_id="user-1", key_id="key-7")
|
||||
context = SimpleNamespace(
|
||||
db=MagicMock(),
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
ensure_json_body=lambda: {"name": "Renamed Key", "rate_limit": 12},
|
||||
add_audit_metadata=lambda **_: None,
|
||||
)
|
||||
|
||||
result = await adapter.handle(context)
|
||||
|
||||
assert result["id"] == "key-7"
|
||||
assert captured == {
|
||||
"user_id": "user-1",
|
||||
"key_id": "key-7",
|
||||
"name": "Renamed Key",
|
||||
"rate_limit": 12,
|
||||
}
|
||||
|
||||
|
||||
def test_update_user_key_route_path_smoke(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def _update_user_key_sync(
|
||||
user_id: str, key_id: str, request: object
|
||||
) -> tuple[dict[str, object], dict[str, object]]:
|
||||
captured["user_id"] = user_id
|
||||
captured["key_id"] = key_id
|
||||
captured["name"] = getattr(request, "name", None)
|
||||
captured["rate_limit"] = getattr(request, "rate_limit", None)
|
||||
return {"id": key_id, "name": captured["name"], "rate_limit": captured["rate_limit"]}, {}
|
||||
|
||||
monkeypatch.setattr("src.api.admin.users.routes._update_user_key_sync", _update_user_key_sync)
|
||||
client = _build_admin_users_app(MagicMock(), monkeypatch)
|
||||
|
||||
response = client.put(
|
||||
"/api/admin/users/user-2/api-keys/key-8",
|
||||
json={"name": "Updated", "rate_limit": 9},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["rate_limit"] == 9
|
||||
assert captured == {
|
||||
"user_id": "user-2",
|
||||
"key_id": "key-8",
|
||||
"name": "Updated",
|
||||
"rate_limit": 9,
|
||||
}
|
||||
|
||||
|
||||
def test_standalone_lock_route_removed(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
client = _build_admin_api_keys_app(MagicMock(), monkeypatch)
|
||||
response = client.patch("/api/admin/api-keys/key-1/lock")
|
||||
@@ -228,6 +318,7 @@ def test_standalone_lock_route_removed(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
|
||||
def test_standalone_list_route_does_not_expose_is_locked(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
db = MagicMock()
|
||||
_patch_get_db_context(monkeypatch, db)
|
||||
api_key = SimpleNamespace(
|
||||
id="sa-key-1",
|
||||
user_id="admin-1",
|
||||
@@ -306,6 +397,7 @@ async def test_create_standalone_key_adapter_preserves_empty_restriction_lists(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
_patch_get_db_context(monkeypatch, db)
|
||||
captured: dict[str, object] = {}
|
||||
created_key = SimpleNamespace(
|
||||
id="sa-key-3",
|
||||
@@ -365,6 +457,7 @@ async def test_update_standalone_key_adapter_preserves_empty_restriction_lists(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
_patch_get_db_context(monkeypatch, db)
|
||||
existing_key = SimpleNamespace(id="sa-key-4", is_standalone=True)
|
||||
_mock_query_first(db, existing_key)
|
||||
|
||||
|
||||
@@ -49,6 +49,7 @@ def test_list_users_uses_wallet_batch_lookup(monkeypatch: pytest.MonkeyPatch) ->
|
||||
allowed_providers=None,
|
||||
allowed_api_formats=None,
|
||||
allowed_models=None,
|
||||
rate_limit=None,
|
||||
is_active=True,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
@@ -62,6 +63,7 @@ def test_list_users_uses_wallet_batch_lookup(monkeypatch: pytest.MonkeyPatch) ->
|
||||
allowed_providers=None,
|
||||
allowed_api_formats=None,
|
||||
allowed_models=None,
|
||||
rate_limit=None,
|
||||
is_active=True,
|
||||
created_at=now,
|
||||
updated_at=None,
|
||||
@@ -101,24 +103,14 @@ async def test_create_user_adapter_preserves_empty_restriction_lists(
|
||||
db = MagicMock()
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
def _create_user(**kwargs: Any) -> SimpleNamespace:
|
||||
captured.update(kwargs)
|
||||
return SimpleNamespace(
|
||||
id="user-3",
|
||||
email="u3@example.com",
|
||||
username="user3",
|
||||
role=SimpleNamespace(value="user"),
|
||||
is_active=True,
|
||||
allowed_providers=[],
|
||||
allowed_api_formats=[],
|
||||
allowed_models=[],
|
||||
)
|
||||
def _create_user_sync(request: Any, role: Any) -> tuple[dict[str, Any], dict[str, Any]]:
|
||||
captured["allowed_providers"] = request.allowed_providers
|
||||
captured["allowed_api_formats"] = request.allowed_api_formats
|
||||
captured["allowed_models"] = request.allowed_models
|
||||
captured["role"] = role
|
||||
return {"id": "user-3"}, {}
|
||||
|
||||
monkeypatch.setattr("src.api.admin.users.routes.UserService.create_user", _create_user)
|
||||
monkeypatch.setattr(
|
||||
"src.api.admin.users.routes._serialize_user",
|
||||
lambda _db, user: {"id": user.id},
|
||||
)
|
||||
monkeypatch.setattr("src.api.admin.users.routes._create_user_sync", _create_user_sync)
|
||||
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
|
||||
120
tests/api/test_monitoring_user_rate_limit.py
Normal file
120
tests/api/test_monitoring_user_rate_limit.py
Normal file
@@ -0,0 +1,120 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from src.api.monitoring.user import UserRateLimitStatusAdapter
|
||||
|
||||
|
||||
def _build_query_with_keys(keys: list[object]) -> MagicMock:
|
||||
query = MagicMock()
|
||||
query.filter.return_value.order_by.return_value.all.return_value = keys
|
||||
return query
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rate_limit_status_adapter_reports_user_and_key_layers(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
now = datetime(2026, 3, 13, 12, 0, tzinfo=timezone.utc)
|
||||
db = MagicMock()
|
||||
user = SimpleNamespace(id="user-1", rate_limit=None)
|
||||
key = SimpleNamespace(id="key-1", name="Primary", is_standalone=False, rate_limit=10)
|
||||
standalone = SimpleNamespace(
|
||||
id="skey-1", name="Standalone", is_standalone=True, rate_limit=None
|
||||
)
|
||||
|
||||
db.query.return_value = _build_query_with_keys([key, standalone])
|
||||
|
||||
limiter = MagicMock()
|
||||
limiter.bucket_seconds = 60
|
||||
limiter.get_reset_at.return_value = now
|
||||
limiter.get_user_rpm_key.return_value = "rpm:user:user-1:bucket"
|
||||
limiter.get_standalone_rpm_key.return_value = "rpm:ukey:skey-1:bucket"
|
||||
limiter.get_key_rpm_key.side_effect = lambda key_id: f"rpm:key:{key_id}:bucket"
|
||||
|
||||
async def _get_scope_count(scope_key: str) -> int:
|
||||
counts = {
|
||||
"rpm:user:user-1:bucket": 55,
|
||||
"rpm:key:key-1:bucket": 7,
|
||||
"rpm:ukey:skey-1:bucket": 12,
|
||||
}
|
||||
return counts[scope_key]
|
||||
|
||||
limiter.get_scope_count = AsyncMock(side_effect=_get_scope_count)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.monitoring.user.get_user_rpm_limiter",
|
||||
AsyncMock(return_value=limiter),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.monitoring.user.SystemConfigService.get_config",
|
||||
lambda *_a, **_k: 60,
|
||||
)
|
||||
|
||||
context = SimpleNamespace(db=db, user=user)
|
||||
result = await UserRateLimitStatusAdapter().handle(context)
|
||||
|
||||
assert result["user_id"] == "user-1"
|
||||
assert result["api_keys"][0] == {
|
||||
"api_key_name": "Primary",
|
||||
"limit": 10,
|
||||
"remaining": 3,
|
||||
"scope": "key",
|
||||
"reset_time": now.isoformat(),
|
||||
"window": "60s",
|
||||
"user_limit": 60,
|
||||
"user_remaining": 5,
|
||||
"key_limit": 10,
|
||||
"key_remaining": 3,
|
||||
}
|
||||
assert result["api_keys"][1] == {
|
||||
"api_key_name": "Standalone",
|
||||
"limit": 60,
|
||||
"remaining": 48,
|
||||
"scope": "user",
|
||||
"reset_time": now.isoformat(),
|
||||
"window": "60s",
|
||||
"user_limit": 60,
|
||||
"user_remaining": 48,
|
||||
"key_limit": None,
|
||||
"key_remaining": None,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rate_limit_status_adapter_reports_unlimited_key_without_counts(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
db = MagicMock()
|
||||
user = SimpleNamespace(id="user-1", rate_limit=0)
|
||||
key = SimpleNamespace(id="key-1", name="Unlimited", is_standalone=False, rate_limit=0)
|
||||
|
||||
db.query.return_value = _build_query_with_keys([key])
|
||||
|
||||
limiter = MagicMock()
|
||||
limiter.bucket_seconds = 60
|
||||
limiter.get_reset_at.return_value = datetime.now(timezone.utc)
|
||||
limiter.get_user_rpm_key.return_value = "rpm:user:user-1:bucket"
|
||||
limiter.get_key_rpm_key.return_value = "rpm:key:key-1:bucket"
|
||||
limiter.get_scope_count = AsyncMock()
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.monitoring.user.get_user_rpm_limiter",
|
||||
AsyncMock(return_value=limiter),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.monitoring.user.SystemConfigService.get_config",
|
||||
lambda *_a, **_k: 60,
|
||||
)
|
||||
|
||||
context = SimpleNamespace(db=db, user=user)
|
||||
result = await UserRateLimitStatusAdapter().handle(context)
|
||||
|
||||
assert result["api_keys"][0]["limit"] is None
|
||||
assert result["api_keys"][0]["remaining"] is None
|
||||
assert result["api_keys"][0]["scope"] is None
|
||||
limiter.get_scope_count.assert_not_awaited()
|
||||
@@ -18,6 +18,7 @@ from src.api.base.adapter import ApiMode
|
||||
from src.api.base.pipeline import ApiRequestPipeline
|
||||
from src.core.enums import UserRole
|
||||
from src.core.modules.hooks import AUTH_TOKEN_PREFIX_AUTHENTICATORS
|
||||
from src.services.rate_limit.user_rpm_limiter import RpmCheckResult
|
||||
|
||||
|
||||
class TestPipelineBalanceCalculation:
|
||||
@@ -499,6 +500,166 @@ class TestPipelineAuthentication:
|
||||
assert "锁定" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
class TestPipelineUserRateLimit:
|
||||
@pytest.fixture
|
||||
def pipeline(self) -> ApiRequestPipeline:
|
||||
return ApiRequestPipeline()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_user_rate_limit_uses_system_default_for_user_scope(
|
||||
self, pipeline: ApiRequestPipeline, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
request = MagicMock()
|
||||
request.state = MagicMock()
|
||||
db = MagicMock()
|
||||
user = MagicMock(id="user-1", rate_limit=None)
|
||||
api_key = MagicMock(id="key-1", is_standalone=False, rate_limit=0)
|
||||
|
||||
limiter = MagicMock()
|
||||
limiter.get_user_rpm_key.return_value = "rpm:user:user-1:1"
|
||||
limiter.get_key_rpm_key.return_value = "rpm:key:key-1:1"
|
||||
limiter.check_and_consume = AsyncMock(return_value=RpmCheckResult(allowed=True))
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.base.pipeline.get_user_rpm_limiter",
|
||||
AsyncMock(return_value=limiter),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.base.pipeline.SystemConfigService.get_config",
|
||||
lambda *_a, **_k: 60,
|
||||
)
|
||||
|
||||
await pipeline._check_user_rate_limit(request, db, user, api_key)
|
||||
|
||||
limiter.check_and_consume.assert_awaited_once_with(
|
||||
user_rpm_key="rpm:user:user-1:1",
|
||||
user_rpm_limit=60,
|
||||
key_rpm_key="rpm:key:key-1:1",
|
||||
key_rpm_limit=0,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_user_rate_limit_returns_429_with_scope_header(
|
||||
self, pipeline: ApiRequestPipeline, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
request = MagicMock()
|
||||
request.state = MagicMock()
|
||||
db = MagicMock()
|
||||
user = MagicMock(id="user-1", rate_limit=100)
|
||||
api_key = MagicMock(id="key-1", is_standalone=False, rate_limit=10)
|
||||
|
||||
limiter = MagicMock()
|
||||
limiter.get_user_rpm_key.return_value = "rpm:user:user-1:1"
|
||||
limiter.get_key_rpm_key.return_value = "rpm:key:key-1:1"
|
||||
limiter.get_retry_after.return_value = 17
|
||||
limiter.check_and_consume = AsyncMock(
|
||||
return_value=RpmCheckResult(
|
||||
allowed=False,
|
||||
scope="key",
|
||||
limit=10,
|
||||
remaining=0,
|
||||
retry_after=17,
|
||||
)
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.base.pipeline.get_user_rpm_limiter",
|
||||
AsyncMock(return_value=limiter),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.base.pipeline.SystemConfigService.get_config",
|
||||
lambda *_a, **_k: 60,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await pipeline._check_user_rate_limit(request, db, user, api_key)
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
assert exc_info.value.headers == {
|
||||
"Retry-After": "17",
|
||||
"X-RateLimit-Limit": "10",
|
||||
"X-RateLimit-Remaining": "0",
|
||||
"X-RateLimit-Scope": "key",
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_user_rate_limit_uses_system_default_for_standalone_key(
|
||||
self, pipeline: ApiRequestPipeline, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
request = MagicMock()
|
||||
request.state = MagicMock()
|
||||
db = MagicMock()
|
||||
user = MagicMock(id="user-1", rate_limit=999)
|
||||
api_key = MagicMock(id="standalone-1", is_standalone=True, rate_limit=None)
|
||||
|
||||
limiter = MagicMock()
|
||||
limiter.get_standalone_rpm_key.return_value = "rpm:ukey:standalone-1:1"
|
||||
limiter.get_key_rpm_key.return_value = "rpm:key:standalone-1:1"
|
||||
limiter.check_and_consume = AsyncMock(return_value=RpmCheckResult(allowed=True))
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.base.pipeline.get_user_rpm_limiter",
|
||||
AsyncMock(return_value=limiter),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.base.pipeline.SystemConfigService.get_config",
|
||||
lambda *_a, **_k: 60,
|
||||
)
|
||||
|
||||
await pipeline._check_user_rate_limit(request, db, user, api_key)
|
||||
|
||||
limiter.check_and_consume.assert_awaited_once_with(
|
||||
user_rpm_key="rpm:ukey:standalone-1:1",
|
||||
user_rpm_limit=60,
|
||||
key_rpm_key="rpm:key:standalone-1:1",
|
||||
key_rpm_limit=0,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_user_rate_limit_returns_429_with_user_scope_header(
|
||||
self, pipeline: ApiRequestPipeline, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
request = MagicMock()
|
||||
request.state = MagicMock()
|
||||
db = MagicMock()
|
||||
user = MagicMock(id="user-1", rate_limit=3)
|
||||
api_key = MagicMock(id="key-1", is_standalone=False, rate_limit=10)
|
||||
|
||||
limiter = MagicMock()
|
||||
limiter.get_user_rpm_key.return_value = "rpm:user:user-1:1"
|
||||
limiter.get_key_rpm_key.return_value = "rpm:key:key-1:1"
|
||||
limiter.get_retry_after.return_value = 23
|
||||
limiter.check_and_consume = AsyncMock(
|
||||
return_value=RpmCheckResult(
|
||||
allowed=False,
|
||||
scope="user",
|
||||
limit=3,
|
||||
remaining=0,
|
||||
retry_after=23,
|
||||
)
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.api.base.pipeline.get_user_rpm_limiter",
|
||||
AsyncMock(return_value=limiter),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"src.api.base.pipeline.SystemConfigService.get_config",
|
||||
lambda *_a, **_k: 60,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await pipeline._check_user_rate_limit(request, db, user, api_key)
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
assert exc_info.value.headers == {
|
||||
"Retry-After": "23",
|
||||
"X-RateLimit-Limit": "3",
|
||||
"X-RateLimit-Remaining": "0",
|
||||
"X-RateLimit-Scope": "user",
|
||||
}
|
||||
|
||||
|
||||
class TestPipelineTokenPrefixAuth:
|
||||
"""Tests token-prefix auth isolation."""
|
||||
|
||||
|
||||
110
tests/api/test_user_me_api_key_routes.py
Normal file
110
tests/api/test_user_me_api_key_routes.py
Normal file
@@ -0,0 +1,110 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from src.api.user_me.routes import UpdateMyApiKeyAdapter
|
||||
from src.api.user_me.routes import router as me_router
|
||||
from src.database import get_db
|
||||
|
||||
|
||||
def _build_me_app(db: MagicMock, monkeypatch: Any) -> TestClient:
|
||||
app = FastAPI()
|
||||
app.include_router(me_router)
|
||||
app.dependency_overrides[get_db] = lambda: db
|
||||
|
||||
async def _fake_pipeline_run(
|
||||
*, adapter: object, http_request: object, db: MagicMock, mode: object
|
||||
) -> object:
|
||||
_ = http_request, mode
|
||||
try:
|
||||
payload = await http_request.json()
|
||||
except Exception:
|
||||
payload = {}
|
||||
context = SimpleNamespace(
|
||||
db=db,
|
||||
user=SimpleNamespace(id="user-1", email="u@example.com"),
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
ensure_json_body=lambda: payload,
|
||||
add_audit_metadata=lambda **_: None,
|
||||
)
|
||||
return await adapter.handle(context)
|
||||
|
||||
monkeypatch.setattr("src.api.user_me.routes.pipeline.run", _fake_pipeline_run)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
async def _fake_update_my_api_key_sync(
|
||||
user_id: str,
|
||||
key_id: str,
|
||||
request: object,
|
||||
captured: dict[str, object],
|
||||
) -> dict[str, object]:
|
||||
captured["user_id"] = user_id
|
||||
captured["key_id"] = key_id
|
||||
captured["name"] = getattr(request, "name", None)
|
||||
captured["rate_limit"] = getattr(request, "rate_limit", None)
|
||||
return {"id": key_id, "name": captured["name"], "rate_limit": captured["rate_limit"]}
|
||||
|
||||
|
||||
def test_update_my_api_key_route_path_smoke(monkeypatch: Any) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def _sync(user_id: str, key_id: str, request: object) -> dict[str, object]:
|
||||
captured["user_id"] = user_id
|
||||
captured["key_id"] = key_id
|
||||
captured["name"] = getattr(request, "name", None)
|
||||
captured["rate_limit"] = getattr(request, "rate_limit", None)
|
||||
return {"id": key_id, "name": captured["name"], "rate_limit": captured["rate_limit"]}
|
||||
|
||||
monkeypatch.setattr("src.api.user_me.routes._update_my_api_key_sync", _sync)
|
||||
client = _build_me_app(MagicMock(), monkeypatch)
|
||||
|
||||
response = client.put("/api/users/me/api-keys/key-1", json={"name": "Edited", "rate_limit": 6})
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["rate_limit"] == 6
|
||||
assert captured == {
|
||||
"user_id": "user-1",
|
||||
"key_id": "key-1",
|
||||
"name": "Edited",
|
||||
"rate_limit": 6,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_my_api_key_adapter_passes_rate_limit_and_name(monkeypatch: Any) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def _sync(user_id: str, key_id: str, request: object) -> dict[str, object]:
|
||||
captured["user_id"] = user_id
|
||||
captured["key_id"] = key_id
|
||||
captured["name"] = getattr(request, "name", None)
|
||||
captured["rate_limit"] = getattr(request, "rate_limit", None)
|
||||
return {"id": key_id, "name": captured["name"], "rate_limit": captured["rate_limit"]}
|
||||
|
||||
monkeypatch.setattr("src.api.user_me.routes._update_my_api_key_sync", _sync)
|
||||
|
||||
adapter = UpdateMyApiKeyAdapter(key_id="key-2")
|
||||
context = SimpleNamespace(
|
||||
db=MagicMock(),
|
||||
user=SimpleNamespace(id="user-1"),
|
||||
request=SimpleNamespace(state=SimpleNamespace()),
|
||||
ensure_json_body=lambda: {"name": "Edited Again", "rate_limit": 15},
|
||||
add_audit_metadata=lambda **_: None,
|
||||
)
|
||||
|
||||
result = await adapter.handle(context)
|
||||
|
||||
assert result["id"] == "key-2"
|
||||
assert captured == {
|
||||
"user_id": "user-1",
|
||||
"key_id": "key-2",
|
||||
"name": "Edited Again",
|
||||
"rate_limit": 15,
|
||||
}
|
||||
148
tests/services/test_user_rpm_limiter.py
Normal file
148
tests/services/test_user_rpm_limiter.py
Normal file
@@ -0,0 +1,148 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncGenerator
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.rate_limit.user_rpm_limiter import RpmCheckResult, UserRpmLimiter
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def limiter(monkeypatch: pytest.MonkeyPatch) -> AsyncGenerator[UserRpmLimiter]:
|
||||
limiter = UserRpmLimiter()
|
||||
limiter._redis = None
|
||||
limiter._memory_counts.clear()
|
||||
await limiter.close()
|
||||
# 阻止 initialize() 连接真实 Redis,确保纯内存模式
|
||||
monkeypatch.setattr(limiter, "initialize", AsyncMock())
|
||||
yield limiter
|
||||
limiter._memory_counts.clear()
|
||||
limiter._redis = None
|
||||
await limiter.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_and_consume_enforces_user_scope_in_memory(
|
||||
limiter: UserRpmLimiter,
|
||||
) -> None:
|
||||
user_key = limiter.get_user_rpm_key("user-1")
|
||||
key_key = limiter.get_key_rpm_key("key-1")
|
||||
|
||||
first = await limiter.check_and_consume(
|
||||
user_rpm_key=user_key,
|
||||
user_rpm_limit=1,
|
||||
key_rpm_key=key_key,
|
||||
key_rpm_limit=0,
|
||||
)
|
||||
second = await limiter.check_and_consume(
|
||||
user_rpm_key=user_key,
|
||||
user_rpm_limit=1,
|
||||
key_rpm_key=key_key,
|
||||
key_rpm_limit=0,
|
||||
)
|
||||
|
||||
assert first == RpmCheckResult(allowed=True, remaining=0)
|
||||
assert second.allowed is False
|
||||
assert second.scope == "user"
|
||||
assert second.limit == 1
|
||||
assert second.remaining == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_and_consume_enforces_key_scope_in_memory(
|
||||
limiter: UserRpmLimiter,
|
||||
) -> None:
|
||||
user_key = limiter.get_user_rpm_key("user-1")
|
||||
key_key = limiter.get_key_rpm_key("key-1")
|
||||
|
||||
first = await limiter.check_and_consume(
|
||||
user_rpm_key=user_key,
|
||||
user_rpm_limit=3,
|
||||
key_rpm_key=key_key,
|
||||
key_rpm_limit=1,
|
||||
)
|
||||
second = await limiter.check_and_consume(
|
||||
user_rpm_key=user_key,
|
||||
user_rpm_limit=3,
|
||||
key_rpm_key=key_key,
|
||||
key_rpm_limit=1,
|
||||
)
|
||||
|
||||
assert first.allowed is True
|
||||
assert second.allowed is False
|
||||
assert second.scope == "key"
|
||||
assert second.limit == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_and_consume_redis_failure_falls_back_to_memory_when_fail_close(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
limiter: UserRpmLimiter,
|
||||
) -> None:
|
||||
fake_redis = MagicMock()
|
||||
fake_redis.eval = AsyncMock(side_effect=RuntimeError("redis down"))
|
||||
limiter._redis = fake_redis
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.services.rate_limit.user_rpm_limiter.config.rate_limit_fail_open", False
|
||||
)
|
||||
|
||||
user_key = limiter.get_user_rpm_key("user-1")
|
||||
key_key = limiter.get_key_rpm_key("key-1")
|
||||
|
||||
first = await limiter.check_and_consume(
|
||||
user_rpm_key=user_key,
|
||||
user_rpm_limit=1,
|
||||
key_rpm_key=key_key,
|
||||
key_rpm_limit=0,
|
||||
)
|
||||
second = await limiter.check_and_consume(
|
||||
user_rpm_key=user_key,
|
||||
user_rpm_limit=1,
|
||||
key_rpm_key=key_key,
|
||||
key_rpm_limit=0,
|
||||
)
|
||||
|
||||
assert first.allowed is True
|
||||
assert second.allowed is False
|
||||
assert second.scope == "user"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_and_consume_redis_failure_fail_open_allows_request(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
limiter: UserRpmLimiter,
|
||||
) -> None:
|
||||
fake_redis = MagicMock()
|
||||
fake_redis.eval = AsyncMock(side_effect=RuntimeError("redis down"))
|
||||
limiter._redis = fake_redis
|
||||
|
||||
monkeypatch.setattr(
|
||||
"src.services.rate_limit.user_rpm_limiter.config.rate_limit_fail_open", True
|
||||
)
|
||||
|
||||
result = await limiter.check_and_consume(
|
||||
user_rpm_key=limiter.get_user_rpm_key("user-1"),
|
||||
user_rpm_limit=1,
|
||||
key_rpm_key=limiter.get_key_rpm_key("key-1"),
|
||||
key_rpm_limit=0,
|
||||
)
|
||||
|
||||
assert result.allowed is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_and_consume_skips_when_all_limits_are_zero(
|
||||
limiter: UserRpmLimiter,
|
||||
) -> None:
|
||||
result = await limiter.check_and_consume(
|
||||
user_rpm_key=limiter.get_user_rpm_key("user-1"),
|
||||
user_rpm_limit=0,
|
||||
key_rpm_key=limiter.get_key_rpm_key("key-1"),
|
||||
key_rpm_limit=0,
|
||||
)
|
||||
|
||||
assert result.allowed is True
|
||||
assert result.scope is None
|
||||
assert limiter._memory_counts == {}
|
||||
@@ -53,3 +53,29 @@ def test_import_user_api_key_material_keeps_legacy_encrypted_payload() -> None:
|
||||
|
||||
assert key_hash == legacy_hash
|
||||
assert key_encrypted == legacy_encrypted
|
||||
|
||||
|
||||
def test_import_user_rate_limit_defaults_to_none_when_missing() -> None:
|
||||
assert AdminImportUsersAdapter._normalize_imported_user_rate_limit({}) is None
|
||||
|
||||
|
||||
def test_import_legacy_standalone_key_null_rate_limit_becomes_unlimited() -> None:
|
||||
assert (
|
||||
AdminImportUsersAdapter._normalize_imported_api_key_rate_limit(
|
||||
{"rate_limit": None},
|
||||
is_standalone=True,
|
||||
legacy_export=True,
|
||||
)
|
||||
== 0
|
||||
)
|
||||
|
||||
|
||||
def test_import_new_standalone_key_null_rate_limit_keeps_inherit_semantics() -> None:
|
||||
assert (
|
||||
AdminImportUsersAdapter._normalize_imported_api_key_rate_limit(
|
||||
{"rate_limit": None},
|
||||
is_standalone=True,
|
||||
legacy_export=False,
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user