1use 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#[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 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 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
52fn 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
76pub 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
84pub mod request_builder {
86 use super::*;
87
88 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 if let Ok(auth_value) = HeaderValue::from_str(&format!("Bearer {api_key}")) {
97 headers.insert(AUTHORIZATION, auth_value);
98 }
99
100 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 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 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#[async_trait]
164trait BaseProvider: Send + Sync {
165 fn config(&self) -> &ProviderConfig;
167
168 fn build_request(&self, request: &LLMRequest) -> Result<reqwest::Request, LLMError>;
170
171 fn parse_response(&self, response: Value) -> Result<LLMResponse, LLMError>;
173
174 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 match serde_json::from_str::<Value>(&text) {
199 Ok(json_value) => {
200 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 return self.parse_response(json_value);
221 }
222 Err(_) => {
223 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 Err(last_error.unwrap_or_else(|| LLMError::Network {
273 message: "All retries exhausted".to_string(),
274 metadata: None,
275 }))
276 }
277}
278
279fn 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
292fn 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}