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
122 changes: 122 additions & 0 deletions codex-rs/ext/guardian-v2/src/extension_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,9 @@ use codex_features::Feature;
use codex_history::RolloutItem;
use codex_login::AuthManager;
use codex_login::CodexAuth;
use codex_login::ExternalAuth;
use codex_login::ExternalAuthFuture;
use codex_login::ExternalAuthRefreshContext;
use codex_model_provider_info::ModelProviderInfo;
use codex_protocol::ResponseItemId;
use codex_protocol::models::ContentItem;
Expand Down Expand Up @@ -55,6 +58,125 @@ const TEST_GUARDIAN_POLICY: &str =
const TEST_CATALOG_GUARDIAN_POLICY: &str =
"Require review before sending organization data to third-party services.";

struct RefreshableAuth(std::sync::Mutex<&'static str>);

impl ExternalAuth for RefreshableAuth {
fn resolve(&self) -> ExternalAuthFuture<'_, CodexAuth> {
Box::pin(async { Ok(CodexAuth::from_api_key(*self.0.lock().expect("auth"))) })
}

fn refresh(&self, _: ExternalAuthRefreshContext) -> ExternalAuthFuture<'_, CodexAuth> {
*self.0.lock().expect("auth") = "refreshed";
self.resolve()
}
}

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

let thread_server = responses::start_mock_server().await;
let test = test_codex().build_with_auto_env(&thread_server).await?;
let events = vec![
ev_assistant_message("sample", r#"{"scores":{"action_risk":0.25}}"#),
ev_completed("response-1"),
];
// Keep the sampled connection open for another request so only auth
// invalidation, not a server close, forces the next handshake.
let server = responses::start_websocket_server(vec![
Vec::new(),
vec![events.clone(), events.clone()],
vec![events],
])
.await;
let auth_manager = AuthManager::from_auth_for_testing(CodexAuth::from_api_key("original"));
auth_manager
.set_external_auth(Arc::new(RefreshableAuth(std::sync::Mutex::new("original"))))
.await?;
let mut config = test.config.clone();
config.model_provider = ModelProviderInfo::create_openai_provider(Some(format!(
"http://{}/v1",
server.uri().trim_start_matches("ws://")
)));
config.features.enable(Feature::GuardianV2)?;
let mut builder = ExtensionRegistryBuilder::new();
crate::install(
&mut builder,
auth_manager.clone(),
Arc::downgrade(&test.thread_manager),
);
let registry = builder.build();
let session_store = ExtensionData::new("session-1");
let thread_store = test.codex.thread_extension_data();
registry.thread_lifecycle_contributors()[0]
.on_thread_start(ThreadStartInput {
config: &config,
session_source: &SessionSource::Exec,
persistent_thread_state_available: false,
environments: &[],
mcp_resource_client: None,
extension_metrics: None,
session_store: &session_store,
thread_store,
})
.await;
let progress = thread_store
.get::<GuardianV2ScoreProgress>()
.expect("Guardian v2 should initialize");
let turn_store = ExtensionData::new("turn-1");
let tool_name = ToolName::plain("read_file");
let payload = ToolPayload::Function {
arguments: r#"{"path":"README.md"}"#.to_owned(),
};

for (call_index, call_id) in [(1, "call-1"), (2, "call-2")] {
if call_index == 2 {
auth_manager.refresh_token_from_authority().await?;
}
registry.tool_lifecycle_contributors()[0]
.on_tool_start(ToolStartInput {
session_store: &session_store,
thread_store,
turn_store: &turn_store,
turn_id: "turn-1",
call_id,
tool_name: &tool_name,
payload: &payload,
conversation_history: Arc::new(TestConversationHistory(Vec::new())),
source: ToolCallSource::Direct,
})
.await;
tokio::time::timeout(Duration::from_secs(5), async {
while progress.latest_scored_tool_call.load(Ordering::Acquire) < call_index {
tokio::task::yield_now().await;
}
})
.await?;
}

assert_eq!(
server
.handshakes()
.iter()
.map(|handshake| handshake.header("authorization"))
.collect::<Vec<_>>(),
vec![
Some("Bearer original".to_owned()),
Some("Bearer original".to_owned()),
Some("Bearer refreshed".to_owned()),
]
);
assert_eq!(
server
.connections()
.iter()
.map(Vec::len)
.collect::<Vec<_>>(),
vec![0, 1, 1]
);
Ok(())
}

struct TestConversationHistory(Vec<ResponseItem>);

