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