Skip to main content

rig_gemini_grpc/
completion.rs

1//! The Gemini `GenerateContent` completion wire over gRPC.
2//!
3//! ```no_run
4//! use rig_core::Model;
5//! use rig_gemini_grpc::{GeminiGrpc, completion::{GEMINI_2_5_FLASH, GenerateContent}};
6//!
7//! # async fn example() -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
8//! let model = GeminiGrpc::new("API_KEY").await?.completion(GEMINI_2_5_FLASH);
9//! # let _ = model;
10//! # Ok(())
11//! # }
12//! ```
13
14/// `gemini-2.5-flash` completion model
15pub const GEMINI_2_5_FLASH: &str = "gemini-2.5-flash";
16/// `gemini-2.0-flash-lite` completion model
17pub const GEMINI_2_0_FLASH_LITE: &str = "gemini-2.0-flash-lite";
18/// `gemini-2.0-flash` completion model
19pub 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/// The `GenerateContent` endpoint for one model: `GenerateContent` for a
41/// unary call, `StreamGenerateContent` for a streamed one.
42#[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    /// The Gemini service issues this wire's reasoning, over gRPC or REST,
66    /// so that is the reasoning a request may replay.
67    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                        // Stop receiving after a tonic failure.
101                        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
120/// Stable descriptor name reported on normalized responses from this provider.
121pub const PROVIDER_NAME: &str = "gemini-grpc";
122
123/// The issuer this transport's reasoning records: the Gemini API service,
124/// which also serves the REST transport, so thought signatures move between
125/// the two.
126pub const REASONING_ISSUER: &str = rig_core::providers::gemini::completion::PROVIDER_NAME;
127
128/// [`REASONING_ISSUER`], the only issuer whose reasoning this wire replays.
129const ISSUER: message::Issuer = message::Issuer::from_static(REASONING_ISSUER);
130
131/// Map Gemini's protobuf `finishReason` onto rig's normalized vocabulary.
132///
133/// The wire value is a prost enum discriminant; `as_str_name` recovers the
134/// SCREAMING_SNAKE proto spelling the shared Google table keys on, and a
135/// discriminant this proto does not model keeps its numeric identity so a
136/// reason Google adds later surfaces rather than reading as a natural stop.
137pub 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
149/// Returns a response error for malformed calls, unexpected calls, or exceeded
150/// tool-call limits, including the supplied finish message. Other discriminants
151/// return `None`.
152pub 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
171/// Build a non-thought `proto::Part` around the given data payload.
172pub(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
181/// Build a plain (non-thought) text `proto::Part`.
182pub(crate) fn text_part(text: String) -> proto::Part {
183    data_part(proto::part::Data::Text(text))
184}
185
186/// Preserves tonic status display text with RPC code and retry classification.
187/// Transport failures use the same provider-body representation.
188pub(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
194/// The gRPC status code's canonical name (`UNAVAILABLE`): the code the
195/// provider answered with, kept apart from the message so a report can
196/// key on it.
197pub(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
209/// Recognizes transient gRPC codes; other codes are non-transient.
210pub(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                // Reasoning another service issued is not replayed.
322                .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            // Replay the function name and only provider-issued IDs; local
368            // correlation handles must not reach the wire.
369            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                    // Only a provider-issued id may travel back on the
449                    // wire; rig-issued ids stay internal.
450                    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    // Allow `data:<mime>;base64,<data>` inputs.
481    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
512/// Map Gemini's `UsageMetadata` onto rig's normalized `Usage`.
513///
514/// Rig's input is the prompt plus the tool-use prompt, its output the
515/// candidates plus the thoughts, and its total their sum, which is Gemini's
516/// `total_token_count`. Proto3 cannot tell an unsent count from zero, so the
517/// tool-use and reasoning counts are always reported; Gemini reports no
518/// cache-write count, which stays `None`.
519pub(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
616/// Converts tool parameters to protobuf schema through the shared Gemini conversion.
617/// Empty object schemas map to `None`.
618fn 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;