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::{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
18/// Runtime components for delegated sub-agent routing.
19pub struct SubagentRouterConfig {
20    /// Targets the delegated-work classifier may select.
21    pub targets: Vec<ModelId>,
22    /// Classifier invoked for delegated work according to `classify_trigger`.
23    pub classifier: Arc<dyn Classifier<State>>,
24    /// Child target used when `classifier` abstains.
25    pub default_target: ModelId,
26    /// Controls whether each child is classified once or on every request.
27    pub classify_trigger: ClassifyTrigger,
28    /// Unsupported for child routing because child identity must come from harness metadata.
29    pub message_hash_fallback: bool,
30}
31
32impl SubagentRouterConfig {
33    /// Routes delegated work directly to one fixed target.
34    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
46/// Routes delegated work independently while preserving the parent algorithm for other traffic.
47pub struct SubagentRouter {
48    parent: Arc<dyn Algorithm>,
49    subagent: FallThrough<State>,
50}
51
52impl SubagentRouter {
53    /// Wraps `parent` with the configured delegated-work route.
54    ///
55    /// # Errors
56    ///
57    /// Returns an error when the child default is not a child target or when the affinity
58    /// settings cannot identify delegated children safely.
59    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}