switchyard_libsy/algorithms/
model_as_tool.rs1use std::sync::Arc;
7
8use serde_json::json;
9use switchyard_protocol::{
10 AggLlmResponse, ContentBlock, Decision, LlmResponse, Metadata, ModelId, Request, Response,
11 ToolCall, ToolChoice, ToolDefinition, text_request,
12};
13
14use crate::core::algorithm::{Algorithm, Driver};
15use crate::{LibsyError, Result};
16
17pub const GENERATE_MEDIA_TOOL_NAME: &str = "generate_media";
19
20pub struct ModelAsTool {
26 primary_target: ModelId,
27 media_target: ModelId,
28}
29
30impl ModelAsTool {
31 pub fn new(primary_target: impl Into<ModelId>, media_target: impl Into<ModelId>) -> Self {
33 Self {
34 primary_target: primary_target.into(),
35 media_target: media_target.into(),
36 }
37 }
38}
39
40#[async_trait::async_trait]
41impl Algorithm for ModelAsTool {
42 fn name(&self) -> &str {
43 "model_as_tool"
44 }
45
46 async fn route(self: Arc<Self>, driver: Driver, mut request: Request) -> Result<Response> {
47 reject_reserved_tool_collision(&request)?;
48
49 let stream_requested = request.llm_request.stream;
50 let request_metadata = request.metadata.clone();
51 request.llm_request.stream = false;
53 request
54 .llm_request
55 .extensions
56 .fields
57 .remove("stream_options");
58 request
59 .llm_request
60 .extensions
61 .fields
62 .insert("parallel_tool_calls".to_string(), json!(false));
63 request.llm_request.tools.push(media_tool());
64 request
65 .llm_request
66 .tool_choice
67 .get_or_insert(ToolChoice::Auto);
68 request.llm_request.preservation.requests.clear();
70
71 tracing::info!(
72 target = %self.primary_target,
73 tool = GENERATE_MEDIA_TOOL_NAME,
74 "offering specialized model as a tool"
75 );
76 let primary_decision = Decision::new(self.primary_target.clone(), true);
77 driver.decide(primary_decision.clone()).await?;
78 let primary_response = driver.call_model(request, primary_decision).await?;
79 let Response {
80 llm_response,
81 metadata,
82 } = primary_response;
83 let aggregate = llm_response
84 .into_agg()
85 .await
86 .map_err(|error| LibsyError::external("inspecting model-as-tool response", error))?;
87
88 let Some(tool_call) = find_media_tool_call(&aggregate)? else {
89 return Ok(response_from_aggregate(
90 aggregate,
91 metadata,
92 stream_requested,
93 ));
94 };
95 let prompt = media_prompt(tool_call)?;
96 tracing::info!(
97 target = %self.media_target,
98 tool = GENERATE_MEDIA_TOOL_NAME,
99 "dispatching selected specialized model tool"
100 );
101
102 let media_request = Request {
103 llm_request: text_request(None, prompt),
104 raw_request: None,
105 metadata: request_metadata,
106 };
107 let media_decision = Decision::new(self.media_target.clone(), true);
108 driver.decide(media_decision.clone()).await?;
109 let media_response = driver.call_model(media_request, media_decision).await?;
110 if !stream_requested {
111 return Ok(media_response);
112 }
113
114 let Response {
115 llm_response,
116 metadata,
117 } = media_response;
118 let aggregate = llm_response
119 .into_agg()
120 .await
121 .map_err(|error| LibsyError::external("streaming model-as-tool response", error))?;
122 Ok(response_from_aggregate(aggregate, metadata, true))
123 }
124}
125
126fn reject_reserved_tool_collision(request: &Request) -> Result<()> {
127 if request
128 .llm_request
129 .tools
130 .iter()
131 .any(|tool| tool.name == GENERATE_MEDIA_TOOL_NAME)
132 {
133 return Err(LibsyError::AlgorithmError {
134 message: format!("request already defines reserved tool {GENERATE_MEDIA_TOOL_NAME:?}"),
135 });
136 }
137 Ok(())
138}
139
140fn media_tool() -> ToolDefinition {
141 ToolDefinition {
142 name: GENERATE_MEDIA_TOOL_NAME.to_string(),
143 description: Some(
144 "Generate an image with a specialized local visual model. Call this tool alone only when visual output materially improves the answer. Supply a self-contained prompt describing the scene, composition, and style."
145 .to_string(),
146 ),
147 parameters: json!({
148 "type": "object",
149 "properties": {
150 "prompt": {
151 "type": "string",
152 "description": "A complete image generation prompt."
153 }
154 },
155 "required": ["prompt"],
156 "additionalProperties": false
157 }),
158 strict: Some(true),
159 }
160}
161
162fn find_media_tool_call(response: &AggLlmResponse) -> Result<Option<&ToolCall>> {
163 let tool_calls = response
164 .outputs
165 .iter()
166 .flat_map(|output| &output.content)
167 .filter_map(|block| match block {
168 ContentBlock::ToolCall(tool_call) => Some(tool_call),
169 _ => None,
170 })
171 .collect::<Vec<_>>();
172 let media_call = tool_calls
173 .iter()
174 .copied()
175 .find(|tool_call| tool_call.name == GENERATE_MEDIA_TOOL_NAME);
176 if media_call.is_some() && tool_calls.len() != 1 {
177 return Err(LibsyError::AlgorithmError {
178 message: format!(
179 "{GENERATE_MEDIA_TOOL_NAME} must be called alone; primary model returned {} tool calls",
180 tool_calls.len()
181 ),
182 });
183 }
184 Ok(media_call)
185}
186
187fn media_prompt(tool_call: &ToolCall) -> Result<String> {
188 tool_call
189 .arguments
190 .get("prompt")
191 .and_then(serde_json::Value::as_str)
192 .map(str::trim)
193 .filter(|prompt| !prompt.is_empty())
194 .map(str::to_string)
195 .ok_or_else(|| LibsyError::AlgorithmError {
196 message: format!(
197 "{GENERATE_MEDIA_TOOL_NAME} call {} requires a non-empty string prompt",
198 tool_call.id
199 ),
200 })
201}
202
203fn response_from_aggregate(
204 aggregate: AggLlmResponse,
205 metadata: Option<Metadata>,
206 stream: bool,
207) -> Response {
208 Response {
209 llm_response: if stream {
210 LlmResponse::Stream(aggregate.into_stream())
211 } else {
212 LlmResponse::Agg(aggregate)
213 },
214 metadata,
215 }
216}
217
218#[cfg(test)]
219mod tests {
220 use std::sync::{Arc, Mutex};
221
222 use serde_json::json;
223 use switchyard_protocol::{
224 AggLlmResponse, ContentBlock, LlmRequest, LlmResponse, Message, Request, Response,
225 ResponseOutput, Role, StopReason, ToolCall, ToolChoice, completion_text, text_response,
226 };
227
228 use super::{GENERATE_MEDIA_TOOL_NAME, ModelAsTool};
229 use crate::core::algorithm::Algorithm;
230 use crate::core::testing::test_drive;
231
232 fn request(stream: bool) -> Request {
233 Request {
234 llm_request: LlmRequest {
235 model: Some("auto".to_string()),
236 messages: vec![Message::text(Role::User, "Make a cinematic launch image")],
237 stream,
238 ..LlmRequest::default()
239 },
240 raw_request: None,
241 metadata: None,
242 }
243 }
244
245 fn tool_call_response(prompt: serde_json::Value) -> Response {
246 Response {
247 llm_response: LlmResponse::Agg(AggLlmResponse {
248 outputs: vec![ResponseOutput {
249 role: Role::Assistant,
250 content: vec![ContentBlock::ToolCall(ToolCall {
251 id: "media-1".to_string(),
252 name: GENERATE_MEDIA_TOOL_NAME.to_string(),
253 arguments: json!({"prompt": prompt}),
254 })],
255 stop_reason: Some(StopReason::ToolUse),
256 }],
257 ..AggLlmResponse::default()
258 }),
259 metadata: None,
260 }
261 }
262
263 #[tokio::test]
264 async fn normal_answer_passes_through_after_tool_injection() -> crate::Result<()> {
265 let recorded = Arc::new(Mutex::new(None));
266 let captured = Arc::clone(&recorded);
267 let algorithm: Arc<dyn Algorithm> = Arc::new(ModelAsTool::new("primary", "cosmos"));
268 let mut streamed_request = request(true);
269 streamed_request
270 .llm_request
271 .extensions
272 .fields
273 .insert("stream_options".to_string(), json!({"include_usage": true}));
274 let (trace, response) =
275 test_drive(algorithm, streamed_request, move |_decision, request| {
276 let captured = Arc::clone(&captured);
277 async move {
278 *captured.lock().expect("recording lock") = Some(request);
279 Ok(Response {
280 llm_response: LlmResponse::Agg(text_response(
281 Some("primary".to_string()),
282 "plain answer",
283 )),
284 metadata: None,
285 })
286 }
287 })
288 .await?;
289
290 let request = recorded
291 .lock()
292 .expect("recording lock")
293 .take()
294 .expect("primary request");
295 assert_eq!(request.llm_request.tools.len(), 1);
296 assert_eq!(request.llm_request.tools[0].name, GENERATE_MEDIA_TOOL_NAME);
297 assert_eq!(request.llm_request.tool_choice, Some(ToolChoice::Auto));
298 assert!(!request.llm_request.stream);
299 assert!(
300 !request
301 .llm_request
302 .extensions
303 .fields
304 .contains_key("stream_options")
305 );
306 assert_eq!(
307 request
308 .llm_request
309 .extensions
310 .fields
311 .get("parallel_tool_calls"),
312 Some(&json!(false))
313 );
314 assert!(request.llm_request.preservation.requests.is_empty());
315 assert_eq!(trace.len(), 1);
316 assert!(trace[0].is_answer_call());
317 let aggregate = response
318 .llm_response
319 .into_agg()
320 .await
321 .map_err(|error| crate::LibsyError::external("aggregating test response", error))?;
322 assert_eq!(completion_text(&aggregate), "plain answer");
323 Ok(())
324 }
325
326 #[tokio::test]
327 async fn selected_tool_dispatches_prompt_to_media_target() -> crate::Result<()> {
328 let calls = Arc::new(Mutex::new(Vec::new()));
329 let captured = Arc::clone(&calls);
330 let algorithm: Arc<dyn Algorithm> = Arc::new(ModelAsTool::new("primary", "cosmos"));
331 let (trace, response) = test_drive(
332 algorithm,
333 request(false),
334 move |decision: switchyard_protocol::Decision, request: Request| {
335 let captured = Arc::clone(&captured);
336 async move {
337 captured
338 .lock()
339 .expect("recording lock")
340 .push((decision.selected_model_id().to_string(), request.clone()));
341 match decision.selected_model_id().as_str() {
342 "primary" => Ok(tool_call_response(json!("A chrome robot in rain"))),
343 "cosmos" => Ok(Response {
344 llm_response: LlmResponse::Agg(text_response(
345 Some("cosmos".to_string()),
346 "Image: output.png",
347 )),
348 metadata: None,
349 }),
350 other => panic!("unexpected target {other}"),
351 }
352 }
353 },
354 )
355 .await?;
356
357 assert_eq!(
358 trace
359 .iter()
360 .map(|decision| decision.selected_model_id())
361 .collect::<Vec<_>>(),
362 ["primary", "cosmos"]
363 );
364 assert!(trace[0].is_answer_call());
365 assert!(trace[1].is_answer_call());
366 let media_prompt = {
367 let calls = calls.lock().expect("recording lock");
368 switchyard_protocol::prompt_text(&calls[1].1.llm_request)
369 };
370 assert_eq!(media_prompt, "A chrome robot in rain");
371 let aggregate = response
372 .llm_response
373 .into_agg()
374 .await
375 .map_err(|error| crate::LibsyError::external("aggregating test response", error))?;
376 assert!(completion_text(&aggregate).contains("output.png"));
377 Ok(())
378 }
379
380 #[tokio::test]
381 async fn rejects_reserved_tool_collision() {
382 let mut request = request(false);
383 request.llm_request.tools.push(super::media_tool());
384 let algorithm: Arc<dyn Algorithm> = Arc::new(ModelAsTool::new("primary", "cosmos"));
385 let result = test_drive(algorithm, request, |_decision, _request| async move {
386 unreachable!("collision must fail before a model call")
387 })
388 .await;
389
390 assert!(matches!(
391 result,
392 Err(crate::LibsyError::AlgorithmError { message }) if message.contains("reserved tool")
393 ));
394 }
395
396 #[tokio::test]
397 async fn rejects_empty_prompt_and_parallel_media_call() {
398 for response in [tool_call_response(json!(" ")), {
399 let mut response = tool_call_response(json!("A robot"));
400 let LlmResponse::Agg(aggregate) = &mut response.llm_response else {
401 unreachable!()
402 };
403 aggregate.outputs[0]
404 .content
405 .push(ContentBlock::ToolCall(ToolCall {
406 id: "shell-1".to_string(),
407 name: "shell".to_string(),
408 arguments: json!({"command": "pwd"}),
409 }));
410 response
411 }] {
412 let response = Mutex::new(Some(response));
413 let algorithm: Arc<dyn Algorithm> = Arc::new(ModelAsTool::new("primary", "cosmos"));
414 let result = test_drive(algorithm, request(false), move |_decision, _request| {
415 let response = response.lock().expect("response lock").take();
416 async move { Ok(response.expect("one primary call")) }
417 })
418 .await;
419 assert!(matches!(
420 result,
421 Err(crate::LibsyError::AlgorithmError { .. })
422 ));
423 }
424 }
425}