"""Typed persistent-session client for MATH OpenEnv.""" from __future__ import annotations import asyncio from typing import Any from urllib.parse import urlsplit from openenv.core import EnvClient from openenv.core.client_types import StepResult from websockets.asyncio.client import connect as ws_connect from websockets.exceptions import SecurityError from math_env.models import ( MathAction, MathObservation, MathState, ) class _AuthenticatedConnect(ws_connect): def process_redirect(self, exc: Exception) -> Exception | str: result = super().process_redirect(exc) if isinstance(result, str): return SecurityError("Authenticated WebSocket redirects are disabled") return result def _authenticated_ws_connect( uri: str, *, bearer_token: str, **kwargs: Any, ): return _AuthenticatedConnect( uri, additional_headers={"Authorization": f"Bearer {bearer_token}"}, **kwargs, ) class MathClient(EnvClient[MathAction, MathObservation, MathState]): """Typed WebSocket client for one connection-owned episode.""" def __init__( self, *args: Any, bearer_token: str | None = None, **kwargs: Any, ) -> None: base_url = kwargs.get("base_url", args[0] if args else None) if bearer_token is not None: scheme = urlsplit(str(base_url)).scheme.lower() if scheme not in {"https", "wss"}: raise ValueError( "Bearer authentication requires an HTTPS or WSS endpoint" ) self._bearer_token = bearer_token super().__init__(*args, **kwargs) async def _connect_async(self) -> "MathClient": if self._bearer_token is None: await super()._connect_async() return self if self._ws is not None: if self._ws_loop is asyncio.get_running_loop(): return self self._ws = None self._ws_loop = None try: self._start_provider_if_needed() except Exception: await self.close() raise assert self._ws_url is not None try: self._ws = await _authenticated_ws_connect( self._ws_url, bearer_token=self._bearer_token, open_timeout=self._connect_timeout, max_size=self._max_message_size, ping_interval=self._websocket_ping_interval_s, ping_timeout=self._websocket_ping_timeout_s, ) self._ws_loop = asyncio.get_running_loop() except Exception as error: await self.close() raise ConnectionError( f"Failed to connect to {self._ws_url}: {error}" ) from error return self def _step_payload(self, action: MathAction) -> dict[str, Any]: return action.model_dump() def _parse_result( self, payload: dict[str, Any], ) -> StepResult[MathObservation]: observation_data = dict(payload.get("observation", {})) reward_value = payload.get("reward") reward = None if reward_value is None else float(reward_value) metadata = payload.get("metadata", observation_data.get("metadata")) observation_data.update( reward=reward, metadata={} if metadata is None else metadata, ) observation = MathObservation.model_validate(observation_data) return StepResult( observation=observation, reward=reward, done=observation.done, metadata=metadata, ) def _parse_state(self, payload: dict[str, Any]) -> MathState: return MathState.model_validate(payload)