1use std::sync::Arc;
7
8use switchyard_protocol::{Category, Metadata, Request};
9
10use super::fall_through::FallThrough;
11use super::util::affinity::{AffinityRouter, ClassifyTrigger};
12use super::util::subagent::SubagentGate;
13use crate::algorithms::llm_class::DefaultCategoryClassifier;
14use crate::core::algorithm::{Algorithm, Driver};
15use crate::core::classifier::Classifier;
16use crate::core::state::State;
17use crate::{LibsyError, Result, RoutingOutcome};
18
19pub struct SubagentRouterConfig {
21 pub classifier: Arc<dyn Classifier<State>>,
23 pub default_target: Category,
25 pub classify_trigger: ClassifyTrigger,
27 pub message_hash_fallback: bool,
29}
30
31impl SubagentRouterConfig {
32 pub fn fixed_target() -> Self {
34 Self {
35 classifier: Arc::new(DefaultCategoryClassifier(Category::Any)),
36 default_target: Category::Any,
37 classify_trigger: ClassifyTrigger::EveryRequest,
38 message_hash_fallback: false,
39 }
40 }
41}
42
43pub struct SubagentRouter {
45 parent: Arc<dyn Algorithm>,
46 subagent: FallThrough<State>,
47}
48
49impl SubagentRouter {
50 pub fn new(parent: Arc<dyn Algorithm>, config: SubagentRouterConfig) -> Result<Self> {
56 if config.message_hash_fallback {
57 return Err(LibsyError::AlgorithmError {
58 message: "sub-agent routing cannot use message_hash_fallback".to_string(),
59 });
60 }
61
62 let mut subagent = match config.classify_trigger {
63 ClassifyTrigger::EveryRequest => FallThrough::new_with_state().with_name("subagent"),
64 ClassifyTrigger::NewSession => {
65 let affinity = Arc::new(AffinityRouter::for_subagents());
66 FallThrough::new_with_state()
67 .with_name("subagent")
68 .with_processor(affinity.clone())
69 .with_classifier(affinity)
70 }
71 ClassifyTrigger::UserTurn => {
72 return Err(LibsyError::AlgorithmError {
73 message: "sub-agent routing cannot use classify_trigger = user_turn"
74 .to_string(),
75 });
76 }
77 };
78 subagent = subagent
79 .with_classifier(Arc::new(SubagentGate::new(config.classifier)))
80 .with_classifier(Arc::new(DefaultCategoryClassifier(config.default_target)));
81
82 Ok(Self { parent, subagent })
83 }
84}
85
86#[async_trait::async_trait]
87impl Algorithm for SubagentRouter {
88 fn name(&self) -> &str {
89 self.parent.name()
90 }
91
92 async fn route(self: Arc<Self>, driver: Driver, request: Request) -> Result<RoutingOutcome> {
93 if request
94 .metadata
95 .as_ref()
96 .is_some_and(Metadata::is_subagent_work)
97 {
98 self.subagent.execute(driver.for_subagent()?, request).await
100 } else {
101 self.parent.clone().route(driver, request).await
102 }
103 }
104}
105
106#[cfg(test)]
107mod tests {
108 use std::sync::Arc;
109 use std::sync::atomic::{AtomicUsize, Ordering};
110
111 use async_trait::async_trait;
112 use parking_lot::Mutex;
113 use serde_json::json;
114 use switchyard_protocol::{
115 Category, ContentBlock, InstructionBlock, Message, Metadata, ModelId, Request, Response,
116 Role, text_request,
117 };
118
119 use super::{SubagentRouter, SubagentRouterConfig};
120 use crate::algorithms::passthrough::Passthrough;
121 use crate::core::classifier::{Classification, Classifier, Score};
122 use crate::core::testing::{echo, reply, test_drive_with_models};
123 use crate::{
124 ClassifyTrigger, CustomClassifierConfig, CustomClassifierPolicy, Driver,
125 LlmClassifierConfig, LlmTaskClassifier, RuntimeModels, State,
126 };
127
128 struct ScriptedClassifier {
129 calls: AtomicUsize,
130 }
131
132 #[async_trait]
133 impl Classifier<State> for ScriptedClassifier {
134 async fn score(
135 &self,
136 _state: &mut State,
137 _request: &mut Request,
138 driver: &Driver,
139 ) -> crate::Result<(Classification, Option<Response>)> {
140 let category = match self.calls.fetch_add(1, Ordering::Relaxed) {
141 0 => Some(Category::Capable),
142 1 => Some(Category::Efficient),
143 _ => None,
144 };
145 let scores = match category {
146 Some(category) => vec![Score {
147 confidence: 1.0,
148 target: driver.first_model_for(&category)?.clone(),
149 category: Some(category),
150 }],
151 None => Vec::new(),
152 };
153 Ok((Classification::Scores(scores), None))
154 }
155 }
156
157 fn request(metadata: Option<Metadata>) -> Request {
158 Request {
159 llm_request: text_request(Some("auto".to_string()), "hi"),
160 raw_request: None,
161 metadata,
162 }
163 }
164
165 fn child(agent_id: &str) -> Request {
166 request(Some(Metadata {
167 session_id: Some("session-1".to_string()),
168 agent_id: Some(agent_id.to_string()),
169 is_subagent: true,
170 is_delegated_work: true,
171 ..Metadata::default()
172 }))
173 }
174
175 fn configured(classifier: Arc<dyn Classifier<State>>) -> crate::Result<Arc<SubagentRouter>> {
176 Ok(Arc::new(SubagentRouter::new(
177 Arc::new(Passthrough),
178 SubagentRouterConfig {
179 classifier,
180 default_target: Category::Capable,
181 classify_trigger: ClassifyTrigger::NewSession,
182 message_hash_fallback: false,
183 },
184 )?))
185 }
186
187 #[tokio::test]
188 async fn routes_parent_and_children_with_affinity_and_default() -> crate::Result<()> {
189 let classifier = Arc::new(ScriptedClassifier {
190 calls: AtomicUsize::new(0),
191 });
192 let router = configured(classifier.clone())?;
193
194 let models = RuntimeModels::new([(Category::Any, vec![ModelId::from("parent")])].into())
196 .with_subagent(
197 [
198 (
199 Category::Any,
200 vec![ModelId::from("worker"), ModelId::from("reviewer")],
201 ),
202 (Category::Capable, vec![ModelId::from("worker")]),
203 (Category::Efficient, vec![ModelId::from("reviewer")]),
204 ]
205 .into(),
206 );
207 let (selected_parent, _) =
208 test_drive_with_models(router.clone(), request(None), models.clone(), echo()).await?;
209 let (first, _) =
210 test_drive_with_models(router.clone(), child("child-1"), models.clone(), echo())
211 .await?;
212 let (same_child, _) =
213 test_drive_with_models(router.clone(), child("child-1"), models.clone(), echo())
214 .await?;
215 let (sibling, _) =
216 test_drive_with_models(router.clone(), child("child-2"), models.clone(), echo())
217 .await?;
218 let (defaulted, _) =
219 test_drive_with_models(router.clone(), child("child-3"), models.clone(), echo())
220 .await?;
221 let maintenance = request(Some(Metadata {
222 session_id: Some("session-1".to_string()),
223 agent_id: Some("child-1".to_string()),
224 is_subagent: true,
225 is_delegated_work: false,
226 ..Metadata::default()
227 }));
228 let (maintenance, _) =
229 test_drive_with_models(router, maintenance, models.clone(), echo()).await?;
230
231 assert_eq!(selected_parent, "parent");
232 assert_eq!(first, "worker");
233 assert_eq!(same_child, "worker");
234 assert_eq!(sibling, "reviewer");
235 assert_eq!(defaulted, "worker");
236 assert_eq!(maintenance, "parent");
237 assert_eq!(classifier.calls.load(Ordering::Relaxed), 3);
238
239 let fixed = Arc::new(SubagentRouter::new(
240 Arc::new(Passthrough),
241 SubagentRouterConfig::fixed_target(),
242 )?);
243 let (fixed, _) = test_drive_with_models(fixed, child("fixed"), models, echo()).await?;
244 assert_eq!(fixed, "worker");
245 Ok(())
246 }
247
248 #[tokio::test]
249 async fn custom_classifier_receives_only_the_delegated_prompt() -> crate::Result<()> {
250 let classifier = LlmTaskClassifier::new(LlmClassifierConfig::Custom {
251 default_target: Category::Capable,
252 config: CustomClassifierConfig::new(
253 "classify the delegated task",
254 json!({
255 "type": "object",
256 "properties": {
257 "target": {"type": "string", "enum": ["capable", "efficient"]}
258 },
259 "required": ["target"],
260 "additionalProperties": false
261 }),
262 CustomClassifierPolicy::target_selector("/target"),
263 ),
264 })?;
265 let router = configured(Arc::new(classifier))?;
266 let mut request = child("child-1");
267 request.llm_request.instructions = vec![InstructionBlock {
268 role: Role::System,
269 content: Message::text(Role::System, "child system instructions").content,
270 }];
271 request.llm_request.messages = vec![
272 Message::text(Role::User, "harness context"),
273 Message {
274 role: Role::User,
275 content: vec![
276 ContentBlock::Text {
277 text: "<system-reminder>tool context</system-reminder>".to_string(),
278 },
279 ContentBlock::Text {
280 text: "review this parser".to_string(),
281 },
282 ],
283 },
284 ];
285 let calls = Arc::new(Mutex::new(Vec::new()));
286 let served_calls = calls.clone();
287
288 let models = RuntimeModels::new([(Category::Any, vec![ModelId::from("parent")])].into())
289 .with_subagent(
290 [
291 (Category::Judge, vec![ModelId::from("judge")]),
292 (Category::Capable, vec![ModelId::from("worker")]),
293 (Category::Efficient, vec![ModelId::from("reviewer")]),
294 (
295 Category::Any,
296 vec![ModelId::from("worker"), ModelId::from("reviewer")],
297 ),
298 ]
299 .into(),
300 );
301 let (selected, _) =
302 test_drive_with_models(router, request, models, move |target, request| {
303 let calls = served_calls.clone();
304 async move {
305 let completion = if target == "judge" {
306 r#"{"target":"efficient"}"#
307 } else {
308 "child answer"
309 };
310 calls.lock().push((target, request));
311 Ok(reply(completion))
312 }
313 })
314 .await?;
315
316 assert_eq!(selected, "reviewer");
317 let calls = calls.lock();
318 assert_eq!(calls.len(), 2);
319 assert_eq!(calls[0].0, "judge");
320 assert_eq!(
321 calls[0].1.llm_request.instructions[0].content,
322 Message::text(Role::System, "classify the delegated task").content
323 );
324 assert_eq!(
325 calls[0].1.llm_request.messages,
326 vec![Message::text(Role::User, "review this parser")]
327 );
328 assert_eq!(calls[1].0, "reviewer");
329 assert_eq!(calls[1].1.llm_request.instructions.len(), 1);
330 assert_eq!(calls[1].1.llm_request.messages.len(), 2);
331 Ok(())
332 }
333}