edgequake-llm 0.10.1

Multi-provider LLM abstraction library with caching, rate limiting, and cost tracking
Documentation
//! OpenAI image generation provider (gpt-image-2).

use std::time::Instant;

use async_trait::async_trait;
use base64::engine::general_purpose::STANDARD as BASE64;
use base64::Engine;
use reqwest::Client;
use serde::Deserialize;
use serde_json::json;
use tracing::{debug, warn};

use crate::imagegen::error::{ImageGenError, Result};
use crate::imagegen::traits::ImageGenProvider;
use crate::imagegen::types::{
    AspectRatio, GeneratedImage, ImageGenData, ImageGenRequest, ImageGenResponse,
};

const DEFAULT_BASE_URL: &str = "https://api.openai.com/v1";
const DEFAULT_MODEL: &str = "gpt-image-2";

/// OpenAI image generation provider using gpt-image-2.
#[derive(Debug, Clone)]
pub struct OpenAIImageGen {
    api_key: String,
    base_url: String,
    http_client: Client,
}

#[derive(Debug, Deserialize)]
struct OpenAIImageResponse {
    data: Vec<OpenAIImageItem>,
}

#[derive(Debug, Deserialize)]
struct OpenAIImageItem {
    #[serde(default)]
    b64_json: Option<String>,
    #[serde(default)]
    url: Option<String>,
    #[serde(default)]
    revised_prompt: Option<String>,
}

impl OpenAIImageGen {
    pub fn new(api_key: impl Into<String>) -> Self {
        Self {
            api_key: api_key.into(),
            base_url: DEFAULT_BASE_URL.to_string(),
            http_client: Client::new(),
        }
    }

    pub fn from_env() -> Result<Self> {
        let api_key = std::env::var("OPENAI_API_KEY")
            .map_err(|_| ImageGenError::ConfigError("OPENAI_API_KEY must be set".to_string()))?;
        Ok(Self::new(api_key))
    }

    fn active_model<'a>(&'a self, request: &'a ImageGenRequest) -> &'a str {
        request.model.as_deref().unwrap_or(DEFAULT_MODEL)
    }

    fn endpoint_url(&self) -> String {
        format!("{}/images/generations", self.base_url)
    }

    fn map_aspect_ratio_to_size(ratio: AspectRatio) -> &'static str {
        match ratio {
            AspectRatio::Auto | AspectRatio::Square | AspectRatio::SquareHd => "1024x1024",
            AspectRatio::Landscape169
            | AspectRatio::Ultrawide
            | AspectRatio::Extreme41
            | AspectRatio::Extreme81 => "1792x1024",
            AspectRatio::Portrait169 | AspectRatio::Extreme14 | AspectRatio::Extreme18 => {
                "1024x1792"
            }
            AspectRatio::Landscape43 | AspectRatio::Frame54 | AspectRatio::Print32 => "1536x1024",
            AspectRatio::Portrait43 | AspectRatio::Frame45 | AspectRatio::Print23 => "1024x1536",
        }
    }

    fn resolve_size(&self, request: &ImageGenRequest) -> &'static str {
        let ratio = request.options.aspect_ratio_or_default();
        Self::map_aspect_ratio_to_size(ratio)
    }

    fn parse_size_dimensions(size: &str) -> (u32, u32) {
        let parts: Vec<&str> = size.split('x').collect();
        if parts.len() == 2 {
            let w = parts[0].parse().unwrap_or(1024);
            let h = parts[1].parse().unwrap_or(1024);
            (w, h)
        } else {
            (1024, 1024)
        }
    }

    fn build_request_body(&self, request: &ImageGenRequest) -> serde_json::Value {
        let model = self.active_model(request);
        let n = request.options.count_or_default();
        let size = self.resolve_size(request);

        let quality = request
            .options
            .extra
            .get("quality")
            .and_then(|v| v.as_str())
            .unwrap_or("auto");

        json!({
            "model": model,
            "prompt": request.prompt,
            "n": n,
            "size": size,
            "quality": quality,
            "response_format": "b64_json"
        })
    }

    fn parse_response(
        &self,
        response: OpenAIImageResponse,
        size: &str,
    ) -> Result<(Vec<GeneratedImage>, Option<String>)> {
        let (width, height) = Self::parse_size_dimensions(size);
        let mut images = Vec::new();
        let mut enhanced_prompt = None;

        for item in response.data {
            if enhanced_prompt.is_none() {
                enhanced_prompt = item.revised_prompt;
            }

            let data = if let Some(b64) = item.b64_json {
                let bytes = BASE64
                    .decode(&b64)
                    .map_err(|e| ImageGenError::InvalidResponse(format!("base64 decode: {e}")))?;
                ImageGenData::Bytes(bytes)
            } else if let Some(url) = item.url {
                ImageGenData::Url(url)
            } else {
                return Err(ImageGenError::InvalidResponse(
                    "response item has neither b64_json nor url".to_string(),
                ));
            };

            images.push(GeneratedImage {
                data,
                width,
                height,
                mime_type: "image/png".to_string(),
                seed: None,
            });
        }

        Ok((images, enhanced_prompt))
    }
}

#[async_trait]
impl ImageGenProvider for OpenAIImageGen {
    fn name(&self) -> &str {
        "openai"
    }

    fn default_model(&self) -> &str {
        DEFAULT_MODEL
    }

    fn available_models(&self) -> Vec<&str> {
        vec!["gpt-image-2", "gpt-image-1", "gpt-image-1-mini"]
    }

