openenv/repl
1
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 