fix(agent): wait for pool close before reconnect

Closes #5251
This commit is contained in:
t8y2 2026-08-04 03:17:20 +08:00
parent 6f0b4666f8
commit a7c4ab6714
No known key found for this signature in database
1 changed files with 103 additions and 5 deletions

View File

@ -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;