fix: handle split streaming terminal events

This commit is contained in:
root
2026-05-25 20:58:17 +08:00
committed by hemo94931
parent d3249485fa
commit 431311979a
@@ -117,6 +117,7 @@ const OPENAI_IMAGE_STREAM_PLAN_KIND: &str = "openai_image_stream";
const SSE_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(15);
const SSE_KEEPALIVE_BYTES: &[u8] = b": aether-keepalive\n\n";
const SSE_CONTROL_FILTER_MAX_BUFFER_BYTES: usize = 1024 * 1024;
const SSE_TERMINAL_DETECTOR_MAX_LINE_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);
@@ -1571,42 +1572,120 @@ fn sse_buffer_has_data_line(buffer: &[u8]) -> bool {
.any(|line| line.trim_start().starts_with("data:"))
}
fn stream_chunk_contains_sse_done(chunk: &[u8]) -> bool {
std::str::from_utf8(chunk).ok().is_some_and(|text| {
text.lines().any(|line| {
let line = line.trim();
if matches!(
line,
"data: [DONE]"
| "event: message_stop"
| "event: response.completed"
| "event: response.failed"
| "event: response.incomplete"
| "event: error"
) {
return true;
#[derive(Default)]
struct ClientVisibleStreamCompletionTracker {
line_buffer: Vec<u8>,
event_type: Option<String>,
data_payload: String,
has_data_payload: bool,
skip_next_lf: bool,
completed: bool,
}
impl ClientVisibleStreamCompletionTracker {
fn observe_chunk(&mut self, chunk: &[u8]) -> bool {
if self.completed {
return true;
}
if chunk.is_empty() {
return false;
}
for byte in chunk {
if self.skip_next_lf {
self.skip_next_lf = false;
if *byte == b'\n' {
continue;
}
}
let Some(data) = line.strip_prefix("data:").map(str::trim) else {
return false;
};
data == "[DONE]"
|| serde_json::from_str::<serde_json::Value>(data).is_ok_and(|value| {
value
.get("type")
.and_then(serde_json::Value::as_str)
.is_some_and(|event_type| {
matches!(
event_type,
"message_stop"
| "response.completed"
| "response.failed"
| "response.incomplete"
| "error"
)
})
})
match *byte {
b'\n' => self.finish_line(),
b'\r' => {
self.finish_line();
self.skip_next_lf = true;
}
_ => {
self.line_buffer.push(*byte);
if self.line_buffer.len() > SSE_TERMINAL_DETECTOR_MAX_LINE_BYTES {
self.line_buffer.clear();
}
}
}
if self.completed {
break;
}
}
self.completed
}
fn finish_line(&mut self) {
let line = std::mem::take(&mut self.line_buffer);
let Ok(line) = std::str::from_utf8(&line) else {
self.reset_current_event();
return;
};
let line = line.trim();
if line.is_empty() {
self.completed = self.current_event_is_terminal();
self.reset_current_event();
return;
}
if let Some(event_type) = line.strip_prefix("event:").map(str::trim) {
self.event_type = Some(event_type.to_string());
return;
}
if let Some(data) = line.strip_prefix("data:").map(str::trim) {
if data.is_empty() {
return;
}
if self.has_data_payload {
self.data_payload.push('\n');
}
self.data_payload.push_str(data);
self.has_data_payload = true;
}
}
fn current_event_is_terminal(&self) -> bool {
self.event_type
.as_deref()
.is_some_and(is_terminal_sse_event_type)
|| (self.has_data_payload && sse_data_payload_is_terminal(&self.data_payload))
}
fn reset_current_event(&mut self) {
self.event_type = None;
self.data_payload.clear();
self.has_data_payload = false;
}
}
fn is_terminal_sse_event_type(event_type: &str) -> bool {
matches!(
event_type,
"message_stop" | "response.completed" | "response.failed" | "response.incomplete" | "error"
)
}
fn sse_data_payload_is_terminal(data: &str) -> bool {
data == "[DONE]"
|| serde_json::from_str::<serde_json::Value>(data).is_ok_and(|value| {
value
.get("type")
.and_then(serde_json::Value::as_str)
.is_some_and(is_terminal_sse_event_type)
})
})
}
fn stream_chunk_contains_sse_done(chunk: &[u8]) -> bool {
let mut tracker = ClientVisibleStreamCompletionTracker::default();
tracker.observe_chunk(chunk)
}
async fn next_stream_frame<R>(
@@ -2680,8 +2759,9 @@ async fn execute_stream_from_frame_stream(
max_stream_body_buffer_bytes,
&mut client_body_truncated,
);
let mut client_stream_completion_tracker = ClientVisibleStreamCompletionTracker::default();
let mut client_visible_stream_completed =
stream_chunk_contains_sse_done(&prefetched_body_for_report);
client_stream_completion_tracker.observe_chunk(&prefetched_body_for_report);
let mut usage_stream_telemetry: Option<ExecutionTelemetry> = initial_telemetry.clone();
let mut telemetry: Option<ExecutionTelemetry> = initial_telemetry;
let reached_eof = initial_reached_eof;
@@ -3065,12 +3145,11 @@ async fn execute_stream_from_frame_stream(
);
let rewritten_chunk_len =
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() {
let rewritten_chunk = Bytes::from(rewritten_chunk);
if tx.send(Ok(rewritten_chunk.clone())).await.is_err() {
warn!(
event_name = "stream_execution_downstream_disconnected",
log_type = "ops",
@@ -3081,7 +3160,8 @@ async fn execute_stream_from_frame_stream(
);
downstream_dropped = true;
} else {
client_visible_stream_completed |= chunk_completed_stream;
client_visible_stream_completed |= client_stream_completion_tracker
.observe_chunk(rewritten_chunk.as_ref());
client_stream_bytes.fetch_add(rewritten_chunk_len, Ordering::Relaxed);
last_client_chunk_elapsed_ms.store(
stream_started_at_for_report
@@ -3210,9 +3290,8 @@ async fn execute_stream_from_frame_stream(
);
let rewritten_chunk_len =
u64::try_from(rewritten_chunk.len()).unwrap_or(u64::MAX);
let chunk_completed_stream =
stream_chunk_contains_sse_done(&rewritten_chunk);
if tx.send(Ok(Bytes::from(rewritten_chunk))).await.is_err() {
let rewritten_chunk = Bytes::from(rewritten_chunk);
if tx.send(Ok(rewritten_chunk.clone())).await.is_err() {
warn!(
event_name = "stream_execution_downstream_flush_disconnected",
log_type = "ops",
@@ -3223,7 +3302,8 @@ async fn execute_stream_from_frame_stream(
);
downstream_dropped = true;
} else {
client_visible_stream_completed |= chunk_completed_stream;
client_visible_stream_completed |= client_stream_completion_tracker
.observe_chunk(rewritten_chunk.as_ref());
client_stream_bytes
.fetch_add(rewritten_chunk_len, Ordering::Relaxed);
last_client_chunk_elapsed_ms.store(
@@ -3281,8 +3361,8 @@ async fn execute_stream_from_frame_stream(
);
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() {
let flushed_chunk = Bytes::from(flushed_chunk);
if tx.send(Ok(flushed_chunk.clone())).await.is_err() {
warn!(
event_name = "stream_execution_downstream_rewrite_flush_disconnected",
log_type = "ops",
@@ -3293,7 +3373,8 @@ async fn execute_stream_from_frame_stream(
);
downstream_dropped = true;
} else {
client_visible_stream_completed |= chunk_completed_stream;
client_visible_stream_completed |= client_stream_completion_tracker
.observe_chunk(flushed_chunk.as_ref());
client_stream_bytes.fetch_add(flushed_chunk_len, Ordering::Relaxed);
last_client_chunk_elapsed_ms.store(
stream_started_at_for_report
@@ -3724,6 +3805,7 @@ mod tests {
stream_requires_observed_terminal_event, stream_terminal_summary_missing_observed_finish,
stream_terminal_summary_missing_observed_finish_with_requirement,
stream_terminal_summary_represents_failure_with_requirement,
ClientVisibleStreamCompletionTracker,
};
use crate::control::GatewayControlDecision;
use crate::tunnel::{tunnel_protocol, TunnelProxyConn};
@@ -3757,6 +3839,20 @@ mod tests {
));
}
#[test]
fn detects_client_visible_sse_terminal_events_across_chunks() {
let mut tracker = ClientVisibleStreamCompletionTracker::default();
assert!(!tracker.observe_chunk(b"data: [DO"));
assert!(!tracker.observe_chunk(b"NE]\n"));
assert!(tracker.observe_chunk(b"\n"));
let mut tracker = ClientVisibleStreamCompletionTracker::default();
assert!(!tracker.observe_chunk(b"event: response.comp"));
assert!(!tracker.observe_chunk(b"leted\r\n"));
assert!(tracker
.observe_chunk(b"data: {\"type\":\"response.completed\",\"response\":{}}\r\n\r\n"));
}
fn tunnel_proxy_snapshot(base_url: String) -> aether_contracts::ProxySnapshot {
aether_contracts::ProxySnapshot {
enabled: Some(true),
@@ -5238,6 +5334,155 @@ mod tests {
);
}
#[tokio::test]
async fn split_done_then_downstream_close_is_recorded_success() {
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
Arc::clone(&request_candidate_repository),
Arc::clone(&usage_repository),
),
)
.with_usage_runtime_for_tests(UsageRuntimeConfig {
enabled: true,
..UsageRuntimeConfig::default()
});
let plan = ExecutionPlan {
request_id: "req-split-done-close-success".into(),
candidate_id: Some("cand-split-done-close-success".into()),
provider_name: Some("openai".into()),
provider_id: "prov-1".into(),
endpoint_id: "ep-1".into(),
key_id: "key-1".into(),
method: "POST".into(),
url: "https://example.com/v1/chat/completions".into(),
headers: BTreeMap::from([
("content-type".into(), "application/json".into()),
("accept".into(), "text/event-stream".into()),
]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({
"model": "gpt-5.4",
"messages": [],
"stream": true
})),
stream: true,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
model_name: Some("gpt-5.4".into()),
proxy: None,
transport_profile: None,
timeouts: None,
};
let release_eof = Arc::new(Notify::new());
let release_eof_for_stream = Arc::clone(&release_eof);
let frame_stream = 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\\\",\\\"object\\\":\\\"chat.completion.chunk\\\",\\\"model\\\":\\\"gpt-5.4\\\",\\\"choices\\\":[{\\\"index\\\":0,\\\"delta\\\":{\\\"content\\\":\\\"hi\\\"},\\\"finish_reason\\\":null}]}\\n\\n\"}}\n",
));
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\\n\"}}\n",
));
yield Ok::<Bytes, std::io::Error>(Bytes::from_static(
b"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: [DO\"}}\n",
));
yield Ok::<Bytes, std::io::Error>(Bytes::from_static(
b"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"NE]\\n\\n\"}}\n",
));
release_eof_for_stream.notified().await;
}
.boxed();
let response = execute_stream_from_frame_stream(
&state,
plan,
"trace-split-done-close-success",
&test_decision(),
"openai_chat_stream",
None,
Some(json!({
"request_id": "req-split-done-close-success",
"candidate_id": "cand-split-done-close-success",
"candidate_index": 0,
"retry_index": 0,
"provider_api_format": "openai:chat",
"client_api_format": "openai:chat"
})),
crate::clock::current_unix_ms(),
Instant::now(),
frame_stream,
None,
)
.await
.expect("execution should succeed")
.expect("execution should return a client response");
let mut body_stream = response.into_body().into_data_stream();
let mut body = Vec::new();
tokio::time::timeout(Duration::from_secs(1), async {
while !String::from_utf8_lossy(&body).contains("data: [DONE]") {
let chunk = body_stream
.next()
.await
.expect("body should yield until done")
.expect("chunk should be ok");
body.extend_from_slice(&chunk);
}
})
.await
.expect("final DONE should arrive");
drop(body_stream);
release_eof.notify_one();
let candidates = tokio::time::timeout(Duration::from_secs(1), async {
loop {
let candidates = request_candidate_repository
.list_by_request_id("req-split-done-close-success")
.await
.expect("request candidates should read");
if candidates
.first()
.is_some_and(|candidate| candidate.status == RequestCandidateStatus::Success)
{
break candidates;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("candidate should be marked success");
assert_eq!(candidates[0].status_code, Some(200));
let stored_usage = tokio::time::timeout(Duration::from_secs(1), async {
loop {
let usage = usage_repository
.find_by_request_id("req-split-done-close-success")
.await
.expect("usage should read");
if usage
.as_ref()
.is_some_and(|usage| usage.status == "completed")
{
break usage.expect("completed usage should exist");
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("usage should be marked completed");
assert_eq!(stored_usage.status_code, Some(200));
assert_eq!(stored_usage.input_tokens, 7);
assert_eq!(stored_usage.output_tokens, 11);
assert_eq!(stored_usage.total_tokens, 18);
}
#[tokio::test]
async fn image_stream_downstream_close_after_done_is_recorded_success() {
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());