mermaid_cli/providers/model/
gemini.rs1use async_trait::async_trait;
9
10use mermaid_domain::ChatRequest;
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
18pub 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 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 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 (relay_tx, relay_handle) = super::stream_bridge::ordered_relay(ctx.sink.clone());
75 let callback = super::stream_bridge::forward_callback(relay_tx.clone());
76 let chat_fut = self
77 .adapter
78 .chat(&request.messages, &config, Some(callback));
79
80 let response = tokio::select! {
81 biased;
82 _ = ctx.token.cancelled() => {
83 return Err(ModelError::Cancelled);
84 },
85 r = chat_fut => r?,
86 };
87
88 let usage = response.usage.clone();
89 let stop_reason = response.stop_reason.clone();
90 let _ = relay_tx.send(StreamEvent::Done {
92 usage: usage.clone(),
93 provider_continuation: None,
94 stop_reason: stop_reason.clone(),
95 });
96 drop(relay_tx);
97 mermaid_model::utils::join_logged(relay_handle.take(), "stream_relay").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.tools.iter().map(|t| t.to_openai_json()).collect(),
117 output_schema: request.output_schema.clone(),
118 ..Default::default()
119 }
120}
121
122#[cfg(test)]
123mod tests {
124 use super::*;
125
126 #[test]
127 fn build_model_config_maps_fields() {
128 let req = ChatRequest {
129 model_id: "gemini/gemini-3.1-pro-preview".to_string(),
130 messages: vec![],
131 system_prompt: "sys".to_string(),
132 instructions: None,
133 reasoning: mermaid_model::models::ReasoningLevel::High,
134 temperature: 0.5,
135 max_tokens: 4096,
136 tools: vec![],
137
138 ollama_num_ctx: None,
139 ollama_allow_ram_offload: None,
140 resolved_context_window: None,
141 resolved_max_output: None,
142 output_schema: None,
143 suppress_auto_compact: false,
144 suppressed_builtin_tools: Vec::new(),
145 };
146 let cfg = build_model_config(&req);
147 assert_eq!(cfg.reasoning, mermaid_model::models::ReasoningLevel::High);
148 assert_eq!(cfg.temperature, 0.5);
149 assert!(cfg.dynamic_system_suffix.is_none());
150 }
151}