Spaces:
Paused
Paused
Download src/math_env/client.py from emrekuruu/math-openenv: direct link, hf CLI and curl.
- Browser
- Download file 3.8 kB
-
https://huggingface.co/spaces/emrekuruu/math-openenv/resolve/main/src/math_env/client.py
- Command line
-
hf download hf://spaces/emrekuruu/math-openenv/src/math_env/client.py
-
curl -L -o client.py https://huggingface.co/spaces/emrekuruu/math-openenv/resolve/main/src/math_env/client.py
3.8 kB
| """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) |