diff --git a/server-rs/crates/api-server/src/llm/mod.rs b/server-rs/crates/api-server/src/llm/mod.rs index 7a51a3524..208a48854 100644 --- a/server-rs/crates/api-server/src/llm/mod.rs +++ b/server-rs/crates/api-server/src/llm/mod.rs @@ -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::>(); + 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 { + 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 { + 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 { 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");