1pub const GEMINI_2_5_FLASH: &str = "gemini-2.5-flash";
16pub const GEMINI_2_0_FLASH_LITE: &str = "gemini-2.0-flash-lite";
18pub const GEMINI_2_0_FLASH: &str = "gemini-2.0-flash";
20
21use base64::Engine as _;
22use futures::StreamExt;
23use rig_core::completion::{self, CompletionRequest};
24use rig_core::driver::{Exchange, Opened, Opening, Transport};
25use rig_core::error::EncodeError;
26use rig_core::error::ProviderError;
27use rig_core::message::{self, MimeType};
28use rig_core::operation::Completion;
29use rig_core::providers::gemini::completion::gemini_api_types::{
30 Schema as GeminiSchema, map_google_finish_reason, tool_parameters_to_schema,
31};
32use rig_core::providers::gemini::text_thought_signature;
33use rig_core::wire::{Descriptor, Mode, Wire};
34use std::convert::TryFrom;
35
36use super::GeminiGrpc;
37use super::proto::{self, GenerateContentRequest, GenerateContentResponse};
38use super::streaming::GrpcAdapter;
39
40#[derive(Clone, Debug, PartialEq)]
43pub struct GenerateContent {
44 pub model: String,
45}
46
47impl GenerateContent {
48 pub fn new(model: impl Into<String>) -> Self {
49 Self {
50 model: model.into(),
51 }
52 }
53}
54
55impl Wire for GenerateContent {
56 type Op = Completion;
57 type Payload = GenerateContentRequest;
58 type Frame = GenerateContentResponse;
59 type Decoder<'id> = GrpcAdapter<'id>;
60
61 fn describe(&self) -> Descriptor<'_> {
62 Descriptor::new(PROVIDER_NAME).model(self.model.as_str())
63 }
64
65 fn encode(
68 &self,
69 request: CompletionRequest,
70 _mode: Mode,
71 ) -> Result<GenerateContentRequest, EncodeError> {
72 create_grpc_request(&self.model, request.replayable_to(&[ISSUER])?)
73 }
74
75 fn decoder<'id>(&self) -> Self::Decoder<'id> {
76 GrpcAdapter::default()
77 }
78}
79
80impl Transport<GenerateContent> for GeminiGrpc {
81 fn send(
82 &self,
83 request: GenerateContentRequest,
84 exchange: Exchange,
85 ) -> Opening<GenerateContentResponse> {
86 let mode = exchange.mode;
87 let mut client = match self.grpc_client() {
88 Ok(client) => client,
89 Err(error) => return Opening::failed(ProviderError::Provider(error.to_string())),
90 };
91 Opening::new(async move {
92 Ok(match mode {
93 Mode::Unary => match client.generate_content(request).await {
94 Ok(response) => Opened::new(futures::stream::iter([Ok(response.into_inner())])),
95 Err(status) => Opened::failed(rpc_error(&status)),
96 },
97 Mode::Streaming => match client.stream_generate_content(request).await {
98 Ok(response) => {
99 let mut chunks = response.into_inner();
100 Opened::new(async_stream::stream! {
102 while let Some(item) = chunks.next().await {
103 match item {
104 Ok(chunk) => yield Ok(chunk),
105 Err(status) => {
106 yield Err(rpc_error(&status));
107 break;
108 }
109 }
110 }
111 })
112 }
113 Err(status) => Opened::failed(rpc_error(&status)),
114 },
115 })
116 })
117 }
118}
119
120pub const PROVIDER_NAME: &str = "gemini-grpc";
122
123pub const REASONING_ISSUER: &str = rig_core::providers::gemini::completion::PROVIDER_NAME;
127
128const ISSUER: message::Issuer = message::Issuer::from_static(REASONING_ISSUER);
130
131pub fn map_finish_reason(reason: i32) -> Option<completion::FinishReason> {
138 use proto::candidate::FinishReason as Wire;
139
140 let Ok(reason) = Wire::try_from(reason) else {
141 return Some(completion::FinishReason::Other(format!(
142 "FINISH_REASON_{reason}"
143 )));
144 };
145
146 map_google_finish_reason(reason.as_str_name())
147}
148
149pub fn tool_protocol_finish_reason_error(
153 reason: i32,
154 finish_message: Option<&str>,
155) -> Option<ProviderError> {
156 use proto::candidate::FinishReason as Wire;
157
158 let reason = Wire::try_from(reason).ok()?;
159 match reason {
160 Wire::MalformedFunctionCall | Wire::UnexpectedToolCall | Wire::TooManyToolCalls => {
161 let message = finish_message.unwrap_or("no finish message provided");
162 Some(ProviderError::Response(format!(
163 "Gemini stopped with finish_reason={}: {message}",
164 reason.as_str_name()
165 )))
166 }
167 _ => None,
168 }
169}
170
171pub(crate) fn data_part(data: proto::part::Data) -> proto::Part {
173 proto::Part {
174 data: Some(data),
175 thought: false,
176 thought_signature: Vec::new(),
177 part_metadata: None,
178 }
179}
180
181pub(crate) fn text_part(text: String) -> proto::Part {
183 data_part(proto::part::Data::Text(text))
184}
185
186pub(crate) fn rpc_error(status: &tonic::Status) -> ProviderError {
189 ProviderError::from_provider_body(status.to_string())
190 .with_provider_code(Some(grpc_code_name(status.code())))
191 .with_transient(Some(transient_grpc_code(status.code())))
192}
193
194pub(crate) fn grpc_code_name(code: tonic::Code) -> String {
198 format!("{code:?}")
199 .chars()
200 .fold(String::new(), |mut name, c| {
201 if c.is_ascii_uppercase() && !name.is_empty() {
202 name.push('_');
203 }
204 name.push(c.to_ascii_uppercase());
205 name
206 })
207}
208
209pub(crate) fn transient_grpc_code(code: tonic::Code) -> bool {
211 matches!(
212 code,
213 tonic::Code::Unavailable
214 | tonic::Code::ResourceExhausted
215 | tonic::Code::DeadlineExceeded
216 | tonic::Code::Aborted
217 )
218}
219
220pub(crate) fn create_grpc_request(
221 model: &str,
222 completion_request: CompletionRequest,
223) -> Result<GenerateContentRequest, EncodeError> {
224 let CompletionRequest {
225 model: _,
226 chat_history,
227 documents: _,
228 tools,
229 temperature,
230 max_tokens,
231 tool_choice: _,
232 additional_params: _,
233 output_schema: _,
234 record_telemetry_content: _,
235 } = completion_request;
236
237 let (history_system, chat_history) = split_system_messages_from_history(chat_history);
238 let mut contents = Vec::new();
239
240 for msg in chat_history {
241 contents.push(rig_message_to_grpc_content(msg)?);
242 }
243
244 let mut system_parts = Vec::new();
245 for content in history_system {
246 if !content.is_empty() {
247 system_parts.push(text_part(content));
248 }
249 }
250 let system_instruction = if system_parts.is_empty() {
251 None
252 } else {
253 Some(proto::Content {
254 parts: system_parts,
255 role: "model".to_string(),
256 })
257 };
258
259 let generation_config = if temperature.is_some() || max_tokens.is_some() {
260 Some(proto::GenerationConfig {
261 temperature: temperature.map(|t| t as f32),
262 max_output_tokens: max_tokens.map(|t| t as i32),
263 ..Default::default()
264 })
265 } else {
266 None
267 };
268
269 let tools = if !tools.is_empty() {
270 let function_declarations = tools
271 .into_iter()
272 .map(|tool| {
273 Ok(proto::FunctionDeclaration {
274 name: tool.name,
275 description: tool.description,
276 parameters: tool_parameters_to_proto_schema(&tool.parameters)?,
277 ..Default::default()
278 })
279 })
280 .collect::<Result<Vec<_>, EncodeError>>()?;
281
282 vec![proto::Tool {
283 function_declarations,
284 code_execution: None,
285 }]
286 } else {
287 vec![]
288 };
289
290 Ok(GenerateContentRequest {
291 model: format!("models/{model}"),
292 contents,
293 tools,
294 safety_settings: vec![],
295 generation_config,
296 tool_config: None,
297 system_instruction,
298 cached_content: String::new(),
299 })
300}
301
302fn rig_message_to_grpc_content(msg: message::Message) -> Result<proto::Content, EncodeError> {
303 match msg {
304 message::Message::System { .. } => Err(EncodeError::request(
305 "System messages must be sent via Gemini gRPC system_instruction",
306 )),
307 message::Message::User { content } => {
308 let parts = content
309 .into_iter()
310 .map(rig_user_content_to_grpc_part)
311 .collect::<Result<Vec<_>, _>>()?;
312
313 Ok(proto::Content {
314 parts,
315 role: "user".to_string(),
316 })
317 }
318 message::Message::Assistant { content, .. } => {
319 let parts = content
320 .into_iter()
321 .filter(|part| match part {
323 message::AssistantContent::Reasoning(reasoning) => {
324 reasoning.open(&ISSUER).is_some()
325 }
326 _ => true,
327 })
328 .map(rig_assistant_content_to_grpc_part)
329 .collect::<Result<Vec<_>, _>>()?;
330
331 Ok(proto::Content {
332 parts,
333 role: "model".to_string(),
334 })
335 }
336 }
337}
338
339use rig_core::providers::gemini::completion::split_system_messages_from_history;
340
341fn rig_user_content_to_grpc_part(
342 content: message::UserContent,
343) -> Result<proto::Part, EncodeError> {
344 match content {
345 message::UserContent::Text(message::Text { text, .. }) => Ok(text_part(text)),
346 message::UserContent::ToolResult(result) => {
347 let mut values = result
348 .content
349 .into_iter()
350 .map(|content| match content {
351 message::ToolResultContent::Text(t) => Ok(serde_json::Value::String(t.text)),
352 message::ToolResultContent::Json { value } => Ok(value),
353 message::ToolResultContent::Image(_) => Err(EncodeError::request(
354 "Gemini gRPC does not support images in tool results",
355 )),
356 })
357 .collect::<Result<Vec<_>, _>>()?;
358 let result_value = if values.len() == 1 {
359 values.remove(0)
360 } else {
361 serde_json::Value::Array(values)
362 };
363
364 let response_struct =
365 json_to_prost_struct(serde_json::json!({ "result": result_value }))?;
366
367 Ok(data_part(proto::part::Data::FunctionResponse(
370 proto::FunctionResponse {
371 name: result.name.into(),
372 response: Some(response_struct),
373 id: result
374 .call
375 .provider()
376 .map(|provider| provider.call_id.clone())
377 .unwrap_or_default(),
378 },
379 )))
380 }
381 message::UserContent::Image(img) => {
382 let Some(media_type) = img.media_type else {
383 return Err(EncodeError::request(
384 "Media type for image is required for Gemini",
385 ));
386 };
387
388 match media_type {
389 message::ImageMediaType::JPEG
390 | message::ImageMediaType::PNG
391 | message::ImageMediaType::WEBP
392 | message::ImageMediaType::HEIC
393 | message::ImageMediaType::HEIF => {}
394 _ => {
395 return Err(EncodeError::request(format!(
396 "Unsupported image media type {media_type:?}"
397 )));
398 }
399 }
400
401 let mime_type = media_type.to_mime_type().to_string();
402
403 let data = match img.data {
404 message::DocumentSourceKind::Url(file_uri) => {
405 return Ok(data_part(proto::part::Data::FileData(proto::FileData {
406 mime_type,
407 file_uri,
408 })));
409 }
410 message::DocumentSourceKind::Raw(bytes) => bytes,
411 message::DocumentSourceKind::Base64(data)
412 | message::DocumentSourceKind::String(data) => decode_base64_bytes(&data)?,
413 message::DocumentSourceKind::Unknown => {
414 return Err(EncodeError::request("Image content has no body"));
415 }
416 _ => {
417 return Err(EncodeError::request("Unsupported document source kind"));
418 }
419 };
420
421 Ok(data_part(proto::part::Data::InlineData(proto::Blob {
422 mime_type,
423 data,
424 })))
425 }
426 _ => Err(EncodeError::request("Unsupported user content type")),
427 }
428}
429
430fn rig_assistant_content_to_grpc_part(
431 content: message::AssistantContent,
432) -> Result<proto::Part, EncodeError> {
433 match content {
434 message::AssistantContent::Text(text) => Ok(proto::Part {
435 thought_signature: decode_optional_base64(
436 text_thought_signature(&text).map(str::to_owned),
437 )?,
438 ..text_part(text.text)
439 }),
440 message::AssistantContent::ToolCall(tool_call) => {
441 let args = json_to_prost_struct(tool_call.function.arguments)?;
442
443 Ok(proto::Part {
444 thought_signature: decode_optional_base64(tool_call.signature)?,
445 ..data_part(proto::part::Data::FunctionCall(proto::FunctionCall {
446 name: tool_call.function.name.into(),
447 args: Some(args),
448 id: tool_call
451 .id
452 .provider()
453 .map(|provider| provider.call_id.clone())
454 .unwrap_or_default(),
455 }))
456 })
457 }
458 message::AssistantContent::Reasoning(reasoning) => {
459 let reasoning = reasoning.open(&ISSUER).ok_or_else(|| {
460 EncodeError::request("Gemini cannot replay reasoning another service issued")
461 })?;
462 Ok(proto::Part {
463 data: Some(proto::part::Data::Text(reasoning.display_text())),
464 thought: true,
465 thought_signature: decode_optional_base64(
466 reasoning
467 .first_signature()
468 .map(std::string::ToString::to_string),
469 )?,
470 part_metadata: None,
471 })
472 }
473 _ => Err(EncodeError::request("Unsupported assistant content type")),
474 }
475}
476
477fn decode_base64_bytes(input: &str) -> Result<Vec<u8>, EncodeError> {
478 let data = input.trim();
479
480 let data = if let Some(rest) = data.strip_prefix("data:") {
482 rest.split_once(',').map_or(data, |(_, b64)| b64)
483 } else {
484 data
485 };
486
487 let mut last_err: Option<String> = None;
488
489 for engine in [
490 &base64::engine::general_purpose::STANDARD,
491 &base64::engine::general_purpose::URL_SAFE,
492 &base64::engine::general_purpose::STANDARD_NO_PAD,
493 &base64::engine::general_purpose::URL_SAFE_NO_PAD,
494 ] {
495 match engine.decode(data) {
496 Ok(bytes) => return Ok(bytes),
497 Err(err) => last_err = Some(err.to_string()),
498 }
499 }
500
501 let err = last_err.unwrap_or_else(|| "unknown base64 decode error".to_string());
502 Err(EncodeError::request(format!("Invalid base64 data: {err}")))
503}
504
505fn decode_optional_base64(sig: Option<String>) -> Result<Vec<u8>, EncodeError> {
506 let Some(sig) = sig else {
507 return Ok(Vec::new());
508 };
509 decode_base64_bytes(&sig)
510}
511
512pub(crate) fn map_usage(usage: Option<&proto::UsageMetadata>) -> completion::Usage {
520 usage
521 .map(|usage| {
522 let count = |count: i32| count as u64;
523 let input = count(usage.prompt_token_count) + count(usage.tool_use_prompt_token_count);
524 let output = count(usage.candidates_token_count) + count(usage.thoughts_token_count);
525 completion::Usage {
526 input_tokens: Some(input),
527 output_tokens: Some(output),
528 total_tokens: Some(input + output),
529 cached_input_tokens: Some(count(usage.cached_content_token_count)),
530 cache_creation_input_tokens: None,
531 tool_use_prompt_tokens: Some(count(usage.tool_use_prompt_token_count)),
532 reasoning_tokens: Some(count(usage.thoughts_token_count)),
533 }
534 })
535 .unwrap_or_default()
536}
537
538pub(crate) fn encode_optional_base64(bytes: &[u8]) -> Option<String> {
539 if bytes.is_empty() {
540 None
541 } else {
542 Some(base64::engine::general_purpose::STANDARD.encode(bytes))
543 }
544}
545
546fn json_to_prost_struct(value: serde_json::Value) -> Result<proto::Struct, EncodeError> {
547 match value {
548 serde_json::Value::Object(map) => Ok(proto::Struct {
549 fields: map
550 .into_iter()
551 .map(|(k, v)| (k, json_to_prost_value(v)))
552 .collect(),
553 }),
554 _ => Err(EncodeError::request(
555 "Expected a JSON object for google.protobuf.Struct",
556 )),
557 }
558}
559
560fn json_to_prost_value(value: serde_json::Value) -> proto::Value {
561 match value {
562 serde_json::Value::Null => proto::Value {
563 kind: Some(proto::value::Kind::NullValue(
564 proto::NullValue::NullValue as i32,
565 )),
566 },
567 serde_json::Value::Bool(b) => proto::Value {
568 kind: Some(proto::value::Kind::BoolValue(b)),
569 },
570 serde_json::Value::Number(n) => proto::Value {
571 kind: Some(proto::value::Kind::NumberValue(
572 n.as_f64().unwrap_or_default(),
573 )),
574 },
575 serde_json::Value::String(s) => proto::Value {
576 kind: Some(proto::value::Kind::StringValue(s)),
577 },
578 serde_json::Value::Array(items) => proto::Value {
579 kind: Some(proto::value::Kind::ListValue(proto::ListValue {
580 values: items.into_iter().map(json_to_prost_value).collect(),
581 })),
582 },
583 serde_json::Value::Object(map) => proto::Value {
584 kind: Some(proto::value::Kind::StructValue(proto::Struct {
585 fields: map
586 .into_iter()
587 .map(|(k, v)| (k, json_to_prost_value(v)))
588 .collect(),
589 })),
590 },
591 }
592}
593
594pub(crate) fn prost_struct_to_json(st: &proto::Struct) -> serde_json::Value {
595 let mut out = serde_json::Map::with_capacity(st.fields.len());
596 for (k, v) in &st.fields {
597 out.insert(k.clone(), prost_value_to_json(v));
598 }
599 serde_json::Value::Object(out)
600}
601
602fn prost_value_to_json(v: &proto::Value) -> serde_json::Value {
603 match &v.kind {
604 None | Some(proto::value::Kind::NullValue(_)) => serde_json::Value::Null,
605 Some(proto::value::Kind::BoolValue(b)) => serde_json::Value::Bool(*b),
606 Some(proto::value::Kind::NumberValue(n)) => serde_json::Number::from_f64(*n)
607 .map_or(serde_json::Value::Null, serde_json::Value::Number),
608 Some(proto::value::Kind::StringValue(s)) => serde_json::Value::String(s.clone()),
609 Some(proto::value::Kind::StructValue(st)) => prost_struct_to_json(st),
610 Some(proto::value::Kind::ListValue(list)) => {
611 serde_json::Value::Array(list.values.iter().map(prost_value_to_json).collect())
612 }
613 }
614}
615
616fn tool_parameters_to_proto_schema(
619 value: &serde_json::Value,
620) -> Result<Option<proto::Schema>, EncodeError> {
621 tool_parameters_to_schema(value.clone()).map(|schema| schema.map(gemini_schema_to_proto_schema))
622}
623
624fn gemini_schema_to_proto_schema(schema: GeminiSchema) -> proto::Schema {
625 proto::Schema {
626 r#type: json_type_to_proto_type(&schema.r#type) as i32,
627 format: schema.format.unwrap_or_default(),
628 description: schema.description.unwrap_or_default(),
629 nullable: schema.nullable.unwrap_or(false),
630 r#enum: schema.r#enum.unwrap_or_default(),
631 items: schema
632 .items
633 .map(|items| Box::new(gemini_schema_to_proto_schema(*items))),
634 properties: schema
635 .properties
636 .unwrap_or_default()
637 .into_iter()
638 .map(|(name, schema)| (name, gemini_schema_to_proto_schema(schema)))
639 .collect(),
640 required: schema.required.unwrap_or_default(),
641 }
642}
643
644fn json_type_to_proto_type(t: &str) -> proto::Type {
645 match t {
646 "string" => proto::Type::String,
647 "number" => proto::Type::Number,
648 "integer" => proto::Type::Integer,
649 "boolean" => proto::Type::Boolean,
650 "array" => proto::Type::Array,
651 "object" => proto::Type::Object,
652 "null" => proto::Type::Null,
653 _ => proto::Type::Unspecified,
654 }
655}
656
657#[cfg(test)]
658#[allow(clippy::expect_used, clippy::unwrap_used)]
659pub(crate) mod tests;