fix(agent): reserve shared runtime atomically

This commit is contained in:
zipg 2026-07-15 13:20:16 +08:00 committed by GitHub
parent 55f15c8e4b
commit ebebc4fccd
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 166 additions and 14 deletions

View File

@ -67,28 +67,90 @@ pub async fn spawn_shared_connection_client(
let jre_key = state.installed_drivers.get(key).map(|driver| driver.jre.as_str()).unwrap_or(DEFAULT_JRE_KEY);
let launch = manager.resolve_agent_launch_spec_with_extra_args(&state, key, jre_key, extra_java_args)?;
let runtime_key = shared_runtime_key(key, &launch);
let runtime_cell = {
let mut runtimes = manager.connection_runtimes.lock().await;
if runtimes.get(&runtime_key).and_then(|cell| cell.get()).is_some_and(|runtime| runtime.is_failed()) {
runtimes.remove(&runtime_key);
}
runtimes.entry(runtime_key).or_insert_with(|| std::sync::Arc::new(tokio::sync::OnceCell::new())).clone()
};
let runtime =
runtime_cell.get_or_try_init(|| AgentRuntimeClient::spawn(launch, manager.agent_app_version())).await?.clone();
let mut session_params = connect_params;
session_params
.as_object_mut()
.ok_or_else(|| "Agent connect parameters must be an object".to_string())?
.insert("agentSessionId".to_string(), serde_json::Value::String(agent_session_id.clone()));
runtime
let (runtime_cell, runtime) = loop {
let runtime_cell = {
let mut runtimes = manager.connection_runtimes.lock().await;
if runtimes.get(&runtime_key).and_then(|cell| cell.get()).is_some_and(|runtime| runtime.is_failed()) {
runtimes.remove(&runtime_key);
}
runtimes
.entry(runtime_key.clone())
.or_insert_with(|| std::sync::Arc::new(tokio::sync::OnceCell::new()))
.clone()
};
let runtime = runtime_cell
.get_or_try_init(|| AgentRuntimeClient::spawn(launch.clone(), manager.agent_app_version()))
.await?
.clone();
let mut runtimes = manager.connection_runtimes.lock().await;
if reserve_runtime_locked(&mut runtimes, &runtime_key, &runtime_cell, &runtime) {
break (runtime_cell, runtime);
}
};
if let Err(err) = runtime
.call::<serde_json::Value>(AgentMethod::OpenSession.as_str(), session_params, Some(connect_timeout), None)
.await?;
runtime.increment_session_count();
.await
{
forget_unused_runtime_after_failed_open(manager, &runtime_key, &runtime_cell, &runtime).await;
return Err(err);
}
Ok(AgentDriverClient::shared_session(runtime, agent_session_id))
}
async fn forget_unused_runtime_after_failed_open(
manager: &AgentManager,
runtime_key: &str,
runtime_cell: &std::sync::Arc<tokio::sync::OnceCell<std::sync::Arc<AgentRuntimeClient>>>,
runtime: &std::sync::Arc<AgentRuntimeClient>,
) {
if AgentRuntimeClient::decrement_session_count(runtime) != 0 {
return;
}
remove_unused_runtime_if_current(manager, runtime_key, runtime_cell, runtime).await;
}
async fn remove_unused_runtime_if_current(
manager: &AgentManager,
runtime_key: &str,
runtime_cell: &std::sync::Arc<tokio::sync::OnceCell<std::sync::Arc<AgentRuntimeClient>>>,
runtime: &std::sync::Arc<AgentRuntimeClient>,
) {
// Keep reservation and map-entry validation under the same lock as openers.
let mut runtimes = manager.connection_runtimes.lock().await;
if runtime.active_session_count() == 0
&& runtimes.get(runtime_key).is_some_and(|current| std::sync::Arc::ptr_eq(current, runtime_cell))
{
runtimes.remove(runtime_key);
}
}
fn reserve_runtime_locked(
runtimes: &mut std::collections::HashMap<
String,
std::sync::Arc<tokio::sync::OnceCell<std::sync::Arc<AgentRuntimeClient>>>,
>,
runtime_key: &str,
runtime_cell: &std::sync::Arc<tokio::sync::OnceCell<std::sync::Arc<AgentRuntimeClient>>>,
runtime: &std::sync::Arc<AgentRuntimeClient>,
) -> bool {
if runtimes.get(runtime_key).is_some_and(|current| std::sync::Arc::ptr_eq(current, runtime_cell))
&& !runtime.is_failed()
{
runtime.increment_session_count();
true
} else {
false
}
}
fn shared_runtime_key(agent_key: &str, launch: &crate::db::agent_driver::AgentLaunchSpec) -> String {
format!(
"{}|{}|{}|{}",
@ -228,6 +290,41 @@ mod tests {
use super::*;
use std::path::PathBuf;
async fn test_shared_runtime(
name: &str,
) -> (
AgentManager,
std::sync::Arc<tokio::sync::OnceCell<std::sync::Arc<AgentRuntimeClient>>>,
std::sync::Arc<AgentRuntimeClient>,
PathBuf,
) {
let test_id = uuid::Uuid::new_v4();
let script_path = std::env::temp_dir().join(format!("dbx-agent-runtime-{name}-{test_id}.py"));
std::fs::write(
&script_path,
r#"import json, sys
print(json.dumps({'ready': True}), flush=True)
for line in sys.stdin:
req = json.loads(line)
result = {'protocolVersion': 2, 'agentProtocolVersion': 2, 'capabilities': ['multi_session']} if req['method'] == 'handshake' else {'ok': True}
print(json.dumps({'jsonrpc': '2.0', 'id': req['id'], 'result': result}), flush=True)
"#,
)
.unwrap();
let runtime = AgentRuntimeClient::spawn(
crate::db::agent_driver::AgentLaunchSpec::new("python3")
.with_args([script_path.to_string_lossy().to_string()]),
"test",
)
.await
.unwrap();
let cell = std::sync::Arc::new(tokio::sync::OnceCell::new());
assert!(cell.set(runtime.clone()).is_ok());
let manager_dir = std::env::temp_dir().join(format!("dbx-agent-manager-{name}-{test_id}"));
let manager = AgentManager::new_with_base_dir_and_app_version(manager_dir, "test");
(manager, cell, runtime, script_path)
}
#[test]
fn prestosql_does_not_use_agent_driver() {
assert_eq!(runtime_agent_key_candidates(&DatabaseType::PrestoSql, None), None);
@ -250,4 +347,57 @@ mod tests {
assert_eq!(shared_runtime_key("oracle", &base), shared_runtime_key("oracle", &base));
assert_ne!(shared_runtime_key("oracle", &base), shared_runtime_key("oracle", &different_args));
}
#[tokio::test]
async fn failed_open_forgets_runtime_when_no_other_session_uses_it() {
let (manager, cell, runtime, script_path) = test_shared_runtime("failed-open-unused").await;
let runtime_key = "kingbase|test";
manager.connection_runtimes.lock().await.insert(runtime_key.to_string(), cell.clone());
runtime.increment_session_count();
forget_unused_runtime_after_failed_open(&manager, runtime_key, &cell, &runtime).await;
assert!(!manager.connection_runtimes.lock().await.contains_key(runtime_key));
assert_eq!(runtime.active_session_count(), 0);
runtime.kill();
let _ = std::fs::remove_file(script_path);
}
#[tokio::test]
async fn failed_open_keeps_runtime_while_another_session_is_reserved() {
let (manager, cell, runtime, script_path) = test_shared_runtime("failed-open-in-use").await;
let runtime_key = "oracle|test";
manager.connection_runtimes.lock().await.insert(runtime_key.to_string(), cell.clone());
runtime.increment_session_count();
runtime.increment_session_count();
forget_unused_runtime_after_failed_open(&manager, runtime_key, &cell, &runtime).await;
assert!(manager.connection_runtimes.lock().await.contains_key(runtime_key));
assert_eq!(runtime.active_session_count(), 1);
AgentRuntimeClient::decrement_session_count(&runtime);
runtime.kill();
let _ = std::fs::remove_file(script_path);
}
#[tokio::test]
async fn failed_open_cleanup_cannot_remove_runtime_after_reservation() {
let (manager, cell, runtime, script_path) = test_shared_runtime("failed-open-race").await;
let runtime_key = "oracle|test";
manager.connection_runtimes.lock().await.insert(runtime_key.to_string(), cell.clone());
runtime.increment_session_count();
assert_eq!(AgentRuntimeClient::decrement_session_count(&runtime), 0);
let mut runtimes = manager.connection_runtimes.lock().await;
assert!(reserve_runtime_locked(&mut runtimes, runtime_key, &cell, &runtime));
drop(runtimes);
remove_unused_runtime_if_current(&manager, runtime_key, &cell, &runtime).await;
assert!(manager.connection_runtimes.lock().await.contains_key(runtime_key));
assert_eq!(runtime.active_session_count(), 1);
AgentRuntimeClient::decrement_session_count(&runtime);
runtime.kill();
let _ = std::fs::remove_file(script_path);
}
}

View File

@ -203,11 +203,12 @@ impl AgentRuntimeClient {
self.active_sessions.fetch_add(1, Ordering::AcqRel);
}
pub fn decrement_session_count(runtime: &Arc<Self>) {
pub fn decrement_session_count(runtime: &Arc<Self>) -> u64 {
let previous = runtime
.active_sessions
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |value| Some(value.saturating_sub(1)))
.unwrap_or_default();
let remaining = previous.saturating_sub(1);
if previous <= 1 {
let runtime = runtime.clone();
tokio::spawn(async move {
@ -217,6 +218,7 @@ impl AgentRuntimeClient {
}
});
}
remaining
}
pub fn kill(&self) {