switchyard_libsy/algorithms/
composite.rs1use std::collections::HashMap;
10use std::sync::Arc;
11
12use async_trait::async_trait;
13use parking_lot::Mutex;
14
15use super::fall_through::FallThrough;
16use super::llm_class::{LlmClassifierConfig, LlmTaskClassifier, TaskClassifierConfig};
17use super::stage::{StageRouterConfig, build_stage_route};
18use super::util::affinity::{ClassifyTrigger, evict_if_full, has_new_user_turn, retention_key};
19use super::util::stage::{Tier, set_fall_open};
20use crate::core::algorithm::{Algorithm, Driver, RoutingIdentity};
21use crate::core::classifier::Classifier;
22use crate::core::processor::{Event, Processor};
23use crate::core::state::State;
24use crate::{LibsyError, Result};
25use switchyard_protocol::{Category, Request};
26
27const COMPOSITE: &str = "composite";
28
29struct TierSetter {
35 judge: Arc<dyn Classifier<State>>,
36 trigger: ClassifyTrigger,
37 message_hash_fallback: bool,
38 tiers: Mutex<HashMap<RoutingIdentity, Tier>>,
39}
40
41impl TierSetter {
42 fn tier_for(category: Option<Category>) -> Option<Tier> {
43 match category {
44 Some(Category::Capable) => Some(Tier::Capable),
45 Some(Category::Efficient) => Some(Tier::Efficient),
46 _ => None,
47 }
48 }
49
50 fn is_due(&self, identity: Option<&RoutingIdentity>, request: &Request) -> bool {
53 match self.trigger {
54 ClassifyTrigger::UserTurn => has_new_user_turn(&request.llm_request.messages),
55 ClassifyTrigger::NewSession => {
57 identity.is_none_or(|identity| !self.tiers.lock().contains_key(identity))
58 }
59 ClassifyTrigger::EveryRequest => true,
61 }
62 }
63
64 fn retain(&self, identity: RoutingIdentity, tier: Tier) {
65 let mut tiers = self.tiers.lock();
66 let writable = self.trigger == ClassifyTrigger::UserTurn || !tiers.contains_key(&identity);
69 if writable {
70 evict_if_full(&mut tiers);
71 tiers.insert(identity, tier);
72 }
73 }
74}
75
76#[async_trait]
77impl Processor<State> for TierSetter {
78 async fn process(&self, state: &mut State, event: Event<'_>) -> Result<()> {
79 let Event::Request { request, driver } = event else {
80 return Ok(());
81 };
82 let identity = retention_key(request, self.message_hash_fallback);
83 if self.is_due(identity.as_ref(), request) {
84 let (classification, _) = self.judge.score(state, request, driver).await?;
85 if let Some(winner) = classification.argmax(false)?
86 && let Some(tier) = Self::tier_for(winner.category)
87 {
88 set_fall_open(state, tier);
89 if let Some(identity) = identity {
90 self.retain(identity, tier);
91 }
92 return Ok(());
93 }
94 }
95 if let Some(tier) = identity.and_then(|identity| self.tiers.lock().get(&identity).copied())
98 {
99 set_fall_open(state, tier);
100 driver.set_evidence_if_empty(serde_json::json!({"source": "retained"}));
101 }
102 Ok(())
103 }
104}
105
106pub struct CompositeRouterConfig {
108 pub judge: TaskClassifierConfig,
110 pub stage: StageRouterConfig,
112}
113
114pub struct CompositeRouter {
116 route: FallThrough<State>,
117}
118
119impl CompositeRouter {
120 pub fn new(config: CompositeRouterConfig) -> Result<Self> {
127 if config.judge.classify_trigger == ClassifyTrigger::EveryRequest {
128 return Err(LibsyError::AlgorithmError {
129 message: "composite: classify_trigger must be user_turn or new_session".to_string(),
130 });
131 }
132 let trigger = config.judge.classify_trigger;
133 let message_hash_fallback = config.judge.message_hash_fallback;
134 let judge_config = TaskClassifierConfig {
138 classify_trigger: ClassifyTrigger::EveryRequest,
139 message_hash_fallback: false,
140 ..config.judge
141 };
142 let judge = LlmTaskClassifier::new(LlmClassifierConfig::Capability {
143 config: judge_config,
144 })?;
145 let setter = TierSetter {
146 judge: Arc::new(judge),
147 trigger,
148 message_hash_fallback,
149 tiers: Mutex::new(HashMap::new()),
150 };
151 let route = build_stage_route(config.stage)?
152 .with_name(COMPOSITE)
153 .with_processor(Arc::new(setter));
154 Ok(Self { route })
155 }
156}
157
158#[async_trait]
159impl Algorithm for CompositeRouter {
160 fn name(&self) -> &str {
161 COMPOSITE
162 }
163
164 async fn route(
165 self: Arc<Self>,
166 driver: Driver,
167 request: Request,
168 ) -> Result<crate::RoutingOutcome> {
169 self.route.execute(driver, request).await
170 }
171}
172
173#[cfg(test)]
174mod tests {
175 use std::collections::HashMap;
176 use std::sync::Arc;
177
178 use switchyard_protocol::{Category, Message, ModelId, Role};
179
180 use super::*;
181 use crate::algorithms::util::stage::PickerMode;
182 use crate::algorithms::util::tier_fixtures::{JUDGE, Recorder, turn_request};
183 use crate::core::testing::test_drive_with_models;
184
185 fn runtime_models() -> HashMap<Category, Vec<ModelId>> {
186 [
187 (Category::Judge, vec![ModelId::from(JUDGE)]),
188 (Category::Efficient, vec![ModelId::from("weak")]),
189 (Category::Capable, vec![ModelId::from("strong")]),
190 (
191 Category::Any,
192 vec![ModelId::from("strong"), ModelId::from("weak")],
193 ),
194 ]
195 .into()
196 }
197
198 fn user_turn_request() -> Request {
199 let mut request = turn_request(false);
200 request
201 .llm_request
202 .messages
203 .push(Message::text(Role::User, "now rewrite the parser"));
204 request
205 }
206
207 fn unkeyed(mut request: Request) -> Request {
209 if let Some(metadata) = request.metadata.as_mut() {
210 metadata.session_id = None;
211 }
212 request
213 }
214
215 fn hash_keyed_router() -> Result<Arc<CompositeRouter>> {
216 Ok(Arc::new(CompositeRouter::new(CompositeRouterConfig {
217 judge: TaskClassifierConfig {
218 base_threshold: 0.5,
219 classify_trigger: ClassifyTrigger::UserTurn,
220 message_hash_fallback: true,
221 ..Default::default()
222 },
223 stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
224 })?))
225 }
226
227 fn router() -> Result<Arc<CompositeRouter>> {
228 Ok(Arc::new(CompositeRouter::new(CompositeRouterConfig {
229 judge: TaskClassifierConfig {
230 base_threshold: 0.5,
231 classify_trigger: ClassifyTrigger::UserTurn,
232 ..Default::default()
233 },
234 stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
235 })?))
236 }
237
238 #[test]
239 fn rejects_every_request_as_a_trigger() {
240 let config = CompositeRouterConfig {
241 judge: TaskClassifierConfig::default(),
242 stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
243 };
244 assert!(matches!(
245 CompositeRouter::new(config),
246 Err(LibsyError::AlgorithmError { .. })
247 ));
248 }
249
250 #[tokio::test]
254 async fn a_judge_sharing_the_capable_model_still_latches_the_tier() -> Result<()> {
255 let models: HashMap<Category, Vec<ModelId>> = [
256 (Category::Judge, vec![ModelId::from("strong")]),
257 (Category::Capable, vec![ModelId::from("strong")]),
258 (Category::Efficient, vec![ModelId::from("weak")]),
259 (
260 Category::Any,
261 vec![ModelId::from("strong"), ModelId::from("weak")],
262 ),
263 ]
264 .into();
265 let calls = Arc::new(Mutex::new(0u32));
267 let serve = {
268 let calls = Arc::clone(&calls);
269 move |target: ModelId, _request: Request| {
270 let calls = Arc::clone(&calls);
271 async move {
272 let mut calls = calls.lock();
273 *calls += 1;
274 let completion = if *calls == 1 {
275 r#"{"crux":"bounded task","primary_rule":"SUP-1","capability_boundary":"supported","p_solve":0.1}"#.to_string()
278 } else {
279 target.to_string()
280 };
281 Ok(crate::core::testing::reply(completion))
282 }
283 }
284 };
285 let router = router()?;
286
287 test_drive_with_models(
288 router.clone(),
289 user_turn_request(),
290 models.clone(),
291 serve.clone(),
292 )
293 .await?;
294 let (selected, _) =
296 test_drive_with_models(router, turn_request(false), models, serve).await?;
297
298 assert_eq!(selected, ModelId::from("strong"));
299 Ok(())
300 }
301
302 #[tokio::test]
303 async fn a_session_without_an_id_keys_on_the_message_hash() -> Result<()> {
304 let recorder = Arc::new(Recorder::default());
305 *recorder.judge_p_solve.lock() = 0.1;
306 let router = hash_keyed_router()?;
307
308 test_drive_with_models(
309 router.clone(),
310 unkeyed(user_turn_request()),
311 runtime_models(),
312 recorder.serve(),
313 )
314 .await?;
315 test_drive_with_models(
316 router.clone(),
317 unkeyed(turn_request(false)),
318 runtime_models(),
319 recorder.serve(),
320 )
321 .await?;
322
323 assert_eq!(
324 recorder.judge_calls(),
325 1,
326 "a tool step is not a user turn, session id or not"
327 );
328 assert_eq!(
329 recorder.routed()[1].target,
330 "strong",
331 "and the tier survives the tool step"
332 );
333 Ok(())
334 }
335
336 #[tokio::test]
337 async fn the_judge_sets_the_tier_once_a_turn_and_the_signals_run_within_it() -> Result<()> {
338 let recorder = Arc::new(Recorder::default());
339 *recorder.judge_p_solve.lock() = 0.1;
340 let router = router()?;
341
342 test_drive_with_models(
343 router.clone(),
344 user_turn_request(),
345 runtime_models(),
346 recorder.serve(),
347 )
348 .await?;
349 test_drive_with_models(
350 router.clone(),
351 turn_request(false),
352 runtime_models(),
353 recorder.serve(),
354 )
355 .await?;
356
357 let routed = recorder.routed();
358 assert_eq!(
359 routed[0].target, "strong",
360 "a quiet turn falls open to the verdict"
361 );
362 assert_eq!(
363 routed[1].target, "strong",
364 "which holds across the tool steps after it"
365 );
366 assert_eq!(
367 recorder.judge_calls(),
368 1,
369 "a tool step is not a new user turn"
370 );
371 Ok(())
372 }
373}