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