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 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: if task is None or task.provider_id != self.provider_id:
raise HTTPException(status_code=404, detail="Task not found") raise HTTPException(status_code=404, detail="Task not found")
return BatchDeleteTaskResponse( return BatchDeleteTaskResponse(
@@ -1111,7 +1111,7 @@ class AdminBatchActionKeysAdapter(AdminApiAdapter):
from src.services.provider_keys.batch_delete_task import submit_batch_delete from src.services.provider_keys.batch_delete_task import submit_batch_delete
key_ids = list(dict.fromkeys(self.body.key_ids)) 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" admin_name = context.user.username if context.user else "admin"
logger.info( logger.info(
"Pool batch delete submitted by {}: provider={}, keys={}, task_id={}", "Pool batch delete submitted by {}: provider={}, keys={}, task_id={}",

View File

@@ -1,18 +1,21 @@
"""Pool Key 批量删除异步任务。 """Pool Key 批量删除异步任务。
接口立即返回 task_id后台执行删除前端轮询进度。 接口立即返回 task_id后台执行删除前端轮询进度。
任务状态存内存字典,重启丢失无影响 任务状态存储在 Redis 中,支持多 worker 进程共享
""" """
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import json
import time import time
import uuid 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 sqlalchemy import delete as sa_delete
from src.clients.redis_client import get_redis_client
from src.core.logger import logger from src.core.logger import logger
# 任务状态 # 任务状态
@@ -21,62 +24,137 @@ STATUS_RUNNING = "running"
STATUS_COMPLETED = "completed" STATUS_COMPLETED = "completed"
STATUS_FAILED = "failed" STATUS_FAILED = "failed"
# 任务完成后保留时长(秒) # 任务完成后保留时长(秒),也用作 Redis key 的 TTL
_TASK_RETAIN_SECONDS = 600 _TASK_RETAIN_SECONDS = 600
# 删除 ProviderAPIKey 时的分批大小 # 删除 ProviderAPIKey 时的分批大小
_DELETE_BATCH_SIZE = 500 _DELETE_BATCH_SIZE = 500
# 同时存在的最大任务数(含已完成但未过期的) # Redis key 前缀
_MAX_TASKS = 100 _REDIS_KEY_PREFIX = "batch_delete_task"
# 持有后台 asyncio.Task 引用,防止 GC 回收(进程内部,不需要跨进程共享)
@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() _running_tasks: set[asyncio.Task[None]] = set()
def _evict_finished_tasks() -> None: def _task_key(task_id: str) -> str:
"""清除已过期的已完成任务,为新任务腾出空间。""" return f"{_REDIS_KEY_PREFIX}:{task_id}"
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: 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。 """提交批量删除任务,返回 task_id。
Raises: Raises:
RuntimeError: 并发任务数超过 _MAX_TASKS 上限 RuntimeError: Redis 不可用时无法追踪任务状态
""" """
_evict_finished_tasks() r = await get_redis_client(require_redis=False)
if len(_tasks) >= _MAX_TASKS: if not r:
raise RuntimeError(f"too many batch-delete tasks ({len(_tasks)}), try again later") raise RuntimeError("Redis is required for batch delete tasks but is not available")
task_id = uuid.uuid4().hex[:16] task_id = uuid.uuid4().hex[:16]
task = BatchDeleteTask( task = BatchDeleteTaskInfo(
task_id=task_id, task_id=task_id,
provider_id=provider_id, provider_id=provider_id,
total=len(key_ids), total=len(key_ids),
) )
_tasks[task_id] = task await _save_task(task, ttl=_TASK_RETAIN_SECONDS * 2, r=r)
bg = asyncio.create_task( bg = asyncio.create_task(
_run_batch_delete(task_id, provider_id, key_ids), _run_batch_delete(task_id, provider_id, key_ids),
name=f"batch-delete-{task_id}", 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 return task_id
def get_batch_delete_task(task_id: str) -> BatchDeleteTask | None: async def get_batch_delete_task(task_id: str) -> BatchDeleteTaskInfo | None:
return _tasks.get(task_id) return await _load_task(task_id)
def _sync_delete( def _sync_delete(
provider_id: str, provider_id: str,
key_ids: list[str], key_ids: list[str],
task: BatchDeleteTask, progress_callback: Callable[[int], None] | None = None,
) -> int: ) -> int:
"""在线程中执行的同步删除逻辑,避免阻塞事件循环。 """在线程中执行的同步删除逻辑,避免阻塞事件循环。
@@ -125,10 +203,9 @@ def _sync_delete(
except Exception: except Exception:
pass pass
raise raise
# NOTE: task.deleted is written from the worker thread and read # 通过回调通知进度(在事件循环线程中更新 Redis
# from the event-loop thread. Safe on CPython (GIL protects simple if progress_callback is not None:
# attribute assignment) but not guaranteed by the language spec. progress_callback(affected)
task.deleted = affected
return affected return affected
finally: finally:
try: try:
@@ -142,14 +219,28 @@ async def _run_batch_delete(
provider_id: str, provider_id: str,
key_ids: list[str], key_ids: list[str],
) -> None: ) -> None:
task = _tasks.get(task_id) r = await get_redis_client(require_redis=False)
if task is None: await _update_task_field(task_id, r=r, status=STATUS_RUNNING)
return
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: 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: if affected > 0:
@@ -168,9 +259,13 @@ async def _run_batch_delete(
except Exception as exc: except Exception as exc:
logger.error("batch delete side effects failed: {}", exc) logger.error("batch delete side effects failed: {}", exc)
task.status = STATUS_COMPLETED await _update_task_field(
task.deleted = affected task_id,
task.message = f"{affected} keys deleted" r=r,
status=STATUS_COMPLETED,
deleted=affected,
message=f"{affected} keys deleted",
)
logger.info( logger.info(
"[BATCH_DELETE_TASK] completed task={} provider={} total={} deleted={}", "[BATCH_DELETE_TASK] completed task={} provider={} total={} deleted={}",
task_id, task_id,
@@ -180,29 +275,27 @@ async def _run_batch_delete(
) )
except asyncio.CancelledError: except asyncio.CancelledError:
task.status = STATUS_FAILED await _update_task_field(
task.message = "task cancelled (shutdown)" task_id,
r=r,
status=STATUS_FAILED,
message="task cancelled (shutdown)",
)
logger.warning( logger.warning(
"[BATCH_DELETE_TASK] cancelled task={} provider={} deleted={}", "[BATCH_DELETE_TASK] cancelled task={} provider={}",
task_id, task_id,
provider_id[:8], provider_id[:8],
task.deleted,
) )
except Exception as exc: except Exception as exc:
task.status = STATUS_FAILED await _update_task_field(
task.message = str(exc) task_id,
r=r,
status=STATUS_FAILED,
message=str(exc),
)
logger.error( logger.error(
"[BATCH_DELETE_TASK] failed task={} provider={}: {}", "[BATCH_DELETE_TASK] failed task={} provider={}: {}",
task_id, task_id,
provider_id[:8], provider_id[:8],
exc, exc,
) )
finally:
task.finished_at = time.monotonic()
# 延迟清除已完成的任务
try:
await asyncio.sleep(_TASK_RETAIN_SECONDS)
except asyncio.CancelledError:
pass
_tasks.pop(task_id, None)