openenv/coding_env
21
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 