mirror of
https://github.com/openai/codex.git
synced 2026-08-23 13:09:46 +00:00
[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`
This commit is contained in:
committed by
GitHub
parent
98d28aab54
commit
be33f80bc6
@@ -121,21 +121,11 @@ pub struct SafetyBuffering {
|
||||
pub reasons: Vec<String>,
|
||||
#[serde(skip)]
|
||||
pub show_buffering_ui: bool,
|
||||
#[serde(skip)]
|
||||
pub faster_model: Option<String>,
|
||||
}
|
||||
|
||||
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<String>,
|
||||
}
|
||||
|
||||
|
||||
@@ -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<crate::common::SafetyBuffering> {
|
||||
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()),
|
||||
}
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<SafetyBufferingTreatment> {
|
||||
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()),
|
||||
})
|
||||
);
|
||||
|
||||
@@ -231,8 +231,18 @@ impl ResponsesStreamEvent {
|
||||
.map(|metadata| TurnModerationMetadataEvent { metadata })
|
||||
}
|
||||
|
||||
pub(crate) fn safety_buffering(&self) -> Option<SafetyBuffering> {
|
||||
serde_json::from_value(self.safety_buffering.as_ref()?.clone()).ok()
|
||||
pub(crate) fn safety_buffering(
|
||||
&self,
|
||||
treatment: &SafetyBufferingTreatment,
|
||||
) -> Option<SafetyBuffering> {
|
||||
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!({
|
||||
|
||||
@@ -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(())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user