openenv/coding_env
21
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the BSD-style license found in the5# LICENSE file in the root directory of this source tree.6 7import inspect8from abc import ABC, abstractmethod9from typing import Any, Generic, Optional, Protocol, TYPE_CHECKING, TypeVar10 11from typing_extensions import TypedDict12 13from .types import Action, EnvironmentMetadata, Observation, State14 15if TYPE_CHECKING:16 from openenv.core.rubrics import Rubric17 18ActT = TypeVar("ActT", bound=Action)19ObsT = TypeVar("ObsT", bound=Observation)20StateT = TypeVar("StateT", bound=State)21 22 23class Message(TypedDict):24 """A message in a conversation.25 26 Compatible with Huggingface chat template format.27 """28 29 role: str30 content: str31 32 33class ModelTokenizer(Protocol):34 """Protocol for tokenizers that support chat templates.35 36 This protocol defines the interface that tokenizers must implement37 to work with chat-based environments. It's compatible with38 Huggingface transformers tokenizers.39 """40 41 def apply_chat_template(42 self,43 conversation: list[Message],44 tokenize: bool = True,45 return_tensors: str | None = None,46 **kwargs: Any,47 ) -> Any:48 """Apply a chat template to format and optionally tokenize a conversation.49 50 Args:51 conversation: List of message dictionaries with 'role' and 'content'52 tokenize: Whether to tokenize the output53 return_tensors: Format for returned tensors ('pt' for PyTorch)54 **kwargs: Additional arguments55 56 Returns:57 Formatted and optionally tokenized conversation58 """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: Token IDs to decode68 skip_special_tokens: Whether to skip special tokens in output69 **kwargs: Additional arguments70 71 Returns:72 Decoded text string73 """74 ...75 76 77class Transform(ABC, Generic[ObsT]):78 """Transform observations to add rewards, metrics, or other modifications.79 80 Transforms follow the TorchRL pattern where they take an observation81 and return a (potentially modified) observation. This allows for82 flexible reward computation and observation augmentation.83 """84 85 @abstractmethod86 def __call__(self, observation: ObsT) -> ObsT:87 """Transform an observation.88 89 Args:90 observation: The input observation91 92 Returns:93 The transformed observation94 """95 pass96 97 98class Environment(ABC, Generic[ActT, ObsT, StateT]):99 """Base class for all environment servers following Gym/Gymnasium API.100 101 Args:102 transform: Optional transform to apply to observations103 rubric: Optional rubric for reward computation. When provided, the104 rubric's output can be used to set the observation's reward in step().105 106 Class Attributes:107 SUPPORTS_CONCURRENT_SESSIONS: Whether this environment supports concurrent sessions.108 When True, multiple WebSocket connections can each have their own109 environment instance (up to max_concurrent_envs). When False (default),110 the environment should only be used with a single session at a time.111 112 Set this to True in your Environment subclass if:113 - The environment uses proper session isolation (e.g., unique working dirs)114 - No shared mutable state exists between instances115 - External resources (databases, APIs) can handle concurrent access116 117 Attributes:118 rubric: Optional rubric for computing rewards. Environments can set this119 in __init__ and use it in step() to compute observation rewards.120 Training infrastructure can access it for introspection:121 for name, r in env.rubric.named_rubrics():122 print(f"{name}: {r.last_score}")123 124 See RFC 004 for rubric design: rfcs/004-rubrics.md125 """126 127 # Class-level flag indicating whether this environment supports concurrent sessions128 SUPPORTS_CONCURRENT_SESSIONS: bool = False129 130 # Optional rubric for reward computation131 rubric: Optional["Rubric"]132 133 def __init__(134 self,135 transform: Optional[Transform[ObsT]] = None,136 rubric: Optional["Rubric"] = None,137 ):138 self.transform = transform139 self.rubric = rubric140 141 @abstractmethod142 def reset(143 self,144 seed: Optional[int] = None,145 episode_id: Optional[str] = None,146 **kwargs: Any,147 ) -> ObsT:148 """Reset the environment and return initial observation."""149 pass150 151 async def reset_async(152 self,153 seed: Optional[int] = None,154 episode_id: Optional[str] = None,155 **kwargs: Any,156 ) -> ObsT:157 """Async version of reset. Default implementation calls sync reset.158 159 Override to provide true async implementation.160 """161 return self.reset(seed=seed, episode_id=episode_id, **kwargs)162 163 @abstractmethod164 def step(165 self,166 action: ActT,167 timeout_s: Optional[float] = None,168 **kwargs: Any,169 ) -> ObsT:170 """Take a step in the environment."""171 pass172 173 async def step_async(174 self,175 action: ActT,176 timeout_s: Optional[float] = None,177 **kwargs: Any,178 ) -> ObsT:179 """Async version of step. Default implementation calls sync step.180 181 Override to provide true async implementation.182 """183 return self.step(action, timeout_s=timeout_s, **kwargs)184 185 @property186 @abstractmethod187 def state(self) -> StateT:188 """Get the current environment state."""189 pass190 191 def get_metadata(self) -> EnvironmentMetadata:192 """193 Get metadata about this environment.194 195 Override this method to provide custom metadata for the environment.196 Default implementation returns basic metadata derived from class name.197 198 Returns:199 EnvironmentMetadata with environment information200 """201 return EnvironmentMetadata(202 name=self.__class__.__name__,203 description=f"{self.__class__.__name__} environment",204 version="1.0.0",205 )206 207 def _apply_transform(self, observation: ObsT) -> ObsT:208 """Apply transform if one is provided."""209 if self.transform is not None:210 return self.transform(observation)211 return observation212 213 def _apply_rubric(self, action: ActT, observation: ObsT) -> float:214 """Apply rubric if one is provided.215 216 Args:217 action: The action taken by the agent.218 observation: The resulting observation.219 220 Returns:221 Reward value from the rubric, or 0.0 if no rubric is set.222 223 Usage in step():224 def step(self, action: MyAction, ...) -> MyObservation:225 # ... execute action and create observation ...226 observation.reward = self._apply_rubric(action, observation)227 return observation228 """229 if self.rubric is not None:230 return self.rubric(action, observation)231 return 0.0232 233 async def _apply_rubric_async(self, action: ActT, observation: ObsT) -> float:234 """Apply rubric asynchronously if one is provided.235 236 Args:237 action: The action taken by the agent.238 observation: The resulting observation.239 240 Returns:241 Reward value from the rubric, or 0.0 if no rubric is set.242 243 Usage in step_async():244 async def step_async(self, action: MyAction, ...) -> MyObservation:245 # ... execute action and create observation ...246 observation.reward = await self._apply_rubric_async(action, observation)247 return observation248 """249 if self.rubric is not None:250 result = self.rubric(action, observation)251 # If rubric returns a coroutine, await it252 if inspect.iscoroutine(result):253 return await result254 return result255 return 0.0256 257 def _reset_rubric(self) -> None:258 """Reset the rubric state if one is provided.259 260 Call this in reset() to clear any trajectory state in the rubric.261 262 Usage in reset():263 def reset(self, ...) -> MyObservation:264 self._reset_rubric()265 # ... create initial observation ...266 return observation267 """268 if self.rubric is not None:269 self.rubric.reset()270 271 async def _reset_rubric_async(self) -> None:272 """Reset the rubric state asynchronously if one is provided.273 274 Call this in reset_async() to clear any trajectory state in the rubric.275 276 Usage in reset_async():277 async def reset_async(self, ...) -> MyObservation:278 await self._reset_rubric_async()279 # ... create initial observation ...280 return observation281 """282 if self.rubric is not None:283 # Check if rubric has async reset method284 if hasattr(self.rubric, "reset_async"):285 result = self.rubric.reset_async()286 if inspect.iscoroutine(result):287 await result288 else:289 self.rubric.reset()290 291 def close(self) -> None:292 """Clean up resources used by the environment.293 294 Override this method to implement custom cleanup logic.295 Called when the environment is being destroyed or reset.296 """297 pass298 