Skip to main content

rig_gemini_grpc/
completion.rs

1// ================================================================
2//! Google Gemini gRPC Completion Integration
3// ================================================================
4
5/// `gemini-2.5-flash` completion model
6pub const GEMINI_2_5_FLASH: &str = "gemini-2.5-flash";
7/// `gemini-2.0-flash-lite` completion model
8pub const GEMINI_2_0_FLASH_LITE: &str = "gemini-2.0-flash-lite";
9/// `gemini-2.0-flash` completion model
10pub 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// =================================================================
26// Rig Implementation Types
27// =================================================================
28
29#[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
44/// Stable descriptor name reported on normalized responses from this provider.
45pub const PROVIDER_NAME: &str = "gemini-grpc";
46
47/// Map Gemini's protobuf `finishReason` onto rig's normalized vocabulary.
48///
49/// The wire value is a prost enum discriminant; `as_str_name` recovers the
50/// SCREAMING_SNAKE proto spelling the shared Google table keys on, and a
51/// discriminant this proto does not model keeps its numeric identity so a
52/// reason Google adds later surfaces rather than reading as a natural stop.
53pub 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
65/// Turn a tool-protocol terminal `finishReason` into an error, mirroring the
66/// REST wire's `function_call_finish_reason_error`.
67///
68/// These reasons mean the turn ABORTED inside the tool protocol: the model
69/// emitted a call the API could not parse, called a tool that was not
70/// offered, or exceeded the per-turn call budget. The candidate that carries
71/// them has no usable tool call, so reporting the turn as merely "finished
72/// for some other reason" lets an agent loop read an aborted turn as a
73/// complete one. The REST surface has always failed here; the gRPC surface
74/// must not diverge.
75///
76/// Only the reasons this proto models are matched — REST's
77/// `MISSING_THOUGHT_SIGNATURE` / `MALFORMED_RESPONSE` have no protobuf
78/// discriminant in `v1beta`, so an unmapped value cannot masquerade as one.
79pub 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    /// Execute a completion and return Gemini's own protobuf response.
100    ///
101    /// This is the escape hatch for fields rig does not normalize;
102    /// [`completion::CompletionModel::completion`] calls it and maps the
103    /// result, so there is exactly one RPC either way.
104    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    /// Open a stream whose terminal record stays Gemini's own protobuf
125    /// response.
126    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        // Capture before `try_into` consumes the raw value.
143        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
157/// Build a non-thought `proto::Part` around the given data payload.
158pub(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
167/// Build a plain (non-thought) text `proto::Part`.
168pub(crate) fn text_part(text: String) -> proto::Part {
169    data_part(proto::part::Data::Text(text))
170}
171
172// Map a failed gRPC call into a `CompletionError` that preserves the provider's
173// error payload verbatim. gRPC is a non-HTTP transport, so there is no
174// `http::StatusCode`; the body is preserved via `from_provider_body` (status:
175// None) rather than a Rig-prefixed `ProviderError` diagnostic. Note: tonic does
176// not distinguish a server-returned gRPC error from a transport/connection
177// failure, so a pure connection error is also preserved here rather than gated
178// out as a Rig diagnostic the way Bedrock's typed service errors are.
179pub(crate) fn rpc_error(status: tonic::Status) -> CompletionError {
180    CompletionError::from_provider_body(status.to_string())
181}
182
183// Helper function to create gRPC request from Rig's CompletionRequest
184pub(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    // functionResponse.name keys the replay: cross-provider ingested
204    // results arrive with an empty name and their call carries it.
205    rig_core::providers::internal::resolve_empty_tool_result_names(&mut chat_history);
206    let mut contents = Vec::new();
207
208    // Convert chat history to gRPC Content messages
209    for msg in chat_history {
210        contents.push(rig_message_to_grpc_content(msg)?);
211    }
212
213    // Handle system instruction (preamble)
214    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    // Handle generation config
235    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    // Handle tools (functions)
246    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
279// Convert Rig message to gRPC Content
280fn 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
312// Convert Rig UserContent to gRPC Part
313fn 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            // `FunctionResponse.name` is the executed function's name —
340            // required data on the result. Only a provider-issued id may
341            // travel back on the wire (the proto field is optional-empty).
342            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
408// Convert Rig AssistantContent to gRPC Part
409fn 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                    // Only a provider-issued id may travel back on the
423                    // wire; minted correlation handles stay internal.
424                    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
445// Convert gRPC GenerateContentResponse to Rig CompletionResponse
446impl 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        // Same helper (and therefore the same message) as the streaming path,
455        // so a tool-protocol abort reads identically on both surfaces.
456        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                    // An id-less call mints its correlation handle —
511                    // never name-as-id, which collides two same-tool calls.
512                    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            // The wire hangs a `thoughtSignature` on a trailing part carrying
530            // no `thought` flag, and this crate's own streaming adapter keeps
531            // it (`streaming.rs`, the non-thought text arm) while this mapper
532            // dropped it — the same blocking/streaming asymmetry the REST wire
533            // had. One shared rule places it on both transports.
534            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
562// Implement ProviderResponseExt for telemetry
563impl 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                    // `thought` marks the model's chain-of-thought, which the
589                    // completion mapper above routes to `Reasoning`. A reader
590                    // that wants the response *text* must skip it, or it
591                    // reports reasoning as the answer — the same defect the
592                    // REST wire carried.
593                    .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    // Allow `data:<mime>;base64,<data>` inputs.
621    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
654/// Map Gemini's `UsageMetadata` onto rig's normalized `Usage`.
655///
656/// Known gap (unchanged here): `tool_use_prompt_token_count` and
657/// `thoughts_token_count` are not yet surfaced, so `tool_use_prompt_tokens`
658/// and `reasoning_tokens` read as 0.
659pub(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
752// Convert the JSON Schema carried by `ToolDefinition.parameters` into the typed
753// `proto::Schema` expected by `FunctionDeclaration.parameters`.
754//
755// Without this, every tool was sent to Gemini with `parameters = None`, which
756// caused the model to invoke tools with no argument shape (issue #1710).
757//
758// An empty object schema (`{"type": "object", "properties": {}}`, the default
759// when a tool takes no arguments) is mapped to `None` rather than a vacuous
760// schema, matching the convention used by `rig-core::providers::gemini`.
761fn 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    // ============================================================
806    // rpc_error — pins the from_provider_body usage on the RPC error path
807    // ============================================================
808
809    #[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        // The raw provider error text is preserved verbatim, and there is no
817        // HTTP status because gRPC is a non-HTTP transport.
818        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    // ============================================================
855    // tool_parameters_to_proto_schema — regression coverage for #1710
856    // ============================================================
857
858    #[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    /// `FunctionResponse.name` is the executed function's name: read from
1033    /// the required `ToolResult::name` — never an identifier, no matter how
1034    /// identifier-shaped the correlation handles are.
1035    #[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                    // Driver-built: the executed name travels as data (a
1068                    // repair hook renamed the call: `sum` ran, not `add`).
1069                    call("call_1", "add"),
1070                    result("call_1", "sum"),
1071                    // Cross-provider history with an OpenAI-shaped id —
1072                    // `call_abc` must never travel as the name.
1073                    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        // The name is the executed tool's name; the proto `id` is the
1089        // provider-issued call id (never rig's minted handle).
1090        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        // The regression in #1710 was `parameters: None` here.
1150        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    /// The gRPC wire carries the model's chain-of-thought in the same `parts`
1157    /// array as the answer, flagged by `thought` — same shape as the REST
1158    /// wire, where reading it as output text was a live-confirmed defect.
1159    /// There is no cassette harness for this transport (it is protobuf over
1160    /// gRPC, not HTTP), so the wire shape is stated directly.
1161    #[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    /// The wire hangs a `thoughtSignature` on a trailing part with no
1195    /// `thought` flag. This crate's streaming adapter has always kept it; the
1196    /// unary mapper dropped it, the same asymmetry the REST wire carried. The
1197    /// signature belongs to the chain-of-thought block that precedes it.
1198    #[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    /// The load-bearing property behind `CompletionResponse::raw` for the
1245    /// gRPC provider: the captured value is
1246    /// `serde_json::to_value(&GenerateContentResponse)` — the prost message
1247    /// `raw_completion` returns, with the serde derives `build.rs` attaches to
1248    /// every generated type — and a consumer must be able to read it back as
1249    /// the same message and get the same JSON. There is no cassette harness
1250    /// for gRPC, so this is the unit-form pin. Fields rig never normalizes
1251    /// (`cached_content_token_count` under `usage_metadata`, the candidate's
1252    /// `finish_message`) survive both directions, and normalizing the
1253    /// restored message agrees with normalizing the original.
1254    #[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}