Merge pull request #858 from stabey/codex/fix-sse-prefetch-handoff

fix(stream): preserve parser state across SSE prefetch handoff
This commit is contained in:
ZheFox
2026-09-29 10:06:04 +08:00
committed by GitHub
5 changed files with 429 additions and 171 deletions
@@ -1,3 +1,4 @@
use std::borrow::Cow;
use std::collections::BTreeMap;
use serde_json::{json, Map, Value};
@@ -226,7 +227,7 @@ enum AiSurfaceStreamRewriteState {
}
pub struct AiSurfaceStreamRewriter<'a> {
report_context: &'a Value,
report_context: Cow<'a, Value>,
buffered: Vec<u8>,
state: AiSurfaceStreamRewriteState,
}
@@ -271,30 +272,39 @@ pub fn maybe_build_ai_surface_stream_rewriter<'a>(
};
Some(AiSurfaceStreamRewriter {
report_context,
report_context: Cow::Borrowed(report_context),
buffered: Vec::new(),
state,
})
}
impl AiSurfaceStreamRewriter<'_> {
/// Move parser state across task boundaries without replaying captured bytes.
pub fn into_owned(self) -> AiSurfaceStreamRewriter<'static> {
AiSurfaceStreamRewriter {
report_context: Cow::Owned(self.report_context.into_owned()),
buffered: self.buffered,
state: self.state,
}
}
pub fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
match &mut self.state {
AiSurfaceStreamRewriteState::OpenAiImage(state) => {
state.push_chunk(self.report_context, chunk)
state.push_chunk(self.report_context.as_ref(), chunk)
}
AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(state) => {
state.push_chunk(self.report_context, chunk)
state.push_chunk(self.report_context.as_ref(), chunk)
}
AiSurfaceStreamRewriteState::ClaudeReadToolSanitize(state) => {
state.push_chunk(self.report_context, chunk)
state.push_chunk(self.report_context.as_ref(), chunk)
}
AiSurfaceStreamRewriteState::KiroToClaudeCli(state) => {
state.push_chunk(self.report_context, chunk)
state.push_chunk(self.report_context.as_ref(), chunk)
}
AiSurfaceStreamRewriteState::KiroToClaudeCliThenStandard { kiro, standard } => {
let claude_bytes = kiro.push_chunk(self.report_context, chunk)?;
transform_standard_bytes(standard, self.report_context, claude_bytes)
let claude_bytes = kiro.push_chunk(self.report_context.as_ref(), chunk)?;
transform_standard_bytes(standard, self.report_context.as_ref(), claude_bytes)
}
AiSurfaceStreamRewriteState::EnvelopeUnwrap
| AiSurfaceStreamRewriteState::ModelDirectiveDisplay
@@ -313,23 +323,25 @@ impl AiSurfaceStreamRewriter<'_> {
pub fn finish(&mut self) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
match &mut self.state {
AiSurfaceStreamRewriteState::OpenAiImage(state) => state.finish(self.report_context),
AiSurfaceStreamRewriteState::OpenAiImage(state) => {
state.finish(self.report_context.as_ref())
}
AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(state) => {
state.finish(self.report_context)
state.finish(self.report_context.as_ref())
}
AiSurfaceStreamRewriteState::ClaudeReadToolSanitize(state) => {
state.finish(self.report_context)
state.finish(self.report_context.as_ref())
}
AiSurfaceStreamRewriteState::KiroToClaudeCli(state) => {
state.finish(self.report_context)
state.finish(self.report_context.as_ref())
}
AiSurfaceStreamRewriteState::KiroToClaudeCliThenStandard { kiro, standard } => {
let mut output = transform_standard_bytes(
standard,
self.report_context,
kiro.finish(self.report_context)?,
self.report_context.as_ref(),
kiro.finish(self.report_context.as_ref())?,
)?;
output.extend(standard.finish(self.report_context)?);
output.extend(standard.finish(self.report_context.as_ref())?);
Ok(output)
}
AiSurfaceStreamRewriteState::EnvelopeUnwrap
@@ -338,14 +350,14 @@ impl AiSurfaceStreamRewriter<'_> {
| AiSurfaceStreamRewriteState::Standard(_) => {
if self.buffered.is_empty() {
if let AiSurfaceStreamRewriteState::Standard(state) = &mut self.state {
return state.finish(self.report_context);
return state.finish(self.report_context.as_ref());
}
return Ok(Vec::new());
}
let line = std::mem::take(&mut self.buffered);
let mut output = self.transform_line(line)?;
if let AiSurfaceStreamRewriteState::Standard(state) = &mut self.state {
output.extend(state.finish(self.report_context)?);
output.extend(state.finish(self.report_context.as_ref())?);
}
Ok(output)
}
@@ -365,18 +377,19 @@ impl AiSurfaceStreamRewriter<'_> {
fn transform_line(&mut self, line: Vec<u8>) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
match &mut self.state {
AiSurfaceStreamRewriteState::EnvelopeUnwrap => {
let output = transform_provider_private_stream_line(self.report_context, line)
.map_err(AiSurfaceFinalizeError::from)?;
rewrite_model_directive_stream_line(self.report_context, output)
let output =
transform_provider_private_stream_line(self.report_context.as_ref(), line)
.map_err(AiSurfaceFinalizeError::from)?;
rewrite_model_directive_stream_line(self.report_context.as_ref(), output)
}
AiSurfaceStreamRewriteState::ModelDirectiveDisplay => {
rewrite_model_directive_stream_line(self.report_context, line)
rewrite_model_directive_stream_line(self.report_context.as_ref(), line)
}
AiSurfaceStreamRewriteState::OpenAiResponsesCompat => {
rewrite_openai_responses_compat_stream_line(self.report_context, line)
rewrite_openai_responses_compat_stream_line(self.report_context.as_ref(), line)
}
AiSurfaceStreamRewriteState::Standard(state) => {
transform_standard_line(state, self.report_context, line)
transform_standard_line(state, self.report_context.as_ref(), line)
}
AiSurfaceStreamRewriteState::OpenAiImage(_)
| AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(_)
@@ -892,7 +905,7 @@ fn is_standard_cli_client_api_format(api_format: &str) -> bool {
#[cfg(test)]
mod tests {
use serde_json::json;
use serde_json::{json, Value};
use super::{
maybe_build_ai_surface_stream_rewriter, resolve_finalize_stream_rewrite_mode,
@@ -1067,6 +1080,50 @@ data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_123\",\"object\
assert!(!output.contains("\"model\":\"gpt-5.5\""));
}
#[test]
fn owned_handoff_preserves_partial_utf8_and_conversion_state() {
for client in ["openai:responses", "openai:chat"] {
let text = "界".repeat(12_000);
let delta = format!(
"data: {}\n\n",
json!({
"type":"response.output_text.delta", "response_id":"resp_handoff",
"item_id":"msg_handoff", "output_index":0, "content_index":0, "delta":text,
})
);
let split = delta.find('界').unwrap() + 17_002;
assert!(!delta.is_char_boundary(split));
let (mut owned, mut output) = {
let context = json!({"provider_api_format":"openai:responses",
"client_api_format":client, "needs_conversion":client == "openai:chat"});
let mut parser = maybe_build_ai_surface_stream_rewriter(Some(&context)).unwrap();
let output = parser.push_chunk(&delta.as_bytes()[..split]).unwrap();
(parser.into_owned(), output)
};
output.extend(owned.push_chunk(&delta.as_bytes()[split..]).unwrap());
output.extend(owned.finish().unwrap());
let output = String::from_utf8(output).unwrap();
let events: Vec<Value> = output
.lines()
.filter_map(|l| l.strip_prefix("data: "))
.filter(|p| *p != "[DONE]")
.map(|p| serde_json::from_str(p).unwrap())
.collect();
let recovered: String = events
.iter()
.filter_map(|e| {
if client == "openai:responses" {
e["delta"].as_str()
} else {
e.pointer("/choices/0/delta/content")
.and_then(Value::as_str)
}
})
.collect();
assert_eq!(recovered, text);
}
}
#[test]
fn standard_rewriter_converts_openai_responses_reasoning_delta_to_chat() {
let report_context = json!({
@@ -1,3 +1,4 @@
use std::borrow::Cow;
use std::collections::BTreeMap;
use serde_json::Value;
@@ -354,7 +355,7 @@ enum ProviderPrivateStreamNormalizeMode {
}
pub struct ProviderPrivateStreamNormalizer<'a> {
report_context: &'a Value,
report_context: Cow<'a, Value>,
buffered: Vec<u8>,
current_event_type: Option<String>,
mode: ProviderPrivateStreamNormalizeMode,
@@ -401,7 +402,7 @@ pub fn maybe_build_provider_private_stream_normalizer<'a>(
return None;
};
Some(ProviderPrivateStreamNormalizer {
report_context,
report_context: Cow::Borrowed(report_context),
buffered: Vec::new(),
current_event_type: None,
mode,
@@ -422,10 +423,20 @@ pub fn extract_provider_private_stream_error_body(
}
impl ProviderPrivateStreamNormalizer<'_> {
/// Move parser state across task boundaries without replaying captured bytes.
pub fn into_owned(self) -> ProviderPrivateStreamNormalizer<'static> {
ProviderPrivateStreamNormalizer {
report_context: Cow::Owned(self.report_context.into_owned()),
buffered: self.buffered,
current_event_type: self.current_event_type,
mode: self.mode,
}
}
pub fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
match &mut self.mode {
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
state.push_chunk(self.report_context, chunk)
state.push_chunk(self.report_context.as_ref(), chunk)
}
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
let next_len = self
@@ -441,7 +452,7 @@ impl ProviderPrivateStreamNormalizer<'_> {
)));
}
self.buffered.extend_from_slice(chunk);
if report_context_is_windsurf_envelope(self.report_context)
if report_context_is_windsurf_envelope(self.report_context.as_ref())
&& buffer_looks_like_connect_frame(&self.buffered)
{
return drain_windsurf_connect_json_frames(&mut self.buffered);
@@ -451,7 +462,7 @@ impl ProviderPrivateStreamNormalizer<'_> {
let line = self.buffered.drain(..=line_end).collect::<Vec<_>>();
output.extend(
transform_provider_private_stream_line_with_event_state(
self.report_context,
self.report_context.as_ref(),
line,
&mut self.current_event_type,
)
@@ -466,20 +477,20 @@ impl ProviderPrivateStreamNormalizer<'_> {
pub fn finish(&mut self) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
match &mut self.mode {
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
state.finish(self.report_context)
state.finish(self.report_context.as_ref())
}
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
if self.buffered.is_empty() {
return Ok(Vec::new());
}
if report_context_is_windsurf_envelope(self.report_context)
if report_context_is_windsurf_envelope(self.report_context.as_ref())
&& buffer_looks_like_connect_frame(&self.buffered)
{
return drain_windsurf_connect_json_frames(&mut self.buffered);
}
let line = std::mem::take(&mut self.buffered);
transform_provider_private_stream_line_with_event_state(
self.report_context,
self.report_context.as_ref(),
line,
&mut self.current_event_type,
)
@@ -939,7 +950,7 @@ fn postprocess_private_response_value(data: &mut Value, report_context: &Value)
#[cfg(test)]
mod tests {
use serde_json::json;
use serde_json::{json, Value};
use super::{
extract_provider_private_stream_error_body, maybe_build_provider_private_stream_normalizer,
@@ -1116,6 +1127,44 @@ mod tests {
assert!(text.contains(r#""content":"chunk""#));
}
#[test]
fn owned_handoff_preserves_private_binary_frame() {
let text = "frame".repeat(10_000);
let framed = connect_json_frame(
0,
&serde_json::to_vec(&json!({
"responseId":"ws-handoff", "response":{"text":text}
}))
.unwrap(),
);
let split = 17_735;
let mut normalizer = {
let context = json!({"has_envelope":true,
"envelope_name":"windsurf:GetChatMessage", "provider_api_format":"openai:chat"});
let mut normalizer =
maybe_build_provider_private_stream_normalizer(Some(&context)).unwrap();
assert!(normalizer.push_chunk(&framed[..split]).unwrap().is_empty());
normalizer.into_owned()
};
let mut output = normalizer.push_chunk(&framed[split..]).unwrap();
output.extend(normalizer.finish().unwrap());
let output = String::from_utf8(output).unwrap();
let events: Vec<Value> = output
.lines()
.filter_map(|l| l.strip_prefix("data: "))
.filter(|p| *p != "[DONE]")
.map(|p| serde_json::from_str(p).unwrap())
.collect();
let recovered: String = events
.iter()
.filter_map(|e| {
e.pointer("/choices/0/delta/content")
.and_then(Value::as_str)
})
.collect();
assert_eq!(recovered, text);
}
#[test]
fn unwraps_windsurf_connect_json_stream_frames() {
let report_context = json!({
+33 -2
View File
@@ -3052,14 +3052,26 @@ fn parse_sse_body_for_storage(text: &str) -> Option<Value> {
let mut chunks = Vec::new();
let mut total_chunks = 0_u64;
let mut saw_done = false;
let mut first_parse_error = None;
for_each_sse_payload(text, |payload| {
if payload == "[DONE]" {
saw_done = true;
return;
}
total_chunks += 1;
if let Ok(json_body) = serde_json::from_str::<Value>(payload) {
chunks.push(json_body);
match serde_json::from_str::<Value>(payload) {
Ok(json_body) => chunks.push(json_body),
Err(error) if first_parse_error.is_none() => {
// A later valid event must not hide an earlier broken one.
// Store diagnostics only, without duplicating raw user content.
first_parse_error = Some(json!({
"chunk_index": total_chunks - 1,
"line": error.line(),
"column": error.column(),
"message": error.to_string(),
}));
}
Err(_) => {}
}
});
if total_chunks == 0 && !saw_done {
@@ -3076,6 +3088,11 @@ fn parse_sse_body_for_storage(text: &str) -> Option<Value> {
if saw_done {
metadata.insert("has_completion".to_string(), Value::Bool(true));
}
if let Some(error) = first_parse_error {
// Capture truncation can also cause a parse error; this describes the
// captured payload, not an assertion that the provider sent bad JSON.
metadata.insert("first_parse_error".to_string(), error);
}
if stored_chunks < total_chunks {
metadata.insert(
"dropped_chunks".to_string(),
@@ -7104,6 +7121,20 @@ mod tests {
);
}
#[test]
fn parse_sse_body_for_storage_reports_bad_event_before_valid_terminal() {
let body = concat!(
"data: {\"tools\":[}\n\n",
"data: {\"type\":\"response.completed\"}\n\n",
);
let parsed = parse_sse_body_for_storage(body).unwrap();
assert_eq!(parsed["metadata"]["dropped_chunks"], 1);
assert_eq!(parsed["metadata"]["first_parse_error"]["chunk_index"], 0);
assert!(parsed["metadata"]["first_parse_error"]["message"].is_string());
assert_eq!(parsed["chunks"][0]["type"], "response.completed");
assert!(parsed.get("raw_response").is_none());
}
#[test]
fn extract_token_counts_from_value_handles_crlf_and_cr_sse_text() {
let sse_body = concat!(