diff --git a/crates/dbx-core/assets/agent-protocol-v1.json b/crates/dbx-core/assets/agent-protocol-v1.json new file mode 100644 index 000000000..7d9abc7e2 --- /dev/null +++ b/crates/dbx-core/assets/agent-protocol-v1.json @@ -0,0 +1,48 @@ +{ + "protocolVersion": 1, + "handshakeMethod": "handshake", + "handshakeResponseFields": [ + "protocolVersion", + "agentProtocolVersion", + "capabilities" + ], + "capabilities": [ + "connect", + "test_connection", + "metadata", + "query", + "paged_query", + "transaction", + "ddl" + ], + "commonMethods": [ + "handshake", + "connect", + "test_connection", + "list_databases", + "list_schemas", + "list_tables", + "list_objects", + "get_object_source", + "get_table_ddl", + "get_columns", + "list_indexes", + "list_foreign_keys", + "list_triggers", + "execute_query", + "execute_query_page", + "fetch_query_page", + "close_query_session", + "execute_transaction", + "disconnect", + "shutdown" + ], + "mongoLegacyMethods": [ + "list_databases", + "list_collections", + "find_documents", + "insert_document", + "update_document", + "delete_document" + ] +} diff --git a/crates/dbx-core/src/db/agent_driver.rs b/crates/dbx-core/src/db/agent_driver.rs index d6a790002..c75aa9847 100644 --- a/crates/dbx-core/src/db/agent_driver.rs +++ b/crates/dbx-core/src/db/agent_driver.rs @@ -5,7 +5,7 @@ use std::sync::{Arc, Mutex}; use std::time::Duration; use serde::de::DeserializeOwned; -use serde::Deserialize; +use serde::{Deserialize, Serialize}; use serde_json::Value; pub const AGENT_PROTOCOL_VERSION: u32 = 1; @@ -81,6 +81,7 @@ pub enum AgentMethod { ListSchemas, ListTables, ListObjects, + GetObjectSource, GetColumns, ListIndexes, ListForeignKeys, @@ -96,6 +97,29 @@ pub enum AgentMethod { } impl AgentMethod { + pub const ALL: [Self; 20] = [ + Self::Handshake, + Self::Connect, + Self::TestConnection, + Self::ListDatabases, + Self::ListSchemas, + Self::ListTables, + Self::ListObjects, + Self::GetObjectSource, + Self::GetTableDdl, + Self::GetColumns, + Self::ListIndexes, + Self::ListForeignKeys, + Self::ListTriggers, + Self::ExecuteQuery, + Self::ExecuteQueryPage, + Self::FetchQueryPage, + Self::CloseQuerySession, + Self::ExecuteTransaction, + Self::Disconnect, + Self::Shutdown, + ]; + pub fn as_str(self) -> &'static str { match self { Self::Handshake => "handshake", @@ -105,11 +129,12 @@ impl AgentMethod { Self::ListSchemas => "list_schemas", Self::ListTables => "list_tables", Self::ListObjects => "list_objects", + Self::GetObjectSource => "get_object_source", + Self::GetTableDdl => "get_table_ddl", 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", @@ -132,6 +157,15 @@ pub enum MongoAgentMethod { } impl MongoAgentMethod { + pub const ALL: [Self; 6] = [ + Self::ListDatabases, + Self::ListCollections, + Self::FindDocuments, + Self::InsertDocument, + Self::UpdateDocument, + Self::DeleteDocument, + ]; + pub fn as_str(self) -> &'static str { match self { Self::ListDatabases => "list_databases", @@ -354,6 +388,15 @@ impl AgentDriverClient { self.call_method(AgentMethod::ListObjects, agent_schema_params(schema)).await } + pub async fn get_object_source( + &mut self, + schema: &str, + name: &str, + object_type: &K, + ) -> Result { + self.call_method(AgentMethod::GetObjectSource, agent_object_source_params(schema, name, object_type)).await + } + pub async fn get_columns( &mut self, schema: &str, @@ -544,6 +587,10 @@ pub fn agent_schema_table_params(schema: &str, table: &str) -> Value { serde_json::json!({ "schema": schema, "table": table }) } +pub fn agent_object_source_params(schema: &str, name: &str, object_type: &K) -> Value { + serde_json::json!({ "schema": schema, "name": name, "object_type": object_type }) +} + pub fn agent_close_query_session_params(session_id: &str) -> Value { serde_json::json!({ "sessionId": session_id }) } @@ -674,11 +721,11 @@ impl Drop for AgentDriverClient { #[cfg(test)] mod tests { use super::{ - 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, + agent_close_query_session_params, agent_handshake_params, agent_java_args, agent_object_source_params, + 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, }; use std::io::Cursor; @@ -794,6 +841,7 @@ mod tests { 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::GetObjectSource.as_str(), "get_object_source"); 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"); @@ -824,6 +872,7 @@ mod tests { let _list_schemas = AgentDriverClient::list_schemas::; let _list_tables = AgentDriverClient::list_tables::; let _list_objects = AgentDriverClient::list_objects::; + let _get_object_source = AgentDriverClient::get_object_source::; let _get_columns = AgentDriverClient::get_columns::; let _list_indexes = AgentDriverClient::list_indexes::; let _list_foreign_keys = AgentDriverClient::list_foreign_keys::; @@ -866,6 +915,10 @@ mod tests { agent_schema_table_params("public", "orders"), serde_json::json!({ "schema": "public", "table": "orders" }) ); + assert_eq!( + agent_object_source_params("public", "active_users", &"VIEW"), + serde_json::json!({ "schema": "public", "name": "active_users", "object_type": "VIEW" }) + ); 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")), @@ -873,6 +926,31 @@ mod tests { ); } + #[test] + fn agent_protocol_matches_contract_file() { + let contract: serde_json::Value = + serde_json::from_str(include_str!("../../assets/agent-protocol-v1.json")).unwrap(); + + assert_eq!(contract["protocolVersion"], AGENT_PROTOCOL_VERSION); + assert_eq!(contract["handshakeMethod"], AgentMethod::Handshake.as_str()); + assert_eq!( + string_array(&contract["handshakeResponseFields"]), + vec!["protocolVersion", "agentProtocolVersion", "capabilities"] + ); + assert_eq!( + string_array(&contract["capabilities"]), + AgentCapability::ALL.iter().map(|method| method.as_str()).collect::>() + ); + assert_eq!( + string_array(&contract["commonMethods"]), + AgentMethod::ALL.iter().map(|method| method.as_str()).collect::>() + ); + assert_eq!( + string_array(&contract["mongoLegacyMethods"]), + MongoAgentMethod::ALL.iter().map(|method| method.as_str()).collect::>() + ); + } + #[test] fn checks_handshake_capability_support() { let handshake = AgentHandshake { @@ -891,4 +969,8 @@ mod tests { assert!(is_unsupported_handshake_error("Agent RPC error (-1): Unknown method: handshake")); assert!(!is_unsupported_handshake_error("Agent RPC error (-1): Connection failed")); } + + fn string_array(value: &serde_json::Value) -> Vec<&str> { + value.as_array().unwrap().iter().map(|item| item.as_str().unwrap()).collect() + } } diff --git a/crates/dbx-core/src/schema.rs b/crates/dbx-core/src/schema.rs index b376fc22c..998f62d70 100644 --- a/crates/dbx-core/src/schema.rs +++ b/crates/dbx-core/src/schema.rs @@ -935,12 +935,7 @@ pub async fn get_object_source_core( } else if let Some(client) = extract_agent(&connections, &pool_key) { drop(connections); let mut client = client.lock().await; - let result: db::ObjectSource = client - .call( - "get_object_source", - serde_json::json!({"schema": schema, "name": name, "object_type": object_type}), - ) - .await?; + let result: db::ObjectSource = client.get_object_source(schema, name, &object_type).await?; return Ok(result); } else { match connections.get(&pool_key).ok_or("Pool not found")? {