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 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 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 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 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 let result = self.send_request(&provider_req).await;
74
75 match result {
76 Ok((status, body)) => {
77 if (200..300).contains(&status) {
78 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 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 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 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 pub async fn execute_stream(
200 &self,
201 request: &ChatRequest,
202 config: &AppConfig,
203 ) -> Result<(String, reqwest::Response)> {
204 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 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 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 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 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}