mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 08:27:46 +08:00
fix(stream): preserve parser state across SSE prefetch handoff
This commit is contained in:
@@ -18,6 +18,12 @@ pub(crate) fn maybe_build_local_stream_rewriter<'a>(
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl LocalStreamRewriter<'_> {
|
impl LocalStreamRewriter<'_> {
|
||||||
|
pub(crate) fn into_owned(self) -> LocalStreamRewriter<'static> {
|
||||||
|
LocalStreamRewriter {
|
||||||
|
inner: self.inner.into_owned(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, GatewayError> {
|
pub(crate) fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, GatewayError> {
|
||||||
self.inner.push_chunk(chunk).map_err(map_surface_error)
|
self.inner.push_chunk(chunk).map_err(map_surface_error)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -6462,6 +6462,21 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
|||||||
|
|
||||||
let normalized_stream_report_context =
|
let normalized_stream_report_context =
|
||||||
normalize_provider_private_report_context(report_context.as_ref());
|
normalize_provider_private_report_context(report_context.as_ref());
|
||||||
|
// Observers follow the live protocol stream across prefetch and transfer.
|
||||||
|
// Diagnostic capture limits must never determine parser state.
|
||||||
|
let stream_usage_report_context = normalized_stream_report_context.clone().or_else(|| {
|
||||||
|
Some(json!({
|
||||||
|
"provider_api_format": plan.provider_api_format.as_str(),
|
||||||
|
"client_api_format": plan.client_api_format.as_str(),
|
||||||
|
}))
|
||||||
|
});
|
||||||
|
let mut stream_usage_observer = stream_usage_report_context
|
||||||
|
.as_ref()
|
||||||
|
.map(|_| StreamingStandardTerminalObserver::default());
|
||||||
|
let mut stream_usage_observer_buffered =
|
||||||
|
StreamUsageObservationBuffer::new(max_stream_body_buffer_bytes);
|
||||||
|
let mut provider_error_inspection = ProviderStreamErrorInspection::default();
|
||||||
|
let mut prefetched_provider_error = None;
|
||||||
let upstream_headers = headers.clone();
|
let upstream_headers = headers.clone();
|
||||||
let mut private_stream_normalizer =
|
let mut private_stream_normalizer =
|
||||||
maybe_build_provider_private_stream_normalizer(report_context.as_ref());
|
maybe_build_provider_private_stream_normalizer(report_context.as_ref());
|
||||||
@@ -6562,7 +6577,8 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
|||||||
stream_commit_gate.commit();
|
stream_commit_gate.commit();
|
||||||
}
|
}
|
||||||
let mut prefetched_chunks: Vec<Bytes> = Vec::new();
|
let mut prefetched_chunks: Vec<Bytes> = Vec::new();
|
||||||
let mut provider_prefetched_body = Vec::new();
|
let mut provider_prefetched_body = StreamBodyCapture::default();
|
||||||
|
let mut provider_prefetched_bytes = 0_u64;
|
||||||
let mut provider_prefetched_body_truncated = false;
|
let mut provider_prefetched_body_truncated = false;
|
||||||
let mut prefetched_body = Vec::new();
|
let mut prefetched_body = Vec::new();
|
||||||
let mut prefetched_inspection_body = Vec::new();
|
let mut prefetched_inspection_body = Vec::new();
|
||||||
@@ -6815,10 +6831,12 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
append_stream_capture_bytes(
|
provider_prefetched_bytes =
|
||||||
|
provider_prefetched_bytes.saturating_add(chunk.len() as u64);
|
||||||
|
append_budgeted_stream_capture_bytes(
|
||||||
&mut provider_prefetched_body,
|
&mut provider_prefetched_body,
|
||||||
&chunk,
|
&chunk,
|
||||||
MAX_STREAM_PREFETCH_BYTES,
|
max_stream_body_buffer_bytes,
|
||||||
&mut provider_prefetched_body_truncated,
|
&mut provider_prefetched_body_truncated,
|
||||||
);
|
);
|
||||||
append_stream_capture_bytes(
|
append_stream_capture_bytes(
|
||||||
@@ -7063,6 +7081,22 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
|||||||
} else {
|
} else {
|
||||||
chunk
|
chunk
|
||||||
};
|
};
|
||||||
|
if let Some(error) = provider_error_inspection
|
||||||
|
.observe(stream_usage_report_context.as_ref(), &normalized_chunk)
|
||||||
|
{
|
||||||
|
prefetched_provider_error.get_or_insert(error);
|
||||||
|
}
|
||||||
|
if let (Some(observer), Some(context)) = (
|
||||||
|
stream_usage_observer.as_mut(),
|
||||||
|
stream_usage_report_context.as_ref(),
|
||||||
|
) {
|
||||||
|
observe_stream_usage_bytes(
|
||||||
|
observer,
|
||||||
|
context,
|
||||||
|
&mut stream_usage_observer_buffered,
|
||||||
|
&normalized_chunk,
|
||||||
|
);
|
||||||
|
}
|
||||||
let rewritten_chunk = if let Some(rewriter) = local_stream_rewriter.as_mut() {
|
let rewritten_chunk = if let Some(rewriter) = local_stream_rewriter.as_mut() {
|
||||||
match rewriter.push_chunk(&normalized_chunk) {
|
match rewriter.push_chunk(&normalized_chunk) {
|
||||||
Ok(rewritten_chunk) => rewritten_chunk,
|
Ok(rewritten_chunk) => rewritten_chunk,
|
||||||
@@ -7193,17 +7227,21 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
|||||||
if stream_commit_gate.is_uncommitted() {
|
if stream_commit_gate.is_uncommitted() {
|
||||||
stream_commit_gate.commit();
|
stream_commit_gate.commit();
|
||||||
}
|
}
|
||||||
let prefetched_response_history_persisted = if let Some(record) = local_stream_rewriter
|
if let Some(record) = local_stream_rewriter
|
||||||
.as_mut()
|
.as_mut()
|
||||||
.and_then(|rewriter| rewriter.take_response_history_record())
|
.and_then(|rewriter| rewriter.take_response_history_record())
|
||||||
{
|
{
|
||||||
crate::ai_serving::persist_response_history_record(state, record).await;
|
crate::ai_serving::persist_response_history_record(state, record).await;
|
||||||
true
|
}
|
||||||
} else {
|
// Keep partial records and conversion state; replaying the bounded
|
||||||
false
|
// inspection/capture prefix loses any bytes consumed beyond that prefix.
|
||||||
};
|
let mut private_stream_normalizer = private_stream_normalizer.map(|parser| parser.into_owned());
|
||||||
drop(private_stream_normalizer);
|
let mut local_stream_rewriter = local_stream_rewriter.map(|parser| parser.into_owned());
|
||||||
drop(local_stream_rewriter);
|
if sync_json_stream_bridge_active {
|
||||||
|
private_stream_normalizer = None;
|
||||||
|
local_stream_rewriter = None;
|
||||||
|
stream_usage_observer = None;
|
||||||
|
}
|
||||||
|
|
||||||
let initial_usage_telemetry = prefetched_usage_telemetry.clone().or_else(|| {
|
let initial_usage_telemetry = prefetched_usage_telemetry.clone().or_else(|| {
|
||||||
prefetched_telemetry
|
prefetched_telemetry
|
||||||
@@ -7246,7 +7284,6 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
|||||||
let headers_for_report = headers.clone();
|
let headers_for_report = headers.clone();
|
||||||
let report_kind_owned = report_kind;
|
let report_kind_owned = report_kind;
|
||||||
let report_context_owned = report_context;
|
let report_context_owned = report_context;
|
||||||
let normalized_stream_report_context_owned = normalized_stream_report_context;
|
|
||||||
let lifecycle_seed_for_report = lifecycle_seed;
|
let lifecycle_seed_for_report = lifecycle_seed;
|
||||||
let provider_prefetched_body_for_report = provider_prefetched_body;
|
let provider_prefetched_body_for_report = provider_prefetched_body;
|
||||||
let prefetched_body_for_report = prefetched_body;
|
let prefetched_body_for_report = prefetched_body;
|
||||||
@@ -7288,40 +7325,10 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
|||||||
let _stream_total_guard =
|
let _stream_total_guard =
|
||||||
StageElapsedGuard::from_started_at("stream_total", stream_started_at_for_report);
|
StageElapsedGuard::from_started_at("stream_total", stream_started_at_for_report);
|
||||||
let _provider_pool_in_flight_guard = provider_pool_in_flight_guard_for_report;
|
let _provider_pool_in_flight_guard = provider_pool_in_flight_guard_for_report;
|
||||||
let mut provider_buffered_body = StreamBodyCapture::default();
|
let mut provider_buffered_body = provider_prefetched_body_for_report;
|
||||||
let mut buffered_body = StreamBodyCapture::default();
|
let mut buffered_body = StreamBodyCapture::default();
|
||||||
let mut provider_body_truncated = false;
|
let mut provider_body_truncated = provider_prefetched_body_truncated;
|
||||||
let mut client_body_truncated = false;
|
let mut client_body_truncated = false;
|
||||||
let mut private_stream_normalizer = if sync_json_stream_bridge_active_for_report {
|
|
||||||
None
|
|
||||||
} else {
|
|
||||||
maybe_build_provider_private_stream_normalizer(report_context_owned.as_ref())
|
|
||||||
};
|
|
||||||
let mut local_stream_rewriter = if sync_json_stream_bridge_active_for_report {
|
|
||||||
None
|
|
||||||
} else {
|
|
||||||
maybe_build_stream_response_rewriter(normalized_stream_report_context_owned.as_ref())
|
|
||||||
};
|
|
||||||
let stream_usage_report_context =
|
|
||||||
normalized_stream_report_context_owned.clone().or_else(|| {
|
|
||||||
Some(serde_json::json!({
|
|
||||||
"provider_api_format": plan_for_report.provider_api_format.as_str(),
|
|
||||||
"client_api_format": plan_for_report.client_api_format.as_str(),
|
|
||||||
}))
|
|
||||||
});
|
|
||||||
let mut stream_usage_observer = stream_usage_report_context
|
|
||||||
.as_ref()
|
|
||||||
.filter(|_| !sync_json_stream_bridge_active_for_report)
|
|
||||||
.map(|_| StreamingStandardTerminalObserver::default());
|
|
||||||
let mut stream_usage_observer_buffered =
|
|
||||||
StreamUsageObservationBuffer::new(max_stream_body_buffer_bytes);
|
|
||||||
let mut provider_error_inspection = ProviderStreamErrorInspection::default();
|
|
||||||
append_budgeted_stream_capture_bytes(
|
|
||||||
&mut provider_buffered_body,
|
|
||||||
&provider_prefetched_body_for_report,
|
|
||||||
max_stream_body_buffer_bytes,
|
|
||||||
&mut provider_body_truncated,
|
|
||||||
);
|
|
||||||
append_budgeted_stream_capture_bytes(
|
append_budgeted_stream_capture_bytes(
|
||||||
&mut buffered_body,
|
&mut buffered_body,
|
||||||
&prefetched_body_for_report,
|
&prefetched_body_for_report,
|
||||||
@@ -7365,9 +7372,7 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
|||||||
} else {
|
} else {
|
||||||
initial_elapsed_ms
|
initial_elapsed_ms
|
||||||
}));
|
}));
|
||||||
let provider_stream_bytes = Arc::new(AtomicU64::new(
|
let provider_stream_bytes = Arc::new(AtomicU64::new(provider_prefetched_bytes));
|
||||||
u64::try_from(provider_prefetched_body_for_report.len()).unwrap_or(u64::MAX),
|
|
||||||
));
|
|
||||||
let client_stream_bytes = Arc::new(AtomicU64::new(
|
let client_stream_bytes = Arc::new(AtomicU64::new(
|
||||||
u64::try_from(prefetched_body_for_report.len()).unwrap_or(u64::MAX),
|
u64::try_from(prefetched_body_for_report.len()).unwrap_or(u64::MAX),
|
||||||
));
|
));
|
||||||
@@ -7463,96 +7468,20 @@ async fn execute_stream_from_frame_stream_with_retry_scope(
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
};
|
};
|
||||||
if !provider_prefetched_body_for_report.is_empty() {
|
if let Some(error_body_json) = prefetched_provider_error {
|
||||||
let normalized_prefetched_chunk = if let Some(normalizer) =
|
provider_error_forwarded_to_client = !prefetched_body_for_report.is_empty();
|
||||||
private_stream_normalizer.as_mut()
|
let error_status_code = resolve_provider_stream_error_status_code(
|
||||||
{
|
plan_for_report.provider_api_format.as_str(),
|
||||||
match normalizer.push_chunk(&provider_prefetched_body_for_report) {
|
status_code,
|
||||||
Ok(normalized_chunk) => Some(normalized_chunk),
|
&error_body_json,
|
||||||
Err(err) => {
|
);
|
||||||
warn!(
|
terminal_failure = Some(build_stream_failure_from_provider_error_body(
|
||||||
event_name = "stream_execution_prefetch_normalize_restore_failed",
|
error_status_code,
|
||||||
log_type = "ops",
|
&error_body_json,
|
||||||
trace_id = %trace_id_owned,
|
));
|
||||||
request_id = %request_id_for_report_log,
|
|
||||||
candidate_id = ?candidate_id_for_report.as_deref(),
|
|
||||||
error_category = "stream_normalization_restore_failed",
|
|
||||||
"gateway failed to restore private stream normalization state after prefetch"
|
|
||||||
);
|
|
||||||
terminal_failure = Some(build_stream_failure_report(
|
|
||||||
"execution_runtime_stream_rewrite_error",
|
|
||||||
format!(
|
|
||||||
"failed to restore private stream normalization state after prefetch: {err:?}"
|
|
||||||
),
|
|
||||||
502,
|
|
||||||
));
|
|
||||||
None
|
|
||||||
}
|
|
||||||
}
|
|
||||||
} else {
|
|
||||||
None
|
|
||||||
};
|
|
||||||
let replay_chunk = normalized_prefetched_chunk
|
|
||||||
.as_deref()
|
|
||||||
.unwrap_or(provider_prefetched_body_for_report.as_slice());
|
|
||||||
if let Some(error_body_json) = provider_error_inspection
|
|
||||||
.observe(stream_usage_report_context.as_ref(), replay_chunk)
|
|
||||||
{
|
|
||||||
provider_error_forwarded_to_client = !prefetched_body_for_report.is_empty();
|
|
||||||
let error_status_code = resolve_provider_stream_error_status_code(
|
|
||||||
plan_for_report.provider_api_format.as_str(),
|
|
||||||
status_code,
|
|
||||||
&error_body_json,
|
|
||||||
);
|
|
||||||
terminal_failure = Some(build_stream_failure_from_provider_error_body(
|
|
||||||
error_status_code,
|
|
||||||
&error_body_json,
|
|
||||||
));
|
|
||||||
}
|
|
||||||
if let (Some(observer), Some(report_context)) = (
|
|
||||||
stream_usage_observer.as_mut(),
|
|
||||||
stream_usage_report_context.as_ref(),
|
|
||||||
) {
|
|
||||||
observe_stream_usage_bytes(
|
|
||||||
observer,
|
|
||||||
report_context,
|
|
||||||
&mut stream_usage_observer_buffered,
|
|
||||||
replay_chunk,
|
|
||||||
);
|
|
||||||
}
|
|
||||||
if terminal_failure.is_none() {
|
|
||||||
if let Some(rewriter) = local_stream_rewriter.as_mut() {
|
|
||||||
if let Err(err) = rewriter.push_chunk(replay_chunk) {
|
|
||||||
warn!(
|
|
||||||
event_name = "stream_execution_prefetch_rewrite_restore_failed",
|
|
||||||
log_type = "ops",
|
|
||||||
trace_id = %trace_id_owned,
|
|
||||||
request_id = %request_id_for_report_log,
|
|
||||||
candidate_id = ?candidate_id_for_report.as_deref(),
|
|
||||||
error_category = "stream_rewrite_restore_failed",
|
|
||||||
"gateway failed to restore local stream rewrite state after prefetch"
|
|
||||||
);
|
|
||||||
terminal_failure = Some(build_stream_failure_report(
|
|
||||||
"execution_runtime_stream_rewrite_error",
|
|
||||||
format!(
|
|
||||||
"failed to restore local stream rewrite state after prefetch: {err:?}"
|
|
||||||
),
|
|
||||||
502,
|
|
||||||
));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if prefetched_response_history_persisted {
|
|
||||||
if let Some(rewriter) = local_stream_rewriter.as_mut() {
|
|
||||||
let _ = rewriter.take_response_history_record();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
// Parser state is already current and capture owns its budgeted bytes.
|
||||||
// These buffers restore parser/rewriter state above. Audit capture owns
|
// This output prefix is needed only to initialize client-side trackers.
|
||||||
// its budgeted copies; retaining semantic prefetch duplicates for the
|
|
||||||
// rest of the stream would bypass the capture memory limit.
|
|
||||||
drop(provider_prefetched_body_for_report);
|
|
||||||
drop(prefetched_body_for_report);
|
drop(prefetched_body_for_report);
|
||||||
|
|
||||||
if terminal_failure.is_none() && !reached_eof {
|
if terminal_failure.is_none() && !reached_eof {
|
||||||
@@ -9312,6 +9241,188 @@ mod tests {
|
|||||||
.unwrap()
|
.unwrap()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn prefetch_handoff_preserves_large_responses_setup_event() {
|
||||||
|
let event = format!(
|
||||||
|
"event: response.created\ndata: {}\n\n",
|
||||||
|
json!({"type":"response.created", "response": {
|
||||||
|
"id":"resp-large-setup", "status":"in_progress", "output":[],
|
||||||
|
"tools":[{"name":"write", "description":"x".repeat(64 * 1024)}]
|
||||||
|
}})
|
||||||
|
);
|
||||||
|
let done = "event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp-large-setup\",\"status\":\"completed\",\"output\":[],\"usage\":{\"input_tokens\":7,\"output_tokens\":2}}}\n\n";
|
||||||
|
// Include the two observed transport boundaries, exact/near budget
|
||||||
|
// boundaries, and multiple prefetch chunks crossing the budget.
|
||||||
|
for cuts in [
|
||||||
|
vec![16_383],
|
||||||
|
vec![16_384],
|
||||||
|
vec![17_735],
|
||||||
|
vec![17_741],
|
||||||
|
vec![8_192, 17_735],
|
||||||
|
] {
|
||||||
|
let mut chunks = Vec::new();
|
||||||
|
let mut start = 0;
|
||||||
|
for end in cuts {
|
||||||
|
chunks.push(&event[start..end]);
|
||||||
|
start = end;
|
||||||
|
}
|
||||||
|
chunks.push(&event[start..]);
|
||||||
|
chunks.push(done);
|
||||||
|
let response = execute_generic_sse_precommit(chunks, json!({}), None, false)
|
||||||
|
.await
|
||||||
|
.expect("large setup should commit at the bounded prefetch limit");
|
||||||
|
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||||
|
let body = String::from_utf8(body.to_vec()).unwrap();
|
||||||
|
assert!(
|
||||||
|
body.starts_with(&event),
|
||||||
|
"setup bytes lost or duplicated at split {start}"
|
||||||
|
);
|
||||||
|
let events: Vec<Value> = body
|
||||||
|
.lines()
|
||||||
|
.filter_map(|line| line.strip_prefix("data: "))
|
||||||
|
.filter(|payload| *payload != "[DONE]")
|
||||||
|
.map(|payload| {
|
||||||
|
serde_json::from_str(payload).expect("every SSE payload must be valid JSON")
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
assert_eq!(events.len(), 2, "events must be forwarded exactly once");
|
||||||
|
assert_eq!(events[1]["type"], "response.completed");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn prefetch_handoff_keeps_audit_usage_and_private_conversion() {
|
||||||
|
for private in [false, true] {
|
||||||
|
let request_id = format!("handoff-audit-{}", uuid::Uuid::new_v4());
|
||||||
|
let mut plan = if private {
|
||||||
|
antigravity_gemini_stream_plan(&request_id)
|
||||||
|
} else {
|
||||||
|
native_anthropic_stream_plan(&request_id)
|
||||||
|
};
|
||||||
|
if !private {
|
||||||
|
plan.provider_api_format = "openai:responses".into();
|
||||||
|
plan.client_api_format = "openai:responses".into();
|
||||||
|
}
|
||||||
|
let context = json!({
|
||||||
|
"request_id": request_id, "candidate_id": plan.candidate_id,
|
||||||
|
"candidate_index":0, "retry_index":0,
|
||||||
|
"provider_api_format": plan.provider_api_format,
|
||||||
|
"client_api_format": plan.client_api_format,
|
||||||
|
"needs_conversion": private, "has_envelope": private,
|
||||||
|
"envelope_name": if private { "antigravity:v1internal" } else { "" },
|
||||||
|
});
|
||||||
|
let repository = Arc::new(InMemoryUsageReadRepository::default());
|
||||||
|
let catalog = provider_catalog_for_plan(&plan, None);
|
||||||
|
let state = AppState::new()
|
||||||
|
.unwrap()
|
||||||
|
.with_data_state_for_tests(
|
||||||
|
crate::data::GatewayDataState::with_usage_repository_for_tests(Arc::clone(
|
||||||
|
&repository,
|
||||||
|
))
|
||||||
|
.with_provider_catalog_reader(Arc::new(catalog))
|
||||||
|
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY)
|
||||||
|
.with_system_config_values_for_tests([(
|
||||||
|
"request_record_level".into(),
|
||||||
|
json!("full"),
|
||||||
|
)]),
|
||||||
|
)
|
||||||
|
.with_usage_runtime_for_tests(UsageRuntimeConfig {
|
||||||
|
enabled: true,
|
||||||
|
..Default::default()
|
||||||
|
});
|
||||||
|
let text = "hello".repeat(12_000);
|
||||||
|
let payload = if private {
|
||||||
|
json!({"response":{"candidates":[{"content":{"role":"model","parts":[{"text":text}]},
|
||||||
|
"finishReason":"STOP"}],"usageMetadata":{"promptTokenCount":1234,"candidatesTokenCount":567},
|
||||||
|
"modelVersion":"gemini-3.7-flash-tiered"}})
|
||||||
|
} else {
|
||||||
|
json!({"type":"response.completed","response":{"id":"resp-handoff-usage","status":"completed",
|
||||||
|
"output":[{"type":"message","id":"msg-handoff","role":"assistant","status":"completed",
|
||||||
|
"content":[{"type":"output_text","text":text,"annotations":[]}]}],
|
||||||
|
"usage":{"input_tokens":1234,"output_tokens":567,"total_tokens":1801}}})
|
||||||
|
};
|
||||||
|
let input = format!("data: {payload}\n\n");
|
||||||
|
// One complete large chunk exercises an already-emitted prefetch
|
||||||
|
// result; the private path exercises incomplete normalization too.
|
||||||
|
let chunks = if private {
|
||||||
|
vec![input[..17_735].to_string(), input[17_735..].to_string()]
|
||||||
|
} else {
|
||||||
|
vec![input.clone()]
|
||||||
|
};
|
||||||
|
let frames = stream! {
|
||||||
|
yield Ok::<Bytes, std::io::Error>(ndjson_frame(StreamFrame {
|
||||||
|
frame_type:StreamFrameType::Headers,
|
||||||
|
payload:StreamFramePayload::Headers { status_code:200,
|
||||||
|
headers:BTreeMap::from([("content-type".into(),"text/event-stream".into())]),
|
||||||
|
response_observation:None },
|
||||||
|
}));
|
||||||
|
for chunk in chunks {
|
||||||
|
yield Ok(ndjson_frame(StreamFrame { frame_type:StreamFrameType::Data,
|
||||||
|
payload:StreamFramePayload::Data { text:Some(chunk),chunk_b64:None } }));
|
||||||
|
}
|
||||||
|
yield Ok(ndjson_frame(StreamFrame::eof()));
|
||||||
|
}.boxed();
|
||||||
|
let response = execute_stream_from_frame_stream(
|
||||||
|
&state,
|
||||||
|
plan,
|
||||||
|
"trace-handoff-audit",
|
||||||
|
&test_decision(),
|
||||||
|
OPENAI_RESPONSES_STREAM_PLAN_KIND,
|
||||||
|
Some("openai_responses_stream_success".into()),
|
||||||
|
Some(context),
|
||||||
|
crate::clock::current_unix_ms(),
|
||||||
|
Instant::now(),
|
||||||
|
RequestStageTrace::from_env(),
|
||||||
|
false,
|
||||||
|
frames,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.unwrap();
|
||||||
|
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
|
||||||
|
let body = String::from_utf8(body.to_vec()).unwrap();
|
||||||
|
let events: Vec<Value> = body
|
||||||
|
.lines()
|
||||||
|
.filter_map(|l| l.strip_prefix("data: "))
|
||||||
|
.filter(|p| *p != "[DONE]")
|
||||||
|
.map(|p| serde_json::from_str(p).unwrap())
|
||||||
|
.collect();
|
||||||
|
assert_eq!(
|
||||||
|
events
|
||||||
|
.iter()
|
||||||
|
.filter(|e| e["type"] == "response.completed")
|
||||||
|
.count(),
|
||||||
|
1
|
||||||
|
);
|
||||||
|
assert!(body.contains(&text));
|
||||||
|
let usage = tokio::time::timeout(Duration::from_secs(3), async {
|
||||||
|
loop {
|
||||||
|
if let Some(u) = repository
|
||||||
|
.find_by_request_id(&request_id)
|
||||||
|
.await
|
||||||
|
.unwrap()
|
||||||
|
.filter(|u| u.status == "completed" || u.status == "failed")
|
||||||
|
{
|
||||||
|
break u;
|
||||||
|
}
|
||||||
|
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.await
|
||||||
|
.expect("usage should finalize");
|
||||||
|
assert_eq!(usage.status, "completed", "{:?}", usage.error_message);
|
||||||
|
assert_eq!(usage.input_tokens, 1234);
|
||||||
|
assert_eq!(usage.output_tokens, 567);
|
||||||
|
let captured = usage.response_body.as_ref().expect("provider capture");
|
||||||
|
assert!(
|
||||||
|
captured["metadata"].get("dropped_chunks").is_none(),
|
||||||
|
"{captured}"
|
||||||
|
);
|
||||||
|
assert_eq!(captured["chunks"].as_array().unwrap(), &vec![payload]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn generic_stream_success_regex_matches_fragmented_plain_body() {
|
async fn generic_stream_success_regex_matches_fragmented_plain_body() {
|
||||||
for chunks in [
|
for chunks in [
|
||||||
@@ -9850,7 +9961,7 @@ mod tests {
|
|||||||
let mut buffer = super::StreamUsageObservationBuffer::new(32 * 1024);
|
let mut buffer = super::StreamUsageObservationBuffer::new(32 * 1024);
|
||||||
let mut rewriter = super::maybe_build_stream_response_rewriter(Some(&context)).unwrap();
|
let mut rewriter = super::maybe_build_stream_response_rewriter(Some(&context)).unwrap();
|
||||||
let mut delivered = Vec::new();
|
let mut delivered = Vec::new();
|
||||||
for chunk in chunks {
|
for (index, chunk) in chunks.into_iter().enumerate() {
|
||||||
provider.append(chunk, 32 * 1024, &mut provider_truncated);
|
provider.append(chunk, 32 * 1024, &mut provider_truncated);
|
||||||
super::observe_stream_usage_bytes(
|
super::observe_stream_usage_bytes(
|
||||||
observer.as_mut().unwrap(),
|
observer.as_mut().unwrap(),
|
||||||
@@ -9861,6 +9972,10 @@ mod tests {
|
|||||||
let output = rewriter.push_chunk(chunk).unwrap();
|
let output = rewriter.push_chunk(chunk).unwrap();
|
||||||
client.append(&output, 32 * 1024, &mut client_truncated);
|
client.append(&output, 32 * 1024, &mut client_truncated);
|
||||||
delivered.extend(output);
|
delivered.extend(output);
|
||||||
|
if index == 0 {
|
||||||
|
// Task handoff must also work when audit admits no bytes.
|
||||||
|
rewriter = rewriter.into_owned();
|
||||||
|
}
|
||||||
}
|
}
|
||||||
let tail = rewriter.finish().unwrap();
|
let tail = rewriter.finish().unwrap();
|
||||||
client.append(&tail, 32 * 1024, &mut client_truncated);
|
client.append(&tail, 32 * 1024, &mut client_truncated);
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
use std::borrow::Cow;
|
||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
|
|
||||||
use serde_json::{json, Map, Value};
|
use serde_json::{json, Map, Value};
|
||||||
@@ -226,7 +227,7 @@ enum AiSurfaceStreamRewriteState {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub struct AiSurfaceStreamRewriter<'a> {
|
pub struct AiSurfaceStreamRewriter<'a> {
|
||||||
report_context: &'a Value,
|
report_context: Cow<'a, Value>,
|
||||||
buffered: Vec<u8>,
|
buffered: Vec<u8>,
|
||||||
state: AiSurfaceStreamRewriteState,
|
state: AiSurfaceStreamRewriteState,
|
||||||
}
|
}
|
||||||
@@ -271,30 +272,39 @@ pub fn maybe_build_ai_surface_stream_rewriter<'a>(
|
|||||||
};
|
};
|
||||||
|
|
||||||
Some(AiSurfaceStreamRewriter {
|
Some(AiSurfaceStreamRewriter {
|
||||||
report_context,
|
report_context: Cow::Borrowed(report_context),
|
||||||
buffered: Vec::new(),
|
buffered: Vec::new(),
|
||||||
state,
|
state,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
impl AiSurfaceStreamRewriter<'_> {
|
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> {
|
pub fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||||
match &mut self.state {
|
match &mut self.state {
|
||||||
AiSurfaceStreamRewriteState::OpenAiImage(state) => {
|
AiSurfaceStreamRewriteState::OpenAiImage(state) => {
|
||||||
state.push_chunk(self.report_context, chunk)
|
state.push_chunk(self.report_context.as_ref(), chunk)
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(state) => {
|
AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(state) => {
|
||||||
state.push_chunk(self.report_context, chunk)
|
state.push_chunk(self.report_context.as_ref(), chunk)
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::ClaudeReadToolSanitize(state) => {
|
AiSurfaceStreamRewriteState::ClaudeReadToolSanitize(state) => {
|
||||||
state.push_chunk(self.report_context, chunk)
|
state.push_chunk(self.report_context.as_ref(), chunk)
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::KiroToClaudeCli(state) => {
|
AiSurfaceStreamRewriteState::KiroToClaudeCli(state) => {
|
||||||
state.push_chunk(self.report_context, chunk)
|
state.push_chunk(self.report_context.as_ref(), chunk)
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::KiroToClaudeCliThenStandard { kiro, standard } => {
|
AiSurfaceStreamRewriteState::KiroToClaudeCliThenStandard { kiro, standard } => {
|
||||||
let claude_bytes = kiro.push_chunk(self.report_context, chunk)?;
|
let claude_bytes = kiro.push_chunk(self.report_context.as_ref(), chunk)?;
|
||||||
transform_standard_bytes(standard, self.report_context, claude_bytes)
|
transform_standard_bytes(standard, self.report_context.as_ref(), claude_bytes)
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::EnvelopeUnwrap
|
AiSurfaceStreamRewriteState::EnvelopeUnwrap
|
||||||
| AiSurfaceStreamRewriteState::ModelDirectiveDisplay
|
| AiSurfaceStreamRewriteState::ModelDirectiveDisplay
|
||||||
@@ -313,23 +323,25 @@ impl AiSurfaceStreamRewriter<'_> {
|
|||||||
|
|
||||||
pub fn finish(&mut self) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
pub fn finish(&mut self) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||||
match &mut self.state {
|
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) => {
|
AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(state) => {
|
||||||
state.finish(self.report_context)
|
state.finish(self.report_context.as_ref())
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::ClaudeReadToolSanitize(state) => {
|
AiSurfaceStreamRewriteState::ClaudeReadToolSanitize(state) => {
|
||||||
state.finish(self.report_context)
|
state.finish(self.report_context.as_ref())
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::KiroToClaudeCli(state) => {
|
AiSurfaceStreamRewriteState::KiroToClaudeCli(state) => {
|
||||||
state.finish(self.report_context)
|
state.finish(self.report_context.as_ref())
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::KiroToClaudeCliThenStandard { kiro, standard } => {
|
AiSurfaceStreamRewriteState::KiroToClaudeCliThenStandard { kiro, standard } => {
|
||||||
let mut output = transform_standard_bytes(
|
let mut output = transform_standard_bytes(
|
||||||
standard,
|
standard,
|
||||||
self.report_context,
|
self.report_context.as_ref(),
|
||||||
kiro.finish(self.report_context)?,
|
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)
|
Ok(output)
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::EnvelopeUnwrap
|
AiSurfaceStreamRewriteState::EnvelopeUnwrap
|
||||||
@@ -338,14 +350,14 @@ impl AiSurfaceStreamRewriter<'_> {
|
|||||||
| AiSurfaceStreamRewriteState::Standard(_) => {
|
| AiSurfaceStreamRewriteState::Standard(_) => {
|
||||||
if self.buffered.is_empty() {
|
if self.buffered.is_empty() {
|
||||||
if let AiSurfaceStreamRewriteState::Standard(state) = &mut self.state {
|
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());
|
return Ok(Vec::new());
|
||||||
}
|
}
|
||||||
let line = std::mem::take(&mut self.buffered);
|
let line = std::mem::take(&mut self.buffered);
|
||||||
let mut output = self.transform_line(line)?;
|
let mut output = self.transform_line(line)?;
|
||||||
if let AiSurfaceStreamRewriteState::Standard(state) = &mut self.state {
|
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)
|
Ok(output)
|
||||||
}
|
}
|
||||||
@@ -365,18 +377,19 @@ impl AiSurfaceStreamRewriter<'_> {
|
|||||||
fn transform_line(&mut self, line: Vec<u8>) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
fn transform_line(&mut self, line: Vec<u8>) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||||
match &mut self.state {
|
match &mut self.state {
|
||||||
AiSurfaceStreamRewriteState::EnvelopeUnwrap => {
|
AiSurfaceStreamRewriteState::EnvelopeUnwrap => {
|
||||||
let output = transform_provider_private_stream_line(self.report_context, line)
|
let output =
|
||||||
.map_err(AiSurfaceFinalizeError::from)?;
|
transform_provider_private_stream_line(self.report_context.as_ref(), line)
|
||||||
rewrite_model_directive_stream_line(self.report_context, output)
|
.map_err(AiSurfaceFinalizeError::from)?;
|
||||||
|
rewrite_model_directive_stream_line(self.report_context.as_ref(), output)
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::ModelDirectiveDisplay => {
|
AiSurfaceStreamRewriteState::ModelDirectiveDisplay => {
|
||||||
rewrite_model_directive_stream_line(self.report_context, line)
|
rewrite_model_directive_stream_line(self.report_context.as_ref(), line)
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::OpenAiResponsesCompat => {
|
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) => {
|
AiSurfaceStreamRewriteState::Standard(state) => {
|
||||||
transform_standard_line(state, self.report_context, line)
|
transform_standard_line(state, self.report_context.as_ref(), line)
|
||||||
}
|
}
|
||||||
AiSurfaceStreamRewriteState::OpenAiImage(_)
|
AiSurfaceStreamRewriteState::OpenAiImage(_)
|
||||||
| AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(_)
|
| AiSurfaceStreamRewriteState::OpenAiImageToOpenAiChat(_)
|
||||||
@@ -892,7 +905,7 @@ fn is_standard_cli_client_api_format(api_format: &str) -> bool {
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use serde_json::json;
|
use serde_json::{json, Value};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
maybe_build_ai_surface_stream_rewriter, resolve_finalize_stream_rewrite_mode,
|
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\""));
|
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]
|
#[test]
|
||||||
fn standard_rewriter_converts_openai_responses_reasoning_delta_to_chat() {
|
fn standard_rewriter_converts_openai_responses_reasoning_delta_to_chat() {
|
||||||
let report_context = json!({
|
let report_context = json!({
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
use std::borrow::Cow;
|
||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
|
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
@@ -354,7 +355,7 @@ enum ProviderPrivateStreamNormalizeMode {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub struct ProviderPrivateStreamNormalizer<'a> {
|
pub struct ProviderPrivateStreamNormalizer<'a> {
|
||||||
report_context: &'a Value,
|
report_context: Cow<'a, Value>,
|
||||||
buffered: Vec<u8>,
|
buffered: Vec<u8>,
|
||||||
current_event_type: Option<String>,
|
current_event_type: Option<String>,
|
||||||
mode: ProviderPrivateStreamNormalizeMode,
|
mode: ProviderPrivateStreamNormalizeMode,
|
||||||
@@ -401,7 +402,7 @@ pub fn maybe_build_provider_private_stream_normalizer<'a>(
|
|||||||
return None;
|
return None;
|
||||||
};
|
};
|
||||||
Some(ProviderPrivateStreamNormalizer {
|
Some(ProviderPrivateStreamNormalizer {
|
||||||
report_context,
|
report_context: Cow::Borrowed(report_context),
|
||||||
buffered: Vec::new(),
|
buffered: Vec::new(),
|
||||||
current_event_type: None,
|
current_event_type: None,
|
||||||
mode,
|
mode,
|
||||||
@@ -422,10 +423,20 @@ pub fn extract_provider_private_stream_error_body(
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl ProviderPrivateStreamNormalizer<'_> {
|
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> {
|
pub fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||||
match &mut self.mode {
|
match &mut self.mode {
|
||||||
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
|
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
|
||||||
state.push_chunk(self.report_context, chunk)
|
state.push_chunk(self.report_context.as_ref(), chunk)
|
||||||
}
|
}
|
||||||
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
|
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
|
||||||
let next_len = self
|
let next_len = self
|
||||||
@@ -441,7 +452,7 @@ impl ProviderPrivateStreamNormalizer<'_> {
|
|||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
self.buffered.extend_from_slice(chunk);
|
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)
|
&& buffer_looks_like_connect_frame(&self.buffered)
|
||||||
{
|
{
|
||||||
return drain_windsurf_connect_json_frames(&mut 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<_>>();
|
let line = self.buffered.drain(..=line_end).collect::<Vec<_>>();
|
||||||
output.extend(
|
output.extend(
|
||||||
transform_provider_private_stream_line_with_event_state(
|
transform_provider_private_stream_line_with_event_state(
|
||||||
self.report_context,
|
self.report_context.as_ref(),
|
||||||
line,
|
line,
|
||||||
&mut self.current_event_type,
|
&mut self.current_event_type,
|
||||||
)
|
)
|
||||||
@@ -466,20 +477,20 @@ impl ProviderPrivateStreamNormalizer<'_> {
|
|||||||
pub fn finish(&mut self) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
pub fn finish(&mut self) -> Result<Vec<u8>, AiSurfaceFinalizeError> {
|
||||||
match &mut self.mode {
|
match &mut self.mode {
|
||||||
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
|
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
|
||||||
state.finish(self.report_context)
|
state.finish(self.report_context.as_ref())
|
||||||
}
|
}
|
||||||
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
|
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
|
||||||
if self.buffered.is_empty() {
|
if self.buffered.is_empty() {
|
||||||
return Ok(Vec::new());
|
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)
|
&& buffer_looks_like_connect_frame(&self.buffered)
|
||||||
{
|
{
|
||||||
return drain_windsurf_connect_json_frames(&mut self.buffered);
|
return drain_windsurf_connect_json_frames(&mut self.buffered);
|
||||||
}
|
}
|
||||||
let line = std::mem::take(&mut self.buffered);
|
let line = std::mem::take(&mut self.buffered);
|
||||||
transform_provider_private_stream_line_with_event_state(
|
transform_provider_private_stream_line_with_event_state(
|
||||||
self.report_context,
|
self.report_context.as_ref(),
|
||||||
line,
|
line,
|
||||||
&mut self.current_event_type,
|
&mut self.current_event_type,
|
||||||
)
|
)
|
||||||
@@ -939,7 +950,7 @@ fn postprocess_private_response_value(data: &mut Value, report_context: &Value)
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use serde_json::json;
|
use serde_json::{json, Value};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
extract_provider_private_stream_error_body, maybe_build_provider_private_stream_normalizer,
|
extract_provider_private_stream_error_body, maybe_build_provider_private_stream_normalizer,
|
||||||
@@ -1116,6 +1127,44 @@ mod tests {
|
|||||||
assert!(text.contains(r#""content":"chunk""#));
|
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]
|
#[test]
|
||||||
fn unwraps_windsurf_connect_json_stream_frames() {
|
fn unwraps_windsurf_connect_json_stream_frames() {
|
||||||
let report_context = json!({
|
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 chunks = Vec::new();
|
||||||
let mut total_chunks = 0_u64;
|
let mut total_chunks = 0_u64;
|
||||||
let mut saw_done = false;
|
let mut saw_done = false;
|
||||||
|
let mut first_parse_error = None;
|
||||||
for_each_sse_payload(text, |payload| {
|
for_each_sse_payload(text, |payload| {
|
||||||
if payload == "[DONE]" {
|
if payload == "[DONE]" {
|
||||||
saw_done = true;
|
saw_done = true;
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
total_chunks += 1;
|
total_chunks += 1;
|
||||||
if let Ok(json_body) = serde_json::from_str::<Value>(payload) {
|
match serde_json::from_str::<Value>(payload) {
|
||||||
chunks.push(json_body);
|
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 {
|
if total_chunks == 0 && !saw_done {
|
||||||
@@ -3076,6 +3088,11 @@ fn parse_sse_body_for_storage(text: &str) -> Option<Value> {
|
|||||||
if saw_done {
|
if saw_done {
|
||||||
metadata.insert("has_completion".to_string(), Value::Bool(true));
|
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 {
|
if stored_chunks < total_chunks {
|
||||||
metadata.insert(
|
metadata.insert(
|
||||||
"dropped_chunks".to_string(),
|
"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]
|
#[test]
|
||||||
fn extract_token_counts_from_value_handles_crlf_and_cr_sse_text() {
|
fn extract_token_counts_from_value_handles_crlf_and_cr_sse_text() {
|
||||||
let sse_body = concat!(
|
let sse_body = concat!(
|
||||||
|
|||||||
Reference in New Issue
Block a user