Team Ai
Apppublic

openenv/repl

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
recursive_backends.py252 linesDownload Raw Back to root
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"""8Recursive backend abstractions for repl_env.9 10This module keeps direct LM calls and recursive child spawning out of the11runner and environment. The runner owns the iterative loop; the backend owns12query/query_batched/child-recursion behavior.13"""14 15from __future__ import annotations16 17import threading18import time19from concurrent.futures import as_completed, ThreadPoolExecutor20from dataclasses import dataclass, field21from typing import Callable, Protocol22 23 24ChatFn = Callable[..., str]25 26 27class RecursiveBackend(Protocol):28    max_depth: int29    depth: int30    child_traces: list["ChildTrace"]31 32    def query(self, prompt: str, model: str | None = None) -> str: ...33 34    def query_batched(35        self, prompts: list[str], model: str | None = None36    ) -> list[str]: ...37 38    def recursive_query(self, prompt: str, model: str | None = None) -> str: ...39 40    def recursive_query_batched(41        self, prompts: list[str], model: str | None = None42    ) -> list[str]: ...43 44 45@dataclass46class BackendLimits:47    max_depth: int = 148    max_batch_workers: int = 849    max_children_total: int | None = None50    max_children_per_batch: int | None = None51    result_truncation_limit: int | None = None52    # Cooperative timeout: checked between iterations, not during LLM calls.53    # A slow LLM call within an iteration will not be interrupted — the timeout54    # fires at the next iteration boundary. For mid-call cancellation, use55    # process-based isolation instead.56    per_child_timeout_s: float | None = None57    # Tree-global child counter shared across all recursion depths58    _children_spawned: int = field(default=0, init=False, repr=False)59    _children_lock: threading.Lock = field(60        default_factory=threading.Lock, init=False, repr=False61    )62 63 64@dataclass65class ChildTrace:66    depth: int67    duration_s: float68    prompt_preview: str69    result_preview: str | None70    error: str | None71 72 73class DirectLMBackend:74    """Direct LM backend with no child recursion beyond fallback to itself."""75 76    def __init__(77        self,78        llm_chat_fn: ChatFn,79        *,80        depth: int = 0,81        limits: BackendLimits | None = None,82    ) -> None:83        self.llm_chat_fn = llm_chat_fn84        self.depth = depth85        self.limits = limits or BackendLimits()86        self.max_depth = self.limits.max_depth87        self.child_traces: list[ChildTrace] = []88 89    def query(self, prompt: str, model: str | None = None) -> str:90        try:91            result = self.llm_chat_fn([{"role": "user", "content": prompt}], model)92        except TypeError:93            result = self.llm_chat_fn([{"role": "user", "content": prompt}])94        return self._truncate(result)95 96    def query_batched(self, prompts: list[str], model: str | None = None) -> list[str]:97        if not prompts:98            return []99        max_workers = min(len(prompts), self.limits.max_batch_workers)100        results: list[str] = [""] * len(prompts)101        with ThreadPoolExecutor(max_workers=max_workers) as executor:102            future_to_idx = {103                executor.submit(self.query, prompt, model): idx104                for idx, prompt in enumerate(prompts)105            }106            for future in as_completed(future_to_idx):107                idx = future_to_idx[future]108                try:109                    results[idx] = future.result()110                except Exception as exc:111                    results[idx] = f"Error: {exc}"112        return results113 114    def recursive_query(self, prompt: str, model: str | None = None) -> str:115        return self.query(prompt, model)116 117    def recursive_query_batched(118        self, prompts: list[str], model: str | None = None119    ) -> list[str]:120        return self.query_batched(prompts, model)121 122    def _truncate(self, result: str) -> str:123        limit = self.limits.result_truncation_limit124        if limit is not None and len(result) > limit:125            return result[:limit]126        return result127 128 129class LocalChildRLMBackend(DirectLMBackend):130    """Recursive backend that spawns child LocalRLMRunner instances."""131 132    def __init__(133        self,134        llm_chat_fn: ChatFn,135        *,136        runner_factory: Callable[..., object],137        system_prompt: str,138        max_iterations: int,139        env_max_iterations_multiplier: int,140        depth: int = 0,141        limits: BackendLimits | None = None,142        on_subcall_start: Callable[[int, str, str], None] | None = None,143        on_subcall_complete: Callable[[int, str, float, str | None], None]144        | None = None,145    ) -> None:146        super().__init__(llm_chat_fn, depth=depth, limits=limits)147        self.runner_factory = runner_factory148        self.system_prompt = system_prompt149        self.max_iterations = max_iterations150        self.env_max_iterations_multiplier = env_max_iterations_multiplier151        self.on_subcall_start = on_subcall_start152        self.on_subcall_complete = on_subcall_complete153 154    def recursive_query(self, prompt: str, model: str | None = None) -> str:155        next_depth = self.depth + 1156        if next_depth >= self.max_depth:157            return self.query(prompt, model)158        with self.limits._children_lock:159            if self.limits.max_children_total is not None:160                if self.limits._children_spawned >= self.limits.max_children_total:161                    return "Error: max_children_total exceeded"162            self.limits._children_spawned += 1163        start = time.perf_counter()164        error: str | None = None165        result_text = ""166        resolved_model = model or "default"167        if self.on_subcall_start is not None:168            try:169                self.on_subcall_start(next_depth, str(resolved_model), prompt[:80])170            except Exception:171                pass172        try:173            child = self.runner_factory(174                self.llm_chat_fn,175                system_prompt=self.system_prompt,176                max_iterations=self.max_iterations,177                max_depth=self.max_depth,178                depth=next_depth,179                env_max_iterations_multiplier=self.env_max_iterations_multiplier,180                max_batch_workers=self.limits.max_batch_workers,181                backend_factory=self._child_backend_factory,182                on_subcall_start=self.on_subcall_start,183                on_subcall_complete=self.on_subcall_complete,184            )185            result = child.run(186                prompt, prompt, model=model, timeout_s=self.limits.per_child_timeout_s187            )188            result_text = self._truncate(result.final_answer or "")189            return result_text190        except Exception as exc:191            error = str(exc)192            raise193        finally:194            duration = time.perf_counter() - start195            self.child_traces.append(196                ChildTrace(197                    depth=next_depth,198                    duration_s=duration,199                    prompt_preview=prompt[:80],200                    result_preview=(result_text[:80] if result_text else None),201                    error=error,202                )203            )204            if self.on_subcall_complete is not None:205                try:206                    self.on_subcall_complete(207                        next_depth,208                        str(resolved_model),209                        duration,210                        error,211                    )212                except Exception:213                    pass214 215    def recursive_query_batched(216        self, prompts: list[str], model: str | None = None217    ) -> list[str]:218        if not prompts:219            return []220        batch_limit = self.limits.max_children_per_batch221        if batch_limit is not None:222            prompts = prompts[:batch_limit]223        max_workers = min(len(prompts), self.limits.max_batch_workers)224        results: list[str] = [""] * len(prompts)225        with ThreadPoolExecutor(max_workers=max_workers) as executor:226            future_to_idx = {227                executor.submit(self.recursive_query, prompt, model): idx228                for idx, prompt in enumerate(prompts)229            }230            for future in as_completed(future_to_idx):231                idx = future_to_idx[future]232                try:233                    results[idx] = future.result()234                except Exception as exc:235                    results[idx] = f"Error: {exc}"236        return results237 238    def _child_backend_factory(239        self, llm_chat_fn: ChatFn, **kwargs240    ) -> "LocalChildRLMBackend":241        return LocalChildRLMBackend(242            llm_chat_fn,243            runner_factory=self.runner_factory,244            system_prompt=kwargs["system_prompt"],245            max_iterations=kwargs["max_iterations"],246            env_max_iterations_multiplier=kwargs["env_max_iterations_multiplier"],247            depth=kwargs["depth"],248            limits=self.limits,249            on_subcall_start=self.on_subcall_start,250            on_subcall_complete=self.on_subcall_complete,251        )252