mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-07 01:47:47 +08:00
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:
@@ -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!({
|
||||
|
||||
@@ -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!(
|
||||
|
||||
Reference in New Issue
Block a user