perf: Pool 批量删除改为异步任务模式,避免大批量删除阻塞请求

- 新增 batch_delete_task 模块,提交删除后立即返回 task_id,后台线程分批执行
- 新增查询任务进度的 API 端点,前端轮询展示实时进度
- RequestCandidate 大表清理改为按行数分批删除,防止单条语句超时
This commit is contained in:
fawney19
2026-03-09 01:35:01 +08:00
parent 0379f01ce8
commit e4476d0bc6
6 changed files with 386 additions and 80 deletions

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

View File

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