Skip to content

Commit 3ad6855

Browse files
authored
feat(llm-client): prepare routed completion candidates (#463)
Prepares routed completion requests in `libsy-llm-client`, after the algorithm has selected a target and immediately before the client call. The selected target and every fallback are prepared independently, so one target's model ID or request policy cannot leak into the next candidate. This is the second PR for [#496](#496) / [SWITCH-1253](https://linear.app/nvidia/issue/SWITCH-1253/move-using-systempromptprocessor-out-of-stage-router-making-it). #455 is merged; #464 adds the target-level TOML configuration. Signed-off-by: Alex Fournier <afournier@nvidia.com>
1 parent 17b0ff6 commit 3ad6855

3 files changed

Lines changed: 211 additions & 72 deletions

File tree

‎crates/libsy-llm-client/src/lib.rs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,7 @@ pub use client::{ModelConfig, TranslatingLlmClient};
3030
pub use error::{LlmClientError, Result};
3131
pub use observation::{LlmCallObservation, RunObservation, RunObserver};
3232
pub use raw::RawResponse;
33-
pub use run::{ClientRouter, run};
33+
pub use run::{ClientRouter, decide, run};
3434
pub use switchyard_translation::RawEventStream;
3535

3636
/// Registers process-wide compatibility gauges with the global meter provider.

‎crates/libsy-llm-client/src/run.rs‎

Lines changed: 205 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -8,21 +8,22 @@
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
1616
use std::collections::HashMap;
1717
use std::sync::Arc;
1818
use std::time::Instant;
1919

2020
use http::StatusCode;
2121
use parking_lot::Mutex;
22-
use switchyard_libsy::{Algorithm, CallModel, LibsyError, Result, drive};
22+
use switchyard_libsy::{Algorithm, CallModel, LibsyError, Result, RoutingOutcome, drive};
2323
use switchyard_protocol::{
2424
LlmClientError, ModelId, Request, Response, RoutedLlmClient, RoutingFallbackReason,
2525
};
26+
use switchyard_translation::prepare_request_for_target;
2627

2728
use crate::observation::{LlmCallObservation, RunObservation, RunObserver};
2829
use 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.
107128
fn 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.
154181
async 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)]
317341
pub 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

321351
enum Routing {
@@ -328,8 +358,25 @@ enum Routing {
328358
impl 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

368437
impl 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

Comments
 (0)