Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
119 changes: 103 additions & 16 deletions codex-rs/ext/guardian-v2/src/sampler.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,11 +15,13 @@ use codex_api::ResponsesApiRequest;
use codex_api::ResponsesWebsocketClient;
use codex_api::ResponsesWebsocketConnection;
use codex_api::ResponsesWsRequest;
use codex_api::TransportError;
use codex_api::build_session_headers;
use codex_api::create_text_param_for_request;
use codex_http_client::HttpClientFactory;
use codex_login::AgentIdentityAuthPolicy;
use codex_login::CodexAuth;
use codex_login::UnauthorizedRecovery;
use codex_login::default_client::add_originator_header;
use codex_login::default_client::default_headers;
use codex_model_provider::AgentIdentitySessionFallback;
Expand All @@ -31,6 +33,7 @@ use codex_protocol::models::ResponseItem;
use codex_protocol::openai_models::ReasoningEffort;
use codex_protocol::protocol::SessionSource;
use http::HeaderValue;
use http::StatusCode;
use serde_json::Value;
use thiserror::Error;
use tokio::sync::OwnedSemaphorePermit;
Expand All @@ -41,6 +44,7 @@ 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 = 16;
const MAX_SAMPLING_RETRIES: usize = 2;
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 =
Expand Down Expand Up @@ -153,10 +157,9 @@ impl LunaSampler {
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 {
for _ in 0..INITIAL_WEBSOCKET_CONNECTIONS {
let connection = match sampler.open_connection().await {
Ok(connection) => connection,
Err(error) if index == 0 => return Err(error),
Err(_) => break,
};
sampler
Expand Down Expand Up @@ -277,6 +280,63 @@ impl LunaSampler {
})
}

async fn retry_after_failure(
&self,
error: &LunaSamplerError,
auth_recovery: &mut Option<UnauthorizedRecovery>,
retries: &mut usize,
) -> bool {
let retryable = match error {
LunaSamplerError::ConnectionTimeout
| LunaSamplerError::Api(
ApiError::Retryable { .. } | ApiError::Stream(_) | ApiError::ServerOverloaded,
)
| LunaSamplerError::Api(ApiError::Transport(
TransportError::RetryLimit
| TransportError::Timeout
| TransportError::Connection(_)
| TransportError::Network(_),
)) => true,
LunaSamplerError::Api(ApiError::Transport(TransportError::Http { status, .. }))
| LunaSamplerError::Api(ApiError::Api { status, .. }) => {
if *status == StatusCode::UNAUTHORIZED {
let Some(recovery) = auth_recovery.as_mut() else {
return false;
};
if !recovery.has_next() || recovery.next().await.is_err() {
return false;
}
self.idle_connections
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clear();
return true;
} else {
status.is_server_error() || *status == StatusCode::TOO_MANY_REQUESTS
}
}
LunaSamplerError::Provider(_)
| LunaSamplerError::MissingOutput
| LunaSamplerError::OutputTooLarge
| LunaSamplerError::Superseded
| LunaSamplerError::Api(
ApiError::Transport(TransportError::Build(_))
| ApiError::ContextWindowExceeded
| ApiError::QuotaExceeded
| ApiError::UsageNotIncluded
| ApiError::RateLimit(_)
| ApiError::InvalidRequest { .. }
| ApiError::MisalignmentPolicyViolation { .. }
| ApiError::CyberPolicy { .. },
) => false,
};
if retryable && *retries < MAX_SAMPLING_RETRIES {
*retries += 1;
return true;
}
false
}

/// Sends one structured, tool-less request on an exclusively leased WebSocket.
pub async fn sample(&self, request: LunaSamplingRequest) -> Result<String, LunaSamplerError> {
let metadata = HashMap::from([
Expand Down Expand Up @@ -377,13 +437,29 @@ impl LunaSampler {
});
}
let mut retries = 0;
let mut auth_recovery = self
.config
.provider
.auth_manager()
.map(|manager| manager.unauthorized_recovery());
'retry: loop {
let lease = tokio::select! {
let lease = match tokio::select! {
biased;
_ = &mut superseded => return Err(LunaSamplerError::Superseded),
lease = self.lease_connection() => lease?,
lease = self.lease_connection() => lease,
} {
Ok(lease) => lease,
Err(error) => {
if self
.retry_after_failure(&error, &mut auth_recovery, &mut retries)
.await
{
continue;
}
return Err(error);
}
};
let mut stream = lease
let mut stream = match lease
.connection
.connection
.stream_request(
Expand All @@ -392,7 +468,19 @@ impl LunaSampler {
/*turn_state*/ None,
)
.await
.map_err(LunaSamplerError::Api)?;
{
Ok(stream) => stream,
Err(error) => {
let error = LunaSamplerError::Api(error);
if self
.retry_after_failure(&error, &mut auth_recovery, &mut retries)
.await
{
continue;
}
return Err(error);
}
};

let mut output = String::new();
let mut deltas = String::new();
Expand All @@ -409,17 +497,16 @@ impl LunaSampler {
} {
let event = match event {
Ok(event) => event,
Err(error)
if retries < INITIAL_WEBSOCKET_CONNECTIONS
&& matches!(
error,
ApiError::Retryable { .. } | ApiError::Stream(_)
) =>
{
retries += 1;
continue 'retry;
Err(error) => {
let error = LunaSamplerError::Api(error);
if self
.retry_after_failure(&error, &mut auth_recovery, &mut retries)
.await
{
continue 'retry;
}
return Err(error);
}
Err(error) => return Err(LunaSamplerError::Api(error)),
};
match event {
ResponseEvent::OutputTextDelta(delta) => {
Expand Down
126 changes: 126 additions & 0 deletions codex-rs/ext/guardian-v2/src/sampler_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -20,13 +20,16 @@ use core_test_support::skip_if_no_network;
use pretty_assertions::assert_eq;
use serde_json::json;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::time::Duration;
use tokio::net::TcpListener;
use tokio::net::TcpStream;

use super::LunaSampler;
use super::LunaSamplerConfig;
use super::LunaSamplingRequest;
use super::MAX_SAMPLING_RETRIES;
use super::MAX_WEBSOCKET_CONNECTIONS;

async fn proxy_websocket_servers(servers: &[&responses::WebSocketTestServer]) -> Result<String> {
Expand Down Expand Up @@ -366,6 +369,51 @@ async fn sampler_returns_complete_json_before_terminal_response_events() -> Resu
Ok(())
}

#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn sampler_recovers_after_initial_prewarm_failures() -> Result<()> {
skip_if_no_network!(Ok(()));

let server = responses::start_websocket_server(vec![vec![vec![
ev_assistant_message("recovered", r#"{"score":0.25}"#),
ev_completed("recovered"),
]]])
.await;
let listener = TcpListener::bind("127.0.0.1:0").await?;
let address = listener.local_addr()?;
let target = server.uri().trim_start_matches("ws://").to_owned();
let failed_connections = Arc::new(AtomicUsize::new(0));
let observed_failures = Arc::clone(&failed_connections);
tokio::spawn(async move {
for _ in 0..=MAX_SAMPLING_RETRIES {
let Ok((connection, _)) = listener.accept().await else {
return;
};
observed_failures.fetch_add(/*val*/ 1, Ordering::Relaxed);
drop(connection);
}
while let Ok((mut incoming, _)) = listener.accept().await {
let target = target.clone();
tokio::spawn(async move {
let Ok(mut outgoing) = TcpStream::connect(target).await else {
return;
};
let _ = tokio::io::copy_bidirectional(&mut incoming, &mut outgoing).await;
});
}
});

let sampler = LunaSampler::connect(sampler_config(format!("http://{address}/v1"))).await?;
assert_eq!(failed_connections.load(Ordering::Relaxed), 1);

assert_eq!(
sampler.sample(sample_request("turn-1")).await?,
r#"{"score":0.25}"#
);
assert_eq!(server.handshakes().len(), 1);

Ok(())
}

#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn sampler_remains_available_when_second_prewarm_fails() -> Result<()> {
skip_if_no_network!(Ok(()));
Expand Down Expand Up @@ -560,3 +608,81 @@ async fn sampler_retries_expired_websockets_on_another_warm_connection() -> Resu
);
Ok(())
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn sampler_reconnects_after_transient_service_failures() -> Result<()> {
skip_if_no_network!(Ok(()));

let unavailable = || {
vec![vec![vec![json!({
"type": "error",
"status": 503,
"error": {
"type": "server_error",
"message": "temporarily unavailable"
}
})]]]
};
let first = responses::start_websocket_server(unavailable()).await;
let second = responses::start_websocket_server(unavailable()).await;
let recovered = responses::start_websocket_server(vec![vec![vec![
ev_assistant_message("recovered", r#"{"score":0.25}"#),
ev_completed("recovered"),
]]])
.await;
let sampler = LunaSampler::connect(sampler_config(
proxy_websocket_servers(&[&first, &second, &recovered]).await?,
))
.await?;

assert_eq!(
sampler.sample(sample_request("turn-1")).await?,
r#"{"score":0.25}"#
);
assert_eq!(first.single_connection().len(), 1);
assert_eq!(second.single_connection().len(), 1);
assert_eq!(recovered.single_connection().len(), 1);
assert_eq!(
recovered.single_handshake().header("authorization"),
Some("Bearer test-api-key".to_owned())
);

Ok(())
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn sampler_limits_transient_recovery_attempts() -> Result<()> {
skip_if_no_network!(Ok(()));

let unavailable = || {
vec![vec![vec![json!({
"type": "error",
"status": 503,
"error": {
"type": "server_error",
"message": "temporarily unavailable"
}
})]]]
};
let first = responses::start_websocket_server(unavailable()).await;
let second = responses::start_websocket_server(unavailable()).await;
let third = responses::start_websocket_server(unavailable()).await;
let unused = responses::start_websocket_server(unavailable()).await;
let sampler = LunaSampler::connect(sampler_config(
proxy_websocket_servers(&[&first, &second, &third, &unused]).await?,
))
.await?;

let error = sampler
.sample(sample_request("turn-1"))
.await
.expect_err("sampling should stop after the bounded retries");

assert!(error.to_string().contains("503"));
assert_eq!(first.single_connection().len(), 1);
assert_eq!(second.single_connection().len(), 1);
assert_eq!(third.single_connection().len(), 1);
assert!(unused.handshakes().is_empty());

Ok(())
}
Loading