impl ConversationHistorySnapshot for TestConversationHistory {
Expand Down
18 changes: 17 additions & 1 deletion codex-rs/ext/guardian-v2/src/sampler.rs
Original file line number Diff line number Diff line change
Expand Up @@ -118,6 +118,7 @@ pub enum LunaSamplerError {
struct PooledConnection {
connection: ResponsesWebsocketConnection,
connected_at: Instant,
auth_changes: Option<tokio::sync::watch::Receiver<u64>>,
}

struct ConnectionLease {
Expand Down Expand Up @@ -178,6 +179,8 @@ impl LunaSampler {
.api_provider()
.await
.map_err(LunaSamplerError::Provider)?;
let auth_manager = self.config.provider.auth_manager();
let auth_changes = auth_manager.map(|manager| manager.auth_change_receiver());
let auth = self
.config
.provider
Expand Down Expand Up @@ -244,10 +247,19 @@ impl LunaSampler {
.await
.map_err(|_| LunaSamplerError::ConnectionTimeout)?
.map_err(LunaSamplerError::Api)?;
if auth_changes
.as_ref()
.is_some_and(|auth| auth.has_changed().unwrap_or(true))
{
return Err(LunaSamplerError::Api(ApiError::Stream(
"authentication changed while connecting".into(),
)));
}

Ok(PooledConnection {
connection,
connected_at: Instant::now(),
auth_changes,
})
}

Expand All @@ -264,7 +276,11 @@ impl LunaSampler {
.pop();
match idle {
Some(connection)
if connection.connected_at.elapsed() < MAX_WEBSOCKET_AGE
if connection
.auth_changes
.as_ref()
.is_none_or(|auth| !auth.has_changed().unwrap_or(true))
&& connection.connected_at.elapsed() < MAX_WEBSOCKET_AGE
&& !connection.connection.is_closed().await =>
{
break connection;
Expand Down
37 changes: 32 additions & 5 deletions codex-rs/ext/guardian-v2/src/sampler_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,9 @@ use codex_http_client::OutboundProxyPolicy;
use codex_login::AgentIdentityAuthPolicy;
use codex_login::AuthManager;
use codex_login::CodexAuth;
use codex_login::ExternalAuth;
use codex_login::ExternalAuthFuture;
use codex_login::ExternalAuthRefreshContext;
use codex_model_provider::create_model_provider;
use codex_model_provider_info::ModelProviderInfo;
use codex_protocol::ResponseItemId;
Expand Down Expand Up @@ -92,6 +95,17 @@ fn sample_request(turn_id: &str) -> LunaSamplingRequest {
}
}

struct RefreshableAuth(std::sync::Mutex<&'static str>);
impl ExternalAuth for RefreshableAuth {
fn resolve(&self) -> ExternalAuthFuture<'_, CodexAuth> {
Box::pin(async { Ok(CodexAuth::from_api_key(*self.0.lock().expect("auth"))) })
}
fn refresh(&self, _: ExternalAuthRefreshContext) -> ExternalAuthFuture<'_, CodexAuth> {
*self.0.lock().expect("auth") = "refreshed";
self.resolve()
}
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn preconnected_sampler_reuses_authenticated_websocket_for_structured_requests() -> Result<()>
{
Expand All @@ -111,13 +125,19 @@ async fn preconnected_sampler_reuses_authenticated_websocket_for_structured_requ
],
];
let idle_server = responses::start_websocket_server(vec![scripted_requests.clone()]).await;
let refreshed =
responses::start_websocket_server(vec![vec![scripted_requests[1].clone()]]).await;
let server = responses::start_websocket_server(vec![scripted_requests]).await;
let base_url = proxy_websocket_servers(&[&idle_server, &server]).await?;
let base_url = proxy_websocket_servers(&[&idle_server, &server, &refreshed]).await?;
let manager = AuthManager::from_auth_for_testing(CodexAuth::from_api_key("test-api-key"));
manager
.set_external_auth(Arc::new(RefreshableAuth(std::sync::Mutex::new(
"test-api-key",
))))
.await?;
let provider = create_model_provider(
ModelProviderInfo::create_openai_provider(Some(base_url)),
Some(AuthManager::from_auth_for_testing(CodexAuth::from_api_key(
"test-api-key",
))),
Some(manager.clone()),
);

let sampler = LunaSampler::connect(LunaSamplerConfig {
Expand Down Expand Up @@ -193,6 +213,7 @@ async fn preconnected_sampler_reuses_authenticated_websocket_for_structured_requ
}
})
.await?;
manager.refresh_token_from_authority().await?;
let second = sampler
.sample(LunaSamplingRequest {
instructions: "Return a risk score.".to_owned(),
Expand All @@ -208,7 +229,13 @@ async fn preconnected_sampler_reuses_authenticated_websocket_for_structured_requ

assert_eq!(first, r#"{"score":0.25}"#);
assert_eq!(second, r#"{"score":0.75}"#);
let requests = server.single_connection();
let mut requests = server.single_connection();
assert_eq!(requests.len(), 1);
assert_eq!(
refreshed.single_handshake().header("authorization"),
Some("Bearer refreshed".to_owned())
);
requests.extend(refreshed.single_connection());
assert_eq!(requests.len(), 2);
assert_eq!(
requests[0].body_json()["input"][2]["content"],
Expand Down
Loading