"""Typed OpenEnv wire models for MATH.""" from __future__ import annotations from typing import Any from openenv.core.env_server.types import Action, Observation, State from pydantic import BaseModel, ConfigDict, Field, model_validator class MathSubmitAnswer(BaseModel): """Submit the final answer for the current MATH problem.""" model_config = ConfigDict(extra="forbid") answer: str = Field( min_length=1, description=( "One complete final answer string (a number, expression, or set). " "No explanation is expected. Plain text or LaTeX is accepted." ), ) class MathAction(Action, MathSubmitAnswer): """The single final-answer string submitted for an episode.""" @model_validator(mode="before") @classmethod def _unwrap_native_tool_call(cls, values: Any) -> Any: if isinstance(values, dict) and set(values) == {"name", "arguments"}: return native_tool_action( values["name"], values["arguments"] ).model_dump() return values def native_tool_specs() -> tuple[dict[str, Any], ...]: """Return the ordered model-facing MATH native-tool catalog.""" return ( { "name": "submit_answer", "description": MathSubmitAnswer.__doc__, "input_schema": MathSubmitAnswer.model_json_schema(), }, ) def native_tool_action(name: str, arguments: dict[str, Any]) -> MathAction: """Validate a native call and construct the OpenEnv transport action.""" if name != "submit_answer": raise ValueError(f"unknown MATH native tool: {name!r}") call = MathSubmitAnswer.model_validate(arguments) return MathAction(**call.model_dump()) class MathObservation(Observation): """Actor-visible problem or terminal summary without hidden grading evidence. Official class and difficulty are post-hoc metadata and are deliberately NOT exposed to the solver; they appear only in the evaluator-facing ``MathState``. """ task_id: int text: str goal: str success: bool progress: float = Field(ge=0.0, le=1.0) terminated: bool truncated: bool step_limit: int = 1 @model_validator(mode="before") @classmethod def _derive_done(cls, values: Any) -> Any: if isinstance(values, dict): values = dict(values) values["done"] = bool( values.get("terminated", False) or values.get("truncated", False) ) return values class MathState(State): """Evaluator-facing state for one connection-owned episode. Official class and difficulty are retained here as post-hoc metadata for analysis; they are never part of the actor-visible observation. """ task_id: int | None = None goal: str | None = None source_class: str | None = None level: str | None = None reward: float = 0.0 success: bool = False progress: float = Field(default=0.0, ge=0.0, le=1.0) terminated: bool = False truncated: bool = False step_limit: int = 1 diagnostics: dict[str, str] = Field(default_factory=dict) closed: bool = False