test(agent): validate protocol contract

This commit is contained in:
t8y2 2026-05-19 16:05:25 +08:00
parent 613c2f981f
commit fa2e710a7b
3 changed files with 138 additions and 13 deletions

View File

@ -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"
]
}

View File

@ -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<T: DeserializeOwned + Send + 'static, K: Serialize>(
&mut self,
schema: &str,
name: &str,
object_type: &K,
) -> Result<T, String> {
self.call_method(AgentMethod::GetObjectSource, agent_object_source_params(schema, name, object_type)).await
}
pub async fn get_columns<T: DeserializeOwned + Send + 'static>(
&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<K: Serialize>(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::<serde_json::Value>;
let _list_tables = AgentDriverClient::list_tables::<serde_json::Value>;
let _list_objects = AgentDriverClient::list_objects::<serde_json::Value>;
let _get_object_source = AgentDriverClient::get_object_source::<serde_json::Value, serde_json::Value>;
let _get_columns = AgentDriverClient::get_columns::<serde_json::Value>;
let _list_indexes = AgentDriverClient::list_indexes::<serde_json::Value>;
let _list_foreign_keys = AgentDriverClient::list_foreign_keys::<serde_json::Value>;
@ -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::<Vec<_>>()
);
assert_eq!(
string_array(&contract["commonMethods"]),
AgentMethod::ALL.iter().map(|method| method.as_str()).collect::<Vec<_>>()
);
assert_eq!(
string_array(&contract["mongoLegacyMethods"]),
MongoAgentMethod::ALL.iter().map(|method| method.as_str()).collect::<Vec<_>>()
);
}
#[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()
}
}

View File

@ -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")? {