From 693c4fe202b0bffe417cebb8d1d6f5a8aba8c2ff Mon Sep 17 00:00:00 2001 From: zgq Date: Wed, 5 Aug 2026 15:42:25 +0800 Subject: [PATCH 1/5] refactor(core): add native driver registry --- .../src/datasource_compatibility.rs | 32 +- crates/chat2db-core/src/lib.rs | 62 ++- crates/chat2db-core/src/native_driver.rs | 421 ++++++++++++++++++ crates/chat2db-core/src/native_mysql.rs | 27 +- crates/chat2db-core/src/query.rs | 126 ++++-- crates/chat2db-core/src/ssh.rs | 20 +- 6 files changed, 603 insertions(+), 85 deletions(-) create mode 100644 crates/chat2db-core/src/native_driver.rs diff --git a/crates/chat2db-core/src/datasource_compatibility.rs b/crates/chat2db-core/src/datasource_compatibility.rs index 01ee5b3..fe18253 100644 --- a/crates/chat2db-core/src/datasource_compatibility.rs +++ b/crates/chat2db-core/src/datasource_compatibility.rs @@ -182,11 +182,8 @@ impl Application { ) -> Result { let database_type = if database_type.trim().is_empty() { let datasource = self.get_datasource(datasource_id).await?; - if self.is_native_mysql_driver(&datasource.driver_id) { - "MYSQL".to_owned() - } else { - datasource.driver_id - } + self.native_database_type_for_driver(&datasource.driver_id) + .unwrap_or(datasource.driver_id) } else { database_type.trim().to_owned() }; @@ -243,27 +240,30 @@ impl Application { }) } - /// Returns an explicit no-JAR result for `MySQL` driver mutations handled by `mysql_async`. + /// Returns an explicit no-artifact result for a registered native Rust driver. /// /// # Errors /// - /// Returns invalid-request for non-MySQL database types. + /// Returns invalid-request for database types without a registered native driver. pub fn native_driver_compatibility( &self, database_type: &str, action: NativeDriverAction, ) -> Result { - if !database_type.trim().eq_ignore_ascii_case("mysql") { - return Err(AppError::invalid( - "native_driver_not_available", - "the requested database type is not implemented by a native Rust driver", - )); - } + let driver = self + .native_driver_for_database_type(database_type) + .ok_or_else(|| { + AppError::invalid( + "native_driver_not_available", + "the requested database type is not implemented by a native Rust driver", + ) + })?; + let descriptor = driver.descriptor(); Ok(NativeDriverCompatibility { - database_type: "MYSQL".to_owned(), - driver_id: "mysql".to_owned(), + database_type: driver.database_types()[0].to_owned(), + driver_id: descriptor.driver_id, action, - implementation: "mysql_async".to_owned(), + implementation: driver.implementation().to_owned(), artifact_required: false, changed: false, }) diff --git a/crates/chat2db-core/src/lib.rs b/crates/chat2db-core/src/lib.rs index 9ce9f5c..53d8935 100644 --- a/crates/chat2db-core/src/lib.rs +++ b/crates/chat2db-core/src/lib.rs @@ -17,6 +17,7 @@ mod mysql_dashboard; pub mod mysql_ddl; mod mysql_schema_diff; mod mysql_workspace; +mod native_driver; mod native_mysql; mod operation; mod query; @@ -61,6 +62,7 @@ pub use large_value::{ pub use legacy_community_import::LegacyCommunityImportOutcome; pub use operation::OperationSubscription; pub use query::{MysqlConsoleCancellation, MysqlConsoleRequest, MysqlConsoleResult}; +pub use query::{NativeConsoleCancellation, NativeConsoleRequest, NativeConsoleResult}; pub use transfer::TransferArtifactDownload; use engine_manager::{ @@ -84,6 +86,7 @@ pub(crate) struct ApplicationInner { engine: EngineProvider, drivers: Vec, managed_driver_ids: Option>, + native_drivers: native_driver::NativeDriverRegistry, large_values: large_value::LargeValueStore, agent_runs: AgentRunHub, operations: OperationHub, @@ -254,6 +257,7 @@ impl Application { .collect() }); let drivers = drivers.unwrap_or_default(); + let native_drivers = native_driver::NativeDriverRegistry::built_in(); Self { inner: Arc::new(ApplicationInner { started_at: Instant::now(), @@ -262,6 +266,7 @@ impl Application { engine, drivers, managed_driver_ids, + native_drivers, large_values: large_value::LargeValueStore::default(), agent_runs: AgentRunHub::new(), operations: OperationHub::new(), @@ -483,11 +488,13 @@ impl Application { #[must_use] pub fn list_drivers(&self) -> JdbcDriverList { let mut items = self.inner.drivers.clone(); - if !items - .iter() - .any(|driver| driver.driver_id.eq_ignore_ascii_case("mysql")) - { - items.push(datasource_compatibility::native_mysql_driver()); + for descriptor in self.inner.native_drivers.descriptors() { + if !items + .iter() + .any(|driver| driver.driver_id.eq_ignore_ascii_case(&descriptor.driver_id)) + { + items.push(descriptor); + } } JdbcDriverList { items } } @@ -511,8 +518,8 @@ impl Application { )); } self.require_managed_driver(driver_id)?; - if self.is_native_mysql_driver(driver_id) { - return native_mysql::test_connection(&connection).await; + if let Some(driver) = self.native_driver_for_driver_id(driver_id) { + return driver.connection().test_connection(&connection).await; } let engine = self.require_engine().await?; let session = datasource_session::open_datasource_session( @@ -621,7 +628,7 @@ impl Application { } fn require_managed_driver(&self, driver_id: &str) -> Result<(), AppError> { - if self.is_native_mysql_driver(driver_id) { + if self.native_driver_for_driver_id(driver_id).is_some() { return Ok(()); } match &self.inner.managed_driver_ids { @@ -630,22 +637,28 @@ impl Application { } } - pub(crate) fn is_native_mysql_driver(&self, driver_id: &str) -> bool { - if driver_id.eq_ignore_ascii_case("mysql") { - return true; - } + pub(crate) fn native_driver_for_driver_id( + &self, + driver_id: &str, + ) -> Option> { self.inner - .drivers - .iter() - .find(|driver| driver.driver_id == driver_id) - .is_some_and(|driver| { - format!( - "{} {} {} {}", - driver.pack_id, driver.name, driver.driver_id, driver.driver_class - ) - .to_ascii_lowercase() - .contains("mysql") - }) + .native_drivers + .driver_for_driver_id(driver_id, &self.inner.drivers) + } + + pub(crate) fn native_driver_for_database_type( + &self, + database_type: &str, + ) -> Option> { + self.inner + .native_drivers + .driver_for_database_type(database_type) + } + + pub(crate) fn native_database_type_for_driver(&self, driver_id: &str) -> Option { + self.native_driver_for_driver_id(driver_id) + .and_then(|driver| driver.database_types().first().copied()) + .map(str::to_owned) } async fn require_managed_driver_for_update( @@ -654,6 +667,9 @@ impl Application { datasource_id: &str, driver_id: &str, ) -> Result<(), AppError> { + if self.native_driver_for_driver_id(driver_id).is_some() { + return Ok(()); + } let Some(driver_ids) = &self.inner.managed_driver_ids else { return Ok(()); }; diff --git a/crates/chat2db-core/src/native_driver.rs b/crates/chat2db-core/src/native_driver.rs new file mode 100644 index 0000000..ef017e9 --- /dev/null +++ b/crates/chat2db-core/src/native_driver.rs @@ -0,0 +1,421 @@ +use std::{collections::HashSet, sync::Arc}; + +use async_trait::async_trait; +use chat2db_contract::{DatasourceConnection, JdbcDriver, ResultMetadata}; +use chat2db_storage::Storage; +use tokio::sync::watch; +use tokio_util::sync::CancellationToken; + +use crate::{ + AppError, Application, + datasource_session::ResolvedDatasourceConnection, + native_mysql, + operation::CancellationRequest, + query::{ + DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, PreparedQuery, + QueryTaskError, + }, +}; + +/// Database connection operations implemented by one native Rust driver. +#[async_trait] +pub(crate) trait NativeConnectionDriver: Send + Sync { + async fn test_connection(&self, connection: &DatasourceConnection) -> Result<(), AppError>; + + async fn test_connection_with_local_port( + &self, + _connection: &DatasourceConnection, + ) -> Result, AppError> { + Err(AppError::invalid( + "ssh_driver_not_supported", + "SSH forwarding is not supported by this native Rust driver", + )) + } +} + +/// Query and Console operations implemented by one native Rust driver. +#[async_trait] +pub(crate) trait NativeQueryDriver: Send + Sync { + fn is_read_candidate(&self, sql: &str) -> Result; + + fn validate_query(&self, query: &PreparedQuery) -> Result<(), AppError>; + + async fn execute_query_task( + &self, + application: &Application, + operation_id: &str, + cancellation: watch::Receiver, + query: PreparedQuery, + storage: Storage, + resolved: ResolvedDatasourceConnection, + ) -> Result; + + async fn execute_update( + &self, + resolved: ResolvedDatasourceConnection, + sql: String, + cancellation: CancellationToken, + ) -> Result; + + async fn execute_console( + &self, + application: &Application, + request: NativeConsoleRequest, + cancellation: watch::Receiver, + force_read_only: bool, + ) -> Result, AppError>; +} + +/// Runtime-polymorphic native Rust database driver. +/// +/// Optional capability accessors allow a driver to participate only in the +/// product surfaces it implements. Additional capability traits are attached +/// here as native metadata and dialect services are migrated. +pub(crate) trait NativeDriver: Send + Sync { + fn id(&self) -> &'static str; + + fn implementation(&self) -> &'static str; + + fn database_types(&self) -> &'static [&'static str]; + + fn descriptor(&self) -> JdbcDriver; + + fn matches_driver(&self, driver_id: &str, descriptor: Option<&JdbcDriver>) -> bool; + + fn connection(&self) -> &dyn NativeConnectionDriver; + + fn query(&self) -> Option<&dyn NativeQueryDriver> { + None + } +} + +/// Immutable registry used to select native implementations at runtime. +#[derive(Clone)] +pub(crate) struct NativeDriverRegistry { + drivers: Arc<[Arc]>, +} + +impl NativeDriverRegistry { + pub(crate) fn built_in() -> Self { + Self::try_new(vec![Arc::new(MysqlNativeDriver)]) + .expect("built-in native drivers must have unique identities") + } + + fn try_new(drivers: Vec>) -> Result { + let mut ids = HashSet::new(); + let mut database_types = HashSet::new(); + for driver in &drivers { + let id = driver.id().trim().to_ascii_lowercase(); + if id.is_empty() || !ids.insert(id) { + return Err(AppError::invalid( + "invalid_native_driver_registry", + "native driver ids must be non-empty and unique", + )); + } + if driver.database_types().is_empty() { + return Err(AppError::invalid( + "invalid_native_driver_registry", + "native drivers must declare at least one database type", + )); + } + for database_type in driver.database_types() { + let database_type = database_type.trim().to_ascii_lowercase(); + if database_type.is_empty() || !database_types.insert(database_type) { + return Err(AppError::invalid( + "invalid_native_driver_registry", + "native database types must be non-empty and unique", + )); + } + } + } + Ok(Self { + drivers: drivers.into(), + }) + } + + pub(crate) fn descriptors(&self) -> impl Iterator + '_ { + self.drivers.iter().map(|driver| driver.descriptor()) + } + + pub(crate) fn driver_for_database_type( + &self, + database_type: &str, + ) -> Option> { + self.drivers + .iter() + .find(|driver| { + driver + .database_types() + .iter() + .any(|candidate| candidate.eq_ignore_ascii_case(database_type.trim())) + }) + .cloned() + } + + pub(crate) fn driver_for_driver_id( + &self, + driver_id: &str, + managed_drivers: &[JdbcDriver], + ) -> Option> { + let descriptor = managed_drivers + .iter() + .find(|driver| driver.driver_id == driver_id); + self.drivers + .iter() + .find(|driver| driver.matches_driver(driver_id, descriptor)) + .cloned() + } +} + +struct MysqlNativeDriver; + +impl NativeDriver for MysqlNativeDriver { + fn id(&self) -> &'static str { + "mysql" + } + + fn implementation(&self) -> &'static str { + "mysql_async" + } + + fn database_types(&self) -> &'static [&'static str] { + &["MYSQL"] + } + + fn descriptor(&self) -> JdbcDriver { + crate::datasource_compatibility::native_mysql_driver() + } + + fn matches_driver(&self, driver_id: &str, descriptor: Option<&JdbcDriver>) -> bool { + if driver_id.eq_ignore_ascii_case(self.id()) { + return true; + } + descriptor.is_some_and(|driver| { + format!( + "{} {} {} {}", + driver.pack_id, driver.name, driver.driver_id, driver.driver_class + ) + .to_ascii_lowercase() + .contains("mysql") + }) + } + + fn connection(&self) -> &dyn NativeConnectionDriver { + self + } + + fn query(&self) -> Option<&dyn NativeQueryDriver> { + Some(self) + } +} + +#[async_trait] +impl NativeConnectionDriver for MysqlNativeDriver { + async fn test_connection(&self, connection: &DatasourceConnection) -> Result<(), AppError> { + native_mysql::test_connection(connection).await + } + + async fn test_connection_with_local_port( + &self, + connection: &DatasourceConnection, + ) -> Result, AppError> { + native_mysql::test_connection_with_local_port(connection).await + } +} + +#[async_trait] +impl NativeQueryDriver for MysqlNativeDriver { + fn is_read_candidate(&self, sql: &str) -> Result { + native_mysql::is_native_read_candidate(sql) + } + + fn validate_query(&self, query: &PreparedQuery) -> Result<(), AppError> { + native_mysql::validate_query(query) + } + + async fn execute_query_task( + &self, + application: &Application, + operation_id: &str, + cancellation: watch::Receiver, + query: PreparedQuery, + storage: Storage, + resolved: ResolvedDatasourceConnection, + ) -> Result { + native_mysql::execute_query_task( + application, + operation_id, + cancellation, + query, + storage, + resolved, + ) + .await + } + + async fn execute_update( + &self, + resolved: ResolvedDatasourceConnection, + sql: String, + cancellation: CancellationToken, + ) -> Result { + native_mysql::execute_update(resolved, sql, cancellation).await + } + + async fn execute_console( + &self, + application: &Application, + request: NativeConsoleRequest, + cancellation: watch::Receiver, + force_read_only: bool, + ) -> Result, AppError> { + native_mysql::execute_console(application, request, cancellation, force_read_only).await + } +} + +#[cfg(test)] +mod tests { + use super::*; + + struct FakePostgresDriver; + + impl NativeDriver for FakePostgresDriver { + fn id(&self) -> &'static str { + "postgresql" + } + + fn implementation(&self) -> &'static str { + "fake_postgres" + } + + fn database_types(&self) -> &'static [&'static str] { + &["POSTGRESQL", "POSTGRES"] + } + + fn descriptor(&self) -> JdbcDriver { + JdbcDriver { + pack_id: "native:fake_postgres".to_owned(), + name: "PostgreSQL (test native Rust)".to_owned(), + version: "test".to_owned(), + driver_id: self.id().to_owned(), + driver_class: "rust:fake_postgres".to_owned(), + artifact_count: 0, + artifact_bytes: "0".to_owned(), + } + } + + fn matches_driver(&self, driver_id: &str, descriptor: Option<&JdbcDriver>) -> bool { + driver_id.eq_ignore_ascii_case(self.id()) + || descriptor.is_some_and(|driver| { + driver + .driver_class + .eq_ignore_ascii_case("org.postgresql.Driver") + }) + } + + fn connection(&self) -> &dyn NativeConnectionDriver { + self + } + } + + #[async_trait] + impl NativeConnectionDriver for FakePostgresDriver { + async fn test_connection( + &self, + _connection: &DatasourceConnection, + ) -> Result<(), AppError> { + Ok(()) + } + } + + #[test] + fn registry_selects_runtime_driver_by_database_type_and_driver_id() { + let registry = NativeDriverRegistry::try_new(vec![Arc::new(FakePostgresDriver)]) + .expect("registry is valid"); + + assert_eq!( + registry + .driver_for_database_type("postgres") + .expect("database type resolves") + .id(), + "postgresql" + ); + assert_eq!( + registry + .driver_for_driver_id("POSTGRESQL", &[]) + .expect("driver id resolves") + .implementation(), + "fake_postgres" + ); + } + + #[test] + fn registry_uses_managed_descriptor_aliases_without_owning_driver_jars() { + let registry = NativeDriverRegistry::try_new(vec![Arc::new(FakePostgresDriver)]) + .expect("registry is valid"); + let managed = vec![JdbcDriver { + pack_id: "postgresql-42".to_owned(), + name: "PostgreSQL JDBC".to_owned(), + version: "42".to_owned(), + driver_id: "managed-pg".to_owned(), + driver_class: "org.postgresql.Driver".to_owned(), + artifact_count: 1, + artifact_bytes: "1".to_owned(), + }]; + + assert_eq!( + registry + .driver_for_driver_id("managed-pg", &managed) + .expect("managed descriptor resolves") + .id(), + "postgresql" + ); + } + + #[test] + fn registry_rejects_duplicate_database_type_ownership() { + struct DuplicatePostgresDriver; + + impl NativeDriver for DuplicatePostgresDriver { + fn id(&self) -> &'static str { + "duplicate-postgresql" + } + + fn implementation(&self) -> &'static str { + "duplicate" + } + + fn database_types(&self) -> &'static [&'static str] { + &["postgresql"] + } + + fn descriptor(&self) -> JdbcDriver { + FakePostgresDriver.descriptor() + } + + fn matches_driver(&self, _driver_id: &str, _descriptor: Option<&JdbcDriver>) -> bool { + false + } + + fn connection(&self) -> &dyn NativeConnectionDriver { + self + } + } + + #[async_trait] + impl NativeConnectionDriver for DuplicatePostgresDriver { + async fn test_connection( + &self, + _connection: &DatasourceConnection, + ) -> Result<(), AppError> { + Ok(()) + } + } + + let result = NativeDriverRegistry::try_new(vec![ + Arc::new(FakePostgresDriver), + Arc::new(DuplicatePostgresDriver), + ]); + assert!(result.is_err()); + } +} diff --git a/crates/chat2db-core/src/native_mysql.rs b/crates/chat2db-core/src/native_mysql.rs index daf423f..c7a0e08 100644 --- a/crates/chat2db-core/src/native_mysql.rs +++ b/crates/chat2db-core/src/native_mysql.rs @@ -40,8 +40,8 @@ use crate::{ datasource_session::{ResolvedDatasourceConnection, resolve_datasource_connection}, operation::CancellationRequest, query::{ - DatabaseWriteError, MysqlConsoleRequest, MysqlConsoleResult, PreparedQuery, QueryTaskError, - RetainedWriter, + DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, PreparedQuery, + QueryTaskError, RetainedWriter, }, ssh::{SshTunnel, SshTunnelIdentity, mysql_target, rewrite_mysql_target}, }; @@ -1772,7 +1772,7 @@ pub(crate) async fn start_table_preview( } struct ConsoleStatementExecution { - results: Vec, + results: Vec, failure: Option, } @@ -1783,10 +1783,10 @@ enum ConsoleExecutionError { pub(crate) async fn execute_console( application: &Application, - request: MysqlConsoleRequest, + request: NativeConsoleRequest, mut cancellation: watch::Receiver, force_read_only: bool, -) -> Result, AppError> { +) -> Result, AppError> { let (statements, page_offset, page_end) = prepare_console_statements(&request)?; let initial_cancellation = { cancellation.borrow().clone() }; @@ -1879,7 +1879,7 @@ pub(crate) async fn execute_console( } fn prepare_console_statements( - request: &MysqlConsoleRequest, + request: &NativeConsoleRequest, ) -> Result<(Vec, u64, u64), AppError> { let (page_offset, page_end) = validate_console_request(request)?; let mut statements = if request.single { @@ -2136,7 +2136,7 @@ async fn execute_console_statement( } if retain { - results.push(MysqlConsoleResult { + results.push(NativeConsoleResult { statement_sequence, result_set_id: current_result_set_id, sql: statement.to_owned(), @@ -2158,7 +2158,7 @@ async fn execute_console_statement( } if results.is_empty() && selected_result_set_id.is_none() { - results.push(MysqlConsoleResult { + results.push(NativeConsoleResult { statement_sequence, result_set_id: None, sql: statement.to_owned(), @@ -2237,9 +2237,9 @@ fn console_failure_result( sql: String, error: &AppError, duration_ms: u64, -) -> MysqlConsoleResult { +) -> NativeConsoleResult { let api_error = error.api_error(); - MysqlConsoleResult { + NativeConsoleResult { statement_sequence, result_set_id: None, sql, @@ -2255,7 +2255,7 @@ fn console_failure_result( } } -fn validate_console_request(request: &MysqlConsoleRequest) -> Result<(u64, u64), AppError> { +fn validate_console_request(request: &NativeConsoleRequest) -> Result<(u64, u64), AppError> { if request.datasource_id.trim().is_empty() { return Err(AppError::invalid( "invalid_mysql_console_request", @@ -4547,7 +4547,10 @@ pub(crate) async fn resolve_native_connection( ) -> Result { let storage = application.require_storage()?; let resolved = resolve_datasource_connection(&storage, datasource_id).await?; - if !application.is_native_mysql_driver(&resolved.driver_id) { + if application + .native_driver_for_driver_id(&resolved.driver_id) + .is_none_or(|driver| driver.id() != "mysql") + { return Err(AppError::invalid( "mysql_driver_mismatch", "The datasource is not configured with a MySQL driver", diff --git a/crates/chat2db-core/src/query.rs b/crates/chat2db-core/src/query.rs index 65b61e1..c1bcc66 100644 --- a/crates/chat2db-core/src/query.rs +++ b/crates/chat2db-core/src/query.rs @@ -36,7 +36,7 @@ pub(crate) struct PreparedQuery { #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] #[allow(clippy::struct_excessive_bools)] #[serde(rename_all = "camelCase")] -pub struct MysqlConsoleRequest { +pub struct NativeConsoleRequest { /// Opaque datasource id resolved by Core. pub datasource_id: String, /// Optional `MySQL` database selected on the Console connection. @@ -68,7 +68,7 @@ pub struct MysqlConsoleRequest { /// One statement result emitted by native `MySQL` Console execution. #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] #[serde(rename_all = "camelCase")] -pub struct MysqlConsoleResult { +pub struct NativeConsoleResult { /// One-based statement position in the submitted script. pub statement_sequence: u32, /// One-based tabular result-set position within the statement. @@ -99,11 +99,11 @@ pub struct MysqlConsoleResult { /// Cloneable cancellation source for one native `MySQL` Console execution. #[derive(Debug, Clone)] -pub struct MysqlConsoleCancellation { +pub struct NativeConsoleCancellation { sender: watch::Sender, } -impl MysqlConsoleCancellation { +impl NativeConsoleCancellation { /// Creates an active cancellation source. #[must_use] pub fn new() -> Self { @@ -134,7 +134,7 @@ impl MysqlConsoleCancellation { } } -impl Default for MysqlConsoleCancellation { +impl Default for NativeConsoleCancellation { fn default() -> Self { Self::new() } @@ -144,12 +144,17 @@ const fn default_console_error_continue() -> bool { true } +pub type MysqlConsoleRequest = NativeConsoleRequest; +pub type MysqlConsoleResult = NativeConsoleResult; +pub type MysqlConsoleCancellation = NativeConsoleCancellation; + enum QueryBackend { Java { engine: EngineLease, resolved: ResolvedDatasourceConnection, }, - NativeMysql { + Native { + driver: std::sync::Arc, resolved: ResolvedDatasourceConnection, }, } @@ -193,7 +198,30 @@ impl Application { request: MysqlConsoleRequest, cancellation: MysqlConsoleCancellation, ) -> Result, AppError> { - crate::native_mysql::execute_console(self, request, cancellation.subscribe(), false).await + self.execute_native_console(request, cancellation).await + } + + /// Executes a native-driver Console request on the runtime-selected driver. + /// + /// # Errors + /// + /// Returns datasource, capability, validation, connection, or execution errors. + pub async fn execute_native_console( + &self, + request: NativeConsoleRequest, + cancellation: NativeConsoleCancellation, + ) -> Result, AppError> { + self.execute_native_console_with_mode(request, cancellation, false) + .await + } + + pub(crate) async fn execute_native_read_console( + &self, + request: NativeConsoleRequest, + cancellation: NativeConsoleCancellation, + ) -> Result, AppError> { + self.execute_native_console_with_mode(request, cancellation, true) + .await } pub(crate) async fn execute_mysql_read_console( @@ -201,7 +229,35 @@ impl Application { request: MysqlConsoleRequest, cancellation: MysqlConsoleCancellation, ) -> Result, AppError> { - crate::native_mysql::execute_console(self, request, cancellation.subscribe(), true).await + self.execute_native_read_console(request, cancellation) + .await + } + + async fn execute_native_console_with_mode( + &self, + request: NativeConsoleRequest, + cancellation: NativeConsoleCancellation, + force_read_only: bool, + ) -> Result, AppError> { + let storage = self.require_storage()?; + let resolved = resolve_datasource_connection(&storage, &request.datasource_id).await?; + let driver = self + .native_driver_for_driver_id(&resolved.driver_id) + .ok_or_else(|| { + AppError::invalid( + "native_driver_not_available", + "The datasource does not have a native Rust driver", + ) + })?; + let query = driver.query().ok_or_else(|| { + AppError::invalid( + "native_query_not_supported", + "The native Rust driver does not implement query execution", + ) + })?; + query + .execute_console(self, request, cancellation.subscribe(), force_read_only) + .await } /// Accepts a query for asynchronous execution and returns its operation id. @@ -286,11 +342,23 @@ impl Application { let _engine = self.require_engine().await?; } let resolved = resolve_datasource_connection(&storage, &prepared.datasource_id).await?; - let backend = if self.is_native_mysql_driver(&resolved.driver_id) - && crate::native_mysql::is_native_read_candidate(&prepared.sql)? - { - crate::native_mysql::validate_query(&prepared)?; - QueryBackend::NativeMysql { resolved } + let backend = if let Some(driver) = self.native_driver_for_driver_id(&resolved.driver_id) { + if let Some(query) = driver.query() { + if query.is_read_candidate(&prepared.sql)? { + query.validate_query(&prepared)?; + QueryBackend::Native { driver, resolved } + } else { + QueryBackend::Java { + engine: self.require_engine().await?, + resolved, + } + } + } else { + QueryBackend::Java { + engine: self.require_engine().await?, + resolved, + } + } } else { QueryBackend::Java { engine: self.require_engine().await?, @@ -371,17 +439,21 @@ impl Application { ) .await } - QueryBackend::NativeMysql { resolved } => { - crate::native_mysql::execute_query_task( - self, - &operation_id, - cancellation, - query, - storage, - resolved, - ) - .await - } + QueryBackend::Native { driver, resolved } => match driver.query() { + Some(native_query) => { + native_query + .execute_query_task( + self, + &operation_id, + cancellation, + query, + storage, + resolved, + ) + .await + } + None => Err(QueryTaskError::Failed(AppError::internal())), + }, }; match outcome { Ok(result) => { @@ -602,8 +674,10 @@ impl Application { ), ))); } - if self.is_native_mysql_driver(&resolved.driver_id) { - return crate::native_mysql::execute_update(resolved, sql, cancellation).await; + if let Some(driver) = self.native_driver_for_driver_id(&resolved.driver_id) + && let Some(query) = driver.query() + { + return query.execute_update(resolved, sql, cancellation).await; } Err(DatabaseWriteError::not_started(AppError::invalid( "mysql_driver_mismatch", diff --git a/crates/chat2db-core/src/ssh.rs b/crates/chat2db-core/src/ssh.rs index aa5c894..94e14e9 100644 --- a/crates/chat2db-core/src/ssh.rs +++ b/crates/chat2db-core/src/ssh.rs @@ -24,7 +24,7 @@ use tokio::{ }; use url::Url; -use crate::{AppError, Application, native_mysql}; +use crate::{AppError, Application}; const SSH_CONNECT_TIMEOUT: Duration = Duration::from_secs(15); const SSH_AUTH_TIMEOUT: Duration = Duration::from_secs(15); @@ -210,15 +210,19 @@ impl Application { }); }; self.require_managed_driver(&request.driver_id)?; - if !self.is_native_mysql_driver(&request.driver_id) { - return Err(AppError::invalid( - "ssh_driver_not_supported", - "SSH forwarding is currently implemented for native MySQL only", - )); - } + let driver = self + .native_driver_for_driver_id(&request.driver_id) + .ok_or_else(|| { + AppError::invalid( + "ssh_driver_not_supported", + "SSH forwarding requires a native Rust driver", + ) + })?; let mut forwarded = request.connection; forwarded.ssh = Some(ssh); - let local_port = native_mysql::test_connection_with_local_port(&forwarded) + let local_port = driver + .connection() + .test_connection_with_local_port(&forwarded) .await? .ok_or_else(AppError::internal)?; Ok(SshDatasourcePreConnectResult { From f5420cd9221be59d996b57fc778b070df51d4d95 Mon Sep 17 00:00:00 2001 From: zgq Date: Wed, 5 Aug 2026 16:00:07 +0800 Subject: [PATCH 2/5] refactor(core): route metadata through native driver spi --- crates/chat2db-core/src/community.rs | 326 +++++------ crates/chat2db-core/src/lib.rs | 17 + crates/chat2db-core/src/mysql_workspace.rs | 26 +- crates/chat2db-core/src/native_driver.rs | 541 ++++++++++++++++++ .../chat2db-core/src/native_driver_types.rs | 360 ++++++++++++ crates/chat2db-core/src/native_mysql.rs | 97 ++-- .../src/transfer/class_generation.rs | 40 +- 7 files changed, 1165 insertions(+), 242 deletions(-) create mode 100644 crates/chat2db-core/src/native_driver_types.rs diff --git a/crates/chat2db-core/src/community.rs b/crates/chat2db-core/src/community.rs index 1a9f38c..411c739 100644 --- a/crates/chat2db-core/src/community.rs +++ b/crates/chat2db-core/src/community.rs @@ -93,7 +93,6 @@ use crate::{ AppError, Application, datasource_session::{SessionReadOnly, open_datasource_session, resolve_datasource_connection}, engine_manager::EngineLease, - native_mysql, }; impl Application { @@ -122,8 +121,10 @@ impl Application { &self, request: ListCommunitySchemasRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_schemas(self, &request.datasource_id).await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_schemas(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -160,8 +161,10 @@ impl Application { &self, request: ListCommunityDatabasesRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_databases(self, &request.datasource_id).await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_databases(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -197,14 +200,10 @@ impl Application { &self, request: ListCommunityTablesRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_tables( - self, - &request.datasource_id, - &request.database_name, - &request.table_name_pattern, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_tables(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -250,15 +249,10 @@ impl Application { &self, request: ListCommunityColumnsRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_columns( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &request.table_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_columns(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -307,14 +301,18 @@ impl Application { table_name: &str, column_names: &[String], ) -> Result<(), AppError> { - native_mysql::validate_column_reorder( - self, - datasource_id, - database_name, - table_name, - column_names, - ) - .await + let driver = self + .require_native_driver_for_datasource(datasource_id) + .await?; + let tables = driver.tables().ok_or_else(|| { + AppError::invalid( + "native_table_capability_not_available", + "The native Rust driver does not implement table operations", + ) + })?; + tables + .validate_column_reorder(self, datasource_id, database_name, table_name, column_names) + .await } /// Reads a `MySQL` table definition through the native driver without starting Java. @@ -329,7 +327,18 @@ impl Application { schema_name: &str, table_name: &str, ) -> Result { - native_mysql::table_ddl(self, data_source_id, database_name, schema_name, table_name).await + let driver = self + .require_native_driver_for_datasource(data_source_id) + .await?; + let tables = driver.tables().ok_or_else(|| { + AppError::invalid( + "native_table_capability_not_available", + "The native Rust driver does not implement table operations", + ) + })?; + tables + .table_ddl(self, data_source_id, database_name, schema_name, table_name) + .await } /// Lists indexes through Community metadata using a forced read-only session. @@ -341,15 +350,10 @@ impl Application { &self, request: ListCommunityIndexesRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_indexes( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &request.table_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_indexes(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -395,15 +399,10 @@ impl Application { &self, request: ListCommunityViewsRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_views( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &request.view_name_pattern, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_views(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -450,15 +449,10 @@ impl Application { request: ListCommunityViewsRequest, ) -> Result { let view_name = request.view_name_pattern.clone(); - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::get_view( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &view_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.get_view(self, request.into()).await; } self.list_community_views(request) .await? @@ -482,14 +476,10 @@ impl Application { &self, request: ListCommunityTableKeysRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_imported_keys( - self, - &request.datasource_id, - &request.database_name, - &request.table_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_imported_keys(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -535,14 +525,10 @@ impl Application { &self, request: ListCommunityTableKeysRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_exported_keys( - self, - &request.datasource_id, - &request.database_name, - &request.table_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_exported_keys(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -588,15 +574,10 @@ impl Application { &self, request: ListCommunityTableKeysRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_primary_keys( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &request.table_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_primary_keys(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -642,14 +623,10 @@ impl Application { &self, request: ListCommunityFunctionsRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_functions( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_functions(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -687,15 +664,10 @@ impl Application { &self, request: GetCommunityFunctionRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::get_function( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &request.function_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.get_function(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -739,15 +711,12 @@ impl Application { &self, request: GetCommunityFunctionRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_function_parameters( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &request.function_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata + .list_function_parameters(self, request.into()) + .await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -796,14 +765,10 @@ impl Application { &self, request: ListCommunityProceduresRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_procedures( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_procedures(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -841,15 +806,10 @@ impl Application { &self, request: GetCommunityProcedureRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::get_procedure( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &request.procedure_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.get_procedure(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -893,15 +853,12 @@ impl Application { &self, request: GetCommunityProcedureRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_procedure_parameters( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &request.procedure_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata + .list_procedure_parameters(self, request.into()) + .await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -950,13 +907,21 @@ impl Application { &self, request: PreviewCommunityRoutineInvocationRequest, ) -> Result { - if !native_mysql::is_mysql_database_type(&request.database_type) { - return Err(AppError::invalid( - "invalid_community_routine_invocation_request", - "routine invocation preview supports only MySQL", - )); - } - native_mysql::preview_routine_invocation(self, request).await + let driver = self + .native_driver_for_database_type(&request.database_type) + .ok_or_else(|| { + AppError::invalid( + "invalid_community_routine_invocation_request", + "routine invocation preview requires a native Rust driver", + ) + })?; + let routines = driver.routines().ok_or_else(|| { + AppError::invalid( + "native_routine_capability_not_available", + "The native Rust driver does not implement routine operations", + ) + })?; + routines.preview_invocation(self, request.into()).await } /// Previews the compensating `MySQL` routine-replacement script. @@ -969,13 +934,21 @@ impl Application { &self, request: &CommunityRoutineMigrationRequest, ) -> Result { - if !native_mysql::is_mysql_database_type(&request.database_type) { - return Err(AppError::invalid( - "invalid_community_routine_migration_request", - "routine migration supports only MySQL", - )); - } - native_mysql::preview_routine_migration(request) + let driver = self + .native_driver_for_database_type(&request.database_type) + .ok_or_else(|| { + AppError::invalid( + "invalid_community_routine_migration_request", + "routine migration requires a native Rust driver", + ) + })?; + let routines = driver.routines().ok_or_else(|| { + AppError::invalid( + "native_routine_capability_not_available", + "The native Rust driver does not implement routine operations", + ) + })?; + routines.preview_migration(request.clone().into()) } /// Replaces one `MySQL` routine and restores its before-image when apply fails. @@ -989,13 +962,21 @@ impl Application { &self, request: CommunityRoutineMigrationRequest, ) -> Result { - if !native_mysql::is_mysql_database_type(&request.database_type) { - return Err(AppError::invalid( - "invalid_community_routine_migration_request", - "routine migration supports only MySQL", - )); - } - native_mysql::execute_routine_migration(self, request).await + let driver = self + .native_driver_for_database_type(&request.database_type) + .ok_or_else(|| { + AppError::invalid( + "invalid_community_routine_migration_request", + "routine migration requires a native Rust driver", + ) + })?; + let routines = driver.routines().ok_or_else(|| { + AppError::invalid( + "native_routine_capability_not_available", + "The native Rust driver does not implement routine operations", + ) + })?; + routines.execute_migration(self, request.into()).await } /// Lists triggers through Community metadata using a forced read-only session. @@ -1007,14 +988,10 @@ impl Application { &self, request: ListCommunityTriggersRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::list_triggers( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.list_triggers(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -1052,15 +1029,10 @@ impl Application { &self, request: GetCommunityTriggerRequest, ) -> Result { - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::get_trigger( - self, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &request.trigger_name, - ) - .await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(metadata) = driver.metadata() + { + return metadata.get_trigger(self, request.into()).await; } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -1169,8 +1141,12 @@ impl Application { "datasourceId cannot be empty", )); } - if native_mysql::is_mysql_database_type(&request.database_type) { - return native_mysql::start_table_preview(self, request, row_limit).await; + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(tables) = driver.tables() + { + return tables + .start_table_preview(self, request.into(), row_limit) + .await; } let engine = self.require_community_engine().await?; diff --git a/crates/chat2db-core/src/lib.rs b/crates/chat2db-core/src/lib.rs index 53d8935..ac25b45 100644 --- a/crates/chat2db-core/src/lib.rs +++ b/crates/chat2db-core/src/lib.rs @@ -18,6 +18,7 @@ pub mod mysql_ddl; mod mysql_schema_diff; mod mysql_workspace; mod native_driver; +mod native_driver_types; mod native_mysql; mod operation; mod query; @@ -661,6 +662,22 @@ impl Application { .map(str::to_owned) } + pub(crate) async fn require_native_driver_for_datasource( + &self, + datasource_id: &str, + ) -> Result, AppError> { + let storage = self.require_storage()?; + let resolved = + datasource_session::resolve_datasource_connection(&storage, datasource_id).await?; + self.native_driver_for_driver_id(&resolved.driver_id) + .ok_or_else(|| { + AppError::invalid( + "native_driver_not_available", + "The datasource does not have a native Rust driver", + ) + }) + } + async fn require_managed_driver_for_update( &self, storage: &Storage, diff --git a/crates/chat2db-core/src/mysql_workspace.rs b/crates/chat2db-core/src/mysql_workspace.rs index 53300d8..8d95e7b 100644 --- a/crates/chat2db-core/src/mysql_workspace.rs +++ b/crates/chat2db-core/src/mysql_workspace.rs @@ -3,7 +3,7 @@ use chat2db_contract::{ CommunityPinnedTableList, CommunityPinnedTableRequest, }; -use crate::{AppError, Application, native_mysql, storage_call}; +use crate::{AppError, Application, storage_call}; impl Application { /// Pins one `MySQL` table in the local workspace. @@ -82,13 +82,23 @@ impl Application { &self, request: CommunityErQueryRequest, ) -> Result { - let tables = native_mysql::load_er_tables( - self, - &request.data_source_id, - &request.database_name, - &request.schema_name, - ) - .await?; + let driver = self + .require_native_driver_for_datasource(&request.data_source_id) + .await?; + let table_driver = driver.tables().ok_or_else(|| { + AppError::invalid( + "native_table_capability_not_available", + "The native Rust driver does not implement table operations", + ) + })?; + let tables = table_driver + .load_er_tables( + self, + &request.data_source_id, + &request.database_name, + &request.schema_name, + ) + .await?; let storage = self.require_storage()?; let position = storage_call(move || { storage.mysql_er_position( diff --git a/crates/chat2db-core/src/native_driver.rs b/crates/chat2db-core/src/native_driver.rs index ef017e9..c2ddda8 100644 --- a/crates/chat2db-core/src/native_driver.rs +++ b/crates/chat2db-core/src/native_driver.rs @@ -9,6 +9,16 @@ use tokio_util::sync::CancellationToken; use crate::{ AppError, Application, datasource_session::ResolvedDatasourceConnection, + native_driver_types::{ + ColumnList, DatabaseList, EntityRelationTable, ForeignKeyList, FunctionList, + FunctionMetadata, FunctionParameterList, IndexList, ListColumnsRequest, + ListDatabasesRequest, ListIndexesRequest, ListRoutinesRequest, ListSchemasRequest, + ListTableKeysRequest, ListTablesRequest, ListTriggersRequest, ListViewsRequest, ObjectRef, + PrimaryKeyList, ProcedureList, ProcedureMetadata, ProcedureParameterList, + RoutineInvocationPreview, RoutineInvocationRequest, RoutineMigrationExecution, + RoutineMigrationRequest, SchemaList, TableList, TableMetadata, TablePreviewAccepted, + TablePreviewRequest, TriggerList, TriggerMetadata, ViewList, + }, native_mysql, operation::CancellationRequest, query::{ @@ -66,6 +76,176 @@ pub(crate) trait NativeQueryDriver: Send + Sync { ) -> Result, AppError>; } +/// Relational metadata operations exposed through the native driver SPI. +#[async_trait] +pub(crate) trait NativeMetadataDriver: Send + Sync { + async fn list_schemas( + &self, + application: &Application, + request: ListSchemasRequest, + ) -> Result; + + async fn list_databases( + &self, + application: &Application, + request: ListDatabasesRequest, + ) -> Result; + + async fn list_tables( + &self, + application: &Application, + request: ListTablesRequest, + ) -> Result; + + async fn list_columns( + &self, + application: &Application, + request: ListColumnsRequest, + ) -> Result; + + async fn list_indexes( + &self, + application: &Application, + request: ListIndexesRequest, + ) -> Result; + + async fn list_views( + &self, + application: &Application, + request: ListViewsRequest, + ) -> Result; + + async fn get_view( + &self, + application: &Application, + request: ObjectRef, + ) -> Result; + + async fn list_imported_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result; + + async fn list_exported_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result; + + async fn list_primary_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result; + + async fn list_functions( + &self, + application: &Application, + request: ListRoutinesRequest, + ) -> Result; + + async fn get_function( + &self, + application: &Application, + request: ObjectRef, + ) -> Result; + + async fn list_function_parameters( + &self, + application: &Application, + request: ObjectRef, + ) -> Result; + + async fn list_procedures( + &self, + application: &Application, + request: ListRoutinesRequest, + ) -> Result; + + async fn get_procedure( + &self, + application: &Application, + request: ObjectRef, + ) -> Result; + + async fn list_procedure_parameters( + &self, + application: &Application, + request: ObjectRef, + ) -> Result; + + async fn list_triggers( + &self, + application: &Application, + request: ListTriggersRequest, + ) -> Result; + + async fn get_trigger( + &self, + application: &Application, + request: ObjectRef, + ) -> Result; +} + +/// Table-specific native capabilities that are not plain metadata listings. +#[async_trait] +pub(crate) trait NativeTableDriver: Send + Sync { + async fn load_er_tables( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + ) -> Result, AppError>; + + async fn validate_column_reorder( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + table_name: &str, + column_names: &[String], + ) -> Result<(), AppError>; + + async fn table_ddl( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, + ) -> Result; + + async fn start_table_preview( + &self, + application: &Application, + request: TablePreviewRequest, + row_limit: u32, + ) -> Result; +} + +/// Stored-routine operations implemented by a native Rust driver. +#[async_trait] +pub(crate) trait NativeRoutineDriver: Send + Sync { + async fn preview_invocation( + &self, + application: &Application, + request: RoutineInvocationRequest, + ) -> Result; + + fn preview_migration( + &self, + request: RoutineMigrationRequest, + ) -> Result; + + async fn execute_migration( + &self, + application: &Application, + request: RoutineMigrationRequest, + ) -> Result; +} + /// Runtime-polymorphic native Rust database driver. /// /// Optional capability accessors allow a driver to participate only in the @@ -87,6 +267,18 @@ pub(crate) trait NativeDriver: Send + Sync { fn query(&self) -> Option<&dyn NativeQueryDriver> { None } + + fn metadata(&self) -> Option<&dyn NativeMetadataDriver> { + None + } + + fn tables(&self) -> Option<&dyn NativeTableDriver> { + None + } + + fn routines(&self) -> Option<&dyn NativeRoutineDriver> { + None + } } /// Immutable registry used to select native implementations at runtime. @@ -207,6 +399,18 @@ impl NativeDriver for MysqlNativeDriver { fn query(&self) -> Option<&dyn NativeQueryDriver> { Some(self) } + + fn metadata(&self) -> Option<&dyn NativeMetadataDriver> { + Some(self) + } + + fn tables(&self) -> Option<&dyn NativeTableDriver> { + Some(self) + } + + fn routines(&self) -> Option<&dyn NativeRoutineDriver> { + Some(self) + } } #[async_trait] @@ -273,6 +477,343 @@ impl NativeQueryDriver for MysqlNativeDriver { } } +#[async_trait] +impl NativeMetadataDriver for MysqlNativeDriver { + async fn list_schemas( + &self, + application: &Application, + request: ListSchemasRequest, + ) -> Result { + native_mysql::list_schemas(application, &request.datasource_id).await + } + + async fn list_databases( + &self, + application: &Application, + request: ListDatabasesRequest, + ) -> Result { + native_mysql::list_databases(application, &request.datasource_id).await + } + + async fn list_tables( + &self, + application: &Application, + request: ListTablesRequest, + ) -> Result { + native_mysql::list_tables( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.name_pattern, + ) + .await + } + + async fn list_columns( + &self, + application: &Application, + request: ListColumnsRequest, + ) -> Result { + native_mysql::list_columns( + application, + &request.table.scope.datasource_id, + &request.table.scope.database_name, + &request.table.scope.schema_name, + &request.table.table_name, + ) + .await + } + + async fn list_indexes( + &self, + application: &Application, + request: ListIndexesRequest, + ) -> Result { + native_mysql::list_indexes( + application, + &request.table.scope.datasource_id, + &request.table.scope.database_name, + &request.table.scope.schema_name, + &request.table.table_name, + ) + .await + } + + async fn list_views( + &self, + application: &Application, + request: ListViewsRequest, + ) -> Result { + native_mysql::list_views( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name_pattern, + ) + .await + } + + async fn get_view( + &self, + application: &Application, + request: ObjectRef, + ) -> Result { + native_mysql::get_view( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name, + ) + .await + } + + async fn list_imported_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + native_mysql::list_imported_keys( + application, + &request.table.scope.datasource_id, + &request.table.scope.database_name, + &request.table.table_name, + ) + .await + } + + async fn list_exported_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + native_mysql::list_exported_keys( + application, + &request.table.scope.datasource_id, + &request.table.scope.database_name, + &request.table.table_name, + ) + .await + } + + async fn list_primary_keys( + &self, + application: &Application, + request: ListTableKeysRequest, + ) -> Result { + native_mysql::list_primary_keys( + application, + &request.table.scope.datasource_id, + &request.table.scope.database_name, + &request.table.scope.schema_name, + &request.table.table_name, + ) + .await + } + + async fn list_functions( + &self, + application: &Application, + request: ListRoutinesRequest, + ) -> Result { + native_mysql::list_functions( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + ) + .await + } + + async fn get_function( + &self, + application: &Application, + request: ObjectRef, + ) -> Result { + native_mysql::get_function( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name, + ) + .await + } + + async fn list_function_parameters( + &self, + application: &Application, + request: ObjectRef, + ) -> Result { + native_mysql::list_function_parameters( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name, + ) + .await + } + + async fn list_procedures( + &self, + application: &Application, + request: ListRoutinesRequest, + ) -> Result { + native_mysql::list_procedures( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + ) + .await + } + + async fn get_procedure( + &self, + application: &Application, + request: ObjectRef, + ) -> Result { + native_mysql::get_procedure( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name, + ) + .await + } + + async fn list_procedure_parameters( + &self, + application: &Application, + request: ObjectRef, + ) -> Result { + native_mysql::list_procedure_parameters( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name, + ) + .await + } + + async fn list_triggers( + &self, + application: &Application, + request: ListTriggersRequest, + ) -> Result { + native_mysql::list_triggers( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + ) + .await + } + + async fn get_trigger( + &self, + application: &Application, + request: ObjectRef, + ) -> Result { + native_mysql::get_trigger( + application, + &request.scope.datasource_id, + &request.scope.database_name, + &request.scope.schema_name, + &request.name, + ) + .await + } +} + +#[async_trait] +impl NativeTableDriver for MysqlNativeDriver { + async fn load_er_tables( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + ) -> Result, AppError> { + native_mysql::load_er_tables(application, datasource_id, database_name, schema_name).await + } + + async fn validate_column_reorder( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + table_name: &str, + column_names: &[String], + ) -> Result<(), AppError> { + native_mysql::validate_column_reorder( + application, + datasource_id, + database_name, + table_name, + column_names, + ) + .await + } + + async fn table_ddl( + &self, + application: &Application, + datasource_id: &str, + database_name: &str, + schema_name: &str, + table_name: &str, + ) -> Result { + native_mysql::table_ddl( + application, + datasource_id, + database_name, + schema_name, + table_name, + ) + .await + } + + async fn start_table_preview( + &self, + application: &Application, + request: TablePreviewRequest, + row_limit: u32, + ) -> Result { + native_mysql::start_table_preview(application, request, row_limit).await + } +} + +#[async_trait] +impl NativeRoutineDriver for MysqlNativeDriver { + async fn preview_invocation( + &self, + application: &Application, + request: RoutineInvocationRequest, + ) -> Result { + native_mysql::preview_routine_invocation(application, request).await + } + + fn preview_migration( + &self, + request: RoutineMigrationRequest, + ) -> Result { + native_mysql::preview_routine_migration(&request) + } + + async fn execute_migration( + &self, + application: &Application, + request: RoutineMigrationRequest, + ) -> Result { + native_mysql::execute_routine_migration(application, request).await + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/crates/chat2db-core/src/native_driver_types.rs b/crates/chat2db-core/src/native_driver_types.rs new file mode 100644 index 0000000..874aef6 --- /dev/null +++ b/crates/chat2db-core/src/native_driver_types.rs @@ -0,0 +1,360 @@ +use chat2db_contract::{ + CommunityDatabaseList, CommunityErTable, CommunityForeignKeyList, CommunityFunction, + CommunityFunctionList, CommunityFunctionParameterList, CommunityPrimaryKeyList, + CommunityProcedure, CommunityProcedureList, CommunityProcedureParameterList, + CommunityRoutineInvocationPreview, CommunityRoutineMigrationExecution, + CommunityRoutineMigrationRequest, CommunitySchemaList, CommunityTable, + CommunityTableColumnList, CommunityTableIndexList, CommunityTableList, + CommunityTablePreviewAccepted, CommunityTrigger, CommunityTriggerList, CommunityViewList, + GetCommunityFunctionRequest, GetCommunityProcedureRequest, GetCommunityTriggerRequest, + ListCommunityColumnsRequest, ListCommunityDatabasesRequest, ListCommunityFunctionsRequest, + ListCommunityIndexesRequest, ListCommunityProceduresRequest, ListCommunitySchemasRequest, + ListCommunityTableKeysRequest, ListCommunityTablesRequest, ListCommunityTriggersRequest, + ListCommunityViewsRequest, PreviewCommunityRoutineInvocationRequest, + StartCommunityTablePreviewRequest, +}; + +pub(crate) type DatabaseList = CommunityDatabaseList; +pub(crate) type SchemaList = CommunitySchemaList; +pub(crate) type TableMetadata = CommunityTable; +pub(crate) type TableList = CommunityTableList; +pub(crate) type ColumnList = CommunityTableColumnList; +pub(crate) type IndexList = CommunityTableIndexList; +pub(crate) type ViewList = CommunityViewList; +pub(crate) type ForeignKeyList = CommunityForeignKeyList; +pub(crate) type PrimaryKeyList = CommunityPrimaryKeyList; +pub(crate) type FunctionMetadata = CommunityFunction; +pub(crate) type FunctionList = CommunityFunctionList; +pub(crate) type FunctionParameterList = CommunityFunctionParameterList; +pub(crate) type ProcedureMetadata = CommunityProcedure; +pub(crate) type ProcedureList = CommunityProcedureList; +pub(crate) type ProcedureParameterList = CommunityProcedureParameterList; +pub(crate) type TriggerMetadata = CommunityTrigger; +pub(crate) type TriggerList = CommunityTriggerList; +pub(crate) type EntityRelationTable = CommunityErTable; +pub(crate) type TablePreviewAccepted = CommunityTablePreviewAccepted; +pub(crate) type RoutineInvocationPreview = CommunityRoutineInvocationPreview; +pub(crate) type RoutineMigrationExecution = CommunityRoutineMigrationExecution; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct MetadataScope { + pub(crate) datasource_id: String, + pub(crate) database_name: String, + pub(crate) schema_name: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct TableRef { + pub(crate) scope: MetadataScope, + pub(crate) table_name: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ObjectRef { + pub(crate) scope: MetadataScope, + pub(crate) name: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListDatabasesRequest { + pub(crate) datasource_id: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListSchemasRequest { + pub(crate) datasource_id: String, + pub(crate) database_name: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListTablesRequest { + pub(crate) scope: MetadataScope, + pub(crate) name_pattern: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListColumnsRequest { + pub(crate) table: TableRef, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListIndexesRequest { + pub(crate) table: TableRef, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListViewsRequest { + pub(crate) scope: MetadataScope, + pub(crate) name_pattern: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListTableKeysRequest { + pub(crate) table: TableRef, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListRoutinesRequest { + pub(crate) scope: MetadataScope, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct ListTriggersRequest { + pub(crate) scope: MetadataScope, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct TablePreviewRequest { + pub(crate) table: TableRef, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct RoutineInvocationRequest { + pub(crate) scope: MetadataScope, + pub(crate) routine_type: String, + pub(crate) routine_name: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct RoutineMigrationRequest { + pub(crate) scope: MetadataScope, + pub(crate) database_type: String, + pub(crate) routine_type: String, + pub(crate) routine_name: String, + pub(crate) ddl: String, +} + +impl From for ListDatabasesRequest { + fn from(request: ListCommunityDatabasesRequest) -> Self { + Self { + datasource_id: request.datasource_id, + } + } +} + +impl From for ListSchemasRequest { + fn from(request: ListCommunitySchemasRequest) -> Self { + Self { + datasource_id: request.datasource_id, + database_name: request.database_name, + } + } +} + +impl From for ListTablesRequest { + fn from(request: ListCommunityTablesRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + name_pattern: request.table_name_pattern, + } + } +} + +impl From for ListColumnsRequest { + fn from(request: ListCommunityColumnsRequest) -> Self { + Self { + table: TableRef { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + table_name: request.table_name, + }, + } + } +} + +impl From for ListIndexesRequest { + fn from(request: ListCommunityIndexesRequest) -> Self { + Self { + table: TableRef { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + table_name: request.table_name, + }, + } + } +} + +impl From for ListViewsRequest { + fn from(request: ListCommunityViewsRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + name_pattern: request.view_name_pattern, + } + } +} + +impl From for ObjectRef { + fn from(request: ListCommunityViewsRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + name: request.view_name_pattern, + } + } +} + +impl From for ListTableKeysRequest { + fn from(request: ListCommunityTableKeysRequest) -> Self { + Self { + table: TableRef { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + table_name: request.table_name, + }, + } + } +} + +impl From for ListRoutinesRequest { + fn from(request: ListCommunityFunctionsRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + } + } +} + +impl From for ObjectRef { + fn from(request: GetCommunityFunctionRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + name: request.function_name, + } + } +} + +impl From for ListRoutinesRequest { + fn from(request: ListCommunityProceduresRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + } + } +} + +impl From for ObjectRef { + fn from(request: GetCommunityProcedureRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + name: request.procedure_name, + } + } +} + +impl From for ListTriggersRequest { + fn from(request: ListCommunityTriggersRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + } + } +} + +impl From for ObjectRef { + fn from(request: GetCommunityTriggerRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + name: request.trigger_name, + } + } +} + +impl From for TablePreviewRequest { + fn from(request: StartCommunityTablePreviewRequest) -> Self { + Self { + table: TableRef { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + table_name: request.table_name, + }, + } + } +} + +impl From for RoutineInvocationRequest { + fn from(request: PreviewCommunityRoutineInvocationRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + routine_type: request.routine_type, + routine_name: request.routine_name, + } + } +} + +impl From for RoutineMigrationRequest { + fn from(request: CommunityRoutineMigrationRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + database_type: request.database_type, + routine_type: request.routine_type, + routine_name: request.routine_name, + ddl: request.ddl, + } + } +} + +impl From for CommunityRoutineMigrationRequest { + fn from(request: RoutineMigrationRequest) -> Self { + Self { + datasource_id: request.scope.datasource_id, + database_type: request.database_type, + database_name: request.scope.database_name, + schema_name: request.scope.schema_name, + routine_type: request.routine_type, + routine_name: request.routine_name, + ddl: request.ddl, + } + } +} diff --git a/crates/chat2db-core/src/native_mysql.rs b/crates/chat2db-core/src/native_mysql.rs index c7a0e08..8b651a9 100644 --- a/crates/chat2db-core/src/native_mysql.rs +++ b/crates/chat2db-core/src/native_mysql.rs @@ -4,14 +4,11 @@ use chat2db_contract::{ CommunityErTable, CommunityForeignKey, CommunityForeignKeyList, CommunityFunction, CommunityFunctionList, CommunityFunctionParameter, CommunityFunctionParameterList, CommunityPrimaryKey, CommunityPrimaryKeyList, CommunityProcedure, CommunityProcedureList, - CommunityProcedureParameter, CommunityProcedureParameterList, - CommunityRoutineInvocationPreview, CommunityRoutineMigrationExecution, - CommunityRoutineMigrationRequest, CommunitySchemaList, CommunityTable, CommunityTableColumn, - CommunityTableColumnList, CommunityTableIndex, CommunityTableIndexColumn, - CommunityTableIndexList, CommunityTableList, CommunityTablePreviewAccepted, CommunityTrigger, + CommunityProcedureParameter, CommunityProcedureParameterList, CommunitySchemaList, + CommunityTable, CommunityTableColumn, CommunityTableColumnList, CommunityTableIndex, + CommunityTableIndexColumn, CommunityTableIndexList, CommunityTableList, CommunityTrigger, CommunityTriggerList, CommunityViewList, DatasourceConnection, JdbcValue, JdbcValueType, - PreviewCommunityRoutineInvocationRequest, QueryLimits, ResultColumn, ResultMetadata, ResultRow, - StartCommunityTablePreviewRequest, StartQueryRequest, + QueryLimits, ResultColumn, ResultMetadata, ResultRow, StartQueryRequest, }; use chat2db_engine_protocol::wire; use chat2db_java_bridge::{JdbcParameter, JdbcValue as BridgeJdbcValue, QueryOptions}; @@ -38,6 +35,10 @@ use url::Url; use crate::{ AppError, AppErrorKind, Application, datasource_session::{ResolvedDatasourceConnection, resolve_datasource_connection}, + native_driver_types::{ + RoutineInvocationPreview, RoutineInvocationRequest, RoutineMigrationExecution, + RoutineMigrationRequest, TablePreviewAccepted, TablePreviewRequest, + }, operation::CancellationRequest, query::{ DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, PreparedQuery, @@ -287,10 +288,6 @@ struct PreparedMysqlConnection { tunnel: Option, } -pub(crate) fn is_mysql_database_type(database_type: &str) -> bool { - database_type.trim().eq_ignore_ascii_case("mysql") -} - pub(crate) async fn test_connection(connection: &DatasourceConnection) -> Result<(), AppError> { test_connection_with_local_port(connection) .await @@ -1047,15 +1044,15 @@ pub(crate) async fn list_procedure_parameters( pub(crate) async fn preview_routine_invocation( application: &Application, - request: PreviewCommunityRoutineInvocationRequest, -) -> Result { + request: RoutineInvocationRequest, +) -> Result { let routine_type = normalize_mysql_routine_type(&request.routine_type)?; let routine_name = request.routine_name.trim().to_owned(); let routine_lookup_name = mysql_routine_lookup_name(&routine_name); - let database_name = request.database_name.trim().to_owned(); + let database_name = request.scope.database_name.trim().to_owned(); validate_metadata_identifier(&routine_name, "routineName")?; validate_metadata_identifier(&database_name, "databaseName")?; - let resolved = resolve_native_connection(application, &request.datasource_id).await?; + let resolved = resolve_native_connection(application, &request.scope.datasource_id).await?; let mut conn = open_resolved_connection(&resolved).await?; let query = "SELECT ORDINAL_POSITION, PARAMETER_MODE, PARAMETER_NAME, DATA_TYPE \ FROM information_schema.PARAMETERS \ @@ -1072,7 +1069,7 @@ pub(crate) async fn preview_routine_invocation( .filter_map(|row| routine_invocation_parameter(routine_type, row)) .collect::>(); parameters.sort_by_key(|parameter| parameter.ordinal_position); - CommunityRoutineInvocationPreview { + RoutineInvocationPreview { sql: render_routine_invocation_preview(routine_type, &routine_name, ¶meters), } }); @@ -1080,20 +1077,20 @@ pub(crate) async fn preview_routine_invocation( } pub(crate) fn preview_routine_migration( - request: &CommunityRoutineMigrationRequest, -) -> Result { + request: &RoutineMigrationRequest, +) -> Result { let plan = routine_migration_plan(request)?; - Ok(CommunityRoutineInvocationPreview { + Ok(RoutineInvocationPreview { sql: plan.preview_sql, }) } pub(crate) async fn execute_routine_migration( application: &Application, - request: CommunityRoutineMigrationRequest, -) -> Result { + request: RoutineMigrationRequest, +) -> Result { let plan = routine_migration_plan(&request)?; - let resolved = resolve_native_connection(application, &request.datasource_id).await?; + let resolved = resolve_native_connection(application, &request.scope.datasource_id).await?; let mut conn = open_resolved_connection(&resolved).await?; let result = execute_routine_migration_with_connection(&mut conn, &plan).await; finish_connection(conn, Ok(result)).await @@ -1102,7 +1099,7 @@ pub(crate) async fn execute_routine_migration( async fn execute_routine_migration_with_connection( conn: &mut Conn, plan: &RoutineMigrationPlan, -) -> CommunityRoutineMigrationExecution { +) -> RoutineMigrationExecution { let selected_database = quote_identifier(&plan.database_name, "databaseName") .expect("validated migration database name must remain valid"); if let Err(error) = metadata_query(conn.query_drop(format!("USE {selected_database}"))).await { @@ -1145,7 +1142,7 @@ async fn execute_routine_migration_with_connection( } match metadata_query(conn.query_drop(&plan.create_sql)).await { - Ok(()) => CommunityRoutineMigrationExecution { + Ok(()) => RoutineMigrationExecution { success: true, message: "Statement executed successfully".to_owned(), sql: plan.preview_sql.clone(), @@ -1234,10 +1231,10 @@ async fn capture_previous_routine( } fn routine_migration_plan( - request: &CommunityRoutineMigrationRequest, + request: &RoutineMigrationRequest, ) -> Result { let routine_type = normalize_mysql_routine_type(&request.routine_type)?; - let database_name = request.database_name.trim().to_owned(); + let database_name = request.scope.database_name.trim().to_owned(); let routine_name = mysql_routine_lookup_name(request.routine_name.trim()); validate_metadata_identifier(&database_name, "databaseName")?; validate_metadata_identifier(&routine_name, "routineName")?; @@ -1276,8 +1273,8 @@ fn routine_migration_failure( failure_stage: &str, restore_attempted: bool, restore_succeeded: bool, -) -> CommunityRoutineMigrationExecution { - CommunityRoutineMigrationExecution { +) -> RoutineMigrationExecution { + RoutineMigrationExecution { success: false, message, sql: plan.preview_sql.clone(), @@ -1744,15 +1741,15 @@ fn invalid_sql_lexeme(detail: &str) -> AppError { pub(crate) async fn start_table_preview( application: &Application, - request: StartCommunityTablePreviewRequest, + request: TablePreviewRequest, row_limit: u32, -) -> Result { - let database_name = quote_identifier(&request.database_name, "databaseName")?; - let table_name = quote_identifier(&request.table_name, "tableName")?; +) -> Result { + let database_name = quote_identifier(&request.table.scope.database_name, "databaseName")?; + let table_name = quote_identifier(&request.table.table_name, "tableName")?; let sql = format!("SELECT * FROM {database_name}.{table_name} LIMIT {row_limit}"); let accepted = application .start_read_query(StartQueryRequest { - datasource_id: request.datasource_id, + datasource_id: request.table.scope.datasource_id, sql: sql.clone(), parameters: Vec::new(), limits: QueryLimits { @@ -1764,7 +1761,7 @@ pub(crate) async fn start_table_preview( }, }) .await?; - Ok(CommunityTablePreviewAccepted { + Ok(TablePreviewAccepted { operation_id: accepted.operation_id, sql, row_limit, @@ -4839,8 +4836,7 @@ fn mysql_query_error(error: MysqlError) -> AppError { #[cfg(test)] mod tests { use chat2db_contract::{ - CommunityRoutineMigrationRequest, DatasourceConnection, DatasourceConnectionProperty, - JdbcValue, ResultRow, + DatasourceConnection, DatasourceConnectionProperty, JdbcValue, ResultRow, }; use mysql_async::{Conn, Opts}; use tokio::sync::watch; @@ -4849,8 +4845,8 @@ mod tests { ColumnRow, ConsoleExecutionError, ConsoleStatementExecution, MAX_CONSOLE_PAGE_SIZE, MAX_CONSOLE_RESULT_BYTES, community_column, community_foreign_key, community_function_parameter, community_indexes, community_procedure_parameter, - connection_opts, execute_console_statement, is_mysql_database_type, - is_native_read_candidate, mysql_column_reorder_hazard, mysql_identifier_is_backtick_quoted, + connection_opts, execute_console_statement, is_native_read_candidate, + mysql_column_reorder_hazard, mysql_identifier_is_backtick_quoted, mysql_metadata_column_type, mysql_routine_default_value, mysql_routine_invocation_name, mysql_routine_lookup_name, normalize_mysql_routine_type, normalize_table_type, open_connection_with_opts, qualified_identifier, quote_identifier, @@ -4860,6 +4856,7 @@ mod tests { validate_read_sql, validate_single_write_sql, }; use super::{MysqlRoutineType, RoutineInvocationParameter}; + use crate::native_driver_types::{MetadataScope, RoutineMigrationRequest}; use crate::{MysqlConsoleRequest, operation::CancellationRequest}; #[test] @@ -5208,9 +5205,7 @@ mod tests { } #[test] - fn mysql_detection_and_table_types_are_closed() { - assert!(is_mysql_database_type(" mysql ")); - assert!(!is_mysql_database_type("mariadb")); + fn mysql_table_types_are_closed() { assert_eq!(normalize_table_type("VIEW"), "VIEW"); assert_eq!(normalize_table_type("BASE TABLE"), "TABLE"); } @@ -5344,11 +5339,13 @@ mod tests { #[test] fn mysql_routine_migration_preview_is_qualified_and_terminated() { - let plan = routine_migration_plan(&CommunityRoutineMigrationRequest { - datasource_id: "mysql-local".to_owned(), + let plan = routine_migration_plan(&RoutineMigrationRequest { + scope: MetadataScope { + datasource_id: "mysql-local".to_owned(), + database_name: "inventory".to_owned(), + schema_name: String::new(), + }, database_type: "MYSQL".to_owned(), - database_name: "inventory".to_owned(), - schema_name: String::new(), routine_type: " function ".to_owned(), routine_name: "`odd``name`".to_owned(), ddl: "CREATE FUNCTION `odd``name`() RETURNS INT RETURN 2".to_owned(), @@ -5364,11 +5361,13 @@ mod tests { #[test] fn mysql_routine_migration_rejects_missing_ddl() { - let error = routine_migration_plan(&CommunityRoutineMigrationRequest { - datasource_id: "mysql-local".to_owned(), + let error = routine_migration_plan(&RoutineMigrationRequest { + scope: MetadataScope { + datasource_id: "mysql-local".to_owned(), + database_name: "inventory".to_owned(), + schema_name: String::new(), + }, database_type: "MYSQL".to_owned(), - database_name: "inventory".to_owned(), - schema_name: String::new(), routine_type: "PROCEDURE".to_owned(), routine_name: "refresh_items".to_owned(), ddl: " ".to_owned(), diff --git a/crates/chat2db-core/src/transfer/class_generation.rs b/crates/chat2db-core/src/transfer/class_generation.rs index f54d5cf..ff6da5f 100644 --- a/crates/chat2db-core/src/transfer/class_generation.rs +++ b/crates/chat2db-core/src/transfer/class_generation.rs @@ -10,7 +10,11 @@ use chat2db_storage::TransferArtifactRecord; use uuid::Uuid; use zip::{CompressionMethod, ZipWriter, write::SimpleFileOptions}; -use crate::{AppError, Application, native_mysql, now_millis}; +use crate::{ + AppError, Application, + native_driver_types::{ListColumnsRequest, MetadataScope, TableRef}, + now_millis, +}; const CLASS_ARCHIVE_TTL_MS: i64 = 24 * 60 * 60 * 1_000; const MAX_CLASS_ARCHIVE_BYTES: u64 = 16 * 1024 * 1024; @@ -79,15 +83,31 @@ async fn render_request( application: &Application, request: &GenerateMysqlClassRequest, ) -> Result { - let columns = native_mysql::list_columns( - application, - &request.datasource_id, - &request.database_name, - &request.schema_name, - &request.table_name, - ) - .await? - .items; + let driver = application + .require_native_driver_for_datasource(&request.datasource_id) + .await?; + let metadata = driver.metadata().ok_or_else(|| { + AppError::invalid( + "native_metadata_capability_not_available", + "The native Rust driver does not implement metadata operations", + ) + })?; + let columns = metadata + .list_columns( + application, + ListColumnsRequest { + table: TableRef { + scope: MetadataScope { + datasource_id: request.datasource_id.clone(), + database_name: request.database_name.clone(), + schema_name: request.schema_name.clone(), + }, + table_name: request.table_name.clone(), + }, + }, + ) + .await? + .items; if columns.is_empty() { return Err(AppError::not_found( "mysql_table_not_found", From 2ac51626b67db295514cf0e6c8398c4de69470ed Mon Sep 17 00:00:00 2001 From: zgq Date: Wed, 5 Aug 2026 18:47:20 +0800 Subject: [PATCH 3/5] refactor(core): complete native driver capability spi --- crates/chat2db-core/src/community.rs | 215 ++- crates/chat2db-core/src/lib.rs | 33 +- crates/chat2db-core/src/mysql_account.rs | 847 +++++++----- crates/chat2db-core/src/mysql_dashboard.rs | 84 +- crates/chat2db-core/src/mysql_ddl.rs | 1212 ++++++++++++++++- crates/chat2db-core/src/mysql_schema_diff.rs | 89 +- .../src/native_administration_types.rs | 222 +++ crates/chat2db-core/src/native_api_adapter.rs | 381 ++++++ crates/chat2db-core/src/native_driver.rs | 281 +++- .../chat2db-core/src/native_driver_types.rs | 152 ++- .../src/native_schema_diff_types.rs | 28 + crates/chat2db-core/src/query.rs | 9 - crates/chat2db-core/src/transfer/mod.rs | 1047 +++++++++++--- crates/chat2db-core/src/transfer/mysql.rs | 23 +- .../chat2db-core/src/transfer/mysql_impl.rs | 100 ++ 15 files changed, 4070 insertions(+), 653 deletions(-) create mode 100644 crates/chat2db-core/src/native_administration_types.rs create mode 100644 crates/chat2db-core/src/native_api_adapter.rs create mode 100644 crates/chat2db-core/src/native_schema_diff_types.rs create mode 100644 crates/chat2db-core/src/transfer/mysql_impl.rs diff --git a/crates/chat2db-core/src/community.rs b/crates/chat2db-core/src/community.rs index 411c739..07ea352 100644 --- a/crates/chat2db-core/src/community.rs +++ b/crates/chat2db-core/src/community.rs @@ -68,6 +68,12 @@ use chat2db_java_bridge::{ }; use chat2db_storage::Storage; +use crate::native_driver_types::{ + CreateSchemaSqlRequest, DatabaseDefinition, DmlAssignment, DmlColumn, DmlRow, DmlSqlRequest, + DmlStatement, DmlTarget, DmlTemporalKind, DmlValue, NamespaceSqlOperation, NamespaceSqlRequest, + SchemaDefinition, +}; + const FIXED_COMMUNITY_CLASSPATH_LOCK: &str = include_str!("../../../third_party/community-h2-classpath.lock"); const DEFAULT_TABLE_PREVIEW_ROWS: u32 = 200; @@ -1077,6 +1083,15 @@ impl Application { &self, request: BuildCommunityCreateSchemaRequest, ) -> Result { + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(dialect) = driver.dialect() + { + return dialect + .build_create_schema(CreateSchemaSqlRequest { + schema: native_schema(request.schema), + }) + .map(|built| CommunityBuiltSql { sql: built.sql }); + } let engine = self.require_community_engine().await?; let client = engine.community_client().map_err(AppError::from)?; client @@ -1096,6 +1111,13 @@ impl Application { &self, request: BuildCommunityNamespaceSqlRequest, ) -> Result { + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(dialect) = driver.dialect() + { + return dialect + .build_namespace_sql(native_namespace_request(request)) + .map(|built| CommunityBuiltSql { sql: built.sql }); + } let engine = self.require_community_engine().await?; let client = engine.community_client().map_err(AppError::from)?; client @@ -1115,6 +1137,13 @@ impl Application { &self, request: BuildCommunityDmlRequest, ) -> Result { + if let Some(driver) = self.native_driver_for_database_type(&request.database_type) + && let Some(dialect) = driver.dialect() + { + return dialect + .build_dml(native_dml_request(request)?) + .map(|built| CommunityBuiltSql { sql: built.sql }); + } let engine = self.require_community_engine().await?; let client = engine.community_client().map_err(AppError::from)?; client @@ -1718,6 +1747,169 @@ fn bridge_database(database: CommunityDatabase) -> BridgeCommunityDatabase { } } +fn native_schema(schema: CommunitySchema) -> SchemaDefinition { + SchemaDefinition { + database_name: schema.database_name, + name: schema.name, + comment: schema.comment, + owner: schema.owner, + system: schema.system, + } +} + +fn native_database(database: CommunityDatabase) -> DatabaseDefinition { + DatabaseDefinition { + name: database.name, + comment: database.comment, + charset: database.charset, + collation: database.collation, + owner: database.owner, + system: database.system, + } +} + +fn native_namespace_request(request: BuildCommunityNamespaceSqlRequest) -> NamespaceSqlRequest { + NamespaceSqlRequest { + operation: match request.operation { + CommunityNamespaceSqlOperation::CreateDatabase { database } => { + NamespaceSqlOperation::CreateDatabase { + database: native_database(database), + } + } + CommunityNamespaceSqlOperation::AlterDatabase { + old_database, + new_database, + } => NamespaceSqlOperation::AlterDatabase { + old_database: native_database(old_database), + new_database: native_database(new_database), + }, + CommunityNamespaceSqlOperation::DropDatabase { database_name } => { + NamespaceSqlOperation::DropDatabase { database_name } + } + CommunityNamespaceSqlOperation::UseDatabase { database_name } => { + NamespaceSqlOperation::UseDatabase { database_name } + } + CommunityNamespaceSqlOperation::CreateSchema { schema } => { + NamespaceSqlOperation::CreateSchema { + schema: native_schema(schema), + } + } + CommunityNamespaceSqlOperation::AlterSchema { + old_schema_name, + new_schema_name, + } => NamespaceSqlOperation::AlterSchema { + old_schema_name, + new_schema_name, + }, + CommunityNamespaceSqlOperation::DropSchema { schema_name } => { + NamespaceSqlOperation::DropSchema { schema_name } + } + }, + } +} + +fn native_dml_request(request: BuildCommunityDmlRequest) -> Result { + Ok(DmlSqlRequest { + target: DmlTarget { + database_name: request.target.database_name, + schema_name: request.target.schema_name, + table_name: request.target.table_name, + }, + statement: native_dml_statement(request.statement)?, + }) +} + +fn native_dml_statement(statement: CommunityDmlStatement) -> Result { + Ok(match statement { + CommunityDmlStatement::SingleInsert { columns, row } => DmlStatement::SingleInsert { + columns: columns.into_iter().map(native_dml_column).collect(), + row: native_dml_row(row)?, + }, + CommunityDmlStatement::MultiInsert { columns, rows } => DmlStatement::MultiInsert { + columns: columns.into_iter().map(native_dml_column).collect(), + rows: rows + .into_iter() + .map(native_dml_row) + .collect::>()?, + }, + CommunityDmlStatement::Update { + assignments, + predicates, + } => DmlStatement::Update { + assignments: assignments + .into_iter() + .map(native_dml_assignment) + .collect::>()?, + predicates: predicates + .into_iter() + .map(native_dml_assignment) + .collect::>()?, + }, + }) +} + +fn native_dml_column(column: CommunityDmlColumn) -> DmlColumn { + DmlColumn { + name: column.name, + data_type_name: column.data_type_name, + precision: column.precision, + scale: column.scale, + } +} + +fn native_dml_row(row: CommunityDmlRow) -> Result { + Ok(DmlRow { + values: row + .values + .into_iter() + .map(native_dml_value) + .collect::>()?, + }) +} + +fn native_dml_assignment(assignment: CommunityDmlAssignment) -> Result { + Ok(DmlAssignment { + column: native_dml_column(assignment.column), + value: native_dml_value(assignment.value)?, + }) +} + +fn native_dml_value(value: CommunityDmlValue) -> Result { + Ok(match value { + CommunityDmlValue::Null => DmlValue::Null, + CommunityDmlValue::String { value } => DmlValue::String(value), + CommunityDmlValue::Decimal { value } => DmlValue::Decimal(value), + CommunityDmlValue::Boolean { value } => DmlValue::Boolean(value), + CommunityDmlValue::Temporal { + temporal_kind, + value, + } => DmlValue::Temporal { + kind: match temporal_kind { + CommunityDmlTemporalKind::Date => DmlTemporalKind::Date, + CommunityDmlTemporalKind::Time => DmlTemporalKind::Time, + CommunityDmlTemporalKind::LocalDatetime => DmlTemporalKind::LocalDatetime, + CommunityDmlTemporalKind::OffsetDatetime => DmlTemporalKind::OffsetDatetime, + }, + iso8601: value, + }, + CommunityDmlValue::Binary { base64 } => { + let bytes = STANDARD.decode(&base64).map_err(|_| { + AppError::invalid( + "community.dml_invalid_value", + "Community DML binary values must use canonical standard base64", + ) + })?; + if STANDARD.encode(&bytes) != base64 { + return Err(AppError::invalid( + "community.dml_invalid_value", + "Community DML binary values must use canonical standard base64", + )); + } + DmlValue::Binary(bytes) + } + }) +} + fn bridge_namespace_request( request: BuildCommunityNamespaceSqlRequest, ) -> BridgeBuildCommunityNamespaceSqlRequest { @@ -2151,8 +2343,9 @@ mod tests { community_function_parameter, community_plugin_catalog, community_primary_key, community_procedure, community_procedure_parameter, community_schema, community_sql_analysis, community_sql_validation, community_table, community_table_column, - community_table_index, community_trigger, preserve_primary_result, run_cancellation_safe, - run_cancellation_safe_with_cleanup, table_preview_row_limit, validate_table_preview_sql, + community_table_index, community_trigger, native_dml_request, preserve_primary_result, + run_cancellation_safe, run_cancellation_safe_with_cleanup, table_preview_row_limit, + validate_table_preview_sql, }; use crate::{AppError, AppErrorKind}; @@ -2265,11 +2458,29 @@ mod tests { ])] ); + let native = native_dml_request(binary_dml_request("AAH/")) + .expect("the native adapter must accept the same canonical bytes"); + let crate::native_driver_types::DmlStatement::SingleInsert { row, .. } = native.statement + else { + panic!("single insert must retain its native variant"); + }; + assert_eq!( + row.values, + vec![crate::native_driver_types::DmlValue::Binary(vec![ + 0, 1, 255 + ])] + ); + for invalid in ["AAH_", "AAH/==", "not base64"] { let error = bridge_dml_request(binary_dml_request(invalid)) .expect_err("noncanonical binary JSON must fail before engine access"); assert_eq!(error.api_error().code, "community.dml_invalid_value"); assert_eq!(error.kind(), AppErrorKind::InvalidRequest); + + let native_error = native_dml_request(binary_dml_request(invalid)) + .expect_err("the native adapter must reject the same noncanonical input"); + assert_eq!(native_error.api_error().code, "community.dml_invalid_value"); + assert_eq!(native_error.kind(), AppErrorKind::InvalidRequest); } } diff --git a/crates/chat2db-core/src/lib.rs b/crates/chat2db-core/src/lib.rs index ac25b45..8593791 100644 --- a/crates/chat2db-core/src/lib.rs +++ b/crates/chat2db-core/src/lib.rs @@ -17,9 +17,12 @@ mod mysql_dashboard; pub mod mysql_ddl; mod mysql_schema_diff; mod mysql_workspace; +mod native_administration_types; +mod native_api_adapter; mod native_driver; mod native_driver_types; mod native_mysql; +mod native_schema_diff_types; mod operation; mod query; mod ssh; @@ -250,6 +253,22 @@ impl Application { storage: Option, engine: EngineProvider, drivers: Option>, + ) -> Self { + Self::compose_with_native_drivers( + runtime_status, + storage, + engine, + drivers, + native_driver::NativeDriverRegistry::built_in(), + ) + } + + fn compose_with_native_drivers( + runtime_status: RuntimeStatus, + storage: Option, + engine: EngineProvider, + drivers: Option>, + native_drivers: native_driver::NativeDriverRegistry, ) -> Self { let managed_driver_ids = drivers.as_ref().map(|drivers| { drivers @@ -258,7 +277,6 @@ impl Application { .collect() }); let drivers = drivers.unwrap_or_default(); - let native_drivers = native_driver::NativeDriverRegistry::built_in(); Self { inner: Arc::new(ApplicationInner { started_at: Instant::now(), @@ -280,6 +298,19 @@ impl Application { } } + #[cfg(test)] + pub(crate) fn with_native_drivers_for_test( + native_drivers: native_driver::NativeDriverRegistry, + ) -> Self { + Self::compose_with_native_drivers( + RuntimeStatus::Ready, + None, + EngineProvider::Disabled, + None, + native_drivers, + ) + } + /// Returns local storage only when runtime composition initialized it. #[must_use] pub fn storage(&self) -> Option<&Storage> { diff --git a/crates/chat2db-core/src/mysql_account.rs b/crates/chat2db-core/src/mysql_account.rs index eb69d20..4a1f1bf 100644 --- a/crates/chat2db-core/src/mysql_account.rs +++ b/crates/chat2db-core/src/mysql_account.rs @@ -5,18 +5,18 @@ use std::{ time::{Duration, Instant}, }; -use chat2db_contract::{ - ApiError, CommunityAccount, CommunityAccountAction, CommunityAccountCapability, - CommunityAccountCommandRequest, CommunityAccountExecution, CommunityAccountGrantList, - CommunityAccountGrantsRequest, CommunityAccountList, CommunityAccountPreview, - CommunityAccountPrivilegeScope, CommunityMysqlPrivilege, DatasourceConnection, -}; +use chat2db_contract::{ApiError, DatasourceConnection}; use mysql_async::{Conn, Error as MysqlError, prelude::Queryable}; use sha2::{Digest, Sha256}; use uuid::Uuid; use crate::{ AppError, AppErrorKind, Application, + native_administration_types::{ + AdministrationAction, AdministrationCapability, AdministrationCommand, + AdministrationExecution, AdministrationPreview, Principal, PrincipalGrantList, + PrincipalGrantsRequest, PrincipalList, PrincipalRef, PrivilegeScope, + }, native_mysql::{finish_connection, open_resolved_connection, resolve_native_connection}, }; @@ -38,6 +38,62 @@ const SELECT_ACCOUNTS: &str = "SELECT User, Host, plugin FROM mysql.user ORDER B const SELECT_ACCOUNTS_WITH_LOCK: &str = "SELECT User, Host, plugin, account_locked FROM mysql.user ORDER BY User, Host"; +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +enum MysqlPrivilege { + Select, + Insert, + Update, + Delete, + Create, + Drop, + Alter, + Index, + References, + Execute, + ShowView, + Trigger, + Event, + CreateTemporaryTables, +} + +impl MysqlPrivilege { + const ALL: [Self; 14] = [ + Self::Select, + Self::Insert, + Self::Update, + Self::Delete, + Self::Create, + Self::Drop, + Self::Alter, + Self::Index, + Self::References, + Self::Execute, + Self::ShowView, + Self::Trigger, + Self::Event, + Self::CreateTemporaryTables, + ]; + + const fn wire_name(self) -> &'static str { + match self { + Self::Select => "SELECT", + Self::Insert => "INSERT", + Self::Update => "UPDATE", + Self::Delete => "DELETE", + Self::Create => "CREATE", + Self::Drop => "DROP", + Self::Alter => "ALTER", + Self::Index => "INDEX", + Self::References => "REFERENCES", + Self::Execute => "EXECUTE", + Self::ShowView => "SHOW_VIEW", + Self::Trigger => "TRIGGER", + Self::Event => "EVENT", + Self::CreateTemporaryTables => "CREATE_TEMPORARY_TABLES", + } + } +} + enum AccountQueryFailure { Timeout, Mysql(MysqlError), @@ -105,60 +161,72 @@ impl AccountPreviewRegistry { && binding.datasource_id == datasource_id && binding.sql_sha256 == sha256(sql.as_bytes()) } -} -impl Application { - /// Returns `MySQL` account-administration capability for one datasource. - /// - /// # Errors - /// - /// Returns datasource, secret, driver, connection, or cleanup errors. - pub async fn mysql_account_capability( - &self, - datasource_id: &str, - ) -> Result { - let resolved = resolve_native_connection(self, datasource_id).await?; - let connection_user = configured_connection_user(&resolved.connection); - let mut conn = open_resolved_connection(&resolved).await?; - - let account_list_readable = match timed_query(conn.query_drop(PROBE_ACCOUNT_LIST)).await { - Ok(()) => true, - Err(AccountQueryFailure::Mysql(_)) => false, - Err(AccountQueryFailure::Timeout) => { - return finish_connection( - conn, - Ok(capability_with_message( - connection_user, - false, - false, - "The MySQL account capability query timed out", - )), - ) - .await; - } + fn authorizes(&self, token: &str, datasource_id: &str, sql: &str) -> bool { + if token.len() != 64 || !token.bytes().all(|byte| byte.is_ascii_hexdigit()) { + return false; + } + let token_sha256 = sha256(token.as_bytes()); + let Ok(pending) = self.pending.lock() else { + return false; }; - let account_lock_supported = match timed_query(conn.query_drop(PROBE_ACCOUNT_LOCK)).await { - Ok(()) => true, - Err(AccountQueryFailure::Mysql(_)) => false, - Err(AccountQueryFailure::Timeout) => { - return finish_connection( - conn, - Ok(capability_with_message( - connection_user, - account_list_readable, - false, - "The MySQL account capability query timed out", - )), - ) - .await; - } + let Some(binding) = pending.get(&token_sha256) else { + return false; }; + binding.expires_at > Instant::now() + && binding.datasource_id == datasource_id + && binding.sql_sha256 == sha256(sql.as_bytes()) + } +} - let (product_version, current_user, message) = match timed_query( - conn.query_first::<(String, String), _>(SELECT_CURRENT_ACCOUNT), - ) - .await - { +/// Returns `MySQL` account-administration capability for one datasource. +/// +/// # Errors +/// +/// Returns datasource, secret, driver, connection, or cleanup errors. +pub(crate) async fn mysql_account_capability( + application: &Application, + datasource_id: &str, +) -> Result { + let resolved = resolve_native_connection(application, datasource_id).await?; + let connection_user = configured_connection_user(&resolved.connection); + let mut conn = open_resolved_connection(&resolved).await?; + + let account_list_readable = match timed_query(conn.query_drop(PROBE_ACCOUNT_LIST)).await { + Ok(()) => true, + Err(AccountQueryFailure::Mysql(_)) => false, + Err(AccountQueryFailure::Timeout) => { + return finish_connection( + conn, + Ok(capability_with_message( + connection_user, + false, + false, + "The MySQL account capability query timed out", + )), + ) + .await; + } + }; + let account_lock_supported = match timed_query(conn.query_drop(PROBE_ACCOUNT_LOCK)).await { + Ok(()) => true, + Err(AccountQueryFailure::Mysql(_)) => false, + Err(AccountQueryFailure::Timeout) => { + return finish_connection( + conn, + Ok(capability_with_message( + connection_user, + account_list_readable, + false, + "The MySQL account capability query timed out", + )), + ) + .await; + } + }; + + let (product_version, current_user, message) = + match timed_query(conn.query_first::<(String, String), _>(SELECT_CURRENT_ACCOUNT)).await { Ok(Some((version, current_user))) => (Some(version), Some(current_user), None), Ok(None) => (None, None, None), Err(AccountQueryFailure::Timeout) => ( @@ -170,208 +238,190 @@ impl Application { (None, None, Some(safe_query_message(&error))) } }; - finish_connection( - conn, - Ok(CommunityAccountCapability { - db_type: "MYSQL".to_owned(), - product_name: "MySQL".to_owned(), - product_version, - current_user, - connection_user, - account_list_readable, - account_lock_supported, - editable_privileges: CommunityMysqlPrivilege::ALL - .into_iter() - .map(|privilege| privilege.wire_name().to_owned()) - .collect(), - message, - }), - ) - .await - } + finish_connection( + conn, + Ok(AdministrationCapability { + database_type: "MYSQL".to_owned(), + product_name: "MySQL".to_owned(), + product_version, + current_principal: current_user, + connection_principal: connection_user, + principal_list_readable: account_list_readable, + principal_lock_supported: account_lock_supported, + editable_privileges: MysqlPrivilege::ALL + .into_iter() + .map(|privilege| privilege.wire_name().to_owned()) + .collect(), + message, + }), + ) + .await +} - /// Lists `MySQL` accounts in stable user and host order. - /// - /// # Errors - /// - /// Returns datasource, connection, permission, query, or cleanup errors. - pub async fn list_mysql_accounts( - &self, - datasource_id: &str, - ) -> Result { - let resolved = resolve_native_connection(self, datasource_id).await?; - let mut conn = open_resolved_connection(&resolved).await?; - let with_lock = timed_query( - conn.query::<(String, String, Option, Option), _>( - SELECT_ACCOUNTS_WITH_LOCK, - ), - ) - .await; - let result = match with_lock { - Ok(rows) => Ok(CommunityAccountList { - items: rows - .into_iter() - .map(|(user, host, plugin, locked)| { - account(user, host, plugin, locked.as_deref()) - }) - .collect(), - }), - Err(AccountQueryFailure::Timeout) => Err(account_query_unavailable( - ACCOUNT_LIST_UNAVAILABLE, - "The MySQL account list query timed out", - )), - Err(AccountQueryFailure::Mysql(_)) => { - match timed_query( - conn.query::<(String, String, Option), _>(SELECT_ACCOUNTS), - ) +/// Lists `MySQL` accounts in stable user and host order. +/// +/// # Errors +/// +/// Returns datasource, connection, permission, query, or cleanup errors. +pub(crate) async fn list_mysql_accounts( + application: &Application, + datasource_id: &str, +) -> Result { + let resolved = resolve_native_connection(application, datasource_id).await?; + let mut conn = open_resolved_connection(&resolved).await?; + let with_lock = timed_query( + conn.query::<(String, String, Option, Option), _>( + SELECT_ACCOUNTS_WITH_LOCK, + ), + ) + .await; + let result = match with_lock { + Ok(rows) => Ok(PrincipalList { + items: rows + .into_iter() + .map(|(user, host, plugin, locked)| account(user, host, plugin, locked.as_deref())) + .collect(), + }), + Err(AccountQueryFailure::Timeout) => Err(account_query_unavailable( + ACCOUNT_LIST_UNAVAILABLE, + "The MySQL account list query timed out", + )), + Err(AccountQueryFailure::Mysql(_)) => { + match timed_query(conn.query::<(String, String, Option), _>(SELECT_ACCOUNTS)) .await - { - Ok(rows) => Ok(CommunityAccountList { - items: rows - .into_iter() - .map(|(user, host, plugin)| account(user, host, plugin, None)) - .collect(), - }), - Err(_) => Err(account_query_unavailable( - ACCOUNT_LIST_UNAVAILABLE, - "The MySQL account list is unavailable", - )), - } + { + Ok(rows) => Ok(PrincipalList { + items: rows + .into_iter() + .map(|(user, host, plugin)| account(user, host, plugin, None)) + .collect(), + }), + Err(_) => Err(account_query_unavailable( + ACCOUNT_LIST_UNAVAILABLE, + "The MySQL account list is unavailable", + )), } - }; - finish_connection(conn, result).await - } - - /// Returns `SHOW GRANTS` rows for one `MySQL` account. - /// - /// # Errors - /// - /// Returns validation, datasource, connection, permission, query, or cleanup errors. - pub async fn mysql_account_grants( - &self, - request: &CommunityAccountGrantsRequest, - ) -> Result { - let account = account_literal(&request.user, &request.host)?; - let resolved = resolve_native_connection(self, &request.datasource_id).await?; - let mut conn = open_resolved_connection(&resolved).await?; - let sql = format!("SHOW GRANTS FOR {account}"); - let result = match timed_query(query_account_grants(&mut conn, &sql)).await { - Ok(items) => Ok(CommunityAccountGrantList { items }), - Err(_) => Err(account_query_unavailable( - ACCOUNT_GRANTS_UNAVAILABLE, - "The MySQL grants are unavailable", - )), - }; - finish_connection(conn, result).await - } - - /// Builds a masked account-operation preview without opening `MySQL` or starting Java. - /// - /// # Errors - /// - /// Returns a field-specific account validation error. - pub fn preview_mysql_account( - &self, - request: &CommunityAccountCommandRequest, - ) -> Result { - preview_account(&self.inner.account_previews, request) - } - - /// Executes one preview-authorized `MySQL` account operation through `mysql_async`. - /// - /// SQL execution errors are returned in [`CommunityAccountExecution`]. Datasource, - /// validation, preview-token, connection, read-only, and cleanup failures remain errors. - /// - /// # Errors - /// - /// Returns validation, token, datasource, connection, read-only, or cleanup errors. - pub async fn execute_mysql_account( - &self, - request: &CommunityAccountCommandRequest, - ) -> Result { - let execution_sql = build_account_sql(request, false)?; - let preview_sql = build_account_sql(request, true)?; - let supplied_token = request.preview_token.as_deref().unwrap_or_default(); - if !self.inner.account_previews.consume( - supplied_token, - &request.datasource_id, - &execution_sql, - ) { - return Err(AppError::new( - AppErrorKind::Conflict, - ApiError::new( - ACCOUNT_PREVIEW_TOKEN_MISMATCH, - "The MySQL account preview token does not match this operation", - ), - )); } + }; + finish_connection(conn, result).await +} - let resolved = resolve_native_connection(self, &request.datasource_id).await?; - if resolved.connection.read_only { - return Err(AppError::new( - AppErrorKind::Conflict, - ApiError::new( - "datasource_read_only", - "The datasource connection is configured as read-only", - ), - )); - } - let mut conn = open_resolved_connection(&resolved).await?; - let query_result = timed_query(execute_account_sql(&mut conn, &execution_sql)).await; - drop(execution_sql); +/// Returns `SHOW GRANTS` rows for one `MySQL` account. +/// +/// # Errors +/// +/// Returns validation, datasource, connection, permission, query, or cleanup errors. +pub(crate) async fn mysql_account_grants( + application: &Application, + request: &PrincipalGrantsRequest, +) -> Result { + let account = account_literal(&request.principal)?; + let resolved = resolve_native_connection(application, &request.datasource_id).await?; + let mut conn = open_resolved_connection(&resolved).await?; + let sql = format!("SHOW GRANTS FOR {account}"); + let result = match timed_query(query_account_grants(&mut conn, &sql)).await { + Ok(items) => Ok(PrincipalGrantList { items }), + Err(_) => Err(account_query_unavailable( + ACCOUNT_GRANTS_UNAVAILABLE, + "The MySQL grants are unavailable", + )), + }; + finish_connection(conn, result).await +} - let response = match query_result { - Ok(()) => CommunityAccountExecution { - action_type: request.action_type, +/// Builds a masked account-operation preview without opening `MySQL` or starting Java. +/// +/// # Errors +/// +/// Returns a field-specific account validation error. +pub(crate) fn preview_mysql_account( + application: &Application, + request: &AdministrationCommand, +) -> Result { + preview_account(&application.inner.account_previews, request) +} + +/// Executes one preview-authorized `MySQL` account operation through `mysql_async`. +/// +/// SQL execution errors are returned in [`AdministrationExecution`]. Datasource, +/// validation, preview-token, connection, read-only, and cleanup failures remain errors. +/// +/// # Errors +/// +/// Returns validation, token, datasource, connection, read-only, or cleanup errors. +pub(crate) async fn execute_mysql_account( + application: &Application, + request: &AdministrationCommand, +) -> Result { + let execution_sql = + authorize_account_preview(&application.inner.account_previews, request, true)?; + let preview_sql = build_account_sql(request, true)?; + + let resolved = resolve_native_connection(application, &request.datasource_id).await?; + if resolved.connection.read_only { + return Err(AppError::new( + AppErrorKind::Conflict, + ApiError::new( + "datasource_read_only", + "The datasource connection is configured as read-only", + ), + )); + } + let mut conn = open_resolved_connection(&resolved).await?; + let query_result = timed_query(execute_account_sql(&mut conn, &execution_sql)).await; + drop(execution_sql); + + let response = match query_result { + Ok(()) => AdministrationExecution { + action: request.action, + sql: preview_sql.clone(), + success: true, + message: Some("OK".to_owned()), + failure_code: None, + error_code: None, + sql_state: None, + }, + Err(AccountQueryFailure::Mysql(MysqlError::Server(server))) => { + AdministrationExecution { + action: request.action, sql: preview_sql.clone(), - success: true, - message: Some("OK".to_owned()), - failure_code: None, - error_code: None, - sql_state: None, - }, - Err(AccountQueryFailure::Mysql(MysqlError::Server(server))) => { - CommunityAccountExecution { - action_type: request.action_type, - sql: preview_sql.clone(), - success: false, - message: Some(redact_password( - &server.message, - request.password.as_deref(), - )), - failure_code: Some(ACCOUNT_EXECUTE_FAILED.to_owned()), - error_code: Some(server.code), - sql_state: Some(server.state), - } + success: false, + message: Some(redact_password( + &server.message, + request.credential.as_deref(), + )), + failure_code: Some(ACCOUNT_EXECUTE_FAILED.to_owned()), + error_code: Some(server.code), + sql_state: Some(server.state), } - Err(AccountQueryFailure::Mysql(_)) => account_outcome_unknown( - request.action_type, - preview_sql.clone(), - "The MySQL connection ended after dispatch; the account-operation outcome is unknown and must not be retried blindly".to_owned(), - ), - Err(AccountQueryFailure::Timeout) => account_outcome_unknown( - request.action_type, - preview_sql, - "The MySQL account operation timed out after dispatch; its outcome is unknown and must not be retried blindly".to_owned(), - ), - }; - if let Err(error) = finish_connection(conn, Ok(())).await { - tracing::warn!( - error = %error, - "native MySQL account connection cleanup failed after a settled operation" - ); } - Ok(response) + Err(AccountQueryFailure::Mysql(_)) => account_outcome_unknown( + request.action, + preview_sql.clone(), + "The MySQL connection ended after dispatch; the account-operation outcome is unknown and must not be retried blindly".to_owned(), + ), + Err(AccountQueryFailure::Timeout) => account_outcome_unknown( + request.action, + preview_sql, + "The MySQL account operation timed out after dispatch; its outcome is unknown and must not be retried blindly".to_owned(), + ), + }; + if let Err(error) = finish_connection(conn, Ok(())).await { + tracing::warn!( + error = %error, + "native MySQL account connection cleanup failed after a settled operation" + ); } + Ok(response) } fn account_outcome_unknown( - action_type: CommunityAccountAction, + action: AdministrationAction, sql: String, message: String, -) -> CommunityAccountExecution { - CommunityAccountExecution { - action_type, +) -> AdministrationExecution { + AdministrationExecution { + action, sql, success: false, message: Some(message), @@ -383,37 +433,61 @@ fn account_outcome_unknown( fn preview_account( registry: &AccountPreviewRegistry, - request: &CommunityAccountCommandRequest, -) -> Result { + request: &AdministrationCommand, +) -> Result { let execution_sql = build_account_sql(request, false)?; let sql = build_account_sql(request, true)?; let preview_token = registry.issue(&request.datasource_id, &execution_sql)?; drop(execution_sql); - Ok(CommunityAccountPreview { - action_type: request.action_type, + Ok(AdministrationPreview { + action: request.action, sql, preview_token, }) } +fn authorize_account_preview( + registry: &AccountPreviewRegistry, + request: &AdministrationCommand, + consume: bool, +) -> Result { + let execution_sql = build_account_sql(request, false)?; + let supplied_token = request.preview_token.as_deref().unwrap_or_default(); + let authorized = if consume { + registry.consume(supplied_token, &request.datasource_id, &execution_sql) + } else { + registry.authorizes(supplied_token, &request.datasource_id, &execution_sql) + }; + if !authorized { + return Err(AppError::new( + AppErrorKind::Conflict, + ApiError::new( + ACCOUNT_PREVIEW_TOKEN_MISMATCH, + "The MySQL account preview token does not match this operation", + ), + )); + } + Ok(execution_sql) +} + fn build_account_sql( - request: &CommunityAccountCommandRequest, + request: &AdministrationCommand, mask_sensitive: bool, ) -> Result { - let account = account_literal(&request.user, &request.host)?; - match request.action_type { - CommunityAccountAction::CreateUser => Ok(format!( + let account = account_literal(&request.principal)?; + match request.action { + AdministrationAction::CreatePrincipal => Ok(format!( "CREATE USER {account} IDENTIFIED BY {}", password_literal(request, mask_sensitive)? )), - CommunityAccountAction::AlterPassword => Ok(format!( + AdministrationAction::AlterCredential => Ok(format!( "ALTER USER {account} IDENTIFIED BY {}", password_literal(request, mask_sensitive)? )), - CommunityAccountAction::LockAccount => Ok(format!("ALTER USER {account} ACCOUNT LOCK")), - CommunityAccountAction::UnlockAccount => Ok(format!("ALTER USER {account} ACCOUNT UNLOCK")), - CommunityAccountAction::DropUser => Ok(format!("DROP USER {account}")), - CommunityAccountAction::GrantPrivilege => { + AdministrationAction::LockPrincipal => Ok(format!("ALTER USER {account} ACCOUNT LOCK")), + AdministrationAction::UnlockPrincipal => Ok(format!("ALTER USER {account} ACCOUNT UNLOCK")), + AdministrationAction::DropPrincipal => Ok(format!("DROP USER {account}")), + AdministrationAction::GrantPrivileges => { let privileges = privilege_list(&request.privileges)?; let scope = privilege_scope(request)?; let grant_option = if request.grant_option { @@ -425,7 +499,7 @@ fn build_account_sql( "GRANT {privileges} ON {scope} TO {account}{grant_option}" )) } - CommunityAccountAction::RevokePrivilege => { + AdministrationAction::RevokePrivileges => { let privileges = privilege_list(&request.privileges)?; let scope = privilege_scope(request)?; Ok(format!("REVOKE {privileges} ON {scope} FROM {account}")) @@ -434,10 +508,13 @@ fn build_account_sql( } fn password_literal( - request: &CommunityAccountCommandRequest, + request: &AdministrationCommand, mask_sensitive: bool, ) -> Result { - let password = request.password.as_deref().filter(|value| !is_blank(value)); + let password = request + .credential + .as_deref() + .filter(|value| !is_blank(value)); if password.is_none() { return Err(account_validation_error( "mysql.account.passwordRequired", @@ -451,7 +528,9 @@ fn password_literal( } } -fn account_literal(user: &str, host: &str) -> Result { +fn account_literal(principal: &PrincipalRef) -> Result { + let user = principal.name.as_str(); + let host = principal.qualifier.as_deref().unwrap_or_default(); validate_account_part( user, "mysql.account.userRequired", @@ -482,33 +561,39 @@ fn validate_account_part( Ok(()) } -fn privilege_scope(request: &CommunityAccountCommandRequest) -> Result { - match request.scope { - Some(CommunityAccountPrivilegeScope::Global) => Ok("*.*".to_owned()), - Some(CommunityAccountPrivilegeScope::Database) => { +fn privilege_scope(request: &AdministrationCommand) -> Result { + let Some(target) = request.target.as_ref() else { + return Err(account_validation_error( + "mysql.account.scopeRequired", + "A privilege scope is required", + )); + }; + match target.scope { + PrivilegeScope::Global => Ok("*.*".to_owned()), + PrivilegeScope::Database => { let database = required_identifier( - request.database_name.as_deref(), + target.database_name.as_deref(), "mysql.account.databaseRequired", "A database name is required for database privileges", )?; Ok(format!("{database}.*")) } - Some(CommunityAccountPrivilegeScope::Table) => { + PrivilegeScope::Table => { let database = required_identifier( - request.database_name.as_deref(), + target.database_name.as_deref(), "mysql.account.databaseRequired", "A database name is required for table privileges", )?; let table = required_identifier( - request.table_name.as_deref(), + target.object_name.as_deref(), "mysql.account.tableRequired", "A table name is required for table privileges", )?; Ok(format!("{database}.{table}")) } - None => Err(account_validation_error( - "mysql.account.scopeRequired", - "A privilege scope is required", + PrivilegeScope::Schema => Err(account_validation_error( + "mysql.account.scopeUnsupported", + "MySQL does not support schema-scoped account privileges", )), } } @@ -606,22 +691,22 @@ fn privilege_list(privileges: &[String]) -> Result { .join(", ")) } -fn parse_privilege(value: &str) -> Result { +fn parse_privilege(value: &str) -> Result { match value.trim().to_ascii_uppercase().as_str() { - "SELECT" => Ok(CommunityMysqlPrivilege::Select), - "INSERT" => Ok(CommunityMysqlPrivilege::Insert), - "UPDATE" => Ok(CommunityMysqlPrivilege::Update), - "DELETE" => Ok(CommunityMysqlPrivilege::Delete), - "CREATE" => Ok(CommunityMysqlPrivilege::Create), - "DROP" => Ok(CommunityMysqlPrivilege::Drop), - "ALTER" => Ok(CommunityMysqlPrivilege::Alter), - "INDEX" => Ok(CommunityMysqlPrivilege::Index), - "REFERENCES" => Ok(CommunityMysqlPrivilege::References), - "EXECUTE" => Ok(CommunityMysqlPrivilege::Execute), - "SHOW_VIEW" => Ok(CommunityMysqlPrivilege::ShowView), - "TRIGGER" => Ok(CommunityMysqlPrivilege::Trigger), - "EVENT" => Ok(CommunityMysqlPrivilege::Event), - "CREATE_TEMPORARY_TABLES" => Ok(CommunityMysqlPrivilege::CreateTemporaryTables), + "SELECT" => Ok(MysqlPrivilege::Select), + "INSERT" => Ok(MysqlPrivilege::Insert), + "UPDATE" => Ok(MysqlPrivilege::Update), + "DELETE" => Ok(MysqlPrivilege::Delete), + "CREATE" => Ok(MysqlPrivilege::Create), + "DROP" => Ok(MysqlPrivilege::Drop), + "ALTER" => Ok(MysqlPrivilege::Alter), + "INDEX" => Ok(MysqlPrivilege::Index), + "REFERENCES" => Ok(MysqlPrivilege::References), + "EXECUTE" => Ok(MysqlPrivilege::Execute), + "SHOW_VIEW" => Ok(MysqlPrivilege::ShowView), + "TRIGGER" => Ok(MysqlPrivilege::Trigger), + "EVENT" => Ok(MysqlPrivilege::Event), + "CREATE_TEMPORARY_TABLES" => Ok(MysqlPrivilege::CreateTemporaryTables), _ => Err(account_validation_error( "mysql.account.privilegeUnsupported", "The requested MySQL privilege is not supported", @@ -629,10 +714,10 @@ fn parse_privilege(value: &str) -> Result { } } -const fn privilege_sql_name(privilege: CommunityMysqlPrivilege) -> &'static str { +const fn privilege_sql_name(privilege: MysqlPrivilege) -> &'static str { match privilege { - CommunityMysqlPrivilege::ShowView => "SHOW VIEW", - CommunityMysqlPrivilege::CreateTemporaryTables => "CREATE TEMPORARY TABLES", + MysqlPrivilege::ShowView => "SHOW VIEW", + MysqlPrivilege::CreateTemporaryTables => "CREATE TEMPORARY TABLES", other => other.wire_name(), } } @@ -646,12 +731,12 @@ fn account( host: String, authentication_plugin: Option, locked: Option<&str>, -) -> CommunityAccount { - CommunityAccount { +) -> Principal { + Principal { display_name: format!("{user}@{host}"), - user, - host, - authentication_plugin, + name: user, + qualifier: Some(host), + authentication_method: authentication_plugin, locked: locked .and_then(|value| (!value.trim().is_empty()).then(|| value.eq_ignore_ascii_case("Y"))), } @@ -674,16 +759,16 @@ fn capability_with_message( account_list_readable: bool, account_lock_supported: bool, message: &str, -) -> CommunityAccountCapability { - CommunityAccountCapability { - db_type: "MYSQL".to_owned(), +) -> AdministrationCapability { + AdministrationCapability { + database_type: "MYSQL".to_owned(), product_name: "MySQL".to_owned(), product_version: None, - current_user: None, - connection_user, - account_list_readable, - account_lock_supported, - editable_privileges: CommunityMysqlPrivilege::ALL + current_principal: None, + connection_principal: connection_user, + principal_list_readable: account_list_readable, + principal_lock_supported: account_lock_supported, + editable_privileges: MysqlPrivilege::ALL .into_iter() .map(|privilege| privilege.wire_name().to_owned()) .collect(), @@ -731,61 +816,65 @@ fn is_blank(value: &str) -> bool { #[cfg(test)] mod tests { - use chat2db_contract::{ - CommunityAccountAction, CommunityAccountCommandRequest, CommunityAccountPrivilegeScope, + use crate::{ + Application, + native_administration_types::{ + AdministrationAction, AdministrationCommand, PrincipalRef, PrivilegeScope, + PrivilegeTarget, + }, }; - use crate::Application; - use super::{ - AccountPreviewRegistry, account_outcome_unknown, build_account_sql, preview_account, - redact_password, sql_mode_with_no_backslash_escapes, string_literal, + AccountPreviewRegistry, account_outcome_unknown, authorize_account_preview, + build_account_sql, execute_mysql_account as execute_mysql_account_impl, preview_account, + preview_mysql_account as preview_mysql_account_impl, redact_password, + sql_mode_with_no_backslash_escapes, string_literal, }; #[test] - fn every_account_action_matches_community_sql() { + fn every_administration_action_matches_mysql_sql() { assert_eq!( - sql(CommunityAccountAction::CreateUser), + sql(AdministrationAction::CreatePrincipal), "CREATE USER 'reader'@'%' IDENTIFIED BY 'pa''ss\\word'" ); assert_eq!( - sql(CommunityAccountAction::AlterPassword), + sql(AdministrationAction::AlterCredential), "ALTER USER 'reader'@'%' IDENTIFIED BY 'pa''ss\\word'" ); assert_eq!( - sql(CommunityAccountAction::LockAccount), + sql(AdministrationAction::LockPrincipal), "ALTER USER 'reader'@'%' ACCOUNT LOCK" ); assert_eq!( - sql(CommunityAccountAction::UnlockAccount), + sql(AdministrationAction::UnlockPrincipal), "ALTER USER 'reader'@'%' ACCOUNT UNLOCK" ); assert_eq!( - sql(CommunityAccountAction::DropUser), + sql(AdministrationAction::DropPrincipal), "DROP USER 'reader'@'%'" ); assert_eq!( - sql(CommunityAccountAction::GrantPrivilege), + sql(AdministrationAction::GrantPrivileges), "GRANT SELECT, SHOW VIEW, CREATE TEMPORARY TABLES ON `odd``db`.`order``item` TO 'reader'@'%' WITH GRANT OPTION" ); assert_eq!( - sql(CommunityAccountAction::RevokePrivilege), + sql(AdministrationAction::RevokePrivileges), "REVOKE SELECT, SHOW VIEW, CREATE TEMPORARY TABLES ON `odd``db`.`order``item` FROM 'reader'@'%'" ); } #[test] - fn scopes_and_account_literals_match_community_escaping() { - let mut request = command(CommunityAccountAction::GrantPrivilege); - request.user = "o'brien\\ops".to_owned(); - request.host = "local'host".to_owned(); - request.scope = Some(CommunityAccountPrivilegeScope::Global); + fn scopes_and_account_literals_use_stable_mysql_escaping() { + let mut request = command(AdministrationAction::GrantPrivileges); + request.principal.name = "o'brien\\ops".to_owned(); + request.principal.qualifier = Some("local'host".to_owned()); + request.target.as_mut().expect("target").scope = PrivilegeScope::Global; assert_eq!( build_account_sql(&request, false).expect("global grant"), "GRANT SELECT, SHOW VIEW, CREATE TEMPORARY TABLES ON *.* TO 'o''brien\\ops'@'local''host' WITH GRANT OPTION" ); - request.scope = Some(CommunityAccountPrivilegeScope::Database); + request.target.as_mut().expect("target").scope = PrivilegeScope::Database; assert_eq!( build_account_sql(&request, false).expect("database grant"), "GRANT SELECT, SHOW VIEW, CREATE TEMPORARY TABLES ON `odd``db`.* TO 'o''brien\\ops'@'local''host' WITH GRANT OPTION" @@ -820,7 +909,7 @@ mod tests { #[test] fn preview_masks_password_and_issues_an_opaque_token() { let registry = AccountPreviewRegistry::default(); - let request = command(CommunityAccountAction::CreateUser); + let request = command(AdministrationAction::CreatePrincipal); let preview = preview_account(®istry, &request).expect("valid account preview"); assert_eq!( @@ -843,7 +932,7 @@ mod tests { ); assert_ne!( preview.preview_token, - preview_account(®istry, &command_with_password("different")) + preview_account(®istry, &command_with_credential("different")) .expect("second preview") .preview_token ); @@ -852,7 +941,7 @@ mod tests { #[test] fn preview_tokens_are_datasource_bound_exact_and_single_use() { let registry = AccountPreviewRegistry::default(); - let request = command(CommunityAccountAction::DropUser); + let request = command(AdministrationAction::DropPrincipal); let sql = build_account_sql(&request, false).expect("account SQL"); let wrong_datasource = @@ -877,9 +966,28 @@ mod tests { assert!(!registry.consume(&valid.preview_token, &request.datasource_id, &sql)); } + #[test] + fn registry_preflight_does_not_consume_the_preview_token() { + let application = Application::new(); + let mut request = command(AdministrationAction::DropPrincipal); + request.preview_token = Some( + preview_mysql_account_impl(&application, &request) + .expect("account preview") + .preview_token, + ); + + authorize_account_preview(&application.inner.account_previews, &request, false) + .expect("non-consuming preflight"); + authorize_account_preview(&application.inner.account_previews, &request, true) + .expect("single execution authorization"); + let error = authorize_account_preview(&application.inner.account_previews, &request, true) + .expect_err("execution authorization must remain single-use"); + assert_eq!(error.api_error().code, "mysql.account.previewTokenMismatch"); + } + #[test] fn duplicate_privileges_are_removed_in_first_seen_order() { - let mut request = command(CommunityAccountAction::GrantPrivilege); + let mut request = command(AdministrationAction::GrantPrivileges); request.privileges = vec![ "select".to_owned(), " SELECT ".to_owned(), @@ -892,52 +1000,66 @@ mod tests { } #[test] - fn invalid_fields_return_community_error_codes() { - let mut request = command(CommunityAccountAction::CreateUser); - request.user.clear(); + fn invalid_fields_return_mysql_error_codes() { + let mut request = command(AdministrationAction::CreatePrincipal); + request.principal.name.clear(); assert_code(&request, "mysql.account.userRequired"); - request.user = "reader\0hidden".to_owned(); + request.principal.name = "reader\0hidden".to_owned(); assert_code(&request, "mysql.account.invalidAccountName"); - request.user = "reader".to_owned(); - request.password = Some(" ".to_owned()); + request.principal.name = "reader".to_owned(); + request.credential = Some(" ".to_owned()); assert_code(&request, "mysql.account.passwordRequired"); - request = command(CommunityAccountAction::GrantPrivilege); - request.scope = None; + request = command(AdministrationAction::GrantPrivileges); + request.target = None; assert_code(&request, "mysql.account.scopeRequired"); - request.scope = Some(CommunityAccountPrivilegeScope::Database); - request.database_name = None; + request.target = Some(PrivilegeTarget { + scope: PrivilegeScope::Database, + database_name: None, + schema_name: None, + object_name: None, + }); assert_code(&request, "mysql.account.databaseRequired"); - request.scope = Some(CommunityAccountPrivilegeScope::Table); - request.database_name = Some("inventory".to_owned()); - request.table_name = None; + request.target = Some(PrivilegeTarget { + scope: PrivilegeScope::Table, + database_name: Some("inventory".to_owned()), + schema_name: None, + object_name: None, + }); assert_code(&request, "mysql.account.tableRequired"); - request.table_name = Some("orders".to_owned()); + request.target.as_mut().expect("target").object_name = Some("orders".to_owned()); request.privileges = vec!["ROLE_ADMIN".to_owned()]; assert_code(&request, "mysql.account.privilegeUnsupported"); + + request.privileges = vec!["SELECT".to_owned()]; + request.target.as_mut().expect("target").scope = PrivilegeScope::Schema; + assert_code(&request, "mysql.account.scopeUnsupported"); } #[tokio::test] async fn token_mismatch_is_rejected_before_storage_or_mysql_access() { - let mut request = command(CommunityAccountAction::DropUser); + let mut request = command(AdministrationAction::DropPrincipal); request.preview_token = Some("not-the-preview-token".to_owned()); - let error = Application::new() - .execute_mysql_account(&request) + let application = Application::new(); + let direct_error = execute_mysql_account_impl(&application, &request) .await - .expect_err("token mismatch must fail before datasource resolution"); - assert_eq!(error.api_error().code, "mysql.account.previewTokenMismatch"); + .expect_err("direct token mismatch must fail before datasource resolution"); + assert_eq!( + direct_error.api_error().code, + "mysql.account.previewTokenMismatch" + ); } #[test] fn interrupted_account_writes_are_explicitly_non_retryable_unknown_outcomes() { let result = account_outcome_unknown( - CommunityAccountAction::AlterPassword, + AdministrationAction::AlterCredential, "ALTER USER 'reader'@'%' IDENTIFIED BY '******'".to_owned(), "The outcome is unknown and must not be retried blindly".to_owned(), ); @@ -968,37 +1090,42 @@ mod tests { assert!(redacted.contains("[REDACTED]")); } - fn sql(action: CommunityAccountAction) -> String { + fn sql(action: AdministrationAction) -> String { build_account_sql(&command(action), false).expect("account SQL") } - fn command(action_type: CommunityAccountAction) -> CommunityAccountCommandRequest { - CommunityAccountCommandRequest { + fn command(action: AdministrationAction) -> AdministrationCommand { + AdministrationCommand { datasource_id: "42".to_owned(), - user: "reader".to_owned(), - host: "%".to_owned(), - action_type, - scope: Some(CommunityAccountPrivilegeScope::Table), - database_name: Some("odd`db".to_owned()), - table_name: Some("order`item".to_owned()), + principal: PrincipalRef { + name: "reader".to_owned(), + qualifier: Some("%".to_owned()), + }, + action, + target: Some(PrivilegeTarget { + scope: PrivilegeScope::Table, + database_name: Some("odd`db".to_owned()), + schema_name: None, + object_name: Some("order`item".to_owned()), + }), privileges: vec![ "SELECT".to_owned(), "SHOW_VIEW".to_owned(), "CREATE_TEMPORARY_TABLES".to_owned(), ], grant_option: true, - password: Some("pa'ss\\word".to_owned()), + credential: Some("pa'ss\\word".to_owned()), preview_token: None, } } - fn command_with_password(password: &str) -> CommunityAccountCommandRequest { - let mut request = command(CommunityAccountAction::CreateUser); - request.password = Some(password.to_owned()); + fn command_with_credential(credential: &str) -> AdministrationCommand { + let mut request = command(AdministrationAction::CreatePrincipal); + request.credential = Some(credential.to_owned()); request } - fn assert_code(request: &CommunityAccountCommandRequest, expected: &str) { + fn assert_code(request: &AdministrationCommand, expected: &str) { let error = build_account_sql(request, false).expect_err("request must be invalid"); assert_eq!(error.api_error().code, expected); } diff --git a/crates/chat2db-core/src/mysql_dashboard.rs b/crates/chat2db-core/src/mysql_dashboard.rs index b248816..154a19c 100644 --- a/crates/chat2db-core/src/mysql_dashboard.rs +++ b/crates/chat2db-core/src/mysql_dashboard.rs @@ -15,8 +15,8 @@ use sqlparser::{ }; use crate::{ - AppError, AppErrorKind, Application, MysqlConsoleCancellation, MysqlConsoleRequest, - MysqlConsoleResult, now_millis, storage_call, + AppError, AppErrorKind, Application, NativeConsoleCancellation, NativeConsoleRequest, + NativeConsoleResult, now_millis, storage_call, }; const CHART_PAGE_SIZE: u32 = 200; @@ -96,11 +96,11 @@ impl Application { storage_call(move || storage.get_community_chart(id)).await } - /// Returns a detached chart copy, optionally refreshing its result through native `MySQL`. + /// Returns a detached chart copy, optionally refreshing its result through a native driver. /// /// # Errors /// - /// Returns validation, datasource, `MySQL`, result-limit, or durable-storage failures. + /// Returns validation, datasource, native-driver, result-limit, or durable-storage failures. pub async fn get_community_chart_detail( &self, id: i64, @@ -116,9 +116,18 @@ impl Application { return Ok(Some(chart)); }; + let database_type = match self.native_chart_database_type(&context).await { + Ok(database_type) => database_type, + Err(error) => { + self.record_chart_history(&chart, &context, None, None, Some(&error)) + .await; + return Err(error); + } + }; + let execution = self - .execute_mysql_read_console( - MysqlConsoleRequest { + .execute_native_read_console( + NativeConsoleRequest { datasource_id: context.datasource_id.clone(), database_name: context.database_name.clone().unwrap_or_default(), sql: context.sql.clone(), @@ -130,7 +139,7 @@ impl Application { explain: false, error_continue: false, }, - MysqlConsoleCancellation::new(), + NativeConsoleCancellation::new(), ) .await; @@ -141,8 +150,14 @@ impl Application { "chart_query_incomplete", "The chart query completed without a result", ); - self.record_chart_history(&chart, &context, None, Some(&error)) - .await; + self.record_chart_history( + &chart, + &context, + Some(&database_type), + None, + Some(&error), + ) + .await; return Err(error); }; if result.success { @@ -155,24 +170,49 @@ impl Application { .as_ref() .map_or_else(|| result.message.clone(), |error| error.message.clone()), ); - self.record_chart_history(&chart, &context, Some(&result), Some(&error)) - .await; + self.record_chart_history( + &chart, + &context, + Some(&database_type), + Some(&result), + Some(&error), + ) + .await; return Err(error); } } Err(error) => { - self.record_chart_history(&chart, &context, None, Some(&error)) - .await; + self.record_chart_history( + &chart, + &context, + Some(&database_type), + None, + Some(&error), + ) + .await; return Err(error); } }; - let header_metadata = self.chart_header_metadata(&context).await; + let header_metadata = self.chart_header_metadata(&context, &database_type).await; chart.meta_data = Some(chart_metadata(&result, header_metadata.as_ref())?); - self.record_chart_history(&chart, &context, Some(&result), None) + self.record_chart_history(&chart, &context, Some(&database_type), Some(&result), None) .await; Ok(Some(chart)) } + async fn native_chart_database_type( + &self, + context: &ChartRefreshContext, + ) -> Result { + self.require_native_driver_for_datasource(&context.datasource_id) + .await? + .database_types() + .first() + .copied() + .map(str::to_owned) + .ok_or_else(AppError::internal) + } + /// Creates one durable Community chart and returns its numeric id. /// /// # Errors @@ -214,7 +254,8 @@ impl Application { &self, chart: &CommunityChart, context: &ChartRefreshContext, - result: Option<&MysqlConsoleResult>, + database_type: Option<&str>, + result: Option<&NativeConsoleResult>, error: Option<&AppError>, ) { let Some(storage) = self.storage().cloned() else { @@ -233,7 +274,7 @@ impl Application { data_source_name: chart.data_source_name.clone(), connectable: Some(true), database_name: context.database_name.clone(), - database_type: Some("MYSQL".to_owned()), + database_type: database_type.map(str::to_owned), ddl: context.sql.clone(), status: if error.is_none() { "success" } else { "fail" }.to_owned(), operation_rows: result.and_then(|result| i64::try_from(result.row_count).ok()), @@ -256,6 +297,7 @@ impl Application { async fn chart_header_metadata( &self, context: &ChartRefreshContext, + database_type: &str, ) -> Option> { let table = chart_editable_table(&context.sql)?; let database_name = table @@ -268,7 +310,7 @@ impl Application { let columns = match self .list_community_columns(ListCommunityColumnsRequest { datasource_id: context.datasource_id.clone(), - database_type: "MYSQL".to_owned(), + database_type: database_type.to_owned(), database_name, schema_name, table_name: table.table_name, @@ -407,7 +449,7 @@ fn non_editable_projection(item: &SelectItem) -> bool { } fn chart_metadata( - result: &MysqlConsoleResult, + result: &NativeConsoleResult, header_metadata: Option<&HashMap>, ) -> Result { let metadata = json!({ @@ -614,7 +656,7 @@ mod tests { #[test] fn chart_metadata_matches_community_display_shape() { - let result = crate::MysqlConsoleResult { + let result = crate::NativeConsoleResult { statement_sequence: 1, result_set_id: Some(1), sql: "SELECT amount, note".to_owned(), @@ -689,7 +731,7 @@ mod tests { "DATETIME" ); - let result = crate::MysqlConsoleResult { + let result = crate::NativeConsoleResult { statement_sequence: 1, result_set_id: Some(1), sql: "SELECT CAST('2024-01-02' AS DATETIME)".to_owned(), diff --git a/crates/chat2db-core/src/mysql_ddl.rs b/crates/chat2db-core/src/mysql_ddl.rs index 6f9aa47..0b8af3b 100644 --- a/crates/chat2db-core/src/mysql_ddl.rs +++ b/crates/chat2db-core/src/mysql_ddl.rs @@ -1,11 +1,19 @@ -//! MySQL-specific structured SQL builders used by the retained Community API. +//! MySQL-specific structured SQL builders. use std::fmt::Write as _; use base64::{Engine as _, engine::general_purpose::STANDARD}; +use chrono::{DateTime, NaiveDate, NaiveDateTime, NaiveTime, Timelike as _}; use serde::{Deserialize, Serialize}; -use crate::AppError; +use crate::{ + AppError, + native_driver_types::{ + BuiltSql, CreateSchemaSqlRequest, DatabaseDefinition, DmlAssignment, DmlColumn, DmlRow, + DmlSqlRequest, DmlStatement, DmlTarget, DmlTemporalKind, DmlValue, NamespaceSqlOperation, + NamespaceSqlRequest, SchemaDefinition, + }, +}; pub const MYSQL_RESULT_DEFAULT_PLACEHOLDER: &str = "CHAT2DB_UPDATE_TABLE_DATA_USER_FILLED_DEFAULT"; pub const MYSQL_RESULT_GENERATED_PLACEHOLDER: &str = @@ -15,6 +23,19 @@ pub const MYSQL_PARTIAL_LARGE_VALUE_PREFIX: &str = "CHAT2DB_LARGE_VALUE_PREVIEW: const MAX_IDENTIFIER_CHARS: usize = 64; const MAX_VIEW_BODY_BYTES: usize = 1024 * 1024; const MAX_COMMENT_BYTES: usize = 2048; +const MAX_DML_COLUMNS: usize = 2_048; +const MAX_DML_ROWS: usize = 4_096; +const MAX_DML_VALUES: usize = 32_768; +const MAX_DML_IDENTIFIER_BYTES: usize = 512; +const MAX_DML_DATA_TYPE_NAME_BYTES: usize = 256; +const MAX_DML_DECIMAL_BYTES: usize = 1_024; +const MAX_DML_TEMPORAL_BYTES: usize = 128; +const MAX_DML_VALUE_BYTES: usize = 262_144; +const MAX_NAMESPACE_IDENTIFIER_BYTES: usize = 512; +const MAX_NAMESPACE_PROPERTY_BYTES: usize = 4_096; +const MAX_NAMESPACE_COMMENT_BYTES: usize = 65_536; +const MAX_REQUEST_BYTES: usize = 8 * 1_024 * 1_024; +const MAX_BUILT_SQL_BYTES: usize = 1_024 * 1_024; #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] @@ -577,6 +598,101 @@ pub fn build_mysql_external_in_values(values: &[String]) -> Result Result { + let CreateSchemaSqlRequest { schema } = request; + validate_namespace_schema(&schema)?; + finish_built_sql(format!( + "CREATE SCHEMA {}", + quote_namespace_identifier(&schema.name) + )) +} + +/// Builds one native `MySQL` database/schema namespace statement. +/// +/// # Errors +/// +/// Returns [`AppError`] when an identifier or option is invalid, or when the requested namespace +/// operation is not supported by `MySQL`. +pub(crate) fn build_mysql_namespace_request( + request: NamespaceSqlRequest, +) -> Result { + validate_namespace_request(&request)?; + let sql = match request.operation { + NamespaceSqlOperation::CreateDatabase { database } => { + let mut sql = format!( + "CREATE DATABASE {}", + quote_namespace_identifier(&database.name) + ); + if let Some(charset) = non_blank_owned(database.charset) { + write!( + &mut sql, + " DEFAULT CHARACTER SET={}", + mysql_namespace_property(&charset)? + ) + .expect("writing to a String cannot fail"); + } + if let Some(collation) = non_blank_owned(database.collation) { + write!( + &mut sql, + " COLLATE={}", + mysql_namespace_property(&collation)? + ) + .expect("writing to a String cannot fail"); + } + sql + } + NamespaceSqlOperation::AlterDatabase { .. } | NamespaceSqlOperation::AlterSchema { .. } => { + return Err(namespace_builder_not_supported()); + } + NamespaceSqlOperation::DropDatabase { database_name } => { + format!( + "DROP DATABASE {}", + quote_namespace_identifier(&database_name) + ) + } + NamespaceSqlOperation::UseDatabase { database_name } => { + format!("USE {}", quote_namespace_identifier(&database_name)) + } + NamespaceSqlOperation::CreateSchema { schema } => { + format!("CREATE SCHEMA {}", quote_namespace_identifier(&schema.name)) + } + NamespaceSqlOperation::DropSchema { schema_name } => { + format!("DROP SCHEMA {}", quote_namespace_identifier(&schema_name)) + } + }; + finish_built_sql(sql) +} + +/// Builds one typed native `MySQL` INSERT or UPDATE statement. +/// +/// # Errors +/// +/// Returns [`AppError`] when identifiers, row shapes, typed values, or predicates are invalid. +pub(crate) fn build_mysql_dml_request(request: DmlSqlRequest) -> Result { + validate_dml_request(&request)?; + let target = quote_dml_target(&request.target)?; + let sql = match request.statement { + DmlStatement::SingleInsert { columns, row } => { + build_typed_insert(&target, &columns, std::slice::from_ref(&row))? + } + DmlStatement::MultiInsert { columns, rows } => { + build_typed_insert(&target, &columns, &rows)? + } + DmlStatement::Update { + assignments, + predicates, + } => build_typed_update(&target, &assignments, &predicates)?, + }; + finish_built_sql(sql) +} + /// Wraps exactly one `MySQL` read statement in a bounded count query. /// /// # Errors @@ -1461,6 +1577,832 @@ fn is_decimal_literal(value: &str, allow_exponent: bool) -> bool { && (!whole.is_empty() || fraction.is_some_and(|fraction| !fraction.is_empty())) } +fn non_blank_owned(value: String) -> Option { + (!value.trim().is_empty()).then_some(value) +} + +fn invalid_database_request(message: impl Into) -> AppError { + AppError::invalid("invalid_database_request", message) +} + +fn protocol_limit_exceeded(message: impl Into) -> AppError { + AppError::invalid("protocol.limit_exceeded", message) +} + +fn namespace_builder_not_supported() -> AppError { + AppError::invalid( + "community.namespace_builder_not_supported", + "the selected Community plugin does not support this namespace operation", + ) +} + +fn dml_value_not_supported() -> AppError { + AppError::invalid( + "community.dml_value_not_supported", + "the selected Community value processor cannot preserve this typed value", + ) +} + +fn finish_built_sql(sql: String) -> Result { + if sql.trim().is_empty() || sql.len() > MAX_BUILT_SQL_BYTES { + return Err(protocol_limit_exceeded(format!( + "built SQL cannot exceed {MAX_BUILT_SQL_BYTES} bytes" + ))); + } + Ok(BuiltSql { sql }) +} + +fn validate_namespace_request(request: &NamespaceSqlRequest) -> Result<(), AppError> { + let mut bytes = 32_usize; + match &request.operation { + NamespaceSqlOperation::CreateDatabase { database } => { + bytes = bytes.saturating_add(validate_namespace_database(database)?); + } + NamespaceSqlOperation::AlterDatabase { + old_database, + new_database, + } => { + bytes = bytes + .saturating_add(validate_namespace_database(old_database)?) + .saturating_add(validate_namespace_database(new_database)?); + } + NamespaceSqlOperation::DropDatabase { database_name } + | NamespaceSqlOperation::UseDatabase { database_name } => { + validate_namespace_identifier(database_name, "database name")?; + bytes = bytes.saturating_add(database_name.len()); + } + NamespaceSqlOperation::CreateSchema { schema } => { + bytes = bytes.saturating_add(validate_namespace_schema(schema)?); + } + NamespaceSqlOperation::AlterSchema { + old_schema_name, + new_schema_name, + } => { + validate_namespace_identifier(old_schema_name, "old schema name")?; + validate_namespace_identifier(new_schema_name, "new schema name")?; + bytes = bytes + .saturating_add(old_schema_name.len()) + .saturating_add(new_schema_name.len()); + } + NamespaceSqlOperation::DropSchema { schema_name } => { + validate_namespace_identifier(schema_name, "schema name")?; + bytes = bytes.saturating_add(schema_name.len()); + } + } + if bytes > MAX_REQUEST_BYTES { + return Err(invalid_database_request(format!( + "Community namespace request cannot exceed {MAX_REQUEST_BYTES} encoded bytes" + ))); + } + Ok(()) +} + +fn validate_namespace_database(database: &DatabaseDefinition) -> Result { + validate_namespace_identifier(&database.name, "database name")?; + validate_utf8_limit( + &database.comment, + MAX_NAMESPACE_COMMENT_BYTES, + "database comment", + )?; + validate_namespace_property_preflight(&database.charset, "database charset")?; + validate_namespace_property_preflight(&database.collation, "database collation")?; + validate_namespace_property_preflight(&database.owner, "database owner")?; + validate_namespace_comment_runtime(&database.comment)?; + validate_namespace_property_runtime(&database.charset)?; + validate_namespace_property_runtime(&database.collation)?; + validate_namespace_property_runtime(&database.owner)?; + Ok(database + .name + .len() + .saturating_add(database.comment.len()) + .saturating_add(database.charset.len()) + .saturating_add(database.collation.len()) + .saturating_add(database.owner.len()) + .saturating_add(64)) +} + +fn validate_namespace_schema(schema: &SchemaDefinition) -> Result { + if !schema.database_name.is_empty() { + validate_namespace_identifier(&schema.database_name, "schema database name")?; + } + validate_namespace_identifier(&schema.name, "schema name")?; + validate_utf8_limit( + &schema.comment, + MAX_NAMESPACE_COMMENT_BYTES, + "schema comment", + )?; + validate_namespace_property_preflight(&schema.owner, "schema owner")?; + validate_namespace_comment_runtime(&schema.comment)?; + validate_namespace_property_runtime(&schema.owner)?; + Ok(schema + .database_name + .len() + .saturating_add(schema.name.len()) + .saturating_add(schema.comment.len()) + .saturating_add(schema.owner.len()) + .saturating_add(64)) +} + +fn validate_namespace_identifier(value: &str, field: &str) -> Result<(), AppError> { + validate_non_blank_utf8_limit(value, MAX_NAMESPACE_IDENTIFIER_BYTES, field)?; + if value.trim() != value + || value.chars().any(char::is_control) + || value.contains(['.', ';', '\'', '"', '`', '[', ']']) + || contains_comment_marker(value) + { + return Err(invalid_database_request(format!( + "Community namespace {field} contains unsafe identifier syntax" + ))); + } + Ok(()) +} + +fn validate_namespace_property_preflight(value: &str, field: &str) -> Result<(), AppError> { + validate_utf8_limit(value, MAX_NAMESPACE_PROPERTY_BYTES, field)?; + if !value.is_empty() + && (value.trim() != value + || value.chars().any(char::is_control) + || value.contains([';', '\'', '"', '`', '[', ']']) + || contains_comment_marker(value)) + { + return Err(invalid_database_request(format!( + "Community namespace {field} contains unsafe property syntax" + ))); + } + Ok(()) +} + +fn validate_namespace_property_runtime(value: &str) -> Result<(), AppError> { + if !value.is_empty() + && !value.chars().all(|character| { + character.is_alphanumeric() || matches!(character, '_' | '-' | '$' | '@') + }) + { + return Err(AppError::invalid( + "community.namespace_property_invalid", + "a Community namespace property contains unsafe syntax", + )); + } + Ok(()) +} + +fn validate_namespace_comment_runtime(value: &str) -> Result<(), AppError> { + if !value.is_empty() + && (value.chars().any(char::is_control) + || value.contains(['\'', '\\']) + || contains_comment_marker(value)) + { + return Err(AppError::invalid( + "community.namespace_comment_invalid", + "a Community namespace comment contains unsafe syntax", + )); + } + Ok(()) +} + +fn mysql_namespace_property(value: &str) -> Result<&str, AppError> { + if value + .bytes() + .all(|byte| byte.is_ascii_alphanumeric() || byte == b'_') + { + Ok(value) + } else { + Err(AppError::invalid( + "community.namespace_builder_failed", + "the Community namespace builder failed internally", + )) + } +} + +fn contains_comment_marker(value: &str) -> bool { + value.contains("--") || value.contains("/*") || value.contains("*/") +} + +fn validate_non_blank_utf8_limit(value: &str, maximum: usize, field: &str) -> Result<(), AppError> { + if value.trim().is_empty() || value.len() > maximum { + return Err(invalid_database_request(format!( + "{field} must be non-blank and cannot exceed {maximum} UTF-8 bytes" + ))); + } + Ok(()) +} + +fn validate_utf8_limit(value: &str, maximum: usize, field: &str) -> Result<(), AppError> { + if value.len() > maximum { + return Err(invalid_database_request(format!( + "{field} cannot exceed {maximum} UTF-8 bytes" + ))); + } + Ok(()) +} + +fn quote_native_identifier(value: &str) -> Result { + let quoted = format!("`{value}`"); + if quoted.len() > MAX_DML_IDENTIFIER_BYTES { + return Err(protocol_limit_exceeded(format!( + "quoted identifier cannot exceed {MAX_DML_IDENTIFIER_BYTES} bytes" + ))); + } + Ok(quoted) +} + +fn quote_namespace_identifier(value: &str) -> String { + format!("`{value}`") +} + +fn validate_dml_request(request: &DmlSqlRequest) -> Result<(), AppError> { + let mut bytes = validate_dml_target(&request.target)?.saturating_add(32); + let value_count = match &request.statement { + DmlStatement::SingleInsert { columns, row } => { + bytes = bytes.saturating_add(validate_dml_columns(columns, "insert columns")?); + if row.values.len() != columns.len() { + return Err(invalid_database_request( + "Community DML insert row width must equal its column count", + )); + } + bytes = bytes.saturating_add(validate_dml_values(&row.values, false)?); + row.values.len() + } + DmlStatement::MultiInsert { columns, rows } => { + bytes = bytes.saturating_add(validate_dml_columns(columns, "batch insert columns")?); + if rows.is_empty() || rows.len() > MAX_DML_ROWS { + return Err(invalid_database_request(format!( + "Community DML batch row count must be between 1 and {MAX_DML_ROWS}" + ))); + } + let total = columns + .len() + .checked_mul(rows.len()) + .ok_or_else(|| invalid_database_request("Community DML value count overflowed"))?; + for row in rows { + if row.values.len() != columns.len() { + return Err(invalid_database_request( + "Community DML batch row width must equal its column count", + )); + } + bytes = bytes + .saturating_add(validate_dml_values(&row.values, false)?) + .saturating_add(8); + } + total + } + DmlStatement::Update { + assignments, + predicates, + } => { + bytes = bytes + .saturating_add(validate_dml_assignments( + assignments, + "update assignments", + false, + )?) + .saturating_add(validate_dml_assignments( + predicates, + "update predicates", + true, + )?); + assignments.len().saturating_add(predicates.len()) + } + }; + if value_count > MAX_DML_VALUES { + return Err(invalid_database_request(format!( + "Community DML cannot contain more than {MAX_DML_VALUES} values" + ))); + } + if bytes > MAX_REQUEST_BYTES { + return Err(invalid_database_request(format!( + "Community DML request cannot exceed {MAX_REQUEST_BYTES} encoded bytes" + ))); + } + Ok(()) +} + +fn validate_dml_target(target: &DmlTarget) -> Result { + let mut bytes = target.table_name.len().saturating_add(16); + validate_dml_identifier(&target.table_name, "table name")?; + if let Some(database_name) = &target.database_name { + validate_dml_identifier(database_name, "database name")?; + bytes = bytes.saturating_add(database_name.len()).saturating_add(8); + } + if let Some(schema_name) = &target.schema_name { + validate_dml_identifier(schema_name, "schema name")?; + bytes = bytes.saturating_add(schema_name.len()).saturating_add(8); + } + Ok(bytes) +} + +fn validate_dml_columns(columns: &[DmlColumn], field: &str) -> Result { + if columns.is_empty() || columns.len() > MAX_DML_COLUMNS { + return Err(invalid_database_request(format!( + "Community DML {field} count must be between 1 and {MAX_DML_COLUMNS}" + ))); + } + let mut names = std::collections::HashSet::with_capacity(columns.len()); + let mut bytes = 0_usize; + for column in columns { + bytes = bytes.saturating_add(validate_dml_column(column)?); + if !names.insert(column.name.as_str()) { + return Err(invalid_database_request(format!( + "Community DML {field} cannot contain duplicate column names" + ))); + } + } + Ok(bytes) +} + +fn validate_dml_assignments( + assignments: &[DmlAssignment], + field: &str, + reject_null: bool, +) -> Result { + if assignments.is_empty() || assignments.len() > MAX_DML_COLUMNS { + return Err(invalid_database_request(format!( + "Community DML {field} count must be between 1 and {MAX_DML_COLUMNS}" + ))); + } + let mut names = std::collections::HashSet::with_capacity(assignments.len()); + let mut bytes = 0_usize; + for assignment in assignments { + bytes = bytes + .saturating_add(validate_dml_column(&assignment.column)?) + .saturating_add(validate_dml_value_preflight( + &assignment.value, + reject_null, + )?) + .saturating_add(8); + if !names.insert(assignment.column.name.as_str()) { + return Err(invalid_database_request(format!( + "Community DML {field} cannot contain duplicate column names" + ))); + } + } + Ok(bytes) +} + +fn validate_dml_column(column: &DmlColumn) -> Result { + validate_dml_identifier(&column.name, "column name")?; + validate_non_blank_utf8_limit( + &column.data_type_name, + MAX_DML_DATA_TYPE_NAME_BYTES, + "DML data type name", + )?; + if column.data_type_name.chars().any(char::is_control) { + return Err(invalid_database_request( + "Community DML data type names cannot contain control characters", + )); + } + if column + .precision + .is_some_and(|value| value > i32::MAX as u32) + { + return Err(invalid_database_request( + "Community DML precision cannot exceed the Java Integer range", + )); + } + Ok(column + .name + .len() + .saturating_add(column.data_type_name.len()) + .saturating_add(24)) +} + +fn validate_dml_values(values: &[DmlValue], reject_null: bool) -> Result { + if values.len() > MAX_DML_VALUES { + return Err(invalid_database_request(format!( + "Community DML cannot contain more than {MAX_DML_VALUES} values" + ))); + } + let mut bytes = 0_usize; + for value in values { + bytes = bytes + .saturating_add(validate_dml_value_preflight(value, reject_null)?) + .saturating_add(8); + } + Ok(bytes) +} + +fn validate_dml_value_preflight(value: &DmlValue, reject_null: bool) -> Result { + match value { + DmlValue::Null if reject_null => Err(invalid_database_request( + "Community DML equality predicates cannot compare NULL", + )), + DmlValue::String(value) => { + validate_utf8_limit(value, MAX_DML_VALUE_BYTES, "DML string value")?; + Ok(value.len()) + } + DmlValue::Decimal(value) => { + validate_plain_decimal(value)?; + Ok(value.len()) + } + DmlValue::Temporal { iso8601, .. } => { + validate_non_blank_utf8_limit(iso8601, MAX_DML_TEMPORAL_BYTES, "DML temporal value")?; + if iso8601.chars().any(char::is_control) { + return Err(invalid_database_request( + "Community DML temporal values cannot contain control characters", + )); + } + Ok(iso8601.len()) + } + DmlValue::Binary(value) if value.len() > MAX_DML_VALUE_BYTES => { + Err(invalid_database_request(format!( + "Community DML binary values cannot exceed {MAX_DML_VALUE_BYTES} bytes" + ))) + } + DmlValue::Null | DmlValue::Boolean(_) => Ok(1), + DmlValue::Binary(value) => Ok(value.len()), + } +} + +fn validate_dml_identifier(value: &str, field: &str) -> Result<(), AppError> { + validate_non_blank_utf8_limit(value, MAX_DML_IDENTIFIER_BYTES, field)?; + if value.trim() != value + || value.chars().any(char::is_control) + || value.contains(['.', ';', '\'', '"', '`', '[', ']']) + || contains_comment_marker(value) + { + return Err(invalid_database_request(format!( + "Community DML {field} contains unsafe identifier syntax" + ))); + } + Ok(()) +} + +fn validate_plain_decimal(value: &str) -> Result<(), AppError> { + validate_non_blank_utf8_limit(value, MAX_DML_DECIMAL_BYTES, "DML decimal value")?; + let bytes = value.as_bytes(); + let mut index = usize::from(bytes.first() == Some(&b'-')); + let integer_start = index; + while bytes.get(index).is_some_and(u8::is_ascii_digit) { + index += 1; + } + if index == integer_start { + return Err(invalid_database_request( + "Community DML decimal values must use plain base-10 notation", + )); + } + if bytes.get(index) == Some(&b'.') { + index += 1; + let fraction_start = index; + while bytes.get(index).is_some_and(u8::is_ascii_digit) { + index += 1; + } + if index == fraction_start { + return Err(invalid_database_request( + "Community DML decimal fractions require at least one digit", + )); + } + } + if index != bytes.len() { + return Err(invalid_database_request( + "Community DML decimal values must use plain base-10 notation", + )); + } + Ok(()) +} + +fn quote_dml_target(target: &DmlTarget) -> Result { + let mut segments = Vec::with_capacity(3); + if let Some(database_name) = &target.database_name { + segments.push(quote_native_identifier(database_name)?); + } + if let Some(schema_name) = &target.schema_name { + segments.push(quote_native_identifier(schema_name)?); + } + segments.push(quote_native_identifier(&target.table_name)?); + Ok(segments.join(".")) +} + +fn build_typed_insert( + target: &str, + columns: &[DmlColumn], + rows: &[DmlRow], +) -> Result { + let column_sql = columns + .iter() + .map(|column| quote_native_identifier(&column.name)) + .collect::, _>>()? + .join(", "); + let values = rows + .iter() + .map(|row| build_typed_row(columns, row)) + .collect::, _>>()? + .join(", "); + Ok(format!( + "INSERT INTO {target} ({column_sql}) VALUES {values}" + )) +} + +fn build_typed_row(columns: &[DmlColumn], row: &DmlRow) -> Result { + let values = columns + .iter() + .zip(&row.values) + .map(|(column, value)| serialize_dml_value(column, value)) + .collect::, _>>()?; + Ok(format!("({})", values.join(", "))) +} + +fn build_typed_update( + target: &str, + assignments: &[DmlAssignment], + predicates: &[DmlAssignment], +) -> Result { + let assignments = assignments + .iter() + .map(|assignment| { + Ok(format!( + "{} = {}", + quote_native_identifier(&assignment.column.name)?, + serialize_dml_value(&assignment.column, &assignment.value)? + )) + }) + .collect::, AppError>>()? + .join(", "); + let predicates = predicates + .iter() + .map(build_typed_predicate) + .collect::, _>>()? + .join(" AND "); + Ok(format!( + "UPDATE {target} SET {assignments} WHERE {predicates}" + )) +} + +fn build_typed_predicate(assignment: &DmlAssignment) -> Result { + let column = quote_native_identifier(&assignment.column.name)?; + Ok(format!( + "{column} = {}", + serialize_dml_value(&assignment.column, &assignment.value)? + )) +} + +fn serialize_dml_value(column: &DmlColumn, value: &DmlValue) -> Result { + match value { + DmlValue::Null => Ok("NULL".to_owned()), + DmlValue::String(value) => quote_dml_string(value), + DmlValue::Decimal(value) => { + if !is_compatible_decimal_type(&column.data_type_name) { + return Err(dml_value_not_supported()); + } + Ok(canonical_decimal(value)) + } + DmlValue::Boolean(value) => { + if !is_compatible_boolean_type(&column.data_type_name) { + return Err(dml_value_not_supported()); + } + let bit = if *value { "1" } else { "0" }; + Ok(if base_data_type(&column.data_type_name) == "BIT" { + format!("b'{bit}'") + } else { + format!("'{bit}'") + }) + } + DmlValue::Temporal { kind, iso8601 } => { + let canonical = canonical_temporal(*kind, iso8601)?; + if !is_compatible_temporal_type(*kind, &column.data_type_name) + || !is_compatible_temporal_scale(*kind, column.scale.unwrap_or(0), &canonical) + { + return Err(dml_value_not_supported()); + } + quote_dml_string(&canonical) + } + DmlValue::Binary(bytes) => { + if !is_compatible_binary_type(&column.data_type_name) { + return Err(dml_value_not_supported()); + } + let mut literal = String::with_capacity(bytes.len().saturating_mul(2) + 2); + literal.push_str("0x"); + for byte in bytes { + write!(&mut literal, "{byte:02X}").expect("writing to a String cannot fail"); + } + Ok(literal) + } + } +} + +fn quote_dml_string(value: &str) -> Result { + if value.contains(['\\', '\0']) { + return Err(dml_value_not_supported()); + } + Ok(format!("'{}'", value.replace('\'', "''"))) +} + +fn canonical_decimal(value: &str) -> String { + let negative = value.starts_with('-'); + let unsigned = value.strip_prefix('-').unwrap_or(value); + let (whole, fraction) = unsigned + .split_once('.') + .map_or((unsigned, ""), |(whole, fraction)| (whole, fraction)); + let whole = whole.trim_start_matches('0'); + let whole = if whole.is_empty() { "0" } else { whole }; + let fraction = fraction.trim_end_matches('0'); + let zero = whole == "0" && fraction.is_empty(); + let mut canonical = if fraction.is_empty() { + whole.to_owned() + } else { + format!("{whole}.{fraction}") + }; + if negative && !zero { + canonical.insert(0, '-'); + } + canonical +} + +fn base_data_type(data_type_name: &str) -> String { + let normalized = data_type_name.trim().to_ascii_uppercase(); + let end = normalized.find(['(', ' ']).unwrap_or(normalized.len()); + normalized[..end].trim().to_owned() +} + +fn is_compatible_decimal_type(data_type_name: &str) -> bool { + matches!( + base_data_type(data_type_name).as_str(), + "TINYINT" + | "SMALLINT" + | "MEDIUMINT" + | "BIGINT" + | "INTEGER" + | "INT" + | "INT2" + | "INT4" + | "INT8" + | "DECIMAL" + | "DEC" + | "NUMERIC" + | "NUMBER" + | "FLOAT" + | "FLOAT4" + | "FLOAT8" + | "DOUBLE" + | "REAL" + | "MONEY" + | "SMALLMONEY" + | "SERIAL" + | "BIGSERIAL" + | "BINARY_FLOAT" + | "BINARY_DOUBLE" + ) +} + +fn is_compatible_boolean_type(data_type_name: &str) -> bool { + let normalized = data_type_name.trim().to_ascii_uppercase(); + !normalized.contains("VARYING") + && matches!( + base_data_type(&normalized).as_str(), + "BOOL" | "BOOLEAN" | "BIT" + ) +} + +fn is_compatible_binary_type(data_type_name: &str) -> bool { + let normalized = data_type_name.trim().to_ascii_uppercase(); + normalized.starts_with("BIT VARYING") + || normalized.starts_with("LONG RAW") + || matches!( + base_data_type(&normalized).as_str(), + "BINARY" + | "VARBINARY" + | "LONGVARBINARY" + | "BLOB" + | "TINYBLOB" + | "MEDIUMBLOB" + | "LONGBLOB" + | "BYTEA" + | "RAW" + | "IMAGE" + ) +} + +fn canonical_temporal(kind: DmlTemporalKind, value: &str) -> Result { + let invalid = || { + AppError::invalid( + "community.dml_temporal_invalid", + "the Community DML temporal value is not valid ISO-8601", + ) + }; + match kind { + DmlTemporalKind::Date => NaiveDate::parse_from_str(value, "%Y-%m-%d") + .map(|date| date.format("%Y-%m-%d").to_string()) + .map_err(|_| invalid()), + DmlTemporalKind::Time => parse_iso_time(value) + .map(format_iso_time) + .ok_or_else(invalid), + DmlTemporalKind::LocalDatetime => parse_iso_local_datetime(value) + .map(format_iso_local_datetime) + .ok_or_else(invalid), + DmlTemporalKind::OffsetDatetime => { + let normalized = normalize_offset_datetime(value).ok_or_else(&invalid)?; + let parsed = DateTime::parse_from_rfc3339(&normalized).map_err(|_| invalid())?; + let offset = if parsed.offset().local_minus_utc() == 0 { + "Z".to_owned() + } else { + parsed.offset().to_string() + }; + Ok(format!( + "{} {offset}", + format_iso_local_datetime(parsed.naive_local()) + )) + } + } +} + +fn parse_iso_time(value: &str) -> Option { + ["%H:%M", "%H:%M:%S", "%H:%M:%S%.f"] + .into_iter() + .find_map(|format| NaiveTime::parse_from_str(value, format).ok()) +} + +fn parse_iso_local_datetime(value: &str) -> Option { + [ + "%Y-%m-%dT%H:%M", + "%Y-%m-%dT%H:%M:%S", + "%Y-%m-%dT%H:%M:%S%.f", + ] + .into_iter() + .find_map(|format| NaiveDateTime::parse_from_str(value, format).ok()) +} + +fn format_iso_time(value: NaiveTime) -> String { + let mut rendered = value.format("%H:%M").to_string(); + if value.second() != 0 || value.nanosecond() != 0 { + write!(&mut rendered, ":{:02}", value.second()).expect("writing to a String cannot fail"); + } + if value.nanosecond() != 0 { + let fraction = format!("{:09}", value.nanosecond()); + rendered.push('.'); + rendered.push_str(fraction.trim_end_matches('0')); + } + rendered +} + +fn format_iso_local_datetime(value: NaiveDateTime) -> String { + format!( + "{} {}", + value.date().format("%Y-%m-%d"), + format_iso_time(value.time()) + ) +} + +fn normalize_offset_datetime(value: &str) -> Option { + let offset_start = if value.ends_with('Z') { + value.len().checked_sub(1)? + } else { + value + .char_indices() + .skip_while(|(index, _)| *index <= 10) + .find_map(|(index, character)| matches!(character, '+' | '-').then_some(index))? + }; + let (local, offset) = value.split_at(offset_start); + let local = parse_iso_local_datetime(local)?; + Some(format!( + "{}T{:02}:{:02}:{:02}.{:09}{offset}", + local.date().format("%Y-%m-%d"), + local.hour(), + local.minute(), + local.second(), + local.nanosecond() + )) +} + +fn is_compatible_temporal_type(kind: DmlTemporalKind, data_type_name: &str) -> bool { + let normalized = data_type_name.trim().to_ascii_uppercase(); + let offset = normalized.contains("WITH TIME ZONE") + || normalized.contains("TIMESTAMPTZ") + || normalized.contains("DATETIMEOFFSET") + || normalized.contains("TIMESTAMP_TZ"); + let local_time_zone = normalized.contains("WITH LOCAL TIME ZONE"); + let date_time = normalized.contains("TIMESTAMP") + || normalized.contains("DATETIME") + || normalized.contains("SMALLDATETIME"); + let date = + normalized == "DATE" || normalized.starts_with("DATE(") || normalized.starts_with("DATE "); + let time = (normalized == "TIME" + || normalized.starts_with("TIME(") + || normalized.starts_with("TIME ")) + && !normalized.contains("TIMESTAMP"); + match kind { + DmlTemporalKind::Date => date, + DmlTemporalKind::Time => time && !offset, + DmlTemporalKind::LocalDatetime => (date_time && (!offset || local_time_zone)) || date, + DmlTemporalKind::OffsetDatetime => offset && !local_time_zone, + } +} + +fn is_compatible_temporal_scale(kind: DmlTemporalKind, scale: i32, canonical: &str) -> bool { + if !(0..=9).contains(&scale) { + return false; + } + let scale = usize::try_from(scale).expect("a validated temporal scale is non-negative"); + if kind == DmlTemporalKind::Date { + return true; + } + canonical.find('.').is_none_or(|decimal_point| { + canonical[decimal_point + 1..] + .bytes() + .take_while(u8::is_ascii_digit) + .count() + <= scale + }) +} + fn build_create_namespace(database: &MysqlDatabaseDefinition) -> Result { let mut sql = format!( "CREATE DATABASE {}{}", @@ -2291,6 +3233,272 @@ mod tests { assert!(quote_qualified_name(&mismatched).is_err()); } + #[test] + fn native_dialect_builds_structured_namespace_sql() { + let create = build_mysql_namespace_request(NamespaceSqlRequest { + operation: NamespaceSqlOperation::CreateDatabase { + database: DatabaseDefinition { + name: "app_db".to_owned(), + comment: String::new(), + charset: "utf8mb4".to_owned(), + collation: "utf8mb4_0900_ai_ci".to_owned(), + owner: String::new(), + system: false, + }, + }, + }) + .expect("native namespace SQL"); + assert_eq!( + create.sql, + "CREATE DATABASE `app_db` DEFAULT CHARACTER SET=utf8mb4 COLLATE=utf8mb4_0900_ai_ci" + ); + + let use_database = build_mysql_namespace_request(NamespaceSqlRequest { + operation: NamespaceSqlOperation::UseDatabase { + database_name: "app_db".to_owned(), + }, + }) + .expect("native USE SQL"); + assert_eq!(use_database.sql, "USE `app_db`"); + + let rename = build_mysql_namespace_request(NamespaceSqlRequest { + operation: NamespaceSqlOperation::AlterSchema { + old_schema_name: "before".to_owned(), + new_schema_name: "after".to_owned(), + }, + }) + .expect_err("MySQL schema rename must fail closed"); + assert_eq!( + rename.api_error().code, + "community.namespace_builder_not_supported" + ); + } + + #[test] + fn native_dialect_builds_typed_insert_and_update_sql() { + let columns = vec![ + DmlColumn { + name: "id".to_owned(), + data_type_name: "BIGINT".to_owned(), + precision: None, + scale: None, + }, + DmlColumn { + name: "label_value".to_owned(), + data_type_name: "VARCHAR".to_owned(), + precision: Some(255), + scale: None, + }, + DmlColumn { + name: "payload".to_owned(), + data_type_name: "BLOB".to_owned(), + precision: None, + scale: None, + }, + ]; + let insert = build_mysql_dml_request(DmlSqlRequest { + target: DmlTarget { + database_name: Some("app".to_owned()), + schema_name: None, + table_name: "items".to_owned(), + }, + statement: DmlStatement::SingleInsert { + columns, + row: DmlRow { + values: vec![ + DmlValue::Decimal("42.00".to_owned()), + DmlValue::String("O'Reilly".to_owned()), + DmlValue::Binary(vec![0, 255, 16]), + ], + }, + }, + }) + .expect("typed INSERT"); + assert_eq!( + insert.sql, + "INSERT INTO `app`.`items` (`id`, `label_value`, `payload`) VALUES (42, 'O''Reilly', 0x00FF10)" + ); + + let update = build_mysql_dml_request(DmlSqlRequest { + target: DmlTarget { + database_name: None, + schema_name: None, + table_name: "items".to_owned(), + }, + statement: DmlStatement::Update { + assignments: vec![DmlAssignment { + column: DmlColumn { + name: "active".to_owned(), + data_type_name: "BOOLEAN".to_owned(), + precision: None, + scale: None, + }, + value: DmlValue::Boolean(true), + }], + predicates: vec![DmlAssignment { + column: DmlColumn { + name: "id".to_owned(), + data_type_name: "BIGINT".to_owned(), + precision: None, + scale: None, + }, + value: DmlValue::Decimal("1".to_owned()), + }], + }, + }) + .expect("typed UPDATE"); + assert_eq!( + update.sql, + "UPDATE `items` SET `active` = '1' WHERE `id` = 1" + ); + } + + #[test] + fn native_dialect_rejects_invalid_predicates_and_typed_values() { + let target = DmlTarget { + database_name: None, + schema_name: None, + table_name: "items".to_owned(), + }; + let column = DmlColumn { + name: "payload".to_owned(), + data_type_name: "BLOB".to_owned(), + precision: None, + scale: None, + }; + let unbounded = build_mysql_dml_request(DmlSqlRequest { + target: target.clone(), + statement: DmlStatement::Update { + assignments: vec![DmlAssignment { + column: column.clone(), + value: DmlValue::Null, + }], + predicates: Vec::new(), + }, + }) + .expect_err("unbounded UPDATE"); + assert_eq!(unbounded.api_error().code, "invalid_database_request"); + + let null_predicate = build_mysql_dml_request(DmlSqlRequest { + target: target.clone(), + statement: DmlStatement::Update { + assignments: vec![DmlAssignment { + column: column.clone(), + value: DmlValue::Null, + }], + predicates: vec![DmlAssignment { + column: column.clone(), + value: DmlValue::Null, + }], + }, + }) + .expect_err("NULL equality predicate"); + assert_eq!(null_predicate.api_error().code, "invalid_database_request"); + + let incompatible = build_mysql_dml_request(DmlSqlRequest { + target: target.clone(), + statement: DmlStatement::SingleInsert { + columns: vec![DmlColumn { + data_type_name: "VARCHAR".to_owned(), + ..column.clone() + }], + row: DmlRow { + values: vec![DmlValue::Binary(vec![0, 255])], + }, + }, + }) + .expect_err("binary value for VARCHAR"); + assert_eq!( + incompatible.api_error().code, + "community.dml_value_not_supported" + ); + + let malformed_temporal = build_mysql_dml_request(DmlSqlRequest { + target, + statement: DmlStatement::SingleInsert { + columns: vec![DmlColumn { + name: "created_at".to_owned(), + data_type_name: "DATETIME".to_owned(), + precision: None, + scale: None, + }], + row: DmlRow { + values: vec![DmlValue::Temporal { + kind: DmlTemporalKind::LocalDatetime, + iso8601: "2026-13-99T99:99:99".to_owned(), + }], + }, + }, + }) + .expect_err("malformed temporal"); + assert_eq!( + malformed_temporal.api_error().code, + "community.dml_temporal_invalid" + ); + } + + #[test] + fn native_dialect_enforces_bridge_limits_and_sql_mode_independent_strings() { + let target = DmlTarget { + database_name: None, + schema_name: None, + table_name: "items".to_owned(), + }; + let string_column = DmlColumn { + name: "label".to_owned(), + data_type_name: "VARCHAR".to_owned(), + precision: Some(255), + scale: None, + }; + let unsupported_string = build_mysql_dml_request(DmlSqlRequest { + target: target.clone(), + statement: DmlStatement::SingleInsert { + columns: vec![string_column.clone()], + row: DmlRow { + values: vec![DmlValue::String("back\\slash".to_owned())], + }, + }, + }) + .expect_err("backslashes cannot be preserved across every MySQL SQL mode"); + assert_eq!( + unsupported_string.api_error().code, + "community.dml_value_not_supported" + ); + + let invalid_decimal = build_mysql_dml_request(DmlSqlRequest { + target: target.clone(), + statement: DmlStatement::SingleInsert { + columns: vec![DmlColumn { + data_type_name: "DECIMAL".to_owned(), + ..string_column.clone() + }], + row: DmlRow { + values: vec![DmlValue::Decimal("+1".to_owned())], + }, + }, + }) + .expect_err("the Bridge accepts only plain decimal notation"); + assert_eq!(invalid_decimal.api_error().code, "invalid_database_request"); + + let oversized_identifier = build_mysql_dml_request(DmlSqlRequest { + target, + statement: DmlStatement::SingleInsert { + columns: vec![DmlColumn { + name: "x".repeat(MAX_DML_IDENTIFIER_BYTES + 1), + ..string_column + }], + row: DmlRow { + values: vec![DmlValue::String("value".to_owned())], + }, + }, + }) + .expect_err("identifiers above the Bridge limit must fail before rendering"); + assert_eq!( + oversized_identifier.api_error().code, + "invalid_database_request" + ); + } + #[test] fn grid_insert_ignores_row_number_and_generated_value() { let operation = MysqlResultGridOperation { diff --git a/crates/chat2db-core/src/mysql_schema_diff.rs b/crates/chat2db-core/src/mysql_schema_diff.rs index 8d2698f..913050b 100644 --- a/crates/chat2db-core/src/mysql_schema_diff.rs +++ b/crates/chat2db-core/src/mysql_schema_diff.rs @@ -3,14 +3,13 @@ use std::{ time::Duration, }; -use chat2db_contract::{ - ApiError, CommunitySchemaDiffEndpoint, CommunitySchemaDiffRequest, CommunitySchemaDiffSql, -}; +use chat2db_contract::ApiError; use mysql_async::{Conn, Error as MysqlError, prelude::Queryable}; use crate::{ AppError, AppErrorKind, Application, native_mysql::{finish_connection, open_resolved_connection, resolve_native_connection}, + native_schema_diff_types::{SchemaDiffEndpoint, SchemaDiffRequest, SchemaDiffSql}, }; const SCHEMA_DIFF_TIMEOUT: Duration = Duration::from_secs(30); @@ -64,10 +63,10 @@ struct ForeignKeyDefinition { referenced_table: String, } -/// Pinned Community parity intentionally compares only these existing-table options. +/// Pinned compatibility behavior intentionally compares only these existing-table options. /// /// CHECK constraints, partition definitions, and additional `MySQL` table options are preserved -/// when a missing table is created, but the Community schema-diff contract does not alter them on +/// when a missing table is created, but the retained schema-diff behavior does not alter them on /// an existing table. The SHOW CREATE `AUTO_INCREMENT=N` next counter is runtime state and is /// deliberately excluded from both comparison and generated CREATE statements. #[derive(Debug, Default)] @@ -85,32 +84,27 @@ struct ExistingTableDiff { foreign_key_adds: Vec, } -impl Application { - /// Previews SQL that changes the target `MySQL` namespace to match the source. - /// - /// This method only reads metadata. Generated SQL is never executed automatically. - /// - /// # Errors - /// - /// Returns validation, datasource, connection, metadata, parse, resource-limit, or cleanup - /// errors. - pub async fn preview_mysql_schema_diff( - &self, - request: &CommunitySchemaDiffRequest, - ) -> Result { - validate_endpoint(&request.source, "source")?; - validate_endpoint(&request.target, "target")?; - - let source = load_endpoint_snapshot(self, &request.source).await?; - let target = load_endpoint_snapshot(self, &request.target).await?; - build_schema_diff(&source, &target).map(CommunitySchemaDiffSql::new) - } +/// Previews SQL that changes the target `MySQL` namespace to match the source. +/// +/// This function only reads metadata. Generated SQL is never executed automatically. +/// +/// # Errors +/// +/// Returns validation, datasource, connection, metadata, parse, resource-limit, or cleanup +/// errors. +pub(crate) async fn preview_mysql_schema_diff( + application: &Application, + request: &SchemaDiffRequest, +) -> Result { + validate_endpoint(&request.source, "source")?; + validate_endpoint(&request.target, "target")?; + + let source = load_endpoint_snapshot(application, &request.source).await?; + let target = load_endpoint_snapshot(application, &request.target).await?; + build_schema_diff(&source, &target).map(SchemaDiffSql::new) } -fn validate_endpoint( - endpoint: &CommunitySchemaDiffEndpoint, - role: &'static str, -) -> Result<(), AppError> { +fn validate_endpoint(endpoint: &SchemaDiffEndpoint, role: &'static str) -> Result<(), AppError> { if endpoint.datasource_id.trim().is_empty() { return Err(invalid_schema_diff(format!( "The {role} datasource id is required" @@ -131,7 +125,7 @@ fn validate_endpoint( async fn load_endpoint_snapshot( application: &Application, - endpoint: &CommunitySchemaDiffEndpoint, + endpoint: &SchemaDiffEndpoint, ) -> Result { let resolved = resolve_native_connection(application, &endpoint.datasource_id).await?; let mut conn = open_resolved_connection(&resolved).await?; @@ -315,7 +309,7 @@ fn parse_table_snapshot(table_name: &str, ddl: &str) -> Result AppError { } fn invalid_schema_diff(message: impl Into) -> AppError { - AppError::invalid("invalid_community_schema_diff_request", message) + AppError::invalid("invalid_schema_diff_request", message) } fn schema_diff_resource_limit(message: &'static str) -> AppError { @@ -1475,13 +1469,15 @@ fn schema_diff_resource_limit(message: &'static str) -> AppError { #[cfg(test)] mod tests { - use chat2db_contract::{CommunitySchemaDiffEndpoint, CommunitySchemaDiffRequest}; - - use crate::Application; + use crate::{ + Application, + native_schema_diff_types::{SchemaDiffEndpoint, SchemaDiffRequest}, + }; use super::{ NO_DIFFERENCES_SQL, SchemaSnapshot, ViewSnapshot, build_schema_diff, - canonicalize_column_definition, parse_table_snapshot, rewrite_qualified_catalog, + canonicalize_column_definition, parse_table_snapshot, + preview_mysql_schema_diff as preview_mysql_schema_diff_impl, rewrite_qualified_catalog, }; fn snapshot(database_name: &str) -> SchemaSnapshot { @@ -1492,7 +1488,7 @@ mod tests { } #[test] - fn no_difference_uses_the_pinned_community_comment() { + fn no_difference_uses_the_pinned_compatibility_comment() { let mut snapshot = snapshot("same_db"); snapshot.tables.insert( "items".to_owned(), @@ -1883,17 +1879,14 @@ mod tests { #[tokio::test] async fn invalid_request_fails_before_storage_or_java_access() { - let error = Application::new() - .preview_mysql_schema_diff(&CommunitySchemaDiffRequest { - source: CommunitySchemaDiffEndpoint::default(), - target: CommunitySchemaDiffEndpoint::default(), - }) + let application = Application::new(); + let request = SchemaDiffRequest { + source: SchemaDiffEndpoint::default(), + target: SchemaDiffEndpoint::default(), + }; + let direct_error = preview_mysql_schema_diff_impl(&application, &request) .await - .expect_err("missing source endpoint must fail"); - - assert_eq!( - error.api_error().code, - "invalid_community_schema_diff_request" - ); + .expect_err("direct missing source endpoint must fail"); + assert_eq!(direct_error.api_error().code, "invalid_schema_diff_request"); } } diff --git a/crates/chat2db-core/src/native_administration_types.rs b/crates/chat2db-core/src/native_administration_types.rs new file mode 100644 index 0000000..b2b03e0 --- /dev/null +++ b/crates/chat2db-core/src/native_administration_types.rs @@ -0,0 +1,222 @@ +use std::fmt::{Debug, Formatter}; + +/// Database-neutral administration operation exposed by a native driver. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AdministrationAction { + CreatePrincipal, + AlterCredential, + LockPrincipal, + UnlockPrincipal, + DropPrincipal, + GrantPrivileges, + RevokePrivileges, +} + +/// Identifies a database principal without imposing one database's identity model. +/// +/// `MySQL` uses `name@qualifier` where the qualifier is the host. Databases whose +/// principals have no qualifier leave it unset. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PrincipalRef { + pub(crate) name: String, + pub(crate) qualifier: Option, +} + +/// Database-neutral privilege target category. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum PrivilegeScope { + Global, + Database, + #[allow( + dead_code, + reason = "reserved for PostgreSQL and other schema-scoped administration drivers" + )] + Schema, + Table, +} + +/// Target of a grant or revoke operation. +/// +/// Optional path segments let each driver validate the path required by its +/// own privilege model without leaking that model into the SPI. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PrivilegeTarget { + pub(crate) scope: PrivilegeScope, + pub(crate) database_name: Option, + pub(crate) schema_name: Option, + pub(crate) object_name: Option, +} + +/// Input for previewing or executing one administration operation. +#[derive(Clone, PartialEq, Eq)] +pub(crate) struct AdministrationCommand { + pub(crate) datasource_id: String, + pub(crate) principal: PrincipalRef, + pub(crate) action: AdministrationAction, + pub(crate) target: Option, + pub(crate) privileges: Vec, + pub(crate) grant_option: bool, + pub(crate) credential: Option, + pub(crate) preview_token: Option, +} + +impl Debug for AdministrationCommand { + fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdministrationCommand") + .field("datasource_id", &self.datasource_id) + .field("principal", &self.principal) + .field("action", &self.action) + .field("target", &self.target) + .field("privileges", &self.privileges) + .field("grant_option", &self.grant_option) + .field( + "credential", + &self.credential.as_ref().map(|_| "[REDACTED]"), + ) + .field( + "preview_token", + &self.preview_token.as_ref().map(|_| "[REDACTED]"), + ) + .finish() + } +} + +/// Input for listing the grants assigned to one principal. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PrincipalGrantsRequest { + pub(crate) datasource_id: String, + pub(crate) principal: PrincipalRef, +} + +/// Server and permission capabilities for native administration. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct AdministrationCapability { + pub(crate) database_type: String, + pub(crate) product_name: String, + pub(crate) product_version: Option, + pub(crate) current_principal: Option, + pub(crate) connection_principal: Option, + pub(crate) principal_list_readable: bool, + pub(crate) principal_lock_supported: bool, + pub(crate) editable_privileges: Vec, + pub(crate) message: Option, +} + +/// One principal projected by a native administration driver. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct Principal { + pub(crate) name: String, + pub(crate) qualifier: Option, + pub(crate) display_name: String, + pub(crate) authentication_method: Option, + pub(crate) locked: Option, +} + +/// Stable principal collection. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PrincipalList { + pub(crate) items: Vec, +} + +/// Stable collection of grants returned by a native administration driver. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct PrincipalGrantList { + pub(crate) items: Vec, +} + +/// Masked SQL preview and authorization token for one administration operation. +#[derive(Clone, PartialEq, Eq)] +pub(crate) struct AdministrationPreview { + pub(crate) action: AdministrationAction, + pub(crate) sql: String, + pub(crate) preview_token: String, +} + +impl Debug for AdministrationPreview { + fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdministrationPreview") + .field("action", &self.action) + .field("sql", &"[REDACTED]") + .field("preview_token", &"[REDACTED]") + .finish() + } +} + +/// Result of executing one preview-authorized administration operation. +#[derive(Clone, PartialEq, Eq)] +pub(crate) struct AdministrationExecution { + pub(crate) action: AdministrationAction, + pub(crate) sql: String, + pub(crate) success: bool, + pub(crate) message: Option, + pub(crate) failure_code: Option, + pub(crate) error_code: Option, + pub(crate) sql_state: Option, +} + +impl Debug for AdministrationExecution { + fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("AdministrationExecution") + .field("action", &self.action) + .field("sql", &"[REDACTED]") + .field("success", &self.success) + .field("message", &self.message.as_ref().map(|_| "[REDACTED]")) + .field("failure_code", &self.failure_code) + .field("error_code", &self.error_code) + .field("sql_state", &self.sql_state) + .finish() + } +} + +#[cfg(test)] +mod tests { + use super::{ + AdministrationAction, AdministrationCommand, AdministrationExecution, + AdministrationPreview, PrincipalRef, + }; + + #[test] + fn sensitive_administration_values_are_redacted_from_debug_output() { + let command = AdministrationCommand { + datasource_id: "42".to_owned(), + principal: PrincipalRef { + name: "reader".to_owned(), + qualifier: Some("%".to_owned()), + }, + action: AdministrationAction::CreatePrincipal, + target: None, + privileges: Vec::new(), + grant_option: false, + credential: Some("plain-secret".to_owned()), + preview_token: Some("sensitive-token".to_owned()), + }; + let command_debug = format!("{command:?}"); + assert!(!command_debug.contains("plain-secret")); + assert!(!command_debug.contains("sensitive-token")); + + let preview = AdministrationPreview { + action: AdministrationAction::CreatePrincipal, + sql: "CREATE USER reader IDENTIFIED BY plain-secret".to_owned(), + preview_token: "sensitive-token".to_owned(), + }; + let preview_debug = format!("{preview:?}"); + assert!(!preview_debug.contains("CREATE USER")); + assert!(!preview_debug.contains("sensitive-token")); + + let execution = AdministrationExecution { + action: AdministrationAction::CreatePrincipal, + sql: preview.sql, + success: false, + message: Some("near plain-secret".to_owned()), + failure_code: Some("driver.execute_failed".to_owned()), + error_code: Some(1064), + sql_state: Some("42000".to_owned()), + }; + let execution_debug = format!("{execution:?}"); + assert!(!execution_debug.contains("CREATE USER")); + assert!(!execution_debug.contains("plain-secret")); + } +} diff --git a/crates/chat2db-core/src/native_api_adapter.rs b/crates/chat2db-core/src/native_api_adapter.rs new file mode 100644 index 0000000..6ecc7e6 --- /dev/null +++ b/crates/chat2db-core/src/native_api_adapter.rs @@ -0,0 +1,381 @@ +//! Adapts retained public API contracts to the database-neutral native driver SPI. + +use chat2db_contract::{ + CommunityAccount, CommunityAccountAction, CommunityAccountCapability, + CommunityAccountCommandRequest, CommunityAccountExecution, CommunityAccountGrantList, + CommunityAccountGrantsRequest, CommunityAccountList, CommunityAccountPreview, + CommunityAccountPrivilegeScope, CommunitySchemaDiffEndpoint, CommunitySchemaDiffRequest, + CommunitySchemaDiffSql, +}; + +use crate::{ + AppError, Application, + native_administration_types::{ + AdministrationAction, AdministrationCapability, AdministrationCommand, + AdministrationExecution, AdministrationPreview, Principal, PrincipalGrantList, + PrincipalGrantsRequest, PrincipalList, PrincipalRef, PrivilegeScope, PrivilegeTarget, + }, + native_schema_diff_types::{SchemaDiffEndpoint, SchemaDiffRequest, SchemaDiffSql}, +}; + +impl Application { + /// Returns the retained account-capability response through the selected native driver. + /// + /// # Errors + /// + /// Returns datasource-resolution, capability, connection, query, or cleanup failures. + pub async fn mysql_account_capability( + &self, + datasource_id: &str, + ) -> Result { + let driver = self + .require_native_driver_for_datasource(datasource_id) + .await?; + let administration = driver + .administration() + .ok_or_else(administration_unavailable)?; + administration + .administration_capability(self, datasource_id) + .await + .map(community_account_capability) + } + + /// Lists retained account rows through the selected native driver. + /// + /// # Errors + /// + /// Returns datasource-resolution, capability, connection, query, or cleanup failures. + pub async fn list_mysql_accounts( + &self, + datasource_id: &str, + ) -> Result { + let driver = self + .require_native_driver_for_datasource(datasource_id) + .await?; + let administration = driver + .administration() + .ok_or_else(administration_unavailable)?; + administration + .list_principals(self, datasource_id) + .await + .map(community_account_list) + } + + /// Returns retained grant rows through the selected native driver. + /// + /// # Errors + /// + /// Returns validation, datasource-resolution, capability, query, or cleanup failures. + pub async fn mysql_account_grants( + &self, + request: &CommunityAccountGrantsRequest, + ) -> Result { + let driver = self + .require_native_driver_for_datasource(&request.datasource_id) + .await?; + let administration = driver + .administration() + .ok_or_else(administration_unavailable)?; + administration + .principal_grants(self, &principal_grants_request(request)) + .await + .map(community_account_grants) + } + + /// Builds the retained `MySQL` account preview through the native dialect capability. + /// + /// # Errors + /// + /// Returns capability or account-command validation failures. + pub fn preview_mysql_account( + &self, + request: &CommunityAccountCommandRequest, + ) -> Result { + let driver = self + .native_driver_for_database_type("MYSQL") + .ok_or_else(administration_unavailable)?; + let administration = driver + .administration() + .ok_or_else(administration_unavailable)?; + administration + .preview_administration(self, &administration_command(request)) + .map(community_account_preview) + } + + /// Executes the retained `MySQL` account command through the selected native driver. + /// + /// # Errors + /// + /// Returns validation, authorization, datasource, connection, or cleanup failures. + pub async fn execute_mysql_account( + &self, + request: &CommunityAccountCommandRequest, + ) -> Result { + let driver = self + .require_native_driver_for_datasource(&request.datasource_id) + .await?; + let administration = driver + .administration() + .ok_or_else(administration_unavailable)?; + administration + .execute_administration(self, &administration_command(request)) + .await + .map(community_account_execution) + } + + /// Builds the retained schema-diff response through the source datasource's native driver. + /// + /// # Errors + /// + /// Returns validation, driver-selection, metadata, resource-limit, or cleanup failures. + pub async fn preview_mysql_schema_diff( + &self, + request: &CommunitySchemaDiffRequest, + ) -> Result { + let request = schema_diff_request(request); + validate_schema_diff_selection(&request)?; + + let source_driver = self + .require_native_driver_for_datasource(&request.source.datasource_id) + .await?; + let target_driver = self + .require_native_driver_for_datasource(&request.target.datasource_id) + .await?; + if !source_driver.id().eq_ignore_ascii_case(target_driver.id()) { + return Err(AppError::invalid( + "invalid_community_schema_diff_request", + "Source and target datasources must use the same native driver", + )); + } + let schema_diff = source_driver + .schema_diff() + .ok_or_else(schema_diff_unavailable)?; + schema_diff + .preview_schema_diff(self, &request) + .await + .map(community_schema_diff_sql) + .map_err(community_schema_diff_error) + } +} + +fn administration_unavailable() -> AppError { + AppError::invalid( + "native_administration_capability_not_available", + "The native Rust driver does not implement database administration", + ) +} + +fn schema_diff_unavailable() -> AppError { + AppError::invalid( + "native_schema_diff_capability_not_available", + "The native Rust driver does not implement schema comparison", + ) +} + +fn administration_action(action: CommunityAccountAction) -> AdministrationAction { + match action { + CommunityAccountAction::CreateUser => AdministrationAction::CreatePrincipal, + CommunityAccountAction::AlterPassword => AdministrationAction::AlterCredential, + CommunityAccountAction::LockAccount => AdministrationAction::LockPrincipal, + CommunityAccountAction::UnlockAccount => AdministrationAction::UnlockPrincipal, + CommunityAccountAction::DropUser => AdministrationAction::DropPrincipal, + CommunityAccountAction::GrantPrivilege => AdministrationAction::GrantPrivileges, + CommunityAccountAction::RevokePrivilege => AdministrationAction::RevokePrivileges, + } +} + +fn community_account_action(action: AdministrationAction) -> CommunityAccountAction { + match action { + AdministrationAction::CreatePrincipal => CommunityAccountAction::CreateUser, + AdministrationAction::AlterCredential => CommunityAccountAction::AlterPassword, + AdministrationAction::LockPrincipal => CommunityAccountAction::LockAccount, + AdministrationAction::UnlockPrincipal => CommunityAccountAction::UnlockAccount, + AdministrationAction::DropPrincipal => CommunityAccountAction::DropUser, + AdministrationAction::GrantPrivileges => CommunityAccountAction::GrantPrivilege, + AdministrationAction::RevokePrivileges => CommunityAccountAction::RevokePrivilege, + } +} + +fn principal(user: &str, host: &str) -> PrincipalRef { + PrincipalRef { + name: user.to_owned(), + qualifier: Some(host.to_owned()), + } +} + +fn principal_grants_request(request: &CommunityAccountGrantsRequest) -> PrincipalGrantsRequest { + PrincipalGrantsRequest { + datasource_id: request.datasource_id.clone(), + principal: principal(&request.user, &request.host), + } +} + +fn administration_command(request: &CommunityAccountCommandRequest) -> AdministrationCommand { + AdministrationCommand { + datasource_id: request.datasource_id.clone(), + principal: principal(&request.user, &request.host), + action: administration_action(request.action_type), + target: request.scope.map(|scope| PrivilegeTarget { + scope: match scope { + CommunityAccountPrivilegeScope::Global => PrivilegeScope::Global, + CommunityAccountPrivilegeScope::Database => PrivilegeScope::Database, + CommunityAccountPrivilegeScope::Table => PrivilegeScope::Table, + }, + database_name: request.database_name.clone(), + schema_name: None, + object_name: request.table_name.clone(), + }), + privileges: request.privileges.clone(), + grant_option: request.grant_option, + credential: request.password.clone(), + preview_token: request.preview_token.clone(), + } +} + +fn community_account_capability( + capability: AdministrationCapability, +) -> CommunityAccountCapability { + CommunityAccountCapability { + db_type: capability.database_type, + product_name: capability.product_name, + product_version: capability.product_version, + current_user: capability.current_principal, + connection_user: capability.connection_principal, + account_list_readable: capability.principal_list_readable, + account_lock_supported: capability.principal_lock_supported, + editable_privileges: capability.editable_privileges, + message: capability.message, + } +} + +fn community_account(principal: Principal) -> CommunityAccount { + CommunityAccount { + user: principal.name, + host: principal.qualifier.unwrap_or_default(), + display_name: principal.display_name, + authentication_plugin: principal.authentication_method, + locked: principal.locked, + } +} + +fn community_account_list(accounts: PrincipalList) -> CommunityAccountList { + CommunityAccountList { + items: accounts.items.into_iter().map(community_account).collect(), + } +} + +fn community_account_grants(grants: PrincipalGrantList) -> CommunityAccountGrantList { + CommunityAccountGrantList { + items: grants.items, + } +} + +fn community_account_preview(preview: AdministrationPreview) -> CommunityAccountPreview { + CommunityAccountPreview { + action_type: community_account_action(preview.action), + sql: preview.sql, + preview_token: preview.preview_token, + } +} + +fn community_account_execution(execution: AdministrationExecution) -> CommunityAccountExecution { + CommunityAccountExecution { + action_type: community_account_action(execution.action), + sql: execution.sql, + success: execution.success, + message: execution.message, + failure_code: execution.failure_code, + error_code: execution.error_code, + sql_state: execution.sql_state, + } +} + +fn schema_diff_endpoint(endpoint: &CommunitySchemaDiffEndpoint) -> SchemaDiffEndpoint { + SchemaDiffEndpoint { + datasource_id: endpoint.datasource_id.clone(), + database_name: endpoint.database_name.clone(), + schema_name: endpoint.schema_name.clone(), + } +} + +fn schema_diff_request(request: &CommunitySchemaDiffRequest) -> SchemaDiffRequest { + SchemaDiffRequest { + source: schema_diff_endpoint(&request.source), + target: schema_diff_endpoint(&request.target), + } +} + +fn validate_schema_diff_selection(request: &SchemaDiffRequest) -> Result<(), AppError> { + for (role, endpoint) in [("source", &request.source), ("target", &request.target)] { + if endpoint.datasource_id.trim().is_empty() { + return Err(AppError::invalid( + "invalid_community_schema_diff_request", + format!("The {role} datasource id is required"), + )); + } + if endpoint.database_name.trim().is_empty() { + return Err(AppError::invalid( + "invalid_community_schema_diff_request", + format!("The {role} database name is required"), + )); + } + if endpoint.database_name.contains('\0') { + return Err(AppError::invalid( + "invalid_community_schema_diff_request", + format!("The {role} database name cannot contain NUL"), + )); + } + } + Ok(()) +} + +fn community_schema_diff_sql(sql: SchemaDiffSql) -> CommunitySchemaDiffSql { + CommunitySchemaDiffSql::new(sql.into_inner()) +} + +fn community_schema_diff_error(error: AppError) -> AppError { + let api = error.api_error(); + if api.code == "invalid_schema_diff_request" { + AppError::invalid("invalid_community_schema_diff_request", api.message) + } else { + error + } +} + +#[cfg(test)] +mod tests { + use chat2db_contract::{CommunityAccountAction, CommunityAccountCommandRequest}; + + use super::{administration_command, community_account_action}; + + #[test] + fn account_actions_round_trip_across_the_compatibility_boundary() { + for action in [ + CommunityAccountAction::CreateUser, + CommunityAccountAction::AlterPassword, + CommunityAccountAction::LockAccount, + CommunityAccountAction::UnlockAccount, + CommunityAccountAction::DropUser, + CommunityAccountAction::GrantPrivilege, + CommunityAccountAction::RevokePrivilege, + ] { + let request = CommunityAccountCommandRequest { + datasource_id: "datasource-1".to_owned(), + user: "reader".to_owned(), + host: "%".to_owned(), + action_type: action, + scope: None, + database_name: None, + table_name: None, + privileges: Vec::new(), + grant_option: false, + password: None, + preview_token: None, + }; + assert_eq!( + community_account_action(administration_command(&request).action), + action + ); + } + } +} diff --git a/crates/chat2db-core/src/native_driver.rs b/crates/chat2db-core/src/native_driver.rs index c2ddda8..d757443 100644 --- a/crates/chat2db-core/src/native_driver.rs +++ b/crates/chat2db-core/src/native_driver.rs @@ -9,17 +9,24 @@ use tokio_util::sync::CancellationToken; use crate::{ AppError, Application, datasource_session::ResolvedDatasourceConnection, + native_administration_types::{ + AdministrationCapability, AdministrationCommand, AdministrationExecution, + AdministrationPreview, PrincipalGrantList, PrincipalGrantsRequest, PrincipalList, + }, native_driver_types::{ - ColumnList, DatabaseList, EntityRelationTable, ForeignKeyList, FunctionList, - FunctionMetadata, FunctionParameterList, IndexList, ListColumnsRequest, - ListDatabasesRequest, ListIndexesRequest, ListRoutinesRequest, ListSchemasRequest, - ListTableKeysRequest, ListTablesRequest, ListTriggersRequest, ListViewsRequest, ObjectRef, + BuiltSql, ColumnList, CreateSchemaSqlRequest, DatabaseList, DmlExportTransferRequest, + DmlSqlRequest, EntityRelationTable, ExportArtifact, ForeignKeyList, FunctionList, + FunctionMetadata, FunctionParameterList, ImportTransferRequest, IndexList, + ListColumnsRequest, ListDatabasesRequest, ListIndexesRequest, ListRoutinesRequest, + ListSchemasRequest, ListTableKeysRequest, ListTablesRequest, ListTriggersRequest, + ListViewsRequest, NamespaceSqlRequest, ObjectRef, OtherExportTransferRequest, PrimaryKeyList, ProcedureList, ProcedureMetadata, ProcedureParameterList, RoutineInvocationPreview, RoutineInvocationRequest, RoutineMigrationExecution, - RoutineMigrationRequest, SchemaList, TableList, TableMetadata, TablePreviewAccepted, - TablePreviewRequest, TriggerList, TriggerMetadata, ViewList, + RoutineMigrationRequest, SchemaList, SqlExportTransferRequest, TableList, TableMetadata, + TablePreviewAccepted, TablePreviewRequest, TriggerList, TriggerMetadata, ViewList, }, native_mysql, + native_schema_diff_types::{SchemaDiffRequest, SchemaDiffSql}, operation::CancellationRequest, query::{ DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, PreparedQuery, @@ -246,6 +253,87 @@ pub(crate) trait NativeRoutineDriver: Send + Sync { ) -> Result; } +/// Import and export operations implemented by a native Rust driver. +#[async_trait] +pub(crate) trait NativeTransferDriver: Send + Sync { + async fn import_file( + &self, + application: &Application, + request: ImportTransferRequest, + ) -> Result; + + async fn export_sql_file( + &self, + application: &Application, + request: SqlExportTransferRequest, + ) -> Result; + + async fn export_other_file( + &self, + application: &Application, + request: OtherExportTransferRequest, + ) -> Result; + + async fn export_dml( + &self, + application: &Application, + request: DmlExportTransferRequest, + ) -> Result; +} + +/// Structured SQL builders supplied by one native database dialect. +pub(crate) trait NativeDialectDriver: Send + Sync { + fn build_create_schema(&self, request: CreateSchemaSqlRequest) -> Result; + + fn build_namespace_sql(&self, request: NamespaceSqlRequest) -> Result; + + fn build_dml(&self, request: DmlSqlRequest) -> Result; +} + +/// Database account and role administration implemented by a native driver. +#[async_trait] +pub(crate) trait NativeAdministrationDriver: Send + Sync { + async fn administration_capability( + &self, + application: &Application, + datasource_id: &str, + ) -> Result; + + async fn list_principals( + &self, + application: &Application, + datasource_id: &str, + ) -> Result; + + async fn principal_grants( + &self, + application: &Application, + request: &PrincipalGrantsRequest, + ) -> Result; + + fn preview_administration( + &self, + application: &Application, + request: &AdministrationCommand, + ) -> Result; + + async fn execute_administration( + &self, + application: &Application, + request: &AdministrationCommand, + ) -> Result; +} + +/// Schema-comparison operations implemented by a native Rust driver. +#[async_trait] +pub(crate) trait NativeSchemaDiffDriver: Send + Sync { + async fn preview_schema_diff( + &self, + application: &Application, + request: &SchemaDiffRequest, + ) -> Result; +} + /// Runtime-polymorphic native Rust database driver. /// /// Optional capability accessors allow a driver to participate only in the @@ -279,6 +367,22 @@ pub(crate) trait NativeDriver: Send + Sync { fn routines(&self) -> Option<&dyn NativeRoutineDriver> { None } + + fn transfer(&self) -> Option<&dyn NativeTransferDriver> { + None + } + + fn dialect(&self) -> Option<&dyn NativeDialectDriver> { + None + } + + fn administration(&self) -> Option<&dyn NativeAdministrationDriver> { + None + } + + fn schema_diff(&self) -> Option<&dyn NativeSchemaDiffDriver> { + None + } } /// Immutable registry used to select native implementations at runtime. @@ -411,6 +515,22 @@ impl NativeDriver for MysqlNativeDriver { fn routines(&self) -> Option<&dyn NativeRoutineDriver> { Some(self) } + + fn transfer(&self) -> Option<&dyn NativeTransferDriver> { + Some(self) + } + + fn dialect(&self) -> Option<&dyn NativeDialectDriver> { + Some(self) + } + + fn administration(&self) -> Option<&dyn NativeAdministrationDriver> { + Some(self) + } + + fn schema_diff(&self) -> Option<&dyn NativeSchemaDiffDriver> { + Some(self) + } } #[async_trait] @@ -814,6 +934,109 @@ impl NativeRoutineDriver for MysqlNativeDriver { } } +#[async_trait] +impl NativeTransferDriver for MysqlNativeDriver { + async fn import_file( + &self, + application: &Application, + request: ImportTransferRequest, + ) -> Result { + crate::transfer::mysql_impl::import_file(application, request).await + } + + async fn export_sql_file( + &self, + application: &Application, + request: SqlExportTransferRequest, + ) -> Result { + crate::transfer::mysql_impl::export_sql_file(application, request).await + } + + async fn export_other_file( + &self, + application: &Application, + request: OtherExportTransferRequest, + ) -> Result { + crate::transfer::mysql_impl::export_other_file(application, request).await + } + + async fn export_dml( + &self, + application: &Application, + request: DmlExportTransferRequest, + ) -> Result { + crate::transfer::mysql_impl::export_dml(application, request).await + } +} + +impl NativeDialectDriver for MysqlNativeDriver { + fn build_create_schema(&self, request: CreateSchemaSqlRequest) -> Result { + crate::mysql_ddl::build_mysql_create_schema_request(request) + } + + fn build_namespace_sql(&self, request: NamespaceSqlRequest) -> Result { + crate::mysql_ddl::build_mysql_namespace_request(request) + } + + fn build_dml(&self, request: DmlSqlRequest) -> Result { + crate::mysql_ddl::build_mysql_dml_request(request) + } +} + +#[async_trait] +impl NativeAdministrationDriver for MysqlNativeDriver { + async fn administration_capability( + &self, + application: &Application, + datasource_id: &str, + ) -> Result { + crate::mysql_account::mysql_account_capability(application, datasource_id).await + } + + async fn list_principals( + &self, + application: &Application, + datasource_id: &str, + ) -> Result { + crate::mysql_account::list_mysql_accounts(application, datasource_id).await + } + + async fn principal_grants( + &self, + application: &Application, + request: &PrincipalGrantsRequest, + ) -> Result { + crate::mysql_account::mysql_account_grants(application, request).await + } + + fn preview_administration( + &self, + application: &Application, + request: &AdministrationCommand, + ) -> Result { + crate::mysql_account::preview_mysql_account(application, request) + } + + async fn execute_administration( + &self, + application: &Application, + request: &AdministrationCommand, + ) -> Result { + crate::mysql_account::execute_mysql_account(application, request).await + } +} + +#[async_trait] +impl NativeSchemaDiffDriver for MysqlNativeDriver { + async fn preview_schema_diff( + &self, + application: &Application, + request: &SchemaDiffRequest, + ) -> Result { + crate::mysql_schema_diff::preview_mysql_schema_diff(application, request).await + } +} + #[cfg(test)] mod tests { use super::*; @@ -857,6 +1080,10 @@ mod tests { fn connection(&self) -> &dyn NativeConnectionDriver { self } + + fn dialect(&self) -> Option<&dyn NativeDialectDriver> { + Some(self) + } } #[async_trait] @@ -869,6 +1096,29 @@ mod tests { } } + impl NativeDialectDriver for FakePostgresDriver { + fn build_create_schema( + &self, + request: CreateSchemaSqlRequest, + ) -> Result { + Ok(BuiltSql { + sql: format!("fake-postgres:create-schema:{}", request.schema.name), + }) + } + + fn build_namespace_sql(&self, _request: NamespaceSqlRequest) -> Result { + Ok(BuiltSql { + sql: "fake-postgres:namespace".to_owned(), + }) + } + + fn build_dml(&self, _request: DmlSqlRequest) -> Result { + Ok(BuiltSql { + sql: "fake-postgres:dml".to_owned(), + }) + } + } + #[test] fn registry_selects_runtime_driver_by_database_type_and_driver_id() { let registry = NativeDriverRegistry::try_new(vec![Arc::new(FakePostgresDriver)]) @@ -890,6 +1140,25 @@ mod tests { ); } + #[tokio::test] + async fn application_dispatches_a_postgres_capability_through_the_registry() { + let registry = NativeDriverRegistry::try_new(vec![Arc::new(FakePostgresDriver)]) + .expect("registry is valid"); + let application = Application::with_native_drivers_for_test(registry); + + let built = application + .build_community_namespace_sql(chat2db_contract::BuildCommunityNamespaceSqlRequest { + database_type: "POSTGRESQL".to_owned(), + operation: chat2db_contract::CommunityNamespaceSqlOperation::UseDatabase { + database_name: "inventory".to_owned(), + }, + }) + .await + .expect("application must dispatch to the fake PostgreSQL capability"); + + assert_eq!(built.sql, "fake-postgres:namespace"); + } + #[test] fn registry_uses_managed_descriptor_aliases_without_owning_driver_jars() { let registry = NativeDriverRegistry::try_new(vec![Arc::new(FakePostgresDriver)]) diff --git a/crates/chat2db-core/src/native_driver_types.rs b/crates/chat2db-core/src/native_driver_types.rs index 874aef6..1cf977a 100644 --- a/crates/chat2db-core/src/native_driver_types.rs +++ b/crates/chat2db-core/src/native_driver_types.rs @@ -6,12 +6,13 @@ use chat2db_contract::{ CommunityRoutineMigrationRequest, CommunitySchemaList, CommunityTable, CommunityTableColumnList, CommunityTableIndexList, CommunityTableList, CommunityTablePreviewAccepted, CommunityTrigger, CommunityTriggerList, CommunityViewList, - GetCommunityFunctionRequest, GetCommunityProcedureRequest, GetCommunityTriggerRequest, - ListCommunityColumnsRequest, ListCommunityDatabasesRequest, ListCommunityFunctionsRequest, - ListCommunityIndexesRequest, ListCommunityProceduresRequest, ListCommunitySchemasRequest, - ListCommunityTableKeysRequest, ListCommunityTablesRequest, ListCommunityTriggersRequest, - ListCommunityViewsRequest, PreviewCommunityRoutineInvocationRequest, - StartCommunityTablePreviewRequest, + DmlExportRequest, GetCommunityFunctionRequest, GetCommunityProcedureRequest, + GetCommunityTriggerRequest, ImportFileRequest, ListCommunityColumnsRequest, + ListCommunityDatabasesRequest, ListCommunityFunctionsRequest, ListCommunityIndexesRequest, + ListCommunityProceduresRequest, ListCommunitySchemasRequest, ListCommunityTableKeysRequest, + ListCommunityTablesRequest, ListCommunityTriggersRequest, ListCommunityViewsRequest, + OtherFileExportRequest, PreviewCommunityRoutineInvocationRequest, SqlFileExportRequest, + StartCommunityTablePreviewRequest, TransferArtifact, }; pub(crate) type DatabaseList = CommunityDatabaseList; @@ -35,6 +36,145 @@ pub(crate) type EntityRelationTable = CommunityErTable; pub(crate) type TablePreviewAccepted = CommunityTablePreviewAccepted; pub(crate) type RoutineInvocationPreview = CommunityRoutineInvocationPreview; pub(crate) type RoutineMigrationExecution = CommunityRoutineMigrationExecution; +pub(crate) type ImportTransferRequest = ImportFileRequest; +pub(crate) type SqlExportTransferRequest = SqlFileExportRequest; +pub(crate) type OtherExportTransferRequest = OtherFileExportRequest; +pub(crate) type DmlExportTransferRequest = DmlExportRequest; +pub(crate) type ExportArtifact = TransferArtifact; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct BuiltSql { + pub(crate) sql: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct DatabaseDefinition { + pub(crate) name: String, + pub(crate) comment: String, + pub(crate) charset: String, + pub(crate) collation: String, + pub(crate) owner: String, + pub(crate) system: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct SchemaDefinition { + pub(crate) database_name: String, + pub(crate) name: String, + pub(crate) comment: String, + pub(crate) owner: String, + pub(crate) system: bool, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct CreateSchemaSqlRequest { + pub(crate) schema: SchemaDefinition, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct NamespaceSqlRequest { + pub(crate) operation: NamespaceSqlOperation, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum NamespaceSqlOperation { + CreateDatabase { + database: DatabaseDefinition, + }, + AlterDatabase { + old_database: DatabaseDefinition, + new_database: DatabaseDefinition, + }, + DropDatabase { + database_name: String, + }, + UseDatabase { + database_name: String, + }, + CreateSchema { + schema: SchemaDefinition, + }, + AlterSchema { + old_schema_name: String, + new_schema_name: String, + }, + DropSchema { + schema_name: String, + }, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct DmlSqlRequest { + pub(crate) target: DmlTarget, + pub(crate) statement: DmlStatement, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +#[allow( + clippy::struct_field_names, + reason = "qualified table segments are clearer with their database, schema, and table suffixes" +)] +pub(crate) struct DmlTarget { + pub(crate) database_name: Option, + pub(crate) schema_name: Option, + pub(crate) table_name: String, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct DmlColumn { + pub(crate) name: String, + pub(crate) data_type_name: String, + pub(crate) precision: Option, + pub(crate) scale: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum DmlTemporalKind { + Date, + Time, + LocalDatetime, + OffsetDatetime, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum DmlValue { + Null, + String(String), + Decimal(String), + Boolean(bool), + Temporal { + kind: DmlTemporalKind, + iso8601: String, + }, + Binary(Vec), +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct DmlRow { + pub(crate) values: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct DmlAssignment { + pub(crate) column: DmlColumn, + pub(crate) value: DmlValue, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum DmlStatement { + SingleInsert { + columns: Vec, + row: DmlRow, + }, + MultiInsert { + columns: Vec, + rows: Vec, + }, + Update { + assignments: Vec, + predicates: Vec, + }, +} #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) struct MetadataScope { diff --git a/crates/chat2db-core/src/native_schema_diff_types.rs b/crates/chat2db-core/src/native_schema_diff_types.rs new file mode 100644 index 0000000..430c475 --- /dev/null +++ b/crates/chat2db-core/src/native_schema_diff_types.rs @@ -0,0 +1,28 @@ +/// One source or target namespace selected for native schema comparison. +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct SchemaDiffEndpoint { + pub(crate) datasource_id: String, + pub(crate) database_name: String, + pub(crate) schema_name: String, +} + +/// Desired source and migration target for native schema comparison. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct SchemaDiffRequest { + pub(crate) source: SchemaDiffEndpoint, + pub(crate) target: SchemaDiffEndpoint, +} + +/// SQL generated by a native schema-comparison capability. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct SchemaDiffSql(String); + +impl SchemaDiffSql { + pub(crate) fn new(sql: impl Into) -> Self { + Self(sql.into()) + } + + pub(crate) fn into_inner(self) -> String { + self.0 + } +} diff --git a/crates/chat2db-core/src/query.rs b/crates/chat2db-core/src/query.rs index c1bcc66..42e07a0 100644 --- a/crates/chat2db-core/src/query.rs +++ b/crates/chat2db-core/src/query.rs @@ -224,15 +224,6 @@ impl Application { .await } - pub(crate) async fn execute_mysql_read_console( - &self, - request: MysqlConsoleRequest, - cancellation: MysqlConsoleCancellation, - ) -> Result, AppError> { - self.execute_native_read_console(request, cancellation) - .await - } - async fn execute_native_console_with_mode( &self, request: NativeConsoleRequest, diff --git a/crates/chat2db-core/src/transfer/mod.rs b/crates/chat2db-core/src/transfer/mod.rs index 4f51e05..46a33ff 100644 --- a/crates/chat2db-core/src/transfer/mod.rs +++ b/crates/chat2db-core/src/transfer/mod.rs @@ -1,13 +1,11 @@ mod class_generation; mod format; mod mysql; +pub(crate) mod mysql_impl; use std::{ - collections::HashMap, - fmt::Write as _, - fs::File, - path::{Path, PathBuf}, - time::Duration, + collections::HashMap, fmt::Write as _, fs::File, future::Future, path::PathBuf, pin::Pin, + sync::Arc, time::Duration, }; use chat2db_contract::{ @@ -21,23 +19,52 @@ use chat2db_storage::{ }; use tokio::{ sync::{Mutex, oneshot}, - task::JoinHandle, + task::{AbortHandle, JoinHandle}, }; use tokio_util::sync::CancellationToken; -use crate::{AppError, Application, native_mysql, storage_call}; +use crate::{AppError, Application, storage_call}; const MAX_TASK_PAGE_SIZE: u32 = 100; +const MAX_TRANSFER_FAILURE_MESSAGE_BYTES: usize = 64 * 1024; +const TRANSFER_FAILURE_TRUNCATION_SUFFIX: &str = "\n[truncated]"; +const TERMINAL_RETRY_INITIAL_DELAY: Duration = Duration::from_millis(25); +const TERMINAL_RETRY_MAX_DELAY: Duration = Duration::from_secs(1); +const TERMINAL_RECOVERY_MESSAGE: &str = + "Transfer terminal state will be recovered when the runtime restarts"; pub(crate) struct TransferTaskHub { tasks: Mutex>, } struct ActiveTransferTask { - cancellation: CancellationToken, + control: TransferTaskControl, handle: JoinHandle<()>, } +struct AbortTaskOnDrop(AbortHandle); + +impl Drop for AbortTaskOnDrop { + fn drop(&mut self) { + self.0.abort(); + } +} + +#[derive(Clone)] +struct TransferTaskControl { + cancellation: CancellationToken, + terminal_gate: Arc>, +} + +impl TransferTaskControl { + fn new() -> Self { + Self { + cancellation: CancellationToken::new(), + terminal_gate: Arc::new(Mutex::new(())), + } + } +} + pub struct TransferArtifactDownload { pub artifact: TransferArtifact, pub path: PathBuf, @@ -65,37 +92,34 @@ impl TransferTaskHub { async fn insert( &self, task_id: i64, - cancellation: CancellationToken, + control: TransferTaskControl, handle: JoinHandle<()>, ) -> Option { - self.tasks.lock().await.insert( - task_id, - ActiveTransferTask { - cancellation, - handle, - }, - ) + self.tasks + .lock() + .await + .insert(task_id, ActiveTransferTask { control, handle }) } async fn remove(&self, task_id: i64) { self.tasks.lock().await.remove(&task_id); } - async fn cancel(&self, task_id: i64) -> bool { - let tasks = self.tasks.lock().await; - let Some(task) = tasks.get(&task_id) else { - return false; - }; - task.cancellation.cancel(); - true + async fn control(&self, task_id: i64) -> Option { + self.tasks + .lock() + .await + .get(&task_id) + .map(|task| task.control.clone()) } - async fn cancel_all(&self) -> Vec { - let tasks = self.tasks.lock().await; - for task in tasks.values() { - task.cancellation.cancel(); - } - tasks.keys().copied().collect() + async fn controls(&self) -> Vec<(i64, TransferTaskControl)> { + self.tasks + .lock() + .await + .iter() + .map(|(task_id, task)| (*task_id, task.control.clone())) + .collect() } async fn take_all(&self) -> HashMap { @@ -103,7 +127,7 @@ impl TransferTaskHub { } } -pub(super) struct TransferContext { +pub(crate) struct TransferContext { storage: Storage, task_id: i64, cancellation: CancellationToken, @@ -118,11 +142,11 @@ impl TransferContext { } } - pub(super) fn cancellation(&self) -> &CancellationToken { + pub(crate) fn cancellation(&self) -> &CancellationToken { &self.cancellation } - pub(super) fn check_cancelled(&self) -> Result<(), TransferRunError> { + pub(crate) fn check_cancelled(&self) -> Result<(), TransferRunError> { if self.cancellation.is_cancelled() { Err(TransferRunError::Cancelled) } else { @@ -130,7 +154,7 @@ impl TransferContext { } } - pub(super) fn begin_artifact( + pub(crate) fn begin_artifact( &self, file_name: &str, media_type: &str, @@ -150,7 +174,7 @@ impl TransferContext { .map_err(TransferRunError::from) } - pub(super) async fn progress( + pub(crate) async fn progress( &self, current: u64, total: Option, @@ -170,13 +194,13 @@ impl TransferContext { } } -pub(super) enum TransferRunError { +pub(crate) enum TransferRunError { Cancelled, Failed(AppError), } impl TransferRunError { - pub(super) fn into_app_error(self) -> AppError { + pub(crate) fn into_app_error(self) -> AppError { match self { Self::Cancelled => { AppError::unavailable("transfer_cancelled", "The transfer operation was cancelled") @@ -198,19 +222,135 @@ impl From for TransferRunError { } } -pub(super) enum TaskCompletion { +pub(crate) enum TaskCompletion { WithoutArtifact(String), - Artifact(TransferArtifactRecord), + Artifact(PendingTransferArtifact), +} + +type TransferJobFuture = + Pin> + Send + 'static>>; + +type TransferJobRunner = + Box TransferJobFuture + Send + 'static>; + +type TransferArtifactFuture = + Pin> + Send + 'static>>; + +type TransferArtifactFinalizer = Box TransferArtifactFuture + Send + 'static>; + +pub(crate) struct PendingTransferArtifact { + finalizer: TransferArtifactFinalizer, +} + +impl PendingTransferArtifact { + pub(crate) fn new(finalizer: Finalizer) -> Self + where + Finalizer: FnOnce() -> FinalizerFuture + Send + 'static, + FinalizerFuture: Future> + Send + 'static, + { + Self { + finalizer: Box::new(move || Box::pin(finalizer())), + } + } + + async fn finalize(self) -> Result { + (self.finalizer)().await + } +} + +#[derive(Clone, Copy)] +pub(crate) enum TransferJobKind { + ImportFile, + ExportSql, + ExportFile, +} + +pub(crate) struct TransferJobSpec { + datasource_id: String, + database_name: String, + schema_name: String, + table_name: Option, + kind: TransferJobKind, + task_name: String, + runner: TransferJobRunner, +} + +impl TransferJobSpec { + pub(crate) fn new( + datasource_id: String, + database_name: String, + schema_name: String, + table_name: Option, + kind: TransferJobKind, + task_name: String, + runner: Runner, + ) -> Self + where + Runner: FnOnce(Application, TransferContext) -> RunnerFuture + Send + 'static, + RunnerFuture: Future> + Send + 'static, + { + Self { + datasource_id, + database_name, + schema_name, + table_name, + kind, + task_name, + runner: Box::new(move |application, context| Box::pin(runner(application, context))), + } + } + + fn into_parts(self) -> (CreateTransferTask, TransferJobRunner) { + let kind = match self.kind { + TransferJobKind::ImportFile => StoredTransferTaskKind::ImportFile, + TransferJobKind::ExportSql => StoredTransferTaskKind::ExportSql, + TransferJobKind::ExportFile => StoredTransferTaskKind::ExportFile, + }; + ( + CreateTransferTask { + datasource_id: self.datasource_id, + database_name: self.database_name, + schema_name: self.schema_name, + table_name: self.table_name, + kind, + task_name: self.task_name, + }, + self.runner, + ) + } } -enum TransferJob { - Import(ImportFileRequest), - SqlExport(SqlFileExportRequest), - OtherExport(OtherFileExportRequest), +enum TransferTerminalState { + Succeeded(String), + Artifact(PendingTransferArtifact), + Cancelled(String), + Failed(String), } impl Application { - /// Starts a durable native-MySQL CSV, XLS, XLSX, or SQL import task. + /// Starts a durable import task through the datasource's native driver. + /// + /// # Errors + /// + /// Returns validation, datasource, storage, or runtime-shutdown failures. + pub async fn import_file( + &self, + request: ImportFileRequest, + ) -> Result { + let driver = self + .require_native_driver_for_datasource(&request.datasource_id) + .await?; + let transfer = driver.transfer().ok_or_else(|| { + AppError::invalid( + "native_transfer_capability_not_available", + "The native Rust driver does not implement import and export operations", + ) + })?; + let spec = transfer.import_file(self, request).await?; + self.start_transfer_job(spec).await + } + + /// Retained `MySQL` compatibility name for [`Self::import_file`]. /// /// # Errors /// @@ -219,27 +359,32 @@ impl Application { &self, request: ImportFileRequest, ) -> Result { - validate_import_request(&request)?; - native_mysql::resolve_native_connection(self, &request.datasource_id).await?; - let file_name = Path::new(&request.file_path) - .file_name() - .and_then(|value| value.to_str()) - .unwrap_or("file"); - self.start_transfer_job( - CreateTransferTask { - datasource_id: request.datasource_id.clone(), - database_name: request.database_name.clone(), - schema_name: request.schema_name.clone(), - table_name: request.table_name.clone(), - kind: StoredTransferTaskKind::ImportFile, - task_name: format!("Import {file_name}"), - }, - TransferJob::Import(request), - ) - .await + self.import_file(request).await + } + + /// Starts a durable SQL dump export through the datasource's native driver. + /// + /// # Errors + /// + /// Returns validation, datasource, storage, or runtime-shutdown failures. + pub async fn export_sql_file( + &self, + request: SqlFileExportRequest, + ) -> Result { + let driver = self + .require_native_driver_for_datasource(&request.datasource_id) + .await?; + let transfer = driver.transfer().ok_or_else(|| { + AppError::invalid( + "native_transfer_capability_not_available", + "The native Rust driver does not implement import and export operations", + ) + })?; + let spec = transfer.export_sql_file(self, request).await?; + self.start_transfer_job(spec).await } - /// Starts a durable native-MySQL SQL dump export task. + /// Retained `MySQL` compatibility name for [`Self::export_sql_file`]. /// /// # Errors /// @@ -248,27 +393,32 @@ impl Application { &self, request: SqlFileExportRequest, ) -> Result { - validate_transfer_scope( - &request.datasource_id, - &request.database_name, - request.export_path.as_deref(), - )?; - native_mysql::resolve_native_connection(self, &request.datasource_id).await?; - self.start_transfer_job( - CreateTransferTask { - datasource_id: request.datasource_id.clone(), - database_name: request.database_name.clone(), - schema_name: request.schema_name.clone(), - table_name: single_table(&request.table_names), - kind: StoredTransferTaskKind::ExportSql, - task_name: format!("Export SQL {}", request.database_name), - }, - TransferJob::SqlExport(request), - ) - .await + self.export_sql_file(request).await + } + + /// Starts a durable table-file export through the datasource's native driver. + /// + /// # Errors + /// + /// Returns validation, datasource, storage, or runtime-shutdown failures. + pub async fn export_other_file( + &self, + request: OtherFileExportRequest, + ) -> Result { + let driver = self + .require_native_driver_for_datasource(&request.datasource_id) + .await?; + let transfer = driver.transfer().ok_or_else(|| { + AppError::invalid( + "native_transfer_capability_not_available", + "The native Rust driver does not implement import and export operations", + ) + })?; + let spec = transfer.export_other_file(self, request).await?; + self.start_transfer_job(spec).await } - /// Starts a durable native-MySQL table file export task. + /// Retained `MySQL` compatibility name for [`Self::export_other_file`]. /// /// # Errors /// @@ -277,34 +427,7 @@ impl Application { &self, request: OtherFileExportRequest, ) -> Result { - validate_transfer_scope( - &request.datasource_id, - &request.database_name, - request.export_path.as_deref(), - )?; - if request.table_names.is_empty() { - return Err(AppError::invalid( - "missing_export_tables", - "tableNames must contain at least one table", - )); - } - native_mysql::resolve_native_connection(self, &request.datasource_id).await?; - self.start_transfer_job( - CreateTransferTask { - datasource_id: request.datasource_id.clone(), - database_name: request.database_name.clone(), - schema_name: request.schema_name.clone(), - table_name: single_table(&request.table_names), - kind: StoredTransferTaskKind::ExportFile, - task_name: format!( - "Export {} {} table(s)", - request.format.extension().to_ascii_uppercase(), - request.table_names.len() - ), - }, - TransferJob::OtherExport(request), - ) - .await + self.export_other_file(request).await } /// Lists retained transfer tasks newest first. @@ -324,7 +447,7 @@ impl Application { /// Lists retained transfer tasks after applying an optional status set. /// /// An empty status set selects every task. Filtering happens before - /// pagination so legacy Community task tabs keep accurate totals. + /// pagination so legacy task tabs keep accurate totals. /// /// # Errors /// @@ -389,10 +512,10 @@ impl Application { /// Returns not-found or durable-storage failures. pub async fn stop_transfer_task(&self, task_id: i64) -> Result<(), AppError> { let storage = self.require_storage()?; - let changed = storage_call(move || storage.request_transfer_cancel(task_id)).await?; - if changed { - self.inner.transfer_tasks.cancel(task_id).await; + if let Some(control) = self.inner.transfer_tasks.control(task_id).await { + control.cancellation.cancel(); } + persist_cancel_request(&storage, task_id).await?; Ok(()) } @@ -442,14 +565,32 @@ impl Application { /// # Errors /// /// Returns SQL analysis, datasource, query, format, or storage failures. + pub async fn export_dml( + &self, + request: DmlExportRequest, + ) -> Result { + let driver = self + .require_native_driver_for_datasource(&request.datasource_id) + .await?; + let transfer = driver.transfer().ok_or_else(|| { + AppError::invalid( + "native_transfer_capability_not_available", + "The native Rust driver does not implement import and export operations", + ) + })?; + transfer.export_dml(self, request).await + } + + /// Retained `MySQL` compatibility name for [`Self::export_dml`]. + /// + /// # Errors + /// + /// Returns SQL analysis, datasource, query, format, or storage failures. pub async fn export_mysql_dml( &self, request: DmlExportRequest, ) -> Result { - native_mysql::resolve_native_connection(self, &request.datasource_id).await?; - mysql::export_dml(self, request) - .await - .map(transfer_artifact) + self.export_dml(request).await } /// Generates `MyBatis` Plus entity, Mapper, and Mapper XML files from native `MySQL` metadata. @@ -461,7 +602,8 @@ impl Application { &self, request: GenerateMysqlClassRequest, ) -> Result { - native_mysql::resolve_native_connection(self, &request.datasource_id).await?; + self.require_native_driver_for_datasource(&request.datasource_id) + .await?; class_generation::generate(self, request).await } @@ -474,16 +616,16 @@ impl Application { &self, request: GenerateMysqlClassRequest, ) -> Result { - native_mysql::resolve_native_connection(self, &request.datasource_id).await?; + self.require_native_driver_for_datasource(&request.datasource_id) + .await?; class_generation::generate_archive(self, request) .await .map(transfer_artifact) } - async fn start_transfer_job( + pub(crate) async fn start_transfer_job( &self, - task: CreateTransferTask, - job: TransferJob, + spec: TransferJobSpec, ) -> Result { let accepting_work = self.inner.accepting_work.lock().await; if !*accepting_work { @@ -492,6 +634,7 @@ impl Application { "The Chat2DB runtime is shutting down", )); } + let (task, runner) = spec.into_parts(); let storage = self.require_storage()?; let task_record = storage_call({ let storage = storage.clone(); @@ -499,23 +642,57 @@ impl Application { }) .await?; let task_id = task_record.id; - let cancellation = CancellationToken::new(); - let run_cancellation = cancellation.clone(); + let control = TransferTaskControl::new(); + let run_control = control.clone(); + let run_storage = storage.clone(); let application = self.clone(); let (registered, wait_for_registration) = oneshot::channel(); let handle = tokio::spawn(async move { if wait_for_registration.await.is_err() { return; } - application - .run_transfer_task(task_id, job, run_cancellation) - .await; + let run_application = application.clone(); + let recovery_storage = run_storage.clone(); + let recovery_control = run_control.clone(); + let mut worker = tokio::spawn(async move { + run_application + .run_transfer_task(task_id, runner, run_storage, run_control) + .await; + }); + let _abort_worker = AbortTaskOnDrop(worker.abort_handle()); + match (&mut worker).await { + Ok(()) => {} + Err(error) if error.is_panic() => { + tracing::error!(task_id, "transfer worker panicked"); + finalize_transfer_task( + &recovery_storage, + task_id, + &recovery_control, + TransferTerminalState::Failed("Transfer worker panicked".to_owned()), + None, + ) + .await; + } + Err(error) => { + tracing::error!(task_id, %error, "transfer worker stopped unexpectedly"); + finalize_transfer_task( + &recovery_storage, + task_id, + &recovery_control, + TransferTerminalState::Failed( + "Transfer worker stopped unexpectedly".to_owned(), + ), + None, + ) + .await; + } + } application.inner.transfer_tasks.remove(task_id).await; }); let replaced = self .inner .transfer_tasks - .insert(task_id, cancellation, handle) + .insert(task_id, control.clone(), handle) .await; debug_assert!(replaced.is_none(), "transfer task ids must be unique"); if registered.send(()).is_err() { @@ -528,11 +705,15 @@ impl Application { .remove(&task_id) { task.handle.abort(); + let _ = task.handle.await; } - let storage = storage.clone(); - let _ = storage_call(move || { - storage.fail_transfer_task(task_id, "Transfer task registration failed") - }) + finalize_transfer_task( + &storage, + task_id, + &control, + TransferTerminalState::Failed("Transfer task registration failed".to_owned()), + None, + ) .await; return Err(AppError::internal()); } @@ -543,67 +724,57 @@ impl Application { async fn run_transfer_task( &self, task_id: i64, - job: TransferJob, - cancellation: CancellationToken, + runner: TransferJobRunner, + storage: Storage, + control: TransferTaskControl, ) { - let Some(storage) = self.storage().cloned() else { - return; - }; - if cancellation.is_cancelled() { - let _ = storage_call(move || storage.request_transfer_cancel(task_id)).await; + if control.cancellation.is_cancelled() { + finalize_transfer_task( + &storage, + task_id, + &control, + TransferTerminalState::Cancelled("Transfer cancelled before startup".to_owned()), + None, + ) + .await; return; } let start_storage = storage.clone(); if let Err(error) = storage_call(move || start_storage.start_transfer_task(task_id)).await { - if !cancellation.is_cancelled() { - tracing::warn!(task_id, %error, "transfer task could not enter running state"); - } + tracing::warn!(task_id, %error, "transfer task could not enter running state"); + finalize_transfer_task( + &storage, + task_id, + &control, + TransferTerminalState::Failed("Transfer task could not start".to_owned()), + None, + ) + .await; return; } - let context = TransferContext::new(storage.clone(), task_id, cancellation.clone()); - let result = match job { - TransferJob::Import(request) => mysql::import_file(self, request, &context).await, - TransferJob::SqlExport(request) => mysql::export_sql(self, request, &context).await, - TransferJob::OtherExport(request) => mysql::export_other(self, request, &context).await, - }; - match result { + let context = TransferContext::new(storage.clone(), task_id, control.cancellation.clone()); + let result = runner(self.clone(), context).await; + let terminal = match result { Ok(TaskCompletion::WithoutArtifact(message)) => { - let complete_storage = storage.clone(); - if let Err(error) = - storage_call(move || complete_storage.complete_transfer_task(task_id, &message)) - .await - { - tracing::warn!(task_id, %error, "transfer task completion could not be persisted"); - } - } - Ok(TaskCompletion::Artifact(artifact)) => { - debug_assert_eq!(artifact.task_id, Some(task_id)); + TransferTerminalState::Succeeded(message) } + Ok(TaskCompletion::Artifact(artifact)) => TransferTerminalState::Artifact(artifact), Err(TransferRunError::Cancelled) => { - let cancel_storage = storage.clone(); - let _ = storage_call(move || { - cancel_storage.cancel_transfer_task(task_id, "Transfer cancelled by request") - }) - .await; + TransferTerminalState::Cancelled("Transfer cancelled by request".to_owned()) } Err(TransferRunError::Failed(error)) => { - let message = error.api_error().message; - tracing::warn!(task_id, code = %error.api_error().code, "transfer task failed"); - let fail_storage = storage.clone(); - let _ = - storage_call(move || fail_storage.fail_transfer_task(task_id, &message)).await; + let error = error.api_error(); + tracing::warn!(task_id, code = %error.code, "transfer task failed"); + TransferTerminalState::Failed(truncate_transfer_failure_message(error.message)) } - } + }; + finalize_transfer_task(&storage, task_id, &control, terminal, None).await; } pub(crate) async fn begin_transfer_shutdown(&self) { - let task_ids = self.inner.transfer_tasks.cancel_all().await; - let Some(storage) = self.storage().cloned() else { - return; - }; - for task_id in task_ids { - let storage = storage.clone(); - let _ = storage_call(move || storage.request_transfer_cancel(task_id)).await; + let controls = self.inner.transfer_tasks.controls().await; + for (_, control) in &controls { + control.cancellation.cancel(); } } @@ -612,23 +783,199 @@ impl Application { let deadline = tokio::time::Instant::now() + timeout; let storage = self.storage().cloned(); for (task_id, mut task) in tasks { - let terminal_message = match tokio::time::timeout_at(deadline, &mut task.handle).await { + let terminal = match tokio::time::timeout_at(deadline, &mut task.handle).await { Ok(Ok(())) => None, - Ok(Err(_)) => Some("Transfer worker stopped unexpectedly"), + Ok(Err(error)) => { + tracing::error!(task_id, %error, "transfer worker monitor stopped unexpectedly"); + Some(TransferTerminalState::Failed( + "Transfer worker stopped unexpectedly".to_owned(), + )) + } Err(_) => { task.handle.abort(); - Some("Transfer stopped during runtime shutdown") + let _ = task.handle.await; + Some(TransferTerminalState::Cancelled( + "Transfer stopped during runtime shutdown".to_owned(), + )) + } + }; + if let (Some(storage), Some(terminal)) = (storage.as_ref(), terminal) { + let finalized = finalize_transfer_task( + storage, + task_id, + &task.control, + terminal, + Some(deadline), + ) + .await; + if !finalized { + tracing::error!(task_id, "{TERMINAL_RECOVERY_MESSAGE}"); + } + } + } + } +} + +async fn persist_cancel_request(storage: &Storage, task_id: i64) -> Result { + let storage = storage.clone(); + storage_call(move || storage.request_transfer_cancel(task_id)).await +} + +async fn load_transfer_task( + storage: &Storage, + task_id: i64, +) -> Result, AppError> { + let storage = storage.clone(); + storage_call(move || storage.get_transfer_task(task_id)).await +} + +const fn stored_transfer_is_terminal(status: StoredTransferTaskStatus) -> bool { + matches!( + status, + StoredTransferTaskStatus::Succeeded + | StoredTransferTaskStatus::Failed + | StoredTransferTaskStatus::Cancelled + | StoredTransferTaskStatus::Interrupted + ) +} + +async fn finalize_transfer_task( + storage: &Storage, + task_id: i64, + control: &TransferTaskControl, + mut terminal: TransferTerminalState, + retry_deadline: Option, +) -> bool { + let mut attempt = 0_u32; + let mut retry_delay = TERMINAL_RETRY_INITIAL_DELAY; + loop { + attempt = attempt.saturating_add(1); + let finalization = async { + let _terminal = control.terminal_gate.lock().await; + finalize_transfer_once(storage, task_id, control, &mut terminal).await + }; + let result = if let Some(deadline) = retry_deadline { + match tokio::time::timeout_at(deadline, finalization).await { + Ok(result) => result, + Err(_) => return false, + } + } else { + finalization.await + }; + match result { + Ok(()) => return true, + Err(error) => { + tracing::warn!( + task_id, + attempt, + %error, + "transfer terminal state could not be persisted" + ); + if retry_deadline.is_some_and(|deadline| tokio::time::Instant::now() >= deadline) { + return false; } + } + } + let delay = retry_deadline.map_or(retry_delay, |deadline| { + retry_delay.min(deadline.saturating_duration_since(tokio::time::Instant::now())) + }); + if !delay.is_zero() { + tokio::time::sleep(delay).await; + } + retry_delay = retry_delay.saturating_mul(2).min(TERMINAL_RETRY_MAX_DELAY); + } +} + +fn truncate_transfer_failure_message(mut message: String) -> String { + if message.len() <= MAX_TRANSFER_FAILURE_MESSAGE_BYTES { + return message; + } + let mut boundary = + MAX_TRANSFER_FAILURE_MESSAGE_BYTES - TRANSFER_FAILURE_TRUNCATION_SUFFIX.len(); + while !message.is_char_boundary(boundary) { + boundary -= 1; + } + message.truncate(boundary); + message.push_str(TRANSFER_FAILURE_TRUNCATION_SUFFIX); + message +} + +async fn finalize_transfer_once( + storage: &Storage, + task_id: i64, + control: &TransferTaskControl, + terminal: &mut TransferTerminalState, +) -> Result<(), AppError> { + let Some(task) = load_transfer_task(storage, task_id).await? else { + tracing::warn!(task_id, "transfer task disappeared before finalization"); + return Ok(()); + }; + if stored_transfer_is_terminal(task.status) { + return Ok(()); + } + if control.cancellation.is_cancelled() || task.cancel_requested { + *terminal = TransferTerminalState::Cancelled( + "Transfer cancelled before terminal persistence".to_owned(), + ); + } + + match terminal { + TransferTerminalState::Succeeded(message) => { + let storage = storage.clone(); + let message = message.clone(); + let result = + storage_call(move || storage.complete_transfer_task(task_id, &message)).await; + if result.is_err() { + *terminal = TransferTerminalState::Failed( + "Transfer completed but its success state could not be persisted".to_owned(), + ); + } + result + } + TransferTerminalState::Artifact(_) => { + let TransferTerminalState::Artifact(artifact) = std::mem::replace( + terminal, + TransferTerminalState::Failed( + "Transfer artifact could not be finalized".to_owned(), + ), + ) else { + unreachable!("artifact terminal state must contain an artifact finalizer"); + }; + let record = artifact.finalize().await?; + if record.task_id != Some(task_id) { + return Err(AppError::internal()); + } + let task = load_transfer_task(storage, task_id) + .await? + .ok_or_else(AppError::internal)?; + if stored_transfer_is_terminal(task.status) { + Ok(()) + } else { + Err(AppError::internal()) + } + } + TransferTerminalState::Cancelled(message) => { + persist_cancel_request(storage, task_id).await?; + let Some(task) = load_transfer_task(storage, task_id).await? else { + return Ok(()); }; - if let (Some(storage), Some(message)) = (storage.clone(), terminal_message) { - let _ = storage_call(move || storage.cancel_transfer_task(task_id, message)).await; + if stored_transfer_is_terminal(task.status) { + return Ok(()); } + let storage = storage.clone(); + let message = message.clone(); + storage_call(move || storage.cancel_transfer_task(task_id, &message)).await + } + TransferTerminalState::Failed(message) => { + let storage = storage.clone(); + let message = message.clone(); + storage_call(move || storage.fail_transfer_task(task_id, &message)).await } } } fn validate_import_request(request: &ImportFileRequest) -> Result<(), AppError> { - validate_transfer_scope(&request.datasource_id, &request.database_name, None)?; + validate_transfer_scope(&request.datasource_id, None)?; if request.file_path.trim().is_empty() || request.file_path.contains('\0') { return Err(AppError::invalid( "invalid_import_file", @@ -646,18 +993,13 @@ fn validate_import_request(request: &ImportFileRequest) -> Result<(), AppError> Ok(()) } -fn validate_transfer_scope( - datasource_id: &str, - database_name: &str, - export_path: Option<&str>, -) -> Result<(), AppError> { +fn validate_transfer_scope(datasource_id: &str, export_path: Option<&str>) -> Result<(), AppError> { if datasource_id.trim().is_empty() { return Err(AppError::invalid( "invalid_transfer_request", "datasourceId cannot be empty", )); } - native_mysql::quote_identifier(database_name, "databaseName")?; if export_path.is_some_and(|path| path.trim().is_empty() || path.contains('\0')) { return Err(AppError::invalid( "invalid_export_path", @@ -737,9 +1079,247 @@ fn sha256_hex(digest: &[u8; 32]) -> String { #[cfg(test)] mod tests { + use std::{future::Future, path::Path, time::Duration}; + + use base64::{Engine as _, engine::general_purpose::STANDARD}; + use chat2db_java_bridge::{EngineCommand, EngineConfig}; use chat2db_storage::{StoredTransferTaskKind, StoredTransferTaskStatus, TransferTaskRecord}; + use tempfile::TempDir; + use tokio::sync::oneshot; + + use super::{ + MAX_TRANSFER_FAILURE_MESSAGE_BYTES, TaskCompletion, TransferContext, TransferJobKind, + TransferJobSpec, TransferRunError, TransferTaskControl, TransferTerminalState, + finalize_transfer_task, transfer_task, + }; + use crate::{AppError, Application, RuntimeConfig, RuntimeHost}; + + #[test] + fn transfer_job_spec_accepts_a_send_async_runner() { + fn assert_send(_: &T) {} + + let spec = test_job("send runner", |_application, _context| async { + Ok(TaskCompletion::WithoutArtifact("done".to_owned())) + }); + assert_send(&spec); + } + + #[tokio::test] + async fn cancellation_wins_over_a_successful_runner_completion() { + let directory = TempDir::new().expect("temporary transfer runtime"); + let mut host = RuntimeHost::open(test_runtime_config(directory.path())) + .await + .expect("transfer runtime must open"); + let application = host.application(); + + let (cancel_started, cancel_ready) = oneshot::channel(); + let (release_success, wait_for_success) = oneshot::channel(); + let cancelled = application + .start_transfer_job(test_job( + "explicit cancel", + move |_application, _context| async move { + let _ = cancel_started.send(()); + let _ = wait_for_success.await; + Ok(TaskCompletion::WithoutArtifact("done".to_owned())) + }, + )) + .await + .expect("generic transfer must start"); + wait_until_started(cancel_ready).await; + application + .stop_transfer_task(cancelled.task_id) + .await + .expect("generic transfer must accept cancellation"); + release_success + .send(()) + .expect("successful runner must be released"); + wait_for_status( + &application, + cancelled.task_id, + chat2db_contract::TransferTaskStatus::Cancelled, + ) + .await; + wait_for_hub_removal(&application, cancelled.task_id).await; + + host.shutdown() + .await + .expect("transfer runtime must shut down cleanly"); + } + + #[tokio::test] + async fn panicking_runner_is_failed_and_removed_from_the_hub() { + let directory = TempDir::new().expect("temporary transfer runtime"); + let mut host = RuntimeHost::open(test_runtime_config(directory.path())) + .await + .expect("transfer runtime must open"); + let application = host.application(); + + let failed = application + .start_transfer_job(test_job("panic", panicking_runner)) + .await + .expect("panicking transfer must be admitted"); + wait_for_status( + &application, + failed.task_id, + chat2db_contract::TransferTaskStatus::Failed, + ) + .await; + wait_for_hub_removal(&application, failed.task_id).await; + let task = application + .transfer_task(failed.task_id) + .await + .expect("failed transfer must remain readable"); + assert!(task.error_log.contains("Transfer worker panicked")); - use super::transfer_task; + host.shutdown() + .await + .expect("transfer runtime must shut down cleanly"); + } + + #[tokio::test] + async fn success_persistence_error_falls_back_to_a_failed_terminal_state() { + let directory = TempDir::new().expect("temporary transfer runtime"); + let mut host = RuntimeHost::open(test_runtime_config(directory.path())) + .await + .expect("transfer runtime must open"); + let application = host.application(); + + let failed = application + .start_transfer_job(test_job( + "invalid completion", + |_application, _context| async { + Ok(TaskCompletion::WithoutArtifact("x".repeat(300 * 1024))) + }, + )) + .await + .expect("transfer with invalid completion text must be admitted"); + wait_for_status( + &application, + failed.task_id, + chat2db_contract::TransferTaskStatus::Failed, + ) + .await; + wait_for_hub_removal(&application, failed.task_id).await; + let task = application + .transfer_task(failed.task_id) + .await + .expect("fallback failure must remain readable"); + assert!( + task.error_log + .contains("success state could not be persisted") + ); + + host.shutdown() + .await + .expect("transfer runtime must shut down cleanly"); + } + + #[tokio::test] + async fn oversized_runner_failure_is_bounded_and_removed_from_the_hub() { + let directory = TempDir::new().expect("temporary transfer runtime"); + let mut host = RuntimeHost::open(test_runtime_config(directory.path())) + .await + .expect("transfer runtime must open"); + let application = host.application(); + + let failed = application + .start_transfer_job(test_job( + "oversized failure", + |_application, _context| async { + Err(TransferRunError::Failed(AppError::invalid( + "transfer_test_failed", + "x".repeat(300 * 1024), + ))) + }, + )) + .await + .expect("failing transfer must be admitted"); + wait_for_status( + &application, + failed.task_id, + chat2db_contract::TransferTaskStatus::Failed, + ) + .await; + wait_for_hub_removal(&application, failed.task_id).await; + let task = application + .transfer_task(failed.task_id) + .await + .expect("failed transfer must remain readable"); + assert!(task.error_log.len() <= MAX_TRANSFER_FAILURE_MESSAGE_BYTES); + assert!(task.error_log.ends_with("[truncated]")); + + host.shutdown() + .await + .expect("transfer runtime must shut down cleanly"); + } + + #[tokio::test] + async fn terminal_deadline_covers_waiting_for_the_terminal_gate() { + let directory = TempDir::new().expect("temporary transfer runtime"); + let mut host = RuntimeHost::open(test_runtime_config(directory.path())) + .await + .expect("transfer runtime must open"); + let application = host.application(); + let storage = application + .storage() + .cloned() + .expect("transfer runtime must have storage"); + let control = TransferTaskControl::new(); + let terminal_guard = control.terminal_gate.lock().await; + let deadline = tokio::time::Instant::now() + Duration::from_millis(25); + + let finalized = tokio::time::timeout( + Duration::from_secs(1), + finalize_transfer_task( + &storage, + -1, + &control, + TransferTerminalState::Failed("shutdown".to_owned()), + Some(deadline), + ), + ) + .await + .expect("terminal finalization must respect its deadline"); + assert!(!finalized); + drop(terminal_guard); + + host.shutdown() + .await + .expect("transfer runtime must shut down cleanly"); + } + + #[tokio::test] + async fn generic_transfer_runner_preserves_shutdown_semantics() { + let directory = TempDir::new().expect("temporary transfer runtime"); + let mut host = RuntimeHost::open(test_runtime_config(directory.path())) + .await + .expect("transfer runtime must open"); + let application = host.application(); + + let (shutdown_started, shutdown_ready) = oneshot::channel(); + let interrupted = application + .start_transfer_job(waiting_job("shutdown cancel", shutdown_started)) + .await + .expect("generic transfer must start before shutdown"); + wait_until_started(shutdown_ready).await; + host.shutdown() + .await + .expect("transfer runtime must shut down cleanly"); + let task = application + .transfer_task(interrupted.task_id) + .await + .expect("shutdown transfer task must remain readable"); + assert_eq!(task.status, chat2db_contract::TransferTaskStatus::Cancelled); + assert!( + !application + .inner + .transfer_tasks + .tasks + .lock() + .await + .contains_key(&interrupted.task_id) + ); + } #[test] fn durable_interrupted_status_is_preserved_for_transport_projection() { @@ -768,4 +1348,91 @@ mod tests { chat2db_contract::TransferTaskStatus::Interrupted ); } + + fn test_runtime_config(directory: &Path) -> RuntimeConfig { + RuntimeConfig::new(EngineConfig::new(EngineCommand::new( + directory.join("missing-java"), + ))) + .with_data_dir(directory.join("data")) + .with_vault_master_key_base64(STANDARD.encode([0x74; 32])) + } + + fn test_job(task_name: &str, runner: Runner) -> TransferJobSpec + where + Runner: FnOnce(Application, TransferContext) -> RunnerFuture + Send + 'static, + RunnerFuture: Future> + Send + 'static, + { + TransferJobSpec::new( + "test-driver".to_owned(), + "test-database".to_owned(), + String::new(), + None, + TransferJobKind::ImportFile, + task_name.to_owned(), + runner, + ) + } + + fn waiting_job(task_name: &str, started: oneshot::Sender<()>) -> TransferJobSpec { + test_job(task_name, move |_application, context| async move { + let _ = started.send(()); + context.cancellation().cancelled().await; + Err(TransferRunError::Cancelled) + }) + } + + async fn panicking_runner( + _application: Application, + _context: TransferContext, + ) -> Result { + panic!("intentional transfer runner panic"); + } + + async fn wait_until_started(started: oneshot::Receiver<()>) { + tokio::time::timeout(Duration::from_secs(2), started) + .await + .expect("generic transfer must start before timeout") + .expect("generic transfer runner must signal startup"); + } + + async fn wait_for_status( + application: &Application, + task_id: i64, + expected: chat2db_contract::TransferTaskStatus, + ) { + tokio::time::timeout(Duration::from_secs(2), async { + loop { + let task = application + .transfer_task(task_id) + .await + .expect("generic transfer task must remain readable"); + if task.status == expected { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("generic transfer status must settle before timeout"); + } + + async fn wait_for_hub_removal(application: &Application, task_id: i64) { + tokio::time::timeout(Duration::from_secs(2), async { + loop { + if !application + .inner + .transfer_tasks + .tasks + .lock() + .await + .contains_key(&task_id) + { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("terminal transfer must leave the active hub before timeout"); + } } diff --git a/crates/chat2db-core/src/transfer/mysql.rs b/crates/chat2db-core/src/transfer/mysql.rs index 0d58aee..b3262bd 100644 --- a/crates/chat2db-core/src/transfer/mysql.rs +++ b/crates/chat2db-core/src/transfer/mysql.rs @@ -27,7 +27,7 @@ use tokio_util::sync::CancellationToken; use uuid::Uuid; use zip::{CompressionMethod, ZipWriter, write::SimpleFileOptions}; -use super::{TaskCompletion, TransferContext, TransferRunError, format}; +use super::{PendingTransferArtifact, TaskCompletion, TransferContext, TransferRunError, format}; use crate::{AppError, AppErrorKind, Application, native_mysql}; const IMPORT_BATCH_ROWS: usize = 256; @@ -1131,13 +1131,20 @@ async fn publish_task_artifact( None => None, }; context.check_cancelled()?; - let artifact = writer.finish().map_err(AppError::from)?; - if let Some(pending) = pending - && let Err(error) = pending.publish().await - { - tracing::warn!(%error, artifact_id = %artifact.id, "managed transfer succeeded but exportPath publication failed"); - } - Ok(TaskCompletion::Artifact(artifact)) + Ok(TaskCompletion::Artifact(PendingTransferArtifact::new( + move || async move { + let artifact = tokio::task::spawn_blocking(move || writer.finish()) + .await + .map_err(|_| AppError::internal())? + .map_err(AppError::from)?; + if let Some(pending) = pending + && let Err(error) = pending.publish().await + { + tracing::warn!(%error, artifact_id = %artifact.id, "managed transfer succeeded but exportPath publication failed"); + } + Ok(artifact) + }, + ))) } struct PendingUserCopy { diff --git a/crates/chat2db-core/src/transfer/mysql_impl.rs b/crates/chat2db-core/src/transfer/mysql_impl.rs new file mode 100644 index 0000000..a251c45 --- /dev/null +++ b/crates/chat2db-core/src/transfer/mysql_impl.rs @@ -0,0 +1,100 @@ +use std::path::Path; + +use chat2db_contract::{ + DmlExportRequest, ImportFileRequest, OtherFileExportRequest, SqlFileExportRequest, + TransferArtifact, +}; + +use super::{ + TransferJobKind, TransferJobSpec, mysql, single_table, transfer_artifact, + validate_import_request, validate_transfer_scope, +}; +use crate::{AppError, Application, native_mysql}; + +pub(crate) async fn import_file( + application: &Application, + request: ImportFileRequest, +) -> Result { + validate_import_request(&request)?; + validate_mysql_database(&request.database_name)?; + native_mysql::resolve_native_connection(application, &request.datasource_id).await?; + let file_name = Path::new(&request.file_path) + .file_name() + .and_then(|value| value.to_str()) + .unwrap_or("file"); + Ok(TransferJobSpec::new( + request.datasource_id.clone(), + request.database_name.clone(), + request.schema_name.clone(), + request.table_name.clone(), + TransferJobKind::ImportFile, + format!("Import {file_name}"), + move |application, context| async move { + mysql::import_file(&application, request, &context).await + }, + )) +} + +pub(crate) async fn export_sql_file( + application: &Application, + request: SqlFileExportRequest, +) -> Result { + validate_transfer_scope(&request.datasource_id, request.export_path.as_deref())?; + validate_mysql_database(&request.database_name)?; + native_mysql::resolve_native_connection(application, &request.datasource_id).await?; + Ok(TransferJobSpec::new( + request.datasource_id.clone(), + request.database_name.clone(), + request.schema_name.clone(), + single_table(&request.table_names), + TransferJobKind::ExportSql, + format!("Export SQL {}", request.database_name), + move |application, context| async move { + mysql::export_sql(&application, request, &context).await + }, + )) +} + +pub(crate) async fn export_other_file( + application: &Application, + request: OtherFileExportRequest, +) -> Result { + validate_transfer_scope(&request.datasource_id, request.export_path.as_deref())?; + validate_mysql_database(&request.database_name)?; + if request.table_names.is_empty() { + return Err(AppError::invalid( + "missing_export_tables", + "tableNames must contain at least one table", + )); + } + native_mysql::resolve_native_connection(application, &request.datasource_id).await?; + Ok(TransferJobSpec::new( + request.datasource_id.clone(), + request.database_name.clone(), + request.schema_name.clone(), + single_table(&request.table_names), + TransferJobKind::ExportFile, + format!( + "Export {} {} table(s)", + request.format.extension().to_ascii_uppercase(), + request.table_names.len() + ), + move |application, context| async move { + mysql::export_other(&application, request, &context).await + }, + )) +} + +pub(crate) async fn export_dml( + application: &Application, + request: DmlExportRequest, +) -> Result { + native_mysql::resolve_native_connection(application, &request.datasource_id).await?; + mysql::export_dml(application, request) + .await + .map(transfer_artifact) +} + +fn validate_mysql_database(database_name: &str) -> Result<(), AppError> { + native_mysql::quote_identifier(database_name, "databaseName").map(|_| ()) +} From 11b6f26ce54b1f1b8ad723981047256cf75a1edf Mon Sep 17 00:00:00 2001 From: zgq Date: Thu, 6 Aug 2026 00:06:31 +0800 Subject: [PATCH 4/5] refactor(core): decouple native driver spi contracts --- apps/chat2db-desktop/src/lib.rs | 14 +- apps/chat2db-web/src/legacy.rs | 20 +- crates/chat2db-core/src/community.rs | 201 +++- crates/chat2db-core/src/convert.rs | 133 ++- .../src/datasource_compatibility.rs | 137 ++- crates/chat2db-core/src/lib.rs | 37 +- crates/chat2db-core/src/mysql_dashboard.rs | 67 +- crates/chat2db-core/src/mysql_workspace.rs | 1 + crates/chat2db-core/src/native_api_adapter.rs | 918 +++++++++++++++++- crates/chat2db-core/src/native_driver.rs | 452 +++++---- .../chat2db-core/src/native_driver_types.rs | 622 ++++++------ crates/chat2db-core/src/native_mysql.rs | 304 +++--- crates/chat2db-core/src/query.rs | 110 ++- crates/chat2db-core/src/ssh.rs | 2 +- .../src/transfer/class_generation.rs | 8 +- crates/chat2db-core/src/transfer/mod.rs | 214 +++- crates/chat2db-core/src/transfer/mysql.rs | 40 +- .../{mysql_impl.rs => mysql_driver.rs} | 21 +- .../tests/native_mysql_console_docker.rs | 26 +- .../tests/native_mysql_product.rs | 10 +- .../tests/native_mysql_ssh_tunnel_docker.rs | 18 +- 21 files changed, 2475 insertions(+), 880 deletions(-) rename crates/chat2db-core/src/transfer/{mysql_impl.rs => mysql_driver.rs} (84%) diff --git a/apps/chat2db-desktop/src/lib.rs b/apps/chat2db-desktop/src/lib.rs index ca1c493..5fbf309 100644 --- a/apps/chat2db-desktop/src/lib.rs +++ b/apps/chat2db-desktop/src/lib.rs @@ -41,7 +41,7 @@ use chat2db_contract::{ UpdateDatasourceRequest, UpdateProviderProfileRequest, ValidateCommunitySqlRequest, }; use chat2db_core::{ - AppError, Application, MysqlConsoleCancellation, RuntimeConfig, RuntimeHost, + AppError, Application, NativeConsoleCancellation, RuntimeConfig, RuntimeHost, load_fixed_community_classpath, }; use chat2db_java_bridge::{BridgeError, EngineCommand, EngineConfig}; @@ -222,11 +222,11 @@ impl DesktopStartup { #[derive(Default)] struct LegacySqlCancellationRegistry { - cancellations: Mutex>, + cancellations: Mutex>, } impl LegacySqlCancellationRegistry { - async fn insert(&self, execution_id: String, cancellation: MysqlConsoleCancellation) { + async fn insert(&self, execution_id: String, cancellation: NativeConsoleCancellation) { self.cancellations .lock() .await @@ -1054,7 +1054,7 @@ async fn legacy_client_command_for( }) .map(|id| format!("mysql-console-{id}")) .map_err(|_| "No MySQL Console execution ids remain".to_owned())?; - let cancellation = MysqlConsoleCancellation::new(); + let cancellation = NativeConsoleCancellation::new(); state .legacy_sql_cancellations .insert(execution_id.clone(), cancellation.clone()) @@ -1278,7 +1278,7 @@ async fn forward_native_mysql_sql_execution( request_uuid: String, execution_id: String, request: chat2db_web::legacy::LegacySqlExecuteRequest, - cancellation: MysqlConsoleCancellation, + cancellation: NativeConsoleCancellation, ) { let started_at = Instant::now(); let mut sequence = 0_u64; @@ -2766,7 +2766,7 @@ mod tests { OperationEvent, OperationEventEnvelope, OperationStreamMessage, StartCommunityTablePreviewRequest, ValidateCommunitySqlRequest, }; - use chat2db_core::{AppError, Application, MysqlConsoleCancellation}; + use chat2db_core::{AppError, Application, NativeConsoleCancellation}; use tokio::sync::oneshot; use super::{ @@ -2989,7 +2989,7 @@ mod tests { #[tokio::test] async fn native_mysql_cancellation_registry_owns_execution_lifecycle() { let registry = LegacySqlCancellationRegistry::default(); - let cancellation = MysqlConsoleCancellation::new(); + let cancellation = NativeConsoleCancellation::new(); registry .insert("mysql-console-1".to_owned(), cancellation.clone()) .await; diff --git a/apps/chat2db-web/src/legacy.rs b/apps/chat2db-web/src/legacy.rs index 480211c..2baa2b4 100644 --- a/apps/chat2db-web/src/legacy.rs +++ b/apps/chat2db-web/src/legacy.rs @@ -57,7 +57,7 @@ use chat2db_contract::{ }; use chat2db_core::{ AppError, AppErrorKind, Application, LargeValueChunk, LargeValueEncoding, LargeValuePreview, - LargeValueType, MysqlConsoleCancellation, MysqlConsoleRequest, MysqlConsoleResult, + LargeValueType, NativeConsoleCancellation, NativeConsoleRequest, NativeConsoleResult, TransferArtifactDownload, mysql_ddl::{ MysqlColumnAlter, MysqlColumnDefinition, MysqlColumnPosition, MysqlDatabaseDefinition, @@ -5381,7 +5381,7 @@ pub async fn execute_sql( return Box::pin(execute_mysql_sql( application, request, - MysqlConsoleCancellation::new(), + NativeConsoleCancellation::new(), &execution_id, "SQL_EDITOR_HTTP", )) @@ -5452,8 +5452,8 @@ pub(crate) async fn count_mysql_rows( } let count_sql = build_mysql_count_query(&request.sql)?; let results = application - .execute_mysql_console( - MysqlConsoleRequest { + .execute_native_console( + NativeConsoleRequest { datasource_id, database_name: request.database_name.clone(), sql: count_sql, @@ -5465,7 +5465,7 @@ pub(crate) async fn count_mysql_rows( explain: false, error_continue: false, }, - MysqlConsoleCancellation::new(), + NativeConsoleCancellation::new(), ) .await?; let result = results.into_iter().next().ok_or_else(|| { @@ -5518,7 +5518,7 @@ pub(crate) async fn count_mysql_rows( pub async fn execute_mysql_sql( application: &Application, request: &LegacySqlExecuteRequest, - cancellation: MysqlConsoleCancellation, + cancellation: NativeConsoleCancellation, execution_id: &str, history_source: &str, ) -> LegacyResult> { @@ -5541,8 +5541,8 @@ pub async fn execute_mysql_sql( .map(|columns| columns.items) }; - let execution = application.execute_mysql_console( - MysqlConsoleRequest { + let execution = application.execute_native_console( + NativeConsoleRequest { datasource_id, database_name: request.database_name.clone(), sql: request.sql.clone(), @@ -5950,10 +5950,10 @@ fn mysql_console_result( application: &Application, large_value_owner: &str, request: &LegacySqlExecuteRequest, - result: MysqlConsoleResult, + result: NativeConsoleResult, editable_columns: Option<&[CommunityTableColumn]>, ) -> LegacyManageResult { - let MysqlConsoleResult { + let NativeConsoleResult { statement_sequence, result_set_id, sql, diff --git a/crates/chat2db-core/src/community.rs b/crates/chat2db-core/src/community.rs index 07ea352..9634a55 100644 --- a/crates/chat2db-core/src/community.rs +++ b/crates/chat2db-core/src/community.rs @@ -130,7 +130,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_schemas(self, request.into()).await; + return metadata + .list_schemas(self, request.into()) + .await + .map(crate::native_api_adapter::schema_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -170,7 +173,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_databases(self, request.into()).await; + return metadata + .list_databases(self, request.into()) + .await + .map(crate::native_api_adapter::database_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -209,7 +215,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_tables(self, request.into()).await; + return metadata + .list_tables(self, request.into()) + .await + .map(crate::native_api_adapter::table_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -258,7 +267,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_columns(self, request.into()).await; + return metadata + .list_columns(self, request.into()) + .await + .map(crate::native_api_adapter::column_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -359,7 +371,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_indexes(self, request.into()).await; + return metadata + .list_indexes(self, request.into()) + .await + .map(crate::native_api_adapter::index_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -408,7 +423,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_views(self, request.into()).await; + return metadata + .list_views(self, request.into()) + .await + .map(crate::native_api_adapter::view_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -458,7 +476,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.get_view(self, request.into()).await; + return metadata + .get_view(self, request.into()) + .await + .map(crate::native_api_adapter::table_response); } self.list_community_views(request) .await? @@ -485,7 +506,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_imported_keys(self, request.into()).await; + return metadata + .list_imported_keys(self, request.into()) + .await + .map(crate::native_api_adapter::foreign_key_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -534,7 +558,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_exported_keys(self, request.into()).await; + return metadata + .list_exported_keys(self, request.into()) + .await + .map(crate::native_api_adapter::foreign_key_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -583,7 +610,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_primary_keys(self, request.into()).await; + return metadata + .list_primary_keys(self, request.into()) + .await + .map(crate::native_api_adapter::primary_key_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -632,7 +662,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_functions(self, request.into()).await; + return metadata + .list_functions(self, request.into()) + .await + .map(crate::native_api_adapter::function_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -673,7 +706,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.get_function(self, request.into()).await; + return metadata + .get_function(self, request.into()) + .await + .map(crate::native_api_adapter::function_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -722,7 +758,8 @@ impl Application { { return metadata .list_function_parameters(self, request.into()) - .await; + .await + .map(crate::native_api_adapter::function_parameter_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -774,7 +811,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_procedures(self, request.into()).await; + return metadata + .list_procedures(self, request.into()) + .await + .map(crate::native_api_adapter::procedure_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -815,7 +855,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.get_procedure(self, request.into()).await; + return metadata + .get_procedure(self, request.into()) + .await + .map(crate::native_api_adapter::procedure_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -864,7 +907,8 @@ impl Application { { return metadata .list_procedure_parameters(self, request.into()) - .await; + .await + .map(crate::native_api_adapter::procedure_parameter_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -927,7 +971,11 @@ impl Application { "The native Rust driver does not implement routine operations", ) })?; - routines.preview_invocation(self, request.into()).await + routines + .preview_invocation(self, request.into()) + .await + .map(crate::native_api_adapter::routine_invocation_response) + .map_err(crate::native_api_adapter::compatibility_api_error) } /// Previews the compensating `MySQL` routine-replacement script. @@ -954,7 +1002,10 @@ impl Application { "The native Rust driver does not implement routine operations", ) })?; - routines.preview_migration(request.clone().into()) + routines + .preview_migration(request.clone().into()) + .map(crate::native_api_adapter::routine_invocation_response) + .map_err(crate::native_api_adapter::compatibility_api_error) } /// Replaces one `MySQL` routine and restores its before-image when apply fails. @@ -982,7 +1033,11 @@ impl Application { "The native Rust driver does not implement routine operations", ) })?; - routines.execute_migration(self, request.into()).await + routines + .execute_migration(self, request.into()) + .await + .map(crate::native_api_adapter::routine_migration_execution_response) + .map_err(crate::native_api_adapter::compatibility_api_error) } /// Lists triggers through Community metadata using a forced read-only session. @@ -997,7 +1052,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.list_triggers(self, request.into()).await; + return metadata + .list_triggers(self, request.into()) + .await + .map(crate::native_api_adapter::trigger_list_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -1038,7 +1096,10 @@ impl Application { if let Some(driver) = self.native_driver_for_database_type(&request.database_type) && let Some(metadata) = driver.metadata() { - return metadata.get_trigger(self, request.into()).await; + return metadata + .get_trigger(self, request.into()) + .await + .map(crate::native_api_adapter::trigger_response); } let storage = self.require_storage()?; let engine = self.require_community_engine().await?; @@ -1175,7 +1236,9 @@ impl Application { { return tables .start_table_preview(self, request.into(), row_limit) - .await; + .await + .map(crate::native_api_adapter::table_preview_response) + .map_err(crate::native_api_adapter::compatibility_api_error); } let engine = self.require_community_engine().await?; @@ -2295,6 +2358,7 @@ fn preserve_primary_result( #[cfg(test)] mod tests { + use async_trait::async_trait; use chat2db_contract::{ BuildCommunityDmlRequest, BuildCommunityNamespaceSqlRequest, CommunityDatabase, CommunityDmlColumn, CommunityDmlRow, CommunityDmlStatement, CommunityDmlTarget, @@ -2304,9 +2368,10 @@ mod tests { CommunityPluginServices, CommunityPrimaryKey, CommunityProcedure, CommunityProcedureParameter, CommunitySchema, CommunitySqlAnalysis, CommunitySqlDiagnostic, CommunitySqlValidation, CommunityTable, CommunityTableColumn, CommunityTableIndex, - CommunityTableIndexColumn, CommunityTrigger, ListCommunityColumnsRequest, - ListCommunityDatabasesRequest, ListCommunityIndexesRequest, ListCommunitySchemasRequest, - ListCommunityTableKeysRequest, ListCommunityTablesRequest, ListCommunityViewsRequest, + CommunityTableIndexColumn, CommunityTrigger, DatasourceConnection, + ListCommunityColumnsRequest, ListCommunityDatabasesRequest, ListCommunityIndexesRequest, + ListCommunitySchemasRequest, ListCommunityTableKeysRequest, ListCommunityTablesRequest, + ListCommunityViewsRequest, }; use chat2db_java_bridge::{ CommunityDatabase as BridgeCommunityDatabase, @@ -2347,7 +2412,91 @@ mod tests { run_cancellation_safe, run_cancellation_safe_with_cleanup, table_preview_row_limit, validate_table_preview_sql, }; - use crate::{AppError, AppErrorKind}; + use crate::{ + AppError, AppErrorKind, + native_driver::{ + NativeConnectionDriver, NativeDialectDriver, NativeDriver, NativeDriverRegistry, + }, + native_driver_types::{ + BuiltSql, CreateSchemaSqlRequest, DmlSqlRequest, NamespaceSqlRequest, + NativeDriverDescriptor, + }, + }; + + struct FakePostgresDriver; + + const FAKE_POSTGRES_DESCRIPTOR: NativeDriverDescriptor = NativeDriverDescriptor { + id: "postgresql", + implementation: "fake_postgres", + database_types: &["POSTGRESQL"], + compatibility_aliases: &["postgresql", "org.postgresql.Driver"], + }; + + impl NativeDriver for FakePostgresDriver { + fn descriptor(&self) -> &'static NativeDriverDescriptor { + &FAKE_POSTGRES_DESCRIPTOR + } + + fn connection(&self) -> &dyn NativeConnectionDriver { + self + } + + fn dialect(&self) -> Option<&dyn NativeDialectDriver> { + Some(self) + } + } + + #[async_trait] + impl NativeConnectionDriver for FakePostgresDriver { + async fn test_connection( + &self, + _connection: &DatasourceConnection, + ) -> Result<(), AppError> { + Ok(()) + } + } + + impl NativeDialectDriver for FakePostgresDriver { + fn build_create_schema( + &self, + request: CreateSchemaSqlRequest, + ) -> Result { + Ok(BuiltSql { + sql: format!("fake-postgres:create-schema:{}", request.schema.name), + }) + } + + fn build_namespace_sql(&self, _request: NamespaceSqlRequest) -> Result { + Ok(BuiltSql { + sql: "fake-postgres:namespace".to_owned(), + }) + } + + fn build_dml(&self, _request: DmlSqlRequest) -> Result { + Ok(BuiltSql { + sql: "fake-postgres:dml".to_owned(), + }) + } + } + + #[tokio::test] + async fn compatibility_namespace_api_dispatches_through_the_native_registry() { + let registry = NativeDriverRegistry::try_new(vec![Arc::new(FakePostgresDriver)]) + .expect("registry is valid"); + let application = Application::with_native_drivers_for_test(registry); + + let response = application + .build_community_namespace_sql(BuildCommunityNamespaceSqlRequest { + database_type: "POSTGRESQL".to_owned(), + operation: CommunityNamespaceSqlOperation::UseDatabase { + database_name: "inventory".to_owned(), + }, + }) + .await + .expect("the compatibility API must use the registered native capability"); + + assert_eq!(response.sql, "fake-postgres:namespace"); + } fn binary_dml_request(base64: &str) -> BuildCommunityDmlRequest { BuildCommunityDmlRequest { diff --git a/crates/chat2db-core/src/convert.rs b/crates/chat2db-core/src/convert.rs index 1957c3a..091463b 100644 --- a/crates/chat2db-core/src/convert.rs +++ b/crates/chat2db-core/src/convert.rs @@ -4,7 +4,10 @@ use chat2db_engine_protocol::wire; use chat2db_java_bridge as bridge; use chat2db_storage as storage; -use crate::AppError; +use crate::{ + AppError, + query::{DatabaseValue, QueryExecutionOptions, QueryParameter}, +}; pub(crate) const fn provider_kind_to_storage( kind: contract::ProviderKind, @@ -125,58 +128,93 @@ pub(crate) fn result_page(page: storage::ResultPage) -> Result Result { - Ok(bridge::JdbcParameter { +) -> Result { + Ok(QueryParameter { position: parameter.position, - value: request_value(parameter.value)?, - jdbc_type: None, - jdbc_type_name: None, + value: database_value_from_contract(parameter.value)?, }) } -fn request_value(value: contract::JdbcValue) -> Result { +fn database_value_from_contract(value: contract::JdbcValue) -> Result { Ok(match value { - contract::JdbcValue::Null => bridge::JdbcValue::Null, - contract::JdbcValue::Boolean { value } => bridge::JdbcValue::Boolean(value), + contract::JdbcValue::Null => DatabaseValue::Null, + contract::JdbcValue::Boolean { value } => DatabaseValue::Boolean(value), contract::JdbcValue::SignedInteger { value } => { - bridge::JdbcValue::SignedInteger(parse_number(&value, "signed integer")?) + DatabaseValue::SignedInteger(parse_number(&value, "signed integer")?) } contract::JdbcValue::UnsignedInteger { value } => { - bridge::JdbcValue::UnsignedInteger(parse_number(&value, "unsigned integer")?) - } - contract::JdbcValue::Float32 { value } => { - bridge::JdbcValue::Float32(parse_float32(&value)?) - } - contract::JdbcValue::Float64 { value } => { - bridge::JdbcValue::Float64(parse_float64(&value)?) + DatabaseValue::UnsignedInteger(parse_number(&value, "unsigned integer")?) } - contract::JdbcValue::Decimal { value } => bridge::JdbcValue::Decimal(value), - contract::JdbcValue::Text { value } => bridge::JdbcValue::Text(value), + contract::JdbcValue::Float32 { value } => DatabaseValue::Float32(parse_float32(&value)?), + contract::JdbcValue::Float64 { value } => DatabaseValue::Float64(parse_float64(&value)?), + contract::JdbcValue::Decimal { value } => DatabaseValue::Decimal(value), + contract::JdbcValue::Text { value } => DatabaseValue::Text(value), contract::JdbcValue::Binary { value } => { - bridge::JdbcValue::Binary(BASE64_STANDARD.decode(value).map_err(|_| { + DatabaseValue::Binary(BASE64_STANDARD.decode(value).map_err(|_| { AppError::invalid( "invalid_query_parameter", "Binary parameters must use base64", ) })?) } - contract::JdbcValue::Date { value } => bridge::JdbcValue::Date(value), - contract::JdbcValue::Time { value } => bridge::JdbcValue::Time(value), - contract::JdbcValue::Timestamp { value } => bridge::JdbcValue::Timestamp(value), + contract::JdbcValue::Date { value } => DatabaseValue::Date(value), + contract::JdbcValue::Time { value } => DatabaseValue::Time(value), + contract::JdbcValue::Timestamp { value } => DatabaseValue::Timestamp(value), contract::JdbcValue::TimestampWithTimeZone { value } => { - bridge::JdbcValue::TimestampWithTimeZone(value) + DatabaseValue::TimestampWithTimeZone(value) } - contract::JdbcValue::Json { value } => bridge::JdbcValue::Json(value), - contract::JdbcValue::Uuid { value } => bridge::JdbcValue::Uuid(value), + contract::JdbcValue::Json { value } => DatabaseValue::Json(value), + contract::JdbcValue::Uuid { value } => DatabaseValue::Uuid(value), contract::JdbcValue::Opaque { .. } => { return Err(AppError::invalid( "invalid_query_parameter", - "Opaque JDBC values cannot be query parameters", + "Opaque result values cannot be query parameters", )); } }) } +pub(crate) fn query_parameter_to_java(parameter: QueryParameter) -> bridge::JdbcParameter { + bridge::JdbcParameter { + position: parameter.position, + value: database_value_to_java(parameter.value), + jdbc_type: None, + jdbc_type_name: None, + } +} + +pub(crate) const fn query_options_to_java(options: QueryExecutionOptions) -> bridge::QueryOptions { + bridge::QueryOptions { + max_rows: options.max_rows, + target_batch_rows: options.target_batch_rows, + target_batch_bytes: options.target_batch_bytes, + initial_batch_credits: options.initial_batch_credits, + max_result_bytes: options.max_result_bytes, + } +} + +fn database_value_to_java(value: DatabaseValue) -> bridge::JdbcValue { + match value { + DatabaseValue::Null => bridge::JdbcValue::Null, + DatabaseValue::Boolean(value) => bridge::JdbcValue::Boolean(value), + DatabaseValue::SignedInteger(value) => bridge::JdbcValue::SignedInteger(value), + DatabaseValue::UnsignedInteger(value) => bridge::JdbcValue::UnsignedInteger(value), + DatabaseValue::Float32(value) => bridge::JdbcValue::Float32(value), + DatabaseValue::Float64(value) => bridge::JdbcValue::Float64(value), + DatabaseValue::Decimal(value) => bridge::JdbcValue::Decimal(value), + DatabaseValue::Text(value) => bridge::JdbcValue::Text(value), + DatabaseValue::Binary(value) => bridge::JdbcValue::Binary(value), + DatabaseValue::Date(value) => bridge::JdbcValue::Date(value), + DatabaseValue::Time(value) => bridge::JdbcValue::Time(value), + DatabaseValue::Timestamp(value) => bridge::JdbcValue::Timestamp(value), + DatabaseValue::TimestampWithTimeZone(value) => { + bridge::JdbcValue::TimestampWithTimeZone(value) + } + DatabaseValue::Json(value) => bridge::JdbcValue::Json(value), + DatabaseValue::Uuid(value) => bridge::JdbcValue::Uuid(value), + } +} + fn result_row(row: wire::JdbcRow) -> Result { Ok(contract::ResultRow { values: row @@ -350,7 +388,11 @@ mod tests { use chat2db_contract::JdbcValue; use chat2db_engine_protocol::wire; - use super::{request_value, result_value}; + use super::{ + database_value_from_contract, database_value_to_java, query_options_to_java, + query_parameter_to_java, result_value, + }; + use crate::query::{DatabaseValue, QueryExecutionOptions, QueryParameter}; #[test] fn all_sixteen_portable_jdbc_values_round_trip() { @@ -399,7 +441,8 @@ mod tests { ]; for value in values { - let bridge = request_value(value.clone()).expect("request conversion"); + let neutral = database_value_from_contract(value.clone()).expect("request conversion"); + let bridge = database_value_to_java(neutral); let wire = wire::JdbcValue::from(bridge); assert_eq!(result_value(wire).expect("result conversion"), value); } @@ -419,17 +462,45 @@ mod tests { #[test] fn opaque_and_noncanonical_infinity_are_rejected_as_inputs() { assert!( - request_value(JdbcValue::Opaque { + database_value_from_contract(JdbcValue::Opaque { type_name: "vendor.Type".to_owned(), display_value: "value".to_owned(), }) .is_err() ); assert!( - request_value(JdbcValue::Float64 { + database_value_from_contract(JdbcValue::Float64 { value: "inf".to_owned(), }) .is_err() ); } + + #[test] + fn java_query_adapter_preserves_parameter_and_execution_options() { + let parameter = query_parameter_to_java(QueryParameter { + position: 7, + value: DatabaseValue::Text("inventory".to_owned()), + }); + assert_eq!(parameter.position, 7); + assert!(matches!( + parameter.value, + chat2db_java_bridge::JdbcValue::Text(value) if value == "inventory" + )); + assert_eq!(parameter.jdbc_type, None); + assert_eq!(parameter.jdbc_type_name, None); + + let options = query_options_to_java(QueryExecutionOptions { + max_rows: 101, + target_batch_rows: 102, + target_batch_bytes: 103, + initial_batch_credits: 104, + max_result_bytes: 105, + }); + assert_eq!(options.max_rows, 101); + assert_eq!(options.target_batch_rows, 102); + assert_eq!(options.target_batch_bytes, 103); + assert_eq!(options.initial_batch_credits, 104); + assert_eq!(options.max_result_bytes, 105); + } } diff --git a/crates/chat2db-core/src/datasource_compatibility.rs b/crates/chat2db-core/src/datasource_compatibility.rs index fe18253..90f00fe 100644 --- a/crates/chat2db-core/src/datasource_compatibility.rs +++ b/crates/chat2db-core/src/datasource_compatibility.rs @@ -1,4 +1,4 @@ -use std::collections::HashSet; +use std::{collections::HashSet, sync::Arc}; use chat2db_contract::{ CloneDatasourceRequest, CommunityDatasourceExport, CommunityDatasourceImportResult, @@ -13,8 +13,12 @@ use chat2db_storage::{CreateDatasource, SecretValue, StorageError}; use url::Url; use crate::{ - AppError, Application, convert, datasource_edit::project_ssh, - datasource_session::resolve_datasource_connection, now_millis, storage_call, + AppError, Application, convert, + datasource_edit::project_ssh, + datasource_session::resolve_datasource_connection, + native_driver::{NativeDriver, NativeDriverRegistry}, + native_driver_types::NativeDriverDescriptor, + now_millis, storage_call, }; const COMMUNITY_DATASOURCE_DOCUMENT_VERSION: u32 = 1; @@ -182,7 +186,7 @@ impl Application { ) -> Result { let database_type = if database_type.trim().is_empty() { let datasource = self.get_datasource(datasource_id).await?; - self.native_database_type_for_driver(&datasource.driver_id) + self.native_database_type_for_datasource_driver_id(&datasource.driver_id) .unwrap_or(datasource.driver_id) } else { database_type.trim().to_owned() @@ -260,28 +264,88 @@ impl Application { })?; let descriptor = driver.descriptor(); Ok(NativeDriverCompatibility { - database_type: driver.database_types()[0].to_owned(), - driver_id: descriptor.driver_id, + database_type: descriptor.database_types[0].to_owned(), + driver_id: descriptor.id.to_owned(), action, - implementation: driver.implementation().to_owned(), + implementation: descriptor.implementation.to_owned(), artifact_required: false, changed: false, }) } } -pub(crate) fn native_mysql_driver() -> JdbcDriver { +/// Converts native identity metadata into the historical JDBC-shaped HTTP contract. +pub(crate) fn jdbc_driver_from_descriptor(descriptor: &NativeDriverDescriptor) -> JdbcDriver { + let display_name = match descriptor.id.to_ascii_lowercase().as_str() { + "mysql" => "MySQL".to_owned(), + _ => descriptor + .database_types + .first() + .copied() + .unwrap_or(descriptor.id) + .to_owned(), + }; JdbcDriver { - pack_id: "native:mysql_async".to_owned(), - name: "MySQL (native Rust)".to_owned(), + pack_id: format!("native:{}", descriptor.implementation), + name: format!("{display_name} (native Rust)"), version: "native".to_owned(), - driver_id: "mysql".to_owned(), - driver_class: "rust:mysql_async".to_owned(), + driver_id: descriptor.id.to_owned(), + driver_class: format!("rust:{}", descriptor.implementation), artifact_count: 0, artifact_bytes: "0".to_owned(), } } +/// Resolves the native implementation for a persisted datasource driver ID. +pub(crate) fn native_driver_for_datasource_driver_id( + registry: &NativeDriverRegistry, + datasource_driver_id: &str, + managed_drivers: &[JdbcDriver], +) -> Option> { + if let Some(driver) = registry.driver_for_datasource_driver_id(datasource_driver_id) { + return Some(driver); + } + + let managed_descriptor = managed_drivers.iter().find(|driver| { + driver + .driver_id + .eq_ignore_ascii_case(datasource_driver_id.trim()) + })?; + let mut matches = registry + .descriptors() + .filter(|descriptor| jdbc_driver_matches_descriptor(managed_descriptor, descriptor)); + let driver_id = matches.next()?.id; + if matches.next().is_some() { + return None; + } + registry.driver_for_datasource_driver_id(driver_id) +} + +fn jdbc_driver_matches_descriptor( + jdbc_driver: &JdbcDriver, + descriptor: &NativeDriverDescriptor, +) -> bool { + let compatibility_values = [ + jdbc_driver.pack_id.as_str(), + jdbc_driver.name.as_str(), + jdbc_driver.driver_id.as_str(), + jdbc_driver.driver_class.as_str(), + ]; + descriptor + .compatibility_aliases + .iter() + .copied() + .filter_map(|alias| { + let alias = alias.trim().to_ascii_lowercase(); + (!alias.is_empty()).then_some(alias) + }) + .any(|alias| { + compatibility_values + .iter() + .any(|value| value.to_ascii_lowercase().contains(&alias)) + }) +} + fn copy_name(name: &str) -> String { let candidate = format!("{name} Copy"); if candidate.len() <= 512 { @@ -431,14 +495,16 @@ mod tests { use chat2db_contract::{ CreateDatasourceRequest, DatasourceConnection, DatasourceConnectionProperty, - ExportCommunityDatasourcesRequest, NativeDriverAction, SshAuthentication, + ExportCommunityDatasourcesRequest, JdbcDriver, NativeDriverAction, SshAuthentication, SshAuthenticationType, SshHostKeyVerification, SshTunnelConfig, }; use chat2db_storage::{SecretRef, SecretValue, SecretVault, SecretVaultError, Storage}; use tempfile::TempDir; - use super::CloneDatasourceRequest; - use crate::Application; + use super::{ + CloneDatasourceRequest, jdbc_driver_from_descriptor, native_driver_for_datasource_driver_id, + }; + use crate::{Application, native_driver::NativeDriverRegistry}; #[derive(Debug, Default)] struct MemoryVault { @@ -665,4 +731,45 @@ mod tests { assert!(!compatibility.artifact_required); assert!(!compatibility.changed); } + + #[test] + fn native_descriptor_preserves_the_existing_mysql_jdbc_wire_shape() { + let registry = NativeDriverRegistry::built_in(); + let descriptor = registry + .descriptors() + .next() + .expect("built-in MySQL descriptor exists"); + + assert_eq!( + jdbc_driver_from_descriptor(descriptor), + JdbcDriver { + pack_id: "native:mysql_async".to_owned(), + name: "MySQL (native Rust)".to_owned(), + version: "native".to_owned(), + driver_id: "mysql".to_owned(), + driver_class: "rust:mysql_async".to_owned(), + artifact_count: 0, + artifact_bytes: "0".to_owned(), + } + ); + } + + #[test] + fn managed_jdbc_driver_id_resolves_through_the_datasource_compatibility_boundary() { + let registry = NativeDriverRegistry::built_in(); + let managed_drivers = vec![JdbcDriver { + pack_id: "mysql-connector-j".to_owned(), + name: "MySQL JDBC".to_owned(), + version: "9".to_owned(), + driver_id: "managed-mysql".to_owned(), + driver_class: "com.mysql.cj.jdbc.Driver".to_owned(), + artifact_count: 1, + artifact_bytes: "1".to_owned(), + }]; + + let driver = + native_driver_for_datasource_driver_id(®istry, "managed-mysql", &managed_drivers) + .expect("managed MySQL descriptor resolves to the native implementation"); + assert_eq!(driver.descriptor().id, "mysql"); + } } diff --git a/crates/chat2db-core/src/lib.rs b/crates/chat2db-core/src/lib.rs index 8593791..26c3717 100644 --- a/crates/chat2db-core/src/lib.rs +++ b/crates/chat2db-core/src/lib.rs @@ -65,7 +65,6 @@ pub use large_value::{ }; pub use legacy_community_import::LegacyCommunityImportOutcome; pub use operation::OperationSubscription; -pub use query::{MysqlConsoleCancellation, MysqlConsoleRequest, MysqlConsoleResult}; pub use query::{NativeConsoleCancellation, NativeConsoleRequest, NativeConsoleResult}; pub use transfer::TransferArtifactDownload; @@ -521,6 +520,7 @@ impl Application { pub fn list_drivers(&self) -> JdbcDriverList { let mut items = self.inner.drivers.clone(); for descriptor in self.inner.native_drivers.descriptors() { + let descriptor = datasource_compatibility::jdbc_driver_from_descriptor(descriptor); if !items .iter() .any(|driver| driver.driver_id.eq_ignore_ascii_case(&descriptor.driver_id)) @@ -550,7 +550,7 @@ impl Application { )); } self.require_managed_driver(driver_id)?; - if let Some(driver) = self.native_driver_for_driver_id(driver_id) { + if let Some(driver) = self.native_driver_for_datasource_driver_id(driver_id) { return driver.connection().test_connection(&connection).await; } let engine = self.require_engine().await?; @@ -660,7 +660,10 @@ impl Application { } fn require_managed_driver(&self, driver_id: &str) -> Result<(), AppError> { - if self.native_driver_for_driver_id(driver_id).is_some() { + if self + .native_driver_for_datasource_driver_id(driver_id) + .is_some() + { return Ok(()); } match &self.inner.managed_driver_ids { @@ -669,13 +672,15 @@ impl Application { } } - pub(crate) fn native_driver_for_driver_id( + pub(crate) fn native_driver_for_datasource_driver_id( &self, - driver_id: &str, + datasource_driver_id: &str, ) -> Option> { - self.inner - .native_drivers - .driver_for_driver_id(driver_id, &self.inner.drivers) + datasource_compatibility::native_driver_for_datasource_driver_id( + &self.inner.native_drivers, + datasource_driver_id, + &self.inner.drivers, + ) } pub(crate) fn native_driver_for_database_type( @@ -687,9 +692,12 @@ impl Application { .driver_for_database_type(database_type) } - pub(crate) fn native_database_type_for_driver(&self, driver_id: &str) -> Option { - self.native_driver_for_driver_id(driver_id) - .and_then(|driver| driver.database_types().first().copied()) + pub(crate) fn native_database_type_for_datasource_driver_id( + &self, + datasource_driver_id: &str, + ) -> Option { + self.native_driver_for_datasource_driver_id(datasource_driver_id) + .and_then(|driver| driver.descriptor().database_types.first().copied()) .map(str::to_owned) } @@ -700,7 +708,7 @@ impl Application { let storage = self.require_storage()?; let resolved = datasource_session::resolve_datasource_connection(&storage, datasource_id).await?; - self.native_driver_for_driver_id(&resolved.driver_id) + self.native_driver_for_datasource_driver_id(&resolved.driver_id) .ok_or_else(|| { AppError::invalid( "native_driver_not_available", @@ -715,7 +723,10 @@ impl Application { datasource_id: &str, driver_id: &str, ) -> Result<(), AppError> { - if self.native_driver_for_driver_id(driver_id).is_some() { + if self + .native_driver_for_datasource_driver_id(driver_id) + .is_some() + { return Ok(()); } let Some(driver_ids) = &self.inner.managed_driver_ids else { diff --git a/crates/chat2db-core/src/mysql_dashboard.rs b/crates/chat2db-core/src/mysql_dashboard.rs index 154a19c..cd0d237 100644 --- a/crates/chat2db-core/src/mysql_dashboard.rs +++ b/crates/chat2db-core/src/mysql_dashboard.rs @@ -2,9 +2,8 @@ use std::collections::HashMap; use chat2db_contract::{ ApiError, CommunityChart, CommunityDashboard, CommunityDashboardListQuery, - CommunityDashboardPage, CommunityTableColumn, CreateCommunityChartRequest, - CreateCommunityDashboardRequest, JdbcValue, ListCommunityColumnsRequest, ResultColumn, - UpdateCommunityChartRequest, UpdateCommunityDashboardRequest, + CommunityDashboardPage, CreateCommunityChartRequest, CreateCommunityDashboardRequest, + JdbcValue, ResultColumn, UpdateCommunityChartRequest, UpdateCommunityDashboardRequest, }; use chat2db_storage::CreateOperationLog; use serde_json::{Map, Value, json}; @@ -16,7 +15,9 @@ use sqlparser::{ use crate::{ AppError, AppErrorKind, Application, NativeConsoleCancellation, NativeConsoleRequest, - NativeConsoleResult, now_millis, storage_call, + NativeConsoleResult, + native_driver_types::{ColumnMetadata, ListColumnsRequest, MetadataScope, TableRef}, + now_millis, storage_call, }; const CHART_PAGE_SIZE: u32 = 200; @@ -193,7 +194,7 @@ impl Application { return Err(error); } }; - let header_metadata = self.chart_header_metadata(&context, &database_type).await; + let header_metadata = self.chart_header_metadata(&context).await; chart.meta_data = Some(chart_metadata(&result, header_metadata.as_ref())?); self.record_chart_history(&chart, &context, Some(&database_type), Some(&result), None) .await; @@ -206,7 +207,8 @@ impl Application { ) -> Result { self.require_native_driver_for_datasource(&context.datasource_id) .await? - .database_types() + .descriptor() + .database_types .first() .copied() .map(str::to_owned) @@ -297,8 +299,7 @@ impl Application { async fn chart_header_metadata( &self, context: &ChartRefreshContext, - database_type: &str, - ) -> Option> { + ) -> Option> { let table = chart_editable_table(&context.sql)?; let database_name = table .database_name @@ -307,15 +308,33 @@ impl Application { .schema_name .clone() .unwrap_or_else(|| database_name.clone()); - let columns = match self - .list_community_columns(ListCommunityColumnsRequest { - datasource_id: context.datasource_id.clone(), - database_type: database_type.to_owned(), - database_name, - schema_name, - table_name: table.table_name, - }) - .await + let columns = match async { + let driver = self + .require_native_driver_for_datasource(&context.datasource_id) + .await?; + let metadata = driver.metadata().ok_or_else(|| { + AppError::invalid( + "native_metadata_capability_not_available", + "The native Rust driver does not implement metadata operations", + ) + })?; + metadata + .list_columns( + self, + ListColumnsRequest { + table: TableRef { + scope: MetadataScope { + datasource_id: context.datasource_id.clone(), + database_name, + schema_name, + }, + table_name: table.table_name, + }, + }, + ) + .await + } + .await { Ok(columns) => columns, Err(error) => { @@ -450,7 +469,7 @@ fn non_editable_projection(item: &SelectItem) -> bool { fn chart_metadata( result: &NativeConsoleResult, - header_metadata: Option<&HashMap>, + header_metadata: Option<&HashMap>, ) -> Result { let metadata = json!({ "dataList": result @@ -483,7 +502,7 @@ fn chart_metadata( Ok(metadata) } -fn chart_header(column: &ResultColumn, metadata: Option<&CommunityTableColumn>) -> Value { +fn chart_header(column: &ResultColumn, metadata: Option<&ColumnMetadata>) -> Value { let column_type = metadata.map_or(column.jdbc_type_name.as_str(), |column| { column.column_type.as_str() }); @@ -581,11 +600,11 @@ fn chart_editor_type(type_name: &str, jdbc_type: i32) -> &'static str { mod tests { use std::collections::HashMap; - use chat2db_contract::{ - ColumnNullability, CommunityTableColumn, JdbcValue, JdbcValueType, ResultColumn, ResultRow, - }; + use chat2db_contract::{ColumnNullability, JdbcValue, JdbcValueType, ResultColumn, ResultRow}; use serde_json::json; + use crate::native_driver_types::ColumnMetadata; + use super::{chart_data_type, chart_editable_table, chart_metadata, chart_refresh_context}; #[test] @@ -691,7 +710,7 @@ mod tests { }; let header_metadata = HashMap::from([( "amount".to_owned(), - CommunityTableColumn { + ColumnMetadata { name: "amount".to_owned(), column_type: "DECIMAL".to_owned(), auto_increment: Some(false), @@ -700,7 +719,7 @@ mod tests { column_size: Some(10), decimal_digits: Some(2), nullable: Some(1), - ..CommunityTableColumn::default() + ..ColumnMetadata::default() }, )]); let metadata = chart_metadata(&result, Some(&header_metadata)).expect("chart metadata"); diff --git a/crates/chat2db-core/src/mysql_workspace.rs b/crates/chat2db-core/src/mysql_workspace.rs index 8d95e7b..db3db7a 100644 --- a/crates/chat2db-core/src/mysql_workspace.rs +++ b/crates/chat2db-core/src/mysql_workspace.rs @@ -99,6 +99,7 @@ impl Application { &request.schema_name, ) .await?; + let tables = crate::native_api_adapter::er_tables_response(tables); let storage = self.require_storage()?; let position = storage_call(move || { storage.mysql_er_position( diff --git a/crates/chat2db-core/src/native_api_adapter.rs b/crates/chat2db-core/src/native_api_adapter.rs index 6ecc7e6..e705006 100644 --- a/crates/chat2db-core/src/native_api_adapter.rs +++ b/crates/chat2db-core/src/native_api_adapter.rs @@ -4,8 +4,23 @@ use chat2db_contract::{ CommunityAccount, CommunityAccountAction, CommunityAccountCapability, CommunityAccountCommandRequest, CommunityAccountExecution, CommunityAccountGrantList, CommunityAccountGrantsRequest, CommunityAccountList, CommunityAccountPreview, - CommunityAccountPrivilegeScope, CommunitySchemaDiffEndpoint, CommunitySchemaDiffRequest, - CommunitySchemaDiffSql, + CommunityAccountPrivilegeScope, CommunityDatabase, CommunityDatabaseList, CommunityErColumn, + CommunityErForeignKey, CommunityErTable, CommunityForeignKey, CommunityForeignKeyList, + CommunityFunction, CommunityFunctionList, CommunityFunctionParameter, + CommunityFunctionParameterList, CommunityPrimaryKey, CommunityPrimaryKeyList, + CommunityProcedure, CommunityProcedureList, CommunityProcedureParameter, + CommunityProcedureParameterList, CommunityRoutineInvocationPreview, + CommunityRoutineMigrationExecution, CommunityRoutineMigrationRequest, CommunitySchema, + CommunitySchemaDiffEndpoint, CommunitySchemaDiffRequest, CommunitySchemaDiffSql, + CommunitySchemaList, CommunityTable, CommunityTableColumn, CommunityTableColumnList, + CommunityTableIndex, CommunityTableIndexColumn, CommunityTableIndexList, CommunityTableList, + CommunityTablePreviewAccepted, CommunityTrigger, CommunityTriggerList, CommunityViewList, + GetCommunityFunctionRequest, GetCommunityProcedureRequest, GetCommunityTriggerRequest, + ListCommunityColumnsRequest, ListCommunityDatabasesRequest, ListCommunityFunctionsRequest, + ListCommunityIndexesRequest, ListCommunityProceduresRequest, ListCommunitySchemasRequest, + ListCommunityTableKeysRequest, ListCommunityTablesRequest, ListCommunityTriggersRequest, + ListCommunityViewsRequest, PreviewCommunityRoutineInvocationRequest, + StartCommunityTablePreviewRequest, }; use crate::{ @@ -15,6 +30,19 @@ use crate::{ AdministrationExecution, AdministrationPreview, Principal, PrincipalGrantList, PrincipalGrantsRequest, PrincipalList, PrincipalRef, PrivilegeScope, PrivilegeTarget, }, + native_driver_types::{ + ColumnList, ColumnMetadata, DatabaseList, DatabaseMetadata, EntityRelationColumn, + EntityRelationForeignKey, EntityRelationTable, ForeignKeyList, ForeignKeyMetadata, + FunctionList, FunctionMetadata, FunctionParameterList, FunctionParameterMetadata, + IndexColumnMetadata, IndexList, IndexMetadata, ListColumnsRequest, ListDatabasesRequest, + ListIndexesRequest, ListRoutinesRequest, ListSchemasRequest, ListTableKeysRequest, + ListTablesRequest, ListTriggersRequest, ListViewsRequest, MetadataObjectRef, MetadataScope, + PrimaryKeyList, PrimaryKeyMetadata, ProcedureList, ProcedureMetadata, + ProcedureParameterList, ProcedureParameterMetadata, RoutineInvocationPreview, + RoutineInvocationRequest, RoutineMigrationExecution, RoutineMigrationRequest, SchemaList, + SchemaMetadata, TableList, TableMetadata, TablePreviewAccepted, TablePreviewRequest, + TableRef, TriggerList, TriggerMetadata, ViewList, + }, native_schema_diff_types::{SchemaDiffEndpoint, SchemaDiffRequest, SchemaDiffSql}, }; @@ -141,7 +169,11 @@ impl Application { let target_driver = self .require_native_driver_for_datasource(&request.target.datasource_id) .await?; - if !source_driver.id().eq_ignore_ascii_case(target_driver.id()) { + if !source_driver + .descriptor() + .id + .eq_ignore_ascii_case(target_driver.descriptor().id) + { return Err(AppError::invalid( "invalid_community_schema_diff_request", "Source and target datasources must use the same native driver", @@ -158,6 +190,654 @@ impl Application { } } +impl From for ListDatabasesRequest { + fn from(request: ListCommunityDatabasesRequest) -> Self { + Self { + datasource_id: request.datasource_id, + } + } +} + +impl From for ListSchemasRequest { + fn from(request: ListCommunitySchemasRequest) -> Self { + Self { + datasource_id: request.datasource_id, + database_name: request.database_name, + } + } +} + +impl From for ListTablesRequest { + fn from(request: ListCommunityTablesRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + name_pattern: request.table_name_pattern, + } + } +} + +impl From for ListColumnsRequest { + fn from(request: ListCommunityColumnsRequest) -> Self { + Self { + table: TableRef { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + table_name: request.table_name, + }, + } + } +} + +impl From for ListIndexesRequest { + fn from(request: ListCommunityIndexesRequest) -> Self { + Self { + table: TableRef { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + table_name: request.table_name, + }, + } + } +} + +impl From for ListViewsRequest { + fn from(request: ListCommunityViewsRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + name_pattern: request.view_name_pattern, + } + } +} + +impl From for MetadataObjectRef { + fn from(request: ListCommunityViewsRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + object_name: request.view_name_pattern, + } + } +} + +impl From for ListTableKeysRequest { + fn from(request: ListCommunityTableKeysRequest) -> Self { + Self { + table: TableRef { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + table_name: request.table_name, + }, + } + } +} + +impl From for ListRoutinesRequest { + fn from(request: ListCommunityFunctionsRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + } + } +} + +impl From for MetadataObjectRef { + fn from(request: GetCommunityFunctionRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + object_name: request.function_name, + } + } +} + +impl From for ListRoutinesRequest { + fn from(request: ListCommunityProceduresRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + } + } +} + +impl From for MetadataObjectRef { + fn from(request: GetCommunityProcedureRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + object_name: request.procedure_name, + } + } +} + +impl From for ListTriggersRequest { + fn from(request: ListCommunityTriggersRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + } + } +} + +impl From for MetadataObjectRef { + fn from(request: GetCommunityTriggerRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + object_name: request.trigger_name, + } + } +} + +impl From for TablePreviewRequest { + fn from(request: StartCommunityTablePreviewRequest) -> Self { + Self { + table: TableRef { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + table_name: request.table_name, + }, + } + } +} + +impl From for RoutineInvocationRequest { + fn from(request: PreviewCommunityRoutineInvocationRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + routine_type: request.routine_type, + routine_name: request.routine_name, + } + } +} + +impl From for RoutineMigrationRequest { + fn from(request: CommunityRoutineMigrationRequest) -> Self { + Self { + scope: MetadataScope { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + }, + database_type: request.database_type, + routine_type: request.routine_type, + routine_name: request.routine_name, + ddl: request.ddl, + } + } +} + +impl From for CommunityRoutineMigrationRequest { + fn from(request: RoutineMigrationRequest) -> Self { + Self { + datasource_id: request.scope.datasource_id, + database_type: request.database_type, + database_name: request.scope.database_name, + schema_name: request.scope.schema_name, + routine_type: request.routine_type, + routine_name: request.routine_name, + ddl: request.ddl, + } + } +} + +pub(crate) fn schema_list_response(response: SchemaList) -> CommunitySchemaList { + CommunitySchemaList { + items: response.items.into_iter().map(schema_response).collect(), + } +} + +pub(crate) fn database_list_response(response: DatabaseList) -> CommunityDatabaseList { + CommunityDatabaseList { + items: response.items.into_iter().map(database_response).collect(), + } +} + +pub(crate) fn table_list_response(response: TableList) -> CommunityTableList { + CommunityTableList { + items: response.items.into_iter().map(table_response).collect(), + } +} + +pub(crate) fn column_list_response(response: ColumnList) -> CommunityTableColumnList { + CommunityTableColumnList { + items: response.items.into_iter().map(column_response).collect(), + } +} + +pub(crate) fn index_list_response(response: IndexList) -> CommunityTableIndexList { + CommunityTableIndexList { + items: response.items.into_iter().map(index_response).collect(), + } +} + +pub(crate) fn view_list_response(response: ViewList) -> CommunityViewList { + CommunityViewList { + items: response.items.into_iter().map(table_response).collect(), + } +} + +pub(crate) fn table_response(response: TableMetadata) -> CommunityTable { + CommunityTable { + database_name: response.database_name, + schema_name: response.schema_name, + name: response.name, + table_type: response.table_type, + comment: response.comment, + database_type: response.database_type, + pinned: response.pinned, + ddl: response.ddl, + engine: response.engine, + charset: response.charset, + collation: response.collation, + increment_value: response.increment_value, + partition: response.partition, + tablespace: response.tablespace, + rows: response.rows, + data_length: response.data_length, + create_time: response.create_time, + update_time: response.update_time, + } +} + +pub(crate) fn foreign_key_list_response(response: ForeignKeyList) -> CommunityForeignKeyList { + CommunityForeignKeyList { + items: response + .items + .into_iter() + .map(foreign_key_response) + .collect(), + } +} + +pub(crate) fn primary_key_list_response(response: PrimaryKeyList) -> CommunityPrimaryKeyList { + CommunityPrimaryKeyList { + items: response + .items + .into_iter() + .map(primary_key_response) + .collect(), + } +} + +pub(crate) fn function_list_response(response: FunctionList) -> CommunityFunctionList { + CommunityFunctionList { + items: response.items.into_iter().map(function_response).collect(), + } +} + +pub(crate) fn function_response(response: FunctionMetadata) -> CommunityFunction { + CommunityFunction { + database_name: response.database_name, + schema_name: response.schema_name, + name: response.name, + remarks: response.remarks, + function_type: response.function_type, + specific_name: response.specific_name, + body: response.body, + template: response.template, + } +} + +pub(crate) fn function_parameter_list_response( + response: FunctionParameterList, +) -> CommunityFunctionParameterList { + CommunityFunctionParameterList { + items: response + .items + .into_iter() + .map(function_parameter_response) + .collect(), + } +} + +pub(crate) fn procedure_list_response(response: ProcedureList) -> CommunityProcedureList { + CommunityProcedureList { + items: response.items.into_iter().map(procedure_response).collect(), + } +} + +pub(crate) fn procedure_response(response: ProcedureMetadata) -> CommunityProcedure { + CommunityProcedure { + database_name: response.database_name, + schema_name: response.schema_name, + name: response.name, + remarks: response.remarks, + procedure_type: response.procedure_type, + specific_name: response.specific_name, + body: response.body, + } +} + +pub(crate) fn procedure_parameter_list_response( + response: ProcedureParameterList, +) -> CommunityProcedureParameterList { + CommunityProcedureParameterList { + items: response + .items + .into_iter() + .map(procedure_parameter_response) + .collect(), + } +} + +pub(crate) fn trigger_list_response(response: TriggerList) -> CommunityTriggerList { + CommunityTriggerList { + items: response.items.into_iter().map(trigger_response).collect(), + } +} + +pub(crate) fn trigger_response(response: TriggerMetadata) -> CommunityTrigger { + CommunityTrigger { + database_name: response.database_name, + schema_name: response.schema_name, + name: response.name, + event_manipulation: response.event_manipulation, + body: response.body, + } +} + +pub(crate) fn er_tables_response(response: Vec) -> Vec { + response.into_iter().map(er_table_response).collect() +} + +pub(crate) fn table_preview_response( + response: TablePreviewAccepted, +) -> CommunityTablePreviewAccepted { + CommunityTablePreviewAccepted { + operation_id: response.operation_id, + sql: response.sql, + row_limit: response.row_limit, + } +} + +pub(crate) fn routine_invocation_response( + response: RoutineInvocationPreview, +) -> CommunityRoutineInvocationPreview { + CommunityRoutineInvocationPreview { sql: response.sql } +} + +pub(crate) fn routine_migration_execution_response( + response: RoutineMigrationExecution, +) -> CommunityRoutineMigrationExecution { + CommunityRoutineMigrationExecution { + success: response.success, + message: response.message, + sql: response.sql, + failure_stage: response.failure_stage, + restore_attempted: response.restore_attempted, + restore_succeeded: response.restore_succeeded, + } +} + +pub(crate) fn compatibility_api_error(error: AppError) -> AppError { + let api = error.api_error(); + let compatibility_code = match api.code.as_str() { + "invalid_routine_invocation_request" => "invalid_community_routine_invocation_request", + "invalid_routine_migration_request" => "invalid_community_routine_migration_request", + "invalid_table_preview_request" => "invalid_community_table_preview_request", + _ => return error, + }; + AppError::invalid(compatibility_code, api.message) +} + +fn schema_response(response: SchemaMetadata) -> CommunitySchema { + CommunitySchema { + database_name: response.database_name, + name: response.name, + comment: response.comment, + owner: response.owner, + system: response.system, + } +} + +fn database_response(response: DatabaseMetadata) -> CommunityDatabase { + CommunityDatabase { + name: response.name, + comment: response.comment, + charset: response.charset, + collation: response.collation, + owner: response.owner, + system: response.system, + } +} + +fn column_response(response: ColumnMetadata) -> CommunityTableColumn { + CommunityTableColumn { + database_name: response.database_name, + schema_name: response.schema_name, + table_name: response.table_name, + name: response.name, + column_type: response.column_type, + data_type: response.data_type, + default_value: response.default_value, + auto_increment: response.auto_increment, + comment: response.comment, + primary_key: response.primary_key, + primary_key_name: response.primary_key_name, + primary_key_order: response.primary_key_order, + column_size: response.column_size, + buffer_length: response.buffer_length, + decimal_digits: response.decimal_digits, + num_prec_radix: response.num_prec_radix, + sql_data_type: response.sql_data_type, + sql_datetime_sub: response.sql_datetime_sub, + char_octet_length: response.char_octet_length, + ordinal_position: response.ordinal_position, + nullable: response.nullable, + generated_column: response.generated_column, + extent: response.extent, + charset: response.charset, + collation: response.collation, + unit: response.unit, + sparse: response.sparse, + default_constraint_name: response.default_constraint_name, + seed: response.seed, + increment: response.increment, + on_update_current_timestamp: response.on_update_current_timestamp, + } +} + +fn index_column_response(response: IndexColumnMetadata) -> CommunityTableIndexColumn { + CommunityTableIndexColumn { + database_name: response.database_name, + schema_name: response.schema_name, + table_name: response.table_name, + index_name: response.index_name, + column_name: response.column_name, + column_type: response.column_type, + comment: response.comment, + ordinal_position: response.ordinal_position, + collation: response.collation, + non_unique: response.non_unique, + index_qualifier: response.index_qualifier, + sort_order: response.sort_order, + cardinality: response.cardinality, + pages: response.pages, + filter_condition: response.filter_condition, + sub_part: response.sub_part, + } +} + +fn index_response(response: IndexMetadata) -> CommunityTableIndex { + CommunityTableIndex { + database_name: response.database_name, + schema_name: response.schema_name, + table_name: response.table_name, + name: response.name, + index_type: response.index_type, + unique: response.unique, + comment: response.comment, + columns: response + .columns + .into_iter() + .map(index_column_response) + .collect(), + concurrently: response.concurrently, + method: response.method, + foreign_schema_name: response.foreign_schema_name, + foreign_table_name: response.foreign_table_name, + foreign_column_names: response.foreign_column_names, + } +} + +fn foreign_key_response(response: ForeignKeyMetadata) -> CommunityForeignKey { + CommunityForeignKey { + primary_table_database: response.primary_table_database, + primary_table_schema: response.primary_table_schema, + primary_table_name: response.primary_table_name, + primary_column_name: response.primary_column_name, + foreign_table_database: response.foreign_table_database, + foreign_table_schema: response.foreign_table_schema, + foreign_table_name: response.foreign_table_name, + foreign_column_name: response.foreign_column_name, + key_sequence: response.key_sequence, + update_rule: response.update_rule, + delete_rule: response.delete_rule, + foreign_key_name: response.foreign_key_name, + primary_key_name: response.primary_key_name, + deferrability: response.deferrability, + } +} + +fn primary_key_response(response: PrimaryKeyMetadata) -> CommunityPrimaryKey { + CommunityPrimaryKey { + database_name: response.database_name, + schema_name: response.schema_name, + table_name: response.table_name, + column_name: response.column_name, + name: response.name, + } +} + +fn function_parameter_response(response: FunctionParameterMetadata) -> CommunityFunctionParameter { + CommunityFunctionParameter { + function_database: response.function_database, + function_schema: response.function_schema, + function_name: response.function_name, + column_name: response.column_name, + column_type: response.column_type, + data_type: response.data_type, + type_name: response.type_name, + precision: response.precision, + length: response.length, + scale: response.scale, + radix: response.radix, + nullable: response.nullable, + remarks: response.remarks, + char_octet_length: response.char_octet_length, + ordinal_position: response.ordinal_position, + is_nullable: response.is_nullable, + specific_name: response.specific_name, + } +} + +fn procedure_parameter_response( + response: ProcedureParameterMetadata, +) -> CommunityProcedureParameter { + CommunityProcedureParameter { + procedure_database: response.procedure_database, + procedure_schema: response.procedure_schema, + procedure_name: response.procedure_name, + column_name: response.column_name, + column_type: response.column_type, + data_type: response.data_type, + type_name: response.type_name, + precision: response.precision, + length: response.length, + scale: response.scale, + radix: response.radix, + nullable: response.nullable, + remarks: response.remarks, + column_default: response.column_default, + sql_data_type: response.sql_data_type, + sql_datetime_sub: response.sql_datetime_sub, + char_octet_length: response.char_octet_length, + ordinal_position: response.ordinal_position, + is_nullable: response.is_nullable, + specific_name: response.specific_name, + } +} + +fn er_table_response(response: EntityRelationTable) -> CommunityErTable { + CommunityErTable { + name: response.name, + comment: response.comment, + column_list: response + .columns + .into_iter() + .map(er_column_response) + .collect(), + foreign_key_list: response + .foreign_keys + .into_iter() + .map(er_foreign_key_response) + .collect(), + } +} + +fn er_column_response(response: EntityRelationColumn) -> CommunityErColumn { + CommunityErColumn { + name: response.name, + column_type: response.column_type, + primary_key: response.primary_key, + comment: response.comment, + } +} + +fn er_foreign_key_response(response: EntityRelationForeignKey) -> CommunityErForeignKey { + CommunityErForeignKey { + pk_table_name: response.primary_table, + pk_column_name: response.primary_column, + fk_table_name: response.foreign_table, + fk_column_name: response.foreign_column, + } +} + fn administration_unavailable() -> AppError { AppError::invalid( "native_administration_capability_not_available", @@ -344,9 +1024,237 @@ fn community_schema_diff_error(error: AppError) -> AppError { #[cfg(test)] mod tests { - use chat2db_contract::{CommunityAccountAction, CommunityAccountCommandRequest}; + use chat2db_contract::{ + CommunityAccountAction, CommunityAccountCommandRequest, CommunityErColumn, + CommunityErForeignKey, CommunityErTable, CommunityTableColumn, CommunityTableColumnList, + CommunityTableIndex, CommunityTableIndexColumn, CommunityTableIndexList, + }; - use super::{administration_command, community_account_action}; + use super::{ + administration_command, column_list_response, community_account_action, + compatibility_api_error, er_tables_response, index_list_response, + }; + use crate::{ + AppError, + native_driver_types::{ + ColumnList, ColumnMetadata, EntityRelationColumn, EntityRelationForeignKey, + EntityRelationTable, IndexColumnMetadata, IndexList, IndexMetadata, + }, + }; + + #[test] + fn neutral_validation_errors_preserve_compatibility_wire_codes() { + for (neutral, legacy) in [ + ( + "invalid_routine_invocation_request", + "invalid_community_routine_invocation_request", + ), + ( + "invalid_routine_migration_request", + "invalid_community_routine_migration_request", + ), + ( + "invalid_table_preview_request", + "invalid_community_table_preview_request", + ), + ] { + let error = compatibility_api_error(AppError::invalid(neutral, "invalid request")); + let api = error.api_error(); + assert_eq!(api.code, legacy); + assert_eq!(api.message, "invalid request"); + } + } + + #[test] + fn column_metadata_maps_every_field_to_the_legacy_response() { + let response = column_list_response(ColumnList { + items: vec![ColumnMetadata { + database_name: "catalog".to_owned(), + schema_name: "schema".to_owned(), + table_name: "orders".to_owned(), + name: "total".to_owned(), + column_type: "DECIMAL(12,2)".to_owned(), + data_type: Some(3), + default_value: Some("0.00".to_owned()), + auto_increment: Some(false), + comment: "order total".to_owned(), + primary_key: Some(false), + primary_key_name: "pk_orders".to_owned(), + primary_key_order: 2, + column_size: Some(12), + buffer_length: Some(16), + decimal_digits: Some(2), + num_prec_radix: Some(10), + sql_data_type: Some(3), + sql_datetime_sub: Some(4), + char_octet_length: Some(48), + ordinal_position: Some(5), + nullable: Some(1), + generated_column: Some(false), + extent: "extent".to_owned(), + charset: "utf8mb4".to_owned(), + collation: "utf8mb4_bin".to_owned(), + unit: "bytes".to_owned(), + sparse: Some(true), + default_constraint_name: "df_orders_total".to_owned(), + seed: Some(7), + increment: Some(8), + on_update_current_timestamp: Some(true), + }], + }); + + assert_eq!( + response, + CommunityTableColumnList { + items: vec![CommunityTableColumn { + database_name: "catalog".to_owned(), + schema_name: "schema".to_owned(), + table_name: "orders".to_owned(), + name: "total".to_owned(), + column_type: "DECIMAL(12,2)".to_owned(), + data_type: Some(3), + default_value: Some("0.00".to_owned()), + auto_increment: Some(false), + comment: "order total".to_owned(), + primary_key: Some(false), + primary_key_name: "pk_orders".to_owned(), + primary_key_order: 2, + column_size: Some(12), + buffer_length: Some(16), + decimal_digits: Some(2), + num_prec_radix: Some(10), + sql_data_type: Some(3), + sql_datetime_sub: Some(4), + char_octet_length: Some(48), + ordinal_position: Some(5), + nullable: Some(1), + generated_column: Some(false), + extent: "extent".to_owned(), + charset: "utf8mb4".to_owned(), + collation: "utf8mb4_bin".to_owned(), + unit: "bytes".to_owned(), + sparse: Some(true), + default_constraint_name: "df_orders_total".to_owned(), + seed: Some(7), + increment: Some(8), + on_update_current_timestamp: Some(true), + }], + } + ); + } + + #[test] + fn index_metadata_preserves_nested_columns() { + let indexes = index_list_response(IndexList { + items: vec![IndexMetadata { + database_name: "catalog".to_owned(), + schema_name: "schema".to_owned(), + table_name: "orders".to_owned(), + name: "idx_orders_customer".to_owned(), + index_type: "BTREE".to_owned(), + unique: Some(true), + comment: "customer lookup".to_owned(), + columns: vec![IndexColumnMetadata { + database_name: "catalog".to_owned(), + schema_name: "schema".to_owned(), + table_name: "orders".to_owned(), + index_name: "idx_orders_customer".to_owned(), + column_name: "customer_id".to_owned(), + column_type: "BIGINT".to_owned(), + comment: "customer".to_owned(), + ordinal_position: Some(1), + collation: "A".to_owned(), + non_unique: Some(false), + index_qualifier: "catalog".to_owned(), + sort_order: "ASC".to_owned(), + cardinality: Some("99".to_owned()), + pages: Some("4".to_owned()), + filter_condition: "active = 1".to_owned(), + sub_part: Some("8".to_owned()), + }], + concurrently: Some(false), + method: "btree".to_owned(), + foreign_schema_name: "crm".to_owned(), + foreign_table_name: "customers".to_owned(), + foreign_column_names: vec!["id".to_owned()], + }], + }); + assert_eq!( + indexes, + CommunityTableIndexList { + items: vec![CommunityTableIndex { + database_name: "catalog".to_owned(), + schema_name: "schema".to_owned(), + table_name: "orders".to_owned(), + name: "idx_orders_customer".to_owned(), + index_type: "BTREE".to_owned(), + unique: Some(true), + comment: "customer lookup".to_owned(), + columns: vec![CommunityTableIndexColumn { + database_name: "catalog".to_owned(), + schema_name: "schema".to_owned(), + table_name: "orders".to_owned(), + index_name: "idx_orders_customer".to_owned(), + column_name: "customer_id".to_owned(), + column_type: "BIGINT".to_owned(), + comment: "customer".to_owned(), + ordinal_position: Some(1), + collation: "A".to_owned(), + non_unique: Some(false), + index_qualifier: "catalog".to_owned(), + sort_order: "ASC".to_owned(), + cardinality: Some("99".to_owned()), + pages: Some("4".to_owned()), + filter_condition: "active = 1".to_owned(), + sub_part: Some("8".to_owned()), + }], + concurrently: Some(false), + method: "btree".to_owned(), + foreign_schema_name: "crm".to_owned(), + foreign_table_name: "customers".to_owned(), + foreign_column_names: vec!["id".to_owned()], + }], + } + ); + } + + #[test] + fn er_metadata_preserves_nested_relationships() { + assert_eq!( + er_tables_response(vec![EntityRelationTable { + name: "orders".to_owned(), + comment: "orders table".to_owned(), + columns: vec![EntityRelationColumn { + name: "customer_id".to_owned(), + column_type: "BIGINT".to_owned(), + primary_key: false, + comment: "customer".to_owned(), + }], + foreign_keys: vec![EntityRelationForeignKey { + primary_table: "customers".to_owned(), + primary_column: "id".to_owned(), + foreign_table: "orders".to_owned(), + foreign_column: "customer_id".to_owned(), + }], + }]), + vec![CommunityErTable { + name: "orders".to_owned(), + comment: "orders table".to_owned(), + column_list: vec![CommunityErColumn { + name: "customer_id".to_owned(), + column_type: "BIGINT".to_owned(), + primary_key: false, + comment: "customer".to_owned(), + }], + foreign_key_list: vec![CommunityErForeignKey { + pk_table_name: "customers".to_owned(), + pk_column_name: "id".to_owned(), + fk_table_name: "orders".to_owned(), + fk_column_name: "customer_id".to_owned(), + }], + }] + ); + } #[test] fn account_actions_round_trip_across_the_compatibility_boundary() { diff --git a/crates/chat2db-core/src/native_driver.rs b/crates/chat2db-core/src/native_driver.rs index d757443..cc01d48 100644 --- a/crates/chat2db-core/src/native_driver.rs +++ b/crates/chat2db-core/src/native_driver.rs @@ -1,7 +1,12 @@ -use std::{collections::HashSet, sync::Arc}; +use std::{ + collections::{HashMap, HashSet}, + sync::Arc, +}; use async_trait::async_trait; -use chat2db_contract::{DatasourceConnection, JdbcDriver, ResultMetadata}; +use chat2db_contract::{ + DatasourceConnection, ImportFileRequest, ResultMetadata, SqlFileExportRequest, TransferArtifact, +}; use chat2db_storage::Storage; use tokio::sync::watch; use tokio_util::sync::CancellationToken; @@ -14,15 +19,14 @@ use crate::{ AdministrationPreview, PrincipalGrantList, PrincipalGrantsRequest, PrincipalList, }, native_driver_types::{ - BuiltSql, ColumnList, CreateSchemaSqlRequest, DatabaseList, DmlExportTransferRequest, - DmlSqlRequest, EntityRelationTable, ExportArtifact, ForeignKeyList, FunctionList, - FunctionMetadata, FunctionParameterList, ImportTransferRequest, IndexList, - ListColumnsRequest, ListDatabasesRequest, ListIndexesRequest, ListRoutinesRequest, - ListSchemasRequest, ListTableKeysRequest, ListTablesRequest, ListTriggersRequest, - ListViewsRequest, NamespaceSqlRequest, ObjectRef, OtherExportTransferRequest, - PrimaryKeyList, ProcedureList, ProcedureMetadata, ProcedureParameterList, - RoutineInvocationPreview, RoutineInvocationRequest, RoutineMigrationExecution, - RoutineMigrationRequest, SchemaList, SqlExportTransferRequest, TableList, TableMetadata, + BuiltSql, ColumnList, CreateSchemaSqlRequest, DatabaseList, DmlSqlRequest, + EntityRelationTable, ForeignKeyList, FunctionList, FunctionMetadata, FunctionParameterList, + IndexList, ListColumnsRequest, ListDatabasesRequest, ListIndexesRequest, + ListRoutinesRequest, ListSchemasRequest, ListTableKeysRequest, ListTablesRequest, + ListTriggersRequest, ListViewsRequest, MetadataObjectRef, NamespaceSqlRequest, + NativeDriverDescriptor, PrimaryKeyList, ProcedureList, ProcedureMetadata, + ProcedureParameterList, RoutineInvocationPreview, RoutineInvocationRequest, + RoutineMigrationExecution, RoutineMigrationRequest, SchemaList, TableList, TableMetadata, TablePreviewAccepted, TablePreviewRequest, TriggerList, TriggerMetadata, ViewList, }, native_mysql, @@ -32,6 +36,7 @@ use crate::{ DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, PreparedQuery, QueryTaskError, }, + transfer::{QueryResultExportRequest, TableFileExportRequest}, }; /// Database connection operations implemented by one native Rust driver. @@ -125,7 +130,7 @@ pub(crate) trait NativeMetadataDriver: Send + Sync { async fn get_view( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result; async fn list_imported_keys( @@ -155,13 +160,13 @@ pub(crate) trait NativeMetadataDriver: Send + Sync { async fn get_function( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result; async fn list_function_parameters( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result; async fn list_procedures( @@ -173,13 +178,13 @@ pub(crate) trait NativeMetadataDriver: Send + Sync { async fn get_procedure( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result; async fn list_procedure_parameters( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result; async fn list_triggers( @@ -191,7 +196,7 @@ pub(crate) trait NativeMetadataDriver: Send + Sync { async fn get_trigger( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result; } @@ -259,26 +264,26 @@ pub(crate) trait NativeTransferDriver: Send + Sync { async fn import_file( &self, application: &Application, - request: ImportTransferRequest, + request: ImportFileRequest, ) -> Result; async fn export_sql_file( &self, application: &Application, - request: SqlExportTransferRequest, + request: SqlFileExportRequest, ) -> Result; - async fn export_other_file( + async fn export_table_file( &self, application: &Application, - request: OtherExportTransferRequest, + request: TableFileExportRequest, ) -> Result; - async fn export_dml( + async fn export_query_result( &self, application: &Application, - request: DmlExportTransferRequest, - ) -> Result; + request: QueryResultExportRequest, + ) -> Result; } /// Structured SQL builders supplied by one native database dialect. @@ -340,15 +345,7 @@ pub(crate) trait NativeSchemaDiffDriver: Send + Sync { /// product surfaces it implements. Additional capability traits are attached /// here as native metadata and dialect services are migrated. pub(crate) trait NativeDriver: Send + Sync { - fn id(&self) -> &'static str; - - fn implementation(&self) -> &'static str; - - fn database_types(&self) -> &'static [&'static str]; - - fn descriptor(&self) -> JdbcDriver; - - fn matches_driver(&self, driver_id: &str, descriptor: Option<&JdbcDriver>) -> bool; + fn descriptor(&self) -> &'static NativeDriverDescriptor; fn connection(&self) -> &dyn NativeConnectionDriver; @@ -397,31 +394,59 @@ impl NativeDriverRegistry { .expect("built-in native drivers must have unique identities") } - fn try_new(drivers: Vec>) -> Result { - let mut ids = HashSet::new(); - let mut database_types = HashSet::new(); + pub(crate) fn try_new(drivers: Vec>) -> Result { + let mut driver_ids = HashSet::new(); + let mut identifier_owners = HashMap::::new(); for driver in &drivers { - let id = driver.id().trim().to_ascii_lowercase(); - if id.is_empty() || !ids.insert(id) { + let descriptor = driver.descriptor(); + let trimmed_driver_id = descriptor.id.trim(); + let driver_id = trimmed_driver_id.to_ascii_lowercase(); + if driver_id.is_empty() + || descriptor.id != trimmed_driver_id + || !driver_ids.insert(driver_id.clone()) + { return Err(AppError::invalid( "invalid_native_driver_registry", - "native driver ids must be non-empty and unique", + "native driver ids must be trimmed, non-empty, and unique", )); } - if driver.database_types().is_empty() { + if descriptor.implementation.is_empty() + || descriptor.implementation != descriptor.implementation.trim() + { + return Err(AppError::invalid( + "invalid_native_driver_registry", + "native driver implementation names must be trimmed and non-empty", + )); + } + if descriptor.database_types.is_empty() { return Err(AppError::invalid( "invalid_native_driver_registry", "native drivers must declare at least one database type", )); } - for database_type in driver.database_types() { - let database_type = database_type.trim().to_ascii_lowercase(); - if database_type.is_empty() || !database_types.insert(database_type) { + + for identifier in std::iter::once(descriptor.id) + .chain(descriptor.database_types.iter().copied()) + .chain(descriptor.compatibility_aliases.iter().copied()) + { + let trimmed_identifier = identifier.trim(); + if trimmed_identifier.is_empty() || identifier != trimmed_identifier { return Err(AppError::invalid( "invalid_native_driver_registry", - "native database types must be non-empty and unique", + "native driver identifiers and aliases must be trimmed and non-empty", )); } + let identifier = trimmed_identifier.to_ascii_lowercase(); + if let Some(owner) = identifier_owners.get(&identifier) { + if owner != &driver_id { + return Err(AppError::invalid( + "invalid_native_driver_registry", + "native driver identifiers and aliases must have one owner", + )); + } + } else { + identifier_owners.insert(identifier, driver_id.clone()); + } } } Ok(Self { @@ -429,71 +454,58 @@ impl NativeDriverRegistry { }) } - pub(crate) fn descriptors(&self) -> impl Iterator + '_ { + pub(crate) fn descriptors(&self) -> impl Iterator + '_ { self.drivers.iter().map(|driver| driver.descriptor()) } - pub(crate) fn driver_for_database_type( + /// Resolves a persisted datasource driver ID to its native implementation. + pub(crate) fn driver_for_datasource_driver_id( &self, - database_type: &str, + datasource_driver_id: &str, ) -> Option> { + let datasource_driver_id = datasource_driver_id.trim(); self.drivers .iter() .find(|driver| { - driver - .database_types() - .iter() - .any(|candidate| candidate.eq_ignore_ascii_case(database_type.trim())) + let descriptor = driver.descriptor(); + descriptor.id.eq_ignore_ascii_case(datasource_driver_id) + || descriptor + .compatibility_aliases + .iter() + .any(|alias| alias.eq_ignore_ascii_case(datasource_driver_id)) }) .cloned() } - pub(crate) fn driver_for_driver_id( + pub(crate) fn driver_for_database_type( &self, - driver_id: &str, - managed_drivers: &[JdbcDriver], + database_type: &str, ) -> Option> { - let descriptor = managed_drivers - .iter() - .find(|driver| driver.driver_id == driver_id); self.drivers .iter() - .find(|driver| driver.matches_driver(driver_id, descriptor)) + .find(|driver| { + driver + .descriptor() + .database_types + .iter() + .any(|candidate| candidate.eq_ignore_ascii_case(database_type.trim())) + }) .cloned() } } struct MysqlNativeDriver; -impl NativeDriver for MysqlNativeDriver { - fn id(&self) -> &'static str { - "mysql" - } - - fn implementation(&self) -> &'static str { - "mysql_async" - } - - fn database_types(&self) -> &'static [&'static str] { - &["MYSQL"] - } - - fn descriptor(&self) -> JdbcDriver { - crate::datasource_compatibility::native_mysql_driver() - } +const MYSQL_DRIVER_DESCRIPTOR: NativeDriverDescriptor = NativeDriverDescriptor { + id: "mysql", + implementation: "mysql_async", + database_types: &["MYSQL"], + compatibility_aliases: &["mysql", "mysql_async", "com.mysql"], +}; - fn matches_driver(&self, driver_id: &str, descriptor: Option<&JdbcDriver>) -> bool { - if driver_id.eq_ignore_ascii_case(self.id()) { - return true; - } - descriptor.is_some_and(|driver| { - format!( - "{} {} {} {}", - driver.pack_id, driver.name, driver.driver_id, driver.driver_class - ) - .to_ascii_lowercase() - .contains("mysql") - }) +impl NativeDriver for MysqlNativeDriver { + fn descriptor(&self) -> &'static NativeDriverDescriptor { + &MYSQL_DRIVER_DESCRIPTOR } fn connection(&self) -> &dyn NativeConnectionDriver { @@ -677,14 +689,14 @@ impl NativeMetadataDriver for MysqlNativeDriver { async fn get_view( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result { native_mysql::get_view( application, &request.scope.datasource_id, &request.scope.database_name, &request.scope.schema_name, - &request.name, + &request.object_name, ) .await } @@ -749,14 +761,14 @@ impl NativeMetadataDriver for MysqlNativeDriver { async fn get_function( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result { native_mysql::get_function( application, &request.scope.datasource_id, &request.scope.database_name, &request.scope.schema_name, - &request.name, + &request.object_name, ) .await } @@ -764,14 +776,14 @@ impl NativeMetadataDriver for MysqlNativeDriver { async fn list_function_parameters( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result { native_mysql::list_function_parameters( application, &request.scope.datasource_id, &request.scope.database_name, &request.scope.schema_name, - &request.name, + &request.object_name, ) .await } @@ -793,14 +805,14 @@ impl NativeMetadataDriver for MysqlNativeDriver { async fn get_procedure( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result { native_mysql::get_procedure( application, &request.scope.datasource_id, &request.scope.database_name, &request.scope.schema_name, - &request.name, + &request.object_name, ) .await } @@ -808,14 +820,14 @@ impl NativeMetadataDriver for MysqlNativeDriver { async fn list_procedure_parameters( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result { native_mysql::list_procedure_parameters( application, &request.scope.datasource_id, &request.scope.database_name, &request.scope.schema_name, - &request.name, + &request.object_name, ) .await } @@ -837,14 +849,14 @@ impl NativeMetadataDriver for MysqlNativeDriver { async fn get_trigger( &self, application: &Application, - request: ObjectRef, + request: MetadataObjectRef, ) -> Result { native_mysql::get_trigger( application, &request.scope.datasource_id, &request.scope.database_name, &request.scope.schema_name, - &request.name, + &request.object_name, ) .await } @@ -939,33 +951,33 @@ impl NativeTransferDriver for MysqlNativeDriver { async fn import_file( &self, application: &Application, - request: ImportTransferRequest, + request: ImportFileRequest, ) -> Result { - crate::transfer::mysql_impl::import_file(application, request).await + crate::transfer::mysql_driver::import_file(application, request).await } async fn export_sql_file( &self, application: &Application, - request: SqlExportTransferRequest, + request: SqlFileExportRequest, ) -> Result { - crate::transfer::mysql_impl::export_sql_file(application, request).await + crate::transfer::mysql_driver::export_sql_file(application, request).await } - async fn export_other_file( + async fn export_table_file( &self, application: &Application, - request: OtherExportTransferRequest, + request: TableFileExportRequest, ) -> Result { - crate::transfer::mysql_impl::export_other_file(application, request).await + crate::transfer::mysql_driver::export_table_file(application, request).await } - async fn export_dml( + async fn export_query_result( &self, application: &Application, - request: DmlExportTransferRequest, - ) -> Result { - crate::transfer::mysql_impl::export_dml(application, request).await + request: QueryResultExportRequest, + ) -> Result { + crate::transfer::mysql_driver::export_query_result(application, request).await } } @@ -1040,54 +1052,91 @@ impl NativeSchemaDiffDriver for MysqlNativeDriver { #[cfg(test)] mod tests { use super::*; + use crate::native_driver_types::NamespaceSqlOperation; + + #[test] + fn native_spi_sources_do_not_depend_on_the_compatibility_wire_namespace() { + let compatibility_prefix = ["Commu", "nity"].concat(); + let compatibility_identifier = ["commu", "nity_"].concat(); + for (path, source) in [ + ("native_driver.rs", include_str!("native_driver.rs")), + ( + "native_driver_types.rs", + include_str!("native_driver_types.rs"), + ), + ("native_mysql.rs", include_str!("native_mysql.rs")), + ( + "native_administration_types.rs", + include_str!("native_administration_types.rs"), + ), + ( + "native_schema_diff_types.rs", + include_str!("native_schema_diff_types.rs"), + ), + ( + "transfer/class_generation.rs", + include_str!("transfer/class_generation.rs"), + ), + ( + "transfer/mysql_driver.rs", + include_str!("transfer/mysql_driver.rs"), + ), + ] { + assert!( + !source.contains(&compatibility_prefix) + && !source.contains(&compatibility_identifier), + "{path} must stay independent from the compatibility wire namespace" + ); + } + } struct FakePostgresDriver; - impl NativeDriver for FakePostgresDriver { - fn id(&self) -> &'static str { - "postgresql" - } + struct DescriptorOnlyDriver(&'static NativeDriverDescriptor); - fn implementation(&self) -> &'static str { - "fake_postgres" + const FAKE_POSTGRES_DESCRIPTOR: NativeDriverDescriptor = NativeDriverDescriptor { + id: "postgresql", + implementation: "fake_postgres", + database_types: &["POSTGRESQL", "POSTGRES"], + compatibility_aliases: &["postgresql", "org.postgresql.Driver"], + }; + + impl NativeDriver for FakePostgresDriver { + fn descriptor(&self) -> &'static NativeDriverDescriptor { + &FAKE_POSTGRES_DESCRIPTOR } - fn database_types(&self) -> &'static [&'static str] { - &["POSTGRESQL", "POSTGRES"] + fn connection(&self) -> &dyn NativeConnectionDriver { + self } - fn descriptor(&self) -> JdbcDriver { - JdbcDriver { - pack_id: "native:fake_postgres".to_owned(), - name: "PostgreSQL (test native Rust)".to_owned(), - version: "test".to_owned(), - driver_id: self.id().to_owned(), - driver_class: "rust:fake_postgres".to_owned(), - artifact_count: 0, - artifact_bytes: "0".to_owned(), - } + fn dialect(&self) -> Option<&dyn NativeDialectDriver> { + Some(self) } + } - fn matches_driver(&self, driver_id: &str, descriptor: Option<&JdbcDriver>) -> bool { - driver_id.eq_ignore_ascii_case(self.id()) - || descriptor.is_some_and(|driver| { - driver - .driver_class - .eq_ignore_ascii_case("org.postgresql.Driver") - }) + impl NativeDriver for DescriptorOnlyDriver { + fn descriptor(&self) -> &'static NativeDriverDescriptor { + self.0 } fn connection(&self) -> &dyn NativeConnectionDriver { self } + } - fn dialect(&self) -> Option<&dyn NativeDialectDriver> { - Some(self) + #[async_trait] + impl NativeConnectionDriver for FakePostgresDriver { + async fn test_connection( + &self, + _connection: &DatasourceConnection, + ) -> Result<(), AppError> { + Ok(()) } } #[async_trait] - impl NativeConnectionDriver for FakePostgresDriver { + impl NativeConnectionDriver for DescriptorOnlyDriver { async fn test_connection( &self, _connection: &DatasourceConnection, @@ -1128,83 +1177,62 @@ mod tests { registry .driver_for_database_type("postgres") .expect("database type resolves") - .id(), + .descriptor() + .id, "postgresql" ); assert_eq!( registry - .driver_for_driver_id("POSTGRESQL", &[]) + .driver_for_datasource_driver_id("POSTGRESQL") .expect("driver id resolves") - .implementation(), + .descriptor() + .implementation, "fake_postgres" ); + assert!( + registry + .driver_for_datasource_driver_id("org.postgresql.Driver") + .is_some(), + "persisted compatibility aliases must resolve to their native implementation" + ); } - #[tokio::test] - async fn application_dispatches_a_postgres_capability_through_the_registry() { + #[test] + fn application_dispatches_a_postgres_capability_through_the_registry() { let registry = NativeDriverRegistry::try_new(vec![Arc::new(FakePostgresDriver)]) .expect("registry is valid"); let application = Application::with_native_drivers_for_test(registry); - let built = application - .build_community_namespace_sql(chat2db_contract::BuildCommunityNamespaceSqlRequest { - database_type: "POSTGRESQL".to_owned(), - operation: chat2db_contract::CommunityNamespaceSqlOperation::UseDatabase { + let driver = application + .native_driver_for_database_type("POSTGRESQL") + .expect("application must resolve the fake PostgreSQL driver"); + let dialect = driver + .dialect() + .expect("fake PostgreSQL must expose its dialect capability"); + let built = dialect + .build_namespace_sql(NamespaceSqlRequest { + operation: NamespaceSqlOperation::UseDatabase { database_name: "inventory".to_owned(), }, }) - .await .expect("application must dispatch to the fake PostgreSQL capability"); assert_eq!(built.sql, "fake-postgres:namespace"); } - #[test] - fn registry_uses_managed_descriptor_aliases_without_owning_driver_jars() { - let registry = NativeDriverRegistry::try_new(vec![Arc::new(FakePostgresDriver)]) - .expect("registry is valid"); - let managed = vec![JdbcDriver { - pack_id: "postgresql-42".to_owned(), - name: "PostgreSQL JDBC".to_owned(), - version: "42".to_owned(), - driver_id: "managed-pg".to_owned(), - driver_class: "org.postgresql.Driver".to_owned(), - artifact_count: 1, - artifact_bytes: "1".to_owned(), - }]; - - assert_eq!( - registry - .driver_for_driver_id("managed-pg", &managed) - .expect("managed descriptor resolves") - .id(), - "postgresql" - ); - } - #[test] fn registry_rejects_duplicate_database_type_ownership() { struct DuplicatePostgresDriver; impl NativeDriver for DuplicatePostgresDriver { - fn id(&self) -> &'static str { - "duplicate-postgresql" - } - - fn implementation(&self) -> &'static str { - "duplicate" - } - - fn database_types(&self) -> &'static [&'static str] { - &["postgresql"] - } - - fn descriptor(&self) -> JdbcDriver { - FakePostgresDriver.descriptor() - } - - fn matches_driver(&self, _driver_id: &str, _descriptor: Option<&JdbcDriver>) -> bool { - false + fn descriptor(&self) -> &'static NativeDriverDescriptor { + static DESCRIPTOR: NativeDriverDescriptor = NativeDriverDescriptor { + id: "duplicate-postgresql", + implementation: "duplicate", + database_types: &["postgresql"], + compatibility_aliases: &[], + }; + &DESCRIPTOR } fn connection(&self) -> &dyn NativeConnectionDriver { @@ -1228,4 +1256,52 @@ mod tests { ]); assert!(result.is_err()); } + + #[test] + fn registry_rejects_duplicate_driver_ids() { + let result = NativeDriverRegistry::try_new(vec![ + Arc::new(FakePostgresDriver), + Arc::new(FakePostgresDriver), + ]); + assert!(result.is_err()); + } + + #[test] + fn registry_rejects_non_canonical_descriptor_text() { + static DESCRIPTORS: [NativeDriverDescriptor; 4] = [ + NativeDriverDescriptor { + id: " postgresql", + implementation: "fake_postgres", + database_types: &["POSTGRESQL"], + compatibility_aliases: &[], + }, + NativeDriverDescriptor { + id: "postgresql", + implementation: "fake_postgres ", + database_types: &["POSTGRESQL"], + compatibility_aliases: &[], + }, + NativeDriverDescriptor { + id: "postgresql", + implementation: "fake_postgres", + database_types: &[" POSTGRESQL"], + compatibility_aliases: &[], + }, + NativeDriverDescriptor { + id: "postgresql", + implementation: "fake_postgres", + database_types: &["POSTGRESQL"], + compatibility_aliases: &["org.postgresql.Driver "], + }, + ]; + + for descriptor in &DESCRIPTORS { + let result = + NativeDriverRegistry::try_new(vec![Arc::new(DescriptorOnlyDriver(descriptor))]); + assert!( + result.is_err(), + "descriptor must be canonical: {descriptor:?}" + ); + } + } } diff --git a/crates/chat2db-core/src/native_driver_types.rs b/crates/chat2db-core/src/native_driver_types.rs index 1cf977a..09edda9 100644 --- a/crates/chat2db-core/src/native_driver_types.rs +++ b/crates/chat2db-core/src/native_driver_types.rs @@ -1,46 +1,343 @@ -use chat2db_contract::{ - CommunityDatabaseList, CommunityErTable, CommunityForeignKeyList, CommunityFunction, - CommunityFunctionList, CommunityFunctionParameterList, CommunityPrimaryKeyList, - CommunityProcedure, CommunityProcedureList, CommunityProcedureParameterList, - CommunityRoutineInvocationPreview, CommunityRoutineMigrationExecution, - CommunityRoutineMigrationRequest, CommunitySchemaList, CommunityTable, - CommunityTableColumnList, CommunityTableIndexList, CommunityTableList, - CommunityTablePreviewAccepted, CommunityTrigger, CommunityTriggerList, CommunityViewList, - DmlExportRequest, GetCommunityFunctionRequest, GetCommunityProcedureRequest, - GetCommunityTriggerRequest, ImportFileRequest, ListCommunityColumnsRequest, - ListCommunityDatabasesRequest, ListCommunityFunctionsRequest, ListCommunityIndexesRequest, - ListCommunityProceduresRequest, ListCommunitySchemasRequest, ListCommunityTableKeysRequest, - ListCommunityTablesRequest, ListCommunityTriggersRequest, ListCommunityViewsRequest, - OtherFileExportRequest, PreviewCommunityRoutineInvocationRequest, SqlFileExportRequest, - StartCommunityTablePreviewRequest, TransferArtifact, -}; - -pub(crate) type DatabaseList = CommunityDatabaseList; -pub(crate) type SchemaList = CommunitySchemaList; -pub(crate) type TableMetadata = CommunityTable; -pub(crate) type TableList = CommunityTableList; -pub(crate) type ColumnList = CommunityTableColumnList; -pub(crate) type IndexList = CommunityTableIndexList; -pub(crate) type ViewList = CommunityViewList; -pub(crate) type ForeignKeyList = CommunityForeignKeyList; -pub(crate) type PrimaryKeyList = CommunityPrimaryKeyList; -pub(crate) type FunctionMetadata = CommunityFunction; -pub(crate) type FunctionList = CommunityFunctionList; -pub(crate) type FunctionParameterList = CommunityFunctionParameterList; -pub(crate) type ProcedureMetadata = CommunityProcedure; -pub(crate) type ProcedureList = CommunityProcedureList; -pub(crate) type ProcedureParameterList = CommunityProcedureParameterList; -pub(crate) type TriggerMetadata = CommunityTrigger; -pub(crate) type TriggerList = CommunityTriggerList; -pub(crate) type EntityRelationTable = CommunityErTable; -pub(crate) type TablePreviewAccepted = CommunityTablePreviewAccepted; -pub(crate) type RoutineInvocationPreview = CommunityRoutineInvocationPreview; -pub(crate) type RoutineMigrationExecution = CommunityRoutineMigrationExecution; -pub(crate) type ImportTransferRequest = ImportFileRequest; -pub(crate) type SqlExportTransferRequest = SqlFileExportRequest; -pub(crate) type OtherExportTransferRequest = OtherFileExportRequest; -pub(crate) type DmlExportTransferRequest = DmlExportRequest; -pub(crate) type ExportArtifact = TransferArtifact; +/// Stable identity and runtime-selection metadata for one native Rust driver. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct NativeDriverDescriptor { + /// Stable identifier owned by the native driver implementation. + pub(crate) id: &'static str, + /// Rust crate or library that implements the database protocol. + pub(crate) implementation: &'static str, + /// Product database types routed to this driver. + pub(crate) database_types: &'static [&'static str], + /// Historical driver names, package identifiers, or classes accepted at the compatibility boundary. + pub(crate) compatibility_aliases: &'static [&'static str], +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct SchemaMetadata { + pub(crate) database_name: String, + pub(crate) name: String, + pub(crate) comment: String, + pub(crate) owner: String, + pub(crate) system: bool, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct SchemaList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct DatabaseMetadata { + pub(crate) name: String, + pub(crate) comment: String, + pub(crate) charset: String, + pub(crate) collation: String, + pub(crate) owner: String, + pub(crate) system: bool, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct DatabaseList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct TableMetadata { + pub(crate) database_name: String, + pub(crate) schema_name: String, + pub(crate) name: String, + pub(crate) table_type: String, + pub(crate) comment: String, + pub(crate) database_type: String, + pub(crate) pinned: bool, + pub(crate) ddl: String, + pub(crate) engine: String, + pub(crate) charset: String, + pub(crate) collation: String, + pub(crate) increment_value: Option, + pub(crate) partition: String, + pub(crate) tablespace: String, + pub(crate) rows: Option, + pub(crate) data_length: Option, + pub(crate) create_time: String, + pub(crate) update_time: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct TableList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ViewList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ColumnMetadata { + pub(crate) database_name: String, + pub(crate) schema_name: String, + pub(crate) table_name: String, + pub(crate) name: String, + pub(crate) column_type: String, + pub(crate) data_type: Option, + pub(crate) default_value: Option, + pub(crate) auto_increment: Option, + pub(crate) comment: String, + pub(crate) primary_key: Option, + pub(crate) primary_key_name: String, + pub(crate) primary_key_order: i32, + pub(crate) column_size: Option, + pub(crate) buffer_length: Option, + pub(crate) decimal_digits: Option, + pub(crate) num_prec_radix: Option, + pub(crate) sql_data_type: Option, + pub(crate) sql_datetime_sub: Option, + pub(crate) char_octet_length: Option, + pub(crate) ordinal_position: Option, + pub(crate) nullable: Option, + pub(crate) generated_column: Option, + pub(crate) extent: String, + pub(crate) charset: String, + pub(crate) collation: String, + pub(crate) unit: String, + pub(crate) sparse: Option, + pub(crate) default_constraint_name: String, + pub(crate) seed: Option, + pub(crate) increment: Option, + pub(crate) on_update_current_timestamp: Option, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ColumnList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct IndexColumnMetadata { + pub(crate) database_name: String, + pub(crate) schema_name: String, + pub(crate) table_name: String, + pub(crate) index_name: String, + pub(crate) column_name: String, + pub(crate) column_type: String, + pub(crate) comment: String, + pub(crate) ordinal_position: Option, + pub(crate) collation: String, + pub(crate) non_unique: Option, + pub(crate) index_qualifier: String, + pub(crate) sort_order: String, + pub(crate) cardinality: Option, + pub(crate) pages: Option, + pub(crate) filter_condition: String, + pub(crate) sub_part: Option, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct IndexMetadata { + pub(crate) database_name: String, + pub(crate) schema_name: String, + pub(crate) table_name: String, + pub(crate) name: String, + pub(crate) index_type: String, + pub(crate) unique: Option, + pub(crate) comment: String, + pub(crate) columns: Vec, + pub(crate) concurrently: Option, + pub(crate) method: String, + pub(crate) foreign_schema_name: String, + pub(crate) foreign_table_name: String, + pub(crate) foreign_column_names: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct IndexList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ForeignKeyMetadata { + pub(crate) primary_table_database: String, + pub(crate) primary_table_schema: String, + pub(crate) primary_table_name: String, + pub(crate) primary_column_name: String, + pub(crate) foreign_table_database: String, + pub(crate) foreign_table_schema: String, + pub(crate) foreign_table_name: String, + pub(crate) foreign_column_name: String, + pub(crate) key_sequence: i32, + pub(crate) update_rule: i32, + pub(crate) delete_rule: i32, + pub(crate) foreign_key_name: String, + pub(crate) primary_key_name: String, + pub(crate) deferrability: i32, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ForeignKeyList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct PrimaryKeyMetadata { + pub(crate) database_name: String, + pub(crate) schema_name: String, + pub(crate) table_name: String, + pub(crate) column_name: String, + pub(crate) name: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct PrimaryKeyList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct FunctionMetadata { + pub(crate) database_name: String, + pub(crate) schema_name: String, + pub(crate) name: String, + pub(crate) remarks: String, + pub(crate) function_type: Option, + pub(crate) specific_name: String, + pub(crate) body: String, + pub(crate) template: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct FunctionList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct FunctionParameterMetadata { + pub(crate) function_database: String, + pub(crate) function_schema: String, + pub(crate) function_name: String, + pub(crate) column_name: String, + pub(crate) column_type: Option, + pub(crate) data_type: Option, + pub(crate) type_name: String, + pub(crate) precision: Option, + pub(crate) length: Option, + pub(crate) scale: Option, + pub(crate) radix: Option, + pub(crate) nullable: Option, + pub(crate) remarks: String, + pub(crate) char_octet_length: Option, + pub(crate) ordinal_position: Option, + pub(crate) is_nullable: String, + pub(crate) specific_name: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct FunctionParameterList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ProcedureMetadata { + pub(crate) database_name: String, + pub(crate) schema_name: String, + pub(crate) name: String, + pub(crate) remarks: String, + pub(crate) procedure_type: Option, + pub(crate) specific_name: String, + pub(crate) body: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ProcedureList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ProcedureParameterMetadata { + pub(crate) procedure_database: String, + pub(crate) procedure_schema: String, + pub(crate) procedure_name: String, + pub(crate) column_name: String, + pub(crate) column_type: Option, + pub(crate) data_type: Option, + pub(crate) type_name: String, + pub(crate) precision: Option, + pub(crate) length: Option, + pub(crate) scale: Option, + pub(crate) radix: Option, + pub(crate) nullable: Option, + pub(crate) remarks: String, + pub(crate) column_default: String, + pub(crate) sql_data_type: Option, + pub(crate) sql_datetime_sub: Option, + pub(crate) char_octet_length: Option, + pub(crate) ordinal_position: Option, + pub(crate) is_nullable: String, + pub(crate) specific_name: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct ProcedureParameterList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct TriggerMetadata { + pub(crate) database_name: String, + pub(crate) schema_name: String, + pub(crate) name: String, + pub(crate) event_manipulation: String, + pub(crate) body: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct TriggerList { + pub(crate) items: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct EntityRelationColumn { + pub(crate) name: String, + pub(crate) column_type: String, + pub(crate) primary_key: bool, + pub(crate) comment: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct EntityRelationForeignKey { + pub(crate) primary_table: String, + pub(crate) primary_column: String, + pub(crate) foreign_table: String, + pub(crate) foreign_column: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct EntityRelationTable { + pub(crate) name: String, + pub(crate) comment: String, + pub(crate) columns: Vec, + pub(crate) foreign_keys: Vec, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct TablePreviewAccepted { + pub(crate) operation_id: String, + pub(crate) sql: String, + pub(crate) row_limit: u32, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct RoutineInvocationPreview { + pub(crate) sql: String, +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub(crate) struct RoutineMigrationExecution { + pub(crate) success: bool, + pub(crate) message: String, + pub(crate) sql: String, + pub(crate) failure_stage: Option, + pub(crate) restore_attempted: bool, + pub(crate) restore_succeeded: bool, +} #[derive(Debug, Clone, PartialEq, Eq)] pub(crate) struct BuiltSql { @@ -190,9 +487,9 @@ pub(crate) struct TableRef { } #[derive(Debug, Clone, PartialEq, Eq)] -pub(crate) struct ObjectRef { +pub(crate) struct MetadataObjectRef { pub(crate) scope: MetadataScope, - pub(crate) name: String, + pub(crate) object_name: String, } #[derive(Debug, Clone, PartialEq, Eq)] @@ -263,238 +560,3 @@ pub(crate) struct RoutineMigrationRequest { pub(crate) routine_name: String, pub(crate) ddl: String, } - -impl From for ListDatabasesRequest { - fn from(request: ListCommunityDatabasesRequest) -> Self { - Self { - datasource_id: request.datasource_id, - } - } -} - -impl From for ListSchemasRequest { - fn from(request: ListCommunitySchemasRequest) -> Self { - Self { - datasource_id: request.datasource_id, - database_name: request.database_name, - } - } -} - -impl From for ListTablesRequest { - fn from(request: ListCommunityTablesRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - name_pattern: request.table_name_pattern, - } - } -} - -impl From for ListColumnsRequest { - fn from(request: ListCommunityColumnsRequest) -> Self { - Self { - table: TableRef { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - table_name: request.table_name, - }, - } - } -} - -impl From for ListIndexesRequest { - fn from(request: ListCommunityIndexesRequest) -> Self { - Self { - table: TableRef { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - table_name: request.table_name, - }, - } - } -} - -impl From for ListViewsRequest { - fn from(request: ListCommunityViewsRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - name_pattern: request.view_name_pattern, - } - } -} - -impl From for ObjectRef { - fn from(request: ListCommunityViewsRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - name: request.view_name_pattern, - } - } -} - -impl From for ListTableKeysRequest { - fn from(request: ListCommunityTableKeysRequest) -> Self { - Self { - table: TableRef { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - table_name: request.table_name, - }, - } - } -} - -impl From for ListRoutinesRequest { - fn from(request: ListCommunityFunctionsRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - } - } -} - -impl From for ObjectRef { - fn from(request: GetCommunityFunctionRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - name: request.function_name, - } - } -} - -impl From for ListRoutinesRequest { - fn from(request: ListCommunityProceduresRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - } - } -} - -impl From for ObjectRef { - fn from(request: GetCommunityProcedureRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - name: request.procedure_name, - } - } -} - -impl From for ListTriggersRequest { - fn from(request: ListCommunityTriggersRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - } - } -} - -impl From for ObjectRef { - fn from(request: GetCommunityTriggerRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - name: request.trigger_name, - } - } -} - -impl From for TablePreviewRequest { - fn from(request: StartCommunityTablePreviewRequest) -> Self { - Self { - table: TableRef { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - table_name: request.table_name, - }, - } - } -} - -impl From for RoutineInvocationRequest { - fn from(request: PreviewCommunityRoutineInvocationRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - routine_type: request.routine_type, - routine_name: request.routine_name, - } - } -} - -impl From for RoutineMigrationRequest { - fn from(request: CommunityRoutineMigrationRequest) -> Self { - Self { - scope: MetadataScope { - datasource_id: request.datasource_id, - database_name: request.database_name, - schema_name: request.schema_name, - }, - database_type: request.database_type, - routine_type: request.routine_type, - routine_name: request.routine_name, - ddl: request.ddl, - } - } -} - -impl From for CommunityRoutineMigrationRequest { - fn from(request: RoutineMigrationRequest) -> Self { - Self { - datasource_id: request.scope.datasource_id, - database_type: request.database_type, - database_name: request.scope.database_name, - schema_name: request.scope.schema_name, - routine_type: request.routine_type, - routine_name: request.routine_name, - ddl: request.ddl, - } - } -} diff --git a/crates/chat2db-core/src/native_mysql.rs b/crates/chat2db-core/src/native_mysql.rs index 8b651a9..f8a3338 100644 --- a/crates/chat2db-core/src/native_mysql.rs +++ b/crates/chat2db-core/src/native_mysql.rs @@ -1,17 +1,9 @@ use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64_STANDARD}; use chat2db_contract::{ - ApiError, CommunityDatabase, CommunityDatabaseList, CommunityErColumn, CommunityErForeignKey, - CommunityErTable, CommunityForeignKey, CommunityForeignKeyList, CommunityFunction, - CommunityFunctionList, CommunityFunctionParameter, CommunityFunctionParameterList, - CommunityPrimaryKey, CommunityPrimaryKeyList, CommunityProcedure, CommunityProcedureList, - CommunityProcedureParameter, CommunityProcedureParameterList, CommunitySchemaList, - CommunityTable, CommunityTableColumn, CommunityTableColumnList, CommunityTableIndex, - CommunityTableIndexColumn, CommunityTableIndexList, CommunityTableList, CommunityTrigger, - CommunityTriggerList, CommunityViewList, DatasourceConnection, JdbcValue, JdbcValueType, - QueryLimits, ResultColumn, ResultMetadata, ResultRow, StartQueryRequest, + ApiError, DatasourceConnection, JdbcValue, JdbcValueType, QueryLimits, ResultColumn, + ResultMetadata, ResultRow, StartQueryRequest, }; use chat2db_engine_protocol::wire; -use chat2db_java_bridge::{JdbcParameter, JdbcValue as BridgeJdbcValue, QueryOptions}; use chat2db_storage::Storage; use chrono::{DateTime, Datelike, NaiveDate, NaiveDateTime, Timelike, Utc}; use mysql_async::{ @@ -36,13 +28,19 @@ use crate::{ AppError, AppErrorKind, Application, datasource_session::{ResolvedDatasourceConnection, resolve_datasource_connection}, native_driver_types::{ + ColumnList, ColumnMetadata, DatabaseList, DatabaseMetadata, EntityRelationColumn, + EntityRelationForeignKey, EntityRelationTable, ForeignKeyList, ForeignKeyMetadata, + FunctionList, FunctionMetadata, FunctionParameterList, FunctionParameterMetadata, + IndexColumnMetadata, IndexList, IndexMetadata, PrimaryKeyList, PrimaryKeyMetadata, + ProcedureList, ProcedureMetadata, ProcedureParameterList, ProcedureParameterMetadata, RoutineInvocationPreview, RoutineInvocationRequest, RoutineMigrationExecution, - RoutineMigrationRequest, TablePreviewAccepted, TablePreviewRequest, + RoutineMigrationRequest, SchemaList, TableList, TableMetadata, TablePreviewAccepted, + TablePreviewRequest, TriggerList, TriggerMetadata, ViewList, }, operation::CancellationRequest, query::{ - DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, PreparedQuery, - QueryTaskError, RetainedWriter, + DatabaseValue, DatabaseWriteError, NativeConsoleRequest, NativeConsoleResult, + PreparedQuery, QueryExecutionOptions, QueryParameter, QueryTaskError, RetainedWriter, }, ssh::{SshTunnel, SshTunnelIdentity, mysql_target, rewrite_mysql_target}, }; @@ -307,7 +305,7 @@ pub(crate) async fn test_connection_with_local_port( pub(crate) async fn list_databases( application: &Application, datasource_id: &str, -) -> Result { +) -> Result { let resolved = resolve_native_connection(application, datasource_id).await?; let mut conn = open_resolved_connection(&resolved).await?; let result = metadata_query(conn.query::<(String, String, String), _>( @@ -315,15 +313,15 @@ pub(crate) async fn list_databases( FROM information_schema.SCHEMATA ORDER BY SCHEMA_NAME", )) .await - .map(|rows| CommunityDatabaseList { + .map(|rows| DatabaseList { items: rows .into_iter() - .map(|(name, charset, collation)| CommunityDatabase { + .map(|(name, charset, collation)| DatabaseMetadata { system: is_system_database(&name), name, charset, collation, - ..CommunityDatabase::default() + ..DatabaseMetadata::default() }) .collect(), }); @@ -333,10 +331,10 @@ pub(crate) async fn list_databases( pub(crate) async fn list_schemas( application: &Application, datasource_id: &str, -) -> Result { +) -> Result { let resolved = resolve_native_connection(application, datasource_id).await?; let conn = open_resolved_connection(&resolved).await?; - finish_connection(conn, Ok(CommunitySchemaList::default())).await + finish_connection(conn, Ok(SchemaList::default())).await } pub(crate) async fn list_tables( @@ -344,7 +342,7 @@ pub(crate) async fn list_tables( datasource_id: &str, database_name: &str, table_name_pattern: &str, -) -> Result { +) -> Result { if database_name.trim().is_empty() { return Err(AppError::invalid( "invalid_mysql_metadata_request", @@ -369,7 +367,7 @@ pub(crate) async fn list_tables( conn.exec::(query, (database_name.to_owned(), pattern.clone(), pattern)), ) .await - .map(|rows| CommunityTableList { + .map(|rows| TableList { items: rows .into_iter() .map( @@ -385,7 +383,7 @@ pub(crate) async fn list_tables( data_length, create_time, update_time, - )| CommunityTable { + )| TableMetadata { database_name, name, table_type: normalize_table_type(&table_type).to_owned(), @@ -401,7 +399,7 @@ pub(crate) async fn list_tables( data_length, create_time: create_time.unwrap_or_default(), update_time: update_time.unwrap_or_default(), - ..CommunityTable::default() + ..TableMetadata::default() }, ) .collect(), @@ -415,7 +413,7 @@ pub(crate) async fn list_columns( database_name: &str, schema_name: &str, table_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(table_name, "tableName")?; let resolved = resolve_native_connection(application, datasource_id).await?; @@ -437,10 +435,10 @@ pub(crate) async fn list_columns( conn.exec::(query, (database_name.to_owned(), table_name.to_owned())), ) .await - .map(|rows| CommunityTableColumnList { + .map(|rows| ColumnList { items: rows .into_iter() - .map(|row| community_column(database_name, schema_name, table_name, row)) + .map(|row| column_metadata(database_name, schema_name, table_name, row)) .collect(), }); finish_connection(conn, result).await @@ -451,7 +449,7 @@ pub(crate) async fn load_er_tables( datasource_id: &str, database_name: &str, schema_name: &str, -) -> Result, AppError> { +) -> Result, AppError> { validate_metadata_identifier(database_name, "databaseName")?; let resolved = resolve_native_connection(application, datasource_id).await?; let mut conn = open_resolved_connection(&resolved).await?; @@ -469,7 +467,7 @@ pub(crate) async fn load_er_tables( ) .await?; - let mut columns_by_table = HashMap::>::new(); + let mut columns_by_table = HashMap::>::new(); for row in column_rows { let ErColumnRow { table_name, @@ -487,7 +485,7 @@ pub(crate) async fn load_er_tables( collation, primary_key_order, } = row; - let column = community_column( + let column = column_metadata( database_name, schema_name, &table_name, @@ -510,7 +508,7 @@ pub(crate) async fn load_er_tables( columns_by_table .entry(table_name) .or_default() - .push(CommunityErColumn { + .push(EntityRelationColumn { name: column.name, column_type: column.column_type, primary_key: column.primary_key.unwrap_or(false), @@ -518,29 +516,29 @@ pub(crate) async fn load_er_tables( }); } - let mut foreign_keys_by_table = HashMap::>::new(); + let mut foreign_keys_by_table = HashMap::>::new(); for row in foreign_key_rows { let table_name = row.4.clone(); foreign_keys_by_table .entry(table_name) .or_default() - .push(CommunityErForeignKey { - pk_table_name: row.1, - pk_column_name: row.2, - fk_table_name: row.4, - fk_column_name: row.5, + .push(EntityRelationForeignKey { + primary_table: row.1, + primary_column: row.2, + foreign_table: row.4, + foreign_column: row.5, }); } Ok(table_rows .into_iter() .map(|row| { - let table = community_table(row, schema_name); + let table = table_metadata(row, schema_name); let name = table.name; - CommunityErTable { + EntityRelationTable { comment: table.comment, - column_list: columns_by_table.remove(&name).unwrap_or_default(), - foreign_key_list: foreign_keys_by_table.remove(&name).unwrap_or_default(), + columns: columns_by_table.remove(&name).unwrap_or_default(), + foreign_keys: foreign_keys_by_table.remove(&name).unwrap_or_default(), name, } }) @@ -606,7 +604,7 @@ pub(crate) async fn list_indexes( database_name: &str, schema_name: &str, table_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(table_name, "tableName")?; let resolved = resolve_native_connection(application, datasource_id).await?; @@ -621,8 +619,8 @@ pub(crate) async fn list_indexes( conn.exec::(query, (database_name.to_owned(), table_name.to_owned())), ) .await - .map(|rows| CommunityTableIndexList { - items: community_indexes(rows, schema_name), + .map(|rows| IndexList { + items: index_metadata(rows, schema_name), }); finish_connection(conn, result).await } @@ -633,7 +631,7 @@ pub(crate) async fn list_views( database_name: &str, schema_name: &str, view_name_pattern: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; let resolved = resolve_native_connection(application, datasource_id).await?; let mut conn = open_resolved_connection(&resolved).await?; @@ -651,10 +649,10 @@ pub(crate) async fn list_views( conn.exec::(query, (database_name.to_owned(), pattern.clone(), pattern)), ) .await - .map(|rows| CommunityViewList { + .map(|rows| ViewList { items: rows .into_iter() - .map(|row| community_table(row, schema_name)) + .map(|row| table_metadata(row, schema_name)) .collect(), }); finish_connection(conn, result).await @@ -666,7 +664,7 @@ pub(crate) async fn get_view( database_name: &str, schema_name: &str, view_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(view_name, "viewName")?; let qualified_name = @@ -679,14 +677,14 @@ pub(crate) async fn get_view( .and_then(|row| { let row = row.ok_or_else(|| metadata_not_found("view", database_name, view_name))?; - Ok(CommunityTable { + Ok(TableMetadata { database_name: database_name.to_owned(), schema_name: schema_name.to_owned(), name: view_name.to_owned(), table_type: "VIEW".to_owned(), database_type: "MYSQL".to_owned(), ddl: row_string_at(&row, 1)?, - ..CommunityTable::default() + ..TableMetadata::default() }) }); finish_connection(conn, result).await @@ -721,7 +719,7 @@ pub(crate) async fn list_imported_keys( datasource_id: &str, database_name: &str, table_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(table_name, "tableName")?; let resolved = resolve_native_connection(application, datasource_id).await?; @@ -742,8 +740,8 @@ pub(crate) async fn list_imported_keys( conn.exec::(query, (database_name.to_owned(), table_name.to_owned())), ) .await - .map(|rows| CommunityForeignKeyList { - items: rows.into_iter().map(community_foreign_key).collect(), + .map(|rows| ForeignKeyList { + items: rows.into_iter().map(foreign_key_metadata).collect(), }); finish_connection(conn, result).await } @@ -753,7 +751,7 @@ pub(crate) async fn list_exported_keys( datasource_id: &str, database_name: &str, table_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(table_name, "tableName")?; let resolved = resolve_native_connection(application, datasource_id).await?; @@ -774,8 +772,8 @@ pub(crate) async fn list_exported_keys( conn.exec::(query, (database_name.to_owned(), table_name.to_owned())), ) .await - .map(|rows| CommunityForeignKeyList { - items: rows.into_iter().map(community_foreign_key).collect(), + .map(|rows| ForeignKeyList { + items: rows.into_iter().map(foreign_key_metadata).collect(), }); finish_connection(conn, result).await } @@ -786,7 +784,7 @@ pub(crate) async fn list_primary_keys( database_name: &str, schema_name: &str, table_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(table_name, "tableName")?; let resolved = resolve_native_connection(application, datasource_id).await?; @@ -800,11 +798,11 @@ pub(crate) async fn list_primary_keys( (database_name.to_owned(), table_name.to_owned()), )) .await - .map(|rows| CommunityPrimaryKeyList { + .map(|rows| PrimaryKeyList { items: rows .into_iter() .map( - |(database_name, table_name, column_name, name)| CommunityPrimaryKey { + |(database_name, table_name, column_name, name)| PrimaryKeyMetadata { database_name, schema_name: schema_name.to_owned(), table_name, @@ -822,7 +820,7 @@ pub(crate) async fn list_functions( datasource_id: &str, database_name: &str, schema_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; let resolved = resolve_native_connection(application, datasource_id).await?; let mut conn = open_resolved_connection(&resolved).await?; @@ -835,18 +833,18 @@ pub(crate) async fn list_functions( conn.exec::<(String, String, String, String), _, _>(query, (database_name.to_owned(),)), ) .await - .map(|rows| CommunityFunctionList { + .map(|rows| FunctionList { items: rows .into_iter() .map( - |(database_name, name, specific_name, remarks)| CommunityFunction { + |(database_name, name, specific_name, remarks)| FunctionMetadata { database_name, schema_name: schema_name.to_owned(), name, remarks, function_type: Some(1), specific_name, - ..CommunityFunction::default() + ..FunctionMetadata::default() }, ) .collect(), @@ -860,7 +858,7 @@ pub(crate) async fn get_function( database_name: &str, schema_name: &str, function_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(function_name, "functionName")?; let qualified_name = @@ -882,7 +880,7 @@ pub(crate) async fn get_function( ) .await? .ok_or_else(|| metadata_not_found("function", database_name, function_name))?; - Ok(CommunityFunction { + Ok(FunctionMetadata { database_name: metadata.0, schema_name: schema_name.to_owned(), name: metadata.1, @@ -903,7 +901,7 @@ pub(crate) async fn list_function_parameters( database_name: &str, schema_name: &str, function_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(function_name, "functionName")?; let resolved = resolve_native_connection(application, datasource_id).await?; @@ -919,10 +917,10 @@ pub(crate) async fn list_function_parameters( (database_name.to_owned(), function_name.to_owned()), )) .await - .map(|rows| CommunityFunctionParameterList { + .map(|rows| FunctionParameterList { items: rows .into_iter() - .map(|row| community_function_parameter(row, schema_name)) + .map(|row| function_parameter_metadata(row, schema_name)) .collect(), }); finish_connection(conn, result).await @@ -933,7 +931,7 @@ pub(crate) async fn list_procedures( datasource_id: &str, database_name: &str, schema_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; let resolved = resolve_native_connection(application, datasource_id).await?; let mut conn = open_resolved_connection(&resolved).await?; @@ -946,11 +944,11 @@ pub(crate) async fn list_procedures( conn.exec::<(String, String, String, String), _, _>(query, (database_name.to_owned(),)), ) .await - .map(|rows| CommunityProcedureList { + .map(|rows| ProcedureList { items: rows .into_iter() .map( - |(database_name, name, specific_name, remarks)| CommunityProcedure { + |(database_name, name, specific_name, remarks)| ProcedureMetadata { database_name, schema_name: schema_name.to_owned(), name, @@ -971,7 +969,7 @@ pub(crate) async fn get_procedure( database_name: &str, schema_name: &str, procedure_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(procedure_name, "procedureName")?; let qualified_name = qualified_identifier( @@ -997,7 +995,7 @@ pub(crate) async fn get_procedure( ) .await? .ok_or_else(|| metadata_not_found("procedure", database_name, procedure_name))?; - Ok(CommunityProcedure { + Ok(ProcedureMetadata { database_name: metadata.0, schema_name: schema_name.to_owned(), name: metadata.1, @@ -1017,7 +1015,7 @@ pub(crate) async fn list_procedure_parameters( database_name: &str, schema_name: &str, procedure_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(procedure_name, "procedureName")?; let resolved = resolve_native_connection(application, datasource_id).await?; @@ -1033,10 +1031,10 @@ pub(crate) async fn list_procedure_parameters( (database_name.to_owned(), procedure_name.to_owned()), )) .await - .map(|rows| CommunityProcedureParameterList { + .map(|rows| ProcedureParameterList { items: rows .into_iter() - .map(|row| community_procedure_parameter(row, schema_name)) + .map(|row| procedure_parameter_metadata(row, schema_name)) .collect(), }); finish_connection(conn, result).await @@ -1241,7 +1239,7 @@ fn routine_migration_plan( let ddl = request.ddl.trim(); if ddl.is_empty() || ddl.len() > MAX_SQL_BYTES || ddl.contains('\0') { return Err(AppError::invalid( - "invalid_community_routine_migration_request", + "invalid_routine_migration_request", "ddl is invalid", )); } @@ -1293,7 +1291,7 @@ pub(crate) async fn list_triggers( datasource_id: &str, database_name: &str, schema_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; let resolved = resolve_native_connection(application, datasource_id).await?; let mut conn = open_resolved_connection(&resolved).await?; @@ -1304,11 +1302,11 @@ pub(crate) async fn list_triggers( conn.exec::<(String, String, String), _, _>(query, (database_name.to_owned(),)), ) .await - .map(|rows| CommunityTriggerList { + .map(|rows| TriggerList { items: rows .into_iter() .map( - |(database_name, name, event_manipulation)| CommunityTrigger { + |(database_name, name, event_manipulation)| TriggerMetadata { database_name, schema_name: schema_name.to_owned(), name, @@ -1327,7 +1325,7 @@ pub(crate) async fn get_trigger( database_name: &str, schema_name: &str, trigger_name: &str, -) -> Result { +) -> Result { validate_metadata_identifier(database_name, "databaseName")?; validate_metadata_identifier(trigger_name, "triggerName")?; let qualified_name = @@ -1348,7 +1346,7 @@ pub(crate) async fn get_trigger( ) .await? .ok_or_else(|| metadata_not_found("trigger", database_name, trigger_name))?; - Ok(CommunityTrigger { + Ok(TriggerMetadata { database_name: metadata.0, schema_name: schema_name.to_owned(), name: metadata.1, @@ -1379,7 +1377,7 @@ pub(crate) fn validate_query(query: &PreparedQuery) -> Result<(), AppError> { validate_query_options(query.options) } -fn mysql_query_parameters(parameters: &[JdbcParameter]) -> Result { +fn mysql_query_parameters(parameters: &[QueryParameter]) -> Result { if parameters.is_empty() { return Ok(Params::Empty); } @@ -1406,32 +1404,28 @@ fn mysql_query_parameters(parameters: &[JdbcParameter]) -> Result Result { +fn mysql_query_value(value: &DatabaseValue) -> Result { match value { - BridgeJdbcValue::Null => Ok(Value::NULL), - BridgeJdbcValue::Boolean(value) => Ok(Value::Int(i64::from(*value))), - BridgeJdbcValue::SignedInteger(value) => Ok(Value::Int(*value)), - BridgeJdbcValue::UnsignedInteger(value) => Ok(Value::UInt(*value)), - BridgeJdbcValue::Float32(value) => Ok(Value::Float(*value)), - BridgeJdbcValue::Float64(value) => Ok(Value::Double(*value)), - BridgeJdbcValue::Decimal(value) => { + DatabaseValue::Null => Ok(Value::NULL), + DatabaseValue::Boolean(value) => Ok(Value::Int(i64::from(*value))), + DatabaseValue::SignedInteger(value) => Ok(Value::Int(*value)), + DatabaseValue::UnsignedInteger(value) => Ok(Value::UInt(*value)), + DatabaseValue::Float32(value) => Ok(Value::Float(*value)), + DatabaseValue::Float64(value) => Ok(Value::Double(*value)), + DatabaseValue::Decimal(value) => { validate_mysql_decimal(value)?; mysql_query_bytes(value.as_bytes(), "decimal") } - BridgeJdbcValue::Text(value) => mysql_query_bytes(value.as_bytes(), "text"), - BridgeJdbcValue::Binary(value) => mysql_query_bytes(value, "binary"), - BridgeJdbcValue::Date(value) => mysql_date_parameter(value), - BridgeJdbcValue::Time(value) => mysql_time_parameter(value), - BridgeJdbcValue::Timestamp(value) => mysql_timestamp_parameter(value), - BridgeJdbcValue::TimestampWithTimeZone(value) => { + DatabaseValue::Text(value) => mysql_query_bytes(value.as_bytes(), "text"), + DatabaseValue::Binary(value) => mysql_query_bytes(value, "binary"), + DatabaseValue::Date(value) => mysql_date_parameter(value), + DatabaseValue::Time(value) => mysql_time_parameter(value), + DatabaseValue::Timestamp(value) => mysql_timestamp_parameter(value), + DatabaseValue::TimestampWithTimeZone(value) => { mysql_timestamp_with_time_zone_parameter(value) } - BridgeJdbcValue::Json(value) => mysql_query_bytes(value.as_bytes(), "JSON"), - BridgeJdbcValue::Uuid(value) => mysql_query_bytes(value.as_bytes(), "UUID"), - BridgeJdbcValue::Opaque { .. } => Err(AppError::invalid( - "invalid_query_parameter", - "Opaque JDBC values cannot be MySQL query parameters", - )), + DatabaseValue::Json(value) => mysql_query_bytes(value.as_bytes(), "JSON"), + DatabaseValue::Uuid(value) => mysql_query_bytes(value.as_bytes(), "UUID"), } } @@ -3607,7 +3601,7 @@ fn mysql_display(value: Value) -> String { mysql_text(value).unwrap_or_else(|_| "[unavailable]".to_owned()) } -fn community_table( +fn table_metadata( ( database_name, name, @@ -3622,8 +3616,8 @@ fn community_table( update_time, ): TableRow, schema_name: &str, -) -> CommunityTable { - CommunityTable { +) -> TableMetadata { + TableMetadata { database_name, schema_name: schema_name.to_owned(), name, @@ -3640,11 +3634,11 @@ fn community_table( data_length, create_time: create_time.unwrap_or_default(), update_time: update_time.unwrap_or_default(), - ..CommunityTable::default() + ..TableMetadata::default() } } -fn community_column( +fn column_metadata( database_name: &str, schema_name: &str, table_name: &str, @@ -3663,11 +3657,11 @@ fn community_column( collation, primary_key_order, }: ColumnRow, -) -> CommunityTableColumn { +) -> ColumnMetadata { let data_type = data_type.to_ascii_uppercase(); let (column_size, decimal_digits) = mysql_column_size(&data_type, &column_definition, numeric_scale); - CommunityTableColumn { + ColumnMetadata { database_name: database_name.to_owned(), schema_name: schema_name.to_owned(), table_name: table_name.to_owned(), @@ -3686,7 +3680,7 @@ fn community_column( charset: charset.unwrap_or_default(), collation: collation.unwrap_or_default(), on_update_current_timestamp: Some(extra.contains("on update CURRENT_TIMESTAMP")), - ..CommunityTableColumn::default() + ..ColumnMetadata::default() } } @@ -3792,8 +3786,8 @@ fn mysql_column_size( (column_size, decimal_digits) } -fn community_indexes(rows: Vec, schema_name: &str) -> Vec { - let mut indexes: Vec = Vec::new(); +fn index_metadata(rows: Vec, schema_name: &str) -> Vec { + let mut indexes: Vec = Vec::new(); for ( database_name, table_name, @@ -3810,7 +3804,7 @@ fn community_indexes(rows: Vec, schema_name: &str) -> Vec, schema_name: &str) -> Vec, schema_name: &str) -> Vec) -> String { } } -fn community_foreign_key( +fn foreign_key_metadata( ( primary_table_database, primary_table_name, @@ -3882,8 +3876,8 @@ fn community_foreign_key( foreign_key_name, primary_key_name, ): ForeignKeyRow, -) -> CommunityForeignKey { - CommunityForeignKey { +) -> ForeignKeyMetadata { + ForeignKeyMetadata { primary_table_database, primary_table_name, primary_column_name, @@ -3896,7 +3890,7 @@ fn community_foreign_key( foreign_key_name, primary_key_name: primary_key_name.unwrap_or_default(), deferrability: 7, - ..CommunityForeignKey::default() + ..ForeignKeyMetadata::default() } } @@ -3915,7 +3909,7 @@ fn normalize_mysql_routine_type(routine_type: &str) -> Result Ok(MysqlRoutineType::Function), "PROCEDURE" => Ok(MysqlRoutineType::Procedure), _ => Err(AppError::invalid( - "invalid_community_routine_invocation_request", + "invalid_routine_invocation_request", "routineType must be FUNCTION or PROCEDURE", )), } @@ -4120,10 +4114,10 @@ fn mysql_routine_default_value(data_type: &str) -> &'static str { } } -fn community_function_parameter( +fn function_parameter_metadata( row: RoutineParameterRow, schema_name: &str, -) -> CommunityFunctionParameter { +) -> FunctionParameterMetadata { let RoutineParameterProjection { database_name, routine_name, @@ -4137,7 +4131,7 @@ fn community_function_parameter( radix, char_octet_length, } = routine_parameter_projection(row); - CommunityFunctionParameter { + FunctionParameterMetadata { function_database: database_name, function_schema: schema_name.to_owned(), function_name: routine_name.clone(), @@ -4157,14 +4151,14 @@ fn community_function_parameter( ordinal_position: Some(ordinal_position), is_nullable: "YES".to_owned(), specific_name: routine_name, - ..CommunityFunctionParameter::default() + ..FunctionParameterMetadata::default() } } -fn community_procedure_parameter( +fn procedure_parameter_metadata( row: RoutineParameterRow, schema_name: &str, -) -> CommunityProcedureParameter { +) -> ProcedureParameterMetadata { let RoutineParameterProjection { database_name, routine_name, @@ -4178,7 +4172,7 @@ fn community_procedure_parameter( radix, char_octet_length, } = routine_parameter_projection(row); - CommunityProcedureParameter { + ProcedureParameterMetadata { procedure_database: database_name, procedure_schema: schema_name.to_owned(), procedure_name: routine_name.clone(), @@ -4195,7 +4189,7 @@ fn community_procedure_parameter( ordinal_position: Some(ordinal_position), is_nullable: "YES".to_owned(), specific_name: routine_name, - ..CommunityProcedureParameter::default() + ..ProcedureParameterMetadata::default() } } @@ -4353,7 +4347,7 @@ fn metadata_not_found(kind: &str, database_name: &str, object_name: &str) -> App ) } -fn validate_query_options(options: QueryOptions) -> Result<(), AppError> { +fn validate_query_options(options: QueryExecutionOptions) -> Result<(), AppError> { if options.target_batch_rows > MAX_BATCH_ROWS { return Err(AppError::invalid( "invalid_query_limits", @@ -4380,7 +4374,7 @@ fn validate_query_options(options: QueryOptions) -> Result<(), AppError> { pub(crate) fn quote_identifier(value: &str, field: &str) -> Result { if value.trim().is_empty() || value.len() > MAX_IDENTIFIER_BYTES || value.contains('\0') { return Err(AppError::invalid( - "invalid_community_table_preview_request", + "invalid_table_preview_request", format!("{field} is invalid"), )); } @@ -4545,8 +4539,8 @@ pub(crate) async fn resolve_native_connection( let storage = application.require_storage()?; let resolved = resolve_datasource_connection(&storage, datasource_id).await?; if application - .native_driver_for_driver_id(&resolved.driver_id) - .is_none_or(|driver| driver.id() != "mysql") + .native_driver_for_datasource_driver_id(&resolved.driver_id) + .is_none_or(|driver| driver.descriptor().id != "mysql") { return Err(AppError::invalid( "mysql_driver_mismatch", @@ -4843,21 +4837,20 @@ mod tests { use super::{ ColumnRow, ConsoleExecutionError, ConsoleStatementExecution, MAX_CONSOLE_PAGE_SIZE, - MAX_CONSOLE_RESULT_BYTES, community_column, community_foreign_key, - community_function_parameter, community_indexes, community_procedure_parameter, - connection_opts, execute_console_statement, is_native_read_candidate, - mysql_column_reorder_hazard, mysql_identifier_is_backtick_quoted, + MAX_CONSOLE_RESULT_BYTES, column_metadata, connection_opts, execute_console_statement, + foreign_key_metadata, function_parameter_metadata, index_metadata, + is_native_read_candidate, mysql_column_reorder_hazard, mysql_identifier_is_backtick_quoted, mysql_metadata_column_type, mysql_routine_default_value, mysql_routine_invocation_name, mysql_routine_lookup_name, normalize_mysql_routine_type, normalize_table_type, - open_connection_with_opts, qualified_identifier, quote_identifier, - render_routine_invocation_preview, reserve_console_result_bytes, + open_connection_with_opts, procedure_parameter_metadata, qualified_identifier, + quote_identifier, render_routine_invocation_preview, reserve_console_result_bytes, routine_invocation_parameter, routine_migration_plan, split_mysql_script, validate_console_request, validate_forced_read_console, validate_read_only_console, validate_read_sql, validate_single_write_sql, }; use super::{MysqlRoutineType, RoutineInvocationParameter}; use crate::native_driver_types::{MetadataScope, RoutineMigrationRequest}; - use crate::{MysqlConsoleRequest, operation::CancellationRequest}; + use crate::{NativeConsoleRequest, operation::CancellationRequest}; #[test] fn mysql_console_splitter_respects_literals_identifiers_and_comments() { @@ -5220,7 +5213,7 @@ mod tests { .expect_err("unsupported routine types must fail closed"); assert_eq!( unsupported.api_error().code, - "invalid_community_routine_invocation_request" + "invalid_routine_invocation_request" ); assert!( routine_invocation_parameter( @@ -5267,7 +5260,7 @@ mod tests { } #[test] - fn mysql_routine_preview_matches_community_sql_rendering() { + fn mysql_routine_preview_matches_native_sql_rendering() { let parameters = vec![ RoutineInvocationParameter { name: "input-value".to_owned(), @@ -5309,7 +5302,7 @@ mod tests { } #[test] - fn mysql_routine_preview_quotes_identifiers_and_defaults_like_community() { + fn mysql_routine_preview_quotes_identifiers_and_uses_compatibility_defaults() { assert!(mysql_identifier_is_backtick_quoted("`odd``name`")); assert!(!mysql_identifier_is_backtick_quoted("``")); assert!(!mysql_identifier_is_backtick_quoted("`odd`name`")); @@ -5373,10 +5366,7 @@ mod tests { ddl: " ".to_owned(), }) .expect_err("empty ddl must fail"); - assert_eq!( - error.api_error().code, - "invalid_community_routine_migration_request" - ); + assert_eq!(error.api_error().code, "invalid_routine_migration_request"); } #[test] @@ -5548,7 +5538,7 @@ mod tests { #[test] fn console_page_size_all_uses_the_bounded_complete_window() { - let mut request = MysqlConsoleRequest { + let mut request = NativeConsoleRequest { datasource_id: "datasource-1".to_owned(), database_name: "inventory".to_owned(), sql: "SELECT 1".to_owned(), @@ -5572,8 +5562,8 @@ mod tests { } #[test] - fn mysql_column_metadata_preserves_community_projection() { - let column = community_column( + fn mysql_column_metadata_preserves_native_projection() { + let column = column_metadata( "inventory", "", "items", @@ -5622,7 +5612,7 @@ mod tests { "('read','write','close)later')", ), ] { - let column = community_column( + let column = column_metadata( "inventory", "", "items", @@ -5695,8 +5685,8 @@ mod tests { } #[test] - fn mysql_index_metadata_groups_columns_and_uses_community_types() { - let indexes = community_indexes( + fn mysql_index_metadata_groups_columns_and_uses_native_types() { + let indexes = index_metadata( vec![ ( "inventory".to_owned(), @@ -5756,7 +5746,7 @@ mod tests { #[test] fn mysql_relation_and_routine_metadata_match_jdbc_constants() { - let key = community_foreign_key(( + let key = foreign_key_metadata(( "inventory".to_owned(), "parent".to_owned(), "id".to_owned(), @@ -5773,7 +5763,7 @@ mod tests { assert_eq!(key.delete_rule, 2); assert_eq!(key.deferrability, 7); - let function_return = community_function_parameter( + let function_return = function_parameter_metadata( ( "inventory".to_owned(), "calculate_total".to_owned(), @@ -5795,7 +5785,7 @@ mod tests { assert_eq!(function_return.precision, Some(12)); assert_eq!(function_return.scale, Some(2)); - let procedure_output = community_procedure_parameter( + let procedure_output = procedure_parameter_metadata( ( "inventory".to_owned(), "load_total".to_owned(), @@ -5816,7 +5806,7 @@ mod tests { assert_eq!(procedure_output.data_type, Some(4)); assert_eq!(procedure_output.radix, Some(10)); - let long_text_input = community_procedure_parameter( + let long_text_input = procedure_parameter_metadata( ( "inventory".to_owned(), "store_text".to_owned(), diff --git a/crates/chat2db-core/src/query.rs b/crates/chat2db-core/src/query.rs index 42e07a0..78d8e35 100644 --- a/crates/chat2db-core/src/query.rs +++ b/crates/chat2db-core/src/query.rs @@ -6,8 +6,8 @@ use chat2db_contract::{ }; use chat2db_engine_protocol::wire; use chat2db_java_bridge::{ - BridgeError, CancelDisposition as BridgeCancelDisposition, JdbcParameter, QueryEvent, - QueryOptions, QueryRequest, QueryStream, + BridgeError, CancelDisposition as BridgeCancelDisposition, QueryEvent, QueryRequest, + QueryStream, }; use chat2db_storage::{ResultWriter, Storage, StorageError}; use tokio::sync::{oneshot, watch}; @@ -26,23 +26,57 @@ use crate::{ pub(crate) struct PreparedQuery { pub(crate) datasource_id: String, pub(crate) sql: String, - pub(crate) parameters: Vec, - pub(crate) options: QueryOptions, + pub(crate) parameters: Vec, + pub(crate) options: QueryExecutionOptions, pub(crate) retention: Duration, pub(crate) force_read_only: bool, } -/// One native `MySQL` Console execution request. +#[derive(Clone, Debug, PartialEq)] +pub(crate) struct QueryParameter { + pub(crate) position: u32, + pub(crate) value: DatabaseValue, +} + +#[derive(Clone, Debug, PartialEq)] +pub(crate) enum DatabaseValue { + Null, + Boolean(bool), + SignedInteger(i64), + UnsignedInteger(u64), + Float32(f32), + Float64(f64), + Decimal(String), + Text(String), + Binary(Vec), + Date(String), + Time(String), + Timestamp(String), + TimestampWithTimeZone(String), + Json(String), + Uuid(String), +} + +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub(crate) struct QueryExecutionOptions { + pub(crate) max_rows: u64, + pub(crate) target_batch_rows: u32, + pub(crate) target_batch_bytes: u32, + pub(crate) initial_batch_credits: u32, + pub(crate) max_result_bytes: u64, +} + +/// One native-driver Console execution request. #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] #[allow(clippy::struct_excessive_bools)] #[serde(rename_all = "camelCase")] pub struct NativeConsoleRequest { /// Opaque datasource id resolved by Core. pub datasource_id: String, - /// Optional `MySQL` database selected on the Console connection. + /// Optional database selected on the Console connection. #[serde(default)] pub database_name: String, - /// One statement or a semicolon-delimited `MySQL` script. + /// One statement or a semicolon-delimited script. pub sql: String, /// One-based result page number. pub page_no: u32, @@ -65,7 +99,7 @@ pub struct NativeConsoleRequest { pub error_continue: bool, } -/// One statement result emitted by native `MySQL` Console execution. +/// One statement result emitted by native-driver Console execution. #[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)] #[serde(rename_all = "camelCase")] pub struct NativeConsoleResult { @@ -74,7 +108,7 @@ pub struct NativeConsoleResult { /// One-based tabular result-set position within the statement. #[serde(skip_serializing_if = "Option::is_none")] pub result_set_id: Option, - /// The exact statement sent to `MySQL`. + /// The exact statement sent to the database. pub sql: String, /// Whether this individual statement result succeeded. pub success: bool, @@ -97,7 +131,7 @@ pub struct NativeConsoleResult { pub error: Option, } -/// Cloneable cancellation source for one native `MySQL` Console execution. +/// Cloneable cancellation source for one native-driver Console execution. #[derive(Debug, Clone)] pub struct NativeConsoleCancellation { sender: watch::Sender, @@ -144,10 +178,6 @@ const fn default_console_error_continue() -> bool { true } -pub type MysqlConsoleRequest = NativeConsoleRequest; -pub type MysqlConsoleResult = NativeConsoleResult; -pub type MysqlConsoleCancellation = NativeConsoleCancellation; - enum QueryBackend { Java { engine: EngineLease, @@ -186,21 +216,6 @@ pub(crate) struct RetainedWriter { } impl Application { - /// Executes an arbitrary `MySQL` Console script natively on one connection. - /// - /// # Errors - /// - /// Returns validation, datasource, connection, cancellation, or unrecoverable - /// protocol failures. Recoverable `MySQL` statement errors are represented in - /// the returned result list so `error_continue` can be honored. - pub async fn execute_mysql_console( - &self, - request: MysqlConsoleRequest, - cancellation: MysqlConsoleCancellation, - ) -> Result, AppError> { - self.execute_native_console(request, cancellation).await - } - /// Executes a native-driver Console request on the runtime-selected driver. /// /// # Errors @@ -233,7 +248,7 @@ impl Application { let storage = self.require_storage()?; let resolved = resolve_datasource_connection(&storage, &request.datasource_id).await?; let driver = self - .native_driver_for_driver_id(&resolved.driver_id) + .native_driver_for_datasource_driver_id(&resolved.driver_id) .ok_or_else(|| { AppError::invalid( "native_driver_not_available", @@ -333,7 +348,9 @@ impl Application { let _engine = self.require_engine().await?; } let resolved = resolve_datasource_connection(&storage, &prepared.datasource_id).await?; - let backend = if let Some(driver) = self.native_driver_for_driver_id(&resolved.driver_id) { + let backend = if let Some(driver) = + self.native_driver_for_datasource_driver_id(&resolved.driver_id) + { if let Some(query) = driver.query() { if query.is_read_candidate(&prepared.sql)? { query.validate_query(&prepared)?; @@ -482,6 +499,7 @@ impl Application { SessionReadOnly::Configured }; let session = open_datasource_session(&engine, resolved, read_only).await?; + let retention = query.retention; let cancellation_request = { cancellation.borrow().clone() }; if let CancellationRequest::Requested { reason } = cancellation_request { @@ -496,9 +514,13 @@ impl Application { let stream = match session .execute_query(QueryRequest { sql: query.sql, - parameters: query.parameters, + parameters: query + .parameters + .into_iter() + .map(convert::query_parameter_to_java) + .collect(), transaction_id: None, - options: query.options, + options: convert::query_options_to_java(query.options), }) .await { @@ -514,13 +536,7 @@ impl Application { }; let result = self - .consume_stream( - operation_id, - &mut cancellation, - stream, - storage, - query.retention, - ) + .consume_stream(operation_id, &mut cancellation, stream, storage, retention) .await; let close_result = session.close().await.map_err(AppError::from); preserve_primary_outcome( @@ -665,7 +681,7 @@ impl Application { ), ))); } - if let Some(driver) = self.native_driver_for_driver_id(&resolved.driver_id) + if let Some(driver) = self.native_driver_for_datasource_driver_id(&resolved.driver_id) && let Some(query) = driver.query() { return query.execute_update(resolved, sql, cancellation).await; @@ -774,7 +790,7 @@ impl TryFrom for PreparedQuery { .into_iter() .map(convert::query_parameter) .collect::>()?, - options: QueryOptions { + options: QueryExecutionOptions { max_rows: parse_u64(&max_rows, "maxRows")?, target_batch_rows: batch_rows, target_batch_bytes: batch_bytes, @@ -987,13 +1003,13 @@ mod tests { use chat2db_java_bridge::{BridgeError, RemoteEngineError}; use super::{ - AppError, MysqlConsoleCancellation, MysqlConsoleRequest, QueryTaskError, + AppError, NativeConsoleCancellation, NativeConsoleRequest, QueryTaskError, is_inactive_credit_grant_error, preserve_primary_outcome, validate_credit_grant, }; #[test] - fn mysql_console_cancellation_preserves_the_first_reason() { - let cancellation = MysqlConsoleCancellation::new(); + fn native_console_cancellation_preserves_the_first_reason() { + let cancellation = NativeConsoleCancellation::new(); let receiver = cancellation.subscribe(); assert!(cancellation.cancel(Some("first".to_owned()))); @@ -1007,8 +1023,8 @@ mod tests { } #[test] - fn mysql_console_json_defaults_to_continuing_after_statement_errors() { - let request: MysqlConsoleRequest = serde_json::from_value(serde_json::json!({ + fn native_console_json_defaults_to_continuing_after_statement_errors() { + let request: NativeConsoleRequest = serde_json::from_value(serde_json::json!({ "datasourceId": "mysql-1", "sql": "SELECT 1", "pageNo": 1, diff --git a/crates/chat2db-core/src/ssh.rs b/crates/chat2db-core/src/ssh.rs index 94e14e9..802c5a1 100644 --- a/crates/chat2db-core/src/ssh.rs +++ b/crates/chat2db-core/src/ssh.rs @@ -211,7 +211,7 @@ impl Application { }; self.require_managed_driver(&request.driver_id)?; let driver = self - .native_driver_for_driver_id(&request.driver_id) + .native_driver_for_datasource_driver_id(&request.driver_id) .ok_or_else(|| { AppError::invalid( "ssh_driver_not_supported", diff --git a/crates/chat2db-core/src/transfer/class_generation.rs b/crates/chat2db-core/src/transfer/class_generation.rs index ff6da5f..a57eda4 100644 --- a/crates/chat2db-core/src/transfer/class_generation.rs +++ b/crates/chat2db-core/src/transfer/class_generation.rs @@ -5,14 +5,14 @@ use std::{ path::{Path, PathBuf}, }; -use chat2db_contract::{CommunityTableColumn, GenerateMysqlClassRequest, GeneratedMysqlClassSet}; +use chat2db_contract::{GenerateMysqlClassRequest, GeneratedMysqlClassSet}; use chat2db_storage::TransferArtifactRecord; use uuid::Uuid; use zip::{CompressionMethod, ZipWriter, write::SimpleFileOptions}; use crate::{ AppError, Application, - native_driver_types::{ListColumnsRequest, MetadataScope, TableRef}, + native_driver_types::{ColumnMetadata, ListColumnsRequest, MetadataScope, TableRef}, now_millis, }; @@ -172,7 +172,7 @@ fn write_class_set( fn render_class_set( table_name: &str, - columns: &[CommunityTableColumn], + columns: &[ColumnMetadata], ) -> Result { let class_name = format!("{}DO", upper_camel(table_name)); let entity_name = format!("{class_name}.java"); @@ -232,7 +232,7 @@ fn safe_archive_component(value: &str) -> String { } } -fn render_entity(class_name: &str, table_name: &str, columns: &[CommunityTableColumn]) -> String { +fn render_entity(class_name: &str, table_name: &str, columns: &[ColumnMetadata]) -> String { let mut imports = BTreeSet::from([ "com.baomidou.mybatisplus.annotation.TableField", "com.baomidou.mybatisplus.annotation.TableName", diff --git a/crates/chat2db-core/src/transfer/mod.rs b/crates/chat2db-core/src/transfer/mod.rs index 46a33ff..ae4d38b 100644 --- a/crates/chat2db-core/src/transfer/mod.rs +++ b/crates/chat2db-core/src/transfer/mod.rs @@ -1,7 +1,7 @@ mod class_generation; mod format; mod mysql; -pub(crate) mod mysql_impl; +pub(crate) mod mysql_driver; use std::{ collections::HashMap, fmt::Write as _, fs::File, future::Future, path::PathBuf, pin::Pin, @@ -9,9 +9,10 @@ use std::{ }; use chat2db_contract::{ - DmlExportRequest, GenerateMysqlClassRequest, GeneratedMysqlClassSet, ImportFileRequest, - OtherFileExportRequest, SqlFileExportRequest, TransferArtifact, TransferTask, - TransferTaskAccepted, TransferTaskKind, TransferTaskPage, TransferTaskStatus, + DmlExportFormat, DmlExportRequest, DmlExportSize, GenerateMysqlClassRequest, + GeneratedMysqlClassSet, ImportFileRequest, OtherFileExportRequest, SqlFileExportRequest, + TransferArtifact, TransferFileFormat, TransferTask, TransferTaskAccepted, TransferTaskKind, + TransferTaskPage, TransferTaskStatus, }; use chat2db_storage::{ CreateTransferTask, ResolvedTransferArtifact, Storage, StorageError, StoredTransferTaskKind, @@ -71,6 +72,90 @@ pub struct TransferArtifactDownload { pub file: File, } +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct TableFileExportRequest { + pub(crate) datasource_id: String, + pub(crate) database_name: String, + pub(crate) schema_name: String, + pub(crate) table_names: Vec, + pub(crate) format: TransferFileFormat, + pub(crate) contains_header: bool, + pub(crate) export_path: Option, +} + +impl From for TableFileExportRequest { + fn from(request: OtherFileExportRequest) -> Self { + Self { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + table_names: request.table_names, + format: request.format, + contains_header: request.contains_header, + export_path: request.export_path, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum QueryResultExportScope { + CurrentPage, + All, +} + +impl From for QueryResultExportScope { + fn from(scope: DmlExportSize) -> Self { + match scope { + DmlExportSize::CurrentPage => Self::CurrentPage, + DmlExportSize::All => Self::All, + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum QueryResultExportFormat { + Csv, + Xlsx, + Insert, +} + +impl From for QueryResultExportFormat { + fn from(format: DmlExportFormat) -> Self { + match format { + DmlExportFormat::Csv => Self::Csv, + DmlExportFormat::Xlsx => Self::Xlsx, + DmlExportFormat::Insert => Self::Insert, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct QueryResultExportRequest { + pub(crate) datasource_id: String, + pub(crate) database_name: String, + pub(crate) schema_name: String, + pub(crate) sql: String, + pub(crate) original_sql: String, + pub(crate) result_set_id: Option, + pub(crate) scope: QueryResultExportScope, + pub(crate) format: QueryResultExportFormat, +} + +impl From for QueryResultExportRequest { + fn from(request: DmlExportRequest) -> Self { + Self { + datasource_id: request.datasource_id, + database_name: request.database_name, + schema_name: request.schema_name, + sql: request.sql, + original_sql: request.original_sql, + result_set_id: request.result_set_id, + scope: request.export_size.into(), + format: request.format.into(), + } + } +} + impl std::fmt::Debug for TransferArtifactDownload { fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { formatter @@ -401,9 +486,9 @@ impl Application { /// # Errors /// /// Returns validation, datasource, storage, or runtime-shutdown failures. - pub async fn export_other_file( + pub(crate) async fn export_table_file( &self, - request: OtherFileExportRequest, + request: TableFileExportRequest, ) -> Result { let driver = self .require_native_driver_for_datasource(&request.datasource_id) @@ -414,11 +499,23 @@ impl Application { "The native Rust driver does not implement import and export operations", ) })?; - let spec = transfer.export_other_file(self, request).await?; + let spec = transfer.export_table_file(self, request).await?; self.start_transfer_job(spec).await } - /// Retained `MySQL` compatibility name for [`Self::export_other_file`]. + /// Retained compatibility entry point for [`Self::export_table_file`]. + /// + /// # Errors + /// + /// Returns validation, datasource, storage, or runtime-shutdown failures. + pub async fn export_other_file( + &self, + request: OtherFileExportRequest, + ) -> Result { + self.export_table_file(request.into()).await + } + + /// Retained `MySQL` compatibility name for [`Self::export_table_file`]. /// /// # Errors /// @@ -560,14 +657,14 @@ impl Application { self.transfer_artifact_download(&artifact_id).await } - /// Streams one DML result into a temporary managed CSV, XLSX, or INSERT artifact. + /// Streams one query result into a temporary managed CSV, XLSX, or INSERT artifact. /// /// # Errors /// /// Returns SQL analysis, datasource, query, format, or storage failures. - pub async fn export_dml( + pub(crate) async fn export_query_result( &self, - request: DmlExportRequest, + request: QueryResultExportRequest, ) -> Result { let driver = self .require_native_driver_for_datasource(&request.datasource_id) @@ -578,10 +675,22 @@ impl Application { "The native Rust driver does not implement import and export operations", ) })?; - transfer.export_dml(self, request).await + transfer.export_query_result(self, request).await } - /// Retained `MySQL` compatibility name for [`Self::export_dml`]. + /// Retained compatibility entry point for [`Self::export_query_result`]. + /// + /// # Errors + /// + /// Returns SQL analysis, datasource, query, format, or storage failures. + pub async fn export_dml( + &self, + request: DmlExportRequest, + ) -> Result { + self.export_query_result(request.into()).await + } + + /// Retained `MySQL` compatibility name for [`Self::export_query_result`]. /// /// # Errors /// @@ -1082,15 +1191,20 @@ mod tests { use std::{future::Future, path::Path, time::Duration}; use base64::{Engine as _, engine::general_purpose::STANDARD}; + use chat2db_contract::{ + DmlExportFormat, DmlExportRequest, DmlExportSize, OtherFileExportRequest, + TransferFileFormat, + }; use chat2db_java_bridge::{EngineCommand, EngineConfig}; use chat2db_storage::{StoredTransferTaskKind, StoredTransferTaskStatus, TransferTaskRecord}; use tempfile::TempDir; use tokio::sync::oneshot; use super::{ - MAX_TRANSFER_FAILURE_MESSAGE_BYTES, TaskCompletion, TransferContext, TransferJobKind, - TransferJobSpec, TransferRunError, TransferTaskControl, TransferTerminalState, - finalize_transfer_task, transfer_task, + MAX_TRANSFER_FAILURE_MESSAGE_BYTES, QueryResultExportFormat, QueryResultExportRequest, + QueryResultExportScope, TableFileExportRequest, TaskCompletion, TransferContext, + TransferJobKind, TransferJobSpec, TransferRunError, TransferTaskControl, + TransferTerminalState, finalize_transfer_task, transfer_task, }; use crate::{AppError, Application, RuntimeConfig, RuntimeHost}; @@ -1104,6 +1218,74 @@ mod tests { assert_send(&spec); } + #[test] + fn compatibility_table_file_request_maps_every_field() { + let request = TableFileExportRequest::from(OtherFileExportRequest { + datasource_id: "datasource".to_owned(), + database_name: "database".to_owned(), + schema_name: "schema".to_owned(), + table_names: vec!["first".to_owned(), "second".to_owned()], + format: TransferFileFormat::Xlsx, + contains_header: false, + export_path: Some("/tmp/export".to_owned()), + }); + + assert_eq!(request.datasource_id, "datasource"); + assert_eq!(request.database_name, "database"); + assert_eq!(request.schema_name, "schema"); + assert_eq!(request.table_names, ["first", "second"]); + assert_eq!(request.format, TransferFileFormat::Xlsx); + assert!(!request.contains_header); + assert_eq!(request.export_path.as_deref(), Some("/tmp/export")); + } + + #[test] + fn compatibility_query_result_request_maps_every_field() { + let request = QueryResultExportRequest::from(DmlExportRequest { + datasource_id: "datasource".to_owned(), + database_name: "database".to_owned(), + schema_name: "schema".to_owned(), + sql: "SELECT * FROM table LIMIT 10".to_owned(), + original_sql: "SELECT * FROM table".to_owned(), + result_set_id: Some(2), + export_size: DmlExportSize::CurrentPage, + format: DmlExportFormat::Insert, + }); + + assert_eq!(request.datasource_id, "datasource"); + assert_eq!(request.database_name, "database"); + assert_eq!(request.schema_name, "schema"); + assert_eq!(request.sql, "SELECT * FROM table LIMIT 10"); + assert_eq!(request.original_sql, "SELECT * FROM table"); + assert_eq!(request.result_set_id, Some(2)); + assert_eq!(request.scope, QueryResultExportScope::CurrentPage); + assert_eq!(request.format, QueryResultExportFormat::Insert); + } + + #[test] + fn compatibility_query_result_enums_map_exhaustively() { + assert_eq!( + QueryResultExportScope::from(DmlExportSize::CurrentPage), + QueryResultExportScope::CurrentPage + ); + assert_eq!( + QueryResultExportScope::from(DmlExportSize::All), + QueryResultExportScope::All + ); + assert_eq!( + QueryResultExportFormat::from(DmlExportFormat::Csv), + QueryResultExportFormat::Csv + ); + assert_eq!( + QueryResultExportFormat::from(DmlExportFormat::Xlsx), + QueryResultExportFormat::Xlsx + ); + assert_eq!( + QueryResultExportFormat::from(DmlExportFormat::Insert), + QueryResultExportFormat::Insert + ); + } + #[tokio::test] async fn cancellation_wins_over_a_successful_runner_completion() { let directory = TempDir::new().expect("temporary transfer runtime"); diff --git a/crates/chat2db-core/src/transfer/mysql.rs b/crates/chat2db-core/src/transfer/mysql.rs index b3262bd..46fa72f 100644 --- a/crates/chat2db-core/src/transfer/mysql.rs +++ b/crates/chat2db-core/src/transfer/mysql.rs @@ -8,8 +8,8 @@ use std::{ }; use chat2db_contract::{ - DmlExportFormat, DmlExportRequest, DmlExportSize, ImportFileRequest, OtherFileExportRequest, - SqlFileExportRequest, TabularImportEncoding, TransferFileFormat, TransferSqlScope, + ImportFileRequest, SqlFileExportRequest, TabularImportEncoding, TransferFileFormat, + TransferSqlScope, }; use chat2db_storage::{TransferArtifactRecord, TransferArtifactWriter}; use mysql_async::{ @@ -27,12 +27,16 @@ use tokio_util::sync::CancellationToken; use uuid::Uuid; use zip::{CompressionMethod, ZipWriter, write::SimpleFileOptions}; -use super::{PendingTransferArtifact, TaskCompletion, TransferContext, TransferRunError, format}; +use super::{ + PendingTransferArtifact, QueryResultExportFormat, QueryResultExportRequest, + QueryResultExportScope, TableFileExportRequest, TaskCompletion, TransferContext, + TransferRunError, format, +}; use crate::{AppError, AppErrorKind, Application, native_mysql}; const IMPORT_BATCH_ROWS: usize = 256; const PROGRESS_ROW_INTERVAL: u64 = 250; -const DML_ARTIFACT_TTL_MS: i64 = 24 * 60 * 60 * 1_000; +const QUERY_RESULT_ARTIFACT_TTL_MS: i64 = 24 * 60 * 60 * 1_000; pub(super) async fn import_file( application: &Application, @@ -112,9 +116,9 @@ pub(super) async fn export_sql( publish_task_artifact(writer, request.export_path.as_deref(), &file_name, context).await } -pub(super) async fn export_other( +pub(super) async fn export_table_file( application: &Application, - request: OtherFileExportRequest, + request: TableFileExportRequest, context: &TransferContext, ) -> Result { validate_database_name(&request.database_name).map_err(TransferRunError::into_app_error)?; @@ -193,24 +197,26 @@ pub(super) async fn export_other( publish_task_artifact(writer, request.export_path.as_deref(), &file_name, context).await } -pub(super) async fn export_dml( +pub(super) async fn export_query_result( application: &Application, - request: DmlExportRequest, + request: QueryResultExportRequest, ) -> Result { validate_database_name(&request.database_name).map_err(TransferRunError::into_app_error)?; - let sql = match request.export_size { - DmlExportSize::CurrentPage if !request.sql.trim().is_empty() => request.sql.trim(), - DmlExportSize::CurrentPage | DmlExportSize::All => request.original_sql.trim(), + let sql = match request.scope { + QueryResultExportScope::CurrentPage if !request.sql.trim().is_empty() => request.sql.trim(), + QueryResultExportScope::CurrentPage | QueryResultExportScope::All => { + request.original_sql.trim() + } }; let table_name = select_table_name(sql)?; let (format, extension, media_type) = match request.format { - DmlExportFormat::Csv => (TransferFileFormat::Csv, "csv", "text/csv; charset=utf-8"), - DmlExportFormat::Xlsx => ( + QueryResultExportFormat::Csv => (TransferFileFormat::Csv, "csv", "text/csv; charset=utf-8"), + QueryResultExportFormat::Xlsx => ( TransferFileFormat::Xlsx, "xlsx", "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet", ), - DmlExportFormat::Insert => ( + QueryResultExportFormat::Insert => ( TransferFileFormat::Sql, "sql", "application/sql; charset=utf-8", @@ -221,7 +227,7 @@ pub(super) async fn export_dml( .unwrap_or(request.database_name.as_str()); let file_name = timestamped_file_name(stem, extension); let storage = application.require_storage()?; - let expires_at = now_millis()?.saturating_add(DML_ARTIFACT_TTL_MS); + let expires_at = now_millis()?.saturating_add(QUERY_RESULT_ARTIFACT_TTL_MS); let mut writer = storage .begin_transfer_artifact( None, @@ -238,7 +244,7 @@ pub(super) async fn export_dml( .map_err(TransferRunError::into_app_error)?; let selected_result_set = request.result_set_id.unwrap_or(0); let write_result = match request.format { - DmlExportFormat::Csv | DmlExportFormat::Xlsx => write_query_tabular( + QueryResultExportFormat::Csv | QueryResultExportFormat::Xlsx => write_query_tabular( &mut conn, writer.file_mut(), sql, @@ -250,7 +256,7 @@ pub(super) async fn export_dml( ) .await .map(|_| ()), - DmlExportFormat::Insert => { + QueryResultExportFormat::Insert => { let table_name = table_name.ok_or_else(|| { AppError::invalid( "sql_analysis_error", diff --git a/crates/chat2db-core/src/transfer/mysql_impl.rs b/crates/chat2db-core/src/transfer/mysql_driver.rs similarity index 84% rename from crates/chat2db-core/src/transfer/mysql_impl.rs rename to crates/chat2db-core/src/transfer/mysql_driver.rs index a251c45..b609084 100644 --- a/crates/chat2db-core/src/transfer/mysql_impl.rs +++ b/crates/chat2db-core/src/transfer/mysql_driver.rs @@ -1,13 +1,10 @@ use std::path::Path; -use chat2db_contract::{ - DmlExportRequest, ImportFileRequest, OtherFileExportRequest, SqlFileExportRequest, - TransferArtifact, -}; +use chat2db_contract::{ImportFileRequest, SqlFileExportRequest, TransferArtifact}; use super::{ - TransferJobKind, TransferJobSpec, mysql, single_table, transfer_artifact, - validate_import_request, validate_transfer_scope, + QueryResultExportRequest, TableFileExportRequest, TransferJobKind, TransferJobSpec, mysql, + single_table, transfer_artifact, validate_import_request, validate_transfer_scope, }; use crate::{AppError, Application, native_mysql}; @@ -55,9 +52,9 @@ pub(crate) async fn export_sql_file( )) } -pub(crate) async fn export_other_file( +pub(crate) async fn export_table_file( application: &Application, - request: OtherFileExportRequest, + request: TableFileExportRequest, ) -> Result { validate_transfer_scope(&request.datasource_id, request.export_path.as_deref())?; validate_mysql_database(&request.database_name)?; @@ -80,17 +77,17 @@ pub(crate) async fn export_other_file( request.table_names.len() ), move |application, context| async move { - mysql::export_other(&application, request, &context).await + mysql::export_table_file(&application, request, &context).await }, )) } -pub(crate) async fn export_dml( +pub(crate) async fn export_query_result( application: &Application, - request: DmlExportRequest, + request: QueryResultExportRequest, ) -> Result { native_mysql::resolve_native_connection(application, &request.datasource_id).await?; - mysql::export_dml(application, request) + mysql::export_query_result(application, request) .await .map(transfer_artifact) } diff --git a/crates/chat2db-core/tests/native_mysql_console_docker.rs b/crates/chat2db-core/tests/native_mysql_console_docker.rs index c86197a..9d8a54c 100644 --- a/crates/chat2db-core/tests/native_mysql_console_docker.rs +++ b/crates/chat2db-core/tests/native_mysql_console_docker.rs @@ -6,8 +6,8 @@ use chat2db_contract::{ JdbcValue, }; use chat2db_core::{ - Application, MysqlConsoleCancellation, MysqlConsoleRequest, MysqlConsoleResult, RuntimeConfig, - RuntimeHost, + Application, NativeConsoleCancellation, NativeConsoleRequest, NativeConsoleResult, + RuntimeConfig, RuntimeHost, }; use chat2db_java_bridge::{EngineCommand, EngineConfig}; use futures_util::FutureExt as _; @@ -539,14 +539,14 @@ async fn verify_read_only( assert!(inspected[0].success); let error = application - .execute_mysql_console( + .execute_native_console( request( &datasource.id, database_name, format!("INSERT INTO `{table_name}` (`id`, `label`, `score`) VALUES (900, 'blocked', 900)"), false, ), - MysqlConsoleCancellation::new(), + NativeConsoleCancellation::new(), ) .await .expect_err("read-only datasource must reject writes"); @@ -563,7 +563,7 @@ async fn verify_cancellation( datasource_id: &str, database_name: &str, ) { - let cancellation = MysqlConsoleCancellation::new(); + let cancellation = NativeConsoleCancellation::new(); let execution_cancellation = cancellation.clone(); let execution_application = application.clone(); let execution_request = request( @@ -574,7 +574,7 @@ async fn verify_cancellation( ); let execution = tokio::spawn(async move { execution_application - .execute_mysql_console(execution_request, execution_cancellation) + .execute_native_console(execution_request, execution_cancellation) .await }); @@ -610,10 +610,10 @@ async fn query_count( async fn execute( application: &Application, - request: MysqlConsoleRequest, -) -> Vec { + request: NativeConsoleRequest, +) -> Vec { application - .execute_mysql_console(request, MysqlConsoleCancellation::new()) + .execute_native_console(request, NativeConsoleCancellation::new()) .await .expect("native MySQL Console request must complete") } @@ -623,8 +623,8 @@ fn request( database_name: &str, sql: String, error_continue: bool, -) -> MysqlConsoleRequest { - MysqlConsoleRequest { +) -> NativeConsoleRequest { + NativeConsoleRequest { datasource_id: datasource_id.to_owned(), database_name: database_name.to_owned(), sql, @@ -638,7 +638,7 @@ fn request( } } -fn assert_single_success(results: &[MysqlConsoleResult], update_count: u64) { +fn assert_single_success(results: &[NativeConsoleResult], update_count: u64) { assert_eq!(results.len(), 1); assert!(results[0].success); assert_eq!(results[0].statement_sequence, 1); @@ -646,7 +646,7 @@ fn assert_single_success(results: &[MysqlConsoleResult], update_count: u64) { assert!(results[0].error.is_none()); } -fn statement_value(result: &MysqlConsoleResult) -> &str { +fn statement_value(result: &NativeConsoleResult) -> &str { let row = result.rows.first().expect("result must contain one row"); let value = row.values.first().expect("result must contain one column"); scalar_text(value) diff --git a/crates/chat2db-core/tests/native_mysql_product.rs b/crates/chat2db-core/tests/native_mysql_product.rs index 111f1f2..f2823a5 100644 --- a/crates/chat2db-core/tests/native_mysql_product.rs +++ b/crates/chat2db-core/tests/native_mysql_product.rs @@ -15,7 +15,7 @@ use chat2db_contract::{ ResultPageRequest, StartCommunityTablePreviewRequest, StartQueryRequest, }; use chat2db_core::{ - Application, MysqlConsoleCancellation, MysqlConsoleRequest, RuntimeConfig, RuntimeHost, + Application, NativeConsoleCancellation, NativeConsoleRequest, RuntimeConfig, RuntimeHost, }; use chat2db_java_bridge::{EngineCommand, EngineConfig}; use futures_util::FutureExt as _; @@ -780,10 +780,10 @@ async fn execute_console_preview( datasource_id: &str, database_name: &str, sql: String, -) -> Vec { +) -> Vec { let results = application - .execute_mysql_console( - MysqlConsoleRequest { + .execute_native_console( + NativeConsoleRequest { datasource_id: datasource_id.to_owned(), database_name: database_name.to_owned(), sql, @@ -795,7 +795,7 @@ async fn execute_console_preview( explain: false, error_continue: false, }, - MysqlConsoleCancellation::new(), + NativeConsoleCancellation::new(), ) .await .expect("generated routine invocation SQL must execute through Console"); diff --git a/crates/chat2db-core/tests/native_mysql_ssh_tunnel_docker.rs b/crates/chat2db-core/tests/native_mysql_ssh_tunnel_docker.rs index 24ff15a..c2dafca 100644 --- a/crates/chat2db-core/tests/native_mysql_ssh_tunnel_docker.rs +++ b/crates/chat2db-core/tests/native_mysql_ssh_tunnel_docker.rs @@ -6,8 +6,8 @@ use chat2db_contract::{ SshAuthentication, SshHostKeyVerification, SshTunnelConfig, }; use chat2db_core::{ - Application, MysqlConsoleCancellation, MysqlConsoleRequest, MysqlConsoleResult, RuntimeConfig, - RuntimeHost, + Application, NativeConsoleCancellation, NativeConsoleRequest, NativeConsoleResult, + RuntimeConfig, RuntimeHost, }; use chat2db_java_bridge::{EngineCommand, EngineConfig}; use tempfile::TempDir; @@ -135,11 +135,11 @@ async fn native_mysql_concurrent_queries_share_one_fixed_ssh_tunnel() { fn spawn_sleep_query( application: Application, datasource_id: String, -) -> tokio::task::JoinHandle, chat2db_core::AppError>> { +) -> tokio::task::JoinHandle, chat2db_core::AppError>> { tokio::spawn(async move { application - .execute_mysql_console( - MysqlConsoleRequest { + .execute_native_console( + NativeConsoleRequest { datasource_id, database_name: String::new(), sql: "SELECT SLEEP(2), CONNECTION_ID()".to_owned(), @@ -151,16 +151,16 @@ fn spawn_sleep_query( explain: false, error_continue: false, }, - MysqlConsoleCancellation::new(), + NativeConsoleCancellation::new(), ) .await }) } async fn await_query( - task: tokio::task::JoinHandle, chat2db_core::AppError>>, + task: tokio::task::JoinHandle, chat2db_core::AppError>>, label: &str, -) -> Vec { +) -> Vec { tokio::time::timeout(QUERY_TIMEOUT, task) .await .unwrap_or_else(|_| panic!("{label} tunneled query timed out")) @@ -168,7 +168,7 @@ async fn await_query( .unwrap_or_else(|error| panic!("{label} tunneled query failed: {error}")) } -fn assert_query_succeeded(results: &[MysqlConsoleResult]) { +fn assert_query_succeeded(results: &[NativeConsoleResult]) { assert_eq!(results.len(), 1); assert!(results[0].success); assert_eq!(results[0].rows.len(), 1); From 51e39bce1d6d02ad50320621c3c8955c6d7fd0fb Mon Sep 17 00:00:00 2001 From: zgq Date: Thu, 6 Aug 2026 11:48:14 +0800 Subject: [PATCH 5/5] fix(core): satisfy macos clippy for temporal dml --- crates/chat2db-core/src/mysql_ddl.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/chat2db-core/src/mysql_ddl.rs b/crates/chat2db-core/src/mysql_ddl.rs index 0b8af3b..5e64145 100644 --- a/crates/chat2db-core/src/mysql_ddl.rs +++ b/crates/chat2db-core/src/mysql_ddl.rs @@ -2289,7 +2289,7 @@ fn canonical_temporal(kind: DmlTemporalKind, value: &str) -> Result { - let normalized = normalize_offset_datetime(value).ok_or_else(&invalid)?; + let normalized = normalize_offset_datetime(value).ok_or_else(invalid)?; let parsed = DateTime::parse_from_rfc3339(&normalized).map_err(|_| invalid())?; let offset = if parsed.offset().local_minus_utc() == 0 { "Z".to_owned()