From a7c4ab67146826dcbb308272558cb5d98adb9f9e Mon Sep 17 00:00:00 2001 From: t8y2 <1156263951@qq.com> Date: Tue, 4 Aug 2026 03:17:20 +0800 Subject: [PATCH] fix(agent): wait for pool close before reconnect Closes #5251 --- crates/dbx-core/src/connection.rs | 108 ++++++++++++++++++++++++++++-- 1 file changed, 103 insertions(+), 5 deletions(-) diff --git a/crates/dbx-core/src/connection.rs b/crates/dbx-core/src/connection.rs index 9255a029c..ea8725ca6 100644 --- a/crates/dbx-core/src/connection.rs +++ b/crates/dbx-core/src/connection.rs @@ -3038,11 +3038,7 @@ impl AppState { self.postgres_cancel_contexts.write().await.remove(&pool_key); let removed = self.connections.write().await.remove(&pool_key); if let Some(pool) = removed { - if matches!(&pool, PoolKind::Agent(_)) { - self.pool_routing_control().close_removed_in_background(vec![(pool_key.clone(), pool)]); - } else { - self.pool_routing_control().close_pool_with_timeout(pool_key.clone(), pool).await; - } + self.pool_routing_control().close_pool_with_timeout(pool_key.clone(), pool).await; } } self.get_or_create_pool_for_session_inner( @@ -6699,6 +6695,108 @@ for line in sys.stdin: let _ = std::fs::remove_dir_all(dir); } + #[tokio::test] + async fn reconnect_waits_for_agent_runtime_replacement_before_publishing_new_pool() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + + let (state, dir) = test_app_state().await; + let script_path = dir.join("delayed-close-replace-runtime-agent.py"); + let close_finished_path = dir.join("close-finished"); + let close_finished = serde_json::to_string(&close_finished_path.to_string_lossy()).unwrap(); + std::fs::write( + &script_path, + format!( + r#"import json, pathlib, sys, time +close_finished = pathlib.Path({close_finished}) +print(json.dumps({{'ready': True}}), flush=True) +for line in sys.stdin: + req = json.loads(line) + if req['method'] == 'handshake': + response = {{ + 'jsonrpc': '2.0', + 'id': req['id'], + 'result': {{'protocolVersion': 2, 'agentProtocolVersion': 2, 'capabilities': ['multi_session']}} + }} + elif req['method'] == 'close_session': + time.sleep(0.5) + close_finished.write_text('done') + response = {{ + 'jsonrpc': '2.0', + 'id': req['id'], + 'error': {{ + 'code': -1, + 'message': 'Agent runtime resource limit reached', + 'data': {{ + 'category': 'resource', + 'retryable': False, + 'sessionDisposition': 'replace_runtime', + 'stage': 'close' + }} + }} + }} + else: + response = {{'jsonrpc': '2.0', 'id': req['id'], 'result': {{}}}} + print(json.dumps(response), flush=True) +"# + ), + ) + .unwrap(); + + let python = if cfg!(windows) { "python" } else { "python3" }; + let runtime = crate::db::agent_driver::AgentRuntimeClient::spawn( + crate::db::agent_driver::AgentLaunchSpec::new(python) + .with_args([script_path.to_string_lossy().to_string()]), + "test", + ) + .await + .unwrap(); + runtime.increment_session_count(); + let client = std::sync::Arc::new(crate::db::agent_driver::PooledAgentClient::new( + crate::db::agent_driver::AgentDriverClient::shared_session( + runtime.clone(), + "metadata-agent-session".to_string(), + ), + )); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut request = Vec::new(); + let mut chunk = [0_u8; 1024]; + while !request.windows(4).any(|window| window == b"\r\n\r\n") { + let read = socket.read(&mut chunk).await.unwrap(); + assert!(read > 0, "rqlite probe ended before the request headers were complete"); + request.extend_from_slice(&chunk[..read]); + } + assert!(String::from_utf8_lossy(&request).starts_with("POST /db/query HTTP/1.1")); + let body = r#"{"results":[{"columns":["1"],"values":[[1]]}]}"#; + let response = format!( + "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}", + body.len() + ); + socket.write_all(response.as_bytes()).await.unwrap(); + }); + let mut config = mysql_config(None); + config.id = "conn".to_string(); + config.db_type = DatabaseType::Rqlite; + config.host = address.ip().to_string(); + config.port = address.port(); + config.keepalive_interval_secs = 0; + state.configs.write().await.insert(config.id.clone(), config); + state.connections.write().await.insert("conn".to_string(), PoolKind::Agent(client)); + + let pool_key = state.reconnect_metadata_pool_for_session("conn", None, None).await.unwrap(); + server.await.unwrap(); + + assert_eq!(pool_key, "conn"); + assert!(close_finished_path.exists(), "reconnect returned before the old Agent close completed"); + assert!(runtime.is_failed()); + assert!(matches!(state.connections.read().await.get("conn"), Some(PoolKind::Rqlite(_)))); + + state.shutdown(Duration::from_secs(1)).await; + let _ = std::fs::remove_dir_all(dir); + } + #[tokio::test] async fn inserting_agent_pool_rejects_runtime_failed_before_publish() { let (state, dir) = test_app_state().await;