feat: Provider 异步删除、可配置密码策略、Hub 超时优化及多项改进

- 新增 Provider 异步删除任务系统,后台分阶段删除子资源并清理残留引用
- 新增可配置密码策略等级(weak/medium/strong),支持系统设置面板调整
- aether-hub 升级至 0.1.4,idle timeout 支持禁用(设为 0),worker 默认超时调整为 120s
- OAuth 手动续期增加 Redis 分布式锁,防止并发刷新冲突
- ProxyNode 心跳检测改为 asyncio.to_thread,避免阻塞事件循环
- 删除 ModelMultiSelect 和 useInvalidModels,MultiSelect 组件通用化
- 明确 allowed_providers/allowed_api_formats 的 NULL 与空数组语义
- 前端 StandaloneKeyFormDialog、UserFormDialog 等多处 UI 优化
- 新增 Alembic 迁移脚本清理 Provider 删除后的残留引用
- 补充相关测试用例
This commit is contained in:
fawney19
2026-03-12 01:11:35 +08:00
parent 0d770d1c4d
commit 6e51a3f45d
55 changed files with 3219 additions and 862 deletions

View File

@@ -447,7 +447,7 @@ class AdminUpdateApiKeyAdapter(AdminApiAdapter):
):
update_data["auto_delete_on_expiry"] = self.key_data.auto_delete_on_expiry
# 访问限制配置(允许设置为空数组来清除限制
# 访问限制配置(NULL=不限制,空数组=[]=全部禁用
if hasattr(self.key_data, "allowed_providers"):
update_data["allowed_providers"] = self.key_data.allowed_providers
if hasattr(self.key_data, "allowed_api_formats"):

View File

