mermaid_cli/providers/model/
gemini.rs1use async_trait::async_trait;
9
10use crate::domain::ChatRequest;
11use crate::models::adapters::gemini::GeminiAdapter;
12use crate::models::{Model, ModelConfig, ModelError, Result};
13
14use super::super::capabilities::Capabilities;
15use super::super::ctx::{FinalResponse, StreamContext, StreamEvent};
16use super::{ContextSizing, ModelProvider, resolve_limits_cached};
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: Capabilities,
28}
29
30impl GeminiProvider {
31 pub fn new(api_key: String, model_name: String, base_url: String) -> Result<Self> {
32 let adapter = GeminiAdapter::new(api_key, model_name, base_url)?;
33 let capabilities = Capabilities::from_legacy(adapter.capabilities());
34 Ok(Self {
35 adapter,
36 capabilities,
37 })
38 }
39}
40
41#[async_trait]
42impl ModelProvider for GeminiProvider {
43 fn capabilities(&self) -> &Capabilities {
44 &self.capabilities
45 }
46
47 async fn resolve_context_window(&self, request: &ChatRequest) -> ContextSizing {
52 let _ = request;
53 let model = Model::name(&self.adapter).to_string();
54 let limits =
55 resolve_limits_cached("gemini", &model, || self.adapter.fetch_model_limits()).await;
56 let window = limits.as_ref().and_then(|l| l.max_context_tokens);
57 ContextSizing {
58 model_max: window,
59 effective: window,
60 source: None,
61 max_output: limits.as_ref().and_then(|l| l.max_output_tokens),
62 }
63 }
64
65 async fn chat(&self, request: ChatRequest, ctx: StreamContext) -> Result<FinalResponse> {
66 let config = build_model_config(&request);
67 let (relay_tx, relay_handle) = super::stream_bridge::ordered_relay(ctx.sink.clone());
68 let callback = super::stream_bridge::forward_callback(relay_tx.clone());
69 let chat_fut = self
70 .adapter
71 .chat(&request.messages, &config, Some(callback));
72
73 let response = tokio::select! {
74 biased;
75 _ = ctx.token.cancelled() => {
76 return Err(ModelError::Cancelled);
77 },
78 r = chat_fut => r?,
79 };
80
81 let usage = response.usage.clone();
82 let stop_reason = response.stop_reason.clone();
83 let _ = relay_tx.send(StreamEvent::Done {
85 usage: usage.clone(),
86 provider_continuation: None,
87 stop_reason: stop_reason.clone(),
88 });
89 drop(relay_tx);
90 crate::utils::join_logged(relay_handle.take(), "stream_relay").await;
91
92 Ok(FinalResponse {
93 usage,
94 provider_continuation: None,
95 tool_calls: response.tool_calls.unwrap_or_default(),
96 stop_reason,
97 })
98 }
99}
100
101fn build_model_config(request: &ChatRequest) -> ModelConfig {
102 ModelConfig {
103 model: request.model_id.clone(),
104 temperature: request.temperature,
105 max_tokens: request.max_tokens,
106 reasoning: request.reasoning,
107 system_prompt: Some(request.system_prompt.clone()),
108 dynamic_system_suffix: request.instructions.clone(),
109 tools: request.tools.iter().map(|t| t.to_openai_json()).collect(),
110 output_schema: request.output_schema.clone(),
111 ..Default::default()
112 }
113}
114
115#[cfg(test)]
116mod tests {
117 use super::*;
118
119 #[test]
120 fn build_model_config_maps_fields() {
121 let req = ChatRequest {
122 model_id: "gemini/gemini-3.1-pro-preview".to_string(),
123 messages: vec![],
124 system_prompt: "sys".to_string(),
125 instructions: None,
126 reasoning: crate::models::ReasoningLevel::High,
127 temperature: 0.5,
128 max_tokens: 4096,
129 tools: vec![],
130
131 ollama_num_ctx: None,
132 ollama_allow_ram_offload: None,
133 resolved_context_window: None,
134 resolved_max_output: None,
135 output_schema: None,
136 suppress_auto_compact: false,
137 suppressed_builtin_tools: Vec::new(),
138 };
139 let cfg = build_model_config(&req);
140 assert_eq!(cfg.reasoning, crate::models::ReasoningLevel::High);
141 assert_eq!(cfg.temperature, 0.5);
142 assert!(cfg.dynamic_system_suffix.is_none());
143 }
144}