diff --git a/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediateStatement.java b/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediateStatement.java index 03a2a8a3e..32bca8141 100644 --- a/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediateStatement.java +++ b/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediateStatement.java @@ -40,6 +40,7 @@ import com.google.cloud.spanner.pgadapter.wireoutput.WireOutput; import com.google.common.annotations.VisibleForTesting; import com.google.common.base.Preconditions; +import com.google.common.base.Strings; import java.io.DataOutputStream; import java.util.concurrent.ExecutionException; import java.util.concurrent.Future; @@ -114,7 +115,17 @@ protected IntermediateStatement( this.parsedStatement = parsedStatement; } this.connection = connectionHandler.getSpannerConnection(); - this.command = parseCommand(this.originalStatement.getSql()); + String parsedCommand = parseCommand(this.originalStatement.getSql()); + // Fall back to "SELECT" if SimpleParser could not determine the command tag, but Spanner's + // statement parser identified this statement as a query. This ensures that any query that + // produces a result set is completed with a "SELECT" command tag rather than triggering an + // EmptyQueryResponse in ControlMessage. + this.command = + Strings.isNullOrEmpty(parsedCommand) + && this.parsedStatement != null + && this.parsedStatement.isQuery() + ? "SELECT" + : parsedCommand; this.commandTag = this.command; this.outputStream = connectionHandler.getConnectionMetadata().getOutputStream(); } diff --git a/src/main/java/com/google/cloud/spanner/pgadapter/statements/SimpleParser.java b/src/main/java/com/google/cloud/spanner/pgadapter/statements/SimpleParser.java index ac747977a..21f099cb7 100644 --- a/src/main/java/com/google/cloud/spanner/pgadapter/statements/SimpleParser.java +++ b/src/main/java/com/google/cloud/spanner/pgadapter/statements/SimpleParser.java @@ -366,7 +366,9 @@ public static boolean isEmpty(String sql) { /** Returns the command tag of the given SQL string */ public static String parseCommand(String sql) { SimpleParser parser = new SimpleParser(sql); + parser.skipOpeningParentheses(); if (parser.eatKeyword("with")) { + parser.eatKeyword("recursive"); do { if (!parser.skipCommonTableExpression()) { // Return WITH as the command tag if we encounter an invalid CTE. This is for safety, as @@ -375,10 +377,12 @@ public static String parseCommand(String sql) { return "WITH"; } } while (parser.eatToken(",")); + parser.skipOpeningParentheses(); } CharSequence keyword = parser.readKeyword(); if (keyword.length() == 0) { parser.setPos(0); + parser.skipOpeningParentheses(); keyword = parser.readKeyword(); } return keyword.toString().toUpperCase(); @@ -388,13 +392,16 @@ public static String parseCommand(String sql) { public static boolean isCommand(String command, String query) { Preconditions.checkNotNull(command); Preconditions.checkNotNull(query); - return new SimpleParser(query).peekKeyword(command); + SimpleParser parser = new SimpleParser(query); + parser.skipOpeningParentheses(); + return parser.peekKeyword(command); } public static boolean isCommand(ImmutableList commands, String query) { Preconditions.checkNotNull(commands); Preconditions.checkNotNull(query); SimpleParser parser = new SimpleParser(query); + parser.skipOpeningParentheses(); for (String command : commands) { if (!parser.eatKeyword(command)) { return false; @@ -457,10 +464,18 @@ boolean skipCommonTableExpression() { if (!eatKeyword("as")) { return false; } + eatKeyword("not"); + eatKeyword("materialized"); if (!eatToken("(")) { return false; } - parseExpressionUntilKeyword(ImmutableList.of()); + // Parse the CTE query expression until the matching closing parenthesis of 'AS (...)'. + // Do not stop at commas, as the CTE query may contain comma-separated expressions or columns. + parseExpressionUntilKeyword( + ImmutableList.of(), + /* sameParensLevelAsStart= */ false, + /* stopAtEndOfExpression= */ true, + /* stopAtComma= */ false); if (!eatToken(")")) { return false; } @@ -763,6 +778,13 @@ boolean eatToken(String token) { return eat(true, false, token); } + /** Skips any opening parentheses at the current position. */ + private void skipOpeningParentheses() { + while (eatToken("(")) { + // Opening parentheses are ignored when parsing the command. + } + } + /** Eats everything until an end parentheses at the same level as the current level. */ String eatSubExpression() { int start = pos; diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/JdbcMockServerTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/JdbcMockServerTest.java index 399e81bb3..c1f01f45d 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/JdbcMockServerTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/JdbcMockServerTest.java @@ -2618,6 +2618,60 @@ public void testMultipleQueriesInTransaction() throws SQLException { } } + @Test + public void testQueryStartingWithParenthesis() throws SQLException { + String[] queries = new String[] {"(SELECT 1)", "((SELECT 1))", "/* comment */ (SELECT 1)"}; + for (String sql : queries) { + mockSpanner.putStatementResult(StatementResult.query(Statement.of(sql), SELECT1_RESULTSET)); + } + + try (Connection connection = DriverManager.getConnection(createUrl())) { + for (String sql : queries) { + try (java.sql.Statement statement = connection.createStatement(); + ResultSet resultSet = statement.executeQuery(sql)) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + try (java.sql.Statement statement = connection.createStatement()) { + assertTrue(statement.execute(sql)); + try (ResultSet resultSet = statement.getResultSet()) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + } + try (PreparedStatement statement = connection.prepareStatement(sql); + ResultSet resultSet = statement.executeQuery()) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + } + + // Also verify parameterized query starting with parenthesis. + String parameterizedSql = "(SELECT ?)"; + mockSpanner.putStatementResult( + StatementResult.query( + Statement.newBuilder("(SELECT $1)").bind("p1").to(1L).build(), SELECT1_RESULTSET)); + try (PreparedStatement statement = connection.prepareStatement(parameterizedSql)) { + statement.setLong(1, 1L); + try (ResultSet resultSet = statement.executeQuery()) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + } + } + + assertEquals(10, mockSpanner.countRequestsOfType(ExecuteSqlRequest.class)); + for (ExecuteSqlRequest request : mockSpanner.getRequestsOfType(ExecuteSqlRequest.class)) { + assertEquals(QueryMode.NORMAL, request.getQueryMode()); + assertTrue(request.getTransaction().hasSingleUse()); + assertTrue(request.getTransaction().getSingleUse().hasReadOnly()); + } + } + @Test public void testTransactionAbortedWithPreparedStatements() throws SQLException { String sql = "SELECT 1"; diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/JdbcSimpleModeMockServerTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/JdbcSimpleModeMockServerTest.java index ae6a56c2e..06c8c5bfc 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/JdbcSimpleModeMockServerTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/JdbcSimpleModeMockServerTest.java @@ -134,6 +134,61 @@ public void testQuery() throws SQLException { assertTrue(request.getTransaction().getSingleUse().hasReadOnly()); } + @Test + public void testQueryStartingWithParenthesis() throws SQLException { + String sql = "(SELECT 1)"; + mockSpanner.putStatementResult( + StatementResult.query(com.google.cloud.spanner.Statement.of(sql), SELECT1_RESULTSET)); + + try (Connection connection = DriverManager.getConnection(createUrl())) { + try (Statement statement = connection.createStatement(); + ResultSet resultSet = statement.executeQuery(sql)) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + try (Statement statement = connection.createStatement()) { + assertTrue(statement.execute(sql)); + try (ResultSet resultSet = statement.getResultSet()) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + } + } + + assertEquals(2, mockSpanner.countRequestsOfType(ExecuteSqlRequest.class)); + for (ExecuteSqlRequest request : mockSpanner.getRequestsOfType(ExecuteSqlRequest.class)) { + assertEquals(QueryMode.NORMAL, request.getQueryMode()); + assertEquals(sql, request.getSql()); + assertTrue(request.getTransaction().hasSingleUse()); + assertTrue(request.getTransaction().getSingleUse().hasReadOnly()); + } + + // Verify multi-statement execution with parentheses. + String multiSql = "(SELECT 1); (SELECT 2)"; + mockSpanner.putStatementResult( + StatementResult.query( + com.google.cloud.spanner.Statement.of("(SELECT 2)"), SELECT2_RESULTSET)); + try (Connection connection = DriverManager.getConnection(createUrl()); + Statement statement = connection.createStatement()) { + assertTrue(statement.execute(multiSql)); + try (ResultSet resultSet = statement.getResultSet()) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + assertTrue(statement.getMoreResults()); + try (ResultSet resultSet = statement.getResultSet()) { + assertTrue(resultSet.next()); + assertEquals(2L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + assertFalse(statement.getMoreResults()); + assertEquals(-1, statement.getUpdateCount()); + } + } + @Test public void testQueryHint() throws SQLException { String sql = "/* @OPTIMIZER_VERSION=1 */ SELECT 1"; diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/statements/IntermediateStatementTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/statements/IntermediateStatementTest.java index 44bce6c5b..571063924 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/statements/IntermediateStatementTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/statements/IntermediateStatementTest.java @@ -170,4 +170,31 @@ public void testInterruptedWhileWaitingForResult() throws Exception { PGException pgException = statement.getException(); assertEquals(SQLState.QueryCanceled, pgException.getSQLState()); } + + @Test + public void testCommandFallbackForQuery() { + when(connectionHandler.getSpannerConnection()).thenReturn(connection); + when(connectionHandler.getConnectionMetadata()).thenReturn(connectionMetadata); + + ParsedStatement parsedQuery = mock(ParsedStatement.class); + when(parsedQuery.isQuery()).thenReturn(true); + + IntermediateStatement statement = + new IntermediateStatement( + mock(OptionsMetadata.class), parsedQuery, Statement.of(""), connectionHandler); + assertEquals("SELECT", statement.getCommand()); + assertEquals("SELECT", statement.getCommandTag()); + } + + @Test + public void testCommandWithNullParsedStatement() { + when(connectionHandler.getSpannerConnection()).thenReturn(connection); + when(connectionHandler.getConnectionMetadata()).thenReturn(connectionMetadata); + + IntermediateStatement statement = + new IntermediateStatement( + mock(OptionsMetadata.class), null, Statement.of(""), connectionHandler); + assertEquals("", statement.getCommand()); + assertEquals("", statement.getCommandTag()); + } } diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/statements/SimpleParserTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/statements/SimpleParserTest.java index 349b8c666..72ac3efce 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/statements/SimpleParserTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/statements/SimpleParserTest.java @@ -606,6 +606,74 @@ public void testParseCommand() { "UPDATE", parseCommand("with my_cte as (select a * b from foo) update bar set col1='one'")); assertEquals( "UPDATE", parseCommand("with my_cte as (select a - b from foo) update bar set col1='one'")); + + assertEquals("SELECT", parseCommand("(select * from foo)")); + assertEquals("SELECT", parseCommand("((select * from foo))")); + assertEquals("SELECT", parseCommand("(((select * from foo)))")); + assertEquals("SELECT", parseCommand("( ( ( select 1 ) ) )")); + assertEquals("SELECT", parseCommand("/* this is a comment */ (select * from foo)")); + assertEquals("SELECT", parseCommand("(/* this is a comment */ select * from foo)")); + assertEquals("SELECT", parseCommand("(select 1) union (select 2)")); + assertEquals("SELECT", parseCommand("(select 1) order by 1")); + assertEquals( + "SELECT", parseCommand("with my_cte as (select * from foo) (select * from my_cte)")); + assertEquals( + "SELECT", parseCommand("with my_cte as (select * from foo) ((select * from my_cte))")); + assertEquals( + "SELECT", parseCommand("(with my_cte as (select * from foo) select * from my_cte)")); + assertEquals( + "SELECT", parseCommand("((with my_cte as (select * from foo) ((select * from my_cte))))")); + assertEquals( + "SELECT", parseCommand("with recursive my_cte as (select 1) select * from my_cte")); + assertEquals( + "SELECT", parseCommand("(with recursive my_cte as (select 1) (select * from my_cte))")); + assertEquals("WITH", parseCommand("(with my_cte as (select * from foo))")); + assertEquals("UPDATE", parseCommand("with my_cte as (select 1, 2) update bar set col1='one'")); + assertEquals("SELECT", parseCommand("with my_cte as (select 1, 2) select * from my_cte")); + assertEquals( + "SELECT", + parseCommand( + "with my_cte as (select 1, 2), my_cte2 as (select 3, 4) select * from my_cte")); + assertEquals( + "UPDATE", + parseCommand( + "(with my_cte as (select 1, 2), my_cte2 as (select 3, 4) update bar set col1='one')")); + assertEquals( + "UPDATE", + parseCommand("with my_cte as (select count(*) from foo) update bar set col1='one'")); + assertEquals( + "UPDATE", parseCommand("with my_cte as (select (1) from foo) update bar set col1='one'")); + assertEquals( + "SELECT", parseCommand("with my_cte as materialized (select 1) select * from my_cte")); + assertEquals( + "SELECT", parseCommand("with my_cte as not materialized (select 1) select * from my_cte")); + assertEquals( + "SELECT", parseCommand("(with my_cte as materialized (select 1) select * from my_cte)")); + assertEquals("", parseCommand("()")); + assertEquals("", parseCommand("(( ))")); + assertEquals("", parseCommand("(/* only a comment */)")); + } + + @Test + public void testIsCommand() { + assertTrue(SimpleParser.isCommand("select", "select 1")); + assertTrue(SimpleParser.isCommand("select", "(select 1)")); + assertTrue(SimpleParser.isCommand("select", "((select 1))")); + assertTrue(SimpleParser.isCommand("select", "/* comment */ (select 1)")); + assertFalse(SimpleParser.isCommand("update", "(select 1)")); + + assertTrue( + SimpleParser.isCommand( + ImmutableList.of("select", "current_setting"), "select current_setting('foo')")); + assertTrue( + SimpleParser.isCommand( + ImmutableList.of("select", "current_setting"), "(select current_setting('foo'))")); + assertTrue( + SimpleParser.isCommand( + ImmutableList.of("select", "current_setting"), "((select current_setting('foo')))")); + assertFalse( + SimpleParser.isCommand( + ImmutableList.of("select", "current_setting"), "(select other_function('foo'))")); } @Test