"""Genuine OpenEnv adapter around one MATH core instance.""" from __future__ import annotations from collections.abc import Callable from typing import Any from uuid import uuid4 from openenv.core.env_server.interfaces import Environment from openenv.core.env_server.types import EnvironmentMetadata from math_env.catalog import get_split from math_env.core import CoreObservation, MathCore from math_env.interface import environment_description from math_env.models import ( MathAction, MathObservation, MathState, ) _CoreFactory = Callable[[], MathCore] class MathEnvironment( Environment[MathAction, MathObservation, MathState] ): """One connection-owned, typed MATH episode.""" SUPPORTS_CONCURRENT_SESSIONS = True REQUIRES_SINGLE_THREAD_EXECUTOR = False def __init__( self, core_factory: _CoreFactory = MathCore, *, default_task_id: int | None = None, ) -> None: super().__init__() self._core_factory = core_factory self._core: MathCore | None = None self._default_task_id = default_task_id self._episode_id: str | None = None self._closed = False def _get_core(self) -> MathCore: if self._closed: raise RuntimeError("MATH environment is closed") if self._core is None: self._core = self._core_factory() return self._core def get_metadata(self) -> EnvironmentMetadata: return EnvironmentMetadata( name="MATH", description=environment_description(), version="0.1.0", documentation_url="https://huggingface.co/datasets/HuggingFaceH4/MATH", ) def list_splits(self) -> list[dict[str, Any]]: return [ { "name": name, "role": get_split(name).role, "num_tasks": len(get_split(name).tasks), "source_revision": get_split(name).source_revision, "sha256": get_split(name).sha256, } for name in ("development", "test") ] def list_tasks(self, split: str) -> list[dict[str, Any]]: task_split = get_split(split) return [ { "id": task.task_id, "source_id": task.source_id, "instruction": task.instruction, "task_type": task.source_class, "source_class": task.source_class, "level": task.level, "split": task_split.role, "source_revision": task_split.source_revision, "catalog_sha256": task_split.sha256, } for task in task_split.tasks ] def num_tasks(self, split: str) -> int: return len(get_split(split).tasks) def get_task(self, split: str, index: int) -> dict[str, Any]: tasks = self.list_tasks(split) try: return tasks[index] except IndexError as error: raise IndexError( f"MATH {split} task index out of range: {index}" ) from error def get_task_range( self, split: str, start: int | None = None, stop: int | None = None, ) -> list[dict[str, Any]]: return self.list_tasks(split)[slice(start, stop)] def reset( self, seed: int | None = None, episode_id: str | None = None, *, task_id: int | None = None, **kwargs: Any, ) -> MathObservation: del seed, kwargs resolved = self._default_task_id if task_id is None else int(task_id) if resolved is None: raise ValueError("MATH reset requires task_id") core_observation = self._get_core().reset(resolved) self._episode_id = episode_id or str(uuid4()) return self._observation(core_observation, reward=0.0, terminated=False) def step( self, action: MathAction, timeout_s: float | None = None, **kwargs: Any, ) -> MathObservation: del timeout_s, kwargs transition = self._get_core().step(action.answer) return self._observation( transition.observation, reward=transition.reward, terminated=transition.terminated, ) @property def state(self) -> MathState: if self._core is None: return MathState( episode_id=self._episode_id, step_count=0, closed=self._closed, ) state = self._core.state() return MathState( episode_id=self._episode_id, step_count=state.step_count, task_id=state.task_id, goal=state.goal, source_class=state.source_class, level=state.level, reward=state.reward, success=state.success, progress=state.progress, terminated=state.terminated, truncated=False, step_limit=state.step_limit, diagnostics=state.diagnostics, closed=state.closed, ) def close(self) -> None: if self._closed: return if self._core is not None: self._core.close() self._closed = True def _observation( self, observation: CoreObservation, *, reward: float, terminated: bool, ) -> MathObservation: return MathObservation( task_id=observation.task_id, text=observation.text, goal=observation.goal, success=observation.success, progress=observation.progress, terminated=terminated, truncated=False, done=terminated, reward=reward, step_limit=observation.step_limit, )