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