1use 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
30pub(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
47const CHILD_TIMEOUT_SECS: u64 = 120;
49pub const MAX_BATCH: usize = 16;
51
52pub(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
109pub struct RlmBridge {
111 client: Arc<dyn RlmLlmClient>,
112 child_model: String,
113 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 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 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 return self.dispatch_llm(prompt, None, None, None).await;
257 }
258
259 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 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}