refactor: 批量删除任务状态存储从内存字典迁移到 Redis

- BatchDeleteTask 改为 BatchDeleteTaskInfo,状态序列化存入 Redis
- submit_batch_delete / get_batch_delete_task 改为 async 函数
- 删除进度通过限频回调写入 Redis,替代直接修改内存属性
- 移除内存任务注册表和过期清理逻辑,依赖 Redis TTL 自动过期
- 路由层对应调整为 await 调用
This commit is contained in:
fawney19
2026-03-09 02:09:24 +08:00
parent e4476d0bc6
commit 84cf07b7a2
2 changed files with 163 additions and 70 deletions

View File

@@ -491,7 +491,7 @@ class AdminBatchDeleteTaskStatusAdapter(AdminApiAdapter):
from src.services.provider_keys.batch_delete_task import get_batch_delete_task
task = get_batch_delete_task(self.task_id)
task = await 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(
@@ -1111,7 +1111,7 @@ class AdminBatchActionKeysAdapter(AdminApiAdapter):
from src.services.provider_keys.batch_delete_task import submit_batch_delete
key_ids = list(dict.fromkeys(self.body.key_ids))
task_id = submit_batch_delete(pid, key_ids)
task_id = await submit_batch_delete(pid, key_ids)
admin_name = context.user.username if context.user else "admin"
logger.info(
"Pool batch delete submitted by {}: provider={}, keys={}, task_id={}",

View File

@@ -1,18 +1,21 @@
"""Pool Key 批量删除异步任务。
接口立即返回 task_id后台执行删除前端轮询进度。
任务状态存内存字典,重启丢失无影响
任务状态存储在 Redis 中,支持多 worker 进程共享
"""
from __future__ import annotations
import asyncio
import json
import time
import uuid
from dataclasses import dataclass, field
from collections.abc import Callable
import redis.asyncio as aioredis
from sqlalchemy import delete as sa_delete
from src.clients.redis_client import get_redis_client
from src.core.logger import logger
# 任务状态
@@ -21,62 +24,137 @@ STATUS_RUNNING = "running"
STATUS_COMPLETED = "completed"
STATUS_FAILED = "failed"
# 任务完成后保留时长(秒)
# 任务完成后保留时长(秒),也用作 Redis key 的 TTL
_TASK_RETAIN_SECONDS = 600
# 删除 ProviderAPIKey 时的分批大小
_DELETE_BATCH_SIZE = 500
# 同时存在的最大任务数(含已完成但未过期的)
_MAX_TASKS = 100
# Redis key 前缀
_REDIS_KEY_PREFIX = "batch_delete_task"
@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 回收
# 持有后台 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 _task_key(task_id: str) -> str:
return f"{_REDIS_KEY_PREFIX}:{task_id}"
def submit_batch_delete(provider_id: str, key_ids: list[str]) -> str:
class BatchDeleteTaskInfo:
"""任务状态数据对象(从 Redis 反序列化)。"""
__slots__ = ("task_id", "provider_id", "status", "total", "deleted", "message")
def __init__(
self,
task_id: str,
provider_id: str,
status: str = STATUS_PENDING,
total: int = 0,
deleted: int = 0,
message: str = "",
):
self.task_id = task_id
self.provider_id = provider_id
self.status = status
self.total = total
self.deleted = deleted
self.message = message
def to_dict(self) -> dict:
return {
"task_id": self.task_id,
"provider_id": self.provider_id,
"status": self.status,
"total": self.total,
"deleted": self.deleted,
"message": self.message,
}
@classmethod
def from_dict(cls, data: dict) -> BatchDeleteTaskInfo:
return cls(
task_id=data["task_id"],
provider_id=data["provider_id"],
status=data.get("status", STATUS_PENDING),
total=int(data.get("total", 0)),
deleted=int(data.get("deleted", 0)),
message=data.get("message", ""),
)
async def _save_task(
task: BatchDeleteTaskInfo,
ttl: int = _TASK_RETAIN_SECONDS,
r: aioredis.Redis | None = None,
) -> None:
"""将任务状态写入 Redis。"""
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()))
except Exception as e:
logger.warning("Failed to save batch delete task to Redis: {}", e)
async def _load_task(task_id: str, r: aioredis.Redis | None = None) -> BatchDeleteTaskInfo | None:
"""从 Redis 加载任务状态。"""
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 BatchDeleteTaskInfo.from_dict(json.loads(data))
except Exception as e:
logger.warning("Failed to load batch delete task from Redis: {}", e)
return None
async def _update_task_field(
task_id: str, r: aioredis.Redis | None = None, **fields: object
) -> None:
"""局部更新任务状态字段read-modify-write"""
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 k, v in fields.items():
setattr(task, k, v)
# 已完成/失败的任务只保留 _TASK_RETAIN_SECONDS
if task.status in (STATUS_COMPLETED, STATUS_FAILED):
await _save_task(task, ttl=_TASK_RETAIN_SECONDS, r=r)
else:
# 运行中的任务使用更长的 TTL防止超长任务过期
await _save_task(task, ttl=_TASK_RETAIN_SECONDS * 2, r=r)
async def submit_batch_delete(provider_id: str, key_ids: list[str]) -> str:
"""提交批量删除任务,返回 task_id。
Raises:
RuntimeError: 并发任务数超过 _MAX_TASKS 上限
RuntimeError: Redis 不可用时无法追踪任务状态
"""
_evict_finished_tasks()
if len(_tasks) >= _MAX_TASKS:
raise RuntimeError(f"too many batch-delete tasks ({len(_tasks)}), try again later")
r = await get_redis_client(require_redis=False)
if not r:
raise RuntimeError("Redis is required for batch delete tasks but is not available")
task_id = uuid.uuid4().hex[:16]
task = BatchDeleteTask(
task = BatchDeleteTaskInfo(
task_id=task_id,
provider_id=provider_id,
total=len(key_ids),
)
_tasks[task_id] = task
await _save_task(task, ttl=_TASK_RETAIN_SECONDS * 2, r=r)
bg = asyncio.create_task(
_run_batch_delete(task_id, provider_id, key_ids),
name=f"batch-delete-{task_id}",
@@ -86,14 +164,14 @@ def submit_batch_delete(provider_id: str, key_ids: list[str]) -> str:
return task_id
def get_batch_delete_task(task_id: str) -> BatchDeleteTask | None:
return _tasks.get(task_id)
async def get_batch_delete_task(task_id: str) -> BatchDeleteTaskInfo | None:
return await _load_task(task_id)
def _sync_delete(
provider_id: str,
key_ids: list[str],
task: BatchDeleteTask,
progress_callback: Callable[[int], None] | None = None,
) -> int:
"""在线程中执行的同步删除逻辑,避免阻塞事件循环。
@@ -125,10 +203,9 @@ def _sync_delete(
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
# 通过回调通知进度(在事件循环线程中更新 Redis
if progress_callback is not None:
progress_callback(affected)
return affected
finally:
try:
@@ -142,14 +219,28 @@ async def _run_batch_delete(
provider_id: str,
key_ids: list[str],
) -> None:
task = _tasks.get(task_id)
if task is None:
return
r = await get_redis_client(require_redis=False)
await _update_task_field(task_id, r=r, status=STATUS_RUNNING)
task.status = STATUS_RUNNING
# 用于从工作线程安全地触发 Redis 进度更新(按时间限频,最多每 2 秒写一次)
loop = asyncio.get_running_loop()
last_report_time = [0.0]
_PROGRESS_INTERVAL = 2.0
def on_progress(current: int) -> None:
now = time.monotonic()
if now - last_report_time[0] < _PROGRESS_INTERVAL:
return
last_report_time[0] = now
try:
asyncio.run_coroutine_threadsafe(
_update_task_field(task_id, r=r, deleted=current), loop
)
except RuntimeError:
pass
try:
affected = await asyncio.to_thread(_sync_delete, provider_id, key_ids, task)
affected = await asyncio.to_thread(_sync_delete, provider_id, key_ids, on_progress)
# 副作用(缓存失效等)是异步操作,在事件循环中执行
if affected > 0:
@@ -168,9 +259,13 @@ async def _run_batch_delete(
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"
await _update_task_field(
task_id,
r=r,
status=STATUS_COMPLETED,
deleted=affected,
message=f"{affected} keys deleted",
)
logger.info(
"[BATCH_DELETE_TASK] completed task={} provider={} total={} deleted={}",
task_id,
@@ -180,29 +275,27 @@ async def _run_batch_delete(
)
except asyncio.CancelledError:
task.status = STATUS_FAILED
task.message = "task cancelled (shutdown)"
await _update_task_field(
task_id,
r=r,
status=STATUS_FAILED,
message="task cancelled (shutdown)",
)
logger.warning(
"[BATCH_DELETE_TASK] cancelled task={} provider={} deleted={}",
"[BATCH_DELETE_TASK] cancelled task={} provider={}",
task_id,
provider_id[:8],
task.deleted,
)
except Exception as exc:
task.status = STATUS_FAILED
task.message = str(exc)
await _update_task_field(
task_id,
r=r,
status=STATUS_FAILED,
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)