Team Ai
Apppublic

openenv/repl

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
repl_environment.py606 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"""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