diff --git a/Cargo.lock b/Cargo.lock index 9bcbf4e21..31561fde2 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4294,6 +4294,7 @@ dependencies = [ "hyper-tls", "hyper-util", "metrics", + "mime", "notify", "openai-reassembler", "opentelemetry", diff --git a/dwctl/src/sync/onwards_config/mod.rs b/dwctl/src/sync/onwards_config/mod.rs index 25fc1cc5c..7eee2994a 100644 --- a/dwctl/src/sync/onwards_config/mod.rs +++ b/dwctl/src/sync/onwards_config/mod.rs @@ -938,6 +938,7 @@ fn convert_composite_to_target_spec( .and_then(|n| usize::try_from(n).ok().filter(|&v| v >= 1)), backoff, max_total_backoff_ms, + stream_continuation: None, }) } else { None @@ -1202,6 +1203,7 @@ fn convert_to_config_file( .and_then(|n| usize::try_from(n).ok().filter(|&v| v >= 1)), backoff, max_total_backoff_ms, + stream_continuation: None, }) } else { None diff --git a/onwards/CHANGELOG.md b/onwards/CHANGELOG.md index 9cfdef39b..be501aca1 100644 --- a/onwards/CHANGELOG.md +++ b/onwards/CHANGELOG.md @@ -23,6 +23,14 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- add best-effort prefix continuation for eligible Completions, Chat Completions, and Responses streams + +### Changed + +- **Rust API migration (next minor):** `FallbackConfig` now includes `stream_continuation`; callers using exhaustive struct literals must initialize it, usually with `None` + ## [0.35.6](https://github.com/doublewordai/onwards/compare/v0.35.5...v0.35.6) - 2026-07-21 ### Fixed diff --git a/onwards/Cargo.toml b/onwards/Cargo.toml index e0bd76af9..7ad7013d1 100644 --- a/onwards/Cargo.toml +++ b/onwards/Cargo.toml @@ -67,6 +67,7 @@ bon = "3.6.5" subtle = "2.6.1" axum-prometheus = "0.10.0" metrics = "0.24" +mime = "0.3" governor = "0.10.1" rand = "0.9" uuid = { version = "1", features = ["v4"] } diff --git a/onwards/README.md b/onwards/README.md index af6a6d907..0129494ca 100644 --- a/onwards/README.md +++ b/onwards/README.md @@ -51,12 +51,91 @@ curl -X POST http://localhost:3000/v1/chat/completions \ - Rate limiting and concurrency limiting (per-target and per-key) - Load balancing with weighted random and priority strategies - Automatic failover across multiple providers +- Best-effort continuation of interrupted text-generation streams - Strict mode for request validation and error standardization - Response sanitization for OpenAI schema compliance - Prometheus metrics - Custom response headers - Optional multi-step Open Responses orchestration loop (`multi-step` feature) +## Stream Continuation + +Stream continuation is an opt-in fallback mode for eligible streaming requests +to `/v1/completions`, `/v1/chat/completions`, and `/v1/responses`. If an +upstream stops before a terminal event, Onwards sends the text already emitted +to another provider selected by the pool, which may be the same provider. + +Configure it inside the pool's `fallback` settings: + +```json +{ + "targets": { + "gpt-4": { + "fallback": { + "enabled": true, + "on_status": [429, 5], + "stream_continuation": { + "enabled": true, + "endpoints": [ + "/v1/completions", + "/v1/chat/completions", + "/v1/responses" + ], + "max_attempts": 1, + "max_buffered_bytes": 1048576, + "idle_timeout_ms": 30000 + } + }, + "providers": [ + { "url": "https://primary.example.com", "onwards_key": "sk-primary" }, + { "url": "https://backup.example.com", "onwards_key": "sk-backup" } + ] + } + } +} +``` + +| Option | Default | Description | +|--------|---------|-------------| +| `enabled` | `false` | Enables continuation, subject to `fallback.enabled` | +| `endpoints` | `[]` | Exact supported request paths eligible for continuation | +| `max_attempts` | `1` | Fresh continuation-attempt budget after the response starts | +| `max_buffered_bytes` | `1048576` | Maximum emitted-text bytes retained for a continuation request | +| `idle_timeout_ms` | `null` | Maximum idle time between upstream events; `null` disables the timeout | + +`fallback.max_attempts` controls retries before response headers are committed. +`stream_continuation.max_attempts` is a separate budget after streaming starts. +Continuation reuses the pool's provider selection, request headers, status +matching, local rate-limit behavior, and backoff configuration. Its cumulative +backoff budget is also fresh. + +Only narrow, single-output text requests are eligible. Completions require a +string prompt, one choice, no echo, and `stream: true`. Chat Completions require +ordinary message content, one choice, no logprobs, and `stream: true`. Responses +require string or message-only input, plain-text foreground output, no storage, +and `stream: true`. Tool calls, reasoning, structured output, multi-choice or +multi-item output, and provider-specific guided generation are excluded. +Unsupported requests pass through normally. If an eligible stream later emits +an unknown event shape, continuation is disabled rather than guessing how to +splice it. + +Completions append emitted text to the original prompt. Chat Completions append +one assistant message and suppress a repeated assistant-role chunk. Responses +append an assistant `output_text` message while preserving response and item +identities, sequence numbers, and the cumulative terminal text snapshot. + +Eligible upstream responses must be successful, identity-encoded SSE using the +`text/event-stream` media type. Onwards requests `Accept-Encoding: identity` and +rejects an encoded or non-SSE continuation response instead of splicing it. + +This is prompt-prefix continuation, not native token-offset resume. A provider +may repeat text, add a preamble, or diverge. The emitted prefix is buffered only +in process memory and is lost on restart. Exceeding the buffer cap disables +continuation for that request. Exhausted retries can leave the client with a +partial stream and no synthetic terminal event. Usage comes from the provider +that supplies the terminal stream and is not aggregate usage or billing across +attempts. + ## Multi-Step Open Responses The `multi-step` Cargo feature adds an orchestration loop that drives diff --git a/onwards/docs/src/load-balancing.md b/onwards/docs/src/load-balancing.md index a6f197431..eb47b90a2 100644 --- a/onwards/docs/src/load-balancing.md +++ b/onwards/docs/src/load-balancing.md @@ -37,6 +37,11 @@ Controls automatic retry on other providers when requests fail: | `enabled` | bool | `false` | Master switch for fallback | | `on_status` | int[] | -- | Status codes that trigger fallback (supports wildcards) | | `on_rate_limit` | bool | `false` | Fallback when hitting local rate limits | +| `with_replacement` | bool | `false` | For `weighted_random`, allow a provider to be selected again | +| `max_attempts` | int? | provider count | Maximum pre-response failover attempts | +| `backoff` | object? | `null` | Delay between attempts; no delay when omitted | +| `max_total_backoff_ms` | int? | `null` | Maximum cumulative time spent sleeping between attempts | +| `stream_continuation` | object? | `null` | Best-effort continuation for eligible interrupted text-generation streams | Status code wildcards: @@ -46,6 +51,105 @@ Status code wildcards: When fallback triggers, the next provider is selected based on strategy (weighted random resamples from remaining pool; priority uses definition order). +### Stream continuation + +Stream continuation is opt-in for eligible `POST /v1/completions`, +`POST /v1/chat/completions`, and `POST /v1/responses` streams. When an upstream +stops before a terminal event, Onwards sends the text already emitted to a newly +selected provider, which may be the same provider. The prefix is kept in memory +for the lifetime of the request; it is not persisted and is lost if the process +restarts. An interruption or exhausted continuation budget can leave the client +with a partial stream, and the proxy does not synthesize a terminal event. + +Configure it inside `fallback`: + +```json +{ + "targets": { + "gpt-4": { + "fallback": { + "enabled": true, + "on_status": [429, 5], + "on_rate_limit": true, + "stream_continuation": { + "enabled": true, + "endpoints": [ + "/v1/completions", + "/v1/chat/completions", + "/v1/responses" + ], + "max_attempts": 1, + "max_buffered_bytes": 1048576, + "idle_timeout_ms": 30000 + } + }, + "providers": [ + { "url": "https://primary.example.com", "onwards_key": "sk-primary" }, + { "url": "https://backup.example.com", "onwards_key": "sk-backup" } + ] + } + } +} +``` + +| Option | Type | Default | Description | +|--------|------|---------|-------------| +| `enabled` | bool | `false` | Enables continuation, subject to the parent `fallback.enabled` switch | +| `endpoints` | string[] | `[]` | Exact request paths eligible for continuation; supported values are `/v1/completions`, `/v1/chat/completions`, and `/v1/responses` | +| `max_attempts` | int | `1` | Fresh post-commit continuation-attempt budget | +| `max_buffered_bytes` | int | `1048576` | Maximum generated-text bytes retained for a continuation request | +| `idle_timeout_ms` | int? | `null` | Maximum idle time between upstream stream events; `null` disables this timeout | + +`fallback.max_attempts` controls failover before the response starts and its +headers are committed. `stream_continuation.max_attempts` is a separate, fresh +budget after a successful stream has committed its response. Continuation +reuses the parent fallback status matching, local rate-limit behavior, provider +selection, request/header handling, and backoff settings, including +`max_total_backoff_ms`. Stream continuation also starts a fresh cumulative +backoff budget; the parent `max_total_backoff_ms` limit applies independently +to the post-commit continuation attempts. + +Only narrow, single-output text shapes are supported. Completions require a +string `prompt`, `stream: true`, `n` omitted or `1`, and `echo` omitted or +`false`. Chat Completions require ordinary message content, `stream: true`, one +choice, and no logprobs. Responses require string or message-only input, +`stream: true`, plain-text output, and foreground, non-stored execution. Tool +calls, reasoning, structured output, multi-choice or multi-item output, and +provider-specific guided generation controls are excluded. Unsupported requests +pass through normally without continuation. If an initially eligible stream +later emits an unrecognized shape, continuation is disabled rather than guessing +how to splice it. + +Continuation is protocol-specific. For Completions, Onwards appends the emitted +text to the original prompt. For Chat Completions, it appends one assistant +message containing the emitted text and suppresses a repeated assistant-role +chunk. For Responses, it appends an assistant `output_text` message, suppresses +repeated lifecycle setup events, and keeps response IDs, item IDs, and sequence +numbers continuous. Terminal Responses snapshots are rewritten with the full +cumulative text. + +The upstream response must be successful, use the exact `text/event-stream` +media type (case-insensitive; parameters are allowed), and have no content +encoding or `Content-Encoding: identity`. Onwards sends +`Accept-Encoding: identity` for eligible requests and continues only unencoded +responses. An encoded eligible stream is left untouched; an encoded or non-SSE +continuation response is rejected rather than spliced into the client stream. +In strict mode, body sanitization cannot be guaranteed when an eligible encoded +response ignores the identity request because preserving its encoded +representation requires forwarding it unwrapped. + +This is prefix-based continuation, not native token-offset resume. Chat +Completions and Responses do not expose a continuation cursor: the emitted text +is supplied as prior assistant output in a new request. The next model call may +add a preamble, repeat text, or diverge from the interrupted generation. + +The buffer cap is measured in UTF-8 bytes. Once appending a recognized text +chunk would exceed the cap, continuation is disabled for that request and the +overflow text is not retained. `idle_timeout_ms` applies between upstream +events and treats an idle stream as interrupted. Final usage is forwarded from +the provider that supplies the terminal stream; it is not aggregate usage or +aggregate billing across the initial and continuation providers. + ## Pool-level options Settings that apply to the entire alias: diff --git a/onwards/src/handlers.rs b/onwards/src/handlers.rs index 72c017da1..eda289206 100644 --- a/onwards/src/handlers.rs +++ b/onwards/src/handlers.rs @@ -7,21 +7,26 @@ use crate::AppState; use crate::auth; use crate::client::HttpClient; use crate::errors::{ErrorResponseBody, OnwardsErrorResponse}; +use crate::load_balancer::ProviderPool; use crate::models::ListModelResponse; -use crate::sse::SseBufferedStream; +use crate::sse::{CheckedSseStream, SseStreamError}; use crate::target::{ConcurrencyGuard, RoutingAction, Target}; use axum::{ Json, extract::Request, extract::State, http::{ - HeaderMap, HeaderName, HeaderValue, StatusCode, Uri, - header::{CONTENT_LENGTH, TRANSFER_ENCODING}, + HeaderMap, HeaderName, HeaderValue, Method, StatusCode, Uri, + header::{ + ACCEPT_ENCODING, CONTENT_ENCODING, CONTENT_LENGTH, CONTENT_RANGE, CONTENT_TYPE, + TRAILER, TRANSFER_ENCODING, + }, }, response::{IntoResponse, Response}, }; use opentelemetry::propagation::{Extractor, Injector, TextMapPropagator}; use serde_json::map::Entry; +use std::collections::HashMap; use tracing::{Instrument, debug, error, instrument, trace, warn}; /// Adapter to extract W3C trace context from an axum HeaderMap. @@ -90,7 +95,7 @@ enum SseEventKind { } /// Classify a single SSE event (a complete frame already reassembled by -/// [`SseBufferedStream`]). +/// the checked SSE framer). /// /// Some providers open a `200 OK` stream and send the error as the /// first `data:` frame (`data: {"error":{"code":429,...}}`) rather than content. @@ -99,24 +104,14 @@ enum SseEventKind { /// multi-line, and non-spaced (`data:{...}`) framing are all handled. A frame /// with no `data:` field is a comment/keep-alive. fn classify_sse_event(chunk: &[u8]) -> SseEventKind { - let Ok(text) = std::str::from_utf8(chunk) else { - // Non-UTF-8 can't be an error envelope we understand — forward as-is. + let data = match crate::sse::parse_sse_event(chunk) { + crate::sse::ParsedSseEvent::Comment => return SseEventKind::Comment, + crate::sse::ParsedSseEvent::Data { data, .. } => data, + crate::sse::ParsedSseEvent::Invalid => return SseEventKind::Data, + }; + let Ok(data) = std::str::from_utf8(&data) else { return SseEventKind::Data; }; - let mut data = String::new(); - let mut has_data = false; - for line in text.lines() { - if let Some(rest) = line.strip_prefix("data:") { - has_data = true; - if !data.is_empty() { - data.push('\n'); - } - data.push_str(rest.strip_prefix(' ').unwrap_or(rest)); - } - } - if !has_data { - return SseEventKind::Comment; - } if data.trim() == "[DONE]" { return SseEventKind::Data; } @@ -192,7 +187,7 @@ impl Drop for InflightGuard { /// outlives the handler. struct GuardedStream { inner: S, - _guard: ConcurrencyGuard, + _guard: Option, _inflight_guard: InflightGuard, } @@ -230,10 +225,12 @@ pub(crate) struct ResolvedTrust(pub(crate) bool); /// answered", not "the request succeeded" — check the status code for that. /// Absent when no upstream produced a response (auth/validation rejections, /// exhausted fallbacks, gateway-generated errors). For streaming tool loops the -/// response is returned before follow-up iterations run, so it names the -/// provider of the initial iteration. Integrators (e.g. request-logging -/// middleware) can read it to attribute traffic to a concrete upstream without -/// re-deriving the routing decision. +/// response is returned before follow-up iterations run, and for stream +/// continuation the response extensions are committed before the body can +/// select another provider, so it names the provider of the initial stream in +/// both cases. Integrators (e.g. request-logging middleware) can read it to +/// attribute traffic to a concrete upstream without re-deriving the routing +/// decision. #[derive(Clone, Debug, PartialEq, Eq)] pub struct ServedBy { /// Full URL of the upstream target that served the request. @@ -351,12 +348,301 @@ fn filter_headers_for_upstream(headers: &mut HeaderMap, target: &Target) { headers.insert("x-forwarded-proto", "https".parse().unwrap()); } -/// The main handler responsible for forwarding requests to targets -/// TODO(fergus): Better error messages beyond raw status codes. +#[derive(Clone)] +pub(crate) struct UpstreamRequestMetadata { + method: Method, + path_and_query: String, + headers: HeaderMap, + trace_context: Option, + pool_trusted: bool, +} + +#[derive(Clone)] +pub(crate) struct CanonicalRequestPath(pub(crate) &'static str); + +#[derive(Clone, Copy)] +pub(crate) struct ContinuationEligibilityPath(pub(crate) &'static str); + +#[derive(Clone, Copy)] +pub(crate) struct PreserveEncodedStream; + +#[derive(Clone, Copy)] +pub(crate) struct SseBufferLimit(pub(crate) usize); + +#[derive(Clone)] +struct CompositeResponsePolicy { + trusted: bool, + sanitize_response: bool, + response_headers: Option>, +} + +impl CompositeResponsePolicy { + fn for_pool(pool: &ProviderPool) -> Self { + let mut response_headers = pool + .providers() + .first() + .and_then(|provider| provider.target.response_headers.clone()) + .unwrap_or_default(); + for provider in pool.providers().iter().skip(1) { + let provider_headers = provider.target.response_headers.as_ref(); + response_headers.retain(|name, value| { + provider_headers.and_then(|headers| headers.get(name)) == Some(value) + }); + } + + Self { + trusted: pool + .providers() + .iter() + .all(|provider| provider.target.trusted.unwrap_or_else(|| pool.is_trusted())), + sanitize_response: pool + .providers() + .iter() + .any(|provider| provider.target.sanitize_response), + response_headers: (!response_headers.is_empty()).then_some(response_headers), + } + } +} + +fn canonical_generation_path(path: &str) -> &str { + match path { + "/completions" | "/v1/completions" => "/v1/completions", + "/chat/completions" | "/v1/chat/completions" => "/v1/chat/completions", + "/responses" | "/v1/responses" => "/v1/responses", + _ => path, + } +} + +impl UpstreamRequestMetadata { + fn new(method: Method, path_and_query: String, headers: HeaderMap, pool_trusted: bool) -> Self { + Self { + method, + path_and_query, + headers, + trace_context: None, + pool_trusted, + } + } + + fn with_current_trace_context(mut self) -> Self { + use tracing_opentelemetry::OpenTelemetrySpanExt; + self.trace_context = Some(tracing::Span::current().context()); + self + } + + pub(crate) fn for_child_span(&self, span: &tracing::Span) -> Self { + use tracing_opentelemetry::OpenTelemetrySpanExt; + if let Some(parent) = self.trace_context.clone() { + let _ = span.set_parent(parent); + } + let mut child = self.clone(); + child.trace_context = Some(span.context()); + child + } + + fn with_identity_encoding(mut self) -> Self { + self.headers + .insert(ACCEPT_ENCODING, HeaderValue::from_static("identity")); + self + } +} + +fn clear_composite_representation_headers(headers: &mut HeaderMap) { + for name in [ + CONTENT_LENGTH, + CONTENT_ENCODING, + CONTENT_RANGE, + TRANSFER_ENCODING, + TRAILER, + ] { + headers.remove(name); + } + for name in [ + "connection", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "te", + "upgrade", + "content-md5", + "digest", + "content-digest", + "repr-digest", + "representation-digest", + "etag", + "accept-ranges", + "last-modified", + ] { + headers.remove(name); + } +} + +pub(crate) fn build_upstream_request( + target: &Target, + metadata: &UpstreamRequestMetadata, + body: bytes::Bytes, +) -> Result<(Request, String), OnwardsErrorResponse> { + let request_path = metadata + .path_and_query + .strip_prefix('/') + .unwrap_or(&metadata.path_and_query); + let target_path = target.url.path().trim_end_matches('/'); + let path_to_join = if !target_path.is_empty() && target_path != "/" { + let target_path_no_slash = &target_path[1..]; + if let Some(rest) = request_path.strip_prefix(target_path_no_slash) { + if rest.is_empty() || rest.starts_with('/') { + rest.strip_prefix('/').unwrap_or(rest) + } else { + request_path + } + } else { + request_path + } + } else { + request_path + }; + let upstream_uri = target + .url + .join(path_to_join) + .map_err(|_| OnwardsErrorResponse::internal())? + .to_string(); + let upstream_uri_parsed = Uri::try_from(&upstream_uri).map_err(|_| { + error!("Invalid URI: {}", upstream_uri); + OnwardsErrorResponse::internal() + })?; + + let mut headers = metadata.headers.clone(); + if let Some(host) = upstream_uri_parsed.host() { + let host_value = if let Some(port) = upstream_uri_parsed.port_u16() { + format!("{host}:{port}") + } else { + host.to_string() + }; + headers.insert("host", host_value.parse().unwrap()); + } + headers.insert( + CONTENT_LENGTH, + body.len() + .to_string() + .parse() + .expect("Content-Length should be valid"), + ); + headers.remove(TRANSFER_ENCODING); + filter_headers_for_upstream(&mut headers, target); + if resolve_trace_propagation(target, metadata.pool_trusted) { + use tracing_opentelemetry::OpenTelemetrySpanExt; + let trace_context = metadata + .trace_context + .clone() + .unwrap_or_else(|| tracing::Span::current().context()); + let propagator = opentelemetry_sdk::propagation::TraceContextPropagator::new(); + propagator.inject_context(&trace_context, &mut HeaderInjector(&mut headers)); + } else { + withhold_trace_context(&mut headers); + } + + let request = Request::builder() + .method(metadata.method.clone()) + .uri(upstream_uri_parsed) + .body(axum::body::Body::from(body)) + .map_err(|_| OnwardsErrorResponse::internal())?; + let (mut parts, body) = request.into_parts(); + parts.headers = headers; + Ok((Request::from_parts(parts, body), upstream_uri)) +} + +trait ContinuationBodyFactory { + const ENABLED: bool; + + fn wrap( + initial_body: axum::body::Body, + initial_guard: ConcurrencyGuard, + continuation: crate::stream_continuation::StreamContinuation, + config: crate::target::StreamContinuationConfig, + pool: ProviderPool, + http_client: T, + request_metadata: UpstreamRequestMetadata, + ) -> axum::body::Body; +} + +struct ContinuationDisabled; + +impl ContinuationBodyFactory for ContinuationDisabled { + const ENABLED: bool = false; + + fn wrap( + _initial_body: axum::body::Body, + _initial_guard: ConcurrencyGuard, + _continuation: crate::stream_continuation::StreamContinuation, + _config: crate::target::StreamContinuationConfig, + _pool: ProviderPool, + _http_client: T, + _request_metadata: UpstreamRequestMetadata, + ) -> axum::body::Body { + unreachable!("legacy handler never activates stream continuation") + } +} + +struct ContinuationEnabled; + +impl ContinuationBodyFactory for ContinuationEnabled +where + T: HttpClient + Send + 'static, +{ + const ENABLED: bool = true; + + fn wrap( + initial_body: axum::body::Body, + initial_guard: ConcurrencyGuard, + continuation: crate::stream_continuation::StreamContinuation, + config: crate::target::StreamContinuationConfig, + pool: ProviderPool, + http_client: T, + request_metadata: UpstreamRequestMetadata, + ) -> axum::body::Body { + crate::stream_continuation::wrap_generation_stream( + initial_body, + initial_guard, + continuation, + config, + pool, + http_client, + request_metadata, + ) + } +} + +/// Forward a request using the legacy public handler contract. +/// +/// Stream continuation is installed by the supported router entry points, +/// whose client bounds allow the client to live inside the downstream body. pub async fn target_message_handler( + state: State>, + req: axum::extract::Request, +) -> Result { + target_message_handler_core::(state, req).await +} + +pub(crate) async fn target_message_handler_with_continuation( + state: State>, + req: axum::extract::Request, +) -> Result +where + T: HttpClient + Send + 'static, +{ + target_message_handler_core::(state, req).await +} + +/// The shared handler core responsible for forwarding requests to targets. +/// TODO(fergus): Better error messages beyond raw status codes. +async fn target_message_handler_core( State(state): State>, mut req: axum::extract::Request, -) -> Result { +) -> Result +where + T: HttpClient, + C: ContinuationBodyFactory, +{ // Create the tracing span BEFORE entering it so that set_parent() correctly // updates the OTel parent context. With #[instrument], the OTel span is started // on the first enter() which happens before the function body runs — making @@ -382,6 +668,8 @@ pub async fn target_message_handler( async move { + let mut http_client = Some(state.http_client); + // Track inflight requests for observability. The guard is moved into GuardedStream // on the success path so the gauge stays incremented for the full lifetime of // streaming response bodies. @@ -654,10 +942,59 @@ pub async fn target_message_handler( .map(|v| v.as_str()) .unwrap_or(req.uri().path()) .to_string(); + let canonical_request_path = req + .extensions() + .get::() + .map(|path| path.0) + .unwrap_or(req.uri().path()) + .to_string(); + let continuation_eligibility_path = req + .extensions() + .get::() + .map(|path| path.0) + .unwrap_or(&canonical_request_path) + .to_string(); + let continuation_protocol_path = canonical_generation_path(req.uri().path()).to_string(); + let is_completion_path = canonical_request_path == "/v1/completions"; // Prepare original headers and method for potential retries let original_headers = req.headers().clone(); let method = req.method().clone(); + let mut stream_continuation = if C::ENABLED { + pool.fallback() + .filter(|fallback| fallback.enabled) + .and_then(|fallback| fallback.stream_continuation.as_ref()) + .and_then(|config| { + crate::stream_continuation::StreamContinuation::from_request_with_resolved_model( + &continuation_eligibility_path, + &continuation_protocol_path, + &method, + &body_bytes, + config, + Some(&model_name), + ) + .map(|continuation| (continuation, config.clone())) + }) + } else { + None + }; + let composite_response_policy = stream_continuation + .as_ref() + .map(|_| CompositeResponsePolicy::for_pool(&pool)); + let request_metadata = UpstreamRequestMetadata::new( + method.clone(), + path_and_query.clone(), + original_headers.clone(), + pool.is_trusted(), + ) + .with_current_trace_context(); + let request_metadata = if stream_continuation.is_some() + || (state.targets.strict_mode && is_completion_path) + { + request_metadata.with_identity_encoding() + } else { + request_metadata + }; // Track last error for fallback scenarios let mut last_error: Option = None; @@ -671,6 +1008,7 @@ pub async fn target_message_handler( let mut total_backoff_ms: u64 = 0; let pool_max_attempts = pool.fallback_max_attempts(); for (_idx, target, connection_guard) in pool.select_iter() { + let mut connection_guard = Some(connection_guard); any_attempted = true; attempt_number += 1; @@ -788,98 +1126,15 @@ pub async fn target_message_handler( }; } - // Build the upstream URI for this target - let request_path = path_and_query.strip_prefix('/').unwrap_or(&path_and_query); - let target_path = target.url.path().trim_end_matches('/'); - - let path_to_join = if !target_path.is_empty() && target_path != "/" { - let target_path_no_slash = &target_path[1..]; - if let Some(rest) = request_path.strip_prefix(target_path_no_slash) { - if rest.is_empty() || rest.starts_with('/') { - rest.strip_prefix('/').unwrap_or(rest) - } else { - request_path - } - } else { - request_path - } - } else { - request_path + let (attempt_req, upstream_uri) = match build_upstream_request( + target, + &request_metadata, + attempt_body, + ) { + Ok(request) => request, + Err(error) => return LoopAction::Done(Err(error)), }; - let upstream_uri = match target.url.join(path_to_join) { - Ok(url) => url.to_string(), - Err(_) => return LoopAction::Done(Err(OnwardsErrorResponse::internal())), - }; - let upstream_uri_parsed = match Uri::try_from(&upstream_uri) { - Ok(uri) => uri, - Err(_) => { - error!("Invalid URI: {}", upstream_uri); - return LoopAction::Done(Err(OnwardsErrorResponse::internal())); - } - }; - - // Build request for this provider - let mut attempt_headers = original_headers.clone(); - - // Update host header - if let Some(host) = upstream_uri_parsed.host() { - let host_value = if let Some(port) = upstream_uri_parsed.port_u16() { - format!("{host}:{port}") - } else { - host.to_string() - }; - attempt_headers.insert("host", host_value.parse().unwrap()); - } - - // Set Content-Length and remove Transfer-Encoding - attempt_headers.insert( - CONTENT_LENGTH, - attempt_body - .len() - .to_string() - .parse() - .expect("Content-Length should be valid"), - ); - attempt_headers.remove(TRANSFER_ENCODING); - - // Filter headers for upstream forwarding - filter_headers_for_upstream(&mut attempt_headers, target); - - // Apply W3C trace-context policy for this upstream, gated on the - // per-target propagate_trace_context flag (defaults to the resolved - // trusted value). - // - // - Enabled: inject the current span context, propagating the trace to - // upstreams that participate in our distributed tracing fabric - // (typically self-hosted, marked trusted). inject_context overwrites - // any inbound traceparent/tracestate. - // - Disabled: strip any *inbound* traceparent/tracestate so they are - // not forwarded to an untrusted third party. These headers are not - // in HEADERS_TO_STRIP (they're handled here, next to the inject - // decision), so without this branch an inbound trace context from - // the caller would pass straight through, defeating the opt-out and - // letting the third party re-emit our trace IDs on its own outbound - // calls. - if resolve_trace_propagation(target, pool.is_trusted()) { - use tracing_opentelemetry::OpenTelemetrySpanExt; - let ctx = tracing::Span::current().context(); - let propagator = opentelemetry_sdk::propagation::TraceContextPropagator::new(); - propagator.inject_context(&ctx, &mut HeaderInjector(&mut attempt_headers)); - } else { - withhold_trace_context(&mut attempt_headers); - } - - // Build the request - let attempt_req = axum::extract::Request::builder() - .method(method.clone()) - .uri(upstream_uri_parsed) - .body(axum::body::Body::from(attempt_body)) - .unwrap(); - let (mut parts, body) = attempt_req.into_parts(); - parts.headers = attempt_headers; - let attempt_req = axum::extract::Request::from_parts(parts, body); - trace!( "Outgoing request to provider:\n URI: {}", upstream_uri @@ -899,7 +1154,13 @@ pub async fn target_message_handler( let request_result = async { if let Some(timeout_secs) = target.request_timeout_secs { let timeout_duration = std::time::Duration::from_secs(timeout_secs); - match tokio::time::timeout(timeout_duration, state.http_client.request(attempt_req)) + match tokio::time::timeout( + timeout_duration, + http_client + .as_ref() + .expect("HTTP client remains available before response success") + .request(attempt_req), + ) .await { Err(_) => { @@ -914,7 +1175,12 @@ pub async fn target_message_handler( } } else { // No timeout configured - state.http_client.request(attempt_req).await.map_err(UpstreamOutcome::Error) + http_client + .as_ref() + .expect("HTTP client remains available before response success") + .request(attempt_req) + .await + .map_err(UpstreamOutcome::Error) } } .instrument(upstream_span.clone()) @@ -1015,7 +1281,19 @@ pub async fn target_message_handler( .and_then(|v| v.to_str().ok()) .unwrap_or("") .to_string(); - let is_sse = content_type.contains("text/event-stream"); + let is_sse = crate::stream_continuation::is_event_stream(response.headers()); + let is_identity_encoded = + crate::stream_continuation::has_identity_content_encoding(response.headers()); + let preserve_encoded_stream = (200..300).contains(&status) + && is_sse + && !is_identity_encoded + && (stream_continuation.is_some() + || (state.targets.strict_mode && is_completion_path)); + if preserve_encoded_stream { + response + .extensions_mut() + .insert(PreserveEncodedStream); + } // Some upstreams return HTTP 200 while embedding the // real error in the body — `{"error":{"code":429,...}}` for a unary @@ -1031,7 +1309,7 @@ pub async fn target_message_handler( // the status-based fallback above. // // Gated to strict_mode: the only mode where the body is *already* buffered - // (unary) or routed through `SseBufferedStream` (streaming) by the strict + // (unary) or routed through checked SSE framing (streaming) by the strict // sanitizer downstream. So this adds no new buffering and no change to // streaming behaviour — it just inspects the body a little earlier so a // retry stays possible. Pure-passthrough / sanitize-without-transform @@ -1050,9 +1328,14 @@ pub async fn target_message_handler( /// 502. Keyed on stream termination, never a time budget, so a /// valid-but-slow stream is forwarded (Clean), not retried. EmptyBody, + /// Checked SSE framing failed before a content frame. The original + /// response is forwarded with its typed body error so neither the + /// pre-response fallback loop nor continuation can treat it as EOF. + Framing, } let scan: Scan2xx = if (200..300).contains(&status) && state.targets.strict_mode + && !preserve_encoded_stream { if is_sse { // Peek the leading SSE events to find the first *real* frame, @@ -1070,7 +1353,14 @@ pub async fn target_message_handler( const SSE_PEEK_MAX_EVENTS: usize = 4; let (parts, body) = response.into_parts(); - let mut events = SseBufferedStream::new(body.into_data_stream()); + let body_stream = body.into_data_stream(); + let mut events = match stream_continuation.as_ref() { + Some(_) => CheckedSseStream::with_max_buffer_size( + body_stream, + crate::stream_continuation::event_buffer_size(), + ), + None => CheckedSseStream::new(body_stream), + }; // `peeked` / `embedded` live outside the peek future so a partial // peek survives a budget timeout (consumed frames are not lost). let mut peeked = Vec::new(); @@ -1080,6 +1370,7 @@ pub async fn target_message_handler( // before any content — a *terminal* empty, safe to retry. let mut saw_data = false; let mut stream_ended = false; + let mut framing_failed = false; let peek = async { for _ in 0..SSE_PEEK_MAX_EVENTS { match events.next().await { @@ -1099,7 +1390,10 @@ pub async fn target_message_handler( SseEventKind::Comment => peeked.push(Ok(chunk)), }, Some(Err(e)) => { - stream_ended = true; + match &e { + SseStreamError::Source(_) => stream_ended = true, + SseStreamError::Framing(_) => framing_failed = true, + } peeked.push(Err(e)); break; } @@ -1121,6 +1415,8 @@ pub async fn target_message_handler( response = Response::from_parts(parts, axum::body::Body::from_stream(rest)); if let Some(status) = embedded { Scan2xx::Embedded(status) + } else if framing_failed { + Scan2xx::Framing } else if !timed_out && stream_ended && !saw_data { // Stream closed/errored before any content frame: nothing was // forwarded, so retrying is safe. A *timeout* (stream still open, @@ -1239,15 +1535,63 @@ pub async fn target_message_handler( record_response_status(503); return LoopAction::Done(Err(OnwardsErrorResponse::service_unavailable())); } + Scan2xx::Framing => { + warn!( + http_status = status, + upstream = %target.url, + "Upstream returned malformed SSE framing; forwarding the checked body error" + ); + } Scan2xx::Clean => {} } + let mut composite_content_type = None; + if is_sse + && is_identity_encoded + && (200..300).contains(&status) + && let Some((continuation, config)) = stream_continuation.take() + { + let (mut parts, body) = response.into_parts(); + composite_content_type = parts.headers.get(CONTENT_TYPE).cloned(); + clear_composite_representation_headers(&mut parts.headers); + let event_buffer_limit = crate::stream_continuation::event_buffer_size(); + let body = C::wrap( + body, + connection_guard + .take() + .expect("initial provider guard moved once"), + continuation, + config, + pool.clone(), + http_client + .take() + .expect("HTTP client moved once into composite response"), + request_metadata.clone(), + ); + response = Response::from_parts(parts, body); + response + .extensions_mut() + .insert(SseBufferLimit(event_buffer_limit)); + } + let response_sanitize = composite_content_type + .as_ref() + .and(composite_response_policy.as_ref()) + .map_or(target.sanitize_response, |policy| policy.sanitize_response); + let response_trusted = composite_content_type + .as_ref() + .and(composite_response_policy.as_ref()) + .map_or_else( + || target.trusted.unwrap_or_else(|| pool.is_trusted()), + |policy| policy.trusted, + ); + // Determine if SSE buffering is needed for non-strict sanitization // Note: Strict mode handlers apply their own buffering before their sanitizers, // so we skip buffering here to avoid double-wrapping let needs_sse_buffering = !state.targets.strict_mode && state.response_transform_fn.is_some() - && target.sanitize_response + && response_sanitize + && !preserve_encoded_stream && (200..300).contains(&status); // Wrap SSE streams with buffering to ensure complete events (delimited by \n\n). @@ -1255,9 +1599,16 @@ pub async fn target_message_handler( // Providers may send partial chunks that split events across network packets. if is_sse && needs_sse_buffering { debug!("Wrapping SSE response with buffered stream for non-strict sanitization"); + let buffer_limit = response + .extensions() + .get::() + .map(|limit| limit.0); let (parts, body) = response.into_parts(); let byte_stream = body.into_data_stream(); - let buffered = SseBufferedStream::new(byte_stream); + let buffered = match buffer_limit { + Some(limit) => CheckedSseStream::with_max_buffer_size(byte_stream, limit), + None => CheckedSseStream::new(byte_stream), + }; let new_body = axum::body::Body::from_stream(buffered); response = Response::from_parts(parts, new_body); } @@ -1266,9 +1617,10 @@ pub async fn target_message_handler( // Per-target opt-in via sanitize_response flag, only for 2xx responses // Skip if strict mode is enabled - strict handlers do their own sanitization if let Some(ref transform_fn) = state.response_transform_fn - && target.sanitize_response + && response_sanitize && (200..300).contains(&status) && !state.targets.strict_mode + && !preserve_encoded_stream { debug!( "Attempting response sanitization for status {}, path {}", @@ -1295,7 +1647,12 @@ pub async fn target_message_handler( match chunk_result { Ok(chunk) => { // Sanitize this chunk - match sanitizer.sanitize_streaming(&chunk) { + let sanitized = if is_completion_path { + sanitizer.sanitize_completion_streaming(&chunk) + } else { + sanitizer.sanitize_streaming(&chunk) + }; + match sanitized { Ok(Some(sanitized)) => Ok::<_, std::io::Error>(sanitized), Ok(None) => Ok(chunk), Err(e) => { @@ -1387,15 +1744,21 @@ pub async fn target_message_handler( if let Some(ref header_name) = state.response_id_header && crate::response_id::path_supports_id_override(&path_and_query) && (200..300).contains(&status) - { - if let Some(override_id) = + && let Some(override_id) = crate::response_id::extract_override_id(&original_headers, header_name) - { - crate::response_id::patch_response_body_id(&mut response, override_id).await; - } + { + crate::response_id::patch_response_body_id(&mut response, override_id).await; } - // Add custom response headers + // Composite streams can span providers, so only retain response headers + // whose values are identical for every provider in the pool. + let response_headers = if composite_content_type.is_some() { + composite_response_policy + .as_ref() + .and_then(|policy| policy.response_headers.as_ref()) + } else { + response_headers.as_ref() + }; if let Some(headers) = response_headers { for (key, value) in headers.iter() { if let (Ok(header_name), Ok(header_value)) = @@ -1411,6 +1774,11 @@ pub async fn target_message_handler( ); } + if let Some(content_type) = composite_content_type { + clear_composite_representation_headers(response.headers_mut()); + response.headers_mut().insert(CONTENT_TYPE, content_type); + } + record_response_status(response.status().as_u16()); debug!( "Returning response with status {}, content-length: {:?}, strict_mode: {}", @@ -1418,10 +1786,9 @@ pub async fn target_message_handler( response.headers().get(CONTENT_LENGTH), state.targets.strict_mode ); - let resolved_trust = target.trusted.unwrap_or_else(|| pool.is_trusted()); response .extensions_mut() - .insert(ResolvedTrust(resolved_trust)); + .insert(ResolvedTrust(response_trusted)); response.extensions_mut().insert(ServedBy { url: target.url.to_string(), onwards_model: target.onwards_model.clone(), @@ -1554,6 +1921,76 @@ pub async fn models( mod tests { use super::*; + #[derive(Debug)] + struct NonCloneHttpClient; + + #[async_trait::async_trait] + impl HttpClient for NonCloneHttpClient { + async fn request( + &self, + _req: axum::extract::Request, + ) -> Result> { + unreachable!("compile-time API compatibility test") + } + } + + #[test] + fn target_message_handler_preserves_literal_http_client_bound() { + fn assert_literal_bound() { + let _handler = target_message_handler::; + } + + assert_literal_bound::(); + } + + #[test] + fn continuation_attempt_metadata_injects_attempt_span_id() { + use opentelemetry::trace::{TraceContextExt, TracerProvider as _}; + use tracing_opentelemetry::OpenTelemetrySpanExt; + use tracing_subscriber::prelude::*; + + let provider = opentelemetry_sdk::trace::SdkTracerProvider::builder().build(); + let tracer = provider.tracer("continuation-attempt-metadata-test"); + let subscriber = + tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer)); + let _subscriber_guard = tracing::subscriber::set_default(subscriber); + + let mut inbound_headers = HeaderMap::new(); + inbound_headers.insert( + "traceparent", + "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01" + .parse() + .unwrap(), + ); + let propagator = opentelemetry_sdk::propagation::TraceContextPropagator::new(); + let parent_context = propagator.extract(&HeaderExtractor(&inbound_headers)); + let metadata = UpstreamRequestMetadata { + method: Method::POST, + path_and_query: "/v1/completions".to_string(), + headers: HeaderMap::new(), + trace_context: Some(parent_context), + pool_trusted: false, + }; + let attempt_span = tracing::info_span!("test.continuation_attempt"); + let attempt_metadata = metadata.for_child_span(&attempt_span); + let attempt_context = attempt_span.context(); + let attempt_span_context = attempt_context.span().span_context().clone(); + let target = target_with_trace_flags(Some(false), Some(true)); + + let (request, _) = + build_upstream_request(&target, &attempt_metadata, bytes::Bytes::new()).unwrap(); + let traceparent = request + .headers() + .get("traceparent") + .unwrap() + .to_str() + .unwrap(); + let fields = traceparent.split('-').collect::>(); + + assert_eq!(fields[1], attempt_span_context.trace_id().to_string()); + assert_eq!(fields[2], attempt_span_context.span_id().to_string()); + } + #[test] fn embedded_error_status_detects_provider_envelope() { // Whole-body error envelope (the 200-with-error pattern). @@ -1640,6 +2077,11 @@ mod tests { Data ); assert_eq!(classify_sse_event(b"data: [DONE]\n\n"), Data); + assert_eq!(classify_sse_event(b"\xef\xbb\xbfdata:[DONE]\r\r"), Data); + assert_eq!( + classify_sse_event(b"\xef\xbb\xbfdata:{\"error\":\rdata:{\"code\":429}}\r\r"), + Error(429) + ); // Comment / keep-alive frames carry no `data:` field. assert_eq!(classify_sse_event(b": keep-alive\n\n"), Comment); diff --git a/onwards/src/lib.rs b/onwards/src/lib.rs index 798a373d2..422b7fc93 100644 --- a/onwards/src/lib.rs +++ b/onwards/src/lib.rs @@ -58,6 +58,7 @@ pub mod response_id; pub mod response_loop; pub mod response_sanitizer; pub mod sse; +pub mod stream_continuation; #[cfg(feature = "multi-step")] pub mod streaming; pub mod strict; @@ -67,7 +68,7 @@ pub mod traits; use client::{HttpClient, HyperClient}; pub use handlers::ServedBy; -use handlers::{models as models_handler, target_message_handler}; +use handlers::{models as models_handler, target_message_handler_with_continuation}; use models::ExtractedModel; #[cfg(feature = "multi-step")] pub use response_loop::{LoopConfig, LoopError, UpstreamTarget, run_response_loop}; @@ -523,7 +524,7 @@ pub fn build_router(state: AppSta Router::new() .route("/models", get(models_handler)) .route("/v1/models", get(models_handler)) - .route("/{*path}", any(target_message_handler)) + .route("/{*path}", any(target_message_handler_with_continuation)) // The wildcard handler buffers the body itself (bounded by // `state.body_limit`), but raise the extractor-level default too so // any extractor-based route added later shares the same limit. @@ -652,6 +653,33 @@ pub mod test_utils { custom_headers: Arc>>, } + #[derive(Clone)] + pub enum MockStreamEvent { + Data(String), + Bytes(Vec), + Error(String), + Pending, + } + + #[derive(Clone)] + pub struct MockStreamingResponse { + pub status: StatusCode, + pub content_type: Option, + pub headers: Vec<(String, String)>, + pub events: Vec, + } + + impl MockStreamingResponse { + pub fn sse(status: StatusCode, events: Vec) -> Self { + Self { + status, + content_type: Some("text/event-stream".to_string()), + headers: Vec::new(), + events, + } + } + } + #[derive(Debug, Clone)] pub struct MockRequest { pub method: String, @@ -734,6 +762,59 @@ pub mod test_utils { } } + pub fn new_streaming_response_sequence(responses: Vec) -> Self { + let counter = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + Self { + requests: Arc::new(Mutex::new(Vec::new())), + custom_headers: Arc::new(Mutex::new(Vec::new())), + response_builder: Arc::new(move || { + use axum::body::Body; + use std::collections::VecDeque; + use std::task::Poll; + + let idx = counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed); + let response = responses + .get(idx) + .cloned() + .unwrap_or(MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("text/event-stream".to_string()), + headers: Vec::new(), + events: Vec::new(), + }); + let mut events = VecDeque::from(response.events); + let stream = futures_util::stream::poll_fn(move |_cx| match events.front() { + Some(MockStreamEvent::Pending) => Poll::Pending, + Some(_) => match events.pop_front().expect("front checked above") { + MockStreamEvent::Data(data) => { + Poll::Ready(Some(Ok::<_, std::io::Error>(data.into_bytes()))) + } + MockStreamEvent::Bytes(data) => { + Poll::Ready(Some(Ok::<_, std::io::Error>(data))) + } + MockStreamEvent::Error(message) => { + Poll::Ready(Some(Err(std::io::Error::other(message)))) + } + MockStreamEvent::Pending => unreachable!("handled above"), + }, + None => Poll::Ready(None), + }); + + let mut builder = axum::response::Response::builder() + .status(response.status) + .header("cache-control", "no-cache") + .header("connection", "keep-alive"); + if let Some(content_type) = response.content_type { + builder = builder.header("content-type", content_type); + } + for (name, value) in response.headers { + builder = builder.header(name, value); + } + builder.body(Body::from_stream(stream)).unwrap() + }), + } + } + pub fn get_requests(&self) -> Vec { self.requests.lock().unwrap().clone() } @@ -928,13 +1009,20 @@ pub mod test_utils { mod tests { use super::*; use crate::load_balancer::{Provider, ProviderPool}; - use crate::target::{ConcurrencyLimiter, Target, Targets}; + use crate::target::{ + ConcurrencyLimiter, FallbackConfig, LoadBalanceStrategy, OpenResponsesConfig, + StreamContinuationConfig, Target, Targets, + }; + use axum::body::Body; use axum::http::StatusCode; use axum_test::TestServer; use dashmap::DashMap; + use futures_util::StreamExt; + use http_body_util::BodyExt; use serde_json::json; use std::sync::Arc; - use test_utils::MockHttpClient; + use test_utils::{MockHttpClient, MockStreamEvent, MockStreamingResponse}; + use tower::ServiceExt; /// Helper to create a single-provider pool from a target fn pool(target: Target) -> ProviderPool { @@ -1196,6 +1284,2091 @@ mod tests { } } + fn stream_continuation_targets( + alias: &str, + provider_models: &[Option<&str>], + fallback: FallbackConfig, + ) -> target::Targets { + stream_continuation_targets_with_options(alias, provider_models, fallback, false, false) + } + + fn stream_continuation_targets_with_options( + alias: &str, + provider_models: &[Option<&str>], + fallback: FallbackConfig, + strict_mode: bool, + sanitize_response: bool, + ) -> target::Targets { + let providers = provider_models + .iter() + .enumerate() + .map(|(i, model)| { + let builder = Target::builder() + .url(format!("https://stream-{i}.example.com/").parse().unwrap()) + .sanitize_response(sanitize_response); + let target = match model { + Some(model) => builder.onwards_model((*model).to_string()).build(), + None => builder.build(), + }; + Provider::new(target, 1) + }) + .collect(); + let pool = ProviderPool::with_config( + providers, + None, + None, + None, + Some(fallback), + LoadBalanceStrategy::Priority, + false, + Vec::new(), + ); + let targets_map = Arc::new(DashMap::new()); + targets_map.insert(alias.to_string(), pool); + target::Targets { + targets: targets_map, + key_rate_limiters: Arc::new(DashMap::new()), + key_concurrency_limiters: Arc::new(DashMap::new()), + key_labels: Arc::new(DashMap::new()), + strict_mode, + http_pool_config: None, + } + } + + fn stream_continuation_targets_from_targets( + alias: &str, + targets: Vec, + fallback: FallbackConfig, + strict_mode: bool, + pool_trusted: bool, + ) -> target::Targets { + let providers = targets + .into_iter() + .map(|target| Provider::new(target, 1)) + .collect(); + let pool = ProviderPool::with_config( + providers, + None, + None, + None, + Some(fallback), + LoadBalanceStrategy::Priority, + pool_trusted, + Vec::new(), + ); + let targets_map = Arc::new(DashMap::new()); + targets_map.insert(alias.to_string(), pool); + target::Targets { + targets: targets_map, + key_rate_limiters: Arc::new(DashMap::new()), + key_concurrency_limiters: Arc::new(DashMap::new()), + key_labels: Arc::new(DashMap::new()), + strict_mode, + http_pool_config: None, + } + } + + fn strict_stream_continuation_router( + targets: target::Targets, + mock: MockHttpClient, + ) -> axum::Router { + axum::Router::new().nest( + "/v1", + crate::strict::build_strict_router(AppState::with_client(targets, mock)), + ) + } + + fn stream_continuation_fallback( + enabled: bool, + pre_response_attempts: usize, + continuation_attempts: usize, + idle_timeout_ms: Option, + endpoints: Vec<&str>, + ) -> FallbackConfig { + FallbackConfig { + enabled, + on_status: vec![502], + max_attempts: Some(pre_response_attempts), + stream_continuation: Some(StreamContinuationConfig { + enabled: true, + endpoints: endpoints.into_iter().map(str::to_string).collect(), + max_attempts: continuation_attempts, + max_buffered_bytes: 1024, + idle_timeout_ms, + }), + ..Default::default() + } + } + + fn completion_event(id: &str, model: &str, text: &str, finish_reason: &str) -> String { + format!( + "data: {{\"id\":\"{id}\",\"object\":\"text_completion\",\"created\":1,\"model\":\"{model}\",\"choices\":[{{\"index\":0,\"text\":{text:?},\"finish_reason\":{finish_reason}}}]}}\n\n" + ) + } + + fn completion_sse_values(body: &str) -> Vec { + body.lines() + .filter_map(|line| line.strip_prefix("data: ")) + .filter(|data| *data != "[DONE]") + .map(|data| serde_json::from_str(data).unwrap()) + .collect() + } + + fn gzip_bytes(data: &[u8]) -> Vec { + use flate2::Compression; + use flate2::write::GzEncoder; + use std::io::Write; + + let mut encoder = GzEncoder::new(Vec::new(), Compression::fast()); + encoder.write_all(data).unwrap(); + encoder.finish().unwrap() + } + + fn brotli_bytes(data: &[u8]) -> Vec { + use std::io::Write; + + let mut compressed = Vec::new(); + { + let mut writer = brotli::CompressorWriter::new(&mut compressed, 4096, 4, 22); + writer.write_all(data).unwrap(); + } + compressed + } + + #[tokio::test] + async fn stream_continuation_interrupted_completion_continues_from_exact_emitted_prefix() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let second = completion_event("cmpl-second", "second", " world", "null"); + let stop = completion_event("cmpl-second", "second", "", "\"stop\""); + let usage = "data: {\"id\":\"cmpl-second\",\"object\":\"text_completion\",\"created\":2,\"model\":\"second\",\"choices\":[],\"usage\":{\"completion_tokens\":1}}\n\n".to_string(); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![ + vec![first], + vec![second, stop, usage, "data: [DONE]\n\n".to_string()], + ], + ); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + })) + .await; + + assert_eq!(response.status_code(), 200); + let body = response.text(); + assert!(body.contains("Hello")); + assert!(body.contains(" world")); + let requests = mock.get_requests(); + assert_eq!(requests.len(), 2); + let second_body: serde_json::Value = serde_json::from_slice(&requests[1].body).unwrap(); + assert_eq!(second_body["prompt"], "Say hello: Hello"); + assert!(body.contains("\"id\":\"cmpl-first\"")); + assert!(!body.contains("cmpl-second")); + } + + #[tokio::test] + async fn stream_continuation_terminal_finish_forwards_usage_and_does_not_retry() { + let text = completion_event("cmpl-first", "first", "Hello", "null"); + let stop = completion_event("cmpl-first", "first", "", "\"stop\""); + let usage = "data: {\"id\":\"cmpl-first\",\"object\":\"text_completion\",\"created\":1,\"model\":\"first\",\"choices\":[],\"usage\":{\"completion_tokens\":1}}\n\n".to_string(); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![text, stop, usage, "data: [DONE]\n\n".to_string()]], + ); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + })) + .await; + + let body = response.text(); + assert!(body.contains("completion_tokens")); + assert!(body.contains("[DONE]")); + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_crlf_terminal_stream_is_not_spuriously_continued() { + let event = + completion_event("cmpl-first", "first", "Hello", "\"stop\"").replace('\n', "\r\n"); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![event, "data:[DONE]\r\n\r\n".to_string()]], + ); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({"model":"requested-model","prompt":"P: ","stream":true})) + .await; + let body = response.text(); + + assert!(body.contains("Hello")); + assert!(body.contains("data:[DONE]")); + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_strict_chat_stream_resumes_as_one_assistant_message() { + let first_role = "data: {\"id\":\"chat-first\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"first\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\n".to_string(); + let first_text = "data: {\"id\":\"chat-first\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"first\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\n".to_string(); + let second_role = "data: {\"id\":\"chat-second\",\"object\":\"chat.completion.chunk\",\"created\":2,\"model\":\"second\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\n".to_string(); + let second_text = "data: {\"id\":\"chat-second\",\"object\":\"chat.completion.chunk\",\"created\":2,\"model\":\"second\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":null}]}\n\n".to_string(); + let stop = "data: {\"id\":\"chat-second\",\"object\":\"chat.completion.chunk\",\"created\":2,\"model\":\"second\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n".to_string(); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![ + vec![first_role, first_text], + vec![ + second_role, + second_text, + stop, + "data: [DONE]\n\n".to_string(), + ], + ], + ); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None, None], + stream_continuation_fallback( + true, + 2, + 1, + None, + vec!["/v1/completions", "/v1/chat/completions"], + ), + true, + false, + ); + let server = + TestServer::new(strict_stream_continuation_router(targets, mock.clone())).unwrap(); + + let response = server + .post("/v1/chat/completions") + .json(&json!({ + "model": "requested-model", + "messages": [{"role": "user", "content": "Hello"}], + "stream": true + })) + .await; + + assert_eq!(response.status_code(), 200); + let body = response.text(); + assert!(body.contains("Hello")); + assert!(body.contains(" world")); + assert!(body.contains("chat-first")); + assert!(!body.contains("chat-second")); + + let requests = mock.get_requests(); + assert_eq!(requests.len(), 2); + let continuation: serde_json::Value = serde_json::from_slice(&requests[1].body).unwrap(); + let messages = continuation["messages"].as_array().unwrap(); + assert_eq!(messages.last().unwrap()["role"], "assistant"); + assert_eq!(messages.last().unwrap()["content"], "Hello"); + } + + #[tokio::test] + async fn stream_continuation_strict_responses_stream_stitches_text_events() { + let first_created = "event: response.created\ndata: {\"type\":\"response.created\",\"sequence_number\":0,\"response\":{\"id\":\"resp_first\",\"model\":\"first\",\"output\":[]}}\n\n".to_string(); + let first_item = "event: response.output_item.added\ndata: {\"type\":\"response.output_item.added\",\"sequence_number\":1,\"output_index\":0,\"item\":{\"id\":\"msg_first\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[]}}\n\n".to_string(); + let first_part = "event: response.content_part.added\ndata: {\"type\":\"response.content_part.added\",\"sequence_number\":2,\"item_id\":\"msg_first\",\"output_index\":0,\"content_index\":0,\"part\":{\"type\":\"output_text\",\"text\":\"\"}}\n\n".to_string(); + let first_delta = "event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"sequence_number\":3,\"item_id\":\"msg_first\",\"output_index\":0,\"content_index\":0,\"delta\":\"Hello\"}\n\n".to_string(); + let second_created = first_created.replace("resp_first", "resp_second"); + let second_item = first_item.replace("msg_first", "msg_second"); + let second_part = first_part.replace("msg_first", "msg_second"); + let second_delta = "event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"sequence_number\":3,\"item_id\":\"msg_second\",\"output_index\":0,\"content_index\":0,\"delta\":\" world\"}\n\n".to_string(); + let completed = "event: response.completed\ndata: {\"type\":\"response.completed\",\"sequence_number\":4,\"response\":{\"id\":\"resp_second\",\"model\":\"second\",\"output\":[{\"id\":\"msg_second\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\" world\"}]}]}}\n\n".to_string(); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![ + vec![first_created, first_item, first_part, first_delta], + vec![ + second_created, + second_item, + second_part, + second_delta, + completed, + ], + ], + ); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/responses"]), + true, + false, + ); + let server = + TestServer::new(strict_stream_continuation_router(targets, mock.clone())).unwrap(); + + let response = server + .post("/v1/responses") + .json(&json!({ + "model": "requested-model", + "input": "Hello", + "stream": true + })) + .await; + + assert_eq!(response.status_code(), 200); + let body = response.text(); + assert!(body.contains("Hello")); + assert!(body.contains(" world")); + assert!(body.contains("resp_first")); + assert!(!body.contains("resp_second")); + + let requests = mock.get_requests(); + assert_eq!(requests.len(), 2); + let continuation: serde_json::Value = serde_json::from_slice(&requests[1].body).unwrap(); + let input = continuation["input"].as_array().unwrap(); + assert_eq!(input.last().unwrap()["role"], "assistant"); + assert_eq!(input.last().unwrap()["content"][0]["text"], "Hello"); + } + + #[tokio::test] + async fn stream_continuation_adapter_responses_uses_external_endpoint_and_chat_protocol() { + let first_role = "data: {\"id\":\"chat-first\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"first\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\n".to_string(); + let first_text = "data: {\"id\":\"chat-first\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"first\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\n".to_string(); + let second_role = "data: {\"id\":\"chat-second\",\"object\":\"chat.completion.chunk\",\"created\":2,\"model\":\"second\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\n".to_string(); + let second_text = "data: {\"id\":\"chat-second\",\"object\":\"chat.completion.chunk\",\"created\":2,\"model\":\"second\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":null}]}\n\n".to_string(); + let stop = "data: {\"id\":\"chat-second\",\"object\":\"chat.completion.chunk\",\"created\":2,\"model\":\"second\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n".to_string(); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![ + vec![first_role, first_text], + vec![ + second_role, + second_text, + stop, + "data: [DONE]\n\n".to_string(), + ], + ], + ); + let adapter_targets = (0..2) + .map(|index| { + Target::builder() + .url( + format!("https://adapter-{index}.example.com/") + .parse() + .unwrap(), + ) + .open_responses(OpenResponsesConfig { adapter: true }) + .build() + }) + .collect(); + let targets = stream_continuation_targets_from_targets( + "requested-model", + adapter_targets, + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/responses"]), + true, + false, + ); + let server = + TestServer::new(strict_stream_continuation_router(targets, mock.clone())).unwrap(); + + let response = server + .post("/v1/responses") + .json(&json!({ + "model": "requested-model", + "input": "Hello", + "stream": true + })) + .await; + + assert_eq!(response.status_code(), 200); + let body = response.text(); + assert!(body.contains("Hello")); + assert!(body.contains(" world")); + let requests = mock.get_requests(); + assert_eq!(requests.len(), 2); + let continuation: serde_json::Value = serde_json::from_slice(&requests[1].body).unwrap(); + let messages = continuation["messages"].as_array().unwrap(); + assert_eq!(messages.last().unwrap()["role"], "assistant"); + assert_eq!(messages.last().unwrap()["content"], "Hello"); + } + + #[tokio::test] + async fn stream_continuation_edge_translated_responses_preserves_external_endpoint() { + let first = "data: {\"id\":\"chat-first\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"first\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\n".to_string(); + let second = "data: {\"id\":\"chat-second\",\"object\":\"chat.completion.chunk\",\"created\":2,\"model\":\"second\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":\"stop\"}]}\n\n".to_string(); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![first], vec![second, "data: [DONE]\n\n".to_string()]], + ); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/responses"]), + true, + false, + ); + let state = AppState::with_client(targets, mock.clone()); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/chat/completions") + .header("content-type", "application/json") + .body(Body::from( + json!({ + "model": "requested-model", + "messages": [{"role": "user", "content": "Hello"}], + "stream": true + }) + .to_string(), + )) + .unwrap(); + + let response = crate::strict::handlers::responses_handler( + axum::extract::State(state), + HeaderMap::new(), + request, + ) + .await; + let body = response.into_body().collect().await.unwrap().to_bytes(); + let body = String::from_utf8(body.to_vec()).unwrap(); + + assert!(body.contains("Hello")); + assert!(body.contains(" world")); + let requests = mock.get_requests(); + assert_eq!(requests.len(), 2); + let continuation: serde_json::Value = serde_json::from_slice(&requests[1].body).unwrap(); + let messages = continuation["messages"].as_array().unwrap(); + assert_eq!(messages.last().unwrap()["role"], "assistant"); + assert_eq!(messages.last().unwrap()["content"], "Hello"); + } + + #[tokio::test] + async fn stream_continuation_accepts_large_responses_terminal_snapshot_beyond_prefix_cap() { + let chunk = "\"".repeat(1024); + let full_text = chunk.repeat(70); + let mut events = vec![ + "event: response.created\ndata: {\"type\":\"response.created\",\"sequence_number\":0,\"response\":{\"id\":\"resp_first\",\"model\":\"first\",\"output\":[]}}\n\n".to_string(), + "event: response.output_item.added\ndata: {\"type\":\"response.output_item.added\",\"sequence_number\":1,\"output_index\":0,\"item\":{\"id\":\"msg_first\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[]}}\n\n".to_string(), + "event: response.content_part.added\ndata: {\"type\":\"response.content_part.added\",\"sequence_number\":2,\"item_id\":\"msg_first\",\"output_index\":0,\"content_index\":0,\"part\":{\"type\":\"output_text\",\"text\":\"\"}}\n\n".to_string(), + ]; + for sequence in 3..73 { + events.push(format!( + "event: response.output_text.delta\ndata: {{\"type\":\"response.output_text.delta\",\"sequence_number\":{sequence},\"item_id\":\"msg_first\",\"output_index\":0,\"content_index\":0,\"delta\":{chunk:?}}}\n\n" + )); + } + events.push(format!( + "event: response.completed\ndata: {{\"type\":\"response.completed\",\"sequence_number\":73,\"response\":{{\"id\":\"resp_first\",\"model\":\"first\",\"output\":[{{\"id\":\"msg_first\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[{{\"type\":\"output_text\",\"text\":{full_text:?}}}]}}]}}}}\n\n" + )); + let mock = MockHttpClient::new_streaming_sequence(StatusCode::OK, vec![events]); + let mut fallback = stream_continuation_fallback(true, 1, 0, None, vec!["/v1/responses"]); + fallback + .stream_continuation + .as_mut() + .unwrap() + .max_buffered_bytes = 1024; + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None], + fallback, + true, + false, + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/responses") + .header("content-type", "application/json") + .body(Body::from( + json!({"model":"requested-model","input":"Hello","stream":true}).to_string(), + )) + .unwrap(); + + let response = strict_stream_continuation_router(targets, mock) + .oneshot(request) + .await + .unwrap(); + let body = response.into_body().collect().await.unwrap().to_bytes(); + let body = String::from_utf8(body.to_vec()).unwrap(); + let completed: serde_json::Value = body + .lines() + .filter_map(|line| line.strip_prefix("data:")) + .filter_map(|data| serde_json::from_str(data.trim()).ok()) + .find(|event: &serde_json::Value| event["type"] == "response.completed") + .unwrap(); + assert_eq!( + completed["response"]["output"][0]["content"][0]["text"], + full_text + ); + } + + #[tokio::test] + async fn stream_continuation_fallback_master_switch_disables_continuation() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let mock = MockHttpClient::new_streaming_sequence(StatusCode::OK, vec![vec![first]]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(false, 2, 1, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + server + .post("/v1/completions") + .json(&json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + })) + .await; + + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_body_error_resumes_and_finishes() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let second = completion_event("cmpl-second", "second", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_response_sequence(vec![ + MockStreamingResponse::sse( + StatusCode::OK, + vec![ + MockStreamEvent::Data(first), + MockStreamEvent::Error("upstream reset".to_string()), + ], + ), + MockStreamingResponse::sse( + StatusCode::OK, + vec![ + MockStreamEvent::Data(second), + MockStreamEvent::Data("data: [DONE]\n\n".to_string()), + ], + ), + ]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + })) + .await; + + let body = response.text(); + assert!(body.contains("Hello")); + assert!(body.contains(" world")); + assert_eq!(mock.get_requests().len(), 2); + } + + #[tokio::test] + async fn stream_continuation_exhausted_plain_eof_has_no_synthetic_done() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let second = completion_event("cmpl-second", "second", " world", "null"); + let mock = + MockHttpClient::new_streaming_sequence(StatusCode::OK, vec![vec![first], vec![second]]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + })) + .await; + + let body = response.text(); + assert!(body.contains("Hello")); + assert!(body.contains(" world")); + assert!(!body.contains("[DONE]")); + assert_eq!(mock.get_requests().len(), 2); + } + + #[tokio::test] + async fn stream_continuation_exhausted_body_error_remains_downstream_error() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let second = completion_event("cmpl-second", "second", " world", "null"); + let mock = MockHttpClient::new_streaming_response_sequence(vec![ + MockStreamingResponse::sse( + StatusCode::OK, + vec![ + MockStreamEvent::Data(first), + MockStreamEvent::Error("first reset".to_string()), + ], + ), + MockStreamingResponse::sse( + StatusCode::OK, + vec![ + MockStreamEvent::Data(second), + MockStreamEvent::Error("final reset".to_string()), + ], + ), + ]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + }) + .to_string(), + )) + .unwrap(); + + let response = build_router(AppState::with_client(targets, mock.clone())) + .oneshot(request) + .await + .unwrap(); + let result = response.into_body().collect().await; + + assert!(result.is_err()); + assert_eq!(mock.get_requests().len(), 2); + } + + #[tokio::test] + async fn stream_continuation_idle_timeout_resumes_pending_body() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let second = completion_event("cmpl-second", "second", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_response_sequence(vec![ + MockStreamingResponse::sse( + StatusCode::OK, + vec![MockStreamEvent::Data(first), MockStreamEvent::Pending], + ), + MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("text/event-stream".to_string()), + headers: vec![("x-upstream".to_string(), "continuation".to_string())], + events: vec![ + MockStreamEvent::Data(second), + MockStreamEvent::Data("data: [DONE]\n\n".to_string()), + ], + }, + ]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, Some(10), vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + })) + .await; + + assert!(response.text().contains(" world")); + assert_eq!(mock.get_requests().len(), 2); + } + + #[tokio::test] + async fn stream_continuation_rewrites_model_for_selected_provider() { + let first = completion_event("cmpl-first", "provider-a", "Hello", "null"); + let second = completion_event("cmpl-second", "provider-b", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![first], vec![second, "data: [DONE]\n\n".to_string()]], + ); + let targets = stream_continuation_targets( + "requested-model", + &[Some("provider-a"), Some("provider-b")], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + server + .post("/v1/completions") + .json(&json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + })) + .await; + + let requests = mock.get_requests(); + assert_eq!(requests.len(), 2); + let first_body: serde_json::Value = serde_json::from_slice(&requests[0].body).unwrap(); + let second_body: serde_json::Value = serde_json::from_slice(&requests[1].body).unwrap(); + assert_eq!(first_body["model"], "provider-a"); + assert_eq!(second_body["model"], "provider-a"); + assert!(requests[1].uri.contains("stream-0.example.com")); + } + + #[tokio::test] + async fn stream_continuation_releases_provider_guard_before_reselection() { + let first = completion_event("cmpl-first", "provider-a", "Hello", "null"); + let second = completion_event("cmpl-second", "provider-a", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![first], vec![second, "data: [DONE]\n\n".to_string()]], + ); + let target = Target::builder() + .url("https://stream-0.example.com/".parse().unwrap()) + .onwards_model("provider-a".to_string()) + .build(); + let pool = ProviderPool::with_config( + vec![Provider::with_concurrency_limit(target, 1, 1)], + None, + None, + None, + Some(stream_continuation_fallback( + true, + 1, + 1, + None, + vec!["/v1/completions"], + )), + LoadBalanceStrategy::Priority, + false, + Vec::new(), + ); + let targets_map = Arc::new(DashMap::new()); + targets_map.insert("requested-model".to_string(), pool); + let targets = target::Targets { + targets: targets_map, + key_rate_limiters: Arc::new(DashMap::new()), + key_concurrency_limiters: Arc::new(DashMap::new()), + key_labels: Arc::new(DashMap::new()), + strict_mode: false, + http_pool_config: None, + }; + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + })) + .await; + + assert!(response.text().contains(" world")); + assert_eq!(mock.get_requests().len(), 2); + } + + #[tokio::test] + async fn stream_continuation_uses_fresh_post_response_attempt_budget() { + let first = completion_event("cmpl-first", "provider-b", "Hello", "null"); + let second = completion_event("cmpl-second", "provider-a", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_response_sequence(vec![ + MockStreamingResponse::sse(StatusCode::BAD_GATEWAY, Vec::new()), + MockStreamingResponse::sse(StatusCode::OK, vec![MockStreamEvent::Data(first)]), + MockStreamingResponse::sse( + StatusCode::OK, + vec![ + MockStreamEvent::Data(second), + MockStreamEvent::Data("data: [DONE]\n\n".to_string()), + ], + ), + ]); + let targets = stream_continuation_targets( + "requested-model", + &[Some("provider-a"), Some("provider-b")], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + })) + .await; + + assert!(response.text().contains(" world")); + assert_eq!(mock.get_requests().len(), 3); + } + + #[tokio::test] + async fn stream_continuation_requests_identity_and_clears_composite_framing() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let second = completion_event("cmpl-second", "second", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_response_sequence(vec![ + MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("Text/Event-Stream; charset=utf-8".to_string()), + headers: vec![ + ("content-length".to_string(), first.len().to_string()), + ("content-encoding".to_string(), "IDENTITY".to_string()), + ("transfer-encoding".to_string(), "chunked".to_string()), + ("trailer".to_string(), "x-checksum".to_string()), + ("etag".to_string(), "\"initial-etag\"".to_string()), + ("digest".to_string(), "sha-256=YWJj".to_string()), + ("content-digest".to_string(), "sha-256=:YWJj:".to_string()), + ("repr-digest".to_string(), "sha-256=:YWJj:".to_string()), + ( + "representation-digest".to_string(), + "sha-256=:YWJj:".to_string(), + ), + ("accept-ranges".to_string(), "bytes".to_string()), + ("content-range".to_string(), "bytes 0-1/2".to_string()), + ( + "last-modified".to_string(), + "Mon, 01 Jan 2024 00:00:00 GMT".to_string(), + ), + ("x-upstream".to_string(), "initial".to_string()), + ], + events: vec![MockStreamEvent::Data(first)], + }, + MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("text/event-stream".to_string()), + headers: vec![("x-upstream".to_string(), "continuation".to_string())], + events: vec![ + MockStreamEvent::Data(second), + MockStreamEvent::Data("data: [DONE]\n\n".to_string()), + ], + }, + ]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .header("accept-encoding", "gzip, br") + .body(axum::body::Body::from( + json!({ + "model": "requested-model", + "prompt": "Say hello: ", + "stream": true + }) + .to_string(), + )) + .unwrap(); + + let response = build_router(AppState::with_client(targets, mock.clone())) + .oneshot(request) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::OK); + assert_eq!(response.headers().get("x-upstream").unwrap(), "initial"); + assert_eq!(response.headers().get("cache-control").unwrap(), "no-cache"); + for removed in [ + "content-length", + "content-encoding", + "transfer-encoding", + "trailer", + "etag", + "digest", + "content-digest", + "repr-digest", + "representation-digest", + "accept-ranges", + "content-range", + "last-modified", + "connection", + ] { + assert!( + response.headers().get(removed).is_none(), + "composite response retained {removed}" + ); + } + let body = response.into_body().collect().await.unwrap().to_bytes(); + let body = String::from_utf8(body.to_vec()).unwrap(); + assert!(body.contains("Hello")); + assert!(body.contains(" world")); + + let requests = mock.get_requests(); + assert_eq!(requests.len(), 2); + for request in requests { + assert_eq!( + request + .headers + .iter() + .find(|(name, _)| name == "accept-encoding") + .map(|(_, value)| value.as_str()), + Some("identity") + ); + } + } + + #[tokio::test] + async fn stream_continuation_preserves_encoded_initial_response_unwrapped() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let compressed = gzip_bytes(first.as_bytes()); + let mock = MockHttpClient::new_streaming_response_sequence(vec![MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("text/event-stream".to_string()), + headers: vec![ + ("content-encoding".to_string(), "gzip".to_string()), + ("content-length".to_string(), compressed.len().to_string()), + ], + events: vec![MockStreamEvent::Bytes(compressed.clone())], + }]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + + let response = build_router(AppState::with_client(targets, mock.clone())) + .oneshot(request) + .await + .unwrap(); + assert_eq!(response.headers().get("content-encoding").unwrap(), "gzip"); + assert_eq!( + response + .headers() + .get("content-length") + .unwrap() + .to_str() + .unwrap(), + compressed.len().to_string().as_str() + ); + let body = response.into_body().collect().await.unwrap().to_bytes(); + assert_eq!(body.as_ref(), compressed.as_slice()); + assert_ne!(body.as_ref(), first.as_bytes()); + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_encoded_chat_still_uses_baseline_strict_processing() { + let error = "data: {\"error\":{\"code\":429,\"message\":\"rate limited\"}}\n\n"; + let compressed = gzip_bytes(error.as_bytes()); + let responses = (0..2) + .map(|_| MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("text/event-stream".to_string()), + headers: vec![("content-encoding".to_string(), "gzip".to_string())], + events: vec![MockStreamEvent::Bytes(compressed.clone())], + }) + .collect(); + let mock = MockHttpClient::new_streaming_response_sequence(responses); + let targets = fallback_targets("gpt-4", 2, vec![429]); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/chat/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({ + "model": "gpt-4", + "messages": [{"role":"user","content":"Hello"}], + "stream": true + }) + .to_string(), + )) + .unwrap(); + let response = strict_stream_continuation_router(targets, mock.clone()) + .oneshot(request) + .await; + + let response = response.unwrap(); + assert!(response.into_body().collect().await.is_err()); + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_encoded_ineligible_strict_completion_is_preserved_safely() { + let chunk = "data: {\"id\":\"cmpl-one\",\"object\":\"text_completion\",\"created\":1,\"model\":\"provider-model\",\"choices\":[{\"index\":0,\"text\":\"Hello\",\"finish_reason\":\"stop\"}],\"provider_field\":\"remove\"}\n\ndata: [DONE]\n\n"; + let compressed = gzip_bytes(chunk.as_bytes()); + let mock = MockHttpClient::new_streaming_response_sequence(vec![MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("text/event-stream".to_string()), + headers: vec![ + ("content-encoding".to_string(), "gzip".to_string()), + ("content-length".to_string(), compressed.len().to_string()), + ], + events: vec![MockStreamEvent::Bytes(compressed.clone())], + }]); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None], + stream_continuation_fallback(false, 1, 1, None, vec!["/v1/completions"]), + true, + false, + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + let response = strict_stream_continuation_router(targets, mock.clone()) + .oneshot(request) + .await + .unwrap(); + + assert_eq!(response.headers().get("content-encoding").unwrap(), "gzip"); + let body = response.into_body().collect().await.unwrap().to_bytes(); + assert_eq!(body.as_ref(), compressed.as_slice()); + assert_eq!(mock.get_requests().len(), 1); + assert_eq!( + mock.get_requests()[0] + .headers + .iter() + .find(|(name, _)| name == "accept-encoding") + .map(|(_, value)| value.as_str()), + Some("identity") + ); + } + + #[tokio::test] + async fn stream_continuation_encoded_strict_chat_and_responses_are_preserved_unwrapped() { + for (path, request_body) in [ + ( + "/v1/chat/completions", + json!({ + "model": "requested-model", + "messages": [{"role": "user", "content": "Hello"}], + "stream": true + }), + ), + ( + "/v1/responses", + json!({ + "model": "requested-model", + "input": "Hello", + "stream": true + }), + ), + ] { + let compressed = gzip_bytes(b"encoded-stream"); + let mock = + MockHttpClient::new_streaming_response_sequence(vec![MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("text/event-stream".to_string()), + headers: vec![ + ("content-encoding".to_string(), "gzip".to_string()), + ("content-length".to_string(), compressed.len().to_string()), + ], + events: vec![MockStreamEvent::Bytes(compressed.clone())], + }]); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None], + stream_continuation_fallback(true, 1, 1, None, vec![path]), + true, + false, + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri(path) + .header("content-type", "application/json") + .body(axum::body::Body::from(request_body.to_string())) + .unwrap(); + + let response = strict_stream_continuation_router(targets, mock.clone()) + .oneshot(request) + .await + .unwrap(); + + assert_eq!(response.headers().get("content-encoding").unwrap(), "gzip"); + let body = response.into_body().collect().await.unwrap().to_bytes(); + assert_eq!(body.as_ref(), compressed.as_slice()); + assert_eq!(mock.get_requests().len(), 1); + assert_eq!( + mock.get_requests()[0] + .headers + .iter() + .find(|(name, _)| name == "accept-encoding") + .map(|(_, value)| value.as_str()), + Some("identity") + ); + } + } + + #[tokio::test] + async fn stream_continuation_encoded_unrelated_response_still_uses_non_strict_transform() { + let upstream = gzip_bytes(br#"{"object":"list","data":[]}"#); + let transformed = br#"{"transformed":true}"#.to_vec(); + let mock = MockHttpClient::new_streaming_response_sequence(vec![MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("application/json".to_string()), + headers: vec![("content-encoding".to_string(), "gzip".to_string())], + events: vec![MockStreamEvent::Bytes(upstream)], + }]); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None], + stream_continuation_fallback(true, 1, 1, None, vec!["/v1/completions"]), + false, + true, + ); + let transformed_for_closure = transformed.clone(); + let transform: ResponseTransformFn = Arc::new(move |path, _, _, _| { + (path == "/v1/embeddings") + .then(|| axum::body::Bytes::from(transformed_for_closure.clone())) + .map(Some) + .ok_or_else(|| "unexpected path".to_string()) + }); + let state = AppState::with_client(targets, mock.clone()).with_response_transform(transform); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/embeddings") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","input":"Hello"}).to_string(), + )) + .unwrap(); + + let response = build_router(state).oneshot(request).await.unwrap(); + let body = response.into_body().collect().await.unwrap().to_bytes(); + + assert_eq!(body.as_ref(), transformed.as_slice()); + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_rejects_encoded_and_lookalike_continuations() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let encoded = completion_event("cmpl-encoded", "encoded", " encoded", "null"); + let encoded_bytes = brotli_bytes(encoded.as_bytes()); + let lookalike = completion_event("cmpl-lookalike", "lookalike", " lookalike", "null"); + let malformed = completion_event("cmpl-malformed", "malformed", " malformed", "null"); + let final_event = completion_event("cmpl-final", "final", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_response_sequence(vec![ + MockStreamingResponse::sse(StatusCode::OK, vec![MockStreamEvent::Data(first)]), + MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("text/event-stream".to_string()), + headers: vec![("content-encoding".to_string(), "br".to_string())], + events: vec![MockStreamEvent::Bytes(encoded_bytes)], + }, + MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("application/x-text/event-streamish".to_string()), + headers: Vec::new(), + events: vec![MockStreamEvent::Data(lookalike)], + }, + MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("text/event-stream; charset".to_string()), + headers: Vec::new(), + events: vec![MockStreamEvent::Data(malformed)], + }, + MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("TEXT/EVENT-STREAM; CHARSET=UTF-8".to_string()), + headers: Vec::new(), + events: vec![ + MockStreamEvent::Data(final_event), + MockStreamEvent::Data("data: [DONE]\n\n".to_string()), + ], + }, + ]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None, None, None, None], + stream_continuation_fallback(true, 2, 4, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({"model":"requested-model","prompt":"P: ","stream":true})) + .await; + let body = response.text(); + + assert!(body.contains("Hello")); + assert!(body.contains(" world")); + assert!(!body.contains(" encoded")); + assert!(!body.contains(" lookalike")); + assert!(!body.contains(" malformed")); + assert_eq!(mock.get_requests().len(), 5); + } + + #[tokio::test] + async fn stream_continuation_initial_lookalike_media_type_is_not_wrapped() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let mock = MockHttpClient::new_streaming_response_sequence(vec![MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("application/x-text/event-streamish".to_string()), + headers: Vec::new(), + events: vec![MockStreamEvent::Data(first.clone())], + }]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({"model":"requested-model","prompt":"P: ","stream":true})) + .await; + + assert_eq!(response.text(), first); + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_strict_mode_uses_external_completion_endpoint() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let second = completion_event("cmpl-second", "second", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_response_sequence(vec![ + MockStreamingResponse::sse(StatusCode::OK, vec![MockStreamEvent::Data(first)]), + MockStreamingResponse::sse( + StatusCode::OK, + vec![ + MockStreamEvent::Data(second), + MockStreamEvent::Data("data: [DONE]\n\n".to_string()), + ], + ), + MockStreamingResponse { + status: StatusCode::OK, + content_type: Some("application/json".to_string()), + headers: Vec::new(), + events: vec![MockStreamEvent::Data( + "{\"object\":\"list\",\"data\":[],\"model\":\"requested-model\",\"usage\":{\"prompt_tokens\":1,\"total_tokens\":1}}".to_string(), + )], + }, + ]); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + true, + false, + ); + let server = + TestServer::new(strict_stream_continuation_router(targets, mock.clone())).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({"model":"requested-model","prompt":"P: ","stream":true})) + .await; + let body = response.text(); + + assert!(body.contains("Hello")); + assert!(body.contains(" world")); + let continuation_uri: axum::http::Uri = mock.get_requests()[1].uri.parse().unwrap(); + assert_eq!(continuation_uri.path(), "/completions"); + assert_eq!( + continuation_uri.path_and_query().unwrap().as_str(), + "/completions" + ); + + let unrelated = server + .post("/v1/embeddings") + .json(&json!({"model":"requested-model","input":"Hello"})) + .await; + assert_eq!(unrelated.status_code(), StatusCode::OK); + + let requests = mock.get_requests(); + assert_eq!(requests.len(), 3); + let unrelated_uri: axum::http::Uri = requests[2].uri.parse().unwrap(); + assert_eq!(unrelated_uri.path(), "/embeddings"); + assert_eq!( + unrelated_uri.path_and_query().unwrap().as_str(), + "/embeddings" + ); + } + + #[tokio::test] + async fn stream_continuation_non_strict_sanitizer_preserves_completion_text() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let second = completion_event("cmpl-second", "second", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![first], vec![second, "data: [DONE]\n\n".to_string()]], + ); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + false, + true, + ); + let state = AppState::with_client(targets, mock.clone()) + .with_response_transform(create_openai_sanitizer()); + let server = TestServer::new(build_router(state)).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({"model":"requested-model","prompt":"P: ","stream":true})) + .await; + let body = response.text(); + let values = completion_sse_values(&body); + let text: String = values + .iter() + .filter_map(|value| value["choices"][0]["text"].as_str()) + .collect(); + + assert_eq!(text, "Hello world"); + assert!(body.contains("data: [DONE]")); + assert_eq!(mock.get_requests().len(), 2); + } + + #[tokio::test] + async fn stream_continuation_strict_missing_identity_stays_stable() { + let first = "data: {\"object\":\"text_completion\",\"choices\":[{\"index\":0,\"text\":\"Hello\",\"finish_reason\":null}]}\n\n".to_string(); + let second = completion_event("cmpl-provider", "provider-model", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![first], vec![second, "data: [DONE]\n\n".to_string()]], + ); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + true, + false, + ); + let server = + TestServer::new(strict_stream_continuation_router(targets, mock.clone())).unwrap(); + + let response = server + .post("/v1/completions") + .add_header("model-override", "requested-model") + .json(&json!({"model":"body-model","prompt":"P: ","stream":true})) + .await; + let body = response.text(); + let values = completion_sse_values(&body); + let ids: std::collections::HashSet<_> = values.iter().map(|value| &value["id"]).collect(); + let models: std::collections::HashSet<_> = + values.iter().map(|value| &value["model"]).collect(); + let created: std::collections::HashSet<_> = + values.iter().map(|value| &value["created"]).collect(); + + assert_eq!(ids.len(), 1); + assert_eq!(models.len(), 1); + assert_eq!(created.len(), 1); + assert_eq!(values[0]["model"], "requested-model"); + assert_ne!(values[0]["id"], "cmpl-provider"); + assert_eq!(mock.get_requests().len(), 2); + } + + #[tokio::test] + async fn stream_continuation_accumulates_bytes_across_two_body_failures() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let second = completion_event("cmpl-second", "second", " brave", "null"); + let third = completion_event("cmpl-third", "third", " world", "\"stop\""); + let mock = MockHttpClient::new_streaming_response_sequence(vec![ + MockStreamingResponse::sse( + StatusCode::OK, + vec![ + MockStreamEvent::Data(first), + MockStreamEvent::Error("first reset".to_string()), + ], + ), + MockStreamingResponse::sse( + StatusCode::OK, + vec![ + MockStreamEvent::Data(second), + MockStreamEvent::Error("second reset".to_string()), + ], + ), + MockStreamingResponse::sse( + StatusCode::OK, + vec![ + MockStreamEvent::Data(third), + MockStreamEvent::Data("data: [DONE]\n\n".to_string()), + ], + ), + ]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None, None], + stream_continuation_fallback(true, 2, 2, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({"model":"requested-model","prompt":"P: ","stream":true})) + .await; + let body = response.text(); + let requests = mock.get_requests(); + + assert!(body.contains("Hello")); + assert!(body.contains(" brave")); + assert!(body.contains(" world")); + assert_eq!(requests.len(), 3); + let third_body: serde_json::Value = serde_json::from_slice(&requests[2].body).unwrap(); + assert_eq!(third_body["prompt"], "P: Hello brave"); + } + + #[tokio::test] + async fn stream_continuation_finish_reason_makes_later_error_clean() { + let finished = completion_event("cmpl-first", "first", "Hello", "\"stop\""); + let mock = + MockHttpClient::new_streaming_response_sequence(vec![MockStreamingResponse::sse( + StatusCode::OK, + vec![ + MockStreamEvent::Data(finished), + MockStreamEvent::Error("late reset".to_string()), + ], + )]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + + let response = build_router(AppState::with_client(targets, mock.clone())) + .oneshot(request) + .await + .unwrap(); + let body = response.into_body().collect().await.unwrap().to_bytes(); + + assert!(String::from_utf8_lossy(&body).contains("Hello")); + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_finish_reason_makes_later_timeout_clean() { + let finished = completion_event("cmpl-first", "first", "Hello", "\"stop\""); + let mock = + MockHttpClient::new_streaming_response_sequence(vec![MockStreamingResponse::sse( + StatusCode::OK, + vec![MockStreamEvent::Data(finished), MockStreamEvent::Pending], + )]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, Some(10), vec!["/v1/completions"]), + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + + let response = build_router(AppState::with_client(targets, mock.clone())) + .oneshot(request) + .await + .unwrap(); + let body = response.into_body().collect().await.unwrap().to_bytes(); + + assert!(String::from_utf8_lossy(&body).contains("Hello")); + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_finish_reason_makes_later_framing_errors_clean() { + let finished = completion_event("cmpl-first", "first", "Hello", "\"stop\""); + + for tail in [ + MockStreamEvent::Data("data: incomplete".to_string()), + MockStreamEvent::Bytes(vec![b'x'; 64 * 1024 + 1]), + ] { + let mock = + MockHttpClient::new_streaming_response_sequence(vec![MockStreamingResponse::sse( + StatusCode::OK, + vec![MockStreamEvent::Data(finished.clone()), tail], + )]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + + let response = build_router(AppState::with_client(targets, mock.clone())) + .oneshot(request) + .await + .unwrap(); + let body = response.into_body().collect().await.unwrap().to_bytes(); + + assert!(String::from_utf8_lossy(&body).contains("Hello")); + assert_eq!(mock.get_requests().len(), 1); + } + } + + #[tokio::test] + async fn stream_continuation_strict_late_framing_errors_never_retry() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let fallback = completion_event("cmpl-second", "second", " leaked", "\"stop\""); + + for tail in [ + MockStreamEvent::Data("data: incomplete".to_string()), + MockStreamEvent::Bytes(vec![b'x'; 64 * 1024 + 1]), + ] { + let mock = MockHttpClient::new_streaming_response_sequence(vec![ + MockStreamingResponse::sse( + StatusCode::OK, + vec![MockStreamEvent::Data(first.clone()), tail], + ), + MockStreamingResponse::sse( + StatusCode::OK, + vec![MockStreamEvent::Data(fallback.clone())], + ), + ]); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + true, + false, + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + + let response = strict_stream_continuation_router(targets, mock.clone()) + .oneshot(request) + .await + .unwrap(); + + assert!(response.into_body().collect().await.is_err()); + assert_eq!(mock.get_requests().len(), 1); + } + } + + #[tokio::test] + async fn stream_continuation_strict_first_frame_framing_errors_never_retry() { + let fallback = completion_event("cmpl-second", "second", " leaked", "\"stop\""); + + for first_event in [ + MockStreamEvent::Data("data: incomplete".to_string()), + MockStreamEvent::Bytes(vec![b'x'; 64 * 1024 + 1]), + ] { + let mock = MockHttpClient::new_streaming_response_sequence(vec![ + MockStreamingResponse::sse(StatusCode::OK, vec![first_event]), + MockStreamingResponse::sse( + StatusCode::OK, + vec![MockStreamEvent::Data(fallback.clone())], + ), + ]); + let targets = stream_continuation_targets_with_options( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + true, + false, + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + + let response = strict_stream_continuation_router(targets, mock.clone()) + .oneshot(request) + .await + .unwrap(); + + assert!(response.into_body().collect().await.is_err()); + assert_eq!(mock.get_requests().len(), 1); + } + } + + #[tokio::test] + async fn stream_continuation_framing_errors_fail_closed_without_retry() { + let incomplete = completion_event("cmpl-first", "first", "Hello", "null") + .trim_end_matches('\n') + .to_string(); + let oversized = vec![b'x'; 64 * 1024 + 1]; + + for event in [ + MockStreamEvent::Data(incomplete), + MockStreamEvent::Bytes(oversized), + ] { + let mock = + MockHttpClient::new_streaming_response_sequence(vec![MockStreamingResponse::sse( + StatusCode::OK, + vec![event], + )]); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + + let response = build_router(AppState::with_client(targets, mock.clone())) + .oneshot(request) + .await + .unwrap(); + + assert!(response.into_body().collect().await.is_err()); + assert_eq!(mock.get_requests().len(), 1); + } + } + + #[tokio::test] + async fn stream_continuation_unsafe_initial_events_never_retry_incomplete_prefixes() { + for unsafe_event in [ + "data: {not-json}\n\n", + "data: {\"error\":{\"message\":\"secret\"}}\n\n", + "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"x\"},\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":0,\"message\":{\"content\":\"x\"},\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":0,\"text\":\"x\",\"finish_reason\":null,\"tool_calls\":[]}]}\n\n", + "data: {\"choices\":[{\"text\":\"x\",\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":1,\"text\":\"x\",\"finish_reason\":null}]}\n\n", + "data: {\"type\":\"notification\",\"payload\":{}}\n\n", + "event: completion\ndata: {\"choices\":[{\"index\":0,\"text\":\"x\",\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":0,\"text\":\"x\",\"finish_reason\":null}],\"payload\":{}}\n\n", + "data: {\"choices\":[{\"index\":0,\"text\":\"x\",\"finish_reason\":null}],\"tool_calls\":[]}\n\n", + ] { + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![unsafe_event.to_string()]], + ); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let server = + TestServer::new(build_router(AppState::with_client(targets, mock.clone()))) + .unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({"model":"requested-model","prompt":"P: ","stream":true})) + .await; + + assert!(response.text().contains(unsafe_event.trim())); + assert_eq!(mock.get_requests().len(), 1, "unsafe event: {unsafe_event}"); + } + } + + #[tokio::test] + async fn stream_continuation_unsafe_continuation_events_are_not_spliced() { + for unsafe_event in [ + "data: {not-json}\n\n", + "data: {\"error\":{\"message\":\"continuation-secret\"}}\n\n", + "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"continuation-secret\"},\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":0,\"message\":{\"content\":\"continuation-secret\"},\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":1,\"text\":\"continuation-secret\",\"finish_reason\":null}]}\n\n", + "data: {\"type\":\"notification\",\"payload\":{\"value\":\"continuation-secret\"}}\n\n", + "event: completion\ndata: {\"choices\":[{\"index\":0,\"text\":\"continuation-secret\",\"finish_reason\":null}]}\n\n", + "data: {\"choices\":[{\"index\":0,\"text\":\"continuation-secret\",\"finish_reason\":null}],\"payload\":{}}\n\n", + "data: {\"choices\":[{\"index\":0,\"text\":\"continuation-secret\",\"finish_reason\":null}],\"tool_calls\":[]}\n\n", + ] { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![first], vec![unsafe_event.to_string()]], + ); + let targets = stream_continuation_targets( + "requested-model", + &[None, None], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + let response = build_router(AppState::with_client(targets, mock.clone())) + .oneshot(request) + .await + .unwrap(); + let mut stream = response.into_body().into_data_stream(); + let mut visible = Vec::new(); + let mut saw_error = false; + while let Some(item) = stream.next().await { + match item { + Ok(bytes) => visible.extend_from_slice(&bytes), + Err(_) => { + saw_error = true; + break; + } + } + } + + let visible = String::from_utf8(visible).unwrap(); + assert!(visible.contains("Hello")); + assert!(!visible.contains("continuation-secret")); + assert!(saw_error, "unsafe event: {unsafe_event}"); + assert_eq!(mock.get_requests().len(), 2); + } + } + + #[tokio::test] + async fn stream_continuation_mixed_trust_uses_least_privilege_strict_policy() { + let first = completion_event("cmpl-first", "first", "Hello", "null"); + let error = "data: {\"error\":{\"code\":500,\"message\":\"provider-secret\"}}\n\n"; + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![first, error.to_string()]], + ); + let trusted = Target::builder() + .url("https://trusted.example.com/".parse().unwrap()) + .trusted(true) + .build(); + let untrusted = Target::builder() + .url("https://untrusted.example.com/".parse().unwrap()) + .trusted(false) + .build(); + let targets = stream_continuation_targets_from_targets( + "requested-model", + vec![trusted, untrusted], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + true, + true, + ); + let server = + TestServer::new(strict_stream_continuation_router(targets, mock.clone())).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({"model":"requested-model","prompt":"P: ","stream":true})) + .await; + let body = response.text(); + + assert!(body.contains("Hello")); + assert!(!body.contains("provider-secret")); + assert!(body.contains("Internal server error")); + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_mixed_sanitize_policy_removes_provider_fields() { + let chunk = "data: {\"id\":\"cmpl-first\",\"object\":\"text_completion\",\"created\":1,\"model\":\"provider\",\"choices\":[{\"index\":0,\"text\":\"Hello\",\"finish_reason\":\"stop\"}],\"provider_secret\":\"remove-me\"}\n\ndata:[DONE]\n\n"; + let mock = + MockHttpClient::new_streaming_sequence(StatusCode::OK, vec![vec![chunk.to_string()]]); + let permissive = Target::builder() + .url("https://permissive.example.com/".parse().unwrap()) + .sanitize_response(false) + .build(); + let sanitizing = Target::builder() + .url("https://sanitizing.example.com/".parse().unwrap()) + .sanitize_response(true) + .build(); + let targets = stream_continuation_targets_from_targets( + "requested-model", + vec![permissive, sanitizing], + stream_continuation_fallback(true, 2, 1, None, vec!["/v1/completions"]), + false, + false, + ); + let state = AppState::with_client(targets, mock.clone()) + .with_response_transform(create_openai_sanitizer()); + let server = TestServer::new(build_router(state)).unwrap(); + + let response = server + .post("/v1/completions") + .json(&json!({"model":"requested-model","prompt":"P: ","stream":true})) + .await; + let body = response.text(); + + assert!(body.contains("Hello")); + assert!(body.contains("data:[DONE]")); + assert!(!body.contains("provider_secret")); + assert!(!body.contains("remove-me")); + assert_eq!(mock.get_requests().len(), 1); + } + + #[tokio::test] + async fn stream_continuation_custom_headers_cannot_restore_composite_metadata() { + let forbidden = [ + "content-length", + "content-encoding", + "transfer-encoding", + "trailer", + "connection", + "etag", + "digest", + "content-digest", + "repr-digest", + "representation-digest", + "accept-ranges", + "content-range", + "last-modified", + ]; + let custom_headers = json!({ + "content-length": "1", + "content-encoding": "gzip", + "transfer-encoding": "chunked", + "trailer": "x-checksum", + "connection": "keep-alive", + "etag": "\"custom\"", + "digest": "sha-256=YWJj", + "content-digest": "sha-256=:YWJj:", + "repr-digest": "sha-256=:YWJj:", + "representation-digest": "sha-256=:YWJj:", + "accept-ranges": "bytes", + "content-range": "bytes 0-1/2", + "last-modified": "Mon, 01 Jan 2024 00:00:00 GMT", + "content-type": "application/json", + "x-safe-custom": "retained" + }); + + for scope in ["provider", "pool"] { + let mut provider = json!({"url": "https://provider.example.com/"}); + let mut pool = json!({ + "providers": [provider.clone()], + "fallback": { + "enabled": true, + "max_attempts": 1, + "stream_continuation": { + "enabled": true, + "endpoints": ["/v1/completions"], + "max_attempts": 0, + "max_buffered_bytes": 1024 + } + } + }); + if scope == "provider" { + provider["response_headers"] = custom_headers.clone(); + pool["providers"] = json!([provider]); + } else { + pool["response_headers"] = custom_headers.clone(); + } + let config: target::ConfigFile = serde_json::from_value(json!({ + "targets": {"requested-model": pool}, + "auth": null, + "strict_mode": false + })) + .unwrap(); + let targets = target::Targets::from_config(config).unwrap(); + let terminal = completion_event("cmpl-first", "first", "Hello", "\"stop\""); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![terminal, "data:[DONE]\n\n".to_string()]], + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + + let response = build_router(AppState::with_client(targets, mock)) + .oneshot(request) + .await + .unwrap(); + + assert_eq!( + response.headers().get("content-type").unwrap(), + "text/event-stream", + "scope: {scope}" + ); + assert_eq!( + response.headers().get("x-safe-custom").unwrap(), + "retained", + "scope: {scope}" + ); + assert_eq!(response.headers().get("cache-control").unwrap(), "no-cache"); + for name in forbidden { + assert!( + response.headers().get(name).is_none(), + "scope {scope} restored {name}" + ); + } + let body = response.into_body().collect().await.unwrap().to_bytes(); + assert!(String::from_utf8_lossy(&body).contains("Hello")); + } + } + + #[tokio::test] + async fn stream_continuation_retains_only_response_headers_common_to_all_providers() { + let config: target::ConfigFile = serde_json::from_value(json!({ + "targets": { + "requested-model": { + "response_headers": {"x-common": "retained"}, + "providers": [ + { + "url": "https://first.example.com/", + "response_headers": {"x-provider-price": "1"} + }, + { + "url": "https://second.example.com/", + "response_headers": {"x-provider-price": "2"} + } + ], + "fallback": { + "enabled": true, + "max_attempts": 2, + "stream_continuation": { + "enabled": true, + "endpoints": ["/v1/completions"], + "max_attempts": 1, + "max_buffered_bytes": 1024 + } + } + } + }, + "auth": null, + "strict_mode": false + })) + .unwrap(); + let targets = target::Targets::from_config(config).unwrap(); + let terminal = completion_event("cmpl-first", "first", "Hello", "\"stop\""); + let mock = MockHttpClient::new_streaming_sequence( + StatusCode::OK, + vec![vec![terminal, "data:[DONE]\n\n".to_string()]], + ); + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + + let response = build_router(AppState::with_client(targets, mock)) + .oneshot(request) + .await + .unwrap(); + + assert_eq!(response.headers().get("x-common").unwrap(), "retained"); + assert!(response.headers().get("x-provider-price").is_none()); + let body = response.into_body().collect().await.unwrap().to_bytes(); + assert!(String::from_utf8_lossy(&body).contains("Hello")); + } + + #[tokio::test] + async fn stream_continuation_body_cancellation_releases_provider_guard() { + let mock = + MockHttpClient::new_streaming_response_sequence(vec![MockStreamingResponse::sse( + StatusCode::OK, + vec![MockStreamEvent::Pending], + )]); + let target = Target::builder() + .url("https://stream.example.com/".parse().unwrap()) + .build(); + let pool = ProviderPool::with_config( + vec![Provider::with_concurrency_limit(target, 1, 1)], + None, + None, + None, + Some(stream_continuation_fallback( + true, + 1, + 1, + None, + vec!["/v1/completions"], + )), + LoadBalanceStrategy::Priority, + false, + Vec::new(), + ); + let observed_pool = pool.clone(); + let targets_map = Arc::new(DashMap::new()); + targets_map.insert("requested-model".to_string(), pool); + let targets = target::Targets { + targets: targets_map, + key_rate_limiters: Arc::new(DashMap::new()), + key_concurrency_limiters: Arc::new(DashMap::new()), + key_labels: Arc::new(DashMap::new()), + strict_mode: false, + http_pool_config: None, + }; + let request = axum::extract::Request::builder() + .method("POST") + .uri("/v1/completions") + .header("content-type", "application/json") + .body(axum::body::Body::from( + json!({"model":"requested-model","prompt":"P: ","stream":true}).to_string(), + )) + .unwrap(); + + let response = build_router(AppState::with_client(targets, mock)) + .oneshot(request) + .await + .unwrap(); + assert_eq!(observed_pool.providers()[0].active_connections(), 1); + + drop(response); + + assert_eq!(observed_pool.providers()[0].active_connections(), 0); + } + /// Retry on an upstream 429. Used by the embedded-error tests below. fn embedded_error_targets(alias: &str, n: usize) -> target::Targets { fallback_targets(alias, n, vec![429]) diff --git a/onwards/src/load_balancer.rs b/onwards/src/load_balancer.rs index 06d5b3c66..6c93bd2c2 100644 --- a/onwards/src/load_balancer.rs +++ b/onwards/src/load_balancer.rs @@ -159,13 +159,59 @@ impl ProviderPool { SelectIter { pool: self, - excluded: HashSet::new(), - max_attempts, - attempts: 0, - with_replacement, + state: SelectionState::new(max_attempts, with_replacement), } } + /// Select the next provider using caller-owned fallback selection state. + pub(crate) fn select_next( + &self, + state: &mut SelectionState, + ) -> Option<(usize, &Target, ConcurrencyGuard)> { + if state.attempts >= state.max_attempts { + return None; + } + state.attempts += 1; + + // Ask the LB strategy for the next eligible provider for the current + // exclusions. `select_excluding` returns `None` when no provider is + // eligible — either every provider has been tried in this pass, or the + // untried ones are all at their concurrency limit. + // + // If that happens but the attempt budget still allows it, start a fresh + // pass: clear the exclusions (re-including already-tried providers) and + // cascade through the strategy's options again. This is what puts the + // configured retry budget *above* the LB strategy — every strategy + // (including a single-provider Priority pool) keeps retrying until + // `max_attempts` is spent, rather than stopping after one cascade. + // + // When `excluded` is already empty there is nothing to re-include, so a + // `None` there means an empty pool or every provider at capacity: end + // the selection rather than re-running the same scan. + let result = match self.select_excluding(&state.excluded) { + Some(result) => result, + None if state.excluded.is_empty() => return None, + None => { + state.excluded.clear(); + self.select_excluding(&state.excluded)? + } + }; + + // For priority strategy, exclude the provider just tried so the next + // step advances through the list within this pass. with_replacement + // only applies to weighted random selection (sample with replacement + // within a pass); cross-pass retries are driven by `max_attempts` above. + let should_exclude = match self.strategy { + LoadBalanceStrategy::Priority => true, + LoadBalanceStrategy::WeightedRandom => !state.with_replacement, + }; + if should_exclude { + state.excluded.insert(result.0); + } + + Some(result) + } + /// Internal: select excluding specific provider indices fn select_excluding( &self, @@ -413,58 +459,33 @@ impl ProviderPool { /// ensuring the most up-to-date load information is used for each attempt. pub struct SelectIter<'a> { pool: &'a ProviderPool, + state: SelectionState, +} + +/// Owned state for repeated provider selection. +pub(crate) struct SelectionState { excluded: HashSet, max_attempts: usize, attempts: usize, with_replacement: bool, } +impl SelectionState { + pub(crate) fn new(max_attempts: usize, with_replacement: bool) -> Self { + Self { + excluded: HashSet::new(), + max_attempts, + attempts: 0, + with_replacement, + } + } +} + impl<'a> Iterator for SelectIter<'a> { type Item = (usize, &'a Target, ConcurrencyGuard); fn next(&mut self) -> Option { - if self.attempts >= self.max_attempts { - return None; - } - self.attempts += 1; - - // Ask the LB strategy for the next eligible provider for the current - // exclusions. `select_excluding` returns `None` when no provider is - // eligible — either every provider has been tried in this pass, or the - // untried ones are all at their concurrency limit. - // - // If that happens but the attempt budget still allows it, start a fresh - // pass: clear the exclusions (re-including already-tried providers) and - // cascade through the strategy's options again. This is what puts the - // configured retry budget *above* the LB strategy — every strategy - // (including a single-provider Priority pool) keeps retrying until - // `max_attempts` is spent, rather than stopping after one cascade. - // - // When `excluded` is already empty there is nothing to re-include, so a - // `None` there means an empty pool or every provider at capacity: end the - // iterator rather than re-running the same scan. - let result = match self.pool.select_excluding(&self.excluded) { - Some(result) => result, - None if self.excluded.is_empty() => return None, - None => { - self.excluded.clear(); - self.pool.select_excluding(&self.excluded)? - } - }; - - // For priority strategy, exclude the provider just tried so the next step - // advances through the list within this pass. with_replacement only - // applies to weighted random selection (sample with replacement within a - // pass); cross-pass retries are driven by `max_attempts` above. - let should_exclude = match self.pool.strategy { - LoadBalanceStrategy::Priority => true, - LoadBalanceStrategy::WeightedRandom => !self.with_replacement, - }; - if should_exclude { - self.excluded.insert(result.0); - } - - Some(result) + self.pool.select_next(&mut self.state) } } @@ -478,6 +499,33 @@ mod tests { Target::builder().url(url.parse().unwrap()).build() } + fn priority_pool_with_three_targets() -> ProviderPool { + ProviderPool::with_config( + vec![ + Provider::new(create_test_target("https://p0.example.com"), 1), + Provider::new(create_test_target("https://p1.example.com"), 1), + Provider::new(create_test_target("https://p2.example.com"), 1), + ], + None, + None, + None, + None, + LoadBalanceStrategy::Priority, + false, + Vec::new(), + ) + } + + #[test] + fn test_selection_state_uses_independent_attempt_budget() { + let pool = priority_pool_with_three_targets(); + let mut state = SelectionState::new(2, false); + let first = pool.select_next(&mut state).unwrap().0; + let second = pool.select_next(&mut state).unwrap().0; + assert_eq!((first, second), (0, 1)); + assert!(pool.select_next(&mut state).is_none()); + } + #[test] fn test_single_provider_pool() { let target = create_test_target("https://api.example.com"); diff --git a/onwards/src/main.rs b/onwards/src/main.rs index cbbae894e..4b1ddec68 100644 --- a/onwards/src/main.rs +++ b/onwards/src/main.rs @@ -19,10 +19,10 @@ pub async fn main() -> anyhow::Result<()> { let result = run().await; // Flush pending spans before exit - if let Some(provider) = tracer_provider { - if let Err(e) = provider.shutdown() { - eprintln!("Failed to shutdown tracer provider: {e}"); - } + if let Some(provider) = tracer_provider + && let Err(e) = provider.shutdown() + { + eprintln!("Failed to shutdown tracer provider: {e}"); } result diff --git a/onwards/src/response_loop.rs b/onwards/src/response_loop.rs index 4de35a82b..48475a984 100644 --- a/onwards/src/response_loop.rs +++ b/onwards/src/response_loop.rs @@ -155,6 +155,7 @@ impl From for LoopError { /// `tool_ctx` is the `RequestContext` passed to `ToolExecutor::tools` /// and `::execute` — carries the per-request resolved tool set for /// dwctl's middleware-driven model. +#[allow(clippy::too_many_arguments)] pub fn run_response_loop<'a, S, T, H>( store: &'a S, tool_executor: &'a T, @@ -678,7 +679,7 @@ fn delta_loop_events(event: &StreamEvent, sequence: i64) -> Vec v, Err(_) => return out, }; diff --git a/onwards/src/response_sanitizer.rs b/onwards/src/response_sanitizer.rs index 9647edf36..ba49ea8cb 100644 --- a/onwards/src/response_sanitizer.rs +++ b/onwards/src/response_sanitizer.rs @@ -152,6 +152,39 @@ struct LenientStreamChoice { _extra: HashMap, } +#[derive(Debug, Deserialize, Serialize)] +struct LenientCompletionStreamChunk { + #[serde(default = "default_completion_id")] + id: String, + #[serde(default = "default_completion_object")] + object: String, + #[serde(default = "default_created")] + created: u32, + #[serde(default)] + model: String, + choices: Vec, + #[serde(default)] + usage: Option, + #[serde(default)] + system_fingerprint: Option, + #[serde(flatten, skip_serializing)] + _extra: HashMap, +} + +#[derive(Debug, Deserialize, Serialize)] +struct LenientCompletionStreamChoice { + #[serde(default)] + text: String, + #[serde(default)] + index: u64, + #[serde(default)] + finish_reason: Option, + #[serde(default)] + logprobs: Option, + #[serde(flatten, skip_serializing)] + _extra: HashMap, +} + fn default_id() -> String { "chatcmpl-unknown".to_string() } @@ -164,6 +197,14 @@ fn default_stream_object() -> String { "chat.completion.chunk".to_string() } +fn default_completion_id() -> String { + "cmpl-unknown".to_string() +} + +fn default_completion_object() -> String { + "text_completion".to_string() +} + fn default_created() -> u32 { std::time::SystemTime::now() .duration_since(std::time::UNIX_EPOCH) @@ -211,7 +252,8 @@ impl ResponseSanitizer { let is_streaming = headers .get("content-type") .and_then(|v| v.to_str().ok()) - .map(|v| v.contains("text/event-stream")) + .and_then(|v| v.split(';').next()) + .map(|v| v.trim().eq_ignore_ascii_case("text/event-stream")) .unwrap_or(false); if is_streaming { @@ -258,8 +300,8 @@ impl ResponseSanitizer { let mut sanitized_lines = Vec::new(); for line in body_str.lines() { - if let Some(data_part) = line.strip_prefix("data: ") { - // Skip "data: " prefix + if let Some(data_part) = line.strip_prefix("data:") { + let data_part = data_part.strip_prefix(' ').unwrap_or(data_part); if data_part.trim() == "[DONE]" { // Preserve [DONE] marker as-is @@ -326,6 +368,58 @@ impl ResponseSanitizer { Ok(Some(Bytes::from(sanitized_body))) } + + /// Sanitizes legacy completion SSE chunks while preserving `choices[].text`. + pub fn sanitize_completion_streaming(&self, body: &[u8]) -> Result, String> { + let body_str = std::str::from_utf8(body) + .map_err(|e| format!("Invalid UTF-8 in streaming response: {}", e))?; + let mut sanitized_lines = Vec::new(); + + for line in body_str.lines() { + if let Some(data_part) = line.strip_prefix("data:") { + let data_part = data_part.strip_prefix(' ').unwrap_or(data_part); + if data_part.trim() == "[DONE]" { + sanitized_lines.push(line.to_string()); + continue; + } + + let value: serde_json::Value = serde_json::from_str(data_part) + .map_err(|e| format!("Failed to parse stream chunk: {}", e))?; + if let Some(error_event) = extract_embedded_error_envelope(&value) { + tracing::warn!( + data_len = data_part.len(), + "Provider returned error envelope in completion SSE stream" + ); + sanitized_lines.push(error_event); + continue; + } + + let mut chunk: LenientCompletionStreamChunk = serde_json::from_value(value) + .map_err(|e| format!("Failed to parse completion stream chunk: {}", e))?; + if let Some(ref original) = self.original_model { + chunk.model = original.clone(); + } + let sanitized_json = serde_json::to_string(&chunk) + .map_err(|e| format!("Failed to serialize completion stream chunk: {}", e))?; + sanitized_lines.push(format!("data: {}", sanitized_json)); + } else if line.is_empty() { + sanitized_lines.push(String::new()); + } + } + + let mut sanitized_body = sanitized_lines.join("\n"); + let input_trailing = body_str.chars().rev().take_while(|&c| c == '\n').count(); + let output_trailing = sanitized_body + .chars() + .rev() + .take_while(|&c| c == '\n') + .count(); + for _ in output_trailing..input_trailing { + sanitized_body.push('\n'); + } + + Ok(Some(Bytes::from(sanitized_body))) + } } /// If `value` carries a provider `error` object — either as the entire @@ -481,6 +575,20 @@ mod tests { ); } + #[test] + fn streaming_sanitizers_preserve_done_without_a_post_colon_space() { + let sanitizer = ResponseSanitizer { + original_model: None, + }; + + for result in [ + sanitizer.sanitize_streaming(b"data:[DONE]\n\n"), + sanitizer.sanitize_completion_streaming(b"data:[DONE]\n\n"), + ] { + assert_eq!(result.unwrap().unwrap().as_ref(), b"data:[DONE]\n\n"); + } + } + #[test] fn test_streaming_multiple_chunks() { let sanitizer = ResponseSanitizer { diff --git a/onwards/src/sse.rs b/onwards/src/sse.rs index 3c229285f..a9b606f6a 100644 --- a/onwards/src/sse.rs +++ b/onwards/src/sse.rs @@ -1,12 +1,14 @@ //! SSE (Server-Sent Events) stream buffering //! //! This module provides a stream wrapper that buffers incomplete SSE events. -//! Some AI providers send partial chunks that split JSON data across multiple -//! network packets. This buffer accumulates bytes until a complete SSE event -//! (terminated by `\n\n`) is received before forwarding. +//! Some AI providers send partial chunks that split JSON data and line endings +//! across network packets. This buffer accumulates bytes until a complete SSE +//! event is received before forwarding it unchanged. use bytes::{Bytes, BytesMut}; use futures_util::Stream; +use std::error::Error; +use std::fmt; use std::pin::Pin; use std::task::{Context, Poll}; @@ -18,75 +20,271 @@ use std::task::{Context, Poll}; /// ~64MB for 1000 concurrent streams. const MAX_SSE_BUFFER_SIZE: usize = 64 * 1024; -/// A stream wrapper that buffers SSE events until they are complete. -/// -/// SSE events are delimited by `\n\n`. This wrapper accumulates incoming -/// bytes and only yields complete events, preventing consumers from -/// receiving partial JSON data. -/// -/// The buffer is capped at [`MAX_SSE_BUFFER_SIZE`] bytes to prevent memory -/// exhaustion from malicious or buggy upstream providers. -pub struct SseBufferedStream { +/// A framing failure detected while reassembling an SSE stream. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SseFramingError { + /// The source ended with bytes that did not form a complete SSE event. + IncompleteEvent, + /// A single event exceeded the bounded reassembly buffer. + BufferOverflow { limit: usize }, +} + +impl fmt::Display for SseFramingError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::IncompleteEvent => { + write!(f, "upstream SSE stream ended with an incomplete event") + } + Self::BufferOverflow { limit } => { + write!(f, "upstream SSE event exceeded the {limit}-byte limit") + } + } + } +} + +impl Error for SseFramingError {} + +/// An error from either the source body or SSE framing validation. +#[derive(Debug)] +pub enum SseStreamError { + Source(E), + Framing(SseFramingError), +} + +impl fmt::Display for SseStreamError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Source(error) => error.fmt(f), + Self::Framing(error) => error.fmt(f), + } + } +} + +impl Error for SseStreamError { + fn source(&self) -> Option<&(dyn Error + 'static)> { + match self { + Self::Source(error) => Some(error), + Self::Framing(error) => Some(error), + } + } +} + +#[derive(Default)] +struct EventBoundaryScanner { + line_has_bytes: bool, + pending_cr: bool, +} + +enum ScanOutcome { + Boundary(usize), + NeedMore(usize), + Limit, +} + +impl EventBoundaryScanner { + fn scan(&mut self, input: &[u8], remaining: usize) -> ScanOutcome { + let mut consumed = 0; + + while consumed < input.len() { + let byte = input[consumed]; + if self.pending_cr { + if byte == b'\n' { + if consumed == remaining { + return ScanOutcome::Limit; + } + let blank_line = !self.line_has_bytes; + self.pending_cr = false; + self.line_has_bytes = false; + consumed += 1; + if blank_line { + return ScanOutcome::Boundary(consumed); + } + continue; + } + + let blank_line = !self.line_has_bytes; + self.pending_cr = false; + self.line_has_bytes = false; + if blank_line { + return ScanOutcome::Boundary(consumed); + } + } + + if consumed == remaining { + return ScanOutcome::Limit; + } + + match byte { + b'\r' => self.pending_cr = true, + b'\n' => { + let blank_line = !self.line_has_bytes; + self.line_has_bytes = false; + consumed += 1; + if blank_line { + return ScanOutcome::Boundary(consumed); + } + continue; + } + _ => self.line_has_bytes = true, + } + consumed += 1; + } + + ScanOutcome::NeedMore(consumed) + } + + fn finish_eof(&mut self) -> bool { + if !self.pending_cr { + return false; + } + let blank_line = !self.line_has_bytes; + self.pending_cr = false; + self.line_has_bytes = false; + blank_line + } +} + +/// Internal SSE framing that keeps source failures distinct from framing failures. +pub(crate) struct CheckedSseStream { inner: S, + max_buffer_size: usize, buffer: BytesMut, + pending_chunk: Option, + pending_offset: usize, + scanner: EventBoundaryScanner, + source_done: bool, + terminated: bool, + incomplete_bytes: Option, + #[cfg(test)] + scanned_bytes: usize, } -impl SseBufferedStream { - /// Wrap an existing stream with SSE buffering. - pub fn new(inner: S) -> Self { +impl CheckedSseStream { + pub(crate) fn new(inner: S) -> Self { + Self::with_max_buffer_size(inner, MAX_SSE_BUFFER_SIZE) + } + + pub(crate) fn with_max_buffer_size(inner: S, max_buffer_size: usize) -> Self { Self { inner, + max_buffer_size, buffer: BytesMut::new(), + pending_chunk: None, + pending_offset: 0, + scanner: EventBoundaryScanner::default(), + source_done: false, + terminated: false, + incomplete_bytes: None, + #[cfg(test)] + scanned_bytes: 0, } } + + fn append_pending(&mut self) -> ScanOutcome { + let chunk = self.pending_chunk.as_ref().expect("pending chunk checked"); + let input = &chunk[self.pending_offset..]; + let outcome = self.scanner.scan( + input, + self.max_buffer_size.saturating_sub(self.buffer.len()), + ); + let consumed = match outcome { + ScanOutcome::Boundary(consumed) | ScanOutcome::NeedMore(consumed) => consumed, + ScanOutcome::Limit => 0, + }; + + if consumed > 0 { + self.buffer.extend_from_slice(&input[..consumed]); + self.pending_offset += consumed; + #[cfg(test)] + { + self.scanned_bytes += consumed; + } + } + if self.pending_offset == chunk.len() { + self.pending_chunk = None; + self.pending_offset = 0; + } + outcome + } + + fn take_incomplete_bytes(&mut self) -> Option { + self.incomplete_bytes.take() + } + + #[cfg(test)] + fn scanned_bytes_for_test(&self) -> usize { + self.scanned_bytes + } + + #[cfg(test)] + fn buffer_capacity_for_test(&self) -> usize { + self.buffer.capacity() + } } -impl Stream for SseBufferedStream +impl Stream for CheckedSseStream where S: Stream> + Unpin, { - type Item = Result; + type Item = Result>; fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { let this = &mut *self; loop { - // Check if buffer contains a complete event (ends with \n\n) - if let Some(pos) = find_event_boundary(&this.buffer) { - // Extract the complete event(s) up to and including the \n\n - let complete = this.buffer.split_to(pos + 2); - return Poll::Ready(Some(Ok(complete.freeze()))); + if this.terminated { + return Poll::Ready(None); } - // Need more data - poll the inner stream - match Pin::new(&mut this.inner).poll_next(cx) { - Poll::Ready(Some(Ok(chunk))) => { - this.buffer.extend_from_slice(&chunk); - - // Check buffer size limit to prevent memory exhaustion - // If exceeded, log error and terminate stream by returning None - if this.buffer.len() > MAX_SSE_BUFFER_SIZE { - tracing::error!( - "SSE buffer exceeded maximum size of {} bytes, terminating stream", - MAX_SSE_BUFFER_SIZE - ); + if this.pending_chunk.is_some() { + match this.append_pending() { + ScanOutcome::Boundary(_) => { + this.scanner = EventBoundaryScanner::default(); + return Poll::Ready(Some(Ok(this.buffer.split().freeze()))); + } + ScanOutcome::NeedMore(_) => continue, + ScanOutcome::Limit => { this.buffer.clear(); - return Poll::Ready(None); + this.pending_chunk = None; + this.terminated = true; + return Poll::Ready(Some(Err(SseStreamError::Framing( + SseFramingError::BufferOverflow { + limit: this.max_buffer_size, + }, + )))); } + } + } - // Loop back to check for complete events + if this.source_done { + if this.buffer.is_empty() { + this.terminated = true; + return Poll::Ready(None); + } + if this.scanner.finish_eof() { + this.scanner = EventBoundaryScanner::default(); + return Poll::Ready(Some(Ok(this.buffer.split().freeze()))); + } + this.incomplete_bytes = Some(this.buffer.split().freeze()); + this.terminated = true; + return Poll::Ready(Some(Err(SseStreamError::Framing( + SseFramingError::IncompleteEvent, + )))); + } + + match Pin::new(&mut this.inner).poll_next(cx) { + Poll::Ready(Some(Ok(chunk))) => { + if !chunk.is_empty() { + this.pending_chunk = Some(chunk); + } } Poll::Ready(Some(Err(e))) => { - return Poll::Ready(Some(Err(e))); + this.buffer.clear(); + this.terminated = true; + return Poll::Ready(Some(Err(SseStreamError::Source(e)))); } Poll::Ready(None) => { - // Stream ended - flush any remaining data - if this.buffer.is_empty() { - return Poll::Ready(None); - } - // Return whatever is left (may be incomplete, but stream is done) - let remaining = this.buffer.split().freeze(); - return Poll::Ready(Some(Ok(remaining))); + this.source_done = true; } Poll::Pending => { return Poll::Pending; @@ -96,9 +294,137 @@ where } } -/// Find the position of `\n\n` in the buffer, returning the index of the first `\n`. -fn find_event_boundary(buf: &[u8]) -> Option { - buf.windows(2).position(|window| window == b"\n\n") +/// A stream wrapper that buffers SSE events until they are complete. +/// +/// This public wrapper preserves its original source error type. Internal proxy +/// paths use checked framing when incomplete or oversized events must be +/// distinguishable from source EOF. +pub struct SseBufferedStream { + checked: CheckedSseStream, +} + +impl SseBufferedStream { + /// Wrap an existing stream with SSE buffering. + pub fn new(inner: S) -> Self { + Self { + checked: CheckedSseStream::new(inner), + } + } +} + +impl Stream for SseBufferedStream +where + S: Stream> + Unpin, +{ + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + match Pin::new(&mut self.checked).poll_next(cx) { + Poll::Ready(Some(Ok(event))) => Poll::Ready(Some(Ok(event))), + Poll::Ready(Some(Err(SseStreamError::Source(error)))) => Poll::Ready(Some(Err(error))), + Poll::Ready(Some(Err(SseStreamError::Framing(SseFramingError::IncompleteEvent)))) => { + Poll::Ready(self.checked.take_incomplete_bytes().map(Ok)) + } + Poll::Ready(Some(Err(SseStreamError::Framing(SseFramingError::BufferOverflow { + .. + })))) => Poll::Ready(None), + Poll::Ready(None) => Poll::Ready(None), + Poll::Pending => Poll::Pending, + } + } +} + +pub(crate) fn framing_error_in_chain(error: &(dyn Error + 'static)) -> Option { + let mut current = Some(error); + while let Some(source) = current { + if let Some(framing) = source.downcast_ref::() { + return Some(*framing); + } + current = source.source(); + } + None +} + +#[derive(Debug, PartialEq, Eq)] +pub(crate) enum ParsedSseEvent { + Comment, + Data { + data: Vec, + event_type: Option>, + }, + Invalid, +} + +/// Parse one framed event without normalizing the bytes forwarded downstream. +pub(crate) fn parse_sse_event(event: &[u8]) -> ParsedSseEvent { + let mut position = if event.starts_with(b"\xef\xbb\xbf") { + 3 + } else { + 0 + }; + let mut data = Vec::new(); + let mut found_data = false; + let mut event_type = None; + + while position < event.len() { + let line_start = position; + while position < event.len() && !matches!(event[position], b'\r' | b'\n') { + position += 1; + } + let line = &event[line_start..position]; + + if position < event.len() { + if event[position] == b'\r' + && position + 1 < event.len() + && event[position + 1] == b'\n' + { + position += 2; + } else { + position += 1; + } + } + + if line.is_empty() || line.starts_with(b":") { + continue; + } + + if line == b"event" { + if event_type.replace(Vec::new()).is_some() { + return ParsedSseEvent::Invalid; + } + continue; + } + if let Some(value) = line.strip_prefix(b"event:") { + if event_type + .replace(value.strip_prefix(b" ").unwrap_or(value).to_vec()) + .is_some() + { + return ParsedSseEvent::Invalid; + } + continue; + } + + let value = if line == b"data" { + &b""[..] + } else if let Some(value) = line.strip_prefix(b"data:") { + value.strip_prefix(b" ").unwrap_or(value) + } else { + return ParsedSseEvent::Invalid; + }; + if found_data { + data.push(b'\n'); + } + found_data = true; + data.extend_from_slice(value); + } + + if found_data { + ParsedSseEvent::Data { data, event_type } + } else if event_type.is_some() { + ParsedSseEvent::Invalid + } else { + ParsedSseEvent::Comment + } } #[cfg(test)] @@ -114,6 +440,72 @@ mod tests { futures_util::stream::iter(chunks.into_iter().map(|c| Ok(Bytes::from_static(c)))) } + #[test] + fn public_buffered_stream_preserves_source_error_item_type() { + fn assert_item_type(_stream: &SseBufferedStream) + where + S: Stream> + Unpin, + SseBufferedStream: Stream>, + { + } + + let stream = SseBufferedStream::new(chunks_to_stream(Vec::new())); + assert_item_type::<_, Infallible>(&stream); + } + + #[tokio::test] + async fn checked_stream_handles_highly_fragmented_input_in_linear_state() { + let expected = Bytes::from_static(b"data: fragmented\r\n\r\n"); + let chunks = expected + .iter() + .copied() + .map(|byte| Ok::<_, Infallible>(Bytes::from(vec![byte]))); + let mut stream = CheckedSseStream::new(futures_util::stream::iter(chunks)); + + let event = stream.next().await.unwrap().unwrap(); + assert_eq!(event, expected); + assert_eq!(stream.scanned_bytes_for_test(), expected.len()); + assert!(stream.next().await.is_none()); + } + + #[tokio::test] + async fn checked_stream_does_not_copy_a_huge_source_chunk() { + let huge = Bytes::from(vec![b'x'; MAX_SSE_BUFFER_SIZE * 8]); + let mut stream = + CheckedSseStream::new(futures_util::stream::iter([Ok::<_, Infallible>(huge)])); + + assert!(matches!( + stream.next().await, + Some(Err(SseStreamError::Framing( + SseFramingError::BufferOverflow { + limit: MAX_SSE_BUFFER_SIZE + } + ))) + )); + assert!(stream.buffer_capacity_for_test() <= MAX_SSE_BUFFER_SIZE); + } + + #[tokio::test] + async fn checked_stream_copies_only_through_a_boundary_in_a_huge_chunk() { + let mut source = b"data: first\n\n".to_vec(); + source.extend(std::iter::repeat_n(b'x', MAX_SSE_BUFFER_SIZE * 8)); + let mut stream = CheckedSseStream::new(futures_util::stream::iter([Ok::<_, Infallible>( + Bytes::from(source), + )])); + + assert_eq!( + stream.next().await.unwrap().unwrap().as_ref(), + b"data: first\n\n" + ); + assert!(stream.buffer_capacity_for_test() < MAX_SSE_BUFFER_SIZE); + assert!(matches!( + stream.next().await, + Some(Err(SseStreamError::Framing( + SseFramingError::BufferOverflow { .. } + ))) + )); + } + #[tokio::test] async fn test_complete_event_passes_through() { let chunks = vec![b"data: {\"hello\": \"world\"}\n\n".as_slice()]; @@ -182,14 +574,16 @@ mod tests { } #[tokio::test] - async fn test_incomplete_event_at_stream_end() { - // Stream ends without final \n\n + async fn test_incomplete_event_at_stream_end_is_a_framing_error() { let chunks = vec![b"data: incomplete".as_slice()]; - let stream = SseBufferedStream::new(chunks_to_stream(chunks)); + let stream = CheckedSseStream::new(chunks_to_stream(chunks)); let results: Vec<_> = stream.collect().await; assert_eq!(results.len(), 1); - assert_eq!(results[0].as_ref().unwrap().as_ref(), b"data: incomplete"); + assert!(matches!( + results[0], + Err(SseStreamError::Framing(SseFramingError::IncompleteEvent)) + )); } #[tokio::test] @@ -224,18 +618,69 @@ mod tests { } #[tokio::test] - async fn test_handles_crlf_events() { - // \r\n\r\n does NOT contain \n\n (it's [0d 0a 0d 0a], not [0a 0a]) - // So we only flush at end of stream. Real SSE servers that use CRLF - // typically send \r\n\r\n which our buffer treats as incomplete until EOF. - // This is acceptable since the data will be flushed when stream ends. - let chunks = vec![b"data: test\r\n\r\n".as_slice()]; + async fn test_fragmented_lf_crlf_cr_and_mixed_boundaries_preserve_exact_bytes() { + for (chunks, expected) in [ + ( + vec![b"data: lf\n".as_slice(), b"\n".as_slice()], + b"data: lf\n\n".as_slice(), + ), + ( + vec![ + b"data: crlf\r".as_slice(), + b"\n\r".as_slice(), + b"\n".as_slice(), + ], + b"data: crlf\r\n\r\n".as_slice(), + ), + ( + vec![b"data: cr\r".as_slice(), b"\r".as_slice()], + b"data: cr\r\r".as_slice(), + ), + ( + vec![b"data: mixed\r\n\r".as_slice(), b"\n".as_slice()], + b"data: mixed\r\n\r\n".as_slice(), + ), + ] { + let stream = SseBufferedStream::new(chunks_to_stream(chunks)); + let results: Vec<_> = stream.collect().await; + + assert_eq!(results.len(), 1); + assert_eq!(results[0].as_ref().unwrap().as_ref(), expected); + } + } + + #[tokio::test] + async fn test_crlf_event_is_available_before_source_eof() { + let source = futures_util::stream::once(async { + Ok::<_, Infallible>(Bytes::from_static(b"data: ready\r\n\r\n")) + }) + .chain(futures_util::stream::pending()); + let mut stream = SseBufferedStream::new(Box::pin(source)); + + let event = tokio::time::timeout(std::time::Duration::from_millis(50), stream.next()) + .await + .expect("complete CRLF event should not wait for EOF") + .expect("event") + .expect("valid frame"); + + assert_eq!(event.as_ref(), b"data: ready\r\n\r\n"); + } + + #[tokio::test] + async fn test_utf8_bom_is_preserved_at_stream_start() { + let chunks = vec![ + b"\xef".as_slice(), + b"\xbb\xbfdata:value\r".as_slice(), + b"\r".as_slice(), + ]; let stream = SseBufferedStream::new(chunks_to_stream(chunks)); let results: Vec<_> = stream.collect().await; - // No \n\n found, so entire content flushed at stream end assert_eq!(results.len(), 1); - assert_eq!(results[0].as_ref().unwrap().as_ref(), b"data: test\r\n\r\n"); + assert_eq!( + results[0].as_ref().unwrap().as_ref(), + b"\xef\xbb\xbfdata:value\r\r" + ); } #[tokio::test] @@ -253,19 +698,39 @@ mod tests { } #[tokio::test] - async fn test_buffer_overflow_terminates_stream() { + async fn test_buffer_overflow_is_a_framing_error() { // Create a chunk larger than MAX_SSE_BUFFER_SIZE without \n\n let large_chunk = vec![b'x'; MAX_SSE_BUFFER_SIZE + 1]; let chunks: Vec<&[u8]> = vec![&large_chunk]; - let stream = SseBufferedStream::new(futures_util::stream::iter( + let stream = CheckedSseStream::new(futures_util::stream::iter( chunks .into_iter() .map(|c| Ok::<_, Infallible>(Bytes::from(c.to_vec()))), )); let results: Vec<_> = stream.collect().await; - // Stream terminates immediately when buffer exceeded (returns None) - assert_eq!(results.len(), 0); + assert_eq!(results.len(), 1); + assert!(matches!( + results[0], + Err(SseStreamError::Framing(SseFramingError::BufferOverflow { + limit: MAX_SSE_BUFFER_SIZE + })) + )); + } + + #[tokio::test] + async fn test_source_body_error_remains_distinct_from_framing_errors() { + let source = futures_util::stream::iter(vec![Err::(std::io::Error::other( + "source failed", + ))]); + let results: Vec<_> = CheckedSseStream::new(source).collect().await; + + assert_eq!(results.len(), 1); + assert!(matches!(results[0], Err(SseStreamError::Source(_)))); + assert_eq!( + results[0].as_ref().unwrap_err().to_string(), + "source failed" + ); } #[tokio::test] diff --git a/onwards/src/stream_continuation.rs b/onwards/src/stream_continuation.rs new file mode 100644 index 000000000..af7cebcad --- /dev/null +++ b/onwards/src/stream_continuation.rs @@ -0,0 +1,2884 @@ +use axum::{ + body::Body, + http::{ + HeaderMap, Method, + header::{CONTENT_ENCODING, CONTENT_TYPE}, + }, +}; +use bytes::Bytes; +use futures_util::StreamExt; +use serde_json::Value; +use std::fmt; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use crate::client::HttpClient; +use crate::handlers::{UpstreamRequestMetadata, build_upstream_request}; +use crate::load_balancer::{ProviderPool, SelectionState}; +use crate::sse::{ + CheckedSseStream, ParsedSseEvent, SseFramingError, SseStreamError, framing_error_in_chain, + parse_sse_event, +}; +use crate::target::ConcurrencyGuard; +use crate::target::StreamContinuationConfig; + +const CONTINUATION_MAX_SSE_EVENT_BYTES: usize = 8 * 1024 * 1024; + +pub(crate) fn event_buffer_size() -> usize { + CONTINUATION_MAX_SSE_EVENT_BYTES +} + +/// State retained while forwarding an eligible text-generation stream. +pub struct StreamContinuation { + request: Value, + protocol: ContinuationProtocol, + generated_text: String, + max_buffered_bytes: usize, + continuable: bool, + terminal: bool, +} + +enum ContinuationProtocol { + Completion(TextStreamIdentity), + Chat(TextStreamIdentity), + Responses(ResponsesStreamState), +} + +struct TextStreamIdentity { + id: Option, + model: Option, + created: Option, + identity_established: bool, +} + +struct ResponsesStreamState { + response_id: String, + model: String, + item_id: String, + response_identity_established: bool, + item_identity_established: bool, + last_sequence: Option, + emit_event_field: Option, + continuation_prefix_len: Option, +} + +pub(crate) fn is_event_stream(headers: &HeaderMap) -> bool { + let mut content_types = headers.get_all(CONTENT_TYPE).iter(); + let (Some(content_type), None) = (content_types.next(), content_types.next()) else { + return false; + }; + content_type + .to_str() + .ok() + .and_then(|value| { + value + .parse::() + .ok() + .map(|media_type| (value, media_type)) + }) + .is_some_and(|(raw_value, media_type)| { + media_type.type_() == mime::TEXT + && media_type + .subtype() + .as_str() + .eq_ignore_ascii_case("event-stream") + && (!raw_value.contains(';') || media_type.params().next().is_some()) + && media_type + .params() + .all(|(_, value)| !value.as_str().is_empty()) + }) +} + +pub(crate) fn has_identity_content_encoding(headers: &HeaderMap) -> bool { + let mut encodings = headers.get_all(CONTENT_ENCODING).iter(); + match (encodings.next(), encodings.next()) { + (None, None) => true, + (Some(value), None) => value + .to_str() + .is_ok_and(|encoding| encoding.trim().eq_ignore_ascii_case("identity")), + _ => false, + } +} + +/// The forwarded SSE event and whether the stream reached a terminal state. +pub struct EventObservation { + pub event: Bytes, + pub forward: bool, + pub terminal: bool, + pub done: bool, + pub safe: bool, + pub accepted: bool, +} + +/// Errors raised while building a continuation request or rewriting an event. +#[derive(Debug)] +pub enum ContinuationError { + NotContinuable, + Serialize(serde_json::Error), +} + +impl fmt::Display for ContinuationError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::NotContinuable => write!(f, "generation stream cannot be continued"), + Self::Serialize(_) => write!(f, "failed to serialize stream continuation data"), + } + } +} + +impl std::error::Error for ContinuationError { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + match self { + Self::NotContinuable => None, + Self::Serialize(error) => Some(error), + } + } +} + +impl StreamContinuation { + /// Creates protocol state when a request can be safely resumed from text output. + pub fn from_request( + path: &str, + method: &Method, + body: &[u8], + config: &StreamContinuationConfig, + ) -> Option { + Self::from_request_with_resolved_model(path, path, method, body, config, None) + } + + pub(crate) fn from_request_with_resolved_model( + eligible_path: &str, + protocol_path: &str, + method: &Method, + body: &[u8], + config: &StreamContinuationConfig, + resolved_model: Option<&str>, + ) -> Option { + if method != Method::POST || !config.enabled_for_path(eligible_path) { + return None; + } + + let request: Value = serde_json::from_slice(body).ok()?; + let request_object = request.as_object()?; + let fallback_model = resolved_model + .map(|model| Value::String(model.to_owned())) + .or_else(|| request.get("model").cloned()) + .unwrap_or_else(|| Value::String("unknown".to_string())); + let fallback_created = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs(); + let protocol = match protocol_path { + "/v1/completions" if eligible_completion_request(request_object) => { + ContinuationProtocol::Completion(TextStreamIdentity::new( + "cmpl", + fallback_model, + fallback_created, + )) + } + "/v1/chat/completions" if eligible_chat_request(request_object) => { + ContinuationProtocol::Chat(TextStreamIdentity::new( + "chatcmpl", + fallback_model, + fallback_created, + )) + } + "/v1/responses" if eligible_responses_request(request_object) => { + ContinuationProtocol::Responses(ResponsesStreamState::new( + fallback_model.as_str().unwrap_or("unknown").to_string(), + )) + } + _ => return None, + }; + Some(Self { + request, + protocol, + generated_text: String::new(), + max_buffered_bytes: config.max_buffered_bytes, + continuable: true, + terminal: false, + }) + } + + /// Records a complete SSE event and optionally normalizes its stream identity. + pub fn observe_event( + &mut self, + event: &[u8], + rewrite_identity: bool, + ) -> Result { + let (data, event_type) = match parse_sse_event(event) { + ParsedSseEvent::Comment => { + return Ok(EventObservation { + event: Bytes::copy_from_slice(event), + forward: true, + terminal: self.terminal, + done: false, + safe: true, + accepted: false, + }); + } + ParsedSseEvent::Data { data, event_type } => (data, event_type), + ParsedSseEvent::Invalid => { + return Ok(self.reject_event(event)); + } + }; + + let Ok(data) = std::str::from_utf8(&data) else { + return Ok(self.reject_event(event)); + }; + if data == "[DONE]" { + if event_type.is_some() { + return Ok(self.reject_event(event)); + } + self.terminal = true; + return Ok(EventObservation { + event: Bytes::copy_from_slice(event), + forward: true, + terminal: true, + done: true, + safe: true, + accepted: true, + }); + } + + let Ok(mut value) = serde_json::from_str::(data) else { + return Ok(self.reject_event(event)); + }; + let outcome = match &mut self.protocol { + ContinuationProtocol::Completion(identity) => observe_completion_event( + identity, + &mut value, + event_type.as_deref(), + rewrite_identity, + ), + ContinuationProtocol::Chat(identity) => observe_chat_event( + identity, + &mut value, + event_type.as_deref(), + rewrite_identity, + ), + ContinuationProtocol::Responses(state) => observe_responses_event( + state, + &mut value, + event_type.as_deref(), + rewrite_identity, + &self.generated_text, + ), + }; + let ProtocolObservation { + text, + terminal, + accepted, + forward, + event, + } = match outcome { + Some(outcome) => outcome, + None => return Ok(self.reject_event(event)), + }; + + if let Some(text) = text.as_deref() { + append_generated_text( + &mut self.generated_text, + &mut self.continuable, + self.max_buffered_bytes, + text, + ); + } + self.terminal |= terminal; + + Ok(EventObservation { + event, + forward, + terminal: self.terminal, + done: false, + safe: true, + accepted, + }) + } + + /// Returns whether an interrupted response may issue another generation request. + pub fn is_continuable(&self) -> bool { + self.continuable && !self.terminal + } + + pub fn is_terminal(&self) -> bool { + self.terminal + } + + fn begin_continuation_attempt(&mut self) { + if let ContinuationProtocol::Responses(state) = &mut self.protocol { + state.continuation_prefix_len = Some(self.generated_text.len()); + } + } + + /// Builds the next protocol request with the emitted text supplied as a prefix. + pub fn request_body(&self, model_override: Option<&str>) -> Result { + if !self.is_continuable() { + return Err(ContinuationError::NotContinuable); + } + + let mut request = self.request.clone(); + let request = request + .as_object_mut() + .expect("eligible continuation requests are JSON objects"); + match self.protocol { + ContinuationProtocol::Completion(_) => { + let prompt = request + .get("prompt") + .and_then(Value::as_str) + .expect("eligible completion prompt is a string"); + request.insert( + "prompt".to_owned(), + Value::String(format!("{prompt}{}", self.generated_text)), + ); + } + ContinuationProtocol::Chat(_) => { + request + .get_mut("messages") + .and_then(Value::as_array_mut) + .expect("eligible chat messages are an array") + .push(serde_json::json!({ + "role": "assistant", + "content": self.generated_text, + })); + } + ContinuationProtocol::Responses(_) => { + let original_input = request + .remove("input") + .expect("eligible Responses input is present"); + let mut input = match original_input { + Value::String(text) => vec![serde_json::json!({ + "type": "message", + "role": "user", + "content": text, + })], + Value::Array(items) => items, + _ => unreachable!("eligible Responses input is text or an array"), + }; + input.push(serde_json::json!({ + "type": "message", + "role": "assistant", + "content": [{ + "type": "output_text", + "text": self.generated_text, + }], + })); + request.insert("input".to_owned(), Value::Array(input)); + } + } + if let Some(model) = model_override { + request.insert("model".to_owned(), Value::String(model.to_owned())); + } + + serde_json::to_vec(&request) + .map(Bytes::from) + .map_err(ContinuationError::Serialize) + } + + fn reject_event(&mut self, event: &[u8]) -> EventObservation { + self.continuable = false; + EventObservation { + event: Bytes::copy_from_slice(event), + forward: true, + terminal: self.terminal, + done: false, + safe: false, + accepted: false, + } + } +} + +struct ProtocolObservation { + text: Option, + terminal: bool, + accepted: bool, + forward: bool, + event: Bytes, +} + +impl TextStreamIdentity { + fn new(prefix: &str, model: Value, created: u64) -> Self { + Self { + id: Some(Value::String(format!("{prefix}-{}", uuid::Uuid::new_v4()))), + model: Some(model), + created: Some(Value::from(created)), + identity_established: false, + } + } + + fn establish(&mut self, value: &Value) { + self.id = value.get("id").cloned().or_else(|| self.id.take()); + self.model = value.get("model").cloned().or_else(|| self.model.take()); + self.created = value + .get("created") + .cloned() + .or_else(|| self.created.take()); + self.identity_established = true; + } + + fn rewrite(&self, value: &mut Value) { + let Some(object) = value.as_object_mut() else { + return; + }; + if let Some(id) = &self.id { + object.insert("id".to_owned(), id.clone()); + } + if let Some(model) = &self.model { + object.insert("model".to_owned(), model.clone()); + } + if let Some(created) = &self.created { + object.insert("created".to_owned(), created.clone()); + } + } +} + +impl ResponsesStreamState { + fn new(model: String) -> Self { + Self { + response_id: format!("resp_{}", uuid::Uuid::new_v4()), + model, + item_id: format!("msg_{}", uuid::Uuid::new_v4()), + response_identity_established: false, + item_identity_established: false, + last_sequence: None, + emit_event_field: None, + continuation_prefix_len: None, + } + } +} + +fn eligible_completion_request(request: &serde_json::Map) -> bool { + !has_unsupported_controls( + request, + &[ + "tools", + "tool_choice", + "functions", + "function_call", + "response_format", + "json_schema", + "grammar", + ], + ) && request.get("prompt").is_some_and(Value::is_string) + && request.get("stream") == Some(&Value::Bool(true)) + && supports_single_choice(request.get("n")) + && matches!(request.get("echo"), None | Some(Value::Bool(false))) + && matches!(request.get("logprobs"), None | Some(Value::Null)) +} + +fn eligible_chat_request(request: &serde_json::Map) -> bool { + !has_unsupported_controls( + request, + &[ + "tools", + "tool_choice", + "functions", + "function_call", + "response_format", + "json_schema", + "grammar", + "modalities", + "audio", + "parallel_tool_calls", + "reasoning", + "reasoning_effort", + "include_reasoning", + "prediction", + ], + ) && request.get("messages").is_some_and(eligible_chat_messages) + && request.get("stream") == Some(&Value::Bool(true)) + && supports_single_choice(request.get("n")) + && matches!(request.get("logprobs"), None | Some(Value::Bool(false))) + && !request.contains_key("top_logprobs") +} + +fn eligible_responses_request(request: &serde_json::Map) -> bool { + let input_supported = request.get("input").is_some_and(eligible_responses_input); + let plain_text = match request.get("text") { + None => true, + Some(Value::Object(text)) => match text.get("format") { + None => true, + Some(Value::Object(format)) => { + format.get("type").and_then(Value::as_str) == Some("text") + } + Some(_) => false, + }, + Some(_) => false, + }; + input_supported + && plain_text + && request.get("stream") == Some(&Value::Bool(true)) + && matches!(request.get("background"), None | Some(Value::Bool(false))) + && matches!(request.get("store"), None | Some(Value::Bool(false))) + && !has_unsupported_controls( + request, + &[ + "tools", + "tool_choice", + "max_tool_calls", + "parallel_tool_calls", + "reasoning", + "previous_response_id", + "conversation", + "include", + "json_schema", + "grammar", + "top_logprobs", + "logprobs", + ], + ) +} + +fn eligible_chat_messages(value: &Value) -> bool { + const ALLOWED_FIELDS: &[&str] = &["role", "content", "name"]; + value.as_array().is_some_and(|messages| { + !messages.is_empty() + && messages.iter().all(|message| { + let Some(message) = message.as_object() else { + return false; + }; + message + .keys() + .all(|key| ALLOWED_FIELDS.contains(&key.as_str())) + && matches!( + message.get("role").and_then(Value::as_str), + Some("system" | "developer" | "user" | "assistant") + ) + && message.get("content").is_some_and(eligible_chat_content) + }) + }) +} + +fn eligible_chat_content(value: &Value) -> bool { + match value { + Value::String(_) => true, + Value::Array(parts) => { + !parts.is_empty() + && parts + .iter() + .all(|part| match part.get("type").and_then(Value::as_str) { + Some("text") => { + keys_allowed(part, &["type", "text"]) + && part.get("text").is_some_and(Value::is_string) + } + Some("image_url") => { + keys_allowed(part, &["type", "image_url"]) + && part.get("image_url").is_some_and(Value::is_object) + } + _ => false, + }) + } + _ => false, + } +} + +fn eligible_responses_input(value: &Value) -> bool { + match value { + Value::String(_) => true, + Value::Array(items) => { + !items.is_empty() + && items.iter().all(|item| { + let Some(item) = item.as_object() else { + return false; + }; + item.keys().all(|key| { + ["type", "id", "role", "content", "status"].contains(&key.as_str()) + }) && item + .get("type") + .is_none_or(|item_type| item_type.as_str() == Some("message")) + && matches!( + item.get("role").and_then(Value::as_str), + Some("system" | "developer" | "user" | "assistant") + ) + && item + .get("content") + .is_some_and(eligible_responses_message_content) + }) + } + _ => false, + } +} + +fn eligible_responses_message_content(value: &Value) -> bool { + match value { + Value::String(_) => true, + Value::Array(parts) => { + !parts.is_empty() + && parts + .iter() + .all(|part| match part.get("type").and_then(Value::as_str) { + Some("input_text") => { + keys_allowed(part, &["type", "text"]) + && part.get("text").is_some_and(Value::is_string) + } + Some("output_text") => { + keys_allowed(part, &["type", "text", "annotations", "logprobs"]) + && part.get("text").is_some_and(Value::is_string) + } + Some("input_image") => keys_allowed(part, &["type", "image_url", "detail"]), + Some("input_file") => keys_allowed(part, &["type", "file_id", "filename"]), + _ => false, + }) + } + _ => false, + } +} + +fn has_unsupported_controls(request: &serde_json::Map, controls: &[&str]) -> bool { + controls.iter().any(|key| request.contains_key(*key)) + || request.keys().any(|key| key.starts_with("guided_")) +} + +fn supports_single_choice(n: Option<&Value>) -> bool { + match n { + None => true, + Some(Value::Number(n)) => n.as_u64() == Some(1), + Some(_) => false, + } +} + +fn append_generated_text( + generated_text: &mut String, + continuable: &mut bool, + max_buffered_bytes: usize, + text: &str, +) { + if !*continuable { + return; + } + if generated_text.len().saturating_add(text.len()) > max_buffered_bytes { + *continuable = false; + metrics::counter!("onwards_stream_continuation_buffer_exhausted_total").increment(1); + return; + } + generated_text.push_str(text); +} + +fn observe_completion_event( + identity: &mut TextStreamIdentity, + value: &mut Value, + event_type: Option<&[u8]>, + rewrite_identity: bool, +) -> Option { + if event_type.is_some() { + return None; + } + let (text, terminal) = match completion_event(value) { + CompletionEvent::Unsafe => return None, + CompletionEvent::Recognized { text, terminal } => (text.map(str::to_owned), terminal), + }; + if !rewrite_identity && !identity.identity_established { + identity.establish(value); + } + identity.rewrite(value); + Some(ProtocolObservation { + accepted: text.is_some() || terminal, + text, + terminal, + forward: true, + event: Bytes::from(serialize_sse(value).ok()?), + }) +} + +fn observe_chat_event( + identity: &mut TextStreamIdentity, + value: &mut Value, + event_type: Option<&[u8]>, + rewrite_identity: bool, +) -> Option { + if event_type.is_some() { + return None; + } + let event = chat_event(value)?; + if !rewrite_identity && !identity.identity_established { + identity.establish(value); + } + identity.rewrite(value); + + let (text, terminal, role_only) = match event { + ChatEvent::Usage => (None, false, false), + ChatEvent::Choice { + text, + terminal, + role_present, + } => { + if rewrite_identity && role_present { + value + .get_mut("choices")? + .get_mut(0)? + .get_mut("delta")? + .as_object_mut()? + .remove("role"); + } + let role_only = role_present && text.is_none() && !terminal; + (text, terminal, role_only) + } + }; + let forward = !(rewrite_identity && role_only); + Some(ProtocolObservation { + accepted: text.is_some() || terminal, + text, + terminal, + forward, + event: if forward { + Bytes::from(serialize_sse(value).ok()?) + } else { + Bytes::new() + }, + }) +} + +fn observe_responses_event( + state: &mut ResponsesStreamState, + value: &mut Value, + event_field: Option<&[u8]>, + rewrite_identity: bool, + generated_text: &str, +) -> Option { + if rewrite_identity && state.continuation_prefix_len.is_none() { + state.continuation_prefix_len = Some(generated_text.len()); + } + let event_type = value.get("type")?.as_str()?.to_owned(); + let incoming_sequence = value.get("sequence_number")?.as_u64()?; + if let Some(event_field) = event_field + && std::str::from_utf8(event_field).ok()? != event_type + { + return None; + } + let has_event_field = event_field.is_some(); + if !rewrite_identity { + match state.emit_event_field { + None => state.emit_event_field = Some(has_event_field), + Some(expected) if expected != has_event_field => return None, + Some(_) => {} + } + establish_responses_identity(state, value); + } + + let (text, terminal, forward) = match event_type.as_str() { + "response.created" => { + if !keys_allowed(value, &["type", "sequence_number", "response"]) + || !rewrite_response_snapshot( + value.get_mut("response")?, + state, + generated_text, + false, + ) + { + return None; + } + (None, false, !rewrite_identity) + } + "response.in_progress" => { + if !keys_allowed(value, &["type", "sequence_number", "response"]) + || !rewrite_response_snapshot( + value.get_mut("response")?, + state, + generated_text, + false, + ) + { + return None; + } + (None, false, !rewrite_identity) + } + "response.output_item.added" => { + if !keys_allowed(value, &["type", "sequence_number", "output_index", "item"]) + || value.get("output_index").and_then(Value::as_u64) != Some(0) + || !rewrite_response_item(value.get_mut("item")?, state, generated_text, false) + { + return None; + } + (None, false, !rewrite_identity) + } + "response.content_part.added" => { + if !keys_allowed( + value, + &[ + "type", + "sequence_number", + "item_id", + "output_index", + "content_index", + "part", + ], + ) || !valid_responses_indexes(value) + || !rewrite_item_id(value, state) + || !rewrite_output_text_part(value.get_mut("part")?, state, generated_text, false) + { + return None; + } + (None, false, !rewrite_identity) + } + "response.output_text.delta" => { + if !keys_allowed( + value, + &[ + "type", + "sequence_number", + "item_id", + "output_index", + "content_index", + "delta", + "logprobs", + ], + ) || !valid_responses_indexes(value) + || !rewrite_item_id(value, state) + || !empty_logprobs(value.get("logprobs")) + { + return None; + } + let text = value.get("delta")?.as_str()?.to_owned(); + (Some(text), false, true) + } + "response.output_text.done" => { + if !keys_allowed( + value, + &[ + "type", + "sequence_number", + "item_id", + "output_index", + "content_index", + "text", + "logprobs", + ], + ) || !valid_responses_indexes(value) + || !rewrite_item_id(value, state) + || !value.get("text").is_some_and(Value::is_string) + || !empty_logprobs(value.get("logprobs")) + { + return None; + } + if !snapshot_text_matches(value.get("text")?.as_str()?, state, generated_text) { + return None; + } + value + .as_object_mut()? + .insert("text".to_owned(), Value::String(generated_text.to_owned())); + (None, false, true) + } + "response.content_part.done" => { + if !keys_allowed( + value, + &[ + "type", + "sequence_number", + "item_id", + "output_index", + "content_index", + "part", + ], + ) || !valid_responses_indexes(value) + || !rewrite_item_id(value, state) + || !rewrite_output_text_part(value.get_mut("part")?, state, generated_text, true) + { + return None; + } + (None, false, true) + } + "response.output_item.done" => { + if !keys_allowed(value, &["type", "sequence_number", "output_index", "item"]) + || value.get("output_index").and_then(Value::as_u64) != Some(0) + || !rewrite_response_item(value.get_mut("item")?, state, generated_text, true) + { + return None; + } + (None, false, true) + } + "response.completed" | "response.incomplete" => { + if !keys_allowed(value, &["type", "sequence_number", "response"]) + || !rewrite_response_snapshot( + value.get_mut("response")?, + state, + generated_text, + true, + ) + { + return None; + } + (None, true, true) + } + "response.failed" | "response.cancelled" | "response.canceled" => { + if !keys_allowed(value, &["type", "sequence_number", "response"]) + || !rewrite_response_snapshot( + value.get_mut("response")?, + state, + generated_text, + false, + ) + { + return None; + } + (None, true, true) + } + _ => return None, + }; + + if !forward { + return Some(ProtocolObservation { + text: None, + terminal, + accepted: false, + forward: false, + event: Bytes::new(), + }); + } + if rewrite_identity { + let sequence = state.last_sequence.map_or(0, |last| last.saturating_add(1)); + value + .as_object_mut()? + .insert("sequence_number".to_owned(), Value::from(sequence)); + state.last_sequence = Some(sequence); + } else { + if state + .last_sequence + .is_some_and(|last| incoming_sequence <= last) + { + return None; + } + state.last_sequence = Some(incoming_sequence); + } + + Some(ProtocolObservation { + accepted: text.is_some() || terminal, + text, + terminal, + forward: true, + event: Bytes::from( + serialize_response_sse( + &event_type, + value, + state.emit_event_field.unwrap_or(has_event_field), + ) + .ok()?, + ), + }) +} + +fn keys_allowed(value: &Value, allowed: &[&str]) -> bool { + value + .as_object() + .is_some_and(|object| object.keys().all(|key| allowed.contains(&key.as_str()))) +} + +fn empty_logprobs(value: Option<&Value>) -> bool { + matches!(value, None | Some(Value::Null)) + || value.and_then(Value::as_array).is_some_and(Vec::is_empty) +} + +fn establish_responses_identity(state: &mut ResponsesStreamState, value: &Value) { + if !state.response_identity_established + && let Some(response) = value.get("response").and_then(Value::as_object) + { + if let Some(id) = response.get("id").and_then(Value::as_str) { + state.response_id = id.to_owned(); + } + if let Some(model) = response.get("model").and_then(Value::as_str) { + state.model = model.to_owned(); + } + state.response_identity_established = true; + } + if !state.item_identity_established { + let item = value.get("item").and_then(Value::as_object); + let response_item = value + .get("response") + .and_then(|response| response.get("output")) + .and_then(Value::as_array) + .and_then(|output| output.first()) + .and_then(Value::as_object); + let item_id = value + .get("item_id") + .and_then(Value::as_str) + .or_else(|| item.and_then(|item| item.get("id")).and_then(Value::as_str)) + .or_else(|| { + response_item + .and_then(|item| item.get("id")) + .and_then(Value::as_str) + }); + if let Some(item_id) = item_id { + state.item_id = item_id.to_owned(); + } + if value.get("item_id").is_some() || item.is_some() || response_item.is_some() { + state.item_identity_established = true; + } + } +} + +fn valid_responses_indexes(value: &Value) -> bool { + value.get("output_index").and_then(Value::as_u64) == Some(0) + && value.get("content_index").and_then(Value::as_u64) == Some(0) +} + +fn rewrite_item_id(value: &mut Value, state: &ResponsesStreamState) -> bool { + let Some(object) = value.as_object_mut() else { + return false; + }; + if !object.get("item_id").is_some_and(Value::is_string) { + return false; + } + object.insert("item_id".to_owned(), Value::String(state.item_id.clone())); + true +} + +fn rewrite_response_snapshot( + response: &mut Value, + state: &ResponsesStreamState, + generated_text: &str, + require_output: bool, +) -> bool { + let Some(object) = response.as_object_mut() else { + return false; + }; + object.insert("id".to_owned(), Value::String(state.response_id.clone())); + object.insert("model".to_owned(), Value::String(state.model.clone())); + let Some(output) = object.get_mut("output") else { + return !require_output; + }; + let Some(output) = output.as_array_mut() else { + return false; + }; + if output.is_empty() { + return !require_output; + } + output.len() == 1 + && rewrite_response_item(&mut output[0], state, generated_text, require_output) +} + +fn rewrite_response_item( + item: &mut Value, + state: &ResponsesStreamState, + generated_text: &str, + require_content: bool, +) -> bool { + let Some(object) = item.as_object_mut() else { + return false; + }; + if object.get("type").and_then(Value::as_str) != Some("message") + || object.get("role").and_then(Value::as_str) != Some("assistant") + { + return false; + } + object.insert("id".to_owned(), Value::String(state.item_id.clone())); + let Some(content) = object.get_mut("content") else { + return !require_content; + }; + let Some(content) = content.as_array_mut() else { + return false; + }; + if content.is_empty() { + return !require_content; + } + content.len() == 1 + && rewrite_output_text_part(&mut content[0], state, generated_text, require_content) +} + +fn rewrite_output_text_part( + part: &mut Value, + state: &ResponsesStreamState, + generated_text: &str, + require_text: bool, +) -> bool { + let Some(object) = part.as_object_mut() else { + return false; + }; + if object.get("type").and_then(Value::as_str) != Some("output_text") { + return false; + } + if require_text { + let Some(text) = object.get("text").and_then(Value::as_str) else { + return false; + }; + if !snapshot_text_matches(text, state, generated_text) { + return false; + } + object.insert("text".to_owned(), Value::String(generated_text.to_owned())); + } + true +} + +fn snapshot_text_matches( + snapshot_text: &str, + state: &ResponsesStreamState, + generated_text: &str, +) -> bool { + snapshot_text == generated_text + || state + .continuation_prefix_len + .and_then(|prefix_len| generated_text.get(prefix_len..)) + .is_some_and(|attempt_text| snapshot_text == attempt_text) +} + +fn serialize_response_sse( + event_type: &str, + value: &Value, + emit_event_field: bool, +) -> Result, ContinuationError> { + let data = serde_json::to_vec(value).map_err(ContinuationError::Serialize)?; + let mut event = Vec::with_capacity(data.len() + event_type.len() + 24); + if emit_event_field { + event.extend_from_slice(b"event: "); + event.extend_from_slice(event_type.as_bytes()); + event.push(b'\n'); + } + event.extend_from_slice(b"data: "); + event.extend_from_slice(&data); + event.extend_from_slice(b"\n\n"); + Ok(event) +} + +enum ChatEvent { + Usage, + Choice { + text: Option, + terminal: bool, + role_present: bool, + }, +} + +fn chat_event(value: &Value) -> Option { + const TOP_LEVEL_FIELDS: &[&str] = &[ + "id", + "object", + "created", + "model", + "choices", + "usage", + "system_fingerprint", + "service_tier", + ]; + let object = value.as_object()?; + if object + .keys() + .any(|key| !TOP_LEVEL_FIELDS.contains(&key.as_str())) + || object + .get("object") + .is_some_and(|object| object.as_str() != Some("chat.completion.chunk")) + { + return None; + } + let choices = object.get("choices")?.as_array()?; + if choices.is_empty() { + return object + .get("usage") + .filter(|usage| usage.is_object()) + .map(|_| ChatEvent::Usage); + } + if choices.len() != 1 { + return None; + } + let choice = choices[0].as_object()?; + const CHOICE_FIELDS: &[&str] = &["index", "delta", "finish_reason", "logprobs"]; + if choice + .keys() + .any(|key| !CHOICE_FIELDS.contains(&key.as_str())) + || choice.get("index").and_then(Value::as_u64) != Some(0) + || !matches!(choice.get("logprobs"), None | Some(Value::Null)) + { + return None; + } + let delta = choice.get("delta")?.as_object()?; + if delta + .keys() + .any(|key| !matches!(key.as_str(), "role" | "content")) + { + return None; + } + let role_present = match delta.get("role") { + None => false, + Some(Value::String(role)) if role == "assistant" => true, + Some(_) => return None, + }; + let text = match delta.get("content") { + None | Some(Value::Null) => None, + Some(Value::String(text)) => Some(text.clone()), + Some(_) => return None, + }; + let terminal = match choice.get("finish_reason") { + None | Some(Value::Null) => false, + Some(Value::String(_)) => true, + Some(_) => return None, + }; + Some(ChatEvent::Choice { + text, + terminal, + role_present, + }) +} + +enum StreamInterruption { + Done, + Eof, + Body(std::io::Error), + Framing(SseFramingError), + IdleTimeout, +} + +#[derive(Default)] +struct TerminalMetricState { + recorded: bool, +} + +impl TerminalMetricState { + fn observe_finish_reason(&mut self) -> Option<&'static str> { + self.record("finish_reason") + } + + fn observe_done(&mut self) -> Option<&'static str> { + self.record("done") + } + + fn record(&mut self, reason: &'static str) -> Option<&'static str> { + if self.recorded { + None + } else { + self.recorded = true; + Some(reason) + } + } +} + +impl StreamInterruption { + fn into_error(self) -> Option { + match self { + Self::Done | Self::Eof => None, + Self::Body(error) => Some(error), + Self::Framing(error) => Some(std::io::Error::other(error)), + Self::IdleTimeout => Some(std::io::Error::new( + std::io::ErrorKind::TimedOut, + "upstream generation stream idle timeout", + )), + } + } + + fn reason(&self) -> &'static str { + match self { + Self::Done => "done", + Self::Eof => "eof", + Self::Body(_) => "body_error", + Self::Framing(_) => "framing_error", + Self::IdleTimeout => "idle_timeout", + } + } + + fn exhausted_reason(&self) -> &'static str { + match self { + Self::Done => "done", + Self::Eof => "exhausted_eof", + Self::Body(_) => "exhausted_body_error", + Self::Framing(_) => "framing_error", + Self::IdleTimeout => "exhausted_idle_timeout", + } + } +} + +/// Combines an initial generation body with independently selected continuation bodies. +pub(crate) fn wrap_generation_stream( + initial_body: Body, + initial_guard: ConcurrencyGuard, + mut continuation: StreamContinuation, + config: StreamContinuationConfig, + pool: ProviderPool, + http_client: T, + request_metadata: UpstreamRequestMetadata, +) -> Body +where + T: HttpClient + Send + 'static, +{ + let fallback = pool.fallback().cloned().unwrap_or_default(); + let mut selection = SelectionState::new(config.max_attempts, fallback.with_replacement); + let max_event_bytes = event_buffer_size(); + let stream = async_stream::stream! { + use tracing::Instrument; + + let mut current_body = initial_body; + let mut current_guard = Some(initial_guard); + let mut rewrite_identity = false; + let mut awaiting_resumption_event = false; + let mut terminal_metric = TerminalMetricState::default(); + let mut continuation_attempt = 0_u32; + let mut total_backoff_ms = 0_u64; + + loop { + let mut events = CheckedSseStream::with_max_buffer_size( + current_body.into_data_stream(), + max_event_bytes, + ); + let interruption = loop { + let next = if let Some(timeout_ms) = config.idle_timeout_ms { + match tokio::time::timeout(Duration::from_millis(timeout_ms), events.next()).await { + Ok(next) => next, + Err(_) => break StreamInterruption::IdleTimeout, + } + } else { + events.next().await + }; + + match next { + Some(Ok(event)) => { + let observation = match continuation.observe_event(&event, rewrite_identity) { + Ok(observation) => observation, + Err(error) => { + yield Err::(std::io::Error::other(error)); + return; + } + }; + if rewrite_identity && !observation.safe { + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "unsafe_continuation_event" + ) + .increment(1); + yield Err::(std::io::Error::other( + "unsafe generation event from continuation provider", + )); + return; + } + if !rewrite_identity && !observation.safe { + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "unsafe_initial_event" + ) + .increment(1); + } + if rewrite_identity && awaiting_resumption_event && observation.accepted { + awaiting_resumption_event = false; + metrics::counter!("onwards_stream_continuation_resumptions_total") + .increment(1); + } + if observation.terminal + && !observation.done + && let Some(reason) = terminal_metric.observe_finish_reason() + { + metrics::counter!( + "onwards_stream_continuation_terminal_total", + "reason" => reason + ) + .increment(1); + } + let done = observation.done; + if observation.forward { + yield Ok::(observation.event); + } + if done { + break StreamInterruption::Done; + } + } + Some(Err(SseStreamError::Source(error))) => { + if let Some(framing) = framing_error_in_chain(&error) { + break StreamInterruption::Framing(framing); + } + break StreamInterruption::Body(std::io::Error::other(error)); + } + Some(Err(SseStreamError::Framing(error))) => { + break StreamInterruption::Framing(error); + } + None => break StreamInterruption::Eof, + } + }; + + drop(events); + drop(current_guard.take()); + + tracing::debug!( + reason = interruption.reason(), + attempt = continuation_attempt, + "Generation stream ended" + ); + + if matches!(interruption, StreamInterruption::Done) { + if let Some(reason) = terminal_metric.observe_done() { + metrics::counter!( + "onwards_stream_continuation_terminal_total", + "reason" => reason + ) + .increment(1); + } + break; + } + + if continuation.is_terminal() { + break; + } + + if matches!(interruption, StreamInterruption::Framing(_)) { + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "framing_error" + ) + .increment(1); + if let Some(error) = interruption.into_error() { + yield Err::(error); + } + break; + } + + let exhausted_reason = interruption.exhausted_reason(); + let final_error = interruption.into_error(); + if !continuation.is_continuable() { + if let Some(error) = final_error { + yield Err::(error); + } + break; + } + + let mut next_body = None; + let mut next_guard = None; + loop { + if continuation_attempt as usize >= config.max_attempts { + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "max_attempts" + ) + .increment(1); + break; + } + let retry_index = continuation_attempt.saturating_add(1); + if let Some(backoff) = fallback.backoff.as_ref() { + let delay = backoff.delay(retry_index); + let delay_ms = delay.as_millis() as u64; + let next_total = total_backoff_ms.saturating_add(delay_ms); + if fallback + .max_total_backoff_ms + .is_some_and(|max_total| next_total > max_total) + { + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "backoff_budget" + ) + .increment(1); + break; + } + total_backoff_ms = next_total; + tokio::time::sleep(delay).await; + } + + let Some((_index, target, guard)) = pool.select_next(&mut selection) else { + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "selection_exhausted" + ) + .increment(1); + break; + }; + continuation_attempt = continuation_attempt.saturating_add(1); + metrics::counter!("onwards_stream_continuation_attempts_total").increment(1); + let target = target.clone(); + let attempt_span = tracing::info_span!( + "onwards.stream_continuation_attempt", + attempt = continuation_attempt, + outcome = tracing::field::Empty, + http.response.status_code = tracing::field::Empty, + ); + let attempt_request_metadata = request_metadata.for_child_span(&attempt_span); + + if target + .limiter + .as_ref() + .is_some_and(|limiter| limiter.check().is_err()) + { + attempt_span.record("outcome", "rate_limit"); + drop(guard); + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "rate_limit" + ) + .increment(1); + if pool.should_fallback_on_rate_limit() { + continue; + } + break; + } + + let request_body = match continuation.request_body(target.onwards_model.as_deref()) { + Ok(body) => body, + Err(error) => { + drop(guard); + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "request_body" + ) + .increment(1); + tracing::error!(error = %error, "Failed to build stream continuation request"); + break; + } + }; + let (request, upstream_uri) = match build_upstream_request( + &target, + &attempt_request_metadata, + request_body, + ) { + Ok(request) => request, + Err(_error) => { + drop(guard); + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "request_build" + ) + .increment(1); + tracing::error!( + "Failed to build stream continuation upstream request" + ); + break; + } + }; + + tracing::debug!( + attempt = continuation_attempt, + upstream = %upstream_uri, + "Requesting generation stream continuation" + ); + let request_future = http_client + .request(request) + .instrument(attempt_span.clone()); + let response = if let Some(timeout_secs) = target.request_timeout_secs { + match tokio::time::timeout( + Duration::from_secs(timeout_secs), + request_future, + ) + .await + { + Ok(response) => response, + Err(_) => { + attempt_span.record("outcome", "header_timeout"); + drop(guard); + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "header_timeout" + ) + .increment(1); + continue; + } + } + } else { + request_future.await + }; + + let response = match response { + Ok(response) => response, + Err(error) => { + attempt_span.record("outcome", "network_error"); + drop(guard); + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "network_error" + ) + .increment(1); + tracing::warn!( + upstream = %target.url, + error = %error, + "Stream continuation request failed" + ); + continue; + } + }; + let status = response.status().as_u16(); + attempt_span.record("http.response.status_code", status); + if !(200..300).contains(&status) { + attempt_span.record("outcome", "status"); + drop(response); + drop(guard); + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "status" + ) + .increment(1); + if pool.should_fallback_on_status(status) { + continue; + } + break; + } + if !is_event_stream(response.headers()) { + attempt_span.record("outcome", "content_type"); + drop(response); + drop(guard); + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "content_type" + ) + .increment(1); + continue; + } + if !has_identity_content_encoding(response.headers()) { + attempt_span.record("outcome", "content_encoding"); + drop(response); + drop(guard); + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => "content_encoding" + ) + .increment(1); + continue; + } + + next_body = Some(response.into_body()); + next_guard = Some(guard); + attempt_span.record("outcome", "headers_accepted"); + break; + } + + match (next_body, next_guard) { + (Some(body), Some(guard)) => { + continuation.begin_continuation_attempt(); + current_body = body; + current_guard = Some(guard); + rewrite_identity = true; + awaiting_resumption_event = true; + } + _ => { + metrics::counter!( + "onwards_stream_continuation_failures_total", + "reason" => exhausted_reason + ) + .increment(1); + if let Some(error) = final_error { + yield Err::(error); + } + break; + } + } + } + }; + + Body::from_stream(stream) +} + +enum CompletionEvent<'a> { + Recognized { + text: Option<&'a str>, + terminal: bool, + }, + Unsafe, +} + +fn completion_event(completion: &Value) -> CompletionEvent<'_> { + let Some(object) = completion.as_object() else { + return CompletionEvent::Unsafe; + }; + const SUPPORTED_TOP_LEVEL_FIELDS: &[&str] = &[ + "id", + "object", + "created", + "model", + "choices", + "usage", + "system_fingerprint", + "error", + ]; + if object + .keys() + .any(|key| !SUPPORTED_TOP_LEVEL_FIELDS.contains(&key.as_str())) + { + return CompletionEvent::Unsafe; + } + if object.contains_key("error") { + return CompletionEvent::Unsafe; + } + if object + .get("object") + .is_some_and(|value| value.as_str() != Some("text_completion")) + { + return CompletionEvent::Unsafe; + } + + let Some(choices_value) = object.get("choices") else { + return if choice_less_completion_metadata(completion) { + CompletionEvent::Recognized { + text: None, + terminal: false, + } + } else { + CompletionEvent::Unsafe + }; + }; + let Some(choices) = choices_value.as_array() else { + return CompletionEvent::Unsafe; + }; + + if choices.is_empty() && choice_less_completion_metadata(completion) { + return CompletionEvent::Recognized { + text: None, + terminal: false, + }; + } + + if choices.len() != 1 { + return CompletionEvent::Unsafe; + } + + let Some(choice) = choices.first().and_then(Value::as_object) else { + return CompletionEvent::Unsafe; + }; + const SUPPORTED_CHOICE_FIELDS: &[&str] = &["index", "text", "finish_reason", "logprobs"]; + if choice + .keys() + .any(|key| !SUPPORTED_CHOICE_FIELDS.contains(&key.as_str())) + || choice.get("index").and_then(Value::as_u64) != Some(0) + || !matches!(choice.get("logprobs"), None | Some(Value::Null)) + { + return CompletionEvent::Unsafe; + } + let Some(text) = choice.get("text").and_then(Value::as_str) else { + return CompletionEvent::Unsafe; + }; + let terminal = match choice.get("finish_reason") { + None | Some(Value::Null) => false, + Some(Value::String(_)) => true, + Some(_) => return CompletionEvent::Unsafe, + }; + + CompletionEvent::Recognized { + text: Some(text), + terminal, + } +} + +fn choice_less_completion_metadata(completion: &Value) -> bool { + let Some(completion) = completion.as_object() else { + return false; + }; + completion.get("object").and_then(Value::as_str) == Some("text_completion") + && completion.contains_key("id") + && completion.contains_key("model") + && completion.contains_key("created") +} + +fn serialize_sse(completion: &Value) -> Result, ContinuationError> { + let data = serde_json::to_vec(completion).map_err(ContinuationError::Serialize)?; + let mut event = Vec::with_capacity(b"data: \n\n".len() + data.len()); + event.extend_from_slice(b"data: "); + event.extend_from_slice(&data); + event.extend_from_slice(b"\n\n"); + Ok(event) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::target::StreamContinuationConfig; + use axum::http::Method; + use serde_json::Value; + + fn eligible_state(max_buffered_bytes: usize) -> StreamContinuation { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/completions".to_string()], + max_attempts: 1, + max_buffered_bytes, + idle_timeout_ms: None, + }; + StreamContinuation::from_request( + "/v1/completions", + &Method::POST, + br#"{"model":"requested-m","prompt":"Say hello: ","stream":true}"#, + &config, + ) + .unwrap() + } + + #[test] + fn accepts_only_supported_completion_requests() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/completions".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let eligible = br#"{"prompt":"hello","stream":true}"#; + + assert!( + StreamContinuation::from_request("/v1/completions", &Method::POST, eligible, &config,) + .is_some() + ); + for (path, method, body) in [ + ("/v1/chat/completions", Method::POST, eligible.as_slice()), + ("/v1/completions", Method::GET, eligible.as_slice()), + ( + "/v1/completions", + Method::POST, + br#"{"prompt":["hello"],"stream":true}"#.as_slice(), + ), + ( + "/v1/completions", + Method::POST, + br#"{"prompt":"hello","stream":false}"#.as_slice(), + ), + ( + "/v1/completions", + Method::POST, + br#"{"prompt":"hello","stream":true,"n":2}"#.as_slice(), + ), + ( + "/v1/completions", + Method::POST, + br#"{"prompt":"hello","stream":true,"echo":true}"#.as_slice(), + ), + ( + "/v1/completions", + Method::POST, + br#"{"prompt":"hello","stream":true,"logprobs":1}"#.as_slice(), + ), + ] { + assert!(StreamContinuation::from_request(path, &method, body, &config).is_none()); + } + + assert!( + StreamContinuation::from_request( + "/v1/completions", + &Method::POST, + br#"{"prompt":"hello","stream":true,"logprobs":null}"#, + &config, + ) + .is_some() + ); + } + + #[test] + fn advanced_generation_controls_are_ineligible() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/completions".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + + for key in [ + "tools", + "tool_choice", + "functions", + "function_call", + "response_format", + "json_schema", + "grammar", + "guided_json", + "guided_regex", + "guided_choice", + "guided_grammar", + "guided_options_request", + ] { + let mut request = serde_json::json!({ + "prompt": "hello", + "stream": true + }); + request[key] = serde_json::json!({"enabled": true}); + let body = serde_json::to_vec(&request).unwrap(); + + assert!( + StreamContinuation::from_request("/v1/completions", &Method::POST, &body, &config,) + .is_none(), + "{key} must disable continuation" + ); + } + } + + #[test] + fn interrupted_completion_builds_prefix_request() { + let mut state = eligible_state(1024); + state + .observe_event( + br#"data: {"id":"cmpl-a","created":1,"model":"m","choices":[{"index":0,"text":"hello","finish_reason":null}]} + +"#, + false, + ) + .unwrap(); + let body: Value = + serde_json::from_slice(&state.request_body(Some("upstream-m")).unwrap()).unwrap(); + assert_eq!(body["prompt"], "Say hello: hello"); + assert_eq!(body["model"], "upstream-m"); + } + + #[test] + fn terminal_events_prevent_continuation() { + let mut state = eligible_state(1024); + assert!( + state + .observe_event(b"data: [DONE]\n\n", false) + .unwrap() + .terminal + ); + assert!(!state.is_continuable()); + + let mut state = eligible_state(1024); + assert!( + state + .observe_event( + br#"data: {"choices":[{"index":0,"text":"","finish_reason":"stop"}]} + +"#, + false, + ) + .unwrap() + .terminal + ); + assert!(!state.is_continuable()); + } + + #[test] + fn multiline_continuation_event_reuses_original_identity() { + let mut state = eligible_state(1024); + state + .observe_event( + br#"data: {"id":"cmpl-first","created":1,"model":"first","choices":[{"index":0,"text":"one","finish_reason":null}]} + +"#, + false, + ) + .unwrap(); + + let observation = state + .observe_event( + b"data: {\"id\":\"cmpl-next\",\"created\":2,\ndata: \"model\":\"next\",\"choices\":[{\"index\":0,\"text\":\"two\",\"finish_reason\":null}]}\n\n", + + true, + ) + .unwrap(); + + let rewritten = String::from_utf8(observation.event.to_vec()).unwrap(); + assert_eq!( + rewritten, + "data: {\"choices\":[{\"finish_reason\":null,\"index\":0,\"text\":\"two\"}],\"created\":1,\"id\":\"cmpl-first\",\"model\":\"first\"}\n\n" + ); + } + + #[test] + fn comments_are_safe_but_malformed_events_disable_continuation() { + let mut state = eligible_state(1024); + let comment = b": keepalive\n\n"; + let observation = state.observe_event(comment, false).unwrap(); + assert_eq!(observation.event.as_ref(), comment); + assert!(state.is_continuable()); + + let malformed = b"data: {not-json}\n\n"; + let observation = state.observe_event(malformed, false).unwrap(); + assert_eq!(observation.event.as_ref(), malformed); + assert!(!state.is_continuable()); + } + + #[test] + fn multiple_choices_disable_continuation_without_appending_a_prefix() { + let mut state = eligible_state(1024); + let event = br#"data: {"choices":[{"index":0,"text":"one","finish_reason":null},{"index":1,"text":"two","finish_reason":"stop"}]} + +"#; + let observation = state.observe_event(event, true).unwrap(); + + assert_eq!(observation.event.as_ref(), event); + assert!(!observation.terminal); + assert!(!state.is_continuable()); + assert!(state.generated_text.is_empty()); + } + + #[test] + fn buffer_exhaustion_disables_continuation_without_retaining_new_text() { + let mut state = eligible_state(5); + state + .observe_event( + br#"data: {"choices":[{"index":0,"text":"hello","finish_reason":null}]} + +"#, + false, + ) + .unwrap(); + state + .observe_event( + br#"data: {"choices":[{"index":0,"text":"!","finish_reason":null}]} + +"#, + false, + ) + .unwrap(); + + assert!(!state.is_continuable()); + assert!(state.request_body(None).is_err()); + } + + #[test] + fn initial_recognized_event_receives_stable_fallback_identity() { + let mut state = eligible_state(1024); + let event = br#"data: {"choices":[{"index":0,"text":"hello","finish_reason":null}]} + +"#; + let observation = state.observe_event(event, false).unwrap(); + let rewritten: Value = + serde_json::from_slice(&observation.event[6..observation.event.len() - 2]).unwrap(); + assert!(rewritten["id"].as_str().unwrap().starts_with("cmpl-")); + assert_eq!(rewritten["model"], "requested-m"); + assert!(rewritten["created"].is_u64()); + assert_eq!(rewritten["choices"][0]["text"], "hello"); + assert!(!observation.terminal); + } + + #[test] + fn unrecognized_choice_events_pass_through_and_disable_continuation() { + for event in [ + br#"data: {"id":"chat-next","choices":[{"index":0,"delta":{"content":"two"},"finish_reason":"stop"}]} + +"# + .as_slice(), + br#"data: {"id":"empty-choice","choices":[{}]} + +"# + .as_slice(), + ] { + let mut state = eligible_state(1024); + state + .observe_event( + br#"data: {"id":"cmpl-first","created":1,"model":"first","choices":[{"index":0,"text":"one","finish_reason":null}]} + +"#, + false, + ) + .unwrap(); + let observation = state.observe_event(event, true).unwrap(); + assert_eq!(observation.event.as_ref(), event); + assert!(!observation.terminal, "unsafe finish reasons are not trusted"); + assert!(!state.is_continuable()); + } + } + + #[test] + fn unrecognized_json_without_choices_passes_through_and_disables_continuation() { + let mut state = eligible_state(1024); + let event = br#"data: {"id":"other","type":"notification","payload":{"value":1}} + +"#; + + let observation = state.observe_event(event, true).unwrap(); + assert_eq!(observation.event.as_ref(), event); + assert!(!observation.terminal); + assert!(!state.is_continuable()); + } + + #[test] + fn hybrid_sse_fields_disable_continuation_without_changing_initial_bytes() { + for field in [ + "event: completion", + "id: provider-event", + "retry: 1000", + "provider-field: unsafe", + ] { + let mut state = eligible_state(1024); + let event = format!( + "{field}\ndata: {{\"choices\":[{{\"index\":0,\"text\":\"x\",\"finish_reason\":null}}]}}\n\n" + ); + + let observation = state.observe_event(event.as_bytes(), false).unwrap(); + + assert_eq!(observation.event.as_ref(), event.as_bytes()); + assert!(!observation.safe, "unsupported SSE field: {field}"); + assert!(!state.is_continuable(), "unsupported SSE field: {field}"); + } + } + + #[test] + fn unsafe_top_level_completion_keys_disable_continuation() { + for key in ["payload", "tool_calls", "unknown_provider_state"] { + let mut state = eligible_state(1024); + let mut value = serde_json::json!({ + "choices": [{"index": 0, "text": "x", "finish_reason": null}] + }); + value[key] = serde_json::json!({"unsafe": true}); + let event = format!("data: {value}\n\n"); + + let observation = state.observe_event(event.as_bytes(), false).unwrap(); + + assert!(!observation.safe, "unsafe key: {key}"); + assert!(!state.is_continuable(), "unsafe key: {key}"); + } + } + + #[test] + fn safe_completion_metadata_remains_continuable() { + let mut state = eligible_state(1024); + let event = b"data: {\"id\":\"cmpl-a\",\"object\":\"text_completion\",\"created\":1,\"model\":\"m\",\"system_fingerprint\":\"fp\",\"choices\":[{\"index\":0,\"text\":\"x\",\"finish_reason\":null}],\"usage\":{\"completion_tokens\":1}}\n\n"; + + let observation = state.observe_event(event, false).unwrap(); + + assert!(observation.safe); + assert!(state.is_continuable()); + } + + #[test] + fn terminal_metric_reason_is_emitted_only_once() { + let mut state = TerminalMetricState::default(); + + assert_eq!(state.observe_finish_reason(), Some("finish_reason")); + assert_eq!(state.observe_done(), None); + assert_eq!(state.observe_finish_reason(), None); + } + + #[test] + fn unsafe_completion_shapes_fail_closed_without_changing_initial_bytes() { + for event in [ + b"data: {not-json}\n\n".as_slice(), + br#"data: {"error":{"message":"failed"}} + +"# + .as_slice(), + br#"data: {"choices":[{"index":0,"text":"x","finish_reason":null}],"error":{"message":"failed"}} + +"# + .as_slice(), + br#"data: {"choices":[{"index":0,"delta":{"content":"x"},"finish_reason":null}]} + +"# + .as_slice(), + br#"data: {"choices":[{"index":0,"message":{"content":"x"},"finish_reason":null}]} + +"# + .as_slice(), + br#"data: {"choices":[{"index":0,"text":"x","finish_reason":null,"tool_calls":[]}]} + +"# + .as_slice(), + br#"data: {"choices":[{"text":"x","finish_reason":null}]} + +"# + .as_slice(), + br#"data: {"choices":[{"index":1,"text":"x","finish_reason":null}]} + +"# + .as_slice(), + br#"data: {"type":"notification","payload":{"value":1}} + +"# + .as_slice(), + ] { + let mut state = eligible_state(1024); + let observation = state.observe_event(event, false).unwrap(); + + assert_eq!(observation.event.as_ref(), event); + assert!(!state.is_continuable(), "unsafe event: {event:?}"); + } + } + + #[test] + fn done_without_post_colon_space_and_bom_are_recognized() { + let mut state = eligible_state(1024); + let event = b"\xef\xbb\xbfdata:[DONE]\r\n\r\n"; + let observation = state.observe_event(event, false).unwrap(); + + assert_eq!(observation.event.as_ref(), event); + assert!(observation.done); + assert!(observation.terminal); + assert!(!state.is_continuable()); + } + + #[test] + fn observation_flags_distinguish_comments_content_and_unsafe_events() { + let mut state = eligible_state(1024); + let comment = state.observe_event(b": ping\n\n", false).unwrap(); + assert!(comment.safe); + assert!(!comment.accepted); + + let content = state + .observe_event( + b"data:{\"choices\":[{\"index\":0,\"text\":\"x\",\"finish_reason\":null}]}\n\n", + false, + ) + .unwrap(); + assert!(content.safe); + assert!(content.accepted); + + let unsafe_event = state.observe_event(b"data:{not-json}\n\n", false).unwrap(); + assert!(!unsafe_event.safe); + assert!(!unsafe_event.accepted); + } + + #[test] + fn missing_initial_identity_is_stable_across_continuations() { + let mut state = eligible_state(1024); + let initial = state + .observe_event( + br#"data: {"id":"cmpl-first","choices":[{"index":0,"text":"one","finish_reason":null}]} + +"#, + false, + ) + .unwrap(); + let initial: Value = + serde_json::from_slice(&initial.event[6..initial.event.len() - 2]).unwrap(); + + for (model, created) in [("provider-a", 2), ("provider-b", 3)] { + let event = format!( + "data: {{\"id\":\"cmpl-next\",\"created\":{created},\"model\":\"{model}\",\"choices\":[{{\"index\":0,\"text\":\"two\",\"finish_reason\":null}}]}}\n\n" + ); + let observation = state.observe_event(event.as_bytes(), true).unwrap(); + let rewritten: Value = + serde_json::from_slice(&observation.event[6..observation.event.len() - 2]).unwrap(); + assert_eq!(rewritten["id"], "cmpl-first"); + assert_eq!(rewritten["model"], initial["model"]); + assert_eq!(rewritten["created"], initial["created"]); + } + } + + #[test] + fn response_representation_checks_are_exact_and_case_insensitive() { + let mut headers = HeaderMap::new(); + headers.insert( + "content-type", + "Text/Event-Stream; charset=utf-8".parse().unwrap(), + ); + assert!(is_event_stream(&headers)); + assert!(has_identity_content_encoding(&headers)); + + headers.insert(CONTENT_ENCODING, "IDENTITY".parse().unwrap()); + assert!(has_identity_content_encoding(&headers)); + + headers.insert( + "content-type", + "application/x-text/event-streamish".parse().unwrap(), + ); + headers.insert(CONTENT_ENCODING, "gzip".parse().unwrap()); + assert!(!is_event_stream(&headers)); + assert!(!has_identity_content_encoding(&headers)); + } + + #[test] + fn event_stream_media_type_requires_one_fully_valid_mime_value() { + for valid in [ + "text/event-stream", + "Text/Event-Stream; charset=utf-8", + "TEXT/EVENT-STREAM; Charset=\"utf-8\"", + ] { + let mut headers = HeaderMap::new(); + headers.insert(CONTENT_TYPE, valid.parse().unwrap()); + assert!(is_event_stream(&headers), "valid MIME rejected: {valid}"); + } + + for invalid in [ + "text/event-stream;", + "text/event-stream; charset", + "text/event-stream; charset=", + "text/event-stream, application/json", + "text/event-stream; charset=utf-8, application/json", + ] { + let mut headers = HeaderMap::new(); + headers.insert(CONTENT_TYPE, invalid.parse().unwrap()); + assert!( + !is_event_stream(&headers), + "invalid MIME accepted: {invalid}" + ); + } + + let mut headers = HeaderMap::new(); + headers.append(CONTENT_TYPE, "text/event-stream".parse().unwrap()); + headers.append(CONTENT_TYPE, "text/event-stream".parse().unwrap()); + assert!(!is_event_stream(&headers)); + } + + #[test] + fn unambiguous_choice_less_completion_metadata_captures_initial_identity() { + let mut state = eligible_state(1024); + state + .observe_event( + br#"data: {"id":"cmpl-first","object":"text_completion","created":1,"model":"first"} + +"#, + false, + ) + .unwrap(); + + let observation = state + .observe_event( + br#"data: {"id":"cmpl-next","created":2,"model":"next","choices":[{"index":0,"text":"two","finish_reason":null}]} + +"#, + true, + ) + .unwrap(); + let rewritten: Value = + serde_json::from_slice(&observation.event[6..observation.event.len() - 2]).unwrap(); + assert_eq!(rewritten["id"], "cmpl-first"); + assert_eq!(rewritten["model"], "first"); + assert_eq!(rewritten["created"], 1); + } + + #[test] + fn multibyte_text_at_the_buffer_boundary_is_retained_but_overflow_is_not() { + let event = + "data: {\"choices\":[{\"index\":0,\"text\":\"éé\",\"finish_reason\":null}]}\n\n"; + + let mut exact_boundary = eligible_state(4); + exact_boundary + .observe_event(event.as_bytes(), false) + .unwrap(); + let body: Value = + serde_json::from_slice(&exact_boundary.request_body(None).unwrap()).unwrap(); + assert_eq!(body["prompt"], "Say hello: éé"); + + let mut overflow = eligible_state(3); + overflow.observe_event(event.as_bytes(), false).unwrap(); + assert!(!overflow.is_continuable()); + assert!(overflow.generated_text.is_empty()); + } + + #[test] + fn chat_continuation_appends_the_emitted_assistant_prefix() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/chat/completions".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let mut state = StreamContinuation::from_request( + "/v1/chat/completions", + &Method::POST, + br#"{"model":"chat-model","messages":[{"role":"user","content":"Hello"}],"stream":true}"#, + &config, + ) + .expect("text-only chat streams should be continuable"); + + state + .observe_event( + br#"data: {"id":"chatcmpl-first","object":"chat.completion.chunk","created":1,"model":"chat-model","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]} + +"#, + false, + ) + .unwrap(); + state + .observe_event( + br#"data: {"id":"chatcmpl-first","object":"chat.completion.chunk","created":1,"model":"chat-model","choices":[{"index":0,"delta":{"content":"Partial answer"},"finish_reason":null}]} + +"#, + false, + ) + .unwrap(); + + let body: Value = serde_json::from_slice(&state.request_body(None).unwrap()).unwrap(); + assert_eq!(body["messages"].as_array().unwrap().len(), 2); + assert_eq!(body["messages"][1]["role"], "assistant"); + assert_eq!(body["messages"][1]["content"], "Partial answer"); + } + + #[test] + fn responses_continuation_appends_an_assistant_output_item() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/responses".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let mut state = StreamContinuation::from_request( + "/v1/responses", + &Method::POST, + br#"{"model":"responses-model","input":"Hello","stream":true}"#, + &config, + ) + .expect("text-only Responses streams should be continuable"); + + state + .observe_event( + b"event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"sequence_number\":3,\"item_id\":\"msg_first\",\"output_index\":0,\"content_index\":0,\"delta\":\"Partial answer\"}\n\n", + false, + ) + .unwrap(); + + let body: Value = serde_json::from_slice(&state.request_body(None).unwrap()).unwrap(); + let input = body["input"].as_array().unwrap(); + assert_eq!(input.len(), 2); + assert_eq!(input[0]["role"], "user"); + assert_eq!(input[0]["content"], "Hello"); + assert_eq!(input[1]["role"], "assistant"); + assert_eq!(input[1]["content"][0]["type"], "output_text"); + assert_eq!(input[1]["content"][0]["text"], "Partial answer"); + } + + #[test] + fn chat_continuation_suppresses_repeated_role_and_reuses_identity() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/chat/completions".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let mut state = StreamContinuation::from_request( + "/v1/chat/completions", + &Method::POST, + br#"{"model":"requested","messages":[{"role":"user","content":"Hello"}],"stream":true}"#, + &config, + ) + .unwrap(); + state + .observe_event( + b"data: {\"id\":\"chat-first\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"first\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\n", + false, + ) + .unwrap(); + state + .observe_event( + b"data: {\"id\":\"chat-first\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"first\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\n", + false, + ) + .unwrap(); + + let repeated_role = state + .observe_event( + b"data: {\"id\":\"chat-second\",\"object\":\"chat.completion.chunk\",\"created\":2,\"model\":\"second\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\n", + true, + ) + .unwrap(); + assert!(!repeated_role.forward); + + let continued = state + .observe_event( + b"data: {\"id\":\"chat-second\",\"object\":\"chat.completion.chunk\",\"created\":2,\"model\":\"second\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":null}]}\n\n", + true, + ) + .unwrap(); + let continued: Value = + serde_json::from_slice(&continued.event[6..continued.event.len() - 2]).unwrap(); + assert_eq!(continued["id"], "chat-first"); + assert_eq!(continued["model"], "first"); + assert_eq!(continued["created"], 1); + assert_eq!(continued["choices"][0]["delta"]["content"], " world"); + } + + #[test] + fn chat_tool_or_reasoning_events_disable_continuation() { + for delta in [ + serde_json::json!({"tool_calls": []}), + serde_json::json!({"reasoning_content": "thinking"}), + ] { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/chat/completions".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let mut state = StreamContinuation::from_request( + "/v1/chat/completions", + &Method::POST, + br#"{"messages":[{"role":"user","content":"Hello"}],"stream":true}"#, + &config, + ) + .unwrap(); + let event = format!( + "data: {}\n\n", + serde_json::json!({ + "id": "chat-first", + "object": "chat.completion.chunk", + "created": 1, + "model": "first", + "choices": [{"index": 0, "delta": delta, "finish_reason": null}] + }) + ); + + let observation = state.observe_event(event.as_bytes(), false).unwrap(); + assert!(!observation.safe); + assert!(!state.is_continuable()); + } + } + + #[test] + fn completion_chat_and_responses_logprobs_events_disable_continuation() { + let mut completion = eligible_state(1024); + let completion_event = b"data: {\"id\":\"cmpl-first\",\"object\":\"text_completion\",\"created\":1,\"model\":\"first\",\"choices\":[{\"index\":0,\"text\":\"Hello\",\"finish_reason\":null,\"logprobs\":{\"tokens\":[\"Hello\"]}}]}\n\n"; + assert!( + !completion + .observe_event(completion_event, false) + .unwrap() + .safe + ); + assert!(!completion.is_continuable()); + + let chat_config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/chat/completions".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let mut chat = StreamContinuation::from_request( + "/v1/chat/completions", + &Method::POST, + br#"{"messages":[{"role":"user","content":"Hello"}],"stream":true}"#, + &chat_config, + ) + .unwrap(); + let chat_event = b"data: {\"id\":\"chat-first\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"first\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null,\"logprobs\":{\"content\":[]}}]}\n\n"; + assert!(!chat.observe_event(chat_event, false).unwrap().safe); + assert!(!chat.is_continuable()); + + let responses_config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/responses".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let mut responses = StreamContinuation::from_request( + "/v1/responses", + &Method::POST, + br#"{"input":"Hello","stream":true}"#, + &responses_config, + ) + .unwrap(); + let responses_event = b"event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"sequence_number\":0,\"item_id\":\"msg_first\",\"output_index\":0,\"content_index\":0,\"delta\":\"Hello\",\"logprobs\":[{}]}\n\n"; + assert!( + !responses + .observe_event(responses_event, false) + .unwrap() + .safe + ); + assert!(!responses.is_continuable()); + } + + #[test] + fn chat_and_responses_reject_non_text_continuation_context() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec![ + "/v1/chat/completions".to_string(), + "/v1/responses".to_string(), + ], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + + for request in [ + serde_json::json!({ + "messages": [{"role": "assistant", "content": null, "tool_calls": []}], + "stream": true + }), + serde_json::json!({ + "messages": [{ + "role": "user", + "content": [{"type": "tool_call", "name": "lookup"}] + }], + "stream": true + }), + serde_json::json!({ + "messages": [{"role": "user", "content": "Hello"}], + "stream": true, + "parallel_tool_calls": true + }), + serde_json::json!({ + "messages": [{"role": "user", "content": "Hello"}], + "stream": true, + "reasoning_effort": "high" + }), + ] { + assert!( + StreamContinuation::from_request( + "/v1/chat/completions", + &Method::POST, + &serde_json::to_vec(&request).unwrap(), + &config, + ) + .is_none(), + "unsafe Chat request was eligible: {request}" + ); + } + + for request in [ + serde_json::json!({ + "input": [{"type": "function_call", "name": "lookup", "arguments": "{}"}], + "stream": true + }), + serde_json::json!({ + "input": [{"type": "reasoning", "summary": []}], + "stream": true + }), + serde_json::json!({ + "input": [{ + "type": "message", + "role": "user", + "content": [{"type": "reasoning_text", "text": "thinking"}] + }], + "stream": true + }), + serde_json::json!({ + "input": [{ + "type": "message", + "role": "assistant", + "content": "Hello", + "tool_calls": [] + }], + "stream": true + }), + serde_json::json!({"input": [], "stream": true}), + serde_json::json!({"input": "Hello", "stream": true, "text": "invalid"}), + serde_json::json!({ + "input": "Hello", + "stream": true, + "text": {"format": {"type": "json_schema"}} + }), + ] { + assert!( + StreamContinuation::from_request( + "/v1/responses", + &Method::POST, + &serde_json::to_vec(&request).unwrap(), + &config, + ) + .is_none(), + "unsafe Responses request was eligible: {request}" + ); + } + } + + #[test] + fn chat_and_responses_accept_supported_message_context() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec![ + "/v1/chat/completions".to_string(), + "/v1/responses".to_string(), + ], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let chat = serde_json::json!({ + "messages": [{ + "role": "user", + "content": [ + {"type": "text", "text": "Describe this"}, + {"type": "image_url", "image_url": {"url": "https://example.com/image.png"}} + ] + }], + "stream": true + }); + let responses = serde_json::json!({ + "input": [ + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "Hello"}] + }, + { + "type": "message", + "id": "msg_prior", + "status": "completed", + "role": "assistant", + "content": [{ + "type": "output_text", + "text": "Hi", + "annotations": [], + "logprobs": [] + }] + } + ], + "stream": true + }); + + assert!( + StreamContinuation::from_request( + "/v1/chat/completions", + &Method::POST, + &serde_json::to_vec(&chat).unwrap(), + &config, + ) + .is_some() + ); + assert!( + StreamContinuation::from_request( + "/v1/responses", + &Method::POST, + &serde_json::to_vec(&responses).unwrap(), + &config, + ) + .is_some() + ); + } + + #[test] + fn responses_continuation_stitches_lifecycle_identity_and_sequence() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/responses".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let mut state = StreamContinuation::from_request( + "/v1/responses", + &Method::POST, + br#"{"model":"requested","input":"Hello","stream":true}"#, + &config, + ) + .unwrap(); + for event in [ + "event: response.created\ndata: {\"type\":\"response.created\",\"sequence_number\":0,\"response\":{\"id\":\"resp_first\",\"model\":\"first\",\"output\":[]}}\n\n", + "event: response.output_item.added\ndata: {\"type\":\"response.output_item.added\",\"sequence_number\":1,\"output_index\":0,\"item\":{\"id\":\"msg_first\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[]}}\n\n", + "event: response.content_part.added\ndata: {\"type\":\"response.content_part.added\",\"sequence_number\":2,\"item_id\":\"msg_first\",\"output_index\":0,\"content_index\":0,\"part\":{\"type\":\"output_text\",\"text\":\"\"}}\n\n", + "event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"sequence_number\":3,\"item_id\":\"msg_first\",\"output_index\":0,\"content_index\":0,\"delta\":\"Hello\"}\n\n", + ] { + assert!( + state + .observe_event(event.as_bytes(), false) + .unwrap() + .forward + ); + } + + for event in [ + "event: response.created\ndata: {\"type\":\"response.created\",\"sequence_number\":0,\"response\":{\"id\":\"resp_second\",\"model\":\"second\",\"output\":[]}}\n\n", + "event: response.output_item.added\ndata: {\"type\":\"response.output_item.added\",\"sequence_number\":1,\"output_index\":0,\"item\":{\"id\":\"msg_second\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[]}}\n\n", + "event: response.content_part.added\ndata: {\"type\":\"response.content_part.added\",\"sequence_number\":2,\"item_id\":\"msg_second\",\"output_index\":0,\"content_index\":0,\"part\":{\"type\":\"output_text\",\"text\":\"\"}}\n\n", + ] { + assert!(!state.observe_event(event.as_bytes(), true).unwrap().forward); + } + + let delta = state + .observe_event( + b"event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"sequence_number\":3,\"item_id\":\"msg_second\",\"output_index\":0,\"content_index\":0,\"delta\":\" world\"}\n\n", + true, + ) + .unwrap(); + let delta = response_event_json(&delta.event); + assert_eq!(delta["sequence_number"], 4); + assert_eq!(delta["item_id"], "msg_first"); + + let completed = state + .observe_event( + b"event: response.completed\ndata: {\"type\":\"response.completed\",\"sequence_number\":4,\"response\":{\"id\":\"resp_second\",\"model\":\"second\",\"output\":[{\"id\":\"msg_second\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\" world\"}]}]}}\n\n", + true, + ) + .unwrap(); + let completed_json = response_event_json(&completed.event); + assert!(completed.terminal); + assert_eq!(completed_json["sequence_number"], 5); + assert_eq!(completed_json["response"]["id"], "resp_first"); + assert_eq!(completed_json["response"]["model"], "first"); + assert_eq!(completed_json["response"]["output"][0]["id"], "msg_first"); + assert_eq!( + completed_json["response"]["output"][0]["content"][0]["text"], + "Hello world" + ); + } + + #[test] + fn responses_fallback_identity_is_frozen_when_initial_events_omit_ids() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/responses".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let mut state = StreamContinuation::from_request( + "/v1/responses", + &Method::POST, + br#"{"model":"requested","input":"Hello","stream":true}"#, + &config, + ) + .unwrap(); + + let created = state + .observe_event( + b"event: response.created\ndata: {\"type\":\"response.created\",\"sequence_number\":0,\"response\":{\"model\":\"first\",\"output\":[]}}\n\n", + false, + ) + .unwrap(); + let response_id = response_event_json(&created.event)["response"]["id"] + .as_str() + .unwrap() + .to_string(); + + let in_progress = state + .observe_event( + b"event: response.in_progress\ndata: {\"type\":\"response.in_progress\",\"sequence_number\":1,\"response\":{\"id\":\"resp_late\",\"model\":\"late\",\"output\":[]}}\n\n", + false, + ) + .unwrap(); + assert_eq!( + response_event_json(&in_progress.event)["response"]["id"], + response_id + ); + + let item = state + .observe_event( + b"event: response.output_item.added\ndata: {\"type\":\"response.output_item.added\",\"sequence_number\":2,\"output_index\":0,\"item\":{\"type\":\"message\",\"role\":\"assistant\",\"content\":[]}}\n\n", + false, + ) + .unwrap(); + let item_id = response_event_json(&item.event)["item"]["id"] + .as_str() + .unwrap() + .to_string(); + + let delta = state + .observe_event( + b"event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"sequence_number\":3,\"item_id\":\"msg_late\",\"output_index\":0,\"content_index\":0,\"delta\":\"Hello\"}\n\n", + false, + ) + .unwrap(); + assert_eq!(response_event_json(&delta.event)["item_id"], item_id); + } + + #[test] + fn responses_retry_setup_is_suppressed_after_a_sparse_initial_stream() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/responses".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let mut state = StreamContinuation::from_request( + "/v1/responses", + &Method::POST, + br#"{"model":"requested","input":"Hello","stream":true}"#, + &config, + ) + .unwrap(); + state + .observe_event( + b"event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"sequence_number\":0,\"item_id\":\"msg_first\",\"output_index\":0,\"content_index\":0,\"delta\":\"Hello\"}\n\n", + false, + ) + .unwrap(); + + let created = state + .observe_event( + b"event: response.created\ndata: {\"type\":\"response.created\",\"sequence_number\":0,\"response\":{\"id\":\"resp_second\",\"model\":\"second\",\"output\":[]}}\n\n", + true, + ) + .unwrap(); + + assert!(!created.forward); + } + + #[test] + fn responses_done_only_text_is_preserved_and_disables_continuation() { + let config = StreamContinuationConfig { + enabled: true, + endpoints: vec!["/v1/responses".to_string()], + max_attempts: 1, + max_buffered_bytes: 1024, + idle_timeout_ms: None, + }; + let mut state = StreamContinuation::from_request( + "/v1/responses", + &Method::POST, + br#"{"model":"requested","input":"Hello","stream":true}"#, + &config, + ) + .unwrap(); + let event = b"event: response.output_text.done\ndata: {\"type\":\"response.output_text.done\",\"sequence_number\":0,\"item_id\":\"msg_first\",\"output_index\":0,\"content_index\":0,\"text\":\"Hello\"}\n\n"; + + let observation = state.observe_event(event, false).unwrap(); + + assert!(!observation.safe); + assert_eq!(observation.event.as_ref(), event); + assert!(!state.is_continuable()); + } + + fn response_event_json(event: &[u8]) -> Value { + let data = event + .split(|byte| *byte == b'\n') + .find_map(|line| line.strip_prefix(b"data: ")) + .unwrap(); + serde_json::from_slice(data).unwrap() + } +} diff --git a/onwards/src/strict/handlers.rs b/onwards/src/strict/handlers.rs index 2911edc50..0cb5fe2a7 100644 --- a/onwards/src/strict/handlers.rs +++ b/onwards/src/strict/handlers.rs @@ -29,7 +29,7 @@ use crate::AppState; use crate::client::HttpClient; use crate::errors::OnwardsErrorResponse; use crate::extract_model_from_request; -use crate::handlers::{ResolvedTrust, target_message_handler}; +use crate::handlers::{ResolvedTrust, target_message_handler_with_continuation}; use crate::traits::RequestContext; use axum::Json; use axum::body::Body; @@ -87,7 +87,16 @@ pub async fn models_handler( pub async fn chat_completions_handler( State(state): State>, headers: HeaderMap, - Json(mut request): Json, + Json(request): Json, +) -> Response { + chat_completions_handler_inner(state, headers, request, None).await +} + +async fn chat_completions_handler_inner( + state: AppState, + headers: HeaderMap, + mut request: ChatCompletionRequest, + continuation_eligibility_path: Option<&'static str>, ) -> Response { request.scrub_request_id_fields(); @@ -120,13 +129,24 @@ pub async fn chat_completions_handler() + .is_some(); if response_is_sse && !is_streaming { debug!( @@ -135,8 +155,10 @@ pub async fn chat_completions_handler( // the normal adapter/passthrough logic below. if req.uri().path().ends_with("/chat/completions") { return match Json::::from_request(req, &state).await { - Ok(chat_request) => chat_completions_handler(State(state), headers, chat_request).await, + Ok(Json(chat_request)) => { + chat_completions_handler_inner(state, headers, chat_request, Some("/v1/responses")) + .await + } Err(rejection) => rejection.into_response(), }; } @@ -282,7 +307,6 @@ pub async fn responses_handler( // OpenAI requires additionalProperties: false in tool schemas even for /v1/responses // Add it if missing to ensure compatibility - let mut request = request; if let Some(ref mut tools) = request.tools { for tool in tools.iter_mut() { if let super::schemas::responses::Tool::Function { parameters, .. } = tool @@ -323,8 +347,12 @@ pub async fn responses_handler( // Internal errors (generated by onwards itself) are never sanitized let final_response = if response.status().is_success() { let response_is_sse = response_is_sse(&response); + let preserve_encoded_stream = response + .extensions() + .get::() + .is_some(); - if is_streaming || response_is_sse { + if (is_streaming || response_is_sse) && !preserve_encoded_stream { sanitize_streaming_responses_response( response, resolved_model, @@ -332,6 +360,8 @@ pub async fn responses_handler( response_id_override, ) .await + } else if is_streaming || response_is_sse { + response } else { sanitize_responses_response(response, resolved_model, response_id_override).await } @@ -492,9 +522,15 @@ pub async fn completions_handler( if response.status().is_success() { let response_is_sse = response_is_sse(&response); + let preserve_encoded_stream = response + .extensions() + .get::() + .is_some(); - if is_streaming || response_is_sse { + if (is_streaming || response_is_sse) && !preserve_encoded_stream { sanitize_streaming_completions_response(response, resolved_model, trusted).await + } else if is_streaming || response_is_sse { + response } else { sanitize_completions_response(response, resolved_model).await } @@ -954,14 +990,13 @@ async fn handle_adapter_request( ); // Mark as completed in the response store - if let Ok(response_value) = serde_json::to_value(&responses_response) { - if let Err(e) = state + if let Ok(response_value) = serde_json::to_value(&responses_response) + && let Err(e) = state .response_store .complete(&response_id, &response_value, parts.status.as_u16()) .await - { - warn!(error = %e, response_id = %response_id, "Failed to mark response as completed"); - } + { + warn!(error = %e, response_id = %response_id, "Failed to mark response as completed"); } // Return the converted response @@ -1098,7 +1133,8 @@ async fn handle_streaming_adapter_request( { Ok(mut req) => { *req.headers_mut() = headers; + req.extensions_mut() + .insert(crate::handlers::ContinuationEligibilityPath( + "/v1/responses", + )); if let Some(reasoning) = canonical_reasoning { req.extensions_mut().insert(reasoning.clone()); } @@ -1300,7 +1340,7 @@ async fn forward_request_raw( }; // Forward using the standard handler - match target_message_handler(State(state), request).await { + match target_message_handler_with_continuation(State(state), request).await { Ok(response) => response, Err(err) => err.into_response(), } @@ -1308,10 +1348,20 @@ async fn forward_request_raw( /// Forward a validated request to the upstream provider async fn forward_request( + state: AppState, + headers: HeaderMap, + path: &str, + body_bytes: Vec, +) -> ForwardResult { + forward_request_with_eligibility(state, headers, path, body_bytes, None).await +} + +async fn forward_request_with_eligibility( state: AppState, mut headers: HeaderMap, path: &str, body_bytes: Vec, + continuation_eligibility_path: Option<&'static str>, ) -> ForwardResult { // Ensure content-type is set headers.insert( @@ -1331,7 +1381,25 @@ async fn forward_request( } let request = match request_builder.body(Body::from(body_bytes)) { - Ok(req) => req, + Ok(mut req) => { + let canonical_path = match path { + "/chat/completions" => Some("/v1/chat/completions"), + "/responses" => Some("/v1/responses"), + "/completions" => Some("/v1/completions"), + _ => None, + }; + if let Some(canonical_path) = canonical_path { + req.extensions_mut() + .insert(crate::handlers::CanonicalRequestPath(canonical_path)); + } + if let Some(eligibility_path) = continuation_eligibility_path { + req.extensions_mut() + .insert(crate::handlers::ContinuationEligibilityPath( + eligibility_path, + )); + } + req + } Err(e) => { error!(error = %e, "Failed to build request"); let response = error_response( @@ -1348,10 +1416,11 @@ async fn forward_request( }; // Use the existing target message handler - let (response, internal_error) = match target_message_handler(State(state), request).await { - Ok(response) => (response, false), - Err(err) => (err.into_response(), true), - }; + let (response, internal_error) = + match target_message_handler_with_continuation(State(state), request).await { + Ok(response) => (response, false), + Err(err) => (err.into_response(), true), + }; let trusted = response .extensions() @@ -1536,12 +1605,19 @@ async fn sanitize_streaming_chat_response( original_model: String, trusted: bool, ) -> Response { - // Wrap with SseBufferedStream to ensure we receive complete SSE events (delimited by \n\n). + // Use checked framing so sanitization receives complete SSE events. // Providers may send partial chunks that split JSON across network packets. // This buffering ensures we can successfully parse JSON in each event. + let buffer_limit = response + .extensions() + .get::() + .map(|limit| limit.0); let body_stream = http_body_util::BodyExt::into_data_stream(std::mem::take(response.body_mut())); - let buffered_stream = crate::sse::SseBufferedStream::new(body_stream); + let buffered_stream = match buffer_limit { + Some(limit) => crate::sse::CheckedSseStream::with_max_buffer_size(body_stream, limit), + None => crate::sse::CheckedSseStream::new(body_stream), + }; let stream_fallback_id = generated_chat_completion_id(); let sanitized_stream = buffered_stream.map(move |chunk_result| { @@ -1553,7 +1629,8 @@ async fn sanitize_streaming_chat_response( // Process SSE line-by-line for streaming chunks, // preserving empty lines and ignoring comment/event-type lines for line in chunk_str.lines() { - if let Some(data_part) = line.strip_prefix("data: ") { + if let Some(data_part) = line.strip_prefix("data:") { + let data_part = data_part.strip_prefix(' ').unwrap_or(data_part); // This is a data line // Handle [DONE] marker @@ -1754,9 +1831,16 @@ async fn sanitize_streaming_completions_response( original_model: String, trusted: bool, ) -> Response { + let buffer_limit = response + .extensions() + .get::() + .map(|limit| limit.0); let body_stream = http_body_util::BodyExt::into_data_stream(std::mem::take(response.body_mut())); - let buffered_stream = crate::sse::SseBufferedStream::new(body_stream); + let buffered_stream = match buffer_limit { + Some(limit) => crate::sse::CheckedSseStream::with_max_buffer_size(body_stream, limit), + None => crate::sse::CheckedSseStream::new(body_stream), + }; let stream_fallback_id = generated_completion_id(); let sanitized_stream = buffered_stream.map(move |chunk_result| { @@ -1766,7 +1850,8 @@ async fn sanitize_streaming_completions_response( let mut sanitized_lines = Vec::new(); for line in chunk_str.lines() { - if let Some(data_part) = line.strip_prefix("data: ") { + if let Some(data_part) = line.strip_prefix("data:") { + let data_part = data_part.strip_prefix(' ').unwrap_or(data_part); if data_part.trim() == "[DONE]" { sanitized_lines.push(line.to_string()); continue; @@ -2032,12 +2117,19 @@ async fn sanitize_streaming_responses_response( trusted: bool, response_id_override: Option, ) -> Response { - // Wrap with SseBufferedStream to ensure we receive complete SSE events (delimited by \n\n). + // Use checked framing so sanitization receives complete SSE events. // Providers may send partial chunks that split JSON across network packets. // This buffering ensures we can successfully parse JSON in each event. + let buffer_limit = response + .extensions() + .get::() + .map(|limit| limit.0); let body_stream = http_body_util::BodyExt::into_data_stream(std::mem::take(response.body_mut())); - let buffered_stream = crate::sse::SseBufferedStream::new(body_stream); + let buffered_stream = match buffer_limit { + Some(limit) => crate::sse::CheckedSseStream::with_max_buffer_size(body_stream, limit), + None => crate::sse::CheckedSseStream::new(body_stream), + }; let response_id_override = response_id_override.clone(); let stream_fallback_response_id = response_id_override .clone() @@ -2051,7 +2143,8 @@ async fn sanitize_streaming_responses_response( // Process SSE line-by-line for streaming chunks for line in chunk_str.lines() { - if let Some(data_part) = line.strip_prefix("data: ") { + if let Some(data_part) = line.strip_prefix("data:") { + let data_part = data_part.strip_prefix(' ').unwrap_or(data_part); // This is a data line // Handle [DONE] marker @@ -2398,6 +2491,30 @@ mod tests { use std::sync::Arc; use tower::ServiceExt; + #[tokio::test] + async fn strict_streaming_sanitizers_preserve_done_without_post_colon_space() { + let completion = Response::builder() + .body(Body::from("data:[DONE]\n\n")) + .unwrap(); + let chat = Response::builder() + .body(Body::from("data:[DONE]\n\n")) + .unwrap(); + let responses = Response::builder() + .body(Body::from("data:[DONE]\n\n")) + .unwrap(); + + for response in [ + sanitize_streaming_completions_response(completion, "m".to_string(), false).await, + sanitize_streaming_chat_response(chat, "m".to_string(), false).await, + sanitize_streaming_responses_response(responses, "m".to_string(), false, None).await, + ] { + let body = axum::body::to_bytes(response.into_body(), usize::MAX) + .await + .unwrap(); + assert_eq!(body.as_ref(), b"data:[DONE]\n\n"); + } + } + // --- ZDR no-payload-logging regression tests (COR-497) --- /// A `MakeWriter` that appends every emitted log byte into a shared buffer, diff --git a/onwards/src/target.rs b/onwards/src/target.rs index b8a61aeeb..b0bb56507 100644 --- a/onwards/src/target.rs +++ b/onwards/src/target.rs @@ -125,6 +125,34 @@ pub enum LoadBalanceStrategy { Priority, } +fn default_stream_continuation_attempts() -> usize { + 1 +} + +fn default_stream_continuation_buffer_bytes() -> usize { + 1024 * 1024 +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct StreamContinuationConfig { + #[serde(default)] + pub enabled: bool, + #[serde(default)] + pub endpoints: Vec, + #[serde(default = "default_stream_continuation_attempts")] + pub max_attempts: usize, + #[serde(default = "default_stream_continuation_buffer_bytes")] + pub max_buffered_bytes: usize, + #[serde(default)] + pub idle_timeout_ms: Option, +} + +impl StreamContinuationConfig { + pub fn enabled_for_path(&self, path: &str) -> bool { + self.enabled && self.endpoints.iter().any(|endpoint| endpoint == path) + } +} + /// Configuration for fallback behavior when requests fail #[derive(Debug, Clone, Serialize, Deserialize, Default)] pub struct FallbackConfig { @@ -166,6 +194,9 @@ pub struct FallbackConfig { /// Bounds only the inter-attempt sleeps, not upstream request time. #[serde(default)] pub max_total_backoff_ms: Option, + + #[serde(default)] + pub stream_continuation: Option, } impl FallbackConfig { @@ -1077,6 +1108,7 @@ impl Targets { // Convert provider specs to providers // Pool-level sanitize_response enables sanitization for all providers let pool_sanitize = pool_config.sanitize_response; + let pool_response_headers = pool_config.response_headers; let providers: Vec = pool_config .providers .into_iter() @@ -1088,6 +1120,12 @@ impl Targets { .map(|cl| cl.max_concurrent_requests); // Enable sanitization if either pool or provider level is true spec.sanitize_response = pool_sanitize || spec.sanitize_response; + let mut response_headers = pool_response_headers.clone().unwrap_or_default(); + if let Some(provider_headers) = spec.response_headers.take() { + response_headers.extend(provider_headers); + } + spec.response_headers = + (!response_headers.is_empty()).then_some(response_headers); let target: Target = spec.into(); match concurrency_limit { Some(limit) => Provider::with_concurrency_limit(target, weight, limit), @@ -1254,6 +1292,37 @@ mod tests { use dashmap::DashMap; use std::collections::HashSet; use std::sync::Arc; + + #[test] + fn stream_continuation_config_defaults_are_conservative() { + let config: FallbackConfig = serde_json::from_str( + r#"{ + "enabled": true, + "stream_continuation": {"enabled": true} + }"#, + ) + .unwrap(); + let stream = config.stream_continuation.unwrap(); + assert_eq!(stream.max_attempts, 1); + assert_eq!(stream.max_buffered_bytes, 1024 * 1024); + assert_eq!(stream.idle_timeout_ms, None); + assert!(!stream.enabled_for_path("/v1/completions")); + } + + #[test] + fn stream_continuation_matches_only_configured_exact_path() { + let config: StreamContinuationConfig = serde_json::from_str( + r#"{ + "enabled": true, + "endpoints": ["/v1/completions"] + }"#, + ) + .unwrap(); + assert!(config.enabled_for_path("/v1/completions")); + assert!(!config.enabled_for_path("/v1/chat/completions")); + assert!(!config.enabled_for_path("/v1/completions/extra")); + } + pub struct MockConfigWatcher { configs: Vec>, }