From be33f80bc65159c094ecd06bf155afa3061ce23d Mon Sep 17 00:00:00 2001 From: Francis Chalissery Date: Sat, 4 Jul 2026 17:07:53 -0700 Subject: [PATCH] [codex] Read buffering metadata from response events (#31064) ## Summary - read optional faster-model metadata from streamed buffering payloads - use the buffering payload itself to determine whether buffering UI should be shown - retain the existing header value as a compatibility fallback when the payload omits the field ## Behavior An object-valued buffering signal now enables the buffering UI. The response event's faster-model field takes precedence when present, while omitted fields fall back to existing response metadata. An explicit null leaves the retry target unset. ## Validation - `just test -p codex-api` - `just fix -p codex-api` - `cargo fmt --all -- --check` - `git diff --check` --- codex-rs/codex-api/src/common.rs | 10 -- .../src/endpoint/responses_websocket.rs | 95 +++++++++++++++++-- codex-rs/codex-api/src/safety_buffering.rs | 47 +++++---- codex-rs/codex-api/src/sse/responses.rs | 66 +++++++++++-- codex-rs/core/tests/suite/safety_buffering.rs | 57 ++++++++++- 5 files changed, 231 insertions(+), 44 deletions(-) diff --git a/codex-rs/codex-api/src/common.rs b/codex-rs/codex-api/src/common.rs index bd037deb58..f00cd742e6 100644 --- a/codex-rs/codex-api/src/common.rs +++ b/codex-rs/codex-api/src/common.rs @@ -121,21 +121,11 @@ pub struct SafetyBuffering { pub reasons: Vec, #[serde(skip)] pub show_buffering_ui: bool, - #[serde(skip)] pub faster_model: Option, } -impl SafetyBuffering { - pub(crate) fn with_treatment(mut self, treatment: &SafetyBufferingTreatment) -> Self { - self.show_buffering_ui = treatment.show_buffering_ui; - self.faster_model.clone_from(&treatment.faster_model); - self - } -} - #[derive(Debug, Clone, Default, PartialEq, Eq)] pub(crate) struct SafetyBufferingTreatment { - pub show_buffering_ui: bool, pub faster_model: Option, } diff --git a/codex-rs/codex-api/src/endpoint/responses_websocket.rs b/codex-rs/codex-api/src/endpoint/responses_websocket.rs index df41608184..c4c617ac99 100644 --- a/codex-rs/codex-api/src/endpoint/responses_websocket.rs +++ b/codex-rs/codex-api/src/endpoint/responses_websocket.rs @@ -689,17 +689,10 @@ async fn run_websocket_response_stream( { let _ = turn_state.set(response_turn_state); } - if let Some(headers) = event.headers.as_ref().and_then(Value::as_object) - && let Some(treatment) = - treatment_from_headers(&json_headers_to_http_headers(headers)) - { - safety_buffering_treatment = treatment; - } let model_verifications = event.model_verifications(); let turn_moderation_metadata = event.turn_moderation_metadata(); - let safety_buffering = event - .safety_buffering() - .map(|buffering| buffering.with_treatment(&safety_buffering_treatment)); + let safety_buffering = + safety_buffering_for_event(&event, &mut safety_buffering_treatment); if event.kind() == "codex.rate_limits" { if let Some(snapshot) = parse_rate_limit_event(&text) { let _ = tx_event.send(Ok(ResponseEvent::RateLimits(snapshot))).await; @@ -774,6 +767,19 @@ async fn run_websocket_response_stream( Ok(()) } +fn safety_buffering_for_event( + event: &ResponsesStreamEvent, + treatment: &mut SafetyBufferingTreatment, +) -> Option { + if let Some(headers) = event.headers.as_ref().and_then(Value::as_object) + && let Some(updated_treatment) = + treatment_from_headers(&json_headers_to_http_headers(headers)) + { + *treatment = updated_treatment; + } + event.safety_buffering(treatment) +} + async fn send_websocket_request( ws_stream: &WsStream, request_text: String, @@ -1039,4 +1045,75 @@ mod tests { Some(&HeaderValue::from_static("default-only")) ); } + + #[test] + fn websocket_safety_buffering_uses_event_before_header_fallback() { + let metadata: ResponsesStreamEvent = serde_json::from_value(json!({ + "type": "codex.response.metadata", + "headers": { + "x-codex-safety-buffering-enabled": "true", + "x-codex-safety-buffering-faster-model": "gpt-fast-header" + } + })) + .expect("deserialize treatment metadata"); + let event: ResponsesStreamEvent = serde_json::from_value(json!({ + "type": "response.output_text.delta", + "safety_buffering": { + "use_cases": ["cyber"], + "reasons": ["user_risk"], + "faster_model": "gpt-fast-wire" + } + })) + .expect("deserialize safety buffering event"); + let mut treatment = SafetyBufferingTreatment::default(); + + assert!(safety_buffering_for_event(&metadata, &mut treatment).is_none()); + let buffering = safety_buffering_for_event(&event, &mut treatment) + .expect("expected safety buffering payload"); + + assert_eq!( + buffering, + crate::common::SafetyBuffering { + use_cases: vec!["cyber".to_string()], + reasons: vec!["user_risk".to_string()], + show_buffering_ui: true, + faster_model: Some("gpt-fast-wire".to_string()), + } + ); + } + + #[test] + fn websocket_safety_buffering_event_controls_visibility_when_header_disables_it() { + let metadata: ResponsesStreamEvent = serde_json::from_value(json!({ + "type": "codex.response.metadata", + "headers": { + "x-codex-safety-buffering-enabled": "false", + "x-codex-safety-buffering-faster-model": "gpt-fast-header" + } + })) + .expect("deserialize treatment metadata"); + let event: ResponsesStreamEvent = serde_json::from_value(json!({ + "type": "response.output_text.delta", + "safety_buffering": { + "use_cases": ["cyber"], + "reasons": ["user_risk"] + } + })) + .expect("deserialize safety buffering event"); + let mut treatment = SafetyBufferingTreatment::default(); + + assert!(safety_buffering_for_event(&metadata, &mut treatment).is_none()); + let buffering = safety_buffering_for_event(&event, &mut treatment) + .expect("expected safety buffering payload"); + + assert_eq!( + buffering, + crate::common::SafetyBuffering { + use_cases: vec!["cyber".to_string()], + reasons: vec!["user_risk".to_string()], + show_buffering_ui: true, + faster_model: Some("gpt-fast-header".to_string()), + } + ); + } } diff --git a/codex-rs/codex-api/src/safety_buffering.rs b/codex-rs/codex-api/src/safety_buffering.rs index e5f3dcd79e..aaf09c8eb4 100644 --- a/codex-rs/codex-api/src/safety_buffering.rs +++ b/codex-rs/codex-api/src/safety_buffering.rs @@ -6,23 +6,17 @@ pub(crate) const X_CODEX_SAFETY_BUFFERING_FASTER_MODEL_HEADER: &str = "x-codex-safety-buffering-faster-model"; pub(crate) fn treatment_from_headers(headers: &HeaderMap) -> Option { - let show_buffering_ui = headers - .get(X_CODEX_SAFETY_BUFFERING_ENABLED_HEADER) - .and_then(|value| value.to_str().ok())? - .eq_ignore_ascii_case("true"); - let faster_model = if show_buffering_ui { - headers - .get(X_CODEX_SAFETY_BUFFERING_FASTER_MODEL_HEADER) - .and_then(|value| value.to_str().ok()) - .map(str::to_string) - } else { - None - }; + if !headers.contains_key(X_CODEX_SAFETY_BUFFERING_ENABLED_HEADER) + && !headers.contains_key(X_CODEX_SAFETY_BUFFERING_FASTER_MODEL_HEADER) + { + return None; + } + let faster_model = headers + .get(X_CODEX_SAFETY_BUFFERING_FASTER_MODEL_HEADER) + .and_then(|value| value.to_str().ok()) + .map(str::to_string); - Some(SafetyBufferingTreatment { - show_buffering_ui, - faster_model, - }) + Some(SafetyBufferingTreatment { faster_model }) } #[cfg(test)] @@ -46,7 +40,26 @@ mod tests { assert_eq!( treatment_from_headers(&headers), Some(SafetyBufferingTreatment { - show_buffering_ui: true, + faster_model: Some("faster-model".to_string()), + }) + ); + } + + #[test] + fn buffering_enabled_header_does_not_gate_the_faster_model_fallback() { + let mut headers = HeaderMap::new(); + headers.insert( + X_CODEX_SAFETY_BUFFERING_ENABLED_HEADER, + HeaderValue::from_static("false"), + ); + headers.insert( + X_CODEX_SAFETY_BUFFERING_FASTER_MODEL_HEADER, + HeaderValue::from_static("faster-model"), + ); + + assert_eq!( + treatment_from_headers(&headers), + Some(SafetyBufferingTreatment { faster_model: Some("faster-model".to_string()), }) ); diff --git a/codex-rs/codex-api/src/sse/responses.rs b/codex-rs/codex-api/src/sse/responses.rs index 71f73c0232..bd742dfa14 100644 --- a/codex-rs/codex-api/src/sse/responses.rs +++ b/codex-rs/codex-api/src/sse/responses.rs @@ -231,8 +231,18 @@ impl ResponsesStreamEvent { .map(|metadata| TurnModerationMetadataEvent { metadata }) } - pub(crate) fn safety_buffering(&self) -> Option { - serde_json::from_value(self.safety_buffering.as_ref()?.clone()).ok() + pub(crate) fn safety_buffering( + &self, + treatment: &SafetyBufferingTreatment, + ) -> Option { + let value = self.safety_buffering.as_ref()?; + let faster_model_present = value.as_object()?.contains_key("faster_model"); + let mut buffering: SafetyBuffering = serde_json::from_value(value.clone()).ok()?; + buffering.show_buffering_ui = true; + if !faster_model_present { + buffering.faster_model.clone_from(&treatment.faster_model); + } + Some(buffering) } } @@ -515,9 +525,7 @@ async fn process_sse_with_treatment( }; let model_verifications = event.model_verifications(); let turn_moderation_metadata = event.turn_moderation_metadata(); - let safety_buffering = event - .safety_buffering() - .map(|buffering| buffering.with_treatment(&safety_buffering_treatment)); + let safety_buffering = event.safety_buffering(&safety_buffering_treatment); if let Some(model) = event.response_model() && last_server_model.as_deref() != Some(model.as_str()) @@ -1354,7 +1362,8 @@ mod tests { "delta": "hello", "safety_buffering": { "use_cases": ["cyber"], - "reasons": ["user_risk"] + "reasons": ["user_risk"], + "faster_model": "gpt-fast-wire" } }), json!({ @@ -1381,7 +1390,10 @@ mod tests { assert_matches!( &events[1], ResponseEvent::SafetyBuffering(buffering) - if buffering.use_cases == ["cyber"] && buffering.reasons == ["user_risk"] + if buffering.use_cases == ["cyber"] + && buffering.reasons == ["user_risk"] + && buffering.show_buffering_ui + && buffering.faster_model.as_deref() == Some("gpt-fast-wire") ); assert_matches!(&events[2], ResponseEvent::OutputTextDelta(delta) if delta == "hello"); assert_matches!( @@ -1398,6 +1410,46 @@ mod tests { assert_matches!(&events[6], ResponseEvent::Completed { response_id, .. } if response_id == "resp-1"); } + #[test] + fn safety_buffering_prefers_wire_faster_model_and_only_falls_back_when_omitted() { + let treatment = SafetyBufferingTreatment { + faster_model: Some("gpt-fast-header".to_string()), + }; + + for (faster_model, expected_faster_model) in [ + (None, Some("gpt-fast-header")), + (Some(Value::Null), None), + (Some(json!("gpt-fast-wire")), Some("gpt-fast-wire")), + ] { + let mut event = json!({ + "type": "response.output_text.delta", + "safety_buffering": { + "use_cases": ["cyber"], + "reasons": ["user_risk"] + } + }); + if let Some(faster_model) = faster_model { + event["safety_buffering"]["faster_model"] = faster_model; + } + let event: ResponsesStreamEvent = + serde_json::from_value(event).expect("deserialize safety buffering event"); + + let buffering = event + .safety_buffering(&treatment) + .expect("expected safety buffering payload"); + + assert_eq!( + buffering, + SafetyBuffering { + use_cases: vec!["cyber".to_string()], + reasons: vec!["user_risk".to_string()], + show_buffering_ui: true, + faster_model: expected_faster_model.map(str::to_string), + } + ); + } + } + #[test] fn responses_stream_event_response_model_reads_top_level_headers() { let ev: ResponsesStreamEvent = serde_json::from_value(json!({ diff --git a/codex-rs/core/tests/suite/safety_buffering.rs b/codex-rs/core/tests/suite/safety_buffering.rs index 6e41b23766..5a8393f864 100644 --- a/codex-rs/core/tests/suite/safety_buffering.rs +++ b/codex-rs/core/tests/suite/safety_buffering.rs @@ -19,7 +19,7 @@ use serde_json::json; const FASTER_MODEL: &str = "faster-model"; #[tokio::test(flavor = "multi_thread", worker_threads = 2)] -async fn emits_safety_buffering_with_the_requested_model() -> anyhow::Result<()> { +async fn emits_safety_buffering_with_the_header_fallback_model() -> anyhow::Result<()> { skip_if_no_network!(Ok(())); let server = start_mock_server().await; @@ -72,3 +72,58 @@ async fn emits_safety_buffering_with_the_requested_model() -> anyhow::Result<()> Ok(()) } + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn emits_safety_buffering_with_the_responses_api_model_without_header_gating() +-> anyhow::Result<()> { + skip_if_no_network!(Ok(())); + + let server = start_mock_server().await; + let mut created = ev_response_created("resp-1"); + created["safety_buffering"] = json!({ + "use_cases": ["cyber"], + "reasons": ["policy-check"], + "faster_model": FASTER_MODEL, + }); + mount_response_once( + &server, + sse_response(sse(vec![created, ev_completed("resp-1")])), + ) + .await; + + let test = test_codex().build(&server).await?; + test.codex + .submit(Op::UserInput { + items: vec![UserInput::Text { + text: "Check this request".into(), + text_elements: Vec::new(), + }], + final_output_json_schema: None, + responsesapi_client_metadata: None, + additional_context: Default::default(), + thread_settings: Default::default(), + }) + .await?; + + let event = wait_for_event_match(&test.codex, |event| match event { + EventMsg::SafetyBuffering(event) => Some(event.clone()), + _ => None, + }) + .await; + assert_eq!( + event, + SafetyBufferingEvent { + model: test.session_configured.model.clone(), + use_cases: vec!["cyber".to_string()], + reasons: vec!["policy-check".to_string()], + show_buffering_ui: true, + faster_model: Some(FASTER_MODEL.to_string()), + } + ); + wait_for_event(&test.codex, |event| { + matches!(event, EventMsg::TurnComplete(_)) + }) + .await; + + Ok(()) +}