Skip to main content

vtcode_llm/
provider_base.rs

1//! Base traits and utilities for LLM providers
2//!
3//! This module provides shared functionality to eliminate duplicate code
4//! across the 15+ LLM provider implementations.
5
6use anyhow::{Context, Result};
7use async_trait::async_trait;
8use futures::StreamExt;
9use reqwest::Client as HttpClient;
10use serde_json::Value;
11use std::time::Duration;
12
13use crate::provider::{LLMError, LLMRequest, LLMStreamEvent};
14use vtcode_config::TimeoutsConfig;
15
16/// Default timeout configurations
17const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_secs(120);
18const DEFAULT_STREAM_TIMEOUT: Duration = Duration::from_secs(300);
19
20/// Base configuration shared by all providers
21#[derive(Debug, Clone)]
22pub struct BaseProviderConfig {
23    api_key: String,
24    base_url: String,
25    model: String,
26    http_client: HttpClient,
27    prompt_cache_enabled: bool,
28    request_timeout: Duration,
29    stream_timeout: Duration,
30}
31
32impl BaseProviderConfig {
33    /// Create base configuration from common parameters
34    fn from_options(
35        api_key: Option<String>,
36        model: Option<String>,
37        base_url: Option<String>,
38        default_model: &'static str,
39        default_url: &'static str,
40        env_var: &'static str,
41        timeouts: Option<TimeoutsConfig>,
42    ) -> Result<Self> {
43        let api_key_value = api_key.unwrap_or_default();
44        let model_value = model.unwrap_or_else(|| default_model.to_string());
45        let base_url_value = Self::resolve_base_url(base_url, default_url, env_var)?;
46
47        let timeout_config = timeouts.unwrap_or_default();
48        let http_timeout = timeout_config
49            .ceiling_duration(timeout_config.streaming_ceiling_seconds)
50            .unwrap_or(DEFAULT_REQUEST_TIMEOUT);
51        let http_client = HttpClient::builder()
52            .timeout(http_timeout)
53            .build()
54            .context("Failed to build HTTP client")?;
55
56        Ok(Self {
57            api_key: api_key_value,
58            base_url: base_url_value,
59            model: model_value,
60            http_client,
61            prompt_cache_enabled: false,
62            request_timeout: http_timeout,
63            stream_timeout: timeout_config
64                .ceiling_duration(timeout_config.streaming_ceiling_seconds)
65                .unwrap_or(DEFAULT_STREAM_TIMEOUT),
66        })
67    }
68
69    /// Resolve base URL with environment variable fallback
70    fn resolve_base_url(base_url: Option<String>, default_url: &'static str, env_var: &'static str) -> Result<String> {
71        if let Some(url) = base_url {
72            Ok(url.trim().to_string())
73        } else if let Ok(env_val) = std::env::var(env_var) {
74            Ok(env_val.trim().to_string())
75        } else {
76            Ok(default_url.to_string())
77        }
78    }
79
80    /// Validate that required API key is present
81    pub fn validate_api_key(&self) -> Result<()> {
82        if self.api_key.is_empty() {
83            anyhow::bail!("API key is required")
84        }
85        Ok(())
86    }
87}
88
89/// Trait for providers that support standard OpenAI-compatible APIs
90#[async_trait]
91trait OpenAICompatibleProvider: Send + Sync {
92    fn provider_name(&self) -> &'static str;
93    fn supports_prompt_caching(&self) -> bool;
94
95    /// Parse request from OpenAI format
96    fn parse_openai_request(&self, value: &Value, default_model: &str) -> Option<LLMRequest> {
97        crate::utils::parse_chat_request_openai_format(value, default_model)
98    }
99
100    /// Serialize messages to OpenAI format
101    fn serialize_openai_messages(&self, request: &LLMRequest) -> Value {
102        use crate::providers::common::serialize_messages_openai_format;
103        match serialize_messages_openai_format(request, self.provider_name()) {
104            Ok(messages) => serde_json::json!({ "messages": messages }),
105            Err(_) => serde_json::json!({ "messages": [] }),
106        }
107    }
108
109    /// Parse response from OpenAI format
110    fn parse_openai_response(
111        &self,
112        response: Value,
113        model: String,
114        include_cache: bool,
115    ) -> Result<crate::provider::LLMResponse> {
116        crate::utils::parse_response_openai_format(response, self.provider_name(), model, include_cache, None)
117    }
118}
119
120/// Shared error handling utilities
121struct ErrorHandler {
122    _provider_name: &'static str,
123}
124
125impl ErrorHandler {
126    fn new(provider_name: &'static str) -> Self {
127        Self { _provider_name: provider_name }
128    }
129
130    /// Handle HTTP errors consistently across providers
131    fn handle_http_error(&self, status: reqwest::StatusCode, error_text: &str) -> LLMError {
132        use reqwest::StatusCode;
133
134        let error_message = match status {
135            StatusCode::UNAUTHORIZED => "Authentication failed: Invalid API key".to_string(),
136            StatusCode::TOO_MANY_REQUESTS => "Rate limit exceeded".to_string(),
137            StatusCode::BAD_REQUEST => format!("Bad request: {}", error_text.trim()),
138            s if s.as_u16() == 402 => "Insufficient balance".to_string(),
139            _ => format!("HTTP {}: {}", status, error_text.trim()),
140        };
141
142        let formatted_error = crate::error_display::format_llm_error(self._provider_name, &error_message);
143
144        // Handle different error types based on status code
145        if status == StatusCode::TOO_MANY_REQUESTS {
146            LLMError::RateLimit { metadata: None }
147        } else {
148            LLMError::Provider { message: formatted_error, metadata: None }
149        }
150    }
151
152    /// Handle request validation errors
153    pub fn validate_request(&self, request: &LLMRequest) -> Result<()> {
154        if request.messages.is_empty() {
155            anyhow::bail!("Request must contain at least one message")
156        }
157
158        if request.model.is_empty() {
159            anyhow::bail!("Request must specify a model")
160        }
161
162        // Check if model is supported (this would need to be customized per provider)
163        if !self.is_model_supported(&request.model) {
164            anyhow::bail!("Unsupported model: {}", request.model)
165        }
166
167        Ok(())
168    }
169
170    /// Check if model is supported (default implementation, override as needed)
171    fn is_model_supported(&self, model: &str) -> bool {
172        // Default implementation assumes all models are supported
173        // Individual providers should override this with their specific model lists
174        !model.is_empty()
175    }
176}
177
178/// Shared streaming utilities
179pub struct StreamProcessor {
180    provider_name: &'static str,
181    supports_reasoning: bool,
182}
183
184impl StreamProcessor {
185    pub fn new(provider_name: &'static str, supports_reasoning: bool) -> Self {
186        Self { provider_name, supports_reasoning }
187    }
188
189    /// Process SSE stream chunk consistently
190    pub fn process_stream_chunk(&self, chunk: &str) -> Vec<LLMStreamEvent> {
191        let mut events = Vec::new();
192
193        for line in chunk.lines() {
194            let line = line.trim();
195            if line.is_empty() {
196                continue;
197            }
198
199            if let Some(data) = line.strip_prefix("data: ") {
200                if data == "[DONE]" {
201                    // Stream completion indicated by DONE marker
202                    continue;
203                }
204
205                match serde_json::from_str::<Value>(data) {
206                    Ok(json) => {
207                        if let Some(event) = self.parse_stream_event(json) {
208                            events.push(event);
209                        }
210                    }
211                    Err(_) => {
212                        // Skip invalid JSON
213                        continue;
214                    }
215                }
216            }
217        }
218
219        events
220    }
221
222    /// Parse individual stream event (override for provider-specific logic)
223    fn parse_stream_event(&self, json: Value) -> Option<LLMStreamEvent> {
224        // Default implementation for OpenAI-compatible providers
225        crate::utils::parse_stream_event_openai_format(json, self.provider_name)
226    }
227
228    /// Extract reasoning content if supported
229    pub fn extract_reasoning(&self, content: &str) -> (Vec<String>, Option<String>) {
230        if !self.supports_reasoning {
231            return (Vec::new(), None);
232        }
233
234        // Default implementation - providers can override
235        crate::utils::extract_reasoning_content(content)
236    }
237}
238
239/// Unified authentication header handling
240pub struct AuthHandler {
241    auth_type: AuthType,
242    api_key: String,
243}
244
245#[derive(Debug, Clone, Copy)]
246pub enum AuthType {
247    BearerToken,
248    ApiKeyHeader(&'static str),
249    QueryParam(&'static str),
250}
251
252impl AuthHandler {
253    pub fn new(auth_type: AuthType, api_key: String) -> Self {
254        Self { auth_type, api_key }
255    }
256
257    /// Apply authentication to request builder
258    pub fn apply_auth(&self, builder: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
259        match self.auth_type {
260            AuthType::BearerToken => builder.bearer_auth(&self.api_key),
261            AuthType::ApiKeyHeader(header_name) => builder.header(header_name, &self.api_key),
262            AuthType::QueryParam(param_name) => builder.query(&[(param_name, &self.api_key)]),
263        }
264    }
265}
266
267/// Shared request/response processing utilities
268pub struct RequestProcessor {
269    provider_name: &'static str,
270}
271
272impl RequestProcessor {
273    pub fn new(provider_name: &'static str) -> Self {
274        Self { provider_name }
275    }
276
277    /// Build HTTP request with consistent error handling
278    pub async fn build_request(
279        &self,
280        client: &HttpClient,
281        method: reqwest::Method,
282        url: String,
283        auth: Option<&AuthHandler>,
284        body: Option<Value>,
285    ) -> Result<reqwest::RequestBuilder> {
286        let mut builder = client.request(method, &url);
287
288        if let Some(auth_handler) = auth {
289            builder = auth_handler.apply_auth(builder);
290        }
291
292        builder = builder
293            .header("Content-Type", "application/json")
294            .header("User-Agent", "VT Code/1.0");
295
296        if let Some(body_value) = body {
297            builder = builder.json(&body_value);
298        }
299
300        Ok(builder)
301    }
302
303    /// Handle response with consistent error processing
304    pub async fn handle_response(&self, response: reqwest::Response) -> Result<Value> {
305        let status = response.status();
306
307        if !status.is_success() {
308            let error_text = crate::providers::common::read_provider_error_body(response).await;
309            let error_handler = ErrorHandler::new(self.provider_name);
310            return Err(error_handler.handle_http_error(status, &error_text).into());
311        }
312
313        let response_text = response.text().await.context("Failed to read response body")?;
314
315        serde_json::from_str(&response_text).context("Failed to parse JSON response")
316    }
317
318    /// Handle streaming response
319    pub async fn handle_stream_response(
320        &self,
321        response: reqwest::Response,
322    ) -> Result<impl futures::Stream<Item = Result<String>>> {
323        let status = response.status();
324
325        if !status.is_success() {
326            let error_text = crate::providers::common::read_provider_error_body(response).await;
327            let error_handler = ErrorHandler::new(self.provider_name);
328            return Err(error_handler.handle_http_error(status, &error_text).into());
329        }
330
331        Ok(response.bytes_stream().map(|result| {
332            result
333                .map(|bytes| String::from_utf8_lossy(&bytes).into_owned())
334                .map_err(|e| anyhow::anyhow!("Stream error: {e}"))
335        }))
336    }
337}
338
339#[cfg(test)]
340mod tests {
341    use super::*;
342
343    /// Test-only model resolution utilities. The production `ModelResolver`
344    /// lives in [`crate::model_resolver`] — this private duplicate exists only
345    /// to exercise fallback/validation logic in unit tests below.
346    struct ModelResolver {
347        provider_name: &'static str,
348        default_model: &'static str,
349        supported_models: &'static [&'static str],
350    }
351
352    impl ModelResolver {
353        fn new(
354            provider_name: &'static str,
355            default_model: &'static str,
356            supported_models: &'static [&'static str],
357        ) -> Self {
358            Self { provider_name, default_model, supported_models }
359        }
360
361        /// Resolve model with fallback to default
362        fn resolve_model(&self, model: Option<String>) -> String {
363            model.unwrap_or_else(|| self.default_model.to_string())
364        }
365
366        /// Validate model is supported
367        fn validate_model(&self, model: &str) -> Result<()> {
368            if self.supported_models.is_empty() {
369                // If no specific supported models listed, accept any non-empty model
370                if model.is_empty() {
371                    anyhow::bail!("Model cannot be empty")
372                }
373                return Ok(());
374            }
375
376            if !self.supported_models.contains(&model) {
377                anyhow::bail!("Unsupported model: {}. Supported models: {:?}", model, self.supported_models)
378            }
379
380            Ok(())
381        }
382    }
383
384    #[test]
385    fn test_base_provider_config() {
386        let config = BaseProviderConfig::from_options(
387            Some("test_key".to_string()),
388            Some("test_model".to_string()),
389            None,
390            "default_model",
391            "https://api.example.com",
392            "TEST_API_KEY",
393            None,
394        )
395        .unwrap();
396
397        assert_eq!(config.api_key, "test_key");
398        assert_eq!(config.model, "test_model");
399        assert_eq!(config.base_url, "https://api.example.com");
400    }
401
402    #[test]
403    fn test_error_handler() {
404        let handler = ErrorHandler::new("test_provider");
405
406        let unauthorized = handler.handle_http_error(reqwest::StatusCode::UNAUTHORIZED, "Invalid API key");
407        let rate_limited = handler.handle_http_error(reqwest::StatusCode::TOO_MANY_REQUESTS, "");
408
409        assert!(matches!(unauthorized, LLMError::Provider { message: _, metadata: _ }));
410        assert!(matches!(rate_limited, LLMError::RateLimit { metadata: _ }));
411    }
412
413    #[test]
414    fn test_model_resolver() {
415        let resolver = ModelResolver::new("test_provider", "default-model", &["model1", "model2"]);
416
417        assert_eq!(resolver.resolve_model(None), "default-model");
418        assert_eq!(resolver.resolve_model(Some("custom".to_string())), "custom");
419
420        resolver.validate_model("model1").unwrap();
421        assert!(resolver.validate_model("unsupported").is_err());
422    }
423}