Skip to main content

synapse/routing/
executor.rs

1//! Standard-lane executor: walk a route's legs with per-leg breaker + retry.
2
3use 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/// Normalised result of one completed LLM call.
17#[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
28/// Build a genai `ChatRequest` from the gateway request: messages (incl. tool
29/// calls / tool results) plus tool definitions. `tool_choice` is not expressible
30/// on genai 0.6 `ChatRequest` and is dropped on this lane (documented).
31fn 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
83/// `{type:function, function:{name, description, parameters}}` -> genai `Tool`.
84fn 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
97/// `{id, function:{name, arguments}}` -> genai `ToolCall`. `arguments` is an
98/// OpenAI JSON STRING; parse to Value (fall back to a string Value if not JSON).
99fn 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
121/// True if a genai error is worth advancing the chain for (transient/5xx/timeout).
122/// 4xx (auth, bad request) are NOT retryable — abort the chain immediately.
123///
124/// Implementation note: we use STRUCTURED matching on the genai error enum rather than
125/// string inspection. The previous implementation checked for `" 500"`, `" 502"`, etc.
126/// (space before digits), but genai's `webc::Error::ResponseFailedStatus` Display format
127/// is `"Request failed with status code '503 ...'` — digits are preceded by a single-quote,
128/// not a space — so the old checks never matched and 5xx errors were treated as non-retryable,
129/// causing `execute_chain` to break out of the fallback loop instead of advancing to the
130/// next leg.
131///
132/// Structured approach: match the two web-call wrapper variants
133/// (`WebModelCall` / `WebAdapterCall`) that carry a `genai::webc::Error`, then match
134/// `ResponseFailedStatus { status, .. }` and call `status.is_server_error()` on the
135/// typed `reqwest::StatusCode`. Timeout and connection errors are detected via
136/// `webc::Error::Reqwest(e)` with `e.is_timeout() || e.is_connect()`.
137/// `genai::Error::HttpError { status, .. }` is also matched for completeness.
138fn is_genai_retryable(e: &genai::Error) -> bool {
139    /// Check whether a `genai::webc::Error` represents a transient/server error.
140    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/// Errors from a single streaming leg. `Start` failures are fallback-eligible.
157#[derive(Debug)]
158pub enum LegError {
159    /// Failed before/at stream start (connection, 5xx, etc.).
160    Start(String),
161    /// Failed after items began flowing.
162    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
174/// Map a genai `StopReason` to our `FinishReason`.
175fn 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/// Stateful call_id -> stable 0-based index assignment for tool-call streaming.
184#[derive(Default)]
185struct ToolIndexer {
186    ids: Vec<String>,
187}
188impl ToolIndexer {
189    /// Returns (index, is_first_time_seen).
190    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
200/// Standard-lane per-leg primitive: open a genai stream, normalize to `StreamItem`.
201/// Outer `Err(LegError::Start)` = stream could not begin (fallback-eligible).
202pub 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    // State threaded across the whole stream. Use `unfold` to own the state
219    // cleanly (a stateful `flat_map` closure can also work; unfold avoids
220    // borrow-checker friction with FnMut captures).
221    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                    _ => {} // Start / ReasoningChunk / ThoughtSignatureChunk ignored; loop for next
275                },
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        // OpenAI-style SSE: two content deltas, a finish, then [DONE].
332        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        // Leg 1: stream starts then the body is cut after one delta (no Done) -> mid-stream failure.
391        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            ) // no [DONE]/usage
399            .mount(&bad)
400            .await;
401        // Leg 2: clean stream.
402        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        // Leg 1 fails to start (500) -> fall back. Leg 2 streams.
450        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        // Leg 1: the entire response (including response head) is delayed 400ms.
505        // With a 150ms first-chunk timeout applied to stream.next(), leg 1 is
506        // abandoned before any item arrives and the chain falls back to leg 2.
507        //
508        // Note: wiremock's set_delay delays BEFORE the response head is sent,
509        // so exec_chat_stream's connection .await will block for 400ms. The
510        // stream returned by stream_one_leg_standard won't yield its first item
511        // within that window. The first-chunk timeout in buffer_one_leg_timed
512        // wraps stream.next() — it fires at 150ms, classifying the failure as
513        // LegError::Start("first-chunk timeout"), which is fallback-eligible.
514        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
596/// Walk legs in order. Retryable failure (or open breaker) advances; the first
597/// non-retryable failure aborts. Returns `AllLegsFailed` if every leg fails.
598pub 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; // non-retryable: abort the chain
630                }
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/// Time bounds for a streaming leg.
644#[derive(Debug, Clone, Copy)]
645pub struct StreamTimeouts {
646    /// Max time to the first item (time-to-first-token).
647    pub first_chunk: std::time::Duration,
648    /// Max gap between successive items.
649    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
661/// Like `buffer_one_leg` but bounded by first-chunk and inter-chunk idle timeouts.
662async 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
706/// Buffered executor with explicit first-chunk and idle timeouts. Walk legs,
707/// fully buffering each. Any failure (start, timeout, or mid-stream) advances
708/// to the next leg. Nothing is ever flushed to the client, so full-chain
709/// fallback is preserved.
710pub 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
734/// Buffered (non-streaming) executor: walk legs, fully buffering each. Any
735/// failure (start or mid-stream) advances to the next leg. Nothing is ever
736/// flushed to the client, so full-chain fallback is preserved.
737///
738/// Unlike [`execute_chain`], the connection phase here is NOT yet wrapped in the
739/// per-leg circuit-breaker/retry — a tracked follow-up. The streaming path trades
740/// that resilience plumbing for guaranteed whole-response fallback.
741///
742/// Delegates to [`execute_buffered_with_timeouts`] with [`StreamTimeouts::default`].
743pub 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
752/// A committed streaming leg: the winning provider/model plus the remaining
753/// item stream (first item already re-prepended).
754pub struct CommittedStream {
755    pub provider: String,
756    pub model: String,
757    pub stream: BoxStream<'static, Result<StreamItem, LegError>>,
758}
759
760impl CommittedStream {
761    /// Wrap an already-obtained stream as a single committed leg (used by the
762    /// native Vertex lane, which has no standard-lane fallback).
763    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
776/// Streaming executor with explicit first-chunk timeout. For each leg, start
777/// the stream and peek the first item within `t.first_chunk`. First leg to
778/// yield an item is committed; failures (including timeout) before the first
779/// item fall back. After commitment there is no fallback.
780pub 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
851/// Streaming executor: for each leg, start the stream and peek the first item.
852/// First leg to yield an item is committed; failures before the first item fall
853/// back. After commitment there is no fallback (caller surfaces mid-stream
854/// errors as SSE error events).
855///
856/// Delegates to [`execute_streaming_with_timeouts`] with [`StreamTimeouts::default`].
857pub 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}