"""Task curriculum helpers for MiniWoB training.""" import random import threading from typing import List, Optional, Sequence MINIWOB_FORM_TASKS = ( "login-user", "enter-text", "enter-password", "enter-date", "enter-time", "choose-list", "use-autocomplete", "multi-layouts", "multi-orderings", "click-button", "click-widget", ) class TaskCurriculumScheduler: """Shuffle tasks by cycle and return each task once before repeating.""" def __init__( self, tasks: Sequence[str] = MINIWOB_FORM_TASKS, rng: Optional[random.Random] = None, ) -> None: if not tasks: raise ValueError("Task curriculum requires at least one task") self._tasks = tuple(tasks) self._rng = rng or random.Random() self._queue: List[str] = [] self.cycle = 0 self._lock = threading.RLock() self._reshuffle() @property def tasks(self) -> tuple[str, ...]: """Return the configured task list.""" return self._tasks def next_task(self) -> str: """Return the next task, reshuffling after each complete cycle.""" with self._lock: if not self._queue: self._reshuffle() return self._queue.pop(0) def snapshot(self) -> tuple[list[str], int]: """Return scheduler state for rollback after failed env creation.""" with self._lock: return list(self._queue), self.cycle def restore(self, snapshot: tuple[list[str], int]) -> None: """Restore scheduler state after a failed env creation.""" queue, cycle = snapshot with self._lock: self._queue = list(queue) self.cycle = cycle def _reshuffle(self) -> None: self._queue = list(self._tasks) self._rng.shuffle(self._queue) self.cycle += 1 def should_enable_task_curriculum(benchmark: str, enabled: bool) -> bool: """Return whether curriculum mode should affect this benchmark.""" return enabled and benchmark == "miniwob" def select_reset_task( current_task: Optional[str], scheduler: Optional[TaskCurriculumScheduler], explicit_task: Optional[str] = None, ) -> Optional[str]: """Select the reset task, giving explicit overrides highest priority.""" if explicit_task is not None: return explicit_task if scheduler is not None: return scheduler.next_task() return current_task