Skip to main content

weft_core/pipeline/
mod.rs

1pub mod context;
2
3use crate::config::AppConfig;
4use crate::layers::key_selector::ApiKeyState;
5use crate::layers::{
6    ErrorAction, ErrorHandlerLayer, KeySelectorLayer, RequestError, RouterLayer,
7};
8use crate::types::{ChatRequest, ChatResponse};
9use anyhow::{bail, Context, Result};
10use bytes::Bytes;
11use context::RequestContext;
12use std::sync::Arc;
13
14pub struct Pipeline {
15    pub router: Arc<dyn RouterLayer>,
16    pub key_selector: Arc<dyn KeySelectorLayer>,
17    pub transforms: Arc<crate::defaults::transforms::TransformRegistry>,
18    pub error_handler: Arc<dyn ErrorHandlerLayer>,
19    pub http_client: reqwest::Client,
20}
21
22impl Pipeline {
23    /// Execute a non-streaming chat request through the full pipeline.
24    pub async fn execute(&self, request: &ChatRequest, config: &AppConfig) -> Result<ChatResponse> {
25        let mut ctx = RequestContext::new();
26        let max_attempts = config.fallback.retry_count + 1;
27        let mut providers_tried: Vec<String> = vec![];
28
29        loop {
30            // 1. Route
31            let provider_name = self
32                .router
33                .route(request, &config.providers)
34                .await
35                .context("Routing failed")?;
36            ctx.selected_provider = Some(provider_name.clone());
37
38            let provider = config
39                .providers
40                .iter()
41                .find(|p| p.name == provider_name)
42                .ok_or_else(|| anyhow::anyhow!("Provider '{}' not found", provider_name))?;
43
44            // 2. Select key
45            let key_states: Vec<ApiKeyState> = provider
46                .keys
47                .iter()
48                .enumerate()
49                .map(|(i, k)| ApiKeyState::from((i, k)))
50                .collect();
51
52            if key_states.is_empty() {
53                bail!("No API keys configured for provider '{}'", provider_name);
54            }
55
56            let key_index = self
57                .key_selector
58                .select(&provider_name, &key_states)
59                .await
60                .context("Key selection failed")?;
61            ctx.selected_key_index = Some(key_index);
62            ctx.selected_key_value = Some(key_states[key_index].value.clone());
63
64            // 3. Transform request
65            let provider_req = self
66                .transforms
67                .for_format(&provider.format)
68                .transform_request(request, &key_states[key_index].value, provider)
69                .await
70                .context("Request transform failed")?;
71
72            // 4. Send HTTP request
73            let result = self.send_request(&provider_req).await;
74
75            match result {
76                Ok((status, body)) => {
77                    if (200..300).contains(&status) {
78                        // 5. Transform response
79                        let resp = self
80                            .transforms
81                            .for_format(&provider.format)
82                            .transform_response(status, body, provider)
83                            .await
84                            .context("Response transform failed")?;
85                        self.key_selector.mark_success(&provider_name, key_index);
86                        return Ok(resp);
87                    }
88
89                    // Error from provider
90                    let error = RequestError {
91                        status: Some(status),
92                        message: String::from_utf8_lossy(&body).to_string(),
93                        provider: provider_name.clone(),
94                        key_index,
95                        retry_count: ctx.retry_count,
96                    };
97
98                    match self.error_handler.handle(&error).await {
99                        ErrorAction::Retry { delay_ms } => {
100                            ctx.retry_count += 1;
101                            if ctx.retry_count > max_attempts {
102                                bail!("Max retries exceeded: {}", error.message);
103                            }
104                            tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await;
105                            continue;
106                        }
107                        ErrorAction::SwitchKey => {
108                            self.key_selector.mark_failed(&provider_name, key_index);
109                            ctx.retry_count += 1;
110                            if ctx.retry_count > max_attempts {
111                                bail!("Max retries exceeded after key switch: {}", error.message);
112                            }
113                            continue;
114                        }
115                        ErrorAction::SwitchProvider => {
116                            providers_tried.push(provider_name.clone());
117                            // Find next provider from fallback priority
118                            let next = config
119                                .fallback
120                                .priority
121                                .iter()
122                                .find(|p| !providers_tried.contains(p));
123                            if let Some(_next_provider) = next {
124                                ctx.retry_count += 1;
125                                continue;
126                            }
127                            bail!("All providers exhausted: {}", error.message);
128                        }
129                        ErrorAction::Fail { message } => {
130                            bail!("{}", message);
131                        }
132                    }
133                }
134                Err(e) => {
135                    // Network error
136                    let error = RequestError {
137                        status: None,
138                        message: e.to_string(),
139                        provider: provider_name.clone(),
140                        key_index,
141                        retry_count: ctx.retry_count,
142                    };
143
144                    match self.error_handler.handle(&error).await {
145                        ErrorAction::Retry { delay_ms } => {
146                            ctx.retry_count += 1;
147                            if ctx.retry_count > max_attempts {
148                                bail!("Max retries exceeded: {}", e);
149                            }
150                            tokio::time::sleep(tokio::time::Duration::from_millis(delay_ms)).await;
151                            continue;
152                        }
153                        ErrorAction::SwitchProvider => {
154                            providers_tried.push(provider_name);
155                            let next = config
156                                .fallback
157                                .priority
158                                .iter()
159                                .find(|p| !providers_tried.contains(p));
160                            if next.is_some() {
161                                ctx.retry_count += 1;
162                                continue;
163                            }
164                            bail!("All providers exhausted: {}", e);
165                        }
166                        _ => bail!("Request failed: {}", e),
167                    }
168                }
169            }
170        }
171    }
172
173    async fn send_request(
174        &self,
175        req: &crate::layers::transform::ProviderRequest,
176    ) -> Result<(u16, Bytes)> {
177        let mut builder = match req.method.as_str() {
178            "POST" => self.http_client.post(&req.url),
179            "GET" => self.http_client.get(&req.url),
180            _ => bail!("Unsupported method: {}", req.method),
181        };
182
183        for (k, v) in &req.headers {
184            builder = builder.header(k.as_str(), v.as_str());
185        }
186
187        let resp = builder
188            .body(req.body.clone())
189            .send()
190            .await
191            .context("HTTP request failed")?;
192
193        let status = resp.status().as_u16();
194        let body = resp.bytes().await.context("Failed to read response body")?;
195        Ok((status, body))
196    }
197
198    /// Execute a streaming chat request. Returns the provider name and raw reqwest::Response.
199    pub async fn execute_stream(
200        &self,
201        request: &ChatRequest,
202        config: &AppConfig,
203    ) -> Result<(String, reqwest::Response)> {
204        // 1. Route
205        let provider_name = self
206            .router
207            .route(request, &config.providers)
208            .await
209            .context("Routing failed")?;
210
211        let provider = config
212            .providers
213            .iter()
214            .find(|p| p.name == provider_name)
215            .ok_or_else(|| anyhow::anyhow!("Provider '{}' not found", provider_name))?;
216
217        // 2. Select key
218        let key_states: Vec<ApiKeyState> = provider
219            .keys
220            .iter()
221            .enumerate()
222            .map(|(i, k)| ApiKeyState::from((i, k)))
223            .collect();
224
225        if key_states.is_empty() {
226            bail!("No API keys configured for provider '{}'", provider_name);
227        }
228
229        let key_index = self
230            .key_selector
231            .select(&provider_name, &key_states)
232            .await
233            .context("Key selection failed")?;
234
235        // 3. Transform request (ensure stream=true)
236        let mut stream_request = request.clone();
237        stream_request.stream = true;
238
239        let provider_req = self
240            .transforms
241            .for_format(&provider.format)
242            .transform_request(&stream_request, &key_states[key_index].value, provider)
243            .await
244            .context("Request transform failed")?;
245
246        // 4. Send HTTP request with streaming response
247        let mut builder = match provider_req.method.as_str() {
248            "POST" => self.http_client.post(&provider_req.url),
249            "GET" => self.http_client.get(&provider_req.url),
250            _ => bail!("Unsupported method: {}", provider_req.method),
251        };
252
253        for (k, v) in &provider_req.headers {
254            builder = builder.header(k.as_str(), v.as_str());
255        }
256
257        let resp = builder
258            .body(provider_req.body.clone())
259            .send()
260            .await
261            .context("Streaming HTTP request failed")?;
262
263        let status = resp.status().as_u16();
264        if status >= 400 {
265            let body = resp.bytes().await.unwrap_or_default();
266            bail!(
267                "Provider returned status {}: {}",
268                status,
269                String::from_utf8_lossy(&body)
270            );
271        }
272
273        Ok((provider_name, resp))
274    }
275}
276
277#[cfg(test)]
278mod tests {
279    use super::*;
280    use crate::config::*;
281    use crate::defaults::*;
282    use crate::types::*;
283
284    fn test_config() -> AppConfig {
285        AppConfig {
286            core: CoreConfig::default(),
287            providers: vec![ProviderConfig {
288                name: "test".into(),
289                base_url: "http://localhost:9999".into(),
290                format: "openai".into(),
291                api: ProviderApi::ChatCompletions,
292                keys: vec![ApiKeyConfig {
293                    value: "sk-test".into(),
294                    label: None,
295                    enabled: true,
296                }],
297                models: vec!["test-model".into()],
298            }],
299            routing: RoutingConfig {
300                default_provider: Some("test".into()),
301                default_model: Some("test-model".into()),
302                ..Default::default()
303            },
304            key_strategy: KeyStrategyConfig::default(),
305            fallback: FallbackConfig {
306                retry_count: 0,
307                switch_key: false,
308                switch_provider: false,
309                priority: vec![],
310            },
311            virtual_keys: vec![],
312            services: vec![],
313            packages: vec![],
314            registry: RegistryConfig::default(),
315            package_aliases: Default::default(),
316            web_search: Default::default(),
317            team: Default::default(),
318        }
319    }
320
321    fn test_request() -> ChatRequest {
322        ChatRequest {
323            model: "test-model".into(),
324            messages: vec![ChatMessage {
325                role: "user".into(),
326                content: "hello".into(),
327                tool_calls: None,
328                tool_call_id: None,
329            }],
330            stream: false,
331            temperature: None,
332            max_tokens: None,
333            top_p: None,
334            tools: None,
335            tool_choice: None,
336            response_format: None,
337            x_provider: None,
338        }
339    }
340
341    fn test_pipeline() -> Pipeline {
342        Pipeline {
343            router: Arc::new(DefaultRouter {
344                default_provider: "test".into(),
345            }),
346            key_selector: Arc::new(FailoverSelector),
347            transforms: Arc::new(crate::defaults::transforms::TransformRegistry::with_defaults()),
348            error_handler: Arc::new(DefaultErrorHandler { max_retries: 2 }),
349            http_client: reqwest::Client::builder()
350                .timeout(std::time::Duration::from_secs(2))
351                .connect_timeout(std::time::Duration::from_secs(1))
352                .build()
353                .unwrap(),
354        }
355    }
356
357    #[tokio::test]
358    async fn test_pipeline_routes_correctly() {
359        let pipeline = test_pipeline();
360        let config = test_config();
361        let result = pipeline.execute(&test_request(), &config).await;
362        // Should fail with connection error (no server running)
363        assert!(result.is_err());
364        let err = result.unwrap_err().to_string();
365        assert!(
366            err.contains("error") || err.contains("connect") || err.contains("failed"),
367            "Unexpected error: {}",
368            err
369        );
370    }
371}