rig_core/providers/xai/
image_generation.rs1use 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
10pub 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
41pub 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 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}