Skip to main content

roder_core/
media_generation.rs

1//! Provider-neutral image generation service.
2//!
3//! Resolves the requested or configured [`MediaGeneratorProvider`], enforces
4//! request limits, resolves input artifacts into inline images, persists every
5//! generated output through [`MediaArtifactStore`], and exposes the canonical
6//! `media_generate_image` tool used by the runtime tool pipeline.
7
8use 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/// Runtime configuration for image generation, mapped from
24/// `[media.image_generation]` user config by the host.
25#[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    /// Reads referenced artifacts from the store and inlines them so
176    /// providers never touch artifact storage directly.
177    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
257/// Deterministic offline image generator used when no live provider is
258/// configured and as the reference implementation in tests.
259pub struct FakeImageProvider;
260
261/// 1x1 transparent PNG used by deterministic fake image generation.
262pub 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
325/// Canonical `media_generate_image` tool. Registered by the runtime, replacing
326/// the offline-only fake tool contributed by `roder-tools`.
327pub 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}