Skip to main content

switchyard_libsy/algorithms/
composite.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Routing that stacks a judge over a stage router.
5//!
6//! The judge runs as a [`Processor`]: it sets configuration the stage router reads,
7//! and picks no target itself.
8
9use 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
29/// Sets the stage router's fall-open tier from a judge verdict.
30///
31/// Retains the tier per routing identity, so it survives requests that carry no
32/// session ID when `message_hash_fallback` is on. The retained tier is replayed
33/// into state on every request so the cascade below reads it.
34struct 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    /// Two requests for one identity can both pass this and both judge, since a
51    /// judge call sits between here and [`retain`](Self::retain). The later wins.
52    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            // Unkeyed requests cannot be told apart, so every one is a new session.
56            ClassifyTrigger::NewSession => {
57                identity.is_none_or(|identity| !self.tiers.lock().contains_key(identity))
58            }
59            // Rejected by the constructor, and only in the enum for the standalone route.
60            ClassifyTrigger::EveryRequest => true,
61        }
62    }
63
64    fn retain(&self, identity: RoutingIdentity, tier: Tier) {
65        let mut tiers = self.tiers.lock();
66        // A user turn re-decides, so it overwrites. A session keeps its first verdict,
67        // matching how affinity retains an assignment.
68        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        // Either not this request's turn or the judge had no verdict, so the last
96        // tier stands rather than dropping back to the picker default.
97        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
106/// A judge stacked over a stage router.
107pub struct CompositeRouterConfig {
108    /// Judge settings, including how often `classify_trigger` runs it.
109    pub judge: TaskClassifierConfig,
110    /// Serves the turns, with the tier the judge picked as its fall-open default.
111    pub stage: StageRouterConfig,
112}
113
114/// Runs a stage router with a tier the judge picks.
115pub struct CompositeRouter {
116    route: FallThrough<State>,
117}
118
119impl CompositeRouter {
120    /// Stacks the judge over a stage router across the same tier pair.
121    ///
122    /// Errors on a configuration either algorithm rejects and on `every_request`.
123    ///
124    /// A stage router carrying its own judge is allowed, but that judge sits ahead
125    /// of the fall-open tier and so answers most of the turns this one set a tier for.
126    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        // Only the judge's Classifier face is used, so its own affinity never runs.
135        // Leaving these set would apply the standalone route's pairing rules to a
136        // trigger this router implements itself.
137        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    /// The same request shape with no session ID, so only the hash can key it.
209    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    /// A route may point its judge at the same target it serves capable turns on.
258    /// The tier then cannot be recovered by looking the served model up in the
259    /// runtime groups — one id, two categories — so it comes off the verdict itself.
260    #[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        // The judge runs first as a request-side processor, so the opening call is its own.
273        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                        // p_solve below the threshold: the judge does not trust the
283                        // efficient tier, so the verdict is capable.
284                        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        // A tool step is not a user turn, so nothing re-judges and the latched tier decides.
302        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}