Team Ai
Apppublic

openenv/repl

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
python_executor.py271 linesDownload Raw Back to 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"""8Sandboxed Python code executor for the REPL environment.9 10Uses smolagents.LocalPythonExecutor as the backend for battle-tested sandboxed11execution, with RLM-specific features on top:12- Context loading (set_context)13- Variable access (get_variable, list_variables)14- Function injection (inject_function for llm_query, llm_query_batched)15- Output capped at 8,192 characters per turn (configurable)16- Persistent namespace across code blocks17"""18 19import json20import logging21import time22import traceback23from collections.abc import Callable24from typing import Any, Dict, List, Optional25 26from smolagents import LocalPythonExecutor27 28logger = logging.getLogger(__name__)29logger.addHandler(logging.NullHandler())30 31 32class PythonExecutor:33    """Sandboxed Python code executor with persistent namespace.34 35    Wraps smolagents.LocalPythonExecutor with RLM-specific features:36    - Context loading for RLM tasks37    - Variable tracking for observation38    - Function injection for llm_query, llm_query_batched39    - Configurable output length limit (default 8192 chars per Prime Intellect)40    """41 42    def __init__(43        self,44        max_output_length: int = 8192,45        allowed_imports: Optional[List[str]] = None,46    ):47        """Initialize the executor.48 49        Args:50            max_output_length: Maximum characters for stdout/stderr (default 8192)51            allowed_imports: List of allowed module names for import52 53        Note:54            smolagents.LocalPythonExecutor does NOT support wall-clock timeouts.55            Instead, it limits operations (10M ops) and while iterations (1M).56        """57        self.max_output_length = max_output_length58 59        # Default allowed imports for RLM tasks60        default_imports = [61            "re",62            "json",63            "math",64            "random",65            "collections",66            "itertools",67            "functools",68            "operator",69            "string",70            "textwrap",71            "difflib",72            "statistics",73            "decimal",74            "fractions",75            "datetime",76            "copy",77            "pprint",78            "typing",79            "dataclasses",80            "enum",81            "bisect",82            "heapq",83            "array",84            "struct",85            "base64",86            "hashlib",87            "hmac",88            "uuid",89        ]90 91        self.allowed_imports = allowed_imports or default_imports92 93        # Initialize the smolagents executor94        self._executor = LocalPythonExecutor(95            additional_authorized_imports=self.allowed_imports96        )97 98        # Track variables we've set (for list_variables)99        self._user_variables: set[str] = set()100 101        # Track callable functions to register with send_tools102        self._callable_tools: Dict[str, Callable[..., Any]] = {}103 104        # Register helper utilities105        self._register_helpers()106 107    def _register_helpers(self) -> None:108        """Register helper functions with the executor."""109        helpers = {110            "format_exc": traceback.format_exc,111            "safe_json_dumps": lambda obj: json.dumps(obj, default=lambda o: repr(o)),112        }113        # Register helpers as callable tools114        for name, func in helpers.items():115            self.inject_function(name, func)116 117    def _sync_callable_tools(self) -> None:118        """Sync callable functions with the executor via send_tools."""119        if self._callable_tools:120            try:121                # Type ignore: smolagents accepts callables despite Tool type hint122                self._executor.send_tools(self._callable_tools)  # type: ignore[arg-type]123            except Exception:124                logger.debug(125                    "send_tools failed; continuing without extra tools",126                    exc_info=True,127                )128 129    def set_context(self, context: str, variable_name: str = "context") -> None:130        """Load context into namespace as a variable.131 132        Args:133            context: The context string to load134            variable_name: Name of the variable (default "context")135        """136        self.set_variable(variable_name, context)137 138    def set_variable(self, name: str, value: Any) -> None:139        """Set a variable in the namespace.140 141        Args:142            name: Variable name143            value: Variable value144        """145        self._executor.send_variables({name: value})146        self._user_variables.add(name)147 148    def get_variable(self, name: str) -> Optional[Any]:149        """Retrieve a variable from namespace.150 151        Args:152            name: Variable name153 154        Returns:155            The variable value or None if not found156        """157        return self._executor.state.get(name)158 159    def list_variables(self) -> List[str]:160        """List non-private variables in namespace.161 162        Returns:163            List of variable names (excluding private and builtins)164        """165        variables = {key for key in self._executor.state if not key.startswith("_")}166        variables.update(self._user_variables)167        return list(variables)168 169    def execute(self, code: str) -> Dict[str, Any]:170        """Execute Python code and return results.171 172        Args:173            code: Python code to execute174 175        Returns:176            Dictionary with stdout, stderr, locals_snapshot, execution_time,177            success, and exception fields178        """179        start_time = time.time()180        success = True181        exception_msg = None182        new_locals: Dict[str, str] = {}183 184        # Track state before execution185        pre_state_keys = set()186        if hasattr(self._executor, "state"):187            pre_state_keys = set(self._executor.state.keys())188 189        stdout_parts: list[str] = []190        stderr_parts: list[str] = []191 192        try:193            exec_result = self._executor(code)194 195            # CodeOutput has: logs (str), output (Any), is_final_answer (bool)196            if exec_result.logs:197                stdout_parts.append(str(exec_result.logs))198 199            if exec_result.output is not None:200                try:201                    stdout_parts.append(json.dumps(exec_result.output))202                except Exception:203                    stdout_parts.append(repr(exec_result.output))204 205        except Exception as e:206            success = False207            exception_msg = f"{type(e).__name__}: {str(e)}\n{traceback.format_exc()}"208            stderr_parts.append(exception_msg)209 210        execution_time = time.time() - start_time211 212        # Capture new/modified variables213        for key in self._executor.state:214            if key not in pre_state_keys and not key.startswith("_"):215                try:216                    val = self._executor.state[key]217                    val_repr = repr(val)218                    if len(val_repr) > 500:219                        val_repr = val_repr[:500] + "..."220                    new_locals[key] = val_repr221                    self._user_variables.add(key)222                except Exception:223                    new_locals[key] = "<unrepresentable>"224 225        # Compose stdout/stderr226        stdout = "\n".join(part for part in stdout_parts if part)227        stderr = "\n".join(part for part in stderr_parts if part)228 229        # Truncate output to max_output_length230        if len(stdout) > self.max_output_length:231            stdout = (232                stdout[: self.max_output_length]233                + f"\n... (truncated, total {len(stdout)} chars)"234            )235 236        if len(stderr) > self.max_output_length:237            stderr = (238                stderr[: self.max_output_length]239                + f"\n... (truncated, total {len(stderr)} chars)"240            )241 242        return {243            "stdout": stdout,244            "stderr": stderr,245            "locals_snapshot": new_locals,246            "execution_time": execution_time,247            "success": success,248            "exception": exception_msg,249        }250 251    def reset(self) -> None:252        """Reset namespace to initial state."""253        # Create a new executor instance254        self._executor = LocalPythonExecutor(255            additional_authorized_imports=self.allowed_imports256        )257        self._user_variables.clear()258        self._callable_tools.clear()259        self._register_helpers()260 261    def inject_function(self, name: str, func: Callable[..., Any]) -> None:262        """Inject a callable function into the namespace.263 264        Args:265            name: Function name in namespace266            func: The callable to inject267        """268        self._callable_tools[name] = func269        self._user_variables.add(name)270        self._sync_callable_tools()271