mermaid_cli/providers/model/
gemini.rs1use 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 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 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 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}