Skip to main content

mermaid_cli/providers/model/
gemini.rs

1//! Gemini provider — wraps `models::adapters::gemini::GeminiAdapter`.
2//!
3//! Google's Gemini family uses a different wire format from OpenAI-
4//! compat (`:streamGenerateContent?alt=sse` + protobuf-ish JSON
5//! shape). The adapter handles all of that; this wrapper just
6//! forwards.
7
8use std::sync::Arc;
9
10use async_trait::async_trait;
11
12use crate::domain::ChatRequest;
13use crate::models::adapters::gemini::GeminiAdapter;
14use crate::models::{
15    Model, ModelConfig, ModelError, ReasoningChunk, Result, StreamCallback,
16    StreamEvent as ModelStreamEvent,
17};
18
19use super::super::capabilities::Capabilities;
20use super::super::ctx::{FinalResponse, StreamContext, StreamEvent};
21use super::ModelProvider;
22
23pub struct GeminiProvider {
24    adapter: GeminiAdapter,
25    capabilities: Capabilities,
26}
27
28impl GeminiProvider {
29    pub fn new(api_key: String, model_name: String, base_url: String) -> Result<Self> {
30        let adapter = GeminiAdapter::new(api_key, model_name, base_url)?;
31        let capabilities = Capabilities::from_legacy(adapter.capabilities());
32        Ok(Self {
33            adapter,
34            capabilities,
35        })
36    }
37}
38
39#[async_trait]
40impl ModelProvider for GeminiProvider {
41    fn capabilities(&self) -> &Capabilities {
42        &self.capabilities
43    }
44
45    async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse> {
46        let config = build_model_config(&request);
47        let (relay_tx, relay_handle) = super::stream_bridge::ordered_relay(ctx.sink.clone());
48        let callback = forward_callback(relay_tx.clone());
49        let chat_fut = self
50            .adapter
51            .chat(&request.messages, &config, Some(callback));
52
53        let response = tokio::select! {
54            biased;
55            _ = ctx.token.cancelled() => {
56                return Err(ModelError::Cancelled);
57            },
58            r = chat_fut => r?,
59        };
60
61        let usage = response.usage.clone();
62        let stop_reason = response.stop_reason.clone();
63        // Terminal Done through the ordered relay, then drain (see openai_compat).
64        let _ = relay_tx.send(StreamEvent::Done {
65            usage: usage.clone(),
66            thinking_signature: None,
67            stop_reason: stop_reason.clone(),
68        });
69        drop(relay_tx);
70        let _ = relay_handle.await;
71
72        Ok(FinalResponse {
73            usage,
74            thinking_signature: None,
75            tool_calls: response.tool_calls.unwrap_or_default(),
76            stop_reason,
77        })
78    }
79}
80
81fn build_model_config(request: &ChatRequest) -> ModelConfig {
82    ModelConfig {
83        model: request.model_id.clone(),
84        temperature: request.temperature,
85        max_tokens: request.max_tokens,
86        reasoning: request.reasoning,
87        system_prompt: Some(request.system_prompt.clone()),
88        dynamic_system_suffix: request.instructions.clone(),
89        tools: request.tools.iter().map(|t| t.to_openai_json()).collect(),
90        ..Default::default()
91    }
92}
93
94fn forward_callback(sink: tokio::sync::mpsc::UnboundedSender<StreamEvent>) -> StreamCallback {
95    Arc::new(move |event: ModelStreamEvent| {
96        let mapped = match event {
97            ModelStreamEvent::Text(s) => StreamEvent::Text(s),
98            ModelStreamEvent::Reasoning(chunk) => StreamEvent::Reasoning(ReasoningChunk {
99                text: chunk.text,
100                signature: chunk.signature,
101            }),
102            ModelStreamEvent::ToolCall(tc) => StreamEvent::ToolCall(tc),
103            ModelStreamEvent::Done { tokens } => StreamEvent::Done {
104                usage: if tokens > 0 {
105                    Some(crate::models::TokenUsage::provider(0, tokens, tokens))
106                } else {
107                    None
108                },
109                thinking_signature: None,
110                stop_reason: None,
111            },
112        };
113        let _ = sink.send(mapped);
114    })
115}
116
117#[cfg(test)]
118mod tests {
119    use super::*;
120
121    #[test]
122    fn build_model_config_maps_fields() {
123        let req = ChatRequest {
124            model_id: "gemini/gemini-3.1-pro-preview".to_string(),
125            messages: vec![],
126            system_prompt: "sys".to_string(),
127            instructions: None,
128            reasoning: crate::models::ReasoningLevel::High,
129            temperature: 0.5,
130            max_tokens: 4096,
131            tools: vec![],
132
133            ollama_num_ctx: None,
134            ollama_allow_ram_offload: None,
135        };
136        let cfg = build_model_config(&req);
137        assert_eq!(cfg.reasoning, crate::models::ReasoningLevel::High);
138        assert_eq!(cfg.temperature, 0.5);
139        assert!(cfg.dynamic_system_suffix.is_none());
140    }
141}