diff --git a/codex-rs/exec-server/src/client_recovery.rs b/codex-rs/exec-server/src/client_recovery.rs index 553c6928c755..8b670aeb25c4 100644 --- a/codex-rs/exec-server/src/client_recovery.rs +++ b/codex-rs/exec-server/src/client_recovery.rs @@ -539,6 +539,19 @@ impl ExecServerClient { return; }; 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; + } + } RpcClientEvent::Notification(notification) => { if let Err(error) = handle_server_notification(&inner, notification).await { rpc_client.close_transport().await; diff --git a/codex-rs/exec-server/src/lib.rs b/codex-rs/exec-server/src/lib.rs index 920171de7e3e..803302b37339 100644 --- a/codex-rs/exec-server/src/lib.rs +++ b/codex-rs/exec-server/src/lib.rs @@ -15,6 +15,7 @@ mod fs_helper_main; mod fs_sandbox; mod local_file_system; mod local_process; +mod network_policy_decisions; mod noise_channel; mod noise_relay; mod process; @@ -27,6 +28,7 @@ mod remote_file_system; mod remote_process; mod resolved_capability; mod rpc; +mod rpc_server_requests; mod runtime_paths; mod sandboxed_file_system; mod server; diff --git a/codex-rs/exec-server/src/local_process.rs b/codex-rs/exec-server/src/local_process.rs index c4687ec1c058..96dfe0cee703 100644 --- a/codex-rs/exec-server/src/local_process.rs +++ b/codex-rs/exec-server/src/local_process.rs @@ -22,6 +22,7 @@ use tokio::sync::Mutex; use tokio::sync::Notify; use tokio::sync::mpsc; use tokio::sync::watch; +use tokio_util::sync::CancellationToken; use crate::ExecBackend; use crate::ExecBackendFuture; @@ -33,6 +34,7 @@ use crate::ExecServerError; use crate::ExecServerRuntimePaths; use crate::ProcessId; use crate::StartedExecProcess; +use crate::network_policy_decisions::network_policy_decider; use crate::process::ExecProcessEventLog; use crate::process_sandbox::prepare_exec_request; use crate::protocol::EXEC_CLOSED_METHOD; @@ -43,6 +45,7 @@ use crate::protocol::ExecOutputDeltaNotification; use crate::protocol::ExecOutputStream; use crate::protocol::ExecParams; use crate::protocol::ExecResponse; +use crate::protocol::MAX_NETWORK_POLICY_PROCESS_ID_BYTES; use crate::protocol::ProcessOutputChunk; use crate::protocol::ProcessSignal; use crate::protocol::ReadParams; @@ -59,6 +62,7 @@ use crate::rpc::RpcServerOutboundMessage; use crate::rpc::internal_error; use crate::rpc::invalid_params; use crate::rpc::invalid_request; +use crate::rpc_server_requests::RpcServerRequestSender; use crate::telemetry::ExecServerTelemetry; use crate::telemetry::ProcessMetricGuard; @@ -101,6 +105,7 @@ struct RunningProcess { sandbox: SandboxType, sandbox_denied: bool, network_proxy_handle: Option, + network_policy_shutdown: Option, } /// Bounded cache of stdin write ids that have already been accepted for one process. @@ -143,6 +148,7 @@ enum ProcessEntry { struct Inner { notifications: std::sync::RwLock>, + requests: Arc>>, processes: Mutex>, telemetry: ExecServerTelemetry, } @@ -195,9 +201,11 @@ impl LocalProcess { telemetry: ExecServerTelemetry, runtime_paths: Option, ) -> Self { + let requests = notifications.request_sender(); Self { inner: Arc::new(Inner { notifications: std::sync::RwLock::new(Some(notifications)), + requests: Arc::new(std::sync::RwLock::new(Some(requests))), processes: Mutex::new(HashMap::new()), telemetry, }), @@ -217,6 +225,9 @@ impl LocalProcess { .collect::>() }; for mut process in remaining { + if let Some(network_policy_shutdown) = process.network_policy_shutdown.take() { + network_policy_shutdown.cancel(); + } if let Some(metrics) = process.metrics.take() { metrics.finish("terminated"); } @@ -225,12 +236,26 @@ impl LocalProcess { } pub(crate) fn set_notification_sender(&self, notifications: Option) { + let requests = notifications + .as_ref() + .map(RpcNotificationSender::request_sender); let mut notification_sender = self .inner .notifications .write() .unwrap_or_else(std::sync::PoisonError::into_inner); *notification_sender = notifications; + let previous_requests = std::mem::replace( + &mut *self + .inner + .requests + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner), + requests, + ); + if let Some(previous_requests) = previous_requests { + previous_requests.close(); + } } async fn start_process( @@ -238,8 +263,32 @@ impl LocalProcess { params: ExecParams, ) -> Result<(ExecResponse, watch::Sender, ExecProcessEventLog), JSONRPCErrorError> { let process_id = params.process_id.clone(); - let prepared = - prepare_exec_request(¶ms, child_env(¶ms), self.runtime_paths.as_ref()).await?; + let request_policy_decisions = params + .network_proxy + .as_ref() + .is_some_and(|launch| launch.proxy.request_policy_decisions); + if request_policy_decisions + && (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 prepared = prepare_exec_request( + ¶ms, + child_env(¶ms), + self.runtime_paths.as_ref(), + network_policy_decider, + ) + .await?; if prepared.command.is_empty() { return Err(invalid_params("argv must not be empty".to_string())); } @@ -325,6 +374,7 @@ impl LocalProcess { sandbox: prepared.sandbox, sandbox_denied: false, network_proxy_handle: prepared.network_proxy_handle, + network_policy_shutdown, })), ); } @@ -541,6 +591,9 @@ impl LocalProcess { let mut process_map = self.inner.processes.lock().await; match process_map.get_mut(¶ms.process_id) { Some(ProcessEntry::Running(process)) => { + if let Some(network_policy_shutdown) = &process.network_policy_shutdown { + network_policy_shutdown.cancel(); + } if process.exit_code.is_some() { return Ok(TerminateResponse { running: false }); } @@ -919,6 +972,9 @@ async fn maybe_emit_closed(process_id: ProcessId, inner: Arc) { } process.closed = true; + if let Some(network_policy_shutdown) = process.network_policy_shutdown.take() { + network_policy_shutdown.cancel(); + } let seq = process.next_seq; process.next_seq += 1; let _ = process.wake_tx.send(seq); @@ -989,9 +1045,22 @@ mod tests { use opentelemetry_sdk::metrics::data::AggregatedMetrics; use opentelemetry_sdk::metrics::data::MetricData; use pretty_assertions::assert_eq; + #[cfg(not(target_os = "windows"))] + use tokio::io::AsyncReadExt; + #[cfg(not(target_os = "windows"))] + use tokio::io::AsyncWriteExt; use tokio::sync::oneshot; use tokio::time::timeout; + #[cfg(not(target_os = "windows"))] + use crate::protocol::ExecServerNetworkPolicyDecision; + #[cfg(not(target_os = "windows"))] + use crate::protocol::NETWORK_POLICY_REQUEST_METHOD; + #[cfg(not(target_os = "windows"))] + use crate::protocol::NetworkPolicyRequestParams; + #[cfg(not(target_os = "windows"))] + use crate::protocol::NetworkPolicyRequestResponse; + fn test_exec_params(env: HashMap) -> ExecParams { ExecParams { process_id: ProcessId::from("env-test"), @@ -1090,6 +1159,68 @@ mod tests { assert_eq!(error, expected); } + #[tokio::test] + async fn callback_enabled_start_bounds_process_id_before_proxy_launch() { + let mut 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 expected = invalid_params(format!( + "callback-enabled process ID must be non-empty and at most {MAX_NETWORK_POLICY_PROCESS_ID_BYTES} bytes" + )); + + for process_id in [ + String::new(), + "p".repeat(MAX_NETWORK_POLICY_PROCESS_ID_BYTES + 1), + ] { + let mut params = test_exec_params(HashMap::new()); + params.process_id = ProcessId::from(process_id); + params.network_proxy = Some(proxy.clone()); + let error = LocalProcess::default() + .start_process(params) + .await + .err() + .expect("invalid callback process ID should be rejected"); + + assert_eq!(error, expected); + } + + let mut boundary = test_exec_params(HashMap::new()); + boundary.process_id = ProcessId::from("p".repeat(MAX_NETWORK_POLICY_PROCESS_ID_BYTES)); + boundary.network_proxy = Some(proxy); + boundary.argv.clear(); + let error = LocalProcess::default() + .start_process(boundary) + .await + .err() + .expect("valid boundary process ID should proceed to process preparation"); + #[cfg(not(target_os = "windows"))] + assert!( + error + .message + .contains("executor-local network proxy launch requires an enabled proxy") + ); + #[cfg(target_os = "windows")] + assert_eq!(error, invalid_params("argv must not be empty".to_string())); + + for process_id in [ + String::new(), + "p".repeat(MAX_NETWORK_POLICY_PROCESS_ID_BYTES + 1), + ] { + let mut ordinary = test_exec_params(HashMap::new()); + ordinary.process_id = ProcessId::from(process_id); + ordinary.argv.clear(); + let error = LocalProcess::default() + .start_process(ordinary) + .await + .err() + .expect("empty argv should be rejected after ID validation"); + + assert_eq!(error, invalid_params("argv must not be empty".to_string())); + } + } + #[test] fn child_env_defaults_to_exact_env() { let params = test_exec_params(HashMap::from([("ONLY_THIS".to_string(), "1".to_string())])); @@ -1158,6 +1289,43 @@ mod tests { assert_finished_process_result(metrics, &exporter, "terminated"); } + #[tokio::test] + async fn termination_request_after_exit_cancels_network_policy_decisions() { + let backend = LocalProcess::default(); + let mut process = spawn_test_process(&backend, "terminate-after-exit").await; + let network_policy_shutdown = CancellationToken::new(); + { + let mut processes = backend.inner.processes.lock().await; + let Some(ProcessEntry::Running(running)) = processes.get_mut(&process.process_id) + else { + panic!("test process should be running"); + }; + running.network_policy_shutdown = Some(network_policy_shutdown.clone()); + } + + process.exit(/*exit_code*/ 0); + let response = + read_process_until_change(&backend, &process.process_id, /*after_seq*/ None).await; + assert!(response.exited); + assert!(!response.closed); + assert!(!network_policy_shutdown.is_cancelled()); + assert_eq!( + backend + .terminate_process(TerminateParams { + process_id: process.process_id.clone(), + }) + .await + .expect("terminate exited process"), + TerminateResponse { running: false }, + ); + assert!(network_policy_shutdown.is_cancelled()); + + drop(process.stdout_tx); + drop(process.stderr_tx); + let _ = read_process_until_closed(&backend, &process.process_id).await; + backend.shutdown().await; + } + #[tokio::test] async fn shutdown_before_exit_records_terminated() { let (backend, metrics, exporter) = telemetry_backend(); @@ -1328,6 +1496,23 @@ mod tests { async fn exited_process_keeps_network_proxy_until_inherited_streams_close() { let backend = LocalProcess::default(); let mut process = spawn_test_process(&backend, "proc-background-child").await; + let (outgoing_tx, outgoing_rx) = mpsc::channel(NOTIFICATION_CHANNEL_CAPACITY); + #[cfg(not(target_os = "windows"))] + let mut outgoing_rx = outgoing_rx; + #[cfg(target_os = "windows")] + let _outgoing_rx = outgoing_rx; + let requests = RpcNotificationSender::new(outgoing_tx).request_sender(); + *backend + .inner + .requests + .write() + .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(requests.clone()); + let network_policy_shutdown = CancellationToken::new(); + let decider = network_policy_decider( + process.process_id.clone(), + Arc::clone(&backend.inner.requests), + network_policy_shutdown.clone(), + ); let config = NetworkProxyConfig { enabled: true, ..Default::default() @@ -1340,6 +1525,7 @@ mod tests { .expect("build network proxy state"); let proxy = NetworkProxy::builder() .state(Arc::new(state)) + .policy_decider_arc(decider) .build() .await .expect("build network proxy"); @@ -1362,6 +1548,7 @@ mod tests { panic!("test process should be running"); }; running.network_proxy_handle = Some(handle); + running.network_policy_shutdown = Some(network_policy_shutdown.clone()); } process.exit(/*exit_code*/ 0); @@ -1369,11 +1556,51 @@ mod tests { read_process_until_change(&backend, &process.process_id, /*after_seq*/ None).await; assert!(exit_response.exited); assert!(!exit_response.closed); - tokio::net::TcpStream::connect(proxy_addr) + assert!(!network_policy_shutdown.is_cancelled()); + let stream = tokio::net::TcpStream::connect(proxy_addr) .await .expect("proxy should remain available to a child holding inherited output streams"); #[cfg(target_os = "windows")] - assert!(proxy.network_proxy_restricting_sid(None).is_some()); + { + assert!(proxy.network_proxy_restricting_sid(None).is_some()); + drop(stream); + } + #[cfg(not(target_os = "windows"))] + { + let mut stream = stream; + stream + .write_all(b"CONNECT 8.8.8.8:443 HTTP/1.1\r\nHost: 8.8.8.8:443\r\n\r\n") + .await + .expect("write CONNECT request"); + let outbound = timeout(Duration::from_secs(1), outgoing_rx.recv()) + .await + .expect("policy request should arrive") + .expect("policy request"); + let RpcServerOutboundMessage::Request(request) = outbound else { + panic!("expected policy request"); + }; + assert_eq!(request.method, NETWORK_POLICY_REQUEST_METHOD); + let params: NetworkPolicyRequestParams = + serde_json::from_value(request.params.expect("request params")) + .expect("deserialize policy request"); + assert_eq!(params.process_id, process.process_id); + assert_eq!(params.request.host, "8.8.8.8"); + requests.complete( + request.id, + Ok(serde_json::to_value(NetworkPolicyRequestResponse { + decision: ExecServerNetworkPolicyDecision::Deny { + reason: "not_allowed".to_string(), + }, + }) + .expect("serialize policy response")), + ); + let mut response = [0_u8; 256]; + let response_len = timeout(Duration::from_secs(1), stream.read(&mut response)) + .await + .expect("proxy response timeout") + .expect("read proxy response"); + assert!(String::from_utf8_lossy(&response[..response_len]).starts_with("HTTP/1.1 403")); + } drop(process.stdout_tx); drop(process.stderr_tx); @@ -1384,6 +1611,7 @@ mod tests { .await .expect("process should close"); assert!(closed_response.closed); + assert!(network_policy_shutdown.is_cancelled()); #[cfg(target_os = "windows")] assert_eq!(proxy.network_proxy_restricting_sid(None), None); #[cfg(not(target_os = "windows"))] @@ -1476,6 +1704,7 @@ mod tests { sandbox: SandboxType::None, sandbox_denied: false, network_proxy_handle: None, + network_policy_shutdown: None, })), ); assert!(previous.is_none()); diff --git a/codex-rs/exec-server/src/network_policy_decisions.rs b/codex-rs/exec-server/src/network_policy_decisions.rs new file mode 100644 index 000000000000..ba21efe292af --- /dev/null +++ b/codex-rs/exec-server/src/network_policy_decisions.rs @@ -0,0 +1,95 @@ +use std::sync::Arc; +use std::sync::RwLock; +use std::time::Duration; + +use codex_network_proxy::NetworkDecision; +use codex_network_proxy::NetworkPolicyDecider; +use codex_network_proxy::NetworkPolicyRequest; +use codex_network_proxy::NetworkProtocol; +use tokio_util::sync::CancellationToken; + +use crate::ProcessId; +use crate::protocol::ExecServerNetworkPolicyDecision; +use crate::protocol::ExecServerNetworkPolicyRequest; +use crate::protocol::ExecServerNetworkProtocol; +use crate::protocol::MAX_NETWORK_POLICY_HOST_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::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); + +pub(crate) fn network_policy_decider( + process_id: ProcessId, + requests: Arc>>, + process_shutdown: CancellationToken, +) -> Arc { + Arc::new(move |request: NetworkPolicyRequest| { + let process_id = process_id.clone(); + let requests = Arc::clone(&requests); + let process_shutdown = process_shutdown.clone(); + async move { + let host = request.host.as_str(); + if host.is_empty() + || host.len() > MAX_NETWORK_POLICY_HOST_BYTES + || host.chars().any(char::is_control) + || host.chars().any(char::is_whitespace) + { + return NetworkDecision::deny("not_allowed"); + } + let requests = requests + .read() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .clone(); + let Some(requests) = requests else { + return NetworkDecision::deny("not_allowed"); + }; + let params = NetworkPolicyRequestParams { + process_id, + request: ExecServerNetworkPolicyRequest { + protocol: match request.protocol { + NetworkProtocol::Http => ExecServerNetworkProtocol::Http, + NetworkProtocol::HttpsConnect => ExecServerNetworkProtocol::HttpsConnect, + NetworkProtocol::Socks5Tcp => ExecServerNetworkProtocol::Socks5Tcp, + NetworkProtocol::Socks5Udp => ExecServerNetworkProtocol::Socks5Udp, + }, + host: request.host, + port: request.port, + }, + }; + tokio::select! { + biased; + _ = process_shutdown.cancelled() => NetworkDecision::deny("not_allowed"), + response = requests.call_with_timeout::<_, NetworkPolicyRequestResponse>( + NETWORK_POLICY_REQUEST_METHOD, + ¶ms, + NETWORK_POLICY_REQUEST_TIMEOUT, + ) => response + .map(|response| match response.decision { + ExecServerNetworkPolicyDecision::Allow => NetworkDecision::Allow, + ExecServerNetworkPolicyDecision::Deny { reason } + | ExecServerNetworkPolicyDecision::Ask { reason } + if reason.len() > MAX_NETWORK_POLICY_REASON_BYTES + || reason.chars().any(char::is_control) => + { + NetworkDecision::deny("not_allowed") + } + ExecServerNetworkPolicyDecision::Deny { reason } => { + NetworkDecision::deny(reason) + } + ExecServerNetworkPolicyDecision::Ask { reason } => { + NetworkDecision::ask(reason) + } + }) + .unwrap_or_else(|_| NetworkDecision::deny("not_allowed")), + } + } + }) +} + +#[cfg(test)] +#[path = "network_policy_decisions_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/network_policy_decisions_tests.rs b/codex-rs/exec-server/src/network_policy_decisions_tests.rs new file mode 100644 index 000000000000..550011539675 --- /dev/null +++ b/codex-rs/exec-server/src/network_policy_decisions_tests.rs @@ -0,0 +1,221 @@ +use std::sync::Arc; +use std::sync::RwLock; +use std::time::Duration; + +use codex_network_proxy::NetworkDecision; +use codex_network_proxy::NetworkPolicyRequest; +use codex_network_proxy::NetworkPolicyRequestArgs; +use codex_network_proxy::NetworkProtocol; +use pretty_assertions::assert_eq; +use tokio::sync::mpsc; +use tokio::task::JoinHandle; +use tokio::time::timeout; + +use super::*; +use crate::protocol::ExecServerNetworkPolicyDecision; +use crate::protocol::NetworkPolicyRequestParams; +use crate::protocol::NetworkPolicyRequestResponse; +use crate::rpc::RpcServerOutboundMessage; + +struct DeciderHarness { + requests: RpcServerRequestSender, + outgoing: mpsc::Receiver, + process_shutdown: CancellationToken, +} + +impl DeciderHarness { + fn new() -> Self { + let (outgoing_tx, outgoing) = mpsc::channel(/*buffer*/ 8); + Self { + requests: RpcServerRequestSender::new(outgoing_tx), + outgoing, + process_shutdown: CancellationToken::new(), + } + } + + fn request(&self, host: &str) -> JoinHandle { + let decider = network_policy_decider( + ProcessId::from("process"), + Arc::new(RwLock::new(Some(self.requests.clone()))), + self.process_shutdown.clone(), + ); + let request = NetworkPolicyRequest::new(NetworkPolicyRequestArgs { + protocol: NetworkProtocol::HttpsConnect, + host: host.to_string(), + port: 443, + environment_id: None, + client_addr: None, + method: None, + command: None, + exec_policy_hint: None, + }); + tokio::spawn(async move { decider.decide(request).await }) + } + + async fn next_request( + &mut self, + ) -> ( + codex_exec_server_protocol::RequestId, + NetworkPolicyRequestParams, + ) { + let outbound = timeout(Duration::from_secs(1), self.outgoing.recv()) + .await + .expect("policy request should arrive") + .expect("policy request"); + let RpcServerOutboundMessage::Request(request) = outbound else { + panic!("expected policy request"); + }; + assert_eq!(request.method, NETWORK_POLICY_REQUEST_METHOD); + let params = serde_json::from_value(request.params.expect("request params")) + .expect("deserialize policy request"); + (request.id, params) + } +} + +async fn await_decision(decision: JoinHandle) -> NetworkDecision { + timeout(Duration::from_secs(1), decision) + .await + .expect("network policy decision should resolve") + .expect("network policy decision task") +} + +#[tokio::test] +async fn returns_client_policy_decision() { + let mut harness = DeciderHarness::new(); + let decision = harness.request("example.com"); + let (request_id, params) = harness.next_request().await; + assert_eq!(params.process_id, ProcessId::from("process")); + assert_eq!(params.request.host, "example.com"); + harness.requests.complete( + request_id, + Ok(serde_json::to_value(NetworkPolicyRequestResponse { + decision: ExecServerNetworkPolicyDecision::Allow, + }) + .expect("serialize policy response")), + ); + + assert_eq!(await_decision(decision).await, NetworkDecision::Allow); +} + +#[tokio::test] +async fn policy_response_reasons_are_bounded_and_fail_closed() { + let mut harness = DeciderHarness::new(); + let boundary_reason = "d".repeat(MAX_NETWORK_POLICY_REASON_BYTES); + let cases = [ + ( + ExecServerNetworkPolicyDecision::Deny { + reason: boundary_reason.clone(), + }, + NetworkDecision::deny(boundary_reason), + ), + ( + ExecServerNetworkPolicyDecision::Deny { + reason: "d".repeat(MAX_NETWORK_POLICY_REASON_BYTES + 1), + }, + NetworkDecision::deny("not_allowed"), + ), + ( + ExecServerNetworkPolicyDecision::Ask { + reason: "ask permission".to_string(), + }, + NetworkDecision::ask("ask permission"), + ), + ( + ExecServerNetworkPolicyDecision::Ask { + reason: "ask\npermission".to_string(), + }, + NetworkDecision::deny("not_allowed"), + ), + ]; + + for (response, expected) in cases { + let decision = harness.request("example.com"); + let (request_id, _) = harness.next_request().await; + harness.requests.complete( + request_id, + Ok( + serde_json::to_value(NetworkPolicyRequestResponse { decision: response }) + .expect("serialize policy response"), + ), + ); + assert_eq!(await_decision(decision).await, expected); + } +} + +#[tokio::test] +async fn boundary_host_is_relayed() { + let mut harness = DeciderHarness::new(); + let host = "h".repeat(MAX_NETWORK_POLICY_HOST_BYTES); + let decision = harness.request(&host); + let (request_id, params) = harness.next_request().await; + assert_eq!(params.request.host, host); + harness.requests.complete( + request_id, + Ok(serde_json::to_value(NetworkPolicyRequestResponse { + decision: ExecServerNetworkPolicyDecision::Allow, + }) + .expect("serialize policy response")), + ); + + assert_eq!(await_decision(decision).await, NetworkDecision::Allow); +} + +#[tokio::test] +async fn invalid_hosts_fail_closed_before_reverse_rpc() { + let mut harness = DeciderHarness::new(); + let invalid_hosts = [ + String::new(), + "host name".to_string(), + "host\u{0000}name".to_string(), + "h".repeat(MAX_NETWORK_POLICY_HOST_BYTES + 1), + ]; + + for host in invalid_hosts { + assert_eq!( + await_decision(harness.request(&host)).await, + NetworkDecision::deny("not_allowed") + ); + assert!(harness.outgoing.try_recv().is_err()); + assert_eq!(harness.requests.pending_request_count(), 0); + } +} + +#[tokio::test] +async fn process_exit_and_disconnect_fail_closed() { + let mut process_exit = DeciderHarness::new(); + let process_decision = process_exit.request("process-exit.example.com"); + process_exit.next_request().await; + process_exit.process_shutdown.cancel(); + assert_eq!( + await_decision(process_decision).await, + NetworkDecision::deny("not_allowed") + ); + assert_eq!(process_exit.requests.pending_request_count(), 0); + + let mut disconnect = DeciderHarness::new(); + let disconnect_decision = disconnect.request("disconnect.example.com"); + disconnect.next_request().await; + disconnect.requests.close(); + assert_eq!( + await_decision(disconnect_decision).await, + NetworkDecision::deny("not_allowed") + ); + assert_eq!(disconnect.requests.pending_request_count(), 0); +} + +#[tokio::test(start_paused = true)] +async fn decision_timeout_fails_closed() { + let mut harness = DeciderHarness::new(); + let decision = harness.request("timeout.example.com"); + harness.next_request().await; + + tokio::time::advance(Duration::from_secs(99)).await; + assert!(!decision.is_finished()); + + tokio::time::advance(Duration::from_secs(1)).await; + assert_eq!( + await_decision(decision).await, + NetworkDecision::deny("not_allowed") + ); + assert_eq!(harness.requests.pending_request_count(), 0); +} diff --git a/codex-rs/exec-server/src/process_sandbox.rs b/codex-rs/exec-server/src/process_sandbox.rs index 61e77277c840..864592c1a26f 100644 --- a/codex-rs/exec-server/src/process_sandbox.rs +++ b/codex-rs/exec-server/src/process_sandbox.rs @@ -4,6 +4,7 @@ use std::sync::Arc; use codex_exec_server_protocol::JSONRPCErrorError; use codex_network_proxy::CUSTOM_CA_ENV_KEYS; use codex_network_proxy::ManagedNetworkSandboxContext; +use codex_network_proxy::NetworkPolicyDecider; use codex_network_proxy::NetworkProxy; use codex_network_proxy::NetworkProxyHandle; use codex_network_proxy::NetworkProxyState; @@ -76,6 +77,7 @@ pub(crate) async fn prepare_exec_request( params: &ExecParams, env: HashMap, runtime_paths: Option<&ExecServerRuntimePaths>, + network_policy_decider: Option>, ) -> Result { #[cfg(target_os = "windows")] let mut env = env; @@ -94,7 +96,13 @@ pub(crate) async fn prepare_exec_request( let network_proxy = params.network_proxy.as_ref(); let (env, managed_network, network_proxy_handle, network_proxy_restricting_sid) = - prepare_managed_network(params.managed_network.as_ref(), network_proxy, env).await?; + prepare_managed_network( + params.managed_network.as_ref(), + network_proxy, + env, + network_policy_decider, + ) + .await?; let Some(sandbox_context) = params.sandbox.as_ref() else { return Ok(PreparedExecRequest { command: params.argv.clone(), @@ -284,6 +292,7 @@ async fn prepare_managed_network( managed_network: Option<&ManagedNetworkSandboxContext>, network_proxy: Option<&RemoteNetworkProxyLaunchConfig>, env: HashMap, + network_policy_decider: Option>, ) -> Result< ( HashMap, @@ -298,8 +307,11 @@ async fn prepare_managed_network( }; let state = NetworkProxyState::from_remote_launch_config(network_proxy) .map_err(|err| invalid_params(format!("invalid network proxy config: {err}")))?; - let proxy = NetworkProxy::builder() - .state(Arc::new(state)) + let mut builder = NetworkProxy::builder().state(Arc::new(state)); + if let Some(network_policy_decider) = network_policy_decider { + builder = builder.policy_decider_arc(network_policy_decider); + } + let proxy = builder .build() .await .map_err(|err| internal_error(format!("failed to build executor network proxy: {err}")))?; diff --git a/codex-rs/exec-server/src/process_sandbox_tests.rs b/codex-rs/exec-server/src/process_sandbox_tests.rs index 1f930ff3ca5e..23ce7871f263 100644 --- a/codex-rs/exec-server/src/process_sandbox_tests.rs +++ b/codex-rs/exec-server/src/process_sandbox_tests.rs @@ -65,9 +65,14 @@ async fn sandbox_request_wraps_native_argv_on_executor() { network_proxy: None, }; - let prepared = prepare_exec_request(¶ms, HashMap::new(), Some(&runtime_paths)) - .await - .expect("prepare sandboxed request"); + let prepared = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + ) + .await + .expect("prepare sandboxed request"); assert_ne!(prepared.command, params.argv); assert_eq!(prepared.cwd, cwd); @@ -128,9 +133,14 @@ async fn sandbox_request_routes_custom_arg0_to_inner_helper() { network_proxy: None, }; - let prepared = prepare_exec_request(¶ms, HashMap::new(), Some(&runtime_paths)) - .await - .expect("prepare sandboxed request"); + let prepared = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + ) + .await + .expect("prepare sandboxed request"); let helper_mode = prepared .command .iter() @@ -186,9 +196,14 @@ async fn sandbox_request_allows_prepared_managed_proxy_port() { network_proxy: None, }; - let prepared = prepare_exec_request(¶ms, HashMap::new(), Some(&runtime_paths)) - .await - .expect("prepare managed-network sandbox request"); + let prepared = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + ) + .await + .expect("prepare managed-network sandbox request"); let policy = prepared .command .windows(2) @@ -221,9 +236,14 @@ async fn native_request_preserves_native_launch_fields() { network_proxy: None, }; - let prepared = prepare_exec_request(¶ms, env.clone(), /*runtime_paths*/ None) - .await - .expect("prepare native request"); + let prepared = prepare_exec_request( + ¶ms, + env.clone(), + /*runtime_paths*/ None, + /*network_policy_decider*/ None, + ) + .await + .expect("prepare native request"); assert_eq!(prepared.command, params.argv); assert_eq!(prepared.cwd, cwd); @@ -271,9 +291,11 @@ async fn native_request_handles_remote_proxy_config_for_platform() { ), ]); - let prepared = prepare_exec_request(¶ms, env, /*runtime_paths*/ None) - .await - .expect("prepare request with executor-local proxy"); + let prepared = prepare_exec_request( + ¶ms, env, /*runtime_paths*/ None, /*network_policy_decider*/ None, + ) + .await + .expect("prepare request with executor-local proxy"); if cfg!(target_os = "windows") { assert_eq!(prepared.env.get("HTTP_PROXY"), None); @@ -338,10 +360,15 @@ async fn disabled_remote_proxy_config_is_rejected_before_exporting_ports() { network_proxy: Some(RemoteNetworkProxyLaunchConfig::new(proxy_config)), }; - let error = prepare_exec_request(¶ms, HashMap::new(), /*runtime_paths*/ None) - .await - .err() - .expect("disabled executor proxy launch must fail closed"); + let error = prepare_exec_request( + ¶ms, + HashMap::new(), + /*runtime_paths*/ None, + /*network_policy_decider*/ None, + ) + .await + .err() + .expect("disabled executor proxy launch must fail closed"); assert_eq!(error.code, -32602); assert!( @@ -388,9 +415,14 @@ async fn managed_network_selects_elevated_windows_spawn() { network_proxy: Some(RemoteNetworkProxyLaunchConfig::new(proxy_config)), }; - let mut prepared = prepare_exec_request(¶ms, HashMap::new(), Some(&runtime_paths)) - .await - .expect("prepare sandboxed request"); + let mut prepared = prepare_exec_request( + ¶ms, + HashMap::new(), + Some(&runtime_paths), + /*network_policy_decider*/ None, + ) + .await + .expect("prepare sandboxed request"); { let spawn = prepared .windows_sandbox_spawn_request() diff --git a/codex-rs/exec-server/src/rpc.rs b/codex-rs/exec-server/src/rpc.rs index de5d05ecc4f5..dd565b3df63f 100644 --- a/codex-rs/exec-server/src/rpc.rs +++ b/codex-rs/exec-server/src/rpc.rs @@ -29,6 +29,7 @@ use tokio::time::timeout; use crate::connection::JsonRpcConnection; use crate::connection::JsonRpcConnectionEvent; use crate::connection::JsonRpcTransport; +use crate::rpc_server_requests::RpcServerRequestSender; pub(crate) const SESSION_ALREADY_ATTACHED_ERROR_CODE: i64 = -32010; const MAX_IN_FLIGHT_REGULAR_CALLS: usize = 1024; @@ -63,12 +64,14 @@ enum RpcCallTimeout { #[derive(Debug)] pub(crate) enum RpcClientEvent { + Request(JSONRPCRequest), Notification(JSONRPCNotification), Disconnected { reason: Option }, } #[derive(Debug, Clone, PartialEq)] pub(crate) enum RpcServerOutboundMessage { + Request(JSONRPCRequest), Response { request_id: RequestId, result: Value, @@ -83,11 +86,20 @@ pub(crate) enum RpcServerOutboundMessage { #[derive(Clone)] pub(crate) struct RpcNotificationSender { outgoing_tx: mpsc::Sender, + requests: RpcServerRequestSender, } impl RpcNotificationSender { pub(crate) fn new(outgoing_tx: mpsc::Sender) -> Self { - Self { outgoing_tx } + let requests = RpcServerRequestSender::new(outgoing_tx.clone()); + Self { + outgoing_tx, + requests, + } + } + + pub(crate) fn request_sender(&self) -> RpcServerRequestSender { + self.requests.clone() } pub(crate) async fn response( @@ -339,6 +351,23 @@ impl RpcClient { .map_err(|_| RpcCallError::Closed) } + pub(crate) async fn respond_error( + &self, + request_id: RequestId, + error: JSONRPCErrorError, + ) -> Result<(), RpcCallError> { + if self.closed.load(Ordering::Acquire) || *self.disconnected_rx.borrow() { + return Err(RpcCallError::Closed); + } + self.write_tx + .send(JSONRPCMessage::Error(JSONRPCError { + id: request_id, + error, + })) + .await + .map_err(|_| RpcCallError::Closed) + } + pub(crate) fn is_disconnected(&self) -> bool { self.closed.load(Ordering::Acquire) || *self.disconnected_rx.borrow() } @@ -522,6 +551,7 @@ pub(crate) fn encode_server_message( message: RpcServerOutboundMessage, ) -> Result { match message { + RpcServerOutboundMessage::Request(request) => Ok(JSONRPCMessage::Request(request)), RpcServerOutboundMessage::Response { request_id, result } => { Ok(JSONRPCMessage::Response(JSONRPCResponse { id: request_id, @@ -642,10 +672,10 @@ async fn handle_server_message( .await; } JSONRPCMessage::Request(request) => { - return Err(format!( - "unexpected JSON-RPC request from remote server: {}", - request.method - )); + event_tx + .send(RpcClientEvent::Request(request)) + .await + .map_err(|_| "RPC client event receiver closed".to_string())?; } } diff --git a/codex-rs/exec-server/src/rpc_server_requests.rs b/codex-rs/exec-server/src/rpc_server_requests.rs new file mode 100644 index 000000000000..f718b2726657 --- /dev/null +++ b/codex-rs/exec-server/src/rpc_server_requests.rs @@ -0,0 +1,170 @@ +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::Mutex; +use std::sync::atomic::AtomicI64; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::RequestId; +use serde::Serialize; +use serde::de::DeserializeOwned; +use serde_json::Value; +use tokio::sync::Semaphore; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::time::timeout; +use tokio_util::sync::CancellationToken; + +use crate::rpc::RpcCallError; +use crate::rpc::RpcServerOutboundMessage; + +const MAX_IN_FLIGHT_SERVER_CALLS: usize = 256; + +type PendingRequest = oneshot::Sender>; + +#[derive(Clone)] +pub(crate) struct RpcServerRequestSender { + inner: Arc, +} + +struct RpcServerRequestSenderInner { + outgoing_tx: mpsc::Sender, + pending: Mutex>, + call_slots: Semaphore, + next_request_id: AtomicI64, + closed: CancellationToken, +} + +struct PendingServerRequestGuard { + inner: Arc, + request_id: RequestId, +} + +impl Drop for PendingServerRequestGuard { + fn drop(&mut self) { + self.inner + .pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .remove(&self.request_id); + } +} + +impl RpcServerRequestSender { + pub(crate) fn new(outgoing_tx: mpsc::Sender) -> Self { + Self { + inner: Arc::new(RpcServerRequestSenderInner { + outgoing_tx, + pending: Mutex::new(HashMap::new()), + call_slots: Semaphore::new(MAX_IN_FLIGHT_SERVER_CALLS), + next_request_id: AtomicI64::new(1), + closed: CancellationToken::new(), + }), + } + } + + pub(crate) async fn call_with_timeout( + &self, + method: &str, + params: &P, + call_timeout: Duration, + ) -> Result + where + P: Serialize, + T: DeserializeOwned, + { + let _call_slot = self.inner.call_slots.try_acquire().map_err(|_| { + RpcCallError::PendingRequestLimitExceeded { + limit: MAX_IN_FLIGHT_SERVER_CALLS, + } + })?; + let params = serde_json::to_value(params).map_err(RpcCallError::Json)?; + let request_id = + RequestId::Integer(self.inner.next_request_id.fetch_add(1, Ordering::SeqCst)); + let (response_tx, response_rx) = oneshot::channel(); + { + let mut pending = self + .inner + .pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + if self.inner.closed.is_cancelled() { + return Err(RpcCallError::Closed); + } + pending.insert(request_id.clone(), response_tx); + } + let _pending = PendingServerRequestGuard { + inner: Arc::clone(&self.inner), + request_id: request_id.clone(), + }; + let request = RpcServerOutboundMessage::Request(JSONRPCRequest { + id: request_id, + method: method.to_string(), + params: Some(params), + trace: codex_otel::current_span_w3c_trace_context(), + }); + + let response = timeout(call_timeout, async { + tokio::select! { + biased; + _ = self.inner.closed.cancelled() => return Err(RpcCallError::Closed), + result = self.inner.outgoing_tx.send(request) => { + result.map_err(|_| RpcCallError::Closed)?; + } + } + response_rx.await.map_err(|_| RpcCallError::Closed)? + }) + .await + .map_err(|_| RpcCallError::TimedOut { + method: method.to_string(), + timeout: call_timeout, + })??; + serde_json::from_value(response).map_err(RpcCallError::Json) + } + + pub(crate) fn complete( + &self, + request_id: RequestId, + result: Result, + ) -> bool { + if let Some(pending) = self + .inner + .pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .remove(&request_id) + { + let _ = pending.send(result); + true + } else { + matches!( + request_id, + RequestId::Integer(id) + if id > 0 && id < self.inner.next_request_id.load(Ordering::Acquire) + ) + } + } + + pub(crate) fn close(&self) { + self.inner.closed.cancel(); + let pending = { + let mut pending = self + .inner + .pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner); + pending + .drain() + .map(|(_, pending)| pending) + .collect::>() + }; + for pending in pending { + let _ = pending.send(Err(RpcCallError::Closed)); + } + } +} + +#[cfg(test)] +#[path = "rpc_server_requests_tests.rs"] +mod tests; diff --git a/codex-rs/exec-server/src/rpc_server_requests_tests.rs b/codex-rs/exec-server/src/rpc_server_requests_tests.rs new file mode 100644 index 000000000000..05692c32aaae --- /dev/null +++ b/codex-rs/exec-server/src/rpc_server_requests_tests.rs @@ -0,0 +1,178 @@ +use std::sync::Arc; +use std::time::Duration; + +use codex_exec_server_protocol::JSONRPCRequest; +use codex_exec_server_protocol::RequestId; +use pretty_assertions::assert_eq; +use tokio::sync::mpsc; +use tokio::task::JoinSet; +use tokio::time::timeout; + +use super::MAX_IN_FLIGHT_SERVER_CALLS; +use super::RpcServerRequestSender; +use crate::rpc::RpcCallError; +use crate::rpc::RpcServerOutboundMessage; + +async fn receive_server_request( + outgoing_rx: &mut mpsc::Receiver, +) -> JSONRPCRequest { + let message = timeout(Duration::from_secs(1), outgoing_rx.recv()) + .await + .expect("server request should arrive") + .expect("server request"); + match message { + RpcServerOutboundMessage::Request(request) => request, + other => panic!("expected server request, got {other:?}"), + } +} + +impl RpcServerRequestSender { + pub(crate) fn pending_request_count(&self) -> usize { + self.inner + .pending + .lock() + .unwrap_or_else(std::sync::PoisonError::into_inner) + .len() + } +} + +#[tokio::test] +async fn rpc_server_sender_matches_out_of_order_responses_by_request_id() { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(/*buffer*/ 8); + let requests = RpcServerRequestSender::new(outgoing_tx); + let slow_requests = requests.clone(); + let slow = tokio::spawn(async move { + slow_requests + .call_with_timeout::<_, serde_json::Value>( + "slow", + &serde_json::json!({ "n": 1 }), + Duration::from_secs(1), + ) + .await + }); + let fast_requests = requests.clone(); + let fast = tokio::spawn(async move { + fast_requests + .call_with_timeout::<_, serde_json::Value>( + "fast", + &serde_json::json!({ "n": 2 }), + Duration::from_secs(1), + ) + .await + }); + + let first = receive_server_request(&mut outgoing_rx).await; + let second = receive_server_request(&mut outgoing_rx).await; + let (slow_request, fast_request) = if first.method == "slow" { + (first, second) + } else { + (second, first) + }; + requests.complete(fast_request.id, Ok(serde_json::json!({ "value": "fast" }))); + requests.complete(slow_request.id, Ok(serde_json::json!({ "value": "slow" }))); + + assert_eq!( + slow.await.expect("slow task").expect("slow server request"), + serde_json::json!({ "value": "slow" }) + ); + assert_eq!( + fast.await.expect("fast task").expect("fast server request"), + serde_json::json!({ "value": "fast" }) + ); + assert_eq!(requests.pending_request_count(), 0); +} + +#[tokio::test] +async fn rpc_server_sender_preserves_response_received_before_close() { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(/*buffer*/ 1); + let requests = RpcServerRequestSender::new(outgoing_tx); + let caller = requests.clone(); + let call = tokio::spawn(async move { + caller + .call_with_timeout::<_, serde_json::Value>( + "ordered", + &serde_json::json!({}), + Duration::from_secs(1), + ) + .await + }); + + let request = receive_server_request(&mut outgoing_rx).await; + requests.complete(request.id, Ok(serde_json::json!({ "value": "accepted" }))); + requests.close(); + + assert_eq!( + call.await + .expect("server request task should join") + .expect("server request should preserve its response"), + serde_json::json!({ "value": "accepted" }) + ); + assert_eq!(requests.pending_request_count(), 0); +} + +#[tokio::test(start_paused = true)] +async fn rpc_server_sender_timeout_removes_pending_request() { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(/*buffer*/ 1); + let requests = RpcServerRequestSender::new(outgoing_tx); + let call_timeout = Duration::from_secs(1); + let params = serde_json::json!({}); + let call = requests.call_with_timeout::<_, serde_json::Value>("slow", ¶ms, call_timeout); + tokio::pin!(call); + assert!(futures::poll!(call.as_mut()).is_pending()); + let request = receive_server_request(&mut outgoing_rx).await; + + tokio::time::advance(call_timeout).await; + assert!(matches!( + call.await, + Err(RpcCallError::TimedOut { method, timeout }) + if method == "slow" && timeout == call_timeout + )); + assert_eq!(requests.pending_request_count(), 0); + assert!(requests.complete(request.id, Ok(serde_json::Value::Null))); + assert!(!requests.complete(RequestId::Integer(2), Ok(serde_json::Value::Null))); +} + +#[tokio::test] +async fn rpc_server_sender_bounds_and_drains_pending_requests_on_close() { + let (outgoing_tx, mut outgoing_rx) = mpsc::channel(MAX_IN_FLIGHT_SERVER_CALLS); + let requests = Arc::new(RpcServerRequestSender::new(outgoing_tx)); + let mut calls = JoinSet::new(); + for index in 0..MAX_IN_FLIGHT_SERVER_CALLS { + let requests = Arc::clone(&requests); + calls.spawn(async move { + requests + .call_with_timeout::<_, serde_json::Value>( + "pending", + &serde_json::json!({ "index": index }), + Duration::from_secs(30), + ) + .await + }); + } + for _ in 0..MAX_IN_FLIGHT_SERVER_CALLS { + receive_server_request(&mut outgoing_rx).await; + } + assert_eq!(requests.pending_request_count(), MAX_IN_FLIGHT_SERVER_CALLS); + + let overflow = requests + .call_with_timeout::<_, serde_json::Value>( + "overflow", + &serde_json::json!({}), + Duration::from_secs(1), + ) + .await; + assert!(matches!( + overflow, + Err(RpcCallError::PendingRequestLimitExceeded { limit }) + if limit == MAX_IN_FLIGHT_SERVER_CALLS + )); + + requests.close(); + assert_eq!(requests.pending_request_count(), 0); + while let Some(call) = calls.join_next().await { + assert!(matches!( + call.expect("pending server request task should join"), + Err(RpcCallError::Closed) + )); + } +} diff --git a/codex-rs/exec-server/src/server/processor.rs b/codex-rs/exec-server/src/server/processor.rs index 0d98dcd54bfa..681c4bd67a0a 100644 --- a/codex-rs/exec-server/src/server/processor.rs +++ b/codex-rs/exec-server/src/server/processor.rs @@ -10,6 +10,7 @@ use crate::ExecServerRuntimePaths; use crate::connection::CHANNEL_CAPACITY; use crate::connection::JsonRpcConnection; use crate::connection::JsonRpcConnectionEvent; +use crate::rpc::RpcCallError; use crate::rpc::RpcNotificationSender; use crate::rpc::RpcServerOutboundMessage; use crate::rpc::encode_server_message; @@ -84,6 +85,7 @@ async fn run_connection( let (outgoing_tx, mut outgoing_rx) = mpsc::channel::(CHANNEL_CAPACITY); let notifications = RpcNotificationSender::new(outgoing_tx.clone()); + let requests = notifications.request_sender(); let handler = Arc::new(ExecServerHandler::new( session_registry, notifications, @@ -208,18 +210,23 @@ async fn run_connection( } } codex_exec_server_protocol::JSONRPCMessage::Response(response) => { - warn!( - "closing exec-server connection after unexpected client response: {:?}", - response.id - ); - break; + if !requests.complete(response.id.clone(), Ok(response.result)) { + warn!( + "closing exec-server connection after unexpected client response: {:?}", + response.id + ); + break; + } } codex_exec_server_protocol::JSONRPCMessage::Error(error) => { - warn!( - "closing exec-server connection after unexpected client error: {:?}", - error.id - ); - break; + if !requests.complete(error.id.clone(), Err(RpcCallError::Server(error.error))) + { + warn!( + "closing exec-server connection after unexpected client error: {:?}", + error.id + ); + break; + } } }, JsonRpcConnectionEvent::Disconnected { reason } => { @@ -231,8 +238,10 @@ async fn run_connection( } } + requests.close(); handler.shutdown().await; drop(handler); + drop(requests); drop(outgoing_tx); for task in connection_tasks { task.abort(); @@ -265,7 +274,9 @@ fn request_result(message: &Option) -> &'static str { match message { Some(RpcServerOutboundMessage::Error { .. }) => "error", Some( - RpcServerOutboundMessage::Response { .. } | RpcServerOutboundMessage::Notification(_), + RpcServerOutboundMessage::Request(_) + | RpcServerOutboundMessage::Response { .. } + | RpcServerOutboundMessage::Notification(_), ) | None => "success", }