Nexus CHASM async completion (1/2) (#9951)

## What changed?

Support Nexus async completion in CHASM.

PS: completion-before-start will be handled in a follow-up PR.

## How did you test it?
- [ ] built
- [ ] run locally and tested manually
- [ ] covered by existing tests
- [x] added new unit test(s)
- [x] added new functional test(s)

## Potential risks

Is behind feature flag.
This commit is contained in:
Stephan Behnke
2026-04-21 08:46:15 -07:00
committed by GitHub
parent 813d1fab57
commit b78391b445
21 changed files with 925 additions and 388 deletions

View File

@@ -8592,7 +8592,9 @@ type CompleteNexusOperationChasmRequest struct {
// *CompleteNexusOperationChasmRequest_Failure
Outcome isCompleteNexusOperationChasmRequest_Outcome `protobuf_oneof:"outcome"`
// Time when the operation was closed.
CloseTime *timestamppb.Timestamp `protobuf:"bytes,4,opt,name=close_time,json=closeTime,proto3" json:"close_time,omitempty"`
CloseTime *timestamppb.Timestamp `protobuf:"bytes,4,opt,name=close_time,json=closeTime,proto3" json:"close_time,omitempty"`
// Links from the Nexus completion callback (e.g. references to the handler workflow).
Links []*v14.Link `protobuf:"bytes,5,rep,name=links,proto3" json:"links,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -8666,6 +8668,13 @@ func (x *CompleteNexusOperationChasmRequest) GetCloseTime() *timestamppb.Timesta
return nil
}
func (x *CompleteNexusOperationChasmRequest) GetLinks() []*v14.Link {
if x != nil {
return x.Links
}
return nil
}
type isCompleteNexusOperationChasmRequest_Outcome interface {
isCompleteNexusOperationChasmRequest_Outcome()
}
@@ -11067,7 +11076,7 @@ const file_temporal_server_api_historyservice_v1_request_response_proto_rawDesc
"\x10ListTasksRequest\x12V\n" +
"\arequest\x18\x01 \x01(\v2<.temporal.server.api.adminservice.v1.ListHistoryTasksRequestR\arequest:\x16\x92\xc4\x03\x12\x1a\x10request.shard_id\"n\n" +
"\x11ListTasksResponse\x12Y\n" +
"\bresponse\x18\x01 \x01(\v2=.temporal.server.api.adminservice.v1.ListHistoryTasksResponseR\bresponse\"\xdd\x02\n" +
"\bresponse\x18\x01 \x01(\v2=.temporal.server.api.adminservice.v1.ListHistoryTasksResponseR\bresponse\"\x91\x03\n" +
"\"CompleteNexusOperationChasmRequest\x12V\n" +
"\n" +
"completion\x18\x01 \x01(\v26.temporal.server.api.token.v1.NexusOperationCompletionR\n" +
@@ -11075,7 +11084,8 @@ const file_temporal_server_api_historyservice_v1_request_response_proto_rawDesc
"\asuccess\x18\x02 \x01(\v2\x1f.temporal.api.common.v1.PayloadH\x00R\asuccess\x12<\n" +
"\afailure\x18\x03 \x01(\v2 .temporal.api.failure.v1.FailureH\x00R\afailure\x129\n" +
"\n" +
"close_time\x18\x04 \x01(\v2\x1a.google.protobuf.TimestampR\tcloseTime:\x1e\x92\xc4\x03\x1aB\x18completion.component_refB\t\n" +
"close_time\x18\x04 \x01(\v2\x1a.google.protobuf.TimestampR\tcloseTime\x122\n" +
"\x05links\x18\x05 \x03(\v2\x1c.temporal.api.common.v1.LinkR\x05links:\x1e\x92\xc4\x03\x1aB\x18completion.component_refB\t\n" +
"\aoutcome\"%\n" +
"#CompleteNexusOperationChasmResponse\"\xe0\x03\n" +
"\x1dCompleteNexusOperationRequest\x12V\n" +
@@ -11685,48 +11695,49 @@ var file_temporal_server_api_historyservice_v1_request_response_proto_depIdxs =
262, // 220: temporal.server.api.historyservice.v1.CompleteNexusOperationChasmRequest.success:type_name -> temporal.api.common.v1.Payload
173, // 221: temporal.server.api.historyservice.v1.CompleteNexusOperationChasmRequest.failure:type_name -> temporal.api.failure.v1.Failure
171, // 222: temporal.server.api.historyservice.v1.CompleteNexusOperationChasmRequest.close_time:type_name -> google.protobuf.Timestamp
261, // 223: temporal.server.api.historyservice.v1.CompleteNexusOperationRequest.completion:type_name -> temporal.server.api.token.v1.NexusOperationCompletion
262, // 224: temporal.server.api.historyservice.v1.CompleteNexusOperationRequest.success:type_name -> temporal.api.common.v1.Payload
263, // 225: temporal.server.api.historyservice.v1.CompleteNexusOperationRequest.failure:type_name -> temporal.api.nexus.v1.Failure
171, // 226: temporal.server.api.historyservice.v1.CompleteNexusOperationRequest.start_time:type_name -> google.protobuf.Timestamp
185, // 227: temporal.server.api.historyservice.v1.CompleteNexusOperationRequest.links:type_name -> temporal.api.common.v1.Link
264, // 228: temporal.server.api.historyservice.v1.InvokeStateMachineMethodRequest.ref:type_name -> temporal.server.api.persistence.v1.StateMachineRef
265, // 229: temporal.server.api.historyservice.v1.DeepHealthCheckResponse.state:type_name -> temporal.server.api.enums.v1.HealthState
266, // 230: temporal.server.api.historyservice.v1.DeepHealthCheckResponse.checks:type_name -> temporal.server.api.health.v1.HealthCheck
186, // 231: temporal.server.api.historyservice.v1.SyncWorkflowStateRequest.execution:type_name -> temporal.api.common.v1.WorkflowExecution
188, // 232: temporal.server.api.historyservice.v1.SyncWorkflowStateRequest.versioned_transition:type_name -> temporal.server.api.persistence.v1.VersionedTransition
192, // 233: temporal.server.api.historyservice.v1.SyncWorkflowStateRequest.version_histories:type_name -> temporal.server.api.history.v1.VersionHistories
267, // 234: temporal.server.api.historyservice.v1.SyncWorkflowStateResponse.versioned_transition_artifact:type_name -> temporal.server.api.replication.v1.VersionedTransitionArtifact
268, // 235: temporal.server.api.historyservice.v1.UpdateActivityOptionsRequest.update_request:type_name -> temporal.api.workflowservice.v1.UpdateActivityOptionsRequest
269, // 236: temporal.server.api.historyservice.v1.UpdateActivityOptionsResponse.activity_options:type_name -> temporal.api.activity.v1.ActivityOptions
270, // 237: temporal.server.api.historyservice.v1.PauseActivityRequest.frontend_request:type_name -> temporal.api.workflowservice.v1.PauseActivityRequest
271, // 238: temporal.server.api.historyservice.v1.UnpauseActivityRequest.frontend_request:type_name -> temporal.api.workflowservice.v1.UnpauseActivityRequest
272, // 239: temporal.server.api.historyservice.v1.ResetActivityRequest.frontend_request:type_name -> temporal.api.workflowservice.v1.ResetActivityRequest
273, // 240: temporal.server.api.historyservice.v1.UpdateWorkflowExecutionOptionsRequest.update_request:type_name -> temporal.api.workflowservice.v1.UpdateWorkflowExecutionOptionsRequest
274, // 241: temporal.server.api.historyservice.v1.UpdateWorkflowExecutionOptionsResponse.workflow_execution_options:type_name -> temporal.api.workflow.v1.WorkflowExecutionOptions
275, // 242: temporal.server.api.historyservice.v1.PauseWorkflowExecutionRequest.pause_request:type_name -> temporal.api.workflowservice.v1.PauseWorkflowExecutionRequest
276, // 243: temporal.server.api.historyservice.v1.UnpauseWorkflowExecutionRequest.unpause_request:type_name -> temporal.api.workflowservice.v1.UnpauseWorkflowExecutionRequest
277, // 244: temporal.server.api.historyservice.v1.StartNexusOperationRequest.request:type_name -> temporal.api.nexus.v1.StartOperationRequest
278, // 245: temporal.server.api.historyservice.v1.StartNexusOperationResponse.response:type_name -> temporal.api.nexus.v1.StartOperationResponse
279, // 246: temporal.server.api.historyservice.v1.CancelNexusOperationRequest.request:type_name -> temporal.api.nexus.v1.CancelOperationRequest
280, // 247: temporal.server.api.historyservice.v1.CancelNexusOperationResponse.response:type_name -> temporal.api.nexus.v1.CancelOperationResponse
1, // 248: temporal.server.api.historyservice.v1.ExecuteMultiOperationRequest.Operation.start_workflow:type_name -> temporal.server.api.historyservice.v1.StartWorkflowExecutionRequest
105, // 249: temporal.server.api.historyservice.v1.ExecuteMultiOperationRequest.Operation.update_workflow:type_name -> temporal.server.api.historyservice.v1.UpdateWorkflowExecutionRequest
2, // 250: temporal.server.api.historyservice.v1.ExecuteMultiOperationResponse.Response.start_workflow:type_name -> temporal.server.api.historyservice.v1.StartWorkflowExecutionResponse
106, // 251: temporal.server.api.historyservice.v1.ExecuteMultiOperationResponse.Response.update_workflow:type_name -> temporal.server.api.historyservice.v1.UpdateWorkflowExecutionResponse
281, // 252: temporal.server.api.historyservice.v1.RecordWorkflowTaskStartedResponse.QueriesEntry.value:type_name -> temporal.api.query.v1.WorkflowQuery
281, // 253: temporal.server.api.historyservice.v1.RecordWorkflowTaskStartedResponseWithRawHistory.QueriesEntry.value:type_name -> temporal.api.query.v1.WorkflowQuery
282, // 254: temporal.server.api.historyservice.v1.GetReplicationMessagesResponse.ShardMessagesEntry.value:type_name -> temporal.server.api.replication.v1.ReplicationMessages
98, // 255: temporal.server.api.historyservice.v1.ShardReplicationStatus.RemoteClustersEntry.value:type_name -> temporal.server.api.historyservice.v1.ShardReplicationStatusPerCluster
97, // 256: temporal.server.api.historyservice.v1.ShardReplicationStatus.HandoverNamespacesEntry.value:type_name -> temporal.server.api.historyservice.v1.HandoverNamespaceInfo
226, // 257: temporal.server.api.historyservice.v1.AddTasksRequest.Task.blob:type_name -> temporal.api.common.v1.DataBlob
283, // 258: temporal.server.api.historyservice.v1.routing:extendee -> google.protobuf.MessageOptions
0, // 259: temporal.server.api.historyservice.v1.routing:type_name -> temporal.server.api.historyservice.v1.RoutingOptions
260, // [260:260] is the sub-list for method output_type
260, // [260:260] is the sub-list for method input_type
259, // [259:260] is the sub-list for extension type_name
258, // [258:259] is the sub-list for extension extendee
0, // [0:258] is the sub-list for field type_name
185, // 223: temporal.server.api.historyservice.v1.CompleteNexusOperationChasmRequest.links:type_name -> temporal.api.common.v1.Link
261, // 224: temporal.server.api.historyservice.v1.CompleteNexusOperationRequest.completion:type_name -> temporal.server.api.token.v1.NexusOperationCompletion
262, // 225: temporal.server.api.historyservice.v1.CompleteNexusOperationRequest.success:type_name -> temporal.api.common.v1.Payload
263, // 226: temporal.server.api.historyservice.v1.CompleteNexusOperationRequest.failure:type_name -> temporal.api.nexus.v1.Failure
171, // 227: temporal.server.api.historyservice.v1.CompleteNexusOperationRequest.start_time:type_name -> google.protobuf.Timestamp
185, // 228: temporal.server.api.historyservice.v1.CompleteNexusOperationRequest.links:type_name -> temporal.api.common.v1.Link
264, // 229: temporal.server.api.historyservice.v1.InvokeStateMachineMethodRequest.ref:type_name -> temporal.server.api.persistence.v1.StateMachineRef
265, // 230: temporal.server.api.historyservice.v1.DeepHealthCheckResponse.state:type_name -> temporal.server.api.enums.v1.HealthState
266, // 231: temporal.server.api.historyservice.v1.DeepHealthCheckResponse.checks:type_name -> temporal.server.api.health.v1.HealthCheck
186, // 232: temporal.server.api.historyservice.v1.SyncWorkflowStateRequest.execution:type_name -> temporal.api.common.v1.WorkflowExecution
188, // 233: temporal.server.api.historyservice.v1.SyncWorkflowStateRequest.versioned_transition:type_name -> temporal.server.api.persistence.v1.VersionedTransition
192, // 234: temporal.server.api.historyservice.v1.SyncWorkflowStateRequest.version_histories:type_name -> temporal.server.api.history.v1.VersionHistories
267, // 235: temporal.server.api.historyservice.v1.SyncWorkflowStateResponse.versioned_transition_artifact:type_name -> temporal.server.api.replication.v1.VersionedTransitionArtifact
268, // 236: temporal.server.api.historyservice.v1.UpdateActivityOptionsRequest.update_request:type_name -> temporal.api.workflowservice.v1.UpdateActivityOptionsRequest
269, // 237: temporal.server.api.historyservice.v1.UpdateActivityOptionsResponse.activity_options:type_name -> temporal.api.activity.v1.ActivityOptions
270, // 238: temporal.server.api.historyservice.v1.PauseActivityRequest.frontend_request:type_name -> temporal.api.workflowservice.v1.PauseActivityRequest
271, // 239: temporal.server.api.historyservice.v1.UnpauseActivityRequest.frontend_request:type_name -> temporal.api.workflowservice.v1.UnpauseActivityRequest
272, // 240: temporal.server.api.historyservice.v1.ResetActivityRequest.frontend_request:type_name -> temporal.api.workflowservice.v1.ResetActivityRequest
273, // 241: temporal.server.api.historyservice.v1.UpdateWorkflowExecutionOptionsRequest.update_request:type_name -> temporal.api.workflowservice.v1.UpdateWorkflowExecutionOptionsRequest
274, // 242: temporal.server.api.historyservice.v1.UpdateWorkflowExecutionOptionsResponse.workflow_execution_options:type_name -> temporal.api.workflow.v1.WorkflowExecutionOptions
275, // 243: temporal.server.api.historyservice.v1.PauseWorkflowExecutionRequest.pause_request:type_name -> temporal.api.workflowservice.v1.PauseWorkflowExecutionRequest
276, // 244: temporal.server.api.historyservice.v1.UnpauseWorkflowExecutionRequest.unpause_request:type_name -> temporal.api.workflowservice.v1.UnpauseWorkflowExecutionRequest
277, // 245: temporal.server.api.historyservice.v1.StartNexusOperationRequest.request:type_name -> temporal.api.nexus.v1.StartOperationRequest
278, // 246: temporal.server.api.historyservice.v1.StartNexusOperationResponse.response:type_name -> temporal.api.nexus.v1.StartOperationResponse
279, // 247: temporal.server.api.historyservice.v1.CancelNexusOperationRequest.request:type_name -> temporal.api.nexus.v1.CancelOperationRequest
280, // 248: temporal.server.api.historyservice.v1.CancelNexusOperationResponse.response:type_name -> temporal.api.nexus.v1.CancelOperationResponse
1, // 249: temporal.server.api.historyservice.v1.ExecuteMultiOperationRequest.Operation.start_workflow:type_name -> temporal.server.api.historyservice.v1.StartWorkflowExecutionRequest
105, // 250: temporal.server.api.historyservice.v1.ExecuteMultiOperationRequest.Operation.update_workflow:type_name -> temporal.server.api.historyservice.v1.UpdateWorkflowExecutionRequest
2, // 251: temporal.server.api.historyservice.v1.ExecuteMultiOperationResponse.Response.start_workflow:type_name -> temporal.server.api.historyservice.v1.StartWorkflowExecutionResponse
106, // 252: temporal.server.api.historyservice.v1.ExecuteMultiOperationResponse.Response.update_workflow:type_name -> temporal.server.api.historyservice.v1.UpdateWorkflowExecutionResponse
281, // 253: temporal.server.api.historyservice.v1.RecordWorkflowTaskStartedResponse.QueriesEntry.value:type_name -> temporal.api.query.v1.WorkflowQuery
281, // 254: temporal.server.api.historyservice.v1.RecordWorkflowTaskStartedResponseWithRawHistory.QueriesEntry.value:type_name -> temporal.api.query.v1.WorkflowQuery
282, // 255: temporal.server.api.historyservice.v1.GetReplicationMessagesResponse.ShardMessagesEntry.value:type_name -> temporal.server.api.replication.v1.ReplicationMessages
98, // 256: temporal.server.api.historyservice.v1.ShardReplicationStatus.RemoteClustersEntry.value:type_name -> temporal.server.api.historyservice.v1.ShardReplicationStatusPerCluster
97, // 257: temporal.server.api.historyservice.v1.ShardReplicationStatus.HandoverNamespacesEntry.value:type_name -> temporal.server.api.historyservice.v1.HandoverNamespaceInfo
226, // 258: temporal.server.api.historyservice.v1.AddTasksRequest.Task.blob:type_name -> temporal.api.common.v1.DataBlob
283, // 259: temporal.server.api.historyservice.v1.routing:extendee -> google.protobuf.MessageOptions
0, // 260: temporal.server.api.historyservice.v1.routing:type_name -> temporal.server.api.historyservice.v1.RoutingOptions
261, // [261:261] is the sub-list for method output_type
261, // [261:261] is the sub-list for method input_type
260, // [260:261] is the sub-list for extension type_name
259, // [259:260] is the sub-list for extension extendee
0, // [0:259] is the sub-list for field type_name
}
func init() { file_temporal_server_api_historyservice_v1_request_response_proto_init() }

View File

@@ -606,7 +606,9 @@ type ChasmNexusCompletion struct {
CloseTime *timestamppb.Timestamp `protobuf:"bytes,3,opt,name=close_time,json=closeTime,proto3" json:"close_time,omitempty"`
// Request ID embedded in the NexusOperationScheduledEvent.
// Allows completing a started operation after a workflow has been reset.
RequestId string `protobuf:"bytes,4,opt,name=request_id,json=requestId,proto3" json:"request_id,omitempty"`
RequestId string `protobuf:"bytes,4,opt,name=request_id,json=requestId,proto3" json:"request_id,omitempty"`
// Links from the Nexus completion callback (e.g. references to the handler workflow).
Links []*v1.Link `protobuf:"bytes,5,rep,name=links,proto3" json:"links,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
@@ -680,6 +682,13 @@ func (x *ChasmNexusCompletion) GetRequestId() string {
return ""
}
func (x *ChasmNexusCompletion) GetLinks() []*v1.Link {
if x != nil {
return x.Links
}
return nil
}
type isChasmNexusCompletion_Outcome interface {
isChasmNexusCompletion_Outcome()
}
@@ -849,14 +858,15 @@ const file_temporal_server_api_persistence_v1_chasm_proto_rawDesc = "" +
"\farchetype_id\x18\x04 \x01(\rR\varchetypeId\x12}\n" +
"\x1eexecution_versioned_transition\x18\x05 \x01(\v27.temporal.server.api.persistence.v1.VersionedTransitionR\x1cexecutionVersionedTransition\x12%\n" +
"\x0ecomponent_path\x18\x06 \x03(\tR\rcomponentPath\x12\x8c\x01\n" +
"&component_initial_versioned_transition\x18\a \x01(\v27.temporal.server.api.persistence.v1.VersionedTransitionR#componentInitialVersionedTransition\"\xf6\x01\n" +
"&component_initial_versioned_transition\x18\a \x01(\v27.temporal.server.api.persistence.v1.VersionedTransitionR#componentInitialVersionedTransition\"\xaa\x02\n" +
"\x14ChasmNexusCompletion\x12;\n" +
"\asuccess\x18\x01 \x01(\v2\x1f.temporal.api.common.v1.PayloadH\x00R\asuccess\x12<\n" +
"\afailure\x18\x02 \x01(\v2 .temporal.api.failure.v1.FailureH\x00R\afailure\x129\n" +
"\n" +
"close_time\x18\x03 \x01(\v2\x1a.google.protobuf.TimestampR\tcloseTime\x12\x1d\n" +
"\n" +
"request_id\x18\x04 \x01(\tR\trequestIdB\t\n" +
"request_id\x18\x04 \x01(\tR\trequestId\x122\n" +
"\x05links\x18\x05 \x03(\v2\x1c.temporal.api.common.v1.LinkR\x05linksB\t\n" +
"\aoutcomeB6Z4go.temporal.io/server/api/persistence/v1;persistenceb\x06proto3"
var (
@@ -888,6 +898,7 @@ var file_temporal_server_api_persistence_v1_chasm_proto_goTypes = []any{
(*v1.Payload)(nil), // 12: temporal.api.common.v1.Payload
(*v11.Failure)(nil), // 13: temporal.api.failure.v1.Failure
(*timestamppb.Timestamp)(nil), // 14: google.protobuf.Timestamp
(*v1.Link)(nil), // 15: temporal.api.common.v1.Link
}
var file_temporal_server_api_persistence_v1_chasm_proto_depIdxs = []int32{
1, // 0: temporal.server.api.persistence.v1.ChasmNode.metadata:type_name -> temporal.server.api.persistence.v1.ChasmNodeMetadata
@@ -908,14 +919,15 @@ var file_temporal_server_api_persistence_v1_chasm_proto_depIdxs = []int32{
12, // 15: temporal.server.api.persistence.v1.ChasmNexusCompletion.success:type_name -> temporal.api.common.v1.Payload
13, // 16: temporal.server.api.persistence.v1.ChasmNexusCompletion.failure:type_name -> temporal.api.failure.v1.Failure
14, // 17: temporal.server.api.persistence.v1.ChasmNexusCompletion.close_time:type_name -> google.protobuf.Timestamp
14, // 18: temporal.server.api.persistence.v1.ChasmComponentAttributes.Task.scheduled_time:type_name -> google.protobuf.Timestamp
10, // 19: temporal.server.api.persistence.v1.ChasmComponentAttributes.Task.data:type_name -> temporal.api.common.v1.DataBlob
11, // 20: temporal.server.api.persistence.v1.ChasmComponentAttributes.Task.versioned_transition:type_name -> temporal.server.api.persistence.v1.VersionedTransition
21, // [21:21] is the sub-list for method output_type
21, // [21:21] is the sub-list for method input_type
21, // [21:21] is the sub-list for extension type_name
21, // [21:21] is the sub-list for extension extendee
0, // [0:21] is the sub-list for field type_name
15, // 18: temporal.server.api.persistence.v1.ChasmNexusCompletion.links:type_name -> temporal.api.common.v1.Link
14, // 19: temporal.server.api.persistence.v1.ChasmComponentAttributes.Task.scheduled_time:type_name -> google.protobuf.Timestamp
10, // 20: temporal.server.api.persistence.v1.ChasmComponentAttributes.Task.data:type_name -> temporal.api.common.v1.DataBlob
11, // 21: temporal.server.api.persistence.v1.ChasmComponentAttributes.Task.versioned_transition:type_name -> temporal.server.api.persistence.v1.VersionedTransition
22, // [22:22] is the sub-list for method output_type
22, // [22:22] is the sub-list for method input_type
22, // [22:22] is the sub-list for extension type_name
22, // [22:22] is the sub-list for extension extendee
0, // [0:22] is the sub-list for field type_name
}
func init() { file_temporal_server_api_persistence_v1_chasm_proto_init() }

View File

@@ -288,7 +288,11 @@ func UpdateWithStartExecution[C RootComponent, I any, O any](
// comment of the NewRef method in MutableContext.
//
// UpdateComponent applies updateFn to the component identified by the supplied component reference.
// It returns the result, along with the new component reference. opts are currently ignored.
// opts are currently ignored.
//
// It returns the result, along with the new component reference. The returned reference may be
// nil when updateFn deletes the component in the same transaction and the component is not the
// root component.
func UpdateComponent[C any, R []byte | ComponentRef, I any, O any](
ctx context.Context,
r R,

View File

@@ -6,6 +6,7 @@ import (
"github.com/nexus-rpc/sdk-go/nexus"
"go.temporal.io/api/serviceerror"
persistencespb "go.temporal.io/server/api/persistence/v1"
"go.temporal.io/server/common"
"go.temporal.io/server/common/cluster"
"go.temporal.io/server/common/log"
@@ -42,16 +43,40 @@ func routeSystemCallbackRequest(
logger.Error("failed to decode completion from token", tag.Error(err))
return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid callback token")
}
ns, err := namespaceRegistry.GetNamespaceByID(namespace.ID(completion.NamespaceId))
// Normalize to support two possible token shapes:
// - legacy HSM tokens carry namespace/workflow IDs directly
// - CHASM tokens carry an encoded component ref instead
namespaceID := completion.GetNamespaceId()
businessID := completion.GetWorkflowId()
if namespaceID == "" && len(completion.GetComponentRef()) > 0 {
ref := &persistencespb.ChasmComponentRef{}
if err := ref.Unmarshal(completion.GetComponentRef()); err != nil {
logger.Error("failed to decode CHASM component ref from callback token", tag.Error(err))
return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid callback token")
}
if ref.GetNamespaceId() == "" {
logger.Error("decoded CHASM component ref is missing namespace ID")
return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid callback token")
}
if ref.GetBusinessId() == "" {
logger.Error("decoded CHASM component ref is missing business ID")
return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid callback token")
}
namespaceID = ref.GetNamespaceId()
businessID = ref.GetBusinessId()
}
ns, err := namespaceRegistry.GetNamespaceByID(namespace.ID(namespaceID))
if err != nil {
logger.Error("failed to get namespace for nexus completion request", tag.WorkflowNamespaceID(completion.NamespaceId), tag.Error(err))
logger.Error("failed to get namespace for nexus completion request", tag.WorkflowNamespaceID(namespaceID), tag.Error(err))
var nfe *serviceerror.NamespaceNotFound
if errors.As(err, &nfe) {
return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeNotFound, "namespace %q not found", completion.NamespaceId)
return nil, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeNotFound, "namespace %q not found", namespaceID)
}
return nil, commonnexus.ConvertGRPCError(err, false)
}
clusterName := ns.ActiveClusterName(namespace.RoutingKey{ID: completion.GetWorkflowId()})
clusterName := ns.ActiveClusterName(namespace.RoutingKey{ID: businessID})
if clusterMetadata.GetCurrentClusterName() == clusterName {
frontendClient = localClient
} else {

View File

@@ -7,7 +7,6 @@ import (
"testing"
"github.com/nexus-rpc/sdk-go/nexus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.temporal.io/api/serviceerror"
persistencespb "go.temporal.io/server/api/persistence/v1"
@@ -129,8 +128,9 @@ func TestRouteRequest_SourceHeaderUnknownCluster(t *testing.T) {
func TestRouteSystemCallbackRequest_NilHeaders(t *testing.T) {
// When the request has nil headers, it should fall back to the local client.
var gotPath string
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, commonnexus.PathCompletionCallbackNoIdentifier, r.URL.Path)
gotPath = r.URL.Path
w.WriteHeader(http.StatusOK)
}))
defer ts.Close()
@@ -155,6 +155,7 @@ func TestRouteSystemCallbackRequest_NilHeaders(t *testing.T) {
require.NoError(t, err)
defer func() { _ = resp.Body.Close() }()
require.Equal(t, http.StatusOK, resp.StatusCode)
require.Equal(t, commonnexus.PathCompletionCallbackNoIdentifier, gotPath)
}
func TestRouteSystemCallbackRequest_InvalidToken(t *testing.T) {
@@ -209,6 +210,8 @@ func TestRouteSystemCallbackRequest_NamespaceNotFound(t *testing.T) {
tokenStr, err := tokenGen.Tokenize(&tokenspb.NexusOperationCompletion{
NamespaceId: "ns-id-1",
WorkflowId: "wf-1",
RunId: "run-1",
Ref: &persistencespb.StateMachineRef{},
})
require.NoError(t, err)
@@ -235,63 +238,150 @@ func TestRouteSystemCallbackRequest_NamespaceNotFound(t *testing.T) {
require.Equal(t, nexus.HandlerErrorTypeNotFound, handlerErr.Type)
}
func TestRouteSystemCallbackRequest_InvalidChasmComponentRef(t *testing.T) {
for _, tc := range []struct {
name string
ref *persistencespb.ChasmComponentRef
}{
{
name: "missing namespace id",
ref: &persistencespb.ChasmComponentRef{
BusinessId: "wf-1",
RunId: "run-1",
},
},
{
name: "missing business id",
ref: &persistencespb.ChasmComponentRef{
NamespaceId: "ns-id-1",
RunId: "run-1",
},
},
} {
t.Run(tc.name, func(t *testing.T) {
tokenGen := commonnexus.NewCallbackTokenGenerator()
ref, err := tc.ref.Marshal()
require.NoError(t, err)
tokenStr, err := tokenGen.Tokenize(&tokenspb.NexusOperationCompletion{ComponentRef: ref})
require.NoError(t, err)
r, err := http.NewRequest(http.MethodPost, commonnexus.SystemCallbackURL, nil)
require.NoError(t, err)
r.Header.Set(commonnexus.CallbackTokenHeader, tokenStr)
_, err = routeSystemCallbackRequest(
r,
nil,
nil,
nil,
tokenGen,
nil,
log.NewNoopLogger(),
)
require.Error(t, err)
var handlerErr *nexus.HandlerError
require.ErrorAs(t, err, &handlerErr)
require.Equal(t, nexus.HandlerErrorTypeBadRequest, handlerErr.Type)
require.Contains(t, handlerErr.Error(), "invalid callback token")
})
}
}
func TestRouteSystemCallbackRequest_Success(t *testing.T) {
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, commonnexus.PathCompletionCallbackNoIdentifier, r.URL.Path)
w.WriteHeader(http.StatusOK)
}))
defer ts.Close()
for _, tc := range []struct {
name string
completionToken func(*commonnexus.CallbackTokenGenerator) (string, error)
}{
{
name: "HSM",
completionToken: func(tokenGen *commonnexus.CallbackTokenGenerator) (string, error) {
return tokenGen.Tokenize(&tokenspb.NexusOperationCompletion{
// HSM sets the deprecated execution fields and ref.
NamespaceId: "ns-id-1",
WorkflowId: "wf-1",
RunId: "run-1",
Ref: &persistencespb.StateMachineRef{},
})
},
},
{
name: "CHASM",
completionToken: func(tokenGen *commonnexus.CallbackTokenGenerator) (string, error) {
ref, err := (&persistencespb.ChasmComponentRef{
NamespaceId: "ns-id-1",
BusinessId: "wf-1",
RunId: "run-1",
}).Marshal()
if err != nil {
return "", err
}
return tokenGen.Tokenize(&tokenspb.NexusOperationCompletion{
// CHASM sets ComponentRef.
ComponentRef: ref,
})
},
},
} {
t.Run(tc.name, func(t *testing.T) {
var gotPath string
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotPath = r.URL.Path
w.WriteHeader(http.StatusOK)
}))
defer ts.Close()
ctrl := gomock.NewController(t)
clusterMeta := cluster.NewMockMetadata(ctrl)
nsRegistry := namespace.NewMockRegistry(ctrl)
ctrl := gomock.NewController(t)
clusterMeta := cluster.NewMockMetadata(ctrl)
nsRegistry := namespace.NewMockRegistry(ctrl)
tokenGen := commonnexus.NewCallbackTokenGenerator()
tokenStr, err := tokenGen.Tokenize(&tokenspb.NexusOperationCompletion{
NamespaceId: "ns-id-1",
WorkflowId: "wf-1",
})
require.NoError(t, err)
tokenGen := commonnexus.NewCallbackTokenGenerator()
tokenStr, err := tc.completionToken(tokenGen)
require.NoError(t, err)
testNS := namespace.NewLocalNamespaceForTest(
&persistencespb.NamespaceInfo{Id: "ns-id-1", Name: "test-ns"},
nil,
"cluster-A",
)
nsRegistry.EXPECT().GetNamespaceByID(namespace.ID("ns-id-1")).Return(testNS, nil)
testNS := namespace.NewLocalNamespaceForTest(
&persistencespb.NamespaceInfo{Id: "ns-id-1", Name: "test-ns"},
nil,
"cluster-A",
)
nsRegistry.EXPECT().GetNamespaceByID(namespace.ID("ns-id-1")).Return(testNS, nil)
// httpClientCache.Get will fail for "cluster-A", so it falls back to localClient.
clusterMeta.EXPECT().GetCurrentClusterName().Return("cluster-A").AnyTimes()
clusterMeta.EXPECT().GetAllClusterInfo().Return(map[string]cluster.ClusterInformation{}).AnyTimes()
clusterMeta.EXPECT().RegisterMetadataChangeCallback(gomock.Any(), gomock.Any())
// httpClientCache.Get will fail for "cluster-A", so it falls back to localClient.
clusterMeta.EXPECT().GetCurrentClusterName().Return("cluster-A").AnyTimes()
clusterMeta.EXPECT().GetAllClusterInfo().Return(map[string]cluster.ClusterInformation{}).AnyTimes()
clusterMeta.EXPECT().RegisterMetadataChangeCallback(gomock.Any(), gomock.Any())
localClient := newTestFrontendHTTPClient(ts)
localClient := newTestFrontendHTTPClient(ts)
// Create a cache that will fail for the requested cluster since we don't set up metadata fully.
httpClientCache := cluster.NewFrontendHTTPClientCache(clusterMeta, nil)
// Create a cache that will fail for the requested cluster since we don't set up metadata fully.
httpClientCache := cluster.NewFrontendHTTPClientCache(clusterMeta, nil)
r, err := http.NewRequest(http.MethodPost, commonnexus.SystemCallbackURL, nil)
require.NoError(t, err)
r.Header.Set(commonnexus.CallbackTokenHeader, tokenStr)
r, err := http.NewRequest(http.MethodPost, commonnexus.SystemCallbackURL, nil)
require.NoError(t, err)
r.Header.Set(commonnexus.CallbackTokenHeader, tokenStr)
resp, err := routeSystemCallbackRequest(
r,
clusterMeta,
nsRegistry,
httpClientCache,
tokenGen,
localClient,
log.NewNoopLogger(),
)
require.NoError(t, err)
defer func() { _ = resp.Body.Close() }()
require.Equal(t, http.StatusOK, resp.StatusCode)
resp, err := routeSystemCallbackRequest(
r,
clusterMeta,
nsRegistry,
httpClientCache,
tokenGen,
localClient,
log.NewNoopLogger(),
)
require.NoError(t, err)
defer func() { _ = resp.Body.Close() }()
require.Equal(t, http.StatusOK, resp.StatusCode)
require.Equal(t, commonnexus.PathCompletionCallbackNoIdentifier, gotPath)
})
}
}
func TestRouteRequest_SystemCallback(t *testing.T) {
// Verify that routeRequest delegates to routeSystemCallbackRequest for system callback URLs.
var gotPath string
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
assert.Equal(t, commonnexus.PathCompletionCallbackNoIdentifier, r.URL.Path)
gotPath = r.URL.Path
w.WriteHeader(http.StatusOK)
}))
defer ts.Close()
@@ -321,4 +411,5 @@ func TestRouteRequest_SystemCallback(t *testing.T) {
require.NoError(t, err)
defer func() { _ = resp.Body.Close() }()
require.Equal(t, http.StatusOK, resp.StatusCode)
require.Equal(t, commonnexus.PathCompletionCallbackNoIdentifier, gotPath)
}

