1use std::sync::Arc;
4
5use tap::Pipe;
6
7use crate::error::{GatewayError, LegFailure};
8use crate::providers::genai_provider::Provider;
9use crate::providers::Catalog;
10use crate::resilience::{run_with_classifier, ResilienceError};
11use crate::routing::effort::Effort;
12use crate::routing::request::ChatRequest;
13use crate::routing::stream::{Accumulator, FinishReason, StreamItem, ToolCallOut};
14use crate::routing::table::ChainLeg;
15use futures::stream::{BoxStream, Stream, StreamExt};
16
17#[derive(Debug, Clone)]
19pub struct Completion {
20 pub provider: String,
21 pub model: String,
22 pub content: String,
23 pub tool_calls: Vec<ToolCallOut>,
24 pub finish_reason: FinishReason,
25 pub input_tokens: u64,
26 pub output_tokens: u64,
27}
28
29fn to_genai_request(req: &ChatRequest) -> genai::chat::ChatRequest {
33 use genai::chat::{ChatMessage, Tool, ToolCall, ToolResponse};
34
35 let mut chat = genai::chat::ChatRequest::default();
36
37 for m in &req.messages {
38 match m.role.as_str() {
39 "assistant" if m.tool_calls.is_some() => {
40 let calls: Vec<ToolCall> = m
41 .tool_calls
42 .as_ref()
43 .unwrap()
44 .iter()
45 .filter_map(openai_tool_call_to_genai)
46 .collect();
47 chat = chat.append_message(ChatMessage::from(calls));
48 }
49 "tool" => {
50 let call_id = m.tool_call_id.clone().unwrap_or_default();
51 let content = m
52 .content
53 .as_str()
54 .map(str::to_string)
55 .unwrap_or_else(|| m.content.to_string());
56 chat = chat.append_message(ChatMessage::from(ToolResponse { call_id, content }));
57 }
58 role => {
59 let msg = match role {
60 "system" => ChatMessage::system(genai::chat::MessageContent::from_parts(
61 crate::routing::content_parts::content_to_genai_parts(&m.content),
62 )),
63 "assistant" => ChatMessage::assistant(genai::chat::MessageContent::from_parts(
64 crate::routing::content_parts::content_to_genai_parts(&m.content),
65 )),
66 _ => ChatMessage::user(genai::chat::MessageContent::from_parts(
67 crate::routing::content_parts::content_to_genai_parts(&m.content),
68 )),
69 };
70 chat = chat.append_message(msg);
71 }
72 }
73 }
74
75 if let Some(tools) = &req.tools {
76 let mapped: Vec<Tool> = tools.iter().filter_map(openai_tool_to_genai).collect();
77 if !mapped.is_empty() {
78 chat = chat.with_tools(mapped);
79 }
80 }
81
82 chat
83}
84
85fn openai_tool_to_genai(v: &serde_json::Value) -> Option<genai::chat::Tool> {
87 let f = v.get("function")?;
88 let name = f.get("name")?.as_str()?.to_string();
89 let mut tool = genai::chat::Tool::new(name);
90 if let Some(desc) = f.get("description").and_then(|d| d.as_str()) {
91 tool = tool.with_description(desc);
92 }
93 if let Some(params) = f.get("parameters") {
94 tool = tool.with_schema(params.clone());
95 }
96 Some(tool)
97}
98
99fn openai_tool_call_to_genai(v: &serde_json::Value) -> Option<genai::chat::ToolCall> {
102 let call_id = v.get("id")?.as_str()?.to_string();
103 let f = v.get("function")?;
104 let fn_name = f.get("name")?.as_str()?.to_string();
105 let raw_args = f.get("arguments").and_then(|a| a.as_str()).unwrap_or("{}");
106 let fn_arguments =
107 serde_json::from_str(raw_args).unwrap_or_else(|_| serde_json::json!(raw_args));
108 Some(genai::chat::ToolCall {
109 call_id,
110 fn_name,
111 fn_arguments,
112 thought_signatures: None,
113 })
114}
115
116fn to_genai_options(req: &ChatRequest, effort: Option<Effort>) -> genai::chat::ChatOptions {
117 genai::chat::ChatOptions::default()
118 .pipe(|o| match req.temperature {
119 Some(t) => o.with_temperature(t as f64),
120 None => o,
121 })
122 .pipe(|o| match &req.response_format {
123 Some(rf) => match rf.kind.as_str() {
124 "json_object" => o.with_response_format(genai::chat::ChatResponseFormat::JsonMode),
125 "json_schema" => {
126 if let Some(spec) = rf.json_schema.clone() {
127 o.with_response_format(genai::chat::ChatResponseFormat::JsonSpec(
128 genai::chat::JsonSpec::new("synapse", spec),
129 ))
130 } else {
131 o
132 }
133 }
134 _ => o,
135 },
136 None => o,
137 })
138 .pipe(
139 |o| match client_effort(req).or(effort).and_then(Effort::to_genai) {
140 Some(e) => o.with_reasoning_effort(e),
141 None => o,
142 },
143 )
144}
145
146pub(crate) fn client_effort(req: &ChatRequest) -> Option<Effort> {
148 req.passthrough
149 .get("reasoning_effort")
150 .and_then(Effort::from_value)
151}
152
153fn is_genai_retryable(e: &genai::Error) -> bool {
171 fn status_retryable(status: reqwest::StatusCode) -> bool {
173 status.is_server_error()
174 || status == reqwest::StatusCode::TOO_MANY_REQUESTS
175 || status == reqwest::StatusCode::REQUEST_TIMEOUT
176 }
177
178 fn webc_retryable(we: &genai::webc::Error) -> bool {
180 match we {
181 genai::webc::Error::ResponseFailedStatus { status, .. } => status_retryable(*status),
182 genai::webc::Error::Reqwest(re) => re.is_timeout() || re.is_connect(),
183 _ => false,
184 }
185 }
186
187 match e {
188 genai::Error::WebModelCall { webc_error, .. } => webc_retryable(webc_error),
189 genai::Error::WebAdapterCall { webc_error, .. } => webc_retryable(webc_error),
190 genai::Error::HttpError { status, .. } => status_retryable(*status),
191 _ => false,
192 }
193}
194
195#[derive(Debug)]
197pub enum LegError {
198 Start(String),
200 MidStream(String),
202}
203
204impl std::fmt::Display for LegError {
205 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
206 match self {
207 LegError::Start(s) => write!(f, "start: {s}"),
208 LegError::MidStream(s) => write!(f, "mid-stream: {s}"),
209 }
210 }
211}
212
213fn map_stop_reason(sr: Option<&genai::chat::StopReason>) -> FinishReason {
215 match sr {
216 Some(genai::chat::StopReason::ToolCall(_)) => FinishReason::ToolCalls,
217 Some(genai::chat::StopReason::MaxTokens(_)) => FinishReason::Length,
218 _ => FinishReason::Stop,
219 }
220}
221
222#[derive(Default)]
224struct ToolIndexer {
225 ids: Vec<String>,
226}
227impl ToolIndexer {
228 fn index_of(&mut self, call_id: &str) -> (u32, bool) {
230 if let Some(pos) = self.ids.iter().position(|c| c == call_id) {
231 (pos as u32, false)
232 } else {
233 self.ids.push(call_id.to_string());
234 ((self.ids.len() - 1) as u32, true)
235 }
236 }
237}
238
239pub async fn stream_one_leg_standard(
242 provider: &Arc<Provider>,
243 model: &str,
244 req: &ChatRequest,
245 effort: Option<Effort>,
246) -> Result<impl Stream<Item = Result<StreamItem, LegError>>, LegError> {
247 let chat_req = to_genai_request(req);
248 let opts = to_genai_options(req, effort)
249 .with_capture_usage(true)
250 .with_capture_tool_calls(true);
251
252 let resp = provider
253 .client
254 .exec_chat_stream(model.to_string(), chat_req, Some(&opts))
255 .await
256 .map_err(|e| LegError::Start(e.to_string()))?;
257
258 struct St<S> {
262 inner: S,
263 indexer: ToolIndexer,
264 input_tokens: u64,
265 output_tokens: u64,
266 }
267 let state = St {
268 inner: Box::pin(resp.stream),
269 indexer: ToolIndexer::default(),
270 input_tokens: 0,
271 output_tokens: 0,
272 };
273
274 let normalized = futures::stream::unfold(state, |mut st| async move {
275 loop {
276 match st.inner.next().await {
277 None => return None,
278 Some(Err(e)) => return Some((Err(LegError::MidStream(e.to_string())), st)),
279 Some(Ok(ev)) => match ev {
280 genai::chat::ChatStreamEvent::Chunk(c) => {
281 if !c.content.is_empty() {
282 return Some((Ok(StreamItem::Delta(c.content)), st));
283 }
284 }
285 genai::chat::ChatStreamEvent::ToolCallChunk(tc) => {
286 let call = tc.tool_call;
287 let (index, first) = st.indexer.index_of(&call.call_id);
288 let args = match &call.fn_arguments {
289 serde_json::Value::String(s) => s.clone(),
290 other => other.to_string(),
291 };
292 return Some((
293 Ok(StreamItem::ToolCallDelta {
294 index,
295 id: if first { Some(call.call_id) } else { None },
296 name: if first { Some(call.fn_name) } else { None },
297 args_fragment: args,
298 }),
299 st,
300 ));
301 }
302 genai::chat::ChatStreamEvent::End(end) => {
303 if let Some(u) = &end.captured_usage {
304 st.input_tokens = u.prompt_tokens.unwrap_or(0).max(0) as u64;
305 st.output_tokens = u.completion_tokens.unwrap_or(0).max(0) as u64;
306 }
307 let done = StreamItem::Done {
308 input_tokens: st.input_tokens,
309 output_tokens: st.output_tokens,
310 finish_reason: map_stop_reason(end.captured_stop_reason.as_ref()),
311 };
312 return Some((Ok(done), st));
313 }
314 _ => {} },
316 }
317 }
318 });
319
320 Ok(normalized)
321}
322
323#[cfg(test)]
324mod tests {
325 use super::*;
326
327 fn req(body: serde_json::Value) -> ChatRequest {
328 serde_json::from_value(body).unwrap()
329 }
330
331 #[test]
332 fn maps_tools_into_genai_request() {
333 let r = req(serde_json::json!({
334 "model": "m",
335 "messages": [{"role": "user", "content": "hi"}],
336 "tools": [{"type": "function", "function": {"name": "get_weather",
337 "description": "Lookup", "parameters": {"type": "object"}}}]
338 }));
339 let g = to_genai_request(&r);
340 let tools = g.tools.expect("tools mapped");
341 assert_eq!(tools.len(), 1);
342 assert_eq!(tools[0].name.to_string(), "get_weather");
343 assert_eq!(tools[0].description.as_deref(), Some("Lookup"));
344 }
345
346 #[test]
347 fn maps_assistant_tool_calls_and_tool_results() {
348 let r = req(serde_json::json!({
349 "model": "m",
350 "messages": [
351 {"role": "user", "content": "weather?"},
352 {"role": "assistant", "content": null, "tool_calls": [
353 {"id": "call_0", "type": "function", "function": {"name": "f", "arguments": "{\"c\":\"SF\"}"}}]},
354 {"role": "tool", "tool_call_id": "call_0", "content": "21C"}
355 ]
356 }));
357 let g = to_genai_request(&r);
358 assert_eq!(g.messages.len(), 3);
359 }
360
361 use crate::providers::genai_provider::{build_openai_compat_provider, OpenAiCompatConfig};
362 use crate::routing::stream::{FinishReason, StreamItem};
363 use futures::StreamExt;
364 use std::time::Duration;
365
366 #[tokio::test]
367 async fn standard_lane_streams_text_then_done() {
368 use wiremock::matchers::{method, path};
369 use wiremock::{Mock, MockServer, ResponseTemplate};
370
371 let sse = "data: {\"choices\":[{\"delta\":{\"content\":\"He\"}}]}\n\n\
373 data: {\"choices\":[{\"delta\":{\"content\":\"llo\"}}]}\n\n\
374 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":2,\"total_tokens\":5}}\n\n\
375 data: [DONE]\n\n";
376 let mock = MockServer::start().await;
377 Mock::given(method("POST"))
378 .and(path("/v1/chat/completions"))
379 .respond_with(
380 ResponseTemplate::new(200)
381 .insert_header("content-type", "text/event-stream")
382 .set_body_string(sse),
383 )
384 .mount(&mock)
385 .await;
386
387 let provider = Arc::new(
388 build_openai_compat_provider(
389 "oai",
390 OpenAiCompatConfig {
391 base_url: format!("{}/v1", mock.uri()),
392 api_key: "k".into(),
393 request_timeout: Duration::from_secs(5),
394 endpoint_override: None,
395 },
396 )
397 .unwrap(),
398 );
399
400 let req = req(
401 serde_json::json!({"model":"m","messages":[{"role":"user","content":"hi"}],"stream":true}),
402 );
403 let mut stream = std::pin::pin!(stream_one_leg_standard(&provider, "m", &req, None)
404 .await
405 .expect("stream starts"));
406 let mut items = Vec::new();
407 while let Some(it) = stream.next().await {
408 items.push(it.expect("no mid-stream error"));
409 }
410 assert!(items
411 .iter()
412 .any(|i| matches!(i, StreamItem::Delta(t) if t == "He")));
413 assert!(matches!(
414 items.last().unwrap(),
415 StreamItem::Done {
416 input_tokens: 3,
417 output_tokens: 2,
418 finish_reason: FinishReason::Stop
419 }
420 ));
421 }
422
423 #[tokio::test]
424 async fn execute_buffered_falls_back_on_midstream_failure() {
425 use crate::providers::Catalog;
426 use crate::routing::table::ChainLeg;
427 use wiremock::matchers::{method, path};
428 use wiremock::{Mock, MockServer, ResponseTemplate};
429
430 let bad = MockServer::start().await;
432 Mock::given(method("POST"))
433 .and(path("/v1/chat/completions"))
434 .respond_with(
435 ResponseTemplate::new(200)
436 .insert_header("content-type", "text/event-stream")
437 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"par\"}}]}\n\n"),
438 ) .mount(&bad)
440 .await;
441 let good = MockServer::start().await;
443 Mock::given(method("POST"))
444 .and(path("/v1/chat/completions"))
445 .respond_with(
446 ResponseTemplate::new(200)
447 .insert_header("content-type", "text/event-stream")
448 .set_body_string(
449 "data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}]}\n\n\
450 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
451 data: [DONE]\n\n",
452 ),
453 )
454 .mount(&good)
455 .await;
456
457 let catalog = Catalog::for_test(vec![
458 ("p1", format!("{}/v1", bad.uri())),
459 ("p2", format!("{}/v1", good.uri())),
460 ]);
461 let legs = vec![
462 ChainLeg {
463 provider: "p1".into(),
464 model: "m".into(),
465 ..Default::default()
466 },
467 ChainLeg {
468 provider: "p2".into(),
469 model: "m".into(),
470 ..Default::default()
471 },
472 ];
473 let r =
474 req(serde_json::json!({"model":"route","messages":[{"role":"user","content":"hi"}]}));
475 let c = execute_buffered(&catalog, "route", &legs, &r)
476 .await
477 .unwrap();
478 assert_eq!(c.content, "ok");
479 assert_eq!(c.provider, "p2");
480 }
481
482 #[tokio::test]
483 async fn execute_streaming_commits_first_leg_with_items() {
484 use crate::providers::Catalog;
485 use crate::routing::stream::StreamItem;
486 use crate::routing::table::ChainLeg;
487 use futures::StreamExt;
488 use wiremock::matchers::{method, path};
489 use wiremock::{Mock, MockServer, ResponseTemplate};
490
491 let bad = MockServer::start().await;
493 Mock::given(method("POST"))
494 .and(path("/v1/chat/completions"))
495 .respond_with(ResponseTemplate::new(500))
496 .mount(&bad)
497 .await;
498 let good = MockServer::start().await;
499 Mock::given(method("POST")).and(path("/v1/chat/completions"))
500 .respond_with(ResponseTemplate::new(200)
501 .insert_header("content-type", "text/event-stream")
502 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"go\"}}]}\n\n\
503 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
504 data: [DONE]\n\n"))
505 .mount(&good).await;
506
507 let catalog = Catalog::for_test(vec![
508 ("p1", format!("{}/v1", bad.uri())),
509 ("p2", format!("{}/v1", good.uri())),
510 ]);
511 let legs = vec![
512 ChainLeg {
513 provider: "p1".into(),
514 model: "m".into(),
515 ..Default::default()
516 },
517 ChainLeg {
518 provider: "p2".into(),
519 model: "m".into(),
520 ..Default::default()
521 },
522 ];
523 let r = req(
524 serde_json::json!({"model":"route","messages":[{"role":"user","content":"hi"}],"stream":true}),
525 );
526 let committed = execute_streaming(&catalog, "route", &legs, &r)
527 .await
528 .unwrap();
529 assert_eq!(committed.provider, "p2");
530 let mut stream = committed.stream;
531 let mut items = Vec::new();
532 while let Some(i) = stream.next().await {
533 items.push(i.unwrap());
534 }
535 assert!(items
536 .iter()
537 .any(|i| matches!(i, StreamItem::Delta(t) if t == "go")));
538 }
539
540 #[tokio::test]
541 async fn first_chunk_timeout_falls_back() {
542 use crate::providers::Catalog;
543 use crate::routing::table::ChainLeg;
544 use std::time::Duration;
545 use wiremock::matchers::{method, path};
546 use wiremock::{Mock, MockServer, ResponseTemplate};
547
548 let slow = MockServer::start().await;
559 Mock::given(method("POST")).and(path("/v1/chat/completions"))
560 .respond_with(ResponseTemplate::new(200)
561 .insert_header("content-type", "text/event-stream")
562 .set_delay(Duration::from_millis(400))
563 .set_body_string("data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{}}\n\ndata: [DONE]\n\n"))
564 .mount(&slow).await;
565 let good = MockServer::start().await;
566 Mock::given(method("POST")).and(path("/v1/chat/completions"))
567 .respond_with(ResponseTemplate::new(200)
568 .insert_header("content-type", "text/event-stream")
569 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}]}\n\n\
570 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
571 data: [DONE]\n\n")).mount(&good).await;
572
573 let catalog = Catalog::for_test(vec![
574 ("p1", format!("{}/v1", slow.uri())),
575 ("p2", format!("{}/v1", good.uri())),
576 ]);
577 let legs = vec![
578 ChainLeg {
579 provider: "p1".into(),
580 model: "m".into(),
581 ..Default::default()
582 },
583 ChainLeg {
584 provider: "p2".into(),
585 model: "m".into(),
586 ..Default::default()
587 },
588 ];
589 let r =
590 req(serde_json::json!({"model":"route","messages":[{"role":"user","content":"hi"}]}));
591 let timeouts = StreamTimeouts {
592 first_chunk: Duration::from_millis(150),
593 idle: Duration::from_secs(5),
594 };
595 let c = execute_buffered_with_timeouts(&catalog, "route", &legs, &r, timeouts)
596 .await
597 .unwrap();
598 assert_eq!(c.provider, "p2");
599 }
600
601 fn opts_req(extra: serde_json::Value) -> ChatRequest {
602 let mut body = serde_json::json!({
603 "model": "m",
604 "messages": [{"role": "user", "content": "hi"}]
605 });
606 body.as_object_mut()
607 .unwrap()
608 .extend(extra.as_object().cloned().unwrap_or_default());
609 serde_json::from_value(body).unwrap()
610 }
611
612 fn effort_name(o: &genai::chat::ChatOptions) -> Option<&'static str> {
613 o.reasoning_effort.as_ref().map(|e| e.variant_name())
614 }
615
616 #[test]
617 fn leg_effort_becomes_reasoning_effort() {
618 let o = to_genai_options(
619 &opts_req(serde_json::json!({})),
620 Some(crate::routing::effort::Effort::Medium),
621 );
622 assert_eq!(effort_name(&o), Some("medium"));
623 }
624
625 #[test]
626 fn effort_none_sends_no_reasoning_effort() {
627 let o = to_genai_options(
628 &opts_req(serde_json::json!({})),
629 Some(crate::routing::effort::Effort::None),
630 );
631 assert_eq!(effort_name(&o), None);
632 }
633
634 #[test]
635 fn client_reasoning_effort_is_forwarded() {
636 let o = to_genai_options(
637 &opts_req(serde_json::json!({"reasoning_effort": "high"})),
638 None,
639 );
640 assert_eq!(effort_name(&o), Some("high"));
641 }
642
643 #[test]
644 fn client_reasoning_effort_wins_over_leg_effort() {
645 let o = to_genai_options(
646 &opts_req(serde_json::json!({"reasoning_effort": "low"})),
647 Some(crate::routing::effort::Effort::Max),
648 );
649 assert_eq!(effort_name(&o), Some("low"));
650 }
651
652 #[test]
653 fn client_none_is_not_forwarded_and_leg_effort_does_not_replace_it() {
654 let o = to_genai_options(
655 &opts_req(serde_json::json!({"reasoning_effort": "none"})),
656 Some(Effort::High),
657 );
658 assert_eq!(effort_name(&o), None);
659 }
660
661 #[test]
662 fn unparseable_client_effort_falls_back_to_leg_effort() {
663 let o = to_genai_options(
664 &opts_req(serde_json::json!({"reasoning_effort": "extreme"})),
665 Some(Effort::Low),
666 );
667 assert_eq!(effort_name(&o), Some("low"));
668 }
669
670 #[test]
671 fn no_effort_anywhere_sends_nothing() {
672 assert_eq!(
673 effort_name(&to_genai_options(&opts_req(serde_json::json!({})), None)),
674 None
675 );
676 }
677}
678
679async fn run_one_leg(
680 provider: &Arc<Provider>,
681 leg: &ChainLeg,
682 req: &ChatRequest,
683) -> Result<Completion, ResilienceError<genai::Error>> {
684 let client = provider.client.clone();
685 let model = leg.model.clone();
686 let chat_req = to_genai_request(req);
687 let opts = to_genai_options(req, leg.effort);
688
689 let resp = run_with_classifier(
690 move || {
691 let (client, model, chat_req, opts) = (
692 client.clone(),
693 model.clone(),
694 chat_req.clone(),
695 opts.clone(),
696 );
697 async move { client.exec_chat(model, chat_req, Some(&opts)).await }
698 },
699 provider.profile,
700 &provider.breaker,
701 provider.label,
702 is_genai_retryable,
703 )
704 .await?;
705
706 let content = resp.first_text().unwrap_or_default().to_string();
707 let usage = &resp.usage;
708 Ok(Completion {
709 provider: leg.provider.clone(),
710 model: leg.model.clone(),
711 content,
712 tool_calls: Vec::new(),
713 finish_reason: FinishReason::Stop,
714 input_tokens: usage.prompt_tokens.unwrap_or(0).max(0) as u64,
715 output_tokens: usage.completion_tokens.unwrap_or(0).max(0) as u64,
716 })
717}
718
719pub async fn execute_chain(
722 catalog: &Catalog,
723 route_name: &str,
724 legs: &[ChainLeg],
725 req: &ChatRequest,
726) -> Result<Completion, GatewayError> {
727 let mut failures: Vec<LegFailure> = Vec::new();
728 let mut all_circuit_open = true;
729 for leg in legs {
730 let provider = catalog.get(&leg.provider).ok_or_else(|| {
731 GatewayError::BadRequest(format!(
732 "route '{route_name}' references unbuilt provider '{}'",
733 leg.provider
734 ))
735 })?;
736 match run_one_leg(provider, leg, req).await {
737 Ok(c) => return Ok(c),
738 Err(ResilienceError::CircuitOpen { name }) => failures.push(LegFailure {
739 provider: leg.provider.clone(),
740 model: leg.model.clone(),
741 message: format!("circuit open: {name}"),
742 }),
743 Err(ResilienceError::Exhausted(e)) => {
744 all_circuit_open = false;
745 let retryable = is_genai_retryable(&e);
746 failures.push(LegFailure {
747 provider: leg.provider.clone(),
748 model: leg.model.clone(),
749 message: e.to_string(),
750 });
751 if !retryable {
752 break; }
754 }
755 }
756 }
757 if all_circuit_open && !failures.is_empty() {
758 return Err(GatewayError::AllCircuitsOpen(route_name.to_string()));
759 }
760 Err(GatewayError::AllLegsFailed {
761 route: route_name.to_string(),
762 failures,
763 })
764}
765
766#[derive(Debug, Clone, Copy)]
768pub struct StreamTimeouts {
769 pub first_chunk: std::time::Duration,
771 pub idle: std::time::Duration,
773}
774
775impl Default for StreamTimeouts {
776 fn default() -> Self {
777 Self {
778 first_chunk: std::time::Duration::from_secs(120),
779 idle: std::time::Duration::from_secs(60),
780 }
781 }
782}
783
784async fn buffer_one_leg_timed(
786 catalog: &Catalog,
787 leg: &ChainLeg,
788 req: &ChatRequest,
789 t: StreamTimeouts,
790) -> Result<Completion, LegError> {
791 let provider = catalog
792 .get(&leg.provider)
793 .ok_or_else(|| LegError::Start(format!("unbuilt provider '{}'", leg.provider)))?;
794 let stream = stream_one_leg_standard(provider, &leg.model, req, leg.effort).await?;
795 let mut stream = std::pin::pin!(stream);
796 let mut acc = Accumulator::default();
797 let mut first = true;
798 loop {
799 let budget = if first { t.first_chunk } else { t.idle };
800 match tokio::time::timeout(budget, stream.next()).await {
801 Err(_) => {
802 return Err(if first {
803 LegError::Start("first-chunk timeout".into())
804 } else {
805 LegError::MidStream("idle timeout".into())
806 })
807 }
808 Ok(None) => break,
809 Ok(Some(item)) => {
810 acc.push(item?);
811 first = false;
812 }
813 }
814 }
815 if !acc.got_done {
816 return Err(LegError::MidStream("stream ended before completion".into()));
817 }
818 Ok(Completion {
819 provider: leg.provider.clone(),
820 model: leg.model.clone(),
821 content: acc.content,
822 tool_calls: acc.tool_calls,
823 finish_reason: acc.finish_reason,
824 input_tokens: acc.input_tokens,
825 output_tokens: acc.output_tokens,
826 })
827}
828
829pub async fn execute_buffered_with_timeouts(
834 catalog: &Catalog,
835 route_name: &str,
836 legs: &[ChainLeg],
837 req: &ChatRequest,
838 t: StreamTimeouts,
839) -> Result<Completion, GatewayError> {
840 let mut failures: Vec<LegFailure> = Vec::new();
841 for leg in legs {
842 match buffer_one_leg_timed(catalog, leg, req, t).await {
843 Ok(c) => return Ok(c),
844 Err(e) => failures.push(LegFailure {
845 provider: leg.provider.clone(),
846 model: leg.model.clone(),
847 message: e.to_string(),
848 }),
849 }
850 }
851 Err(GatewayError::AllLegsFailed {
852 route: route_name.to_string(),
853 failures,
854 })
855}
856
857pub async fn execute_buffered(
867 catalog: &Catalog,
868 route_name: &str,
869 legs: &[ChainLeg],
870 req: &ChatRequest,
871) -> Result<Completion, GatewayError> {
872 execute_buffered_with_timeouts(catalog, route_name, legs, req, StreamTimeouts::default()).await
873}
874
875pub struct CommittedStream {
878 pub provider: String,
879 pub model: String,
880 pub stream: BoxStream<'static, Result<StreamItem, LegError>>,
881}
882
883impl CommittedStream {
884 pub fn single(
887 provider: String,
888 model: String,
889 stream: impl Stream<Item = Result<StreamItem, LegError>> + Send + 'static,
890 ) -> Self {
891 Self {
892 provider,
893 model,
894 stream: stream.boxed(),
895 }
896 }
897}
898
899pub async fn execute_streaming_with_timeouts(
904 catalog: &Catalog,
905 route_name: &str,
906 legs: &[ChainLeg],
907 req: &ChatRequest,
908 t: StreamTimeouts,
909) -> Result<CommittedStream, GatewayError> {
910 let mut failures: Vec<LegFailure> = Vec::new();
911
912 for leg in legs {
913 let provider = match catalog.get(&leg.provider) {
914 Some(p) => p,
915 None => {
916 failures.push(LegFailure {
917 provider: leg.provider.clone(),
918 model: leg.model.clone(),
919 message: format!("unbuilt provider '{}'", leg.provider),
920 });
921 continue;
922 }
923 };
924 let started = match stream_one_leg_standard(provider, &leg.model, req, leg.effort).await {
925 Ok(s) => s,
926 Err(e) => {
927 failures.push(LegFailure {
928 provider: leg.provider.clone(),
929 model: leg.model.clone(),
930 message: e.to_string(),
931 });
932 continue;
933 }
934 };
935 let mut stream = Box::pin(started);
936 match tokio::time::timeout(t.first_chunk, stream.next()).await {
937 Err(_) => {
938 failures.push(LegFailure {
939 provider: leg.provider.clone(),
940 model: leg.model.clone(),
941 message: "first-chunk timeout".into(),
942 });
943 }
944 Ok(Some(Ok(first))) => {
945 let rest = futures::stream::once(async move { Ok(first) }).chain(stream);
946 return Ok(CommittedStream {
947 provider: leg.provider.clone(),
948 model: leg.model.clone(),
949 stream: rest.boxed(),
950 });
951 }
952 Ok(Some(Err(e))) => {
953 failures.push(LegFailure {
954 provider: leg.provider.clone(),
955 model: leg.model.clone(),
956 message: e.to_string(),
957 });
958 }
959 Ok(None) => {
960 failures.push(LegFailure {
961 provider: leg.provider.clone(),
962 model: leg.model.clone(),
963 message: "empty stream".into(),
964 });
965 }
966 }
967 }
968 Err(GatewayError::AllLegsFailed {
969 route: route_name.to_string(),
970 failures,
971 })
972}
973
974pub async fn execute_streaming(
981 catalog: &Catalog,
982 route_name: &str,
983 legs: &[ChainLeg],
984 req: &ChatRequest,
985) -> Result<CommittedStream, GatewayError> {
986 execute_streaming_with_timeouts(catalog, route_name, legs, req, StreamTimeouts::default()).await
987}