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::{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
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 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 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
101pub struct CompositeRouterConfig {
103 pub judge_target: ModelId,
105 pub judge: TaskClassifierConfig,
107 pub stage: StageRouterConfig,
109}
110
111pub struct CompositeRouter {
113 route: FallThrough<State>,
114}
115
116impl CompositeRouter {
117 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 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 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}