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::{ContextSizing, ModelProvider, resolve_limits_cached};
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    /// Live limit discovery via Gemini's models endpoint (`GET {base}/models/
46    /// {id}` → `inputTokenLimit` window + `outputTokenLimit` output ceiling).
47    /// Cache-first via `provider_probes` (TTL-bounded), one live fetch on a
48    /// miss; a fetch failure resolves all-`None`.
49    async fn resolve_context_window(&self, request: &ChatRequest) -> ContextSizing {
50        let _ = request;
51        let model = Model::name(&self.adapter).to_string();
52        let limits =
53            resolve_limits_cached("gemini", &model, || self.adapter.fetch_model_limits()).await;
54        let window = limits.as_ref().and_then(|l| l.max_context_tokens);
55        ContextSizing {
56            model_max: window,
57            effective: window,
58            source: None,
59            max_output: limits.as_ref().and_then(|l| l.max_output_tokens),
60        }
61    }
62
63    async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse> {
64        let config = build_model_config(&request);
65        let (relay_tx, relay_handle) = super::stream_bridge::ordered_relay(ctx.sink.clone());
66        let callback = forward_callback(relay_tx.clone());
67        let chat_fut = self
68            .adapter
69            .chat(&request.messages, &config, Some(callback));
70
71        let response = tokio::select! {
72            biased;
73            _ = ctx.token.cancelled() => {
74                return Err(ModelError::Cancelled);
75            },
76            r = chat_fut => r?,
77        };
78
79        let usage = response.usage.clone();
80        let stop_reason = response.stop_reason.clone();
81        // Terminal Done through the ordered relay, then drain (see openai_compat).
82        let _ = relay_tx.send(StreamEvent::Done {
83            usage: usage.clone(),
84            provider_continuation: None,
85            stop_reason: stop_reason.clone(),
86        });
87        drop(relay_tx);
88        crate::utils::join_logged(relay_handle.take(), "stream_relay").await;
89
90        Ok(FinalResponse {
91            usage,
92            provider_continuation: None,
93            tool_calls: response.tool_calls.unwrap_or_default(),
94            stop_reason,
95        })
96    }
97}
98
99fn build_model_config(request: &ChatRequest) -> ModelConfig {
100    ModelConfig {
101        model: request.model_id.clone(),
102        temperature: request.temperature,
103        max_tokens: request.max_tokens,
104        reasoning: request.reasoning,
105        system_prompt: Some(request.system_prompt.clone()),
106        dynamic_system_suffix: request.instructions.clone(),
107        tools: request.tools.iter().map(|t| t.to_openai_json()).collect(),
108        output_schema: request.output_schema.clone(),
109        ..Default::default()
110    }
111}
112
113fn forward_callback(sink: tokio::sync::mpsc::UnboundedSender<StreamEvent>) -> StreamCallback {
114    Arc::new(move |event: ModelStreamEvent| {
115        let mapped = match event {
116            ModelStreamEvent::Text(s) => StreamEvent::Text(s),
117            ModelStreamEvent::Reasoning(chunk) => StreamEvent::Reasoning(ReasoningChunk {
118                text: chunk.text,
119                signature: chunk.signature,
120            }),
121            ModelStreamEvent::ToolCall(tc) => StreamEvent::ToolCall(tc),
122            ModelStreamEvent::Status(s) => StreamEvent::Status(s),
123            // No adapter emits `Done` through this callback — the wrapper
124            // sends the authoritative terminal `Done` built from the
125            // returned `ModelResponse` (F3). Map defensively without
126            // inventing usage (the old placeholder misfiled everything
127            // as completion tokens).
128            ModelStreamEvent::Done { .. } => StreamEvent::Done {
129                usage: None,
130                provider_continuation: None,
131                stop_reason: None,
132            },
133        };
134        let _ = sink.send(mapped);
135    })
136}
137
138#[cfg(test)]
139mod tests {
140    use super::*;
141
142    #[test]
143    fn build_model_config_maps_fields() {
144        let req = ChatRequest {
145            model_id: "gemini/gemini-3.1-pro-preview".to_string(),
146            messages: vec![],
147            system_prompt: "sys".to_string(),
148            instructions: None,
149            reasoning: crate::models::ReasoningLevel::High,
150            temperature: 0.5,
151            max_tokens: 4096,
152            tools: vec![],
153
154            ollama_num_ctx: None,
155            ollama_allow_ram_offload: None,
156            resolved_context_window: None,
157            resolved_max_output: None,
158            output_schema: None,
159            suppress_auto_compact: false,
160        };
161        let cfg = build_model_config(&req);
162        assert_eq!(cfg.reasoning, crate::models::ReasoningLevel::High);
163        assert_eq!(cfg.temperature, 0.5);
164        assert!(cfg.dynamic_system_suffix.is_none());
165    }
166}