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 "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
115pub 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
211pub 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
279pub 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}