Skip to main content

codewhale_tui/rlm/
bridge.rs

1//! RPC bridge that services `llm_query` / `rlm_query` calls coming back
2//! from the long-lived Python REPL during an RLM turn.
3//!
4//! This is the spiritual successor to the HTTP sidecar from earlier
5//! versions — except instead of binding a localhost port and routing
6//! through `urllib`, requests come in through stdin/stdout and we just
7//! call the LLM client directly here in Rust.
8//!
9//! The bridge tracks cumulative token usage and the recursion budget. For
10//! `Rlm` / `RlmBatch` requests it recursively calls `run_rlm_turn_inner`
11//! at depth-1; the future-type cycle (bridge → run_rlm_turn_inner →
12//! bridge) is broken by `run_rlm_turn_inner` returning a boxed dyn future.
13
14use std::sync::Arc;
15use std::time::Duration;
16use std::{future::Future, pin::Pin};
17
18use anyhow::Result;
19use futures_util::future::join_all;
20use tokio::sync::Mutex;
21
22use crate::llm_client::LlmClient;
23use crate::models::{
24    ContentBlock, Message, MessageRequest, MessageResponse, SystemPrompt, Usage,
25    is_incomplete_stop_reason, stop_reason_detail,
26};
27use crate::repl::runtime::{BatchResp, RpcDispatcher, RpcRequest, RpcResponse, SingleResp};
28use crate::utils::spawn_supervised;
29
30/// Object-safe runtime-model adapter for a working kernel.
31///
32/// The normal turn loop owns a `SharedModelClient`, while the original RLM
33/// bridge predates that boundary and accepts the concrete [`LlmClient`] trait.
34/// Keeping this small adapter here means a persistent kernel follows exactly
35/// the selected model route (including custom providers) without teaching the
36/// kernel about provider transports or falling back to a side channel.
37pub(crate) struct ModelClientRlmAdapter {
38    client: crate::core::model_client::SharedModelClient,
39}
40
41impl ModelClientRlmAdapter {
42    pub(crate) fn new(client: crate::core::model_client::SharedModelClient) -> Self {
43        Self { client }
44    }
45}
46
47/// Per-child completion timeout — same as the previous sidecar default.
48const CHILD_TIMEOUT_SECS: u64 = 120;
49/// Hard cap on prompts per batch RPC.
50pub const MAX_BATCH: usize = 16;
51
52/// Object-safe slice of the LLM client interface that the RLM bridge needs.
53///
54/// `LlmClient` itself uses native async trait methods, which are not dyn-safe.
55/// The bridge only needs non-streaming completions, so this boxed-future shim
56/// gives tests a clean mock seam without changing the wider provider trait.
57pub(crate) trait RlmLlmClient: Send + Sync {
58    fn effective_route_envelope(
59        &self,
60        requested_model: &str,
61        dispatched_at: chrono::DateTime<chrono::Utc>,
62    ) -> crate::cost_status::EffectiveRouteEnvelope;
63
64    fn create_message_boxed(
65        &self,
66        request: MessageRequest,
67    ) -> Pin<Box<dyn Future<Output = Result<MessageResponse>> + Send + '_>>;
68}
69
70impl RlmLlmClient for ModelClientRlmAdapter {
71    fn effective_route_envelope(
72        &self,
73        requested_model: &str,
74        dispatched_at: chrono::DateTime<chrono::Utc>,
75    ) -> crate::cost_status::EffectiveRouteEnvelope {
76        self.client
77            .effective_route_envelope(requested_model, dispatched_at)
78    }
79
80    fn create_message_boxed(
81        &self,
82        request: MessageRequest,
83    ) -> Pin<Box<dyn Future<Output = Result<MessageResponse>> + Send + '_>> {
84        let client = Arc::clone(&self.client);
85        Box::pin(async move { client.create_message(request).await })
86    }
87}
88
89impl<T> RlmLlmClient for T
90where
91    T: LlmClient + Send + Sync,
92{
93    fn effective_route_envelope(
94        &self,
95        requested_model: &str,
96        dispatched_at: chrono::DateTime<chrono::Utc>,
97    ) -> crate::cost_status::EffectiveRouteEnvelope {
98        LlmClient::effective_route_envelope(self, requested_model, dispatched_at)
99    }
100
101    fn create_message_boxed(
102        &self,
103        request: MessageRequest,
104    ) -> Pin<Box<dyn Future<Output = Result<MessageResponse>> + Send + '_>> {
105        Box::pin(self.create_message(request))
106    }
107}
108
109/// State shared with the bridge across all RPC calls in one turn.
110pub struct RlmBridge {
111    client: Arc<dyn RlmLlmClient>,
112    child_model: String,
113    /// Recursion budget remaining for `Rlm` / `RlmBatch` requests. When
114    /// zero, those requests fall back to plain `Llm` completions.
115    depth_remaining: u32,
116    usage: Arc<Mutex<Usage>>,
117}
118
119impl RlmBridge {
120    pub(crate) fn new(
121        client: Arc<dyn RlmLlmClient>,
122        child_model: String,
123        depth_remaining: u32,
124    ) -> Self {
125        Self {
126            client,
127            child_model,
128            depth_remaining,
129            usage: Arc::new(Mutex::new(Usage::default())),
130        }
131    }
132
133    pub fn usage_handle(&self) -> Arc<Mutex<Usage>> {
134        Arc::clone(&self.usage)
135    }
136
137    async fn dispatch_llm(
138        &self,
139        prompt: String,
140        _model: Option<String>,
141        max_tokens: Option<u32>,
142        system: Option<String>,
143    ) -> SingleResp {
144        let request_route = self
145            .client
146            .effective_route_envelope(&self.child_model, chrono::Utc::now());
147        let route_max_tokens = crate::route_budget::effective_max_output_tokens_for_route(
148            request_route.provider,
149            &request_route.model,
150            None,
151        );
152        let request = MessageRequest {
153            // The Python helper accepts `model=` for older snippets, but it is
154            // intentionally not authoritative. RLM child calls are pinned to
155            // the tool's configured child model so model-generated Python
156            // cannot silently upgrade cheap fanout work to an expensive model.
157            model: self.child_model.clone(),
158            messages: vec![Message {
159                role: "user".to_string(),
160                content: vec![ContentBlock::Text {
161                    text: prompt,
162                    cache_control: None,
163                }],
164            }],
165            // An explicit RLM helper bound remains authoritative, but the
166            // default is the selected route's ordinary allowance rather than
167            // a hidden 4K ceiling.
168            max_tokens: max_tokens.map_or(route_max_tokens, |limit| limit.min(route_max_tokens)),
169            system: system.map(SystemPrompt::Text),
170            tools: None,
171            tool_choice: None,
172            metadata: None,
173            thinking: None,
174            reasoning_effort: None,
175            stream: Some(false),
176            temperature: None,
177            top_p: None,
178        };
179
180        let fut = self.client.create_message_boxed(request);
181        let response =
182            match tokio::time::timeout(Duration::from_secs(CHILD_TIMEOUT_SECS), fut).await {
183                Ok(Ok(r)) => r,
184                Ok(Err(e)) => {
185                    return SingleResp {
186                        text: String::new(),
187                        error: Some(format!("llm_query failed: {e}")),
188                    };
189                }
190                Err(_) => {
191                    return SingleResp {
192                        text: String::new(),
193                        error: Some(format!("llm_query timed out after {CHILD_TIMEOUT_SECS}s")),
194                    };
195                }
196            };
197
198        {
199            let mut u = self.usage.lock().await;
200            super::add_usage_with_prompt_cache(&mut u, &response.usage);
201        }
202
203        if is_incomplete_stop_reason(response.stop_reason.as_deref()) {
204            return SingleResp {
205                text: String::new(),
206                error: Some(format!(
207                    "llm_query response incomplete: provider stop reason `{}`; partial output was not accepted.",
208                    stop_reason_detail(response.stop_reason.as_deref())
209                )),
210            };
211        }
212
213        let text = response
214            .content
215            .iter()
216            .filter_map(|b| match b {
217                ContentBlock::Text { text, .. } => Some(text.as_str()),
218                _ => None,
219            })
220            .collect::<Vec<_>>()
221            .join("\n");
222
223        SingleResp { text, error: None }
224    }
225
226    async fn dispatch_llm_batch(
227        &self,
228        prompts: Vec<String>,
229        _model: Option<String>,
230        dependency_mode: Option<String>,
231    ) -> BatchResp {
232        if let Some(resp) = batch_guard(prompts.len(), dependency_mode.as_deref()) {
233            return resp;
234        }
235
236        let model = Arc::new(self.child_model.clone());
237
238        let futures = prompts.into_iter().map(|prompt| {
239            let model = Arc::clone(&model);
240            async move {
241                self.dispatch_llm((*prompt).to_string(), Some((*model).clone()), None, None)
242                    .await
243            }
244        });
245
246        BatchResp {
247            results: join_all(futures).await,
248        }
249    }
250
251    async fn dispatch_rlm(&self, prompt: String, _model: Option<String>) -> SingleResp {
252        if self.depth_remaining == 0 {
253            // Budget exhausted — fall back to a one-shot child completion
254            // rather than returning an error. Matches the paper's behaviour
255            // ("sub_RLM gracefully degrades to llm_query at depth=0").
256            return self.dispatch_llm(prompt, None, None, None).await;
257        }
258
259        // Build a drain channel to absorb status events from the nested
260        // turn (we don't surface them; this dispatch is invisible to the
261        // outer agent stream).
262        let (tx, mut rx) = tokio::sync::mpsc::channel(64);
263        let drain = spawn_supervised(
264            "rlm-bridge-drain",
265            std::panic::Location::caller(),
266            async move { while rx.recv().await.is_some() {} },
267        );
268
269        let child_model = self.child_model.clone();
270
271        // Recursive call. The dyn-erasure on `run_rlm_turn_inner` breaks
272        // the `bridge → turn → bridge` opaque-future cycle.
273        let result = super::turn::run_rlm_turn_inner(
274            Arc::clone(&self.client),
275            child_model.clone(),
276            prompt,
277            None,
278            child_model,
279            tx,
280            self.depth_remaining.saturating_sub(1),
281        )
282        .await;
283
284        drain.abort();
285
286        {
287            let mut u = self.usage.lock().await;
288            super::add_usage_with_prompt_cache(&mut u, &result.usage);
289        }
290
291        SingleResp {
292            text: result.answer,
293            error: result.error,
294        }
295    }
296
297    async fn dispatch_rlm_batch(
298        &self,
299        prompts: Vec<String>,
300        _model: Option<String>,
301        dependency_mode: Option<String>,
302    ) -> BatchResp {
303        if let Some(resp) = batch_guard(prompts.len(), dependency_mode.as_deref()) {
304            return resp;
305        }
306
307        let futures = prompts
308            .into_iter()
309            .map(|p| async move { self.dispatch_rlm(p, None).await });
310        BatchResp {
311            results: join_all(futures).await,
312        }
313    }
314}
315
316fn batch_guard(prompt_count: usize, dependency_mode: Option<&str>) -> Option<BatchResp> {
317    if prompt_count == 0 {
318        return Some(BatchResp { results: vec![] });
319    }
320    if prompt_count > MAX_BATCH {
321        return Some(BatchResp {
322            results: (0..prompt_count)
323                .map(|_| SingleResp {
324                    text: String::new(),
325                    error: Some(format!("batch too large: {prompt_count} > {MAX_BATCH}")),
326                })
327                .collect(),
328        });
329    }
330    let mode = dependency_mode
331        .unwrap_or_default()
332        .trim()
333        .to_ascii_lowercase()
334        .replace(['-', ' '], "_");
335    if !matches!(
336        mode.as_str(),
337        "independent" | "parallel_safe" | "map_reduce"
338    ) {
339        return Some(BatchResp {
340            results: (0..prompt_count)
341                .map(|_| SingleResp {
342                    text: String::new(),
343                    error: Some(
344                        "batch requires dependency_mode='independent'; use sub_query_sequence or sequential sub_query calls for dependent work"
345                            .to_string(),
346                    ),
347                })
348                .collect(),
349        });
350    }
351    None
352}
353
354impl RpcDispatcher for RlmBridge {
355    fn dispatch<'a>(
356        &'a self,
357        req: RpcRequest,
358    ) -> std::pin::Pin<Box<dyn std::future::Future<Output = RpcResponse> + Send + 'a>> {
359        Box::pin(async move {
360            match req {
361                RpcRequest::Llm {
362                    prompt,
363                    model,
364                    max_tokens,
365                    system,
366                } => {
367                    RpcResponse::Single(self.dispatch_llm(prompt, model, max_tokens, system).await)
368                }
369                RpcRequest::LlmBatch {
370                    prompts,
371                    model,
372                    dependency_mode,
373                    safety_note: _,
374                } => RpcResponse::Batch(
375                    self.dispatch_llm_batch(prompts, model, dependency_mode)
376                        .await,
377                ),
378                RpcRequest::Rlm { prompt, model } => {
379                    RpcResponse::Single(self.dispatch_rlm(prompt, model).await)
380                }
381                RpcRequest::RlmBatch {
382                    prompts,
383                    model,
384                    dependency_mode,
385                    safety_note: _,
386                } => RpcResponse::Batch(
387                    self.dispatch_rlm_batch(prompts, model, dependency_mode)
388                        .await,
389                ),
390            }
391        })
392    }
393}
394
395#[cfg(test)]
396mod tests {
397    use super::*;
398    use crate::llm_client::mock::MockLlmClient;
399
400    fn mock_response_with_usage(text: &str, usage: Usage) -> MessageResponse {
401        MessageResponse {
402            id: "mock_msg".to_string(),
403            r#type: "message".to_string(),
404            role: "assistant".to_string(),
405            content: vec![ContentBlock::Text {
406                text: text.to_string(),
407                cache_control: None,
408            }],
409            model: "mock-model".to_string(),
410            stop_reason: Some("end_turn".to_string()),
411            stop_sequence: None,
412            container: None,
413            usage,
414        }
415    }
416
417    fn mock_response(text: &str, input_tokens: u32, output_tokens: u32) -> MessageResponse {
418        mock_response_with_usage(
419            text,
420            Usage {
421                input_tokens,
422                output_tokens,
423                ..Usage::default()
424            },
425        )
426    }
427
428    fn bridge_for(mock: Arc<MockLlmClient>, depth_remaining: u32) -> RlmBridge {
429        let client: Arc<dyn RlmLlmClient> = mock;
430        RlmBridge::new(client, "child-model".to_string(), depth_remaining)
431    }
432
433    #[test]
434    fn batch_guard_allows_non_empty_batches_at_the_cap() {
435        assert!(batch_guard(MAX_BATCH, Some("independent")).is_none());
436    }
437
438    #[test]
439    fn batch_guard_returns_empty_response_for_empty_batches() {
440        let response = batch_guard(0, None).expect("empty batch should be handled");
441        assert!(response.results.is_empty());
442    }
443
444    #[test]
445    fn batch_guard_returns_one_error_per_oversized_prompt() {
446        let response = batch_guard(MAX_BATCH + 2, Some("independent"))
447            .expect("oversized batch should be handled");
448        assert_eq!(response.results.len(), MAX_BATCH + 2);
449        assert!(response.results.iter().all(|result| {
450            result.text.is_empty()
451                && result
452                    .error
453                    .as_deref()
454                    .is_some_and(|err| err.contains("batch too large"))
455        }));
456    }
457
458    #[test]
459    fn batch_guard_requires_explicit_independence_for_parallel_work() {
460        let response = batch_guard(2, None).expect("missing dependency mode should be handled");
461        assert_eq!(response.results.len(), 2);
462        assert!(response.results.iter().all(|result| {
463            result.text.is_empty()
464                && result
465                    .error
466                    .as_deref()
467                    .is_some_and(|err| err.contains("dependency_mode='independent'"))
468        }));
469
470        let response = batch_guard(2, Some("sequential"))
471            .expect("dependent dependency mode should be handled");
472        assert!(response.results.iter().all(|result| {
473            result
474                .error
475                .as_deref()
476                .is_some_and(|err| err.contains("sub_query_sequence"))
477        }));
478    }
479
480    #[tokio::test]
481    async fn llm_dispatch_pins_configured_child_model() {
482        let mock = Arc::new(MockLlmClient::new(Vec::new()));
483        mock.push_message_response(mock_response("child answer", 7, 11));
484        let bridge = bridge_for(Arc::clone(&mock), 1);
485
486        let response = bridge
487            .dispatch(RpcRequest::Llm {
488                prompt: "child prompt".to_string(),
489                model: Some("override-model".to_string()),
490                max_tokens: Some(123),
491                system: Some("child system".to_string()),
492            })
493            .await;
494
495        match response {
496            RpcResponse::Single(single) => {
497                assert_eq!(single.text, "child answer");
498                assert!(single.error.is_none());
499            }
500            other => panic!("expected single response, got {other:?}"),
501        }
502
503        let captured = mock.captured_requests();
504        assert_eq!(captured.len(), 1);
505        assert_eq!(captured[0].model, "child-model");
506        assert_eq!(captured[0].max_tokens, 123);
507        assert_eq!(
508            captured[0].system,
509            Some(SystemPrompt::Text("child system".to_string()))
510        );
511
512        let usage = bridge.usage.lock().await;
513        assert_eq!(usage.input_tokens, 7);
514        assert_eq!(usage.output_tokens, 11);
515    }
516
517    #[tokio::test]
518    async fn llm_dispatch_preserves_prompt_cache_usage() {
519        let mock = Arc::new(MockLlmClient::new(Vec::new()));
520        mock.push_message_response(mock_response_with_usage(
521            "cached child answer",
522            Usage {
523                input_tokens: 1000,
524                output_tokens: 100,
525                prompt_cache_hit_tokens: Some(800),
526                prompt_cache_miss_tokens: Some(200),
527                ..Usage::default()
528            },
529        ));
530        let bridge = bridge_for(Arc::clone(&mock), 1);
531
532        let response = bridge
533            .dispatch(RpcRequest::Llm {
534                prompt: "child prompt".to_string(),
535                model: None,
536                max_tokens: None,
537                system: None,
538            })
539            .await;
540
541        match response {
542            RpcResponse::Single(single) => {
543                assert_eq!(single.text, "cached child answer");
544                assert!(single.error.is_none());
545            }
546            other => panic!("expected single response, got {other:?}"),
547        }
548
549        let usage = bridge.usage.lock().await;
550        assert_eq!(usage.input_tokens, 1000);
551        assert_eq!(usage.output_tokens, 100);
552        assert_eq!(usage.prompt_cache_hit_tokens, Some(800));
553        assert_eq!(usage.prompt_cache_miss_tokens, Some(200));
554    }
555
556    #[tokio::test]
557    async fn llm_dispatch_rejects_max_tokens_partial_output_after_charging_usage() {
558        let mock = Arc::new(MockLlmClient::new(Vec::new()));
559        let usage = Usage {
560            input_tokens: 23,
561            output_tokens: 4096,
562            reasoning_tokens: Some(4000),
563            ..Usage::default()
564        };
565        let mut response = mock_response_with_usage(
566            "FINAL('partial answer')\n```repl\nFINAL('also partial')\n```",
567            usage.clone(),
568        );
569        response.stop_reason = Some("max_tokens".to_string());
570        mock.push_message_response(response);
571        let bridge = bridge_for(Arc::clone(&mock), 1);
572
573        let response = bridge
574            .dispatch(RpcRequest::Llm {
575                prompt: "child prompt".to_string(),
576                model: None,
577                max_tokens: None,
578                system: None,
579            })
580            .await;
581
582        match response {
583            RpcResponse::Single(single) => {
584                assert!(
585                    single.text.is_empty(),
586                    "partial output must not be accepted"
587                );
588                let error = single.error.expect("truncation must surface as an error");
589                assert!(error.contains("incomplete"), "{error}");
590                assert!(error.contains("max_tokens"), "{error}");
591            }
592            other => panic!("expected single response, got {other:?}"),
593        }
594
595        assert_eq!(*bridge.usage.lock().await, usage);
596        assert_eq!(mock.call_count(), 1, "truncation must not retry");
597    }
598
599    #[tokio::test]
600    async fn llm_batch_dispatch_pins_configured_child_model() {
601        let mock = Arc::new(MockLlmClient::new(Vec::new()));
602        mock.push_message_response(mock_response("one", 1, 2));
603        mock.push_message_response(mock_response("two", 3, 4));
604        mock.push_message_response(mock_response("three", 5, 6));
605        let bridge = bridge_for(Arc::clone(&mock), 1);
606
607        let response = bridge
608            .dispatch(RpcRequest::LlmBatch {
609                prompts: vec!["a".to_string(), "b".to_string(), "c".to_string()],
610                model: Some("batch-model".to_string()),
611                dependency_mode: Some("independent".to_string()),
612                safety_note: Some("test prompts are independent".to_string()),
613            })
614            .await;
615
616        match response {
617            RpcResponse::Batch(batch) => {
618                let texts: Vec<_> = batch
619                    .results
620                    .iter()
621                    .map(|result| result.text.as_str())
622                    .collect();
623                assert_eq!(texts, ["one", "two", "three"]);
624                assert!(batch.results.iter().all(|result| result.error.is_none()));
625            }
626            other => panic!("expected batch response, got {other:?}"),
627        }
628
629        let captured = mock.captured_requests();
630        assert_eq!(captured.len(), 3);
631        assert!(
632            captured
633                .iter()
634                .all(|request| request.model == "child-model")
635        );
636
637        let usage = bridge.usage.lock().await;
638        assert_eq!(usage.input_tokens, 9);
639        assert_eq!(usage.output_tokens, 12);
640    }
641
642    #[tokio::test]
643    async fn rlm_dispatch_at_depth_zero_pins_configured_child_model() {
644        let mock = Arc::new(MockLlmClient::new(Vec::new()));
645        mock.push_message_response(mock_response("fallback answer", 3, 5));
646        let bridge = bridge_for(Arc::clone(&mock), 0);
647
648        let response = bridge
649            .dispatch(RpcRequest::Rlm {
650                prompt: "nested prompt".to_string(),
651                model: Some("override-model".to_string()),
652            })
653            .await;
654
655        match response {
656            RpcResponse::Single(single) => {
657                assert_eq!(single.text, "fallback answer");
658                assert!(single.error.is_none());
659            }
660            other => panic!("expected single response, got {other:?}"),
661        }
662
663        let usage = bridge.usage.lock().await;
664        assert_eq!(usage.input_tokens, 3);
665        assert_eq!(usage.output_tokens, 5);
666
667        let captured = mock.captured_requests();
668        assert_eq!(captured.len(), 1);
669        assert_eq!(captured[0].model, "child-model");
670    }
671}