Team Ai
Apppublic

openenv/echo_env

sourceHugging Faceupdated 1d agoView on Hugging Face
6likes
serialization.py186 linesDownload Raw Back to env_server
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