Skip to main content

vtcode_llm/providers/gemini/wire/client/
mod.rs

1pub mod config;
2pub mod retry;
3
4pub use config::ClientConfig;
5pub use retry::RetryConfig;
6
7use super::models::{GenerateContentRequest, GenerateContentResponse};
8use super::streaming::{StreamingError, StreamingMetrics, StreamingProcessor, StreamingResponse};
9use anyhow::{Context, Result};
10use reqwest::Client as ReqwestClient;
11use reqwest::StatusCode;
12use std::time::Instant;
13use tracing::warn;
14use vtcode_commons::llm::{LLMError, LLMErrorMetadata};
15
16#[derive(Clone)]
17pub struct Client {
18    api_key: String,
19    model: String,
20    http: ReqwestClient,
21    config: ClientConfig,
22    retry_config: RetryConfig,
23    metrics: StreamingMetrics,
24}
25
26impl Client {
27    pub fn new(api_key: String, model: String) -> Self {
28        Self::with_config(api_key, model, ClientConfig::default())
29    }
30
31    /// Create a client with custom configuration
32    fn with_config(api_key: String, model: String, config: ClientConfig) -> Self {
33        let http_client = ReqwestClient::builder()
34            .pool_max_idle_per_host(config.pool_max_idle_per_host)
35            .pool_idle_timeout(config.pool_idle_timeout)
36            .tcp_keepalive(config.tcp_keepalive)
37            .timeout(config.request_timeout)
38            .connect_timeout(config.connect_timeout)
39            .user_agent(&config.user_agent)
40            .build()
41            .unwrap_or_else(|error| {
42                warn!(error = %error, "Failed to build Gemini HTTP client; using default client");
43                ReqwestClient::new()
44            });
45
46        Self {
47            api_key,
48            model,
49            http: http_client,
50            config,
51            retry_config: RetryConfig::default(),
52            metrics: StreamingMetrics::default(),
53        }
54    }
55
56    /// Get current client configuration
57    pub fn config(&self) -> &ClientConfig {
58        &self.config
59    }
60
61    /// Set retry configuration
62    pub fn with_retry_config(mut self, retry_config: RetryConfig) -> Self {
63        self.retry_config = retry_config;
64        self
65    }
66
67    /// Get current retry configuration
68    pub fn retry_config(&self) -> &RetryConfig {
69        &self.retry_config
70    }
71
72    /// Get streaming metrics
73    pub fn metrics(&self) -> &StreamingMetrics {
74        &self.metrics
75    }
76
77    /// Reset streaming metrics
78    pub fn reset_metrics(&mut self) {
79        self.metrics = StreamingMetrics::default();
80    }
81
82    /// Classify error to determine if it's retryable
83    fn classify_error(&self, error: &anyhow::Error) -> StreamingError {
84        let decision = vtcode_commons::retry::RetryPolicy::default().classify_anyhow(error);
85        let message = error.to_string();
86
87        match decision.category {
88            vtcode_commons::ErrorCategory::RateLimit => StreamingError::ApiError {
89                status_code: 429,
90                message,
91                is_retryable: decision.retryable,
92            },
93            vtcode_commons::ErrorCategory::ServiceUnavailable => StreamingError::ApiError {
94                status_code: 503,
95                message,
96                is_retryable: decision.retryable,
97            },
98            _ => StreamingError::NetworkError { message, is_retryable: decision.retryable },
99        }
100    }
101
102    fn classify_api_error(&self, status: StatusCode, message: String) -> StreamingError {
103        let decision = vtcode_commons::retry::RetryPolicy::default().classify_status(status.as_u16());
104
105        StreamingError::ApiError {
106            status_code: status.as_u16(),
107            message,
108            is_retryable: decision.retryable,
109        }
110    }
111
112    /// Generate content with the Gemini API
113    pub async fn generate(&mut self, request: &GenerateContentRequest) -> Result<GenerateContentResponse> {
114        let start_time = Instant::now();
115
116        let url = format!("https://generativelanguage.googleapis.com/v1beta/models/{}:generateContent", self.model);
117
118        let response = self
119            .http
120            .post(&url)
121            .header("x-api-key", &self.api_key)
122            .json(request)
123            .send()
124            .await
125            .context("Failed to send request")?;
126
127        if !response.status().is_success() {
128            let status = response.status();
129            let error_text = crate::providers::common::read_provider_error_body(response).await;
130            let error = llm_error_for_status(status, &error_text);
131            return Err(anyhow::Error::new(error));
132        }
133
134        let response_data: GenerateContentResponse = response.json().await.context("Failed to parse response")?;
135
136        self.metrics.total_requests += 1;
137        self.metrics.total_response_time += start_time.elapsed();
138
139        Ok(response_data)
140    }
141
142    /// Generate content with the Gemini API using streaming
143    pub async fn generate_stream<F>(
144        &mut self,
145        request: &GenerateContentRequest,
146        on_chunk: F,
147    ) -> Result<StreamingResponse, StreamingError>
148    where
149        F: FnMut(&str) -> Result<(), StreamingError>,
150    {
151        let start_time = Instant::now();
152
153        let url =
154            format!("https://generativelanguage.googleapis.com/v1beta/models/{}:streamGenerateContent", self.model);
155
156        let response = self
157            .http
158            .post(&url)
159            .header("x-api-key", &self.api_key)
160            .json(request)
161            .send()
162            .await
163            .map_err(|e| {
164                let error = anyhow::Error::new(e);
165                self.classify_error(&error)
166            })?;
167
168        if !response.status().is_success() {
169            let status = response.status();
170            let error_text = crate::providers::common::read_provider_error_body(response).await;
171            return Err(self.classify_api_error(status, error_text));
172        }
173
174        // Process the streaming response
175        let mut processor = StreamingProcessor::new();
176        let result = processor.process_stream(response, on_chunk).await;
177
178        self.metrics.total_requests += 1;
179        self.metrics.total_response_time += start_time.elapsed();
180
181        result
182    }
183}
184
185fn llm_error_for_status(status: StatusCode, message: &str) -> LLMError {
186    let metadata = Some(LLMErrorMetadata::new(
187        "gemini",
188        Some(status.as_u16()),
189        None,
190        None,
191        None,
192        None,
193        Some(message.to_string()),
194    ));
195
196    match status {
197        StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => {
198            LLMError::Authentication { message: message.to_string(), metadata }
199        }
200        StatusCode::TOO_MANY_REQUESTS => LLMError::RateLimit { metadata },
201        StatusCode::BAD_REQUEST | StatusCode::UNPROCESSABLE_ENTITY => {
202            LLMError::InvalidRequest { message: message.to_string(), metadata }
203        }
204        _ => LLMError::Provider { message: message.to_string(), metadata },
205    }
206}