    async fn generate(&self, request: &ImageGenRequest) -> Result<ImageGenResponse> {
        if request.prompt.trim().is_empty() {
            return Err(ImageGenError::InvalidRequest(
                "prompt must not be empty".to_string(),
            ));
        }

        let model = self.active_model(request).to_string();
        let size = self.resolve_size(request);
        let body = self.build_request_body(request);

        debug!(provider = "openai", model = %model, "generating image");

        let started = Instant::now();
        let response = self
            .http_client
            .post(self.endpoint_url())
            .header("Authorization", format!("Bearer {}", self.api_key))
            .header("Content-Type", "application/json")
            .json(&body)
            .send()
            .await?;

        let status = response.status();
        let response_text = response.text().await?;

        if !status.is_success() {
            warn!(provider = "openai", status = %status, "image generation failed");
            return Err(match status.as_u16() {
                400 => ImageGenError::InvalidRequest(response_text),
                401 | 403 => ImageGenError::AuthError(response_text),
                429 => ImageGenError::RateLimited { retry_after: None },
                _ => ImageGenError::ProviderError(format!(
                    "HTTP {}: {}",
                    status.as_u16(),
                    response_text
                )),
            });
        }

        let latency_ms = started.elapsed().as_millis() as u64;
        let payload: OpenAIImageResponse = serde_json::from_str(&response_text)?;

        if payload.data.is_empty() {
            return Err(ImageGenError::InvalidResponse(
                "OpenAI returned empty image data".to_string(),
            ));
        }

        let (images, enhanced_prompt) = self.parse_response(payload, size)?;

        Ok(ImageGenResponse {
            images,
            provider: self.name().to_string(),
            model,
            latency_ms,
            enhanced_prompt,
        })
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::imagegen::types::{AspectRatio, ImageGenOptions, ImageGenRequest};

    #[test]
    fn test_build_request_body_default() {
        let provider = OpenAIImageGen::new("test-key");
        let request = ImageGenRequest::new("A sunset over mountains");
        let body = provider.build_request_body(&request);

        assert_eq!(body["model"], "gpt-image-2");
        assert_eq!(body["prompt"], "A sunset over mountains");
        assert_eq!(body["n"], 1);
        assert_eq!(body["size"], "1024x1024");
        assert_eq!(body["quality"], "auto");
        assert_eq!(body["response_format"], "b64_json");
    }

    #[test]
    fn test_build_request_body_landscape() {
        let provider = OpenAIImageGen::new("test-key");
        let request = ImageGenRequest::new("panoramic view").with_options(ImageGenOptions {
            aspect_ratio: Some(AspectRatio::Landscape169),
            count: Some(2),
            ..Default::default()
        });
        let body = provider.build_request_body(&request);

        assert_eq!(body["size"], "1792x1024");
        assert_eq!(body["n"], 2);
    }

    #[test]
    fn test_build_request_body_portrait() {
        let provider = OpenAIImageGen::new("test-key");
        let request = ImageGenRequest::new("tall building").with_options(ImageGenOptions {
            aspect_ratio: Some(AspectRatio::Portrait169),
            ..Default::default()
        });
        let body = provider.build_request_body(&request);

        assert_eq!(body["size"], "1024x1792");
    }

    #[test]
    fn test_build_request_body_with_model_override() {
        let provider = OpenAIImageGen::new("test-key");
        let request = ImageGenRequest::new("a cat").with_model("gpt-image-1");
        let body = provider.build_request_body(&request);

        assert_eq!(body["model"], "gpt-image-1");
    }

    #[test]
    fn test_size_mapping_all_ratios() {
        assert_eq!(
            OpenAIImageGen::map_aspect_ratio_to_size(AspectRatio::Square),
            "1024x1024"
        );
        assert_eq!(
            OpenAIImageGen::map_aspect_ratio_to_size(AspectRatio::Landscape43),
            "1536x1024"
        );
        assert_eq!(
            OpenAIImageGen::map_aspect_ratio_to_size(AspectRatio::Portrait43),
            "1024x1536"
        );
        assert_eq!(
            OpenAIImageGen::map_aspect_ratio_to_size(AspectRatio::Frame54),
            "1536x1024"
        );
        assert_eq!(
            OpenAIImageGen::map_aspect_ratio_to_size(AspectRatio::Frame45),
            "1024x1536"
        );
    }

    #[test]
    fn test_parse_size_dimensions() {
        assert_eq!(
            OpenAIImageGen::parse_size_dimensions("1024x1024"),
            (1024, 1024)
        );
        assert_eq!(
            OpenAIImageGen::parse_size_dimensions("1792x1024"),
            (1792, 1024)
        );
        assert_eq!(
            OpenAIImageGen::parse_size_dimensions("1024x1792"),
            (1024, 1792)
        );
    }

    #[test]
    fn test_parse_response_b64() {
        let provider = OpenAIImageGen::new("test-key");
        let raw = OpenAIImageResponse {
            data: vec![OpenAIImageItem {
                b64_json: Some(BASE64.encode(b"fake-image-bytes")),
                url: None,
                revised_prompt: Some("enhanced prompt".to_string()),
            }],
        };
        let (images, enhanced) = provider.parse_response(raw, "1024x1024").unwrap();

        assert_eq!(images.len(), 1);
        assert_eq!(images[0].width, 1024);
        assert_eq!(images[0].height, 1024);
        assert_eq!(
            images[0].data,
            ImageGenData::Bytes(b"fake-image-bytes".to_vec())
        );
        assert_eq!(enhanced, Some("enhanced prompt".to_string()));
    }

    #[test]
    fn test_available_models() {
        let provider = OpenAIImageGen::new("key");
        let models = provider.available_models();
        assert!(models.contains(&"gpt-image-2"));
        assert!(models.contains(&"gpt-image-1"));
        assert!(models.contains(&"gpt-image-1-mini"));
    }
}