View File

@@ -7,6 +7,7 @@ import (
commonpb "go.temporal.io/api/common/v1"
failurepb "go.temporal.io/api/failure/v1"
"go.temporal.io/api/serviceerror"
persistencespb "go.temporal.io/server/api/persistence/v1"
"go.temporal.io/server/chasm"
nexusoperationpb "go.temporal.io/server/chasm/lib/nexusoperation/gen/nexusoperationpb/v1"
"go.temporal.io/server/common/backoff"
@@ -15,6 +16,7 @@ import (
)
var _ chasm.StateMachine[nexusoperationpb.OperationStatus] = (*Operation)(nil)
var _ chasm.NexusCompletionHandler = (*Operation)(nil)
// ErrCancellationAlreadyRequested is returned when a cancellation has already been requested for an operation.
var ErrCancellationAlreadyRequested = serviceerror.NewFailedPrecondition("cancellation already requested")
@@ -217,6 +219,7 @@ type saveInvocationResultInput struct {
retryPolicy backoff.RetryPolicy
}
// saveInvocationResult handles the outcome of the initial start call.
func (o *Operation) saveInvocationResult(
ctx chasm.MutableContext,
input saveInvocationResultInput,
@@ -225,6 +228,8 @@ func (o *Operation) saveInvocationResult(
case startResultOK:
links := convertResponseLinks(r.response.Links, ctx.Logger())
if r.response.Pending != nil {
// An async operation transitions to STARTED here;
// HandleNexusCompletion will apply its outcome from the completion callback.
return nil, o.onStarted(ctx, r.response.Pending.Token, links)
}
return nil, o.onCompleted(ctx, r.response.Successful, links)
@@ -244,6 +249,31 @@ func (o *Operation) saveInvocationResult(
}
}
// HandleNexusCompletion handles the outcome of an asynchronous completion callback.
func (o *Operation) HandleNexusCompletion(
ctx chasm.MutableContext,
completion *persistencespb.ChasmNexusCompletion,
) error {
// TODO: support completion-before-start
// Request ID lets us reject a stale or misrouted completion.
if completion.GetRequestId() != "" && o.GetRequestId() != completion.GetRequestId() {
return serviceerror.NewNotFound("operation not found")
}
switch outcome := completion.Outcome.(type) {
case *persistencespb.ChasmNexusCompletion_Success:
return o.onCompleted(ctx, outcome.Success, completion.GetLinks())
case *persistencespb.ChasmNexusCompletion_Failure:
if outcome.Failure.GetCanceledFailureInfo() != nil {
return o.onCanceled(ctx, outcome.Failure)
}
return o.onFailed(ctx, outcome.Failure)
default:
return serviceerror.NewInvalidArgument("invalid completion outcome")
}
}
func (o *Operation) resolveUnsuccessfully(ctx chasm.MutableContext, failure *failurepb.Failure, closeTime time.Time) error {
// When we transition from scheduled to failed it is always due to the attempt failing with a non
// retryable failure. The failure should be recorded in the state for standalone Nexus operations.

View File

@@ -11,6 +11,7 @@ import (
enumspb "go.temporal.io/api/enums/v1"
failurepb "go.temporal.io/api/failure/v1"
"go.temporal.io/api/serviceerror"
persistencespb "go.temporal.io/server/api/persistence/v1"
tokenspb "go.temporal.io/server/api/token/v1"
"go.temporal.io/server/chasm"
nexusoperationpb "go.temporal.io/server/chasm/lib/nexusoperation/gen/nexusoperationpb/v1"
@@ -181,17 +182,7 @@ func (h *operationInvocationTaskHandler) Execute(
endpoint, err := h.lookupEndpoint(ctx, ns.ID(), args.endpointID, args.endpointName)
if err != nil {
if _, ok := errors.AsType[*serviceerror.NotFound](err); ok {
h.logger.Error("endpoint not found while processing invocation task", tag.Error(err))
handlerErr := nexus.NewHandlerErrorf(nexus.HandlerErrorTypeNotFound, "endpoint not registered")
result, err := newStartResult(nil, handlerErr)
if err != nil {
return fmt.Errorf("failed to construct invocation result: %w", err)
}
_, _, err = chasm.UpdateComponent(ctx, opRef, (*Operation).saveInvocationResult, saveInvocationResultInput{
result: result,
retryPolicy: h.config.RetryPolicy(),
})
return err
return h.handleMissingEndpoint(ctx, opRef, err)
}
return err
}
@@ -302,6 +293,24 @@ func (h *operationInvocationTaskHandler) Execute(
return saveErr
}
func (h *operationInvocationTaskHandler) handleMissingEndpoint(
ctx context.Context,
opRef chasm.ComponentRef,
lookupErr error,
) error {
h.logger.Error("endpoint not found while processing invocation task", tag.Error(lookupErr))
handlerErr := nexus.NewHandlerErrorf(nexus.HandlerErrorTypeNotFound, "endpoint not registered")
result, err := newStartResult(nil, handlerErr)
if err != nil {
return fmt.Errorf("failed to construct invocation result: %w", err)
}
_, _, err = chasm.UpdateComponent(ctx, opRef, (*Operation).saveInvocationResult, saveInvocationResultInput{
result: result,
retryPolicy: h.config.RetryPolicy(),
})
return err
}
func (h *operationInvocationTaskHandler) validateStartResult(
ns *namespace.Namespace,
result *nexusrpc.ClientStartOperationResponse[*commonpb.Payload],
@@ -324,8 +333,23 @@ func (h *operationInvocationTaskHandler) generateCallbackToken(
serializedRef []byte,
requestID string,
) (string, error) {
// TODO: replace this with selective CHASM consistency solution once available
ref := &persistencespb.ChasmComponentRef{}
if err := ref.Unmarshal(serializedRef); err != nil {
return "", fmt.Errorf("%w: %w", queueserrors.NewUnprocessableTaskError("failed to decode component ref for callback token"), err)
}
// Both VT becomes stale after workflow mutations between token minting and completion arrival.
ref.ExecutionVersionedTransition = nil
ref.ComponentInitialVersionedTransition = nil
stableRef, err := ref.Marshal()
if err != nil {
return "", fmt.Errorf("%w: %w", queueserrors.NewUnprocessableTaskError("failed to encode component ref for callback token"), err)
}
token, err := h.callbackTokenGenerator.Tokenize(&tokenspb.NexusOperationCompletion{
ComponentRef: serializedRef,
ComponentRef: stableRef,
RequestId: requestID,
})
if err != nil {

View File

@@ -0,0 +1,95 @@
package nexusoperation
import (
"testing"
"time"
"github.com/stretchr/testify/require"
failurepb "go.temporal.io/api/failure/v1"
"go.temporal.io/api/serviceerror"
persistencespb "go.temporal.io/server/api/persistence/v1"
"go.temporal.io/server/chasm"
nexusoperationpb "go.temporal.io/server/chasm/lib/nexusoperation/gen/nexusoperationpb/v1"
)
func newScheduledTestOperation(t *testing.T, ctx *chasm.MockMutableContext) *Operation {
t.Helper()
op := newTestOperation()
require.NoError(t, TransitionScheduled.Apply(op, ctx, EventScheduled{}))
return op
}
func TestHandleNexusCompletion(t *testing.T) {
newStartedOp := func(t *testing.T, ctx *chasm.MockMutableContext) *Operation {
t.Helper()
op := newScheduledTestOperation(t, ctx)
require.NoError(t, TransitionStarted.Apply(op, ctx, EventStarted{OperationToken: "tok"}))
return op
}
newCtx := func() *chasm.MockMutableContext {
return &chasm.MockMutableContext{
MockContext: chasm.MockContext{
HandleNow: func(chasm.Component) time.Time { return defaultTime },
},
}
}
t.Run("Success", func(t *testing.T) {
ctx := newCtx()
op := newStartedOp(t, ctx)
err := op.HandleNexusCompletion(ctx, &persistencespb.ChasmNexusCompletion{
RequestId: op.GetRequestId(),
Outcome: &persistencespb.ChasmNexusCompletion_Success{
Success: mustToPayload(t, "result"),
},
})
require.NoError(t, err)
require.Equal(t, nexusoperationpb.OPERATION_STATUS_SUCCEEDED, op.GetStatus())
})
t.Run("Failure", func(t *testing.T) {
ctx := newCtx()
op := newStartedOp(t, ctx)
err := op.HandleNexusCompletion(ctx, &persistencespb.ChasmNexusCompletion{
RequestId: op.GetRequestId(),
Outcome: &persistencespb.ChasmNexusCompletion_Failure{
Failure: &failurepb.Failure{Message: "oops"},
},
})
require.NoError(t, err)
require.Equal(t, nexusoperationpb.OPERATION_STATUS_FAILED, op.GetStatus())
})
t.Run("Canceled", func(t *testing.T) {
ctx := newCtx()
op := newStartedOp(t, ctx)
err := op.HandleNexusCompletion(ctx, &persistencespb.ChasmNexusCompletion{
RequestId: op.GetRequestId(),
Outcome: &persistencespb.ChasmNexusCompletion_Failure{
Failure: &failurepb.Failure{
Message: "canceled",
FailureInfo: &failurepb.Failure_CanceledFailureInfo{
CanceledFailureInfo: &failurepb.CanceledFailureInfo{},
},
},
},
})
require.NoError(t, err)
require.Equal(t, nexusoperationpb.OPERATION_STATUS_CANCELED, op.GetStatus())
})
t.Run("RequestIDMismatch", func(t *testing.T) {
ctx := newCtx()
op := newStartedOp(t, ctx)
err := op.HandleNexusCompletion(ctx, &persistencespb.ChasmNexusCompletion{
RequestId: "wrong-request-id",
Outcome: &persistencespb.ChasmNexusCompletion_Success{
Success: mustToPayload(t, "result"),
},
})
require.Error(t, err)
var notFound *serviceerror.NotFound
require.ErrorAs(t, err, &notFound)
require.Equal(t, nexusoperationpb.OPERATION_STATUS_STARTED, op.GetStatus())
})
}

View File

@@ -4,6 +4,7 @@ import (
"encoding/base64"
"encoding/json"
"go.temporal.io/api/serviceerror"
tokenspb "go.temporal.io/server/api/token/v1"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
@@ -58,7 +59,36 @@ func (g *CallbackTokenGenerator) DecodeCompletion(token *CallbackToken) (*tokens
}
completion := &tokenspb.NexusOperationCompletion{}
return completion, proto.Unmarshal(plaintext, completion)
if err := proto.Unmarshal(plaintext, completion); err != nil {
return nil, err
}
if err := validateCompletion(completion); err != nil {
return nil, err
}
return completion, nil
}
func validateCompletion(completion *tokenspb.NexusOperationCompletion) error {
hasCHASMRef := len(completion.GetComponentRef()) > 0
hasHSMRef := completion.GetNamespaceId() != "" ||
completion.GetWorkflowId() != "" ||
completion.GetRunId() != "" ||
completion.GetRef() != nil
isCompleteHSM := completion.GetNamespaceId() != "" &&
completion.GetWorkflowId() != "" &&
completion.GetRunId() != "" &&
completion.GetRef() != nil
switch {
case hasCHASMRef && hasHSMRef:
return serviceerror.NewInvalidArgument("callback token contains both HSM and CHASM fields")
case hasCHASMRef:
return nil
case isCompleteHSM:
return nil
default:
return serviceerror.NewInvalidArgument("callback token must contain either all HSM fields or a component ref")
}
}
// DecodeCallbackToken unmarshals the given token applying minimal data verification.

View File

@@ -0,0 +1,91 @@
package nexus
import (
"testing"
"github.com/stretchr/testify/require"
"go.temporal.io/api/serviceerror"
persistencespb "go.temporal.io/server/api/persistence/v1"
tokenspb "go.temporal.io/server/api/token/v1"
)
func TestCallbackTokenGenerator_DecodeCompletion(t *testing.T) {
g := NewCallbackTokenGenerator()
for _, tc := range []struct {
name string
completion *tokenspb.NexusOperationCompletion
wantErr string
}{
{
name: "valid HSM",
completion: &tokenspb.NexusOperationCompletion{
NamespaceId: "ns-id",
WorkflowId: "wf-id",
RunId: "run-id",
Ref: &persistencespb.StateMachineRef{},
},
},
{
name: "valid CHASM",
completion: &tokenspb.NexusOperationCompletion{
ComponentRef: []byte("component-ref"),
},
},
{
name: "mixed with namespace id",
completion: &tokenspb.NexusOperationCompletion{
NamespaceId: "ns-id",
ComponentRef: []byte("component-ref"),
},
wantErr: "both HSM and CHASM",
},
{
name: "mixed with workflow id",
completion: &tokenspb.NexusOperationCompletion{
WorkflowId: "wf-id",
ComponentRef: []byte("component-ref"),
},
wantErr: "both HSM and CHASM",
},
{
name: "mixed with run id",
completion: &tokenspb.NexusOperationCompletion{
RunId: "run-id",
ComponentRef: []byte("component-ref"),
},
wantErr: "both HSM and CHASM",
},
{
name: "mixed with ref",
completion: &tokenspb.NexusOperationCompletion{
Ref: &persistencespb.StateMachineRef{},
ComponentRef: []byte("component-ref"),
},
wantErr: "both HSM and CHASM",
},
{
name: "empty",
completion: &tokenspb.NexusOperationCompletion{},
wantErr: "either all HSM fields or a component ref",
},
} {
tokenString, err := g.Tokenize(tc.completion)
require.NoError(t, err, tc.name)
token, err := DecodeCallbackToken(tokenString)
require.NoError(t, err, tc.name)
completion, err := g.DecodeCompletion(token)
if tc.wantErr == "" {
require.NoError(t, err, tc.name)
require.NotNil(t, completion, tc.name)
continue
}
require.ErrorContains(t, err, tc.wantErr, tc.name)
var invalidArgumentErr *serviceerror.InvalidArgument
require.ErrorAs(t, err, &invalidArgumentErr, tc.name)
require.Nil(t, completion, tc.name)
}
}

View File

@@ -209,6 +209,8 @@ func TestRouteSystemCallbackRequest_NamespaceNotFound(t *testing.T) {
tokenStr, err := tokenGen.Tokenize(&tokenspb.NexusOperationCompletion{
NamespaceId: "ns-id-1",
WorkflowId: "wf-1",
RunId: "run-1",
Ref: &persistencespb.StateMachineRef{},
})
require.NoError(t, err)
@@ -250,6 +252,8 @@ func TestRouteSystemCallbackRequest_Success(t *testing.T) {
tokenStr, err := tokenGen.Tokenize(&tokenspb.NexusOperationCompletion{
NamespaceId: "ns-id-1",
WorkflowId: "wf-1",
RunId: "run-1",
Ref: &persistencespb.StateMachineRef{},
})
require.NoError(t, err)

View File

@@ -1,57 +0,0 @@
package frontend
import (
"net/http"
"github.com/gorilla/mux"
"go.temporal.io/server/common/dynamicconfig"
"go.temporal.io/server/common/headers"
"go.temporal.io/server/common/log"
"go.temporal.io/server/common/metrics"
commonnexus "go.temporal.io/server/common/nexus"
"go.temporal.io/server/common/nexus/nexusrpc"
"go.temporal.io/server/common/rpc"
"go.temporal.io/server/components/nexusoperations"
"go.uber.org/fx"
)
var Module = fx.Module(
"component.nexusoperations.frontend",
fx.Provide(ConfigProvider),
fx.Provide(commonnexus.NewCallbackTokenGenerator),
fx.Invoke(RegisterHTTPHandler),
)
func ConfigProvider(coll *dynamicconfig.Collection) *Config {
return &Config{
PayloadSizeLimit: dynamicconfig.BlobSizeLimitError.Get(coll),
ForwardingEnabledForNamespace: dynamicconfig.EnableNamespaceNotActiveAutoForwarding.Get(coll),
MaxOperationTokenLength: nexusoperations.MaxOperationTokenLength.Get(coll),
}
}
func RegisterHTTPHandler(options HandlerOptions, logger log.Logger, router *mux.Router) {
h := nexusrpc.NewCompletionHTTPHandler(nexusrpc.CompletionHandlerOptions{
Handler: &completionHandler{
options,
headers.NewDefaultVersionChecker(),
options.MetricsHandler.Counter(metrics.NexusCompletionRequestPreProcessErrors.Name()),
},
Logger: log.NewSlogLogger(logger),
Serializer: commonnexus.PayloadSerializer,
})
router.Path("/" + commonnexus.RouteCompletionCallback.Representation()).HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Limit the request body to max allowed Payload size.
// Content headers are transformed to Payload metadata and contribute to the Payload size as well. A separate
// limit is enforced on top of this in the CompleteOperation method.
r.Body = http.MaxBytesReader(w, r.Body, rpc.MaxNexusAPIRequestBodyBytes)
h.ServeHTTP(w, r)
})
router.Path(commonnexus.PathCompletionCallbackNoIdentifier).HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Limit the request body to max allowed Payload size.
// Content headers are transformed to Payload metadata and contribute to the Payload size as well. A separate
// limit is enforced on top of this in the CompleteOperation method.
r.Body = http.MaxBytesReader(w, r.Body, rpc.MaxNexusAPIRequestBodyBytes)
h.ServeHTTP(w, r)
})
}

View File

@@ -1205,6 +1205,8 @@ message CompleteNexusOperationChasmRequest {
}
// Time when the operation was closed.
google.protobuf.Timestamp close_time = 4;
// Links from the Nexus completion callback (e.g. references to the handler workflow).
repeated temporal.api.common.v1.Link links = 5;
}
message CompleteNexusOperationChasmResponse {}

View File

@@ -128,4 +128,6 @@ message ChasmNexusCompletion {
// Request ID embedded in the NexusOperationScheduledEvent.
// Allows completing a started operation after a workflow has been reset.
string request_id = 4;
// Links from the Nexus completion callback (e.g. references to the handler workflow).
repeated temporal.api.common.v1.Link links = 5;
}

View File

@@ -9,6 +9,7 @@ import (
"go.temporal.io/server/chasm"
"go.temporal.io/server/chasm/lib/activity"
"go.temporal.io/server/chasm/lib/callback"
chasmnexus "go.temporal.io/server/chasm/lib/nexusoperation"
"go.temporal.io/server/chasm/lib/scheduler/gen/schedulerpb/v1"
"go.temporal.io/server/client"
"go.temporal.io/server/common"
@@ -25,7 +26,6 @@ import (
"go.temporal.io/server/common/metrics"
"go.temporal.io/server/common/namespace"
"go.temporal.io/server/common/namespace/nsreplication"
"go.temporal.io/server/common/nexus"
"go.temporal.io/server/common/persistence"
"go.temporal.io/server/common/persistence/serialization"
"go.temporal.io/server/common/persistence/visibility"
@@ -42,7 +42,6 @@ import (
"go.temporal.io/server/common/searchattribute"
"go.temporal.io/server/common/telemetry"
hsmcallbacks "go.temporal.io/server/components/callbacks"
nexusfrontend "go.temporal.io/server/components/nexusoperations/frontend"
"go.temporal.io/server/service"
"go.temporal.io/server/service/frontend/configs"
"go.temporal.io/server/service/history/tasks"
@@ -110,16 +109,18 @@ var Module = fx.Options(
fx.Provide(OperatorHandlerProvider),
fx.Provide(NewVersionChecker),
fx.Provide(ServiceResolverProvider),
fx.Invoke(RegisterNexusHTTPHandler),
fx.Provide(newNexusCompletionHandler),
fx.Provide(NewNexusOperationHTTPHandler),
fx.Provide(newNexusCompletionHTTPHandler),
fx.Invoke(RegisterNexusOperationHTTPHandler),
fx.Invoke(RegisterNexusCompletionHTTPHandler),
fx.Invoke(RegisterOpenAPIHTTPHandler),
fx.Provide(HTTPAPIServerProvider),
fx.Provide(NewServiceProvider),
fx.Provide(NexusEndpointClientProvider),
fx.Provide(NexusEndpointRegistryProvider),
fx.Invoke(ServiceLifetimeHooks),
fx.Invoke(EndpointRegistryLifetimeHooks),
fx.Provide(schedulerpb.NewSchedulerServiceLayeredClient),
nexusfrontend.Module,
chasmnexus.Module,
activity.FrontendModule,
fx.Provide(visibility.ChasmVisibilityManagerProvider),
fx.Provide(chasm.ChasmVisibilityInterceptorProvider),
@@ -887,46 +888,17 @@ func HandlerProvider(
return wfHandler
}
func RegisterNexusHTTPHandler(
serviceConfig *Config,
serviceName primitives.ServiceName,
matchingClient resource.MatchingClient,
metricsHandler metrics.Handler,
clusterMetadata cluster.Metadata,
clientCache *cluster.FrontendHTTPClientCache,
namespaceRegistry namespace.Registry,
endpointRegistry nexus.EndpointRegistry,
authInterceptor *authorization.Interceptor,
telemetryInterceptor *interceptor.TelemetryInterceptor,
requestErrorHandler *interceptor.RequestErrorHandler,
redirectionInterceptor *interceptor.Redirection,
namespaceRateLimiterInterceptor interceptor.NamespaceRateLimitInterceptor,
namespaceCountLimiterInterceptor *interceptor.ConcurrentRequestLimitInterceptor,
namespaceValidatorInterceptor *interceptor.NamespaceValidatorInterceptor,
rateLimitInterceptor *interceptor.RateLimitInterceptor,
logger log.Logger,
func RegisterNexusOperationHTTPHandler(
h *NexusOperationHTTPHandler,
router *mux.Router,
) {
h.RegisterRoutes(router)
}
func RegisterNexusCompletionHTTPHandler(
h *nexusCompletionHTTPHandler,
router *mux.Router,
httpTraceProvider nexus.HTTPClientTraceProvider,
) {
h := NewNexusHTTPHandler(
serviceConfig,
matchingClient,
metricsHandler,
clusterMetadata,
clientCache,
namespaceRegistry,
endpointRegistry,
authInterceptor,
telemetryInterceptor,
requestErrorHandler,
redirectionInterceptor,
namespaceValidatorInterceptor,
namespaceRateLimiterInterceptor,
namespaceCountLimiterInterceptor,
rateLimitInterceptor,
logger,
httpTraceProvider,
)
h.RegisterRoutes(router)
}
@@ -1009,27 +981,6 @@ func NexusEndpointClientProvider(
)
}
func NexusEndpointRegistryProvider(
matchingClient resource.MatchingClient,
nexusEndpointManager persistence.NexusEndpointManager,
dc *dynamicconfig.Collection,
logger log.Logger,
metricsHandler metrics.Handler,
) nexus.EndpointRegistry {
registryConfig := nexus.NewEndpointRegistryConfig(dc)
return nexus.NewEndpointRegistry(
registryConfig,
matchingClient,
nexusEndpointManager,
logger,
metricsHandler,
)
}
func EndpointRegistryLifetimeHooks(lc fx.Lifecycle, registry nexus.EndpointRegistry) {
lc.Append(fx.StartStopHook(registry.StartLifecycle, registry.StopLifecycle))
}
func ServiceLifetimeHooks(lc fx.Lifecycle, svc *Service) {
lc.Append(fx.StartStopHook(svc.Start, svc.Stop))
}

View File

@@ -17,10 +17,11 @@ import (
commonpb "go.temporal.io/api/common/v1"
"go.temporal.io/api/serviceerror"
"go.temporal.io/server/api/historyservice/v1"
persistencespb "go.temporal.io/server/api/persistence/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/dynamicconfig"
"go.temporal.io/server/common/headers"
"go.temporal.io/server/common/log"
"go.temporal.io/server/common/log/tag"
@@ -29,32 +30,19 @@ import (
commonnexus "go.temporal.io/server/common/nexus"
"go.temporal.io/server/common/nexus/nexusrpc"
"go.temporal.io/server/common/resource"
"go.temporal.io/server/common/rpc"
"go.temporal.io/server/common/rpc/interceptor"
"go.temporal.io/server/service/frontend/configs"
"go.uber.org/fx"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/metadata"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/timestamppb"
)
var apiName = configs.CompleteNexusOperation
const (
methodNameForMetrics = "CompleteNexusOperation"
// user-agent header contains Nexus SDK client info in the form <sdk-name>/v<sdk-version>
headerUserAgent = "user-agent"
clientNameVersionDelim = "/v"
)
type Config struct {
MaxOperationTokenLength dynamicconfig.IntPropertyFnWithNamespaceFilter
PayloadSizeLimit dynamicconfig.IntPropertyFnWithNamespaceFilter
ForwardingEnabledForNamespace dynamicconfig.BoolPropertyFnWithNamespaceFilter
}
type HandlerOptions struct {
fx.In
const nexusCompletionAPIName = configs.CompleteNexusOperation
const nexusCompletionMethodNameForMetrics = "CompleteNexusOperation"
type nexusCompletionHandler struct {
ClusterMetadata cluster.Metadata
NamespaceRegistry namespace.Registry
Logger log.Logger
@@ -72,17 +60,69 @@ type HandlerOptions struct {
RedirectionInterceptor *interceptor.Redirection
ForwardingClients *cluster.FrontendHTTPClientCache
HTTPTraceProvider commonnexus.HTTPClientTraceProvider
clientVersionChecker headers.VersionChecker
preProcessErrorsCounter metrics.CounterIface
}
type completionHandler struct {
HandlerOptions
clientVersionChecker headers.VersionChecker
preProcessErrorsCounter metrics.CounterIface
type nexusCompletionHTTPHandler struct {
httpHandler http.Handler
}
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,
) *nexusCompletionHandler {
return &nexusCompletionHandler{
ClusterMetadata: clusterMetadata,
NamespaceRegistry: namespaceRegistry,
Logger: logger,
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()),
}
}
func newNexusCompletionHTTPHandler(handler *nexusCompletionHandler, logger log.Logger) *nexusCompletionHTTPHandler {
return &nexusCompletionHTTPHandler{
httpHandler: nexusrpc.NewCompletionHTTPHandler(nexusrpc.CompletionHandlerOptions{
Handler: handler,
Logger: log.NewSlogLogger(logger),
Serializer: commonnexus.PayloadSerializer,
}),
}
}
// CompleteOperation implements nexus.CompletionHandler.
// nolint:revive // (cyclomatic complexity) This function is long but the complexity is justified.
func (h *completionHandler) CompleteOperation(ctx context.Context, r *nexusrpc.CompletionRequest) (retErr error) {
func (h *nexusCompletionHandler) CompleteOperation(ctx context.Context, r *nexusrpc.CompletionRequest) (retErr error) {
startTime := time.Now()
token, err := commonnexus.DecodeCallbackToken(r.HTTPRequest.Header.Get(commonnexus.CallbackTokenHeader))
if err != nil {
@@ -95,30 +135,46 @@ func (h *completionHandler) CompleteOperation(ctx context.Context, r *nexusrpc.C
h.Logger.Error("failed to decode completion from token", tag.Error(err))
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid callback token")
}
ns, err := h.NamespaceRegistry.GetNamespaceByID(namespace.ID(completion.NamespaceId))
// Determine the target namespace, workflow, and run ID from the completion token.
targetNamespaceID := completion.GetNamespaceId()
targetBusinessID := completion.GetWorkflowId()
targetRunID := completion.GetRunId()
if len(completion.GetComponentRef()) > 0 {
ref := &persistencespb.ChasmComponentRef{}
if err := proto.Unmarshal(completion.GetComponentRef(), ref); err != nil {
h.Logger.Error("failed to unmarshal CHASM component ref", tag.Error(err))
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid callback token")
}
targetNamespaceID = ref.GetNamespaceId()
targetBusinessID = ref.GetBusinessId()
targetRunID = ref.GetRunId()
}
ns, err := h.NamespaceRegistry.GetNamespaceByID(namespace.ID(targetNamespaceID))
if err != nil {
h.Logger.Error("failed to get namespace for nexus completion request", tag.WorkflowNamespaceID(completion.NamespaceId), tag.Error(err))
h.Logger.Error("failed to get namespace for nexus completion request", tag.WorkflowNamespaceID(targetNamespaceID), tag.Error(err))
h.preProcessErrorsCounter.Record(1)
var nfe *serviceerror.NamespaceNotFound
if errors.As(err, &nfe) {
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeNotFound, "namespace %q not found", completion.NamespaceId)
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeNotFound, "namespace %q not found", targetNamespaceID)
}
return commonnexus.ConvertGRPCError(err, false)
}
logger := log.With(
h.Logger,
tag.WorkflowNamespace(ns.Name().String()),
tag.WorkflowID(completion.GetWorkflowId()),
tag.WorkflowRunID(completion.GetRunId()),
tag.WorkflowID(targetBusinessID),
tag.WorkflowRunID(targetRunID),
)
rCtx := &requestContext{
completionHandler: h,
namespace: ns,
workflowID: completion.GetWorkflowId(),
logger: log.With(h.Logger, tag.WorkflowNamespace(ns.Name().String())),
metricsHandler: h.MetricsHandler.WithTags(metrics.NamespaceTag(ns.Name().String())),
nexusCompletionHandler: h,
namespace: ns,
businessID: targetBusinessID,
logger: log.With(h.Logger, tag.WorkflowNamespace(ns.Name().String())),
metricsHandler: h.MetricsHandler.WithTags(metrics.NamespaceTag(ns.Name().String())),
metricsHandlerForInterceptors: h.MetricsHandler.WithTags(
metrics.OperationTag(methodNameForMetrics),
metrics.OperationTag(nexusCompletionMethodNameForMetrics),
metrics.NamespaceTag(ns.Name().String()),
),
requestStartTime: startTime,
@@ -128,6 +184,7 @@ func (h *completionHandler) CompleteOperation(ctx context.Context, r *nexusrpc.C
}
ctx = rCtx.augmentContext(ctx, r.HTTPRequest.Header)
defer rCtx.capturePanicAndRecordMetrics(&ctx, &retErr)
if r.HTTPRequest.URL.Path != commonnexus.PathCompletionCallbackNoIdentifier {
nsNameEscaped := commonnexus.RouteCompletionCallback.Deserialize(mux.Vars(r.HTTPRequest))
nsName, err := url.PathUnescape(nsNameEscaped)
@@ -141,7 +198,7 @@ func (h *completionHandler) CompleteOperation(ctx context.Context, r *nexusrpc.C
"namespace ID in token doesn't match the token",
tag.WorkflowNamespaceID(ns.ID().String()),
tag.Error(err),
tag.String("completion-namespace-id", completion.GetNamespaceId()),
tag.String("completion-namespace-id", targetNamespaceID),
)
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid callback token")
}
@@ -154,7 +211,7 @@ func (h *completionHandler) CompleteOperation(ctx context.Context, r *nexusrpc.C
}
return err
}
tokenLimit := h.Config.MaxOperationTokenLength(ns.Name().String())
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)
}
@@ -183,72 +240,144 @@ func (h *completionHandler) CompleteOperation(ctx context.Context, r *nexusrpc.C
h.Logger.Warn(fmt.Sprintf("invalid link data type: %q", nexusLink.Type))
}
}
hr := &historyservice.CompleteNexusOperationRequest{
Completion: completion,
State: string(r.State),
OperationToken: r.OperationToken,
StartTime: timestamppb.New(r.StartTime),
Links: links,
}
var successPayload *commonpb.Payload
switch r.State { // nolint:exhaustive
case nexus.OperationStateFailed, nexus.OperationStateCanceled:
hr.Outcome = &historyservice.CompleteNexusOperationRequest_Failure{
Failure: commonnexus.NexusFailureToProtoFailure(*r.Error.OriginalFailure),
}
// no validation needed
case nexus.OperationStateSucceeded:
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")
}
if result.Size() > h.Config.PayloadSizeLimit(ns.Name().String()) {
if result.Size() > h.Config.BlobSizeLimitError(ns.Name().String()) {
logger.Error("payload size exceeds error limit for Nexus CompleteOperation request", tag.WorkflowNamespace(ns.Name().String()))
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "result exceeds size limit")
}
hr.Outcome = &historyservice.CompleteNexusOperationRequest_Success{
Success: result,
}
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)), tag.Error(err))
logger.Error("invalid operation state in completion request", tag.String("state", string(r.State)))
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid completion state")
}
_, err = h.HistoryClient.CompleteNexusOperation(ctx, hr)
if err != nil {
logger.Error("failed to process nexus completion request", tag.Error(err))
var namespaceInactiveErr *serviceerror.NamespaceNotActive
if errors.As(err, &namespaceInactiveErr) {
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnavailable, "cluster inactive")
}
var notFoundErr *serviceerror.NotFound
if errors.As(err, &notFoundErr) {
return commonnexus.ConvertGRPCError(err, true)
}
return commonnexus.ConvertGRPCError(err, false)
if len(completion.GetComponentRef()) > 0 {
err = h.completeChasmOperation(ctx, logger, completion, successPayload, r, links)
} else {
err = h.completeHSMOperation(ctx, completion, successPayload, r, links)
}
return nil
if err == nil {
return nil
}
logger.Error("failed to process nexus completion request", tag.Error(err))
var namespaceInactiveErr *serviceerror.NamespaceNotActive
if errors.As(err, &namespaceInactiveErr) {
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnavailable, "cluster inactive")
}
var notFoundErr *serviceerror.NotFound
if errors.As(err, &notFoundErr) {
return commonnexus.ConvertGRPCError(err, true)
}
return commonnexus.ConvertGRPCError(err, false)
}
func (h *completionHandler) forwardCompleteOperation(ctx context.Context, r *nexusrpc.CompletionRequest, rCtx *requestContext) error {
client, err := h.ForwardingClients.Get(rCtx.namespace.ActiveClusterName(namespace.RoutingKey{ID: rCtx.workflowID}))
func (h *nexusCompletionHandler) completeHSMOperation(
ctx context.Context,
completion *tokenspb.NexusOperationCompletion,
successPayload *commonpb.Payload,
req *nexusrpc.CompletionRequest,
links []*commonpb.Link,
) error {
hr := &historyservice.CompleteNexusOperationRequest{
Completion: completion,
State: string(req.State),
OperationToken: req.OperationToken,
StartTime: timestamppb.New(req.StartTime),
Links: links,
}
switch req.State { // nolint:exhaustive
case nexus.OperationStateFailed, nexus.OperationStateCanceled:
hr.Outcome = &historyservice.CompleteNexusOperationRequest_Failure{
Failure: commonnexus.NexusFailureToProtoFailure(*req.Error.OriginalFailure),
}
case nexus.OperationStateSucceeded:
hr.Outcome = &historyservice.CompleteNexusOperationRequest_Success{
Success: successPayload,
}
default:
// Should be unreachable as validated earlier.
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid completion state")
}
_, err := h.HistoryClient.CompleteNexusOperation(ctx, hr)
return err
}
func (h *nexusCompletionHandler) completeChasmOperation(
ctx context.Context,
logger log.Logger,
completion *tokenspb.NexusOperationCompletion,
successPayload *commonpb.Payload,
req *nexusrpc.CompletionRequest,
links []*commonpb.Link,
) error {
hr := &historyservice.CompleteNexusOperationChasmRequest{
Completion: &tokenspb.NexusOperationCompletion{
RequestId: completion.GetRequestId(),
ComponentRef: completion.GetComponentRef(),
},
Links: links,
}
if !req.CloseTime.IsZero() {
hr.CloseTime = timestamppb.New(req.CloseTime)
}
switch req.State { // nolint:exhaustive
case nexus.OperationStateFailed, nexus.OperationStateCanceled:
failure, err := commonnexus.NexusFailureToTemporalFailure(*req.Error.OriginalFailure)
if err != nil {
logger.Error("cannot convert nexus failure from completion request", tag.Error(err))
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid failure content")
}
hr.Outcome = &historyservice.CompleteNexusOperationChasmRequest_Failure{
Failure: failure,
}
case nexus.OperationStateSucceeded:
hr.Outcome = &historyservice.CompleteNexusOperationChasmRequest_Success{
Success: successPayload,
}
default:
// Should be unreachable as validated earlier.
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeBadRequest, "invalid completion state")
}
_, err := h.HistoryClient.CompleteNexusOperationChasm(ctx, hr)
return err
}
func (h *nexusCompletionHandler) forwardCompleteOperation(ctx context.Context, r *nexusrpc.CompletionRequest, rCtx *requestContext) error {
targetCluster := rCtx.namespace.ActiveClusterName(namespace.RoutingKey{ID: rCtx.businessID})
client, err := h.ForwardingClients.Get(targetCluster)
if err != nil {
h.Logger.Error("unable to get HTTP client for forward request", tag.Operation(apiName), tag.WorkflowNamespace(rCtx.namespace.Name().String()), tag.Error(err), tag.SourceCluster(h.ClusterMetadata.GetCurrentClusterName()), tag.TargetCluster(rCtx.namespace.ActiveClusterName(namespace.RoutingKey{ID: rCtx.workflowID})))
h.Logger.Error("unable to get HTTP client for forward request", tag.Operation(nexusCompletionAPIName), tag.WorkflowNamespace(rCtx.namespace.Name().String()), tag.Error(err), tag.SourceCluster(h.ClusterMetadata.GetCurrentClusterName()), tag.TargetCluster(targetCluster))
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "internal error")
}
forwardURL, err := url.JoinPath(client.BaseURL(), commonnexus.RouteCompletionCallback.Path(rCtx.namespace.Name().String()))
if err != nil {
h.Logger.Error("failed to construct forwarding request URL", tag.Operation(apiName), tag.WorkflowNamespace(rCtx.namespace.Name().String()), tag.Error(err), tag.TargetCluster(rCtx.namespace.ActiveClusterName(namespace.RoutingKey{ID: rCtx.workflowID})))
h.Logger.Error("failed to construct forwarding request URL", tag.Operation(nexusCompletionAPIName), tag.WorkflowNamespace(rCtx.namespace.Name().String()), tag.Error(err), tag.TargetCluster(targetCluster))
return nexus.NewHandlerErrorf(nexus.HandlerErrorTypeInternal, "internal error")
}
if h.HTTPTraceProvider != nil {
traceLogger := log.With(h.Logger,
tag.Operation(apiName),
tag.Operation(nexusCompletionAPIName),
tag.WorkflowNamespace(rCtx.namespace.Name().String()),
tag.AttemptStart(time.Now().UTC()),
tag.SourceCluster(h.ClusterMetadata.GetCurrentClusterName()),
tag.TargetCluster(rCtx.namespace.ActiveClusterName(namespace.RoutingKey{ID: rCtx.workflowID})),
tag.TargetCluster(targetCluster),
)
if trace := h.HTTPTraceProvider.NewForwardingTrace(traceLogger); trace != nil {
ctx = httptrace.WithClientTrace(ctx, trace)
@@ -256,7 +385,6 @@ func (h *completionHandler) forwardCompleteOperation(ctx context.Context, r *nex
}
var completion nexusrpc.CompleteOperationOptions
switch r.State {
case nexus.OperationStateSucceeded:
completion = nexusrpc.CompleteOperationOptions{
@@ -290,6 +418,17 @@ func (h *completionHandler) forwardCompleteOperation(ctx context.Context, r *nex
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)
h.httpHandler.ServeHTTP(w, r)
})
r.Path(commonnexus.PathCompletionCallbackNoIdentifier).HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
r.Body = http.MaxBytesReader(w, r.Body, rpc.MaxNexusAPIRequestBodyBytes)
h.httpHandler.ServeHTTP(w, r)
})
}
type forwardingHTTPHeaderWrapper struct {
client *common.FrontendHTTPClient
originalRequestHeaders http.Header
@@ -302,17 +441,16 @@ func (f *forwardingHTTPHeaderWrapper) Do(req *http.Request) (*http.Response, err
req.Header.Set(k, v[0])
}
}
return f.client.Do(req)
}
type requestContext struct {
*completionHandler
*nexusCompletionHandler
logger log.Logger
metricsHandler metrics.Handler
metricsHandlerForInterceptors metrics.Handler
namespace *namespace.Namespace
workflowID string
businessID string
cleanupFunctions []func(error)
requestStartTime time.Time
outcomeTag metrics.Tag
@@ -326,7 +464,7 @@ func (c *requestContext) augmentContext(ctx context.Context, header http.Header)
ctx = interceptor.PopulateCallerInfo(
ctx,
func() string { return c.namespace.Name().String() },
func() string { return methodNameForMetrics },
func() string { return nexusCompletionMethodNameForMetrics },
)
if userAgent := header.Get(headerUserAgent); userAgent != "" {
// Preserve original strict behavior: only process if exactly one delimiter present.
@@ -355,7 +493,6 @@ func (c *requestContext) capturePanicAndRecordMetrics(ctxPtr *context.Context, e
}
st := string(debug.Stack())
c.logger.Error("Panic captured", tag.SysStackTrace(st), tag.Error(err))
*errPtr = err
}
@@ -415,7 +552,7 @@ func (c *requestContext) interceptRequest(ctx context.Context, request *nexusrpc
}
_, err = c.AuthInterceptor.Authorize(ctx, claims, &authorization.CallTarget{
APIName: apiName,
APIName: nexusCompletionAPIName,
Namespace: c.namespace.Name().String(),
Request: request,
})
@@ -433,21 +570,21 @@ func (c *requestContext) interceptRequest(ctx context.Context, request *nexusrpc
return commonnexus.ConvertGRPCError(err, false)
}
if err := c.NamespaceValidationInterceptor.ValidateState(c.namespace, apiName, c.workflowID); err != nil {
if err := c.NamespaceValidationInterceptor.ValidateState(c.namespace, nexusCompletionAPIName, c.businessID); err != nil {
c.outcomeTag = metrics.OutcomeTag("invalid_namespace_state")
return commonnexus.ConvertGRPCError(err, false)
}
// Redirect if current cluster is passive for this namespace.
if c.namespace.ActiveClusterName(namespace.RoutingKey{ID: c.workflowID}) != c.ClusterMetadata.GetCurrentClusterName() {
if c.shouldForwardRequest(ctx, request.HTTPRequest.Header, c.workflowID) {
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(methodNameForMetrics)
handler, forwardStartTime := c.RedirectionInterceptor.BeforeCall(nexusCompletionMethodNameForMetrics)
c.cleanupFunctions = append(c.cleanupFunctions, func(retErr error) {
c.RedirectionInterceptor.AfterCall(handler, forwardStartTime, c.namespace.ActiveClusterName(namespace.RoutingKey{ID: c.workflowID}), c.namespace.Name().String(), retErr)
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.workflowID}))
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")
@@ -459,26 +596,26 @@ func (c *requestContext) interceptRequest(ctx context.Context, request *nexusrpc
request,
"",
c.metricsHandlerForInterceptors,
[]tag.Tag{tag.Operation(methodNameForMetrics), tag.WorkflowNamespace(c.namespace.Name().String())},
[]tag.Tag{tag.Operation(nexusCompletionMethodNameForMetrics), tag.WorkflowNamespace(c.namespace.Name().String())},
retErr,
c.namespace.Name(),
)
}
})
cleanup, err := c.NamespaceConcurrencyLimitInterceptor.Allow(c.namespace.Name(), apiName, c.metricsHandlerForInterceptors, request)
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(c.namespace.Name(), apiName, request.HTTPRequest.Header); err != nil {
if err := c.NamespaceRateLimitInterceptor.Allow(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(apiName, request.HTTPRequest.Header); err != nil {
if err := c.RateLimitInterceptor.Allow(nexusCompletionAPIName, request.HTTPRequest.Header); err != nil {
c.outcomeTag = metrics.OutcomeTag("global_rate_limited")
return commonnexus.ConvertGRPCError(err, true)
}
@@ -505,5 +642,5 @@ func (c *requestContext) shouldForwardRequest(ctx context.Context, header http.H
return redirectAllowed &&
c.RedirectionInterceptor.RedirectionAllowed(ctx) &&
c.namespace.IsGlobalNamespace() &&
c.Config.ForwardingEnabledForNamespace(c.namespace.Name().String())
c.Config.EnableNamespaceNotActiveAutoForwarding(c.namespace.Name().String())
}

View File

@@ -21,6 +21,7 @@ import (
"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/resource"
"go.temporal.io/server/common/routing"
"go.temporal.io/server/common/rpc"
"go.temporal.io/server/common/rpc/interceptor"
@@ -31,7 +32,7 @@ import (
)
// Small wrapper that does some pre-processing before handing requests over to the Nexus SDK's HTTP handler.
type NexusHTTPHandler struct {
type NexusOperationHTTPHandler struct {
base nexusrpc.BaseHTTPHandler
logger log.Logger
nexusHandler http.Handler
@@ -45,9 +46,9 @@ type NexusHTTPHandler struct {
rateLimitInterceptor *interceptor.RateLimitInterceptor
}
func NewNexusHTTPHandler(
func NewNexusOperationHTTPHandler(
serviceConfig *Config,
matchingClient matchingservice.MatchingServiceClient,
matchingClient resource.MatchingClient,
metricsHandler metrics.Handler,
clusterMetadata cluster.Metadata,
clientCache *cluster.FrontendHTTPClientCache,
@@ -59,12 +60,12 @@ func NewNexusHTTPHandler(
redirectionInterceptor *interceptor.Redirection,
namespaceValidationInterceptor *interceptor.NamespaceValidatorInterceptor,
namespaceRateLimitInterceptor interceptor.NamespaceRateLimitInterceptor,
namespaceConcurrencyLimitIntercptor *interceptor.ConcurrentRequestLimitInterceptor,
namespaceConcurrencyLimitInterceptor *interceptor.ConcurrentRequestLimitInterceptor,
rateLimitInterceptor *interceptor.RateLimitInterceptor,
logger log.Logger,
httpTraceProvider commonnexus.HTTPClientTraceProvider,
) *NexusHTTPHandler {
return &NexusHTTPHandler{
) *NexusOperationHTTPHandler {
return &NexusOperationHTTPHandler{
base: nexusrpc.BaseHTTPHandler{
Logger: log.NewSlogLogger(logger),
FailureConverter: nexusrpc.DefaultFailureConverter(),
@@ -75,7 +76,7 @@ func NewNexusHTTPHandler(
auth: authInterceptor,
namespaceValidationInterceptor: namespaceValidationInterceptor,
namespaceRateLimitInterceptor: namespaceRateLimitInterceptor,
namespaceConcurrencyLimitInterceptor: namespaceConcurrencyLimitIntercptor,
namespaceConcurrencyLimitInterceptor: namespaceConcurrencyLimitInterceptor,
rateLimitInterceptor: rateLimitInterceptor,
preprocessErrorCounter: metricsHandler.Counter(metrics.NexusRequestPreProcessErrors.Name()).Record,
nexusHandler: nexusrpc.NewHTTPHandler(nexusrpc.HandlerOptions{
@@ -84,7 +85,7 @@ func NewNexusHTTPHandler(
metricsHandler: metricsHandler,
clusterMetadata: clusterMetadata,
namespaceRegistry: namespaceRegistry,
matchingClient: matchingClient,
matchingClient: matchingservice.MatchingServiceClient(matchingClient),
auth: authInterceptor,
telemetryInterceptor: telemetryInterceptor,
requestErrorHandler: requestErrorHandler,
@@ -104,20 +105,20 @@ func NewNexusHTTPHandler(
}
}
func (h *NexusHTTPHandler) RegisterRoutes(r *mux.Router) {
func (h *NexusOperationHTTPHandler) RegisterRoutes(r *mux.Router) {
r.PathPrefix("/" + commonnexus.RouteDispatchNexusTaskByNamespaceAndTaskQueue.Representation() + "/").
HandlerFunc(h.dispatchNexusTaskByNamespaceAndTaskQueue)
r.PathPrefix("/" + commonnexus.RouteDispatchNexusTaskByEndpoint.Representation() + "/").
HandlerFunc(h.dispatchNexusTaskByEndpoint)
}
func (h *NexusHTTPHandler) writeFailure(writer http.ResponseWriter, r *http.Request, err error) {
func (h *NexusOperationHTTPHandler) writeFailure(writer http.ResponseWriter, r *http.Request, err error) {
h.preprocessErrorCounter.Record(1)
h.base.WriteFailure(writer, r, err)
}
// Handler for [nexushttp.RouteSet.DispatchNexusTaskByNamespaceAndTaskQueue].
func (h *NexusHTTPHandler) dispatchNexusTaskByNamespaceAndTaskQueue(w http.ResponseWriter, r *http.Request) {
func (h *NexusOperationHTTPHandler) dispatchNexusTaskByNamespaceAndTaskQueue(w http.ResponseWriter, r *http.Request) {
var err error
nc := h.baseNexusContext(configs.DispatchNexusTaskByNamespaceAndTaskQueueAPIName, r.Header)
params := prepareRequest(commonnexus.RouteDispatchNexusTaskByNamespaceAndTaskQueue, w, r)
@@ -138,7 +139,7 @@ func (h *NexusHTTPHandler) dispatchNexusTaskByNamespaceAndTaskQueue(w http.Respo
return
}
rWithAuthCtx, err := h.parseTlsAndAuthInfo(r, nc)
rWithAuthCtx, err := h.parseTLSAndAuthInfo(r, nc)
if err != nil {
h.logger.Error("failed to get claims", tag.Error(err))
h.writeFailure(w, r, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnauthenticated, "unauthorized"))
@@ -157,7 +158,7 @@ func (h *NexusHTTPHandler) dispatchNexusTaskByNamespaceAndTaskQueue(w http.Respo
}
// Handler for [nexushttp.RouteSet.DispatchNexusTaskByEndpoint].
func (h *NexusHTTPHandler) dispatchNexusTaskByEndpoint(w http.ResponseWriter, r *http.Request) {
func (h *NexusOperationHTTPHandler) dispatchNexusTaskByEndpoint(w http.ResponseWriter, r *http.Request) {
endpointIDEscaped := prepareRequest(commonnexus.RouteDispatchNexusTaskByEndpoint, w, r)
endpointID, err := url.PathUnescape(endpointIDEscaped)
@@ -198,7 +199,7 @@ func (h *NexusHTTPHandler) dispatchNexusTaskByEndpoint(w http.ResponseWriter, r
return
}
rWithAuthCtx, err := h.parseTlsAndAuthInfo(r, nc)
rWithAuthCtx, err := h.parseTLSAndAuthInfo(r, nc)
if err != nil {
h.logger.Error("failed to get claims", tag.Error(err))
h.writeFailure(w, r, nexus.NewHandlerErrorf(nexus.HandlerErrorTypeUnauthenticated, "unauthorized"))
@@ -216,7 +217,7 @@ func (h *NexusHTTPHandler) dispatchNexusTaskByEndpoint(w http.ResponseWriter, r
h.serveResolvedURL(w, r, u, nc)
}
func (h *NexusHTTPHandler) baseNexusContext(apiName string, header http.Header) *nexusContext {
func (h *NexusOperationHTTPHandler) baseNexusContext(apiName string, header http.Header) *nexusContext {
return &nexusContext{
namespaceValidationInterceptor: h.namespaceValidationInterceptor,
namespaceRateLimitInterceptor: h.namespaceRateLimitInterceptor,
@@ -233,7 +234,7 @@ func (h *NexusHTTPHandler) baseNexusContext(apiName string, header http.Header)
// endpoint is valid for dispatching.
// For security reasons, at the moment only worker target endpoints are considered valid, in the future external
// endpoints may also be supported.
func (h *NexusHTTPHandler) nexusContextFromEndpoint(entry *persistencespb.NexusEndpointEntry, w http.ResponseWriter, r *http.Request) (*nexusContext, bool) {
func (h *NexusOperationHTTPHandler) nexusContextFromEndpoint(entry *persistencespb.NexusEndpointEntry, w http.ResponseWriter, r *http.Request) (*nexusContext, bool) {
switch v := entry.Endpoint.Spec.GetTarget().GetVariant().(type) {
case *persistencespb.NexusEndpointTarget_Worker_:
nsName, err := h.namespaceRegistry.GetNamespaceName(namespace.ID(v.Worker.GetNamespaceId()))
@@ -273,7 +274,7 @@ func prepareRequest[T any](route routing.Route[T], w http.ResponseWriter, r *htt
return route.Deserialize(vars)
}
func (h *NexusHTTPHandler) parseTlsAndAuthInfo(r *http.Request, nc *nexusContext) (*http.Request, error) {
func (h *NexusOperationHTTPHandler) parseTLSAndAuthInfo(r *http.Request, nc *nexusContext) (*http.Request, error) {
var tlsInfo *credentials.TLSInfo
if r.TLS != nil {
tlsInfo = &credentials.TLSInfo{
@@ -299,7 +300,7 @@ func (h *NexusHTTPHandler) parseTlsAndAuthInfo(r *http.Request, nc *nexusContext
return r, nil
}
func (h *NexusHTTPHandler) serveResolvedURL(w http.ResponseWriter, r *http.Request, u *url.URL, nc *nexusContext) {
func (h *NexusOperationHTTPHandler) serveResolvedURL(w http.ResponseWriter, r *http.Request, u *url.URL, nc *nexusContext) {
// Attach Nexus context to response writer and request context.
nc.originalRequestHeaders = r.Header.Clone()
w = newNexusHTTPResponseWriter(w, nc)

View File

@@ -44,12 +44,12 @@ func (f *fakeNamespaceRegistry) GetNamespaceName(id namespace.ID) (namespace.Nam
return f.getNamespaceName(id)
}
func newTestNexusHTTPHandler(
func newTestNexusOperationHTTPHandler(
endpointRegistry commonnexus.EndpointRegistry,
namespaceRegistry namespace.Registry,
) (*NexusHTTPHandler, *mux.Router) {
) (*NexusOperationHTTPHandler, *mux.Router) {
logger := log.NewTestLogger()
h := &NexusHTTPHandler{
h := &NexusOperationHTTPHandler{
base: nexusrpc.BaseHTTPHandler{
Logger: log.NewSlogLogger(logger),
FailureConverter: nexusrpc.DefaultFailureConverter(),
@@ -79,7 +79,7 @@ func TestDispatchNexusTaskByEndpoint_NotFound_NonRetryable(t *testing.T) {
return nil, serviceerror.NewNotFound("endpoint not found")
},
}
_, router := newTestNexusHTTPHandler(reg, nil)
_, router := newTestNexusOperationHTTPHandler(reg, nil)
rec := doNexusHTTPRequest(t, router, "test-endpoint-id")
@@ -97,7 +97,7 @@ func TestDispatchNexusTaskByEndpoint_NotFound_Retryable(t *testing.T) {
return nil, &retryableNotFoundError{msg: "endpoint temporarily unavailable"}
},
}
_, router := newTestNexusHTTPHandler(reg, nil)
_, router := newTestNexusOperationHTTPHandler(reg, nil)
rec := doNexusHTTPRequest(t, router, "test-endpoint-id")
@@ -138,7 +138,7 @@ func TestDispatchNexusTaskByEndpoint_NamespaceNotFound_Retryable(t *testing.T) {
},
}
_, router := newTestNexusHTTPHandler(reg, nsReg)
_, router := newTestNexusOperationHTTPHandler(reg, nsReg)
rec := doNexusHTTPRequest(t, router, "test-endpoint-id")

View File

@@ -447,6 +447,11 @@ func (e *ChasmEngine) applyUpdateWithLease(
serializedRef, err := mutableContext.Ref(component)
if err != nil {
if errors.As(err, new(*serviceerror.NotFound)) {
// The update may legitimately delete the addressed component, in which case
// there is no new ref to return even though the transition succeeded.
return nil, nil
}
return nil, serviceerror.NewInternalf("componentRef: %+v: %s", ref, err)
}

View File

@@ -2252,6 +2252,7 @@ func (h *Handler) CompleteNexusOperationChasm(
completion := &persistencespb.ChasmNexusCompletion{
CloseTime: request.CloseTime,
RequestId: request.Completion.RequestId,
Links: request.Links,
}
switch variant := request.Outcome.(type) {
case *historyservice.CompleteNexusOperationChasmRequest_Failure:

View File

@@ -524,9 +524,6 @@ func (s *NexusWorkflowTestSuite) TestNexusOperationSyncCompletion_LargePayload(c
}
func (s *NexusWorkflowTestSuite) TestNexusOperationAsyncCompletion(chasmEnabled bool) {
if chasmEnabled {
s.T().Skip("Blocked on CHASM Nexus async completion support")
}
env := s.newNexusWorkflowTestEnv(chasmEnabled)
ctx := testcore.NewContext()
taskQueue := testcore.RandomizeStr(s.T().Name())
@@ -722,35 +719,103 @@ func (s *NexusWorkflowTestSuite) TestNexusOperationAsyncCompletion(chasmEnabled
completionToken, err := gen.DecodeCompletion(decodedToken)
s.NoError(err)
// Request fails if the workflow is not found.
workflowNotFoundToken := common.CloneProto(completionToken)
workflowNotFoundToken.WorkflowId = "not-found"
callbackToken, err = gen.Tokenize(workflowNotFoundToken)
s.NoError(err)
completion.Header = nexus.Header{commonnexus.CallbackTokenHeader: callbackToken}
assertInvalidCompletionTokenRejected := func(
caseName string,
mutate func(*tokenspb.NexusOperationCompletion) *tokenspb.NexusOperationCompletion,
expectedErrorType nexus.HandlerErrorType,
) {
s.T().Helper()
capture = env.StartNamespaceMetricCapture()
capture = env.StartNamespaceMetricCapture()
err = s.sendNexusCompletionRequest(ctx, publicCallbackURL, completion)
completionRequests = capture.Metric("nexus_completion_requests")
s.ErrorAs(err, &handlerErr)
s.Equal(nexus.HandlerErrorTypeNotFound, handlerErr.Type)
s.Len(completionRequests, 1)
s.Subset(completionRequests[0].Tags, map[string]string{"namespace": env.Namespace().String(), "outcome": "error_not_found"})
// Mutate the mutatedCompletionToken and tokenize it to get a callback mutatedCompletionToken.
mutatedCompletionToken := mutate(common.CloneProto(completionToken))
mutatedCallbackToken, err := gen.Tokenize(mutatedCompletionToken)
s.NoError(err, caseName)
completion.Header = nexus.Header{commonnexus.CallbackTokenHeader: mutatedCallbackToken}
// Request fails if the state machine reference is stale.
staleToken := common.CloneProto(completionToken)
staleToken.Ref.MachineInitialVersionedTransition.NamespaceFailoverVersion++
callbackToken, err = gen.Tokenize(staleToken)
s.NoError(err)
completion.Header = nexus.Header{commonnexus.CallbackTokenHeader: callbackToken}
// Send the completion request and verify the error type.
err = s.sendNexusCompletionRequest(ctx, publicCallbackURL, completion)
var handlerErr *nexus.HandlerError
s.ErrorAs(err, &handlerErr, caseName)
s.Equal(expectedErrorType, handlerErr.Type, caseName)
capture = env.StartNamespaceMetricCapture()
err = s.sendNexusCompletionRequest(ctx, publicCallbackURL, completion)
completionRequests = capture.Metric("nexus_completion_requests")
s.ErrorAs(err, &handlerErr)
s.Equal(nexus.HandlerErrorTypeNotFound, handlerErr.Type)
s.Len(completionRequests, 1)
s.Subset(completionRequests[0].Tags, map[string]string{"namespace": env.Namespace().String(), "outcome": "error_not_found"})
// Verify metrics.
completionRequests = capture.Metric("nexus_completion_requests")
if expectedErrorType == nexus.HandlerErrorTypeNotFound {
s.Len(completionRequests, 1, caseName)
s.Subset(completionRequests[0].Tags, map[string]string{"outcome": "error_not_found"}, caseName)
} else {
s.Empty(completionRequests, caseName)
}
}
if chasmEnabled {
assertInvalidCompletionTokenRejected(
"missing execution",
func(token *tokenspb.NexusOperationCompletion) *tokenspb.NexusOperationCompletion {
s.mutateCompletionComponentRef(token, func(ref *persistencespb.ChasmComponentRef) {
ref.BusinessId = "not-found"
})
return token
},
nexus.HandlerErrorTypeNotFound,
)
assertInvalidCompletionTokenRejected(
"missing run",
func(token *tokenspb.NexusOperationCompletion) *tokenspb.NexusOperationCompletion {
s.mutateCompletionComponentRef(token, func(ref *persistencespb.ChasmComponentRef) {
ref.RunId = uuid.NewString()
})
return token
},
nexus.HandlerErrorTypeNotFound,
)
assertInvalidCompletionTokenRejected(
"wrong archetype",
func(token *tokenspb.NexusOperationCompletion) *tokenspb.NexusOperationCompletion {
s.mutateCompletionComponentRef(token, func(ref *persistencespb.ChasmComponentRef) {
ref.ArchetypeId = chasm.SchedulerArchetypeID
})
return token
},
nexus.HandlerErrorTypeNotFound,
)
assertInvalidCompletionTokenRejected(
"empty component ref",
func(token *tokenspb.NexusOperationCompletion) *tokenspb.NexusOperationCompletion {
token.ComponentRef = nil
return token
},
nexus.HandlerErrorTypeBadRequest,
)
assertInvalidCompletionTokenRejected(
"malformed component ref",
func(token *tokenspb.NexusOperationCompletion) *tokenspb.NexusOperationCompletion {
token.ComponentRef = []byte("not-a-proto")
return token
},
nexus.HandlerErrorTypeBadRequest,
)
} else {
assertInvalidCompletionTokenRejected(
"workflow not found",
func(token *tokenspb.NexusOperationCompletion) *tokenspb.NexusOperationCompletion {
token.WorkflowId = "not-found"
return token
},
nexus.HandlerErrorTypeNotFound,
)
// Request fails if the state machine reference is stale.
assertInvalidCompletionTokenRejected(
"stale state machine ref",
func(token *tokenspb.NexusOperationCompletion) *tokenspb.NexusOperationCompletion {
token.Ref.MachineInitialVersionedTransition.NamespaceFailoverVersion++
return token
},
nexus.HandlerErrorTypeNotFound,
)
}
callbackToken, err = gen.Tokenize(completionToken)
s.NoError(err)
@@ -868,7 +933,7 @@ func (s *NexusWorkflowTestSuite) TestNexusOperationAsyncCompletion(chasmEnabled
func (s *NexusWorkflowTestSuite) TestNexusOperationAsyncCompletionBeforeStart(chasmEnabled bool) {
if chasmEnabled {
s.T().Skip("Blocked on CHASM Nexus async completion before start support")
s.T().Skip("Blocked on CHASM Nexus completion-before-start support")
}
env := s.newNexusWorkflowTestEnv(chasmEnabled)
ctx := testcore.NewContext()
@@ -1121,9 +1186,6 @@ func (s *NexusWorkflowTestSuite) TestNexusOperationAsyncCompletionBeforeStart(ch
}
func (s *NexusWorkflowTestSuite) TestNexusOperationAsyncFailure(chasmEnabled bool) {
if chasmEnabled {
s.T().Skip("Blocked on CHASM Nexus async completion support")
}
env := s.newNexusWorkflowTestEnv(chasmEnabled)
ctx := testcore.NewContext()
taskQueue := testcore.RandomizeStr(s.T().Name())
@@ -3270,3 +3332,19 @@ func (s *NexusWorkflowTestSuite) TestNexusOperationSystemEndpoint(chasmEnabled b
s.NoError(run.Get(ctx, &response))
s.Equal("Hello, Temporal", response)
}
func (s *NexusWorkflowTestSuite) mutateCompletionComponentRef(
token *tokenspb.NexusOperationCompletion,
mutate func(*persistencespb.ChasmComponentRef),
) {
s.T().Helper()
ref := &persistencespb.ChasmComponentRef{}
s.NoError(ref.Unmarshal(token.GetComponentRef()))
mutate(ref)
var err error
token.ComponentRef, err = ref.Marshal()
s.NoError(err)
}