Skip to main content

xz_provider/layer/
retry.rs

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
11/// 重试中间件 —— 指数退避 + 抖动,只重试 Transient 错误
12pub 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        // Stream retry: only retry the initial connection, not the stream itself
65        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
87/// 日志/指标中间件 —— 自动记录请求耗时、token 用量、错误
88pub 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}