221 lines
6.3 KiB
Python
221 lines
6.3 KiB
Python
"""Tests for AI outfit batch orchestration."""
|
|||
|
|
import sys
|
||
|
|
import threading
|
||
|
|
import time
|
||
|
|
import unittest
|
||
|
|
from pathlib import Path
|
||
|
|
|
||
|
|
sys.path.insert(0, str(Path(__file__).parent.parent / "src"))
|
||
|
|
|
||
|
|
|
||
|
|
def _task(row, product_id=None):
|
||
|
|
from core.models import OutfitTask
|
||
|
|
|
||
|
|
return OutfitTask(
|
||
|
|
row_index=row,
|
||
|
|
title="title{}".format(row),
|
||
|
|
product_id=product_id or "TY{:03d}".format(row),
|
||
|
|
garment_path="g{}.png".format(row),
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def _result(task, success=True, error=""):
|
||
|
|
from core.models import OutfitResult
|
||
|
|
|
||
|
|
return OutfitResult(task=task, success=success, error=error)
|
||
|
|
|
||
|
|
|
||
|
|
class TestRateLimiter(unittest.TestCase):
|
||
|
|
def test_wait_enforces_interval(self):
|
||
|
|
from core.outfit_batch import RateLimiter
|
||
|
|
|
||
|
|
now = [0.0]
|
||
|
|
sleeps = []
|
||
|
|
|
||
|
|
def fake_now():
|
||
|
|
return now[0]
|
||
|
|
|
||
|
|
def fake_sleep(seconds):
|
||
|
|
sleeps.append(seconds)
|
||
|
|
now[0] += seconds
|
||
|
|
|
||
|
|
limiter = RateLimiter(2.0, now_func=fake_now, sleep_func=fake_sleep)
|
||
|
|
limiter.wait()
|
||
|
|
limiter.wait()
|
||
|
|
|
||
|
|
self.assertEqual(sleeps, [2.0])
|
||
|
|
|
||
|
|
|
||
|
|
class TestOutfitBatchRunner(unittest.TestCase):
|
||
|
|
def test_run_success_summary_and_progress(self):
|
||
|
|
from core.outfit_batch import OutfitBatchOptions, OutfitBatchRunner
|
||
|
|
|
||
|
|
tasks = [_task(2), _task(3)]
|
||
|
|
progress = []
|
||
|
|
|
||
|
|
runner = OutfitBatchRunner(
|
||
|
|
tasks,
|
||
|
|
generate_func=lambda task: _result(task, True),
|
||
|
|
options=OutfitBatchOptions(concurrency=1),
|
||
|
|
progress_callback=lambda done, total, result: progress.append((done, total, result.task.row_index)),
|
||
|
|
)
|
||
|
|
|
||
|
|
summary = runner.run()
|
||
|
|
|
||
|
|
self.assertEqual(summary.total, 2)
|
||
|
|
self.assertEqual(summary.success_count, 2)
|
||
|
|
self.assertEqual(summary.failure_count, 0)
|
||
|
|
self.assertEqual(progress, [(1, 2, 2), (2, 2, 3)])
|
||
|
|
|
||
|
|
def test_retry_until_success(self):
|
||
|
|
from core.outfit_batch import OutfitBatchOptions, OutfitBatchRunner
|
||
|
|
|
||
|
|
calls = []
|
||
|
|
|
||
|
|
def generate(task):
|
||
|
|
calls.append(task.row_index)
|
||
|
|
if len(calls) == 1:
|
||
|
|
return _result(task, False, "temporary")
|
||
|
|
return _result(task, True)
|
||
|
|
|
||
|
|
runner = OutfitBatchRunner(
|
||
|
|
[_task(2)],
|
||
|
|
generate_func=generate,
|
||
|
|
options=OutfitBatchOptions(concurrency=1, retry_count=1),
|
||
|
|
sleep_func=lambda _seconds: None,
|
||
|
|
)
|
||
|
|
|
||
|
|
summary = runner.run()
|
||
|
|
|
||
|
|
self.assertEqual(len(calls), 2)
|
||
|
|
self.assertEqual(summary.success_count, 1)
|
||
|
|
self.assertEqual(summary.results[0].attempts, 2)
|
||
|
|
|
||
|
|
def test_concurrency_runs_more_than_one_task_at_once(self):
|
||
|
|
from core.outfit_batch import OutfitBatchOptions, OutfitBatchRunner
|
||
|
|
|
||
|
|
active = [0]
|
||
|
|
max_active = [0]
|
||
|
|
lock = threading.Lock()
|
||
|
|
|
||
|
|
def generate(task):
|
||
|
|
with lock:
|
||
|
|
active[0] += 1
|
||
|
|
max_active[0] = max(max_active[0], active[0])
|
||
|
|
time.sleep(0.05)
|
||
|
|
with lock:
|
||
|
|
active[0] -= 1
|
||
|
|
return _result(task, True)
|
||
|
|
|
||
|
|
runner = OutfitBatchRunner(
|
||
|
|
[_task(2), _task(3), _task(4), _task(5)],
|
||
|
|
generate_func=generate,
|
||
|
|
options=OutfitBatchOptions(concurrency=2),
|
||
|
|
)
|
||
|
|
|
||
|
|
summary = runner.run()
|
||
|
|
|
||
|
|
self.assertEqual(summary.success_count, 4)
|
||
|
|
self.assertGreaterEqual(max_active[0], 2)
|
||
|
|
|
||
|
|
def test_task_cooldown_is_applied(self):
|
||
|
|
from core.outfit_batch import OutfitBatchOptions, OutfitBatchRunner
|
||
|
|
|
||
|
|
sleeps = []
|
||
|
|
runner = OutfitBatchRunner(
|
||
|
|
[_task(2)],
|
||
|
|
generate_func=lambda task: _result(task, True),
|
||
|
|
options=OutfitBatchOptions(concurrency=1, task_cooldown=1.5),
|
||
|
|
sleep_func=sleeps.append,
|
||
|
|
)
|
||
|
|
|
||
|
|
runner.run()
|
||
|
|
|
||
|
|
self.assertEqual(sleeps, [1.5])
|
||
|
|
|
||
|
|
def test_failed_after_retries(self):
|
||
|
|
from core.outfit_batch import OutfitBatchOptions, OutfitBatchRunner
|
||
|
|
|
||
|
|
runner = OutfitBatchRunner(
|
||
|
|
[_task(2)],
|
||
|
|
generate_func=lambda task: _result(task, False, "boom"),
|
||
|
|
options=OutfitBatchOptions(concurrency=1, retry_count=2),
|
||
|
|
sleep_func=lambda _seconds: None,
|
||
|
|
)
|
||
|
|
|
||
|
|
summary = runner.run()
|
||
|
|
|
||
|
|
self.assertEqual(summary.failure_count, 1)
|
||
|
|
self.assertEqual(summary.results[0].attempts, 3)
|
||
|
|
|
||
|
|
def test_gentle_stop_allows_inflight_to_finish_without_new_tasks(self):
|
||
|
|
from core.outfit_batch import OutfitBatchOptions, OutfitBatchRunner
|
||
|
|
|
||
|
|
started = []
|
||
|
|
|
||
|
|
def generate(task):
|
||
|
|
started.append(task.row_index)
|
||
|
|
runner.stop()
|
||
|
|
return _result(task, True)
|
||
|
|
|
||
|
|
runner = OutfitBatchRunner(
|
||
|
|
[_task(2), _task(3), _task(4)],
|
||
|
|
generate_func=generate,
|
||
|
|
options=OutfitBatchOptions(concurrency=1),
|
||
|
|
)
|
||
|
|
|
||
|
|
summary = runner.run()
|
||
|
|
|
||
|
|
self.assertEqual(started, [2])
|
||
|
|
self.assertTrue(summary.stopped)
|
||
|
|
self.assertEqual(summary.total, 3)
|
||
|
|
self.assertEqual(len(summary.results), 1)
|
||
|
|
|
||
|
|
def test_heartbeat_log_for_long_running_task(self):
|
||
|
|
from core.outfit_batch import OutfitBatchOptions, OutfitBatchRunner
|
||
|
|
|
||
|
|
logs = []
|
||
|
|
task_started = threading.Event()
|
||
|
|
release = threading.Event()
|
||
|
|
|
||
|
|
def generate(task):
|
||
|
|
task_started.set()
|
||
|
|
release.wait(2)
|
||
|
|
return _result(task, True)
|
||
|
|
|
||
|
|
runner = OutfitBatchRunner(
|
||
|
|
[_task(2)],
|
||
|
|
generate_func=generate,
|
||
|
|
options=OutfitBatchOptions(
|
||
|
|
concurrency=1,
|
||
|
|
heartbeat_after=0.01,
|
||
|
|
heartbeat_interval=0.01,
|
||
|
|
),
|
||
|
|
log_callback=logs.append,
|
||
|
|
)
|
||
|
|
|
||
|
|
thread = threading.Thread(target=runner.run)
|
||
|
|
thread.start()
|
||
|
|
task_started.wait(1)
|
||
|
|
time.sleep(0.08)
|
||
|
|
release.set()
|
||
|
|
thread.join(1)
|
||
|
|
|
||
|
|
self.assertTrue(any("等待中转响应" in line for line in logs))
|
||
|
|
|
||
|
|
def test_retry_delay_rate_limit_errors_back_off_more(self):
|
||
|
|
from core.outfit_batch import retry_delay_seconds
|
||
|
|
|
||
|
|
self.assertGreater(retry_delay_seconds("429 rate limit", 2), retry_delay_seconds("boom", 2))
|
||
|
|
|
||
|
|
|
||
|
|
class TestAiOutfitWorkerImport(unittest.TestCase):
|
||
|
|
def test_worker_symbols_import(self):
|
||
|
|
from app.workers.ai_outfit_worker import AiOutfitWorker
|
||
|
|
|
||
|
|
self.assertIsNotNone(AiOutfitWorker)
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
unittest.main()
|