ad671d94c5
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m25s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m56s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 6m40s
CI/CD Pipeline / Staging E2E Tests (push) Successful in 4m30s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 5m35s
fix(voice-clone): source_audio_url 不做预签名转换,原样返回用户输入
352 lines
14 KiB
Python
352 lines
14 KiB
Python
"""任务队列限流防护单元测试。"""
|
|
|
|
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,
|
|
)
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# 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 # 用户级不检查,全局没超限 → 通过
|