mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +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:
@@ -22,7 +22,7 @@ from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db, get_db_context
|
||||
from src.models.api import CreateApiKeyRequest
|
||||
from src.models.database import ApiKey, Wallet
|
||||
from src.models.database import ApiKey, Usage, Wallet
|
||||
from src.services.user.apikey import ApiKeyService
|
||||
from src.services.user.bulk_cleanup import pre_clean_api_key
|
||||
from src.services.wallet import WalletService
|
||||
@@ -70,7 +70,9 @@ router = APIRouter(prefix="/api/admin/api-keys", tags=["Admin - API Keys (Standa
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
def _serialize_standalone_key_item(api_key: ApiKey) -> dict[str, Any]:
|
||||
def _serialize_standalone_key_item(
|
||||
api_key: ApiKey, *, total_tokens: int | None = None
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"id": api_key.id,
|
||||
"user_id": api_key.user_id,
|
||||
@@ -79,6 +81,7 @@ def _serialize_standalone_key_item(api_key: ApiKey) -> dict[str, Any]:
|
||||
"is_active": api_key.is_active,
|
||||
"is_standalone": api_key.is_standalone,
|
||||
"total_requests": api_key.total_requests,
|
||||
"total_tokens": int(total_tokens or 0),
|
||||
"total_cost_usd": float(api_key.total_cost_usd or 0),
|
||||
"rate_limit": api_key.rate_limit,
|
||||
"allowed_providers": api_key.allowed_providers,
|
||||
@@ -116,8 +119,27 @@ def _list_standalone_api_keys_sync(
|
||||
for api_key in api_keys:
|
||||
db.refresh(api_key)
|
||||
|
||||
token_map: dict[str, int] = {}
|
||||
if api_keys:
|
||||
stats_rows = (
|
||||
db.query(
|
||||
Usage.api_key_id,
|
||||
func.sum(Usage.total_tokens).label("total_tokens"),
|
||||
)
|
||||
.filter(Usage.api_key_id.in_([api_key.id for api_key in api_keys]))
|
||||
.group_by(Usage.api_key_id)
|
||||
.all()
|
||||
)
|
||||
token_map = {row.api_key_id: int(row.total_tokens or 0) for row in stats_rows}
|
||||
|
||||
return {
|
||||
"api_keys": [_serialize_standalone_key_item(api_key) for api_key in api_keys],
|
||||
"api_keys": [
|
||||
_serialize_standalone_key_item(
|
||||
api_key,
|
||||
total_tokens=token_map.get(api_key.id, 0),
|
||||
)
|
||||
for api_key in api_keys
|
||||
],
|
||||
"total": total,
|
||||
"limit": limit,
|
||||
"skip": skip,
|
||||
@@ -385,7 +407,7 @@ async def create_standalone_api_key(
|
||||
- `allowed_providers`: 可选,允许使用的提供商列表
|
||||
- `allowed_api_formats`: 可选,允许使用的 API 格式列表
|
||||
- `allowed_models`: 可选,允许使用的模型列表
|
||||
- `rate_limit`: 可选,速率限制配置(请求数/秒)
|
||||
- `rate_limit`: 可选,每分钟请求限制(null 表示跟随系统默认,0 表示不限制)
|
||||
- `expire_days`: 可选,过期天数(与 expires_at 二选一)
|
||||
- `expires_at`: 可选,过期时间(ISO 格式或 YYYY-MM-DD 格式,优先级高于 expire_days)
|
||||
- `auto_delete_on_expiry`: 可选,过期后是否自动删除
|
||||
@@ -421,7 +443,7 @@ async def update_api_key(
|
||||
**请求体字段**:
|
||||
- `name`: 可选,API Key 的名称
|
||||
- `unlimited_balance`: 可选,是否无限余额(true=无限,false=有限,不修改余额数值)
|
||||
- `rate_limit`: 可选,速率限制配置(null 表示无限制)
|
||||
- `rate_limit`: 可选,每分钟请求限制(null 表示跟随系统默认,0 表示不限制)
|
||||
- `allowed_providers`: 可选,允许使用的提供商列表
|
||||
- `allowed_api_formats`: 可选,允许使用的 API 格式列表
|
||||
- `allowed_models`: 可选,允许使用的模型列表
|
||||
@@ -683,6 +705,14 @@ class AdminGetKeyDetailAdapter(AdminApiAdapter):
|
||||
"is_active": api_key.is_active,
|
||||
"is_standalone": api_key.is_standalone,
|
||||
"total_requests": api_key.total_requests,
|
||||
"total_tokens": int(
|
||||
(
|
||||
db.query(func.sum(Usage.total_tokens))
|
||||
.filter(Usage.api_key_id == api_key.id)
|
||||
.scalar()
|
||||
)
|
||||
or 0
|
||||
),
|
||||
"total_cost_usd": float(api_key.total_cost_usd or 0),
|
||||
"rate_limit": api_key.rate_limit,
|
||||
"allowed_providers": api_key.allowed_providers,
|
||||
|
||||
@@ -2252,6 +2252,7 @@ class AdminExportUsersAdapter(AdminApiAdapter):
|
||||
"allowed_providers": user.allowed_providers,
|
||||
"allowed_api_formats": user.allowed_api_formats,
|
||||
"allowed_models": user.allowed_models,
|
||||
"rate_limit": user.rate_limit,
|
||||
"model_capability_settings": user.model_capability_settings,
|
||||
"unlimited": wallet_service.is_unlimited_wallet(wallet),
|
||||
"wallet": (wallet_service.serialize_wallet_summary(wallet) if wallet else None),
|
||||
@@ -2265,7 +2266,7 @@ class AdminExportUsersAdapter(AdminApiAdapter):
|
||||
standalone_keys_data = [self._serialize_api_key(key, db=db) for key in standalone_keys]
|
||||
|
||||
return {
|
||||
"version": "1.2",
|
||||
"version": "1.3",
|
||||
"exported_at": datetime.now(timezone.utc).isoformat(),
|
||||
"users": users_data,
|
||||
"standalone_keys": standalone_keys_data,
|
||||
@@ -2273,6 +2274,17 @@ class AdminExportUsersAdapter(AdminApiAdapter):
|
||||
|
||||
|
||||
class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
@staticmethod
|
||||
def _is_legacy_users_export(version: object) -> bool:
|
||||
if version is None:
|
||||
return True
|
||||
normalized = str(version).strip()
|
||||
try:
|
||||
parts = normalized.split(".")
|
||||
return (int(parts[0]), int(parts[1])) < (1, 3)
|
||||
except Exception:
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _resolve_api_key_material(key_data: dict[str, Any]) -> tuple[str | None, str | None]:
|
||||
"""解析用户 API Key 导入材料,优先使用明文 key。"""
|
||||
@@ -2289,6 +2301,30 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
key_encrypted = key_data.get("key_encrypted")
|
||||
return key_hash, key_encrypted
|
||||
|
||||
@staticmethod
|
||||
def _normalize_imported_user_rate_limit(user_data: dict[str, Any]) -> int | None:
|
||||
if "rate_limit" not in user_data:
|
||||
return None
|
||||
value = user_data.get("rate_limit")
|
||||
return int(value) if value is not None else None
|
||||
|
||||
@staticmethod
|
||||
def _normalize_imported_api_key_rate_limit(
|
||||
key_data: dict[str, Any],
|
||||
*,
|
||||
is_standalone: bool,
|
||||
legacy_export: bool,
|
||||
) -> int | None:
|
||||
if "rate_limit" not in key_data:
|
||||
return None if is_standalone and not legacy_export else 0
|
||||
|
||||
value = key_data.get("rate_limit")
|
||||
if value is None:
|
||||
if is_standalone and not legacy_export:
|
||||
return None
|
||||
return 0
|
||||
return int(value)
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
"""导入用户数据"""
|
||||
import uuid
|
||||
@@ -2306,6 +2342,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
|
||||
# 获取导入选项
|
||||
merge_mode = payload.get("merge_mode", "skip") # skip, overwrite, error
|
||||
legacy_export = self._is_legacy_users_export(payload.get("version"))
|
||||
users_data = payload.get("users", [])
|
||||
standalone_keys_data = payload.get("standalone_keys", [])
|
||||
|
||||
@@ -2358,7 +2395,11 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
allowed_providers=key_data.get("allowed_providers"),
|
||||
allowed_api_formats=key_data.get("allowed_api_formats"),
|
||||
allowed_models=key_data.get("allowed_models"),
|
||||
rate_limit=key_data.get("rate_limit"),
|
||||
rate_limit=self._normalize_imported_api_key_rate_limit(
|
||||
key_data,
|
||||
is_standalone=is_standalone or key_data.get("is_standalone", False),
|
||||
legacy_export=legacy_export,
|
||||
),
|
||||
concurrent_limit=key_data.get("concurrent_limit", 5),
|
||||
force_capabilities=key_data.get("force_capabilities"),
|
||||
is_active=key_data.get("is_active", True),
|
||||
@@ -2396,6 +2437,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
and wallet_payload.get("limit_mode") in {"finite", "unlimited"}
|
||||
else ("unlimited" if user_data.get("unlimited") else "finite")
|
||||
)
|
||||
imported_user_rate_limit = self._normalize_imported_user_rate_limit(user_data)
|
||||
|
||||
if existing_user:
|
||||
user_id = existing_user.id
|
||||
@@ -2413,6 +2455,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
existing_user.allowed_providers = user_data.get("allowed_providers")
|
||||
existing_user.allowed_api_formats = user_data.get("allowed_api_formats")
|
||||
existing_user.allowed_models = user_data.get("allowed_models")
|
||||
existing_user.rate_limit = imported_user_rate_limit
|
||||
existing_user.model_capability_settings = user_data.get(
|
||||
"model_capability_settings"
|
||||
)
|
||||
@@ -2449,6 +2492,7 @@ class AdminImportUsersAdapter(AdminApiAdapter):
|
||||
allowed_providers=user_data.get("allowed_providers"),
|
||||
allowed_api_formats=user_data.get("allowed_api_formats"),
|
||||
allowed_models=user_data.get("allowed_models"),
|
||||
rate_limit=imported_user_rate_limit,
|
||||
model_capability_settings=user_data.get("model_capability_settings"),
|
||||
is_active=user_data.get("is_active", True),
|
||||
)
|
||||
|
||||
@@ -18,7 +18,7 @@ from src.core.exceptions import InvalidRequestException, NotFoundException, tran
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db, get_db_context
|
||||
from src.models.admin_requests import UpdateUserRequest
|
||||
from src.models.api import CreateApiKeyRequest, CreateUserRequest
|
||||
from src.models.api import CreateApiKeyRequest, CreateUserRequest, UpdateMyApiKeyRequest
|
||||
from src.models.database import ApiKey, User, UserRole, Wallet
|
||||
from src.services.cache.user_cache import UserCacheService
|
||||
from src.services.system.config import SystemConfigService
|
||||
@@ -57,6 +57,7 @@ def _serialize_user(
|
||||
"allowed_providers": user.allowed_providers,
|
||||
"allowed_api_formats": user.allowed_api_formats,
|
||||
"allowed_models": user.allowed_models,
|
||||
"rate_limit": user.rate_limit,
|
||||
"unlimited": WalletService.is_unlimited_wallet(resolved_wallet),
|
||||
"is_active": user.is_active,
|
||||
"created_at": user.created_at.isoformat(),
|
||||
@@ -89,6 +90,7 @@ def _create_user_sync(
|
||||
allowed_providers=request.allowed_providers,
|
||||
allowed_api_formats=request.allowed_api_formats,
|
||||
allowed_models=request.allowed_models,
|
||||
rate_limit=request.rate_limit,
|
||||
)
|
||||
return _serialize_user(db, user), {
|
||||
"action": "create_user",
|
||||
@@ -255,6 +257,61 @@ def _delete_user_key_sync(user_id: str, key_id: str) -> tuple[dict[str, Any], di
|
||||
}
|
||||
|
||||
|
||||
def _update_user_key_sync(
|
||||
user_id: str,
|
||||
key_id: str,
|
||||
request: UpdateMyApiKeyRequest,
|
||||
) -> tuple[dict[str, Any], dict[str, Any]]:
|
||||
with get_db_context() as db:
|
||||
api_key = (
|
||||
db.query(ApiKey)
|
||||
.filter(
|
||||
ApiKey.id == key_id,
|
||||
ApiKey.user_id == user_id,
|
||||
ApiKey.is_standalone == False,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if not api_key:
|
||||
raise NotFoundException("API Key不存在或不属于该用户", "api_key")
|
||||
|
||||
update_data = request.model_dump(exclude_unset=True)
|
||||
if "rate_limit" in update_data and update_data["rate_limit"] is None:
|
||||
update_data["rate_limit"] = 0
|
||||
|
||||
updated_key = ApiKeyService.update_api_key(db, key_id, **update_data)
|
||||
if not updated_key:
|
||||
raise NotFoundException("API Key不存在或不属于该用户", "api_key")
|
||||
|
||||
return (
|
||||
{
|
||||
"id": updated_key.id,
|
||||
"name": updated_key.name,
|
||||
"key_display": updated_key.get_display_key(),
|
||||
"is_active": updated_key.is_active,
|
||||
"is_locked": updated_key.is_locked,
|
||||
"total_requests": updated_key.total_requests,
|
||||
"total_cost_usd": float(updated_key.total_cost_usd or 0),
|
||||
"rate_limit": updated_key.rate_limit,
|
||||
"expires_at": (
|
||||
updated_key.expires_at.isoformat() if updated_key.expires_at else None
|
||||
),
|
||||
"last_used_at": (
|
||||
updated_key.last_used_at.isoformat() if updated_key.last_used_at else None
|
||||
),
|
||||
"created_at": updated_key.created_at.isoformat(),
|
||||
"message": "API Key更新成功",
|
||||
},
|
||||
{
|
||||
"action": "update_user_api_key",
|
||||
"target_user_id": user_id,
|
||||
"key_id": key_id,
|
||||
"updated_fields": list(update_data.keys()),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _toggle_user_key_lock_sync(
|
||||
user_id: str,
|
||||
key_id: str,
|
||||
@@ -452,6 +509,26 @@ async def delete_user_api_key(
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.put("/{user_id}/api-keys/{key_id}")
|
||||
async def update_user_api_key(
|
||||
user_id: str,
|
||||
key_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""
|
||||
更新用户的 API 密钥
|
||||
|
||||
更新指定用户的普通 API 密钥基础配置。
|
||||
|
||||
**路径参数**:
|
||||
- `user_id`: 用户 ID (UUID)
|
||||
- `key_id`: 密钥 ID
|
||||
"""
|
||||
adapter = AdminUpdateUserKeyAdapter(user_id=user_id, key_id=key_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.patch("/{user_id}/api-keys/{key_id}/lock")
|
||||
async def toggle_user_api_key_lock(
|
||||
user_id: str,
|
||||
@@ -700,6 +777,33 @@ class AdminDeleteUserKeyAdapter(AdminApiAdapter):
|
||||
return response
|
||||
|
||||
|
||||
class AdminUpdateUserKeyAdapter(AdminApiAdapter):
|
||||
"""更新用户的普通 API Key"""
|
||||
|
||||
def __init__(self, user_id: str, key_id: str):
|
||||
self.user_id = user_id
|
||||
self.key_id = key_id
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
payload = context.ensure_json_body()
|
||||
try:
|
||||
request = UpdateMyApiKeyRequest.model_validate(payload)
|
||||
except ValidationError as e:
|
||||
errors = e.errors()
|
||||
if errors:
|
||||
raise InvalidRequestException(translate_pydantic_error(errors[0]))
|
||||
raise InvalidRequestException("请求数据验证失败")
|
||||
|
||||
response, audit_meta = await run_in_threadpool(
|
||||
_update_user_key_sync,
|
||||
self.user_id,
|
||||
self.key_id,
|
||||
request,
|
||||
)
|
||||
context.add_audit_metadata(**audit_meta)
|
||||
return response
|
||||
|
||||
|
||||
class AdminToggleUserKeyLockAdapter(AdminApiAdapter):
|
||||
"""切换用户普通 API Key 的锁定状态"""
|
||||
|
||||
|
||||
@@ -19,7 +19,9 @@ from src.core.logger import logger
|
||||
from src.database.database import create_session
|
||||
from src.models.database import ApiKey, AuditEventType, User
|
||||
from src.services.auth.service import AuthService
|
||||
from src.services.rate_limit.user_rpm_limiter import SYSTEM_RPM_CONFIG_KEY, get_user_rpm_limiter
|
||||
from src.services.system.audit import AuditService
|
||||
from src.services.system.config import SystemConfigService
|
||||
from src.services.usage.service import UsageService
|
||||
from src.services.wallet import WalletService
|
||||
from src.utils.perf import PerfRecorder
|
||||
@@ -123,6 +125,9 @@ class ApiRequestPipeline:
|
||||
auth_duration = PerfRecorder.stop(auth_start, "pipeline_auth", labels=perf_labels)
|
||||
_record_perf_metric("auth_ms", auth_duration)
|
||||
|
||||
if mode in {ApiMode.STANDARD, ApiMode.PROXY} and api_key and user:
|
||||
await self._check_user_rate_limit(http_request, db, user, api_key)
|
||||
|
||||
raw_body = None
|
||||
should_eager_read_body = http_request.method in {"POST", "PUT", "PATCH"} and getattr(
|
||||
adapter, "eager_request_body", True
|
||||
@@ -267,6 +272,55 @@ class ApiRequestPipeline:
|
||||
# Internal helpers
|
||||
# --------------------------------------------------------------------- #
|
||||
|
||||
async def _check_user_rate_limit(
|
||||
self,
|
||||
request: Request,
|
||||
db: Session,
|
||||
user: User,
|
||||
api_key: ApiKey,
|
||||
) -> None:
|
||||
limiter = await get_user_rpm_limiter()
|
||||
system_default_raw = SystemConfigService.get_config(db, SYSTEM_RPM_CONFIG_KEY, default=0)
|
||||
system_default = max(int(system_default_raw or 0), 0)
|
||||
|
||||
if api_key.is_standalone:
|
||||
effective_user_limit = (
|
||||
max(int(api_key.rate_limit or 0), 0)
|
||||
if api_key.rate_limit is not None
|
||||
else system_default
|
||||
)
|
||||
user_rpm_key = limiter.get_standalone_rpm_key(api_key.id)
|
||||
key_rpm_limit = 0
|
||||
else:
|
||||
effective_user_limit = (
|
||||
max(int(user.rate_limit or 0), 0) if user.rate_limit is not None else system_default
|
||||
)
|
||||
user_rpm_key = limiter.get_user_rpm_key(user.id)
|
||||
key_rpm_limit = max(int(api_key.rate_limit or 0), 0)
|
||||
|
||||
result = await limiter.check_and_consume(
|
||||
user_rpm_key=user_rpm_key,
|
||||
user_rpm_limit=effective_user_limit,
|
||||
key_rpm_key=limiter.get_key_rpm_key(api_key.id),
|
||||
key_rpm_limit=key_rpm_limit,
|
||||
)
|
||||
|
||||
if result.allowed:
|
||||
return
|
||||
|
||||
scope = result.scope or "user"
|
||||
limit = result.limit or (effective_user_limit if scope == "user" else key_rpm_limit)
|
||||
retry_after = result.retry_after or limiter.get_retry_after()
|
||||
|
||||
headers = {
|
||||
"Retry-After": str(retry_after),
|
||||
"X-RateLimit-Limit": str(limit),
|
||||
"X-RateLimit-Remaining": "0",
|
||||
"X-RateLimit-Scope": scope,
|
||||
}
|
||||
request.state.rate_limit_scope = scope
|
||||
raise HTTPException(status_code=429, detail="请求过于频繁,请稍后重试", headers=headers)
|
||||
|
||||
async def _authenticate_client(
|
||||
self, request: Request, db: Session, adapter: ApiAdapter, **_kw: object
|
||||
) -> tuple[User, ApiKey]:
|
||||
|
||||
@@ -16,7 +16,8 @@ from src.api.base.pipeline import get_pipeline
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import ApiKey, AuditLog
|
||||
from src.plugins.manager import get_plugin_manager
|
||||
from src.services.rate_limit.user_rpm_limiter import SYSTEM_RPM_CONFIG_KEY, get_user_rpm_limiter
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
router = APIRouter(prefix="/api/monitoring", tags=["Monitoring"])
|
||||
pipeline = get_pipeline()
|
||||
@@ -146,10 +147,6 @@ class UserRateLimitStatusAdapter(AuthenticatedApiAdapter):
|
||||
if not user:
|
||||
raise HTTPException(status_code=401, detail="未登录")
|
||||
|
||||
rate_limiter = _get_rate_limit_plugin()
|
||||
if not rate_limiter or not hasattr(rate_limiter, "get_rate_limit_headers"):
|
||||
raise HTTPException(status_code=503, detail="速率限制插件未启用或不支持状态查询")
|
||||
|
||||
api_keys = (
|
||||
db.query(ApiKey)
|
||||
.filter(ApiKey.user_id == user.id, ApiKey.is_active.is_(True))
|
||||
@@ -157,31 +154,84 @@ class UserRateLimitStatusAdapter(AuthenticatedApiAdapter):
|
||||
.all()
|
||||
)
|
||||
|
||||
try:
|
||||
limiter = await get_user_rpm_limiter()
|
||||
system_default_raw = SystemConfigService.get_config(
|
||||
db, SYSTEM_RPM_CONFIG_KEY, default=0
|
||||
)
|
||||
system_default = max(int(system_default_raw or 0), 0)
|
||||
reset_at = limiter.get_reset_at()
|
||||
window = f"{limiter.bucket_seconds}s"
|
||||
except Exception as exc:
|
||||
logger.warning("读取新 RPM 限流状态失败,回退插件状态接口: {}", exc)
|
||||
limiter = None
|
||||
system_default = 0
|
||||
reset_at = None
|
||||
window = None
|
||||
|
||||
rate_limit_info = []
|
||||
for key in api_keys:
|
||||
try:
|
||||
headers = rate_limiter.get_rate_limit_headers(key)
|
||||
except Exception as exc:
|
||||
logger.warning(f"无法获取Key {key.id} 的限流信息: {exc}")
|
||||
headers = {}
|
||||
if limiter is not None:
|
||||
if key.is_standalone:
|
||||
user_limit = key.rate_limit if key.rate_limit is not None else system_default
|
||||
user_scope_key = limiter.get_standalone_rpm_key(key.id)
|
||||
key_limit = 0
|
||||
else:
|
||||
user_limit = user.rate_limit if user.rate_limit is not None else system_default
|
||||
user_scope_key = limiter.get_user_rpm_key(user.id)
|
||||
key_limit = max(int(key.rate_limit or 0), 0)
|
||||
|
||||
user_count = (
|
||||
await limiter.get_scope_count(user_scope_key)
|
||||
if user_limit and user_limit > 0
|
||||
else 0
|
||||
)
|
||||
key_count = (
|
||||
await limiter.get_scope_count(limiter.get_key_rpm_key(key.id))
|
||||
if key_limit > 0
|
||||
else 0
|
||||
)
|
||||
user_remaining = max(user_limit - user_count, 0) if user_limit > 0 else None
|
||||
key_remaining = max(key_limit - key_count, 0) if key_limit > 0 else None
|
||||
|
||||
scoped_statuses: list[tuple[str, int, int]] = []
|
||||
if user_limit > 0 and user_remaining is not None:
|
||||
scoped_statuses.append(("user", user_limit, user_remaining))
|
||||
if key_limit > 0 and key_remaining is not None:
|
||||
scoped_statuses.append(("key", key_limit, key_remaining))
|
||||
|
||||
primary_scope = (
|
||||
min(scoped_statuses, key=lambda item: item[2]) if scoped_statuses else None
|
||||
)
|
||||
rate_limit_info.append(
|
||||
{
|
||||
"api_key_name": key.name or f"Key-{key.id}",
|
||||
"limit": primary_scope[1] if primary_scope else None,
|
||||
"remaining": primary_scope[2] if primary_scope else None,
|
||||
"scope": primary_scope[0] if primary_scope else None,
|
||||
"reset_time": reset_at.isoformat() if reset_at else None,
|
||||
"window": window,
|
||||
"user_limit": user_limit,
|
||||
"user_remaining": user_remaining,
|
||||
"key_limit": key_limit if key_limit > 0 else None,
|
||||
"key_remaining": key_remaining,
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
rate_limit_info.append(
|
||||
{
|
||||
"api_key_name": key.name or f"Key-{key.id}",
|
||||
"limit": headers.get("X-RateLimit-Limit"),
|
||||
"remaining": headers.get("X-RateLimit-Remaining"),
|
||||
"reset_time": headers.get("X-RateLimit-Reset"),
|
||||
"window": headers.get("X-RateLimit-Window"),
|
||||
"limit": None,
|
||||
"remaining": None,
|
||||
"scope": None,
|
||||
"reset_time": None,
|
||||
"window": None,
|
||||
"user_limit": None,
|
||||
"user_remaining": None,
|
||||
"key_limit": None,
|
||||
"key_remaining": None,
|
||||
}
|
||||
)
|
||||
|
||||
return {"user_id": user.id, "api_keys": rate_limit_info}
|
||||
|
||||
|
||||
def _get_rate_limit_plugin() -> Any:
|
||||
try:
|
||||
plugin_manager = get_plugin_manager()
|
||||
return plugin_manager.get_plugin("rate_limit")
|
||||
except Exception as exc:
|
||||
logger.warning(f"获取速率限制插件失败: {exc}")
|
||||
return None
|
||||
|
||||
@@ -33,6 +33,7 @@ from src.models.api import (
|
||||
PublicGlobalModelListResponse,
|
||||
PublicGlobalModelResponse,
|
||||
UpdateApiKeyProvidersRequest,
|
||||
UpdateMyApiKeyRequest,
|
||||
UpdatePreferencesRequest,
|
||||
UpdateProfileRequest,
|
||||
)
|
||||
@@ -146,6 +147,7 @@ def _create_my_api_key_sync(user_id: str, request: CreateMyApiKeyRequest) -> dic
|
||||
db=db,
|
||||
user_id=user_id,
|
||||
name=request.name,
|
||||
rate_limit=request.rate_limit,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise InvalidRequestException(str(exc)) from exc
|
||||
@@ -154,6 +156,7 @@ def _create_my_api_key_sync(user_id: str, request: CreateMyApiKeyRequest) -> dic
|
||||
"name": api_key.name,
|
||||
"key": plain_key,
|
||||
"key_display": api_key.get_display_key(),
|
||||
"rate_limit": api_key.rate_limit,
|
||||
"message": "API密钥创建成功",
|
||||
}
|
||||
|
||||
@@ -189,6 +192,50 @@ def _toggle_my_api_key_sync(user_id: str, key_id: str) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _update_my_api_key_sync(
|
||||
user_id: str,
|
||||
key_id: str,
|
||||
request: UpdateMyApiKeyRequest,
|
||||
) -> dict[str, Any]:
|
||||
with get_db_context() as db:
|
||||
api_key = (
|
||||
db.query(ApiKey)
|
||||
.filter(
|
||||
ApiKey.id == key_id,
|
||||
ApiKey.user_id == user_id,
|
||||
ApiKey.is_standalone == False,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if not api_key:
|
||||
raise NotFoundException("API密钥不存在", "api_key")
|
||||
if api_key.is_locked:
|
||||
raise ForbiddenException("该密钥已被管理员锁定,无法修改")
|
||||
|
||||
update_data = request.model_dump(exclude_unset=True)
|
||||
if "rate_limit" in update_data and update_data["rate_limit"] is None:
|
||||
update_data["rate_limit"] = 0
|
||||
|
||||
updated = ApiKeyService.update_api_key(db, key_id, **update_data)
|
||||
if not updated:
|
||||
raise NotFoundException("API密钥不存在", "api_key")
|
||||
|
||||
return {
|
||||
"id": updated.id,
|
||||
"name": updated.name,
|
||||
"key_display": updated.get_display_key(),
|
||||
"is_active": updated.is_active,
|
||||
"is_locked": updated.is_locked,
|
||||
"allowed_providers": updated.allowed_providers,
|
||||
"force_capabilities": updated.force_capabilities,
|
||||
"rate_limit": updated.rate_limit,
|
||||
"last_used_at": updated.last_used_at.isoformat() if updated.last_used_at else None,
|
||||
"expires_at": updated.expires_at.isoformat() if updated.expires_at else None,
|
||||
"created_at": updated.created_at.isoformat(),
|
||||
"message": "API密钥已更新",
|
||||
}
|
||||
|
||||
|
||||
def _update_api_key_providers_sync(
|
||||
user_id: str,
|
||||
api_key_id: str,
|
||||
@@ -484,6 +531,20 @@ async def delete_my_api_key(key_id: str, request: Request, db: Session = Depends
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.put("/api-keys/{key_id}")
|
||||
async def update_my_api_key(key_id: str, request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
"""
|
||||
更新 API 密钥
|
||||
|
||||
更新指定 API 密钥的基础配置。
|
||||
|
||||
**路径参数**:
|
||||
- `key_id`: 密钥 ID
|
||||
"""
|
||||
adapter = UpdateMyApiKeyAdapter(key_id=key_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.patch("/api-keys/{key_id}")
|
||||
async def toggle_my_api_key(key_id: str, request: Request, db: Session = Depends(get_db)) -> Any:
|
||||
"""
|
||||
@@ -877,6 +938,7 @@ class ListMyApiKeysAdapter(AuthenticatedApiAdapter):
|
||||
"created_at": key.created_at.isoformat(),
|
||||
"total_requests": real_stats["total_requests"],
|
||||
"total_cost_usd": real_stats["total_cost_usd"],
|
||||
"rate_limit": key.rate_limit,
|
||||
"allowed_providers": key.allowed_providers,
|
||||
"force_capabilities": key.force_capabilities,
|
||||
}
|
||||
@@ -965,6 +1027,27 @@ class GetMyApiKeyDetailAdapter(AuthenticatedApiAdapter):
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpdateMyApiKeyAdapter(AuthenticatedApiAdapter):
|
||||
"""更新 API 密钥基础配置的适配器"""
|
||||
|
||||
key_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
payload = context.ensure_json_body()
|
||||
try:
|
||||
request = UpdateMyApiKeyRequest.model_validate(payload)
|
||||
except ValidationError as e:
|
||||
errors = e.errors()
|
||||
if errors:
|
||||
raise InvalidRequestException(translate_pydantic_error(errors[0]))
|
||||
raise InvalidRequestException("请求数据验证失败")
|
||||
|
||||
return await run_in_threadpool(
|
||||
_update_my_api_key_sync, context.user.id, self.key_id, request
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DeleteMyApiKeyAdapter(AuthenticatedApiAdapter):
|
||||
"""删除 API 密钥的适配器"""
|
||||
|
||||
@@ -109,8 +109,6 @@ class Config:
|
||||
# 支付回调安全配置(公开回调入口必须携带该共享密钥)
|
||||
self.payment_callback_secret = os.getenv("PAYMENT_CALLBACK_SECRET", "").strip()
|
||||
|
||||
# LLM API 速率限制配置(每分钟请求数)
|
||||
self.llm_api_rate_limit = int(os.getenv("LLM_API_RATE_LIMIT", "100"))
|
||||
self.public_api_rate_limit = int(os.getenv("PUBLIC_API_RATE_LIMIT", "60"))
|
||||
|
||||
# 异常处理配置
|
||||
|
||||
11
src/main.py
11
src/main.py
@@ -45,6 +45,7 @@ if TYPE_CHECKING:
|
||||
from src.services.model.fetch_scheduler import ModelFetchScheduler
|
||||
from src.services.provider_keys.pool_quota_probe_scheduler import PoolQuotaProbeScheduler
|
||||
from src.services.rate_limit.concurrency_manager import ConcurrencyManager
|
||||
from src.services.rate_limit.user_rpm_limiter import UserRpmLimiter
|
||||
from src.services.system.maintenance_scheduler import MaintenanceScheduler
|
||||
from src.services.system.scheduler import TaskScheduler
|
||||
from src.services.task.polling.task_poller import TaskPollerService
|
||||
@@ -99,6 +100,7 @@ class LifecycleState:
|
||||
|
||||
redis_client: Redis | None = None
|
||||
concurrency_manager: ConcurrencyManager | None = None
|
||||
user_rpm_limiter: UserRpmLimiter | None = None
|
||||
plugin_manager: PluginManager | None = None
|
||||
available_modules: list[ModuleDefinition] = field(default_factory=list)
|
||||
task_coordinator: StartupTaskCoordinator | None = None
|
||||
@@ -181,6 +183,11 @@ async def _initialize_core_infrastructure(state: LifecycleState) -> None:
|
||||
|
||||
state.concurrency_manager = await get_concurrency_manager()
|
||||
|
||||
logger.info("初始化用户/API Key RPM 限流器...")
|
||||
from src.services.rate_limit.user_rpm_limiter import get_user_rpm_limiter
|
||||
|
||||
state.user_rpm_limiter = await get_user_rpm_limiter()
|
||||
|
||||
# 初始化批量提交器(提升数据库并发能力)
|
||||
logger.info("初始化批量提交器...")
|
||||
from src.core.batch_committer import init_batch_committer
|
||||
@@ -541,6 +548,10 @@ async def _run_shutdown(state: LifecycleState) -> None:
|
||||
if state.concurrency_manager:
|
||||
await state.concurrency_manager.close()
|
||||
|
||||
logger.info("关闭用户/API Key RPM 限流器...")
|
||||
if state.user_rpm_limiter:
|
||||
await state.user_rpm_limiter.close()
|
||||
|
||||
# 关闭全局Redis客户端
|
||||
logger.info("关闭全局Redis客户端...")
|
||||
from src.clients.redis_client import close_redis_client
|
||||
|
||||
@@ -8,7 +8,6 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
@@ -50,7 +49,6 @@ class PluginMiddleware:
|
||||
self._notification_cache_expires: float = 0.0
|
||||
|
||||
# 从配置读取速率限制值
|
||||
self.llm_api_rate_limit = config.llm_api_rate_limit
|
||||
self.public_api_rate_limit = config.public_api_rate_limit
|
||||
|
||||
# 完全跳过限流的路径(静态资源、文档等)
|
||||
@@ -69,14 +67,6 @@ class PluginMiddleware:
|
||||
"/api/monitoring/", # 监控端点
|
||||
]
|
||||
|
||||
# LLM API 端点(需要特殊的速率限制策略)
|
||||
self.llm_api_paths = [
|
||||
"/v1/messages",
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses",
|
||||
"/v1/completions",
|
||||
]
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
"""ASGI 入口点"""
|
||||
if scope["type"] != "http":
|
||||
@@ -291,13 +281,6 @@ class PluginMiddleware:
|
||||
|
||||
return "unknown"
|
||||
|
||||
def _is_llm_api_path(self, path: str) -> bool:
|
||||
"""检查是否为 LLM API 端点"""
|
||||
for llm_path in self.llm_api_paths:
|
||||
if path.startswith(llm_path):
|
||||
return True
|
||||
return False
|
||||
|
||||
async def _get_rate_limit_key_and_config(
|
||||
self, request: Request
|
||||
) -> tuple[str | None, int | None]:
|
||||
@@ -305,7 +288,6 @@ class PluginMiddleware:
|
||||
获取速率限制的key和配置
|
||||
|
||||
策略说明:
|
||||
- /v1/messages, /v1/chat/completions 等 LLM API: 按 API Key 限流
|
||||
- /api/public/* 端点: 使用服务器级别 IP 限制
|
||||
- /api/admin/* 端点: 跳过(在 skip_rate_limit_paths 中跳过)
|
||||
- /api/auth/* 端点: 跳过(由路由层的 IPRateLimiter 处理)
|
||||
@@ -315,30 +297,6 @@ class PluginMiddleware:
|
||||
"""
|
||||
path = request.url.path
|
||||
|
||||
# LLM API 端点: 按 API Key 或 IP 限流
|
||||
if self._is_llm_api_path(path):
|
||||
# 尝试从请求头获取 API Key
|
||||
auth_header = request.headers.get("authorization", "")
|
||||
api_key = request.headers.get("x-api-key", "")
|
||||
|
||||
if auth_header.lower().startswith("bearer "):
|
||||
api_key = auth_header[7:]
|
||||
|
||||
if api_key:
|
||||
# 使用 API Key 的哈希作为限制 key(避免日志泄露完整 key)
|
||||
key_hash = hashlib.sha256(api_key.encode()).hexdigest()[:16]
|
||||
key = f"llm_api_key:{key_hash}"
|
||||
request.state.rate_limit_key_type = "api_key"
|
||||
else:
|
||||
# 无 API Key 时使用 IP 限制(更严格)
|
||||
client_ip = self._get_client_ip(request)
|
||||
key = f"llm_ip:{client_ip}"
|
||||
request.state.rate_limit_key_type = "ip"
|
||||
|
||||
rate_limit = self.llm_api_rate_limit
|
||||
request.state.rate_limit_value = rate_limit
|
||||
return key, rate_limit
|
||||
|
||||
# /api/public/* 端点: 使用服务器级别 IP 地址作为限制 key
|
||||
if path.startswith("/api/public/"):
|
||||
client_ip = self._get_client_ip(request)
|
||||
|
||||
@@ -692,6 +692,11 @@ class UpdateUserRequest(BaseModel):
|
||||
allowed_providers: list[str] | None = Field(None, description="允许使用的提供商 ID 列表")
|
||||
allowed_api_formats: list[str] | None = Field(None, description="允许使用的 API 格式列表")
|
||||
allowed_models: list[str] | None = Field(None, description="允许使用的模型名称列表")
|
||||
rate_limit: int | None = Field(
|
||||
None,
|
||||
ge=0,
|
||||
description="每分钟请求限制;null 表示继承系统默认,0 表示不限制",
|
||||
)
|
||||
|
||||
@field_validator("username")
|
||||
@classmethod
|
||||
|
||||
@@ -252,6 +252,11 @@ class CreateUserRequest(BaseModel):
|
||||
allowed_models: list[str] | None = Field(
|
||||
default=None, description="允许使用的模型名称列表,null表示无限制"
|
||||
)
|
||||
rate_limit: int | None = Field(
|
||||
default=None,
|
||||
ge=0,
|
||||
description="每分钟请求限制;null 表示继承系统默认,0 表示不限制",
|
||||
)
|
||||
|
||||
@field_validator("initial_gift_usd", mode="before")
|
||||
@classmethod
|
||||
@@ -335,6 +340,11 @@ class UpdateUserRequest(BaseModel):
|
||||
allowed_providers: list[str] | None = None # 允许使用的提供商 ID 列表
|
||||
allowed_api_formats: list[str] | None = None # 允许使用的 API 格式列表
|
||||
allowed_models: list[str] | None = None # 允许使用的模型名称列表
|
||||
rate_limit: int | None = Field(
|
||||
default=None,
|
||||
ge=0,
|
||||
description="每分钟请求限制;null 表示继承系统默认,0 表示不限制",
|
||||
)
|
||||
is_active: bool | None = None
|
||||
|
||||
@field_validator("allowed_api_formats")
|
||||
@@ -351,7 +361,11 @@ class CreateApiKeyRequest(BaseModel):
|
||||
allowed_providers: list[str] | None = None # 允许使用的提供商 ID 列表
|
||||
allowed_api_formats: list[str] | None = None # 允许使用的 API 格式列表
|
||||
allowed_models: list[str] | None = None # 允许使用的模型名称列表
|
||||
rate_limit: int | None = None # None = 无限制
|
||||
rate_limit: int | None = Field(
|
||||
None,
|
||||
ge=0,
|
||||
description="每分钟请求限制;独立Key: null=继承系统默认,0=不限制;普通Key: 0=不限制",
|
||||
)
|
||||
expire_days: int | None = None # None = 永不过期,数字 = 多少天后过期
|
||||
expires_at: str | None = None # ISO 日期字符串,如 "2025-12-31",优先于 expire_days
|
||||
initial_balance_usd: float | None = Field(
|
||||
@@ -382,6 +396,7 @@ class UserResponse(BaseModel):
|
||||
allowed_providers: list[str] | None = None # 允许使用的提供商 ID 列表
|
||||
allowed_api_formats: list[str] | None = None # 允许使用的 API 格式列表
|
||||
allowed_models: list[str] | None = None # 允许使用的模型名称列表
|
||||
rate_limit: int | None = None
|
||||
unlimited: bool = False
|
||||
is_active: bool
|
||||
created_at: datetime
|
||||
@@ -402,7 +417,7 @@ class ApiKeyResponse(BaseModel):
|
||||
total_cost_usd: float
|
||||
allowed_providers: list[str] | None
|
||||
allowed_models: list[str] | None
|
||||
rate_limit: int
|
||||
rate_limit: int | None
|
||||
is_active: bool
|
||||
expires_at: datetime | None = None
|
||||
is_standalone: bool = False
|
||||
@@ -772,6 +787,18 @@ class CreateMyApiKeyRequest(BaseModel):
|
||||
"""创建我的API密钥请求"""
|
||||
|
||||
name: str
|
||||
rate_limit: int = Field(0, ge=0, description="该 Key 的每分钟请求限制,0 表示不限制")
|
||||
|
||||
|
||||
class UpdateMyApiKeyRequest(BaseModel):
|
||||
"""更新我的 API 密钥请求"""
|
||||
|
||||
name: str | None = None
|
||||
rate_limit: int | None = Field(
|
||||
None,
|
||||
ge=0,
|
||||
description="该 Key 的每分钟请求限制;0 表示不限制,null 表示不修改",
|
||||
)
|
||||
|
||||
|
||||
class ProviderConfig(BaseModel):
|
||||
|
||||
@@ -110,6 +110,9 @@ class User(Base):
|
||||
allowed_providers = Column(JSON, nullable=True) # 允许使用的提供商 ID 列表
|
||||
allowed_api_formats = Column(JSON, nullable=True) # 允许使用的 API 格式列表
|
||||
allowed_models = Column(JSON, nullable=True) # 允许使用的模型名称列表
|
||||
rate_limit = Column(
|
||||
Integer, nullable=True, default=None
|
||||
) # 每分钟请求限制,NULL=继承系统默认,0=不限制,N=N RPM
|
||||
|
||||
# Key 能力配置
|
||||
model_capability_settings = Column(JSON, nullable=True) # 用户针对特定模型的能力配置
|
||||
@@ -209,7 +212,9 @@ class ApiKey(Base):
|
||||
allowed_providers = Column(JSON, nullable=True) # 允许使用的提供商 ID 列表
|
||||
allowed_api_formats = Column(JSON, nullable=True) # 允许使用的 API 格式列表
|
||||
allowed_models = Column(JSON, nullable=True) # 允许使用的模型名称列表
|
||||
rate_limit = Column(Integer, default=None, nullable=True) # 每分钟请求限制,None = 无限制
|
||||
rate_limit = Column(
|
||||
Integer, default=None, nullable=True
|
||||
) # 每分钟请求限制;独立Key: NULL=继承系统默认,普通Key: 0=不限制
|
||||
concurrent_limit = Column(Integer, default=5, nullable=True) # 并发请求限制
|
||||
|
||||
# Key 能力配置
|
||||
|
||||
@@ -12,6 +12,7 @@ from src.services.rate_limit.adaptive_rpm import (
|
||||
from src.services.rate_limit.concurrency_manager import ConcurrencyManager
|
||||
from src.services.rate_limit.detector import RateLimitDetector
|
||||
from src.services.rate_limit.ip_limiter import IPRateLimiter
|
||||
from src.services.rate_limit.user_rpm_limiter import UserRpmLimiter, get_user_rpm_limiter
|
||||
|
||||
__all__ = [
|
||||
"AdaptiveConcurrencyManager", # 向后兼容
|
||||
@@ -19,5 +20,7 @@ __all__ = [
|
||||
"ConcurrencyManager",
|
||||
"IPRateLimiter",
|
||||
"RateLimitDetector",
|
||||
"UserRpmLimiter",
|
||||
"get_adaptive_rpm_manager",
|
||||
"get_user_rpm_limiter",
|
||||
]
|
||||
|
||||
354
src/services/rate_limit/user_rpm_limiter.py
Normal file
354
src/services/rate_limit/user_rpm_limiter.py
Normal file
@@ -0,0 +1,354 @@
|
||||
"""
|
||||
用户/API Key RPM 限制器
|
||||
|
||||
支持两层叠加限流:
|
||||
1. 用户级(或独立 Key 级)总 RPM
|
||||
2. 普通 Key 子限制 RPM
|
||||
|
||||
实现策略:
|
||||
- Redis 可用时使用分钟桶 + Lua 脚本原子检查/消费
|
||||
- Redis 不可用时降级为内存计数(仅适用于单实例)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import redis.asyncio as aioredis
|
||||
|
||||
from src.config.settings import config
|
||||
from src.core.logger import logger
|
||||
|
||||
SYSTEM_RPM_CONFIG_KEY = "rate_limit_per_minute"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class RpmCheckResult:
|
||||
"""RPM 检查结果。"""
|
||||
|
||||
allowed: bool
|
||||
scope: str | None = None
|
||||
limit: int | None = None
|
||||
remaining: int | None = None
|
||||
retry_after: int | None = None
|
||||
|
||||
|
||||
class UserRpmLimiter:
|
||||
"""用户/API Key 双层 RPM 限制器。
|
||||
|
||||
通过模块级 ``get_user_rpm_limiter()`` 工厂函数获取唯一实例,
|
||||
不要直接调用构造函数。
|
||||
"""
|
||||
|
||||
_CHECK_AND_CONSUME_SCRIPT = """
|
||||
local user_key = KEYS[1]
|
||||
local key_key = KEYS[2]
|
||||
local user_limit = tonumber(ARGV[1])
|
||||
local key_limit = tonumber(ARGV[2])
|
||||
local ttl = tonumber(ARGV[3])
|
||||
local retry_after = tonumber(ARGV[4])
|
||||
|
||||
local user_count = 0
|
||||
if user_limit > 0 then
|
||||
user_count = tonumber(redis.call('GET', user_key) or '0')
|
||||
if user_count >= user_limit then
|
||||
return {0, 1, user_limit, 0, retry_after}
|
||||
end
|
||||
end
|
||||
|
||||
local key_count = 0
|
||||
if key_limit > 0 then
|
||||
key_count = tonumber(redis.call('GET', key_key) or '0')
|
||||
if key_count >= key_limit then
|
||||
return {0, 2, key_limit, 0, retry_after}
|
||||
end
|
||||
end
|
||||
|
||||
local remaining = -1
|
||||
if user_limit > 0 then
|
||||
user_count = redis.call('INCR', user_key)
|
||||
redis.call('EXPIRE', user_key, ttl)
|
||||
remaining = user_limit - user_count
|
||||
end
|
||||
|
||||
if key_limit > 0 then
|
||||
key_count = redis.call('INCR', key_key)
|
||||
redis.call('EXPIRE', key_key, ttl)
|
||||
local key_remaining = key_limit - key_count
|
||||
if remaining == -1 or key_remaining < remaining then
|
||||
remaining = key_remaining
|
||||
end
|
||||
end
|
||||
|
||||
return {1, 0, 0, remaining, 0}
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._redis: aioredis.Redis | None = None
|
||||
self._bucket_seconds = int(config.rpm_bucket_seconds)
|
||||
self._key_ttl_seconds = int(config.rpm_key_ttl_seconds)
|
||||
self._cleanup_interval_seconds = int(config.rpm_cleanup_interval_seconds)
|
||||
self._memory_lock: asyncio.Lock = asyncio.Lock()
|
||||
self._memory_counts: dict[str, tuple[int, int]] = {}
|
||||
self._cleanup_task: asyncio.Task | None = None
|
||||
|
||||
async def initialize(self) -> None:
|
||||
if self._redis is not None:
|
||||
return
|
||||
|
||||
try:
|
||||
from src.clients.redis_client import get_redis_client
|
||||
|
||||
self._redis = await get_redis_client(require_redis=False)
|
||||
if self._redis:
|
||||
logger.info("[OK] UserRpmLimiter 已复用全局 Redis 客户端")
|
||||
return
|
||||
except Exception as exc:
|
||||
logger.warning("初始化 UserRpmLimiter Redis 客户端失败,降级为内存模式: {}", exc)
|
||||
|
||||
self._redis = None
|
||||
self._start_background_cleanup()
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._cleanup_task is not None:
|
||||
self._cleanup_task.cancel()
|
||||
try:
|
||||
await self._cleanup_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._cleanup_task = None
|
||||
|
||||
@property
|
||||
def bucket_seconds(self) -> int:
|
||||
return self._bucket_seconds
|
||||
|
||||
def get_user_rpm_key(self, user_id: str, bucket: int | None = None) -> str:
|
||||
b = bucket if bucket is not None else self._get_rpm_bucket()
|
||||
return f"rpm:user:{user_id}:{b}"
|
||||
|
||||
def get_standalone_rpm_key(self, api_key_id: str, bucket: int | None = None) -> str:
|
||||
b = bucket if bucket is not None else self._get_rpm_bucket()
|
||||
return f"rpm:ukey:{api_key_id}:{b}"
|
||||
|
||||
def get_key_rpm_key(self, api_key_id: str, bucket: int | None = None) -> str:
|
||||
b = bucket if bucket is not None else self._get_rpm_bucket()
|
||||
return f"rpm:key:{api_key_id}:{b}"
|
||||
|
||||
def get_retry_after(self, now_ts: float | None = None) -> int:
|
||||
ts = now_ts if now_ts is not None else time.time()
|
||||
elapsed = int(ts % self._bucket_seconds)
|
||||
return max(1, self._bucket_seconds - elapsed)
|
||||
|
||||
def get_reset_at(self, now_ts: float | None = None) -> datetime:
|
||||
ts = now_ts if now_ts is not None else time.time()
|
||||
bucket = self._get_rpm_bucket(ts)
|
||||
reset_ts = (bucket + 1) * self._bucket_seconds
|
||||
return datetime.fromtimestamp(reset_ts, tz=timezone.utc)
|
||||
|
||||
async def get_scope_count(self, scope_key: str) -> int:
|
||||
await self.initialize()
|
||||
|
||||
if self._redis is None:
|
||||
async with self._memory_lock:
|
||||
bucket = self._get_rpm_bucket()
|
||||
self._cleanup_expired_memory_counts(bucket)
|
||||
return self._get_memory_count(scope_key)
|
||||
|
||||
try:
|
||||
result = await self._redis.get(scope_key)
|
||||
return int(result) if result else 0
|
||||
except Exception as exc:
|
||||
logger.warning("读取 RPM 计数失败,回退内存模式: {}", exc)
|
||||
if config.rate_limit_fail_open:
|
||||
return 0
|
||||
async with self._memory_lock:
|
||||
bucket = self._get_rpm_bucket()
|
||||
self._cleanup_expired_memory_counts(bucket)
|
||||
return self._get_memory_count(scope_key)
|
||||
|
||||
async def check_and_consume(
|
||||
self,
|
||||
*,
|
||||
user_rpm_key: str,
|
||||
user_rpm_limit: int,
|
||||
key_rpm_key: str,
|
||||
key_rpm_limit: int,
|
||||
) -> RpmCheckResult:
|
||||
"""原子检查并消费两层 RPM 配额。"""
|
||||
|
||||
await self.initialize()
|
||||
|
||||
normalized_user_limit = max(int(user_rpm_limit or 0), 0)
|
||||
normalized_key_limit = max(int(key_rpm_limit or 0), 0)
|
||||
|
||||
if normalized_user_limit <= 0 and normalized_key_limit <= 0:
|
||||
return RpmCheckResult(allowed=True)
|
||||
|
||||
if self._redis is None:
|
||||
return await self._check_and_consume_memory(
|
||||
user_rpm_key=user_rpm_key,
|
||||
user_rpm_limit=normalized_user_limit,
|
||||
key_rpm_key=key_rpm_key,
|
||||
key_rpm_limit=normalized_key_limit,
|
||||
)
|
||||
|
||||
retry_after = self.get_retry_after()
|
||||
|
||||
try:
|
||||
raw_result = await self._redis.eval(
|
||||
self._CHECK_AND_CONSUME_SCRIPT,
|
||||
2,
|
||||
user_rpm_key,
|
||||
key_rpm_key,
|
||||
normalized_user_limit,
|
||||
normalized_key_limit,
|
||||
self._key_ttl_seconds,
|
||||
retry_after,
|
||||
)
|
||||
return self._parse_redis_result(raw_result)
|
||||
except Exception as exc:
|
||||
logger.warning("Redis RPM 检查失败: {}", exc)
|
||||
if config.rate_limit_fail_open:
|
||||
return RpmCheckResult(allowed=True)
|
||||
return await self._check_and_consume_memory(
|
||||
user_rpm_key=user_rpm_key,
|
||||
user_rpm_limit=normalized_user_limit,
|
||||
key_rpm_key=key_rpm_key,
|
||||
key_rpm_limit=normalized_key_limit,
|
||||
)
|
||||
|
||||
def _parse_redis_result(self, raw_result: object) -> RpmCheckResult:
|
||||
values = list(raw_result) if isinstance(raw_result, (list, tuple)) else [raw_result]
|
||||
allowed = int(values[0]) == 1
|
||||
scope_code = int(values[1]) if len(values) > 1 else 0
|
||||
limit = int(values[2]) if len(values) > 2 and values[2] is not None else None
|
||||
remaining = int(values[3]) if len(values) > 3 and values[3] is not None else None
|
||||
retry_after = int(values[4]) if len(values) > 4 and values[4] is not None else None
|
||||
scope = {1: "user", 2: "key"}.get(scope_code)
|
||||
return RpmCheckResult(
|
||||
allowed=allowed,
|
||||
scope=scope,
|
||||
limit=limit,
|
||||
remaining=remaining,
|
||||
retry_after=retry_after,
|
||||
)
|
||||
|
||||
async def _check_and_consume_memory(
|
||||
self,
|
||||
*,
|
||||
user_rpm_key: str,
|
||||
user_rpm_limit: int,
|
||||
key_rpm_key: str,
|
||||
key_rpm_limit: int,
|
||||
) -> RpmCheckResult:
|
||||
async with self._memory_lock:
|
||||
bucket = self._get_rpm_bucket()
|
||||
self._cleanup_expired_memory_counts(bucket)
|
||||
|
||||
user_count = self._get_memory_count(user_rpm_key)
|
||||
if user_rpm_limit > 0 and user_count >= user_rpm_limit:
|
||||
return RpmCheckResult(
|
||||
allowed=False,
|
||||
scope="user",
|
||||
limit=user_rpm_limit,
|
||||
remaining=0,
|
||||
retry_after=self.get_retry_after(),
|
||||
)
|
||||
|
||||
key_count = self._get_memory_count(key_rpm_key)
|
||||
if key_rpm_limit > 0 and key_count >= key_rpm_limit:
|
||||
return RpmCheckResult(
|
||||
allowed=False,
|
||||
scope="key",
|
||||
limit=key_rpm_limit,
|
||||
remaining=0,
|
||||
retry_after=self.get_retry_after(),
|
||||
)
|
||||
|
||||
remaining_candidates: list[int] = []
|
||||
|
||||
if user_rpm_limit > 0:
|
||||
user_count += 1
|
||||
self._set_memory_count(user_rpm_key, user_count)
|
||||
remaining_candidates.append(user_rpm_limit - user_count)
|
||||
|
||||
if key_rpm_limit > 0:
|
||||
key_count += 1
|
||||
self._set_memory_count(key_rpm_key, key_count)
|
||||
remaining_candidates.append(key_rpm_limit - key_count)
|
||||
|
||||
remaining = min(remaining_candidates) if remaining_candidates else None
|
||||
return RpmCheckResult(allowed=True, remaining=remaining)
|
||||
|
||||
def _start_background_cleanup(self) -> None:
|
||||
if self._cleanup_task is not None:
|
||||
return
|
||||
|
||||
async def cleanup_loop() -> None:
|
||||
while True:
|
||||
try:
|
||||
await asyncio.sleep(self._bucket_seconds)
|
||||
async with self._memory_lock:
|
||||
self._cleanup_expired_memory_counts(self._get_rpm_bucket())
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as exc:
|
||||
logger.debug("UserRpmLimiter 后台清理异常: {}", exc)
|
||||
|
||||
try:
|
||||
self._cleanup_task = asyncio.create_task(cleanup_loop())
|
||||
except RuntimeError:
|
||||
self._cleanup_task = None
|
||||
|
||||
def _get_rpm_bucket(self, now_ts: float | None = None) -> int:
|
||||
ts = now_ts if now_ts is not None else time.time()
|
||||
return int(ts // self._bucket_seconds)
|
||||
|
||||
def _split_scope_key(self, scope_key: str) -> tuple[str, int]:
|
||||
base_key, bucket_str = scope_key.rsplit(":", 1)
|
||||
return base_key, int(bucket_str)
|
||||
|
||||
def _get_memory_count(self, scope_key: str) -> int:
|
||||
base_key, bucket = self._split_scope_key(scope_key)
|
||||
stored = self._memory_counts.get(base_key)
|
||||
if not stored:
|
||||
return 0
|
||||
stored_bucket, count = stored
|
||||
if stored_bucket != bucket:
|
||||
self._memory_counts.pop(base_key, None)
|
||||
return 0
|
||||
return count
|
||||
|
||||
def _set_memory_count(self, scope_key: str, count: int) -> None:
|
||||
base_key, bucket = self._split_scope_key(scope_key)
|
||||
self._memory_counts[base_key] = (bucket, count)
|
||||
|
||||
def _cleanup_expired_memory_counts(self, current_bucket: int) -> None:
|
||||
expired_keys = [
|
||||
base_key
|
||||
for base_key, (bucket, _count) in self._memory_counts.items()
|
||||
if bucket < current_bucket
|
||||
]
|
||||
for base_key in expired_keys:
|
||||
self._memory_counts.pop(base_key, None)
|
||||
|
||||
if expired_keys:
|
||||
logger.debug(
|
||||
"[CLEANUP] 清理了 {} 个过期的用户/API Key RPM 计数(interval={}s)",
|
||||
len(expired_keys),
|
||||
self._cleanup_interval_seconds,
|
||||
)
|
||||
|
||||
|
||||
_user_rpm_limiter: UserRpmLimiter | None = None
|
||||
|
||||
|
||||
async def get_user_rpm_limiter() -> UserRpmLimiter:
|
||||
global _user_rpm_limiter
|
||||
if _user_rpm_limiter is None:
|
||||
_user_rpm_limiter = UserRpmLimiter()
|
||||
await _user_rpm_limiter.initialize()
|
||||
return _user_rpm_limiter
|
||||
@@ -61,6 +61,10 @@ class ApiKeyService:
|
||||
if final_expires_at is None and expire_days:
|
||||
final_expires_at = datetime.now(timezone.utc) + timedelta(days=expire_days)
|
||||
|
||||
normalized_rate_limit = rate_limit
|
||||
if not is_standalone and normalized_rate_limit is None:
|
||||
normalized_rate_limit = 0
|
||||
|
||||
api_key = ApiKey(
|
||||
user_id=user_id,
|
||||
key_hash=key_hash,
|
||||
@@ -69,7 +73,7 @@ class ApiKeyService:
|
||||
allowed_providers=allowed_providers,
|
||||
allowed_api_formats=allowed_api_formats,
|
||||
allowed_models=allowed_models,
|
||||
rate_limit=rate_limit,
|
||||
rate_limit=normalized_rate_limit,
|
||||
concurrent_limit=concurrent_limit,
|
||||
expires_at=final_expires_at,
|
||||
is_standalone=is_standalone,
|
||||
@@ -144,7 +148,8 @@ class ApiKeyService:
|
||||
# 允许显式设置为空数组/None 的字段(NULL=不限制,[]=全部禁用)
|
||||
nullable_list_fields = {"allowed_providers", "allowed_api_formats", "allowed_models"}
|
||||
|
||||
# 允许显式设置为 None 的字段(如 expires_at=None 表示永不过期,rate_limit=None 表示无限制)
|
||||
# 允许显式设置为 None 的字段(如 expires_at=None 表示永不过期;
|
||||
# standalone rate_limit=None 表示继承系统默认)
|
||||
nullable_fields = {"expires_at", "rate_limit"}
|
||||
|
||||
for field, value in kwargs.items():
|
||||
@@ -154,7 +159,9 @@ class ApiKeyService:
|
||||
if field in nullable_list_fields:
|
||||
setattr(api_key, field, value)
|
||||
elif field in nullable_fields:
|
||||
# 这些字段允许显式设置为 None
|
||||
if field == "rate_limit" and not api_key.is_standalone and value is None:
|
||||
setattr(api_key, field, 0)
|
||||
continue
|
||||
setattr(api_key, field, value)
|
||||
elif value is not None:
|
||||
setattr(api_key, field, value)
|
||||
@@ -180,39 +187,6 @@ class ApiKeyService:
|
||||
logger.info(f"删除API密钥: ID {key_id}")
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def check_rate_limit(db: Session, api_key: ApiKey, window_minutes: int = 1) -> tuple[bool, int]:
|
||||
"""检查速率限制
|
||||
|
||||
Returns:
|
||||
(is_allowed, remaining): 是否允许请求,剩余可用次数
|
||||
当 rate_limit 为 None 时表示不限制,返回 (True, -1)
|
||||
"""
|
||||
# 如果 rate_limit 为 None,表示不限制
|
||||
if api_key.rate_limit is None:
|
||||
return True, -1 # -1 表示无限制
|
||||
|
||||
# 计算时间窗口
|
||||
window_start = datetime.now(timezone.utc) - timedelta(minutes=window_minutes)
|
||||
|
||||
# 统计窗口内的请求数
|
||||
request_count = (
|
||||
db.query(func.count(Usage.id))
|
||||
.filter(Usage.api_key_id == api_key.id, Usage.created_at >= window_start)
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
|
||||
# 检查是否超限
|
||||
is_allowed = request_count < api_key.rate_limit
|
||||
|
||||
if not is_allowed:
|
||||
logger.warning(
|
||||
f"API密钥速率限制: Key ID {api_key.id}, 请求数 {request_count}/{api_key.rate_limit}"
|
||||
)
|
||||
|
||||
return is_allowed, api_key.rate_limit - request_count
|
||||
|
||||
@staticmethod
|
||||
def cleanup_expired_keys(db: Session, auto_delete: bool = False) -> int:
|
||||
"""清理过期的API密钥
|
||||
|
||||
@@ -38,6 +38,7 @@ class UserService:
|
||||
allowed_providers: list[str] | None = None,
|
||||
allowed_api_formats: list[str] | None = None,
|
||||
allowed_models: list[str] | None = None,
|
||||
rate_limit: int | None = None,
|
||||
) -> User:
|
||||
"""创建新用户。"""
|
||||
|
||||
@@ -74,6 +75,7 @@ class UserService:
|
||||
allowed_providers=allowed_providers,
|
||||
allowed_api_formats=allowed_api_formats,
|
||||
allowed_models=allowed_models,
|
||||
rate_limit=rate_limit,
|
||||
)
|
||||
user.set_password(password)
|
||||
|
||||
@@ -218,6 +220,7 @@ class UserService:
|
||||
"allowed_providers",
|
||||
"allowed_api_formats",
|
||||
"allowed_models",
|
||||
"rate_limit",
|
||||
]
|
||||
|
||||
# 允许设置为 None 的字段(表示无限制)
|
||||
@@ -225,6 +228,7 @@ class UserService:
|
||||
"allowed_providers",
|
||||
"allowed_api_formats",
|
||||
"allowed_models",
|
||||
"rate_limit",
|
||||
]
|
||||
|
||||
for field, value in kwargs.items():
|
||||
|
||||
Reference in New Issue
Block a user