From 2ebd72ab6dfdd73e0c80b8163799f3ccd35c1991 Mon Sep 17 00:00:00 2001 From: t8y2 <1156263951@qq.com> Date: Tue, 19 May 2026 16:07:25 +0800 Subject: [PATCH] refactor(agent): retain handshake capabilities --- crates/dbx-core/src/db/agent_driver.rs | 37 ++++++++++++++++++++++---- 1 file changed, 32 insertions(+), 5 deletions(-) diff --git a/crates/dbx-core/src/db/agent_driver.rs b/crates/dbx-core/src/db/agent_driver.rs index c75aa9847..08714a621 100644 --- a/crates/dbx-core/src/db/agent_driver.rs +++ b/crates/dbx-core/src/db/agent_driver.rs @@ -20,6 +20,7 @@ pub struct AgentDriverClient { stdin: Option>, stdout: Option>, stderr_tail: Arc>, + handshake: Option, next_id: u64, } @@ -276,7 +277,7 @@ impl AgentDriverClient { } }; - Ok(Self { child, stdin: Some(stdin), stdout: Some(ready_stdout), stderr_tail, next_id: 0 }) + Ok(Self { child, stdin: Some(stdin), stdout: Some(ready_stdout), stderr_tail, handshake: None, next_id: 0 }) } /// Send a JSON-RPC 2.0 request and wait for the response. @@ -523,6 +524,7 @@ impl AgentDriverClient { handshake.agent_protocol_version, handshake.capabilities ); + self.handshake = Some(handshake.clone()); Some(handshake) } Err(err) if is_unsupported_handshake_error(&err) => { @@ -536,6 +538,14 @@ impl AgentDriverClient { } } + pub fn handshake(&self) -> Option<&AgentHandshake> { + self.handshake.as_ref() + } + + pub fn supports_capability(&self, capability: AgentCapability) -> bool { + agent_supports_capability(self.handshake.as_ref(), capability) + } + /// Send a shutdown message to the agent and wait for the process to exit. pub async fn shutdown(&mut self) { // Try to send a shutdown RPC; ignore errors if the agent is already gone @@ -579,6 +589,10 @@ pub fn is_unsupported_handshake_error(error: &str) -> bool { || error.contains("method not found: handshake") } +pub fn agent_supports_capability(handshake: Option<&AgentHandshake>, capability: AgentCapability) -> bool { + handshake.map(|value| value.supports(capability)).unwrap_or(true) +} + pub fn agent_schema_params(schema: &str) -> Value { serde_json::json!({ "schema": schema }) } @@ -722,10 +736,10 @@ impl Drop for AgentDriverClient { mod tests { use super::{ 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, + agent_proxy_env_vars, agent_schema_params, agent_schema_table_params, agent_supports_capability, + 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; @@ -964,6 +978,19 @@ mod tests { assert!(!handshake.supports(AgentCapability::Query)); } + #[test] + fn treats_missing_handshake_as_legacy_capability_support() { + let handshake = AgentHandshake { + protocol_version: AGENT_PROTOCOL_VERSION, + agent_protocol_version: AGENT_PROTOCOL_VERSION, + capabilities: vec!["connect".to_string()], + }; + + assert!(agent_supports_capability(None, AgentCapability::Query)); + assert!(agent_supports_capability(Some(&handshake), AgentCapability::Connect)); + assert!(!agent_supports_capability(Some(&handshake), AgentCapability::Query)); + } + #[test] fn treats_unknown_handshake_method_as_compatible_fallback() { assert!(is_unsupported_handshake_error("Agent RPC error (-1): Unknown method: handshake"));