Skip to main content

ferrin_google/
image.rs

1//! Gemini image generation (`generateContent` with the `IMAGE` modality).
2
3use 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
36/// Default maximum images per call.
37pub const DEFAULT_MAX_IMAGES_PER_CALL: usize = 10;
38
39/// Image model backed by a Gemini image-capable language model.
40#[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    /// Creates the model.
50    #[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    /// Overrides the advertised maximum; Gemini still rejects `n > 1`.
61    #[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    /// Builds the language model call for `options`.
68    ///
69    /// # Errors
70    ///
71    /// Returns [`ProviderError::InvalidArgument`] for non-Gemini model ids,
72    /// masks and `n > 1`.
73    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/// Converts the language model result of an image call into an image result:
197/// image files become the generated images, the `google` metadata gains an
198/// `images` array and the token usage is summed.
199#[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}