openenv/echo_env
6
1# SPDX-License-Identifier: BSD-3-Clause2 3import inspect4from abc import ABC, abstractmethod5from typing import Any, Generic, Optional, Protocol, TYPE_CHECKING, TypeVar6 7from typing_extensions import TypedDict8 9from .types import Action, EnvironmentMetadata, Observation, State10 11if TYPE_CHECKING:12 from openenv.core.rubrics import Rubric13 14ActT = TypeVar("ActT", bound=Action)15ObsT = TypeVar("ObsT", bound=Observation)16StateT = TypeVar("StateT", bound=State)17 18 19class Message(TypedDict):20 """A message in a conversation.21 22 Compatible with Huggingface chat template format.23 """24 25 role: str26 content: str27 28 29class ModelTokenizer(Protocol):30 """Protocol for tokenizers that support chat templates.31 32 This protocol defines the interface that tokenizers must implement33 to work with chat-based environments. It's compatible with34 Huggingface transformers tokenizers.35 """36 37 def apply_chat_template(38 self,39 conversation: list[Message],40 tokenize: bool = True,41 return_tensors: str | None = None,42 **kwargs: Any,43 ) -> Any:44 """Apply a chat template to format and optionally tokenize a conversation.45 46 Args:47 conversation (`list[Message]`):48 List of message dictionaries with 'role' and 'content'.49 tokenize (`bool`, *optional*, defaults to `True`):50 Whether to tokenize the output.51 return_tensors (`str`, *optional*):52 Format for returned tensors ('pt' for PyTorch).53 **kwargs:54 Additional arguments.55 56 Returns:57 Formatted and optionally tokenized conversation.58 """59 ...60 61 def decode(62 self, token_ids: Any, skip_special_tokens: bool = False, **kwargs: Any63 ) -> str:64 """Decode token IDs back to text.65 66 Args:67 token_ids (`Any`):68 Token IDs to decode.69 skip_special_tokens (`bool`, *optional*, defaults to `False`):70 Whether to skip special tokens in output.71 **kwargs:72 Additional arguments.73 74 Returns:75 `str`: Decoded text string.76 """77 ...78 79 80class TaskProvider(Protocol):81 """82 Optional task discovery API for dataset-backed environments.83 84 An environment implements this protocol structurally — declare the methods on85 an [`~openenv.core.env_server.interfaces.Environment`] subclass, without86 inheriting from `TaskProvider`. When the methods are present,87 [`~openenv.core.env_server.http_server.HTTPEnvServer`] exposes them as HTTP88 routes under `/{env_name}/…`; when they are absent, those routes return89 `501 Not Implemented`. Each method may be sync or async.90 91 Task provider methods are for metadata/discovery only and should be92 side-effect-free. They must be callable on a freshly constructed93 environment instance because HTTP compatibility routes may create a94 short-lived instance solely for task discovery.95 96 Selecting a task is not part of this protocol — pass the chosen split and97 index to `reset()` instead. See the98 [Task API guide](https://huggingface.co/docs/openenv/guides/task-api).99 100 Examples:101 102 ```python103 env.list_splits() # ["train", "test"]104 env.num_tasks("test") # 7595105 env.get_task("test", 12) # {"id": "test-12", "index": 12}106 env.reset(split="test", index=12)107 ```108 """109 110 def list_splits(self) -> list[Any]:111 """112 Return task split descriptors supported by this environment.113 114 Returns:115 `list[Any]`: Split descriptors. Plain strings, dicts, and Pydantic116 models are all accepted; the server normalizes each entry to117 `{"name": ..., "type": ...}`.118 """119 ...120 121 def list_tasks(self, split: str) -> list[Any]:122 """123 Return all task specs for a split.124 125 Args:126 split (`str`):127 Task split name.128 129 Returns:130 `list[Any]`: Task specs for the split. Environments backed by very131 large or streamed splits may return a bounded preview, but132 `num_tasks` should still report the true total.133 """134 ...135 136 def num_tasks(self, split: str) -> int:137 """138 Return the number of task specs in a split.139 140 Args:141 split (`str`):142 Task split name.143 144 Returns:145 `int`: Number of task specs available in the split.146 """147 ...148 149 def get_task(self, split: str, index: int) -> Any:150 """151 Return one task spec by split and index.152 153 Args:154 split (`str`):155 Task split name.156 index (`int`):157 Task index within the split.158 159 Returns:160 `Any`: The task spec at that position.161 162 Raises:163 `IndexError`: If `index` is out of range for the split. The HTTP164 route converts this to a `400 Bad Request`.165 """166 ...167 168 def get_task_range(169 self,170 split: str,171 start: Optional[int] = None,172 stop: Optional[int] = None,173 ) -> list[Any]:174 """175 Return task specs for Python slice-style range bounds.176 177 Args:178 split (`str`):179 Task split name.180 start (`int`, *optional*):181 Inclusive start index. Defaults to the beginning of the split.182 stop (`int`, *optional*):183 Exclusive stop index. Defaults to the end of the split.184 185 Returns:186 `list[Any]`: Task specs in `[start, stop)`.187 """188 ...189 190 191class Transform(ABC, Generic[ObsT]):192 """Transform observations to add rewards, metrics, or other modifications.193 194 Transforms follow the TorchRL pattern where they take an observation195 and return a (potentially modified) observation. This allows for196 flexible reward computation and observation augmentation.197 """198 199 @abstractmethod200 def __call__(self, observation: ObsT) -> ObsT:201 """Transform an observation.202 203 Args:204 observation (`ObsT`):205 The input observation.206 207 Returns:208 `ObsT`: The transformed observation.209 """210 pass211 212 213class Environment(ABC, Generic[ActT, ObsT, StateT]):214 """Base class for all environment servers following Gym/Gymnasium API.215 216 See [rfcs/004-rubrics.md](https://github.com/huggingface/OpenEnv/blob/main/rfcs/004-rubrics.md) for rubric design details.217 218 Args:219 transform (`Transform`, *optional*):220 Optional transform to apply to observations.221 rubric (`Rubric`, *optional*):222 Optional rubric for reward computation. When provided, the223 rubric's output can be used to set the observation's reward in step().224 225 Attributes:226 SUPPORTS_CONCURRENT_SESSIONS (`bool`):227 Whether this environment supports concurrent sessions. When ``True``,228 multiple WebSocket connections can each have their own environment229 instance (up to ``max_concurrent_envs``). When ``False`` (default),230 the environment should only be used with a single session at a time.231 232 Set this to ``True`` in your subclass if the environment uses proper233 session isolation (unique working dirs, no shared mutable state, and234 external resources that can handle concurrent access).235 rubric (`Rubric`, *optional*):236 Optional rubric for computing rewards. Set in ``__init__`` and use in237 ``step()`` to compute observation rewards. Training infrastructure can238 access it for introspection:239 240 ```python241 for name, r in env.rubric.named_rubrics():242 print(f"{name}: {r.last_score}")243 ```244 """245 246 # Class-level flag indicating whether this environment supports concurrent sessions247 SUPPORTS_CONCURRENT_SESSIONS: bool = False248 249 REQUIRES_SINGLE_THREAD_EXECUTOR: bool = False250 251 # Optional rubric for reward computation252 rubric: Optional["Rubric"]253 254 def __init__(255 self,256 transform: Optional[Transform[ObsT]] = None,257 rubric: Optional["Rubric"] = None,258 ):259 self.transform = transform260 self.rubric = rubric261 262 @abstractmethod263 def reset(264 self,265 seed: Optional[int] = None,266 episode_id: Optional[str] = None,267 **kwargs: Any,268 ) -> ObsT:269 """Reset the environment and return initial observation."""270 pass271 272 async def reset_async(273 self,274 seed: Optional[int] = None,275 episode_id: Optional[str] = None,276 **kwargs: Any,277 ) -> ObsT:278 """Async version of reset. Default implementation calls sync reset.279 280 Override to provide true async implementation.281 """282 return self.reset(seed=seed, episode_id=episode_id, **kwargs)283 284 @abstractmethod285 def step(286 self,287 action: ActT,288 timeout_s: Optional[float] = None,289 **kwargs: Any,290 ) -> ObsT:291 """Take a step in the environment."""292 pass293 294 async def step_async(295 self,296 action: ActT,297 timeout_s: Optional[float] = None,298 **kwargs: Any,299 ) -> ObsT:300 """Async version of step. Default implementation calls sync step.301 302 Override to provide true async implementation.303 """304 return self.step(action, timeout_s=timeout_s, **kwargs)305 306 @property307 @abstractmethod308 def state(self) -> StateT:309 """Get the current environment state."""310 pass311 312 def get_metadata(self) -> EnvironmentMetadata:313 """314 Get metadata about this environment.315 316 Override this method to provide custom metadata for the environment.317 Default implementation returns basic metadata derived from class name.318 319 Returns:320 [`EnvironmentMetadata`] with environment information.321 """322 return EnvironmentMetadata(323 name=self.__class__.__name__,324 description=f"{self.__class__.__name__} environment",325 version="1.0.0",326 )327 328 def _apply_transform(self, observation: ObsT) -> ObsT:329 """Apply transform if one is provided."""330 if self.transform is not None:331 return self.transform(observation)332 return observation333 334 def _apply_rubric(self, action: ActT, observation: ObsT) -> float:335 """Apply rubric if one is provided.336 337 Args:338 action (`ActT`):339 The action taken by the agent.340 observation (`ObsT`):341 The resulting observation.342 343 Returns:344 `float`: Reward value from the rubric, or 0.0 if no rubric is set.345 346 Call this in `step()` to compute and assign the reward:347 348 ```python349 def step(self, action: MyAction, ...) -> MyObservation:350 # ... execute action and create observation ...351 observation.reward = self._apply_rubric(action, observation)352 return observation353 ```354 """355 if self.rubric is not None:356 return self.rubric(action, observation)357 return 0.0358 359 async def _apply_rubric_async(self, action: ActT, observation: ObsT) -> float:360 """Apply rubric asynchronously if one is provided.361 362 Args:363 action (`ActT`):364 The action taken by the agent.365 observation (`ObsT`):366 The resulting observation.367 368 Returns:369 `float`: Reward value from the rubric, or 0.0 if no rubric is set.370 371 Call this in `step_async()` to compute and assign the reward:372 373 ```python374 async def step_async(self, action: MyAction, ...) -> MyObservation:375 # ... execute action and create observation ...376 observation.reward = await self._apply_rubric_async(action, observation)377 return observation378 ```379 """380 if self.rubric is not None:381 result = self.rubric(action, observation)382 # If rubric returns a coroutine, await it383 if inspect.iscoroutine(result):384 return await result385 return result386 return 0.0387 388 def _reset_rubric(self) -> None:389 """Reset the rubric state if one is provided.390 391 Call this in `reset()` to clear any trajectory state in the rubric:392 393 ```python394 def reset(self, ...) -> MyObservation:395 self._reset_rubric()396 # ... create initial observation ...397 return observation398 ```399 """400 if self.rubric is not None:401 self.rubric.reset()402 403 async def _reset_rubric_async(self) -> None:404 """Reset the rubric state asynchronously if one is provided.405 406 Call this in `reset_async()` to clear any trajectory state in the rubric:407 408 ```python409 async def reset_async(self, ...) -> MyObservation:410 await self._reset_rubric_async()411 # ... create initial observation ...412 return observation413 ```414 """415 if self.rubric is not None:416 # Check if rubric has async reset method417 if hasattr(self.rubric, "reset_async"):418 result = self.rubric.reset_async()419 if inspect.iscoroutine(result):420 await result421 else:422 self.rubric.reset()423 424 def close(self) -> None:425 """Clean up resources used by the environment.426 427 Override this method to implement custom cleanup logic.428 Called when the environment is being destroyed or reset.429 """430 pass431 