Skip to main content

apollo/tools/
media.rs

1//! Media generation tools backed by `rs_ai` — image, text-to-speech and
2//! speech-to-text.
3
4use std::path::{Path, PathBuf};
5
6use async_trait::async_trait;
7use base64::Engine as _;
8use serde::Deserialize;
9
10use crate::redaction;
11use crate::tools::{Tool, ToolResult, ToolSpec};
12
13fn resolve_api_key(
14    provider: &str,
15    arg_key: Option<String>,
16    default: &str,
17) -> anyhow::Result<String> {
18    if let Some(k) = arg_key {
19        if !k.is_empty() {
20            return Ok(k);
21        }
22    }
23    let env_var = match provider {
24        "chatgpt" | "openai" => "OPENAI_API_KEY",
25        "gemini" => "GOOGLE_API_KEY",
26        "xai" | "grok" => "XAI_API_KEY",
27        "claude" => "ANTHROPIC_API_KEY",
28        _ => "OPENAI_API_KEY",
29    };
30    if let Ok(k) = std::env::var(env_var) {
31        if !k.is_empty() {
32            return Ok(k);
33        }
34    }
35    if !default.is_empty() {
36        return Ok(default.to_string());
37    }
38    anyhow::bail!("no API key available for provider {provider}")
39}
40
41fn default_image_model(provider: &str) -> &str {
42    match provider {
43        "chatgpt" | "openai" => "dall-e-3",
44        "gemini" => "imagen-3.0-generate-002",
45        "xai" | "grok" => "grok-2a-image",
46        _ => "dall-e-3",
47    }
48}
49
50fn default_tts_model(provider: &str) -> &str {
51    match provider {
52        "chatgpt" | "openai" => "tts-1",
53        "gemini" => "gemini-2.5-flash",
54        "xai" | "grok" => "grok-4.20-realtime",
55        _ => "tts-1",
56    }
57}
58
59fn default_stt_model(provider: &str) -> &str {
60    match provider {
61        "chatgpt" | "openai" => "whisper-1",
62        "gemini" => "gemini-2.5-flash",
63        "xai" | "grok" => "grok-4.20-realtime",
64        _ => "whisper-1",
65    }
66}
67
68fn normalize_provider(p: &str) -> &str {
69    match p {
70        "openai" => "chatgpt",
71        "grok" => "xai",
72        other => other,
73    }
74}
75
76fn build_client(
77    provider: &str,
78    api_key: &str,
79    model: &str,
80) -> anyhow::Result<rs_ai::ClientBuilder> {
81    let builder = match normalize_provider(provider) {
82        "chatgpt" => rs_ai::chatgpt(),
83        "gemini" => rs_ai::gemini(),
84        "xai" => rs_ai::xai(),
85        "claude" => rs_ai::claude(),
86        _ => rs_ai::chatgpt(),
87    };
88    Ok(builder.api_key(api_key).model(model))
89}
90
91async fn save_generated_file(
92    workspace: &Path,
93    bytes: Vec<u8>,
94    media_type: &str,
95) -> anyhow::Result<String> {
96    let ext = match media_type {
97        "image/png" => "png",
98        "image/jpeg" | "image/jpg" => "jpg",
99        "image/webp" => "webp",
100        "image/gif" => "gif",
101        "audio/mp3" | "audio/mpeg" => "mp3",
102        "audio/wav" => "wav",
103        "audio/ogg" => "ogg",
104        "audio/webm" => "webm",
105        _ => "bin",
106    };
107    let dir = workspace.join(".apollo").join("generated");
108    tokio::fs::create_dir_all(&dir).await?;
109    let name = format!("{}.{}", uuid::Uuid::new_v4(), ext);
110    let path = dir.join(name);
111    tokio::fs::write(&path, bytes).await?;
112    Ok(path.to_string_lossy().to_string())
113}
114
115/// Generate an image from a text prompt.
116pub struct ImageGenerationTool {
117    workspace: PathBuf,
118    default_api_key: String,
119    default_provider: String,
120}
121
122impl ImageGenerationTool {
123    pub fn new(workspace: PathBuf, default_api_key: String, default_provider: String) -> Self {
124        Self {
125            workspace,
126            default_api_key,
127            default_provider,
128        }
129    }
130}
131
132#[derive(Deserialize)]
133struct GenerateImageArgs {
134    prompt: String,
135    #[serde(default)]
136    provider: Option<String>,
137    #[serde(default)]
138    model: Option<String>,
139    #[serde(default)]
140    size: Option<String>,
141    #[serde(default)]
142    n: Option<u32>,
143    #[serde(default)]
144    aspect_ratio: Option<String>,
145    #[serde(default)]
146    api_key: Option<String>,
147}
148
149#[async_trait]
150impl Tool for ImageGenerationTool {
151    fn name(&self) -> &str {
152        "generate_image"
153    }
154
155    fn spec(&self) -> ToolSpec {
156        ToolSpec {
157            name: "generate_image".to_string(),
158            description: "Generate an image from a text prompt using ChatGPT DALL-E, Gemini Imagen, or xAI Grok Imagine.".to_string(),
159            parameters: serde_json::json!({
160                "type": "object",
161                "required": ["prompt"],
162                "properties": {
163                    "prompt": { "type": "string", "description": "image description" },
164                    "provider": { "type": "string", "enum": ["chatgpt", "gemini", "xai"], "description": "provider to use" },
165                    "model": { "type": "string", "description": "specific model id" },
166                    "size": { "type": "string", "description": "size like 1024x1024" },
167                    "n": { "type": "integer", "description": "number of images" },
168                    "aspect_ratio": { "type": "string", "description": "aspect ratio like 16:9" },
169                    "api_key": { "type": "string", "description": "optional api key override" }
170                }
171            }),
172        }
173    }
174
175    async fn execute(&self, arguments: &str) -> anyhow::Result<ToolResult> {
176        let args: GenerateImageArgs = serde_json::from_str(arguments)?;
177        let provider = args.provider.as_deref().unwrap_or(&self.default_provider);
178        let model = args
179            .model
180            .as_deref()
181            .unwrap_or(default_image_model(provider));
182        let api_key = resolve_api_key(provider, args.api_key, &self.default_api_key)?;
183
184        let mut options = rs_ai_core::ImageGenerationOptions::default();
185        if let Some(n) = args.n {
186            options.n = Some(n);
187        }
188        if let Some(size) = args.size {
189            options.size = Some(size);
190        }
191        if let Some(ar) = args.aspect_ratio {
192            options.aspect_ratio = Some(ar);
193        }
194
195        let result = build_client(provider, &api_key, model)?
196            .generate_image(&args.prompt, options)
197            .await
198            .map_err(|e| anyhow::anyhow!("image generation failed: {e}"))?;
199
200        let file = result.image;
201        let path = save_generated_file(&self.workspace, file.bytes, &file.media_type).await?;
202        let output = serde_json::json!({
203            "file_path": redaction::redact_text(&path),
204            "media_type": file.media_type,
205            "images_generated": result.images.len().max(1),
206        });
207        Ok(ToolResult::success(output.to_string()))
208    }
209}
210
211/// Synthesize speech from text.
212pub struct TextToSpeechTool {
213    workspace: PathBuf,
214    default_api_key: String,
215    default_provider: String,
216}
217
218impl TextToSpeechTool {
219    pub fn new(workspace: PathBuf, default_api_key: String, default_provider: String) -> Self {
220        Self {
221            workspace,
222            default_api_key,
223            default_provider,
224        }
225    }
226}
227
228#[derive(Deserialize)]
229struct TtsArgs {
230    text: String,
231    #[serde(default)]
232    provider: Option<String>,
233    #[serde(default)]
234    model: Option<String>,
235    #[serde(default)]
236    api_key: Option<String>,
237}
238
239#[async_trait]
240impl Tool for TextToSpeechTool {
241    fn name(&self) -> &str {
242        "text_to_speech"
243    }
244
245    fn spec(&self) -> ToolSpec {
246        ToolSpec {
247            name: "text_to_speech".to_string(),
248            description: "Synthesize speech from text and return an audio file path.".to_string(),
249            parameters: serde_json::json!({
250                "type": "object",
251                "required": ["text"],
252                "properties": {
253                    "text": { "type": "string" },
254                    "provider": { "type": "string", "enum": ["chatgpt", "gemini", "xai"], "description": "provider to use" },
255                    "model": { "type": "string", "description": "model id like tts-1" },
256                    "api_key": { "type": "string" }
257                }
258            }),
259        }
260    }
261
262    async fn execute(&self, arguments: &str) -> anyhow::Result<ToolResult> {
263        let args: TtsArgs = serde_json::from_str(arguments)?;
264        let provider = args.provider.as_deref().unwrap_or(&self.default_provider);
265        let model = args.model.as_deref().unwrap_or(default_tts_model(provider));
266        let api_key = resolve_api_key(provider, args.api_key, &self.default_api_key)?;
267
268        let bytes = build_client(provider, &api_key, model)?
269            .speak(&args.text)
270            .await
271            .map_err(|e| anyhow::anyhow!("text-to-speech failed: {e}"))?;
272
273        let path = save_generated_file(&self.workspace, bytes, "audio/mp3").await?;
274        let output = serde_json::json!({"file_path": redaction::redact_text(&path)});
275        Ok(ToolResult::success(output.to_string()))
276    }
277}
278
279/// Transcribe audio to text.
280pub struct SpeechToTextTool {
281    default_api_key: String,
282    default_provider: String,
283}
284
285impl SpeechToTextTool {
286    pub fn new(default_api_key: String, default_provider: String) -> Self {
287        Self {
288            default_api_key,
289            default_provider,
290        }
291    }
292}
293
294#[derive(Deserialize)]
295struct SttArgs {
296    #[serde(default)]
297    audio_path: Option<String>,
298    #[serde(default)]
299    audio_base64: Option<String>,
300    #[serde(default)]
301    mime_type: Option<String>,
302    #[serde(default)]
303    provider: Option<String>,
304    #[serde(default)]
305    model: Option<String>,
306    #[serde(default)]
307    api_key: Option<String>,
308}
309
310#[async_trait]
311impl Tool for SpeechToTextTool {
312    fn name(&self) -> &str {
313        "speech_to_text"
314    }
315
316    fn spec(&self) -> ToolSpec {
317        ToolSpec {
318            name: "speech_to_text".to_string(),
319            description: "Transcribe audio from a file path or base64 string to text.".to_string(),
320            parameters: serde_json::json!({
321                "type": "object",
322                "properties": {
323                    "audio_path": { "type": "string", "description": "path to audio file" },
324                    "audio_base64": { "type": "string", "description": "base64-encoded audio bytes" },
325                    "mime_type": { "type": "string", "description": "audio mime type like audio/mp3" },
326                    "provider": { "type": "string", "enum": ["chatgpt", "gemini", "xai"], "description": "provider to use" },
327                    "model": { "type": "string", "description": "model id like whisper-1" },
328                    "api_key": { "type": "string" }
329                }
330            }),
331        }
332    }
333
334    async fn execute(&self, arguments: &str) -> anyhow::Result<ToolResult> {
335        let args: SttArgs = serde_json::from_str(arguments)?;
336        let provider = args.provider.as_deref().unwrap_or(&self.default_provider);
337        let model = args.model.as_deref().unwrap_or(default_stt_model(provider));
338        let api_key = resolve_api_key(provider, args.api_key, &self.default_api_key)?;
339
340        let audio = if let Some(path) = args.audio_path {
341            tokio::fs::read(&path).await?
342        } else if let Some(b64) = args.audio_base64 {
343            base64::engine::general_purpose::STANDARD
344                .decode(b64.trim())
345                .map_err(|e| anyhow::anyhow!("invalid base64 audio: {e}"))?
346        } else {
347            anyhow::bail!("provide either audio_path or audio_base64")
348        };
349
350        let mime_type = args.mime_type.unwrap_or_else(|| "audio/mp3".to_string());
351        let text = build_client(provider, &api_key, model)?
352            .transcribe(audio, &mime_type)
353            .await
354            .map_err(|e| anyhow::anyhow!("speech-to-text failed: {e}"))?;
355
356        Ok(ToolResult::success(text))
357    }
358}