switchyard_libsy/algorithms/util/
subagent.rs1use async_trait::async_trait;
19
20use crate::Result;
21use crate::core::algorithm::Driver;
22use crate::core::classifier::{Classification, Classifier, Score};
23use switchyard_protocol::ModelId;
24use switchyard_protocol::{Metadata, Request, Response};
25
26pub struct SubagentOverride {
28 worker: ModelId,
30}
31
32impl SubagentOverride {
33 pub fn new(worker: impl Into<ModelId>) -> Self {
38 Self {
39 worker: worker.into(),
40 }
41 }
42}
43
44#[async_trait]
45impl<S> Classifier<S> for SubagentOverride
46where
47 S: Send + 'static,
48{
49 async fn score(
50 &self,
51 _state: &mut S,
52 request: &mut Request,
53 _driver: Option<&Driver>,
54 ) -> Result<(Classification, Option<Response>)> {
55 let is_delegated_work = request
58 .metadata
59 .as_ref()
60 .is_some_and(Metadata::is_subagent_work);
61 Ok((
62 Classification::Scores(if is_delegated_work {
63 vec![Score {
64 confidence: 1.0,
65 target: self.worker.clone(),
66 }]
67 } else {
68 Vec::new()
69 }),
70 None,
71 ))
72 }
73}
74
75#[cfg(test)]
76mod tests {
77 use super::*;
78 use switchyard_protocol::{slice_to_header_map, text_request};
79
80 fn request(headers: &[(&str, &str)]) -> Request {
81 let metadata =
82 (!headers.is_empty()).then(|| Metadata::from_headers(&slice_to_header_map(headers)));
83 Request {
84 llm_request: text_request(Some(ModelId::from("auto").to_string()), "hi"),
85 raw_request: None,
86 metadata,
87 }
88 }
89
90 async fn selected(headers: &[(&str, &str)]) -> Result<Option<ModelId>> {
92 let mut state = ();
93 let classification = SubagentOverride::new("worker")
94 .score(&mut state, &mut request(headers), None)
95 .await?;
96 Ok(classification.0.argmax(false)?.map(|score| score.target))
97 }
98
99 #[tokio::test]
100 async fn requests_without_metadata_abstain() -> Result<()> {
101 assert_eq!(selected(&[]).await?, None);
102 Ok(())
103 }
104
105 #[tokio::test]
106 async fn subagent_work_scores_the_worker() -> Result<()> {
107 let claude = &[
109 ("x-claude-code-session-id", "root"),
110 ("x-claude-code-agent-id", "child-1"),
111 ];
112 assert_eq!(selected(claude).await?, Some(ModelId::from("worker")));
113
114 assert_eq!(
116 selected(&[("x-openai-subagent", "review")]).await?,
117 Some(ModelId::from("worker"))
118 );
119 assert_eq!(
120 selected(&[("x-openai-subagent", "collab_spawn")]).await?,
121 Some(ModelId::from("worker"))
122 );
123 Ok(())
124 }
125
126 #[tokio::test]
127 async fn harness_maintenance_turns_abstain() -> Result<()> {
128 assert_eq!(selected(&[("x-openai-subagent", "compact")]).await?, None);
129 assert_eq!(
130 selected(&[("x-switchyard-is-subagent", "false")]).await?,
131 None
132 );
133 Ok(())
134 }
135
136 #[tokio::test]
137 async fn delegated_work_is_scored_definitively() -> Result<()> {
138 let mut state = ();
141 let classification = SubagentOverride::new("worker")
142 .score(
143 &mut state,
144 &mut request(&[("x-openai-subagent", "review")]),
145 None,
146 )
147 .await?;
148 match classification.0 {
149 Classification::Scores(scores) => {
150 assert_eq!(scores.len(), 1);
151 assert_eq!(scores[0].confidence, 1.0);
152 }
153 Classification::Ambiguous(_) => panic!("override must score definitively"),
154 }
155 Ok(())
156 }
157}