diff --git a/src/cmd/src/frontend.rs b/src/cmd/src/frontend.rs index 9d5d0bab1e..c2cddfd588 100644 --- a/src/cmd/src/frontend.rs +++ b/src/cmd/src/frontend.rs @@ -13,7 +13,9 @@ // limitations under the License. use std::fmt::Debug; +use std::future::Future; use std::path::Path; +use std::pin::Pin; use std::sync::Arc; use std::time::Duration; @@ -61,6 +63,7 @@ use crate::options::{GlobalOptions, GreptimeOptions}; use crate::{App, create_resource_limit_metrics, log_versions, maybe_activate_heap_profile}; type FrontendOptions = GreptimeOptions; +type HeartbeatExtensionSetupFuture<'a> = Pin> + Send + 'a>>; pub struct Instance { frontend: Frontend, @@ -117,7 +120,35 @@ pub struct Command { impl Command { pub async fn build(&self, opts: FrontendOptions) -> Result { - self.subcmd.build(opts).await + self.build_with_heartbeat_extensions(opts, FrontendHeartbeatExtensions::default()) + .await + } + + /// Builds a frontend with pre-registered heartbeat extensions. + /// + /// Register extensions before calling this method. The supplied registry is installed before + /// normal plugin setup. After plugin heartbeat setup completes, the registry is frozen before + /// heartbeat handlers and the heartbeat task are built, so later registrations are rejected + /// and every heartbeat consumer observes the same extension membership. + pub async fn build_with_heartbeat_extensions( + &self, + opts: FrontendOptions, + heartbeat_extensions: FrontendHeartbeatExtensions, + ) -> Result { + let plugins = plugins_with_heartbeat_extensions(heartbeat_extensions); + self.forward_build_with_plugins(opts, plugins, |command, opts, plugins| { + command.build_with_plugins(opts, plugins) + }) + .await + } + + fn forward_build_with_plugins<'a, T>( + &'a self, + opts: FrontendOptions, + plugins: Plugins, + build: impl FnOnce(&'a StartCommand, FrontendOptions, Plugins) -> T, + ) -> T { + self.subcmd.forward_build_with_plugins(opts, plugins, build) } pub fn load_options(&self, global_options: &GlobalOptions) -> Result { @@ -131,9 +162,14 @@ pub enum SubCommand { } impl SubCommand { - async fn build(&self, opts: FrontendOptions) -> Result { + fn forward_build_with_plugins<'a, T>( + &'a self, + opts: FrontendOptions, + plugins: Plugins, + build: impl FnOnce(&'a StartCommand, FrontendOptions, Plugins) -> T, + ) -> T { match self { - SubCommand::Start(cmd) => cmd.build(opts).await, + SubCommand::Start(cmd) => build(cmd, opts, plugins), } } @@ -329,7 +365,11 @@ impl StartCommand { Ok(()) } - async fn build(&self, opts: FrontendOptions) -> Result { + async fn build_with_plugins( + &self, + opts: FrontendOptions, + plugins: Plugins, + ) -> Result { common_runtime::init_global_runtimes(&opts.runtime); let guard = common_telemetry::init_global_logging( @@ -380,15 +420,8 @@ impl StartCommand { .await .context(error::MetaClientInitSnafu)?; - let mut plugins = Plugins::new(); - plugins::setup_frontend_plugins_pre_build( - &mut plugins, - &plugin_opts, - &opts, - Some(&meta_config), - ) - .await - .context(error::StartFrontendSnafu)?; + let mut plugins = + prepare_frontend_plugins(plugins, &plugin_opts, &opts, Some(&meta_config)).await?; // now initialize the meta_client with plugins let meta_client = meta_client::create_meta_client( @@ -497,12 +530,20 @@ impl StartCommand { let instance = Arc::new(instance); - plugins::setup_frontend_heartbeat_extensions(&mut plugins, &plugin_opts, &instance) - .await - .context(error::StartFrontendSnafu)?; - let heartbeat_extensions = plugins - .get::() - .unwrap_or_default(); + let heartbeat_instance = instance.clone(); + let heartbeat_extensions = + setup_and_freeze_frontend_heartbeat_extensions(&mut plugins, move |plugins| { + Box::pin(async move { + plugins::setup_frontend_heartbeat_extensions( + plugins, + &plugin_opts, + &heartbeat_instance, + ) + .await + .context(error::StartFrontendSnafu) + }) + }) + .await?; let heartbeat_task = Some(create_heartbeat_task_with_extensions( &opts, meta_client, @@ -524,6 +565,36 @@ impl StartCommand { } } +async fn prepare_frontend_plugins( + mut plugins: Plugins, + plugin_opts: &[PluginOptions], + opts: &frontend::frontend::FrontendOptions, + meta_config: Option<&[PluginOptions]>, +) -> Result { + plugins::setup_frontend_plugins_pre_build(&mut plugins, plugin_opts, opts, meta_config) + .await + .context(error::StartFrontendSnafu)?; + Ok(plugins) +} + +fn plugins_with_heartbeat_extensions(heartbeat_extensions: FrontendHeartbeatExtensions) -> Plugins { + let plugins = Plugins::new(); + plugins.insert(heartbeat_extensions); + plugins +} + +async fn setup_and_freeze_frontend_heartbeat_extensions( + plugins: &mut Plugins, + setup: impl for<'a> FnOnce(&'a mut Plugins) -> HeartbeatExtensionSetupFuture<'a>, +) -> Result { + setup(plugins).await?; + let extensions = plugins + .get::() + .unwrap_or_default(); + extensions.freeze(); + Ok(extensions) +} + pub fn create_heartbeat_task( options: &frontend::frontend::FrontendOptions, meta_client: MetaClientRef, @@ -563,6 +634,7 @@ fn create_heartbeat_task_with_extensions( #[cfg(test)] mod tests { use std::io::Write; + use std::sync::Arc; use std::time::Duration; use auth::{Identity, Password, UserProviderRef}; @@ -576,6 +648,105 @@ mod tests { use super::*; use crate::options::GlobalOptions; + #[derive(Debug)] + struct TestExtension(&'static str); + + #[async_trait] + impl frontend::heartbeat::FrontendHeartbeatExtension for TestExtension { + fn name(&self) -> &'static str { + self.0 + } + } + + #[test] + fn test_command_forwards_heartbeat_extensions_to_start_build() { + let command = Command { + subcmd: SubCommand::Start(StartCommand::default()), + }; + let extensions = FrontendHeartbeatExtensions::default(); + let shared_extensions = extensions.clone(); + + let forwarded_plugins = command.forward_build_with_plugins( + FrontendOptions::default(), + plugins_with_heartbeat_extensions(extensions), + |_, _, plugins| plugins, + ); + let forwarded_extensions = forwarded_plugins + .get::() + .unwrap(); + + assert_registry_identity(&shared_extensions, &forwarded_extensions); + } + + #[tokio::test] + async fn test_prefilled_heartbeat_extensions_survive_pre_build_setup() { + let extensions = FrontendHeartbeatExtensions::default(); + let shared_extensions = extensions.clone(); + + let plugins = prepare_frontend_plugins( + plugins_with_heartbeat_extensions(extensions), + &[], + &frontend::frontend::FrontendOptions::default(), + None, + ) + .await + .unwrap(); + + let setup_extensions = plugins.get_or_insert(FrontendHeartbeatExtensions::default); + assert_registry_identity(&shared_extensions, &setup_extensions); + } + + #[tokio::test] + async fn test_heartbeat_extension_setup_completes_before_freeze() { + let extensions = FrontendHeartbeatExtensions::default(); + assert_eq!( + extensions.try_register(Arc::new(TestExtension("caller"))), + Ok(()) + ); + let mut plugins = plugins_with_heartbeat_extensions(extensions.clone()); + + let finalized = setup_and_freeze_frontend_heartbeat_extensions(&mut plugins, |plugins| { + Box::pin(async move { + let setup_extensions = plugins.get_or_insert(FrontendHeartbeatExtensions::default); + assert_eq!( + setup_extensions.try_register(Arc::new(TestExtension("plugin"))), + Ok(()) + ); + Ok(()) + }) + }) + .await + .unwrap(); + + assert_eq!( + finalized + .extensions() + .iter() + .map(|extension| extension.name()) + .collect::>(), + ["caller", "plugin"] + ); + assert_eq!( + extensions.try_register(Arc::new(TestExtension("late"))), + Err(frontend::heartbeat::RegistrationError::Frozen) + ); + assert_eq!(finalized.len(), 2); + } + + fn assert_registry_identity( + expected: &FrontendHeartbeatExtensions, + actual: &FrontendHeartbeatExtensions, + ) { + let extension = Arc::new(TestExtension("cmd-test-extension")); + assert!(expected.register(extension.clone())); + let registered = actual.extensions(); + assert_eq!(registered.len(), 1); + assert!(Arc::ptr_eq( + &(extension as Arc), + ®istered[0] + )); + } + #[test] fn test_try_from_start_command() { let command = StartCommand { diff --git a/src/frontend/src/heartbeat.rs b/src/frontend/src/heartbeat.rs index 0d9cde74fb..b7e3a7aff4 100644 --- a/src/frontend/src/heartbeat.rs +++ b/src/frontend/src/heartbeat.rs @@ -16,8 +16,9 @@ mod tests; use std::collections::{HashMap, HashSet}; +use std::fmt::{Display, Formatter}; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; -use std::sync::{Arc, Mutex as StdMutex}; +use std::sync::{Arc, Mutex as StdMutex, MutexGuard, PoisonError}; use api::v1::meta::heartbeat_request::NodeWorkloads; use api::v1::meta::{FrontendWorkloads, HeartbeatRequest, HeartbeatResponse, NodeInfo, Peer}; @@ -93,34 +94,88 @@ pub trait FrontendHeartbeatExtension: Send + Sync { } /// A shareable, ordered registry of frontend heartbeat extensions. +/// +/// Registration is available until [`FrontendHeartbeatExtensions::freeze`] is called. After that, +/// membership is immutable so every heartbeat consumer observes the same extensions. #[derive(Clone, Default)] pub struct FrontendHeartbeatExtensions { inner: Arc>, } +/// The reason a frontend heartbeat extension was not registered. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum RegistrationError { + /// An extension with the same name is already registered. + Duplicate, + /// The registry is frozen and no longer accepts registrations. + Frozen, +} + +impl Display for RegistrationError { + fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result { + match self { + Self::Duplicate => formatter.write_str("heartbeat extension is already registered"), + Self::Frozen => formatter.write_str("heartbeat extension registry is frozen"), + } + } +} + +impl std::error::Error for RegistrationError {} + #[derive(Default)] struct FrontendHeartbeatExtensionsInner { names: HashSet, extensions: Vec>, + frozen: bool, } impl FrontendHeartbeatExtensions { + fn lock(&self) -> MutexGuard<'_, FrontendHeartbeatExtensionsInner> { + self.inner.lock().unwrap_or_else(PoisonError::into_inner) + } + /// Registers an extension without replacing an existing extension of the same name. /// - /// Returns `true` for a new registration and `false` for an idempotent duplicate. + /// Returns `true` for a new registration and `false` for an idempotent duplicate or when the + /// registry is frozen. Use [`FrontendHeartbeatExtensions::try_register`] when the rejection + /// reason must not be discarded. pub fn register(&self, extension: Arc) -> bool { - let mut inner = self.inner.lock().unwrap(); + self.try_register(extension).is_ok() + } + + /// Registers an extension and returns the reason when registration is rejected. + /// + /// If the registry is already frozen, this method returns [`RegistrationError::Frozen`] + /// without calling [`FrontendHeartbeatExtension::name`]. If freezing races with name + /// evaluation, the registry state is checked again before insertion. + pub fn try_register( + &self, + extension: Arc, + ) -> std::result::Result<(), RegistrationError> { + if self.lock().frozen { + return Err(RegistrationError::Frozen); + } + let name = extension.name().to_string(); + let mut inner = self.lock(); + if inner.frozen { + return Err(RegistrationError::Frozen); + } if !inner.names.insert(name) { - return false; + return Err(RegistrationError::Duplicate); } inner.extensions.push(extension); - true + Ok(()) + } + + /// Prevents further registrations while preserving the current registration order. + pub fn freeze(&self) { + self.lock().frozen = true; } /// Returns the registered extensions in registration order. pub fn extensions(&self) -> Vec> { - self.inner.lock().unwrap().extensions.clone() + self.lock().extensions.clone() } /// Returns response handlers in extension registration order. @@ -133,7 +188,7 @@ impl FrontendHeartbeatExtensions { /// Returns the number of registered extensions. pub fn len(&self) -> usize { - self.inner.lock().unwrap().extensions.len() + self.lock().extensions.len() } /// Returns whether no extension is registered. diff --git a/src/frontend/src/heartbeat/tests.rs b/src/frontend/src/heartbeat/tests.rs index d3c651de32..8f5bfa614b 100644 --- a/src/frontend/src/heartbeat/tests.rs +++ b/src/frontend/src/heartbeat/tests.rs @@ -14,8 +14,9 @@ use std::collections::{HashMap, VecDeque}; use std::future::pending; +use std::panic::{AssertUnwindSafe, catch_unwind}; use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; -use std::sync::{Arc, Mutex}; +use std::sync::{Arc, Barrier, Mutex}; use std::time::Duration; use api::v1::meta::heartbeat_request::NodeWorkloads; @@ -494,6 +495,170 @@ impl FrontendHeartbeatExtension for TestExtension { } } +struct PanickingNameExtension; + +#[async_trait] +impl FrontendHeartbeatExtension for PanickingNameExtension { + fn name(&self) -> &str { + panic!("extension name panic") + } +} + +#[test] +fn test_panicking_extension_name_does_not_corrupt_registry() { + let extensions = FrontendHeartbeatExtensions::default(); + + let result = catch_unwind(AssertUnwindSafe(|| { + extensions.register(Arc::new(PanickingNameExtension)); + })); + assert!(result.is_err()); + + let valid = TestExtension::new("valid"); + assert!(extensions.register(valid.clone())); + let registered = extensions.extensions(); + assert_eq!(registered.len(), 1); + assert!(Arc::ptr_eq( + &(valid as Arc), + ®istered[0] + )); + assert_eq!(extensions.len(), 1); +} + +#[test] +fn test_registry_recovers_from_poisoned_lock() { + let extensions = FrontendHeartbeatExtensions::default(); + let poisoned = extensions.clone(); + + let result = catch_unwind(AssertUnwindSafe(move || { + let _guard = poisoned.lock(); + panic!("poison registry lock") + })); + assert!(result.is_err()); + + let valid = TestExtension::new("valid"); + assert!(extensions.register(valid.clone())); + let registered = extensions.extensions(); + assert_eq!(registered.len(), 1); + assert!(Arc::ptr_eq( + &(valid as Arc), + ®istered[0] + )); + assert_eq!(extensions.len(), 1); +} + +#[test] +fn test_registry_freeze_rejects_late_registration_without_changing_membership() { + let extensions = FrontendHeartbeatExtensions::default(); + let calls = Arc::new(AtomicUsize::new(0)); + let handler: HeartbeatResponseHandlerRef = Arc::new(ShortCircuitHandler { + result: ShortCircuitResult::Done, + calls, + }); + let registered = TestExtension::with_handler("registered", handler.clone()); + + assert!(extensions.register(registered.clone())); + assert!(!extensions.register(TestExtension::new("registered"))); + extensions.freeze(); + + assert!(!extensions.register(TestExtension::new("late"))); + let frozen_extensions = extensions.extensions(); + assert_eq!(frozen_extensions.len(), 1); + assert!(Arc::ptr_eq( + &(registered as Arc), + &frozen_extensions[0] + )); + let frozen_handlers = extensions.response_handlers(); + assert_eq!(frozen_handlers.len(), 1); + assert!(Arc::ptr_eq(&handler, &frozen_handlers[0])); + assert_eq!(extensions.len(), 1); +} + +#[test] +fn test_try_register_distinguishes_duplicate_and_frozen() { + let extensions = FrontendHeartbeatExtensions::default(); + + assert_eq!( + extensions.try_register(TestExtension::new("registered")), + Ok(()) + ); + assert_eq!( + extensions.try_register(TestExtension::new("registered")), + Err(RegistrationError::Duplicate) + ); + + extensions.freeze(); + assert_eq!( + extensions.try_register(TestExtension::new("late")), + Err(RegistrationError::Frozen) + ); +} + +#[test] +fn test_frozen_registration_does_not_call_extension_name() { + let extensions = FrontendHeartbeatExtensions::default(); + extensions.freeze(); + + let compatible_result = catch_unwind(AssertUnwindSafe(|| { + extensions.register(Arc::new(PanickingNameExtension)) + })); + assert!(!compatible_result.unwrap()); + + let typed_result = catch_unwind(AssertUnwindSafe(|| { + extensions.try_register(Arc::new(PanickingNameExtension)) + })); + + assert_eq!(typed_result.unwrap(), Err(RegistrationError::Frozen)); + assert!(extensions.is_empty()); +} + +struct BlockingNameExtension { + entered_name: Arc, + release_name: Arc, +} + +#[async_trait] +impl FrontendHeartbeatExtension for BlockingNameExtension { + fn name(&self) -> &str { + self.entered_name.wait(); + self.release_name.wait(); + "blocking" + } +} + +#[test] +fn test_freeze_wins_while_registration_name_is_running() { + let extensions = FrontendHeartbeatExtensions::default(); + let entered_name = Arc::new(Barrier::new(2)); + let release_name = Arc::new(Barrier::new(2)); + let registering = extensions.clone(); + let extension = Arc::new(BlockingNameExtension { + entered_name: entered_name.clone(), + release_name: release_name.clone(), + }); + + let registration = std::thread::spawn(move || registering.try_register(extension)); + entered_name.wait(); + let freezing = extensions.clone(); + let (freeze_done_tx, freeze_done_rx) = std::sync::mpsc::sync_channel(1); + let freeze = std::thread::spawn(move || { + freezing.freeze(); + let _ = freeze_done_tx.send(()); + }); + if let Err(error) = freeze_done_rx.recv_timeout(Duration::from_secs(2)) { + release_name.wait(); + let registration = registration.join(); + let freeze = freeze.join(); + panic!( + "freeze did not complete while extension name was blocked: {error:?}; registration: {registration:?}; freeze: {freeze:?}" + ); + } + release_name.wait(); + + freeze.join().unwrap(); + assert_eq!(registration.join().unwrap(), Err(RegistrationError::Frozen)); + assert!(extensions.is_empty()); +} + fn test_task( connector: Arc, extensions: FrontendHeartbeatExtensions,