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 async_trait::async_trait;
9
10use mermaid_domain::{ChatRequest, ToolDefinition};
11use mermaid_model::models::adapters::gemini::GeminiAdapter;
12use mermaid_model::models::{Model, ModelConfig, ModelError, Result};
13
14use super::super::ctx::{FinalResponse, StreamContext, StreamEvent};
15use super::{ContextSizing, ModelProvider, resolve_limits_cached};
16use mermaid_model::models::ModelCapabilities;
17
18/// Gemini's AI Studio root, and the env vars its key lives in. `LEGACY_API_KEY_ENV`
19/// predates Google's rename and is still accepted when `GOOGLE_API_KEY` is unset —
20/// but only when the user has not pointed at a specific var themselves.
21pub const DEFAULT_BASE_URL: &str = "https://generativelanguage.googleapis.com/v1beta";
22pub const DEFAULT_API_KEY_ENV: &str = "GOOGLE_API_KEY";
23pub const LEGACY_API_KEY_ENV: &str = "GEMINI_API_KEY";
24
25pub struct GeminiProvider {
26    adapter: GeminiAdapter,
27    capabilities: ModelCapabilities,
28}
29
30impl GeminiProvider {
31    /// Wrap a fresh [`GeminiAdapter`] as a `ModelProvider`.
32    ///
33    /// # Errors
34    ///
35    /// Only [`GeminiAdapter::new`]'s — the HTTP client build. The API is not
36    /// contacted here, so an invalid key or unreachable `base_url` still
37    /// constructs and fails on the first request.
38    pub fn new(api_key: String, model_name: String, base_url: String) -> Result<Self> {
39        let adapter = GeminiAdapter::new(api_key, model_name, base_url)?;
40        let capabilities = adapter.capabilities().clone();
41        Ok(Self {
42            adapter,
43            capabilities,
44        })
45    }
46}
47
48#[async_trait]
49impl ModelProvider for GeminiProvider {
50    fn capabilities(&self) -> &ModelCapabilities {
51        &self.capabilities
52    }
53
54    /// Live limit discovery via Gemini's models endpoint (`GET {base}/models/
55    /// {id}` → `inputTokenLimit` window + `outputTokenLimit` output ceiling).
56    /// Cache-first via `provider_probes` (TTL-bounded), one live fetch on a
57    /// miss; a fetch failure resolves all-`None`.
58    async fn resolve_context_window(&self, request: &ChatRequest) -> ContextSizing {
59        let _ = request;
60        let model = Model::name(&self.adapter).to_string();
61        let limits =
62            resolve_limits_cached("gemini", &model, || self.adapter.fetch_model_limits()).await;
63        let window = limits.as_ref().and_then(|l| l.max_context_tokens);
64        ContextSizing {
65            model_max: window,
66            effective: window,
67            source: None,
68            max_output: limits.as_ref().and_then(|l| l.max_output_tokens),
69        }
70    }
71
72    async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse> {
73        let config = build_model_config(&request);
74        let chat_fut = self
75            .adapter
76            .chat(&request.messages, &config, Some(ctx.sink.clone()));
77
78        let response = tokio::select! {
79            biased;
80            _ = ctx.token.cancelled() => {
81                return Err(ModelError::Cancelled);
82            },
83            r = chat_fut => r?,
84        };
85
86        let usage = response.usage.clone();
87        let stop_reason = response.stop_reason.clone();
88        // The terminal Done goes on the same sink the adapter just finished
89        // writing to, so it cannot overtake a still-queued ToolCall.
90        let _ = ctx
91            .sink
92            .send(StreamEvent::Done {
93                usage: usage.clone(),
94                provider_continuation: None,
95                stop_reason: stop_reason.clone(),
96            })
97            .await;
98
99        Ok(FinalResponse {
100            usage,
101            provider_continuation: None,
102            tool_calls: response.tool_calls.unwrap_or_default(),
103            stop_reason,
104        })
105    }
106}
107
108fn build_model_config(request: &ChatRequest) -> ModelConfig {
109    ModelConfig {
110        model: request.model_id.clone(),
111        temperature: request.temperature,
112        max_tokens: request.max_tokens,
113        reasoning: request.reasoning,
114        system_prompt: Some(request.system_prompt.clone()),
115        dynamic_system_suffix: request.instructions.clone(),
116        tools: request
117            .tools
118            .iter()
119            .map(ToolDefinition::to_openai_json)
120            .collect(),
121        output_schema: request.output_schema.clone(),
122        ..Default::default()
123    }
124}
125
126#[cfg(test)]
127mod tests {
128    use super::*;
129
130    #[test]
131    fn build_model_config_maps_fields() {
132        let req = ChatRequest {
133            model_id: "gemini/gemini-3.1-pro-preview".to_string(),
134            messages: vec![],
135            system_prompt: "sys".to_string(),
136            instructions: None,
137            reasoning: mermaid_model::models::ReasoningLevel::High,
138            temperature: 0.5,
139            max_tokens: 4096,
140            tools: vec![],
141
142            ollama_num_ctx: None,
143            ollama_allow_ram_offload: None,
144            resolved_context_window: None,
145            resolved_max_output: None,
146            output_schema: None,
147            suppress_auto_compact: false,
148            suppressed_builtin_tools: Vec::new(),
149        };
150        let cfg = build_model_config(&req);
151        assert_eq!(cfg.reasoning, mermaid_model::models::ReasoningLevel::High);
152        assert_eq!(cfg.temperature, 0.5);
153        assert!(cfg.dynamic_system_suffix.is_none());
154    }
155}