mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
@@ -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"):
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
# Antigravity:enrich_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)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
# 额外复杂度评分
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
# ========== 使用统计 ==========
|
||||
|
||||
@@ -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 = False:access 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: {}",
|
||||
|
||||
293
src/services/provider/delete_cleanup.py
Normal file
293
src/services/provider/delete_cleanup.py
Normal 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),
|
||||
}
|
||||
540
src/services/provider/delete_task.py
Normal file
540
src/services/provider/delete_task.py
Normal 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,
|
||||
)
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
"""清理超过保留期的连接事件记录(在线程池中执行,避免阻塞事件循环)"""
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user