diff --git a/NEXT_CHANGELOG.md b/NEXT_CHANGELOG.md index e12c8875f4..1a2dfbc59d 100644 --- a/NEXT_CHANGELOG.md +++ b/NEXT_CHANGELOG.md @@ -3,6 +3,7 @@ ## [Unreleased] ### Added +- Added session-version exchange for SQL Exec API connections. ### Updated - `DatabaseMetaData.getColumns(...)` with a `null` catalog now issues a single `SHOW COLUMNS IN ALL CATALOGS` statement (consistent with `getSchemas`/`getTables`) instead of enumerating every catalog and issuing a per-catalog `SHOW COLUMNS`. Older DBR versions that do not support the syntax transparently fall back to the previous enumerate-and-fan-out behavior. diff --git a/src/main/java/com/databricks/jdbc/api/impl/DatabricksResultSet.java b/src/main/java/com/databricks/jdbc/api/impl/DatabricksResultSet.java index cde481ccf3..a48e6bea8d 100644 --- a/src/main/java/com/databricks/jdbc/api/impl/DatabricksResultSet.java +++ b/src/main/java/com/databricks/jdbc/api/impl/DatabricksResultSet.java @@ -429,6 +429,8 @@ private void startHeartbeatIfEnabled() { // connection is GC'd without close(), heartbeat RPCs will fail and self-stop after // maxConsecutiveFailures (10 ticks, ~10 min at 60s interval). Acceptable tradeoff. final IDatabricksClient client = conn.getSession().getDatabricksClient(); + final IDatabricksSession capturedSession = conn.getSession(); + final String originatingSessionId = parentStatement.getOriginatingSessionId(); final StatementId capturedStatementId = this.statementId; final int maxConsecutiveFailures = 10; final java.util.concurrent.atomic.AtomicInteger consecutiveFailures = @@ -449,7 +451,9 @@ private void startHeartbeatIfEnabled() { return; // client/session may be closed, skip RPC } try { - boolean alive = client.checkStatementAlive(capturedStatementId); + boolean alive = + client.checkStatementAlive( + capturedStatementId, capturedSession, originatingSessionId); consecutiveFailures.set(0); // reset on success if (!alive) { LOGGER.info( diff --git a/src/main/java/com/databricks/jdbc/api/impl/DatabricksSession.java b/src/main/java/com/databricks/jdbc/api/impl/DatabricksSession.java index bb33fa1189..c2b9b9fde8 100644 --- a/src/main/java/com/databricks/jdbc/api/impl/DatabricksSession.java +++ b/src/main/java/com/databricks/jdbc/api/impl/DatabricksSession.java @@ -21,6 +21,7 @@ import com.databricks.jdbc.exception.DatabricksTemporaryRedirectException; import com.databricks.jdbc.log.JdbcLogger; import com.databricks.jdbc.log.JdbcLoggerFactory; +import com.databricks.jdbc.model.core.SessionVersion; import com.databricks.jdbc.model.telemetry.enums.DatabricksDriverErrorCode; import com.databricks.jdbc.telemetry.TelemetryHelper; import com.databricks.jdbc.telemetry.latency.DatabricksMetricsTimedProcessor; @@ -29,6 +30,7 @@ import java.sql.SQLException; import java.util.HashMap; import java.util.Map; +import java.util.concurrent.atomic.AtomicReference; import javax.annotation.Nullable; /** @@ -43,6 +45,7 @@ public class DatabricksSession implements IDatabricksSession { private final IDatabricksComputeResource computeResource; private boolean isSessionOpen; private ImmutableSessionInfo sessionInfo; + private final AtomicReference sessionVersion = new AtomicReference<>(); /** For context based commands */ private String catalog; @@ -111,6 +114,37 @@ public ImmutableSessionInfo getSessionInfo() { return sessionInfo; } + @Nullable + @Override + public SessionVersion getSessionVersion() { + Long versionId = sessionVersion.get(); + return versionId == null ? null : new SessionVersion().setVersionId(versionId); + } + + @Override + public void updateSessionVersion( + @Nullable String expectedSessionId, @Nullable SessionVersion newSessionVersion) { + if (expectedSessionId == null + || newSessionVersion == null + || newSessionVersion.getVersionId() == null) { + return; + } + synchronized (this) { + if (!isSessionOpen + || sessionInfo == null + || !expectedSessionId.equals(sessionInfo.sessionId())) { + return; + } + Long newVersionId = newSessionVersion.getVersionId(); + sessionVersion.accumulateAndGet( + newVersionId, + (currentVersion, candidateVersion) -> + currentVersion == null || candidateVersion > currentVersion + ? candidateVersion + : currentVersion); + } + } + @Override public IDatabricksComputeResource getComputeResource() { LOGGER.debug("public String getComputeResource()"); @@ -217,6 +251,7 @@ public void open() throws SQLException { throw e; } } + this.sessionVersion.set(sessionInfo == null ? null : sessionInfo.sessionVersion()); this.isSessionOpen = true; } } @@ -240,6 +275,7 @@ public void close() throws SQLException { } finally { // Always clean up local state this.sessionInfo = null; + this.sessionVersion.set(null); this.isSessionOpen = false; } } @@ -406,6 +442,7 @@ public void forceClose() { } catch (SQLException e) { LOGGER.error("Error closing session resources, but marking the session as closed."); } finally { + this.sessionVersion.set(null); this.isSessionOpen = false; } } diff --git a/src/main/java/com/databricks/jdbc/api/impl/DatabricksStatement.java b/src/main/java/com/databricks/jdbc/api/impl/DatabricksStatement.java index d1dd0d30e1..52ff897309 100644 --- a/src/main/java/com/databricks/jdbc/api/impl/DatabricksStatement.java +++ b/src/main/java/com/databricks/jdbc/api/impl/DatabricksStatement.java @@ -25,6 +25,7 @@ import java.util.List; import java.util.Map; import java.util.concurrent.*; +import javax.annotation.Nullable; import org.apache.http.entity.InputStreamEntity; public class DatabricksStatement implements IDatabricksStatement, IDatabricksStatementInternal { @@ -44,6 +45,7 @@ public class DatabricksStatement implements IDatabricksStatement, IDatabricksSta protected final DatabricksConnection connection; DatabricksResultSet resultSet; private volatile StatementId statementId; // volatile: cancel() reads from a different thread + private volatile String originatingSessionId; private boolean isClosed; private boolean closeOnCompletion; private SQLWarning warnings = null; @@ -69,6 +71,7 @@ public DatabricksStatement(DatabricksConnection connection) throws DatabricksVal this.connection = connection; this.resultSet = null; this.statementId = null; + this.originatingSessionId = null; this.isClosed = false; this.timeoutInSeconds = DEFAULT_STATEMENT_TIMEOUT_SECONDS; this.databricksBatchExecutor = @@ -79,6 +82,7 @@ public DatabricksStatement(DatabricksConnection connection, StatementId statemen throws DatabricksValidationException { this.connection = connection; this.statementId = statementId; + this.originatingSessionId = null; this.resultSet = null; this.isClosed = false; this.timeoutInSeconds = DEFAULT_STATEMENT_TIMEOUT_SECONDS; @@ -644,8 +648,14 @@ public void handleResultSetClose(IDatabricksResultSet resultSet) throws Databric @Override public void setStatementId(StatementId statementId) { + setStatementId(statementId, null); + } + + @Override + public void setStatementId(StatementId statementId, @Nullable String originatingSessionId) { LOGGER.debug("void setStatementId(Statement statementId = {})", statementId); this.statementId = statementId; + this.originatingSessionId = originatingSessionId; } @Override @@ -658,6 +668,12 @@ public Statement getStatement() { return this; } + @Override + @Nullable + public String getOriginatingSessionId() { + return originatingSessionId; + } + @Override public void allowInputStreamForVolumeOperation(boolean allowInputStream) throws DatabricksSQLException { @@ -1097,6 +1113,7 @@ private void resetForNewExecution() { // Null out statementId so that if the new execution fails before setStatementId(), // close() takes the statementId==null branch instead of sending closeStatement(stale-id) statementId = null; + originatingSessionId = null; } /** diff --git a/src/main/java/com/databricks/jdbc/api/impl/SessionInfo.java b/src/main/java/com/databricks/jdbc/api/impl/SessionInfo.java index b09909b6ed..799ab83e57 100644 --- a/src/main/java/com/databricks/jdbc/api/impl/SessionInfo.java +++ b/src/main/java/com/databricks/jdbc/api/impl/SessionInfo.java @@ -12,6 +12,9 @@ public interface SessionInfo { IDatabricksComputeResource computeResource(); + @Nullable + Long sessionVersion(); + @Nullable TSessionHandle sessionHandle(); // This field is set only for all-purpose cluster compute } diff --git a/src/main/java/com/databricks/jdbc/api/internal/IDatabricksSession.java b/src/main/java/com/databricks/jdbc/api/internal/IDatabricksSession.java index 0229444674..405568f7f1 100644 --- a/src/main/java/com/databricks/jdbc/api/internal/IDatabricksSession.java +++ b/src/main/java/com/databricks/jdbc/api/internal/IDatabricksSession.java @@ -6,6 +6,7 @@ import com.databricks.jdbc.dbclient.IDatabricksClient; import com.databricks.jdbc.dbclient.IDatabricksMetadataClient; import com.databricks.jdbc.exception.DatabricksSQLException; +import com.databricks.jdbc.model.core.SessionVersion; import java.sql.SQLException; import java.util.Map; import javax.annotation.Nullable; @@ -24,6 +25,14 @@ public interface IDatabricksSession { @Nullable ImmutableSessionInfo getSessionInfo(); + @Nullable + default SessionVersion getSessionVersion() { + return null; + } + + default void updateSessionVersion( + @Nullable String expectedSessionId, @Nullable SessionVersion sessionVersion) {} + /** * Get the warehouse associated with the session. * diff --git a/src/main/java/com/databricks/jdbc/api/internal/IDatabricksStatementInternal.java b/src/main/java/com/databricks/jdbc/api/internal/IDatabricksStatementInternal.java index 62b893c6b9..eb9a0be45a 100644 --- a/src/main/java/com/databricks/jdbc/api/internal/IDatabricksStatementInternal.java +++ b/src/main/java/com/databricks/jdbc/api/internal/IDatabricksStatementInternal.java @@ -4,6 +4,7 @@ import com.databricks.jdbc.dbclient.impl.common.StatementId; import com.databricks.jdbc.exception.DatabricksSQLException; import java.sql.Statement; +import javax.annotation.Nullable; import org.apache.http.entity.InputStreamEntity; /** Extended callback handle for java.sql.Statement interface */ @@ -19,10 +20,19 @@ public interface IDatabricksStatementInternal { void setStatementId(StatementId statementId); + default void setStatementId(StatementId statementId, @Nullable String originatingSessionId) { + setStatementId(statementId); + } + StatementId getStatementId(); Statement getStatement(); + @Nullable + default String getOriginatingSessionId() { + return null; + } + void allowInputStreamForVolumeOperation(boolean allowedInputStream) throws DatabricksSQLException; boolean isAllowedInputStreamForVolumeOperation() throws DatabricksSQLException; diff --git a/src/main/java/com/databricks/jdbc/dbclient/IDatabricksClient.java b/src/main/java/com/databricks/jdbc/dbclient/IDatabricksClient.java index 71e7900743..c234006949 100644 --- a/src/main/java/com/databricks/jdbc/dbclient/IDatabricksClient.java +++ b/src/main/java/com/databricks/jdbc/dbclient/IDatabricksClient.java @@ -16,6 +16,7 @@ import com.databricks.sdk.core.DatabricksConfig; import java.sql.SQLException; import java.util.Map; +import javax.annotation.Nullable; /** Interface for Databricks client which abstracts the integration with Databricks server. */ public interface IDatabricksClient { @@ -120,6 +121,12 @@ default boolean checkStatementAlive(StatementId statementId) throws SQLException throw new java.sql.SQLFeatureNotSupportedException("Heartbeat not supported by this client"); } + default boolean checkStatementAlive( + StatementId statementId, IDatabricksSession session, @Nullable String originatingSessionId) + throws SQLException { + return checkStatementAlive(statementId); + } + /** * Fetches result for underlying statement-Id * diff --git a/src/main/java/com/databricks/jdbc/dbclient/impl/sqlexec/DatabricksSdkClient.java b/src/main/java/com/databricks/jdbc/dbclient/impl/sqlexec/DatabricksSdkClient.java index a62a4c6318..6b565a19af 100644 --- a/src/main/java/com/databricks/jdbc/dbclient/impl/sqlexec/DatabricksSdkClient.java +++ b/src/main/java/com/databricks/jdbc/dbclient/impl/sqlexec/DatabricksSdkClient.java @@ -37,6 +37,8 @@ import com.databricks.jdbc.model.core.ExternalLink; import com.databricks.jdbc.model.core.ResultData; import com.databricks.jdbc.model.core.ResultManifest; +import com.databricks.jdbc.model.core.SessionExecutionMode; +import com.databricks.jdbc.model.core.SessionVersion; import com.databricks.jdbc.model.core.StatementStatus; import com.databricks.jdbc.model.telemetry.enums.DatabricksDriverErrorCode; import com.databricks.sdk.WorkspaceClient; @@ -55,6 +57,7 @@ import java.util.*; import java.util.stream.Collectors; import java.util.stream.Stream; +import javax.annotation.Nullable; import javax.net.ssl.SSLHandshakeException; import org.apache.http.HttpStatus; @@ -116,7 +119,9 @@ public ImmutableSessionInfo createSession( schema, sessionConf); CreateSessionRequest request = - new CreateSessionRequest().setWarehouseId(((Warehouse) warehouse).getWarehouseId()); + new CreateSessionRequest() + .setWarehouseId(((Warehouse) warehouse).getWarehouseId()) + .setExecutionMode(SessionExecutionMode.FAST); if (catalog != null) { request.setCatalog(catalog); } @@ -153,11 +158,20 @@ public ImmutableSessionInfo createSession( LOGGER.error(errorMessage, e); throw new DatabricksSQLException(errorMessage, e, DatabricksDriverErrorCode.SDK_CLIENT_ERROR); } - DatabricksThreadContextHolder.setSessionId(createSessionResponse.getSessionId()); - return ImmutableSessionInfo.builder() - .computeResource(warehouse) - .sessionId(createSessionResponse.getSessionId()) - .build(); + String sessionId = createSessionResponse == null ? null : createSessionResponse.getSessionId(); + if (sessionId == null || sessionId.isEmpty()) { + throw new DatabricksSQLException( + "Create session response did not include session_id", + DatabricksDriverErrorCode.CONNECTION_ERROR); + } + DatabricksThreadContextHolder.setSessionId(sessionId); + ImmutableSessionInfo.Builder sessionInfo = + ImmutableSessionInfo.builder().computeResource(warehouse).sessionId(sessionId); + SessionVersion initialVersion = createSessionResponse.getSessionVersion(); + if (initialVersion != null && initialVersion.getVersionId() != null) { + sessionInfo.sessionVersion(initialVersion.getVersionId()); + } + return sessionInfo.build(); } @Override @@ -200,7 +214,8 @@ public DatabricksResultSet executeStatement( session, parentStatement, metadataOperationType); - DatabricksThreadContextHolder.setSessionId(session.getSessionId()); + String requestSessionId = session.getSessionId(); + DatabricksThreadContextHolder.setSessionId(requestSessionId); long pollCount = 0; long executionStartTime = Instant.now().toEpochMilli(); DatabricksThreadContextHolder.setStatementType(statementType); @@ -210,6 +225,7 @@ public DatabricksResultSet executeStatement( sql, ((Warehouse) computeResource).getWarehouseId(), session, + requestSessionId, parameters, parentStatement, false); @@ -223,6 +239,7 @@ public DatabricksResultSet executeStatement( } req.withHeaders(getHeaders("executeStatement", statementType, false, additionalHeaders)); response = apiClient.execute(req, ExecuteStatementResponse.class); + updateSessionVersion(session, requestSessionId, response.getStatus()); } catch (IOException e) { String errorMessage = "Error while processing the execute statement request"; LOGGER.error(errorMessage, e); @@ -246,7 +263,7 @@ public DatabricksResultSet executeStatement( StatementId typedStatementId = new StatementId(statementId); DatabricksThreadContextHolder.setStatementId(typedStatementId); if (parentStatement != null) { - parentStatement.setStatementId(typedStatementId); + parentStatement.setStatementId(typedStatementId, requestSessionId); } int timeoutInSeconds; @@ -268,6 +285,7 @@ public DatabricksResultSet executeStatement( TimeoutHandler.forStatement(timeoutInSeconds, typedStatementId, this, timeoutErrorCode); StatementState responseState = response.getStatus().getState(); + GetStatementRequest getStatementRequest = new GetStatementRequest().setStatementId(statementId); while (responseState == StatementState.PENDING || responseState == StatementState.RUNNING) { // Check for timeout timeoutHandler.checkTimeout(); @@ -286,9 +304,11 @@ public DatabricksResultSet executeStatement( } String getStatusPath = String.format(STATEMENT_PATH_WITH_ID, statementId); try { - Request req = new Request(Request.GET, getStatusPath, apiClient.serialize(request)); + Request req = + new Request(Request.GET, getStatusPath, apiClient.serialize(getStatementRequest)); req.withHeaders(getHeaders("getStatement")); response = wrapGetStatementResponse(apiClient.execute(req, GetStatementResponse.class)); + updateSessionVersion(session, requestSessionId, response.getStatus()); } catch (IOException e) { String errorMessage = "Error while processing the get statement response"; LOGGER.error(errorMessage, e); @@ -369,7 +389,8 @@ public DatabricksResultSet executeStatementAsync( computeResource.toString(), session, parentStatement); - DatabricksThreadContextHolder.setSessionId(session.getSessionId()); + String requestSessionId = session.getSessionId(); + DatabricksThreadContextHolder.setSessionId(requestSessionId); StatementType statementType = StatementType.SQL; ExecuteStatementRequest request = getRequest( @@ -377,6 +398,7 @@ public DatabricksResultSet executeStatementAsync( sql, ((Warehouse) computeResource).getWarehouseId(), session, + requestSessionId, parameters, parentStatement, true); @@ -385,6 +407,7 @@ public DatabricksResultSet executeStatementAsync( Request req = new Request(Request.POST, STATEMENT_PATH, apiClient.serialize(request)); req.withHeaders(getHeaders("executeStatement", statementType, true)); response = apiClient.execute(req, ExecuteStatementResponse.class); + updateSessionVersion(session, requestSessionId, response.getStatus()); } catch (IOException e) { String errorMessage = "Error while processing the execute statement async request"; LOGGER.error(errorMessage, e); @@ -398,7 +421,7 @@ public DatabricksResultSet executeStatementAsync( StatementId typedStatementId = new StatementId(statementId); DatabricksThreadContextHolder.setStatementId(typedStatementId); if (parentStatement != null) { - parentStatement.setStatementId(typedStatementId); + parentStatement.setStatementId(typedStatementId, requestSessionId); } LOGGER.debug("Executed sql [{}] with status [{}]", sql, response.getStatus().getState()); @@ -417,6 +440,15 @@ public DatabricksResultSet executeStatementAsync( @Override public boolean checkStatementAlive(StatementId typedStatementId) throws SQLException { + return checkStatementAlive(typedStatementId, null, null); + } + + @Override + public boolean checkStatementAlive( + StatementId typedStatementId, + @Nullable IDatabricksSession session, + @Nullable String originatingSessionId) + throws SQLException { String statementId = typedStatementId.toSQLExecStatementId(); // Use lightweight /status endpoint (~100 bytes) instead of full GetStatement (~21KB) String statusPath = String.format(STATEMENT_STATUS_PATH_WITH_ID, statementId); @@ -424,6 +456,7 @@ public boolean checkStatementAlive(StatementId typedStatementId) throws SQLExcep Request req = new Request(Request.GET, statusPath, (String) null); req.withHeaders(getHeaders("getStatementStatus")); StatementStatus status = apiClient.execute(req, StatementStatus.class); + updateSessionVersion(session, originatingSessionId, status); StatementState state = status.getState(); // Terminal states mean the operation is no longer alive return state != StatementState.CANCELED @@ -445,7 +478,8 @@ public DatabricksResultSet getStatementResult( IDatabricksStatementInternal parentStatement) throws SQLException { DatabricksThreadContextHolder.setStatementId(typedStatementId); - DatabricksThreadContextHolder.setSessionId(session.getSessionId()); + String requestSessionId = session.getSessionId(); + DatabricksThreadContextHolder.setSessionId(requestSessionId); String statementId = typedStatementId.toSQLExecStatementId(); GetStatementRequest request = new GetStatementRequest().setStatementId(statementId); String getStatusPath = String.format(STATEMENT_PATH_WITH_ID, statementId); @@ -454,6 +488,9 @@ public DatabricksResultSet getStatementResult( Request req = new Request(Request.GET, getStatusPath, apiClient.serialize(request)); req.withHeaders(getHeaders("getStatement")); response = apiClient.execute(req, GetStatementResponse.class); + String originatingSessionId = + parentStatement == null ? null : parentStatement.getOriginatingSessionId(); + updateSessionVersion(session, originatingSessionId, response.getStatus()); } catch (IOException e) { String errorMessage = "Error while processing the get statement result request"; LOGGER.error(errorMessage, e); @@ -712,6 +749,7 @@ private ExecuteStatementRequest getRequest( String sql, String warehouseId, IDatabricksSession session, + String requestSessionId, Map parameters, IDatabricksStatementInternal parentStatement, boolean executeAsync) @@ -734,13 +772,17 @@ private ExecuteStatementRequest getRequest( parameters.values().stream().map(this::mapToParameterListItem).collect(Collectors.toList()); ExecuteStatementRequest request = new ExecuteStatementRequest() - .setSessionId(session.getSessionId()) + .setSessionId(requestSessionId) .setStatement(sql) .setWarehouseId(warehouseId) .setDisposition(disposition) .setFormat(format) .setResultCompression(compressionCodec) .setParameters(parameterListItems); + SessionVersion sessionVersion = session.getSessionVersion(); + if (sessionVersion != null) { + request.setSessionVersion(sessionVersion); + } if (executeAsync) { request.setWaitTimeout(ASYNC_TIMEOUT_VALUE); } else { @@ -823,6 +865,13 @@ private ExecuteStatementResponse wrapGetStatementResponse( .setResult(getStatementResponse.getResult()); } + private void updateSessionVersion( + IDatabricksSession session, String requestSessionId, StatementStatus status) { + if (session != null && status != null) { + session.updateSessionVersion(requestSessionId, status.getSessionVersion()); + } + } + /** * Builds actionable error messages for SSL handshake failures. Returns a generic message if the * error is not SSL-related. diff --git a/src/main/java/com/databricks/jdbc/model/client/sqlexec/CreateSessionRequest.java b/src/main/java/com/databricks/jdbc/model/client/sqlexec/CreateSessionRequest.java index 0897ede538..dbf1888d51 100644 --- a/src/main/java/com/databricks/jdbc/model/client/sqlexec/CreateSessionRequest.java +++ b/src/main/java/com/databricks/jdbc/model/client/sqlexec/CreateSessionRequest.java @@ -1,5 +1,6 @@ package com.databricks.jdbc.model.client.sqlexec; +import com.databricks.jdbc.model.core.SessionExecutionMode; import com.fasterxml.jackson.annotation.JsonProperty; import java.util.Map; @@ -19,6 +20,9 @@ public class CreateSessionRequest { @JsonProperty("session_confs") private Map sessionConfigs; + @JsonProperty("execution_mode") + private SessionExecutionMode executionMode; + public CreateSessionRequest setWarehouseId(String warehouseId) { this.warehouseId = warehouseId; return this; @@ -54,4 +58,13 @@ public CreateSessionRequest setSessionConfigs(Map sessionConfigs public Map getSessionConfigs() { return sessionConfigs; } + + public CreateSessionRequest setExecutionMode(SessionExecutionMode executionMode) { + this.executionMode = executionMode; + return this; + } + + public SessionExecutionMode getExecutionMode() { + return executionMode; + } } diff --git a/src/main/java/com/databricks/jdbc/model/client/sqlexec/CreateSessionResponse.java b/src/main/java/com/databricks/jdbc/model/client/sqlexec/CreateSessionResponse.java index 7ef7ef0ff4..724d25e009 100644 --- a/src/main/java/com/databricks/jdbc/model/client/sqlexec/CreateSessionResponse.java +++ b/src/main/java/com/databricks/jdbc/model/client/sqlexec/CreateSessionResponse.java @@ -1,5 +1,6 @@ package com.databricks.jdbc.model.client.sqlexec; +import com.databricks.jdbc.model.core.SessionVersion; import com.fasterxml.jackson.annotation.JsonProperty; /** @@ -13,6 +14,9 @@ public class CreateSessionResponse { @JsonProperty("session_id") private String sessionId; + @JsonProperty("session_version") + private SessionVersion sessionVersion; + public CreateSessionResponse setSessionId(String sessionId) { this.sessionId = sessionId; return this; @@ -21,4 +25,13 @@ public CreateSessionResponse setSessionId(String sessionId) { public String getSessionId() { return sessionId; } + + public CreateSessionResponse setSessionVersion(SessionVersion sessionVersion) { + this.sessionVersion = sessionVersion; + return this; + } + + public SessionVersion getSessionVersion() { + return sessionVersion; + } } diff --git a/src/main/java/com/databricks/jdbc/model/client/sqlexec/ExecuteStatementRequest.java b/src/main/java/com/databricks/jdbc/model/client/sqlexec/ExecuteStatementRequest.java index 24968d2cf4..a1c73831c4 100644 --- a/src/main/java/com/databricks/jdbc/model/client/sqlexec/ExecuteStatementRequest.java +++ b/src/main/java/com/databricks/jdbc/model/client/sqlexec/ExecuteStatementRequest.java @@ -2,6 +2,7 @@ import com.databricks.jdbc.common.CompressionCodec; import com.databricks.jdbc.model.core.Disposition; +import com.databricks.jdbc.model.core.SessionVersion; import com.databricks.sdk.service.sql.ExecuteStatementRequestOnWaitTimeout; import com.databricks.sdk.service.sql.Format; import com.databricks.sdk.service.sql.StatementParameterListItem; @@ -46,6 +47,9 @@ public class ExecuteStatementRequest { @JsonProperty("result_compression") private CompressionCodec resultCompression; + @JsonProperty("session_version") + private SessionVersion sessionVersion; + public String getStatement() { return statement; } @@ -86,6 +90,10 @@ public CompressionCodec getResultCompression() { return resultCompression; } + public SessionVersion getSessionVersion() { + return sessionVersion; + } + // Setters public ExecuteStatementRequest setStatement(String statement) { this.statement = statement; @@ -138,6 +146,11 @@ public ExecuteStatementRequest setParameters(Collection + session.updateSessionVersion( + SESSION_ID, new SessionVersion().setVersionId(version))); + session.updateSessionVersion(SESSION_ID, new SessionVersion().setVersionId(500L)); + session.updateSessionVersion(SESSION_ID, new SessionVersion()); + session.updateSessionVersion(SESSION_ID, null); + + assertEquals(1000L, session.getSessionVersion().getVersionId()); + session.close(); + assertNull(session.getSessionVersion()); + } + + @Test + public void testSessionVersionRejectsLateUpdatesAfterCloseAndReopen() throws SQLException { + setupWarehouse(false /* useThrift */); + String replacementSessionId = "replacement_session_id"; + when(sdkClient.createSession(any(), any(), any(), any())) + .thenReturn( + ImmutableSessionInfo.builder() + .sessionId(SESSION_ID) + .sessionVersion(10L) + .computeResource(WAREHOUSE_COMPUTE) + .build()) + .thenReturn( + ImmutableSessionInfo.builder() + .sessionId(replacementSessionId) + .sessionVersion(20L) + .computeResource(WAREHOUSE_COMPUTE) + .build()); + DatabricksSession session = new DatabricksSession(connectionContext, sdkClient); + session.open(); + session.close(); + + session.updateSessionVersion(SESSION_ID, new SessionVersion().setVersionId(99L)); + assertNull(session.getSessionVersion()); + + session.open(); + session.updateSessionVersion(SESSION_ID, new SessionVersion().setVersionId(99L)); + assertEquals(20L, session.getSessionVersion().getVersionId()); + + session.updateSessionVersion(replacementSessionId, new SessionVersion().setVersionId(21L)); + assertEquals(21L, session.getSessionVersion().getVersionId()); + } + @Test public void testOpenRedirectedThriftSession() throws SQLException { setupWarehouse(false /* useThrift */); diff --git a/src/test/java/com/databricks/jdbc/api/impl/DatabricksStatementTest.java b/src/test/java/com/databricks/jdbc/api/impl/DatabricksStatementTest.java index e1dc979841..bfae63ef10 100644 --- a/src/test/java/com/databricks/jdbc/api/impl/DatabricksStatementTest.java +++ b/src/test/java/com/databricks/jdbc/api/impl/DatabricksStatementTest.java @@ -369,6 +369,22 @@ public void testGetStatementId() throws DatabricksSQLException { when(mockConnection.getConnectionContext()).thenReturn(connectionContext); DatabricksStatement statement = new DatabricksStatement(mockConnection, STATEMENT_ID); assertEquals(STATEMENT_ID, statement.getStatementId()); + assertNull(statement.getOriginatingSessionId()); + } + + @Test + public void testStatementRecordsOriginatingSessionId() throws DatabricksSQLException { + DatabricksConnection mockConnection = mock(DatabricksConnection.class); + IDatabricksConnectionContext connectionContext = + DatabricksConnectionContextFactory.create(JDBC_URL, new Properties()); + when(mockConnection.getConnectionContext()).thenReturn(connectionContext); + DatabricksStatement statement = new DatabricksStatement(mockConnection); + + statement.setStatementId(STATEMENT_ID, SESSION_ID); + assertEquals(SESSION_ID, statement.getOriginatingSessionId()); + + statement.setStatementId(STATEMENT_ID); + assertNull(statement.getOriginatingSessionId()); } @Test diff --git a/src/test/java/com/databricks/jdbc/api/impl/ResultSetHeartbeatEligibilityTest.java b/src/test/java/com/databricks/jdbc/api/impl/ResultSetHeartbeatEligibilityTest.java index 4b9425f13e..0c7b12ef1e 100644 --- a/src/test/java/com/databricks/jdbc/api/impl/ResultSetHeartbeatEligibilityTest.java +++ b/src/test/java/com/databricks/jdbc/api/impl/ResultSetHeartbeatEligibilityTest.java @@ -3,10 +3,15 @@ import static org.junit.jupiter.api.Assertions.*; import static org.mockito.Mockito.*; +import com.databricks.jdbc.api.internal.IDatabricksSession; +import com.databricks.jdbc.api.internal.IDatabricksStatementInternal; import com.databricks.jdbc.common.StatementType; +import com.databricks.jdbc.dbclient.IDatabricksClient; import com.databricks.jdbc.dbclient.impl.common.StatementId; import com.databricks.jdbc.model.core.StatementStatus; import com.databricks.sdk.service.sql.StatementState; +import java.sql.Statement; +import java.util.concurrent.atomic.AtomicBoolean; import org.junit.jupiter.api.Test; /** @@ -173,4 +178,54 @@ void testAsyncRunningNotEligible() { assertFalse( rs.isHeartbeatEligible(), "Async RUNNING — user controls polling via getExecutionResult"); } + + @Test + void testHeartbeatForwardsStatementOwnership() throws Exception { + assertHeartbeatForwardsStatementOwnership("session-id"); + assertHeartbeatForwardsStatementOwnership(null); + } + + private void assertHeartbeatForwardsStatementOwnership(String originatingSessionId) + throws Exception { + StatementId statementId = new StatementId("test-stmt"); + IDatabricksStatementInternal parentStatement = mock(IDatabricksStatementInternal.class); + Statement jdbcStatement = mock(Statement.class); + DatabricksConnection connection = mock(DatabricksConnection.class); + ResultHeartbeatManager heartbeatManager = mock(ResultHeartbeatManager.class); + IDatabricksSession session = mock(IDatabricksSession.class); + IDatabricksClient client = mock(IDatabricksClient.class); + + when(parentStatement.getStatement()).thenReturn(jdbcStatement); + when(parentStatement.getOriginatingSessionId()).thenReturn(originatingSessionId); + when(jdbcStatement.getConnection()).thenReturn(connection); + when(connection.getHeartbeatManager()).thenReturn(heartbeatManager); + when(connection.getSession()).thenReturn(session); + when(session.getDatabricksClient()).thenReturn(client); + when(heartbeatManager.getStoppedFlag(statementId)).thenReturn(new AtomicBoolean(false)); + when(client.checkStatementAlive(statementId, session, originatingSessionId)).thenReturn(true); + doAnswer( + invocation -> { + invocation.getArgument(1).run(); + return null; + }) + .when(heartbeatManager) + .startHeartbeat(eq(statementId), any(Runnable.class)); + + DatabricksResultSet resultSet = + new DatabricksResultSet( + new StatementStatus().setState(StatementState.SUCCEEDED), + statementId, + StatementType.QUERY, + parentStatement, + mock(IExecutionResult.class), + null, + false); + setField(resultSet, "resultSetType", DatabricksResultSet.ResultSetType.THRIFT_INLINE); + java.lang.reflect.Method startHeartbeat = + DatabricksResultSet.class.getDeclaredMethod("startHeartbeatIfEnabled"); + startHeartbeat.setAccessible(true); + startHeartbeat.invoke(resultSet); + + verify(client).checkStatementAlive(statementId, session, originatingSessionId); + } } diff --git a/src/test/java/com/databricks/jdbc/dbclient/impl/sqlexec/DatabricksSdkClientTest.java b/src/test/java/com/databricks/jdbc/dbclient/impl/sqlexec/DatabricksSdkClientTest.java index 4a6e266295..8f71a0c126 100644 --- a/src/test/java/com/databricks/jdbc/dbclient/impl/sqlexec/DatabricksSdkClientTest.java +++ b/src/test/java/com/databricks/jdbc/dbclient/impl/sqlexec/DatabricksSdkClientTest.java @@ -32,6 +32,8 @@ import com.databricks.jdbc.model.core.ResultData; import com.databricks.jdbc.model.core.ResultManifest; import com.databricks.jdbc.model.core.ResultSchema; +import com.databricks.jdbc.model.core.SessionExecutionMode; +import com.databricks.jdbc.model.core.SessionVersion; import com.databricks.jdbc.model.core.StatementStatus; import com.databricks.jdbc.model.telemetry.enums.DatabricksDriverErrorCode; import com.databricks.sdk.core.ApiClient; @@ -59,6 +61,8 @@ public class DatabricksSdkClientTest { // Reference to MetadataOperationType to ensure import is not removed private static final MetadataOperationType SAMPLE_OP_TYPE = MetadataOperationType.GET_CATALOGS; private static final String SESSION_ID = "session_id"; + private static final long INITIAL_SESSION_VERSION = 10L; + private static final long UPDATED_SESSION_VERSION = 12L; private static final StatementId STATEMENT_ID = new StatementId("statementId"); private static final String STATEMENT = "SELECT * FROM orders WHERE user_id = ? AND shard = ? AND region_code = ? AND namespace = ?"; @@ -76,13 +80,29 @@ public class DatabricksSdkClientTest { } }; + private static SessionVersion sessionVersion(long versionId) { + return new SessionVersion().setVersionId(versionId); + } + + private static CreateSessionResponse createSessionResponse() { + return new CreateSessionResponse() + .setSessionId(SESSION_ID) + .setSessionVersion(sessionVersion(INITIAL_SESSION_VERSION)); + } + private void setupSessionMocks() throws IOException { - CreateSessionResponse response = new CreateSessionResponse().setSessionId(SESSION_ID); + CreateSessionResponse response = createSessionResponse(); when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) .thenReturn(response); } private void setupClientMocks(boolean includeResults, boolean async) throws IOException { + setupClientMocks(includeResults, async, true); + } + + private void setupClientMocks( + boolean includeResults, boolean async, boolean includeInitialSessionVersion) + throws IOException { List params = new ArrayList<>() { { @@ -93,7 +113,10 @@ private void setupClientMocks(boolean includeResults, boolean async) throws IOEx } }; - StatementStatus statementStatus = new StatementStatus().setState(StatementState.SUCCEEDED); + StatementStatus statementStatus = + new StatementStatus() + .setState(StatementState.SUCCEEDED) + .setSessionVersion(sessionVersion(UPDATED_SESSION_VERSION)); ExecuteStatementRequest executeStatementRequest = new ExecuteStatementRequest() .setSessionId(SESSION_ID) @@ -102,6 +125,7 @@ private void setupClientMocks(boolean includeResults, boolean async) throws IOEx .setDisposition(Disposition.INLINE_OR_EXTERNAL_LINKS) .setFormat(Format.ARROW_STREAM) .setRowLimit(100L) + .setSessionVersion(sessionVersion(INITIAL_SESSION_VERSION)) .setParameters(params); if (async) { executeStatementRequest.setWaitTimeout("0s"); @@ -131,7 +155,9 @@ private void setupClientMocks(boolean includeResults, boolean async) throws IOEx if (req.getUrl().equals(STATEMENT_PATH)) { return response; } else if (req.getUrl().equals(SESSION_PATH)) { - return new CreateSessionResponse().setSessionId(SESSION_ID); + return includeInitialSessionVersion + ? createSessionResponse() + : new CreateSessionResponse().setSessionId(SESSION_ID); } return null; }); @@ -148,6 +174,50 @@ public void testCreateSession() throws DatabricksSQLException, IOException { databricksSdkClient.createSession(warehouse, null, null, null); assertEquals(sessionInfo.sessionId(), SESSION_ID); assertEquals(sessionInfo.computeResource(), warehouse); + assertEquals(INITIAL_SESSION_VERSION, sessionInfo.sessionVersion()); + verify(apiClient) + .serialize( + argThat( + request -> + request instanceof CreateSessionRequest + && ((CreateSessionRequest) request).getExecutionMode() + == SessionExecutionMode.FAST)); + } + + @Test + public void testCreateSessionWithoutInitialSessionVersion() throws Exception { + when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) + .thenReturn(new CreateSessionResponse().setSessionId(SESSION_ID)); + IDatabricksConnectionContext connectionContext = + DatabricksConnectionContext.parse(JDBC_URL, new Properties()); + DatabricksSdkClient databricksSdkClient = + new DatabricksSdkClient(connectionContext, statementExecutionService, apiClient); + + ImmutableSessionInfo sessionInfo = + databricksSdkClient.createSession(warehouse, null, null, null); + + assertEquals(SESSION_ID, sessionInfo.sessionId()); + assertNull(sessionInfo.sessionVersion()); + verify(apiClient, never()).execute(any(Request.class), eq(Void.class)); + } + + @Test + public void testCreateSessionFailsWithoutSessionId() throws Exception { + when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) + .thenReturn( + new CreateSessionResponse().setSessionVersion(sessionVersion(INITIAL_SESSION_VERSION))); + IDatabricksConnectionContext connectionContext = + DatabricksConnectionContext.parse(JDBC_URL, new Properties()); + DatabricksSdkClient databricksSdkClient = + new DatabricksSdkClient(connectionContext, statementExecutionService, apiClient); + + DatabricksSQLException exception = + assertThrows( + DatabricksSQLException.class, + () -> databricksSdkClient.createSession(warehouse, null, null, null)); + + assertTrue(exception.getMessage().contains("session_id")); + verify(apiClient, never()).execute(any(Request.class), eq(Void.class)); } @Test @@ -216,10 +286,19 @@ public void testExecuteStatement() throws Exception { statement, null); assertEquals(STATEMENT_ID, statement.getStatementId()); + assertEquals(SESSION_ID, statement.getOriginatingSessionId()); assertNotNull(resultSet.getMetaData()); + assertEquals( + UPDATED_SESSION_VERSION, connection.getSession().getSessionVersion().getVersionId()); // Verify a Request with POST method is created and executed - verify(apiClient, atLeastOnce()).serialize(any(ExecuteStatementRequest.class)); + verify(apiClient, atLeastOnce()) + .serialize( + argThat( + request -> + request instanceof ExecuteStatementRequest + && sessionVersion(INITIAL_SESSION_VERSION) + .equals(((ExecuteStatementRequest) request).getSessionVersion()))); verify(apiClient, atLeastOnce()) .execute( argThat( @@ -227,6 +306,56 @@ public void testExecuteStatement() throws Exception { eq(ExecuteStatementResponse.class)); } + @Test + public void testExecuteStatementWithoutInitialSessionVersion() throws Exception { + setupClientMocks(true, false, false); + IDatabricksConnectionContext connectionContext = + DatabricksConnectionContext.parse(JDBC_URL, new Properties()); + DatabricksSdkClient databricksSdkClient = + new DatabricksSdkClient(connectionContext, statementExecutionService, apiClient); + DatabricksConnection connection = + new DatabricksConnection(connectionContext, databricksSdkClient); + connection.open(); + DatabricksStatement statement = new DatabricksStatement(connection); + + databricksSdkClient.executeStatement( + STATEMENT, + warehouse, + sqlParams, + StatementType.QUERY, + connection.getSession(), + statement, + null); + + verify(apiClient, atLeastOnce()) + .serialize( + argThat( + request -> + request instanceof ExecuteStatementRequest + && ((ExecuteStatementRequest) request).getSessionVersion() == null)); + assertEquals( + UPDATED_SESSION_VERSION, connection.getSession().getSessionVersion().getVersionId()); + + clearInvocations(apiClient); + DatabricksStatement nextStatement = new DatabricksStatement(connection); + databricksSdkClient.executeStatement( + STATEMENT, + warehouse, + sqlParams, + StatementType.QUERY, + connection.getSession(), + nextStatement, + null); + + verify(apiClient, atLeastOnce()) + .serialize( + argThat( + request -> + request instanceof ExecuteStatementRequest + && sessionVersion(UPDATED_SESSION_VERSION) + .equals(((ExecuteStatementRequest) request).getSessionVersion()))); + } + @Test public void testExecuteStatementAsync() throws Exception { setupClientMocks(false, true); @@ -244,10 +373,19 @@ public void testExecuteStatementAsync() throws Exception { databricksSdkClient.executeStatementAsync( STATEMENT, warehouse, sqlParams, connection.getSession(), statement); assertEquals(STATEMENT_ID, statement.getStatementId()); + assertEquals(SESSION_ID, statement.getOriginatingSessionId()); assertNull(resultSet.getMetaData()); + assertEquals( + UPDATED_SESSION_VERSION, connection.getSession().getSessionVersion().getVersionId()); // Verify a Request with POST method is created and executed - verify(apiClient).serialize(any(ExecuteStatementRequest.class)); + verify(apiClient) + .serialize( + argThat( + request -> + request instanceof ExecuteStatementRequest + && sessionVersion(INITIAL_SESSION_VERSION) + .equals(((ExecuteStatementRequest) request).getSessionVersion()))); verify(apiClient) .execute( argThat( @@ -400,6 +538,128 @@ public void testGetStatementResult_CancelledState_ThrowsWithHY008() throws Excep assertEquals(1008, exception.getErrorCode()); // EXECUTE_STATEMENT_CANCELLED stable code } + @Test + public void testGetStatementResultUpdatesSessionVersionForMatchingSession() throws Exception { + IDatabricksConnectionContext connectionContext = + DatabricksConnectionContext.parse(JDBC_URL, new Properties()); + DatabricksSdkClient databricksSdkClient = + new DatabricksSdkClient(connectionContext, statementExecutionService, apiClient); + DatabricksConnection connection = + new DatabricksConnection(connectionContext, databricksSdkClient); + when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) + .thenReturn(createSessionResponse()); + connection.open(); + + GetStatementResponse response = + new GetStatementResponse() + .setStatementId(STATEMENT_ID.toSQLExecStatementId()) + .setStatus( + new StatementStatus() + .setState(StatementState.SUCCEEDED) + .setSessionVersion(sessionVersion(UPDATED_SESSION_VERSION))); + when(apiClient.execute(any(Request.class), eq(GetStatementResponse.class))) + .thenReturn(response); + DatabricksStatement statement = new DatabricksStatement(connection); + statement.setStatementId(STATEMENT_ID, SESSION_ID); + + databricksSdkClient.getStatementResult(STATEMENT_ID, connection.getSession(), statement); + + assertEquals( + UPDATED_SESSION_VERSION, connection.getSession().getSessionVersion().getVersionId()); + } + + @Test + public void testGetStatementResultDoesNotUpdateSessionVersionForDifferentSession() + throws Exception { + IDatabricksConnectionContext connectionContext = + DatabricksConnectionContext.parse(JDBC_URL, new Properties()); + DatabricksSdkClient databricksSdkClient = + new DatabricksSdkClient(connectionContext, statementExecutionService, apiClient); + DatabricksConnection connection = + new DatabricksConnection(connectionContext, databricksSdkClient); + when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) + .thenReturn(createSessionResponse()); + connection.open(); + + GetStatementResponse response = + new GetStatementResponse() + .setStatementId(STATEMENT_ID.toSQLExecStatementId()) + .setStatus( + new StatementStatus() + .setState(StatementState.SUCCEEDED) + .setSessionVersion(sessionVersion(UPDATED_SESSION_VERSION))); + when(apiClient.execute(any(Request.class), eq(GetStatementResponse.class))) + .thenReturn(response); + DatabricksStatement statement = new DatabricksStatement(connection); + statement.setStatementId(STATEMENT_ID, "different_session_id"); + + databricksSdkClient.getStatementResult(STATEMENT_ID, connection.getSession(), statement); + + assertEquals( + INITIAL_SESSION_VERSION, connection.getSession().getSessionVersion().getVersionId()); + } + + @Test + public void testGetStatementResultDoesNotUpdateSessionVersionForReattachedStatement() + throws Exception { + IDatabricksConnectionContext connectionContext = + DatabricksConnectionContext.parse(JDBC_URL, new Properties()); + DatabricksSdkClient databricksSdkClient = + new DatabricksSdkClient(connectionContext, statementExecutionService, apiClient); + DatabricksConnection connection = + new DatabricksConnection(connectionContext, databricksSdkClient); + when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) + .thenReturn(createSessionResponse()); + connection.open(); + + GetStatementResponse response = + new GetStatementResponse() + .setStatementId(STATEMENT_ID.toSQLExecStatementId()) + .setStatus( + new StatementStatus() + .setState(StatementState.SUCCEEDED) + .setSessionVersion(sessionVersion(UPDATED_SESSION_VERSION))); + when(apiClient.execute(any(Request.class), eq(GetStatementResponse.class))) + .thenReturn(response); + DatabricksStatement reattachedStatement = + (DatabricksStatement) connection.getStatement(STATEMENT_ID.toString()); + + databricksSdkClient.getStatementResult( + STATEMENT_ID, connection.getSession(), reattachedStatement); + + assertEquals( + INITIAL_SESSION_VERSION, connection.getSession().getSessionVersion().getVersionId()); + } + + @Test + public void testGetStatementResultDoesNotUpdateSessionVersionWithoutStatementOwnership() + throws Exception { + IDatabricksConnectionContext connectionContext = + DatabricksConnectionContext.parse(JDBC_URL, new Properties()); + DatabricksSdkClient databricksSdkClient = + new DatabricksSdkClient(connectionContext, statementExecutionService, apiClient); + DatabricksConnection connection = + new DatabricksConnection(connectionContext, databricksSdkClient); + when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) + .thenReturn(createSessionResponse()); + connection.open(); + + GetStatementResponse response = + new GetStatementResponse() + .setStatementId(STATEMENT_ID.toSQLExecStatementId()) + .setStatus( + new StatementStatus() + .setState(StatementState.SUCCEEDED) + .setSessionVersion(sessionVersion(UPDATED_SESSION_VERSION))); + when(apiClient.execute(any(Request.class), eq(GetStatementResponse.class))) + .thenReturn(response); + + databricksSdkClient.getStatementResult(STATEMENT_ID, connection.getSession(), null); + + assertEquals( + INITIAL_SESSION_VERSION, connection.getSession().getSessionVersion().getVersionId()); + } + @Test public void testDisposition_arrowAndCloudFetchEnabled_usesExternalLinks() throws Exception { setupClientMocks(true, false); @@ -467,7 +727,7 @@ public void testExecuteStatementWithTimeout() throws Exception { new DatabricksConnection(connectionContext, databricksSdkClient); // Mock session creation - CreateSessionResponse sessionResponse = new CreateSessionResponse().setSessionId(SESSION_ID); + CreateSessionResponse sessionResponse = createSessionResponse(); when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) .thenReturn(sessionResponse); connection.open(); @@ -487,7 +747,10 @@ public void testExecuteStatementWithTimeout() throws Exception { .setStatus(new StatementStatus().setState(StatementState.RUNNING)); GetStatementResponse successStatementResponse = new GetStatementResponse() - .setStatus(new StatementStatus().setState(StatementState.SUCCEEDED)); + .setStatus( + new StatementStatus() + .setState(StatementState.SUCCEEDED) + .setSessionVersion(sessionVersion(UPDATED_SESSION_VERSION))); // Set up response sequence for execute() calls when(apiClient.execute( @@ -504,6 +767,13 @@ public void testExecuteStatementWithTimeout() throws Exception { .thenReturn(runningStatementResponse) .thenReturn(runningStatementResponse) .thenReturn(successStatementResponse); + String getStatementBody = + String.format("{\"statement_id\":\"%s\"}", STATEMENT_ID.toSQLExecStatementId()); + doAnswer( + invocation -> + invocation.getArgument(0) instanceof GetStatementRequest ? getStatementBody : null) + .when(apiClient) + .serialize(any()); assertDoesNotThrow( () -> @@ -515,6 +785,15 @@ public void testExecuteStatementWithTimeout() throws Exception { connection.getSession(), statement, null)); + assertEquals( + UPDATED_SESSION_VERSION, connection.getSession().getSessionVersion().getVersionId()); + + ArgumentCaptor pollRequestCaptor = ArgumentCaptor.forClass(Request.class); + verify(apiClient, times(3)) + .execute(pollRequestCaptor.capture(), eq(GetStatementResponse.class)); + assertTrue( + pollRequestCaptor.getAllValues().stream() + .allMatch(request -> getStatementBody.equals(request.getBodyString()))); // Verify no cancellation occurred due to timeout verify(apiClient, atLeastOnce()) @@ -542,7 +821,7 @@ public void testExecuteStatementWithTimeoutExpired() throws Exception { new DatabricksConnection(connectionContext, databricksSdkClient); // Mock session creation - CreateSessionResponse sessionResponse = new CreateSessionResponse().setSessionId(SESSION_ID); + CreateSessionResponse sessionResponse = createSessionResponse(); when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) .thenReturn(sessionResponse); connection.open(); @@ -620,7 +899,7 @@ public void testMetadataOperationUsesMetadataTimeout() throws Exception { DatabricksConnection connection = new DatabricksConnection(connectionContext, databricksSdkClient); - CreateSessionResponse sessionResponse = new CreateSessionResponse().setSessionId(SESSION_ID); + CreateSessionResponse sessionResponse = createSessionResponse(); when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) .thenReturn(sessionResponse); connection.open(); @@ -681,7 +960,7 @@ public void testNonMetadataWithNullParentHasNoTimeout() throws Exception { DatabricksConnection connection = new DatabricksConnection(connectionContext, databricksSdkClient); - CreateSessionResponse sessionResponse = new CreateSessionResponse().setSessionId(SESSION_ID); + CreateSessionResponse sessionResponse = createSessionResponse(); when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) .thenReturn(sessionResponse); connection.open(); @@ -721,7 +1000,7 @@ public void testServerSideTimeoutThrowsTimeoutException() throws Exception { new DatabricksConnection(connectionContext, databricksSdkClient); // Mock session creation - CreateSessionResponse sessionResponse = new CreateSessionResponse().setSessionId(SESSION_ID); + CreateSessionResponse sessionResponse = createSessionResponse(); when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) .thenReturn(sessionResponse); connection.open(); @@ -1175,7 +1454,7 @@ public void testExecuteStatementWithClosedStatus() throws Exception { new DatabricksConnection(connectionContext, databricksSdkClient); // Mock session creation - CreateSessionResponse sessionResponse = new CreateSessionResponse().setSessionId(SESSION_ID); + CreateSessionResponse sessionResponse = createSessionResponse(); when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) .thenReturn(sessionResponse); connection.open(); @@ -1233,7 +1512,7 @@ public void testExecuteStatementWithClosedStatusAndNoParentStatement() throws Ex new DatabricksConnection(connectionContext, databricksSdkClient); // Mock session creation - CreateSessionResponse sessionResponse = new CreateSessionResponse().setSessionId(SESSION_ID); + CreateSessionResponse sessionResponse = createSessionResponse(); when(apiClient.execute(any(Request.class), eq(CreateSessionResponse.class))) .thenReturn(sessionResponse); connection.open(); @@ -1334,6 +1613,25 @@ public void testCheckStatementAlive_succeededState_returnsTrue() throws Exceptio assertTrue(databricksSdkClient.checkStatementAlive(STATEMENT_ID)); } + @Test + public void testCheckStatementAliveUpdatesSessionVersion() throws Exception { + IDatabricksConnectionContext connectionContext = + DatabricksConnectionContext.parse(JDBC_URL, new Properties()); + DatabricksSdkClient databricksSdkClient = + new DatabricksSdkClient(connectionContext, statementExecutionService, apiClient); + DatabricksSession session = mock(DatabricksSession.class); + SessionVersion updatedVersion = sessionVersion(UPDATED_SESSION_VERSION); + StatementStatus status = + new StatementStatus().setState(StatementState.SUCCEEDED).setSessionVersion(updatedVersion); + when(apiClient.execute(any(Request.class), eq(StatementStatus.class))).thenReturn(status); + + String originatingSessionId = "originating_session_id"; + assertTrue( + databricksSdkClient.checkStatementAlive(STATEMENT_ID, session, originatingSessionId)); + + verify(session).updateSessionVersion(originatingSessionId, updatedVersion); + } + @Test public void testCheckStatementAlive_runningState_returnsTrue() throws Exception { IDatabricksConnectionContext connectionContext = diff --git a/src/test/java/com/databricks/jdbc/model/core/SessionVersionSerializationTest.java b/src/test/java/com/databricks/jdbc/model/core/SessionVersionSerializationTest.java new file mode 100644 index 0000000000..e09327ad35 --- /dev/null +++ b/src/test/java/com/databricks/jdbc/model/core/SessionVersionSerializationTest.java @@ -0,0 +1,45 @@ +package com.databricks.jdbc.model.core; + +import static org.junit.jupiter.api.Assertions.assertEquals; + +import com.databricks.jdbc.model.client.sqlexec.CreateSessionRequest; +import com.databricks.jdbc.model.client.sqlexec.CreateSessionResponse; +import com.databricks.jdbc.model.client.sqlexec.ExecuteStatementRequest; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import org.junit.jupiter.api.Test; + +public class SessionVersionSerializationTest { + private final ObjectMapper objectMapper = new ObjectMapper(); + + @Test + public void testSessionExecutionModeAndVersionUseProtoJsonFieldNames() throws Exception { + JsonNode createRequest = + objectMapper.valueToTree( + new CreateSessionRequest() + .setWarehouseId("warehouse") + .setExecutionMode(SessionExecutionMode.FAST)); + JsonNode executeRequest = + objectMapper.valueToTree( + new ExecuteStatementRequest() + .setSessionVersion(new SessionVersion().setVersionId(42L))); + + assertEquals("FAST", createRequest.get("execution_mode").asText()); + assertEquals(42L, executeRequest.get("session_version").get("version_id").asLong()); + } + + @Test + public void testSessionVersionsDeserializeFromCreateAndStatusResponses() throws Exception { + CreateSessionResponse createResponse = + objectMapper.readValue( + "{\"session_id\":\"session\",\"session_version\":{\"version_id\":7}}", + CreateSessionResponse.class); + StatementStatus status = + objectMapper.readValue( + "{\"state\":\"SUCCEEDED\",\"session_version\":{\"version_id\":9}}", + StatementStatus.class); + + assertEquals(7L, createResponse.getSessionVersion().getVersionId()); + assertEquals(9L, status.getSessionVersion().getVersionId()); + } +}