@@ -62,17 +62,17 @@ pub async fn run(
6262 . and_then ( |outcome| outcome. response . as_ref ( ) )
6363 . and_then ( Response :: served_model) ;
6464 emit_routing_observations ( & observer, & routing_observations, answered_model) ;
65- let outcome = outcome?;
65+ let mut outcome = outcome?;
6666 let overhead = run_started. elapsed ( ) ;
6767 metrics:: record_routing_overhead ( & algorithm_name, overhead) ;
6868
69- let selected_model_id = outcome. selected_model_id ;
70- let ( result, answer_duration) = if let Some ( response) = outcome. response {
69+ let selected_model_id = outcome. selected_model_id . clone ( ) ;
70+ let ( result, answer_duration) = if let Some ( response) = outcome. response . take ( ) {
7171 ( Ok ( response) , None )
7272 } else {
7373 let mut models = Vec :: with_capacity ( 1 + outcome. fallback_models . len ( ) ) ;
7474 models. push ( selected_model_id. clone ( ) ) ;
75- models. extend ( outcome. fallback_models ) ;
75+ models. extend ( outcome. fallback_models . iter ( ) . cloned ( ) ) ;
7676 let answer_started = Instant :: now ( ) ;
7777 let observe = |observation| {
7878 if let Some ( observer) = & observer {
@@ -82,8 +82,8 @@ pub async fn run(
8282 let result = call_first_available (
8383 & clients,
8484 & algorithm_name,
85- & outcome. request ,
8685 & models,
86+ move |target| outcome. request_for ( target) ,
8787 & observe,
8888 )
8989 . await ;
@@ -142,8 +142,8 @@ async fn serve(
142142 let result = call_first_available (
143143 & clients,
144144 & call. algorithm ,
145- & call. request ,
146145 & call. models ,
146+ |target| call. request_for ( target) ,
147147 & observe,
148148 )
149149 . await ;
@@ -154,12 +154,12 @@ async fn serve(
154154async fn call_first_available (
155155 clients : & ClientRouter ,
156156 algorithm : & str ,
157- request : & Request ,
158157 models : & [ ModelId ] ,
158+ request_for : impl Fn ( & ModelId ) -> Result < Request > + Send ,
159159 observe : & ( dyn Fn ( LlmCallObservation ) + Send + Sync ) ,
160160) -> Result < Response > {
161161 for ( index, target) in models. iter ( ) . enumerate ( ) {
162- let request = request_for ( request , target) ;
162+ let request = request_for ( target) ? ;
163163 match call_one (
164164 clients,
165165 target,
@@ -298,13 +298,6 @@ fn fallback_reason(error: &LibsyError) -> Option<RoutingFallbackReason> {
298298 }
299299}
300300
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-
308301/// Resolves a routed call's selected model to the client that serves it.
309302///
310303/// An algorithm routes among named targets; which provider each target lives on is the
@@ -379,10 +372,10 @@ mod tests {
379372 use async_trait:: async_trait;
380373 use futures:: StreamExt ;
381374 use http:: StatusCode ;
382- use switchyard_libsy:: { Driver , RoutingOutcome } ;
375+ use switchyard_libsy:: { Driver , RoutingOutcome , TargetPrompts , with_target_prompts } ;
383376 use switchyard_protocol:: {
384- LlmResponse , LlmResponseChunk , LlmResponseStreamEvent , completion_text, text_request ,
385- text_response,
377+ ContentBlock , LlmResponse , LlmResponseChunk , LlmResponseStreamEvent , completion_text,
378+ text_request , text_response,
386379 } ;
387380 use wiremock:: matchers:: method;
388381 use wiremock:: { Mock , MockServer , ResponseTemplate } ;
@@ -429,7 +422,7 @@ mod tests {
429422 request : Request ,
430423 ) -> Result < RoutingOutcome > {
431424 let response = driver
432- . call_model ( request. clone ( ) , vec ! [ self . model. clone( ) ] )
425+ . call_answer_model ( request. clone ( ) , self . model . clone ( ) )
433426 . await ?;
434427 Ok ( RoutingOutcome :: answered (
435428 self . model . clone ( ) ,
@@ -449,6 +442,7 @@ mod tests {
449442
450443 struct CandidateClient {
451444 calls : Mutex < Vec < ModelId > > ,
445+ prompts : Mutex < Vec < Vec < String > > > ,
452446 first : FirstOutcome ,
453447 }
454448
@@ -457,6 +451,18 @@ mod tests {
457451 async fn call ( & self , request : Request ) -> std:: result:: Result < Response , LlmClientError > {
458452 let model = request. model_id ( ) . unwrap_or_default ( ) ;
459453 self . calls . lock ( ) . push ( model. clone ( ) ) ;
454+ self . prompts . lock ( ) . push (
455+ request
456+ . llm_request
457+ . instructions
458+ . iter ( )
459+ . flat_map ( |instruction| & instruction. content )
460+ . filter_map ( |block| match block {
461+ ContentBlock :: Text { text } => Some ( text. clone ( ) ) ,
462+ _ => None ,
463+ } )
464+ . collect ( ) ,
465+ ) ;
460466 if model == "weak" {
461467 return match self . first {
462468 FirstOutcome :: ContextWindow => Err ( LlmClientError :: ContextWindowExceeded {
@@ -521,11 +527,16 @@ mod tests {
521527 ) -> ( Arc < CandidateClient > , Result < ( ModelId , Response ) > ) {
522528 let client = Arc :: new ( CandidateClient {
523529 calls : Mutex :: new ( Vec :: new ( ) ) ,
530+ prompts : Mutex :: new ( Vec :: new ( ) ) ,
524531 first,
525532 } ) ;
526- let algorithm = Arc :: new ( CandidateAlgorithm {
533+ let inner : Arc < dyn Algorithm > = Arc :: new ( CandidateAlgorithm {
527534 models : vec ! [ "weak" . into( ) , "strong" . into( ) ] ,
528535 } ) ;
536+ let prompts = TargetPrompts :: default ( )
537+ . with ( "weak" , "weak prompt" )
538+ . with ( "strong" , "strong prompt" ) ;
539+ let algorithm = with_target_prompts ( inner, prompts) ;
529540 let result = run (
530541 algorithm,
531542 ClientRouter :: single ( client. clone ( ) ) ,
@@ -540,6 +551,7 @@ mod tests {
540551 async fn answered_outcome_does_not_make_a_second_model_call ( ) -> Result < ( ) > {
541552 let client = Arc :: new ( CandidateClient {
542553 calls : Mutex :: new ( Vec :: new ( ) ) ,
554+ prompts : Mutex :: new ( Vec :: new ( ) ) ,
543555 first : FirstOutcome :: StreamSuccess ,
544556 } ) ;
545557 let observations = Arc :: new ( Mutex :: new ( Vec :: new ( ) ) ) ;
@@ -621,6 +633,13 @@ mod tests {
621633 & * client. calls. lock( ) ,
622634 & [ ModelId :: from( "weak" ) , "strong" . into( ) ]
623635 ) ;
636+ assert_eq ! (
637+ & * client. prompts. lock( ) ,
638+ & [
639+ vec![ "weak prompt" . to_string( ) ] ,
640+ vec![ "strong prompt" . to_string( ) ]
641+ ]
642+ ) ;
624643 assert_eq ! (
625644 response
626645 . llm_response
0 commit comments