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::request::ChatRequest;
12use crate::routing::stream::{Accumulator, FinishReason, StreamItem, ToolCallOut};
13use crate::routing::table::ChainLeg;
14use futures::stream::{BoxStream, Stream, StreamExt};
15
16#[derive(Debug, Clone)]
18pub struct Completion {
19 pub provider: String,
20 pub model: String,
21 pub content: String,
22 pub tool_calls: Vec<ToolCallOut>,
23 pub finish_reason: FinishReason,
24 pub input_tokens: u64,
25 pub output_tokens: u64,
26}
27
28fn to_genai_request(req: &ChatRequest) -> genai::chat::ChatRequest {
32 use genai::chat::{ChatMessage, Tool, ToolCall, ToolResponse};
33
34 let mut chat = genai::chat::ChatRequest::default();
35
36 for m in &req.messages {
37 match m.role.as_str() {
38 "assistant" if m.tool_calls.is_some() => {
39 let calls: Vec<ToolCall> = m
40 .tool_calls
41 .as_ref()
42 .unwrap()
43 .iter()
44 .filter_map(openai_tool_call_to_genai)
45 .collect();
46 chat = chat.append_message(ChatMessage::from(calls));
47 }
48 "tool" => {
49 let call_id = m.tool_call_id.clone().unwrap_or_default();
50 let content = m
51 .content
52 .as_str()
53 .map(str::to_string)
54 .unwrap_or_else(|| m.content.to_string());
55 chat = chat.append_message(ChatMessage::from(ToolResponse { call_id, content }));
56 }
57 role => {
58 let text = m
59 .content
60 .as_str()
61 .map(str::to_string)
62 .unwrap_or_else(|| m.content.to_string());
63 let msg = match role {
64 "system" => ChatMessage::system(text),
65 "assistant" => ChatMessage::assistant(text),
66 _ => ChatMessage::user(text),
67 };
68 chat = chat.append_message(msg);
69 }
70 }
71 }
72
73 if let Some(tools) = &req.tools {
74 let mapped: Vec<Tool> = tools.iter().filter_map(openai_tool_to_genai).collect();
75 if !mapped.is_empty() {
76 chat = chat.with_tools(mapped);
77 }
78 }
79
80 chat
81}
82
83fn openai_tool_to_genai(v: &serde_json::Value) -> Option<genai::chat::Tool> {
85 let f = v.get("function")?;
86 let name = f.get("name")?.as_str()?.to_string();
87 let mut tool = genai::chat::Tool::new(name);
88 if let Some(desc) = f.get("description").and_then(|d| d.as_str()) {
89 tool = tool.with_description(desc);
90 }
91 if let Some(params) = f.get("parameters") {
92 tool = tool.with_schema(params.clone());
93 }
94 Some(tool)
95}
96
97fn openai_tool_call_to_genai(v: &serde_json::Value) -> Option<genai::chat::ToolCall> {
100 let call_id = v.get("id")?.as_str()?.to_string();
101 let f = v.get("function")?;
102 let fn_name = f.get("name")?.as_str()?.to_string();
103 let raw_args = f.get("arguments").and_then(|a| a.as_str()).unwrap_or("{}");
104 let fn_arguments =
105 serde_json::from_str(raw_args).unwrap_or_else(|_| serde_json::json!(raw_args));
106 Some(genai::chat::ToolCall {
107 call_id,
108 fn_name,
109 fn_arguments,
110 thought_signatures: None,
111 })
112}
113
114fn to_genai_options(req: &ChatRequest) -> genai::chat::ChatOptions {
115 genai::chat::ChatOptions::default().pipe(|o| match req.temperature {
116 Some(t) => o.with_temperature(t as f64),
117 None => o,
118 })
119}
120
121fn is_genai_retryable(e: &genai::Error) -> bool {
139 fn webc_retryable(we: &genai::webc::Error) -> bool {
141 match we {
142 genai::webc::Error::ResponseFailedStatus { status, .. } => status.is_server_error(),
143 genai::webc::Error::Reqwest(re) => re.is_timeout() || re.is_connect(),
144 _ => false,
145 }
146 }
147
148 match e {
149 genai::Error::WebModelCall { webc_error, .. } => webc_retryable(webc_error),
150 genai::Error::WebAdapterCall { webc_error, .. } => webc_retryable(webc_error),
151 genai::Error::HttpError { status, .. } => status.is_server_error(),
152 _ => false,
153 }
154}
155
156#[derive(Debug)]
158pub enum LegError {
159 Start(String),
161 MidStream(String),
163}
164
165impl std::fmt::Display for LegError {
166 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
167 match self {
168 LegError::Start(s) => write!(f, "start: {s}"),
169 LegError::MidStream(s) => write!(f, "mid-stream: {s}"),
170 }
171 }
172}
173
174fn map_stop_reason(sr: Option<&genai::chat::StopReason>) -> FinishReason {
176 match sr {
177 Some(genai::chat::StopReason::ToolCall(_)) => FinishReason::ToolCalls,
178 Some(genai::chat::StopReason::MaxTokens(_)) => FinishReason::Length,
179 _ => FinishReason::Stop,
180 }
181}
182
183#[derive(Default)]
185struct ToolIndexer {
186 ids: Vec<String>,
187}
188impl ToolIndexer {
189 fn index_of(&mut self, call_id: &str) -> (u32, bool) {
191 if let Some(pos) = self.ids.iter().position(|c| c == call_id) {
192 (pos as u32, false)
193 } else {
194 self.ids.push(call_id.to_string());
195 ((self.ids.len() - 1) as u32, true)
196 }
197 }
198}
199
200pub async fn stream_one_leg_standard(
203 provider: &Arc<Provider>,
204 model: &str,
205 req: &ChatRequest,
206) -> Result<impl Stream<Item = Result<StreamItem, LegError>>, LegError> {
207 let chat_req = to_genai_request(req);
208 let opts = to_genai_options(req)
209 .with_capture_usage(true)
210 .with_capture_tool_calls(true);
211
212 let resp = provider
213 .client
214 .exec_chat_stream(model.to_string(), chat_req, Some(&opts))
215 .await
216 .map_err(|e| LegError::Start(e.to_string()))?;
217
218 struct St<S> {
222 inner: S,
223 indexer: ToolIndexer,
224 input_tokens: u64,
225 output_tokens: u64,
226 }
227 let state = St {
228 inner: Box::pin(resp.stream),
229 indexer: ToolIndexer::default(),
230 input_tokens: 0,
231 output_tokens: 0,
232 };
233
234 let normalized = futures::stream::unfold(state, |mut st| async move {
235 loop {
236 match st.inner.next().await {
237 None => return None,
238 Some(Err(e)) => return Some((Err(LegError::MidStream(e.to_string())), st)),
239 Some(Ok(ev)) => match ev {
240 genai::chat::ChatStreamEvent::Chunk(c) => {
241 if !c.content.is_empty() {
242 return Some((Ok(StreamItem::Delta(c.content)), st));
243 }
244 }
245 genai::chat::ChatStreamEvent::ToolCallChunk(tc) => {
246 let call = tc.tool_call;
247 let (index, first) = st.indexer.index_of(&call.call_id);
248 let args = match &call.fn_arguments {
249 serde_json::Value::String(s) => s.clone(),
250 other => other.to_string(),
251 };
252 return Some((
253 Ok(StreamItem::ToolCallDelta {
254 index,
255 id: if first { Some(call.call_id) } else { None },
256 name: if first { Some(call.fn_name) } else { None },
257 args_fragment: args,
258 }),
259 st,
260 ));
261 }
262 genai::chat::ChatStreamEvent::End(end) => {
263 if let Some(u) = &end.captured_usage {
264 st.input_tokens = u.prompt_tokens.unwrap_or(0).max(0) as u64;
265 st.output_tokens = u.completion_tokens.unwrap_or(0).max(0) as u64;
266 }
267 let done = StreamItem::Done {
268 input_tokens: st.input_tokens,
269 output_tokens: st.output_tokens,
270 finish_reason: map_stop_reason(end.captured_stop_reason.as_ref()),
271 };
272 return Some((Ok(done), st));
273 }
274 _ => {} },
276 }
277 }
278 });
279
280 Ok(normalized)
281}
282
283#[cfg(test)]
284mod tests {
285 use super::*;
286
287 fn req(body: serde_json::Value) -> ChatRequest {
288 serde_json::from_value(body).unwrap()
289 }
290
291 #[test]
292 fn maps_tools_into_genai_request() {
293 let r = req(serde_json::json!({
294 "model": "m",
295 "messages": [{"role": "user", "content": "hi"}],
296 "tools": [{"type": "function", "function": {"name": "get_weather",
297 "description": "Lookup", "parameters": {"type": "object"}}}]
298 }));
299 let g = to_genai_request(&r);
300 let tools = g.tools.expect("tools mapped");
301 assert_eq!(tools.len(), 1);
302 assert_eq!(tools[0].name.to_string(), "get_weather");
303 assert_eq!(tools[0].description.as_deref(), Some("Lookup"));
304 }
305
306 #[test]
307 fn maps_assistant_tool_calls_and_tool_results() {
308 let r = req(serde_json::json!({
309 "model": "m",
310 "messages": [
311 {"role": "user", "content": "weather?"},
312 {"role": "assistant", "content": null, "tool_calls": [
313 {"id": "call_0", "type": "function", "function": {"name": "f", "arguments": "{\"c\":\"SF\"}"}}]},
314 {"role": "tool", "tool_call_id": "call_0", "content": "21C"}
315 ]
316 }));
317 let g = to_genai_request(&r);
318 assert_eq!(g.messages.len(), 3);
319 }
320
321 use crate::providers::genai_provider::{build_openai_compat_provider, OpenAiCompatConfig};
322 use crate::routing::stream::{FinishReason, StreamItem};
323 use futures::StreamExt;
324 use std::time::Duration;
325
326 #[tokio::test]
327 async fn standard_lane_streams_text_then_done() {
328 use wiremock::matchers::{method, path};
329 use wiremock::{Mock, MockServer, ResponseTemplate};
330
331 let sse = "data: {\"choices\":[{\"delta\":{\"content\":\"He\"}}]}\n\n\
333 data: {\"choices\":[{\"delta\":{\"content\":\"llo\"}}]}\n\n\
334 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":2,\"total_tokens\":5}}\n\n\
335 data: [DONE]\n\n";
336 let mock = MockServer::start().await;
337 Mock::given(method("POST"))
338 .and(path("/v1/chat/completions"))
339 .respond_with(
340 ResponseTemplate::new(200)
341 .insert_header("content-type", "text/event-stream")
342 .set_body_string(sse),
343 )
344 .mount(&mock)
345 .await;
346
347 let provider = Arc::new(
348 build_openai_compat_provider(
349 "oai",
350 OpenAiCompatConfig {
351 base_url: format!("{}/v1", mock.uri()),
352 api_key: "k".into(),
353 request_timeout: Duration::from_secs(5),
354 endpoint_override: None,
355 },
356 )
357 .unwrap(),
358 );
359
360 let req = req(
361 serde_json::json!({"model":"m","messages":[{"role":"user","content":"hi"}],"stream":true}),
362 );
363 let mut stream = std::pin::pin!(stream_one_leg_standard(&provider, "m", &req)
364 .await
365 .expect("stream starts"));
366 let mut items = Vec::new();
367 while let Some(it) = stream.next().await {
368 items.push(it.expect("no mid-stream error"));
369 }
370 assert!(items
371 .iter()
372 .any(|i| matches!(i, StreamItem::Delta(t) if t == "He")));
373 assert!(matches!(
374 items.last().unwrap(),
375 StreamItem::Done {
376 input_tokens: 3,
377 output_tokens: 2,
378 finish_reason: FinishReason::Stop
379 }
380 ));
381 }
382
383 #[tokio::test]
384 async fn execute_buffered_falls_back_on_midstream_failure() {
385 use crate::providers::Catalog;
386 use crate::routing::table::ChainLeg;
387 use wiremock::matchers::{method, path};
388 use wiremock::{Mock, MockServer, ResponseTemplate};
389
390 let bad = MockServer::start().await;
392 Mock::given(method("POST"))
393 .and(path("/v1/chat/completions"))
394 .respond_with(
395 ResponseTemplate::new(200)
396 .insert_header("content-type", "text/event-stream")
397 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"par\"}}]}\n\n"),
398 ) .mount(&bad)
400 .await;
401 let good = MockServer::start().await;
403 Mock::given(method("POST"))
404 .and(path("/v1/chat/completions"))
405 .respond_with(
406 ResponseTemplate::new(200)
407 .insert_header("content-type", "text/event-stream")
408 .set_body_string(
409 "data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}]}\n\n\
410 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
411 data: [DONE]\n\n",
412 ),
413 )
414 .mount(&good)
415 .await;
416
417 let catalog = Catalog::for_test(vec![
418 ("p1", format!("{}/v1", bad.uri())),
419 ("p2", format!("{}/v1", good.uri())),
420 ]);
421 let legs = vec![
422 ChainLeg {
423 provider: "p1".into(),
424 model: "m".into(),
425 },
426 ChainLeg {
427 provider: "p2".into(),
428 model: "m".into(),
429 },
430 ];
431 let r =
432 req(serde_json::json!({"model":"route","messages":[{"role":"user","content":"hi"}]}));
433 let c = execute_buffered(&catalog, "route", &legs, &r)
434 .await
435 .unwrap();
436 assert_eq!(c.content, "ok");
437 assert_eq!(c.provider, "p2");
438 }
439
440 #[tokio::test]
441 async fn execute_streaming_commits_first_leg_with_items() {
442 use crate::providers::Catalog;
443 use crate::routing::stream::StreamItem;
444 use crate::routing::table::ChainLeg;
445 use futures::StreamExt;
446 use wiremock::matchers::{method, path};
447 use wiremock::{Mock, MockServer, ResponseTemplate};
448
449 let bad = MockServer::start().await;
451 Mock::given(method("POST"))
452 .and(path("/v1/chat/completions"))
453 .respond_with(ResponseTemplate::new(500))
454 .mount(&bad)
455 .await;
456 let good = MockServer::start().await;
457 Mock::given(method("POST")).and(path("/v1/chat/completions"))
458 .respond_with(ResponseTemplate::new(200)
459 .insert_header("content-type", "text/event-stream")
460 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"go\"}}]}\n\n\
461 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
462 data: [DONE]\n\n"))
463 .mount(&good).await;
464
465 let catalog = Catalog::for_test(vec![
466 ("p1", format!("{}/v1", bad.uri())),
467 ("p2", format!("{}/v1", good.uri())),
468 ]);
469 let legs = vec![
470 ChainLeg {
471 provider: "p1".into(),
472 model: "m".into(),
473 },
474 ChainLeg {
475 provider: "p2".into(),
476 model: "m".into(),
477 },
478 ];
479 let r = req(
480 serde_json::json!({"model":"route","messages":[{"role":"user","content":"hi"}],"stream":true}),
481 );
482 let committed = execute_streaming(&catalog, "route", &legs, &r)
483 .await
484 .unwrap();
485 assert_eq!(committed.provider, "p2");
486 let mut stream = committed.stream;
487 let mut items = Vec::new();
488 while let Some(i) = stream.next().await {
489 items.push(i.unwrap());
490 }
491 assert!(items
492 .iter()
493 .any(|i| matches!(i, StreamItem::Delta(t) if t == "go")));
494 }
495
496 #[tokio::test]
497 async fn first_chunk_timeout_falls_back() {
498 use crate::providers::Catalog;
499 use crate::routing::table::ChainLeg;
500 use std::time::Duration;
501 use wiremock::matchers::{method, path};
502 use wiremock::{Mock, MockServer, ResponseTemplate};
503
504 let slow = MockServer::start().await;
515 Mock::given(method("POST")).and(path("/v1/chat/completions"))
516 .respond_with(ResponseTemplate::new(200)
517 .insert_header("content-type", "text/event-stream")
518 .set_delay(Duration::from_millis(400))
519 .set_body_string("data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{}}\n\ndata: [DONE]\n\n"))
520 .mount(&slow).await;
521 let good = MockServer::start().await;
522 Mock::given(method("POST")).and(path("/v1/chat/completions"))
523 .respond_with(ResponseTemplate::new(200)
524 .insert_header("content-type", "text/event-stream")
525 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}]}\n\n\
526 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
527 data: [DONE]\n\n")).mount(&good).await;
528
529 let catalog = Catalog::for_test(vec![
530 ("p1", format!("{}/v1", slow.uri())),
531 ("p2", format!("{}/v1", good.uri())),
532 ]);
533 let legs = vec![
534 ChainLeg {
535 provider: "p1".into(),
536 model: "m".into(),
537 },
538 ChainLeg {
539 provider: "p2".into(),
540 model: "m".into(),
541 },
542 ];
543 let r =
544 req(serde_json::json!({"model":"route","messages":[{"role":"user","content":"hi"}]}));
545 let timeouts = StreamTimeouts {
546 first_chunk: Duration::from_millis(150),
547 idle: Duration::from_secs(5),
548 };
549 let c = execute_buffered_with_timeouts(&catalog, "route", &legs, &r, timeouts)
550 .await
551 .unwrap();
552 assert_eq!(c.provider, "p2");
553 }
554}
555
556async fn run_one_leg(
557 provider: &Arc<Provider>,
558 leg: &ChainLeg,
559 req: &ChatRequest,
560) -> Result<Completion, ResilienceError<genai::Error>> {
561 let client = provider.client.clone();
562 let model = leg.model.clone();
563 let chat_req = to_genai_request(req);
564 let opts = to_genai_options(req);
565
566 let resp = run_with_classifier(
567 move || {
568 let (client, model, chat_req, opts) = (
569 client.clone(),
570 model.clone(),
571 chat_req.clone(),
572 opts.clone(),
573 );
574 async move { client.exec_chat(model, chat_req, Some(&opts)).await }
575 },
576 provider.profile,
577 &provider.breaker,
578 provider.label,
579 is_genai_retryable,
580 )
581 .await?;
582
583 let content = resp.first_text().unwrap_or_default().to_string();
584 let usage = &resp.usage;
585 Ok(Completion {
586 provider: leg.provider.clone(),
587 model: leg.model.clone(),
588 content,
589 tool_calls: Vec::new(),
590 finish_reason: FinishReason::Stop,
591 input_tokens: usage.prompt_tokens.unwrap_or(0).max(0) as u64,
592 output_tokens: usage.completion_tokens.unwrap_or(0).max(0) as u64,
593 })
594}
595
596pub async fn execute_chain(
599 catalog: &Catalog,
600 route_name: &str,
601 legs: &[ChainLeg],
602 req: &ChatRequest,
603) -> Result<Completion, GatewayError> {
604 let mut failures: Vec<LegFailure> = Vec::new();
605 let mut all_circuit_open = true;
606 for leg in legs {
607 let provider = catalog.get(&leg.provider).ok_or_else(|| {
608 GatewayError::BadRequest(format!(
609 "route '{route_name}' references unbuilt provider '{}'",
610 leg.provider
611 ))
612 })?;
613 match run_one_leg(provider, leg, req).await {
614 Ok(c) => return Ok(c),
615 Err(ResilienceError::CircuitOpen { name }) => failures.push(LegFailure {
616 provider: leg.provider.clone(),
617 model: leg.model.clone(),
618 message: format!("circuit open: {name}"),
619 }),
620 Err(ResilienceError::Exhausted(e)) => {
621 all_circuit_open = false;
622 let retryable = is_genai_retryable(&e);
623 failures.push(LegFailure {
624 provider: leg.provider.clone(),
625 model: leg.model.clone(),
626 message: e.to_string(),
627 });
628 if !retryable {
629 break; }
631 }
632 }
633 }
634 if all_circuit_open && !failures.is_empty() {
635 return Err(GatewayError::AllCircuitsOpen(route_name.to_string()));
636 }
637 Err(GatewayError::AllLegsFailed {
638 route: route_name.to_string(),
639 failures,
640 })
641}
642
643#[derive(Debug, Clone, Copy)]
645pub struct StreamTimeouts {
646 pub first_chunk: std::time::Duration,
648 pub idle: std::time::Duration,
650}
651
652impl Default for StreamTimeouts {
653 fn default() -> Self {
654 Self {
655 first_chunk: std::time::Duration::from_secs(120),
656 idle: std::time::Duration::from_secs(60),
657 }
658 }
659}
660
661async fn buffer_one_leg_timed(
663 catalog: &Catalog,
664 leg: &ChainLeg,
665 req: &ChatRequest,
666 t: StreamTimeouts,
667) -> Result<Completion, LegError> {
668 let provider = catalog
669 .get(&leg.provider)
670 .ok_or_else(|| LegError::Start(format!("unbuilt provider '{}'", leg.provider)))?;
671 let stream = stream_one_leg_standard(provider, &leg.model, req).await?;
672 let mut stream = std::pin::pin!(stream);
673 let mut acc = Accumulator::default();
674 let mut first = true;
675 loop {
676 let budget = if first { t.first_chunk } else { t.idle };
677 match tokio::time::timeout(budget, stream.next()).await {
678 Err(_) => {
679 return Err(if first {
680 LegError::Start("first-chunk timeout".into())
681 } else {
682 LegError::MidStream("idle timeout".into())
683 })
684 }
685 Ok(None) => break,
686 Ok(Some(item)) => {
687 acc.push(item?);
688 first = false;
689 }
690 }
691 }
692 if !acc.got_done {
693 return Err(LegError::MidStream("stream ended before completion".into()));
694 }
695 Ok(Completion {
696 provider: leg.provider.clone(),
697 model: leg.model.clone(),
698 content: acc.content,
699 tool_calls: acc.tool_calls,
700 finish_reason: acc.finish_reason,
701 input_tokens: acc.input_tokens,
702 output_tokens: acc.output_tokens,
703 })
704}
705
706pub async fn execute_buffered_with_timeouts(
711 catalog: &Catalog,
712 route_name: &str,
713 legs: &[ChainLeg],
714 req: &ChatRequest,
715 t: StreamTimeouts,
716) -> Result<Completion, GatewayError> {
717 let mut failures: Vec<LegFailure> = Vec::new();
718 for leg in legs {
719 match buffer_one_leg_timed(catalog, leg, req, t).await {
720 Ok(c) => return Ok(c),
721 Err(e) => failures.push(LegFailure {
722 provider: leg.provider.clone(),
723 model: leg.model.clone(),
724 message: e.to_string(),
725 }),
726 }
727 }
728 Err(GatewayError::AllLegsFailed {
729 route: route_name.to_string(),
730 failures,
731 })
732}
733
734pub async fn execute_buffered(
744 catalog: &Catalog,
745 route_name: &str,
746 legs: &[ChainLeg],
747 req: &ChatRequest,
748) -> Result<Completion, GatewayError> {
749 execute_buffered_with_timeouts(catalog, route_name, legs, req, StreamTimeouts::default()).await
750}
751
752pub struct CommittedStream {
755 pub provider: String,
756 pub model: String,
757 pub stream: BoxStream<'static, Result<StreamItem, LegError>>,
758}
759
760impl CommittedStream {
761 pub fn single(
764 provider: String,
765 model: String,
766 stream: impl Stream<Item = Result<StreamItem, LegError>> + Send + 'static,
767 ) -> Self {
768 Self {
769 provider,
770 model,
771 stream: stream.boxed(),
772 }
773 }
774}
775
776pub async fn execute_streaming_with_timeouts(
781 catalog: &Catalog,
782 route_name: &str,
783 legs: &[ChainLeg],
784 req: &ChatRequest,
785 t: StreamTimeouts,
786) -> Result<CommittedStream, GatewayError> {
787 let mut failures: Vec<LegFailure> = Vec::new();
788
789 for leg in legs {
790 let provider = match catalog.get(&leg.provider) {
791 Some(p) => p,
792 None => {
793 failures.push(LegFailure {
794 provider: leg.provider.clone(),
795 model: leg.model.clone(),
796 message: format!("unbuilt provider '{}'", leg.provider),
797 });
798 continue;
799 }
800 };
801 let started = match stream_one_leg_standard(provider, &leg.model, req).await {
802 Ok(s) => s,
803 Err(e) => {
804 failures.push(LegFailure {
805 provider: leg.provider.clone(),
806 model: leg.model.clone(),
807 message: e.to_string(),
808 });
809 continue;
810 }
811 };
812 let mut stream = Box::pin(started);
813 match tokio::time::timeout(t.first_chunk, stream.next()).await {
814 Err(_) => {
815 failures.push(LegFailure {
816 provider: leg.provider.clone(),
817 model: leg.model.clone(),
818 message: "first-chunk timeout".into(),
819 });
820 }
821 Ok(Some(Ok(first))) => {
822 let rest = futures::stream::once(async move { Ok(first) }).chain(stream);
823 return Ok(CommittedStream {
824 provider: leg.provider.clone(),
825 model: leg.model.clone(),
826 stream: rest.boxed(),
827 });
828 }
829 Ok(Some(Err(e))) => {
830 failures.push(LegFailure {
831 provider: leg.provider.clone(),
832 model: leg.model.clone(),
833 message: e.to_string(),
834 });
835 }
836 Ok(None) => {
837 failures.push(LegFailure {
838 provider: leg.provider.clone(),
839 model: leg.model.clone(),
840 message: "empty stream".into(),
841 });
842 }
843 }
844 }
845 Err(GatewayError::AllLegsFailed {
846 route: route_name.to_string(),
847 failures,
848 })
849}
850
851pub async fn execute_streaming(
858 catalog: &Catalog,
859 route_name: &str,
860 legs: &[ChainLeg],
861 req: &ChatRequest,
862) -> Result<CommittedStream, GatewayError> {
863 execute_streaming_with_timeouts(catalog, route_name, legs, req, StreamTimeouts::default()).await
864}