Skip to main content

rig_core/providers/gemini/
extension.rs

1//! Typed request options and reply extras for the Gemini API, keyed by
2//! [`PROVIDER_NAME`](super::PROVIDER_NAME). One [`GeminiOptions`] serves both
3//! routes: its `"*"` section goes to GenerateContent and Interactions alike,
4//! and each route section only to its own route.
5//!
6//! [`GeminiExtras`] reads the reply document of either route. A field the
7//! route taken does not return reads `None`.
8//!
9//! ```
10//! use rig_core::completion::CompletionRequest;
11//! use rig_core::providers::gemini::extension::GeminiOptions;
12//!
13//! let options = GeminiOptions::new().top_k(40);
14//! let request = CompletionRequest::new("hi").provider_option(options);
15//! # let _ = request;
16//! ```
17
18use std::collections::BTreeMap;
19
20use serde::de::Error as _;
21use serde::ser::SerializeMap;
22use serde::{Deserialize, Serialize, Serializer};
23use serde_json::Value;
24
25use crate::completion::provider_options::{ExtensionOptions, ProviderExtension, ReplyExtras};
26use crate::message::Api;
27
28/// The API name of the GenerateContent route.
29const GENERATE_CONTENT: &str = "gemini.generate_content";
30
31/// The API name of the Interactions route.
32const INTERACTIONS: &str = "gemini.interactions";
33
34/// The Gemini API's extension: [`GeminiOptions`] and [`GeminiExtras`].
35#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
36pub struct GeminiExt;
37
38impl ProviderExtension for GeminiExt {
39    const PROVIDER: &'static str = super::PROVIDER_NAME;
40    type Options = GeminiOptions;
41    type Extras = GeminiExtras;
42}
43
44/// Whether `value` is its type's default, so a section or config that sets
45/// nothing writes no key.
46fn is_default<T: Default + PartialEq>(value: &T) -> bool {
47    *value == T::default()
48}
49
50/// The Gemini API's request options: a section both routes read, and one
51/// per route.
52///
53/// The `generationConfig` setters here (`top_k`, `include_thoughts`, ...)
54/// write the same entry as the nested [`GenerateContentOptions`] form, so
55/// only GenerateContent sends them. [`generate_content`](Self::generate_content)
56/// replaces the whole section, including entries set before it. Safety
57/// settings stay per section, because each route spells them differently.
58#[non_exhaustive]
59#[derive(Clone, Debug, Default, PartialEq, Serialize)]
60pub struct GeminiOptions {
61    /// Fields both routes spell the same.
62    #[serde(rename = "*")]
63    pub shared: GeminiShared,
64    /// Fields only GenerateContent reads.
65    #[serde(rename = "gemini.generate_content")]
66    pub generate_content: GenerateContentOptions,
67    /// Fields only Interactions reads.
68    #[serde(rename = "gemini.interactions")]
69    pub interactions: InteractionsOptions,
70}
71
72impl GeminiOptions {
73    /// No field set.
74    pub fn new() -> Self {
75        Self::default()
76    }
77
78    /// Whether the API stores the request and its reply (`store`).
79    pub fn store(mut self, store: bool) -> Self {
80        self.shared.store = Some(store);
81        self
82    }
83
84    /// Add the label `key` = `value` (`labels`), user metadata the API
85    /// reports billing by.
86    pub fn label(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
87        self.shared.labels.insert(key.into(), value.into());
88        self
89    }
90
91    /// The GenerateContent section.
92    pub fn generate_content(mut self, section: GenerateContentOptions) -> Self {
93        self.generate_content = section;
94        self
95    }
96
97    /// The Interactions section.
98    pub fn interactions(mut self, section: InteractionsOptions) -> Self {
99        self.interactions = section;
100        self
101    }
102
103    /// Apply `set` to the GenerateContent `generationConfig` entries.
104    fn generation_config_entry(
105        mut self,
106        set: impl FnOnce(GenerationConfig) -> GenerationConfig,
107    ) -> Self {
108        let common = std::mem::take(&mut self.generate_content.generation_config.common);
109        self.generate_content.generation_config.common = set(common);
110        self
111    }
112
113    /// GenerateContent `generationConfig.thinkingConfig.includeThoughts`,
114    /// as [`GenerationConfig::include_thoughts`] in the GenerateContent
115    /// section. Interactions reads its own `thinking_summaries` instead.
116    pub fn include_thoughts(self, include: bool) -> Self {
117        self.generation_config_entry(|config| config.include_thoughts(include))
118    }
119
120    /// GenerateContent `generationConfig.topK`, as
121    /// [`GenerationConfig::top_k`] in the GenerateContent section.
122    pub fn top_k(self, top_k: u32) -> Self {
123        self.generation_config_entry(|config| config.top_k(top_k))
124    }
125
126    /// GenerateContent `generationConfig.presencePenalty`, as
127    /// [`GenerationConfig::presence_penalty`] in the GenerateContent section.
128    pub fn presence_penalty(self, penalty: f64) -> Self {
129        self.generation_config_entry(|config| config.presence_penalty(penalty))
130    }
131
132    /// GenerateContent `generationConfig.frequencyPenalty`, as
133    /// [`GenerationConfig::frequency_penalty`] in the GenerateContent
134    /// section.
135    pub fn frequency_penalty(self, penalty: f64) -> Self {
136        self.generation_config_entry(|config| config.frequency_penalty(penalty))
137    }
138
139    /// GenerateContent `generationConfig.responseLogprobs`, as
140    /// [`GenerationConfig::response_logprobs`] in the GenerateContent
141    /// section.
142    pub fn response_logprobs(self, enable: bool) -> Self {
143        self.generation_config_entry(|config| config.response_logprobs(enable))
144    }
145
146    /// GenerateContent `generationConfig.logprobs`, as
147    /// [`GenerationConfig::logprobs`] in the GenerateContent section.
148    pub fn logprobs(self, top: u32) -> Self {
149        self.generation_config_entry(|config| config.logprobs(top))
150    }
151
152    /// GenerateContent `generationConfig.candidateCount`, as
153    /// [`GenerationConfig::candidate_count`] in the GenerateContent section.
154    pub fn candidate_count(self, count: CandidateCount) -> Self {
155        self.generation_config_entry(|config| config.candidate_count(count))
156    }
157
158    /// GenerateContent `generationConfig.responseModalities`, as
159    /// [`GenerationConfig::response_modalities`] in the GenerateContent
160    /// section.
161    pub fn response_modalities(
162        self,
163        modalities: impl IntoIterator<Item = ResponseModality>,
164    ) -> Self {
165        self.generation_config_entry(|config| config.response_modalities(modalities))
166    }
167
168    /// GenerateContent `generationConfig.mediaResolution`, as
169    /// [`GenerationConfig::media_resolution`] in the GenerateContent
170    /// section.
171    pub fn media_resolution(self, resolution: MediaResolution) -> Self {
172        self.generation_config_entry(|config| config.media_resolution(resolution))
173    }
174}
175
176impl ExtensionOptions for GeminiOptions {
177    type Ext = GeminiExt;
178}
179
180/// The fields both Gemini routes read, at the top level of the body.
181#[non_exhaustive]
182#[derive(Clone, Debug, Default, PartialEq, Serialize)]
183pub struct GeminiShared {
184    /// `store`: whether the API stores the request and its reply.
185    #[serde(skip_serializing_if = "Option::is_none")]
186    pub store: Option<bool>,
187    /// `labels`: user metadata, each key and value at most 63 lowercase
188    /// letters, digits, underscores and dashes.
189    #[serde(skip_serializing_if = "BTreeMap::is_empty")]
190    pub labels: BTreeMap<String, String>,
191}
192
193/// The GenerateContent-only fields: `generationConfig` entries and
194/// `safetySettings`.
195#[non_exhaustive]
196#[derive(Clone, Debug, Default, PartialEq, Serialize)]
197pub struct GenerateContentOptions {
198    /// `generationConfig` entries, merged key by key with the ones the
199    /// request and its generation options write.
200    #[serde(rename = "generationConfig", skip_serializing_if = "is_default")]
201    pub generation_config: GeminiGenerationConfig,
202    /// `safetySettings`, which take the place of the `null` sent otherwise.
203    #[serde(rename = "safetySettings", skip_serializing_if = "Vec::is_empty")]
204    pub safety_settings: Vec<SafetySetting>,
205}
206
207impl GenerateContentOptions {
208    /// No field set.
209    pub fn new() -> Self {
210        Self::default()
211    }
212
213    /// The `generationConfig` entries every GenerateContent route takes.
214    pub fn generation_config(mut self, config: GenerationConfig) -> Self {
215        self.generation_config.common = config;
216        self
217    }
218
219    /// `generationConfig.enableEnhancedCivicAnswers`.
220    pub fn enable_enhanced_civic_answers(mut self, enable: bool) -> Self {
221        self.generation_config.enable_enhanced_civic_answers = Some(enable);
222        self
223    }
224
225    /// Add a safety setting: block `category` at `threshold`.
226    pub fn safety_setting(mut self, category: HarmCategory, threshold: HarmBlockThreshold) -> Self {
227        self.safety_settings
228            .push(SafetySetting::new(category, threshold));
229        self
230    }
231}
232
233/// The Gemini API's `generationConfig` entries: the ones every
234/// GenerateContent route takes and the API's own.
235#[non_exhaustive]
236#[derive(Clone, Debug, Default, PartialEq, Serialize)]
237pub struct GeminiGenerationConfig {
238    /// The entries every GenerateContent route takes.
239    #[serde(flatten)]
240    pub common: GenerationConfig,
241    /// `enableEnhancedCivicAnswers`.
242    #[serde(
243        rename = "enableEnhancedCivicAnswers",
244        skip_serializing_if = "Option::is_none"
245    )]
246    pub enable_enhanced_civic_answers: Option<bool>,
247}
248
249/// The `generationConfig` entries the Gemini API, Vertex AI and the gRPC
250/// API all take. None of them is one a request or its generation options
251/// write: `temperature`, `maxOutputTokens`, `topP`, `seed`, `stopSequences`
252/// and the thinking level or budget are set there.
253#[non_exhaustive]
254#[derive(Clone, Debug, Default, PartialEq, Serialize)]
255#[serde(rename_all = "camelCase")]
256pub struct GenerationConfig {
257    /// `thinkingConfig.includeThoughts`, next to the thinking level or
258    /// budget the request's reasoning sets.
259    #[serde(skip_serializing_if = "is_default")]
260    pub thinking_config: ThinkingConfig,
261    /// `topK`.
262    #[serde(skip_serializing_if = "Option::is_none")]
263    pub top_k: Option<u32>,
264    /// `presencePenalty`.
265    #[serde(skip_serializing_if = "Option::is_none")]
266    pub presence_penalty: Option<f64>,
267    /// `frequencyPenalty`.
268    #[serde(skip_serializing_if = "Option::is_none")]
269    pub frequency_penalty: Option<f64>,
270    /// `responseLogprobs`: return the chosen tokens' log probabilities.
271    #[serde(skip_serializing_if = "Option::is_none")]
272    pub response_logprobs: Option<bool>,
273    /// `logprobs`: how many top candidates to return per token.
274    #[serde(skip_serializing_if = "Option::is_none")]
275    pub logprobs: Option<u32>,
276    /// `candidateCount`.
277    #[serde(skip_serializing_if = "Option::is_none")]
278    pub candidate_count: Option<CandidateCount>,
279    /// `responseModalities`.
280    #[serde(skip_serializing_if = "Vec::is_empty")]
281    pub response_modalities: Vec<ResponseModality>,
282    /// `imageConfig`.
283    #[serde(skip_serializing_if = "Option::is_none")]
284    pub image_config: Option<ImageConfig>,
285    /// `speechConfig`.
286    #[serde(skip_serializing_if = "Option::is_none")]
287    pub speech_config: Option<SpeechConfig>,
288    /// `mediaResolution`.
289    #[serde(skip_serializing_if = "Option::is_none")]
290    pub media_resolution: Option<MediaResolution>,
291}
292
293impl GenerationConfig {
294    /// No entry set.
295    pub fn new() -> Self {
296        Self::default()
297    }
298
299    /// `thinkingConfig.includeThoughts`: return thought summaries.
300    pub fn include_thoughts(mut self, include: bool) -> Self {
301        self.thinking_config.include_thoughts = Some(include);
302        self
303    }
304
305    /// `topK`.
306    pub fn top_k(mut self, top_k: u32) -> Self {
307        self.top_k = Some(top_k);
308        self
309    }
310
311    /// `presencePenalty`.
312    pub fn presence_penalty(mut self, penalty: f64) -> Self {
313        self.presence_penalty = Some(penalty);
314        self
315    }
316
317    /// `frequencyPenalty`.
318    pub fn frequency_penalty(mut self, penalty: f64) -> Self {
319        self.frequency_penalty = Some(penalty);
320        self
321    }
322
323    /// `responseLogprobs`.
324    pub fn response_logprobs(mut self, enable: bool) -> Self {
325        self.response_logprobs = Some(enable);
326        self
327    }
328
329    /// `logprobs`.
330    pub fn logprobs(mut self, top: u32) -> Self {
331        self.logprobs = Some(top);
332        self
333    }
334
335    /// `candidateCount`.
336    pub fn candidate_count(mut self, count: CandidateCount) -> Self {
337        self.candidate_count = Some(count);
338        self
339    }
340
341    /// `responseModalities`.
342    pub fn response_modalities(
343        mut self,
344        modalities: impl IntoIterator<Item = ResponseModality>,
345    ) -> Self {
346        self.response_modalities = modalities.into_iter().collect();
347        self
348    }
349
350    /// `imageConfig`.
351    pub fn image_config(mut self, config: ImageConfig) -> Self {
352        self.image_config = Some(config);
353        self
354    }
355
356    /// `speechConfig`.
357    pub fn speech_config(mut self, config: SpeechConfig) -> Self {
358        self.speech_config = Some(config);
359        self
360    }
361
362    /// `mediaResolution`.
363    pub fn media_resolution(mut self, resolution: MediaResolution) -> Self {
364        self.media_resolution = Some(resolution);
365        self
366    }
367}
368
369/// The provider half of `thinkingConfig`.
370#[non_exhaustive]
371#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)]
372#[serde(rename_all = "camelCase")]
373pub struct ThinkingConfig {
374    /// `includeThoughts`.
375    #[serde(skip_serializing_if = "Option::is_none")]
376    pub include_thoughts: Option<bool>,
377}
378
379/// How many candidates to generate. Rig reads only the first candidate of
380/// a reply, so one is the only count it can ask for.
381#[non_exhaustive]
382#[derive(Clone, Copy, Debug, PartialEq, Eq)]
383pub enum CandidateCount {
384    /// One candidate, sent as `1`.
385    One,
386}
387
388impl Serialize for CandidateCount {
389    fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
390        match self {
391            Self::One => serializer.serialize_u32(1),
392        }
393    }
394}
395
396/// A modality of the reply.
397#[non_exhaustive]
398#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
399#[serde(rename_all = "SCREAMING_SNAKE_CASE")]
400pub enum ResponseModality {
401    /// Text.
402    Text,
403    /// Images.
404    Image,
405    /// Audio.
406    Audio,
407}
408
409/// The resolution media inputs are read at.
410#[non_exhaustive]
411#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
412pub enum MediaResolution {
413    /// `MEDIA_RESOLUTION_LOW`.
414    #[serde(rename = "MEDIA_RESOLUTION_LOW")]
415    Low,
416    /// `MEDIA_RESOLUTION_MEDIUM`.
417    #[serde(rename = "MEDIA_RESOLUTION_MEDIUM")]
418    Medium,
419    /// `MEDIA_RESOLUTION_HIGH`.
420    #[serde(rename = "MEDIA_RESOLUTION_HIGH")]
421    High,
422}
423
424/// How generated images are shaped.
425#[non_exhaustive]
426#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)]
427#[serde(rename_all = "camelCase")]
428pub struct ImageConfig {
429    /// `aspectRatio`, such as `"16:9"`.
430    #[serde(skip_serializing_if = "Option::is_none")]
431    pub aspect_ratio: Option<String>,
432    /// `imageSize`, such as `"2K"`.
433    #[serde(skip_serializing_if = "Option::is_none")]
434    pub image_size: Option<String>,
435}
436
437impl ImageConfig {
438    /// No entry set.
439    pub fn new() -> Self {
440        Self::default()
441    }
442
443    /// `aspectRatio`.
444    pub fn aspect_ratio(mut self, ratio: impl Into<String>) -> Self {
445        self.aspect_ratio = Some(ratio.into());
446        self
447    }
448
449    /// `imageSize`.
450    pub fn image_size(mut self, size: impl Into<String>) -> Self {
451        self.image_size = Some(size.into());
452        self
453    }
454}
455
456/// How generated speech sounds: one voice, or one per speaker.
457#[non_exhaustive]
458#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
459#[serde(rename_all = "camelCase")]
460pub struct SpeechConfig {
461    /// `voiceConfig`: the one voice.
462    #[serde(skip_serializing_if = "Option::is_none")]
463    pub voice_config: Option<VoiceConfig>,
464    /// `multiSpeakerVoiceConfig`: a voice per speaker.
465    #[serde(skip_serializing_if = "Option::is_none")]
466    pub multi_speaker_voice_config: Option<MultiSpeakerVoiceConfig>,
467    /// `languageCode`, such as `"en-US"`.
468    #[serde(skip_serializing_if = "Option::is_none")]
469    pub language_code: Option<String>,
470}
471
472impl SpeechConfig {
473    /// Speak with the prebuilt voice `voice_name`.
474    pub fn voice(voice_name: impl Into<String>) -> Self {
475        Self {
476            voice_config: Some(VoiceConfig::prebuilt(voice_name)),
477            multi_speaker_voice_config: None,
478            language_code: None,
479        }
480    }
481
482    /// Speak each `(speaker, voice_name)` pair's lines with its prebuilt
483    /// voice.
484    pub fn speakers<S: Into<String>, V: Into<String>>(
485        speakers: impl IntoIterator<Item = (S, V)>,
486    ) -> Self {
487        let speaker_voice_configs = speakers
488            .into_iter()
489            .map(|(speaker, voice)| SpeakerVoiceConfig {
490                speaker: speaker.into(),
491                voice_config: VoiceConfig::prebuilt(voice),
492            })
493            .collect();
494        Self {
495            voice_config: None,
496            multi_speaker_voice_config: Some(MultiSpeakerVoiceConfig {
497                speaker_voice_configs,
498            }),
499            language_code: None,
500        }
501    }
502
503    /// `languageCode`.
504    pub fn language_code(mut self, code: impl Into<String>) -> Self {
505        self.language_code = Some(code.into());
506        self
507    }
508}
509
510/// A voice.
511#[non_exhaustive]
512#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
513#[serde(rename_all = "camelCase")]
514pub struct VoiceConfig {
515    /// `prebuiltVoiceConfig`.
516    pub prebuilt_voice_config: PrebuiltVoiceConfig,
517}
518
519impl VoiceConfig {
520    /// The prebuilt voice `voice_name`, such as `"Kore"`.
521    pub fn prebuilt(voice_name: impl Into<String>) -> Self {
522        Self {
523            prebuilt_voice_config: PrebuiltVoiceConfig {
524                voice_name: voice_name.into(),
525            },
526        }
527    }
528}
529
530/// A prebuilt voice, by name.
531#[non_exhaustive]
532#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
533#[serde(rename_all = "camelCase")]
534pub struct PrebuiltVoiceConfig {
535    /// `voiceName`.
536    pub voice_name: String,
537}
538
539/// A voice per speaker.
540#[non_exhaustive]
541#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
542#[serde(rename_all = "camelCase")]
543pub struct MultiSpeakerVoiceConfig {
544    /// `speakerVoiceConfigs`.
545    pub speaker_voice_configs: Vec<SpeakerVoiceConfig>,
546}
547
548/// One speaker's voice.
549#[non_exhaustive]
550#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
551#[serde(rename_all = "camelCase")]
552pub struct SpeakerVoiceConfig {
553    /// `speaker`: the name the prompt gives the speaker.
554    pub speaker: String,
555    /// `voiceConfig`.
556    pub voice_config: VoiceConfig,
557}
558
559/// A harm category a safety setting applies to.
560#[non_exhaustive]
561#[derive(Clone, Copy, Debug, PartialEq, Eq)]
562pub enum HarmCategory {
563    /// Hate speech.
564    HateSpeech,
565    /// Dangerous content.
566    DangerousContent,
567    /// Harassment.
568    Harassment,
569    /// Sexually explicit content.
570    SexuallyExplicit,
571    /// Content that may be used to harm civic integrity.
572    CivicIntegrity,
573    /// Jailbreak attempts.
574    Jailbreak,
575}
576
577impl HarmCategory {
578    /// The GenerateContent spelling, such as `HARM_CATEGORY_HATE_SPEECH`.
579    pub fn as_str(self) -> &'static str {
580        match self {
581            Self::HateSpeech => "HARM_CATEGORY_HATE_SPEECH",
582            Self::DangerousContent => "HARM_CATEGORY_DANGEROUS_CONTENT",
583            Self::Harassment => "HARM_CATEGORY_HARASSMENT",
584            Self::SexuallyExplicit => "HARM_CATEGORY_SEXUALLY_EXPLICIT",
585            Self::CivicIntegrity => "HARM_CATEGORY_CIVIC_INTEGRITY",
586            Self::Jailbreak => "HARM_CATEGORY_JAILBREAK",
587        }
588    }
589
590    /// The Interactions spelling, such as `hate_speech`.
591    fn interactions(self) -> &'static str {
592        match self {
593            Self::HateSpeech => "hate_speech",
594            Self::DangerousContent => "dangerous_content",
595            Self::Harassment => "harassment",
596            Self::SexuallyExplicit => "sexually_explicit",
597            Self::CivicIntegrity => "civic_integrity",
598            Self::Jailbreak => "jailbreak",
599        }
600    }
601}
602
603/// The probability at and above which content is blocked.
604#[non_exhaustive]
605#[derive(Clone, Copy, Debug, PartialEq, Eq)]
606pub enum HarmBlockThreshold {
607    /// Block low probability and above.
608    BlockLowAndAbove,
609    /// Block medium probability and above.
610    BlockMediumAndAbove,
611    /// Block only high probability.
612    BlockOnlyHigh,
613    /// Block nothing.
614    BlockNone,
615    /// Turn the safety filter off.
616    Off,
617}
618
619impl HarmBlockThreshold {
620    /// The GenerateContent spelling, such as `BLOCK_ONLY_HIGH`.
621    pub fn as_str(self) -> &'static str {
622        match self {
623            Self::BlockLowAndAbove => "BLOCK_LOW_AND_ABOVE",
624            Self::BlockMediumAndAbove => "BLOCK_MEDIUM_AND_ABOVE",
625            Self::BlockOnlyHigh => "BLOCK_ONLY_HIGH",
626            Self::BlockNone => "BLOCK_NONE",
627            Self::Off => "OFF",
628        }
629    }
630
631    /// The Interactions spelling, such as `block_only_high`.
632    fn interactions(self) -> &'static str {
633        match self {
634            Self::BlockLowAndAbove => "block_low_and_above",
635            Self::BlockMediumAndAbove => "block_medium_and_above",
636            Self::BlockOnlyHigh => "block_only_high",
637            Self::BlockNone => "block_none",
638            Self::Off => "off",
639        }
640    }
641}
642
643/// Block `category` at `threshold`. Sent as
644/// `{"category": "HARM_CATEGORY_…", "threshold": "BLOCK_…"}` on
645/// GenerateContent and `{"type": "…", "threshold": "block_…"}` on
646/// Interactions.
647#[non_exhaustive]
648#[derive(Clone, Copy, Debug, PartialEq, Eq)]
649pub struct SafetySetting {
650    /// The category.
651    pub category: HarmCategory,
652    /// The threshold.
653    pub threshold: HarmBlockThreshold,
654}
655
656impl SafetySetting {
657    /// Block `category` at `threshold`.
658    pub fn new(category: HarmCategory, threshold: HarmBlockThreshold) -> Self {
659        Self {
660            category,
661            threshold,
662        }
663    }
664}
665
666impl Serialize for SafetySetting {
667    fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
668        let mut map = serializer.serialize_map(Some(2))?;
669        map.serialize_entry("category", self.category.as_str())?;
670        map.serialize_entry("threshold", self.threshold.as_str())?;
671        map.end()
672    }
673}
674
675/// `settings` in the Interactions spelling.
676fn interactions_safety<S: Serializer>(
677    settings: &[SafetySetting],
678    serializer: S,
679) -> Result<S::Ok, S::Error> {
680    serializer.collect_seq(settings.iter().map(|setting| {
681        BTreeMap::from([
682            ("type", setting.category.interactions()),
683            ("threshold", setting.threshold.interactions()),
684        ])
685    }))
686}
687
688/// The Interactions-only fields, at the top level of the create body but
689/// for `generation_config`.
690#[non_exhaustive]
691#[derive(Clone, Debug, Default, PartialEq, Serialize)]
692pub struct InteractionsOptions {
693    /// `agent`: the agent to run in place of the model, which the body then
694    /// leaves out.
695    #[serde(skip_serializing_if = "Option::is_none")]
696    pub agent: Option<String>,
697    /// `agent_config`.
698    #[serde(skip_serializing_if = "Option::is_none")]
699    pub agent_config: Option<AgentConfig>,
700    /// `background`: run the interaction in the background.
701    #[serde(skip_serializing_if = "Option::is_none")]
702    pub background: Option<bool>,
703    /// `previous_interaction_id`: continue a stored interaction.
704    #[serde(skip_serializing_if = "Option::is_none")]
705    pub previous_interaction_id: Option<String>,
706    /// `safety_settings`.
707    #[serde(
708        skip_serializing_if = "Vec::is_empty",
709        serialize_with = "interactions_safety"
710    )]
711    pub safety_settings: Vec<SafetySetting>,
712    /// `generation_config` entries, merged key by key with the ones the
713    /// request and its generation options write.
714    #[serde(skip_serializing_if = "is_default")]
715    pub generation_config: InteractionsGenerationConfig,
716}
717
718impl InteractionsOptions {
719    /// No field set.
720    pub fn new() -> Self {
721        Self::default()
722    }
723
724    /// `agent`.
725    pub fn agent(mut self, agent: impl Into<String>) -> Self {
726        self.agent = Some(agent.into());
727        self
728    }
729
730    /// `agent_config`.
731    pub fn agent_config(mut self, config: AgentConfig) -> Self {
732        self.agent_config = Some(config);
733        self
734    }
735
736    /// `background`.
737    pub fn background(mut self, background: bool) -> Self {
738        self.background = Some(background);
739        self
740    }
741
742    /// `previous_interaction_id`.
743    pub fn previous_interaction_id(mut self, id: impl Into<String>) -> Self {
744        self.previous_interaction_id = Some(id.into());
745        self
746    }
747
748    /// Add a safety setting: block `category` at `threshold`.
749    pub fn safety_setting(mut self, category: HarmCategory, threshold: HarmBlockThreshold) -> Self {
750        self.safety_settings
751            .push(SafetySetting::new(category, threshold));
752        self
753    }
754
755    /// `generation_config.thinking_summaries`.
756    pub fn thinking_summaries(mut self, summaries: ThinkingSummaries) -> Self {
757        self.generation_config.thinking_summaries = Some(summaries);
758        self
759    }
760
761    /// Add a voice to `generation_config.speech_config`.
762    pub fn speech(mut self, speech: InteractionSpeech) -> Self {
763        self.generation_config.speech_config.push(speech);
764        self
765    }
766}
767
768/// The provider entries of the Interactions `generation_config`.
769#[non_exhaustive]
770#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)]
771pub struct InteractionsGenerationConfig {
772    /// `thinking_summaries`.
773    #[serde(skip_serializing_if = "Option::is_none")]
774    pub thinking_summaries: Option<ThinkingSummaries>,
775    /// `speech_config`.
776    #[serde(skip_serializing_if = "Vec::is_empty")]
777    pub speech_config: Vec<InteractionSpeech>,
778}
779
780/// Whether the reply carries summaries of the model's thoughts.
781#[non_exhaustive]
782#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize)]
783#[serde(rename_all = "snake_case")]
784pub enum ThinkingSummaries {
785    /// `auto`.
786    Auto,
787    /// `none`.
788    None,
789}
790
791/// An agent's configuration, tagged by `type`.
792#[non_exhaustive]
793#[derive(Clone, Debug, PartialEq, Eq, Serialize)]
794#[serde(tag = "type", rename_all = "kebab-case")]
795pub enum AgentConfig {
796    /// `dynamic`.
797    Dynamic,
798    /// `deep-research`.
799    DeepResearch {
800        /// `thinking_summaries`.
801        #[serde(skip_serializing_if = "Option::is_none")]
802        thinking_summaries: Option<ThinkingSummaries>,
803    },
804}
805
806/// One voice of the Interactions `speech_config`.
807#[non_exhaustive]
808#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)]
809pub struct InteractionSpeech {
810    /// `voice`.
811    #[serde(skip_serializing_if = "Option::is_none")]
812    pub voice: Option<String>,
813    /// `language`.
814    #[serde(skip_serializing_if = "Option::is_none")]
815    pub language: Option<String>,
816    /// `speaker`: the name the prompt gives the speaker.
817    #[serde(skip_serializing_if = "Option::is_none")]
818    pub speaker: Option<String>,
819}
820
821impl InteractionSpeech {
822    /// Speak with `voice`.
823    pub fn voice(voice: impl Into<String>) -> Self {
824        Self {
825            voice: Some(voice.into()),
826            ..Self::default()
827        }
828    }
829
830    /// `language`.
831    pub fn language(mut self, language: impl Into<String>) -> Self {
832        self.language = Some(language.into());
833        self
834    }
835
836    /// `speaker`.
837    pub fn speaker(mut self, speaker: impl Into<String>) -> Self {
838        self.speaker = Some(speaker.into());
839        self
840    }
841}
842
843/// The Gemini API's reply fields rig does not normalize, from the reply
844/// document of either route. Each field says which route returns it; on the
845/// other route it reads `None`, as does a field the reply left out. A
846/// candidate field is the first candidate's, the one rig reads.
847#[non_exhaustive]
848#[derive(Clone, Debug, Default, PartialEq)]
849pub struct GeminiExtras {
850    /// GenerateContent `modelVersion`.
851    pub model_version: Option<String>,
852    /// GenerateContent `responseId`.
853    pub response_id: Option<String>,
854    /// The tier that served the request: GenerateContent
855    /// `usageMetadata.serviceTier`, Interactions `service_tier`.
856    pub service_tier: Option<String>,
857    /// GenerateContent `usageMetadata.promptTokensDetails`.
858    pub prompt_tokens_details: Option<Vec<ModalityTokenCount>>,
859    /// GenerateContent `usageMetadata.cacheTokensDetails`.
860    pub cache_tokens_details: Option<Vec<ModalityTokenCount>>,
861    /// GenerateContent `usageMetadata.candidatesTokensDetails`.
862    pub candidates_tokens_details: Option<Vec<ModalityTokenCount>>,
863    /// GenerateContent `usageMetadata.toolUsePromptTokensDetails`.
864    pub tool_use_prompt_tokens_details: Option<Vec<ModalityTokenCount>>,
865    /// GenerateContent `promptFeedback`.
866    pub prompt_feedback: Option<PromptFeedback>,
867    /// GenerateContent `safetyRatings` of the candidate.
868    pub safety_ratings: Option<Vec<SafetyRating>>,
869    /// GenerateContent `finishMessage` of the candidate.
870    pub finish_message: Option<String>,
871    /// GenerateContent `citationMetadata` of the candidate.
872    pub citation_metadata: Option<CitationMetadata>,
873    /// GenerateContent `groundingMetadata` of the candidate.
874    pub grounding_metadata: Option<GroundingMetadata>,
875    /// GenerateContent `urlContextMetadata` of the candidate.
876    pub url_context_metadata: Option<UrlContextMetadata>,
877    /// GenerateContent `avgLogprobs` of the candidate.
878    pub avg_logprobs: Option<f64>,
879    /// GenerateContent `logprobsResult` of the candidate.
880    pub logprobs_result: Option<LogprobsResult>,
881    /// Interactions `id`.
882    pub id: Option<String>,
883    /// Interactions `status`.
884    pub status: Option<InteractionStatus>,
885    /// Interactions `created`, an RFC 3339 timestamp.
886    pub created: Option<String>,
887    /// Interactions `updated`, an RFC 3339 timestamp.
888    pub updated: Option<String>,
889    /// Interactions `usage.input_tokens_by_modality`.
890    pub input_tokens_by_modality: Option<Vec<ModalityTokens>>,
891    /// Interactions `usage.output_tokens_by_modality`.
892    pub output_tokens_by_modality: Option<Vec<ModalityTokens>>,
893    /// Interactions `usage.cached_tokens_by_modality`.
894    pub cached_tokens_by_modality: Option<Vec<ModalityTokens>>,
895    /// Interactions `usage.grounding_tool_count`.
896    pub grounding_tool_count: Option<Vec<GroundingToolCount>>,
897}
898
899impl ReplyExtras for GeminiExtras {
900    fn from_reply(api: &Api, raw: &Value) -> Result<Self, serde_json::Error> {
901        match api.as_str() {
902            GENERATE_CONTENT => Ok(GenerateContentReply::deserialize(raw)?.into()),
903            INTERACTIONS => Ok(InteractionReply::deserialize(raw)?.into()),
904            other => Err(serde_json::Error::custom(format!(
905                "the Gemini API returns no `{other}` reply"
906            ))),
907        }
908    }
909}
910
911/// The parts of a `generateContent` reply [`GeminiExtras`] reads.
912#[derive(Deserialize)]
913#[serde(rename_all = "camelCase")]
914struct GenerateContentReply {
915    model_version: Option<String>,
916    response_id: Option<String>,
917    usage_metadata: Option<UsageDetails>,
918    prompt_feedback: Option<PromptFeedback>,
919    #[serde(default)]
920    candidates: Vec<CandidateDetails>,
921}
922
923#[derive(Deserialize)]
924#[serde(rename_all = "camelCase")]
925struct UsageDetails {
926    service_tier: Option<String>,
927    prompt_tokens_details: Option<Vec<ModalityTokenCount>>,
928    cache_tokens_details: Option<Vec<ModalityTokenCount>>,
929    candidates_tokens_details: Option<Vec<ModalityTokenCount>>,
930    tool_use_prompt_tokens_details: Option<Vec<ModalityTokenCount>>,
931}
932
933#[derive(Deserialize)]
934#[serde(rename_all = "camelCase")]
935struct CandidateDetails {
936    safety_ratings: Option<Vec<SafetyRating>>,
937    finish_message: Option<String>,
938    citation_metadata: Option<CitationMetadata>,
939    grounding_metadata: Option<GroundingMetadata>,
940    url_context_metadata: Option<UrlContextMetadata>,
941    avg_logprobs: Option<f64>,
942    logprobs_result: Option<LogprobsResult>,
943}
944
945impl From<GenerateContentReply> for GeminiExtras {
946    fn from(reply: GenerateContentReply) -> Self {
947        let mut extras = Self {
948            model_version: reply.model_version,
949            response_id: reply.response_id,
950            prompt_feedback: reply.prompt_feedback,
951            ..Self::default()
952        };
953        if let Some(usage) = reply.usage_metadata {
954            extras.service_tier = usage.service_tier;
955            extras.prompt_tokens_details = usage.prompt_tokens_details;
956            extras.cache_tokens_details = usage.cache_tokens_details;
957            extras.candidates_tokens_details = usage.candidates_tokens_details;
958            extras.tool_use_prompt_tokens_details = usage.tool_use_prompt_tokens_details;
959        }
960        if let Some(candidate) = reply.candidates.into_iter().next() {
961            extras.safety_ratings = candidate.safety_ratings;
962            extras.finish_message = candidate.finish_message;
963            extras.citation_metadata = candidate.citation_metadata;
964            extras.grounding_metadata = candidate.grounding_metadata;
965            extras.url_context_metadata = candidate.url_context_metadata;
966            extras.avg_logprobs = candidate.avg_logprobs;
967            extras.logprobs_result = candidate.logprobs_result;
968        }
969        extras
970    }
971}
972
973/// The parts of an interaction resource [`GeminiExtras`] reads.
974#[derive(Deserialize)]
975struct InteractionReply {
976    id: Option<String>,
977    status: Option<InteractionStatus>,
978    service_tier: Option<String>,
979    created: Option<String>,
980    updated: Option<String>,
981    usage: Option<InteractionUsage>,
982}
983
984#[derive(Deserialize)]
985struct InteractionUsage {
986    input_tokens_by_modality: Option<Vec<ModalityTokens>>,
987    output_tokens_by_modality: Option<Vec<ModalityTokens>>,
988    cached_tokens_by_modality: Option<Vec<ModalityTokens>>,
989    grounding_tool_count: Option<Vec<GroundingToolCount>>,
990}
991
992impl From<InteractionReply> for GeminiExtras {
993    fn from(reply: InteractionReply) -> Self {
994        let mut extras = Self {
995            id: reply.id,
996            status: reply.status,
997            service_tier: reply.service_tier,
998            created: reply.created,
999            updated: reply.updated,
1000            ..Self::default()
1001        };
1002        if let Some(usage) = reply.usage {
1003            extras.input_tokens_by_modality = usage.input_tokens_by_modality;
1004            extras.output_tokens_by_modality = usage.output_tokens_by_modality;
1005            extras.cached_tokens_by_modality = usage.cached_tokens_by_modality;
1006            extras.grounding_tool_count = usage.grounding_tool_count;
1007        }
1008        extras
1009    }
1010}
1011
1012/// A token count of one modality, as GenerateContent reports it.
1013#[non_exhaustive]
1014#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1015#[serde(rename_all = "camelCase")]
1016pub struct ModalityTokenCount {
1017    /// `modality`, such as `TEXT`.
1018    pub modality: Option<String>,
1019    /// `tokenCount`.
1020    pub token_count: Option<u64>,
1021}
1022
1023/// Why the prompt was blocked, and how it rated.
1024#[non_exhaustive]
1025#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1026#[serde(rename_all = "camelCase")]
1027pub struct PromptFeedback {
1028    /// `blockReason`, such as `SAFETY`, when the prompt was blocked.
1029    pub block_reason: Option<String>,
1030    /// `safetyRatings`.
1031    pub safety_ratings: Option<Vec<SafetyRating>>,
1032}
1033
1034/// How a prompt or candidate rated in one harm category.
1035#[non_exhaustive]
1036#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1037#[serde(rename_all = "camelCase")]
1038pub struct SafetyRating {
1039    /// `category`, such as `HARM_CATEGORY_HARASSMENT`.
1040    pub category: Option<String>,
1041    /// `probability`, such as `NEGLIGIBLE`.
1042    pub probability: Option<String>,
1043    /// `blocked`: whether this rating blocked the content.
1044    pub blocked: Option<bool>,
1045}
1046
1047/// The sources a candidate recites.
1048#[non_exhaustive]
1049#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1050#[serde(rename_all = "camelCase")]
1051pub struct CitationMetadata {
1052    /// `citationSources`.
1053    pub citation_sources: Option<Vec<CitationSource>>,
1054}
1055
1056/// One recited source.
1057#[non_exhaustive]
1058#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1059#[serde(rename_all = "camelCase")]
1060pub struct CitationSource {
1061    /// `startIndex` of the reciting text.
1062    pub start_index: Option<u64>,
1063    /// `endIndex` of the reciting text.
1064    pub end_index: Option<u64>,
1065    /// `uri` of the source.
1066    pub uri: Option<String>,
1067    /// `license` of the source.
1068    pub license: Option<String>,
1069}
1070
1071/// What grounded a candidate in search or retrieval.
1072#[non_exhaustive]
1073#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
1074#[serde(rename_all = "camelCase")]
1075pub struct GroundingMetadata {
1076    /// `webSearchQueries`.
1077    pub web_search_queries: Option<Vec<String>>,
1078    /// `searchEntryPoint`.
1079    pub search_entry_point: Option<SearchEntryPoint>,
1080    /// `groundingChunks`.
1081    pub grounding_chunks: Option<Vec<GroundingChunk>>,
1082    /// `groundingSupports`.
1083    pub grounding_supports: Option<Vec<GroundingSupport>>,
1084}
1085
1086/// The search suggestion to show with a grounded answer.
1087#[non_exhaustive]
1088#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1089#[serde(rename_all = "camelCase")]
1090pub struct SearchEntryPoint {
1091    /// `renderedContent`: HTML and CSS to embed.
1092    pub rendered_content: Option<String>,
1093}
1094
1095/// One grounding source.
1096#[non_exhaustive]
1097#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1098#[serde(rename_all = "camelCase")]
1099pub struct GroundingChunk {
1100    /// `web`: a web page.
1101    pub web: Option<WebChunk>,
1102    /// `retrievedContext`: a retrieved document.
1103    pub retrieved_context: Option<RetrievedContext>,
1104}
1105
1106/// A web page a candidate is grounded in.
1107#[non_exhaustive]
1108#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1109pub struct WebChunk {
1110    /// `uri`.
1111    pub uri: Option<String>,
1112    /// `title`.
1113    pub title: Option<String>,
1114}
1115
1116/// A retrieved document a candidate is grounded in.
1117#[non_exhaustive]
1118#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1119pub struct RetrievedContext {
1120    /// `uri`.
1121    pub uri: Option<String>,
1122    /// `title`.
1123    pub title: Option<String>,
1124    /// `text`.
1125    pub text: Option<String>,
1126}
1127
1128/// A span of the candidate and the chunks that support it.
1129#[non_exhaustive]
1130#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
1131#[serde(rename_all = "camelCase")]
1132pub struct GroundingSupport {
1133    /// `segment`.
1134    pub segment: Option<Segment>,
1135    /// `groundingChunkIndices`, into `groundingChunks`.
1136    pub grounding_chunk_indices: Option<Vec<u32>>,
1137    /// `confidenceScores`, one per index.
1138    pub confidence_scores: Option<Vec<f64>>,
1139}
1140
1141/// A span of a candidate's content.
1142#[non_exhaustive]
1143#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1144#[serde(rename_all = "camelCase")]
1145pub struct Segment {
1146    /// `partIndex`.
1147    pub part_index: Option<u32>,
1148    /// `startIndex`, in bytes.
1149    pub start_index: Option<u64>,
1150    /// `endIndex`, in bytes.
1151    pub end_index: Option<u64>,
1152    /// `text`.
1153    pub text: Option<String>,
1154}
1155
1156/// The URLs the URL context tool retrieved.
1157#[non_exhaustive]
1158#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1159#[serde(rename_all = "camelCase")]
1160pub struct UrlContextMetadata {
1161    /// `urlMetadata`.
1162    pub url_metadata: Option<Vec<UrlMetadata>>,
1163}
1164
1165/// One URL the URL context tool retrieved.
1166#[non_exhaustive]
1167#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1168#[serde(rename_all = "camelCase")]
1169pub struct UrlMetadata {
1170    /// `retrievedUrl`.
1171    pub retrieved_url: Option<String>,
1172    /// `urlRetrievalStatus`, such as `URL_RETRIEVAL_STATUS_SUCCESS`.
1173    pub url_retrieval_status: Option<String>,
1174}
1175
1176/// The log probabilities of a candidate's tokens.
1177#[non_exhaustive]
1178#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
1179#[serde(rename_all = "camelCase")]
1180pub struct LogprobsResult {
1181    /// `topCandidates`, one entry per decoding step.
1182    pub top_candidates: Option<Vec<TopCandidates>>,
1183    /// `chosenCandidates`, one per decoding step.
1184    pub chosen_candidates: Option<Vec<LogprobsCandidate>>,
1185    /// `logProbabilitySum`.
1186    pub log_probability_sum: Option<f64>,
1187}
1188
1189/// The most likely tokens of one decoding step.
1190#[non_exhaustive]
1191#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
1192pub struct TopCandidates {
1193    /// `candidates`, most likely first.
1194    pub candidates: Option<Vec<LogprobsCandidate>>,
1195}
1196
1197/// A token and its log probability.
1198#[non_exhaustive]
1199#[derive(Clone, Debug, Default, PartialEq, Deserialize)]
1200#[serde(rename_all = "camelCase")]
1201pub struct LogprobsCandidate {
1202    /// `token`.
1203    pub token: Option<String>,
1204    /// `tokenId`.
1205    pub token_id: Option<i64>,
1206    /// `logProbability`.
1207    pub log_probability: Option<f64>,
1208}
1209
1210/// Where an interaction stands.
1211#[non_exhaustive]
1212#[derive(Clone, Debug, PartialEq, Eq, Deserialize)]
1213#[serde(rename_all = "snake_case")]
1214pub enum InteractionStatus {
1215    /// `in_progress`.
1216    InProgress,
1217    /// `requires_action`.
1218    RequiresAction,
1219    /// `incomplete`.
1220    Incomplete,
1221    /// `budget_exceeded`.
1222    BudgetExceeded,
1223    /// `completed`.
1224    Completed,
1225    /// `failed`.
1226    Failed,
1227    /// `cancelled`.
1228    Cancelled,
1229    /// A status this crate does not know yet, as the API spelled it.
1230    #[serde(untagged)]
1231    Unknown(String),
1232}
1233
1234/// A token count of one modality, as Interactions reports it.
1235#[non_exhaustive]
1236#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1237pub struct ModalityTokens {
1238    /// `modality`, such as `text`.
1239    pub modality: Option<String>,
1240    /// `tokens`.
1241    pub tokens: Option<u64>,
1242}
1243
1244/// How often a grounding tool ran.
1245#[non_exhaustive]
1246#[derive(Clone, Debug, Default, PartialEq, Eq, Deserialize)]
1247pub struct GroundingToolCount {
1248    /// `type`, such as `google_search`.
1249    #[serde(rename = "type")]
1250    pub kind: Option<String>,
1251    /// `count`.
1252    pub count: Option<u64>,
1253    /// `search_query_count`.
1254    pub search_query_count: Option<u64>,
1255}
1256
1257#[cfg(test)]
1258mod tests;