Skip to main content

llm/providers/gemini/
provider.rs

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