llm/providers/gemini/
provider.rs1use crate::provider::{get_context_window, validate_reasoning};
2use crate::provider_connection::DEFAULT_STREAM_IDLE_TIMEOUT;
3use crate::providers::http::{http_client, openai_client};
4use crate::providers::openai_compatible::{AetherOpenAiConfig, build_chat_request, create_custom_stream_generic};
5use crate::providers::response_stream::error_stream;
6use crate::{
7 Context, LlmError, LlmResponseStream, ProviderAuthMode, ProviderConnectionConfig, ProviderFactory, Result,
8 StreamingModelProvider,
9};
10use std::env::var;
11use std::future::ready;
12use std::time::Duration;
13
14pub const GEMINI_API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta/openai/";
15
16#[derive(Clone)]
17pub struct GeminiProvider {
18 api_key: Option<String>,
19 base_url: Option<String>,
20 auth_mode: ProviderAuthMode,
21 model: String,
22 http: reqwest::Client,
23 idle_timeout: Duration,
24}
25
26impl GeminiProvider {
27 pub fn new(api_key: Option<String>) -> Self {
28 Self {
29 api_key,
30 base_url: None,
31 auth_mode: ProviderAuthMode::Default,
32 model: String::new(),
33 http: http_client(),
34 idle_timeout: DEFAULT_STREAM_IDLE_TIMEOUT,
35 }
36 }
37
38 pub fn with_connection(mut self, connection: ProviderConnectionConfig) -> Self {
39 self.base_url = connection.base_url;
40 self.auth_mode = connection.auth_mode;
41 self.idle_timeout = connection.idle_timeout;
42 self
43 }
44
45 fn get_api_key(&self) -> Result<String> {
46 if self.auth_mode == ProviderAuthMode::None {
47 return Ok(String::new());
48 }
49 if let Some(key) = &self.api_key {
50 return Ok(key.clone());
51 }
52
53 if let Ok(api_key) = var("GEMINI_API_KEY") {
54 return Ok(api_key);
55 }
56
57 Err(LlmError::MissingApiKey(
58 "GEMINI_API_KEY not set. Set the environment variable or provide an API key.".to_string(),
59 ))
60 }
61
62 fn build_openai_client(&self, api_key: &str) -> async_openai::Client<AetherOpenAiConfig> {
63 let api_base = self.base_url.as_deref().unwrap_or(GEMINI_API_BASE);
64 let config = async_openai::config::OpenAIConfig::new().with_api_key(api_key).with_api_base(api_base);
65 openai_client(AetherOpenAiConfig::new(config, self.auth_mode), self.http.clone())
66 }
67
68 fn try_stream_response(&self, context: &Context) -> Result<LlmResponseStream> {
69 validate_reasoning(context, self.model().as_ref())?;
70 let api_key = self.get_api_key()?;
71 let request = build_chat_request(&self.model, context, None)?;
72
73 tracing::info!("Using Gemini API with API key (OpenAI-compatible endpoint)");
74 Ok(create_custom_stream_generic(&self.build_openai_client(&api_key), request, self.idle_timeout))
75 }
76}
77
78impl ProviderFactory for GeminiProvider {
79 fn from_env() -> impl Future<Output = Result<Self>> + Send {
80 ready(Ok(Self::new(None)))
81 }
82
83 fn from_env_with_connection(connection: ProviderConnectionConfig) -> impl Future<Output = Result<Self>> + Send {
84 ready(Ok(Self::new(None).with_connection(connection)))
85 }
86
87 fn with_model(mut self, model: &str) -> Self {
88 self.model = model.to_string();
89 self
90 }
91}
92
93impl StreamingModelProvider for GeminiProvider {
94 fn model(&self) -> Option<crate::LlmModel> {
95 format!("gemini:{}", self.model).parse().ok()
96 }
97
98 fn context_window(&self) -> Option<u32> {
99 get_context_window("gemini", &self.model)
100 }
101
102 fn stream_response(&self, context: &Context) -> LlmResponseStream {
103 self.try_stream_response(context).unwrap_or_else(error_stream)
104 }
105
106 fn display_name(&self) -> String {
107 format!("Gemini ({})", self.model)
108 }
109}
110
111#[cfg(test)]
112mod tests {
113 use super::*;
114 use async_openai::config::Config;
115 use futures::StreamExt;
116 use reqwest::header::AUTHORIZATION;
117
118 #[tokio::test]
119 async fn disabled_uses_none_not_minimal() {
120 use crate::providers::test_capture_server::CaptureServer;
121 let model = crate::LlmModel::all()
122 .iter()
123 .find(|model| model.provider_enum() == crate::catalog::Provider::Gemini && model.supports_reasoning_off())
124 .unwrap();
125 let mut server = CaptureServer::start_chat_completions().await;
126 let provider =
127 GeminiProvider::new(None).with_model(&model.model_id()).with_connection(ProviderConnectionConfig {
128 base_url: Some(server.base_url.clone()),
129 auth_mode: ProviderAuthMode::None,
130 ..Default::default()
131 });
132 let mut context = Context::new(vec![crate::ChatMessage::user("Hello")], vec![]);
133 context.set_reasoning_effort(crate::ReasoningEffort::Disabled);
134 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
135 assert!(responses.iter().all(Result::is_ok), "{responses:?}");
136 assert_eq!(server.captured().await.body["reasoning_effort"], "none");
137 }
138
139 #[test]
140 fn test_provider_display_name() {
141 let provider = GeminiProvider::new(None).with_model("gemini-2.0-flash");
142 assert_eq!(provider.display_name(), "Gemini (gemini-2.0-flash)");
143 }
144
145 #[test]
146 fn get_api_key_returns_empty_when_auth_is_none() {
147 let provider = GeminiProvider::new(Some("real-key".to_string()))
148 .with_connection(ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() });
149 assert_eq!(provider.get_api_key().unwrap(), "");
150 }
151
152 #[test]
153 fn build_openai_client_strips_authorization_when_auth_is_none() {
154 let provider = GeminiProvider::new(Some("real-key".to_string()))
155 .with_connection(ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() });
156 let api_key = provider.get_api_key().unwrap();
157 let client = provider.build_openai_client(&api_key);
158 assert!(!client.config().headers().contains_key(AUTHORIZATION));
159 }
160}