mirror of
https://github.com/openai/codex.git
synced 2026-08-23 13:09:46 +00:00
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
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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<AtomicBool>,
|
||||
}
|
||||
|
||||
/// A bounded pool of authenticated Responses WebSockets dedicated to Luna sampling.
|
||||
pub struct LunaSampler {
|
||||
config: LunaSamplerConfig,
|
||||
idle_connections: Arc<Mutex<Vec<PooledConnection>>>,
|
||||
capacity: Arc<Semaphore>,
|
||||
active_requests: Mutex<VecDeque<ActiveRequest>>,
|
||||
}
|
||||
|
||||
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::<serde_json::Map<String, Value>>(&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::<serde_json::Map<String, Value>>(&output).is_ok() {
|
||||
scored.store(true, Ordering::Relaxed);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
if serde_json::from_str::<serde_json::Map<String, Value>>(&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);
|
||||
}
|
||||
|
||||
@@ -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<String> {
|
||||
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::<Vec<_>>();
|
||||
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(()));
|
||||
|
||||
Reference in New Issue
Block a user