| 1 | //! Long-lived Python REPL runtime. |
| 2 | //! |
| 3 | //! One Python subprocess lives for the duration of an RLM turn (or an |
| 4 | //! inline `repl` block sequence in the agent loop). Code blocks are sent |
| 5 | //! over stdin framed by `__RLM_RUN__`/`__RLM_END__` sentinels; the bootstrap |
| 6 | //! `exec()`s them into the same global namespace so variables, imports, |
| 7 | //! and even open file handles persist naturally across rounds. |
| 8 | //! |
| 9 | //! Sub-LLM helpers (`sub_query`, `sub_query_batch`, `sub_rlm`, plus legacy |
| 10 | //! `llm_query`, `llm_query_batched`, `rlm_query`, `rlm_query_batched`) are |
| 11 | //! wired through a stdin/stdout RPC protocol: |
| 12 | //! Python emits `__RLM_REQ_<sid>__::{json}` on stdout, Rust dispatches the |
| 13 | //! request and writes `__RLM_RESP_<sid>__::{json}` back on stdin. No HTTP |
| 14 | //! sidecar, no temp ports — the same pipes carry both control and data. |
| 15 | //! |
| 16 | //! The session id (`<sid>`) is a UUID generated per spawn, so user output |
| 17 | //! that happens to contain "REQ" or "FINAL" can't be confused with control |
| 18 | //! messages. |
| 19 | |
| 20 | use std::ffi::OsString; |
| 21 | use std::path::{Path, PathBuf}; |
| 22 | use std::process::Stdio; |
| 23 | use std::time::{Duration, Instant}; |
| 24 | |
| 25 | use serde::{Deserialize, Serialize}; |
| 26 | use serde_json::Value; |
| 27 | use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; |
| 28 | use tokio::process::{Child, ChildStdin, ChildStdout}; |
| 29 | use uuid::Uuid; |
| 30 | |
| 31 | use crate::child_env; |
| 32 | use crate::dependencies::ExternalTool; |
| 33 | |
| 34 | // --------------------------------------------------------------------------- |
| 35 | // Public types |
| 36 | // --------------------------------------------------------------------------- |
| 37 | |
| 38 | /// Result of executing one code block. |
| 39 | #[derive(Debug, Clone)] |
| 40 | pub struct ReplRound { |
| 41 | /// Stdout shown to the model as metadata next round. |
| 42 | pub stdout: String, |
| 43 | /// Full stdout (with sentinels stripped, but otherwise raw). |
| 44 | pub full_stdout: String, |
| 45 | /// Stderr from this round (if any). |
| 46 | pub stderr: String, |
| 47 | /// `True` if the user code raised an unhandled Python exception. |
| 48 | pub has_error: bool, |
| 49 | /// Captured `finalize(value, confidence=...)` payload, if any. |
| 50 | pub final_value: Option<String>, |
| 51 | /// Captured final value before string fallback. Structured `finalize` |
| 52 | /// payloads use this so `handle_read` can expose JSON instead of a Python |
| 53 | /// repr string. |
| 54 | pub final_json: Option<Value>, |
| 55 | /// Optional confidence supplied to `finalize(...)`. |
| 56 | pub final_confidence: Option<Value>, |
| 57 | /// Number of `sub_query`/`sub_rlm` RPCs the round issued. |
| 58 | pub rpc_count: u32, |
| 59 | /// Wall-clock duration of the round. |
| 60 | pub elapsed: Duration, |
| 61 | } |
| 62 | |
| 63 | /// One RPC request emitted by Python during a round. |
| 64 | #[derive(Debug, Clone, Serialize, Deserialize)] |
| 65 | #[serde(tag = "type", rename_all = "snake_case")] |
| 66 | pub enum RpcRequest { |
| 67 | /// `llm_query(prompt, model=None, max_tokens=None, system=None)` |
| 68 | Llm { |
| 69 | prompt: String, |
| 70 | #[serde(default)] |
| 71 | model: Option<String>, |
| 72 | #[serde(default)] |
| 73 | max_tokens: Option<u32>, |
| 74 | #[serde(default)] |
| 75 | system: Option<String>, |
| 76 | }, |
| 77 | /// `llm_query_batched(prompts, model=None, dependency_mode="independent")` |
| 78 | LlmBatch { |
| 79 | prompts: Vec<String>, |
| 80 | #[serde(default)] |
| 81 | model: Option<String>, |
| 82 | #[serde(default)] |
| 83 | dependency_mode: Option<String>, |
| 84 | #[serde(default)] |
| 85 | safety_note: Option<String>, |
| 86 | }, |
| 87 | /// `rlm_query(prompt, model=None)` — recursive sub-RLM (paper's `sub_RLM`). |
| 88 | Rlm { |
| 89 | prompt: String, |
| 90 | #[serde(default)] |
| 91 | model: Option<String>, |
| 92 | }, |
| 93 | /// `rlm_query_batched(prompts, model=None, dependency_mode="independent")` |
| 94 | RlmBatch { |
| 95 | prompts: Vec<String>, |
| 96 | #[serde(default)] |
| 97 | model: Option<String>, |
| 98 | #[serde(default)] |
| 99 | dependency_mode: Option<String>, |
| 100 | #[serde(default)] |
| 101 | safety_note: Option<String>, |
| 102 | }, |
| 103 | } |
| 104 | |
| 105 | /// Response for one RPC request. |
| 106 | #[derive(Debug, Clone, Serialize, Deserialize)] |
| 107 | #[serde(untagged)] |
| 108 | pub enum RpcResponse { |
| 109 | /// Single-text reply (Llm / Rlm). |
| 110 | Single(SingleResp), |
| 111 | /// Batch reply (LlmBatch / RlmBatch). |
| 112 | Batch(BatchResp), |
| 113 | } |
| 114 | |
| 115 | #[derive(Debug, Clone, Serialize, Deserialize)] |
| 116 | pub struct SingleResp { |
| 117 | #[serde(default)] |
| 118 | pub text: String, |
| 119 | #[serde(default, skip_serializing_if = "Option::is_none")] |
| 120 | pub error: Option<String>, |
| 121 | } |
| 122 | |
| 123 | #[derive(Debug, Clone, Serialize, Deserialize)] |
| 124 | pub struct BatchResp { |
| 125 | pub results: Vec<SingleResp>, |
| 126 | } |
| 127 | |
| 128 | /// Trait-object handle for dispatching Python RPCs back into Rust. |
| 129 | /// |
| 130 | /// Each RLM turn supplies one. Implementations forward to the LLM client |
| 131 | /// (and recursively through the same captured Engine for `Rlm` / `RlmBatch`). |
| 132 | pub trait RpcDispatcher: Send + Sync { |
| 133 | fn dispatch<'a>( |
| 134 | &'a self, |
| 135 | req: RpcRequest, |
| 136 | ) -> std::pin::Pin<Box<dyn std::future::Future<Output = RpcResponse> + Send + 'a>>; |
| 137 | } |
| 138 | |
| 139 | // --------------------------------------------------------------------------- |
| 140 | // Constants |
| 141 | // --------------------------------------------------------------------------- |
| 142 | |
| 143 | const DEFAULT_STDOUT_LIMIT: usize = 8_192; |
| 144 | /// Longest single stdout line (protocol or plain output) held in memory. An |
| 145 | /// RPC request is one line, so this also bounds what a request can make Rust |
| 146 | /// deserialize before the bridge's own caps apply; an over-long request is |
| 147 | /// truncated, fails to parse, and is answered with an error. |
| 148 | const MAX_STDOUT_LINE_BYTES: usize = 32 * 1024 * 1024; |
| 149 | /// Most stdout one round retains (`full_stdout`); the rest is counted and |
| 150 | /// reported, never buffered. |
| 151 | const MAX_ROUND_STDOUT_BYTES: usize = 32 * 1024 * 1024; |
| 152 | const ROUND_TIMEOUT: Duration = Duration::from_secs(180); |
| 153 | #[cfg(not(windows))] |
| 154 | const SPAWN_READY_TIMEOUT: Duration = Duration::from_secs(10); |
| 155 | #[cfg(windows)] |
| 156 | const SPAWN_READY_TIMEOUT: Duration = Duration::from_secs(30); |
| 157 | |
| 158 | // --------------------------------------------------------------------------- |
| 159 | // PythonRuntime |
| 160 | // --------------------------------------------------------------------------- |
| 161 | |
| 162 | /// Long-lived Python REPL. |
| 163 | #[derive(Debug)] |
| 164 | pub struct PythonRuntime { |
| 165 | child: Child, |
| 166 | stdin: ChildStdin, |
| 167 | stdout: BufReader<ChildStdout>, |
| 168 | /// Per-spawn session id used in protocol sentinels. |
| 169 | session_id: String, |
| 170 | /// Path to the file holding `context` (kept around for cleanup). |
| 171 | context_path: Option<PathBuf>, |
| 172 | stdout_limit: usize, |
| 173 | round_count: u64, |
| 174 | started: Instant, |
| 175 | round_timeout: Option<Duration>, |
| 176 | /// Set while a round has not yet read its own DONE sentinel, and kept |
| 177 | /// when a round ends without it (timeout, I/O failure, or a caller that |
| 178 | /// dropped the round mid-flight). The protocol stream is then out of step |
| 179 | /// — the interrupted round may still print — so a failed round kills the |
| 180 | /// interpreter and every later round refuses instead of reading another |
| 181 | /// round's output as its own. |
| 182 | broken: Option<String>, |
| 183 | } |
| 184 | |
| 185 | impl PythonRuntime { |
| 186 | /// Spawn a REPL with no `context` variable and no LLM helpers wired up. |
| 187 | /// Used by the agent loop for inline `repl` blocks the model emits in |
| 188 | /// regular conversation. |
| 189 | pub async fn new() -> Result<Self, String> { |
| 190 | Self::spawn_inner(None, Some(ROUND_TIMEOUT)).await |
| 191 | } |
| 192 | |
| 193 | /// Compatibility shim — older RLM code path used to pass a state file. |
| 194 | /// The state file is no longer used, but the path doubles as an extra |
| 195 | /// scratch location callers can rely on for cleanup symmetry. |
| 196 | pub fn with_state_path(_path: PathBuf) -> Self { |
| 197 | // Synchronous constructor is no longer meaningful: spawning Python |
| 198 | // is async. Callers in turn.rs already use `spawn_with_context` — |
| 199 | // this stub is kept only so the public surface compiles for any |
| 200 | // out-of-tree user. It returns a deliberately broken runtime that |
| 201 | // panics on first use, which is preferable to silently lying. |
| 202 | unreachable!( |
| 203 | "PythonRuntime::with_state_path is deprecated — \ |
| 204 | use PythonRuntime::new() or PythonRuntime::spawn_with_context()" |
| 205 | ) |
| 206 | } |
| 207 | |
| 208 | /// Spawn a REPL with the long input preloaded from a file. Used by the |
| 209 | /// RLM turn loop. |
| 210 | pub async fn spawn_with_context(context_path: &Path) -> Result<Self, String> { |
| 211 | Self::spawn_inner(Some(context_path), None).await |
| 212 | } |
| 213 | |
| 214 | async fn spawn_inner( |
| 215 | context_path: Option<&Path>, |
| 216 | round_timeout: Option<Duration>, |
| 217 | ) -> Result<Self, String> { |
| 218 | let session_id = Uuid::new_v4().simple().to_string(); |
| 219 | let bootstrap = render_bootstrap(&session_id); |
| 220 | |
| 221 | let mut cmd = crate::dependencies::Python::tokio_command().ok_or_else(|| { |
| 222 | "no Python interpreter found on PATH (tried python3, python, py -3). \ |
| 223 | Install Python 3 and restart codewhale." |
| 224 | .to_string() |
| 225 | })?; |
| 226 | cmd.arg("-u") |
| 227 | .arg("-c") |
| 228 | .arg(&bootstrap) |
| 229 | .stdin(Stdio::piped()) |
| 230 | .stdout(Stdio::piped()) |
| 231 | .stderr(Stdio::piped()) |
| 232 | .kill_on_drop(true); |
| 233 | |
| 234 | let context_env = context_path |
| 235 | .map(|path| { |
| 236 | vec![( |
| 237 | OsString::from("RLM_CONTEXT_FILE"), |
| 238 | path.as_os_str().to_os_string(), |
| 239 | )] |
| 240 | }) |
| 241 | .unwrap_or_default(); |
| 242 | child_env::apply_to_tokio_command(&mut cmd, context_env); |
| 243 | |
| 244 | let mut child = cmd |
| 245 | .spawn() |
| 246 | .map_err(|e| format!("failed to spawn Python interpreter: {e}"))?; |
| 247 | |
| 248 | let stdin = child |
| 249 | .stdin |
| 250 | .take() |
| 251 | .ok_or_else(|| "Python interpreter stdin pipe missing".to_string())?; |
| 252 | let raw_stdout = child |
| 253 | .stdout |
| 254 | .take() |
| 255 | .ok_or_else(|| "Python interpreter stdout pipe missing".to_string())?; |
| 256 | let stdout = BufReader::new(raw_stdout); |
| 257 | |
| 258 | let mut rt = Self { |
| 259 | child, |
| 260 | stdin, |
| 261 | stdout, |
| 262 | session_id: session_id.clone(), |
| 263 | context_path: context_path.map(Path::to_path_buf), |
| 264 | stdout_limit: DEFAULT_STDOUT_LIMIT, |
| 265 | round_count: 0, |
| 266 | started: Instant::now(), |
| 267 | round_timeout, |
| 268 | broken: None, |
| 269 | }; |
| 270 | |
| 271 | // Wait for `__RLM_READY_<sid>__` before handing control back. If |
| 272 | // Python failed to start (missing module, syntax error in the |
| 273 | // bootstrap, etc.), this is where we'll find out. |
| 274 | let ready_sentinel = format!("__RLM_READY_{session_id}__"); |
| 275 | match tokio::time::timeout(SPAWN_READY_TIMEOUT, rt.read_until_ready(&ready_sentinel)).await |
| 276 | { |
| 277 | Ok(Ok(())) => Ok(rt), |
| 278 | Ok(Err(e)) => { |
| 279 | let _ = rt.child.kill().await; |
| 280 | Err(format!("Python interpreter bootstrap failed: {e}")) |
| 281 | } |
| 282 | Err(_) => { |
| 283 | let _ = rt.child.kill().await; |
| 284 | Err(format!( |
| 285 | "Python interpreter bootstrap did not signal ready within {}s", |
| 286 | SPAWN_READY_TIMEOUT.as_secs() |
| 287 | )) |
| 288 | } |
| 289 | } |
| 290 | } |
| 291 | |
| 292 | async fn read_until_ready(&mut self, ready_sentinel: &str) -> Result<(), String> { |
| 293 | loop { |
| 294 | let line = match self.read_stdout_line_lossy().await? { |
| 295 | Some(line) => line, |
| 296 | None => { |
| 297 | return Err("Python interpreter closed stdout before ready signal".to_string()); |
| 298 | } |
| 299 | }; |
| 300 | let trimmed = line.trim_end_matches(['\n', '\r']); |
| 301 | if trimmed == ready_sentinel { |
| 302 | return Ok(()); |
| 303 | } |
| 304 | // Pre-ready output is rare; ignore it. |
| 305 | } |
| 306 | } |
| 307 | |
| 308 | /// One stdout line, holding at most [`MAX_STDOUT_LINE_BYTES`] of it: the |
| 309 | /// rest of an over-long line is consumed and replaced by an omission note. |
| 310 | async fn read_stdout_line_lossy(&mut self) -> Result<Option<String>, String> { |
| 311 | let mut buf = Vec::new(); |
| 312 | let mut omitted = 0usize; |
| 313 | loop { |
| 314 | let available = self |
| 315 | .stdout |
| 316 | .fill_buf() |
| 317 | .await |
| 318 | .map_err(|e| format!("stdout read: {e}"))?; |
| 319 | if available.is_empty() { |
| 320 | break; |
| 321 | } |
| 322 | let (take, ends_line) = match available.iter().position(|&byte| byte == b'\n') { |
| 323 | Some(index) => (index + 1, true), |
| 324 | None => (available.len(), false), |
| 325 | }; |
| 326 | let keep = take.min(MAX_STDOUT_LINE_BYTES.saturating_sub(buf.len())); |
| 327 | buf.extend_from_slice(&available[..keep]); |
| 328 | omitted += take - keep; |
| 329 | self.stdout.consume(take); |
| 330 | if ends_line { |
| 331 | break; |
| 332 | } |
| 333 | } |
| 334 | if buf.is_empty() && omitted == 0 { |
| 335 | return Ok(None); |
| 336 | } |
| 337 | let mut line = String::from_utf8_lossy(&buf).into_owned(); |
| 338 | if omitted > 0 { |
| 339 | line.truncate(line.trim_end_matches(['\n', '\r']).len()); |
| 340 | line.push_str(&format!( |
| 341 | "[... {omitted} bytes of this line omitted by the REPL line limit ...]\n" |
| 342 | )); |
| 343 | } |
| 344 | Ok(Some(line)) |
| 345 | } |
| 346 | |
| 347 | /// Execute a Python code block with no RPC dispatcher. Used for inline |
| 348 | /// `repl` blocks where `llm_query()` should fall back to a sentinel. |
| 349 | pub async fn execute(&mut self, code: &str) -> Result<ReplRound, String> { |
| 350 | self.run(code, None::<&dyn RpcDispatcher>).await |
| 351 | } |
| 352 | |
| 353 | /// Replace the long context visible to bounded helpers without restarting |
| 354 | /// the Python process. User-created variables, imports, and handles stay |
| 355 | /// alive across the refresh, which is what lets the normal agent loop use |
| 356 | /// one working kernel instead of rebuilding a throwaway REPL each turn. |
| 357 | /// |
| 358 | /// The payload travels through a temporary file rather than a generated |
| 359 | /// Python string so large transcripts neither bloat the command stream nor |
| 360 | /// acquire quoting semantics. The old owned file is released only after |
| 361 | /// the kernel has successfully loaded the replacement. |
| 362 | pub async fn replace_context(&mut self, body: &str) -> Result<(), String> { |
| 363 | let path = crate::rlm::session::write_context_file(body) |
| 364 | .map_err(|e| format!("write refreshed REPL context: {e}"))?; |
| 365 | let path_literal = serde_json::to_string(&path.to_string_lossy()) |
| 366 | .map_err(|e| format!("encode refreshed REPL context path: {e}"))?; |
| 367 | let code = format!("_replace_context_file({path_literal})"); |
| 368 | |
| 369 | match self.execute(&code).await { |
| 370 | Ok(round) if !round.has_error => { |
| 371 | if let Some(previous) = self.context_path.replace(path) { |
| 372 | let _ = tokio::fs::remove_file(previous).await; |
| 373 | } |
| 374 | Ok(()) |
| 375 | } |
| 376 | Ok(round) => { |
| 377 | let _ = tokio::fs::remove_file(&path).await; |
| 378 | Err(format!( |
| 379 | "refresh REPL context failed: {}{}", |
| 380 | round.stdout, |
| 381 | if round.stderr.is_empty() { |
| 382 | String::new() |
| 383 | } else { |
| 384 | format!("\nstderr: {}", round.stderr) |
| 385 | } |
| 386 | )) |
| 387 | } |
| 388 | Err(error) => { |
| 389 | let _ = tokio::fs::remove_file(&path).await; |
| 390 | Err(error) |
| 391 | } |
| 392 | } |
| 393 | } |
| 394 | |
| 395 | /// Execute a code block, dispatching any sub-LLM RPCs through `bridge`. |
| 396 | /// |
| 397 | /// Returns once Python emits `__RLM_DONE_<sid>__` or the round timeout |
| 398 | /// elapses (whichever happens first). |
| 399 | /// |
| 400 | /// Known limitations (F02-07): the sentinels share stdout with the code |
| 401 | /// they frame, and `_SID`/`_DONE` are ordinary Python globals, so code in |
| 402 | /// this round (or a subprocess inheriting fd 1) can still print its own |
| 403 | /// round's DONE line and end the round early; only another round's DONE |
| 404 | /// is refused. A failed or timed-out round kills the interpreter process |
| 405 | /// itself, not its process group, so a grandchild it spawned may outlive |
| 406 | /// the kernel. Closing either needs an out-of-band protocol channel. |
| 407 | pub async fn run<D>(&mut self, code: &str, bridge: Option<&D>) -> Result<ReplRound, String> |
| 408 | where |
| 409 | D: RpcDispatcher + ?Sized, |
| 410 | { |
| 411 | if let Some(reason) = &self.broken { |
| 412 | return Err(format!( |
| 413 | "Python REPL is unusable after an earlier round failed ({reason}); start a new REPL" |
| 414 | )); |
| 415 | } |
| 416 | // Pessimistic until this round reads its own DONE: a caller that |
| 417 | // drops this future mid-round leaves the stream out of step too. |
| 418 | self.broken = Some("an earlier round was cancelled before it finished".to_string()); |
| 419 | let started = Instant::now(); |
| 420 | self.round_count += 1; |
| 421 | let round_id = self.round_count; |
| 422 | let round_tag = round_id.to_string(); |
| 423 | |
| 424 | // Send the code header + body + end marker in one write. |
| 425 | let header = format!("__RLM_RUN_{}__::{round_id}\n", self.session_id); |
| 426 | let footer = format!("__RLM_END_{}__\n", self.session_id); |
| 427 | let payload = format!("{header}{code}\n{footer}"); |
| 428 | self.stdin |
| 429 | .write_all(payload.as_bytes()) |
| 430 | .await |
| 431 | .map_err(|e| format!("stdin write: {e}"))?; |
| 432 | self.stdin |
| 433 | .flush() |
| 434 | .await |
| 435 | .map_err(|e| format!("stdin flush: {e}"))?; |
| 436 | |
| 437 | // Sentinels for this session. |
| 438 | let req_prefix = format!("__RLM_REQ_{}__::", self.session_id); |
| 439 | let final_prefix = format!("__RLM_FINAL_{}__::", self.session_id); |
| 440 | let err_prefix = format!("__RLM_ERR_{}__::", self.session_id); |
| 441 | let done_prefix = format!("__RLM_DONE_{}__::", self.session_id); |
| 442 | |
| 443 | let mut stdout_buf = String::new(); |
| 444 | let mut stdout_omitted = 0usize; |
| 445 | let mut final_value: Option<String> = None; |
| 446 | let mut final_json: Option<Value> = None; |
| 447 | let mut final_confidence: Option<Value> = None; |
| 448 | let mut had_error = false; |
| 449 | let mut rpc_count: u32 = 0; |
| 450 | let round_timeout = self.round_timeout; |
| 451 | |
| 452 | let read_loop = async { |
| 453 | loop { |
| 454 | let line = match self.read_stdout_line_lossy().await? { |
| 455 | Some(line) => line, |
| 456 | None => { |
| 457 | return Err("Python interpreter closed stdout mid-round".to_string()); |
| 458 | } |
| 459 | }; |
| 460 | let trimmed = line.trim_end_matches(['\n', '\r']); |
| 461 | |
| 462 | if let Some(rest) = trimmed.strip_prefix(&done_prefix) { |
| 463 | if rest == round_tag { |
| 464 | break; |
| 465 | } |
| 466 | // Another round's completion (or user output imitating |
| 467 | // one) never ends this round. |
| 468 | push_bounded(&mut stdout_buf, &line, &mut stdout_omitted); |
| 469 | continue; |
| 470 | } |
| 471 | if let Some(rest) = trimmed.strip_prefix(&final_prefix) { |
| 472 | // New sessions emit an object with value/confidence; |
| 473 | // legacy helpers emitted a JSON string. |
| 474 | match serde_json::from_str::<Value>(rest) { |
| 475 | Ok(Value::Object(map)) => { |
| 476 | let value_json = map |
| 477 | .get("value") |
| 478 | .cloned() |
| 479 | .unwrap_or(Value::String(rest.to_string())); |
| 480 | let value = value_json |
| 481 | .as_str() |
| 482 | .map(str::to_string) |
| 483 | .unwrap_or_else(|| value_json.to_string()); |
| 484 | final_json = Some(value_json); |
| 485 | final_value = Some(value); |
| 486 | final_confidence = map.get("confidence").cloned(); |
| 487 | } |
| 488 | Ok(Value::String(value)) => { |
| 489 | final_json = Some(Value::String(value.clone())); |
| 490 | final_value = Some(value); |
| 491 | final_confidence = None; |
| 492 | } |
| 493 | Ok(other) => { |
| 494 | final_json = Some(other.clone()); |
| 495 | final_value = Some(other.to_string()); |
| 496 | final_confidence = None; |
| 497 | } |
| 498 | Err(_) => { |
| 499 | final_value = Some(rest.to_string()); |
| 500 | final_confidence = None; |
| 501 | } |
| 502 | } |
| 503 | continue; |
| 504 | } |
| 505 | if let Some(rest) = trimmed.strip_prefix(&err_prefix) { |
| 506 | let traceback = |
| 507 | serde_json::from_str::<String>(rest).unwrap_or_else(|_| rest.to_string()); |
| 508 | had_error = true; |
| 509 | push_bounded( |
| 510 | &mut stdout_buf, |
| 511 | &format!("[traceback]\n{traceback}\n"), |
| 512 | &mut stdout_omitted, |
| 513 | ); |
| 514 | continue; |
| 515 | } |
| 516 | if let Some(rest) = trimmed.strip_prefix(&req_prefix) { |
| 517 | rpc_count = rpc_count.saturating_add(1); |
| 518 | let req: RpcRequest = match serde_json::from_str(rest) { |
| 519 | Ok(r) => r, |
| 520 | Err(e) => { |
| 521 | // Send an error response so Python isn't blocked. |
| 522 | self.send_resp(&RpcResponse::Single(SingleResp { |
| 523 | text: String::new(), |
| 524 | error: Some(format!("malformed RPC: {e}")), |
| 525 | })) |
| 526 | .await?; |
| 527 | continue; |
| 528 | } |
| 529 | }; |
| 530 | let resp = match bridge { |
| 531 | Some(b) => b.dispatch(req).await, |
| 532 | None => RpcResponse::Single(SingleResp { |
| 533 | text: String::new(), |
| 534 | error: Some("no LLM bridge bound to this REPL".to_string()), |
| 535 | }), |
| 536 | }; |
| 537 | self.send_resp(&resp).await?; |
| 538 | continue; |
| 539 | } |
| 540 | |
| 541 | push_bounded(&mut stdout_buf, &line, &mut stdout_omitted); |
| 542 | } |
| 543 | Ok::<_, String>(()) |
| 544 | }; |
| 545 | |
| 546 | let outcome = match round_timeout { |
| 547 | Some(round_timeout) => tokio::time::timeout(round_timeout, read_loop) |
| 548 | .await |
| 549 | .unwrap_or_else(|_| { |
| 550 | Err(format!( |
| 551 | "REPL round timed out after {}s", |
| 552 | round_timeout.as_secs() |
| 553 | )) |
| 554 | }), |
| 555 | None => read_loop.await, |
| 556 | }; |
| 557 | if let Err(error) = outcome { |
| 558 | // The interpreter may still be executing this round. Kill it now |
| 559 | // (the handle is reaped on drop) rather than leave it running and |
| 560 | // let its late output land in a later round. |
| 561 | let _ = self.child.start_kill(); |
| 562 | self.broken = Some(error.clone()); |
| 563 | return Err(error); |
| 564 | } |
| 565 | self.broken = None; |
| 566 | if stdout_omitted > 0 { |
| 567 | stdout_buf.push_str(&format!( |
| 568 | "\n[... REPL round output capped: {stdout_omitted} bytes omitted ...]\n" |
| 569 | )); |
| 570 | } |
| 571 | |
| 572 | let stderr = self.drain_stderr().await; |
| 573 | let display = truncate_stdout(stdout_buf.trim_end_matches('\n'), self.stdout_limit); |
| 574 | |
| 575 | Ok(ReplRound { |
| 576 | stdout: display, |
| 577 | full_stdout: stdout_buf, |
| 578 | stderr, |
| 579 | has_error: had_error, |
| 580 | final_value, |
| 581 | final_json, |
| 582 | final_confidence, |
| 583 | rpc_count, |
| 584 | elapsed: started.elapsed(), |
| 585 | }) |
| 586 | } |
| 587 | |
| 588 | async fn send_resp(&mut self, resp: &RpcResponse) -> Result<(), String> { |
| 589 | let body = serde_json::to_string(resp).map_err(|e| format!("encode rpc resp: {e}"))?; |
| 590 | let line = format!("__RLM_RESP_{}__::{body}\n", self.session_id); |
| 591 | self.stdin |
| 592 | .write_all(line.as_bytes()) |
| 593 | .await |
| 594 | .map_err(|e| format!("stdin write resp: {e}"))?; |
| 595 | self.stdin |
| 596 | .flush() |
| 597 | .await |
| 598 | .map_err(|e| format!("stdin flush resp: {e}"))?; |
| 599 | Ok(()) |
| 600 | } |
| 601 | |
| 602 | async fn drain_stderr(&mut self) -> String { |
| 603 | // We don't continuously read stderr — drain whatever's pending after |
| 604 | // a round so it can show up in error reports without deadlocking |
| 605 | // anything during normal operation. |
| 606 | let Some(stderr) = self.child.stderr.as_mut() else { |
| 607 | return String::new(); |
| 608 | }; |
| 609 | use tokio::io::AsyncReadExt; |
| 610 | let mut buf = Vec::new(); |
| 611 | // Best-effort read with a tight deadline; we don't want to block. |
| 612 | let fut = async { |
| 613 | let mut chunk = [0u8; 4096]; |
| 614 | loop { |
| 615 | match tokio::time::timeout(Duration::from_millis(20), stderr.read(&mut chunk)).await |
| 616 | { |
| 617 | Ok(Ok(0)) => break, |
| 618 | Ok(Ok(n)) => buf.extend_from_slice(&chunk[..n]), |
| 619 | _ => break, |
| 620 | } |
| 621 | } |
| 622 | }; |
| 623 | let _ = fut.await; |
| 624 | String::from_utf8_lossy(&buf).to_string() |
| 625 | } |
| 626 | |
| 627 | /// True once a round failed or was dropped before reading its own DONE |
| 628 | /// sentinel. Every later round refuses, so a holder that wants to keep |
| 629 | /// going replaces the kernel instead of reusing it. |
| 630 | pub fn is_broken(&self) -> bool { |
| 631 | self.broken.is_some() |
| 632 | } |
| 633 | |
| 634 | /// Total rounds executed. |
| 635 | pub fn round_count(&self) -> u64 { |
| 636 | self.round_count |
| 637 | } |
| 638 | |
| 639 | /// Current per-round timeout policy. RLM context runs intentionally return |
| 640 | /// `None` so long map-reduce jobs are not killed by the old 180s cap. |
| 641 | pub fn round_timeout(&self) -> Option<Duration> { |
| 642 | self.round_timeout |
| 643 | } |
| 644 | |
| 645 | /// Wall-clock uptime since spawn. |
| 646 | pub fn uptime(&self) -> Duration { |
| 647 | self.started.elapsed() |
| 648 | } |
| 649 | |
| 650 | /// Cleanly tear down the subprocess. |
| 651 | pub async fn shutdown(mut self) { |
| 652 | let _ = self.stdin.shutdown().await; |
| 653 | let _ = self.child.kill().await; |
| 654 | if let Some(path) = self.context_path.take() { |
| 655 | let _ = tokio::fs::remove_file(path).await; |
| 656 | } |
| 657 | } |
| 658 | } |
| 659 | |
| 660 | impl Drop for PythonRuntime { |
| 661 | fn drop(&mut self) { |
| 662 | // tokio sets `kill_on_drop(true)` on the child; the context file |
| 663 | // (if any) is removed on `shutdown()` — drop is best-effort. |
| 664 | if let Some(path) = self.context_path.take() { |
| 665 | let _ = std::fs::remove_file(path); |
| 666 | } |
| 667 | } |
| 668 | } |
| 669 | |
| 670 | // --------------------------------------------------------------------------- |
| 671 | // Bootstrap script |
| 672 | // --------------------------------------------------------------------------- |
| 673 | |
| 674 | /// Render the Python bootstrap with session-specific sentinels baked in. |
| 675 | /// The sentinels include a UUID to prevent user prints from being mistaken |
| 676 | /// for control messages. |
| 677 | fn render_bootstrap(session_id: &str) -> String { |
| 678 | BOOTSTRAP_TEMPLATE.replace("__SID__", session_id) |
| 679 | } |
| 680 | |
| 681 | const BOOTSTRAP_TEMPLATE: &str = r#" |
| 682 | import json as _json |
| 683 | import os as _os |
| 684 | import re as _re |
| 685 | import sys as _sys |
| 686 | import traceback as _traceback |
| 687 | |
| 688 | _SID = "__SID__" |
| 689 | _REQ = f"__RLM_REQ_{_SID}__::" |
| 690 | _RESP = f"__RLM_RESP_{_SID}__::" |
| 691 | _FINAL = f"__RLM_FINAL_{_SID}__::" |
| 692 | _ERR = f"__RLM_ERR_{_SID}__::" |
| 693 | _RUN = f"__RLM_RUN_{_SID}__::" |
| 694 | _END = f"__RLM_END_{_SID}__" |
| 695 | _DONE = f"__RLM_DONE_{_SID}__::" |
| 696 | _READY = f"__RLM_READY_{_SID}__" |
| 697 | |
| 698 | def _rpc(req): |
| 699 | _sys.stdout.write(_REQ + _json.dumps(req) + "\n") |
| 700 | _sys.stdout.flush() |
| 701 | line = _sys.stdin.readline() |
| 702 | if not line: |
| 703 | return {"error": "rust driver closed stdin"} |
| 704 | if line.startswith(_RESP): |
| 705 | try: |
| 706 | return _json.loads(line[len(_RESP):]) |
| 707 | except Exception as e: |
| 708 | return {"error": f"malformed rpc resp: {e}"} |
| 709 | return {"error": f"unexpected protocol line: {line[:120]!r}"} |
| 710 | |
| 711 | def llm_query(prompt, model=None, max_tokens=None, system=None): |
| 712 | """One-shot sub-LLM call. The model arg is accepted for compatibility but ignored by Rust.""" |
| 713 | resp = _rpc({"type":"llm","prompt":str(prompt),"model":model, |
| 714 | "max_tokens":max_tokens,"system":system}) |
| 715 | if isinstance(resp, dict) and resp.get("error"): |
| 716 | return f"[llm_query error: {resp['error']}]" |
| 717 | if isinstance(resp, dict): |
| 718 | return resp.get("text","") |
| 719 | return str(resp) |
| 720 | |
| 721 | def _normalize_dependency_mode(mode): |
| 722 | if mode is None: |
| 723 | return "" |
| 724 | return str(mode).strip().lower().replace("-", "_").replace(" ", "_") |
| 725 | |
| 726 | def _batch_dependency_error(helper, prompts, dependency_mode): |
| 727 | mode = _normalize_dependency_mode(dependency_mode) |
| 728 | if mode in ("independent", "parallel_safe", "map_reduce"): |
| 729 | return None |
| 730 | if mode in ("sequential", "dependent", "ordered", "chain", "serial"): |
| 731 | return ( |
| 732 | f"[{helper}: refused parallel batch because dependency_mode={dependency_mode!r}. " |
| 733 | "Use sub_query_sequence(...) or an explicit for-loop with sub_query(...) so each step can consume the previous result.]" |
| 734 | ) |
| 735 | return ( |
| 736 | f"[{helper}: batch helpers require dependency_mode='independent'. " |
| 737 | "Use only for independent slices/items; for A->B dependencies, global-state refactors, migrations, or rollback-sensitive work, use sub_query_sequence(...).]" |
| 738 | ) |
| 739 | |
| 740 | def llm_query_batched(prompts, model=None, dependency_mode=None, safety_note=None): |
| 741 | """Run independent sub-LLM calls concurrently. Declare dependency_mode='independent'.""" |
| 742 | if not isinstance(prompts, (list, tuple)): |
| 743 | return ["[llm_query_batched: prompts must be a list]"] |
| 744 | err = _batch_dependency_error("llm_query_batched", prompts, dependency_mode) |
| 745 | if err is not None: |
| 746 | return [err for _ in prompts] |
| 747 | resp = _rpc({ |
| 748 | "type":"llm_batch", |
| 749 | "prompts":[str(p) for p in prompts], |
| 750 | "model":model, |
| 751 | "dependency_mode":dependency_mode, |
| 752 | "safety_note":safety_note, |
| 753 | }) |
| 754 | if isinstance(resp, dict) and resp.get("error"): |
| 755 | return [f"[llm_query_batched: {resp['error']}]" for _ in prompts] |
| 756 | results = (resp or {}).get("results", []) if isinstance(resp, dict) else [] |
| 757 | if len(results) != len(prompts): |
| 758 | return [f"[llm_query_batched: size mismatch ({len(results)}/{len(prompts)})]" for _ in prompts] |
| 759 | out = [] |
| 760 | for r in results: |
| 761 | if r.get("error"): |
| 762 | out.append(f"[child err: {r['error']}]") |
| 763 | else: |
| 764 | out.append(r.get("text","")) |
| 765 | return out |
| 766 | |
| 767 | def _rlm_child_text(r, err_label): |
| 768 | # A sub-RLM can end with a partial answer AND an error (it ran out of |
| 769 | # rounds before FINAL). Keep the text and mark it incomplete; only an |
| 770 | # error with no text collapses to the error marker. |
| 771 | text = r.get("text") or "" |
| 772 | err = r.get("error") |
| 773 | if not err: |
| 774 | return text |
| 775 | if text.strip(): |
| 776 | return f"{text}\n[rlm_query incomplete: {err}]" |
| 777 | return f"[{err_label}: {err}]" |
| 778 | |
| 779 | def rlm_query(prompt, model=None): |
| 780 | """Recursive sub-RLM. The model arg is accepted for compatibility but ignored by Rust.""" |
| 781 | resp = _rpc({"type":"rlm","prompt":str(prompt),"model":model}) |
| 782 | if isinstance(resp, dict): |
| 783 | return _rlm_child_text(resp, "rlm_query error") |
| 784 | return str(resp) |
| 785 | |
| 786 | def rlm_query_batched(prompts, model=None, dependency_mode=None, safety_note=None): |
| 787 | """Run independent recursive sub-RLMs in parallel. Declare dependency_mode='independent'.""" |
| 788 | if not isinstance(prompts, (list, tuple)): |
| 789 | return ["[rlm_query_batched: prompts must be a list]"] |
| 790 | err = _batch_dependency_error("rlm_query_batched", prompts, dependency_mode) |
| 791 | if err is not None: |
| 792 | return [err for _ in prompts] |
| 793 | resp = _rpc({ |
| 794 | "type":"rlm_batch", |
| 795 | "prompts":[str(p) for p in prompts], |
| 796 | "model":model, |
| 797 | "dependency_mode":dependency_mode, |
| 798 | "safety_note":safety_note, |
| 799 | }) |
| 800 | if isinstance(resp, dict) and resp.get("error"): |
| 801 | return [f"[rlm_query_batched: {resp['error']}]" for _ in prompts] |
| 802 | results = (resp or {}).get("results", []) if isinstance(resp, dict) else [] |
| 803 | if len(results) != len(prompts): |
| 804 | return [f"[rlm_query_batched: size mismatch ({len(results)}/{len(prompts)})]" for _ in prompts] |
| 805 | return [_rlm_child_text(r, "child err") for r in results] |
| 806 | |
| 807 | def _slice_text(slice_value): |
| 808 | if slice_value is None: |
| 809 | return "" |
| 810 | if isinstance(slice_value, dict): |
| 811 | if "text" in slice_value: |
| 812 | return str(slice_value["text"]) |
| 813 | return _json.dumps(slice_value, ensure_ascii=False) |
| 814 | return str(slice_value) |
| 815 | |
| 816 | def _prompt_with_slice(prompt, slice_value): |
| 817 | text = _slice_text(slice_value) |
| 818 | if not text: |
| 819 | return str(prompt) |
| 820 | if isinstance(slice_value, dict) and ("index" in slice_value or ("start" in slice_value and "end" in slice_value)): |
| 821 | label = f"slice index={slice_value.get('index', '?')} range={slice_value.get('start', '?')}:{slice_value.get('end', '?')}" |
| 822 | else: |
| 823 | label = "slice" |
| 824 | return f"{prompt}\n\n--- {label} ---\n{text}" |
| 825 | |
| 826 | def sub_query(prompt, slice=None, timeout_secs=None, **kwargs): |
| 827 | """One child LLM call, optionally scoped to a bounded slice.""" |
| 828 | return llm_query(_prompt_with_slice(prompt, slice)) |
| 829 | |
| 830 | def sub_query_batch(prompt, slices, timeout_secs=None, dependency_mode=None, safety_note=None, **kwargs): |
| 831 | """Apply one prompt to many independent bounded slices concurrently.""" |
| 832 | if not isinstance(slices, (list, tuple)): |
| 833 | return ["[sub_query_batch: slices must be a list]"] |
| 834 | return llm_query_batched( |
| 835 | [_prompt_with_slice(prompt, s) for s in slices], |
| 836 | dependency_mode=dependency_mode, |
| 837 | safety_note=safety_note, |
| 838 | ) |
| 839 | |
| 840 | def sub_query_map(prompts, slices=None, timeout_secs=None, dependency_mode=None, safety_note=None, **kwargs): |
| 841 | """Run N distinct independent prompts, optionally paired with N bounded slices.""" |
| 842 | if not isinstance(prompts, (list, tuple)): |
| 843 | return ["[sub_query_map: prompts must be a list]"] |
| 844 | if slices is None: |
| 845 | return llm_query_batched( |
| 846 | [str(p) for p in prompts], |
| 847 | dependency_mode=dependency_mode, |
| 848 | safety_note=safety_note, |
| 849 | ) |
| 850 | if not isinstance(slices, (list, tuple)): |
| 851 | return ["[sub_query_map: slices must be a list]"] |
| 852 | if len(prompts) != len(slices): |
| 853 | return [f"[sub_query_map: size mismatch ({len(prompts)}/{len(slices)})]" for _ in prompts] |
| 854 | return llm_query_batched( |
| 855 | [_prompt_with_slice(p, s) for p, s in zip(prompts, slices)], |
| 856 | dependency_mode=dependency_mode, |
| 857 | safety_note=safety_note, |
| 858 | ) |
| 859 | |
| 860 | def sub_query_sequence(prompt, slices, carry_prompt=None, timeout_secs=None, **kwargs): |
| 861 | """Apply one prompt to slices sequentially, feeding each result into the next step.""" |
| 862 | if not isinstance(slices, (list, tuple)): |
| 863 | return ["[sub_query_sequence: slices must be a list]"] |
| 864 | out = [] |
| 865 | previous = "" |
| 866 | carry = str(carry_prompt or "Previous step result; treat it as required input for this step:") |
| 867 | total = len(slices) |
| 868 | for i, s in enumerate(slices): |
| 869 | step_prompt = _prompt_with_slice(prompt, s) |
| 870 | if previous: |
| 871 | step_prompt = ( |
| 872 | f"{step_prompt}\n\n--- dependency_state step {i}/{total} ---\n" |
| 873 | f"{carry}\n{previous}" |
| 874 | ) |
| 875 | result = llm_query(step_prompt) |
| 876 | out.append(result) |
| 877 | previous = result |
| 878 | return out |
| 879 | |
| 880 | def sub_rlm(prompt, source=None, timeout_secs=None, **kwargs): |
| 881 | """Recursive sub-RLM call for tasks that need their own decomposition.""" |
| 882 | return rlm_query(_prompt_with_slice(prompt, source)) |
| 883 | |
| 884 | def _json_safe(value): |
| 885 | try: |
| 886 | _json.dumps(value, ensure_ascii=False) |
| 887 | return value |
| 888 | except Exception: |
| 889 | return str(value) |
| 890 | |
| 891 | def _emit_final(value, confidence=None): |
| 892 | safe_value = _json_safe(value) |
| 893 | _sys.stdout.write(_FINAL + _json.dumps({ |
| 894 | "value": safe_value, |
| 895 | "confidence": confidence, |
| 896 | }, ensure_ascii=False) + "\n") |
| 897 | _sys.stdout.flush() |
| 898 | |
| 899 | def FINAL(value): |
| 900 | """Legacy compatibility alias for finalize(value).""" |
| 901 | _emit_final(value) |
| 902 | |
| 903 | def FINAL_VAR(name): |
| 904 | """Legacy compatibility alias for finalize(repl_get(name)).""" |
| 905 | name_str = str(name).strip().strip("'\"") |
| 906 | if name_str in globals(): |
| 907 | _emit_final(globals()[name_str]) |
| 908 | else: |
| 909 | print(f"FINAL_VAR error: variable '{name_str}' not found. " |
| 910 | f"Use SHOW_VARS() to list available variables.", flush=True) |
| 911 | |
| 912 | def SHOW_VARS(): |
| 913 | """Return a dict of {name: type-name} for all user variables in the REPL.""" |
| 914 | out = {} |
| 915 | for k, v in list(globals().items()): |
| 916 | if k.startswith('_') or k in _BOOTSTRAP_NAMES: |
| 917 | continue |
| 918 | out[k] = type(v).__name__ |
| 919 | return out |
| 920 | |
| 921 | def repl_get(name, default=None): |
| 922 | return globals().get(str(name), default) |
| 923 | |
| 924 | def repl_set(name, value): |
| 925 | globals()[str(name)] = value |
| 926 | |
| 927 | def context_meta(): |
| 928 | """Return bounded metadata about the loaded input; never includes the full text.""" |
| 929 | text = _context |
| 930 | line_count = 0 if text == "" else text.count("\n") + (0 if text.endswith("\n") else 1) |
| 931 | return { |
| 932 | "chars": len(text), |
| 933 | "lines": line_count, |
| 934 | "preview": text[:500], |
| 935 | "tail_preview": text[-500:] if len(text) > 500 else text, |
| 936 | } |
| 937 | |
| 938 | def _slice_chars(start, end): |
| 939 | total = len(_context) |
| 940 | s = max(0, int(start)) |
| 941 | e = max(s, min(total, int(end))) |
| 942 | return _context[s:e] |
| 943 | |
| 944 | def _slice_lines(start, end): |
| 945 | lines = _context.splitlines() |
| 946 | s = max(0, int(start)) |
| 947 | e = max(s, min(len(lines), int(end))) |
| 948 | return "\n".join(lines[s:e]) |
| 949 | |
| 950 | def peek(start, end, unit="chars"): |
| 951 | """Return a bounded slice of the input by char offsets or line numbers.""" |
| 952 | if str(unit).lower() in ("line", "lines"): |
| 953 | return _slice_lines(start, end) |
| 954 | if str(unit).lower() not in ("char", "chars"): |
| 955 | raise ValueError("unit must be 'chars' or 'lines'") |
| 956 | return _slice_chars(start, end) |
| 957 | |
| 958 | def search(pattern, max_hits=100): |
| 959 | """Regex-search the input and return bounded hit records with snippets.""" |
| 960 | max_hits = max(0, int(max_hits)) |
| 961 | hits = [] |
| 962 | if max_hits == 0: |
| 963 | return hits |
| 964 | rx = _re.compile(str(pattern), _re.MULTILINE) |
| 965 | for i, m in enumerate(rx.finditer(_context)): |
| 966 | if i >= max_hits: |
| 967 | break |
| 968 | start, end = m.span() |
| 969 | snippet_start = max(0, start - 120) |
| 970 | snippet_end = min(len(_context), end + 120) |
| 971 | hits.append({ |
| 972 | "index": i, |
| 973 | "start": start, |
| 974 | "end": end, |
| 975 | "match": m.group(0), |
| 976 | "snippet": _context[snippet_start:snippet_end], |
| 977 | }) |
| 978 | return hits |
| 979 | |
| 980 | def chunk(max_chars=20000, overlap=0): |
| 981 | """Return full-coverage input chunks with index/start/end/text fields.""" |
| 982 | max_chars = int(max_chars) |
| 983 | overlap = max(0, int(overlap)) |
| 984 | if max_chars <= 0: |
| 985 | raise ValueError("max_chars must be > 0") |
| 986 | if overlap >= max_chars: |
| 987 | raise ValueError("overlap must be smaller than max_chars") |
| 988 | chunks = [] |
| 989 | start = 0 |
| 990 | idx = 0 |
| 991 | total = len(_context) |
| 992 | while start < total: |
| 993 | end = min(total, start + max_chars) |
| 994 | chunks.append({"index": idx, "start": start, "end": end, "text": _context[start:end]}) |
| 995 | idx += 1 |
| 996 | if end >= total: |
| 997 | break |
| 998 | start = end - overlap |
| 999 | return chunks |
| 1000 | |
| 1001 | def chunk_context(max_chars=20000, overlap=0): |
| 1002 | """Compatibility alias for chunk().""" |
| 1003 | return chunk(max_chars=max_chars, overlap=overlap) |
| 1004 | |
| 1005 | def chunk_coverage(chunks): |
| 1006 | """Summarize coverage for chunks produced by chunk().""" |
| 1007 | spans = [] |
| 1008 | for c in chunks: |
| 1009 | try: |
| 1010 | spans.append((int(c["start"]), int(c["end"]))) |
| 1011 | except Exception: |
| 1012 | continue |
| 1013 | spans.sort() |
| 1014 | covered = 0 |
| 1015 | cursor = 0 |
| 1016 | gaps = [] |
| 1017 | for start, end in spans: |
| 1018 | if start > cursor: |
| 1019 | gaps.append((cursor, start)) |
| 1020 | if end > cursor: |
| 1021 | covered += end - max(start, cursor) |
| 1022 | cursor = end |
| 1023 | if cursor < len(_context): |
| 1024 | gaps.append((cursor, len(_context))) |
| 1025 | return { |
| 1026 | "chunks": len(chunks), |
| 1027 | "context_chars": len(_context), |
| 1028 | "input_chars": len(_context), |
| 1029 | "covered_chars": covered, |
| 1030 | "gaps": gaps, |
| 1031 | "complete": covered >= len(_context) and not gaps, |
| 1032 | } |
| 1033 | |
| 1034 | def finalize(value, confidence=None): |
| 1035 | """Signal the session's final answer and persist confidence metadata.""" |
| 1036 | global final_answer, final_confidence, final_result |
| 1037 | final_answer = _json_safe(value) |
| 1038 | final_confidence = confidence |
| 1039 | final_result = { |
| 1040 | "value": final_answer, |
| 1041 | "confidence": confidence, |
| 1042 | } |
| 1043 | _emit_final(final_answer, confidence=confidence) |
| 1044 | return final_answer |
| 1045 | |
| 1046 | def evaluate_progress(): |
| 1047 | """Return lightweight state useful before deciding the next REPL step.""" |
| 1048 | vars_now = SHOW_VARS() |
| 1049 | return { |
| 1050 | "has_final_answer": "final_answer" in globals(), |
| 1051 | "final_confidence": globals().get("final_confidence", None), |
| 1052 | "user_variables": vars_now, |
| 1053 | } |
| 1054 | |
| 1055 | # Load the long input from a file. This keeps the big string out of the |
| 1056 | # process command-line and out of the LLM's window. |
| 1057 | _ctx_file = _os.environ.get("RLM_CONTEXT_FILE","") |
| 1058 | _context = "" |
| 1059 | if _ctx_file: |
| 1060 | try: |
| 1061 | with open(_ctx_file, "r", encoding="utf-8", errors="replace") as f: |
| 1062 | _context = f.read() |
| 1063 | except Exception as e: |
| 1064 | _sys.stderr.write(f"[bootstrap] failed to load context: {e}\n") |
| 1065 | content = _context |
| 1066 | |
| 1067 | def _replace_context_file(path): |
| 1068 | """Atomically switch bounded helpers to a freshly written context file.""" |
| 1069 | global _context, content |
| 1070 | with open(path, "r", encoding="utf-8", errors="replace") as f: |
| 1071 | _context = f.read() |
| 1072 | content = _context |
| 1073 | return context_meta() |
| 1074 | |
| 1075 | _BOOTSTRAP_NAMES = { |
| 1076 | "_SID","_REQ","_RESP","_FINAL","_ERR","_RUN","_END","_DONE","_READY", |
| 1077 | "_rpc","_ctx_file","_context","_slice_chars","_slice_lines","_replace_context_file","_BOOTSTRAP_NAMES","_main_loop", |
| 1078 | "_emit_final","_json_safe","_slice_text","_prompt_with_slice", |
| 1079 | "_normalize_dependency_mode","_batch_dependency_error", |
| 1080 | "llm_query","llm_query_batched","rlm_query","rlm_query_batched", |
| 1081 | "sub_query","sub_query_batch","sub_query_map","sub_query_sequence","sub_rlm", |
| 1082 | "FINAL","FINAL_VAR","SHOW_VARS","repl_get","repl_set", |
| 1083 | "context_meta","peek","search","chunk","chunk_context","chunk_coverage", |
| 1084 | "finalize","evaluate_progress","content", |
| 1085 | "_json","_os","_re","_sys","_traceback", |
| 1086 | } |
| 1087 | |
| 1088 | def _main_loop(): |
| 1089 | _sys.stdout.write(_READY + "\n") |
| 1090 | _sys.stdout.flush() |
| 1091 | while True: |
| 1092 | header = _sys.stdin.readline() |
| 1093 | if not header: |
| 1094 | return |
| 1095 | if not header.startswith(_RUN): |
| 1096 | continue |
| 1097 | round_id = header.rstrip("\n")[len(_RUN):] |
| 1098 | code_lines = [] |
| 1099 | while True: |
| 1100 | line = _sys.stdin.readline() |
| 1101 | if not line: |
| 1102 | return |
| 1103 | if line.rstrip("\n") == _END: |
| 1104 | break |
| 1105 | code_lines.append(line) |
| 1106 | code = "".join(code_lines) |
| 1107 | try: |
| 1108 | exec(compile(code, f"<repl-{round_id}>", "exec"), globals()) |
| 1109 | except SystemExit: |
| 1110 | _sys.stdout.write(_DONE + round_id + "\n") |
| 1111 | _sys.stdout.flush() |
| 1112 | return |
| 1113 | except BaseException: |
| 1114 | tb = _traceback.format_exc() |
| 1115 | _sys.stdout.write(_ERR + _json.dumps(tb) + "\n") |
| 1116 | _sys.stdout.flush() |
| 1117 | _sys.stdout.write(_DONE + round_id + "\n") |
| 1118 | _sys.stdout.flush() |
| 1119 | |
| 1120 | _main_loop() |
| 1121 | "#; |
| 1122 | |
| 1123 | // --------------------------------------------------------------------------- |
| 1124 | // Helpers |
| 1125 | // --------------------------------------------------------------------------- |
| 1126 | |
| 1127 | /// Append to a round's retained stdout without exceeding |
| 1128 | /// [`MAX_ROUND_STDOUT_BYTES`]; whatever does not fit is only counted. |
| 1129 | fn push_bounded(buf: &mut String, text: &str, omitted: &mut usize) { |
| 1130 | let room = MAX_ROUND_STDOUT_BYTES.saturating_sub(buf.len()); |
| 1131 | if text.len() <= room { |
| 1132 | buf.push_str(text); |
| 1133 | return; |
| 1134 | } |
| 1135 | let cut = (0..=room) |
| 1136 | .rev() |
| 1137 | .find(|&index| text.is_char_boundary(index)) |
| 1138 | .unwrap_or(0); |
| 1139 | buf.push_str(&text[..cut]); |
| 1140 | *omitted += text.len() - cut; |
| 1141 | } |
| 1142 | |
| 1143 | fn truncate_stdout(stdout: &str, limit: usize) -> String { |
| 1144 | if stdout.len() <= limit { |
| 1145 | return stdout.to_string(); |
| 1146 | } |
| 1147 | let take = limit.saturating_sub(80); |
| 1148 | let mut out: String = stdout.chars().take(take).collect(); |
| 1149 | let omitted = stdout.len().saturating_sub(out.len()); |
| 1150 | out.push_str(&format!( |
| 1151 | "\n\n[... REPL output truncated: {omitted} bytes omitted ...]\n" |
| 1152 | )); |
| 1153 | out |
| 1154 | } |
| 1155 | |
| 1156 | // --------------------------------------------------------------------------- |
| 1157 | // Tests |
| 1158 | // --------------------------------------------------------------------------- |
| 1159 | |
| 1160 | #[cfg(test)] |
| 1161 | mod tests { |
| 1162 | use super::*; |
| 1163 | use std::sync::Arc; |
| 1164 | use std::sync::atomic::{AtomicU32, Ordering}; |
| 1165 | use tokio::sync::Mutex; |
| 1166 | |
| 1167 | /// In-process dispatcher that records what was asked and replies with |
| 1168 | /// canned text. Lets tests verify the round-trip without real network. |
| 1169 | struct StubBridge { |
| 1170 | calls: Arc<Mutex<Vec<RpcRequest>>>, |
| 1171 | canned: Arc<AtomicU32>, |
| 1172 | } |
| 1173 | |
| 1174 | impl StubBridge { |
| 1175 | fn new() -> Self { |
| 1176 | Self { |
| 1177 | calls: Arc::new(Mutex::new(Vec::new())), |
| 1178 | canned: Arc::new(AtomicU32::new(0)), |
| 1179 | } |
| 1180 | } |
| 1181 | } |
| 1182 | |
| 1183 | impl RpcDispatcher for StubBridge { |
| 1184 | fn dispatch<'a>( |
| 1185 | &'a self, |
| 1186 | req: RpcRequest, |
| 1187 | ) -> std::pin::Pin<Box<dyn std::future::Future<Output = RpcResponse> + Send + 'a>> { |
| 1188 | Box::pin(async move { |
| 1189 | self.calls.lock().await.push(req.clone()); |
| 1190 | let n = self.canned.fetch_add(1, Ordering::Relaxed); |
| 1191 | match req { |
| 1192 | RpcRequest::Llm { prompt, .. } | RpcRequest::Rlm { prompt, .. } => { |
| 1193 | RpcResponse::Single(SingleResp { |
| 1194 | text: format!("stub#{n}: {prompt}"), |
| 1195 | error: None, |
| 1196 | }) |
| 1197 | } |
| 1198 | RpcRequest::LlmBatch { prompts, .. } | RpcRequest::RlmBatch { prompts, .. } => { |
| 1199 | let results = prompts |
| 1200 | .into_iter() |
| 1201 | .enumerate() |
| 1202 | .map(|(i, p)| SingleResp { |
| 1203 | text: format!("stub#{n}.{i}: {p}"), |
| 1204 | error: None, |
| 1205 | }) |
| 1206 | .collect(); |
| 1207 | RpcResponse::Batch(BatchResp { results }) |
| 1208 | } |
| 1209 | } |
| 1210 | }) |
| 1211 | } |
| 1212 | } |
| 1213 | |
| 1214 | fn write_temp_context(body: &str) -> std::path::PathBuf { |
| 1215 | let dir = std::env::temp_dir().join("deepseek_repl_runtime_tests"); |
| 1216 | std::fs::create_dir_all(&dir).unwrap(); |
| 1217 | let path = dir.join(format!("ctx_{}_{}.txt", std::process::id(), Uuid::new_v4())); |
| 1218 | std::fs::write(&path, body).unwrap(); |
| 1219 | path |
| 1220 | } |
| 1221 | |
| 1222 | #[tokio::test] |
| 1223 | async fn spawns_and_executes_simple_print() { |
| 1224 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1225 | let round = rt.execute("print('hello world')").await.expect("execute"); |
| 1226 | assert!(round.stdout.contains("hello world")); |
| 1227 | assert!(!round.has_error); |
| 1228 | assert!(round.final_value.is_none()); |
| 1229 | assert_eq!(round.rpc_count, 0); |
| 1230 | rt.shutdown().await; |
| 1231 | } |
| 1232 | |
| 1233 | #[tokio::test] |
| 1234 | async fn non_utf8_stdout_decodes_lossy_and_runtime_survives() { |
| 1235 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1236 | let round = rt |
| 1237 | .execute( |
| 1238 | "import sys\n\ |
| 1239 | sys.stdout.buffer.write(b'bad:\\xff\\n')\n\ |
| 1240 | sys.stdout.buffer.flush()\n\ |
| 1241 | print('after invalid')", |
| 1242 | ) |
| 1243 | .await |
| 1244 | .expect("execute"); |
| 1245 | |
| 1246 | assert!(round.stdout.contains("bad:\u{fffd}"), "{}", round.stdout); |
| 1247 | assert!(round.stdout.contains("after invalid"), "{}", round.stdout); |
| 1248 | rt.shutdown().await; |
| 1249 | } |
| 1250 | |
| 1251 | #[tokio::test] |
| 1252 | async fn variables_persist_across_rounds() { |
| 1253 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1254 | rt.execute("x = [1, 2, 3]").await.expect("r1"); |
| 1255 | rt.execute("x.append(99)").await.expect("r2"); |
| 1256 | let round = rt.execute("print(x)").await.expect("r3"); |
| 1257 | assert!(round.stdout.contains("[1, 2, 3, 99]")); |
| 1258 | rt.shutdown().await; |
| 1259 | } |
| 1260 | |
| 1261 | #[tokio::test] |
| 1262 | async fn imports_persist_across_rounds() { |
| 1263 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1264 | rt.execute("import math").await.expect("r1"); |
| 1265 | let round = rt.execute("print(math.pi)").await.expect("r2"); |
| 1266 | assert!(round.stdout.contains("3.14")); |
| 1267 | rt.shutdown().await; |
| 1268 | } |
| 1269 | |
| 1270 | #[tokio::test] |
| 1271 | async fn context_loads_from_file() { |
| 1272 | let path = write_temp_context("the quick brown fox"); |
| 1273 | let mut rt = PythonRuntime::spawn_with_context(&path) |
| 1274 | .await |
| 1275 | .expect("spawn"); |
| 1276 | let round = rt |
| 1277 | .execute("print(context_meta()['chars'], peek(0, 5))") |
| 1278 | .await |
| 1279 | .expect("execute"); |
| 1280 | assert!(round.stdout.contains("19")); |
| 1281 | assert!(round.stdout.contains("the q")); |
| 1282 | rt.shutdown().await; |
| 1283 | } |
| 1284 | |
| 1285 | #[tokio::test] |
| 1286 | async fn context_aliases_keep_common_content_name_bounded() { |
| 1287 | let path = write_temp_context("aleph-style"); |
| 1288 | let mut rt = PythonRuntime::spawn_with_context(&path) |
| 1289 | .await |
| 1290 | .expect("spawn"); |
| 1291 | let round = rt |
| 1292 | .execute("print(content == _context, 'context' in globals(), 'ctx' in globals())") |
| 1293 | .await |
| 1294 | .expect("execute"); |
| 1295 | assert!(round.stdout.contains("True False False")); |
| 1296 | rt.shutdown().await; |
| 1297 | } |
| 1298 | |
| 1299 | #[tokio::test] |
| 1300 | async fn replacing_context_keeps_kernel_variables_and_refreshes_helpers() { |
| 1301 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1302 | rt.execute("remembered = {'answer': 42}") |
| 1303 | .await |
| 1304 | .expect("seed persistent variable"); |
| 1305 | rt.replace_context("fresh transcript\nwith a needle") |
| 1306 | .await |
| 1307 | .expect("refresh context"); |
| 1308 | |
| 1309 | let round = rt |
| 1310 | .execute( |
| 1311 | "print(remembered['answer'])\n\ |
| 1312 | print(context_meta()['chars'])\n\ |
| 1313 | print(search('needle')[0]['match'])", |
| 1314 | ) |
| 1315 | .await |
| 1316 | .expect("inspect refreshed context"); |
| 1317 | |
| 1318 | assert!(round.stdout.contains("42"), "{}", round.stdout); |
| 1319 | assert!(round.stdout.contains("30"), "{}", round.stdout); |
| 1320 | assert!(round.stdout.contains("needle"), "{}", round.stdout); |
| 1321 | rt.shutdown().await; |
| 1322 | } |
| 1323 | |
| 1324 | #[tokio::test] |
| 1325 | async fn context_chunk_helpers_report_full_coverage() { |
| 1326 | let path = write_temp_context("abcdefghijklmnopqrstuvwxyz"); |
| 1327 | let mut rt = PythonRuntime::spawn_with_context(&path) |
| 1328 | .await |
| 1329 | .expect("spawn"); |
| 1330 | let round = rt |
| 1331 | .execute( |
| 1332 | "chunks = chunk_context(max_chars=10)\n\ |
| 1333 | coverage = chunk_coverage(chunks)\n\ |
| 1334 | print(len(chunks), coverage['covered_chars'], coverage['complete'])", |
| 1335 | ) |
| 1336 | .await |
| 1337 | .expect("execute"); |
| 1338 | assert!(round.stdout.contains("3 26 True"), "{}", round.stdout); |
| 1339 | rt.shutdown().await; |
| 1340 | } |
| 1341 | |
| 1342 | #[tokio::test] |
| 1343 | async fn bounded_input_helpers_work() { |
| 1344 | let path = write_temp_context("alpha\nbeta needle\ngamma needle\nomega"); |
| 1345 | let mut rt = PythonRuntime::spawn_with_context(&path) |
| 1346 | .await |
| 1347 | .expect("spawn"); |
| 1348 | let round = rt |
| 1349 | .execute( |
| 1350 | "meta = context_meta()\n\ |
| 1351 | hits = search('needle', max_hits=1)\n\ |
| 1352 | print(meta['chars'], meta['lines'])\n\ |
| 1353 | print(peek(6, 17))\n\ |
| 1354 | print(peek(1, 3, unit='lines'))\n\ |
| 1355 | print(len(hits), hits[0]['match'], hits[0]['start'])", |
| 1356 | ) |
| 1357 | .await |
| 1358 | .expect("execute"); |
| 1359 | let stdout = round.stdout.replace("\r\n", "\n"); |
| 1360 | assert!(stdout.contains("36 4"), "{stdout}"); |
| 1361 | assert!(stdout.contains("beta needle"), "{stdout}"); |
| 1362 | assert!(stdout.contains("beta needle\ngamma needle"), "{stdout}"); |
| 1363 | assert!(stdout.contains("1 needle 11"), "{stdout}"); |
| 1364 | rt.shutdown().await; |
| 1365 | } |
| 1366 | |
| 1367 | #[tokio::test] |
| 1368 | async fn new_chunk_helper_reports_full_coverage() { |
| 1369 | let path = write_temp_context("abcdefghijklmnopqrstuvwxyz"); |
| 1370 | let mut rt = PythonRuntime::spawn_with_context(&path) |
| 1371 | .await |
| 1372 | .expect("spawn"); |
| 1373 | let round = rt |
| 1374 | .execute( |
| 1375 | "chunks = chunk(max_chars=10)\n\ |
| 1376 | coverage = chunk_coverage(chunks)\n\ |
| 1377 | print(len(chunks), coverage['input_chars'], coverage['covered_chars'], coverage['complete'])", |
| 1378 | ) |
| 1379 | .await |
| 1380 | .expect("execute"); |
| 1381 | assert!(round.stdout.contains("3 26 26 True"), "{}", round.stdout); |
| 1382 | rt.shutdown().await; |
| 1383 | } |
| 1384 | |
| 1385 | #[tokio::test] |
| 1386 | async fn finalize_helper_is_captured_directly() { |
| 1387 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1388 | let round = rt |
| 1389 | .execute("finalize('computed answer', confidence='high')") |
| 1390 | .await |
| 1391 | .expect("execute"); |
| 1392 | assert_eq!(round.final_value.as_deref(), Some("computed answer")); |
| 1393 | assert_eq!( |
| 1394 | round.final_json.as_ref().and_then(Value::as_str), |
| 1395 | Some("computed answer") |
| 1396 | ); |
| 1397 | assert_eq!( |
| 1398 | round.final_confidence.as_ref().and_then(Value::as_str), |
| 1399 | Some("high") |
| 1400 | ); |
| 1401 | rt.shutdown().await; |
| 1402 | } |
| 1403 | |
| 1404 | #[tokio::test] |
| 1405 | async fn finalize_preserves_json_values_for_handles() { |
| 1406 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1407 | let round = rt |
| 1408 | .execute("finalize({'answer': 42, 'items': ['a', 'b']})") |
| 1409 | .await |
| 1410 | .expect("execute"); |
| 1411 | |
| 1412 | assert_eq!( |
| 1413 | round.final_value.as_deref(), |
| 1414 | Some(r#"{"answer":42,"items":["a","b"]}"#) |
| 1415 | ); |
| 1416 | assert_eq!( |
| 1417 | round.final_json, |
| 1418 | Some(serde_json::json!({"answer": 42, "items": ["a", "b"]})) |
| 1419 | ); |
| 1420 | rt.shutdown().await; |
| 1421 | } |
| 1422 | |
| 1423 | #[tokio::test] |
| 1424 | async fn sub_query_accepts_timeout_keyword_for_agent_guesses() { |
| 1425 | let bridge = StubBridge::new(); |
| 1426 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1427 | let round = rt |
| 1428 | .run( |
| 1429 | "answer = sub_query('summarize', timeout_secs=2)\nprint(answer)", |
| 1430 | Some(&bridge), |
| 1431 | ) |
| 1432 | .await |
| 1433 | .expect("execute"); |
| 1434 | |
| 1435 | assert!(!round.has_error, "{}", round.stdout); |
| 1436 | assert!( |
| 1437 | round.stdout.contains("stub#0: summarize"), |
| 1438 | "{}", |
| 1439 | round.stdout |
| 1440 | ); |
| 1441 | rt.shutdown().await; |
| 1442 | } |
| 1443 | |
| 1444 | #[tokio::test] |
| 1445 | async fn rlm_context_runtime_has_no_fixed_round_timeout() { |
| 1446 | let path = write_temp_context("long input"); |
| 1447 | let rt = PythonRuntime::spawn_with_context(&path) |
| 1448 | .await |
| 1449 | .expect("spawn"); |
| 1450 | assert!( |
| 1451 | rt.round_timeout().is_none(), |
| 1452 | "RLM context runs must not inherit the old 180s REPL round timeout" |
| 1453 | ); |
| 1454 | rt.shutdown().await; |
| 1455 | } |
| 1456 | |
| 1457 | #[tokio::test] |
| 1458 | async fn timed_out_round_kills_the_interpreter_and_refuses_later_rounds() { |
| 1459 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1460 | rt.round_timeout = Some(Duration::from_millis(300)); |
| 1461 | let error = rt |
| 1462 | .execute("import time\ntime.sleep(30)\nprint('late output')") |
| 1463 | .await |
| 1464 | .expect_err("the round must time out"); |
| 1465 | assert!(error.contains("timed out"), "{error}"); |
| 1466 | |
| 1467 | let refused = rt |
| 1468 | .execute("print('next round')") |
| 1469 | .await |
| 1470 | .expect_err("a desynchronized REPL must not run another round"); |
| 1471 | assert!(refused.contains("unusable"), "{refused}"); |
| 1472 | |
| 1473 | let deadline = Instant::now() + Duration::from_secs(5); |
| 1474 | while rt.child.try_wait().expect("child status").is_none() { |
| 1475 | assert!( |
| 1476 | Instant::now() < deadline, |
| 1477 | "the timed-out interpreter must be killed" |
| 1478 | ); |
| 1479 | tokio::time::sleep(Duration::from_millis(25)).await; |
| 1480 | } |
| 1481 | } |
| 1482 | |
| 1483 | #[tokio::test] |
| 1484 | async fn another_rounds_done_sentinel_does_not_end_this_round() { |
| 1485 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1486 | let round = rt |
| 1487 | .execute("print(_DONE + '999')\nprint('still this round')") |
| 1488 | .await |
| 1489 | .expect("execute"); |
| 1490 | assert!( |
| 1491 | round.stdout.contains("still this round"), |
| 1492 | "{}", |
| 1493 | round.stdout |
| 1494 | ); |
| 1495 | let next = rt.execute("print('second')").await.expect("next round"); |
| 1496 | assert!(next.stdout.contains("second"), "{}", next.stdout); |
| 1497 | assert!(!next.stdout.contains("still this round"), "{}", next.stdout); |
| 1498 | rt.shutdown().await; |
| 1499 | } |
| 1500 | |
| 1501 | #[tokio::test] |
| 1502 | async fn inline_runtime_keeps_bounded_round_timeout() { |
| 1503 | let rt = PythonRuntime::new().await.expect("spawn"); |
| 1504 | assert_eq!(rt.round_timeout(), Some(ROUND_TIMEOUT)); |
| 1505 | rt.shutdown().await; |
| 1506 | } |
| 1507 | |
| 1508 | #[tokio::test] |
| 1509 | async fn legacy_final_is_captured() { |
| 1510 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1511 | let round = rt |
| 1512 | .execute("FINAL('the answer is 42')") |
| 1513 | .await |
| 1514 | .expect("execute"); |
| 1515 | assert_eq!(round.final_value.as_deref(), Some("the answer is 42")); |
| 1516 | rt.shutdown().await; |
| 1517 | } |
| 1518 | |
| 1519 | #[tokio::test] |
| 1520 | async fn legacy_final_var_is_captured() { |
| 1521 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1522 | rt.execute("answer = 'computed'").await.expect("r1"); |
| 1523 | let round = rt.execute("FINAL_VAR('answer')").await.expect("r2"); |
| 1524 | assert_eq!(round.final_value.as_deref(), Some("computed")); |
| 1525 | rt.shutdown().await; |
| 1526 | } |
| 1527 | |
| 1528 | #[tokio::test] |
| 1529 | async fn errors_are_reported_without_killing_runtime() { |
| 1530 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1531 | let r1 = rt.execute("raise ValueError('boom')").await.expect("r1"); |
| 1532 | assert!(r1.has_error); |
| 1533 | assert!(r1.full_stdout.contains("boom") || r1.stdout.contains("boom")); |
| 1534 | // The runtime is still alive — next round should work. |
| 1535 | let r2 = rt.execute("print('still here')").await.expect("r2"); |
| 1536 | assert!(r2.stdout.contains("still here")); |
| 1537 | rt.shutdown().await; |
| 1538 | } |
| 1539 | |
| 1540 | #[tokio::test] |
| 1541 | async fn rpc_dispatcher_round_trips_llm_query() { |
| 1542 | let bridge = StubBridge::new(); |
| 1543 | let calls = Arc::clone(&bridge.calls); |
| 1544 | |
| 1545 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1546 | let round = rt |
| 1547 | .run("print(llm_query('hello'))", Some(&bridge)) |
| 1548 | .await |
| 1549 | .expect("execute"); |
| 1550 | assert!( |
| 1551 | round.stdout.contains("stub#0: hello"), |
| 1552 | "stdout: {:?}", |
| 1553 | round.stdout |
| 1554 | ); |
| 1555 | assert_eq!(round.rpc_count, 1); |
| 1556 | |
| 1557 | let recorded = calls.lock().await; |
| 1558 | assert_eq!(recorded.len(), 1); |
| 1559 | match &recorded[0] { |
| 1560 | RpcRequest::Llm { prompt, .. } => assert_eq!(prompt, "hello"), |
| 1561 | other => panic!("expected Llm request, got {other:?}"), |
| 1562 | } |
| 1563 | drop(recorded); |
| 1564 | rt.shutdown().await; |
| 1565 | } |
| 1566 | |
| 1567 | /// A sub-RLM that ran out of rounds returns its last answer AND an |
| 1568 | /// error. Answers every recursive request with that shape. |
| 1569 | struct ExhaustedRlmBridge; |
| 1570 | |
| 1571 | impl RpcDispatcher for ExhaustedRlmBridge { |
| 1572 | fn dispatch<'a>( |
| 1573 | &'a self, |
| 1574 | req: RpcRequest, |
| 1575 | ) -> std::pin::Pin<Box<dyn std::future::Future<Output = RpcResponse> + Send + 'a>> { |
| 1576 | Box::pin(async move { |
| 1577 | let partial = |i: usize| SingleResp { |
| 1578 | text: if i == 1 { |
| 1579 | String::new() |
| 1580 | } else { |
| 1581 | format!("partial answer {i}") |
| 1582 | }, |
| 1583 | error: Some("RLM loop exhausted after 25 iterations without FINAL".to_string()), |
| 1584 | }; |
| 1585 | match req { |
| 1586 | RpcRequest::RlmBatch { prompts, .. } => RpcResponse::Batch(BatchResp { |
| 1587 | results: (0..prompts.len()).map(partial).collect(), |
| 1588 | }), |
| 1589 | _ => RpcResponse::Single(partial(0)), |
| 1590 | } |
| 1591 | }) |
| 1592 | } |
| 1593 | } |
| 1594 | |
| 1595 | /// #6511: the Rust loop kept the exhausted sub-RLM's last answer, but the |
| 1596 | /// Python helpers returned only the error marker whenever `error` was |
| 1597 | /// set, so the calling code never saw the text. |
| 1598 | #[tokio::test] |
| 1599 | async fn rlm_query_keeps_partial_answer_from_an_exhausted_child() { |
| 1600 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1601 | let round = rt |
| 1602 | .run( |
| 1603 | "print(rlm_query('go'))\nfor r in rlm_query_batched(['a', 'b'], dependency_mode='independent'):\n print(r)", |
| 1604 | Some(&ExhaustedRlmBridge), |
| 1605 | ) |
| 1606 | .await |
| 1607 | .expect("execute"); |
| 1608 | // Windows Python prints CRLF; compare line content, not line endings. |
| 1609 | let out = &round.stdout.replace("\r\n", "\n"); |
| 1610 | assert!( |
| 1611 | out.contains("partial answer 0\n[rlm_query incomplete: RLM loop exhausted"), |
| 1612 | "{out}" |
| 1613 | ); |
| 1614 | // Batched: a child with text keeps it; one with none keeps the old marker. |
| 1615 | assert!(out.matches("partial answer 0").count() == 2, "{out}"); |
| 1616 | assert!(out.contains("[child err: RLM loop exhausted"), "{out}"); |
| 1617 | rt.shutdown().await; |
| 1618 | } |
| 1619 | |
| 1620 | #[tokio::test] |
| 1621 | async fn rpc_dispatcher_round_trips_sub_query_alias() { |
| 1622 | let bridge = StubBridge::new(); |
| 1623 | let calls = Arc::clone(&bridge.calls); |
| 1624 | |
| 1625 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1626 | let round = rt |
| 1627 | .run("print(sub_query('hello from sub'))", Some(&bridge)) |
| 1628 | .await |
| 1629 | .expect("execute"); |
| 1630 | assert!( |
| 1631 | round.stdout.contains("stub#0: hello from sub"), |
| 1632 | "stdout: {:?}", |
| 1633 | round.stdout |
| 1634 | ); |
| 1635 | assert_eq!(round.rpc_count, 1); |
| 1636 | |
| 1637 | let recorded = calls.lock().await; |
| 1638 | assert_eq!(recorded.len(), 1); |
| 1639 | match &recorded[0] { |
| 1640 | RpcRequest::Llm { prompt, .. } => assert_eq!(prompt, "hello from sub"), |
| 1641 | other => panic!("expected Llm request, got {other:?}"), |
| 1642 | } |
| 1643 | drop(recorded); |
| 1644 | rt.shutdown().await; |
| 1645 | } |
| 1646 | |
| 1647 | #[tokio::test] |
| 1648 | async fn rpc_dispatcher_round_trips_batch() { |
| 1649 | let bridge = StubBridge::new(); |
| 1650 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1651 | let round = rt |
| 1652 | .run( |
| 1653 | "outs = llm_query_batched(['a','b','c'], dependency_mode='independent', safety_note='same independent classification')\n\ |
| 1654 | print('|'.join(outs))", |
| 1655 | Some(&bridge), |
| 1656 | ) |
| 1657 | .await |
| 1658 | .expect("execute"); |
| 1659 | assert!(round.stdout.contains("stub#0.0: a")); |
| 1660 | assert!(round.stdout.contains("stub#0.1: b")); |
| 1661 | assert!(round.stdout.contains("stub#0.2: c")); |
| 1662 | assert_eq!(round.rpc_count, 1); |
| 1663 | rt.shutdown().await; |
| 1664 | } |
| 1665 | |
| 1666 | #[tokio::test] |
| 1667 | async fn batched_helpers_require_independence_declaration() { |
| 1668 | let bridge = StubBridge::new(); |
| 1669 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1670 | let round = rt |
| 1671 | .run( |
| 1672 | "outs = sub_query_batch('summarize', [{'text': 'a'}, {'text': 'b'}])\n\ |
| 1673 | print(outs[0])", |
| 1674 | Some(&bridge), |
| 1675 | ) |
| 1676 | .await |
| 1677 | .expect("execute"); |
| 1678 | |
| 1679 | assert!( |
| 1680 | round.stdout.contains("dependency_mode='independent'"), |
| 1681 | "{}", |
| 1682 | round.stdout |
| 1683 | ); |
| 1684 | assert_eq!(round.rpc_count, 0); |
| 1685 | rt.shutdown().await; |
| 1686 | } |
| 1687 | |
| 1688 | #[tokio::test] |
| 1689 | async fn dependent_batch_mode_points_to_sequence_helper() { |
| 1690 | let bridge = StubBridge::new(); |
| 1691 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1692 | let round = rt |
| 1693 | .run( |
| 1694 | "outs = llm_query_batched(['migrate A', 'migrate B'], dependency_mode='sequential')\n\ |
| 1695 | print(outs[0])", |
| 1696 | Some(&bridge), |
| 1697 | ) |
| 1698 | .await |
| 1699 | .expect("execute"); |
| 1700 | |
| 1701 | assert!( |
| 1702 | round.stdout.contains("sub_query_sequence"), |
| 1703 | "{}", |
| 1704 | round.stdout |
| 1705 | ); |
| 1706 | assert_eq!(round.rpc_count, 0); |
| 1707 | rt.shutdown().await; |
| 1708 | } |
| 1709 | |
| 1710 | #[tokio::test] |
| 1711 | async fn sub_query_sequence_feeds_prior_result_into_next_prompt() { |
| 1712 | let bridge = StubBridge::new(); |
| 1713 | let calls = Arc::clone(&bridge.calls); |
| 1714 | |
| 1715 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1716 | let round = rt |
| 1717 | .run( |
| 1718 | "outs = sub_query_sequence('process this step', [{'text': 'A'}, {'text': 'B'}])\n\ |
| 1719 | print(len(outs))", |
| 1720 | Some(&bridge), |
| 1721 | ) |
| 1722 | .await |
| 1723 | .expect("execute"); |
| 1724 | |
| 1725 | assert!(round.stdout.contains("2"), "{}", round.stdout); |
| 1726 | assert_eq!(round.rpc_count, 2); |
| 1727 | |
| 1728 | let recorded = calls.lock().await; |
| 1729 | assert_eq!(recorded.len(), 2); |
| 1730 | let second_prompt = match &recorded[1] { |
| 1731 | RpcRequest::Llm { prompt, .. } => prompt, |
| 1732 | other => panic!("expected second Llm request, got {other:?}"), |
| 1733 | }; |
| 1734 | assert!(second_prompt.contains("--- dependency_state step 1/2 ---")); |
| 1735 | assert!(second_prompt.contains("stub#0: process this step")); |
| 1736 | drop(recorded); |
| 1737 | rt.shutdown().await; |
| 1738 | } |
| 1739 | |
| 1740 | #[tokio::test] |
| 1741 | async fn no_dispatcher_returns_unavailable_sentinel() { |
| 1742 | let mut rt = PythonRuntime::new().await.expect("spawn"); |
| 1743 | let round = rt.execute("print(llm_query('hi'))").await.expect("execute"); |
| 1744 | assert!( |
| 1745 | round.stdout.contains("[llm_query error:") || round.stdout.contains("no LLM bridge"), |
| 1746 | "stdout: {:?}", |
| 1747 | round.stdout |
| 1748 | ); |
| 1749 | rt.shutdown().await; |
| 1750 | } |
| 1751 | |
| 1752 | #[test] |
| 1753 | fn truncate_keeps_short_unchanged() { |
| 1754 | assert_eq!(truncate_stdout("hello", 100), "hello"); |
| 1755 | } |
| 1756 | |
| 1757 | #[test] |
| 1758 | fn truncate_clips_long() { |
| 1759 | let long = "a".repeat(10_000); |
| 1760 | let out = truncate_stdout(&long, 1024); |
| 1761 | assert!(out.len() < 1500); |
| 1762 | assert!(out.contains("truncated")); |
| 1763 | } |
| 1764 | } |
| 1765 |