mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
perf: Pool 批量删除改为异步任务模式,避免大批量删除阻塞请求
- 新增 batch_delete_task 模块,提交删除后立即返回 task_id,后台线程分批执行 - 新增查询任务进度的 API 端点,前端轮询展示实时进度 - RequestCandidate 大表清理改为按行数分批删除,防止单条语句超时
This commit is contained in:
@@ -46,6 +46,7 @@ from src.services.provider_keys.quota_reader import get_quota_reader
|
||||
from .schemas import (
|
||||
BatchActionRequest,
|
||||
BatchActionResponse,
|
||||
BatchDeleteTaskResponse,
|
||||
BatchImportError,
|
||||
BatchImportRequest,
|
||||
BatchImportResponse,
|
||||
@@ -449,6 +450,21 @@ async def batch_action_keys(
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/{provider_id}/keys/batch-delete-task/{task_id}",
|
||||
response_model=BatchDeleteTaskResponse,
|
||||
)
|
||||
async def get_batch_delete_task_status(
|
||||
provider_id: str,
|
||||
task_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> BatchDeleteTaskResponse:
|
||||
"""Query the progress of an async batch-delete task."""
|
||||
adapter = AdminBatchDeleteTaskStatusAdapter(provider_id=provider_id, task_id=task_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/{provider_id}/keys/cleanup-banned", response_model=BatchActionResponse)
|
||||
async def cleanup_banned_keys(
|
||||
provider_id: str,
|
||||
@@ -465,6 +481,28 @@ async def cleanup_banned_keys(
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminBatchDeleteTaskStatusAdapter(AdminApiAdapter):
|
||||
provider_id: str = ""
|
||||
task_id: str = ""
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
from fastapi import HTTPException
|
||||
|
||||
from src.services.provider_keys.batch_delete_task import get_batch_delete_task
|
||||
|
||||
task = get_batch_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 BatchDeleteTaskResponse(
|
||||
task_id=task.task_id,
|
||||
status=task.status,
|
||||
total=task.total,
|
||||
deleted=task.deleted,
|
||||
message=task.message,
|
||||
)
|
||||
|
||||
|
||||
class AdminListSchedulingPresetsAdapter(AdminApiAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
items: list[PresetDimensionMetaResponse] = [
|
||||
@@ -1070,86 +1108,23 @@ class AdminBatchActionKeysAdapter(AdminApiAdapter):
|
||||
affected = 0
|
||||
|
||||
if self.body.action == "delete":
|
||||
delete_started_at = time.perf_counter()
|
||||
sql_delete_ms = 0.0
|
||||
cleanup_ms = 0.0
|
||||
commit_ms = 0.0
|
||||
side_effects_ms = 0.0
|
||||
from src.services.provider_keys.batch_delete_task import submit_batch_delete
|
||||
|
||||
key_ids = list(dict.fromkeys(self.body.key_ids))
|
||||
delete_batch_size = _resolve_delete_batch_size(db)
|
||||
delete_batch_count = 0
|
||||
try:
|
||||
# 先清理关联表,避免 CASCADE 级联删除导致超时
|
||||
cleanup_started_at = time.perf_counter()
|
||||
cleanup_key_references(db, key_ids)
|
||||
cleanup_ms = (time.perf_counter() - cleanup_started_at) * 1000.0
|
||||
|
||||
for batch in _iter_batches(key_ids, delete_batch_size):
|
||||
batch_started_at = time.perf_counter()
|
||||
result = db.execute(
|
||||
sa_delete(ProviderAPIKey).where(
|
||||
ProviderAPIKey.provider_id == pid,
|
||||
ProviderAPIKey.id.in_(batch),
|
||||
)
|
||||
)
|
||||
sql_delete_ms += (time.perf_counter() - batch_started_at) * 1000.0
|
||||
delete_batch_count += 1
|
||||
rowcount = getattr(result, "rowcount", 0) or 0
|
||||
if rowcount > 0:
|
||||
affected += int(rowcount)
|
||||
commit_started_at = time.perf_counter()
|
||||
db.commit()
|
||||
commit_ms = (time.perf_counter() - commit_started_at) * 1000.0
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
total_ms = (time.perf_counter() - delete_started_at) * 1000.0
|
||||
logger.error(
|
||||
"batch delete commit failed: {} | provider={} requested={} batches={} sql_ms={:.2f} cleanup_ms={:.2f} commit_ms={:.2f} total_ms={:.2f}",
|
||||
exc,
|
||||
pid[:8],
|
||||
len(key_ids),
|
||||
delete_batch_count,
|
||||
sql_delete_ms,
|
||||
cleanup_ms,
|
||||
commit_ms,
|
||||
total_ms,
|
||||
)
|
||||
return BatchActionResponse(affected=0, message=f"commit failed: {exc}")
|
||||
|
||||
if affected > 0:
|
||||
from src.services.provider_keys.key_side_effects import (
|
||||
run_delete_key_side_effects,
|
||||
)
|
||||
|
||||
try:
|
||||
side_effects_started_at = time.perf_counter()
|
||||
await run_delete_key_side_effects(
|
||||
db=db,
|
||||
provider_id=pid,
|
||||
deleted_key_allowed_models=None,
|
||||
)
|
||||
side_effects_ms = (time.perf_counter() - side_effects_started_at) * 1000.0
|
||||
except Exception as exc:
|
||||
side_effects_ms = (time.perf_counter() - side_effects_started_at) * 1000.0
|
||||
logger.error("batch delete side effects failed: {}", exc)
|
||||
|
||||
total_ms = (time.perf_counter() - delete_started_at) * 1000.0
|
||||
task_id = submit_batch_delete(pid, key_ids)
|
||||
admin_name = context.user.username if context.user else "admin"
|
||||
logger.info(
|
||||
"[POOL_BATCH_DELETE_TIMING] provider={} requested={} affected={} batches={} batch_size={} cleanup_ms={:.2f} sql_ms={:.2f} commit_ms={:.2f} side_effects_ms={:.2f} total_ms={:.2f}",
|
||||
"Pool batch delete submitted by {}: provider={}, keys={}, task_id={}",
|
||||
admin_name,
|
||||
pid[:8],
|
||||
len(key_ids),
|
||||
affected,
|
||||
delete_batch_count,
|
||||
delete_batch_size,
|
||||
cleanup_ms,
|
||||
sql_delete_ms,
|
||||
commit_ms,
|
||||
side_effects_ms,
|
||||
total_ms,
|
||||
task_id,
|
||||
)
|
||||
return BatchActionResponse(
|
||||
affected=0,
|
||||
message=f"delete task submitted ({len(key_ids)} keys)",
|
||||
task_id=task_id,
|
||||
)
|
||||
|
||||
admin_name = context.user.username if context.user else "admin"
|
||||
affected_ids = [kid[:8] for kid in key_ids[:20]]
|
||||
else:
|
||||
keys = (
|
||||
db.query(ProviderAPIKey)
|
||||
|
||||
@@ -172,3 +172,12 @@ class BatchActionRequest(BaseModel):
|
||||
class BatchActionResponse(BaseModel):
|
||||
affected: int = 0
|
||||
message: str = ""
|
||||
task_id: str | None = None
|
||||
|
||||
|
||||
class BatchDeleteTaskResponse(BaseModel):
|
||||
task_id: str
|
||||
status: str # pending / running / completed / failed
|
||||
total: int = 0
|
||||
deleted: int = 0
|
||||
message: str = ""
|
||||
|
||||
208
src/services/provider_keys/batch_delete_task.py
Normal file
208
src/services/provider_keys/batch_delete_task.py
Normal file
@@ -0,0 +1,208 @@
|
||||
"""Pool Key 批量删除异步任务。
|
||||
|
||||
接口立即返回 task_id,后台执行删除,前端轮询进度。
|
||||
任务状态存内存字典,重启丢失无影响。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from sqlalchemy import delete as sa_delete
|
||||
|
||||
from src.core.logger import logger
|
||||
|
||||
# 任务状态
|
||||
STATUS_PENDING = "pending"
|
||||
STATUS_RUNNING = "running"
|
||||
STATUS_COMPLETED = "completed"
|
||||
STATUS_FAILED = "failed"
|
||||
|
||||
# 任务完成后保留时长(秒)
|
||||
_TASK_RETAIN_SECONDS = 600
|
||||
|
||||
# 删除 ProviderAPIKey 时的分批大小
|
||||
_DELETE_BATCH_SIZE = 500
|
||||
|
||||
# 同时存在的最大任务数(含已完成但未过期的)
|
||||
_MAX_TASKS = 100
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchDeleteTask:
|
||||
task_id: str
|
||||
provider_id: str
|
||||
status: str = STATUS_PENDING
|
||||
total: int = 0
|
||||
deleted: int = 0
|
||||
message: str = ""
|
||||
created_at: float = field(default_factory=time.monotonic)
|
||||
finished_at: float | None = None
|
||||
|
||||
|
||||
# 模块级任务注册表
|
||||
_tasks: dict[str, BatchDeleteTask] = {}
|
||||
# 持有后台 asyncio.Task 引用,防止 GC 回收
|
||||
_running_tasks: set[asyncio.Task[None]] = set()
|
||||
|
||||
|
||||
def _evict_finished_tasks() -> None:
|
||||
"""清除已过期的已完成任务,为新任务腾出空间。"""
|
||||
now = time.monotonic()
|
||||
expired = [
|
||||
tid
|
||||
for tid, t in _tasks.items()
|
||||
if t.finished_at is not None and (now - t.finished_at) > _TASK_RETAIN_SECONDS
|
||||
]
|
||||
for tid in expired:
|
||||
_tasks.pop(tid, None)
|
||||
|
||||
|
||||
def submit_batch_delete(provider_id: str, key_ids: list[str]) -> str:
|
||||
"""提交批量删除任务,返回 task_id。
|
||||
|
||||
Raises:
|
||||
RuntimeError: 并发任务数超过 _MAX_TASKS 上限。
|
||||
"""
|
||||
_evict_finished_tasks()
|
||||
if len(_tasks) >= _MAX_TASKS:
|
||||
raise RuntimeError(f"too many batch-delete tasks ({len(_tasks)}), try again later")
|
||||
task_id = uuid.uuid4().hex[:16]
|
||||
task = BatchDeleteTask(
|
||||
task_id=task_id,
|
||||
provider_id=provider_id,
|
||||
total=len(key_ids),
|
||||
)
|
||||
_tasks[task_id] = task
|
||||
bg = asyncio.create_task(
|
||||
_run_batch_delete(task_id, provider_id, key_ids),
|
||||
name=f"batch-delete-{task_id}",
|
||||
)
|
||||
_running_tasks.add(bg)
|
||||
bg.add_done_callback(_running_tasks.discard)
|
||||
return task_id
|
||||
|
||||
|
||||
def get_batch_delete_task(task_id: str) -> BatchDeleteTask | None:
|
||||
return _tasks.get(task_id)
|
||||
|
||||
|
||||
def _sync_delete(
|
||||
provider_id: str,
|
||||
key_ids: list[str],
|
||||
task: BatchDeleteTask,
|
||||
) -> int:
|
||||
"""在线程中执行的同步删除逻辑,避免阻塞事件循环。
|
||||
|
||||
每批 cleanup + DELETE + commit 作为独立事务,减少锁持有时间。
|
||||
"""
|
||||
from src.database import create_session
|
||||
from src.models.database import ProviderAPIKey
|
||||
from src.services.provider_keys.key_side_effects import cleanup_key_references
|
||||
|
||||
db = create_session()
|
||||
try:
|
||||
affected = 0
|
||||
for i in range(0, len(key_ids), _DELETE_BATCH_SIZE):
|
||||
batch = key_ids[i : i + _DELETE_BATCH_SIZE]
|
||||
try:
|
||||
cleanup_key_references(db, batch)
|
||||
result = db.execute(
|
||||
sa_delete(ProviderAPIKey).where(
|
||||
ProviderAPIKey.provider_id == provider_id,
|
||||
ProviderAPIKey.id.in_(batch),
|
||||
)
|
||||
)
|
||||
rowcount = getattr(result, "rowcount", 0) or 0
|
||||
affected += int(rowcount)
|
||||
db.commit()
|
||||
except Exception:
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
raise
|
||||
# NOTE: task.deleted is written from the worker thread and read
|
||||
# from the event-loop thread. Safe on CPython (GIL protects simple
|
||||
# attribute assignment) but not guaranteed by the language spec.
|
||||
task.deleted = affected
|
||||
return affected
|
||||
finally:
|
||||
try:
|
||||
db.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def _run_batch_delete(
|
||||
task_id: str,
|
||||
provider_id: str,
|
||||
key_ids: list[str],
|
||||
) -> None:
|
||||
task = _tasks.get(task_id)
|
||||
if task is None:
|
||||
return
|
||||
|
||||
task.status = STATUS_RUNNING
|
||||
|
||||
try:
|
||||
affected = await asyncio.to_thread(_sync_delete, provider_id, key_ids, task)
|
||||
|
||||
# 副作用(缓存失效等)是异步操作,在事件循环中执行
|
||||
if affected > 0:
|
||||
try:
|
||||
from src.database import get_db_context
|
||||
from src.services.provider_keys.key_side_effects import (
|
||||
run_delete_key_side_effects,
|
||||
)
|
||||
|
||||
with get_db_context() as db:
|
||||
await run_delete_key_side_effects(
|
||||
db=db,
|
||||
provider_id=provider_id,
|
||||
deleted_key_allowed_models=None,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error("batch delete side effects failed: {}", exc)
|
||||
|
||||
task.status = STATUS_COMPLETED
|
||||
task.deleted = affected
|
||||
task.message = f"{affected} keys deleted"
|
||||
logger.info(
|
||||
"[BATCH_DELETE_TASK] completed task={} provider={} total={} deleted={}",
|
||||
task_id,
|
||||
provider_id[:8],
|
||||
len(key_ids),
|
||||
affected,
|
||||
)
|
||||
|
||||
except asyncio.CancelledError:
|
||||
task.status = STATUS_FAILED
|
||||
task.message = "task cancelled (shutdown)"
|
||||
logger.warning(
|
||||
"[BATCH_DELETE_TASK] cancelled task={} provider={} deleted={}",
|
||||
task_id,
|
||||
provider_id[:8],
|
||||
task.deleted,
|
||||
)
|
||||
except Exception as exc:
|
||||
task.status = STATUS_FAILED
|
||||
task.message = str(exc)
|
||||
logger.error(
|
||||
"[BATCH_DELETE_TASK] failed task={} provider={}: {}",
|
||||
task_id,
|
||||
provider_id[:8],
|
||||
exc,
|
||||
)
|
||||
finally:
|
||||
task.finished_at = time.monotonic()
|
||||
|
||||
# 延迟清除已完成的任务
|
||||
try:
|
||||
await asyncio.sleep(_TASK_RETAIN_SECONDS)
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
_tasks.pop(task_id, None)
|
||||
@@ -4,7 +4,10 @@ Provider Key 写操作后的副作用处理。
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import delete as sa_delete
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
@@ -15,6 +18,10 @@ from src.services.cache.provider_cache import ProviderCacheService
|
||||
_SQLITE_BATCH_SIZE = 900
|
||||
_DEFAULT_BATCH_SIZE = 2000
|
||||
|
||||
# RequestCandidate 等大表每轮 DELETE 的行数上限,
|
||||
# 避免单条 DELETE 语句扫描过多行导致超时。
|
||||
_ROW_DELETE_BATCH = 5000
|
||||
|
||||
|
||||
def cleanup_key_references(db: Session, key_ids: list[str]) -> None:
|
||||
"""在删除 ProviderAPIKey 前,先清理关联表记录,避免 CASCADE 级联删除超时。"""
|
||||
@@ -22,11 +29,32 @@ def cleanup_key_references(db: Session, key_ids: list[str]) -> None:
|
||||
return
|
||||
batch_size = _resolve_batch_size(db)
|
||||
for batch in _iter_batches(key_ids, batch_size):
|
||||
db.execute(sa_delete(RequestCandidate).where(RequestCandidate.key_id.in_(batch)))
|
||||
# RequestCandidate 数据量通常最大,分批按行数删除
|
||||
_delete_in_row_batches(db, RequestCandidate, RequestCandidate.key_id, batch)
|
||||
db.execute(sa_delete(GeminiFileMapping).where(GeminiFileMapping.key_id.in_(batch)))
|
||||
db.execute(sa_delete(VideoTask).where(VideoTask.key_id.in_(batch)))
|
||||
|
||||
|
||||
def _delete_in_row_batches(db: Session, model: Any, fk_col: Any, key_id_batch: list[str]) -> int:
|
||||
"""分批删除大表记录,每次最多删除 _ROW_DELETE_BATCH 行,避免单条语句超时。"""
|
||||
total = 0
|
||||
while True:
|
||||
row_ids = [
|
||||
r[0]
|
||||
for r in db.execute(
|
||||
select(model.id).where(fk_col.in_(key_id_batch)).limit(_ROW_DELETE_BATCH)
|
||||
)
|
||||
]
|
||||
if not row_ids:
|
||||
break
|
||||
result = db.execute(sa_delete(model).where(model.id.in_(row_ids)))
|
||||
deleted = getattr(result, "rowcount", 0) or 0
|
||||
total += deleted
|
||||
if len(row_ids) < _ROW_DELETE_BATCH:
|
||||
break
|
||||
return total
|
||||
|
||||
|
||||
def _resolve_batch_size(db: Session) -> int:
|
||||
try:
|
||||
bind = db.get_bind()
|
||||
|
||||
Reference in New Issue
Block a user