From 742edd6c16939185b424beb68c8eb46878b323a3 Mon Sep 17 00:00:00 2001 From: jif Date: Fri, 14 Aug 2026 15:49:42 +0000 Subject: [PATCH] Prioritize new Guardian classifications under load (#38596) ## What changed - Expand the Guardian sampling pool from 8 to 16 WebSocket connections. - When the pool is full, supersede the oldest request that has already produced a score before replacing an unfinished classification. Treat superseded classifications as a no-op in the extension. - Stop response WebSocket work when its event consumer is dropped, including while waiting for the connection lock or draining a completed sample. - Allow retries across both initially warmed connections for retryable stream failures. ## Testing - Add a concurrent sampler test that fills the pool and verifies scored drains are replaced before an unfinished classification. GitOrigin-RevId: 9e4df7b01516094726ff3e11878f8ca742d21b03 --- .../src/endpoint/responses_websocket.rs | 34 ++++-- codex-rs/ext/guardian-v2/src/extension.rs | 9 +- codex-rs/ext/guardian-v2/src/sampler.rs | 109 +++++++++++++----- codex-rs/ext/guardian-v2/src/sampler_tests.rs | 101 ++++++++++++++++ 4 files changed, 214 insertions(+), 39 deletions(-) diff --git a/codex-rs/codex-api/src/endpoint/responses_websocket.rs b/codex-rs/codex-api/src/endpoint/responses_websocket.rs index 733658fb39..96c2e04a3e 100644 --- a/codex-rs/codex-api/src/endpoint/responses_websocket.rs +++ b/codex-rs/codex-api/src/endpoint/responses_websocket.rs @@ -282,7 +282,14 @@ impl ResponsesWebsocketConnection { .send(Ok(ResponseEvent::ServerReasoningIncluded(true))) .await; } - let mut guard = stream.lock().await; + let mut guard = tokio::select! { + biased; + _ = tx_event.closed() => return, + guard = stream.lock() => guard, + }; + if tx_event.is_closed() { + return; + } let result = { let Some(ws_stream) = guard.as_mut() else { let _ = tx_event @@ -293,16 +300,21 @@ impl ResponsesWebsocketConnection { return; }; - run_websocket_response_stream( - ws_stream, - tx_event.clone(), - request_text, - idle_timeout, - telemetry, - turn_state.as_deref(), - &timing_log_context, - ) - .await + tokio::select! { + biased; + result = run_websocket_response_stream( + ws_stream, + tx_event.clone(), + request_text, + idle_timeout, + telemetry, + turn_state.as_deref(), + &timing_log_context, + ) => result, + _ = tx_event.closed() => Err(ApiError::Stream( + "response event consumer dropped".to_string(), + )), + } }; if let Err(err) = result { diff --git a/codex-rs/ext/guardian-v2/src/extension.rs b/codex-rs/ext/guardian-v2/src/extension.rs index a9b6696988..55e08b55e3 100644 --- a/codex-rs/ext/guardian-v2/src/extension.rs +++ b/codex-rs/ext/guardian-v2/src/extension.rs @@ -34,6 +34,7 @@ use serde_json::json; use crate::LunaSampler; use crate::LunaSamplerConfig; +use crate::LunaSamplerError; use crate::LunaSamplingRequest; use crate::sampler::MODEL; use crate::transcript::TranscriptConfig; @@ -317,7 +318,7 @@ impl ToolLifecycleContributor for GuardianV2Extension { ">>> APPROVAL REQUEST END\n".to_owned(), ]); let result: Result<(), String> = async { - let output = sampler + let output = match sampler .sample(LunaSamplingRequest { instructions: CLASSIFIER_INSTRUCTIONS.to_owned(), input: classification_input, @@ -346,7 +347,11 @@ impl ToolLifecycleContributor for GuardianV2Extension { turn_id: turn_id.clone(), }) .await - .map_err(|error| error.to_string())?; + { + Ok(output) => output, + Err(LunaSamplerError::Superseded) => return Ok(()), + Err(error) => return Err(error.to_string()), + }; let output: serde_json::Value = serde_json::from_str(&output).map_err(|error| error.to_string())?; let scores = output diff --git a/codex-rs/ext/guardian-v2/src/sampler.rs b/codex-rs/ext/guardian-v2/src/sampler.rs index f806b6aa3b..a02acccb70 100644 --- a/codex-rs/ext/guardian-v2/src/sampler.rs +++ b/codex-rs/ext/guardian-v2/src/sampler.rs @@ -1,6 +1,9 @@ use std::collections::HashMap; +use std::collections::VecDeque; use std::sync::Arc; use std::sync::Mutex; +use std::sync::atomic::AtomicBool; +use std::sync::atomic::Ordering; use std::time::Duration; use std::time::Instant; @@ -32,11 +35,12 @@ use serde_json::Value; use thiserror::Error; use tokio::sync::OwnedSemaphorePermit; use tokio::sync::Semaphore; +use tokio::sync::oneshot; pub(crate) const MODEL: &str = "gpt-5.6-luna"; const MAX_OUTPUT_BYTES: usize = 8 * 1024; const INITIAL_WEBSOCKET_CONNECTIONS: usize = 2; -const MAX_WEBSOCKET_CONNECTIONS: usize = 8; +const MAX_WEBSOCKET_CONNECTIONS: usize = 16; const MAX_WEBSOCKET_AGE: Duration = Duration::from_secs(55 * 60); const RESPONSES_WEBSOCKETS_BETA: &str = "responses_websockets=2026-02-06"; const RESPONSES_LITE_METADATA_KEY: &str = @@ -100,6 +104,9 @@ pub enum LunaSamplerError { /// The response exceeded the bounded output limit. #[error("Luna response exceeded the output limit")] OutputTooLarge, + /// A newer classification replaced this request when the pool was full. + #[error("Luna request was superseded by a newer classification")] + Superseded, } struct PooledConnection { @@ -122,11 +129,17 @@ impl ConnectionLease { } } +struct ActiveRequest { + supersede: oneshot::Sender<()>, + scored: Arc, +} + /// A bounded pool of authenticated Responses WebSockets dedicated to Luna sampling. pub struct LunaSampler { config: LunaSamplerConfig, idle_connections: Arc>>, capacity: Arc, + active_requests: Mutex>, } impl LunaSampler { @@ -136,6 +149,7 @@ impl LunaSampler { config, idle_connections: Arc::new(Mutex::new(Vec::with_capacity(MAX_WEBSOCKET_CONNECTIONS))), capacity: Arc::new(Semaphore::new(MAX_WEBSOCKET_CONNECTIONS)), + active_requests: Mutex::new(VecDeque::with_capacity(MAX_WEBSOCKET_CONNECTIONS)), }; for index in 0..INITIAL_WEBSOCKET_CONNECTIONS { let connection = match sampler.open_connection().await { @@ -332,9 +346,35 @@ impl LunaSampler { ), client_metadata: Some(metadata), }; - let mut retried = false; + let (supersede, mut superseded) = oneshot::channel(); + let scored = Arc::new(AtomicBool::new(false)); + { + let mut active_requests = self + .active_requests + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + active_requests.retain(|request| !request.supersede.is_closed()); + if active_requests.len() == MAX_WEBSOCKET_CONNECTIONS { + let oldest_scored = active_requests + .iter() + .position(|request| request.scored.load(Ordering::Relaxed)) + .unwrap_or(0); + if let Some(oldest) = active_requests.remove(oldest_scored) { + let _ = oldest.supersede.send(()); + } + } + active_requests.push_back(ActiveRequest { + supersede, + scored: Arc::clone(&scored), + }); + } + let mut retries = 0; 'retry: loop { - let lease = self.lease_connection().await?; + let lease = tokio::select! { + biased; + _ = &mut superseded => return Err(LunaSamplerError::Superseded), + lease = self.lease_connection() => lease?, + }; let mut stream = lease .connection .connection @@ -348,17 +388,27 @@ impl LunaSampler { let mut output = String::new(); let mut deltas = String::new(); - while let Some(event) = stream.rx_event.recv().await { + while let Some(event) = tokio::select! { + biased; + _ = &mut superseded => { + return if scored.load(Ordering::Relaxed) && !output.is_empty() { + Ok(output) + } else { + Err(LunaSamplerError::Superseded) + }; + } + event = stream.rx_event.recv() => event, + } { let event = match event { Ok(event) => event, Err(error) - if !retried + if retries < INITIAL_WEBSOCKET_CONNECTIONS && matches!( error, ApiError::Retryable { .. } | ApiError::Stream(_) ) => { - retried = true; + retries += 1; continue 'retry; } Err(error) => return Err(LunaSamplerError::Api(error)), @@ -366,26 +416,6 @@ impl LunaSampler { match event { ResponseEvent::OutputTextDelta(delta) => { deltas.push_str(&delta); - if deltas.len() > MAX_OUTPUT_BYTES { - return Err(LunaSamplerError::OutputTooLarge); - } - - if serde_json::from_str::>(&deltas).is_ok() { - let mut remaining_events = stream.rx_event; - tokio::spawn(async move { - while let Some(event) = remaining_events.recv().await { - match event { - Ok(ResponseEvent::Completed { .. }) => { - lease.reuse(); - break; - } - Err(_) => break, - _ => {} - } - } - }); - return Ok(deltas); - } } ResponseEvent::OutputItemDone(ResponseItem::Message { role, content, .. @@ -411,6 +441,33 @@ impl LunaSampler { if output.len() > MAX_OUTPUT_BYTES || deltas.len() > MAX_OUTPUT_BYTES { return Err(LunaSamplerError::OutputTooLarge); } + if !output.is_empty() { + if serde_json::from_str::>(&output).is_ok() { + scored.store(true, Ordering::Relaxed); + } + continue; + } + if serde_json::from_str::>(&deltas).is_ok() { + scored.store(true, Ordering::Relaxed); + let mut remaining_events = stream.rx_event; + tokio::spawn(async move { + while let Some(event) = tokio::select! { + biased; + _ = &mut superseded => None, + event = remaining_events.recv() => event, + } { + match event { + Ok(ResponseEvent::Completed { .. }) => { + lease.reuse(); + break; + } + Err(_) => break, + _ => {} + } + } + }); + return Ok(deltas); + } } return Err(LunaSamplerError::MissingOutput); } diff --git a/codex-rs/ext/guardian-v2/src/sampler_tests.rs b/codex-rs/ext/guardian-v2/src/sampler_tests.rs index 5c7848a599..5e27d62c9f 100644 --- a/codex-rs/ext/guardian-v2/src/sampler_tests.rs +++ b/codex-rs/ext/guardian-v2/src/sampler_tests.rs @@ -19,6 +19,7 @@ use core_test_support::responses::ev_output_text_delta; use core_test_support::skip_if_no_network; use pretty_assertions::assert_eq; use serde_json::json; +use std::sync::Arc; use std::time::Duration; use tokio::net::TcpListener; use tokio::net::TcpStream; @@ -26,6 +27,7 @@ use tokio::net::TcpStream; use super::LunaSampler; use super::LunaSamplerConfig; use super::LunaSamplingRequest; +use super::MAX_WEBSOCKET_CONNECTIONS; async fn proxy_websocket_servers(servers: &[&responses::WebSocketTestServer]) -> Result { let listener = TcpListener::bind("127.0.0.1:0").await?; @@ -418,6 +420,105 @@ async fn sampler_grows_its_pool_for_overlapping_requests() -> Result<()> { Ok(()) } +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn sampler_replaces_scored_drains_before_unfinished_classifications() -> Result<()> { + skip_if_no_network!(Ok(())); + + let incomplete_response = WebSocketConnectionConfig { + requests: vec![vec![ev_output_text_delta(r#"{"score":0.25}"#)]], + response_headers: Vec::new(), + accept_delay: None, + close_after_requests: false, + }; + let scored_response = WebSocketConnectionConfig { + requests: vec![vec![ev_assistant_message("scored", r#"{"score":0.25}"#)]], + ..incomplete_response.clone() + }; + let stalled_response = WebSocketConnectionConfig { + requests: vec![Vec::new()], + response_headers: Vec::new(), + accept_delay: None, + close_after_requests: false, + }; + let mut servers = Vec::with_capacity(MAX_WEBSOCKET_CONNECTIONS + 2); + servers.push(responses::start_websocket_server_with_headers(vec![scored_response]).await); + servers.push(responses::start_websocket_server_with_headers(vec![stalled_response]).await); + for _ in 2..=MAX_WEBSOCKET_CONNECTIONS { + servers.push( + responses::start_websocket_server_with_headers(vec![incomplete_response.clone()]).await, + ); + } + servers.push( + responses::start_websocket_server(vec![vec![vec![ + ev_assistant_message("newest", r#"{"score":0.75}"#), + ev_completed("newest"), + ]]]) + .await, + ); + let server_refs = servers.iter().collect::>(); + let sampler = Arc::new( + LunaSampler::connect(sampler_config(proxy_websocket_servers(&server_refs).await?)).await?, + ); + + let oldest_sampler = Arc::clone(&sampler); + let oldest = tokio::spawn(async move { oldest_sampler.sample(sample_request("oldest")).await }); + tokio::time::timeout( + Duration::from_secs(2), + servers[1].wait_for_request(/*connection_index*/ 0, /*request_index*/ 0), + ) + .await?; + + let scored_sampler = Arc::clone(&sampler); + let scored_request = + tokio::spawn(async move { scored_sampler.sample(sample_request("scored")).await }); + tokio::time::timeout( + Duration::from_secs(2), + servers[0].wait_for_request(/*connection_index*/ 0, /*request_index*/ 0), + ) + .await?; + + for index in 0..MAX_WEBSOCKET_CONNECTIONS - 2 { + assert_eq!( + sampler + .sample(sample_request(&format!("turn-{index}"))) + .await?, + r#"{"score":0.25}"# + ); + } + + assert_eq!( + tokio::time::timeout( + Duration::from_secs(2), + sampler.sample(sample_request("replace-oldest")), + ) + .await??, + r#"{"score":0.25}"# + ); + assert_eq!( + tokio::time::timeout(Duration::from_secs(2), scored_request).await???, + r#"{"score":0.25}"# + ); + assert!(!oldest.is_finished()); + + assert_eq!( + tokio::time::timeout( + Duration::from_secs(2), + sampler.sample(sample_request("replace-oldest-drain")), + ) + .await??, + r#"{"score":0.75}"# + ); + + assert!(!oldest.is_finished()); + oldest.abort(); + let _ = oldest.await; + drop(sampler); + for server in servers { + server.shutdown().await; + } + Ok(()) +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] async fn sampler_retries_expired_websockets_on_another_warm_connection() -> Result<()> { skip_if_no_network!(Ok(()));