from concurrent.futures import ThreadPoolExecutor from server.task_curriculum import ( TaskCurriculumScheduler, select_reset_task, should_enable_task_curriculum, ) class ReverseShuffle: def __init__(self) -> None: self.calls = 0 def shuffle(self, values): self.calls += 1 values.reverse() def test_scheduler_visits_every_task_once_before_repetition(): tasks = ("a", "b", "c") scheduler = TaskCurriculumScheduler(tasks=tasks, rng=ReverseShuffle()) first_cycle = [scheduler.next_task() for _ in tasks] assert sorted(first_cycle) == sorted(tasks) next_task = scheduler.next_task() assert next_task in tasks assert scheduler.cycle == 2 def test_scheduler_reshuffles_after_complete_cycle(): tasks = ("a", "b", "c") rng = ReverseShuffle() scheduler = TaskCurriculumScheduler(tasks=tasks, rng=rng) assert rng.calls == 1 for _ in tasks: scheduler.next_task() scheduler.next_task() assert rng.calls == 2 def test_curriculum_enabled_non_miniwob_does_not_sample(): scheduler = TaskCurriculumScheduler(tasks=("scheduled",), rng=ReverseShuffle()) scheduler_before = scheduler.cycle enabled = should_enable_task_curriculum("webarena", enabled=True) selected_task = select_reset_task( current_task="0", scheduler=scheduler if enabled else None, ) assert enabled is False assert selected_task == "0" assert scheduler.cycle == scheduler_before def test_explicit_reset_task_overrides_scheduler_selection(): scheduler = TaskCurriculumScheduler(tasks=("scheduled",), rng=ReverseShuffle()) scheduler_before = scheduler.cycle selected_task = select_reset_task( current_task="click-test", scheduler=scheduler, explicit_task="enter-text", ) assert selected_task == "enter-text" assert scheduler.cycle == scheduler_before def test_scheduler_is_thread_safe_across_complete_cycles(): tasks = tuple(str(index) for index in range(10)) scheduler = TaskCurriculumScheduler(tasks=tasks, rng=ReverseShuffle()) with ThreadPoolExecutor(max_workers=20) as executor: results = list(executor.map(lambda _: scheduler.next_task(), range(100))) assert sorted(results) == sorted(tasks * 10)