fix(gateway): drain downstream-disconnected streams and stop inferring cancelled usage

This commit is contained in:
zhefox
2026-05-21 20:30:44 +08:00
parent 71f9afc526
commit 77c2d91eb0
4 changed files with 186 additions and 466 deletions

View File

@@ -85,7 +85,7 @@ fn injects_stable_prompt_cache_key_for_codex_requests() {
assert_eq!(
body["prompt_cache_key"],
"172c39e6-c0a0-5a70-8b63-e0f8e0d185a3"
"53363264-dbb0-5f9d-b9c7-3e92c45c5bdf"
);
}

View File

@@ -209,7 +209,7 @@ fn local_openai_responses_compact_wrapper_strips_include_for_codex_requests() {
assert_eq!(provider_request_body["instructions"], "");
assert_eq!(
provider_request_body["prompt_cache_key"],
"172c39e6-c0a0-5a70-8b63-e0f8e0d185a3"
"3d2e2842-74cb-55dd-803a-b8940b3500c2"
);
}
@@ -355,7 +355,7 @@ fn injects_codex_prompt_cache_key_for_openai_responses_cross_format_requests() {
assert_eq!(
provider_request_body["prompt_cache_key"],
"172c39e6-c0a0-5a70-8b63-e0f8e0d185a3"
"b4dfeb75-b105-544c-a706-39b92f0bddb0"
);
}
@@ -385,6 +385,6 @@ fn injects_codex_prompt_cache_key_for_openai_chat_cross_format_requests() {
assert_eq!(
provider_request_body["prompt_cache_key"],
"172c39e6-c0a0-5a70-8b63-e0f8e0d185a3"
"4ee6ea6e-3ac6-5a18-8cb8-1f8b956419e5"
);
}

View File

