diff --git a/service/history/nexus_dispatch_response.go b/common/nexus/dispatch_response.go similarity index 84% rename from service/history/nexus_dispatch_response.go rename to common/nexus/dispatch_response.go index 63b9b07a991..143525b85d0 100644 --- a/service/history/nexus_dispatch_response.go +++ b/common/nexus/dispatch_response.go @@ -1,4 +1,4 @@ -package history +package nexus import ( "github.com/nexus-rpc/sdk-go/nexus" @@ -7,14 +7,14 @@ import ( "go.temporal.io/server/api/matchingservice/v1" ) -// dispatchResponseToError converts a DispatchNexusTaskResponse proto into a Go error. +// DispatchResponseToError converts a DispatchNexusTaskResponse proto into a Go error. // Returns nil if the response indicates success. // // For failure cases (worker explicitly returned an error), the Temporal SDK's failure // converter is used to produce standard Go errors (ApplicationError, CanceledError). // For transport-level issues (timeout, internal), a nexus.HandlerError is returned // so the caller can check Retryable(). -func dispatchResponseToError(resp *matchingservice.DispatchNexusTaskResponse) error { +func DispatchResponseToError(resp *matchingservice.DispatchNexusTaskResponse) error { switch t := resp.GetOutcome().(type) { case *matchingservice.DispatchNexusTaskResponse_Failure: // Worker received the task and explicitly failed it (via RespondNexusTaskFailed). @@ -22,15 +22,15 @@ func dispatchResponseToError(resp *matchingservice.DispatchNexusTaskResponse) er case *matchingservice.DispatchNexusTaskResponse_RequestTimeout: return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUpstreamTimeout, "upstream timeout") case *matchingservice.DispatchNexusTaskResponse_Response: - return startOperationResponseToError(t.Response.GetStartOperation()) + return StartOperationResponseToError(t.Response.GetStartOperation()) default: return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "empty or unknown dispatch outcome") } } -// startOperationResponseToError converts a StartOperationResponse proto into a Go error. +// StartOperationResponseToError converts a StartOperationResponse proto into a Go error. // Returns nil for success variants (SyncSuccess, AsyncSuccess). -func startOperationResponseToError(resp *nexuspb.StartOperationResponse) error { +func StartOperationResponseToError(resp *nexuspb.StartOperationResponse) error { switch t := resp.GetVariant().(type) { case *nexuspb.StartOperationResponse_SyncSuccess: return nil diff --git a/service/history/nexus_dispatch_response_test.go b/common/nexus/dispatch_response_test.go similarity index 93% rename from service/history/nexus_dispatch_response_test.go rename to common/nexus/dispatch_response_test.go index 76b1b9568ba..b98829fcca1 100644 --- a/service/history/nexus_dispatch_response_test.go +++ b/common/nexus/dispatch_response_test.go @@ -1,4 +1,4 @@ -package history +package nexus import ( "testing" @@ -25,7 +25,7 @@ func TestDispatchResponseToError_SyncSuccess(t *testing.T) { }, }, } - err := dispatchResponseToError(resp) + err := DispatchResponseToError(resp) require.NoError(t, err) } @@ -45,7 +45,7 @@ func TestDispatchResponseToError_AsyncSuccess(t *testing.T) { }, }, } - err := dispatchResponseToError(resp) + err := DispatchResponseToError(resp) require.NoError(t, err) } @@ -55,7 +55,7 @@ func TestDispatchResponseToError_RequestTimeout(t *testing.T) { RequestTimeout: &matchingservice.DispatchNexusTaskResponse_Timeout{}, }, } - err := dispatchResponseToError(resp) + err := DispatchResponseToError(resp) require.Error(t, err) var handlerErr *nexus.HandlerError @@ -76,7 +76,7 @@ func TestDispatchResponseToError_WorkerFailure(t *testing.T) { }, }, } - err := dispatchResponseToError(resp) + err := DispatchResponseToError(resp) require.Error(t, err) var appErr *temporal.ApplicationError @@ -105,7 +105,7 @@ func TestDispatchResponseToError_OperationFailure_ApplicationError(t *testing.T) }, }, } - err := dispatchResponseToError(resp) + err := DispatchResponseToError(resp) require.Error(t, err) var appErr *temporal.ApplicationError @@ -132,7 +132,7 @@ func TestDispatchResponseToError_OperationFailure_CanceledError(t *testing.T) { }, }, } - err := dispatchResponseToError(resp) + err := DispatchResponseToError(resp) require.Error(t, err) var cancelErr *temporal.CanceledError @@ -141,7 +141,7 @@ func TestDispatchResponseToError_OperationFailure_CanceledError(t *testing.T) { func TestDispatchResponseToError_EmptyOutcome(t *testing.T) { resp := &matchingservice.DispatchNexusTaskResponse{} - err := dispatchResponseToError(resp) + err := DispatchResponseToError(resp) require.Error(t, err) var handlerErr *nexus.HandlerError @@ -151,7 +151,7 @@ func TestDispatchResponseToError_EmptyOutcome(t *testing.T) { func TestStartOperationResponseToError_EmptyVariant(t *testing.T) { resp := &nexuspb.StartOperationResponse{} - err := startOperationResponseToError(resp) + err := StartOperationResponseToError(resp) require.Error(t, err) var handlerErr *nexus.HandlerError diff --git a/common/workercommands/dispatch.go b/common/workercommands/dispatch.go new file mode 100644 index 00000000000..7c771c15143 --- /dev/null +++ b/common/workercommands/dispatch.go @@ -0,0 +1,139 @@ +package workercommands + +import ( + "context" + "errors" + "fmt" + + "github.com/nexus-rpc/sdk-go/nexus" + commonpb "go.temporal.io/api/common/v1" + enumspb "go.temporal.io/api/enums/v1" + nexuspb "go.temporal.io/api/nexus/v1" + workerservicepb "go.temporal.io/api/nexusservices/workerservice/v1" + taskqueuepb "go.temporal.io/api/taskqueue/v1" + workerpb "go.temporal.io/api/worker/v1" + "go.temporal.io/server/api/matchingservice/v1" + "go.temporal.io/server/common/log" + "go.temporal.io/server/common/log/tag" + "go.temporal.io/server/common/metrics" + commonnexus "go.temporal.io/server/common/nexus" + "go.temporal.io/server/common/resource" + "google.golang.org/protobuf/proto" +) + +const ( + // Nexus service and operation names for worker commands. + // TODO: Replace with workerservicepb.WorkerService.ServiceName and + // workerservicepb.WorkerService.ExecuteCommands.Name() once the Nexus service + // descriptor is published in go.temporal.io/api. + ServiceName = "temporal.api.nexusservices.workerservice.v1.WorkerService" + OperationName = "ExecuteCommands" +) + +// DispatchToWorker dispatches worker commands to a worker's control queue via Nexus. +// It encodes the commands as binary/protobuf (which SDK Core can decode natively via prost), +// sends them via DispatchNexusTask to matching, and handles the response. +// Returns nil on success or permanent (non-retryable) errors. Returns an error for +// retryable failures so the caller can retry. +func DispatchToWorker( + ctx context.Context, + matchingClient resource.MatchingClient, + metricsHandler metrics.Handler, + logger log.Logger, + namespaceID string, + controlQueue string, + commands []*workerpb.WorkerCommand, +) error { + request := &workerservicepb.ExecuteCommandsRequest{ + Commands: commands, + } + // Encode as binary/protobuf using the standard Temporal payload format. + // Worker commands are handled directly by SDK Core (not by lang-SDK Nexus handlers), + // so we use binary/protobuf which Core can decode natively via prost. The standard + // payload.Encode() uses json/protobuf encoding, which Core does not support because + // it normally delegates Nexus payload deserialization to the lang SDK. + requestData, err := proto.Marshal(request) + if err != nil { + return fmt.Errorf("failed to encode worker commands request: %w", err) + } + requestPayload := &commonpb.Payload{ + Metadata: map[string][]byte{ + "encoding": []byte("binary/protobuf"), + }, + Data: requestData, + } + + nexusRequest := &nexuspb.Request{ + Header: map[string]string{}, + Variant: &nexuspb.Request_StartOperation{ + StartOperation: &nexuspb.StartOperationRequest{ + Service: ServiceName, + Operation: OperationName, + Payload: requestPayload, + }, + }, + } + + resp, err := matchingClient.DispatchNexusTask(ctx, &matchingservice.DispatchNexusTaskRequest{ + NamespaceId: namespaceID, + TaskQueue: &taskqueuepb.TaskQueue{ + Name: controlQueue, + Kind: enumspb.TASK_QUEUE_KIND_WORKER_COMMANDS, + }, + Request: nexusRequest, + }) + if err != nil { + logger.Warn("Failed to dispatch worker commands", + tag.NewStringTag("control_queue", controlQueue), + tag.Error(err)) + metrics.WorkerCommandsSent.With(metricsHandler).Record(1, metrics.OutcomeTag("rpc_error")) + return err + } + + nexusErr := commonnexus.DispatchResponseToError(resp) + if nexusErr == nil { + metrics.WorkerCommandsSent.With(metricsHandler).Record(1, metrics.OutcomeTag("success")) + return nil + } + + return HandleDispatchError(nexusErr, controlQueue, metricsHandler, logger) +} + +// HandleDispatchError classifies a Nexus dispatch error and records the appropriate metric. +// Returns nil for permanent errors (caller should not retry) and the original error for +// retryable failures. +func HandleDispatchError(nexusErr error, controlQueue string, metricsHandler metrics.Handler, logger log.Logger) error { + var handlerErr *nexus.HandlerError + if errors.As(nexusErr, &handlerErr) { + if handlerErr.Type == nexus.HandlerErrorTypeUpstreamTimeout { + logger.Warn("No worker polling control queue", + tag.NewStringTag("control_queue", controlQueue)) + metrics.WorkerCommandsSent.With(metricsHandler).Record(1, metrics.OutcomeTag("no_poller")) + return nexusErr + } + + if !handlerErr.Retryable() { + logger.Error("Worker commands non-retryable handler error", + tag.NewStringTag("control_queue", controlQueue), + tag.Error(nexusErr)) + metrics.WorkerCommandsSent.With(metricsHandler).Record(1, metrics.OutcomeTag("non_retryable_error")) + return nil + } + + logger.Warn("Worker commands transport failure", + tag.NewStringTag("control_queue", controlQueue), + tag.Error(nexusErr)) + metrics.WorkerCommandsSent.With(metricsHandler).Record(1, metrics.OutcomeTag("transport_error")) + return nexusErr + } + + // Worker-returned failure (ApplicationError, CanceledError, etc.). The worker received + // and processed the request but returned an error. Permanent — the worker contract + // requires success for all defined commands, so this indicates a bug or version + // incompatibility. Retrying won't help. + logger.Error("Worker returned failure for worker commands", + tag.NewStringTag("control_queue", controlQueue), + tag.Error(nexusErr)) + metrics.WorkerCommandsSent.With(metricsHandler).Record(1, metrics.OutcomeTag("worker_error")) + return nil +} diff --git a/common/workercommands/dispatch_test.go b/common/workercommands/dispatch_test.go new file mode 100644 index 00000000000..b53f83d35da --- /dev/null +++ b/common/workercommands/dispatch_test.go @@ -0,0 +1,321 @@ +package workercommands + +import ( + "context" + "errors" + "testing" + + "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" + workerservicepb "go.temporal.io/api/nexusservices/workerservice/v1" + workerpb "go.temporal.io/api/worker/v1" + "go.temporal.io/sdk/temporal" + "go.temporal.io/server/api/matchingservice/v1" + "go.temporal.io/server/api/matchingservicemock/v1" + "go.temporal.io/server/common/log" + "go.temporal.io/server/common/metrics" + "go.temporal.io/server/common/metrics/metricstest" + "go.uber.org/mock/gomock" + "google.golang.org/protobuf/proto" +) + +func testCommands() []*workerpb.WorkerCommand { + return []*workerpb.WorkerCommand{ + {Type: &workerpb.WorkerCommand_CancelActivity{ + CancelActivity: &workerpb.CancelActivityCommand{TaskToken: []byte("token1")}, + }}, + } +} + +func requireMetric(t *testing.T, snap map[string][]*metricstest.CapturedRecording, expectedOutcome string) { + t.Helper() + recordings := snap[metrics.WorkerCommandsSent.Name()] + require.Len(t, recordings, 1, "expected exactly 1 metric recording") + require.Equal(t, expectedOutcome, recordings[0].Tags["outcome"]) +} + +// ---- DispatchToWorker tests ---- + +func TestDispatchToWorker_Success(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := matchingservicemock.NewMockMatchingServiceClient(ctrl) + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + var capturedReq *matchingservice.DispatchNexusTaskRequest + mockClient.EXPECT().DispatchNexusTask(gomock.Any(), gomock.Any()).DoAndReturn( + func(_ context.Context, req *matchingservice.DispatchNexusTaskRequest, _ ...any) (*matchingservice.DispatchNexusTaskResponse, error) { + capturedReq = req + return syncSuccessResponse(), nil + }) + + err := DispatchToWorker( + context.Background(), mockClient, metricsHandler, log.NewNoopLogger(), + "test-ns-id", "test-control-queue", testCommands(), + ) + require.NoError(t, err) + requireMetric(t, capture.Snapshot(), "success") + + // Verify request shape. + require.NotNil(t, capturedReq) + require.Equal(t, "test-ns-id", capturedReq.NamespaceId) + require.Equal(t, "test-control-queue", capturedReq.TaskQueue.Name) + require.Equal(t, enumspb.TASK_QUEUE_KIND_WORKER_COMMANDS, capturedReq.TaskQueue.Kind) + + // Verify Nexus request contents. + startOp := capturedReq.Request.GetStartOperation() + require.NotNil(t, startOp) + require.Equal(t, ServiceName, startOp.Service) + require.Equal(t, OperationName, startOp.Operation) + + // Verify payload is binary/protobuf-encoded ExecuteCommandsRequest. + require.Equal(t, "binary/protobuf", string(startOp.Payload.Metadata["encoding"])) + var decoded workerservicepb.ExecuteCommandsRequest + require.NoError(t, proto.Unmarshal(startOp.Payload.Data, &decoded)) + require.Len(t, decoded.Commands, 1) + require.NotNil(t, decoded.Commands[0].GetCancelActivity()) + require.Equal(t, []byte("token1"), decoded.Commands[0].GetCancelActivity().TaskToken) +} + +func TestDispatchToWorker_RPCError(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := matchingservicemock.NewMockMatchingServiceClient(ctrl) + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + mockClient.EXPECT().DispatchNexusTask(gomock.Any(), gomock.Any()).Return( + nil, errors.New("connection refused")) + + err := DispatchToWorker( + context.Background(), mockClient, metricsHandler, log.NewNoopLogger(), + "test-ns-id", "test-control-queue", testCommands(), + ) + require.Error(t, err) + require.Contains(t, err.Error(), "connection refused") + requireMetric(t, capture.Snapshot(), "rpc_error") +} + +func TestDispatchToWorker_UpstreamTimeout(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := matchingservicemock.NewMockMatchingServiceClient(ctrl) + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + mockClient.EXPECT().DispatchNexusTask(gomock.Any(), gomock.Any()).Return( + &matchingservice.DispatchNexusTaskResponse{ + Outcome: &matchingservice.DispatchNexusTaskResponse_RequestTimeout{ + RequestTimeout: &matchingservice.DispatchNexusTaskResponse_Timeout{}, + }, + }, nil) + + err := DispatchToWorker( + context.Background(), mockClient, metricsHandler, log.NewNoopLogger(), + "test-ns-id", "test-control-queue", testCommands(), + ) + require.Error(t, err) + + var he *nexus.HandlerError + require.ErrorAs(t, err, &he) + require.Equal(t, nexus.HandlerErrorTypeUpstreamTimeout, he.Type) + requireMetric(t, capture.Snapshot(), "no_poller") +} + +func TestDispatchToWorker_WorkerFailure_SwallowedAsPermanent(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := matchingservicemock.NewMockMatchingServiceClient(ctrl) + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + mockClient.EXPECT().DispatchNexusTask(gomock.Any(), gomock.Any()).Return( + workerFailureResponse("worker bug"), nil) + + err := DispatchToWorker( + context.Background(), mockClient, metricsHandler, log.NewNoopLogger(), + "test-ns-id", "test-control-queue", testCommands(), + ) + require.NoError(t, err, "worker-returned failures are permanent and should be swallowed") + requireMetric(t, capture.Snapshot(), "worker_error") +} + +func TestDispatchToWorker_NonRetryableHandlerError_SwallowedAsPermanent(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := matchingservicemock.NewMockMatchingServiceClient(ctrl) + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + mockClient.EXPECT().DispatchNexusTask(gomock.Any(), gomock.Any()).Return( + handlerErrorResponse(nexus.HandlerErrorTypeBadRequest, "bad request"), nil) + + err := DispatchToWorker( + context.Background(), mockClient, metricsHandler, log.NewNoopLogger(), + "test-ns-id", "test-control-queue", testCommands(), + ) + require.NoError(t, err, "non-retryable handler errors should be swallowed") + requireMetric(t, capture.Snapshot(), "non_retryable_error") +} + +func TestDispatchToWorker_RetryableTransportError(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := matchingservicemock.NewMockMatchingServiceClient(ctrl) + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + mockClient.EXPECT().DispatchNexusTask(gomock.Any(), gomock.Any()).Return( + handlerErrorResponse(nexus.HandlerErrorTypeInternal, "something broke"), nil) + + err := DispatchToWorker( + context.Background(), mockClient, metricsHandler, log.NewNoopLogger(), + "test-ns-id", "test-control-queue", testCommands(), + ) + require.Error(t, err, "internal handler errors are retryable") + requireMetric(t, capture.Snapshot(), "transport_error") +} + +func TestDispatchToWorker_MultipleCommands(t *testing.T) { + ctrl := gomock.NewController(t) + mockClient := matchingservicemock.NewMockMatchingServiceClient(ctrl) + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + var capturedReq *matchingservice.DispatchNexusTaskRequest + mockClient.EXPECT().DispatchNexusTask(gomock.Any(), gomock.Any()).DoAndReturn( + func(_ context.Context, req *matchingservice.DispatchNexusTaskRequest, _ ...any) (*matchingservice.DispatchNexusTaskResponse, error) { + capturedReq = req + return syncSuccessResponse(), nil + }) + + commands := []*workerpb.WorkerCommand{ + {Type: &workerpb.WorkerCommand_CancelActivity{ + CancelActivity: &workerpb.CancelActivityCommand{TaskToken: []byte("token1")}, + }}, + {Type: &workerpb.WorkerCommand_CancelActivity{ + CancelActivity: &workerpb.CancelActivityCommand{TaskToken: []byte("token2")}, + }}, + } + + err := DispatchToWorker( + context.Background(), mockClient, metricsHandler, log.NewNoopLogger(), + "test-ns-id", "test-control-queue", commands, + ) + require.NoError(t, err) + + // Verify both commands are in the payload. + startOp := capturedReq.Request.GetStartOperation() + var decoded workerservicepb.ExecuteCommandsRequest + require.NoError(t, proto.Unmarshal(startOp.Payload.Data, &decoded)) + require.Len(t, decoded.Commands, 2) + require.Equal(t, []byte("token1"), decoded.Commands[0].GetCancelActivity().TaskToken) + require.Equal(t, []byte("token2"), decoded.Commands[1].GetCancelActivity().TaskToken) +} + +// ---- HandleDispatchError tests ---- + +func TestHandleDispatchError_WorkerError_ReturnsNil(t *testing.T) { + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + workerErr := temporal.NewApplicationError("worker bug", "SomeType", nil) + err := HandleDispatchError(workerErr, "test-control-queue", metricsHandler, log.NewNoopLogger()) + require.NoError(t, err, "worker-returned errors are permanent and should be swallowed") + requireMetric(t, capture.Snapshot(), "worker_error") +} + +func TestHandleDispatchError_UpstreamTimeout_ReturnsError(t *testing.T) { + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + handlerErr := nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUpstreamTimeout, "upstream timeout") + err := HandleDispatchError(handlerErr, "test-control-queue", metricsHandler, log.NewNoopLogger()) + require.Error(t, err, "upstream timeout should be retried") + + var he *nexus.HandlerError + require.ErrorAs(t, err, &he) + require.Equal(t, nexus.HandlerErrorTypeUpstreamTimeout, he.Type) + requireMetric(t, capture.Snapshot(), "no_poller") +} + +func TestHandleDispatchError_NonRetryableHandlerError_ReturnsNil(t *testing.T) { + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + handlerErr := nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "bad request") + err := HandleDispatchError(handlerErr, "test-control-queue", metricsHandler, log.NewNoopLogger()) + require.NoError(t, err, "non-retryable handler errors should be swallowed") + requireMetric(t, capture.Snapshot(), "non_retryable_error") +} + +func TestHandleDispatchError_RetryableHandlerError_ReturnsError(t *testing.T) { + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + handlerErr := nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "something broke") + err := HandleDispatchError(handlerErr, "test-control-queue", metricsHandler, log.NewNoopLogger()) + require.Error(t, err, "retryable handler errors should be returned") + requireMetric(t, capture.Snapshot(), "transport_error") +} + +func TestHandleDispatchError_CanceledError_ReturnsNil(t *testing.T) { + metricsHandler := metricstest.NewCaptureHandler() + capture := metricsHandler.StartCapture() + defer metricsHandler.StopCapture(capture) + + canceledErr := temporal.NewCanceledError() + err := HandleDispatchError(canceledErr, "test-control-queue", metricsHandler, log.NewNoopLogger()) + require.NoError(t, err, "CanceledError is permanent and should be swallowed") + requireMetric(t, capture.Snapshot(), "worker_error") +} + +// ---- helpers ---- + +func syncSuccessResponse() *matchingservice.DispatchNexusTaskResponse { + return &matchingservice.DispatchNexusTaskResponse{ + Outcome: &matchingservice.DispatchNexusTaskResponse_Response{ + Response: &nexuspb.Response{ + Variant: &nexuspb.Response_StartOperation{ + StartOperation: &nexuspb.StartOperationResponse{ + Variant: &nexuspb.StartOperationResponse_SyncSuccess{ + SyncSuccess: &nexuspb.StartOperationResponse_Sync{}, + }, + }, + }, + }, + }, + } +} + +func workerFailureResponse(message string) *matchingservice.DispatchNexusTaskResponse { + failure := temporal.GetDefaultFailureConverter().ErrorToFailure( + temporal.NewApplicationError(message, "TestErrorType", nil), + ) + return &matchingservice.DispatchNexusTaskResponse{ + Outcome: &matchingservice.DispatchNexusTaskResponse_Failure{ + Failure: failure, + }, + } +} + +func handlerErrorResponse(errType nexus.HandlerErrorType, message string) *matchingservice.DispatchNexusTaskResponse { + // Simulate the response shape that DispatchResponseToError converts into a HandlerError. + // For handler-level errors, matching returns a Response with a StartOperation failure + // containing the error type in the failure metadata. + handlerErr := nexus.NewHandlerErrorf(errType, "%s", message) + failure := temporal.GetDefaultFailureConverter().ErrorToFailure(handlerErr) + return &matchingservice.DispatchNexusTaskResponse{ + Outcome: &matchingservice.DispatchNexusTaskResponse_Failure{ + Failure: failure, + }, + } +} diff --git a/service/history/worker_commands_task_dispatcher.go b/service/history/worker_commands_task_dispatcher.go index 847ddfbbc51..20ffd3ae46d 100644 --- a/service/history/worker_commands_task_dispatcher.go +++ b/service/history/worker_commands_task_dispatcher.go @@ -2,37 +2,21 @@ package history import ( "context" - "errors" - "fmt" "time" - "github.com/nexus-rpc/sdk-go/nexus" - commonpb "go.temporal.io/api/common/v1" - enumspb "go.temporal.io/api/enums/v1" - nexuspb "go.temporal.io/api/nexus/v1" - workerservicepb "go.temporal.io/api/nexusservices/workerservice/v1" - taskqueuepb "go.temporal.io/api/taskqueue/v1" - "go.temporal.io/server/api/matchingservice/v1" "go.temporal.io/server/common/debug" "go.temporal.io/server/common/log" "go.temporal.io/server/common/log/tag" "go.temporal.io/server/common/metrics" "go.temporal.io/server/common/resource" + "go.temporal.io/server/common/workercommands" "go.temporal.io/server/service/history/configs" "go.temporal.io/server/service/history/tasks" - "google.golang.org/protobuf/proto" ) const ( workerCommandsTaskTimeout = time.Second * 10 * debug.TimeoutMultiplier workerCommandsMaxTaskAttempt = 3 - - // Nexus service and operation names for worker commands. - // TODO: Replace with workerservicepb.WorkerService.ServiceName and - // workerservicepb.WorkerService.ExecuteCommands.Name() once the Nexus service - // descriptor is published in go.temporal.io/api. - workerCommandsServiceName = "temporal.api.nexusservices.workerservice.v1.WorkerService" - workerCommandsOperationName = "ExecuteCommands" ) // workerCommandsTaskDispatcher dispatches worker commands to workers via Nexus. @@ -108,102 +92,13 @@ func (d *workerCommandsTaskDispatcher) execute( ctx, cancel := context.WithTimeout(ctx, workerCommandsTaskTimeout) defer cancel() - return d.dispatchToWorker(ctx, task) -} - -func (d *workerCommandsTaskDispatcher) dispatchToWorker( - ctx context.Context, - task *tasks.WorkerCommandsTask, -) error { - request := &workerservicepb.ExecuteCommandsRequest{ - Commands: task.Commands, - } - // Encode as binary/protobuf using the standard Temporal payload format. - // Worker commands are handled directly by SDK Core (not by lang-SDK Nexus handlers), - // so we use binary/protobuf which Core can decode natively via prost. The standard - // payload.Encode() uses json/protobuf encoding, which Core does not support because - // it normally delegates Nexus payload deserialization to the lang SDK. - requestData, err := proto.Marshal(request) - if err != nil { - return fmt.Errorf("failed to encode worker commands request: %w", err) - } - requestPayload := &commonpb.Payload{ - Metadata: map[string][]byte{ - "encoding": []byte("binary/protobuf"), - }, - Data: requestData, - } - - nexusRequest := &nexuspb.Request{ - Header: map[string]string{}, - Variant: &nexuspb.Request_StartOperation{ - StartOperation: &nexuspb.StartOperationRequest{ - Service: workerCommandsServiceName, - Operation: workerCommandsOperationName, - Payload: requestPayload, - }, - }, - } - - resp, err := d.matchingClient.DispatchNexusTask(ctx, &matchingservice.DispatchNexusTaskRequest{ - NamespaceId: task.NamespaceID, - TaskQueue: &taskqueuepb.TaskQueue{ - Name: task.Destination, - Kind: enumspb.TASK_QUEUE_KIND_WORKER_COMMANDS, - }, - Request: nexusRequest, - }) - if err != nil { - metrics.WorkerCommandsSent.With(d.metricsHandler).Record(1, metrics.OutcomeTag("rpc_error")) - return fmt.Errorf("failed to dispatch worker commands to control queue %s: %w", task.Destination, err) - } - - nexusErr := dispatchResponseToError(resp) - if nexusErr == nil { - metrics.WorkerCommandsSent.With(d.metricsHandler).Record(1, metrics.OutcomeTag("success")) - return nil - } - - return d.handleError(nexusErr, task) -} - -func (d *workerCommandsTaskDispatcher) handleError(nexusErr error, task *tasks.WorkerCommandsTask) error { - var handlerErr *nexus.HandlerError - if errors.As(nexusErr, &handlerErr) { - // Handler-level error (transport, timeout, internal). These are constructed by - // dispatchResponseToError for non-worker-returned failures. - if handlerErr.Type == nexus.HandlerErrorTypeUpstreamTimeout { - d.logger.Warn("No worker polling control queue", - tag.NewStringTag("control_queue", task.Destination)) - metrics.WorkerCommandsSent.With(d.metricsHandler).Record(1, metrics.OutcomeTag("no_poller")) - return nexusErr - } - - if !handlerErr.Retryable() { - d.logger.Error("Worker commands non-retryable handler error", - tag.NewStringTag("control_queue", task.Destination), - tag.Error(nexusErr)) - metrics.WorkerCommandsSent.With(d.metricsHandler).Record(1, metrics.OutcomeTag("non_retryable_error")) - return nil - } - - d.logger.Warn("Worker commands transport failure", - tag.NewStringTag("control_queue", task.Destination), - tag.Error(nexusErr)) - metrics.WorkerCommandsSent.With(d.metricsHandler).Record(1, metrics.OutcomeTag("transport_error")) - return nexusErr - } - - // Worker-returned failure (ApplicationError, CanceledError, etc.). The worker received - // and processed the request but returned an error. Permanent — the worker contract - // requires success for all defined commands, so this indicates a bug or version - // incompatibility. Retrying won't help. - d.logger.Error("Worker returned failure for worker commands", - tag.WorkflowID(task.WorkflowID), - tag.WorkflowRunID(task.RunID), - tag.NewStringTag("control_queue", task.Destination), - tag.NewInt("command_count", len(task.Commands)), - tag.Error(nexusErr)) - metrics.WorkerCommandsSent.With(d.metricsHandler).Record(1, metrics.OutcomeTag("worker_error")) - return nil + return workercommands.DispatchToWorker( + ctx, + d.matchingClient, + d.metricsHandler, + d.logger, + task.NamespaceID, + task.Destination, + task.Commands, + ) } diff --git a/service/history/worker_commands_task_dispatcher_test.go b/service/history/worker_commands_task_dispatcher_test.go index 36e94996dc9..e81badc5225 100644 --- a/service/history/worker_commands_task_dispatcher_test.go +++ b/service/history/worker_commands_task_dispatcher_test.go @@ -10,7 +10,6 @@ import ( enumspb "go.temporal.io/api/enums/v1" nexuspb "go.temporal.io/api/nexus/v1" workerpb "go.temporal.io/api/worker/v1" - "go.temporal.io/sdk/temporal" "go.temporal.io/server/api/matchingservice/v1" "go.temporal.io/server/api/matchingservicemock/v1" "go.temporal.io/server/common/definition" @@ -234,79 +233,4 @@ func TestExecute_UpstreamTimeout(t *testing.T) { requireMetricValue(t, capture.Snapshot(), "no_poller") } -func TestHandleError_WorkerError_ReturnNil(t *testing.T) { - metricsHandler := metricstest.NewCaptureHandler() - capture := metricsHandler.StartCapture() - defer metricsHandler.StopCapture(capture) - - d := &workerCommandsTaskDispatcher{ - metricsHandler: metricsHandler, - logger: log.NewNoopLogger(), - } - - // Worker-returned errors (ApplicationError, CanceledError) are permanent. - workerErr := temporal.NewApplicationError("worker bug", "SomeType", nil) - task := testWorkerCommandsTask() - err := d.handleError(workerErr, task) - require.NoError(t, err, "worker-returned errors are permanent and should be swallowed") - - requireMetricValue(t, capture.Snapshot(), "worker_error") -} - -func TestHandleError_UpstreamTimeout_ReturnRetryable(t *testing.T) { - metricsHandler := metricstest.NewCaptureHandler() - capture := metricsHandler.StartCapture() - defer metricsHandler.StopCapture(capture) - - d := &workerCommandsTaskDispatcher{ - metricsHandler: metricsHandler, - logger: log.NewNoopLogger(), - } - - handlerErr := nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUpstreamTimeout, "upstream timeout") - task := testWorkerCommandsTask() - err := d.handleError(handlerErr, task) - require.Error(t, err, "upstream timeout should be retried") - - var he *nexus.HandlerError - require.ErrorAs(t, err, &he) - require.Equal(t, nexus.HandlerErrorTypeUpstreamTimeout, he.Type) - - requireMetricValue(t, capture.Snapshot(), "no_poller") -} - -func TestHandleError_NonRetryableHandlerError_ReturnNil(t *testing.T) { - metricsHandler := metricstest.NewCaptureHandler() - capture := metricsHandler.StartCapture() - defer metricsHandler.StopCapture(capture) - - d := &workerCommandsTaskDispatcher{ - metricsHandler: metricsHandler, - logger: log.NewNoopLogger(), - } - - handlerErr := nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "bad request") - task := testWorkerCommandsTask() - err := d.handleError(handlerErr, task) - require.NoError(t, err, "non-retryable handler errors should be swallowed") - - requireMetricValue(t, capture.Snapshot(), "non_retryable_error") -} - -func TestHandleError_OtherHandlerError_ReturnRetryable(t *testing.T) { - metricsHandler := metricstest.NewCaptureHandler() - capture := metricsHandler.StartCapture() - defer metricsHandler.StopCapture(capture) - - d := &workerCommandsTaskDispatcher{ - metricsHandler: metricsHandler, - logger: log.NewNoopLogger(), - } - - handlerErr := nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "something broke") - task := testWorkerCommandsTask() - err := d.handleError(handlerErr, task) - require.Error(t, err, "transport errors should be retried") - - requireMetricValue(t, capture.Snapshot(), "transport_error") -} +// HandleDispatchError tests are in common/workercommands/dispatch_test.go.