Skip to content

Commit 6aa6613

Browse files
committed
feat(libsy): Pass models at runtime not construction time- #630
Instead of giving the available models to the algorithm in `new` we pass them alongside the request in `run_stream` where they go in the `Driver`. See #588 Assisted-by: Codex:GPT 5.6 Sol high Signed-off-by: Graham King <grahamk@nvidia.com>
1 parent 8dc8911 commit 6aa6613

47 files changed

Lines changed: 1802 additions & 1660 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎README.md‎

Lines changed: 3 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -171,17 +171,15 @@ switchyard-protocol = { git = "https://github.com/NVIDIA-NeMo/Switchyard.git", b
171171
tokio = { version = "1", features = ["macros", "rt"] }
172172
```
173173

174-
**2. Construct an algorithm.** Target names are whatever your harness calls its
175-
models. This is the stage router from the benchmark; `random`,
176-
`llm_task_classifier`, and `llm_classifier` are built the same way.
174+
**2. Construct an algorithm.** Models are supplied when each request runs. This
175+
is the stage router from the benchmark; `random`, `llm_task_classifier`, and
176+
`llm_classifier` are built the same way.
177177

178178
```python
179179
from switchyard.libsy import LlmResponse, Step
180180
from switchyard.libsy.algorithms import stage_router
181181

182182
algorithm = stage_router(
183-
"capable",
184-
"efficient",
185183
picker="efficient_first",
186184
confidence_threshold=0.5,
187185
)

‎benchmark/routing-profiles/tau2-telecom-custom-opus-qwen-aggressive.toml‎

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -50,19 +50,18 @@ llm_client = "openrouter"
5050
[routes.switchyard]
5151
id = "switchyard"
5252
type = "llm_classifier"
53-
classifier_target = "classifier"
5453
mode = "custom"
5554
recent_turn_window = 6
5655
# Re-classify when the user speaks again, and hold that target across the tool calls
5756
# in between, so a tool chain never switches tier mid-task.
5857
classify_trigger = "user_turn"
59-
targets = ["weak", "strong"]
60-
default_target = "strong"
58+
models = { judge = ["classifier"], capable = ["strong"], efficient = ["weak"], any = ["strong", "weak"] }
59+
default_target = "capable"
6160
response_schema = '''
6261
{
6362
"type": "object",
6463
"properties": {
65-
"route": { "type": "string", "enum": ["weak", "strong"] },
64+
"route": { "type": "string", "enum": ["efficient", "capable"] },
6665
"confidence": { "type": "number" },
6766
"abstain": { "type": "boolean" }
6867
},
@@ -74,10 +73,10 @@ prompt = '''
7473
You are a routing classifier inside a customer-service agent. Return exactly
7574
one JSON object:
7675
77-
{"route": "weak" or "strong", "confidence": number 0..1, "abstain": boolean}
76+
{"route": "efficient" or "capable", "confidence": number 0..1, "abstain": boolean}
7877
79-
State the route DIRECTLY: "weak" = the on-device assistant handles this turn;
80-
"strong" = escalate this turn to the frontier model.
78+
State the route DIRECTLY: "efficient" = the on-device assistant handles this turn;
79+
"capable" = escalate this turn to the frontier model.
8180
8281
ROUTING BIAS: the WEAK tier is the DEFAULT — it handles nearly all support
8382
work end-to-end: lookups, standard actions, troubleshooting with known steps,

‎benchmark/routing-profiles/tau2-telecom-custom-opus-qwen-balanced.toml‎

Lines changed: 8 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -50,19 +50,18 @@ llm_client = "openrouter"
5050
[routes.switchyard]
5151
id = "switchyard"
5252
type = "llm_classifier"
53-
classifier_target = "classifier"
5453
mode = "custom"
5554
recent_turn_window = 6
5655
# Re-classify when the user speaks again, and hold that target across the tool calls
5756
# in between, so a tool chain never switches tier mid-task.
5857
classify_trigger = "user_turn"
59-
targets = ["weak", "strong"]
60-
default_target = "strong"
58+
models = { judge = ["classifier"], capable = ["strong"], efficient = ["weak"], any = ["strong", "weak"] }
59+
default_target = "capable"
6160
response_schema = '''
6261
{
6362
"type": "object",
6463
"properties": {
65-
"route": { "type": "string", "enum": ["weak", "strong"] },
64+
"route": { "type": "string", "enum": ["efficient", "capable"] },
6665
"confidence": { "type": "number" },
6766
"abstain": { "type": "boolean" }
6867
},
@@ -76,20 +75,20 @@ You see a condensed view of the conversation: the original request, recent turns
7675
(including tool results), and the customer's newest message. Return exactly one
7776
JSON object:
7877
79-
{"route": "weak" or "strong", "confidence": number 0..1, "abstain": boolean}
78+
{"route": "efficient" or "capable", "confidence": number 0..1, "abstain": boolean}
8079
81-
State the route DIRECTLY: "weak" = the on-device assistant handles this turn;
82-
"strong" = escalate this turn to the frontier model. Decide for the customer's
80+
State the route DIRECTLY: "efficient" = the on-device assistant handles this turn;
81+
"capable" = escalate this turn to the frontier model. Decide for the customer's
8382
NEWEST request, using the recent turns as context.
8483
85-
Route "weak" when the newest request is ROUTINE — the procedure is
84+
Route "efficient" when the newest request is ROUTINE — the procedure is
8685
clear and it's about executing it:
8786
account/order/status lookups, reading or relaying tool results, standard
8887
single-step actions (toggle a setting, resend a code, restart a service),
8988
collecting information from the customer, confirmations, pleasantries,
9089
straightforward troubleshooting with an obvious next step.
9190
92-
Route "strong" when the newest request needs NON-OBVIOUS
91+
Route "capable" when the newest request needs NON-OBVIOUS
9392
JUDGMENT the routine tier may get wrong:
9493
applying or reconciling POLICY with multiple conditions (eligibility, refunds,
9594
exceptions, proration), conflicts between what the customer wants and what

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

Lines changed: 45 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@ use http::StatusCode;
2121
use parking_lot::Mutex;
2222
use switchyard_libsy::{Algorithm, CallModel, LibsyError, Result, RoutingOutcome, drive};
2323
use switchyard_protocol::{
24-
LlmClientError, ModelId, Request, Response, RoutedLlmClient, RoutingFallbackReason,
24+
Category, LlmClientError, ModelId, Request, Response, RoutedLlmClient, RoutingFallbackReason,
2525
};
2626
use switchyard_translation::prepare_request_for_target;
2727

@@ -44,6 +44,7 @@ pub async fn run(
4444
algorithm: Arc<dyn Algorithm>,
4545
clients: ClientRouter,
4646
request: Request,
47+
models: HashMap<Category, Vec<ModelId>>,
4748
observer: Option<RunObserver>,
4849
) -> Result<(ModelId, Response)> {
4950
let algorithm_name = algorithm.name().to_string();
@@ -52,7 +53,7 @@ pub async fn run(
5253
// This says if we have an observer, put Some(..) in routing_observations.
5354
// No observer means we don't want any routing_observations.
5455
let routing_observations = observer.as_ref().map(|_| Arc::new(Mutex::new(Vec::new())));
55-
let outcome = drive(algorithm, request, {
56+
let outcome = drive(algorithm, request, models, {
5657
let routing_observations = routing_observations.clone();
5758
move |call| serve(routing_clients.clone(), call, routing_observations.clone())
5859
})
@@ -110,9 +111,10 @@ pub async fn decide(
110111
algorithm: Arc<dyn Algorithm>,
111112
clients: ClientRouter,
112113
request: Request,
114+
models: HashMap<Category, Vec<ModelId>>,
113115
) -> Result<RoutingOutcome> {
114116
let routing_clients = clients.clone();
115-
let mut outcome = drive(algorithm, request, move |call| {
117+
let mut outcome = drive(algorithm, request, models, move |call| {
116118
serve(routing_clients.clone(), call, None)
117119
})
118120
.await?;
@@ -453,9 +455,7 @@ mod tests {
453455

454456
use crate::{Backend, HttpBackendConfig, ModelConfig, TranslatingLlmClient};
455457

456-
struct CandidateAlgorithm {
457-
models: Vec<ModelId>,
458-
}
458+
struct CandidateAlgorithm {}
459459

460460
struct AnsweredAlgorithm {
461461
model: ModelId,
@@ -469,13 +469,14 @@ mod tests {
469469

470470
async fn route(
471471
self: Arc<Self>,
472-
_driver: Driver,
472+
driver: Driver,
473473
request: Request,
474474
) -> Result<RoutingOutcome> {
475-
let selected_model = self.models.first().cloned().ok_or(LibsyError::NoTargets)?;
475+
let models = driver.models_for(Category::Any);
476+
let selected_model = models.first().cloned().ok_or(LibsyError::NoTargets)?;
476477
Ok(RoutingOutcome::route_to(
477478
selected_model,
478-
self.models.iter().skip(1).cloned().collect(),
479+
models.iter().skip(1).cloned().collect(),
479480
request,
480481
))
481482
}
@@ -603,19 +604,27 @@ mod tests {
603604
requests: Mutex::new(Vec::new()),
604605
first,
605606
});
606-
let algorithm = Arc::new(CandidateAlgorithm {
607-
models: vec!["weak".into(), "strong".into()],
608-
});
607+
let algorithm = Arc::new(CandidateAlgorithm {});
608+
let models = to_category_map(&["weak", "strong"]);
609609
let result = run(
610610
algorithm,
611611
ClientRouter::single(client.clone()),
612612
request(),
613+
models,
613614
None,
614615
)
615616
.await;
616617
(client, result)
617618
}
618619

620+
fn to_category_map(names: &[&str]) -> HashMap<Category, Vec<ModelId>> {
621+
[(
622+
Category::Any,
623+
names.iter().map(|name| ModelId::from(*name)).collect(),
624+
)]
625+
.into()
626+
}
627+
619628
#[tokio::test]
620629
async fn answered_outcome_does_not_make_a_second_model_call() -> Result<()> {
621630
let client = Arc::new(CandidateClient {
@@ -633,6 +642,7 @@ mod tests {
633642
}),
634643
ClientRouter::single(client.clone()),
635644
request(),
645+
HashMap::new(),
636646
Some(observer),
637647
)
638648
.await?;
@@ -697,11 +707,10 @@ mod tests {
697707
);
698708

699709
run(
700-
Arc::new(CandidateAlgorithm {
701-
models: vec!["weak".into(), "strong".into()],
702-
}),
710+
Arc::new(CandidateAlgorithm {}),
703711
clients,
704712
request(),
713+
to_category_map(&["weak", "strong"]),
705714
None,
706715
)
707716
.await?;
@@ -733,6 +742,7 @@ mod tests {
733742
}),
734743
clients,
735744
request(),
745+
HashMap::new(),
736746
)
737747
.await?;
738748

@@ -763,11 +773,10 @@ mod tests {
763773
);
764774

765775
let outcome = decide(
766-
Arc::new(CandidateAlgorithm {
767-
models: vec!["weak".into(), "strong".into()],
768-
}),
776+
Arc::new(CandidateAlgorithm {}),
769777
clients,
770778
request(),
779+
to_category_map(&["weak", "strong"]),
771780
)
772781
.await?;
773782

@@ -904,10 +913,15 @@ mod tests {
904913
])
905914
.map_err(|error| LibsyError::external("building test client", error))?,
906915
);
907-
let algorithm = Arc::new(CandidateAlgorithm {
908-
models: vec!["weak".into(), "strong".into()],
909-
});
910-
run(algorithm, ClientRouter::single(client), request(), None).await?;
916+
let algorithm = Arc::new(CandidateAlgorithm {});
917+
run(
918+
algorithm,
919+
ClientRouter::single(client),
920+
request(),
921+
to_category_map(&["weak", "strong"]),
922+
None,
923+
)
924+
.await?;
911925

912926
assert_eq!(&*calls.lock(), &["weak", "weak", "weak", "strong"]);
913927
Ok(())
@@ -995,17 +1009,22 @@ mod tests {
9951009
])
9961010
.expect("building test client"),
9971011
);
998-
let algorithm = Arc::new(CandidateAlgorithm {
999-
models: vec!["weak".into(), "strong".into()],
1000-
});
1012+
let algorithm = Arc::new(CandidateAlgorithm {});
10011013
let mut llm_request = text_request(Some("auto".to_string()), "hello".to_string());
10021014
llm_request.stream = true;
10031015
let request = Request {
10041016
llm_request,
10051017
raw_request: None,
10061018
metadata: None,
10071019
};
1008-
let result = run(algorithm, ClientRouter::single(client), request, None).await;
1020+
let result = run(
1021+
algorithm,
1022+
ClientRouter::single(client),
1023+
request,
1024+
to_category_map(&["weak", "strong"]),
1025+
None,
1026+
)
1027+
.await;
10091028
(server, calls, result)
10101029
}
10111030

0 commit comments

Comments
 (0)