1use std::sync::Arc;
12use std::time::Duration;
13
14use async_trait::async_trait;
15
16use llm_trait::{
17 CallMode, Capabilities, ChatRequest, ChatResponse, ChatStream, HttpClient, LlmError,
18 LlmProvider, ProviderInfo, RawAdapter, RawRequest, ReqwestHttpClient,
19};
20
21use crate::model_registry::ModelProfile;
22
23#[derive(Debug, Clone)]
25pub struct ProviderConfig {
26 pub connect_timeout: Duration,
27 pub request_timeout: Duration,
28 pub max_retries: u32,
29 pub retry_delay: Duration,
30 pub client: Option<reqwest::Client>,
32}
33
34impl Default for ProviderConfig {
35 fn default() -> Self {
36 Self {
37 connect_timeout: Duration::from_secs(15),
38 request_timeout: Duration::from_secs(120),
39 max_retries: 3,
40 retry_delay: Duration::from_secs(1),
41 client: None,
42 }
43 }
44}
45
46pub struct GenericProvider {
51 adapter: Box<dyn RawAdapter>,
52 client: Arc<dyn HttpClient>,
53 config: ProviderConfig,
54}
55
56impl GenericProvider {
57 pub fn new(adapter: Box<dyn RawAdapter>) -> Self {
58 Self::with_config(adapter, ProviderConfig::default())
59 }
60
61 pub fn with_config(adapter: Box<dyn RawAdapter>, config: ProviderConfig) -> Self {
62 let reqwest_client = config.client.clone().unwrap_or_else(|| {
63 reqwest::Client::builder()
64 .connect_timeout(config.connect_timeout)
65 .read_timeout(config.request_timeout)
66 .build()
67 .expect("Failed to build HTTP client")
68 });
69
70 Self {
71 adapter,
72 client: Arc::new(ReqwestHttpClient::new(reqwest_client)),
73 config,
74 }
75 }
76
77 pub fn with_http_client(
83 adapter: Box<dyn RawAdapter>,
84 client: Arc<dyn HttpClient>,
85 config: ProviderConfig,
86 ) -> Self {
87 Self {
88 adapter,
89 client,
90 config,
91 }
92 }
93
94 pub fn adapter(&self) -> &dyn RawAdapter {
96 self.adapter.as_ref()
97 }
98
99 async fn execute_once(&self, request: RawRequest) -> Result<ChatResponse, LlmError> {
101 let response = self.send_request(&request).await?;
102 let body = response.text().await;
103 self.adapter.parse_response(body.as_bytes())
104 }
105
106 async fn execute_stream(&self, request: RawRequest) -> Result<ChatStream, LlmError> {
111 let mut last_err = None;
112
113 for attempt in 0..=self.config.max_retries {
114 if attempt > 0 {
115 let delay = self.calculate_backoff(attempt);
116 tokio::time::sleep(delay).await;
117 }
118
119 match self.client.send(&request).await {
120 Ok(response) => {
121 if response.is_success() {
122 return self
124 .adapter
125 .parse_sse_stream(self.client.as_ref(), request, response)
126 .await;
127 }
128
129 let status = response.status();
130 let body = response.text().await;
131
132 let is_retryable = status == 429 || status >= 500;
134 if !is_retryable || attempt == self.config.max_retries {
135 tracing::error!(
136 status = status,
137 url = %request.url,
138 error_body = %body,
139 "Stream HTTP error with full request context"
140 );
141 return Err(LlmError::api(status, body));
142 }
143
144 tracing::warn!(attempt, status, "Stream request failed, retrying");
145 last_err = Some(LlmError::api(status, body));
146 }
147 Err(e) => {
148 if attempt == self.config.max_retries {
149 return Err(e);
150 }
151 tracing::warn!(attempt, error = %e, "Stream request failed, retrying");
152 last_err = Some(e);
153 }
154 }
155 }
156
157 Err(last_err.unwrap_or_else(|| LlmError::llm("Stream request failed after retries")))
158 }
159
160 async fn send_request(
162 &self,
163 request: &RawRequest,
164 ) -> Result<llm_trait::HttpResponse, LlmError> {
165 let mut last_err = None;
166
167 for attempt in 0..=self.config.max_retries {
168 if attempt > 0 {
169 let delay = self.calculate_backoff(attempt);
170 tokio::time::sleep(delay).await;
171 }
172
173 match self.client.send(request).await {
174 Ok(response) => {
175 if response.is_success() {
176 return Ok(response);
177 }
178
179 let status = response.status();
180 let body = response.text().await;
181
182 let is_retryable = status == 429 || status >= 500;
184 if !is_retryable || attempt == self.config.max_retries {
185 tracing::error!(
186 status = status,
187 url = %request.url,
188 error_body = %body,
189 request_body = %serde_json::to_string(&request.body).unwrap_or_default(),
190 "HTTP error with full request context"
191 );
192 return Err(LlmError::api(status, body));
193 }
194
195 tracing::warn!(attempt, status, "Request failed, retrying");
196 last_err = Some(LlmError::api(status, body));
197 }
198 Err(e) => {
199 if attempt == self.config.max_retries {
200 return Err(e);
201 }
202 tracing::warn!(attempt, error = %e, "Request failed, retrying");
203 last_err = Some(e);
204 }
205 }
206 }
207
208 Err(last_err.unwrap_or_else(|| LlmError::llm("Request failed after retries")))
209 }
210
211 fn calculate_backoff(&self, attempt: u32) -> Duration {
212 let base = self.config.retry_delay.as_millis() as u64;
213 let exponential = base * 2u64.pow(attempt.saturating_sub(1));
214 let jitter = rand::random::<u64>() % 100;
215 Duration::from_millis((exponential + jitter).min(30_000))
216 }
217}
218
219#[async_trait]
220impl LlmProvider for GenericProvider {
221 async fn stream(&self, request: ChatRequest) -> Result<ChatStream, LlmError> {
222 let modes = self.adapter.supported_modes();
223 if !modes.contains(&CallMode::Stream) {
224 return Err(LlmError::llm("Streaming not supported by this adapter"));
225 }
226
227 let raw_request = self.adapter.build_request(&request, CallMode::Stream)?;
228 self.execute_stream(raw_request).await
229 }
230
231 async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, LlmError> {
232 let modes = self.adapter.supported_modes();
233
234 if modes.contains(&CallMode::Once) {
235 let raw_request = self.adapter.build_request(&request, CallMode::Once)?;
236 return self.execute_once(raw_request).await;
237 }
238
239 let stream = self.stream(request).await?;
241 stream.collect_response().await
242 }
243
244 fn capabilities(&self) -> Capabilities {
245 self.adapter.capabilities()
246 }
247
248 fn info(&self) -> ProviderInfo {
249 self.adapter.info()
250 }
251}
252
253pub struct ProfiledProvider {
258 inner: GenericProvider,
259 profile: ModelProfile,
260}
261
262impl ProfiledProvider {
263 pub fn new(inner: GenericProvider, profile: ModelProfile) -> Self {
264 Self { inner, profile }
265 }
266}
267
268#[async_trait]
269impl LlmProvider for ProfiledProvider {
270 async fn stream(&self, request: ChatRequest) -> Result<ChatStream, LlmError> {
271 self.inner.stream(request).await
272 }
273
274 async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, LlmError> {
275 self.inner.chat(request).await
276 }
277
278 fn capabilities(&self) -> Capabilities {
279 self.profile.capabilities.clone()
280 }
281
282 fn info(&self) -> ProviderInfo {
283 ProviderInfo {
284 name: self.profile.provider_name.to_string(),
285 model: self.inner.adapter().info().model.clone(),
286 version: None,
287 }
288 }
289}
290
291#[cfg(test)]
292mod tests {
293 use super::*;
294 use llm_trait::{ChatMessage, FinishReason, HttpMethod, StreamChunk, UsageInfo};
295 use std::sync::Arc;
296
297 struct MockHttpClient {
299 responses: std::sync::Mutex<Vec<Result<llm_trait::HttpResponse, LlmError>>>,
300 calls: std::sync::atomic::AtomicUsize,
301 }
302
303 impl MockHttpClient {
304 fn new(responses: Vec<llm_trait::HttpResponse>) -> Self {
305 Self {
306 responses: std::sync::Mutex::new(responses.into_iter().map(Ok).collect()),
307 calls: std::sync::atomic::AtomicUsize::new(0),
308 }
309 }
310
311 fn calls(&self) -> usize {
312 self.calls.load(std::sync::atomic::Ordering::SeqCst)
313 }
314 }
315
316 #[async_trait]
317 impl HttpClient for MockHttpClient {
318 async fn send(&self, _request: &RawRequest) -> Result<llm_trait::HttpResponse, LlmError> {
319 self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
320 let mut queue = self.responses.lock().unwrap();
321 if queue.is_empty() {
322 return Err(LlmError::llm("MockHttpClient: no responses left"));
323 }
324 queue.remove(0)
325 }
326 }
327
328 struct MockAdapter;
330
331 #[async_trait]
332 impl RawAdapter for MockAdapter {
333 fn build_request(
334 &self,
335 _request: &ChatRequest,
336 mode: CallMode,
337 ) -> Result<RawRequest, LlmError> {
338 Ok(RawRequest {
339 url: "https://api.example.com/v1/messages".to_string(),
340 method: HttpMethod::Post,
341 headers: Default::default(),
342 body: serde_json::json!({"model": "test"}),
343 stream: mode == CallMode::Stream,
344 })
345 }
346
347 async fn execute_stream(
348 &self,
349 _client: &dyn HttpClient,
350 _request: RawRequest,
351 ) -> Result<ChatStream, LlmError> {
352 let chunks = vec![
353 Ok(StreamChunk::Text("hello".into())),
354 Ok(StreamChunk::Stop {
355 finish_reason: Some("stop".into()),
356 }),
357 ];
358 Ok(ChatStream::new(Box::pin(futures_util::stream::iter(
359 chunks,
360 ))))
361 }
362
363 async fn parse_sse_stream(
364 &self,
365 _client: &dyn HttpClient,
366 _request: RawRequest,
367 _response: llm_trait::HttpResponse,
368 ) -> Result<ChatStream, LlmError> {
369 let chunks = vec![
370 Ok(StreamChunk::Text("hello".into())),
371 Ok(StreamChunk::Stop {
372 finish_reason: Some("stop".into()),
373 }),
374 ];
375 Ok(ChatStream::new(Box::pin(futures_util::stream::iter(
376 chunks,
377 ))))
378 }
379
380 fn parse_response(&self, _body: &[u8]) -> Result<ChatResponse, LlmError> {
381 Ok(ChatResponse {
382 content: "mock response".to_string(),
383 reasoning_content: None,
384 thinking_signature: None,
385 tool_calls: vec![],
386 usage: UsageInfo::default(),
387 finish_reason: FinishReason::Stop,
388 raw: None,
389 })
390 }
391
392 fn capabilities(&self) -> Capabilities {
393 Capabilities {
394 supports_streaming: true,
395 supports_tools: true,
396 ..Default::default()
397 }
398 }
399
400 fn info(&self) -> ProviderInfo {
401 ProviderInfo {
402 name: "mock".to_string(),
403 model: "mock-model".to_string(),
404 version: None,
405 }
406 }
407
408 fn supported_modes(&self) -> &[CallMode] {
409 &[CallMode::Stream, CallMode::Once]
410 }
411 }
412
413 #[test]
414 fn generic_provider_info() {
415 let provider = GenericProvider::new(Box::new(MockAdapter));
416 let info = provider.info();
417 assert_eq!(info.name, "mock");
418 assert_eq!(info.model, "mock-model");
419 }
420
421 #[test]
422 fn generic_provider_capabilities() {
423 let provider = GenericProvider::new(Box::new(MockAdapter));
424 let caps = provider.capabilities();
425 assert!(caps.supports_streaming);
426 assert!(caps.supports_tools);
427 }
428
429 #[tokio::test]
430 async fn chat_uses_injected_http_client() {
431 let client = Arc::new(MockHttpClient::new(vec![
434 llm_trait::HttpResponse::from_text(200, "{}".to_string()),
435 ]));
436 let provider = GenericProvider::with_http_client(
437 Box::new(MockAdapter),
438 client.clone(),
439 ProviderConfig::default(),
440 );
441
442 let response = provider
443 .chat(ChatRequest::new(vec![ChatMessage::user("hi")]))
444 .await
445 .unwrap();
446
447 assert_eq!(response.content, "mock response");
448 assert_eq!(client.calls(), 1);
449 }
450
451 #[tokio::test]
452 async fn retries_5xx_then_succeeds() {
453 let client = Arc::new(MockHttpClient::new(vec![
454 llm_trait::HttpResponse::from_text(503, "overloaded".to_string()),
455 llm_trait::HttpResponse::from_text(200, "{}".to_string()),
456 ]));
457 let config = ProviderConfig {
458 retry_delay: Duration::ZERO,
459 max_retries: 2,
460 ..Default::default()
461 };
462 let provider =
463 GenericProvider::with_http_client(Box::new(MockAdapter), client.clone(), config);
464
465 let response = provider
466 .chat(ChatRequest::new(vec![ChatMessage::user("hi")]))
467 .await
468 .unwrap();
469
470 assert_eq!(response.content, "mock response");
471 assert_eq!(client.calls(), 2, "503 should be retried once");
472 }
473
474 #[tokio::test]
475 async fn does_not_retry_4xx() {
476 let client = Arc::new(MockHttpClient::new(vec![
477 llm_trait::HttpResponse::from_text(401, "bad key".to_string()),
478 ]));
479 let config = ProviderConfig {
480 retry_delay: Duration::ZERO,
481 max_retries: 3,
482 ..Default::default()
483 };
484 let provider =
485 GenericProvider::with_http_client(Box::new(MockAdapter), client.clone(), config);
486
487 let err = provider
488 .chat(ChatRequest::new(vec![ChatMessage::user("hi")]))
489 .await
490 .unwrap_err();
491
492 assert_eq!(err.status(), Some(401));
493 assert_eq!(client.calls(), 1, "401 must not be retried");
494 }
495
496 #[test]
501 fn profiled_provider_info() {
502 let profile = ModelProfile {
503 protocol: llm_trait::Protocol::OpenAi,
504 provider_name: "deepseek",
505 capabilities: Capabilities::default(),
506 reasoning_mode: llm_trait::ReasoningMode::Effort,
507 supported_extra_params: &[],
508 };
509 let provider = ProfiledProvider::new(GenericProvider::new(Box::new(MockAdapter)), profile);
510 let info = provider.info();
511 assert_eq!(info.name, "deepseek");
512 assert_eq!(info.model, "mock-model");
513 }
514
515 #[test]
516 fn provider_config_default() {
517 let config = ProviderConfig::default();
518 assert_eq!(config.connect_timeout, Duration::from_secs(15));
519 assert_eq!(config.request_timeout, Duration::from_secs(120));
520 assert_eq!(config.max_retries, 3);
521 }
522
523 #[test]
524 fn profiled_provider_capabilities() {
525 let caps = Capabilities {
526 supports_streaming: true,
527 supports_tools: false,
528 supports_vision: true,
529 ..Default::default()
530 };
531 let profile = ModelProfile {
532 protocol: llm_trait::Protocol::OpenAi,
533 provider_name: "test",
534 capabilities: caps.clone(),
535 reasoning_mode: llm_trait::ReasoningMode::Effort,
536 supported_extra_params: &[],
537 };
538 let provider = ProfiledProvider::new(GenericProvider::new(Box::new(MockAdapter)), profile);
539 let got = provider.capabilities();
540 assert!(got.supports_streaming);
541 assert!(!got.supports_tools);
542 assert!(got.supports_vision);
543 }
544
545 #[test]
546 fn calculate_backoff_respects_max() {
547 let provider = GenericProvider::new(Box::new(MockAdapter));
548 let delay = provider.calculate_backoff(20);
550 assert!(delay <= Duration::from_millis(30_100)); }
552
553 #[test]
554 fn calculate_backoff_increases_with_attempt() {
555 let provider = GenericProvider::new(Box::new(MockAdapter));
556 let mut delays: Vec<u64> = (1..=5)
558 .map(|a| provider.calculate_backoff(a).as_millis() as u64)
559 .collect();
560 delays.sort();
561 let d1 = provider.calculate_backoff(1).as_millis() as u64;
563 let d5 = provider.calculate_backoff(5).as_millis() as u64;
564 assert!(d5 > d1, "d5={} should be > d1={}", d5, d1);
566 }
567
568 #[tokio::test]
569 async fn profiled_provider_delegates_stream() {
570 let profile = ModelProfile {
571 protocol: llm_trait::Protocol::OpenAi,
572 provider_name: "test",
573 capabilities: Capabilities::default(),
574 reasoning_mode: llm_trait::ReasoningMode::Effort,
575 supported_extra_params: &[],
576 };
577 let provider = ProfiledProvider::new(GenericProvider::new(Box::new(MockAdapter)), profile);
578 let req = ChatRequest::new(vec![ChatMessage::user("hi")]);
579 let result = provider.stream(req).await;
583 assert!(result.is_ok() || result.is_err());
586 }
587}