Skip to main content

vtcode_llm/providers/
base.rs

1//! Base trait and common implementations for LLM providers
2//!
3//! This module provides a unified foundation for all LLM providers to eliminate
4//! code duplication across provider implementations.
5
6use crate::provider::{LLMError, LLMRequest, LLMResponse, Message, ToolDefinition};
7use crate::providers::common::read_provider_error_body;
8use async_trait::async_trait;
9use hashbrown::HashMap;
10use reqwest::{Client as HttpClient, StatusCode};
11use serde_json::Value;
12use std::sync::{Arc, LazyLock, Mutex};
13use std::time::Duration;
14use tokio::sync::{OwnedSemaphorePermit, Semaphore};
15use tokio::time::{sleep, timeout};
16use vtcode_commons::sanitizer::sanitize_provider_diagnostic;
17
18const DEFAULT_MAX_INFLIGHT_PER_MODEL: usize = 4;
19const RATE_LIMIT_ACQUIRE_TIMEOUT: Duration = Duration::from_secs(10);
20
21static MODEL_LIMITERS: LazyLock<Mutex<HashMap<String, Arc<Semaphore>>>> = LazyLock::new(|| Mutex::new(HashMap::new()));
22
23/// Base configuration shared by all providers
24#[derive(Debug, Clone)]
25pub struct ProviderConfig {
26    api_key: String,
27    base_url: String,
28    model: String,
29    timeout: Duration,
30    max_retries: u32,
31}
32
33impl ProviderConfig {
34    /// Create provider config with sensible defaults
35    pub fn new(api_key: String, base_url: String, model: String) -> Self {
36        Self {
37            api_key,
38            base_url,
39            model,
40            timeout: Duration::from_secs(120),
41            max_retries: 3,
42        }
43    }
44
45    /// Build HTTP client with provider-specific configuration
46    fn build_http_client(&self) -> Result<HttpClient, LLMError> {
47        use crate::http_client::HttpClientFactory;
48        Ok(HttpClientFactory::with_timeouts(self.timeout, Duration::from_secs(30)))
49    }
50}
51
52/// Common HTTP error handling for all providers
53fn handle_http_error(status: StatusCode, error_text: &str, _model: &str) -> LLMError {
54    let error_text = sanitize_provider_diagnostic(error_text.as_bytes());
55    match status {
56        StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => LLMError::Authentication {
57            message: format!("Authentication failed ({status}): {error_text}"),
58            metadata: None,
59        },
60        StatusCode::TOO_MANY_REQUESTS => LLMError::RateLimit { metadata: None },
61        StatusCode::REQUEST_TIMEOUT => LLMError::Network {
62            message: format!("Request timeout ({status}): {error_text}"),
63            metadata: None,
64        },
65        _ if status.is_server_error() => LLMError::Provider {
66            message: format!("Server error ({status}): {error_text}"),
67            metadata: None,
68        },
69        _ => LLMError::Network {
70            message: format!("HTTP error ({status}): {error_text}"),
71            metadata: None,
72        },
73    }
74}
75
76/// Check if error indicates model not found (common across providers)
77pub fn is_model_not_found(status: StatusCode, error_text: &str) -> bool {
78    status == StatusCode::NOT_FOUND
79        || error_text.contains("model_not_found")
80        || (error_text.to_ascii_lowercase().contains("does not exist")
81            && error_text.to_ascii_lowercase().contains("model"))
82}
83
84/// Common request building utilities
85pub mod request_builder {
86    use super::*;
87
88    /// Build standard headers for API requests
89    pub fn build_headers(api_key: &str, provider_headers: Option<Vec<(&str, &str)>>) -> reqwest::header::HeaderMap {
90        use reqwest::header::{AUTHORIZATION, CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue};
91
92        let mut headers = HeaderMap::new();
93        headers.insert(CONTENT_TYPE, HeaderValue::from_static("application/json"));
94
95        // Default authorization header (can be overridden by providers)
96        if let Ok(auth_value) = HeaderValue::from_str(&format!("Bearer {api_key}")) {
97            headers.insert(AUTHORIZATION, auth_value);
98        }
99
100        // Add provider-specific headers
101        if let Some(custom_headers) = provider_headers {
102            for (key, value) in custom_headers {
103                if let (Ok(name), Ok(val)) = (HeaderName::from_bytes(key.as_bytes()), HeaderValue::from_str(value)) {
104                    headers.insert(name, val);
105                }
106            }
107        }
108
109        headers
110    }
111
112    /// Convert tools to OpenAI-compatible format (used by many providers)
113    pub fn serialize_tools_openai(tools: &[ToolDefinition]) -> Option<Vec<Value>> {
114        if tools.is_empty() {
115            return None;
116        }
117        Some(tools.iter().map(|tool| serde_json::json!(tool)).collect())
118    }
119
120    /// Build standard request body structure
121    pub fn build_request_body(
122        messages: &[Message],
123        model: &str,
124        max_tokens: Option<u32>,
125        temperature: Option<f32>,
126        tools: Option<Vec<Value>>,
127        stream: bool,
128        reasoning_effort: Option<String>,
129    ) -> Value {
130        let mut body = serde_json::json!({
131            "model": model,
132            "messages": messages.iter().map(|msg| serde_json::json!({
133                "role": msg.role.to_string().to_lowercase(),
134                "content": msg.content,
135            })).collect::<Vec<_>>(),
136        });
137
138        if let Some(max_tokens_val) = max_tokens {
139            body["max_tokens"] = serde_json::json!(max_tokens_val);
140        }
141
142        if let Some(temp) = temperature {
143            body["temperature"] = serde_json::json!(crate::providers::common::sampling_param_f64(temp));
144        }
145
146        if let Some(val) = tools {
147            body["tools"] = serde_json::json!(val);
148        }
149
150        if let Some(effort) = reasoning_effort {
151            body["reasoning_effort"] = serde_json::json!(effort);
152        }
153
154        if stream {
155            body["stream"] = serde_json::json!(true);
156        }
157
158        body
159    }
160}
161
162/// Base provider trait with common functionality
163#[async_trait]
164trait BaseProvider: Send + Sync {
165    /// Get provider configuration
166    fn config(&self) -> &ProviderConfig;
167
168    /// Build HTTP request for the provider
169    fn build_request(&self, request: &LLMRequest) -> Result<reqwest::Request, LLMError>;
170
171    /// Parse response from the provider
172    fn parse_response(&self, response: Value) -> Result<LLMResponse, LLMError>;
173
174    /// Execute LLM request with common error handling and retry logic
175    async fn execute_request(&self, request: LLMRequest) -> Result<LLMResponse, LLMError> {
176        let _permit = acquire_model_permit(&self.config().model).await?;
177        let client = self.config().build_http_client()?;
178        let max_retries = self.config().max_retries;
179
180        let mut last_error = None;
181
182        for attempt in 0..=max_retries {
183            match self.build_request(&request) {
184                Ok(http_request) => {
185                    match client.execute(http_request).await {
186                        Ok(response) => {
187                            let status = response.status();
188
189                            let response_text = if status.is_success() {
190                                response.text().await
191                            } else {
192                                Ok(read_provider_error_body(response).await)
193                            };
194
195                            match response_text {
196                                Ok(text) => {
197                                    // Try to parse as JSON first
198                                    match serde_json::from_str::<Value>(&text) {
199                                        Ok(json_value) => {
200                                            // Check for provider-specific error format
201                                            if let Some(error_obj) = json_value.get("error") {
202                                                let error_text = error_obj.to_string();
203                                                if attempt < max_retries && should_retry_status(status) {
204                                                    sleep(backoff_duration(attempt)).await;
205                                                    last_error = Some(handle_http_error(
206                                                        status,
207                                                        &error_text,
208                                                        &self.config().model,
209                                                    ));
210                                                    continue;
211                                                }
212                                                return Err(handle_http_error(
213                                                    status,
214                                                    &error_text,
215                                                    &self.config().model,
216                                                ));
217                                            }
218
219                                            // Success - parse response
220                                            return self.parse_response(json_value);
221                                        }
222                                        Err(_) => {
223                                            // Not JSON - treat as error text
224                                            if attempt < max_retries && should_retry_status(status) {
225                                                sleep(backoff_duration(attempt)).await;
226                                                last_error =
227                                                    Some(handle_http_error(status, &text, &self.config().model));
228                                                continue;
229                                            }
230                                            return Err(handle_http_error(status, &text, &self.config().model));
231                                        }
232                                    }
233                                }
234                                Err(e) => {
235                                    let error = LLMError::Network {
236                                        message: format!("Failed to read response: {e}"),
237                                        metadata: None,
238                                    };
239                                    if attempt < max_retries {
240                                        last_error = Some(error);
241                                        continue;
242                                    }
243                                    return Err(error);
244                                }
245                            }
246                        }
247                        Err(e) => {
248                            let error = LLMError::Network {
249                                message: format!("Request failed: {e}"),
250                                metadata: None,
251                            };
252                            if attempt < max_retries {
253                                sleep(backoff_duration(attempt)).await;
254                                last_error = Some(error);
255                                continue;
256                            }
257                            return Err(error);
258                        }
259                    }
260                }
261                Err(e) => {
262                    if attempt < max_retries {
263                        last_error = Some(e);
264                        continue;
265                    }
266                    return Err(e);
267                }
268            }
269        }
270
271        // All retries exhausted
272        Err(last_error.unwrap_or_else(|| LLMError::Network {
273            message: "All retries exhausted".to_string(),
274            metadata: None,
275        }))
276    }
277}
278
279/// Determine if a status code should trigger a retry
280fn should_retry_status(status: StatusCode) -> bool {
281    matches!(
282        status,
283        StatusCode::REQUEST_TIMEOUT
284            | StatusCode::TOO_MANY_REQUESTS
285            | StatusCode::INTERNAL_SERVER_ERROR
286            | StatusCode::BAD_GATEWAY
287            | StatusCode::SERVICE_UNAVAILABLE
288            | StatusCode::GATEWAY_TIMEOUT
289    )
290}
291
292/// Exponential backoff with an upper bound to reduce provider hammering
293fn backoff_duration(attempt: u32) -> Duration {
294    let capped_attempt = attempt.min(5);
295    const BASE_MS: u64 = 200;
296    let backoff_ms = BASE_MS.saturating_mul(2_u64.saturating_pow(capped_attempt));
297    Duration::from_millis(backoff_ms.min(5_000))
298}
299
300fn limiter_for_model(model: &str) -> Arc<Semaphore> {
301    if let Ok(mut guard) = MODEL_LIMITERS.lock() {
302        guard
303            .entry(model.to_string())
304            .or_insert_with(|| Arc::new(Semaphore::new(DEFAULT_MAX_INFLIGHT_PER_MODEL)))
305            .clone()
306    } else {
307        Arc::new(Semaphore::new(DEFAULT_MAX_INFLIGHT_PER_MODEL))
308    }
309}
310
311async fn acquire_model_permit(model: &str) -> Result<OwnedSemaphorePermit, LLMError> {
312    let limiter = limiter_for_model(model);
313    match timeout(RATE_LIMIT_ACQUIRE_TIMEOUT, limiter.acquire_owned()).await {
314        Ok(Ok(permit)) => Ok(permit),
315        Ok(Err(_)) => Err(LLMError::RateLimit { metadata: None }),
316        Err(_) => Err(LLMError::RateLimit { metadata: None }),
317    }
318}