1use std::pin::Pin;
2
3use async_trait::async_trait;
4use futures::Stream;
5
6use super::{LayerService, ProviderLayer};
7use crate::error::{ProviderError, RetryStrategy};
8use crate::traits::LlmProvider;
9use crate::types::{CompletionRequest, CompletionResponse, RequestOptions, StreamEvent};
10
11pub struct RetryLayer {
13 strategy: RetryStrategy,
14}
15
16impl RetryLayer {
17 pub fn new(strategy: RetryStrategy) -> Self {
18 Self { strategy }
19 }
20}
21
22impl<S: LlmProvider + 'static> ProviderLayer<S> for RetryLayer {
23 type Stack = super::Layered<RetryLayerService, S>;
24
25 fn wrap(self, inner: S) -> Self::Stack {
26 super::Layered::new(RetryLayerService { strategy: self.strategy }, inner)
27 }
28}
29
30pub struct RetryLayerService {
31 strategy: RetryStrategy,
32}
33
34#[async_trait]
35impl<S: LlmProvider> LayerService<S> for RetryLayerService {
36 async fn complete(
37 &self,
38 inner: &S,
39 request: CompletionRequest,
40 options: RequestOptions,
41 ) -> Result<CompletionResponse, ProviderError> {
42 let mut attempt = 0u32;
43 loop {
44 match inner.complete(request.clone(), options.clone()).await {
45 Ok(resp) => return Ok(resp),
46 Err(e) if e.is_retryable() && attempt < self.strategy.max_retries => {
47 let delay = self.strategy.delay_for_attempt(attempt);
48 tokio::time::sleep(delay).await;
49 attempt += 1;
50 continue;
51 }
52 Err(e) => return Err(e),
53 }
54 }
55 }
56
57 async fn complete_stream(
58 &self,
59 inner: &S,
60 request: CompletionRequest,
61 options: RequestOptions,
62 ) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, ProviderError>> + Send>>, ProviderError>
63 {
64 let mut attempt = 0u32;
66 loop {
67 match inner.complete_stream(request.clone(), options.clone()).await {
68 Ok(stream) => return Ok(stream),
69 Err(e) if e.is_retryable() && attempt < self.strategy.max_retries => {
70 let delay = self.strategy.delay_for_attempt(attempt);
71 tokio::time::sleep(delay).await;
72 attempt += 1;
73 continue;
74 }
75 Err(e) => return Err(e),
76 }
77 }
78 }
79}
80
81impl Default for RetryLayer {
82 fn default() -> Self {
83 Self::new(RetryStrategy::default())
84 }
85}
86
87pub struct TelemetryLayer;
89
90impl<S: LlmProvider + 'static> ProviderLayer<S> for TelemetryLayer {
91 type Stack = super::Layered<TelemetryLayerService, S>;
92
93 fn wrap(self, inner: S) -> Self::Stack {
94 super::Layered::new(TelemetryLayerService, inner)
95 }
96}
97
98pub struct TelemetryLayerService;
99
100#[async_trait]
101impl<S: LlmProvider> LayerService<S> for TelemetryLayerService {
102 async fn complete(
103 &self,
104 inner: &S,
105 request: CompletionRequest,
106 options: RequestOptions,
107 ) -> Result<CompletionResponse, ProviderError> {
108 let start = std::time::Instant::now();
109 let model = request.model.clone().unwrap_or_default();
110 let has_tools = request.tools.is_some();
111 let tool_count = request.tools.as_ref().map(|t| t.len()).unwrap_or(0);
112
113 let _span = tracing::info_span!(
114 "llm_completion",
115 provider = %inner.name(),
116 model = %model,
117 has_tools = has_tools,
118 tool_count = tool_count,
119 stream = false,
120 );
121
122 let result = inner.complete(request, options).await;
123 let latency = start.elapsed();
124
125 match &result {
126 Ok(resp) => {
127 tracing::info!(
128 target: "xz_provider",
129 prompt_tokens = resp.usage.prompt_tokens,
130 completion_tokens = resp.usage.completion_tokens,
131 cached_tokens = resp.usage.cached_tokens.unwrap_or(0),
132 latency_ms = latency.as_millis() as u64,
133 finish_reason = ?resp.finish_reason,
134 tool_calls = resp.tool_calls.len(),
135 "llm_completion_completed"
136 );
137 }
138 Err(e) => {
139 tracing::warn!(
140 target: "xz_provider",
141 error = %e,
142 is_retryable = e.is_retryable(),
143 "llm_completion_failed"
144 );
145 }
146 }
147
148 result
149 }
150
151 async fn complete_stream(
152 &self,
153 inner: &S,
154 request: CompletionRequest,
155 options: RequestOptions,
156 ) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, ProviderError>> + Send>>, ProviderError>
157 {
158 inner.complete_stream(request, options).await
159 }
160}
161
162#[cfg(test)]
163mod tests {
164 use super::*;
165 use crate::types::*;
166 use std::time::Duration;
167
168 #[derive(Debug)]
169 struct MockProvider {
170 name: String,
171 behavior: std::sync::Mutex<Vec<Result<CompletionResponse, ProviderError>>>,
172 }
173
174 #[async_trait]
175 impl LlmProvider for MockProvider {
176 fn name(&self) -> &str {
177 &self.name
178 }
179 fn models(&self) -> &[ModelInfo] {
180 &[]
181 }
182 async fn complete(
183 &self,
184 _request: CompletionRequest,
185 _options: RequestOptions,
186 ) -> Result<CompletionResponse, ProviderError> {
187 let mut b = self.behavior.lock().unwrap();
188 if b.is_empty() {
189 Ok(CompletionResponse {
190 content: Some("ok".into()),
191 thinking: None,
192 tool_calls: vec![],
193 usage: TokenUsage::new(0, 0),
194 model: "mock".into(),
195 finish_reason: FinishReason::Stop,
196 latency_ms: 0,
197 cache_info: None,
198 ..Default::default()
199 })
200 } else {
201 b.remove(0)
202 }
203 }
204 async fn complete_stream(
205 &self,
206 _request: CompletionRequest,
207 _options: RequestOptions,
208 ) -> Result<
209 Pin<Box<dyn Stream<Item = Result<StreamEvent, ProviderError>> + Send>>,
210 ProviderError,
211 > {
212 Err(ProviderError::Config("not implemented".into()))
213 }
214 }
215
216 #[tokio::test]
217 async fn test_retry_layer_success() {
218 let mock = MockProvider { name: "test".into(), behavior: std::sync::Mutex::new(vec![]) };
219 let layer = RetryLayer::new(RetryStrategy {
220 max_retries: 2,
221 base_delay: Duration::from_millis(1),
222 max_delay: Duration::from_millis(10),
223 jitter: false,
224 });
225 let stacked = layer.wrap(mock);
226 let resp = stacked
227 .complete(CompletionRequest::default(), RequestOptions::default())
228 .await
229 .unwrap();
230 assert_eq!(resp.content.as_deref(), Some("ok"));
231 }
232
233 #[tokio::test]
234 async fn test_retry_layer_retries_then_succeeds() {
235 let mock = MockProvider {
236 name: "test".into(),
237 behavior: std::sync::Mutex::new(vec![
238 Err(ProviderError::Overloaded),
239 Err(ProviderError::Timeout { timeout_ms: 1000 }),
240 Ok(CompletionResponse {
241 content: Some("recovered".into()),
242 thinking: None,
243 tool_calls: vec![],
244 usage: TokenUsage::new(0, 0),
245 model: "mock".into(),
246 finish_reason: FinishReason::Stop,
247 latency_ms: 0,
248 cache_info: None,
249 ..Default::default()
250 }),
251 ]),
252 };
253 let layer = RetryLayer::new(RetryStrategy {
254 max_retries: 3,
255 base_delay: Duration::from_millis(1),
256 max_delay: Duration::from_millis(10),
257 jitter: false,
258 });
259 let stacked = layer.wrap(mock);
260 let resp = stacked
261 .complete(CompletionRequest::default(), RequestOptions::default())
262 .await
263 .unwrap();
264 assert_eq!(resp.content.as_deref(), Some("recovered"));
265 }
266
267 #[tokio::test]
268 async fn test_retry_layer_fatal_error_not_retried() {
269 let mock = MockProvider {
270 name: "test".into(),
271 behavior: std::sync::Mutex::new(vec![
272 Err(ProviderError::Auth("bad key".into())),
273 Ok(CompletionResponse {
274 content: Some("should not reach".into()),
275 thinking: None,
276 tool_calls: vec![],
277 usage: TokenUsage::new(0, 0),
278 model: "mock".into(),
279 finish_reason: FinishReason::Stop,
280 latency_ms: 0,
281 cache_info: None,
282 ..Default::default()
283 }),
284 ]),
285 };
286 let layer = RetryLayer::new(RetryStrategy {
287 max_retries: 3,
288 base_delay: Duration::from_millis(1),
289 max_delay: Duration::from_millis(10),
290 jitter: false,
291 });
292 let stacked = layer.wrap(mock);
293 let err = stacked
294 .complete(CompletionRequest::default(), RequestOptions::default())
295 .await
296 .unwrap_err();
297 assert!(matches!(err, ProviderError::Auth(_)));
298 }
299}