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