Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down
Original file line number Diff line number Diff line change
@@ -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;
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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";
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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);

}
Original file line number Diff line number Diff line change
Expand Up @@ -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<jbyte*>(&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);
Expand Down Expand Up @@ -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;
Expand All @@ -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);
}
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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.*;
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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());
Expand Down Expand Up @@ -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());
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand Down
Loading
Loading