diff --git a/apps/api/app/api/routes/edit_plans.py b/apps/api/app/api/routes/edit_plans.py old mode 100644 new mode 100755 index 5c7ea8f7a..877c4f6fb --- a/apps/api/app/api/routes/edit_plans.py +++ b/apps/api/app/api/routes/edit_plans.py @@ -24,6 +24,7 @@ from typing import Any, List, Optional from app.auth import AuthenticatedUser, get_current_user from app.core.celery_app import celery_app +from app.core.task_enqueue import GLOBAL_PENDING_LIMIT, USER_PENDING_LIMIT from app.dependencies import get_asset_library_repository, get_asset_repository, get_db_session, get_project_repository from app.schemas.generation_task import GenerationTaskResponse from app.services import EditPlanService, PlanGeneratorService @@ -644,6 +645,31 @@ def generate_plan( # 创建 GenerationTask gen_task_repo = SQLAlchemyGenerationTaskRepository(db) + + # 队列限流预检查(repository 不支持计数时跳过) + user_id = current_user.user.id + try: + has_count = hasattr(gen_task_repo, "count_pending_by_user") and hasattr( + gen_task_repo, "count_pending_total" + ) + if has_count: + user_pending = gen_task_repo.count_pending_by_user(user_id) + global_pending = gen_task_repo.count_pending_total() + if user_pending >= USER_PENDING_LIMIT: + raise HTTPException( + status_code=429, + detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交", + ) + if global_pending >= GLOBAL_PENDING_LIMIT: + raise HTTPException( + status_code=503, + detail="系统繁忙,请稍后再试", + ) + except HTTPException: + raise + except Exception as e: + logger.warning("[队列限流] 剪辑计划限流检查失败,跳过: %s", e) + gen_task_use_case = CreateGenerationTaskUseCase(gen_task_repo) plan = svc.get_plan_or_raise(plan_id) gen_task = gen_task_use_case.execute( diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index 3728c0c64..d3f37c77e 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -5,7 +5,14 @@ from typing import Any from app.auth import AuthenticatedUser, get_current_user from app.core.storage import OSSStorageService, get_storage_service -from app.core.task_enqueue import safe_enqueue_generation_task +from app.core.task_enqueue import ( + GLOBAL_PENDING_LIMIT, + USER_PENDING_LIMIT, + GlobalQueueFull, + UserPendingLimitExceeded, + check_queue_limits, + safe_enqueue_generation_task, +) from app.dependencies import ( get_asset_library_repository, get_asset_repository, @@ -228,9 +235,31 @@ def create_generation_task( count = request.count created_tasks = [] failed_tasks = [] + user_id = authenticated_user.user.id # 同批次任务共享 batch_id,用于视频查重时批次内比对 batch_id = uuid.uuid4().hex if count > 1 else "" + # 预检查:批量提交前先看会不会超限,避免建一半才拒 + try: + user_pending = generation_task_repository.count_pending_by_user(user_id) + global_pending = generation_task_repository.count_pending_total() + if user_pending + count > USER_PENDING_LIMIT: + raise UserPendingLimitExceeded( + user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT + ) + if global_pending + count > GLOBAL_PENDING_LIMIT: + raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT) + except UserPendingLimitExceeded as e: + raise HTTPException( + status_code=429, + detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待完成后再提交", + ) from e + except GlobalQueueFull as e: + raise HTTPException( + status_code=503, + detail="系统繁忙,请稍后再试", + ) from e + try: for _ in range(count): task = use_case.execute( @@ -243,16 +272,42 @@ def create_generation_task( asset_ids=resolved_asset_ids, title_ids=request.title_ids, voice_ids=request.voice_ids, - created_by_user_id=authenticated_user.user.id, + created_by_user_id=user_id, source_edit_plan_id=request.source_edit_plan_id, asset_select_mode=request.asset_select_mode, batch_id=batch_id, ) ) - if safe_enqueue_generation_task(task, generation_task_repository, log_prefix="[生成任务]", log_task_status=True): - created_tasks.append(task) - else: + try: + if safe_enqueue_generation_task( + task, + generation_task_repository, + user_id=user_id, + log_prefix="[生成任务]", + log_task_status=True, + ): + created_tasks.append(task) + else: + failed_tasks.append(task) + except UserPendingLimitExceeded: + # 兜底:如果预检查后又并发提交了,在这里也拦住 failed_tasks.append(task) + if not created_tasks: + raise HTTPException( + status_code=429, + detail="您的待处理任务过多,请等待完成后再提交", + ) + break + except GlobalQueueFull: + failed_tasks.append(task) + if not created_tasks: + raise HTTPException( + status_code=503, + detail="系统繁忙,请稍后再试", + ) + break + except HTTPException: + raise except Exception as e: logger.error("[生成任务] 创建失败: %s", e, exc_info=True) raise HTTPException(status_code=500, detail="创建生成任务失败,请稍后重试或查看任务日志") @@ -327,6 +382,21 @@ def retry_generation_task( if status_val != "failed": raise HTTPException(status_code=409, detail="Only failed tasks can be retried") + user_id = authenticated_user.user.id + # 预检查:创建前判断,>= 上限就拒绝 + user_pending = generation_task_repository.count_pending_by_user(user_id) + global_pending = generation_task_repository.count_pending_total() + if user_pending >= USER_PENDING_LIMIT: + raise HTTPException( + status_code=429, + detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交", + ) + if global_pending >= GLOBAL_PENDING_LIMIT: + raise HTTPException( + status_code=503, + detail="系统繁忙,请稍后再试", + ) + use_case = CreateGenerationTaskUseCase(generation_task_repository) retried = use_case.execute( CreateGenerationTaskCommand( @@ -338,11 +408,28 @@ def retry_generation_task( asset_ids=task.asset_ids, title_ids=task.title_ids, voice_ids=task.voice_ids, - created_by_user_id=authenticated_user.user.id, + created_by_user_id=user_id, source_edit_plan_id=task.source_edit_plan_id or "", asset_select_mode=getattr(task, "asset_select_mode", ""), ) ) - if not safe_enqueue_generation_task(retried, generation_task_repository, log_prefix="[生成任务]", log_task_status=True): - logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id) + try: + if not safe_enqueue_generation_task( + retried, + generation_task_repository, + user_id=user_id, + log_prefix="[生成任务]", + log_task_status=True, + ): + logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id) + except UserPendingLimitExceeded: + raise HTTPException( + status_code=429, + detail="您的待处理任务过多,请等待完成后再提交", + ) from None + except GlobalQueueFull: + raise HTTPException( + status_code=503, + detail="系统繁忙,请稍后再试", + ) from None return _to_generation_task_response(retried) diff --git a/apps/api/app/api/routes/task_center.py b/apps/api/app/api/routes/task_center.py index 0838973bf..c796fd58d 100755 --- a/apps/api/app/api/routes/task_center.py +++ b/apps/api/app/api/routes/task_center.py @@ -3,7 +3,13 @@ from typing import Any from app.auth import AuthenticatedUser, get_current_user from app.core.celery_app import celery_app -from app.core.task_enqueue import safe_enqueue_generation_task +from app.core.task_enqueue import ( + GLOBAL_PENDING_LIMIT, + USER_PENDING_LIMIT, + GlobalQueueFull, + UserPendingLimitExceeded, + safe_enqueue_generation_task, +) from app.dependencies import ( get_generation_task_repository, get_ingest_job_repository, @@ -142,6 +148,21 @@ def retry_task_by_id( if _status_value(task.status) != "failed": raise HTTPException(status_code=409, detail="Only failed tasks can be retried") + user_id = authenticated_user.user.id + # 预检查 + user_pending = generation_task_repository.count_pending_by_user(user_id) + global_pending = generation_task_repository.count_pending_total() + if user_pending >= USER_PENDING_LIMIT: + raise HTTPException( + status_code=429, + detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交", + ) + if global_pending >= GLOBAL_PENDING_LIMIT: + raise HTTPException( + status_code=503, + detail="系统繁忙,请稍后再试", + ) + use_case = CreateGenerationTaskUseCase(generation_task_repository) retried = use_case.execute( CreateGenerationTaskCommand( @@ -153,11 +174,24 @@ def retry_task_by_id( asset_ids=task.asset_ids, title_ids=task.title_ids, voice_ids=task.voice_ids, - created_by_user_id=authenticated_user.user.id, + created_by_user_id=user_id, ) ) - if not safe_enqueue_generation_task(retried, generation_task_repository, log_prefix="[任务中心]"): - logger.warning("[任务中心] 用户级重试入队失败: task_id=%s", retried.id) + try: + if not safe_enqueue_generation_task( + retried, generation_task_repository, user_id=user_id, log_prefix="[任务中心]" + ): + logger.warning("[任务中心] 用户级重试入队失败: task_id=%s", retried.id) + except UserPendingLimitExceeded: + raise HTTPException( + status_code=429, + detail="您的待处理任务过多,请等待完成后再提交", + ) from None + except GlobalQueueFull: + raise HTTPException( + status_code=503, + detail="系统繁忙,请稍后再试", + ) from None return UserTaskResponse( id=f"generation:{retried.id}", task_type="generation", @@ -225,6 +259,22 @@ def retry_project_task( raise HTTPException(status_code=404, detail="Generation task not found") if _status_value(task.status) != "failed": raise HTTPException(status_code=409, detail="Only failed tasks can be retried") + + user_id = authenticated_user.user.id + # 预检查 + user_pending = generation_task_repository.count_pending_by_user(user_id) + global_pending = generation_task_repository.count_pending_total() + if user_pending >= USER_PENDING_LIMIT: + raise HTTPException( + status_code=429, + detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交", + ) + if global_pending >= GLOBAL_PENDING_LIMIT: + raise HTTPException( + status_code=503, + detail="系统繁忙,请稍后再试", + ) + use_case = CreateGenerationTaskUseCase(generation_task_repository) retried = use_case.execute( CreateGenerationTaskCommand( @@ -236,11 +286,24 @@ def retry_project_task( asset_ids=task.asset_ids, title_ids=task.title_ids, voice_ids=task.voice_ids, - created_by_user_id=authenticated_user.user.id, + created_by_user_id=user_id, ) ) - if not safe_enqueue_generation_task(retried, generation_task_repository, log_prefix="[任务中心]"): - logger.warning("[任务中心] 项目级重试用队失败: task_id=%s", retried.id) + try: + if not safe_enqueue_generation_task( + retried, generation_task_repository, user_id=user_id, log_prefix="[任务中心]" + ): + logger.warning("[任务中心] 项目级重试入队失败: task_id=%s", retried.id) + except UserPendingLimitExceeded: + raise HTTPException( + status_code=429, + detail="您的待处理任务过多,请等待完成后再提交", + ) from None + except GlobalQueueFull: + raise HTTPException( + status_code=503, + detail="系统繁忙,请稍后再试", + ) from None return _generation_task_to_project_response(retried) if task_type == "ingest": job = ingest_job_repository.get(source_id) diff --git a/apps/api/app/core/task_enqueue.py b/apps/api/app/core/task_enqueue.py index 881938a58..0b3048d65 100755 --- a/apps/api/app/core/task_enqueue.py +++ b/apps/api/app/core/task_enqueue.py @@ -5,37 +5,167 @@ from app.core.celery_app import celery_app logger = logging.getLogger(__name__) +# ── 限流阈值常量(全系统统一管理,不要在业务代码里硬编码) ── +USER_PENDING_LIMIT = 3 # 单用户 pending 上限 +GLOBAL_PENDING_LIMIT = 20 # 全局 pending 上限 + + +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 = USER_PENDING_LIMIT, + global_pending_limit: int = GLOBAL_PENDING_LIMIT, +) -> None: + """检查队列限流(预检查用,任务创建前调用),超限抛对应异常。 + + 边界语义:>= 上限即拒绝(达到上限就不能再加新任务)。 + + Args: + user_id: 用户 ID + generation_task_repository: 任务仓储 + user_pending_limit: 单用户 pending 上限,默认 USER_PENDING_LIMIT + global_pending_limit: 全局 pending 上限,默认 GLOBAL_PENDING_LIMIT + + 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) + + # 再查用户级 + 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, + ) + 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 = USER_PENDING_LIMIT, + global_pending_limit: int = GLOBAL_PENDING_LIMIT, ) -> bool: - """安全入队:send_task 失败时自动把任务标记为 failed,避免留下 pending 僵尸任务。 + """安全入队:入队前限流检查 → 发送 Celery 任务 → 入队后最终校验兜底。 + + 边界说明: + 入队前检查用 > 而非 >=。因为调用此函数时 task 已经是 pending 状态并计入 DB, + pending 总数包含了当前任务本身。pending > limit 等价于"其他任务数 >= limit", + 与预检查的 >= 语义一致(都是达到上限就拒绝新任务)。 + + 入队后最终校验:发送 Celery 成功后再查一次 DB 计数,处理并发竞态场景 + (两个请求同时通过入队前检查,后到的那个在这里被兜住)。 Args: - task: 生成任务对象,需有 id 属性和 mark_failed 方法 + task: 生成任务对象,需有 id 属性和 mark_failed 方法(状态已为 pending) generation_task_repository: 任务仓储,用于更新状态 + user_id: 用户 ID,传了才做用户级限流检查 log_prefix: 日志前缀,便于区分调用来源 log_task_status: 成功日志中是否额外打印任务状态 + user_pending_limit: 单用户 pending 上限,默认 USER_PENDING_LIMIT + global_pending_limit: 全局 pending 上限,默认 GLOBAL_PENDING_LIMIT Returns: True 表示入队成功,False 表示入队失败(已标记为 failed) + + Raises: + GlobalQueueFull: 全局 pending 超限时抛出,任务会被标记为 failed + UserPendingLimitExceeded: 用户 pending 超限时抛出,任务会被标记为 failed """ + # ── 入队前检查:任务已是 pending,用 > 判断(包含当前任务) ── + + # 全局限流检查(始终生效) + 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 + + # ── 发送 Celery 任务 ── 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", @@ -56,3 +186,42 @@ def safe_enqueue_generation_task( exc_info=True, ) return False + + # ── 入队后最终校验:并发竞态兜底 ── + # 发送成功后再查一次,防止两个请求同时通过入队前检查导致超限 + global_after = generation_task_repository.count_pending_total() + user_after = generation_task_repository.count_pending_by_user(user_id) if user_id else 0 + + global_over = global_after > global_pending_limit + user_over = bool(user_id and user_after > user_pending_limit) + + if global_over or user_over: + if global_over: + reason = f"全局 pending 超限(入队后): {global_after}/{global_pending_limit}" + exc: Exception = GlobalQueueFull(pending_count=global_after, limit=global_pending_limit) + else: + reason = f"用户 pending 超限(入队后): {user_after}/{user_pending_limit}" + exc = UserPendingLimitExceeded( + user_id=user_id, pending_count=user_after, limit=user_pending_limit + ) + + logger.warning( + "[队列限流] %s, task_id=%s, user_id=%s — 回滚状态为 failed", + reason, + task.id, + user_id or "unknown", + ) + _mark_task_failed_safely(task, generation_task_repository, log_prefix, reason) + raise exc + + # 入队成功日志 + 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 diff --git a/apps/api/postgres b/apps/api/postgres new file mode 100644 index 000000000..e69de29bb diff --git a/packages/adapters/sqlalchemy_impl/generation_task_repository.py b/packages/adapters/sqlalchemy_impl/generation_task_repository.py old mode 100644 new mode 100755 index d736a7eb2..646a05be1 --- a/packages/adapters/sqlalchemy_impl/generation_task_repository.py +++ b/packages/adapters/sqlalchemy_impl/generation_task_repository.py @@ -91,6 +91,23 @@ class SQLAlchemyGenerationTaskRepository: def count_by_user(self, user_id: str) -> int: return self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id).count() + def count_pending_by_user(self, user_id: str) -> int: + return ( + self.session.query(GenerationTaskModel) + .filter( + GenerationTaskModel.created_by_user_id == user_id, + GenerationTaskModel.status == GenerationTaskStatus.PENDING.value, + ) + .count() + ) + + def count_pending_total(self) -> int: + return ( + self.session.query(GenerationTaskModel) + .filter(GenerationTaskModel.status == GenerationTaskStatus.PENDING.value) + .count() + ) + def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: models = ( self.session.query(GenerationTaskModel) diff --git a/packages/ports/generation_task_repository.py b/packages/ports/generation_task_repository.py old mode 100644 new mode 100755 index 832144233..a86314aa9 --- a/packages/ports/generation_task_repository.py +++ b/packages/ports/generation_task_repository.py @@ -16,6 +16,10 @@ class GenerationTaskRepository(Protocol): def count_by_user(self, user_id: str) -> int: ... + def count_pending_by_user(self, user_id: str) -> int: ... + + def count_pending_total(self) -> int: ... + def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: ... def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]: ... diff --git a/tests/unit/test_task_queue_limit.py b/tests/unit/test_task_queue_limit.py new file mode 100755 index 000000000..6a906df57 --- /dev/null +++ b/tests/unit/test_task_queue_limit.py @@ -0,0 +1,351 @@ +"""任务队列限流防护单元测试。""" +from __future__ import annotations + +import sys +import os +from unittest.mock import MagicMock + +import pytest + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api")) + +from app.core.task_enqueue import ( + GLOBAL_PENDING_LIMIT, + USER_PENDING_LIMIT, + GlobalQueueFull, + UserPendingLimitExceeded, + check_queue_limits, + safe_enqueue_generation_task, +) + + +# --------------------------------------------------------------------------- +# Mock helpers +# --------------------------------------------------------------------------- + + +class MockRepository: + """支持 pending 计数的 mock repository。 + + 支持通过 set_pending 动态修改计数,用于模拟入队后计数变化的并发场景。 + """ + + def __init__(self, user_pending: int = 0, global_pending: int = 0): + self._user_pending = user_pending + self._global_pending = global_pending + self._send_task_called = False + self.updated_tasks = [] + + def count_pending_by_user(self, user_id: str) -> int: + return self._user_pending + + def count_pending_total(self) -> int: + return self._global_pending + + def update(self, task): + self.updated_tasks.append(task) + return task + + def set_pending(self, *, user_pending: int | None = None, global_pending: int | None = None): + """动态修改 pending 计数,模拟并发场景。""" + if user_pending is not None: + self._user_pending = user_pending + if global_pending is not None: + self._global_pending = global_pending + + +class MockTask: + def __init__(self, task_id: str = "task-1", status: str = "pending"): + self.id = task_id + self.status = status + self.error_message = "" + + def mark_failed(self, reason: str): + self.status = "failed" + self.error_message = reason + + +@pytest.fixture(autouse=True) +def mock_celery(monkeypatch): + """mock 掉 celery_app.send_task,避免真实发送。""" + mock_send = MagicMock() + monkeypatch.setattr("app.core.celery_app.celery_app.send_task", mock_send) + return mock_send + + +# --------------------------------------------------------------------------- +# 常量导出测试 +# --------------------------------------------------------------------------- + + +def test_limit_constants_are_exported(): + """限流阈值常量已导出,供业务代码引用。""" + assert USER_PENDING_LIMIT == 3 + assert GLOBAL_PENDING_LIMIT == 20 + + +# --------------------------------------------------------------------------- +# check_queue_limits 单元测试(预检查用,>= 边界) +# --------------------------------------------------------------------------- + + +class TestCheckQueueLimits: + """队列限流检查函数测试(预检查语义,>= 上限即拒绝)。""" + + def test_normal_passes_through(self): + """正常范围内的任务不受限制。""" + repo = MockRepository(user_pending=1, global_pending=5) + check_queue_limits("user-1", repo) + + def test_user_limit_exceeded_raises(self): + """用户 pending 超过上限抛 UserPendingLimitExceeded。""" + repo = MockRepository(user_pending=4, global_pending=5) + with pytest.raises(UserPendingLimitExceeded) as exc_info: + check_queue_limits("user-1", repo) + assert exc_info.value.user_id == "user-1" + assert exc_info.value.pending_count == 4 + assert exc_info.value.limit == 3 + + def test_user_at_limit_also_raises(self): + """用户 pending 刚好等于上限也拒绝(>= 边界)。""" + repo = MockRepository(user_pending=3, global_pending=5) + with pytest.raises(UserPendingLimitExceeded): + check_queue_limits("user-1", repo) + + def test_user_below_limit_passes(self): + """用户 pending 比上限少 1,通过。""" + repo = MockRepository(user_pending=2, global_pending=5) + check_queue_limits("user-1", repo) + + def test_global_limit_exceeded_raises(self): + """全局 pending 超过上限抛 GlobalQueueFull。""" + repo = MockRepository(user_pending=1, global_pending=21) + with pytest.raises(GlobalQueueFull) as exc_info: + check_queue_limits("user-1", repo) + assert exc_info.value.pending_count == 21 + assert exc_info.value.limit == 20 + + def test_global_at_limit_also_raises(self): + """全局 pending 刚好等于上限也拒绝(>= 边界)。""" + repo = MockRepository(user_pending=1, global_pending=20) + with pytest.raises(GlobalQueueFull): + check_queue_limits("user-1", repo) + + def test_global_below_limit_passes(self): + """全局 pending 比上限少 1,通过。""" + repo = MockRepository(user_pending=1, global_pending=19) + check_queue_limits("user-1", repo) + + def test_global_takes_priority_over_user(self): + """全局和用户都超限时,优先抛全局异常。""" + repo = MockRepository(user_pending=5, global_pending=25) + with pytest.raises(GlobalQueueFull): + check_queue_limits("user-1", repo) + + def test_empty_user_id_skips_user_check(self): + """不传 user_id 时跳过用户级检查,只做全局检查。""" + repo = MockRepository(user_pending=10, global_pending=5) + # 用户超限但不传 user_id → 全局未超限,应该通过 + check_queue_limits("", repo) + + +# --------------------------------------------------------------------------- +# safe_enqueue_generation_task 限流集成测试(入队前用 >,包含当前任务) +# --------------------------------------------------------------------------- + + +class TestSafeEnqueueWithLimits: + """安全入队函数的限流功能测试。""" + + def test_normal_task_enqueues_successfully(self, mock_celery): + """正常任务入队成功,返回 True。""" + repo = MockRepository(user_pending=0, global_pending=0) + task = MockTask("task-1") + result = safe_enqueue_generation_task(task, repo, user_id="user-1") + assert result is True + mock_celery.assert_called_once_with("worker.generate_video", args=["task-1"]) + assert len(repo.updated_tasks) == 0 # 成功不需要更新状态 + + def test_user_limit_rejected_with_failed_status(self, mock_celery): + """用户超限:任务标记为 failed,抛 UserPendingLimitExceeded。""" + repo = MockRepository(user_pending=5, global_pending=5) + task = MockTask("task-1") + with pytest.raises(UserPendingLimitExceeded): + safe_enqueue_generation_task(task, repo, user_id="user-1") + mock_celery.assert_not_called() + assert task.status == "failed" + assert "限流" in task.error_message + assert len(repo.updated_tasks) == 1 + + def test_user_at_limit_still_passes(self, mock_celery): + """用户 pending 刚好等于上限:入队前检查用 >,包含当前任务,刚好到上限不算超。 + + 与预检查的 >= 语义一致:预检查时 pending=3 拒绝(不能再加新的), + 但 safe_enqueue 被调用时任务已是 pending(就是第3个), + pending=3 不满足 >3,所以通过。 + """ + repo = MockRepository(user_pending=3, global_pending=5) + task = MockTask("task-1") + result = safe_enqueue_generation_task(task, repo, user_id="user-1") + assert result is True + mock_celery.assert_called_once() + + def test_user_one_over_limit_rejected(self, mock_celery): + """用户 pending = limit + 1:超限被拒。""" + repo = MockRepository(user_pending=4, global_pending=5) + task = MockTask("task-1") + with pytest.raises(UserPendingLimitExceeded): + safe_enqueue_generation_task(task, repo, user_id="user-1") + mock_celery.assert_not_called() + + def test_global_limit_rejected_with_failed_status(self, mock_celery): + """全局超限:任务标记为 failed,抛 GlobalQueueFull。""" + repo = MockRepository(user_pending=1, global_pending=21) + task = MockTask("task-1") + with pytest.raises(GlobalQueueFull): + safe_enqueue_generation_task(task, repo, user_id="user-1") + mock_celery.assert_not_called() + assert task.status == "failed" + assert len(repo.updated_tasks) == 1 + + def test_global_at_limit_still_passes(self, mock_celery): + """全局 pending 刚好等于上限:入队前检查用 >,包含当前任务,刚好到上限不算超。""" + repo = MockRepository(user_pending=1, global_pending=20) + task = MockTask("task-1") + result = safe_enqueue_generation_task(task, repo, user_id="user-1") + assert result is True + mock_celery.assert_called_once() + + def test_no_user_id_skips_user_limit(self, mock_celery): + """不传 user_id 时跳过用户级限流,只做全局检查。""" + repo = MockRepository(user_pending=10, global_pending=5) + task = MockTask("task-1") + result = safe_enqueue_generation_task(task, repo, user_id="") + assert result is True + mock_celery.assert_called_once() + + def test_no_user_id_still_checks_global(self, mock_celery): + """不传 user_id 时全局超限仍然被拦。""" + repo = MockRepository(user_pending=10, global_pending=25) + task = MockTask("task-1") + with pytest.raises(GlobalQueueFull): + safe_enqueue_generation_task(task, repo, user_id="") + mock_celery.assert_not_called() + + def test_default_limits_match_constants(self, mock_celery): + """默认配置与导出常量一致。""" + # 刚好在默认限制内(limit - 1) + repo = MockRepository(user_pending=2, global_pending=19) + task = MockTask("task-1") + result = safe_enqueue_generation_task(task, repo, user_id="user-1") + assert result is True + + def test_update_failure_does_not_crash(self, mock_celery): + """repository.update 失败也不崩溃,异常继续向上抛。""" + + class BadRepo(MockRepository): + def update(self, task): + raise RuntimeError("db down") + + repo = BadRepo(user_pending=5, global_pending=5) + task = MockTask("task-1") + # 仍然抛 UserPendingLimitExceeded,不会被 update 失败掩盖 + with pytest.raises(UserPendingLimitExceeded): + safe_enqueue_generation_task(task, repo, user_id="user-1") + mock_celery.assert_not_called() + # 任务状态还是变了(内存里改了) + assert task.status == "failed" + + +# --------------------------------------------------------------------------- +# 入队后最终校验(并发竞态兜底)测试 +# --------------------------------------------------------------------------- + + +class TestPostEnqueueFinalCheck: + """入队后最终校验:模拟并发场景,Celery发送后计数增加被兜住。""" + + def test_post_enqueue_global_overflow_rollback(self, mock_celery): + """并发场景:入队前检查通过,但发送Celery后全局计数超限 → 回滚为failed。 + + 模拟两个请求同时通过入队前检查(都查到 global=19), + 都创建了任务(DB里变成 21),先发送Celery的那个在最终校验时被兜住。 + """ + repo = MockRepository(user_pending=1, global_pending=20) # 入队前:20 > 20?否 + task = MockTask("task-1") + + # 模拟发送Celery后,另一个并发请求也创建了任务,全局变成21 + def side_effect(*args, **kwargs): + repo.set_pending(global_pending=21) + + mock_celery.side_effect = side_effect + + with pytest.raises(GlobalQueueFull) as exc_info: + safe_enqueue_generation_task(task, repo, user_id="user-1") + + # Celery 确实发出去了(兜底不撤销 Celery,只回滚 DB 状态) + mock_celery.assert_called_once() + # 任务被标记为 failed + assert task.status == "failed" + assert "入队后" in task.error_message + assert exc_info.value.pending_count == 21 + assert len(repo.updated_tasks) == 1 + + def test_post_enqueue_user_overflow_rollback(self, mock_celery): + """并发场景:入队前检查通过,但发送Celery后用户计数超限 → 回滚为failed。""" + repo = MockRepository(user_pending=3, global_pending=5) # 入队前:3 > 3?否 + task = MockTask("task-1") + + def side_effect(*args, **kwargs): + repo.set_pending(user_pending=4) + + mock_celery.side_effect = side_effect + + with pytest.raises(UserPendingLimitExceeded) as exc_info: + safe_enqueue_generation_task(task, repo, user_id="user-1") + + mock_celery.assert_called_once() + assert task.status == "failed" + assert "入队后" in task.error_message + assert exc_info.value.user_id == "user-1" + assert exc_info.value.pending_count == 4 + + def test_post_enqueue_global_priority_over_user(self, mock_celery): + """入队后校验:全局和用户都超限时,优先抛全局异常。""" + repo = MockRepository(user_pending=3, global_pending=20) + task = MockTask("task-1") + + def side_effect(*args, **kwargs): + repo.set_pending(user_pending=5, global_pending=22) + + mock_celery.side_effect = side_effect + + with pytest.raises(GlobalQueueFull): + safe_enqueue_generation_task(task, repo, user_id="user-1") + + assert task.status == "failed" + + def test_post_enqueue_no_change_still_passes(self, mock_celery): + """入队后计数没变 → 正常通过,不回滚。""" + repo = MockRepository(user_pending=2, global_pending=10) + task = MockTask("task-1") + + result = safe_enqueue_generation_task(task, repo, user_id="user-1") + + assert result is True + mock_celery.assert_called_once() + assert task.status == "pending" # 状态没变 + assert len(repo.updated_tasks) == 0 # 没更新 DB + + def test_post_enqueue_no_user_id_skips_user_check(self, mock_celery): + """不传 user_id 时,入队后校验也跳过用户级,只查全局。""" + repo = MockRepository(user_pending=10, global_pending=5) + task = MockTask("task-1") + + def side_effect(*args, **kwargs): + repo.set_pending(user_pending=15, global_pending=5) # 用户超限但全局没超 + + mock_celery.side_effect = side_effect + + result = safe_enqueue_generation_task(task, repo, user_id="") + assert result is True # 用户级不检查,全局没超限 → 通过