diff --git a/crates/dbx-core/src/db/agent_driver.rs b/crates/dbx-core/src/db/agent_driver.rs index af04b31e4..d6a790002 100644 --- a/crates/dbx-core/src/db/agent_driver.rs +++ b/crates/dbx-core/src/db/agent_driver.rs @@ -79,9 +79,18 @@ pub enum AgentMethod { TestConnection, ListDatabases, ListSchemas, + ListTables, + ListObjects, + GetColumns, + ListIndexes, + ListForeignKeys, + ListTriggers, + GetTableDdl, ExecuteQuery, ExecuteQueryPage, FetchQueryPage, + CloseQuerySession, + ExecuteTransaction, Disconnect, Shutdown, } @@ -94,9 +103,18 @@ impl AgentMethod { Self::TestConnection => "test_connection", Self::ListDatabases => "list_databases", Self::ListSchemas => "list_schemas", + Self::ListTables => "list_tables", + Self::ListObjects => "list_objects", + Self::GetColumns => "get_columns", + Self::ListIndexes => "list_indexes", + Self::ListForeignKeys => "list_foreign_keys", + Self::ListTriggers => "list_triggers", + Self::GetTableDdl => "get_table_ddl", Self::ExecuteQuery => "execute_query", Self::ExecuteQueryPage => "execute_query_page", Self::FetchQueryPage => "fetch_query_page", + Self::CloseQuerySession => "close_query_session", + Self::ExecuteTransaction => "execute_transaction", Self::Disconnect => "disconnect", Self::Shutdown => "shutdown", } @@ -328,6 +346,54 @@ impl AgentDriverClient { self.call_method(AgentMethod::ListSchemas, serde_json::json!({ "database": database })).await } + pub async fn list_tables(&mut self, schema: &str) -> Result { + self.call_method(AgentMethod::ListTables, agent_schema_params(schema)).await + } + + pub async fn list_objects(&mut self, schema: &str) -> Result { + self.call_method(AgentMethod::ListObjects, agent_schema_params(schema)).await + } + + pub async fn get_columns( + &mut self, + schema: &str, + table: &str, + ) -> Result { + self.call_method(AgentMethod::GetColumns, agent_schema_table_params(schema, table)).await + } + + pub async fn list_indexes( + &mut self, + schema: &str, + table: &str, + ) -> Result { + self.call_method(AgentMethod::ListIndexes, agent_schema_table_params(schema, table)).await + } + + pub async fn list_foreign_keys( + &mut self, + schema: &str, + table: &str, + ) -> Result { + self.call_method(AgentMethod::ListForeignKeys, agent_schema_table_params(schema, table)).await + } + + pub async fn list_triggers( + &mut self, + schema: &str, + table: &str, + ) -> Result { + self.call_method(AgentMethod::ListTriggers, agent_schema_table_params(schema, table)).await + } + + pub async fn get_table_ddl( + &mut self, + schema: &str, + table: &str, + ) -> Result { + self.call_method(AgentMethod::GetTableDdl, agent_schema_table_params(schema, table)).await + } + pub async fn execute_query(&mut self, params: Value) -> Result { self.call_method(AgentMethod::ExecuteQuery, params).await } @@ -343,6 +409,21 @@ impl AgentDriverClient { self.call_method(AgentMethod::FetchQueryPage, params).await } + pub async fn close_query_session( + &mut self, + session_id: &str, + ) -> Result { + self.call_method(AgentMethod::CloseQuerySession, agent_close_query_session_params(session_id)).await + } + + pub async fn execute_transaction( + &mut self, + statements: &[String], + schema: Option<&str>, + ) -> Result { + self.call_method(AgentMethod::ExecuteTransaction, agent_transaction_params(statements, schema)).await + } + pub async fn call_mongo_method( &mut self, method: MongoAgentMethod, @@ -455,6 +536,25 @@ pub fn is_unsupported_handshake_error(error: &str) -> bool { || error.contains("method not found: handshake") } +pub fn agent_schema_params(schema: &str) -> Value { + serde_json::json!({ "schema": schema }) +} + +pub fn agent_schema_table_params(schema: &str, table: &str) -> Value { + serde_json::json!({ "schema": schema, "table": table }) +} + +pub fn agent_close_query_session_params(session_id: &str) -> Value { + serde_json::json!({ "sessionId": session_id }) +} + +pub fn agent_transaction_params(statements: &[String], schema: Option<&str>) -> Value { + serde_json::json!({ + "statements": statements, + "schema": schema, + }) +} + pub fn mongo_database_params(database: &str) -> Value { serde_json::json!({ "database": database }) } @@ -574,7 +674,8 @@ impl Drop for AgentDriverClient { #[cfg(test)] mod tests { use super::{ - agent_handshake_params, agent_java_args, agent_proxy_env_vars, format_agent_process_error, + agent_close_query_session_params, agent_handshake_params, agent_java_args, agent_proxy_env_vars, + agent_schema_params, agent_schema_table_params, agent_transaction_params, format_agent_process_error, is_unsupported_handshake_error, mongo_collection_params, mongo_database_params, mongo_document_id_params, read_agent_line, AgentCapability, AgentDriverClient, AgentHandshake, AgentMethod, MongoAgentMethod, StderrTail, AGENT_PROTOCOL_VERSION, @@ -691,9 +792,18 @@ mod tests { assert_eq!(AgentMethod::TestConnection.as_str(), "test_connection"); assert_eq!(AgentMethod::ListDatabases.as_str(), "list_databases"); assert_eq!(AgentMethod::ListSchemas.as_str(), "list_schemas"); + assert_eq!(AgentMethod::ListTables.as_str(), "list_tables"); + assert_eq!(AgentMethod::ListObjects.as_str(), "list_objects"); + assert_eq!(AgentMethod::GetColumns.as_str(), "get_columns"); + assert_eq!(AgentMethod::ListIndexes.as_str(), "list_indexes"); + assert_eq!(AgentMethod::ListForeignKeys.as_str(), "list_foreign_keys"); + assert_eq!(AgentMethod::ListTriggers.as_str(), "list_triggers"); + assert_eq!(AgentMethod::GetTableDdl.as_str(), "get_table_ddl"); assert_eq!(AgentMethod::ExecuteQuery.as_str(), "execute_query"); assert_eq!(AgentMethod::ExecuteQueryPage.as_str(), "execute_query_page"); assert_eq!(AgentMethod::FetchQueryPage.as_str(), "fetch_query_page"); + assert_eq!(AgentMethod::CloseQuerySession.as_str(), "close_query_session"); + assert_eq!(AgentMethod::ExecuteTransaction.as_str(), "execute_transaction"); assert_eq!(AgentMethod::Disconnect.as_str(), "disconnect"); assert_eq!(AgentMethod::Shutdown.as_str(), "shutdown"); } @@ -712,9 +822,18 @@ mod tests { fn exposes_schema_and_query_protocol_wrappers() { let _list_databases = AgentDriverClient::list_databases::; let _list_schemas = AgentDriverClient::list_schemas::; + let _list_tables = AgentDriverClient::list_tables::; + let _list_objects = AgentDriverClient::list_objects::; + let _get_columns = AgentDriverClient::get_columns::; + let _list_indexes = AgentDriverClient::list_indexes::; + let _list_foreign_keys = AgentDriverClient::list_foreign_keys::; + let _list_triggers = AgentDriverClient::list_triggers::; + let _get_table_ddl = AgentDriverClient::get_table_ddl::; let _execute_query = AgentDriverClient::execute_query::; let _execute_query_page = AgentDriverClient::execute_query_page::; let _fetch_query_page = AgentDriverClient::fetch_query_page::; + let _close_query_session = AgentDriverClient::close_query_session::; + let _execute_transaction = AgentDriverClient::execute_transaction::; } #[test] @@ -740,6 +859,20 @@ mod tests { ); } + #[test] + fn builds_schema_table_and_transaction_params() { + assert_eq!(agent_schema_params("public"), serde_json::json!({ "schema": "public" })); + assert_eq!( + agent_schema_table_params("public", "orders"), + serde_json::json!({ "schema": "public", "table": "orders" }) + ); + assert_eq!(agent_close_query_session_params("session-1"), serde_json::json!({ "sessionId": "session-1" })); + assert_eq!( + agent_transaction_params(&["BEGIN".to_string(), "COMMIT".to_string()], Some("public")), + serde_json::json!({ "statements": ["BEGIN", "COMMIT"], "schema": "public" }) + ); + } + #[test] fn checks_handshake_capability_support() { let handshake = AgentHandshake { diff --git a/crates/dbx-core/src/query.rs b/crates/dbx-core/src/query.rs index ecdd67913..4ca1be5b0 100644 --- a/crates/dbx-core/src/query.rs +++ b/crates/dbx-core/src/query.rs @@ -548,7 +548,7 @@ pub async fn close_query_session( let client = client.clone(); drop(connections); let mut client = client.lock().await; - client.call("close_query_session", agent_close_query_session_params(session_id)).await + client.close_query_session(session_id).await } _ => Ok(false), } @@ -953,11 +953,7 @@ async fn exec_tx_explicit_inner( let conns = state.connections.read().await; if let Some(crate::connection::PoolKind::Agent(client)) = conns.get(pool_key) { let mut client = client.lock().await; - let params = serde_json::json!({ - "statements": statements, - "schema": schema, - }); - let result: db::QueryResult = client.call("execute_transaction", params).await?; + let result: db::QueryResult = client.execute_transaction(statements, schema).await?; return Ok(db::QueryResult { execution_time_ms: start.elapsed().as_millis(), ..result }); } drop(conns); diff --git a/crates/dbx-core/src/schema.rs b/crates/dbx-core/src/schema.rs index 411086e33..b376fc22c 100644 --- a/crates/dbx-core/src/schema.rs +++ b/crates/dbx-core/src/schema.rs @@ -373,7 +373,7 @@ pub async fn list_tables_core( if let Some(client) = extract_agent(&connections, &pool_key) { drop(connections); let mut client = client.lock().await; - return client.call("list_tables", serde_json::json!({"schema": schema})).await; + return client.list_tables(schema).await; } } @@ -507,7 +507,7 @@ pub async fn list_objects_core( if let Some(client) = extract_agent(&connections, &pool_key) { drop(connections); let mut client = client.lock().await; - return client.call("list_objects", serde_json::json!({"schema": schema})).await; + return client.list_objects(schema).await; } } @@ -596,7 +596,7 @@ pub async fn get_columns_core( if let Some(client) = extract_agent(&connections, &pool_key) { drop(connections); let mut client = client.lock().await; - return client.call("get_columns", serde_json::json!({"schema": schema, "table": table})).await; + return client.get_columns(schema, table).await; } } @@ -636,7 +636,7 @@ pub async fn list_indexes_core( if let Some(client) = extract_agent(&connections, &pool_key) { drop(connections); let mut client = client.lock().await; - return client.call("list_indexes", serde_json::json!({"schema": schema, "table": table})).await; + return client.list_indexes(schema, table).await; } } @@ -676,7 +676,7 @@ pub async fn list_foreign_keys_core( if let Some(client) = extract_agent(&connections, &pool_key) { drop(connections); let mut client = client.lock().await; - return client.call("list_foreign_keys", serde_json::json!({"schema": schema, "table": table})).await; + return client.list_foreign_keys(schema, table).await; } } @@ -716,7 +716,7 @@ pub async fn list_triggers_core( if let Some(client) = extract_agent(&connections, &pool_key) { drop(connections); let mut client = client.lock().await; - return client.call("list_triggers", serde_json::json!({"schema": schema, "table": table})).await; + return client.list_triggers(schema, table).await; } } @@ -786,7 +786,7 @@ pub async fn get_table_ddl_core( if let Some(client) = extract_agent(&connections, &pool_key) { drop(connections); let mut client = client.lock().await; - return client.call("get_table_ddl", serde_json::json!({"schema": schema, "table": table})).await; + return client.get_table_ddl(schema, table).await; } }