diff --git a/src/main/java/com/google/cloud/spanner/pgadapter/metadata/DescribeResult.java b/src/main/java/com/google/cloud/spanner/pgadapter/metadata/DescribeResult.java index be35a44eab..c7e7f806c3 100644 --- a/src/main/java/com/google/cloud/spanner/pgadapter/metadata/DescribeResult.java +++ b/src/main/java/com/google/cloud/spanner/pgadapter/metadata/DescribeResult.java @@ -20,6 +20,7 @@ import com.google.cloud.spanner.pgadapter.error.PGExceptionFactory; import com.google.cloud.spanner.pgadapter.error.SQLState; import com.google.cloud.spanner.pgadapter.parsers.Parser; +import com.google.common.base.Preconditions; import com.google.spanner.v1.StructType; import java.util.Arrays; import javax.annotation.Nullable; @@ -31,6 +32,10 @@ public class DescribeResult { @Nullable private final Type columns; private final int[] parameters; + public static DescribeResult of(int[] givenParameterTypes, @Nullable Type columns) { + return new DescribeResult(givenParameterTypes, columns); + } + public DescribeResult(int[] givenParameterTypes, @Nullable ResultSet resultMetadata) { this( givenParameterTypes, @@ -40,11 +45,15 @@ public DescribeResult(int[] givenParameterTypes, @Nullable ResultSet resultMetad resultMetadata == null ? null : resultMetadata.getType()); } + private DescribeResult(int[] givenParameterTypes, @Nullable Type columns) { + this(givenParameterTypes, null, columns); + } + private DescribeResult( int[] givenParameterTypes, @Nullable StructType undeclaredParameters, @Nullable Type columns) { - this.givenParameterTypes = givenParameterTypes; + this.givenParameterTypes = Preconditions.checkNotNull(givenParameterTypes); this.undeclaredParameters = undeclaredParameters; this.parameters = extractParameters(givenParameterTypes, undeclaredParameters); this.columns = columns; diff --git a/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediatePortalStatement.java b/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediatePortalStatement.java index 8a17c03a3e..c9915a58df 100644 --- a/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediatePortalStatement.java +++ b/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediatePortalStatement.java @@ -65,10 +65,12 @@ public IntermediatePreparedStatement getPreparedStatement() { } public short getParameterFormatCode(int index) { - if (this.parameterFormatCodes.length == 0) { + if (this.parameterFormatCodes == null || this.parameterFormatCodes.length == 0) { return 0; - } else if (index >= this.parameterFormatCodes.length) { + } else if (this.parameterFormatCodes.length == 1) { return this.parameterFormatCodes[0]; + } else if (index < 0 || index >= this.parameterFormatCodes.length) { + return 0; } else { return this.parameterFormatCodes[index]; } @@ -80,6 +82,8 @@ public short getResultFormatCode(int index) { return super.getResultFormatCode(index); } else if (this.resultFormatCodes.length == 1) { return this.resultFormatCodes[0]; + } else if (index < 0 || index >= this.resultFormatCodes.length) { + return super.getResultFormatCode(index); } else { return this.resultFormatCodes[index]; } diff --git a/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediatePreparedStatement.java b/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediatePreparedStatement.java index 733b360ef3..747145d9c7 100644 --- a/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediatePreparedStatement.java +++ b/src/main/java/com/google/cloud/spanner/pgadapter/statements/IntermediatePreparedStatement.java @@ -144,7 +144,7 @@ public Future describeAsync(BackendConnection backendConnection new DescribeResult(this.givenParameterDataTypes, result.getResultSet()); result.getResultSet().close(); } else { - describeResult = new DescribeResult(this.givenParameterDataTypes, null); + describeResult = DescribeResult.of(this.givenParameterDataTypes, null); } return describeResult; }, @@ -158,7 +158,7 @@ public DescribeResult describe() { if (this.describeResult == null) { // Just return a DescribeResult that contains whatever information we were given in the // PARSE message. - return new DescribeResult(this.givenParameterDataTypes, null); + return DescribeResult.of(this.givenParameterDataTypes, null); } try { return this.describeResult.get(); 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 73000ea6b1..697d284425 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/InvalidMessagesTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/InvalidMessagesTest.java @@ -22,8 +22,13 @@ import com.google.cloud.spanner.pgadapter.wireprotocol.SSLMessage; import com.google.cloud.spanner.pgadapter.wireprotocol.StartupMessage; import com.google.cloud.spanner.pgadapter.wireprotocol.WireMessage; +import com.google.common.collect.ImmutableList; +import com.google.protobuf.ListValue; +import com.google.protobuf.Value; import com.google.spanner.v1.CommitRequest; import com.google.spanner.v1.ExecuteSqlRequest; +import com.google.spanner.v1.ResultSetStats; +import com.google.spanner.v1.TypeCode; import io.grpc.Status; import io.opentelemetry.api.OpenTelemetry; import java.io.DataInputStream; @@ -778,4 +783,919 @@ public void testParseMessageDuplicateStatementName() throws IOException { } } } + + @Test + public void testPreparedStatementReturningWithConcurrentSchemaChange() throws IOException { + String sql = "UPDATE t SET v = 'new_val' WHERE id = 1 RETURNING *"; + com.google.spanner.v1.ResultSet sevenColumnsResultSet = + com.google.spanner.v1.ResultSet.newBuilder() + .setMetadata( + createMetadata( + ImmutableList.of( + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64), + ImmutableList.of("c1", "c2", "c3", "c4", "c5", "c6", "c7"))) + .setStats(ResultSetStats.newBuilder().setRowCountExact(1L).build()) + .addRows( + ListValue.newBuilder() + .addValues(Value.newBuilder().setStringValue("1").build()) + .addValues(Value.newBuilder().setStringValue("v2").build()) + .addValues(Value.newBuilder().setStringValue("3").build()) + .addValues(Value.newBuilder().setStringValue("v4").build()) + .addValues(Value.newBuilder().setStringValue("5").build()) + .addValues(Value.newBuilder().setStringValue("v6").build()) + .addValues(Value.newBuilder().setStringValue("7").build()) + .build()) + .build(); + mockSpanner.putStatementResult(StatementResult.query(Statement.of(sql), sevenColumnsResultSet)); + + try (Socket socket = new Socket("localhost", pgServer.getLocalPort())) { + try (DataInputStream inputStream = new DataInputStream(socket.getInputStream()); + DataOutputStream outputStream = new DataOutputStream(socket.getOutputStream())) { + // 1. Startup message. + 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(); + + // Drain startup response until ReadyForQuery ('Z') + while (true) { + byte message = inputStream.readByte(); + int length = inputStream.readInt(); + inputStream.readFully(new byte[length - 4]); + if (message == 'Z') { + break; + } + } + + String statementName = "s1"; + String portalName = "p1"; + byte[] statementNameBytes = statementName.getBytes(StandardCharsets.UTF_8); + byte[] portalNameBytes = portalName.getBytes(StandardCharsets.UTF_8); + byte[] sqlBytes = sql.getBytes(StandardCharsets.UTF_8); + + // 2. PARSE statement s1 + outputStream.writeByte('P'); + outputStream.writeInt(4 + statementNameBytes.length + 1 + sqlBytes.length + 1 + 2); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.write(sqlBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); // 0 parameter types + + // DESCRIBE statement s1 + outputStream.writeByte('D'); + outputStream.writeInt(4 + 1 + statementNameBytes.length + 1); + outputStream.writeByte('S'); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + + // BIND portal p1 with 7 result format codes (matching initial 7 columns) + outputStream.writeByte('B'); + outputStream.writeInt( + 4 + portalNameBytes.length + 1 + statementNameBytes.length + 1 + 2 + 2 + 2 + 7 * 2); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); // 0 parameter format codes + outputStream.writeShort(0); // 0 parameter values + outputStream.writeShort(7); // 7 result format codes + for (int i = 0; i < 7; i++) { + outputStream.writeShort(0); // text format code + } + + // EXECUTE portal p1 + outputStream.writeByte('E'); + outputStream.writeInt(4 + portalNameBytes.length + 1 + 4); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.writeInt(0); // return all rows + + // SYNC + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + // Verify initial response: + // '1' ParseComplete + assertEquals('1', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + + // 't' ParameterDescription + assertEquals('t', inputStream.readByte()); + int paramDescLength = inputStream.readInt(); + inputStream.readFully(new byte[paramDescLength - 4]); + + // 'T' RowDescription (7 columns) + assertEquals('T', inputStream.readByte()); + int rowDescLength = inputStream.readInt(); + short numFields = inputStream.readShort(); + assertEquals(7, numFields); + inputStream.readFully(new byte[rowDescLength - 4 - 2]); + + // '2' BindComplete + assertEquals('2', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + + // 'D' DataRow (7 columns) + assertEquals('D', inputStream.readByte()); + int dataRowLength = inputStream.readInt(); + short numDataCols = inputStream.readShort(); + assertEquals(7, numDataCols); + inputStream.readFully(new byte[dataRowLength - 4 - 2]); + + // 'C' CommandComplete ("UPDATE 1") + assertEquals('C', inputStream.readByte()); + int cmdCompleteLength = inputStream.readInt(); + byte[] cmdCompleteBytes = new byte[cmdCompleteLength - 4]; + inputStream.readFully(cmdCompleteBytes); + assertEquals( + "UPDATE 1", + new String(cmdCompleteBytes, 0, cmdCompleteBytes.length - 1, StandardCharsets.UTF_8)); + + // 'Z' ReadyForQuery + assertEquals('Z', inputStream.readByte()); + assertEquals(5, inputStream.readInt()); + assertEquals('I', inputStream.readByte()); + + // 3. Schema change! Another session added a column 'c8'. + // MockSpanner now returns 8 columns for this SQL: + com.google.spanner.v1.ResultSet eightColumnsResultSet = + com.google.spanner.v1.ResultSet.newBuilder() + .setMetadata( + createMetadata( + ImmutableList.of( + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING), + ImmutableList.of("c1", "c2", "c3", "c4", "c5", "c6", "c7", "c8"))) + .setStats(ResultSetStats.newBuilder().setRowCountExact(1L).build()) + .addRows( + ListValue.newBuilder() + .addValues(Value.newBuilder().setStringValue("1").build()) + .addValues(Value.newBuilder().setStringValue("v2").build()) + .addValues(Value.newBuilder().setStringValue("3").build()) + .addValues(Value.newBuilder().setStringValue("v4").build()) + .addValues(Value.newBuilder().setStringValue("5").build()) + .addValues(Value.newBuilder().setStringValue("v6").build()) + .addValues(Value.newBuilder().setStringValue("7").build()) + .addValues(Value.newBuilder().setStringValue("new_col_val").build()) + .build()) + .build(); + mockSpanner.putStatementResult( + StatementResult.query(Statement.of(sql), eightColumnsResultSet)); + + // 4. Client re-executes prepared statement s1. + // Client still specifies 7 result format codes in Bind because it cached 7 columns. + outputStream.writeByte('B'); + outputStream.writeInt( + 4 + portalNameBytes.length + 1 + statementNameBytes.length + 1 + 2 + 2 + 2 + 7 * 2); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); // 0 parameter format codes + outputStream.writeShort(0); // 0 parameter values + outputStream.writeShort(7); // 7 result format codes (for 8 actual columns!) + for (int i = 0; i < 7; i++) { + outputStream.writeShort(0); + } + + // EXECUTE portal p1 + outputStream.writeByte('E'); + outputStream.writeInt(4 + portalNameBytes.length + 1 + 4); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.writeInt(0); + + // SYNC + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + // 5. Verify re-execution: + // Must succeed without throwing ArrayIndexOutOfBoundsException or returning ErrorResponse! + assertEquals('2', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + + // 'D' DataRow should have 8 columns + assertEquals('D', inputStream.readByte()); + int secondDataRowLength = inputStream.readInt(); + short secondNumDataColumns = inputStream.readShort(); + assertEquals(8, secondNumDataColumns); + for (int col = 0; col < 8; col++) { + int valueLength = inputStream.readInt(); + byte[] valueBytes = new byte[valueLength]; + inputStream.readFully(valueBytes); + if (col == 7) { + assertEquals("new_col_val", new String(valueBytes, StandardCharsets.UTF_8)); + } + } + + // 'C' CommandComplete + assertEquals('C', inputStream.readByte()); + int secondCommandCompleteLength = inputStream.readInt(); + byte[] secondCommandCompleteBytes = new byte[secondCommandCompleteLength - 4]; + inputStream.readFully(secondCommandCompleteBytes); + assertEquals( + "UPDATE 1", + new String( + secondCommandCompleteBytes, + 0, + secondCommandCompleteBytes.length - 1, + StandardCharsets.UTF_8)); + + // 'Z' ReadyForQuery + assertEquals('Z', inputStream.readByte()); + assertEquals(5, inputStream.readInt()); + assertEquals('I', inputStream.readByte()); + + // 6. Describe statement s1 against Spanner now reflects the updated 8 columns. + outputStream.writeByte('D'); + outputStream.writeInt(4 + 1 + statementNameBytes.length + 1); + outputStream.writeByte('S'); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + + // SYNC + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + // 't' ParameterDescription + assertEquals('t', inputStream.readByte()); + int secondParameterDescriptionLength = inputStream.readInt(); + inputStream.readFully(new byte[secondParameterDescriptionLength - 4]); + + // 'T' RowDescription should now describe 8 columns! + assertEquals('T', inputStream.readByte()); + int secondRowDescriptionLength = inputStream.readInt(); + short secondNumFields = inputStream.readShort(); + assertEquals(8, secondNumFields); + inputStream.readFully(new byte[secondRowDescriptionLength - 4 - 2]); + + // 'Z' ReadyForQuery + assertEquals('Z', inputStream.readByte()); + assertEquals(5, inputStream.readInt()); + assertEquals('I', inputStream.readByte()); + } + } + } + + @Test + public void testPreparedStatementReturning_MixedFormatCodes_OutOfBoundsHandledSafely() + throws IOException { + String sql = "UPDATE t2 SET v = 'new_val' WHERE id = 1 RETURNING *"; + com.google.spanner.v1.ResultSet eightColumnsResultSet = + com.google.spanner.v1.ResultSet.newBuilder() + .setMetadata( + createMetadata( + ImmutableList.of( + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING), + ImmutableList.of("c1", "c2", "c3", "c4", "c5", "c6", "c7", "c8"))) + .setStats(ResultSetStats.newBuilder().setRowCountExact(1L).build()) + .addRows( + ListValue.newBuilder() + .addValues(Value.newBuilder().setStringValue("100").build()) + .addValues(Value.newBuilder().setStringValue("v2").build()) + .addValues(Value.newBuilder().setStringValue("300").build()) + .addValues(Value.newBuilder().setStringValue("v4").build()) + .addValues(Value.newBuilder().setStringValue("500").build()) + .addValues(Value.newBuilder().setStringValue("v6").build()) + .addValues(Value.newBuilder().setStringValue("700").build()) + .addValues(Value.newBuilder().setStringValue("v8").build()) + .build()) + .build(); + mockSpanner.putStatementResult(StatementResult.query(Statement.of(sql), eightColumnsResultSet)); + + try (Socket socket = new Socket("localhost", pgServer.getLocalPort())) { + try (DataInputStream inputStream = new DataInputStream(socket.getInputStream()); + DataOutputStream outputStream = new DataOutputStream(socket.getOutputStream())) { + // 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(); + + while (true) { + byte message = inputStream.readByte(); + int length = inputStream.readInt(); + inputStream.readFully(new byte[length - 4]); + if (message == 'Z') { + break; + } + } + + String statementName = "s2"; + String portalName = "p2"; + byte[] statementNameBytes = statementName.getBytes(StandardCharsets.UTF_8); + byte[] portalNameBytes = portalName.getBytes(StandardCharsets.UTF_8); + byte[] sqlBytes = sql.getBytes(StandardCharsets.UTF_8); + + // PARSE + outputStream.writeByte('P'); + outputStream.writeInt(4 + statementNameBytes.length + 1 + sqlBytes.length + 1 + 2); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.write(sqlBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); + + // BIND with 7 format codes: {1, 0, 1, 0, 1, 0, 1} + outputStream.writeByte('B'); + outputStream.writeInt( + 4 + portalNameBytes.length + 1 + statementNameBytes.length + 1 + 2 + 2 + 2 + 7 * 2); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); // param format codes + outputStream.writeShort(0); // param values + outputStream.writeShort(7); // 7 result format codes + outputStream.writeShort(1); // binary for c1 (INT64) + outputStream.writeShort(0); // text for c2 (STRING) + outputStream.writeShort(1); // binary for c3 (INT64) + outputStream.writeShort(0); // text for c4 (STRING) + outputStream.writeShort(1); // binary for c5 (INT64) + outputStream.writeShort(0); // text for c6 (STRING) + outputStream.writeShort(1); // binary for c7 (INT64) + // Note: c8 has no format code specified (index 7 out of bounds of format codes) + + // EXECUTE + outputStream.writeByte('E'); + outputStream.writeInt(4 + portalNameBytes.length + 1 + 4); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.writeInt(0); + + // SYNC + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + // Responses: + // '1' ParseComplete + assertEquals('1', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + + // '2' BindComplete + assertEquals('2', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + + // 'D' DataRow (8 columns) + assertEquals('D', inputStream.readByte()); + int dataRowLength = inputStream.readInt(); + short numDataColumns = inputStream.readShort(); + assertEquals(8, numDataColumns); + + // Col 0 (c1): binary format -> 8 bytes, int64 100L + int column0Length = inputStream.readInt(); + assertEquals(8, column0Length); + assertEquals(100L, inputStream.readLong()); + + // Col 1 (c2): text format -> 2 bytes, "v2" + int column1Length = inputStream.readInt(); + byte[] column1Bytes = new byte[column1Length]; + inputStream.readFully(column1Bytes); + assertEquals("v2", new String(column1Bytes, StandardCharsets.UTF_8)); + + // Col 2 (c3): binary format -> 8 bytes, int64 300L + int column2Length = inputStream.readInt(); + assertEquals(8, column2Length); + assertEquals(300L, inputStream.readLong()); + + // Col 3 (c4): text format -> 2 bytes, "v4" + int column3Length = inputStream.readInt(); + byte[] column3Bytes = new byte[column3Length]; + inputStream.readFully(column3Bytes); + assertEquals("v4", new String(column3Bytes, StandardCharsets.UTF_8)); + + // Col 4 (c5): binary format -> 8 bytes, int64 500L + int column4Length = inputStream.readInt(); + assertEquals(8, column4Length); + assertEquals(500L, inputStream.readLong()); + + // Col 5 (c6): text format -> 2 bytes, "v6" + int column5Length = inputStream.readInt(); + byte[] column5Bytes = new byte[column5Length]; + inputStream.readFully(column5Bytes); + assertEquals("v6", new String(column5Bytes, StandardCharsets.UTF_8)); + + // Col 6 (c7): binary format -> 8 bytes, int64 700L + int column6Length = inputStream.readInt(); + assertEquals(8, column6Length); + assertEquals(700L, inputStream.readLong()); + + // Col 7 (c8): fallback to text format (0) -> 2 bytes, "v8" + int column7Length = inputStream.readInt(); + byte[] column7Bytes = new byte[column7Length]; + inputStream.readFully(column7Bytes); + assertEquals("v8", new String(column7Bytes, StandardCharsets.UTF_8)); + + // 'C' CommandComplete + assertEquals('C', inputStream.readByte()); + int commandCompleteLength = inputStream.readInt(); + byte[] commandCompleteBytes = new byte[commandCompleteLength - 4]; + inputStream.readFully(commandCompleteBytes); + assertEquals( + "UPDATE 1", + new String( + commandCompleteBytes, 0, commandCompleteBytes.length - 1, StandardCharsets.UTF_8)); + + // 'Z' ReadyForQuery + assertEquals('Z', inputStream.readByte()); + assertEquals(5, inputStream.readInt()); + assertEquals('I', inputStream.readByte()); + } + } + } + + @Test + public void testPreparedStatementSelect_SchemaExpansionHandledSafely() throws IOException { + String sql = "SELECT * FROM users_expansion WHERE id = 1"; + com.google.spanner.v1.ResultSet sevenColumnsResultSet = + com.google.spanner.v1.ResultSet.newBuilder() + .setMetadata( + createMetadata( + ImmutableList.of( + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64), + ImmutableList.of("c1", "c2", "c3", "c4", "c5", "c6", "c7"))) + .setStats(ResultSetStats.newBuilder().setRowCountExact(1L).build()) + .addRows( + ListValue.newBuilder() + .addValues(Value.newBuilder().setStringValue("1").build()) + .addValues(Value.newBuilder().setStringValue("v2").build()) + .addValues(Value.newBuilder().setStringValue("3").build()) + .addValues(Value.newBuilder().setStringValue("v4").build()) + .addValues(Value.newBuilder().setStringValue("5").build()) + .addValues(Value.newBuilder().setStringValue("v6").build()) + .addValues(Value.newBuilder().setStringValue("7").build()) + .build()) + .build(); + mockSpanner.putStatementResult(StatementResult.query(Statement.of(sql), sevenColumnsResultSet)); + + try (Socket socket = new Socket("localhost", pgServer.getLocalPort())) { + try (DataInputStream inputStream = new DataInputStream(socket.getInputStream()); + DataOutputStream outputStream = new DataOutputStream(socket.getOutputStream())) { + // 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(); + + while (true) { + byte message = inputStream.readByte(); + int length = inputStream.readInt(); + inputStream.readFully(new byte[length - 4]); + if (message == 'Z') { + break; + } + } + + String statementName = "s_select"; + String portalName = "p_select"; + byte[] statementNameBytes = statementName.getBytes(StandardCharsets.UTF_8); + byte[] portalNameBytes = portalName.getBytes(StandardCharsets.UTF_8); + byte[] sqlBytes = sql.getBytes(StandardCharsets.UTF_8); + + // PARSE + outputStream.writeByte('P'); + outputStream.writeInt(4 + statementNameBytes.length + 1 + sqlBytes.length + 1 + 2); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.write(sqlBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); + + // BIND with 7 format codes (all text) + outputStream.writeByte('B'); + outputStream.writeInt( + 4 + portalNameBytes.length + 1 + statementNameBytes.length + 1 + 2 + 2 + 2 + 7 * 2); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); + outputStream.writeShort(0); + outputStream.writeShort(7); + for (int index = 0; index < 7; index++) { + outputStream.writeShort(0); + } + + // EXECUTE + outputStream.writeByte('E'); + outputStream.writeInt(4 + portalNameBytes.length + 1 + 4); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.writeInt(0); + + // SYNC + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + assertEquals('1', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + assertEquals('2', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + + // 'D' DataRow (7 columns) + assertEquals('D', inputStream.readByte()); + int firstDataRowLength = inputStream.readInt(); + short firstColumnCount = inputStream.readShort(); + assertEquals(7, firstColumnCount); + inputStream.readFully(new byte[firstDataRowLength - 4 - 2]); + + // 'C' CommandComplete ("SELECT 1") + assertEquals('C', inputStream.readByte()); + int firstCommandCompleteLength = inputStream.readInt(); + byte[] firstCommandCompleteBytes = new byte[firstCommandCompleteLength - 4]; + inputStream.readFully(firstCommandCompleteBytes); + assertEquals( + "SELECT 1", + new String( + firstCommandCompleteBytes, + 0, + firstCommandCompleteBytes.length - 1, + StandardCharsets.UTF_8)); + + // 'Z' ReadyForQuery + assertEquals('Z', inputStream.readByte()); + assertEquals(5, inputStream.readInt()); + assertEquals('I', inputStream.readByte()); + + // Concurrent schema change: table now has 8 columns! + com.google.spanner.v1.ResultSet eightColumnsResultSet = + com.google.spanner.v1.ResultSet.newBuilder() + .setMetadata( + createMetadata( + ImmutableList.of( + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING), + ImmutableList.of("c1", "c2", "c3", "c4", "c5", "c6", "c7", "c8"))) + .setStats(ResultSetStats.newBuilder().setRowCountExact(1L).build()) + .addRows( + ListValue.newBuilder() + .addValues(Value.newBuilder().setStringValue("1").build()) + .addValues(Value.newBuilder().setStringValue("v2").build()) + .addValues(Value.newBuilder().setStringValue("3").build()) + .addValues(Value.newBuilder().setStringValue("v4").build()) + .addValues(Value.newBuilder().setStringValue("5").build()) + .addValues(Value.newBuilder().setStringValue("v6").build()) + .addValues(Value.newBuilder().setStringValue("7").build()) + .addValues(Value.newBuilder().setStringValue("v8_added").build()) + .build()) + .build(); + mockSpanner.putStatementResult( + StatementResult.query(Statement.of(sql), eightColumnsResultSet)); + + // Re-execute prepared statement with 7 result format codes (cached from prior describe) + outputStream.writeByte('B'); + outputStream.writeInt( + 4 + portalNameBytes.length + 1 + statementNameBytes.length + 1 + 2 + 2 + 2 + 7 * 2); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); + outputStream.writeShort(0); + outputStream.writeShort(7); + for (int index = 0; index < 7; index++) { + outputStream.writeShort(0); + } + + outputStream.writeByte('E'); + outputStream.writeInt(4 + portalNameBytes.length + 1 + 4); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.writeInt(0); + + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + // Must succeed with 8 columns, 8th column formatted safely as text (format code 0) + assertEquals('2', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + + assertEquals('D', inputStream.readByte()); + int secondDataRowLength = inputStream.readInt(); + short secondColumnCount = inputStream.readShort(); + assertEquals(8, secondColumnCount); + for (int columnIndex = 0; columnIndex < 8; columnIndex++) { + int columnLength = inputStream.readInt(); + byte[] columnBytes = new byte[columnLength]; + inputStream.readFully(columnBytes); + if (columnIndex == 7) { + assertEquals("v8_added", new String(columnBytes, StandardCharsets.UTF_8)); + } + } + + assertEquals('C', inputStream.readByte()); + int secondCommandCompleteLength = inputStream.readInt(); + byte[] secondCommandCompleteBytes = new byte[secondCommandCompleteLength - 4]; + inputStream.readFully(secondCommandCompleteBytes); + assertEquals( + "SELECT 1", + new String( + secondCommandCompleteBytes, + 0, + secondCommandCompleteBytes.length - 1, + StandardCharsets.UTF_8)); + + assertEquals('Z', inputStream.readByte()); + assertEquals(5, inputStream.readInt()); + assertEquals('I', inputStream.readByte()); + } + } + } + + @Test + public void testPreparedStatement_SchemaContractionHandledSafely() throws IOException { + String sql = "SELECT * FROM items_contraction WHERE id = 1"; + com.google.spanner.v1.ResultSet eightColumnsResultSet = + com.google.spanner.v1.ResultSet.newBuilder() + .setMetadata( + createMetadata( + ImmutableList.of( + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING), + ImmutableList.of("c1", "c2", "c3", "c4", "c5", "c6", "c7", "c8"))) + .setStats(ResultSetStats.newBuilder().setRowCountExact(1L).build()) + .addRows( + ListValue.newBuilder() + .addValues(Value.newBuilder().setStringValue("10").build()) + .addValues(Value.newBuilder().setStringValue("v2").build()) + .addValues(Value.newBuilder().setStringValue("30").build()) + .addValues(Value.newBuilder().setStringValue("v4").build()) + .addValues(Value.newBuilder().setStringValue("50").build()) + .addValues(Value.newBuilder().setStringValue("v6").build()) + .addValues(Value.newBuilder().setStringValue("70").build()) + .addValues(Value.newBuilder().setStringValue("v8").build()) + .build()) + .build(); + mockSpanner.putStatementResult(StatementResult.query(Statement.of(sql), eightColumnsResultSet)); + + try (Socket socket = new Socket("localhost", pgServer.getLocalPort())) { + try (DataInputStream inputStream = new DataInputStream(socket.getInputStream()); + DataOutputStream outputStream = new DataOutputStream(socket.getOutputStream())) { + // 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(); + + while (true) { + byte message = inputStream.readByte(); + int length = inputStream.readInt(); + inputStream.readFully(new byte[length - 4]); + if (message == 'Z') { + break; + } + } + + String statementName = "s_contract"; + String portalName = "p_contract"; + byte[] statementNameBytes = statementName.getBytes(StandardCharsets.UTF_8); + byte[] portalNameBytes = portalName.getBytes(StandardCharsets.UTF_8); + byte[] sqlBytes = sql.getBytes(StandardCharsets.UTF_8); + + // PARSE + outputStream.writeByte('P'); + outputStream.writeInt(4 + statementNameBytes.length + 1 + sqlBytes.length + 1 + 2); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.write(sqlBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); + + // BIND with 8 format codes: {1, 0, 1, 0, 1, 0, 1, 0} + outputStream.writeByte('B'); + outputStream.writeInt( + 4 + portalNameBytes.length + 1 + statementNameBytes.length + 1 + 2 + 2 + 2 + 8 * 2); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); + outputStream.writeShort(0); + outputStream.writeShort(8); + outputStream.writeShort(1); + outputStream.writeShort(0); + outputStream.writeShort(1); + outputStream.writeShort(0); + outputStream.writeShort(1); + outputStream.writeShort(0); + outputStream.writeShort(1); + outputStream.writeShort(0); + + // EXECUTE + outputStream.writeByte('E'); + outputStream.writeInt(4 + portalNameBytes.length + 1 + 4); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.writeInt(0); + + // SYNC + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + assertEquals('1', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + assertEquals('2', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + + assertEquals('D', inputStream.readByte()); + int firstDataRowLength = inputStream.readInt(); + short firstColumnCount = inputStream.readShort(); + assertEquals(8, firstColumnCount); + inputStream.readFully(new byte[firstDataRowLength - 4 - 2]); + + assertEquals('C', inputStream.readByte()); + int firstCommandCompleteLength = inputStream.readInt(); + byte[] firstCommandCompleteBytes = new byte[firstCommandCompleteLength - 4]; + inputStream.readFully(firstCommandCompleteBytes); + assertEquals( + "SELECT 1", + new String( + firstCommandCompleteBytes, + 0, + firstCommandCompleteBytes.length - 1, + StandardCharsets.UTF_8)); + + assertEquals('Z', inputStream.readByte()); + assertEquals(5, inputStream.readInt()); + assertEquals('I', inputStream.readByte()); + + // Schema change: Column 8 is dropped! MockSpanner now returns only 7 columns. + com.google.spanner.v1.ResultSet sevenColumnsResultSet = + com.google.spanner.v1.ResultSet.newBuilder() + .setMetadata( + createMetadata( + ImmutableList.of( + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64), + ImmutableList.of("c1", "c2", "c3", "c4", "c5", "c6", "c7"))) + .setStats(ResultSetStats.newBuilder().setRowCountExact(1L).build()) + .addRows( + ListValue.newBuilder() + .addValues(Value.newBuilder().setStringValue("20").build()) + .addValues(Value.newBuilder().setStringValue("v2").build()) + .addValues(Value.newBuilder().setStringValue("40").build()) + .addValues(Value.newBuilder().setStringValue("v4").build()) + .addValues(Value.newBuilder().setStringValue("60").build()) + .addValues(Value.newBuilder().setStringValue("v6").build()) + .addValues(Value.newBuilder().setStringValue("80").build()) + .build()) + .build(); + mockSpanner.putStatementResult( + StatementResult.query(Statement.of(sql), sevenColumnsResultSet)); + + // Re-execute prepared statement with 8 format codes (more format codes than columns in + // result) + outputStream.writeByte('B'); + outputStream.writeInt( + 4 + portalNameBytes.length + 1 + statementNameBytes.length + 1 + 2 + 2 + 2 + 8 * 2); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.write(statementNameBytes); + outputStream.writeByte(0); + outputStream.writeShort(0); + outputStream.writeShort(0); + outputStream.writeShort(8); + outputStream.writeShort(1); + outputStream.writeShort(0); + outputStream.writeShort(1); + outputStream.writeShort(0); + outputStream.writeShort(1); + outputStream.writeShort(0); + outputStream.writeShort(1); + outputStream.writeShort(0); + + outputStream.writeByte('E'); + outputStream.writeInt(4 + portalNameBytes.length + 1 + 4); + outputStream.write(portalNameBytes); + outputStream.writeByte(0); + outputStream.writeInt(0); + + outputStream.writeByte('S'); + outputStream.writeInt(4); + outputStream.flush(); + + // Must succeed with 7 columns without any errors + assertEquals('2', inputStream.readByte()); + assertEquals(4, inputStream.readInt()); + + assertEquals('D', inputStream.readByte()); + int secondDataRowLength = inputStream.readInt(); + short secondColumnCount = inputStream.readShort(); + assertEquals(7, secondColumnCount); + + // Verify format codes for the 7 columns: + // Col 0: binary format (int64 20L) + assertEquals(8, inputStream.readInt()); + assertEquals(20L, inputStream.readLong()); + + // Col 1: text format ("v2") + int col1Length = inputStream.readInt(); + byte[] col1Bytes = new byte[col1Length]; + inputStream.readFully(col1Bytes); + assertEquals("v2", new String(col1Bytes, StandardCharsets.UTF_8)); + + // Col 2: binary format (int64 40L) + assertEquals(8, inputStream.readInt()); + assertEquals(40L, inputStream.readLong()); + + // Col 3: text format ("v4") + int col3Length = inputStream.readInt(); + byte[] col3Bytes = new byte[col3Length]; + inputStream.readFully(col3Bytes); + assertEquals("v4", new String(col3Bytes, StandardCharsets.UTF_8)); + + // Col 4: binary format (int64 60L) + assertEquals(8, inputStream.readInt()); + assertEquals(60L, inputStream.readLong()); + + // Col 5: text format ("v6") + int col5Length = inputStream.readInt(); + byte[] col5Bytes = new byte[col5Length]; + inputStream.readFully(col5Bytes); + assertEquals("v6", new String(col5Bytes, StandardCharsets.UTF_8)); + + // Col 6: binary format (int64 80L) + assertEquals(8, inputStream.readInt()); + assertEquals(80L, inputStream.readLong()); + + assertEquals('C', inputStream.readByte()); + int secondCommandCompleteLength = inputStream.readInt(); + byte[] secondCommandCompleteBytes = new byte[secondCommandCompleteLength - 4]; + inputStream.readFully(secondCommandCompleteBytes); + assertEquals( + "SELECT 1", + new String( + secondCommandCompleteBytes, + 0, + secondCommandCompleteBytes.length - 1, + StandardCharsets.UTF_8)); + + assertEquals('Z', inputStream.readByte()); + assertEquals(5, inputStream.readInt()); + assertEquals('I', inputStream.readByte()); + } + } + } } 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 535d6a5348..14d023dbf0 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/JdbcMockServerTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/JdbcMockServerTest.java @@ -5431,6 +5431,128 @@ public void testDmlReturningMultipleRows() throws SQLException { assertEquals(1, mockSpanner.countRequestsOfType(CommitRequest.class)); } + @Test + public void testPreparedStatementReturningWithConcurrentSchemaChange() throws SQLException { + String sql = "update test_table set value=? where id=? RETURNING *"; + String pgSql = "update test_table set value=$1 where id=$2 RETURNING *"; + + com.google.spanner.v1.ResultSet sevenColumnsResultSet = + com.google.spanner.v1.ResultSet.newBuilder() + .setMetadata( + createMetadata( + ImmutableList.of( + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64), + ImmutableList.of("c1", "c2", "c3", "c4", "c5", "c6", "c7"))) + .setStats(ResultSetStats.newBuilder().setRowCountExact(1L).build()) + .addRows( + ListValue.newBuilder() + .addValues(Value.newBuilder().setStringValue("1").build()) + .addValues(Value.newBuilder().setStringValue("v2").build()) + .addValues(Value.newBuilder().setStringValue("3").build()) + .addValues(Value.newBuilder().setStringValue("v4").build()) + .addValues(Value.newBuilder().setStringValue("5").build()) + .addValues(Value.newBuilder().setStringValue("v6").build()) + .addValues(Value.newBuilder().setStringValue("7").build()) + .build()) + .build(); + mockSpanner.putStatementResult( + StatementResult.query(Statement.of(pgSql), sevenColumnsResultSet)); + mockSpanner.putStatementResult( + StatementResult.query( + Statement.newBuilder(pgSql).bind("p1").to("v1").bind("p2").to(1L).build(), + sevenColumnsResultSet)); + + try (Connection connection = + DriverManager.getConnection( + createUrl() + "&prepareThreshold=1&binaryTransferEnable=int8")) { + try (PreparedStatement preparedStatement = connection.prepareStatement(sql)) { + preparedStatement.setString(1, "v1"); + preparedStatement.setLong(2, 1L); + try (ResultSet resultSet = preparedStatement.executeQuery()) { + assertTrue(resultSet.next()); + assertEquals(1L, resultSet.getLong(1)); + assertEquals("v2", resultSet.getString(2)); + assertFalse(resultSet.next()); + } + + // Schema change: MockSpanner now returns 8 columns for the statement. + com.google.spanner.v1.ResultSet eightColumnsResultSet = + com.google.spanner.v1.ResultSet.newBuilder() + .setMetadata( + createMetadata( + ImmutableList.of( + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING, + TypeCode.INT64, + TypeCode.STRING), + ImmutableList.of("c1", "c2", "c3", "c4", "c5", "c6", "c7", "c8"))) + .setStats(ResultSetStats.newBuilder().setRowCountExact(1L).build()) + .addRows( + ListValue.newBuilder() + .addValues(Value.newBuilder().setStringValue("2").build()) + .addValues(Value.newBuilder().setStringValue("v2").build()) + .addValues(Value.newBuilder().setStringValue("3").build()) + .addValues(Value.newBuilder().setStringValue("v4").build()) + .addValues(Value.newBuilder().setStringValue("5").build()) + .addValues(Value.newBuilder().setStringValue("v6").build()) + .addValues(Value.newBuilder().setStringValue("7").build()) + .addValues(Value.newBuilder().setStringValue("new_col").build()) + .build()) + .build(); + mockSpanner.putStatementResult( + StatementResult.query(Statement.of(pgSql), eightColumnsResultSet)); + mockSpanner.putStatementResult( + StatementResult.query( + Statement.newBuilder(pgSql).bind("p1").to("v2").bind("p2").to(2L).build(), + eightColumnsResultSet)); + + // Re-execute prepared statement with executeQuery + preparedStatement.setString(1, "v2"); + preparedStatement.setLong(2, 2L); + try (ResultSet resultSet = preparedStatement.executeQuery()) { + assertTrue(resultSet.next()); + assertEquals(2L, resultSet.getLong(1)); + assertEquals("v2", resultSet.getString(2)); + assertEquals(3L, resultSet.getLong(3)); + assertEquals("v4", resultSet.getString(4)); + assertEquals(5L, resultSet.getLong(5)); + assertEquals("v6", resultSet.getString(6)); + assertEquals(7L, resultSet.getLong(7)); + // The JDBC driver cached 7 columns on the client side for this specific PreparedStatement + // instance, so getMetaData().getColumnCount() is 7 and accessing column 8 on this old + // PreparedStatement throws a client-side SQLException. + assertEquals(7, resultSet.getMetaData().getColumnCount()); + assertThrows(SQLException.class, () -> resultSet.getString(8)); + assertFalse(resultSet.next()); + } + + // A new prepared statement on the same connection reflects the updated 8 columns because + // prepareStatement describes the statement against Spanner, returning the updated schema. + try (PreparedStatement secondPreparedStatement = connection.prepareStatement(sql)) { + secondPreparedStatement.setString(1, "v2"); + secondPreparedStatement.setLong(2, 2L); + try (ResultSet resultSet = secondPreparedStatement.executeQuery()) { + assertEquals(8, resultSet.getMetaData().getColumnCount()); + assertTrue(resultSet.next()); + assertEquals(2L, resultSet.getLong(1)); + assertEquals("new_col", resultSet.getString(8)); + assertFalse(resultSet.next()); + } + } + } + } + } + @Test public void testUUIDParameter() throws SQLException { String jdbcSql = "SELECT * FROM all_types WHERE col_uuid=?"; diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/metadata/DescribeResultTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/metadata/DescribeResultTest.java index 17ecdd4334..d32abbf214 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/metadata/DescribeResultTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/metadata/DescribeResultTest.java @@ -17,6 +17,7 @@ import static com.google.cloud.spanner.pgadapter.metadata.DescribeResult.extractParameterTypes; import static org.junit.Assert.assertArrayEquals; import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNull; import static org.junit.Assert.assertSame; import static org.junit.Assert.assertThrows; import static org.mockito.Mockito.mock; @@ -116,6 +117,37 @@ public void testExtractParameterTypes() { assertEquals("Invalid parameter name: foo", exception.getMessage()); } + @Test + public void testOf() { + com.google.cloud.spanner.Type columnsType = + com.google.cloud.spanner.Type.struct( + com.google.cloud.spanner.Type.StructField.of( + "c1", com.google.cloud.spanner.Type.int64()), + com.google.cloud.spanner.Type.StructField.of( + "c2", com.google.cloud.spanner.Type.string())); + int[] parameters = new int[] {Oid.INT8, Oid.VARCHAR}; + + DescribeResult describeResult = DescribeResult.of(parameters, columnsType); + assertArrayEquals(parameters, describeResult.getParameters()); + assertEquals(columnsType, describeResult.getColumns()); + + DescribeResult nullColumnsResult = DescribeResult.of(parameters, null); + assertArrayEquals(parameters, nullColumnsResult.getParameters()); + assertNull(nullColumnsResult.getColumns()); + } + + @Test + public void testNullParametersThrows() { + com.google.cloud.spanner.Type columnsType = + com.google.cloud.spanner.Type.struct( + com.google.cloud.spanner.Type.StructField.of( + "c1", com.google.cloud.spanner.Type.int64())); + + assertThrows(NullPointerException.class, () -> DescribeResult.of(null, columnsType)); + assertThrows(NullPointerException.class, () -> DescribeResult.of(null, null)); + assertThrows(NullPointerException.class, () -> new DescribeResult(null, (ResultSet) null)); + } + @Test public void testWithGivenParameterTypes() { ResultSetMetadata metadata = @@ -150,7 +182,7 @@ public void testWithGivenParameterTypes() { DescribeResult updated = result.withGivenParameterTypes(new int[] {Oid.INT8, Oid.UNSPECIFIED}); assertArrayEquals(new int[] {Oid.INT8, Oid.INT8}, updated.getParameters()); - DescribeResult nullMetadataResult = new DescribeResult(initialGivenTypes, null); + DescribeResult nullMetadataResult = new DescribeResult(initialGivenTypes, (ResultSet) null); assertArrayEquals( new int[] {Oid.INT8, Oid.UNSPECIFIED}, nullMetadataResult diff --git a/src/test/java/com/google/cloud/spanner/pgadapter/statements/StatementTest.java b/src/test/java/com/google/cloud/spanner/pgadapter/statements/StatementTest.java index ab9ebbffd6..2a8e431ac0 100644 --- a/src/test/java/com/google/cloud/spanner/pgadapter/statements/StatementTest.java +++ b/src/test/java/com/google/cloud/spanner/pgadapter/statements/StatementTest.java @@ -260,6 +260,175 @@ public void testBasicNoResultStatement() throws Exception { Mockito.verify(resultSet, never()).close(); } + @Test + public void testGetResultFormatCode_NullOrEmpty() { + when(connectionHandler.getConnectionMetadata()).thenReturn(connectionMetadata); + String sql = "SELECT * FROM users"; + IntermediatePreparedStatement preparedStatement = + new IntermediatePreparedStatement( + connectionHandler, options, "", NO_PARAMETER_TYPES, parse(sql), Statement.of(sql)); + + IntermediatePortalStatement nullCodesPortal = + new IntermediatePortalStatement("", preparedStatement, NO_PARAMS, NO_FORMAT_CODES, null); + assertEquals(0, nullCodesPortal.getResultFormatCode(0)); + assertEquals(0, nullCodesPortal.getResultFormatCode(5)); + + IntermediatePortalStatement emptyCodesPortal = + new IntermediatePortalStatement( + "", preparedStatement, NO_PARAMS, NO_FORMAT_CODES, NO_FORMAT_CODES); + assertEquals(0, emptyCodesPortal.getResultFormatCode(0)); + assertEquals(0, emptyCodesPortal.getResultFormatCode(5)); + } + + @Test + public void testGetResultFormatCode_SingleCode() { + when(connectionHandler.getConnectionMetadata()).thenReturn(connectionMetadata); + String sql = "SELECT * FROM users"; + IntermediatePreparedStatement preparedStatement = + new IntermediatePreparedStatement( + connectionHandler, options, "", NO_PARAMETER_TYPES, parse(sql), Statement.of(sql)); + + short[] singleBinaryCode = new short[] {1}; + IntermediatePortalStatement portal = + new IntermediatePortalStatement( + "", preparedStatement, NO_PARAMS, NO_FORMAT_CODES, singleBinaryCode); + + assertEquals(1, portal.getResultFormatCode(0)); + assertEquals(1, portal.getResultFormatCode(1)); + assertEquals(1, portal.getResultFormatCode(7)); + assertEquals(1, portal.getResultFormatCode(100)); + } + + @Test + public void testGetResultFormatCode_MultipleCodes_OutOfBoundsHandledSafely() { + when(connectionHandler.getConnectionMetadata()).thenReturn(connectionMetadata); + String sql = "UPDATE t SET v = $1 WHERE id = $2 RETURNING *"; + IntermediatePreparedStatement preparedStatement = + new IntermediatePreparedStatement( + connectionHandler, + options, + "", + new int[] {Oid.VARCHAR, Oid.INT8}, + parse(sql), + Statement.of(sql)); + + short[] formatCodes = new short[] {0, 1, 0, 1, 0, 1, 0}; // 7 codes + IntermediatePortalStatement portal = + new IntermediatePortalStatement( + "", preparedStatement, NO_PARAMS, NO_FORMAT_CODES, formatCodes); + + // Within bounds + assertEquals(0, portal.getResultFormatCode(0)); + assertEquals(1, portal.getResultFormatCode(1)); + assertEquals(0, portal.getResultFormatCode(2)); + assertEquals(1, portal.getResultFormatCode(3)); + assertEquals(0, portal.getResultFormatCode(4)); + assertEquals(1, portal.getResultFormatCode(5)); + assertEquals(0, portal.getResultFormatCode(6)); + + // Out of bounds: 8th and 9th column (indexes 7 and 8) should safely default to 0 without + // throwing + assertEquals(0, portal.getResultFormatCode(7)); + assertEquals(0, portal.getResultFormatCode(8)); + + // Negative index should also safely return 0 without throwing + assertEquals(0, portal.getResultFormatCode(-1)); + } + + @Test + public void testGetResultFormatCode_ContractingSchema_FewerColumnsThanCodes() { + when(connectionHandler.getConnectionMetadata()).thenReturn(connectionMetadata); + String sql = "SELECT * FROM users"; + IntermediatePreparedStatement preparedStatement = + new IntermediatePreparedStatement( + connectionHandler, options, "", NO_PARAMETER_TYPES, parse(sql), Statement.of(sql)); + + short[] formatCodes = new short[] {0, 1, 0, 1, 0, 1, 0, 1}; // 8 codes + IntermediatePortalStatement portal = + new IntermediatePortalStatement( + "", preparedStatement, NO_PARAMS, NO_FORMAT_CODES, formatCodes); + + // Schema contracted: only first 6 columns needed + for (int index = 0; index < 6; index++) { + assertEquals(formatCodes[index], portal.getResultFormatCode(index)); + } + } + + @Test + public void testGetParameterFormatCode_NullOrEmpty() { + when(connectionHandler.getConnectionMetadata()).thenReturn(connectionMetadata); + String sql = "UPDATE t SET v = $1 WHERE id = $2"; + IntermediatePreparedStatement preparedStatement = + new IntermediatePreparedStatement( + connectionHandler, + options, + "", + new int[] {Oid.VARCHAR, Oid.INT8}, + parse(sql), + Statement.of(sql)); + + IntermediatePortalStatement nullCodesPortal = + new IntermediatePortalStatement("", preparedStatement, NO_PARAMS, null, NO_FORMAT_CODES); + assertEquals(0, nullCodesPortal.getParameterFormatCode(0)); + assertEquals(0, nullCodesPortal.getParameterFormatCode(1)); + assertEquals(0, nullCodesPortal.getParameterFormatCode(-1)); + + IntermediatePortalStatement emptyCodesPortal = + new IntermediatePortalStatement( + "", preparedStatement, NO_PARAMS, NO_FORMAT_CODES, NO_FORMAT_CODES); + assertEquals(0, emptyCodesPortal.getParameterFormatCode(0)); + assertEquals(0, emptyCodesPortal.getParameterFormatCode(1)); + assertEquals(0, emptyCodesPortal.getParameterFormatCode(-1)); + } + + @Test + public void testGetParameterFormatCode_SingleCode() { + when(connectionHandler.getConnectionMetadata()).thenReturn(connectionMetadata); + String sql = "UPDATE t SET v = $1 WHERE id = $2"; + IntermediatePreparedStatement preparedStatement = + new IntermediatePreparedStatement( + connectionHandler, + options, + "", + new int[] {Oid.VARCHAR, Oid.INT8}, + parse(sql), + Statement.of(sql)); + + short[] singleBinaryCode = new short[] {1}; + IntermediatePortalStatement portal = + new IntermediatePortalStatement( + "", preparedStatement, NO_PARAMS, singleBinaryCode, NO_FORMAT_CODES); + + assertEquals(1, portal.getParameterFormatCode(0)); + assertEquals(1, portal.getParameterFormatCode(1)); + assertEquals(1, portal.getParameterFormatCode(5)); + assertEquals(1, portal.getParameterFormatCode(100)); + } + + @Test + public void testGetParameterFormatCode_MultipleCodes_OutOfBoundsHandledSafely() { + when(connectionHandler.getConnectionMetadata()).thenReturn(connectionMetadata); + String sql = "UPDATE t SET v = $1 WHERE id = $2"; + IntermediatePreparedStatement preparedStatement = + new IntermediatePreparedStatement( + connectionHandler, + options, + "", + new int[] {Oid.VARCHAR, Oid.INT8}, + parse(sql), + Statement.of(sql)); + + short[] parameterFormatCodes = new short[] {1, 0}; + IntermediatePortalStatement portal = + new IntermediatePortalStatement( + "", preparedStatement, NO_PARAMS, parameterFormatCodes, NO_FORMAT_CODES); + + assertEquals(1, portal.getParameterFormatCode(0)); + assertEquals(0, portal.getParameterFormatCode(1)); + assertEquals(0, portal.getParameterFormatCode(2)); + assertEquals(0, portal.getParameterFormatCode(-1)); + } + @Test public void testDescribeBasicStatementThrowsException() { when(connectionHandler.getConnectionMetadata()).thenReturn(connectionMetadata);