88//! consumer — it drives the stream with [`switchyard_libsy::drive`], hands routing-time calls to
99//! a [`RoutedLlmClient`], and serves the terminal routing outcome.
1010//!
11- //! libsy owns the stream mechanics; what this module adds is ordered candidate fallback and the
12- //! `libsy.client_call` span around each candidate. Each candidate exhausts its backend retry
13- //! budget before fallback advances, so the worst case is `candidates × (max_retries + 1)`
14- //! upstream attempts plus every candidate's backoff.
11+ //! libsy owns the stream mechanics; what this module adds is per-target request preparation,
12+ //! ordered candidate fallback, and the `libsy.client_call` span around each candidate. Each
13+ //! candidate exhausts its backend retry budget before fallback advances, so the worst case is
14+ //! `candidates × (max_retries + 1)` upstream attempts plus every candidate's backoff.
1515
1616use std:: collections:: HashMap ;
1717use std:: sync:: Arc ;
1818use std:: time:: Instant ;
1919
2020use http:: StatusCode ;
2121use parking_lot:: Mutex ;
22- use switchyard_libsy:: { Algorithm , CallModel , LibsyError , Result , drive} ;
22+ use switchyard_libsy:: { Algorithm , CallModel , LibsyError , Result , RoutingOutcome , drive} ;
2323use switchyard_protocol:: {
2424 LlmClientError , ModelId , Request , Response , RoutedLlmClient , RoutingFallbackReason ,
2525} ;
26+ use switchyard_translation:: prepare_request_for_target;
2627
2728use crate :: observation:: { LlmCallObservation , RunObservation , RunObserver } ;
2829use crate :: { metrics, observability} ;
@@ -84,6 +85,7 @@ pub async fn run(
8485 & algorithm_name,
8586 & outcome. request ,
8687 & models,
88+ CallPhase :: Completion ,
8789 & observe,
8890 )
8991 . await ;
@@ -103,6 +105,25 @@ pub async fn run(
103105 result. map ( |response| ( selected_model_id, response) )
104106}
105107
108+ /// Run an algorithm to a routing decision without serving its terminal completion.
109+ ///
110+ /// Routing-time calls are served normally. The returned request is prepared only for the selected
111+ /// target; use [`run`] to execute selected and fallback candidates.
112+ pub async fn decide (
113+ algorithm : Arc < dyn Algorithm > ,
114+ clients : ClientRouter ,
115+ request : Request ,
116+ ) -> Result < RoutingOutcome > {
117+ let routing_clients = clients. clone ( ) ;
118+ let mut outcome = drive ( algorithm, request, move |call| {
119+ serve ( routing_clients. clone ( ) , call, None )
120+ } )
121+ . await ?;
122+ outcome. request =
123+ clients. prepare_completion_request ( outcome. request , & outcome. selected_model_id ) ;
124+ Ok ( outcome)
125+ }
126+
106127/// Emits completed routing calls after the outcome reveals whether one response became the answer.
107128fn emit_routing_observations (
108129 observer : & Option < RunObserver > ,
@@ -144,22 +165,32 @@ async fn serve(
144165 & call. algorithm ,
145166 & call. request ,
146167 & call. models ,
168+ CallPhase :: Routing ,
147169 & observe,
148170 )
149171 . await ;
150172 call. respond ( result)
151173}
152174
175+ enum CallPhase {
176+ Routing ,
177+ Completion ,
178+ }
179+
153180/// Try candidates in order until one succeeds or a failure stops fallback.
154181async fn call_first_available (
155182 clients : & ClientRouter ,
156183 algorithm : & str ,
157184 request : & Request ,
158185 models : & [ ModelId ] ,
186+ phase : CallPhase ,
159187 observe : & ( dyn Fn ( LlmCallObservation ) + Send + Sync ) ,
160188) -> Result < Response > {
161189 for ( index, target) in models. iter ( ) . enumerate ( ) {
162- let request = request_for ( request, target) ;
190+ let request = match phase {
191+ CallPhase :: Routing => clients. prepare_routing_request ( request. clone ( ) , target) ,
192+ CallPhase :: Completion => clients. prepare_completion_request ( request. clone ( ) , target) ,
193+ } ;
163194 match call_one (
164195 clients,
165196 target,
@@ -298,13 +329,6 @@ fn fallback_reason(error: &LibsyError) -> Option<RoutingFallbackReason> {
298329 }
299330}
300331
301- /// Clone a request and stamp the candidate model that should receive it.
302- fn request_for ( request : & Request , target : & ModelId ) -> Request {
303- let mut request = request. clone ( ) ;
304- request. llm_request . model = Some ( target. to_string ( ) ) ;
305- request
306- }
307-
308332/// Resolves a routed call's selected model to the client that serves it.
309333///
310334/// An algorithm routes among named targets; which provider each target lives on is the
@@ -315,7 +339,13 @@ fn request_for(request: &Request, target: &ModelId) -> Request {
315339/// Cloning is cheap — the mapping is shared, so one router can serve every request.
316340#[ derive( Clone ) ]
317341pub struct ClientRouter {
318- routing : Arc < Routing > ,
342+ inner : Arc < ClientRouting > ,
343+ }
344+
345+ struct ClientRouting {
346+ routing : Routing ,
347+ target_prompts : HashMap < ModelId , String > ,
348+ routing_answer_target : Option < ModelId > ,
319349}
320350
321351enum Routing {
@@ -328,8 +358,25 @@ enum Routing {
328358impl ClientRouter {
329359 /// Build a router over `model name -> client`, for targets spread across providers.
330360 pub fn new ( by_model : HashMap < ModelId , Arc < dyn RoutedLlmClient > > ) -> Self {
361+ Self :: new_with_target_prompts ( by_model, HashMap :: new ( ) , None )
362+ }
363+
364+ /// Build a router with target prompts used by [`run`] and [`decide`].
365+ ///
366+ /// `routing_answer_target` identifies a routing-time call whose response may become the
367+ /// answer. Pass `None` when routing only selects a later completion target. Prompt keys and
368+ /// `routing_answer_target` use the resolved model IDs stored in `by_model`.
369+ pub fn new_with_target_prompts (
370+ by_model : HashMap < ModelId , Arc < dyn RoutedLlmClient > > ,
371+ target_prompts : HashMap < ModelId , String > ,
372+ routing_answer_target : Option < ModelId > ,
373+ ) -> Self {
331374 Self {
332- routing : Arc :: new ( Routing :: ByModel ( by_model) ) ,
375+ inner : Arc :: new ( ClientRouting {
376+ routing : Routing :: ByModel ( by_model) ,
377+ target_prompts,
378+ routing_answer_target,
379+ } ) ,
333380 }
334381 }
335382
@@ -340,7 +387,11 @@ impl ClientRouter {
340387 /// only duplicate that.
341388 pub fn single ( client : Arc < dyn RoutedLlmClient > ) -> Self {
342389 Self {
343- routing : Arc :: new ( Routing :: Single ( client) ) ,
390+ inner : Arc :: new ( ClientRouting {
391+ routing : Routing :: Single ( client) ,
392+ target_prompts : HashMap :: new ( ) ,
393+ routing_answer_target : None ,
394+ } ) ,
344395 }
345396 }
346397
@@ -352,7 +403,7 @@ impl ClientRouter {
352403 & self ,
353404 model : & ModelId ,
354405 ) -> std:: result:: Result < & Arc < dyn RoutedLlmClient > , LlmClientError > {
355- match self . routing . as_ref ( ) {
406+ match & self . inner . routing {
356407 Routing :: Single ( client) => Ok ( client) ,
357408 Routing :: ByModel ( by_model) => {
358409 by_model
@@ -363,6 +414,24 @@ impl ClientRouter {
363414 }
364415 }
365416 }
417+
418+ /// Prepare a completion candidate with its configured target prompt.
419+ fn prepare_completion_request ( & self , mut request : Request , target : & ModelId ) -> Request {
420+ let prompt = self . inner . target_prompts . get ( target) . map ( String :: as_str) ;
421+ prepare_request_for_target ( & mut request. llm_request , target, prompt) ;
422+ request
423+ }
424+
425+ /// Prepare a routing call, adding a target prompt only when it generates a candidate answer.
426+ fn prepare_routing_request ( & self , mut request : Request , target : & ModelId ) -> Request {
427+ let prompt = if self . inner . routing_answer_target . as_ref ( ) == Some ( target) {
428+ self . inner . target_prompts . get ( target) . map ( String :: as_str)
429+ } else {
430+ None
431+ } ;
432+ prepare_request_for_target ( & mut request. llm_request , target, prompt) ;
433+ request
434+ }
366435}
367436
368437impl FromIterator < ( ModelId , Arc < dyn RoutedLlmClient > ) > for ClientRouter {
@@ -381,8 +450,8 @@ mod tests {
381450 use http:: StatusCode ;
382451 use switchyard_libsy:: { Driver , RoutingOutcome } ;
383452 use switchyard_protocol:: {
384- LlmResponse , LlmResponseChunk , LlmResponseStreamEvent , completion_text, text_request ,
385- text_response,
453+ ContentBlock , LlmResponse , LlmResponseChunk , LlmResponseStreamEvent , completion_text,
454+ text_request , text_response,
386455 } ;
387456 use wiremock:: matchers:: method;
388457 use wiremock:: { Mock , MockServer , ResponseTemplate } ;
@@ -449,6 +518,7 @@ mod tests {
449518
450519 struct CandidateClient {
451520 calls : Mutex < Vec < ModelId > > ,
521+ requests : Mutex < Vec < Request > > ,
452522 first : FirstOutcome ,
453523 }
454524
@@ -457,6 +527,7 @@ mod tests {
457527 async fn call ( & self , request : Request ) -> std:: result:: Result < Response , LlmClientError > {
458528 let model = request. model_id ( ) . unwrap_or_default ( ) ;
459529 self . calls . lock ( ) . push ( model. clone ( ) ) ;
530+ self . requests . lock ( ) . push ( request) ;
460531 if model == "weak" {
461532 return match self . first {
462533 FirstOutcome :: ContextWindow => Err ( LlmClientError :: ContextWindowExceeded {
@@ -516,11 +587,25 @@ mod tests {
516587 }
517588 }
518589
590+ fn instruction_text ( request : & Request ) -> Vec < & str > {
591+ request
592+ . llm_request
593+ . instructions
594+ . iter ( )
595+ . flat_map ( |instruction| & instruction. content )
596+ . filter_map ( |block| match block {
597+ ContentBlock :: Text { text } => Some ( text. as_str ( ) ) ,
598+ _ => None ,
599+ } )
600+ . collect ( )
601+ }
602+
519603 async fn run_candidates (
520604 first : FirstOutcome ,
521605 ) -> ( Arc < CandidateClient > , Result < ( ModelId , Response ) > ) {
522606 let client = Arc :: new ( CandidateClient {
523607 calls : Mutex :: new ( Vec :: new ( ) ) ,
608+ requests : Mutex :: new ( Vec :: new ( ) ) ,
524609 first,
525610 } ) ;
526611 let algorithm = Arc :: new ( CandidateAlgorithm {
@@ -540,6 +625,7 @@ mod tests {
540625 async fn answered_outcome_does_not_make_a_second_model_call ( ) -> Result < ( ) > {
541626 let client = Arc :: new ( CandidateClient {
542627 calls : Mutex :: new ( Vec :: new ( ) ) ,
628+ requests : Mutex :: new ( Vec :: new ( ) ) ,
543629 first : FirstOutcome :: StreamSuccess ,
544630 } ) ;
545631 let observations = Arc :: new ( Mutex :: new ( Vec :: new ( ) ) ) ;
@@ -569,6 +655,106 @@ mod tests {
569655 Ok ( ( ) )
570656 }
571657
658+ #[ tokio:: test]
659+ async fn each_fallback_candidate_receives_only_its_own_prompt ( ) -> Result < ( ) > {
660+ let client = Arc :: new ( CandidateClient {
661+ calls : Mutex :: new ( Vec :: new ( ) ) ,
662+ requests : Mutex :: new ( Vec :: new ( ) ) ,
663+ first : FirstOutcome :: ContextWindow ,
664+ } ) ;
665+ let routed_client: Arc < dyn RoutedLlmClient > = client. clone ( ) ;
666+ let clients = ClientRouter :: new_with_target_prompts (
667+ HashMap :: from ( [
668+ ( ModelId :: from ( "weak" ) , Arc :: clone ( & routed_client) ) ,
669+ ( ModelId :: from ( "strong" ) , routed_client) ,
670+ ] ) ,
671+ HashMap :: from ( [
672+ ( "weak" . into ( ) , "weak prompt" . to_string ( ) ) ,
673+ ( "strong" . into ( ) , "strong prompt" . to_string ( ) ) ,
674+ ] ) ,
675+ None ,
676+ ) ;
677+
678+ run (
679+ Arc :: new ( CandidateAlgorithm {
680+ models : vec ! [ "weak" . into( ) , "strong" . into( ) ] ,
681+ } ) ,
682+ clients,
683+ request ( ) ,
684+ None ,
685+ )
686+ . await ?;
687+
688+ let calls = client. requests . lock ( ) ;
689+ assert_eq ! ( calls. len( ) , 2 ) ;
690+ assert_eq ! ( instruction_text( & calls[ 0 ] ) , [ "weak prompt" ] ) ;
691+ assert_eq ! ( instruction_text( & calls[ 1 ] ) , [ "strong prompt" ] ) ;
692+ Ok ( ( ) )
693+ }
694+
695+ #[ tokio:: test]
696+ async fn decision_prompts_a_routing_response_target ( ) -> Result < ( ) > {
697+ let client = Arc :: new ( CandidateClient {
698+ calls : Mutex :: new ( Vec :: new ( ) ) ,
699+ requests : Mutex :: new ( Vec :: new ( ) ) ,
700+ first : FirstOutcome :: StreamSuccess ,
701+ } ) ;
702+ let routed_client: Arc < dyn RoutedLlmClient > = client. clone ( ) ;
703+ let clients = ClientRouter :: new_with_target_prompts (
704+ HashMap :: from ( [ ( ModelId :: from ( "answer" ) , routed_client) ] ) ,
705+ HashMap :: from ( [ ( "answer" . into ( ) , "answer prompt" . to_string ( ) ) ] ) ,
706+ Some ( "answer" . into ( ) ) ,
707+ ) ;
708+
709+ decide (
710+ Arc :: new ( AnsweredAlgorithm {
711+ model : "answer" . into ( ) ,
712+ } ) ,
713+ clients,
714+ request ( ) ,
715+ )
716+ . await ?;
717+
718+ let calls = client. requests . lock ( ) ;
719+ assert_eq ! ( calls. len( ) , 1 ) ;
720+ assert_eq ! ( instruction_text( & calls[ 0 ] ) , [ "answer prompt" ] ) ;
721+ Ok ( ( ) )
722+ }
723+
724+ #[ tokio:: test]
725+ async fn decision_prepares_only_the_selected_request ( ) -> Result < ( ) > {
726+ let client = Arc :: new ( CandidateClient {
727+ calls : Mutex :: new ( Vec :: new ( ) ) ,
728+ requests : Mutex :: new ( Vec :: new ( ) ) ,
729+ first : FirstOutcome :: StreamSuccess ,
730+ } ) ;
731+ let routed_client: Arc < dyn RoutedLlmClient > = client;
732+ let clients = ClientRouter :: new_with_target_prompts (
733+ HashMap :: from ( [
734+ ( ModelId :: from ( "weak" ) , Arc :: clone ( & routed_client) ) ,
735+ ( ModelId :: from ( "strong" ) , routed_client) ,
736+ ] ) ,
737+ HashMap :: from ( [
738+ ( ModelId :: from ( "weak" ) , "weak prompt" . to_string ( ) ) ,
739+ ( ModelId :: from ( "strong" ) , "strong prompt" . to_string ( ) ) ,
740+ ] ) ,
741+ None ,
742+ ) ;
743+
744+ let outcome = decide (
745+ Arc :: new ( CandidateAlgorithm {
746+ models : vec ! [ "weak" . into( ) , "strong" . into( ) ] ,
747+ } ) ,
748+ clients,
749+ request ( ) ,
750+ )
751+ . await ?;
752+
753+ assert_eq ! ( outcome. fallback_models, [ ModelId :: from( "strong" ) ] ) ;
754+ assert_eq ! ( instruction_text( & outcome. request) , [ "weak prompt" ] ) ;
755+ Ok ( ( ) )
756+ }
757+
572758 #[ test]
573759 fn fallback_only_accepts_context_and_unavailable_failures ( ) {
574760 let error = |source| LibsyError :: client_call ( "target" , source) ;
0 commit comments