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"""8REPL Environment Implementation.9 10A Python REPL environment for training language models on code execution tasks,11based on the Recursive Language Models (RLM) paradigm.12 13References:14- RLM Paper: https://arxiv.org/abs/2512.2460115- Prime Intellect Blog: https://www.primeintellect.ai/blog/rlm16- Alex Zhang Blog: https://alexzhang13.github.io/blog/2025/rlm/17"""18 19import os20import re21from collections.abc import Callable22from typing import Any, List, Optional23from uuid import uuid424 25try:26 from openenv.core.env_server.interfaces import Environment27 from openenv.core.env_server.types import EnvironmentMetadata28except ImportError:29 from openenv.core.env_server.interfaces import Environment30 from openenv.core.env_server.types import EnvironmentMetadata31 32try:33 from ..models import CodeBlockResult, REPLAction, REPLObservation, REPLState34except ImportError:35 try:36 from repl_env.models import CodeBlockResult, REPLAction, REPLObservation, REPLState37 except ImportError:38 from models import CodeBlockResult, REPLAction, REPLObservation, REPLState39 40try:41 from ..recursive_controller import create_server_recursive_controller42 from ..rubrics import REPLRubric43 from .python_executor import PythonExecutor44except ImportError:45 try:46 from repl_env.recursive_controller import create_server_recursive_controller47 from repl_env.rubrics import REPLRubric48 from .python_executor import PythonExecutor49 except ImportError:50 from .python_executor import PythonExecutor51 from recursive_controller import create_server_recursive_controller52 from rubrics import REPLRubric53 54 55class REPLEnvironment(Environment):56 """57 A REPL environment for training language models to use code execution.58 59 Based on the Recursive Language Models (RLM) paradigm, this environment allows60 language models to:61 - Execute Python code in a sandboxed REPL62 - Work with large contexts loaded as variables63 - Finalize answers via FINAL(), FINAL_VAR(), or answer dict pattern64 - Optionally make recursive LLM calls via llm_query() / llm_query_batched()65 66 Supports two finalization patterns:67 1. RLM-style: print('FINAL(answer)') or print('FINAL_VAR(var_name)')68 2. Prime Intellect style: answer = {"content": "...", "ready": True}69 70 Example:71 >>> env = REPLEnvironment(context="Hello World", task_prompt="Count chars")72 >>> obs = env.reset()73 >>> print(obs.context_preview) # "Hello World"74 >>>75 >>> obs = env.step(REPLAction(code="result = len(context)"))76 >>> print(obs.result.success) # True77 >>> print(obs.available_variables) # ["context", "result", "answer"]78 >>>79 >>> obs = env.step(REPLAction(code="print(f'FINAL({result})')"))80 >>> print(obs.done) # True81 >>> print(obs.metadata["final_answer"]) # "11"82 """83 84 SUPPORTS_CONCURRENT_SESSIONS = True85 86 def __init__(87 self,88 context: Optional[str] = None,89 task_prompt: Optional[str] = None,90 max_iterations: int = 30,91 max_output_length: int = 8192,92 context_preview_length: int = 500,93 rubric: Optional[REPLRubric] = None,94 llm_query_fn: Optional[Callable[[str], str]] = None,95 llm_batch_fn: Optional[Callable[[List[str]], List[str]]] = None,96 subcall_fn: Optional[Callable[[str, Optional[str]], str]] = None,97 subcall_batch_fn: Optional[98 Callable[[List[str], Optional[str]], List[str]]99 ] = None,100 rlm_max_depth: int = 1,101 rlm_max_iterations: int | None = None,102 ):103 """Initialize the REPL environment.104 105 Args:106 context: Initial context to load (can also be set via REPL_CONTEXT env var)107 task_prompt: Task description (can also be set via REPL_TASK_PROMPT env var)108 max_iterations: Maximum steps per episode (default 30, env var REPL_MAX_ITERATIONS)109 max_output_length: Max chars for stdout/stderr per turn (default 8192)110 context_preview_length: Chars to show in context preview (default 500)111 rubric: Optional REPLRubric for reward computation (default: REPLRubric())112 llm_query_fn: Optional function for llm_query() support113 llm_batch_fn: Optional function for llm_query_batched() support114 subcall_fn: Optional function for recursive rlm_query() support115 subcall_batch_fn: Optional function for recursive rlm_query_batched() support116 rlm_max_depth: Max recursion depth for server-backed rlm_query()117 rlm_max_iterations: Max iterations for recursive child runners118 """119 self.initial_context = context or os.environ.get("REPL_CONTEXT", "")120 self.initial_task_prompt = task_prompt or os.environ.get("REPL_TASK_PROMPT", "")121 self.max_iterations = int(os.environ.get("REPL_MAX_ITERATIONS", max_iterations))122 self.max_output_length = max_output_length123 self.context_preview_length = context_preview_length124 125 # Rubric for reward computation (OpenEnv RFC 004)126 self.rubric = rubric or REPLRubric()127 128 # Optional LLM functions for recursive calls129 self.llm_query_fn = llm_query_fn130 self.llm_batch_fn = llm_batch_fn131 self.subcall_fn = subcall_fn132 self.subcall_batch_fn = subcall_batch_fn133 self.rlm_max_depth = rlm_max_depth134 self.rlm_max_iterations = rlm_max_iterations or max_iterations135 136 # State (initialized on reset)137 self._state: Optional[REPLState] = None138 self._executor: Optional[PythonExecutor] = None139 self._runtime_controller = None140 self._runtime_controller_chat_fn: Optional[Callable[..., str]] = None141 142 @staticmethod143 def _build_hf_chat_fn(144 hf_token: Optional[str] = None,145 llm_model: Optional[str] = None,146 ) -> Callable[..., str]:147 try:148 from huggingface_hub import InferenceClient, InferenceTimeoutError149 except ImportError:150 raise RuntimeError("huggingface_hub is required for HF-backed recursion")151 152 default_model = llm_model or os.environ.get("LLM_MODEL", "Qwen/Qwen3.5-9B")153 client = InferenceClient(model=default_model, token=hf_token, timeout=300)154 155 def chat_fn(messages: list[dict[str, str]], model: str | None = None) -> str:156 try:157 response = client.chat.completions.create(158 model=model or default_model,159 messages=messages,160 max_tokens=2048,161 # Qwen3.5 non-thinking mode for precise coding tasks (from model card)162 temperature=0.6,163 top_p=0.95,164 presence_penalty=0.0,165 extra_body={166 "top_k": 20,167 "min_p": 0.0,168 "repetition_penalty": 1.0,169 "chat_template_kwargs": {"enable_thinking": False},170 },171 )172 return response.choices[0].message.content or ""173 except InferenceTimeoutError:174 return "Error: LLM inference timed out"175 except Exception as e:176 return f"Error: {e}"177 178 return chat_fn179 180 def _create_llm_functions(181 self,182 hf_token: Optional[str],183 llm_model: Optional[str] = None,184 ) -> None:185 """Create LLM/subcall functions dynamically using client-provided token."""186 try:187 chat_fn = self._build_hf_chat_fn(hf_token, llm_model)188 except RuntimeError:189 return190 191 self._runtime_controller_chat_fn = chat_fn192 self._runtime_controller = create_server_recursive_controller(193 chat_fn,194 max_depth=self.rlm_max_depth,195 max_iterations=self.rlm_max_iterations,196 )197 self.llm_query_fn = self._runtime_controller.llm_query_fn198 self.llm_batch_fn = self._runtime_controller.llm_batch_fn199 self.subcall_fn = self._runtime_controller.rlm_query_fn200 self.subcall_batch_fn = self._runtime_controller.rlm_batch_fn201 202 def reset(203 self,204 seed: Optional[int] = None,205 episode_id: Optional[str] = None,206 context: Optional[str] = None,207 task_prompt: Optional[str] = None,208 hf_token: Optional[str] = None,209 llm_model: Optional[str] = None,210 **kwargs: Any,211 ) -> REPLObservation:212 """Reset the environment with optional new context.213 214 Args:215 seed: Optional random seed (for reproducibility)216 episode_id: Optional episode identifier (if not provided, one is generated)217 context: Context to load (overrides initial_context)218 task_prompt: Task description (overrides initial_task_prompt)219 hf_token: Optional HuggingFace token for llm_query/llm_query_batched.220 If provided, creates LLM functions using this token.221 Security: Token is NOT stored in state or logged.222 llm_model: Optional model name for LLM functions (default: from env or Qwen3.5-9B)223 **kwargs: Additional reset parameters including:224 expected_answer: Ground truth for rubric-based reward scoring225 rlm_max_depth: Override max recursion depth226 rlm_max_iterations: Override max iterations for recursive child runners227 228 Returns:229 Initial REPLObservation with environment ready message230 """231 effective_context = context or self.initial_context232 effective_task_prompt = task_prompt or self.initial_task_prompt233 234 # Set expected answer for rubric-based reward computation235 expected_answer = kwargs.get("expected_answer")236 self.rubric.reset()237 if expected_answer is not None:238 self.rubric.set_expected(expected_answer)239 240 runtime_rlm_max_depth = kwargs.get("rlm_max_depth")241 if runtime_rlm_max_depth is None:242 runtime_rlm_max_depth = self.rlm_max_depth243 runtime_rlm_max_depth = int(runtime_rlm_max_depth)244 245 runtime_rlm_max_iterations = kwargs.get("rlm_max_iterations")246 if runtime_rlm_max_iterations is None:247 runtime_rlm_max_iterations = self.rlm_max_iterations248 runtime_rlm_max_iterations = int(runtime_rlm_max_iterations)249 250 # Detect if recursion config changed — controller must be rebuilt251 depth_changed = (252 runtime_rlm_max_depth != self.rlm_max_depth253 or runtime_rlm_max_iterations != self.rlm_max_iterations254 )255 self.rlm_max_depth = runtime_rlm_max_depth256 self.rlm_max_iterations = runtime_rlm_max_iterations257 258 # Create or rebuild LLM functions when needed.259 # Token resolution: explicit hf_token > HF_TOKEN env var > cached HF login.260 if not self.llm_query_fn:261 effective_token = (262 hf_token if hf_token is not None else os.environ.get("HF_TOKEN")263 )264 self._create_llm_functions(effective_token, llm_model)265 elif depth_changed and self._runtime_controller is not None:266 # Rebuild controller with new depth/iteration config but reuse267 # the existing chat_fn — don't require re-providing credentials.268 self._runtime_controller.close()269 self._runtime_controller = create_server_recursive_controller(270 self._runtime_controller_chat_fn,271 max_depth=self.rlm_max_depth,272 max_iterations=self.rlm_max_iterations,273 )274 self.llm_query_fn = self._runtime_controller.llm_query_fn275 self.llm_batch_fn = self._runtime_controller.llm_batch_fn276 self.subcall_fn = self._runtime_controller.rlm_query_fn277 self.subcall_batch_fn = self._runtime_controller.rlm_batch_fn278 279 # Initialize state280 self._state = REPLState(281 episode_id=episode_id or str(uuid4()),282 step_count=0,283 context=effective_context,284 task_prompt=effective_task_prompt,285 iteration=0,286 max_iterations=self.max_iterations,287 namespace_keys=[],288 final_answer=None,289 total_execution_time=0.0,290 )291 292 # Initialize executor293 self._executor = PythonExecutor(max_output_length=self.max_output_length)294 295 # Initialize answer dict (Prime Intellect style)296 self._executor.set_variable("answer", {"content": "", "ready": False})297 298 # Load context into namespace if provided299 if effective_context:300 self._executor.set_context(effective_context)301 302 def _call_single_query(prompt: str, model: str | None = None) -> str:303 if not self.llm_query_fn:304 raise RuntimeError("llm_query is not configured")305 try:306 return self.llm_query_fn(prompt, model) # type: ignore[misc]307 except TypeError:308 return self.llm_query_fn(prompt) # type: ignore[misc]309 310 def _call_batched_query(311 prompts: List[str], model: str | None = None312 ) -> List[str]:313 if not self.llm_batch_fn:314 raise RuntimeError("llm_query_batched is not configured")315 try:316 return self.llm_batch_fn(prompts, model) # type: ignore[misc]317 except TypeError:318 return self.llm_batch_fn(prompts) # type: ignore[misc]319 320 def _call_recursive_query(prompt: str, model: str | None = None) -> str:321 if self.subcall_fn is None:322 return _call_single_query(prompt, model)323 return self.subcall_fn(prompt, model)324 325 def _call_recursive_batched(326 prompts: List[str], model: str | None = None327 ) -> List[str]:328 if not prompts:329 return []330 if self.subcall_batch_fn is not None:331 return self.subcall_batch_fn(prompts, model)332 return _call_batched_query(prompts, model)333 334 # Inject LLM functions if provided335 # Names: llm_query (single), llm_query_batched (official RLM), llm_batch (alias)336 if self.llm_query_fn:337 self._executor.inject_function("llm_query", _call_single_query)338 if self.llm_batch_fn:339 self._executor.inject_function(340 "llm_query_batched", _call_batched_query341 ) # Official name342 self._executor.inject_function("llm_batch", _call_batched_query) # Alias343 if self.llm_query_fn or self.subcall_fn:344 self._executor.inject_function("rlm_query", _call_recursive_query)345 if self.llm_batch_fn or self.subcall_batch_fn:346 self._executor.inject_function("rlm_query_batched", _call_recursive_batched)347 348 # Inject FINAL helper function so both FINAL(x) and print(f'FINAL({x})') work349 # Returns the FINAL pattern as a string so it appears in stdout for detection350 def final_helper(value):351 """Helper that returns FINAL(value) string for detection."""352 return f"FINAL({value})"353 354 self._executor.inject_function("FINAL", final_helper)355 356 # Inject FINAL_VAR helper that looks up variable and returns FINAL(value)357 # This matches official RLM behavior - strips quotes from var_name and looks up in namespace358 executor = self._executor # Capture for closure359 360 def final_var_helper(var_name: str):361 """Look up variable by name and return FINAL(value) for detection."""362 # Strip quotes if present (handles both FINAL_VAR("x") and FINAL_VAR(x))363 var_name_clean = str(var_name).strip().strip("\"'")364 # Look up variable in executor namespace365 value = executor.get_variable(var_name_clean)366 if value is not None:367 return f"FINAL({value})"368 return f"FINAL_VAR({var_name_clean})" # Fallback for regex detection369 370 self._executor.inject_function("FINAL_VAR", final_var_helper)371 372 def show_vars_helper():373 """Return the current non-private variables in the namespace."""374 return sorted(executor.list_variables())375 376 self._executor.inject_function("SHOW_VARS", show_vars_helper)377 378 # Update namespace keys379 self._state.namespace_keys = self._executor.list_variables()380 381 # Build initial message382 message_parts = ["REPL environment initialized."]383 if effective_context:384 message_parts.append(385 f"Context loaded ({len(effective_context)} chars). Use 'context' variable to access it."386 )387 if effective_task_prompt:388 message_parts.append(f"Task: {effective_task_prompt}")389 message_parts.append(390 "Use answer['content'] to store your answer, and set answer['ready'] = True when done."391 )392 393 return REPLObservation(394 result=CodeBlockResult(395 stdout="\n".join(message_parts),396 stderr="",397 locals_snapshot={},398 execution_time=0.0,399 success=True,400 exception=None,401 ),402 context_preview=(403 effective_context[: self.context_preview_length]404 if effective_context405 else None406 ),407 context_length=len(effective_context) if effective_context else 0,408 available_variables=self._state.namespace_keys,409 iteration=0,410 max_iterations=self.max_iterations,411 done=False,412 metadata={413 "task_prompt": effective_task_prompt,414 "message": "Environment ready.",415 },416 )417 418 def step(419 self,420 action: REPLAction,421 timeout_s: Optional[float] = None,422 **kwargs: Any,423 ) -> REPLObservation:424 """Execute code and return observation.425 426 Args:427 action: REPLAction containing code to execute428 timeout_s: Optional timeout in seconds (not currently used)429 **kwargs: Additional step parameters430 431 Returns:432 REPLObservation with execution results433 """434 if self._state is None or self._executor is None:435 raise RuntimeError("Environment not initialized. Call reset() first.")436 437 self._state.step_count += 1438 self._state.iteration += 1439 440 # Check if agent explicitly signals final answer441 if action.is_final:442 self._state.final_answer = action.final_answer or ""443 obs = self._create_final_observation(444 success=True,445 message="Final answer submitted.",446 )447 obs.reward = self._apply_rubric(action, obs)448 return obs449 450 # Check iteration limit451 if self._state.iteration >= self.max_iterations:452 # Check if there's a partial answer in the answer dict453 answer_var = self._executor.get_variable("answer")454 if isinstance(answer_var, dict) and answer_var.get("content"):455 self._state.final_answer = str(answer_var.get("content", ""))456 obs = self._create_final_observation(457 success=False,458 message=f"Maximum iterations ({self.max_iterations}) reached.",459 )460 obs.reward = self._apply_rubric(action, obs)461 return obs462 463 # Execute code464 result = self._executor.execute(action.code)465 self._state.total_execution_time += result["execution_time"]466 self._state.namespace_keys = self._executor.list_variables()467 468 # Check for final answer patterns469 final_answer = self._extract_final_answer(result["stdout"])470 done = final_answer is not None471 472 if done:473 self._state.final_answer = final_answer474 475 obs = REPLObservation(476 result=CodeBlockResult(477 stdout=result["stdout"],478 stderr=result["stderr"],479 locals_snapshot=result["locals_snapshot"],480 execution_time=result["execution_time"],481 success=result["success"],482 exception=result["exception"],483 ),484 context_preview=(485 self._state.context[: self.context_preview_length]486 if self._state.context487 else None488 ),489 context_length=len(self._state.context) if self._state.context else 0,490 available_variables=self._state.namespace_keys,491 iteration=self._state.iteration,492 max_iterations=self.max_iterations,493 done=done,494 metadata={495 "task_prompt": self._state.task_prompt,496 "final_answer": final_answer,497 "execution_time": result["execution_time"],498 },499 )500 obs.reward = self._apply_rubric(action, obs)501 return obs502 503 def _extract_final_answer(self, stdout: str) -> Optional[str]:504 """Extract final answer from output.505 506 Supports multiple patterns:507 1. RLM-style: FINAL(answer) in stdout508 2. RLM-style: FINAL_VAR(variable_name) in stdout509 3. Prime Intellect style: answer = {"content": "...", "ready": True} in namespace510 511 Args:512 stdout: Standard output from code execution513 514 Returns:515 Final answer string or None if not found516 """517 # Pattern 1: RLM-style FINAL(answer)518 final_match = re.search(r"FINAL\((.*?)\)", stdout, re.DOTALL)519 if final_match:520 return final_match.group(1).strip()521 522 # Pattern 2: RLM-style FINAL_VAR(variable_name)523 final_var_match = re.search(r"FINAL_VAR\((\w+)\)", stdout)524 if final_var_match and self._executor:525 var_name = final_var_match.group(1)526 value = self._executor.get_variable(var_name)527 if value is not None:528 return str(value)529 530 # Pattern 3: Prime Intellect style answer dict531 if self._executor:532 answer_var = self._executor.get_variable("answer")533 if isinstance(answer_var, dict):534 if answer_var.get("ready", False):535 return str(answer_var.get("content", ""))536 537 return None538 539 def _create_final_observation(self, success: bool, message: str) -> REPLObservation:540 """Create observation for episode termination.541 542 Args:543 success: Whether the episode ended successfully544 message: Termination message545 546 Returns:547 Final REPLObservation with done=True (reward set by rubric)548 """549 return REPLObservation(550 result=CodeBlockResult(551 stdout=message,552 stderr="",553 locals_snapshot={},554 execution_time=0.0,555 success=success,556 exception=None,557 ),558 context_preview=None,559 context_length=0,560 available_variables=[],561 iteration=self._state.iteration if self._state else 0,562 max_iterations=self.max_iterations,563 done=True,564 metadata={565 "final_answer": self._state.final_answer if self._state else None,566 "total_execution_time": (567 self._state.total_execution_time if self._state else 0568 ),569 "total_iterations": self._state.iteration if self._state else 0,570 },571 )572 573 @property574 def state(self) -> REPLState:575 """Get the current environment state.576 577 Returns:578 Current REPLState579 580 Raises:581 RuntimeError: If environment not initialized582 """583 if self._state is None:584 raise RuntimeError("Environment not initialized. Call reset() first.")585 return self._state586 587 def close(self) -> None:588 """Cleanup resources."""589 if self._runtime_controller is not None:590 self._runtime_controller.close()591 self._runtime_controller = None592 self._executor = None593 self._state = None594 595 def get_metadata(self) -> EnvironmentMetadata:596 """Get environment metadata.597 598 Returns:599 EnvironmentMetadata with environment info600 """601 return EnvironmentMetadata(602 name="repl_env",603 description="Python REPL environment for RLM-style code execution",604 version="0.1.0",605 )606 