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 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
84/// `{type:function, function:{name, description, parameters}}` -> genai `Tool`.
85fn 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
98/// `{id, function:{name, arguments}}` -> genai `ToolCall`. `arguments` is an
99/// OpenAI JSON STRING; parse to Value (fall back to a string Value if not JSON).
100fn 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
139/// True if a genai error is worth advancing the chain for (transient/5xx/timeout).
140/// 4xx (auth, bad request) are NOT retryable — abort the chain immediately.
141///
142/// Implementation note: we use STRUCTURED matching on the genai error enum rather than
143/// string inspection. The previous implementation checked for `" 500"`, `" 502"`, etc.
144/// (space before digits), but genai's `webc::Error::ResponseFailedStatus` Display format
145/// is `"Request failed with status code '503 ...'` — digits are preceded by a single-quote,
146/// not a space — so the old checks never matched and 5xx errors were treated as non-retryable,
147/// causing `execute_chain` to break out of the fallback loop instead of advancing to the
148/// next leg.
149///
150/// Structured approach: match the two web-call wrapper variants
151/// (`WebModelCall` / `WebAdapterCall`) that carry a `genai::webc::Error`, then match
152/// `ResponseFailedStatus { status, .. }` and call `status.is_server_error()` on the
153/// typed `reqwest::StatusCode`. Timeout and connection errors are detected via
154/// `webc::Error::Reqwest(e)` with `e.is_timeout() || e.is_connect()`.
155/// `genai::Error::HttpError { status, .. }` is also matched for completeness.
156fn is_genai_retryable(e: &genai::Error) -> bool {
157    /// Align with `resilience::is_retryable_reqwest`: 5xx, 429, request timeout.
158    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    /// Check whether a `genai::webc::Error` represents a transient/server error.
165    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/// Errors from a single streaming leg. `Start` failures are fallback-eligible.
182#[derive(Debug)]
183pub enum LegError {
184    /// Failed before/at stream start (connection, 5xx, etc.).
185    Start(String),
186    /// Failed after items began flowing.
187    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
199/// Map a genai `StopReason` to our `FinishReason`.
200fn 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/// Stateful call_id -> stable 0-based index assignment for tool-call streaming.
209#[derive(Default)]
210struct ToolIndexer {
211    ids: Vec<String>,
212}
213impl ToolIndexer {
214    /// Returns (index, is_first_time_seen).
215    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
225/// Standard-lane per-leg primitive: open a genai stream, normalize to `StreamItem`.
226/// Outer `Err(LegError::Start)` = stream could not begin (fallback-eligible).
227pub 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    // State threaded across the whole stream. Use `unfold` to own the state
244    // cleanly (a stateful `flat_map` closure can also work; unfold avoids
245    // borrow-checker friction with FnMut captures).
246    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                    _ => {} // Start / ReasoningChunk / ThoughtSignatureChunk ignored; loop for next
300                },
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        // OpenAI-style SSE: two content deltas, a finish, then [DONE].
357        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        // Leg 1: stream starts then the body is cut after one delta (no Done) -> mid-stream failure.
416        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            ) // no [DONE]/usage
424            .mount(&bad)
425            .await;
426        // Leg 2: clean stream.
427        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        // Leg 1 fails to start (500) -> fall back. Leg 2 streams.
477        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        // Leg 1: the entire response (including response head) is delayed 400ms.
534        // With a 150ms first-chunk timeout applied to stream.next(), leg 1 is
535        // abandoned before any item arrives and the chain falls back to leg 2.
536        //
537        // Note: wiremock's set_delay delays BEFORE the response head is sent,
538        // so exec_chat_stream's connection .await will block for 400ms. The
539        // stream returned by stream_one_leg_standard won't yield its first item
540        // within that window. The first-chunk timeout in buffer_one_leg_timed
541        // wraps stream.next() — it fires at 150ms, classifying the failure as
542        // LegError::Start("first-chunk timeout"), which is fallback-eligible.
543        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
627/// Walk legs in order. Retryable failure (or open breaker) advances; the first
628/// non-retryable failure aborts. Returns `AllLegsFailed` if every leg fails.
629pub 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; // non-retryable: abort the chain
661                }
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/// Time bounds for a streaming leg.
675#[derive(Debug, Clone, Copy)]
676pub struct StreamTimeouts {
677    /// Max time to the first item (time-to-first-token).
678    pub first_chunk: std::time::Duration,
679    /// Max gap between successive items.
680    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
692/// Like `buffer_one_leg` but bounded by first-chunk and inter-chunk idle timeouts.
693async 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
737/// Buffered executor with explicit first-chunk and idle timeouts. Walk legs,
738/// fully buffering each. Any failure (start, timeout, or mid-stream) advances
739/// to the next leg. Nothing is ever flushed to the client, so full-chain
740/// fallback is preserved.
741pub 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
765/// Buffered (non-streaming) executor: walk legs, fully buffering each. Any
766/// failure (start or mid-stream) advances to the next leg. Nothing is ever
767/// flushed to the client, so full-chain fallback is preserved.
768///
769/// Unlike [`execute_chain`], the connection phase here is NOT yet wrapped in the
770/// per-leg circuit-breaker/retry — a tracked follow-up. The streaming path trades
771/// that resilience plumbing for guaranteed whole-response fallback.
772///
773/// Delegates to [`execute_buffered_with_timeouts`] with [`StreamTimeouts::default`].
774pub 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
783/// A committed streaming leg: the winning provider/model plus the remaining
784/// item stream (first item already re-prepended).
785pub struct CommittedStream {
786    pub provider: String,
787    pub model: String,
788    pub stream: BoxStream<'static, Result<StreamItem, LegError>>,
789}
790
791impl CommittedStream {
792    /// Wrap an already-obtained stream as a single committed leg (used by the
793    /// native Vertex lane, which has no standard-lane fallback).
794    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
807/// Streaming executor with explicit first-chunk timeout. For each leg, start
808/// the stream and peek the first item within `t.first_chunk`. First leg to
809/// yield an item is committed; failures (including timeout) before the first
810/// item fall back. After commitment there is no fallback.
811pub 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
882/// Streaming executor: for each leg, start the stream and peek the first item.
883/// First leg to yield an item is committed; failures before the first item fall
884/// back. After commitment there is no fallback (caller surfaces mid-stream
885/// errors as SSE error events).
886///
887/// Delegates to [`execute_streaming_with_timeouts`] with [`StreamTimeouts::default`].
888pub 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}