Skip to main content

rig_core/providers/xai/
image_generation.rs

1use crate::image_generation;
2use crate::image_generation::{ImageGenerationError, ImageGenerationRequest};
3use crate::json_utils::merge_inplace;
4use crate::providers::internal::image_generation::{
5    GenericImageGenerationModel, JsonImageGenerationProvider, decode_base64_image,
6};
7use serde::Deserialize;
8use serde_json::json;
9
10// ================================================================
11// xAI Image Generation API
12// ================================================================
13pub const GROK_IMAGINE_IMAGE: &str = "grok-imagine-image";
14pub const GROK_IMAGINE_IMAGE_PRO: &str = "grok-imagine-image-pro";
15
16#[derive(Debug, Deserialize)]
17pub struct ImageGenerationData {
18    pub b64_json: String,
19}
20
21#[derive(Debug, Deserialize)]
22pub struct ImageGenerationResponse {
23    pub data: Vec<ImageGenerationData>,
24}
25
26impl TryFrom<ImageGenerationResponse>
27    for image_generation::ImageGenerationResponse<ImageGenerationResponse>
28{
29    type Error = ImageGenerationError;
30
31    fn try_from(value: ImageGenerationResponse) -> Result<Self, Self::Error> {
32        decode_base64_image(
33            value,
34            |response| response.data.first().map(|image| image.b64_json.as_str()),
35            "No image data returned",
36            Some("Base64 decode error: "),
37        )
38    }
39}
40
41/// xAI image generation model.
42pub type ImageGenerationModel<T = reqwest::Client> =
43    GenericImageGenerationModel<super::client::XAiExt, T>;
44
45impl JsonImageGenerationProvider for super::client::XAiExt {
46    const IMAGE_GENERATION_PATH: &'static str = "/v1/images/generations";
47    type Response = ImageGenerationResponse;
48
49    fn image_generation_request_body(
50        model: &str,
51        generation_request: ImageGenerationRequest,
52    ) -> Result<serde_json::Value, ImageGenerationError> {
53        let mut request = json!({
54            "model": model,
55            "prompt": generation_request.prompt,
56            "response_format": "b64_json",
57            "aspect_ratio": "1:1",
58        });
59
60        if let Some(additional_params) = generation_request.additional_params {
61            merge_inplace(&mut request, additional_params);
62        }
63
64        Ok(request)
65    }
66}
67
68#[cfg(test)]
69mod tests {
70    use super::*;
71    use crate::client::image_generation::ImageGenerationClient;
72    use crate::image_generation::ImageGenerationModel as _;
73
74    fn request() -> ImageGenerationRequest {
75        ImageGenerationRequest {
76            prompt: "draw a cat".to_string(),
77            width: 256,
78            height: 256,
79            additional_params: None,
80        }
81    }
82
83    #[tokio::test]
84    async fn image_generation_non_success_preserves_status_and_body() {
85        use crate::test_utils::RecordingHttpClient;
86
87        let body = r#"{"error":"boom","code":"503"}"#;
88        let http_client =
89            RecordingHttpClient::with_error_response(http::StatusCode::SERVICE_UNAVAILABLE, body);
90        let client = crate::providers::xai::Client::builder()
91            .api_key("test-key")
92            .http_client(http_client)
93            .build()
94            .expect("build client");
95        let model = client.image_generation_model(GROK_IMAGINE_IMAGE);
96
97        let error = model
98            .image_generation(request())
99            .await
100            .expect_err("should fail with non-success status");
101
102        assert!(matches!(error, ImageGenerationError::HttpError(_)));
103        assert_eq!(
104            error.provider_response_status(),
105            Some(http::StatusCode::SERVICE_UNAVAILABLE)
106        );
107        assert_eq!(error.provider_response_body(), Some(body));
108    }
109
110    #[tokio::test]
111    async fn image_generation_2xx_error_envelope_preserves_status_and_body() {
112        use crate::test_utils::RecordingHttpClient;
113
114        // Deserializes to `ApiResponse::Err(ApiErrorResponse)` on a 200 OK.
115        let body = r#"{"error":"boom","code":"503"}"#;
116        let http_client = RecordingHttpClient::new(body);
117        let client = crate::providers::xai::Client::builder()
118            .api_key("test-key")
119            .http_client(http_client)
120            .build()
121            .expect("build client");
122        let model = client.image_generation_model(GROK_IMAGINE_IMAGE);
123
124        let error = model
125            .image_generation(request())
126            .await
127            .expect_err("should fail with provider error envelope");
128
129        match &error {
130            ImageGenerationError::ProviderResponse(stored) => {
131                assert_eq!(stored.body, body);
132                assert_eq!(stored.status, Some(http::StatusCode::OK));
133            }
134            other => panic!("expected ProviderResponse, got {other:?}"),
135        }
136    }
137}