rig_core/
image_generation.rs1use crate::markers::{Missing, Provided};
4use crate::wasm_compat::{WasmCompatSend, WasmCompatSync};
5use serde_json::Value;
6
7crate::provider_response::provider_error_enum!(
8 ImageGenerationError, "image generation" {
9 #[error("RequestError: {0}")]
11 RequestError(#[from] Box<dyn std::error::Error + Send + Sync + 'static>),
12 }
13);
14
15#[derive(Debug)]
17pub struct ImageGenerationResponse<T> {
18 pub image: Vec<u8>,
19 pub response: T,
20}
21
22pub trait ImageGenerationModel: Clone + WasmCompatSend + WasmCompatSync {
23 type Response: WasmCompatSend + WasmCompatSync;
24
25 type Client;
26
27 fn make(client: &Self::Client, model: impl Into<String>) -> Self;
28
29 fn image_generation(
30 &self,
31 request: ImageGenerationRequest,
32 ) -> impl std::future::Future<
33 Output = Result<ImageGenerationResponse<Self::Response>, ImageGenerationError>,
34 > + WasmCompatSend;
35
36 fn image_generation_request(&self) -> ImageGenerationRequestBuilder<Self, Missing> {
37 ImageGenerationRequestBuilder::new(self.clone())
38 }
39}
40pub struct ImageGenerationRequest {
42 pub prompt: String,
43 pub width: u32,
44 pub height: u32,
45 pub additional_params: Option<Value>,
46}
47
48pub struct ImageGenerationRequestBuilder<M, P = Missing>
51where
52 M: ImageGenerationModel,
53{
54 model: M,
55 prompt: P,
56 width: u32,
57 height: u32,
58 additional_params: Option<Value>,
59}
60
61impl<M> ImageGenerationRequestBuilder<M, Missing>
62where
63 M: ImageGenerationModel,
64{
65 pub fn new(model: M) -> Self {
66 Self {
67 model,
68 prompt: Missing,
69 height: 256,
70 width: 256,
71 additional_params: None,
72 }
73 }
74}
75
76impl<M, P> ImageGenerationRequestBuilder<M, P>
77where
78 M: ImageGenerationModel,
79{
80 pub fn prompt(self, prompt: &str) -> ImageGenerationRequestBuilder<M, Provided<String>> {
82 ImageGenerationRequestBuilder {
83 model: self.model,
84 prompt: Provided(prompt.to_string()),
85 width: self.width,
86 height: self.height,
87 additional_params: self.additional_params,
88 }
89 }
90
91 pub fn width(mut self, width: u32) -> Self {
93 self.width = width;
94 self
95 }
96
97 pub fn height(mut self, height: u32) -> Self {
99 self.height = height;
100 self
101 }
102
103 pub fn additional_params(mut self, params: Value) -> Self {
105 self.additional_params = Some(params);
106 self
107 }
108}
109
110impl<M> ImageGenerationRequestBuilder<M, Provided<String>>
111where
112 M: ImageGenerationModel,
113{
114 pub fn build(self) -> ImageGenerationRequest {
115 ImageGenerationRequest {
116 prompt: self.prompt.0,
117 width: self.width,
118 height: self.height,
119 additional_params: self.additional_params,
120 }
121 }
122
123 pub async fn send(self) -> Result<ImageGenerationResponse<M::Response>, ImageGenerationError> {
124 let model = self.model.clone();
125
126 model.image_generation(self.build()).await
127 }
128}
129
130#[cfg(test)]
131mod provider_response_tests {
132 use super::*;
133 use crate::{http_client, provider_response};
134 use http::StatusCode;
135
136 #[test]
137 fn image_generation_error_provider_response_helpers_with_preserved_json_body() {
138 let body = r#"{"error":{"message":"content policy"}}"#;
139 let error = ImageGenerationError::ProviderResponse(
140 provider_response::ProviderResponseError::without_status(body.to_string()),
141 );
142
143 assert_eq!(error.provider_response_body(), Some(body));
144 assert_eq!(error.provider_response_status(), None);
145 assert_eq!(
146 error.provider_response_json().expect("valid JSON"),
147 Some(serde_json::json!({ "error": { "message": "content policy" } }))
148 );
149 }
150
151 #[test]
152 fn image_generation_error_provider_response_helpers_with_http_non_success() {
153 let body = r#"{"error":{"message":"bad request"}}"#;
154 let error =
155 ImageGenerationError::HttpError(http_client::Error::InvalidStatusCodeWithMessage(
156 StatusCode::BAD_REQUEST,
157 body.to_string(),
158 ));
159
160 assert_eq!(error.provider_response_body(), Some(body));
161 assert_eq!(
162 error.provider_response_status(),
163 Some(StatusCode::BAD_REQUEST)
164 );
165 assert_eq!(
166 error.provider_response_json().expect("valid JSON"),
167 Some(serde_json::json!({ "error": { "message": "bad request" } }))
168 );
169 }
170
171 #[test]
172 fn image_generation_error_provider_error_is_not_a_provider_response() {
173 let error = ImageGenerationError::ProviderError("internal diagnostic".to_string());
174
175 assert_eq!(error.provider_response_body(), None);
176 assert_eq!(error.provider_response_status(), None);
177 assert_eq!(error.provider_response_json().expect("no body"), None);
178 }
179
180 #[test]
181 fn image_generation_error_provider_response_helpers_with_unrelated_variant() {
182 let error = ImageGenerationError::ResponseError("parse failed".to_string());
183
184 assert_eq!(error.provider_response_body(), None);
185 assert_eq!(error.provider_response_status(), None);
186 assert_eq!(error.provider_response_json().expect("no body"), None);
187 }
188}