Skip to main content

llm/providers/gemini/
provider.rs

1use 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}