"""任务队列限流防护单元测试 (#2098: 用户级改为软上限,仅全局硬拒).""" from __future__ import annotations import os import sys 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, ) class MockRepository: 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): 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_send = MagicMock() monkeypatch.setattr("app.core.celery_app.celery_app.send_task", mock_send) return mock_send def test_limit_constants_are_exported(): """#2098: USER_PENDING_LIMIT 从 3 提到 20 作为软上限;GLOBAL_PENDING_LIMIT 保持 20 为硬上限。""" assert USER_PENDING_LIMIT == 20 assert GLOBAL_PENDING_LIMIT == 20 class TestCheckQueueLimits: """check_queue_limits 预检查:仅全局硬上限拒绝,用户级改为软提示。""" 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_no_longer_raises(self): """#2098: 用户 pending 超过软上限不再抛异常。""" repo = MockRepository(user_pending=100, global_pending=5) check_queue_limits("user-1", repo) # 不抛即通过 def test_user_at_limit_no_longer_raises(self): """#2098: 用户 pending 等于软上限也不拒绝。""" repo = MockRepository(user_pending=USER_PENDING_LIMIT, global_pending=5) check_queue_limits("user-1", repo) def test_user_below_limit_passes(self): repo = MockRepository(user_pending=2, global_pending=5) check_queue_limits("user-1", repo) def test_global_limit_exceeded_raises(self): 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): 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): repo = MockRepository(user_pending=1, global_pending=19) check_queue_limits("user-1", repo) def test_empty_user_id_skips_user_check(self): repo = MockRepository(user_pending=10, global_pending=5) check_queue_limits("", repo) class TestSafeEnqueueWithLimits: """safe_enqueue_generation_task:用户超限仅 warning 仍入队;全局超限硬拒。""" def test_normal_task_enqueues_successfully(self, mock_celery): 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) == 1 assert task.celery_task_id def test_user_limit_exceeded_still_enqueues(self, mock_celery, caplog): """#2098: 用户远超软上限仍入队,任务不被标记 failed。""" import logging caplog.set_level(logging.WARNING) repo = MockRepository(user_pending=100, 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() assert task.status == "pending" # 没被标记 failed assert any("超过软上限" in r.message for r in caplog.records) def test_user_at_limit_still_passes(self, mock_celery): repo = MockRepository(user_pending=USER_PENDING_LIMIT, 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_global_limit_rejected_with_failed_status(self, mock_celery): 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): 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): 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): 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): 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 class TestPostEnqueueFinalCheck: """入队后校验:仅全局超限回滚;用户超限仅 warning。""" def test_post_enqueue_global_overflow_rollback(self, mock_celery): repo = MockRepository(user_pending=1, global_pending=20) task = MockTask("task-1") 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") mock_celery.assert_called_once() 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_no_rollback(self, mock_celery, caplog): """#2098: 入队后用户超软上限仅 warning,不回滚。""" import logging caplog.set_level(logging.WARNING) repo = MockRepository(user_pending=USER_PENDING_LIMIT, global_pending=5) task = MockTask("task-1") def side_effect(*args, **kwargs): repo.set_pending(user_pending=USER_PENDING_LIMIT + 1) mock_celery.side_effect = side_effect 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 any("超软上限(入队后)" in r.message for r in caplog.records) 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) == 1 assert task.celery_task_id def test_post_enqueue_no_user_id_skips_user_check(self, mock_celery): 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 def test_user_pending_limit_exceeded_class_still_exists(): """UserPendingLimitExceeded 保留用于兼容历史 import/except(#2098 后不再主动 raise)。""" exc = UserPendingLimitExceeded("u1", 5, 3) assert exc.user_id == "u1" assert exc.pending_count == 5 assert exc.limit == 3