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:
fawney19
2026-03-15 14:22:59 +08:00
parent 920a383136
commit f92b0943b5
35 changed files with 2051 additions and 238 deletions

View File

@@ -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,

View File

@@ -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),
)

View File

@@ -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 的锁定状态"""

View File

@@ -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]:

View File

@@ -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

View File

@@ -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 密钥的适配器"""

View File

@@ -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"))
# 异常处理配置

View File

@@ -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

View File

@@ -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)

View File

@@ -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

View File

@@ -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):

View File

@@ -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 能力配置

View File

@@ -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",
]

View 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

View File

@@ -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密钥

View File

@@ -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():