openenv/echo_env
6
1# SPDX-License-Identifier: BSD-3-Clause2 3"""4Shared serialization and deserialization utilities for OpenEnv HTTP servers.5 6This module provides common utilities for converting between JSON dictionaries7and Pydantic models (Action/Observation) to eliminate code duplication across8HTTP server and web interface implementations.9"""10 11import json12from typing import Any, Dict, Type13 14from .mcp_types import CallToolAction, ListToolsAction15from .types import Action, Observation16 17# MCP action types keyed by their "type" discriminator value.18# These are checked before the environment's own action_cls so that19# ListToolsAction / CallToolAction payloads are never rejected by an20# unrelated Pydantic model.21_MCP_ACTION_TYPES: Dict[str, Type[Action]] = {22 "list_tools": ListToolsAction,23 "call_tool": CallToolAction,24}25 26 27def _deserialize_mcp_action(28 action_data: Dict[str, Any], action_cls: Type[Action]29) -> Action | None:30 # Only intercept when action_cls is the generic Action base or itself an31 # MCP type. This keeps env-specific action validation authoritative.32 action_type = action_data.get("type")33 if action_type not in _MCP_ACTION_TYPES:34 return None35 36 mcp_cls = _MCP_ACTION_TYPES[action_type]37 if action_cls is Action or action_cls in _MCP_ACTION_TYPES.values():38 return mcp_cls.model_validate(action_data)39 40 return None41 42 43def deserialize_action(action_data: Dict[str, Any], action_cls: Type[Action]) -> Action:44 """45 Convert JSON dict to Action instance using Pydantic validation.46 47 MCP action types (``list_tools``, ``call_tool``) are recognised48 automatically via the ``"type"`` discriminator field, regardless of49 the environment's configured ``action_cls``. All other payloads50 fall through to ``action_cls.model_validate()``.51 52 For special cases (e.g., tensor fields, custom type conversions),53 use deserialize_action_with_preprocessing().54 55 Args:56 action_data (`dict`):57 Dictionary containing action data.58 action_cls (`type`):59 The Action subclass to instantiate.60 61 Returns:62 `Action` instance.63 64 Raises:65 `ValidationError`: If `action_data` is invalid for the action class.66 """67 mcp_action = _deserialize_mcp_action(action_data, action_cls)68 if mcp_action is not None:69 return mcp_action70 71 return action_cls.model_validate(action_data)72 73 74def deserialize_action_with_preprocessing(75 action_data: Dict[str, Any], action_cls: Type[Action]76) -> Action:77 """78 Convert JSON dict to Action instance with preprocessing for special types.79 80 This version handles common type conversions needed for web interfaces:81 - Converting JSON string arguments to dict for MCP call_tool actions82 - Converting lists/strings to tensors for 'tokens' field83 - Converting string action_id to int84 - Other custom preprocessing as needed85 86 Args:87 action_data (`dict`):88 Dictionary containing action data.89 action_cls (`type`):90 The Action subclass to instantiate.91 92 Returns:93 `Action` instance.94 95 Raises:96 `ValidationError`: If `action_data` is invalid for the action class.97 """98 mcp_data = action_data99 if action_data.get("type") == "call_tool" and isinstance(100 action_data.get("arguments"), str101 ):102 mcp_data = dict(action_data)103 try:104 mcp_data["arguments"] = json.loads(action_data["arguments"])105 except Exception:106 pass107 108 mcp_action = _deserialize_mcp_action(mcp_data, action_cls)109 if mcp_action is not None:110 return mcp_action111 112 processed_data = {}113 114 for key, value in action_data.items():115 if key == "tokens" and isinstance(value, (list, str)):116 # Convert list or string to tensor117 if isinstance(value, str):118 # If it's a string, try to parse it as a list of numbers119 try:120 value = json.loads(value)121 except Exception:122 # If parsing fails, treat as empty list123 value = []124 if isinstance(value, list):125 try:126 import torch # type: ignore127 128 processed_data[key] = torch.tensor(value, dtype=torch.long)129 except ImportError:130 # If torch not available, keep as list131 processed_data[key] = value132 else:133 processed_data[key] = value134 elif key == "action_id" and isinstance(value, str):135 # Convert action_id from string to int136 try:137 processed_data[key] = int(value)138 except ValueError:139 # If conversion fails, keep original value140 processed_data[key] = value141 else:142 processed_data[key] = value143 144 return action_cls.model_validate(processed_data)145 146 147def serialize_observation(observation: Observation) -> Dict[str, Any]:148 """149 Convert Observation instance to JSON-compatible dict using Pydantic.150 151 Args:152 observation (`Observation`):153 Observation instance to serialize.154 155 Returns:156 `dict` compatible with `EnvClient._parse_result()`, with keys:157 - `observation` (`dict`): Observation fields.158 - `reward` (`float` or `None`): Reward value.159 - `done` (`bool`): Whether the episode is done.160 - `metadata` (`dict`, *optional*): Additional observation metadata.161 """162 # Keep metadata in the nested observation payload for backwards163 # compatibility with typed clients, and also surface it as a top-level164 # sibling for clients that read the generic wire format directly.165 obs_dict = observation.model_dump(166 exclude={167 "reward",168 "done",169 } # Exclude these from observation dict170 )171 172 # Extract reward, done, and metadata directly from the observation.173 reward = observation.reward174 done = observation.done175 metadata = observation.metadata176 177 # Return in EnvClient expected format.178 result = {179 "observation": obs_dict,180 "reward": reward,181 "done": done,182 }183 if metadata:184 result["metadata"] = metadata185 return result186 