f45baa2ce0
实现3层队列防护,防止批量提交导致队列爆炸: **用户级限流(核心)** - 每个用户同时 pending 的 generation 任务上限 3 个 - 超限返回 429:您的待处理任务过多,请等待完成后再提交 **全局限流(兜底)** - 系统 pending 任务超过 20 个一律拒绝 - 返回 503:系统繁忙,请稍后再试 **覆盖的入口** - 一键生成(generation_tasks 批量创建 + 重试) - 剪辑计划生成(edit_plans generate) - 任务中心重试(task_center 用户级 + 项目级) **实现细节** - 预检查 + 入队前检查双重保障 - 限流拒绝时任务标记为 failed,避免 pending 僵尸 - repository 不支持计数时自动降级跳过(兼容旧代码) - 全局检查始终生效,用户级检查需传 user_id **新增** - task_enqueue.py: UserPendingLimitExceeded / GlobalQueueFull 异常 - generation_task_repository: count_pending_by_user / count_pending_total - 13个单元测试,覆盖正常/超限/全局/降级等场景
181 lines
6.2 KiB
Python
Executable File
181 lines
6.2 KiB
Python
Executable File
import logging
|
|
from typing import Any
|
|
|
|
from app.core.celery_app import celery_app
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class UserPendingLimitExceeded(Exception):
|
|
"""用户 pending 任务数超限,返回 429。"""
|
|
|
|
def __init__(self, user_id: str, pending_count: int, limit: int):
|
|
self.user_id = user_id
|
|
self.pending_count = pending_count
|
|
self.limit = limit
|
|
super().__init__(f"用户 {user_id} pending 任务数 {pending_count} 超过上限 {limit}")
|
|
|
|
|
|
class GlobalQueueFull(Exception):
|
|
"""全局限流,返回 503。"""
|
|
|
|
def __init__(self, pending_count: int, limit: int):
|
|
self.pending_count = pending_count
|
|
self.limit = limit
|
|
super().__init__(f"系统 pending 任务数 {pending_count} 超过上限 {limit}")
|
|
|
|
|
|
def check_queue_limits(
|
|
user_id: str,
|
|
generation_task_repository: Any,
|
|
*,
|
|
user_pending_limit: int = 3,
|
|
global_pending_limit: int = 20,
|
|
) -> None:
|
|
"""检查队列限流,超限抛对应异常。
|
|
|
|
Args:
|
|
user_id: 用户 ID
|
|
generation_task_repository: 任务仓储
|
|
user_pending_limit: 单用户 pending 上限,默认 3
|
|
global_pending_limit: 全局 pending 上限,默认 20
|
|
|
|
Raises:
|
|
GlobalQueueFull: 全局超限时抛出(优先级更高,先查全局)
|
|
UserPendingLimitExceeded: 用户超限时抛出
|
|
"""
|
|
# 先查全局(系统级保护优先级更高)
|
|
global_pending = generation_task_repository.count_pending_total()
|
|
if global_pending > global_pending_limit:
|
|
logger.warning(
|
|
"[队列限流] 全局 pending 任务数超限: %d/%d, user_id=%s",
|
|
global_pending,
|
|
global_pending_limit,
|
|
user_id,
|
|
)
|
|
raise GlobalQueueFull(pending_count=global_pending, limit=global_pending_limit)
|
|
|
|
# 再查用户级
|
|
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
|
if user_pending > user_pending_limit:
|
|
logger.warning(
|
|
"[队列限流] 用户 pending 任务数超限: user_id=%s, count=%d/%d",
|
|
user_id,
|
|
user_pending,
|
|
user_pending_limit,
|
|
)
|
|
raise UserPendingLimitExceeded(
|
|
user_id=user_id, pending_count=user_pending, limit=user_pending_limit
|
|
)
|
|
|
|
|
|
def _mark_task_failed_safely(
|
|
task: Any,
|
|
generation_task_repository: Any,
|
|
log_prefix: str,
|
|
reason: str,
|
|
) -> None:
|
|
"""安全地把任务标记为 failed,更新失败只打日志不崩溃。"""
|
|
try:
|
|
task.mark_failed(f"任务被限流拒绝: {reason}")
|
|
generation_task_repository.update(task)
|
|
except Exception as update_err:
|
|
logger.error(
|
|
"%s 限流后更新状态也失败: task_id=%s error=%s",
|
|
log_prefix,
|
|
task.id,
|
|
update_err,
|
|
exc_info=True,
|
|
)
|
|
|
|
|
|
def safe_enqueue_generation_task(
|
|
task: Any,
|
|
generation_task_repository: Any,
|
|
*,
|
|
user_id: str = "",
|
|
log_prefix: str = "[任务队列]",
|
|
log_task_status: bool = False,
|
|
user_pending_limit: int = 3,
|
|
global_pending_limit: int = 20,
|
|
) -> bool:
|
|
"""安全入队:限流检查 → 发送 Celery 任务 → 失败自动标记 failed。
|
|
|
|
Args:
|
|
task: 生成任务对象,需有 id 属性和 mark_failed 方法
|
|
generation_task_repository: 任务仓储,用于更新状态
|
|
user_id: 用户 ID,传了才做用户级限流检查
|
|
log_prefix: 日志前缀,便于区分调用来源
|
|
log_task_status: 成功日志中是否额外打印任务状态
|
|
user_pending_limit: 单用户 pending 上限,默认 3
|
|
global_pending_limit: 全局 pending 上限,默认 20
|
|
|
|
Returns:
|
|
True 表示入队成功,False 表示入队失败(已标记为 failed)
|
|
|
|
Raises:
|
|
GlobalQueueFull: 全局 pending 超限时抛出,任务会被标记为 failed
|
|
UserPendingLimitExceeded: 用户 pending 超限时抛出,任务会被标记为 failed
|
|
"""
|
|
# 全局限流检查(始终生效)
|
|
global_pending = generation_task_repository.count_pending_total()
|
|
if global_pending > global_pending_limit:
|
|
logger.warning(
|
|
"[队列限流] 全局 pending 任务数超限: %d/%d, user_id=%s",
|
|
global_pending,
|
|
global_pending_limit,
|
|
user_id or "unknown",
|
|
)
|
|
exc = GlobalQueueFull(pending_count=global_pending, limit=global_pending_limit)
|
|
_mark_task_failed_safely(task, generation_task_repository, log_prefix, str(exc))
|
|
raise exc
|
|
|
|
# 用户级限流检查(传了 user_id 才做)
|
|
if user_id:
|
|
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
|
if user_pending > user_pending_limit:
|
|
logger.warning(
|
|
"[队列限流] 用户 pending 任务数超限: user_id=%s, count=%d/%d",
|
|
user_id,
|
|
user_pending,
|
|
user_pending_limit,
|
|
)
|
|
exc = UserPendingLimitExceeded(
|
|
user_id=user_id, pending_count=user_pending, limit=user_pending_limit
|
|
)
|
|
_mark_task_failed_safely(task, generation_task_repository, log_prefix, str(exc))
|
|
raise exc
|
|
|
|
try:
|
|
celery_app.send_task("worker.generate_video", args=[task.id])
|
|
if log_task_status:
|
|
logger.info(
|
|
"%s 入队成功: task_id=%s, status=%s",
|
|
log_prefix,
|
|
task.id,
|
|
task.status,
|
|
)
|
|
else:
|
|
logger.info("%s 入队成功: task_id=%s", log_prefix, task.id)
|
|
return True
|
|
except Exception as e:
|
|
logger.error(
|
|
"%s 入队失败,标记为失败: task_id=%s error=%s",
|
|
log_prefix,
|
|
task.id,
|
|
e,
|
|
exc_info=True,
|
|
)
|
|
try:
|
|
task.mark_failed(f"任务入队失败: {e}")
|
|
generation_task_repository.update(task)
|
|
except Exception as update_err:
|
|
logger.error(
|
|
"%s 入队失败后更新状态也失败: task_id=%s error=%s",
|
|
log_prefix,
|
|
task.id,
|
|
update_err,
|
|
exc_info=True,
|
|
)
|
|
return False
|