diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 63234ecaff..c40c2ecaed 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -3306,6 +3306,7 @@ dependencies = [ "codex-api", "codex-core", "codex-extension-api", + "codex-features", "codex-history", "codex-http-client", "codex-login", diff --git a/codex-rs/ext/guardian-v2/Cargo.toml b/codex-rs/ext/guardian-v2/Cargo.toml index 2cbd2d9dca..15e9d2cf29 100644 --- a/codex-rs/ext/guardian-v2/Cargo.toml +++ b/codex-rs/ext/guardian-v2/Cargo.toml @@ -16,6 +16,7 @@ workspace = true codex-api = { workspace = true } codex-core = { workspace = true } codex-extension-api = { workspace = true } +codex-features = { workspace = true } codex-history = { workspace = true } codex-http-client = { workspace = true } codex-login = { workspace = true } diff --git a/codex-rs/ext/guardian-v2/src/extension.rs b/codex-rs/ext/guardian-v2/src/extension.rs index 49f1d0ae15..16b203c55d 100644 --- a/codex-rs/ext/guardian-v2/src/extension.rs +++ b/codex-rs/ext/guardian-v2/src/extension.rs @@ -2,20 +2,30 @@ use std::sync::Arc; use std::sync::Weak; use codex_core::ThreadManager; +use codex_core::config::Config; use codex_extension_api::ExtensionEventSink; +use codex_extension_api::ExtensionFuture; use codex_extension_api::ExtensionRegistryBuilder; use codex_extension_api::ExtensionWarning; +use codex_extension_api::ThreadLifecycleContributor; +use codex_extension_api::ThreadOriginator; +use codex_extension_api::ThreadStartInput; use codex_extension_api::ToolLifecycleContributor; use codex_extension_api::ToolLifecycleFuture; use codex_extension_api::ToolPayload; use codex_extension_api::ToolStartInput; +use codex_features::Feature; use codex_history::RolloutItem; +use codex_login::AgentIdentityAuthPolicy; +use codex_login::AuthManager; +use codex_model_provider::create_model_provider; use codex_protocol::ThreadId; use codex_protocol::openai_models::ReasoningEffort; use codex_protocol::security_risk::SecurityRiskScore; use serde_json::json; use crate::LunaSampler; +use crate::LunaSamplerConfig; use crate::LunaSamplingRequest; use crate::transcript::TranscriptConfig; @@ -26,14 +36,63 @@ Return an action_risk score from 0.0 (safe and authorized) to 1.0 (dangerous or #[derive(Clone)] struct GuardianV2Extension { - sampler: Arc, + auth_manager: Arc, event_sink: Arc, thread_manager: Weak, } +impl ThreadLifecycleContributor for GuardianV2Extension { + fn on_thread_start<'a>( + &'a self, + input: ThreadStartInput<'a, Config>, + ) -> ExtensionFuture<'a, ()> { + Box::pin(async move { + if !input.config.features.enabled(Feature::GuardianV2) { + return; + } + + let thread_id = input.thread_store.level_id().to_string(); + let sampler = LunaSampler::connect(LunaSamplerConfig { + provider: create_model_provider( + input.config.model_provider.clone(), + Some(Arc::clone(&self.auth_manager)), + ), + http_client_factory: input.config.http_client_factory(), + agent_identity_policy: if input.config.features.enabled(Feature::UseAgentIdentity) { + AgentIdentityAuthPolicy::ChatGptAuth + } else { + AgentIdentityAuthPolicy::JwtOnly + }, + session_source: input.session_source.clone(), + session_id: input.session_store.level_id().to_string(), + thread_id: thread_id.clone(), + originator: input + .thread_store + .get::() + .map(|originator| originator.0.clone()), + service_tier: input.config.service_tier.clone(), + }) + .await; + + match sampler { + Ok(sampler) => { + input.thread_store.insert(sampler); + } + Err(error) => self.event_sink.emit_warning(ExtensionWarning { + thread_id, + turn_id: None, + message: format!("Guardian V2 Luna initialization failed: {error}"), + }), + } + }) + } +} + impl ToolLifecycleContributor for GuardianV2Extension { fn on_tool_start<'a>(&'a self, input: ToolStartInput<'a>) -> ToolLifecycleFuture<'a> { - let sampler = Arc::clone(&self.sampler); + let Some(sampler) = input.thread_store.get::() else { + return Box::pin(std::future::ready(())); + }; let event_sink = Arc::clone(&self.event_sink); let thread_manager = self.thread_manager.clone(); let thread_id = input.thread_store.level_id().to_owned(); @@ -154,17 +213,19 @@ impl ToolLifecycleContributor for GuardianV2Extension { } } -/// Installs Guardian V2 tool classification over a caller-owned Luna sampler. -pub fn install( - registry: &mut ExtensionRegistryBuilder, - sampler: Arc, +/// Installs feature-gated Guardian V2 tool classification for each thread. +pub fn install( + registry: &mut ExtensionRegistryBuilder, + auth_manager: Arc, thread_manager: Weak, ) { - registry.tool_lifecycle_contributor(Arc::new(GuardianV2Extension { - sampler, + let extension = Arc::new(GuardianV2Extension { + auth_manager, event_sink: registry.event_sink(), thread_manager, - })); + }); + registry.thread_lifecycle_contributor(extension.clone()); + registry.tool_lifecycle_contributor(extension); } #[cfg(test)] diff --git a/codex-rs/ext/guardian-v2/src/extension_tests.rs b/codex-rs/ext/guardian-v2/src/extension_tests.rs index 9b0bbc675e..504cb3a870 100644 --- a/codex-rs/ext/guardian-v2/src/extension_tests.rs +++ b/codex-rs/ext/guardian-v2/src/extension_tests.rs @@ -6,17 +6,15 @@ use codex_extension_api::ConversationHistorySnapshot; use codex_extension_api::ExtensionData; use codex_extension_api::ExtensionRegistryBuilder; use codex_extension_api::ResponseItem; +use codex_extension_api::ThreadStartInput; use codex_extension_api::ToolCallSource; use codex_extension_api::ToolName; use codex_extension_api::ToolPayload; use codex_extension_api::ToolStartInput; +use codex_features::Feature; use codex_history::RolloutItem; -use codex_http_client::HttpClientFactory; -use codex_http_client::OutboundProxyPolicy; -use codex_login::AgentIdentityAuthPolicy; use codex_login::AuthManager; use codex_login::CodexAuth; -use codex_model_provider::create_model_provider; use codex_model_provider_info::ModelProviderInfo; use codex_protocol::models::ContentItem; use codex_protocol::models::FunctionCallOutputPayload; @@ -31,9 +29,6 @@ use core_test_support::test_codex::test_codex; use pretty_assertions::assert_eq; use serde_json::json; -use crate::LunaSampler; -use crate::LunaSamplerConfig; - struct TestConversationHistory(Vec); impl ConversationHistorySnapshot for TestConversationHistory { @@ -58,31 +53,31 @@ async fn contributor_samples_tool_calls_with_the_existing_luna_pool() -> Result< "http://{}/v1", server.uri().trim_start_matches("ws://") ))); - let sampler = LunaSampler::connect(LunaSamplerConfig { - provider: create_model_provider( - provider_info, - Some(AuthManager::from_auth_for_testing(CodexAuth::from_api_key( - "test-api-key", - ))), - ), - http_client_factory: HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), - agent_identity_policy: AgentIdentityAuthPolicy::JwtOnly, - session_source: SessionSource::Exec, - session_id: "session-1".to_owned(), - thread_id: thread_id.to_string(), - originator: None, - service_tier: None, - }) - .await?; - let mut builder = ExtensionRegistryBuilder::<()>::new(); + let auth_manager = AuthManager::from_auth_for_testing(CodexAuth::from_api_key("test-api-key")); + let mut config = test.config.clone(); + config.model_provider = provider_info; + config.features.enable(Feature::GuardianV2)?; + let mut builder = ExtensionRegistryBuilder::new(); crate::install( &mut builder, - Arc::new(sampler), + auth_manager, 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 turn_store = ExtensionData::new("turn-1"); let tool_name = ToolName::plain("read_file"); let tool_payload = ToolPayload::Function { diff --git a/codex-rs/ext/guardian-v2/src/sampler_tests.rs b/codex-rs/ext/guardian-v2/src/sampler_tests.rs index 72c9f3e8d2..0dfaec47dc 100644 --- a/codex-rs/ext/guardian-v2/src/sampler_tests.rs +++ b/codex-rs/ext/guardian-v2/src/sampler_tests.rs @@ -122,6 +122,12 @@ async fn preconnected_sampler_reuses_authenticated_websocket_for_structured_requ .await?; let handshake = server.single_handshake(); + tokio::time::timeout(Duration::from_secs(2), async { + while server.connections().is_empty() { + tokio::task::yield_now().await; + } + }) + .await?; assert!(server.single_connection().is_empty()); assert_eq!( handshake.header("authorization"),