From b3bc2901fd47a60981d9dfe286063834d3aefbdd Mon Sep 17 00:00:00 2001 From: Lukas Wirth Date: Fri, 24 Jul 2026 21:37:40 +0200 Subject: [PATCH] proc-macro-api: Implement client side timeout handling --- Cargo.lock | 5 +- crates/load-cargo/src/lib.rs | 2 + crates/proc-macro-api/Cargo.toml | 1 + .../src/bidirectional_protocol.rs | 11 +- crates/proc-macro-api/src/legacy_protocol.rs | 17 +- crates/proc-macro-api/src/lib.rs | 97 +++- crates/proc-macro-api/src/pool.rs | 116 ++++- crates/proc-macro-api/src/process.rs | 430 +++++++++++++++++- crates/rust-analyzer/src/reload.rs | 9 +- 9 files changed, 604 insertions(+), 84 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 0e12d80b7e18..f02bdabd693e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -863,7 +863,6 @@ dependencies = [ "rustc-hash 2.1.2", "rustc_apfloat", "salsa", - "salsa-macros", "smallvec", "span", "stdx", @@ -892,7 +891,6 @@ dependencies = [ "parser", "rustc-hash 2.1.2", "salsa", - "salsa-macros", "smallvec", "span", "stdx", @@ -934,7 +932,6 @@ dependencies = [ "rustc-hash 2.1.2", "rustc_apfloat", "salsa", - "salsa-macros", "serde", "serde_derive", "smallvec", @@ -1135,7 +1132,6 @@ dependencies = [ "rayon", "rustc-hash 2.1.2", "salsa", - "salsa-macros", "smallvec", "span", "stdx", @@ -1869,6 +1865,7 @@ version = "0.0.0" dependencies = [ "indexmap", "intern", + "parking_lot", "paths", "postcard", "proc-macro-srv", diff --git a/crates/load-cargo/src/lib.rs b/crates/load-cargo/src/lib.rs index f5c5cb432d55..e60a16278e49 100644 --- a/crates/load-cargo/src/lib.rs +++ b/crates/load-cargo/src/lib.rs @@ -116,6 +116,7 @@ pub fn load_workspace_into_db( extra_env, ws.toolchain.as_ref(), load_config.proc_macro_processes, + Some(proc_macro_api::DEFAULT_EXPANSION_TIMEOUT), ) .map_err(Into::into) }) @@ -127,6 +128,7 @@ pub fn load_workspace_into_db( extra_env, ws.toolchain.as_ref(), load_config.proc_macro_processes, + Some(proc_macro_api::DEFAULT_EXPANSION_TIMEOUT), ) .map_err(|e| ProcMacroLoadingError::ProcMacroSrvError(e.to_string().into_boxed_str())), ), diff --git a/crates/proc-macro-api/Cargo.toml b/crates/proc-macro-api/Cargo.toml index 7342e0ecdcf6..d7f69274430d 100644 --- a/crates/proc-macro-api/Cargo.toml +++ b/crates/proc-macro-api/Cargo.toml @@ -19,6 +19,7 @@ serde_json = { workspace = true, features = ["unbounded_depth"] } tracing.workspace = true rustc-hash.workspace = true indexmap.workspace = true +parking_lot = "0.12.4" # local deps paths = { workspace = true, features = ["serde1"] } diff --git a/crates/proc-macro-api/src/bidirectional_protocol.rs b/crates/proc-macro-api/src/bidirectional_protocol.rs index f070b1c9a334..e7ada60b35d4 100644 --- a/crates/proc-macro-api/src/bidirectional_protocol.rs +++ b/crates/proc-macro-api/src/bidirectional_protocol.rs @@ -100,7 +100,7 @@ pub(crate) fn version_check( ) -> Result { let request = BidirectionalMessage::Request(Request::ApiVersionCheck(ApiVersionCheck {})); - let response_payload = run_request(srv, request, callback)?; + let response_payload = run_request(srv, request, callback, None)?; match response_payload { BidirectionalMessage::Response(Response::ApiVersionCheck(version)) => Ok(version), @@ -119,7 +119,7 @@ pub(crate) fn enable_rust_analyzer_spans( span_mode: SpanMode::RustAnalyzer, })); - let response_payload = run_request(srv, request, callback)?; + let response_payload = run_request(srv, request, callback, None)?; match response_payload { BidirectionalMessage::Response(Response::SetConfig(ServerConfig { span_mode })) => { @@ -139,7 +139,7 @@ pub(crate) fn find_proc_macros( dylib_path: dylib_path.to_path_buf().into(), })); - let response_payload = run_request(srv, request, callback)?; + let response_payload = run_request(srv, request, callback, None)?; match response_payload { BidirectionalMessage::Response(Response::ListMacros(it)) => Ok(it), @@ -182,7 +182,7 @@ pub(crate) fn expand( current_dir: Some(current_dir), }))); - let response_payload = run_request(process, task, callback)?; + let response_payload = run_request(process, task, callback, Some(proc_macro.name()))?; match response_payload { BidirectionalMessage::Response(Response::ExpandMacro(it)) => Ok(it @@ -202,11 +202,12 @@ fn run_request( srv: &ProcMacroServerProcess, msg: BidirectionalMessage, callback: SubCallback<'_>, + macro_name: Option<&str>, ) -> Result { if let Some(err) = srv.exited() { return Err(err.clone()); } - srv.run_bidirectional(msg, callback) + srv.run_bidirectional(msg, callback, macro_name) } pub fn reject_subrequests(req: SubRequest) -> Result { diff --git a/crates/proc-macro-api/src/legacy_protocol.rs b/crates/proc-macro-api/src/legacy_protocol.rs index ee1795d39c2e..17c7858f0024 100644 --- a/crates/proc-macro-api/src/legacy_protocol.rs +++ b/crates/proc-macro-api/src/legacy_protocol.rs @@ -37,7 +37,7 @@ impl std::fmt::Debug for SpanId { pub(crate) fn version_check(srv: &ProcMacroServerProcess) -> Result { let request = Request::ApiVersionCheck {}; - let response = send_task(srv, request)?; + let response = send_task(srv, request, None)?; match response { Response::ApiVersionCheck(version) => Ok(version), @@ -50,7 +50,7 @@ pub(crate) fn enable_rust_analyzer_spans( srv: &ProcMacroServerProcess, ) -> Result { let request = Request::SetConfig(ServerConfig { span_mode: SpanMode::RustAnalyzer }); - let response = send_task(srv, request)?; + let response = send_task(srv, request, None)?; match response { Response::SetConfig(ServerConfig { span_mode }) => Ok(span_mode), @@ -65,7 +65,7 @@ pub(crate) fn find_proc_macros( ) -> Result, String>, ServerError> { let request = Request::ListMacros { dylib_path: dylib_path.to_path_buf().into() }; - let response = send_task(srv, request)?; + let response = send_task(srv, request, None)?; match response { Response::ListMacros(it) => Ok(it), @@ -112,7 +112,8 @@ pub(crate) fn expand( current_dir: Some(current_dir), }; - let response = send_task(process, Request::ExpandMacro(Box::new(task)))?; + let response = + send_task(process, Request::ExpandMacro(Box::new(task)), Some(proc_macro.name()))?; match response { Response::ExpandMacro(it) => Ok(it @@ -142,12 +143,16 @@ pub(crate) fn expand( } /// Sends a request to the proc-macro server and waits for a response. -fn send_task(srv: &ProcMacroServerProcess, req: Request) -> Result { +fn send_task( + srv: &ProcMacroServerProcess, + req: Request, + macro_name: Option<&str>, +) -> Result { if let Some(server_error) = srv.exited() { return Err(server_error.clone()); } - srv.send_task_legacy::<_, _>(send_request, req) + srv.send_task_legacy::<_, _>(send_request, req, macro_name) } /// Sends a request to the server and reads the response. diff --git a/crates/proc-macro-api/src/lib.rs b/crates/proc-macro-api/src/lib.rs index 7b2e209b8e03..8616a79658d2 100644 --- a/crates/proc-macro-api/src/lib.rs +++ b/crates/proc-macro-api/src/lib.rs @@ -24,12 +24,23 @@ pub mod transport; use paths::{AbsPath, AbsPathBuf}; use semver::Version; use span::{ErasedFileAstId, FIXUP_ERASED_FILE_AST_ID_MARKER, Span}; -use std::{fmt, io, sync::Arc, time::SystemTime}; +use std::{ + ffi::OsString, + fmt, io, + sync::Arc, + time::{Duration, SystemTime}, +}; use crate::{ - bidirectional_protocol::SubCallback, pool::ProcMacroServerPool, process::ProcMacroServerProcess, + bidirectional_protocol::SubCallback, + pool::{ProcMacroServerPool, ProcessFactory}, + process::{ProcMacroServerProcess, ProcessWatchdog}, }; +/// How long a single proc-macro expansion may take before the server process is killed +/// and restarted. +pub const DEFAULT_EXPANSION_TIMEOUT: Duration = Duration::from_secs(30); + /// The versions of the server protocol pub mod version { pub const NO_VERSION_CHECK_VERSION: u32 = 0; @@ -108,7 +119,7 @@ impl MacroDylib { /// we share a single expander process for all macros within a workspace. #[derive(Debug, Clone)] pub struct ProcMacro { - pool: ProcMacroServerPool, + pool: Arc, dylib_path: Arc, name: Box, kind: ProcMacroKind, @@ -149,19 +160,40 @@ impl ProcMacroClient { process_path: &AbsPath, env: impl IntoIterator< Item = (impl AsRef, &'a Option>), - > + Clone, + >, version: Option<&Version>, num_process: usize, + expansion_timeout: Option, ) -> io::Result { - let pool_size = num_process; - let mut workers = Vec::with_capacity(pool_size); - for _ in 0..pool_size { - let worker = ProcMacroServerProcess::spawn(process_path, env.clone(), version)?; - workers.push(worker); - } + let process_path = process_path.to_owned(); + let env: Arc<[(OsString, Option)]> = env + .into_iter() + .map(|(key, value)| { + ( + key.as_ref().to_os_string(), + value.as_ref().map(|value| value.as_ref().to_os_string()), + ) + }) + .collect(); + let version = version.cloned(); + let watchdog = expansion_timeout + .map(|timeout| io::Result::Ok((ProcessWatchdog::spawn()?, timeout))) + .transpose()?; + let spawn: ProcessFactory = Box::new({ + let process_path = process_path.clone(); + move || { + ProcMacroServerProcess::spawn( + &process_path, + env.iter().map(|(key, value)| (key, value)), + version.as_ref(), + watchdog.clone(), + ) + } + }); + let workers = (0..num_process).map(|_| spawn()).collect::>>()?; - let pool = ProcMacroServerPool::new(workers); - Ok(ProcMacroClient { pool: Arc::new(pool), path: process_path.to_owned() }) + let pool = Arc::new(ProcMacroServerPool::new(workers, spawn)); + Ok(ProcMacroClient { pool, path: process_path }) } /// Invokes `spawn` and returns a client connected to the resulting read and write handles. @@ -175,20 +207,30 @@ impl ProcMacroClient { Box, Box, Box, - )> + Clone, + )> + Clone + + Send + + Sync + + 'static, version: Option<&Version>, num_process: usize, + expansion_timeout: Option, ) -> io::Result { - let pool_size = num_process; - let mut workers = Vec::with_capacity(pool_size); - for _ in 0..pool_size { - let worker = - ProcMacroServerProcess::run(spawn.clone(), version, || "".to_owned())?; - workers.push(worker); - } + let version = version.cloned(); + let watchdog = expansion_timeout + .map(|timeout| io::Result::Ok((ProcessWatchdog::spawn()?, timeout))) + .transpose()?; + let spawn: ProcessFactory = Box::new(move || { + ProcMacroServerProcess::run( + spawn.clone(), + version.as_ref(), + || "".to_owned(), + watchdog.clone(), + ) + }); + let workers = (0..num_process).map(|_| spawn()).collect::>>()?; - let pool = ProcMacroServerPool::new(workers); - Ok(ProcMacroClient { pool: Arc::new(pool), path: process_path.to_owned() }) + let pool = Arc::new(ProcMacroServerPool::new(workers, spawn)); + Ok(ProcMacroClient { pool, path: process_path.to_owned() }) } /// Returns the absolute path to the proc-macro server. @@ -202,7 +244,7 @@ impl ProcMacroClient { } /// Checks if the proc-macro server has exited. - pub fn exited(&self) -> Option<&ServerError> { + pub fn exited(&self) -> Option { self.pool.exited() } } @@ -263,7 +305,8 @@ impl ProcMacro { } } - self.pool.pick_process()?.expand( + let process = self.pool.pick_process()?; + let result = process.expand( self, subtree, attr, @@ -273,6 +316,10 @@ impl ProcMacro { mixed_site, current_dir, callback, - ) + ); + if process.timed_out() { + self.pool.replace_timed_out_process_in_background(process); + } + result } } diff --git a/crates/proc-macro-api/src/pool.rs b/crates/proc-macro-api/src/pool.rs index e6541823da58..2d052aa54650 100644 --- a/crates/proc-macro-api/src/pool.rs +++ b/crates/proc-macro-api/src/pool.rs @@ -1,44 +1,62 @@ //! A pool of proc-macro server processes -use std::sync::Arc; +use std::{io, panic::RefUnwindSafe, sync::Arc}; +use parking_lot::Mutex; use rayon::iter::{IntoParallelIterator, ParallelIterator}; use crate::{MacroDylib, ProcMacro, ServerError, process::ProcMacroServerProcess}; -#[derive(Debug, Clone)] +pub(crate) type ProcessFactory = Box io::Result + Send + Sync>; + +/// A fixed-size pool of proc-macro server processes. pub(crate) struct ProcMacroServerPool { - workers: Arc<[ProcMacroServerProcess]>, + workers: Mutex]>>, + spawn: ProcessFactory, version: u32, } -impl ProcMacroServerPool { - pub(crate) fn new(workers: Vec) -> Self { - let version = workers[0].version(); - Self { workers: workers.into(), version } +impl RefUnwindSafe for ProcMacroServerPool {} + +impl std::fmt::Debug for ProcMacroServerPool { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ProcMacroServerPool") + .field("workers", &self.workers) + .field("version", &self.version) + .finish() } } impl ProcMacroServerPool { - pub(crate) fn exited(&self) -> Option<&ServerError> { - for worker in &*self.workers { - worker.exited()?; + pub(crate) fn new(workers: Vec, spawn: ProcessFactory) -> Self { + let version = workers[0].version(); + let workers = workers.into_iter().map(Arc::new).collect::>().into_boxed_slice(); + Self { workers: Mutex::new(workers), spawn, version } + } + + pub(crate) fn exited(&self) -> Option { + let workers = self.workers.lock(); + if workers.iter().any(|worker| worker.exited().is_none()) { + return None; } - self.workers[0].exited() + workers.first()?.exited().cloned() } - pub(crate) fn pick_process(&self) -> Result<&ProcMacroServerProcess, ServerError> { - let mut best: Option<&ProcMacroServerProcess> = None; + pub(crate) fn pick_process(&self) -> Result, ServerError> { + let mut best: Option> = None; let mut best_load = u32::MAX; - for w in self.workers.iter().filter(|w| w.exited().is_none()) { - let load = w.number_of_active_req(); + for worker in self.workers.lock().iter() { + if worker.exited().is_some() { + continue; + } + let load = worker.number_of_active_req(); if load == 0 { - return Ok(w); + return Ok(worker.clone()); } if load < best_load { - best = Some(w); + best = Some(worker.clone()); best_load = load; } } @@ -49,14 +67,72 @@ impl ProcMacroServerPool { }) } - pub(crate) fn load_dylib(&self, dylib: &MacroDylib) -> Result, ServerError> { + pub(crate) fn replace_timed_out_process_in_background( + self: &Arc, + process: Arc, + ) { + if !process.claim_replacement() { + return; + } + + let thread = std::thread::Builder::new().name("proc-macro-restarter".into()).spawn({ + let pool = self.clone(); + let process = process.clone(); + move || { + if let Err(error) = pool.replace_timed_out_process(&process) { + process.replacement_failed(); + tracing::error!(%error, "failed to replace timed out proc-macro server"); + } + } + }); + if let Err(error) = thread { + process.replacement_failed(); + tracing::error!(%error, "failed to spawn proc-macro server restarter"); + } + } + + fn replace_timed_out_process( + &self, + process: &Arc, + ) -> Result<(), ServerError> { + let slot = self.workers.lock().iter().position(|worker| Arc::ptr_eq(worker, process)); + let Some(slot) = slot else { + return Ok(()); + }; + + let replacement = (self.spawn)().map_err(|error| ServerError { + message: "failed to restart proc-macro server after expansion timeout".into(), + io: Some(Arc::new(error)), + })?; + if replacement.version() != self.version { + return Err(ServerError { + message: format!( + "restarted proc-macro server changed protocol version from {} to {}", + self.version, + replacement.version() + ), + io: None, + }); + } + + let mut workers = self.workers.lock(); + if Arc::ptr_eq(&workers[slot], process) { + workers[slot] = Arc::new(replacement); + } + Ok(()) + } + + pub(crate) fn load_dylib( + self: &Arc, + dylib: &MacroDylib, + ) -> Result, ServerError> { let _span = tracing::info_span!("ProcMacroServer::load_dylib").entered(); let dylib_path = Arc::new(dylib.path.clone()); let dylib_last_modified = std::fs::metadata(dylib_path.as_path()).ok().and_then(|m| m.modified().ok()); - - let (first, rest) = self.workers.split_first().expect("worker pool must not be empty"); + let workers = self.workers.lock().iter().cloned().collect::>(); + let (first, rest) = workers.split_first().expect("worker pool must not be empty"); let macros = first .find_proc_macros(&dylib.path)? diff --git a/crates/proc-macro-api/src/process.rs b/crates/proc-macro-api/src/process.rs index 035c12669c8f..2efcb54c9771 100644 --- a/crates/proc-macro-api/src/process.rs +++ b/crates/proc-macro-api/src/process.rs @@ -3,14 +3,16 @@ use std::{ fmt::Debug, io::{self, BufRead, BufReader, Read, Write}, - panic::AssertUnwindSafe, process::{Child, ChildStdin, ChildStdout, Command, Stdio}, sync::{ - Arc, Mutex, OnceLock, - atomic::{AtomicU32, Ordering}, + Arc, OnceLock, + atomic::{AtomicU8, AtomicU32, Ordering}, + mpsc, }, + time::{Duration, Instant}, }; +use parking_lot::Mutex; use paths::AbsPath; use semver::Version; use span::Span; @@ -32,10 +34,14 @@ pub(crate) struct ProcMacroServerProcess { /// The state of the proc-macro server process, the protocol is currently strictly sequential /// hence the lock on the state. state: Mutex, + /// The process handle and its health, shared with the watchdog thread so that it can kill + /// the process when an expansion times out. + control: Arc, + /// When set, each expansion is raced against the timeout by the watchdog, which kills the + /// process if it does not respond in time. + watchdog: Option<(ProcessWatchdog, Duration)>, version: u32, protocol: Protocol, - /// Populated when the server exits. - exited: OnceLock>, active: AtomicU32, } @@ -44,7 +50,7 @@ impl std::fmt::Debug for ProcMacroServerProcess { f.debug_struct("ProcMacroServerProcess") .field("version", &self.version) .field("protocol", &self.protocol) - .field("exited", &self.exited) + .field("exited", &self.control.exited) .finish() } } @@ -57,6 +63,170 @@ pub(crate) enum Protocol { pub trait ProcessExit: Send + Sync { fn exit_err(&mut self) -> Option; + fn kill(&mut self) -> io::Result<()>; +} + +/// Health of a server process with respect to expansion timeouts. +/// +/// The watchdog moves a process from `Healthy` to `TimedOut` when an expansion misses its +/// deadline. The pool claims a timed out process (`TimedOut` -> `BeingReplaced`) before spawning +/// its replacement, releasing the claim (back to `TimedOut`) if that fails so that a later +/// expansion can retry the replacement. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[repr(u8)] +enum ProcessStatus { + Healthy = 0, + TimedOut = 1, + BeingReplaced = 2, +} + +/// The part of a server process that is shared with the watchdog thread. +struct ProcessControl { + process: Mutex>, + /// Populated when the server exits, whether on its own or killed by the watchdog. + exited: OnceLock, + /// A [`ProcessStatus`]. + status: AtomicU8, +} + +impl ProcessControl { + fn new(process: Box) -> Arc { + Arc::new(ProcessControl { + process: Mutex::new(process), + exited: OnceLock::new(), + status: AtomicU8::new(ProcessStatus::Healthy as u8), + }) + } + + fn time_out(&self, macro_name: &str, timeout: Duration) { + let error = ServerError { + message: format!("proc-macro `{macro_name}` expansion timed out after {timeout:?}"), + io: None, + }; + if self.exited.set(error).is_ok() { + self.status.store(ProcessStatus::TimedOut as u8, Ordering::Release); + _ = self.process.lock().kill(); + } + } + + fn timed_out(&self) -> bool { + self.status.load(Ordering::Acquire) != ProcessStatus::Healthy as u8 + } + + fn claim_replacement(&self) -> bool { + self.status + .compare_exchange( + ProcessStatus::TimedOut as u8, + ProcessStatus::BeingReplaced as u8, + Ordering::AcqRel, + Ordering::Acquire, + ) + .is_ok() + } + + fn replacement_failed(&self) { + self.status.store(ProcessStatus::TimedOut as u8, Ordering::Release); + } +} + +/// A handle to the watchdog thread which kills server processes whose expansion requests +/// exceed their deadline. The thread exits once all handles to it have been dropped. +#[derive(Clone, Debug)] +pub(crate) struct ProcessWatchdog { + sender: mpsc::Sender, +} + +struct WatchdogRequest { + deadline: Instant, + timeout: Duration, + macro_name: Box, + control: Arc, +} + +enum WatchdogCommand { + Arm(WatchdogRequest), + Disarm(Arc), +} + +/// Disarms the pending watchdog request on drop. +struct WatchdogGuard { + control: Arc, + sender: mpsc::Sender, +} + +impl Drop for WatchdogGuard { + fn drop(&mut self) { + _ = self.sender.send(WatchdogCommand::Disarm(self.control.clone())); + } +} + +impl ProcessWatchdog { + pub(crate) fn spawn() -> io::Result { + let (sender, receiver) = mpsc::channel(); + std::thread::Builder::new() + .name("proc-macro-watchdog".into()) + .spawn(move || run_watchdog(receiver))?; + Ok(ProcessWatchdog { sender }) + } + + /// Arms a timeout for a single expansion; dropping the returned guard disarms it. + /// + /// A process has at most one request armed at a time as its I/O is strictly sequential, + /// which makes `control` a unique key for the request. + fn arm( + &self, + control: &Arc, + macro_name: &str, + timeout: Duration, + ) -> WatchdogGuard { + let guard = WatchdogGuard { control: control.clone(), sender: self.sender.clone() }; + let Some(deadline) = Instant::now().checked_add(timeout) else { + return guard; + }; + let request = WatchdogRequest { + deadline, + timeout, + macro_name: macro_name.into(), + control: control.clone(), + }; + if self.sender.send(WatchdogCommand::Arm(request)).is_err() { + stdx::never!("the proc-macro watchdog thread has died"); + } + guard + } +} + +fn run_watchdog(receiver: mpsc::Receiver) { + let mut requests = Vec::::new(); + loop { + let deadline = requests.iter().map(|request| request.deadline).min(); + let command = match deadline { + Some(deadline) => { + match receiver.recv_timeout(deadline.saturating_duration_since(Instant::now())) { + Ok(command) => Some(command), + Err(mpsc::RecvTimeoutError::Timeout) => None, + Err(mpsc::RecvTimeoutError::Disconnected) => return, + } + } + None => match receiver.recv() { + Ok(command) => Some(command), + Err(mpsc::RecvError) => return, + }, + }; + + match command { + Some(WatchdogCommand::Arm(request)) => requests.push(request), + Some(WatchdogCommand::Disarm(control)) => { + requests.retain(|request| !Arc::ptr_eq(&request.control, &control)); + } + None => { + let now = Instant::now(); + for request in requests.extract_if(.., |request| request.deadline <= now) { + request.control.time_out(&request.macro_name, request.timeout); + } + } + } + } } impl ProcessExit for Process { @@ -80,11 +250,14 @@ impl ProcessExit for Process { } } } + + fn kill(&mut self) -> io::Result<()> { + self.child.kill() + } } /// Maintains the state of the proc-macro server process. pub(crate) struct ProcessSrvState { - process: Box, stdin: Box, stdout: Box, } @@ -97,6 +270,7 @@ impl ProcMacroServerProcess { Item = (impl AsRef, &'a Option>), > + Clone, version: Option<&Version>, + watchdog: Option<(ProcessWatchdog, Duration)>, ) -> io::Result { Self::run( |format| { @@ -118,6 +292,7 @@ impl ProcMacroServerProcess { .map(|output| String::from_utf8_lossy(&output.stdout).trim().to_owned()) .unwrap_or_else(|_| "unknown version".to_owned()) }, + watchdog, ) } @@ -132,6 +307,7 @@ impl ProcMacroServerProcess { )>, version: Option<&Version>, binary_server_version: impl Fn() -> String, + watchdog: Option<(ProcessWatchdog, Duration)>, ) -> io::Result { const VERSION: Version = Version::new(1, 93, 0); // we do `>` for nightly as this started working in the middle of the 1.93 nightly release, so we dont want to break on half of the nightlies @@ -156,7 +332,9 @@ impl ProcMacroServerProcess { let (process, stdin, stdout) = spawn(format)?; io::Result::Ok(ProcMacroServerProcess { - state: Mutex::new(ProcessSrvState { process, stdin, stdout }), + state: Mutex::new(ProcessSrvState { stdin, stdout }), + control: ProcessControl::new(process), + watchdog: watchdog.clone(), version: 0, protocol: match format { Some(ProtocolFormat::BidirectionalPostcardPrototype) => { @@ -166,7 +344,6 @@ impl ProcMacroServerProcess { Protocol::LegacyJson { mode: SpanMode::Id } } }, - exited: OnceLock::new(), active: AtomicU32::new(0), }) }; @@ -229,7 +406,20 @@ impl ProcMacroServerProcess { /// Returns the server error if the process has exited. pub(crate) fn exited(&self) -> Option<&ServerError> { - self.exited.get().map(|it| &it.0) + self.control.exited.get() + } + + /// Whether the process was killed because an expansion timed out. + pub(crate) fn timed_out(&self) -> bool { + self.control.timed_out() + } + + pub(crate) fn claim_replacement(&self) -> bool { + self.control.claim_replacement() + } + + pub(crate) fn replacement_failed(&self) { + self.control.replacement_failed() } /// Retrieves the API version of the proc-macro server. @@ -323,8 +513,9 @@ impl ProcMacroServerProcess { &mut String, ) -> Result, ServerError>, req: Request, + macro_name: Option<&str>, ) -> Result { - self.with_locked_io(String::new(), |writer, reader, buf| { + self.with_locked_io(String::new(), macro_name, |writer, reader, buf| { send(writer, reader, req, buf).and_then(|res| { res.ok_or_else(|| { let message = "proc-macro server did not respond with data".to_owned(); @@ -343,19 +534,35 @@ impl ProcMacroServerProcess { fn with_locked_io( &self, mut buf: B, + macro_name: Option<&str>, f: impl FnOnce(&mut dyn Write, &mut dyn BufRead, &mut B) -> Result, ) -> Result { - let state = &mut *self.state.lock().unwrap(); - f(&mut state.stdin, &mut state.stdout, &mut buf).map_err(|e| { - if e.io.as_ref().map(|it| it.kind()) == Some(io::ErrorKind::BrokenPipe) { - match state.process.exit_err() { - None => e, - Some(server_error) => { - self.exited.get_or_init(|| AssertUnwindSafe(server_error)).0.clone() - } - } - } else { - e + let state = &mut *self.state.lock(); + let watchdog = match (&self.watchdog, macro_name) { + (Some((watchdog, timeout)), Some(macro_name)) => { + Some(watchdog.arm(&self.control, macro_name, *timeout)) + } + (Some(_), None) | (None, Some(_)) | (None, None) => None, + }; + let result = f(&mut state.stdin, &mut state.stdout, &mut buf); + drop(watchdog); + + if self.timed_out() { + return Err(self.exited().unwrap().clone()); + } + + result.map_err(|e| { + let process_exited = matches!( + e.io.as_ref().map(|it| it.kind()), + Some(io::ErrorKind::BrokenPipe | io::ErrorKind::UnexpectedEof) + ); + if !process_exited { + return e; + } + + match self.control.process.lock().exit_err() { + None => e, + Some(server_error) => self.control.exited.get_or_init(|| server_error).clone(), } }) } @@ -364,8 +571,9 @@ impl ProcMacroServerProcess { &self, initial: BidirectionalMessage, callback: SubCallback<'_>, + macro_name: Option<&str>, ) -> Result { - self.with_locked_io(Vec::new(), |writer, reader, buf| { + self.with_locked_io(Vec::new(), macro_name, |writer, reader, buf| { bidirectional_protocol::run_conversation(writer, reader, buf, initial, callback) }) } @@ -437,3 +645,179 @@ fn mk_child<'a>( } cmd.spawn() } + +#[cfg(test)] +mod tests { + use std::{ + io::Read, + sync::atomic::{AtomicBool, AtomicUsize}, + }; + + use parking_lot::Condvar; + + use crate::pool::{ProcMacroServerPool, ProcessFactory}; + + use super::*; + + struct FakeProcess { + killed: Arc, + wake: Arc<(Mutex, Condvar)>, + } + + impl ProcessExit for FakeProcess { + fn exit_err(&mut self) -> Option { + None + } + + fn kill(&mut self) -> io::Result<()> { + self.killed.store(true, Ordering::Release); + let (lock, wake) = &*self.wake; + *lock.lock() = true; + wake.notify_all(); + Ok(()) + } + } + + struct BlockingReader { + wake: Arc<(Mutex, Condvar)>, + } + + impl Read for BlockingReader { + fn read(&mut self, _buf: &mut [u8]) -> io::Result { + let (lock, wake) = &*self.wake; + let mut killed = lock.lock(); + while !*killed { + wake.wait(&mut killed); + } + Ok(0) + } + } + + impl BufRead for BlockingReader { + fn fill_buf(&mut self) -> io::Result<&[u8]> { + let mut byte = [0]; + _ = self.read(&mut byte)?; + Ok(&[]) + } + + fn consume(&mut self, _amt: usize) {} + } + + fn fake_process( + watchdog: Option<(ProcessWatchdog, Duration)>, + killed: Arc, + wake: Arc<(Mutex, Condvar)>, + ) -> ProcMacroServerProcess { + ProcMacroServerProcess { + state: Mutex::new(ProcessSrvState { + stdin: Box::new(io::sink()), + stdout: Box::new(BlockingReader { wake: wake.clone() }), + }), + control: ProcessControl::new(Box::new(FakeProcess { killed, wake })), + watchdog, + version: 0, + protocol: Protocol::LegacyJson { mode: SpanMode::Id }, + active: AtomicU32::new(0), + } + } + + #[test] + fn expansion_timeout_kills_process() { + let killed = Arc::new(AtomicBool::new(false)); + let wake = Arc::new((Mutex::new(false), Condvar::new())); + let watchdog = Some((ProcessWatchdog::spawn().unwrap(), Duration::from_millis(1000))); + let process = fake_process(watchdog, killed.clone(), wake); + + let result: Result<(), ServerError> = process.send_task_legacy( + |_writer, reader, (), _buf| { + let mut byte = [0]; + _ = reader.read(&mut byte).map_err(|error| ServerError { + message: "failed to read response".into(), + io: Some(Arc::new(error)), + })?; + Ok(None) + }, + (), + Some("hang"), + ); + + let error = result.unwrap_err(); + assert!(killed.load(Ordering::Acquire)); + assert!(process.timed_out()); + assert_eq!(error.message, "proc-macro `hang` expansion timed out after 10ms"); + } + + #[test] + fn completed_expansion_is_disarmed() { + let killed = Arc::new(AtomicBool::new(false)); + let watchdog = Some((ProcessWatchdog::spawn().unwrap(), Duration::from_millis(10))); + let process = + fake_process(watchdog, killed.clone(), Arc::new((Mutex::new(false), Condvar::new()))); + + let result: Result<(), ServerError> = + process.send_task_legacy(|_writer, _reader, (), _buf| Ok(Some(())), (), Some("fast")); + + result.unwrap(); + std::thread::sleep(Duration::from_millis(30)); + assert!(!killed.load(Ordering::Acquire)); + assert!(!process.timed_out()); + } + + #[test] + fn watchdog_times_out_multiple_processes() { + let watchdog = ProcessWatchdog::spawn().unwrap(); + let first_killed = Arc::new(AtomicBool::new(false)); + let first = + fake_process(None, first_killed.clone(), Arc::new((Mutex::new(false), Condvar::new()))); + let second_killed = Arc::new(AtomicBool::new(false)); + let second = fake_process( + None, + second_killed.clone(), + Arc::new((Mutex::new(false), Condvar::new())), + ); + + let first_guard = watchdog.arm(&first.control, "first", Duration::from_millis(10)); + let second_guard = watchdog.arm(&second.control, "second", Duration::from_millis(10)); + std::thread::sleep(Duration::from_millis(30)); + + assert!(first_killed.load(Ordering::Acquire)); + assert!(second_killed.load(Ordering::Acquire)); + drop((first_guard, second_guard)); + } + + #[test] + fn timed_out_process_is_replaced_in_background() { + let spawn_count = Arc::new(AtomicUsize::new(0)); + let spawn: ProcessFactory = Box::new({ + let spawn_count = spawn_count.clone(); + move || { + spawn_count.fetch_add(1, Ordering::AcqRel); + Ok(fake_process( + None, + Arc::new(AtomicBool::new(false)), + Arc::new((Mutex::new(false), Condvar::new())), + )) + } + }); + let pool = Arc::new(ProcMacroServerPool::new(vec![spawn().unwrap()], spawn)); + let process = pool.pick_process().unwrap(); + + process.control.time_out("hang", Duration::from_millis(10)); + assert!(process.timed_out()); + pool.replace_timed_out_process_in_background(process.clone()); + pool.replace_timed_out_process_in_background(process.clone()); + + let deadline = Instant::now() + Duration::from_secs(1); + let replacement = loop { + if let Ok(replacement) = pool.pick_process() + && !Arc::ptr_eq(&process, &replacement) + { + break replacement; + } + assert!(Instant::now() < deadline, "timed out waiting for replacement process"); + std::thread::yield_now(); + }; + assert!(!Arc::ptr_eq(&process, &replacement)); + assert_eq!(spawn_count.load(Ordering::Acquire), 2); + } +} diff --git a/crates/rust-analyzer/src/reload.rs b/crates/rust-analyzer/src/reload.rs index 3bf3cd562255..297d41089c4e 100644 --- a/crates/rust-analyzer/src/reload.rs +++ b/crates/rust-analyzer/src/reload.rs @@ -712,7 +712,14 @@ impl GlobalState { info!("Spawning proc-macro server at {path}"); let num_process = self.config.proc_macro_num_processes(); - Some(match ProcMacroClient::spawn(path, env, toolchain.as_ref(), num_process) { + let client = ProcMacroClient::spawn( + path, + env, + toolchain.as_ref(), + num_process, + Some(proc_macro_api::DEFAULT_EXPANSION_TIMEOUT), + ); + Some(match client { Ok(client) => { clients.push((key.clone(), client.clone())); Ok(client)