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:
jif
2026-08-14 15:49:42 +00:00
committed by copyberry
parent 23094236ac
commit 742edd6c16
4 changed files with 214 additions and 39 deletions

View File

@@ -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 {

View File

@@ -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

View File

@@ -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);
}

View File

@@ -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(()));