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 google.remove("responseModalities");
130 let mut image_config = google.remove("imageConfig").and_then(|value| match value {
131 JsonValue::Object(object) => Some(object),
132 _ => None,
133 });
134 if let Some(ratio) = &options.aspect_ratio {
135 image_config
136 .get_or_insert_with(JsonObject::new)
137 .insert("aspectRatio".to_owned(), JsonValue::from(ratio.to_string()));
138 }
139 google.insert(
140 "responseModalities".to_owned(),
141 JsonValue::Array(vec![JsonValue::from("IMAGE")]),
142 );
143 if let Some(image_config) = image_config {
144 google.insert("imageConfig".to_owned(), JsonValue::Object(image_config));
145 }
146 let mut provider_options = ProviderOptions::new();
147 provider_options.insert(CANONICAL_OPTIONS_KEY.to_owned(), google);
148 let mut call = CallOptions::new(vec![PromptMessage::user(content)]);
149 call.seed = options.seed;
150 call.provider_options = provider_options;
151 call.headers = options.headers.clone();
152 call.cancellation = options.cancellation.clone();
153 if let Some(search) = google_search {
154 let args = match search {
155 JsonValue::Object(object) => object,
156 _ => JsonObject::new(),
157 };
158 call.tools.push(ToolDefinition::provider(
159 ids::GOOGLE_SEARCH,
160 "google_search",
161 args,
162 ));
163 }
164 Ok((call, warnings))
165 }
166}
167
168#[must_use]
172pub fn image_result(
173 config: &GoogleConfig,
174 model_id: ModelId,
175 result: GenerateResult,
176 warnings: Vec<Warning>,
177) -> ImageResult {
178 let mut images = Vec::new();
179 for part in &result.content {
180 if let Content::File {
181 data: FileData::Bytes { data },
182 media_type,
183 ..
184 } = part
185 && media_type.as_str().starts_with("image/")
186 {
187 images.push(GeneratedImage {
188 data: data.clone(),
189 media_type: Some(media_type.clone()),
190 });
191 }
192 }
193 let mut metadata = result
194 .provider_metadata
195 .as_ref()
196 .and_then(|metadata| metadata.get(CANONICAL_OPTIONS_KEY))
197 .cloned()
198 .unwrap_or_default();
199 metadata.insert(
200 "images".to_owned(),
201 JsonValue::Array(
202 images
203 .iter()
204 .map(|_| JsonValue::Object(JsonObject::new()))
205 .collect(),
206 ),
207 );
208 let mut provider_metadata = ProviderMetadata::new();
209 if config.options_key() != CANONICAL_OPTIONS_KEY {
210 provider_metadata.insert(config.options_key().to_owned(), metadata.clone());
211 }
212 provider_metadata.insert(CANONICAL_OPTIONS_KEY.to_owned(), metadata);
213 let input = result.usage.input.total;
214 let output = result.usage.output.total;
215 let mut all_warnings = warnings;
216 all_warnings.extend(result.warnings);
217 ImageResult {
218 images,
219 is_retryable: (result.finish_reason.unified == FinishReasonKind::ContentFilter)
220 .then_some(false),
221 warnings: all_warnings,
222 provider_metadata: Some(provider_metadata),
223 response: ResponseMetadata {
224 id: result.response.id,
225 timestamp: Some(chrono::Utc::now()),
226 model_id: Some(model_id),
227 headers: result.response.headers,
228 body: result.response.body,
229 },
230 usage: Some(ImageUsage {
231 input_tokens: input,
232 output_tokens: output,
233 total_tokens: Some(input.unwrap_or_default() + output.unwrap_or_default()),
234 }),
235 }
236}
237
238impl ImageModel for GoogleImageModel {
239 fn provider(&self) -> &ProviderId {
240 &self.provider
241 }
242
243 fn model_id(&self) -> &ModelId {
244 &self.model_id
245 }
246
247 fn max_images_per_call(&self) -> Option<usize> {
248 Some(self.max_images_per_call)
249 }
250
251 #[tracing::instrument(skip_all, fields(model = %self.model_id))]
252 async fn do_generate(&self, options: ImageOptions) -> Result<ImageResult, ProviderError> {
253 let (call, warnings) = self.prepare_call(&options)?;
254 let language_model = GoogleLanguageModel::new(self.config.clone(), self.model_id.clone());
255 let result = language_model.do_generate(call).await?;
256 Ok(image_result(
257 &self.config,
258 self.model_id.clone(),
259 result,
260 warnings,
261 ))
262 }
263}