diff --git a/java-bigquery-jdbc/src/main/java/com/google/cloud/bigquery/jdbc/BigQueryStatement.java b/java-bigquery-jdbc/src/main/java/com/google/cloud/bigquery/jdbc/BigQueryStatement.java index 30c11d13d263..878295f3a8af 100644 --- a/java-bigquery-jdbc/src/main/java/com/google/cloud/bigquery/jdbc/BigQueryStatement.java +++ b/java-bigquery-jdbc/src/main/java/com/google/cloud/bigquery/jdbc/BigQueryStatement.java @@ -602,6 +602,17 @@ ExecuteResult executeJob(QueryJobConfiguration jobConfiguration) return new ExecuteResult(tableResult, job); } + private StatementType getStatementType(ExecuteResult executeResult) { + if (executeResult.tableResult != null + && executeResult.tableResult.getJobStatistics() instanceof QueryStatistics) { + return ((QueryStatistics) executeResult.tableResult.getJobStatistics()).getStatementType(); + } + if (executeResult.job != null && executeResult.job.getStatistics() instanceof QueryStatistics) { + return ((QueryStatistics) executeResult.job.getStatistics()).getStatementType(); + } + return null; + } + /** * Execute the SQL script and sets the reference of the underlying job, passing null querySettings * will result in the FastQueryPath @@ -620,10 +631,7 @@ void runQuery(String query, QueryJobConfiguration jobConfiguration) try { resetStatementFields(); ExecuteResult executeResult = executeJob(jobConfiguration); - StatementType statementType = - executeResult.job == null - ? getStatementType(jobConfiguration) - : ((QueryStatistics) executeResult.job.getStatistics()).getStatementType(); + StatementType statementType = getStatementType(executeResult); SqlType queryType = getQueryType(jobConfiguration, statementType); handleQueryResult(query, executeResult.tableResult, queryType, executeResult.job); } catch (InterruptedException ex) { @@ -709,11 +717,18 @@ void handleQueryResult(String query, TableResult results, SqlType queryType, Job break; case DML: case DML_EXTRA: - QueryStatistics dmlStats = getQueryStatisticsFromJob(results, job); - Long dmlRowCount = - (dmlStats != null && dmlStats.getNumDmlAffectedRows() != null) - ? dmlStats.getNumDmlAffectedRows() - : 0L; + Long dmlRowCount; + if (results != null && results.getJobStatistics() instanceof QueryStatistics) { + Long affectedRows = + ((QueryStatistics) results.getJobStatistics()).getNumDmlAffectedRows(); + dmlRowCount = affectedRows != null ? affectedRows : 0L; + } else { + QueryStatistics dmlStats = getQueryStatisticsFromJob(results, job); + dmlRowCount = + (dmlStats != null && dmlStats.getNumDmlAffectedRows() != null) + ? dmlStats.getNumDmlAffectedRows() + : 0L; + } updateAffectedRowCount(dmlRowCount); break; case TCL: diff --git a/java-bigquery-jdbc/src/test/java/com/google/cloud/bigquery/jdbc/BigQueryStatementTest.java b/java-bigquery-jdbc/src/test/java/com/google/cloud/bigquery/jdbc/BigQueryStatementTest.java index 91ae858a6ab8..f56136b04a56 100644 --- a/java-bigquery-jdbc/src/test/java/com/google/cloud/bigquery/jdbc/BigQueryStatementTest.java +++ b/java-bigquery-jdbc/src/test/java/com/google/cloud/bigquery/jdbc/BigQueryStatementTest.java @@ -20,6 +20,8 @@ import static com.google.common.truth.Truth.assertThat; 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.assertNull; import static org.junit.jupiter.api.Assertions.assertTrue; import static org.junit.jupiter.api.Assertions.fail; import static org.mockito.ArgumentMatchers.any; @@ -169,6 +171,14 @@ private TableResult setupMockQueryResults(JobId jobId, StatementType type, Long TableResult tableResultMock = mock(TableResult.class); doReturn(jobId).when(tableResultMock).getJobId(); doReturn(Schema.of()).when(tableResultMock).getSchema(); + if (type != null || affectedRows != null) { + QueryStatistics queryStatsMock = mock(QueryStatistics.class); + doReturn(type).when(queryStatsMock).getStatementType(); + if (affectedRows != null) { + doReturn(affectedRows).when(queryStatsMock).getNumDmlAffectedRows(); + } + doReturn(queryStatsMock).when(tableResultMock).getJobStatistics(); + } doReturn(tableResultMock) .when(bigquery) .queryWithTimeout(any(QueryJobConfiguration.class), any(), any()); @@ -447,6 +457,9 @@ public void testJoblessQuery() throws SQLException, InterruptedException { TableResult tableResultMock = mock(TableResult.class); doReturn("queryId").when(tableResultMock).getQueryId(); doReturn(null).when(tableResultMock).getJobId(); + QueryStatistics queryStatsMock = mock(QueryStatistics.class); + doReturn(StatementType.SELECT).when(queryStatsMock).getStatementType(); + doReturn(queryStatsMock).when(tableResultMock).getJobStatistics(); doReturn(tableResultMock) .when(bigquery) .queryWithTimeout(any(QueryJobConfiguration.class), any(), any()); @@ -454,17 +467,10 @@ public void testJoblessQuery() throws SQLException, InterruptedException { .when(joblessStatementSpy) .processJsonResultSet(eq(tableResultMock), any()); - Job dryRunJobMock = getJobMock(null, null, StatementType.SELECT); - ArgumentCaptor dryRunCaptor = ArgumentCaptor.forClass(JobInfo.class); - doReturn(dryRunJobMock).when(bigquery).create(dryRunCaptor.capture()); - joblessStatementSpy.executeQuery("SELECT 1"); verify(bigquery).queryWithTimeout(any(QueryJobConfiguration.class), any(), any()); - verify(bigquery).create(any(JobInfo.class)); - assertTrue( - Boolean.TRUE.equals( - ((QueryJobConfiguration) dryRunCaptor.getValue().getConfiguration()).dryRun())); + verify(bigquery, Mockito.never()).create(any(JobInfo.class)); // 2. Test JobCreationMode=1 (jobful) Mockito.reset(bigquery); @@ -914,6 +920,9 @@ public void testExecute_propagatesContextAndBaggage() throws Exception { TableResult tableResultMock = mock(TableResult.class); doReturn(jobId).when(tableResultMock).getJobId(); doReturn(Schema.of()).when(tableResultMock).getSchema(); + QueryStatistics queryStatsMock = mock(QueryStatistics.class); + doReturn(StatementType.SELECT).when(queryStatsMock).getStatementType(); + doReturn(queryStatsMock).when(tableResultMock).getJobStatistics(); return tableResultMock; }) .when(bigquery) @@ -923,8 +932,6 @@ public void testExecute_propagatesContextAndBaggage() throws Exception { // Setup connection mocks to allow the statement to execute successfully doReturn(true).when(bigQueryConnection).getUseStatelessQueryMode(); - Job dryRunJobMock = getJobMock(null, null, StatementType.SELECT); - doReturn(dryRunJobMock).when(bigquery).create(Mockito.any(JobInfo.class)); BigQueryJsonResultSet resultSetMock = mock(BigQueryJsonResultSet.class); doReturn(resultSetMock) @@ -937,6 +944,7 @@ public void testExecute_propagatesContextAndBaggage() throws Exception { // Verify the SDK call actually occurred verify(bigquery) .queryWithTimeout(Mockito.any(QueryJobConfiguration.class), Mockito.any(), Mockito.any()); + verify(bigquery, Mockito.never()).create(Mockito.any(JobInfo.class)); } @Test @@ -1065,13 +1073,12 @@ public void testTemporaryDatasetCreationRespectsConnectionLocation() // 2. Mock bigQuery.getDataset to return null (triggering creation) doReturn(null).when(bigquery).getDataset(eq(DatasetId.of("temp_dataset"))); - // 2b. Mock bigQuery.create for dry run during getStatementType - Job dryRunJobMock = getJobMock(null, null, StatementType.SELECT); - doReturn(dryRunJobMock).when(bigquery).create(any(JobInfo.class)); - // 3. Mock bigquery.queryWithTimeout(...) to return tableResult (so execution doesn't fail on // query execution) TableResult result = mock(TableResult.class); + QueryStatistics queryStatsMock = mock(QueryStatistics.class); + doReturn(StatementType.SELECT).when(queryStatsMock).getStatementType(); + doReturn(queryStatsMock).when(result).getJobStatistics(); doReturn(result) .when(bigquery) .queryWithTimeout(any(QueryJobConfiguration.class), any(JobId.class), any()); @@ -1091,4 +1098,32 @@ public void testTemporaryDatasetCreationRespectsConnectionLocation() assertEquals("temp_dataset", createdDatasetInfo.getDatasetId().getDataset()); assertEquals("europe-west3", createdDatasetInfo.getLocation()); } + + @Test + public void testStatelessQueryExecutionDoesNotInvokeDryRun() throws Exception { + TableResult tableResultMock = setupMockQueryResults(null, StatementType.SELECT, null); + BigQueryStatement statementSpy = Mockito.spy(bigQueryStatement); + doReturn(mock(BigQueryJsonResultSet.class)) + .when(statementSpy) + .processJsonResultSet(eq(tableResultMock), any()); + + boolean hasResultSet = statementSpy.execute("SELECT 1"); + + assertTrue(hasResultSet); + assertNotNull(statementSpy.getResultSet()); + verify(bigquery, Mockito.never()).create(any(JobInfo.class)); + } + + @Test + public void testStatelessDmlExecutionUsesTableResultWithoutDryRunOrGetJob() throws Exception { + setupMockQueryResults(null, StatementType.UPDATE, 15L); + + int updatedCount = bigQueryStatement.executeUpdate("UPDATE dataset.table SET col = 1"); + + assertEquals(15, updatedCount); + assertEquals(15L, bigQueryStatement.getLargeUpdateCount()); + assertNull(bigQueryStatement.getResultSet()); + verify(bigquery, Mockito.never()).create(any(JobInfo.class)); + verify(bigquery, Mockito.never()).getJob(any(JobId.class)); + } }