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