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::{StageTargets, 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::{ModelId, 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    targets: StageTargets,
37    trigger: ClassifyTrigger,
38    message_hash_fallback: bool,
39    tiers: Mutex<HashMap<RoutingIdentity, Tier>>,
40}
41
42impl TierSetter {
43    /// Two requests for one identity can both pass this and both judge, since a
44    /// judge call sits between here and [`retain`](Self::retain). The later wins.
45    fn is_due(&self, identity: Option<&RoutingIdentity>, request: &Request) -> bool {
46        match self.trigger {
47            ClassifyTrigger::UserTurn => has_new_user_turn(&request.llm_request.messages),
48            // Unkeyed requests cannot be told apart, so every one is a new session.
49            ClassifyTrigger::NewSession => {
50                identity.is_none_or(|identity| !self.tiers.lock().contains_key(identity))
51            }
52            // Rejected by the constructor, and only in the enum for the standalone route.
53            ClassifyTrigger::EveryRequest => true,
54        }
55    }
56
57    fn retain(&self, identity: RoutingIdentity, tier: Tier) {
58        let mut tiers = self.tiers.lock();
59        // A user turn re-decides, so it overwrites. A session keeps its first verdict,
60        // matching how affinity retains an assignment.
61        let writable = self.trigger == ClassifyTrigger::UserTurn || !tiers.contains_key(&identity);
62        if writable {
63            evict_if_full(&mut tiers);
64            tiers.insert(identity, tier);
65        }
66    }
67}
68
69#[async_trait]
70impl Processor<State> for TierSetter {
71    async fn process(&self, state: &mut State, event: Event<'_>) -> Result<()> {
72        let Event::Request { request, driver } = event else {
73            return Ok(());
74        };
75        let identity = retention_key(request, self.message_hash_fallback);
76        if self.is_due(identity.as_ref(), request) {
77            let (classification, _) = self.judge.score(state, request, driver).await?;
78            if let Some(winner) = classification.argmax(false)?
79                && let Some(tier) = self.targets.tier_for(&winner.target)
80            {
81                set_fall_open(state, tier);
82                if let Some(identity) = identity {
83                    self.retain(identity, tier);
84                }
85                return Ok(());
86            }
87        }
88        // Either not this request's turn or the judge had no verdict, so the last
89        // tier stands rather than dropping back to the picker default.
90        if let Some(tier) = identity.and_then(|identity| self.tiers.lock().get(&identity).copied())
91        {
92            set_fall_open(state, tier);
93            if let Some(driver) = driver {
94                driver.set_evidence_if_empty(serde_json::json!({"source": "retained"}));
95            }
96        }
97        Ok(())
98    }
99}
100
101/// A judge stacked over a stage router.
102pub struct CompositeRouterConfig {
103    /// Target the judge is called through. Not a routing destination.
104    pub judge_target: ModelId,
105    /// Judge settings, including how often `classify_trigger` runs it.
106    pub judge: TaskClassifierConfig,
107    /// Serves the turns, with the tier the judge picked as its fall-open default.
108    pub stage: StageRouterConfig,
109}
110
111/// Runs a stage router with a tier the judge picks.
112pub struct CompositeRouter {
113    route: FallThrough<State>,
114}
115
116impl CompositeRouter {
117    /// Stacks the judge over a stage router across the same tier pair.
118    ///
119    /// Errors on a configuration either algorithm rejects and on `every_request`.
120    ///
121    /// A stage router carrying its own judge is allowed, but that judge sits ahead
122    /// of the fall-open tier and so answers most of the turns this one set a tier for.
123    pub fn new(
124        capable: ModelId,
125        efficient: ModelId,
126        config: CompositeRouterConfig,
127    ) -> Result<Self> {
128        if config.judge.classify_trigger == ClassifyTrigger::EveryRequest {
129            return Err(LibsyError::AlgorithmError {
130                message: "composite: classify_trigger must be user_turn or new_session".to_string(),
131            });
132        }
133        let trigger = config.judge.classify_trigger;
134        let message_hash_fallback = config.judge.message_hash_fallback;
135        // Only the judge's Classifier face is used, so its own affinity never runs.
136        // Leaving these set would apply the standalone route's pairing rules to a
137        // trigger this router implements itself.
138        let judge_config = TaskClassifierConfig {
139            classify_trigger: ClassifyTrigger::EveryRequest,
140            message_hash_fallback: false,
141            ..config.judge
142        };
143        let judge = LlmTaskClassifier::new(LlmClassifierConfig::Capability {
144            judge_target: config.judge_target,
145            efficient_target: efficient.clone(),
146            capable_target: capable.clone(),
147            config: judge_config,
148        })?;
149        let setter = TierSetter {
150            judge: Arc::new(judge),
151            targets: StageTargets::new(capable.clone(), efficient.clone()),
152            trigger,
153            message_hash_fallback,
154            tiers: Mutex::new(HashMap::new()),
155        };
156        let route = build_stage_route(capable, efficient, config.stage)?
157            .with_name(COMPOSITE)
158            .with_processor(Arc::new(setter));
159        Ok(Self { route })
160    }
161}
162
163#[async_trait]
164impl Algorithm for CompositeRouter {
165    fn name(&self) -> &str {
166        COMPOSITE
167    }
168
169    async fn route(
170        self: Arc<Self>,
171        driver: Driver,
172        request: Request,
173    ) -> Result<crate::RoutingOutcome> {
174        self.route.execute(driver, request).await
175    }
176}
177
178#[cfg(test)]
179mod tests {
180    use std::sync::Arc;
181
182    use switchyard_protocol::{Message, Role};
183
184    use super::*;
185    use crate::algorithms::util::stage::PickerMode;
186    use crate::algorithms::util::tier_fixtures::{JUDGE, Recorder, turn_request};
187    use crate::core::testing::test_drive;
188
189    fn user_turn_request() -> Request {
190        let mut request = turn_request(false);
191        request
192            .llm_request
193            .messages
194            .push(Message::text(Role::User, "now rewrite the parser"));
195        request
196    }
197
198    /// The same request shape with no session ID, so only the hash can key it.
199    fn unkeyed(mut request: Request) -> Request {
200        if let Some(metadata) = request.metadata.as_mut() {
201            metadata.session_id = None;
202        }
203        request
204    }
205
206    fn hash_keyed_router() -> Result<Arc<CompositeRouter>> {
207        Ok(Arc::new(CompositeRouter::new(
208            ModelId::from("strong"),
209            ModelId::from("weak"),
210            CompositeRouterConfig {
211                judge_target: ModelId::from(JUDGE),
212                judge: TaskClassifierConfig {
213                    base_threshold: 0.5,
214                    classify_trigger: ClassifyTrigger::UserTurn,
215                    message_hash_fallback: true,
216                    ..Default::default()
217                },
218                stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
219            },
220        )?))
221    }
222
223    fn router() -> Result<Arc<CompositeRouter>> {
224        Ok(Arc::new(CompositeRouter::new(
225            ModelId::from("strong"),
226            ModelId::from("weak"),
227            CompositeRouterConfig {
228                judge_target: ModelId::from(JUDGE),
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
239    #[test]
240    fn rejects_every_request_as_a_trigger() {
241        let config = CompositeRouterConfig {
242            judge_target: ModelId::from(JUDGE),
243            judge: TaskClassifierConfig::default(),
244            stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
245        };
246        assert!(matches!(
247            CompositeRouter::new(ModelId::from("strong"), ModelId::from("weak"), config),
248            Err(LibsyError::AlgorithmError { .. })
249        ));
250    }
251
252    #[tokio::test]
253    async fn a_session_without_an_id_keys_on_the_message_hash() -> Result<()> {
254        let recorder = Arc::new(Recorder::default());
255        *recorder.judge_p_solve.lock() = 0.1;
256        let router = hash_keyed_router()?;
257
258        test_drive(
259            router.clone(),
260            unkeyed(user_turn_request()),
261            recorder.serve(),
262        )
263        .await?;
264        test_drive(
265            router.clone(),
266            unkeyed(turn_request(false)),
267            recorder.serve(),
268        )
269        .await?;
270
271        assert_eq!(
272            recorder.judge_calls(),
273            1,
274            "a tool step is not a user turn, session id or not"
275        );
276        assert_eq!(
277            recorder.routed()[1].target,
278            "strong",
279            "and the tier survives the tool step"
280        );
281        Ok(())
282    }
283
284    #[tokio::test]
285    async fn the_judge_sets_the_tier_once_a_turn_and_the_signals_run_within_it() -> Result<()> {
286        let recorder = Arc::new(Recorder::default());
287        *recorder.judge_p_solve.lock() = 0.1;
288        let router = router()?;
289
290        test_drive(router.clone(), user_turn_request(), recorder.serve()).await?;
291        test_drive(router.clone(), turn_request(false), recorder.serve()).await?;
292
293        let routed = recorder.routed();
294        assert_eq!(
295            routed[0].target, "strong",
296            "a quiet turn falls open to the verdict"
297        );
298        assert_eq!(
299            routed[1].target, "strong",
300            "which holds across the tool steps after it"
301        );
302        assert_eq!(
303            recorder.judge_calls(),
304            1,
305            "a tool step is not a new user turn"
306        );
307        Ok(())
308    }
309}