Skip to main content

switchyard_libsy/algorithms/
model_as_tool.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Route a model-requested media tool call to a specialized generation model.
5
6use 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
17/// Name of the synthetic tool exposed to the primary model.
18pub const GENERATE_MEDIA_TOOL_NAME: &str = "generate_media";
19
20/// Routes a model-selected media tool call to a specialized generation model.
21///
22/// The primary model receives one additional `generate_media(prompt)` tool. Normal answers pass
23/// through unchanged. A matching tool call becomes a second routed model call whose host client
24/// is responsible for turning the prompt into media.
25pub struct ModelAsTool {
26    primary_target: ModelId,
27    media_target: ModelId,
28}
29
30impl ModelAsTool {
31    /// Creates a router backed by a reasoning model and a specialized media model.
32    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        // The tool choice must be inspected before the algorithm can decide what to return.
52        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        // Same-format replay would encode the preserved request instead of the injected tool.
69        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}