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