1pub const GEMINI_3_1_FLASH_LITE_PREVIEW: &str = "gemini-3.1-flash-lite-preview";
13pub const GEMINI_3_FLASH_PREVIEW: &str = "gemini-3-flash-preview";
15pub const GEMINI_2_5_PRO_PREVIEW_06_05: &str = "gemini-2.5-pro-preview-06-05";
17pub const GEMINI_2_5_PRO_PREVIEW_05_06: &str = "gemini-2.5-pro-preview-05-06";
19pub const GEMINI_2_5_PRO_PREVIEW_03_25: &str = "gemini-2.5-pro-preview-03-25";
21pub const GEMINI_2_5_FLASH_PREVIEW_04_17: &str = "gemini-2.5-flash-preview-04-17";
23pub const GEMINI_2_5_PRO_EXP_03_25: &str = "gemini-2.5-pro-exp-03-25";
25pub const GEMINI_2_5_FLASH: &str = "gemini-2.5-flash";
27#[cfg(feature = "image")]
29#[cfg_attr(docsrs, doc(cfg(feature = "image")))]
30pub const GEMINI_2_5_FLASH_IMAGE: &str = "gemini-2.5-flash-image";
31pub const GEMINI_2_0_FLASH_LITE: &str = "gemini-2.0-flash-lite";
33pub const GEMINI_2_0_FLASH: &str = "gemini-2.0-flash";
35
36use self::gemini_api_types::tool_parameters_to_schema;
37use crate::completion::{self, CompletionRequest};
38use crate::error::EncodeError;
39use crate::error::ProviderError;
40use crate::operation::Completion;
41use crate::providers::gemini::completion::gemini_api_types::{
42 AdditionalParameters, FunctionCallingMode, ToolConfig,
43};
44use crate::telemetry::GenAiOperation;
45use crate::wire::{Body, Descriptor, Encoded, Framing, Mode, Wire};
46use gemini_api_types::{
47 Content, FinishReason, FunctionDeclaration, GenerateContentRequest, GenerationConfig, Part,
48 PartKind, Role, Tool,
49};
50use serde_json::{Map, Value};
51use std::convert::TryFrom;
52
53pub const PROVIDER_NAME: &str = "gcp.gemini";
55
56pub(crate) const ISSUER: crate::message::Issuer =
58 crate::message::Issuer::from_static(PROVIDER_NAME);
59
60#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
63pub struct GenerateContent {
64 pub provider: super::GeminiConfig,
66 pub model: String,
68 pub cached_content: Option<String>,
71 #[serde(default)]
73 pub thought_replay: ThoughtReplay,
74}
75
76#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
91pub enum ThoughtReplay {
92 #[default]
94 All,
95 CurrentTurn,
99}
100
101impl GenerateContent {
102 pub fn new(provider: super::GeminiConfig, model: impl Into<String>) -> Self {
104 Self {
105 provider,
106 model: model.into(),
107 cached_content: None,
108 thought_replay: ThoughtReplay::All,
109 }
110 }
111
112 pub fn thought_replay(mut self, replay: ThoughtReplay) -> Self {
114 self.thought_replay = replay;
115 self
116 }
117
118 pub fn with_cached_content(mut self, name: impl Into<String>) -> Self {
122 self.cached_content = Some(name.into());
123 self
124 }
125}
126
127impl<T> crate::driver::Model<GenerateContent, T> {
128 pub fn thought_replay(mut self, replay: ThoughtReplay) -> Self {
132 self.wire.thought_replay = replay;
133 self
134 }
135}
136
137fn drop_finished_signatures(contents: &mut [Content]) {
141 let current = contents.iter().rposition(|content| {
142 content.role == Some(Role::User)
143 && content
144 .parts
145 .iter()
146 .any(|part| matches!(part.part, PartKind::Text(_)))
147 && !content
148 .parts
149 .iter()
150 .any(|part| matches!(part.part, PartKind::FunctionResponse(_)))
151 });
152 let Some(current) = current else {
153 return;
154 };
155 for content in contents.iter_mut().take(current) {
156 for part in &mut content.parts {
157 part.thought_signature = None;
158 if let Some(Value::Object(extra)) = &mut part.additional_params {
159 extra.remove("thoughtSignature");
160 extra.remove("thought_signature");
161 }
162 }
163 }
164}
165
166impl Wire for GenerateContent {
167 type Op = Completion;
168 type Payload = crate::wire::Encoded;
169 type Frame = crate::wire::WireFrame;
170 type Decoder<'id> = super::streaming::GenerateContentDecoder<'id>;
171
172 fn describe(&self) -> Descriptor<'_> {
173 Descriptor::new(PROVIDER_NAME)
174 .model(self.model.as_str())
175 .telemetry(|mode| match mode {
176 Mode::Unary => GenAiOperation::GenerateContent,
177 Mode::Streaming => GenAiOperation::ChatStreaming,
178 })
179 }
180
181 fn encode(&self, request: CompletionRequest, mode: Mode) -> Result<Encoded, EncodeError> {
182 let request = request.replayable_to(&[ISSUER])?;
183 let model = resolve_request_model(&self.model, &request);
185 let mut body = create_request_body(request)?;
186 if let Some(name) = self.cached_content.as_deref() {
187 body.with_cached_content(name)?;
188 }
189 if self.thought_replay == ThoughtReplay::CurrentTurn {
190 drop_finished_signatures(&mut body.contents);
191 }
192 let (path, framing, target) = match mode {
193 Mode::Unary => (
194 completion_endpoint(&model),
195 Framing::Whole,
196 crate::providers::internal::LogTarget::Completions,
197 ),
198 Mode::Streaming => (
201 format!("{}?alt=sse", streaming_endpoint(&model)),
202 Framing::Sse,
203 crate::providers::internal::LogTarget::Streaming,
204 ),
205 };
206 crate::providers::internal::trace_json(target, "Gemini completion request", &body);
207 let request = http::Request::post(self.provider.uri(&path))
208 .header("Content-Type", "application/json")
209 .body(Body::Bytes(serde_json::to_vec(&body)?))?;
210 Ok(Encoded::new(request, framing)
212 .with_projection(super::streaming::GenerateContentDecoder::project)
213 .with_analysis_only(super::streaming::GenerateContentDecoder::is_analysis_only))
214 }
215
216 fn decoder<'id>(&self) -> Self::Decoder<'id> {
217 super::streaming::GenerateContentDecoder::new()
218 }
219}
220
221pub(crate) fn create_request_body(
222 completion_request: CompletionRequest,
223) -> Result<GenerateContentRequest, EncodeError> {
224 let chat_history = completion_request.chat_history_with_documents();
225
226 let CompletionRequest {
227 model: _,
228 chat_history: _,
229 documents: _,
230 tools: function_tools,
231 temperature,
232 max_tokens,
233 tool_choice,
234 mut additional_params,
235 output_schema,
236 record_telemetry_content: _,
237 } = completion_request;
238
239 let mut full_history = Vec::new();
240 full_history.extend(chat_history);
241 let (history_system, full_history) = split_system_messages_from_history(full_history);
242
243 let mut additional_params_payload = additional_params
244 .take()
245 .unwrap_or_else(|| Value::Object(Map::new()));
246 let mut additional_tools =
247 extract_tools_from_additional_params(&mut additional_params_payload)?;
248 let mut smuggled_cached_content = Vec::new();
251 for spelling in CACHED_CONTENT {
252 let Some(value) = additional_params_payload
253 .as_object_mut()
254 .and_then(|object| object.remove(spelling))
255 else {
256 continue;
257 };
258 match value {
259 Value::String(name) => smuggled_cached_content.push(name),
260 other => {
261 return Err(EncodeError::request(format!(
262 "additional_params.{spelling} should be a string, got {other}"
263 )));
264 }
265 }
266 }
267 let smuggled_system_instruction =
270 smuggled_field(&additional_params_payload, &SYSTEM_INSTRUCTION);
271 let smuggled_tool_config = smuggled_field(&additional_params_payload, &TOOL_CONFIG);
272
273 let AdditionalParameters {
274 mut generation_config,
275 additional_params,
276 } = serde_json::from_value::<AdditionalParameters>(additional_params_payload)?;
277
278 if let Some(schema) = output_schema {
279 let cfg = generation_config.get_or_insert_with(GenerationConfig::default);
280 cfg.response_mime_type = Some("application/json".to_string());
281 cfg.response_json_schema = Some(schema.to_value());
282 }
283
284 if temperature.is_some() || max_tokens.is_some() {
287 let cfg = generation_config.get_or_insert_with(GenerationConfig::default);
288
289 if let Some(temp) = temperature {
290 cfg.temperature = Some(temp);
291 }
292
293 if let Some(max_tokens) = max_tokens {
294 cfg.max_output_tokens = Some(max_tokens);
295 }
296 }
297
298 let mut system_parts: Vec<Part> = Vec::new();
299 for content in history_system {
300 if !content.is_empty() {
301 system_parts.push(content.into());
302 }
303 }
304 let system_instruction = if system_parts.is_empty() {
305 None
306 } else {
307 Some(Content {
308 parts: system_parts,
309 role: Some(Role::Model),
310 })
311 };
312 if let (Some(typed), Some(spelling)) = (&system_instruction, smuggled_system_instruction) {
314 return Err(EncodeError::request(format!(
315 "a Gemini request set the system instruction twice — once as a preamble or \
316 system message ({} part(s)) and once through `additional_params.{spelling}`. \
317 Both would reach the wire, and Gemini rejects that outright: \
318 `system_instruction` is an optional proto field, so a second one is `oneof \
319 field '_system_instruction' is already set`. Set it one way or the other",
320 typed.parts.len()
321 )));
322 }
323
324 let mut tools = if function_tools.is_empty() {
325 Vec::new()
326 } else {
327 vec![serde_json::to_value(Tool::try_from(function_tools)?)?]
328 };
329 tools.append(&mut additional_tools);
330 let tools = if tools.is_empty() { None } else { Some(tools) };
331
332 let tool_config = if let Some(cfg) = tool_choice {
333 Some(ToolConfig {
334 function_calling_config: Some(FunctionCallingMode::try_from(cfg)?),
335 })
336 } else {
337 None
338 };
339 if tool_config.is_some()
342 && let Some(spelling) = smuggled_tool_config
343 {
344 return Err(EncodeError::request(format!(
345 "a Gemini request set the tool choice twice — once as `tool_choice` and once \
346 through `additional_params.{spelling}`. Both would reach the wire, and Gemini \
347 does not take the last — it *merges* them, so the two allowed-function lists \
348 are unioned and the narrower `tool_choice` silently stops restricting anything. \
349 Set it one way or the other"
350 )));
351 }
352
353 let mut request = GenerateContentRequest {
354 contents: full_history
355 .into_iter()
356 .map(|msg| msg.try_into().map_err(EncodeError::request))
357 .collect::<Result<Vec<_>, _>>()?,
358 generation_config,
359 safety_settings: None,
360 tools,
361 tool_config,
362 system_instruction,
363 cached_content: None,
364 additional_params,
365 };
366
367 for name in smuggled_cached_content {
368 request.with_cached_content(&name)?;
369 }
370
371 Ok(request)
372}
373
374pub fn split_system_messages_from_history(
377 history: Vec<completion::Message>,
378) -> (Vec<String>, Vec<completion::Message>) {
379 let mut system = Vec::new();
380 let mut remaining = Vec::new();
381
382 for message in history {
383 match message {
384 completion::Message::System { content } => system.push(content),
385 other => remaining.push(other),
386 }
387 }
388
389 (system, remaining)
390}
391
392const SYSTEM_INSTRUCTION: [&str; 2] = ["systemInstruction", "system_instruction"];
394
395const TOOL_CONFIG: [&str; 2] = ["toolConfig", "tool_config"];
397
398const CACHED_CONTENT: [&str; 2] = ["cachedContent", "cached_content"];
400
401const TOOLS: [&str; 1] = ["tools"];
403
404fn smuggled_field<'a>(payload: &Value, spellings: &[&'a str]) -> Option<&'a str> {
407 let object = payload.as_object()?;
408 spellings
409 .iter()
410 .find(|spelling| object.get(**spelling).is_some_and(|value| !value.is_null()))
411 .copied()
412}
413
414fn extract_tools_from_additional_params(
415 additional_params: &mut Value,
416) -> Result<Vec<Value>, EncodeError> {
417 if let Some(map) = additional_params.as_object_mut()
418 && let Some(raw_tools) = map.remove("tools")
419 {
420 return serde_json::from_value::<Vec<Value>>(raw_tools).map_err(|err| {
421 EncodeError::request(format!(
422 "Invalid Gemini `additional_params.tools` payload: {err}"
423 ))
424 });
425 }
426
427 Ok(Vec::new())
428}
429
430pub(crate) fn resolve_request_model(
431 default_model: &str,
432 completion_request: &CompletionRequest,
433) -> String {
434 completion_request
435 .model
436 .clone()
437 .unwrap_or_else(|| default_model.to_string())
438}
439
440pub(crate) fn completion_endpoint(model: &str) -> String {
441 format!("/v1beta/models/{model}:generateContent")
442}
443
444pub(crate) fn streaming_endpoint(model: &str) -> String {
445 format!("/v1beta/models/{model}:streamGenerateContent")
446}
447
448impl TryFrom<Vec<completion::ToolDefinition>> for Tool {
449 type Error = EncodeError;
450
451 fn try_from(tools: Vec<completion::ToolDefinition>) -> Result<Self, Self::Error> {
452 let mut function_declarations = Vec::new();
453
454 for tool in tools {
455 let parameters = tool_parameters_to_schema(tool.parameters).map_err(|error| {
456 let reason = std::error::Error::source(&error)
458 .map_or_else(|| error.to_string(), ToString::to_string);
459 EncodeError::request(format!(
460 "Tool '{}' could not be converted to a schema: {reason}",
461 tool.name
462 ))
463 })?;
464
465 function_declarations.push(FunctionDeclaration {
466 name: tool.name,
467 description: tool.description,
468 parameters,
469 });
470 }
471
472 Ok(Self {
473 function_declarations,
474 code_execution: None,
475 })
476 }
477}
478
479mod erased_wire {
482 pub(super) trait Wire {
483 fn wire_name(&self) -> String;
484 }
485 impl<T: serde::Serialize> Wire for T {
486 fn wire_name(&self) -> String {
487 match serde_json::to_value(self) {
488 Ok(serde_json::Value::String(name)) => name,
489 Ok(other) => other.to_string(),
490 Err(_) => "<unserializable>".to_owned(),
491 }
492 }
493 }
494}
495
496pub(crate) fn blocked_prompt_error(
499 feedback: &gemini_api_types::PromptFeedback,
500) -> Option<ProviderError> {
501 let reason = match feedback.block_reason.as_ref()? {
502 gemini_api_types::BlockReason::BlockReasonUnspecified => return None,
505 reason => reason,
506 };
507 let wire = |value: &dyn erased_wire::Wire| value.wire_name();
508 let ratings = feedback
509 .safety_ratings
510 .as_ref()
511 .filter(|ratings| !ratings.is_empty())
512 .map(|ratings| {
513 ratings
514 .iter()
515 .map(|rating| format!("{}={}", wire(&rating.category), wire(&rating.probability)))
516 .collect::<Vec<_>>()
517 .join(", ")
518 })
519 .map(|ratings| format!(", safety_ratings=[{ratings}]"))
520 .unwrap_or_default();
521 let message = format!(
522 "Gemini blocked the prompt: block_reason={}{ratings}",
523 reason.as_wire_str()
524 );
525 Some(match reason {
526 gemini_api_types::BlockReason::Safety
527 | gemini_api_types::BlockReason::Blocklist
528 | gemini_api_types::BlockReason::ProhibitedContent
529 | gemini_api_types::BlockReason::BlockReasonUnspecified => ProviderError::ProviderResponse(
530 crate::provider_response::ProviderResponseError::without_status(message)
531 .with_code(Some(reason.as_wire_str().to_owned()))
532 .with_refusal(true),
533 ),
534 gemini_api_types::BlockReason::Other | gemini_api_types::BlockReason::Unknown(_) => {
535 ProviderError::ProviderResponse(
536 crate::provider_response::ProviderResponseError::without_status(message)
537 .with_code(Some(reason.as_wire_str().to_owned()))
538 .with_transient(Some(true)),
539 )
540 }
541 })
542}
543
544pub(crate) fn function_call_finish_reason_error(
545 reason: &FinishReason,
546 finish_message: Option<&str>,
547) -> Option<ProviderError> {
548 match reason {
549 FinishReason::MalformedFunctionCall
550 | FinishReason::UnexpectedToolCall
551 | FinishReason::MissingThoughtSignature
552 | FinishReason::TooManyToolCalls
553 | FinishReason::MalformedResponse => {
554 let message = finish_message.unwrap_or("no finish message provided");
555 Some(ProviderError::Response(format!(
556 "Gemini stopped with finish_reason={reason:?}: {message}"
557 )))
558 }
559 _ => None,
560 }
561}
562
563pub(crate) fn part_kind_name(part: &PartKind) -> &'static str {
565 match part {
566 PartKind::Text(_) => "text",
567 PartKind::InlineData(_) => "inlineData",
568 PartKind::FunctionCall(_) => "functionCall",
569 PartKind::FunctionResponse(_) => "functionResponse",
570 PartKind::FileData(_) => "fileData",
571 PartKind::ExecutableCode(_) => "executableCode",
572 PartKind::CodeExecutionResult(_) => "codeExecutionResult",
573 }
574}
575
576pub mod gemini_api_types {
577 use crate::error::EncodeError;
578 use std::{collections::HashMap, convert::Infallible, str::FromStr};
579
580 use serde::{Deserialize, Serialize};
581 use serde_json::{Value, json};
582
583 use crate::message::{DocumentSourceKind, ImageMediaType, MessageError, MimeType};
584 use crate::{
585 message,
586 providers::gemini::gemini_api_types::{CodeExecutionResult, ExecutableCode},
587 };
588
589 #[derive(Debug, Deserialize, Serialize, Default)]
590 #[serde(rename_all = "camelCase")]
591 pub struct AdditionalParameters {
592 pub generation_config: Option<GenerationConfig>,
594 #[serde(flatten, skip_serializing_if = "Option::is_none")]
596 pub additional_params: Option<serde_json::Value>,
597 }
598
599 impl AdditionalParameters {
600 pub fn with_config(mut self, cfg: GenerationConfig) -> Self {
601 self.generation_config = Some(cfg);
602 self
603 }
604
605 pub fn with_params(mut self, params: serde_json::Value) -> Self {
606 self.additional_params = Some(params);
607 self
608 }
609 }
610
611 #[derive(Debug, Deserialize, Serialize)]
621 #[serde(rename_all = "camelCase")]
622 pub struct GenerateContentResponse {
623 #[serde(default)]
624 pub response_id: String,
625 #[serde(default)]
627 pub candidates: Vec<ContentCandidate>,
628 pub prompt_feedback: Option<PromptFeedback>,
631 pub usage_metadata: Option<UsageMetadata>,
633 pub model_version: Option<String>,
634 #[serde(default, skip_serializing_if = "Option::is_none")]
639 pub error: Option<Value>,
640 }
641
642 pub(crate) fn visible_text_parts(content: &Content) -> impl Iterator<Item = &str> {
645 content.parts.iter().filter_map(|part| match &part.part {
646 PartKind::Text(text) if !part.thought.unwrap_or(false) => Some(text.as_str()),
647 _ => None,
648 })
649 }
650
651 #[derive(Clone, Debug, Deserialize, Serialize)]
653 #[serde(rename_all = "camelCase")]
654 pub struct ContentCandidate {
655 #[serde(skip_serializing_if = "Option::is_none")]
657 pub content: Option<Content>,
658 pub finish_reason: Option<FinishReason>,
661 pub safety_ratings: Option<Vec<SafetyRating>>,
664 pub citation_metadata: Option<CitationMetadata>,
668 pub token_count: Option<i32>,
670 pub avg_logprobs: Option<f64>,
672 pub logprobs_result: Option<LogprobsResult>,
674 pub index: Option<i32>,
676 pub finish_message: Option<String>,
678 }
679
680 #[derive(Clone, Debug, Deserialize, Serialize)]
681 pub struct Content {
682 #[serde(default)]
684 pub parts: Vec<Part>,
685 pub role: Option<Role>,
688 }
689
690 impl TryFrom<message::Message> for Content {
691 type Error = message::MessageError;
692
693 fn try_from(msg: message::Message) -> Result<Self, Self::Error> {
694 Ok(match msg {
695 message::Message::System { content } => Content {
696 parts: vec![content.into()],
697 role: Some(Role::User),
698 },
699 message::Message::User { content } => Content {
700 parts: content
701 .into_iter()
702 .map(std::convert::TryInto::try_into)
703 .collect::<Result<Vec<_>, _>>()?,
704 role: Some(Role::User),
705 },
706 message::Message::Assistant { content, .. } => Content {
707 role: Some(Role::Model),
708 parts: content
709 .into_iter()
710 .filter(|part| match part {
712 message::AssistantContent::Reasoning(reasoning) => reasoning
713 .open(&crate::providers::gemini::completion::ISSUER)
714 .is_some(),
715 _ => true,
716 })
717 .map(std::convert::TryInto::try_into)
718 .collect::<Result<Vec<_>, _>>()?,
719 },
720 })
721 }
722 }
723
724 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
725 #[serde(rename_all = "lowercase")]
726 pub enum Role {
727 User,
728 Model,
729 }
730
731 #[derive(Debug, Default, Deserialize, Serialize, Clone, PartialEq)]
732 #[serde(rename_all = "camelCase")]
733 pub struct Part {
734 #[serde(skip_serializing_if = "Option::is_none")]
736 pub thought: Option<bool>,
737 #[serde(skip_serializing_if = "Option::is_none")]
739 pub thought_signature: Option<String>,
740 #[serde(flatten)]
741 pub part: PartKind,
742 #[serde(flatten, skip_serializing_if = "Option::is_none")]
743 pub additional_params: Option<Value>,
744 }
745
746 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
749 #[serde(rename_all = "camelCase")]
750 pub enum PartKind {
751 Text(String),
752 InlineData(Blob),
753 FunctionCall(FunctionCall),
754 FunctionResponse(FunctionResponse),
755 FileData(FileData),
756 ExecutableCode(ExecutableCode),
757 CodeExecutionResult(CodeExecutionResult),
758 }
759
760 impl Default for PartKind {
761 fn default() -> Self {
762 Self::Text(String::new())
763 }
764 }
765
766 impl From<String> for Part {
767 fn from(text: String) -> Self {
768 Self {
769 thought: Some(false),
770 thought_signature: None,
771 part: PartKind::Text(text),
772 additional_params: None,
773 }
774 }
775 }
776
777 impl From<&str> for Part {
778 fn from(text: &str) -> Self {
779 Self::from(text.to_string())
780 }
781 }
782
783 impl FromStr for Part {
784 type Err = Infallible;
785
786 fn from_str(s: &str) -> Result<Self, Self::Err> {
787 Ok(s.into())
788 }
789 }
790
791 fn media_source_to_part_kind(
795 kind: &str,
796 mime_type: String,
797 source: DocumentSourceKind,
798 string_is_data: bool,
799 ) -> Result<PartKind, message::MessageError> {
800 match source {
801 DocumentSourceKind::Url(file_uri) => Ok(PartKind::FileData(FileData {
802 mime_type: Some(mime_type),
803 file_uri,
804 })),
805 DocumentSourceKind::Base64(data) => Ok(PartKind::InlineData(Blob { mime_type, data })),
806 DocumentSourceKind::String(data) if string_is_data => {
807 Ok(PartKind::InlineData(Blob { mime_type, data }))
808 }
809 DocumentSourceKind::String(_) => Err(message::MessageError::ConversionError(format!(
810 "Strings cannot be used as Gemini {kind} inputs"
811 ))),
812 DocumentSourceKind::Raw(_) => Err(message::MessageError::ConversionError(
813 "Raw files not supported, encode as base64 first".to_string(),
814 )),
815 DocumentSourceKind::FileId(_) => Err(message::MessageError::ConversionError(format!(
816 "Provider file IDs are not supported for Gemini {kind} inputs"
817 ))),
818 DocumentSourceKind::Unknown => Err(message::MessageError::ConversionError(format!(
819 "Gemini {kind} input has no body"
820 ))),
821 }
822 }
823
824 impl TryFrom<(ImageMediaType, DocumentSourceKind)> for PartKind {
825 type Error = message::MessageError;
826 fn try_from(
827 (mime_type, doc_src): (ImageMediaType, DocumentSourceKind),
828 ) -> Result<Self, Self::Error> {
829 media_source_to_part_kind("image", mime_type.to_mime_type().to_string(), doc_src, true)
830 }
831 }
832
833 fn image_to_part(image: message::Image) -> Result<Part, message::MessageError> {
838 let message::Image {
839 data, media_type, ..
840 } = image;
841
842 let Some(media_type) = media_type else {
843 return Err(message::MessageError::ConversionError(
844 "Media type for image is required for Gemini".to_string(),
845 ));
846 };
847
848 match media_type {
849 message::ImageMediaType::JPEG
850 | message::ImageMediaType::PNG
851 | message::ImageMediaType::WEBP
852 | message::ImageMediaType::HEIC
853 | message::ImageMediaType::HEIF => Ok(Part {
854 thought: Some(false),
855 thought_signature: None,
856 part: PartKind::try_from((media_type, data))?,
857 additional_params: None,
858 }),
859 _ => Err(message::MessageError::ConversionError(format!(
860 "Unsupported image media type {media_type:?}"
861 ))),
862 }
863 }
864
865 fn gemini_tool_result_image_mime_type(
866 media_type: Option<&ImageMediaType>,
867 ) -> Result<&'static str, MessageError> {
868 let media_type = media_type.ok_or_else(|| {
869 MessageError::ConversionError(
870 "Image media type is required for Gemini tool results".to_string(),
871 )
872 })?;
873
874 match media_type {
875 ImageMediaType::JPEG | ImageMediaType::PNG | ImageMediaType::WEBP => {
876 Ok(media_type.to_mime_type())
877 }
878 _ => Err(MessageError::ConversionError(format!(
879 "Unsupported image media type {media_type:?} for Gemini tool results; supported types are JPEG, PNG, and WEBP"
880 ))),
881 }
882 }
883
884 impl TryFrom<message::UserContent> for Part {
885 type Error = message::MessageError;
886
887 fn try_from(content: message::UserContent) -> Result<Self, Self::Error> {
888 match content {
889 message::UserContent::Text(message::Text { text, .. }) => Ok(Part {
890 thought: Some(false),
891 thought_signature: None,
892 part: PartKind::Text(text),
893 additional_params: None,
894 }),
895 message::UserContent::ToolResult(message::ToolResult {
896 call,
897 name,
898 content,
899 }) => {
900 let function_name = name;
901 let mut response_values = Vec::new();
902 let mut parts: Vec<FunctionResponsePart> = Vec::new();
903
904 for item in content.iter() {
905 match item {
906 message::ToolResultContent::Text(text) => {
907 response_values.push(json!(&text.text));
908 }
909 message::ToolResultContent::Json { value } => {
910 response_values.push(value.clone());
911 }
912 message::ToolResultContent::Image(image) => {
913 let part = match &image.data {
914 DocumentSourceKind::Base64(b64) => {
915 let mime_type = gemini_tool_result_image_mime_type(
916 image.media_type.as_ref(),
917 )?;
918
919 FunctionResponsePart {
922 inline_data: Some(FunctionResponseInlineData {
923 mime_type: mime_type.to_string(),
924 data: b64.clone(),
925 display_name: None,
926 }),
927 file_data: None,
928 }
929 }
930 DocumentSourceKind::Url(_) => {
931 return Err(message::MessageError::ConversionError(
932 "Gemini tool result images must use base64 inline data; URL-backed images are not supported"
933 .to_string(),
934 ));
935 }
936 _ => {
937 return Err(message::MessageError::ConversionError(
938 "Unsupported image source kind for tool results"
939 .to_string(),
940 ));
941 }
942 };
943 parts.push(part);
944 }
945 }
946 }
947
948 let response_json = if response_values.is_empty() {
949 None
950 } else {
951 let result = if response_values.len() == 1 {
952 response_values.remove(0)
953 } else {
954 serde_json::Value::Array(response_values)
955 };
956 Some(json!({ "result": result }))
957 };
958
959 Ok(Part {
960 thought: Some(false),
961 thought_signature: None,
962 part: PartKind::FunctionResponse(FunctionResponse {
963 name: function_name.into(),
964 id: call.provider().map(|provider| provider.call_id.clone()),
965 response: response_json,
966 parts: if parts.is_empty() { None } else { Some(parts) },
967 }),
968 additional_params: None,
969 })
970 }
971 message::UserContent::Image(image) => image_to_part(image),
972 message::UserContent::Document(message::Document {
973 data, media_type, ..
974 }) => {
975 let Some(media_type) = media_type else {
976 return Err(MessageError::ConversionError(
977 "A mime type is required for document inputs to Gemini".to_string(),
978 ));
979 };
980
981 if matches!(
984 media_type,
985 message::DocumentMediaType::TXT
986 | message::DocumentMediaType::RTF
987 | message::DocumentMediaType::HTML
988 | message::DocumentMediaType::CSS
989 | message::DocumentMediaType::MARKDOWN
990 | message::DocumentMediaType::CSV
991 | message::DocumentMediaType::XML
992 | message::DocumentMediaType::Javascript
993 | message::DocumentMediaType::Python
994 ) {
995 use base64::Engine;
996 let part = match data {
997 DocumentSourceKind::String(text) => PartKind::Text(text),
998 DocumentSourceKind::Base64(data) => {
999 let text = String::from_utf8(
1000 base64::engine::general_purpose::STANDARD
1001 .decode(&data)
1002 .map_err(|e| {
1003 MessageError::ConversionError(format!(
1004 "Failed to decode base64: {e}"
1005 ))
1006 })?,
1007 )
1008 .map_err(|e| {
1009 MessageError::ConversionError(format!(
1010 "Invalid UTF-8 in document: {e}"
1011 ))
1012 })?;
1013 PartKind::Text(text)
1014 }
1015 DocumentSourceKind::Url(file_uri) => PartKind::FileData(FileData {
1016 mime_type: Some(media_type.to_mime_type().to_string()),
1017 file_uri,
1018 }),
1019 DocumentSourceKind::Raw(_) => {
1020 return Err(MessageError::ConversionError(
1021 "Raw files not supported, encode as base64 first".to_string(),
1022 ));
1023 }
1024 DocumentSourceKind::FileId(_) => {
1025 return Err(MessageError::ConversionError(
1026 "Provider file IDs are not supported for Gemini documents"
1027 .to_string(),
1028 ));
1029 }
1030 DocumentSourceKind::Unknown => {
1031 return Err(MessageError::ConversionError(
1032 "Document has no body".to_string(),
1033 ));
1034 }
1035 };
1036
1037 Ok(Part {
1038 thought: Some(false),
1039 part,
1040 ..Default::default()
1041 })
1042 } else if !media_type.is_code() {
1043 let part = media_source_to_part_kind(
1044 "document",
1045 media_type.to_mime_type().to_string(),
1046 data,
1047 true,
1048 )?;
1049
1050 Ok(Part {
1051 thought: Some(false),
1052 part,
1053 ..Default::default()
1054 })
1055 } else {
1056 Err(message::MessageError::ConversionError(format!(
1057 "Unsupported document media type {media_type:?}"
1058 )))
1059 }
1060 }
1061
1062 message::UserContent::Audio(message::Audio {
1063 data, media_type, ..
1064 }) => {
1065 let Some(media_type) = media_type else {
1066 return Err(MessageError::ConversionError(
1067 "A mime type is required for audio inputs to Gemini".to_string(),
1068 ));
1069 };
1070
1071 let part = media_source_to_part_kind(
1072 "audio",
1073 media_type.to_mime_type().to_string(),
1074 data,
1075 false,
1076 )?;
1077
1078 Ok(Part {
1079 thought: Some(false),
1080 part,
1081 ..Default::default()
1082 })
1083 }
1084 message::UserContent::Video(message::Video {
1085 data,
1086 media_type,
1087 additional_params,
1088 ..
1089 }) => {
1090 let mime_type = media_type.map(|media_ty| media_ty.to_mime_type().to_string());
1091
1092 let part = match data {
1093 DocumentSourceKind::Url(file_uri)
1097 if file_uri.starts_with("https://www.youtube.com") =>
1098 {
1099 PartKind::FileData(FileData {
1100 mime_type,
1101 file_uri,
1102 })
1103 }
1104 data => {
1105 let mime_type = mime_type.ok_or_else(|| {
1106 MessageError::ConversionError(
1107 "A mime type is required for non-Youtube video inputs to Gemini"
1108 .to_string(),
1109 )
1110 })?;
1111
1112 media_source_to_part_kind("video", mime_type, data, false)?
1113 }
1114 };
1115
1116 Ok(Part {
1117 thought: Some(false),
1118 thought_signature: None,
1119 part,
1120 additional_params: additional_params.map(Into::into),
1121 })
1122 }
1123 }
1124 }
1125 }
1126
1127 impl TryFrom<message::AssistantContent> for Part {
1128 type Error = message::MessageError;
1129
1130 fn try_from(content: message::AssistantContent) -> Result<Self, Self::Error> {
1131 match content {
1132 message::AssistantContent::Text(text) => {
1133 let thought_signature =
1135 super::super::text_thought_signature(&text).map(str::to_owned);
1136 Ok(Part {
1137 thought_signature,
1138 ..text.text.into()
1139 })
1140 }
1141 message::AssistantContent::Image(image) => image_to_part(image),
1142 message::AssistantContent::ToolCall(tool_call) => Ok(tool_call.into()),
1143 message::AssistantContent::Reasoning(reasoning) => {
1144 let reasoning = reasoning
1145 .open(&crate::providers::gemini::completion::ISSUER)
1146 .ok_or_else(|| {
1147 MessageError::ConversionError(
1148 "Gemini cannot replay reasoning another service issued".to_owned(),
1149 )
1150 })?;
1151 Ok(Part {
1152 thought: Some(true),
1153 thought_signature: reasoning.first_signature().map(str::to_owned),
1154 part: PartKind::Text(reasoning.display_text()),
1155 additional_params: None,
1156 })
1157 }
1158 }
1159 }
1160 }
1161
1162 impl From<message::ToolCall> for Part {
1163 fn from(tool_call: message::ToolCall) -> Self {
1164 Self {
1165 thought: Some(false),
1166 thought_signature: tool_call.signature,
1167 part: PartKind::FunctionCall(FunctionCall {
1168 name: tool_call.function.name.into(),
1169 args: tool_call.function.arguments,
1170 id: tool_call
1173 .id
1174 .provider()
1175 .map(|provider| provider.call_id.clone()),
1176 }),
1177 additional_params: None,
1178 }
1179 }
1180 }
1181
1182 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
1185 #[serde(rename_all = "camelCase")]
1186 pub struct Blob {
1187 pub mime_type: String,
1190 pub data: String,
1192 }
1193
1194 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
1196 pub struct FunctionCall {
1197 pub name: String,
1200 pub args: serde_json::Value,
1202 #[serde(skip_serializing_if = "Option::is_none")]
1204 pub id: Option<String>,
1205 }
1206
1207 impl From<message::ToolCall> for FunctionCall {
1208 fn from(tool_call: message::ToolCall) -> Self {
1209 Self {
1210 name: tool_call.function.name.into(),
1211 args: tool_call.function.arguments,
1212 id: tool_call
1213 .id
1214 .provider()
1215 .map(|provider| provider.call_id.clone()),
1216 }
1217 }
1218 }
1219
1220 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
1222 pub struct FunctionResponse {
1223 pub name: String,
1226 #[serde(skip_serializing_if = "Option::is_none")]
1228 pub id: Option<String>,
1229 #[serde(skip_serializing_if = "Option::is_none")]
1231 pub response: Option<serde_json::Value>,
1232 #[serde(skip_serializing_if = "Option::is_none")]
1234 pub parts: Option<Vec<FunctionResponsePart>>,
1235 }
1236
1237 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
1239 #[serde(rename_all = "camelCase")]
1240 pub struct FunctionResponsePart {
1241 #[serde(skip_serializing_if = "Option::is_none")]
1243 pub inline_data: Option<FunctionResponseInlineData>,
1244 #[serde(skip_serializing_if = "Option::is_none")]
1246 pub file_data: Option<FileData>,
1247 }
1248
1249 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
1251 #[serde(rename_all = "camelCase")]
1252 pub struct FunctionResponseInlineData {
1253 pub mime_type: String,
1255 pub data: String,
1257 #[serde(skip_serializing_if = "Option::is_none")]
1259 pub display_name: Option<String>,
1260 }
1261
1262 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
1264 #[serde(rename_all = "camelCase")]
1265 pub struct FileData {
1266 pub mime_type: Option<String>,
1268 pub file_uri: String,
1270 }
1271
1272 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
1273 pub struct SafetyRating {
1274 pub category: HarmCategory,
1275 pub probability: HarmProbability,
1276 }
1277
1278 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
1279 #[serde(rename_all = "SCREAMING_SNAKE_CASE")]
1280 pub enum HarmProbability {
1281 HarmProbabilityUnspecified,
1282 Negligible,
1283 Low,
1284 Medium,
1285 High,
1286 #[serde(untagged)]
1289 Unknown(String),
1290 }
1291
1292 #[derive(Debug, Deserialize, Serialize, Clone, PartialEq)]
1293 #[serde(rename_all = "SCREAMING_SNAKE_CASE")]
1294 pub enum HarmCategory {
1295 HarmCategoryUnspecified,
1296 HarmCategoryDerogatory,
1297 HarmCategoryToxicity,
1298 HarmCategoryViolence,
1299 HarmCategorySexually,
1300 HarmCategoryMedical,
1301 HarmCategoryDangerous,
1302 HarmCategoryHarassment,
1303 HarmCategoryHateSpeech,
1304 HarmCategorySexuallyExplicit,
1305 HarmCategoryDangerousContent,
1306 HarmCategoryCivicIntegrity,
1307 #[serde(untagged)]
1309 Unknown(String),
1310 }
1311
1312 #[derive(Debug, Deserialize, Clone, Default, Serialize)]
1313 #[serde(rename_all = "camelCase")]
1314 pub struct UsageMetadata {
1315 #[serde(default)]
1316 pub prompt_token_count: i32,
1317 #[serde(skip_serializing_if = "Option::is_none")]
1318 pub cached_content_token_count: Option<i32>,
1319 #[serde(skip_serializing_if = "Option::is_none")]
1320 pub candidates_token_count: Option<i32>,
1321 #[serde(default)]
1322 pub total_token_count: i32,
1323 #[serde(skip_serializing_if = "Option::is_none")]
1324 pub thoughts_token_count: Option<i32>,
1325 #[serde(default, skip_serializing_if = "Option::is_none")]
1326 pub prompt_tokens_details: Option<Vec<ModalityTokenCount>>,
1327 #[serde(default, skip_serializing_if = "Option::is_none")]
1328 pub cache_tokens_details: Option<Vec<ModalityTokenCount>>,
1329 #[serde(default, skip_serializing_if = "Option::is_none")]
1330 pub candidates_tokens_details: Option<Vec<ModalityTokenCount>>,
1331 #[serde(default, skip_serializing_if = "Option::is_none")]
1332 pub tool_use_prompt_token_count: Option<i32>,
1333 #[serde(default, skip_serializing_if = "Option::is_none")]
1334 pub tool_use_prompt_tokens_details: Option<Vec<ModalityTokenCount>>,
1335 #[serde(default, skip_serializing_if = "Option::is_none")]
1336 pub traffic_type: Option<TrafficType>,
1337 }
1338
1339 #[derive(Clone, Debug, Deserialize, Serialize)]
1340 #[serde(rename_all = "camelCase")]
1341 pub struct ModalityTokenCount {
1342 pub modality: Modality,
1343 #[serde(default)]
1344 pub token_count: i32,
1345 }
1346
1347 #[derive(Clone, Debug, Deserialize, Serialize)]
1348 #[serde(rename_all = "SCREAMING_SNAKE_CASE")]
1349 pub enum Modality {
1350 ModalityUnspecified,
1351 Text,
1352 Image,
1353 Video,
1354 Audio,
1355 Document,
1356 }
1357
1358 #[derive(Clone, Debug, Deserialize, Serialize)]
1359 #[serde(rename_all = "SCREAMING_SNAKE_CASE")]
1360 pub enum TrafficType {
1361 TrafficTypeUnspecified,
1362 OnDemand,
1363 ProvisionedThroughput,
1364 }
1365
1366 impl From<&UsageMetadata> for crate::completion::Usage {
1371 fn from(value: &UsageMetadata) -> crate::completion::Usage {
1372 let count = |count: i32| count as u64;
1373 let tool_use = value.tool_use_prompt_token_count.map_or(0, count);
1376 let input = count(value.prompt_token_count).saturating_add(tool_use);
1377 let thoughts = value.thoughts_token_count.map_or(0, count);
1378 let output = value
1379 .candidates_token_count
1380 .map_or(0, count)
1381 .saturating_add(thoughts);
1382 crate::completion::Usage {
1383 input_tokens: Some(input),
1384 output_tokens: Some(output),
1385 cached_input_tokens: value.cached_content_token_count.map(count),
1386 reasoning_tokens: value.thoughts_token_count.map(count),
1387 tool_use_prompt_tokens: value.tool_use_prompt_token_count.map(count),
1388 total_tokens: Some(input.saturating_add(output)),
1389 cache_creation_input_tokens: None,
1390 }
1391 }
1392 }
1393
1394 #[derive(Debug, Deserialize, Serialize)]
1396 #[serde(rename_all = "camelCase")]
1397 pub struct PromptFeedback {
1398 pub block_reason: Option<BlockReason>,
1400 pub safety_ratings: Option<Vec<SafetyRating>>,
1402 }
1403
1404 #[derive(Debug, Deserialize, Serialize)]
1406 #[serde(rename_all = "SCREAMING_SNAKE_CASE")]
1407 pub enum BlockReason {
1408 BlockReasonUnspecified,
1410 Safety,
1412 Other,
1414 Blocklist,
1416 ProhibitedContent,
1418 #[serde(untagged)]
1422 Unknown(String),
1423 }
1424
1425 impl BlockReason {
1426 pub fn as_wire_str(&self) -> &str {
1429 match self {
1430 Self::BlockReasonUnspecified => "BLOCK_REASON_UNSPECIFIED",
1431 Self::Safety => "SAFETY",
1432 Self::Other => "OTHER",
1433 Self::Blocklist => "BLOCKLIST",
1434 Self::ProhibitedContent => "PROHIBITED_CONTENT",
1435 Self::Unknown(raw) => raw.as_str(),
1436 }
1437 }
1438 }
1439
1440 #[derive(Clone, Debug, Deserialize, Serialize)]
1441 #[serde(rename_all = "SCREAMING_SNAKE_CASE")]
1442 pub enum FinishReason {
1443 FinishReasonUnspecified,
1445 Stop,
1447 MaxTokens,
1449 Safety,
1451 Recitation,
1453 Language,
1455 Other,
1457 Blocklist,
1459 ProhibitedContent,
1461 Spii,
1463 MalformedFunctionCall,
1465 UnexpectedToolCall,
1467 MissingThoughtSignature,
1469 TooManyToolCalls,
1471 MalformedResponse,
1473 #[serde(untagged)]
1475 Unknown(String),
1476 }
1477
1478 impl FinishReason {
1479 pub fn as_wire_str(&self) -> &str {
1485 match self {
1486 Self::FinishReasonUnspecified => "FINISH_REASON_UNSPECIFIED",
1487 Self::Stop => "STOP",
1488 Self::MaxTokens => "MAX_TOKENS",
1489 Self::Safety => "SAFETY",
1490 Self::Recitation => "RECITATION",
1491 Self::Language => "LANGUAGE",
1492 Self::Other => "OTHER",
1493 Self::Blocklist => "BLOCKLIST",
1494 Self::ProhibitedContent => "PROHIBITED_CONTENT",
1495 Self::Spii => "SPII",
1496 Self::MalformedFunctionCall => "MALFORMED_FUNCTION_CALL",
1497 Self::UnexpectedToolCall => "UNEXPECTED_TOOL_CALL",
1498 Self::MissingThoughtSignature => "MISSING_THOUGHT_SIGNATURE",
1499 Self::TooManyToolCalls => "TOO_MANY_TOOL_CALLS",
1500 Self::MalformedResponse => "MALFORMED_RESPONSE",
1501 Self::Unknown(reason) => reason,
1502 }
1503 }
1504 }
1505
1506 pub fn map_google_finish_reason(wire_name: &str) -> Option<crate::completion::FinishReason> {
1510 Some(match wire_name {
1511 "FINISH_REASON_UNSPECIFIED" => return None,
1512 "STOP" => crate::completion::FinishReason::Stop,
1513 "MAX_TOKENS" => crate::completion::FinishReason::Length,
1514 "SAFETY" | "BLOCKLIST" | "PROHIBITED_CONTENT" | "SPII" => {
1515 crate::completion::FinishReason::ContentFilter
1516 }
1517 other => crate::completion::FinishReason::Other(other.to_owned()),
1518 })
1519 }
1520
1521 pub(crate) fn map_finish_reason(
1525 reason: &FinishReason,
1526 ) -> Option<crate::completion::FinishReason> {
1527 map_google_finish_reason(reason.as_wire_str())
1528 }
1529
1530 #[derive(Clone, Debug, Deserialize, Serialize)]
1531 #[serde(rename_all = "camelCase")]
1532 pub struct CitationMetadata {
1533 #[serde(default)]
1534 pub citation_sources: Vec<CitationSource>,
1535 }
1536
1537 #[derive(Clone, Debug, Deserialize, Serialize)]
1538 #[serde(rename_all = "camelCase")]
1539 pub struct CitationSource {
1540 #[serde(skip_serializing_if = "Option::is_none")]
1541 pub uri: Option<String>,
1542 #[serde(skip_serializing_if = "Option::is_none")]
1543 pub start_index: Option<i32>,
1544 #[serde(skip_serializing_if = "Option::is_none")]
1545 pub end_index: Option<i32>,
1546 #[serde(skip_serializing_if = "Option::is_none")]
1547 pub license: Option<String>,
1548 }
1549
1550 #[derive(Clone, Debug, Deserialize, Serialize)]
1551 #[serde(rename_all = "camelCase")]
1552 pub struct LogprobsResult {
1553 #[serde(default)]
1554 pub top_candidates: Vec<TopCandidate>,
1555 #[serde(skip_serializing_if = "Option::is_none")]
1556 pub log_probability_sum: Option<f64>,
1557 #[serde(default)]
1558 pub chosen_candidates: Vec<LogProbCandidate>,
1559 }
1560
1561 #[derive(Clone, Debug, Deserialize, Serialize)]
1562 pub struct TopCandidate {
1563 #[serde(default)]
1564 pub candidates: Vec<LogProbCandidate>,
1565 }
1566
1567 #[derive(Clone, Debug, Deserialize, Serialize)]
1568 #[serde(rename_all = "camelCase")]
1569 pub struct LogProbCandidate {
1570 #[serde(skip_serializing_if = "Option::is_none")]
1571 pub token: Option<String>,
1572 #[serde(skip_serializing_if = "Option::is_none")]
1573 pub token_id: Option<i32>,
1574 #[serde(skip_serializing_if = "Option::is_none")]
1575 pub log_probability: Option<f64>,
1576 }
1577
1578 #[derive(Debug, Default, Deserialize, Serialize)]
1582 #[serde(rename_all = "camelCase")]
1583 pub struct GenerationConfig {
1584 #[serde(skip_serializing_if = "Option::is_none")]
1587 pub stop_sequences: Option<Vec<String>>,
1588 #[serde(skip_serializing_if = "Option::is_none")]
1591 pub response_mime_type: Option<String>,
1592 #[serde(skip_serializing_if = "Option::is_none")]
1595 pub response_schema: Option<Schema>,
1596 #[serde(
1602 skip_serializing_if = "Option::is_none",
1603 rename = "_responseJsonSchema"
1604 )]
1605 pub _response_json_schema: Option<Value>,
1606 #[serde(skip_serializing_if = "Option::is_none")]
1608 pub response_json_schema: Option<Value>,
1609 #[serde(skip_serializing_if = "Option::is_none")]
1612 pub candidate_count: Option<i32>,
1613 #[serde(skip_serializing_if = "Option::is_none")]
1615 pub max_output_tokens: Option<u64>,
1616 #[serde(skip_serializing_if = "Option::is_none")]
1618 pub temperature: Option<f64>,
1619 #[serde(skip_serializing_if = "Option::is_none")]
1622 pub top_p: Option<f64>,
1623 #[serde(skip_serializing_if = "Option::is_none")]
1626 pub top_k: Option<i32>,
1627 #[serde(skip_serializing_if = "Option::is_none")]
1630 pub presence_penalty: Option<f64>,
1631 #[serde(skip_serializing_if = "Option::is_none")]
1635 pub frequency_penalty: Option<f64>,
1636 #[serde(skip_serializing_if = "Option::is_none")]
1638 pub response_logprobs: Option<bool>,
1639 #[serde(skip_serializing_if = "Option::is_none")]
1642 pub logprobs: Option<i32>,
1643 #[serde(skip_serializing_if = "Option::is_none")]
1645 pub thinking_config: Option<ThinkingConfig>,
1646 #[serde(skip_serializing_if = "Option::is_none")]
1648 pub response_modalities: Option<Vec<ResponseModality>>,
1649 #[serde(skip_serializing_if = "Option::is_none")]
1650 pub image_config: Option<ImageConfig>,
1651 }
1652
1653 #[derive(Clone, Debug, Deserialize, Serialize, PartialEq)]
1655 #[serde(rename_all = "SCREAMING_SNAKE_CASE")]
1656 pub enum ResponseModality {
1657 Text,
1658 Image,
1659 Audio,
1660 }
1661
1662 #[derive(Clone, Debug, Deserialize, Serialize, PartialEq)]
1664 #[serde(rename_all = "snake_case")]
1665 pub enum ThinkingLevel {
1666 Minimal,
1667 Low,
1668 Medium,
1669 High,
1670 }
1671
1672 #[derive(Debug, Deserialize, Serialize)]
1676 #[serde(rename_all = "camelCase")]
1677 pub struct ThinkingConfig {
1678 #[serde(skip_serializing_if = "Option::is_none")]
1680 pub thinking_budget: Option<u32>,
1681 #[serde(skip_serializing_if = "Option::is_none")]
1683 pub thinking_level: Option<ThinkingLevel>,
1684 #[serde(skip_serializing_if = "Option::is_none")]
1686 pub include_thoughts: Option<bool>,
1687 }
1688
1689 #[derive(Debug, Deserialize, Serialize)]
1690 #[serde(rename_all = "camelCase")]
1691 pub struct ImageConfig {
1692 #[serde(skip_serializing_if = "Option::is_none")]
1693 pub aspect_ratio: Option<String>,
1694 #[serde(skip_serializing_if = "Option::is_none")]
1695 pub image_size: Option<String>,
1696 }
1697
1698 #[derive(Debug, Deserialize, Serialize, Clone)]
1702 pub struct Schema {
1703 pub r#type: String,
1704 #[serde(skip_serializing_if = "Option::is_none")]
1705 pub format: Option<String>,
1706 #[serde(skip_serializing_if = "Option::is_none")]
1707 pub description: Option<String>,
1708 #[serde(skip_serializing_if = "Option::is_none")]
1709 pub nullable: Option<bool>,
1710 #[serde(skip_serializing_if = "Option::is_none")]
1711 pub r#enum: Option<Vec<String>>,
1712 #[serde(skip_serializing_if = "Option::is_none")]
1713 pub max_items: Option<i32>,
1714 #[serde(skip_serializing_if = "Option::is_none")]
1715 pub min_items: Option<i32>,
1716 #[serde(
1718 skip_serializing_if = "Option::is_none",
1719 serialize_with = "crate::json_utils::serialize_optional_map_sorted"
1720 )]
1721 pub properties: Option<HashMap<String, Schema>>,
1722 #[serde(skip_serializing_if = "Option::is_none")]
1723 pub required: Option<Vec<String>>,
1724 #[serde(skip_serializing_if = "Option::is_none")]
1725 pub items: Option<Box<Schema>>,
1726 }
1727
1728 pub fn tool_parameters_to_schema(parameters: Value) -> Result<Option<Schema>, EncodeError> {
1734 if parameters.is_null() || parameters == json!({"type": "object", "properties": {}}) {
1735 Ok(None)
1736 } else {
1737 parameters.try_into().map(Some)
1738 }
1739 }
1740
1741 pub fn flatten_schema(mut schema: Value) -> Result<Value, EncodeError> {
1746 let defs = schema
1747 .as_object()
1748 .and_then(|obj| obj.get("$defs").or_else(|| obj.get("definitions")))
1749 .cloned();
1750
1751 let Some(defs_value) = defs else {
1752 return Ok(schema);
1753 };
1754
1755 let Some(defs_obj) = defs_value.as_object() else {
1756 return Err(EncodeError::request("$defs must be an object"));
1757 };
1758
1759 resolve_refs(&mut schema, defs_obj)?;
1760
1761 if let Some(obj) = schema.as_object_mut() {
1762 obj.remove("$defs");
1763 obj.remove("definitions");
1764 }
1765
1766 Ok(schema)
1767 }
1768
1769 fn resolve_refs(
1772 value: &mut Value,
1773 defs: &serde_json::Map<String, Value>,
1774 ) -> Result<(), EncodeError> {
1775 match value {
1776 Value::Object(obj) => {
1777 if let Some(ref_value) = obj.get("$ref")
1778 && let Some(ref_str) = ref_value.as_str()
1779 {
1780 let def_name = parse_ref_path(ref_str)?;
1781
1782 let def = defs.get(&def_name).ok_or_else(|| {
1783 EncodeError::request(format!("Reference not found: {ref_str}"))
1784 })?;
1785
1786 let mut resolved = def.clone();
1787 resolve_refs(&mut resolved, defs)?;
1788 *value = resolved;
1789 return Ok(());
1790 }
1791
1792 for (_, v) in obj.iter_mut() {
1793 resolve_refs(v, defs)?;
1794 }
1795 }
1796 Value::Array(arr) => {
1797 for item in arr.iter_mut() {
1798 resolve_refs(item, defs)?;
1799 }
1800 }
1801 _ => {}
1802 }
1803
1804 Ok(())
1805 }
1806
1807 fn parse_ref_path(ref_str: &str) -> Result<String, EncodeError> {
1810 if let Some(fragment) = ref_str.strip_prefix('#') {
1811 if let Some(name) = fragment.strip_prefix("/$defs/") {
1812 Ok(name.to_string())
1813 } else if let Some(name) = fragment.strip_prefix("/definitions/") {
1814 Ok(name.to_string())
1815 } else {
1816 Err(EncodeError::request(format!(
1817 "Unsupported reference format: {ref_str}"
1818 )))
1819 }
1820 } else {
1821 Err(EncodeError::request(format!(
1822 "Only fragment references (#/...) are supported: {ref_str}"
1823 )))
1824 }
1825 }
1826
1827 fn extract_type(type_value: &Value) -> Option<String> {
1830 if let Some(t) = type_value.as_str() {
1831 return Some(t.to_string());
1832 }
1833
1834 type_value.as_array().and_then(|arr| {
1835 arr.iter()
1836 .filter_map(|v| v.as_str())
1837 .find(|t| *t != "null")
1838 .or_else(|| arr.iter().find_map(|v| v.as_str()))
1839 .map(str::to_owned)
1840 })
1841 }
1842
1843 fn schema_is_null(obj: &serde_json::Map<String, Value>) -> bool {
1844 obj.get("type")
1845 .and_then(extract_type)
1846 .as_deref()
1847 .is_some_and(|t| t == "null")
1848 }
1849
1850 fn schema_is_nullable(obj: &serde_json::Map<String, Value>) -> bool {
1851 obj.get("nullable")
1852 .and_then(serde_json::Value::as_bool)
1853 .unwrap_or(false)
1854 || obj
1855 .get("type")
1856 .and_then(|v| v.as_array())
1857 .is_some_and(|arr| arr.iter().any(|v| v.as_str() == Some("null")))
1858 || ["anyOf", "oneOf", "allOf"].iter().any(|key| {
1859 obj.get(*key).and_then(|v| v.as_array()).is_some_and(|arr| {
1860 arr.iter()
1861 .filter_map(|schema| schema.as_object())
1862 .any(schema_is_null)
1863 })
1864 })
1865 }
1866
1867 fn extract_type_from_composition(composition: &Value) -> Option<String> {
1870 composition.as_array().and_then(|arr| {
1871 arr.iter().find_map(|schema| {
1872 let obj = schema.as_object()?;
1873 if schema_is_null(obj) {
1874 return None;
1875 }
1876
1877 obj.get("type").and_then(extract_type).or_else(|| {
1878 if obj.contains_key("properties") {
1879 Some("object".to_string())
1880 } else if obj.contains_key("enum") {
1881 Some("string".to_string())
1883 } else {
1884 None
1885 }
1886 })
1887 })
1888 })
1889 }
1890
1891 fn extract_schema_from_composition(
1894 composition: &Value,
1895 ) -> Option<serde_json::Map<String, Value>> {
1896 composition.as_array().and_then(|arr| {
1897 arr.iter().find_map(|schema| {
1898 let obj = schema.as_object()?;
1899 if schema_is_null(obj) {
1900 None
1901 } else {
1902 Some(obj.clone())
1903 }
1904 })
1905 })
1906 }
1907
1908 fn extract_schema_from_composition_obj(
1909 obj: &serde_json::Map<String, Value>,
1910 ) -> Option<serde_json::Map<String, Value>> {
1911 obj.get("anyOf")
1912 .and_then(extract_schema_from_composition)
1913 .or_else(|| obj.get("oneOf").and_then(extract_schema_from_composition))
1914 .or_else(|| obj.get("allOf").and_then(extract_schema_from_composition))
1915 }
1916
1917 fn infer_type(obj: &serde_json::Map<String, Value>) -> String {
1920 if let Some(type_val) = obj.get("type")
1921 && let Some(type_str) = extract_type(type_val)
1922 {
1923 return type_str;
1924 }
1925
1926 if let Some(any_of) = obj.get("anyOf")
1927 && let Some(type_str) = extract_type_from_composition(any_of)
1928 {
1929 return type_str;
1930 }
1931
1932 if let Some(one_of) = obj.get("oneOf")
1933 && let Some(type_str) = extract_type_from_composition(one_of)
1934 {
1935 return type_str;
1936 }
1937
1938 if let Some(all_of) = obj.get("allOf")
1939 && let Some(type_str) = extract_type_from_composition(all_of)
1940 {
1941 return type_str;
1942 }
1943
1944 if obj.contains_key("properties") {
1945 "object".to_string()
1946 } else if obj.contains_key("enum") {
1947 "string".to_string()
1948 } else {
1949 String::new()
1950 }
1951 }
1952
1953 impl TryFrom<Value> for Schema {
1954 type Error = EncodeError;
1955
1956 fn try_from(value: Value) -> Result<Self, Self::Error> {
1957 let flattened_val = flatten_schema(value)?;
1958 if let Some(obj) = flattened_val.as_object() {
1959 let composition_source = extract_schema_from_composition_obj(obj);
1960 let props_source = if obj.get("properties").is_none() {
1961 composition_source.clone().unwrap_or_else(|| obj.clone())
1962 } else {
1963 obj.clone()
1964 };
1965
1966 let schema_type = infer_type(obj);
1967 let items = obj
1968 .get("items")
1969 .or_else(|| props_source.get("items"))
1970 .and_then(|v| v.clone().try_into().ok())
1971 .map(Box::new);
1972
1973 let items = if schema_type == "array" && items.is_none() {
1976 Some(Box::new(Schema {
1977 r#type: "string".to_string(),
1978 format: None,
1979 description: None,
1980 nullable: None,
1981 r#enum: None,
1982 max_items: None,
1983 min_items: None,
1984 properties: None,
1985 required: None,
1986 items: None,
1987 }))
1988 } else {
1989 items
1990 };
1991
1992 Ok(Schema {
1993 r#type: schema_type,
1994 format: obj
1995 .get("format")
1996 .or_else(|| props_source.get("format"))
1997 .and_then(|v| v.as_str())
1998 .map(String::from),
1999 description: obj
2000 .get("description")
2001 .or_else(|| props_source.get("description"))
2002 .and_then(|v| v.as_str())
2003 .map(String::from),
2004 nullable: if schema_is_nullable(obj)
2005 || composition_source.as_ref().is_some_and(schema_is_nullable)
2006 {
2007 Some(true)
2008 } else {
2009 None
2010 },
2011 r#enum: obj
2012 .get("enum")
2013 .or_else(|| props_source.get("enum"))
2014 .and_then(|v| v.as_array())
2015 .map(|arr| {
2016 arr.iter()
2017 .filter_map(|v| v.as_str().map(String::from))
2018 .collect()
2019 }),
2020 max_items: obj
2021 .get("maxItems")
2022 .and_then(serde_json::Value::as_i64)
2023 .map(|v| v as i32),
2024 min_items: obj
2025 .get("minItems")
2026 .and_then(serde_json::Value::as_i64)
2027 .map(|v| v as i32),
2028 properties: props_source
2029 .get("properties")
2030 .and_then(|v| v.as_object())
2031 .map(|map| {
2032 map.iter()
2033 .filter_map(|(k, v)| {
2034 v.clone().try_into().ok().map(|schema| (k.clone(), schema))
2035 })
2036 .collect()
2037 }),
2038 required: props_source
2039 .get("required")
2040 .and_then(|v| v.as_array())
2041 .map(|arr| {
2042 arr.iter()
2043 .filter_map(|v| v.as_str().map(String::from))
2044 .collect()
2045 }),
2046 items,
2047 })
2048 } else {
2049 Err(EncodeError::request("Expected a JSON object for Schema"))
2050 }
2051 }
2052 }
2053
2054 #[derive(Debug, Serialize)]
2055 #[serde(rename_all = "camelCase")]
2056 pub struct GenerateContentRequest {
2057 pub contents: Vec<Content>,
2058 #[serde(skip_serializing_if = "Option::is_none")]
2059 pub tools: Option<Vec<Value>>,
2060 pub tool_config: Option<ToolConfig>,
2061 pub generation_config: Option<GenerationConfig>,
2063 pub safety_settings: Option<Vec<SafetySetting>>,
2068 pub system_instruction: Option<Content>,
2071 #[serde(skip_serializing_if = "Option::is_none")]
2075 pub cached_content: Option<String>,
2076 #[serde(flatten, skip_serializing_if = "Option::is_none")]
2078 pub additional_params: Option<serde_json::Value>,
2079 }
2080
2081 #[derive(Debug, Serialize)]
2082 #[serde(rename_all = "camelCase")]
2083 pub struct Tool {
2084 pub function_declarations: Vec<FunctionDeclaration>,
2085 pub code_execution: Option<CodeExecution>,
2086 }
2087
2088 #[derive(Debug, Serialize, Clone)]
2089 #[serde(rename_all = "camelCase")]
2090 pub struct FunctionDeclaration {
2091 pub name: String,
2092 pub description: String,
2093 #[serde(skip_serializing_if = "Option::is_none")]
2094 pub parameters: Option<Schema>,
2095 }
2096
2097 #[derive(Debug, Serialize, Deserialize)]
2098 #[serde(rename_all = "camelCase")]
2099 pub struct ToolConfig {
2100 pub function_calling_config: Option<FunctionCallingMode>,
2101 }
2102
2103 #[derive(Debug, Serialize, Deserialize, Default)]
2104 #[serde(tag = "mode", rename_all = "UPPERCASE")]
2105 pub enum FunctionCallingMode {
2106 #[default]
2107 Auto,
2108 None,
2109 Any {
2110 #[serde(skip_serializing_if = "Option::is_none")]
2111 allowed_function_names: Option<Vec<String>>,
2112 },
2113 }
2114
2115 impl TryFrom<message::ToolChoice> for FunctionCallingMode {
2116 type Error = EncodeError;
2117 fn try_from(value: message::ToolChoice) -> Result<Self, Self::Error> {
2118 let res = match value {
2119 message::ToolChoice::Auto => Self::Auto,
2120 message::ToolChoice::None => Self::None,
2121 message::ToolChoice::Required => Self::Any {
2122 allowed_function_names: None,
2123 },
2124 message::ToolChoice::Specific { function_names } => Self::Any {
2125 allowed_function_names: Some(function_names),
2126 },
2127 };
2128
2129 Ok(res)
2130 }
2131 }
2132
2133 #[derive(Debug, Serialize)]
2134 pub struct CodeExecution {}
2135
2136 #[derive(Debug, Serialize)]
2137 #[serde(rename_all = "camelCase")]
2138 pub struct SafetySetting {
2139 pub category: HarmCategory,
2140 pub threshold: HarmBlockThreshold,
2141 }
2142
2143 #[derive(Debug, Serialize)]
2144 #[serde(rename_all = "SCREAMING_SNAKE_CASE")]
2145 pub enum HarmBlockThreshold {
2146 HarmBlockThresholdUnspecified,
2147 BlockLowAndAbove,
2148 BlockMediumAndAbove,
2149 BlockOnlyHigh,
2150 BlockNone,
2151 Off,
2152 }
2153}
2154
2155impl gemini_api_types::GenerateContentRequest {
2156 pub fn with_cached_content(&mut self, name: &str) -> Result<(), EncodeError> {
2160 if !name.starts_with("cachedContents/") {
2161 return Err(EncodeError::request(format!(
2162 "gemini cached content handle should look like `cachedContents/<id>`, got \
2163 `{name}`"
2164 )));
2165 }
2166
2167 if let Some(existing) = self.cached_content.as_deref()
2169 && existing != name
2170 {
2171 return Err(EncodeError::request(format!(
2172 "a Gemini request set cached content twice, to `{existing}` and `{name}` — \
2173 set it one way or the other"
2174 )));
2175 }
2176
2177 let blob = self.additional_params.as_ref();
2180 let smuggled = |spellings: &'static [&'static str]| {
2181 blob.and_then(|payload| smuggled_field(payload, spellings))
2182 };
2183
2184 let mut conflicts = Vec::new();
2185 if self.system_instruction.is_some() || smuggled(&SYSTEM_INSTRUCTION).is_some() {
2186 conflicts.push("a system instruction (preamble)");
2187 }
2188 let smuggled_tools = smuggled(&TOOLS);
2189 if self.tools.is_some() || smuggled_tools.is_some() {
2190 conflicts.push("tools");
2191 }
2192 if self.tool_config.is_some() || smuggled(&TOOL_CONFIG).is_some() {
2193 conflicts.push("a tool choice");
2194 }
2195 if !conflicts.is_empty() {
2196 let declares_function = |tool: &Value| {
2199 ["functionDeclarations", "function_declarations"]
2200 .iter()
2201 .any(|spelling| {
2202 tool.get(spelling)
2203 .and_then(Value::as_array)
2204 .is_some_and(|declarations| !declarations.is_empty())
2205 })
2206 };
2207 let declares_functions = self
2208 .tools
2209 .iter()
2210 .flatten()
2211 .chain(
2212 smuggled_tools
2213 .and_then(|spelling| blob?.get(spelling))
2214 .and_then(Value::as_array)
2215 .into_iter()
2216 .flatten(),
2217 )
2218 .any(declares_function);
2219 let tool_caveat = if declares_functions {
2220 " Note that function declarations in a cache are declarations only — rig's \
2221 `Agent` can only dispatch tools it advertised, so a cached function tool set is \
2222 never executable from an agent; it is usable only when you drive \
2223 `GenerateContent` yourself and run the tool loop. (Provider-hosted tools such \
2224 as `codeExecution` run on Gemini's side and are fine to keep in the cache.)"
2225 } else {
2226 ""
2227 };
2228 return Err(EncodeError::request(format!(
2229 "a Gemini request using cached content `{name}` also set {}. The cached \
2230 content already owns the system instruction, tools and tool choice for every \
2231 request that uses it — move them into the cache, or drop the cache \
2232 handle.{tool_caveat}",
2233 conflicts.join(" and ")
2234 )));
2235 }
2236
2237 self.cached_content = Some(name.to_owned());
2238 Ok(())
2239 }
2240}
2241
2242#[cfg(test)]
2243mod tests;
2244
2245#[cfg(test)]
2246mod cached_content_conflict_matrix;
2247#[cfg(test)]
2248mod cached_content_request_tests;