Skip to main content

vtcode_core/llm/
provider_builder.rs

1use crate::config::TimeoutsConfig;
2use crate::config::core::PromptCachingConfig;
3use crate::llm::provider::{LLMError, LLMProvider};
4use std::marker::PhantomData;
5
6/// Generic provider builder to eliminate duplicate provider creation patterns
7pub struct ProviderBuilder<T> {
8    api_key: Option<String>,
9    model: Option<String>,
10    base_url: Option<String>,
11    prompt_cache: Option<PromptCachingConfig>,
12    timeouts: Option<TimeoutsConfig>,
13    _phantom: PhantomData<T>,
14}
15
16impl<T> Default for ProviderBuilder<T> {
17    fn default() -> Self {
18        Self {
19            api_key: None,
20            model: None,
21            base_url: None,
22            prompt_cache: None,
23            timeouts: None,
24            _phantom: PhantomData,
25        }
26    }
27}
28
29impl<T> ProviderBuilder<T>
30where
31    T: ProviderConfig,
32{
33    /// Create a new builder with all fields unset.
34    pub fn new() -> Self {
35        Self::default()
36    }
37
38    /// Set the API key for provider authentication.
39    pub fn api_key(mut self, api_key: String) -> Self {
40        self.api_key = Some(api_key);
41        self
42    }
43
44    /// Set the model identifier to use.
45    pub fn model(mut self, model: String) -> Self {
46        self.model = Some(model);
47        self
48    }
49
50    /// Build the provider, returning an error if creation fails.
51    pub fn try_build(self) -> Result<Box<dyn LLMProvider>, LLMError> {
52        crate::llm::provider_config::create_provider_unified(
53            T::PROVIDER_KEY,
54            self.api_key,
55            self.model,
56            self.base_url,
57            self.prompt_cache,
58            self.timeouts,
59        )
60    }
61
62    /// Build the provider, panicking if creation fails.
63    ///
64    /// This method is intended for use in contexts where provider creation
65    /// should never fail (e.g., after configuration validation). If failure
66    /// is possible, use [`Self::try_build`] instead.
67    ///
68    /// # Panics
69    ///
70    /// Panics if the provider cannot be created. This indicates a bug in the
71    /// configuration validation or provider creation logic.
72    pub fn build(self) -> Box<dyn LLMProvider> {
73        match self.try_build() {
74            Ok(provider) => provider,
75            Err(error) => panic!(
76                "provider builder invariant violated for `{}`: {}. \
77                 This indicates a bug in configuration validation. \
78                 Use try_build() if failure is expected.",
79                T::PROVIDER_KEY,
80                error
81            ),
82        }
83    }
84}
85
86/// Trait for provider-specific configuration and creation
87pub trait ProviderConfig {
88    /// Registry key used to look up this provider in the factory.
89    const PROVIDER_KEY: &'static str;
90    /// Human-readable display name for this provider.
91    const DISPLAY_NAME: &'static str;
92    /// Default model identifier when none is specified.
93    const DEFAULT_MODEL: &'static str;
94    /// Base URL for the provider's API endpoint.
95    const API_BASE_URL: &'static str;
96    /// Optional environment variable that overrides the base URL.
97    const BASE_URL_ENV_VAR: Option<&'static str>;
98
99    /// Construct a boxed [`LLMProvider`] from the given configuration.
100    fn create_provider(
101        api_key: String,
102        model: String,
103        base_url: String,
104        prompt_cache_enabled: bool,
105        prompt_cache_settings: Self::PromptCacheSettings,
106        timeouts: TimeoutsConfig,
107    ) -> Box<dyn LLMProvider>
108    where
109        Self::PromptCacheSettings: Send + Sync + 'static,
110    {
111        let _ = prompt_cache_settings;
112        let prompt_cache = prompt_cache_enabled.then(|| PromptCachingConfig { enabled: true, ..Default::default() });
113
114        match crate::llm::provider_config::create_provider_unified(
115            Self::PROVIDER_KEY,
116            (!api_key.trim().is_empty()).then_some(api_key),
117            (!model.trim().is_empty()).then_some(model),
118            (!base_url.trim().is_empty()).then_some(base_url),
119            prompt_cache,
120            Some(timeouts),
121        ) {
122            Ok(provider) => provider,
123            Err(error) => {
124                panic!("provider config invariant violated for `{}`: {}", Self::PROVIDER_KEY, error)
125            }
126        }
127    }
128
129    /// Provider-specific prompt cache configuration type.
130    type PromptCacheSettings: Clone + Default + Send + Sync + 'static;
131}
132
133/// HTTP client pool to avoid creating new clients for each provider
134mod http_client_pool {
135    use crate::config::TimeoutsConfig;
136    use hashbrown::HashMap;
137    use once_cell::sync::Lazy;
138    use reqwest::Client as HttpClient;
139    use std::sync::{Arc, RwLock};
140    use std::time::Duration;
141
142    type HttpClientPool = Arc<RwLock<HashMap<String, Arc<HttpClient>>>>;
143
144    static CLIENT_POOL: Lazy<HttpClientPool> = Lazy::new(|| {
145        let mut pool = HashMap::new();
146
147        // Default client
148        pool.insert("default".to_string(), Arc::new(HttpClient::new()));
149
150        // Timeout-configured clients
151        pool.insert(
152            "timeout_30s".to_string(),
153            Arc::new(
154                HttpClient::builder()
155                    .timeout(Duration::from_secs(30))
156                    .build()
157                    .unwrap_or_else(|error| {
158                        tracing::warn!(
159                            error = %error,
160                            "Failed to build 30s timeout HTTP client; falling back to default client"
161                        );
162                        HttpClient::new()
163                    }),
164            ),
165        );
166
167        pool.insert(
168            "timeout_120s".to_string(),
169            Arc::new(
170                HttpClient::builder()
171                    .timeout(Duration::from_secs(120))
172                    .build()
173                    .unwrap_or_else(|error| {
174                        tracing::warn!(
175                            error = %error,
176                            "Failed to build 120s timeout HTTP client; falling back to default client"
177                        );
178                        HttpClient::new()
179                    }),
180            ),
181        );
182
183        Arc::new(RwLock::new(pool))
184    });
185
186    /// Retrieve a pooled HTTP client by key, falling back to the default client.
187    pub fn get_http_client(key: &str) -> Arc<HttpClient> {
188        let pool_guard = CLIENT_POOL.read();
189        let pool = match pool_guard {
190            Ok(guard) => guard,
191            Err(poisoned) => {
192                tracing::warn!("HTTP client pool poisoned; continuing with recovered state");
193                poisoned.into_inner()
194            }
195        };
196
197        if let Some(client) = pool.get(key).cloned() {
198            return client;
199        }
200
201        if let Some(default_client) = pool.get("default").cloned() {
202            return default_client;
203        }
204
205        tracing::warn!("HTTP client pool missing default client; constructing transient client");
206        Arc::new(HttpClient::new())
207    }
208
209    /// Select an appropriate pooled HTTP client based on the timeout ceiling.
210    pub fn get_http_client_for_timeouts(timeouts: &TimeoutsConfig) -> Arc<HttpClient> {
211        let key = if timeouts.default_ceiling_seconds >= 120 {
212            "timeout_120s"
213        } else if timeouts.default_ceiling_seconds >= 30 {
214            "timeout_30s"
215        } else {
216            "default"
217        };
218        get_http_client(key)
219    }
220}
221
222pub use http_client_pool::{get_http_client, get_http_client_for_timeouts};