"""Framework-neutral MATH single-attempt episode lifecycle.""" from __future__ import annotations from collections.abc import Callable from dataclasses import dataclass from math_env.catalog import MathTask, get_task from math_env.grader import compare_answers, verification_report from math_env.interface import build_task_prompt, extract_answer @dataclass(frozen=True, slots=True) class CoreObservation: task_id: int text: str goal: str success: bool progress: float step_limit: int = 1 @dataclass(frozen=True, slots=True) class CoreStep: observation: CoreObservation reward: float terminated: bool @dataclass(frozen=True, slots=True) class CoreState: task_id: int | None goal: str | None source_class: str | None level: str | None reward: float success: bool progress: float terminated: bool step_count: int step_limit: int diagnostics: dict[str, str] closed: bool _TaskLookup = Callable[[int], MathTask] class MathCore: """Own one isolated, single-attempt MATH episode. MATH has no simulator, sandbox, or code execution: reset exposes the problem and step grades a submitted answer against the hidden reference. This is the whole runtime. """ def __init__( self, *, task_lookup: _TaskLookup = get_task, ) -> None: self._task_lookup = task_lookup self._task: MathTask | None = None self._reward = 0.0 self._progress = 0.0 self._terminated = False self._step_count = 0 self._diagnostics: dict[str, str] = {} self._closed = False def reset(self, task_id: int) -> CoreObservation: self._ensure_open() task = self._task_lookup(int(task_id)) self._task = task self._reward = 0.0 self._progress = 0.0 self._terminated = False self._step_count = 0 self._diagnostics = {} prompt = build_task_prompt(task.instruction) return self._observation(prompt) def step(self, action_answer: str) -> CoreStep: self._ensure_open() if self._task is None: raise RuntimeError("MATH core must be reset before step") if self._terminated: raise RuntimeError("MATH episode has terminated") answer = extract_answer(action_answer) matched = compare_answers(answer, self._task.answer) self._step_count = 1 self._terminated = True self._progress = 1.0 if matched else 0.0 self._reward = 1.0 if matched else 0.0 self._diagnostics = { "verification": verification_report(answer, self._task.answer), "fail_reason": "" if matched else "answer-mismatch", } summary = "Answer graded: PASS" if matched else "Answer graded: FAIL" observation = self._observation(summary) return CoreStep( observation=observation, reward=self._reward, terminated=True, ) def state(self) -> CoreState: return CoreState( task_id=None if self._task is None else self._task.task_id, goal=None if self._task is None else self._task.instruction, source_class=None if self._task is None else self._task.source_class, level=None if self._task is None else self._task.level, reward=self._reward, success=self._terminated and self._reward == 1.0, progress=self._progress, terminated=self._terminated, step_count=self._step_count, step_limit=1, diagnostics=dict(self._diagnostics), closed=self._closed, ) def close(self) -> None: self._closed = True def _observation(self, text: str) -> CoreObservation: if self._task is None: raise RuntimeError("MATH core has no current task") return CoreObservation( task_id=self._task.task_id, text=text, goal=self._task.instruction, success=self._terminated and self._reward == 1.0, progress=self._progress, ) def _ensure_open(self) -> None: if self._closed: raise RuntimeError("MATH core is closed")