@@ -953,6 +953,8 @@ async def refresh_oauth(
db: Session = Depends(get_db),
_: User = Depends(require_admin),
) -> CompleteOAuthResponse:
from src.services.provider.auth import _acquire_refresh_lock, _release_refresh_lock
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise NotFoundException("Key 不存在", "key")
@@ -964,12 +966,57 @@ async def refresh_oauth(
raise NotFoundException("Provider 不存在", "provider")
provider_type = _require_fixed_provider(provider)
# Kiro 使用自定义 token refresh 机制
if provider_type == ProviderType.KIRO.value:
from datetime import datetime, timezone
redis, got_lock = await _acquire_refresh_lock(key_id)
if redis is not None and not got_lock:
raise InvalidRequestException("该 Key 正在续期,请稍后重试")
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
from src.services.provider.adapters.kiro.token_manager import refresh_access_token
try:
# Kiro 使用自定义 token refresh 机制
if provider_type == ProviderType.KIRO.value:
from datetime import datetime, timezone
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
from src.services.provider.adapters.kiro.token_manager import refresh_access_token
encrypted_auth_config = getattr(key, "auth_config", None)
if not encrypted_auth_config:
raise InvalidRequestException("缺少 auth_config无法 refresh")
decrypted = crypto_service.decrypt(encrypted_auth_config)
parsed = json.loads(decrypted)
cfg = KiroAuthConfig.from_dict(parsed)
cfg.provider_type = ProviderType.KIRO.value
from src.services.proxy_node.resolver import resolve_effective_proxy
proxy_config = resolve_effective_proxy(
getattr(provider, "proxy", None), getattr(key, "proxy", None)
)
try:
access_token, new_cfg = await refresh_access_token(cfg, proxy_config=proxy_config)
except Exception as e:
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = f"[REFRESH_FAILED] Token 续期失败: {e}"
db.commit()
logger.warning("Kiro Key {} token 刷新失败,已标记为刷新失效: {}", key_id, e)
raise InvalidRequestException("Kiro token refresh 失败,请检查凭据是否有效")
key.api_key = crypto_service.encrypt(access_token)
key.auth_config = crypto_service.encrypt(json.dumps(new_cfg.to_dict()))
key.oauth_invalid_at = None
key.oauth_invalid_reason = None
key.is_active = True
db.commit()
return CompleteOAuthResponse(
provider_type=provider_type,
expires_at=new_cfg.expires_at or None,
has_refresh_token=bool(new_cfg.refresh_token),
email=None,
)
template = _require_oauth_template(provider_type)
encrypted_auth_config = getattr(key, "auth_config", None)
if not encrypted_auth_config:
@@ -977,181 +1024,138 @@ async def refresh_oauth(
decrypted = crypto_service.decrypt(encrypted_auth_config)
parsed = json.loads(decrypted)
refresh_token = str(parsed.get("refresh_token") or "")
if not refresh_token:
raise InvalidRequestException("缺少 refresh_token需要重新授权")
cfg = KiroAuthConfig.from_dict(parsed)
cfg.provider_type = ProviderType.KIRO.value
token_url = template.oauth.token_url
is_json = "anthropic.com" in token_url
scope_str = " ".join(template.oauth.scopes) if template.oauth.scopes else ""
if is_json:
body: dict[str, Any] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
body["scope"] = scope_str
headers = {"Content-Type": "application/json", "Accept": "application/json"}
data = None
json_body = body
else:
form: dict[str, str] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
form["scope"] = scope_str
if template.oauth.client_secret:
form["client_secret"] = template.oauth.client_secret
headers = {
"Content-Type": "application/x-www-form-urlencoded",
"Accept": "application/json",
}
data = form
json_body = None
from src.services.proxy_node.resolver import resolve_effective_proxy
proxy_config = resolve_effective_proxy(
getattr(provider, "proxy", None), getattr(key, "proxy", None)
)
try:
access_token, new_cfg = await refresh_access_token(cfg, proxy_config=proxy_config)
except Exception as e:
# 标记为失效
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = f"[REFRESH_FAILED] Token 续期失败: {e}"
db.commit()
logger.warning("Kiro Key {} token 刷新失败,已标记为刷新失效: {}", key_id, e)
raise InvalidRequestException("Kiro token refresh 失败,请检查凭据是否有效")
# 更新 key
resp = await post_oauth_token(
provider_type=provider_type,
token_url=token_url,
headers=headers,
data=data,
json_body=json_body,
proxy_config=proxy_config,
timeout_seconds=30.0,
)
if resp.status_code < 200 or resp.status_code >= 300:
error_reason = f"HTTP {resp.status_code}"
try:
error_body = resp.json()
if "error" in error_body:
error_reason = str(
error_body.get("error_description") or error_body.get("error")
)
except Exception:
error_reason = resp.text[:100] if resp.text else f"HTTP {resp.status_code}"
if resp.status_code in (400, 401, 403):
from datetime import datetime, timezone
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = (
f"[REFRESH_FAILED] Token 续期失败 ({resp.status_code}): {error_reason}"
)
db.commit()
logger.warning(
"Key {} OAuth token 刷新失败,已标记为刷新失效: {}", key_id, error_reason
)
raise InvalidRequestException(f"token refresh 失败: {error_reason}")
token = resp.json()
access_token = str(token.get("access_token") or "")
new_refresh_token = str(token.get("refresh_token") or "")
expires_in = token.get("expires_in")
expires_at: int | None = None
try:
if expires_in is not None:
expires_at = int(time.time()) + int(expires_in)
except Exception:
expires_at = None
if not access_token:
raise InvalidRequestException("token refresh 返回缺少 access_token")
key.api_key = crypto_service.encrypt(access_token)
key.auth_config = crypto_service.encrypt(json.dumps(new_cfg.to_dict()))
key.oauth_invalid_at = None
key.oauth_invalid_reason = None
key.is_active = True
parsed["token_type"] = token.get("token_type")
if new_refresh_token:
parsed["refresh_token"] = new_refresh_token
parsed["expires_at"] = expires_at
parsed["scope"] = token.get("scope")
parsed["updated_at"] = int(time.time())
parsed = await enrich_auth_config(
provider_type=provider_type,
auth_config=parsed,
token_response=token,
access_token=access_token,
proxy_config=proxy_config,
)
if provider_type == ProviderType.ANTIGRAVITY and not parsed.get("project_id"):
logger.warning(
"[OAUTH_REFRESH] Antigravity key {} 刷新成功但 project_id 仍缺失,"
"下次刷新将继续尝试获取",
key_id,
)
key.auth_config = crypto_service.encrypt(json.dumps(parsed))
from src.services.provider.oauth_token import is_account_level_block
if not is_account_level_block(getattr(key, "oauth_invalid_reason", None)):
key.oauth_invalid_at = None
key.oauth_invalid_reason = None
key.is_active = True
db.commit()
return CompleteOAuthResponse(
provider_type=provider_type,
expires_at=new_cfg.expires_at or None,
has_refresh_token=bool(new_cfg.refresh_token),
email=None,
expires_at=expires_at,
has_refresh_token=bool(parsed.get("refresh_token")),
email=parsed.get("email"),
)
template = _require_oauth_template(provider_type)
encrypted_auth_config = getattr(key, "auth_config", None)
if not encrypted_auth_config:
raise InvalidRequestException("缺少 auth_config无法 refresh")
decrypted = crypto_service.decrypt(encrypted_auth_config)
parsed = json.loads(decrypted)
refresh_token = str(parsed.get("refresh_token") or "")
if not refresh_token:
raise InvalidRequestException("缺少 refresh_token需要重新授权")
token_url = template.oauth.token_url
is_json = "anthropic.com" in token_url
scope_str = " ".join(template.oauth.scopes) if template.oauth.scopes else ""
if is_json:
body: dict[str, Any] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
body["scope"] = scope_str
headers = {"Content-Type": "application/json", "Accept": "application/json"}
data = None
json_body = body
else:
form: dict[str, str] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if template.oauth.client_secret:
form["client_secret"] = template.oauth.client_secret
headers = {
"Content-Type": "application/x-www-form-urlencoded",
"Accept": "application/json",
}
data = form
json_body = None
from src.services.proxy_node.resolver import resolve_effective_proxy
proxy_config = resolve_effective_proxy(
getattr(provider, "proxy", None), getattr(key, "proxy", None)
)
resp = await post_oauth_token(
provider_type=provider_type,
token_url=token_url,
headers=headers,
data=data,
json_body=json_body,
proxy_config=proxy_config,
timeout_seconds=30.0,
)
if resp.status_code < 200 or resp.status_code >= 300:
# 解析错误原因
error_reason = f"HTTP {resp.status_code}"
try:
error_body = resp.json()
if "error" in error_body:
error_reason = str(error_body.get("error_description") or error_body.get("error"))
except Exception:
error_reason = resp.text[:100] if resp.text else f"HTTP {resp.status_code}"
# 标记为失效400/401/403 通常表示永久性错误)
if resp.status_code in (400, 401, 403):
from datetime import datetime, timezone
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = (
f"[REFRESH_FAILED] Token 续期失败 ({resp.status_code}): {error_reason}"
)
db.commit()
logger.warning(
"Key {} OAuth token 刷新失败,已标记为刷新失效: {}", key_id, error_reason
)
raise InvalidRequestException(f"token refresh 失败: {error_reason}")
token = resp.json()
access_token = str(token.get("access_token") or "")
new_refresh_token = str(token.get("refresh_token") or "")
expires_in = token.get("expires_in")
expires_at: int | None = None
try:
if expires_in is not None:
expires_at = int(time.time()) + int(expires_in)
except Exception:
expires_at = None
if not access_token:
raise InvalidRequestException("token refresh 返回缺少 access_token")
# store
key.api_key = crypto_service.encrypt(access_token)
parsed["token_type"] = token.get("token_type")
if new_refresh_token:
parsed["refresh_token"] = new_refresh_token
parsed["expires_at"] = expires_at
parsed["scope"] = token.get("scope")
parsed["updated_at"] = int(time.time())
parsed = await enrich_auth_config(
provider_type=provider_type,
auth_config=parsed,
token_response=token,
access_token=access_token,
proxy_config=proxy_config,
)
# Antigravityenrich_auth_config 会自动尝试补 project_id
# 即使本次仍未获取到也不阻断刷新token 已成功更新),下次刷新会继续重试
if provider_type == ProviderType.ANTIGRAVITY and not parsed.get("project_id"):
logger.warning(
"[OAUTH_REFRESH] Antigravity key {} 刷新成功但 project_id 仍缺失,"
"下次刷新将继续尝试获取",
key_id,
)
key.auth_config = crypto_service.encrypt(json.dumps(parsed))
# 刷新成功,清除 token 级别的失效标记
# 但保留账号级别的失效标记(以 OAUTH_ACCOUNT_BLOCK_PREFIX 开头),
# 这种不是 token 问题,刷新 token 解决不了
from src.services.provider.oauth_token import is_account_level_block
if not is_account_level_block(getattr(key, "oauth_invalid_reason", None)):
key.oauth_invalid_at = None
key.oauth_invalid_reason = None
key.is_active = True
db.commit()
return CompleteOAuthResponse(
provider_type=provider_type,
expires_at=expires_at,
has_refresh_token=bool(parsed.get("refresh_token")),
email=parsed.get("email"),
)
finally:
if got_lock:
await _release_refresh_lock(redis, key_id)
# ==============================================================================

View File

@@ -29,6 +29,7 @@ from src.models.database import GlobalModel, Provider, ProviderAPIKey, ProviderE
from src.models.endpoint_models import ProviderWithEndpointsSummary
from src.services.cache.model_cache import ModelCacheService
from src.services.cache.provider_cache import ProviderCacheService
from src.services.provider.delete_task import get_provider_delete_task, submit_provider_delete
from src.utils.cache_decorator import cache_result
from .summary import _build_provider_summary
@@ -232,6 +233,24 @@ class ProviderMappingPreviewResponse(BaseModel):
model_config = ConfigDict(from_attributes=True)
class ProviderDeleteSubmitResponse(BaseModel):
task_id: str
status: str = "pending"
message: str = ""
class ProviderDeleteTaskResponse(BaseModel):
task_id: str
provider_id: str
status: str
stage: str = "queued"
total_keys: int = 0
deleted_keys: int = 0
total_endpoints: int = 0
deleted_endpoints: int = 0
message: str = ""
@router.get("/")
async def list_providers(
request: Request,
@@ -335,10 +354,10 @@ async def update_provider(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.delete("/{provider_id}")
@router.delete("/{provider_id}", response_model=ProviderDeleteSubmitResponse)
async def delete_provider(
provider_id: str, request: Request, db: Session = Depends(get_db)
) -> None:
) -> ProviderDeleteSubmitResponse:
"""
删除提供商
@@ -348,12 +367,26 @@ async def delete_provider(
- `provider_id`: 提供商 ID
**返回字段**:
- `message`: 删除成功提示信息
- `task_id`: 后台删除任务 ID
- `status`: 任务状态
- `message`: 提交结果提示
"""
adapter = AdminDeleteProviderAdapter(provider_id=provider_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/{provider_id}/delete-task/{task_id}", response_model=ProviderDeleteTaskResponse)
async def get_delete_provider_task_status(
provider_id: str,
task_id: str,
request: Request,
db: Session = Depends(get_db),
) -> ProviderDeleteTaskResponse:
"""查询 Provider 删除任务状态。"""
adapter = AdminProviderDeleteTaskStatusAdapter(provider_id=provider_id, task_id=task_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
class AdminListProvidersAdapter(AdminApiAdapter):
def __init__(self, skip: int, limit: int, is_active: bool | None):
self.skip = skip
@@ -722,16 +755,50 @@ class AdminDeleteProviderAdapter(AdminApiAdapter):
provider_id=provider.id,
provider_name=provider.name,
)
db.delete(provider)
db.commit()
# 清除 /v1/models 列表缓存
await invalidate_models_list_cache()
task_id = await submit_provider_delete(provider.id)
# 清除 GlobalModel 解析缓存(删除 Provider 会影响模型解析结果)
await ModelCacheService.invalidate_all_resolve_cache()
provider_was_active = bool(provider.is_active)
if provider_was_active:
provider.is_active = False
db.commit()
await invalidate_models_list_cache()
await ModelCacheService.invalidate_all_resolve_cache()
await ProviderCacheService.invalidate_provider_cache(provider.id)
return {"message": "提供商已删除"}
context.add_audit_metadata(
task_id=task_id,
provider_deactivated=provider_was_active,
)
return {
"task_id": task_id,
"status": "pending",
"message": "删除任务已提交,提供商已进入后台删除队列",
}
class AdminProviderDeleteTaskStatusAdapter(AdminApiAdapter):
def __init__(self, provider_id: str, task_id: str):
self.provider_id = provider_id
self.task_id = task_id
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from fastapi import HTTPException
task = await get_provider_delete_task(self.task_id)
if task is None or task.provider_id != self.provider_id:
raise HTTPException(status_code=404, detail="Task not found")
return ProviderDeleteTaskResponse(
task_id=task.task_id,
provider_id=task.provider_id,
status=task.status,
stage=task.stage,
total_keys=task.total_keys,
deleted_keys=task.deleted_keys,
total_endpoints=task.total_endpoints,
deleted_endpoints=task.deleted_endpoints,
message=task.message,
)
@router.get(

View File

@@ -571,6 +571,7 @@ class AdminGetSystemSettingsAdapter(AdminApiAdapter):
default_provider=default_provider,
default_model=default_model,
enable_usage_tracking=enable_usage_tracking,
password_policy_level=SystemConfigService.get_password_policy_level(db),
)
@@ -620,6 +621,13 @@ class AdminUpdateSystemSettingsAdapter(AdminApiAdapter):
str(settings_request.enable_usage_tracking).lower(),
)
if settings_request.password_policy_level is not None:
SystemConfigService.set_config(
db,
"password_policy_level",
settings_request.password_policy_level,
)
return {"message": "系统设置更新成功"}
@@ -662,12 +670,15 @@ class AdminSetSystemConfigAdapter(AdminApiAdapter):
value = crypto_service.encrypt(value)
config = SystemConfigService.set_config(
context.db,
self.key,
value,
payload.get("description"),
)
try:
config = SystemConfigService.set_config(
context.db,
self.key,
value,
payload.get("description"),
)
except ValueError as exc:
raise InvalidRequestException(str(exc))
# 如果更新的是签到任务时间,动态更新调度器
if self.key == "provider_checkin_time" and value:

View File

@@ -295,10 +295,10 @@ class AdminCreateUserAdapter(AdminApiAdapter):
db, "default_user_initial_gift_usd", default=None
)
# 处理访问权限字段:空数组转为 None表示无限制
allowed_providers = request.allowed_providers if request.allowed_providers else None
allowed_api_formats = request.allowed_api_formats if request.allowed_api_formats else None
allowed_models = request.allowed_models if request.allowed_models else None
# 访问限制语义NULL=不限制,空数组=[]=全部禁用
allowed_providers = request.allowed_providers
allowed_api_formats = request.allowed_api_formats
allowed_models = request.allowed_models
try:
user = UserService.create_user(

View File

@@ -17,6 +17,7 @@ from src.api.base.context import ApiRequestContext
from src.api.base.pipeline import ApiRequestPipeline
from src.core.exceptions import InvalidRequestException
from src.core.logger import logger
from src.core.validators import PasswordValidator
from src.database import get_db
from src.models.api import (
LoginRequest,
@@ -415,6 +416,7 @@ class AuthRegistrationSettingsAdapter(AuthPublicAdapter):
enable_registration=bool(enable_registration),
require_email_verification=bool(require_verification),
email_configured=email_configured,
password_policy_level=SystemConfigService.get_password_policy_level(db),
).model_dump()
@@ -619,8 +621,10 @@ class AuthChangePasswordAdapter(AuthenticatedApiAdapter):
user = context.user
if not user.verify_password(old_password):
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="旧密码错误")
if len(new_password) < 6:
raise InvalidRequestException("密码长度至少6位")
policy_level = SystemConfigService.get_password_policy_level(context.db)
valid, error_msg = PasswordValidator.validate(new_password, policy=policy_level)
if not valid:
raise InvalidRequestException(error_msg or "密码格式无效")
user.set_password(new_password)
context.db.commit()
context.request.state.tx_committed_by_route = True

View File

@@ -24,6 +24,7 @@ from src.core.exceptions import (
translate_pydantic_error,
)
from src.core.logger import logger
from src.core.validators import PasswordValidator
from src.database import get_db
from src.models.api import (
ChangePasswordRequest,
@@ -43,6 +44,7 @@ from src.models.database import (
User,
UserModelUsageCount,
)
from src.services.system.config import SystemConfigService
from src.services.system.time_range import TimeRangeParams
from src.services.usage.service import UsageService
from src.services.user.apikey import ApiKeyService
@@ -539,8 +541,10 @@ class ChangePasswordAdapter(AuthenticatedApiAdapter):
raise InvalidRequestException("旧密码错误")
# 无密码(如 OAuth 用户首次设置):无需旧密码
if len(request.new_password) < 6:
raise InvalidRequestException("密码长度至少6位")
policy_level = SystemConfigService.get_password_policy_level(db)
valid, error_msg = PasswordValidator.validate(request.new_password, policy=policy_level)
if not valid:
raise InvalidRequestException(error_msg or "密码格式无效")
user.set_password(request.new_password)
user.updated_at = datetime.now(timezone.utc)
@@ -1480,7 +1484,7 @@ class UpdateApiKeyProvidersAdapter(AuthenticatedApiAdapter):
# 因为 allowed_providers 字段设计为存储 provider ID 字符串列表
api_key.allowed_providers = (
[cfg.provider_id for cfg in request.allowed_providers]
if request.allowed_providers
if request.allowed_providers is not None
else None
)
api_key.updated_at = datetime.now(timezone.utc)

View File

@@ -3,29 +3,48 @@
包含密码复杂度验证和其他输入验证
"""
from __future__ import annotations
import re
from enum import Enum
class PasswordPolicyLevel(str, Enum):
"""密码策略级别。"""
WEAK = "weak"
MEDIUM = "medium"
STRONG = "strong"
class PasswordValidator:
"""密码复杂度验证器"""
MIN_LENGTH = 6 # 降低到6位
MIN_LENGTH = 6
MEDIUM_MIN_LENGTH = 8
STRONG_MIN_LENGTH = 8
MAX_LENGTH = 128
@classmethod
def validate(cls, password: str) -> tuple[bool, str | None]:
"""
验证密码复杂度
def normalize_policy(cls, policy: PasswordPolicyLevel | str | None) -> PasswordPolicyLevel:
"""规范化密码策略级别,异常值回退为弱策略。"""
if isinstance(policy, PasswordPolicyLevel):
return policy
if isinstance(policy, str):
normalized = policy.strip().lower()
try:
return PasswordPolicyLevel(normalized)
except ValueError:
pass
return PasswordPolicyLevel.WEAK
要求:
- 长度至少6个字符
Args:
password: 待验证的密码
Returns:
(是否通过, 错误消息)
"""
@classmethod
def validate(
cls,
password: str,
policy: PasswordPolicyLevel | str | None = None,
) -> tuple[bool, str | None]:
"""验证密码复杂度。"""
if not password:
return False, "密码不能为空"
@@ -35,22 +54,26 @@ class PasswordValidator:
if len(password) > cls.MAX_LENGTH:
return False, f"密码长度不能超过{cls.MAX_LENGTH}个字符"
# 简化密码复杂度要求 - 只检查长度
# 不再要求大小写字母、数字和特殊字符
policy_level = cls.normalize_policy(policy)
# 检查常见弱密码
weak_passwords = [
"password123",
"admin123",
"12345678",
"qwerty123",
"password@123",
"admin@123",
"Password123!",
"Admin123!",
]
if password.lower() in [p.lower() for p in weak_passwords]:
return False, "密码过于简单,请使用更复杂的密码"
if policy_level == PasswordPolicyLevel.MEDIUM:
if len(password) < cls.MEDIUM_MIN_LENGTH:
return False, f"密码长度至少为{cls.MEDIUM_MIN_LENGTH}个字符"
if not re.search(r"[A-Za-z]", password):
return False, "密码必须包含至少一个字母"
if not re.search(r"\d", password):
return False, "密码必须包含至少一个数字"
elif policy_level == PasswordPolicyLevel.STRONG:
if len(password) < cls.STRONG_MIN_LENGTH:
return False, f"密码长度至少为{cls.STRONG_MIN_LENGTH}个字符"
if not re.search(r"[A-Z]", password):
return False, "密码必须包含至少一个大写字母"
if not re.search(r"[a-z]", password):
return False, "密码必须包含至少一个小写字母"
if not re.search(r"\d", password):
return False, "密码必须包含至少一个数字"
if not re.search(r"[!@#$%^&*()_+\-=\[\]{};:'\",.<>?/\\|`~]", password):
return False, "密码必须包含至少一个特殊字符"
return True, None
@@ -85,7 +108,7 @@ class PasswordValidator:
score += 1
if re.search(r"\d", password):
score += 1
if re.search(r'[!@#$%^&*()_+\-=\[\]{};:\'",.<>?/\\|`~]', password):
if re.search(r"[!@#$%^&*()_+\-=\[\]{};:\'\",.<>?/\\|`~]", password):
score += 2
# 额外复杂度评分

View File

@@ -11,6 +11,7 @@ from typing import Any, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from ..core.enums import UserRole
from ..core.validators import PasswordValidator
# ========== 认证相关 ==========
@@ -112,15 +113,10 @@ class RegisterRequest(BaseModel):
@classmethod
@field_validator("password")
def validate_password(cls, v: Any) -> Any:
"""验证密码强度"""
if len(v) < 6:
raise ValueError("密码至少需要6个字符")
if not re.search(r"[A-Z]", v):
raise ValueError("密码必须包含至少一个大写字母")
if not re.search(r"[a-z]", v):
raise ValueError("密码必须包含至少一个小写字母")
if not re.search(r"\d", v):
raise ValueError("密码必须包含至少一个数字")
"""基础校验(长度上下限),策略级别校验在服务层按系统配置执行。"""
valid, error_msg = PasswordValidator.validate(v)
if not valid:
raise ValueError(error_msg or "密码格式无效")
return v
@@ -231,6 +227,7 @@ class RegistrationSettingsResponse(BaseModel):
enable_registration: bool
require_email_verification: bool
email_configured: bool = Field(description="是否配置了邮箱服务")
password_policy_level: str = Field(description="密码策略等级weak/medium/strong")
# ========== 用户管理 ==========
@@ -320,15 +317,10 @@ class CreateUserRequest(BaseModel):
@classmethod
@field_validator("password")
def validate_password(cls, v: Any) -> Any:
"""验证密码强度"""
if len(v) < 6:
raise ValueError("密码至少需要6个字符")
if not re.search(r"[A-Z]", v):
raise ValueError("密码必须包含至少一个大写字母")
if not re.search(r"[a-z]", v):
raise ValueError("密码必须包含至少一个小写字母")
if not re.search(r"\d", v):
raise ValueError("密码必须包含至少一个数字")
"""基础校验(长度上下限),策略级别校验在服务层按系统配置执行。"""
valid, error_msg = PasswordValidator.validate(v)
if not valid:
raise ValueError(error_msg or "密码格式无效")
return v
@@ -645,6 +637,7 @@ class SystemSettingsRequest(BaseModel):
default_provider: str | None = None
default_model: str | None = None
enable_usage_tracking: bool | None = None
password_policy_level: Literal["weak", "medium", "strong"] | None = None
class SystemSettingsResponse(BaseModel):
@@ -653,6 +646,7 @@ class SystemSettingsResponse(BaseModel):
default_provider: str | None
default_model: str | None
enable_usage_tracking: bool
password_policy_level: str
# ========== 使用统计 ==========

View File

@@ -54,6 +54,44 @@ async def _release_refresh_lock(redis: Any, key_id: str) -> None:
pass
def _safe_object_session(key: Any) -> Any | None:
try:
return object_session(key)
except Exception:
return None
def _persist_detached_oauth_invalid_state(
key: Any,
*,
invalid_at: Any,
invalid_reason: str,
) -> None:
key.oauth_invalid_at = invalid_at
key.oauth_invalid_reason = invalid_reason
sess = _safe_object_session(key)
if sess is not None:
sess.add(key)
sess.commit()
return
key_id = str(getattr(key, "id", "") or "").strip()
if not key_id:
raise ValueError("OAuth key missing id")
from src.database import create_session
from src.models.database import ProviderAPIKey
with create_session() as db:
row = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if row is None:
raise ValueError(f"OAuth key not found: {key_id}")
row.oauth_invalid_at = invalid_at
row.oauth_invalid_reason = invalid_reason
db.commit()
def _persist_refreshed_token(
key: Any,
access_token: str,
@@ -71,7 +109,7 @@ def _persist_refreshed_token(
key.oauth_invalid_at = None
key.oauth_invalid_reason = None
sess = object_session(key)
sess = _safe_object_session(key)
if sess is not None:
sess.add(key)
sess.commit()
@@ -120,19 +158,16 @@ def _mark_refresh_token_invalid(
reason = f"{reason}: {detail}"
try:
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = reason
# 不设置 is_active = Falseaccess token 过期前 key 仍可调度
sess = object_session(key)
if sess is not None:
sess.add(key)
sess.commit()
logger.info(
"[OAUTH_REFRESH] key {} marked refresh_token invalid: {}",
str(getattr(key, "id", "?"))[:8],
reason[:120],
)
_persist_detached_oauth_invalid_state(
key,
invalid_at=datetime.now(timezone.utc),
invalid_reason=reason,
)
logger.info(
"[OAUTH_REFRESH] key {} marked refresh_token invalid: {}",
str(getattr(key, "id", "?"))[:8],
reason[:120],
)
except Exception as exc:
logger.warning(
"[OAUTH_REFRESH] failed to mark key {} refresh invalid: {}",
@@ -158,17 +193,15 @@ def _mark_oauth_token_expired(key: Any, expires_at: Any) -> None:
reason = f"[OAUTH_EXPIRED] Token 已过期且续期失败 (expired_at={expires_at})"
try:
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = reason
sess = object_session(key)
if sess is not None:
sess.add(key)
sess.commit()
logger.info(
"[OAUTH_EXPIRED] key {} token expired and refresh failed, blocking scheduling",
str(getattr(key, "id", "?"))[:8],
)
_persist_detached_oauth_invalid_state(
key,
invalid_at=datetime.now(timezone.utc),
invalid_reason=reason,
)
logger.info(
"[OAUTH_EXPIRED] key {} token expired and refresh failed, blocking scheduling",
str(getattr(key, "id", "?"))[:8],
)
except Exception as exc:
logger.warning(
"[OAUTH_EXPIRED] failed to mark key {} as expired: {}",

View File

@@ -0,0 +1,293 @@
from __future__ import annotations
from collections.abc import Iterable, Sequence
from typing import Any
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import (
ApiKey,
Model,
Provider,
ProviderAPIKey,
ProviderEndpoint,
RequestCandidate,
Usage,
User,
UserPreference,
VideoTask,
)
from src.models.database_extensions import ApiKeyProviderMapping, ProviderUsageTracking
from src.services.provider_keys.key_side_effects import cleanup_key_references
_BATCH_SIZE = 2000
def _empty_cleanup_stats() -> dict[str, int]:
return {
"users": 0,
"api_keys": 0,
"user_preferences": 0,
"usage_provider": 0,
"usage_endpoint": 0,
"video_tasks_provider": 0,
"video_tasks_endpoint": 0,
"request_candidates_provider": 0,
"request_candidates_endpoint": 0,
}
def _empty_delete_stats() -> dict[str, int]:
return {
"api_key_mappings": 0,
"usage_tracking": 0,
"models": 0,
"api_keys": 0,
"endpoints": 0,
"providers": 0,
}
def _iter_batches(items: Sequence[str], batch_size: int = _BATCH_SIZE) -> list[list[str]]:
if not items:
return []
if batch_size <= 0:
return [list(items)]
return [list(items[i : i + batch_size]) for i in range(0, len(items), batch_size)]
def _collect_provider_child_ids(db: Session, provider_id: str) -> tuple[list[str], list[str]]:
endpoint_ids = [
endpoint_id
for endpoint_id, in db.query(ProviderEndpoint.id)
.filter(ProviderEndpoint.provider_id == provider_id)
.all()
]
key_ids = [
key_id
for key_id, in db.query(ProviderAPIKey.id)
.filter(ProviderAPIKey.provider_id == provider_id)
.all()
]
return endpoint_ids, key_ids
def prune_allowed_provider_list(
allowed_providers: Any, provider_id: str
) -> tuple[list[str] | None | Any, bool]:
"""从访问限制列表中移除指定 Provider ID。"""
if not isinstance(allowed_providers, list):
return allowed_providers, False
if provider_id not in allowed_providers:
return allowed_providers, False
next_allowed = [value for value in allowed_providers if value != provider_id]
return next_allowed, True
def prune_allowed_provider_refs(records: Iterable[Any], provider_id: str) -> int:
"""批量移除记录中的 allowed_providers 引用。"""
updated = 0
for record in records:
next_allowed, changed = prune_allowed_provider_list(
getattr(record, "allowed_providers", None),
provider_id,
)
if not changed:
continue
record.allowed_providers = next_allowed
updated += 1
return updated
def cleanup_deleted_provider_references(
db: Session,
provider_id: str,
*,
endpoint_ids: Sequence[str] | None = None,
key_ids: Sequence[str] | None = None,
) -> dict[str, int]:
"""清理 Provider 删除时的大扇出引用,避免依赖数据库级联导致慢删。"""
if not provider_id:
return _empty_cleanup_stats()
if endpoint_ids is None or key_ids is None:
resolved_endpoint_ids, resolved_key_ids = _collect_provider_child_ids(db, provider_id)
endpoint_ids = resolved_endpoint_ids if endpoint_ids is None else list(endpoint_ids)
key_ids = resolved_key_ids if key_ids is None else list(key_ids)
else:
endpoint_ids = list(endpoint_ids)
key_ids = list(key_ids)
updated_users = prune_allowed_provider_refs(
db.query(User).filter(User.allowed_providers.isnot(None)).all(),
provider_id,
)
updated_api_keys = prune_allowed_provider_refs(
db.query(ApiKey).filter(ApiKey.allowed_providers.isnot(None)).all(),
provider_id,
)
cleared_preferences = int(
db.query(UserPreference)
.filter(UserPreference.default_provider_id == provider_id)
.update({UserPreference.default_provider_id: None}, synchronize_session=False)
or 0
)
cleared_usage_providers = int(
db.query(Usage)
.filter(Usage.provider_id == provider_id)
.update({Usage.provider_id: None}, synchronize_session=False)
or 0
)
cleared_video_task_providers = int(
db.query(VideoTask)
.filter(VideoTask.provider_id == provider_id)
.update({VideoTask.provider_id: None}, synchronize_session=False)
or 0
)
if key_ids:
cleanup_key_references(db, list(key_ids))
cleared_usage_endpoints = 0
cleared_video_task_endpoints = 0
deleted_request_candidates_endpoints = 0
for batch in _iter_batches(endpoint_ids):
cleared_usage_endpoints += int(
db.query(Usage)
.filter(Usage.provider_endpoint_id.in_(batch))
.update({Usage.provider_endpoint_id: None}, synchronize_session=False)
or 0
)
cleared_video_task_endpoints += int(
db.query(VideoTask)
.filter(VideoTask.endpoint_id.in_(batch))
.update({VideoTask.endpoint_id: None}, synchronize_session=False)
or 0
)
deleted_request_candidates_endpoints += int(
db.query(RequestCandidate)
.filter(RequestCandidate.endpoint_id.in_(batch))
.delete(synchronize_session=False)
or 0
)
deleted_request_candidates_provider = int(
db.query(RequestCandidate)
.filter(RequestCandidate.provider_id == provider_id)
.delete(synchronize_session=False)
or 0
)
stats = {
"users": updated_users,
"api_keys": updated_api_keys,
"user_preferences": cleared_preferences,
"usage_provider": cleared_usage_providers,
"usage_endpoint": cleared_usage_endpoints,
"video_tasks_provider": cleared_video_task_providers,
"video_tasks_endpoint": cleared_video_task_endpoints,
"request_candidates_provider": deleted_request_candidates_provider,
"request_candidates_endpoint": deleted_request_candidates_endpoints,
}
if any(stats.values()) or key_ids:
logger.info(
"Provider 删除引用清理: provider_id={}, key_refs={}, users={}, api_keys={}, "
"user_preferences={}, usage_provider={}, usage_endpoint={}, "
"video_tasks_provider={}, video_tasks_endpoint={}, "
"request_candidates_provider={}, request_candidates_endpoint={}",
provider_id,
len(key_ids),
stats["users"],
stats["api_keys"],
stats["user_preferences"],
stats["usage_provider"],
stats["usage_endpoint"],
stats["video_tasks_provider"],
stats["video_tasks_endpoint"],
stats["request_candidates_provider"],
stats["request_candidates_endpoint"],
)
return stats
def delete_provider_tree(db: Session, provider_id: str) -> dict[str, Any]:
"""分阶段删除 Provider 及其子资源,降低 ORM/FK 级联导致的超时风险。"""
if not provider_id:
return {
"cleanup": _empty_cleanup_stats(),
"deleted": _empty_delete_stats(),
"key_count": 0,
"endpoint_count": 0,
}
endpoint_ids, key_ids = _collect_provider_child_ids(db, provider_id)
cleanup_stats = cleanup_deleted_provider_references(
db,
provider_id,
endpoint_ids=endpoint_ids,
key_ids=key_ids,
)
deleted_stats = {
"api_key_mappings": int(
db.query(ApiKeyProviderMapping)
.filter(ApiKeyProviderMapping.provider_id == provider_id)
.delete(synchronize_session=False)
or 0
),
"usage_tracking": int(
db.query(ProviderUsageTracking)
.filter(ProviderUsageTracking.provider_id == provider_id)
.delete(synchronize_session=False)
or 0
),
"models": int(
db.query(Model)
.filter(Model.provider_id == provider_id)
.delete(synchronize_session=False)
or 0
),
"api_keys": int(
db.query(ProviderAPIKey)
.filter(ProviderAPIKey.provider_id == provider_id)
.delete(synchronize_session=False)
or 0
),
"endpoints": int(
db.query(ProviderEndpoint)
.filter(ProviderEndpoint.provider_id == provider_id)
.delete(synchronize_session=False)
or 0
),
"providers": int(
db.query(Provider).filter(Provider.id == provider_id).delete(synchronize_session=False)
or 0
),
}
logger.info(
"Provider 分阶段删除: provider_id={}, key_count={}, endpoint_count={}, "
"deleted_mappings={}, deleted_usage_tracking={}, deleted_models={}, "
"deleted_api_keys={}, deleted_endpoints={}, deleted_providers={}",
provider_id,
len(key_ids),
len(endpoint_ids),
deleted_stats["api_key_mappings"],
deleted_stats["usage_tracking"],
deleted_stats["models"],
deleted_stats["api_keys"],
deleted_stats["endpoints"],
deleted_stats["providers"],
)
return {
"cleanup": cleanup_stats,
"deleted": deleted_stats,
"key_count": len(key_ids),
"endpoint_count": len(endpoint_ids),
}

View File

@@ -0,0 +1,540 @@
"""Provider 异步删除任务。
接口提交后立即返回 task_id后台分阶段删除 provider 及其子资源。
任务状态存储在 Redis 中,支持多 worker 进程共享。
"""
from __future__ import annotations
import asyncio
import json
import time
import uuid
from collections.abc import Callable, Sequence
from concurrent.futures import Future
from typing import Any
import redis.asyncio as aioredis
from sqlalchemy import text
from src.clients.redis_client import get_redis_client
from src.core.logger import logger
from src.database import create_session
from src.models.database import (
ApiKey,
Model,
Provider,
ProviderAPIKey,
ProviderEndpoint,
RequestCandidate,
Usage,
User,
UserPreference,
VideoTask,
)
from src.models.database_extensions import ApiKeyProviderMapping, ProviderUsageTracking
from src.services.cache.model_cache import ModelCacheService
from src.services.cache.model_list_cache import invalidate_models_list_cache
from src.services.cache.provider_cache import ProviderCacheService
from src.services.provider.delete_cleanup import prune_allowed_provider_refs
from src.services.provider_keys.key_side_effects import cleanup_key_references
STATUS_PENDING = "pending"
STATUS_RUNNING = "running"
STATUS_COMPLETED = "completed"
STATUS_FAILED = "failed"
_TASK_RETAIN_SECONDS = 600
_KEY_BATCH_SIZE = 50
_ENDPOINT_BATCH_SIZE = 200
_BATCH_STATEMENT_TIMEOUT_S = 30
_BATCH_LOCK_TIMEOUT_S = 5
_TASK_TIMEOUT_S = 1800
_REDIS_KEY_PREFIX = "provider_delete_task"
_running_tasks: set[asyncio.Task[None]] = set()
def _task_key(task_id: str) -> str:
return f"{_REDIS_KEY_PREFIX}:{task_id}"
def _provider_lock_key(provider_id: str) -> str:
return f"{_REDIS_KEY_PREFIX}:provider:{provider_id}"
class ProviderDeleteTaskInfo:
__slots__ = (
"task_id",
"provider_id",
"status",
"stage",
"total_keys",
"deleted_keys",
"total_endpoints",
"deleted_endpoints",
"message",
)
def __init__(
self,
task_id: str,
provider_id: str,
status: str = STATUS_PENDING,
stage: str = "queued",
total_keys: int = 0,
deleted_keys: int = 0,
total_endpoints: int = 0,
deleted_endpoints: int = 0,
message: str = "",
) -> None:
self.task_id = task_id
self.provider_id = provider_id
self.status = status
self.stage = stage
self.total_keys = total_keys
self.deleted_keys = deleted_keys
self.total_endpoints = total_endpoints
self.deleted_endpoints = deleted_endpoints
self.message = message
def to_dict(self) -> dict[str, Any]:
return {
"task_id": self.task_id,
"provider_id": self.provider_id,
"status": self.status,
"stage": self.stage,
"total_keys": self.total_keys,
"deleted_keys": self.deleted_keys,
"total_endpoints": self.total_endpoints,
"deleted_endpoints": self.deleted_endpoints,
"message": self.message,
}
@classmethod
def from_dict(cls, data: dict[str, Any]) -> "ProviderDeleteTaskInfo":
return cls(
task_id=str(data["task_id"]),
provider_id=str(data["provider_id"]),
status=str(data.get("status", STATUS_PENDING)),
stage=str(data.get("stage", "queued")),
total_keys=int(data.get("total_keys", 0)),
deleted_keys=int(data.get("deleted_keys", 0)),
total_endpoints=int(data.get("total_endpoints", 0)),
deleted_endpoints=int(data.get("deleted_endpoints", 0)),
message=str(data.get("message", "")),
)
async def _save_task(
task: ProviderDeleteTaskInfo,
ttl: int = _TASK_RETAIN_SECONDS,
r: aioredis.Redis | None = None,
) -> None:
if r is None:
r = await get_redis_client(require_redis=False)
if not r:
return
try:
await r.setex(_task_key(task.task_id), ttl, json.dumps(task.to_dict()))
await r.setex(_provider_lock_key(task.provider_id), ttl, task.task_id)
except Exception as exc:
logger.warning("Failed to save provider delete task: {}", exc)
async def _load_task(
task_id: str,
r: aioredis.Redis | None = None,
) -> ProviderDeleteTaskInfo | None:
if r is None:
r = await get_redis_client(require_redis=False)
if not r:
return None
try:
data = await r.get(_task_key(task_id))
if data is None:
return None
return ProviderDeleteTaskInfo.from_dict(json.loads(data))
except Exception as exc:
logger.warning("Failed to load provider delete task: {}", exc)
return None
async def _update_task_field(
task_id: str,
r: aioredis.Redis | None = None,
**fields: object,
) -> None:
if r is None:
r = await get_redis_client(require_redis=False)
if not r:
return
task = await _load_task(task_id, r=r)
if task is None:
return
for key, value in fields.items():
setattr(task, key, value)
ttl = (
_TASK_RETAIN_SECONDS
if task.status in (STATUS_COMPLETED, STATUS_FAILED)
else _TASK_RETAIN_SECONDS * 2
)
await _save_task(task, ttl=ttl, r=r)
async def submit_provider_delete(provider_id: str) -> str:
r = await get_redis_client(require_redis=False)
if not r:
raise RuntimeError("Redis is required for provider delete tasks but is not available")
existing_task_id = await r.get(_provider_lock_key(provider_id))
if isinstance(existing_task_id, bytes):
existing_task_id = existing_task_id.decode()
if existing_task_id:
existing_task = await _load_task(str(existing_task_id), r=r)
if existing_task and existing_task.status in (STATUS_PENDING, STATUS_RUNNING):
return existing_task.task_id
task_id = uuid.uuid4().hex[:16]
task = ProviderDeleteTaskInfo(
task_id=task_id,
provider_id=provider_id,
message="delete task submitted",
)
await _save_task(task, ttl=_TASK_RETAIN_SECONDS * 2, r=r)
def _on_task_done(task: asyncio.Task[None]) -> None:
_running_tasks.discard(task)
if not task.cancelled() and task.exception():
logger.error("[PROVIDER_DELETE_TASK] unhandled error: {}", task.exception())
bg = asyncio.create_task(
_run_provider_delete(task_id, provider_id),
name=f"provider-delete-{task_id}",
)
_running_tasks.add(bg)
bg.add_done_callback(_on_task_done)
return task_id
async def get_provider_delete_task(task_id: str) -> ProviderDeleteTaskInfo | None:
return await _load_task(task_id)
def _iter_batches(items: Sequence[str], batch_size: int) -> list[list[str]]:
if not items:
return []
if batch_size <= 0:
return [list(items)]
return [list(items[i : i + batch_size]) for i in range(0, len(items), batch_size)]
def _apply_statement_timeouts(db: Any) -> None:
db.execute(text(f"SET LOCAL statement_timeout = '{_BATCH_STATEMENT_TIMEOUT_S * 1000}'"))
db.execute(text(f"SET LOCAL lock_timeout = '{_BATCH_LOCK_TIMEOUT_S * 1000}'"))
def _collect_ids(db: Any, provider_id: str) -> tuple[list[str], list[str]]:
endpoint_ids = [
endpoint_id
for endpoint_id, in db.query(ProviderEndpoint.id)
.filter(ProviderEndpoint.provider_id == provider_id)
.all()
]
key_ids = [
key_id
for key_id, in db.query(ProviderAPIKey.id)
.filter(ProviderAPIKey.provider_id == provider_id)
.all()
]
return endpoint_ids, key_ids
def _sync_delete_provider(
provider_id: str,
progress_callback: Callable[[dict[str, object]], None] | None = None,
) -> dict[str, Any]:
db = create_session()
try:
task_start = time.monotonic()
provider = db.query(Provider).filter(Provider.id == provider_id).first()
if not provider:
raise RuntimeError("provider not found")
_apply_statement_timeouts(db)
endpoint_ids, key_ids = _collect_ids(db, provider_id)
if progress_callback is not None:
progress_callback(
{
"stage": "preparing",
"total_keys": len(key_ids),
"total_endpoints": len(endpoint_ids),
"message": f"preparing delete for {len(key_ids)} keys and {len(endpoint_ids)} endpoints",
}
)
if getattr(provider, "is_active", False):
provider.is_active = False
db.commit()
if progress_callback is not None:
progress_callback(
{
"stage": "disabling",
"message": "provider disabled; starting cleanup",
}
)
_apply_statement_timeouts(db)
updated_users = prune_allowed_provider_refs(
db.query(User).filter(User.allowed_providers.isnot(None)).all(),
provider_id,
)
updated_api_keys = prune_allowed_provider_refs(
db.query(ApiKey).filter(ApiKey.allowed_providers.isnot(None)).all(),
provider_id,
)
db.commit()
if progress_callback is not None:
progress_callback(
{
"stage": "cleaning_restrictions",
"message": f"cleaned access restrictions (users={updated_users}, api_keys={updated_api_keys})",
}
)
_apply_statement_timeouts(db)
db.query(UserPreference).filter(UserPreference.default_provider_id == provider_id).update(
{UserPreference.default_provider_id: None},
synchronize_session=False,
)
db.query(Usage).filter(Usage.provider_id == provider_id).update(
{Usage.provider_id: None},
synchronize_session=False,
)
db.query(VideoTask).filter(VideoTask.provider_id == provider_id).update(
{VideoTask.provider_id: None},
synchronize_session=False,
)
db.query(RequestCandidate).filter(RequestCandidate.provider_id == provider_id).delete(
synchronize_session=False,
)
db.commit()
if progress_callback is not None:
progress_callback(
{
"stage": "cleaning_provider_refs",
"message": "cleaned provider-wide history references",
}
)
deleted_keys = 0
key_batches = _iter_batches(key_ids, _KEY_BATCH_SIZE)
for index, batch in enumerate(key_batches, start=1):
if time.monotonic() - task_start > _TASK_TIMEOUT_S:
raise RuntimeError(f"task timeout after {_TASK_TIMEOUT_S}s")
_apply_statement_timeouts(db)
cleanup_key_references(db, batch)
deleted_batch = int(
db.query(ProviderAPIKey)
.filter(ProviderAPIKey.provider_id == provider_id, ProviderAPIKey.id.in_(batch))
.delete(synchronize_session=False)
or 0
)
db.commit()
deleted_keys += deleted_batch
if progress_callback is not None:
progress_callback(
{
"stage": "deleting_keys",
"deleted_keys": deleted_keys,
"message": f"deleted key batch {index}/{max(len(key_batches), 1)}",
}
)
deleted_endpoints = 0
endpoint_batches = _iter_batches(endpoint_ids, _ENDPOINT_BATCH_SIZE)
for index, batch in enumerate(endpoint_batches, start=1):
if time.monotonic() - task_start > _TASK_TIMEOUT_S:
raise RuntimeError(f"task timeout after {_TASK_TIMEOUT_S}s")
_apply_statement_timeouts(db)
db.query(Usage).filter(Usage.provider_endpoint_id.in_(batch)).update(
{Usage.provider_endpoint_id: None},
synchronize_session=False,
)
db.query(VideoTask).filter(VideoTask.endpoint_id.in_(batch)).update(
{VideoTask.endpoint_id: None},
synchronize_session=False,
)
db.query(RequestCandidate).filter(RequestCandidate.endpoint_id.in_(batch)).delete(
synchronize_session=False,
)
deleted_batch = int(
db.query(ProviderEndpoint)
.filter(ProviderEndpoint.provider_id == provider_id, ProviderEndpoint.id.in_(batch))
.delete(synchronize_session=False)
or 0
)
db.commit()
deleted_endpoints += deleted_batch
if progress_callback is not None:
progress_callback(
{
"stage": "deleting_endpoints",
"deleted_endpoints": deleted_endpoints,
"message": f"deleted endpoint batch {index}/{max(len(endpoint_batches), 1)}",
}
)
_apply_statement_timeouts(db)
deleted_models = int(
db.query(Model)
.filter(Model.provider_id == provider_id)
.delete(synchronize_session=False)
or 0
)
deleted_mappings = int(
db.query(ApiKeyProviderMapping)
.filter(ApiKeyProviderMapping.provider_id == provider_id)
.delete(synchronize_session=False)
or 0
)
deleted_usage_tracking = int(
db.query(ProviderUsageTracking)
.filter(ProviderUsageTracking.provider_id == provider_id)
.delete(synchronize_session=False)
or 0
)
deleted_provider = int(
db.query(Provider).filter(Provider.id == provider_id).delete(synchronize_session=False)
or 0
)
db.commit()
return {
"provider_id": provider_id,
"total_keys": len(key_ids),
"deleted_keys": deleted_keys,
"total_endpoints": len(endpoint_ids),
"deleted_endpoints": deleted_endpoints,
"deleted_models": deleted_models,
"deleted_mappings": deleted_mappings,
"deleted_usage_tracking": deleted_usage_tracking,
"deleted_provider": deleted_provider,
"elapsed_seconds": time.monotonic() - task_start,
}
except Exception:
try:
db.rollback()
except Exception:
pass
raise
finally:
try:
db.close()
except Exception:
pass
async def _run_provider_delete(task_id: str, provider_id: str) -> None:
r = await get_redis_client(require_redis=False)
await _update_task_field(
task_id,
r=r,
status=STATUS_RUNNING,
stage="queued",
message="delete task started",
)
loop = asyncio.get_running_loop()
progress_futures: list[Future[object]] = []
def on_progress(fields: dict[str, object]) -> None:
try:
future = asyncio.run_coroutine_threadsafe(
_update_task_field(task_id, r=r, **fields),
loop,
)
progress_futures.append(future)
except RuntimeError:
pass
async def _drain_progress() -> None:
if progress_futures:
await asyncio.gather(
*(asyncio.wrap_future(f) for f in progress_futures),
return_exceptions=True,
)
try:
summary = await asyncio.wait_for(
asyncio.to_thread(_sync_delete_provider, provider_id, on_progress),
timeout=_TASK_TIMEOUT_S + 60,
)
await _drain_progress()
try:
await invalidate_models_list_cache()
await ModelCacheService.invalidate_all_resolve_cache()
await ProviderCacheService.invalidate_provider_cache(provider_id)
except Exception as exc:
logger.error("provider delete cache invalidation failed: {}", exc)
await _update_task_field(
task_id,
r=r,
status=STATUS_COMPLETED,
stage="completed",
total_keys=int(summary.get("total_keys", 0)),
deleted_keys=int(summary.get("deleted_keys", 0)),
total_endpoints=int(summary.get("total_endpoints", 0)),
deleted_endpoints=int(summary.get("deleted_endpoints", 0)),
message=(
f"provider deleted: keys={summary.get('deleted_keys', 0)}, "
f"endpoints={summary.get('deleted_endpoints', 0)}"
),
)
logger.info(
"[PROVIDER_DELETE_TASK] completed task={} provider={} keys={}/{} endpoints={}/{} elapsed={:.1f}s",
task_id,
provider_id[:8],
summary.get("deleted_keys", 0),
summary.get("total_keys", 0),
summary.get("deleted_endpoints", 0),
summary.get("total_endpoints", 0),
float(summary.get("elapsed_seconds", 0.0) or 0.0),
)
except asyncio.TimeoutError:
await _drain_progress()
msg = f"task timeout after {_TASK_TIMEOUT_S + 60}s"
await _update_task_field(task_id, r=r, status=STATUS_FAILED, stage="failed", message=msg)
logger.error("[PROVIDER_DELETE_TASK] {} task={} provider={}", msg, task_id, provider_id[:8])
except asyncio.CancelledError:
await _drain_progress()
await _update_task_field(
task_id,
r=r,
status=STATUS_FAILED,
stage="failed",
message="task cancelled (shutdown)",
)
logger.warning(
"[PROVIDER_DELETE_TASK] cancelled task={} provider={}",
task_id,
provider_id[:8],
)
except Exception as exc:
await _drain_progress()
await _update_task_field(
task_id,
r=r,
status=STATUS_FAILED,
stage="failed",
message=str(exc),
)
logger.error(
"[PROVIDER_DELETE_TASK] failed task={} provider={} error={}",
task_id,
provider_id[:8],
exc,
)

View File

@@ -19,6 +19,7 @@ from sqlalchemy import text
from src.clients.redis_client import get_redis_client
from src.core.logger import logger
from src.services.provider_keys.key_side_effects import cleanup_key_references
# 任务状态
STATUS_PENDING = "pending"
@@ -183,7 +184,7 @@ def _sync_delete(
) -> int:
"""在线程中执行的同步删除逻辑,避免阻塞事件循环。
按小批次_CLEANUP_BATCH_SIZE直接删除 key依赖数据库 CASCADE/SET NULL 自动清理关联表
按小批次_CLEANUP_BATCH_SIZE先显式处理关联表,再删除 key
每个批次独立事务,单批失败跳过并继续。
"""
from src.database import create_session
@@ -212,6 +213,7 @@ def _sync_delete(
# 设置 statement_timeout防止单条 SQL 无限等锁
timeout_ms = _BATCH_STATEMENT_TIMEOUT_S * 1000
db.execute(text(f"SET LOCAL statement_timeout = '{timeout_ms}'"))
cleanup_key_references(db, batch)
result = db.execute(
sa_delete(ProviderAPIKey).where(
ProviderAPIKey.provider_id == provider_id,

View File

@@ -31,6 +31,7 @@ from src.services.provider.fingerprint import generate_fingerprint, normalize_fi
from src.services.provider_keys.auth_type import normalize_auth_type
from src.services.provider_keys.duplicate_check import check_duplicate_key
from src.services.provider_keys.key_side_effects import (
cleanup_key_references,
run_create_key_side_effects,
run_delete_key_side_effects,
run_update_key_side_effects,
@@ -472,6 +473,7 @@ def _delete_endpoint_key(db: Session, key_id: str) -> _DeleteKeyResult:
provider_id = key.provider_id
deleted_key_allowed_models = key.allowed_models # 保存被删除 Key 的 allowed_models
try:
cleanup_key_references(db, [key_id])
db.delete(key)
db.commit()
except Exception as exc:
@@ -512,10 +514,11 @@ async def batch_delete_endpoint_keys_response(db: Session, key_ids: list[str]) -
# 收集受影响的 provider_id
affected_provider_ids = {key.provider_id for key in keys if key.provider_id}
# 批量 SQL DELETE,依赖数据库 CASCADE/SET NULL 自动清理关联表
# 批量 SQL DELETE 前先显式处理关联表,降低大批量删除时的级联成本
success_count = 0
try:
found_id_list = list(found_ids)
cleanup_key_references(db, found_id_list)
db.execute(sa_delete(ProviderAPIKey).where(ProviderAPIKey.id.in_(found_id_list)))
db.commit()
success_count = len(found_ids)

View File

@@ -5,10 +5,17 @@ Provider Key 写操作后的副作用处理。
from __future__ import annotations
from sqlalchemy import delete as sa_delete
from sqlalchemy import update as sa_update
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import GeminiFileMapping, ProviderAPIKey, RequestCandidate, VideoTask
from src.models.database import (
GeminiFileMapping,
ProviderAPIKey,
RequestCandidate,
Usage,
VideoTask,
)
from src.services.cache.model_list_cache import invalidate_models_list_cache
from src.services.cache.provider_cache import ProviderCacheService
@@ -17,9 +24,12 @@ _DEFAULT_BATCH_SIZE = 2000
def cleanup_key_references(db: Session, key_ids: list[str]) -> None:
"""在删除 ProviderAPIKey 前,先理关联表记录,避免 CASCADE 级联删除超时
"""在删除 ProviderAPIKey 前,先显式处理关联表引用,降低级联删除/置空成本
PostgreSQL 下直接使用 DELETE WHERE key_id IN (...),利用 key_id 索引高效删除
- request_candidates / gemini_file_mappings: 直接删除
- usage / video_tasks: 先置空外键,保留快照与历史记录
PostgreSQL 下直接按 key_id/provider_api_key_id 批量处理;
SQLite 下按 key_id 批次拆分,避免超出变量数限制。
"""
if not key_ids:
@@ -28,7 +38,12 @@ def cleanup_key_references(db: Session, key_ids: list[str]) -> None:
for batch in _iter_batches(key_ids, batch_size):
db.execute(sa_delete(RequestCandidate).where(RequestCandidate.key_id.in_(batch)))
db.execute(sa_delete(GeminiFileMapping).where(GeminiFileMapping.key_id.in_(batch)))
db.execute(sa_delete(VideoTask).where(VideoTask.key_id.in_(batch)))
db.execute(
sa_update(Usage)
.where(Usage.provider_api_key_id.in_(batch))
.values(provider_api_key_id=None)
)
db.execute(sa_update(VideoTask).where(VideoTask.key_id.in_(batch)).values(key_id=None))
def _resolve_batch_size(db: Session) -> int:

View File

@@ -83,51 +83,59 @@ class ProxyNodeHealthScheduler:
await self._cleanup_old_events()
async def _check_heartbeats(self) -> None:
db = create_session()
try:
now = datetime.now(timezone.utc)
# 检查所有非手动节点(手动节点无心跳,始终保持 ONLINE
# 包括 OFFLINE 节点:心跳恢复后可自愈
nodes = (
db.query(ProxyNode)
.filter(
ProxyNode.is_manual == False, # noqa: E712
)
.all()
)
if not nodes:
return
import asyncio
changed = 0
for node in nodes:
if heartbeat_is_stale(node, now):
if node.tunnel_connected:
node.tunnel_connected = False
node.tunnel_connected_at = now
changed += 1
if node.status != ProxyNodeStatus.OFFLINE:
node.status = ProxyNodeStatus.OFFLINE
def _sync_check() -> None:
db = create_session()
try:
now = datetime.now(timezone.utc)
# 检查所有非手动节点(手动节点无心跳,始终保持 ONLINE
# 包括 OFFLINE 节点:心跳恢复后可自愈
nodes = (
db.query(ProxyNode)
.filter(
ProxyNode.is_manual == False, # noqa: E712
)
.all()
)
if not nodes:
return
changed = 0
for node in nodes:
if heartbeat_is_stale(node, now):
if node.tunnel_connected:
node.tunnel_connected = False
node.tunnel_connected_at = now
changed += 1
if node.status != ProxyNodeStatus.OFFLINE:
node.status = ProxyNodeStatus.OFFLINE
node.updated_at = now
changed += 1
continue
# 心跳正常且连接状态为已连时,确保 ONLINE自愈状态不一致
if node.tunnel_connected and node.status != ProxyNodeStatus.ONLINE:
node.status = ProxyNodeStatus.ONLINE
node.updated_at = now
changed += 1
continue
# 心跳正常且连接状态为已连时,确保 ONLINE自愈状态不一致
if node.tunnel_connected and node.status != ProxyNodeStatus.ONLINE:
node.status = ProxyNodeStatus.ONLINE
node.updated_at = now
changed += 1
if changed:
db.commit()
logger.info("ProxyNode 心跳状态已更新: {} 个节点", changed)
except Exception as e:
try:
db.rollback()
except Exception:
pass
logger.exception("ProxyNode 心跳检测失败: {}", e)
finally:
db.close()
if changed:
db.commit()
logger.info("ProxyNode 心跳状态已更新: {} 个节点", changed)
try:
await asyncio.to_thread(_sync_check)
except Exception as e:
try:
db.rollback()
except Exception:
pass
logger.exception("ProxyNode 心跳检测失败: {}", e)
finally:
db.close()
logger.warning("ProxyNode 心跳检测线程执行失败: {}", e)
async def _cleanup_old_events(self) -> None:
"""清理超过保留期的连接事件记录(在线程池中执行,避免阻塞事件循环)"""

View File

@@ -12,6 +12,7 @@ from typing import Any
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.core.validators import PasswordPolicyLevel
from src.models.database import Provider, SystemConfig
REQUEST_RECORD_LEVEL_KEY = "request_record_level"
@@ -98,6 +99,10 @@ class SystemConfigService:
"value": 10.0,
"description": "新用户默认初始赠款(美元)",
},
"password_policy_level": {
"value": PasswordPolicyLevel.WEAK.value,
"description": "密码策略等级weak(弱密码), medium(中等强度), strong(强密码)",
},
REQUEST_RECORD_LEVEL_KEY: {
"value": RequestRecordLevel.BASIC.value,
"description": "请求记录级别basic(基本信息), headers(含请求/响应头), full(完整请求/响应)",
@@ -296,6 +301,13 @@ class SystemConfigService:
db: Session, key: str, value: Any, description: str | None = None
) -> SystemConfig:
"""设置系统配置值"""
if key == "password_policy_level":
normalized = str(value).strip().lower() if value is not None else ""
value = (
PasswordPolicyLevel(normalized).value
if normalized
else PasswordPolicyLevel.WEAK.value
)
# Backward-compatible alias: request_log_level -> request_record_level
if key in {REQUEST_RECORD_LEVEL_KEY, _LEGACY_REQUEST_LOG_LEVEL_KEY}:
config = (
@@ -354,6 +366,18 @@ class SystemConfigService:
return config
@staticmethod
def get_password_policy_level(db: Session) -> str:
"""获取密码策略等级,异常值自动回退为弱策略。"""
value = SystemConfigService.get_config(
db, "password_policy_level", PasswordPolicyLevel.WEAK.value
)
return (
PasswordPolicyLevel(value).value
if value in PasswordPolicyLevel._value2member_map_
else PasswordPolicyLevel.WEAK.value
)
@staticmethod
def get_default_provider(db: Session) -> str | None:
"""

View File

@@ -61,15 +61,14 @@ class ApiKeyService:
if final_expires_at is None and expire_days:
final_expires_at = datetime.now(timezone.utc) + timedelta(days=expire_days)
# 空数组转为 None表示不限制
api_key = ApiKey(
user_id=user_id,
key_hash=key_hash,
key_encrypted=key_encrypted,
name=name or f"API Key {datetime.now(timezone.utc).strftime('%Y%m%d%H%M%S')}",
allowed_providers=allowed_providers or None,
allowed_api_formats=allowed_api_formats or None,
allowed_models=allowed_models or None,
allowed_providers=allowed_providers,
allowed_api_formats=allowed_api_formats,
allowed_models=allowed_models,
rate_limit=rate_limit,
concurrent_limit=concurrent_limit,
expires_at=final_expires_at,
@@ -142,7 +141,7 @@ class ApiKeyService:
"auto_delete_on_expiry",
]
# 允许显式设置为空数组/None 的字段(空数组会转为 None表示"全部"
# 允许显式设置为空数组/None 的字段(NULL=不限制,[]=全部禁用
nullable_list_fields = {"allowed_providers", "allowed_api_formats", "allowed_models"}
# 允许显式设置为 None 的字段(如 expires_at=None 表示永不过期rate_limit=None 表示无限制)
@@ -151,11 +150,9 @@ class ApiKeyService:
for field, value in kwargs.items():
if field not in updatable_fields:
continue
# 对于 nullable_list_fields空数组应该转为 None表示不限制
# 对于 nullable_list_fields保留 None/[] 的语义差异
if field in nullable_list_fields:
if value is not None:
# 空数组转为 None表示允许全部
setattr(api_key, field, value if value else None)
setattr(api_key, field, value)
elif field in nullable_fields:
# 这些字段允许显式设置为 None
setattr(api_key, field, value)

View File

@@ -14,6 +14,7 @@ from src.core.logger import logger
from src.core.validators import EmailValidator, PasswordValidator, UsernameValidator
from src.models.database import ApiKey, GlobalModel, Model, Provider, Usage, User, UserRole
from src.services.cache.user_cache import UserCacheService
from src.services.system.config import SystemConfigService
from src.services.user.bulk_cleanup import batch_nullify_fk, pre_clean_api_key
from src.utils.async_utils import safe_create_task
from src.utils.transaction_manager import retry_on_database_error, transactional
@@ -55,7 +56,8 @@ class UserService:
raise ValueError(error_msg)
# 验证密码复杂度
valid, error_msg = PasswordValidator.validate(password)
policy_level = SystemConfigService.get_password_policy_level(db)
valid, error_msg = PasswordValidator.validate(password, policy=policy_level)
if not valid:
raise ValueError(error_msg)
@@ -237,7 +239,8 @@ class UserService:
# 如果提供了新密码
if "password" in kwargs and kwargs["password"]:
# 验证新密码复杂度
valid, error_msg = PasswordValidator.validate(kwargs["password"])
policy_level = SystemConfigService.get_password_policy_level(db)
valid, error_msg = PasswordValidator.validate(kwargs["password"], policy=policy_level)
if not valid:
raise ValueError(error_msg)
user.set_password(kwargs["password"])
@@ -381,7 +384,8 @@ class UserService:
return False, "旧密码错误"
# 验证新密码复杂度
valid, error_msg = PasswordValidator.validate(new_password)
policy_level = SystemConfigService.get_password_policy_level(db)
valid, error_msg = PasswordValidator.validate(new_password, policy=policy_level)
if not valid:
return False, error_msg