diff --git a/internal/lambda-managed-instances/aws-lambda-rie/internal/app.go b/internal/lambda-managed-instances/aws-lambda-rie/internal/app.go index 18dbe8e0..f8f540c3 100644 --- a/internal/lambda-managed-instances/aws-lambda-rie/internal/app.go +++ b/internal/lambda-managed-instances/aws-lambda-rie/internal/app.go @@ -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) @@ -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{} diff --git a/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/responder.go b/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/responder.go index 2192ee57..aa2e5369 100644 --- a/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/responder.go +++ b/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/responder.go @@ -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, } } @@ -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) } } @@ -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) } } diff --git a/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/responder_test.go b/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/responder_test.go index 1e3e1479..15db0d53 100644 --- a/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/responder_test.go +++ b/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/responder_test.go @@ -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" @@ -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) @@ -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")) @@ -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")), "") @@ -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) @@ -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") @@ -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) @@ -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 { diff --git a/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/rie_invoke_request.go b/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/rie_invoke_request.go index a5d1d3d1..f3a7d73c 100644 --- a/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/rie_invoke_request.go +++ b/internal/lambda-managed-instances/aws-lambda-rie/internal/invoke/rie_invoke_request.go @@ -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) { @@ -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 } @@ -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) @@ -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 } diff --git a/internal/lambda-managed-instances/aws-lambda-rie/internal/mock_raptor_app.go b/internal/lambda-managed-instances/aws-lambda-rie/internal/mock_raptor_app.go index 1132f8d1..121e4263 100644 --- a/internal/lambda-managed-instances/aws-lambda-rie/internal/mock_raptor_app.go +++ b/internal/lambda-managed-instances/aws-lambda-rie/internal/mock_raptor_app.go @@ -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" @@ -37,8 +38,8 @@ 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") @@ -46,24 +47,31 @@ func (_m *mockRaptorApp) Invoke(ctx context.Context, msg interop.InvokeRequest, 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 { diff --git a/internal/lambda-managed-instances/aws-lambda-rie/internal/run.go b/internal/lambda-managed-instances/aws-lambda-rie/internal/run.go index 2010ebb5..c0ced2b6 100644 --- a/internal/lambda-managed-instances/aws-lambda-rie/internal/run.go +++ b/internal/lambda-managed-instances/aws-lambda-rie/internal/run.go @@ -7,6 +7,7 @@ import ( "context" "fmt" "log/slog" + "net/http" "os" "time" @@ -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{ @@ -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), } diff --git a/internal/lambda-managed-instances/interop/mock_init_static_data_provider.go b/internal/lambda-managed-instances/interop/mock_init_static_data_provider.go index bebc4b80..971dcc76 100644 --- a/internal/lambda-managed-instances/interop/mock_init_static_data_provider.go +++ b/internal/lambda-managed-instances/interop/mock_init_static_data_provider.go @@ -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() diff --git a/internal/lambda-managed-instances/interop/mock_invoke_metrics.go b/internal/lambda-managed-instances/interop/mock_invoke_metrics.go index 51e57586..5a9c9601 100644 --- a/internal/lambda-managed-instances/interop/mock_invoke_metrics.go +++ b/internal/lambda-managed-instances/interop/mock_invoke_metrics.go @@ -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() } diff --git a/internal/lambda-managed-instances/interop/mock_invoke_request.go b/internal/lambda-managed-instances/interop/mock_invoke_request.go index b6aef276..53f3d5e4 100644 --- a/internal/lambda-managed-instances/interop/mock_invoke_request.go +++ b/internal/lambda-managed-instances/interop/mock_invoke_request.go @@ -5,12 +5,9 @@ package interop import ( io "io" - http "net/http" - - mock "github.com/stretchr/testify/mock" - time "time" + mock "github.com/stretchr/testify/mock" model "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/rapid/model" ) @@ -18,10 +15,6 @@ type MockInvokeRequest struct { mock.Mock } -func (_m *MockInvokeRequest) AddResponseHeader(_a0 string, _a1 string) { - _m.Called(_a0, _a1) -} - func (_m *MockInvokeRequest) BodyReader() io.Reader { ret := _m.Called() @@ -177,6 +170,25 @@ func (_m *MockInvokeRequest) InvokeID() string { return r0 } +func (_m *MockInvokeRequest) LongPollingConfig() *LongPollingConfig { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for LongPollingConfig") + } + + var r0 *LongPollingConfig + if rf, ok := ret.Get(0).(func() *LongPollingConfig); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*LongPollingConfig) + } + } + + return r0 +} + func (_m *MockInvokeRequest) MaxPayloadSize() int64 { ret := _m.Called() @@ -194,11 +206,11 @@ func (_m *MockInvokeRequest) MaxPayloadSize() int64 { return r0 } -func (_m *MockInvokeRequest) ResponseBandwidthBurstRate() int64 { +func (_m *MockInvokeRequest) ResolvedFunctionTimeoutMs() int64 { ret := _m.Called() if len(ret) == 0 { - panic("no return value specified for ResponseBandwidthBurstRate") + panic("no return value specified for ResolvedFunctionTimeoutMs") } var r0 int64 @@ -211,11 +223,28 @@ func (_m *MockInvokeRequest) ResponseBandwidthBurstRate() int64 { return r0 } -func (_m *MockInvokeRequest) ResponseBandwidthRate() int64 { +func (_m *MockInvokeRequest) ResolvedInvokeTimeout() time.Duration { ret := _m.Called() if len(ret) == 0 { - panic("no return value specified for ResponseBandwidthRate") + panic("no return value specified for ResolvedInvokeTimeout") + } + + var r0 time.Duration + if rf, ok := ret.Get(0).(func() time.Duration); ok { + r0 = rf() + } else { + r0 = ret.Get(0).(time.Duration) + } + + return r0 +} + +func (_m *MockInvokeRequest) ResponseBandwidthBurstRate() int64 { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for ResponseBandwidthBurstRate") } var r0 int64 @@ -228,46 +257,40 @@ func (_m *MockInvokeRequest) ResponseBandwidthRate() int64 { return r0 } -func (_m *MockInvokeRequest) ResponseMode() string { +func (_m *MockInvokeRequest) ResponseBandwidthRate() int64 { ret := _m.Called() if len(ret) == 0 { - panic("no return value specified for ResponseMode") + panic("no return value specified for ResponseBandwidthRate") } - var r0 string - if rf, ok := ret.Get(0).(func() string); ok { + var r0 int64 + if rf, ok := ret.Get(0).(func() int64); ok { r0 = rf() } else { - r0 = ret.Get(0).(string) + r0 = ret.Get(0).(int64) } return r0 } -func (_m *MockInvokeRequest) ResponseWriter() http.ResponseWriter { +func (_m *MockInvokeRequest) ResponseMode() string { ret := _m.Called() if len(ret) == 0 { - panic("no return value specified for ResponseWriter") + panic("no return value specified for ResponseMode") } - var r0 http.ResponseWriter - if rf, ok := ret.Get(0).(func() http.ResponseWriter); ok { + var r0 string + if rf, ok := ret.Get(0).(func() string); ok { r0 = rf() } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(http.ResponseWriter) - } + r0 = ret.Get(0).(string) } return r0 } -func (_m *MockInvokeRequest) SetResponseHeader(_a0 string, _a1 string) { - _m.Called(_a0, _a1) -} - func (_m *MockInvokeRequest) TraceId() string { ret := _m.Called() @@ -304,10 +327,6 @@ func (_m *MockInvokeRequest) UpdateFromInitData(_a0 InitStaticDataProvider) mode return r0 } -func (_m *MockInvokeRequest) WriteResponseHeaders(_a0 int) { - _m.Called(_a0) -} - func NewMockInvokeRequest(t interface { mock.TestingT Cleanup(func()) diff --git a/internal/lambda-managed-instances/interop/mock_rapid_context.go b/internal/lambda-managed-instances/interop/mock_rapid_context.go index 49a616b2..a6cfbce8 100644 --- a/internal/lambda-managed-instances/interop/mock_rapid_context.go +++ b/internal/lambda-managed-instances/interop/mock_rapid_context.go @@ -5,10 +5,13 @@ package interop import ( context "context" - netip "net/netip" + http "net/http" mock "github.com/stretchr/testify/mock" + model "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/rapid/model" + + netip "net/netip" ) type MockRapidContext struct { @@ -34,8 +37,8 @@ func (_m *MockRapidContext) HandleInit(ctx context.Context, initData InitExecuti return r0 } -func (_m *MockRapidContext) HandleInvoke(ctx context.Context, invokeRequest InvokeRequest, invokeMetrics InvokeMetrics) (model.AppError, bool) { - ret := _m.Called(ctx, invokeRequest, invokeMetrics) +func (_m *MockRapidContext) HandleInvoke(ctx context.Context, invokeRequest InvokeRequest, invokeMetrics InvokeMetrics, responseWriter http.ResponseWriter) (model.AppError, bool, bool) { + ret := _m.Called(ctx, invokeRequest, invokeMetrics, responseWriter) if len(ret) == 0 { panic("no return value specified for HandleInvoke") @@ -43,24 +46,48 @@ func (_m *MockRapidContext) HandleInvoke(ctx context.Context, invokeRequest Invo var r0 model.AppError var r1 bool - if rf, ok := ret.Get(0).(func(context.Context, InvokeRequest, InvokeMetrics) (model.AppError, bool)); ok { - return rf(ctx, invokeRequest, invokeMetrics) + var r2 bool + if rf, ok := ret.Get(0).(func(context.Context, InvokeRequest, InvokeMetrics, http.ResponseWriter) (model.AppError, bool, bool)); ok { + return rf(ctx, invokeRequest, invokeMetrics, responseWriter) } - if rf, ok := ret.Get(0).(func(context.Context, InvokeRequest, InvokeMetrics) model.AppError); ok { - r0 = rf(ctx, invokeRequest, invokeMetrics) + if rf, ok := ret.Get(0).(func(context.Context, InvokeRequest, InvokeMetrics, http.ResponseWriter) model.AppError); ok { + r0 = rf(ctx, invokeRequest, invokeMetrics, responseWriter) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(model.AppError) } } - if rf, ok := ret.Get(1).(func(context.Context, InvokeRequest, InvokeMetrics) bool); ok { - r1 = rf(ctx, invokeRequest, invokeMetrics) + if rf, ok := ret.Get(1).(func(context.Context, InvokeRequest, InvokeMetrics, http.ResponseWriter) bool); ok { + r1 = rf(ctx, invokeRequest, invokeMetrics, responseWriter) } else { r1 = ret.Get(1).(bool) } - return r0, r1 + if rf, ok := ret.Get(2).(func(context.Context, InvokeRequest, InvokeMetrics, http.ResponseWriter) bool); ok { + r2 = rf(ctx, invokeRequest, invokeMetrics, responseWriter) + } else { + r2 = ret.Get(2).(bool) + } + + return r0, r1, r2 +} + +func (_m *MockRapidContext) HandleReconnect(ctx context.Context, invokeID string, responseWriter http.ResponseWriter, metrics ReconnectMetrics) ReconnectResult { + ret := _m.Called(ctx, invokeID, responseWriter, metrics) + + if len(ret) == 0 { + panic("no return value specified for HandleReconnect") + } + + var r0 ReconnectResult + if rf, ok := ret.Get(0).(func(context.Context, string, http.ResponseWriter, ReconnectMetrics) ReconnectResult); ok { + r0 = rf(ctx, invokeID, responseWriter, metrics) + } else { + r0 = ret.Get(0).(ReconnectResult) + } + + return r0 } func (_m *MockRapidContext) HandleShutdown(shutdownCause model.AppError, metrics ShutdownMetrics) model.AppError { diff --git a/internal/lambda-managed-instances/interop/mock_reconnect_metrics.go b/internal/lambda-managed-instances/interop/mock_reconnect_metrics.go new file mode 100644 index 00000000..cf792bb0 --- /dev/null +++ b/internal/lambda-managed-instances/interop/mock_reconnect_metrics.go @@ -0,0 +1,50 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 + +package interop + +import ( + time "time" + + mock "github.com/stretchr/testify/mock" +) + +type MockReconnectMetrics struct { + mock.Mock +} + +func (_m *MockReconnectMetrics) SetFunctionDoneTime(t time.Time) { + _m.Called(t) +} + +func (_m *MockReconnectMetrics) TriggerConnectionGap(lastDisconnectTime *time.Time) { + _m.Called(lastDisconnectTime) +} + +func (_m *MockReconnectMetrics) TriggerPollEnd() { + _m.Called() +} + +func (_m *MockReconnectMetrics) TriggerPollStart() { + _m.Called() +} + +func (_m *MockReconnectMetrics) TriggerResponseReplayDone(size int) { + _m.Called(size) +} + +func (_m *MockReconnectMetrics) TriggerResponseReplayStart() { + _m.Called() +} + +func NewMockReconnectMetrics(t interface { + mock.TestingT + Cleanup(func()) +}) *MockReconnectMetrics { + mock := &MockReconnectMetrics{} + mock.Mock.Test(t) + + t.Cleanup(func() { mock.AssertExpectations(t) }) + + return mock +} diff --git a/internal/lambda-managed-instances/interop/sandbox_model.go b/internal/lambda-managed-instances/interop/sandbox_model.go index 8f50ea45..6c81b666 100644 --- a/internal/lambda-managed-instances/interop/sandbox_model.go +++ b/internal/lambda-managed-instances/interop/sandbox_model.go @@ -239,6 +239,10 @@ func (i *InitExecutionData) RuntimeVersion() string { return i.StaticData.RuntimeVersion } +func (i *InitExecutionData) RuntimeRelease() string { + return i.StaticData.RuntimeRelease +} + func (i *InitExecutionData) AvailabilityZoneId() string { return i.StaticData.AvailabilityZoneId } @@ -263,11 +267,53 @@ type TelemetrySubscriptionConfig struct { APIAddr netip.AddrPort } +type ReconnectOutcome string + +const ( + ReconnectOutcomeCompleted ReconnectOutcome = "completed" + ReconnectOutcomePendingHit ReconnectOutcome = "pending_hit" + ReconnectOutcomeTimeout ReconnectOutcome = "timeout" + ReconnectOutcomeNotFound ReconnectOutcome = "not_found" + ReconnectOutcomeDisplaced ReconnectOutcome = "displaced" + ReconnectOutcomeError ReconnectOutcome = "error" +) + +type ReconnectResult struct { + InvokeMetrics InvokeMetrics + + FunctionDoneTime time.Time + Outcome ReconnectOutcome + Err model.AppError + WasResponseSent bool +} + +func (r *ReconnectResult) FinalizeInvokeMetrics() (totalMs time.Duration, runMs *time.Duration, initData InitStaticDataProvider, ok bool) { + if r.InvokeMetrics == nil || !r.WasResponseSent { + return 0, nil, nil, false + } + if !r.FunctionDoneTime.IsZero() { + r.InvokeMetrics.SetResponseWaitTime(time.Since(r.FunctionDoneTime)) + } + r.InvokeMetrics.SetResponseDeliverySent() + totalMs, runMs, initData = r.InvokeMetrics.TriggerInvokeDone() + return totalMs, runMs, initData, true +} + +type ReconnectMetrics interface { + TriggerConnectionGap(lastDisconnectTime *time.Time) + TriggerPollStart() + TriggerPollEnd() + TriggerResponseReplayStart() + TriggerResponseReplayDone(size int) + SetFunctionDoneTime(t time.Time) +} + type RapidContext interface { HandleInit(ctx context.Context, initData InitExecutionData, initMetrics InitMetrics) (err model.AppError) HandleShutdown(shutdownCause model.AppError, metrics ShutdownMetrics) model.AppError - HandleInvoke(ctx context.Context, invokeRequest InvokeRequest, invokeMetrics InvokeMetrics) (err model.AppError, wasResponseSent bool) + HandleInvoke(ctx context.Context, invokeRequest InvokeRequest, invokeMetrics InvokeMetrics, responseWriter http.ResponseWriter) (err model.AppError, wasResponseSent bool, invokePending bool) + HandleReconnect(ctx context.Context, invokeID InvokeID, responseWriter http.ResponseWriter, metrics ReconnectMetrics) ReconnectResult RuntimeAPIAddrPort() netip.AddrPort ProcessTerminationNotifier() <-chan model.AppError @@ -294,18 +340,21 @@ type InvokeRequest interface { ResponseMode() string BodyReader() io.Reader - ResponseWriter() http.ResponseWriter - - SetResponseHeader(string, string) - AddResponseHeader(string, string) - WriteResponseHeaders(int) UpdateFromInitData(InitStaticDataProvider) model.AppError FunctionVersionID() string + ResolvedFunctionTimeoutMs() int64 + ResolvedInvokeTimeout() time.Duration + LongPollingConfig() *LongPollingConfig InternalInvocationID() string } +type LongPollingConfig struct { + ConnectionHoldTimeoutMs int64 + ResponseHoldTimeoutMs int64 +} + type InitStaticDataProvider interface { FunctionARN() string FunctionVersion() string @@ -319,11 +368,12 @@ type InitStaticDataProvider interface { ArtefactType() intmodel.ArtefactType AmiId() string RuntimeVersion() string + RuntimeRelease() string AvailabilityZoneId() string } type InvokeMetrics interface { - TriggerGetRequest() + TriggerGetRequest() time.Time AttachInvokeRequest(InvokeRequest) AttachDependencies(InitStaticDataProvider, EventsAPI) UpdateConcurrencyMetrics(inflightInvokes, idleRuntimesCount int) @@ -335,6 +385,10 @@ type InvokeMetrics interface { TriggerInvokeDone() (totalMs time.Duration, runMs *time.Duration, initData InitStaticDataProvider) SetReservationUsed(wasReserved bool) + SetResponseWaitTime(d time.Duration) + SetResponseDeliveryLost() + SetResponseDeliverySent() + SetInvokeMode(mode string) SendInvokeStartEvent(*TracingCtx) error SendInvokeFinishedEvent(tracingCtx *TracingCtx, xrayErrorCause json.RawMessage) error diff --git a/internal/lambda-managed-instances/invoke/invoke_router.go b/internal/lambda-managed-instances/invoke/invoke_router.go index 17be7f86..829551c7 100644 --- a/internal/lambda-managed-instances/invoke/invoke_router.go +++ b/internal/lambda-managed-instances/invoke/invoke_router.go @@ -34,9 +34,11 @@ type RuntimeResponseRequest interface { BodyReader() io.Reader + Cancel() + TrailerError() ErrorForInvoker - InvocationID() string + InvocationID() *string } type RuntimeErrorRequest interface { @@ -50,11 +52,11 @@ type RuntimeErrorRequest interface { ErrorDetails() string GetXrayErrorCause() json.RawMessage - InvocationID() string + InvocationID() *string } type runningInvoke interface { - RunInvokeAndSendResult(context.Context, interop.InitStaticDataProvider, interop.InvokeRequest, interop.InvokeMetrics) model.AppError + RunInvokeAndSendResult(context.Context, interop.InitStaticDataProvider, interop.InvokeRequest, interop.InvokeMetrics, InvokeResponseSender) model.AppError RuntimeNextWait(context.Context) model.AppError RuntimeResponse(context.Context, RuntimeResponseRequest) model.AppError RuntimeError(context.Context, RuntimeErrorRequest) model.AppError @@ -83,7 +85,6 @@ type InvokeRouter struct { func NewInvokeRouter( runtimePoolSize int, telemetryEventsApi interop.EventsAPI, - responderFactoryFunc ResponderFactoryFunc, timeoutCache timeoutCache, ) *InvokeRouter { return &InvokeRouter{ @@ -92,13 +93,13 @@ func NewInvokeRouter( eventsApi: telemetryEventsApi, timeoutCache: timeoutCache, createRunningInvoke: func(runtimeNext http.ResponseWriter) runningInvoke { - r := newRunningInvoke(runtimeNext, responderFactoryFunc, timeoutCache) + r := newRunningInvoke(runtimeNext, timeoutCache) return &r }, } } -func (ir *InvokeRouter) Invoke(ctx context.Context, initData interop.InitStaticDataProvider, invokeReq interop.InvokeRequest, metrics interop.InvokeMetrics) (err model.AppError, wasResponseSent bool) { +func (ir *InvokeRouter) Invoke(ctx context.Context, initData interop.InitStaticDataProvider, invokeReq interop.InvokeRequest, metrics interop.InvokeMetrics, sender InvokeResponseSender) (err model.AppError, wasResponseSent bool) { logging.Debug(ctx, "InvokeRouter: received Invoke") ir.wg.Add(1) defer ir.wg.Done() @@ -124,7 +125,7 @@ func (ir *InvokeRouter) Invoke(ctx context.Context, initData interop.InitStaticD ir.runningInvokes.Set(invokeReq.InvokeID(), idleRuntime) - return idleRuntime.RunInvokeAndSendResult(ctx, initData, invokeReq, metrics), true + return idleRuntime.RunInvokeAndSendResult(ctx, initData, invokeReq, metrics, sender), true } func (ir *InvokeRouter) RuntimeNext(ctx context.Context, runtimeReq http.ResponseWriter) (model.RuntimeNextWaiter, model.AppError) { diff --git a/internal/lambda-managed-instances/invoke/invoke_router_test.go b/internal/lambda-managed-instances/invoke/invoke_router_test.go index 86685947..42b55b72 100644 --- a/internal/lambda-managed-instances/invoke/invoke_router_test.go +++ b/internal/lambda-managed-instances/invoke/invoke_router_test.go @@ -62,7 +62,7 @@ func hijackInvokeRouterDeps(router *InvokeRouter, mocks *invokeRouterMocks) { func createMocksAndInitRouter() (*invokeRouterMocks, *InvokeRouter) { mocks := newInvokeRouterMocks() - router := NewInvokeRouter(testInvokeRouterMaxIdleRuntime, &telemetry.NoOpEventsAPI{}, nil, mocks.timeoutCache) + router := NewInvokeRouter(testInvokeRouterMaxIdleRuntime, &telemetry.NoOpEventsAPI{}, mocks.timeoutCache) hijackInvokeRouterDeps(router, &mocks) return &mocks, router @@ -96,7 +96,7 @@ func TestInvokeSuccess(t *testing.T) { mocks.invokeMetrics.On("UpdateConcurrencyMetrics", 0, 1) mocks.eaInvokeRequest.On("InvokeID").Return("123456") - mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &mocks.staticData, &mocks.eaInvokeRequest, mock.Anything).Run(func(args mock.Arguments) { + mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &mocks.staticData, &mocks.eaInvokeRequest, mock.Anything, mock.Anything).Run(func(args mock.Arguments) { close(syncChan) @@ -112,7 +112,7 @@ func TestInvokeSuccess(t *testing.T) { wg.Add(1) go func() { defer wg.Done() - err, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics) + err, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics, nil) assert.NoError(t, err) assert.True(t, wasResponseSent) }() @@ -138,7 +138,7 @@ func TestInvokeFailure_NoIdleRuntime(t *testing.T) { mocks.eaInvokeRequest.On("InvokeID").Return("123456") - err, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics) + err, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics, nil) assert.Error(t, err) assert.False(t, wasResponseSent) assert.Equal(t, model.ErrorRuntimeUnavailable, err.ErrorType()) @@ -165,7 +165,7 @@ func TestInvokeFailure_DublicatedInvokeId(t *testing.T) { mocks.invokeMetrics.On("SetReservationUsed", mock.AnythingOfType("bool")).Maybe() mocks.eaInvokeRequest.On("InvokeID").Return("123456") - mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &mocks.staticData, &mocks.eaInvokeRequest, mock.Anything).Return(nil).WaitUntil(respChannel).Once() + mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &mocks.staticData, &mocks.eaInvokeRequest, mock.Anything, mock.Anything).Return(nil).WaitUntil(respChannel).Once() wg := new(sync.WaitGroup) ch := make(chan model.AppError, 2) @@ -174,7 +174,7 @@ func TestInvokeFailure_DublicatedInvokeId(t *testing.T) { wg.Add(1) go func() { defer wg.Done() - err, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics) + err, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics, nil) if wasResponseSent { wasResponseSentCnt.Add(1) } @@ -184,7 +184,7 @@ func TestInvokeFailure_DublicatedInvokeId(t *testing.T) { wg.Add(1) go func() { defer wg.Done() - err, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics) + err, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics, nil) if wasResponseSent { wasResponseSentCnt.Add(1) } @@ -400,6 +400,7 @@ func TestReserveIdleRuntime_DuplicateInvokeID(t *testing.T) { } func TestReserveIdleRuntime_DuplicateAgainstRunningInvoke(t *testing.T) { + t.Parallel() mocks, router := createMocksAndInitRouter() @@ -457,9 +458,9 @@ func TestReserveIdleRuntime_InvokeConsumesReservation(t *testing.T) { mocks.invokeMetrics.On("UpdateConcurrencyMetrics", mock.AnythingOfType("int"), mock.AnythingOfType("int")) mocks.invokeMetrics.On("SetReservationUsed", true) mocks.eaInvokeRequest.On("InvokeID").Return("reserve-then-invoke") - mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &mocks.staticData, &mocks.eaInvokeRequest, mock.Anything).Return(nil) + mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &mocks.staticData, &mocks.eaInvokeRequest, mock.Anything, mock.Anything).Return(nil) - invokeErr, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics) + invokeErr, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics, nil) assert.NoError(t, invokeErr) assert.True(t, wasResponseSent) @@ -479,9 +480,9 @@ func TestReserveIdleRuntime_InvokeWithoutReservation(t *testing.T) { mocks.invokeMetrics.On("UpdateConcurrencyMetrics", mock.AnythingOfType("int"), mock.AnythingOfType("int")) mocks.invokeMetrics.On("SetReservationUsed", false) mocks.eaInvokeRequest.On("InvokeID").Return("no-reservation-invoke") - mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &mocks.staticData, &mocks.eaInvokeRequest, mock.Anything).Return(nil) + mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &mocks.staticData, &mocks.eaInvokeRequest, mock.Anything, mock.Anything).Return(nil) - invokeErr, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics) + invokeErr, wasResponseSent := router.Invoke(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.invokeMetrics, nil) assert.NoError(t, invokeErr) assert.True(t, wasResponseSent) } diff --git a/internal/lambda-managed-instances/invoke/long_invoke.go b/internal/lambda-managed-instances/invoke/long_invoke.go new file mode 100644 index 00000000..e790f7a4 --- /dev/null +++ b/internal/lambda-managed-instances/invoke/long_invoke.go @@ -0,0 +1,298 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 + +package invoke + +import ( + "context" + "errors" + "log/slog" + "net/http" + "sync" + "sync/atomic" + "time" + + cmap "github.com/orcaman/concurrent-map" + + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/interop" + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/logging" + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/rapid/model" +) + +var ( + ErrInvokeSessionNotFound = errors.New("invoke session not found") + ErrPendingResponseExpired = errors.New("pending response expired") +) + +const ( + DefaultPendingResponseTTL = 10 * time.Second + DefaultLongInvokeThresholdMs = int64(15 * time.Minute / time.Millisecond) + + maxShutdownDrainTimeout = 15 * time.Second + + headerInvokeID = "invoke-id" + headerFunctionVersionID = "invoked-function-version" + headerWaitingReason = "Invoke-Wait-Reason" +) + +type WaitingReason string + +const ( + WaitReasonStillRunning WaitingReason = "still_running" + WaitReasonDisplaced WaitingReason = "displaced" +) + +type pendingInvoke struct { + resultCh chan invokeResult + + preemptCh atomic.Value + + lastDisconnectTime atomic.Value + + functionVersionID string + + connectionHoldTimeout time.Duration + + responseHoldTimeout time.Duration +} + +func newPendingInvoke(functionVersionID string, connectionHoldTimeout time.Duration, responseHoldTimeout time.Duration) *pendingInvoke { + pi := &pendingInvoke{ + resultCh: make(chan invokeResult), + functionVersionID: functionVersionID, + connectionHoldTimeout: connectionHoldTimeout, + responseHoldTimeout: responseHoldTimeout, + } + pi.preemptCh.Store(make(chan struct{})) + return pi +} + +func (p *pendingInvoke) preempt() <-chan struct{} { + newCh := make(chan struct{}) + old := p.preemptCh.Swap(newCh).(chan struct{}) + close(old) + return newCh +} + +func (r invokeResult) outcome(pollStart time.Time) interop.ReconnectOutcome { + if !r.functionDoneTime.IsZero() && r.functionDoneTime.Before(pollStart) { + return interop.ReconnectOutcomePendingHit + } + return interop.ReconnectOutcomeCompleted +} + +func (p *pendingInvoke) recordDisconnectIfPending(outcome interop.ReconnectOutcome) bool { + pending := outcome == interop.ReconnectOutcomeTimeout || outcome == interop.ReconnectOutcomeDisplaced + if pending { + p.lastDisconnectTime.Store(time.Now()) + } + return pending +} + +func (p *pendingInvoke) getLastDisconnectTime() *time.Time { + if v := p.lastDisconnectTime.Load(); v != nil { + t := v.(time.Time) + return &t + } + return nil +} + +type LongInvokerRouter struct { + innerRouter *InvokeRouter + responderFactory ResponderFactoryFunc + + LongInvokeThresholdMs int64 + + pendingInvokes cmap.ConcurrentMap + + wg sync.WaitGroup +} + +func NewLongInvokerRouter(invokeRouter *InvokeRouter, responderFactory ResponderFactoryFunc) *LongInvokerRouter { + return &LongInvokerRouter{ + innerRouter: invokeRouter, + responderFactory: responderFactory, + LongInvokeThresholdMs: DefaultLongInvokeThresholdMs, + pendingInvokes: cmap.New(), + } +} + +func (l *LongInvokerRouter) InnerRouter() *InvokeRouter { return l.innerRouter } + +type invokeResult struct { + recorder *ResponseWriterRecorder + err model.AppError + metrics interop.InvokeMetrics + + functionDoneTime time.Time +} + +func newInvokeResult(recorder *ResponseWriterRecorder, err model.AppError, metrics interop.InvokeMetrics) invokeResult { + return invokeResult{recorder: recorder, err: err, metrics: metrics, functionDoneTime: time.Now()} +} + +func (l *LongInvokerRouter) Invoke(ctx context.Context, initData interop.InitStaticDataProvider, invokeReq interop.InvokeRequest, metrics interop.InvokeMetrics, responseWriter http.ResponseWriter) (err model.AppError, wasResponseSent bool, invokePending bool) { + + lpCfg := invokeReq.LongPollingConfig() + isLongInvoke := lpCfg != nil && invokeReq.ResolvedFunctionTimeoutMs() > l.LongInvokeThresholdMs + if !isLongInvoke { + metrics.SetInvokeMode(string(InvokeModeNormal)) + directResponder := l.responderFactory(ctx, invokeReq, responseWriter) + err, wasResponseSent := l.innerRouter.Invoke(ctx, initData, invokeReq, metrics, directResponder) + return err, wasResponseSent, false + } + + metrics.SetInvokeMode(string(InvokeModeLong)) + + logging.Info(ctx, "LongInvokerRouter: starting long invoke") + + recorder := &ResponseWriterRecorder{} + directResponder := l.responderFactory(ctx, invokeReq, recorder) + + connectionHoldTimeout := time.Duration(lpCfg.ConnectionHoldTimeoutMs) * time.Millisecond + responseHoldTimeout := time.Duration(lpCfg.ResponseHoldTimeoutMs) * time.Millisecond + + pending := newPendingInvoke(invokeReq.FunctionVersionID(), connectionHoldTimeout, responseHoldTimeout) + + if !l.pendingInvokes.SetIfAbsent(invokeReq.InvokeID(), pending) { + logging.Error(ctx, "LongInvokerRouter error: duplicated invokeId") + return model.NewClientError(ErrInvokeIdAlreadyExists, model.ErrorSeverityError, model.ErrorDuplicatedInvokeId), false, false + } + + bgCtx := context.WithoutCancel(ctx) + l.wg.Add(1) + go func() { + defer l.wg.Done() + invokeErr, wasResponseSent := l.innerRouter.Invoke(bgCtx, initData, invokeReq, metrics, directResponder) + if wasResponseSent { + l.sendInvokeResult(invokeReq.InvokeID(), newInvokeResult(recorder, invokeErr, metrics)) + } else { + l.sendInvokeResult(invokeReq.InvokeID(), newInvokeResult(nil, invokeErr, metrics)) + } + }() + + pollCtx, pollCancel := context.WithTimeout(ctx, connectionHoldTimeout) + defer pollCancel() + + pollResult := l.longPollInvoke(pollCtx, invokeReq.InvokeID(), pending.functionVersionID, responseWriter, pending.resultCh, pending.preemptCh.Load().(chan struct{}), NoopReconnectMetrics()) + invokePending = pending.recordDisconnectIfPending(pollResult.Outcome) + return pollResult.Err, pollResult.WasResponseSent, invokePending +} + +func (l *LongInvokerRouter) Reconnect(ctx context.Context, invokeID interop.InvokeID, responseWriter http.ResponseWriter, reconnectMetrics interop.ReconnectMetrics) interop.ReconnectResult { + val, ok := l.pendingInvokes.Get(invokeID) + if !ok { + return interop.ReconnectResult{Outcome: interop.ReconnectOutcomeNotFound, Err: model.NewClientError(ErrInvokeSessionNotFound, model.ErrorSeverityInvalid, model.ErrorInvalidInvokeId)} + } + pending := val.(*pendingInvoke) + + reconnectMetrics.TriggerConnectionGap(pending.getLastDisconnectTime()) + + preemptCh := pending.preempt() + + pollCtx, pollCancel := context.WithTimeout(ctx, pending.connectionHoldTimeout) + defer pollCancel() + + result := l.longPollInvoke(pollCtx, invokeID, pending.functionVersionID, responseWriter, pending.resultCh, preemptCh, reconnectMetrics) + pending.recordDisconnectIfPending(result.Outcome) + + return result +} + +func (l *LongInvokerRouter) sendStatusAccepted(w http.ResponseWriter, invokeID interop.InvokeID, functionVersionID string, reason WaitingReason) { + w.Header().Set(headerInvokeID, string(invokeID)) + w.Header().Set(headerFunctionVersionID, functionVersionID) + w.Header().Set(headerWaitingReason, string(reason)) + w.WriteHeader(http.StatusAccepted) +} + +func (l *LongInvokerRouter) longPollInvoke(ctx context.Context, invokeID interop.InvokeID, functionVersionID string, w http.ResponseWriter, resultCh <-chan invokeResult, preemptCh <-chan struct{}, reconnectMetrics interop.ReconnectMetrics) interop.ReconnectResult { + pollStart := time.Now() + reconnectMetrics.TriggerPollStart() + select { + case result := <-resultCh: + reconnectMetrics.TriggerPollEnd() + reconnectMetrics.SetFunctionDoneTime(result.functionDoneTime) + if result.recorder == nil { + if result.err != nil { + + return interop.ReconnectResult{InvokeMetrics: result.metrics, FunctionDoneTime: result.functionDoneTime, Outcome: interop.ReconnectOutcomeCompleted, Err: result.err} + } + return interop.ReconnectResult{Outcome: interop.ReconnectOutcomeNotFound, Err: model.NewClientError(ErrPendingResponseExpired, model.ErrorSeverityInvalid, model.ErrorInvalidInvokeId)} + } + logging.Info(ctx, "longPollInvoke: delivering recorded response", "invokeId", invokeID) + reconnectMetrics.TriggerResponseReplayStart() + if err := result.recorder.WriteTo(w); err != nil { + reconnectMetrics.TriggerResponseReplayDone(result.recorder.BodySize()) + return interop.ReconnectResult{InvokeMetrics: result.metrics, FunctionDoneTime: result.functionDoneTime, Outcome: interop.ReconnectOutcomeError, Err: model.NewPlatformError(err, model.ErrorResponseReplayFailed), WasResponseSent: true} + } + reconnectMetrics.TriggerResponseReplayDone(result.recorder.BodySize()) + return interop.ReconnectResult{InvokeMetrics: result.metrics, FunctionDoneTime: result.functionDoneTime, Outcome: result.outcome(pollStart), Err: result.err, WasResponseSent: true} + case <-preemptCh: + reconnectMetrics.TriggerPollEnd() + + logging.Info(ctx, "longPollInvoke: preempted by new reconnect, returning 202", "invokeId", invokeID) + l.sendStatusAccepted(w, invokeID, functionVersionID, WaitReasonDisplaced) + return interop.ReconnectResult{Outcome: interop.ReconnectOutcomeDisplaced, WasResponseSent: true} + case <-ctx.Done(): + reconnectMetrics.TriggerPollEnd() + + if errors.Is(context.Cause(ctx), context.DeadlineExceeded) { + logging.Info(ctx, "longPollInvoke: connection hold timeout, returning 202", "invokeId", invokeID) + l.sendStatusAccepted(w, invokeID, functionVersionID, WaitReasonStillRunning) + return interop.ReconnectResult{Outcome: interop.ReconnectOutcomeTimeout, WasResponseSent: true} + } + logging.Warn(ctx, "longPollInvoke: client disconnected", "invokeId", invokeID) + return interop.ReconnectResult{Outcome: interop.ReconnectOutcomeTimeout} + } +} + +func (l *LongInvokerRouter) sendInvokeResult(invokeID interop.InvokeID, result invokeResult) { + val, ok := l.pendingInvokes.Get(invokeID) + if !ok { + + slog.Error("sendInvokeResult: invoke not found, possible duplicate completion or race with abort", "invokeId", invokeID) + return + } + pending := val.(*pendingInvoke) + + select { + case pending.resultCh <- result: + + case <-time.After(pending.responseHoldTimeout): + slog.Error("sendInvokeResult: timed out waiting for poller, response lost", "invokeId", invokeID) + if result.metrics != nil { + result.metrics.SetResponseWaitTime(pending.responseHoldTimeout) + result.metrics.SetResponseDeliveryLost() + result.metrics.TriggerInvokeDone() + if err := result.metrics.SendMetrics(result.err); err != nil { + slog.Error("sendInvokeResult: failed to send metrics on TTL expiry", "invokeId", invokeID, "error", err) + } + } + } + + l.pendingInvokes.RemoveCb(invokeID, func(key string, v interface{}, exists bool) bool { + if exists { + close(v.(*pendingInvoke).resultCh) + } + return true + }) +} + +func (l *LongInvokerRouter) AbortRunningInvokes(metrics interop.ShutdownMetrics, err model.AppError) { + + l.innerRouter.AbortRunningInvokes(metrics, err) + + done := make(chan struct{}) + go func() { l.wg.Wait(); close(done) }() + select { + case <-done: + case <-time.After(maxShutdownDrainTimeout): + slog.Error("AbortRunningInvokes: safety cap reached, proceeding with shutdown", + "timeout", maxShutdownDrainTimeout) + } +} + +func (l *LongInvokerRouter) GetActiveRuntimeCount() int { + return l.innerRouter.GetRuntimePoolCounts().Total + l.innerRouter.GetRunningInvokesCount() +} diff --git a/internal/lambda-managed-instances/invoke/long_invoke_test.go b/internal/lambda-managed-instances/invoke/long_invoke_test.go new file mode 100644 index 00000000..192d7e3b --- /dev/null +++ b/internal/lambda-managed-instances/invoke/long_invoke_test.go @@ -0,0 +1,548 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 + +package invoke + +import ( + "context" + "net/http" + "net/http/httptest" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" + + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/interop" + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/rapid/model" + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/telemetry" +) + +func newTestLongInvokerRouter(t *testing.T) *LongInvokerRouter { + t.Helper() + tc := newMockTimeoutCache(t) + router := NewInvokeRouter(1, &telemetry.NoOpEventsAPI{}, tc) + return NewLongInvokerRouter(router, func(_ context.Context, _ interop.InvokeRequest, _ http.ResponseWriter) InvokeResponseSender { + return &MockInvokeResponseSender{} + }) +} + +func TestLongPollInvoke_ReceivesResponse_WritesToWriter(t *testing.T) { + t.Parallel() + l := newTestLongInvokerRouter(t) + w := httptest.NewRecorder() + + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + + rec := &ResponseWriterRecorder{statusCode: http.StatusOK, body: []byte("hello")} + rec.header = http.Header{"Content-Type": []string{"application/json"}} + go func() { pi.resultCh <- invokeResult{recorder: rec} }() + + res := l.longPollInvoke(context.Background(), "invoke-1", "test-version", w, pi.resultCh, pi.preemptCh.Load().(chan struct{}), NoopReconnectMetrics()) + + assert.True(t, res.WasResponseSent) + assert.NoError(t, res.Err) + assert.Equal(t, http.StatusOK, w.Code) + assert.Equal(t, interop.ReconnectOutcomeCompleted, res.Outcome) + assert.Equal(t, "hello", w.Body.String()) + assert.Equal(t, "application/json", w.Header().Get("Content-Type")) +} + +func TestLongPollInvoke_ContextTimeout_Returns202(t *testing.T) { + t.Parallel() + l := newTestLongInvokerRouter(t) + w := httptest.NewRecorder() + + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + + res := l.longPollInvoke(ctx, "invoke-1", "test-version", w, pi.resultCh, pi.preemptCh.Load().(chan struct{}), NoopReconnectMetrics()) + + assert.True(t, res.WasResponseSent) + assert.Equal(t, interop.ReconnectOutcomeTimeout, res.Outcome) + assert.NoError(t, res.Err) + assert.Equal(t, http.StatusAccepted, w.Code) + assert.Equal(t, "invoke-1", w.Header().Get(headerInvokeID)) + assert.Equal(t, "test-version", w.Header().Get(headerFunctionVersionID)) + assert.Equal(t, string(WaitReasonStillRunning), w.Header().Get(headerWaitingReason)) +} + +func TestLongPollInvoke_ClosedChannel_ReturnsPendingExpired(t *testing.T) { + t.Parallel() + l := newTestLongInvokerRouter(t) + w := httptest.NewRecorder() + + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + close(pi.resultCh) + + res := l.longPollInvoke(context.Background(), "invoke-1", "test-version", w, pi.resultCh, pi.preemptCh.Load().(chan struct{}), NoopReconnectMetrics()) + + assert.False(t, res.WasResponseSent) + assert.Equal(t, interop.ReconnectOutcomeNotFound, res.Outcome) + require.Error(t, res.Err) + assert.Equal(t, model.ErrorInvalidInvokeId, res.Err.ErrorType()) +} + +func TestLongPollInvoke_InnerRouterError_ReturnsFalse(t *testing.T) { + t.Parallel() + l := newTestLongInvokerRouter(t) + w := httptest.NewRecorder() + + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + + routerErr := model.NewClientError(ErrInvokeNoReadyRuntime, model.ErrorSeverityError, model.ErrorRuntimeUnavailable) + go func() { pi.resultCh <- invokeResult{err: routerErr} }() + + res := l.longPollInvoke(context.Background(), "invoke-1", "test-version", w, pi.resultCh, pi.preemptCh.Load().(chan struct{}), NoopReconnectMetrics()) + + assert.False(t, res.WasResponseSent) + assert.Equal(t, interop.ReconnectOutcomeCompleted, res.Outcome) + require.Error(t, res.Err) + assert.Equal(t, model.ErrorRuntimeUnavailable, res.Err.ErrorType()) +} + +type longInvokeTestHarness struct { + mocks invokeRouterMocks + innerRouter *InvokeRouter + longRouter *LongInvokerRouter + mockResponder *MockInvokeResponseSender + originalWriter *httptest.ResponseRecorder + capturedWriter http.ResponseWriter +} + +func newLongInvokeTestHarness(t *testing.T) *longInvokeTestHarness { + t.Helper() + h := &longInvokeTestHarness{ + mocks: newInvokeRouterMocks(), + originalWriter: httptest.NewRecorder(), + } + h.mockResponder = &MockInvokeResponseSender{} + h.innerRouter = NewInvokeRouter(1, &telemetry.NoOpEventsAPI{}, h.mocks.timeoutCache) + hijackInvokeRouterDeps(h.innerRouter, &h.mocks) + h.longRouter = NewLongInvokerRouter(h.innerRouter, func(_ context.Context, _ interop.InvokeRequest, rw http.ResponseWriter) InvokeResponseSender { + h.capturedWriter = rw + return h.mockResponder + }) + h.longRouter.LongInvokeThresholdMs = 0 + h.mocks.eaInvokeRequest.On("FunctionVersionID").Return("test-version") + h.mocks.eaInvokeRequest.On("LongPollingConfig").Return(&interop.LongPollingConfig{ + ConnectionHoldTimeoutMs: 50, + ResponseHoldTimeoutMs: 50, + }) + return h +} + +func (h *longInvokeTestHarness) prepareIdleRuntime(t *testing.T) { + t.Helper() + h.mocks.runnningInvoke.On("RuntimeNextWait", mock.Anything).Return(nil).Once() + waiter, err := h.innerRouter.RuntimeNext(h.mocks.ctx, h.mocks.runtimeNextRequest) + require.NoError(t, err) + require.NoError(t, waiter.RuntimeNextWait(h.mocks.ctx)) +} + +func TestLongInvokerRouter_Invoke_NonLongInvoke_Passthrough(t *testing.T) { + t.Parallel() + h := newLongInvokeTestHarness(t) + h.prepareIdleRuntime(t) + + h.mocks.eaInvokeRequest.On("ResolvedFunctionTimeoutMs").Return(int64(0)) + h.mocks.eaInvokeRequest.On("InvokeID").Return("normal-1") + h.mocks.invokeMetrics.On("UpdateConcurrencyMetrics", mock.Anything, mock.Anything) + h.mocks.invokeMetrics.On("SetInvokeMode", mock.Anything).Maybe() + h.mocks.invokeMetrics.On("SetReservationUsed", false) + h.mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &h.mocks.staticData, &h.mocks.eaInvokeRequest, mock.Anything, mock.Anything).Return(nil) + + err, sent, pending := h.longRouter.Invoke(h.mocks.ctx, &h.mocks.staticData, &h.mocks.eaInvokeRequest, &h.mocks.invokeMetrics, h.originalWriter) + assert.NoError(t, err) + assert.True(t, sent) + assert.False(t, pending) +} + +func TestLongInvokerRouter_Invoke_LongInvoke_HappyPath(t *testing.T) { + t.Parallel() + h := newLongInvokeTestHarness(t) + h.prepareIdleRuntime(t) + + h.mocks.eaInvokeRequest.On("ResolvedFunctionTimeoutMs").Return(int64(900000)) + h.mocks.eaInvokeRequest.On("InvokeID").Return("long-1") + h.mocks.invokeMetrics.On("UpdateConcurrencyMetrics", mock.Anything, mock.Anything) + h.mocks.invokeMetrics.On("SetInvokeMode", mock.Anything).Maybe() + h.mocks.invokeMetrics.On("SetReservationUsed", false) + + h.mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &h.mocks.staticData, &h.mocks.eaInvokeRequest, mock.Anything, mock.Anything). + Run(func(args mock.Arguments) { + h.capturedWriter.Header().Set("Content-Type", "application/json") + h.capturedWriter.WriteHeader(http.StatusOK) + _, _ = h.capturedWriter.Write([]byte("hello")) + sender := args.Get(4).(InvokeResponseSender) + sender.SendRuntimeResponseTrailers(nil) + }).Return(nil) + h.mockResponder.On("SendRuntimeResponseTrailers", mock.Anything).Return() + + err, sent, pending := h.longRouter.Invoke(h.mocks.ctx, &h.mocks.staticData, &h.mocks.eaInvokeRequest, &h.mocks.invokeMetrics, h.originalWriter) + assert.NoError(t, err) + assert.True(t, sent) + + assert.Equal(t, http.StatusOK, h.originalWriter.Code) + assert.Equal(t, "hello", h.originalWriter.Body.String()) + assert.Equal(t, "application/json", h.originalWriter.Header().Get("Content-Type")) + assert.False(t, pending) +} + +func TestLongInvokerRouter_Invoke_LongInvoke_202Timeout(t *testing.T) { + t.Parallel() + h := newLongInvokeTestHarness(t) + h.prepareIdleRuntime(t) + + h.mocks.eaInvokeRequest.On("ResolvedFunctionTimeoutMs").Return(int64(900000)) + h.mocks.eaInvokeRequest.On("InvokeID").Return("long-timeout-1") + h.mocks.invokeMetrics.On("UpdateConcurrencyMetrics", mock.Anything, mock.Anything) + h.mocks.invokeMetrics.On("SetInvokeMode", mock.Anything).Maybe() + h.mocks.invokeMetrics.On("SetReservationUsed", false) + h.mocks.invokeMetrics.On("SetResponseWaitTime", mock.Anything).Maybe() + h.mocks.invokeMetrics.On("SetResponseDeliveryLost").Maybe() + h.mocks.invokeMetrics.On("SetResponseDeliverySent").Maybe() + h.mocks.invokeMetrics.On("TriggerInvokeDone").Return(time.Duration(0), (*time.Duration)(nil), interop.InitStaticDataProvider(nil)).Maybe() + h.mocks.invokeMetrics.On("SendMetrics", mock.Anything).Return(nil).Maybe() + + invokeStarted := make(chan struct{}) + invokeDone := make(chan struct{}) + h.mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &h.mocks.staticData, &h.mocks.eaInvokeRequest, mock.Anything, mock.Anything). + Run(func(args mock.Arguments) { + close(invokeStarted) + <-invokeDone + }).Return(nil) + + err, sent, pending := h.longRouter.Invoke(h.mocks.ctx, &h.mocks.staticData, &h.mocks.eaInvokeRequest, &h.mocks.invokeMetrics, h.originalWriter) + + assert.NoError(t, err) + assert.True(t, sent) + assert.True(t, pending) + assert.Equal(t, http.StatusAccepted, h.originalWriter.Code) + + <-invokeStarted + close(invokeDone) + + time.Sleep(100 * time.Millisecond) +} + +func TestLongInvokerRouter_Invoke_LongInvoke_DuplicateInvokeID(t *testing.T) { + t.Parallel() + h := newLongInvokeTestHarness(t) + + h.innerRouter = NewInvokeRouter(2, &telemetry.NoOpEventsAPI{}, h.mocks.timeoutCache) + hijackInvokeRouterDeps(h.innerRouter, &h.mocks) + h.longRouter = NewLongInvokerRouter(h.innerRouter, func(_ context.Context, _ interop.InvokeRequest, rw http.ResponseWriter) InvokeResponseSender { + h.capturedWriter = rw + return h.mockResponder + }) + h.longRouter.LongInvokeThresholdMs = 0 + + h.mocks.runnningInvoke.On("RuntimeNextWait", mock.Anything).Return(nil).Twice() + w1, err := h.innerRouter.RuntimeNext(h.mocks.ctx, h.mocks.runtimeNextRequest) + require.NoError(t, err) + require.NoError(t, w1.RuntimeNextWait(h.mocks.ctx)) + w2, err := h.innerRouter.RuntimeNext(h.mocks.ctx, h.mocks.runtimeNextRequest) + require.NoError(t, err) + require.NoError(t, w2.RuntimeNextWait(h.mocks.ctx)) + + h.mocks.eaInvokeRequest.On("ResolvedFunctionTimeoutMs").Return(int64(900001)) + h.mocks.eaInvokeRequest.On("InvokeID").Return("dup-1") + h.mocks.invokeMetrics.On("UpdateConcurrencyMetrics", mock.Anything, mock.Anything) + h.mocks.invokeMetrics.On("SetInvokeMode", mock.Anything).Maybe() + h.mocks.invokeMetrics.On("SetReservationUsed", false) + + firstStarted := make(chan struct{}) + firstDone := make(chan struct{}) + h.mocks.runnningInvoke.On("RunInvokeAndSendResult", mock.Anything, &h.mocks.staticData, &h.mocks.eaInvokeRequest, mock.Anything, mock.Anything). + Run(func(args mock.Arguments) { + close(firstStarted) + <-firstDone + h.capturedWriter.WriteHeader(http.StatusOK) + sender := args.Get(4).(InvokeResponseSender) + sender.SendRuntimeResponseTrailers(nil) + }).Return(nil).Once() + h.mockResponder.On("SendRuntimeResponseTrailers", mock.Anything).Return().Maybe() + + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + _, _, _ = h.longRouter.Invoke(h.mocks.ctx, &h.mocks.staticData, &h.mocks.eaInvokeRequest, &h.mocks.invokeMetrics, h.originalWriter) + }() + <-firstStarted + + err2, sent2, pending2 := h.longRouter.Invoke(h.mocks.ctx, &h.mocks.staticData, &h.mocks.eaInvokeRequest, &h.mocks.invokeMetrics, h.originalWriter) + assert.Error(t, err2) + assert.False(t, sent2) + assert.False(t, pending2) + assert.Equal(t, model.ErrorDuplicatedInvokeId, err2.ErrorType()) + + close(firstDone) + wg.Wait() +} + +func TestLongInvokerRouter_Invoke_LongInvoke_InnerRouterError(t *testing.T) { + t.Parallel() + h := newLongInvokeTestHarness(t) + + h.mocks.eaInvokeRequest.On("ResolvedFunctionTimeoutMs").Return(int64(900001)) + h.mocks.eaInvokeRequest.On("InvokeID").Return("err-1") + h.mocks.invokeMetrics.On("UpdateConcurrencyMetrics", mock.Anything, mock.Anything) + h.mocks.invokeMetrics.On("SetInvokeMode", mock.Anything) + + err, sent, pending := h.longRouter.Invoke(h.mocks.ctx, &h.mocks.staticData, &h.mocks.eaInvokeRequest, &h.mocks.invokeMetrics, h.originalWriter) + + assert.Error(t, err) + assert.False(t, sent) + assert.False(t, pending) + assert.Equal(t, model.ErrorRuntimeUnavailable, err.ErrorType()) +} + +func TestSendInvokeResult_TTLExpiry_NoReader(t *testing.T) { + t.Parallel() + + l := newTestLongInvokerRouter(t) + + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + l.pendingInvokes.Set("ttl-1", pi) + + rec := ResponseWriterRecorder{statusCode: http.StatusOK, body: []byte("lost")} + + start := time.Now() + l.sendInvokeResult("ttl-1", invokeResult{recorder: &rec}) + elapsed := time.Since(start) + + assert.GreaterOrEqual(t, elapsed, 40*time.Millisecond) + assert.Less(t, elapsed, 500*time.Millisecond) + + _, exists := l.pendingInvokes.Get("ttl-1") + assert.False(t, exists) +} + +func TestReconnect_UnknownInvokeID_ReturnsNotFound(t *testing.T) { + t.Parallel() + l := newTestLongInvokerRouter(t) + w := httptest.NewRecorder() + + res := l.Reconnect(context.Background(), "unknown-id", w, NoopReconnectMetrics()) + + assert.False(t, res.WasResponseSent) + require.Error(t, res.Err) + assert.Equal(t, model.ErrorInvalidInvokeId, res.Err.ErrorType()) +} + +func TestReconnect_DeliversBufferedResponse(t *testing.T) { + t.Parallel() + l := newTestLongInvokerRouter(t) + w := httptest.NewRecorder() + + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + l.pendingInvokes.Set("invoke-1", pi) + + rec := &ResponseWriterRecorder{statusCode: http.StatusOK, body: []byte("buffered")} + rec.header = http.Header{"Content-Type": []string{"application/json"}} + go func() { pi.resultCh <- invokeResult{recorder: rec} }() + + res := l.Reconnect(context.Background(), "invoke-1", w, NoopReconnectMetrics()) + + assert.True(t, res.WasResponseSent) + assert.NoError(t, res.Err) + assert.Equal(t, http.StatusOK, w.Code) + assert.Equal(t, "buffered", w.Body.String()) + assert.Equal(t, "application/json", w.Header().Get("Content-Type")) +} + +func TestReconnect_HoldTimeout_Returns202(t *testing.T) { + t.Parallel() + l := newTestLongInvokerRouter(t) + w := httptest.NewRecorder() + + pi := newPendingInvoke("test-version", 50*time.Millisecond, 10*time.Second) + l.pendingInvokes.Set("invoke-1", pi) + + res := l.Reconnect(context.Background(), "invoke-1", w, NoopReconnectMetrics()) + + assert.True(t, res.WasResponseSent) + assert.NoError(t, res.Err) + assert.Equal(t, http.StatusAccepted, w.Code) +} + +func TestReconnect_ErrorResult_ReturnsError(t *testing.T) { + t.Parallel() + l := newTestLongInvokerRouter(t) + w := httptest.NewRecorder() + + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + l.pendingInvokes.Set("invoke-1", pi) + + routerErr := model.NewClientError(ErrInvokeNoReadyRuntime, model.ErrorSeverityError, model.ErrorRuntimeUnavailable) + go func() { pi.resultCh <- invokeResult{err: routerErr} }() + + res := l.Reconnect(context.Background(), "invoke-1", w, NoopReconnectMetrics()) + + assert.False(t, res.WasResponseSent) + require.Error(t, res.Err) + assert.Equal(t, model.ErrorRuntimeUnavailable, res.Err.ErrorType()) +} + +func TestReconnect_NewWins_PreemptsExistingPreempt(t *testing.T) { + t.Parallel() + l := newTestLongInvokerRouter(t) + + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + l.pendingInvokes.Set("invoke-1", pi) + + wA := httptest.NewRecorder() + resumeADone := make(chan struct{}) + go func() { + l.Reconnect(context.Background(), "invoke-1", wA, NoopReconnectMetrics()) + close(resumeADone) + }() + + time.Sleep(50 * time.Millisecond) + + wB := httptest.NewRecorder() + rec := &ResponseWriterRecorder{statusCode: http.StatusOK, body: []byte("hello")} + go func() { + time.Sleep(50 * time.Millisecond) + pi.resultCh <- invokeResult{recorder: rec} + }() + + resB := l.Reconnect(context.Background(), "invoke-1", wB, NoopReconnectMetrics()) + + <-resumeADone + assert.Equal(t, http.StatusAccepted, wA.Code) + + assert.True(t, resB.WasResponseSent) + assert.NoError(t, resB.Err) + assert.Equal(t, http.StatusOK, wB.Code) + assert.Equal(t, "hello", wB.Body.String()) +} + +func TestLongPollInvoke_Preempted_Returns202(t *testing.T) { + t.Parallel() + l := newTestLongInvokerRouter(t) + w := httptest.NewRecorder() + + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + preemptCh := pi.preemptCh.Load().(chan struct{}) + close(preemptCh) + + res := l.longPollInvoke(context.Background(), "invoke-1", "test-version", w, pi.resultCh, preemptCh, NoopReconnectMetrics()) + + assert.True(t, res.WasResponseSent) + assert.Equal(t, interop.ReconnectOutcomeDisplaced, res.Outcome) + assert.NoError(t, res.Err) + assert.Equal(t, http.StatusAccepted, w.Code) + assert.Equal(t, "invoke-1", w.Header().Get(headerInvokeID)) + assert.Equal(t, "test-version", w.Header().Get(headerFunctionVersionID)) + assert.Equal(t, string(WaitReasonDisplaced), w.Header().Get(headerWaitingReason)) +} + +func TestSwapPreempt_ClosesOldChannel(t *testing.T) { + t.Parallel() + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + oldCh := pi.preemptCh.Load().(chan struct{}) + + pi.preempt() + + select { + case <-oldCh: + default: + t.Fatal("old preemptCh should be closed") + } +} + +func TestSwapPreempt_ReturnsNewChannel(t *testing.T) { + t.Parallel() + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + + newCh := pi.preempt() + + select { + case <-newCh: + t.Fatal("new preemptCh should not be closed") + default: + } +} + +func TestSwapPreempt_ConcurrentCallsNoRace(t *testing.T) { + t.Parallel() + pi := newPendingInvoke("test-version", 5*time.Second, 50*time.Millisecond) + + var wg sync.WaitGroup + for i := 0; i < 10; i++ { + wg.Add(1) + go func() { + defer wg.Done() + pi.preempt() + }() + } + wg.Wait() +} + +func TestGetActiveRuntimeCount(t *testing.T) { + t.Parallel() + + tc := newMockTimeoutCache(t) + router := NewInvokeRouter(10, &telemetry.NoOpEventsAPI{}, tc) + l := NewLongInvokerRouter(router, func(_ context.Context, _ interop.InvokeRequest, _ http.ResponseWriter) InvokeResponseSender { + return &MockInvokeResponseSender{} + }) + + assert.Equal(t, 0, l.GetActiveRuntimeCount()) + + require.NoError(t, router.runtimePool.Add(newMockRunningInvoke(t))) + require.NoError(t, router.runtimePool.Add(newMockRunningInvoke(t))) + require.NoError(t, router.runtimePool.Add(newMockRunningInvoke(t))) + + assert.Equal(t, 3, l.GetActiveRuntimeCount()) + + router.runningInvokes.Set("invoke-1", newMockRunningInvoke(t)) + + assert.Equal(t, 4, l.GetActiveRuntimeCount()) + + router.runningInvokes.Remove("invoke-1") + + assert.Equal(t, 3, l.GetActiveRuntimeCount()) + + require.NoError(t, router.runtimePool.Add(newMockRunningInvoke(t))) + + assert.Equal(t, 4, l.GetActiveRuntimeCount()) +} + +func TestAbortRunningInvokes_DrainsWithLargeResponseHoldTimeout(t *testing.T) { + t.Parallel() + + l := newTestLongInvokerRouter(t) + + pi := newPendingInvoke("test-version", 5*time.Second, 200*time.Millisecond) + l.pendingInvokes.Set("large-timeout-1", pi) + + l.wg.Add(1) + go func() { + defer l.wg.Done() + rec := &ResponseWriterRecorder{statusCode: http.StatusOK, body: []byte("result")} + l.sendInvokeResult("large-timeout-1", invokeResult{recorder: rec}) + }() + + var shutdownMetrics interop.MockShutdownMetrics + var durationMetric interop.MockDurationMetricTimer + shutdownMetrics.On("CreateDurationMetric", interop.ShutdownAbortInvokesDurationMetric).Return(&durationMetric) + durationMetric.On("Done").Return() + + start := time.Now() + l.AbortRunningInvokes(&shutdownMetrics, nil) + elapsed := time.Since(start) + + assert.GreaterOrEqual(t, elapsed, 150*time.Millisecond) + assert.Less(t, elapsed, 1*time.Second) + + _, exists := l.pendingInvokes.Get("large-timeout-1") + assert.False(t, exists) +} diff --git a/internal/lambda-managed-instances/invoke/metrics.go b/internal/lambda-managed-instances/invoke/metrics.go index cf856ddf..732af45b 100644 --- a/internal/lambda-managed-instances/invoke/metrics.go +++ b/internal/lambda-managed-instances/invoke/metrics.go @@ -7,6 +7,7 @@ import ( "encoding/json" "fmt" "log/slog" + "strconv" "time" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/interop" @@ -27,6 +28,22 @@ const ( RequestResponseModeDimension = "RequestMode" ResponseModeDimension = "ResponseMode" + InvokeModeDimension = "InvokeMode" + ResponseDeliveryDimension = "ResponseDelivery" + + PendingAgeMetric = "PendingAge" +) + +type InvokeMode string + +const ( + InvokeModeLong InvokeMode = "LongInvoke" + InvokeModeNormal InvokeMode = "Invoke" +) + +const ( + ResponseDeliverySent = "Sent" + ResponseDeliveryLost = "Lost" RequestSendDurationMetric = "RequestSendDuration" @@ -73,6 +90,13 @@ type invokeMetrics struct { timeSentResponse time.Time timeInvokeDone time.Time + responseWaitTime time.Duration + + responseDeliveryLost bool + responseDeliverySent bool + + invokeMode string + responseMetrics *interop.InvokeResponseMetrics requestPayloadBytes int64 @@ -108,8 +132,25 @@ func (e *invokeMetrics) AttachDependencies(initData interop.InitStaticDataProvid e.telemetryEventsAPI = telemetryEventsAPI } -func (e *invokeMetrics) TriggerGetRequest() { +func (e *invokeMetrics) TriggerGetRequest() time.Time { e.timeGetRequest = e.getCurrentTime() + return e.timeGetRequest +} + +func (e *invokeMetrics) SetResponseWaitTime(d time.Duration) { + e.responseWaitTime = d +} + +func (e *invokeMetrics) SetResponseDeliveryLost() { + e.responseDeliveryLost = true +} + +func (e *invokeMetrics) SetResponseDeliverySent() { + e.responseDeliverySent = true +} + +func (e *invokeMetrics) SetInvokeMode(mode string) { + e.invokeMode = mode } func (e *invokeMetrics) UpdateConcurrencyMetrics(inflightInvokes, idleRuntimesCount int) { @@ -275,6 +316,13 @@ func (e *invokeMetrics) buildProperties() []servicelogs.Property { Value: e.invokeReq.InvokeID(), }, ) + + if resolvedMs := e.invokeReq.ResolvedFunctionTimeoutMs(); resolvedMs > 0 { + props = append(props, servicelogs.Property{ + Name: InvokeTimeoutProperty, + Value: strconv.Itoa(int(time.Duration(resolvedMs) * time.Millisecond / time.Second)), + }) + } } return props @@ -290,6 +338,15 @@ func (e *invokeMetrics) buildDimensions() []servicelogs.Dimension { Value: e.invokeReq.ResponseMode(), }, ) + + invokeMode := InvokeModeNormal + if e.invokeMode != "" { + invokeMode = InvokeMode(e.invokeMode) + } + dim = append(dim, servicelogs.Dimension{ + Name: InvokeModeDimension, + Value: string(invokeMode), + }) } if e.responseMetrics != nil { @@ -299,6 +356,18 @@ func (e *invokeMetrics) buildDimensions() []servicelogs.Dimension { }) } + if e.responseDeliverySent { + dim = append(dim, servicelogs.Dimension{ + Name: ResponseDeliveryDimension, + Value: ResponseDeliverySent, + }) + } else if e.responseDeliveryLost { + dim = append(dim, servicelogs.Dimension{ + Name: ResponseDeliveryDimension, + Value: ResponseDeliveryLost, + }) + } + return dim } @@ -308,7 +377,7 @@ func (e *invokeMetrics) buildMetrics() []servicelogs.Metric { if !e.timeSentResponse.IsZero() { runDuration = e.timeSentResponse.Sub(e.timeStartRequest) } - platformOverhead := totalDuration - runDuration + platformOverhead := totalDuration - runDuration - e.responseWaitTime metrics := []servicelogs.Metric{ servicelogs.Timer(interop.TotalDurationMetric, totalDuration), @@ -317,6 +386,10 @@ func (e *invokeMetrics) buildMetrics() []servicelogs.Metric { servicelogs.Counter(IdleRuntimesCountMetric, float64(e.idleRuntimesCount)), } + if e.responseWaitTime > 0 { + metrics = append(metrics, servicelogs.Timer(PendingAgeMetric, e.responseWaitTime)) + } + if e.wasReserved { metrics = append(metrics, servicelogs.Counter(ReserveUsedMetric, 1), diff --git a/internal/lambda-managed-instances/invoke/metrics_test.go b/internal/lambda-managed-instances/invoke/metrics_test.go index 839a3e3b..6c9dd442 100644 --- a/internal/lambda-managed-instances/invoke/metrics_test.go +++ b/internal/lambda-managed-instances/invoke/metrics_test.go @@ -88,6 +88,8 @@ func createInvokeEventsMocks(t *testing.T) *invokeMetricsMocks { mocks.invokeReq.On("InvokeID").Return(eventInvokeId) mocks.invokeReq.On("ResponseMode").Return("Streaming").Maybe() + mocks.invokeReq.On("LongPollingConfig").Return((*interop.LongPollingConfig)(nil)).Maybe() + mocks.invokeReq.On("ResolvedFunctionTimeoutMs").Return(int64(0)).Maybe() return &mocks } @@ -375,6 +377,7 @@ func Test_invokeMetrics_ServiceLogs(t *testing.T) { }, expectedDims: []servicelogs.Dimension{ {Name: "RequestMode", Value: "Streaming"}, + {Name: "InvokeMode", Value: "Invoke"}, }, expectedMetrics: []servicelogs.Metric{ {Type: servicelogs.TimerType, Key: "TotalDuration", Value: 1000000}, @@ -404,6 +407,7 @@ func Test_invokeMetrics_ServiceLogs(t *testing.T) { }, expectedDims: []servicelogs.Dimension{ {Name: "RequestMode", Value: "Streaming"}, + {Name: "InvokeMode", Value: "Invoke"}, }, expectedMetrics: []servicelogs.Metric{ {Type: servicelogs.TimerType, Key: "TotalDuration", Value: 1000000}, @@ -433,6 +437,7 @@ func Test_invokeMetrics_ServiceLogs(t *testing.T) { }, expectedDims: []servicelogs.Dimension{ {Name: "RequestMode", Value: "Streaming"}, + {Name: "InvokeMode", Value: "Invoke"}, }, expectedMetrics: []servicelogs.Metric{ {Type: servicelogs.TimerType, Key: "TotalDuration", Value: 1000000}, @@ -468,6 +473,7 @@ func Test_invokeMetrics_ServiceLogs(t *testing.T) { }, expectedDims: []servicelogs.Dimension{ {Name: "RequestMode", Value: "Streaming"}, + {Name: "InvokeMode", Value: "Invoke"}, }, expectedMetrics: []servicelogs.Metric{ {Type: servicelogs.TimerType, Key: "TotalDuration", Value: 4000000}, @@ -509,6 +515,7 @@ func Test_invokeMetrics_ServiceLogs(t *testing.T) { }, expectedDims: []servicelogs.Dimension{ {Name: "RequestMode", Value: "Streaming"}, + {Name: "InvokeMode", Value: "Invoke"}, {Name: "ResponseMode", Value: "Streaming"}, }, expectedMetrics: []servicelogs.Metric{ @@ -558,6 +565,7 @@ func Test_invokeMetrics_ServiceLogs(t *testing.T) { }, expectedDims: []servicelogs.Dimension{ {Name: "RequestMode", Value: "Streaming"}, + {Name: "InvokeMode", Value: "Invoke"}, {Name: "ResponseMode", Value: "Streaming"}, }, expectedMetrics: []servicelogs.Metric{ @@ -607,6 +615,7 @@ func Test_invokeMetrics_ServiceLogs(t *testing.T) { }, expectedDims: []servicelogs.Dimension{ {Name: "RequestMode", Value: "Streaming"}, + {Name: "InvokeMode", Value: "Invoke"}, {Name: "ResponseMode", Value: "streaming"}, }, expectedMetrics: []servicelogs.Metric{ @@ -673,6 +682,8 @@ func Test_invokeMetrics_ServiceLogs(t *testing.T) { mocks.invokeReq.On("InvokeID").Return("invoke-id").Maybe() mocks.invokeReq.On("ResponseMode").Return("Streaming").Maybe().Maybe() mocks.invokeReq.On("TraceId").Return("Root=12345;Parent=67890;Sampled=1;Lineage=22222").Maybe() + mocks.invokeReq.On("ResolvedFunctionTimeoutMs").Return(int64(0)).Maybe() + mocks.invokeReq.On("LongPollingConfig").Return((*interop.LongPollingConfig)(nil)).Maybe() mocks.initData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough).Maybe() mocks.initData.On("MemorySizeMB").Return(uint64(128)).Maybe() diff --git a/internal/lambda-managed-instances/invoke/mock_responder_factory_func.go b/internal/lambda-managed-instances/invoke/mock_responder_factory_func.go index a1d020d7..b4ac563c 100644 --- a/internal/lambda-managed-instances/invoke/mock_responder_factory_func.go +++ b/internal/lambda-managed-instances/invoke/mock_responder_factory_func.go @@ -5,6 +5,7 @@ package invoke 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" @@ -14,16 +15,16 @@ type MockResponderFactoryFunc struct { mock.Mock } -func (_m *MockResponderFactoryFunc) Execute(_a0 context.Context, _a1 interop.InvokeRequest) InvokeResponseSender { - ret := _m.Called(_a0, _a1) +func (_m *MockResponderFactoryFunc) Execute(_a0 context.Context, _a1 interop.InvokeRequest, _a2 http.ResponseWriter) InvokeResponseSender { + ret := _m.Called(_a0, _a1, _a2) if len(ret) == 0 { panic("no return value specified for Execute") } var r0 InvokeResponseSender - if rf, ok := ret.Get(0).(func(context.Context, interop.InvokeRequest) InvokeResponseSender); ok { - r0 = rf(_a0, _a1) + if rf, ok := ret.Get(0).(func(context.Context, interop.InvokeRequest, http.ResponseWriter) InvokeResponseSender); ok { + r0 = rf(_a0, _a1, _a2) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(InvokeResponseSender) diff --git a/internal/lambda-managed-instances/invoke/mock_running_invoke.go b/internal/lambda-managed-instances/invoke/mock_running_invoke.go index 8b40260a..c1559883 100644 --- a/internal/lambda-managed-instances/invoke/mock_running_invoke.go +++ b/internal/lambda-managed-instances/invoke/mock_running_invoke.go @@ -20,16 +20,16 @@ func (_m *mockRunningInvoke) CancelAsync(_a0 model.AppError) { _m.Called(_a0) } -func (_m *mockRunningInvoke) RunInvokeAndSendResult(_a0 context.Context, _a1 interop.InitStaticDataProvider, _a2 interop.InvokeRequest, _a3 interop.InvokeMetrics) model.AppError { - ret := _m.Called(_a0, _a1, _a2, _a3) +func (_m *mockRunningInvoke) RunInvokeAndSendResult(_a0 context.Context, _a1 interop.InitStaticDataProvider, _a2 interop.InvokeRequest, _a3 interop.InvokeMetrics, _a4 InvokeResponseSender) model.AppError { + ret := _m.Called(_a0, _a1, _a2, _a3, _a4) if len(ret) == 0 { panic("no return value specified for RunInvokeAndSendResult") } var r0 model.AppError - if rf, ok := ret.Get(0).(func(context.Context, interop.InitStaticDataProvider, interop.InvokeRequest, interop.InvokeMetrics) model.AppError); ok { - r0 = rf(_a0, _a1, _a2, _a3) + if rf, ok := ret.Get(0).(func(context.Context, interop.InitStaticDataProvider, interop.InvokeRequest, interop.InvokeMetrics, InvokeResponseSender) model.AppError); ok { + r0 = rf(_a0, _a1, _a2, _a3, _a4) } else { if ret.Get(0) != nil { r0 = ret.Get(0).(model.AppError) diff --git a/internal/lambda-managed-instances/invoke/mock_runtime_error_request.go b/internal/lambda-managed-instances/invoke/mock_runtime_error_request.go index 38b39284..9efeb84c 100644 --- a/internal/lambda-managed-instances/invoke/mock_runtime_error_request.go +++ b/internal/lambda-managed-instances/invoke/mock_runtime_error_request.go @@ -120,6 +120,25 @@ func (_m *MockRuntimeErrorRequest) GetXrayErrorCause() json.RawMessage { return r0 } +func (_m *MockRuntimeErrorRequest) InvocationID() *string { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for InvocationID") + } + + var r0 *string + if rf, ok := ret.Get(0).(func() *string); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*string) + } + } + + return r0 +} + func (_m *MockRuntimeErrorRequest) InvokeID() string { ret := _m.Called() @@ -171,23 +190,6 @@ func (_m *MockRuntimeErrorRequest) ReturnCode() int { return r0 } -func (_m *MockRuntimeErrorRequest) InvocationID() string { - ret := _m.Called() - - if len(ret) == 0 { - panic("no return value specified for InvocationID") - } - - var r0 string - if rf, ok := ret.Get(0).(func() string); ok { - r0 = rf() - } else { - r0 = ret.Get(0).(string) - } - - return r0 -} - func NewMockRuntimeErrorRequest(t interface { mock.TestingT Cleanup(func()) diff --git a/internal/lambda-managed-instances/invoke/mock_runtime_response_request.go b/internal/lambda-managed-instances/invoke/mock_runtime_response_request.go index bc550bc8..850d21da 100644 --- a/internal/lambda-managed-instances/invoke/mock_runtime_response_request.go +++ b/internal/lambda-managed-instances/invoke/mock_runtime_response_request.go @@ -33,6 +33,10 @@ func (_m *MockRuntimeResponseRequest) BodyReader() io.Reader { return r0 } +func (_m *MockRuntimeResponseRequest) Cancel() { + _m.Called() +} + func (_m *MockRuntimeResponseRequest) ContentType() string { ret := _m.Called() @@ -50,6 +54,25 @@ func (_m *MockRuntimeResponseRequest) ContentType() string { return r0 } +func (_m *MockRuntimeResponseRequest) InvocationID() *string { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for InvocationID") + } + + var r0 *string + if rf, ok := ret.Get(0).(func() *string); ok { + r0 = rf() + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*string) + } + } + + return r0 +} + func (_m *MockRuntimeResponseRequest) InvokeID() string { ret := _m.Called() @@ -122,23 +145,6 @@ func (_m *MockRuntimeResponseRequest) TrailerError() ErrorForInvoker { return r0 } -func (_m *MockRuntimeResponseRequest) InvocationID() string { - ret := _m.Called() - - if len(ret) == 0 { - panic("no return value specified for InvocationID") - } - - var r0 string - if rf, ok := ret.Get(0).(func() string); ok { - r0 = rf() - } else { - r0 = ret.Get(0).(string) - } - - return r0 -} - func NewMockRuntimeResponseRequest(t interface { mock.TestingT Cleanup(func()) diff --git a/internal/lambda-managed-instances/invoke/reconnect_metrics.go b/internal/lambda-managed-instances/invoke/reconnect_metrics.go new file mode 100644 index 00000000..ef0e8b5d --- /dev/null +++ b/internal/lambda-managed-instances/invoke/reconnect_metrics.go @@ -0,0 +1,125 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 + +package invoke + +import ( + "time" + + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/interop" + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/servicelogs" +) + +type ReconnectMetrics struct { + startTime time.Time + invokeID interop.InvokeID + reconnectRequestID string + logger servicelogs.Logger + functionDoneTime time.Time + connectionGap time.Duration + responseReplayStart time.Time + responseReplayDuration time.Duration + responseReplaySize int + pollStartTime time.Time + pollEndTime time.Time +} + +func NewReconnectMetrics(invokeID interop.InvokeID, reconnectID string, logger servicelogs.Logger) *ReconnectMetrics { + return &ReconnectMetrics{startTime: time.Now(), invokeID: invokeID, reconnectRequestID: reconnectID, logger: logger} +} + +func (m *ReconnectMetrics) TriggerConnectionGap(lastDisconnectTime *time.Time) { + if lastDisconnectTime != nil { + m.connectionGap = time.Since(*lastDisconnectTime) + } +} + +func (m *ReconnectMetrics) TriggerResponseReplayStart() { + m.responseReplayStart = time.Now() +} + +func (m *ReconnectMetrics) TriggerResponseReplayDone(size int) { + m.responseReplayDuration = time.Since(m.responseReplayStart) + m.responseReplaySize = size +} + +func (m *ReconnectMetrics) SetFunctionDoneTime(t time.Time) { m.functionDoneTime = t } + +func (m *ReconnectMetrics) TriggerPollStart() { m.pollStartTime = time.Now() } + +func (m *ReconnectMetrics) TriggerPollEnd() { m.pollEndTime = time.Now() } + +type noopReconnectMetrics struct{} + +func (noopReconnectMetrics) TriggerConnectionGap(*time.Time) {} +func (noopReconnectMetrics) TriggerPollStart() {} +func (noopReconnectMetrics) TriggerPollEnd() {} +func (noopReconnectMetrics) TriggerResponseReplayStart() {} +func (noopReconnectMetrics) TriggerResponseReplayDone(int) {} +func (noopReconnectMetrics) SetFunctionDoneTime(time.Time) {} + +func NoopReconnectMetrics() interop.ReconnectMetrics { return noopReconnectMetrics{} } + +func (m *ReconnectMetrics) SendMetrics(outcome interop.ReconnectOutcome) { + totalDuration := time.Since(m.startTime) + + props := []servicelogs.Property{ + {Name: interop.RequestIdProperty, Value: string(m.invokeID)}, + {Name: "Outcome", Value: string(outcome)}, + } + if m.reconnectRequestID != "" { + props = append(props, servicelogs.Property{Name: "reconnectId", Value: m.reconnectRequestID}) + } + + metrics := []servicelogs.Metric{ + servicelogs.Timer(interop.TotalDurationMetric, totalDuration), + } + if m.connectionGap > 0 { + metrics = append(metrics, servicelogs.Timer("ConnectionGap", m.connectionGap)) + } + + if !m.functionDoneTime.IsZero() { + pendingAge := m.startTime.Sub(m.functionDoneTime) + if pendingAge > 0 { + metrics = append(metrics, servicelogs.Timer("PendingAge", pendingAge)) + } + } + if m.responseReplayDuration > 0 { + metrics = append(metrics, servicelogs.Timer("ResponseReplayDuration", m.responseReplayDuration)) + metrics = append(metrics, servicelogs.Counter("ResponseReplaySizeBytes", float64(m.responseReplaySize))) + bytesPerSec := float64(m.responseReplaySize) / m.responseReplayDuration.Seconds() + metrics = append(metrics, servicelogs.Counter("ResponseReplaySpeedBytesPerSec", bytesPerSec)) + } + + var overhead time.Duration + if !m.pollStartTime.IsZero() && !m.pollEndTime.IsZero() { + pollDuration := m.pollEndTime.Sub(m.pollStartTime) + overhead = totalDuration - pollDuration + } else { + + overhead = totalDuration + } + metrics = append(metrics, servicelogs.Timer("ReconnectOverhead", overhead)) + + var platformErrCnt, clientErrCnt float64 + switch outcome { + case interop.ReconnectOutcomeError: + platformErrCnt = 1 + case interop.ReconnectOutcomeNotFound: + clientErrCnt = 1 + } + metrics = append(metrics, + servicelogs.Counter(interop.PlatformErrorMetric, platformErrCnt), + servicelogs.Counter(interop.ClientErrorMetric, clientErrCnt), + ) + + m.logger.Log(servicelogs.ReconnectOp, m.startTime, props, nil, metrics) +} + +func SendInvokePendingMetrics(logger servicelogs.Logger, invokeID interop.InvokeID, invokeStart time.Time) { + logger.Log(servicelogs.InvokePendingOp, invokeStart, + []servicelogs.Property{{Name: interop.RequestIdProperty, Value: string(invokeID)}}, + nil, + []servicelogs.Metric{servicelogs.Timer(interop.TotalDurationMetric, time.Since(invokeStart))}, + ) +} diff --git a/internal/lambda-managed-instances/invoke/reconnect_metrics_test.go b/internal/lambda-managed-instances/invoke/reconnect_metrics_test.go new file mode 100644 index 00000000..4879464d --- /dev/null +++ b/internal/lambda-managed-instances/invoke/reconnect_metrics_test.go @@ -0,0 +1,174 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 + +package invoke + +import ( + "testing" + "time" + + "github.com/stretchr/testify/assert" + + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/interop" + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/servicelogs" +) + +type capturedLog struct { + op servicelogs.Operation + opStart time.Time + props []servicelogs.Property + dims []servicelogs.Dimension + metrics []servicelogs.Metric +} + +type capturingLogger struct { + logs []capturedLog +} + +func (l *capturingLogger) Log(op servicelogs.Operation, opStart time.Time, props []servicelogs.Property, dims []servicelogs.Dimension, metrics []servicelogs.Metric) { + l.logs = append(l.logs, capturedLog{op: op, opStart: opStart, props: props, dims: dims, metrics: metrics}) +} +func (l *capturingLogger) Close() error { return nil } + +func findMetric(metrics []servicelogs.Metric, key string) *servicelogs.Metric { + for i := range metrics { + if metrics[i].Key == key { + return &metrics[i] + } + } + return nil +} + +func findProp(props []servicelogs.Property, name string) string { + for _, p := range props { + if p.Name == name { + return p.Value + } + } + return "" +} + +func TestReconnectMetrics_SendMetrics_BasicFields(t *testing.T) { + logger := &capturingLogger{} + m := NewReconnectMetrics("invoke-1", "reconnect-1", logger) + + m.SendMetrics(interop.ReconnectOutcomeCompleted) + + assert.Len(t, logger.logs, 1) + log := logger.logs[0] + assert.Equal(t, servicelogs.ReconnectOp, log.op) + assert.Equal(t, "invoke-1", findProp(log.props, interop.RequestIdProperty)) + assert.Equal(t, "reconnect-1", findProp(log.props, "reconnectId")) + assert.NotNil(t, findMetric(log.metrics, interop.TotalDurationMetric)) + assert.Equal(t, "completed", findProp(log.props, "Outcome")) + assert.NotNil(t, findMetric(log.metrics, "ReconnectOverhead")) + + assert.Equal(t, float64(0), findMetric(log.metrics, interop.PlatformErrorMetric).Value) + assert.Equal(t, float64(0), findMetric(log.metrics, interop.ClientErrorMetric).Value) +} + +func TestReconnectMetrics_SendMetrics_WithConnectionGap(t *testing.T) { + logger := &capturingLogger{} + m := NewReconnectMetrics("invoke-1", "", logger) + past := time.Now().Add(-500 * time.Millisecond) + m.TriggerConnectionGap(&past) + + m.SendMetrics(interop.ReconnectOutcomeCompleted) + + log := logger.logs[0] + assert.Equal(t, "", findProp(log.props, "reconnectId")) + gap := findMetric(log.metrics, "ConnectionGap") + assert.NotNil(t, gap) + assert.Greater(t, gap.Value, float64(0)) +} + +func TestReconnectMetrics_SendMetrics_WithResponseReplay(t *testing.T) { + logger := &capturingLogger{} + m := NewReconnectMetrics("invoke-1", "r-1", logger) + m.TriggerResponseReplayStart() + time.Sleep(10 * time.Millisecond) + m.TriggerResponseReplayDone(4096) + + m.SendMetrics(interop.ReconnectOutcomeCompleted) + + log := logger.logs[0] + assert.NotNil(t, findMetric(log.metrics, "ResponseReplayDuration")) + size := findMetric(log.metrics, "ResponseReplaySizeBytes") + assert.NotNil(t, size) + assert.Equal(t, float64(4096), size.Value) + speed := findMetric(log.metrics, "ResponseReplaySpeedBytesPerSec") + assert.NotNil(t, speed) + assert.Greater(t, speed.Value, float64(0)) +} + +func TestReconnectMetrics_SendMetrics_PendingAge(t *testing.T) { + logger := &capturingLogger{} + m := &ReconnectMetrics{ + startTime: time.Now(), + invokeID: "invoke-1", + logger: logger, + + functionDoneTime: time.Now().Add(-200 * time.Millisecond), + } + + m.SendMetrics(interop.ReconnectOutcomeCompleted) + + log := logger.logs[0] + pending := findMetric(log.metrics, "PendingAge") + assert.NotNil(t, pending) + assert.Greater(t, pending.Value, float64(0)) +} + +func TestReconnectMetrics_SendMetrics_TimeoutOutcome(t *testing.T) { + logger := &capturingLogger{} + m := NewReconnectMetrics("invoke-1", "", logger) + m.TriggerPollStart() + time.Sleep(1 * time.Millisecond) + m.TriggerPollEnd() + + m.SendMetrics(interop.ReconnectOutcomeTimeout) + + log := logger.logs[0] + assert.Equal(t, "timeout", findProp(log.props, "Outcome")) + + overhead := findMetric(log.metrics, "ReconnectOverhead") + assert.NotNil(t, overhead) + assert.Less(t, overhead.Value, float64(10000)) +} + +func TestReconnectMetrics_SendMetrics_ErrorOutcome_PlatformError(t *testing.T) { + logger := &capturingLogger{} + m := NewReconnectMetrics("invoke-1", "", logger) + + m.SendMetrics(interop.ReconnectOutcomeError) + + log := logger.logs[0] + assert.Equal(t, float64(1), findMetric(log.metrics, interop.PlatformErrorMetric).Value) + assert.Equal(t, float64(0), findMetric(log.metrics, interop.ClientErrorMetric).Value) +} + +func TestReconnectMetrics_SendMetrics_NotFoundOutcome_ClientError(t *testing.T) { + logger := &capturingLogger{} + m := NewReconnectMetrics("invoke-1", "", logger) + + m.SendMetrics(interop.ReconnectOutcomeNotFound) + + log := logger.logs[0] + assert.Equal(t, float64(0), findMetric(log.metrics, interop.PlatformErrorMetric).Value) + assert.Equal(t, float64(1), findMetric(log.metrics, interop.ClientErrorMetric).Value) +} + +func TestSendInvokePendingMetrics(t *testing.T) { + logger := &capturingLogger{} + invokeStart := time.Now().Add(-10 * time.Second) + + SendInvokePendingMetrics(logger, "invoke-1", invokeStart) + + assert.Len(t, logger.logs, 1) + log := logger.logs[0] + assert.Equal(t, servicelogs.InvokePendingOp, log.op) + assert.Equal(t, "invoke-1", findProp(log.props, interop.RequestIdProperty)) + dur := findMetric(log.metrics, interop.TotalDurationMetric) + assert.NotNil(t, dur) + assert.Greater(t, dur.Value, float64(0)) +} diff --git a/internal/lambda-managed-instances/invoke/response_writer_recorder.go b/internal/lambda-managed-instances/invoke/response_writer_recorder.go new file mode 100644 index 00000000..c5aa5a1d --- /dev/null +++ b/internal/lambda-managed-instances/invoke/response_writer_recorder.go @@ -0,0 +1,68 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 + +package invoke + +import ( + "log/slog" + "net/http" + "sync/atomic" +) + +type ResponseWriterRecorder struct { + statusCode int + header http.Header + trailer http.Header + body []byte + headersSent atomic.Bool +} + +func (r *ResponseWriterRecorder) Header() http.Header { + if r.headersSent.Load() { + if r.trailer == nil { + r.trailer = make(http.Header) + } + return r.trailer + } + if r.header == nil { + r.header = make(http.Header) + } + return r.header +} + +func (r *ResponseWriterRecorder) Write(bytes []byte) (int, error) { + r.headersSent.Store(true) + r.body = append(r.body, bytes...) + return len(bytes), nil +} + +func (r *ResponseWriterRecorder) WriteHeader(statusCode int) { + r.statusCode = statusCode + r.headersSent.Store(true) +} + +func (r *ResponseWriterRecorder) Flush() {} + +func (r *ResponseWriterRecorder) BodySize() int { return len(r.body) } + +func (r *ResponseWriterRecorder) WriteTo(w http.ResponseWriter) error { + for k, vs := range r.header { + for _, v := range vs { + w.Header().Add(k, v) + } + } + if r.statusCode != 0 { + w.WriteHeader(r.statusCode) + } + + if _, err := w.Write(r.body); err != nil { + slog.Error("ResponseWriterRecorder: failed to write body", "err", err) + return err + } + for k, vs := range r.trailer { + for _, v := range vs { + w.Header().Add(k, v) + } + } + return nil +} diff --git a/internal/lambda-managed-instances/invoke/response_writer_recorder_test.go b/internal/lambda-managed-instances/invoke/response_writer_recorder_test.go new file mode 100644 index 00000000..2eb2c3c1 --- /dev/null +++ b/internal/lambda-managed-instances/invoke/response_writer_recorder_test.go @@ -0,0 +1,182 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 + +package invoke + +import ( + "errors" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestResponseWriterRecorder_HeadersBeforeWrite(t *testing.T) { + t.Parallel() + rec := &ResponseWriterRecorder{} + + rec.Header().Set("Content-Type", "text/plain") + rec.WriteHeader(http.StatusOK) + _, err := rec.Write([]byte("body")) + require.NoError(t, err) + + assert.Equal(t, "text/plain", rec.header.Get("Content-Type")) + assert.Equal(t, http.StatusOK, rec.statusCode) + assert.Equal(t, []byte("body"), rec.body) +} + +func TestResponseWriterRecorder_TrailersAfterWrite(t *testing.T) { + t.Parallel() + rec := &ResponseWriterRecorder{} + + rec.Header().Set("Content-Type", "text/plain") + rec.WriteHeader(http.StatusOK) + _, err := rec.Write([]byte("body")) + require.NoError(t, err) + + rec.Header().Set("End-Of-Response", "Complete") + + assert.Equal(t, "text/plain", rec.header.Get("Content-Type")) + assert.Equal(t, "", rec.header.Get("End-Of-Response")) + assert.Equal(t, "Complete", rec.trailer.Get("End-Of-Response")) +} + +func TestResponseWriterRecorder_WriteTo_ReplaysAll(t *testing.T) { + t.Parallel() + rec := &ResponseWriterRecorder{} + + rec.Header().Set("X-Header", "val") + rec.WriteHeader(http.StatusOK) + _, err := rec.Write([]byte("body")) + require.NoError(t, err) + rec.Header().Set("X-Trailer", "tval") + + w := httptest.NewRecorder() + _ = rec.WriteTo(w) + + assert.Equal(t, http.StatusOK, w.Code) + assert.Equal(t, "val", w.Header().Get("X-Header")) + assert.Equal(t, "body", w.Body.String()) + assert.Equal(t, "tval", w.Header().Get("X-Trailer")) +} + +func TestResponseWriterRecorder_WriteTo_EmptyRecorder(t *testing.T) { + t.Parallel() + rec := &ResponseWriterRecorder{} + w := httptest.NewRecorder() + + _ = rec.WriteTo(w) + + assert.Equal(t, http.StatusOK, w.Code) + assert.Empty(t, w.Body.String()) +} + +func TestResponseWriterRecorder_WriteTo_WriteError(t *testing.T) { + t.Parallel() + rec := &ResponseWriterRecorder{} + rec.WriteHeader(http.StatusOK) + _, _ = rec.Write([]byte("payload")) + + w := &failingWriter{header: make(http.Header)} + err := rec.WriteTo(w) + assert.Error(t, err) +} + +func TestResponseWriterRecorder_Flush_NoOp(t *testing.T) { + t.Parallel() + rec := &ResponseWriterRecorder{} + rec.Flush() +} + +func TestResponseWriterRecorder_WriteHeaderSwitchesToTrailerMap(t *testing.T) { + t.Parallel() + rec := &ResponseWriterRecorder{} + + rec.Header().Set("Invoke-Id", "abc-123") + rec.WriteHeader(http.StatusOK) + rec.Header().Set("End-Of-Response", "Complete") + rec.Header().Set("Error-Category", "RuntimeError") + + assert.Equal(t, "abc-123", rec.header.Get("Invoke-Id")) + + assert.Equal(t, "", rec.header.Get("End-Of-Response")) + assert.Equal(t, "", rec.header.Get("Error-Category")) + assert.Equal(t, "Complete", rec.trailer.Get("End-Of-Response")) + assert.Equal(t, "RuntimeError", rec.trailer.Get("Error-Category")) +} + +func TestResponseWriterRecorder_WriteTo_ErrorWithoutBody(t *testing.T) { + t.Parallel() + rec := &ResponseWriterRecorder{} + + rec.Header().Set("Content-Type", "text/plain") + rec.Header().Set("Invoke-Id", "abc-123") + rec.WriteHeader(http.StatusOK) + + rec.Header().Set("End-Of-Response", "Complete") + rec.Header().Set("Error-Category", "RuntimeError") + + w := httptest.NewRecorder() + _ = rec.WriteTo(w) + + assert.Equal(t, http.StatusOK, w.Code) + assert.Equal(t, "text/plain", w.Header().Get("Content-Type")) + assert.Equal(t, "abc-123", w.Header().Get("Invoke-Id")) + + assert.Equal(t, "Complete", w.Header().Get("End-Of-Response")) + assert.Equal(t, "RuntimeError", w.Header().Get("Error-Category")) +} + +func TestResponseWriterRecorder_WriteTo_EmptyBodyTriggersWrite(t *testing.T) { + t.Parallel() + rec := &ResponseWriterRecorder{} + + rec.Header().Set("Invoke-Id", "abc-123") + rec.WriteHeader(http.StatusOK) + + spy := &spyWriter{ResponseRecorder: httptest.NewRecorder()} + _ = rec.WriteTo(spy) + + assert.True(t, spy.writeCalled, "Write should have been called even with empty body") + assert.Equal(t, http.StatusOK, spy.Code) +} + +type spyWriter struct { + *httptest.ResponseRecorder + writeCalled bool +} + +func (s *spyWriter) Write(b []byte) (int, error) { + s.writeCalled = true + return s.ResponseRecorder.Write(b) +} + +func TestResponseWriterRecorder_WriteTo_SuccessWithBody(t *testing.T) { + t.Parallel() + rec := &ResponseWriterRecorder{} + + rec.Header().Set("Content-Type", "application/json") + rec.Header().Set("Invoke-Id", "abc-123") + rec.WriteHeader(http.StatusOK) + _, err := rec.Write([]byte(`{"result":"ok"}`)) + require.NoError(t, err) + rec.Header().Set("End-Of-Response", "Complete") + + w := httptest.NewRecorder() + _ = rec.WriteTo(w) + + assert.Equal(t, http.StatusOK, w.Code) + assert.Equal(t, "application/json", w.Header().Get("Content-Type")) + assert.Equal(t, `{"result":"ok"}`, w.Body.String()) + assert.Equal(t, "Complete", w.Header().Get("End-Of-Response")) +} + +type failingWriter struct { + header http.Header +} + +func (f *failingWriter) Header() http.Header { return f.header } +func (f *failingWriter) WriteHeader(_ int) {} +func (f *failingWriter) Write(_ []byte) (int, error) { return 0, errors.New("connection reset") } diff --git a/internal/lambda-managed-instances/invoke/running_invoke.go b/internal/lambda-managed-instances/invoke/running_invoke.go index e08b12c7..12aad5a3 100644 --- a/internal/lambda-managed-instances/invoke/running_invoke.go +++ b/internal/lambda-managed-instances/invoke/running_invoke.go @@ -45,7 +45,7 @@ type InvokeResponseSender interface { ErrorPayloadSizeBytes() int } -type ResponderFactoryFunc func(context.Context, interop.InvokeRequest) InvokeResponseSender +type ResponderFactoryFunc func(context.Context, interop.InvokeRequest, http.ResponseWriter) InvokeResponseSender type SendResponseBodyResult struct { Metrics interop.InvokeResponseMetrics @@ -70,14 +70,12 @@ type runningInvokeImpl struct { internalInvocationID string - responderFactoryFunc ResponderFactoryFunc - sendInvokeToRuntime func(context.Context, interop.InitStaticDataProvider, interop.InvokeRequest, http.ResponseWriter, string) (int64, time.Duration, time.Duration, model.AppError) - createTracingData func(traceId string, tracingMode intmodel.XrayTracingMode, segmentIDGenerator func() string) (downstreamTraceId string, tracingCtx *interop.TracingCtx) + sendInvokeToRuntime func(context.Context, interop.InitStaticDataProvider, interop.InvokeRequest, http.ResponseWriter, string) (int64, time.Duration, time.Duration, model.AppError) + createTracingData func(traceId string, tracingMode intmodel.XrayTracingMode, segmentIDGenerator func() string) (downstreamTraceId string, tracingCtx *interop.TracingCtx) } func newRunningInvoke( runtimeNext http.ResponseWriter, - responderFactoryFunc ResponderFactoryFunc, timeoutCache timeoutCache, ) runningInvokeImpl { ctx, cancel := context.WithCancelCause(context.Background()) @@ -92,13 +90,12 @@ func newRunningInvoke( runtimeErrorChan: make(chan RuntimeErrorRequest, 1), runtimeNext: runtimeNext, - responderFactoryFunc: responderFactoryFunc, - sendInvokeToRuntime: sendInvokeToRuntime, - createTracingData: xray.CreateTracingData, + sendInvokeToRuntime: sendInvokeToRuntime, + createTracingData: xray.CreateTracingData, } } -func (r *runningInvokeImpl) RunInvokeAndSendResult(ctx context.Context, initData interop.InitStaticDataProvider, invokeReq interop.InvokeRequest, metrics interop.InvokeMetrics) model.AppError { +func (r *runningInvokeImpl) RunInvokeAndSendResult(ctx context.Context, initData interop.InitStaticDataProvider, invokeReq interop.InvokeRequest, metrics interop.InvokeMetrics, sender InvokeResponseSender) model.AppError { downstreamTraceId, tracingCtx := r.createTracingData(invokeReq.TraceId(), initData.XRayTracingMode(), xray.GenerateSegmentID) r.internalInvocationID = invokeReq.InternalInvocationID() @@ -108,9 +105,9 @@ func (r *runningInvokeImpl) RunInvokeAndSendResult(ctx context.Context, initData logging.Error(ctx, "Failed to send InvokeStartEvent", "err", err) } - r.invokeRespSender = r.responderFactoryFunc(ctx, invokeReq) + r.invokeRespSender = sender - ctx, cancel := r.getInvokeCtx(ctx, initData.FunctionTimeout()) + ctx, cancel := r.getInvokeCtx(ctx, invokeReq.ResolvedInvokeTimeout()) defer cancel() logging.Debug(ctx, "Sending Invoke to Runtime") @@ -147,7 +144,7 @@ func (r *runningInvokeImpl) RunInvokeAndSendResult(ctx context.Context, initData r.invokeRespSender.SendError(err, initData) break } - runtimeAnswerSent, invokeResponseMetrics, err, xrayErrorCause = r.sendRuntimeResponse(ctx, initData, runtimeResp) + runtimeAnswerSent, invokeResponseMetrics, err, xrayErrorCause = r.sendRuntimeResponse(ctx, initData, runtimeResp, invokeReq.ResolvedInvokeTimeout()) case runtimeErr := <-r.runtimeErrorChan: metrics.TriggerGetResponse() err = runtimeErr.GetError() @@ -158,7 +155,7 @@ func (r *runningInvokeImpl) RunInvokeAndSendResult(ctx context.Context, initData case <-ctx.Done(): r.responseState.Store(stateGotError) - err = BuildInvokeAppError(context.Cause(ctx), initData.FunctionTimeout()) + err = BuildInvokeAppError(context.Cause(ctx), invokeReq.ResolvedInvokeTimeout()) logging.Info(ctx, "Received ctx cancellation", "err", err) if err.ErrorType() == model.ErrorSandboxTimedout { @@ -198,7 +195,7 @@ func (r *runningInvokeImpl) getInvokeCtx(ctx context.Context, timeout time.Durat } } -func (r *runningInvokeImpl) sendRuntimeResponse(ctx context.Context, initData interop.InitStaticDataProvider, runtimeResp RuntimeResponseRequest) (bool, *interop.InvokeResponseMetrics, model.AppError, json.RawMessage) { +func (r *runningInvokeImpl) sendRuntimeResponse(ctx context.Context, initData interop.InitStaticDataProvider, runtimeResp RuntimeResponseRequest, resolvedTimeout time.Duration) (bool, *interop.InvokeResponseMetrics, model.AppError, json.RawMessage) { logging.Debug(ctx, "Sending Runtime response headers") r.invokeRespSender.SendRuntimeResponseHeaders(initData, runtimeResp.ContentType(), runtimeResp.ResponseMode()) @@ -208,11 +205,10 @@ func (r *runningInvokeImpl) sendRuntimeResponse(ctx context.Context, initData in sendBodyRes := make(chan SendResponseBodyResult) logging.Debug(ctx, "Sending Runtime response body") go func() { - sendBodyRes <- r.invokeRespSender.SendRuntimeResponseBody(childCtx, runtimeResp, initData.FunctionTimeout()) + sendBodyRes <- r.invokeRespSender.SendRuntimeResponseBody(childCtx, runtimeResp, resolvedTimeout) }() select { - case res := <-sendBodyRes: if res.Err != nil { logging.Err(ctx, "Failed sending body", res.Err) @@ -294,10 +290,10 @@ func (r *runningInvokeImpl) RuntimeResponse(ctx context.Context, runtimeRespReq return model.NewCustomerError(model.ErrorRuntimeInvokeResponseInProgress) } - if echoedID := runtimeRespReq.InvocationID(); echoedID != "" && r.internalInvocationID != "" { - if echoedID != r.internalInvocationID { + if echoedID := runtimeRespReq.InvocationID(); echoedID != nil && r.internalInvocationID != "" { + if *echoedID != r.internalInvocationID { logging.Warn(ctx, "Cross-wiring detected: invocation ID mismatch on response", - "expected", r.internalInvocationID, "received", echoedID) + "expected", r.internalInvocationID, "received", *echoedID) r.responseState.CompareAndSwap(stateGotResponse, stateNoResponse) return model.NewCustomerError(model.ErrorRuntimeInvokeTimeout) } @@ -308,10 +304,11 @@ func (r *runningInvokeImpl) RuntimeResponse(ctx context.Context, runtimeRespReq } func (r *runningInvokeImpl) RuntimeError(ctx context.Context, runtimeErrReq RuntimeErrorRequest) model.AppError { - if echoedID := runtimeErrReq.InvocationID(); echoedID != "" && r.internalInvocationID != "" { - if echoedID != r.internalInvocationID { + + if echoedID := runtimeErrReq.InvocationID(); echoedID != nil && r.internalInvocationID != "" { + if *echoedID != r.internalInvocationID { logging.Warn(ctx, "Cross-wiring detected: invocation ID mismatch on error", - "expected", r.internalInvocationID, "received", echoedID) + "expected", r.internalInvocationID, "received", *echoedID) return model.NewCustomerError(model.ErrorRuntimeInvokeTimeout) } } diff --git a/internal/lambda-managed-instances/invoke/running_invoke_test.go b/internal/lambda-managed-instances/invoke/running_invoke_test.go index 4f69ae68..cf68771d 100644 --- a/internal/lambda-managed-instances/invoke/running_invoke_test.go +++ b/internal/lambda-managed-instances/invoke/running_invoke_test.go @@ -66,15 +66,10 @@ func hijackRunningInvokeDeps(ri *runningInvokeImpl, mocks *runningInvokeMocks) { func createMocksAndInitRunningInvoke(t *testing.T) (*runningInvokeMocks, *runningInvokeImpl) { mocks := newRunningInvokeMocks(t) - mocks.eaInvokeRequest.On("InternalInvocationID").Return("").Maybe() - mocks.runtimeRespReq.On("InvocationID").Return("").Maybe() - mocks.runtimeErrorReq.On("InvocationID").Return("").Maybe() + mocks.eaInvokeRequest.On("InternalInvocationID").Return("") ri := newRunningInvoke( mocks.runtimeNextRequest, - func(ctx context.Context, ir interop.InvokeRequest) InvokeResponseSender { - return &mocks.eaInvokeResponder - }, mocks.timeoutCache, ) hijackRunningInvokeDeps(&ri, &mocks) @@ -120,12 +115,13 @@ func TestRunInvokeAndSendResultSuccess_RuntimeResponse(t *testing.T) { mocks, runInvoke := createMocksAndInitRunningInvoke(t) - mocks.staticData.On("FunctionTimeout").Return(time.Second) mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Second) mocks.runtimeRespReq.On("ParsingError").Return(nil) + mocks.runtimeRespReq.On("InvocationID").Return((*string)(nil)) mocks.runtimeRespReq.On("ContentType").Return("") mocks.runtimeRespReq.On("ResponseMode").Return("") mocks.runtimeRespReq.On("TrailerError").Return(nil) @@ -145,7 +141,7 @@ func TestRunInvokeAndSendResultSuccess_RuntimeResponse(t *testing.T) { assert.NoError(t, err) }() - err := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics) + err := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) assert.NoError(t, err) wg.Wait() @@ -158,12 +154,13 @@ func TestRunInvokeAndSendResultSuccess_RuntimeError(t *testing.T) { mocks, runInvoke := createMocksAndInitRunningInvoke(t) err := model.NewCustomerError(model.ErrorFunctionUnknown) - mocks.staticData.On("FunctionTimeout").Return(time.Second) mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Second) mocks.errorPayloadSizeBytes = 100 + mocks.runtimeErrorReq.On("InvocationID").Return((*string)(nil)) mocks.runtimeErrorReq.On("GetError").Return(err) mocks.runtimeErrorReq.On("GetXrayErrorCause").Return(json.RawMessage(nil)) mocks.eaInvokeResponder.On("SendRuntimeResponseHeaders", &mocks.staticData, mock.Anything, mock.Anything).Return().Once() @@ -182,7 +179,7 @@ func TestRunInvokeAndSendResultSuccess_RuntimeError(t *testing.T) { assert.NoError(t, err) }() - invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics) + invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) assert.Error(t, invokeErr) wg.Wait() @@ -197,15 +194,16 @@ func TestRunInvokeAndSendResultSuccess_RuntimeTrailerError(t *testing.T) { trailerErrorType := model.ErrorType("Function.Unknown") expectedTrailerErr := model.NewCustomerError(trailerErrorType) - mocks.staticData.On("FunctionTimeout").Return(time.Second) mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Second) mocks.runtimeRespReq.On("ParsingError").Return(nil) + mocks.runtimeRespReq.On("InvocationID").Return((*string)(nil)) mocks.runtimeRespReq.On("ContentType").Return("") mocks.runtimeRespReq.On("ResponseMode").Return("") - mocks.runtimeRespReq.On("TrailerError").Return(expectedTrailerErr, NewMockErrorForInvoker(t)) + mocks.runtimeRespReq.On("TrailerError").Return(expectedTrailerErr) mocks.eaInvokeResponder.On("SendRuntimeResponseHeaders", &mocks.staticData, mock.Anything, mock.Anything).Return() mocks.eaInvokeResponder.On("SendRuntimeResponseBody", mock.Anything, &mocks.runtimeRespReq, mock.Anything).Return(SendResponseBodyResult{}) @@ -224,7 +222,7 @@ func TestRunInvokeAndSendResultSuccess_RuntimeTrailerError(t *testing.T) { assert.Equal(t, trailerErrorType, err.ErrorType()) }() - err := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics) + err := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) assert.Error(t, err) assert.Equal(t, trailerErrorType, err.ErrorType()) @@ -242,11 +240,11 @@ func TestRuntimeErrorFailure_SendInvokeToRuntime_Error(t *testing.T) { return 0, 0, 0, err } - mocks.staticData.On("FunctionTimeout").Return(time.Second) mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") - mocks.eaInvokeResponder.On("SendError", err, &mocks.staticData, mock.Anything).Return() + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Second) + mocks.eaInvokeResponder.On("SendError", err, &mocks.staticData).Return() mocks.eaInvokeResponder.On("ErrorPayloadSizeBytes").Return(mocks.errorPayloadSizeBytes) mocks.metrics.On("TriggerStartRequest") @@ -254,7 +252,7 @@ func TestRuntimeErrorFailure_SendInvokeToRuntime_Error(t *testing.T) { mocks.metrics.On("TriggerSentResponse", false, err, mock.Anything, 0).Return() mocks.metrics.On("SendInvokeFinishedEvent", mock.AnythingOfType("*interop.TracingCtx"), mock.AnythingOfType("json.RawMessage")).Return(nil) - invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics) + invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) assert.Error(t, invokeErr) checkRunningInvokeMockExpectations(t, mocks) @@ -274,11 +272,11 @@ func TestRuntimeErrorFailure_SendInvokeToRuntime_Timeout(t *testing.T) { return 0, 0, 0, err } - mocks.staticData.On("FunctionTimeout").Return(time.Second) mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") - mocks.eaInvokeResponder.On("SendError", err, &mocks.staticData, mock.Anything).Return() + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Second) + mocks.eaInvokeResponder.On("SendError", err, &mocks.staticData).Return() mocks.eaInvokeResponder.On("ErrorPayloadSizeBytes").Return(mocks.errorPayloadSizeBytes) mocks.metrics.On("TriggerStartRequest") @@ -286,7 +284,7 @@ func TestRuntimeErrorFailure_SendInvokeToRuntime_Timeout(t *testing.T) { mocks.metrics.On("TriggerSentResponse", false, err, mock.Anything, 0).Return() mocks.metrics.On("SendInvokeFinishedEvent", mock.AnythingOfType("*interop.TracingCtx"), mock.AnythingOfType("json.RawMessage")).Return(nil) - invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics) + invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) assert.Error(t, invokeErr) checkRunningInvokeMockExpectations(t, mocks) @@ -301,15 +299,15 @@ func TestRunInvokeAndSendResultFailure_Timeout(t *testing.T) { mocks.timeoutCache.On("Register", invokeID) mocks.eaInvokeRequest.On("InvokeID").Return(invokeID) - mocks.staticData.On("FunctionTimeout").Return(time.Millisecond) mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") - mocks.eaInvokeResponder.On("SendError", mock.Anything, &mocks.staticData, mock.Anything).Return() + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Millisecond) + mocks.eaInvokeResponder.On("SendError", mock.Anything, &mocks.staticData).Return() mocks.eaInvokeResponder.On("ErrorPayloadSizeBytes").Return(mocks.errorPayloadSizeBytes) mockMetricsUnfinished(mocks, mock.Anything) - invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics) + invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) assert.Error(t, invokeErr) checkRunningInvokeMockExpectations(t, mocks) @@ -321,12 +319,13 @@ func TestRunInvokeAndSendResultFailure_TimeoutWhileResponse(t *testing.T) { mocks, runInvoke := createMocksAndInitRunningInvoke(t) timeoutErr := model.NewCustomerError(model.ErrorSandboxTimedout) - mocks.staticData.On("FunctionTimeout").Return(time.Second) mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Second) mocks.runtimeRespReq.On("ParsingError").Return(nil) + mocks.runtimeRespReq.On("InvocationID").Return((*string)(nil)) mocks.runtimeRespReq.On("ContentType").Return("") mocks.runtimeRespReq.On("ResponseMode").Return("") @@ -345,7 +344,7 @@ func TestRunInvokeAndSendResultFailure_TimeoutWhileResponse(t *testing.T) { assert.Error(t, err) }() - err := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics) + err := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) assert.Error(t, err) wg.Wait() @@ -359,11 +358,11 @@ func TestRunInvokeAndSendResultFailure_ContextCancelled(t *testing.T) { mocks, runInvoke := createMocksAndInitRunningInvoke(t) err := model.NewPlatformError(nil, "test fatal error") - mocks.staticData.On("FunctionTimeout").Return(time.Second) mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") - mocks.eaInvokeResponder.On("SendError", err, &mocks.staticData, mock.Anything).Return(nil) + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Second) + mocks.eaInvokeResponder.On("SendError", err, &mocks.staticData).Return(nil) mocks.eaInvokeResponder.On("ErrorPayloadSizeBytes").Return(mocks.errorPayloadSizeBytes) mockMetricsUnfinished(mocks, mock.Anything) @@ -375,7 +374,7 @@ func TestRunInvokeAndSendResultFailure_ContextCancelled(t *testing.T) { }() close(ch) - invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics) + invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) assert.Error(t, invokeErr) checkRunningInvokeMockExpectations(t, mocks) @@ -387,12 +386,13 @@ func TestRuntimeResponseFailure_ResponseWhileResponse(t *testing.T) { mocks, runInvoke := createMocksAndInitRunningInvoke(t) syncChan := make(chan time.Time) - mocks.staticData.On("FunctionTimeout").Return(5 * time.Second) mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(5 * time.Second) mocks.runtimeRespReq.On("ParsingError").Return(nil) + mocks.runtimeRespReq.On("InvocationID").Return((*string)(nil)) mocks.runtimeRespReq.On("ContentType").Return("") mocks.runtimeRespReq.On("ResponseMode").Return("") mocks.runtimeRespReq.On("TrailerError").Return(nil) @@ -421,7 +421,7 @@ func TestRuntimeResponseFailure_ResponseWhileResponse(t *testing.T) { wg.Add(1) go func() { defer wg.Done() - err := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics) + err := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) assert.NoError(t, err) }() @@ -444,13 +444,14 @@ func TestRuntimeErrorFailure_ErrorWhileError(t *testing.T) { err := model.NewCustomerError(model.ErrorFunctionUnknown) - mocks.staticData.On("FunctionTimeout").Return(time.Second) mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Second) mocks.errorPayloadSizeBytes = 100 + mocks.runtimeErrorReq.On("InvocationID").Return((*string)(nil)) mocks.runtimeErrorReq.On("GetError").Return(err) mocks.runtimeErrorReq.On("GetXrayErrorCause").Return(json.RawMessage(nil)) @@ -465,7 +466,7 @@ func TestRuntimeErrorFailure_ErrorWhileError(t *testing.T) { wg.Add(1) go func() { defer wg.Done() - err := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics) + err := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) assert.Error(t, err) assert.Equal(t, model.ErrorFunctionUnknown, err.ErrorType()) }() @@ -480,3 +481,56 @@ func TestRuntimeErrorFailure_ErrorWhileError(t *testing.T) { assert.Error(t, customerErr) assert.Equal(t, model.ErrorRuntimeInvokeErrorInProgress, customerErr.ErrorType()) } + +func TestRunInvokeAndSendResult_ResolvedTimeoutOverridesCausesTimeout(t *testing.T) { + t.Parallel() + + mocks, runInvoke := createMocksAndInitRunningInvoke(t) + + invokeID := "invoke-resolved-timeout" + mocks.timeoutCache.On("Register", invokeID) + mocks.eaInvokeRequest.On("InvokeID").Return(invokeID) + + mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) + + mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Millisecond) + + mocks.eaInvokeResponder.On("SendError", mock.Anything, &mocks.staticData).Return() + mocks.eaInvokeResponder.On("ErrorPayloadSizeBytes").Return(mocks.errorPayloadSizeBytes) + mockMetricsUnfinished(mocks, mock.Anything) + + start := time.Now() + invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) + elapsed := time.Since(start) + + assert.Error(t, invokeErr) + + assert.Less(t, elapsed, 5*time.Second) + + checkRunningInvokeMockExpectations(t, mocks) +} + +func TestRunInvokeAndSendResult_ResolvedTimeoutZeroFallsBackToFunctionTimeout(t *testing.T) { + t.Parallel() + + mocks, runInvoke := createMocksAndInitRunningInvoke(t) + + invokeID := "invoke-fallback-timeout" + mocks.timeoutCache.On("Register", invokeID) + mocks.eaInvokeRequest.On("InvokeID").Return(invokeID) + + mocks.staticData.On("XRayTracingMode").Return(intmodel.XRayTracingModePassThrough) + + mocks.eaInvokeRequest.On("TraceId").Return("Root=root1;Parent=parent1;Sampled=1;Lineage=foo:1|bar:65535") + mocks.eaInvokeRequest.On("ResolvedInvokeTimeout").Return(time.Millisecond) + + mocks.eaInvokeResponder.On("SendError", mock.Anything, &mocks.staticData).Return() + mocks.eaInvokeResponder.On("ErrorPayloadSizeBytes").Return(mocks.errorPayloadSizeBytes) + mockMetricsUnfinished(mocks, mock.Anything) + + invokeErr := runInvoke.RunInvokeAndSendResult(mocks.ctx, &mocks.staticData, &mocks.eaInvokeRequest, &mocks.metrics, &mocks.eaInvokeResponder) + assert.Error(t, invokeErr) + + checkRunningInvokeMockExpectations(t, mocks) +} diff --git a/internal/lambda-managed-instances/invoke/runtime_error_request.go b/internal/lambda-managed-instances/invoke/runtime_error_request.go index 990bf105..cb4e17d3 100644 --- a/internal/lambda-managed-instances/invoke/runtime_error_request.go +++ b/internal/lambda-managed-instances/invoke/runtime_error_request.go @@ -29,7 +29,7 @@ type runtimeError struct { invokeID interop.InvokeID errorType model.ErrorType errorCategory model.ErrorCategory - invocationID string + invocationID *string errorDetails string xrayErrorCause json.RawMessage @@ -43,7 +43,7 @@ func NewRuntimeError(ctx context.Context, request *http.Request, invokeID intero invokeID: invokeID, errorType: model.GetValidRuntimeOrFunctionErrorType(request.Header.Get(RuntimeErrorTypeHeader)), errorCategory: RuntimeErrorCategory, - invocationID: request.Header.Get(RuntimeInvocationIdHeader), + invocationID: invocationIDFromRequest(request), errorDetails: errorDetails, xrayErrorCause: getValidatedErrorCause(ctx, request.Header.Get(LambdaXRayErrorCauseHeader)), } @@ -91,7 +91,7 @@ func (r *runtimeError) GetXrayErrorCause() json.RawMessage { return r.xrayErrorCause } -func (r *runtimeError) InvocationID() string { +func (r *runtimeError) InvocationID() *string { return r.invocationID } diff --git a/internal/lambda-managed-instances/invoke/runtime_response_request.go b/internal/lambda-managed-instances/invoke/runtime_response_request.go index 7e51d52a..473f9d45 100644 --- a/internal/lambda-managed-instances/invoke/runtime_response_request.go +++ b/internal/lambda-managed-instances/invoke/runtime_response_request.go @@ -31,7 +31,7 @@ type runtimeResponse struct { contentType string invokeID interop.InvokeID responseMode string - invocationID string + invocationID *string } func NewRuntimeResponse(ctx context.Context, request *http.Request, writer http.ResponseWriter, invokeID interop.InvokeID) runtimeResponse { @@ -45,7 +45,7 @@ func NewRuntimeResponse(ctx context.Context, request *http.Request, writer http. rc: http.NewResponseController(writer), contentType: contentType, invokeID: invokeID, - invocationID: request.Header.Get(RuntimeInvocationIdHeader), + invocationID: invocationIDFromRequest(request), } switch mode := request.Header.Get(RuntimeResponseModeHeader); mode { @@ -89,7 +89,7 @@ func (r *runtimeResponse) ResponseMode() string { return r.responseMode } -func (r *runtimeResponse) InvocationID() string { +func (r *runtimeResponse) InvocationID() *string { return r.invocationID } @@ -135,3 +135,14 @@ func (t trailerError) ErrorType() model.ErrorType { func (t trailerError) ErrorDetails() string { return t.details } + +func invocationIDFromRequest(request *http.Request) *string { + if values, ok := request.Header[http.CanonicalHeaderKey(RuntimeInvocationIdHeader)]; ok { + v := "" + if len(values) > 0 { + v = values[0] + } + return &v + } + return nil +} diff --git a/internal/lambda-managed-instances/invoke/runtime_response_sender.go b/internal/lambda-managed-instances/invoke/runtime_response_sender.go index 88a8d72b..a20f4438 100644 --- a/internal/lambda-managed-instances/invoke/runtime_response_sender.go +++ b/internal/lambda-managed-instances/invoke/runtime_response_sender.go @@ -36,11 +36,11 @@ func sendInvokeToRuntime(ctx context.Context, initData interop.InitStaticDataPro if traceId != "" { runtimeReq.Header().Set(RuntimeTraceIdHeader, traceId) } - if cc := invokeReq.ClientContext(); cc != "" { - runtimeReq.Header().Set(RuntimeClientContextHeader, cc) + if clientCtx := invokeReq.ClientContext(); clientCtx != "" { + runtimeReq.Header().Set(RuntimeClientContextHeader, clientCtx) } - if cogId := buildCognitoIdentifyHeader(invokeReq); cogId != "" { - runtimeReq.Header().Set(RuntimeCognitoIdentifyHeader, cogId) + if cognitoId := buildCognitoIdentifyHeader(invokeReq); cognitoId != "" { + runtimeReq.Header().Set(RuntimeCognitoIdentifyHeader, cognitoId) } if internalInvocationID := invokeReq.InternalInvocationID(); internalInvocationID != "" { runtimeReq.Header().Set(RuntimeInvocationIdHeader, internalInvocationID) @@ -69,7 +69,7 @@ func sendInvokeToRuntime(ctx context.Context, initData interop.InitStaticDataPro select { case <-ctx.Done(): - return 0, 0, 0, BuildInvokeAppError(context.Cause(ctx), initData.FunctionTimeout()) + return 0, 0, 0, BuildInvokeAppError(context.Cause(ctx), invokeReq.ResolvedInvokeTimeout()) case err := <-resChan: if err != nil { return 0, timedWriter.TotalDuration, timedWriter.TotalDuration, model.NewCustomerError(model.ErrorRuntimeUnknown, model.WithCause(err)) diff --git a/internal/lambda-managed-instances/invoke/runtime_response_sender_test.go b/internal/lambda-managed-instances/invoke/runtime_response_sender_test.go index a8796876..ab1d2eea 100644 --- a/internal/lambda-managed-instances/invoke/runtime_response_sender_test.go +++ b/internal/lambda-managed-instances/invoke/runtime_response_sender_test.go @@ -46,6 +46,7 @@ func createMocksAndRuntimeResponder() *runtimeResponseSenderMocks { } mocks.initData.On("FunctionTimeout").Return(time.Duration(0)).Maybe() + mocks.invokeReq.On("InternalInvocationID").Return("").Maybe() return &mocks } @@ -62,7 +63,6 @@ func buildInvokeReqMocks(invokeReq *interop.MockInvokeRequest) { invokeReq.On("ClientContext").Return("client-context-example") invokeReq.On("CognitoId").Return("cognito_id_12345") invokeReq.On("CognitoPoolId").Return("cognito_pool_id_6789") - invokeReq.On("InternalInvocationID").Return("") } func buildInitDataMocks(initData *interop.MockInitStaticDataProvider) { @@ -98,6 +98,7 @@ func TestSendResponseFailure_Timeout(t *testing.T) { PayloadSize: 100, WaitBeforeRead: time.Second, }) + mocks.invokeReq.On("ResolvedInvokeTimeout").Return(time.Second) ctx, cancel := context.WithTimeout(context.Background(), time.Millisecond) defer cancel() @@ -122,6 +123,7 @@ func TestSendResponseFailure_CtxCancelled(t *testing.T) { PayloadSize: 100, WaitBeforeRead: time.Second, }) + mocks.invokeReq.On("ResolvedInvokeTimeout").Return(time.Second) ctx, cancel := context.WithCancelCause(context.Background()) cancel(model.NewCustomerError(model.ErrorReasonExtensionExecFailed)) @@ -148,7 +150,6 @@ func TestSendResponse_OmitsEmptyOptionalHeaders(t *testing.T) { mocks.invokeReq.On("ClientContext").Return("") mocks.invokeReq.On("CognitoId").Return("") mocks.invokeReq.On("CognitoPoolId").Return("") - mocks.invokeReq.On("InternalInvocationID").Return("") mocks.invokeReq.On("BodyReader").Return(mocks.reader) recorder := httptest.NewRecorder() @@ -158,9 +159,13 @@ func TestSendResponse_OmitsEmptyOptionalHeaders(t *testing.T) { assert.NoError(t, err) headers := recorder.Header() - assert.Empty(t, headers.Get(RuntimeTraceIdHeader)) - assert.Empty(t, headers.Get(RuntimeClientContextHeader)) - assert.Empty(t, headers.Get(RuntimeCognitoIdentifyHeader)) + + _, traceExists := headers[http.CanonicalHeaderKey(RuntimeTraceIdHeader)] + assert.False(t, traceExists, "Trace-Id header should be absent when empty") + _, clientCtxExists := headers[http.CanonicalHeaderKey(RuntimeClientContextHeader)] + assert.False(t, clientCtxExists, "Client-Context header should be absent when empty") + _, cognitoExists := headers[http.CanonicalHeaderKey(RuntimeCognitoIdentifyHeader)] + assert.False(t, cognitoExists, "Cognito-Identity header should be absent when empty") assert.NotEmpty(t, headers.Get(RuntimeRequestIdHeader)) assert.NotEmpty(t, headers.Get(RuntimeDeadlineHeader)) diff --git a/internal/lambda-managed-instances/logging/contextual_logger.go b/internal/lambda-managed-instances/logging/contextual_logger.go index 0580ff0a..14c45bb1 100644 --- a/internal/lambda-managed-instances/logging/contextual_logger.go +++ b/internal/lambda-managed-instances/logging/contextual_logger.go @@ -50,6 +50,10 @@ func WithInvokeID(ctx context.Context, invokeID interop.InvokeID) context.Contex return WithFields(ctx, interop.RequestIdProperty, invokeID) } +func WithReconnectID(ctx context.Context, reconnectID string) context.Context { + return WithFields(ctx, "reconnectId", reconnectID) +} + func Debug(ctx context.Context, msg string, args ...any) { FromContext(ctx).Debug(msg, args...) } diff --git a/internal/lambda-managed-instances/rapid/handlers.go b/internal/lambda-managed-instances/rapid/handlers.go index 44fb6b94..96fd701c 100644 --- a/internal/lambda-managed-instances/rapid/handlers.go +++ b/internal/lambda-managed-instances/rapid/handlers.go @@ -16,7 +16,6 @@ import ( "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/appctx" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/core" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/interop" - "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/invoke" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/logging" internalmodel "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/model" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/rapi" @@ -52,7 +51,7 @@ type rapidContext struct { telemetrySubscriptionAPI telemetry.SubscriptionAPI logsEgressAPI telemetry.StdLogsEgressAPI eventsAPI interop.EventsAPI - invokeRouter *invoke.InvokeRouter + invokeRouter InvokeRouter initMetrics interop.InitMetrics shutdownContext *shutdownContext diff --git a/internal/lambda-managed-instances/rapid/model/client_error.go b/internal/lambda-managed-instances/rapid/model/client_error.go index f2c8e5e2..08af1df3 100644 --- a/internal/lambda-managed-instances/rapid/model/client_error.go +++ b/internal/lambda-managed-instances/rapid/model/client_error.go @@ -22,6 +22,7 @@ const ( ErrorInvalidConnectionHoldTimeout ErrorType = "ErrInvalidConnectionHoldTimeout" ErrorInvalidResponseHoldTimeout ErrorType = "ErrInvalidResponseHoldTimeout" + ErrorIncompleteLongPollingConfig ErrorType = "ErrIncompleteLongPollingConfig" ) type ClientError struct { diff --git a/internal/lambda-managed-instances/rapid/model/error_types.go b/internal/lambda-managed-instances/rapid/model/error_types.go index 81000fb1..f050526f 100644 --- a/internal/lambda-managed-instances/rapid/model/error_types.go +++ b/internal/lambda-managed-instances/rapid/model/error_types.go @@ -116,7 +116,7 @@ func (e *appError) ReturnCode() int { } func GetValidRuntimeOrFunctionErrorType(errorType string) ErrorType { - match, _ := regexp.MatchString("(Runtime|Function)\\.[A-Z][a-zA-Z]+", errorType) + match, _ := regexp.MatchString("^(Runtime|Function)\\.[A-Z][a-zA-Z]+$", errorType) if match { return ErrorType(errorType) } @@ -129,7 +129,7 @@ func GetValidRuntimeOrFunctionErrorType(errorType string) ErrorType { } func GetValidExtensionErrorType(errorType string, defaultErrorType ErrorType) ErrorType { - match, _ := regexp.MatchString("Extension\\.[A-Z][a-zA-Z]+", errorType) + match, _ := regexp.MatchString("^Extension\\.[A-Z][a-zA-Z]+$", errorType) if match { return ErrorType(errorType) } diff --git a/internal/lambda-managed-instances/rapid/model/error_types_fuzz_test.go b/internal/lambda-managed-instances/rapid/model/error_types_fuzz_test.go new file mode 100644 index 00000000..3d54e86d --- /dev/null +++ b/internal/lambda-managed-instances/rapid/model/error_types_fuzz_test.go @@ -0,0 +1,83 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 + +package model + +import ( + "strings" + "testing" +) + +func isStrictErrorType(prefix, s string) bool { + rest, ok := strings.CutPrefix(s, prefix+".") + if !ok || len(rest) < 2 { + return false + } + if rest[0] < 'A' || rest[0] > 'Z' { + return false + } + for i := 1; i < len(rest); i++ { + c := rest[i] + if (c < 'a' || c > 'z') && (c < 'A' || c > 'Z') { + return false + } + } + return true +} + +func FuzzGetValidRuntimeOrFunctionErrorType(f *testing.F) { + seeds := []string{ + "", + "Runtime.MyError", + "Function.MyError", + "Sandbox.Failure Runtime.Ab", + "Runtime.MyError\n", + "Runtime.My🚀Error", + " Function.MyError ", + } + for _, s := range seeds { + f.Add(s) + } + + f.Fuzz(func(t *testing.T, input string) { + got := GetValidRuntimeOrFunctionErrorType(input) + switch got { + case ErrorRuntimeUnknown, ErrorFunctionUnknown: + + case ErrorType(input): + if !isStrictErrorType("Runtime", input) && !isStrictErrorType("Function", input) { + t.Errorf("accepted malformed input %q", input) + } + default: + t.Errorf("returned %q which is neither a fallback nor the input %q", got, input) + } + }) +} + +func FuzzGetValidExtensionErrorType(f *testing.F) { + seeds := []string{ + "", + "Extension.AA", + "Sandbox.Failure Extension.AA", + "Extension.AA\n", + "Extension.A🚀", + " Extension.AA ", + } + for _, s := range seeds { + f.Add(s) + } + + f.Fuzz(func(t *testing.T, input string) { + got := GetValidExtensionErrorType(input, ErrorAgentExit) + switch got { + case ErrorAgentExit: + + case ErrorType(input): + if !isStrictErrorType("Extension", input) { + t.Errorf("accepted malformed input %q", input) + } + default: + t.Errorf("returned %q which is neither the default nor the input %q", got, input) + } + }) +} diff --git a/internal/lambda-managed-instances/rapid/model/error_types_test.go b/internal/lambda-managed-instances/rapid/model/error_types_test.go index 3aee46d7..9c14f8de 100644 --- a/internal/lambda-managed-instances/rapid/model/error_types_test.go +++ b/internal/lambda-managed-instances/rapid/model/error_types_test.go @@ -20,6 +20,19 @@ func TestGetValidRuntimeOrFunctionErrorType(t *testing.T) { {"", ErrorRuntimeUnknown}, {"MyCustomError", ErrorRuntimeUnknown}, {"MyCustomError.Error", ErrorRuntimeUnknown}, + {"Sandbox.Failure Runtime.Ab", ErrorRuntimeUnknown}, + {" Runtime.MyError", ErrorRuntimeUnknown}, + {"Runtime.MyError ", ErrorRuntimeUnknown}, + {"Runtime.My Error", ErrorRuntimeUnknown}, + {"Runtime.MyError\n", ErrorRuntimeUnknown}, + {"runtime.MyError", ErrorRuntimeUnknown}, + {"Runtime.error", ErrorRuntimeUnknown}, + {"Runtime.Error404", ErrorRuntimeUnknown}, + {"Runtime.My-Error", ErrorRuntimeUnknown}, + {"Runtime.My🚀Error", ErrorRuntimeUnknown}, + {"Runtime.Ошибка", ErrorRuntimeUnknown}, + {"Runtime..Error", ErrorRuntimeUnknown}, + {"Function.error", ErrorFunctionUnknown}, {"Runtime.MyCustomErrorTypeHere", ErrorType("Runtime.MyCustomErrorTypeHere")}, {"Function.MyCustomErrorTypeHere", ErrorType("Function.MyCustomErrorTypeHere")}, } @@ -49,6 +62,16 @@ func TestGetValidExtensionErrorType(t *testing.T) { {"Extension.", defaultErrorType}, {"Extension.A", defaultErrorType}, {"Extension.az", defaultErrorType}, + {"Sandbox.Failure Extension.AA", defaultErrorType}, + {" Extension.AA", defaultErrorType}, + {"Extension.AA ", defaultErrorType}, + {"Extension.A A", defaultErrorType}, + {"Extension.AA\n", defaultErrorType}, + {"extension.AA", defaultErrorType}, + {"Extension.Error404", defaultErrorType}, + {"Extension.A🚀", defaultErrorType}, + {"Extension.Ошибка", defaultErrorType}, + {"Extension..Error", defaultErrorType}, {"Extension.AA", ErrorType("Extension.AA")}, {"Extension.Az", ErrorType("Extension.Az")}, } diff --git a/internal/lambda-managed-instances/rapid/sandbox.go b/internal/lambda-managed-instances/rapid/sandbox.go index 9a452901..e28e1ef3 100644 --- a/internal/lambda-managed-instances/rapid/sandbox.go +++ b/internal/lambda-managed-instances/rapid/sandbox.go @@ -6,6 +6,7 @@ package rapid import ( "context" "log/slog" + "net/http" "net/netip" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lmds" @@ -18,6 +19,7 @@ import ( rapimodel "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/rapi/model" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/rapi/rendering" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/rapid/model" + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/servicelogs" supvmodel "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/supervisor/model" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/telemetry" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/utils" @@ -32,8 +34,9 @@ type Dependencies struct { EventsAPI interop.EventsAPI Supervisor supvmodel.ProcessSupervisor FileUtils utils.FileUtil - InvokeRouter *invoke.InvokeRouter - MetadataService *lmds.Service + + InvokeRouter *invoke.LongInvokerRouter + MetadataService *lmds.Service RuntimeAPIAddrPort netip.AddrPort } @@ -45,7 +48,7 @@ func Start(ctx context.Context, deps Dependencies) (interop.RapidContext, error) registrationService := core.NewRegistrationService(initFlow) renderingService := rendering.NewRenderingService() - server, err := rapi.NewServer(deps.RuntimeAPIAddrPort, appCtx, registrationService, renderingService, deps.TelemetrySubscriptionAPI, deps.InvokeRouter, deps.MetadataService) + server, err := rapi.NewServer(deps.RuntimeAPIAddrPort, appCtx, registrationService, renderingService, deps.TelemetrySubscriptionAPI, deps.InvokeRouter.InnerRouter(), deps.MetadataService) if err != nil { return nil, err } @@ -61,8 +64,9 @@ func Start(ctx context.Context, deps Dependencies) (interop.RapidContext, error) renderingService: renderingService, shutdownContext: newShutdownContext(), fileUtils: deps.FileUtils, - invokeRouter: deps.InvokeRouter, - processTermChan: make(chan model.AppError), + + invokeRouter: deps.InvokeRouter, + processTermChan: make(chan model.AppError), telemetrySubscriptionAPI: deps.TelemetrySubscriptionAPI, logsEgressAPI: deps.LogsEgressAPI, @@ -90,17 +94,24 @@ func (r *rapidContext) HandleInit(ctx context.Context, initData interop.InitExec return handleInit(ctx, r) } -func (r *rapidContext) HandleInvoke(ctx context.Context, invokeReq interop.InvokeRequest, metrics interop.InvokeMetrics) (err model.AppError, wasResponseSent bool) { +func (r *rapidContext) HandleInvoke(ctx context.Context, invokeReq interop.InvokeRequest, metrics interop.InvokeMetrics, responseWriter http.ResponseWriter) (err model.AppError, wasResponseSent bool, invokePending bool) { if err := invokeReq.UpdateFromInitData(&r.initExecutionData); err != nil { - return err, false + return err, false, false } metrics.AttachDependencies(&r.initExecutionData, r.eventsAPI) - return r.invokeRouter.Invoke(ctx, &r.initExecutionData, invokeReq, metrics) + return r.invokeRouter.Invoke(ctx, &r.initExecutionData, invokeReq, metrics, responseWriter) +} + +func (r *rapidContext) HandleReconnect(ctx context.Context, invokeID interop.InvokeID, responseWriter http.ResponseWriter, metrics interop.ReconnectMetrics) interop.ReconnectResult { + return r.invokeRouter.Reconnect(ctx, invokeID, responseWriter, metrics) } func (r *rapidContext) HandleShutdown(shutdownCause model.AppError, metrics interop.ShutdownMetrics) model.AppError { metrics.SetAgentCount(len(r.registrationService.GetInternalAgents()), len(r.registrationService.GetExternalAgents())) + metrics.AddMetric(servicelogs.Counter(interop.RuntimeNextCountMetric, float64(r.invokeRouter.GetActiveRuntimeCount()))) + metrics.AddMetric(servicelogs.Counter(interop.RuntimeWorkerCountMetric, float64(r.initExecutionData.FunctionMetadata.RuntimeWorkerCount))) + r.invokeRouter.AbortRunningInvokes(metrics, shutdownCause) reason := rapimodel.Spindown @@ -140,3 +151,10 @@ func (r *rapidContext) HandleShutdown(shutdownCause model.AppError, metrics inte func (r *rapidContext) RuntimeAPIAddrPort() netip.AddrPort { return r.server.AddrPort() } + +type InvokeRouter interface { + Invoke(ctx context.Context, initData interop.InitStaticDataProvider, invokeReq interop.InvokeRequest, metrics interop.InvokeMetrics, responseWriter http.ResponseWriter) (err model.AppError, wasResponseSent bool, invokePending bool) + Reconnect(ctx context.Context, invokeID interop.InvokeID, responseWriter http.ResponseWriter, metrics interop.ReconnectMetrics) interop.ReconnectResult + AbortRunningInvokes(metrics interop.ShutdownMetrics, err model.AppError) + GetActiveRuntimeCount() int +} diff --git a/internal/lambda-managed-instances/raptor/app.go b/internal/lambda-managed-instances/raptor/app.go index 60b98140..fc12c36e 100644 --- a/internal/lambda-managed-instances/raptor/app.go +++ b/internal/lambda-managed-instances/raptor/app.go @@ -8,6 +8,7 @@ import ( "errors" "fmt" "log/slog" + "net/http" "net/netip" "sync" "sync/atomic" @@ -51,7 +52,7 @@ func StartApp(deps rapid.Dependencies, telemetryFDSocketPath, metadataToken stri app := &App{ rapidCtx: rapidCtx, - invokeRouter: deps.InvokeRouter, + invokeRouter: deps.InvokeRouter.InnerRouter(), state: internal.NewStateGuard(), doneCh: make(chan struct{}), shutdownStartedCh: make(chan struct{}), @@ -106,11 +107,11 @@ func (a *App) Init(ctx context.Context, init *internalModel.InitRequestMessage, return nil } -func (a *App) Invoke(ctx context.Context, invokeMsg interop.InvokeRequest, metrics interop.InvokeMetrics) (err model.AppError, wasResponseSent bool) { +func (a *App) Invoke(ctx context.Context, invokeMsg interop.InvokeRequest, metrics interop.InvokeMetrics, responseWriter http.ResponseWriter) (err model.AppError, wasResponseSent bool, invokePending bool) { currState := a.state.GetState() switch currState { case internal.Initialized: - return a.rapidCtx.HandleInvoke(ctx, invokeMsg, metrics) + return a.rapidCtx.HandleInvoke(ctx, invokeMsg, metrics, responseWriter) case internal.Idle, internal.Initializing: logging.Error(ctx, "Sandbox not Initialized", "state", currState) return interop.ClientError{ @@ -119,7 +120,7 @@ func (a *App) Invoke(ctx context.Context, invokeMsg interop.InvokeRequest, metri model.ErrorSeverityError, model.ErrorInitIncomplete, ), - }, false + }, false, false case internal.ShuttingDown, internal.Shutdown: logging.Error(ctx, "Invoke while Sandbox shutting down") return interop.ClientError{ @@ -128,7 +129,28 @@ func (a *App) Invoke(ctx context.Context, invokeMsg interop.InvokeRequest, metri model.ErrorSeverityFatal, model.ErrorEnvironmentUnhealthy, ), - }, false + }, false, false + default: + panic(fmt.Sprintf("unknown current state: %d", currState)) + } +} + +func (a *App) Reconnect(ctx context.Context, invokeID interop.InvokeID, responseWriter http.ResponseWriter, reconnectMetrics interop.ReconnectMetrics) interop.ReconnectResult { + currState := a.state.GetState() + switch currState { + case internal.Initialized, internal.ShuttingDown: + + return a.rapidCtx.HandleReconnect(ctx, invokeID, responseWriter, reconnectMetrics) + case internal.Idle, internal.Initializing: + logging.Error(ctx, "Reconnect: sandbox not initialized", "state", currState) + return interop.ReconnectResult{Outcome: interop.ReconnectOutcomeError, Err: interop.ClientError{ + ClientError: model.NewClientError(ErrNotInitialized, model.ErrorSeverityError, model.ErrorInitIncomplete), + }} + case internal.Shutdown: + logging.Error(ctx, "Reconnect: sandbox already shut down", "state", currState) + return interop.ReconnectResult{Outcome: interop.ReconnectOutcomeError, Err: interop.ClientError{ + ClientError: model.NewClientError(ErrorEnvironmentUnhealthy, model.ErrorSeverityError, model.ErrorEnvironmentUnhealthy), + }} default: panic(fmt.Sprintf("unknown current state: %d", currState)) } diff --git a/internal/lambda-managed-instances/raptor/app_test.go b/internal/lambda-managed-instances/raptor/app_test.go index 581712a7..152fd055 100644 --- a/internal/lambda-managed-instances/raptor/app_test.go +++ b/internal/lambda-managed-instances/raptor/app_test.go @@ -14,6 +14,7 @@ import ( "github.com/stretchr/testify/require" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/interop" + "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/invoke" internalModel "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/model" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/rapid/model" "github.com/aws/aws-lambda-runtime-interface-emulator/internal/lambda-managed-instances/raptor/internal" @@ -201,9 +202,10 @@ func TestAppInvokeStateValidation(t *testing.T) { require.NoError(t, app.state.SetState(state)) } - err, wasResponseSent := app.Invoke(context.Background(), invokeMsg, invokeMetrics) + err, wasResponseSent, invokePending := app.Invoke(context.Background(), invokeMsg, invokeMetrics, nil) assert.False(t, wasResponseSent) + assert.False(t, invokePending) assert.ErrorAs(t, err, &interop.ClientError{}) assert.Equal(t, tc.wantErrorType, err.ErrorType()) assert.Equal(t, tc.wantError, err.Unwrap()) @@ -211,6 +213,68 @@ func TestAppInvokeStateValidation(t *testing.T) { } } +func TestAppReconnectStateValidation(t *testing.T) { + testCases := []struct { + name string + states []internal.Status + expectError bool + wantErrorType model.ErrorType + }{ + { + name: "Idle_rejects", + states: []internal.Status{}, + expectError: true, + wantErrorType: model.ErrorInitIncomplete, + }, + { + name: "Initializing_rejects", + states: []internal.Status{internal.Initializing}, + expectError: true, + wantErrorType: model.ErrorInitIncomplete, + }, + { + name: "Initialized_allows", + states: []internal.Status{internal.Initializing, internal.Initialized}, + expectError: false, + }, + { + name: "ShuttingDown_allows", + states: []internal.Status{internal.ShuttingDown}, + expectError: false, + }, + { + name: "Shutdown_rejects", + states: []internal.Status{internal.ShuttingDown, internal.Shutdown}, + expectError: true, + wantErrorType: model.ErrorEnvironmentUnhealthy, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + mockRapidCtx, app, _, _, _, _ := setupAppTest(t) + + for _, state := range tc.states { + require.NoError(t, app.state.SetState(state)) + } + + if !tc.expectError { + mockRapidCtx.On("HandleReconnect", mock.Anything, mock.Anything, mock.Anything, mock.Anything). + Return(interop.ReconnectResult{}) + } + + res := app.Reconnect(context.Background(), "test-invoke", nil, invoke.NoopReconnectMetrics()) + + if tc.expectError { + assert.ErrorAs(t, res.Err, &interop.ClientError{}) + assert.Equal(t, tc.wantErrorType, res.Err.ErrorType()) + } else { + assert.NoError(t, res.Err) + } + }) + } +} + func setupAppTest(t *testing.T) (*interop.MockRapidContext, *App, *internalModel.InitRequestMessage, interop.InitMetrics, interop.InvokeRequest, interop.InvokeMetrics) { mockRapidCtx := interop.NewMockRapidContext(t) diff --git a/internal/lambda-managed-instances/raptor/server_test.go b/internal/lambda-managed-instances/raptor/server_test.go index a574f85c..4d99dce7 100644 --- a/internal/lambda-managed-instances/raptor/server_test.go +++ b/internal/lambda-managed-instances/raptor/server_test.go @@ -53,11 +53,11 @@ func TestStartNewServer_TCP(t *testing.T) { server, err := StartServer(mockShutdownHandler, handler, &TCPAddress{ eaAPIAddrPort, - }, false) + }, true) require.NoError(t, err) assert.Equal(t, eaAPIAddrPort, server.Addr.(*TCPAddress).AddrPort) - _, err = http.Get("http://" + server.Addr.String()) + _, err = testutils.NewH2CClient().Get("http://" + server.Addr.String()) require.NoError(t, err) } @@ -80,7 +80,7 @@ func TestStartNewServe_TCP_ListenError(t *testing.T) { mockShutdownHandler := newMockShutdownHandler(t) handler := mocks.NewNoOpHandler() - _, err := StartServer(mockShutdownHandler, handler, &TCPAddress{eaAPIAddrPort}, false) + _, err := StartServer(mockShutdownHandler, handler, &TCPAddress{eaAPIAddrPort}, true) assert.Error(t, err) assert.Contains(t, err.Error(), "1.1.1.1:49275") } diff --git a/internal/lambda-managed-instances/servicelogs/logger.go b/internal/lambda-managed-instances/servicelogs/logger.go index fa0aebff..24bbd904 100644 --- a/internal/lambda-managed-instances/servicelogs/logger.go +++ b/internal/lambda-managed-instances/servicelogs/logger.go @@ -16,10 +16,12 @@ type Logger interface { type Operation string const ( - InitOp Operation = "Init" - InvokeOp Operation = "Invoke" - ShutdownOp Operation = "Shutdown" - ReserveOp Operation = "Reserve" + InitOp Operation = "Init" + InvokeOp Operation = "Invoke" + ShutdownOp Operation = "Shutdown" + ReserveOp Operation = "Reserve" + ReconnectOp Operation = "Reconnect" + InvokePendingOp Operation = "InvokePending" ) type Tuple struct { diff --git a/internal/lambda-managed-instances/supervisor/local/process.go b/internal/lambda-managed-instances/supervisor/local/process.go index 9e64a501..e6242b6a 100644 --- a/internal/lambda-managed-instances/supervisor/local/process.go +++ b/internal/lambda-managed-instances/supervisor/local/process.go @@ -239,7 +239,7 @@ func kill(p process, name string, deadline time.Time) error { slog.Info("Sending SIGKILL to process", "name", name, "pid", p.pid) } - if (time.Since(deadline)) > 0 { + if time.Since(deadline) > 0 { return fmt.Errorf("invalid timeout while killing %s", name) } diff --git a/internal/lambda-managed-instances/supervisor/local/process_test.go b/internal/lambda-managed-instances/supervisor/local/process_test.go index fe5f83dc..ef483114 100644 --- a/internal/lambda-managed-instances/supervisor/local/process_test.go +++ b/internal/lambda-managed-instances/supervisor/local/process_test.go @@ -231,13 +231,13 @@ func TestTerminateCheckStatus(t *testing.T) { func TestCheckOomKill_OomKilled(t *testing.T) { path := filepath.Join(t.TempDir(), "memory.events") - os.WriteFile(path, []byte("low 0\nhigh 0\nmax 96\noom 1\noom_kill 1\noom_group_kill 0\n"), 0644) + require.NoError(t, os.WriteFile(path, []byte("low 0\nhigh 0\nmax 96\noom 1\noom_kill 1\noom_group_kill 0\n"), 0o644)) assert.True(t, checkOomKill(path)) } func TestCheckOomKill_NoOom(t *testing.T) { path := filepath.Join(t.TempDir(), "memory.events") - os.WriteFile(path, []byte("low 0\nhigh 0\nmax 0\noom 0\noom_kill 0\noom_group_kill 0\n"), 0644) + require.NoError(t, os.WriteFile(path, []byte("low 0\nhigh 0\nmax 0\noom 0\noom_kill 0\noom_group_kill 0\n"), 0o644)) assert.False(t, checkOomKill(path)) } @@ -247,18 +247,18 @@ func TestCheckOomKill_FileNotFound(t *testing.T) { func TestCheckOomKill_MultipleOomKills(t *testing.T) { path := filepath.Join(t.TempDir(), "memory.events") - os.WriteFile(path, []byte("low 0\nhigh 0\nmax 50\noom 3\noom_kill 3\noom_group_kill 0\n"), 0644) + require.NoError(t, os.WriteFile(path, []byte("low 0\nhigh 0\nmax 50\noom 3\noom_kill 3\noom_group_kill 0\n"), 0o644)) assert.True(t, checkOomKill(path)) } func TestCheckOomKill_MalformedCount(t *testing.T) { path := filepath.Join(t.TempDir(), "memory.events") - os.WriteFile(path, []byte("low 0\nhigh 0\noom_kill abc\n"), 0644) + require.NoError(t, os.WriteFile(path, []byte("low 0\nhigh 0\noom_kill abc\n"), 0o644)) assert.False(t, checkOomKill(path)) } func TestCheckOomKill_NoOomKillLine(t *testing.T) { path := filepath.Join(t.TempDir(), "memory.events") - os.WriteFile(path, []byte("low 0\nhigh 0\nmax 0\n"), 0644) + require.NoError(t, os.WriteFile(path, []byte("low 0\nhigh 0\nmax 0\n"), 0o644)) assert.False(t, checkOomKill(path)) } diff --git a/internal/lambda-managed-instances/testutils/functional/process_supervisor.go b/internal/lambda-managed-instances/testutils/functional/process_supervisor.go index 072470f8..be2ef5c7 100644 --- a/internal/lambda-managed-instances/testutils/functional/process_supervisor.go +++ b/internal/lambda-managed-instances/testutils/functional/process_supervisor.go @@ -148,6 +148,7 @@ func (r *RuntimeExecutionEnvironment) executeEnvActions(client *Client, t *testi if a.InvokeID == "" { a.InvokeID = r.InvokeID } + if a.ResponseHeaders == nil && r.InvocationID != "" { a.ResponseHeaders = map[string]string{invoke.RuntimeInvocationIdHeader: r.InvocationID} } @@ -161,6 +162,7 @@ func (r *RuntimeExecutionEnvironment) executeEnvActions(client *Client, t *testi if a.InvokeID == "" { a.InvokeID = r.InvokeID } + if a.ResponseHeaders == nil && r.InvocationID != "" { a.ResponseHeaders = map[string]string{invoke.RuntimeInvocationIdHeader: r.InvocationID} }