diff --git a/codex-rs/exec-server/src/client.rs b/codex-rs/exec-server/src/client.rs index 797c286456..e8a6be1bd9 100644 --- a/codex-rs/exec-server/src/client.rs +++ b/codex-rs/exec-server/src/client.rs @@ -9,7 +9,9 @@ use std::sync::atomic::Ordering; use std::time::Duration; use arc_swap::ArcSwap; +use arc_swap::ArcSwapOption; use codex_exec_server_protocol::JSONRPCNotification; +use codex_network_proxy::NetworkPolicyDecider; use futures::FutureExt; use futures::future::BoxFuture; use serde_json::Value; @@ -18,6 +20,7 @@ use tokio::sync::OnceCell; use tokio::sync::Semaphore; use tokio::sync::mpsc; use tokio::sync::watch; +use tokio_util::sync::CancellationToken; use tokio_util::task::AbortOnDropHandle; use tokio::time::timeout; @@ -112,6 +115,7 @@ use crate::protocol::WriteParams; use crate::protocol::WriteResponse; use crate::rpc::RpcCallError; use crate::rpc::RpcClient; +use crate::rpc_server_requests::MAX_IN_FLIGHT_SERVER_CALLS; use codex_http_client::HttpClientFactory; pub(crate) mod http_client; @@ -180,6 +184,14 @@ pub(crate) struct SessionState { ordered_events: StdMutex, recoverable: AtomicBool, next_write_id: AtomicU64, + network_policy_controller: ArcSwapOption, + network_policy_cancelled: CancellationToken, +} + +#[derive(Clone)] +struct NetworkPolicyDecisionController { + decider: Arc, + timeout: Duration, } #[derive(Default)] @@ -223,6 +235,8 @@ struct Inner { http_body_streams_write_lock: Mutex<()>, http_body_stream_byte_budget: Arc, http_body_stream_next_id: AtomicU64, + // Keep admission shared while recovered transports finish older requests. + rpc_inbound_request_slots: Arc, session_id: OnceLock, reconnect_strategy: Option, } @@ -289,6 +303,21 @@ impl Drop for ActiveProcessStart { } } +struct PendingProcessStartSession { + inner: Arc, + process_id: ProcessId, + state: Arc, + armed: bool, +} + +impl Drop for PendingProcessStartSession { + fn drop(&mut self) { + if self.armed { + self.inner.remove_session_if(&self.process_id, &self.state); + } + } +} + type ConnectionResult = Result>; type ConnectionAttempt = OnceCell; @@ -863,8 +892,15 @@ impl ExecServerClient { let active_start = ActiveProcessStart { inner: Arc::clone(&self.inner), }; + let mut pending_start = PendingProcessStartSession { + inner: Arc::clone(&self.inner), + process_id: process_id.clone(), + state: Arc::clone(&state), + armed: true, + }; let client = self.clone(); let (result_tx, result_rx) = tokio::sync::oneshot::channel(); + let (result_received_tx, result_received_rx) = tokio::sync::oneshot::channel(); let process_start_task = async move { let _active_start = active_start; match client @@ -878,7 +914,9 @@ impl ExecServerClient { process_id: process_id.clone(), state: Arc::clone(&state), }; - if result_tx.send(Ok(session)).is_err() { + // Wait for caller receipt so cancellation after send still triggers cleanup. + if result_tx.send(Ok(session)).is_err() || result_received_rx.await.is_err() + { state.recoverable.store(false, Ordering::Release); tokio::spawn(async move { cleanup_process_start(&client, &process_id, &state).await; @@ -902,7 +940,12 @@ impl ExecServerClient { .in_current_span() .with_current_subscriber(), ); - return result_rx.await.map_err(|_| { + let result = result_rx.await; + if matches!(&result, Ok(Ok(_))) { + pending_start.armed = false; + let _ = result_received_tx.send(()); + } + return result.map_err(|_| { ExecServerError::Protocol("process start task stopped unexpectedly".to_string()) })?; } @@ -980,6 +1023,7 @@ impl ExecServerClient { http_body_streams_write_lock: Mutex::new(()), http_body_stream_byte_budget: Arc::new(Semaphore::new(MAX_QUEUED_HTTP_BODY_BYTES)), http_body_stream_next_id: AtomicU64::new(1), + rpc_inbound_request_slots: Arc::new(Semaphore::new(MAX_IN_FLIGHT_SERVER_CALLS)), session_id, reconnect_strategy, }); @@ -1086,6 +1130,8 @@ impl SessionState { ordered_events: StdMutex::new(OrderedSessionEvents::default()), recoverable: AtomicBool::new(recoverable), next_write_id: AtomicU64::new(1), + network_policy_controller: ArcSwapOption::empty(), + network_policy_cancelled: CancellationToken::new(), } } @@ -1395,6 +1441,8 @@ impl Inner { let mut next_sessions = sessions.as_ref().clone(); next_sessions.remove(process_id); self.sessions.store(Arc::new(next_sessions)); + expected.network_policy_cancelled.cancel(); + expected.network_policy_controller.store(None); } fn take_all_sessions(&self) -> HashMap> { @@ -1433,6 +1481,8 @@ fn fail_all_sessions(inner: &Arc, message: String) { let sessions = inner.take_all_sessions(); for (_, session) in sessions { + session.network_policy_cancelled.cancel(); + session.network_policy_controller.store(None); // Sessions synthesize a closed read response and emit a pushed Failed // event. That covers both polling consumers and streaming consumers // such as environment-backed MCP stdio. @@ -2764,4 +2814,6 @@ mod tests { drop(client); server.await.expect("server task should finish"); } + + mod network_policy_tests; } diff --git a/codex-rs/exec-server/src/client/tests/network_policy_tests.rs b/codex-rs/exec-server/src/client/tests/network_policy_tests.rs new file mode 100644 index 0000000000..4445c959ff --- /dev/null +++ b/codex-rs/exec-server/src/client/tests/network_policy_tests.rs @@ -0,0 +1,345 @@ +use std::sync::Arc; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCMessage; +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::JSONRPCResponse; +use codex_exec_server_protocol::RequestId; +use codex_http_client::HttpClientFactory; +use codex_http_client::OutboundProxyPolicy; +use codex_network_proxy::NetworkDecision; +use codex_network_proxy::NetworkPolicyDecider; +use codex_network_proxy::NetworkPolicyRequest; +use codex_utils_path_uri::PathUri; +use pretty_assertions::assert_eq; +use tokio::net::TcpListener; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::time::timeout; + +use super::super::LazyRemoteExecServerClient; +use super::super::NetworkPolicyDecisionController; +use super::accept_websocket; +use super::complete_websocket_initialize; +use super::read_jsonrpc_websocket; +use super::write_jsonrpc_websocket; +use crate::ProcessId; +use crate::client_api::ExecServerTransportParams; +use crate::protocol::EXEC_METHOD; +use crate::protocol::EXEC_TERMINATE_METHOD; +use crate::protocol::ExecParams; +use crate::protocol::ExecServerNetworkPolicyDecision; +use crate::protocol::ExecServerNetworkPolicyRequest; +use crate::protocol::ExecServerNetworkProtocol; +use crate::protocol::NETWORK_POLICY_REQUEST_METHOD; +use crate::protocol::NetworkPolicyRequestParams; +use crate::protocol::NetworkPolicyRequestResponse; +use crate::rpc_server_requests::MAX_IN_FLIGHT_SERVER_CALLS; + +struct PendingDecisionGuard(mpsc::UnboundedSender<()>); + +impl Drop for PendingDecisionGuard { + fn drop(&mut self) { + let _ = self.0.send(()); + } +} + +fn policy_request(request_id: i64, process_id: ProcessId, host: &str) -> JSONRPCMessage { + JSONRPCMessage::Request(JSONRPCRequest { + id: RequestId::Integer(request_id), + method: NETWORK_POLICY_REQUEST_METHOD.to_string(), + params: Some( + serde_json::to_value(NetworkPolicyRequestParams { + process_id, + request: ExecServerNetworkPolicyRequest { + protocol: ExecServerNetworkProtocol::HttpsConnect, + host: host.to_string(), + port: 443, + }, + }) + .expect("policy request should serialize"), + ), + trace: None, + }) +} + +async fn read_decision( + websocket: &mut tokio_tungstenite::WebSocketStream, + request_id: i64, +) -> ExecServerNetworkPolicyDecision { + let JSONRPCMessage::Response(response) = read_jsonrpc_websocket(websocket).await else { + panic!("expected network policy response"); + }; + assert_eq!(response.id, RequestId::Integer(request_id)); + serde_json::from_value::(response.result) + .expect("policy response should deserialize") + .decision +} + +#[tokio::test] +async fn abandoned_process_start_unregisters_and_cleans_up() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!("ws://{}", listener.local_addr().expect("listener address")); + let (start_seen_tx, start_seen_rx) = oneshot::channel(); + let (finish_start_tx, finish_start_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let mut websocket = accept_websocket(&listener).await; + complete_websocket_initialize(&mut websocket, "p", Default::default()).await; + let JSONRPCMessage::Request(start) = read_jsonrpc_websocket(&mut websocket).await else { + panic!("expected process start request"); + }; + assert_eq!(start.method, EXEC_METHOD); + start_seen_tx.send(()).expect("start should be observed"); + finish_start_rx.await.expect("start should be released"); + write_jsonrpc_websocket( + &mut websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id: start.id, + result: serde_json::json!({"processId": "pending-start"}), + }), + ) + .await; + let JSONRPCMessage::Request(terminate) = read_jsonrpc_websocket(&mut websocket).await + else { + panic!("expected process terminate request"); + }; + assert_eq!(terminate.method, EXEC_TERMINATE_METHOD); + write_jsonrpc_websocket( + &mut websocket, + JSONRPCMessage::Response(JSONRPCResponse { + id: terminate.id, + result: serde_json::json!({"running": true}), + }), + ) + .await; + }); + + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::websocket_url(websocket_url, Duration::from_secs(1)), + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .get() + .await + .expect("client should connect"); + let process_id = ProcessId::from("pending-start"); + let start_client = client.clone(); + let start_process_id = process_id.clone(); + let start = tokio::spawn(async move { + start_client + .start_process(ExecParams { + process_id: start_process_id, + argv: vec!["true".to_string()], + cwd: PathUri::from_host_native_path(std::env::current_dir().expect("cwd")) + .expect("cwd URI"), + env_policy: None, + env: Default::default(), + tty: false, + pipe_stdin: false, + arg0: None, + sandbox: None, + enforce_managed_network: false, + managed_network: None, + network_proxy: None, + }) + .await + }); + start_seen_rx.await.expect("start should be observed"); + let state = client + .inner + .get_session(&process_id) + .expect("pending process should be registered"); + let decider: Arc = + Arc::new(|_request: NetworkPolicyRequest| async { NetworkDecision::Allow }); + let decider_weak = Arc::downgrade(&decider); + state + .network_policy_controller + .store(Some(Arc::new(NetworkPolicyDecisionController { + decider, + timeout: Duration::from_secs(30), + }))); + + start.abort(); + assert!(start.await.is_err_and(|error| error.is_cancelled())); + assert!(state.network_policy_cancelled.is_cancelled()); + assert!(client.inner.get_session(&process_id).is_none()); + assert!(decider_weak.upgrade().is_none()); + + finish_start_tx.send(()).expect("start should be released"); + server.await.expect("server task should finish"); +} + +#[tokio::test] +async fn policy_requests_use_process_decider_and_cancel_on_unregister() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let websocket_url = format!( + "ws://{}", + listener.local_addr().expect("listener should have address") + ); + let process_id = ProcessId::from("policy-process"); + let server_process_id = process_id.clone(); + let (ready_tx, ready_rx) = oneshot::channel(); + let (overflow_checked_tx, overflow_checked_rx) = oneshot::channel(); + let (unregistered_tx, unregistered_rx) = oneshot::channel(); + let server = tokio::spawn(async move { + let mut websocket = accept_websocket(&listener).await; + complete_websocket_initialize( + &mut websocket, + "policy-session", + /*expected_resume_session_id*/ None, + ) + .await; + ready_rx.await.expect("process should be registered"); + + for (request_id, host, expected) in [ + (0, "allowed.example", ExecServerNetworkPolicyDecision::Allow), + ( + 2, + "denied.example", + ExecServerNetworkPolicyDecision::Deny { + reason: "blocked".to_string(), + }, + ), + ( + 3, + "invalid host", + ExecServerNetworkPolicyDecision::Deny { + reason: "not_allowed".to_string(), + }, + ), + ] { + write_jsonrpc_websocket( + &mut websocket, + policy_request(request_id, server_process_id.clone(), host), + ) + .await; + assert_eq!(read_decision(&mut websocket, request_id).await, expected); + } + + let first_pending_request_id = 100; + for offset in 0..MAX_IN_FLIGHT_SERVER_CALLS { + write_jsonrpc_websocket( + &mut websocket, + policy_request( + first_pending_request_id + offset as i64, + server_process_id.clone(), + "pending.example", + ), + ) + .await; + } + let overflow_request_id = first_pending_request_id + MAX_IN_FLIGHT_SERVER_CALLS as i64; + write_jsonrpc_websocket( + &mut websocket, + policy_request( + overflow_request_id, + server_process_id.clone(), + "pending.example", + ), + ) + .await; + assert_eq!( + read_decision(&mut websocket, overflow_request_id).await, + ExecServerNetworkPolicyDecision::Deny { + reason: "not_allowed".to_string(), + } + ); + overflow_checked_tx.send(()).expect("overflow observed"); + + unregistered_rx + .await + .expect("process should be unregistered"); + + let post_unregister_request_id = 900; + write_jsonrpc_websocket( + &mut websocket, + policy_request( + post_unregister_request_id, + server_process_id, + "allowed.example", + ), + ) + .await; + assert_eq!( + read_decision(&mut websocket, post_unregister_request_id).await, + ExecServerNetworkPolicyDecision::Deny { + reason: "not_allowed".to_string(), + } + ); + }); + + let client = LazyRemoteExecServerClient::new( + ExecServerTransportParams::WebSocketUrl { + websocket_url, + connect_timeout: Duration::from_secs(1), + initialize_timeout: Duration::from_secs(1), + }, + HttpClientFactory::new(OutboundProxyPolicy::ReqwestDefault), + ) + .get() + .await + .expect("client should connect"); + let session = client + .register_session(&process_id) + .await + .expect("session should register"); + let (started_tx, mut started_rx) = mpsc::unbounded_channel(); + let (dropped_tx, mut dropped_rx) = mpsc::unbounded_channel(); + let decider: Arc = Arc::new(move |request: NetworkPolicyRequest| { + let started_tx = started_tx.clone(); + let dropped_tx = dropped_tx.clone(); + async move { + match request.host.as_str() { + "allowed.example" => NetworkDecision::Allow, + "denied.example" => NetworkDecision::deny("blocked"), + "pending.example" => { + started_tx.send(()).expect("decision should start"); + let _drop_guard = PendingDecisionGuard(dropped_tx); + std::future::pending().await + } + host => panic!("unexpected policy host: {host}"), + } + } + }); + session.state.network_policy_controller.store(Some(Arc::new( + NetworkPolicyDecisionController { + decider, + timeout: Duration::from_secs(30), + }, + ))); + ready_tx.send(()).expect("server should be waiting"); + timeout(Duration::from_secs(5), async { + for _ in 0..MAX_IN_FLIGHT_SERVER_CALLS { + started_rx + .recv() + .await + .expect("pending decision should start"); + } + }) + .await + .expect("pending decisions should start"); + overflow_checked_rx + .await + .expect("overflow should be observed"); + session.unregister().await; + timeout(Duration::from_secs(5), async { + for _ in 0..MAX_IN_FLIGHT_SERVER_CALLS { + dropped_rx + .recv() + .await + .expect("unregistered decision should be dropped"); + } + }) + .await + .expect("unregistered decisions should be cancelled"); + unregistered_tx + .send(()) + .expect("server should verify late responses"); + timeout(Duration::from_secs(2), server) + .await + .expect("policy routing should finish") + .expect("server task should finish"); +} diff --git a/codex-rs/exec-server/src/client_recovery.rs b/codex-rs/exec-server/src/client_recovery.rs index 58c7fdd910..4318840ab5 100644 --- a/codex-rs/exec-server/src/client_recovery.rs +++ b/codex-rs/exec-server/src/client_recovery.rs @@ -5,10 +5,18 @@ use std::sync::Arc; use std::sync::atomic::Ordering; use std::time::Duration; +use codex_network_proxy::NetworkDecision; +use codex_network_proxy::NetworkPolicyDecision; +use codex_network_proxy::NetworkPolicyRequest; +use codex_network_proxy::NetworkProtocol; +use serde_json::Value; use tokio::sync::mpsc; use tokio::time::Instant; use tokio::time::sleep; +use tokio::time::timeout; use tokio::time::timeout_at; +use tokio_util::sync::CancellationToken; +use tracing::debug; use super::ConnectionStatus; use super::ExecServerClient; @@ -25,13 +33,24 @@ use crate::client_transport::ExecServerReconnectStrategy; use crate::process::ExecProcessEvent; use crate::protocol::EXEC_READ_METHOD; use crate::protocol::EXEC_TERMINATE_METHOD; +use crate::protocol::ExecServerNetworkPolicyDecision; +use crate::protocol::ExecServerNetworkProtocol; +use crate::protocol::MAX_NETWORK_POLICY_HOST_BYTES; +use crate::protocol::MAX_NETWORK_POLICY_PROCESS_ID_BYTES; +use crate::protocol::MAX_NETWORK_POLICY_REASON_BYTES; +use crate::protocol::NETWORK_POLICY_REQUEST_METHOD; +use crate::protocol::NetworkPolicyRequestParams; +use crate::protocol::NetworkPolicyRequestResponse; use crate::protocol::ReadParams; use crate::protocol::ReadResponse; use crate::protocol::TerminateParams; use crate::protocol::TerminateResponse; use crate::rpc::RpcClient; use crate::rpc::RpcClientEvent; +use crate::rpc::RpcInboundRequestAdmissionError; use crate::rpc::SESSION_ALREADY_ATTACHED_ERROR_CODE; +use crate::rpc::invalid_params; +use crate::rpc::method_not_found; #[cfg(test)] const SESSION_RECOVERY_TIMEOUT: Duration = Duration::from_millis(500); @@ -42,6 +61,7 @@ const SESSION_RECOVERY_TIMEOUT: Duration = Duration::from_secs(25); const SESSION_RECOVERY_RETRY_INTERVAL: Duration = Duration::from_millis(100); const REGISTRY_RECOVERY_INITIAL_RETRY_INTERVAL: Duration = Duration::from_millis(500); const REGISTRY_RECOVERY_MAX_RETRY_INTERVAL: Duration = Duration::from_secs(5); +const NETWORK_POLICY_DENIAL_REASON: &str = "not_allowed"; impl SessionState { fn last_published_seq(&self) -> u64 { @@ -531,8 +551,12 @@ impl ExecServerClient { mut events_rx: mpsc::Receiver, ) { let inner = Arc::downgrade(&self.inner); + let rpc_inbound_request_slots = Arc::clone(&self.inner.rpc_inbound_request_slots); let rpc_client = Arc::downgrade(rpc_client); + let connection_cancelled = CancellationToken::new(); + let connection_cancel_guard = connection_cancelled.clone().drop_guard(); tokio::spawn(async move { + let _connection_cancel_guard = connection_cancel_guard; while let Some(event) = events_rx.recv().await { let (Some(inner), Some(rpc_client)) = (inner.upgrade(), rpc_client.upgrade()) else { @@ -540,17 +564,191 @@ impl ExecServerClient { }; match event { RpcClientEvent::Request(request) => { - let error = crate::rpc::method_not_found(format!( - "exec-server client does not implement `{}` yet", - request.method - )); - if rpc_client.respond_error(request.id, error).await.is_err() { - inner.request_recovery( - rpc_client, - disconnected_message(/*reason*/ None), - ); - return; + if request.method != NETWORK_POLICY_REQUEST_METHOD { + let error = method_not_found(format!( + "exec-server client does not implement `{}` yet", + request.method + )); + if rpc_client.respond_error(request.id, error).await.is_err() { + inner.request_recovery( + rpc_client, + disconnected_message(/*reason*/ None), + ); + return; + } + continue; } + + let request_guard = match rpc_client + .admit_inbound_request(&request.id, &rpc_inbound_request_slots) + { + Ok(request_guard) => request_guard, + Err(RpcInboundRequestAdmissionError::InvalidRequestId) => { + rpc_client.close_transport().await; + inner.request_recovery( + rpc_client, + "exec-server sent an invalid request ID".to_string(), + ); + return; + } + Err(RpcInboundRequestAdmissionError::DuplicateRequestId) => { + rpc_client.close_transport().await; + inner.request_recovery( + rpc_client, + "exec-server reused an in-flight request ID".to_string(), + ); + return; + } + Err(RpcInboundRequestAdmissionError::AtCapacity) => { + let response = NetworkPolicyRequestResponse { + decision: ExecServerNetworkPolicyDecision::Deny { + reason: NETWORK_POLICY_DENIAL_REASON.to_string(), + }, + }; + if rpc_client.respond(request.id, &response).await.is_err() { + inner.request_recovery( + rpc_client, + disconnected_message(/*reason*/ None), + ); + return; + } + continue; + } + }; + let request_id = request.id; + let params: NetworkPolicyRequestParams = + match serde_json::from_value(request.params.unwrap_or(Value::Null)) { + Ok(params) => params, + Err(_) => { + let error = invalid_params( + "invalid network policy request params".to_string(), + ); + if rpc_client.respond_error(request_id, error).await.is_err() { + inner.request_recovery( + rpc_client, + disconnected_message(/*reason*/ None), + ); + return; + } + continue; + } + }; + let process_id = params.process_id; + let request = params.request; + let process_id_valid = !process_id.is_empty() + && process_id.len() <= MAX_NETWORK_POLICY_PROCESS_ID_BYTES; + let host_valid = !request.host.is_empty() + && request.host.len() <= MAX_NETWORK_POLICY_HOST_BYTES + && !request.host.chars().any(char::is_control) + && !request.host.chars().any(char::is_whitespace); + let session = (process_id_valid && host_valid) + .then(|| inner.get_session(&process_id)) + .flatten(); + let controller = session + .as_ref() + .and_then(|session| session.network_policy_controller.load_full()); + let process_cancelled = session + .as_ref() + .map(|session| session.network_policy_cancelled.clone()); + let expected_session = session.as_ref().map(Arc::downgrade); + let policy_request = + (process_id_valid && host_valid).then_some(NetworkPolicyRequest { + protocol: match request.protocol { + ExecServerNetworkProtocol::Http => NetworkProtocol::Http, + ExecServerNetworkProtocol::HttpsConnect => { + NetworkProtocol::HttpsConnect + } + ExecServerNetworkProtocol::Socks5Tcp => { + NetworkProtocol::Socks5Tcp + } + ExecServerNetworkProtocol::Socks5Udp => { + NetworkProtocol::Socks5Udp + } + }, + host: request.host, + port: request.port, + environment_id: None, + client_addr: None, + method: None, + command: None, + exec_policy_hint: None, + execution_id: None, + }); + let inner = Arc::downgrade(&inner); + let rpc_client = Arc::downgrade(&rpc_client); + let connection_cancelled = connection_cancelled.clone(); + tokio::spawn(async move { + let _request_guard = request_guard; + let decision = match (controller, policy_request, process_cancelled) { + (Some(controller), Some(request), Some(process_cancelled)) => { + // Core's pending-approval guard makes dropping this + // future on process removal or deadline fail closed. + tokio::select! { + biased; + _ = connection_cancelled.cancelled() => return, + _ = process_cancelled.cancelled() => { + NetworkDecision::deny(NETWORK_POLICY_DENIAL_REASON) + } + decision = timeout( + controller.timeout, + controller.decider.decide(request), + ) => decision.unwrap_or_else(|_| { + NetworkDecision::deny(NETWORK_POLICY_DENIAL_REASON) + }), + } + } + (None, _, _) | (_, None, _) | (_, _, None) => { + NetworkDecision::deny(NETWORK_POLICY_DENIAL_REASON) + } + }; + if let Some(expected_session) = expected_session { + let (Some(inner), Some(expected_session)) = + (inner.upgrade(), expected_session.upgrade()) + else { + return; + }; + if !inner + .get_session(&process_id) + .is_some_and(|session| Arc::ptr_eq(&session, &expected_session)) + { + return; + } + } + let Some(rpc_client) = rpc_client.upgrade() else { + return; + }; + let decision = match decision { + NetworkDecision::Allow => ExecServerNetworkPolicyDecision::Allow, + NetworkDecision::Deny { + reason, decision, .. + } if reason.len() <= MAX_NETWORK_POLICY_REASON_BYTES + && !reason.chars().any(char::is_control) => + { + match decision { + NetworkPolicyDecision::Deny => { + ExecServerNetworkPolicyDecision::Deny { reason } + } + NetworkPolicyDecision::Ask => { + ExecServerNetworkPolicyDecision::Ask { reason } + } + } + } + NetworkDecision::Deny { .. } => { + ExecServerNetworkPolicyDecision::Deny { + reason: NETWORK_POLICY_DENIAL_REASON.to_string(), + } + } + }; + if let Err(error) = rpc_client + .respond(request_id, &NetworkPolicyRequestResponse { decision }) + .await + { + debug!( + ?error, + "failed to send network policy decision to exec-server" + ); + } + }); } RpcClientEvent::Notification(notification) => { if let Err(error) = handle_server_notification(&inner, notification).await { diff --git a/codex-rs/exec-server/src/local_process.rs b/codex-rs/exec-server/src/local_process.rs index 96dfe0cee7..a6f795d6ed 100644 --- a/codex-rs/exec-server/src/local_process.rs +++ b/codex-rs/exec-server/src/local_process.rs @@ -263,25 +263,35 @@ impl LocalProcess { params: ExecParams, ) -> Result<(ExecResponse, watch::Sender, ExecProcessEventLog), JSONRPCErrorError> { let process_id = params.process_id.clone(); - let request_policy_decisions = params + let policy_decision_timeout_ms = params .network_proxy .as_ref() - .is_some_and(|launch| launch.proxy.request_policy_decisions); - if request_policy_decisions + .and_then(|launch| launch.policy_decision_timeout_ms); + if policy_decision_timeout_ms == Some(0) { + return Err(invalid_params( + "network policy decision callback timeout must be nonzero".to_string(), + )); + } + if policy_decision_timeout_ms.is_some() && (process_id.is_empty() || process_id.len() > MAX_NETWORK_POLICY_PROCESS_ID_BYTES) { return Err(invalid_params(format!( "callback-enabled process ID must be non-empty and at most {MAX_NETWORK_POLICY_PROCESS_ID_BYTES} bytes" ))); } - let network_policy_shutdown = request_policy_decisions.then(CancellationToken::new); - let network_policy_decider = network_policy_shutdown.as_ref().map(|process_shutdown| { - network_policy_decider( - process_id.clone(), - Arc::clone(&self.inner.requests), - process_shutdown.clone(), - ) - }); + let policy_decision_timeout = policy_decision_timeout_ms.map(Duration::from_millis); + let network_policy_shutdown = policy_decision_timeout.map(|_| CancellationToken::new()); + let network_policy_decider = network_policy_shutdown + .as_ref() + .zip(policy_decision_timeout) + .map(|(process_shutdown, controller_timeout)| { + network_policy_decider( + process_id.clone(), + Arc::clone(&self.inner.requests), + controller_timeout, + process_shutdown.clone(), + ) + }); let prepared = prepare_exec_request( ¶ms, child_env(¶ms), @@ -1161,11 +1171,11 @@ mod tests { #[tokio::test] async fn callback_enabled_start_bounds_process_id_before_proxy_launch() { - let mut proxy_config = + let proxy_config = RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig::default()) .expect("remote proxy config"); - proxy_config.request_policy_decisions = true; - let proxy = RemoteNetworkProxyLaunchConfig::new(proxy_config); + let mut proxy = RemoteNetworkProxyLaunchConfig::new(proxy_config); + proxy.policy_decision_timeout_ms = Some(1_000); let expected = invalid_params(format!( "callback-enabled process ID must be non-empty and at most {MAX_NETWORK_POLICY_PROCESS_ID_BYTES} bytes" )); @@ -1511,6 +1521,7 @@ mod tests { let decider = network_policy_decider( process.process_id.clone(), Arc::clone(&backend.inner.requests), + Duration::from_secs(30), network_policy_shutdown.clone(), ); let config = NetworkProxyConfig { diff --git a/codex-rs/exec-server/src/network_policy_decisions.rs b/codex-rs/exec-server/src/network_policy_decisions.rs index ba21efe292..129ae30c33 100644 --- a/codex-rs/exec-server/src/network_policy_decisions.rs +++ b/codex-rs/exec-server/src/network_policy_decisions.rs @@ -19,14 +19,16 @@ use crate::protocol::NetworkPolicyRequestParams; use crate::protocol::NetworkPolicyRequestResponse; use crate::rpc_server_requests::RpcServerRequestSender; -// Leave transport overhead outside the client-side 95-second decision window. -const NETWORK_POLICY_REQUEST_TIMEOUT: Duration = Duration::from_secs(100); +const NETWORK_POLICY_TRANSPORT_TIMEOUT_MARGIN: Duration = Duration::from_secs(5); pub(crate) fn network_policy_decider( process_id: ProcessId, requests: Arc>>, + controller_timeout: Duration, process_shutdown: CancellationToken, ) -> Arc { + let request_timeout = + controller_timeout.saturating_add(NETWORK_POLICY_TRANSPORT_TIMEOUT_MARGIN); Arc::new(move |request: NetworkPolicyRequest| { let process_id = process_id.clone(); let requests = Arc::clone(&requests); @@ -66,7 +68,7 @@ pub(crate) fn network_policy_decider( response = requests.call_with_timeout::<_, NetworkPolicyRequestResponse>( NETWORK_POLICY_REQUEST_METHOD, ¶ms, - NETWORK_POLICY_REQUEST_TIMEOUT, + request_timeout, ) => response .map(|response| match response.decision { ExecServerNetworkPolicyDecision::Allow => NetworkDecision::Allow, diff --git a/codex-rs/exec-server/src/network_policy_decisions_tests.rs b/codex-rs/exec-server/src/network_policy_decisions_tests.rs index 5500115396..ace0c3e1f3 100644 --- a/codex-rs/exec-server/src/network_policy_decisions_tests.rs +++ b/codex-rs/exec-server/src/network_policy_decisions_tests.rs @@ -20,6 +20,7 @@ use crate::rpc::RpcServerOutboundMessage; struct DeciderHarness { requests: RpcServerRequestSender, outgoing: mpsc::Receiver, + controller_timeout: Duration, process_shutdown: CancellationToken, } @@ -29,6 +30,7 @@ impl DeciderHarness { Self { requests: RpcServerRequestSender::new(outgoing_tx), outgoing, + controller_timeout: Duration::from_secs(60), process_shutdown: CancellationToken::new(), } } @@ -37,6 +39,7 @@ impl DeciderHarness { let decider = network_policy_decider( ProcessId::from("process"), Arc::new(RwLock::new(Some(self.requests.clone()))), + self.controller_timeout, self.process_shutdown.clone(), ); let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { @@ -204,12 +207,13 @@ async fn process_exit_and_disconnect_fail_closed() { } #[tokio::test(start_paused = true)] -async fn decision_timeout_fails_closed() { +async fn configured_decision_timeout_fails_closed() { let mut harness = DeciderHarness::new(); + harness.controller_timeout = Duration::from_secs(17); let decision = harness.request("timeout.example.com"); harness.next_request().await; - tokio::time::advance(Duration::from_secs(99)).await; + tokio::time::advance(Duration::from_secs(21)).await; assert!(!decision.is_finished()); tokio::time::advance(Duration::from_secs(1)).await; diff --git a/codex-rs/exec-server/src/rpc.rs b/codex-rs/exec-server/src/rpc.rs index dd565b3df6..363d41bdd0 100644 --- a/codex-rs/exec-server/src/rpc.rs +++ b/codex-rs/exec-server/src/rpc.rs @@ -1,7 +1,9 @@ use std::collections::HashMap; +use std::collections::HashSet; use std::future::Future; use std::pin::Pin; use std::sync::Arc; +use std::sync::Mutex as StdMutex; use std::sync::atomic::AtomicBool; use std::sync::atomic::AtomicI64; use std::sync::atomic::Ordering; @@ -18,6 +20,7 @@ use serde::Serialize; use serde::de::DeserializeOwned; use serde_json::Value; use tokio::sync::Mutex; +use tokio::sync::OwnedSemaphorePermit; use tokio::sync::Semaphore; use tokio::sync::SemaphorePermit; use tokio::sync::mpsc; @@ -69,6 +72,27 @@ pub(crate) enum RpcClientEvent { Disconnected { reason: Option }, } +pub(crate) enum RpcInboundRequestAdmissionError { + InvalidRequestId, + DuplicateRequestId, + AtCapacity, +} + +pub(crate) struct RpcInboundRequestGuard { + request_id: RequestId, + request_ids: Arc>>, + _call_slot: OwnedSemaphorePermit, +} + +impl Drop for RpcInboundRequestGuard { + fn drop(&mut self) { + self.request_ids + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .remove(&self.request_id); + } +} + #[derive(Debug, Clone, PartialEq)] pub(crate) enum RpcServerOutboundMessage { Request(JSONRPCRequest), @@ -253,6 +277,7 @@ where pub(crate) struct RpcClient { write_tx: mpsc::Sender, pending: Arc>>, + inbound_request_ids: Arc>>, // Shared transport state from `JsonRpcConnection`. Calls use this to fail // immediately when the socket closes, even if no JSON-RPC error response // can be delivered for their request id. @@ -320,6 +345,7 @@ impl RpcClient { Self { write_tx, pending, + inbound_request_ids: Arc::new(StdMutex::new(HashSet::new())), disconnected_rx, closed, shared_call_slots: Semaphore::new(MAX_IN_FLIGHT_REGULAR_CALLS), @@ -333,6 +359,44 @@ impl RpcClient { ) } + pub(crate) fn admit_inbound_request( + &self, + request_id: &RequestId, + call_slots: &Arc, + ) -> Result { + let request_id = match request_id { + RequestId::Integer(request_id) if *request_id >= 0 => request_id, + RequestId::Integer(_) | RequestId::String(_) => { + return Err(RpcInboundRequestAdmissionError::InvalidRequestId); + } + }; + let request_id = RequestId::Integer(*request_id); + { + let mut request_ids = self + .inbound_request_ids + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if !request_ids.insert(request_id.clone()) { + return Err(RpcInboundRequestAdmissionError::DuplicateRequestId); + } + } + let call_slot = match Arc::clone(call_slots).try_acquire_owned() { + Ok(call_slot) => call_slot, + Err(_) => { + self.inbound_request_ids + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .remove(&request_id); + return Err(RpcInboundRequestAdmissionError::AtCapacity); + } + }; + Ok(RpcInboundRequestGuard { + request_id, + request_ids: Arc::clone(&self.inbound_request_ids), + _call_slot: call_slot, + }) + } + pub(crate) async fn notify( &self, method: &str, @@ -351,6 +415,24 @@ impl RpcClient { .map_err(|_| RpcCallError::Closed) } + pub(crate) async fn respond( + &self, + request_id: RequestId, + result: &T, + ) -> Result<(), RpcCallError> { + let result = serde_json::to_value(result).map_err(RpcCallError::Json)?; + if self.closed.load(Ordering::Acquire) || *self.disconnected_rx.borrow() { + return Err(RpcCallError::Closed); + } + self.write_tx + .send(JSONRPCMessage::Response(JSONRPCResponse { + id: request_id, + result, + })) + .await + .map_err(|_| RpcCallError::Closed) + } + pub(crate) async fn respond_error( &self, request_id: RequestId, diff --git a/codex-rs/exec-server/src/rpc_server_requests.rs b/codex-rs/exec-server/src/rpc_server_requests.rs index f718b27266..b7d158802b 100644 --- a/codex-rs/exec-server/src/rpc_server_requests.rs +++ b/codex-rs/exec-server/src/rpc_server_requests.rs @@ -19,7 +19,7 @@ use tokio_util::sync::CancellationToken; use crate::rpc::RpcCallError; use crate::rpc::RpcServerOutboundMessage; -const MAX_IN_FLIGHT_SERVER_CALLS: usize = 256; +pub(crate) const MAX_IN_FLIGHT_SERVER_CALLS: usize = 256; type PendingRequest = oneshot::Sender>; diff --git a/codex-rs/network-proxy/src/proxy.rs b/codex-rs/network-proxy/src/proxy.rs index 610187bdd4..e1a8ce443b 100644 --- a/codex-rs/network-proxy/src/proxy.rs +++ b/codex-rs/network-proxy/src/proxy.rs @@ -831,6 +831,7 @@ impl NetworkProxy { audit_metadata: self.state.audit_metadata().clone(), environment_id, execution_id, + policy_decision_timeout_ms: None, }) } diff --git a/codex-rs/network-proxy/src/remote_config.rs b/codex-rs/network-proxy/src/remote_config.rs index f405e0ede4..3bb5101a61 100644 --- a/codex-rs/network-proxy/src/remote_config.rs +++ b/codex-rs/network-proxy/src/remote_config.rs @@ -24,6 +24,9 @@ pub struct RemoteNetworkProxyLaunchConfig { pub environment_id: Option, #[serde(default)] pub execution_id: Option, + /// Controller-side policy decision budget. The executor adds transport overhead. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub policy_decision_timeout_ms: Option, } impl RemoteNetworkProxyLaunchConfig { @@ -33,6 +36,7 @@ impl RemoteNetworkProxyLaunchConfig { audit_metadata: NetworkProxyAuditMetadata::default(), environment_id: None, execution_id: None, + policy_decision_timeout_ms: None, } } @@ -66,9 +70,6 @@ pub struct RemoteNetworkProxyConfig { pub domains: Option, pub unix_sockets: Option, pub allow_local_binding: bool, - /// Whether the executor sends domain policy decisions back to the client. - #[serde(default, skip_serializing_if = "std::ops::Not::not")] - pub request_policy_decisions: bool, } impl RemoteNetworkProxyConfig { @@ -91,7 +92,6 @@ impl RemoteNetworkProxyConfig { domains: config.domains.clone(), unix_sockets: config.unix_sockets.clone(), allow_local_binding: config.allow_local_binding, - request_policy_decisions: false, }) } diff --git a/codex-rs/network-proxy/src/remote_config_tests.rs b/codex-rs/network-proxy/src/remote_config_tests.rs index 1b7ea713f3..5094acfaad 100644 --- a/codex-rs/network-proxy/src/remote_config_tests.rs +++ b/codex-rs/network-proxy/src/remote_config_tests.rs @@ -111,6 +111,7 @@ fn launch_config_materializes_audit_and_execution_attribution() { audit_metadata: audit_metadata.clone(), environment_id: Some("remote".to_string()), execution_id: Some("execution-1".to_string()), + policy_decision_timeout_ms: None, }) .expect("remote launch state"); @@ -120,24 +121,18 @@ fn launch_config_materializes_audit_and_execution_attribution() { } #[test] -fn policy_decision_callback_opt_in_is_backward_compatible() { - let mut config = - RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig::default()) - .expect("supported remote config"); - let legacy = serde_json::to_value(&config).expect("serialize legacy config"); - assert_eq!(legacy.get("requestPolicyDecisions"), None); - assert!( - !serde_json::from_value::(legacy) - .expect("deserialize legacy remote config") - .request_policy_decisions - ); - - config.request_policy_decisions = true; - let enabled = serde_json::to_value(&config).expect("serialize callback-enabled config"); - assert_eq!(enabled["requestPolicyDecisions"], true); +fn policy_decision_callback_timeout_round_trips() { + let config = RemoteNetworkProxyConfig::from_effective_config(&NetworkProxyConfig::default()) + .expect("supported remote config"); + let mut launch = RemoteNetworkProxyLaunchConfig::new(config); + let without_timeout = serde_json::to_value(&launch).expect("serialize launch config"); + assert_eq!(without_timeout.get("policyDecisionTimeoutMs"), None); + launch.policy_decision_timeout_ms = Some(900_000); + let with_timeout = serde_json::to_value(&launch).expect("serialize launch timeout"); + assert_eq!(with_timeout["policyDecisionTimeoutMs"], 900_000); assert_eq!( - serde_json::from_value::(enabled) - .expect("deserialize callback-enabled config"), - config + serde_json::from_value::(with_timeout) + .expect("deserialize launch timeout"), + launch ); } diff --git a/codex-rs/network-proxy/src/runtime.rs b/codex-rs/network-proxy/src/runtime.rs index ed2ee513aa..10f1367fa9 100644 --- a/codex-rs/network-proxy/src/runtime.rs +++ b/codex-rs/network-proxy/src/runtime.rs @@ -284,6 +284,7 @@ impl NetworkProxyState { audit_metadata, environment_id, execution_id, + policy_decision_timeout_ms: _, } = launch; anyhow::ensure!( proxy.enabled,