Skip to content
Open
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 @@ -69,7 +69,7 @@ func (h *HTTPHandler) invoke(w http.ResponseWriter, r *http.Request) {

metrics := invoke.NewInvokeMetrics(nil, &noOpCounter{})
metrics.AttachInvokeRequest(invokeReq)
if err, responseSent := h.app.Invoke(ctx, invokeReq, metrics); err != nil {
if err, responseSent, _ := h.app.Invoke(ctx, invokeReq, metrics, w); err != nil {
logging.Err(ctx, "invoke failed", err)
if !responseSent {
h.respondWithError(w, err)
Expand All @@ -93,7 +93,7 @@ func (h *HTTPHandler) respondWithError(w http.ResponseWriter, err rapidmodel.App

type raptorApp interface {
Init(ctx context.Context, req *intmodel.InitRequestMessage, metrics interop.InitMetrics) rapidmodel.AppError
Invoke(ctx context.Context, msg interop.InvokeRequest, metrics interop.InvokeMetrics) (err rapidmodel.AppError, responseSent bool)
Invoke(ctx context.Context, msg interop.InvokeRequest, metrics interop.InvokeMetrics, responseWriter http.ResponseWriter) (err rapidmodel.AppError, responseSent bool, invokePending bool)
}

type noOpCounter struct{}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,15 +16,15 @@ import (
)

type Responder struct {
invokeReq interop.InvokeRequest
body []byte
rw http.ResponseWriter
invokeReq interop.InvokeRequest
body []byte
responseWriter http.ResponseWriter
}

func NewResponder(invokeReq interop.InvokeRequest) *Responder {
func NewResponder(invokeReq interop.InvokeRequest, responseWriter http.ResponseWriter) *Responder {
return &Responder{
invokeReq: invokeReq,
rw: invokeReq.ResponseWriter(),
invokeReq: invokeReq,
responseWriter: responseWriter,
}
}

Expand Down Expand Up @@ -60,9 +60,9 @@ func (s *Responder) SendRuntimeResponseTrailers(request invoke.RuntimeResponseRe
s.SendErrorTrailers(trailerError, "")
return
}
s.rw.Header().Set(invoke.СontentTypeHeader, request.ContentType())
s.rw.Header().Set(invoke.RuntimeResponseModeHeader, request.ResponseMode())
if _, err := s.rw.Write(s.body); err != nil {
s.responseWriter.Header().Set(invoke.СontentTypeHeader, request.ContentType())
s.responseWriter.Header().Set(invoke.RuntimeResponseModeHeader, request.ResponseMode())
if _, err := s.responseWriter.Write(s.body); err != nil {
slog.Error("could not write invoke response", "err", err)
}
}
Expand All @@ -72,10 +72,10 @@ func (s *Responder) SendError(err invoke.ErrorForInvoker, _ interop.InitStaticDa
}

func (s *Responder) SendErrorTrailers(err invoke.ErrorForInvoker, _ invoke.InvokeBodyResponseStatus) {
s.rw.Header().Set("Error-Type", err.ErrorType().String())
s.responseWriter.Header().Set("Error-Type", err.ErrorType().String())

s.rw.WriteHeader(err.ReturnCode())
if _, err := s.rw.Write([]byte(err.ErrorDetails())); err != nil {
s.responseWriter.WriteHeader(err.ReturnCode())
if _, err := s.responseWriter.Write([]byte(err.ErrorDetails())); err != nil {
slog.Error("could not write invoke error response", "err", err)
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,6 @@ func TestResponder_CompleteSuccessFlow(t *testing.T) {
mockRuntimeReq := invoke.NewMockRuntimeResponseRequest(t)
mockInitData := interop.NewMockInitStaticDataProvider(t)

mockInvokeReq.On("ResponseWriter").Return(recorder)
mockInvokeReq.On("MaxPayloadSize").Return(int64(1024))

responseBody := "test response body"
Expand All @@ -34,7 +33,7 @@ func TestResponder_CompleteSuccessFlow(t *testing.T) {
mockRuntimeReq.On("ResponseMode").Return("buffered")
mockRuntimeReq.On("TrailerError").Return(nil)

responder := NewResponder(mockInvokeReq)
responder := NewResponder(mockInvokeReq, recorder)
responder.SendRuntimeResponseHeaders(mockInitData, "", "")
result := responder.SendRuntimeResponseBody(context.Background(), mockRuntimeReq, 0)
assert.NoError(t, result.Err)
Expand All @@ -55,12 +54,10 @@ func TestResponder_SendErrorFlow(t *testing.T) {
mockInvokeReq := interop.NewMockInvokeRequest(t)
mockInitData := interop.NewMockInitStaticDataProvider(t)

mockInvokeReq.On("ResponseWriter").Return(recorder)

baseErr := io.ErrUnexpectedEOF
appError := model.NewCustomerError("Function.TestError", model.WithCause(baseErr), model.WithErrorMessage("test error"))

responder := NewResponder(mockInvokeReq)
responder := NewResponder(mockInvokeReq, recorder)
responder.SendError(appError, mockInitData)

assert.Equal(t, "Function.TestError", recorder.Header().Get("Error-Type"))
Expand All @@ -76,9 +73,7 @@ func TestResponder_RuntimeInvocationErrorFlow(t *testing.T) {
mockInvokeReq := interop.NewMockInvokeRequest(t)
mockInitData := interop.NewMockInitStaticDataProvider(t)

mockInvokeReq.On("ResponseWriter").Return(recorder)

responder := NewResponder(mockInvokeReq)
responder := NewResponder(mockInvokeReq, recorder)
responder.SendRuntimeResponseHeaders(mockInitData, "", "")
responder.SendErrorTrailers(model.NewCustomerError("Runtime.TestError", model.WithErrorMessage("trailer error")), "")

Expand All @@ -96,13 +91,12 @@ func TestResponder_ErrorInTheMiddleOfResponse(t *testing.T) {
mockRuntimeReq := invoke.NewMockRuntimeResponseRequest(t)
mockInitData := interop.NewMockInitStaticDataProvider(t)

mockInvokeReq.On("ResponseWriter").Return(recorder)
mockInvokeReq.On("MaxPayloadSize").Return(int64(1024))

responseBody := "test response body"
mockRuntimeReq.On("BodyReader").Return(strings.NewReader(responseBody))

responder := NewResponder(mockInvokeReq)
responder := NewResponder(mockInvokeReq, recorder)
responder.SendRuntimeResponseHeaders(mockInitData, "", "")
result := responder.SendRuntimeResponseBody(context.Background(), mockRuntimeReq, 0)
assert.NoError(t, result.Err)
Expand All @@ -123,7 +117,6 @@ func TestResponder_RuntimeResponseTrailerError(t *testing.T) {
mockRuntimeReq := invoke.NewMockRuntimeResponseRequest(t)
mockInitData := interop.NewMockInitStaticDataProvider(t)

mockInvokeReq.On("ResponseWriter").Return(recorder)
mockInvokeReq.On("MaxPayloadSize").Return(int64(1024))

errorType := model.ErrorType("Function.TrailerError")
Expand All @@ -138,7 +131,7 @@ func TestResponder_RuntimeResponseTrailerError(t *testing.T) {
mockRuntimeReq.On("BodyReader").Return(strings.NewReader(responseBody))
mockRuntimeReq.On("TrailerError").Return(trailerError)

responder := NewResponder(mockInvokeReq)
responder := NewResponder(mockInvokeReq, recorder)
responder.SendRuntimeResponseHeaders(mockInitData, "", "")
result := responder.SendRuntimeResponseBody(context.Background(), mockRuntimeReq, 0)
assert.NoError(t, result.Err)
Expand Down Expand Up @@ -188,10 +181,9 @@ func TestResponder_SendRuntimeResponseBody(t *testing.T) {
mockInvokeReq := interop.NewMockInvokeRequest(t)
mockRuntimeReq := invoke.NewMockRuntimeResponseRequest(t)

mockInvokeReq.On("ResponseWriter").Return(recorder)
tt.setupMocks(mockInvokeReq, mockRuntimeReq)

responder := NewResponder(mockInvokeReq)
responder := NewResponder(mockInvokeReq, recorder)
result := responder.SendRuntimeResponseBody(context.Background(), mockRuntimeReq, 0)

if tt.expectError {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,8 @@ type rieInvokeRequest struct {
responseMode string
internalInvocationID string

functionVersionID string
functionVersionID string
resolvedInvokeTimeout time.Duration
}

func NewRieInvokeRequest(request *http.Request, writer http.ResponseWriter) (*rieInvokeRequest, model.AppError) {
Expand Down Expand Up @@ -149,22 +150,6 @@ func (r *rieInvokeRequest) BodyReader() io.Reader {
return r.request.Body
}

func (r *rieInvokeRequest) ResponseWriter() http.ResponseWriter {
return r.writer
}

func (r *rieInvokeRequest) SetResponseHeader(key string, val string) {
r.writer.Header().Set(key, val)
}

func (r *rieInvokeRequest) AddResponseHeader(key string, val string) {
r.writer.Header().Add(key, val)
}

func (r *rieInvokeRequest) WriteResponseHeaders(status int) {
r.writer.WriteHeader(status)
}

func (r *rieInvokeRequest) ResponseMode() string {
return r.responseMode
}
Expand All @@ -174,7 +159,8 @@ func (r *rieInvokeRequest) UpdateFromInitData(initData interop.InitStaticDataPro
return model.NewClientError(errors.New("sandbox is not initialized"), model.ErrorSeverityError, model.ErrorInitIncomplete)
}

r.deadline = time.Now().Add(time.Duration(initData.FunctionTimeout()) * time.Millisecond)
r.resolvedInvokeTimeout = initData.FunctionTimeout()
r.deadline = time.Now().Add(r.resolvedInvokeTimeout)

if r.functionVersionID != initData.FunctionVersionID() {
return model.NewClientError(nil, model.ErrorSeverityInvalid, model.ErrorInvalidFunctionVersion)
Expand All @@ -187,6 +173,14 @@ func (r *rieInvokeRequest) FunctionVersionID() string {
return r.functionVersionID
}

func (r *rieInvokeRequest) InternalInvocationID() string {
return r.internalInvocationID
func (r *rieInvokeRequest) ResolvedFunctionTimeoutMs() int64 {
return 0
}

func (r *rieInvokeRequest) ResolvedInvokeTimeout() time.Duration {
return r.resolvedInvokeTimeout
}

func (r *rieInvokeRequest) LongPollingConfig() *interop.LongPollingConfig { return nil }

func (r *rieInvokeRequest) InternalInvocationID() string { return r.internalInvocationID }
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ package internal

import (
context "context"
http "net/http"

mock "github.com/stretchr/testify/mock"
interop "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/interop"
Expand Down Expand Up @@ -37,33 +38,40 @@ func (_m *mockRaptorApp) Init(ctx context.Context, req *model.InitRequestMessage
return r0
}

func (_m *mockRaptorApp) Invoke(ctx context.Context, msg interop.InvokeRequest, metrics interop.InvokeMetrics) (rapidmodel.AppError, bool) {
ret := _m.Called(ctx, msg, metrics)
func (_m *mockRaptorApp) Invoke(ctx context.Context, msg interop.InvokeRequest, metrics interop.InvokeMetrics, responseWriter http.ResponseWriter) (rapidmodel.AppError, bool, bool) {
ret := _m.Called(ctx, msg, metrics, responseWriter)

if len(ret) == 0 {
panic("no return value specified for Invoke")
}

var r0 rapidmodel.AppError
var r1 bool
if rf, ok := ret.Get(0).(func(context.Context, interop.InvokeRequest, interop.InvokeMetrics) (rapidmodel.AppError, bool)); ok {
return rf(ctx, msg, metrics)
var r2 bool
if rf, ok := ret.Get(0).(func(context.Context, interop.InvokeRequest, interop.InvokeMetrics, http.ResponseWriter) (rapidmodel.AppError, bool, bool)); ok {
return rf(ctx, msg, metrics, responseWriter)
}
if rf, ok := ret.Get(0).(func(context.Context, interop.InvokeRequest, interop.InvokeMetrics) rapidmodel.AppError); ok {
r0 = rf(ctx, msg, metrics)
if rf, ok := ret.Get(0).(func(context.Context, interop.InvokeRequest, interop.InvokeMetrics, http.ResponseWriter) rapidmodel.AppError); ok {
r0 = rf(ctx, msg, metrics, responseWriter)
} else {
if ret.Get(0) != nil {
r0 = ret.Get(0).(rapidmodel.AppError)
}
}

if rf, ok := ret.Get(1).(func(context.Context, interop.InvokeRequest, interop.InvokeMetrics) bool); ok {
r1 = rf(ctx, msg, metrics)
if rf, ok := ret.Get(1).(func(context.Context, interop.InvokeRequest, interop.InvokeMetrics, http.ResponseWriter) bool); ok {
r1 = rf(ctx, msg, metrics, responseWriter)
} else {
r1 = ret.Get(1).(bool)
}

return r0, r1
if rf, ok := ret.Get(2).(func(context.Context, interop.InvokeRequest, interop.InvokeMetrics, http.ResponseWriter) bool); ok {
r2 = rf(ctx, msg, metrics, responseWriter)
} else {
r2 = ret.Get(2).(bool)
}

return r0, r1, r2
}

func newMockRaptorApp(t interface {
Expand Down
11 changes: 6 additions & 5 deletions internal/lambda-managed-instances/aws-lambda-rie/internal/run.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import (
"context"
"fmt"
"log/slog"
"net/http"
"os"
"time"

Expand Down Expand Up @@ -47,10 +48,10 @@ func Run(supv supvmodel.ProcessSupervisor, args []string, fileUtil utils.FileUti
telemetryAPIRelay := telemetry.NewRelay()
eventsAPI := telemetry.NewEventsAPI(telemetryAPIRelay)

responderFactoryFunc := func(_ context.Context, invokeReq interop.InvokeRequest) invoke.InvokeResponseSender {
return rieinvoke.NewResponder(invokeReq)
}
invokeRouter := invoke.NewInvokeRouter(rapid.RuntimePoolSize, eventsAPI, responderFactoryFunc, timeout.NewRecentCache())
invokeRouter := invoke.NewInvokeRouter(rapid.RuntimePoolSize, eventsAPI, timeout.NewRecentCache())
longPollRouter := invoke.NewLongInvokerRouter(invokeRouter, func(_ context.Context, invokeReq interop.InvokeRequest, responseWriter http.ResponseWriter) invoke.InvokeResponseSender {
return rieinvoke.NewResponder(invokeReq, responseWriter)
})

metadataToken := uuid.NewString()
deps := rapid.Dependencies{
Expand All @@ -60,7 +61,7 @@ func Run(supv supvmodel.ProcessSupervisor, args []string, fileUtil utils.FileUti
Supervisor: supv,
RuntimeAPIAddrPort: runtimeAPIAddr,
FileUtils: fileUtil,
InvokeRouter: invokeRouter,
InvokeRouter: longPollRouter,
MetadataService: lmds.NewService(metadataToken),
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -201,6 +201,23 @@ func (_m *MockInitStaticDataProvider) MemorySizeMB() uint64 {
return r0
}

func (_m *MockInitStaticDataProvider) RuntimeRelease() string {
ret := _m.Called()

if len(ret) == 0 {
panic("no return value specified for RuntimeRelease")
}

var r0 string
if rf, ok := ret.Get(0).(func() string); ok {
r0 = rf()
} else {
r0 = ret.Get(0).(string)
}

return r0
}

func (_m *MockInitStaticDataProvider) RuntimeVersion() string {
ret := _m.Called()

Expand Down
31 changes: 30 additions & 1 deletion internal/lambda-managed-instances/interop/mock_invoke_metrics.go
Original file line number Diff line number Diff line change
Expand Up @@ -74,14 +74,43 @@ func (_m *MockInvokeMetrics) SendMetrics(_a0 model.AppError) error {
return r0
}

func (_m *MockInvokeMetrics) SetInvokeMode(mode string) {
_m.Called(mode)
}

func (_m *MockInvokeMetrics) SetReservationUsed(wasReserved bool) {
_m.Called(wasReserved)
}

func (_m *MockInvokeMetrics) TriggerGetRequest() {
func (_m *MockInvokeMetrics) SetResponseDeliveryLost() {
_m.Called()
}

func (_m *MockInvokeMetrics) SetResponseDeliverySent() {
_m.Called()
}

func (_m *MockInvokeMetrics) SetResponseWaitTime(d time.Duration) {
_m.Called(d)
}

func (_m *MockInvokeMetrics) TriggerGetRequest() time.Time {
ret := _m.Called()

if len(ret) == 0 {
panic("no return value specified for TriggerGetRequest")
}

var r0 time.Time
if rf, ok := ret.Get(0).(func() time.Time); ok {
r0 = rf()
} else {
r0 = ret.Get(0).(time.Time)
}

return r0
}

func (_m *MockInvokeMetrics) TriggerGetResponse() {
_m.Called()
}
Expand Down
Loading
Loading