Team Ai
Apppublic

openenv/coding_env

sourceHugging Faceupdated 3mo agoView on Hugging Face
21likes
serialization.py172 linesDownload Raw Back to env_server
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