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 status_retryable(status: reqwest::StatusCode) -> bool {
159 status.is_server_error()
160 || status == reqwest::StatusCode::TOO_MANY_REQUESTS
161 || status == reqwest::StatusCode::REQUEST_TIMEOUT
162 }
163
164 fn webc_retryable(we: &genai::webc::Error) -> bool {
166 match we {
167 genai::webc::Error::ResponseFailedStatus { status, .. } => status_retryable(*status),
168 genai::webc::Error::Reqwest(re) => re.is_timeout() || re.is_connect(),
169 _ => false,
170 }
171 }
172
173 match e {
174 genai::Error::WebModelCall { webc_error, .. } => webc_retryable(webc_error),
175 genai::Error::WebAdapterCall { webc_error, .. } => webc_retryable(webc_error),
176 genai::Error::HttpError { status, .. } => status_retryable(*status),
177 _ => false,
178 }
179}
180
181#[derive(Debug)]
183pub enum LegError {
184 Start(String),
186 MidStream(String),
188}
189
190impl std::fmt::Display for LegError {
191 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
192 match self {
193 LegError::Start(s) => write!(f, "start: {s}"),
194 LegError::MidStream(s) => write!(f, "mid-stream: {s}"),
195 }
196 }
197}
198
199fn map_stop_reason(sr: Option<&genai::chat::StopReason>) -> FinishReason {
201 match sr {
202 Some(genai::chat::StopReason::ToolCall(_)) => FinishReason::ToolCalls,
203 Some(genai::chat::StopReason::MaxTokens(_)) => FinishReason::Length,
204 _ => FinishReason::Stop,
205 }
206}
207
208#[derive(Default)]
210struct ToolIndexer {
211 ids: Vec<String>,
212}
213impl ToolIndexer {
214 fn index_of(&mut self, call_id: &str) -> (u32, bool) {
216 if let Some(pos) = self.ids.iter().position(|c| c == call_id) {
217 (pos as u32, false)
218 } else {
219 self.ids.push(call_id.to_string());
220 ((self.ids.len() - 1) as u32, true)
221 }
222 }
223}
224
225pub async fn stream_one_leg_standard(
228 provider: &Arc<Provider>,
229 model: &str,
230 req: &ChatRequest,
231) -> Result<impl Stream<Item = Result<StreamItem, LegError>>, LegError> {
232 let chat_req = to_genai_request(req);
233 let opts = to_genai_options(req)
234 .with_capture_usage(true)
235 .with_capture_tool_calls(true);
236
237 let resp = provider
238 .client
239 .exec_chat_stream(model.to_string(), chat_req, Some(&opts))
240 .await
241 .map_err(|e| LegError::Start(e.to_string()))?;
242
243 struct St<S> {
247 inner: S,
248 indexer: ToolIndexer,
249 input_tokens: u64,
250 output_tokens: u64,
251 }
252 let state = St {
253 inner: Box::pin(resp.stream),
254 indexer: ToolIndexer::default(),
255 input_tokens: 0,
256 output_tokens: 0,
257 };
258
259 let normalized = futures::stream::unfold(state, |mut st| async move {
260 loop {
261 match st.inner.next().await {
262 None => return None,
263 Some(Err(e)) => return Some((Err(LegError::MidStream(e.to_string())), st)),
264 Some(Ok(ev)) => match ev {
265 genai::chat::ChatStreamEvent::Chunk(c) => {
266 if !c.content.is_empty() {
267 return Some((Ok(StreamItem::Delta(c.content)), st));
268 }
269 }
270 genai::chat::ChatStreamEvent::ToolCallChunk(tc) => {
271 let call = tc.tool_call;
272 let (index, first) = st.indexer.index_of(&call.call_id);
273 let args = match &call.fn_arguments {
274 serde_json::Value::String(s) => s.clone(),
275 other => other.to_string(),
276 };
277 return Some((
278 Ok(StreamItem::ToolCallDelta {
279 index,
280 id: if first { Some(call.call_id) } else { None },
281 name: if first { Some(call.fn_name) } else { None },
282 args_fragment: args,
283 }),
284 st,
285 ));
286 }
287 genai::chat::ChatStreamEvent::End(end) => {
288 if let Some(u) = &end.captured_usage {
289 st.input_tokens = u.prompt_tokens.unwrap_or(0).max(0) as u64;
290 st.output_tokens = u.completion_tokens.unwrap_or(0).max(0) as u64;
291 }
292 let done = StreamItem::Done {
293 input_tokens: st.input_tokens,
294 output_tokens: st.output_tokens,
295 finish_reason: map_stop_reason(end.captured_stop_reason.as_ref()),
296 };
297 return Some((Ok(done), st));
298 }
299 _ => {} },
301 }
302 }
303 });
304
305 Ok(normalized)
306}
307
308#[cfg(test)]
309mod tests {
310 use super::*;
311
312 fn req(body: serde_json::Value) -> ChatRequest {
313 serde_json::from_value(body).unwrap()
314 }
315
316 #[test]
317 fn maps_tools_into_genai_request() {
318 let r = req(serde_json::json!({
319 "model": "m",
320 "messages": [{"role": "user", "content": "hi"}],
321 "tools": [{"type": "function", "function": {"name": "get_weather",
322 "description": "Lookup", "parameters": {"type": "object"}}}]
323 }));
324 let g = to_genai_request(&r);
325 let tools = g.tools.expect("tools mapped");
326 assert_eq!(tools.len(), 1);
327 assert_eq!(tools[0].name.to_string(), "get_weather");
328 assert_eq!(tools[0].description.as_deref(), Some("Lookup"));
329 }
330
331 #[test]
332 fn maps_assistant_tool_calls_and_tool_results() {
333 let r = req(serde_json::json!({
334 "model": "m",
335 "messages": [
336 {"role": "user", "content": "weather?"},
337 {"role": "assistant", "content": null, "tool_calls": [
338 {"id": "call_0", "type": "function", "function": {"name": "f", "arguments": "{\"c\":\"SF\"}"}}]},
339 {"role": "tool", "tool_call_id": "call_0", "content": "21C"}
340 ]
341 }));
342 let g = to_genai_request(&r);
343 assert_eq!(g.messages.len(), 3);
344 }
345
346 use crate::providers::genai_provider::{build_openai_compat_provider, OpenAiCompatConfig};
347 use crate::routing::stream::{FinishReason, StreamItem};
348 use futures::StreamExt;
349 use std::time::Duration;
350
351 #[tokio::test]
352 async fn standard_lane_streams_text_then_done() {
353 use wiremock::matchers::{method, path};
354 use wiremock::{Mock, MockServer, ResponseTemplate};
355
356 let sse = "data: {\"choices\":[{\"delta\":{\"content\":\"He\"}}]}\n\n\
358 data: {\"choices\":[{\"delta\":{\"content\":\"llo\"}}]}\n\n\
359 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":3,\"completion_tokens\":2,\"total_tokens\":5}}\n\n\
360 data: [DONE]\n\n";
361 let mock = MockServer::start().await;
362 Mock::given(method("POST"))
363 .and(path("/v1/chat/completions"))
364 .respond_with(
365 ResponseTemplate::new(200)
366 .insert_header("content-type", "text/event-stream")
367 .set_body_string(sse),
368 )
369 .mount(&mock)
370 .await;
371
372 let provider = Arc::new(
373 build_openai_compat_provider(
374 "oai",
375 OpenAiCompatConfig {
376 base_url: format!("{}/v1", mock.uri()),
377 api_key: "k".into(),
378 request_timeout: Duration::from_secs(5),
379 endpoint_override: None,
380 },
381 )
382 .unwrap(),
383 );
384
385 let req = req(
386 serde_json::json!({"model":"m","messages":[{"role":"user","content":"hi"}],"stream":true}),
387 );
388 let mut stream = std::pin::pin!(stream_one_leg_standard(&provider, "m", &req)
389 .await
390 .expect("stream starts"));
391 let mut items = Vec::new();
392 while let Some(it) = stream.next().await {
393 items.push(it.expect("no mid-stream error"));
394 }
395 assert!(items
396 .iter()
397 .any(|i| matches!(i, StreamItem::Delta(t) if t == "He")));
398 assert!(matches!(
399 items.last().unwrap(),
400 StreamItem::Done {
401 input_tokens: 3,
402 output_tokens: 2,
403 finish_reason: FinishReason::Stop
404 }
405 ));
406 }
407
408 #[tokio::test]
409 async fn execute_buffered_falls_back_on_midstream_failure() {
410 use crate::providers::Catalog;
411 use crate::routing::table::ChainLeg;
412 use wiremock::matchers::{method, path};
413 use wiremock::{Mock, MockServer, ResponseTemplate};
414
415 let bad = MockServer::start().await;
417 Mock::given(method("POST"))
418 .and(path("/v1/chat/completions"))
419 .respond_with(
420 ResponseTemplate::new(200)
421 .insert_header("content-type", "text/event-stream")
422 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"par\"}}]}\n\n"),
423 ) .mount(&bad)
425 .await;
426 let good = MockServer::start().await;
428 Mock::given(method("POST"))
429 .and(path("/v1/chat/completions"))
430 .respond_with(
431 ResponseTemplate::new(200)
432 .insert_header("content-type", "text/event-stream")
433 .set_body_string(
434 "data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}]}\n\n\
435 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
436 data: [DONE]\n\n",
437 ),
438 )
439 .mount(&good)
440 .await;
441
442 let catalog = Catalog::for_test(vec![
443 ("p1", format!("{}/v1", bad.uri())),
444 ("p2", format!("{}/v1", good.uri())),
445 ]);
446 let legs = vec![
447 ChainLeg {
448 provider: "p1".into(),
449 model: "m".into(),
450 ..Default::default()
451 },
452 ChainLeg {
453 provider: "p2".into(),
454 model: "m".into(),
455 ..Default::default()
456 },
457 ];
458 let r =
459 req(serde_json::json!({"model":"route","messages":[{"role":"user","content":"hi"}]}));
460 let c = execute_buffered(&catalog, "route", &legs, &r)
461 .await
462 .unwrap();
463 assert_eq!(c.content, "ok");
464 assert_eq!(c.provider, "p2");
465 }
466
467 #[tokio::test]
468 async fn execute_streaming_commits_first_leg_with_items() {
469 use crate::providers::Catalog;
470 use crate::routing::stream::StreamItem;
471 use crate::routing::table::ChainLeg;
472 use futures::StreamExt;
473 use wiremock::matchers::{method, path};
474 use wiremock::{Mock, MockServer, ResponseTemplate};
475
476 let bad = MockServer::start().await;
478 Mock::given(method("POST"))
479 .and(path("/v1/chat/completions"))
480 .respond_with(ResponseTemplate::new(500))
481 .mount(&bad)
482 .await;
483 let good = MockServer::start().await;
484 Mock::given(method("POST")).and(path("/v1/chat/completions"))
485 .respond_with(ResponseTemplate::new(200)
486 .insert_header("content-type", "text/event-stream")
487 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"go\"}}]}\n\n\
488 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
489 data: [DONE]\n\n"))
490 .mount(&good).await;
491
492 let catalog = Catalog::for_test(vec![
493 ("p1", format!("{}/v1", bad.uri())),
494 ("p2", format!("{}/v1", good.uri())),
495 ]);
496 let legs = vec![
497 ChainLeg {
498 provider: "p1".into(),
499 model: "m".into(),
500 ..Default::default()
501 },
502 ChainLeg {
503 provider: "p2".into(),
504 model: "m".into(),
505 ..Default::default()
506 },
507 ];
508 let r = req(
509 serde_json::json!({"model":"route","messages":[{"role":"user","content":"hi"}],"stream":true}),
510 );
511 let committed = execute_streaming(&catalog, "route", &legs, &r)
512 .await
513 .unwrap();
514 assert_eq!(committed.provider, "p2");
515 let mut stream = committed.stream;
516 let mut items = Vec::new();
517 while let Some(i) = stream.next().await {
518 items.push(i.unwrap());
519 }
520 assert!(items
521 .iter()
522 .any(|i| matches!(i, StreamItem::Delta(t) if t == "go")));
523 }
524
525 #[tokio::test]
526 async fn first_chunk_timeout_falls_back() {
527 use crate::providers::Catalog;
528 use crate::routing::table::ChainLeg;
529 use std::time::Duration;
530 use wiremock::matchers::{method, path};
531 use wiremock::{Mock, MockServer, ResponseTemplate};
532
533 let slow = 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_delay(Duration::from_millis(400))
548 .set_body_string("data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{}}\n\ndata: [DONE]\n\n"))
549 .mount(&slow).await;
550 let good = MockServer::start().await;
551 Mock::given(method("POST")).and(path("/v1/chat/completions"))
552 .respond_with(ResponseTemplate::new(200)
553 .insert_header("content-type", "text/event-stream")
554 .set_body_string("data: {\"choices\":[{\"delta\":{\"content\":\"ok\"}}]}\n\n\
555 data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":1}}\n\n\
556 data: [DONE]\n\n")).mount(&good).await;
557
558 let catalog = Catalog::for_test(vec![
559 ("p1", format!("{}/v1", slow.uri())),
560 ("p2", format!("{}/v1", good.uri())),
561 ]);
562 let legs = vec![
563 ChainLeg {
564 provider: "p1".into(),
565 model: "m".into(),
566 ..Default::default()
567 },
568 ChainLeg {
569 provider: "p2".into(),
570 model: "m".into(),
571 ..Default::default()
572 },
573 ];
574 let r =
575 req(serde_json::json!({"model":"route","messages":[{"role":"user","content":"hi"}]}));
576 let timeouts = StreamTimeouts {
577 first_chunk: Duration::from_millis(150),
578 idle: Duration::from_secs(5),
579 };
580 let c = execute_buffered_with_timeouts(&catalog, "route", &legs, &r, timeouts)
581 .await
582 .unwrap();
583 assert_eq!(c.provider, "p2");
584 }
585}
586
587async fn run_one_leg(
588 provider: &Arc<Provider>,
589 leg: &ChainLeg,
590 req: &ChatRequest,
591) -> Result<Completion, ResilienceError<genai::Error>> {
592 let client = provider.client.clone();
593 let model = leg.model.clone();
594 let chat_req = to_genai_request(req);
595 let opts = to_genai_options(req);
596
597 let resp = run_with_classifier(
598 move || {
599 let (client, model, chat_req, opts) = (
600 client.clone(),
601 model.clone(),
602 chat_req.clone(),
603 opts.clone(),
604 );
605 async move { client.exec_chat(model, chat_req, Some(&opts)).await }
606 },
607 provider.profile,
608 &provider.breaker,
609 provider.label,
610 is_genai_retryable,
611 )
612 .await?;
613
614 let content = resp.first_text().unwrap_or_default().to_string();
615 let usage = &resp.usage;
616 Ok(Completion {
617 provider: leg.provider.clone(),
618 model: leg.model.clone(),
619 content,
620 tool_calls: Vec::new(),
621 finish_reason: FinishReason::Stop,
622 input_tokens: usage.prompt_tokens.unwrap_or(0).max(0) as u64,
623 output_tokens: usage.completion_tokens.unwrap_or(0).max(0) as u64,
624 })
625}
626
627pub async fn execute_chain(
630 catalog: &Catalog,
631 route_name: &str,
632 legs: &[ChainLeg],
633 req: &ChatRequest,
634) -> Result<Completion, GatewayError> {
635 let mut failures: Vec<LegFailure> = Vec::new();
636 let mut all_circuit_open = true;
637 for leg in legs {
638 let provider = catalog.get(&leg.provider).ok_or_else(|| {
639 GatewayError::BadRequest(format!(
640 "route '{route_name}' references unbuilt provider '{}'",
641 leg.provider
642 ))
643 })?;
644 match run_one_leg(provider, leg, req).await {
645 Ok(c) => return Ok(c),
646 Err(ResilienceError::CircuitOpen { name }) => failures.push(LegFailure {
647 provider: leg.provider.clone(),
648 model: leg.model.clone(),
649 message: format!("circuit open: {name}"),
650 }),
651 Err(ResilienceError::Exhausted(e)) => {
652 all_circuit_open = false;
653 let retryable = is_genai_retryable(&e);
654 failures.push(LegFailure {
655 provider: leg.provider.clone(),
656 model: leg.model.clone(),
657 message: e.to_string(),
658 });
659 if !retryable {
660 break; }
662 }
663 }
664 }
665 if all_circuit_open && !failures.is_empty() {
666 return Err(GatewayError::AllCircuitsOpen(route_name.to_string()));
667 }
668 Err(GatewayError::AllLegsFailed {
669 route: route_name.to_string(),
670 failures,
671 })
672}
673
674#[derive(Debug, Clone, Copy)]
676pub struct StreamTimeouts {
677 pub first_chunk: std::time::Duration,
679 pub idle: std::time::Duration,
681}
682
683impl Default for StreamTimeouts {
684 fn default() -> Self {
685 Self {
686 first_chunk: std::time::Duration::from_secs(120),
687 idle: std::time::Duration::from_secs(60),
688 }
689 }
690}
691
692async fn buffer_one_leg_timed(
694 catalog: &Catalog,
695 leg: &ChainLeg,
696 req: &ChatRequest,
697 t: StreamTimeouts,
698) -> Result<Completion, LegError> {
699 let provider = catalog
700 .get(&leg.provider)
701 .ok_or_else(|| LegError::Start(format!("unbuilt provider '{}'", leg.provider)))?;
702 let stream = stream_one_leg_standard(provider, &leg.model, req).await?;
703 let mut stream = std::pin::pin!(stream);
704 let mut acc = Accumulator::default();
705 let mut first = true;
706 loop {
707 let budget = if first { t.first_chunk } else { t.idle };
708 match tokio::time::timeout(budget, stream.next()).await {
709 Err(_) => {
710 return Err(if first {
711 LegError::Start("first-chunk timeout".into())
712 } else {
713 LegError::MidStream("idle timeout".into())
714 })
715 }
716 Ok(None) => break,
717 Ok(Some(item)) => {
718 acc.push(item?);
719 first = false;
720 }
721 }
722 }
723 if !acc.got_done {
724 return Err(LegError::MidStream("stream ended before completion".into()));
725 }
726 Ok(Completion {
727 provider: leg.provider.clone(),
728 model: leg.model.clone(),
729 content: acc.content,
730 tool_calls: acc.tool_calls,
731 finish_reason: acc.finish_reason,
732 input_tokens: acc.input_tokens,
733 output_tokens: acc.output_tokens,
734 })
735}
736
737pub async fn execute_buffered_with_timeouts(
742 catalog: &Catalog,
743 route_name: &str,
744 legs: &[ChainLeg],
745 req: &ChatRequest,
746 t: StreamTimeouts,
747) -> Result<Completion, GatewayError> {
748 let mut failures: Vec<LegFailure> = Vec::new();
749 for leg in legs {
750 match buffer_one_leg_timed(catalog, leg, req, t).await {
751 Ok(c) => return Ok(c),
752 Err(e) => failures.push(LegFailure {
753 provider: leg.provider.clone(),
754 model: leg.model.clone(),
755 message: e.to_string(),
756 }),
757 }
758 }
759 Err(GatewayError::AllLegsFailed {
760 route: route_name.to_string(),
761 failures,
762 })
763}
764
765pub async fn execute_buffered(
775 catalog: &Catalog,
776 route_name: &str,
777 legs: &[ChainLeg],
778 req: &ChatRequest,
779) -> Result<Completion, GatewayError> {
780 execute_buffered_with_timeouts(catalog, route_name, legs, req, StreamTimeouts::default()).await
781}
782
783pub struct CommittedStream {
786 pub provider: String,
787 pub model: String,
788 pub stream: BoxStream<'static, Result<StreamItem, LegError>>,
789}
790
791impl CommittedStream {
792 pub fn single(
795 provider: String,
796 model: String,
797 stream: impl Stream<Item = Result<StreamItem, LegError>> + Send + 'static,
798 ) -> Self {
799 Self {
800 provider,
801 model,
802 stream: stream.boxed(),
803 }
804 }
805}
806
807pub async fn execute_streaming_with_timeouts(
812 catalog: &Catalog,
813 route_name: &str,
814 legs: &[ChainLeg],
815 req: &ChatRequest,
816 t: StreamTimeouts,
817) -> Result<CommittedStream, GatewayError> {
818 let mut failures: Vec<LegFailure> = Vec::new();
819
820 for leg in legs {
821 let provider = match catalog.get(&leg.provider) {
822 Some(p) => p,
823 None => {
824 failures.push(LegFailure {
825 provider: leg.provider.clone(),
826 model: leg.model.clone(),
827 message: format!("unbuilt provider '{}'", leg.provider),
828 });
829 continue;
830 }
831 };
832 let started = match stream_one_leg_standard(provider, &leg.model, req).await {
833 Ok(s) => s,
834 Err(e) => {
835 failures.push(LegFailure {
836 provider: leg.provider.clone(),
837 model: leg.model.clone(),
838 message: e.to_string(),
839 });
840 continue;
841 }
842 };
843 let mut stream = Box::pin(started);
844 match tokio::time::timeout(t.first_chunk, stream.next()).await {
845 Err(_) => {
846 failures.push(LegFailure {
847 provider: leg.provider.clone(),
848 model: leg.model.clone(),
849 message: "first-chunk timeout".into(),
850 });
851 }
852 Ok(Some(Ok(first))) => {
853 let rest = futures::stream::once(async move { Ok(first) }).chain(stream);
854 return Ok(CommittedStream {
855 provider: leg.provider.clone(),
856 model: leg.model.clone(),
857 stream: rest.boxed(),
858 });
859 }
860 Ok(Some(Err(e))) => {
861 failures.push(LegFailure {
862 provider: leg.provider.clone(),
863 model: leg.model.clone(),
864 message: e.to_string(),
865 });
866 }
867 Ok(None) => {
868 failures.push(LegFailure {
869 provider: leg.provider.clone(),
870 model: leg.model.clone(),
871 message: "empty stream".into(),
872 });
873 }
874 }
875 }
876 Err(GatewayError::AllLegsFailed {
877 route: route_name.to_string(),
878 failures,
879 })
880}
881
882pub async fn execute_streaming(
889 catalog: &Catalog,
890 route_name: &str,
891 legs: &[ChainLeg],
892 req: &ChatRequest,
893) -> Result<CommittedStream, GatewayError> {
894 execute_streaming_with_timeouts(catalog, route_name, legs, req, StreamTimeouts::default()).await
895}