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 7"""8Shared serialization and deserialization utilities for OpenEnv HTTP servers.9 10This module provides common utilities for converting between JSON dictionaries11and Pydantic models (Action/Observation) to eliminate code duplication across12HTTP server and web interface implementations.13"""14 15from typing import Any, Dict, Type16 17from .mcp_types import CallToolAction, ListToolsAction18from .types import Action, Observation19 20# MCP action types keyed by their "type" discriminator value.21# These are checked before the environment's own action_cls so that22# ListToolsAction / CallToolAction payloads are never rejected by an23# unrelated Pydantic model.24_MCP_ACTION_TYPES: Dict[str, Type[Action]] = {25 "list_tools": ListToolsAction,26 "call_tool": CallToolAction,27}28 29 30def deserialize_action(action_data: Dict[str, Any], action_cls: Type[Action]) -> Action:31 """32 Convert JSON dict to Action instance using Pydantic validation.33 34 MCP action types (``list_tools``, ``call_tool``) are recognised35 automatically via the ``"type"`` discriminator field, regardless of36 the environment's configured ``action_cls``. All other payloads37 fall through to ``action_cls.model_validate()``.38 39 For special cases (e.g., tensor fields, custom type conversions),40 use deserialize_action_with_preprocessing().41 42 Args:43 action_data: Dictionary containing action data44 action_cls: The Action subclass to instantiate45 46 Returns:47 Action instance48 49 Raises:50 ValidationError: If action_data is invalid for the action class51 52 Note:53 This uses Pydantic's model_validate() for automatic validation.54 """55 # Route MCP action types before falling through to the env action_cls.56 # Only intercept when action_cls is the generic Action base or itself an57 # MCP type (i.e. the server hosts an MCP environment). This avoids58 # silently bypassing env-specific validation for non-MCP environments59 # that happen to use "call_tool" / "list_tools" as a type discriminator.60 action_type = action_data.get("type")61 if action_type in _MCP_ACTION_TYPES:62 mcp_cls = _MCP_ACTION_TYPES[action_type]63 if action_cls is Action or action_cls in _MCP_ACTION_TYPES.values():64 return mcp_cls.model_validate(action_data)65 66 return action_cls.model_validate(action_data)67 68 69def deserialize_action_with_preprocessing(70 action_data: Dict[str, Any], action_cls: Type[Action]71) -> Action:72 """73 Convert JSON dict to Action instance with preprocessing for special types.74 75 This version handles common type conversions needed for web interfaces:76 - Converting lists/strings to tensors for 'tokens' field77 - Converting string action_id to int78 - Other custom preprocessing as needed79 80 Args:81 action_data: Dictionary containing action data82 action_cls: The Action subclass to instantiate83 84 Returns:85 Action instance86 87 Raises:88 ValidationError: If action_data is invalid for the action class89 """90 # Route MCP action types before preprocessing (they don't need it).91 # Same guard as deserialize_action: only intercept when action_cls is92 # the generic Action base or itself an MCP type.93 action_type = action_data.get("type")94 if action_type in _MCP_ACTION_TYPES:95 mcp_cls = _MCP_ACTION_TYPES[action_type]96 if action_cls is Action or action_cls in _MCP_ACTION_TYPES.values():97 return mcp_cls.model_validate(action_data)98 99 processed_data = {}100 101 for key, value in action_data.items():102 if key == "tokens" and isinstance(value, (list, str)):103 # Convert list or string to tensor104 if isinstance(value, str):105 # If it's a string, try to parse it as a list of numbers106 try:107 import json108 109 value = json.loads(value)110 except Exception:111 # If parsing fails, treat as empty list112 value = []113 if isinstance(value, list):114 try:115 import torch # type: ignore116 117 processed_data[key] = torch.tensor(value, dtype=torch.long)118 except ImportError:119 # If torch not available, keep as list120 processed_data[key] = value121 else:122 processed_data[key] = value123 elif key == "action_id" and isinstance(value, str):124 # Convert action_id from string to int125 try:126 processed_data[key] = int(value)127 except ValueError:128 # If conversion fails, keep original value129 processed_data[key] = value130 else:131 processed_data[key] = value132 133 return action_cls.model_validate(processed_data)134 135 136def serialize_observation(observation: Observation) -> Dict[str, Any]:137 """138 Convert Observation instance to JSON-compatible dict using Pydantic.139 140 Args:141 observation: Observation instance142 143 Returns:144 Dictionary compatible with EnvClient._parse_result()145 146 The format matches what EnvClient expects:147 {148 "observation": {...}, # Observation fields149 "reward": float | None,150 "done": bool,151 }152 """153 # Use Pydantic's model_dump() for serialization154 obs_dict = observation.model_dump(155 exclude={156 "reward",157 "done",158 "metadata",159 } # Exclude these from observation dict160 )161 162 # Extract reward and done directly from the observation163 reward = observation.reward164 done = observation.done165 166 # Return in EnvClient expected format167 return {168 "observation": obs_dict,169 "reward": reward,170 "done": done,171 }172 