From cbe03f39a7f11103eb79e3853c08e1ebeb09671e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Knut=20Olav=20L=C3=B8ite?= Date: Thu, 24 Sep 2026 16:54:59 +0200 Subject: [PATCH 1/2] fix: do not reserve statement name on prepare failure Deregister prepared statements from the connection session if statement analysis fails during PREPARE. This ensures that a failed PREPARE does not reserve the statement name or block subsequent attempts. Also align duplicate prepared statement handling with PostgreSQL: - Return SQLSTATE 42P05 (DuplicatePreparedStatement) instead of dropping the connection on duplicate statement names. - Fold unquoted statement names to lowercase and strip identifier quotes in PREPARE to match EXECUTE and DEALLOCATE behavior. - Reject zero-length delimited identifiers (`PREPARE ""`). --- .../statements/IntermediateStatement.java | 8 + .../statements/PrepareStatement.java | 94 +++++++-- .../pgadapter/wireprotocol/ParseMessage.java | 25 ++- .../pgadapter/EmulatedPsqlMockServerTest.java | 183 ++++++++++++++++++ .../pgadapter/InvalidMessagesTest.java | 121 ++++++++++++ .../statements/PrepareStatementTest.java | 18 ++ .../wireprotocol/ParseMessageTest.java | 50 +++++ .../pgadapter/wireprotocol/ProtocolTest.java | 15 +- 8 files changed, 488 insertions(+), 26 deletions(-) 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 d273df8129..03a2a8a3eb 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 @@ -301,6 +301,14 @@ public StatementType getStatementType() { return this.parsedStatement.getType(); } + public ParsedStatement getParsedStatement() { + return this.parsedStatement; + } + + public Statement getOriginalStatement() { + return this.originalStatement; + } + public String getSql() { return this.originalStatement.getSql(); } diff --git a/src/main/java/com/google/cloud/spanner/pgadapter/statements/PrepareStatement.java b/src/main/java/com/google/cloud/spanner/pgadapter/statements/PrepareStatement.java index 36c1898f20..365235d035 100644 --- a/src/main/java/com/google/cloud/spanner/pgadapter/statements/PrepareStatement.java +++ b/src/main/java/com/google/cloud/spanner/pgadapter/statements/PrepareStatement.java @@ -14,6 +14,8 @@ package com.google.cloud.spanner.pgadapter.statements; +import static com.google.cloud.spanner.pgadapter.statements.SimpleParser.unquoteOrFoldIdentifier; + import com.google.api.core.InternalApi; import com.google.cloud.spanner.Dialect; import com.google.cloud.spanner.Statement; @@ -22,20 +24,21 @@ import com.google.cloud.spanner.connection.AbstractStatementParser.StatementType; import com.google.cloud.spanner.connection.StatementResult; import com.google.cloud.spanner.pgadapter.ConnectionHandler; +import com.google.cloud.spanner.pgadapter.error.PGException; import com.google.cloud.spanner.pgadapter.error.PGExceptionFactory; import com.google.cloud.spanner.pgadapter.error.SQLState; import com.google.cloud.spanner.pgadapter.metadata.OptionsMetadata; import com.google.cloud.spanner.pgadapter.statements.BackendConnection.NoResult; import com.google.cloud.spanner.pgadapter.statements.SimpleParser.TableOrIndexName; import com.google.cloud.spanner.pgadapter.statements.SimpleParser.TypeDefinition; -import com.google.cloud.spanner.pgadapter.wireprotocol.ControlMessage.ManuallyCreatedToken; -import com.google.cloud.spanner.pgadapter.wireprotocol.ControlMessage.PreparedType; -import com.google.cloud.spanner.pgadapter.wireprotocol.DescribeMessage; import com.google.cloud.spanner.pgadapter.wireprotocol.ParseMessage; import com.google.common.base.Preconditions; import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; +import com.google.common.util.concurrent.FutureCallback; import com.google.common.util.concurrent.Futures; +import com.google.common.util.concurrent.ListenableFuture; +import com.google.common.util.concurrent.MoreExecutors; import java.util.List; import java.util.concurrent.Future; import java.util.stream.Collectors; @@ -144,28 +147,74 @@ public StatementType getStatementType() { public void executeAsync(BackendConnection backendConnection) { if (!this.executed) { this.executed = true; + if (this.connectionHandler.hasStatement(this.preparedStatement.name)) { + PGException exception = + PGExceptionFactory.newPGException( + String.format( + "prepared statement \"%s\" already exists", this.preparedStatement.name), + SQLState.DuplicatePreparedStatement); + setFutureStatementResult(Futures.immediateFailedFuture(exception)); + backendConnection.execute( + new InvalidStatement( + this.connectionHandler, + this.options, + this.parsedStatement, + this.originalStatement, + exception)); + return; + } + IntermediatePreparedStatement statement = + ParseMessage.createStatement( + this.connectionHandler, + this.preparedStatement.name, + this.preparedStatement.parsedPreparedStatement, + this.preparedStatement.originalPreparedStatement, + this.preparedStatement.dataTypes); + if (statement instanceof InvalidStatement) { + setFutureStatementResult(Futures.immediateFailedFuture(statement.getException())); + backendConnection.execute((InvalidStatement) statement); + return; + } + ListenableFuture analyzeResult; try { - new ParseMessage( - connectionHandler, - preparedStatement.name, - preparedStatement.dataTypes, - preparedStatement.parsedPreparedStatement, - preparedStatement.originalPreparedStatement) - .send(); - new DescribeMessage( - connectionHandler, - PreparedType.Statement, - preparedStatement.name, - ManuallyCreatedToken.MANUALLY_CREATED_TOKEN) - .send(); + Future describeFuture = statement.describeAsync(backendConnection); + Preconditions.checkState( + describeFuture instanceof ListenableFuture, + "describeAsync must return an instance of ListenableFuture"); + analyzeResult = (ListenableFuture) describeFuture; } catch (Exception exception) { - setFutureStatementResult(Futures.immediateFailedFuture(exception)); + PGException pgException = PGExceptionFactory.toPGException(exception); + setFutureStatementResult(Futures.immediateFailedFuture(pgException)); backendConnection.execute( new InvalidStatement( - connectionHandler, options, parsedStatement, originalStatement, exception)); + this.connectionHandler, + this.options, + this.parsedStatement, + this.originalStatement, + pgException)); return; } - setFutureStatementResult(Futures.immediateFuture(new NoResult(getCommandTag()))); + this.connectionHandler.registerStatement(this.preparedStatement.name, statement); + Futures.addCallback( + analyzeResult, + new FutureCallback() { + @Override + public void onSuccess(StatementResult result) {} + + @Override + public void onFailure(Throwable throwable) { + if (connectionHandler.hasStatement(preparedStatement.name) + && connectionHandler.getStatement(preparedStatement.name) == statement) { + connectionHandler.closeStatement(preparedStatement.name); + } + } + }, + MoreExecutors.directExecutor()); + setFutureStatementResult( + Futures.transform( + analyzeResult, + ignored -> new NoResult(getCommandTag()), + MoreExecutors.directExecutor())); } } @@ -196,6 +245,11 @@ static ParsedPreparedStatement parse(String sql) { throw PGExceptionFactory.newPGException( "invalid prepared statement name", SQLState.InvalidSqlStatementName); } + String statementName = unquoteOrFoldIdentifier(name.name); + if (statementName == null || statementName.isEmpty()) { + throw PGExceptionFactory.newPGException( + "zero-length delimited identifier", SQLState.InvalidSqlStatementName); + } ImmutableList.Builder dataTypesBuilder = ImmutableList.builder(); if (parser.eatToken("(")) { List dataTypesNames = parser.parseExpressionList(); @@ -216,7 +270,7 @@ static ParsedPreparedStatement parse(String sql) { "missing 'AS' keyword in PREPARE statement: " + sql, SQLState.SyntaxError); } return new ParsedPreparedStatement( - name.name, + statementName, dataTypesBuilder.build().stream().mapToInt(i -> i).toArray(), parser.getSql().substring(parser.getPos()).trim()); } diff --git a/src/main/java/com/google/cloud/spanner/pgadapter/wireprotocol/ParseMessage.java b/src/main/java/com/google/cloud/spanner/pgadapter/wireprotocol/ParseMessage.java index f05abec5a4..f9c1b41c75 100644 --- a/src/main/java/com/google/cloud/spanner/pgadapter/wireprotocol/ParseMessage.java +++ b/src/main/java/com/google/cloud/spanner/pgadapter/wireprotocol/ParseMessage.java @@ -41,6 +41,8 @@ import com.google.cloud.spanner.connection.AbstractStatementParser.ParsedStatement; import com.google.cloud.spanner.connection.AbstractStatementParser.StatementType; import com.google.cloud.spanner.pgadapter.ConnectionHandler; +import com.google.cloud.spanner.pgadapter.error.PGExceptionFactory; +import com.google.cloud.spanner.pgadapter.error.SQLState; import com.google.cloud.spanner.pgadapter.statements.BackendConnection; import com.google.cloud.spanner.pgadapter.statements.CloseStatement; import com.google.cloud.spanner.pgadapter.statements.CopyStatement; @@ -74,7 +76,7 @@ public class ParseMessage extends AbstractQueryProtocolMessage { protected static final char IDENTIFIER = 'P'; private final String name; - private final IntermediatePreparedStatement statement; + private IntermediatePreparedStatement statement; private final int[] parameterDataTypes; public ParseMessage(ConnectionHandler connection) throws Exception { @@ -117,7 +119,7 @@ public ParseMessage( createStatement(connection, name, parsedStatement, originalStatement, parameterDataTypes); } - static IntermediatePreparedStatement createStatement( + public static IntermediatePreparedStatement createStatement( ConnectionHandler connectionHandler, String name, ParsedStatement parsedStatement, @@ -276,10 +278,25 @@ static IntermediatePreparedStatement createStatement( @Override void buffer(BackendConnection backendConnection) { if (!Strings.isNullOrEmpty(this.name) && this.connection.hasStatement(this.name)) { - throw new IllegalStateException("Must close statement before reusing name."); + this.statement = + new InvalidStatement( + this.connection, + this.connection.getServer().getOptions(), + this.name, + this.statement.getParsedStatement(), + this.statement.getOriginalStatement(), + PGExceptionFactory.newPGException( + String.format("prepared statement \"%s\" already exists", this.name), + SQLState.DuplicatePreparedStatement)); + if (backendConnection != null) { + backendConnection.execute((InvalidStatement) this.statement); + } + return; } if (this.statement instanceof InvalidStatement) { - backendConnection.execute((InvalidStatement) this.statement); + if (backendConnection != null) { + backendConnection.execute((InvalidStatement) this.statement); + } } this.connection.registerStatement(this.name, this.statement); } diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/EmulatedPsqlMockServerTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/EmulatedPsqlMockServerTest.java index d31cbf0954..7b0e6230c2 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/EmulatedPsqlMockServerTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/EmulatedPsqlMockServerTest.java @@ -25,6 +25,7 @@ import com.google.cloud.spanner.MockSpannerServiceImpl.StatementResult; import com.google.cloud.spanner.Statement; import com.google.cloud.spanner.pgadapter.error.PGException; +import com.google.cloud.spanner.pgadapter.error.SQLState; import com.google.cloud.spanner.pgadapter.statements.IntermediateStatement; import com.google.cloud.spanner.pgadapter.utils.ClientAutoDetector; import com.google.cloud.spanner.pgadapter.utils.ClientAutoDetector.WellKnownClient; @@ -394,6 +395,188 @@ public void testPrepareInvalidStatement() throws SQLException { } } + @Test + public void testPrepareStatementWithErrorDoesNotReserveName() throws SQLException { + mockSpanner.putStatementResult( + StatementResult.exception( + Statement.of("select bad_statement"), + Status.INVALID_ARGUMENT + .withDescription("column \"bad_statement\" does not exist") + .asRuntimeException())); + + try (Connection connection = DriverManager.getConnection(createUrl("my-db"))) { + try (java.sql.Statement statement = connection.createStatement()) { + PSQLException exception1 = + assertThrows( + PSQLException.class, + () -> + statement.execute( + "prepare test_q(timestamptz, timestamptz) as select bad_statement")); + assertTrue(exception1.getMessage().contains("column \"bad_statement\" does not exist")); + + // Running the same failed prepare statement again should return the same error, + // and NOT complain that the statement name is already reserved or that the connection is + // closed. + PSQLException exception2 = + assertThrows( + PSQLException.class, + () -> + statement.execute( + "prepare test_q(timestamptz, timestamptz) as select bad_statement")); + assertTrue(exception2.getMessage().contains("column \"bad_statement\" does not exist")); + + // Now prepare a valid statement using the same name. + statement.execute("prepare test_q as SELECT 1"); + try (ResultSet resultSet = statement.executeQuery("execute test_q")) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + + statement.execute("deallocate test_q"); + } + } + } + + @Test + public void testPrepareDuplicateStatementName() throws SQLException { + try (Connection connection = DriverManager.getConnection(createUrl("my-db"))) { + try (java.sql.Statement statement = connection.createStatement()) { + statement.execute("prepare test_dup as SELECT 1"); + + PSQLException exception = + assertThrows( + PSQLException.class, () -> statement.execute("prepare test_dup as SELECT 2")); + assertEquals("42P05", exception.getSQLState()); + assertTrue( + exception.getMessage().contains("prepared statement \"test_dup\" already exists")); + + // Verify that the original statement is still intact and runnable. + try (ResultSet resultSet = statement.executeQuery("execute test_dup")) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + + statement.execute("deallocate test_dup"); + } + } + } + + @Test + public void testPrepareCaseInsensitiveAndQuotedIdentifier() throws SQLException { + try (Connection connection = DriverManager.getConnection(createUrl("my-db"))) { + try (java.sql.Statement statement = connection.createStatement()) { + // Unquoted mixed-case identifier folds to lowercase. + statement.execute("prepare MyStmt as SELECT 1"); + try (ResultSet resultSet = statement.executeQuery("execute mystmt")) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + try (ResultSet resultSet = statement.executeQuery("execute MyStmt")) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + statement.execute("deallocate MyStmt"); + + // Quoted identifier preserves case. + statement.execute("prepare \"MyQuotedStmt\" as SELECT 2"); + try (ResultSet resultSet = statement.executeQuery("execute \"MyQuotedStmt\"")) { + assertTrue(resultSet.next()); + assertEquals(2L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + statement.execute("deallocate \"MyQuotedStmt\""); + } + } + } + + @Test + public void testPrepareZeroLengthDelimitedIdentifier() throws SQLException { + try (Connection connection = DriverManager.getConnection(createUrl("my-db"))) { + try (java.sql.Statement statement = connection.createStatement()) { + SQLException exception = + assertThrows(SQLException.class, () -> statement.execute("prepare \"\" as SELECT 1")); + assertEquals(SQLState.InvalidSqlStatementName.toString(), exception.getSQLState()); + assertTrue(exception.getMessage().contains("zero-length delimited identifier")); + } + } + } + + @Test + public void testPrepareSelectCurrentSetting() throws SQLException { + try (Connection connection = DriverManager.getConnection(createUrl("my-db"))) { + try (java.sql.Statement statement = connection.createStatement()) { + // 1. Prepare and execute SELECT current_setting(...) + statement.execute("prepare test_setting as select current_setting('application_name')"); + try (ResultSet resultSet = statement.executeQuery("execute test_setting")) { + assertTrue(resultSet.next()); + assertEquals("PostgreSQL JDBC Driver", resultSet.getString(1)); + assertFalse(resultSet.next()); + } + statement.execute("deallocate test_setting"); + + // 2. Prepare and execute SELECT set_config(...) + statement.execute( + "prepare test_set_config as select set_config('application_name', 'my-custom-app', false)"); + try (ResultSet resultSet = statement.executeQuery("execute test_set_config")) { + assertTrue(resultSet.next()); + assertEquals("my-custom-app", resultSet.getString(1)); + assertFalse(resultSet.next()); + } + statement.execute("deallocate test_set_config"); + + // 3. Prepare an invalid client-side statement (syntax error for current_setting). + SQLException exception = + assertThrows( + SQLException.class, + () -> + statement.execute("prepare test_invalid_setting as select current_setting()")); + assertTrue(exception.getMessage().contains("Invalid quote character")); + + // Verify that the failed client-side statement did not reserve the name. + statement.execute("prepare test_invalid_setting as SELECT 1"); + try (ResultSet resultSet = statement.executeQuery("execute test_invalid_setting")) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + statement.execute("deallocate test_invalid_setting"); + } + } + } + + @Test + public void testPrepareFailureInExplicitTransactionAbortsTransaction() throws SQLException { + try (Connection connection = DriverManager.getConnection(createUrl("my-db"))) { + connection.setAutoCommit(false); + try (java.sql.Statement statement = connection.createStatement()) { + SQLException prepareException = + assertThrows( + SQLException.class, + () -> + statement.execute("prepare test_invalid as select * from non_existing_table")); + assertTrue(prepareException.getMessage().contains("non_existing_table")); + + // Subsequent statements in the transaction must fail with 25P02 (transaction aborted). + SQLException abortedException = + assertThrows(SQLException.class, () -> statement.execute("SELECT 1")); + assertEquals(SQLState.InFailedSqlTransaction.toString(), abortedException.getSQLState()); + + connection.rollback(); + + // After rollback, the connection should be healthy again. + try (ResultSet resultSet = statement.executeQuery("SELECT 1")) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertFalse(resultSet.next()); + } + } + } + } + @Test public void testRoundParamValueForPreparedStatement() throws SQLException { mockSpanner.putStatementResult( diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/InvalidMessagesTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/InvalidMessagesTest.java index 1b403556f8..73000ea6b1 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/InvalidMessagesTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/InvalidMessagesTest.java @@ -657,4 +657,125 @@ public void testCreateInvalidPreparedStatement() throws IOException { } } } + + @Test + public void testParseMessageDuplicateStatementName() throws IOException { + try (Socket socket = new Socket("localhost", pgServer.getLocalPort())) { + try (DataInputStream inputStream = new DataInputStream(socket.getInputStream()); + DataOutputStream outputStream = new DataOutputStream(socket.getOutputStream())) { + // Request startup. + outputStream.writeInt(17); + outputStream.writeInt(StartupMessage.PROTOCOL_VERSION_3_0_IDENTIFIER); + outputStream.writeBytes("user"); + outputStream.writeByte(0); + outputStream.writeBytes("foo"); + outputStream.writeByte(0); + outputStream.flush(); + + // Verify that the server responds with auth OK. + assertEquals('R', inputStream.readByte()); + assertEquals(8, inputStream.readInt()); + assertEquals(0, inputStream.readInt()); // 0 == success + + // Receive key data. + assertEquals('K', inputStream.readByte()); + assertEquals(12, inputStream.readInt()); + inputStream.readInt(); + inputStream.readInt(); + + // Just skip parameter data and wait for 'Z' (ready for query) + while (true) { + byte message = inputStream.readByte(); + int length = inputStream.readInt(); + inputStream.readFully(new byte[length - 4]); + if (message == 'Z') { + break; + } + } + + // Send first PARSE for "my_stmt" + byte[] sqlBytes = "SELECT 1".getBytes(StandardCharsets.UTF_8); + byte[] nameBytes = "my_stmt".getBytes(StandardCharsets.UTF_8); + outputStream.writeByte('P'); + outputStream.writeInt(4 + nameBytes.length + 1 + sqlBytes.length + 1 + 2); + outputStream.write(nameBytes); + outputStream.writeByte(0); + outputStream.write(sqlBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); + // SYNC + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + // Wait for '1' (ParseComplete) and 'Z' (ReadyForQuery) + assertEquals('1', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + assertEquals('Z', inputStream.readByte()); + assertEquals(5, inputStream.readInt()); + inputStream.readByte(); + + // Now send duplicate PARSE for "my_stmt" + byte[] sql2Bytes = "SELECT 2".getBytes(StandardCharsets.UTF_8); + outputStream.writeByte('P'); + outputStream.writeInt(4 + nameBytes.length + 1 + sql2Bytes.length + 1 + 2); + outputStream.write(nameBytes); + outputStream.writeByte(0); + outputStream.write(sql2Bytes); + outputStream.writeByte(0); + outputStream.writeShort(0); + // SYNC + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + // Verify that we receive 'E' (ErrorResponse) followed by 'Z' (ReadyForQuery) + assertEquals('E', inputStream.readByte()); + int length = inputStream.readInt(); + byte[] errorBytes = new byte[length - 4]; + inputStream.readFully(errorBytes); + String errorPayload = new String(errorBytes, StandardCharsets.UTF_8); + assertTrue(errorPayload.contains("prepared statement \"my_stmt\" already exists")); + assertTrue(errorPayload.contains("42P05")); + + assertEquals('Z', inputStream.readByte()); + assertEquals(5, inputStream.readInt()); + inputStream.readByte(); + + // Bind and execute "my_stmt" to confirm the original statement is still valid. + outputStream.writeByte('B'); + outputStream.writeInt(4 + 1 + nameBytes.length + 1 + 2 + 2 + 2); + outputStream.writeByte(0); // Unnamed portal + outputStream.write(nameBytes); + outputStream.writeByte(0); // Statement name "my_stmt" + outputStream.writeShort(0); // Zero parameter format codes + outputStream.writeShort(0); // Zero parameter values + outputStream.writeShort(0); // Zero result format codes + + outputStream.writeByte('E'); + outputStream.writeInt(4 + 1 + 4); + outputStream.writeByte(0); // Unnamed portal + outputStream.writeInt(0); // Return all rows + + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + assertEquals('2', inputStream.readByte()); // BindComplete + assertEquals(4, inputStream.readInt()); + + assertEquals('D', inputStream.readByte()); // DataRow + int dataRowLength = inputStream.readInt(); + inputStream.readFully(new byte[dataRowLength - 4]); + + assertEquals('C', inputStream.readByte()); // CommandComplete + int commandCompleteLength = inputStream.readInt(); + inputStream.readFully(new byte[commandCompleteLength - 4]); + + assertEquals('Z', inputStream.readByte()); // ReadyForQuery + assertEquals(5, inputStream.readInt()); + inputStream.readByte(); + } + } + } } diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/statements/PrepareStatementTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/statements/PrepareStatementTest.java index ee0f0124a8..e7753f97ec 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/statements/PrepareStatementTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/statements/PrepareStatementTest.java @@ -21,6 +21,7 @@ import static org.junit.Assert.assertThrows; import com.google.cloud.spanner.pgadapter.error.PGException; +import com.google.cloud.spanner.pgadapter.error.SQLState; import com.google.cloud.spanner.pgadapter.statements.PrepareStatement.ParsedPreparedStatement; import org.junit.Test; import org.junit.runner.RunWith; @@ -126,4 +127,21 @@ public void testParse() { assertThrows(PGException.class, () -> parse("prepare (bigint) as select 1")); assertThrows(PGException.class, () -> parse("prepare as select 1")); } + + @Test + public void testParseIdentifierCaseAndQuotes() { + ParsedPreparedStatement statement = parse("prepare MyStmt as select 1"); + assertEquals("mystmt", statement.name); + + statement = parse("prepare \"MyStmt\" as select 1"); + assertEquals("MyStmt", statement.name); + + statement = parse("prepare \"my_stmt\" as select 1"); + assertEquals("my_stmt", statement.name); + + PGException exception = + assertThrows(PGException.class, () -> parse("prepare \"\" as select 1")); + assertEquals(SQLState.InvalidSqlStatementName, exception.getSQLState()); + assertEquals("zero-length delimited identifier", exception.getMessage()); + } } diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/wireprotocol/ParseMessageTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/wireprotocol/ParseMessageTest.java index 63f92c4a38..9646318cf5 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/wireprotocol/ParseMessageTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/wireprotocol/ParseMessageTest.java @@ -36,6 +36,7 @@ import com.google.cloud.spanner.pgadapter.ConnectionHandler; import com.google.cloud.spanner.pgadapter.ProxyServer; import com.google.cloud.spanner.pgadapter.error.PGException; +import com.google.cloud.spanner.pgadapter.error.SQLState; import com.google.cloud.spanner.pgadapter.metadata.ConnectionMetadata; import com.google.cloud.spanner.pgadapter.metadata.OptionsMetadata; import com.google.cloud.spanner.pgadapter.statements.BackendConnection; @@ -191,4 +192,53 @@ public void testParseMessageUnnamedAndInvalidLifecycle() throws Exception { assertTrue(parseMessage.isReturnedErrorResponse()); assertTrue(outputBytes.size() > 0); } + + @Test + public void testParseMessageDuplicateStatementName() throws Exception { + ConnectionHandler localConnection = mock(ConnectionHandler.class); + ProxyServer server = mock(ProxyServer.class); + OptionsMetadata options = mock(OptionsMetadata.class); + ConnectionMetadata metadata = mock(ConnectionMetadata.class); + ExtendedQueryProtocolHandler protocolHandler = mock(ExtendedQueryProtocolHandler.class); + BackendConnection backendConnection = mock(BackendConnection.class); + IntermediatePreparedStatement existingStatement = mock(IntermediatePreparedStatement.class); + ByteArrayOutputStream outputBytes = new ByteArrayOutputStream(); + DataOutputStream outputStream = new DataOutputStream(outputBytes); + + when(server.getOptions()).thenReturn(options); + when(localConnection.getServer()).thenReturn(server); + when(localConnection.getWellKnownClient()).thenReturn(WellKnownClient.UNSPECIFIED); + when(localConnection.getConnectionMetadata()).thenReturn(metadata); + when(metadata.getOutputStream()).thenReturn(outputStream); + when(localConnection.getExtendedQueryProtocolHandler()).thenReturn(protocolHandler); + when(protocolHandler.getBackendConnection()).thenReturn(backendConnection); + when(localConnection.hasStatement("test_stmt")).thenReturn(true); + when(localConnection.getStatement("test_stmt")).thenReturn(existingStatement); + + Statement statement = Statement.of("select 1"); + ParseMessage parseMessage = + new ParseMessage( + localConnection, "test_stmt", new int[0], PARSER.parse(statement), statement); + + // buffer() should not throw IllegalStateException, should not overwrite existingStatement, + // and should execute InvalidStatement with SQLState.DuplicatePreparedStatement. + parseMessage.buffer(backendConnection); + assertTrue(parseMessage.getStatement() instanceof InvalidStatement); + assertEquals( + SQLState.DuplicatePreparedStatement, + ((InvalidStatement) parseMessage.getStatement()).getException().getSQLState()); + verify(backendConnection).execute((InvalidStatement) parseMessage.getStatement()); + verify(localConnection, never()).registerStatement(eq("test_stmt"), any()); + + // flush() should send ErrorResponse without removing the existing statement from + // ConnectionHandler. + parseMessage.flush(); + verify(localConnection, never()).closeStatement("test_stmt"); + assertTrue(parseMessage.isReturnedErrorResponse()); + assertTrue(outputBytes.size() > 0); + + // abort() should also not close the existing statement. + parseMessage.abort(); + verify(localConnection, never()).closeStatement("test_stmt"); + } } diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/wireprotocol/ProtocolTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/wireprotocol/ProtocolTest.java index a5481e9d6d..ef8aa5f7c2 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/wireprotocol/ProtocolTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/wireprotocol/ProtocolTest.java @@ -62,6 +62,7 @@ import com.google.cloud.spanner.pgadapter.statements.ExtendedQueryProtocolHandler; import com.google.cloud.spanner.pgadapter.statements.IntermediatePortalStatement; import com.google.cloud.spanner.pgadapter.statements.IntermediatePreparedStatement; +import com.google.cloud.spanner.pgadapter.statements.InvalidStatement; import com.google.cloud.spanner.pgadapter.utils.ClientAutoDetector.WellKnownClient; import com.google.cloud.spanner.pgadapter.utils.Metrics; import com.google.cloud.spanner.pgadapter.utils.MutationWriter; @@ -505,7 +506,12 @@ public void testParseMessageExceptsIfNameIsInUse() throws Exception { WireMessage message = server.getMessageReader().create(connectionHandler); when(connectionHandler.hasStatement(anyString())).thenReturn(true); - assertThrows(IllegalStateException.class, message::send); + message.send(); + assertTrue(((ParseMessage) message).getStatement().hasException()); + assertEquals( + SQLState.DuplicatePreparedStatement, + ((InvalidStatement) ((ParseMessage) message).getStatement()).getException().getSQLState()); + verify(connectionHandler, never()).registerStatement(anyString(), any()); } @Test @@ -552,7 +558,12 @@ public void testParseMessageExceptsIfNameIsNull() throws Exception { WireMessage message = server.getMessageReader().create(connectionHandler); when(connectionHandler.hasStatement(anyString())).thenReturn(true); - assertThrows(IllegalStateException.class, message::send); + message.send(); + assertTrue(((ParseMessage) message).getStatement().hasException()); + assertEquals( + SQLState.DuplicatePreparedStatement, + ((InvalidStatement) ((ParseMessage) message).getStatement()).getException().getSQLState()); + verify(connectionHandler, never()).registerStatement(anyString(), any()); } @Test From b114fe391a4810b36cd02d1d7d15c3b058cbf67b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Knut=20Olav=20L=C3=B8ite?= Date: Mon, 28 Sep 2026 14:51:57 +0200 Subject: [PATCH 2/2] fix: also refuse empty names for execute and deallocate --- .../pgadapter/statements/DeallocateStatement.java | 4 ++++ .../pgadapter/statements/ExecuteStatement.java | 4 ++++ .../pgadapter/EmulatedPsqlMockServerTest.java | 15 ++++++++++++++- .../statements/DeallocateStatementTest.java | 9 +++++++++ .../statements/ExecuteStatementTest.java | 5 +++++ 5 files changed, 36 insertions(+), 1 deletion(-) diff --git a/src/main/java/com/google/cloud/spanner/pgadapter/statements/DeallocateStatement.java b/src/main/java/com/google/cloud/spanner/pgadapter/statements/DeallocateStatement.java index 6ee5a2190a..5438f7763d 100644 --- a/src/main/java/com/google/cloud/spanner/pgadapter/statements/DeallocateStatement.java +++ b/src/main/java/com/google/cloud/spanner/pgadapter/statements/DeallocateStatement.java @@ -122,6 +122,10 @@ static ParsedDeallocateStatement parse(String sql) { "invalid prepared statement name", SQLState.InvalidSqlStatementName); } statementName = unquoteOrFoldIdentifier(name.name); + if (statementName == null || statementName.isEmpty()) { + throw PGExceptionFactory.newPGException( + "zero-length delimited identifier", SQLState.InvalidSqlStatementName); + } } parser.skipWhitespaces(); if (parser.getPos() < parser.getSql().length()) { diff --git a/src/main/java/com/google/cloud/spanner/pgadapter/statements/ExecuteStatement.java b/src/main/java/com/google/cloud/spanner/pgadapter/statements/ExecuteStatement.java index 064a5f1d42..344c85b124 100644 --- a/src/main/java/com/google/cloud/spanner/pgadapter/statements/ExecuteStatement.java +++ b/src/main/java/com/google/cloud/spanner/pgadapter/statements/ExecuteStatement.java @@ -141,6 +141,10 @@ static ParsedExecuteStatement parse(String sql) { "invalid prepared statement name", SQLState.InvalidSqlStatementName); } String statementName = unquoteOrFoldIdentifier(name.name); + if (statementName == null || statementName.isEmpty()) { + throw PGExceptionFactory.newPGException( + "zero-length delimited identifier", SQLState.InvalidSqlStatementName); + } List parameters; if (parser.eatToken("(")) { diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/EmulatedPsqlMockServerTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/EmulatedPsqlMockServerTest.java index 7b0e6230c2..dd383691c3 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/EmulatedPsqlMockServerTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/EmulatedPsqlMockServerTest.java @@ -494,13 +494,26 @@ public void testPrepareCaseInsensitiveAndQuotedIdentifier() throws SQLException } @Test - public void testPrepareZeroLengthDelimitedIdentifier() throws SQLException { + public void testZeroLengthDelimitedIdentifier() throws SQLException { try (Connection connection = DriverManager.getConnection(createUrl("my-db"))) { try (java.sql.Statement statement = connection.createStatement()) { SQLException exception = assertThrows(SQLException.class, () -> statement.execute("prepare \"\" as SELECT 1")); assertEquals(SQLState.InvalidSqlStatementName.toString(), exception.getSQLState()); assertTrue(exception.getMessage().contains("zero-length delimited identifier")); + + exception = assertThrows(SQLException.class, () -> statement.execute("execute \"\"")); + assertEquals(SQLState.InvalidSqlStatementName.toString(), exception.getSQLState()); + assertTrue(exception.getMessage().contains("zero-length delimited identifier")); + + exception = assertThrows(SQLException.class, () -> statement.execute("deallocate \"\"")); + assertEquals(SQLState.InvalidSqlStatementName.toString(), exception.getSQLState()); + assertTrue(exception.getMessage().contains("zero-length delimited identifier")); + + exception = + assertThrows(SQLException.class, () -> statement.execute("deallocate prepare \"\"")); + assertEquals(SQLState.InvalidSqlStatementName.toString(), exception.getSQLState()); + assertTrue(exception.getMessage().contains("zero-length delimited identifier")); } } } diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/statements/DeallocateStatementTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/statements/DeallocateStatementTest.java index d4796dbe8e..92851b74fc 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/statements/DeallocateStatementTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/statements/DeallocateStatementTest.java @@ -26,6 +26,7 @@ import com.google.cloud.spanner.connection.AbstractStatementParser; import com.google.cloud.spanner.pgadapter.ConnectionHandler; import com.google.cloud.spanner.pgadapter.error.PGException; +import com.google.cloud.spanner.pgadapter.error.SQLState; import com.google.cloud.spanner.pgadapter.metadata.ConnectionMetadata; import com.google.cloud.spanner.pgadapter.metadata.OptionsMetadata; import org.junit.Test; @@ -54,6 +55,14 @@ public void testParse() { assertThrows(PGException.class, () -> parse("deallocate prepare")); assertThrows(PGException.class, () -> parse("deallocate foo bar")); assertThrows(PGException.class, () -> parse("deallocate foo.bar")); + + PGException exception = assertThrows(PGException.class, () -> parse("deallocate \"\"")); + assertEquals(SQLState.InvalidSqlStatementName, exception.getSQLState()); + assertEquals("zero-length delimited identifier", exception.getMessage()); + + exception = assertThrows(PGException.class, () -> parse("deallocate prepare \"\"")); + assertEquals(SQLState.InvalidSqlStatementName, exception.getSQLState()); + assertEquals("zero-length delimited identifier", exception.getMessage()); } @Test diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/statements/ExecuteStatementTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/statements/ExecuteStatementTest.java index d07bd8ea21..ca74c2509e 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/statements/ExecuteStatementTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/statements/ExecuteStatementTest.java @@ -27,6 +27,7 @@ import com.google.cloud.spanner.connection.AbstractStatementParser.StatementType; import com.google.cloud.spanner.pgadapter.ConnectionHandler; import com.google.cloud.spanner.pgadapter.error.PGException; +import com.google.cloud.spanner.pgadapter.error.SQLState; import com.google.cloud.spanner.pgadapter.metadata.ConnectionMetadata; import com.google.cloud.spanner.pgadapter.metadata.OptionsMetadata; import java.nio.charset.StandardCharsets; @@ -81,6 +82,10 @@ public void testParse() { assertThrows(PGException.class, () -> parse("execute foo ()")); assertThrows(PGException.class, () -> parse("execute foo (1) bar")); assertThrows(PGException.class, () -> parse("execute foo (1")); + + PGException exception = assertThrows(PGException.class, () -> parse("execute \"\"")); + assertEquals(SQLState.InvalidSqlStatementName, exception.getSQLState()); + assertEquals("zero-length delimited identifier", exception.getMessage()); } static byte[] param(String value) {