switchyard_libsy/algorithms/util/
subagent.rs1use std::sync::Arc;
19
20use async_trait::async_trait;
21
22use crate::Result;
23use crate::core::algorithm::Driver;
24use crate::core::classifier::{Classification, Classifier, Score};
25use switchyard_protocol::{
26 ContentBlock, LlmRequest, Message, Metadata, ModelId, Request, Response, Role,
27};
28
29pub struct SubagentGate<S> {
34 inner: Arc<dyn Classifier<S>>,
35}
36
37impl<S> SubagentGate<S> {
38 pub fn new(inner: Arc<dyn Classifier<S>>) -> Self {
40 Self { inner }
41 }
42}
43
44fn delegated_prompt_request(request: &Request) -> Option<Request> {
46 let prompt = request
48 .llm_request
49 .messages
50 .iter()
51 .rev()
52 .find(|message| message.role == Role::User)?
53 .content
54 .iter()
55 .rev()
56 .find_map(|block| match block {
57 ContentBlock::Text { text } if !text.trim().is_empty() => Some(text.clone()),
58 _ => None,
59 })?;
60
61 Some(Request {
62 llm_request: LlmRequest {
63 model: request.llm_request.model.clone(),
64 messages: vec![Message::text(Role::User, prompt)],
65 ..LlmRequest::default()
66 },
67 raw_request: None,
68 metadata: request.metadata.clone(),
69 })
70}
71
72#[async_trait]
73impl<S> Classifier<S> for SubagentGate<S>
74where
75 S: Send + 'static,
76{
77 fn routing_tier(&self, selected_model_id: &ModelId) -> Option<&'static str> {
78 self.inner.routing_tier(selected_model_id)
79 }
80
81 async fn score(
82 &self,
83 state: &mut S,
84 request: &mut Request,
85 driver: Option<&Driver>,
86 ) -> Result<(Classification, Option<Response>)> {
87 if !request
88 .metadata
89 .as_ref()
90 .is_some_and(Metadata::is_subagent_work)
91 {
92 return Ok((Classification::Scores(Vec::new()), None));
93 }
94 let Some(mut classifier_request) = delegated_prompt_request(request) else {
95 return Ok((Classification::Scores(Vec::new()), None));
96 };
97 self.inner
98 .score(state, &mut classifier_request, driver)
99 .await
100 }
101}
102
103pub struct SubagentOverride {
105 worker: ModelId,
107}
108
109impl SubagentOverride {
110 pub fn new(worker: impl Into<ModelId>) -> Self {
115 Self {
116 worker: worker.into(),
117 }
118 }
119}
120
121#[async_trait]
122impl<S> Classifier<S> for SubagentOverride
123where
124 S: Send + 'static,
125{
126 async fn score(
127 &self,
128 _state: &mut S,
129 request: &mut Request,
130 _driver: Option<&Driver>,
131 ) -> Result<(Classification, Option<Response>)> {
132 let is_delegated_work = request
135 .metadata
136 .as_ref()
137 .is_some_and(Metadata::is_subagent_work);
138 Ok((
139 Classification::Scores(if is_delegated_work {
140 vec![Score {
141 confidence: 1.0,
142 target: self.worker.clone(),
143 }]
144 } else {
145 Vec::new()
146 }),
147 None,
148 ))
149 }
150}
151
152#[cfg(test)]
153mod tests {
154 use super::*;
155 use parking_lot::Mutex;
156 use switchyard_protocol::{slice_to_header_map, text_request};
157
158 #[derive(Default)]
159 struct CapturingClassifier {
160 requests: Mutex<Vec<Request>>,
161 }
162
163 #[async_trait]
164 impl Classifier<()> for CapturingClassifier {
165 async fn score(
166 &self,
167 _state: &mut (),
168 request: &mut Request,
169 _driver: Option<&Driver>,
170 ) -> Result<(Classification, Option<Response>)> {
171 self.requests.lock().push(request.clone());
172 Ok((
173 Classification::Scores(vec![Score {
174 confidence: 1.0,
175 target: ModelId::from("worker"),
176 }]),
177 None,
178 ))
179 }
180 }
181
182 fn request(headers: &[(&str, &str)]) -> Request {
183 let metadata =
184 (!headers.is_empty()).then(|| Metadata::from_headers(&slice_to_header_map(headers)));
185 Request {
186 llm_request: text_request(Some(ModelId::from("auto").to_string()), "hi"),
187 raw_request: None,
188 metadata,
189 }
190 }
191
192 async fn selected(headers: &[(&str, &str)]) -> Result<Option<ModelId>> {
194 let mut state = ();
195 let classification = SubagentOverride::new("worker")
196 .score(&mut state, &mut request(headers), None)
197 .await?;
198 Ok(classification.0.argmax(false)?.map(|score| score.target))
199 }
200
201 #[tokio::test]
202 async fn requests_without_metadata_abstain() -> Result<()> {
203 assert_eq!(selected(&[]).await?, None);
204 Ok(())
205 }
206
207 #[tokio::test]
208 async fn subagent_work_scores_the_worker() -> Result<()> {
209 let claude = &[
211 ("x-claude-code-session-id", "root"),
212 ("x-claude-code-agent-id", "child-1"),
213 ];
214 assert_eq!(selected(claude).await?, Some(ModelId::from("worker")));
215
216 assert_eq!(
218 selected(&[("x-openai-subagent", "review")]).await?,
219 Some(ModelId::from("worker"))
220 );
221 assert_eq!(
222 selected(&[("x-openai-subagent", "collab_spawn")]).await?,
223 Some(ModelId::from("worker"))
224 );
225 Ok(())
226 }
227
228 #[tokio::test]
229 async fn harness_maintenance_turns_abstain() -> Result<()> {
230 assert_eq!(selected(&[("x-openai-subagent", "compact")]).await?, None);
231 assert_eq!(
232 selected(&[("x-switchyard-is-subagent", "false")]).await?,
233 None
234 );
235 Ok(())
236 }
237
238 #[tokio::test]
239 async fn delegated_work_is_scored_definitively() -> Result<()> {
240 let mut state = ();
243 let classification = SubagentOverride::new("worker")
244 .score(
245 &mut state,
246 &mut request(&[("x-openai-subagent", "review")]),
247 None,
248 )
249 .await?;
250 match classification.0 {
251 Classification::Scores(scores) => {
252 assert_eq!(scores.len(), 1);
253 assert_eq!(scores[0].confidence, 1.0);
254 }
255 Classification::Ambiguous(_) => panic!("override must score definitively"),
256 }
257 Ok(())
258 }
259
260 #[tokio::test]
261 async fn gate_abstains_when_delegated_work_has_no_text_prompt() -> Result<()> {
262 let classifier = Arc::new(CapturingClassifier::default());
263 let gate = SubagentGate::new(classifier.clone());
264 let mut request = request(&[("x-openai-subagent", "collab_spawn")]);
265 request.llm_request.messages = vec![Message::text(Role::Assistant, "no user prompt")];
266
267 let mut state = ();
268 let (classification, response) = gate.score(&mut state, &mut request, None).await?;
269
270 assert!(classification.argmax(false)?.is_none());
271 assert!(response.is_none());
272 assert!(classifier.requests.lock().is_empty());
273 Ok(())
274 }
275}