1use ferrin_spec::Content;
4use ferrin_spec::FileData;
5use ferrin_spec::FinishReasonKind;
6use ferrin_spec::JsonObject;
7use ferrin_spec::JsonValue;
8use ferrin_spec::ModelId;
9use ferrin_spec::ProviderId;
10use ferrin_spec::ProviderMetadata;
11use ferrin_spec::ProviderOptions;
12use ferrin_spec::ResponseMetadata;
13use ferrin_spec::ToolDefinition;
14use ferrin_spec::error::InvalidArgumentError;
15use ferrin_spec::error::ProviderError;
16use ferrin_spec::image_model::GeneratedImage;
17use ferrin_spec::image_model::ImageModel;
18use ferrin_spec::image_model::ImageOptions;
19use ferrin_spec::image_model::ImageResult;
20use ferrin_spec::image_model::ImageUsage;
21use ferrin_spec::language_model::CallOptions;
22use ferrin_spec::language_model::GenerateResult;
23use ferrin_spec::language_model::LanguageModel;
24use ferrin_spec::language_model::PromptMessage;
25use ferrin_spec::language_model::prompt::FilePart;
26use ferrin_spec::language_model::prompt::TextPart;
27use ferrin_spec::language_model::prompt::UserPromptPart;
28use ferrin_spec::shared::Warning;
29
30use crate::config::CANONICAL_OPTIONS_KEY;
31use crate::config::GoogleConfig;
32use crate::config::SharedConfig;
33use crate::language_model::GoogleLanguageModel;
34use crate::prepare_tools::ids;
35
36pub const DEFAULT_MAX_IMAGES_PER_CALL: usize = 10;
38
39#[derive(Debug, Clone)]
41pub struct GoogleImageModel {
42 config: SharedConfig,
43 provider: ProviderId,
44 model_id: ModelId,
45 max_images_per_call: usize,
46}
47
48impl GoogleImageModel {
49 #[must_use]
51 pub fn new(config: SharedConfig, model_id: impl Into<ModelId>) -> Self {
52 Self {
53 provider: ProviderId::new(config.name.clone()),
54 config,
55 model_id: model_id.into(),
56 max_images_per_call: DEFAULT_MAX_IMAGES_PER_CALL,
57 }
58 }
59
60 #[must_use]
62 pub fn with_max_images_per_call(mut self, max: usize) -> Self {
63 self.max_images_per_call = max;
64 self
65 }
66
67 pub fn prepare_call(
74 &self,
75 options: &ImageOptions,
76 ) -> Result<(CallOptions, Vec<Warning>), ProviderError> {
77 if !self.model_id.as_str().starts_with("gemini-") {
78 return Err(InvalidArgumentError::new(
79 "model_id",
80 "Google image models other than Gemini are not supported; use a model id that starts with `gemini-`",
81 )
82 .into());
83 }
84 if options.mask.is_some() {
85 return Err(InvalidArgumentError::new(
86 "mask",
87 "Gemini image models do not support mask-based image editing",
88 )
89 .into());
90 }
91 if options.n > 1 {
92 return Err(InvalidArgumentError::new(
93 "n",
94 "Gemini image models do not support generating a set number of images per call; use n=1 or omit n",
95 )
96 .into());
97 }
98 let mut warnings = Vec::new();
99 if options.size.is_some() {
100 warnings.push(Warning::unsupported_with_details(
101 "size",
102 "This model does not support the `size` option. Use `aspectRatio` instead.",
103 ));
104 }
105 let mut content = Vec::new();
106 if let Some(prompt) = &options.prompt {
107 content.push(UserPromptPart::Text(TextPart::new(prompt.clone())));
108 }
109 for file in &options.files {
110 let media_type = match &file.data {
111 FileData::Url { .. } => "image/*".to_owned(),
112 _ => file
113 .media_type
114 .as_ref()
115 .map_or_else(|| "image/*".to_owned(), |media| media.as_str().to_owned()),
116 };
117 content.push(UserPromptPart::File(FilePart::new(
118 file.data.clone(),
119 media_type,
120 )));
121 }
122 let mut google: JsonObject = options
123 .provider_options
124 .get(self.config.options_key())
125 .or_else(|| options.provider_options.get(CANONICAL_OPTIONS_KEY))
126 .cloned()
127 .unwrap_or_default();
128 let google_search = google.remove("googleSearch");
129 if let Some(search) = &google_search {
130 let valid = search.is_object()
131 && search.get("searchTypes").is_none_or(|types| {
132 types.is_object()
133 && ["webSearch", "imageSearch"]
134 .iter()
135 .all(|key| types.get(key).is_none_or(JsonValue::is_object))
136 })
137 && search.get("timeRangeFilter").is_none_or(|range| {
138 range.get("startTime").is_some_and(JsonValue::is_string)
139 && range.get("endTime").is_some_and(JsonValue::is_string)
140 });
141 if !valid {
142 return Err(InvalidArgumentError::new(
143 "googleSearch",
144 "invalid Google Search grounding options",
145 )
146 .into());
147 }
148 }
149 google.remove("responseModalities");
150 if google
151 .get("imageConfig")
152 .is_some_and(|value| !value.is_null() && !value.is_object())
153 {
154 return Err(
155 InvalidArgumentError::new("imageConfig", "imageConfig must be an object").into(),
156 );
157 }
158 let mut image_config = google.remove("imageConfig").and_then(|value| match value {
159 JsonValue::Object(object) => Some(object),
160 _ => None,
161 });
162 if let Some(ratio) = &options.aspect_ratio {
163 image_config
164 .get_or_insert_with(JsonObject::new)
165 .insert("aspectRatio".to_owned(), JsonValue::from(ratio.to_string()));
166 }
167 google.insert(
168 "responseModalities".to_owned(),
169 JsonValue::Array(vec![JsonValue::from("IMAGE")]),
170 );
171 if let Some(image_config) = image_config {
172 google.insert("imageConfig".to_owned(), JsonValue::Object(image_config));
173 }
174 let mut provider_options = ProviderOptions::new();
175 provider_options.insert(CANONICAL_OPTIONS_KEY.to_owned(), google);
176 let mut call = CallOptions::new(vec![PromptMessage::user(content)]);
177 call.seed = options.seed;
178 call.provider_options = provider_options;
179 call.headers = options.headers.clone();
180 call.cancellation = options.cancellation.clone();
181 if let Some(search) = google_search {
182 let args = match search {
183 JsonValue::Object(object) => object,
184 _ => JsonObject::new(),
185 };
186 call.tools.push(ToolDefinition::provider(
187 ids::GOOGLE_SEARCH,
188 "google_search",
189 args,
190 ));
191 }
192 Ok((call, warnings))
193 }
194}
195
196#[must_use]
200pub fn image_result(
201 config: &GoogleConfig,
202 model_id: ModelId,
203 result: GenerateResult,
204 warnings: Vec<Warning>,
205) -> ImageResult {
206 let mut images = Vec::new();
207 for part in &result.content {
208 if let Content::File {
209 data: FileData::Bytes { data },
210 media_type,
211 ..
212 } = part
213 && media_type.as_str().starts_with("image/")
214 {
215 images.push(GeneratedImage {
216 data: data.clone(),
217 media_type: Some(media_type.clone()),
218 });
219 }
220 }
221 let mut metadata = result
222 .provider_metadata
223 .as_ref()
224 .and_then(|metadata| metadata.get(CANONICAL_OPTIONS_KEY))
225 .cloned()
226 .unwrap_or_default();
227 metadata.insert(
228 "images".to_owned(),
229 JsonValue::Array(
230 images
231 .iter()
232 .map(|_| JsonValue::Object(JsonObject::new()))
233 .collect(),
234 ),
235 );
236 let mut provider_metadata = ProviderMetadata::new();
237 if config.options_key() != CANONICAL_OPTIONS_KEY {
238 provider_metadata.insert(config.options_key().to_owned(), metadata.clone());
239 }
240 provider_metadata.insert(CANONICAL_OPTIONS_KEY.to_owned(), metadata);
241 let input = result.usage.input.total;
242 let output = result.usage.output.total;
243 let mut all_warnings = warnings;
244 all_warnings.extend(result.warnings);
245 ImageResult {
246 images,
247 is_retryable: (result.finish_reason.unified == FinishReasonKind::ContentFilter)
248 .then_some(false),
249 warnings: all_warnings,
250 provider_metadata: Some(provider_metadata),
251 response: ResponseMetadata {
252 id: result.response.id,
253 timestamp: Some(chrono::Utc::now()),
254 model_id: Some(model_id),
255 headers: result.response.headers,
256 body: result.response.body,
257 },
258 usage: Some(ImageUsage {
259 input_tokens: input,
260 output_tokens: output,
261 total_tokens: Some(input.unwrap_or_default() + output.unwrap_or_default()),
262 }),
263 }
264}
265
266impl ImageModel for GoogleImageModel {
267 fn provider(&self) -> &ProviderId {
268 &self.provider
269 }
270
271 fn model_id(&self) -> &ModelId {
272 &self.model_id
273 }
274
275 fn max_images_per_call(&self) -> Option<usize> {
276 Some(self.max_images_per_call)
277 }
278
279 #[tracing::instrument(skip_all, fields(model = %self.model_id))]
280 async fn do_generate(&self, options: ImageOptions) -> Result<ImageResult, ProviderError> {
281 let (call, warnings) = self.prepare_call(&options)?;
282 let language_model = GoogleLanguageModel::new(self.config.clone(), self.model_id.clone());
283 let result = language_model.do_generate(call).await?;
284 Ok(image_result(
285 &self.config,
286 self.model_id.clone(),
287 result,
288 warnings,
289 ))
290 }
291}