Fix stream timeout semantics

This commit is contained in:
fawney19
2026-05-25 20:09:37 +08:00
parent 54c5d5803b
commit d3249485fa
7 changed files with 561 additions and 189 deletions
@@ -120,7 +120,6 @@ const SSE_CONTROL_FILTER_MAX_BUFFER_BYTES: usize = 1024 * 1024;
const STREAM_IDLE_LOG_INTERVAL: Duration = Duration::from_secs(60);
const STREAM_IDLE_LOG_INTERVAL_MS: u64 = 60_000;
const REWRITTEN_STREAM_PREFETCH_TIMEOUT: Duration = Duration::from_millis(750);
const OPENAI_IMAGE_STREAM_DEFAULT_TOTAL_TIMEOUT_MS: u64 = 900_000;
fn record_sync_terminal_usage(
state: &AppState,
@@ -1386,22 +1385,6 @@ fn encode_openai_image_failed_event(
Ok(Bytes::from(event))
}
fn resolve_openai_image_stream_total_timeout_ms(
plan_kind: &str,
plan: &ExecutionPlan,
) -> Option<u64> {
if plan_kind != OPENAI_IMAGE_STREAM_PLAN_KIND {
return None;
}
Some(
plan.timeouts
.as_ref()
.and_then(|timeouts| timeouts.total_ms)
.unwrap_or(OPENAI_IMAGE_STREAM_DEFAULT_TOTAL_TIMEOUT_MS)
.max(1),
)
}
fn should_limit_direct_finalize_prefetch(plan_kind: &str, has_local_stream_rewriter: bool) -> bool {
plan_kind == OPENAI_IMAGE_STREAM_PLAN_KIND || has_local_stream_rewriter
}
@@ -2622,8 +2605,6 @@ async fn execute_stream_from_frame_stream(
let candidate_id_for_report = candidate_id.clone();
let candidate_index_for_report = candidate_index.clone();
let is_openai_image_stream_for_report = plan_kind == OPENAI_IMAGE_STREAM_PLAN_KIND;
let openai_image_stream_total_timeout_ms =
resolve_openai_image_stream_total_timeout_ms(plan_kind, &plan);
let plan_for_report = plan;
let emit_passthrough_sse_terminal_error = skip_direct_finalize_prefetch
&& response_headers_indicate_sse(&upstream_headers)
@@ -2883,77 +2864,14 @@ async fn execute_stream_from_frame_stream(
}
if terminal_failure.is_none() && !reached_eof {
let mut image_stream_total_timeout = openai_image_stream_total_timeout_ms
.map(|timeout_ms| Box::pin(tokio::time::sleep(Duration::from_millis(timeout_ms))));
loop {
let next_frame_result = if let Some(timeout_sleep) =
image_stream_total_timeout.as_mut()
{
tokio::select! {
biased;
_ = tx.closed(), if client_visible_stream_completed => {
downstream_dropped = true;
break;
}
result = next_stream_frame(&mut buffered_frames, &mut lines) => result,
_ = timeout_sleep.as_mut() => {
let timeout_ms = openai_image_stream_total_timeout_ms
.unwrap_or(OPENAI_IMAGE_STREAM_DEFAULT_TOTAL_TIMEOUT_MS);
let elapsed_ms = stream_started_at_for_report
.elapsed()
.as_millis()
.min(u128::from(u64::MAX)) as u64;
warn!(
event_name = "openai_image_stream_total_timeout",
log_type = "ops",
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
candidate_index = candidate_index_for_report.as_str(),
plan_kind = plan_kind_for_report.as_str(),
provider_name = plan_for_report.provider_name.as_deref().unwrap_or("-"),
endpoint_id = %plan_for_report.endpoint_id,
key_id = %plan_for_report.key_id,
model_name = plan_for_report.model_name.as_deref().unwrap_or("-"),
elapsed_ms,
timeout_ms,
provider_bytes = provider_stream_bytes.load(Ordering::Relaxed),
client_bytes = client_stream_bytes.load(Ordering::Relaxed),
last_upstream_frame_elapsed_ms = last_upstream_frame_elapsed_ms.load(Ordering::Relaxed),
last_client_chunk_elapsed_ms = last_client_chunk_elapsed_ms.load(Ordering::Relaxed),
"gateway OpenAI image stream exceeded total timeout"
);
telemetry = Some(ExecutionTelemetry {
ttfb_ms: telemetry
.as_ref()
.and_then(|telemetry| telemetry.ttfb_ms)
.or_else(|| {
usage_stream_telemetry
.as_ref()
.and_then(|telemetry| telemetry.ttfb_ms)
}),
elapsed_ms: Some(elapsed_ms),
upstream_bytes: Some(provider_stream_bytes.load(Ordering::Relaxed)),
});
terminal_failure = Some(build_stream_failure_report(
"image_stream_total_timeout",
format!(
"OpenAI image stream exceeded total timeout of {timeout_ms}ms"
),
504,
));
break;
}
}
} else {
tokio::select! {
biased;
_ = tx.closed(), if client_visible_stream_completed => {
downstream_dropped = true;
break;
}
result = next_stream_frame(&mut buffered_frames, &mut lines) => result,
let next_frame_result = tokio::select! {
biased;
_ = tx.closed(), if client_visible_stream_completed => {
downstream_dropped = true;
break;
}
result = next_stream_frame(&mut buffered_frames, &mut lines) => result,
};
let next_frame = match next_frame_result {
Ok(frame) => frame,
@@ -4804,7 +4722,7 @@ mod tests {
}
#[tokio::test]
async fn openai_image_stream_total_timeout_emits_image_failed_event() {
async fn openai_image_stream_ignores_plan_total_timeout() {
let state = AppState::new().expect("app state should build");
let plan = ExecutionPlan {
request_id: "req-image-stream-timeout".into(),
@@ -4876,17 +4794,19 @@ mod tests {
.expect("execution should succeed")
.expect("execution should return a client response");
let body = tokio::time::timeout(
Duration::from_secs(2),
to_bytes(response.into_body(), usize::MAX),
)
.await
.expect("timeout failure should close the response body")
.expect("response body should read");
let text = String::from_utf8(body.to_vec()).expect("response body should be utf8");
assert!(text.contains(": aether-keepalive\n\n"));
assert!(text.contains("event: image_generation.failed"));
assert!(text.contains("\"type\":\"image_stream_total_timeout\""));
let mut body_stream = response.into_body().into_data_stream();
let keepalive = tokio::time::timeout(Duration::from_millis(50), body_stream.next())
.await
.expect("initial keepalive should be emitted")
.expect("body should yield initial keepalive")
.expect("initial keepalive should be ok");
assert_eq!(keepalive.as_ref(), b": aether-keepalive\n\n");
let next_chunk = tokio::time::timeout(Duration::from_millis(100), body_stream.next()).await;
assert!(
next_chunk.is_err(),
"stream total_ms must not synthesize an image failure or close the response body"
);
}
#[tokio::test]
@@ -19,6 +19,7 @@ use base64::Engine as _;
use flate2::read::{DeflateDecoder, GzDecoder};
use flate2::write::GzEncoder;
use flate2::Compression;
use futures_util::StreamExt;
use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
use reqwest::redirect::Policy;
use serde::Serialize;
@@ -40,6 +41,8 @@ const HUB_RELAY_CONTENT_TYPE: &str = "application/vnd.aether.tunnel-envelope";
const HUB_RELAY_ERROR_HEADER: &str = "x-aether-tunnel-error";
const TUNNEL_RELAY_PATH_PREFIX: &str = "/api/internal/tunnel/relay";
const DEFAULT_TUNNEL_TIMEOUT_MS: u64 = 60_000;
const DEFAULT_STREAM_FIRST_BYTE_TIMEOUT_MS: u64 = 30_000;
const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000;
const MIN_TUNNEL_TIMEOUT_SECS: u64 = 1;
const MAX_TUNNEL_TIMEOUT_SECS: u64 = 300;
pub(crate) fn format_upstream_request_error(err: &reqwest::Error) -> String {
@@ -247,7 +250,8 @@ impl DirectSyncExecutionRuntime {
let ttfb_ms = started_at.elapsed().as_millis() as u64;
let status_code = response.status_code();
let headers = response.headers();
let body_bytes = response.bytes().await?;
let (body_bytes, stream_ttfb_ms) =
response.bytes_with_stream_timeout(plan, started_at).await?;
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)
.unwrap_or_else(|| body_bytes.to_vec());
let elapsed_ms = started_at.elapsed().as_millis() as u64;
@@ -267,7 +271,7 @@ impl DirectSyncExecutionRuntime {
headers,
body,
telemetry: Some(ExecutionTelemetry {
ttfb_ms: Some(ttfb_ms),
ttfb_ms: stream_ttfb_ms.or(Some(ttfb_ms)),
elapsed_ms: Some(elapsed_ms),
upstream_bytes: Some(upstream_bytes),
}),
@@ -544,14 +548,8 @@ async fn execute_sync_plan_via_local_tunnel_inner(
let status_code = response.status();
let headers = collect_tunnel_response_headers(response.headers());
let proxy_timing = execution_header_for_log(&headers, "x-proxy-timing").unwrap_or("-");
let mut body_bytes = Vec::new();
while let Some(chunk) = response
.next_chunk()
.await
.map_err(ExecutionRuntimeTransportError::UpstreamRequest)?
{
body_bytes.extend_from_slice(&chunk);
}
let (body_bytes, stream_ttfb_ms) =
collect_local_tunnel_response_body(response, plan, started_at).await?;
let decoded_body_bytes =
decode_response_body_bytes(&headers, &body_bytes).unwrap_or_else(|| body_bytes.clone());
let elapsed_ms = started_at.elapsed().as_millis() as u64;
@@ -600,7 +598,7 @@ async fn execute_sync_plan_via_local_tunnel_inner(
headers,
body,
telemetry: Some(ExecutionTelemetry {
ttfb_ms: Some(ttfb_ms),
ttfb_ms: stream_ttfb_ms.or(Some(ttfb_ms)),
elapsed_ms: Some(elapsed_ms),
upstream_bytes: Some(upstream_bytes),
}),
@@ -608,6 +606,38 @@ async fn execute_sync_plan_via_local_tunnel_inner(
})
}
async fn collect_local_tunnel_response_body(
mut response: tunnel::DirectRelayResponse,
plan: &ExecutionPlan,
started_at: Instant,
) -> Result<(Vec<u8>, Option<u64>), ExecutionRuntimeTransportError> {
let mut body_bytes = Vec::new();
let mut first_byte_ms = None;
let first_byte_timeout = plan
.stream
.then(|| resolve_stream_first_byte_timeout(plan))
.flatten();
loop {
let item = if first_byte_ms.is_none() && plan.stream {
await_stream_body_first_item(response.next_chunk(), started_at, first_byte_timeout)
.await?
} else {
response.next_chunk().await
}
.map_err(ExecutionRuntimeTransportError::UpstreamRequest)?;
let Some(chunk) = item else {
break;
};
if plan.stream && first_byte_ms.is_none() && !chunk.is_empty() {
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
}
body_bytes.extend_from_slice(&chunk);
}
Ok((body_bytes, first_byte_ms))
}
fn build_direct_tunnel_request_meta(
plan: &ExecutionPlan,
headers: &HeaderMap,
@@ -741,6 +771,26 @@ impl DirectHttpResponse {
}
}
async fn bytes_with_stream_timeout(
self,
plan: &ExecutionPlan,
started_at: Instant,
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
if !plan.stream {
return self.bytes().await.map(|bytes| (bytes, None));
}
let first_byte_timeout = resolve_stream_first_byte_timeout(plan);
match self {
DirectHttpResponse::Reqwest(response) => {
collect_reqwest_stream_body(response, started_at, first_byte_timeout).await
}
DirectHttpResponse::BrowserWreq(response) => {
collect_wreq_stream_body(response, started_at, first_byte_timeout).await
}
}
}
fn into_direct_upstream_response(self) -> DirectUpstreamResponse {
match self {
DirectHttpResponse::Reqwest(response) => DirectUpstreamResponse::Reqwest(response),
@@ -751,6 +801,92 @@ impl DirectHttpResponse {
}
}
async fn await_stream_body_first_item<T, F>(
future: F,
started_at: Instant,
timeout: Option<Duration>,
) -> Result<T, ExecutionRuntimeTransportError>
where
F: Future<Output = T>,
{
let Some(timeout) = timeout else {
return Ok(future.await);
};
let Some(remaining) = timeout.checked_sub(started_at.elapsed()) else {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
stream_first_byte_timeout_message(timeout),
));
};
if remaining.is_zero() {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(
stream_first_byte_timeout_message(timeout),
));
}
tokio::time::timeout(remaining, future).await.map_err(|_| {
ExecutionRuntimeTransportError::UpstreamRequest(stream_first_byte_timeout_message(timeout))
})
}
async fn collect_reqwest_stream_body(
response: reqwest::Response,
started_at: Instant,
first_byte_timeout: Option<Duration>,
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
let mut stream = response.bytes_stream();
let mut body_bytes = Vec::new();
let mut first_byte_ms = None;
loop {
let item = if first_byte_ms.is_none() {
await_stream_body_first_item(stream.next(), started_at, first_byte_timeout).await?
} else {
stream.next().await
};
let Some(item) = item else {
break;
};
let chunk = item.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
})?;
if first_byte_ms.is_none() && !chunk.is_empty() {
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
}
body_bytes.extend_from_slice(&chunk);
}
Ok((Bytes::from(body_bytes), first_byte_ms))
}
async fn collect_wreq_stream_body(
response: wreq::Response,
started_at: Instant,
first_byte_timeout: Option<Duration>,
) -> Result<(Bytes, Option<u64>), ExecutionRuntimeTransportError> {
let mut stream = response.bytes_stream();
let mut body_bytes = Vec::new();
let mut first_byte_ms = None;
loop {
let item = if first_byte_ms.is_none() {
await_stream_body_first_item(stream.next(), started_at, first_byte_timeout).await?
} else {
stream.next().await
};
let Some(item) = item else {
break;
};
let chunk = item.map_err(|err| {
ExecutionRuntimeTransportError::BrowserBody(format_wreq_upstream_request_error(&err))
})?;
if first_byte_ms.is_none() && !chunk.is_empty() {
first_byte_ms = Some(started_at.elapsed().as_millis() as u64);
}
body_bytes.extend_from_slice(&chunk);
}
Ok((Bytes::from(body_bytes), first_byte_ms))
}
async fn send_via_browser_wreq_transport(
plan: &ExecutionPlan,
method: reqwest::Method,
@@ -1040,9 +1176,8 @@ fn resolve_relay_timeout_seconds(plan: &ExecutionPlan) -> u64 {
fn resolve_tunnel_first_byte_timeout(plan: &ExecutionPlan) -> Option<Duration> {
plan.stream.then(|| {
Duration::from_millis(
resolve_selected_tunnel_timeout_ms(plan).unwrap_or(DEFAULT_TUNNEL_TIMEOUT_MS),
)
resolve_stream_first_byte_timeout(plan)
.unwrap_or_else(|| Duration::from_millis(DEFAULT_TUNNEL_TIMEOUT_MS))
})
}
@@ -1050,20 +1185,24 @@ fn resolve_non_stream_total_timeout(plan: &ExecutionPlan) -> Option<Duration> {
if plan.stream {
return None;
}
plan.timeouts
let timeout_ms = plan
.timeouts
.as_ref()
.and_then(|timeouts| timeouts.total_ms)
.map(|value| Duration::from_millis(value.max(1)))
.unwrap_or(DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS);
Some(Duration::from_millis(timeout_ms.max(1)))
}
pub(crate) fn resolve_stream_first_byte_timeout(plan: &ExecutionPlan) -> Option<Duration> {
if !plan.stream {
return None;
}
plan.timeouts
let timeout_ms = plan
.timeouts
.as_ref()
.and_then(|timeouts| timeouts.first_byte_ms.or(timeouts.total_ms))
.map(|value| Duration::from_millis(value.max(1)))
.and_then(|timeouts| timeouts.first_byte_ms)
.unwrap_or(DEFAULT_STREAM_FIRST_BYTE_TIMEOUT_MS);
Some(Duration::from_millis(timeout_ms.max(1)))
}
pub(crate) async fn with_non_stream_total_timeout<T, F>(
@@ -1142,32 +1281,31 @@ pub(crate) fn stream_first_byte_timeout_message(timeout: Duration) -> String {
}
fn resolve_tunnel_timeout_metadata(plan: &ExecutionPlan) -> TunnelTimeoutMetadata {
TunnelTimeoutMetadata {
request_timeout_ms: plan
.timeouts
let request_timeout_ms = if plan.stream {
None
} else {
resolve_non_stream_total_timeout(plan)
.map(|timeout| u64::try_from(timeout.as_millis()).unwrap_or(u64::MAX))
};
let stream_first_byte_timeout_ms = if plan.stream {
resolve_stream_first_byte_timeout(plan)
.map(|timeout| u64::try_from(timeout.as_millis()).unwrap_or(u64::MAX))
} else {
plan.timeouts
.as_ref()
.and_then(|timeouts| timeouts.total_ms),
stream_first_byte_timeout_ms: plan
.timeouts
.as_ref()
.and_then(|timeouts| timeouts.first_byte_ms),
legacy_timeout_secs: timeout_ms_to_secs(
resolve_selected_tunnel_timeout_ms(plan).unwrap_or(DEFAULT_TUNNEL_TIMEOUT_MS),
),
}
}
.and_then(|timeouts| timeouts.first_byte_ms)
};
let legacy_timeout_ms = if plan.stream {
stream_first_byte_timeout_ms.unwrap_or(DEFAULT_TUNNEL_TIMEOUT_MS)
} else {
request_timeout_ms.unwrap_or(DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS)
};
fn resolve_selected_tunnel_timeout_ms(plan: &ExecutionPlan) -> Option<u64> {
plan.timeouts
.as_ref()
.and_then(|timeouts| {
if plan.stream {
timeouts.first_byte_ms.or(timeouts.total_ms)
} else {
timeouts.total_ms.or(timeouts.first_byte_ms)
}
})
.map(|value| value.max(1))
TunnelTimeoutMetadata {
request_timeout_ms,
stream_first_byte_timeout_ms,
legacy_timeout_secs: timeout_ms_to_secs(legacy_timeout_ms),
}
}
fn timeout_ms_to_secs(ms: u64) -> u64 {
@@ -1717,7 +1855,9 @@ mod tests {
use axum::http::HeaderMap as AxumHeaderMap;
use axum::routing::{any, post};
use axum::{Json, Router};
use base64::Engine as _;
use serde_json::json;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::sync::watch;
use super::{
@@ -1725,7 +1865,8 @@ mod tests {
build_execution_response_body, build_request_headers, execute_sync_plan,
record_manual_proxy_request_failure, record_manual_proxy_request_outcome,
record_manual_proxy_request_success, record_manual_proxy_stream_error,
resolve_execution_transport_controls, response_body_is_json, DirectSyncExecutionRuntime,
resolve_execution_transport_controls, resolve_non_stream_total_timeout,
resolve_stream_first_byte_timeout, response_body_is_json, DirectSyncExecutionRuntime,
ExecutionRuntimeTransportError, ExecutionTransportControls,
};
use crate::constants::{
@@ -1853,11 +1994,82 @@ mod tests {
);
assert!(meta.stream);
assert_eq!(meta.request_timeout_ms, Some(90_000));
assert_eq!(meta.request_timeout_ms, None);
assert_eq!(meta.stream_first_byte_timeout_ms, Some(12_345));
assert_eq!(meta.timeout, 13);
}
#[test]
fn stream_first_byte_timeout_uses_default_when_unconfigured() {
let mut plan = tunnel_timeout_plan(true);
plan.timeouts = None;
let timeout = resolve_stream_first_byte_timeout(&plan)
.expect("stream plans should have a first-byte default");
let meta = build_direct_tunnel_request_meta(
&plan,
&reqwest::header::HeaderMap::new(),
ExecutionTransportControls::default(),
);
assert_eq!(timeout, std::time::Duration::from_millis(30_000));
assert_eq!(meta.request_timeout_ms, None);
assert_eq!(meta.stream_first_byte_timeout_ms, Some(30_000));
assert_eq!(meta.timeout, 30);
}
#[test]
fn stream_first_byte_timeout_ignores_total_timeout() {
let mut plan = tunnel_timeout_plan(true);
plan.timeouts = Some(ExecutionTimeouts {
total_ms: Some(90_000),
..ExecutionTimeouts::default()
});
let timeout = resolve_stream_first_byte_timeout(&plan)
.expect("stream plans should have a first-byte default");
let meta = build_direct_tunnel_request_meta(
&plan,
&reqwest::header::HeaderMap::new(),
ExecutionTransportControls::default(),
);
assert_eq!(timeout, std::time::Duration::from_millis(30_000));
assert_eq!(meta.request_timeout_ms, None);
assert_eq!(meta.stream_first_byte_timeout_ms, Some(30_000));
assert_eq!(meta.timeout, 30);
}
#[test]
fn non_stream_total_timeout_defaults_to_provider_request_timeout() {
let mut plan = tunnel_timeout_plan(false);
plan.timeouts = None;
let timeout = resolve_non_stream_total_timeout(&plan)
.expect("non-stream plans should have a default total timeout");
assert_eq!(timeout, std::time::Duration::from_secs(300));
}
#[test]
fn tunnel_request_meta_uses_non_stream_default_instead_of_first_byte_default() {
let mut plan = tunnel_timeout_plan(false);
plan.timeouts = Some(ExecutionTimeouts {
first_byte_ms: Some(30_000),
..ExecutionTimeouts::default()
});
let meta = build_direct_tunnel_request_meta(
&plan,
&reqwest::header::HeaderMap::new(),
ExecutionTransportControls::default(),
);
assert!(!meta.stream);
assert_eq!(meta.request_timeout_ms, Some(300_000));
assert_eq!(meta.stream_first_byte_timeout_ms, Some(30_000));
assert_eq!(meta.timeout, 300);
}
fn tunnel_timeout_plan(stream: bool) -> ExecutionPlan {
ExecutionPlan {
request_id: "req-timeout".into(),
@@ -2095,6 +2307,113 @@ mod tests {
);
}
#[tokio::test]
async fn direct_sync_execution_runtime_applies_stream_first_byte_timeout_to_body_after_headers()
{
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("client should connect");
let mut request = [0_u8; 1024];
let _ = socket
.read(&mut request)
.await
.expect("request should read");
socket
.write_all(
b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ntransfer-encoding: chunked\r\n\r\n",
)
.await
.expect("headers should write");
socket.flush().await.expect("headers should flush");
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
let _ = socket.write_all(b"d\r\ndata: hello\n\n\r\n0\r\n\r\n").await;
});
let result = DirectSyncExecutionRuntime::new()
.execute_sync(&direct_timeout_plan(
format!("http://{addr}/chat"),
true,
ExecutionTimeouts {
first_byte_ms: Some(50),
total_ms: Some(5_000),
..ExecutionTimeouts::default()
},
))
.await;
server.abort();
let error = match result {
Ok(_) => panic!("stream sync body should hit first-byte timeout"),
Err(error) => error,
};
assert!(
error
.to_string()
.contains("provider stream first byte timeout after 50 ms"),
"unexpected error: {error}"
);
}
#[tokio::test]
async fn direct_sync_execution_runtime_does_not_apply_total_timeout_after_stream_body_starts() {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("client should connect");
let mut request = [0_u8; 1024];
let _ = socket
.read(&mut request)
.await
.expect("request should read");
socket
.write_all(
b"HTTP/1.1 200 OK\r\ncontent-type: text/event-stream\r\ntransfer-encoding: chunked\r\n\r\n",
)
.await
.expect("headers should write");
socket
.write_all(b"b\r\ndata: one\n\n\r\n")
.await
.expect("first chunk should write");
socket.flush().await.expect("first chunk should flush");
tokio::time::sleep(std::time::Duration::from_millis(120)).await;
socket
.write_all(b"b\r\ndata: two\n\n\r\n0\r\n\r\n")
.await
.expect("second chunk should write");
});
let result = DirectSyncExecutionRuntime::new()
.execute_sync(&direct_timeout_plan(
format!("http://{addr}/chat"),
true,
ExecutionTimeouts {
first_byte_ms: Some(50),
total_ms: Some(25),
..ExecutionTimeouts::default()
},
))
.await
.expect("stream body should not use total timeout after first chunk");
server.abort();
let body = result
.body
.and_then(|body| body.body_bytes_b64)
.and_then(|body| base64::engine::general_purpose::STANDARD.decode(body).ok())
.expect("stream body should be captured as bytes");
let body = String::from_utf8(body).expect("stream body should be utf8");
assert!(body.contains("data: one"));
assert!(body.contains("data: two"));
}
#[tokio::test]
async fn direct_stream_execution_runtime_applies_first_byte_timeout() {
let listener = crate::test_support::bind_loopback_listener()
@@ -2185,6 +2504,46 @@ mod tests {
assert_eq!(execution.status_code, http::StatusCode::OK.as_u16());
}
#[tokio::test]
async fn direct_stream_execution_runtime_ignores_total_timeout_when_first_byte_unset() {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let app = Router::new().route(
"/chat",
post(|| async {
tokio::time::sleep(std::time::Duration::from_millis(15)).await;
axum::response::Response::builder()
.status(http::StatusCode::OK)
.header("content-type", "text/event-stream")
.body(Body::from(Bytes::from_static(b"data: {}\n\n")))
.expect("response should build")
}),
);
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("test server should run");
});
let execution = DirectSyncExecutionRuntime::new()
.execute_stream(&direct_timeout_plan(
format!("http://{addr}/chat"),
true,
ExecutionTimeouts {
total_ms: Some(5),
..ExecutionTimeouts::default()
},
))
.await
.expect("stream should ignore total_ms and use the first-byte default");
server.abort();
assert_eq!(execution.status_code, http::StatusCode::OK.as_u16());
}
#[tokio::test]
async fn direct_sync_execution_runtime_routes_browser_wreq_transport_in_process() {
async fn browser_upstream(headers: AxumHeaderMap, body: Bytes) -> axum::response::Response {
@@ -26,7 +26,8 @@ use crate::request_candidate_runtime::{
};
use crate::{AppState, GatewayError};
const DEFAULT_STREAM_CANDIDATE_WATCHDOG_TIMEOUT_MS: u64 = 300_000;
const DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS: u64 = 30_000;
const DEFAULT_NON_STREAM_CANDIDATE_WATCHDOG_TIMEOUT_MS: u64 = 300_000;
fn attach_redaction_execution_candidate(response: &mut Response<Body>, candidate_id: Option<&str>) {
if let Some(candidate_id) = candidate_id
@@ -534,18 +535,18 @@ fn resolve_stream_candidate_watchdog_timeout(
.and_then(|context| context.get(UPSTREAM_IS_STREAM_KEY))
.and_then(serde_json::Value::as_bool)
.unwrap_or(true);
let timeout_ms = plan
.timeouts
.as_ref()
.and_then(|timeouts| {
if upstream_is_stream {
timeouts.first_byte_ms.or(timeouts.total_ms)
} else {
timeouts.total_ms.or(timeouts.first_byte_ms)
}
})
.unwrap_or(DEFAULT_STREAM_CANDIDATE_WATCHDOG_TIMEOUT_MS)
.max(1);
let timeout_ms = if upstream_is_stream {
plan.timeouts
.as_ref()
.and_then(|timeouts| timeouts.first_byte_ms)
.unwrap_or(DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS)
} else {
plan.timeouts
.as_ref()
.and_then(|timeouts| timeouts.total_ms)
.unwrap_or(DEFAULT_NON_STREAM_CANDIDATE_WATCHDOG_TIMEOUT_MS)
}
.max(1);
Duration::from_millis(timeout_ms)
}
@@ -783,7 +784,24 @@ mod tests {
assert_eq!(
timeout,
Duration::from_millis(DEFAULT_STREAM_CANDIDATE_WATCHDOG_TIMEOUT_MS)
Duration::from_millis(DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS)
);
}
#[test]
fn stream_candidate_watchdog_ignores_total_timeout_for_stream_upstream() {
let report_context = json!({"upstream_is_stream": true});
let timeout = resolve_stream_candidate_watchdog_timeout(
&test_plan(Some(ExecutionTimeouts {
total_ms: Some(90_000),
..ExecutionTimeouts::default()
})),
Some(&report_context),
);
assert_eq!(
timeout,
Duration::from_millis(DEFAULT_STREAM_FIRST_BYTE_WATCHDOG_TIMEOUT_MS)
);
}
@@ -803,17 +821,20 @@ mod tests {
}
#[test]
fn stream_candidate_watchdog_falls_back_to_first_byte_when_upstream_non_stream_lacks_total() {
fn stream_candidate_watchdog_uses_non_stream_default_when_upstream_non_stream_lacks_total() {
let report_context = json!({"upstream_is_stream": false});
let timeout = resolve_stream_candidate_watchdog_timeout(
&test_plan(Some(ExecutionTimeouts {
first_byte_ms: Some(300_000),
first_byte_ms: Some(12_345),
..ExecutionTimeouts::default()
})),
Some(&report_context),
);
assert_eq!(timeout, Duration::from_millis(300_000));
assert_eq!(
timeout,
Duration::from_millis(DEFAULT_NON_STREAM_CANDIDATE_WATCHDOG_TIMEOUT_MS)
);
}
#[test]
@@ -14,6 +14,8 @@ use super::{
provider_catalog_key_supports_format, query_param_value, AppState, GatewayPublicRequestContext,
};
const DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS: u64 = 300_000;
pub(super) async fn maybe_build_local_test_connection_route_response(
state: &AppState,
request_context: &GatewayPublicRequestContext,
@@ -287,12 +289,10 @@ pub(super) async fn maybe_build_local_test_connection_route_response(
for (name, value) in &provider_request_headers {
upstream_request = upstream_request.header(name, value);
}
if let Some(total_ms) =
crate::provider_transport::resolve_transport_execution_timeouts(&transport)
.and_then(|timeouts| timeouts.total_ms.or(timeouts.first_byte_ms))
{
upstream_request = upstream_request.timeout(Duration::from_millis(total_ms));
}
let total_ms = crate::provider_transport::resolve_transport_execution_timeouts(&transport)
.and_then(|timeouts| timeouts.total_ms)
.unwrap_or(DEFAULT_NON_STREAM_TOTAL_TIMEOUT_MS);
upstream_request = upstream_request.timeout(Duration::from_millis(total_ms));
let response = match upstream_request.json(&provider_request_body).send().await {
Ok(response) => response,
@@ -153,7 +153,6 @@ fn map_request_admission_error(error: super::RequestAdmissionError) -> String {
fn relay_header_timeout(meta: &protocol::RequestMeta) -> Duration {
let timeout_ms = if meta.stream {
meta.stream_first_byte_timeout_ms
.or(meta.request_timeout_ms)
.unwrap_or_else(|| meta.timeout.saturating_mul(1_000))
} else {
meta.request_timeout_ms
@@ -482,7 +481,10 @@ fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Respo
mod tests {
use super::super::hub::ProxyConn;
use super::super::{protocol, AppState, ConnConfig, ControlPlaneClient};
use super::{relay_request, Body, Request, SocketAddr, StatusCode, TUNNEL_ERROR_HEADER};
use super::{
relay_header_timeout, relay_request, Body, Request, SocketAddr, StatusCode,
TUNNEL_ERROR_HEADER,
};
use crate::data::GatewayDataState;
use crate::maintenance::start_proxy_upgrade_rollout;
use aether_contracts::tunnel::TUNNEL_RELAY_FORWARDED_BY_HEADER;
@@ -512,6 +514,27 @@ mod tests {
)
}
#[test]
fn relay_header_timeout_ignores_request_timeout_for_stream_requests() {
let meta = protocol::RequestMeta {
provider_id: None,
endpoint_id: None,
key_id: None,
method: "GET".to_string(),
url: "https://example.com/stream".to_string(),
headers: HashMap::new(),
stream: true,
request_timeout_ms: Some(90_000),
stream_first_byte_timeout_ms: None,
timeout: 7,
follow_redirects: None,
http1_only: false,
transport_profile: None,
};
assert_eq!(relay_header_timeout(&meta), Duration::from_secs(7));
}
fn sample_connected_proxy_node(node_id: &str) -> StoredProxyNode {
StoredProxyNode::new(
node_id.to_string(),