From a5a6da186a53aee90b0d6badf2d92cebe0259371 Mon Sep 17 00:00:00 2001 From: discord9 Date: Wed, 2 Sep 2026 14:17:16 +0800 Subject: [PATCH] feat(flow): add enterprise create handler seam Signed-off-by: discord9 --- src/common/meta/src/ddl_manager.rs | 195 ++++++++++++++++++++++++++++- src/operator/src/statement/ddl.rs | 67 ++++++++-- 2 files changed, 251 insertions(+), 11 deletions(-) diff --git a/src/common/meta/src/ddl_manager.rs b/src/common/meta/src/ddl_manager.rs index d88f08918e..8178d68961 100644 --- a/src/common/meta/src/ddl_manager.rs +++ b/src/common/meta/src/ddl_manager.rs @@ -110,6 +110,8 @@ pub struct DdlManager { repartition_procedure_factory: RepartitionProcedureFactoryRef, #[cfg(feature = "enterprise")] trigger_ddl_manager: Option, + #[cfg(feature = "enterprise")] + create_flow_handler: Option, } /// This trait is responsible for handling DDL tasks about triggers. e.g., @@ -141,6 +143,23 @@ pub trait TriggerDdlManager: Send + Sync { #[cfg(feature = "enterprise")] pub type TriggerDdlManagerRef = Arc; +/// This trait is responsible for handling enterprise CREATE FLOW tasks. +#[cfg(feature = "enterprise")] +#[async_trait::async_trait] +pub trait CreateFlowHandler: Send + Sync { + async fn create_flow( + &self, + create_flow_task: CreateFlowTask, + procedure_manager: ProcedureManagerRef, + ddl_context: DdlContext, + query_context: QueryContext, + procedure_context: ProcedureContext, + ) -> Result; +} + +#[cfg(feature = "enterprise")] +pub type CreateFlowHandlerRef = Arc; + macro_rules! procedure_loader_entry { ($procedure:ident) => { ( @@ -235,6 +254,8 @@ impl DdlManager { repartition_procedure_factory, #[cfg(feature = "enterprise")] trigger_ddl_manager: None, + #[cfg(feature = "enterprise")] + create_flow_handler: None, } } @@ -244,6 +265,12 @@ impl DdlManager { self } + #[cfg(feature = "enterprise")] + pub fn with_create_flow_handler(mut self, create_flow_handler: CreateFlowHandlerRef) -> Self { + self.create_flow_handler = Some(create_flow_handler); + self + } + pub fn with_create_database_metadata_committer( mut self, committer: CreateDatabaseMetadataCommitterRef, @@ -1365,6 +1392,19 @@ async fn handle_create_flow_task( query_context: QueryContext, procedure_context: ProcedureContext, ) -> Result { + #[cfg(feature = "enterprise")] + if let Some(handler) = ddl_manager.create_flow_handler.as_ref() { + return handler + .create_flow( + create_flow_task, + ddl_manager.procedure_manager.clone(), + ddl_manager.ddl_context.clone(), + query_context, + procedure_context, + ) + .await; + } + let (id, output) = ddl_manager .submit_create_flow_task(create_flow_task.clone(), query_context, procedure_context) .await?; @@ -1538,6 +1578,7 @@ mod tests { use crate::ddl::create_database::{ AtomicCreateOutcome, CreateDatabaseMetadataCommitter, CreateDatabaseProcedure, }; + use crate::ddl::create_flow::CreateFlowProcedure; use crate::ddl::create_table::CreateTableProcedure; use crate::ddl::drop_table::DropTableProcedure; use crate::ddl::flow_meta::FlowMetadataAllocator; @@ -1555,11 +1596,13 @@ mod tests { use crate::region_registry::LeaderRegionRegistry; #[cfg(feature = "enterprise")] use crate::rpc::ddl::trigger::{CreateTriggerTask, DropTriggerTask}; + #[cfg(feature = "enterprise")] + use crate::rpc::ddl::{ + CreateFlowTask, DdlTask, QueryContext, SubmitDdlTaskRequest, SubmitDdlTaskResponse, + }; use crate::rpc::ddl::{CreatorGrantIntent, UndropTableTask}; #[cfg(not(feature = "enterprise"))] use crate::rpc::ddl::{DdlTask, PurgeDroppedTableTask, QueryContext, SubmitDdlTaskRequest}; - #[cfg(feature = "enterprise")] - use crate::rpc::ddl::{DdlTask, QueryContext, SubmitDdlTaskRequest}; use crate::sequence::SequenceBuilder; use crate::state_store::KvStateStore; use crate::test_util::{MockDatanodeManager, new_ddl_context}; @@ -1632,6 +1675,59 @@ mod tests { procedure_contexts: Mutex>, } + #[cfg(feature = "enterprise")] + #[derive(Default)] + struct RecordingCreateFlowHandler { + tasks: Mutex>, + contexts: Mutex>, + response: Mutex>, + fail: bool, + } + + #[cfg(feature = "enterprise")] + #[async_trait::async_trait] + impl super::CreateFlowHandler for RecordingCreateFlowHandler { + async fn create_flow( + &self, + create_flow_task: CreateFlowTask, + _procedure_manager: ProcedureManagerRef, + _ddl_context: DdlContext, + query_context: QueryContext, + procedure_context: ProcedureContext, + ) -> crate::error::Result { + self.tasks.lock().unwrap().push(create_flow_task); + self.contexts + .lock() + .unwrap() + .push((query_context, procedure_context)); + if self.fail { + return crate::error::UnsupportedSnafu { + operation: "test create flow handler", + } + .fail(); + } + Ok(self.response.lock().unwrap().take().unwrap_or_default()) + } + } + + #[cfg(feature = "enterprise")] + fn test_create_flow_task() -> CreateFlowTask { + CreateFlowTask { + catalog_name: "greptime".to_string(), + flow_name: "test_flow".to_string(), + source_table_names: vec![], + sink_table_name: TableName::new("greptime", "public", "sink"), + or_replace: false, + create_if_not_exists: false, + expire_after: None, + eval_interval_secs: None, + comment: String::new(), + sql: "select 1".to_string(), + flow_options: Default::default(), + eval_schedule: None, + } + } + #[cfg(feature = "enterprise")] #[async_trait::async_trait] impl super::TriggerDdlManager for RecordingTriggerDdlManager { @@ -1805,6 +1901,101 @@ mod tests { ) } + #[cfg(feature = "enterprise")] + #[tokio::test] + async fn test_create_flow_handler_dispatches_without_procedure() { + let response = SubmitDdlTaskResponse { + key: b"enterprise".to_vec(), + ..Default::default() + }; + let handler = Arc::new(RecordingCreateFlowHandler { + response: Mutex::new(Some(response)), + ..Default::default() + }); + let ddl_manager = + build_soft_drop_test_ddl_manager().with_create_flow_handler(handler.clone()); + let procedure_context = ProcedureContext { + actor: Some("test-user".to_string()), + event_context: Some( + PersistentEventContext::new(TriggerReason::Manual).with_protocol("mysql"), + ), + }; + let query_context = QueryContext { + channel: Channel::Mysql as u8, + ..Default::default() + }; + let actual = ddl_manager + .submit_ddl_task( + ExecutorContext { + query_context: Some(query_context.clone()), + actor: procedure_context.actor.clone(), + event_input: Some(ProcedureEventInput::new(TriggerReason::Manual)), + ..Default::default() + }, + SubmitDdlTaskRequest::new(DdlTask::new_create_flow(test_create_flow_task())), + ) + .await + .unwrap(); + assert_eq!(actual.key, b"enterprise"); + assert_eq!(handler.tasks.lock().unwrap().len(), 1); + assert_eq!( + handler.contexts.lock().unwrap().as_slice(), + &[(query_context, procedure_context)] + ); + } + + #[cfg(feature = "enterprise")] + #[tokio::test] + async fn test_create_flow_handler_error_does_not_fallback() { + let handler = Arc::new(RecordingCreateFlowHandler { + fail: true, + ..Default::default() + }); + let ddl_manager = + build_soft_drop_test_ddl_manager().with_create_flow_handler(handler.clone()); + let err = ddl_manager + .submit_ddl_task( + ExecutorContext { + query_context: Some(QueryContext::default()), + ..Default::default() + }, + SubmitDdlTaskRequest::new(DdlTask::new_create_flow(test_create_flow_task())), + ) + .await + .unwrap_err(); + assert!(err.to_string().contains("test create flow handler")); + assert_eq!(handler.tasks.lock().unwrap().len(), 1); + } + + #[cfg(feature = "enterprise")] + #[tokio::test] + async fn test_create_flow_without_handler_uses_ordinary_procedure() { + let ddl_manager = build_soft_drop_test_ddl_manager(); + ddl_manager.procedure_manager.start().await.unwrap(); + let err = ddl_manager + .submit_ddl_task( + ExecutorContext { + query_context: Some(QueryContext::default()), + ..Default::default() + }, + SubmitDdlTaskRequest::new(DdlTask::new_create_flow(test_create_flow_task())), + ) + .await + .unwrap_err(); + assert!(!err.to_string().contains("test create flow handler")); + assert_eq!( + ddl_manager + .procedure_manager + .list_procedures() + .await + .unwrap() + .iter() + .filter(|procedure| procedure.type_name == CreateFlowProcedure::TYPE_NAME) + .count(), + 1 + ); + } + #[cfg(feature = "enterprise")] #[tokio::test] async fn test_trigger_ddl_forwards_procedure_context() { diff --git a/src/operator/src/statement/ddl.rs b/src/operator/src/statement/ddl.rs index 59695a6e7b..06a55c4a71 100644 --- a/src/operator/src/statement/ddl.rs +++ b/src/operator/src/statement/ddl.rs @@ -127,6 +127,7 @@ struct DdlSubmitOptions { timeout: Duration, } +#[cfg(not(feature = "enterprise"))] const ALLOWED_FLOW_OPTIONS: [&str; 2] = [ DEFER_ON_MISSING_SOURCE_KEY, FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY, @@ -171,6 +172,7 @@ fn parse_ddl_options(options: &OptionMap) -> Result { Ok(DdlSubmitOptions { wait, timeout }) } +#[cfg(not(feature = "enterprise"))] fn supported_flow_options() -> String { ALLOWED_FLOW_OPTIONS.join(", ") } @@ -220,6 +222,9 @@ fn validate_and_normalize_flow_options( DEFER_ON_MISSING_SOURCE_KEY | FLOW_EXPERIMENTAL_ENABLE_INCREMENTAL_READ_KEY => { normalize_flow_bool_option(&key, &value)? } + #[cfg(feature = "enterprise")] + _ => value, + #[cfg(not(feature = "enterprise"))] _ => { return InvalidSqlSnafu { err_msg: format!( @@ -3205,6 +3210,7 @@ mod test { ); } + #[cfg(not(feature = "enterprise"))] #[test] fn test_validate_and_normalize_flow_options_unknown_option() { let err = validate_and_normalize_flow_options( @@ -3254,7 +3260,7 @@ mod test { } #[test] - fn test_validate_and_normalize_flow_options_rejects_redacted_invalid_input() { + fn test_parser_flow_option_access_key_id() { let sql = r" CREATE FLOW task_6 SINK TO schema_1.table_1 @@ -3272,11 +3278,20 @@ SELECT max(c1), min(c2) FROM schema_2.table_2;"; }; let expr = expr_helper::to_create_flow_task_expr(create_flow, &QueryContext::arc()).unwrap(); - let err = validate_and_normalize_flow_options(expr.flow_options, None).unwrap_err(); - assert!( - err.to_string() - .contains("unknown flow option 'access_key_id'") + #[cfg(not(feature = "enterprise"))] + { + let err = validate_and_normalize_flow_options(expr.flow_options, None).unwrap_err(); + assert!( + err.to_string() + .contains("unknown flow option 'access_key_id'") + ); + assert!(!err.to_string().contains("true")); + } + #[cfg(feature = "enterprise")] + assert_eq!( + validate_and_normalize_flow_options(expr.flow_options, None).unwrap(), + HashMap::from([("access_key_id".to_string(), "['true']".to_string())]) ); } @@ -3298,7 +3313,18 @@ SELECT max(c1), min(c2) FROM schema_2.table_2;"; } #[test] - fn test_schedule_and_internal_keys_rejected_as_unknown_options() { + fn test_internal_schedule_key_rejected() { + let err = validate_and_normalize_flow_options( + HashMap::from([(INTERNAL_EVAL_SCHEDULE_KEY.to_string(), "value".to_string())]), + Some(300), + ) + .unwrap_err(); + assert!(err.to_string().contains("reserved for internal use")); + } + + #[cfg(not(feature = "enterprise"))] + #[test] + fn test_schedule_keys_rejected_as_unknown_options() { for key in [ "eval_interval_anchor", "eval_interval_start", @@ -3311,15 +3337,38 @@ SELECT max(c1), min(c2) FROM schema_2.table_2;"; Some(300), ) .unwrap_err(); - assert!( err.to_string() - .contains(&format!("unknown flow option '{key}'")), - "unexpected error for {key}: {err}" + .contains(&format!("unknown flow option '{key}'")) ); } } + #[cfg(feature = "enterprise")] + #[test] + fn test_schedule_keys_preserved() { + let options = HashMap::from([ + ("eval_interval_anchor".to_string(), "anchor".to_string()), + ("eval_interval_start".to_string(), "start".to_string()), + ( + "eval_interval_missed_tick_policy".to_string(), + "missed".to_string(), + ), + ( + "eval_interval_catchup_max_runs".to_string(), + "runs".to_string(), + ), + ( + "eval_interval_catchup_max_lag".to_string(), + "lag".to_string(), + ), + ]); + assert_eq!( + validate_and_normalize_flow_options(options.clone(), Some(300)).unwrap(), + options + ); + } + #[test] fn test_internal_transport_keys_rejected_as_reserved() { for key in [