diff --git a/chasm/interceptors.go b/chasm/interceptors.go index 484eb5f0c48..06a58c52de2 100644 --- a/chasm/interceptors.go +++ b/chasm/interceptors.go @@ -5,6 +5,7 @@ import ( "go.temporal.io/server/common/log" "go.temporal.io/server/common/metrics" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc" ) @@ -60,6 +61,15 @@ func (i *ChasmVisibilityInterceptor) Intercept( return handler(ctx, req) } +func (i *ChasmVisibilityInterceptor) InterceptNexus( + ctx context.Context, + in interceptornexus.InterceptorInput, + next interceptornexus.HandlerFunc, +) (any, error) { + ctx = NewVisibilityManagerContext(ctx, i.visibilityMgr) + return next(ctx, in) +} + func ChasmVisibilityInterceptorProvider(visibilityMgr VisibilityManager) *ChasmVisibilityInterceptor { return &ChasmVisibilityInterceptor{ visibilityMgr: visibilityMgr, diff --git a/common/authorization/interceptor.go b/common/authorization/interceptor.go index 48a106f87b5..a2d41746266 100644 --- a/common/authorization/interceptor.go +++ b/common/authorization/interceptor.go @@ -3,8 +3,8 @@ package authorization import ( "cmp" "context" - "crypto/x509" "crypto/x509/pkix" + "errors" "time" commonpb "go.temporal.io/api/common/v1" @@ -18,9 +18,11 @@ import ( "go.temporal.io/server/common/log/tag" "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/namespace" + commonnexus "go.temporal.io/server/common/nexus" + "go.temporal.io/server/common/rpc/interceptor/nexus" + "go.temporal.io/server/common/rpc/tlsinfo" "google.golang.org/grpc" "google.golang.org/grpc/credentials" - "google.golang.org/grpc/peer" ) type ( @@ -54,32 +56,6 @@ var ( AuthHeader contextKeyAuthHeader ) -// TLSInfoFromContext extracts TLS information from the context's peer value. -func TLSInfoFromContext(ctx context.Context) *credentials.TLSInfo { - p, ok := peer.FromContext(ctx) - if !ok { - return nil - } - if tlsInfo, ok := p.AuthInfo.(credentials.TLSInfo); ok { - return &tlsInfo - } - return nil -} - -// PeerCert extracts an x509 certificate from given tlsInfo. -func PeerCert(tlsInfo *credentials.TLSInfo) *x509.Certificate { - if tlsInfo == nil || len(tlsInfo.State.VerifiedChains) == 0 || len(tlsInfo.State.VerifiedChains[0]) == 0 { - return nil - } - // The assumption here is that we only expect a single verified chain of certs (first[0]). - // It's unclear how we should handle a situation when more than one chain is presented, - // which subject to use. It's okay for us to limit ourselves to one chain. - // We can always extend this logic later. - // We take the first element in the chain ([0]) because that's the client cert - // (at the beginning of the chain), not intermediary CAs or the root CA (at the end of the chain). - return tlsInfo.State.VerifiedChains[0][0] -} - type Interceptor struct { claimMapper ClaimMapper authorizer Authorizer @@ -132,7 +108,7 @@ func (a *Interceptor) Intercept( info *grpc.UnaryServerInfo, handler grpc.UnaryHandler, ) (any, error) { - tlsConnection := TLSInfoFromContext(ctx) + tlsConnection := tlsinfo.FromContext(ctx) authInfo := a.GetAuthInfo(tlsConnection, headers.NewGRPCHeaderGetter(ctx), func() string { if a.audienceGetter != nil { @@ -184,6 +160,58 @@ func (a *Interceptor) Intercept( return handler(ctx, req) } +func (a *Interceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + a.logger.Debug("authorizing request") + ctx = headers.StripPrincipal(ctx) + if a.authorizer == nil { + return next(ctx, in) + } + namespaceName := in.NamespaceName() + apiName := in.APIName() + endpointName := in.EndpointName() + claims, _ := ctx.Value(MappedClaims).(*Claims) //nolint:revive // unchecked-type-assertion: empty claims will 403 + ct := &CallTarget{ + APIName: apiName, + NexusEndpointName: endpointName, + Namespace: namespaceName, + Request: in.Request(), + } + principal, err := a.Authorize(ctx, claims, ct) + if err != nil { + if permissionDeniedError, ok := errors.AsType[*serviceerror.PermissionDenied](err); ok { + a.logger.Debug("Request unauthorized") + return nil, &nexus.InterceptorError{ + Err: commonnexus.AdaptAuthorizeError(permissionDeniedError), + Outcome: "unauthorized", + SkipServiceErrorReporting: true, + } + } + logTags := []tag.Tag{ + tag.Operation(api.MethodName(apiName)), + tag.WorkflowNamespace(namespaceName), + tag.Endpoint(endpointName), + tag.Error(err), + } + if operationName := in.OperationName(); operationName != "" { + logTags = append(logTags, tag.NexusOperation(operationName)) + } + a.logger.Error("Authorization internal error with processing nexus request", logTags...) + return nil, &nexus.InterceptorError{ + Err: err, + Outcome: "internal_auth_error", + SkipServiceErrorReporting: true, + } + } + if a.enablePrincipalPropagation != nil && a.enablePrincipalPropagation(namespaceName) && principal != nil { + ctx = headers.SetPrincipal(ctx, principal) + } + return next(ctx, in) +} + // InterceptStream is a gRPC stream server interceptor that enforces authorization. func (a *Interceptor) InterceptStream( srv any, @@ -194,7 +222,7 @@ func (a *Interceptor) InterceptStream( ctx := ss.Context() bypassAuth := a.disableStreamingAuthorizer() if !bypassAuth { - tlsConnection := TLSInfoFromContext(ctx) + tlsConnection := tlsinfo.FromContext(ctx) headerGetter := headers.NewGRPCHeaderGetter(ctx) authInfo := a.GetAuthInfo(tlsConnection, headerGetter, func() string { @@ -260,7 +288,7 @@ func (a *Interceptor) GetAuthInfo(tlsConnection *credentials.TLSInfo, header hea authHeader = header.Get(a.authHeaderName) authExtraHeader = header.Get(a.authExtraHeaderName) } - clientCert := PeerCert(tlsConnection) + clientCert := tlsinfo.PeerCert(tlsConnection) if clientCert != nil { tlsSubject = &clientCert.Subject } diff --git a/common/authorization/interceptor_test.go b/common/authorization/interceptor_test.go index e6ee83e2c7b..05f4499589f 100644 --- a/common/authorization/interceptor_test.go +++ b/common/authorization/interceptor_test.go @@ -8,7 +8,9 @@ import ( "errors" "slices" "testing" + "time" + "github.com/nexus-rpc/sdk-go/nexus" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" commandpb "go.temporal.io/api/command/v1" @@ -16,12 +18,14 @@ import ( enumspb "go.temporal.io/api/enums/v1" "go.temporal.io/api/serviceerror" "go.temporal.io/api/workflowservice/v1" + "go.temporal.io/server/api/matchingservice/v1" "go.temporal.io/server/common/api" "go.temporal.io/server/common/dynamicconfig" "go.temporal.io/server/common/headers" "go.temporal.io/server/common/log" "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/namespace" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.uber.org/mock/gomock" "google.golang.org/grpc" "google.golang.org/grpc/credentials" @@ -68,6 +72,78 @@ func TestAuthorizerInterceptorSuite(t *testing.T) { suite.Run(t, s) } +func (s *authorizerInterceptorSuite) TestInterceptNexus() { + apiName, endpoint := "NexusAPI", "endpoint" + authorizationRequest := &matchingservice.DispatchNexusTaskRequest{} + input := interceptornexus.NewStartOpInput( + "s", + "o", + testNamespace, + time.Now(), + nexus.StartOperationOptions{}, + nil, + interceptornexus.ForwardingInfo{}, + interceptornexus.RequestMetadata{ + APIName: apiName, + EndpointName: endpoint, + Request: authorizationRequest, + }, + ) + expectedTarget := &CallTarget{ + APIName: apiName, + NexusEndpointName: endpoint, + Namespace: testNamespace, + Request: authorizationRequest, + } + for _, tc := range []struct { + name string + ctx context.Context + authorizationResult *Result + nextCalled bool + expectedError error + }{ + { + name: "authorized", + ctx: context.Background(), + authorizationResult: &Result{Decision: DecisionAllow}, + nextCalled: true, + }, + { + name: "unauthorized", + ctx: context.Background(), + authorizationResult: &Result{Decision: DecisionDeny}, + expectedError: &interceptornexus.InterceptorError{ + Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnauthorized, "permission denied"), + Outcome: "unauthorized", + SkipServiceErrorReporting: true, + }, + }, + } { + s.Run(tc.name, func() { + if tc.authorizationResult != nil { + s.mockAuthorizer.EXPECT().Authorize(gomock.Any(), nil, expectedTarget). + Return(*tc.authorizationResult, nil) + if tc.authorizationResult.Decision == DecisionDeny { + s.mockMetricsHandler.EXPECT(). + Counter(metrics.ServiceErrUnauthorizedCounter.Name()). + Return(metrics.NoopCounterMetricFunc) + } + } + + nextCalled := false + _, err := s.interceptor.InterceptNexus( + tc.ctx, + input, + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + nextCalled = true + return nil, nil + }) + s.Equal(tc.expectedError, err) + s.Equal(tc.nextCalled, nextCalled) + }) + } +} + func (s *authorizerInterceptorSuite) SetupTest() { s.Assertions = require.New(s.T()) s.controller = gomock.NewController(s.T()) diff --git a/common/rpc/grpcfaults/interceptor.go b/common/rpc/grpcfaults/interceptor.go index 1d7828f65d7..c5100a1f8d4 100644 --- a/common/rpc/grpcfaults/interceptor.go +++ b/common/rpc/grpcfaults/interceptor.go @@ -3,6 +3,7 @@ package grpcfaults import ( "context" + "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc" ) @@ -43,3 +44,33 @@ func UnaryServerInterceptor(generator Generator) grpc.UnaryServerInterceptor { return resp, err } } + +func NewFaultsInterceptor(generator Generator) *FaultsInterceptor { + return &FaultsInterceptor{ + h: UnaryServerInterceptor(generator), + } +} + +type FaultsInterceptor struct { + h grpc.UnaryServerInterceptor +} + +func (g *FaultsInterceptor) Intercept( + ctx context.Context, + req any, + info *grpc.UnaryServerInfo, + handler grpc.UnaryHandler, +) (any, error) { + if g.h == nil { + return handler(ctx, req) + } + return g.h(ctx, req, info, handler) +} + +func (g *FaultsInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + return next(ctx, in) +} diff --git a/common/rpc/interceptor/caller_info.go b/common/rpc/interceptor/caller_info.go index a8937f7b4b4..e199138e83a 100644 --- a/common/rpc/interceptor/caller_info.go +++ b/common/rpc/interceptor/caller_info.go @@ -6,6 +6,7 @@ import ( "go.temporal.io/server/common/api" "go.temporal.io/server/common/headers" "go.temporal.io/server/common/namespace" + "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc" ) @@ -40,6 +41,20 @@ func (i *CallerInfoInterceptor) Intercept( return handler(ctx, req) } +// InterceptNexus adds caller information for a Nexus request. +func (i *CallerInfoInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + ctx = PopulateCallerInfo( + ctx, + in.NamespaceName, + in.MethodName, + ) + return next(headers.Propagate(ctx), in) +} + // PopulateCallerInfo gets current caller info value from the context and updates any that are missing. // Namespace name and method are passed as functions to avoid expensive lookups if those values are already set. func PopulateCallerInfo( diff --git a/common/rpc/interceptor/caller_info_test.go b/common/rpc/interceptor/caller_info_test.go index e028cb4266b..25f5901e1ad 100644 --- a/common/rpc/interceptor/caller_info_test.go +++ b/common/rpc/interceptor/caller_info_test.go @@ -2,13 +2,18 @@ package interceptor import ( "context" + "net/http" "testing" + "time" + "github.com/nexus-rpc/sdk-go/nexus" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" "go.temporal.io/api/workflowservice/v1" "go.temporal.io/server/common/headers" "go.temporal.io/server/common/namespace" + "go.temporal.io/server/common/nexus/nexusrpc" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.uber.org/mock/gomock" "google.golang.org/grpc" ) @@ -124,6 +129,51 @@ func (s *callerInfoSuite) TestIntercept_CallerName() { } } +func (s *callerInfoSuite) TestInterceptNexus() { + completeInput, err := interceptornexus.NewCompleteOpInput( + testNamespace, + time.Now(), + &nexusrpc.CompletionRequest{HTTPRequest: &http.Request{}}, + nil, + interceptornexus.ForwardingInfo{}, + interceptornexus.RequestMetadata{}, + ) + s.NoError(err) + for _, tc := range []struct { + name string + input interceptornexus.InterceptorInput + callerInfo headers.CallerInfo + expectedOrigin string + }{ + { + name: "start", + input: interceptornexus.NewStartOpInput("s", "o", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, interceptornexus.ForwardingInfo{}, interceptornexus.RequestMetadata{}), + expectedOrigin: "StartNexusOperation", + }, + { + name: "cancel - preserves background origin", + input: interceptornexus.NewCancelOpInput("s", "o", testNamespace, time.Now(), nexus.CancelOperationOptions{}, "t", interceptornexus.ForwardingInfo{}, interceptornexus.RequestMetadata{}), + callerInfo: headers.SystemBackgroundHighCallerInfo, + }, + { + name: "complete", + input: completeInput, + expectedOrigin: "CompleteNexusOperation", + }, + } { + s.Run(tc.name, func() { + ctx := headers.SetCallerInfo(context.Background(), tc.callerInfo) + _, err := s.interceptor.InterceptNexus(ctx, tc.input, func(ctx context.Context, _ interceptornexus.InterceptorInput) (any, error) { + callerInfo := headers.GetCallerInfo(ctx) + s.Equal(testNamespace, callerInfo.CallerName) + s.Equal(tc.expectedOrigin, callerInfo.CallOrigin) + return nil, nil + }) + s.NoError(err) + }) + } +} + func (s *callerInfoSuite) TestIntercept_CallerType() { s.mockRegistry.EXPECT().GetNamespace(gomock.Any()).Return(nil, nil).AnyTimes() diff --git a/common/rpc/interceptor/concurrent_request_limit.go b/common/rpc/interceptor/concurrent_request_limit.go index 0014c794ea6..35d5236ec01 100644 --- a/common/rpc/interceptor/concurrent_request_limit.go +++ b/common/rpc/interceptor/concurrent_request_limit.go @@ -13,6 +13,7 @@ import ( "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/namespace" "go.temporal.io/server/common/quotas/calculator" + "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc" ) @@ -116,6 +117,24 @@ func (ni *ConcurrentRequestLimitInterceptor) Allow( return cleanup, nil } +// InterceptNexus enforces the namespace concurrent-request limit for a Nexus request. +func (ni *ConcurrentRequestLimitInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + metricsHandler := GetMetricsHandlerFromContext(ctx, ni.logger) + cleanup, err := ni.Allow(namespace.Name(in.NamespaceName()), in.APIName(), metricsHandler, in) + defer cleanup() + if err != nil { + return nil, &nexus.InterceptorError{ + Err: err, + Outcome: "namespace_concurrency_limited", + } + } + return next(ctx, in) +} + func (ni *ConcurrentRequestLimitInterceptor) counter( namespace namespace.Name, methodName string, diff --git a/common/rpc/interceptor/concurrent_request_limit_test.go b/common/rpc/interceptor/concurrent_request_limit_test.go index e14c7159204..a96a986cf30 100644 --- a/common/rpc/interceptor/concurrent_request_limit_test.go +++ b/common/rpc/interceptor/concurrent_request_limit_test.go @@ -2,15 +2,20 @@ package interceptor import ( "context" + "errors" "testing" + "time" + "github.com/nexus-rpc/sdk-go/nexus" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "go.temporal.io/api/workflowservice/v1" "go.temporal.io/server/common/dynamicconfig" "go.temporal.io/server/common/log" "go.temporal.io/server/common/namespace" "go.temporal.io/server/common/quotas/calculator" "go.temporal.io/server/common/quotas/quotastest" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.uber.org/mock/gomock" "google.golang.org/grpc" ) @@ -139,6 +144,57 @@ func TestNamespaceCountLimitInterceptor_Intercept(t *testing.T) { } } +func TestConcurrentRequestLimitInterceptor_InterceptNexus(t *testing.T) { + interceptor := NewConcurrentRequestLimitInterceptor( + nil, + quotastest.NewFakeMemberCounter(1), + log.NewNoopLogger(), + dynamicconfig.GetIntPropertyFnFilteredByNamespace(1), + dynamicconfig.GetIntPropertyFnFilteredByNamespace(1), + map[string]int{"NexusAPI": 1}, + ) + input := interceptornexus.NewStartOpInput( + "s", "o", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, + interceptornexus.ForwardingInfo{}, + interceptornexus.RequestMetadata{APIName: "NexusAPI"}, + ) + + ctx := context.Background() + + blockUntilFirstReqStarted := make(chan struct{}) + unblockFirstRequest := make(chan struct{}) + firstReqErrorCh := make(chan error, 1) + + go func() { + _, err := interceptor.InterceptNexus( + ctx, + input, + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + close(blockUntilFirstReqStarted) + <-unblockFirstRequest + return nil, nil + }, + ) + firstReqErrorCh <- err + }() + <-blockUntilFirstReqStarted + // second req should never proceed to calling next + _, err := interceptor.InterceptNexus( + ctx, + input, + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + t.Fatal("second request reached handler") + return nil, errors.New("throttled request reached") + }, + ) + var interceptorErr *interceptornexus.InterceptorError + require.ErrorAs(t, err, &interceptorErr) + require.Equal(t, "namespace_concurrency_limited", interceptorErr.Outcome) + + close(unblockFirstRequest) + require.NoError(t, <-firstReqErrorCh) +} + // run the test case by simulating a bunch of blocked pollers, sending a final request, and verifying that it is either // rate limited or not. func (tc *nsCountLimitTestCase) run(t *testing.T) { diff --git a/common/rpc/interceptor/context_metadata_interceptor.go b/common/rpc/interceptor/context_metadata_interceptor.go index 1368d417d42..d701c6ba46d 100644 --- a/common/rpc/interceptor/context_metadata_interceptor.go +++ b/common/rpc/interceptor/context_metadata_interceptor.go @@ -8,6 +8,7 @@ import ( "go.temporal.io/server/common/contextutil" "go.temporal.io/server/common/log" "go.temporal.io/server/common/log/tag" + "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc" "google.golang.org/grpc/metadata" "google.golang.org/protobuf/proto" @@ -56,13 +57,21 @@ func (c *ContextMetadataInterceptor) Intercept( resp, err := handler(ctx, req) if c.setTrailer { - c.appendContextMetadataToTrailer(ctx, info) + c.appendContextMetadataToTrailer(ctx, info.FullMethod) } return resp, err } -func (c *ContextMetadataInterceptor) appendContextMetadataToTrailer(ctx context.Context, info *grpc.UnaryServerInfo) { +func (c *ContextMetadataInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + return next(ctx, in) +} + +func (c *ContextMetadataInterceptor) appendContextMetadataToTrailer(ctx context.Context, method string) { // If the context is done, the gRPC stream may already be in streamDone state, // and SetTrailer would return ErrIllegalHeaderWrite ("SendHeader called multiple times"). select { @@ -74,7 +83,7 @@ func (c *ContextMetadataInterceptor) appendContextMetadataToTrailer(ctx context. allMetadata := contextutil.ContextMetadataGetAll(ctx) if len(allMetadata) == 0 { c.throttledLogger.Info("ContextMetadataInterceptor: No metadata in context, not setting trailer", - tag.NewStringTag("fullMethod", info.FullMethod), + tag.NewStringTag("fullMethod", method), ) return } @@ -84,13 +93,13 @@ func (c *ContextMetadataInterceptor) appendContextMetadataToTrailer(ctx context. trailer := metadata.Pairs(trailerPairs...) c.throttledLogger.Info("ContextMetadataInterceptor: Setting trailer", tag.NewAnyTag("trailer", trailer), - tag.NewStringTag("fullMethod", info.FullMethod), + tag.NewStringTag("fullMethod", method), ) if err := grpc.SetTrailer(ctx, trailer); err != nil { c.logger.Error("ContextMetadataInterceptor: Failed to set trailer", tag.Error(err), - tag.NewStringTag("fullMethod", info.FullMethod)) + tag.NewStringTag("fullMethod", method)) } } diff --git a/common/rpc/interceptor/context_metadata_interceptor_test.go b/common/rpc/interceptor/context_metadata_interceptor_test.go index a359a5b414e..c16cc10a844 100644 --- a/common/rpc/interceptor/context_metadata_interceptor_test.go +++ b/common/rpc/interceptor/context_metadata_interceptor_test.go @@ -197,7 +197,7 @@ func TestContextMetadataInterceptor_appendContextMetadataToTrailer(t *testing.T) info := &grpc.UnaryServerInfo{ FullMethod: "/test.Service/TestMethod", } - interceptor.appendContextMetadataToTrailer(ctx, info) + interceptor.appendContextMetadataToTrailer(ctx, info.FullMethod) }) } } diff --git a/common/rpc/interceptor/frontend_service_error.go b/common/rpc/interceptor/frontend_service_error.go index faf1422b59b..92720491ec7 100644 --- a/common/rpc/interceptor/frontend_service_error.go +++ b/common/rpc/interceptor/frontend_service_error.go @@ -2,11 +2,13 @@ package interceptor import ( "context" + "errors" "go.temporal.io/api/serviceerror" "go.temporal.io/server/common/api" "go.temporal.io/server/common/log" "go.temporal.io/server/common/log/tag" + "go.temporal.io/server/common/rpc/interceptor/nexus" serviceerrors "go.temporal.io/server/common/serviceerror" "google.golang.org/grpc" "google.golang.org/grpc/metadata" @@ -20,41 +22,77 @@ const ( ResourceExhaustedScopeHeader = "X-Resource-Exhausted-Scope" ) -// NewFrontendServiceErrorInterceptor returns a gRPC interceptor that has two responsibilities: +type FrontendServiceErrorInterceptor struct { + logger log.Logger +} + +// NewFrontendServiceErrorInterceptorWrapper returns interceptors that have two responsibilities: // 1. Mask certain internal service error details. // 2. Propagate resource exhaustion details via gRPC headers. -func NewFrontendServiceErrorInterceptor( - logger log.Logger, -) grpc.UnaryServerInterceptor { - return func( - ctx context.Context, - req any, - info *grpc.UnaryServerInfo, - handler grpc.UnaryHandler, - ) (any, error) { - resp, err := handler(ctx, req) - if err == nil { - return resp, nil - } +func NewFrontendServiceErrorInterceptorWrapper(logger log.Logger) *FrontendServiceErrorInterceptor { + return &FrontendServiceErrorInterceptor{ + logger: logger, + } +} - switch serviceErr := err.(type) { - case *serviceerrors.ShardOwnershipLost: - err = serviceerror.NewUnavailable("shard unavailable, please backoff and retry") - case *serviceerror.DataLoss: - err = serviceerror.NewUnavailable("internal history service error") - case *serviceerror.ResourceExhausted: - if headerErr := grpc.SetHeader(ctx, metadata.Pairs( - ResourceExhaustedCauseHeader, serviceErr.Cause.String(), - ResourceExhaustedScopeHeader, serviceErr.Scope.String(), - )); headerErr != nil { - // So while this is *not* a user-facing error or problem in itself, - // it indicates that there might be larger connection issues at play. - logger.Error("Failed to add Resource-Exhausted headers to response", - tag.Operation(api.MethodName(info.FullMethod)), - tag.Error(headerErr)) - } - } +// NewFrontendServiceErrorInterceptor provides the legacy standalone gRPC Interceptor for existing deployments. +// +// Deprecated: use the unified [NewFrontendServiceErrorInterceptorWrapper] instead. +func NewFrontendServiceErrorInterceptor(logger log.Logger) grpc.UnaryServerInterceptor { + t := NewFrontendServiceErrorInterceptorWrapper(logger) + return t.Intercept +} - return resp, err +func (f *FrontendServiceErrorInterceptor) Intercept( + ctx context.Context, + req any, + info *grpc.UnaryServerInfo, + handler grpc.UnaryHandler, +) (any, error) { + resp, err := handler(ctx, req) + + return resp, f.transformError(ctx, info.FullMethod, err, true) +} + +func (f *FrontendServiceErrorInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + resp, err := next(ctx, in) + if ie, ok := errors.AsType[*nexus.InterceptorError](err); ok { + ie.Err = f.transformError(ctx, in.APIName(), ie.Err, false) + return resp, ie + } + return resp, f.transformError(ctx, in.APIName(), err, false) +} + +func (f *FrontendServiceErrorInterceptor) transformError(ctx context.Context, method string, err error, isGRPC bool) error { + if err == nil { + return nil + } + method = api.MethodName(method) + + switch serviceErr := err.(type) { + case *serviceerrors.ShardOwnershipLost: + err = serviceerror.NewUnavailable("shard unavailable, please backoff and retry") + case *serviceerror.DataLoss: + err = serviceerror.NewUnavailable("internal history service error") + case *serviceerror.ResourceExhausted: + if !isGRPC { + break + } + if headerErr := grpc.SetHeader(ctx, metadata.Pairs( + ResourceExhaustedCauseHeader, serviceErr.Cause.String(), + ResourceExhaustedScopeHeader, serviceErr.Scope.String(), + )); headerErr != nil { + // So while this is *not* a user-facing error or problem in itself, + // it indicates that there might be larger connection issues at play. + f.logger.Error("Failed to add Resource-Exhausted headers to response", + tag.Operation(method), + tag.Error(headerErr)) + } + default: } + return err } diff --git a/common/rpc/interceptor/frontend_service_error_test.go b/common/rpc/interceptor/frontend_service_error_test.go index f50e0e64873..b55ff690491 100644 --- a/common/rpc/interceptor/frontend_service_error_test.go +++ b/common/rpc/interceptor/frontend_service_error_test.go @@ -105,7 +105,7 @@ func TestFrontendServiceErrorInterceptor(t *testing.T) { } ctx := grpc.NewContextWithServerTransportStream(context.Background(), stream) - var interceptorFn = NewFrontendServiceErrorInterceptor(tl) + var interceptorFn = NewFrontendServiceErrorInterceptorWrapper(tl).Intercept info := &grpc.UnaryServerInfo{FullMethod: method} _, err := interceptorFn(ctx, nil, info, func(_ context.Context, _ any) (any, error) { diff --git a/common/rpc/interceptor/health.go b/common/rpc/interceptor/health.go index 3023dd6ff5a..c650b3598ce 100644 --- a/common/rpc/interceptor/health.go +++ b/common/rpc/interceptor/health.go @@ -7,6 +7,7 @@ import ( "go.temporal.io/api/serviceerror" "go.temporal.io/server/common/api" + "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc" ) @@ -32,16 +33,31 @@ func (i *HealthInterceptor) Intercept( info *grpc.UnaryServerInfo, handler grpc.UnaryHandler, ) (any, error) { - // only enforce health check on WorkflowService and OperatorService - if strings.HasPrefix(info.FullMethod, api.WorkflowServicePrefix) || - strings.HasPrefix(info.FullMethod, api.OperatorServicePrefix) { - if !i.healthy.Load() { - return nil, notHealthyErr - } + if i.isNotHealthy(info.FullMethod) { + return nil, notHealthyErr } return handler(ctx, req) } +// InterceptNexus is a no-op as nexus APIs are considered internal +func (i *HealthInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + return next(ctx, in) +} + +func (i *HealthInterceptor) isNotHealthy(methodName string) bool { + if i.healthy.Load() { + return false + } + + // only enforce health check on WorkflowService and OperatorService + return strings.HasPrefix(methodName, api.WorkflowServicePrefix) || + strings.HasPrefix(methodName, api.OperatorServicePrefix) +} + func (i *HealthInterceptor) SetHealthy(healthy bool) { i.healthy.Store(healthy) } diff --git a/common/rpc/interceptor/mask_internal_error.go b/common/rpc/interceptor/mask_internal_error.go index 0676dcdbb96..62bdf02b313 100644 --- a/common/rpc/interceptor/mask_internal_error.go +++ b/common/rpc/interceptor/mask_internal_error.go @@ -2,8 +2,10 @@ package interceptor import ( "context" + "errors" "fmt" + "github.com/nexus-rpc/sdk-go/nexus" "go.temporal.io/api/serviceerror" "go.temporal.io/server/common" "go.temporal.io/server/common/api" @@ -12,6 +14,7 @@ import ( "go.temporal.io/server/common/log/tag" "go.temporal.io/server/common/namespace" "go.temporal.io/server/common/rpc/interceptor/logtags" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.temporal.io/server/common/tasktoken" "google.golang.org/grpc" "google.golang.org/grpc/codes" @@ -56,6 +59,26 @@ func (mi *MaskInternalErrorDetailsInterceptor) Intercept( return resp, err } +func (mi *MaskInternalErrorDetailsInterceptor) InterceptNexus( + ctx context.Context, + in interceptornexus.InterceptorInput, + next interceptornexus.HandlerFunc, +) (any, error) { + + resp, err := next(ctx, in) + + if err == nil || !mi.shouldMaskErrors(in) { + return resp, err + } + if ie, ok := errors.AsType[*interceptornexus.InterceptorError](err); ok { + ie.Err = mi.maskNexusError(in, ie.Err) + err = ie + } else { + err = mi.maskNexusError(in, err) + } + return resp, err +} + func (mi *MaskInternalErrorDetailsInterceptor) shouldMaskErrors(req any) bool { ns := MustGetNamespaceName(mi.namespaceRegistry, req) if ns.IsEmpty() { @@ -64,6 +87,19 @@ func (mi *MaskInternalErrorDetailsInterceptor) shouldMaskErrors(req any) bool { return mi.maskInternalError(ns.String()) } +func (mi *MaskInternalErrorDetailsInterceptor) maskNexusError(in interceptornexus.InterceptorInput, err error) error { + if _, ok := errors.AsType[*nexus.HandlerError](err); ok { + return err + } + if _, ok := errors.AsType[*nexus.OperationError](err); ok { + return err + } + if _, ok := common.GetRPCStatus(err); !ok { + return err + } + return mi.maskUnknownOrInternalErrors(in, in.APIName(), err) +} + func (mi *MaskInternalErrorDetailsInterceptor) maskUnknownOrInternalErrors( req any, fullMethodName string, err error, ) error { diff --git a/common/rpc/interceptor/namespace.go b/common/rpc/interceptor/namespace.go index fdeded48add..11ba12a25b0 100644 --- a/common/rpc/interceptor/namespace.go +++ b/common/rpc/interceptor/namespace.go @@ -4,6 +4,7 @@ import ( "go.temporal.io/api/serviceerror" "go.temporal.io/api/workflowservice/v1" "go.temporal.io/server/common/namespace" + "go.temporal.io/server/common/rpc/interceptor/nexus" ) // gRPC method request must implement either NamespaceNameGetter or NamespaceIDGetter @@ -57,6 +58,13 @@ func GetNamespaceName( } return namespaceName, nil + case nexus.InterceptorInput: + ns, err := request.NamespaceEntry() + if err != nil { + return namespace.EmptyName, err + } + return ns.Name(), nil + default: return namespace.EmptyName, serviceerror.NewInternalf("unable to extract namespace info from request of type %T", req) } diff --git a/common/rpc/interceptor/namespace_handover.go b/common/rpc/interceptor/namespace_handover.go index 4d1962d683a..6328f2554f2 100644 --- a/common/rpc/interceptor/namespace_handover.go +++ b/common/rpc/interceptor/namespace_handover.go @@ -14,6 +14,7 @@ import ( "go.temporal.io/server/common/log" "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/namespace" + "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc" ) @@ -91,6 +92,16 @@ func (i *NamespaceHandoverInterceptor) handlesMethod(fullMethod string) bool { return false } +// InterceptNexus is a no-op: the handover gate only applies to WorkflowService +// methods- see [NamespaceHandoverInterceptor.handlesMethod] for details. +func (i *NamespaceHandoverInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + return next(ctx, in) +} + func (i *NamespaceHandoverInterceptor) Intercept( ctx context.Context, req any, diff --git a/common/rpc/interceptor/namespace_logger.go b/common/rpc/interceptor/namespace_logger.go index 8f5a99f1ec8..9b2365fd006 100644 --- a/common/rpc/interceptor/namespace_logger.go +++ b/common/rpc/interceptor/namespace_logger.go @@ -6,10 +6,11 @@ import ( "fmt" "go.temporal.io/server/common/api" - "go.temporal.io/server/common/authorization" "go.temporal.io/server/common/log" "go.temporal.io/server/common/log/tag" "go.temporal.io/server/common/namespace" + "go.temporal.io/server/common/rpc/interceptor/nexus" + "go.temporal.io/server/common/rpc/tlsinfo" "google.golang.org/grpc" ) @@ -40,12 +41,12 @@ func (nli *NamespaceLogInterceptor) Intercept( if nli.logger != nil { methodName := api.MethodName(info.FullMethod) namespace := MustGetNamespaceName(nli.namespaceRegistry, req) - tlsInfo := authorization.TLSInfoFromContext(ctx) + tlsInfo := tlsinfo.FromContext(ctx) var serverName string var certThumbprint string if tlsInfo != nil { serverName = tlsInfo.State.ServerName - cert := authorization.PeerCert(tlsInfo) + cert := tlsinfo.PeerCert(tlsInfo) if cert != nil { certThumbprint = fmt.Sprintf("%x", md5.Sum(cert.Raw)) } @@ -59,3 +60,31 @@ func (nli *NamespaceLogInterceptor) Intercept( } return handler(ctx, req) } + +func (nli *NamespaceLogInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + if nli.logger != nil { + methodName := api.MethodName(in.APIName()) + namespaceName := MustGetNamespaceName(nli.namespaceRegistry, in) + tlsInfo := tlsinfo.FromContext(ctx) + var serverName string + var certThumbprint string + if tlsInfo != nil { + serverName = tlsInfo.State.ServerName + cert := tlsinfo.PeerCert(tlsInfo) + if cert != nil { + certThumbprint = fmt.Sprintf("%x", md5.Sum(cert.Raw)) + } + } + nli.logger.Debug( + "Frontend method invoked.", + tag.WorkflowNamespace(namespaceName.String()), + tag.Operation(methodName), + tag.ServerName(serverName), + tag.CertThumbprint(certThumbprint)) + } + return next(ctx, in) +} diff --git a/common/rpc/interceptor/namespace_rate_limit.go b/common/rpc/interceptor/namespace_rate_limit.go index 080e9f667b6..d7f29f9f3a9 100644 --- a/common/rpc/interceptor/namespace_rate_limit.go +++ b/common/rpc/interceptor/namespace_rate_limit.go @@ -13,6 +13,7 @@ import ( "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/namespace" "go.temporal.io/server/common/quotas" + "go.temporal.io/server/common/rpc/interceptor/nexus" "go.temporal.io/server/service/frontend/configs" "google.golang.org/grpc" ) @@ -90,6 +91,41 @@ type ( var _ grpc.UnaryServerInterceptor = (*NamespaceRateLimitInterceptorImpl)(nil).Intercept var _ NamespaceRateLimitInterceptor = (*NamespaceRateLimitInterceptorImpl)(nil) +func NewNamespaceRateLimitInterceptorWrapper(ni NamespaceRateLimitInterceptor) *NamespaceRateLimitInterceptorWrapper { + return &NamespaceRateLimitInterceptorWrapper{ + ni: ni, + } +} + +// NamespaceRateLimitInterceptorWrapper is a wrapper on namespace rate limiter +type NamespaceRateLimitInterceptorWrapper struct { + ni NamespaceRateLimitInterceptor +} + +func (n *NamespaceRateLimitInterceptorWrapper) Intercept( + ctx context.Context, + req any, + info *grpc.UnaryServerInfo, + handler grpc.UnaryHandler, +) (resp any, err error) { + return n.ni.Intercept(ctx, req, info, handler) +} + +func (n *NamespaceRateLimitInterceptorWrapper) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (out any, retErr error) { + if err := n.ni.Allow(ctx, namespace.Name(in.NamespaceName()), in.APIName(), in.Header()); err != nil { + return nil, &nexus.InterceptorError{ + Err: err, + Outcome: "namespace_rate_limited", + ExposeDetails: true, + } + } + return next(ctx, in) +} + func NewNamespaceRateLimitInterceptor( namespaceRegistry namespace.Registry, rateLimiter quotas.RequestRateLimiter, diff --git a/common/rpc/interceptor/namespace_rate_limit_test.go b/common/rpc/interceptor/namespace_rate_limit_test.go index ab90219ccd4..b90d761e377 100644 --- a/common/rpc/interceptor/namespace_rate_limit_test.go +++ b/common/rpc/interceptor/namespace_rate_limit_test.go @@ -5,6 +5,7 @@ import ( "testing" "time" + "github.com/nexus-rpc/sdk-go/nexus" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" "go.temporal.io/api/workflowservice/v1" @@ -12,6 +13,7 @@ import ( "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/namespace" "go.temporal.io/server/common/quotas" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.uber.org/mock/gomock" "google.golang.org/grpc" ) @@ -31,6 +33,42 @@ type namespaceRateLimitInterceptorSuite struct { mockRegistry *namespace.MockRegistry } +func (s *namespaceRateLimitInterceptorSuite) TestInterceptNexus() { + for _, tc := range []struct { + name string + input interceptornexus.InterceptorInput + allow bool + nextCalled bool + expectedOutcome string + }{ + {name: "allowed", input: interceptornexus.NewStartOpInput("s", "o", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, interceptornexus.ForwardingInfo{}, interceptornexus.RequestMetadata{APIName: "NexusOperation"}), allow: true, nextCalled: true}, + {name: "rate limited", input: interceptornexus.NewStartOpInput("s", "o", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, interceptornexus.ForwardingInfo{}, interceptornexus.RequestMetadata{APIName: "NexusOperation"}), expectedOutcome: "namespace_rate_limited"}, + } { + s.Run(tc.name, func() { + ctx := context.Background() + s.mockRateLimiter.EXPECT().Allow(gomock.Any(), gomock.Any()).Return(tc.allow) + nextCalled := false + wrapper := NewNamespaceRateLimitInterceptorWrapper(s.newImpl(false)) + _, err := wrapper.InterceptNexus( + ctx, + tc.input, + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + nextCalled = true + return nil, nil + }, + ) + if tc.expectedOutcome != "" { + var interceptorErr *interceptornexus.InterceptorError + s.ErrorAs(err, &interceptorErr) + s.Equal(tc.expectedOutcome, interceptorErr.Outcome) + } else { + s.NoError(err) + } + s.Equal(tc.nextCalled, nextCalled) + }) + } +} + func TestNamespaceRateLimitInterceptorSuite(t *testing.T) { suite.Run(t, &namespaceRateLimitInterceptorSuite{}) } diff --git a/common/rpc/interceptor/namespace_validator.go b/common/rpc/interceptor/namespace_validator.go index 6928e5e0268..4a4c1e02595 100644 --- a/common/rpc/interceptor/namespace_validator.go +++ b/common/rpc/interceptor/namespace_validator.go @@ -12,6 +12,7 @@ import ( "go.temporal.io/server/common/api" "go.temporal.io/server/common/dynamicconfig" "go.temporal.io/server/common/namespace" + "go.temporal.io/server/common/rpc/interceptor/nexus" "go.temporal.io/server/common/tasktoken" "google.golang.org/grpc" ) @@ -29,6 +30,15 @@ type ( maxNamespaceLength dynamicconfig.IntPropertyFn additionalAllowedMethodsDuringHandover map[string]struct{} } + + // NamespaceLengthValidatorInterceptor enforces the namespace name length limit. It is separate + // from NamespaceValidatorInterceptor to allow both to expose cleaner Intercept/InterceptNexus + // methods that are used as gRPC and Nexus interceptors. + NamespaceLengthValidatorInterceptor struct { + namespaceRegistry namespace.Registry + tokenSerializer *tasktoken.Serializer + maxNamespaceLength dynamicconfig.IntPropertyFn + } ) var ( @@ -86,8 +96,8 @@ var ( } ) -var _ grpc.UnaryServerInterceptor = (*NamespaceValidatorInterceptor)(nil).StateValidationIntercept -var _ grpc.UnaryServerInterceptor = (*NamespaceValidatorInterceptor)(nil).NamespaceValidateIntercept +var _ grpc.UnaryServerInterceptor = (*NamespaceValidatorInterceptor)(nil).Intercept +var _ grpc.UnaryServerInterceptor = (*NamespaceLengthValidatorInterceptor)(nil).Intercept func NewNamespaceValidatorInterceptor( namespaceRegistry namespace.Registry, @@ -108,26 +118,59 @@ func NewNamespaceValidatorInterceptor( } } -func (ni *NamespaceValidatorInterceptor) NamespaceValidateIntercept( +func NewNamespaceLengthValidatorInterceptor( + namespaceRegistry namespace.Registry, + maxNamespaceLength dynamicconfig.IntPropertyFn, +) *NamespaceLengthValidatorInterceptor { + return &NamespaceLengthValidatorInterceptor{ + namespaceRegistry: namespaceRegistry, + tokenSerializer: tasktoken.NewSerializer(), + maxNamespaceLength: maxNamespaceLength, + } +} + +func (nsvi *NamespaceLengthValidatorInterceptor) Intercept( ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler, ) (any, error) { - err := ni.setNamespaceIfNotPresent(req) + err := setNamespaceIfNotPresent(nsvi.tokenSerializer, nsvi.namespaceRegistry, req) if err != nil { return nil, err } reqWithNamespace, hasNamespace := req.(NamespaceNameGetter) - if hasNamespace { - if err := ni.ValidateName(reqWithNamespace.GetNamespace()); err != nil { - return nil, err - } + if hasNamespace && len(reqWithNamespace.GetNamespace()) > nsvi.maxNamespaceLength() { + return nil, errNamespaceTooLong } return handler(ctx, req) } +func (nsvi *NamespaceLengthValidatorInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + ns, err := in.NamespaceEntry() + if err != nil { + return nil, &nexus.InterceptorError{ + Err: err, + Outcome: "interceptor_failed", + SkipServiceErrorReporting: true, + } + } + if len(ns.Info().GetName()) > nsvi.maxNamespaceLength() { + return nil, &nexus.InterceptorError{ + Err: errNamespaceTooLong, + Outcome: "interceptor_failed", + SkipServiceErrorReporting: true, + } + } + + return next(ctx, in) +} + // ValidateName validates a namespace name (currently only a max length check). func (ni *NamespaceValidatorInterceptor) ValidateName(ns string) error { if len(ns) > ni.maxNamespaceLength() { @@ -136,17 +179,19 @@ func (ni *NamespaceValidatorInterceptor) ValidateName(ns string) error { return nil } -func (ni *NamespaceValidatorInterceptor) setNamespaceIfNotPresent( +func setNamespaceIfNotPresent( + tokenSerializer *tasktoken.Serializer, + namespaceRegistry namespace.Registry, req any, ) error { switch request := req.(type) { case NamespaceNameGetter: if request.GetNamespace() == "" { - namespaceEntry, err := ni.extractNamespaceFromTaskToken(req) + namespaceEntry, err := extractNamespaceFromTaskToken(tokenSerializer, namespaceRegistry, req) if err != nil { return err } - ni.setNamespace(namespaceEntry, req) + setNamespace(namespaceEntry, req) } return nil default: @@ -154,7 +199,7 @@ func (ni *NamespaceValidatorInterceptor) setNamespaceIfNotPresent( } } -func (ni *NamespaceValidatorInterceptor) setNamespace( +func setNamespace( namespaceEntry *namespace.Namespace, req any, ) { @@ -198,8 +243,8 @@ func (ni *NamespaceValidatorInterceptor) setNamespace( } } -// StateValidationIntercept runs ValidateState - see docstring for that method. -func (ni *NamespaceValidatorInterceptor) StateValidationIntercept( +// Intercept runs ValidateState - see docstring for that method. +func (ni *NamespaceValidatorInterceptor) Intercept( ctx context.Context, req any, info *grpc.UnaryServerInfo, @@ -230,9 +275,33 @@ func (ni *NamespaceValidatorInterceptor) ValidateState(namespaceEntry *namespace return ni.checkReplicationState(namespaceEntry, fullMethod, businessID) } +// InterceptNexus validates the namespace state for a Nexus request. +func (ni *NamespaceValidatorInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + namespaceEntry, err := in.NamespaceEntry() + if err != nil { + return nil, &nexus.InterceptorError{ + Err: err, + Outcome: "interceptor_failed", + SkipServiceErrorReporting: true, + } + } + if err := ni.ValidateState(namespaceEntry, in.APIName(), in.ForwardingInfo().BusinessID); err != nil { + return nil, &nexus.InterceptorError{ + Err: err, + Outcome: "invalid_namespace_state", + SkipServiceErrorReporting: true, + } + } + return next(ctx, in) +} + func (ni *NamespaceValidatorInterceptor) extractNamespace(req any) (*namespace.Namespace, error) { // Token namespace has priority over request namespace. Check it first. - tokenNamespaceEntry, tokenErr := ni.extractNamespaceFromTaskToken(req) + tokenNamespaceEntry, tokenErr := extractNamespaceFromTaskToken(ni.tokenSerializer, ni.namespaceRegistry, req) if tokenErr != nil { return nil, tokenErr } @@ -321,7 +390,11 @@ func (ni *NamespaceValidatorInterceptor) extractNamespaceFromRequest(req any) (* } } -func (ni *NamespaceValidatorInterceptor) extractNamespaceFromTaskToken(req any) (*namespace.Namespace, error) { +func extractNamespaceFromTaskToken( + tokenSerializer *tasktoken.Serializer, + namespaceRegistry namespace.Registry, + req any, +) (*namespace.Namespace, error) { reqWithTaskToken, hasTaskToken := req.(TaskTokenGetter) if !hasTaskToken { return nil, nil @@ -333,13 +406,13 @@ func (ni *NamespaceValidatorInterceptor) extractNamespaceFromTaskToken(req any) var namespaceID namespace.ID // Special case for deprecated RespondQueryTaskCompleted API. if _, ok := req.(*workflowservice.RespondQueryTaskCompletedRequest); ok { - taskToken, err := ni.tokenSerializer.DeserializeQueryTaskToken(taskTokenBytes) + taskToken, err := tokenSerializer.DeserializeQueryTaskToken(taskTokenBytes) if err != nil { return nil, errDeserializingToken } namespaceID = namespace.ID(taskToken.GetNamespaceId()) } else { - taskToken, err := ni.tokenSerializer.Deserialize(taskTokenBytes) + taskToken, err := tokenSerializer.Deserialize(taskTokenBytes) if err != nil { return nil, errDeserializingToken } @@ -349,7 +422,7 @@ func (ni *NamespaceValidatorInterceptor) extractNamespaceFromTaskToken(req any) if namespaceID.IsEmpty() { return nil, errNamespaceNotSet } - return ni.namespaceRegistry.GetNamespaceByID(namespaceID) + return namespaceRegistry.GetNamespaceByID(namespaceID) } func (ni *NamespaceValidatorInterceptor) checkNamespaceMatch(requestNamespace *namespace.Namespace, tokenNamespace *namespace.Namespace) error { diff --git a/common/rpc/interceptor/namespace_validator_test.go b/common/rpc/interceptor/namespace_validator_test.go index f7599e0d44f..f83d02c5715 100644 --- a/common/rpc/interceptor/namespace_validator_test.go +++ b/common/rpc/interceptor/namespace_validator_test.go @@ -4,8 +4,10 @@ import ( "context" "fmt" "testing" + "time" "github.com/google/uuid" + "github.com/nexus-rpc/sdk-go/nexus" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" enumspb "go.temporal.io/api/enums/v1" @@ -19,6 +21,7 @@ import ( "go.temporal.io/server/common/api" "go.temporal.io/server/common/dynamicconfig" "go.temporal.io/server/common/namespace" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.temporal.io/server/common/tasktoken" "go.uber.org/mock/gomock" "google.golang.org/grpc" @@ -96,7 +99,7 @@ func (s *namespaceValidatorSuite) Test_StateValidationIntercept_NamespaceNotSet( for _, testCase := range testCases { handlerCalled := false - _, err := nvi.StateValidationIntercept(context.Background(), testCase.req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err := nvi.Intercept(context.Background(), testCase.req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.StartWorkflowExecutionResponse{}, nil }) @@ -111,6 +114,88 @@ func (s *namespaceValidatorSuite) Test_StateValidationIntercept_NamespaceNotSet( } } +func (s *namespaceValidatorSuite) TestInterceptNexus() { + validator := NewNamespaceValidatorInterceptor( + s.mockRegistry, + dynamicconfig.GetBoolPropertyFn(false), + dynamicconfig.GetIntPropertyFn(100), + nil, + ) + for _, tc := range []struct { + name string + input interceptornexus.InterceptorInput + nextCalled bool + expectedOutcome string + }{ + { + name: "resolved namespace", + input: interceptornexus.NewStartOpInput( + "s", "o", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, + interceptornexus.ForwardingInfo{}, + interceptornexus.RequestMetadata{ + APIName: api.NexusServicePrefix + "DispatchNexusTask", + NamespaceEntry: namespace.NewNamespaceForTest( + &persistencespb.NamespaceInfo{Name: testNamespace, State: enumspb.NAMESPACE_STATE_REGISTERED}, + nil, + false, + nil, + 0, + ), + }, + ), + nextCalled: true, + }, + { + name: "invalid namespace state", + input: interceptornexus.NewStartOpInput( + "s", "o", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, + interceptornexus.ForwardingInfo{}, + interceptornexus.RequestMetadata{ + APIName: api.NexusServicePrefix + "DispatchNexusTask", + NamespaceEntry: namespace.NewNamespaceForTest( + &persistencespb.NamespaceInfo{Name: testNamespace, State: enumspb.NAMESPACE_STATE_DEPRECATED}, + nil, + false, + nil, + 0, + ), + }, + ), + expectedOutcome: "invalid_namespace_state", + }, + { + name: "missing namespace", + input: interceptornexus.NewStartOpInput( + "s", "o", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, + interceptornexus.ForwardingInfo{}, + interceptornexus.RequestMetadata{APIName: "NexusAPI"}, + ), + expectedOutcome: "interceptor_failed", + }, + } { + s.Run(tc.name, func() { + nextCalled := false + _, err := validator.InterceptNexus( + context.Background(), + tc.input, + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + nextCalled = true + return nil, nil + }, + ) + if tc.expectedOutcome != "" { + var interceptorErr *interceptornexus.InterceptorError + s.ErrorAs(err, &interceptorErr) + s.Equal(tc.expectedOutcome, interceptorErr.Outcome) + s.True(interceptorErr.SkipServiceErrorReporting) + } else { + s.NoError(err) + } + s.Equal(tc.nextCalled, nextCalled) + }) + } +} + func (s *namespaceValidatorSuite) Test_StateValidationIntercept_NamespaceNotFound() { nvi := NewNamespaceValidatorInterceptor( @@ -125,7 +210,7 @@ func (s *namespaceValidatorSuite) Test_StateValidationIntercept_NamespaceNotFoun s.mockRegistry.EXPECT().GetNamespace(namespace.Name("not-found-namespace")).Return(nil, serviceerror.NewNamespaceNotFound("missing-namespace")) req := &workflowservice.StartWorkflowExecutionRequest{Namespace: "not-found-namespace"} handlerCalled := false - _, err := nvi.StateValidationIntercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err := nvi.Intercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.StartWorkflowExecutionResponse{}, nil }) @@ -142,7 +227,7 @@ func (s *namespaceValidatorSuite) Test_StateValidationIntercept_NamespaceNotFoun TaskToken: taskToken, } handlerCalled = false - _, err = nvi.StateValidationIntercept(context.Background(), tokenReq, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err = nvi.Intercept(context.Background(), tokenReq, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.RespondWorkflowTaskCompletedResponse{}, nil }) @@ -399,7 +484,7 @@ func (s *namespaceValidatorSuite) Test_StateValidationIntercept_StatusFromNamesp } handlerCalled := false - _, err := nvi.StateValidationIntercept(context.Background(), testCase.req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err := nvi.Intercept(context.Background(), testCase.req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.StartWorkflowExecutionResponse{}, nil }) @@ -476,7 +561,7 @@ func (s *namespaceValidatorSuite) Test_StateValidationIntercept_StatusFromToken( } handlerCalled := false - _, err := nvi.StateValidationIntercept(context.Background(), testCase.req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err := nvi.Intercept(context.Background(), testCase.req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.RespondWorkflowTaskCompletedResponse{}, nil }) @@ -503,7 +588,7 @@ func (s *namespaceValidatorSuite) Test_StateValidationIntercept_DescribeNamespac req := &workflowservice.DescribeNamespaceRequest{Id: "test-namespace-id"} handlerCalled := false - _, err := nvi.StateValidationIntercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err := nvi.Intercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.DescribeNamespaceResponse{}, nil }) @@ -513,7 +598,7 @@ func (s *namespaceValidatorSuite) Test_StateValidationIntercept_DescribeNamespac req = &workflowservice.DescribeNamespaceRequest{} handlerCalled = false - _, err = nvi.StateValidationIntercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err = nvi.Intercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.DescribeNamespaceResponse{}, nil }) @@ -535,7 +620,7 @@ func (s *namespaceValidatorSuite) Test_StateValidationIntercept_GetClusterInfo() // Example of API which doesn't have namespace field. req := &workflowservice.GetClusterInfoRequest{} handlerCalled := false - _, err := nvi.StateValidationIntercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err := nvi.Intercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.GetClusterInfoResponse{}, nil }) @@ -556,7 +641,7 @@ func (s *namespaceValidatorSuite) Test_Intercept_RegisterNamespace() { req := &workflowservice.RegisterNamespaceRequest{Namespace: "new-namespace"} handlerCalled := false - _, err := nvi.StateValidationIntercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err := nvi.Intercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.RegisterNamespaceResponse{}, nil }) @@ -566,7 +651,7 @@ func (s *namespaceValidatorSuite) Test_Intercept_RegisterNamespace() { req = &workflowservice.RegisterNamespaceRequest{} handlerCalled = false - _, err = nvi.StateValidationIntercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err = nvi.Intercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.RegisterNamespaceResponse{}, nil }) @@ -669,11 +754,11 @@ func (s *namespaceValidatorSuite) Test_StateValidationIntercept_TokenNamespaceEn } handlerCalled := false - _, err = nvi.StateValidationIntercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err = nvi.Intercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.RespondWorkflowTaskCompletedResponse{}, nil }) - _, queryErr := nvi.StateValidationIntercept(context.Background(), queryReq, serverInfo, func(ctx context.Context, req any) (any, error) { + _, queryErr := nvi.Intercept(context.Background(), queryReq, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.RespondQueryTaskCompletedResponse{}, nil }) @@ -715,7 +800,7 @@ func (s *namespaceValidatorSuite) Test_Intercept_DescribeHistoryHostRequests() { } handlerCalled := false - _, err := nvi.StateValidationIntercept( + _, err := nvi.Intercept( context.Background(), testCase.req, serverInfo, @@ -801,7 +886,7 @@ func (s *namespaceValidatorSuite) Test_Intercept_SearchAttributeRequests() { } handlerCalled := false - _, err := nvi.StateValidationIntercept( + _, err := nvi.Intercept( context.Background(), testCase.req, serverInfo, @@ -816,11 +901,10 @@ func (s *namespaceValidatorSuite) Test_Intercept_SearchAttributeRequests() { } func (s *namespaceValidatorSuite) Test_NamespaceValidateIntercept() { - nvi := NewNamespaceValidatorInterceptor( + nnvi := NewNamespaceLengthValidatorInterceptor( s.mockRegistry, - dynamicconfig.GetBoolPropertyFn(false), dynamicconfig.GetIntPropertyFn(10), - nil) + ) serverInfo := &grpc.UnaryServerInfo{ FullMethod: api.WorkflowServicePrefix + "random", } @@ -852,7 +936,7 @@ func (s *namespaceValidatorSuite) Test_NamespaceValidateIntercept() { req := &workflowservice.StartWorkflowExecutionRequest{Namespace: "namespace"} handlerCalled := false - _, err = nvi.NamespaceValidateIntercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err = nnvi.Intercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.StartWorkflowExecutionResponse{}, nil }) @@ -861,7 +945,7 @@ func (s *namespaceValidatorSuite) Test_NamespaceValidateIntercept() { req = &workflowservice.StartWorkflowExecutionRequest{Namespace: "namespaceTooLong"} handlerCalled = false - _, err = nvi.NamespaceValidateIntercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { + _, err = nnvi.Intercept(context.Background(), req, serverInfo, func(ctx context.Context, req any) (any, error) { handlerCalled = true return &workflowservice.StartWorkflowExecutionResponse{}, nil }) @@ -885,59 +969,52 @@ func (s *namespaceValidatorSuite) TestSetNamespace() { namespaceEntry, err := namespace.FromPersistentState(detail, factory(detail)) s.NoError(err) - nvi := NewNamespaceValidatorInterceptor( - s.mockRegistry, - dynamicconfig.GetBoolPropertyFn(false), - dynamicconfig.GetIntPropertyFn(10), - nil, - ) - queryReq := &workflowservice.RespondQueryTaskCompletedRequest{} - nvi.setNamespace(namespaceEntry, queryReq) + setNamespace(namespaceEntry, queryReq) s.Equal(namespaceEntryName, queryReq.Namespace) queryReq.Namespace = namespaceRequestName - nvi.setNamespace(namespaceEntry, queryReq) + setNamespace(namespaceEntry, queryReq) s.Equal(namespaceRequestName, queryReq.Namespace) completeWorkflowTaskReq := &workflowservice.RespondWorkflowTaskCompletedRequest{} - nvi.setNamespace(namespaceEntry, completeWorkflowTaskReq) + setNamespace(namespaceEntry, completeWorkflowTaskReq) s.Equal(namespaceEntryName, completeWorkflowTaskReq.Namespace) completeWorkflowTaskReq.Namespace = namespaceRequestName - nvi.setNamespace(namespaceEntry, completeWorkflowTaskReq) + setNamespace(namespaceEntry, completeWorkflowTaskReq) s.Equal(namespaceRequestName, completeWorkflowTaskReq.Namespace) failWorkflowTaskReq := &workflowservice.RespondWorkflowTaskFailedRequest{} - nvi.setNamespace(namespaceEntry, failWorkflowTaskReq) + setNamespace(namespaceEntry, failWorkflowTaskReq) s.Equal(namespaceEntryName, failWorkflowTaskReq.Namespace) failWorkflowTaskReq.Namespace = namespaceRequestName - nvi.setNamespace(namespaceEntry, failWorkflowTaskReq) + setNamespace(namespaceEntry, failWorkflowTaskReq) s.Equal(namespaceRequestName, failWorkflowTaskReq.Namespace) heartbeatActivityTaskReq := &workflowservice.RecordActivityTaskHeartbeatRequest{} - nvi.setNamespace(namespaceEntry, heartbeatActivityTaskReq) + setNamespace(namespaceEntry, heartbeatActivityTaskReq) s.Equal(namespaceEntryName, heartbeatActivityTaskReq.Namespace) heartbeatActivityTaskReq.Namespace = namespaceRequestName - nvi.setNamespace(namespaceEntry, heartbeatActivityTaskReq) + setNamespace(namespaceEntry, heartbeatActivityTaskReq) s.Equal(namespaceRequestName, heartbeatActivityTaskReq.Namespace) cancelActivityTaskReq := &workflowservice.RespondActivityTaskCanceledRequest{} - nvi.setNamespace(namespaceEntry, cancelActivityTaskReq) + setNamespace(namespaceEntry, cancelActivityTaskReq) s.Equal(namespaceEntryName, cancelActivityTaskReq.Namespace) cancelActivityTaskReq.Namespace = namespaceRequestName - nvi.setNamespace(namespaceEntry, cancelActivityTaskReq) + setNamespace(namespaceEntry, cancelActivityTaskReq) s.Equal(namespaceRequestName, cancelActivityTaskReq.Namespace) completeActivityTaskReq := &workflowservice.RespondActivityTaskCompletedRequest{} - nvi.setNamespace(namespaceEntry, completeActivityTaskReq) + setNamespace(namespaceEntry, completeActivityTaskReq) s.Equal(namespaceEntryName, completeActivityTaskReq.Namespace) completeActivityTaskReq.Namespace = namespaceRequestName - nvi.setNamespace(namespaceEntry, completeActivityTaskReq) + setNamespace(namespaceEntry, completeActivityTaskReq) s.Equal(namespaceRequestName, completeActivityTaskReq.Namespace) failActivityTaskReq := &workflowservice.RespondActivityTaskFailedRequest{} - nvi.setNamespace(namespaceEntry, failActivityTaskReq) + setNamespace(namespaceEntry, failActivityTaskReq) s.Equal(namespaceEntryName, failActivityTaskReq.Namespace) failActivityTaskReq.Namespace = namespaceRequestName - nvi.setNamespace(namespaceEntry, failActivityTaskReq) + setNamespace(namespaceEntry, failActivityTaskReq) s.Equal(namespaceRequestName, failActivityTaskReq.Namespace) } diff --git a/common/rpc/interceptor/nexus/nexus.go b/common/rpc/interceptor/nexus/nexus.go new file mode 100644 index 00000000000..b669ae7dcff --- /dev/null +++ b/common/rpc/interceptor/nexus/nexus.go @@ -0,0 +1,342 @@ +package nexus + +import ( + "context" + "errors" + "fmt" + "net/http" + "slices" + "strings" + "sync" + "time" + + "github.com/nexus-rpc/sdk-go/nexus" + tokenspb "go.temporal.io/server/api/token/v1" + "go.temporal.io/server/common/headers" + "go.temporal.io/server/common/metrics" + "go.temporal.io/server/common/namespace" + "go.temporal.io/server/common/nexus/nexusrpc" +) + +type HandlerFunc func(ctx context.Context, in InterceptorInput) (any, error) + +type Interceptor func(ctx context.Context, in InterceptorInput, next HandlerFunc) (any, error) + +type InterceptorInput interface { + ServiceName() string + OperationName() string + NamespaceName() string + ForwardingInfo() ForwardingInfo + APIName() string // analogous to the gRPC FullMethod + NamespaceEntry() (*namespace.Namespace, error) + EndpointName() string + MetricTags() []metrics.Tag + Header() headers.HeaderGetter + MethodName() string + Request() any + Outcome(out any, err error) string + StartTime() time.Time + sealNexusOp() +} + +var ( + _ InterceptorInput = StartOpInput{} + _ InterceptorInput = CancelOpInput{} + _ InterceptorInput = CompleteOpInput{} +) + +// ForwardingInfo contains the request data needed to forward a Nexus operation. +type ForwardingInfo struct { + OriginalRequestHeaders http.Header + TaskQueue string + EndpointID string + EndpointName string + BusinessID string +} + +type InterceptorError struct { + // wrapped error + Err error + // Outcome tag for metrics reporting + Outcome string + // Flag for propagating the error to the caller as-is without conversion + ExposeDetails bool + // SkipServiceErrorReporting prevents reporting the error as a frontend service failure. + SkipServiceErrorReporting bool +} + +func (t *InterceptorError) Error() string { + return fmt.Sprintf("interceptor error (%s): %v", t.Outcome, t.Err) +} + +func (t *InterceptorError) Unwrap() error { + return t.Err +} + +type outcomeOverrideCtxKey struct{} + +// OutcomeOverride lets an inner interceptor that short-circuits the chain(eg. request forwarder) +// replace the success outcome that would otherwise be derived from the response type +type OutcomeOverride struct { + mu sync.Mutex + value string +} + +func (o *OutcomeOverride) Set(v string) { + if o == nil { + return + } + o.mu.Lock() + defer o.mu.Unlock() + o.value = v +} + +func (o *OutcomeOverride) Value() string { + if o == nil { + return "" + } + o.mu.Lock() + defer o.mu.Unlock() + return o.value +} + +func NewOutcomeOverrideContext(ctx context.Context) (context.Context, *OutcomeOverride) { + override := &OutcomeOverride{} + return context.WithValue(ctx, outcomeOverrideCtxKey{}, override), override +} + +func SetOutcomeOverride(ctx context.Context, v string) { + override, ok := ctx.Value(outcomeOverrideCtxKey{}).(*OutcomeOverride) + if !ok { + return + } + override.Set(v) +} + +// RequestMetadata carries request metadata resolved by the handler (e.g. after a +// namespace registry lookup) that is supplied alongside the rest of the params at +// InterceptorInput construction time. +type RequestMetadata struct { + APIName string + NamespaceEntry *namespace.Namespace + EndpointName string + MetricTags []metrics.Tag // handler-resolved frontend dynamic config for the tags to record + Request any // preserves the request shape passed to custom authorizers +} + +// container for ServiceName(), OperationName(), NamespaceName(), ForwardingInfo(), and +// the fields in RequestMetadata. +type nexusOpBase struct { + serviceName, operation, namespaceName, methodName string + + header headers.HeaderGetter + forwardingInfo ForwardingInfo + requestMetadata RequestMetadata + startTime time.Time +} + +func (b nexusOpBase) StartTime() time.Time { + return b.startTime +} + +func (b nexusOpBase) ServiceName() string { + return b.serviceName +} + +func (b nexusOpBase) OperationName() string { + return b.operation +} + +func (b nexusOpBase) NamespaceName() string { + return b.namespaceName +} + +func (b nexusOpBase) ForwardingInfo() ForwardingInfo { + return b.forwardingInfo +} + +func (b nexusOpBase) APIName() string { + return b.requestMetadata.APIName +} + +func (b nexusOpBase) NamespaceEntry() (*namespace.Namespace, error) { + if b.requestMetadata.NamespaceEntry == nil { + return nil, errors.New("namespace not found in request metadata") + } + return b.requestMetadata.NamespaceEntry, nil +} + +func (b nexusOpBase) EndpointName() string { + return b.requestMetadata.EndpointName +} + +func (b nexusOpBase) MetricTags() []metrics.Tag { + return b.requestMetadata.MetricTags +} + +func (b nexusOpBase) Header() headers.HeaderGetter { + return b.header +} + +func (b nexusOpBase) MethodName() string { + return b.methodName +} + +func (b nexusOpBase) Request() any { + return b.requestMetadata.Request +} + +func (nexusOpBase) sealNexusOp() {} + +type StartOpInput struct { + nexusOpBase + StartOperationOptions nexus.StartOperationOptions + StartOperationInput *nexus.LazyValue +} + +func NewStartOpInput( + serviceName string, + operation string, + namespaceName string, + startTime time.Time, + options nexus.StartOperationOptions, + input *nexus.LazyValue, + forwardingInfo ForwardingInfo, + requestMetadata RequestMetadata, +) StartOpInput { + return StartOpInput{ + nexusOpBase: nexusOpBase{ + serviceName: serviceName, + operation: operation, + namespaceName: namespaceName, + header: options.Header, + methodName: "StartNexusOperation", + forwardingInfo: forwardingInfo, + requestMetadata: requestMetadata, + startTime: startTime, + }, + StartOperationOptions: options, + StartOperationInput: input, + } +} + +type CancelOpInput struct { + nexusOpBase + CancelOperationOptions nexus.CancelOperationOptions + CancellationToken string +} + +func NewCancelOpInput( + serviceName string, + operation string, + namespaceName string, + startTime time.Time, + options nexus.CancelOperationOptions, + cancellationToken string, + forwardingInfo ForwardingInfo, + requestMetadata RequestMetadata, +) CancelOpInput { + return CancelOpInput{ + nexusOpBase: nexusOpBase{ + serviceName: serviceName, + operation: operation, + namespaceName: namespaceName, + header: options.Header, + methodName: "CancelNexusOperation", + forwardingInfo: forwardingInfo, + requestMetadata: requestMetadata, + startTime: startTime, + }, + CancelOperationOptions: options, + CancellationToken: cancellationToken, + } +} + +type CompleteOpInput struct { + nexusOpBase + CompletionRequest *nexusrpc.CompletionRequest + Completion *tokenspb.NexusOperationCompletion +} + +func NewCompleteOpInput( + namespaceName string, + startTime time.Time, + request *nexusrpc.CompletionRequest, + completion *tokenspb.NexusOperationCompletion, + forwardingInfo ForwardingInfo, + requestMetadata RequestMetadata, +) (CompleteOpInput, error) { + if request == nil || request.HTTPRequest == nil { + return CompleteOpInput{}, errors.New("nexus completion request not found") + } + requestMetadata.Request = request + return CompleteOpInput{ + nexusOpBase: nexusOpBase{ + namespaceName: namespaceName, + header: request.HTTPRequest.Header, + methodName: "CompleteNexusOperation", + forwardingInfo: forwardingInfo, + requestMetadata: requestMetadata, + startTime: startTime, + }, + CompletionRequest: request, + Completion: completion, + }, nil +} + +func (c CompleteOpInput) Outcome(out any, err error) string { + if err == nil { + return "success" + } + if ie, ok := errors.AsType[*InterceptorError](err); ok { + if ie.Outcome != "" { + return ie.Outcome + } + err = ie.Err + } + // retaining behavior + if handlerErr, ok := errors.AsType[*nexus.HandlerError](err); ok { + return "error_" + strings.ToLower(string(handlerErr.Type)) + } + return "error_internal" +} + +func (s StartOpInput) Outcome(out any, err error) string { + if outcome, ok := errorOutcome(err); ok { + return outcome + } + switch out.(type) { + case *nexus.HandlerStartOperationResultSync[any]: + return "sync_success" + case *nexus.HandlerStartOperationResultAsync: + return "async_success" + } + return "internal_error" +} + +func (c CancelOpInput) Outcome(out any, err error) string { + if outcome, ok := errorOutcome(err); ok { + return outcome + } + return "success" +} + +func errorOutcome(err error) (string, bool) { + if err != nil { + if ie, ok := errors.AsType[*InterceptorError](err); ok && ie.Outcome != "" { + return ie.Outcome, true + } + return "internal_error", true + } + return "", false +} + +func ChainInterceptors(final HandlerFunc, chain []Interceptor) HandlerFunc { + for _, curr := range slices.Backward(chain) { + next := final + final = func(ctx context.Context, opts InterceptorInput) (any, error) { + return curr(ctx, opts, next) + } + } + return final +} diff --git a/common/rpc/interceptor/nexus/nexus_test.go b/common/rpc/interceptor/nexus/nexus_test.go new file mode 100644 index 00000000000..e4f3b147d7e --- /dev/null +++ b/common/rpc/interceptor/nexus/nexus_test.go @@ -0,0 +1,154 @@ +package nexus + +import ( + "context" + "errors" + "net/http" + "testing" + "time" + + "github.com/nexus-rpc/sdk-go/nexus" + "github.com/stretchr/testify/require" + "go.temporal.io/server/common/nexus/nexusrpc" +) + +func TestOperationInputOutcomes(t *testing.T) { + handlerErr := nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid input") + tests := []struct { + name string + input InterceptorInput + out any + err error + outcome string + }{ + { + name: "start synchronous success", + input: StartOpInput{}, + out: &nexus.HandlerStartOperationResultSync[any]{}, + outcome: "sync_success", + }, + { + name: "start asynchronous success", + input: StartOpInput{}, + out: &nexus.HandlerStartOperationResultAsync{}, + outcome: "async_success", + }, + { + name: "start interceptor error", + input: StartOpInput{}, + err: &InterceptorError{Err: errors.New("failed"), Outcome: "custom_outcome"}, + outcome: "custom_outcome", + }, + { + name: "cancel success", + input: CancelOpInput{}, + outcome: "success", + }, + { + name: "cancel unclassified error", + input: CancelOpInput{}, + err: errors.New("failed"), + outcome: "internal_error", + }, + { + name: "completion success", + input: CompleteOpInput{}, + outcome: "success", + }, + { + name: "completion interceptor error", + input: CompleteOpInput{}, + err: &InterceptorError{Err: errors.New("failed"), Outcome: "custom_outcome"}, + outcome: "custom_outcome", + }, + { + name: "completion handler error", + input: CompleteOpInput{}, + err: handlerErr, + outcome: "error_bad_request", + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + require.Equal(t, tc.outcome, tc.input.Outcome(tc.out, tc.err)) + }) + } + + require.Equal(t, "interceptor error (): ", (&InterceptorError{}).Error()) + _, err := NewCompleteOpInput("namespace", time.Now(), nil, nil, ForwardingInfo{}, RequestMetadata{}) + require.EqualError(t, err, "nexus completion request not found") +} + +func TestInterceptorInputRequest(t *testing.T) { + dispatchRequest := &http.Request{Method: http.MethodPost} + requestStartTime := time.Date(2026, time.May, 5, 17, 0, 0, 123456789, time.UTC) + requestMetadata := RequestMetadata{Request: dispatchRequest} + inputs := []InterceptorInput{ + NewStartOpInput("s", "o", "n", requestStartTime, nexus.StartOperationOptions{}, nil, ForwardingInfo{}, requestMetadata), + NewCancelOpInput("s", "o", "n", requestStartTime, nexus.CancelOperationOptions{}, "t", ForwardingInfo{}, requestMetadata), + } + for _, input := range inputs { + require.Same(t, dispatchRequest, input.Request()) + require.True(t, input.StartTime().Equal(requestStartTime)) + } + + completionRequest := &nexusrpc.CompletionRequest{HTTPRequest: &http.Request{}} + completionInput, err := NewCompleteOpInput("n", requestStartTime, completionRequest, nil, ForwardingInfo{}, RequestMetadata{}) + require.NoError(t, err) + require.Same(t, completionRequest, completionInput.Request()) + require.True(t, completionInput.StartTime().Equal(requestStartTime)) +} + +func TestChainNexusInterceptors(t *testing.T) { + var calls []string + chain := []Interceptor{ + func(ctx context.Context, in InterceptorInput, next HandlerFunc) (any, error) { + calls = append(calls, "first-before") + result, err := next(ctx, in) + calls = append(calls, "first-after") + return result, err + }, + func(ctx context.Context, in InterceptorInput, next HandlerFunc) (any, error) { + calls = append(calls, "second-before") + result, err := next(ctx, in) + calls = append(calls, "second-after") + return result, err + }, + } + + result, err := ChainInterceptors(func(context.Context, InterceptorInput) (any, error) { + calls = append(calls, "handler") + return "result", nil + }, chain)(context.Background(), StartOpInput{}) + + require.NoError(t, err) + require.Equal(t, "result", result) + require.Equal(t, []string{ + "first-before", + "second-before", + "handler", + "second-after", + "first-after", + }, calls) +} + +func TestChainNexusInterceptorsShortCircuit(t *testing.T) { + var calls []string + chain := []Interceptor{ + func(context.Context, InterceptorInput, HandlerFunc) (any, error) { + calls = append(calls, "interceptor") + // dont call next - just return + return "intercepted", nil + }, + } + + result, err := ChainInterceptors(func(context.Context, InterceptorInput) (any, error) { + calls = append(calls, "handler") + return "handler", nil + }, chain)(context.Background(), StartOpInput{}) + + require.NoError(t, err) + require.Equal(t, "intercepted", result) + require.Equal(t, []string{"interceptor"}, calls) +} diff --git a/common/rpc/interceptor/rate_limit.go b/common/rpc/interceptor/rate_limit.go index 2594c03cb6a..dccf8008b32 100644 --- a/common/rpc/interceptor/rate_limit.go +++ b/common/rpc/interceptor/rate_limit.go @@ -9,6 +9,7 @@ import ( "go.temporal.io/api/workflowservice/v1" "go.temporal.io/server/common/headers" "go.temporal.io/server/common/quotas" + "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc" ) @@ -90,3 +91,19 @@ func (i *RateLimitInterceptor) Allow( } return nil } + +// InterceptNexus enforces the global rate limit for a Nexus request. +func (i *RateLimitInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + if err := i.Allow(in.APIName(), in.Header()); err != nil { + return nil, &nexus.InterceptorError{ + Err: err, + Outcome: "global_rate_limited", + ExposeDetails: true, + } + } + return next(ctx, in) +} diff --git a/common/rpc/interceptor/rate_limit_test.go b/common/rpc/interceptor/rate_limit_test.go index 0c645c0be50..12d74beeae8 100644 --- a/common/rpc/interceptor/rate_limit_test.go +++ b/common/rpc/interceptor/rate_limit_test.go @@ -3,10 +3,13 @@ package interceptor import ( "context" "testing" + "time" + "github.com/nexus-rpc/sdk-go/nexus" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" "go.temporal.io/server/common/quotas" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.uber.org/mock/gomock" "google.golang.org/grpc" ) @@ -26,6 +29,42 @@ func TestRateLimitInterceptorSuite(t *testing.T) { suite.Run(t, &rateLimitInterceptorSuite{}) } +func (s *rateLimitInterceptorSuite) TestInterceptNexus() { + for _, tc := range []struct { + name string + input interceptornexus.InterceptorInput + allow bool + nextCalled bool + expectedOutcome string + }{ + {name: "allowed", input: interceptornexus.NewStartOpInput("service", "operation", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, interceptornexus.ForwardingInfo{}, interceptornexus.RequestMetadata{APIName: "NexusOperation"}), allow: true, nextCalled: true}, + {name: "rate limited", input: interceptornexus.NewStartOpInput("service", "operation", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, interceptornexus.ForwardingInfo{}, interceptornexus.RequestMetadata{APIName: "NexusOperation"}), expectedOutcome: "global_rate_limited"}, + } { + s.Run(tc.name, func() { + ctx := context.Background() + interceptor := NewRateLimitInterceptor(s.mockRateLimiter, nil) + s.mockRateLimiter.EXPECT().Allow(gomock.Any(), gomock.Any()).Return(tc.allow) + nextCalled := false + _, err := interceptor.InterceptNexus( + ctx, + tc.input, + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + nextCalled = true + return nil, nil + }, + ) + if tc.expectedOutcome != "" { + var interceptorErr *interceptornexus.InterceptorError + s.ErrorAs(err, &interceptorErr) + s.Equal(tc.expectedOutcome, interceptorErr.Outcome) + } else { + s.NoError(err) + } + s.Equal(tc.nextCalled, nextCalled) + }) + } +} + func (s *rateLimitInterceptorSuite) SetupTest() { s.Assertions = require.New(s.T()) s.controller = gomock.NewController(s.T()) diff --git a/common/rpc/interceptor/redirection.go b/common/rpc/interceptor/redirection.go index 80a1cf01423..1210dec1597 100644 --- a/common/rpc/interceptor/redirection.go +++ b/common/rpc/interceptor/redirection.go @@ -27,7 +27,10 @@ const ( DCRedirectionContextHeaderName = "xdc-redirection" DCRedirectionAPIHeaderName = "xdc-redirection-api" DCRedirectionSourceCellHeaderName = "xdc-redirection-source-cell" - dcRedirectionMetricsPrefix = "DCRedirection" + // DCRedirectionMetricsPrefix prefixes the operation tag on redirection metrics so a + // redirected call is distinguishable from the same operation served locally. Exported + // so the Nexus forwarding interceptor in service/frontend follows the same convention. + DCRedirectionMetricsPrefix = "DCRedirection" ) var ( @@ -287,7 +290,7 @@ func (i *Redirection) handleLocalAPIInvocation( handler grpc.UnaryHandler, methodName string, ) (_ any, retError error) { - scope, startTime := i.BeforeCall(dcRedirectionMetricsPrefix + methodName) + scope, startTime := i.BeforeCall(DCRedirectionMetricsPrefix + methodName) defer func() { i.AfterCall(scope, startTime, i.currentClusterName, "local", retError) }() @@ -307,7 +310,7 @@ func (i *Redirection) handleRedirectAPIInvocation( var targetClusterName = i.currentClusterName var err error - scope, startTime := i.BeforeCall(dcRedirectionMetricsPrefix + methodName) + scope, startTime := i.BeforeCall(DCRedirectionMetricsPrefix + methodName) defer func() { i.AfterCall(scope, startTime, targetClusterName, namespaceName.String(), retError) }() diff --git a/common/rpc/interceptor/retry.go b/common/rpc/interceptor/retry.go index fe2f818b893..c99ff66a570 100644 --- a/common/rpc/interceptor/retry.go +++ b/common/rpc/interceptor/retry.go @@ -4,6 +4,7 @@ import ( "context" "go.temporal.io/server/common/backoff" + "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc" ) @@ -42,3 +43,12 @@ func (i *RetryableInterceptor) Intercept( err := backoff.ThrottleRetryContext(ctx, op, i.policy, i.isRetryable) return response, err } + +// InterceptNexus is a no-op as retries are on the caller side +func (i *RetryableInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + return next(ctx, in) +} diff --git a/common/rpc/interceptor/routing_key_interceptor.go b/common/rpc/interceptor/routing_key_interceptor.go index 82b73ce6e83..909ade6b18e 100644 --- a/common/rpc/interceptor/routing_key_interceptor.go +++ b/common/rpc/interceptor/routing_key_interceptor.go @@ -6,6 +6,7 @@ import ( "go.temporal.io/server/common/log" "go.temporal.io/server/common/log/tag" "go.temporal.io/server/common/namespace" + "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc" ) @@ -121,6 +122,20 @@ func (i *RoutingKeyInterceptor) Intercept( return handler(ctx, req) } +func (i *RoutingKeyInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + if in.ForwardingInfo().BusinessID != "" { + key := namespace.RoutingKey{ + ID: in.ForwardingInfo().BusinessID, + } + ctx = AddRoutingKeyToContext(ctx, key) + } + return next(ctx, in) +} + // AddRoutingKeyToContext adds the routing Key to the context func AddRoutingKeyToContext(ctx context.Context, routingKey namespace.RoutingKey) context.Context { return context.WithValue(ctx, routingKeyCtxKey, routingKey) diff --git a/common/rpc/interceptor/sdk_version.go b/common/rpc/interceptor/sdk_version.go index 6fdb34ffa8f..ea7683be22c 100644 --- a/common/rpc/interceptor/sdk_version.go +++ b/common/rpc/interceptor/sdk_version.go @@ -5,6 +5,7 @@ import ( "sync" "go.temporal.io/server/common/headers" + "go.temporal.io/server/common/rpc/interceptor/nexus" "go.temporal.io/server/common/versioninfo" "google.golang.org/grpc" ) @@ -44,6 +45,26 @@ func (vi *SDKVersionInterceptor) Intercept( return handler(ctx, req) } +// InterceptNexus records and validates the SDK version for a Nexus request. +func (vi *SDKVersionInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + sdkName, sdkVersion := headers.GetClientNameAndVersion(ctx) + if sdkName != "" && sdkVersion != "" { + vi.RecordSDKInfo(sdkName, sdkVersion) + } + if err := vi.versionChecker.ClientSupported(ctx); err != nil { + return nil, &nexus.InterceptorError{ + Err: err, + Outcome: "unsupported_client", + ExposeDetails: true, + } + } + return next(ctx, in) +} + // RecordSDKInfo records name and version tuple in memory func (vi *SDKVersionInterceptor) RecordSDKInfo(name, version string) { info := versioninfo.SDKInfo{Name: name, Version: version} diff --git a/common/rpc/interceptor/sdk_version_test.go b/common/rpc/interceptor/sdk_version_test.go index 25ed973100e..e0b40ace2f5 100644 --- a/common/rpc/interceptor/sdk_version_test.go +++ b/common/rpc/interceptor/sdk_version_test.go @@ -4,9 +4,13 @@ import ( "context" "sort" "testing" + "time" + "github.com/nexus-rpc/sdk-go/nexus" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "go.temporal.io/server/common/headers" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.temporal.io/server/common/versioninfo" ) @@ -64,3 +68,57 @@ func TestSDKVersionRecorder(t *testing.T) { assert.Equal(t, headers.ClientNameTypeScriptSDK, info[1].Name) assert.Equal(t, sdkVersion, info[1].Version) } + +func TestSDKVersionInterceptNexus(t *testing.T) { + clientVersion := "1.10.1" + for _, tc := range []struct { + name string + ctx context.Context + expectedOutcome string + }{ + { + name: "supported client", + ctx: headers.SetVersionsForTests( + context.Background(), + clientVersion, + headers.ClientNameGoSDK, + headers.SupportedServerVersions, + headers.AllFeatures, + ), + }, + { + name: "unsupported client", + ctx: headers.SetVersionsForTests( + context.Background(), + "unparseable.client.version", + headers.ClientNameGoSDK, + headers.SupportedServerVersions, + headers.AllFeatures, + ), + expectedOutcome: "unsupported_client", + }, + } { + t.Run(tc.name, func(t *testing.T) { + interceptor := NewSDKVersionInterceptor() + nextCalled := false + _, err := interceptor.InterceptNexus( + tc.ctx, + interceptornexus.NewStartOpInput("s", "o", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, interceptornexus.ForwardingInfo{}, interceptornexus.RequestMetadata{}), + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + nextCalled = true + return nil, nil + }, + ) + if tc.expectedOutcome != "" { + var interceptorErr *interceptornexus.InterceptorError + require.ErrorAs(t, err, &interceptorErr) + require.Equal(t, tc.expectedOutcome, interceptorErr.Outcome) + require.False(t, nextCalled) + } else { + require.True(t, nextCalled) + require.NoError(t, err) + require.Contains(t, interceptor.GetAndResetSDKInfo(), versioninfo.SDKInfo{Name: headers.ClientNameGoSDK, Version: clientVersion}) + } + }) + } +} diff --git a/common/rpc/interceptor/service_error_interceptor.go b/common/rpc/interceptor/service_error_interceptor.go index 67abfc827fd..b99ab3f9fcc 100644 --- a/common/rpc/interceptor/service_error_interceptor.go +++ b/common/rpc/interceptor/service_error_interceptor.go @@ -4,11 +4,14 @@ import ( "context" "errors" + "github.com/nexus-rpc/sdk-go/nexus" "go.temporal.io/api/serviceerror" "go.temporal.io/server/common/dynamicconfig" "go.temporal.io/server/common/log" + "go.temporal.io/server/common/log/tag" "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/persistence/serialization" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.temporal.io/server/common/util" "google.golang.org/grpc" "google.golang.org/grpc/status" @@ -44,6 +47,26 @@ func (i *ServiceErrorInterceptor) Intercept( ) (any, error) { resp, err := i.capturePanicHandler(ctx, req, handler) + return resp, i.transformError(err) +} + +func (i *ServiceErrorInterceptor) InterceptNexus( + ctx context.Context, + in interceptornexus.InterceptorInput, + next interceptornexus.HandlerFunc, +) (any, error) { + resp, err := i.capturePanicHandlerNexus(ctx, in, next) + if ie, ok := errors.AsType[*interceptornexus.InterceptorError](err); ok { + ie.Err = i.transformNexusError(ie.Err) + return resp, ie + } + return resp, i.transformNexusError(err) +} + +func (i *ServiceErrorInterceptor) transformError(err error) error { + if err == nil { + return nil + } var deserializationError *serialization.DeserializationError var serializationError *serialization.SerializationError // convert serialization errors to be captured as serviceerrors across gRPC calls @@ -59,8 +82,7 @@ func (i *ServiceErrorInterceptor) Intercept( p.Message = util.TruncateUTF8(p.Message, maxLength-len(truncatedSuffix)) + truncatedSuffix st = status.FromProto(p) } - - return resp, st.Err() + return st.Err() } func (i *ServiceErrorInterceptor) capturePanicHandler( @@ -71,3 +93,49 @@ func (i *ServiceErrorInterceptor) capturePanicHandler( defer metrics.CapturePanic(i.logger, i.metricsHandler, &retError) return handler(ctx, req) } + +func (i *ServiceErrorInterceptor) capturePanicHandlerNexus( + ctx context.Context, + in interceptornexus.InterceptorInput, + next interceptornexus.HandlerFunc, +) (_ any, retError error) { + logTags := []tag.Tag{ + tag.Operation(in.MethodName()), + tag.WorkflowNamespace(in.NamespaceName()), + } + if endpointName := in.EndpointName(); endpointName != "" { + logTags = append(logTags, tag.Endpoint(endpointName)) + } + if operationName := in.OperationName(); operationName != "" { + logTags = append(logTags, tag.NexusOperation(operationName)) + } + switch input := in.(type) { + case interceptornexus.StartOpInput: + logTags = append(logTags, tag.NexusStageHandlerInbound, tag.RequestID(input.StartOperationOptions.RequestID)) + case interceptornexus.CancelOpInput: + logTags = append(logTags, tag.NexusStageHandlerInbound) + case interceptornexus.CompleteOpInput: + logTags = append(logTags, tag.NexusStageCallerInbound) + if input.Completion != nil && input.Completion.GetRequestId() != "" { + logTags = append(logTags, tag.RequestID(input.Completion.GetRequestId())) + } + default: + } + defer metrics.CapturePanic(log.With(i.logger, logTags...), i.metricsHandler, &retError) + return next(ctx, in) +} + +// transformNexusError only normalizes gRPC-shaped errors. Nexus-native errors +// are returned as-is to preserve existing mappings. +func (i *ServiceErrorInterceptor) transformNexusError(err error) error { + if err == nil { + return nil + } + if _, ok := errors.AsType[*nexus.HandlerError](err); ok { + return err + } + if _, ok := errors.AsType[*nexus.OperationError](err); ok { + return err + } + return i.transformError(err) +} diff --git a/common/rpc/interceptor/slow_request_logger.go b/common/rpc/interceptor/slow_request_logger.go index dd8dddc7871..e0fd88c0f3d 100644 --- a/common/rpc/interceptor/slow_request_logger.go +++ b/common/rpc/interceptor/slow_request_logger.go @@ -9,6 +9,7 @@ import ( "go.temporal.io/server/common/log" "go.temporal.io/server/common/log/tag" "go.temporal.io/server/common/rpc/interceptor/logtags" + "go.temporal.io/server/common/rpc/interceptor/nexus" "go.temporal.io/server/common/tasktoken" "google.golang.org/grpc" ) @@ -36,31 +37,49 @@ func (i *SlowRequestLoggerInterceptor) Intercept( info *grpc.UnaryServerInfo, handler grpc.UnaryHandler, ) (any, error) { - // Long-polled methods aren't useful logged. - if api.GetMethodMetadata(info.FullMethod).Polling == api.PollingNone { - startTime := time.Now() + tracker := i.trackSlowRequestFn(info.FullMethod, request) + defer tracker() + + return handler(ctx, request) +} - defer func() { - elapsed := time.Since(startTime) - if elapsed > i.slowRequestThreshold() { - i.logSlowRequest(request, info, elapsed) - } - }() +func (i *SlowRequestLoggerInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + tracker := i.trackSlowRequestFn(in.APIName(), in) + defer tracker() + return next(ctx, in) +} + +func (i *SlowRequestLoggerInterceptor) trackSlowRequestFn(operationName string, req any) func() { + // Long-polled methods aren't useful logged. + // If it's a polled method, return a no-op function to defer + if api.GetMethodMetadata(operationName).Polling != api.PollingNone { + return func() {} } - return handler(ctx, request) + startTime := time.Now() + + // Return the cleanup closure for the parent to defer + return func() { + elapsed := time.Since(startTime) + if elapsed > i.slowRequestThreshold() { + i.logSlowRequest(req, operationName, elapsed) + } + } } func (i *SlowRequestLoggerInterceptor) logSlowRequest( request any, - info *grpc.UnaryServerInfo, + method string, elapsed time.Duration, ) { - method := info.FullMethod tags := i.workflowTags.Extract(request, method) tags = append(tags, tag.Duration("duration", elapsed)) tags = append(tags, tag.String("method", method)) - i.logger.Warn("Slow gRPC call", tags...) + i.logger.Warn("Slow request", tags...) } diff --git a/common/rpc/interceptor/slow_request_logger_test.go b/common/rpc/interceptor/slow_request_logger_test.go index cc1372c3b8a..9fc7625f937 100644 --- a/common/rpc/interceptor/slow_request_logger_test.go +++ b/common/rpc/interceptor/slow_request_logger_test.go @@ -5,12 +5,14 @@ import ( "testing" "time" + "github.com/nexus-rpc/sdk-go/nexus" "github.com/stretchr/testify/suite" commonpb "go.temporal.io/api/common/v1" "go.temporal.io/api/workflowservice/v1" "go.temporal.io/server/common/dynamicconfig" "go.temporal.io/server/common/log" "go.temporal.io/server/common/rpc/interceptor" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.uber.org/mock/gomock" "google.golang.org/grpc" ) @@ -71,7 +73,7 @@ func (s *slowRequestLoggerSuite) TestIntercept() { s.NoError(err) // Ensure slow requests are logged. - expectedMsg := "Slow gRPC call" + expectedMsg := "Slow request" s.logger.EXPECT().Warn(gomock.Eq(expectedMsg), gomock.Any()).Times(1) _, err = s.interceptor.Intercept(ctx, request, info, slowHandler) s.NoError(err) @@ -96,3 +98,41 @@ func (s *slowRequestLoggerSuite) TestIntercept() { _, err = s.interceptor.Intercept(ctx, nil, info, slowHandler) s.NoError(err) } + +func (s *slowRequestLoggerSuite) TestInterceptNexus() { + ctx := context.Background() + + const nexusDispatchAPIName = "/temporal.api.nexusservice.v1.NexusService/DispatchByNamespaceAndTaskQueue" + + makeNext := func(delay time.Duration) interceptornexus.HandlerFunc { + return func(context.Context, interceptornexus.InterceptorInput) (any, error) { + //nolint:forbidigo // Allow time.Sleep for timeout tests + time.Sleep(delay) + return nil, nil + } + } + fastNext := makeNext(0) + slowNext := makeNext(testThreshold + 1) + + // The operation name here is deliberately not a known API name: the interceptor + // must key off APIName, not OperationName. + input := interceptornexus.NewStartOpInput( + "test-service", + "user-defined-operation", + "namespace-name", + time.Now(), + nexus.StartOperationOptions{}, + nil, + interceptornexus.ForwardingInfo{}, + interceptornexus.RequestMetadata{APIName: nexusDispatchAPIName}, + ) + + // Ensure fast requests aren't logged. + _, err := s.interceptor.InterceptNexus(ctx, input, fastNext) + s.Require().NoError(err) + + // Ensure slow requests are logged. + s.logger.EXPECT().Warn(gomock.Eq("Slow request"), gomock.Any()).Times(1) + _, err = s.interceptor.InterceptNexus(ctx, input, slowNext) + s.Require().NoError(err) +} diff --git a/common/rpc/interceptor/telemetry.go b/common/rpc/interceptor/telemetry.go index f906da64032..fcf675d4d8a 100644 --- a/common/rpc/interceptor/telemetry.go +++ b/common/rpc/interceptor/telemetry.go @@ -16,6 +16,7 @@ import ( "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/namespace" "go.temporal.io/server/common/rpc/interceptor/logtags" + "go.temporal.io/server/common/rpc/interceptor/nexus" "go.temporal.io/server/common/tasktoken" "go.temporal.io/server/service/frontend/configs" "google.golang.org/grpc" @@ -47,7 +48,7 @@ var ( updateResponseMessageBody anypb.Any _ = updateResponseMessageBody.MarshalFrom(&updatepb.Response{}) - _ grpc.UnaryServerInterceptor = (*TelemetryInterceptor)(nil).UnaryIntercept + _ grpc.UnaryServerInterceptor = (*TelemetryInterceptor)(nil).Intercept _ grpc.StreamServerInterceptor = (*TelemetryInterceptor)(nil).StreamIntercept ) @@ -162,7 +163,7 @@ func telemetryOverrideOperationTag(fullName, operation string) string { return operation } -func (ti *TelemetryInterceptor) UnaryIntercept( +func (ti *TelemetryInterceptor) Intercept( ctx context.Context, req any, info *grpc.UnaryServerInfo, @@ -204,6 +205,86 @@ func AddTelemetryContext(ctx context.Context, metricsHandler metrics.Handler) co return context.WithValue(ctx, metricsCtxKey, metricsHandler) } +// InterceptNexus is a no-op as Nexus request telemetry is recorded by +// [*TelemetryInterceptor.InterceptNexusOutermost] +func (ti *TelemetryInterceptor) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + return next(ctx, in) +} + +func (ti *TelemetryInterceptor) InterceptNexusOutermost( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + serviceHandler := ti.metricsHandler.WithTags( + metrics.OperationTag(in.MethodName()), + metrics.NamespaceTag(in.NamespaceName()), + ) + ctx = AddTelemetryContext(ctx, serviceHandler) + metrics.ServiceRequests.With(serviceHandler).Record(1) + + // Installed before calling next so that an inner interceptor that short-circuits the + // chain (e.g. request forwarding) can still override the derived success outcome. + ctx, outcomeOverride := nexus.NewOutcomeOverrideContext(ctx) + + startTime := in.StartTime() + outcome, failed := "internal_error", true + ctx = metrics.AddMetricsContext(ctx) + defer func() { + ti.RecordLatencyMetrics(ctx, startTime, serviceHandler) + ti.recordNexusRequest(in, startTime, outcome, failed) + }() + + out, err := next(ctx, in) + outcome, failed = in.Outcome(out, err), err != nil + + // override outcome if its set - for request forwarding cases. + // error cases are captured by the wrapped InterceptorError + if err == nil { + if override := outcomeOverride.Value(); override != "" { + outcome = override + } + } + return out, err +} + +func (ti *TelemetryInterceptor) recordNexusRequest( + in nexus.InterceptorInput, + startTime time.Time, + outcome string, + failed bool, +) { + if _, ok := in.(nexus.CompleteOpInput); ok { + handler := ti.metricsHandler.WithTags( + metrics.NamespaceTag(in.NamespaceName()), + metrics.OutcomeTag(outcome), + ) + handler.Counter(metrics.NexusCompletionRequests.Name()).Record(1) + handler.Histogram(metrics.NexusCompletionLatencyHistogram.Name(), metrics.Milliseconds). + Record(time.Since(startTime).Milliseconds()) + return + } + + handler := ti.metricsHandler.WithTags( + metrics.NamespaceTag(in.NamespaceName()), + metrics.NexusEndpointTag(in.EndpointName()), + metrics.NexusMethodTag(in.MethodName()), + ) + handler = handler.WithTags(in.MetricTags()...) + // applied last so that a configured tag doesnt shadow the outcome + handler = handler.WithTags(metrics.OutcomeTag(outcome)) + + metrics.NexusRequests.With(handler).Record(1) + metrics.NexusLatency.With(handler).Record(time.Since(startTime)) + if failed { + metrics.NexusRequestErrors.With(handler).Record(1) + } +} + func (ti *TelemetryInterceptor) RecordLatencyMetrics(ctx context.Context, startTime time.Time, metricsHandler metrics.Handler) { userLatencyDuration := time.Duration(0) if val, ok := metrics.ContextCounterGet(ctx, metrics.HistoryWorkflowExecutionCacheLatency.Name()); ok { diff --git a/common/rpc/interceptor/telemetry_test.go b/common/rpc/interceptor/telemetry_test.go index c57f9efb1dd..5edae64d038 100644 --- a/common/rpc/interceptor/telemetry_test.go +++ b/common/rpc/interceptor/telemetry_test.go @@ -2,9 +2,13 @@ package interceptor import ( "context" + "errors" "testing" + "time" + "github.com/nexus-rpc/sdk-go/nexus" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" commandpb "go.temporal.io/api/command/v1" commonpb "go.temporal.io/api/common/v1" enumspb "go.temporal.io/api/enums/v1" @@ -19,13 +23,134 @@ import ( "go.temporal.io/server/common/log" "go.temporal.io/server/common/log/tag" "go.temporal.io/server/common/metrics" + "go.temporal.io/server/common/metrics/metricstest" "go.temporal.io/server/common/namespace" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" serviceerrors "go.temporal.io/server/common/serviceerror" "go.uber.org/mock/gomock" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" ) +func TestTelemetryInterceptNexusOutermost(t *testing.T) { + extraTag := metrics.StringTag("configured", "tag") + input := interceptornexus.NewStartOpInput( + "s", "o", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, + interceptornexus.ForwardingInfo{}, + interceptornexus.RequestMetadata{MetricTags: []metrics.Tag{extraTag}}, + ) + for _, tc := range []struct { + name string + handlerOut any + handlerErr error + setOverride string + expectedOutcome string + expectedErrors int + }{ + { + name: "sync success is derived from the result type", + handlerOut: &nexus.HandlerStartOperationResultSync[any]{}, + expectedOutcome: "sync_success", + }, + { + name: "async success is derived from the result type", + handlerOut: &nexus.HandlerStartOperationResultAsync{}, + expectedOutcome: "async_success", + }, + { + name: "an interceptor's outcome rides on its error", + handlerErr: &interceptornexus.InterceptorError{Err: errors.New("rejected"), Outcome: "rejected"}, + expectedOutcome: "rejected", + expectedErrors: 1, + }, + { + name: "an unclassified error counts as internal", + handlerErr: errors.New("boom"), + expectedOutcome: "internal_error", + expectedErrors: 1, + }, + { + name: "a short-circuiting interceptor overrides the success outcome", + handlerOut: &nexus.HandlerStartOperationResultSync[any]{}, + setOverride: "request_forwarded", + expectedOutcome: "request_forwarded", + }, + { + name: "an error outcome wins over the override", + handlerErr: &interceptornexus.InterceptorError{Err: errors.New("forward failed"), Outcome: "forwarded_request_error"}, + setOverride: "request_forwarded", + expectedOutcome: "forwarded_request_error", + expectedErrors: 1, + }, + } { + t.Run(tc.name, func(t *testing.T) { + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + telemetry := NewTelemetryInterceptor(nil, metricsHandler, log.NewNoopLogger(), nil, nil) + nextCalled := false + out, err := telemetry.InterceptNexusOutermost( + context.Background(), + input, + func(ctx context.Context, _ interceptornexus.InterceptorInput) (any, error) { + nextCalled = true + // Downstream interceptors read the published handler from the context. + require.NotNil(t, GetMetricsHandlerFromContext(ctx, log.NewNoopLogger())) + if tc.setOverride != "" { + interceptornexus.SetOutcomeOverride(ctx, tc.setOverride) + } + return tc.handlerOut, tc.handlerErr + }, + ) + require.True(t, nextCalled) + require.Equal(t, tc.handlerOut, out) + require.Equal(t, tc.handlerErr, err) + + snapshot := capture.Snapshot() + namespaceTag := metrics.NamespaceTag(testNamespace) + + outcomeTag := metrics.OutcomeTag(tc.expectedOutcome) + methodTag := metrics.NexusMethodTag("StartNexusOperation") + nexusRequests := snapshot[metrics.NexusRequests.Name()] + require.Len(t, nexusRequests, 1) + require.Equal(t, outcomeTag.Value, nexusRequests[0].Tags[outcomeTag.Key]) + require.Equal(t, methodTag.Value, nexusRequests[0].Tags[methodTag.Key]) + require.Equal(t, namespaceTag.Value, nexusRequests[0].Tags[namespaceTag.Key]) + require.Equal(t, extraTag.Value, nexusRequests[0].Tags[extraTag.Key]) + require.Len(t, snapshot[metrics.NexusLatency.Name()], 1) + require.Len(t, snapshot[metrics.NexusRequestErrors.Name()], tc.expectedErrors) + + requests := snapshot[metrics.ServiceRequests.Name()] + require.Len(t, requests, 1) + require.Equal(t, "StartNexusOperation", requests[0].Tags[metrics.OperationTagName]) + require.Equal(t, namespaceTag.Value, requests[0].Tags[namespaceTag.Key]) + require.Len(t, snapshot[metrics.ServiceLatency.Name()], 1) + }) + } +} + +// The shared chain position records nothing; InterceptNexusOutermost is the only recorder. +func TestTelemetryInterceptNexusRecordsNothing(t *testing.T) { + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + telemetry := NewTelemetryInterceptor(nil, metricsHandler, log.NewNoopLogger(), nil, nil) + nextCalled := false + _, err := telemetry.InterceptNexus( + context.Background(), + interceptornexus.NewStartOpInput("s", "o", testNamespace, time.Now(), nexus.StartOperationOptions{}, nil, interceptornexus.ForwardingInfo{}, interceptornexus.RequestMetadata{}), + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + nextCalled = true + return nil, nil + }, + ) + require.NoError(t, err) + require.True(t, nextCalled) + require.Empty(t, capture.Snapshot()) +} + const ( startWorkflow = "StartWorkflowExecution" executeMultiOps = "ExecuteMultiOperation" diff --git a/common/rpc/tlsinfo/context.go b/common/rpc/tlsinfo/context.go new file mode 100644 index 00000000000..fae4c73bc5c --- /dev/null +++ b/common/rpc/tlsinfo/context.go @@ -0,0 +1,35 @@ +package tlsinfo + +import ( + "context" + "crypto/x509" + + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/peer" +) + +// FromContext extracts TLS information from the context's peer value. +func FromContext(ctx context.Context) *credentials.TLSInfo { + p, ok := peer.FromContext(ctx) + if !ok { + return nil + } + if tlsInfo, ok := p.AuthInfo.(credentials.TLSInfo); ok { + return &tlsInfo + } + return nil +} + +// PeerCert extracts an x509 certificate from given tlsInfo. +func PeerCert(tlsInfo *credentials.TLSInfo) *x509.Certificate { + if tlsInfo == nil || len(tlsInfo.State.VerifiedChains) == 0 || len(tlsInfo.State.VerifiedChains[0]) == 0 { + return nil + } + // The assumption here is that we only expect a single verified chain of certs (first[0]). + // It's unclear how we should handle a situation when more than one chain is presented, + // which subject to use. It's okay for us to limit ourselves to one chain. + // We can always extend this logic later. + // We take the first element in the chain ([0]) because that's the client cert + // (at the beginning of the chain), not intermediary CAs or the root CA (at the end of the chain). + return tlsInfo.State.VerifiedChains[0][0] +} diff --git a/service/frontend/frontend_interceptors.go b/service/frontend/frontend_interceptors.go new file mode 100644 index 00000000000..86d21b10867 --- /dev/null +++ b/service/frontend/frontend_interceptors.go @@ -0,0 +1,170 @@ +package frontend + +import ( + "context" + + "go.temporal.io/server/chasm" + "go.temporal.io/server/common/authorization" + "go.temporal.io/server/common/metrics" + "go.temporal.io/server/common/rpc/grpcfaults" + "go.temporal.io/server/common/rpc/interceptor" + "go.temporal.io/server/common/rpc/interceptor/nexus" + "google.golang.org/grpc" +) + +// Interceptor is a unified interface for gRPC and Nexus interceptors +type Interceptor interface { + // gRPC Interceptor + Intercept( + ctx context.Context, + req any, + info *grpc.UnaryServerInfo, + handler grpc.UnaryHandler, + ) (any, error) + // Nexus Interceptor + InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, + ) (any, error) +} + +type interceptorsProvider struct { + interceptors []Interceptor + nexusTelemetry nexus.Interceptor // required to be first in the Nexus chain +} + +func newInterceptorsProvider( + maskInternalErrorDetailsInterceptor *interceptor.MaskInternalErrorDetailsInterceptor, + serviceErrorInterceptor *interceptor.ServiceErrorInterceptor, + frontendServiceErrorInterceptor *interceptor.FrontendServiceErrorInterceptor, + businessIDInterceptor *interceptor.RoutingKeyInterceptor, + namespaceValidatorInterceptor *interceptor.NamespaceValidatorInterceptor, + namespaceLogInterceptor *interceptor.NamespaceLogInterceptor, + authInterceptor *authorization.Interceptor, + namespaceHandoverInterceptor *interceptor.NamespaceHandoverInterceptor, + redirectionInterceptor *interceptor.Redirection, + nexusForwarder *nexusForwardingInterceptor, + telemetryInterceptor *interceptor.TelemetryInterceptor, + healthInterceptor *interceptor.HealthInterceptor, + namespaceLengthValidatorInterceptor *interceptor.NamespaceLengthValidatorInterceptor, + namespaceCountLimiterInterceptor *interceptor.ConcurrentRequestLimitInterceptor, + namespaceRateLimiterInterceptorWrapper *interceptor.NamespaceRateLimitInterceptorWrapper, + rateLimitInterceptor *interceptor.RateLimitInterceptor, + sdkVersionInterceptor *interceptor.SDKVersionInterceptor, + callerInfoInterceptor *interceptor.CallerInfoInterceptor, + slowRequestLoggerInterceptor *interceptor.SlowRequestLoggerInterceptor, + chasmRequestVisibilityInterceptor *chasm.ChasmVisibilityInterceptor, + contextMetadataInterceptor *interceptor.ContextMetadataInterceptor, + customGRPCInterceptors []grpc.UnaryServerInterceptor, + customInterceptors []Interceptor, + retryableInterceptor *interceptor.RetryableInterceptor, + faultsInterceptor *grpcfaults.FaultsInterceptor, +) *interceptorsProvider { + + metricsCtxInjectorInterceptor := &interceptorWrapper{ + grpcInterceptor: metrics.NewServerMetricsContextInjectorInterceptor(), + nexusInterceptor: nexusNoOpInterceptor, // added by telemetryInterceptor.InterceptNexusOutermost + } + + // redirectionWrapper is one chain position for both transports: gRPC DC redirection + // and Nexus HTTP forwarding. The implementations stay separate but are wrapped together + // for canonical ordering of interceptors for both gRPC and Nexus + redirectionWrapper := &interceptorWrapper{ + grpcInterceptor: redirectionInterceptor.Intercept, + nexusInterceptor: nexusForwarder.InterceptNexus, + } + + // Order is important. Error interceptors must stay outermost, routing must precede namespace + // access, and telemetry must follow redirection to attribute requests to the serving cluster. + // Nexus interceptors outward of error producers must preserve InterceptorError. + interceptors := []Interceptor{ + maskInternalErrorDetailsInterceptor, + serviceErrorInterceptor, + frontendServiceErrorInterceptor, + businessIDInterceptor, + namespaceLengthValidatorInterceptor, + namespaceLogInterceptor, + metricsCtxInjectorInterceptor, + authInterceptor, + namespaceHandoverInterceptor, + redirectionWrapper, + telemetryInterceptor, + healthInterceptor, + namespaceValidatorInterceptor, + namespaceCountLimiterInterceptor, + namespaceRateLimiterInterceptorWrapper, + rateLimitInterceptor, + sdkVersionInterceptor, + callerInfoInterceptor, + slowRequestLoggerInterceptor, + chasmRequestVisibilityInterceptor, + contextMetadataInterceptor, + } + for _, grpcInterceptor := range customGRPCInterceptors { + interceptors = append(interceptors, &interceptorWrapper{ + grpcInterceptor: grpcInterceptor, + nexusInterceptor: nexusNoOpInterceptor, + }) + } + interceptors = append(interceptors, customInterceptors...) + + interceptors = append(interceptors, faultsInterceptor) + interceptors = append(interceptors, retryableInterceptor) + + return &interceptorsProvider{ + interceptors: interceptors, + nexusTelemetry: telemetryInterceptor.InterceptNexusOutermost, + } +} + +func (n *interceptorsProvider) grpcInterceptors() []grpc.UnaryServerInterceptor { + grpcInterceptors := make([]grpc.UnaryServerInterceptor, 0, len(n.interceptors)) + for _, i := range n.interceptors { + grpcInterceptors = append(grpcInterceptors, i.Intercept) + } + return grpcInterceptors +} + +func (n *interceptorsProvider) nexusInterceptors() []nexus.Interceptor { + nexusInterceptors := make([]nexus.Interceptor, 0, len(n.interceptors)+1) + // telemetry is the outermost in chain for Nexus requests to allow recording + // all metrics and retain behavior. In the future, gRPC will also move telemetry + // to outermost after an impact evaluation- this will allow gRPC to also capture + // all metrics from authz/redirection related failures as well. + nexusInterceptors = append(nexusInterceptors, n.nexusTelemetry) + for _, i := range n.interceptors { + nexusInterceptors = append(nexusInterceptors, i.InterceptNexus) + } + return nexusInterceptors +} + +type interceptorWrapper struct { + grpcInterceptor grpc.UnaryServerInterceptor + nexusInterceptor nexus.Interceptor +} + +func (i interceptorWrapper) Intercept( + ctx context.Context, + req any, + info *grpc.UnaryServerInfo, + handler grpc.UnaryHandler, +) (any, error) { + return i.grpcInterceptor(ctx, req, info, handler) +} + +func (i interceptorWrapper) InterceptNexus( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + return i.nexusInterceptor(ctx, in, next) +} + +func nexusNoOpInterceptor( + ctx context.Context, + in nexus.InterceptorInput, + next nexus.HandlerFunc, +) (any, error) { + return next(ctx, in) +} diff --git a/service/frontend/fx.go b/service/frontend/fx.go index 2ba531cf795..10b3a7afa5d 100644 --- a/service/frontend/fx.go +++ b/service/frontend/fx.go @@ -105,6 +105,7 @@ var Module = fx.Options( fx.Provide(interceptor.NewHealthInterceptor), fx.Provide(NamespaceCountLimitInterceptorProvider), fx.Provide(NamespaceValidatorInterceptorProvider), + fx.Provide(NamespaceLengthValidatorInterceptorProvider), fx.Provide(NamespaceRateLimitersProvider), fx.Provide(NamespaceRateLimitInterceptorProvider), fx.Provide(SDKVersionInterceptorProvider), @@ -126,10 +127,15 @@ var Module = fx.Options( fx.Provide(callbackValidatorProvider), fx.Provide(HandlerProvider), fx.Provide(AdminHandlerProvider), + fx.Provide(FrontendServiceErrorInterceptorProvider), fx.Provide(NamespaceDLQHandlerProvider), fx.Provide(OperatorHandlerProvider), fx.Provide(NewVersionChecker), fx.Provide(ServiceResolverProvider), + fx.Provide(newNexusForwardingInterceptor), + fx.Provide(interceptor.NewNamespaceRateLimitInterceptorWrapper), + fx.Provide(NewFaultsInterceptorProvider), + fx.Provide(newInterceptorsProvider), fx.Provide(newNexusCompletionHandler), fx.Provide(NewNexusOperationHTTPHandler), fx.Provide(newNexusCompletionHTTPHandler), @@ -233,35 +239,15 @@ func (n *namespaceChecker) Exists(name namespace.Name) error { func GrpcServerOptionsProvider( logger log.Logger, - cfg *config.Config, serviceConfig *Config, serviceName primitives.ServiceName, rpcFactory common.RPCFactory, - serviceErrorInterceptor *interceptor.ServiceErrorInterceptor, - namespaceLogInterceptor *interceptor.NamespaceLogInterceptor, - namespaceRateLimiterInterceptor interceptor.NamespaceRateLimitInterceptor, - namespaceCountLimiterInterceptor *interceptor.ConcurrentRequestLimitInterceptor, - namespaceValidatorInterceptor *interceptor.NamespaceValidatorInterceptor, - namespaceHandoverInterceptor *interceptor.NamespaceHandoverInterceptor, - businessIDInterceptor *interceptor.RoutingKeyInterceptor, - redirectionInterceptor *interceptor.Redirection, + interceptorsProvider *interceptorsProvider, telemetryInterceptor *interceptor.TelemetryInterceptor, - retryableInterceptor *interceptor.RetryableInterceptor, - healthInterceptor *interceptor.HealthInterceptor, - rateLimitInterceptor *interceptor.RateLimitInterceptor, traceStatsHandler telemetry.ServerStatsHandler, metricsStatsHandler metrics.ServerStatsHandler, - sdkVersionInterceptor *interceptor.SDKVersionInterceptor, - callerInfoInterceptor *interceptor.CallerInfoInterceptor, authInterceptor *authorization.Interceptor, - maskInternalErrorDetailsInterceptor *interceptor.MaskInternalErrorDetailsInterceptor, - contextMetadataInterceptor *interceptor.ContextMetadataInterceptor, - slowRequestLoggerInterceptor *interceptor.SlowRequestLoggerInterceptor, - chasmRequestVisibilityInterceptor *chasm.ChasmVisibilityInterceptor, - customInterceptors []grpc.UnaryServerInterceptor, customStreamInterceptors []grpc.StreamServerInterceptor, - metricsHandler metrics.Handler, - testHooks testhooks.TestHooks, ) GrpcServerOptions { kep := keepalive.EnforcementPolicy{ MinTime: serviceConfig.KeepAliveMinTime(), @@ -287,46 +273,8 @@ func GrpcServerOptionsProvider( if err != nil { logger.Fatal("creating gRPC server options failed", tag.Error(err)) } - unaryInterceptors := []grpc.UnaryServerInterceptor{ - // Order of interceptors is important - // Mask error interceptor should be the most outer interceptor since it handle the errors format - // Service Error Interceptor should be the next most outer interceptor on error handling - maskInternalErrorDetailsInterceptor.Intercept, - serviceErrorInterceptor.Intercept, - interceptor.NewFrontendServiceErrorInterceptor(logger), - // BusinessID interceptor extracts business ID and adds it to context for use, must be before any interceptor that touches namespaces (namespaceValidator, handoverInterceptor) - businessIDInterceptor.Intercept, - namespaceValidatorInterceptor.NamespaceValidateIntercept, - namespaceLogInterceptor.Intercept, // TODO: Deprecate this with a outer custom interceptor - metrics.NewServerMetricsContextInjectorInterceptor(), - authInterceptor.Intercept, - // Handover interceptor has to above redirection because the request will route to the correct cluster after handover completed. - // And retry cannot be performed before customInterceptors. - namespaceHandoverInterceptor.Intercept, - redirectionInterceptor.Intercept, - // Telemetry interceptor must be after redirection to ensure metrics are recorded in the correct cluster - telemetryInterceptor.UnaryIntercept, - healthInterceptor.Intercept, - namespaceValidatorInterceptor.StateValidationIntercept, - namespaceCountLimiterInterceptor.Intercept, - namespaceRateLimiterInterceptor.Intercept, - rateLimitInterceptor.Intercept, - sdkVersionInterceptor.Intercept, - callerInfoInterceptor.Intercept, - slowRequestLoggerInterceptor.Intercept, - chasmRequestVisibilityInterceptor.Intercept, - contextMetadataInterceptor.Intercept, - } - if len(customInterceptors) > 0 { - // TODO: Deprecate WithChainedFrontendGrpcInterceptors and provide a inner custom interceptor - unaryInterceptors = append(unaryInterceptors, customInterceptors...) - } - faultGenerator := grpcfaultstest.NewGenerator(testHooks) - if faultInterceptor := grpcfaults.UnaryServerInterceptor(faultGenerator); faultInterceptor != nil { - unaryInterceptors = append(unaryInterceptors, faultInterceptor) - } - // retry interceptor should be the most inner interceptor - unaryInterceptors = append(unaryInterceptors, retryableInterceptor.Intercept) + + unaryInterceptors := interceptorsProvider.grpcInterceptors() streamInterceptor := []grpc.StreamServerInterceptor{ authInterceptor.InterceptStream, @@ -692,6 +640,15 @@ func NamespaceValidatorInterceptorProvider( ) } +func NamespaceLengthValidatorInterceptorProvider( + params NamespaceValidatorInterceptorParams, +) *interceptor.NamespaceLengthValidatorInterceptor { + return interceptor.NewNamespaceLengthValidatorInterceptor( + params.NamespaceRegistry, + params.ServiceConfig.MaxIDLengthLimit, + ) +} + func SDKVersionInterceptorProvider() *interceptor.SDKVersionInterceptor { return interceptor.NewSDKVersionInterceptor() } @@ -712,6 +669,12 @@ func SlowRequestLoggerInterceptorProvider( ) } +func FrontendServiceErrorInterceptorProvider( + logger log.Logger, +) *interceptor.FrontendServiceErrorInterceptor { + return interceptor.NewFrontendServiceErrorInterceptorWrapper(logger) +} + func PersistenceRateLimitingParamsProvider( serviceConfig *Config, persistenceLazyLoadedServiceResolver service.PersistenceLazyLoadedServiceResolver, @@ -780,6 +743,12 @@ func FEReplicatorNamespaceReplicationQueueProvider( return replicatorNamespaceReplicationQueue } +func NewFaultsInterceptorProvider(hooks testhooks.TestHooks) *grpcfaults.FaultsInterceptor { + return grpcfaults.NewFaultsInterceptor( + grpcfaultstest.NewGenerator(hooks), + ) +} + func ServiceResolverProvider( membershipMonitor membership.Monitor, serviceName primitives.ServiceName, diff --git a/service/frontend/fx_test.go b/service/frontend/fx_test.go index 39a7ea121d3..e04c3196089 100644 --- a/service/frontend/fx_test.go +++ b/service/frontend/fx_test.go @@ -235,7 +235,7 @@ func TestRateLimitInterceptorProvider(t *testing.T) { svc := &testSvc{} server := grpc.NewServer(grpc.ChainUnaryInterceptor( serviceErrorInterceptor.Intercept, - interceptor.NewFrontendServiceErrorInterceptor(log.NewTestLogger()), + interceptor.NewFrontendServiceErrorInterceptorWrapper(log.NewTestLogger()).Intercept, rateLimitInterceptor.Intercept, )) workflowservice.RegisterWorkflowServiceServer(server, svc) @@ -603,7 +603,7 @@ func TestNamespaceRateLimitInterceptorProvider(t *testing.T) { svc := &testSvc{} server := grpc.NewServer(grpc.ChainUnaryInterceptor( serviceErrorInterceptor.Intercept, - interceptor.NewFrontendServiceErrorInterceptor(log.NewTestLogger()), + interceptor.NewFrontendServiceErrorInterceptorWrapper(log.NewTestLogger()).Intercept, rateLimitInterceptor.Intercept, )) workflowservice.RegisterWorkflowServiceServer(server, svc) @@ -798,7 +798,7 @@ func TestNamespaceRateLimitMetrics(t *testing.T) { svc := &testSvc{} server := grpc.NewServer(grpc.ChainUnaryInterceptor( serviceErrorInterceptor.Intercept, - interceptor.NewFrontendServiceErrorInterceptor(log.NewTestLogger()), + interceptor.NewFrontendServiceErrorInterceptorWrapper(log.NewTestLogger()).Intercept, rateLimitInterceptor.Intercept, )) workflowservice.RegisterWorkflowServiceServer(server, svc) diff --git a/service/frontend/nexus_completion_http_handler.go b/service/frontend/nexus_completion_http_handler.go index 66b323612ba..27811f9e2e3 100644 --- a/service/frontend/nexus_completion_http_handler.go +++ b/service/frontend/nexus_completion_http_handler.go @@ -5,10 +5,8 @@ import ( "errors" "fmt" "net/http" - "net/http/httptrace" "net/url" "runtime/debug" - "strconv" "strings" "time" @@ -18,9 +16,7 @@ import ( "go.temporal.io/api/serviceerror" "go.temporal.io/server/api/historyservice/v1" tokenspb "go.temporal.io/server/api/token/v1" - "go.temporal.io/server/common" "go.temporal.io/server/common/authorization" - "go.temporal.io/server/common/cluster" "go.temporal.io/server/common/headers" "go.temporal.io/server/common/log" "go.temporal.io/server/common/log/tag" @@ -31,6 +27,7 @@ import ( "go.temporal.io/server/common/resource" "go.temporal.io/server/common/rpc" "go.temporal.io/server/common/rpc/interceptor" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "go.temporal.io/server/nexusworkflowref" "go.temporal.io/server/service/frontend/configs" "go.temporal.io/server/service/history/consts" @@ -43,25 +40,17 @@ const nexusCompletionAPIName = configs.CompleteNexusOperation const nexusCompletionMethodName = "CompleteNexusOperation" type nexusCompletionHandler struct { - ClusterMetadata cluster.Metadata - NamespaceRegistry namespace.Registry - Logger log.Logger - MetricsHandler metrics.Handler - Config *Config - CallbackTokenGenerator *commonnexus.CallbackTokenGenerator - HistoryClient resource.HistoryClient - TelemetryInterceptor *interceptor.TelemetryInterceptor - RequestErrorHandler *interceptor.RequestErrorHandler - NamespaceValidationInterceptor *interceptor.NamespaceValidatorInterceptor - NamespaceRateLimitInterceptor interceptor.NamespaceRateLimitInterceptor - NamespaceConcurrencyLimitInterceptor *interceptor.ConcurrentRequestLimitInterceptor - RateLimitInterceptor *interceptor.RateLimitInterceptor - AuthInterceptor *authorization.Interceptor - RedirectionInterceptor *interceptor.Redirection - ForwardingClients *cluster.FrontendHTTPClientCache - HTTPTraceProvider commonnexus.HTTPClientTraceProvider - clientVersionChecker headers.VersionChecker - preProcessErrorsCounter metrics.CounterIface + NamespaceRegistry namespace.Registry + Logger log.Logger + MetricsHandler metrics.Handler + Config *Config + CallbackTokenGenerator *commonnexus.CallbackTokenGenerator + HistoryClient resource.HistoryClient + RequestErrorHandler *interceptor.RequestErrorHandler + telemetryInterceptor *interceptor.TelemetryInterceptor + AuthInterceptor *authorization.Interceptor // required for parsing auth info, not used as an interceptor + preProcessErrorsCounter metrics.CounterIface + chainedHandler interceptornexus.HandlerFunc } type nexusCompletionHTTPHandler struct { @@ -69,45 +58,32 @@ type nexusCompletionHTTPHandler struct { } func newNexusCompletionHandler( - clusterMetadata cluster.Metadata, namespaceRegistry namespace.Registry, logger log.Logger, metricsHandler metrics.Handler, serviceConfig *Config, callbackTokenGenerator *commonnexus.CallbackTokenGenerator, historyClient resource.HistoryClient, - telemetryInterceptor *interceptor.TelemetryInterceptor, requestErrorHandler *interceptor.RequestErrorHandler, - namespaceValidationInterceptor *interceptor.NamespaceValidatorInterceptor, - namespaceRateLimitInterceptor interceptor.NamespaceRateLimitInterceptor, - namespaceConcurrencyLimitInterceptor *interceptor.ConcurrentRequestLimitInterceptor, - rateLimitInterceptor *interceptor.RateLimitInterceptor, authInterceptor *authorization.Interceptor, - redirectionInterceptor *interceptor.Redirection, - forwardingClients *cluster.FrontendHTTPClientCache, - httpTraceProvider commonnexus.HTTPClientTraceProvider, + telemetryInterceptor *interceptor.TelemetryInterceptor, + interceptorsProvider *interceptorsProvider, ) *nexusCompletionHandler { - return &nexusCompletionHandler{ - ClusterMetadata: clusterMetadata, - NamespaceRegistry: namespaceRegistry, - Logger: log.With(logger, tag.NexusStageCallerInbound), - MetricsHandler: metricsHandler, - Config: serviceConfig, - CallbackTokenGenerator: callbackTokenGenerator, - HistoryClient: historyClient, - TelemetryInterceptor: telemetryInterceptor, - RequestErrorHandler: requestErrorHandler, - NamespaceValidationInterceptor: namespaceValidationInterceptor, - NamespaceRateLimitInterceptor: namespaceRateLimitInterceptor, - NamespaceConcurrencyLimitInterceptor: namespaceConcurrencyLimitInterceptor, - RateLimitInterceptor: rateLimitInterceptor, - AuthInterceptor: authInterceptor, - RedirectionInterceptor: redirectionInterceptor, - ForwardingClients: forwardingClients, - HTTPTraceProvider: httpTraceProvider, - clientVersionChecker: headers.NewDefaultVersionChecker(), - preProcessErrorsCounter: metricsHandler.Counter(metrics.NexusCompletionRequestPreProcessErrors.Name()), - } + + h := &nexusCompletionHandler{ + NamespaceRegistry: namespaceRegistry, + Logger: log.With(logger, tag.NexusStageCallerInbound), + MetricsHandler: metricsHandler, + Config: serviceConfig, + CallbackTokenGenerator: callbackTokenGenerator, + HistoryClient: historyClient, + RequestErrorHandler: requestErrorHandler, + AuthInterceptor: authInterceptor, + telemetryInterceptor: telemetryInterceptor, + preProcessErrorsCounter: metricsHandler.Counter(metrics.NexusCompletionRequestPreProcessErrors.Name()), + } + h.chainedHandler = interceptornexus.ChainInterceptors(h.finalCompleteHandler, interceptorsProvider.nexusInterceptors()) + return h } func newNexusCompletionHTTPHandler(handler *nexusCompletionHandler) *nexusCompletionHTTPHandler { @@ -123,7 +99,7 @@ func newNexusCompletionHTTPHandler(handler *nexusCompletionHandler) *nexusComple // CompleteOperation implements nexus.CompletionHandler. // nolint:revive // (cyclomatic complexity) This function is long but the complexity is justified. func (h *nexusCompletionHandler) CompleteOperation(ctx context.Context, r *nexusrpc.CompletionRequest) (retErr error) { - startTime := time.Now() + requestStartTime := time.Now() token, err := commonnexus.DecodeCallbackToken(r.HTTPRequest.Header.Get(commonnexus.CallbackTokenHeader)) if err != nil { h.Logger.Error("failed to decode callback token", tag.Error(err)) @@ -166,20 +142,37 @@ func (h *nexusCompletionHandler) CompleteOperation(ctx context.Context, r *nexus rCtx := &requestContext{ nexusCompletionHandler: h, namespace: ns, - businessID: targetBusinessID, logger: logger, - metricsHandler: h.MetricsHandler.WithTags(metrics.NamespaceTag(ns.Name().String())), metricsHandlerForInterceptors: h.MetricsHandler.WithTags( metrics.OperationTag(nexusCompletionMethodName), metrics.NamespaceTag(ns.Name().String()), ), - requestStartTime: startTime, } if r.HTTPRequest.Header != nil { rCtx.originalHeaders = r.HTTPRequest.Header.Clone() } ctx = rCtx.augmentContext(ctx, r.HTTPRequest.Header) - defer rCtx.capturePanicAndRecordMetrics(&ctx, &retErr) + defer finalizeCompletionRequest(rCtx, &retErr) + + const outcomeBadRequest = "error_bad_request" + + // recordPreInterceptorFailure is for pre-interceptor chain error recording. + recordPreInterceptorFailure := func(outcome string) { + completionMetrics := h.MetricsHandler.WithTags( + metrics.NamespaceTag(ns.Name().String()), + metrics.OutcomeTag(outcome), + ) + completionMetrics.Counter(metrics.NexusCompletionRequests.Name()).Record(1) + completionMetrics.Histogram( + metrics.NexusCompletionLatencyHistogram.Name(), + metrics.Milliseconds, + ).Record(time.Since(requestStartTime).Milliseconds()) + + metrics.ServiceRequests.With(rCtx.metricsHandlerForInterceptors).Record(1) + h.telemetryInterceptor.RecordLatencyMetrics( + ctx, requestStartTime, rCtx.metricsHandlerForInterceptors, + ) + } if r.HTTPRequest.URL.Path != commonnexus.PathCompletionCallbackNoIdentifier { nsNameEscaped := commonnexus.RouteCompletionCallback.Deserialize(mux.Vars(r.HTTPRequest)) @@ -187,26 +180,79 @@ func (h *nexusCompletionHandler) CompleteOperation(ctx context.Context, r *nexus if err != nil { logger.Error("failed to extract namespace from request", tag.Error(err)) h.preProcessErrorsCounter.Record(1) - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid URL") + recordPreInterceptorFailure(outcomeBadRequest) + return &interceptornexus.InterceptorError{ + Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid URL"), + SkipServiceErrorReporting: true, + } } if nsName != ns.Name().String() { logger.Error( "namespace in callback URL doesn't match the completion token", tag.String("url-namespace", nsName), ) - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid callback token") + recordPreInterceptorFailure(outcomeBadRequest) + return &interceptornexus.InterceptorError{ + Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid callback token"), + SkipServiceErrorReporting: true, + } + } + } + ctx, err = rCtx.parseTLSAndAuthInfo(ctx, r) + if err != nil { + recordPreInterceptorFailure("error_internal") + return &interceptornexus.InterceptorError{ + Err: err, + SkipServiceErrorReporting: true, } } - if err := rCtx.interceptRequest(ctx, r); err != nil { - if _, ok := errors.AsType[*serviceerror.NamespaceNotActive](err); ok { - return h.forwardCompleteOperation(ctx, r, rCtx) + interceptorInput, err := interceptornexus.NewCompleteOpInput( + ns.Name().String(), + requestStartTime, + r, + completion, + interceptornexus.ForwardingInfo{ + OriginalRequestHeaders: rCtx.originalHeaders, + BusinessID: targetBusinessID, + }, + interceptornexus.RequestMetadata{ + APIName: nexusCompletionAPIName, + NamespaceEntry: ns, + }, + ) + if err != nil { + logger.Error("invalid nexus completion request", tag.Error(err)) + recordPreInterceptorFailure(outcomeBadRequest) + return &interceptornexus.InterceptorError{ + Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid request"), + SkipServiceErrorReporting: true, } - return err } + ctx = withRequestContext(ctx, rCtx) + _, err = h.chainedHandler(ctx, interceptorInput) + return err +} + +func (h *nexusCompletionHandler) finalCompleteHandler( + ctx context.Context, + in interceptornexus.InterceptorInput, +) (any, error) { + rCtx, ok := requestContextFromContext(ctx) + if !ok { + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "invalid request context for nexus completion") + } + coi, ok := in.(interceptornexus.CompleteOpInput) + if !ok { + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "invalid request for nexus complete operation") + } + logger := rCtx.logger + completion := coi.Completion + r := coi.CompletionRequest + ns := rCtx.namespace tokenLimit := h.Config.MaxNexusOperationTokenLength(ns.Name().String()) if len(r.OperationToken) > tokenLimit { - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "operation token length exceeds allowed limit (%d/%d)", len(r.OperationToken), tokenLimit) + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "operation token length exceeds allowed limit (%d/%d)", len(r.OperationToken), tokenLimit) } links := commonnexus.ConvertNexusLinksToProtoLinks(r.Links, logger) @@ -219,31 +265,37 @@ func (h *nexusCompletionHandler) CompleteOperation(ctx context.Context, r *nexus var result *commonpb.Payload if err := r.Result.Consume(&result); err != nil { logger.Error("cannot deserialize payload from completion result", tag.Error(err)) - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid result content") + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid result content") } if result.Size() > h.Config.BlobSizeLimitError(ns.Name().String()) { - logger.Error("payload size exceeds error limit for Nexus CompleteOperation request") - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "result exceeds size limit") + logger.Error("payload size exceeds error limit for Nexus CompleteOperation request", tag.WorkflowNamespace(ns.Name().String())) + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "result exceeds size limit") } successPayload = result default: // The Nexus SDK ensures this never happens but just in case... logger.Error("invalid operation state in completion request", tag.String("state", string(r.State))) - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid completion state") + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid completion state") } - err = h.completeOperation(ctx, logger, completion, successPayload, r, links, h.Config.EnableChasm(ns.Name().String())) + err := h.completeOperation(ctx, logger, completion, successPayload, r, links, h.Config.EnableChasm(ns.Name().String())) if err == nil { - return nil + return nil, nil } logger.Error("failed to process nexus completion request", tag.Error(err)) if _, ok := errors.AsType[*serviceerror.NamespaceNotActive](err); ok { - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnavailable, "cluster inactive") + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnavailable, "cluster inactive") } if _, ok := errors.AsType[*serviceerror.NotFound](err); ok { - return commonnexus.ConvertGRPCError(err, true) + return nil, &interceptornexus.InterceptorError{Err: err, Outcome: "error_not_found", ExposeDetails: true} } - return commonnexus.ConvertGRPCError(err, false) + // Preserve specific outcome tags on handler errors. + converted := commonnexus.ConvertGRPCError(err, false) + outcome := "error_internal" + if handlerErr, ok := errors.AsType[*nexus.HandlerError](converted); ok { + outcome = "error_" + strings.ToLower(string(handlerErr.Type)) + } + return nil, &interceptornexus.InterceptorError{Err: err, Outcome: outcome} } // completeOperation dispatches the completion to the framework named by its @@ -398,68 +450,6 @@ func (h *nexusCompletionHandler) completeChasmOperation( return err } -func (h *nexusCompletionHandler) forwardCompleteOperation(ctx context.Context, r *nexusrpc.CompletionRequest, rCtx *requestContext) error { - targetCluster := rCtx.namespace.ActiveClusterName(namespace.RoutingKey{ID: rCtx.businessID}) - logger := log.With( - rCtx.logger, - tag.SourceCluster(h.ClusterMetadata.GetCurrentClusterName()), - tag.TargetCluster(targetCluster), - ) - - client, err := h.ForwardingClients.Get(targetCluster) - if err != nil { - logger.Error("unable to get HTTP client for forward request", tag.Error(err)) - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "internal error") - } - - forwardURL, err := url.JoinPath(client.BaseURL(), commonnexus.RouteCompletionCallback.Path(rCtx.namespace.Name().String())) - if err != nil { - logger.Error("failed to construct forwarding request URL", tag.Error(err)) - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "internal error") - } - - if h.HTTPTraceProvider != nil { - traceLogger := log.With(logger, tag.AttemptStart(time.Now().UTC())) - if trace := h.HTTPTraceProvider.NewForwardingTrace(traceLogger); trace != nil { - ctx = httptrace.WithClientTrace(ctx, trace) - } - } - - var completion nexusrpc.CompleteOperationOptions - switch r.State { - case nexus.OperationStateSucceeded: - completion = nexusrpc.CompleteOperationOptions{ - Result: r.Result.Reader, - OperationToken: r.OperationToken, - StartTime: r.StartTime, - CloseTime: r.CloseTime, - Links: r.Links, - } - case nexus.OperationStateFailed, nexus.OperationStateCanceled: - // For unsuccessful operations, the Nexus framework reads and closes the original request body to deserialize - // the failure, so we must construct a new completion to forward. - completion = nexusrpc.CompleteOperationOptions{ - Error: r.Error, - OperationToken: r.OperationToken, - StartTime: r.StartTime, - CloseTime: r.CloseTime, - Links: r.Links, - } - default: - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid operation state: %q", r.State) - } - - rCtx.originalHeaders.Set(interceptor.DCRedirectionAPIHeaderName, "true") - rCtx.originalHeaders.Set(interceptor.DCRedirectionSourceCellHeaderName, h.ClusterMetadata.GetCurrentClusterName()) - cc := nexusrpc.NewCompletionHTTPClient(nexusrpc.CompletionHTTPClientOptions{ - HTTPCaller: (&forwardingHTTPHeaderWrapper{ - client: client, - originalRequestHeaders: rCtx.originalHeaders, - }).Do, - }) - return cc.CompleteOperation(ctx, forwardURL, completion) -} - func (h *nexusCompletionHTTPHandler) RegisterRoutes(r *mux.Router) { r.Path("/" + commonnexus.RouteCompletionCallback.Representation()).HandlerFunc(func(w http.ResponseWriter, r *http.Request) { r.Body = http.MaxBytesReader(w, r.Body, rpc.MaxNexusAPIRequestBodyBytes) @@ -471,43 +461,30 @@ func (h *nexusCompletionHTTPHandler) RegisterRoutes(r *mux.Router) { }) } -type forwardingHTTPHeaderWrapper struct { - client *common.FrontendHTTPClient - originalRequestHeaders http.Header -} - -func (f *forwardingHTTPHeaderWrapper) Do(req *http.Request) (*http.Response, error) { - // For forwarded requests, copy the original HTTP headers without sanitization. - for k, v := range f.originalRequestHeaders { - if req.Header.Get(k) == "" { - req.Header.Set(k, v[0]) - } - } - return f.client.Do(req) -} - type requestContext struct { *nexusCompletionHandler logger log.Logger - metricsHandler metrics.Handler metricsHandlerForInterceptors metrics.Handler - namespace *namespace.Namespace - businessID string - cleanupFunctions []func(error) - requestStartTime time.Time - outcomeTag metrics.Tag - forwarded bool + namespace *namespace.Namespace // required for reporting via handleRequestError originalHeaders http.Header } +// Key to extract a *requestContext from a context.Context. +type requestContextKey struct{} + +func withRequestContext(ctx context.Context, rCtx *requestContext) context.Context { + if rCtx == nil { + return ctx + } + return context.WithValue(ctx, requestContextKey{}, rCtx) +} + +func requestContextFromContext(ctx context.Context) (*requestContext, bool) { + rCtx, ok := ctx.Value(requestContextKey{}).(*requestContext) + return rCtx, ok +} + func (c *requestContext) augmentContext(ctx context.Context, header http.Header) context.Context { - ctx = metrics.AddMetricsContext(ctx) - ctx = interceptor.AddTelemetryContext(ctx, c.metricsHandlerForInterceptors) - ctx = interceptor.PopulateCallerInfo( - ctx, - func() string { return c.namespace.Name().String() }, - func() string { return nexusCompletionMethodName }, - ) if userAgent := header.Get(headerUserAgent); userAgent != "" { // Preserve original strict behavior: only process if exactly one delimiter present. if strings.Count(userAgent, clientNameVersionDelim) == 1 { @@ -523,52 +500,47 @@ func (c *requestContext) augmentContext(ctx context.Context, header http.Header) } } } - return headers.Propagate(ctx) + return ctx +} + +func (c *requestContext) handleRequestError(err error) { + if err == nil { + return + } + if taggedErr, ok := errors.AsType[*interceptornexus.InterceptorError](err); ok { + if taggedErr.SkipServiceErrorReporting { + return + } + err = taggedErr.Err + } + c.RequestErrorHandler.HandleError( + nil, + "", + c.metricsHandlerForInterceptors, + []tag.Tag{tag.Operation(nexusCompletionMethodName), tag.WorkflowNamespace(c.namespace.Name().String())}, + err, + c.namespace.Name(), + ) } -func (c *requestContext) capturePanicAndRecordMetrics(ctxPtr *context.Context, errPtr *error) { - recovered := recover() //nolint:revive - if recovered != nil { +// finalizeCompletionRequest is the single deferred step for a Nexus completion request: capture a +// panic into errPtr, log/classify the (still raw) resulting error, then sanitize it for the +// response. Order matters and must not be split back into separate defers. +func finalizeCompletionRequest(rCtx *requestContext, errPtr *error) { + if recovered := recover(); recovered != nil { //nolint:revive err, ok := recovered.(error) if !ok { err = fmt.Errorf("panic: %v", recovered) } - - st := string(debug.Stack()) - c.logger.Error("Panic captured", tag.SysStackTrace(st), tag.Error(err)) + rCtx.logger.Error("Panic captured", tag.SysStackTrace(string(debug.Stack())), tag.Error(err)) *errPtr = err } - if *errPtr == nil { - if c.forwarded { - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("request_forwarded")) - } else { - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("success")) - } - } else if c.outcomeTag.Key != "" { - c.metricsHandler = c.metricsHandler.WithTags(c.outcomeTag) - } else { - if he, ok := errors.AsType[*nexus.HandlerError](*errPtr); ok { - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("error_" + strings.ToLower(string(he.Type)))) - } else { - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("error_internal")) - } - } - - // Record Nexus-specific metrics - c.metricsHandler.Counter(metrics.NexusCompletionRequests.Name()).Record(1) - c.metricsHandler.Histogram(metrics.NexusCompletionLatencyHistogram.Name(), metrics.Milliseconds).Record(time.Since(c.requestStartTime).Milliseconds()) - - // Record general telemetry metrics - metrics.ServiceRequests.With(c.metricsHandlerForInterceptors).Record(1) - c.TelemetryInterceptor.RecordLatencyMetrics(*ctxPtr, c.requestStartTime, c.metricsHandlerForInterceptors) - - for _, fn := range c.cleanupFunctions { - fn(*errPtr) - } + rCtx.handleRequestError(*errPtr) + *errPtr = convertInterceptorError(*errPtr) } -// TODO(bergundy): Merge this with the interceptRequest method in nexus_handler.go. -func (c *requestContext) interceptRequest(ctx context.Context, request *nexusrpc.CompletionRequest) error { +// enrich context with authInfo +func (c *requestContext) parseTLSAndAuthInfo(ctx context.Context, request *nexusrpc.CompletionRequest) (context.Context, error) { var tlsInfo *credentials.TLSInfo if request.HTTPRequest.TLS != nil { tlsInfo = &credentials.TLSInfo{ @@ -580,112 +552,12 @@ func (c *requestContext) interceptRequest(ctx context.Context, request *nexusrpc authInfo := c.AuthInterceptor.GetAuthInfo(tlsInfo, request.HTTPRequest.Header, func() string { return "" // TODO: support audience getter }) - - var claims *authorization.Claims - var err error - if authInfo != nil { - claims, err = c.AuthInterceptor.GetClaims(authInfo) - if err != nil { - return err - } - // Make the auth info and claims available on the context. - ctx = c.AuthInterceptor.EnhanceContext(ctx, authInfo, claims) - } - - _, err = c.AuthInterceptor.Authorize(ctx, claims, &authorization.CallTarget{ - APIName: nexusCompletionAPIName, - Namespace: c.namespace.Name().String(), - Request: request, - }) - if err != nil { - // If frontend.exposeAuthorizerErrors is false, Authorize err is either an explicitly set reason, or a generic - // "Request unauthorized." message. - // Otherwise, expose the underlying error. - if permissionDeniedError, ok := errors.AsType[*serviceerror.PermissionDenied](err); ok { - c.outcomeTag = metrics.OutcomeTag("unauthorized") - return commonnexus.AdaptAuthorizeError(permissionDeniedError) - } - c.outcomeTag = metrics.OutcomeTag("internal_auth_error") - c.logger.Error("Authorization internal error with processing nexus callback", tag.Error(err)) - return commonnexus.ConvertGRPCError(err, false) - } - - if err := c.NamespaceValidationInterceptor.ValidateState(c.namespace, nexusCompletionAPIName, c.businessID); err != nil { - c.outcomeTag = metrics.OutcomeTag("invalid_namespace_state") - return commonnexus.ConvertGRPCError(err, false) + if authInfo == nil { + return ctx, nil } - - // Redirect if current cluster is passive for this namespace. - if c.namespace.ActiveClusterName(namespace.RoutingKey{ID: c.businessID}) != c.ClusterMetadata.GetCurrentClusterName() { - if c.shouldForwardRequest(ctx, request.HTTPRequest.Header, c.businessID) { - c.forwarded = true - handler, forwardStartTime := c.RedirectionInterceptor.BeforeCall(nexusCompletionMethodName) - c.cleanupFunctions = append(c.cleanupFunctions, func(retErr error) { - c.RedirectionInterceptor.AfterCall(handler, forwardStartTime, c.namespace.ActiveClusterName(namespace.RoutingKey{ID: c.businessID}), c.namespace.Name().String(), retErr) - }) - // Handler methods should have special logic to forward requests if this method returns a serviceerror.NamespaceNotActive error. - return serviceerror.NewNamespaceNotActive(c.namespace.Name().String(), c.ClusterMetadata.GetCurrentClusterName(), c.namespace.ActiveClusterName(namespace.RoutingKey{ID: c.businessID})) - } - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("namespace_inactive_forwarding_disabled")) - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnavailable, "cluster inactive") - } - - c.cleanupFunctions = append(c.cleanupFunctions, func(retErr error) { - if retErr != nil { - c.RequestErrorHandler.HandleError( - request, - "", - c.metricsHandlerForInterceptors, - []tag.Tag{tag.Operation(nexusCompletionMethodName), tag.WorkflowNamespace(c.namespace.Name().String())}, - retErr, - c.namespace.Name(), - ) - } - }) - - cleanup, err := c.NamespaceConcurrencyLimitInterceptor.Allow(c.namespace.Name(), nexusCompletionAPIName, c.metricsHandlerForInterceptors, request) - c.cleanupFunctions = append(c.cleanupFunctions, func(error) { cleanup() }) - if err != nil { - c.outcomeTag = metrics.OutcomeTag("namespace_concurrency_limited") - return commonnexus.ConvertGRPCError(err, false) - } - - if err := c.NamespaceRateLimitInterceptor.Allow( - ctx, - c.namespace.Name(), - nexusCompletionAPIName, - request.HTTPRequest.Header, - ); err != nil { - c.outcomeTag = metrics.OutcomeTag("namespace_rate_limited") - return commonnexus.ConvertGRPCError(err, true) - } - - if err := c.RateLimitInterceptor.Allow(nexusCompletionAPIName, request.HTTPRequest.Header); err != nil { - c.outcomeTag = metrics.OutcomeTag("global_rate_limited") - return commonnexus.ConvertGRPCError(err, true) - } - - if err := c.clientVersionChecker.ClientSupported(ctx); err != nil { - c.outcomeTag = metrics.OutcomeTag("unsupported_client") - return commonnexus.ConvertGRPCError(err, true) - } - - return nil -} - -// TODO: copied from nexus_handler.go; should be combined with other intercept logic. -// Combines logic from RedirectionInterceptor.redirectionAllowed and some from -// SelectedAPIsForwardingRedirectionPolicy.getTargetClusterAndIsNamespaceNotActiveAutoForwarding so all -// redirection conditions can be checked at once. If either of those methods are updated, this should -// be kept in sync. -func (c *requestContext) shouldForwardRequest(ctx context.Context, header http.Header, businessID string) bool { - redirectHeader := header.Get(interceptor.DCRedirectionContextHeaderName) - redirectAllowed, err := strconv.ParseBool(redirectHeader) + claims, err := c.AuthInterceptor.GetClaims(authInfo) if err != nil { - redirectAllowed = true + return nil, err } - return redirectAllowed && - c.RedirectionInterceptor.RedirectionAllowed(ctx) && - c.namespace.IsGlobalNamespace() && - c.Config.EnableNamespaceNotActiveAutoForwarding(c.namespace.Name().String()) + return c.AuthInterceptor.EnhanceContext(ctx, authInfo, claims), nil } diff --git a/service/frontend/nexus_dispatch_result.go b/service/frontend/nexus_dispatch_result.go index ad76571d891..fee7720fbef 100644 --- a/service/frontend/nexus_dispatch_result.go +++ b/service/frontend/nexus_dispatch_result.go @@ -7,6 +7,7 @@ import ( "go.temporal.io/server/common/log/tag" commonnexus "go.temporal.io/server/common/nexus" "go.temporal.io/server/common/nexus/nexusrpc" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" ) // handleStartOperationResponse converts matching's response to a StartOperation dispatch into the result the @@ -20,7 +21,7 @@ func (c *operationContext) handleStartOperationResponse( operation string, ) (nexus.HandlerStartOperationResult[any], []nexus.Link, error) { result := commonnexus.ClassifyStartOperationDispatch(resp) - c.recordDispatchOutcome(result) + c.attributeFailureToWorker(result) switch result.Outcome { case commonnexus.DispatchOutcomeSyncSuccess: @@ -38,16 +39,16 @@ func (c *operationContext) handleStartOperationResponse( // answer, reported to the caller as a Nexus operation error rather than a handler error. cause, internalErr := c.convertWorkerFailure(result.Failure, operation) if internalErr != nil { - return nil, nil, internalErr + return nil, nil, dispatchError(result, internalErr) } state := nexus.OperationStateFailed if result.Failure.GetCanceledFailureInfo() != nil { state = nexus.OperationStateCanceled } - return nil, nil, c.operationError(state, cause, operation) + return nil, nil, dispatchError(result, c.operationError(state, cause, operation)) default: - return nil, nil, c.failedDispatchToNexusError(result, operation) + return nil, nil, dispatchError(result, c.failedDispatchToNexusError(result, operation)) } } @@ -59,12 +60,24 @@ func (c *operationContext) handleCancelOperationResponse( operation string, ) error { result := commonnexus.ClassifyCancelOperationDispatch(resp) - c.recordDispatchOutcome(result) + c.attributeFailureToWorker(result) if result.Outcome == commonnexus.DispatchOutcomeCancelAccepted { return nil } - return c.failedDispatchToNexusError(result, operation) + return dispatchError(result, c.failedDispatchToNexusError(result, operation)) +} + +// dispatchError wraps the error with the result's outcome tag so it can +// be tagged in turn by the nexus telemetry outermost interceptor +func dispatchError(result commonnexus.DispatchResult, err error) error { + if err == nil { + return nil + } + return &interceptornexus.InterceptorError{ + Err: err, + Outcome: result.OutcomeTag().Value, + } } // failedDispatchToNexusError converts the outcomes that mean the task was never handled, or was @@ -140,10 +153,7 @@ func (c *operationContext) operationError( return opErr } -// recordDispatchOutcome tags the request's metrics with the dispatch outcome and, when the dispatch -// did not succeed, attributes the failure to the worker in the response header. -func (c *operationContext) recordDispatchOutcome(result commonnexus.DispatchResult) { - c.metricsHandler = c.metricsHandler.WithTags(result.OutcomeTag()) +func (c *operationContext) attributeFailureToWorker(result commonnexus.DispatchResult) { if !result.Outcome.Succeeded() { c.setFailureSource(commonnexus.FailureSourceWorker) } diff --git a/service/frontend/nexus_dispatch_result_test.go b/service/frontend/nexus_dispatch_result_test.go index 093e038d3b4..3cbba35f029 100644 --- a/service/frontend/nexus_dispatch_result_test.go +++ b/service/frontend/nexus_dispatch_result_test.go @@ -1,8 +1,10 @@ package frontend import ( + "context" "encoding/json" "testing" + "time" "github.com/nexus-rpc/sdk-go/nexus" "github.com/stretchr/testify/require" @@ -11,40 +13,56 @@ import ( failurepb "go.temporal.io/api/failure/v1" nexuspb "go.temporal.io/api/nexus/v1" "go.temporal.io/server/api/matchingservice/v1" + "go.temporal.io/server/common/log" + "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/metrics/metricstest" commonnexus "go.temporal.io/server/common/nexus" + rpcinterceptor "go.temporal.io/server/common/rpc/interceptor" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" ) // These tests pin down how the frontend turns matching's DispatchNexusTaskResponse into the result the // Nexus SDK serializes back to the caller. Every arm of the response oneof is wire-visible: the error // type decides the HTTP status, and the outcome tag and failure-source header are consumed by -// dashboards and by interceptRequest's error-reporting cleanup. They are asserted here so the shared +// dashboards and by request error-reporting cleanup. They are asserted here so the shared // classifier introduced alongside them cannot silently change any of it. -// outcomeTagOf reads the outcome tag accumulated on the context's metrics handler. -func outcomeTagOf(t *testing.T, oc *operationContext) string { - t.Helper() - mh, ok := oc.metricsHandler.(*metricstest.CaptureHandler) - require.True(t, ok, "expected a capture handler") - capture := mh.StartCapture() - oc.metricsHandler.Counter("test").Record(1) - mh.StopCapture(capture) - snap := capture.Snapshot() - require.Len(t, snap["test"], 1) - return snap["test"][0].Tags["outcome"] -} - func failureSourceOf(oc *operationContext) string { return oc.responseHeaders[commonnexus.FailureSourceHeaderName] } -func testOperationContext() *operationContext { - return newOperationContext(contextOptions{ - namespaceState: enumspb.NAMESPACE_STATE_REGISTERED, - quota: 1, - namespaceRateLimitAllow: true, - rateLimitAllow: true, - }) +func requireDispatchOutcome(t *testing.T, err error, outcome string) { + t.Helper() + var interceptorErr *interceptornexus.InterceptorError + require.ErrorAs(t, err, &interceptorErr) + require.Equal(t, outcome, interceptorErr.Outcome) +} + +func requireRecordedDispatchOutcome( + t *testing.T, + input interceptornexus.InterceptorInput, + expectedOutcome string, + handler func(*operationContext) error, +) { + t.Helper() + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + telemetry := rpcinterceptor.NewTelemetryInterceptor(nil, metricsHandler, log.NewNoopLogger(), nil, nil) + _, err := telemetry.InterceptNexusOutermost( + context.Background(), + input, + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + return nil, handler(testOperationContext()) + }, + ) + requireDispatchOutcome(t, err, expectedOutcome) + + snapshot := capture.Snapshot() + require.Len(t, snapshot[metrics.NexusRequests.Name()], 1) + outcomeTag := metrics.OutcomeTag(expectedOutcome) + require.Equal(t, expectedOutcome, snapshot[metrics.NexusRequests.Name()][0].Tags[outcomeTag.Key]) } // startOperationResponse wraps a StartOperationResponse in the matching response envelope. The oneof @@ -59,6 +77,78 @@ func startOperationResponse(sor *nexuspb.StartOperationResponse) *matchingservic } } +func TestDispatchErrorsPreserveOutcomeForTelemetry(t *testing.T) { + handlerFailure := &matchingservice.DispatchNexusTaskResponse{ + Outcome: &matchingservice.DispatchNexusTaskResponse_Failure{ + Failure: &failurepb.Failure{ + FailureInfo: &failurepb.Failure_NexusHandlerFailureInfo{ + NexusHandlerFailureInfo: &failurepb.NexusHandlerFailureInfo{ + Type: string(nexus.HandlerErrorTypeBadRequest), + }, + }, + }, + }, + } + requestTimeout := &matchingservice.DispatchNexusTaskResponse{ + Outcome: &matchingservice.DispatchNexusTaskResponse_RequestTimeout{ + RequestTimeout: &matchingservice.DispatchNexusTaskResponse_Timeout{}, + }, + } + operationFailure := startOperationResponse(&nexuspb.StartOperationResponse{ + Variant: &nexuspb.StartOperationResponse_Failure{ + Failure: &failurepb.Failure{ + FailureInfo: &failurepb.Failure_ApplicationFailureInfo{ + ApplicationFailureInfo: &failurepb.ApplicationFailureInfo{}, + }, + }, + }, + }) + + for _, tc := range []struct { + name string + response *matchingservice.DispatchNexusTaskResponse + outcome string + }{ + {name: "handler failure", response: handlerFailure, outcome: "handler_error:BAD_REQUEST"}, + {name: "request timeout", response: requestTimeout, outcome: "handler_timeout"}, + {name: "operation failure", response: operationFailure, outcome: "failure"}, + {name: "unrecognized outcome", response: &matchingservice.DispatchNexusTaskResponse{}, outcome: "handler_error:EMPTY_OUTCOME"}, + } { + t.Run("start "+tc.name, func(t *testing.T) { + requireRecordedDispatchOutcome( + t, + interceptornexus.NewStartOpInput("s", "o", "n", time.Now(), nexus.StartOperationOptions{}, nil, interceptornexus.ForwardingInfo{}, interceptornexus.RequestMetadata{}), + tc.outcome, + func(oc *operationContext) error { + _, _, err := oc.handleStartOperationResponse(tc.response, "op") + return err + }, + ) + }) + } + + for _, tc := range []struct { + name string + response *matchingservice.DispatchNexusTaskResponse + outcome string + }{ + {name: "handler failure", response: handlerFailure, outcome: "handler_error:BAD_REQUEST"}, + {name: "request timeout", response: requestTimeout, outcome: "handler_timeout"}, + {name: "unrecognized outcome", response: &matchingservice.DispatchNexusTaskResponse{}, outcome: "handler_error:EMPTY_OUTCOME"}, + } { + t.Run("cancel "+tc.name, func(t *testing.T) { + requireRecordedDispatchOutcome( + t, + interceptornexus.NewCancelOpInput("s", "o", "n", time.Now(), nexus.CancelOperationOptions{}, "t", interceptornexus.ForwardingInfo{}, interceptornexus.RequestMetadata{}), + tc.outcome, + func(oc *operationContext) error { + return oc.handleCancelOperationResponse(tc.response, "op") + }, + ) + }) + } +} + func TestHandleStartOperationResponse_SyncSuccess(t *testing.T) { oc := testOperationContext() payload := &commonpb.Payload{Data: []byte("hello")} @@ -82,7 +172,6 @@ func TestHandleStartOperationResponse_SyncSuccess(t *testing.T) { require.Len(t, links, 1) require.Equal(t, "http://links.test/valid", links[0].URL.String()) require.Equal(t, "some.Type", links[0].Type) - require.Equal(t, "sync_success", outcomeTagOf(t, oc)) require.Empty(t, failureSourceOf(oc), "success must not be attributed to the worker") } @@ -100,7 +189,6 @@ func TestHandleStartOperationResponse_SyncSuccess_NoPayloadNoLinks(t *testing.T) require.True(t, ok) require.Nil(t, sync.Value) require.Empty(t, links) - require.Equal(t, "sync_success", outcomeTagOf(t, oc)) } func TestHandleStartOperationResponse_AsyncSuccess_PrefersOperationToken(t *testing.T) { @@ -122,7 +210,6 @@ func TestHandleStartOperationResponse_AsyncSuccess_PrefersOperationToken(t *test require.True(t, ok, "expected an async result, got %T", result) require.Equal(t, "token", async.OperationToken) require.Len(t, links, 1) - require.Equal(t, "async_success", outcomeTagOf(t, oc)) require.Empty(t, failureSourceOf(oc)) } @@ -191,7 +278,7 @@ func TestHandleStartOperationResponse_HandlerFailure(t *testing.T) { require.Equal(t, "handler said no", handlerErr.Message) require.Equal(t, tc.wantRetryable, handlerErr.Retryable()) require.NoError(t, handlerErr.Cause, "no cause on the wire means no cause on the error") - require.Equal(t, "handler_error:BAD_REQUEST", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "handler_error:BAD_REQUEST") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) }) } @@ -255,7 +342,7 @@ func TestHandleStartOperationResponse_WorkerFailure_NotAHandlerError(t *testing. var handlerErr *nexus.HandlerError require.NotErrorAs(t, err, &handlerErr, "not reported as a handler error today") // There is no handler error type to report, so the tag bounds to UNKNOWN. - require.Equal(t, "handler_error:UNKNOWN", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "handler_error:UNKNOWN") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) } @@ -282,7 +369,7 @@ func TestHandleStartOperationResponse_DeprecatedHandlerError(t *testing.T) { deprecatedCause, ok := handlerErr.Cause.(*nexus.FailureError) require.True(t, ok, "expected a Nexus FailureError cause, got %T", handlerErr.Cause) require.Equal(t, "slow down", deprecatedCause.Failure.Message) - require.Equal(t, "handler_error:RESOURCE_EXHAUSTED", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "handler_error:RESOURCE_EXHAUSTED") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) } @@ -301,7 +388,7 @@ func TestHandleStartOperationResponse_RequestTimeout(t *testing.T) { require.ErrorAs(t, err, &handlerErr) require.Equal(t, nexus.HandlerErrorTypeUpstreamTimeout, handlerErr.Type) require.Equal(t, "upstream timeout", handlerErr.Message) - require.Equal(t, "handler_timeout", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "handler_timeout") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) } @@ -351,7 +438,7 @@ func TestHandleStartOperationResponse_OperationFailure(t *testing.T) { require.NotNil(t, opErr.OriginalFailure) require.Equal(t, "true", opErr.OriginalFailure.Metadata["unwrap-error"]) require.NotNil(t, opErr.OriginalFailure.Cause) - require.Equal(t, "failure", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "failure") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) }) } @@ -386,7 +473,7 @@ func TestHandleStartOperationResponse_OperationFailure_UnconvertibleFailureIsInt var opErr *nexus.OperationError require.NotErrorAs(t, err, &opErr, "an unreadable failure is not a legitimate operation error") // The outcome was still classified as an operation failure, so the tag and header stand. - require.Equal(t, "failure", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "failure") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) } @@ -415,7 +502,7 @@ func TestHandleStartOperationResponse_HandlerFailure_UnconvertibleCauseIsInterna require.ErrorAs(t, err, &handlerErr) require.Equal(t, nexus.HandlerErrorTypeInternal, handlerErr.Type, "the worker's own BAD_REQUEST must not survive a failed conversion") - require.Equal(t, "handler_error:BAD_REQUEST", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "handler_error:BAD_REQUEST") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) } @@ -459,7 +546,7 @@ func TestHandleStartOperationResponse_DeprecatedOperationError(t *testing.T) { require.Equal(t, "worker canceled it", cause.Failure.Message) require.NotNil(t, opErr.OriginalFailure) require.Equal(t, "true", opErr.OriginalFailure.Metadata["unwrap-error"]) - require.Equal(t, "operation_error", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "operation_error") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) } @@ -502,7 +589,7 @@ func TestHandleStartOperationResponse_DeprecatedOperationErrorReEncodesWorkerFai require.NoError(t, json.Unmarshal(details[0].GetData(), &workerFailure)) require.Equal(t, map[string]string{"k": "v"}, workerFailure.Metadata) require.JSONEq(t, `"details"`, string(workerFailure.Details)) - require.Equal(t, "operation_error", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "operation_error") } // Anything the frontend cannot interpret is blamed on the worker and reported as an internal error. @@ -545,7 +632,7 @@ func TestHandleStartOperationResponse_UnrecognizedOutcomes(t *testing.T) { require.ErrorAs(t, err, &handlerErr) require.Equal(t, nexus.HandlerErrorTypeInternal, handlerErr.Type) require.Equal(t, "empty outcome", handlerErr.Message) - require.Equal(t, "handler_error:EMPTY_OUTCOME", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "handler_error:EMPTY_OUTCOME") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) }) } @@ -581,7 +668,6 @@ func TestHandleCancelOperationResponse_Success(t *testing.T) { t.Run(tc.name, func(t *testing.T) { oc := testOperationContext() require.NoError(t, oc.handleCancelOperationResponse(tc.resp, "op")) - require.Equal(t, "success", outcomeTagOf(t, oc)) require.Empty(t, failureSourceOf(oc)) }) } @@ -607,7 +693,7 @@ func TestHandleCancelOperationResponse_HandlerFailure(t *testing.T) { require.ErrorAs(t, err, &handlerErr) require.Equal(t, nexus.HandlerErrorTypeNotFound, handlerErr.Type) require.Equal(t, "cannot cancel", handlerErr.Message) - require.Equal(t, "handler_error:NOT_FOUND", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "handler_error:NOT_FOUND") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) } @@ -627,7 +713,7 @@ func TestHandleCancelOperationResponse_DeprecatedHandlerError(t *testing.T) { var handlerErr *nexus.HandlerError require.ErrorAs(t, err, &handlerErr) require.Equal(t, nexus.HandlerErrorTypeNotImplemented, handlerErr.Type) - require.Equal(t, "handler_error:NOT_IMPLEMENTED", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "handler_error:NOT_IMPLEMENTED") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) } @@ -644,7 +730,7 @@ func TestHandleCancelOperationResponse_RequestTimeout(t *testing.T) { require.ErrorAs(t, err, &handlerErr) require.Equal(t, nexus.HandlerErrorTypeUpstreamTimeout, handlerErr.Type) require.Equal(t, "upstream timeout", handlerErr.Message) - require.Equal(t, "handler_timeout", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "handler_timeout") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) } @@ -655,7 +741,7 @@ func TestHandleCancelOperationResponse_UnrecognizedOutcome(t *testing.T) { require.ErrorAs(t, err, &handlerErr) require.Equal(t, nexus.HandlerErrorTypeInternal, handlerErr.Type) require.Equal(t, "empty outcome", handlerErr.Message) - require.Equal(t, "handler_error:EMPTY_OUTCOME", outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, "handler_error:EMPTY_OUTCOME") require.Equal(t, commonnexus.FailureSourceWorker, failureSourceOf(oc)) } @@ -725,7 +811,7 @@ func TestHandleStartOperationResponse_HandlerErrorTypeTagIsBounded(t *testing.T) _, _, err := oc.handleStartOperationResponse(resp, "op") require.Error(t, err) - require.Equal(t, tc.wantTag, outcomeTagOf(t, oc)) + requireDispatchOutcome(t, err, tc.wantTag) // The error itself still carries the worker's real type; only the metric is bounded. var handlerErr *nexus.HandlerError require.ErrorAs(t, err, &handlerErr) @@ -744,6 +830,7 @@ func TestHandleCancelOperationResponse_DeprecatedHandlerErrorTypeTagIsBounded(t }, } - require.Error(t, oc.handleCancelOperationResponse(resp, "op")) - require.Equal(t, "handler_error:UNKNOWN", outcomeTagOf(t, oc)) + err := oc.handleCancelOperationResponse(resp, "op") + require.Error(t, err) + requireDispatchOutcome(t, err, "handler_error:UNKNOWN") } diff --git a/service/frontend/nexus_forward_interceptor.go b/service/frontend/nexus_forward_interceptor.go new file mode 100644 index 00000000000..c7c23413b06 --- /dev/null +++ b/service/frontend/nexus_forward_interceptor.go @@ -0,0 +1,355 @@ +package frontend + +import ( + "context" + "errors" + "net/http" + "net/http/httptrace" + "net/url" + "strconv" + "time" + + "github.com/nexus-rpc/sdk-go/nexus" + "go.temporal.io/server/common" + "go.temporal.io/server/common/api" + "go.temporal.io/server/common/cluster" + "go.temporal.io/server/common/headers" + "go.temporal.io/server/common/log" + "go.temporal.io/server/common/log/tag" + "go.temporal.io/server/common/namespace" + commonnexus "go.temporal.io/server/common/nexus" + "go.temporal.io/server/common/nexus/nexusrpc" + "go.temporal.io/server/common/rpc/interceptor" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" +) + +type nexusForwardingInterceptor struct { + logger log.Logger + clusterMetadata cluster.Metadata + redirectionInterceptor *interceptor.Redirection + forwardingClients frontendHTTPClientCache + serviceConfig *Config + httpTraceProvider commonnexus.HTTPClientTraceProvider +} + +type frontendHTTPClientCache interface { + Get(targetClusterName string) (*common.FrontendHTTPClient, error) +} + +func newNexusForwardingInterceptor( + logger log.Logger, + clusterMetadata cluster.Metadata, + redirectionInterceptor *interceptor.Redirection, + forwardingClients *cluster.FrontendHTTPClientCache, + serviceConfig *Config, + httpTraceProvider commonnexus.HTTPClientTraceProvider, +) *nexusForwardingInterceptor { + return &nexusForwardingInterceptor{ + logger: logger, + clusterMetadata: clusterMetadata, + redirectionInterceptor: redirectionInterceptor, + forwardingClients: forwardingClients, + serviceConfig: serviceConfig, + httpTraceProvider: httpTraceProvider, + } +} + +func (i *nexusForwardingInterceptor) InterceptNexus( + ctx context.Context, + in interceptornexus.InterceptorInput, + next interceptornexus.HandlerFunc, +) (out any, retErr error) { + info := in.ForwardingInfo() + header := in.Header() + namespaceEntry, err := in.NamespaceEntry() + if err != nil { + return nil, &interceptornexus.InterceptorError{ + Err: err, + Outcome: "interceptor_failed", + SkipServiceErrorReporting: true, + } + } + currentCluster := i.clusterMetadata.GetCurrentClusterName() + targetCluster := namespaceEntry.ActiveClusterName(namespace.RoutingKey{ID: info.BusinessID}) + if !namespaceEntry.IsGlobalNamespace() || targetCluster == currentCluster { + return next(ctx, in) + } + if !i.shouldForwardRequest(ctx, header, namespaceEntry) { + return nil, &interceptornexus.InterceptorError{ + Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnavailable, "cluster inactive"), + Outcome: "namespace_inactive_forwarding_disabled", + SkipServiceErrorReporting: true, + } + } + + interceptornexus.SetOutcomeOverride(ctx, "request_forwarded") + + metricsHandler, forwardStartTime := i.redirectionInterceptor.BeforeCall( + interceptor.DCRedirectionMetricsPrefix + api.MethodName(in.APIName()), + ) + defer func() { + redirectionErr := retErr + if taggedErr, ok := errors.AsType[*interceptornexus.InterceptorError](retErr); ok { + redirectionErr = taggedErr.Err + } + i.redirectionInterceptor.AfterCall(metricsHandler, forwardStartTime, targetCluster, namespaceEntry.Name().String(), redirectionErr) + }() + + logTags := []tag.Tag{ + tag.SourceCluster(i.clusterMetadata.GetCurrentClusterName()), + tag.TargetCluster(targetCluster), + tag.Operation(in.MethodName()), + tag.WorkflowNamespace(namespaceEntry.Name().String()), + } + if endpointName := in.EndpointName(); endpointName != "" { + // empty on namespace/task-queue routed requests + logTags = append(logTags, tag.Endpoint(endpointName)) + } + if operationName := in.OperationName(); operationName != "" { + // empty for completion requests + logTags = append(logTags, tag.NexusOperation(operationName)) + } + // Retrieve loggers for operation type-specific tags. + baseLogger := i.logger + if rCtx, ok := requestContextFromContext(ctx); ok { + baseLogger = rCtx.logger + } else if oc, ok := operationContextFromContext(ctx); ok { + baseLogger = oc.logger + } + logger := log.With(baseLogger, logTags...) + + switch request := in.(type) { + case interceptornexus.StartOpInput: + out, retErr = i.forwardStartOperation(ctx, logger, request, info, namespaceEntry, targetCluster) + case interceptornexus.CancelOpInput: + retErr = i.forwardCancelOperation(ctx, logger, request, info, namespaceEntry, targetCluster) + case interceptornexus.CompleteOpInput: + retErr = i.forwardCompleteOperation(ctx, logger, request, info, namespaceEntry, targetCluster) + default: + return nil, &interceptornexus.InterceptorError{ + Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnavailable, "forwarding failed, unknown operation type"), + SkipServiceErrorReporting: true, + } + } + return out, retErr +} + +func (i *nexusForwardingInterceptor) shouldForwardRequest( + ctx context.Context, + header headers.HeaderGetter, + namespaceEntry *namespace.Namespace, +) bool { + redirectAllowed, err := strconv.ParseBool(header.Get(interceptor.DCRedirectionContextHeaderName)) + if err != nil { + redirectAllowed = true + } + return redirectAllowed && + i.redirectionInterceptor.RedirectionAllowed(ctx) && + i.serviceConfig.EnableNamespaceNotActiveAutoForwarding(namespaceEntry.Name().String()) +} + +func (i *nexusForwardingInterceptor) forwardStartOperation( + ctx context.Context, + logger log.Logger, + request interceptornexus.StartOpInput, + info interceptornexus.ForwardingInfo, + namespaceEntry *namespace.Namespace, + targetCluster string, +) (any, error) { + logger = log.With( + logger, + tag.RequestID(request.StartOperationOptions.RequestID), + ) + request.StartOperationOptions.Header[interceptor.DCRedirectionAPIHeaderName] = "true" + request.StartOperationOptions.Header[interceptor.DCRedirectionSourceCellHeaderName] = i.clusterMetadata.GetCurrentClusterName() + client, err := i.nexusClientForActiveCluster(ctx, logger, request.ServiceName(), info, namespaceEntry, targetCluster) + if err != nil { + return nil, err + } + ctx = i.withForwardingTrace(ctx, "StartNexusOperation", request.OperationName(), request.StartOperationOptions.RequestID, info, namespaceEntry, targetCluster) + response, err := client.StartOperation(ctx, request.OperationName(), request.StartOperationInput.Reader, request.StartOperationOptions) + if err != nil { + logger.Error("received error from remote cluster for forwarded Nexus start operation request", tag.Error(err)) + return nil, &interceptornexus.InterceptorError{Err: err, Outcome: "forwarded_request_error", SkipServiceErrorReporting: true} + } + if response.Successful != nil { + return &nexus.HandlerStartOperationResultSync[any]{Value: response.Successful.Reader}, nil + } + return &nexus.HandlerStartOperationResultAsync{OperationToken: response.Pending.Token}, nil +} + +func (i *nexusForwardingInterceptor) forwardCancelOperation( + ctx context.Context, + logger log.Logger, + request interceptornexus.CancelOpInput, + info interceptornexus.ForwardingInfo, + namespaceEntry *namespace.Namespace, + targetCluster string, +) error { + request.CancelOperationOptions.Header[interceptor.DCRedirectionAPIHeaderName] = "true" + request.CancelOperationOptions.Header[interceptor.DCRedirectionSourceCellHeaderName] = i.clusterMetadata.GetCurrentClusterName() + client, err := i.nexusClientForActiveCluster(ctx, logger, request.ServiceName(), info, namespaceEntry, targetCluster) + if err != nil { + return err + } + handle, err := client.NewOperationHandle(request.OperationName(), request.CancellationToken) + if err != nil { + logger.Warn("invalid Nexus cancel operation", tag.Error(err)) + return &interceptornexus.InterceptorError{ + Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid operation"), + Outcome: "error_bad_request", + } + } + ctx = i.withForwardingTrace(ctx, "CancelNexusOperation", request.OperationName(), "", info, namespaceEntry, targetCluster) + if err := handle.Cancel(ctx, request.CancelOperationOptions); err != nil { + logger.Error("received error from remote cluster for forwarded Nexus cancel operation request", tag.Error(err)) + return &interceptornexus.InterceptorError{Err: err, Outcome: "forwarded_request_error", SkipServiceErrorReporting: true} + } + return nil +} + +func (i *nexusForwardingInterceptor) forwardCompleteOperation( + ctx context.Context, + logger log.Logger, + request interceptornexus.CompleteOpInput, + info interceptornexus.ForwardingInfo, + namespaceEntry *namespace.Namespace, + targetCluster string, +) error { + client, err := i.forwardingClients.Get(targetCluster) + if err != nil { + logger.Error("unable to get HTTP client for forward request", tag.Error(err)) + return &interceptornexus.InterceptorError{Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "internal error"), Outcome: "request_forwarding_failed", SkipServiceErrorReporting: true} + } + forwardURL, err := url.JoinPath(client.BaseURL(), commonnexus.RouteCompletionCallback.Path(namespaceEntry.Name().String())) + if err != nil { + logger.Error("failed to construct forwarding request URL", tag.Error(err)) + return &interceptornexus.InterceptorError{Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "internal error"), Outcome: "request_forwarding_failed", SkipServiceErrorReporting: true} + } + info.OriginalRequestHeaders.Set(interceptor.DCRedirectionAPIHeaderName, "true") + info.OriginalRequestHeaders.Set(interceptor.DCRedirectionSourceCellHeaderName, i.clusterMetadata.GetCurrentClusterName()) + completion, err := completeOperationOptions(request.CompletionRequest) + if err != nil { + return &interceptornexus.InterceptorError{Err: err, Outcome: "forwarded_request_error", SkipServiceErrorReporting: true} + } + ctx = i.withForwardingTrace(ctx, "CompleteNexusOperation", "", "", info, namespaceEntry, targetCluster) + err = nexusrpc.NewCompletionHTTPClient(nexusrpc.CompletionHTTPClientOptions{ + // completions dont report a failure source back to the caller through headers + HTTPCaller: (&nexusForwardingHTTPHeaderWrapper{client: client, originalRequestHeaders: info.OriginalRequestHeaders}).Do, + }).CompleteOperation(ctx, forwardURL, completion) + if err != nil { + return &interceptornexus.InterceptorError{Err: err, Outcome: "forwarded_request_error", SkipServiceErrorReporting: true} + } + return nil +} + +func completeOperationOptions(request *nexusrpc.CompletionRequest) (nexusrpc.CompleteOperationOptions, error) { + switch request.State { + case nexus.OperationStateSucceeded: + return nexusrpc.CompleteOperationOptions{Result: request.Result.Reader, OperationToken: request.OperationToken, StartTime: request.StartTime, CloseTime: request.CloseTime, Links: request.Links}, nil + case nexus.OperationStateFailed, nexus.OperationStateCanceled: + return nexusrpc.CompleteOperationOptions{Error: request.Error, OperationToken: request.OperationToken, StartTime: request.StartTime, CloseTime: request.CloseTime, Links: request.Links}, nil + default: + return nexusrpc.CompleteOperationOptions{}, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid operation state: %q", request.State) + } +} + +func (i *nexusForwardingInterceptor) nexusClientForActiveCluster( + ctx context.Context, + logger log.Logger, + service string, + info interceptornexus.ForwardingInfo, + namespaceEntry *namespace.Namespace, + targetCluster string, +) (*nexusrpc.HTTPClient, error) { + var setFailureSource func(string) // required for setting the source in case of a failure + if oc, ok := operationContextFromContext(ctx); ok { + setFailureSource = oc.setFailureSource + } + httpClient, err := i.forwardingClients.Get(targetCluster) + if err != nil { + logger.Error("failed to forward Nexus request: error creating HTTP client", tag.Error(err)) + return nil, &interceptornexus.InterceptorError{Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "request forwarding failed"), Outcome: "request_forwarding_failed", SkipServiceErrorReporting: true} + } + var baseURL string + if i.serviceConfig.NexusForwardRequestUseEndpoint() && info.EndpointID != "" { + baseURL, err = url.JoinPath(httpClient.BaseURL(), commonnexus.RouteDispatchNexusTaskByEndpoint.Path(info.EndpointID)) + } else { + baseURL, err = url.JoinPath(httpClient.BaseURL(), commonnexus.RouteDispatchNexusTaskByNamespaceAndTaskQueue.Path(commonnexus.NamespaceAndTaskQueue{Namespace: namespaceEntry.Name().String(), TaskQueue: info.TaskQueue})) + } + if err != nil { + logger.Error("failed to forward Nexus request: error constructing ServiceBaseURL", tag.URL(httpClient.BaseURL()), tag.WorkflowTaskQueueName(info.TaskQueue), tag.Error(err)) + return nil, &interceptornexus.InterceptorError{Err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "request forwarding failed"), Outcome: "request_forwarding_failed", SkipServiceErrorReporting: true} + } + return nexusrpc.NewHTTPClient(nexusrpc.HTTPClientOptions{ + HTTPCaller: (&nexusForwardingHTTPHeaderWrapper{client: httpClient, originalRequestHeaders: info.OriginalRequestHeaders, setFailureSource: setFailureSource}).Do, + BaseURL: baseURL, + Service: service, + }) +} + +func (i *nexusForwardingInterceptor) withForwardingTrace( + ctx context.Context, + method string, + operation string, + requestID string, + info interceptornexus.ForwardingInfo, + namespaceEntry *namespace.Namespace, + targetCluster string, +) context.Context { + if i.httpTraceProvider == nil { + return ctx + } + traceLogger := i.logger + tags := []tag.Tag{ + tag.AttemptStart(time.Now().UTC()), + tag.SourceCluster(i.clusterMetadata.GetCurrentClusterName()), + tag.TargetCluster(targetCluster), + } + if rCtx, ok := requestContextFromContext(ctx); ok { + traceLogger = rCtx.logger + } else { + tags = append(tags, + tag.Operation(method), + tag.WorkflowNamespace(namespaceEntry.Name().String()), + ) + if requestID != "" { + tags = append(tags, tag.RequestID(requestID)) + } + if operation != "" { + tags = append(tags, tag.NexusOperation(operation)) + } + if info.EndpointName != "" { + tags = append(tags, tag.Endpoint(info.EndpointName)) + } + } + traceLogger = log.With(traceLogger, tags...) + if trace := i.httpTraceProvider.NewForwardingTrace(traceLogger); trace != nil { + return httptrace.WithClientTrace(ctx, trace) + } + return ctx +} + +type nexusForwardingHTTPHeaderWrapper struct { + client *common.FrontendHTTPClient + originalRequestHeaders http.Header + setFailureSource func(string) +} + +func (f *nexusForwardingHTTPHeaderWrapper) Do(request *http.Request) (*http.Response, error) { + // for forwarded requests, copy the original HTTP headers without sanitization. + for name, values := range f.originalRequestHeaders { + if request.Header.Get(name) == "" { + request.Header.Set(name, values[0]) + } + } + response, err := f.client.Do(request) + if err != nil { + return nil, err + } + + if source := response.Header.Get(commonnexus.FailureSourceHeaderName); source != "" && f.setFailureSource != nil { + f.setFailureSource(source) + } + return response, nil +} diff --git a/service/frontend/nexus_forward_interceptor_test.go b/service/frontend/nexus_forward_interceptor_test.go new file mode 100644 index 00000000000..405aa48d361 --- /dev/null +++ b/service/frontend/nexus_forward_interceptor_test.go @@ -0,0 +1,215 @@ +package frontend + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strconv" + "testing" + "time" + + "github.com/nexus-rpc/sdk-go/nexus" + "github.com/stretchr/testify/require" + persistencespb "go.temporal.io/server/api/persistence/v1" + "go.temporal.io/server/common" + "go.temporal.io/server/common/clock" + "go.temporal.io/server/common/cluster" + "go.temporal.io/server/common/cluster/clustertest" + "go.temporal.io/server/common/config" + "go.temporal.io/server/common/dynamicconfig" + "go.temporal.io/server/common/log" + "go.temporal.io/server/common/metrics" + "go.temporal.io/server/common/namespace" + "go.temporal.io/server/common/rpc/interceptor" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" +) + +func TestNexusForwardingInterceptorInterceptNexus(t *testing.T) { + metadata := clustertest.NewMetadataForTest(cluster.NewTestClusterMetadataConfig(true, true)) + currentCluster := cluster.TestCurrentClusterName + remoteCluster := cluster.TestAlternativeClusterName + + type requestDisposition int + const ( + requestFailed requestDisposition = iota + requestHandledLocally + requestForwarded + ) + + var receivedHeaders http.Header + // dummy server to simulate forwarded req + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) { + receivedHeaders = request.Header.Clone() + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusCreated) + _, _ = fmt.Fprint(w, `{"token":"operation-token","state":"running"}`) + })) + defer server.Close() + forwardingClient := testFrontendHTTPClientCache{clients: map[string]*common.FrontendHTTPClient{ + remoteCluster: { + Client: *server.Client(), + Address: server.Listener.Addr().String(), + Scheme: "http", + }, + }} + for _, tc := range []struct { + name string + namespace *namespace.Namespace + forwardingOn bool + redirectAllowed *bool + expectedOutcome string + disposition requestDisposition + }{ + { + name: "local namespace should resolve", + namespace: namespace.NewLocalNamespaceForTest( + &persistencespb.NamespaceInfo{Name: testNamespace}, + nil, + currentCluster, + ), + disposition: requestHandledLocally, + }, + { + name: "global namespace with forwarding enabled should redirect", + namespace: namespace.NewNamespaceForTest( + &persistencespb.NamespaceInfo{Name: testNamespace}, + nil, + true, + &persistencespb.NamespaceReplicationConfig{ActiveClusterName: remoteCluster, Clusters: []string{currentCluster, remoteCluster}}, + 0, + ), + forwardingOn: true, + disposition: requestForwarded, + }, + { + name: "global namespace with forwarding disabled should fail", + namespace: namespace.NewNamespaceForTest( + &persistencespb.NamespaceInfo{Name: testNamespace}, + nil, + true, + &persistencespb.NamespaceReplicationConfig{ActiveClusterName: remoteCluster, Clusters: []string{currentCluster, remoteCluster}}, + 0, + ), + forwardingOn: false, + expectedOutcome: "namespace_inactive_forwarding_disabled", + disposition: requestFailed, + }, + { + name: "global namespace with redirection disabled should fail", + namespace: namespace.NewNamespaceForTest( + &persistencespb.NamespaceInfo{Name: testNamespace}, + nil, + true, + &persistencespb.NamespaceReplicationConfig{ActiveClusterName: remoteCluster, Clusters: []string{currentCluster, remoteCluster}}, + 0, + ), + forwardingOn: true, + redirectAllowed: new(false), + expectedOutcome: "namespace_inactive_forwarding_disabled", + disposition: requestFailed, + }, + { + name: "global namespace with forwarding enabled to unknown cluster fails", + namespace: namespace.NewNamespaceForTest( + &persistencespb.NamespaceInfo{Name: testNamespace}, + nil, + true, + &persistencespb.NamespaceReplicationConfig{ActiveClusterName: "unknown-cluster", Clusters: []string{currentCluster}}, + 0, + ), + forwardingOn: true, + expectedOutcome: "request_forwarding_failed", + disposition: requestFailed, + }, + } { + t.Run(tc.name, func(t *testing.T) { + receivedHeaders = nil + options := nexus.StartOperationOptions{Header: nexus.Header{"X-Request": "request"}} + if tc.redirectAllowed != nil { + options.Header[interceptor.DCRedirectionContextHeaderName] = strconv.FormatBool(*tc.redirectAllowed) + } + requestInput := nexus.NewLazyValue(nexus.DefaultSerializer(), &nexus.Reader{ + ReadCloser: io.NopCloser(bytes.NewBufferString(`"input"`)), + Header: nexus.Header{"type": "json"}, + }) + forwardingInfo := interceptornexus.ForwardingInfo{ + OriginalRequestHeaders: http.Header{"X-Original": {"original"}}, + TaskQueue: "task-queue", + } + forwarder := &nexusForwardingInterceptor{ + logger: log.NewNoopLogger(), + clusterMetadata: metadata, + forwardingClients: forwardingClient, + redirectionInterceptor: interceptor.NewRedirection( + dynamicconfig.GetBoolPropertyFnFilteredByNamespace(true), + dynamicconfig.GetBoolPropertyFnFilteredByNamespace(false), + nil, + config.DCRedirectionPolicy{Policy: interceptor.DCRedirectionPolicyAllAPIsForwarding}, + log.NewNoopLogger(), + nil, + metrics.NoopMetricsHandler, + clock.NewRealTimeSource(), + metadata, + ), + serviceConfig: &Config{ + EnableNamespaceNotActiveAutoForwarding: dynamicconfig.GetBoolPropertyFnFilteredByNamespace(tc.forwardingOn), + NexusForwardRequestUseEndpoint: dynamicconfig.GetBoolPropertyFn(false), + }, + } + in := interceptornexus.NewStartOpInput( + "s", "o", testNamespace, time.Now(), options, requestInput, + forwardingInfo, + interceptornexus.RequestMetadata{NamespaceEntry: tc.namespace}, + ) + ctx := context.Background() + nextCalled := false + result, err := forwarder.InterceptNexus( + ctx, + in, + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + nextCalled = true + return requestHandledLocally, nil + }, + ) + if tc.expectedOutcome != "" { + var interceptorErr *interceptornexus.InterceptorError + require.ErrorAs(t, err, &interceptorErr) + require.Equal(t, tc.expectedOutcome, interceptorErr.Outcome) + require.True(t, interceptorErr.SkipServiceErrorReporting) + } else { + require.NoError(t, err) + } + expectedNextCalled := tc.disposition == requestHandledLocally + require.Equal(t, expectedNextCalled, nextCalled) + switch tc.disposition { + case requestHandledLocally: + require.Equal(t, requestHandledLocally, result) + case requestForwarded: + require.IsType(t, &nexus.HandlerStartOperationResultAsync{}, result) + require.Equal(t, "true", receivedHeaders.Get(interceptor.DCRedirectionAPIHeaderName)) + require.Equal(t, currentCluster, receivedHeaders.Get(interceptor.DCRedirectionSourceCellHeaderName)) + require.Equal(t, "original", receivedHeaders.Get("X-Original")) + case requestFailed: + require.Nil(t, result) + default: + t.Fatal("unexpected disposition") + } + }) + } +} + +type testFrontendHTTPClientCache struct { + clients map[string]*common.FrontendHTTPClient +} + +func (c testFrontendHTTPClientCache) Get(clusterName string) (*common.FrontendHTTPClient, error) { + client, ok := c.clients[clusterName] + if !ok { + return nil, errors.New("unknown cluster") + } + return client, nil +} diff --git a/service/frontend/nexus_handler.go b/service/frontend/nexus_handler.go index 356014f64f7..14d6c52d483 100644 --- a/service/frontend/nexus_handler.go +++ b/service/frontend/nexus_handler.go @@ -5,11 +5,9 @@ import ( "errors" "fmt" "net/http" - "net/http/httptrace" "net/url" "regexp" "runtime/debug" - "strconv" "strings" "sync" "time" @@ -22,9 +20,6 @@ import ( taskqueuepb "go.temporal.io/api/taskqueue/v1" "go.temporal.io/server/api/matchingservice/v1" chasmnexus "go.temporal.io/server/chasm/lib/nexusoperation" - "go.temporal.io/server/common" - "go.temporal.io/server/common/authorization" - "go.temporal.io/server/common/cluster" "go.temporal.io/server/common/dynamicconfig" "go.temporal.io/server/common/headers" "go.temporal.io/server/common/log" @@ -34,6 +29,7 @@ import ( commonnexus "go.temporal.io/server/common/nexus" "go.temporal.io/server/common/nexus/nexusrpc" "go.temporal.io/server/common/rpc/interceptor" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" "google.golang.org/grpc/metadata" "google.golang.org/protobuf/types/known/timestamppb" ) @@ -49,87 +45,28 @@ const ( type nexusContext struct { // Whether to use the new Temporal failure responses path. // Set from the incoming nexus request's "temporal-nexus-failure-support" header. - callerFailureSupport bool - requestStartTime time.Time - apiName string - namespaceName string - taskQueue string - endpointName string - endpointID string - claims *authorization.Claims - namespaceValidationInterceptor *interceptor.NamespaceValidatorInterceptor - namespaceRateLimitInterceptor interceptor.NamespaceRateLimitInterceptor - namespaceConcurrencyLimitInterceptor *interceptor.ConcurrentRequestLimitInterceptor - rateLimitInterceptor *interceptor.RateLimitInterceptor - responseHeaders map[string]string - responseHeadersMutex sync.Mutex - originalRequestHeaders http.Header // Original HTTP request headers to be used for forwarded requests. + callerFailureSupport bool + requestStartTime time.Time + apiName string + namespaceName string + taskQueue string + endpointName string + endpointID string + responseHeaders map[string]string + responseHeadersMutex sync.Mutex + originalRequestHeaders http.Header // Original HTTP request headers to be used for forwarded requests. } // Context for a specific Nexus operation, includes a resolved namespace, and a bound metrics handler and logger. type operationContext struct { *nexusContext - method string - clusterMetadata cluster.Metadata - namespace *namespace.Namespace + method string + namespace *namespace.Namespace // "Special" metrics handler that should only be passed to interceptors, which require a different set of // pre-baked tags than the "normal" metricsHandler. metricsHandlerForInterceptors metrics.Handler - metricsHandler metrics.Handler logger log.Logger - clientVersionChecker headers.VersionChecker - auth *authorization.Interceptor - telemetryInterceptor *interceptor.TelemetryInterceptor requestErrorHandler *interceptor.RequestErrorHandler - redirectionInterceptor *interceptor.Redirection - forwardingEnabledForNamespace dynamicconfig.BoolPropertyFnWithNamespaceFilter - headersBlacklist dynamicconfig.TypedPropertyFn[*regexp.Regexp] - metricTagConfig dynamicconfig.TypedPropertyFn[chasmnexus.NexusMetricTagConfig] - cleanupFunctions []func(map[string]string, error) -} - -func (c *operationContext) annotateServerSpan( - ctx context.Context, - service, operation, requestID string, -) { - nexusrpc.AnnotateServerSpan(trace.SpanFromContext(ctx), nexusrpc.ServerSpanAttributes{ - Endpoint: c.endpointName, - Service: service, - Operation: operation, - RequestID: requestID, - }) -} - -// Panic handler and metrics recording function. -// Used as a deferred statement in Nexus handler methods. -func (c *operationContext) capturePanicAndRecordMetrics(ctxPtr *context.Context, errPtr *error) { - recovered := recover() //nolint:revive - if recovered != nil { - err, ok := recovered.(error) - if !ok { - err = fmt.Errorf("panic: %v", recovered) - } - - st := string(debug.Stack()) - - c.logger.Error("Panic captured", tag.SysStackTrace(st), tag.Error(err)) - *errPtr = err - } - - // Record Nexus-specific metrics - metrics.NexusRequests.With(c.metricsHandler).Record(1) - metrics.NexusLatency.With(c.metricsHandler).Record(time.Since(c.requestStartTime)) - if *errPtr != nil { - metrics.NexusRequestErrors.With(c.metricsHandler).Record(1) - } - - // Record general telemetry metrics - metrics.ServiceRequests.With(c.metricsHandlerForInterceptors).Record(1) - c.telemetryInterceptor.RecordLatencyMetrics(*ctxPtr, c.requestStartTime, c.metricsHandlerForInterceptors) - - for _, fn := range c.cleanupFunctions { - fn(c.responseHeaders, *errPtr) - } } func (c *operationContext) matchingRequest(req *nexuspb.Request) *matchingservice.DispatchNexusTaskRequest { @@ -141,14 +78,16 @@ func (c *operationContext) matchingRequest(req *nexuspb.Request) *matchingservic } } +func (c *operationContext) annotateServerSpan(ctx context.Context, service, operation, requestID string) { + nexusrpc.AnnotateServerSpan(trace.SpanFromContext(ctx), nexusrpc.ServerSpanAttributes{ + Endpoint: c.endpointName, + Service: service, + Operation: operation, + RequestID: requestID, + }) +} + func (c *operationContext) augmentContext(ctx context.Context, header nexus.Header) context.Context { - ctx = metrics.AddMetricsContext(ctx) - ctx = interceptor.AddTelemetryContext(ctx, c.metricsHandlerForInterceptors) - ctx = interceptor.PopulateCallerInfo( - ctx, - func() string { return c.namespaceName }, - func() string { return c.method }, - ) if userAgent, ok := header[headerUserAgent]; ok { // Use SplitN for efficiency but enforce exactly one delimiter to preserve the // original (pre-SplitN) strictness where additional delimiters cause us to ignore @@ -166,163 +105,78 @@ func (c *operationContext) augmentContext(ctx context.Context, header nexus.Head } } } - return headers.Propagate(ctx) + return ctx } -func (c *operationContext) interceptRequest( - ctx context.Context, - request *matchingservice.DispatchNexusTaskRequest, - header nexus.Header, -) error { - _, err := c.auth.Authorize(ctx, c.claims, &authorization.CallTarget{ - APIName: c.apiName, - Namespace: c.namespaceName, - NexusEndpointName: c.endpointName, - Request: request, - }) - if err != nil { - // If frontend.exposeAuthorizerErrors is false, Authorize err is either an explicitly set reason, or a generic - // "Request unauthorized." message. - // Otherwise, expose the underlying error. - if permissionDeniedError, ok := errors.AsType[*serviceerror.PermissionDenied](err); ok { - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("unauthorized")) - return commonnexus.AdaptAuthorizeError(permissionDeniedError) - } - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("internal_auth_error")) - c.logger.Error("Authorization internal error with processing nexus request", tag.Error(err)) - return commonnexus.ConvertGRPCError(err, false) - } - - // Nexus requests are not tied to a business ID, hence the empty string. - if err := c.namespaceValidationInterceptor.ValidateState(c.namespace, c.apiName, namespace.EmptyBusinessID); err != nil { - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("invalid_namespace_state")) - return commonnexus.ConvertGRPCError(err, false) - } - - //nolint:forbidigo // Nexus requests are not tied to a business ID by design (see line 184) - if !c.namespace.ActiveInCluster(c.clusterMetadata.GetCurrentClusterName()) { - if c.shouldForwardRequest(ctx, header) { - // Handler methods should have special logic to forward requests if this method returns - // a serviceerror.NamespaceNotActive error. - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("request_forwarded")) - handler, forwardStartTime := c.redirectionInterceptor.BeforeCall(c.apiName) - c.cleanupFunctions = append(c.cleanupFunctions, func(_ map[string]string, retErr error) { - c.redirectionInterceptor.AfterCall(handler, forwardStartTime, c.namespace.ActiveClusterName(namespace.RoutingKey{}), c.namespace.Name().String(), retErr) - }) - return serviceerror.NewNamespaceNotActive( - c.namespaceName, - c.clusterMetadata.GetCurrentClusterName(), - c.namespace.ActiveClusterName(namespace.RoutingKey{}), - ) - } - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("namespace_inactive_forwarding_disabled")) - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnavailable, "cluster inactive") - } - - c.cleanupFunctions = append(c.cleanupFunctions, func(respHeaders map[string]string, retErr error) { - if retErr != nil { - if source, ok := respHeaders[commonnexus.FailureSourceHeaderName]; ok && source != commonnexus.FailureSourceWorker { - c.requestErrorHandler.HandleError( - request, - "", - c.metricsHandlerForInterceptors, - []tag.Tag{tag.Operation(c.method), tag.WorkflowNamespace(c.namespaceName)}, - retErr, - c.namespace.Name(), - ) - } +func (c *operationContext) handleRequestError(err error) { + if err == nil { + return + } + if taggedErr, ok := errors.AsType[*interceptornexus.InterceptorError](err); ok { + if taggedErr.SkipServiceErrorReporting { + return } - }) - - cleanup, err := c.namespaceConcurrencyLimitInterceptor.Allow( - c.namespace.Name(), - c.apiName, - c.metricsHandlerForInterceptors, - request, - ) - c.cleanupFunctions = append(c.cleanupFunctions, func(map[string]string, error) { cleanup() }) - if err != nil { - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("namespace_concurrency_limited")) - return commonnexus.ConvertGRPCError(err, false) + err = taggedErr.Err } - - if err := c.namespaceRateLimitInterceptor.Allow( - ctx, - c.namespace.Name(), - c.apiName, - header, - ); err != nil { - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("namespace_rate_limited")) - return commonnexus.ConvertGRPCError(err, true) - } - - if err := c.rateLimitInterceptor.Allow(c.apiName, header); err != nil { - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("global_rate_limited")) - return commonnexus.ConvertGRPCError(err, true) + source, ok := c.responseHeaders[commonnexus.FailureSourceHeaderName] + if !ok || source == commonnexus.FailureSourceWorker { + return } + c.requestErrorHandler.HandleError( + nil, + "", + c.metricsHandlerForInterceptors, + []tag.Tag{tag.Operation(c.method), tag.WorkflowNamespace(c.namespaceName)}, + err, + c.namespace.Name(), + ) +} - if err := c.clientVersionChecker.ClientSupported(ctx); err != nil { - c.metricsHandler = c.metricsHandler.WithTags(metrics.OutcomeTag("unsupported_client")) - converted := commonnexus.ConvertGRPCError(err, true) - return converted +// convertInterceptorError converts the error returned by the interceptor chain into the sanitized +// form returned to the Nexus caller, hiding internal error detail. Interceptors intentionally leave +// InterceptorError.Err raw so the boundary can log/classify the full original error via +// [*operationContext.handleRequestError] before this runs and replaces it for the response. +func convertInterceptorError(err error) error { + if err == nil { + return nil } - - // THIS MUST BE THE LAST STEP IN interceptRequest. - // Sanitize headers. - if request.GetRequest().GetHeader() != nil { - // Making a copy to ensure the original map is not modified as it might be used somewhere else. - sanitizedHeaders := make(map[string]string, len(request.Request.Header)) - headersBlacklist := c.headersBlacklist() - for name, value := range request.Request.Header { - if !headersBlacklist.MatchString(name) { - sanitizedHeaders[name] = value - } - } - request.Request.Header = sanitizedHeaders + exposeDetails := false + if taggedErr, ok := errors.AsType[*interceptornexus.InterceptorError](err); ok { + err = taggedErr.Err + exposeDetails = taggedErr.ExposeDetails } - - // DO NOT ADD ANY STEPS HERE. ALL STEPS MUST BE BEFORE HEADERS SANITIZATION. - - return nil + return commonnexus.ConvertGRPCError(err, exposeDetails) } -// Combines logic from RedirectionInterceptor.redirectionAllowed and some from -// SelectedAPIsForwardingRedirectionPolicy.getTargetClusterAndIsNamespaceNotActiveAutoForwarding so all -// redirection conditions can be checked at once. If either of those methods are updated, this should -// be kept in sync. -func (c *operationContext) shouldForwardRequest(ctx context.Context, header nexus.Header) bool { - redirectHeader := header.Get(interceptor.DCRedirectionContextHeaderName) - redirectAllowed, err := strconv.ParseBool(redirectHeader) - if err != nil { - redirectAllowed = true +// finalizeOperationRequest is the single deferred step for a Nexus start/cancel operation: capture +// a panic into errPtr, log/classify the (still raw) resulting error, then sanitize it for the +// response. Order matters and must not be split back into separate defers. +func finalizeOperationRequest(oc *operationContext, errPtr *error) { + if recovered := recover(); recovered != nil { //nolint:revive + err, ok := recovered.(error) + if !ok { + err = fmt.Errorf("panic: %v", recovered) + } + oc.logger.Error("Panic captured", tag.SysStackTrace(string(debug.Stack())), tag.Error(err)) + *errPtr = err } - return redirectAllowed && - c.redirectionInterceptor.RedirectionAllowed(ctx) && - c.namespace.IsGlobalNamespace() && - c.forwardingEnabledForNamespace(c.namespaceName) + oc.handleRequestError(*errPtr) + *errPtr = convertInterceptorError(*errPtr) } -// enrichNexusOperationMetrics enhances metrics with additional Nexus operation context based on configuration. -func (c *operationContext) enrichNexusOperationMetrics(service, operation string, requestHeader nexus.Header) { - conf := c.metricTagConfig() - - var tags []metrics.Tag - - if conf.IncludeServiceTag { - tags = append(tags, metrics.NexusServiceTag(service)) - } - - if conf.IncludeOperationTag { - tags = append(tags, metrics.NexusOperationTag(operation)) - } - - for _, mapping := range conf.HeaderTagMappings { - tags = append(tags, metrics.StringTag(mapping.TargetTag, requestHeader.Get(mapping.SourceHeader))) +func (h *nexusHandler) sanitizeRequestHeaders(request *matchingservice.DispatchNexusTaskRequest) { + if request.GetRequest().GetHeader() == nil { + return } - if len(tags) > 0 { - c.metricsHandler = c.metricsHandler.WithTags(tags...) + sanitizedHeaders := make(map[string]string, len(request.Request.Header)) + headersBlacklist := h.headersBlacklist() + for name, value := range request.Request.Header { + if !headersBlacklist.MatchString(name) { + sanitizedHeaders[name] = value + } } + request.Request.Header = sanitizedHeaders } // enrichNexusOperationLogs adds Nexus operation context to the handler-side logger. @@ -341,26 +195,76 @@ func (c *operationContext) enrichNexusOperationLogs(service, operation, requestI // Key to extract a nexusContext object from a context.Context. type nexusContextKey struct{} +type operationContextKey struct{} + +func withOperationContext(ctx context.Context, oc *operationContext) context.Context { + if oc == nil { + return ctx + } + return context.WithValue(ctx, operationContextKey{}, oc) +} + +func operationContextFromContext(ctx context.Context) (*operationContext, bool) { + oc, ok := ctx.Value(operationContextKey{}).(*operationContext) + return oc, ok +} + // A Nexus Handler implementation. // Dispatches Nexus requests as Nexus tasks to workers via matching. type nexusHandler struct { nexus.UnimplementedHandler - logger log.Logger - metricsHandler metrics.Handler - clusterMetadata cluster.Metadata - namespaceRegistry namespace.Registry - matchingClient matchingservice.MatchingServiceClient - auth *authorization.Interceptor - telemetryInterceptor *interceptor.TelemetryInterceptor - requestErrorHandler *interceptor.RequestErrorHandler - redirectionInterceptor *interceptor.Redirection - forwardingEnabledForNamespace dynamicconfig.BoolPropertyFnWithNamespaceFilter - forwardingClients *cluster.FrontendHTTPClientCache - payloadSizeLimit dynamicconfig.IntPropertyFnWithNamespaceFilter - headersBlacklist dynamicconfig.TypedPropertyFn[*regexp.Regexp] - useForwardByEndpoint dynamicconfig.BoolPropertyFn - metricTagConfig dynamicconfig.TypedPropertyFn[chasmnexus.NexusMetricTagConfig] - httpTraceProvider commonnexus.HTTPClientTraceProvider + logger log.Logger + metricsHandler metrics.Handler + namespaceRegistry namespace.Registry + matchingClient matchingservice.MatchingServiceClient + requestErrorHandler *interceptor.RequestErrorHandler + payloadSizeLimit dynamicconfig.IntPropertyFnWithNamespaceFilter + headersBlacklist dynamicconfig.TypedPropertyFn[*regexp.Regexp] + metricTagConfig dynamicconfig.TypedPropertyFn[chasmnexus.NexusMetricTagConfig] + chainedHandler interceptornexus.HandlerFunc +} + +func newNexusHandler( + logger log.Logger, + metricsHandler metrics.Handler, + namespaceRegistry namespace.Registry, + matchingClient matchingservice.MatchingServiceClient, + requestErrorHandler *interceptor.RequestErrorHandler, + payloadSizeLimit dynamicconfig.IntPropertyFnWithNamespaceFilter, + headersBlacklist dynamicconfig.TypedPropertyFn[*regexp.Regexp], + metricTagConfig dynamicconfig.TypedPropertyFn[chasmnexus.NexusMetricTagConfig], + nexusInterceptors []interceptornexus.Interceptor, +) *nexusHandler { + h := &nexusHandler{ + logger: logger, + metricsHandler: metricsHandler, + namespaceRegistry: namespaceRegistry, + matchingClient: matchingClient, + requestErrorHandler: requestErrorHandler, + payloadSizeLimit: payloadSizeLimit, + headersBlacklist: headersBlacklist, + metricTagConfig: metricTagConfig, + } + h.chainedHandler = interceptornexus.ChainInterceptors(h.finalHandler, nexusInterceptors) + return h +} + +// nexusMetricTags resolves the operator-configurable tags for this request's Nexus metrics. Only the +// frontend can read the configuration, so the tags travel to the telemetry interceptor as request +// metadata rather than being built where they are recorded. +func (h *nexusHandler) nexusMetricTags(service, operation string, header nexus.Header) []metrics.Tag { + conf := h.metricTagConfig() + var tags []metrics.Tag + if conf.IncludeServiceTag { + tags = append(tags, metrics.NexusServiceTag(service)) + } + if conf.IncludeOperationTag { + tags = append(tags, metrics.NexusOperationTag(operation)) + } + for _, mapping := range conf.HeaderTagMappings { + tags = append(tags, metrics.StringTag(mapping.TargetTag, header.Get(mapping.SourceHeader))) + } + return tags } // Extracts a nexusContext from the given ctx and returns an operationContext with tagged metrics and logging. @@ -371,35 +275,23 @@ func (h *nexusHandler) getOperationContext(ctx context.Context, method string) ( return nil, errors.New("no nexus context set on context") } oc := operationContext{ - nexusContext: nc, - method: method, - clusterMetadata: h.clusterMetadata, - clientVersionChecker: headers.NewDefaultVersionChecker(), - auth: h.auth, - telemetryInterceptor: h.telemetryInterceptor, - requestErrorHandler: h.requestErrorHandler, - redirectionInterceptor: h.redirectionInterceptor, - forwardingEnabledForNamespace: h.forwardingEnabledForNamespace, - headersBlacklist: h.headersBlacklist, - metricTagConfig: h.metricTagConfig, - cleanupFunctions: make([]func(map[string]string, error), 0), + nexusContext: nc, + method: method, + requestErrorHandler: h.requestErrorHandler, } oc.metricsHandlerForInterceptors = h.metricsHandler.WithTags( metrics.OperationTag(method), metrics.NamespaceTag(nc.namespaceName), ) - oc.metricsHandler = h.metricsHandler.WithTags( - metrics.NamespaceTag(nc.namespaceName), - metrics.NexusEndpointTag(nc.endpointName), - metrics.NexusMethodTag(method), - // default to internal error unless overridden by handler - metrics.OutcomeTag("internal_error"), - ) var err error if oc.namespace, err = h.namespaceRegistry.GetNamespace(namespace.Name(nc.namespaceName)); err != nil { - metrics.NexusRequests.With(oc.metricsHandler).Record( + // namespace lookup runs before the interceptor chain, so this outcome is recorded here. + metrics.NexusRequests.With(h.metricsHandler).Record( 1, + metrics.NamespaceTag(nc.namespaceName), + metrics.NexusEndpointTag(nc.endpointName), + metrics.NexusMethodTag(method), metrics.OutcomeTag("namespace_not_found"), ) @@ -408,7 +300,6 @@ func (h *nexusHandler) getOperationContext(ctx context.Context, method string) ( } return nil, commonnexus.ConvertGRPCError(err, false) } - oc.forwardingEnabledForNamespace = h.forwardingEnabledForNamespace oc.logger = log.With(h.logger, tag.Operation(method), tag.WorkflowNamespace(nc.namespaceName)) return &oc, nil } @@ -425,11 +316,12 @@ func (h *nexusHandler) StartOperation( return nil, err } ctx = oc.augmentContext(ctx, options.Header) - oc.enrichNexusOperationMetrics(service, operation, options.Header) oc.enrichNexusOperationLogs(service, operation, options.RequestID) oc.annotateServerSpan(ctx, service, operation, options.RequestID) - defer oc.capturePanicAndRecordMetrics(&ctx, &retErr) + // to handle edge case where the operation panics before the interceptor chain is invoked + defer finalizeOperationRequest(oc, &retErr) + ctx = withOperationContext(ctx, oc) var links []*nexuspb.Link for _, nexusLink := range options.Links { links = append(links, &nexuspb.Link{ @@ -437,35 +329,79 @@ func (h *nexusHandler) StartOperation( Type: nexusLink.Type, }) } - - startOperationRequest := nexuspb.StartOperationRequest{ - Service: service, - Operation: operation, - Callback: options.CallbackURL, - CallbackHeader: options.CallbackHeader, - RequestId: options.RequestID, - Links: links, - } request := oc.matchingRequest(&nexuspb.Request{ ScheduledTime: timestamppb.New(oc.requestStartTime), Header: options.Header, Variant: &nexuspb.Request_StartOperation{ - StartOperation: &startOperationRequest, + StartOperation: &nexuspb.StartOperationRequest{ + Service: service, + Operation: operation, + Callback: options.CallbackURL, + CallbackHeader: options.CallbackHeader, + RequestId: options.RequestID, + Links: links, + }, }, Capabilities: &nexuspb.Request_Capabilities{ TemporalFailureResponses: oc.callerFailureSupport, }, }) - if err := oc.interceptRequest(ctx, request, options.Header); err != nil { - if _, ok := errors.AsType[*serviceerror.NamespaceNotActive](err); ok { - return h.forwardStartOperation(ctx, service, operation, input, options, oc) - } + nexusOpInput := interceptornexus.NewStartOpInput( + service, + operation, + oc.namespaceName, + oc.requestStartTime, + options, + input, + interceptornexus.ForwardingInfo{ + OriginalRequestHeaders: oc.originalRequestHeaders, + TaskQueue: oc.taskQueue, + EndpointID: oc.endpointID, + EndpointName: oc.endpointName, + }, + interceptornexus.RequestMetadata{ + APIName: oc.apiName, + NamespaceEntry: oc.namespace, + EndpointName: oc.endpointName, + MetricTags: h.nexusMetricTags(service, operation, options.Header), + Request: request, + }, + ) + out, err := h.chainedHandler(ctx, nexusOpInput) + if err != nil { return nil, err } + res, ok := out.(nexus.HandlerStartOperationResult[any]) + if !ok { + return nil, fmt.Errorf("unexpected Nexus start interceptor result type %T", out) + } + return res, nil +} +//nolint:revive,cognitive-complexity: justified to keep the flow intact +func (h *nexusHandler) finalStartHandler( + ctx context.Context, + in interceptornexus.InterceptorInput, +) (any, error) { + oc, ocok := operationContextFromContext(ctx) + if !ocok { + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "invalid operation context for nexus start operation") + } + operation := in.OperationName() + soi, ok := in.(interceptornexus.StartOpInput) + if !ok { + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "invalid request for nexus start operation") + } + request, ok := soi.Request().(*matchingservice.DispatchNexusTaskRequest) + if !ok || request.GetRequest().GetStartOperation() == nil { + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "invalid dispatch request for nexus start operation") + } + startOperationRequest := request.GetRequest().GetStartOperation() + h.sanitizeRequestHeaders(request) + var err error // Transform nexus Content to temporal Payload with common/nexus PayloadSerializer. - if err = input.Consume(&startOperationRequest.Payload); err != nil { + if err = soi.StartOperationInput.Consume(&startOperationRequest.Payload); err != nil { oc.logger.Warn("invalid input", tag.Error(err)) return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid input") } @@ -479,9 +415,11 @@ func (h *nexusHandler) StartOperation( // RPC. response, err := h.matchingClient.DispatchNexusTask(ctx, request) if err != nil { - oc.metricsHandler = oc.metricsHandler.WithTags(metrics.OutcomeTag("matching_timeout")) oc.logger.Error("received error from matching service for Nexus StartOperation request", tag.Error(err)) - return nil, commonnexus.ConvertGRPCError(err, false) + return nil, &interceptornexus.InterceptorError{ + Err: err, + Outcome: "matching_timeout", + } } // Convert to standard Nexus SDK response. result, handlerLinks, err := oc.handleStartOperationResponse(response, operation) @@ -510,65 +448,16 @@ func parseLinks(links []*nexuspb.Link, logger log.Logger) []nexus.Link { return nexusLinks } -// forwardStartOperation forwards the StartOperation request to the active cluster using an HTTP request. -// Inputs and response values are passed as Reader objects to avoid reading bodies and bypass serialization. -func (h *nexusHandler) forwardStartOperation( - ctx context.Context, - service string, - operation string, - input *nexus.LazyValue, - options nexus.StartOperationOptions, - oc *operationContext, -) (nexus.HandlerStartOperationResult[any], error) { - options.Header[interceptor.DCRedirectionAPIHeaderName] = "true" - options.Header[interceptor.DCRedirectionSourceCellHeaderName] = h.clusterMetadata.GetCurrentClusterName() - - client, err := h.nexusClientForActiveCluster(oc, service) - if err != nil { - return nil, err - } - - if h.httpTraceProvider != nil { - traceLogger := log.With(h.logger, - tag.Operation(oc.method), - tag.WorkflowNamespace(oc.namespaceName), - tag.RequestID(options.RequestID), - tag.NexusOperation(operation), - tag.Endpoint(oc.endpointName), - tag.AttemptStart(time.Now().UTC()), - tag.SourceCluster(h.clusterMetadata.GetCurrentClusterName()), - tag.TargetCluster(oc.namespace.ActiveClusterName(namespace.RoutingKey{})), - ) - if trace := h.httpTraceProvider.NewForwardingTrace(traceLogger); trace != nil { - ctx = httptrace.WithClientTrace(ctx, trace) - } - } - - resp, err := client.StartOperation(ctx, operation, input.Reader, options) - if err != nil { - oc.logger.Error("received error from remote cluster for forwarded Nexus start operation request.", tag.Error(err)) - oc.metricsHandler = oc.metricsHandler.WithTags(metrics.OutcomeTag("forwarded_request_error")) - return nil, err - } - - if resp.Successful != nil { - return &nexus.HandlerStartOperationResultSync[any]{Value: resp.Successful.Reader}, nil - } - // If Nexus client did not return an error, one of Successful or Pending will always be set. - return &nexus.HandlerStartOperationResultAsync{OperationToken: resp.Pending.Token}, nil -} - func (h *nexusHandler) CancelOperation(ctx context.Context, service, operation, token string, options nexus.CancelOperationOptions) (retErr error) { oc, err := h.getOperationContext(ctx, "CancelNexusOperation") if err != nil { return err } ctx = oc.augmentContext(ctx, options.Header) - oc.enrichNexusOperationMetrics(service, operation, options.Header) oc.enrichNexusOperationLogs(service, operation, "") oc.annotateServerSpan(ctx, service, operation, "") - defer oc.capturePanicAndRecordMetrics(&ctx, &retErr) - + // for edge case where the operation panics before the interceptor chain is invoked + defer finalizeOperationRequest(oc, &retErr) request := oc.matchingRequest(&nexuspb.Request{ Header: options.Header, ScheduledTime: timestamppb.New(oc.requestStartTime), @@ -585,118 +474,79 @@ func (h *nexusHandler) CancelOperation(ctx context.Context, service, operation, TemporalFailureResponses: oc.callerFailureSupport, }, }) - if err := oc.interceptRequest(ctx, request, options.Header); err != nil { - if _, ok := errors.AsType[*serviceerror.NamespaceNotActive](err); ok { - return h.forwardCancelOperation(ctx, service, operation, token, options, oc) - } - return err - } - // Dispatch the request to be sync matched with a worker polling on the nexusContext taskQueue. - // matchingClient sets a context timeout of 60 seconds for this request, this should be enough for any Nexus - // RPC. - response, err := h.matchingClient.DispatchNexusTask(ctx, request) - if err != nil { - oc.metricsHandler = oc.metricsHandler.WithTags(metrics.OutcomeTag("matching_timeout")) - oc.logger.Error("received error from matching service for Nexus CancelOperation request", tag.Error(err)) - return commonnexus.ConvertGRPCError(err, false) - } - // Convert to standard Nexus SDK response. - return oc.handleCancelOperationResponse(response, operation) + nexusInterceptorInput := interceptornexus.NewCancelOpInput( + service, + operation, + oc.namespaceName, + oc.requestStartTime, + options, + token, + interceptornexus.ForwardingInfo{ + OriginalRequestHeaders: oc.originalRequestHeaders, + TaskQueue: oc.taskQueue, + EndpointID: oc.endpointID, + EndpointName: oc.endpointName, + }, + interceptornexus.RequestMetadata{ + APIName: oc.apiName, + NamespaceEntry: oc.namespace, + EndpointName: oc.endpointName, + MetricTags: h.nexusMetricTags(service, operation, options.Header), + Request: request, + }, + ) + ctx = withOperationContext(ctx, oc) + _, err = h.chainedHandler(ctx, nexusInterceptorInput) + return err } -func (h *nexusHandler) forwardCancelOperation( +func (h *nexusHandler) finalHandler( ctx context.Context, - service string, - operation string, - id string, - options nexus.CancelOperationOptions, - oc *operationContext, -) error { - options.Header[interceptor.DCRedirectionAPIHeaderName] = "true" - options.Header[interceptor.DCRedirectionSourceCellHeaderName] = h.clusterMetadata.GetCurrentClusterName() - - client, err := h.nexusClientForActiveCluster(oc, service) - if err != nil { - return err - } - - handle, err := client.NewOperationHandle(operation, id) - if err != nil { - oc.logger.Warn("invalid Nexus cancel operation.", tag.Error(err)) - return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid operation") - } - - if h.httpTraceProvider != nil { - traceLogger := log.With(h.logger, - tag.Operation(oc.method), - tag.WorkflowNamespace(oc.namespaceName), - tag.NexusOperation(operation), - tag.Endpoint(oc.endpointName), - tag.AttemptStart(time.Now().UTC()), - tag.SourceCluster(h.clusterMetadata.GetCurrentClusterName()), - tag.TargetCluster(oc.namespace.ActiveClusterName(namespace.RoutingKey{})), - ) - if trace := h.httpTraceProvider.NewForwardingTrace(traceLogger); trace != nil { - ctx = httptrace.WithClientTrace(ctx, trace) - } - } - - err = handle.Cancel(ctx, options) - if err != nil { - oc.logger.Error("received error from remote cluster for forwarded Nexus cancel operation request.", tag.Error(err)) - oc.metricsHandler = oc.metricsHandler.WithTags(metrics.OutcomeTag("forwarded_request_error")) - return err + in interceptornexus.InterceptorInput, +) (any, error) { + switch in.(type) { + case interceptornexus.StartOpInput: + return h.finalStartHandler(ctx, in) + case interceptornexus.CancelOpInput: + return h.finalCancelHandler(ctx, in) + default: + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "unknown operation triggered, expected start/cancel nexus op") } - - return nil } -func (h *nexusHandler) nexusClientForActiveCluster(oc *operationContext, service string) (*nexusrpc.HTTPClient, error) { - httpClient, err := h.forwardingClients.Get(oc.namespace.ActiveClusterName(namespace.RoutingKey{})) - if err != nil { - oc.logger.Error("failed to forward Nexus request. error creating HTTP client", tag.Error(err), tag.SourceCluster(oc.namespace.ActiveClusterName(namespace.RoutingKey{})), tag.TargetCluster(oc.namespace.ActiveClusterName(namespace.RoutingKey{}))) - oc.metricsHandler = oc.metricsHandler.WithTags(metrics.OutcomeTag("request_forwarding_failed")) - return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "request forwarding failed") +func (h *nexusHandler) finalCancelHandler( + ctx context.Context, + in interceptornexus.InterceptorInput, +) (any, error) { + oc, ocok := operationContextFromContext(ctx) + if !ocok { + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "invalid operation context for nexus cancel operation") } - - httpCaller := &forwardingHttpHeaderWrapper{ - client: httpClient, - nc: oc.nexusContext, - originalRequestHeaders: oc.originalRequestHeaders, + coi, ok := in.(interceptornexus.CancelOpInput) + if !ok { + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "invalid request for nexus cancel operation") } - - var baseURL string - if h.useForwardByEndpoint() && oc.endpointID != "" { - // If the request was originally dispatched by endpoint, forward by endpoint as well. - baseURL, err = url.JoinPath(httpClient.BaseURL(), - commonnexus.RouteDispatchNexusTaskByEndpoint.Path(oc.endpointID)) - } else { - // Fallback to dispatch by namespace and task queue since those have already been resolved by this point. - // NOTE: When forwarding by namespace and task queue, the endpoint name is not preserved and cannot be provided to a worker polling. - baseURL, err = url.JoinPath( - httpClient.BaseURL(), - commonnexus.RouteDispatchNexusTaskByNamespaceAndTaskQueue.Path(commonnexus.NamespaceAndTaskQueue{ - Namespace: oc.namespaceName, - TaskQueue: oc.taskQueue, - })) + operation := in.OperationName() + request, ok := coi.Request().(*matchingservice.DispatchNexusTaskRequest) + if !ok || request.GetRequest().GetCancelOperation() == nil { + return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "invalid dispatch request for nexus cancel operation") } + h.sanitizeRequestHeaders(request) + // Dispatch the request to be sync matched with a worker polling on the nexusContext taskQueue. + // matchingClient sets a context timeout of 60 seconds for this request, this should be enough for any Nexus + // RPC. + response, err := h.matchingClient.DispatchNexusTask(ctx, request) if err != nil { - oc.logger.Error("failed to forward Nexus request. error constructing ServiceBaseURL", - tag.URL(httpClient.BaseURL()), - tag.WorkflowNamespace(oc.namespaceName), - tag.WorkflowTaskQueueName(oc.taskQueue), - tag.Error(err)) - oc.metricsHandler = oc.metricsHandler.WithTags(metrics.OutcomeTag("request_forwarding_failed")) - return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "request forwarding failed") - } - - return nexusrpc.NewHTTPClient(nexusrpc.HTTPClientOptions{ - HTTPCaller: httpCaller.Do, - BaseURL: baseURL, - Service: service, - }) + oc.logger.Error("received error from matching service for Nexus CancelOperation request", tag.Error(err)) + return nil, &interceptornexus.InterceptorError{ + Err: err, + Outcome: "matching_timeout", + } + } + // Convert to standard Nexus SDK response. + return nil, oc.handleCancelOperationResponse(response, operation) } func (nc *nexusContext) setFailureSource(source string) { @@ -704,29 +554,3 @@ func (nc *nexusContext) setFailureSource(source string) { defer nc.responseHeadersMutex.Unlock() nc.responseHeaders[commonnexus.FailureSourceHeaderName] = source } - -type forwardingHttpHeaderWrapper struct { - client *common.FrontendHTTPClient - nc *nexusContext - originalRequestHeaders http.Header -} - -func (f *forwardingHttpHeaderWrapper) Do(req *http.Request) (*http.Response, error) { - // For forwarded requests, copy the original HTTP headers without sanitization. - for k, v := range f.originalRequestHeaders { - if req.Header.Get(k) == "" { - req.Header.Set(k, v[0]) - } - } - - response, err := f.client.Do(req) - if err != nil { - return nil, err - } - - if failureSource := response.Header.Get(commonnexus.FailureSourceHeaderName); failureSource != "" { - f.nc.setFailureSource(failureSource) - } - - return response, nil -} diff --git a/service/frontend/nexus_handler_test.go b/service/frontend/nexus_handler_test.go index 17891daea3f..15ce9338f45 100644 --- a/service/frontend/nexus_handler_test.go +++ b/service/frontend/nexus_handler_test.go @@ -1,114 +1,38 @@ package frontend import ( - "context" - "errors" - "testing" - "time" - "github.com/google/uuid" - "github.com/nexus-rpc/sdk-go/nexus" - "github.com/stretchr/testify/require" enumspb "go.temporal.io/api/enums/v1" - nexuspb "go.temporal.io/api/nexus/v1" - "go.temporal.io/api/serviceerror" - "go.temporal.io/server/api/matchingservice/v1" persistencespb "go.temporal.io/server/api/persistence/v1" - "go.temporal.io/server/common/authorization" - "go.temporal.io/server/common/clock" "go.temporal.io/server/common/cluster" - "go.temporal.io/server/common/cluster/clustertest" - "go.temporal.io/server/common/config" - "go.temporal.io/server/common/dynamicconfig" - "go.temporal.io/server/common/headers" "go.temporal.io/server/common/log" - "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/metrics/metricstest" "go.temporal.io/server/common/namespace" "go.temporal.io/server/common/primitives/timestamp" - "go.temporal.io/server/common/quotas" - "go.temporal.io/server/common/rpc/interceptor" - "go.temporal.io/server/common/util" ) -type mockAuthorizer struct{} - -// Authorize implements authorization.Authorizer. -func (mockAuthorizer) Authorize(ctx context.Context, caller *authorization.Claims, target *authorization.CallTarget) (authorization.Result, error) { - return authorization.Result{Decision: authorization.DecisionAllow}, nil -} - -var _ authorization.Authorizer = mockAuthorizer{} - -type mockRateLimiter struct { - allow bool -} - -// Allow implements quotas.RequestRateLimiter. -func (r mockRateLimiter) Allow(now time.Time, request quotas.Request) bool { - return r.allow -} - -// Reserve implements quotas.RequestRateLimiter. -func (mockRateLimiter) Reserve(now time.Time, request quotas.Request) quotas.Reservation { - panic("unimplemented for test") -} - -// Wait implements quotas.RequestRateLimiter. -func (mockRateLimiter) Wait(ctx context.Context, request quotas.Request) error { - panic("unimplemented for test") -} - -var _ quotas.RequestRateLimiter = mockRateLimiter{} - -type mockNamespaceChecker namespace.Name - -func (n mockNamespaceChecker) Exists(name namespace.Name) error { - if name == namespace.Name(n) { - return nil - } - return errors.New("doesn't exist") -} - -type contextOptions struct { - namespaceState enumspb.NamespaceState - namespacePassive bool - quota int - namespaceRateLimitAllow bool - rateLimitAllow bool - redirectAllow bool - headersBlacklist []string -} - -func newOperationContext(options contextOptions) *operationContext { +func testOperationContext() *operationContext { oc := &operationContext{ nexusContext: &nexusContext{}, } oc.logger = log.NewTestLogger() - mh := metricstest.NewCaptureHandler() - oc.metricsHandlerForInterceptors = mh - oc.metricsHandler = mh - oc.clientVersionChecker = headers.NewDefaultVersionChecker() + oc.metricsHandlerForInterceptors = metricstest.NewCaptureHandler() oc.apiName = "/temporal.api.nexusservice.v1.NexusService/DispatchNexusTask" oc.responseHeaders = make(map[string]string) oc.namespaceName = "test-namespace" - activeClusterName := cluster.TestCurrentClusterName - if options.namespacePassive { - activeClusterName = cluster.TestAlternativeClusterName - } oc.namespace = namespace.NewGlobalNamespaceForTest( &persistencespb.NamespaceInfo{ Id: uuid.NewString(), Name: oc.namespaceName, - State: options.namespaceState, + State: enumspb.NAMESPACE_STATE_REGISTERED, }, &persistencespb.NamespaceConfig{ Retention: timestamp.DurationFromDays(1), CustomSearchAttributeAliases: make(map[string]string), }, &persistencespb.NamespaceReplicationConfig{ - ActiveClusterName: activeClusterName, + ActiveClusterName: cluster.TestCurrentClusterName, Clusters: []string{ cluster.TestCurrentClusterName, cluster.TestAlternativeClusterName, @@ -117,266 +41,5 @@ func newOperationContext(options contextOptions) *operationContext { 1, ) - checker := mockNamespaceChecker(oc.namespace.Name()) - oc.auth = authorization.NewInterceptor( - nil, - mockAuthorizer{}, - oc.metricsHandler, - oc.logger, - checker, - nil, - "", - "", - dynamicconfig.GetBoolPropertyFn(false), // exposeAuthorizerErrors - dynamicconfig.GetBoolPropertyFn(false), // enableCrossNamespaceCommands - dynamicconfig.GetBoolPropertyFnFilteredByNamespace(false), // enablePrincipalPropagation - dynamicconfig.GetBoolPropertyFn(false), // disableStreamingAuthorizer - ) - oc.namespaceConcurrencyLimitInterceptor = interceptor.NewConcurrentRequestLimitInterceptor( - nil, - nil, - oc.logger, - func(ns string) int { return options.quota }, - func(ns string) int { return options.quota }, - map[string]int{ - oc.apiName: 1, - }, - ) - oc.namespaceRateLimitInterceptor = interceptor.NewNamespaceRateLimitInterceptor( - nil, - mockRateLimiter{options.namespaceRateLimitAllow}, - map[string]struct{}{}, - dynamicconfig.GetBoolPropertyFnFilteredByNamespace(false), - metrics.NoopMetricsHandler, - ) - oc.rateLimitInterceptor = interceptor.NewRateLimitInterceptor( - mockRateLimiter{options.rateLimitAllow}, - make(map[string]int), - ) - - oc.clusterMetadata = clustertest.NewMetadataForTest( - cluster.NewTestClusterMetadataConfig(true, !options.namespacePassive), - ) - oc.forwardingEnabledForNamespace = dynamicconfig.GetBoolPropertyFnFilteredByNamespace( - options.redirectAllow, - ) - re, err := dynamicconfig.ConvertWildcardStringListToRegexp(options.headersBlacklist) - if err != nil { - panic(err) // nolint:forbidigo - } - oc.headersBlacklist = dynamicconfig.GetTypedPropertyFn(re) - oc.redirectionInterceptor = interceptor.NewRedirection( - nil, - dynamicconfig.GetBoolPropertyFnFilteredByNamespace(false), - nil, - config.DCRedirectionPolicy{Policy: interceptor.DCRedirectionPolicyAllAPIsForwarding}, - oc.logger, - nil, - oc.metricsHandlerForInterceptors, - clock.NewRealTimeSource(), - oc.clusterMetadata, - ) - return oc } - -func TestNexusInterceptRequest_InvalidNamespaceState_ResultsInBadRequest(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - var err error - oc := newOperationContext(contextOptions{ - namespaceState: enumspb.NAMESPACE_STATE_DELETED, - quota: 1, - namespaceRateLimitAllow: true, - rateLimitAllow: true, - }) - err = oc.interceptRequest(ctx, &matchingservice.DispatchNexusTaskRequest{}, nexus.Header{}) - var handlerError *nexus.HandlerError - require.ErrorAs(t, err, &handlerError) - require.Equal(t, nexus.HandlerErrorTypeBadRequest, handlerError.Type) - require.Equal(t, "bad request", handlerError.Message) - mh := oc.metricsHandler.(*metricstest.CaptureHandler) //nolint:revive - capture := mh.StartCapture() - oc.metricsHandler.Counter("test").Record(1) - mh.StopCapture(capture) - snap := capture.Snapshot() - require.Len(t, snap["test"], 1) - require.Equal(t, map[string]string{"outcome": "invalid_namespace_state"}, snap["test"][0].Tags) -} - -func TestNexusInterceptRequest_NamespaceConcurrencyLimited_ResultsInResourceExhausted(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - var err error - oc := newOperationContext(contextOptions{ - namespaceState: enumspb.NAMESPACE_STATE_REGISTERED, - quota: 0, - namespaceRateLimitAllow: true, - rateLimitAllow: true, - }) - err = oc.interceptRequest(ctx, &matchingservice.DispatchNexusTaskRequest{}, nexus.Header{}) - var handlerError *nexus.HandlerError - require.ErrorAs(t, err, &handlerError) - require.Equal(t, nexus.HandlerErrorTypeResourceExhausted, handlerError.Type) - require.Equal(t, "resource exhausted", handlerError.Message) - mh := oc.metricsHandler.(*metricstest.CaptureHandler) //nolint:revive - capture := mh.StartCapture() - oc.metricsHandler.Counter("test").Record(1) - mh.StopCapture(capture) - snap := capture.Snapshot() - require.Len(t, snap["test"], 1) - require.Equal(t, map[string]string{"outcome": "namespace_concurrency_limited"}, snap["test"][0].Tags) -} - -func TestNexusInterceptRequest_NamespaceRateLimited_ResultsInResourceExhausted(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - var err error - oc := newOperationContext(contextOptions{ - namespaceState: enumspb.NAMESPACE_STATE_REGISTERED, - quota: 1, - namespaceRateLimitAllow: false, - rateLimitAllow: true, - }) - err = oc.interceptRequest(ctx, &matchingservice.DispatchNexusTaskRequest{}, nexus.Header{}) - var handlerError *nexus.HandlerError - require.ErrorAs(t, err, &handlerError) - require.Equal(t, nexus.HandlerErrorTypeResourceExhausted, handlerError.Type) - require.Equal(t, "namespace rate limit exceeded", handlerError.Message) - mh := oc.metricsHandler.(*metricstest.CaptureHandler) //nolint:revive - capture := mh.StartCapture() - oc.metricsHandler.Counter("test").Record(1) - mh.StopCapture(capture) - snap := capture.Snapshot() - require.Len(t, snap["test"], 1) - require.Equal(t, map[string]string{"outcome": "namespace_rate_limited"}, snap["test"][0].Tags) -} - -func TestNexusInterceptRequest_GlobalRateLimited_ResultsInResourceExhausted(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - var err error - oc := newOperationContext(contextOptions{ - namespaceState: enumspb.NAMESPACE_STATE_REGISTERED, - quota: 1, - namespaceRateLimitAllow: true, - rateLimitAllow: false, - }) - err = oc.interceptRequest(ctx, &matchingservice.DispatchNexusTaskRequest{}, nexus.Header{}) - var handlerError *nexus.HandlerError - require.ErrorAs(t, err, &handlerError) - require.Equal(t, nexus.HandlerErrorTypeResourceExhausted, handlerError.Type) - require.Equal(t, "service rate limit exceeded", handlerError.Message) - mh := oc.metricsHandler.(*metricstest.CaptureHandler) //nolint:revive - capture := mh.StartCapture() - oc.metricsHandler.Counter("test").Record(1) - mh.StopCapture(capture) - snap := capture.Snapshot() - require.Len(t, snap["test"], 1) - require.Equal(t, map[string]string{"outcome": "global_rate_limited"}, snap["test"][0].Tags) -} - -func TestNexusInterceptRequest_ForwardingDisabled_ResultsInUnavailable(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - var err error - oc := newOperationContext(contextOptions{ - namespaceState: enumspb.NAMESPACE_STATE_REGISTERED, - namespacePassive: true, - quota: 1, - namespaceRateLimitAllow: true, - rateLimitAllow: true, - redirectAllow: false, - }) - err = oc.interceptRequest(ctx, &matchingservice.DispatchNexusTaskRequest{}, nexus.Header{}) - var handlerError *nexus.HandlerError - require.ErrorAs(t, err, &handlerError) - require.Equal(t, nexus.HandlerErrorTypeUnavailable, handlerError.Type) - mh := oc.metricsHandler.(*metricstest.CaptureHandler) //nolint:revive - capture := mh.StartCapture() - oc.metricsHandler.Counter("test").Record(1) - mh.StopCapture(capture) - snap := capture.Snapshot() - require.Len(t, snap["test"], 1) - require.Equal(t, map[string]string{"outcome": "namespace_inactive_forwarding_disabled"}, snap["test"][0].Tags) -} - -func TestNexusInterceptRequest_ForwardingEnabled_ResultsInNotActiveError(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - var err error - oc := newOperationContext(contextOptions{ - namespaceState: enumspb.NAMESPACE_STATE_REGISTERED, - namespacePassive: true, - quota: 1, - namespaceRateLimitAllow: true, - rateLimitAllow: true, - redirectAllow: true, - }) - err = oc.interceptRequest(ctx, &matchingservice.DispatchNexusTaskRequest{}, nexus.Header{}) - var notActiveErr *serviceerror.NamespaceNotActive - require.ErrorAs(t, err, ¬ActiveErr) - mh := oc.metricsHandler.(*metricstest.CaptureHandler) //nolint:revive - capture := mh.StartCapture() - oc.metricsHandler.Counter("test").Record(1) - mh.StopCapture(capture) - snap := capture.Snapshot() - require.Len(t, snap["test"], 1) - require.Equal(t, map[string]string{"outcome": "request_forwarded"}, snap["test"][0].Tags) -} - -func TestNexusInterceptRequest_InvalidSDKVersion_ResultsInBadRequest(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - var err error - oc := newOperationContext(contextOptions{ - namespaceState: enumspb.NAMESPACE_STATE_REGISTERED, - namespacePassive: false, - quota: 1, - namespaceRateLimitAllow: true, - rateLimitAllow: true, - redirectAllow: true, - }) - header := nexus.Header{headerUserAgent: "Nexus-go-sdk/v99.0.0"} - ctx = oc.augmentContext(ctx, header) - err = oc.interceptRequest(ctx, &matchingservice.DispatchNexusTaskRequest{}, header) - var handlerError *nexus.HandlerError - require.ErrorAs(t, err, &handlerError) - require.Equal(t, nexus.HandlerErrorTypeBadRequest, handlerError.Type) - mh := oc.metricsHandler.(*metricstest.CaptureHandler) //nolint:revive - capture := mh.StartCapture() - oc.metricsHandler.Counter("test").Record(1) - mh.StopCapture(capture) - snap := capture.Snapshot() - require.Len(t, snap["test"], 1) - require.Equal(t, map[string]string{"outcome": "unsupported_client"}, snap["test"][0].Tags) -} - -func TestNexusInterceptRequest_HeadersSanitization(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), time.Second) - defer cancel() - var err error - oc := newOperationContext(contextOptions{ - namespaceState: enumspb.NAMESPACE_STATE_REGISTERED, - namespacePassive: false, - quota: 1, - namespaceRateLimitAllow: true, - rateLimitAllow: true, - headersBlacklist: []string{"delete-*", "remove-*"}, - }) - initialHeader := nexus.Header{ - "ok-header": "ok", - "delete-foo": "foo", - "delete-bar": "bar", - "remove-zzz": "zzz", - } - header := util.CloneMapNonNil(initialHeader) - ctx = oc.augmentContext(ctx, header) - request := &matchingservice.DispatchNexusTaskRequest{ - Request: &nexuspb.Request{Header: header}, - } - err = oc.interceptRequest(ctx, request, header) - require.NoError(t, err) - require.Equal(t, initialHeader, header) - require.Equal(t, map[string]string{"ok-header": "ok"}, request.Request.Header) -} diff --git a/service/frontend/nexus_interceptor_chain_test.go b/service/frontend/nexus_interceptor_chain_test.go new file mode 100644 index 00000000000..3df3f5ce2d3 --- /dev/null +++ b/service/frontend/nexus_interceptor_chain_test.go @@ -0,0 +1,272 @@ +package frontend + +import ( + "context" + "errors" + "fmt" + "reflect" + "testing" + "time" + + "github.com/nexus-rpc/sdk-go/nexus" + "github.com/stretchr/testify/require" + "go.temporal.io/server/common/dynamicconfig" + "go.temporal.io/server/common/log" + "go.temporal.io/server/common/metrics" + "go.temporal.io/server/common/metrics/metricstest" + rpcinterceptor "go.temporal.io/server/common/rpc/interceptor" + interceptornexus "go.temporal.io/server/common/rpc/interceptor/nexus" + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +type testGRPCError struct { + status *status.Status +} + +func (e testGRPCError) Error() string { + return e.status.Message() +} + +func (e testGRPCError) GRPCStatus() *status.Status { + return e.status +} + +func TestInterceptorsProviderOrder(t *testing.T) { + customGRPCInterceptor := func(context.Context, any, *grpc.UnaryServerInfo, grpc.UnaryHandler) (any, error) { + return nil, nil + } + customInterceptor := &interceptorWrapper{ + grpcInterceptor: customGRPCInterceptor, + nexusInterceptor: nexusNoOpInterceptor, + } + provider := newInterceptorsProvider( + nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, + []grpc.UnaryServerInterceptor{customGRPCInterceptor}, []Interceptor{customInterceptor}, nil, nil, + ) + + expectedTypes := []string{ + "*interceptor.MaskInternalErrorDetailsInterceptor", + "*interceptor.ServiceErrorInterceptor", + "*interceptor.FrontendServiceErrorInterceptor", + "*interceptor.RoutingKeyInterceptor", + "*interceptor.NamespaceLengthValidatorInterceptor", + "*interceptor.NamespaceLogInterceptor", + "*frontend.interceptorWrapper", + "*authorization.Interceptor", + "*interceptor.NamespaceHandoverInterceptor", + "*frontend.interceptorWrapper", + "*interceptor.TelemetryInterceptor", + "*interceptor.HealthInterceptor", + "*interceptor.NamespaceValidatorInterceptor", + "*interceptor.ConcurrentRequestLimitInterceptor", + "*interceptor.NamespaceRateLimitInterceptorWrapper", + "*interceptor.RateLimitInterceptor", + "*interceptor.SDKVersionInterceptor", + "*interceptor.CallerInfoInterceptor", + "*interceptor.SlowRequestLoggerInterceptor", + "*chasm.ChasmVisibilityInterceptor", + "*interceptor.ContextMetadataInterceptor", + "*frontend.interceptorWrapper", + "*frontend.interceptorWrapper", + "*grpcfaults.FaultsInterceptor", + "*interceptor.RetryableInterceptor", + } + actualTypes := make([]string, 0, len(provider.interceptors)) + for _, current := range provider.interceptors { + actualTypes = append(actualTypes, reflect.TypeOf(current).String()) + } + require.Equal(t, expectedTypes, actualTypes) + + grpcInterceptors := provider.grpcInterceptors() + nexusInterceptors := provider.nexusInterceptors() + require.Len(t, grpcInterceptors, len(expectedTypes)) + require.Len(t, nexusInterceptors, len(grpcInterceptors)+1) + require.Equal( + t, + reflect.ValueOf(provider.nexusTelemetry).Pointer(), + reflect.ValueOf(nexusInterceptors[0]).Pointer(), + "Outermost interceptor for Nexus must be telemetry", + ) +} + +func TestNexusChainPreservesNativeErrors(t *testing.T) { + tests := []struct { + name string + err error + outcome string + wrapError bool + exposeDetails bool + assertErrors func(*testing.T, error, bool) + }{ + { + name: "operation error", + err: &nexus.OperationError{ + Message: "operation failed", + State: nexus.OperationStateFailed, + Cause: errors.New("worker failure"), + }, + outcome: "operation_error", + wrapError: true, + assertErrors: func(t *testing.T, err error, _ bool) { + var operationErr *nexus.OperationError + require.ErrorAs(t, err, &operationErr) + require.Equal(t, "operation failed", operationErr.Message) + + convertedErr := convertInterceptorError(err) + require.ErrorAs(t, convertedErr, &operationErr) + require.Equal(t, nexus.OperationStateFailed, operationErr.State) + require.Equal(t, "operation failed", operationErr.Message) + }, + }, + { + name: "handler error", + err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid input"), + outcome: "handler_error", + wrapError: true, + assertErrors: func(t *testing.T, err error, _ bool) { + var handlerErr *nexus.HandlerError + require.ErrorAs(t, err, &handlerErr) + require.Equal(t, nexus.HandlerErrorTypeBadRequest, handlerErr.Type) + require.Equal(t, "invalid input", handlerErr.Message) + + require.ErrorAs(t, convertInterceptorError(err), &handlerErr) + require.Equal(t, nexus.HandlerErrorTypeBadRequest, handlerErr.Type) + }, + }, + { + name: "bare handler error", + err: nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid input"), + outcome: "internal_error", + wrapError: false, + assertErrors: func(t *testing.T, err error, _ bool) { + var handlerErr *nexus.HandlerError + require.ErrorAs(t, err, &handlerErr) + require.Equal(t, nexus.HandlerErrorTypeBadRequest, handlerErr.Type) + require.Equal(t, "invalid input", handlerErr.Message) + }, + }, + { + name: "internal gRPC error", + err: status.Error(codes.Internal, "worker failure"), + outcome: "internal_error", + wrapError: true, + assertErrors: func(t *testing.T, err error, maskErrors bool) { + require.Equal(t, codes.Internal, status.Code(err)) + if maskErrors { + require.NotContains(t, err.Error(), "worker failure") + } else { + require.ErrorContains(t, err, "worker failure") + } + + var handlerErr *nexus.HandlerError + require.ErrorAs(t, convertInterceptorError(err), &handlerErr) + require.Equal(t, nexus.HandlerErrorTypeInternal, handlerErr.Type) + require.Equal(t, "internal error", handlerErr.Message) + }, + }, + { + name: "resource exhausted error details", + err: testGRPCError{status: status.New(codes.ResourceExhausted, "namespace rate limit exceeded")}, + outcome: "namespace_rate_limited", + wrapError: true, + exposeDetails: true, + assertErrors: func(t *testing.T, err error, _ bool) { + var handlerErr *nexus.HandlerError + require.ErrorAs(t, convertInterceptorError(err), &handlerErr) + require.Equal(t, nexus.HandlerErrorTypeResourceExhausted, handlerErr.Type) + require.Contains(t, handlerErr.Message, "namespace rate limit exceeded") + }, + }, + { + name: "resource exhausted error details masked", + err: testGRPCError{status: status.New(codes.ResourceExhausted, "namespace rate limit exceeded")}, + outcome: "namespace_rate_limited", + wrapError: true, + assertErrors: func(t *testing.T, err error, _ bool) { + var handlerErr *nexus.HandlerError + require.ErrorAs(t, convertInterceptorError(err), &handlerErr) + require.Equal(t, nexus.HandlerErrorTypeResourceExhausted, handlerErr.Type) + require.Equal(t, "resource exhausted", handlerErr.Message) + }, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + for _, maskErrors := range []bool{false, true} { + t.Run(fmt.Sprintf("mask errors=%t", maskErrors), func(t *testing.T) { + t.Parallel() + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + chainedHandler := newTestNexusInterceptorChain(metricsHandler, maskErrors, tc.err, tc.outcome, tc.wrapError, tc.exposeDetails) + _, err := chainedHandler(context.Background(), newTestNexusStartInput()) + + if tc.wrapError { + var interceptorErr *interceptornexus.InterceptorError + require.ErrorAs(t, err, &interceptorErr) + require.Equal(t, tc.outcome, interceptorErr.Outcome) + } + tc.assertErrors(t, err, maskErrors) + + snapshot := capture.Snapshot() + require.Len(t, snapshot[metrics.NexusRequests.Name()], 1) + require.Equal(t, tc.outcome, snapshot[metrics.NexusRequests.Name()][0].Tags["outcome"]) + }) + } + }) + } +} + +func newTestNexusInterceptorChain( + metricsHandler metrics.Handler, + maskErrors bool, + terminalErr error, + outcome string, + wrapError bool, + exposeDetails bool, +) interceptornexus.HandlerFunc { + telemetry := rpcinterceptor.NewTelemetryInterceptor(nil, metricsHandler, log.NewNoopLogger(), nil, nil) + mask := rpcinterceptor.NewMaskInternalErrorDetailsInterceptor( + dynamicconfig.GetBoolPropertyFnFilteredByNamespace(maskErrors), + nil, + log.NewNoopLogger(), + ) + serviceErrors := rpcinterceptor.NewServiceErrorInterceptor( + dynamicconfig.GetIntPropertyFn(4000), + metrics.NoopMetricsHandler, + log.NewNoopLogger(), + ) + frontendServiceErrors := rpcinterceptor.NewFrontendServiceErrorInterceptorWrapper(log.NewNoopLogger()) + + return interceptornexus.ChainInterceptors( + func(context.Context, interceptornexus.InterceptorInput) (any, error) { + if !wrapError { + return nil, terminalErr + } + return nil, &interceptornexus.InterceptorError{Err: terminalErr, Outcome: outcome, ExposeDetails: exposeDetails} + }, + []interceptornexus.Interceptor{ + telemetry.InterceptNexusOutermost, + mask.InterceptNexus, + serviceErrors.InterceptNexus, + frontendServiceErrors.InterceptNexus, + }, + ) +} + +func newTestNexusStartInput() interceptornexus.StartOpInput { + return interceptornexus.NewStartOpInput( + "s", + "o", + testNamespace, + time.Now(), + nexus.StartOperationOptions{}, + nil, + interceptornexus.ForwardingInfo{}, + interceptornexus.RequestMetadata{NamespaceEntry: testOperationContext().namespace}, + ) +} diff --git a/service/frontend/nexus_operation_http_handler.go b/service/frontend/nexus_operation_http_handler.go index 93998d2f17e..62e43430c94 100644 --- a/service/frontend/nexus_operation_http_handler.go +++ b/service/frontend/nexus_operation_http_handler.go @@ -16,7 +16,6 @@ import ( "go.temporal.io/server/api/matchingservice/v1" persistencespb "go.temporal.io/server/api/persistence/v1" "go.temporal.io/server/common/authorization" - "go.temporal.io/server/common/cluster" "go.temporal.io/server/common/log" "go.temporal.io/server/common/log/tag" "go.temporal.io/server/common/metrics" @@ -36,75 +35,56 @@ import ( // Small wrapper that does some pre-processing before handing requests over to the Nexus SDK's HTTP handler. type NexusOperationHTTPHandler struct { - base nexusrpc.BaseHTTPHandler - logger log.Logger - nexusHandler http.Handler - enpointRegistry commonnexus.EndpointRegistry - namespaceRegistry namespace.Registry - preprocessErrorCounter metrics.CounterFunc - auth *authorization.Interceptor - namespaceValidationInterceptor *interceptor.NamespaceValidatorInterceptor - namespaceRateLimitInterceptor interceptor.NamespaceRateLimitInterceptor - namespaceConcurrencyLimitInterceptor *interceptor.ConcurrentRequestLimitInterceptor - rateLimitInterceptor *interceptor.RateLimitInterceptor - httpServerHandlerInstrumenter telemetry.HTTPServerHandlerInstrumenter + base nexusrpc.BaseHTTPHandler + logger log.Logger + nexusHandler http.Handler + enpointRegistry commonnexus.EndpointRegistry + namespaceRegistry namespace.Registry + namespaceValidationInterceptor *interceptor.NamespaceValidatorInterceptor + preprocessErrorCounter metrics.CounterFunc + auth *authorization.Interceptor + httpServerHandlerInstrumenter telemetry.HTTPServerHandlerInstrumenter } func NewNexusOperationHTTPHandler( serviceConfig *Config, matchingClient resource.MatchingClient, metricsHandler metrics.Handler, - clusterMetadata cluster.Metadata, - clientCache *cluster.FrontendHTTPClientCache, namespaceRegistry namespace.Registry, endpointRegistry commonnexus.EndpointRegistry, authInterceptor *authorization.Interceptor, - telemetryInterceptor *interceptor.TelemetryInterceptor, - requestErrorHandler *interceptor.RequestErrorHandler, - redirectionInterceptor *interceptor.Redirection, namespaceValidationInterceptor *interceptor.NamespaceValidatorInterceptor, - namespaceRateLimitInterceptor interceptor.NamespaceRateLimitInterceptor, - namespaceConcurrencyLimitInterceptor *interceptor.ConcurrentRequestLimitInterceptor, - rateLimitInterceptor *interceptor.RateLimitInterceptor, + requestErrorHandler *interceptor.RequestErrorHandler, + interceptorsProvider *interceptorsProvider, logger log.Logger, - httpTraceProvider commonnexus.HTTPClientTraceProvider, httpServerHandlerInstrumenter telemetry.HTTPServerHandlerInstrumenter, ) *NexusOperationHTTPHandler { logger = log.With(logger, tag.NexusStageHandlerInbound) + return &NexusOperationHTTPHandler{ base: nexusrpc.BaseHTTPHandler{ Logger: log.NewSlogLogger(logger), FailureConverter: nexusrpc.DefaultFailureConverter(), }, - logger: logger, - enpointRegistry: endpointRegistry, - namespaceRegistry: namespaceRegistry, - auth: authInterceptor, - namespaceValidationInterceptor: namespaceValidationInterceptor, - namespaceRateLimitInterceptor: namespaceRateLimitInterceptor, - namespaceConcurrencyLimitInterceptor: namespaceConcurrencyLimitInterceptor, - rateLimitInterceptor: rateLimitInterceptor, - preprocessErrorCounter: metricsHandler.Counter(metrics.NexusRequestPreProcessErrors.Name()).Record, - httpServerHandlerInstrumenter: httpServerHandlerInstrumenter, + logger: logger, + enpointRegistry: endpointRegistry, + namespaceRegistry: namespaceRegistry, + auth: authInterceptor, + namespaceValidationInterceptor: namespaceValidationInterceptor, + preprocessErrorCounter: metricsHandler.Counter(metrics.NexusRequestPreProcessErrors.Name()).Record, + httpServerHandlerInstrumenter: httpServerHandlerInstrumenter, nexusHandler: nexusrpc.NewHTTPHandler(nexusrpc.HandlerOptions{ - Handler: &nexusHandler{ - logger: logger, - metricsHandler: metricsHandler, - clusterMetadata: clusterMetadata, - namespaceRegistry: namespaceRegistry, - matchingClient: matchingservice.MatchingServiceClient(matchingClient), - auth: authInterceptor, - telemetryInterceptor: telemetryInterceptor, - requestErrorHandler: requestErrorHandler, - redirectionInterceptor: redirectionInterceptor, - forwardingEnabledForNamespace: serviceConfig.EnableNamespaceNotActiveAutoForwarding, - forwardingClients: clientCache, - payloadSizeLimit: serviceConfig.BlobSizeLimitError, - headersBlacklist: serviceConfig.NexusRequestHeadersBlacklist, - useForwardByEndpoint: serviceConfig.NexusForwardRequestUseEndpoint, - metricTagConfig: serviceConfig.NexusOperationsMetricTagConfig, - httpTraceProvider: httpTraceProvider, - }, + Handler: newNexusHandler( + logger, + metricsHandler, + namespaceRegistry, + matchingservice.MatchingServiceClient(matchingClient), + requestErrorHandler, + serviceConfig.BlobSizeLimitError, + serviceConfig.NexusRequestHeadersBlacklist, + serviceConfig.NexusOperationsMetricTagConfig, + interceptorsProvider.nexusInterceptors(), + ), GetResultTimeout: serviceConfig.KeepAliveMaxConnectionIdle(), Logger: log.NewSlogLogger(logger), Serializer: commonnexus.PayloadSerializer, @@ -163,7 +143,7 @@ func (h *NexusOperationHTTPHandler) dispatchNexusTaskByNamespaceAndTaskQueue(w h return } - rWithAuthCtx, err := h.parseTLSAndAuthInfo(r, nc) + rWithAuthCtx, err := h.parseTLSAndAuthInfo(r) if err != nil { logger.Error("failed to get claims", tag.Error(err)) h.writeFailure(w, r, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnauthenticated, "unauthorized")) @@ -226,7 +206,7 @@ func (h *NexusOperationHTTPHandler) dispatchNexusTaskByEndpoint(w http.ResponseW return } - rWithAuthCtx, err := h.parseTLSAndAuthInfo(r, nc) + rWithAuthCtx, err := h.parseTLSAndAuthInfo(r) if err != nil { logger.Error("failed to get claims", tag.Error(err)) h.writeFailure(w, r, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnauthenticated, "unauthorized")) @@ -246,14 +226,10 @@ func (h *NexusOperationHTTPHandler) dispatchNexusTaskByEndpoint(w http.ResponseW func (h *NexusOperationHTTPHandler) baseNexusContext(apiName string, header http.Header) *nexusContext { return &nexusContext{ - namespaceValidationInterceptor: h.namespaceValidationInterceptor, - namespaceRateLimitInterceptor: h.namespaceRateLimitInterceptor, - namespaceConcurrencyLimitInterceptor: h.namespaceConcurrencyLimitInterceptor, - rateLimitInterceptor: h.rateLimitInterceptor, - apiName: apiName, - requestStartTime: time.Now(), - responseHeaders: make(map[string]string), - callerFailureSupport: header.Get(nexusrpc.HeaderTemporalNexusFailureSupport) == "true", + apiName: apiName, + requestStartTime: time.Now(), + responseHeaders: make(map[string]string), + callerFailureSupport: header.Get(nexusrpc.HeaderTemporalNexusFailureSupport) == "true", } } @@ -310,7 +286,7 @@ func prepareRequest[T any](route routing.Route[T], w http.ResponseWriter, r *htt return route.Deserialize(vars) } -func (h *NexusOperationHTTPHandler) parseTLSAndAuthInfo(r *http.Request, nc *nexusContext) (*http.Request, error) { +func (h *NexusOperationHTTPHandler) parseTLSAndAuthInfo(r *http.Request) (*http.Request, error) { var tlsInfo *credentials.TLSInfo if r.TLS != nil { tlsInfo = &credentials.TLSInfo{ @@ -323,14 +299,13 @@ func (h *NexusOperationHTTPHandler) parseTLSAndAuthInfo(r *http.Request, nc *nex return "" // TODO: support audience getter }) - var err error if authInfo != nil { - nc.claims, err = h.auth.GetClaims(authInfo) + claims, err := h.auth.GetClaims(authInfo) if err != nil { return nil, err } // Make the auth info and claims available on the context. - r = r.WithContext(h.auth.EnhanceContext(r.Context(), authInfo, nc.claims)) + r = r.WithContext(h.auth.EnhanceContext(r.Context(), authInfo, claims)) } return r, nil diff --git a/service/fx.go b/service/fx.go index f41031b62a4..30490aa1b3d 100644 --- a/service/fx.go +++ b/service/fx.go @@ -165,7 +165,7 @@ func getUnaryInterceptors(params GrpcServerOptionsParams) []grpc.UnaryServerInte params.ServiceErrorInterceptor.Intercept, metrics.NewServerMetricsContextInjectorInterceptor(), metrics.NewServerMetricsTrailerPropagatorInterceptor(params.Logger), - params.TelemetryInterceptor.UnaryIntercept, + params.TelemetryInterceptor.Intercept, } interceptors = append(interceptors, params.AdditionalInterceptors...) diff --git a/service/history/history_engine_test.go b/service/history/history_engine_test.go index 7550fdcfad9..05a2751ad13 100644 --- a/service/history/history_engine_test.go +++ b/service/history/history_engine_test.go @@ -5588,7 +5588,7 @@ func (s *engineSuite) TestEagerWorkflowStart_DoesNotCreateTransferTask() { s.mockShard.Resource.Logger, s.config.LogAllReqErrors, s.mockErrorHandler) - response, err := i.UnaryIntercept(context.Background(), nil, &grpc.UnaryServerInfo{FullMethod: "StartWorkflowExecution"}, func(ctx context.Context, req any) (any, error) { + response, err := i.Intercept(context.Background(), nil, &grpc.UnaryServerInfo{FullMethod: "StartWorkflowExecution"}, func(ctx context.Context, req any) (any, error) { response, err := s.historyEngine.StartWorkflowExecution(ctx, &historyservice.StartWorkflowExecutionRequest{ NamespaceId: tests.NamespaceID.String(), Attempt: 1, @@ -5627,7 +5627,7 @@ func (s *engineSuite) TestEagerWorkflowStart_FromCron_SkipsEager() { s.mockShard.Resource.Logger, s.config.LogAllReqErrors, s.mockErrorHandler) - response, err := i.UnaryIntercept(context.Background(), nil, &grpc.UnaryServerInfo{FullMethod: "StartWorkflowExecution"}, func(ctx context.Context, req any) (any, error) { + response, err := i.Intercept(context.Background(), nil, &grpc.UnaryServerInfo{FullMethod: "StartWorkflowExecution"}, func(ctx context.Context, req any) (any, error) { firstWorkflowTaskBackoff := time.Second response, err := s.historyEngine.StartWorkflowExecution(ctx, &historyservice.StartWorkflowExecutionRequest{ NamespaceId: tests.NamespaceID.String(), @@ -5671,7 +5671,7 @@ func (s *engineSuite) TestEagerWorkflowStart_WithSearchAttributes() { s.mockShard.Resource.Logger, s.config.LogAllReqErrors, s.mockErrorHandler) - response, err := i.UnaryIntercept(context.Background(), nil, &grpc.UnaryServerInfo{FullMethod: "StartWorkflowExecution"}, func(ctx context.Context, req any) (any, error) { + response, err := i.Intercept(context.Background(), nil, &grpc.UnaryServerInfo{FullMethod: "StartWorkflowExecution"}, func(ctx context.Context, req any) (any, error) { response, err := s.historyEngine.StartWorkflowExecution(ctx, &historyservice.StartWorkflowExecutionRequest{ NamespaceId: tests.NamespaceID.String(), Attempt: 1, diff --git a/temporal/fx.go b/temporal/fx.go index 221b731335a..398aecd7440 100644 --- a/temporal/fx.go +++ b/temporal/fx.go @@ -121,6 +121,8 @@ type ( TokenProvider auth.TokenProvider ServiceHosts map[primitives.ServiceName]static.Hosts + CustomFrontendUnifiedInterceptors []frontend.Interceptor + // below are things that could be over write by server options or may have default if not supplied by serverOptions. Logger log.Logger ClientFactoryProvider client.FactoryProvider @@ -318,11 +320,12 @@ func ServerOptionsProvider(opts []ServerOption) (serverOptionsProvider, error) { ServiceHosts: so.hostsByService, NamespaceLogger: so.namespaceLogger, - ServiceResolver: so.persistenceServiceResolver, - CustomDataStoreFactory: so.customDataStoreFactory, - CustomVisibilityStore: so.customVisibilityStoreFactory, - CustomHistoryArchiverFactory: so.customHistoryArchiverFactory, - CustomVisibilityArchiverFactory: so.customVisibilityArchiverFactory, + ServiceResolver: so.persistenceServiceResolver, + CustomDataStoreFactory: so.customDataStoreFactory, + CustomVisibilityStore: so.customVisibilityStoreFactory, + CustomHistoryArchiverFactory: so.customHistoryArchiverFactory, + CustomVisibilityArchiverFactory: so.customVisibilityArchiverFactory, + CustomFrontendUnifiedInterceptors: so.customFrontendUnifiedInterceptors, SearchAttributesMapper: so.searchAttributesMapper, CustomFrontendInterceptors: so.customFrontendInterceptors, @@ -380,36 +383,37 @@ type ( ServiceProviderParamsCommon struct { fx.In - Cfg *config.Config - ServiceNames resource.ServiceNames - Logger log.Logger - NamespaceLogger resource.NamespaceLogger - DynamicConfigClient dynamicconfig.Client - MetricsHandler metrics.Handler - EventLoggerProvider otellog.LoggerProvider - EsClient esclient.Client - TlsConfigProvider encryption.TLSConfigProvider //nolint:staticcheck // should be TLSConfigProvider - PersistenceConfig config.Persistence - ClusterMetadata *cluster.Config - ClientFactoryProvider client.FactoryProvider - AudienceGetter authorization.JWTAudienceMapper - PersistenceServiceResolver resolver.ServiceResolver - PersistenceFactoryProvider persistenceClient.FactoryProviderFn - SearchAttributesMapper searchattribute.Mapper - CustomFrontendInterceptors []grpc.UnaryServerInterceptor - AdditionalStreamInterceptors []grpc.StreamServerInterceptor - Authorizer authorization.Authorizer - ClaimMapper authorization.ClaimMapper - TokenProvider auth.TokenProvider - DataStoreFactory persistenceClient.AbstractDataStoreFactory - VisibilityStoreFactory visibility.VisibilityStoreFactory - CustomHistoryArchiverFactory provider.CustomHistoryArchiverFactory - CustomVisibilityArchiverFactory provider.CustomVisibilityArchiverFactory - SpanExporters []otelsdktrace.SpanExporter - InstanceID resource.InstanceID `optional:"true"` - StaticServiceHosts map[primitives.ServiceName]static.Hosts `optional:"true"` - TaskCategoryRegistry tasks.TaskCategoryRegistry - TestHooks testhooks.TestHooks + Cfg *config.Config + ServiceNames resource.ServiceNames + Logger log.Logger + NamespaceLogger resource.NamespaceLogger + DynamicConfigClient dynamicconfig.Client + MetricsHandler metrics.Handler + EventLoggerProvider otellog.LoggerProvider + EsClient esclient.Client + TlsConfigProvider encryption.TLSConfigProvider //nolint:staticcheck // should be TLSConfigProvider + PersistenceConfig config.Persistence + ClusterMetadata *cluster.Config + ClientFactoryProvider client.FactoryProvider + AudienceGetter authorization.JWTAudienceMapper + PersistenceServiceResolver resolver.ServiceResolver + PersistenceFactoryProvider persistenceClient.FactoryProviderFn + SearchAttributesMapper searchattribute.Mapper + CustomFrontendInterceptors []grpc.UnaryServerInterceptor + CustomFrontendUnifiedInterceptors []frontend.Interceptor + AdditionalStreamInterceptors []grpc.StreamServerInterceptor + Authorizer authorization.Authorizer + ClaimMapper authorization.ClaimMapper + TokenProvider auth.TokenProvider + DataStoreFactory persistenceClient.AbstractDataStoreFactory + VisibilityStoreFactory visibility.VisibilityStoreFactory + CustomHistoryArchiverFactory provider.CustomHistoryArchiverFactory + CustomVisibilityArchiverFactory provider.CustomVisibilityArchiverFactory + SpanExporters []otelsdktrace.SpanExporter + InstanceID resource.InstanceID `optional:"true"` + StaticServiceHosts map[primitives.ServiceName]static.Hosts `optional:"true"` + TaskCategoryRegistry tasks.TaskCategoryRegistry + TestHooks testhooks.TestHooks } ) @@ -594,6 +598,7 @@ func genericFrontendServiceProvider( app := fx.New( params.GetCommonServiceOptions(serviceName), fx.Supply(params.CustomFrontendInterceptors), + fx.Supply(params.CustomFrontendUnifiedInterceptors), fx.Decorate(func() authorization.ClaimMapper { switch serviceName { case primitives.FrontendService: diff --git a/temporal/server_option.go b/temporal/server_option.go index b72cfe81e8e..1edd06cda6b 100644 --- a/temporal/server_option.go +++ b/temporal/server_option.go @@ -20,6 +20,7 @@ import ( "go.temporal.io/server/common/rpc/encryption" "go.temporal.io/server/common/searchattribute" "go.temporal.io/server/common/testing/testhooks" + "go.temporal.io/server/service/frontend" "google.golang.org/grpc" ) @@ -201,6 +202,8 @@ func WithSearchAttributesMapper(m searchattribute.Mapper) ServerOption { // Frontend gRPC API calls. The list of custom interceptors will be appended to the end of the internal // ServerInterceptors. The custom interceptors will be invoked in the order as they appear in the supplied list, after // the internal ServerInterceptors. +// +// Deprecated: Use [WithChainedFrontendInterceptors] instead. These options are mutually exclusive. func WithChainedFrontendGrpcInterceptors( interceptors ...grpc.UnaryServerInterceptor, ) ServerOption { @@ -209,6 +212,18 @@ func WithChainedFrontendGrpcInterceptors( }) } +// WithChainedFrontendInterceptors sets an ordered chain of custom gRPC+Nexus interceptors that will be invoked for all +// Frontend gRPC and Nexus API calls respectively. Custom interceptors run after the internal +// interceptors and before the fault-injection and retryable interceptors, in the order supplied. +// Cannot be used with [WithChainedFrontendGrpcInterceptors]- they are mutually exclusive. +func WithChainedFrontendInterceptors( + interceptors ...frontend.Interceptor, +) ServerOption { + return applyFunc(func(s *serverOptions) { + s.customFrontendUnifiedInterceptors = interceptors + }) +} + // WithAdditionalStreamInterceptors sets a chain of ordered custom grpc stream interceptors that will be invoked for all // service gRPC stream calls. The list of custom interceptors will be appended to the end of the internal // ServerInterceptors. The custom interceptors will be invoked in the order as they appear in the supplied list, after diff --git a/temporal/server_options.go b/temporal/server_options.go index 388425738cb..8ad988b7503 100644 --- a/temporal/server_options.go +++ b/temporal/server_options.go @@ -23,6 +23,7 @@ import ( "go.temporal.io/server/common/rpc/encryption" "go.temporal.io/server/common/searchattribute" "go.temporal.io/server/common/testing/testhooks" + "go.temporal.io/server/service/frontend" "google.golang.org/grpc" ) @@ -44,28 +45,29 @@ type ( startupSynchronizationMode synchronizationModeParams - logger log.Logger - namespaceLogger log.Logger - authorizer authorization.Authorizer - tlsConfigProvider encryption.TLSConfigProvider - claimMapper authorization.ClaimMapper - audienceGetter authorization.JWTAudienceMapper - persistenceServiceResolver resolver.ServiceResolver - elasticsearchHttpClient *http.Client //nolint:staticcheck // should be elasticsearchHTTPClient - dynamicConfigClient dynamicconfig.Client - customDataStoreFactory persistenceClient.AbstractDataStoreFactory - customVisibilityStoreFactory visibility.VisibilityStoreFactory - customHistoryArchiverFactory provider.CustomHistoryArchiverFactory - customVisibilityArchiverFactory provider.CustomVisibilityArchiverFactory - clientFactoryProvider client.FactoryProvider - persistenceFactoryProvider persistenceClient.FactoryProviderFn - searchAttributesMapper searchattribute.Mapper - customFrontendInterceptors []grpc.UnaryServerInterceptor - additionalStreamInterceptors []grpc.StreamServerInterceptor - metricHandler metrics.Handler - eventLoggerProvider otellog.LoggerProvider - tokenProvider auth.TokenProvider - testHooks *testhooks.TestHooks + logger log.Logger + namespaceLogger log.Logger + authorizer authorization.Authorizer + tlsConfigProvider encryption.TLSConfigProvider + claimMapper authorization.ClaimMapper + audienceGetter authorization.JWTAudienceMapper + persistenceServiceResolver resolver.ServiceResolver + elasticsearchHttpClient *http.Client //nolint:staticcheck // should be elasticsearchHTTPClient + dynamicConfigClient dynamicconfig.Client + customDataStoreFactory persistenceClient.AbstractDataStoreFactory + customVisibilityStoreFactory visibility.VisibilityStoreFactory + customHistoryArchiverFactory provider.CustomHistoryArchiverFactory + customVisibilityArchiverFactory provider.CustomVisibilityArchiverFactory + clientFactoryProvider client.FactoryProvider + persistenceFactoryProvider persistenceClient.FactoryProviderFn + searchAttributesMapper searchattribute.Mapper + customFrontendInterceptors []grpc.UnaryServerInterceptor + customFrontendUnifiedInterceptors []frontend.Interceptor + additionalStreamInterceptors []grpc.StreamServerInterceptor + metricHandler metrics.Handler + eventLoggerProvider otellog.LoggerProvider + tokenProvider auth.TokenProvider + testHooks *testhooks.TestHooks } ) @@ -130,6 +132,12 @@ func (so *serverOptions) loadConfig() error { } func (so *serverOptions) validateConfig() error { + if len(so.customFrontendInterceptors) > 0 && + len(so.customFrontendUnifiedInterceptors) > 0 { + // Both could be supported as a migration path but intentionally avoided as + // migration itself is as simple as wrapping with no-op Nexus Interceptors. + return errors.New("WithChainedFrontendGrpcInterceptors is deprecated in favor of WithChainedFrontendInterceptors- they cannot both be set") + } if err := so.config.Validate(); err != nil { return err } diff --git a/temporal/server_options_test.go b/temporal/server_options_test.go new file mode 100644 index 00000000000..ea279076ed9 --- /dev/null +++ b/temporal/server_options_test.go @@ -0,0 +1,46 @@ +package temporal + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + "go.temporal.io/server/common/rpc/interceptor/nexus" + "go.temporal.io/server/service/frontend" + "google.golang.org/grpc" +) + +type testFrontendInterceptor struct{} + +func (testFrontendInterceptor) Intercept( + context.Context, + any, + *grpc.UnaryServerInfo, + grpc.UnaryHandler, +) (any, error) { + return nil, nil +} + +func (testFrontendInterceptor) InterceptNexus( + context.Context, + nexus.InterceptorInput, + nexus.HandlerFunc, +) (any, error) { + return nil, nil +} + +var _ frontend.Interceptor = testFrontendInterceptor{} + +func TestServerOptionsRejectsBothFrontendInterceptorOptions(t *testing.T) { + options := serverOptions{ + customFrontendInterceptors: []grpc.UnaryServerInterceptor{ + func(context.Context, any, *grpc.UnaryServerInfo, grpc.UnaryHandler) (any, error) { + return nil, nil + }, + }, + customFrontendUnifiedInterceptors: []frontend.Interceptor{testFrontendInterceptor{}}, + } + + err := options.validateConfig() + require.EqualError(t, err, "WithChainedFrontendGrpcInterceptors is deprecated in favor of WithChainedFrontendInterceptors- they cannot both be set") +} diff --git a/tools/flakereport/report.go b/tools/flakereport/report.go index 874e9ca57a6..18a3266d9ed 100644 --- a/tools/flakereport/report.go +++ b/tools/flakereport/report.go @@ -148,7 +148,7 @@ func generateOccurrenceReportTable(reports []TestReport, nameHeader, countHeader reports = limitReportRows(reports) var sb strings.Builder - sb.WriteString(fmt.Sprintf("| %s | %s | Last Occurrence | Trend | Links |\n", nameHeader, countHeader)) + fmt.Fprintf(&sb, "| %s | %s | Last Occurrence | Trend | Links |\n", nameHeader, countHeader) sb.WriteString("|------|--------------------|-----------------|-------|-------|\n") for _, report := range reports { links := formatLinks(report.GitHubURLs, maxLinks) @@ -156,8 +156,8 @@ func generateOccurrenceReportTable(reports []TestReport, nameHeader, countHeader if !report.LastFailure.IsZero() { lastOccurrence = hoursAgo(report.LastFailure) } - sb.WriteString(fmt.Sprintf("| `%s` | %d | %s | `%s` | %s |\n", - report.TestName, report.FailureCount, lastOccurrence, formatSparkline(report.TrendPoints), links)) + fmt.Fprintf(&sb, "| `%s` | %d | %s | `%s` | %s |\n", + report.TestName, report.FailureCount, lastOccurrence, formatSparkline(report.TrendPoints), links) } return sb.String()