@@ -1546,24 +1546,6 @@ where
read_next_frame(lines).await
}
async fn next_stream_frame_until_downstream_closed<R>(
buffered_frames: &mut VecDeque<StreamFrame>,
lines: &mut FramedRead<R, LinesCodec>,
tx: &mpsc::Sender<Result<Bytes, IoError>>,
) -> Result<Option<StreamFrame>, GatewayError>
where
R: tokio::io::AsyncRead + Unpin,
{
if let Some(frame) = buffered_frames.pop_front() {
return Ok(Some(frame));
}
tokio::select! {
frame = read_next_frame(lines) => frame,
() = tx.closed() => Ok(None),
}
}
fn should_refresh_stream_usage_telemetry(
previous: Option<&ExecutionTelemetry>,
next: &ExecutionTelemetry,
@@ -2764,15 +2746,7 @@ async fn execute_stream_from_frame_stream(
image_stream_total_timeout.as_mut()
{
tokio::select! {
result = next_stream_frame_until_downstream_closed(
&mut buffered_frames,
&mut lines,
&tx,
) => result,
() = tx.closed() => {
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);
@@ -2823,8 +2797,7 @@ async fn execute_stream_from_frame_stream(
}
}
} else {
next_stream_frame_until_downstream_closed(&mut buffered_frames, &mut lines, &tx)
.await
next_stream_frame(&mut buffered_frames, &mut lines).await
};
let next_frame = match next_frame_result {
Ok(frame) => frame,
@@ -3003,6 +2976,9 @@ async fn execute_stream_from_frame_stream(
u64::try_from(rewritten_chunk.len()).unwrap_or(u64::MAX);
let chunk_completed_stream =
stream_chunk_contains_sse_done(&rewritten_chunk);
if downstream_dropped {
continue;
}
if tx.send(Ok(Bytes::from(rewritten_chunk))).await.is_err() {
warn!(
event_name = "stream_execution_downstream_disconnected",
@@ -3010,10 +2986,9 @@ async fn execute_stream_from_frame_stream(
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
"gateway stream downstream dropped; stopping execution runtime stream forwarding"
"gateway stream downstream dropped; continuing to drain execution runtime stream"
);
downstream_dropped = true;
break;
} else {
client_visible_stream_completed |= chunk_completed_stream;
client_stream_bytes.fetch_add(rewritten_chunk_len, Ordering::Relaxed);
@@ -3070,30 +3045,30 @@ async fn execute_stream_from_frame_stream(
}
if downstream_dropped {
drop(lines);
debug!(
event_name = "execution_runtime_stream_flush_skipped",
event_name = "execution_runtime_stream_client_flush_skipped",
log_type = "debug",
debug_context = "redacted",
stream_status = "downstream_disconnected",
trace_id = %trace_id_owned,
"gateway skipped local stream flush after downstream disconnect"
"gateway skipped client stream flush after downstream disconnect"
);
} else {
if let Some(normalizer) = private_stream_normalizer.as_mut() {
match normalizer.finish() {
Ok(normalized_chunk) if !normalized_chunk.is_empty() => {
if let (Some(observer), Some(report_context)) = (
stream_usage_observer.as_mut(),
stream_usage_report_context.as_ref(),
) {
observe_stream_usage_bytes(
observer,
report_context,
&mut stream_usage_observer_buffered,
&normalized_chunk,
);
}
}
if let Some(normalizer) = private_stream_normalizer.as_mut() {
match normalizer.finish() {
Ok(normalized_chunk) if !normalized_chunk.is_empty() => {
if let (Some(observer), Some(report_context)) = (
stream_usage_observer.as_mut(),
stream_usage_report_context.as_ref(),
) {
observe_stream_usage_bytes(
observer,
report_context,
&mut stream_usage_observer_buffered,
&normalized_chunk,
);
}
if !downstream_dropped {
let rewritten_chunk = if let Some(rewriter) = local_stream_rewriter.as_mut()
{
match rewriter.push_chunk(&normalized_chunk) {
@@ -3157,86 +3132,85 @@ async fn execute_stream_from_frame_stream(
}
}
}
}
Ok(_) => {}
Err(err) => {
warn!(
event_name = "stream_execution_normalization_flush_failed",
log_type = "ops",
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
error = ?err,
"gateway failed to flush private stream normalization"
);
terminal_failure.get_or_insert_with(|| {
build_stream_failure_report(
"execution_runtime_stream_rewrite_flush_error",
format!("failed to flush private stream normalization: {err:?}"),
502,
)
});
}
}
}
if !downstream_dropped {
if let Some(rewriter) = local_stream_rewriter.as_mut() {
match rewriter.finish() {
Ok(flushed_chunk) if !flushed_chunk.is_empty() => {
append_stream_capture_bytes(
&mut buffered_body,
&flushed_chunk,
max_stream_body_buffer_bytes,
&mut client_body_truncated,
);
let flushed_chunk_len =
u64::try_from(flushed_chunk.len()).unwrap_or(u64::MAX);
let chunk_completed_stream = stream_chunk_contains_sse_done(&flushed_chunk);
if tx.send(Ok(Bytes::from(flushed_chunk))).await.is_err() {
warn!(
event_name = "stream_execution_downstream_rewrite_flush_disconnected",
log_type = "ops",
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
"gateway stream downstream dropped while flushing local stream rewrite"
);
downstream_dropped = true;
} else {
client_visible_stream_completed |= chunk_completed_stream;
client_stream_bytes.fetch_add(flushed_chunk_len, Ordering::Relaxed);
last_client_chunk_elapsed_ms.store(
stream_started_at_for_report
.elapsed()
.as_millis()
.min(u128::from(u64::MAX))
as u64,
Ordering::Relaxed,
);
}
}
Ok(_) => {}
Err(err) => {
warn!(
event_name = "stream_execution_normalization_flush_failed",
event_name = "stream_execution_rewrite_flush_failed",
log_type = "ops",
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
error = ?err,
"gateway failed to flush private stream normalization"
"gateway failed to flush local stream rewrite"
);
terminal_failure.get_or_insert_with(|| {
build_stream_failure_report(
"execution_runtime_stream_rewrite_flush_error",
format!("failed to flush private stream normalization: {err:?}"),
format!("failed to flush local stream rewrite: {err:?}"),
502,
)
});
}
}
}
if !downstream_dropped {
if let Some(rewriter) = local_stream_rewriter.as_mut() {
match rewriter.finish() {
Ok(flushed_chunk) if !flushed_chunk.is_empty() => {
append_stream_capture_bytes(
&mut buffered_body,
&flushed_chunk,
max_stream_body_buffer_bytes,
&mut client_body_truncated,
);
let flushed_chunk_len =
u64::try_from(flushed_chunk.len()).unwrap_or(u64::MAX);
let chunk_completed_stream =
stream_chunk_contains_sse_done(&flushed_chunk);
if tx.send(Ok(Bytes::from(flushed_chunk))).await.is_err() {
warn!(
event_name = "stream_execution_downstream_rewrite_flush_disconnected",
log_type = "ops",
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
"gateway stream downstream dropped while flushing local stream rewrite"
);
downstream_dropped = true;
} else {
client_visible_stream_completed |= chunk_completed_stream;
client_stream_bytes.fetch_add(flushed_chunk_len, Ordering::Relaxed);
last_client_chunk_elapsed_ms.store(
stream_started_at_for_report
.elapsed()
.as_millis()
.min(u128::from(u64::MAX))
as u64,
Ordering::Relaxed,
);
}
}
Ok(_) => {}
Err(err) => {
warn!(
event_name = "stream_execution_rewrite_flush_failed",
log_type = "ops",
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
error = ?err,
"gateway failed to flush local stream rewrite"
);
terminal_failure.get_or_insert_with(|| {
build_stream_failure_report(
"execution_runtime_stream_rewrite_flush_error",
format!("failed to flush local stream rewrite: {err:?}"),
502,
)
});
}
}
}
}
}
if !downstream_dropped {
@@ -3527,7 +3501,7 @@ async fn execute_stream_from_frame_stream(
None
},
error_message: stream_failed
.then(|| stream_terminal_error_message)
.then_some(stream_terminal_error_message)
.flatten(),
latency_ms: usage_payload
.telemetry
@@ -4655,7 +4629,7 @@ mod tests {
}
#[tokio::test]
async fn execute_stream_from_frame_stream_stops_upstream_when_client_drops_body() {
async fn execute_stream_from_frame_stream_drains_upstream_when_client_drops_body() {
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let state = AppState::new()
@@ -4698,24 +4672,22 @@ mod tests {
transport_profile: None,
timeouts: None,
};
let frame_stream_dropped = Arc::new(Notify::new());
let frame_stream_dropped_for_stream = Arc::clone(&frame_stream_dropped);
let release_terminal = Arc::new(Notify::new());
let terminal_frame_drained = Arc::new(Notify::new());
let release_terminal_for_stream = Arc::clone(&release_terminal);
let terminal_frame_drained_for_stream = Arc::clone(&terminal_frame_drained);
let frame_stream = stream! {
struct NotifyOnDrop(Arc<Notify>);
impl Drop for NotifyOnDrop {
fn drop(&mut self) {
self.0.notify_waiters();
}
}
let _drop_guard = NotifyOnDrop(frame_stream_dropped_for_stream);
yield Ok::<Bytes, std::io::Error>(Bytes::from_static(
b"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
));
yield Ok::<Bytes, std::io::Error>(Bytes::from_static(
b"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"id\\\":\\\"first\\\"}\\n\\n\"}}\n",
));
std::future::pending::<()>().await;
release_terminal_for_stream.notified().await;
yield Ok::<Bytes, std::io::Error>(Bytes::from_static(
b"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"id\\\":\\\"terminal\\\",\\\"object\\\":\\\"chat.completion.chunk\\\",\\\"model\\\":\\\"gpt-5.4\\\",\\\"choices\\\":[{\\\"index\\\":0,\\\"delta\\\":{},\\\"finish_reason\\\":\\\"stop\\\"}],\\\"usage\\\":{\\\"prompt_tokens\\\":7,\\\"completion_tokens\\\":11,\\\"total_tokens\\\":18}}\\n\\ndata: [DONE]\\n\\n\"}}\n",
));
terminal_frame_drained_for_stream.notify_one();
}
.boxed();
@@ -4761,10 +4733,11 @@ mod tests {
assert_eq!(first.as_ref(), b"data: {\"id\":\"first\"}\n\n");
tokio::time::sleep(Duration::from_millis(30)).await;
drop(body_stream);
release_terminal.notify_one();
tokio::time::timeout(Duration::from_secs(1), frame_stream_dropped.notified())
tokio::time::timeout(Duration::from_secs(1), terminal_frame_drained.notified())
.await
.expect("upstream frame stream should be dropped after client disconnect");
.expect("upstream frame stream should be drained after client disconnect");
let candidates = tokio::time::timeout(Duration::from_secs(1), async {
loop {
let candidates = request_candidate_repository
@@ -4807,6 +4780,9 @@ mod tests {
.expect("usage should be marked cancelled");
assert_eq!(stored_usage.billing_status, "pending");
assert_eq!(stored_usage.status_code, Some(499));
assert_eq!(stored_usage.input_tokens, 7);
assert_eq!(stored_usage.output_tokens, 11);
assert_eq!(stored_usage.total_tokens, 18);
let first_byte_time_ms = stored_usage
.first_byte_time_ms
.expect("cancelled stream should retain first byte time");