Skip to main content

switchyard_libsy/algorithms/
subagent.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Delegated sub-agent routing around an arbitrary parent algorithm.
5
6use 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
19/// Runtime components for delegated sub-agent routing.
20pub struct SubagentRouterConfig {
21    /// Classifier invoked for delegated work according to `classify_trigger`.
22    pub classifier: Arc<dyn Classifier<State>>,
23    /// Child model category used when `classifier` abstains.
24    pub default_target: Category,
25    /// Controls whether each child is classified once or on every request.
26    pub classify_trigger: ClassifyTrigger,
27    /// Unsupported for child routing because child identity must come from harness metadata.
28    pub message_hash_fallback: bool,
29}
30
31impl SubagentRouterConfig {
32    /// Routes all delegated work to the first model in the sub-agent `Any` category.
33    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
43/// Routes delegated work independently while preserving the parent algorithm for other traffic.
44pub struct SubagentRouter {
45    parent: Arc<dyn Algorithm>,
46    subagent: FallThrough<State>,
47}
48
49impl SubagentRouter {
50    /// Wraps `parent` with the configured delegated-work route.
51    ///
52    /// # Errors
53    ///
54    /// Returns an error when the affinity settings cannot identify delegated children safely.
55    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            // Delegated work routes over the sub-agent's own models, never the parent's.
99            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        // The parent and its children route over separate model groups.
195        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}