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        let provider = self.clone();
84        let context = context.clone();
85
86        Box::pin(stream! {
87            let api_key = match provider.get_api_key() {
88                Ok(key) => key,
89                Err(e) => {
90                    yield Err(e);
91                    return;
92                }
93            };
94
95            tracing::info!("Using Gemini API with API key (OpenAI-compatible endpoint)");
96            let client = provider.build_openai_client(&api_key);
97            let request = match build_chat_request(&provider.model, &context, None) {
98                Ok(req) => req,
99                Err(e) => {
100                    yield Err(e);
101                    return;
102                }
103            };
104            let mut inner_stream =
105                create_custom_stream_generic(&client, request);
106
107            while let Some(result) = inner_stream.next().await {
108                yield result;
109            }
110        })
111    }
112
113    fn display_name(&self) -> String {
114        format!("Gemini ({})", self.model)
115    }
116}
117
118#[cfg(test)]
119mod tests {
120    use super::*;
121    use async_openai::config::Config;
122    use reqwest::header::AUTHORIZATION;
123
124    #[test]
125    fn test_provider_display_name() {
126        let provider = GeminiProvider::new(None).with_model("gemini-2.0-flash");
127        assert_eq!(provider.display_name(), "Gemini (gemini-2.0-flash)");
128    }
129
130    #[test]
131    fn get_api_key_returns_empty_when_auth_is_none() {
132        let provider = GeminiProvider::new(Some("real-key".to_string()))
133            .with_connection(ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() });
134        assert_eq!(provider.get_api_key().unwrap(), "");
135    }
136
137    #[test]
138    fn build_openai_client_strips_authorization_when_auth_is_none() {
139        let provider = GeminiProvider::new(Some("real-key".to_string()))
140            .with_connection(ProviderConnectionConfig { auth_mode: ProviderAuthMode::None, ..Default::default() });
141        let api_key = provider.get_api_key().unwrap();
142        let client = provider.build_openai_client(&api_key);
143        assert!(!client.config().headers().contains_key(AUTHORIZATION));
144    }
145}