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 crate::{CapabilityJudgeConfig, LlmCapabilityConfig};
176 use std::collections::HashMap;
177 use std::sync::Arc;
178
179 use switchyard_protocol::{Category, Message, ModelId, Role};
180
181 use super::*;
182 use crate::algorithms::util::stage::PickerMode;
183 use crate::algorithms::util::tier_fixtures::{JUDGE, Recorder, turn_request};
184 use crate::core::testing::test_drive_with_models;
185
186 fn runtime_models() -> HashMap<Category, Vec<ModelId>> {
187 [
188 (Category::Judge, vec![ModelId::from(JUDGE)]),
189 (Category::Efficient, vec![ModelId::from("weak")]),
190 (Category::Capable, vec![ModelId::from("strong")]),
191 (
192 Category::Any,
193 vec![ModelId::from("strong"), ModelId::from("weak")],
194 ),
195 ]
196 .into()
197 }
198
199 fn user_turn_request() -> Request {
200 let mut request = turn_request(false);
201 request
202 .llm_request
203 .messages
204 .push(Message::text(Role::User, "now rewrite the parser"));
205 request
206 }
207
208 fn unkeyed(mut request: Request) -> Request {
210 if let Some(metadata) = request.metadata.as_mut() {
211 metadata.session_id = None;
212 }
213 request
214 }
215
216 fn hash_keyed_router() -> Result<Arc<CompositeRouter>> {
217 Ok(Arc::new(CompositeRouter::new(CompositeRouterConfig {
218 judge: TaskClassifierConfig {
219 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
220 base_threshold: 0.5,
221 ..LlmCapabilityConfig::default()
222 }),
223 classify_trigger: ClassifyTrigger::UserTurn,
224 message_hash_fallback: true,
225 ..Default::default()
226 },
227 stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
228 })?))
229 }
230
231 fn router() -> Result<Arc<CompositeRouter>> {
232 Ok(Arc::new(CompositeRouter::new(CompositeRouterConfig {
233 judge: TaskClassifierConfig {
234 judge: CapabilityJudgeConfig::Llm(LlmCapabilityConfig {
235 base_threshold: 0.5,
236 ..LlmCapabilityConfig::default()
237 }),
238 classify_trigger: ClassifyTrigger::UserTurn,
239 ..Default::default()
240 },
241 stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
242 })?))
243 }
244
245 #[test]
246 fn rejects_every_request_as_a_trigger() {
247 let config = CompositeRouterConfig {
248 judge: TaskClassifierConfig::default(),
249 stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
250 };
251 assert!(matches!(
252 CompositeRouter::new(config),
253 Err(LibsyError::AlgorithmError { .. })
254 ));
255 }
256
257 #[tokio::test]
261 async fn a_judge_sharing_the_capable_model_still_latches_the_tier() -> Result<()> {
262 let models: HashMap<Category, Vec<ModelId>> = [
263 (Category::Judge, vec![ModelId::from("strong")]),
264 (Category::Capable, vec![ModelId::from("strong")]),
265 (Category::Efficient, vec![ModelId::from("weak")]),
266 (
267 Category::Any,
268 vec![ModelId::from("strong"), ModelId::from("weak")],
269 ),
270 ]
271 .into();
272 let calls = Arc::new(Mutex::new(0u32));
274 let serve = {
275 let calls = Arc::clone(&calls);
276 move |target: ModelId, _request: Request| {
277 let calls = Arc::clone(&calls);
278 async move {
279 let mut calls = calls.lock();
280 *calls += 1;
281 let completion = if *calls == 1 {
282 r#"{"crux":"bounded task","primary_rule":"SUP-1","capability_boundary":"supported","p_solve":0.1}"#.to_string()
285 } else {
286 target.to_string()
287 };
288 Ok(crate::core::testing::reply(completion))
289 }
290 }
291 };
292 let router = router()?;
293
294 test_drive_with_models(
295 router.clone(),
296 user_turn_request(),
297 models.clone(),
298 serve.clone(),
299 )
300 .await?;
301 let (selected, _) =
303 test_drive_with_models(router, turn_request(false), models, serve).await?;
304
305 assert_eq!(selected, ModelId::from("strong"));
306 Ok(())
307 }
308
309 #[tokio::test]
310 async fn a_session_without_an_id_keys_on_the_message_hash() -> Result<()> {
311 let recorder = Arc::new(Recorder::default());
312 *recorder.judge_p_solve.lock() = 0.1;
313 let router = hash_keyed_router()?;
314
315 test_drive_with_models(
316 router.clone(),
317 unkeyed(user_turn_request()),
318 runtime_models(),
319 recorder.serve(),
320 )
321 .await?;
322 test_drive_with_models(
323 router.clone(),
324 unkeyed(turn_request(false)),
325 runtime_models(),
326 recorder.serve(),
327 )
328 .await?;
329
330 assert_eq!(
331 recorder.judge_calls(),
332 1,
333 "a tool step is not a user turn, session id or not"
334 );
335 assert_eq!(
336 recorder.routed()[1].target,
337 "strong",
338 "and the tier survives the tool step"
339 );
340 Ok(())
341 }
342
343 #[tokio::test]
344 async fn the_judge_sets_the_tier_once_a_turn_and_the_signals_run_within_it() -> Result<()> {
345 let recorder = Arc::new(Recorder::default());
346 *recorder.judge_p_solve.lock() = 0.1;
347 let router = router()?;
348
349 test_drive_with_models(
350 router.clone(),
351 user_turn_request(),
352 runtime_models(),
353 recorder.serve(),
354 )
355 .await?;
356 test_drive_with_models(
357 router.clone(),
358 turn_request(false),
359 runtime_models(),
360 recorder.serve(),
361 )
362 .await?;
363
364 let routed = recorder.routed();
365 assert_eq!(
366 routed[0].target, "strong",
367 "a quiet turn falls open to the verdict"
368 );
369 assert_eq!(
370 routed[1].target, "strong",
371 "which holds across the tool steps after it"
372 );
373 assert_eq!(
374 recorder.judge_calls(),
375 1,
376 "a tool step is not a new user turn"
377 );
378 Ok(())
379 }
380}