diff --git a/Cargo.lock b/Cargo.lock index bc59d90af5..c8fa30c9d7 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1623,7 +1623,7 @@ dependencies = [ "maybe-owned", "rustix 1.0.7", "rustix-linux-procfs", - "windows-sys 0.61.2", + "windows-sys 0.60.2", "winx", ] @@ -2259,7 +2259,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "117725a109d387c937a1533ce01b450cbde6b88abceea8473c4d7a85853cda3c" dependencies = [ "lazy_static", - "windows-sys 0.59.0", + "windows-sys 0.48.0", ] [[package]] @@ -6041,7 +6041,7 @@ dependencies = [ [[package]] name = "greptime-proto" version = "0.1.0" -source = "git+https://github.com/GreptimeTeam/greptime-proto.git?rev=5adb8a1637abe87bcbb455d32136eee5d538cc49#5adb8a1637abe87bcbb455d32136eee5d538cc49" +source = "git+https://github.com/GreptimeTeam/greptime-proto.git?rev=032510ded061277b7d1bbb29d4046abb5b1cbf4b#032510ded061277b7d1bbb29d4046abb5b1cbf4b" dependencies = [ "prost 0.14.1", "prost-types 0.14.1", @@ -6587,7 +6587,7 @@ dependencies = [ "libc", "percent-encoding", "pin-project-lite", - "socket2 0.6.4", + "socket2 0.5.10", "tokio", "tower-service", "tracing", @@ -9095,7 +9095,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -11284,7 +11284,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ac6c3320f9abac597dcbc668774ef006702672474aad53c6d596b62e487b40b1" dependencies = [ "heck 0.5.0", - "itertools 0.14.0", + "itertools 0.10.5", "log", "multimap", "once_cell", @@ -11332,7 +11332,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8a56d757972c98b346a9b766e3f02746cde6dd1cd1d1d563472929fdd74bec4d" dependencies = [ "anyhow", - "itertools 0.14.0", + "itertools 0.10.5", "proc-macro2", "quote", "syn 2.0.117", @@ -11345,7 +11345,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9120690fafc389a67ba3803df527d0ec9cbbc9cc45e4cc20b332996dfb672425" dependencies = [ "anyhow", - "itertools 0.14.0", + "itertools 0.10.5", "proc-macro2", "quote", "syn 2.0.117", @@ -12800,7 +12800,7 @@ dependencies = [ "security-framework 3.7.0", "security-framework-sys", "webpki-root-certs", - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -13712,7 +13712,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "52d1cfed4120b4d927bf7c0f86d2087a4a7d6027c906d9f9d525a80573b9be51" dependencies = [ "libc", - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -14706,7 +14706,7 @@ dependencies = [ "getrandom 0.3.4", "once_cell", "rustix 1.0.7", - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -16459,7 +16459,7 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf221c93e13a30d793f7645a0e7762c55d169dbb0a49671918a2319d289b10bb" dependencies = [ - "windows-sys 0.59.0", + "windows-sys 0.48.0", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 711423b0aa..323f014c5c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -159,7 +159,7 @@ fs2 = "0.4" fst = "0.4.7" futures = "0.3" futures-util = "0.3" -greptime-proto = { git = "https://github.com/GreptimeTeam/greptime-proto.git", rev = "5adb8a1637abe87bcbb455d32136eee5d538cc49" } +greptime-proto = { git = "https://github.com/GreptimeTeam/greptime-proto.git", rev = "032510ded061277b7d1bbb29d4046abb5b1cbf4b" } hex = "0.4" http = "1" humantime = "2.1" diff --git a/src/cmd/src/frontend.rs b/src/cmd/src/frontend.rs index 3a314734c9..9d5d0bab1e 100644 --- a/src/cmd/src/frontend.rs +++ b/src/cmd/src/frontend.rs @@ -32,10 +32,6 @@ use common_base::Plugins; use common_config::{Configurable, DEFAULT_DATA_HOME}; use common_error::ext::BoxedError; use common_meta::cache::{CacheRegistryBuilder, LayeredCacheRegistryBuilder}; -use common_meta::heartbeat::handler::HandlerGroupExecutor; -use common_meta::heartbeat::handler::invalidate_table_cache::InvalidateCacheHandler; -use common_meta::heartbeat::handler::parse_mailbox_message::ParseMailboxMessageHandler; -use common_meta::heartbeat::handler::suspend::SuspendHandler; use common_query::prelude::set_default_prefix; use common_stat::ResourceStatImpl; use common_telemetry::info; @@ -43,7 +39,9 @@ use common_telemetry::logging::{DEFAULT_LOGGING_DIR, TracingOptions}; use common_time::timezone::set_default_timezone; use common_version::{short_version, verbose_version}; use frontend::frontend::Frontend; -use frontend::heartbeat::HeartbeatTask; +use frontend::heartbeat::{ + FrontendHeartbeatExtensions, HeartbeatTask, heartbeat_response_handler_executor, +}; use frontend::instance::builder::FrontendBuilder; use frontend::server::Services; use meta_client::{MetaClientOptions, MetaClientRef, MetaClientType}; @@ -497,10 +495,21 @@ impl StartCommand { .await .context(error::StartFrontendSnafu)?; - let heartbeat_task = Some(create_heartbeat_task(&opts, meta_client, &instance)); - 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_task = Some(create_heartbeat_task_with_extensions( + &opts, + meta_client, + &instance, + heartbeat_extensions, + )); + let servers = Services::new(opts, instance.clone(), plugins) .build() .context(error::StartFrontendSnafu)?; @@ -520,13 +529,20 @@ pub fn create_heartbeat_task( meta_client: MetaClientRef, instance: &frontend::instance::Instance, ) -> HeartbeatTask { - let executor = Arc::new(HandlerGroupExecutor::new(vec![ - Arc::new(ParseMailboxMessageHandler), - Arc::new(SuspendHandler::new(instance.suspend_state())), - Arc::new(InvalidateCacheHandler::new( - instance.cache_invalidator().clone(), - )), - ])); + create_heartbeat_task_with_extensions(options, meta_client, instance, Default::default()) +} + +fn create_heartbeat_task_with_extensions( + options: &frontend::frontend::FrontendOptions, + meta_client: MetaClientRef, + instance: &frontend::instance::Instance, + extensions: FrontendHeartbeatExtensions, +) -> HeartbeatTask { + let executor = heartbeat_response_handler_executor( + &extensions, + instance.suspend_state(), + instance.cache_invalidator().clone(), + ); let stat = { let mut stat = ResourceStatImpl::default(); @@ -541,6 +557,7 @@ pub fn create_heartbeat_task( executor, stat, ) + .with_extensions(extensions) } #[cfg(test)] diff --git a/src/frontend/src/frontend.rs b/src/frontend/src/frontend.rs index 2616846b9f..30fcd39492 100644 --- a/src/frontend/src/frontend.rs +++ b/src/frontend/src/frontend.rs @@ -130,17 +130,26 @@ pub struct Frontend { impl Frontend { pub async fn start(&mut self) -> Result<()> { - if let Some(t) = &self.heartbeat_task { - t.start().await?; + if let Some(t) = &self.heartbeat_task + && let Err(error) = t.start().await + { + t.shutdown().await; + return Err(error); } - self.servers - .start_all() - .await - .context(error::StartServerSnafu) + if let Err(source) = self.servers.start_all().await { + if let Some(t) = &self.heartbeat_task { + t.shutdown().await; + } + return Err(source).context(error::StartServerSnafu); + } + Ok(()) } pub async fn shutdown(&mut self) -> Result<()> { + if let Some(t) = &self.heartbeat_task { + t.shutdown().await; + } self.servers .shutdown_all() .await @@ -154,7 +163,9 @@ impl Frontend { #[cfg(test)] mod tests { - use std::sync::atomic::{AtomicBool, Ordering}; + use std::any::Any; + use std::net::SocketAddr; + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::time::Duration; use api::v1::meta::heartbeat_server::HeartbeatServer; @@ -181,6 +192,7 @@ mod tests { use servers::grpc::{FlightCompression, GRPC_SERVER}; use servers::http::HTTP_SERVER; use servers::http::result::greptime_result_v1::GreptimedbV1Response; + use servers::server::Server; use tokio::sync::mpsc; use tonic::codec::CompressionEncoding; use tonic::codegen::tokio_stream::StreamExt; @@ -188,6 +200,9 @@ mod tests { use tonic::{Request, Response, Status, Streaming}; use super::*; + use crate::heartbeat::{ + FrontendHeartbeatExtension, FrontendHeartbeatExtensionResult, FrontendHeartbeatExtensions, + }; use crate::instance::builder::FrontendBuilder; use crate::server::Services; @@ -209,6 +224,46 @@ mod tests { struct SuspendableHeartbeatServer { suspend: Arc, + fail_heartbeat: bool, + } + + struct FailingServer; + + struct ShutdownTrackingExtension { + shutdown_calls: AtomicUsize, + } + + #[async_trait] + impl FrontendHeartbeatExtension for ShutdownTrackingExtension { + fn name(&self) -> &str { + "shutdown-tracking" + } + + async fn shutdown(&self) -> FrontendHeartbeatExtensionResult<()> { + self.shutdown_calls.fetch_add(1, Ordering::AcqRel); + Ok(()) + } + } + + #[async_trait] + impl Server for FailingServer { + async fn shutdown(&self) -> servers::error::Result<()> { + Ok(()) + } + + async fn start(&mut self, _listening: SocketAddr) -> servers::error::Result<()> { + Err(servers::error::Error::Internal { + err_msg: "mock server start failure".to_string(), + }) + } + + fn name(&self) -> &str { + "FAILING_SERVER" + } + + fn as_any(&self) -> &dyn Any { + self + } } #[async_trait] @@ -219,6 +274,10 @@ mod tests { &self, request: Request>, ) -> std::result::Result, Status> { + if self.fail_heartbeat { + return Err(Status::unavailable("mock initial heartbeat failure")); + } + let (tx, rx) = mpsc::channel(4); common_runtime::spawn_global({ @@ -358,6 +417,93 @@ mod tests { Ok(frontend) } + #[tokio::test] + async fn test_server_start_failure_shuts_down_heartbeat() { + let meta_client_options = MetaClientOptions { + metasrv_addrs: vec!["localhost:0".to_string()], + ..Default::default() + }; + let options = FrontendOptions { + meta_client: Some(meta_client_options.clone()), + ..Default::default() + }; + let heartbeat_server = Arc::new(SuspendableHeartbeatServer { + suspend: Arc::new(AtomicBool::new(false)), + fail_heartbeat: false, + }); + let meta_client = create_meta_client(&meta_client_options, heartbeat_server).await; + let instance = Arc::new( + FrontendBuilder::new_test(&options, meta_client.clone()) + .try_build() + .await + .unwrap(), + ); + let heartbeat_task = HeartbeatTask::new( + instance.frontend_peer_addr().to_string(), + &options, + meta_client, + Arc::new(HandlerGroupExecutor::new(vec![])), + Arc::new(ResourceStatImpl::default()), + ); + let heartbeat_probe = heartbeat_task.clone(); + let servers = ServerHandlers::default(); + servers.insert((Box::new(FailingServer), "127.0.0.1:0".parse().unwrap())); + let mut frontend = Frontend { + instance, + servers, + heartbeat_task: Some(heartbeat_task), + }; + + assert!(frontend.start().await.is_err()); + assert!(heartbeat_probe.is_shutdown()); + } + + #[tokio::test] + async fn test_heartbeat_start_failure_shuts_down_extensions() { + let meta_client_options = MetaClientOptions { + metasrv_addrs: vec!["localhost:0".to_string()], + ..Default::default() + }; + let options = FrontendOptions { + meta_client: Some(meta_client_options.clone()), + ..Default::default() + }; + let heartbeat_server = Arc::new(SuspendableHeartbeatServer { + suspend: Arc::new(AtomicBool::new(false)), + fail_heartbeat: true, + }); + let meta_client = create_meta_client(&meta_client_options, heartbeat_server).await; + let instance = Arc::new( + FrontendBuilder::new_test(&options, meta_client.clone()) + .try_build() + .await + .unwrap(), + ); + let extension = Arc::new(ShutdownTrackingExtension { + shutdown_calls: AtomicUsize::new(0), + }); + let extensions = FrontendHeartbeatExtensions::default(); + assert!(extensions.register(extension.clone())); + let heartbeat_task = HeartbeatTask::new( + instance.frontend_peer_addr().to_string(), + &options, + meta_client, + Arc::new(HandlerGroupExecutor::new(vec![])), + Arc::new(ResourceStatImpl::default()), + ) + .with_extensions(extensions); + let heartbeat_probe = heartbeat_task.clone(); + let mut frontend = Frontend { + instance, + servers: ServerHandlers::default(), + heartbeat_task: Some(heartbeat_task), + }; + + assert!(frontend.start().await.is_err()); + assert!(heartbeat_probe.is_shutdown()); + assert_eq!(extension.shutdown_calls.load(Ordering::Acquire), 1); + } + async fn verify_suspend_state_by_http( frontend: &Frontend, expected: std::result::Result<&str, (StatusCode, &str)>, @@ -454,6 +600,7 @@ mod tests { let server = Arc::new(SuspendableHeartbeatServer { suspend: Arc::new(AtomicBool::new(false)), + fail_heartbeat: false, }); let meta_client = create_meta_client(&meta_client_options, server.clone()).await; let frontend = create_frontend(&options, meta_client).await?; diff --git a/src/frontend/src/heartbeat.rs b/src/frontend/src/heartbeat.rs index 4dca46d9fb..0d9cde74fb 100644 --- a/src/frontend/src/heartbeat.rs +++ b/src/frontend/src/heartbeat.rs @@ -15,119 +15,481 @@ #[cfg(test)] mod tests; -use std::sync::Arc; +use std::collections::{HashMap, HashSet}; +use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; +use std::sync::{Arc, Mutex as StdMutex}; use api::v1::meta::heartbeat_request::NodeWorkloads; -use api::v1::meta::{FrontendWorkloads, HeartbeatRequest, NodeInfo, Peer}; +use api::v1::meta::{FrontendWorkloads, HeartbeatRequest, HeartbeatResponse, NodeInfo, Peer}; +use async_trait::async_trait; +use common_error::ext::BoxedError; +use common_meta::cache_invalidator::CacheInvalidatorRef; use common_meta::datanode::EnvVars; +use common_meta::heartbeat::handler::invalidate_table_cache::InvalidateCacheHandler; +use common_meta::heartbeat::handler::parse_mailbox_message::ParseMailboxMessageHandler; +use common_meta::heartbeat::handler::suspend::SuspendHandler; use common_meta::heartbeat::handler::{ - HeartbeatResponseHandlerContext, HeartbeatResponseHandlerExecutorRef, + HandlerGroupExecutor, HeartbeatResponseHandlerContext, HeartbeatResponseHandlerExecutorRef, + HeartbeatResponseHandlerRef, }; use common_meta::heartbeat::mailbox::{HeartbeatMailbox, MailboxRef, OutgoingMessage}; use common_meta::heartbeat::utils::outgoing_message_to_mailbox_message; use common_stat::ResourceStatRef; use common_telemetry::{debug, error, info, warn}; +use meta_client::client::heartbeat::HeartbeatConfig; use meta_client::client::{HeartbeatSender, HeartbeatStream, MetaClient}; use servers::addrs; use snafu::ResultExt; -use tokio::sync::mpsc; use tokio::sync::mpsc::Receiver; +use tokio::sync::{Mutex, mpsc}; use tokio::time::{Duration, Instant}; +use tokio_util::sync::CancellationToken; use crate::error; use crate::error::Result; use crate::frontend::FrontendOptions; use crate::metrics::{HEARTBEAT_RECV_COUNT, HEARTBEAT_SENT_COUNT}; -/// The frontend heartbeat task which sending `[HeartbeatRequest]` to Metasrv periodically in background. +/// The result type returned by a [`FrontendHeartbeatExtension`]. +pub type FrontendHeartbeatExtensionResult = std::result::Result; + +/// An extension to frontend heartbeat requests, responses, and lifecycle events. +/// +/// [`FrontendHeartbeatExtension::request_extensions`] is called for every heartbeat. A failed call +/// is isolated from the base heartbeat and from other extensions. +/// [`FrontendHeartbeatExtension::connected`] is called once for every successfully established +/// connection generation, including reconnects; implementations must therefore be idempotent. +/// [`FrontendHeartbeatExtension::shutdown`] is called once during task shutdown after heartbeat +/// I/O has stopped. +#[async_trait] +pub trait FrontendHeartbeatExtension: Send + Sync { + /// Returns the stable name used to make registration idempotent. + fn name(&self) -> &str; + + /// Generates request extensions for one heartbeat. + async fn request_extensions( + &self, + ) -> FrontendHeartbeatExtensionResult>> { + Ok(HashMap::new()) + } + + /// Returns a handler to insert into the heartbeat response handler chain. + /// + /// Errors and [`common_meta::heartbeat::handler::HandleControl::Done`] are isolated to this + /// extension so they cannot skip later extensions or mandatory OSS handlers. + fn response_handler(&self) -> Option { + None + } + + /// Notifies the extension that a heartbeat connection generation is ready. + async fn connected(&self, _generation: u64) -> FrontendHeartbeatExtensionResult<()> { + Ok(()) + } + + /// Stops and joins background work owned by the extension. + async fn shutdown(&self) -> FrontendHeartbeatExtensionResult<()> { + Ok(()) + } +} + +/// A shareable, ordered registry of frontend heartbeat extensions. +#[derive(Clone, Default)] +pub struct FrontendHeartbeatExtensions { + inner: Arc>, +} + +#[derive(Default)] +struct FrontendHeartbeatExtensionsInner { + names: HashSet, + extensions: Vec>, +} + +impl FrontendHeartbeatExtensions { + /// Registers an extension without replacing an existing extension of the same name. + /// + /// Returns `true` for a new registration and `false` for an idempotent duplicate. + pub fn register(&self, extension: Arc) -> bool { + let mut inner = self.inner.lock().unwrap(); + let name = extension.name().to_string(); + if !inner.names.insert(name) { + return false; + } + inner.extensions.push(extension); + true + } + + /// Returns the registered extensions in registration order. + pub fn extensions(&self) -> Vec> { + self.inner.lock().unwrap().extensions.clone() + } + + /// Returns response handlers in extension registration order. + pub fn response_handlers(&self) -> Vec { + self.extensions() + .into_iter() + .filter_map(|extension| extension.response_handler()) + .collect() + } + + /// Returns the number of registered extensions. + pub fn len(&self) -> usize { + self.inner.lock().unwrap().extensions.len() + } + + /// Returns whether no extension is registered. + pub fn is_empty(&self) -> bool { + self.len() == 0 + } +} + +/// Builds the frontend heartbeat response handler chain. +/// +/// Mailbox parsing always runs first, followed by extension handlers, suspension state handling, +/// and cache invalidation. +pub fn heartbeat_response_handler_executor( + extensions: &FrontendHeartbeatExtensions, + suspend_state: Arc, + cache_invalidator: CacheInvalidatorRef, +) -> HeartbeatResponseHandlerExecutorRef { + let mut handlers: Vec = vec![Arc::new(ParseMailboxMessageHandler)]; + handlers.extend(extensions.response_handlers().into_iter().map(|handler| { + Arc::new(IsolatedHeartbeatResponseHandler(handler)) as HeartbeatResponseHandlerRef + })); + handlers.extend([ + Arc::new(SuspendHandler::new(suspend_state)) as HeartbeatResponseHandlerRef, + Arc::new(InvalidateCacheHandler::new(cache_invalidator)), + ]); + Arc::new(HandlerGroupExecutor::new(handlers)) +} + +struct IsolatedHeartbeatResponseHandler(HeartbeatResponseHandlerRef); + +#[async_trait] +impl common_meta::heartbeat::handler::HeartbeatResponseHandler + for IsolatedHeartbeatResponseHandler +{ + fn is_acceptable(&self, _ctx: &HeartbeatResponseHandlerContext) -> bool { + true + } + + async fn handle( + &self, + ctx: &mut HeartbeatResponseHandlerContext, + ) -> common_meta::error::Result { + use common_meta::heartbeat::handler::HandleControl; + + if self.0.is_acceptable(ctx) + && let Err(error) = self.0.handle(ctx).await + { + error!(error; "Heartbeat extension response handler failed"); + } + Ok(HandleControl::Continue) + } +} + +#[async_trait] +trait HeartbeatConnector: Send + Sync { + async fn connect(&self) -> Result; +} + +#[async_trait] +trait HeartbeatRequestSender: Send + Sync { + async fn send(&self, request: HeartbeatRequest) -> FrontendHeartbeatExtensionResult<()>; +} + +#[async_trait] +trait HeartbeatResponseStream: Send { + async fn message(&mut self) -> FrontendHeartbeatExtensionResult>; +} + +struct HeartbeatConnection { + sender: Arc, + stream: Box, + config: HeartbeatConfig, +} + +struct MetaHeartbeatConnector { + client: Arc, +} + +#[async_trait] +impl HeartbeatConnector for MetaHeartbeatConnector { + async fn connect(&self) -> Result { + let (sender, stream, config) = self + .client + .heartbeat() + .await + .context(error::CreateMetaHeartbeatStreamSnafu)?; + Ok(HeartbeatConnection { + sender: Arc::new(MetaHeartbeatSender(sender)), + stream: Box::new(MetaHeartbeatStream(stream)), + config, + }) + } +} + +struct MetaHeartbeatSender(HeartbeatSender); + +#[async_trait] +impl HeartbeatRequestSender for MetaHeartbeatSender { + async fn send(&self, request: HeartbeatRequest) -> FrontendHeartbeatExtensionResult<()> { + self.0.send(request).await.map_err(BoxedError::new) + } +} + +struct MetaHeartbeatStream(HeartbeatStream); + +#[async_trait] +impl HeartbeatResponseStream for MetaHeartbeatStream { + async fn message(&mut self) -> FrontendHeartbeatExtensionResult> { + self.0.message().await.map_err(BoxedError::new) + } +} + #[derive(Clone)] -pub struct HeartbeatTask { +struct HeartbeatRunner { peer_addr: String, - meta_client: Arc, + connector: Arc, resp_handler_executor: HeartbeatResponseHandlerExecutorRef, start_time_ms: u64, resource_stat: ResourceStatRef, env_vars: EnvVars, + extensions: FrontendHeartbeatExtensions, + cancellation: CancellationToken, + generation: Arc, } -impl HeartbeatTask { - pub fn new( - peer_addr: String, - opts: &FrontendOptions, - meta_client: Arc, - resp_handler_executor: HeartbeatResponseHandlerExecutorRef, - resource_stat: ResourceStatRef, - ) -> Self { - HeartbeatTask { - peer_addr, - meta_client, - resp_handler_executor, - start_time_ms: common_time::util::current_time_millis() as u64, - resource_stat, - env_vars: EnvVars::from_config(&opts.heartbeat_env_vars), +impl HeartbeatRunner { + async fn connect(&self) -> Result> { + tokio::select! { + _ = self.cancellation.cancelled() => Ok(None), + connection = self.connector.connect() => connection.map(Some), } } - pub async fn start(&self) -> Result<()> { - let (req_sender, resp_stream, config) = self - .meta_client - .heartbeat() - .await - .context(error::CreateMetaHeartbeatStreamSnafu)?; - - info!("Heartbeat started with Metasrv config: {}", config); - - let (outgoing_tx, outgoing_rx) = mpsc::channel(16); - let mailbox = Arc::new(HeartbeatMailbox::new(outgoing_tx)); - - self.start_handle_resp_stream(resp_stream, mailbox, config.retry_interval); - - self.start_heartbeat_report(req_sender, outgoing_rx, config.interval); - - Ok(()) + fn next_generation(&self) -> u64 { + self.generation + .fetch_add(1, Ordering::AcqRel) + .wrapping_add(1) } - fn start_handle_resp_stream( - &self, - mut resp_stream: HeartbeatStream, - mailbox: MailboxRef, - retry_interval: Duration, - ) { - let capture_self = self.clone(); + async fn notify_connected(&self, generation: u64) -> bool { + for extension in self.extensions.extensions() { + let result = tokio::select! { + _ = self.cancellation.cancelled() => return false, + result = extension.connected(generation) => result, + }; + if let Err(error) = result { + error!(error; "Heartbeat extension '{}' failed its connected callback", extension.name()); + } + } + true + } + + async fn shutdown_extensions(&self) { + for extension in self.extensions.extensions() { + if let Err(error) = extension.shutdown().await { + error!(error; "Failed to shut down heartbeat extension '{}'", extension.name()); + } + } + } + + async fn run(self, mut connection: HeartbeatConnection) { + loop { + let retry_interval = connection.config.retry_interval; + if self.run_connection(connection).await == ConnectionEnd::Shutdown { + return; + } - let _handle = common_runtime::spawn_hb(async move { loop { - match resp_stream.message().await { - Ok(Some(resp)) => { - debug!("Receiving heartbeat response: {:?}", resp); - if let Some(message) = &resp.mailbox_message { - info!("Received mailbox message: {message:?}"); - } - let ctx = HeartbeatResponseHandlerContext::new(mailbox.clone(), resp); - if let Err(e) = capture_self.handle_response(ctx).await { - error!(e; "Error while handling heartbeat response"); - HEARTBEAT_RECV_COUNT - .with_label_values(&["processing_error"]) - .inc(); - } else { - HEARTBEAT_RECV_COUNT.with_label_values(&["success"]).inc(); - } - } - Ok(None) => { - warn!("Heartbeat response stream closed"); - capture_self.start_with_retry(retry_interval).await; - break; - } - Err(e) => { - HEARTBEAT_RECV_COUNT.with_label_values(&["error"]).inc(); - error!(e; "Occur error while reading heartbeat response"); - capture_self.start_with_retry(retry_interval).await; + if !self.wait_retry(retry_interval).await { + return; + } + info!("Try to re-establish the heartbeat connection to metasrv."); + match self.connect().await { + Ok(Some(next)) => { + let generation = self.next_generation(); + if !self.notify_connected(generation).await { + return; + } + connection = next; break; } + Ok(None) => return, + Err(error) => { + error!(error; "Failed to re-establish heartbeat connection to metasrv"); + } } } - }); + } + } + + async fn wait_retry(&self, retry_interval: Duration) -> bool { + tokio::select! { + _ = self.cancellation.cancelled() => false, + _ = tokio::time::sleep(retry_interval) => true, + } + } + + async fn run_connection(&self, connection: HeartbeatConnection) -> ConnectionEnd { + let (outgoing_tx, outgoing_rx) = mpsc::channel(16); + let mailbox = Arc::new(HeartbeatMailbox::new(outgoing_tx)); + let report = + self.report_heartbeats(connection.sender, outgoing_rx, connection.config.interval); + let responses = self.handle_responses(connection.stream, mailbox); + tokio::pin!(report); + tokio::pin!(responses); + + tokio::select! { + _ = self.cancellation.cancelled() => ConnectionEnd::Shutdown, + end = &mut report => end, + end = &mut responses => end, + } + } + + async fn handle_responses( + &self, + mut stream: Box, + mailbox: MailboxRef, + ) -> ConnectionEnd { + loop { + let response = tokio::select! { + _ = self.cancellation.cancelled() => return ConnectionEnd::Shutdown, + response = stream.message() => response, + }; + match response { + Ok(Some(response)) => { + debug!("Receiving heartbeat response: {:?}", response); + if let Some(message) = &response.mailbox_message { + info!("Received mailbox message: {message:?}"); + } + let context = HeartbeatResponseHandlerContext::new(mailbox.clone(), response); + let result = tokio::select! { + _ = self.cancellation.cancelled() => return ConnectionEnd::Shutdown, + result = self.handle_response(context) => result, + }; + if let Err(error) = result { + error!(error; "Error while handling heartbeat response"); + HEARTBEAT_RECV_COUNT + .with_label_values(&["processing_error"]) + .inc(); + } else { + HEARTBEAT_RECV_COUNT.with_label_values(&["success"]).inc(); + } + } + Ok(None) => { + warn!("Heartbeat response stream closed"); + return ConnectionEnd::Reconnect; + } + Err(error) => { + HEARTBEAT_RECV_COUNT.with_label_values(&["error"]).inc(); + error!(error; "Occur error while reading heartbeat response"); + return ConnectionEnd::Reconnect; + } + } + } + } + + async fn report_heartbeats( + &self, + sender: Arc, + mut outgoing_rx: Receiver, + report_interval: Duration, + ) -> ConnectionEnd { + let total_cpu_millicores = self.resource_stat.get_total_cpu_millicores(); + let total_memory_bytes = self.resource_stat.get_total_memory_bytes(); + let mut extensions = HashMap::new(); + self.env_vars.into_extensions(&mut extensions); + let heartbeat_request = HeartbeatRequest { + peer: Some(Peer { + // Metasrv calculates the frontend id by hashing this reachable address. + id: 0, + addr: self.peer_addr.clone(), + }), + info: Self::build_node_info( + self.start_time_ms, + total_cpu_millicores, + total_memory_bytes, + ), + node_workloads: Some(NodeWorkloads::Frontend(FrontendWorkloads { types: vec![] })), + extensions, + ..Default::default() + }; + let sleep = tokio::time::sleep(Duration::ZERO); + tokio::pin!(sleep); + + loop { + let request = tokio::select! { + _ = self.cancellation.cancelled() => return ConnectionEnd::Shutdown, + message = outgoing_rx.recv() => { + if let Some(message) = message { + Self::new_heartbeat_request(&heartbeat_request, Some(message), 0, 0) + } else { + warn!("Sender has been dropped, exiting the heartbeat loop"); + return ConnectionEnd::Reconnect; + } + } + _ = &mut sleep => { + sleep.as_mut().reset(Instant::now() + report_interval); + Self::new_heartbeat_request( + &heartbeat_request, + None, + self.resource_stat.get_cpu_usage_millicores(), + self.resource_stat.get_memory_usage_bytes(), + ) + } + }; + + if let Some(mut request) = request { + if !self.add_request_extensions(&mut request).await { + return ConnectionEnd::Shutdown; + } + debug!( + "Sending a heartbeat request to metasrv, content: {:?}", + request + ); + let result = tokio::select! { + _ = self.cancellation.cancelled() => return ConnectionEnd::Shutdown, + result = sender.send(request) => result, + }; + if let Err(error) = result { + error!(error; "Failed to send heartbeat to metasrv"); + return ConnectionEnd::Reconnect; + } + HEARTBEAT_SENT_COUNT.inc(); + } + } + } + + async fn add_request_extensions(&self, request: &mut HeartbeatRequest) -> bool { + for extension in self.extensions.extensions() { + let generated = tokio::select! { + _ = self.cancellation.cancelled() => return false, + generated = extension.request_extensions() => generated, + }; + let generated = match generated { + Ok(generated) => generated, + Err(error) => { + error!(error; "Heartbeat extension '{}' failed to generate request extensions", extension.name()); + continue; + } + }; + + if let Some(key) = generated + .keys() + .find(|key| request.extensions.contains_key(*key)) + { + warn!( + "Heartbeat extension '{}' produced conflicting key '{}'; discarding its output", + extension.name(), + key + ); + continue; + } + request.extensions.extend(generated); + } + true } fn new_heartbeat_request( @@ -138,8 +500,8 @@ impl HeartbeatTask { ) -> Option { let mailbox_message = match message.map(outgoing_message_to_mailbox_message) { Some(Ok(message)) => Some(message), - Some(Err(e)) => { - error!(e; "Failed to encode mailbox messages"); + Some(Err(error)) => { + error!(error; "Failed to encode mailbox messages"); return None; } None => None, @@ -149,12 +511,10 @@ impl HeartbeatTask { mailbox_message, ..heartbeat_request.clone() }; - if let Some(info) = heartbeat_request.info.as_mut() { info.memory_usage_bytes = memory_usage; info.cpu_usage_millicores = cpu_usage; } - Some(heartbeat_request) } @@ -165,7 +525,6 @@ impl HeartbeatTask { total_memory_bytes: i64, ) -> Option { let build_info = common_version::build_info(); - Some(NodeInfo { version: build_info.version.to_string(), git_commit: build_info.commit_short.to_string(), @@ -184,90 +543,156 @@ impl HeartbeatTask { }) } - fn start_heartbeat_report( - &self, - req_sender: HeartbeatSender, - mut outgoing_rx: Receiver, - report_interval: Duration, - ) { - let start_time_ms = self.start_time_ms; - let self_peer = Some(Peer { - // The node id will be actually calculated from its address (by hashing the address - // string) in the metasrv. So it can be set to 0 here, as a placeholder. - id: 0, - addr: self.peer_addr.clone(), - }); - let total_cpu_millicores = self.resource_stat.get_total_cpu_millicores(); - let total_memory_bytes = self.resource_stat.get_total_memory_bytes(); - let resource_stat = self.resource_stat.clone(); - let env_vars = self.env_vars.clone(); - common_runtime::spawn_hb(async move { - let sleep = tokio::time::sleep(Duration::from_millis(0)); - tokio::pin!(sleep); - - let mut extensions = std::collections::HashMap::new(); - env_vars.into_extensions(&mut extensions); - - let heartbeat_request = HeartbeatRequest { - peer: self_peer, - info: Self::build_node_info( - start_time_ms, - total_cpu_millicores, - total_memory_bytes, - ), - node_workloads: Some(NodeWorkloads::Frontend(FrontendWorkloads { types: vec![] })), - extensions, - ..Default::default() - }; - - loop { - let req = tokio::select! { - message = outgoing_rx.recv() => { - if let Some(message) = message { - Self::new_heartbeat_request(&heartbeat_request, Some(message), 0, 0) - } else { - warn!("Sender has been dropped, exiting the heartbeat loop"); - // Receives None that means Sender was dropped, we need to break the current loop - break - } - } - _ = &mut sleep => { - sleep.as_mut().reset(Instant::now() + report_interval); - Self::new_heartbeat_request(&heartbeat_request, None, resource_stat.get_cpu_usage_millicores(), resource_stat.get_memory_usage_bytes()) - } - }; - - if let Some(req) = req { - if let Err(e) = req_sender.send(req.clone()).await { - error!(e; "Failed to send heartbeat to metasrv"); - break; - } else { - HEARTBEAT_SENT_COUNT.inc(); - debug!("Send a heartbeat request to metasrv, content: {:?}", req); - } - } - } - }); - } - - async fn handle_response(&self, ctx: HeartbeatResponseHandlerContext) -> Result<()> { + async fn handle_response(&self, context: HeartbeatResponseHandlerContext) -> Result<()> { self.resp_handler_executor - .handle(ctx) + .handle(context) .await .context(error::HandleHeartbeatResponseSnafu) } +} - async fn start_with_retry(&self, retry_interval: Duration) { - loop { - tokio::time::sleep(retry_interval).await; +#[derive(Debug, PartialEq, Eq)] +enum ConnectionEnd { + Reconnect, + Shutdown, +} - info!("Try to re-establish the heartbeat connection to metasrv."); +/// The frontend task that sends [`HeartbeatRequest`] values to metasrv in the background. +#[derive(Clone)] +pub struct HeartbeatTask { + runner: HeartbeatRunner, + start_lock: Arc>, + shutdown_lock: Arc>, + supervisor: Arc>>>, +} - if self.start().await.is_ok() { - break; - } +impl HeartbeatTask { + pub fn new( + peer_addr: String, + opts: &FrontendOptions, + meta_client: Arc, + resp_handler_executor: HeartbeatResponseHandlerExecutorRef, + resource_stat: ResourceStatRef, + ) -> Self { + Self::new_with_connector( + peer_addr, + opts, + Arc::new(MetaHeartbeatConnector { + client: meta_client, + }), + resp_handler_executor, + resource_stat, + ) + } + + fn new_with_connector( + peer_addr: String, + opts: &FrontendOptions, + connector: Arc, + resp_handler_executor: HeartbeatResponseHandlerExecutorRef, + resource_stat: ResourceStatRef, + ) -> Self { + Self { + runner: HeartbeatRunner { + peer_addr, + connector, + resp_handler_executor, + start_time_ms: common_time::util::current_time_millis() as u64, + resource_stat, + env_vars: EnvVars::from_config(&opts.heartbeat_env_vars), + extensions: FrontendHeartbeatExtensions::default(), + cancellation: CancellationToken::new(), + generation: Arc::new(AtomicU64::new(0)), + }, + start_lock: Arc::new(Mutex::new(())), + shutdown_lock: Arc::new(Mutex::new(())), + supervisor: Arc::new(Mutex::new(None)), } } + + /// Installs the extensions registered before heartbeat startup. + pub fn with_extensions(mut self, extensions: FrontendHeartbeatExtensions) -> Self { + self.runner.extensions = extensions; + self + } + + /// Establishes the initial heartbeat connection and starts its background supervisor. + pub async fn start(&self) -> Result<()> { + let _start_guard = self.start_lock.lock().await; + if self.runner.cancellation.is_cancelled() { + return Ok(()); + } + + let finished = { + let mut supervisor = self.supervisor.lock().await; + match supervisor.as_ref() { + Some(handle) if !handle.is_finished() => return Ok(()), + Some(_) => supervisor.take(), + None => None, + } + }; + if let Some(handle) = finished + && let Err(error) = handle.await + && !error.is_cancelled() + { + error!(error; "Heartbeat supervisor join failed"); + } + + let Some(connection) = self.runner.connect().await? else { + return Ok(()); + }; + info!( + "Heartbeat started with Metasrv config: {}", + connection.config + ); + + let generation = self.runner.next_generation(); + if !self.runner.notify_connected(generation).await { + return Ok(()); + } + + let runner = self.runner.clone(); + let handle = common_runtime::spawn_hb(async move { + runner.run(connection).await; + }); + *self.supervisor.lock().await = Some(handle); + Ok(()) + } + + /// Cancels and joins heartbeat I/O and all registered extension lifecycles. + pub async fn shutdown(&self) { + let _shutdown_guard = self.shutdown_lock.lock().await; + if self.runner.cancellation.is_cancelled() { + return; + } + self.runner.cancellation.cancel(); + + // Wait for a concurrently running handshake or connected callback to observe cancellation. + let _start_guard = self.start_lock.lock().await; + let handle = self.supervisor.lock().await.take(); + if let Some(handle) = handle + && let Err(error) = handle.await + && !error.is_cancelled() + { + error!(error; "Heartbeat supervisor join failed"); + } + self.runner.shutdown_extensions().await; + } + + #[cfg(test)] + fn generation(&self) -> u64 { + self.runner.generation.load(Ordering::Acquire) + } + + #[cfg(test)] + async fn has_supervisor(&self) -> bool { + self.supervisor.lock().await.is_some() + } + + #[cfg(test)] + pub(crate) fn is_shutdown(&self) -> bool { + self.runner.cancellation.is_cancelled() + } } pub(crate) fn frontend_peer_addr(opts: &FrontendOptions) -> String { diff --git a/src/frontend/src/heartbeat/tests.rs b/src/frontend/src/heartbeat/tests.rs index d6c314afba..d3c651de32 100644 --- a/src/frontend/src/heartbeat/tests.rs +++ b/src/frontend/src/heartbeat/tests.rs @@ -12,14 +12,26 @@ // See the License for the specific language governing permissions and // limitations under the License. -use std::collections::HashMap; +use std::collections::{HashMap, VecDeque}; +use std::future::pending; +use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; +use std::time::Duration; -use api::v1::meta::{HeartbeatResponse, Role}; +use api::v1::meta::heartbeat_request::NodeWorkloads; +use api::v1::meta::mailbox_message::Payload; +use api::v1::meta::{ + HeartbeatConfig as ApiHeartbeatConfig, HeartbeatRequest, HeartbeatResponse, MailboxMessage, + RegionLease, ResponseHeader, Role, +}; +use async_trait::async_trait; +use common_error::ext::{BoxedError, PlainError}; +use common_error::status_code::StatusCode; use common_meta::cache_invalidator::KvCacheInvalidator; use common_meta::heartbeat::handler::invalidate_table_cache::InvalidateCacheHandler; use common_meta::heartbeat::handler::{ - HandlerGroupExecutor, HeartbeatResponseHandlerContext, HeartbeatResponseHandlerExecutor, + HandleControl, HandlerGroupExecutor, HeartbeatResponseHandler, HeartbeatResponseHandlerContext, + HeartbeatResponseHandlerExecutor, HeartbeatResponseHandlerRef, }; use common_meta::heartbeat::mailbox::{HeartbeatMailbox, MessageMeta}; use common_meta::instruction::{CacheIdent, Instruction}; @@ -29,9 +41,10 @@ use common_meta::key::table_info::TableInfoKey; use common_stat::ResourceStatImpl; use common_telemetry::tracing_context::TracingContext; use meta_client::client::MetaClient; -use tokio::sync::mpsc; +use prost::Message; +use tokio::sync::{Notify, mpsc}; -use super::HeartbeatTask; +use super::*; use crate::frontend::FrontendOptions; #[derive(Default)] @@ -39,6 +52,18 @@ pub struct MockKvCacheInvalidator { inner: Mutex, i32>>, } +#[derive(Clone, PartialEq, Message)] +struct LegacyHeartbeatResponse { + #[prost(message, optional, tag = "1")] + header: Option, + #[prost(message, optional, tag = "2")] + mailbox_message: Option, + #[prost(message, optional, tag = "3")] + region_lease: Option, + #[prost(message, optional, tag = "4")] + heartbeat_config: Option, +} + #[async_trait::async_trait] impl KvCacheInvalidator for MockKvCacheInvalidator { async fn invalidate_key(&self, key: &[u8]) { @@ -172,5 +197,995 @@ fn test_heartbeat_task_uses_resolved_peer_addr() { stat, ); - assert_eq!(task.peer_addr, "10.0.0.1:4001"); + assert_eq!(task.runner.peer_addr, "10.0.0.1:4001"); +} + +enum ConnectPlan { + Ready(MockConnectionPlan), + Fail, + Pending, +} + +struct MockConnector { + plans: Mutex>, + calls: AtomicUsize, + called: Notify, +} + +impl MockConnector { + fn new(plans: impl IntoIterator) -> Arc { + Arc::new(Self { + plans: Mutex::new(plans.into_iter().collect()), + calls: AtomicUsize::new(0), + called: Notify::new(), + }) + } + + async fn wait_for_calls(&self, expected: usize) { + tokio::time::timeout(Duration::from_secs(2), async { + loop { + if self.calls.load(Ordering::Acquire) >= expected { + return; + } + let notified = self.called.notified(); + if self.calls.load(Ordering::Acquire) >= expected { + return; + } + notified.await; + } + }) + .await + .unwrap(); + } +} + +#[async_trait] +impl HeartbeatConnector for MockConnector { + async fn connect(&self) -> Result { + self.calls.fetch_add(1, Ordering::AcqRel); + self.called.notify_waiters(); + let plan = self.plans.lock().unwrap().pop_front(); + match plan { + Some(ConnectPlan::Ready(plan)) => Ok(plan.into_connection()), + Some(ConnectPlan::Fail) => Err(crate::error::Error::NotSupported { + feat: "mock heartbeat connection".to_string(), + }), + Some(ConnectPlan::Pending) | None => pending().await, + } + } +} + +struct MockConnectionPlan { + sender: MockSender, + response_rx: mpsc::UnboundedReceiver, + config: HeartbeatConfig, +} + +impl MockConnectionPlan { + fn into_connection(self) -> HeartbeatConnection { + HeartbeatConnection { + sender: Arc::new(self.sender), + stream: Box::new(MockStream { + receiver: self.response_rx, + }), + config: self.config, + } + } +} + +#[derive(Clone)] +struct MockConnectionHandle { + requests: Arc>>, + request_added: Arc, + response_tx: mpsc::UnboundedSender, + fail_send: Arc, +} + +impl MockConnectionHandle { + async fn wait_for_requests(&self, expected: usize) -> Vec { + tokio::time::timeout(Duration::from_secs(2), async { + loop { + let requests = self.requests.lock().unwrap().clone(); + if requests.len() >= expected { + return requests; + } + let notified = self.request_added.notified(); + if self.requests.lock().unwrap().len() >= expected { + continue; + } + notified.await; + } + }) + .await + .unwrap() + } + + fn close_responses(&self) { + self.response_tx.send(MockStreamEvent::Close).unwrap(); + } + + fn fail_responses(&self) { + self.response_tx.send(MockStreamEvent::Error).unwrap(); + } + + fn send_response(&self, response: HeartbeatResponse) { + self.response_tx + .send(MockStreamEvent::Response(response)) + .unwrap(); + } +} + +#[derive(Clone)] +struct MockSender { + requests: Arc>>, + request_added: Arc, + fail_send: Arc, +} + +#[async_trait] +impl HeartbeatRequestSender for MockSender { + async fn send(&self, request: HeartbeatRequest) -> FrontendHeartbeatExtensionResult<()> { + if self.fail_send.load(Ordering::Acquire) { + return Err(test_error("mock heartbeat send failure")); + } + self.requests.lock().unwrap().push(request); + self.request_added.notify_waiters(); + Ok(()) + } +} + +enum MockStreamEvent { + Response(HeartbeatResponse), + Close, + Error, +} + +struct MockStream { + receiver: mpsc::UnboundedReceiver, +} + +#[async_trait] +impl HeartbeatResponseStream for MockStream { + async fn message(&mut self) -> FrontendHeartbeatExtensionResult> { + match self.receiver.recv().await { + Some(MockStreamEvent::Response(response)) => Ok(Some(response)), + Some(MockStreamEvent::Close) | None => Ok(None), + Some(MockStreamEvent::Error) => Err(test_error("mock heartbeat receive failure")), + } + } +} + +fn mock_connection( + interval: Duration, + retry_interval: Duration, +) -> (MockConnectionPlan, MockConnectionHandle) { + let requests = Arc::new(Mutex::new(Vec::new())); + let request_added = Arc::new(Notify::new()); + let fail_send = Arc::new(AtomicBool::new(false)); + let (response_tx, response_rx) = mpsc::unbounded_channel(); + let sender = MockSender { + requests: requests.clone(), + request_added: request_added.clone(), + fail_send: fail_send.clone(), + }; + let handle = MockConnectionHandle { + requests, + request_added, + response_tx, + fail_send, + }; + ( + MockConnectionPlan { + sender, + response_rx, + config: HeartbeatConfig { + interval, + retry_interval, + gc_enabled: false, + }, + }, + handle, + ) +} + +fn test_error(message: &str) -> BoxedError { + BoxedError::new(PlainError::new(message.to_string(), StatusCode::Unexpected)) +} + +struct TestExtension { + name: String, + static_extensions: HashMap>, + dynamic_key: Option, + fail_request: AtomicBool, + request_calls: AtomicUsize, + connected_generations: Mutex>, + shutdown_calls: AtomicUsize, + response_handler: Option, +} + +impl TestExtension { + fn new(name: &str) -> Arc { + Arc::new(Self { + name: name.to_string(), + static_extensions: HashMap::new(), + dynamic_key: None, + fail_request: AtomicBool::new(false), + request_calls: AtomicUsize::new(0), + connected_generations: Mutex::new(Vec::new()), + shutdown_calls: AtomicUsize::new(0), + response_handler: None, + }) + } + + fn with_static(name: &str, static_extensions: HashMap>) -> Arc { + Arc::new(Self { + static_extensions, + ..Self::new_fields(name) + }) + } + + fn dynamic(name: &str, key: &str) -> Arc { + Arc::new(Self { + dynamic_key: Some(key.to_string()), + ..Self::new_fields(name) + }) + } + + fn failing(name: &str) -> Arc { + Arc::new(Self { + fail_request: AtomicBool::new(true), + ..Self::new_fields(name) + }) + } + + fn with_handler(name: &str, response_handler: HeartbeatResponseHandlerRef) -> Arc { + Arc::new(Self { + response_handler: Some(response_handler), + ..Self::new_fields(name) + }) + } + + fn new_fields(name: &str) -> Self { + Self { + name: name.to_string(), + static_extensions: HashMap::new(), + dynamic_key: None, + fail_request: AtomicBool::new(false), + request_calls: AtomicUsize::new(0), + connected_generations: Mutex::new(Vec::new()), + shutdown_calls: AtomicUsize::new(0), + response_handler: None, + } + } +} + +#[async_trait] +impl FrontendHeartbeatExtension for TestExtension { + fn name(&self) -> &str { + &self.name + } + + async fn request_extensions( + &self, + ) -> FrontendHeartbeatExtensionResult>> { + let call = self.request_calls.fetch_add(1, Ordering::AcqRel) + 1; + if self.fail_request.load(Ordering::Acquire) { + return Err(test_error("mock extension failure")); + } + let mut extensions = self.static_extensions.clone(); + if let Some(key) = &self.dynamic_key { + extensions.insert(key.clone(), call.to_string().into_bytes()); + } + Ok(extensions) + } + + fn response_handler(&self) -> Option { + self.response_handler.clone() + } + + async fn connected(&self, generation: u64) -> FrontendHeartbeatExtensionResult<()> { + self.connected_generations.lock().unwrap().push(generation); + Ok(()) + } + + async fn shutdown(&self) -> FrontendHeartbeatExtensionResult<()> { + self.shutdown_calls.fetch_add(1, Ordering::AcqRel); + Ok(()) + } +} + +fn test_task( + connector: Arc, + extensions: FrontendHeartbeatExtensions, + options: FrontendOptions, +) -> HeartbeatTask { + let executor = heartbeat_response_handler_executor( + &extensions, + Arc::new(AtomicBool::new(false)), + Arc::new(MockKvCacheInvalidator::default()), + ); + HeartbeatTask::new_with_connector( + "127.0.0.1:4001".to_string(), + &options, + connector, + executor, + Arc::new(ResourceStatImpl::default()), + ) + .with_extensions(extensions) +} + +#[tokio::test] +async fn test_heartbeat_without_extensions_preserves_base_behavior() { + let (plan, handle) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let connector = MockConnector::new([ConnectPlan::Ready(plan)]); + let task = test_task( + connector, + FrontendHeartbeatExtensions::default(), + FrontendOptions::default(), + ); + + task.start().await.unwrap(); + let requests = handle.wait_for_requests(1).await; + let request = &requests[0]; + assert_eq!(request.peer.as_ref().unwrap().addr, "127.0.0.1:4001"); + assert!(request.extensions.is_empty()); + assert!(matches!( + request.node_workloads, + Some(NodeWorkloads::Frontend(_)) + )); + + task.shutdown().await; + assert!(!task.has_supervisor().await); +} + +#[tokio::test] +async fn test_initial_connection_failure_does_not_start_extension_lifecycle() { + let extension = TestExtension::new("lifecycle"); + let extensions = FrontendHeartbeatExtensions::default(); + assert!(extensions.register(extension.clone())); + let connector = MockConnector::new([ConnectPlan::Fail]); + let task = test_task(connector, extensions, FrontendOptions::default()); + + assert!(task.start().await.is_err()); + assert!(extension.connected_generations.lock().unwrap().is_empty()); + assert!(!task.has_supervisor().await); + task.shutdown().await; +} + +#[tokio::test] +async fn test_request_extensions_are_regenerated_for_every_heartbeat() { + let extension = TestExtension::dynamic("dynamic", "dynamic-key"); + let extensions = FrontendHeartbeatExtensions::default(); + extensions.register(extension.clone()); + let (plan, handle) = mock_connection(Duration::from_millis(10), Duration::from_millis(10)); + let task = test_task( + MockConnector::new([ConnectPlan::Ready(plan)]), + extensions, + FrontendOptions::default(), + ); + + task.start().await.unwrap(); + let requests = handle.wait_for_requests(2).await; + assert_eq!(requests[0].extensions["dynamic-key"], b"1"); + assert_eq!(requests[1].extensions["dynamic-key"], b"2"); + assert_eq!(extension.request_calls.load(Ordering::Acquire), 2); + task.shutdown().await; +} + +#[tokio::test] +async fn test_failed_provider_preserves_base_and_other_extensions() { + let failing = TestExtension::failing("failing"); + let successful = TestExtension::with_static( + "successful", + HashMap::from([("other".to_string(), b"value".to_vec())]), + ); + let extensions = FrontendHeartbeatExtensions::default(); + extensions.register(failing); + extensions.register(successful); + let task = test_task( + MockConnector::new([ConnectPlan::Pending]), + extensions, + FrontendOptions::default(), + ); + let mut request = HeartbeatRequest { + extensions: HashMap::from([("base".to_string(), b"keep".to_vec())]), + ..Default::default() + }; + + assert!(task.runner.add_request_extensions(&mut request).await); + assert_eq!(request.extensions["base"], b"keep"); + assert_eq!(request.extensions["other"], b"value"); +} + +#[tokio::test] +async fn test_conflicting_provider_output_is_discarded_atomically() { + let conflicting = TestExtension::with_static( + "conflicting", + HashMap::from([ + ("base".to_string(), b"replace".to_vec()), + ("partial".to_string(), b"discard".to_vec()), + ]), + ); + let successful = TestExtension::with_static( + "successful", + HashMap::from([("other".to_string(), b"value".to_vec())]), + ); + let extensions = FrontendHeartbeatExtensions::default(); + assert!(extensions.register(conflicting.clone())); + assert!(!extensions.register(conflicting)); + assert!(extensions.register(successful)); + let task = test_task( + MockConnector::new([ConnectPlan::Pending]), + extensions, + FrontendOptions::default(), + ); + let mut request = HeartbeatRequest { + extensions: HashMap::from([("base".to_string(), b"keep".to_vec())]), + ..Default::default() + }; + + assert!(task.runner.add_request_extensions(&mut request).await); + assert_eq!(request.extensions["base"], b"keep"); + assert!(!request.extensions.contains_key("partial")); + assert_eq!(request.extensions["other"], b"value"); +} + +type HandlerObservations = Arc>>; + +struct OrderingHandler { + observations: HandlerObservations, + suspend_state: Arc, + cache: Arc, + table_key: Vec, +} + +enum ShortCircuitResult { + Done, + Error, +} + +struct ShortCircuitHandler { + result: ShortCircuitResult, + calls: Arc, +} + +struct PendingHandler { + started: Arc, +} + +#[async_trait] +impl HeartbeatResponseHandler for PendingHandler { + fn is_acceptable(&self, _ctx: &HeartbeatResponseHandlerContext) -> bool { + true + } + + async fn handle( + &self, + _ctx: &mut HeartbeatResponseHandlerContext, + ) -> common_meta::error::Result { + self.started.notify_waiters(); + pending().await + } +} + +#[async_trait] +impl HeartbeatResponseHandler for ShortCircuitHandler { + fn is_acceptable(&self, _ctx: &HeartbeatResponseHandlerContext) -> bool { + true + } + + async fn handle( + &self, + _ctx: &mut HeartbeatResponseHandlerContext, + ) -> common_meta::error::Result { + self.calls.fetch_add(1, Ordering::AcqRel); + match self.result { + ShortCircuitResult::Done => Ok(HandleControl::Done), + ShortCircuitResult::Error => common_meta::error::UnsupportedSnafu { + operation: "mock extension response handler", + } + .fail(), + } + } +} + +#[async_trait] +impl HeartbeatResponseHandler for OrderingHandler { + fn is_acceptable(&self, _ctx: &HeartbeatResponseHandlerContext) -> bool { + true + } + + async fn handle( + &self, + ctx: &mut HeartbeatResponseHandlerContext, + ) -> common_meta::error::Result { + self.observations.lock().unwrap().push(( + ctx.incoming_message.is_some(), + self.suspend_state.load(Ordering::Acquire), + self.cache + .inner + .lock() + .unwrap() + .contains_key(&self.table_key), + ctx.response.extensions.contains_key("response-extension"), + )); + Ok(HandleControl::Continue) + } +} + +#[tokio::test] +async fn test_response_extension_handler_order() { + let table_id = 42; + let table_key = TableInfoKey::new(table_id).to_bytes(); + let cache = Arc::new(MockKvCacheInvalidator { + inner: Mutex::new(HashMap::from([(table_key.clone(), 1)])), + }); + let suspend_state = Arc::new(AtomicBool::new(true)); + let observations = Arc::new(Mutex::new(Vec::new())); + let handler = Arc::new(OrderingHandler { + observations: observations.clone(), + suspend_state: suspend_state.clone(), + cache: cache.clone(), + table_key: table_key.clone(), + }); + let extensions = FrontendHeartbeatExtensions::default(); + extensions.register(TestExtension::with_handler("response", handler)); + let executor = + heartbeat_response_handler_executor(&extensions, suspend_state.clone(), cache.clone()); + let (mailbox_tx, _) = mpsc::channel(1); + let mailbox = Arc::new(HeartbeatMailbox::new(mailbox_tx)); + + executor + .handle(HeartbeatResponseHandlerContext::new( + mailbox.clone(), + HeartbeatResponse { + extensions: HashMap::from([("response-extension".to_string(), b"value".to_vec())]), + ..Default::default() + }, + )) + .await + .unwrap(); + assert!(!suspend_state.load(Ordering::Acquire)); + + let response = HeartbeatResponse { + mailbox_message: Some(MailboxMessage { + payload: Some(Payload::Json( + serde_json::to_string(&Instruction::InvalidateCaches(vec![CacheIdent::TableId( + table_id, + )])) + .unwrap(), + )), + ..Default::default() + }), + extensions: HashMap::from([("response-extension".to_string(), b"value".to_vec())]), + ..Default::default() + }; + executor + .handle(HeartbeatResponseHandlerContext::new(mailbox, response)) + .await + .unwrap(); + + assert_eq!( + observations.lock().unwrap().as_slice(), + &[(false, true, true, true), (true, false, true, true)] + ); + assert!(!suspend_state.load(Ordering::Acquire)); + assert!(!cache.inner.lock().unwrap().contains_key(&table_key)); +} + +#[tokio::test] +async fn test_heartbeat_response_wire_compatibility_preserves_handlers() { + let table_id = 42; + let table_key = TableInfoKey::new(table_id).to_bytes(); + let cache = Arc::new(MockKvCacheInvalidator { + inner: Mutex::new(HashMap::from([(table_key.clone(), 1)])), + }); + let suspend_state = Arc::new(AtomicBool::new(false)); + let executor = heartbeat_response_handler_executor( + &FrontendHeartbeatExtensions::default(), + suspend_state.clone(), + cache.clone(), + ); + let (mailbox_tx, _) = mpsc::channel(1); + let mailbox = Arc::new(HeartbeatMailbox::new(mailbox_tx)); + let heartbeat_config = ApiHeartbeatConfig { + heartbeat_interval_ms: 3_000, + retry_interval_ms: 500, + gc_enabled: true, + }; + + let new_response = HeartbeatResponse { + header: Some(ResponseHeader::success()), + mailbox_message: Some(MailboxMessage { + payload: Some(Payload::Json( + serde_json::to_string(&Instruction::Suspend).unwrap(), + )), + ..Default::default() + }), + region_lease: Some(RegionLease::default()), + heartbeat_config: Some(heartbeat_config), + extensions: HashMap::from([("response-extension".to_string(), b"value".to_vec())]), + }; + let legacy_decoded = + LegacyHeartbeatResponse::decode(new_response.encode_to_vec().as_slice()).unwrap(); + assert_eq!(new_response.header, legacy_decoded.header); + assert_eq!(new_response.mailbox_message, legacy_decoded.mailbox_message); + assert_eq!(new_response.region_lease, legacy_decoded.region_lease); + assert_eq!( + new_response.heartbeat_config, + legacy_decoded.heartbeat_config + ); + + executor + .handle(HeartbeatResponseHandlerContext::new( + mailbox.clone(), + HeartbeatResponse { + header: legacy_decoded.header, + mailbox_message: legacy_decoded.mailbox_message, + region_lease: legacy_decoded.region_lease, + heartbeat_config: legacy_decoded.heartbeat_config, + extensions: HashMap::new(), + }, + )) + .await + .unwrap(); + assert!(suspend_state.load(Ordering::Acquire)); + + let legacy_response = LegacyHeartbeatResponse { + header: Some(ResponseHeader::success()), + mailbox_message: Some(MailboxMessage { + payload: Some(Payload::Json( + serde_json::to_string(&Instruction::InvalidateCaches(vec![CacheIdent::TableId( + table_id, + )])) + .unwrap(), + )), + ..Default::default() + }), + region_lease: Some(RegionLease::default()), + heartbeat_config: Some(heartbeat_config), + }; + let current_decoded = + HeartbeatResponse::decode(legacy_response.encode_to_vec().as_slice()).unwrap(); + assert_eq!(legacy_response.header, current_decoded.header); + assert_eq!( + legacy_response.mailbox_message, + current_decoded.mailbox_message + ); + assert_eq!(legacy_response.region_lease, current_decoded.region_lease); + assert_eq!( + legacy_response.heartbeat_config, + current_decoded.heartbeat_config + ); + assert!(current_decoded.extensions.is_empty()); + + executor + .handle(HeartbeatResponseHandlerContext::new( + mailbox, + current_decoded, + )) + .await + .unwrap(); + assert!(!cache.inner.lock().unwrap().contains_key(&table_key)); +} + +async fn assert_extension_short_circuit_is_isolated(result: ShortCircuitResult) { + let table_id = 42; + let table_key = TableInfoKey::new(table_id).to_bytes(); + let cache = Arc::new(MockKvCacheInvalidator { + inner: Mutex::new(HashMap::from([(table_key.clone(), 1)])), + }); + let suspend_state = Arc::new(AtomicBool::new(true)); + let short_circuit_calls = Arc::new(AtomicUsize::new(0)); + let observations = Arc::new(Mutex::new(Vec::new())); + let extensions = FrontendHeartbeatExtensions::default(); + extensions.register(TestExtension::with_handler( + "short-circuit", + Arc::new(ShortCircuitHandler { + result, + calls: short_circuit_calls.clone(), + }), + )); + extensions.register(TestExtension::with_handler( + "remaining", + Arc::new(OrderingHandler { + observations: observations.clone(), + suspend_state: suspend_state.clone(), + cache: cache.clone(), + table_key: table_key.clone(), + }), + )); + let executor = + heartbeat_response_handler_executor(&extensions, suspend_state.clone(), cache.clone()); + let (mailbox_tx, _) = mpsc::channel(1); + let mailbox = Arc::new(HeartbeatMailbox::new(mailbox_tx)); + + executor + .handle(HeartbeatResponseHandlerContext::new( + mailbox.clone(), + HeartbeatResponse::default(), + )) + .await + .unwrap(); + + let invalidate_response = HeartbeatResponse { + mailbox_message: Some(MailboxMessage { + payload: Some(Payload::Json( + serde_json::to_string(&Instruction::InvalidateCaches(vec![CacheIdent::TableId( + table_id, + )])) + .unwrap(), + )), + ..Default::default() + }), + ..Default::default() + }; + + executor + .handle(HeartbeatResponseHandlerContext::new( + mailbox, + invalidate_response, + )) + .await + .unwrap(); + + assert_eq!(short_circuit_calls.load(Ordering::Acquire), 2); + assert_eq!(observations.lock().unwrap().len(), 2); + assert!(!suspend_state.load(Ordering::Acquire)); + assert!(!cache.inner.lock().unwrap().contains_key(&table_key)); +} + +#[tokio::test] +async fn test_response_extension_done_does_not_skip_remaining_handlers() { + assert_extension_short_circuit_is_isolated(ShortCircuitResult::Done).await; +} + +#[tokio::test] +async fn test_response_extension_error_does_not_skip_remaining_handlers() { + assert_extension_short_circuit_is_isolated(ShortCircuitResult::Error).await; +} + +#[tokio::test] +async fn test_closed_response_stream_reconnects_once() { + let extension = TestExtension::new("lifecycle"); + let extensions = FrontendHeartbeatExtensions::default(); + extensions.register(extension.clone()); + let (first_plan, first) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let (second_plan, second) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let connector = MockConnector::new([ + ConnectPlan::Ready(first_plan), + ConnectPlan::Ready(second_plan), + ]); + let task = test_task(connector.clone(), extensions, FrontendOptions::default()); + + task.start().await.unwrap(); + first.wait_for_requests(1).await; + first.close_responses(); + connector.wait_for_calls(2).await; + second.wait_for_requests(1).await; + assert_eq!(task.generation(), 2); + assert_eq!( + extension.connected_generations.lock().unwrap().as_slice(), + &[1, 2] + ); + task.shutdown().await; +} + +#[tokio::test] +async fn test_response_stream_error_reconnects() { + let (first_plan, first) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let (second_plan, second) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let connector = MockConnector::new([ + ConnectPlan::Ready(first_plan), + ConnectPlan::Ready(second_plan), + ]); + let task = test_task( + connector.clone(), + FrontendHeartbeatExtensions::default(), + FrontendOptions::default(), + ); + + task.start().await.unwrap(); + first.wait_for_requests(1).await; + first.fail_responses(); + connector.wait_for_calls(2).await; + second.wait_for_requests(1).await; + assert_eq!(task.generation(), 2); + task.shutdown().await; +} + +#[tokio::test] +async fn test_failed_reconnect_is_retried() { + let extension = TestExtension::new("lifecycle"); + let extensions = FrontendHeartbeatExtensions::default(); + extensions.register(extension.clone()); + let (first_plan, first) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let (second_plan, second) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let connector = MockConnector::new([ + ConnectPlan::Ready(first_plan), + ConnectPlan::Fail, + ConnectPlan::Ready(second_plan), + ]); + let task = test_task(connector.clone(), extensions, FrontendOptions::default()); + + task.start().await.unwrap(); + first.wait_for_requests(1).await; + first.close_responses(); + connector.wait_for_calls(3).await; + second.wait_for_requests(1).await; + assert_eq!(task.generation(), 2); + assert_eq!( + extension.connected_generations.lock().unwrap().as_slice(), + &[1, 2] + ); + task.shutdown().await; +} + +#[tokio::test] +async fn test_send_failure_reconnects() { + let (first_plan, first) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + first.fail_send.store(true, Ordering::Release); + let (second_plan, second) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let connector = MockConnector::new([ + ConnectPlan::Ready(first_plan), + ConnectPlan::Ready(second_plan), + ]); + let task = test_task( + connector.clone(), + FrontendHeartbeatExtensions::default(), + FrontendOptions::default(), + ); + + task.start().await.unwrap(); + connector.wait_for_calls(2).await; + second.wait_for_requests(1).await; + assert_eq!(task.generation(), 2); + task.shutdown().await; +} + +#[tokio::test] +async fn test_concurrent_send_receive_failure_creates_one_generation() { + let (first_plan, first) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + first.fail_send.store(true, Ordering::Release); + first.fail_responses(); + let (second_plan, second) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let connector = MockConnector::new([ + ConnectPlan::Ready(first_plan), + ConnectPlan::Ready(second_plan), + ]); + let task = test_task( + connector.clone(), + FrontendHeartbeatExtensions::default(), + FrontendOptions::default(), + ); + + task.start().await.unwrap(); + connector.wait_for_calls(2).await; + second.wait_for_requests(1).await; + tokio::time::sleep(Duration::from_millis(30)).await; + assert_eq!(connector.calls.load(Ordering::Acquire), 2); + assert_eq!(task.generation(), 2); + task.shutdown().await; +} + +#[tokio::test] +async fn test_shutdown_cancels_stuck_handshake() { + let extension = TestExtension::new("lifecycle"); + let extensions = FrontendHeartbeatExtensions::default(); + extensions.register(extension.clone()); + let connector = MockConnector::new([ConnectPlan::Pending]); + let task = test_task(connector.clone(), extensions, FrontendOptions::default()); + let start_task = task.clone(); + let start = tokio::spawn(async move { start_task.start().await }); + connector.wait_for_calls(1).await; + + tokio::time::timeout(Duration::from_secs(1), task.shutdown()) + .await + .unwrap(); + start.await.unwrap().unwrap(); + assert!(extension.connected_generations.lock().unwrap().is_empty()); + assert!(!task.has_supervisor().await); +} + +#[tokio::test] +async fn test_shutdown_cancels_retry_sleep_and_joins_supervisor() { + let (plan, handle) = mock_connection(Duration::from_secs(60), Duration::from_secs(60)); + let connector = MockConnector::new([ConnectPlan::Ready(plan)]); + let task = test_task( + connector.clone(), + FrontendHeartbeatExtensions::default(), + FrontendOptions::default(), + ); + task.start().await.unwrap(); + handle.wait_for_requests(1).await; + handle.close_responses(); + tokio::time::sleep(Duration::from_millis(20)).await; + + tokio::time::timeout(Duration::from_secs(1), task.shutdown()) + .await + .unwrap(); + assert_eq!(connector.calls.load(Ordering::Acquire), 1); + assert!(!task.has_supervisor().await); +} + +#[tokio::test] +async fn test_shutdown_cancels_inflight_response_handler() { + let started = Arc::new(Notify::new()); + let extension = TestExtension::with_handler( + "pending-response", + Arc::new(PendingHandler { + started: started.clone(), + }), + ); + let extensions = FrontendHeartbeatExtensions::default(); + extensions.register(extension.clone()); + let (plan, handle) = mock_connection(Duration::from_secs(60), Duration::from_secs(60)); + let task = test_task( + MockConnector::new([ConnectPlan::Ready(plan)]), + extensions, + FrontendOptions::default(), + ); + task.start().await.unwrap(); + handle.wait_for_requests(1).await; + + let handler_started = started.notified(); + handle.send_response(HeartbeatResponse::default()); + tokio::time::timeout(Duration::from_secs(1), handler_started) + .await + .unwrap(); + + tokio::time::timeout(Duration::from_secs(1), task.shutdown()) + .await + .unwrap(); + assert!(!task.has_supervisor().await); + assert_eq!(extension.shutdown_calls.load(Ordering::Acquire), 1); +} + +#[tokio::test] +async fn test_concurrent_start_has_one_active_generation() { + let (plan, handle) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let connector = MockConnector::new([ConnectPlan::Ready(plan)]); + let task = test_task( + connector.clone(), + FrontendHeartbeatExtensions::default(), + FrontendOptions::default(), + ); + let starts = (0..8).map(|_| { + let task = task.clone(); + tokio::spawn(async move { task.start().await }) + }); + for start in starts { + start.await.unwrap().unwrap(); + } + handle.wait_for_requests(1).await; + + assert_eq!(connector.calls.load(Ordering::Acquire), 1); + assert_eq!(task.generation(), 1); + task.shutdown().await; +} + +#[tokio::test] +async fn test_shutdown_prevents_restart_and_extension_callbacks() { + let extension = TestExtension::dynamic("lifecycle", "dynamic"); + let extensions = FrontendHeartbeatExtensions::default(); + extensions.register(extension.clone()); + let (plan, handle) = mock_connection(Duration::from_secs(60), Duration::from_millis(10)); + let connector = MockConnector::new([ConnectPlan::Ready(plan)]); + let task = test_task(connector.clone(), extensions, FrontendOptions::default()); + task.start().await.unwrap(); + handle.wait_for_requests(1).await; + task.shutdown().await; + task.shutdown().await; + + let request_calls = extension.request_calls.load(Ordering::Acquire); + let connected = extension.connected_generations.lock().unwrap().clone(); + task.start().await.unwrap(); + tokio::time::sleep(Duration::from_millis(30)).await; + assert_eq!(connector.calls.load(Ordering::Acquire), 1); + assert_eq!( + extension.request_calls.load(Ordering::Acquire), + request_calls + ); + assert_eq!(*extension.connected_generations.lock().unwrap(), connected); + assert_eq!(extension.shutdown_calls.load(Ordering::Acquire), 1); } diff --git a/src/meta-srv/src/handler.rs b/src/meta-srv/src/handler.rs index 72d5fcdab8..082f902ee1 100644 --- a/src/meta-srv/src/handler.rs +++ b/src/meta-srv/src/handler.rs @@ -405,6 +405,7 @@ impl HeartbeatHandlerGroup { region_lease: acc.region_lease, mailbox_message, heartbeat_config, + extensions: Default::default(), }; Ok(res) } diff --git a/src/plugins/src/frontend.rs b/src/plugins/src/frontend.rs index ce5ace0391..195742936a 100644 --- a/src/plugins/src/frontend.rs +++ b/src/plugins/src/frontend.rs @@ -12,11 +12,14 @@ // See the License for the specific language governing permissions and // limitations under the License. +use std::sync::Arc; + use auth::{DefaultPermissionChecker, PermissionCheckerRef, UserProviderRef}; use common_base::Plugins; use common_meta::cache::CacheRegistryBuilder; use frontend::error::{IllegalAuthConfigSnafu, Result}; use frontend::frontend::FrontendOptions; +use frontend::heartbeat::FrontendHeartbeatExtensions; use frontend::instance::Instance; use frontend::instance::builder::FrontendBuilder; use snafu::ResultExt; @@ -64,6 +67,20 @@ pub async fn setup_frontend_plugins_post_build( Ok(()) } +/// Sets up heartbeat extensions after the frontend [`Instance`] is available. +/// +/// Implementations may use the instance to construct extensions and register them in the +/// [`FrontendHeartbeatExtensions`] stored in `plugins`. Registrations must be idempotent because +/// plugin setup can be invoked more than once. +pub async fn setup_frontend_heartbeat_extensions( + plugins: &mut Plugins, + _plugin_options: &[PluginOptions], + _instance: &Arc, +) -> Result<()> { + plugins.get_or_insert(FrontendHeartbeatExtensions::default); + Ok(()) +} + pub async fn start_frontend_plugins(_instance: &Instance) -> Result<()> { Ok(()) } diff --git a/src/plugins/src/lib.rs b/src/plugins/src/lib.rs index 9fcfdc08c2..be68a043d0 100644 --- a/src/plugins/src/lib.rs +++ b/src/plugins/src/lib.rs @@ -28,7 +28,8 @@ pub use flownode::{ setup_flownode_plugins_post_build, setup_flownode_plugins_pre_build, start_flownode_plugins, }; pub use frontend::{ - setup_frontend_plugins_post_build, setup_frontend_plugins_pre_build, start_frontend_plugins, + setup_frontend_heartbeat_extensions, setup_frontend_plugins_post_build, + setup_frontend_plugins_pre_build, start_frontend_plugins, }; pub use meta_srv::{ setup_metasrv_plugins_post_build, setup_metasrv_plugins_pre_build, start_metasrv_plugins,