1pub const GEMINI_2_5_FLASH: &str = "gemini-2.5-flash";
7pub const GEMINI_2_0_FLASH_LITE: &str = "gemini-2.0-flash-lite";
9pub const GEMINI_2_0_FLASH: &str = "gemini-2.0-flash";
11
12use base64::Engine as _;
13use rig_core::completion::{self, CompletionError, CompletionRequest};
14use rig_core::message::{self, MimeType, Reasoning};
15use rig_core::providers::gemini::completion::attach_trailing_signature;
16use rig_core::providers::gemini::completion::gemini_api_types::{
17 Schema as GeminiSchema, map_google_finish_reason, tool_parameters_to_schema,
18};
19use rig_core::telemetry::ProviderResponseExt;
20use std::convert::TryFrom;
21
22use super::Client;
23use super::proto::{self, GenerateContentRequest, GenerateContentResponse};
24
25#[derive(Clone, Debug)]
30pub struct CompletionModel {
31 pub(crate) client: Client,
32 pub model: String,
33}
34
35impl CompletionModel {
36 pub fn new(client: Client, model: impl Into<String>) -> Self {
37 Self {
38 client,
39 model: model.into(),
40 }
41 }
42}
43
44pub const PROVIDER_NAME: &str = "gemini-grpc";
46
47pub fn map_finish_reason(reason: i32) -> Option<completion::FinishReason> {
54 use proto::candidate::FinishReason as Wire;
55
56 let Ok(reason) = Wire::try_from(reason) else {
57 return Some(completion::FinishReason::Other(format!(
58 "FINISH_REASON_{reason}"
59 )));
60 };
61
62 map_google_finish_reason(reason.as_str_name())
63}
64
65pub fn tool_protocol_finish_reason_error(
80 reason: i32,
81 finish_message: Option<&str>,
82) -> Option<CompletionError> {
83 use proto::candidate::FinishReason as Wire;
84
85 let reason = Wire::try_from(reason).ok()?;
86 match reason {
87 Wire::MalformedFunctionCall | Wire::UnexpectedToolCall | Wire::TooManyToolCalls => {
88 let message = finish_message.unwrap_or("no finish message provided");
89 Some(CompletionError::ResponseError(format!(
90 "Gemini stopped with finish_reason={}: {message}",
91 reason.as_str_name()
92 )))
93 }
94 _ => None,
95 }
96}
97
98impl CompletionModel {
99 pub async fn raw_completion(
105 &self,
106 completion_request: CompletionRequest,
107 ) -> Result<GenerateContentResponse, CompletionError> {
108 let request = create_grpc_request(self.model.clone(), completion_request)?;
109
110 let mut grpc_client = self
111 .client
112 .grpc_client()
113 .map_err(|e| CompletionError::ProviderError(e.to_string()))?;
114
115 let response = grpc_client
116 .generate_content(request)
117 .await
118 .map_err(rpc_error)?
119 .into_inner();
120
121 Ok(response)
122 }
123
124 pub async fn raw_stream(
127 &self,
128 request: CompletionRequest,
129 ) -> Result<
130 rig_core::streaming::RawStreamingResult<super::streaming::StreamingCompletionResponse>,
131 CompletionError,
132 > {
133 super::streaming::raw_stream(self.client.clone(), self.model.clone(), request).await
134 }
135}
136
137impl completion::CompletionModel for CompletionModel {
138 async fn completion(
139 &self,
140 completion_request: CompletionRequest,
141 ) -> Result<completion::CompletionResponse, CompletionError> {
142 let raw = self.raw_completion(completion_request).await?;
144 let captured = serde_json::to_value(&raw)?;
145 let response: completion::CompletionResponse = raw.try_into()?;
146 Ok(response.with_raw(captured))
147 }
148
149 async fn stream(
150 &self,
151 request: CompletionRequest,
152 ) -> Result<rig_core::streaming::StreamingCompletionResponse, CompletionError> {
153 super::streaming::stream(self.client.clone(), self.model.clone(), request).await
154 }
155}
156
157pub(crate) fn data_part(data: proto::part::Data) -> proto::Part {
159 proto::Part {
160 data: Some(data),
161 thought: false,
162 thought_signature: Vec::new(),
163 part_metadata: None,
164 }
165}
166
167pub(crate) fn text_part(text: String) -> proto::Part {
169 data_part(proto::part::Data::Text(text))
170}
171
172pub(crate) fn rpc_error(status: tonic::Status) -> CompletionError {
180 CompletionError::from_provider_body(status.to_string())
181}
182
183pub(crate) fn create_grpc_request(
185 model: String,
186 completion_request: CompletionRequest,
187) -> Result<GenerateContentRequest, CompletionError> {
188 let CompletionRequest {
189 model: _,
190 preamble,
191 chat_history,
192 documents: _,
193 tools,
194 temperature,
195 max_tokens,
196 tool_choice: _,
197 additional_params: _,
198 output_schema: _,
199 record_telemetry_content: _,
200 } = completion_request;
201
202 let (history_system, mut chat_history) = split_system_messages_from_history(chat_history);
203 rig_core::providers::internal::resolve_empty_tool_result_names(&mut chat_history);
206 let mut contents = Vec::new();
207
208 for msg in chat_history {
210 contents.push(rig_message_to_grpc_content(msg)?);
211 }
212
213 let mut system_parts = Vec::new();
215 if let Some(preamble) = preamble
216 && !preamble.is_empty()
217 {
218 system_parts.push(text_part(preamble));
219 }
220 for content in history_system {
221 if !content.is_empty() {
222 system_parts.push(text_part(content));
223 }
224 }
225 let system_instruction = if system_parts.is_empty() {
226 None
227 } else {
228 Some(proto::Content {
229 parts: system_parts,
230 role: "model".to_string(),
231 })
232 };
233
234 let generation_config = if temperature.is_some() || max_tokens.is_some() {
236 Some(proto::GenerationConfig {
237 temperature: temperature.map(|t| t as f32),
238 max_output_tokens: max_tokens.map(|t| t as i32),
239 ..Default::default()
240 })
241 } else {
242 None
243 };
244
245 let tools = if !tools.is_empty() {
247 let function_declarations = tools
248 .into_iter()
249 .map(|tool| {
250 Ok(proto::FunctionDeclaration {
251 name: tool.name,
252 description: tool.description,
253 parameters: tool_parameters_to_proto_schema(&tool.parameters)?,
254 ..Default::default()
255 })
256 })
257 .collect::<Result<Vec<_>, CompletionError>>()?;
258
259 vec![proto::Tool {
260 function_declarations,
261 code_execution: None,
262 }]
263 } else {
264 vec![]
265 };
266
267 Ok(GenerateContentRequest {
268 model: format!("models/{}", model),
269 contents,
270 tools,
271 safety_settings: vec![],
272 generation_config,
273 tool_config: None,
274 system_instruction,
275 cached_content: String::new(),
276 })
277}
278
279fn rig_message_to_grpc_content(msg: message::Message) -> Result<proto::Content, CompletionError> {
281 match msg {
282 message::Message::System { .. } => Err(CompletionError::RequestError(
283 "System messages must be sent via Gemini gRPC system_instruction".into(),
284 )),
285 message::Message::User { content } => {
286 let parts = content
287 .into_iter()
288 .map(rig_user_content_to_grpc_part)
289 .collect::<Result<Vec<_>, _>>()?;
290
291 Ok(proto::Content {
292 parts,
293 role: "user".to_string(),
294 })
295 }
296 message::Message::Assistant { content, .. } => {
297 let parts = content
298 .into_iter()
299 .map(rig_assistant_content_to_grpc_part)
300 .collect::<Result<Vec<_>, _>>()?;
301
302 Ok(proto::Content {
303 parts,
304 role: "model".to_string(),
305 })
306 }
307 }
308}
309
310use rig_core::providers::gemini::completion::split_system_messages_from_history;
311
312fn rig_user_content_to_grpc_part(
314 content: message::UserContent,
315) -> Result<proto::Part, CompletionError> {
316 match content {
317 message::UserContent::Text(message::Text { text, .. }) => Ok(text_part(text)),
318 message::UserContent::ToolResult(result) => {
319 let mut values = result
320 .content
321 .into_iter()
322 .map(|content| match content {
323 message::ToolResultContent::Text(t) => Ok(serde_json::Value::String(t.text)),
324 message::ToolResultContent::Json { value } => Ok(value),
325 message::ToolResultContent::Image(_) => Err(CompletionError::RequestError(
326 "Gemini gRPC does not support images in tool results".into(),
327 )),
328 })
329 .collect::<Result<Vec<_>, _>>()?;
330 let result_value = if values.len() == 1 {
331 values.remove(0)
332 } else {
333 serde_json::Value::Array(values)
334 };
335
336 let response_struct =
337 json_to_prost_struct(serde_json::json!({ "result": result_value }))?;
338
339 Ok(data_part(proto::part::Data::FunctionResponse(
343 proto::FunctionResponse {
344 name: result.name,
345 response: Some(response_struct),
346 id: result
347 .provider
348 .map(|provider| provider.call_id)
349 .unwrap_or_default(),
350 },
351 )))
352 }
353 message::UserContent::Image(img) => {
354 let Some(media_type) = img.media_type else {
355 return Err(CompletionError::RequestError(
356 "Media type for image is required for Gemini".into(),
357 ));
358 };
359
360 match media_type {
361 message::ImageMediaType::JPEG
362 | message::ImageMediaType::PNG
363 | message::ImageMediaType::WEBP
364 | message::ImageMediaType::HEIC
365 | message::ImageMediaType::HEIF => {}
366 _ => {
367 return Err(CompletionError::RequestError(
368 format!("Unsupported image media type {media_type:?}").into(),
369 ));
370 }
371 }
372
373 let mime_type = media_type.to_mime_type().to_string();
374
375 let data = match img.data {
376 message::DocumentSourceKind::Url(file_uri) => {
377 return Ok(data_part(proto::part::Data::FileData(proto::FileData {
378 mime_type,
379 file_uri,
380 })));
381 }
382 message::DocumentSourceKind::Raw(bytes) => bytes,
383 message::DocumentSourceKind::Base64(data)
384 | message::DocumentSourceKind::String(data) => decode_base64_bytes(&data)?,
385 message::DocumentSourceKind::Unknown => {
386 return Err(CompletionError::RequestError(
387 "Image content has no body".into(),
388 ));
389 }
390 _ => {
391 return Err(CompletionError::RequestError(
392 "Unsupported document source kind".into(),
393 ));
394 }
395 };
396
397 Ok(data_part(proto::part::Data::InlineData(proto::Blob {
398 mime_type,
399 data,
400 })))
401 }
402 _ => Err(CompletionError::RequestError(
403 "Unsupported user content type".into(),
404 )),
405 }
406}
407
408fn rig_assistant_content_to_grpc_part(
410 content: message::AssistantContent,
411) -> Result<proto::Part, CompletionError> {
412 match content {
413 message::AssistantContent::Text(message::Text { text, .. }) => Ok(text_part(text)),
414 message::AssistantContent::ToolCall(tool_call) => {
415 let args = json_to_prost_struct(tool_call.function.arguments)?;
416
417 Ok(proto::Part {
418 thought_signature: decode_optional_base64(tool_call.signature)?,
419 ..data_part(proto::part::Data::FunctionCall(proto::FunctionCall {
420 name: tool_call.function.name,
421 args: Some(args),
422 id: tool_call
425 .provider
426 .map(|provider| provider.call_id)
427 .unwrap_or_default(),
428 }))
429 })
430 }
431 message::AssistantContent::Reasoning(reasoning) => Ok(proto::Part {
432 data: Some(proto::part::Data::Text(reasoning.display_text())),
433 thought: true,
434 thought_signature: decode_optional_base64(
435 reasoning.first_signature().map(|s| s.to_string()),
436 )?,
437 part_metadata: None,
438 }),
439 _ => Err(CompletionError::RequestError(
440 "Unsupported assistant content type".into(),
441 )),
442 }
443}
444
445impl TryFrom<GenerateContentResponse> for completion::CompletionResponse {
447 type Error = CompletionError;
448
449 fn try_from(response: GenerateContentResponse) -> Result<Self, Self::Error> {
450 let candidate = response.candidates.first().ok_or_else(|| {
451 CompletionError::ResponseError("No response candidates in response".into())
452 })?;
453
454 if let Some(err) = tool_protocol_finish_reason_error(
457 candidate.finish_reason,
458 candidate.finish_message.as_deref(),
459 ) {
460 return Err(err);
461 }
462
463 let content_ref = candidate.content.as_ref().ok_or_else(|| {
464 CompletionError::ResponseError(format!(
465 "Gemini candidate missing content (finish_reason={})",
466 candidate.finish_reason
467 ))
468 })?;
469
470 let mut assistant_contents = Vec::new();
471
472 for part in &content_ref.parts {
473 let assistant_content = match &part.data {
474 Some(proto::part::Data::Text(text)) => {
475 if part.thought {
476 completion::AssistantContent::Reasoning(Reasoning::new_with_signature(
477 text,
478 encode_optional_base64(&part.thought_signature),
479 ))
480 } else {
481 completion::AssistantContent::text(text)
482 }
483 }
484 Some(proto::part::Data::InlineData(inline_data)) => {
485 let mime_type = message::MediaType::from_mime_type(&inline_data.mime_type);
486 match mime_type {
487 Some(message::MediaType::Image(media_type)) => {
488 let b64 =
489 base64::engine::general_purpose::STANDARD.encode(&inline_data.data);
490 completion::AssistantContent::image_base64(
491 b64,
492 Some(media_type),
493 Some(message::ImageDetail::default()),
494 )
495 }
496 _ => {
497 return Err(CompletionError::ResponseError(format!(
498 "Unsupported media type {mime_type:?}"
499 )));
500 }
501 }
502 }
503 Some(proto::part::Data::FunctionCall(function_call)) => {
504 let args = function_call
505 .args
506 .as_ref()
507 .map(prost_struct_to_json)
508 .unwrap_or(serde_json::Value::Object(serde_json::Map::new()));
509
510 let tool_call = message::ToolCall::from_wire(
513 function_call.id.clone(),
514 message::ToolFunction::new(function_call.name.clone(), args),
515 )
516 .with_signature(encode_optional_base64(&part.thought_signature));
517
518 completion::AssistantContent::ToolCall(tool_call)
519 }
520 _ => {
521 return Err(CompletionError::ResponseError(
522 "Response did not contain a message or tool call".into(),
523 ));
524 }
525 };
526
527 assistant_contents.push(assistant_content);
528
529 if !part.thought
535 && matches!(part.data, Some(proto::part::Data::Text(_)))
536 && let Some(signature) = encode_optional_base64(&part.thought_signature)
537 {
538 attach_trailing_signature(&mut assistant_contents, signature);
539 }
540 }
541
542 let choice = rig_core::message::require_non_empty_response(assistant_contents)?;
543
544 let usage = map_usage(response.usage_metadata.as_ref());
545
546 let finish_reason = response
547 .candidates
548 .first()
549 .and_then(|candidate| map_finish_reason(candidate.finish_reason));
550 let model = Some(response.model_version.clone()).filter(|model| !model.is_empty());
551 Ok(
552 completion::CompletionResponse::new(choice, usage, PROVIDER_NAME)
553 .with_optional_finish_reason(finish_reason)
554 .with_optional_response_id(
555 Some(response.response_id.clone()).filter(|id| !id.is_empty()),
556 )
557 .with_optional_model(model),
558 )
559 }
560}
561
562impl ProviderResponseExt for GenerateContentResponse {
564 type Usage = proto::UsageMetadata;
565
566 fn get_response_id(&self) -> Option<String> {
567 if self.response_id.is_empty() {
568 None
569 } else {
570 Some(self.response_id.clone())
571 }
572 }
573
574 fn get_response_model_name(&self) -> Option<String> {
575 if self.model_version.is_empty() {
576 None
577 } else {
578 Some(self.model_version.clone())
579 }
580 }
581
582 fn get_text_response(&self) -> Option<String> {
583 self.candidates.first().and_then(|c| {
584 c.content.as_ref().and_then(|content| {
585 let text: Vec<String> = content
586 .parts
587 .iter()
588 .filter(|part| !part.thought)
594 .filter_map(|part| {
595 if let Some(proto::part::Data::Text(text)) = &part.data {
596 Some(text.clone())
597 } else {
598 None
599 }
600 })
601 .collect();
602
603 if text.is_empty() {
604 None
605 } else {
606 Some(text.join("\n"))
607 }
608 })
609 })
610 }
611
612 fn get_usage(&self) -> Option<Self::Usage> {
613 self.usage_metadata
614 }
615}
616
617fn decode_base64_bytes(input: &str) -> Result<Vec<u8>, CompletionError> {
618 let data = input.trim();
619
620 let data = if let Some(rest) = data.strip_prefix("data:") {
622 rest.split_once(',').map(|(_, b64)| b64).unwrap_or(data)
623 } else {
624 data
625 };
626
627 let mut last_err: Option<String> = None;
628
629 for engine in [
630 &base64::engine::general_purpose::STANDARD,
631 &base64::engine::general_purpose::URL_SAFE,
632 &base64::engine::general_purpose::STANDARD_NO_PAD,
633 &base64::engine::general_purpose::URL_SAFE_NO_PAD,
634 ] {
635 match engine.decode(data) {
636 Ok(bytes) => return Ok(bytes),
637 Err(err) => last_err = Some(err.to_string()),
638 }
639 }
640
641 let err = last_err.unwrap_or_else(|| "unknown base64 decode error".to_string());
642 Err(CompletionError::RequestError(
643 format!("Invalid base64 data: {err}").into(),
644 ))
645}
646
647fn decode_optional_base64(sig: Option<String>) -> Result<Vec<u8>, CompletionError> {
648 let Some(sig) = sig else {
649 return Ok(Vec::new());
650 };
651 decode_base64_bytes(&sig)
652}
653
654pub(crate) fn map_usage(usage: Option<&proto::UsageMetadata>) -> completion::Usage {
660 usage
661 .map(|usage| completion::Usage {
662 input_tokens: usage.prompt_token_count as u64,
663 output_tokens: usage.candidates_token_count as u64,
664 total_tokens: usage.total_token_count as u64,
665 cached_input_tokens: usage.cached_content_token_count as u64,
666 cache_creation_input_tokens: 0,
667 tool_use_prompt_tokens: 0,
668 reasoning_tokens: 0,
669 })
670 .unwrap_or_default()
671}
672
673pub(crate) fn encode_optional_base64(bytes: &[u8]) -> Option<String> {
674 if bytes.is_empty() {
675 None
676 } else {
677 Some(base64::engine::general_purpose::STANDARD.encode(bytes))
678 }
679}
680
681fn json_to_prost_struct(value: serde_json::Value) -> Result<proto::Struct, CompletionError> {
682 match value {
683 serde_json::Value::Object(map) => Ok(proto::Struct {
684 fields: map
685 .into_iter()
686 .map(|(k, v)| (k, json_to_prost_value(v)))
687 .collect(),
688 }),
689 _ => Err(CompletionError::RequestError(
690 "Expected a JSON object for google.protobuf.Struct".into(),
691 )),
692 }
693}
694
695fn json_to_prost_value(value: serde_json::Value) -> proto::Value {
696 match value {
697 serde_json::Value::Null => proto::Value {
698 kind: Some(proto::value::Kind::NullValue(
699 proto::NullValue::NullValue as i32,
700 )),
701 },
702 serde_json::Value::Bool(b) => proto::Value {
703 kind: Some(proto::value::Kind::BoolValue(b)),
704 },
705 serde_json::Value::Number(n) => proto::Value {
706 kind: Some(proto::value::Kind::NumberValue(
707 n.as_f64().unwrap_or_default(),
708 )),
709 },
710 serde_json::Value::String(s) => proto::Value {
711 kind: Some(proto::value::Kind::StringValue(s)),
712 },
713 serde_json::Value::Array(items) => proto::Value {
714 kind: Some(proto::value::Kind::ListValue(proto::ListValue {
715 values: items.into_iter().map(json_to_prost_value).collect(),
716 })),
717 },
718 serde_json::Value::Object(map) => proto::Value {
719 kind: Some(proto::value::Kind::StructValue(proto::Struct {
720 fields: map
721 .into_iter()
722 .map(|(k, v)| (k, json_to_prost_value(v)))
723 .collect(),
724 })),
725 },
726 }
727}
728
729pub(crate) fn prost_struct_to_json(st: &proto::Struct) -> serde_json::Value {
730 let mut out = serde_json::Map::with_capacity(st.fields.len());
731 for (k, v) in &st.fields {
732 out.insert(k.clone(), prost_value_to_json(v));
733 }
734 serde_json::Value::Object(out)
735}
736
737fn prost_value_to_json(v: &proto::Value) -> serde_json::Value {
738 match &v.kind {
739 None | Some(proto::value::Kind::NullValue(_)) => serde_json::Value::Null,
740 Some(proto::value::Kind::BoolValue(b)) => serde_json::Value::Bool(*b),
741 Some(proto::value::Kind::NumberValue(n)) => serde_json::Number::from_f64(*n)
742 .map(serde_json::Value::Number)
743 .unwrap_or(serde_json::Value::Null),
744 Some(proto::value::Kind::StringValue(s)) => serde_json::Value::String(s.clone()),
745 Some(proto::value::Kind::StructValue(st)) => prost_struct_to_json(st),
746 Some(proto::value::Kind::ListValue(list)) => {
747 serde_json::Value::Array(list.values.iter().map(prost_value_to_json).collect())
748 }
749 }
750}
751
752fn tool_parameters_to_proto_schema(
762 value: &serde_json::Value,
763) -> Result<Option<proto::Schema>, CompletionError> {
764 tool_parameters_to_schema(value.clone()).map(|schema| schema.map(gemini_schema_to_proto_schema))
765}
766
767fn gemini_schema_to_proto_schema(schema: GeminiSchema) -> proto::Schema {
768 proto::Schema {
769 r#type: json_type_to_proto_type(&schema.r#type) as i32,
770 format: schema.format.unwrap_or_default(),
771 description: schema.description.unwrap_or_default(),
772 nullable: schema.nullable.unwrap_or(false),
773 r#enum: schema.r#enum.unwrap_or_default(),
774 items: schema
775 .items
776 .map(|items| Box::new(gemini_schema_to_proto_schema(*items))),
777 properties: schema
778 .properties
779 .unwrap_or_default()
780 .into_iter()
781 .map(|(name, schema)| (name, gemini_schema_to_proto_schema(schema)))
782 .collect(),
783 required: schema.required.unwrap_or_default(),
784 }
785}
786
787fn json_type_to_proto_type(t: &str) -> proto::Type {
788 match t {
789 "string" => proto::Type::String,
790 "number" => proto::Type::Number,
791 "integer" => proto::Type::Integer,
792 "boolean" => proto::Type::Boolean,
793 "array" => proto::Type::Array,
794 "object" => proto::Type::Object,
795 "null" => proto::Type::Null,
796 _ => proto::Type::Unspecified,
797 }
798}
799
800#[cfg(test)]
801#[allow(clippy::expect_used, clippy::unwrap_used, clippy::panic)]
802mod tests {
803 use super::*;
804
805 #[test]
810 fn rpc_error_preserves_status_text_without_http_status() {
811 let status = tonic::Status::unavailable("boom");
812 let expected = status.to_string();
813
814 let err = rpc_error(status);
815
816 assert_eq!(err.provider_response_body(), Some(expected.as_str()));
819 assert_eq!(err.provider_response_status(), None);
820 }
821
822 #[test]
823 fn test_decode_base64_bytes_accepts_url_safe_with_padding() {
824 assert!(matches!(
825 decode_base64_bytes("_-wgVQA="),
826 Ok(bytes) if bytes == vec![0xFF, 0xEC, 0x20, 0x55, 0x00]
827 ));
828 }
829
830 #[test]
831 fn test_decode_base64_bytes_accepts_url_safe_no_pad() {
832 assert!(matches!(
833 decode_base64_bytes("_-wgVQA"),
834 Ok(bytes) if bytes == vec![0xFF, 0xEC, 0x20, 0x55, 0x00]
835 ));
836 }
837
838 #[test]
839 fn test_decode_base64_bytes_accepts_standard_no_pad() {
840 assert!(matches!(
841 decode_base64_bytes("Zg"),
842 Ok(bytes) if bytes == b"f".to_vec()
843 ));
844 }
845
846 #[test]
847 fn test_decode_base64_bytes_accepts_data_uri_prefix() {
848 assert!(matches!(
849 decode_base64_bytes("data:text/plain;base64,Zm9v"),
850 Ok(bytes) if bytes == b"foo".to_vec()
851 ));
852 }
853
854 #[test]
859 fn tool_params_empty_object_maps_to_none() {
860 let v = serde_json::json!({"type": "object", "properties": {}});
861 assert!(tool_parameters_to_proto_schema(&v).unwrap().is_none());
862 }
863
864 #[test]
865 fn tool_params_null_maps_to_none() {
866 assert!(
867 tool_parameters_to_proto_schema(&serde_json::Value::Null)
868 .unwrap()
869 .is_none()
870 );
871 }
872
873 #[test]
874 fn tool_params_object_with_scalar_properties_round_trips() {
875 let v = serde_json::json!({
876 "type": "object",
877 "properties": {
878 "city": { "type": "string", "description": "City name" },
879 "max_price": { "type": "integer", "description": "Cap, USD" }
880 },
881 "required": ["city"]
882 });
883
884 let schema = tool_parameters_to_proto_schema(&v)
885 .expect("schema conversion")
886 .expect("schema");
887 assert_eq!(schema.r#type, proto::Type::Object as i32);
888 assert_eq!(schema.required, vec!["city".to_string()]);
889 assert_eq!(schema.properties.len(), 2);
890
891 let city = schema.properties.get("city").expect("city prop");
892 assert_eq!(city.r#type, proto::Type::String as i32);
893 assert_eq!(city.description, "City name");
894
895 let max_price = schema.properties.get("max_price").expect("max_price prop");
896 assert_eq!(max_price.r#type, proto::Type::Integer as i32);
897 }
898
899 #[test]
900 fn tool_params_array_with_typed_items() {
901 let v = serde_json::json!({
902 "type": "array",
903 "items": { "type": "string" }
904 });
905
906 let schema = tool_parameters_to_proto_schema(&v)
907 .expect("schema conversion")
908 .expect("schema");
909 assert_eq!(schema.r#type, proto::Type::Array as i32);
910 let items = schema.items.expect("items");
911 assert_eq!(items.r#type, proto::Type::String as i32);
912 }
913
914 #[test]
915 fn tool_params_enum_strings_preserved() {
916 let v = serde_json::json!({
917 "type": "string",
918 "enum": ["celsius", "fahrenheit"]
919 });
920
921 let schema = tool_parameters_to_proto_schema(&v)
922 .expect("schema conversion")
923 .expect("schema");
924 assert_eq!(schema.r#type, proto::Type::String as i32);
925 assert_eq!(
926 schema.r#enum,
927 vec!["celsius".to_string(), "fahrenheit".to_string()]
928 );
929 }
930
931 #[test]
932 fn tool_params_resolves_defs_ref_properties() {
933 let v = serde_json::json!({
934 "type": "object",
935 "properties": {
936 "destination": { "$ref": "#/$defs/Destination" }
937 },
938 "required": ["destination"],
939 "$defs": {
940 "Destination": {
941 "type": "object",
942 "properties": {
943 "city": { "type": "string" },
944 "country_code": { "type": "string" }
945 },
946 "required": ["city"]
947 }
948 }
949 });
950
951 let schema = tool_parameters_to_proto_schema(&v)
952 .expect("schema conversion")
953 .expect("schema");
954 let destination = schema
955 .properties
956 .get("destination")
957 .expect("destination prop");
958
959 assert_eq!(destination.r#type, proto::Type::Object as i32);
960 assert_eq!(destination.required, vec!["city".to_string()]);
961 assert_eq!(
962 destination
963 .properties
964 .get("city")
965 .expect("city prop")
966 .r#type,
967 proto::Type::String as i32
968 );
969 }
970
971 #[test]
972 fn tool_params_nullable_type_array_preserves_non_null_type() {
973 let v = serde_json::json!({
974 "type": "object",
975 "properties": {
976 "nickname": { "type": ["null", "string"] }
977 }
978 });
979
980 let schema = tool_parameters_to_proto_schema(&v)
981 .expect("schema conversion")
982 .expect("schema");
983 let nickname = schema.properties.get("nickname").expect("nickname prop");
984
985 assert_eq!(nickname.r#type, proto::Type::String as i32);
986 assert!(nickname.nullable);
987 }
988
989 #[test]
990 fn tool_params_any_of_uses_non_null_schema() {
991 let v = serde_json::json!({
992 "anyOf": [
993 { "type": "null" },
994 {
995 "type": "object",
996 "properties": {
997 "query": { "type": "string" }
998 },
999 "required": ["query"]
1000 }
1001 ]
1002 });
1003
1004 let schema = tool_parameters_to_proto_schema(&v)
1005 .expect("schema conversion")
1006 .expect("schema");
1007
1008 assert_eq!(schema.r#type, proto::Type::Object as i32);
1009 assert!(schema.nullable);
1010 assert_eq!(schema.required, vec!["query".to_string()]);
1011 assert_eq!(
1012 schema.properties.get("query").expect("query prop").r#type,
1013 proto::Type::String as i32
1014 );
1015 }
1016
1017 #[test]
1018 fn tool_params_array_without_items_defaults_to_string_items() {
1019 let v = serde_json::json!({ "type": "array" });
1020
1021 let schema = tool_parameters_to_proto_schema(&v)
1022 .expect("schema conversion")
1023 .expect("schema");
1024
1025 assert_eq!(schema.r#type, proto::Type::Array as i32);
1026 assert_eq!(
1027 schema.items.expect("items").r#type,
1028 proto::Type::String as i32
1029 );
1030 }
1031
1032 #[test]
1036 fn create_grpc_request_sends_the_executed_name_not_an_identifier() {
1037 use rig_core::message::{
1038 AssistantContent, ProviderCallId, ToolCall, ToolCallId, ToolFunction, ToolResult,
1039 ToolResultContent,
1040 };
1041
1042 let call = |wire_id: &str, name: &str| message::Message::Assistant {
1043 id: None,
1044 content: vec![AssistantContent::ToolCall(ToolCall::from_wire(
1045 wire_id,
1046 ToolFunction {
1047 name: name.to_owned(),
1048 arguments: serde_json::json!({}),
1049 },
1050 ))],
1051 };
1052 let result = |wire_id: &str, name: &str| message::Message::User {
1053 content: vec![message::UserContent::ToolResult(ToolResult {
1054 call: ToolCallId::new_or_mint(wire_id),
1055 provider: ProviderCallId::new(wire_id),
1056 name: name.to_owned(),
1057 content: vec![ToolResultContent::text("out")],
1058 })],
1059 };
1060
1061 let req = create_grpc_request(
1062 "gemini-2.5-flash".to_string(),
1063 CompletionRequest {
1064 model: None,
1065 preamble: None,
1066 chat_history: vec![
1067 call("call_1", "add"),
1070 result("call_1", "sum"),
1071 call("call_abc", "get_weather"),
1074 result("call_abc", "get_weather"),
1075 ],
1076 documents: Vec::new(),
1077 tools: Vec::new(),
1078 temperature: None,
1079 max_tokens: None,
1080 tool_choice: None,
1081 additional_params: None,
1082 output_schema: None,
1083 record_telemetry_content: false,
1084 },
1085 )
1086 .expect("request build");
1087
1088 let responses: Vec<(&str, &str)> = req
1091 .contents
1092 .iter()
1093 .flat_map(|content| content.parts.iter())
1094 .filter_map(|part| match &part.data {
1095 Some(proto::part::Data::FunctionResponse(fr)) => {
1096 Some((fr.id.as_str(), fr.name.as_str()))
1097 }
1098 _ => None,
1099 })
1100 .collect();
1101 assert_eq!(
1102 responses,
1103 vec![("call_1", "sum"), ("call_abc", "get_weather")]
1104 );
1105 }
1106
1107 #[test]
1108 fn create_grpc_request_populates_tool_parameters() {
1109 use rig_core::completion::ToolDefinition;
1110
1111 let tool = ToolDefinition {
1112 name: "get_weather".to_string(),
1113 description: "Look up the current weather for a city.".to_string(),
1114 parameters: serde_json::json!({
1115 "type": "object",
1116 "properties": {
1117 "city": { "type": "string", "description": "City name" }
1118 },
1119 "required": ["city"]
1120 }),
1121 };
1122
1123 let req = create_grpc_request(
1124 "gemini-2.5-flash".to_string(),
1125 CompletionRequest {
1126 model: None,
1127 preamble: None,
1128 chat_history: vec![message::Message::user("forecast in Berlin?")],
1129 documents: Vec::new(),
1130 tools: vec![tool],
1131 temperature: None,
1132 max_tokens: None,
1133 tool_choice: None,
1134 additional_params: None,
1135 output_schema: None,
1136 record_telemetry_content: false,
1137 },
1138 )
1139 .expect("request build");
1140
1141 assert_eq!(req.tools.len(), 1);
1142 let tool = req.tools.first().expect("tool entry");
1143 let decl = tool
1144 .function_declarations
1145 .first()
1146 .expect("function declaration");
1147 assert_eq!(decl.name, "get_weather");
1148
1149 let params = decl.parameters.as_ref().expect("parameters populated");
1151 assert_eq!(params.r#type, proto::Type::Object as i32);
1152 assert_eq!(params.required, vec!["city".to_string()]);
1153 assert!(params.properties.contains_key("city"));
1154 }
1155
1156 #[test]
1162 fn get_text_response_skips_thought_parts() {
1163 let response = proto::GenerateContentResponse {
1164 candidates: vec![proto::Candidate {
1165 content: Some(proto::Content {
1166 parts: vec![
1167 proto::Part {
1168 data: Some(proto::part::Data::Text(
1169 "Let me work through this...".to_string(),
1170 )),
1171 thought: true,
1172 ..Default::default()
1173 },
1174 proto::Part {
1175 data: Some(proto::part::Data::Text("The answer is 42.".to_string())),
1176 thought: false,
1177 ..Default::default()
1178 },
1179 ],
1180 ..Default::default()
1181 }),
1182 ..Default::default()
1183 }],
1184 ..Default::default()
1185 };
1186
1187 assert_eq!(
1188 response.get_text_response().as_deref(),
1189 Some("The answer is 42."),
1190 "reasoning must not be reported as the response text"
1191 );
1192 }
1193
1194 #[test]
1199 fn a_trailing_thought_signature_signs_the_reasoning_before_it() {
1200 let response = proto::GenerateContentResponse {
1201 candidates: vec![proto::Candidate {
1202 content: Some(proto::Content {
1203 parts: vec![
1204 proto::Part {
1205 data: Some(proto::part::Data::Text("the chain".to_string())),
1206 thought: true,
1207 ..Default::default()
1208 },
1209 proto::Part {
1210 data: Some(proto::part::Data::Text("answer".to_string())),
1211 thought: false,
1212 thought_signature: b"sig-bytes".to_vec(),
1213 ..Default::default()
1214 },
1215 ],
1216 ..Default::default()
1217 }),
1218 ..Default::default()
1219 }],
1220 ..Default::default()
1221 };
1222
1223 let normalized: completion::CompletionResponse =
1224 response.try_into().expect("payload should normalize");
1225 assert_eq!(
1226 normalized.choice.len(),
1227 2,
1228 "no empty sibling; got {:?}",
1229 normalized.choice
1230 );
1231 assert!(
1232 matches!(
1233 normalized.choice.first(),
1234 Some(completion::AssistantContent::Reasoning(reasoning))
1235 if matches!(reasoning.content.first(),
1236 Some(message::ReasoningContent::Text { text, signature })
1237 if text == "the chain" && signature.is_some())
1238 ),
1239 "the reasoning block must carry the trailing signature, got {:?}",
1240 normalized.choice
1241 );
1242 }
1243
1244 #[test]
1255 fn generate_content_response_round_trips_through_serde_json_value() {
1256 let raw = proto::GenerateContentResponse {
1257 candidates: vec![proto::Candidate {
1258 content: Some(proto::Content {
1259 parts: vec![proto::Part {
1260 data: Some(proto::part::Data::Text("hello".to_string())),
1261 ..Default::default()
1262 }],
1263 role: "model".to_string(),
1264 }),
1265 finish_reason: proto::candidate::FinishReason::Stop as i32,
1266 index: Some(0),
1267 finish_message: Some("done".to_string()),
1268 }],
1269 usage_metadata: Some(proto::UsageMetadata {
1270 prompt_token_count: 10,
1271 candidates_token_count: 20,
1272 total_token_count: 30,
1273 cached_content_token_count: 4,
1274 }),
1275 model_version: "gemini-2.5-flash".to_string(),
1276 response_id: "resp-grpc-1".to_string(),
1277 prompt_feedback: None,
1278 };
1279
1280 let value = serde_json::to_value(&raw).expect("serialize");
1281 assert_eq!(
1282 value.pointer("/usage_metadata/cached_content_token_count"),
1283 Some(&serde_json::json!(4))
1284 );
1285 assert_eq!(
1286 value.pointer("/candidates/0/finish_message"),
1287 Some(&serde_json::json!("done"))
1288 );
1289 assert_eq!(
1290 value.pointer("/model_version"),
1291 Some(&serde_json::json!("gemini-2.5-flash"))
1292 );
1293
1294 let back: proto::GenerateContentResponse =
1295 serde_json::from_value(value.clone()).expect("deserialize");
1296 assert_eq!(
1297 serde_json::to_value(&back).expect("re-serialize"),
1298 value,
1299 "the capture must read back into GenerateContentResponse and re-serialize identically"
1300 );
1301 assert_eq!(back, raw);
1302
1303 let original: completion::CompletionResponse = raw.try_into().expect("original converts");
1304 let restored: completion::CompletionResponse = back.try_into().expect("restored converts");
1305 assert_eq!(restored.identity(), original.identity());
1306 assert_eq!(restored.finish_reason(), original.finish_reason());
1307 assert_eq!(restored.model, original.model);
1308 assert_eq!(restored.usage, original.usage);
1309 assert_eq!(restored.choice, original.choice);
1310 assert_eq!(
1311 restored.identity().response_id.as_deref(),
1312 Some("resp-grpc-1")
1313 );
1314 assert_eq!(
1315 restored.finish_reason(),
1316 Some(completion::FinishReason::Stop)
1317 );
1318 }
1319}