feat: 添加 API Key 锁定功能和请求头规则系统

1. API Key 锁定功能
- 管理员可锁定/解锁用户的 API Key
- 锁定后用户无法使用、修改或删除该 Key
- 前端显示锁定状态并禁用相关操作

2. 请求头规则系统
- 将 endpoint.headers 升级为 header_rules
- 支持 set(设置)、drop(删除)、rename(重命名)操作
- 前端提供可视化规则编辑界面
- 包含数据迁移脚本

Close #37 Close #86 Close #88
This commit is contained in:
fawney19
2026-01-16 01:18:54 +08:00
parent ea45918561
commit 8e11f3864a
22 changed files with 852 additions and 56 deletions

View File

@@ -210,6 +210,25 @@ async def delete_api_key(key_id: str, request: Request, db: Session = Depends(ge
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.patch("/{key_id}/lock")
async def toggle_lock_api_key(key_id: str, request: Request, db: Session = Depends(get_db)):
"""
切换 API Key 锁定状态
锁定/解锁指定的 API Key。锁定后用户无法使用和操作此密钥。
**路径参数**:
- `key_id`: API Key ID
**返回字段**:
- `id`: API Key ID
- `is_locked`: 新的锁定状态
- `message`: 提示信息
"""
adapter = AdminToggleLockApiKeyAdapter(key_id=key_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.patch("/{key_id}/balance")
async def add_balance_to_key(
key_id: str,
@@ -346,6 +365,7 @@ class AdminListStandaloneKeysAdapter(AdminApiAdapter):
"name": api_key.name,
"key_display": api_key.get_display_key(),
"is_active": api_key.is_active,
"is_locked": api_key.is_locked,
"is_standalone": api_key.is_standalone,
"current_balance_usd": api_key.current_balance_usd,
"balance_used_usd": float(api_key.balance_used_usd or 0),
@@ -541,6 +561,39 @@ class AdminToggleApiKeyAdapter(AdminApiAdapter):
}
class AdminToggleLockApiKeyAdapter(AdminApiAdapter):
"""切换API密钥锁定状态"""
def __init__(self, key_id: str):
self.key_id = key_id
async def handle(self, context): # type: ignore[override]
db = context.db
api_key = db.query(ApiKey).filter(ApiKey.id == self.key_id).first()
if not api_key:
raise NotFoundException("API密钥不存在", "api_key")
api_key.is_locked = not api_key.is_locked
api_key.updated_at = datetime.now(timezone.utc)
db.commit()
db.refresh(api_key)
logger.info(f"管理员切换API密钥锁定状态: Key ID {self.key_id}, 新状态 {'锁定' if api_key.is_locked else '解锁'}")
context.add_audit_metadata(
action="toggle_lock_api_key",
target_key_id=api_key.id,
user_id=api_key.user_id,
new_lock_status="locked" if api_key.is_locked else "unlocked",
)
return {
"id": api_key.id,
"is_locked": api_key.is_locked,
"message": f"API密钥已{'锁定' if api_key.is_locked else '解锁'}",
}
class AdminDeleteApiKeyAdapter(AdminApiAdapter):
def __init__(self, key_id: str):
self.key_id = key_id
@@ -663,6 +716,7 @@ class AdminGetKeyDetailAdapter(AdminApiAdapter):
"name": api_key.name,
"key_display": api_key.get_display_key(),
"is_active": api_key.is_active,
"is_locked": api_key.is_locked,
"is_standalone": api_key.is_standalone,
"current_balance_usd": api_key.current_balance_usd,
"balance_used_usd": float(api_key.balance_used_usd or 0),

View File

@@ -102,7 +102,7 @@ async def create_provider_endpoint(
- `api_format`: API 格式(如 claude、openai、gemini 等)
- `base_url`: 基础 URL
- `custom_path`: 自定义路径(可选)
- `headers`: 自定义请求头(可选
- `header_rules`: 请求头规则列表(可选,支持 set/drop/rename 操作
- `timeout`: 超时时间(秒,默认 300
- `max_retries`: 最大重试次数(默认 2
- `config`: 额外配置(可选)
@@ -169,7 +169,7 @@ async def update_endpoint(
**请求体字段**(均为可选):
- `base_url`: 基础 URL
- `custom_path`: 自定义路径
- `headers`: 自定义请求头
- `header_rules`: 请求头规则列表
- `timeout`: 超时时间(秒)
- `max_retries`: 最大重试次数
- `is_active`: 是否活跃
@@ -298,13 +298,14 @@ class AdminCreateProviderEndpointAdapter(AdminApiAdapter):
)
now = datetime.now(timezone.utc)
new_endpoint = ProviderEndpoint(
id=str(uuid.uuid4()),
provider_id=self.provider_id,
api_format=self.endpoint_data.api_format,
base_url=self.endpoint_data.base_url,
custom_path=self.endpoint_data.custom_path,
headers=self.endpoint_data.headers,
header_rules=self.endpoint_data.header_rules,
timeout=self.endpoint_data.timeout,
max_retries=self.endpoint_data.max_retries,
is_active=True,
@@ -398,6 +399,7 @@ class AdminUpdateProviderEndpointAdapter(AdminApiAdapter):
raise NotFoundException(f"Endpoint {self.endpoint_id} 不存在")
update_data = self.endpoint_data.model_dump(exclude_unset=True)
# 把 proxy 转换为 dict 存储,支持显式设置为 None 清除代理
if "proxy" in update_data:
if update_data["proxy"] is not None:

View File

@@ -15,11 +15,13 @@ from src.api.handlers.base.chat_adapter_base import get_adapter_class
from src.api.handlers.base.cli_adapter_base import get_cli_adapter_class
from src.config.constants import TimeoutDefaults
from src.core.crypto import crypto_service
from src.core.headers import get_extra_headers_from_endpoint
from src.core.logger import logger
from src.database.database import get_db
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint, User
from src.utils.auth_utils import get_current_user
router = APIRouter(prefix="/api/admin/provider-query", tags=["Provider Query"])
@@ -130,7 +132,7 @@ async def query_available_models(
"api_key": api_key_value,
"base_url": endpoint.base_url,
"api_format": fmt,
"extra_headers": endpoint.headers,
"extra_headers": get_extra_headers_from_endpoint(endpoint),
})
if not endpoint_configs:
@@ -160,7 +162,7 @@ async def query_available_models(
"api_key": api_key_value,
"base_url": endpoint.base_url,
"api_format": endpoint.api_format,
"extra_headers": endpoint.headers,
"extra_headers": get_extra_headers_from_endpoint(endpoint),
})
break # 只取第一个可用的 Key
@@ -321,7 +323,7 @@ async def test_model(
"api_key_id": api_key.id, # 添加API Key ID用于用量记录
"base_url": endpoint.base_url,
"api_format": endpoint.api_format,
"extra_headers": endpoint.headers,
"extra_headers": get_extra_headers_from_endpoint(endpoint),
"timeout": provider.timeout or TimeoutDefaults.HTTP_REQUEST,
}

View File

@@ -773,7 +773,7 @@ class AdminExportConfigAdapter(AdminApiAdapter):
{
"api_format": ep.api_format,
"base_url": ep.base_url,
"headers": ep.headers,
"header_rules": ep.header_rules,
"timeout": ep.timeout,
"max_retries": ep.max_retries,
"is_active": ep.is_active,
@@ -1063,7 +1063,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
existing_ep.base_url = ep_data.get(
"base_url", existing_ep.base_url
)
existing_ep.headers = ep_data.get("headers")
existing_ep.header_rules = ep_data.get("header_rules")
existing_ep.timeout = ep_data.get("timeout", 300)
existing_ep.max_retries = ep_data.get("max_retries", 2)
existing_ep.is_active = ep_data.get("is_active", True)
@@ -1078,7 +1078,7 @@ class AdminImportConfigAdapter(AdminApiAdapter):
provider_id=provider_id,
api_format=ep_data["api_format"],
base_url=ep_data["base_url"],
headers=ep_data.get("headers"),
header_rules=ep_data.get("header_rules"),
timeout=ep_data.get("timeout", 300),
max_retries=ep_data.get("max_retries", 2),
is_active=ep_data.get("is_active", True),

View File

@@ -489,6 +489,7 @@ class AdminGetUserKeysAdapter(AdminApiAdapter):
"name": key.name,
"key_display": key.get_display_key(),
"is_active": key.is_active,
"is_locked": key.is_locked,
"total_requests": key.total_requests,
"total_cost_usd": float(key.total_cost_usd or 0),
"rate_limit": key.rate_limit,

View File

@@ -140,16 +140,17 @@ class PassthroughRequestBuilder(RequestBuilder):
builder = HeaderBuilder()
# 2. 透传原始头部(排除敏感头部 - 黑名单模式
# 2. 透传原始头部(排除默认敏感头部)
if original_headers:
for name, value in original_headers.items():
if name.lower() in SENSITIVE_HEADERS:
continue
builder.add(name, value)
# 3. 添加 endpoint 配置的额外头部(不能覆盖认证头/Content-Type
if endpoint.headers:
builder.add_protected(endpoint.headers, protected_keys)
# 3. 应用 endpoint 的请求头规则
header_rules = getattr(endpoint, "header_rules", None)
if header_rules:
builder.apply_rules(header_rules, protected_keys)
# 4. 添加额外头部
if extra_headers:

View File

@@ -514,6 +514,7 @@ class ListMyApiKeysAdapter(AuthenticatedApiAdapter):
"name": key.name,
"key_display": key.get_display_key(),
"is_active": key.is_active,
"is_locked": key.is_locked,
"last_used_at": (
real_stats["last_used_at"].isoformat()
if real_stats["last_used_at"]
@@ -614,6 +615,7 @@ class GetMyApiKeyDetailAdapter(AuthenticatedApiAdapter):
"name": api_key.name,
"key_display": api_key.get_display_key(),
"is_active": api_key.is_active,
"is_locked": api_key.is_locked,
"allowed_providers": api_key.allowed_providers,
"force_capabilities": api_key.force_capabilities,
"rate_limit": api_key.rate_limit,
@@ -637,6 +639,8 @@ class DeleteMyApiKeyAdapter(AuthenticatedApiAdapter):
)
if not api_key:
raise NotFoundException("API密钥不存在", "api_key")
if api_key.is_locked:
raise ForbiddenException("该密钥已被管理员锁定,无法删除")
context.db.delete(api_key)
context.db.commit()
return {"message": "API密钥已删除"}
@@ -656,6 +660,8 @@ class ToggleMyApiKeyAdapter(AuthenticatedApiAdapter):
)
if not api_key:
raise NotFoundException("API密钥不存在", "api_key")
if api_key.is_locked:
raise ForbiddenException("该密钥已被管理员锁定,无法修改状态")
api_key.is_active = not api_key.is_active
context.db.commit()
context.db.refresh(api_key)
@@ -1055,6 +1061,8 @@ class UpdateApiKeyProvidersAdapter(AuthenticatedApiAdapter):
)
if not api_key:
raise NotFoundException("API密钥不存在")
if api_key.is_locked:
raise ForbiddenException("该密钥已被管理员锁定,无法修改")
if request.allowed_providers is not None and len(request.allowed_providers) > 0:
provider_ids = [cfg.provider_id for cfg in request.allowed_providers]
@@ -1101,6 +1109,8 @@ class UpdateApiKeyCapabilitiesAdapter(AuthenticatedApiAdapter):
)
if not api_key:
raise NotFoundException("API密钥不存在")
if api_key.is_locked:
raise ForbiddenException("该密钥已被管理员锁定,无法修改")
# 保存旧值用于审计
old_capabilities = api_key.force_capabilities

View File

@@ -213,7 +213,7 @@ class HeaderBuilder:
"""
添加头部但保护指定的 key 不被覆盖
用于 endpoint.headers 不能覆盖认证头的场景。
用于 endpoint 额外请求头不能覆盖认证头的场景。
"""
protected_lower = {k.lower() for k in protected_keys}
for k, v in headers.items():
@@ -227,6 +227,61 @@ class HeaderBuilder:
self._headers.pop(k.lower(), None)
return self
def rename(self, from_key: str, to_key: str) -> "HeaderBuilder":
"""
重命名头部(保留原值)
如果 from_key 不存在,则不做任何操作。
"""
from_lower = from_key.lower()
if from_lower in self._headers:
_, value = self._headers.pop(from_lower)
self._headers[to_key.lower()] = (to_key, value)
return self
def apply_rules(
self,
rules: list[Dict[str, Any]],
protected_keys: Optional[AbstractSet[str]] = None,
) -> "HeaderBuilder":
"""
应用请求头规则
支持的规则类型:
- set: 设置/覆盖头部 {"action": "set", "key": "X-Custom", "value": "fixed"}
- drop: 删除头部 {"action": "drop", "key": "X-Unwanted"}
- rename: 重命名头部 {"action": "rename", "from": "X-Old", "to": "X-New"}
Args:
rules: 规则列表
protected_keys: 受保护的 key不能被 set/drop/rename 修改)
"""
protected_lower = {k.lower() for k in protected_keys} if protected_keys else set()
for rule in rules:
action = rule.get("action")
if action == "set":
key = rule.get("key", "")
value = rule.get("value", "")
if key and key.lower() not in protected_lower:
self.add(key, value)
elif action == "drop":
key = rule.get("key", "")
if key and key.lower() not in protected_lower:
self._headers.pop(key.lower(), None)
elif action == "rename":
from_key = rule.get("from", "")
to_key = rule.get("to", "")
if from_key and to_key:
# 两个 key 都不能是受保护的
if from_key.lower() not in protected_lower and to_key.lower() not in protected_lower:
self.rename(from_key, to_key)
return self
def build(self) -> Dict[str, str]:
"""构建最终的头部字典"""
return {original_key: value for original_key, value in self._headers.values()}
@@ -471,3 +526,53 @@ def get_adapter_protected_keys(api_format: APIFormat) -> tuple[str, ...]:
"""
return tuple(get_protected_keys(api_format))
# =============================================================================
# Header Rules 工具函数
# =============================================================================
def extract_set_headers_from_rules(
header_rules: Optional[list[Dict[str, Any]]],
) -> Optional[Dict[str, str]]:
"""
从 header_rules 中提取 set 操作生成的头部字典
用于需要构造额外请求头的场景(如模型列表查询、模型测试等)。
注意drop 和 rename 操作在这里不适用,因为它们用于修改已存在的头部。
Args:
header_rules: 请求头规则列表 [{"action": "set", "key": "X-Custom", "value": "val"}, ...]
Returns:
set 操作生成的头部字典,如果没有则返回 None
"""
if not header_rules:
return None
headers: Dict[str, str] = {}
for rule in header_rules:
if rule.get("action") == "set":
key = rule.get("key", "")
value = rule.get("value", "")
if key:
headers[key] = value
return headers if headers else None
def get_extra_headers_from_endpoint(endpoint: Any) -> Optional[Dict[str, str]]:
"""
从 endpoint 提取额外请求头
用于需要构造额外请求头的场景(如模型列表查询、模型测试等)。
Args:
endpoint: ProviderEndpoint 对象
Returns:
额外请求头字典,如果没有则返回 None
"""
header_rules = getattr(endpoint, "header_rules", None)
return extract_set_headers_from_rules(header_rules)

View File

@@ -176,6 +176,7 @@ class ApiKey(Base):
# 状态
is_active = Column(Boolean, default=True, nullable=False)
is_locked = Column(Boolean, default=False, nullable=False) # 管理员锁定,用户无法使用/操作
last_used_at = Column(DateTime(timezone=True), nullable=True)
expires_at = Column(DateTime(timezone=True), nullable=True) # 过期时间
auto_delete_on_expiry = Column(Boolean, default=False, nullable=False) # 过期后是否自动删除
@@ -602,7 +603,7 @@ class ProviderEndpoint(Base):
base_url = Column(String(500), nullable=False)
# 请求配置
headers = Column(JSON, nullable=True) # 额外请求头
header_rules = Column(JSON, nullable=True) # 请求头规则 [{action, key, value, from, to}]
timeout = Column(Integer, default=300) # 超时(秒)
max_retries = Column(Integer, default=2) # 最大重试次数

View File

@@ -10,6 +10,16 @@ from pydantic import BaseModel, ConfigDict, Field, field_validator
from src.models.admin_requests import ProxyConfig
# ========== Header Rule 类型定义 ==========
# 请求头规则支持三种操作:
# - set: 设置/覆盖请求头 {"action": "set", "key": "X-Custom", "value": "val"}
# - drop: 删除请求头 {"action": "drop", "key": "X-Unwanted"}
# - rename: 重命名请求头 {"action": "rename", "from": "X-Old", "to": "X-New"}
# 实际验证在 headers.py 的 apply_rules 中处理
HeaderRule = Dict[str, Any]
# ========== ProviderEndpoint CRUD ==========
@@ -21,8 +31,12 @@ class ProviderEndpointCreate(BaseModel):
base_url: str = Field(..., min_length=1, max_length=500, description="API 基础 URL")
custom_path: Optional[str] = Field(default=None, max_length=200, description="自定义请求路径")
# 请求配置
headers: Optional[Dict[str, str]] = Field(default=None, description="自定义请求头")
# 请求配置
header_rules: Optional[List[HeaderRule]] = Field(
default=None,
description="请求头规则列表,支持 set/drop/rename 操作",
)
timeout: int = Field(default=300, ge=10, le=600, description="超时时间(秒)")
max_retries: int = Field(default=2, ge=0, le=10, description="最大重试次数")
@@ -60,7 +74,13 @@ class ProviderEndpointUpdate(BaseModel):
default=None, min_length=1, max_length=500, description="API 基础 URL"
)
custom_path: Optional[str] = Field(default=None, max_length=200, description="自定义请求路径")
headers: Optional[Dict[str, str]] = Field(default=None, description="自定义请求头")
# 请求头配置
header_rules: Optional[List[HeaderRule]] = Field(
default=None,
description="请求头规则列表,支持 set/drop/rename 操作",
)
timeout: Optional[int] = Field(default=None, ge=10, le=600, description="超时时间(秒)")
max_retries: Optional[int] = Field(default=None, ge=0, le=10, description="最大重试次数")
is_active: Optional[bool] = Field(default=None, description="是否启用")
@@ -92,8 +112,11 @@ class ProviderEndpointResponse(BaseModel):
base_url: str
custom_path: Optional[str] = None
# 请求配置
headers: Optional[Dict[str, str]] = None
# 请求配置
header_rules: Optional[List[HeaderRule]] = Field(
default=None, description="请求头规则列表"
)
timeout: int
max_retries: int

View File

@@ -22,6 +22,7 @@ from sqlalchemy.orm import Session, joinedload
from src.config import config
from src.core.logger import logger
from src.core.enums import AuthSource
from src.core.exceptions import ForbiddenException
from src.services.system.config import SystemConfigService
if TYPE_CHECKING:
@@ -385,6 +386,10 @@ class AuthService:
logger.warning("API认证失败 - 密钥已禁用")
return None
if key_record.is_locked:
logger.warning("API认证失败 - 密钥已被管理员锁定")
raise ForbiddenException("该API密钥已被管理员锁定请联系管理员")
# 检查过期时间
if key_record.expires_at:
# 确保 expires_at 是 aware datetime

View File

@@ -19,6 +19,7 @@ import httpx
from sqlalchemy.orm import Session, joinedload
from src.core.crypto import crypto_service
from src.core.headers import get_extra_headers_from_endpoint
from src.core.logger import logger
from src.database import create_session
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
@@ -268,7 +269,7 @@ class ModelFetchScheduler:
"api_key": api_key_value,
"base_url": endpoint.base_url,
"api_format": fmt,
"extra_headers": endpoint.headers,
"extra_headers": get_extra_headers_from_endpoint(endpoint),
}
)