switchyard_libsy/algorithms/
hierarchical.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::{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
29struct 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 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 ClassifyTrigger::NewSession => {
50 identity.is_none_or(|identity| !self.tiers.lock().contains_key(identity))
51 }
52 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 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
90pub struct HierarchicalRouterConfig {
92 pub judge_target: ModelId,
94 pub judge: TaskClassifierConfig,
96 pub stage: StageRouterConfig,
98}
99
100pub struct HierarchicalRouter {
102 route: FallThrough<State>,
103}
104
105impl HierarchicalRouter {
106 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 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 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}