1use std::path::PathBuf;
9use std::sync::Arc;
10
11use base64::Engine;
12use roder_api::media::{
13 FAKE_MEDIA_PROVIDER_ID, GeneratedImage, ImageGenerationAction, ImageGenerationBatch,
14 ImageModelDescriptor, MediaDimensions, MediaGenerationMetadata, MediaGenerationOutput,
15 MediaGenerationRequest, MediaGenerationResponse, MediaGeneratorProvider, MediaImageInput,
16 MediaKind, MediaProviderDescriptor,
17};
18use roder_api::tools::{ToolCall, ToolExecutionContext, ToolExecutor, ToolResult, ToolSpec};
19use serde_json::json;
20
21use crate::media_artifacts::{GeneratedMediaSpec, MediaArtifactStore, default_media_artifact_dir};
22
23#[derive(Debug, Clone)]
26pub struct RuntimeMediaGenerationConfig {
27 pub default_provider: Option<String>,
28 pub default_model: Option<String>,
29 pub max_outputs: u32,
30 pub max_input_images: u32,
31 pub artifacts_dir: Option<PathBuf>,
32 pub max_read_bytes: Option<u64>,
33}
34
35impl Default for RuntimeMediaGenerationConfig {
36 fn default() -> Self {
37 Self {
38 default_provider: None,
39 default_model: None,
40 max_outputs: 4,
41 max_input_images: 16,
42 artifacts_dir: None,
43 max_read_bytes: None,
44 }
45 }
46}
47
48pub struct MediaGenerationService {
49 providers: Vec<Arc<dyn MediaGeneratorProvider>>,
50 config: RuntimeMediaGenerationConfig,
51}
52
53impl MediaGenerationService {
54 pub fn new(
55 mut providers: Vec<Arc<dyn MediaGeneratorProvider>>,
56 config: RuntimeMediaGenerationConfig,
57 ) -> Self {
58 if !providers
59 .iter()
60 .any(|provider| provider.provider_id() == FAKE_MEDIA_PROVIDER_ID)
61 {
62 providers.push(Arc::new(FakeImageProvider));
63 }
64 Self { providers, config }
65 }
66
67 pub fn config(&self) -> &RuntimeMediaGenerationConfig {
68 &self.config
69 }
70
71 pub fn provider_descriptors(&self) -> Vec<MediaProviderDescriptor> {
72 self.providers
73 .iter()
74 .map(|provider| provider.descriptor())
75 .collect()
76 }
77
78 pub fn default_provider_id(&self) -> String {
79 self.config
80 .default_provider
81 .clone()
82 .unwrap_or_else(|| FAKE_MEDIA_PROVIDER_ID.to_string())
83 }
84
85 pub fn store(&self) -> anyhow::Result<MediaArtifactStore> {
86 let root = self
87 .config
88 .artifacts_dir
89 .clone()
90 .or_else(|| std::env::var_os("RODER_MEDIA_ARTIFACT_DIR").map(PathBuf::from))
91 .map(Ok)
92 .unwrap_or_else(default_media_artifact_dir)?;
93 let mut store = MediaArtifactStore::new(root);
94 if let Some(max_read_bytes) = self.config.max_read_bytes {
95 store = store.with_max_read_bytes(max_read_bytes);
96 }
97 Ok(store)
98 }
99
100 pub async fn generate_image(
101 &self,
102 mut request: MediaGenerationRequest,
103 ) -> anyhow::Result<MediaGenerationResponse> {
104 if request.prompt.trim().is_empty() {
105 anyhow::bail!("image generation requires a non-empty prompt");
106 }
107 let count = request.count.unwrap_or(1);
108 if count == 0 || count > self.config.max_outputs {
109 anyhow::bail!(
110 "requested {count} outputs; configured limit is 1..={}",
111 self.config.max_outputs
112 );
113 }
114 let input_count = request.input_artifacts.len() + request.input_images.len();
115 if input_count > self.config.max_input_images as usize {
116 anyhow::bail!(
117 "requested {input_count} input images; configured limit is {}",
118 self.config.max_input_images
119 );
120 }
121 if request.action == Some(ImageGenerationAction::Edit) && input_count == 0 {
122 anyhow::bail!(
123 "the edit action requires at least one input artifact or inline input image"
124 );
125 }
126
127 let provider_id = request
128 .provider
129 .clone()
130 .unwrap_or_else(|| self.default_provider_id());
131 let provider = self
132 .providers
133 .iter()
134 .find(|provider| provider.provider_id() == provider_id)
135 .ok_or_else(|| {
136 anyhow::anyhow!(
137 "image provider {provider_id:?} is not available; installed providers: {}",
138 self.provider_ids().join(", ")
139 )
140 })?;
141
142 if request.model.is_none() && provider_id == self.default_provider_id() {
143 request.model = self.config.default_model.clone();
144 }
145
146 let store = self.store()?;
147 self.resolve_input_artifacts(&store, &mut request)?;
148 request.provider = Some(provider_id);
149
150 let prompt = request.prompt.clone();
151 let batch = provider.generate_image(request).await?;
152 if batch.images.is_empty() {
153 if batch.output_errors.is_empty() {
154 anyhow::bail!(
155 "image provider {} returned no images",
156 provider.provider_id()
157 );
158 }
159 anyhow::bail!(
160 "image provider {} generated no images: {}",
161 provider.provider_id(),
162 batch.output_errors.join("; ")
163 );
164 }
165 self.persist_batch(&store, &prompt, batch)
166 }
167
168 fn provider_ids(&self) -> Vec<String> {
169 self.providers
170 .iter()
171 .map(|provider| provider.provider_id().to_string())
172 .collect()
173 }
174
175 fn resolve_input_artifacts(
178 &self,
179 store: &MediaArtifactStore,
180 request: &mut MediaGenerationRequest,
181 ) -> anyhow::Result<()> {
182 for artifact_id in std::mem::take(&mut request.input_artifacts) {
183 let (artifact, bytes) = store.read(&artifact_id, None).map_err(|error| {
184 anyhow::anyhow!("could not read input artifact {artifact_id}: {error}")
185 })?;
186 if artifact.kind != MediaKind::Image {
187 anyhow::bail!("input artifact {artifact_id} is not an image");
188 }
189 request.input_images.push(MediaImageInput {
190 bytes_base64: base64::engine::general_purpose::STANDARD.encode(bytes),
191 mime_type: artifact.mime_type,
192 });
193 }
194 Ok(())
195 }
196
197 fn persist_batch(
198 &self,
199 store: &MediaArtifactStore,
200 prompt: &str,
201 batch: ImageGenerationBatch,
202 ) -> anyhow::Result<MediaGenerationResponse> {
203 let mut outputs = Vec::with_capacity(batch.images.len());
204 let mut response_revised_prompt = None;
205 let mut response_watermark = None;
206 let mut response_safety = None;
207 for image in &batch.images {
208 let bytes = base64::engine::general_purpose::STANDARD
209 .decode(&image.bytes_base64)
210 .map_err(|error| {
211 anyhow::anyhow!(
212 "image provider {} returned invalid base64 output: {error}",
213 batch.provider
214 )
215 })?;
216 let generation = MediaGenerationMetadata {
217 provider: batch.provider.clone(),
218 model: Some(batch.model.clone()),
219 revised_prompt: image.revised_prompt.clone(),
220 watermark: image.watermark.clone(),
221 safety: image.safety.clone(),
222 provider_response_id: batch.provider_response_id.clone(),
223 };
224 let (artifact, preview) = store.write_generated(&GeneratedMediaSpec {
225 prompt,
226 kind: MediaKind::Image,
227 mime_type: &image.mime_type,
228 provider: &batch.provider,
229 bytes: &bytes,
230 dimensions: image.dimensions.clone(),
231 duration_millis: None,
232 generation: Some(generation),
233 })?;
234 response_revised_prompt = response_revised_prompt.or(image.revised_prompt.clone());
235 response_watermark = response_watermark.or(image.watermark.clone());
236 response_safety = response_safety.or(image.safety.clone());
237 outputs.push(MediaGenerationOutput {
238 artifact,
239 preview,
240 revised_prompt: image.revised_prompt.clone(),
241 });
242 }
243 Ok(MediaGenerationResponse {
244 provider: batch.provider,
245 model: Some(batch.model),
246 outputs,
247 revised_prompt: response_revised_prompt,
248 provider_response_id: batch.provider_response_id,
249 usage: batch.usage,
250 watermark: response_watermark,
251 safety: response_safety,
252 output_errors: batch.output_errors,
253 })
254 }
255}
256
257pub struct FakeImageProvider;
260
261pub const FAKE_IMAGE_PNG_BASE64: &str = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg==";
263
264#[async_trait::async_trait]
265impl MediaGeneratorProvider for FakeImageProvider {
266 fn provider_id(&self) -> &str {
267 FAKE_MEDIA_PROVIDER_ID
268 }
269
270 fn descriptor(&self) -> MediaProviderDescriptor {
271 MediaProviderDescriptor {
272 id: FAKE_MEDIA_PROVIDER_ID.to_string(),
273 display_name: "Fake Media (offline)".to_string(),
274 supports_images: true,
275 supports_videos: false,
276 configured: true,
277 default_model: Some("fake-image".to_string()),
278 image_models: vec![ImageModelDescriptor {
279 id: "fake-image".to_string(),
280 display_name: "Fake Image".to_string(),
281 provider: FAKE_MEDIA_PROVIDER_ID.to_string(),
282 is_default: true,
283 legacy: false,
284 supports_edit: true,
285 supports_multiple_outputs: true,
286 supported_aspect_ratios: Vec::new(),
287 supported_sizes: Vec::new(),
288 supported_image_sizes: Vec::new(),
289 supports_transparent_background: false,
290 supports_partial_images: false,
291 }],
292 }
293 }
294
295 async fn generate_image(
296 &self,
297 request: MediaGenerationRequest,
298 ) -> anyhow::Result<ImageGenerationBatch> {
299 let count = request.count.unwrap_or(1).max(1);
300 let model = request.model.unwrap_or_else(|| "fake-image".to_string());
301 let images = (0..count)
302 .map(|_| GeneratedImage {
303 bytes_base64: FAKE_IMAGE_PNG_BASE64.to_string(),
304 mime_type: "image/png".to_string(),
305 dimensions: Some(MediaDimensions {
306 width: 1,
307 height: 1,
308 }),
309 revised_prompt: None,
310 watermark: None,
311 safety: None,
312 })
313 .collect();
314 Ok(ImageGenerationBatch {
315 provider: FAKE_MEDIA_PROVIDER_ID.to_string(),
316 model,
317 images,
318 provider_response_id: None,
319 usage: None,
320 output_errors: Vec::new(),
321 })
322 }
323}
324
325pub struct MediaGenerateImageTool {
328 service: Arc<MediaGenerationService>,
329}
330
331impl MediaGenerateImageTool {
332 pub fn new(service: Arc<MediaGenerationService>) -> Self {
333 Self { service }
334 }
335}
336
337#[async_trait::async_trait]
338impl ToolExecutor for MediaGenerateImageTool {
339 fn spec(&self) -> ToolSpec {
340 ToolSpec {
341 name: "media_generate_image".to_string(),
342 description:
343 "Generates one or more images with the configured image provider and stores them as Roder media artifacts."
344 .to_string(),
345 parameters: json!({
346 "type": "object",
347 "properties": {
348 "prompt": { "type": "string", "description": "Text description of the image(s) to generate." },
349 "provider": { "type": "string", "description": "Image provider id (e.g. openai, google, fake). Defaults to the configured provider." },
350 "model": { "type": "string", "description": "Image model id, e.g. gpt-image-2 or gemini-3.1-flash-image." },
351 "action": { "type": "string", "enum": ["auto", "generate", "edit"] },
352 "inputArtifacts": { "type": "array", "items": { "type": "string" }, "description": "Roder media artifact ids used as reference/edit inputs." },
353 "count": { "type": "integer", "minimum": 1 },
354 "aspectRatio": { "type": "string", "description": "Aspect ratio such as 16:9 (Gemini models)." },
355 "size": { "type": "string", "description": "Pixel size such as 1536x1024 (OpenAI models)." },
356 "imageSize": { "type": "string", "description": "Resolution tier such as 1K, 2K, or 4K (Gemini models)." },
357 "quality": { "type": "string" },
358 "outputFormat": { "type": "string", "enum": ["png", "jpeg", "webp"] },
359 "background": { "type": "string", "enum": ["auto", "transparent", "opaque"] },
360 "outputCompression": { "type": "integer", "minimum": 0, "maximum": 100 },
361 "moderation": { "type": "string" }
362 },
363 "required": ["prompt"],
364 "additionalProperties": false
365 }),
366 }
367 }
368
369 async fn execute(
370 &self,
371 _ctx: ToolExecutionContext,
372 call: ToolCall,
373 ) -> anyhow::Result<ToolResult> {
374 let request: MediaGenerationRequest = serde_json::from_value(call.arguments.clone())
375 .map_err(|error| anyhow::anyhow!("invalid media_generate_image arguments: {error}"))?;
376 let response = self.service.generate_image(request).await?;
377 let artifact_ids: Vec<&str> = response
378 .outputs
379 .iter()
380 .map(|output| output.artifact.id.as_str())
381 .collect();
382 let artifacts: Vec<_> = response
383 .outputs
384 .iter()
385 .map(|output| output.artifact.clone())
386 .collect();
387 let previews: Vec<_> = response
388 .outputs
389 .iter()
390 .map(|output| output.preview.clone())
391 .collect();
392 let text = format!(
393 "generated {} image artifact(s) with {}{}: {}",
394 response.outputs.len(),
395 response.provider,
396 response
397 .model
398 .as_deref()
399 .map(|model| format!("/{model}"))
400 .unwrap_or_default(),
401 artifact_ids.join(", ")
402 );
403 Ok(ToolResult {
404 id: call.id,
405 name: call.name,
406 text,
407 data: json!({
408 "mediaArtifacts": artifacts,
409 "mediaPreviews": previews,
410 "mediaGeneration": response,
411 }),
412 is_error: false,
413 })
414 }
415}
416
417#[cfg(test)]
418mod tests {
419 use super::*;
420 use roder_api::policy_mode::PolicyMode;
421
422 fn temp_config() -> (RuntimeMediaGenerationConfig, PathBuf) {
423 let dir = std::env::temp_dir().join(format!("roder-media-gen-{}", uuid::Uuid::new_v4()));
424 let config = RuntimeMediaGenerationConfig {
425 artifacts_dir: Some(dir.clone()),
426 ..RuntimeMediaGenerationConfig::default()
427 };
428 (config, dir)
429 }
430
431 fn request(prompt: &str) -> MediaGenerationRequest {
432 MediaGenerationRequest {
433 prompt: prompt.to_string(),
434 ..MediaGenerationRequest::default()
435 }
436 }
437
438 #[tokio::test]
439 async fn fake_image_generation_persists_artifacts_in_store() {
440 let (config, dir) = temp_config();
441 let service = MediaGenerationService::new(Vec::new(), config);
442
443 let response = service
444 .generate_image(MediaGenerationRequest {
445 count: Some(2),
446 ..request("two tiny images")
447 })
448 .await
449 .unwrap();
450
451 assert_eq!(response.provider, FAKE_MEDIA_PROVIDER_ID);
452 assert_eq!(response.outputs.len(), 2);
453 for output in &response.outputs {
454 assert!(
455 output
456 .artifact
457 .store_path
458 .starts_with(&*dir.display().to_string())
459 );
460 assert!(output.artifact.roder_owned);
461 assert_eq!(
462 output
463 .artifact
464 .generation
465 .as_ref()
466 .map(|generation| generation.provider.as_str()),
467 Some(FAKE_MEDIA_PROVIDER_ID)
468 );
469 }
470 assert_eq!(service.store().unwrap().list().unwrap().len(), 2);
471 }
472
473 #[tokio::test]
474 async fn generation_limits_and_missing_provider_fail_with_clear_errors() {
475 let (config, _dir) = temp_config();
476 let service = MediaGenerationService::new(Vec::new(), config);
477
478 let empty_prompt = service.generate_image(request(" ")).await.unwrap_err();
479 assert!(empty_prompt.to_string().contains("non-empty prompt"));
480
481 let too_many = service
482 .generate_image(MediaGenerationRequest {
483 count: Some(5),
484 ..request("too many")
485 })
486 .await
487 .unwrap_err();
488 assert!(too_many.to_string().contains("configured limit is 1..=4"));
489
490 let missing_provider = service
491 .generate_image(MediaGenerationRequest {
492 provider: Some("missing".to_string()),
493 ..request("nope")
494 })
495 .await
496 .unwrap_err();
497 assert!(
498 missing_provider
499 .to_string()
500 .contains("image provider \"missing\" is not available")
501 );
502
503 let edit_without_inputs = service
504 .generate_image(MediaGenerationRequest {
505 action: Some(ImageGenerationAction::Edit),
506 ..request("edit nothing")
507 })
508 .await
509 .unwrap_err();
510 assert!(
511 edit_without_inputs
512 .to_string()
513 .contains("edit action requires at least one input")
514 );
515 }
516
517 #[tokio::test]
518 async fn input_artifacts_are_resolved_into_inline_images_for_the_provider() {
519 struct CapturingProvider {
520 captured: std::sync::Mutex<Option<MediaGenerationRequest>>,
521 }
522
523 #[async_trait::async_trait]
524 impl MediaGeneratorProvider for CapturingProvider {
525 fn provider_id(&self) -> &str {
526 "capture"
527 }
528
529 fn descriptor(&self) -> MediaProviderDescriptor {
530 MediaProviderDescriptor {
531 id: "capture".to_string(),
532 display_name: "Capture".to_string(),
533 supports_images: true,
534 configured: true,
535 ..MediaProviderDescriptor::default()
536 }
537 }
538
539 async fn generate_image(
540 &self,
541 request: MediaGenerationRequest,
542 ) -> anyhow::Result<ImageGenerationBatch> {
543 *self.captured.lock().unwrap() = Some(request);
544 Ok(ImageGenerationBatch {
545 provider: "capture".to_string(),
546 model: "capture-image".to_string(),
547 images: vec![GeneratedImage {
548 bytes_base64: FAKE_IMAGE_PNG_BASE64.to_string(),
549 mime_type: "image/png".to_string(),
550 dimensions: None,
551 revised_prompt: Some("revised".to_string()),
552 watermark: Some("synthid".to_string()),
553 safety: None,
554 }],
555 provider_response_id: Some("resp-1".to_string()),
556 usage: None,
557 output_errors: Vec::new(),
558 })
559 }
560 }
561
562 let provider = Arc::new(CapturingProvider {
563 captured: std::sync::Mutex::new(None),
564 });
565 let (config, _dir) = temp_config();
566 let service = MediaGenerationService::new(vec![provider.clone()], config);
567
568 let (seed_artifact, _) = service
569 .store()
570 .unwrap()
571 .write_generated(&GeneratedMediaSpec {
572 prompt: "seed",
573 kind: MediaKind::Image,
574 mime_type: "image/png",
575 provider: "fake",
576 bytes: b"abc",
577 dimensions: None,
578 duration_millis: None,
579 generation: None,
580 })
581 .unwrap();
582
583 let response = service
584 .generate_image(MediaGenerationRequest {
585 provider: Some("capture".to_string()),
586 action: Some(ImageGenerationAction::Edit),
587 input_artifacts: vec![seed_artifact.id.clone()],
588 ..request("edit the seed")
589 })
590 .await
591 .unwrap();
592
593 let captured = provider.captured.lock().unwrap().clone().unwrap();
594 assert!(captured.input_artifacts.is_empty());
595 assert_eq!(captured.input_images.len(), 1);
596 assert_eq!(captured.input_images[0].mime_type, "image/png");
597 assert_eq!(
598 captured.input_images[0].bytes_base64,
599 base64::engine::general_purpose::STANDARD.encode(b"abc")
600 );
601
602 assert_eq!(response.revised_prompt.as_deref(), Some("revised"));
603 assert_eq!(response.watermark.as_deref(), Some("synthid"));
604 assert_eq!(response.provider_response_id.as_deref(), Some("resp-1"));
605 let generation = response.outputs[0].artifact.generation.clone().unwrap();
606 assert_eq!(generation.watermark.as_deref(), Some("synthid"));
607 assert_eq!(generation.provider_response_id.as_deref(), Some("resp-1"));
608 }
609
610 #[tokio::test]
611 async fn media_generate_image_tool_returns_canonical_payload() {
612 let (config, _dir) = temp_config();
613 let service = Arc::new(MediaGenerationService::new(Vec::new(), config));
614 let tool = MediaGenerateImageTool::new(service);
615
616 let result = tool
617 .execute(
618 ToolExecutionContext::new("thread", "turn", PolicyMode::Default),
619 ToolCall {
620 id: "call".to_string(),
621 name: "media_generate_image".to_string(),
622 arguments: json!({ "prompt": "tiny" }),
623 raw_arguments: "{}".to_string(),
624 thread_id: "thread".to_string(),
625 turn_id: "turn".to_string(),
626 },
627 )
628 .await
629 .unwrap();
630
631 assert!(!result.is_error);
632 assert_eq!(result.data["mediaArtifacts"][0]["kind"], "image");
633 assert_eq!(result.data["mediaPreviews"][0]["strategy"], "thumbnail");
634 assert_eq!(result.data["mediaGeneration"]["provider"], "fake");
635 assert!(result.text.contains("generated 1 image artifact(s)"));
636 }
637}