1use 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
113pub 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
209pub 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
277pub 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}