diff --git a/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/AWSLambda.java b/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/AWSLambda.java index c81433fd..141f2117 100644 --- a/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/AWSLambda.java +++ b/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/AWSLambda.java @@ -327,7 +327,9 @@ private static void startRuntimeLoop(LambdaRequestHandler lambdaRequestHandler, try { ByteArrayOutputStream payload = lambdaRequestHandler.call(request); - runtimeClient.reportInvocationSuccess(request.getId(), payload.toByteArray(), request.getInvocationId()); + // Post straight from the backing array: toByteArray() would copy the whole response. + ResponseBufferViewer view = ResponseBufferViewer.of(payload); + runtimeClient.reportInvocationSuccess(request.getId(), view.array(), view.length(), request.getInvocationId()); // clear interrupted flag in case if it was set by user's code Thread.interrupted(); } catch (Throwable t) { diff --git a/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/ResponseBufferViewer.java b/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/ResponseBufferViewer.java new file mode 100644 index 00000000..f2281843 --- /dev/null +++ b/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/ResponseBufferViewer.java @@ -0,0 +1,62 @@ +/* +Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +SPDX-License-Identifier: Apache-2.0 +*/ + +package com.amazonaws.services.lambda.runtime.api.client; + +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.io.OutputStream; + +/** + * Gets the backing array of a ByteArrayOutputStream without copying it. ByteArrayOutputStream.writeTo(out) is + * specified to call out.write(buf, 0, count) with its own array, under the stream's lock, so this stream keeps that + * reference instead of writing anywhere. The RIC's response buffers are plain ByteArrayOutputStreams + * (EventHandlerLoaderTest checks it); a subclass that overrides writeTo to write differently makes the viewer throw. + * A viewer is only made by of(), serves exactly one writeTo, and is closed afterwards. + */ +final class ResponseBufferViewer extends OutputStream { + private byte[] array; + private int length; + private boolean closed; + + private ResponseBufferViewer() { + } + + /** Runs payload.writeTo on a new viewer, then closes it. */ + static ResponseBufferViewer of(ByteArrayOutputStream payload) throws IOException { + ResponseBufferViewer view = new ResponseBufferViewer(); + payload.writeTo(view); + view.close(); + return view; + } + + @Override + public void write(byte[] b, int off, int len) throws IOException { + if (closed || array != null || off != 0) { + throw new IOException("ResponseBufferViewer takes a single write(buf, 0, count) from writeTo"); + } + array = b; + length = len; + } + + @Override + public void write(int b) throws IOException { + throw new IOException("ResponseBufferViewer takes a single write(buf, 0, count) from writeTo"); + } + + @Override + public void close() { + closed = true; + } + + /** The stream's backing array. Only the first length() bytes are the content. */ + byte[] array() { + return array; + } + + int length() { + return length; + } +} diff --git a/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClient.java b/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClient.java index 042bd257..83002e4b 100644 --- a/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClient.java +++ b/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClient.java @@ -7,6 +7,7 @@ import com.amazonaws.services.lambda.runtime.api.client.logging.LambdaContextLogger; import com.amazonaws.services.lambda.runtime.api.client.runtimeapi.dto.InvocationRequest; import java.io.IOException; +import java.util.Arrays; /** * Java interface for @@ -38,6 +39,17 @@ public interface LambdaRuntimeApiClient { */ void reportInvocationSuccess(String requestId, byte[] response, String invocationId) throws IOException; + /** + * Report invocation success with the first responseLength bytes of response + * @param requestId request id + * @param response byte array whose first responseLength bytes are the response + * @param responseLength number of bytes of response to send + * @param invocationId invocation id for cross-wiring protection (may be null) + */ + default void reportInvocationSuccess(String requestId, byte[] response, int responseLength, String invocationId) throws IOException { + reportInvocationSuccess(requestId, Arrays.copyOf(response, responseLength), invocationId); + } + /** * Report invocation error * @param requestId request id diff --git a/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClientImpl.java b/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClientImpl.java index a87a458e..6684896c 100644 --- a/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClientImpl.java +++ b/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClientImpl.java @@ -130,6 +130,12 @@ public void reportInvocationSuccess(String requestId, byte[] response, String in NativeClient.postInvocationResponse(requestId.getBytes(UTF_8), response, invocationIdBytes); } + @Override + public void reportInvocationSuccess(String requestId, byte[] response, int responseLength, String invocationId) { + byte[] invocationIdBytes = invocationId != null ? invocationId.getBytes(UTF_8) : null; + NativeClient.postInvocationResponseWithLength(requestId.getBytes(UTF_8), response, responseLength, invocationIdBytes); + } + @Override public void reportInvocationError(String requestId, LambdaError error, String invocationId) throws IOException { String endpoint = invocationEndpoint + requestId + "/error"; diff --git a/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/NativeClient.java b/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/NativeClient.java index 5c690814..9bdffde0 100644 --- a/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/NativeClient.java +++ b/aws-lambda-java-runtime-interface-client/src/main/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/NativeClient.java @@ -23,4 +23,7 @@ static void init(String awsLambdaRuntimeApi) { static native void postInvocationResponse(byte[] requestId, byte[] response, byte[] invocationId); + /** Posts the first responseLength bytes of response, so callers can pass a buffer's backing array without copying it. */ + static native void postInvocationResponseWithLength(byte[] requestId, byte[] response, int responseLength, byte[] invocationId); + } diff --git a/aws-lambda-java-runtime-interface-client/src/main/jni/com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient.cpp b/aws-lambda-java-runtime-interface-client/src/main/jni/com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient.cpp index fb6cd3ca..8801f320 100644 --- a/aws-lambda-java-runtime-interface-client/src/main/jni/com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient.cpp +++ b/aws-lambda-java-runtime-interface-client/src/main/jni/com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient.cpp @@ -64,15 +64,19 @@ static void throwLambdaRuntimeClientException(JNIEnv *env, std::string message, env->Throw(lambdaRuntimeException); } -static std::string toNativeString(JNIEnv *env, jbyteArray jArray) { - int length = env->GetArrayLength(jArray); - jbyte* bytes = env->GetByteArrayElements(jArray, NULL); - std::string nativeString = std::string((char *)bytes, length); - env->ReleaseByteArrayElements(jArray, bytes, JNI_ABORT); +// Copies the first length bytes of the array straight into the string. GetByteArrayElements would usually make a +// temporary copy first. A length outside the array raises ArrayIndexOutOfBoundsException; callers check for it. +static std::string toNativeString(JNIEnv *env, jbyteArray jArray, jsize length) { + std::string nativeString(length > 0 ? length : 0, '\0'); + env->GetByteArrayRegion(jArray, 0, length, reinterpret_cast(&nativeString[0])); env->DeleteLocalRef(jArray); return nativeString; } +static std::string toNativeString(JNIEnv *env, jbyteArray jArray) { + return toNativeString(env, jArray, env->GetArrayLength(jArray)); +} + JNIEXPORT void JNICALL Java_com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient_initializeClient(JNIEnv *env, jobject thisObject, jbyteArray userAgent, jbyteArray awsLambdaRuntimeApi) { std::string user_agent = toNativeString(env, userAgent); std::string endpoint = toNativeString(env, awsLambdaRuntimeApi); @@ -129,12 +133,7 @@ JNIEXPORT jobject JNICALL Java_com_amazonaws_services_lambda_runtime_api_client_ return NULL; } -JNIEXPORT void JNICALL Java_com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient_postInvocationResponse - (JNIEnv *env, jobject thisObject, jbyteArray jrequestId, jbyteArray jresponseArray, jbyteArray jinvocationId) { - std::string payload = toNativeString(env, jresponseArray); - if ((env)->ExceptionOccurred()){ - return; - } +static void postInvocationResponse(JNIEnv *env, jbyteArray jrequestId, std::string payload, jbyteArray jinvocationId) { std::string requestId = toNativeString(env, jrequestId); if ((env)->ExceptionOccurred()){ return; @@ -148,10 +147,28 @@ JNIEXPORT void JNICALL Java_com_amazonaws_services_lambda_runtime_api_client_run } } - auto response = aws::lambda_runtime::invocation_response::success(payload, "application/json"); + auto response = aws::lambda_runtime::invocation_response::success(std::move(payload), "application/json"); auto outcome = CLIENT->post_success(requestId, response, invocationId); if (!outcome.is_success()) { std::string errorMessage("Failed to post invocation response."); throwLambdaRuntimeClientException(env, errorMessage, outcome.get_failure()); } } + +JNIEXPORT void JNICALL Java_com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient_postInvocationResponse + (JNIEnv *env, jobject thisObject, jbyteArray jrequestId, jbyteArray jresponseArray, jbyteArray jinvocationId) { + std::string payload = toNativeString(env, jresponseArray); + if ((env)->ExceptionOccurred()){ + return; + } + postInvocationResponse(env, jrequestId, std::move(payload), jinvocationId); +} + +JNIEXPORT void JNICALL Java_com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient_postInvocationResponseWithLength + (JNIEnv *env, jobject thisObject, jbyteArray jrequestId, jbyteArray jresponseArray, jint responseLength, jbyteArray jinvocationId) { + std::string payload = toNativeString(env, jresponseArray, responseLength); + if ((env)->ExceptionOccurred()){ + return; + } + postInvocationResponse(env, jrequestId, std::move(payload), jinvocationId); +} diff --git a/aws-lambda-java-runtime-interface-client/src/main/jni/com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient.h b/aws-lambda-java-runtime-interface-client/src/main/jni/com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient.h index 0f1aaa2c..86087999 100644 --- a/aws-lambda-java-runtime-interface-client/src/main/jni/com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient.h +++ b/aws-lambda-java-runtime-interface-client/src/main/jni/com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient.h @@ -19,6 +19,9 @@ JNIEXPORT jobject JNICALL Java_com_amazonaws_services_lambda_runtime_api_client_ JNIEXPORT void JNICALL Java_com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient_postInvocationResponse (JNIEnv *, jobject, jbyteArray, jbyteArray, jbyteArray); +JNIEXPORT void JNICALL Java_com_amazonaws_services_lambda_runtime_api_client_runtimeapi_NativeClient_postInvocationResponseWithLength + (JNIEnv *, jobject, jbyteArray, jbyteArray, jint, jbyteArray); + #ifdef __cplusplus } #endif diff --git a/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/AWSLambdaTest.java b/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/AWSLambdaTest.java index d747d21f..89104091 100644 --- a/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/AWSLambdaTest.java +++ b/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/AWSLambdaTest.java @@ -8,6 +8,7 @@ import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertTrue; import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyInt; import static org.mockito.ArgumentMatchers.anyString; import static org.mockito.ArgumentMatchers.eq; import static org.mockito.Mockito.*; @@ -242,7 +243,7 @@ void testConcurrentRunWithPlatformThreads() throws Throwable { AWSLambda.startRuntimeLoops(lambdaRequestHandler, lambdaLogger, concurrencyConfig, runtimeClient); // Success Reports Must Equal number of tasks that ran successfully. - verify(runtimeClient, times(7)).reportInvocationSuccess(eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), any()); + verify(runtimeClient, times(7)).reportInvocationSuccess(eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), anyInt(), any()); // Hashmap keys should equal the number of threads (runtime loops). assertEquals(4, SampleHandler.hashMap.size()); // Hashmap total count should equal all tasks that ran * number of iterations per task @@ -284,7 +285,7 @@ void testConcurrentRunWithPlatformThreadsWithFailures() throws Throwable { verify(runtimeClient).reportInvocationError(eq(UserFaultID), any(), any()); // Success Reports Must Equal number of tasks that ran successfully. - verify(runtimeClient, times(2)).reportInvocationSuccess(eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), any()); + verify(runtimeClient, times(2)).reportInvocationSuccess(eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), anyInt(), any()); // Hashmap keys should equal the minumum between(number of threads (runtime loops) AND number of tasks that ran successfully). assertEquals(2, SampleHandler.hashMap.size()); @@ -331,7 +332,7 @@ void testConcurrentModeLoopDoesNotExitExceptForLambdaRuntimeClientMaxRetriesExce verify(runtimeClient).reportInvocationError(eq(IOErrorID), any(), any()); // Success Reports Must Equal number of tasks that ran successfully. - verify(runtimeClient, times(2)).reportInvocationSuccess(eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), any()); + verify(runtimeClient, times(2)).reportInvocationSuccess(eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), anyInt(), any()); // Hashmap keys should equal the minumum between(number of threads (runtime loops) AND number of tasks that ran successfully). assertEquals(1, SampleHandler.hashMap.size()); @@ -520,7 +521,7 @@ void testSequentialWithFatalUserFaultErrorStopsLoop() throws Throwable { verify(runtimeClient).reportInvocationError(eq(UserFaultID), any(), any()); // Success Reports Must Equal number of tasks that ran successfully. And only 2 Error reports for failImmediatelyRequest and userFaultRequest. - verify(runtimeClient, times(2)).reportInvocationSuccess(eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), any()); + verify(runtimeClient, times(2)).reportInvocationSuccess(eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), anyInt(), any()); verify(runtimeClient, times(2)).reportInvocationError(any(), any(), any()); // Hashmap keys should equal one as it is not multithreaded. @@ -566,7 +567,7 @@ void testSequentialWithVirtualMachineErrorStopsLoop() throws Throwable { verify(runtimeClient).reportInvocationError(eq(IOErrorID), any(), any()); // Success Reports Must Equal number of tasks that ran successfully. And only 2 Error reports for failImmediatelyRequest and virtualMachineErrorRequest. - verify(runtimeClient, times(2)).reportInvocationSuccess(eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), any()); + verify(runtimeClient, times(2)).reportInvocationSuccess(eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), anyInt(), any()); verify(runtimeClient, times(2)).reportInvocationError(any(), any(), any()); // Hashmap keys should equal one as it is not multithreaded. @@ -671,7 +672,45 @@ void testInvocationIdIsPassedToReportSuccess() throws Throwable { AWSLambda.startRuntimeLoops(lambdaRequestHandler, lambdaLogger, concurrencyConfig, runtimeClient); verify(runtimeClient).reportInvocationSuccess( - eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), eq("test-inv-uuid-1234")); + eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), any(), anyInt(), eq("test-inv-uuid-1234")); + } + + /** Exposes the backing array so tests can check that it is posted as is. */ + private static final class InspectableByteArrayOutputStream extends ByteArrayOutputStream { + byte[] backingArray() { + return buf; + } + } + + private void runOneInvocation(LambdaRequestHandler handler) throws Throwable { + when(concurrencyConfig.isMultiConcurrent()).thenReturn(false); + + InvocationRequest request = getFakeInvocationRequest(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE); + request.setInvocationId("test-inv-uuid-1234"); + + // Fatal error to stop the loop after one successful invocation + InvocationRequest fatalRequest = mock(InvocationRequest.class); + when(fatalRequest.getId()).thenThrow(UserFault.makeUserFault(new IOError(new Throwable()), true)).thenReturn("fatal"); + + when(runtimeClient.nextInvocation()) + .thenReturn(request) + .thenReturn(fatalRequest); + + AWSLambda.startRuntimeLoops(handler, lambdaLogger, concurrencyConfig, runtimeClient); + } + + @Test + @Timeout(value = 1, unit = TimeUnit.MINUTES) + void testResponseBackingArrayIsPostedWithoutCopy() throws Throwable { + InspectableByteArrayOutputStream output = new InspectableByteArrayOutputStream(); + output.write("\"success\"".getBytes()); + + runOneInvocation(request -> output); + + // The backing array itself is passed, with the content length, instead of a toByteArray() copy. + verify(runtimeClient).reportInvocationSuccess( + eq(SampleHandler.ADD_ENTRY_TO_MAP_ID_OP_MODE), same(output.backingArray()), eq(9), eq("test-inv-uuid-1234")); + verify(runtimeClient, never()).reportInvocationSuccess(anyString(), any(), any()); } @Test diff --git a/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/EventHandlerLoaderTest.java b/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/EventHandlerLoaderTest.java index aae2f1af..89a315ae 100644 --- a/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/EventHandlerLoaderTest.java +++ b/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/EventHandlerLoaderTest.java @@ -72,6 +72,9 @@ private static void assertSuccessfulInvocation(LambdaRequestHandler lambdaReques String result = resultBytes.toString(); assertEquals("\"success\"", result); + // AWSLambda posts this buffer through ResponseBufferViewer, which relies on the JDK's own writeTo. If this + // fails, the buffer became a subclass: check that its writeTo still makes a single write(buf, 0, count). + assertEquals(ByteArrayOutputStream.class, resultBytes.getClass()); } private static InvocationRequest getTestInvocationRequest() { diff --git a/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/ResponseBufferViewerTest.java b/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/ResponseBufferViewerTest.java new file mode 100644 index 00000000..ffd0a60b --- /dev/null +++ b/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/ResponseBufferViewerTest.java @@ -0,0 +1,117 @@ +/* +Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +SPDX-License-Identifier: Apache-2.0 +*/ + +package com.amazonaws.services.lambda.runtime.api.client; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.io.OutputStream; +import java.nio.charset.StandardCharsets; +import java.util.Arrays; +import org.junit.jupiter.api.Test; + +class ResponseBufferViewerTest { + + /** Exposes the backing array so the test can check that the view holds the same instance. */ + private static final class InspectableByteArrayOutputStream extends ByteArrayOutputStream { + InspectableByteArrayOutputStream(int size) { + super(size); + } + + byte[] backingArray() { + return buf; + } + } + + @Test + void capturesTheBackingArrayWithoutCopying() throws IOException { + InspectableByteArrayOutputStream payload = new InspectableByteArrayOutputStream(64); + byte[] content = "response".getBytes(StandardCharsets.UTF_8); + payload.write(content); + + ResponseBufferViewer view = ResponseBufferViewer.of(payload); + + assertSame(payload.backingArray(), view.array()); + assertEquals(content.length, view.length()); + assertArrayEquals(content, Arrays.copyOf(view.array(), view.length())); + } + + @Test + void capturesAnEmptyStream() throws IOException { + InspectableByteArrayOutputStream payload = new InspectableByteArrayOutputStream(16); + + ResponseBufferViewer view = ResponseBufferViewer.of(payload); + + assertSame(payload.backingArray(), view.array()); + assertEquals(0, view.length()); + } + + @Test + void capturesTheCurrentArrayAfterTheStreamGrew() throws IOException { + InspectableByteArrayOutputStream payload = new InspectableByteArrayOutputStream(4); + payload.write(new byte[1000]); + + ResponseBufferViewer view = ResponseBufferViewer.of(payload); + + assertSame(payload.backingArray(), view.array()); + assertEquals(1000, view.length()); + } + + @Test + void throwsWhenWriteToWritesSingleBytes() throws IOException { + ByteArrayOutputStream payload = new ByteArrayOutputStream() { + @Override + public synchronized void writeTo(OutputStream out) throws IOException { + out.write('a'); + } + }; + + assertThrows(IOException.class, () -> ResponseBufferViewer.of(payload)); + } + + @Test + void throwsWhenWriteToWritesMoreThanOneArray() throws IOException { + ByteArrayOutputStream payload = new ByteArrayOutputStream() { + @Override + public synchronized void writeTo(OutputStream out) throws IOException { + out.write(new byte[] {'a'}, 0, 1); + out.write(new byte[] {'b'}, 0, 1); + } + }; + + assertThrows(IOException.class, () -> ResponseBufferViewer.of(payload)); + } + + @Test + void throwsWhenWriteToWritesFromAnOffset() throws IOException { + ByteArrayOutputStream payload = new ByteArrayOutputStream() { + @Override + public synchronized void writeTo(OutputStream out) throws IOException { + out.write(new byte[] {'a', 'b'}, 1, 1); + } + }; + + assertThrows(IOException.class, () -> ResponseBufferViewer.of(payload)); + } + + @Test + void isClosedAfterItsWriteTo() throws IOException { + InspectableByteArrayOutputStream payload = new InspectableByteArrayOutputStream(16); + payload.write("response".getBytes(StandardCharsets.UTF_8)); + ResponseBufferViewer view = ResponseBufferViewer.of(payload); + + assertThrows(IOException.class, () -> payload.writeTo(view)); + assertThrows(IOException.class, () -> view.write('a')); + + // The rejected writes did not replace what the viewer holds. + assertSame(payload.backingArray(), view.array()); + assertEquals(8, view.length()); + } +} diff --git a/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClientImplTest.java b/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClientImplTest.java index 43334e86..346f0753 100644 --- a/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClientImplTest.java +++ b/aws-lambda-java-runtime-interface-client/src/test/java/com/amazonaws/services/lambda/runtime/api/client/runtimeapi/LambdaRuntimeApiClientImplTest.java @@ -13,6 +13,7 @@ import org.junit.jupiter.api.parallel.Resources; import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.junit.jupiter.api.Assertions.assertArrayEquals; import static org.junit.jupiter.api.Assertions.assertNull; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertThrows; @@ -28,6 +29,7 @@ import java.net.InetAddress; import java.net.NetworkInterface; import java.util.ArrayList; +import java.util.Arrays; import java.util.Enumeration; import java.util.List; import java.util.function.Function; @@ -38,6 +40,7 @@ import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.anyString; import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.CALLS_REAL_METHODS; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.times; import static org.mockito.Mockito.verify; @@ -360,6 +363,93 @@ public void reportInvocationSuccessTest() { } } + @Test + public void reportInvocationSuccessEmptyBodyTest() throws Exception { + assertArrayEquals(new byte[0], postSuccessAndGetBody(new byte[0], null).getBody().readByteArray()); + } + + @Test + public void reportInvocationSuccessLargeBinaryBodyTest() throws Exception { + // 6 MB covering every byte value, including zero bytes that a C string would stop at. + byte[] response = new byte[6 * 1024 * 1024]; + for (int i = 0; i < response.length; i++) { + response[i] = (byte) i; + } + assertArrayEquals(response, postSuccessAndGetBody(response, null).getBody().readByteArray()); + } + + @Test + public void reportInvocationSuccessWithInvocationIdTest() throws Exception { + String invocationId = "test-invocation-uuid-1234"; + RecordedRequest recordedRequest = postSuccessAndGetBody("{\"msg\":\"test\"}".getBytes(), invocationId); + assertEquals(invocationId, recordedRequest.getHeader("Lambda-Runtime-Invocation-Id")); + assertEquals("{\"msg\":\"test\"}", recordedRequest.getBody().readUtf8()); + } + + @Test + public void reportInvocationSuccessWithLengthSendsOnlyLengthBytesTest() throws Exception { + // A 6 MB response in an 8 MB backing array, as after the buffer has doubled; the tail must not be sent. + int length = 6 * 1024 * 1024; + byte[] buffer = new byte[8 * 1024 * 1024]; + for (int i = 0; i < buffer.length; i++) { + buffer[i] = (byte) i; + } + RecordedRequest recordedRequest = postSuccessWithLengthAndGetRequest(buffer, length, null); + assertArrayEquals(Arrays.copyOf(buffer, length), recordedRequest.getBody().readByteArray()); + } + + @Test + public void reportInvocationSuccessWithZeroLengthTest() throws Exception { + RecordedRequest recordedRequest = postSuccessWithLengthAndGetRequest("stale".getBytes(), 0, null); + assertArrayEquals(new byte[0], recordedRequest.getBody().readByteArray()); + } + + @Test + public void reportInvocationSuccessWithLengthAndInvocationIdTest() throws Exception { + String invocationId = "test-invocation-uuid-1234"; + RecordedRequest recordedRequest = postSuccessWithLengthAndGetRequest("{\"msg\":\"test\"}stale".getBytes(), 14, invocationId); + assertEquals(invocationId, recordedRequest.getHeader("Lambda-Runtime-Invocation-Id")); + assertEquals("{\"msg\":\"test\"}", recordedRequest.getBody().readUtf8()); + } + + @Test + public void reportInvocationSuccessWithLengthOutOfBoundsTest() { + assertThrows(ArrayIndexOutOfBoundsException.class, + () -> lambdaRuntimeApiClientImpl.reportInvocationSuccess(requestId, new byte[4], 5, null)); + assertThrows(ArrayIndexOutOfBoundsException.class, + () -> lambdaRuntimeApiClientImpl.reportInvocationSuccess(requestId, new byte[4], -1, null)); + assertEquals(0, mockWebServer.getRequestCount()); + } + + @Test + public void reportInvocationSuccessWithLengthDefaultMethodCopiesTest() throws Exception { + LambdaRuntimeApiClient client = mock(LambdaRuntimeApiClient.class, CALLS_REAL_METHODS); + client.reportInvocationSuccess(requestId, "{\"msg\":\"test\"}stale".getBytes(), 14, "id"); + verify(client).reportInvocationSuccess(requestId, "{\"msg\":\"test\"}".getBytes(), "id"); + } + + private RecordedRequest postSuccessWithLengthAndGetRequest(byte[] buffer, int length, String invocationId) throws Exception { + MockResponse mockResponse = new MockResponse(); + mockResponse.setResponseCode(HTTP_ACCEPTED); + mockWebServer.enqueue(mockResponse); + + lambdaRuntimeApiClientImpl.reportInvocationSuccess(requestId, buffer, length, invocationId); + RecordedRequest recordedRequest = mockWebServer.takeRequest(); + assertEquals("/2018-06-01/runtime/invocation/1234/response", recordedRequest.getPath()); + return recordedRequest; + } + + private RecordedRequest postSuccessAndGetBody(byte[] response, String invocationId) throws Exception { + MockResponse mockResponse = new MockResponse(); + mockResponse.setResponseCode(HTTP_ACCEPTED); + mockWebServer.enqueue(mockResponse); + + lambdaRuntimeApiClientImpl.reportInvocationSuccess(requestId, response, invocationId); + RecordedRequest recordedRequest = mockWebServer.takeRequest(); + assertEquals("/2018-06-01/runtime/invocation/1234/response", recordedRequest.getPath()); + return recordedRequest; + } + @Test public void restoreNextTest() { try {