Team Ai
Apppublic

openenv/coding_env

sourceHugging Faceupdated 3mo agoView on Hugging Face
21likes
web_interface.py698 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"""8Web interface for OpenEnv environments.9 10When ENABLE_WEB_INTERFACE is set, the server exposes a Gradio UI at /web for11reset, step, and state observation. Controlled by the CLI enable_interface12option (e.g. openenv push --enable-interface) or ENABLE_WEB_INTERFACE env var.13"""14 15from __future__ import annotations16 17import asyncio18import inspect19import json20from concurrent.futures import ThreadPoolExecutor21from datetime import datetime22from typing import Any, Callable, Dict, List, Optional, Type23 24import gradio as gr25from fastapi import Body, FastAPI, HTTPException, status, WebSocket, WebSocketDisconnect26from fastapi.responses import RedirectResponse27from pydantic import BaseModel, ConfigDict, Field28 29from .gradio_theme import OPENENV_GRADIO_CSS, OPENENV_GRADIO_THEME30from .gradio_ui import build_gradio_app, get_gradio_display_title31from .interfaces import Environment32from .serialization import deserialize_action_with_preprocessing, serialize_observation33from .types import Action, EnvironmentMetadata, Observation, State34 35# Quick Start markdown template; placeholders match init suffixes (__ENV_NAME__, __ENV_CLASS_NAME__*).36DEFAULT_QUICK_START_MARKDOWN = """37### Connect to this environment38 39Connect from Python using `__ENV_CLASS_NAME__Env`:40 41```python42from __ENV_NAME__ import __ENV_CLASS_NAME__Action, __ENV_CLASS_NAME__Env43 44with __ENV_CLASS_NAME__Env.from_env("<SPACE_ID>") as env:45    result = await env.step(__ENV_CLASS_NAME__Action(message="..."))46```47 48Or connect directly to a running server:49 50```python51env = __ENV_CLASS_NAME__Env(base_url="http://localhost:8000")52```53 54### Contribute to this environment55 56Submit improvements via pull request on the Hugging Face Hub.57 58```bash59openenv fork <SPACE_ID> --repo-id <your-username>/<your-repo-name>60```61 62Then make your changes and submit a pull request:63 64```bash65cd <forked-repo>66openenv push <SPACE_ID> --create-pr67```68 69For more information, see the [OpenEnv documentation](https://meta-pytorch.org/OpenEnv/).70"""71 72 73def get_quick_start_markdown(74    metadata: Optional[EnvironmentMetadata],75    action_cls: Type[Action],76    observation_cls: Type[Observation],77) -> str:78    """79    Build Quick Start markdown with class names replaced from current env (init-style suffixes).80 81    Uses the same placeholder names as the init template so that __ENV_CLASS_NAME__Env,82    __ENV_CLASS_NAME__Action, __ENV_CLASS_NAME__Observation and __ENV_NAME__ are83    replaced with the actual class/package names.84    """85    import os86 87    # Prefix from action class (e.g. EchoAction -> Echo)88    action_name = getattr(action_cls, "__name__", "Action")89    if action_name.endswith("Action"):90        prefix = action_name[: -len("Action")]91    else:92        prefix = action_name.replace("Action", "").strip() or "Env"93 94    env_client_name = f"{prefix}Env"95    obs_name = getattr(observation_cls, "__name__", "Observation")96    pkg_name = (metadata.name if metadata else "env").replace(" ", "_").lower()97 98    space_id = os.environ.get("SPACE_ID", "<hf-username>/<hf-repo-name>")99 100    content = DEFAULT_QUICK_START_MARKDOWN101    content = content.replace("__ENV_CLASS_NAME__Env", env_client_name)102    content = content.replace("__ENV_CLASS_NAME__Action", action_name)103    content = content.replace("__ENV_CLASS_NAME__Observation", obs_name)104    content = content.replace("__ENV_CLASS_NAME__", prefix)105    content = content.replace("__ENV_NAME__", pkg_name)106    content = content.replace("<SPACE_ID>", space_id)107    return content.strip()108 109 110def load_environment_metadata(111    env: Environment, env_name: Optional[str] = None112) -> EnvironmentMetadata:113    """114    Load environment metadata including README content.115 116    Args:117        env: The environment instance, class, or factory function.118             - If a class: used as a factory, won't call instance methods119             - If a function: used as a factory, won't call instance methods120             - If an instance: may call get_metadata() if available121        env_name: Optional environment name for README file lookup122 123    Returns:124        EnvironmentMetadata with loaded information125    """126    import inspect127 128    # Determine what type of env we received:129    # 1. A class (used as factory) - e.g., PythonCodeActEnv130    # 2. A function (factory function) - e.g., create_chat_environment131    # 3. An actual instance - e.g., SnakeEnvironment()132    is_class = inspect.isclass(env)133    is_function = inspect.isfunction(env) or inspect.ismethod(env)134    is_factory = is_class or is_function135 136    # Try to get metadata from environment if it's an instance with get_metadata137    if not is_factory and hasattr(env, "get_metadata"):138        return env.get_metadata()139 140    # Determine the class name for default metadata141    if is_class:142        # env is the class itself143        class_name = env.__name__144    elif is_function:145        # env is a factory function - use its name or derive from env_name146        class_name = env_name or env.__name__147    else:148        # env is an instance149        class_name = env.__class__.__name__150 151    # Default metadata152    metadata = EnvironmentMetadata(153        name=env_name or class_name,154        description=f"{class_name} environment",155        version="1.0.0",156    )157 158    # Try to load README from file system159    readme_content = _load_readme_from_filesystem(env_name)160    if readme_content:161        metadata.readme_content = readme_content162 163    return metadata164 165 166def _load_readme_from_filesystem(env_name: Optional[str]) -> Optional[str]:167    """168    Load README content from the filesystem.169 170    Tries multiple locations:171    1. Container filesystem: /app/README.md172    2. Local development: src/envs/{env_name}/README.md173    3. Environment variable: ENV_README_PATH174    """175    import os176    from pathlib import Path177 178    # Try container filesystem first179    container_readme = Path("/app/README.md")180    if container_readme.exists():181        try:182            return container_readme.read_text(encoding="utf-8")183        except Exception:184            pass185 186    # Try environment variable path187    custom_path = os.environ.get("ENV_README_PATH")188    if custom_path and Path(custom_path).exists():189        try:190            return Path(custom_path).read_text(encoding="utf-8")191        except Exception:192            pass193 194    # Try local development path195    if env_name:196        local_readme = Path(f"src/envs/{env_name}/README.md")197        if local_readme.exists():198            try:199                return local_readme.read_text(encoding="utf-8")200            except Exception:201                pass202 203    return None204 205 206class ActionLog(BaseModel):207    """Log entry for an action taken."""208 209    model_config = ConfigDict(extra="forbid", validate_assignment=True)210 211    timestamp: str = Field(description="Timestamp when action was taken")212    action: Dict[str, Any] = Field(description="Action that was taken")213    observation: Dict[str, Any] = Field(description="Observation returned from action")214    reward: Optional[float] = Field(215        default=None, description="Reward received from action"216    )217    done: bool = Field(description="Whether the episode is done after this action")218    step_count: int = Field(description="Step count when this action was taken")219 220 221class EpisodeState(BaseModel):222    """Current episode state for the web interface."""223 224    model_config = ConfigDict(extra="forbid", validate_assignment=True)225 226    episode_id: Optional[str] = Field(default=None, description="Current episode ID")227    step_count: int = Field(description="Current step count in episode")228    current_observation: Optional[Dict[str, Any]] = Field(229        default=None, description="Current observation"230    )231    action_logs: List[ActionLog] = Field(232        default_factory=list, description="List of action logs"233    )234    is_reset: bool = Field(235        default=True, description="Whether the episode has been reset"236    )237 238 239class WebInterfaceManager:240    """Manages the web interface for an environment."""241 242    MAX_ACTION_LOGS = 1000243 244    def __init__(245        self,246        env: Environment,247        action_cls: Type[Action],248        observation_cls: Type[Observation],249        metadata: Optional[EnvironmentMetadata] = None,250    ):251        import inspect252 253        # If env is a class or factory function, instantiate it254        if inspect.isclass(env) or inspect.isfunction(env):255            self.env = env()256        else:257            self.env = env258        self.action_cls = action_cls259        self.observation_cls = observation_cls260        self.metadata = metadata or EnvironmentMetadata(261            name=env.__class__.__name__,262            description=f"{env.__class__.__name__} environment",263        )264        self.episode_state = EpisodeState(265            episode_id=None,266            step_count=0,267            current_observation=None,268            action_logs=[],269        )270        self.connected_clients: List[WebSocket] = []271        # Thread pool for running sync code (e.g., Playwright sync API) in async context272        self._executor = ThreadPoolExecutor(max_workers=1)273 274    @staticmethod275    def _get_valid_kwargs(276        sig: inspect.Signature,277        kwargs: Dict[str, Any],278        skip_params: Optional[set[str]] = None,279    ) -> Dict[str, Any]:280        """Filter kwargs to only those accepted by the target function."""281        skip_params = skip_params or set()282        valid_kwargs: Dict[str, Any] = {}283        has_var_kwargs = any(284            param.kind == inspect.Parameter.VAR_KEYWORD285            for param in sig.parameters.values()286        )287 288        for key, value in kwargs.items():289            if key in skip_params:290                continue291            if key in sig.parameters or has_var_kwargs:292                valid_kwargs[key] = value293 294        return valid_kwargs295 296    async def _run_sync_in_thread_pool(self, func, *args, **kwargs):297        """Run a synchronous function in the thread pool executor.298 299        This is needed for environments using sync libraries (e.g., Playwright sync API)300        that cannot be called directly from an async context.301        """302        loop = asyncio.get_event_loop()303        # Use default arguments to capture values at lambda definition time304        # to avoid closure issues with late binding305        return await loop.run_in_executor(306            self._executor, lambda f=func, a=args, kw=kwargs: f(*a, **kw)307        )308 309    async def connect_websocket(self, websocket: WebSocket):310        """Connect a new WebSocket client."""311        await websocket.accept()312        self.connected_clients.append(websocket)313 314        # Send current state to the new client315        await self._send_state_update()316 317    async def disconnect_websocket(self, websocket: WebSocket):318        """Disconnect a WebSocket client."""319        if websocket in self.connected_clients:320            self.connected_clients.remove(websocket)321 322    async def _send_state_update(self):323        """Send current state to all connected clients."""324        if not self.connected_clients:325            return326 327        state_data = {328            "type": "state_update",329            "episode_state": self.episode_state.model_dump(),330        }331 332        # Send to all connected clients333        disconnected_clients = []334        for client in self.connected_clients:335            try:336                await client.send_text(json.dumps(state_data))337            except Exception:338                disconnected_clients.append(client)339 340        # Remove disconnected clients341        for client in disconnected_clients:342            self.connected_clients.remove(client)343 344    async def reset_environment(345        self, reset_kwargs: Optional[Dict[str, Any]] = None346    ) -> Dict[str, Any]:347        """Reset the environment and update state."""348        reset_kwargs = reset_kwargs or {}349 350        is_async = self.env.reset_async.__func__ is not Environment.reset_async351        sig = inspect.signature(self.env.reset_async if is_async else self.env.reset)352        valid_kwargs = self._get_valid_kwargs(sig, reset_kwargs)353 354        if is_async:355            observation = await self.env.reset_async(**valid_kwargs)356        else:357            # Run sync reset in thread pool to avoid blocking event loop358            # and to support environments using sync libraries (e.g., Playwright)359            observation = await self._run_sync_in_thread_pool(360                self.env.reset, **valid_kwargs361            )362        state: State = self.env.state363 364        # Serialize observation once using shared utility365        serialized = serialize_observation(observation)366 367        # Update episode state368        self.episode_state.episode_id = state.episode_id369        self.episode_state.step_count = 0370        self.episode_state.current_observation = serialized["observation"]371        self.episode_state.action_logs = []372        self.episode_state.is_reset = True373 374        # Send state update375        await self._send_state_update()376 377        return serialized378 379    async def step_environment(self, action_data: Dict[str, Any]) -> Dict[str, Any]:380        """Execute a step in the environment and update state."""381        # Deserialize action with preprocessing for web interface special cases382        action: Action = deserialize_action_with_preprocessing(383            action_data, self.action_cls384        )385 386        # Run sync step in thread pool to avoid blocking event loop387        # and to support environments using sync libraries (e.g., Playwright)388        observation: Observation = await self._run_sync_in_thread_pool(389            self.env.step, action390        )391        state: State = self.env.state392 393        # Serialize observation once using shared utility394        serialized = serialize_observation(observation)395 396        # Create action log397        action_log = ActionLog(398            timestamp=datetime.now().isoformat(),399            action=action.model_dump(exclude={"metadata"}),400            observation=serialized["observation"],401            reward=observation.reward,402            done=observation.done,403            step_count=state.step_count,404        )405 406        # Update episode state407        self.episode_state.episode_id = state.episode_id408        self.episode_state.step_count = state.step_count409        self.episode_state.current_observation = serialized["observation"]410        self.episode_state.action_logs.append(action_log)411        if len(self.episode_state.action_logs) > self.MAX_ACTION_LOGS:412            self.episode_state.action_logs = self.episode_state.action_logs[413                -self.MAX_ACTION_LOGS :414            ]415        self.episode_state.is_reset = False416 417        # Send state update418        await self._send_state_update()419 420        return serialized421 422    def get_state(self) -> Dict[str, Any]:423        """Get current environment state."""424        state: State = self.env.state425        return state.model_dump()426 427 428def create_web_interface_app(429    env: Environment,430    action_cls: Type[Action],431    observation_cls: Type[Observation],432    env_name: Optional[str] = None,433    max_concurrent_envs: Optional[int] = None,434    concurrency_config: Optional[Any] = None,435    gradio_builder: Optional[Callable[..., Any]] = None,436) -> FastAPI:437    """438    Create a FastAPI application with web interface for the given environment.439 440    Args:441        env: The Environment instance to serve442        action_cls: The Action subclass this environment expects443        observation_cls: The Observation subclass this environment returns444        env_name: Optional environment name for README loading445        max_concurrent_envs: Maximum concurrent WebSocket sessions446        concurrency_config: Optional ConcurrencyConfig for advanced concurrency settings447        gradio_builder: Optional callable (web_manager, action_fields, metadata,448            is_chat_env, title, quick_start_md) -> gr.Blocks to use instead of the449            default Gradio UI. Lets envs replace or customize the /web interface.450 451    Returns:452        FastAPI application instance with web interface453    """454    from .http_server import create_fastapi_app455 456    # Create the base environment app457    app = create_fastapi_app(458        env, action_cls, observation_cls, max_concurrent_envs, concurrency_config459    )460 461    # Load environment metadata462    metadata = load_environment_metadata(env, env_name)463 464    # Create web interface manager465    web_manager = WebInterfaceManager(env, action_cls, observation_cls, metadata)466 467    # Web API routes first (so they take precedence over Gradio mount at /web)468    @app.get("/", include_in_schema=False)469    async def web_root():470        """Redirect the app root to the Gradio interface."""471        return RedirectResponse(url="/web/")472 473    @app.get("/web", include_in_schema=False)474    async def web_root_no_slash():475        """Redirect /web to /web/ for mounted Gradio deployments behind proxies."""476        return RedirectResponse(url="/web/")477 478    @app.get("/web/metadata")479    async def web_metadata():480        """Get environment metadata."""481        return web_manager.metadata.model_dump()482 483    @app.websocket("/ws/ui")484    async def websocket_ui_endpoint(websocket: WebSocket):485        """WebSocket endpoint for web UI real-time updates.486 487        Note: Uses /ws/ui to avoid conflict with /ws in http_server.py488        which is used for concurrent environment sessions.489        """490        await web_manager.connect_websocket(websocket)491        try:492            while True:493                # Keep connection alive494                await websocket.receive_text()495        except WebSocketDisconnect:496            await web_manager.disconnect_websocket(websocket)497 498    @app.post("/web/reset")499    async def web_reset(request: Optional[Dict[str, Any]] = Body(default=None)):500        """Reset endpoint for web interface."""501        return await web_manager.reset_environment(request)502 503    @app.post("/web/step")504    async def web_step(request: Dict[str, Any]):505        """Step endpoint for web interface."""506        # Check if this is a message-based request (chat environment)507        if "message" in request:508            message = request["message"]509            if hasattr(web_manager.env, "message_to_action"):510                action = web_manager.env.message_to_action(message)511                if hasattr(action, "tokens"):512                    action_data = {"tokens": action.tokens.tolist()}513                else:514                    action_data = action.model_dump(exclude={"metadata"})515            else:516                action_data = {"message": message}517        else:518            action_data = request.get("action", {})519 520        return await web_manager.step_environment(action_data)521 522    @app.get("/web/state")523    async def web_state():524        """State endpoint for web interface."""525        try:526            return web_manager.get_state()527        except RuntimeError as exc:528            raise HTTPException(529                status_code=status.HTTP_409_CONFLICT,530                detail=str(exc),531            ) from exc532 533    action_fields = _extract_action_fields(action_cls)534    is_chat_env = _is_chat_env(action_cls)535    quick_start_md = get_quick_start_markdown(metadata, action_cls, observation_cls)536 537    default_blocks = build_gradio_app(538        web_manager,539        action_fields,540        metadata,541        is_chat_env,542        title=metadata.name,543        quick_start_md=quick_start_md,544    )545    if gradio_builder is not None:546        custom_blocks = gradio_builder(547            web_manager,548            action_fields,549            metadata,550            is_chat_env,551            metadata.name,552            quick_start_md,553        )554        if not isinstance(custom_blocks, gr.Blocks):555            raise TypeError(556                f"gradio_builder must return a gr.Blocks instance, "557                f"got {type(custom_blocks).__name__}"558            )559        gradio_blocks = gr.TabbedInterface(560            [default_blocks, custom_blocks],561            tab_names=["Playground", "Custom"],562            title=get_gradio_display_title(metadata),563        )564    else:565        gradio_blocks = default_blocks566    app = gr.mount_gradio_app(567        app,568        gradio_blocks,569        path="/web",570        theme=OPENENV_GRADIO_THEME,571        css=OPENENV_GRADIO_CSS,572    )573 574    return app575 576 577def _is_chat_env(action_cls: Type[Action]) -> bool:578    """Return True if the action class is a chat-style env (tokens field)."""579    if hasattr(action_cls, "model_fields"):580        for field_name, field_info in action_cls.model_fields.items():581            if (582                field_name == "tokens"583                and hasattr(field_info.annotation, "__name__")584                and "Tensor" in str(field_info.annotation)585            ):586                return True587    return False588 589 590def _extract_action_fields(action_cls: Type[Action]) -> List[Dict[str, Any]]:591    """Extract enhanced field metadata from Action class for form generation."""592    # Use Pydantic's JSON schema generation for robust metadata extraction593    try:594        schema = action_cls.model_json_schema()595    except AttributeError:596        # Fallback for non-Pydantic v2 models or if something goes wrong597        return []598 599    properties = schema.get("properties", {})600    required_fields = schema.get("required", [])601 602    action_fields = []603 604    for field_name, field_info in properties.items():605        if field_name == "metadata":606            continue607 608        # JSON schema "type" can be a string or list/undefined609        # Determine our internal input type610        input_type = _determine_input_type_from_schema(field_info, field_name)611 612        is_required = field_name in required_fields613 614        action_fields.append(615            {616                "name": field_name,617                "type": input_type,618                "required": is_required,619                "description": field_info.get("description", ""),620                "default_value": field_info.get("default"),621                "choices": field_info.get("enum"),622                "min_value": field_info.get("minimum"),623                "max_value": field_info.get("maximum"),624                "min_length": field_info.get("minLength"),625                "max_length": field_info.get("maxLength"),626                "pattern": field_info.get("pattern"),627                "placeholder": _generate_placeholder(field_name, field_info),628                "help_text": _generate_help_text(field_name, field_info),629            }630        )631 632    return action_fields633 634 635def _determine_input_type_from_schema(636    field_info: Dict[str, Any], field_name: str637) -> str:638    """Determine input type from JSON schema for form generation (Gradio UI)."""639    schema_type = field_info.get("type")640 641    # Check for specific tensor field convention642    if "tokens" in field_name.lower():643        return "tensor"644 645    if "enum" in field_info:646        return "select"647 648    if schema_type == "boolean":649        return "checkbox"650 651    if schema_type == "integer" or schema_type == "number":652        return "number"653 654    if schema_type == "string":655        # Check if it should be a textarea656        if (657            field_info.get("maxLength", 0) > 100658            or "message" in field_name.lower()659            or "code" in field_name.lower()660        ):661            return "textarea"662        return "text"663 664    # Default fallback665    return "text"666 667 668def _generate_placeholder(field_name: str, field_info: Dict[str, Any]) -> str:669    """Generate placeholder text."""670    if "message" in field_name.lower():671        return f"Enter {field_name.replace('_', ' ')}..."672    elif "code" in field_name.lower():673        return "Enter Python code here..."674    elif "tokens" in field_name.lower():675        return "Enter comma-separated token IDs (e.g., 1,2,3,4,5)"676    else:677        return f"Enter {field_name.replace('_', ' ')}..."678 679 680def _generate_help_text(field_name: str, field_info: Dict[str, Any]) -> str:681    """Generate help text."""682    description = field_info.get("description", "")683    if description:684        return description685 686    if "action_id" in field_name.lower():687        return "The action ID to execute in environment"688    elif "game_name" in field_name.lower():689        return "Name of game or environment"690    elif "tokens" in field_name.lower():691        return "Token IDs as a comma-separated list of integers"692    elif "code" in field_name.lower():693        return "Python code to execute in environment"694    elif "message" in field_name.lower():695        return "Text message to send"696 697    return ""698