feat: harden frontend heartbeat extensions

Add a typed `Command::build_with_heartbeat_extensions` seam in `src/cmd/src/frontend.rs`.

Freeze and harden `FrontendHeartbeatExtensions` in `src/frontend/src/heartbeat.rs`, with lifecycle and race coverage in `src/frontend/src/heartbeat/tests.rs`.

Signed-off-by: Lei, HUANG <ratuthomm@gmail.com>
This commit is contained in:
Lei, HUANG
2026-08-08 11:40:55 +08:00
parent d90cca4b75
commit e589a77a90
3 changed files with 418 additions and 27 deletions
+190 -19
View File
@@ -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<frontend::frontend::FrontendOptions>;
type HeartbeatExtensionSetupFuture<'a> = Pin<Box<dyn Future<Output = Result<()>> + Send + 'a>>;
pub struct Instance {
frontend: Frontend,
@@ -117,7 +120,35 @@ pub struct Command {
impl Command {
pub async fn build(&self, opts: FrontendOptions) -> Result<Instance> {
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<Instance> {
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<FrontendOptions> {
@@ -131,9 +162,14 @@ pub enum SubCommand {
}
impl SubCommand {
async fn build(&self, opts: FrontendOptions) -> Result<Instance> {
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<Instance> {
async fn build_with_plugins(
&self,
opts: FrontendOptions,
plugins: Plugins,
) -> Result<Instance> {
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::<FrontendHeartbeatExtensions>()
.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> {
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<FrontendHeartbeatExtensions> {
setup(plugins).await?;
let extensions = plugins
.get::<FrontendHeartbeatExtensions>()
.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::<FrontendHeartbeatExtensions>()
.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::<Vec<_>>(),
["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<dyn frontend::heartbeat::FrontendHeartbeatExtension>),
&registered[0]
));
}
#[test]
fn test_try_from_start_command() {
let command = StartCommand {
+62 -7
View File
@@ -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<StdMutex<FrontendHeartbeatExtensionsInner>>,
}
/// 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<String>,
extensions: Vec<Arc<dyn FrontendHeartbeatExtension>>,
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<dyn FrontendHeartbeatExtension>) -> 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<dyn FrontendHeartbeatExtension>,
) -> 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<Arc<dyn FrontendHeartbeatExtension>> {
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.
+166 -1
View File
@@ -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<dyn FrontendHeartbeatExtension>),
&registered[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<dyn FrontendHeartbeatExtension>),
&registered[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<dyn FrontendHeartbeatExtension>),
&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<Barrier>,
release_name: Arc<Barrier>,
}
#[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<dyn HeartbeatConnector>,
extensions: FrontendHeartbeatExtensions,