"""task_enqueue 单测 — 队列限流 + 安全入队逻辑 (#2098).""" from __future__ import annotations from unittest.mock import MagicMock, patch import pytest 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, global_count: int = 0, user_count: int = 0): self._global = global_count self._user = user_count self.update_called = 0 def count_pending_total(self) -> int: return self._global def count_pending_by_user(self, user_id: str) -> int: return self._user def update(self, task): self.update_called += 1 def make_mock_task(task_id: str = "task-1"): task = MagicMock() task.id = task_id task.status = "pending" task.mark_failed = MagicMock() return task class TestCheckQueueLimits: def test_below_limits_passes(self): repo = MockRepository(global_count=5, user_count=1) check_queue_limits("user-1", repo) def test_global_at_limit_raises(self): repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT, user_count=1) with pytest.raises(GlobalQueueFull) as exc_info: check_queue_limits("user-1", repo) assert exc_info.value.pending_count == GLOBAL_PENDING_LIMIT assert exc_info.value.limit == GLOBAL_PENDING_LIMIT def test_global_over_limit_raises(self): repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT + 1, user_count=1) with pytest.raises(GlobalQueueFull): check_queue_limits("user-1", repo) def test_user_at_limit_no_longer_raises(self): repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT) check_queue_limits("user-1", repo) def test_user_over_limit_no_longer_raises(self): repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 100) check_queue_limits("user-1", repo) def test_empty_user_id_skips_user_check(self): repo = MockRepository(global_count=5, user_count=999) check_queue_limits("", repo) def test_custom_global_limit_still_honored(self): repo = MockRepository(global_count=15, user_count=999) with pytest.raises(GlobalQueueFull): check_queue_limits("u1", repo, user_pending_limit=999, global_pending_limit=10) check_queue_limits("u1", repo, user_pending_limit=999, global_pending_limit=20) class TestSafeEnqueueGenerationTask: @patch("app.core.task_enqueue.celery_app") def test_success_path(self, mock_celery): repo = MockRepository(global_count=1, user_count=1) task = make_mock_task() result = safe_enqueue_generation_task(task, repo, user_id="user-1") assert result is True mock_celery.send_task.assert_called_once_with("worker.generate_video", args=[task.id]) task.mark_failed.assert_not_called() @patch("app.core.task_enqueue.celery_app") def test_no_user_id_skips_user_check(self, mock_celery): repo = MockRepository(global_count=1, user_count=999) task = make_mock_task() result = safe_enqueue_generation_task(task, repo, user_id="") assert result is True @patch("app.core.task_enqueue.celery_app") def test_precheck_global_over_marks_failed(self, mock_celery): repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT + 1, user_count=0) task = make_mock_task() with pytest.raises(GlobalQueueFull): safe_enqueue_generation_task(task, repo, user_id="user-1") task.mark_failed.assert_called_once() mock_celery.send_task.assert_not_called() assert repo.update_called == 1 @patch("app.core.task_enqueue.celery_app") def test_precheck_user_over_still_enqueues(self, mock_celery, caplog): import logging caplog.set_level(logging.WARNING) repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT + 100) task = make_mock_task() result = safe_enqueue_generation_task(task, repo, user_id="user-1") assert result is True mock_celery.send_task.assert_called_once() task.mark_failed.assert_not_called() assert any("超过软上限" in r.message for r in caplog.records) @patch("app.core.task_enqueue.celery_app") def test_celery_send_false_returns_false(self, mock_celery): repo = MockRepository(global_count=1, user_count=1) task = make_mock_task() mock_celery.send_task.side_effect = Exception("celery down") result = safe_enqueue_generation_task(task, repo, user_id="user-1") assert result is False task.mark_failed.assert_called_once() assert "入队失败" in task.mark_failed.call_args[0][0] @patch("app.core.task_enqueue.celery_app") def test_celery_send_failure_update_also_fails(self, mock_celery): repo = MockRepository(global_count=1, user_count=1) repo.update = MagicMock(side_effect=Exception("db down")) task = make_mock_task() mock_celery.send_task.side_effect = Exception("celery down") result = safe_enqueue_generation_task(task, repo, user_id="user-1") assert result is False @patch("app.core.task_enqueue.celery_app") def test_postcheck_global_over_rollback(self, mock_celery): call_count = [0] def count_pending_total_side_effect(): call_count[0] += 1 if call_count[0] == 1: return GLOBAL_PENDING_LIMIT return GLOBAL_PENDING_LIMIT + 1 repo = MockRepository(global_count=GLOBAL_PENDING_LIMIT, user_count=0) repo.count_pending_total = MagicMock(side_effect=count_pending_total_side_effect) task = make_mock_task() with pytest.raises(GlobalQueueFull): safe_enqueue_generation_task(task, repo, user_id="user-1") task.mark_failed.assert_called_once() assert "入队后" in task.mark_failed.call_args[0][0] mock_celery.send_task.assert_called_once() @patch("app.core.task_enqueue.celery_app") def test_postcheck_user_over_does_not_rollback(self, mock_celery, caplog): import logging caplog.set_level(logging.WARNING) repo = MockRepository(global_count=5, user_count=USER_PENDING_LIMIT) call_count = [0] def count_by_user_side_effect(user_id): call_count[0] += 1 if call_count[0] <= 1: return USER_PENDING_LIMIT return USER_PENDING_LIMIT + 1 repo.count_pending_by_user = MagicMock(side_effect=count_by_user_side_effect) task = make_mock_task() result = safe_enqueue_generation_task(task, repo, user_id="user-1") assert result is True mock_celery.send_task.assert_called_once() task.mark_failed.assert_not_called() assert any("超软上限(入队后)" in r.message for r in caplog.records) @patch("app.core.task_enqueue.celery_app") def test_log_task_status_enabled(self, mock_celery): repo = MockRepository(global_count=1, user_count=1) task = make_mock_task() result = safe_enqueue_generation_task(task, repo, user_id="user-1", log_task_status=True) assert result is True @patch("app.core.task_enqueue.celery_app") def test_custom_global_limit_in_enqueue(self, mock_celery): repo = MockRepository(global_count=15, user_count=999) task = make_mock_task() result = safe_enqueue_generation_task(task, repo, user_id="user-1") assert result is True task2 = make_mock_task("task-2") mock_celery.reset_mock() with pytest.raises(GlobalQueueFull): safe_enqueue_generation_task(task2, repo, user_id="user-1", user_pending_limit=999, global_pending_limit=10) class TestExceptionClasses: def test_user_pending_limit_message(self): exc = UserPendingLimitExceeded("u1", 5, 3) assert "u1" in str(exc) assert "5" in str(exc) assert "3" in str(exc) def test_global_queue_full_message(self): exc = GlobalQueueFull(25, 20) assert "25" in str(exc) assert "20" in str(exc) def test_user_pending_limit_constant(self): assert USER_PENDING_LIMIT == 20