rig_core/providers/venice/
image_generation.rs1use serde::{Deserialize, Serialize};
10use serde_json::json;
11
12use crate::image_generation::{self, ImageGenerationError, ImageGenerationRequest};
13use crate::json_utils::merge_inplace;
14use crate::providers::internal::image_generation::{
15 GenericImageGenerationModel, JsonImageGenerationProvider, decode_base64_image,
16};
17
18pub const VENICE_SD35: &str = "venice-sd35";
23pub const Z_IMAGE_TURBO: &str = "z-image-turbo";
25pub const QWEN_IMAGE: &str = "qwen-image";
27pub const FLUX_2_PRO: &str = "flux-2-pro";
29pub const HUNYUAN_IMAGE_V3: &str = "hunyuan-image-v3";
31
32#[derive(Debug, Clone, Copy, Default, Deserialize, Serialize)]
34pub struct ImageGenerationTiming {
35 #[serde(default)]
37 pub inference_duration: f64,
38 #[serde(default, rename = "inferencePreprocessingTime")]
40 pub inference_preprocessing_time: f64,
41 #[serde(default, rename = "inferenceQueueTime")]
43 pub inference_queue_time: f64,
44 #[serde(default)]
46 pub total: f64,
47}
48
49#[derive(Debug, Deserialize, Serialize)]
51pub struct ImageGenerationResponse {
52 pub id: String,
54 pub images: Vec<String>,
56 #[serde(default, skip_serializing_if = "Option::is_none")]
58 pub request: Option<serde_json::Value>,
59 #[serde(default, skip_serializing_if = "Option::is_none")]
61 pub timing: Option<ImageGenerationTiming>,
62}
63
64impl TryFrom<ImageGenerationResponse>
65 for image_generation::ImageGenerationResponse<ImageGenerationResponse>
66{
67 type Error = ImageGenerationError;
68
69 fn try_from(value: ImageGenerationResponse) -> Result<Self, Self::Error> {
70 decode_base64_image(
71 value,
72 |response| response.images.first().map(String::as_str),
73 "No image data returned",
74 Some("Base64 decode error: "),
75 )
76 }
77}
78
79pub type ImageGenerationModel<T = reqwest::Client> =
81 GenericImageGenerationModel<super::client::VeniceExt, T>;
82
83impl JsonImageGenerationProvider for super::client::VeniceExt {
84 const IMAGE_GENERATION_PATH: &'static str = "/image/generate";
85 type Response = ImageGenerationResponse;
86
87 fn image_generation_request_body(
88 model: &str,
89 generation_request: ImageGenerationRequest,
90 ) -> Result<serde_json::Value, ImageGenerationError> {
91 let mut request = json!({
95 "model": model,
96 "prompt": generation_request.prompt,
97 "width": generation_request.width,
98 "height": generation_request.height,
99 });
100
101 if let Some(additional_params) = generation_request.additional_params {
102 merge_inplace(&mut request, additional_params);
103 }
104
105 Ok(request)
106 }
107}
108
109#[cfg(test)]
110mod tests {
111 use super::*;
112 use crate::client::image_generation::ImageGenerationClient;
113 use crate::image_generation::ImageGenerationModel as _;
114
115 fn request() -> ImageGenerationRequest {
116 ImageGenerationRequest {
117 prompt: "a red circle on white".to_string(),
118 width: 256,
119 height: 256,
120 additional_params: None,
121 }
122 }
123
124 #[tokio::test]
128 async fn image_generation_non_success_preserves_status_and_body() {
129 use crate::test_utils::RecordingHttpClient;
130
131 let body = r#"{"error":"Specified model not found: nope"}"#;
132 let http_client =
133 RecordingHttpClient::with_error_response(http::StatusCode::NOT_FOUND, body);
134 let client = crate::providers::venice::Client::builder()
135 .api_key("test-key")
136 .http_client(http_client)
137 .build()
138 .expect("build client");
139 let model = client.image_generation_model(VENICE_SD35);
140
141 let error = model
142 .image_generation(request())
143 .await
144 .expect_err("should fail with non-success status");
145
146 assert!(matches!(error, ImageGenerationError::HttpError(_)));
147 assert_eq!(
148 error.provider_response_status(),
149 Some(http::StatusCode::NOT_FOUND)
150 );
151 assert_eq!(error.provider_response_body(), Some(body));
152 }
153
154 #[tokio::test]
155 async fn image_generation_posts_venice_native_body() {
156 use crate::test_utils::RecordingHttpClient;
157
158 let http_client = RecordingHttpClient::new(r#"{"id":"abc","images":["aGVsbG8="]}"#);
159 let client = crate::providers::venice::Client::builder()
160 .api_key("test-key")
161 .http_client(http_client.clone())
162 .build()
163 .expect("build client");
164 let model = client.image_generation_model(VENICE_SD35);
165
166 let response = model
167 .image_generation(request())
168 .await
169 .expect("image generation should succeed");
170
171 assert_eq!(response.image, b"hello");
172 assert_eq!(response.response.id, "abc");
173
174 let requests = http_client.requests();
175 let recorded = requests.first().expect("one request");
176 assert!(recorded.uri.ends_with("/image/generate"));
177 let body: serde_json::Value =
178 serde_json::from_slice(&recorded.body).expect("body should be JSON");
179 assert_eq!(
180 body,
181 serde_json::json!({
182 "model": VENICE_SD35,
183 "prompt": "a red circle on white",
184 "width": 256,
185 "height": 256,
186 })
187 );
188 }
189}