1use 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
16const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_secs(120);
18const DEFAULT_STREAM_TIMEOUT: Duration = Duration::from_secs(300);
19
20#[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 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 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 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#[async_trait]
91trait OpenAICompatibleProvider: Send + Sync {
92 fn provider_name(&self) -> &'static str;
93 fn supports_prompt_caching(&self) -> bool;
94
95 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 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 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
120struct 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 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 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 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 if !self.is_model_supported(&request.model) {
164 anyhow::bail!("Unsupported model: {}", request.model)
165 }
166
167 Ok(())
168 }
169
170 fn is_model_supported(&self, model: &str) -> bool {
172 !model.is_empty()
175 }
176}
177
178pub 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 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 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 continue;
214 }
215 }
216 }
217 }
218
219 events
220 }
221
222 fn parse_stream_event(&self, json: Value) -> Option<LLMStreamEvent> {
224 crate::utils::parse_stream_event_openai_format(json, self.provider_name)
226 }
227
228 pub fn extract_reasoning(&self, content: &str) -> (Vec<String>, Option<String>) {
230 if !self.supports_reasoning {
231 return (Vec::new(), None);
232 }
233
234 crate::utils::extract_reasoning_content(content)
236 }
237}
238
239pub 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 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
267pub 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 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 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 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 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 fn resolve_model(&self, model: Option<String>) -> String {
363 model.unwrap_or_else(|| self.default_model.to_string())
364 }
365
366 fn validate_model(&self, model: &str) -> Result<()> {
368 if self.supported_models.is_empty() {
369 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}