Team Ai
Apppublic

openenv/atari_env

sourceHugging Faceupdated 6mo agoView on Hugging Face
3likes
gradio_ui.py241 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"""8Gradio-based web UI for OpenEnv environments.9 10Replaces the legacy HTML/JavaScript interface when ENABLE_WEB_INTERFACE is set.11Mount at /web via gr.mount_gradio_app() from create_web_interface_app().12"""13 14from __future__ import annotations15 16import json17import re18from typing import Any, Dict, List, Optional19 20import gradio as gr21 22from .types import EnvironmentMetadata23 24 25def _escape_md(text: str) -> str:26    """Escape Markdown special characters in user-controlled content."""27    return re.sub(r"([\\`*_\{\}\[\]()#+\-.!|~>])", r"\\\1", str(text))28 29 30def _format_observation(data: Dict[str, Any]) -> str:31    """Format reset/step response for Markdown display."""32    lines: List[str] = []33    obs = data.get("observation", {})34    if isinstance(obs, dict):35        if obs.get("prompt"):36            lines.append(f"**Prompt:**\n\n{_escape_md(obs['prompt'])}\n")37        messages = obs.get("messages", [])38        if messages:39            lines.append("**Messages:**\n")40            for msg in messages:41                sender = _escape_md(str(msg.get("sender_id", "?")))42                content = _escape_md(str(msg.get("content", "")))43                cat = _escape_md(str(msg.get("category", "")))44                lines.append(f"- `[{cat}]` Player {sender}: {content}")45            lines.append("")46    reward = data.get("reward")47    done = data.get("done")48    if reward is not None:49        lines.append(f"**Reward:** `{reward}`")50    if done is not None:51        lines.append(f"**Done:** `{done}`")52    return "\n".join(lines) if lines else "*No observation data*"53 54 55def _readme_section(metadata: Optional[EnvironmentMetadata]) -> str:56    """README content for the left panel."""57    if not metadata or not metadata.readme_content:58        return "*No README available.*"59    return metadata.readme_content60 61 62def get_gradio_display_title(63    metadata: Optional[EnvironmentMetadata],64    fallback: str = "OpenEnv Environment",65) -> str:66    """Return the title used for the Gradio app (browser tab and Blocks)."""67    name = metadata.name if metadata else fallback68    return f"OpenEnv Agentic Environment: {name}"69 70 71def build_gradio_app(72    web_manager: Any,73    action_fields: List[Dict[str, Any]],74    metadata: Optional[EnvironmentMetadata],75    is_chat_env: bool,76    title: str = "OpenEnv Environment",77    quick_start_md: Optional[str] = None,78) -> gr.Blocks:79    """80    Build a Gradio Blocks app for the OpenEnv web interface.81 82    Args:83        web_manager: WebInterfaceManager (reset/step_environment, get_state).84        action_fields: Field dicts from _extract_action_fields(action_cls).85        metadata: Environment metadata for README/name.86        is_chat_env: If True, single message textbox; else form from action_fields.87        title: App title (overridden by metadata.name when present; see get_gradio_display_title).88        quick_start_md: Optional Quick Start markdown (class names already replaced).89 90    Returns:91        gr.Blocks to mount with gr.mount_gradio_app(app, blocks, path="/web").92    """93    readme_content = _readme_section(metadata)94    display_title = get_gradio_display_title(metadata, fallback=title)95 96    async def reset_env():97        try:98            data = await web_manager.reset_environment()99            obs_md = _format_observation(data)100            return (101                obs_md,102                json.dumps(data, indent=2),103                "Environment reset successfully.",104            )105        except Exception as e:106            return ("", "", f"Error: {e}")107 108    def _step_with_action(action_data: Dict[str, Any]):109        async def _run():110            try:111                data = await web_manager.step_environment(action_data)112                obs_md = _format_observation(data)113                return (114                    obs_md,115                    json.dumps(data, indent=2),116                    "Step complete.",117                )118            except Exception as e:119                return ("", "", f"Error: {e}")120 121        return _run122 123    async def step_chat(message: str):124        if not (message or str(message).strip()):125            return ("", "", "Please enter an action message.")126        action = {"message": str(message).strip()}127        return await _step_with_action(action)()128 129    def get_state_sync():130        try:131            data = web_manager.get_state()132            return json.dumps(data, indent=2)133        except Exception as e:134            return f"Error: {e}"135 136    with gr.Blocks(title=display_title) as demo:137        with gr.Row():138            with gr.Column(scale=1, elem_classes="col-left"):139                if quick_start_md:140                    with gr.Accordion("Quick Start", open=True):141                        gr.Markdown(quick_start_md)142                with gr.Accordion("README", open=False):143                    gr.Markdown(readme_content)144 145            with gr.Column(scale=2, elem_classes="col-right"):146                obs_display = gr.Markdown(147                    value=("# Playground\n\nClick **Reset** to start a new episode."),148                )149                with gr.Group():150                    if is_chat_env:151                        action_input = gr.Textbox(152                            label="Action message",153                            placeholder="e.g. Enter your message...",154                        )155                        step_inputs = [action_input]156                        step_fn = step_chat157                    else:158                        step_inputs = []159                        for field in action_fields:160                            name = field["name"]161                            field_type = field.get("type", "text")162                            label = name.replace("_", " ").title()163                            placeholder = field.get("placeholder", "")164                            if field_type == "checkbox":165                                inp = gr.Checkbox(label=label)166                            elif field_type == "number":167                                inp = gr.Number(label=label)168                            elif field_type == "select":169                                choices = field.get("choices") or []170                                inp = gr.Dropdown(171                                    choices=choices,172                                    label=label,173                                    allow_custom_value=False,174                                )175                            elif field_type in ("textarea", "tensor"):176                                inp = gr.Textbox(177                                    label=label,178                                    placeholder=placeholder,179                                    lines=3,180                                )181                            else:182                                inp = gr.Textbox(183                                    label=label,184                                    placeholder=placeholder,185                                )186                            step_inputs.append(inp)187 188                        async def step_form(*values):189                            if not action_fields:190                                return await _step_with_action({})()191                            action_data = {}192                            for i, field in enumerate(action_fields):193                                if i >= len(values):194                                    break195                                name = field["name"]196                                val = values[i]197                                if field.get("type") == "checkbox":198                                    action_data[name] = bool(val)199                                elif val is not None and val != "":200                                    action_data[name] = val201                            return await _step_with_action(action_data)()202 203                        step_fn = step_form204 205                    with gr.Row():206                        step_btn = gr.Button("Step", variant="primary")207                        reset_btn = gr.Button("Reset", variant="secondary")208                        state_btn = gr.Button("Get state", variant="secondary")209                    with gr.Row():210                        status = gr.Textbox(211                            label="Status",212                            interactive=False,213                        )214                    raw_json = gr.Code(215                        label="Raw JSON response",216                        language="json",217                        interactive=False,218                    )219 220        reset_btn.click(221            fn=reset_env,222            outputs=[obs_display, raw_json, status],223        )224        step_btn.click(225            fn=step_fn,226            inputs=step_inputs,227            outputs=[obs_display, raw_json, status],228        )229        if is_chat_env:230            action_input.submit(231                fn=step_fn,232                inputs=step_inputs,233                outputs=[obs_display, raw_json, status],234            )235        state_btn.click(236            fn=get_state_sync,237            outputs=[raw_json],238        )239 240    return demo241