Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 15 additions & 4 deletions src/main/java/com/databricks/jdbc/api/impl/BatchParameterSet.java
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,9 @@
import java.sql.Date;
import java.sql.Time;
import java.sql.Timestamp;
import java.util.Collections;
import java.util.Comparator;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
Expand All @@ -12,16 +14,21 @@
/**
* Immutable, position-ordered snapshot of one prepared-statement parameter set.
*
* <p>This model normalizes JDBC's one-based parameter indexes to zero-based wire ordinals. It does
* not validate parameter completeness, index continuity, or consistency with other parameter sets;
* those validations remain the backend's responsibility.
* <p>This model preserves JDBC's one-based parameter indexes. Transport adapters are responsible
* for converting them to protocol-specific wire ordinals. It does not validate parameter
* completeness, index continuity, or consistency with other parameter sets; those validations
* remain the backend's responsibility.
*/
public final class BatchParameterSet {

private final List<ImmutableSqlParameter> parameters;
private final Map<Integer, ImmutableSqlParameter> parameterBindings;

private BatchParameterSet(List<ImmutableSqlParameter> parameters) {
this.parameters = List.copyOf(parameters);
Map<Integer, ImmutableSqlParameter> bindings = new LinkedHashMap<>();
this.parameters.forEach(parameter -> bindings.put(parameter.cardinal(), parameter));
this.parameterBindings = Collections.unmodifiableMap(bindings);
}

public static BatchParameterSet from(Map<Integer, ImmutableSqlParameter> parameterBindings) {
Expand All @@ -38,6 +45,10 @@ public List<ImmutableSqlParameter> getParameters() {
return parameters;
}

public Map<Integer, ImmutableSqlParameter> getParameterBindings() {
return parameterBindings;
}

public int size() {
return parameters.size();
}
Expand All @@ -50,7 +61,7 @@ private static ImmutableSqlParameter snapshotParameter(
Map.Entry<Integer, ImmutableSqlParameter> entry) {
ImmutableSqlParameter parameter = entry.getValue();
return ImmutableSqlParameter.builder()
.cardinal(entry.getKey() - 1)
.cardinal(entry.getKey())
.type(parameter.type())
.value(snapshotValue(parameter.value()))
.build();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ public class DatabricksPreparedStatement extends DatabricksStatement implements
JdbcLoggerFactory.getLogger(DatabricksPreparedStatement.class);
private final String sql;
private DatabricksParameterMetaData databricksParameterMetaData;
private List<DatabricksParameterMetaData> databricksBatchParameterMetaData;
private List<BatchParameterSet> batchParameterSets;
private final boolean interpolateParameters;
private final int CHUNK_SIZE = 8192;

Expand All @@ -43,7 +43,7 @@ public DatabricksPreparedStatement(DatabricksConnection connection, String sql)
this.sql = sql;
this.interpolateParameters = connection.getConnectionContext().supportManyParameters();
this.databricksParameterMetaData = new DatabricksParameterMetaData(sql);
this.databricksBatchParameterMetaData = new ArrayList<>();
this.batchParameterSets = new ArrayList<>();
// Cache whether this statement should return a ResultSet (based on SQL and config)
this.shouldReturnResultSet = shouldReturnResultSetWithConfig(sql);
}
Expand All @@ -58,7 +58,7 @@ public DatabricksPreparedStatement(DatabricksConnection connection, String sql)
this.sql = sql;
this.interpolateParameters = interpolateParameters;
this.databricksParameterMetaData = databricksParameterMetaData;
this.databricksBatchParameterMetaData = new ArrayList<>();
this.batchParameterSets = new ArrayList<>();
// Cache whether this statement should return a ResultSet (based on SQL and config)
this.shouldReturnResultSet = shouldReturnResultSetWithConfig(sql);
}
Expand Down Expand Up @@ -110,7 +110,7 @@ public int[] executeBatch() throws DatabricksBatchUpdateException {
public long[] executeLargeBatch() throws DatabricksBatchUpdateException {
LOGGER.debug("public long executeLargeBatch()");

if (databricksBatchParameterMetaData.isEmpty()) {
if (batchParameterSets.isEmpty()) {
return new long[0];
}

Expand All @@ -123,7 +123,7 @@ public long[] executeLargeBatch() throws DatabricksBatchUpdateException {
(sqlToExecute, params, statementType, closeStatement) ->
executeInternal(sqlToExecute, params, statementType, closeStatement));

long[] updateCounts = batchExecutor.executeBatch(databricksBatchParameterMetaData);
long[] updateCounts = batchExecutor.executeBatch(batchParameterSets);

// Clear the batch after successful execution per JDBC spec
try {
Expand Down Expand Up @@ -371,7 +371,8 @@ public boolean execute() throws SQLException {
@Override
public void addBatch() {
LOGGER.debug("public void addBatch()");
this.databricksBatchParameterMetaData.add(databricksParameterMetaData);
this.batchParameterSets.add(
BatchParameterSet.from(databricksParameterMetaData.getParameterBindings()));
this.databricksParameterMetaData = new DatabricksParameterMetaData(sql);
}

Expand All @@ -380,7 +381,7 @@ public void clearBatch() throws DatabricksSQLException {
LOGGER.debug("public void clearBatch()");
checkIfClosed();
this.databricksParameterMetaData = new DatabricksParameterMetaData(sql);
this.databricksBatchParameterMetaData = new ArrayList<>();
this.batchParameterSets = new ArrayList<>();
}

@Override
Expand Down Expand Up @@ -755,7 +756,7 @@ private void checkLength(long targetLength, long sourceLength) throws SQLExcepti
}

private void checkIfBatchOperation() throws DatabricksSQLException {
if (!this.databricksBatchParameterMetaData.isEmpty()) {
if (!this.batchParameterSets.isEmpty()) {
String errorMessage =
"Batch must either be executed with executeBatch() or cleared with clearBatch()";
LOGGER.error(errorMessage);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -41,18 +41,18 @@ class LegacyPreparedStatementBatchExecutor {
this.statementExecutor = statementExecutor;
}

long[] executeBatch(List<DatabricksParameterMetaData> batchParameterMetaData)
long[] executeBatch(List<BatchParameterSet> batchParameterSets)
throws DatabricksBatchUpdateException {
if (batchParameterMetaData.isEmpty()) {
if (batchParameterSets.isEmpty()) {
return new long[0];
}

// Try to optimize INSERT statements with multi-row batching
if (canUseBatchedInsert()) {
return executeBatchedInsert(batchParameterMetaData);
return executeBatchedInsert(batchParameterSets);
} else {
// Fall back to individual execution for non-INSERT or incompatible statements
return executeIndividualStatements(batchParameterMetaData);
return executeIndividualStatements(batchParameterSets);
}
}

Expand All @@ -76,9 +76,9 @@ private boolean canUseBatchedInsert() {
}
}

private long[] executeBatchedInsert(List<DatabricksParameterMetaData> batchParameterMetaData)
private long[] executeBatchedInsert(List<BatchParameterSet> batchParameterSets)
throws DatabricksBatchUpdateException {
LOGGER.debug("Executing batched INSERT with {} rows", batchParameterMetaData.size());
LOGGER.debug("Executing batched INSERT with {} rows", batchParameterSets.size());

try {
InsertStatementParser.InsertInfo insertInfo = InsertStatementParser.parseInsertStrict(sql);
Expand All @@ -98,7 +98,7 @@ private long[] executeBatchedInsert(List<DatabricksParameterMetaData> batchParam
"BatchInsertSize must be at least 1, got: " + configuredBatchSize,
DatabricksDriverErrorCode.INVALID_STATE);
}
maxRowsPerChunk = Math.min(configuredBatchSize, batchParameterMetaData.size());
maxRowsPerChunk = Math.min(configuredBatchSize, batchParameterSets.size());
} else {
// When using parameterized queries, respect the 256 parameter limit from Databricks
// backend
Expand All @@ -113,13 +113,13 @@ private long[] executeBatchedInsert(List<DatabricksParameterMetaData> batchParam
}
}

long[] allUpdateCounts = new long[batchParameterMetaData.size()];
long[] allUpdateCounts = new long[batchParameterSets.size()];

// Process batches in chunks
for (int startIndex = 0;
startIndex < batchParameterMetaData.size();
startIndex < batchParameterSets.size();
startIndex += maxRowsPerChunk) {
int endIndex = Math.min(startIndex + maxRowsPerChunk, batchParameterMetaData.size());
int endIndex = Math.min(startIndex + maxRowsPerChunk, batchParameterSets.size());
int chunkSize = endIndex - startIndex;

// Build multi-row INSERT for this chunk
Expand All @@ -128,7 +128,7 @@ private long[] executeBatchedInsert(List<DatabricksParameterMetaData> batchParam
int paramIndex = 1;

for (int i = startIndex; i < endIndex; i++) {
DatabricksParameterMetaData batchParams = batchParameterMetaData.get(i);
BatchParameterSet batchParams = batchParameterSets.get(i);
Map<Integer, ImmutableSqlParameter> rowParams = batchParams.getParameterBindings();
for (int j = 1; j <= rowParams.size(); j++) {
if (rowParams.containsKey(j)) {
Expand Down Expand Up @@ -161,7 +161,7 @@ private long[] executeBatchedInsert(List<DatabricksParameterMetaData> batchParam
} catch (Exception e) {
// Unexpected exception - mark all as failed
LOGGER.error("Unexpected error executing batched INSERT: {}", e.getMessage(), e);
long[] failedCounts = new long[batchParameterMetaData.size()];
long[] failedCounts = new long[batchParameterSets.size()];
for (int i = 0; i < failedCounts.length; i++) {
failedCounts[i] = Statement.EXECUTE_FAILED;
}
Expand All @@ -170,22 +170,17 @@ private long[] executeBatchedInsert(List<DatabricksParameterMetaData> batchParam
}
}

private long[] executeIndividualStatements(
List<DatabricksParameterMetaData> batchParameterMetaData)
private long[] executeIndividualStatements(List<BatchParameterSet> batchParameterSets)
throws DatabricksBatchUpdateException {
LOGGER.debug("Executing batch individually with {} statements", batchParameterMetaData.size());
long[] largeUpdateCount = new long[batchParameterMetaData.size()];
LOGGER.debug("Executing batch individually with {} statements", batchParameterSets.size());
long[] largeUpdateCount = new long[batchParameterSets.size()];

for (int sqlQueryIndex = 0; sqlQueryIndex < batchParameterMetaData.size(); sqlQueryIndex++) {
DatabricksParameterMetaData databricksParameterMetaData =
batchParameterMetaData.get(sqlQueryIndex);
for (int sqlQueryIndex = 0; sqlQueryIndex < batchParameterSets.size(); sqlQueryIndex++) {
BatchParameterSet batchParameterSet = batchParameterSets.get(sqlQueryIndex);
try {
DatabricksResultSet resultSet =
statementExecutor.execute(
sql,
databricksParameterMetaData.getParameterBindings(),
StatementType.UPDATE,
false);
sql, batchParameterSet.getParameterBindings(), StatementType.UPDATE, false);
largeUpdateCount[sqlQueryIndex] = resultSet.getUpdateCount();
} catch (Exception e) {
LOGGER.error(
Expand Down
Original file line number Diff line number Diff line change
@@ -1,14 +1,31 @@
package com.databricks.jdbc.api.impl;

import com.databricks.jdbc.common.StatementType;
import com.databricks.jdbc.common.util.InsertStatementParser;
import com.databricks.jdbc.exception.DatabricksBatchUpdateException;
import java.sql.SQLException;
import java.util.List;
import java.util.Map;

class PreparedStatementBatchExecutor {

private static final NativeBatchExecutor UNSUPPORTED_NATIVE_EXECUTOR =
new NativeBatchExecutor() {
@Override
public boolean isSupported() {
return false;
}

@Override
public long[] execute(String sql, List<BatchParameterSet> parameterSets) {
throw new IllegalStateException("Native batch execution is not supported");
}
};

private final String sql;
private final DatabricksConnection connection;
private final LegacyPreparedStatementBatchExecutor legacyExecutor;
private final NativeBatchExecutor nativeExecutor;

@FunctionalInterface
interface StatementExecutor {
Expand All @@ -20,18 +37,47 @@ DatabricksResultSet execute(
throws SQLException;
}

interface NativeBatchExecutor {
boolean isSupported();

long[] execute(String sql, List<BatchParameterSet> parameterSets)
throws DatabricksBatchUpdateException;
}

PreparedStatementBatchExecutor(
String sql,
DatabricksConnection connection,
boolean interpolateParameters,
StatementExecutor statementExecutor) {
this(sql, connection, interpolateParameters, statementExecutor, UNSUPPORTED_NATIVE_EXECUTOR);
}

PreparedStatementBatchExecutor(
String sql,
DatabricksConnection connection,
boolean interpolateParameters,
StatementExecutor statementExecutor,
NativeBatchExecutor nativeExecutor) {
this.sql = sql;
this.connection = connection;
this.legacyExecutor =
new LegacyPreparedStatementBatchExecutor(
sql, connection, interpolateParameters, statementExecutor);
this.nativeExecutor = nativeExecutor;
}

long[] executeBatch(List<DatabricksParameterMetaData> batchParameterMetaData)
long[] executeBatch(List<BatchParameterSet> batchParameterSets)
throws DatabricksBatchUpdateException {
return legacyExecutor.executeBatch(batchParameterMetaData);
if (canUseNativeBatching(batchParameterSets)) {
return nativeExecutor.execute(sql, batchParameterSets);
}
return legacyExecutor.executeBatch(batchParameterSets);
}

private boolean canUseNativeBatching(List<BatchParameterSet> batchParameterSets) {
return !batchParameterSets.isEmpty()
&& connection.getConnectionContext().isNativeBatchingEnabled()
&& InsertStatementParser.isParametrizedInsert(sql)
&& nativeExecutor.isSupported();
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
class BatchParameterSetTest {

@Test
void ordersParametersByJdbcIndexAndUsesZeroBasedOrdinals() {
void ordersParametersAndPreservesJdbcIndexes() {
Map<Integer, ImmutableSqlParameter> bindings = new HashMap<>();
bindings.put(3, parameter(99, "third", ColumnInfoTypeName.STRING));
bindings.put(1, parameter(99, "first", ColumnInfoTypeName.STRING));
Expand All @@ -26,7 +26,8 @@ void ordersParametersByJdbcIndexAndUsesZeroBasedOrdinals() {
BatchParameterSet parameterSet = BatchParameterSet.from(bindings);

assertEquals(List.of("first", "second", "third"), values(parameterSet));
assertEquals(List.of(0, 1, 2), ordinals(parameterSet));
assertEquals(List.of(1, 2, 3), indexes(parameterSet));
assertEquals(List.of(1, 2, 3), List.copyOf(parameterSet.getParameterBindings().keySet()));
}

@Test
Expand All @@ -38,7 +39,7 @@ void preservesSparseIndexesWithoutValidation() {
BatchParameterSet parameterSet = BatchParameterSet.from(bindings);

assertEquals(List.of("first", "third"), values(parameterSet));
assertEquals(List.of(0, 2), ordinals(parameterSet));
assertEquals(List.of(1, 3), indexes(parameterSet));
}

@Test
Expand Down Expand Up @@ -70,6 +71,12 @@ void snapshotsBindingsAndMutableValues() {
assertThrows(
UnsupportedOperationException.class,
() -> parameterSet.getParameters().add(parameter(3, "extra", ColumnInfoTypeName.STRING)));
assertThrows(
UnsupportedOperationException.class,
() ->
parameterSet
.getParameterBindings()
.put(3, parameter(3, "extra", ColumnInfoTypeName.STRING)));
}

@Test
Expand All @@ -80,7 +87,7 @@ void preservesNullValueAndType() {
ImmutableSqlParameter parameter = parameterSet.getParameters().get(0);
assertNull(parameter.value());
assertEquals(ColumnInfoTypeName.DECIMAL, parameter.type());
assertEquals(0, parameter.cardinal());
assertEquals(1, parameter.cardinal());
}

private ImmutableSqlParameter parameter(
Expand All @@ -98,7 +105,7 @@ private List<Object> values(BatchParameterSet parameterSet) {
.collect(java.util.stream.Collectors.toList());
}

private List<Integer> ordinals(BatchParameterSet parameterSet) {
private List<Integer> indexes(BatchParameterSet parameterSet) {
return parameterSet.getParameters().stream()
.map(ImmutableSqlParameter::cardinal)
.collect(java.util.stream.Collectors.toList());
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import static com.databricks.jdbc.TestConstants.*;
import static org.junit.jupiter.api.Assertions.*;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.anyMap;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.Mockito.lenient;
import static org.mockito.Mockito.when;
Expand Down Expand Up @@ -304,7 +305,7 @@ void testBatchExecution() throws Exception {
when(client.executeStatement(
eq(CALL_SQL_AS_EXECUTED),
eq(new Warehouse(WAREHOUSE_ID)),
any(HashMap.class),
anyMap(),
eq(StatementType.UPDATE),
any(IDatabricksSession.class),
eq(stmt),
Expand Down
Loading
Loading