修正Responses流式结算事件顺序

延迟透传上游 [DONE],确保扣费失败时客户端先收到 error 事件

补充 SSE done 标记边界解析回归测试
This commit is contained in:
2026-08-30 16:18:19 +08:00
parent 9f832ee96a
commit 7132ce543e
+87 -1
View File
@@ -461,6 +461,9 @@ fn stream_responses_with_billing(
async_stream::stream! {
let mut pending = String::new();
let mut usage = None;
let mut output_pending = Vec::new();
let mut deferred_done = Vec::new();
let mut done_seen = false;
while let Some(chunk) = upstream.next().await {
match chunk {
Ok(bytes) => {
@@ -473,7 +476,35 @@ fn stream_responses_with_billing(
usage = Some(event_usage);
}
}
yield Ok(bytes);
if done_seen {
deferred_done.extend_from_slice(bytes.as_ref());
continue;
}
output_pending.extend_from_slice(bytes.as_ref());
if let Some(done_start) = find_sse_done_marker(&output_pending) {
if let Some(done_end) = find_sse_event_end(&output_pending, done_start) {
let prefix = output_pending[..done_start].to_vec();
deferred_done.extend_from_slice(&output_pending[done_start..done_end]);
deferred_done.extend_from_slice(&output_pending[done_end..]);
output_pending.clear();
done_seen = true;
if !prefix.is_empty() {
yield Ok(Bytes::from(prefix));
}
continue;
}
}
let keep_len = sse_done_marker_suffix_len(&output_pending);
if output_pending.len() > keep_len {
let flush_len = output_pending.len() - keep_len;
let flushed = output_pending.drain(..flush_len).collect::<Vec<_>>();
if !flushed.is_empty() {
yield Ok(Bytes::from(flushed));
}
}
}
Err(error) => {
yield Err(std::io::Error::other(format!("AGC Router 响应流读取失败:{error}")));
@@ -485,6 +516,16 @@ fn stream_responses_with_billing(
if let Some(event_usage) = extract_llm_usage_from_sse_event(pending.trim()) {
usage = Some(event_usage);
}
if !done_seen {
if let Some(done_start) = find_sse_done_marker(&output_pending) {
deferred_done.extend_from_slice(&output_pending[done_start..]);
output_pending.truncate(done_start);
done_seen = true;
}
}
if !output_pending.is_empty() {
yield Ok(Bytes::from(std::mem::take(&mut output_pending)));
}
if let Err(error) = settle_agc_llm_usage(
&state,
owner_user_id.as_str(),
@@ -504,9 +545,42 @@ fn stream_responses_with_billing(
});
yield Ok(Bytes::from(format!("event: error\ndata: {payload}\n\n")));
}
if done_seen && !deferred_done.is_empty() {
yield Ok(Bytes::from(deferred_done));
}
}
}
const SSE_DONE_MARKER: &[u8] = b"data: [DONE]";
fn find_sse_done_marker(bytes: &[u8]) -> Option<usize> {
bytes
.windows(SSE_DONE_MARKER.len())
.enumerate()
.find_map(|(position, window)| {
(window == SSE_DONE_MARKER && (position == 0 || bytes[position - 1] == b'\n'))
.then_some(position)
})
}
fn find_sse_event_end(bytes: &[u8], event_start: usize) -> Option<usize> {
let event = &bytes[event_start..];
if let Some(offset) = event.windows(2).position(|window| window == b"\n\n") {
return Some(event_start + offset + 2);
}
event
.windows(4)
.position(|window| window == b"\r\n\r\n")
.map(|offset| event_start + offset + 4)
}
fn sse_done_marker_suffix_len(bytes: &[u8]) -> usize {
(1..SSE_DONE_MARKER.len())
.rev()
.find(|&length| bytes.ends_with(&SSE_DONE_MARKER[..length]))
.unwrap_or(0)
}
fn extract_llm_usage_from_sse_event(event: &str) -> Option<LlmTokenUsage> {
let mut data_lines = Vec::new();
for line in event.lines() {
@@ -1108,6 +1182,18 @@ mod tests {
assert!(extract_llm_usage_from_sse_event("event: response.output_text.delta").is_none());
}
#[test]
fn responses_sse_done_marker_is_found_only_at_event_line_start() {
let bytes = b"data: {\"text\":\"data: [DONE]\"}\n\ndata: [DONE]\n\n";
let done_start = find_sse_done_marker(bytes).expect("done event should be found");
assert_eq!(
&bytes[done_start..done_start + SSE_DONE_MARKER.len()],
SSE_DONE_MARKER
);
assert_eq!(find_sse_event_end(bytes, done_start), Some(bytes.len()));
assert_eq!(sse_done_marker_suffix_len(b"data: [DON"), 10);
}
#[test]
fn agc_llm_ledger_id_is_stable_and_does_not_expose_raw_key() {
let first = agc_llm_ledger_id("request-1");