mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor: 批量删除任务状态存储从内存字典迁移到 Redis
- BatchDeleteTask 改为 BatchDeleteTaskInfo,状态序列化存入 Redis - submit_batch_delete / get_batch_delete_task 改为 async 函数 - 删除进度通过限频回调写入 Redis,替代直接修改内存属性 - 移除内存任务注册表和过期清理逻辑,依赖 Redis TTL 自动过期 - 路由层对应调整为 await 调用
This commit is contained in:
@@ -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={}",
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user