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"""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 