diff --git a/google-cloud-spanner-cassandra/src/main/java/com/google/cloud/spanner/adapter/DriverConnectionHandler.java b/google-cloud-spanner-cassandra/src/main/java/com/google/cloud/spanner/adapter/DriverConnectionHandler.java index c812fb85..9bc1e50c 100644 --- a/google-cloud-spanner-cassandra/src/main/java/com/google/cloud/spanner/adapter/DriverConnectionHandler.java +++ b/google-cloud-spanner-cassandra/src/main/java/com/google/cloud/spanner/adapter/DriverConnectionHandler.java @@ -71,6 +71,8 @@ final class DriverConnectionHandler implements Runnable { private static final char WRITE_ACTION_QUERY_ID_PREFIX = 'W'; private static final String ROUTE_TO_LEADER_HEADER_KEY = "x-goog-spanner-route-to-leader"; private static final String MAX_COMMIT_DELAY_ATTACHMENT_KEY = "max_commit_delay"; + private static final String KEYSPACE_ATTACHMENT_KEY = "keyspace"; + private static final ByteBufAllocator byteBufAllocator = ByteBufAllocator.DEFAULT; private static final FrameCodec serverFrameCodec = customServerCodec(byteBufAllocator); private static final FrameCodec clientFrameCodec = @@ -353,6 +355,8 @@ private PreparePayloadResult preparePayload(MessageContext ctx) { return prepareBatchMessage((Batch) decodeFrame(ctx.payload).message, ctx.streamId); case ProtocolConstants.Opcode.QUERY: return prepareQueryMessage((Query) decodeFrame(ctx.payload).message); + case ProtocolConstants.Opcode.PREPARE: + return preparePrepareMessage(); default: return new PreparePayloadResult(DEFAULT_CONTEXT); } @@ -396,19 +400,29 @@ private PreparePayloadResult prepareBatchMessage(Batch message, int streamId) { private PreparePayloadResult prepareQueryMessage(Query message) { ApiCallContext context; - Map attachments = Collections.emptyMap(); + Map attachments = new HashMap<>(); if (startsWith(message.query, "SELECT")) { context = DEFAULT_CONTEXT; } else { context = DEFAULT_CONTEXT_WITH_LAR; if (maxCommitDelayMillis.isPresent()) { - attachments = new HashMap<>(); attachments.put(MAX_COMMIT_DELAY_ATTACHMENT_KEY, maxCommitDelayMillis.get()); } } + Optional keyspace = + adapterClientWrapper.getAttachmentsCache().get(KEYSPACE_ATTACHMENT_KEY); + keyspace.ifPresent(v -> attachments.put(KEYSPACE_ATTACHMENT_KEY, v)); return new PreparePayloadResult(context, attachments); } + private PreparePayloadResult preparePrepareMessage() { + Map attachments = new HashMap<>(); + Optional keyspace = + adapterClientWrapper.getAttachmentsCache().get(KEYSPACE_ATTACHMENT_KEY); + keyspace.ifPresent(v -> attachments.put(KEYSPACE_ATTACHMENT_KEY, v)); + return new PreparePayloadResult(DEFAULT_CONTEXT, attachments); + } + private Optional prepareAttachmentForQueryId( int streamId, Map attachments, byte[] queryId) { String key = constructKey(queryId); diff --git a/google-cloud-spanner-cassandra/src/test/java/com/google/cloud/spanner/adapter/DriverConnectionHandlerTest.java b/google-cloud-spanner-cassandra/src/test/java/com/google/cloud/spanner/adapter/DriverConnectionHandlerTest.java index 2a0d0d37..12d5cb2c 100644 --- a/google-cloud-spanner-cassandra/src/test/java/com/google/cloud/spanner/adapter/DriverConnectionHandlerTest.java +++ b/google-cloud-spanner-cassandra/src/test/java/com/google/cloud/spanner/adapter/DriverConnectionHandlerTest.java @@ -82,6 +82,7 @@ public void setUp() throws IOException { mockSocket = mock(Socket.class); outputStream = new ByteArrayOutputStream(); when(mockSocket.getOutputStream()).thenReturn(outputStream); + when(mockAdapterClient.getAttachmentsCache()).thenReturn(new AttachmentsCache(10)); } @Test @@ -319,6 +320,54 @@ public void shortBody_writesErrorMessageToSocket() throws IOException { verify(mockSocket).close(); } + @Test + public void successfulQueryMessageWithKeyspace() throws IOException { + byte[] validPayload = createQueryMessage(); + ByteString grpcResponse = ByteString.copyFromUtf8("gRPC response"); + when(mockSocket.getInputStream()).thenReturn(new ByteArrayInputStream(validPayload)); + when(mockAdapterClient.sendGrpcRequest(any(byte[].class), any(), any(), any(int.class))) + .thenReturn(grpcResponse); + + AttachmentsCache attachmentsCache = new AttachmentsCache(1); + attachmentsCache.put("keyspace", "test_keyspace"); + when(mockAdapterClient.getAttachmentsCache()).thenReturn(attachmentsCache); + + DriverConnectionHandler handler = + new DriverConnectionHandler(mockSocket, mockAdapterClient, mockMetricsRecorder); + handler.run(); + + assertThat(outputStream.toString(StandardCharsets.UTF_8.name())).isEqualTo("gRPC response"); + verify(mockSocket).close(); + verify(mockAdapterClient) + .sendGrpcRequest( + any(), attachmentsCaptor.capture(), contextCaptor.capture(), any(int.class)); + assertThat(attachmentsCaptor.getValue()).containsExactly("keyspace", "test_keyspace"); + } + + @Test + public void successfulPrepareMessageWithKeyspace() throws IOException { + byte[] validPayload = createPrepareMessage(); + ByteString grpcResponse = ByteString.copyFromUtf8("gRPC response"); + when(mockSocket.getInputStream()).thenReturn(new ByteArrayInputStream(validPayload)); + when(mockAdapterClient.sendGrpcRequest(any(byte[].class), any(), any(), any(int.class))) + .thenReturn(grpcResponse); + + AttachmentsCache attachmentsCache = new AttachmentsCache(1); + attachmentsCache.put("keyspace", "test_keyspace"); + when(mockAdapterClient.getAttachmentsCache()).thenReturn(attachmentsCache); + + DriverConnectionHandler handler = + new DriverConnectionHandler(mockSocket, mockAdapterClient, mockMetricsRecorder); + handler.run(); + + assertThat(outputStream.toString(StandardCharsets.UTF_8.name())).isEqualTo("gRPC response"); + verify(mockSocket).close(); + verify(mockAdapterClient) + .sendGrpcRequest( + any(), attachmentsCaptor.capture(), contextCaptor.capture(), any(int.class)); + assertThat(attachmentsCaptor.getValue()).containsExactly("keyspace", "test_keyspace"); + } + private static byte[] createQueryMessage() { return encodeMessage(new Query("SELECT * FROM ks.T")); }