From 0e414e638ec80ceff135875733ed47ba452d3b15 Mon Sep 17 00:00:00 2001 From: miracle Date: Sat, 1 Aug 2026 01:17:49 +0800 Subject: [PATCH] fix(agent): recover JDBC connections after query timeouts --- .github/workflows/ci.yml | 14 +- .../java/com/dbx/agent/AbstractJdbcAgent.java | 112 +- .../java/com/dbx/agent/AgentRpcError.java | 135 ++ .../java/com/dbx/agent/ConnectParams.java | 15 +- .../dbx/agent/JdbcConnectionPoolRegistry.java | 1381 ++++++++++++- .../java/com/dbx/agent/JdbcSessionRole.java | 10 + .../java/com/dbx/agent/JsonRpcServer.java | 5 + .../dbx/agent/MultiSessionJsonRpcServer.java | 158 +- .../agent/CommonJavaCompatibilityTest.java | 13 +- .../dbx/agent/JdbcConnectionPoolingTest.java | 1707 ++++++++++++++++- agents/docs/agent-protocol-v2.md | 20 +- agents/drivers/rabbitmq/integration_test.go | 20 +- .../connectionStore.metadataLoading.spec.ts | 255 +++ apps/desktop/src/stores/connectionStore.ts | 128 +- crates/dbx-core/src/agent_connection.rs | 45 + crates/dbx-core/src/agent_manager.rs | 2 +- crates/dbx-core/src/agent_runtime.rs | 47 +- crates/dbx-core/src/connection.rs | 1396 ++++++++++++-- crates/dbx-core/src/database_export.rs | 4 +- crates/dbx-core/src/db/agent_driver.rs | 236 ++- crates/dbx-core/src/query.rs | 315 ++- crates/dbx-core/src/schema.rs | 424 +++- crates/dbx-core/src/schema/kingbase.rs | 8 +- src-tauri/src/commands/connection.rs | 6 +- 24 files changed, 5983 insertions(+), 473 deletions(-) create mode 100644 agents/common/src/main/java/com/dbx/agent/AgentRpcError.java create mode 100644 agents/common/src/main/java/com/dbx/agent/JdbcSessionRole.java diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 1b7c59afd..1dcc5224d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -557,24 +557,34 @@ jobs: set -euo pipefail for version in 3.13 4.3; do name="dbx-rabbitmq-${version//./-}" + cookie="dbx-ci-${GITHUB_RUN_ID:-local}-${version//./-}" + # RabbitMQ is disposable in CI. Keep its data on a fresh tmpfs + # owned by the image's rabbitmq user so .erlang.cookie is readable. + docker rm -fv "$name" >/dev/null 2>&1 || true docker run -d --name "$name" \ + --user 999:999 \ + --tmpfs /var/lib/rabbitmq:rw,exec,uid=999,gid=999,mode=700 \ -e RABBITMQ_DEFAULT_USER=dbx \ -e RABBITMQ_DEFAULT_PASS=dbx-password \ + -e RABBITMQ_ERLANG_COOKIE="$cookie" \ -p 5672:5672 -p 15672:15672 \ "rabbitmq:${version}-management" cleanup() { - docker rm -f "$name" >/dev/null 2>&1 || true + docker rm -fv "$name" >/dev/null 2>&1 || true } trap cleanup EXIT ready=false for _ in $(seq 1 60); do - if docker exec "$name" rabbitmq-diagnostics -q ping >/dev/null 2>&1; then + if docker exec "$name" rabbitmq-diagnostics -q check_running >/dev/null 2>&1 \ + && curl --fail --silent --noproxy '*' --user dbx:dbx-password \ + http://127.0.0.1:15672/api/overview >/dev/null; then ready=true break fi sleep 2 done if [ "$ready" != "true" ]; then + docker inspect "$name" --format 'image={{.Config.Image}} user={{.Config.User}} status={{.State.Status}} exit={{.State.ExitCode}}' || true docker logs "$name" exit 1 fi diff --git a/agents/common/src/main/java/com/dbx/agent/AbstractJdbcAgent.java b/agents/common/src/main/java/com/dbx/agent/AbstractJdbcAgent.java index 03acc1b94..67a0945b1 100644 --- a/agents/common/src/main/java/com/dbx/agent/AbstractJdbcAgent.java +++ b/agents/common/src/main/java/com/dbx/agent/AbstractJdbcAgent.java @@ -29,6 +29,7 @@ public abstract class AbstractJdbcAgent extends BaseDatabaseAgent { private boolean requestActive; private boolean leasePinnedAtRequestStart; private boolean sessionAffinity; + private boolean pooledConnectionPoisoned; @Override public final Connection getConnection() { @@ -47,6 +48,7 @@ public abstract class AbstractJdbcAgent extends BaseDatabaseAgent { sessionAffinity = false; requestActive = false; leasePinnedAtRequestStart = false; + pooledConnectionPoisoned = false; loadDriver(params); configuredDatabase = params.getDatabase(); connectParams = params; @@ -256,6 +258,7 @@ public abstract class AbstractJdbcAgent extends BaseDatabaseAgent { sessionAffinity = false; requestActive = false; leasePinnedAtRequestStart = false; + pooledConnectionPoisoned = false; identifierQuote = ""; }); } @@ -271,24 +274,54 @@ public abstract class AbstractJdbcAgent extends BaseDatabaseAgent { return poolRegistry != null; } - final synchronized void beginPooledRequest() throws Exception { - if (poolRegistry == null) { - return; - } - if (connectParams == null) { - throw new IllegalStateException("Not connected"); - } - requestActive = true; - leasePinnedAtRequestStart = pooledLease != null; - if (pooledLease == null) { - try { - pooledLease = borrowPooledConnection(); - } catch (Exception error) { - requestActive = false; - throw error; + final synchronized boolean quarantinePooledConnection() { + pooledConnectionPoisoned = true; + return requestActive && pooledLease != null && pooledLease.quarantine(); + } + + final void beginPooledRequest() throws Exception { + ConnectParams params; + String identity; + JdbcConnectionPoolRegistry registry; + synchronized (this) { + if (poolRegistry == null) { + return; } + if (connectParams == null || poolIdentity == null) { + throw new IllegalStateException("Not connected"); + } + requestActive = true; + leasePinnedAtRequestStart = pooledLease != null; + if (pooledLease != null) { + connection = pooledLease.connection(); + return; + } + params = connectParams; + identity = poolIdentity; + registry = poolRegistry; + } + + JdbcConnectionPoolRegistry.Lease borrowed; + try { + borrowed = borrowPooledConnection(registry, identity, params); + } catch (Exception error) { + synchronized (this) { + requestActive = false; + leasePinnedAtRequestStart = false; + } + throw error; + } + + synchronized (this) { + if (pooledConnectionPoisoned) { + requestActive = false; + leasePinnedAtRequestStart = false; + borrowed.evict(); + throw new IllegalStateException("JDBC Session was quarantined while waiting for a connection"); + } + pooledLease = borrowed; + connection = borrowed.connection(); } - connection = pooledLease.connection(); } final synchronized void finishPooledRequest( @@ -301,6 +334,10 @@ public abstract class AbstractJdbcAgent extends BaseDatabaseAgent { return; } requestActive = false; + if (pooledConnectionPoisoned) { + releasePooledConnection(true); + return; + } if (succeeded && requiresSessionAffinity) { sessionAffinity = true; JdbcSchemaSwitcher.forget(connection); @@ -322,7 +359,14 @@ public abstract class AbstractJdbcAgent extends BaseDatabaseAgent { } final synchronized void releaseIdlePooledConnection(JdbcExecutor executor) { - if (poolRegistry == null || requestActive || sessionAffinity || pooledLease == null) { + if (poolRegistry == null || requestActive || pooledLease == null) { + return; + } + if (pooledConnectionPoisoned) { + releasePooledConnection(true); + return; + } + if (sessionAffinity) { return; } if (executor.hasOpenSessions() || executor.hasActiveStatements()) { @@ -489,18 +533,20 @@ public abstract class AbstractJdbcAgent extends BaseDatabaseAgent { private Connection openInitializedConnection(ConnectParams params) throws Exception { Connection opened = openConnection(params); - boolean initialized = false; try { afterPhysicalConnect(params, opened); - initialized = true; return opened; - } finally { - if (!initialized) { - try { - opened.close(); - } catch (Exception ignored) { - } + } catch (Exception error) { + try { + opened.close(); + } catch (Exception closeError) { + error.addSuppressed(closeError); + throw AgentRpcError.resource( + "close", + new JdbcConnectionPoolRegistry.PhysicalConnectionStateUnknownException(error) + ); } + throw error; } } @@ -509,12 +555,24 @@ public abstract class AbstractJdbcAgent extends BaseDatabaseAgent { if (params == null || poolIdentity == null || poolRegistry == null) { throw new IllegalStateException("Not connected"); } - return poolRegistry.borrow(poolIdentity, () -> openInitializedConnection(params)); + return borrowPooledConnection(poolRegistry, poolIdentity, params); + } + + private JdbcConnectionPoolRegistry.Lease borrowPooledConnection( + JdbcConnectionPoolRegistry registry, + String identity, + ConnectParams params + ) throws Exception { + return registry.borrow( + identity, + JdbcSessionRole.from(params.getSessionRole()), + () -> openInitializedConnection(params) + ); } private void closeCurrentConnection() throws Exception { if (pooledLease != null) { - boolean evict = sessionAffinity || !preparePooledConnectionForReturn(); + boolean evict = pooledConnectionPoisoned || sessionAffinity || !preparePooledConnectionForReturn(); releasePooledConnection(evict); } else if (connection != null) { connection.close(); diff --git a/agents/common/src/main/java/com/dbx/agent/AgentRpcError.java b/agents/common/src/main/java/com/dbx/agent/AgentRpcError.java new file mode 100644 index 000000000..bb356e1b2 --- /dev/null +++ b/agents/common/src/main/java/com/dbx/agent/AgentRpcError.java @@ -0,0 +1,135 @@ +package com.dbx.agent; + +import com.google.gson.JsonObject; + +import java.sql.SQLException; +import java.sql.SQLRecoverableException; +import java.sql.SQLTransientConnectionException; +import java.util.Locale; + +final class AgentRpcError extends RuntimeException { + private final String category; + private final boolean retryable; + private final String disposition; + private final String stage; + + private AgentRpcError( + String message, + String category, + boolean retryable, + String disposition, + String stage, + Throwable cause + ) { + super(message, cause); + this.category = category; + this.retryable = retryable; + this.disposition = disposition; + this.stage = stage; + } + + static AgentRpcError resource(String stage, Throwable cause) { + return new AgentRpcError( + "Agent runtime resource limit reached", + "resource", + false, + "replace_runtime", + stage, + cause + ); + } + + static AgentRpcError backpressure(String stage, Throwable cause) { + return new AgentRpcError( + "Agent request capacity is temporarily exhausted", + "resource", + true, + "keep", + stage, + cause + ); + } + + static JsonObject toJson(Throwable error, String method, String agentSessionId) { + AgentRpcError classified = classify(error, method); + JsonObject rpcError = new JsonObject(); + rpcError.addProperty("code", -1); + rpcError.addProperty("message", message(error)); + JsonObject data = new JsonObject(); + data.addProperty("category", classified.category); + data.addProperty("retryable", classified.retryable); + data.addProperty("sessionDisposition", classified.disposition); + data.addProperty("stage", classified.stage); + if (agentSessionId != null && !agentSessionId.trim().isEmpty()) { + data.addProperty("agentSessionId", agentSessionId); + } + rpcError.add("data", data); + return rpcError; + } + + private static AgentRpcError classify(Throwable error, String method) { + AgentRpcError explicit = find(error, AgentRpcError.class); + if (explicit != null) { + return explicit; + } + SQLException sqlError = find(error, SQLException.class); + if (sqlError != null) { + String sqlState = sqlError.getSQLState(); + String stage = stage(method); + boolean connectionError = "connect".equals(stage) + || "validate".equals(stage) + || sqlError instanceof SQLRecoverableException + || sqlError instanceof SQLTransientConnectionException + || (sqlState != null && sqlState.toUpperCase(Locale.ROOT).startsWith("08")); + boolean operationRetryable = connectionError && ("connect".equals(stage) || "validate".equals(stage)); + String disposition = connectionError && !"connect".equals(stage) ? "quarantine" : "keep"; + return new AgentRpcError( + message(error), + connectionError ? "connection" : "sql", + operationRetryable, + disposition, + stage, + error + ); + } + return new AgentRpcError(message(error), "protocol", false, "keep", stage(method), error); + } + + private static String stage(String method) { + if (method == null) { + return "request"; + } + if (AgentProtocol.METHOD_CONNECT.equals(method) || AgentProtocol.METHOD_OPEN_SESSION.equals(method)) { + return "connect"; + } + if (AgentProtocol.METHOD_VALIDATE_CONNECTION.equals(method) || AgentProtocol.METHOD_VALIDATE_SESSION.equals(method)) { + return "validate"; + } + if (AgentProtocol.METHOD_CANCEL_SESSION.equals(method)) { + return "cancel"; + } + if (AgentProtocol.METHOD_CLOSE_SESSION.equals(method) || AgentProtocol.METHOD_DISCONNECT.equals(method)) { + return "close"; + } + if (AgentProtocol.METHOD_FETCH_QUERY_PAGE.equals(method) + || AgentProtocol.METHOD_FETCH_TABLE_READ_PAGE.equals(method)) { + return "fetch"; + } + return "execute"; + } + + private static String message(Throwable error) { + return error.getMessage() == null ? error.toString() : error.getMessage(); + } + + private static T find(Throwable error, Class type) { + Throwable current = error; + while (current != null) { + if (type.isInstance(current)) { + return type.cast(current); + } + current = current.getCause(); + } + return null; + } +} diff --git a/agents/common/src/main/java/com/dbx/agent/ConnectParams.java b/agents/common/src/main/java/com/dbx/agent/ConnectParams.java index 1e4c911b8..2ba4994c6 100644 --- a/agents/common/src/main/java/com/dbx/agent/ConnectParams.java +++ b/agents/common/src/main/java/com/dbx/agent/ConnectParams.java @@ -22,6 +22,7 @@ public final class ConnectParams { private String client_key_path; private String gbase_server; private String informix_server; + private String sessionRole; public ConnectParams() { this("", 0, "", "", "", "", "", false, "", Collections.emptyList()); @@ -200,6 +201,14 @@ public final class ConnectParams { this.informix_server = informix_server; } + public String getSessionRole() { + return sessionRole; + } + + public void setSessionRole(String sessionRole) { + this.sessionRole = sessionRole; + } + @Override public boolean equals(Object other) { if (this == other) return true; @@ -221,13 +230,14 @@ public final class ConnectParams { && Objects.equals(client_cert_path, that.client_cert_path) && Objects.equals(client_key_path, that.client_key_path) && Objects.equals(gbase_server, that.gbase_server) - && Objects.equals(informix_server, that.informix_server); + && Objects.equals(informix_server, that.informix_server) + && Objects.equals(sessionRole, that.sessionRole); } @Override public int hashCode() { return Objects.hash(host, port, database, username, password, url_params, connection_string, - port_explicit, mysql_compat_mode, jdbc_driver_class, jdbc_driver_paths, ssl, ca_cert_path, client_cert_path, client_key_path, gbase_server, informix_server); + port_explicit, mysql_compat_mode, jdbc_driver_class, jdbc_driver_paths, ssl, ca_cert_path, client_cert_path, client_key_path, gbase_server, informix_server, sessionRole); } @Override @@ -249,6 +259,7 @@ public final class ConnectParams { + ", client_key_path=" + client_key_path + ", gbase_server=" + gbase_server + ", informix_server=" + informix_server + + ", sessionRole=" + sessionRole + ")"; } } diff --git a/agents/common/src/main/java/com/dbx/agent/JdbcConnectionPoolRegistry.java b/agents/common/src/main/java/com/dbx/agent/JdbcConnectionPoolRegistry.java index 885e34abe..bf64bc995 100644 --- a/agents/common/src/main/java/com/dbx/agent/JdbcConnectionPoolRegistry.java +++ b/agents/common/src/main/java/com/dbx/agent/JdbcConnectionPoolRegistry.java @@ -5,18 +5,35 @@ import com.zaxxer.hikari.HikariDataSource; import javax.sql.DataSource; import java.io.PrintWriter; +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Proxy; import java.nio.charset.StandardCharsets; import java.security.MessageDigest; import java.sql.Connection; import java.sql.SQLException; import java.sql.SQLFeatureNotSupportedException; +import java.sql.SQLTransientConnectionException; import java.util.ArrayList; +import java.util.Collections; +import java.util.IdentityHashMap; import java.util.List; import java.util.Map; import java.util.Objects; +import java.util.Set; +import java.util.concurrent.ArrayBlockingQueue; +import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.Executor; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.RejectedExecutionException; +import java.util.concurrent.Semaphore; import java.util.concurrent.ThreadFactory; +import java.util.concurrent.ThreadPoolExecutor; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicReference; import java.util.logging.Logger; @@ -28,8 +45,19 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { private static final long DEFAULT_IDLE_TIMEOUT_MILLIS = 120_000L; private static final long DEFAULT_MAX_LIFETIME_MILLIS = 1_800_000L; private static final long DEFAULT_POOL_RETIRE_MILLIS = 300_000L; + private static final long CHECKOUT_WATCHDOG_GRACE_MILLIS = 50L; + private static final int DEFAULT_METADATA_RESERVE = 2; + private static final int DEFAULT_GLOBAL_MAXIMUM_PHYSICAL_CONNECTIONS = 32; + private static final int DEFAULT_MAX_QUARANTINED_OPERATIONS = 2; private final Map pools = new ConcurrentHashMap<>(); private final PoolSettings settings; + private final PhysicalConnectionBudget physicalConnectionBudget; + private final PhysicalConnectionOpener physicalConnectionOpener; + private final PhysicalConnectionCloser physicalConnectionCloser; + private final ConnectionReleaseExecutor connectionReleaseExecutor; + private final PoolCloseExecutor poolCloseExecutor; + private final JdbcCheckoutExecutor checkoutExecutor; + private final AtomicReference runtimeFailure = new AtomicReference<>(); private final AtomicBoolean closed = new AtomicBoolean(); JdbcConnectionPoolRegistry() { @@ -38,14 +66,34 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { JdbcConnectionPoolRegistry(PoolSettings settings) { this.settings = Objects.requireNonNull(settings, "settings"); + this.physicalConnectionBudget = new PhysicalConnectionBudget(settings.globalMaximumPhysicalConnections); + this.physicalConnectionOpener = new PhysicalConnectionOpener(settings.globalMaximumPhysicalConnections); + this.physicalConnectionCloser = new PhysicalConnectionCloser(settings.globalMaximumPhysicalConnections); + this.connectionReleaseExecutor = new ConnectionReleaseExecutor(settings.globalMaximumPhysicalConnections); + this.poolCloseExecutor = new PoolCloseExecutor( + settings.globalMaximumPhysicalConnections, + runtimeFailure + ); + this.checkoutExecutor = new JdbcCheckoutExecutor( + settings.globalMaximumPhysicalConnections, + connectionReleaseExecutor + ); } Lease borrow(String identity, ConnectionFactory connectionFactory) throws Exception { + return borrow(identity, JdbcSessionRole.WORKLOAD, connectionFactory); + } + + Lease borrow(String identity, JdbcSessionRole role, ConnectionFactory connectionFactory) throws Exception { String key = digest(identity); while (true) { if (closed.get()) { throw new IllegalStateException("JDBC connection pool registry is closed"); } + SQLException failure = runtimeFailure.get(); + if (failure != null) { + throw AgentRpcError.resource("close", failure); + } PoolEntry entry; try { entry = pools.computeIfAbsent(key, ignored -> createPoolEntry(key, connectionFactory)); @@ -53,9 +101,14 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { throw error.unwrap(); } try { - return entry.borrow(); + return entry.borrow(role); } catch (PoolRetiredException ignored) { pools.remove(key, entry); + } catch (AgentRpcError error) { + if (entry.isRetired()) { + pools.remove(key, entry); + } + throw error; } } } @@ -68,14 +121,19 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { return settings.enabled; } + int activePhysicalConnectionCount() { + return physicalConnectionBudget.activeCount(); + } + private PoolEntry createPoolEntry(String key, ConnectionFactory connectionFactory) { - ConnectionFactoryDataSource factoryDataSource = null; try { - Connection initialConnection = Objects.requireNonNull( - connectionFactory.open(), - "JDBC connection factory returned null" + ConnectionFactoryDataSource factoryDataSource = new ConnectionFactoryDataSource( + connectionFactory, + physicalConnectionBudget, + physicalConnectionOpener, + physicalConnectionCloser, + settings.connectionTimeoutMillis ); - factoryDataSource = new ConnectionFactoryDataSource(connectionFactory, initialConnection); HikariConfig config = new HikariConfig(); config.setPoolName("dbx-jdbc-" + key.substring(0, 12)); config.setDataSource(factoryDataSource); @@ -88,11 +146,19 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { config.setInitializationFailTimeout(-1L); config.setIsolateInternalQueries(true); config.setThreadFactory(daemonThreadFactory("dbx-jdbc-pool-" + key.substring(0, 8))); - return new PoolEntry(new HikariDataSource(config), factoryDataSource, settings.poolRetireMillis); + return new PoolEntry( + new HikariDataSource(config), + settings.poolRetireMillis, + settings.connectionTimeoutMillis, + settings.maximumPoolSize, + settings.metadataReserve, + settings.maxQuarantinedOperations, + factoryDataSource, + checkoutExecutor, + connectionReleaseExecutor, + poolCloseExecutor + ); } catch (Exception error) { - if (factoryDataSource != null) { - factoryDataSource.closeUnusedInitialConnection(); - } throw new PoolCreationException(error); } } @@ -126,6 +192,11 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { entry.close(); } pools.clear(); + checkoutExecutor.close(); + connectionReleaseExecutor.close(); + poolCloseExecutor.close(); + physicalConnectionOpener.close(); + physicalConnectionCloser.close(); } private static String digest(String identity) { @@ -150,35 +221,78 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { }; } + private static ExecutorService boundedExecutor(int maximumThreads, String name) { + int threads = Math.max(1, maximumThreads); + return new ThreadPoolExecutor( + threads, + threads, + 0L, + TimeUnit.MILLISECONDS, + new ArrayBlockingQueue<>(threads), + daemonThreadFactory(name), + new ThreadPoolExecutor.AbortPolicy() + ); + } + + private static long hikariCheckoutTimeoutMillis(long physicalOperationTimeoutMillis) { + return addTimeoutMargin(physicalOperationTimeoutMillis, CHECKOUT_WATCHDOG_GRACE_MILLIS); + } + + private static long addTimeoutMargin(long timeoutMillis, long marginMillis) { + return timeoutMillis > Long.MAX_VALUE - marginMillis ? Long.MAX_VALUE : timeoutMillis + marginMillis; + } + @FunctionalInterface interface ConnectionFactory { Connection open() throws Exception; } + private static final class OperationDeadline { + private final long deadlineNanos; + + private OperationDeadline(long timeoutMillis) { + long timeoutNanos = TimeUnit.MILLISECONDS.toNanos(Math.max(1L, timeoutMillis)); + long now = System.nanoTime(); + this.deadlineNanos = now > Long.MAX_VALUE - timeoutNanos ? Long.MAX_VALUE : now + timeoutNanos; + } + + private long remainingNanos() { + return Math.max(0L, deadlineNanos - System.nanoTime()); + } + } + static final class Lease implements AutoCloseable { private final PoolEntry entry; private final Connection connection; + private final boolean workloadPermit; private final AtomicBoolean closed = new AtomicBoolean(); + private final AtomicBoolean quarantined = new AtomicBoolean(); - private Lease(PoolEntry entry, Connection connection) { + private Lease(PoolEntry entry, Connection connection, boolean workloadPermit) { this.entry = entry; this.connection = connection; + this.workloadPermit = workloadPermit; } Connection connection() { return connection; } - void evict() { + synchronized boolean quarantine() { + return !closed.get() && quarantined.compareAndSet(false, true) && entry.markQuarantined(); + } + + synchronized void evict() { if (closed.compareAndSet(false, true)) { - entry.release(connection, true); + entry.release(connection, true, workloadPermit, quarantined.get()); } } @Override - public void close() { + public synchronized void close() { if (closed.compareAndSet(false, true)) { - entry.release(connection, false); + boolean poisoned = quarantined.get(); + entry.release(connection, poisoned, workloadPermit, poisoned); } } } @@ -192,6 +306,9 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { private final long idleTimeoutMillis; private final long maxLifetimeMillis; private final long poolRetireMillis; + private final int metadataReserve; + private final int globalMaximumPhysicalConnections; + private final int maxQuarantinedOperations; PoolSettings( int maximumPoolSize, @@ -210,7 +327,10 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { validationTimeoutMillis, idleTimeoutMillis, maxLifetimeMillis, - poolRetireMillis + poolRetireMillis, + DEFAULT_METADATA_RESERVE, + DEFAULT_GLOBAL_MAXIMUM_PHYSICAL_CONNECTIONS, + DEFAULT_MAX_QUARANTINED_OPERATIONS ); } @@ -223,6 +343,34 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { long idleTimeoutMillis, long maxLifetimeMillis, long poolRetireMillis + ) { + this( + enabled, + maximumPoolSize, + minimumIdle, + connectionTimeoutMillis, + validationTimeoutMillis, + idleTimeoutMillis, + maxLifetimeMillis, + poolRetireMillis, + DEFAULT_METADATA_RESERVE, + DEFAULT_GLOBAL_MAXIMUM_PHYSICAL_CONNECTIONS, + DEFAULT_MAX_QUARANTINED_OPERATIONS + ); + } + + PoolSettings( + boolean enabled, + int maximumPoolSize, + int minimumIdle, + long connectionTimeoutMillis, + long validationTimeoutMillis, + long idleTimeoutMillis, + long maxLifetimeMillis, + long poolRetireMillis, + int metadataReserve, + int globalMaximumPhysicalConnections, + int maxQuarantinedOperations ) { this.enabled = enabled; this.maximumPoolSize = maximumPoolSize; @@ -232,6 +380,11 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { this.idleTimeoutMillis = idleTimeoutMillis; this.maxLifetimeMillis = maxLifetimeMillis; this.poolRetireMillis = poolRetireMillis; + this.metadataReserve = Math.max(0, Math.min(metadataReserve, maximumPoolSize - 1)); + this.globalMaximumPhysicalConnections = Math.max(1, globalMaximumPhysicalConnections); + this.maxQuarantinedOperations = maximumPoolSize == 1 + ? 1 + : Math.max(1, maxQuarantinedOperations); } static PoolSettings fromEnvironment() { @@ -291,10 +444,39 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { "DBX_AGENT_JDBC_POOL_RETIRE_MILLIS", DEFAULT_POOL_RETIRE_MILLIS, 60_000L + ), + intSetting( + "dbx.agent.jdbc.pool.metadataReserve", + "DBX_AGENT_JDBC_POOL_METADATA_RESERVE", + DEFAULT_METADATA_RESERVE, + 0, + Math.max(0, maximumPoolSize - 1) + ), + intSetting( + "dbx.agent.jdbc.pool.globalMaximumPhysicalConnections", + "DBX_AGENT_JDBC_POOL_GLOBAL_MAXIMUM_PHYSICAL_CONNECTIONS", + DEFAULT_GLOBAL_MAXIMUM_PHYSICAL_CONNECTIONS, + 1, + 256 + ), + intSetting( + "dbx.agent.jdbc.pool.maxQuarantinedOperations", + "DBX_AGENT_JDBC_POOL_MAX_QUARANTINED_OPERATIONS", + DEFAULT_MAX_QUARANTINED_OPERATIONS, + 1, + 64 ) ); } + int effectiveMetadataReserve() { + return metadataReserve; + } + + int effectiveMaxQuarantinedOperations() { + return maxQuarantinedOperations; + } + private static int intSetting(String property, String environment, int defaultValue, int minimum, int maximum) { String value = setting(property, environment); if (value == null) { @@ -342,53 +524,162 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { private static final class PoolEntry implements AutoCloseable { private final HikariDataSource dataSource; - private final ConnectionFactoryDataSource factoryDataSource; private final long retireMillis; + private final long connectionTimeoutMillis; + private final Semaphore workloadPermits; + private final Semaphore leasePermits; + private final int maxQuarantinedOperations; + private final ConnectionFactoryDataSource factoryDataSource; + private final JdbcCheckoutExecutor checkoutExecutor; + private final ConnectionReleaseExecutor connectionReleaseExecutor; + private final PoolCloseExecutor poolCloseExecutor; private int activeLeases; + private int quarantinedLeases; private boolean retired; private boolean dataSourceClosed; private volatile long lastReleasedAtMillis = System.currentTimeMillis(); private PoolEntry( HikariDataSource dataSource, + long retireMillis, + long connectionTimeoutMillis, + int maximumPoolSize, + int metadataReserve, + int maxQuarantinedOperations, ConnectionFactoryDataSource factoryDataSource, - long retireMillis + JdbcCheckoutExecutor checkoutExecutor, + ConnectionReleaseExecutor connectionReleaseExecutor, + PoolCloseExecutor poolCloseExecutor ) { this.dataSource = dataSource; - this.factoryDataSource = factoryDataSource; this.retireMillis = retireMillis; + this.connectionTimeoutMillis = connectionTimeoutMillis; + this.workloadPermits = new Semaphore(maximumPoolSize - metadataReserve, true); + this.leasePermits = new Semaphore(maximumPoolSize, true); + this.maxQuarantinedOperations = maxQuarantinedOperations; + this.factoryDataSource = factoryDataSource; + this.checkoutExecutor = checkoutExecutor; + this.connectionReleaseExecutor = connectionReleaseExecutor; + this.poolCloseExecutor = poolCloseExecutor; } - private Lease borrow() throws SQLException { + private Lease borrow(JdbcSessionRole role) throws SQLException { + OperationDeadline deadline = new OperationDeadline(hikariCheckoutTimeoutMillis(connectionTimeoutMillis)); + SQLException causalFailure = factoryDataSource.causalFailure(); + if (causalFailure != null) { + if (requiresRuntimeReplacement(causalFailure)) { + retireAfterCheckoutFailure(deadline); + } + throw AgentRpcError.resource("connect", causalFailure); + } + boolean workloadPermit = false; + boolean leasePermit = false; + if (role == JdbcSessionRole.WORKLOAD) { + workloadPermit = acquirePermit( + workloadPermits, + deadline, + "JDBC workload lease capacity is exhausted" + ); + } + try { + leasePermit = acquirePermit( + leasePermits, + deadline, + "JDBC connection pool lease capacity is exhausted" + ); + } catch (SQLException | RuntimeException error) { + if (leasePermit) { + leasePermits.release(); + } + if (workloadPermit) { + workloadPermits.release(); + } + throw error; + } synchronized (this) { if (retired) { + leasePermits.release(); + if (workloadPermit) { + workloadPermits.release(); + } throw new PoolRetiredException(); } activeLeases += 1; } try { - return new Lease(this, dataSource.getConnection()); + Lease lease = new Lease( + this, + checkoutExecutor.checkout(dataSource, factoryDataSource, deadline), + workloadPermit + ); + return lease; } catch (SQLException | RuntimeException error) { synchronized (this) { activeLeases -= 1; lastReleasedAtMillis = System.currentTimeMillis(); } + leasePermits.release(); + if (workloadPermit) { + workloadPermits.release(); + } + causalFailure = factoryDataSource.causalFailure(); + if (contains(error, PhysicalConnectionLimitException.class) + || contains(causalFailure, PhysicalConnectionLimitException.class)) { + throw AgentRpcError.backpressure("connect", causalFailure == null ? error : causalFailure); + } + if (causalFailure != null || contains(error, PhysicalConnectionStateUnknownException.class)) { + retireAfterCheckoutFailure(deadline); + throw AgentRpcError.resource("connect", causalFailure == null ? error : causalFailure); + } throw error; } } - private void release(Connection connection, boolean evict) { + private boolean acquirePermit( + Semaphore permits, + OperationDeadline deadline, + String exhaustedMessage + ) throws SQLException { try { - if (evict) { - dataSource.evictConnection(connection); + if (permits.tryAcquire(deadline.remainingNanos(), TimeUnit.NANOSECONDS)) { + return true; } - connection.close(); - } catch (Exception ignored) { + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + throw new SQLException("Interrupted while waiting for a JDBC workload lease", error); + } + throw AgentRpcError.backpressure( + "checkout", + new SQLTransientConnectionException(exhaustedMessage) + ); + } + + private synchronized boolean markQuarantined() { + quarantinedLeases += 1; + return quarantinedLeases >= maxQuarantinedOperations; + } + + private void release(Connection connection, boolean evict, boolean workloadPermit, boolean quarantined) { + try { + connectionReleaseExecutor.release( + dataSource, + connection, + evict, + factoryDataSource, + connectionTimeoutMillis + ); } finally { synchronized (this) { lastReleasedAtMillis = System.currentTimeMillis(); activeLeases -= 1; + if (quarantined) { + quarantinedLeases -= 1; + } } + if (workloadPermit) { + workloadPermits.release(); + } + leasePermits.release(); } } @@ -400,15 +691,30 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { return true; } + private synchronized boolean isRetired() { + return retired; + } + + private void retireAfterCheckoutFailure(OperationDeadline deadline) { + synchronized (this) { + retired = true; + } + closeRetired(deadline); + } + private void closeRetired() { + closeRetired(new OperationDeadline(hikariCheckoutTimeoutMillis(connectionTimeoutMillis))); + } + + private void closeRetired(OperationDeadline deadline) { synchronized (this) { if (dataSourceClosed) { return; } dataSourceClosed = true; } - factoryDataSource.closeUnusedInitialConnection(); - dataSource.close(); + factoryDataSource.retire(); + poolCloseExecutor.close(dataSource, factoryDataSource, deadline); } @Override @@ -426,38 +732,350 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { } } + private static final class PhysicalConnectionAttemptCanceledException extends SQLTransientConnectionException { + private PhysicalConnectionAttemptCanceledException() { + super("JDBC physical connection attempt was canceled"); + } + } + private static final class ConnectionFactoryDataSource implements DataSource { private final ConnectionFactory connectionFactory; - private final AtomicReference initialConnection; + private final PhysicalConnectionBudget physicalConnectionBudget; + private final PhysicalConnectionOpener physicalConnectionOpener; + private final PhysicalConnectionCloser physicalConnectionCloser; + private final long connectionTimeoutMillis; + private final AtomicReference causalFailure = new AtomicReference<>(); + private final Set checkoutDeadlines = ConcurrentHashMap.newKeySet(); + private final AtomicBoolean retired = new AtomicBoolean(); + private final AtomicInteger physicalCapacityWaiters = new AtomicInteger(); + private final Object attemptMonitor = new Object(); + private long latestStartedAttempt; + private long latestCompletedAttempt; + private AttemptDisposition latestAttemptDisposition; + private SQLException latestAttemptFailure; - private ConnectionFactoryDataSource(ConnectionFactory connectionFactory, Connection initialConnection) { + private ConnectionFactoryDataSource( + ConnectionFactory connectionFactory, + PhysicalConnectionBudget physicalConnectionBudget, + PhysicalConnectionOpener physicalConnectionOpener, + PhysicalConnectionCloser physicalConnectionCloser, + long connectionTimeoutMillis + ) { this.connectionFactory = connectionFactory; - this.initialConnection = new AtomicReference<>(initialConnection); + this.physicalConnectionBudget = physicalConnectionBudget; + this.physicalConnectionOpener = physicalConnectionOpener; + this.physicalConnectionCloser = physicalConnectionCloser; + this.connectionTimeoutMillis = connectionTimeoutMillis; } @Override public Connection getConnection() throws SQLException { - Connection opened = initialConnection.getAndSet(null); - if (opened != null) { - return opened; + SQLException existingFailure = causalFailure.get(); + if (existingFailure != null) { + throw existingFailure; } + OperationDeadline deadline = currentCheckoutDeadline(); + if (retired.get()) { + throw new PoolRetiredException(); + } + AttemptRegistration attempt = beginAttempt(); + physicalCapacityWaiters.incrementAndGet(); try { - return connectionFactory.open(); + try { + physicalConnectionBudget.acquire(deadline); + } catch (PhysicalConnectionLimitException error) { + attempt.complete(AttemptDisposition.CAPACITY, error); + throw error; + } catch (SQLException error) { + attempt.complete(AttemptDisposition.CAPACITY, error); + throw error; + } + } finally { + physicalCapacityWaiters.decrementAndGet(); + } + if (retired.get()) { + physicalConnectionBudget.release(); + attempt.complete(AttemptDisposition.CANCELED, null); + throw new PhysicalConnectionAttemptCanceledException(); + } + boolean opened = false; + boolean releaseBudget = true; + Connection connection = null; + try { + connection = Objects.requireNonNull( + physicalConnectionOpener.open( + connectionFactory, + physicalConnectionBudget, + physicalConnectionCloser, + this, + deadline, + connectionTimeoutMillis + ), + "JDBC connection factory returned null" + ); + HikariSetupAttempt setupAttempt = new HikariSetupAttempt(attempt); + Connection wrapped = physicalConnectionBudget.wrap( + connection, + physicalConnectionCloser, + this, + connectionTimeoutMillis, + setupAttempt + ); + opened = true; + return wrapped; } catch (SQLException error) { + releaseBudget = !contains(error, PhysicalConnectionStateUnknownException.class); + if (!releaseBudget) { + causalFailure.compareAndSet(null, find(error, PhysicalConnectionStateUnknownException.class)); + attempt.complete(AttemptDisposition.UNKNOWN, error); + } else if (contains(error, JdbcOperationCapacityException.class)) { + causalFailure.compareAndSet(null, find(error, JdbcOperationCapacityException.class)); + attempt.complete(AttemptDisposition.UNKNOWN, error); + } else if (contains(error, PhysicalConnectionLimitException.class)) { + attempt.complete(AttemptDisposition.CAPACITY, error); + } else { + attempt.complete(AttemptDisposition.KNOWN_FAILURE, error); + } throw error; } catch (Exception error) { - throw new SQLException("Failed to open JDBC connection", error); + SQLException failure = new SQLException("Failed to open JDBC connection", error); + releaseBudget = !contains(error, PhysicalConnectionStateUnknownException.class); + if (!releaseBudget) { + causalFailure.compareAndSet(null, find(error, PhysicalConnectionStateUnknownException.class)); + attempt.complete(AttemptDisposition.UNKNOWN, failure); + } else { + attempt.complete(AttemptDisposition.KNOWN_FAILURE, failure); + } + throw failure; + } finally { + if (!opened) { + if (connection != null) { + boolean connectionClosed = physicalConnectionCloser.close( + connection, + this, + connectionTimeoutMillis + ); + releaseBudget = releaseBudget && connectionClosed; + } + if (releaseBudget) { + physicalConnectionBudget.release(); + } + } } } - private void closeUnusedInitialConnection() { - Connection opened = initialConnection.getAndSet(null); - if (opened == null) { - return; + private SQLException causalFailure() { + return causalFailure.get(); + } + + private void completeHikariCheckout(Connection connection) throws SQLException { + if (connection.isWrapperFor(HikariSetupTrackedConnection.class)) { + connection.unwrap(HikariSetupTrackedConnection.class).completeHikariSetup(); } - try { - opened.close(); - } catch (Exception ignored) { + } + + private void retire() { + retired.set(true); + signalAttemptStateChanged(); + } + + private long latestCompletedAttemptGeneration() { + synchronized (attemptMonitor) { + return latestCompletedAttempt; + } + } + + private AttemptRegistration beginAttempt() { + synchronized (attemptMonitor) { + latestStartedAttempt += 1L; + return new AttemptRegistration(latestStartedAttempt); + } + } + + private AttemptSnapshot awaitAttemptAfter(long baseline, OperationDeadline deadline) { + synchronized (attemptMonitor) { + long targetGeneration = latestStartedAttempt; + if (targetGeneration <= baseline) { + return null; + } + while (latestCompletedAttempt < targetGeneration + && causalFailure.get() == null + && !retired.get()) { + long remainingNanos = deadline.remainingNanos(); + if (remainingNanos == 0L) { + break; + } + try { + TimeUnit.NANOSECONDS.timedWait(attemptMonitor, remainingNanos); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + break; + } + } + if (latestCompletedAttempt < targetGeneration) { + return null; + } + return new AttemptSnapshot(latestAttemptDisposition, latestAttemptFailure); + } + } + + private void signalAttemptStateChanged() { + synchronized (attemptMonitor) { + attemptMonitor.notifyAll(); + } + } + + private Throwable classifyCheckoutFailure( + SQLException error, + long attemptBaseline, + OperationDeadline deadline + ) { + SQLException failure = causalFailure.get(); + if (failure != null) { + return AgentRpcError.resource("checkout", failure); + } + if (retired.get()) { + return new PoolRetiredException(); + } + if (contains(error, PhysicalConnectionLimitException.class)) { + return AgentRpcError.backpressure("checkout", error); + } + if (contains(error, PhysicalConnectionStateUnknownException.class) + || contains(error, JdbcOperationCapacityException.class)) { + poison(error); + return AgentRpcError.resource("checkout", error); + } + if (physicalCapacityWaiters.get() > 0) { + return AgentRpcError.backpressure("checkout", error); + } + AttemptSnapshot attempt = awaitAttemptAfter(attemptBaseline, deadline); + failure = causalFailure.get(); + if (failure != null) { + return AgentRpcError.resource("checkout", failure); + } + if (retired.get()) { + return new PoolRetiredException(); + } + if (attempt != null) { + if (attempt.disposition == AttemptDisposition.KNOWN_FAILURE) { + return attempt.failure == null ? error : attempt.failure; + } + if (attempt.disposition == AttemptDisposition.CAPACITY + || attempt.disposition == AttemptDisposition.SUCCESS) { + return AgentRpcError.backpressure("checkout", attempt.failure == null ? error : attempt.failure); + } + } + SQLException unknown = new PhysicalConnectionStateUnknownException(error); + poison(unknown); + return AgentRpcError.resource("checkout", unknown); + } + + private final class HikariSetupAttempt { + private final AtomicBoolean completed = new AtomicBoolean(); + private final AtomicReference failure = new AtomicReference<>(); + private final AttemptRegistration attempt; + + private HikariSetupAttempt(AttemptRegistration attempt) { + this.attempt = attempt; + } + + private void recordFailure(SQLException error) { + failure.compareAndSet(null, error); + } + + private void completeSuccessfully() { + if (completed.compareAndSet(false, true)) { + attempt.complete(AttemptDisposition.SUCCESS, null); + } + } + + private void completeAfterPhysicalClose() { + if (completed.compareAndSet(false, true)) { + SQLException error = failure.get(); + attempt.complete( + AttemptDisposition.KNOWN_FAILURE, + error == null + ? new SQLException("Hikari rejected a physical connection during setup") + : error + ); + } + } + } + + private final class AttemptRegistration { + private final long generation; + private final AtomicBoolean completed = new AtomicBoolean(); + + private AttemptRegistration(long generation) { + this.generation = generation; + } + + private void complete(AttemptDisposition disposition, SQLException failure) { + if (!completed.compareAndSet(false, true)) { + return; + } + synchronized (attemptMonitor) { + if (generation >= latestCompletedAttempt) { + latestCompletedAttempt = generation; + latestAttemptDisposition = disposition; + latestAttemptFailure = failure; + } + attemptMonitor.notifyAll(); + } + } + } + + private enum AttemptDisposition { + SUCCESS, + KNOWN_FAILURE, + CAPACITY, + CANCELED, + UNKNOWN + } + + private static final class AttemptSnapshot { + private final AttemptDisposition disposition; + private final SQLException failure; + + private AttemptSnapshot(AttemptDisposition disposition, SQLException failure) { + this.disposition = disposition; + this.failure = failure; + } + } + + private void poison(SQLException failure) { + causalFailure.compareAndSet(null, failure); + signalAttemptStateChanged(); + } + + private DeadlineRegistration registerCheckout(OperationDeadline deadline) { + checkoutDeadlines.add(deadline); + return new DeadlineRegistration(deadline); + } + + private OperationDeadline currentCheckoutDeadline() throws SQLException { + OperationDeadline earliest = null; + for (OperationDeadline deadline : checkoutDeadlines) { + if (earliest == null || deadline.deadlineNanos < earliest.deadlineNanos) { + earliest = deadline; + } + } + if (earliest == null) { + throw new SQLTransientConnectionException("No active JDBC checkout owns physical connection creation"); + } + return earliest; + } + + private final class DeadlineRegistration implements AutoCloseable { + private final OperationDeadline deadline; + + private DeadlineRegistration(OperationDeadline deadline) { + this.deadline = deadline; + } + + @Override + public void close() { + checkoutDeadlines.remove(deadline); } } @@ -481,7 +1099,9 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { @Override public int getLoginTimeout() { - return 0; + // Hikari waits this long for in-flight add tasks before closing its connection bag. + long timeoutSeconds = (connectionTimeoutMillis + 999L) / 1_000L; + return (int) Math.min(Integer.MAX_VALUE, Math.max(1L, timeoutSeconds)); } @Override @@ -503,6 +1123,679 @@ final class JdbcConnectionPoolRegistry implements AutoCloseable { } } + private static final class JdbcCheckoutExecutor implements AutoCloseable { + private final ExecutorService executor; + private final ConnectionReleaseExecutor releaseExecutor; + + private JdbcCheckoutExecutor(int maximumConcurrentCheckouts, ConnectionReleaseExecutor releaseExecutor) { + this.executor = boundedExecutor(maximumConcurrentCheckouts, "dbx-jdbc-checkout"); + this.releaseExecutor = releaseExecutor; + } + + private Connection checkout( + HikariDataSource dataSource, + ConnectionFactoryDataSource factoryDataSource, + OperationDeadline deadline + ) throws SQLException { + CompletableFuture outcome = new CompletableFuture<>(); + long attemptBaseline = factoryDataSource.latestCompletedAttemptGeneration(); + try (ConnectionFactoryDataSource.DeadlineRegistration ignored = factoryDataSource.registerCheckout(deadline)) { + try { + executor.execute(() -> completeCheckout( + dataSource, + factoryDataSource, + attemptBaseline, + deadline, + outcome + )); + } catch (RejectedExecutionException error) { + JdbcOperationCapacityException failure = new JdbcOperationCapacityException("checkout", error); + factoryDataSource.poison(failure); + throw AgentRpcError.resource("checkout", failure); + } + try { + return outcome.get(deadline.remainingNanos(), TimeUnit.NANOSECONDS); + } catch (TimeoutException error) { + return abandonCheckout(outcome, factoryDataSource, error); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + return abandonCheckout(outcome, factoryDataSource, error); + } catch (ExecutionException error) { + throwCheckoutFailure(error.getCause()); + throw new IllegalStateException("unreachable"); + } + } + } + + private void completeCheckout( + HikariDataSource dataSource, + ConnectionFactoryDataSource factoryDataSource, + long attemptBaseline, + OperationDeadline deadline, + CompletableFuture outcome + ) { + try { + Connection connection = dataSource.getConnection(); + factoryDataSource.completeHikariCheckout(connection); + if (!outcome.complete(connection)) { + releaseExecutor.releaseLate( + dataSource, + connection, + factoryDataSource, + factoryDataSource.connectionTimeoutMillis + ); + } + } catch (SQLException error) { + outcome.completeExceptionally( + factoryDataSource.classifyCheckoutFailure(error, attemptBaseline, deadline) + ); + } catch (Throwable error) { + outcome.completeExceptionally(error); + } + } + + private static Connection abandonCheckout( + CompletableFuture outcome, + ConnectionFactoryDataSource factoryDataSource, + Exception cause + ) throws SQLException { + SQLException causalFailure = factoryDataSource.causalFailure(); + Throwable failure; + if (causalFailure == null) { + SQLException abandonedFailure = new PhysicalConnectionStateUnknownException(cause); + factoryDataSource.poison(abandonedFailure); + failure = AgentRpcError.resource("checkout", abandonedFailure); + } else { + failure = AgentRpcError.resource("checkout", causalFailure); + } + if (outcome.completeExceptionally(failure)) { + throwCheckoutFailure(failure); + throw new IllegalStateException("unreachable"); + } + try { + return outcome.get(); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + throwCheckoutFailure(failure); + throw new IllegalStateException("unreachable"); + } catch (ExecutionException error) { + throwCheckoutFailure(error.getCause()); + throw new IllegalStateException("unreachable"); + } + } + + private static void throwCheckoutFailure(Throwable error) throws SQLException { + if (error instanceof AgentRpcError rpcError) { + throw rpcError; + } + if (error instanceof SQLException sqlError) { + throw sqlError; + } + if (error instanceof RuntimeException runtimeError) { + throw runtimeError; + } + if (error instanceof Error fatal) { + throw fatal; + } + throw new SQLException("Failed to checkout JDBC connection", error); + } + + @Override + public void close() { + executor.shutdownNow(); + } + } + + private static final class ConnectionReleaseExecutor implements AutoCloseable { + private final ExecutorService executor; + + private ConnectionReleaseExecutor(int maximumConcurrentReleases) { + executor = boundedExecutor(maximumConcurrentReleases, "dbx-jdbc-release"); + } + + private void release( + HikariDataSource dataSource, + Connection connection, + boolean evict, + ConnectionFactoryDataSource factoryDataSource, + long timeoutMillis + ) { + release(dataSource, connection, evict, factoryDataSource, timeoutMillis, false); + } + + private void release( + HikariDataSource dataSource, + Connection connection, + boolean evict, + ConnectionFactoryDataSource factoryDataSource, + long timeoutMillis, + boolean attemptInlineWhenRejected + ) { + CompletableFuture outcome = submit( + dataSource, + connection, + evict, + factoryDataSource, + attemptInlineWhenRejected + ); + if (outcome == null) { + return; + } + OperationDeadline deadline = new OperationDeadline(timeoutMillis); + try { + outcome.get(deadline.remainingNanos(), TimeUnit.NANOSECONDS); + } catch (TimeoutException error) { + factoryDataSource.poison(new PhysicalConnectionStateUnknownException(error)); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + factoryDataSource.poison(new PhysicalConnectionStateUnknownException(error)); + } catch (ExecutionException error) { + factoryDataSource.poison(new PhysicalConnectionStateUnknownException(error.getCause())); + } + } + + private void releaseLate( + HikariDataSource dataSource, + Connection connection, + ConnectionFactoryDataSource factoryDataSource, + long timeoutMillis + ) { + release(dataSource, connection, true, factoryDataSource, timeoutMillis, true); + } + + private CompletableFuture submit( + HikariDataSource dataSource, + Connection connection, + boolean evict, + ConnectionFactoryDataSource factoryDataSource, + boolean attemptInlineWhenRejected + ) { + CompletableFuture outcome = new CompletableFuture<>(); + try { + executor.execute(() -> completeRelease(dataSource, connection, evict, outcome)); + return outcome; + } catch (RejectedExecutionException error) { + factoryDataSource.poison(new JdbcOperationCapacityException("release", error)); + if (attemptInlineWhenRejected) { + completeRelease(dataSource, connection, evict, outcome); + return outcome; + } + return null; + } + } + + private static void completeRelease( + HikariDataSource dataSource, + Connection connection, + boolean evict, + CompletableFuture outcome + ) { + Throwable failure = null; + try { + if (evict) { + dataSource.evictConnection(connection); + } + } catch (Throwable error) { + failure = error; + } + try { + connection.close(); + } catch (Throwable error) { + if (failure == null) { + failure = error; + } else { + failure.addSuppressed(error); + } + } + if (failure == null) { + outcome.complete(null); + } else { + outcome.completeExceptionally(failure); + } + } + + @Override + public void close() { + executor.shutdownNow(); + } + } + + private static final class PoolCloseExecutor implements AutoCloseable { + private final ExecutorService executor; + private final AtomicReference runtimeFailure; + + private PoolCloseExecutor(int maximumConcurrentCloses, AtomicReference runtimeFailure) { + this.executor = boundedExecutor(maximumConcurrentCloses, "dbx-jdbc-pool-close"); + this.runtimeFailure = runtimeFailure; + } + + private void close( + HikariDataSource dataSource, + ConnectionFactoryDataSource factoryDataSource, + OperationDeadline deadline + ) { + CompletableFuture outcome = new CompletableFuture<>(); + try { + executor.execute(() -> { + try { + dataSource.close(); + outcome.complete(null); + } catch (Throwable error) { + outcome.completeExceptionally(error); + } + }); + } catch (RejectedExecutionException error) { + poison(factoryDataSource, new JdbcOperationCapacityException("pool_close", error)); + return; + } + try { + outcome.get(deadline.remainingNanos(), TimeUnit.NANOSECONDS); + SQLException failure = factoryDataSource.causalFailure(); + if (failure != null && requiresRuntimeReplacement(failure)) { + runtimeFailure.compareAndSet(null, failure); + } + } catch (TimeoutException error) { + poison(factoryDataSource, new PhysicalConnectionStateUnknownException(error)); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + poison(factoryDataSource, new PhysicalConnectionStateUnknownException(error)); + } catch (ExecutionException error) { + poison(factoryDataSource, new PhysicalConnectionStateUnknownException(error.getCause())); + } + } + + private void poison(ConnectionFactoryDataSource factoryDataSource, SQLException failure) { + factoryDataSource.poison(failure); + runtimeFailure.compareAndSet(null, failure); + } + + @Override + public void close() { + executor.shutdownNow(); + } + } + + private static final class PhysicalConnectionCloser implements AutoCloseable { + private final ExecutorService executor; + + private PhysicalConnectionCloser(int maximumConcurrentCloses) { + executor = boundedExecutor(maximumConcurrentCloses, "dbx-jdbc-physical-close"); + } + + private boolean close( + Connection connection, + ConnectionFactoryDataSource factoryDataSource, + long timeoutMillis + ) { + try { + call("physical_close", () -> { + connection.close(); + return null; + }, factoryDataSource, timeoutMillis); + return true; + } catch (SQLException ignored) { + return false; + } + } + + private boolean abort( + Connection connection, + Executor abortExecutor, + ConnectionFactoryDataSource factoryDataSource, + long timeoutMillis + ) { + try { + call("physical_abort", () -> { + connection.abort(abortExecutor); + return null; + }, factoryDataSource, timeoutMillis); + factoryDataSource.poison(new PhysicalConnectionStateUnknownException( + new SQLException("JDBC connection abort cannot confirm physical resource release") + )); + } catch (SQLException ignored) { + } + return false; + } + + private boolean isClosed( + Connection connection, + ConnectionFactoryDataSource factoryDataSource, + long timeoutMillis + ) throws SQLException { + return call("physical_is_closed", connection::isClosed, factoryDataSource, timeoutMillis); + } + + private void setNetworkTimeout( + Connection connection, + Executor networkTimeoutExecutor, + int networkTimeoutMillis, + ConnectionFactoryDataSource factoryDataSource, + long timeoutMillis + ) throws SQLException { + call("physical_set_network_timeout", () -> { + connection.setNetworkTimeout(networkTimeoutExecutor, networkTimeoutMillis); + return null; + }, factoryDataSource, timeoutMillis); + } + + private T call( + String operation, + PhysicalConnectionCall call, + ConnectionFactoryDataSource factoryDataSource, + long timeoutMillis + ) throws SQLException { + CompletableFuture outcome = new CompletableFuture<>(); + try { + executor.execute(() -> { + try { + outcome.complete(call.run()); + } catch (Throwable error) { + outcome.completeExceptionally(error); + } + }); + } catch (RejectedExecutionException error) { + SQLException failure = new JdbcOperationCapacityException(operation, error); + factoryDataSource.poison(failure); + throw failure; + } + OperationDeadline deadline = new OperationDeadline(timeoutMillis); + try { + return outcome.get(deadline.remainingNanos(), TimeUnit.NANOSECONDS); + } catch (TimeoutException error) { + SQLException failure = new PhysicalConnectionStateUnknownException(error); + factoryDataSource.poison(failure); + throw failure; + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + SQLException failure = new PhysicalConnectionStateUnknownException(error); + factoryDataSource.poison(failure); + throw failure; + } catch (ExecutionException error) { + SQLException failure = new PhysicalConnectionStateUnknownException(error.getCause()); + factoryDataSource.poison(failure); + throw failure; + } + } + + @Override + public void close() { + executor.shutdownNow(); + } + } + + @FunctionalInterface + private interface PhysicalConnectionCall { + T run() throws SQLException; + } + + private interface HikariSetupTrackedConnection { + void completeHikariSetup(); + } + + private static final class PhysicalConnectionOpener implements AutoCloseable { + private final ExecutorService executor; + + private PhysicalConnectionOpener(int maximumConcurrentOpens) { + executor = boundedExecutor(maximumConcurrentOpens, "dbx-jdbc-physical-connect"); + } + + private Connection open( + ConnectionFactory connectionFactory, + PhysicalConnectionBudget physicalConnectionBudget, + PhysicalConnectionCloser physicalConnectionCloser, + ConnectionFactoryDataSource factoryDataSource, + OperationDeadline deadline, + long closeTimeoutMillis + ) throws Exception { + CompletableFuture outcome = new CompletableFuture<>(); + try { + executor.execute(() -> completeOpen( + connectionFactory, + physicalConnectionBudget, + physicalConnectionCloser, + factoryDataSource, + closeTimeoutMillis, + outcome + )); + } catch (RejectedExecutionException error) { + throw new JdbcOperationCapacityException("physical_connect", error); + } + try { + return outcome.get(deadline.remainingNanos(), TimeUnit.NANOSECONDS); + } catch (TimeoutException error) { + return abandon(outcome, error); + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + return abandon(outcome, error); + } catch (ExecutionException error) { + throwOpenFailure(error.getCause()); + throw new IllegalStateException("unreachable"); + } + } + + private static void completeOpen( + ConnectionFactory connectionFactory, + PhysicalConnectionBudget physicalConnectionBudget, + PhysicalConnectionCloser physicalConnectionCloser, + ConnectionFactoryDataSource factoryDataSource, + long closeTimeoutMillis, + CompletableFuture outcome + ) { + Connection connection = null; + try { + connection = Objects.requireNonNull( + connectionFactory.open(), + "JDBC connection factory returned null" + ); + if (!outcome.complete(connection) + && physicalConnectionCloser.close(connection, factoryDataSource, closeTimeoutMillis)) { + physicalConnectionBudget.release(); + } + } catch (Throwable error) { + if (!outcome.completeExceptionally(error) + && !contains(error, PhysicalConnectionStateUnknownException.class)) { + physicalConnectionBudget.release(); + } + } + } + + private static Connection abandon( + CompletableFuture outcome, + Exception cause + ) throws Exception { + PhysicalConnectionStateUnknownException failure = new PhysicalConnectionStateUnknownException(cause); + if (outcome.completeExceptionally(failure)) { + throw failure; + } + try { + return outcome.get(); + } catch (ExecutionException error) { + throwOpenFailure(error.getCause()); + throw new IllegalStateException("unreachable"); + } + } + + private static void throwOpenFailure(Throwable error) throws Exception { + if (error instanceof Error fatal) { + throw fatal; + } + if (error instanceof Exception exception) { + throw exception; + } + throw new SQLException("Failed to open JDBC connection", error); + } + + @Override + public void close() { + executor.shutdownNow(); + } + } + + private static final class PhysicalConnectionBudget { + private final int maximum; + private final Semaphore permits; + + private PhysicalConnectionBudget(int maximum) { + this.maximum = maximum; + this.permits = new Semaphore(maximum, true); + } + + private void acquire(OperationDeadline deadline) throws SQLException { + try { + if (permits.tryAcquire(deadline.remainingNanos(), TimeUnit.NANOSECONDS)) { + return; + } + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + throw new SQLException("Interrupted while waiting for the JDBC physical connection budget", error); + } + throw new PhysicalConnectionLimitException(maximum); + } + + private Connection wrap( + Connection connection, + PhysicalConnectionCloser physicalConnectionCloser, + ConnectionFactoryDataSource factoryDataSource, + long closeTimeoutMillis, + ConnectionFactoryDataSource.HikariSetupAttempt setupAttempt + ) { + AtomicBoolean released = new AtomicBoolean(); + return (Connection) Proxy.newProxyInstance( + JdbcConnectionPoolRegistry.class.getClassLoader(), + new Class[] {Connection.class, HikariSetupTrackedConnection.class}, + (proxy, method, arguments) -> { + if ("completeHikariSetup".equals(method.getName()) && method.getParameterCount() == 0) { + setupAttempt.completeSuccessfully(); + return null; + } + boolean close = "close".equals(method.getName()) && method.getParameterCount() == 0; + boolean abort = "abort".equals(method.getName()) && method.getParameterCount() == 1; + if (close || abort) { + synchronized (released) { + if (released.get()) { + return null; + } + boolean terminated = close + ? physicalConnectionCloser.close( + connection, + factoryDataSource, + closeTimeoutMillis + ) + : physicalConnectionCloser.abort( + connection, + (Executor) arguments[0], + factoryDataSource, + closeTimeoutMillis + ); + if (terminated) { + setupAttempt.completeAfterPhysicalClose(); + if (released.compareAndSet(false, true)) { + release(); + } + return null; + } + throw new PhysicalConnectionStateUnknownException( + new SQLException("Physical JDBC connection termination did not complete") + ); + } + } + if ("isClosed".equals(method.getName()) && method.getParameterCount() == 0) { + return released.get() || physicalConnectionCloser.isClosed( + connection, + factoryDataSource, + closeTimeoutMillis + ); + } + if ("setNetworkTimeout".equals(method.getName()) && method.getParameterCount() == 2) { + physicalConnectionCloser.setNetworkTimeout( + connection, + (Executor) arguments[0], + (Integer) arguments[1], + factoryDataSource, + closeTimeoutMillis + ); + return null; + } + try { + return method.invoke(connection, arguments); + } catch (InvocationTargetException error) { + Throwable cause = error.getCause(); + if (cause instanceof SQLException sqlError) { + setupAttempt.recordFailure(sqlError); + } + throw cause; + } + } + ); + } + + private void release() { + permits.release(); + } + + private int activeCount() { + return maximum - permits.availablePermits(); + } + + } + + static final class PhysicalConnectionStateUnknownException extends SQLException { + PhysicalConnectionStateUnknownException(Throwable cause) { + super("Physical JDBC connection state could not be confirmed", cause); + } + } + + private static final class JdbcOperationCapacityException extends SQLTransientConnectionException { + private JdbcOperationCapacityException(String operation, Throwable cause) { + super("JDBC " + operation + " executor capacity is exhausted", cause); + } + } + + private static final class PhysicalConnectionLimitException extends SQLTransientConnectionException { + private PhysicalConnectionLimitException(int maximum) { + super("Agent runtime JDBC physical connection limit reached: " + maximum); + } + } + + static boolean requiresRuntimeReplacement(Throwable error) { + return contains(error, PhysicalConnectionStateUnknownException.class); + } + + private static T find(Throwable error, Class type) { + java.util.Set visited = Collections.newSetFromMap(new IdentityHashMap<>()); + Throwable current = error; + while (current != null && visited.add(current)) { + if (type.isInstance(current)) { + return type.cast(current); + } + current = current.getCause(); + } + return null; + } + + private static boolean contains(Throwable error, Class type) { + return contains(error, type, Collections.newSetFromMap(new IdentityHashMap<>())); + } + + private static boolean contains( + Throwable error, + Class type, + java.util.Set visited + ) { + Throwable current = error; + while (current != null && visited.add(current)) { + if (type.isInstance(current)) { + return true; + } + if (current instanceof SQLException sqlError && sqlError.getNextException() != null + && contains(sqlError.getNextException(), type, visited)) { + return true; + } + current = current.getCause(); + } + return false; + } + private static final class PoolCreationException extends RuntimeException { private PoolCreationException(Exception cause) { super(cause); diff --git a/agents/common/src/main/java/com/dbx/agent/JdbcSessionRole.java b/agents/common/src/main/java/com/dbx/agent/JdbcSessionRole.java new file mode 100644 index 000000000..7a3417595 --- /dev/null +++ b/agents/common/src/main/java/com/dbx/agent/JdbcSessionRole.java @@ -0,0 +1,10 @@ +package com.dbx.agent; + +enum JdbcSessionRole { + WORKLOAD, + METADATA; + + static JdbcSessionRole from(String value) { + return "metadata".equalsIgnoreCase(value == null ? "" : value.trim()) ? METADATA : WORKLOAD; + } +} diff --git a/agents/common/src/main/java/com/dbx/agent/JsonRpcServer.java b/agents/common/src/main/java/com/dbx/agent/JsonRpcServer.java index 4b2aa0d91..806812c33 100644 --- a/agents/common/src/main/java/com/dbx/agent/JsonRpcServer.java +++ b/agents/common/src/main/java/com/dbx/agent/JsonRpcServer.java @@ -121,6 +121,11 @@ public final class JsonRpcServer { } } + boolean quarantine() { + AbstractJdbcAgent jdbcAgent = pooledJdbcAgent(); + return jdbcAgent != null && jdbcAgent.quarantinePooledConnection(); + } + void expireIdleResources() { jdbcExecutor.expireIdleResources(); releaseIdlePooledConnection(); diff --git a/agents/common/src/main/java/com/dbx/agent/MultiSessionJsonRpcServer.java b/agents/common/src/main/java/com/dbx/agent/MultiSessionJsonRpcServer.java index a8f236096..f9dd6fa4e 100644 --- a/agents/common/src/main/java/com/dbx/agent/MultiSessionJsonRpcServer.java +++ b/agents/common/src/main/java/com/dbx/agent/MultiSessionJsonRpcServer.java @@ -13,22 +13,30 @@ import java.util.Map; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; +import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.SynchronousQueue; +import java.util.concurrent.ThreadPoolExecutor; import java.util.concurrent.ThreadFactory; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; import java.util.concurrent.locks.ReentrantLock; +import java.util.function.Consumer; import java.util.function.Supplier; public final class MultiSessionJsonRpcServer implements AutoCloseable { private static final String LEGACY_SESSION_ID = "__legacy__"; private static final int MAX_SESSIONS = 256; + private static final int MAX_REQUEST_THREADS = 64; + private static final int MAX_CLEANUP_THREADS = 16; private static final long MAINTENANCE_INTERVAL_MILLIS = 60_000L; private final Supplier agentFactory; private final Supplier sessionHandlerFactory; private final Map sessions = new ConcurrentHashMap<>(); - private final ExecutorService requests = Executors.newCachedThreadPool(); + private final ExecutorService requests; + private final ExecutorService cleanup; private final JdbcConnectionPoolRegistry poolRegistry; private final Gson gson = new Gson(); private final PrintStream protocolOutput = System.out; @@ -38,29 +46,43 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { private ScheduledExecutorService maintenance; public MultiSessionJsonRpcServer(Supplier agentFactory) { - this(agentFactory, new JdbcConnectionPoolRegistry()); + this(agentFactory, new JdbcConnectionPoolRegistry(), RuntimeLimits.defaults()); } MultiSessionJsonRpcServer( Supplier agentFactory, JdbcConnectionPoolRegistry.PoolSettings poolSettings ) { - this(agentFactory, new JdbcConnectionPoolRegistry(poolSettings)); + this(agentFactory, new JdbcConnectionPoolRegistry(poolSettings), RuntimeLimits.defaults()); + } + + MultiSessionJsonRpcServer( + Supplier agentFactory, + JdbcConnectionPoolRegistry.PoolSettings poolSettings, + RuntimeLimits runtimeLimits + ) { + this(agentFactory, new JdbcConnectionPoolRegistry(poolSettings), runtimeLimits); } private MultiSessionJsonRpcServer( Supplier agentFactory, - JdbcConnectionPoolRegistry poolRegistry + JdbcConnectionPoolRegistry poolRegistry, + RuntimeLimits runtimeLimits ) { this.agentFactory = agentFactory; this.sessionHandlerFactory = null; this.poolRegistry = poolRegistry; + this.requests = boundedExecutor(runtimeLimits.maximumRequestThreads, "dbx-agent-request"); + this.cleanup = boundedExecutor(runtimeLimits.maximumCleanupThreads, "dbx-agent-cleanup"); } private MultiSessionJsonRpcServer(Supplier sessionHandlerFactory, boolean customHandler) { this.agentFactory = null; this.sessionHandlerFactory = sessionHandlerFactory; this.poolRegistry = new JdbcConnectionPoolRegistry(); + RuntimeLimits runtimeLimits = RuntimeLimits.defaults(); + this.requests = boundedExecutor(runtimeLimits.maximumRequestThreads, "dbx-agent-request"); + this.cleanup = boundedExecutor(runtimeLimits.maximumCleanupThreads, "dbx-agent-cleanup"); } /** Creates a protocol v2 server for a non-JDBC, session-scoped agent. */ @@ -82,7 +104,7 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { writeResponse(handleRequest(request)); return; } - requests.submit(() -> writeResponse(handleRequest(request))); + executeRequest(request, this::writeResponse); } } catch (Exception e) { throw new RuntimeException(e); @@ -95,6 +117,14 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { return gson.toJson(handleRequest(JsonParser.parseString(line).getAsJsonObject())); } + void executeRequest(JsonObject request, Consumer responseConsumer) { + try { + requests.execute(() -> responseConsumer.accept(handleRequest(request))); + } catch (RejectedExecutionException error) { + responseConsumer.accept(errorResponse(request.get("id"), AgentRpcError.backpressure("request", error))); + } + } + private JsonObject handleRequest(JsonObject request) { JsonElement id = request.get("id"); String method = request.get("method").getAsString(); @@ -134,10 +164,7 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { } response.add("result", gson.toJsonTree(result)); } catch (Throwable error) { - JsonObject rpcError = new JsonObject(); - rpcError.addProperty("code", -1); - rpcError.addProperty("message", error.getMessage() == null ? error.toString() : error.getMessage()); - response.add("error", rpcError); + response.add("error", AgentRpcError.toJson(error, method, stringOrNull(params, "agentSessionId"))); } return response; } @@ -165,7 +192,7 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { return session.connect(params); } catch (Exception error) { sessions.remove(sessionId, session); - session.close(); + session.quarantineAndClose(cleanup); throw error; } } @@ -173,7 +200,13 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { private Object closeSession(String sessionId) { Session session = sessions.remove(sessionId); if (session != null) { - session.close(); + boolean replaceRuntime = session.quarantineAndClose(cleanup); + if (replaceRuntime) { + throw AgentRpcError.resource( + "close", + new IllegalStateException("JDBC quarantine operation limit reached") + ); + } } return Collections.singletonMap("ok", true); } @@ -206,7 +239,11 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { private void closeAllSessions() { for (String sessionId : sessions.keySet()) { - closeSession(sessionId); + try { + closeSession(sessionId); + } catch (AgentRpcError ignored) { + // Sessions are already detached; process shutdown remains the final cleanup boundary. + } } } @@ -263,7 +300,16 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { } } closeAllSessions(); - requests.shutdown(); + requests.shutdownNow(); + cleanup.shutdown(); + try { + if (!cleanup.awaitTermination(2, TimeUnit.SECONDS)) { + cleanup.shutdownNow(); + } + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + cleanup.shutdownNow(); + } poolRegistry.close(); } @@ -275,6 +321,47 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { }; } + private static ExecutorService boundedExecutor(int maximumThreads, String threadName) { + return new ThreadPoolExecutor( + 0, + maximumThreads, + 60L, + TimeUnit.SECONDS, + new SynchronousQueue<>(), + daemonThreadFactory(threadName), + new ThreadPoolExecutor.AbortPolicy() + ); + } + + private static JsonObject errorResponse(JsonElement id, Throwable error) { + JsonObject response = new JsonObject(); + response.addProperty("jsonrpc", "2.0"); + response.add("id", id); + response.add("error", AgentRpcError.toJson(error, "request", null)); + return response; + } + + private static String stringOrNull(JsonObject params, String key) { + return params.has(key) && !params.get(key).isJsonNull() ? params.get(key).getAsString() : null; + } + + static final class RuntimeLimits { + private final int maximumRequestThreads; + private final int maximumCleanupThreads; + + RuntimeLimits(int maximumRequestThreads, int maximumCleanupThreads) { + if (maximumRequestThreads <= 0 || maximumCleanupThreads <= 0) { + throw new IllegalArgumentException("Agent runtime thread limits must be positive"); + } + this.maximumRequestThreads = maximumRequestThreads; + this.maximumCleanupThreads = maximumCleanupThreads; + } + + private static RuntimeLimits defaults() { + return new RuntimeLimits(MAX_REQUEST_THREADS, MAX_CLEANUP_THREADS); + } + } + private static String requiredSessionId(JsonObject params) { if (!params.has("agentSessionId") || params.get("agentSessionId").getAsString().trim().isEmpty()) { throw new IllegalArgumentException("agentSessionId is required"); @@ -293,6 +380,8 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { private final JsonRpcServer server; private final SessionRpcHandler handler; private final ReentrantLock lock = new ReentrantLock(); + private final AtomicReference state = new AtomicReference<>(State.ACTIVE); + private final AtomicBoolean cleanupScheduled = new AtomicBoolean(); private Session(JsonRpcServer server) { this.server = server; @@ -305,8 +394,10 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { } private Object handle(String method, JsonObject params) throws Exception { + requireActive(); lock.lock(); try { + requireActive(); return handler == null ? server.dispatchForRuntime(method, params) : handler.handle(method, params); } finally { lock.unlock(); @@ -314,8 +405,10 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { } private Object connect(JsonObject params) throws Exception { + requireActive(); lock.lock(); try { + requireActive(); return handler == null ? server.dispatchForRuntime(AgentProtocol.METHOD_CONNECT, params) : handler.connect(params); @@ -324,9 +417,31 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { } } + private boolean quarantineAndClose(ExecutorService cleanup) { + state.compareAndSet(State.ACTIVE, State.QUARANTINED); + boolean replaceRuntime = server != null && server.quarantine(); + if (!cleanupScheduled.compareAndSet(false, true)) { + return replaceRuntime; + } + try { + cleanup.execute(this::closeWhenIdle); + } catch (RejectedExecutionException error) { + throw AgentRpcError.resource("close", error); + } + return replaceRuntime; + } + + private void closeWhenIdle() { + close(); + } + private void close() { + state.compareAndSet(State.ACTIVE, State.QUARANTINED); lock.lock(); try { + if (state.get() == State.CLOSED) { + return; + } if (handler != null) { handler.close(); } else { @@ -334,6 +449,7 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { } } catch (Exception ignored) { } finally { + state.set(State.CLOSED); lock.unlock(); } } @@ -347,7 +463,7 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { } private void expireIdleResources() { - if (handler != null) { + if (state.get() != State.ACTIVE || handler != null) { return; } if (!lock.tryLock()) { @@ -361,7 +477,7 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { } private void expireIdleResources(long nowMillis, long idleTimeoutMillis) { - if (handler != null) { + if (state.get() != State.ACTIVE || handler != null) { return; } if (!lock.tryLock()) { @@ -373,5 +489,17 @@ public final class MultiSessionJsonRpcServer implements AutoCloseable { lock.unlock(); } } + + private void requireActive() { + if (state.get() != State.ACTIVE) { + throw new IllegalStateException("Agent session is quarantined"); + } + } + + private enum State { + ACTIVE, + QUARANTINED, + CLOSED + } } } diff --git a/agents/common/src/test/java/com/dbx/agent/CommonJavaCompatibilityTest.java b/agents/common/src/test/java/com/dbx/agent/CommonJavaCompatibilityTest.java index 469273a76..d0bdcd98c 100644 --- a/agents/common/src/test/java/com/dbx/agent/CommonJavaCompatibilityTest.java +++ b/agents/common/src/test/java/com/dbx/agent/CommonJavaCompatibilityTest.java @@ -151,7 +151,7 @@ class CommonJavaCompatibilityTest { } @Test - void multiSessionServerCreatesAndClosesIndependentAgents() { + void multiSessionServerCreatesAndClosesIndependentAgents() throws Exception { java.util.List created = new java.util.ArrayList<>(); MultiSessionJsonRpcServer server = new MultiSessionJsonRpcServer(() -> { TrackingAgent agent = new TrackingAgent(); @@ -172,6 +172,7 @@ class CommonJavaCompatibilityTest { assertEquals(1, created.get(1).connectCount); server.handleRequest("{\"jsonrpc\":\"2.0\",\"id\":4,\"method\":\"close_session\",\"params\":{\"agentSessionId\":\"a\"}}"); + awaitCondition(() -> created.get(0).disconnectCount == 1); assertEquals(1, created.get(0).disconnectCount); assertEquals(0, created.get(1).disconnectCount); } @@ -691,7 +692,7 @@ class CommonJavaCompatibilityTest { private static final class TrackingAgent extends MinimalAgent { private int connectCount; - private int disconnectCount; + private volatile int disconnectCount; @Override public void connect(ConnectParams params) { @@ -1078,6 +1079,14 @@ class CommonJavaCompatibilityTest { return false; } + private static void awaitCondition(java.util.function.BooleanSupplier condition) throws InterruptedException { + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(2); + while (!condition.getAsBoolean() && System.nanoTime() < deadline) { + Thread.sleep(10L); + } + assertTrue(condition.getAsBoolean()); + } + private static JsonObject protocolContract(String resourcePath) { InputStream stream = CommonJavaCompatibilityTest.class.getResourceAsStream(resourcePath); if (stream == null) { diff --git a/agents/common/src/test/java/com/dbx/agent/JdbcConnectionPoolingTest.java b/agents/common/src/test/java/com/dbx/agent/JdbcConnectionPoolingTest.java index 77d9b4f1c..a4d16761f 100644 --- a/agents/common/src/test/java/com/dbx/agent/JdbcConnectionPoolingTest.java +++ b/agents/common/src/test/java/com/dbx/agent/JdbcConnectionPoolingTest.java @@ -3,27 +3,36 @@ package com.dbx.agent; import com.google.gson.Gson; import com.google.gson.JsonObject; import com.google.gson.JsonParser; +import org.junit.jupiter.api.RepeatedTest; import org.junit.jupiter.api.Test; import java.sql.Connection; import java.sql.DriverManager; import java.sql.SQLException; +import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Proxy; import java.sql.Statement; import java.util.Collections; import java.util.List; import java.util.UUID; +import java.util.concurrent.BlockingQueue; import java.util.concurrent.CountDownLatch; +import java.util.concurrent.Executor; import java.util.concurrent.ExecutorService; import java.util.concurrent.Executors; +import java.util.concurrent.ExecutionException; import java.util.concurrent.Future; +import java.util.concurrent.LinkedBlockingQueue; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicBoolean; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; import static org.junit.jupiter.api.Assertions.assertNotNull; import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.fail; class JdbcConnectionPoolingTest { private static final Gson GSON = new Gson(); @@ -116,12 +125,12 @@ class JdbcConnectionPoolingTest { } @Test - void registrySurfacesInitialPhysicalConnectionFailureWithoutWaitingForBorrowTimeout() { + void registryBoundsInitialPhysicalConnectionFailureByBorrowTimeout() { JdbcConnectionPoolRegistry.PoolSettings settings = new JdbcConnectionPoolRegistry.PoolSettings( 1, 0, - 5_000L, - 1_000L, + 250L, + 250L, 10_000L, 30_000L, 60_000L @@ -135,11 +144,857 @@ class JdbcConnectionPoolingTest { }) ); long elapsedMillis = System.currentTimeMillis() - startedAtMillis; - assertTrue(elapsedMillis < 1_000L, () -> "initial failure took " + elapsedMillis + "ms"); - assertTrue(error.getMessage().contains("auth failed"), error::toString); + assertTrue(elapsedMillis < 450L, () -> "initial failure took " + elapsedMillis + "ms"); + assertTrue(throwableText(error).contains("auth failed"), error::toString); } } + @Test + void metadataLeaseUsesReservedCapacityWhenWorkloadIsSaturated() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + String url = h2Url("metadata_reserve"); + JdbcConnectionPoolRegistry.PoolSettings settings = poolSettings(4, 1, 32, 2); + ExecutorService worker = Executors.newSingleThreadExecutor(); + + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(settings)) { + JdbcConnectionPoolRegistry.Lease first = registry.borrow( + "same-identity", + JdbcSessionRole.WORKLOAD, + () -> openH2(url, physicalOpens) + ); + JdbcConnectionPoolRegistry.Lease second = registry.borrow( + "same-identity", + JdbcSessionRole.WORKLOAD, + () -> openH2(url, physicalOpens) + ); + JdbcConnectionPoolRegistry.Lease third = registry.borrow( + "same-identity", + JdbcSessionRole.WORKLOAD, + () -> openH2(url, physicalOpens) + ); + Future waitingWorkload = worker.submit(() -> { + try (JdbcConnectionPoolRegistry.Lease ignored = registry.borrow( + "same-identity", + JdbcSessionRole.WORKLOAD, + () -> openH2(url, physicalOpens) + )) { + return null; + } + }); + + Thread.sleep(100L); + assertFalse(waitingWorkload.isDone()); + try (JdbcConnectionPoolRegistry.Lease ignored = registry.borrow( + "same-identity", + JdbcSessionRole.METADATA, + () -> openH2(url, physicalOpens) + )) { + assertEquals(4, physicalOpens.get()); + } + + first.close(); + waitingWorkload.get(2, TimeUnit.SECONDS); + second.close(); + third.close(); + } finally { + worker.shutdownNow(); + } + } + + @Test + void metadataSessionCompletesThroughProtocolWhenWorkloadSessionsAreSaturated() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicInteger agentIndexes = new AtomicInteger(); + AtomicInteger requestIds = new AtomicInteger(); + CountDownLatch workloadsStarted = new CountDownLatch(3); + CountDownLatch releaseWorkloads = new CountDownLatch(1); + ExecutorService workers = Executors.newFixedThreadPool(3); + String url = h2Url("metadata_protocol_reserve"); + try (MultiSessionJsonRpcServer server = new MultiSessionJsonRpcServer( + () -> { + int agentIndex = agentIndexes.incrementAndGet(); + return new H2TestAgent(url, physicalOpens) { + @Override + public List listDatabases() { + if (agentIndex <= 3) { + workloadsStarted.countDown(); + awaitUninterruptibly(releaseWorkloads); + } + return Collections.emptyList(); + } + }; + }, + poolSettings(4, 1, 32, 2) + )) { + List> workloadRequests = new java.util.ArrayList<>(); + for (int index = 0; index < 3; index++) { + String sessionId = "workload-" + index; + openSession(server, requestIds, sessionId); + workloadRequests.add(workers.submit(() -> request( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams(sessionId) + ))); + } + assertTrue(workloadsStarted.await(2, TimeUnit.SECONDS)); + + JsonObject metadataParams = sessionParams("metadata"); + metadataParams.addProperty("sessionRole", "metadata"); + metadataParams.addProperty("database", "pooling"); + metadataParams.addProperty("username", "sa"); + metadataParams.addProperty("password", ""); + result(rawRequest(server, requestIds, AgentProtocol.METHOD_OPEN_SESSION, metadataParams)); + JsonObject metadataResponse = request( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("metadata") + ); + assertTrue(metadataResponse.getAsJsonArray("result").isEmpty()); + assertEquals(4, physicalOpens.get()); + + releaseWorkloads.countDown(); + for (Future workloadRequest : workloadRequests) { + workloadRequest.get(2, TimeUnit.SECONDS); + } + } finally { + releaseWorkloads.countDown(); + workers.shutdownNow(); + } + } + + @Test + void poolSizeOneDisablesMetadataReserveWithoutBlockingWorkload() throws Exception { + JdbcConnectionPoolRegistry.PoolSettings settings = poolSettings(1, 2, 32, 2); + assertEquals(0, settings.effectiveMetadataReserve()); + assertEquals(1, settings.effectiveMaxQuarantinedOperations()); + + AtomicInteger physicalOpens = new AtomicInteger(); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(settings); + JdbcConnectionPoolRegistry.Lease lease = registry.borrow( + "single-connection", + JdbcSessionRole.WORKLOAD, + () -> openH2(h2Url("single_connection_reserve"), physicalOpens) + )) { + assertEquals(1, physicalOpens.get()); + assertTrue(lease.quarantine()); + lease.evict(); + } + } + + @Test + void closingQuarantinedLeaseEvictsItInsteadOfReturningItToPool() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + String url = h2Url("quarantined_close"); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(poolSettings(2))) { + JdbcConnectionPoolRegistry.Lease quarantined = registry.borrow( + "same-identity", + () -> openH2(url, physicalOpens) + ); + assertFalse(quarantined.quarantine()); + quarantined.close(); + + try (JdbcConnectionPoolRegistry.Lease ignored = registry.borrow( + "same-identity", + () -> openH2(url, physicalOpens) + )) { + assertEquals(2, physicalOpens.get()); + } + } + } + + @Test + void registryCapsPhysicalConnectionsAcrossConnectionIdentities() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + JdbcConnectionPoolRegistry.PoolSettings settings = new JdbcConnectionPoolRegistry.PoolSettings( + true, + 4, + 0, + 250L, + 250L, + 10_000L, + 30_000L, + 0L, + 0, + 2, + 2 + ); + JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(settings); + JdbcConnectionPoolRegistry.Lease first = null; + JdbcConnectionPoolRegistry.Lease second = null; + try { + first = registry.borrow("identity-a", () -> openH2(h2Url("global_a"), physicalOpens)); + second = registry.borrow("identity-b", () -> openH2(h2Url("global_b"), physicalOpens)); + assertEquals(2, registry.activePhysicalConnectionCount()); + + AgentRpcError capacity = assertThrows( + AgentRpcError.class, + () -> registry.borrow("identity-c", () -> openH2(h2Url("global_c"), physicalOpens)) + ); + JsonObject capacityData = AgentRpcError.toJson( + capacity, + AgentProtocol.METHOD_OPEN_SESSION, + "capacity" + ).getAsJsonObject("data"); + assertEquals("keep", capacityData.get("sessionDisposition").getAsString()); + assertEquals(2, registry.activePhysicalConnectionCount()); + + first.close(); + first = null; + registry.retireUnusedPools(); + awaitPhysicalConnectionCount(registry, 1); + try (JdbcConnectionPoolRegistry.Lease ignored = registry.borrow( + "identity-c", + () -> openH2(h2Url("global_c"), physicalOpens) + )) { + assertEquals(2, registry.activePhysicalConnectionCount()); + } + } finally { + if (first != null) { + first.close(); + } + if (second != null) { + second.close(); + } + registry.close(); + } + assertEquals(0, registry.activePhysicalConnectionCount()); + } + + @Test + void globalPhysicalCapacityBackpressureKeepsActiveLeaseUsable() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + CountDownLatch abortCalled = new CountDownLatch(1); + String url = h2Url("shared_capacity"); + try (Connection ignored = openH2(url, physicalOpens)) { + // Keep H2 bootstrap outside the 250ms capacity watchdog exercised below. + } + physicalOpens.set(0); + JdbcConnectionPoolRegistry.PoolSettings settings = new JdbcConnectionPoolRegistry.PoolSettings( + true, + 2, + 0, + 250L, + 250L, + 10_000L, + 30_000L, + 60_000L, + 0, + 1, + 2 + ); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(settings); + JdbcConnectionPoolRegistry.Lease active = registry.borrow( + "shared-capacity", + () -> abortTrackingConnection( + openH2(url, physicalOpens), + abortCalled + ) + )) { + AgentRpcError capacity = assertThrows( + AgentRpcError.class, + () -> registry.borrow( + "shared-capacity", + () -> openH2(url, physicalOpens) + ) + ); + JsonObject data = AgentRpcError.toJson( + capacity, + AgentProtocol.METHOD_OPEN_SESSION, + "shared-capacity" + ).getAsJsonObject("data"); + assertEquals("keep", data.get("sessionDisposition").getAsString()); + assertFalse(abortCalled.await(1, TimeUnit.SECONDS)); + assertTrue(active.connection().isValid(1)); + assertEquals(1, registry.activePhysicalConnectionCount()); + } + } + + @Test + void physicalBudgetIsRetainedWhenDriverCloseCannotBeConfirmed() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicInteger closeAttempts = new AtomicInteger(); + JdbcConnectionPoolRegistry.PoolSettings settings = shortTimeoutPoolSettings(2, 1); + JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(settings); + try { + JdbcConnectionPoolRegistry.Lease lease = registry.borrow( + "identity-a", + () -> closeFailingConnection(openH2(h2Url("close_unknown"), physicalOpens), closeAttempts) + ); + lease.quarantine(); + lease.close(); + + awaitCount(closeAttempts, 1); + assertEquals(1, registry.activePhysicalConnectionCount()); + assertThrows( + AgentRpcError.class, + () -> registry.borrow("identity-b", () -> openH2(h2Url("close_unknown_b"), physicalOpens)) + ); + } finally { + registry.close(); + } + } + + @Test + void hikariPreCloseNetworkTimeoutIsBoundedAndPoisonsIdentity() throws Exception { + AtomicBoolean blockNetworkTimeout = new AtomicBoolean(); + CountDownLatch networkTimeoutStarted = new CountDownLatch(1); + CountDownLatch releaseNetworkTimeout = new CountDownLatch(1); + String url = h2Url("pre_close_network_timeout"); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(shortTimeoutPoolSettings(1, 1))) { + JdbcConnectionPoolRegistry.Lease lease = registry.borrow( + "pre-close-network-timeout", + () -> networkTimeoutBlockingConnection( + openH2(url, new AtomicInteger()), + blockNetworkTimeout, + networkTimeoutStarted, + releaseNetworkTimeout + ) + ); + blockNetworkTimeout.set(true); + lease.quarantine(); + lease.close(); + assertTrue(networkTimeoutStarted.await(1, TimeUnit.SECONDS)); + Thread.sleep(350L); + + AgentRpcError retry = assertThrows( + AgentRpcError.class, + () -> registry.borrow( + "pre-close-network-timeout", + () -> openH2(url, new AtomicInteger()) + ) + ); + JsonObject data = AgentRpcError.toJson(retry, AgentProtocol.METHOD_OPEN_SESSION, "retry") + .getAsJsonObject("data"); + assertEquals("replace_runtime", data.get("sessionDisposition").getAsString()); + } finally { + releaseNetworkTimeout.countDown(); + } + } + + @Test + void asynchronousAbortRetainsPhysicalBudgetWhenOnlyLogicalCloseIsConfirmed() throws Exception { + CountDownLatch abortScheduled = new CountDownLatch(1); + CountDownLatch releaseTermination = new CountDownLatch(1); + ExecutorService worker = Executors.newSingleThreadExecutor(); + JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(shortTimeoutPoolSettings(1, 1)); + registry.borrow( + "async-abort", + () -> asynchronousAbortConnection( + openH2(h2Url("async_abort"), new AtomicInteger()), + abortScheduled, + releaseTermination + ) + ); + try { + Future close = worker.submit(registry::close); + assertTrue(abortScheduled.await(1, TimeUnit.SECONDS)); + assertEquals(1, registry.activePhysicalConnectionCount()); + releaseTermination.countDown(); + close.get(2, TimeUnit.SECONDS); + awaitPhysicalConnectionCount(registry, 0); + } finally { + releaseTermination.countDown(); + registry.close(); + worker.shutdownNow(); + } + } + + @Test + void physicalBudgetIsRetainedWhenInitializationReportsUnknownConnectionState() { + JdbcConnectionPoolRegistry.PoolSettings settings = shortTimeoutPoolSettings(2, 1); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(settings)) { + assertThrows( + AgentRpcError.class, + () -> registry.borrow( + "identity-a", + () -> { + throw new JdbcConnectionPoolRegistry.PhysicalConnectionStateUnknownException( + new SQLException("initialization failed and close did not complete") + ); + } + ) + ); + assertEquals(1, registry.activePhysicalConnectionCount()); + } + } + + @Test + void quarantineThresholdIsCountedOncePerLeaseAndReleasedOnClose() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + String url = h2Url("quarantine_threshold"); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(poolSettings(4, 0, 32, 2))) { + JdbcConnectionPoolRegistry.Lease first = registry.borrow( + "same-identity", + () -> openH2(url, physicalOpens) + ); + JdbcConnectionPoolRegistry.Lease second = registry.borrow( + "same-identity", + () -> openH2(url, physicalOpens) + ); + + assertFalse(first.quarantine()); + assertFalse(first.quarantine()); + assertTrue(second.quarantine()); + first.evict(); + second.evict(); + + JdbcConnectionPoolRegistry.Lease replacement = registry.borrow( + "same-identity", + () -> openH2(url, physicalOpens) + ); + assertFalse(replacement.quarantine()); + replacement.evict(); + } + } + + @Test + void missingOrUnknownSessionRoleDefaultsToWorkload() { + assertEquals(JdbcSessionRole.WORKLOAD, JdbcSessionRole.from(null)); + assertEquals(JdbcSessionRole.WORKLOAD, JdbcSessionRole.from("")); + assertEquals(JdbcSessionRole.WORKLOAD, JdbcSessionRole.from("future-role")); + assertEquals(JdbcSessionRole.METADATA, JdbcSessionRole.from(" metadata ")); + } + + @Test + void registryBorrowTimeoutBoundsBlockedPhysicalConnect() throws Exception { + AtomicInteger connectAttempts = new AtomicInteger(); + CountDownLatch connectStarted = new CountDownLatch(1); + CountDownLatch releaseConnect = new CountDownLatch(1); + JdbcConnectionPoolRegistry.PoolSettings settings = new JdbcConnectionPoolRegistry.PoolSettings( + 1, + 0, + 250L, + 250L, + 10_000L, + 30_000L, + 60_000L + ); + JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(settings); + long startedAtMillis = System.currentTimeMillis(); + try { + AgentRpcError initial = assertThrows(AgentRpcError.class, () -> registry.borrow("blocked-connect", () -> { + connectAttempts.incrementAndGet(); + connectStarted.countDown(); + awaitUninterruptibly(releaseConnect); + throw new SQLException("late connect"); + })); + long elapsedMillis = System.currentTimeMillis() - startedAtMillis; + assertTrue(connectStarted.await(1, TimeUnit.SECONDS)); + assertTrue(elapsedMillis < 450L, () -> "blocked connect took " + elapsedMillis + "ms"); + JsonObject initialData = AgentRpcError.toJson(initial, AgentProtocol.METHOD_OPEN_SESSION, "initial") + .getAsJsonObject("data"); + assertEquals("replace_runtime", initialData.get("sessionDisposition").getAsString()); + AgentRpcError retry = assertThrows( + AgentRpcError.class, + () -> registry.borrow("blocked-connect", () -> openH2(h2Url("blocked_retry"), new AtomicInteger())) + ); + JsonObject retryData = AgentRpcError.toJson(retry, AgentProtocol.METHOD_OPEN_SESSION, "retry") + .getAsJsonObject("data"); + assertEquals("replace_runtime", retryData.get("sessionDisposition").getAsString()); + assertEquals(1, connectAttempts.get()); + releaseConnect.countDown(); + awaitPhysicalConnectionCount(registry, 0); + } finally { + releaseConnect.countDown(); + registry.close(); + } + } + + @Test + void registryBorrowTimeoutBoundsBlockedHikariValidationAndPoisonsIdentity() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicBoolean blockValidation = new AtomicBoolean(); + CountDownLatch validationStarted = new CountDownLatch(1); + CountDownLatch releaseValidation = new CountDownLatch(1); + ExecutorService worker = Executors.newSingleThreadExecutor(); + String url = h2Url("blocked_validation"); + try (Connection ignored = openH2(url, physicalOpens)) { + // Keep H2 bootstrap outside the 250ms validation watchdog exercised below. + } + physicalOpens.set(0); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(shortTimeoutPoolSettings(1, 32))) { + try (JdbcConnectionPoolRegistry.Lease ignored = registry.borrow( + "blocked-validation", + () -> validationBlockingConnection( + openH2(url, physicalOpens), + blockValidation, + validationStarted, + releaseValidation + ) + )) { + assertEquals(1, physicalOpens.get()); + } + blockValidation.set(true); + Thread.sleep(600L); // Hikari skips validation for connections used within its 500ms alive-bypass window. + + long startedAtNanos = System.nanoTime(); + Future blocked = worker.submit( + () -> registry.borrow("blocked-validation", () -> openH2(url, physicalOpens)) + ); + assertTrue(validationStarted.await(1, TimeUnit.SECONDS)); + AgentRpcError timeout = futureFailure(blocked, AgentRpcError.class, 2, TimeUnit.SECONDS); + long elapsedMillis = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - startedAtNanos); + assertTrue(elapsedMillis < 450L, () -> "blocked validation took " + elapsedMillis + "ms"); + JsonObject timeoutData = AgentRpcError.toJson(timeout, AgentProtocol.METHOD_OPEN_SESSION, "blocked") + .getAsJsonObject("data"); + assertEquals("replace_runtime", timeoutData.get("sessionDisposition").getAsString()); + + AgentRpcError retry = assertThrows( + AgentRpcError.class, + () -> registry.borrow("blocked-validation", () -> openH2(url, physicalOpens)) + ); + JsonObject retryData = AgentRpcError.toJson(retry, AgentProtocol.METHOD_OPEN_SESSION, "retry") + .getAsJsonObject("data"); + assertEquals("replace_runtime", retryData.get("sessionDisposition").getAsString()); + assertEquals(1, physicalOpens.get()); + releaseValidation.countDown(); + awaitPhysicalConnectionCount(registry, 0); + } finally { + releaseValidation.countDown(); + worker.shutdownNow(); + } + } + + @Test + void physicalBudgetWaitAndConnectShareOneBorrowDeadline() throws Exception { + CountDownLatch connectStarted = new CountDownLatch(1); + CountDownLatch releaseConnect = new CountDownLatch(1); + ExecutorService worker = Executors.newSingleThreadExecutor(); + JdbcConnectionPoolRegistry.PoolSettings settings = new JdbcConnectionPoolRegistry.PoolSettings( + true, + 1, + 0, + 500L, + 500L, + 10_000L, + 30_000L, + 0L, + 0, + 1, + 2 + ); + JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(settings); + JdbcConnectionPoolRegistry.Lease held = registry.borrow( + "deadline-holder", + () -> openH2(h2Url("deadline_holder"), new AtomicInteger()) + ); + try { + long startedAtNanos = System.nanoTime(); + Future blocked = worker.submit(() -> registry.borrow( + "deadline-waiter", + () -> { + connectStarted.countDown(); + awaitUninterruptibly(releaseConnect); + throw new SQLException("late connect"); + } + )); + Thread.sleep(400L); + held.close(); + held = null; + registry.retireUnusedPools(); + assertTrue(connectStarted.await(1, TimeUnit.SECONDS)); + + AgentRpcError timeout = futureFailure(blocked, AgentRpcError.class, 2, TimeUnit.SECONDS); + long elapsedMillis = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - startedAtNanos); + assertTrue(elapsedMillis < 700L, () -> "budget and connect took " + elapsedMillis + "ms"); + JsonObject data = AgentRpcError.toJson(timeout, AgentProtocol.METHOD_OPEN_SESSION, "deadline") + .getAsJsonObject("data"); + assertEquals("replace_runtime", data.get("sessionDisposition").getAsString()); + } finally { + if (held != null) { + held.close(); + } + releaseConnect.countDown(); + worker.shutdownNow(); + registry.close(); + } + } + + @Test + void metadataCapacityTimeoutKeepsIdentityRoutable() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + ExecutorService worker = Executors.newSingleThreadExecutor(); + String url = h2Url("metadata_capacity"); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(shortTimeoutPoolSettings(1, 32))) { + JdbcConnectionPoolRegistry.Lease held = registry.borrow( + "metadata-capacity", + JdbcSessionRole.METADATA, + () -> openH2(url, physicalOpens) + ); + try { + Future waiting = worker.submit(() -> registry.borrow( + "metadata-capacity", + JdbcSessionRole.METADATA, + () -> openH2(url, physicalOpens) + )); + AgentRpcError capacity = futureFailure(waiting, AgentRpcError.class, 2, TimeUnit.SECONDS); + JsonObject data = AgentRpcError.toJson(capacity, AgentProtocol.METHOD_OPEN_SESSION, "metadata") + .getAsJsonObject("data"); + assertTrue(data.get("retryable").getAsBoolean()); + assertEquals("keep", data.get("sessionDisposition").getAsString()); + } finally { + held.close(); + } + + try (JdbcConnectionPoolRegistry.Lease ignored = registry.borrow( + "metadata-capacity", + JdbcSessionRole.METADATA, + () -> openH2(url, physicalOpens) + )) { + assertEquals(1, physicalOpens.get()); + } + } finally { + worker.shutdownNow(); + } + } + + @Test + void blockedValidationDoesNotSerializeHealthyMetadataCheckout() throws Exception { + AtomicBoolean blockNextValidation = new AtomicBoolean(); + AtomicInteger physicalOpens = new AtomicInteger(); + CountDownLatch validationStarted = new CountDownLatch(1); + CountDownLatch releaseValidation = new CountDownLatch(1); + ExecutorService worker = Executors.newSingleThreadExecutor(); + String url = h2Url("parallel_validation"); + JdbcConnectionPoolRegistry.PoolSettings settings = new JdbcConnectionPoolRegistry.PoolSettings( + true, + 2, + 0, + 500L, + 500L, + 10_000L, + 30_000L, + 60_000L, + 1, + 2, + 2 + ); + try (Connection ignored = openH2(url, physicalOpens)) { + // Keep H2 bootstrap outside the concurrent validation watchdog. + } + physicalOpens.set(0); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(settings)) { + JdbcConnectionPoolRegistry.ConnectionFactory factory = () -> validationBlockingOnceConnection( + openH2(url, physicalOpens), + blockNextValidation, + validationStarted, + releaseValidation + ); + JdbcConnectionPoolRegistry.Lease first = registry.borrow("parallel-validation", factory); + first.close(); + Thread.sleep(600L); + + blockNextValidation.set(true); + Future blocked = worker.submit( + () -> registry.borrow("parallel-validation", factory) + ); + assertTrue(validationStarted.await(1, TimeUnit.SECONDS)); + long metadataStartedAt = System.nanoTime(); + try (JdbcConnectionPoolRegistry.Lease metadata = registry.borrow( + "parallel-validation", + JdbcSessionRole.METADATA, + factory + )) { + long elapsedMillis = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - metadataStartedAt); + assertTrue(elapsedMillis < 400L, () -> "metadata checkout took " + elapsedMillis + "ms"); + assertTrue(metadata.connection().isValid(1)); + } + AgentRpcError timeout = futureFailure(blocked, AgentRpcError.class, 2, TimeUnit.SECONDS); + JsonObject timeoutData = AgentRpcError.toJson(timeout, AgentProtocol.METHOD_OPEN_SESSION, "blocked") + .getAsJsonObject("data"); + assertEquals("replace_runtime", timeoutData.get("sessionDisposition").getAsString()); + + AgentRpcError retry = assertThrows( + AgentRpcError.class, + () -> registry.borrow("parallel-validation", factory) + ); + JsonObject retryData = AgentRpcError.toJson(retry, AgentProtocol.METHOD_OPEN_SESSION, "retry") + .getAsJsonObject("data"); + assertEquals("replace_runtime", retryData.get("sessionDisposition").getAsString()); + } finally { + releaseValidation.countDown(); + worker.shutdownNow(); + } + } + + @RepeatedTest(5) + void deterministicHikariSetupFailureDoesNotPoisonIdentity() throws Exception { + AtomicBoolean failSetup = new AtomicBoolean(true); + AtomicInteger physicalOpens = new AtomicInteger(); + String url = h2Url("setup_failure"); + try (Connection ignored = openH2(url, physicalOpens)) { + // Keep H2 bootstrap outside the setup classification watchdog. + } + physicalOpens.set(0); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(shortTimeoutPoolSettings(1, 32))) { + SQLException setupFailure = assertThrows(SQLException.class, () -> registry.borrow( + "setup-failure", + () -> setupFailingConnection(openH2(url, physicalOpens), failSetup) + )); + JsonObject setupData = AgentRpcError.toJson( + setupFailure, + AgentProtocol.METHOD_OPEN_SESSION, + "setup-failure" + ).getAsJsonObject("data"); + assertEquals("keep", setupData.get("sessionDisposition").getAsString()); + + failSetup.set(false); + try (JdbcConnectionPoolRegistry.Lease ignored = registry.borrow( + "setup-failure", + () -> setupFailingConnection(openH2(url, physicalOpens), failSetup) + )) { + assertTrue(physicalOpens.get() >= 2); + } + } + } + + @Test + void blockedSetupAfterKnownFailurePoisonsCurrentAttemptGeneration() throws Exception { + AtomicInteger connectionAttempts = new AtomicInteger(); + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicBoolean failFirstSetup = new AtomicBoolean(true); + CountDownLatch setupStarted = new CountDownLatch(1); + CountDownLatch releaseSetup = new CountDownLatch(1); + ExecutorService worker = Executors.newSingleThreadExecutor(); + String url = h2Url("known_then_blocked_setup"); + try (Connection ignored = openH2(url, physicalOpens)) { + // Keep H2 bootstrap outside the setup generation watchdog. + } + physicalOpens.set(0); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(shortTimeoutPoolSettings(1, 32))) { + JdbcConnectionPoolRegistry.ConnectionFactory factory = () -> { + Connection connection = openH2(url, physicalOpens); + if (connectionAttempts.getAndIncrement() == 0) { + return setupFailingConnection(connection, failFirstSetup); + } + return setupBlockingConnection(connection, setupStarted, releaseSetup); + }; + Future blocked = worker.submit( + () -> registry.borrow("known-then-blocked-setup", factory) + ); + assertTrue(setupStarted.await(1, TimeUnit.SECONDS)); + + AgentRpcError timeout = futureFailure(blocked, AgentRpcError.class, 2, TimeUnit.SECONDS); + JsonObject data = AgentRpcError.toJson(timeout, AgentProtocol.METHOD_OPEN_SESSION, "blocked") + .getAsJsonObject("data"); + assertEquals("replace_runtime", data.get("sessionDisposition").getAsString()); + + AgentRpcError retry = assertThrows( + AgentRpcError.class, + () -> registry.borrow("known-then-blocked-setup", factory) + ); + JsonObject retryData = AgentRpcError.toJson(retry, AgentProtocol.METHOD_OPEN_SESSION, "retry") + .getAsJsonObject("data"); + assertEquals("replace_runtime", retryData.get("sessionDisposition").getAsString()); + } finally { + releaseSetup.countDown(); + worker.shutdownNow(); + } + } + + @Test + void blockedHikariSetupPoisonsIdentityWithinBorrowDeadline() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + CountDownLatch setupStarted = new CountDownLatch(1); + CountDownLatch releaseSetup = new CountDownLatch(1); + ExecutorService worker = Executors.newSingleThreadExecutor(); + String url = h2Url("blocked_setup"); + try (Connection ignored = openH2(url, physicalOpens)) { + // Keep H2 bootstrap outside the 250ms setup watchdog exercised below. + } + physicalOpens.set(0); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(shortTimeoutPoolSettings(1, 32))) { + long startedAtNanos = System.nanoTime(); + Future blocked = worker.submit(() -> registry.borrow( + "blocked-setup", + () -> setupBlockingConnection( + openH2(url, physicalOpens), + setupStarted, + releaseSetup + ) + )); + assertTrue(setupStarted.await(1, TimeUnit.SECONDS)); + AgentRpcError timeout = futureFailure(blocked, AgentRpcError.class, 2, TimeUnit.SECONDS); + long elapsedMillis = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - startedAtNanos); + assertTrue(elapsedMillis < 450L, () -> "blocked setup took " + elapsedMillis + "ms"); + JsonObject data = AgentRpcError.toJson(timeout, AgentProtocol.METHOD_OPEN_SESSION, "blocked-setup") + .getAsJsonObject("data"); + assertEquals("replace_runtime", data.get("sessionDisposition").getAsString()); + + AgentRpcError retry = assertThrows( + AgentRpcError.class, + () -> registry.borrow("blocked-setup", () -> openH2(url, physicalOpens)) + ); + JsonObject retryData = AgentRpcError.toJson(retry, AgentProtocol.METHOD_OPEN_SESSION, "retry") + .getAsJsonObject("data"); + assertEquals("replace_runtime", retryData.get("sessionDisposition").getAsString()); + } finally { + releaseSetup.countDown(); + worker.shutdownNow(); + } + } + + @Test + void poolRetirementDoesNotOutliveBorrowDeadline() throws Exception { + AtomicBoolean blockShutdown = new AtomicBoolean(); + AtomicInteger physicalOpens = new AtomicInteger(); + CountDownLatch shutdownStarted = new CountDownLatch(1); + CountDownLatch releaseShutdown = new CountDownLatch(1); + JdbcConnectionPoolRegistry.PoolSettings settings = new JdbcConnectionPoolRegistry.PoolSettings( + true, + 1, + 0, + 250L, + 250L, + 10_000L, + 30_000L, + 0L, + 0, + 32, + 2 + ); + JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(settings); + try (JdbcConnectionPoolRegistry.Lease ignored = registry.borrow( + "retired-pool", + () -> closeBlockingConnection( + openH2(h2Url("blocked_pool_close"), physicalOpens), + blockShutdown, + shutdownStarted, + releaseShutdown + ) + )) { + assertEquals(1, physicalOpens.get()); + } + try { + blockShutdown.set(true); + long startedAtNanos = System.nanoTime(); + registry.retireUnusedPools(); + assertTrue(shutdownStarted.await(1, TimeUnit.SECONDS)); + long elapsedMillis = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - startedAtNanos); + assertTrue(elapsedMillis < 450L, () -> "pool retirement took " + elapsedMillis + "ms"); + AgentRpcError error = assertThrows(AgentRpcError.class, () -> registry.borrow( + "different-identity", + () -> openH2(h2Url("different_identity"), physicalOpens) + )); + JsonObject data = AgentRpcError.toJson(error, AgentProtocol.METHOD_OPEN_SESSION, "pool-close") + .getAsJsonObject("data"); + assertEquals("replace_runtime", data.get("sessionDisposition").getAsString()); + } finally { + releaseShutdown.countDown(); + registry.close(); + } + } + + @Test + void onlyCausalPhysicalConnectionFailuresRequireRuntimeReplacement() { + assertFalse(JdbcConnectionPoolRegistry.requiresRuntimeReplacement(new SQLException("authentication failed"))); + assertTrue(JdbcConnectionPoolRegistry.requiresRuntimeReplacement( + new JdbcConnectionPoolRegistry.PhysicalConnectionStateUnknownException(new SQLException("connect stuck")) + )); + } + @Test void registryRetiresOnlyUnusedPools() throws Exception { AtomicInteger physicalOpens = new AtomicInteger(); @@ -338,6 +1193,491 @@ class JdbcConnectionPoolingTest { } } + @Test + void closingBusySessionDetachesImmediatelyAndPoisonsItsLease() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicInteger physicalCloses = new AtomicInteger(); + AtomicInteger requestIds = new AtomicInteger(); + AtomicBoolean blockNextMetadataCall = new AtomicBoolean(true); + CountDownLatch requestStarted = new CountDownLatch(1); + CountDownLatch releaseRequest = new CountDownLatch(1); + ExecutorService workers = Executors.newSingleThreadExecutor(); + String url = h2Url("quarantine_busy_session"); + try (MultiSessionJsonRpcServer server = new MultiSessionJsonRpcServer( + () -> new H2TestAgent(url, physicalOpens) { + @Override + protected Connection openConnection(ConnectParams params) throws Exception { + return trackedConnection(openH2(url, physicalOpens), physicalCloses); + } + + @Override + public List listDatabases() { + if (blockNextMetadataCall.compareAndSet(true, false)) { + requestStarted.countDown(); + awaitUninterruptibly(releaseRequest); + } + return Collections.emptyList(); + } + }, + poolSettings(2) + )) { + openSession(server, requestIds, "blocked"); + Future blocked = workers.submit(() -> request( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("blocked") + )); + assertTrue(requestStarted.await(2, TimeUnit.SECONDS)); + + long closeStarted = System.currentTimeMillis(); + closeSession(server, requestIds, "blocked"); + assertTrue(System.currentTimeMillis() - closeStarted < 500L); + + JsonObject staleResponse = rawRequest( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("blocked") + ); + assertTrue(staleResponse.has("error"), staleResponse::toString); + + openSession(server, requestIds, "replacement"); + JsonObject replacement = query(server, requestIds, "replacement", "SELECT 1", null); + assertEquals(1, replacement.getAsJsonArray("rows").get(0).getAsJsonArray().get(0).getAsInt()); + assertEquals(2, physicalOpens.get()); + + releaseRequest.countDown(); + blocked.get(2, TimeUnit.SECONDS); + awaitCount(physicalCloses, 1); + } finally { + releaseRequest.countDown(); + workers.shutdownNow(); + } + } + + @Test + void closingIdlePinnedSingleConnectionSessionDoesNotRequireRuntimeReplacement() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicInteger requestIds = new AtomicInteger(); + String url = h2Url("pinned_single_close"); + try (MultiSessionJsonRpcServer server = new MultiSessionJsonRpcServer( + () -> new H2TestAgent(url, physicalOpens), + poolSettings(1, 0, 32, 2) + )) { + openSession(server, requestIds, "pinned"); + query(server, requestIds, "pinned", "SET SCHEMA PUBLIC", null); + + JsonObject response = rawRequest( + server, + requestIds, + AgentProtocol.METHOD_CLOSE_SESSION, + sessionParams("pinned") + ); + + assertTrue(response.has("result"), response::toString); + openSession(server, requestIds, "replacement"); + JsonObject result = query(server, requestIds, "replacement", "SELECT 1", null); + assertEquals(1, result.getAsJsonArray("rows").get(0).getAsJsonArray().get(0).getAsInt()); + } + } + + @Test + void blockedPinnedConnectionClosePoisonsIdentityForNextSession() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicInteger requestIds = new AtomicInteger(); + AtomicBoolean blockClose = new AtomicBoolean(); + CountDownLatch closeStarted = new CountDownLatch(1); + CountDownLatch releaseClose = new CountDownLatch(1); + String url = h2Url("pinned_blocked_close"); + try (MultiSessionJsonRpcServer server = new MultiSessionJsonRpcServer( + () -> new H2TestAgent(url, physicalOpens) { + @Override + protected Connection openConnection(ConnectParams params) throws Exception { + return closeBlockingConnection( + openH2(url, physicalOpens), + blockClose, + closeStarted, + releaseClose + ); + } + }, + shortTimeoutPoolSettings(1, 32) + )) { + openSession(server, requestIds, "pinned"); + query(server, requestIds, "pinned", "SET SCHEMA PUBLIC", null); + blockClose.set(true); + + JsonObject close = rawRequest( + server, + requestIds, + AgentProtocol.METHOD_CLOSE_SESSION, + sessionParams("pinned") + ); + assertTrue(close.has("result"), close::toString); + assertTrue(closeStarted.await(1, TimeUnit.SECONDS)); + Thread.sleep(400L); + + JsonObject replacementParams = sessionParams("replacement"); + replacementParams.addProperty("database", "pooling"); + replacementParams.addProperty("username", "sa"); + replacementParams.addProperty("password", ""); + JsonObject retry = rawRequest( + server, + requestIds, + AgentProtocol.METHOD_OPEN_SESSION, + replacementParams + ); + assertTrue(retry.has("error"), retry::toString); + JsonObject data = retry.getAsJsonObject("error").getAsJsonObject("data"); + assertEquals("replace_runtime", data.get("sessionDisposition").getAsString()); + } finally { + releaseClose.countDown(); + } + } + + @Test + void closingSessionDoesNotWaitForAnotherRequestBlockedOnPoolCheckout() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicInteger agentIndexes = new AtomicInteger(); + AtomicInteger requestIds = new AtomicInteger(); + CountDownLatch holderStarted = new CountDownLatch(1); + CountDownLatch releaseHolder = new CountDownLatch(1); + ExecutorService workers = Executors.newFixedThreadPool(2); + String url = h2Url("quarantine_checkout_wait"); + try (MultiSessionJsonRpcServer server = new MultiSessionJsonRpcServer( + () -> { + int agentIndex = agentIndexes.incrementAndGet(); + return new H2TestAgent(url, physicalOpens) { + @Override + public List listDatabases() { + if (agentIndex == 1) { + holderStarted.countDown(); + awaitUninterruptibly(releaseHolder); + } + return Collections.emptyList(); + } + }; + }, + poolSettings(1, 0, 32, 1) + )) { + openSession(server, requestIds, "holder"); + openSession(server, requestIds, "waiting"); + Future holder = workers.submit(() -> request( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("holder") + )); + assertTrue(holderStarted.await(2, TimeUnit.SECONDS)); + Future waiting = workers.submit(() -> rawRequest( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("waiting") + )); + Thread.sleep(100L); + + long closeStarted = System.nanoTime(); + closeSession(server, requestIds, "waiting"); + long closeElapsedMillis = TimeUnit.NANOSECONDS.toMillis(System.nanoTime() - closeStarted); + assertTrue(closeElapsedMillis < 500L, () -> "close waited " + closeElapsedMillis + "ms"); + + releaseHolder.countDown(); + holder.get(2, TimeUnit.SECONDS); + assertTrue(waiting.get(2, TimeUnit.SECONDS).has("error")); + } finally { + releaseHolder.countDown(); + workers.shutdownNow(); + } + } + + @Test + void workloadPermitTimeoutKeepsSessionAndRuntimeRoutable() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicInteger agentIndexes = new AtomicInteger(); + AtomicInteger requestIds = new AtomicInteger(); + CountDownLatch holderStarted = new CountDownLatch(1); + CountDownLatch releaseHolder = new CountDownLatch(1); + ExecutorService workers = Executors.newSingleThreadExecutor(); + String url = h2Url("workload_backpressure"); + try (MultiSessionJsonRpcServer server = new MultiSessionJsonRpcServer( + () -> { + int agentIndex = agentIndexes.incrementAndGet(); + return new H2TestAgent(url, physicalOpens) { + @Override + public List listDatabases() { + if (agentIndex == 1) { + holderStarted.countDown(); + awaitUninterruptibly(releaseHolder); + } + return Collections.emptyList(); + } + }; + }, + shortTimeoutPoolSettings(1, 32) + )) { + openSession(server, requestIds, "holder"); + openSession(server, requestIds, "waiting"); + Future holder = workers.submit(() -> request( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("holder") + )); + assertTrue(holderStarted.await(2, TimeUnit.SECONDS)); + + JsonObject saturated = rawRequest( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("waiting") + ); + JsonObject data = saturated.getAsJsonObject("error").getAsJsonObject("data"); + assertEquals("resource", data.get("category").getAsString()); + assertTrue(data.get("retryable").getAsBoolean()); + assertEquals("keep", data.get("sessionDisposition").getAsString()); + + releaseHolder.countDown(); + holder.get(2, TimeUnit.SECONDS); + JsonObject recovered = request( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("waiting") + ); + assertTrue(recovered.getAsJsonArray("result").isEmpty()); + } finally { + releaseHolder.countDown(); + workers.shutdownNow(); + } + } + + @Test + void requestExecutorBackpressureKeepsExistingSessionsAndRuntimeRoutable() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicInteger agentIndexes = new AtomicInteger(); + AtomicInteger requestIds = new AtomicInteger(); + CountDownLatch holderStarted = new CountDownLatch(1); + CountDownLatch releaseHolder = new CountDownLatch(1); + BlockingQueue responses = new LinkedBlockingQueue<>(); + String url = h2Url("request_executor_backpressure"); + try (MultiSessionJsonRpcServer server = new MultiSessionJsonRpcServer( + () -> { + int agentIndex = agentIndexes.incrementAndGet(); + return new H2TestAgent(url, physicalOpens) { + @Override + public List listDatabases() { + if (agentIndex == 1) { + holderStarted.countDown(); + awaitUninterruptibly(releaseHolder); + } + return Collections.emptyList(); + } + }; + }, + poolSettings(2), + new MultiSessionJsonRpcServer.RuntimeLimits(1, 1) + )) { + openSession(server, requestIds, "holder"); + openSession(server, requestIds, "waiting"); + + int holderRequestId = requestIds.incrementAndGet(); + server.executeRequest( + jsonRpcRequest(holderRequestId, AgentProtocol.METHOD_LIST_DATABASES, sessionParams("holder")), + responses::add + ); + assertTrue(holderStarted.await(2, TimeUnit.SECONDS)); + + int waitingRequestId = requestIds.incrementAndGet(); + server.executeRequest( + jsonRpcRequest(waitingRequestId, AgentProtocol.METHOD_LIST_DATABASES, sessionParams("waiting")), + responses::add + ); + JsonObject saturated = responses.poll(2, TimeUnit.SECONDS); + assertNotNull(saturated); + assertEquals(waitingRequestId, saturated.get("id").getAsInt()); + JsonObject data = saturated.getAsJsonObject("error").getAsJsonObject("data"); + assertEquals("resource", data.get("category").getAsString()); + assertTrue(data.get("retryable").getAsBoolean()); + assertEquals("keep", data.get("sessionDisposition").getAsString()); + + releaseHolder.countDown(); + JsonObject holder = responses.poll(2, TimeUnit.SECONDS); + assertNotNull(holder); + assertEquals(holderRequestId, holder.get("id").getAsInt()); + assertFalse(holder.has("error"), () -> holder.toString()); + + request(server, requestIds, AgentProtocol.METHOD_LIST_DATABASES, sessionParams("holder")); + request(server, requestIds, AgentProtocol.METHOD_LIST_DATABASES, sessionParams("waiting")); + assertEquals(2, agentIndexes.get()); + } finally { + releaseHolder.countDown(); + } + } + + @Test + void idleAffinityLeaseDoesNotCountAsQuarantinedOperation() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + String url = h2Url("idle_affinity_quarantine"); + try (JdbcConnectionPoolRegistry registry = new JdbcConnectionPoolRegistry(poolSettings(4, 0, 32, 2))) { + for (int index = 0; index < 2; index++) { + H2TestAgent agent = new H2TestAgent(url, physicalOpens); + agent.attachConnectionPoolRegistry(registry); + JsonRpcServer server = new JsonRpcServer(agent); + JsonObject connectParams = new JsonObject(); + connectParams.addProperty("database", "pooling"); + connectParams.addProperty("username", "sa"); + connectParams.addProperty("password", ""); + server.dispatchForRuntime(AgentProtocol.METHOD_CONNECT, connectParams); + + JsonObject queryParams = new JsonObject(); + queryParams.addProperty("sql", "SET SCHEMA PUBLIC"); + server.dispatchForRuntime(AgentProtocol.METHOD_EXECUTE_QUERY, queryParams); + + assertFalse(server.quarantine()); + server.dispatchForRuntime(AgentProtocol.METHOD_DISCONNECT, new JsonObject()); + } + } + } + + @Test + void cleanupSaturationReturnsStructuredRuntimeReplacementError() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicInteger requestIds = new AtomicInteger(); + CountDownLatch requestsStarted = new CountDownLatch(2); + CountDownLatch releaseRequests = new CountDownLatch(1); + ExecutorService workers = Executors.newFixedThreadPool(2); + String url = h2Url("cleanup_saturation"); + try (MultiSessionJsonRpcServer server = new MultiSessionJsonRpcServer( + () -> new H2TestAgent(url, physicalOpens) { + @Override + public List listDatabases() { + requestsStarted.countDown(); + awaitUninterruptibly(releaseRequests); + return Collections.emptyList(); + } + }, + poolSettings(2, 0, 32, 8), + new MultiSessionJsonRpcServer.RuntimeLimits(4, 1) + )) { + openSession(server, requestIds, "first"); + openSession(server, requestIds, "second"); + Future first = workers.submit(() -> request( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("first") + )); + Future second = workers.submit(() -> request( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("second") + )); + assertTrue(requestsStarted.await(2, TimeUnit.SECONDS)); + + closeSession(server, requestIds, "first"); + JsonObject saturated = rawRequest( + server, + requestIds, + AgentProtocol.METHOD_CLOSE_SESSION, + sessionParams("second") + ); + JsonObject data = saturated.getAsJsonObject("error").getAsJsonObject("data"); + assertEquals("resource", data.get("category").getAsString()); + assertEquals("replace_runtime", data.get("sessionDisposition").getAsString()); + + releaseRequests.countDown(); + first.get(2, TimeUnit.SECONDS); + second.get(2, TimeUnit.SECONDS); + } finally { + releaseRequests.countDown(); + workers.shutdownNow(); + } + } + + @Test + void quarantineThresholdReturnsStructuredRuntimeReplacementError() throws Exception { + AtomicInteger physicalOpens = new AtomicInteger(); + AtomicInteger requestIds = new AtomicInteger(); + CountDownLatch requestsStarted = new CountDownLatch(2); + CountDownLatch releaseRequests = new CountDownLatch(1); + ExecutorService workers = Executors.newFixedThreadPool(2); + String url = h2Url("quarantine_runtime_replacement"); + try (MultiSessionJsonRpcServer server = new MultiSessionJsonRpcServer( + () -> new H2TestAgent(url, physicalOpens) { + @Override + public List listDatabases() { + requestsStarted.countDown(); + awaitUninterruptibly(releaseRequests); + return Collections.emptyList(); + } + }, + poolSettings(4, 0, 32, 2), + new MultiSessionJsonRpcServer.RuntimeLimits(4, 2) + )) { + openSession(server, requestIds, "first"); + openSession(server, requestIds, "second"); + Future first = workers.submit(() -> request( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("first") + )); + Future second = workers.submit(() -> request( + server, + requestIds, + AgentProtocol.METHOD_LIST_DATABASES, + sessionParams("second") + )); + assertTrue(requestsStarted.await(2, TimeUnit.SECONDS)); + + closeSession(server, requestIds, "first"); + JsonObject threshold = rawRequest( + server, + requestIds, + AgentProtocol.METHOD_CLOSE_SESSION, + sessionParams("second") + ); + JsonObject data = threshold.getAsJsonObject("error").getAsJsonObject("data"); + assertEquals("resource", data.get("category").getAsString()); + assertEquals("replace_runtime", data.get("sessionDisposition").getAsString()); + + releaseRequests.countDown(); + first.get(2, TimeUnit.SECONDS); + second.get(2, TimeUnit.SECONDS); + } finally { + releaseRequests.countDown(); + workers.shutdownNow(); + } + } + + @Test + void jdbcConnectionErrorsCarryStructuredQuarantineData() { + JsonObject error = AgentRpcError.toJson( + new RuntimeException(new SQLException("connection lost", "08006")), + AgentProtocol.METHOD_EXECUTE_QUERY, + "session-1" + ); + JsonObject data = error.getAsJsonObject("data"); + assertEquals("connection", data.get("category").getAsString()); + assertFalse(data.get("retryable").getAsBoolean()); + assertEquals("quarantine", data.get("sessionDisposition").getAsString()); + assertEquals("session-1", data.get("agentSessionId").getAsString()); + } + + @Test + void jdbcConnectErrorsRemainRetryableAfterRuntimeRecovery() { + JsonObject error = AgentRpcError.toJson( + new RuntimeException(new SQLException("connection refused", "08001")), + AgentProtocol.METHOD_OPEN_SESSION, + "session-1" + ); + + assertTrue(error.getAsJsonObject("data").get("retryable").getAsBoolean()); + } + @Test void independentSessionsExecuteConcurrentlyThroughSharedPool() throws Exception { AtomicInteger physicalOpens = new AtomicInteger(); @@ -543,16 +1883,30 @@ class JdbcConnectionPoolingTest { String method, JsonObject params ) { - JsonObject request = new JsonObject(); - request.addProperty("jsonrpc", "2.0"); - request.addProperty("id", requestIds.incrementAndGet()); - request.addProperty("method", method); - request.add("params", params); - JsonObject response = JsonParser.parseString(server.handleRequest(GSON.toJson(request))).getAsJsonObject(); + JsonObject response = rawRequest(server, requestIds, method, params); assertFalse(response.has("error"), () -> response.toString()); return response; } + private static JsonObject rawRequest( + MultiSessionJsonRpcServer server, + AtomicInteger requestIds, + String method, + JsonObject params + ) { + JsonObject request = jsonRpcRequest(requestIds.incrementAndGet(), method, params); + return JsonParser.parseString(server.handleRequest(GSON.toJson(request))).getAsJsonObject(); + } + + private static JsonObject jsonRpcRequest(int requestId, String method, JsonObject params) { + JsonObject request = new JsonObject(); + request.addProperty("jsonrpc", "2.0"); + request.addProperty("id", requestId); + request.addProperty("method", method); + request.add("params", params); + return request; + } + private static JsonObject result(JsonObject response) { assertNotNull(response.get("result")); return response.getAsJsonObject("result"); @@ -565,20 +1919,343 @@ class JdbcConnectionPoolingTest { } private static JdbcConnectionPoolRegistry.PoolSettings poolSettings(int maximumPoolSize) { + return poolSettings(maximumPoolSize, 0, 32, 2); + } + + private static Connection openH2(String url, AtomicInteger physicalOpens) throws Exception { + physicalOpens.incrementAndGet(); + return DriverManager.getConnection(url, "sa", ""); + } + + private static JdbcConnectionPoolRegistry.PoolSettings poolSettings( + int maximumPoolSize, + int metadataReserve, + int globalMaximumPhysicalConnections, + int maxQuarantinedOperations + ) { return new JdbcConnectionPoolRegistry.PoolSettings( + true, maximumPoolSize, 0, 2_000L, 1_000L, 10_000L, 30_000L, - 60_000L + 60_000L, + metadataReserve, + globalMaximumPhysicalConnections, + maxQuarantinedOperations ); } - private static Connection openH2(String url, AtomicInteger physicalOpens) throws Exception { - physicalOpens.incrementAndGet(); - return DriverManager.getConnection(url, "sa", ""); + private static JdbcConnectionPoolRegistry.PoolSettings shortTimeoutPoolSettings( + int maximumPoolSize, + int globalMaximumPhysicalConnections + ) { + return new JdbcConnectionPoolRegistry.PoolSettings( + true, + maximumPoolSize, + 0, + 250L, + 250L, + 10_000L, + 30_000L, + 60_000L, + 0, + globalMaximumPhysicalConnections, + 2 + ); + } + + private static Connection trackedConnection(Connection delegate, AtomicInteger closeCount) { + AtomicBoolean closed = new AtomicBoolean(); + return (Connection) Proxy.newProxyInstance( + Connection.class.getClassLoader(), + new Class[] {Connection.class}, + (proxy, method, args) -> { + if ("close".equals(method.getName()) && closed.compareAndSet(false, true)) { + closeCount.incrementAndGet(); + } + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException error) { + throw error.getCause(); + } + } + ); + } + + private static Connection closeBlockingConnection( + Connection delegate, + AtomicBoolean blockClose, + CountDownLatch closeStarted, + CountDownLatch releaseClose + ) { + AtomicBoolean closing = new AtomicBoolean(); + return (Connection) Proxy.newProxyInstance( + Connection.class.getClassLoader(), + new Class[] {Connection.class}, + (proxy, method, args) -> { + if ("close".equals(method.getName()) + && blockClose.get() + && closing.compareAndSet(false, true)) { + closeStarted.countDown(); + awaitUninterruptibly(releaseClose); + } + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException error) { + throw error.getCause(); + } + } + ); + } + + private static Connection validationBlockingConnection( + Connection delegate, + AtomicBoolean blockValidation, + CountDownLatch validationStarted, + CountDownLatch releaseValidation + ) { + return (Connection) Proxy.newProxyInstance( + Connection.class.getClassLoader(), + new Class[] {Connection.class}, + (proxy, method, args) -> { + if ("isValid".equals(method.getName()) && blockValidation.get()) { + validationStarted.countDown(); + awaitUninterruptibly(releaseValidation); + } + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException error) { + throw error.getCause(); + } + } + ); + } + + private static Connection validationBlockingOnceConnection( + Connection delegate, + AtomicBoolean blockNextValidation, + CountDownLatch validationStarted, + CountDownLatch releaseValidation + ) { + return (Connection) Proxy.newProxyInstance( + Connection.class.getClassLoader(), + new Class[] {Connection.class}, + (proxy, method, args) -> { + if ("isValid".equals(method.getName()) && blockNextValidation.compareAndSet(true, false)) { + validationStarted.countDown(); + awaitUninterruptibly(releaseValidation); + } + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException error) { + throw error.getCause(); + } + } + ); + } + + private static Connection setupFailingConnection(Connection delegate, AtomicBoolean failSetup) { + return (Connection) Proxy.newProxyInstance( + Connection.class.getClassLoader(), + new Class[] {Connection.class}, + (proxy, method, args) -> { + if ("getAutoCommit".equals(method.getName()) && failSetup.get()) { + throw new SQLException("simulated deterministic Hikari setup failure"); + } + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException error) { + throw error.getCause(); + } + } + ); + } + + private static Connection setupBlockingConnection( + Connection delegate, + CountDownLatch setupStarted, + CountDownLatch releaseSetup + ) { + return (Connection) Proxy.newProxyInstance( + Connection.class.getClassLoader(), + new Class[] {Connection.class}, + (proxy, method, args) -> { + if ("getAutoCommit".equals(method.getName())) { + setupStarted.countDown(); + awaitUninterruptibly(releaseSetup); + } + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException error) { + throw error.getCause(); + } + } + ); + } + + private static Connection abortTrackingConnection(Connection delegate, CountDownLatch abortCalled) { + return (Connection) Proxy.newProxyInstance( + Connection.class.getClassLoader(), + new Class[] {Connection.class}, + (proxy, method, args) -> { + if ("abort".equals(method.getName())) { + abortCalled.countDown(); + } + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException error) { + throw error.getCause(); + } + } + ); + } + + private static Connection networkTimeoutBlockingConnection( + Connection delegate, + AtomicBoolean blockNetworkTimeout, + CountDownLatch networkTimeoutStarted, + CountDownLatch releaseNetworkTimeout + ) { + return (Connection) Proxy.newProxyInstance( + Connection.class.getClassLoader(), + new Class[] {Connection.class}, + (proxy, method, args) -> { + if ("setNetworkTimeout".equals(method.getName()) && blockNetworkTimeout.get()) { + networkTimeoutStarted.countDown(); + awaitUninterruptibly(releaseNetworkTimeout); + } + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException error) { + throw error.getCause(); + } + } + ); + } + + private static Connection asynchronousAbortConnection( + Connection delegate, + CountDownLatch abortScheduled, + CountDownLatch releaseTermination + ) { + AtomicBoolean closed = new AtomicBoolean(); + return (Connection) Proxy.newProxyInstance( + Connection.class.getClassLoader(), + new Class[] {Connection.class}, + (proxy, method, args) -> { + if ("abort".equals(method.getName())) { + Executor executor = (Executor) args[0]; + closed.set(true); + executor.execute(() -> { + abortScheduled.countDown(); + awaitUninterruptibly(releaseTermination); + try { + delegate.close(); + } catch (SQLException ignored) { + } + }); + return null; + } + if ("close".equals(method.getName())) { + awaitUninterruptibly(releaseTermination); + delegate.close(); + closed.set(true); + return null; + } + if ("isClosed".equals(method.getName())) { + return closed.get(); + } + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException error) { + throw error.getCause(); + } + } + ); + } + + private static Connection closeFailingConnection(Connection delegate, AtomicInteger closeAttempts) { + return (Connection) Proxy.newProxyInstance( + Connection.class.getClassLoader(), + new Class[] {Connection.class}, + (proxy, method, args) -> { + if ("close".equals(method.getName())) { + closeAttempts.incrementAndGet(); + throw new SQLException("simulated close failure"); + } + if ("isClosed".equals(method.getName())) { + return false; + } + try { + return method.invoke(delegate, args); + } catch (InvocationTargetException error) { + throw error.getCause(); + } + } + ); + } + + private static void awaitUninterruptibly(CountDownLatch latch) { + boolean interrupted = false; + while (true) { + try { + latch.await(); + break; + } catch (InterruptedException ignored) { + interrupted = true; + } + } + if (interrupted) { + Thread.currentThread().interrupt(); + } + } + + private static T futureFailure( + Future future, + Class expectedType, + long timeout, + TimeUnit unit + ) throws Exception { + try { + future.get(timeout, unit); + fail("Expected " + expectedType.getSimpleName()); + throw new IllegalStateException("unreachable"); + } catch (ExecutionException error) { + assertTrue(expectedType.isInstance(error.getCause()), () -> "Unexpected failure: " + error.getCause()); + return expectedType.cast(error.getCause()); + } + } + + private static void awaitCount(AtomicInteger counter, int expected) throws InterruptedException { + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(5); + while (counter.get() < expected && System.nanoTime() < deadline) { + Thread.sleep(10L); + } + assertTrue(counter.get() >= expected, () -> "counter remained at " + counter.get()); + } + + private static void awaitPhysicalConnectionCount( + JdbcConnectionPoolRegistry registry, + int expected + ) throws InterruptedException { + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(5); + while (registry.activePhysicalConnectionCount() != expected && System.nanoTime() < deadline) { + Thread.sleep(10L); + } + assertEquals(expected, registry.activePhysicalConnectionCount()); + } + + private static String throwableText(Throwable error) { + StringBuilder text = new StringBuilder(); + Throwable current = error; + while (current != null) { + text.append(current).append('\n'); + current = current.getCause(); + } + return text.toString(); } private static String h2Url(String prefix) { diff --git a/agents/docs/agent-protocol-v2.md b/agents/docs/agent-protocol-v2.md index bd7da125e..9ca6e7dcb 100644 --- a/agents/docs/agent-protocol-v2.md +++ b/agents/docs/agent-protocol-v2.md @@ -4,7 +4,7 @@ Protocol v2 allows one Agent process to serve multiple isolated database session ## Session lifecycle -- `open_session` creates one logical database session. Parameters contain the normal connection fields plus `agentSessionId`. +- `open_session` creates one logical database session. Parameters contain the normal connection fields plus `agentSessionId` and an optional `sessionRole`. - Every connection-scoped RPC contains `agentSessionId`. - `validate_session` validates and, where supported, reconnects only that session. - `cancel_session` cancels active statements and cursor fetches for only that session; other sessions in the runtime continue normally. @@ -13,6 +13,8 @@ Protocol v2 allows one Agent process to serve multiple isolated database session `agentSessionId` identifies a logical database connection. Existing `sessionId` fields remain pagination cursor identifiers and must not be used as logical connection identifiers. +`sessionRole` is `workload` by default. DBX sends `metadata` for object-tree, completion, and other read-only metadata sessions. New runtimes use this role to preserve metadata checkout capacity; older runtimes may ignore the field. + ## Concurrency Requests for different sessions may execute concurrently. Requests for the same session are serialized because connection state, transactions, schema changes, and driver connections are not generally safe for concurrent use. JSON-RPC responses may be returned out of order and are correlated by request `id`. @@ -31,6 +33,22 @@ Etcd and ZooKeeper retain the legacy path because they use the key-value Agent p A runtime accepts at most 256 logical sessions. Closing the final session starts a 30-second grace period before the process exits, preventing rapid tab open/close cycles from repeatedly starting a runtime. Process EOF fails all pending requests; the failed runtime is removed from reuse and recreated on demand. Connection validation and reconnect operate on a single logical session. +JSON-RPC failures may include structured recovery data: + +```json +{ + "category": "timeout|canceled|connection|protocol|resource|sql", + "retryable": false, + "sessionDisposition": "keep|quarantine|replace_runtime", + "agentSessionId": "optional-session-id", + "stage": "checkout|connect|validate|execute|fetch|cancel|close" +} +``` + +`keep` preserves the logical session, `quarantine` removes only that session from routing, and `replace_runtime` requires DBX to atomically remove every pool sharing the runtime before terminating it. Agent code reports the disposition but must not independently terminate a shared runtime because it does not own DBX routing state. Temporary workload checkout backpressure uses `category=resource`, `retryable=true`, and `sessionDisposition=keep`; only unrecoverable runtime or cleanup saturation requests `replace_runtime`. + +The complete JDBC pool checkout runs under a bounded runtime executor, including HikariCP idle-connection validation, physical connection creation, and driver setup. Workload admission, the runtime-wide physical connection budget, physical creation, and checkout consume one absolute deadline rather than restarting the timeout at each stage. Connection return, eviction, and physical close use separate bounded executors so they cannot deadlock checkout or creation. If a driver call outlives its boundary, or cleanup cannot confirm the physical connection state, the connection identity is poisoned and returns `category=resource` with `sessionDisposition=replace_runtime` on the current or next checkout. A late connection must be evicted and closed instead of published, and DBX must not replay the timed-out user operation automatically. + ## Driver author guidance Use `MultiSessionJsonRpcServer(YourAgent::new)` for Java SQL Agents so each logical session receives a new `DatabaseAgent` with isolated connection state. The shared runtime owns the physical JDBC pools. Do not store connection, statement, cursor, transaction, or schema state in static mutable fields. Use the session execution context for paged query resources. Native Agents must provide equivalent per-session state and synchronized stdout writes. diff --git a/agents/drivers/rabbitmq/integration_test.go b/agents/drivers/rabbitmq/integration_test.go index 2ad2c8d21..0dd525e23 100644 --- a/agents/drivers/rabbitmq/integration_test.go +++ b/agents/drivers/rabbitmq/integration_test.go @@ -103,12 +103,20 @@ func TestRabbitMQIntegration(t *testing.T) { t.Fatalf("unexpected messages %#v", messages) } - stats, err := service.getTopicStats(jsonObject{"topic": queue, "virtual_host": vhost}) - if err != nil { - t.Fatal(err) - } - if stats.(jsonObject)["totalMessages"] != int64(1) { - t.Fatalf("unexpected stats %#v", stats) + var stats any + statsDeadline := time.Now().Add(10 * time.Second) + for { + stats, err = service.getTopicStats(jsonObject{"topic": queue, "virtual_host": vhost}) + if err == nil && stats.(jsonObject)["totalMessages"] == int64(1) { + break + } + if time.Now().After(statsDeadline) { + if err != nil { + t.Fatal(err) + } + t.Fatalf("unexpected stats %#v", stats) + } + time.Sleep(250 * time.Millisecond) } config, err := service.getTopicConfig(jsonObject{"topic": queue, "virtual_host": vhost}) if err != nil { diff --git a/apps/desktop/src/stores/__tests__/connectionStore.metadataLoading.spec.ts b/apps/desktop/src/stores/__tests__/connectionStore.metadataLoading.spec.ts index 0ec770ab6..34bb77f19 100644 --- a/apps/desktop/src/stores/__tests__/connectionStore.metadataLoading.spec.ts +++ b/apps/desktop/src/stores/__tests__/connectionStore.metadataLoading.spec.ts @@ -721,6 +721,261 @@ describe("connectionStore metadata loading", () => { expect(store.treeNodes[0]?.children?.[0]?.children?.map((node) => node.label)).toEqual(["public", "tree.extensions"]); }); + it("preserves the last successful tree snapshot when a forced metadata refresh fails", async () => { + const listSchemaInfos = vi.fn().mockRejectedValue(new Error("Agent RPC call timed out (5s)")); + const deleteSchemaCachePrefix = vi.fn().mockResolvedValue(undefined); + + vi.doMock("@/lib/backend/tauriRuntime", () => ({ isTauriRuntime: () => false })); + vi.doMock("@/lib/backend/api", () => ({ + checkConnectionHealth: vi.fn().mockResolvedValue(undefined), + deleteSchemaCachePrefix, + listInstalledAgents: vi.fn().mockResolvedValue([]), + listSchemaInfos, + loadSchemaCache: vi.fn().mockResolvedValue(null), + saveConnections: vi.fn().mockResolvedValue(undefined), + saveSchemaCache: vi.fn().mockResolvedValue(undefined), + saveSidebarLayout: vi.fn().mockResolvedValue(undefined), + })); + + const { useConnectionStore } = await import("@/stores/connectionStore"); + const store = useConnectionStore(); + const connection = postgresConnection(); + const previousSchema: TreeNode = { + id: `${connection.id}:app:public`, + label: "public", + type: "schema", + connectionId: connection.id, + database: "app", + schema: "public", + isExpanded: false, + children: [], + }; + const databaseNode: TreeNode = { + id: `${connection.id}:app`, + label: "app", + type: "database", + connectionId: connection.id, + database: "app", + isExpanded: true, + children: [previousSchema], + }; + store.connections = [connection]; + store.connectedIds.add(connection.id); + store.treeNodes = [ + { + id: connection.id, + label: connection.name, + type: "connection", + connectionId: connection.id, + isExpanded: true, + children: [databaseNode], + }, + ]; + + await expect(store.refreshTreeNode(databaseNode)).rejects.toThrow("Agent RPC call timed out (5s)"); + + expect(databaseNode.children).toEqual([previousSchema]); + expect(databaseNode.isExpanded).toBe(true); + expect(store.connectionErrors[connection.id]).toBe("Agent RPC call timed out (5s)"); + expect(deleteSchemaCachePrefix).toHaveBeenCalledWith("pg-1:app:"); + }); + + it("does not let an older refresh resume after a newer refresh succeeds", async () => { + let resolveOlderMetadata!: (value: Array<{ name: string; comment: null }>) => void; + const olderMetadata = new Promise>((resolve) => { + resolveOlderMetadata = resolve; + }); + const deleteSchemaCachePrefix = vi.fn().mockResolvedValue(undefined); + const listSchemaInfos = vi + .fn() + .mockImplementationOnce(() => olderMetadata) + .mockResolvedValue([{ name: "latest", comment: null }]); + + vi.doMock("@/lib/backend/tauriRuntime", () => ({ isTauriRuntime: () => false })); + vi.doMock("@/lib/backend/api", () => ({ + checkConnectionHealth: vi.fn().mockResolvedValue(undefined), + deleteSchemaCachePrefix, + listInstalledAgents: vi.fn().mockResolvedValue([]), + listSchemaInfos, + loadSchemaCache: vi.fn().mockResolvedValue(null), + saveConnections: vi.fn().mockResolvedValue(undefined), + saveSchemaCache: vi.fn().mockResolvedValue(undefined), + saveSidebarLayout: vi.fn().mockResolvedValue(undefined), + })); + + const { useConnectionStore } = await import("@/stores/connectionStore"); + const store = useConnectionStore(); + const connection = postgresConnection(); + const databaseNode: TreeNode = { + id: `${connection.id}:app`, + label: "app", + type: "database", + connectionId: connection.id, + database: "app", + isExpanded: true, + children: [], + }; + store.connections = [connection]; + store.connectedIds.add(connection.id); + store.treeNodes = [ + { + id: connection.id, + label: connection.name, + type: "connection", + connectionId: connection.id, + isExpanded: true, + children: [databaseNode], + }, + ]; + + const olderRefresh = store.refreshTreeNode(databaseNode); + await vi.waitFor(() => expect(listSchemaInfos).toHaveBeenCalledTimes(1)); + await store.refreshTreeNode(databaseNode); + resolveOlderMetadata([{ name: "stale", comment: null }]); + await olderRefresh; + + expect(listSchemaInfos).toHaveBeenCalledTimes(2); + expect(databaseNode.children?.map((node) => node.label)).toEqual(["latest", "tree.extensions"]); + }); + + it("does not let an older refresh failure overwrite a newer successful refresh", async () => { + let rejectOlderMetadata!: (reason: Error) => void; + const olderMetadata = new Promise>((_, reject) => { + rejectOlderMetadata = reject; + }); + const listSchemaInfos = vi + .fn() + .mockImplementationOnce(() => olderMetadata) + .mockResolvedValue([{ name: "latest", comment: null }]); + + vi.doMock("@/lib/backend/tauriRuntime", () => ({ isTauriRuntime: () => false })); + vi.doMock("@/lib/backend/api", () => ({ + checkConnectionHealth: vi.fn().mockResolvedValue(undefined), + deleteSchemaCachePrefix: vi.fn().mockResolvedValue(undefined), + listInstalledAgents: vi.fn().mockResolvedValue([]), + listSchemaInfos, + loadSchemaCache: vi.fn().mockResolvedValue(null), + saveConnections: vi.fn().mockResolvedValue(undefined), + saveSchemaCache: vi.fn().mockResolvedValue(undefined), + saveSidebarLayout: vi.fn().mockResolvedValue(undefined), + })); + + const { useConnectionStore } = await import("@/stores/connectionStore"); + const store = useConnectionStore(); + const connection = postgresConnection(); + const databaseNode: TreeNode = { + id: `${connection.id}:app`, + label: "app", + type: "database", + connectionId: connection.id, + database: "app", + isExpanded: true, + children: [], + }; + store.connections = [connection]; + store.connectedIds.add(connection.id); + store.treeNodes = [ + { + id: connection.id, + label: connection.name, + type: "connection", + connectionId: connection.id, + isExpanded: true, + children: [databaseNode], + }, + ]; + + const olderRefresh = store.refreshTreeNode(databaseNode); + await vi.waitFor(() => expect(listSchemaInfos).toHaveBeenCalledTimes(1)); + await store.refreshTreeNode(databaseNode); + rejectOlderMetadata(new Error("connection closed")); + await expect(olderRefresh).rejects.toThrow("connection closed"); + + expect(databaseNode.children?.map((node) => node.label)).toEqual(["latest", "tree.extensions"]); + expect(store.connectionErrors[connection.id]).toBeUndefined(); + expect(store.connectedIds.has(connection.id)).toBe(true); + }); + + it("does not restore a pre-disconnect snapshot into a same-id reconnected node", async () => { + let rejectMetadata!: (reason: Error) => void; + const pendingMetadata = new Promise>((_, reject) => { + rejectMetadata = reject; + }); + const listSchemaInfos = vi.fn().mockReturnValue(pendingMetadata); + + vi.doMock("@/lib/backend/tauriRuntime", () => ({ isTauriRuntime: () => false })); + vi.doMock("@/lib/backend/api", () => ({ + checkConnectionHealth: vi.fn().mockResolvedValue(undefined), + deleteSchemaCachePrefix: vi.fn().mockResolvedValue(undefined), + disconnectDb: vi.fn().mockResolvedValue(undefined), + listInstalledAgents: vi.fn().mockResolvedValue([]), + listSchemaInfos, + loadSchemaCache: vi.fn().mockResolvedValue(null), + saveConnections: vi.fn().mockResolvedValue(undefined), + saveSchemaCache: vi.fn().mockResolvedValue(undefined), + saveSidebarLayout: vi.fn().mockResolvedValue(undefined), + })); + + const { useConnectionStore } = await import("@/stores/connectionStore"); + const store = useConnectionStore(); + const connection = postgresConnection(); + const databaseId = `${connection.id}:app`; + const staleSchema: TreeNode = { + id: `${databaseId}:stale`, + label: "stale", + type: "schema", + connectionId: connection.id, + database: "app", + schema: "stale", + children: [], + }; + const databaseNode: TreeNode = { + id: databaseId, + label: "app", + type: "database", + connectionId: connection.id, + database: "app", + isExpanded: true, + children: [staleSchema], + }; + const connectionNode: TreeNode = { + id: connection.id, + label: connection.name, + type: "connection", + connectionId: connection.id, + isExpanded: true, + children: [databaseNode], + }; + store.connections = [connection]; + store.connectedIds.add(connection.id); + store.treeNodes = [connectionNode]; + + const staleRefresh = store.refreshTreeNode(databaseNode); + await vi.waitFor(() => expect(listSchemaInfos).toHaveBeenCalledTimes(1)); + await store.disconnect(connection.id); + + const freshSchema: TreeNode = { + id: `${databaseId}:fresh`, + label: "fresh", + type: "schema", + connectionId: connection.id, + database: "app", + schema: "fresh", + children: [], + }; + const reconnectedDatabaseNode: TreeNode = { + ...databaseNode, + children: [freshSchema], + }; + connectionNode.children = [reconnectedDatabaseNode]; + store.connectedIds.add(connection.id); + + rejectMetadata(new Error("disconnected refresh")); + await expect(staleRefresh).rejects.toThrow("disconnected refresh"); + + expect(reconnectedDatabaseNode.children?.map((child) => child.label)).toEqual(["fresh"]); + }); + it.each(["opengauss", "kingbase"] as const)("reloads %s sidebar schemas when system visibility changes", async (dbType) => { const listSchemaInfos = vi.fn().mockResolvedValue([ { name: "information_schema", comment: null }, diff --git a/apps/desktop/src/stores/connectionStore.ts b/apps/desktop/src/stores/connectionStore.ts index 43e68ddd7..3c21e71fc 100644 --- a/apps/desktop/src/stores/connectionStore.ts +++ b/apps/desktop/src/stores/connectionStore.ts @@ -419,6 +419,8 @@ export const useConnectionStore = defineStore("connection", () => { const connectionGroupPaths = computed(() => buildConnectionGroupPathMap(sidebarLayout.value)); let layoutPersistTimer: ReturnType | null = null; const staleTreeRefreshIds = new Set(); + const activeTreeRefreshGenerations = new Map(); + let nextTreeRefreshGeneration = 0; const metadataLoadCoordinator = new MetadataLoadCoordinator((event) => { console.debug("[DBX][metadata-load:coordinator]", event); }); @@ -913,7 +915,8 @@ export const useConnectionStore = defineStore("connection", () => { } // Metadata loaders keep this internal: match connection-loss errors before recording generic errors. - function recordMetadataLoadError(connectionId: string, error: unknown) { + function recordMetadataLoadError(connectionId: string, error: unknown, load?: TreeNodeLoadHandle) { + if (load && !load.isCurrent()) return; if (recordConnectionLostError(connectionId, error)) return; recordConnectionError(connectionId, error); } @@ -2976,7 +2979,7 @@ export const useConnectionStore = defineStore("connection", () => { if (cacheHit) { // Render the last known metadata immediately; network validation and refresh // continue in the background so opening a connection never waits on them. - void loadDatabases(connectionId, { ...options, force: true }).catch((error) => recordMetadataLoadError(connectionId, error)); + void loadDatabases(connectionId, { ...options, force: true }).catch(() => undefined); return; } } @@ -3053,7 +3056,7 @@ export const useConnectionStore = defineStore("connection", () => { let dorisCatalogs: CatalogInfo[] | null = null; if (connectionIsDorisFamilyCatalogCapable(config)) { dorisCatalogs = await withMetadataLoadTimeout(connectionId, api.listDorisCatalogs(connectionId), "catalogs").catch((error: unknown) => { - recordMetadataLoadError(connectionId, error); + recordMetadataLoadError(connectionId, error, load); return null; }); } @@ -3134,7 +3137,7 @@ export const useConnectionStore = defineStore("connection", () => { if (liveNode) liveNode.isExpanded = true; if (options?.force) void loadSidebarDatabaseStorage(connectionId, { force: true }); } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3203,7 +3206,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3258,7 +3261,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3295,7 +3298,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3349,7 +3352,7 @@ export const useConnectionStore = defineStore("connection", () => { const liveNode = treeNodeLoadTarget(load); if (liveNode) liveNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3393,7 +3396,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3460,7 +3463,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3514,7 +3517,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3550,7 +3553,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3586,7 +3589,7 @@ export const useConnectionStore = defineStore("connection", () => { setChildren(targetNode, isMilvus && database ? collectionChildren : withSavedSqlRoot(connectionId, collectionChildren, targetNode)); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3631,7 +3634,7 @@ export const useConnectionStore = defineStore("connection", () => { setChildren(targetNode, children); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3708,7 +3711,7 @@ export const useConnectionStore = defineStore("connection", () => { await savePersistedTreeChildren(cacheKey, children); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3747,7 +3750,7 @@ export const useConnectionStore = defineStore("connection", () => { await savePersistedTreeChildren(cacheKey, children); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3783,7 +3786,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3821,7 +3824,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3866,7 +3869,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(node.connectionId, e); + recordMetadataLoadError(node.connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -3922,7 +3925,7 @@ export const useConnectionStore = defineStore("connection", () => { setChildren(targetNode, databaseNodes); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4004,7 +4007,7 @@ export const useConnectionStore = defineStore("connection", () => { } targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4031,7 +4034,7 @@ export const useConnectionStore = defineStore("connection", () => { const nodeId = schema ? `${connectionId}:${database}:${schema}` : `${connectionId}:${database}`; const cacheKey = schemaCacheKey(connectionId, database, schema || "", "objects-simple-v6"); if (await hydrateTreeNodeFromCache(findNode(treeNodes.value, nodeId), cacheKey)) { - void loadTables(connectionId, database, schema, { ...options, force: true }).catch((error) => recordMetadataLoadError(connectionId, error)); + void loadTables(connectionId, database, schema, { ...options, force: true }).catch(() => undefined); return; } } @@ -4136,7 +4139,7 @@ export const useConnectionStore = defineStore("connection", () => { }); } } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4162,9 +4165,7 @@ export const useConnectionStore = defineStore("connection", () => { }); if (!options?.force && !searchFilter && !options?.sidebarTableSearchParentId && !tableNameFilterForScope) { if (await hydrateTreeNodeFromCache(node, objectGroupCacheKey(node))) { - void loadObjectGroupChildren(node, { ...options, force: true }).catch((error) => { - if (node.connectionId) recordMetadataLoadError(node.connectionId, error); - }); + void loadObjectGroupChildren(node, { ...options, force: true }).catch(() => undefined); return; } } @@ -4259,7 +4260,7 @@ export const useConnectionStore = defineStore("connection", () => { } targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(node.connectionId, e); + recordMetadataLoadError(node.connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4383,7 +4384,7 @@ export const useConnectionStore = defineStore("connection", () => { await savePersistedTreeChildren(objectGroupCacheKey(targetParent), nextChildren); targetParent.isExpanded = true; } catch (e) { - recordMetadataLoadError(parentConnectionId, e); + recordMetadataLoadError(parentConnectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4418,7 +4419,7 @@ export const useConnectionStore = defineStore("connection", () => { targetNode.objectCount = children.length; targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4553,7 +4554,7 @@ export const useConnectionStore = defineStore("connection", () => { const finishedParent = treeNodeLoadTarget(load); if (finishedParent) finishedParent.isExpanded = true; } catch (e) { - recordMetadataLoadError(parent.connectionId, e); + recordMetadataLoadError(parent.connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4795,7 +4796,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4836,7 +4837,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4881,7 +4882,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4923,7 +4924,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4954,7 +4955,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -4985,7 +4986,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -5016,7 +5017,7 @@ export const useConnectionStore = defineStore("connection", () => { ); targetNode.isExpanded = true; } catch (e) { - recordMetadataLoadError(connectionId, e); + recordMetadataLoadError(connectionId, e, load); throw e; } finally { finishTreeNodeLoad(load); @@ -5114,19 +5115,22 @@ export const useConnectionStore = defineStore("connection", () => { } } - async function restoreExpandedChildren(node: TreeNode, expandedIds: Set, options?: LoadTreeOptions) { + async function restoreExpandedChildren(node: TreeNode, expandedIds: Set, options?: LoadTreeOptions, isCurrent: () => boolean = () => true) { + if (!isCurrent()) return; if (!node.children) return; for (const child of node.children) { + if (!isCurrent()) return; if (!expandedIds.has(child.id)) continue; await loadTreeNodeChildren(child, options); - await restoreExpandedChildren(child, expandedIds, options); + if (!isCurrent()) return; + await restoreExpandedChildren(child, expandedIds, options, isCurrent); } } async function refreshTreeNode(node: TreeNode) { invalidateMetadataCachesForNode(node); if (objectTypesForGroupNode(node.type)) { - clearLoadedChildrenCache(node.id); + clearLoadedChildrenCache(node.id, { deletePersisted: false }); await loadObjectGroupChildren(node, { force: true }); return; } @@ -5144,24 +5148,44 @@ export const useConnectionStore = defineStore("connection", () => { const previousChildren = node.children; const previousHiddenChildren = node.hiddenChildren; const previousObjectCount = node.objectCount; + const previousExpanded = node.isExpanded; const previousLoadedIds = [...loadedTreeNodeChildrenIds.value].filter((id) => id === node.id || id.startsWith(`${node.id}:`)); const previousConfirmedEmptyIds = [...confirmedEmptyTreeNodeIds.value].filter((id) => id === node.id || id.startsWith(`${node.id}:`)); - await clearPersistedTreeCacheForNode(node); - clearLoadedChildrenCache(node.id); - if (node.type !== "connection-group") { - node.children = []; - } + const connectionRevision = node.connectionId ? connectionStateRevision(node.connectionId) : undefined; + const refreshGeneration = ++nextTreeRefreshGeneration; + activeTreeRefreshGenerations.set(node.id, refreshGeneration); + const ownsRefreshGeneration = () => activeTreeRefreshGenerations.get(node.id) === refreshGeneration; + const isCurrentRefresh = () => ownsRefreshGeneration() && (!node.connectionId || connectionStateRevision(node.connectionId) === connectionRevision); try { + await clearPersistedTreeCacheForNode(node); + if (!isCurrentRefresh()) return; + clearLoadedChildrenCache(node.id); + if (node.type !== "connection-group") { + node.children = []; + } await loadTreeNodeChildren(node, { force: true }); - await restoreExpandedChildren(node, expandedIds, { force: true }); + if (isCurrentRefresh()) { + await restoreExpandedChildren(node, expandedIds, { force: true }, isCurrentRefresh); + } } catch (error) { - node.children = previousChildren; - node.hiddenChildren = previousHiddenChildren; - node.objectCount = previousObjectCount; - clearLoadedChildrenCache(node.id, { deletePersisted: false }); - for (const id of previousLoadedIds) loadedTreeNodeChildrenIds.value.add(id); - for (const id of previousConfirmedEmptyIds) confirmedEmptyTreeNodeIds.value.add(id); + // A stale failure must never overwrite a newer successful (including empty) result. + if (isCurrentRefresh()) { + const target = treeNodeInSidebarTree(node); + if (target) { + target.children = previousChildren; + target.hiddenChildren = previousHiddenChildren; + target.objectCount = previousObjectCount; + target.isExpanded = previousExpanded; + clearLoadedChildrenCache(target.id, { deletePersisted: false }); + for (const id of previousLoadedIds) loadedTreeNodeChildrenIds.value.add(id); + for (const id of previousConfirmedEmptyIds) confirmedEmptyTreeNodeIds.value.add(id); + } + } throw error; + } finally { + if (ownsRefreshGeneration()) { + activeTreeRefreshGenerations.delete(node.id); + } } } diff --git a/crates/dbx-core/src/agent_connection.rs b/crates/dbx-core/src/agent_connection.rs index a5f8c0bc4..de94bde59 100644 --- a/crates/dbx-core/src/agent_connection.rs +++ b/crates/dbx-core/src/agent_connection.rs @@ -7,6 +7,22 @@ const OCEANBASE_ORACLE_COMPATIBLE_OJDBC_VERSION_KEY: &str = "compatibleOjdbcVers const OCEANBASE_ORACLE_COMPATIBLE_OJDBC_VERSION_PARAM: &str = "compatibleOjdbcVersion=8"; const ZOOKEEPER_MIN_CONNECTION_TIMEOUT_MS: u64 = 15_000; +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum AgentSessionRole { + #[default] + Workload, + Metadata, +} + +impl AgentSessionRole { + fn as_str(self) -> &'static str { + match self { + Self::Workload => "workload", + Self::Metadata => "metadata", + } + } +} + fn agent_jdbc_driver_class(config: &ConnectionConfig) -> &str { let driver_class = config.jdbc_driver_class.as_deref().unwrap_or(""); if config.db_type == DatabaseType::H2 @@ -19,6 +35,16 @@ fn agent_jdbc_driver_class(config: &ConnectionConfig) -> &str { } pub fn agent_connect_params(config: &ConnectionConfig, host: &str, port: u16, database: &str) -> serde_json::Value { + agent_connect_params_with_role(config, host, port, database, AgentSessionRole::Workload) +} + +pub fn agent_connect_params_with_role( + config: &ConnectionConfig, + host: &str, + port: u16, + database: &str, + session_role: AgentSessionRole, +) -> serde_json::Value { let agent_database = if config.db_type == DatabaseType::MongoDb { mongo_agent_database(config, database) } else if matches!(config.db_type, DatabaseType::Oracle | DatabaseType::OceanbaseOracle) { @@ -83,6 +109,7 @@ pub fn agent_connect_params(config: &ConnectionConfig, host: &str, port: u16, da "informix_server": config.informix_server, "jdbc_driver_class": agent_jdbc_driver_class(config), "jdbc_driver_paths": &config.jdbc_driver_paths, + "sessionRole": session_role.as_str(), }); if config.db_type == DatabaseType::ZooKeeper { params["connection_timeout_ms"] = serde_json::json!( @@ -662,6 +689,24 @@ mod tests { } } + #[test] + fn agent_connect_params_default_to_workload_session_role() { + let params = agent_connect_params(&config(DatabaseType::H2, Some("test")), "127.0.0.1", 9092, "test"); + assert_eq!(params["sessionRole"], "workload"); + } + + #[test] + fn metadata_agent_connect_params_include_metadata_session_role() { + let params = agent_connect_params_with_role( + &config(DatabaseType::H2, Some("test")), + "127.0.0.1", + 9092, + "test", + AgentSessionRole::Metadata, + ); + assert_eq!(params["sessionRole"], "metadata"); + } + #[test] fn mongodb_database_falls_back_to_uri_database() { let mut cfg = config(DatabaseType::MongoDb, None); diff --git a/crates/dbx-core/src/agent_manager.rs b/crates/dbx-core/src/agent_manager.rs index ea99e1fa3..2ab560add 100644 --- a/crates/dbx-core/src/agent_manager.rs +++ b/crates/dbx-core/src/agent_manager.rs @@ -894,7 +894,7 @@ impl AgentManager { agent_session_id: String, connect_params: serde_json::Value, connect_timeout: std::time::Duration, - ) -> Result { + ) -> Result { crate::agent_runtime::spawn_shared_connection_client( self, db_type, diff --git a/crates/dbx-core/src/agent_runtime.rs b/crates/dbx-core/src/agent_runtime.rs index b6090f7e0..5e5d9154e 100644 --- a/crates/dbx-core/src/agent_runtime.rs +++ b/crates/dbx-core/src/agent_runtime.rs @@ -3,9 +3,22 @@ use std::time::Duration; use crate::agent_manager::{AgentManager, DEFAULT_JRE_KEY}; use crate::database_capabilities; -use crate::db::agent_driver::{AgentDriverClient, AgentMethod, AgentRuntimeClient}; +use crate::db::agent_driver::{ + agent_session_disposition, AgentDriverClient, AgentMethod, AgentRuntimeClient, AgentSessionDisposition, +}; use crate::models::connection::DatabaseType; +pub struct SharedConnectionOpenError { + pub(crate) message: String, + pub(crate) runtime: Option>, +} + +impl From for SharedConnectionOpenError { + fn from(message: String) -> Self { + Self { message, runtime: None } + } +} + pub fn db_type_to_agent_key(db_type: &DatabaseType, driver_profile: Option<&str>) -> Option<&'static str> { database_capabilities::agent_key(db_type, driver_profile) } @@ -64,7 +77,7 @@ pub async fn spawn_shared_connection_client( agent_session_id: String, connect_params: serde_json::Value, connect_timeout: Duration, -) -> Result { +) -> Result { let keys = runtime_agent_key_candidates(db_type, driver_profile) .ok_or_else(|| format!("{:?} is not an agent-driven database type", db_type))?; let key = first_installed_agent_key(manager, &keys).unwrap_or(keys[0]); @@ -103,12 +116,22 @@ pub async fn spawn_shared_connection_client( .call::(AgentMethod::OpenSession.as_str(), session_params, Some(connect_timeout), None) .await { + let open_error = shared_connection_open_error(err, runtime.clone()); forget_unused_runtime_after_failed_open(manager, &runtime_key, &runtime_cell, &runtime).await; - return Err(err); + return Err(open_error); } Ok(AgentDriverClient::shared_session(runtime, agent_session_id)) } +fn shared_connection_open_error( + message: String, + runtime: std::sync::Arc, +) -> SharedConnectionOpenError { + let runtime = + (agent_session_disposition(&message) == Some(AgentSessionDisposition::ReplaceRuntime)).then_some(runtime); + SharedConnectionOpenError { message, runtime } +} + async fn forget_unused_runtime_after_failed_open( manager: &AgentManager, runtime_key: &str, @@ -316,8 +339,9 @@ for line in sys.stdin: "#, ) .unwrap(); + let python = if cfg!(windows) { "python" } else { "python3" }; let runtime = AgentRuntimeClient::spawn( - crate::db::agent_driver::AgentLaunchSpec::new("python3") + crate::db::agent_driver::AgentLaunchSpec::new(python) .with_args([script_path.to_string_lossy().to_string()]), "test", ) @@ -385,6 +409,21 @@ for line in sys.stdin: let _ = std::fs::remove_file(script_path); } + #[tokio::test] + async fn replace_runtime_open_error_reports_runtime_without_killing_it_directly() { + let (manager, _cell, runtime, script_path) = test_shared_runtime("replace-runtime-open-error").await; + let error = "Agent RPC error (-1): capacity exhausted\nDBX_AGENT_ERROR_DATA:{\"category\":\"resource\",\"sessionDisposition\":\"replace_runtime\"}"; + + let open_error = shared_connection_open_error(error.to_string(), runtime.clone()); + + assert!(open_error.runtime.as_ref().is_some_and(|failed| std::sync::Arc::ptr_eq(failed, &runtime))); + assert!(!runtime.is_failed()); + + runtime.kill(); + drop(manager); + 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; diff --git a/crates/dbx-core/src/connection.rs b/crates/dbx-core/src/connection.rs index c9e42a0b5..297640606 100644 --- a/crates/dbx-core/src/connection.rs +++ b/crates/dbx-core/src/connection.rs @@ -9,9 +9,10 @@ use mysql_async::prelude::Queryable; use mysql_async::Row as MysqlRow; use crate::agent_connection::{ - agent_connect_params, h2_file_path_from_jdbc_url, is_h2_file_connection, mongo_legacy_error_with_auth_hint, - mongo_uses_legacy_driver, oracle_alternate_connect_config_labels, oracle_alternate_connect_configs, - oracle_error_with_driver_hint, should_retry_mongo_with_legacy_driver, trino_like_jdbc_connection_string, + agent_connect_params, agent_connect_params_with_role, h2_file_path_from_jdbc_url, is_h2_file_connection, + mongo_legacy_error_with_auth_hint, mongo_uses_legacy_driver, oracle_alternate_connect_config_labels, + oracle_alternate_connect_configs, oracle_error_with_driver_hint, should_retry_mongo_with_legacy_driver, + trino_like_jdbc_connection_string, AgentSessionRole, }; use crate::agent_manager::{JavaRuntimeMode, DEFAULT_JRE_KEY}; use crate::database_capabilities; @@ -104,7 +105,7 @@ pub enum PoolKind { HBase(db::hbase_driver::HBaseClient), VectorDb(db::vector_driver::VectorClient), InfluxDb(db::influxdb_driver::InfluxdbClient), - Agent(Arc>), + Agent(Arc), ExternalDriver { driver_id: String, config: Arc, @@ -117,8 +118,21 @@ pub enum PoolKind { Nacos, } +impl PoolKind { + pub fn agent(client: db::agent_driver::AgentDriverClient) -> Self { + Self::Agent(Arc::new(db::agent_driver::PooledAgentClient::new(client))) + } + + fn is_available_for_routing(&self) -> bool { + match self { + Self::Agent(client) => client.is_runtime_available(), + _ => true, + } + } +} + enum ConnectionDatabaseInfoSource { - Agent(Arc>), + Agent(Arc), ExternalDriver { config: Arc, session: Arc }, NativeMysql(db::mysql::MySqlPool), NativeHBase(db::hbase_driver::HBaseClient), @@ -280,12 +294,17 @@ pub struct PoolActivityTouch { task_supervisor: TaskSupervisor, } -pub(crate) struct ClientSessionPoolCleanupGuard { - pool_key: String, +#[derive(Clone)] +struct PoolRoutingControl { connections: Arc>>, pool_activity: Arc>>, postgres_cancel_contexts: Arc>>, task_supervisor: TaskSupervisor, +} + +pub(crate) struct ClientSessionPoolCleanupGuard { + pool_key: String, + routing: PoolRoutingControl, armed: bool, } @@ -319,22 +338,222 @@ impl Drop for ClientSessionPoolCleanupGuard { return; } let pool_key = self.pool_key.clone(); - let connections = self.connections.clone(); - let pool_activity = self.pool_activity.clone(); - let postgres_cancel_contexts = self.postgres_cancel_contexts.clone(); - let task_supervisor = self.task_supervisor.clone(); - task_supervisor.stop(&format!("keepalive:{pool_key}")); - task_supervisor.spawn_once(format!("client-session-cleanup:{pool_key}"), move |_| async move { - pool_activity.write().await.remove(&pool_key); - postgres_cancel_contexts.write().await.remove(&pool_key); - let removed = connections.write().await.remove(&pool_key); - if let Some(pool) = removed { - close_pool_kind_with_timeout(pool_key, pool).await; - } + let routing = self.routing.clone(); + routing.stop_keepalive(&pool_key); + let supervisor = routing.task_supervisor.clone(); + supervisor.spawn_once(format!("client-session-cleanup:{pool_key}"), move |_| async move { + routing.detach_pool_by_key(&pool_key, false).await; }); } } +impl PoolRoutingControl { + fn stop_keepalive(&self, pool_key: &str) { + self.task_supervisor.stop(&format!("keepalive:{pool_key}")); + } + + async fn detach_pool_by_key(&self, pool_key: &str, replace_agent_runtime: bool) -> bool { + let removed = { + let mut connections = self.connections.write().await; + let Some(pool) = connections.remove(pool_key) else { + return false; + }; + let mut removed = vec![(pool_key.to_string(), pool)]; + if replace_agent_runtime { + let sibling_keys = shared_runtime_sibling_keys(&connections, &removed[0].1); + for key in sibling_keys { + if let Some(pool) = connections.remove(&key) { + removed.push((key, pool)); + } + } + fail_stop_removed_agent_pool(pool_key, &removed[0].1); + } + removed + }; + + self.finish_detach(removed).await; + true + } + + async fn detach_agent_pool_if_current( + &self, + pool_key: &str, + expected_client: &Arc, + replace_agent_runtime: bool, + ) -> bool { + let removed = { + let mut connections = self.connections.write().await; + let is_current = matches!( + connections.get(pool_key), + Some(PoolKind::Agent(current)) if Arc::ptr_eq(current, expected_client) + ); + if !is_current && !replace_agent_runtime { + return false; + } + if replace_agent_runtime { + let runtime_keys = connections + .iter() + .filter_map(|(key, pool)| match pool { + PoolKind::Agent(client) if expected_client.shares_runtime_with(client) => Some(key.clone()), + _ => None, + }) + .collect::>(); + let removed = runtime_keys + .into_iter() + .filter_map(|key| connections.remove(&key).map(|pool| (key, pool))) + .collect::>(); + if !expected_client.fail_stop() { + log::warn!("Failed to terminate the shared Agent runtime while detaching pool '{pool_key}'"); + } + removed + } else { + let pool = connections.remove(pool_key).expect("current Agent pool must still exist"); + vec![(pool_key.to_string(), pool)] + } + }; + + let detached = !removed.is_empty(); + self.finish_detach(removed).await; + detached + } + + async fn finish_detach(&self, removed: Vec<(String, PoolKind)>) { + for (key, _) in &removed { + self.stop_keepalive(key); + } + { + let mut activity = self.pool_activity.write().await; + let mut cancel_contexts = self.postgres_cancel_contexts.write().await; + for (key, _) in &removed { + activity.remove(key); + cancel_contexts.remove(key); + } + } + self.close_removed_in_background(removed); + } + + async fn close_pool_with_timeout(&self, pool_key: String, pool: PoolKind) { + let agent_client = match &pool { + PoolKind::Agent(client) => Some(client.clone()), + _ => None, + }; + match tokio::time::timeout(Duration::from_secs(POOL_CLOSE_TIMEOUT_SECS), close_pool_kind(pool)).await { + Ok(Ok(())) => {} + Ok(Err(error)) => { + log::warn!("Failed to close connection pool '{pool_key}': {error}"); + if let Some(client) = agent_client.filter(|_| should_replace_agent_runtime(&error)) { + self.replace_runtime_after_close_failure(&pool_key, &client).await; + } + } + Err(_) => { + log::warn!( + "Timed out closing connection pool '{pool_key}' after {POOL_CLOSE_TIMEOUT_SECS}s; replacing a shared Agent runtime when present." + ); + if let Some(client) = agent_client { + self.replace_runtime_after_close_failure(&pool_key, &client).await; + } + } + } + } + + async fn replace_runtime_after_close_failure( + &self, + closed_pool_key: &str, + failed_client: &Arc, + ) { + let removed = { + let mut connections = self.connections.write().await; + let sibling_keys = connections + .iter() + .filter_map(|(key, pool)| match pool { + PoolKind::Agent(client) if failed_client.shares_runtime_with(client) => Some(key.clone()), + _ => None, + }) + .collect::>(); + let removed = sibling_keys + .into_iter() + .filter_map(|key| connections.remove(&key).map(|pool| (key, pool))) + .collect::>(); + if !failed_client.fail_stop() { + log::warn!( + "Detached Agent pool '{closed_pool_key}' after close failure, but its legacy process could not be terminated without the client lock" + ); + } + removed + }; + self.finish_detach(removed).await; + } + + async fn replace_runtime_after_open_failure(&self, failed_runtime: &Arc) { + let removed = { + let mut connections = self.connections.write().await; + let keys = connections + .iter() + .filter_map(|(key, pool)| match pool { + PoolKind::Agent(client) if client.uses_runtime(failed_runtime) => Some(key.clone()), + _ => None, + }) + .collect::>(); + let removed = + keys.into_iter().filter_map(|key| connections.remove(&key).map(|pool| (key, pool))).collect::>(); + failed_runtime.kill(); + removed + }; + self.finish_detach(removed).await; + } + + async fn close_removed(&self, removed: Vec<(String, PoolKind)>) { + for (pool_key, pool) in removed { + self.close_pool_with_timeout(pool_key, pool).await; + } + } + + fn close_removed_in_background(&self, removed: Vec<(String, PoolKind)>) { + if removed.is_empty() { + return; + } + let pool_count = removed.len(); + let routing = self.clone(); + let task_key = format!("pool-close:{}", uuid::Uuid::new_v4()); + if !self.task_supervisor.spawn_once(task_key, move |_| async move { + for (pool_key, pool) in removed { + routing.close_pool_with_timeout(pool_key, pool).await; + } + }) { + log::debug!("Dropped {pool_count} detached pool handle(s) during application shutdown"); + } + } +} + +fn shared_runtime_sibling_keys(connections: &HashMap, source_pool: &PoolKind) -> Vec { + let PoolKind::Agent(source_client) = source_pool else { + return Vec::new(); + }; + connections + .iter() + .filter_map(|(key, pool)| match pool { + PoolKind::Agent(client) if source_client.shares_runtime_with(client) => Some(key.clone()), + _ => None, + }) + .collect() +} + +fn fail_stop_removed_agent_pool(pool_key: &str, pool: &PoolKind) { + let PoolKind::Agent(client) = pool else { + return; + }; + if !client.fail_stop() { + log::warn!( + "Detached busy legacy Agent pool '{pool_key}', but its process cannot be terminated without the client lock" + ); + } +} + +fn should_replace_agent_runtime(error: &str) -> bool { + crate::db::agent_driver::agent_session_disposition(error) + == Some(crate::db::agent_driver::AgentSessionDisposition::ReplaceRuntime) +} + pub fn metadata_connection_config(config: &ConnectionConfig) -> ConnectionConfig { let mut db_config = config.canonicalized(); if database_capabilities::is_metadata_connection_scoped(&db_config.db_type) { @@ -725,6 +944,51 @@ fn mysql_metadata_fallback_url( } impl AppState { + fn pool_routing_control(&self) -> PoolRoutingControl { + PoolRoutingControl { + connections: self.connections.clone(), + pool_activity: self.pool_activity.clone(), + postgres_cancel_contexts: self.postgres_cancel_contexts.clone(), + task_supervisor: self.task_supervisor.clone(), + } + } + + async fn handle_shared_connection_open_error( + &self, + error: crate::agent_runtime::SharedConnectionOpenError, + ) -> String { + if let Some(runtime) = error.runtime.as_ref() { + self.pool_routing_control().replace_runtime_after_open_failure(runtime).await; + } + error.message + } + + async fn spawn_routed_shared_agent_client( + &self, + db_type: &DatabaseType, + driver_profile: Option<&str>, + extra_java_args: &[String], + agent_session_id: String, + connect_params: serde_json::Value, + connect_timeout: Duration, + ) -> Result { + match self + .agent_manager + .spawn_shared_connection_client( + db_type, + driver_profile, + extra_java_args, + agent_session_id, + connect_params, + connect_timeout, + ) + .await + { + Ok(client) => Ok(client), + Err(error) => Err(self.handle_shared_connection_open_error(error).await), + } + } + pub fn new(storage: Storage) -> Self { Self::new_with_plugin_dir(storage, default_plugin_dir()) } @@ -840,7 +1104,7 @@ impl AppState { // Test the submitted form as a fresh session so unsaved ATTACH/init // changes cannot be masked by a pool created from older settings. let pool = self.create_duckdb_pool(config).await?; - close_pool_kind(pool).await; + close_pool_kind(pool).await?; Ok(()) } @@ -957,7 +1221,7 @@ impl AppState { ) .await .map_err(|err| sqlserver_legacy_driver_error(&err))?; - return Ok(PoolKind::Agent(Arc::new(tokio::sync::Mutex::new(client)))); + return Ok(PoolKind::agent(client)); } let client = db::sqlserver::connect_with_port_explicit( @@ -1019,22 +1283,59 @@ impl AppState { pool: PoolKind, config: &ConnectionConfig, wait_for_drain: bool, - ) { + ) -> Result<(), String> { if wait_for_drain { self.wait_for_pool_drain(&pool_key).await; } - self.stop_keepalive_task(&pool_key).await; - self.pool_activity.write().await.insert(pool_key.clone(), PoolActivity::now()); - self.start_keepalive_task(&pool_key, &pool, config).await; - let previous_key = pool_key.clone(); - let previous = self.connections.write().await.insert(pool_key, pool); + let routing = self.pool_routing_control(); + let previous = loop { + let mut connections = self.connections.write().await; + if !pool.is_available_for_routing() { + break Err(pool); + } + let Ok(mut activity) = self.pool_activity.try_write() else { + // Idle reclamation reads activity before routing. Never await that lock while + // holding the routing lock; release and retry to preserve a single lock order. + drop(connections); + tokio::task::yield_now().await; + continue; + }; + // Abort the old probe while the route cannot change underneath us. A failed + // candidate never reaches this point, so the existing route keeps its state. + routing.stop_keepalive(&pool_key); + activity.insert(pool_key.clone(), PoolActivity::now()); + self.start_keepalive_task(&pool_key, &pool, config); + break Ok(connections.insert(pool_key.clone(), pool)); + }; + let previous = match previous { + Ok(previous) => previous, + Err(pool) => { + if let PoolKind::Agent(client) = &pool { + routing.replace_runtime_after_close_failure(&pool_key, client).await; + } + routing.close_pool_with_timeout(pool_key, pool).await; + return Err("Agent runtime is unavailable while publishing the connection pool".to_string()); + } + }; if let Some(pool) = previous { - close_pool_kind_with_timeout(previous_key, pool).await; + routing.close_pool_with_timeout(pool_key.clone(), pool).await; } + let route_is_available = + self.connections.read().await.get(&pool_key).is_some_and(PoolKind::is_available_for_routing); + if !route_is_available { + routing.detach_pool_by_key(&pool_key, true).await; + return Err("Agent runtime is unavailable while publishing the connection pool".to_string()); + } + Ok(()) } - pub async fn insert_connection_pool(&self, pool_key: String, pool: PoolKind, config: &ConnectionConfig) { - self.insert_connection_pool_inner(pool_key, pool, config, true).await; + pub async fn insert_connection_pool( + &self, + pool_key: String, + pool: PoolKind, + config: &ConnectionConfig, + ) -> Result<(), String> { + self.insert_connection_pool_inner(pool_key, pool, config, true).await } pub async fn begin_connection_attempt(&self, connection_id: &str) -> u64 { @@ -1103,11 +1404,10 @@ impl AppState { config: &ConnectionConfig, ) -> Result<(), String> { if let Err(err) = self.ensure_current_connection_attempt(connection_id, Some(attempt)).await { - close_pool_kind_with_timeout(pool_key, pool).await; + self.pool_routing_control().close_pool_with_timeout(pool_key, pool).await; return Err(err); } - self.insert_connection_pool(pool_key, pool, config).await; - Ok(()) + self.insert_connection_pool(pool_key, pool, config).await } async fn discard_stale_connection_attempt_pool( @@ -1123,10 +1423,10 @@ impl AppState { self.mq_registry.drop_connection(connection_id).await; } self.reset_connection_transport_for_config(connection_id, config).await; - close_pool_kind_with_timeout(pool_key, pool).await; + self.pool_routing_control().close_pool_with_timeout(pool_key, pool).await; } - async fn start_keepalive_task(&self, pool_key: &str, pool: &PoolKind, config: &ConnectionConfig) { + fn start_keepalive_task(&self, pool_key: &str, pool: &PoolKind, config: &ConnectionConfig) { let interval_secs = config.keepalive_interval_secs; let mut target = keepalive_target_from_pool(pool, config); if interval_secs == 0 { @@ -1142,9 +1442,8 @@ impl AppState { let key = pool_key.to_string(); let interval = Duration::from_secs(interval_secs.max(1)); let timeout = Duration::from_secs(config.effective_connect_timeout_secs().max(1)); + let routing = self.pool_routing_control(); let connections = self.connections.clone(); - let pool_activity = self.pool_activity.clone(); - let cancel_contexts = self.postgres_cancel_contexts.clone(); let running_queries = self.running_queries.clone(); self.task_supervisor.spawn_replace(format!("keepalive:{pool_key}"), move |shutdown| async move { loop { @@ -1163,12 +1462,15 @@ impl AppState { Ok(Ok(())) => {} Ok(Err(err)) => { log::warn!("Connection keepalive failed for '{key}': {err}; invalidating pool"); - let removed = remove_keepalive_pool_if_current(&connections, &key, target).await; - if let Some(pool) = removed { - pool_activity.write().await.remove(&key); - cancel_contexts.write().await.remove(&key); - close_pool_kind_with_timeout(key, pool).await; - } else { + if !detach_keepalive_target_if_current( + &routing, + &connections, + &key, + target, + should_replace_agent_runtime(&err), + ) + .await + { log::debug!("Skipping stale keepalive result for replaced pool '{key}'"); } break; @@ -1178,12 +1480,7 @@ impl AppState { "Connection keepalive timed out for '{key}' after {}s; invalidating pool", timeout.as_secs() ); - let removed = remove_keepalive_pool_if_current(&connections, &key, target).await; - if let Some(pool) = removed { - pool_activity.write().await.remove(&key); - cancel_contexts.write().await.remove(&key); - close_pool_kind_with_timeout(key, pool).await; - } else { + if !detach_keepalive_target_if_current(&routing, &connections, &key, target, false).await { log::debug!("Skipping stale keepalive timeout for replaced pool '{key}'"); } break; @@ -1236,9 +1533,10 @@ impl AppState { self.transaction_sessions.write().await.clear(); let shutdown = async { + let routing = self.pool_routing_control(); tokio::join!( self.task_supervisor.shutdown(deadline), - close_removed_pools(removed_pools), + routing.close_removed(removed_pools), self.tunnels.stop_all_tunnels(), self.proxy_tunnels.stop_all_tunnels(), self.http_tunnels.stop_all_tunnels(), @@ -1274,7 +1572,15 @@ impl AppState { database: Option<&str>, attempt: u64, ) -> Result { - self.get_or_create_pool_for_session_inner(connection_id, database, None, None, Some(attempt)).await + self.get_or_create_pool_for_session_inner( + connection_id, + database, + None, + None, + AgentSessionRole::Workload, + Some(attempt), + ) + .await } pub async fn get_or_create_pool_for_session( @@ -1293,7 +1599,32 @@ impl AppState { catalog: Option<&str>, client_session_id: Option<&str>, ) -> Result { - self.get_or_create_pool_for_session_inner(connection_id, database, catalog, client_session_id, None).await + self.get_or_create_pool_for_session_inner( + connection_id, + database, + catalog, + client_session_id, + AgentSessionRole::Workload, + None, + ) + .await + } + + pub(crate) async fn get_or_create_metadata_pool_for_session( + &self, + connection_id: &str, + database: Option<&str>, + client_session_id: Option<&str>, + ) -> Result { + self.get_or_create_pool_for_session_inner( + connection_id, + database, + None, + client_session_id, + AgentSessionRole::Metadata, + None, + ) + .await } async fn get_or_create_pool_for_session_inner( @@ -1302,6 +1633,7 @@ impl AppState { database: Option<&str>, catalog: Option<&str>, client_session_id: Option<&str>, + session_role: AgentSessionRole, connection_attempt: Option, ) -> Result { let config = { @@ -1314,7 +1646,7 @@ impl AppState { let catalog = catalog.map(str::trim).filter(|value| !value.is_empty()); let base_pool_key = base_pool_key_for_with_catalog(db_type, connection_id, database, catalog, false); - let pool_key = session_scoped_pool_key_for(Some(&config), base_pool_key.clone(), client_session_id); + let pool_key = pool_key_for_session_role(Some(&config), base_pool_key.clone(), client_session_id, session_role); loop { self.wait_for_pool_drain(&pool_key).await; @@ -1482,7 +1814,7 @@ impl AppState { let connect_params = serde_json::json!({ "connection": agent_connect_params(&db_config, &host, port, db_config.effective_database().unwrap_or("")) }); let mut client = self.agent_manager.spawn(&DatabaseType::MongoDb, Some("mongodb-legacy")).await?; client.connect(connect_params).await.map_err(|err| mongo_legacy_error_with_auth_hint(&err))?; - PoolKind::Agent(Arc::new(tokio::sync::Mutex::new(client))) + PoolKind::agent(client) } else { let native_err = match db::mongo_driver::connect(&url, connect_timeout, idle_timeout).await { Ok(client) => match db::mongo_driver::test_connection( @@ -1495,7 +1827,9 @@ impl AppState { Ok(()) => { // Re-check: another task may have created the pool while we were connecting. if self.connections.read().await.contains_key(&pool_key) { - close_pool_kind_with_timeout(pool_key.clone(), PoolKind::MongoDb(client)).await; + self.pool_routing_control() + .close_pool_with_timeout(pool_key.clone(), PoolKind::MongoDb(client)) + .await; return Ok(pool_key); } if let Err(err) = @@ -1511,7 +1845,7 @@ impl AppState { return Err(err); } self.insert_connection_pool(pool_key.clone(), PoolKind::MongoDb(client), &db_config) - .await; + .await?; return Ok(pool_key); } Err(e) => e, @@ -1529,7 +1863,7 @@ impl AppState { mongo_legacy_error_with_auth_hint(&err) ) })?; - PoolKind::Agent(Arc::new(tokio::sync::Mutex::new(client))) + PoolKind::agent(client) } else { return Err(native_err); } @@ -1618,13 +1952,17 @@ impl AppState { PoolKind::Nacos } agent_connection_pool_database_type!() => { - let connect_params = - agent_connect_params(&db_config, &host, port, db_config.effective_database().unwrap_or("")); + let connect_params = agent_connect_params_with_role( + &db_config, + &host, + port, + db_config.effective_database().unwrap_or(""), + session_role, + ); if db_config.db_type != DatabaseType::ZooKeeper { let agent_session_id = uuid::Uuid::new_v4().simple().to_string(); let mut initial_result = self - .agent_manager - .spawn_shared_connection_client( + .spawn_routed_shared_agent_client( &db_config.db_type, db_config.driver_profile.as_deref(), &db_config.agent_java_options, @@ -1644,17 +1982,17 @@ impl AppState { for retry_delay_ms in [150, 350] { tokio::time::sleep(Duration::from_millis(retry_delay_ms)).await; initial_result = self - .agent_manager - .spawn_shared_connection_client( + .spawn_routed_shared_agent_client( &db_config.db_type, db_config.driver_profile.as_deref(), &db_config.agent_java_options, agent_session_id.clone(), - agent_connect_params( + agent_connect_params_with_role( &db_config, &host, port, db_config.effective_database().unwrap_or(""), + session_role, ), agent_connect_timeout(&db_config), ) @@ -1681,11 +2019,12 @@ impl AppState { client .call_method_with_timeout::( AgentMethod::Connect, - agent_connect_params( + agent_connect_params_with_role( &db_config, &host, port, db_config.effective_database().unwrap_or(""), + session_role, ), Some(agent_connect_timeout(&db_config)), ) @@ -1703,15 +2042,15 @@ impl AppState { .into_iter() .next() .unwrap_or_else(|| "alternate".to_string()); - let alternate_params = agent_connect_params( + let alternate_params = agent_connect_params_with_role( &alternate_config, &host, port, alternate_config.effective_database().unwrap_or(""), + session_role, ); match self - .agent_manager - .spawn_shared_connection_client( + .spawn_routed_shared_agent_client( &alternate_config.db_type, alternate_config.driver_profile.as_deref(), &alternate_config.agent_java_options, @@ -1737,7 +2076,7 @@ impl AppState { } } }; - PoolKind::Agent(Arc::new(tokio::sync::Mutex::new(client))) + PoolKind::agent(client) } else { // ZooKeeper JVM properties are connection-scoped; shared agent daemons must not inherit them. let mut client = self @@ -1775,11 +2114,12 @@ impl AppState { match client .call_method_with_timeout::( AgentMethod::Connect, - agent_connect_params( + agent_connect_params_with_role( &alternate_config, &host, port, alternate_config.effective_database().unwrap_or(""), + session_role, ), Some(agent_connect_timeout(&alternate_config)), ) @@ -1804,7 +2144,7 @@ impl AppState { return Err(oracle_error_with_driver_hint(&db_config, &err)); } } - PoolKind::Agent(Arc::new(tokio::sync::Mutex::new(client))) + PoolKind::agent(client) } } DatabaseType::PrestoSql => { @@ -1858,7 +2198,7 @@ impl AppState { self.discard_stale_connection_attempt_pool(connection_id, pool_key.clone(), pool, &db_config).await; return Err(err); } - self.insert_connection_pool(pool_key.clone(), pool, &db_config).await; + self.insert_connection_pool(pool_key.clone(), pool, &db_config).await?; Ok(pool_key) } @@ -2572,9 +2912,17 @@ impl AppState { ); false } + Err(err) if should_replace_agent_runtime(&err) => { + log::warn!( + "Agent connection pool '{pool_key}' requested runtime replacement during health probe: {err}" + ); + drop(agent); + return self.detach_agent_pool_if_current(pool_key, &client, true).await; + } Err(err) => { log::warn!("Agent connection pool '{pool_key}' is stale: {err}"); - true + drop(agent); + return self.detach_agent_pool_if_current(pool_key, &client, false).await; } } } @@ -2595,7 +2943,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 { - close_pool_kind_with_timeout(pool_key.to_string(), pool).await; + self.pool_routing_control().close_pool_with_timeout(pool_key.to_string(), pool).await; true } else { false @@ -2621,6 +2969,40 @@ impl AppState { database: Option<&str>, catalog: Option<&str>, client_session_id: Option<&str>, + ) -> Result { + self.reconnect_pool_for_session_with_catalog_and_role( + connection_id, + database, + catalog, + client_session_id, + AgentSessionRole::Workload, + ) + .await + } + + pub(crate) async fn reconnect_metadata_pool_for_session( + &self, + connection_id: &str, + database: Option<&str>, + client_session_id: Option<&str>, + ) -> Result { + self.reconnect_pool_for_session_with_catalog_and_role( + connection_id, + database, + None, + client_session_id, + AgentSessionRole::Metadata, + ) + .await + } + + async fn reconnect_pool_for_session_with_catalog_and_role( + &self, + connection_id: &str, + database: Option<&str>, + catalog: Option<&str>, + client_session_id: Option<&str>, + session_role: AgentSessionRole, ) -> Result { let config = { let configs = self.configs.read().await; @@ -2629,7 +3011,7 @@ impl AppState { let db_type = config.as_ref().map(|config| config.db_type); let catalog = catalog.map(str::trim).filter(|value| !value.is_empty()); let base_pool_key = base_pool_key_for_with_catalog(db_type, connection_id, database, catalog, true); - let pool_key = session_scoped_pool_key_for(config.as_ref(), base_pool_key, client_session_id); + let pool_key = pool_key_for_session_role(config.as_ref(), base_pool_key, client_session_id, session_role); if self.uses_forwarded_transport(connection_id).await { self.remove_connection_pools(connection_id).await; self.reset_connection_transport(connection_id).await; @@ -2639,10 +3021,22 @@ 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 { - close_pool_kind_with_timeout(pool_key.clone(), pool).await; + 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.get_or_create_pool_for_session_with_catalog(connection_id, database, catalog, client_session_id).await + self.get_or_create_pool_for_session_inner( + connection_id, + database, + catalog, + client_session_id, + session_role, + None, + ) + .await } pub async fn close_client_session_pool( @@ -2651,19 +3045,53 @@ impl AppState { database: Option<&str>, client_session_id: &str, ) -> Result { - let Some((pool_key, pool)) = self.take_client_session_pool(connection_id, database, client_session_id).await? + let Some((pool_key, pool)) = self + .take_client_session_pool(connection_id, database, client_session_id, AgentSessionRole::Workload) + .await? else { return Ok(false); }; - close_pool_kind_with_timeout(pool_key, pool).await; + self.pool_routing_control().close_pool_with_timeout(pool_key, pool).await; Ok(true) } - pub(crate) async fn client_session_pool_cleanup_guard( + pub(crate) async fn close_metadata_session_pool( &self, connection_id: &str, database: Option<&str>, client_session_id: &str, + ) -> Result { + let Some((pool_key, pool)) = self + .take_client_session_pool(connection_id, database, client_session_id, AgentSessionRole::Metadata) + .await? + else { + return Ok(false); + }; + self.pool_routing_control().close_pool_with_timeout(pool_key, pool).await; + Ok(true) + } + + pub(crate) async fn metadata_session_pool_cleanup_guard( + &self, + connection_id: &str, + database: Option<&str>, + client_session_id: &str, + ) -> Option { + self.client_session_pool_cleanup_guard_for_role( + connection_id, + database, + client_session_id, + AgentSessionRole::Metadata, + ) + .await + } + + async fn client_session_pool_cleanup_guard_for_role( + &self, + connection_id: &str, + database: Option<&str>, + client_session_id: &str, + session_role: AgentSessionRole, ) -> Option { let config = { let configs = self.configs.read().await; @@ -2671,18 +3099,12 @@ impl AppState { }; let db_type = config.as_ref().map(|config| config.db_type); let base_pool_key = base_pool_key_for(db_type, connection_id, database, false); - let pool_key = session_scoped_pool_key_for(config.as_ref(), base_pool_key.clone(), Some(client_session_id)); + let pool_key = + pool_key_for_session_role(config.as_ref(), base_pool_key.clone(), Some(client_session_id), session_role); if pool_key == base_pool_key { return None; } - Some(ClientSessionPoolCleanupGuard { - pool_key, - connections: self.connections.clone(), - pool_activity: self.pool_activity.clone(), - postgres_cancel_contexts: self.postgres_cancel_contexts.clone(), - task_supervisor: self.task_supervisor.clone(), - armed: true, - }) + Some(ClientSessionPoolCleanupGuard { pool_key, routing: self.pool_routing_control(), armed: true }) } /// Removes a session-scoped pool immediately and schedules the potentially slow driver @@ -2693,18 +3115,86 @@ impl AppState { database: Option<&str>, client_session_id: &str, ) -> Result { - let Some(removed) = self.take_client_session_pool(connection_id, database, client_session_id).await? else { + let Some(removed) = self + .take_client_session_pool(connection_id, database, client_session_id, AgentSessionRole::Workload) + .await? + else { return Ok(false); }; - close_removed_pools_in_background(&self.task_supervisor, vec![removed]); + self.pool_routing_control().close_removed_in_background(vec![removed]); Ok(true) } + #[cfg(test)] + pub(crate) async fn replace_runtime_for_metadata_pool( + &self, + connection_id: &str, + database: Option<&str>, + client_session_id: Option<&str>, + ) -> bool { + let config = { + let configs = self.configs.read().await; + configs.get(connection_id).cloned() + }; + let db_type = config.as_ref().map(|config| config.db_type); + let base_pool_key = base_pool_key_for(db_type, connection_id, database, false); + let pool_key = + pool_key_for_session_role(config.as_ref(), base_pool_key, client_session_id, AgentSessionRole::Metadata); + self.detach_pool_by_key(&pool_key, true).await + } + + pub(crate) async fn detach_metadata_pool_after_error( + &self, + connection_id: &str, + database: Option<&str>, + client_session_id: Option<&str>, + error: &str, + replace_agent_runtime: bool, + ) -> bool { + let config = { + let configs = self.configs.read().await; + configs.get(connection_id).cloned() + }; + let db_type = config.as_ref().map(|config| config.db_type); + let base_pool_key = base_pool_key_for(db_type, connection_id, database, false); + let pool_key = + pool_key_for_session_role(config.as_ref(), base_pool_key, client_session_id, AgentSessionRole::Metadata); + if let Some(session_id) = crate::db::agent_driver::agent_rpc_error_session_id(error) { + let expected_client = { + let connections = self.connections.read().await; + match connections.get(&pool_key) { + Some(PoolKind::Agent(client)) if client.matches_session_id(&session_id) => Some(client.clone()), + Some(PoolKind::Agent(_)) => return false, + _ => None, + } + }; + if let Some(client) = expected_client { + return self.detach_agent_pool_if_current(&pool_key, &client, replace_agent_runtime).await; + } + } + self.detach_pool_by_key(&pool_key, replace_agent_runtime).await + } + + /// Detaches a pool before cleanup so a stuck Agent close cannot delay replacement. + pub async fn detach_pool_by_key(&self, pool_key: &str, replace_agent_runtime: bool) -> bool { + self.pool_routing_control().detach_pool_by_key(pool_key, replace_agent_runtime).await + } + + pub(crate) async fn detach_agent_pool_if_current( + &self, + pool_key: &str, + expected_client: &Arc, + replace_agent_runtime: bool, + ) -> bool { + self.pool_routing_control().detach_agent_pool_if_current(pool_key, expected_client, replace_agent_runtime).await + } + async fn take_client_session_pool( &self, connection_id: &str, database: Option<&str>, client_session_id: &str, + session_role: AgentSessionRole, ) -> Result, String> { let session = normalize_client_session_id(Some(client_session_id)); let Some(session) = session else { @@ -2716,7 +3206,7 @@ impl AppState { }; let db_type = config.as_ref().map(|config| config.db_type); let base_pool_key = base_pool_key_for(db_type, connection_id, database, false); - let pool_key = session_scoped_pool_key_for(config.as_ref(), base_pool_key.clone(), Some(&session)); + let pool_key = pool_key_for_session_role(config.as_ref(), base_pool_key.clone(), Some(&session), session_role); if pool_key == base_pool_key { return Ok(None); } @@ -2733,7 +3223,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 { - close_pool_kind_with_timeout(pool_key.to_string(), pool).await; + self.pool_routing_control().close_pool_with_timeout(pool_key.to_string(), pool).await; true } else { false @@ -2801,6 +3291,11 @@ impl AppState { self.postgres_cancel_contexts.write().await.remove(pool_key); match close_reclaimed_agent_pool(pool).await { Ok(()) => true, + Err((PoolKind::Agent(client), error)) if should_replace_agent_runtime(&error) => { + log::warn!("Reclaimed Agent pool '{pool_key}' requested runtime replacement while closing: {error}"); + self.pool_routing_control().replace_runtime_after_close_failure(pool_key, &client).await; + true + } Err((pool, error)) => { log::warn!("Failed to close reclaimed Agent pool '{pool_key}': {error}; restoring the pool"); self.restore_reclaimed_agent_pool(pool_key, pool).await; @@ -2822,7 +3317,7 @@ impl AppState { } }; if let (Some(config), Some(client)) = (config, client) { - self.start_keepalive_task(pool_key, &PoolKind::Agent(client), &config).await; + self.start_keepalive_task(pool_key, &PoolKind::Agent(client), &config); } } @@ -2832,7 +3327,7 @@ impl AppState { config_for_pool_key(pool_key, &configs).cloned() }; if let Some(config) = config { - self.insert_connection_pool_inner(pool_key.to_string(), pool, &config, false).await; + let _ = self.insert_connection_pool_inner(pool_key.to_string(), pool, &config, false).await; } else { self.pool_activity.write().await.insert(pool_key.to_string(), PoolActivity::now()); self.connections.write().await.insert(pool_key.to_string(), pool); @@ -2853,12 +3348,13 @@ impl AppState { } let base_pool_key = base_pool_key_for(db_type, connection_id, database, false); let session_prefix = format!("{base_pool_key}:session:"); + let metadata_role_key = format!("{base_pool_key}:role:metadata"); let keys_to_remove: Vec = self .connections .read() .await .keys() - .filter(|key| *key == &base_pool_key || key.starts_with(&session_prefix)) + .filter(|key| *key == &base_pool_key || *key == &metadata_role_key || key.starts_with(&session_prefix)) .cloned() .collect(); self.stop_keepalive_tasks(&keys_to_remove).await; @@ -2880,7 +3376,7 @@ impl AppState { drop(conns); let closed = !removed.is_empty(); for (key, pool) in removed { - close_pool_kind_with_timeout(key, pool).await; + self.pool_routing_control().close_pool_with_timeout(key, pool).await; } Ok(closed) } @@ -2952,7 +3448,7 @@ impl AppState { let pool_key = self.get_or_create_pool(connection_id, database).await?; enum IdentifierQuoteSource { NativeGaussdb(deadpool_postgres::Pool), - Agent(Arc>), + Agent(Arc), ExternalDriver { config: Arc, session: Arc }, } let source = { @@ -3162,6 +3658,7 @@ impl AppState { (checks, redis_keys) }; + let mut failed_agent_checks = Vec::new(); let mut dead_pools = Vec::new(); let timeout = crate::db::connection_timeout(); @@ -3299,6 +3796,7 @@ impl AppState { } Err(e) => { log::warn!("Agent connection pool '{key}' is unhealthy: {e}"); + failed_agent_checks.push((key.clone(), client.clone(), should_replace_agent_runtime(&e))); false } } @@ -3310,7 +3808,7 @@ impl AppState { | PoolKind::Nacos => true, PoolKind::Redis(_) => unreachable!("Redis handled separately"), }; - if !healthy { + if !healthy && !matches!(pool, PoolKind::Agent(_)) { dead_pools.push((key.clone(), agent_pool_identity(pool))); } } @@ -3331,6 +3829,10 @@ impl AppState { } } + for (key, client, replace_runtime) in failed_agent_checks { + self.detach_agent_pool_if_current(&key, &client, replace_runtime).await; + } + // Remove dead pools if !dead_pools.is_empty() { let mut conns = self.connections.write().await; @@ -3352,16 +3854,7 @@ impl AppState { } } drop(conns); - - let removed_keys: Vec = removed.iter().map(|(key, _)| key.clone()).collect(); - self.stop_keepalive_tasks(&removed_keys).await; - { - let mut activity = self.pool_activity.write().await; - for key in &removed_keys { - activity.remove(key); - } - } - close_removed_pools(removed).await; + self.pool_routing_control().finish_detach(removed).await; } // Re-establish SSH tunnels that have died @@ -3377,39 +3870,20 @@ impl AppState { pub async fn remove_connection_pools(&self, connection_id: &str) { let removed = self.drain_connection_pools(connection_id).await; - close_removed_pools(removed).await; + self.pool_routing_control().close_removed(removed).await; } pub async fn remove_connection_pools_detached(&self, connection_id: &str) { let removed = self.drain_connection_pools(connection_id).await; - close_removed_pools_in_background(&self.task_supervisor, removed); + self.pool_routing_control().close_removed_in_background(removed); } pub async fn invalidate_agent_pool_if_current( &self, pool_key: &str, - expected: &Arc>, + expected: &Arc, ) -> bool { - let removed = { - let mut pools = self.connections.write().await; - let is_current = matches!( - pools.get(pool_key), - Some(PoolKind::Agent(current)) if Arc::ptr_eq(current, expected) - ); - if is_current { - self.task_supervisor.stop(&format!("keepalive:{pool_key}")); - pools.remove(pool_key) - } else { - None - } - }; - let Some(pool) = removed else { - return false; - }; - self.pool_activity.write().await.remove(pool_key); - self.postgres_cancel_contexts.write().await.remove(pool_key); - close_removed_pools_in_background(&self.task_supervisor, vec![(pool_key.to_string(), pool)]); - true + self.detach_agent_pool_if_current(pool_key, expected, false).await } async fn drain_all_connection_pools(&self) -> Vec<(String, PoolKind)> { @@ -3424,7 +3898,7 @@ impl AppState { #[cfg(feature = "duckdb-sidecar")] async fn remove_duckdb_pools_detached(&self) { let removed = self.drain_duckdb_pools().await; - close_removed_pools_in_background(&self.task_supervisor, removed); + self.pool_routing_control().close_removed_in_background(removed); } #[cfg(not(feature = "duckdb-sidecar"))] @@ -3432,7 +3906,7 @@ impl AppState { pub async fn remove_external_driver_pools(&self, driver_id: &str) { let removed = self.drain_external_driver_pools(driver_id).await; - close_removed_pools(removed).await; + self.pool_routing_control().close_removed(removed).await; } async fn drain_connection_pools(&self, connection_id: &str) -> Vec<(String, PoolKind)> { @@ -3553,7 +4027,7 @@ enum KeepaliveTarget { HBase(db::hbase_driver::HBaseClient), VectorDb(db::vector_driver::VectorClient), InfluxDb(db::influxdb_driver::InfluxdbClient), - Agent(Arc>), + Agent(Arc), } impl KeepaliveTarget { @@ -3582,6 +4056,23 @@ async fn remove_keepalive_pool_if_current( } } +async fn detach_keepalive_target_if_current( + routing: &PoolRoutingControl, + connections: &Arc>>, + pool_key: &str, + target: &KeepaliveTarget, + replace_agent_runtime: bool, +) -> bool { + if let KeepaliveTarget::Agent(expected) = target { + return routing.detach_agent_pool_if_current(pool_key, expected, replace_agent_runtime).await; + } + let Some(pool) = remove_keepalive_pool_if_current(connections, pool_key, target).await else { + return false; + }; + routing.finish_detach(vec![(pool_key.to_string(), pool)]).await; + true +} + fn keepalive_target_from_pool(pool: &PoolKind, config: &ConnectionConfig) -> Option { match pool { PoolKind::Mysql(pool, _) => Some(KeepaliveTarget::Mysql(pool.clone())), @@ -3638,10 +4129,7 @@ async fn ping_keepalive_target(target: &mut KeepaliveTarget, timeout: Duration) match client.validate_connection(Some(timeout)).await { Ok(_) => Ok(()), Err(err) if is_agent_validate_connection_unsupported(&err) => Ok(()), - Err(err) => { - client.kill(); - Err(err) - } + Err(err) => Err(err), } } } @@ -3805,6 +4293,22 @@ fn session_scoped_pool_key_for( session_scoped_pool_key(base_pool_key, client_session_id) } +fn pool_key_for_session_role( + config: Option<&ConnectionConfig>, + base_pool_key: String, + client_session_id: Option<&str>, + session_role: AgentSessionRole, +) -> String { + let pool_key = session_scoped_pool_key_for(config, base_pool_key, client_session_id); + if session_role == AgentSessionRole::Metadata + && config.is_some_and(|config| database_capabilities::is_agent_type(&config.db_type)) + { + format!("{pool_key}:role:metadata") + } else { + pool_key + } +} + fn clone_pool_kind(pool: &PoolKind) -> PoolKind { match pool { PoolKind::Mysql(p, mode) => PoolKind::Mysql(p.clone(), *mode), @@ -3835,7 +4339,7 @@ fn clone_pool_kind(pool: &PoolKind) -> PoolKind { } } -pub async fn close_pool_kind(pool: PoolKind) { +async fn close_pool_kind(pool: PoolKind) -> Result<(), String> { match pool { PoolKind::Mysql(p, _) => { let _ = p.disconnect().await; @@ -3880,7 +4384,7 @@ pub async fn close_pool_kind(pool: PoolKind) { } PoolKind::Agent(client) => { let mut client = client.lock().await; - let _ = client.disconnect().await; + client.disconnect().await?; } PoolKind::ExternalDriver { session, .. } => { session.shutdown().await; @@ -3888,6 +4392,7 @@ pub async fn close_pool_kind(pool: PoolKind) { PoolKind::MessageQueue => {} PoolKind::Nacos => {} } + Ok(()) } async fn close_reclaimed_agent_pool(pool: PoolKind) -> Result<(), (PoolKind, String)> { @@ -3905,36 +4410,6 @@ async fn close_reclaimed_agent_pool(pool: PoolKind) -> Result<(), (PoolKind, Str } } -async fn close_removed_pools(removed: Vec<(String, PoolKind)>) { - for (pool_key, pool) in removed { - close_pool_kind_with_timeout(pool_key, pool).await; - } -} - -fn close_removed_pools_in_background(supervisor: &TaskSupervisor, removed: Vec<(String, PoolKind)>) { - if removed.is_empty() { - return; - } - // Supervision keeps detached cleanup visible to application shutdown instead of leaving an - // untracked Tokio task that may be abandoned silently. - let pool_count = removed.len(); - let task_key = format!("pool-close:{}", uuid::Uuid::new_v4()); - if !supervisor.spawn_once(task_key, move |_| async move { - close_removed_pools(removed).await; - }) { - log::debug!("Dropped {pool_count} detached pool handle(s) during application shutdown"); - } -} - -async fn close_pool_kind_with_timeout(pool_key: String, pool: PoolKind) { - match tokio::time::timeout(Duration::from_secs(POOL_CLOSE_TIMEOUT_SECS), close_pool_kind(pool)).await { - Ok(()) => {} - Err(_) => log::warn!( - "Timed out closing connection pool '{pool_key}' after {POOL_CLOSE_TIMEOUT_SECS}s; cleanup will continue by dropping the pool handle." - ), - } -} - fn extract_auth_token_from_params(params: &str) -> Option { params .trim() @@ -4018,7 +4493,7 @@ fn should_validate_existing_pool_before_reuse(db_type: DatabaseType) -> bool { !matches!(db_type, DatabaseType::Postgres | DatabaseType::Etcd) } -fn agent_pool_identity(pool: &PoolKind) -> Option>> { +fn agent_pool_identity(pool: &PoolKind) -> Option> { match pool { PoolKind::Agent(client) => Some(client.clone()), _ => None, @@ -4805,9 +5280,7 @@ mod tests { } fn agent_pool_stub() -> PoolKind { - PoolKind::Agent(std::sync::Arc::new(tokio::sync::Mutex::new( - crate::db::agent_driver::AgentDriverClient::test_stub(), - ))) + PoolKind::agent(crate::db::agent_driver::AgentDriverClient::test_stub()) } #[tokio::test] @@ -5025,7 +5498,8 @@ mod tests { let pool = db::mysql::connect_bare_with_pool_limit(&url, Duration::from_secs(5), 1).await.unwrap(); state .insert_connection_pool("conn".to_string(), PoolKind::Mysql(pool.clone(), MysqlMode::Normal), &config) - .await; + .await + .unwrap(); let held_connection = pool.get_conn().await.unwrap(); let started = Instant::now(); @@ -5439,6 +5913,57 @@ mod tests { ); } + #[test] + fn agent_metadata_pool_keys_are_isolated_from_workload_keys() { + let mut config = mysql_config(Some("analytics")); + config.db_type = DatabaseType::Dameng; + + let workload = super::pool_key_for_session_role( + Some(&config), + "conn:analytics".to_string(), + Some("task:1"), + crate::agent_connection::AgentSessionRole::Workload, + ); + let metadata = super::pool_key_for_session_role( + Some(&config), + "conn:analytics".to_string(), + Some("task:1"), + crate::agent_connection::AgentSessionRole::Metadata, + ); + let base_metadata = super::pool_key_for_session_role( + Some(&config), + "conn:analytics".to_string(), + None, + crate::agent_connection::AgentSessionRole::Metadata, + ); + + assert_eq!(workload, "conn:analytics:session:task_1"); + assert_eq!(metadata, "conn:analytics:session:task_1:role:metadata"); + assert_eq!(base_metadata, "conn:analytics:role:metadata"); + assert_ne!(metadata, workload); + } + + #[test] + fn non_agent_metadata_role_preserves_existing_pool_keys() { + let config = mysql_config(Some("analytics")); + + let workload = super::pool_key_for_session_role( + Some(&config), + "conn:analytics".to_string(), + Some("task:1"), + crate::agent_connection::AgentSessionRole::Workload, + ); + let metadata = super::pool_key_for_session_role( + Some(&config), + "conn:analytics".to_string(), + Some("task:1"), + crate::agent_connection::AgentSessionRole::Metadata, + ); + + assert_eq!(metadata, workload); + assert_eq!(metadata, "conn:analytics:session:task_1"); + } + #[test] fn redis_sentinel_transport_ids_are_connection_scoped_by_role_and_endpoint() { let endpoint = db::redis_driver::RedisNodeEndpoint { host: "10.0.0.8".to_string(), port: 6379 }; @@ -5639,15 +6164,16 @@ for line in sys.stdin: ) .unwrap(); + let python = if cfg!(windows) { "python" } else { "python3" }; let runtime = crate::db::agent_driver::AgentRuntimeClient::spawn( - crate::db::agent_driver::AgentLaunchSpec::new("python3") + 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(tokio::sync::Mutex::new( + let client = std::sync::Arc::new(crate::db::agent_driver::PooledAgentClient::new( crate::db::agent_driver::AgentDriverClient::shared_session(runtime.clone(), "metadata-session".to_string()), )); state.connections.write().await.insert("conn:analytics".to_string(), PoolKind::Agent(client.clone())); @@ -5698,8 +6224,9 @@ for line in sys.stdin: #[tokio::test] async fn busy_agent_pool_skips_health_probe_without_waiting() { let (state, dir) = test_app_state().await; - let client = - std::sync::Arc::new(tokio::sync::Mutex::new(crate::db::agent_driver::AgentDriverClient::test_stub())); + let client = std::sync::Arc::new(crate::db::agent_driver::PooledAgentClient::new( + crate::db::agent_driver::AgentDriverClient::test_stub(), + )); state.connections.write().await.insert("conn".to_string(), PoolKind::Agent(client.clone())); let _busy = client.lock().await; @@ -5777,7 +6304,7 @@ for line in sys.stdin: .await .insert(pool_key.to_string(), super::PoolActivity::idle_for(std::time::Duration::from_secs(10))); let pool = super::clone_pool_kind(state.connections.read().await.get(pool_key).unwrap()); - state.start_keepalive_task(pool_key, &pool, &config).await; + state.start_keepalive_task(pool_key, &pool, &config); tokio::time::sleep(std::time::Duration::from_millis(20)).await; @@ -5827,6 +6354,500 @@ for line in sys.stdin: let _ = std::fs::remove_dir_all(dir); } + #[tokio::test] + async fn replace_runtime_for_base_metadata_pool_detaches_routing_and_kills_runtime() { + let (state, dir) = test_app_state().await; + let script_path = dir.join("replace-runtime-agent.py"); + let request_started_path = dir.join("replace-runtime-request-started"); + let request_started = serde_json::to_string(&request_started_path.to_string_lossy()).unwrap(); + std::fs::write( + &script_path, + format!( + r#"import json, pathlib, sys, time +request_started = pathlib.Path({request_started}) +print(json.dumps({{'ready': True}}), flush=True) +for line in sys.stdin: + req = json.loads(line) + if req['method'] == 'handshake': + result = {{'protocolVersion': 2, 'agentProtocolVersion': 2, 'capabilities': ['multi_session']}} + elif req['method'] == 'list_databases': + request_started.write_text('started') + time.sleep(30) + result = [] + else: + result = {{}} + print(json.dumps({{'jsonrpc': '2.0', 'id': req['id'], 'result': result}}), 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 metadata_client = std::sync::Arc::new(crate::db::agent_driver::PooledAgentClient::new( + crate::db::agent_driver::AgentDriverClient::shared_session(runtime.clone(), "metadata-session".to_string()), + )); + runtime.increment_session_count(); + let workload_client = std::sync::Arc::new(crate::db::agent_driver::PooledAgentClient::new( + crate::db::agent_driver::AgentDriverClient::shared_session(runtime.clone(), "workload-session".to_string()), + )); + + let mut config = mysql_config(Some("analytics")); + config.id = "conn".to_string(); + config.db_type = DatabaseType::Dameng; + state.configs.write().await.insert(config.id.clone(), config); + let pool_key = "conn:analytics:role:metadata"; + let sibling_pool_key = "conn:analytics:session:workload"; + { + let mut connections = state.connections.write().await; + connections.insert(pool_key.to_string(), PoolKind::Agent(metadata_client.clone())); + connections.insert(sibling_pool_key.to_string(), PoolKind::Agent(workload_client)); + } + { + let mut activity = state.pool_activity.write().await; + activity.insert(pool_key.to_string(), super::PoolActivity::now()); + activity.insert(sibling_pool_key.to_string(), super::PoolActivity::now()); + } + + let blocked_client = metadata_client.clone(); + let blocked_request = tokio::spawn(async move { + blocked_client.lock().await.list_databases::>(Some(Duration::from_secs(30))).await + }); + for _ in 0..100 { + if request_started_path.exists() { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + assert!(request_started_path.exists()); + + let replacement = tokio::time::timeout( + Duration::from_secs(1), + state.replace_runtime_for_metadata_pool("conn", Some("analytics"), None), + ) + .await; + if replacement.is_err() { + runtime.kill(); + } + let _ = tokio::time::timeout(Duration::from_secs(2), blocked_request).await; + + assert!(replacement.expect("runtime replacement must not wait for the Agent client lock")); + assert!(!state.connections.read().await.contains_key(pool_key)); + assert!(!state.connections.read().await.contains_key(sibling_pool_key)); + assert!(!state.pool_activity.read().await.contains_key(pool_key)); + assert!(!state.pool_activity.read().await.contains_key(sibling_pool_key)); + assert!(runtime.is_failed()); + + let _ = std::fs::remove_dir_all(dir); + } + + async fn replace_runtime_on_error_clients( + dir: &std::path::Path, + error_method: &str, + ) -> ( + std::sync::Arc, + std::sync::Arc, + std::sync::Arc, + ) { + let script_path = dir.join("close-replace-runtime-agent.py"); + let script = r#"import json, sys +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'] == '__ERROR_METHOD__': + 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) +"# + .replace("'__ERROR_METHOD__'", &serde_json::to_string(error_method).unwrap()); + std::fs::write(&script_path, script).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 metadata_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(), + ), + )); + runtime.increment_session_count(); + let workload_client = std::sync::Arc::new(crate::db::agent_driver::PooledAgentClient::new( + crate::db::agent_driver::AgentDriverClient::shared_session( + runtime.clone(), + "workload-agent-session".to_string(), + ), + )); + (runtime, metadata_client, workload_client) + } + + #[tokio::test] + async fn metadata_close_replace_runtime_detaches_shared_runtime_siblings() { + let (state, dir) = test_app_state().await; + let (runtime, metadata_client, workload_client) = replace_runtime_on_error_clients(&dir, "close_session").await; + let mut config = mysql_config(Some("analytics")); + config.id = "conn".to_string(); + config.db_type = DatabaseType::Dameng; + state.configs.write().await.insert(config.id.clone(), config); + let metadata_pool_key = "conn:analytics:session:metadata-session:role:metadata"; + let workload_pool_key = "conn:analytics:session:workload-session"; + { + let mut connections = state.connections.write().await; + connections.insert(metadata_pool_key.to_string(), PoolKind::Agent(metadata_client)); + connections.insert(workload_pool_key.to_string(), PoolKind::Agent(workload_client)); + } + { + let mut activity = state.pool_activity.write().await; + activity.insert(metadata_pool_key.to_string(), super::PoolActivity::now()); + activity.insert(workload_pool_key.to_string(), super::PoolActivity::now()); + } + + assert!(state.close_metadata_session_pool("conn", Some("analytics"), "metadata-session").await.unwrap()); + + assert!(!state.connections.read().await.contains_key(metadata_pool_key)); + assert!(!state.connections.read().await.contains_key(workload_pool_key)); + assert!(!state.pool_activity.read().await.contains_key(workload_pool_key)); + assert!(runtime.is_failed()); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn reclaim_close_replace_runtime_never_restores_failed_pool() { + let (state, dir) = test_app_state().await; + let (runtime, reclaimed_client, sibling_client) = replace_runtime_on_error_clients(&dir, "close_session").await; + let mut config = mysql_config(Some("analytics")); + config.id = "conn".to_string(); + config.db_type = DatabaseType::Dameng; + state.configs.write().await.insert(config.id.clone(), config); + let reclaimed_pool_key = "conn:analytics"; + let sibling_pool_key = "conn:analytics:session:workload"; + { + let mut connections = state.connections.write().await; + connections.insert(reclaimed_pool_key.to_string(), PoolKind::Agent(reclaimed_client)); + connections.insert(sibling_pool_key.to_string(), PoolKind::Agent(sibling_client)); + } + state.pool_activity.write().await.insert(reclaimed_pool_key.to_string(), super::PoolActivity::now()); + + let reclaimed = state.try_reclaim_idle_agent_pool(reclaimed_pool_key).await; + + assert!(reclaimed); + assert!(!state.connections.read().await.contains_key(reclaimed_pool_key)); + assert!(!state.connections.read().await.contains_key(sibling_pool_key)); + assert!(runtime.is_failed()); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn inserting_replacement_returns_error_when_previous_close_replaces_shared_runtime() { + let (state, dir) = test_app_state().await; + let (runtime, previous_client, replacement_client) = + replace_runtime_on_error_clients(&dir, "close_session").await; + let mut config = mysql_config(Some("analytics")); + config.id = "conn".to_string(); + config.db_type = DatabaseType::Dameng; + state.configs.write().await.insert(config.id.clone(), config.clone()); + let pool_key = "conn:analytics"; + state.connections.write().await.insert(pool_key.to_string(), PoolKind::Agent(previous_client)); + + let result = + state.insert_connection_pool(pool_key.to_string(), PoolKind::Agent(replacement_client), &config).await; + + assert!(result.is_err()); + assert!(!state.connections.read().await.contains_key(pool_key)); + assert!(runtime.is_failed()); + 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; + let (runtime, client, sibling_client) = replace_runtime_on_error_clients(&dir, "unused").await; + let mut config = mysql_config(Some("analytics")); + config.id = "conn".to_string(); + config.db_type = DatabaseType::Dameng; + state.connections.write().await.insert("conn:billing".to_string(), PoolKind::Agent(sibling_client)); + runtime.kill(); + + let result = state.insert_connection_pool("conn:analytics".to_string(), PoolKind::Agent(client), &config).await; + + assert!(result.is_err()); + assert!(state.connections.read().await.is_empty()); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn inserting_failed_agent_pool_preserves_healthy_existing_route_state() { + let (state, dir) = test_app_state().await; + let (healthy_runtime, healthy_client, _healthy_sibling) = + replace_runtime_on_error_clients(&dir, "unused").await; + let (failed_runtime, failed_client, _failed_sibling) = replace_runtime_on_error_clients(&dir, "unused").await; + let mut config = mysql_config(Some("analytics")); + config.id = "conn".to_string(); + config.db_type = DatabaseType::Dameng; + config.keepalive_interval_secs = 60; + let pool_key = "conn:analytics"; + state + .insert_connection_pool(pool_key.to_string(), PoolKind::Agent(healthy_client.clone()), &config) + .await + .unwrap(); + state + .pool_activity + .write() + .await + .insert(pool_key.to_string(), super::PoolActivity::idle_for(Duration::from_secs(600))); + assert_eq!(state.supervised_task_count(), 1); + failed_runtime.kill(); + + let result = state.insert_connection_pool(pool_key.to_string(), PoolKind::Agent(failed_client), &config).await; + + assert!(result.is_err()); + let connections = state.connections.read().await; + let PoolKind::Agent(routed_client) = connections.get(pool_key).expect("healthy route must remain") else { + panic!("existing route must remain an Agent pool"); + }; + assert!(healthy_client.shares_runtime_with(routed_client)); + drop(connections); + assert!( + state.pool_activity.read().await.get(pool_key).expect("activity must remain").elapsed().as_secs() >= 300 + ); + assert_eq!(state.supervised_task_count(), 1); + assert!(!healthy_runtime.is_failed()); + + state.shutdown(Duration::from_secs(1)).await; + healthy_runtime.kill(); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn inserting_pool_does_not_wait_for_activity_while_holding_route_lock() { + let (state, dir) = test_app_state().await; + let state = std::sync::Arc::new(state); + let activity_guard = state.pool_activity.read().await; + let mut config = mysql_config(Some("analytics")); + config.keepalive_interval_secs = 0; + let publishing_state = state.clone(); + let publish = tokio::spawn(async move { + publishing_state.insert_connection_pool("conn:analytics".to_string(), agent_pool_stub(), &config).await + }); + tokio::time::sleep(Duration::from_millis(20)).await; + + let route_read = tokio::time::timeout(Duration::from_millis(100), state.connections.read()).await; + + assert!(route_read.is_ok(), "pool publication must not await activity while holding the route lock"); + drop(route_read); + drop(activity_guard); + assert!(tokio::time::timeout(Duration::from_secs(1), publish).await.unwrap().unwrap().is_ok()); + state.shutdown(Duration::from_secs(1)).await; + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn open_session_replace_runtime_error_detaches_existing_shared_runtime_pools() { + let (state, dir) = test_app_state().await; + let (runtime, first_client, second_client) = replace_runtime_on_error_clients(&dir, "unused").await; + { + let mut connections = state.connections.write().await; + connections.insert("conn:analytics".to_string(), PoolKind::Agent(first_client)); + connections.insert("conn:billing".to_string(), PoolKind::Agent(second_client)); + } + let error = crate::agent_runtime::SharedConnectionOpenError { + message: "open session requested runtime replacement".to_string(), + runtime: Some(runtime.clone()), + }; + + let message = state.handle_shared_connection_open_error(error).await; + + assert_eq!(message, "open session requested runtime replacement"); + assert!(state.connections.read().await.is_empty()); + assert!(runtime.is_failed()); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn stale_probe_replace_runtime_detaches_shared_runtime_siblings() { + let (state, dir) = test_app_state().await; + let (runtime, target_client, sibling_client) = + replace_runtime_on_error_clients(&dir, "validate_connection").await; + let target_pool_key = "conn:analytics"; + let sibling_pool_key = "conn:analytics:session:workload"; + { + let mut connections = state.connections.write().await; + connections.insert(target_pool_key.to_string(), PoolKind::Agent(target_client)); + connections.insert(sibling_pool_key.to_string(), PoolKind::Agent(sibling_client)); + } + + assert!(state.remove_stale_connection_pool(target_pool_key).await); + + assert!(!state.connections.read().await.contains_key(target_pool_key)); + assert!(!state.connections.read().await.contains_key(sibling_pool_key)); + assert!(runtime.is_failed()); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn global_health_validates_current_session_and_fail_stops_shared_runtime() { + let (state, dir) = test_app_state().await; + let (runtime, target_client, sibling_client) = + replace_runtime_on_error_clients(&dir, "validate_connection").await; + { + let mut connections = state.connections.write().await; + connections.insert("conn:analytics".to_string(), PoolKind::Agent(target_client)); + connections.insert("conn:billing".to_string(), PoolKind::Agent(sibling_client)); + } + + state.refresh_connections().await; + + assert!(state.connections.read().await.is_empty()); + assert!(runtime.is_failed()); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn stale_agent_failure_preserves_newer_pool_generation() { + let (state, dir) = test_app_state().await; + let (stale_runtime, stale_client, _stale_sibling) = replace_runtime_on_error_clients(&dir, "unused").await; + let (current_runtime, current_client, _current_sibling) = + replace_runtime_on_error_clients(&dir, "unused").await; + let pool_key = "conn:analytics"; + state.connections.write().await.insert(pool_key.to_string(), PoolKind::Agent(current_client.clone())); + + assert!(!state.detach_agent_pool_if_current(pool_key, &stale_client, true).await); + + let connections = state.connections.read().await; + let PoolKind::Agent(routed_client) = connections.get(pool_key).expect("new generation must remain routed") + else { + panic!("current route must remain an Agent pool"); + }; + assert!(std::sync::Arc::ptr_eq(routed_client, ¤t_client)); + drop(connections); + assert!(stale_runtime.is_failed()); + assert!(!current_runtime.is_failed()); + current_runtime.kill(); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn stale_agent_failure_detaches_current_routes_on_the_same_runtime() { + let (state, dir) = test_app_state().await; + let (runtime, stale_client, current_client) = replace_runtime_on_error_clients(&dir, "unused").await; + let pool_key = "conn:analytics"; + let sibling_pool_key = "conn:billing"; + { + let mut connections = state.connections.write().await; + connections.insert(pool_key.to_string(), PoolKind::Agent(current_client.clone())); + connections.insert(sibling_pool_key.to_string(), PoolKind::Agent(current_client)); + } + + assert!(state.detach_agent_pool_if_current(pool_key, &stale_client, true).await); + + assert!(!state.connections.read().await.contains_key(pool_key)); + assert!(!state.connections.read().await.contains_key(sibling_pool_key)); + assert!(runtime.is_failed()); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn stale_metadata_error_preserves_newer_session_generation() { + let (state, dir) = test_app_state().await; + let (current_runtime, current_client, _current_sibling) = + replace_runtime_on_error_clients(&dir, "unused").await; + let mut config = mysql_config(Some("analytics")); + config.id = "conn".to_string(); + config.db_type = DatabaseType::Dameng; + state.configs.write().await.insert(config.id.clone(), config); + let pool_key = "conn:analytics:role:metadata"; + state.connections.write().await.insert(pool_key.to_string(), PoolKind::Agent(current_client.clone())); + let stale_error = concat!( + "Agent RPC error (-1): stale metadata failure\nDBX_AGENT_ERROR_DATA:", + r#"{"category":"resource","sessionDisposition":"replace_runtime","agentSessionId":"stale-session"}"# + ); + + assert!(!state.detach_metadata_pool_after_error("conn", Some("analytics"), None, stale_error, true).await); + + let connections = state.connections.read().await; + let PoolKind::Agent(routed_client) = connections.get(pool_key).expect("new metadata generation must remain") + else { + panic!("metadata route must remain an Agent pool"); + }; + assert!(std::sync::Arc::ptr_eq(routed_client, ¤t_client)); + drop(connections); + assert!(!current_runtime.is_failed()); + current_runtime.kill(); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn keepalive_probe_reports_replacement_without_killing_runtime_directly() { + let (_state, dir) = test_app_state().await; + let (runtime, target_client, _sibling_client) = + replace_runtime_on_error_clients(&dir, "validate_connection").await; + let mut target = super::KeepaliveTarget::Agent(target_client); + + let error = super::ping_keepalive_target(&mut target, Duration::from_secs(1)).await.unwrap_err(); + + assert_eq!( + crate::db::agent_driver::agent_session_disposition(&error), + Some(crate::db::agent_driver::AgentSessionDisposition::ReplaceRuntime) + ); + assert!(!runtime.is_failed()); + runtime.kill(); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn detach_pool_by_key_removes_routing_before_background_close() { + let (state, dir) = test_app_state().await; + let pool_key = "conn:session:timeout"; + let pool = crate::db::sqlite::connect_path(":memory:").await.unwrap(); + state.connections.write().await.insert(pool_key.to_string(), PoolKind::Sqlite(pool)); + state.pool_activity.write().await.insert(pool_key.to_string(), super::PoolActivity::now()); + + assert!(state.detach_pool_by_key(pool_key, false).await); + assert!(!state.connections.read().await.contains_key(pool_key)); + assert!(!state.pool_activity.read().await.contains_key(pool_key)); + + for _ in 0..100 { + if state.supervised_task_count() == 0 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + assert_eq!(state.supervised_task_count(), 0); + + let _ = std::fs::remove_dir_all(dir); + } + #[tokio::test] async fn client_session_cleanup_guard_detaches_pool_when_request_is_dropped() { let (state, dir) = test_app_state().await; @@ -5839,8 +6860,15 @@ for line in sys.stdin: state.connections.write().await.insert(pool_key.to_string(), PoolKind::Sqlite(pool)); state.pool_activity.write().await.insert(pool_key.to_string(), super::PoolActivity::now()); - let guard = - state.client_session_pool_cleanup_guard("conn", None, "completion-objects:request-1").await.unwrap(); + let guard = state + .client_session_pool_cleanup_guard_for_role( + "conn", + None, + "completion-objects:request-1", + crate::agent_connection::AgentSessionRole::Workload, + ) + .await + .unwrap(); drop(guard); for _ in 0..100 { @@ -5955,6 +6983,7 @@ for line in sys.stdin: let mut conns = state.connections.write().await; conns.insert("conn".to_string(), PoolKind::Sqlite(pool.clone())); conns.insert("conn:analytics".to_string(), PoolKind::Sqlite(pool.clone())); + conns.insert("conn:analytics:role:metadata".to_string(), PoolKind::Sqlite(pool.clone())); conns.insert("conn:analytics:session:tab-1".to_string(), PoolKind::Sqlite(pool.clone())); conns.insert("conn:billing".to_string(), PoolKind::Sqlite(pool)); } @@ -5964,6 +6993,7 @@ for line in sys.stdin: let conns = state.connections.read().await; assert!(conns.contains_key("conn")); assert!(!conns.contains_key("conn:analytics")); + assert!(!conns.contains_key("conn:analytics:role:metadata")); assert!(!conns.contains_key("conn:analytics:session:tab-1")); assert!(conns.contains_key("conn:billing")); diff --git a/crates/dbx-core/src/database_export.rs b/crates/dbx-core/src/database_export.rs index 2053250d2..1e3cee8a6 100644 --- a/crates/dbx-core/src/database_export.rs +++ b/crates/dbx-core/src/database_export.rs @@ -2371,7 +2371,6 @@ mod tests { #[test] fn concurrent_prefetch_only_allowed_for_multi_connection_pools() { use crate::connection::PoolKind; - use std::sync::Arc; // ChClient::new 只构造 HTTP 客户端,不发起连接 let clickhouse = PoolKind::ClickHouse(crate::db::clickhouse_driver::ChClient::new( @@ -2383,8 +2382,7 @@ mod tests { assert!(concurrent_metadata_prefetch_allowed(Some(&clickhouse))); // Agent(JDBC sidecar)请求超时覆盖排队时间,必须回退串行 - let agent = - PoolKind::Agent(Arc::new(tokio::sync::Mutex::new(crate::db::agent_driver::AgentDriverClient::test_stub()))); + let agent = PoolKind::agent(crate::db::agent_driver::AgentDriverClient::test_stub()); assert!(!concurrent_metadata_prefetch_allowed(Some(&agent))); assert!(!concurrent_metadata_prefetch_allowed(None)); diff --git a/crates/dbx-core/src/db/agent_driver.rs b/crates/dbx-core/src/db/agent_driver.rs index a889d0e15..0ee71a7ee 100644 --- a/crates/dbx-core/src/db/agent_driver.rs +++ b/crates/dbx-core/src/db/agent_driver.rs @@ -285,15 +285,78 @@ impl AgentRuntimeClient { fn decode_agent_response(response: Value) -> Result { if let Some(err) = response.get("error") { - let message = err.get("message").and_then(Value::as_str).unwrap_or("Unknown agent error"); - let code = err.get("code").and_then(Value::as_i64).unwrap_or(-1); - return Err(format!("Agent RPC error ({code}): {message}")); + return Err(format_agent_rpc_error(err)); } let result = response.get("result").ok_or_else(|| "Agent response missing both 'result' and 'error'".to_string())?; serde_json::from_value(result.clone()).map_err(|e| format!("Failed to deserialize agent result: {e}")) } +const AGENT_RPC_ERROR_DATA_MARKER: &str = "\nDBX_AGENT_ERROR_DATA:"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum AgentSessionDisposition { + Keep, + Quarantine, + ReplaceRuntime, +} + +pub fn agent_rpc_error_category(error: &str) -> Option { + agent_rpc_error_data(error)?.get("category")?.as_str().map(str::to_string) +} + +pub fn agent_rpc_error_session_id(error: &str) -> Option { + agent_rpc_error_data(error)?.get("agentSessionId")?.as_str().map(str::to_string) +} + +pub fn agent_session_disposition(error: &str) -> Option { + match agent_rpc_error_data(error)?.get("sessionDisposition")?.as_str()? { + "keep" => Some(AgentSessionDisposition::Keep), + "quarantine" => Some(AgentSessionDisposition::Quarantine), + "replace_runtime" => Some(AgentSessionDisposition::ReplaceRuntime), + _ => None, + } +} + +fn format_agent_rpc_error(error: &Value) -> String { + let message = error.get("message").and_then(Value::as_str).unwrap_or("Unknown agent error"); + let code = error.get("code").and_then(Value::as_i64).unwrap_or(-1); + let mut formatted = format!("Agent RPC error ({code}): {message}"); + if let Some(data) = error.get("data").filter(|data| data.is_object()) { + formatted.push_str(AGENT_RPC_ERROR_DATA_MARKER); + formatted.push_str(&data.to_string()); + } + formatted +} + +fn agent_rpc_error_data(error: &str) -> Option { + let (_, data) = error.rsplit_once(AGENT_RPC_ERROR_DATA_MARKER)?; + serde_json::from_str(data).ok() +} + +fn with_agent_rpc_error_session_id(error: String, agent_session_id: Option<&str>) -> String { + let Some(agent_session_id) = agent_session_id else { + return error; + }; + let Some((message, data)) = error.rsplit_once(AGENT_RPC_ERROR_DATA_MARKER) else { + return format!( + "{error}{AGENT_RPC_ERROR_DATA_MARKER}{}", + serde_json::json!({ "agentSessionId": agent_session_id }) + ); + }; + let Ok(mut data) = serde_json::from_str::(data) else { + return format!( + "{error}{AGENT_RPC_ERROR_DATA_MARKER}{}", + serde_json::json!({ "agentSessionId": agent_session_id }) + ); + }; + let Some(data) = data.as_object_mut() else { + return error; + }; + data.insert("agentSessionId".to_string(), Value::String(agent_session_id.to_string())); + format!("{message}{AGENT_RPC_ERROR_DATA_MARKER}{}", Value::Object(data.clone())) +} + fn deserialize_cached_agent_result(result: Result) -> Result { result .and_then(|value| serde_json::from_value(value).map_err(|e| format!("Failed to deserialize agent result: {e}"))) @@ -329,6 +392,66 @@ pub struct AgentDriverClient { cached_query: Option, } +/// Keeps serialized session RPC access separate from the process-level fail-stop handle. +/// A stuck RPC may hold `client` indefinitely, but it must never prevent terminating the +/// shared Agent runtime after the pool has been removed from routing. +pub struct PooledAgentClient { + client: tokio::sync::Mutex, + shared_runtime: Option>, + agent_session_id: Option, +} + +impl PooledAgentClient { + pub fn new(client: AgentDriverClient) -> Self { + let shared_runtime = client.shared_runtime.clone(); + let agent_session_id = client.agent_session_id.clone(); + Self { client: tokio::sync::Mutex::new(client), shared_runtime, agent_session_id } + } + + pub async fn lock(&self) -> tokio::sync::MutexGuard<'_, AgentDriverClient> { + self.client.lock().await + } + + pub fn try_lock(&self) -> Result, tokio::sync::TryLockError> { + self.client.try_lock() + } + + pub fn shares_runtime_with(&self, other: &Self) -> bool { + match (&self.shared_runtime, &other.shared_runtime) { + (Some(runtime), Some(other_runtime)) => Arc::ptr_eq(runtime, other_runtime), + _ => false, + } + } + + pub fn uses_runtime(&self, runtime: &Arc) -> bool { + self.shared_runtime.as_ref().is_some_and(|current| Arc::ptr_eq(current, runtime)) + } + + pub fn matches_session_id(&self, session_id: &str) -> bool { + self.agent_session_id.as_deref() == Some(session_id) + } + + pub fn is_runtime_available(&self) -> bool { + self.shared_runtime.as_ref().is_none_or(|runtime| !runtime.is_failed()) + } + + /// Terminates a protocol-v2 shared runtime without waiting for the logical session lock. + /// Legacy single-session clients can only be killed immediately when they are not busy. + pub fn fail_stop(&self) -> bool { + if let Some(runtime) = &self.shared_runtime { + runtime.kill(); + return true; + } + match self.client.try_lock() { + Ok(mut client) => { + client.kill(); + true + } + Err(_) => false, + } + } +} + #[derive(Debug, Clone, PartialEq, Eq)] pub struct AgentLaunchSpec { pub program: PathBuf, @@ -936,18 +1059,22 @@ impl AgentDriverClient { cancel_token: Option, ) -> Result { if let Some(runtime) = &self.shared_runtime { + let agent_session_id = self.agent_session_id.clone(); let mut params = params; if method != AgentMethod::Handshake.as_str() && method != AgentMethod::TestConnection.as_str() && method != AgentMethod::Shutdown.as_str() { - let session_id = self.agent_session_id.as_ref().ok_or("Shared Agent session id is missing")?; + let session_id = agent_session_id.as_ref().ok_or("Shared Agent session id is missing")?; params .as_object_mut() .ok_or_else(|| "Agent RPC parameters must be an object".to_string())? .insert("agentSessionId".to_string(), Value::String(session_id.clone())); } - return runtime.call(method, params, timeout_duration, cancel_token).await; + return runtime + .call(method, params, timeout_duration, cancel_token) + .await + .map_err(|error| with_agent_rpc_error_session_id(error, agent_session_id.as_deref())); } self.next_id += 1; let id = self.next_id; @@ -993,9 +1120,7 @@ impl AgentDriverClient { }; let result = if let Some(err) = resp.get("error") { - let msg = err.get("message").and_then(|m| m.as_str()).unwrap_or("Unknown agent error"); - let code = err.get("code").and_then(|c| c.as_i64()).unwrap_or(-1); - Err(format!("Agent RPC error ({code}): {msg}")) + Err(format_agent_rpc_error(err)) } else if let Some(result_val) = resp.get("result") { serde_json::from_value::(result_val.clone()) .map_err(|e| format!("Failed to deserialize agent result: {e}")) @@ -2178,14 +2303,15 @@ impl Drop for AgentDriverClient { mod tests { use super::{ agent_close_query_session_params, agent_handshake_params, agent_java_args, agent_java_args_with_extra, - agent_java_args_with_extra_opts, agent_object_source_params, agent_proxy_env_vars, agent_schema_params, - agent_schema_table_params, agent_supports_capability, agent_transaction_params, format_agent_process_error, + agent_java_args_with_extra_opts, agent_object_source_params, agent_proxy_env_vars, agent_rpc_error_category, + agent_rpc_error_session_id, agent_schema_params, agent_schema_table_params, agent_session_disposition, + agent_supports_capability, agent_transaction_params, decode_agent_response, format_agent_process_error, format_agent_startup_error, is_agent_rpc_response_error, is_unsupported_handshake_error, mongo_collection_params, mongo_database_params, mongo_document_id_params, parse_agent_java_opts, read_agent_line, start_stderr_collector, validate_dameng_java_system_properties, AgentCapability, AgentDriverClient, AgentHandshake, AgentKvMethod, AgentLaunchSpec, AgentMethod, AgentRuntimeClient, - AgentTableReadCloseParams, AgentTableReadPageParams, AgentTableReadStartParams, MongoAgentMethod, StderrTail, - AGENT_PROTOCOL_VERSION, + AgentSessionDisposition, AgentTableReadCloseParams, AgentTableReadPageParams, AgentTableReadStartParams, + MongoAgentMethod, StderrTail, AGENT_PROTOCOL_VERSION, }; use std::io::Cursor; use std::io::Write; @@ -2194,6 +2320,30 @@ mod tests { use std::time::{Duration, Instant}; use tokio_util::sync::CancellationToken; + #[test] + fn structured_agent_error_data_survives_legacy_string_boundary() { + let response = serde_json::json!({ + "error": { + "code": -1, + "message": "connection lost", + "data": { + "category": "connection", + "retryable": true, + "sessionDisposition": "quarantine", + "agentSessionId": "session-generation-1", + "stage": "execute" + } + } + }); + + let error = decode_agent_response::(response).unwrap_err(); + + assert!(error.starts_with("Agent RPC error (-1): connection lost")); + assert_eq!(agent_rpc_error_category(&error).as_deref(), Some("connection")); + assert_eq!(agent_rpc_error_session_id(&error).as_deref(), Some("session-generation-1")); + assert_eq!(agent_session_disposition(&error), Some(AgentSessionDisposition::Quarantine)); + } + #[test] fn agent_java_args_include_oracle_network_compatibility_flags() { let args = agent_java_args("/tmp/dbx-agent-oracle.jar"); @@ -2494,8 +2644,9 @@ for line in sys.stdin: "#, ) .unwrap(); + let python = if cfg!(windows) { "python" } else { "python3" }; let runtime = AgentRuntimeClient::spawn( - AgentLaunchSpec::new("python3").with_args([script_path.to_string_lossy().to_string()]), + AgentLaunchSpec::new(python).with_args([script_path.to_string_lossy().to_string()]), "test", ) .await @@ -2520,6 +2671,7 @@ for line in sys.stdin: .await .unwrap_err(); assert!(error.contains("Agent RPC call timed out")); + assert_eq!(agent_rpc_error_session_id(&error).as_deref(), Some("timeout-session")); let started = Instant::now(); client .call_with_timeout::("probe", serde_json::json!({}), Some(Duration::from_millis(500))) @@ -2536,6 +2688,7 @@ for line in sys.stdin: .await .unwrap_err(); assert!(error.contains("Agent RPC call timed out")); + assert_eq!(agent_rpc_error_session_id(&error).as_deref(), Some("timeout-session")); let started = Instant::now(); client .call_with_timeout::("probe", serde_json::json!({}), Some(Duration::from_millis(500))) @@ -2776,6 +2929,63 @@ for line in sys.stdin: let _ = std::fs::remove_file(script_path); } + #[tokio::test] + async fn disconnect_reports_runtime_replacement_without_owning_runtime_routing() { + let script_path = + std::env::temp_dir().join(format!("dbx-agent-runtime-cleanup-saturation-{}.py", uuid::Uuid::new_v4())); + 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) + if req['method'] == 'handshake': + response = { + 'jsonrpc': '2.0', + 'id': req['id'], + 'result': {'protocolVersion': 2, 'agentProtocolVersion': 2, 'capabilities': ['multi_session']} + } + elif req['method'] == 'close_session': + 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 = AgentRuntimeClient::spawn( + AgentLaunchSpec::new(python).with_args([script_path.to_string_lossy().to_string()]), + "test", + ) + .await + .unwrap(); + runtime.increment_session_count(); + let mut client = AgentDriverClient::shared_session(runtime.clone(), "session-1".to_string()); + + let error = client.disconnect().await.unwrap_err(); + + assert_eq!(agent_session_disposition(&error), Some(AgentSessionDisposition::ReplaceRuntime)); + assert!(!runtime.is_failed()); + runtime.kill(); + drop(client); + let _ = std::fs::remove_file(script_path); + } + #[tokio::test] async fn canceling_one_runtime_request_keeps_other_requests_alive() { let script_path = std::env::temp_dir().join(format!("dbx-agent-cancel-test-{}.py", uuid::Uuid::new_v4())); diff --git a/crates/dbx-core/src/query.rs b/crates/dbx-core/src/query.rs index 7792cd22b..79ecc5ee4 100644 --- a/crates/dbx-core/src/query.rs +++ b/crates/dbx-core/src/query.rs @@ -545,6 +545,9 @@ pub fn agent_close_query_session_params(session_id: &str) -> serde_json::Value { } pub fn is_connection_error(err: &str) -> bool { + if crate::db::agent_driver::agent_rpc_error_category(err).as_deref() == Some("connection") { + return true; + } let lower = err.to_lowercase(); if is_dbx_query_timeout_error(&lower) || is_agent_rpc_timeout_error(&lower) { return false; @@ -571,6 +574,9 @@ pub fn is_connection_error(err: &str) -> bool { || lower.contains("idle") || lower.contains("agent stdin not available") || lower.contains("agent stdout not available") + || lower.contains("agent runtime terminated") + || lower.contains("agent runtime is unavailable") + || lower.contains("agent runtime unavailable") || lower.contains("failed to write to agent stdin") || lower.contains("failed to flush agent stdin") || lower.contains("communicating with the server") @@ -590,6 +596,15 @@ fn is_schema_reset_cleanup_error(lower: &str) -> bool { } fn should_discard_agent_pool_after_error(err: &str) -> bool { + if matches!( + crate::db::agent_driver::agent_session_disposition(err), + Some( + crate::db::agent_driver::AgentSessionDisposition::Quarantine + | crate::db::agent_driver::AgentSessionDisposition::ReplaceRuntime + ) + ) { + return true; + } let lower = err.to_lowercase(); is_dbx_query_timeout_error(&lower) || is_agent_rpc_timeout_error(&lower) @@ -601,6 +616,19 @@ fn should_discard_agent_pool_after_error(err: &str) -> bool { } pub fn pool_error_action(db_type: Option, err: &str) -> PoolErrorAction { + if db_type.is_some_and(|db_type| database_capabilities::is_agent_type(&db_type)) + && matches!( + crate::db::agent_driver::agent_session_disposition(err), + Some( + crate::db::agent_driver::AgentSessionDisposition::Quarantine + | crate::db::agent_driver::AgentSessionDisposition::ReplaceRuntime + ) + ) + { + // The connection may be replaced, but the result of the user operation is unknown. + // Discard the session without replaying SQL, DDL, writes, or transactions. + return PoolErrorAction::Discard; + } let lower = err.to_lowercase(); if db::sqlserver::is_driver_panic_error(err) || (is_dbx_query_timeout_error(&lower) && should_discard_pool_after_query_timeout(db_type)) @@ -680,6 +708,41 @@ pub fn should_discard_pool_after_error(db_type: Option, err: &str) matches!(pool_error_action(db_type, err), PoolErrorAction::Discard | PoolErrorAction::ReconnectAndRetry) } +async fn discard_pool_after_error(state: &AppState, pool_key: &str, db_type: Option, error: &str) { + let action = pool_error_action(db_type, error); + if !matches!(action, PoolErrorAction::Discard | PoolErrorAction::ReconnectAndRetry) { + return; + } + + let replace_agent_runtime = matches!( + crate::db::agent_driver::agent_session_disposition(error), + Some(crate::db::agent_driver::AgentSessionDisposition::ReplaceRuntime) + ); + if replace_agent_runtime { + state.detach_pool_by_key(pool_key, true).await; + } else { + state.remove_pool_by_key(pool_key).await; + } +} + +async fn discard_agent_pool_after_error( + state: &AppState, + pool_key: &str, + client: &Arc, + db_type: Option, + error: &str, +) { + let action = pool_error_action(db_type, error); + if !matches!(action, PoolErrorAction::Discard | PoolErrorAction::ReconnectAndRetry) { + return; + } + let replace_agent_runtime = matches!( + crate::db::agent_driver::agent_session_disposition(error), + Some(crate::db::agent_driver::AgentSessionDisposition::ReplaceRuntime) + ); + state.detach_agent_pool_if_current(pool_key, client, replace_agent_runtime).await; +} + fn query_pool_error_action(db_type: Option, sql: &str, err: &str) -> PoolErrorAction { match pool_error_action(db_type, err) { // A connection error does not prove that the database did not receive @@ -1264,6 +1327,7 @@ pub async fn do_execute( } PoolKind::Agent(client) => { let client = client.clone(); + let source_client = client.clone(); let sql = sql_for_execution_context(pool_db_type, sql, schema); let database = database.map(|s| s.to_string()); let schema = schema_for_execution_context(pool_db_type, schema).map(|s| s.to_string()); @@ -1301,10 +1365,12 @@ pub async fn do_execute( .await .map(|result| truncate_result_with_max_rows(result, max_rows)); if matches!(result.as_ref(), Err(err) if err == QUERY_CANCELED) { - state.remove_pool_by_key(pool_key).await; + state.detach_agent_pool_if_current(pool_key, &source_client, false).await; } - if matches!(result.as_ref(), Err(err) if should_discard_pool_after_error(pool_db_type, err)) { - state.remove_pool_by_key(pool_key).await; + if let Err(err) = result.as_ref() { + if err != QUERY_CANCELED { + discard_agent_pool_after_error(state, pool_key, &source_client, pool_db_type, err).await; + } } result } @@ -1493,7 +1559,11 @@ pub async fn execute_sql_statement_with_options( ) } Some(PoolErrorAction::Discard) => { - state.remove_pool_by_key(&pool_key).await; + // Agent execution owns structured quarantine/runtime replacement before + // returning. Native drivers retain the existing caller-side cleanup. + if !db_type.is_some_and(|db_type| database_capabilities::is_agent_type(&db_type)) { + state.remove_pool_by_key(&pool_key).await; + } with_sql_context(result) } _ => with_sql_context(result), @@ -2169,6 +2239,7 @@ pub async fn execute_statements( } }; if let Some(client) = agent_client { + let source_client = client.clone(); check_read_only_for_connection_multi(state, &pool_key, statements).await?; let db_type = connection_database_type_for_pool_key(state, &pool_key).await; let execution_schema = schema_for_execution_context(db_type, schema); @@ -2183,6 +2254,7 @@ pub async fn execute_statements( let mut client = client.lock().await; let database = if database.trim().is_empty() { None } else { Some(database) }; let result = execute_multi_agent(&mut client, database, statements, execution_schema, timeout_secs).await; + drop(client); match result { Ok(result) => return Ok(db::QueryResult { execution_time_ms: start.elapsed().as_millis(), ..result }), Err(err) => { @@ -2191,12 +2263,8 @@ pub async fn execute_statements( "Agent does not support execute_batch; falling back to statement-by-statement execution" ); } else { - match pool_error_action(connection_database_type(state, connection_id).await, &err) { - PoolErrorAction::ReconnectAndRetry | PoolErrorAction::Discard => { - let _ = state.remove_pool_by_key(&pool_key).await; - } - PoolErrorAction::Keep => {} - } + let db_type = connection_database_type(state, connection_id).await; + discard_agent_pool_after_error(state, &pool_key, &source_client, db_type, &err).await; return Err(query_error_with_omitted_sql_context(&err, sql_ctx)); } } @@ -2220,15 +2288,18 @@ pub async fn execute_statements( total_affected += result.affected_rows; } Err(e) => { - match pool_error_action(connection_database_type(state, connection_id).await, &e) { + let db_type = connection_database_type(state, connection_id).await; + match pool_error_action(db_type, &e) { PoolErrorAction::ReconnectAndRetry => { let db_opt = if database.is_empty() { None } else { Some(database) }; let _ = state.reconnect_pool(connection_id, db_opt).await; } - PoolErrorAction::Discard => { + PoolErrorAction::Discard + if !db_type.is_some_and(|db_type| database_capabilities::is_agent_type(&db_type)) => + { let _ = state.remove_pool_by_key(&pool_key).await; } - PoolErrorAction::Keep => {} + PoolErrorAction::Discard | PoolErrorAction::Keep => {} } return Err(query_error_with_omitted_sql_context( &format!("Statement {} failed: {}. Previous {} statement(s) may have been committed.", i + 1, e, i), @@ -2560,11 +2631,10 @@ pub async fn execute_statements_in_transaction_on_pool( PoolKind::Mysql(mp, _mode) => TxPath::Mysql(mp.clone(), false), PoolKind::Sqlite(sq) => TxPath::Sqlite(sq.clone()), PoolKind::CloudflareD1(client) => TxPath::CloudflareD1(client.clone()), - PoolKind::ClickHouse(_) - | PoolKind::Rqlite(_) - | PoolKind::Turso(_) - | PoolKind::SqlServer(_) - | PoolKind::Agent(_) => TxPath::Explicit, + PoolKind::ClickHouse(_) | PoolKind::Rqlite(_) | PoolKind::Turso(_) | PoolKind::SqlServer(_) => { + TxPath::Explicit + } + PoolKind::Agent(client) => TxPath::Agent(client.clone()), PoolKind::MessageQueue | PoolKind::Nacos | PoolKind::HBase(_) => TxPath::None, PoolKind::DuckDbWorker(_) | PoolKind::Redis(_) @@ -2606,6 +2676,13 @@ pub async fn execute_statements_in_transaction_on_pool( ) .await } + Some(TxPath::Agent(client)) => { + let result = exec_tx_agent_inner(client.clone(), db_type, Some(database), statements, schema, start).await; + if let Err(error) = result.as_ref() { + discard_agent_pool_after_error(state, pool_key, &client, db_type, error).await; + } + return result; + } Some(TxPath::Explicit) => { let mysql_dialect = connection_mysql_query_dialect(state, connection_id).await; exec_tx_explicit_inner(state, pool_key, mysql_dialect, Some(database), statements, schema, start).await @@ -2618,9 +2695,7 @@ pub async fn execute_statements_in_transaction_on_pool( }; if let Err(err) = result.as_ref() { - if matches!(pool_error_action(db_type, err), PoolErrorAction::Discard | PoolErrorAction::ReconnectAndRetry) { - state.remove_pool_by_key(pool_key).await; - } + discard_pool_after_error(state, pool_key, db_type, err).await; } result @@ -2632,6 +2707,7 @@ enum TxPath { Mysql(mysql_async::Pool, bool), Sqlite(db::sqlite::SqliteHandle), CloudflareD1(db::cloudflare_d1_driver::CloudflareD1Client), + Agent(Arc), Explicit, None, } @@ -2859,24 +2935,6 @@ async fn exec_tx_explicit_inner( schema: Option<&str>, start: std::time::Instant, ) -> Result { - let conns = state.connections.read().await; - if let Some(crate::connection::PoolKind::Agent(client)) = conns.get(pool_key) { - let db_type = connection_database_type_for_pool_key(state, pool_key).await; - let execution_schema = schema_for_execution_context(db_type, schema); - let rewritten_statements; - let statements = if qualifies_unqualified_agent_relations(db_type) { - rewritten_statements = - statements.iter().map(|sql| sql_for_execution_context(db_type, sql, schema)).collect::>(); - rewritten_statements.as_slice() - } else { - statements - }; - let mut client = client.lock().await; - let result: db::QueryResult = client.execute_transaction(database, statements, execution_schema).await?; - return Ok(db::QueryResult { execution_time_ms: start.elapsed().as_millis(), ..result }); - } - drop(conns); - do_execute( state, pool_key, @@ -2938,6 +2996,28 @@ async fn exec_tx_explicit_inner( }) } +async fn exec_tx_agent_inner( + client: Arc, + db_type: Option, + database: Option<&str>, + statements: &[String], + schema: Option<&str>, + start: std::time::Instant, +) -> Result { + let execution_schema = schema_for_execution_context(db_type, schema); + let rewritten_statements; + let statements = if qualifies_unqualified_agent_relations(db_type) { + rewritten_statements = + statements.iter().map(|sql| sql_for_execution_context(db_type, sql, schema)).collect::>(); + rewritten_statements.as_slice() + } else { + statements + }; + let mut client = client.lock().await; + let result: db::QueryResult = client.execute_transaction(database, statements, execution_schema).await?; + Ok(db::QueryResult { execution_time_ms: start.elapsed().as_millis(), ..result }) +} + async fn exec_tx_none_inner( state: &AppState, pool_key: &str, @@ -3884,6 +3964,141 @@ for line in sys.stdin: } } + async fn agent_error_state( + disposition: &str, + ) -> (AppState, std::path::PathBuf, std::sync::Arc) { + let dir = std::env::temp_dir().join(format!("dbx-query-agent-error-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&dir).unwrap(); + let script_path = dir.join("agent.py"); + std::fs::write( + &script_path, + format!( + r#"import json, sys +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'] in ('execute_query', 'execute_batch', 'execute_transaction'): + response = {{ + 'jsonrpc': '2.0', + 'id': req['id'], + 'error': {{ + 'code': -1, + 'message': 'injected Agent failure', + 'data': {{ + 'category': 'resource', + 'retryable': False, + 'sessionDisposition': '{disposition}', + 'stage': 'execute' + }} + }} + }} + 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 storage = Storage::open(&dir.join("storage.db")).await.unwrap(); + let state = AppState::new(storage); + state.configs.write().await.insert("conn-1".to_string(), test_connection_config(DatabaseType::Dameng)); + state.connections.write().await.insert( + "conn-1".to_string(), + PoolKind::agent(crate::db::agent_driver::AgentDriverClient::shared_session( + runtime.clone(), + "session-1".to_string(), + )), + ); + + (state, dir, runtime) + } + + #[tokio::test] + async fn agent_query_replace_runtime_error_detaches_pool_and_stops_runtime() { + let (state, dir, runtime) = agent_error_state("replace_runtime").await; + + let error = execute_sql_statement(&state, "conn-1", "", "SELECT 1", None, None).await.unwrap_err(); + + assert!(error.contains("injected Agent failure")); + assert!(!state.connections.read().await.contains_key("conn-1")); + assert!(runtime.is_failed()); + + runtime.kill(); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn agent_transaction_replace_runtime_error_detaches_pool_and_stops_runtime() { + let (state, dir, runtime) = agent_error_state("replace_runtime").await; + + let error = execute_statements_in_transaction_on_pool( + &state, + "conn-1", + "conn-1", + "", + &["UPDATE test_table SET value = 1".to_string()], + None, + None, + ) + .await + .unwrap_err(); + + assert!(error.contains("injected Agent failure")); + assert!(!state.connections.read().await.contains_key("conn-1")); + assert!(runtime.is_failed()); + + runtime.kill(); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn agent_batch_replace_runtime_error_detaches_pool_and_stops_runtime() { + let (state, dir, runtime) = agent_error_state("replace_runtime").await; + + let error = + execute_statements(&state, "conn-1", "", &["UPDATE test_table SET value = 1".to_string()], None, None) + .await + .unwrap_err(); + + assert!(error.contains("injected Agent failure")); + assert!(!state.connections.read().await.contains_key("conn-1")); + assert!(runtime.is_failed()); + + runtime.kill(); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn agent_quarantine_error_removes_only_target_pool() { + let (state, dir, runtime) = agent_error_state("quarantine").await; + + let error = execute_sql_statement(&state, "conn-1", "", "SELECT 1", None, None).await.unwrap_err(); + + assert!(error.contains("injected Agent failure")); + assert!(!state.connections.read().await.contains_key("conn-1")); + assert!(!runtime.is_failed()); + + runtime.kill(); + let _ = std::fs::remove_dir_all(dir); + } + struct FakeMysqlBatchExecutor { outcomes: std::collections::VecDeque>, executed: Vec, @@ -5010,16 +5225,38 @@ for line in sys.stdin: ); } + #[test] + fn structured_agent_disposition_controls_pool_recovery() { + let quarantined = "Agent RPC error (-1): lost\nDBX_AGENT_ERROR_DATA:{\"category\":\"connection\",\"sessionDisposition\":\"quarantine\"}"; + let replace_runtime = "Agent RPC error (-1): saturated\nDBX_AGENT_ERROR_DATA:{\"category\":\"resource\",\"sessionDisposition\":\"replace_runtime\"}"; + + assert!(should_discard_agent_pool_after_error(quarantined)); + assert!(should_discard_agent_pool_after_error(replace_runtime)); + assert!(is_connection_error(quarantined)); + assert_eq!(pool_error_action(Some(DatabaseType::Oracle), quarantined), PoolErrorAction::Discard); + assert_eq!(pool_error_action(Some(DatabaseType::Oracle), replace_runtime), PoolErrorAction::Discard); + } + #[test] fn unavailable_agent_pipes_are_reconnectable_errors() { assert!(should_discard_agent_pool_after_error("Agent stdin not available")); assert!(should_discard_agent_pool_after_error("Agent stdout not available")); assert!(is_connection_error("Agent stdin not available")); assert!(is_connection_error("Agent stdout not available")); + assert!(is_connection_error("Agent runtime terminated")); + assert!(is_connection_error("Agent runtime is unavailable")); assert_eq!( pool_error_action(Some(DatabaseType::Oracle), "Agent stdin not available"), PoolErrorAction::ReconnectAndRetry ); + assert_eq!( + pool_error_action(Some(DatabaseType::Oracle), "Agent runtime terminated"), + PoolErrorAction::ReconnectAndRetry + ); + assert_eq!( + pool_error_action(Some(DatabaseType::Oracle), "Agent runtime is unavailable"), + PoolErrorAction::ReconnectAndRetry + ); } #[test] diff --git a/crates/dbx-core/src/schema.rs b/crates/dbx-core/src/schema.rs index 5d68d6c3f..58f5c27c6 100644 --- a/crates/dbx-core/src/schema.rs +++ b/crates/dbx-core/src/schema.rs @@ -43,7 +43,7 @@ impl EphemeralAgentMetadataSession { let client_session_id = ephemeral_agent_metadata_session_id(db_config.as_ref(), task_kind); let cleanup_guard = match client_session_id.as_deref() { Some(client_session_id) => { - state.client_session_pool_cleanup_guard(connection_id, database, client_session_id).await + state.metadata_session_pool_cleanup_guard(connection_id, database, client_session_id).await } None => None, }; @@ -350,7 +350,7 @@ pub async fn list_sqlserver_linked_server_tables_core( /// use the MySQL protocol, so this is a defensive no-op); the caller's /// flat-sidebar fallback then renders the standard database list. pub async fn list_doris_catalogs_core(state: &AppState, connection_id: &str) -> Result, String> { - let pool_key = state.get_or_create_pool(connection_id, None).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, None, None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; if let Some(PoolKind::Mysql(p, _)) = connections.get(&pool_key) { @@ -373,7 +373,7 @@ pub async fn list_doris_catalog_databases_core( connection_id: &str, catalog: &str, ) -> Result, String> { - let pool_key = state.get_or_create_pool(connection_id, None).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, None, None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; let pool = connections.get(&pool_key).ok_or("Pool not found")?; @@ -417,7 +417,7 @@ pub async fn list_doris_catalog_tables_core( object_types: Option<&[String]>, table_name_filter: Option<&TableNameFilter>, ) -> Result, String> { - let pool_key = state.get_or_create_pool(connection_id, None).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, None, None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; let pool = connections.get(&pool_key).ok_or("Pool not found")?; @@ -440,7 +440,7 @@ pub async fn get_doris_catalog_columns_core( database: &str, table: &str, ) -> Result, String> { - let pool_key = state.get_or_create_pool(connection_id, None).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, None, None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; let pool = connections.get(&pool_key).ok_or("Pool not found")?; @@ -463,7 +463,7 @@ pub async fn get_doris_catalog_table_ddl_core( database: &str, table: &str, ) -> Result { - let pool_key = state.get_or_create_pool(connection_id, None).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, None, None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; let pool = connections.get(&pool_key).ok_or("Pool not found")?; @@ -485,7 +485,7 @@ pub async fn list_doris_catalog_indexes_core( database: &str, table: &str, ) -> Result, String> { - let pool_key = state.get_or_create_pool(connection_id, None).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, None, None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; let pool = connections.get(&pool_key).ok_or("Pool not found")?; @@ -654,7 +654,7 @@ async fn list_schema_infos_once( connection_id: &str, database: &str, ) -> Result, String> { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; let show_system_schemas = db_config.as_ref().is_some_and(|config| config.show_system_schemas); { @@ -674,7 +674,7 @@ pub async fn list_data_types_core( database: &str, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; if let Some(PoolKind::ExternalDriver { config, session, .. }) = connections.get(&pool_key) { @@ -725,7 +725,7 @@ async fn list_schemas_once( database: &str, apply_visible_filter: bool, ) -> Result, String> { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; let show_system_schemas = db_config.as_ref().is_some_and(|config| config.show_system_schemas); let visible_schema_filter = visible_schema_filter(db_config.as_ref(), database, apply_visible_filter); @@ -892,8 +892,13 @@ pub async fn list_vector_collections_core( connection_id: &str, database: &str, ) -> Result, String> { - let pool_key = - state.get_or_create_pool(connection_id, if database.is_empty() { None } else { Some(database) }).await?; + let pool_key = state + .get_or_create_metadata_pool_for_session( + connection_id, + if database.is_empty() { None } else { Some(database) }, + None, + ) + .await?; let client = { let connections = state.connections.read().await; match connections.get(&pool_key) { @@ -911,8 +916,13 @@ pub async fn get_vector_collection_detail_core( database: &str, collection: &str, ) -> Result { - let pool_key = - state.get_or_create_pool(connection_id, if database.is_empty() { None } else { Some(database) }).await?; + let pool_key = state + .get_or_create_metadata_pool_for_session( + connection_id, + if database.is_empty() { None } else { Some(database) }, + None, + ) + .await?; let client = { let connections = state.connections.read().await; match connections.get(&pool_key) { @@ -958,7 +968,8 @@ async fn get_table_comment_core_for_session( client_session_id: Option<&str>, ) -> Result, String> { retry_metadata_connection_for_session(state, connection_id, Some(database), client_session_id, || async { - let pool_key = state.get_or_create_pool_for_session(connection_id, Some(database), client_session_id).await?; + let pool_key = + state.get_or_create_metadata_pool_for_session(connection_id, Some(database), client_session_id).await?; let db_config = connection_config(state, connection_id).await; { @@ -1656,7 +1667,7 @@ async fn load_oracle_table_comments_for_objects( } async fn oracle_agent_list_object_statistics( - client: Arc>, + client: Arc, database: &str, schema: &str, timeout_duration: Option, @@ -1696,7 +1707,7 @@ async fn oracle_agent_list_object_statistics( } async fn dameng_agent_list_object_statistics( - client: Arc>, + client: Arc, database: &str, schema: &str, timeout_duration: Option, @@ -1735,7 +1746,7 @@ async fn dameng_agent_list_object_statistics( } async fn agent_list_object_statistics( - client: Arc>, + client: Arc, database: &str, schema: &str, sql: String, @@ -1778,7 +1789,8 @@ async fn list_tables_once( table_name_filter: Option<&TableNameFilter>, client_session_id: Option<&str>, ) -> Result, String> { - let pool_key = state.get_or_create_pool_for_session(connection_id, Some(database), client_session_id).await?; + let pool_key = + state.get_or_create_metadata_pool_for_session(connection_id, Some(database), client_session_id).await?; let db_config = connection_config(state, connection_id).await; { @@ -2518,18 +2530,19 @@ mod tests { ephemeral_agent_metadata_session_id, external_driver_uses_mysql_ddl, filter_mongodb_agent_collections, filter_mysql_system_databases_for_config, filter_object_infos, filter_table_infos, filter_visible_schema_names, gbase8a_object_statistics_sql, is_agent_postgres_metadata_fallback_config, is_mysql_external_driver_config, - is_retryable_metadata_error, metadata_name_or_comment_matches, mysql_external_driver_ddl_from_query_result, - mysql_external_driver_ddl_sql, mysql_object_source_ddl_column_index, mysql_object_source_sql, - mysql_table_metadata_catalog, normalize_information_schema_table_type, oracle_columns_from_query_result, - oracle_columns_sql, oracle_object_statistics_dba_segments_sql, oracle_object_statistics_from_query_result, + is_retryable_metadata_error, metadata_error_action, metadata_name_or_comment_matches, + mysql_external_driver_ddl_from_query_result, mysql_external_driver_ddl_sql, + mysql_object_source_ddl_column_index, mysql_object_source_sql, mysql_table_metadata_catalog, + normalize_information_schema_table_type, oracle_columns_from_query_result, oracle_columns_sql, + oracle_object_statistics_dba_segments_sql, oracle_object_statistics_from_query_result, oracle_object_statistics_rows_only_sql, oracle_object_statistics_sql, oracle_object_statistics_user_segments_sql, oracle_table_comment_from_query_result, oracle_table_comment_sql, oracle_table_comments_sql, presto_like_columns_from_query_result, presto_like_information_schema_columns_sql, - presto_like_information_schema_tables_sql, presto_like_tables_from_query_result, + presto_like_information_schema_tables_sql, presto_like_tables_from_query_result, replace_metadata_runtime, should_query_oracle_columns_via_sql_first, table_comments_from_query_result, table_name_filter_matches, tdengine_table_comment_like_pattern, tdengine_table_comment_sql, tdengine_table_comments_sql, - uses_mongodb_agent_collection_listing, visible_schema_filter, TableNameFilter, TDENGINE_COMMENT_SEARCH_TIMEOUT, - TDENGINE_LIKE_PATTERN_MAX_BYTES, + uses_mongodb_agent_collection_listing, visible_schema_filter, MetadataErrorAction, TableNameFilter, + TDENGINE_COMMENT_SEARCH_TIMEOUT, TDENGINE_LIKE_PATTERN_MAX_BYTES, }; use super::{list_databases_core, list_tables_core}; use crate::connection::{AppState, PoolKind}; @@ -2868,10 +2881,184 @@ mod tests { assert!(is_retryable_metadata_error("Pool not found")); assert!(is_retryable_metadata_error("connection reset by peer")); assert!(is_retryable_metadata_error("Agent RPC error (-1): dm.jdbc.driver.DMException: 网络通信异常")); + assert!(is_retryable_metadata_error( + "Agent RPC error (-1): connection lost\nDBX_AGENT_ERROR_DATA:{\"category\":\"connection\",\"sessionDisposition\":\"quarantine\"}" + )); + assert!(!is_retryable_metadata_error( + "Agent RPC error (-1): connection text in SQL error\nDBX_AGENT_ERROR_DATA:{\"category\":\"sql\",\"sessionDisposition\":\"keep\"}" + )); + assert!(!is_retryable_metadata_error( + "Agent RPC error (-1): connection kept\nDBX_AGENT_ERROR_DATA:{\"category\":\"connection\",\"sessionDisposition\":\"keep\"}" + )); + assert!(!is_retryable_metadata_error( + "Agent RPC error (-1): runtime saturated\nDBX_AGENT_ERROR_DATA:{\"category\":\"resource\",\"sessionDisposition\":\"replace_runtime\"}" + )); assert!(!is_retryable_metadata_error("Unknown column 'email' in 'field list'")); assert!(!is_retryable_metadata_error("Access denied for user")); } + #[test] + fn metadata_error_action_applies_fail_stop_to_every_attempt() { + let quarantine = "Agent RPC error (-1): connection lost\nDBX_AGENT_ERROR_DATA:{\"category\":\"connection\",\"sessionDisposition\":\"quarantine\"}"; + let replace_runtime = "Agent RPC error (-1): runtime saturated\nDBX_AGENT_ERROR_DATA:{\"category\":\"resource\",\"sessionDisposition\":\"replace_runtime\"}"; + let sql = "Agent RPC error (-1): syntax error\nDBX_AGENT_ERROR_DATA:{\"category\":\"sql\",\"sessionDisposition\":\"keep\"}"; + let db_type = Some(DatabaseType::Dameng); + + assert_eq!(metadata_error_action(db_type, quarantine, false), MetadataErrorAction::Retry); + assert_eq!(metadata_error_action(db_type, quarantine, true), MetadataErrorAction::Discard); + assert_eq!( + metadata_error_action(db_type, "Agent RPC call timed out (30s)", false), + MetadataErrorAction::Discard + ); + assert_eq!(metadata_error_action(db_type, replace_runtime, false), MetadataErrorAction::ReplaceRuntime); + assert_eq!(metadata_error_action(db_type, replace_runtime, true), MetadataErrorAction::ReplaceRuntime); + assert_eq!(metadata_error_action(db_type, sql, false), MetadataErrorAction::Return); + } + + #[tokio::test] + async fn metadata_fail_stop_detaches_base_pool_without_client_session() { + let dir = std::env::temp_dir().join(format!("dbx-schema-metadata-fail-stop-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&dir).unwrap(); + let storage = crate::storage::Storage::open(&dir.join("storage.db")).await.unwrap(); + let state = crate::connection::AppState::new(storage); + let mut config = test_connection_config(DatabaseType::Dameng); + config.id = "conn".to_string(); + state.configs.write().await.insert(config.id.clone(), config); + state.connections.write().await.insert( + "conn:analytics:role:metadata".to_string(), + super::PoolKind::agent(crate::db::agent_driver::AgentDriverClient::test_stub()), + ); + + replace_metadata_runtime(&state, "conn", Some("analytics"), None).await; + + assert!(!state.connections.read().await.contains_key("conn:analytics:role:metadata")); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn metadata_timeout_detaches_pool_without_replaying_operation() { + let dir = std::env::temp_dir().join(format!("dbx-schema-metadata-timeout-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&dir).unwrap(); + let storage = crate::storage::Storage::open(&dir.join("storage.db")).await.unwrap(); + let state = crate::connection::AppState::new(storage); + let mut config = test_connection_config(DatabaseType::Dameng); + config.id = "conn".to_string(); + state.configs.write().await.insert(config.id.clone(), config); + let pool = crate::db::sqlite::connect_path(":memory:").await.unwrap(); + state.connections.write().await.insert("conn:role:metadata".to_string(), super::PoolKind::Sqlite(pool)); + let mut attempts = 0; + + let result = super::retry_metadata_connection_for_session(&state, "conn", None, None, || { + attempts += 1; + async { Err::<(), _>("Agent RPC call timed out (30s)".to_string()) } + }) + .await; + + assert_eq!(result.unwrap_err(), "Agent RPC call timed out (30s)"); + assert_eq!(attempts, 1); + assert!(!state.connections.read().await.contains_key("conn:role:metadata")); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn metadata_second_quarantine_detaches_replacement_pool() { + let dir = std::env::temp_dir().join(format!("dbx-schema-metadata-quarantine-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&dir).unwrap(); + let storage = crate::storage::Storage::open(&dir.join("storage.db")).await.unwrap(); + let state = crate::connection::AppState::new(storage); + let mut config = test_connection_config(DatabaseType::Sqlite); + config.id = "conn".to_string(); + config.host = ":memory:".to_string(); + config.password.clear(); + config.database = None; + state.configs.write().await.insert(config.id.clone(), config); + let pool = crate::db::sqlite::connect_path(":memory:").await.unwrap(); + state.connections.write().await.insert("conn".to_string(), super::PoolKind::Sqlite(pool)); + let mut attempts = 0; + let quarantine = "Agent RPC error (-1): connection lost\nDBX_AGENT_ERROR_DATA:{\"category\":\"connection\",\"sessionDisposition\":\"quarantine\"}"; + + let result = super::retry_metadata_connection_for_session(&state, "conn", None, None, || { + attempts += 1; + async { Err::<(), _>(quarantine.to_string()) } + }) + .await; + + assert_eq!(result.unwrap_err(), quarantine); + assert_eq!(attempts, 2); + assert!(!state.connections.read().await.contains_key("conn")); + let _ = std::fs::remove_dir_all(dir); + } + + #[tokio::test] + async fn table_ddl_timeout_detaches_metadata_pool_without_replay() { + let dir = std::env::temp_dir().join(format!("dbx-schema-table-ddl-timeout-{}", uuid::Uuid::new_v4())); + std::fs::create_dir_all(&dir).unwrap(); + let script_path = dir.join("table-ddl-timeout-agent.py"); + let call_count_path = dir.join("table-ddl-call-count"); + let call_count = serde_json::to_string(&call_count_path.to_string_lossy()).unwrap(); + std::fs::write( + &script_path, + format!( + r#"import json, pathlib, sys +call_count = pathlib.Path({call_count}) +print(json.dumps({{'ready': True}}), flush=True) +for line in sys.stdin: + req = json.loads(line) + if req['method'] == 'handshake': + result = {{'protocolVersion': 2, 'agentProtocolVersion': 2, 'capabilities': ['multi_session']}} + response = {{'jsonrpc': '2.0', 'id': req['id'], 'result': result}} + elif req['method'] in ('validate_session', 'validate_connection'): + response = {{'jsonrpc': '2.0', 'id': req['id'], 'result': {{}}}} + else: + previous = int(call_count.read_text()) if call_count.exists() else 0 + call_count.write_text(str(previous + 1)) + response = {{ + 'jsonrpc': '2.0', + 'id': req['id'], + 'error': {{ + 'code': -1, + 'message': 'metadata timed out', + 'data': {{ + 'category': 'timeout', + 'retryable': False, + 'sessionDisposition': 'quarantine', + 'stage': 'execute' + }} + }} + }} + 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 = + crate::db::agent_driver::AgentDriverClient::shared_session(runtime.clone(), "metadata-session".to_string()); + let storage = crate::storage::Storage::open(&dir.join("storage.db")).await.unwrap(); + let state = crate::connection::AppState::new(storage); + let mut config = test_connection_config(DatabaseType::Dameng); + config.id = "conn".to_string(); + state.configs.write().await.insert(config.id.clone(), config); + let pool_key = "conn:analytics:role:metadata"; + state.connections.write().await.insert(pool_key.to_string(), super::PoolKind::agent(client)); + + let error = super::get_table_ddl_core(&state, "conn", "analytics", "APP", "EVENTS", None).await.unwrap_err(); + + assert_eq!(crate::db::agent_driver::agent_rpc_error_category(&error).as_deref(), Some("timeout")); + assert_eq!(std::fs::read_to_string(call_count_path).unwrap(), "1"); + assert!(!state.connections.read().await.contains_key(pool_key)); + runtime.kill(); + let _ = std::fs::remove_dir_all(dir); + } + #[test] fn visible_schema_filter_only_applies_when_requested() { let mut config = test_connection_config(DatabaseType::Oracle); @@ -4021,7 +4208,7 @@ async fn close_ephemeral_agent_metadata_session( let Some(client_session_id) = client_session_id else { return true; }; - match state.close_client_session_pool(connection_id, database, client_session_id).await { + match state.close_metadata_session_pool(connection_id, database, client_session_id).await { Ok(_) => true, Err(error) => { log::warn!( @@ -4047,7 +4234,9 @@ pub async fn completion_assistant_search_core( request.max_results ); retry_metadata_connection(state, &request.connection_id, Some(&request.database), || async { - let pool_key = state.get_or_create_pool(&request.connection_id, Some(&request.database)).await?; + let pool_key = state + .get_or_create_metadata_pool_for_session(&request.connection_id, Some(&request.database), None) + .await?; log::debug!("[schema][completion_assistant:start] {request_summary}"); { let connections = state.connections.read().await; @@ -4280,7 +4469,7 @@ async fn list_object_statistics_once( database: &str, schema: &str, ) -> Result, String> { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; try_sqlserver!(connections, &pool_key, list_object_statistics, schema); @@ -4370,7 +4559,8 @@ async fn list_objects_once( object_types: Option<&[String]>, client_session_id: Option<&str>, ) -> Result { - let pool_key = state.get_or_create_pool_for_session(connection_id, Some(database), client_session_id).await?; + let pool_key = + state.get_or_create_metadata_pool_for_session(connection_id, Some(database), client_session_id).await?; let db_config = connection_config(state, connection_id).await; let (mysql_limit, mysql_offset) = if filter.is_none_or(|value| value.trim().is_empty()) { (limit, offset) } else { (None, None) }; @@ -4567,7 +4757,8 @@ async fn list_completion_objects_once( schema: &str, client_session_id: Option<&str>, ) -> Result, String> { - let pool_key = state.get_or_create_pool_for_session(connection_id, Some(database), client_session_id).await?; + let pool_key = + state.get_or_create_metadata_pool_for_session(connection_id, Some(database), client_session_id).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; @@ -4726,17 +4917,120 @@ where F: FnMut() -> Fut, Fut: Future>, { - let result = operation().await; - match result { - Err(error) if is_retryable_metadata_error(&error) => { - state.reconnect_pool_for_session(connection_id, database, client_session_id).await?; - operation().await + let db_type = { + let configs = state.configs.read().await; + configs.get(connection_id).map(|config| config.db_type) + }; + let mut retried = false; + loop { + let result = operation().await; + let action = result + .as_ref() + .err() + .map(|error| metadata_error_action(db_type, error, retried)) + .unwrap_or(MetadataErrorAction::Return); + match action { + MetadataErrorAction::ReplaceRuntime => { + state + .detach_metadata_pool_after_error( + connection_id, + database, + client_session_id, + result.as_ref().err().expect("replace-runtime action requires an error"), + true, + ) + .await; + return result; + } + MetadataErrorAction::Discard => { + state + .detach_metadata_pool_after_error( + connection_id, + database, + client_session_id, + result.as_ref().err().expect("discard action requires an error"), + false, + ) + .await; + return result; + } + MetadataErrorAction::Retry => { + retried = true; + if let Err(error) = + state.reconnect_metadata_pool_for_session(connection_id, database, client_session_id).await + { + match metadata_error_action(db_type, &error, true) { + MetadataErrorAction::ReplaceRuntime => { + state + .detach_metadata_pool_after_error( + connection_id, + database, + client_session_id, + &error, + true, + ) + .await; + } + MetadataErrorAction::Retry | MetadataErrorAction::Discard => { + state + .detach_metadata_pool_after_error( + connection_id, + database, + client_session_id, + &error, + false, + ) + .await; + } + MetadataErrorAction::Return => {} + } + return Err(error); + } + } + MetadataErrorAction::Return => return result, } - _ => result, } } +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum MetadataErrorAction { + Retry, + Discard, + ReplaceRuntime, + Return, +} + +fn metadata_error_action(db_type: Option, error: &str, retried: bool) -> MetadataErrorAction { + if crate::db::agent_driver::agent_session_disposition(error) + == Some(crate::db::agent_driver::AgentSessionDisposition::ReplaceRuntime) + { + MetadataErrorAction::ReplaceRuntime + } else if !retried && is_retryable_metadata_error(error) { + MetadataErrorAction::Retry + } else if should_discard_pool_after_error(db_type, error) { + MetadataErrorAction::Discard + } else { + MetadataErrorAction::Return + } +} + +#[cfg(test)] +async fn replace_metadata_runtime( + state: &AppState, + connection_id: &str, + database: Option<&str>, + client_session_id: Option<&str>, +) { + state.replace_runtime_for_metadata_pool(connection_id, database, client_session_id).await; +} + fn is_retryable_metadata_error(error: &str) -> bool { + let category = crate::db::agent_driver::agent_rpc_error_category(error); + if let Some(category) = category { + return category == "connection" + && crate::db::agent_driver::agent_session_disposition(error) + == Some(crate::db::agent_driver::AgentSessionDisposition::Quarantine); + } error == "Pool not found" || crate::query::is_connection_error(error) } @@ -4791,7 +5085,7 @@ async fn get_columns_core_for_session_inner( let context_session_id = if use_client_session_context { client_session_id } else { None }; retry_metadata_connection_for_session(state, connection_id, Some(database), client_session_id, || async { let pool_key = state - .get_or_create_pool_for_session(connection_id, Some(database), client_session_id) + .get_or_create_metadata_pool_for_session(connection_id, Some(database), client_session_id) .await?; let db_config = connection_config(state, connection_id).await; @@ -5074,7 +5368,7 @@ pub async fn get_sqlserver_column_metadata_core( table: &str, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let connections = state.connections.read().await; try_sqlserver!(connections, &pool_key, get_column_metadata, schema, table); Err("SQL Server column metadata requires a native SQL Server connection".to_string()) @@ -5158,7 +5452,8 @@ async fn list_indexes_core_for_session( client_session_id: Option<&str>, ) -> Result, String> { retry_metadata_connection_for_session(state, connection_id, Some(database), client_session_id, || async { - let pool_key = state.get_or_create_pool_for_session(connection_id, Some(database), client_session_id).await?; + let pool_key = + state.get_or_create_metadata_pool_for_session(connection_id, Some(database), client_session_id).await?; let db_config = connection_config(state, connection_id).await; { @@ -5238,7 +5533,8 @@ async fn list_foreign_keys_core_for_session( client_session_id: Option<&str>, ) -> Result, String> { retry_metadata_connection_for_session(state, connection_id, Some(database), client_session_id, || async { - let pool_key = state.get_or_create_pool_for_session(connection_id, Some(database), client_session_id).await?; + let pool_key = + state.get_or_create_metadata_pool_for_session(connection_id, Some(database), client_session_id).await?; let db_config = connection_config(state, connection_id).await; { @@ -5286,7 +5582,7 @@ pub async fn list_triggers_core( return Ok(vec![]); } retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; { @@ -5332,7 +5628,7 @@ pub async fn list_constraints_core( table: &str, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; if let Some(client) = extract_pool!(&connections, &pool_key, Agent) { @@ -5353,7 +5649,7 @@ pub async fn list_partitions_core( table: &str, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; if let Some(client) = extract_pool!(&connections, &pool_key, Agent) { @@ -5374,7 +5670,7 @@ pub async fn list_subpartitions_core( table: &str, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; if let Some(client) = extract_pool!(&connections, &pool_key, Agent) { @@ -5396,7 +5692,7 @@ pub async fn list_functions_core( schema: &str, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let connections = state.connections.read().await; let pool = connections.get(&pool_key).ok_or("Pool not found")?; @@ -5416,7 +5712,7 @@ pub async fn list_sequences_core( with_last_values: bool, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; let connections = state.connections.read().await; let pool = connections.get(&pool_key).ok_or("Pool not found")?; @@ -5439,7 +5735,7 @@ pub async fn list_rules_core( schema: &str, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let connections = state.connections.read().await; let pool = connections.get(&pool_key).ok_or("Pool not found")?; @@ -5458,7 +5754,7 @@ pub async fn list_extensions_core( schema: Option<&str>, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; if db_config.as_ref().is_some_and(|config| config.db_type == DatabaseType::Kingbase) { let connections = state.connections.read().await; @@ -5486,7 +5782,7 @@ pub async fn list_available_extensions_core( database: &str, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; if db_config.as_ref().is_some_and(|config| config.db_type == DatabaseType::Kingbase) { let connections = state.connections.read().await; @@ -5519,7 +5815,7 @@ pub async fn list_owners_core( schema: &str, ) -> Result, String> { retry_metadata_connection(state, connection_id, Some(database), || async { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let connections = state.connections.read().await; let pool = connections.get(&pool_key).ok_or("Pool not found")?; @@ -5600,7 +5896,21 @@ async fn get_table_ddl_core_with_options( return Ok(source.source); } - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + retry_metadata_connection(state, connection_id, Some(database), || { + get_table_ddl_once(state, connection_id, database, schema, table, include_postgres_access) + }) + .await +} + +async fn get_table_ddl_once( + state: &AppState, + connection_id: &str, + database: &str, + schema: &str, + table: &str, + include_postgres_access: bool, +) -> Result { + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; { @@ -6362,7 +6672,7 @@ async fn get_object_source_once( signature: Option<&str>, relation_name: Option<&str>, ) -> Result { - let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?; + let pool_key = state.get_or_create_metadata_pool_for_session(connection_id, Some(database), None).await?; let db_config = connection_config(state, connection_id).await; let source = { let connections = state.connections.read().await; @@ -6525,7 +6835,7 @@ pub fn oracle_list_objects_sql(schema: &str) -> String { } async fn oracle_agent_list_objects( - client: Arc>, + client: Arc, database: &str, schema: &str, timeout_duration: Option, @@ -6565,7 +6875,7 @@ async fn oracle_agent_list_objects( } async fn oracle_agent_object_source( - client: Arc>, + client: Arc, database: &str, schema: &str, name: &str, @@ -6585,7 +6895,7 @@ async fn oracle_agent_object_source( } async fn oracle_agent_table_ddl( - client: Arc>, + client: Arc, database: &str, schema: &str, table: &str, @@ -6666,7 +6976,7 @@ fn append_oracle_comments_to_ddl( } async fn db2_agent_table_ddl( - client: Arc>, + client: Arc, database: &str, schema: &str, table: &str, diff --git a/crates/dbx-core/src/schema/kingbase.rs b/crates/dbx-core/src/schema/kingbase.rs index f016207a5..935fac867 100644 --- a/crates/dbx-core/src/schema/kingbase.rs +++ b/crates/dbx-core/src/schema/kingbase.rs @@ -71,7 +71,7 @@ fn list_available_extensions_sql(catalog: ExtensionCatalog) -> String { } async fn query_result( - client: Arc>, + client: Arc, database: &str, sql: &str, max_rows: usize, @@ -88,7 +88,7 @@ async fn query_result( } async fn query_result_with_catalog_fallback( - client: Arc>, + client: Arc, database: &str, sys_sql: String, pg_sql: String, @@ -104,7 +104,7 @@ async fn query_result_with_catalog_fallback( } pub(super) async fn list_extensions( - client: Arc>, + client: Arc, database: &str, schema: Option<&str>, timeout_duration: Option, @@ -122,7 +122,7 @@ pub(super) async fn list_extensions( } pub(super) async fn list_available_extensions( - client: Arc>, + client: Arc, database: &str, timeout_duration: Option, ) -> Result, String> { diff --git a/src-tauri/src/commands/connection.rs b/src-tauri/src/commands/connection.rs index db67fa192..18b432ba3 100644 --- a/src-tauri/src/commands/connection.rs +++ b/src-tauri/src/commands/connection.rs @@ -162,7 +162,7 @@ async fn connect_agent_pool( } } - Ok(PoolKind::Agent(Arc::new(tokio::sync::Mutex::new(client)))) + Ok(PoolKind::agent(client)) } #[cfg(test)] @@ -1216,7 +1216,7 @@ pub async fn connect_db( .await .map_err(|err| mongo_legacy_error_with_auth_hint(&err))?; state.ensure_current_connection_attempt(&id, Some(attempt)).await?; - PoolKind::Agent(std::sync::Arc::new(tokio::sync::Mutex::new(client))) + PoolKind::agent(client) } else { let native_err = match db::mongo_driver::connect(&url, connect_timeout, idle_timeout).await { Ok(client) => { @@ -1266,7 +1266,7 @@ pub async fn connect_db( mark_mongo_legacy_driver(&mut connected_config); connected_db_config = metadata_connection_config(&connected_config); persist_mongo_legacy_driver_profile(state.inner(), &connected_config).await?; - PoolKind::Agent(std::sync::Arc::new(tokio::sync::Mutex::new(client))) + PoolKind::agent(client) } else { return Err(native_err); }