Skip to main content

switchyard_libsy/algorithms/
hierarchical.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 [`Preroute`]: 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::preroute::Preroute;
23use crate::core::state::State;
24use crate::{LibsyError, Result};
25use switchyard_protocol::{ModelId, Request};
26
27const HIERARCHICAL: &str = "hierarchical";
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        evict_if_full(&mut tiers);
60        tiers.insert(identity, tier);
61    }
62}
63
64#[async_trait]
65impl Preroute<State> for TierSetter {
66    async fn run(&self, state: &mut State, request: &mut Request, driver: &Driver) -> Result<()> {
67        let identity = retention_key(request, self.message_hash_fallback);
68        if self.is_due(identity.as_ref(), request) {
69            let (classification, _) = self.judge.score(state, request, Some(driver)).await?;
70            if let Some(winner) = classification.argmax(false)?
71                && let Some(tier) = self.targets.tier_for(&winner.target)
72            {
73                set_fall_open(state, tier);
74                if let Some(identity) = identity {
75                    self.retain(identity, tier);
76                }
77                return Ok(());
78            }
79        }
80        // Either not this request's turn or the judge had no verdict, so the last
81        // tier stands rather than dropping back to the picker default.
82        if let Some(tier) = identity.and_then(|identity| self.tiers.lock().get(&identity).copied())
83        {
84            set_fall_open(state, tier);
85        }
86        Ok(())
87    }
88}
89
90/// A judge stacked over a stage router.
91pub struct HierarchicalRouterConfig {
92    /// Target the judge is called through. Not a routing destination.
93    pub judge_target: ModelId,
94    /// Judge settings, including how often `classify_trigger` runs it.
95    pub judge: TaskClassifierConfig,
96    /// Serves the turns, with the tier the judge picked as its fall-open default.
97    pub stage: StageRouterConfig,
98}
99
100/// Runs a stage router with a tier the judge picks.
101pub struct HierarchicalRouter {
102    route: FallThrough<State>,
103}
104
105impl HierarchicalRouter {
106    /// Stacks the judge over a stage router across the same tier pair.
107    ///
108    /// Errors on a configuration either algorithm rejects, on `every_request`, and
109    /// on a stage router carrying its own judge, which would decide the very turns
110    /// this judge set a tier for.
111    pub fn new(
112        capable: ModelId,
113        efficient: ModelId,
114        config: HierarchicalRouterConfig,
115    ) -> Result<Self> {
116        if config.judge.classify_trigger == ClassifyTrigger::EveryRequest {
117            return Err(LibsyError::AlgorithmError {
118                message: "hierarchical: classify_trigger must be user_turn or new_session"
119                    .to_string(),
120            });
121        }
122        if config.stage.llm_fallback.is_some() {
123            return Err(LibsyError::AlgorithmError {
124                message: "hierarchical: the stage router cannot also carry a judge".to_string(),
125            });
126        }
127        let trigger = config.judge.classify_trigger;
128        let message_hash_fallback = config.judge.message_hash_fallback;
129        // Only the judge's Classifier face is used, so its own affinity never runs.
130        // Leaving these set would apply the standalone route's pairing rules to a
131        // trigger this router implements itself.
132        let judge_config = TaskClassifierConfig {
133            classify_trigger: ClassifyTrigger::EveryRequest,
134            message_hash_fallback: false,
135            ..config.judge
136        };
137        let judge = LlmTaskClassifier::new(LlmClassifierConfig::Capability {
138            judge_target: config.judge_target,
139            efficient_target: efficient.clone(),
140            capable_target: capable.clone(),
141            config: judge_config,
142        })?;
143        let setter = TierSetter {
144            judge: Arc::new(judge),
145            targets: StageTargets::new(capable.clone(), efficient.clone()),
146            trigger,
147            message_hash_fallback,
148            tiers: Mutex::new(HashMap::new()),
149        };
150        let route = build_stage_route(capable, efficient, config.stage)?
151            .with_name(HIERARCHICAL)
152            .with_preroute(Arc::new(setter));
153        Ok(Self { route })
154    }
155}
156
157#[async_trait]
158impl Algorithm for HierarchicalRouter {
159    fn name(&self) -> &str {
160        HIERARCHICAL
161    }
162
163    async fn route(
164        self: Arc<Self>,
165        driver: Driver,
166        request: Request,
167    ) -> Result<crate::RoutingOutcome> {
168        self.route.execute(driver, request).await
169    }
170}
171
172#[cfg(test)]
173mod tests {
174    use std::sync::Arc;
175
176    use switchyard_protocol::{Message, Role};
177
178    use super::*;
179    use crate::algorithms::stage::LlmFallback;
180    use crate::algorithms::util::stage::PickerMode;
181    use crate::algorithms::util::tier_fixtures::{JUDGE, Recorder, turn_request};
182    use crate::core::testing::test_drive;
183
184    fn user_turn_request() -> Request {
185        let mut request = turn_request(false);
186        request
187            .llm_request
188            .messages
189            .push(Message::text(Role::User, "now rewrite the parser"));
190        request
191    }
192
193    /// The same request shape with no session ID, so only the hash can key it.
194    fn unkeyed(mut request: Request) -> Request {
195        if let Some(metadata) = request.metadata.as_mut() {
196            metadata.session_id = None;
197        }
198        request
199    }
200
201    fn hash_keyed_router() -> Result<Arc<HierarchicalRouter>> {
202        Ok(Arc::new(HierarchicalRouter::new(
203            ModelId::from("strong"),
204            ModelId::from("weak"),
205            HierarchicalRouterConfig {
206                judge_target: ModelId::from(JUDGE),
207                judge: TaskClassifierConfig {
208                    base_threshold: 0.5,
209                    classify_trigger: ClassifyTrigger::UserTurn,
210                    message_hash_fallback: true,
211                    ..Default::default()
212                },
213                stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
214            },
215        )?))
216    }
217
218    fn router() -> Result<Arc<HierarchicalRouter>> {
219        Ok(Arc::new(HierarchicalRouter::new(
220            ModelId::from("strong"),
221            ModelId::from("weak"),
222            HierarchicalRouterConfig {
223                judge_target: ModelId::from(JUDGE),
224                judge: TaskClassifierConfig {
225                    base_threshold: 0.5,
226                    classify_trigger: ClassifyTrigger::UserTurn,
227                    ..Default::default()
228                },
229                stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
230            },
231        )?))
232    }
233
234    #[test]
235    fn rejects_a_stage_router_that_carries_its_own_judge() {
236        let mut stage = StageRouterConfig::new(PickerMode::EfficientFirst, 0.5);
237        stage.llm_fallback = Some(LlmFallback {
238            judge_target: ModelId::from(JUDGE),
239            config: TaskClassifierConfig::default(),
240        });
241        let config = HierarchicalRouterConfig {
242            judge_target: ModelId::from(JUDGE),
243            judge: TaskClassifierConfig::default(),
244            stage,
245        };
246        assert!(matches!(
247            HierarchicalRouter::new(ModelId::from("strong"), ModelId::from("weak"), config),
248            Err(LibsyError::AlgorithmError { .. })
249        ));
250    }
251
252    #[test]
253    fn rejects_every_request_as_a_trigger() {
254        let config = HierarchicalRouterConfig {
255            judge_target: ModelId::from(JUDGE),
256            judge: TaskClassifierConfig::default(),
257            stage: StageRouterConfig::new(PickerMode::EfficientFirst, 0.5),
258        };
259        assert!(matches!(
260            HierarchicalRouter::new(ModelId::from("strong"), ModelId::from("weak"), config),
261            Err(LibsyError::AlgorithmError { .. })
262        ));
263    }
264
265    #[tokio::test]
266    async fn a_session_without_an_id_keys_on_the_message_hash() -> Result<()> {
267        let recorder = Arc::new(Recorder::default());
268        *recorder.judge_p_solve.lock() = 0.1;
269        let router = hash_keyed_router()?;
270
271        test_drive(
272            router.clone(),
273            unkeyed(user_turn_request()),
274            recorder.serve(),
275        )
276        .await?;
277        test_drive(
278            router.clone(),
279            unkeyed(turn_request(false)),
280            recorder.serve(),
281        )
282        .await?;
283
284        assert_eq!(
285            recorder.judge_calls(),
286            1,
287            "new_session judges once without a session id"
288        );
289        assert_eq!(
290            recorder.routed()[1].target,
291            "strong",
292            "and the tier survives the tool step"
293        );
294        Ok(())
295    }
296
297    #[tokio::test]
298    async fn the_judge_sets_the_tier_once_a_turn_and_the_signals_run_within_it() -> Result<()> {
299        let recorder = Arc::new(Recorder::default());
300        *recorder.judge_p_solve.lock() = 0.1;
301        let router = router()?;
302
303        test_drive(router.clone(), user_turn_request(), recorder.serve()).await?;
304        test_drive(router.clone(), turn_request(false), recorder.serve()).await?;
305
306        let routed = recorder.routed();
307        assert_eq!(
308            routed[0].target, "strong",
309            "a quiet turn falls open to the verdict"
310        );
311        assert_eq!(
312            routed[1].target, "strong",
313            "which holds across the tool steps after it"
314        );
315        assert_eq!(
316            recorder.judge_calls(),
317            1,
318            "a tool step is not a new user turn"
319        );
320        Ok(())
321    }
322}