math-openenv / src /math_env /models.py
emrekuruu's picture
Argument descriptions carry format only; environment description is the single home of world rules
95ee698 verified
Raw History Blame Contribute Delete
3.18 kB
"""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