diff --git a/converter/transfer_types.go b/converter/transfer_types.go index 26b50f580..6d58956c2 100644 --- a/converter/transfer_types.go +++ b/converter/transfer_types.go @@ -417,6 +417,20 @@ func (dc *transferAwareDataConverter) WithSerializationContext(ctx Serialization return &result } +// WithTransferWorkflowContext passes ctx to workflow transfer conversion callbacks +// without calling the underlying data converter's WithWorkflowContext method. +// dc must already be transfer-aware, as returned by [MakeTransferAware]. +// The underlying converter's workflow and serialization contexts are unchanged. +// The original converter is not modified. +// +// NOTE: Experimental. +func WithTransferWorkflowContext(dc DataConverter, ctx WorkflowContext) DataConverter { + result := *dc.(*transferAwareDataConverter) + result.context = nil + result.workflowContext = ctx + return &result +} + func (dc *transferAwareDataConverter) WithWorkflowContext(ctx WorkflowContext) DataConverter { result := *dc if parent, ok := dc.parent.(ContextAware); ok { diff --git a/internal/nexus_operation_registry.go b/internal/nexus_operation_registry.go new file mode 100644 index 000000000..eb1fb7883 --- /dev/null +++ b/internal/nexus_operation_registry.go @@ -0,0 +1,26 @@ +package internal + +import ( + "fmt" + + "go.temporal.io/sdk/converter" +) + +type NexusOperationKey struct { + Service, Operation string +} + +type NexusOperationRegistryEntry struct { + SerializationContext func(any) converter.SerializationContext +} + +var nexusOperationRegistry = make(map[NexusOperationKey]NexusOperationRegistryEntry) + +func RegisterNexusOperationRegistry(registry map[NexusOperationKey]NexusOperationRegistryEntry) { + for key, entry := range registry { + if _, ok := nexusOperationRegistry[key]; ok { + panic(fmt.Sprintf("Nexus operation registry already contains service %q operation %q", key.Service, key.Operation)) + } + nexusOperationRegistry[key] = entry + } +} diff --git a/internal/nexus_operation_registry_test.go b/internal/nexus_operation_registry_test.go new file mode 100644 index 000000000..bbe67b395 --- /dev/null +++ b/internal/nexus_operation_registry_test.go @@ -0,0 +1,239 @@ +package internal + +import ( + "context" + "errors" + "fmt" + "testing" + + "github.com/stretchr/testify/require" + commonpb "go.temporal.io/api/common/v1" + "go.temporal.io/sdk/converter" +) + +type registryModel struct{ Value string } + +func (registryModel) TransferTypeConverter() (converter.TransferTypeConverter, error) { + return converter.NewContextualTransferTypeConverter[registryModel, commonpb.Payload](nil, nil, + func(ctx Context, value *registryModel) (*commonpb.Payload, error) { + return GetDataConverterFromWorkflowContext(ctx).ToPayload(value.Value) + }, + func(ctx Context, payload *commonpb.Payload, value *registryModel) error { + return GetDataConverterFromWorkflowContext(ctx).FromPayload(payload, &value.Value) + }) +} + +// Binding records the converter installed in the workflow context. This catches +// accidental rebinding of the envelope parent to the inner payload context. +type registryBindingDC struct { + converter.DataConverter + scope converter.SerializationContext + binding Context + bound *[]*registryBindingDC +} + +func (dc *registryBindingDC) WithSerializationContext(sc converter.SerializationContext) converter.DataConverter { + result := *dc + result.scope = sc + result.DataConverter = converter.WithDataConverterSerializationContext(dc.DataConverter, sc) + return &result +} +func (dc *registryBindingDC) WithWorkflowContext(ctx Context) converter.DataConverter { + result := *dc + result.binding = ctx + *result.bound = append(*result.bound, &result) + return &result +} +func (dc *registryBindingDC) WithContext(context.Context) converter.DataConverter { return dc } + +func TestNexusOperationRegistry(t *testing.T) { + service := "registry-test" + registry := map[NexusOperationKey]NexusOperationRegistryEntry{} + calls := 0 + for _, operation := range []string{"first", "second", "wrapped"} { + registry[NexusOperationKey{service, operation}] = NexusOperationRegistryEntry{ + SerializationContext: func(input any) converter.SerializationContext { + calls++ + return converter.WorkflowSerializationContext{Namespace: operation, WorkflowID: operation + ":" + input.(registryModel).Value} + }, + } + } + RegisterNexusOperationRegistry(registry) + t.Cleanup(func() { + for key := range registry { + delete(nexusOperationRegistry, key) + } + }) + env := new(WorkflowUnitTest).NewTestWorkflowEnvironment() + var bound []*registryBindingDC + env.SetDataConverter(®istryBindingDC{ + DataConverter: converter.NewCodecDataConverter(converter.GetDefaultDataConverter(), &serCtxSigningCodec{}), + bound: &bound, + }) + fc := newNexusCapturingFailureConverter() + env.SetFailureConverter(fc) + interceptor, ctx, err := newWorkflowContext(env.impl, env.impl.GetRegistry().interceptors) + require.NoError(t, err) + capture := &captureNexusSerializationEnv{WorkflowEnvironment: interceptor.env} + interceptor.env = capture + var futures []NexusOperationFuture + var results [3]registryModel + d, _ := newDispatcher(ctx, interceptor, func(ctx Context) { + for i, operation := range []string{"first", "second", "wrapped"} { + client := NewSystemNexusClient(service) + if i == 1 { + client = NewNexusClient("temporal-system", service) + } else if i == 2 { + client = NewNexusClient("ordinary", service) + } + callCtx := ctx + if i == 2 { + wrapper := ®istryWrappingInterceptor{WorkflowOutboundInterceptorBase: WorkflowOutboundInterceptorBase{Next: interceptor}} + callCtx = WithValue(ctx, workflowInterceptorContextKey, wrapper) + } + futures = append(futures, client.ExecuteOperation(callCtx, operation, registryModel{Value: "target"}, NexusOperationOptions{})) + } + for i := len(capture.calls) - 1; i >= 0; i-- { + call := capture.calls[i] + dc := call.params.dataConverter + var outer *registryBindingDC + for _, binding := range bound { + if scope, ok := binding.scope.(converter.NexusSerializationContext); ok && scope.Operation == call.params.operation { + require.Same(t, getWorkflowEnvOptions(ctx).DataConverter, getWorkflowEnvOptions(binding.binding).DataConverter) + outer = binding + } + } + require.NotNil(t, outer) + var wire commonpb.Payload + require.NoError(t, outer.FromPayload(call.params.input, &wire)) + require.Equal(t, call.params.operation+":target", string(wire.Metadata["ctx-signature"])) + payload, err := dc.ToPayload(registryModel{Value: "result"}) + require.NoError(t, err) + call.started("token", nil) + call.completed(payload, nil) + failure := call.params.failureConverter.ErrorToFailure(errors.New("failure")) + require.Error(t, call.params.failureConverter.FailureToError(failure)) + } + for i, future := range futures { + if i == 2 { + require.IsType(t, ®istryWrappedFuture{}, future) + } + require.NoError(t, future.Get(ctx, &results[i])) + } + }, func() bool { return false }) + d.interceptor = interceptor + defer d.Close() + requireNoExecuteErr(t, d.ExecuteUntilAllBlocked(defaultDeadlockDetectionTimeout)) + require.Equal(t, 3, calls) + require.Equal(t, [3]registryModel{{"result"}, {"result"}, {"result"}}, results) + for i, conversion := range fc.captured() { + operation := []string{"wrapped", "second", "first"}[i/2] + require.Equal(t, converter.WorkflowSerializationContext{Namespace: operation, WorkflowID: operation + ":target"}, conversion.context) + } +} + +type registryWrappedFuture struct { + NexusOperationFuture +} + +type registryWrappingInterceptor struct { + WorkflowOutboundInterceptorBase +} + +func (i *registryWrappingInterceptor) ExecuteNexusOperation(ctx Context, input ExecuteNexusOperationInput) NexusOperationFuture { + return ®istryWrappedFuture{NexusOperationFuture: i.Next.ExecuteNexusOperation(ctx, input)} +} + +type registryInputInterceptor struct { + WorkflowOutboundInterceptorBase + seen any +} + +func (i *registryInputInterceptor) ExecuteNexusOperation(ctx Context, input ExecuteNexusOperationInput) NexusOperationFuture { + i.seen = input.Input + input.Input = registryModel{Value: "replacement"} + return i.Next.ExecuteNexusOperation(ctx, input) +} + +func TestNexusOperationRegistryAfterInterceptors(t *testing.T) { + key := NexusOperationKey{Service: "registry-interceptor", Operation: "operation"} + var selected any + registry := map[NexusOperationKey]NexusOperationRegistryEntry{key: { + SerializationContext: func(input any) converter.SerializationContext { + selected = input + return converter.WorkflowSerializationContext{WorkflowID: input.(registryModel).Value} + }, + }} + RegisterNexusOperationRegistry(registry) + t.Cleanup(func() { delete(nexusOperationRegistry, key) }) + env := new(WorkflowUnitTest).NewTestWorkflowEnvironment() + env.SetDataConverter(converter.NewCodecDataConverter(converter.GetDefaultDataConverter(), &serCtxSigningCodec{})) + interceptor, ctx, err := newWorkflowContext(env.impl, nil) + require.NoError(t, err) + capture := &captureNexusSerializationEnv{WorkflowEnvironment: interceptor.env} + interceptor.env = capture + replacer := ®istryInputInterceptor{WorkflowOutboundInterceptorBase: WorkflowOutboundInterceptorBase{Next: interceptor}} + d, _ := newDispatcher(ctx, interceptor, func(ctx Context) { + ctx = WithValue(ctx, workflowInterceptorContextKey, replacer) + future := NewSystemNexusClient(key.Service).ExecuteOperation(ctx, key.Operation, registryModel{Value: "original"}, NexusOperationOptions{}) + require.Len(t, capture.calls, 1) + call := capture.calls[0] + var wire commonpb.Payload + require.NoError(t, call.params.dataConverter.FromPayload(call.params.input, &wire)) + require.Equal(t, "replacement", string(wire.Metadata["ctx-signature"])) + var value string + targetDC := withRootDataConverterSerializationContext(ctx, converter.WorkflowSerializationContext{WorkflowID: "replacement"}) + require.NoError(t, targetDC.FromPayload(&wire, &value)) + require.Equal(t, "replacement", value) + call.started("token", nil) + call.completed(nil, nil) + require.NoError(t, future.Get(ctx, nil)) + require.NoError(t, future.GetNexusOperationExecution().Get(ctx, nil)) + }, func() bool { return false }) + d.interceptor = interceptor + defer d.Close() + requireNoExecuteErr(t, d.ExecuteUntilAllBlocked(defaultDeadlockDetectionTimeout)) + require.Equal(t, registryModel{Value: "original"}, replacer.seen) + require.Equal(t, registryModel{Value: "replacement"}, selected) +} + +func TestNexusOperationRegistryMissingEntry(t *testing.T) { + env := new(WorkflowUnitTest).NewTestWorkflowEnvironment() + interceptor, ctx, err := newWorkflowContext(env.impl, nil) + require.NoError(t, err) + for _, endpoint := range []string{systemNexusEndpoint, "temporal-system", "ordinary"} { + params, err := interceptor.prepareNexusOperationParams(ctx, ExecuteNexusOperationInput{ + Client: nexusClient{endpoint, t.Name()}, Operation: "missing", Input: "native", + }) + require.NoError(t, err) + var value string + require.NoError(t, params.dataConverter.FromPayload(params.input, &value)) + require.Equal(t, "native", value) + } +} + +func TestNexusOperationRegistryMergesEntries(t *testing.T) { + first := NexusOperationKey{Service: t.Name(), Operation: "first"} + second := NexusOperationKey{Service: t.Name(), Operation: "second"} + t.Cleanup(func() { + delete(nexusOperationRegistry, first) + delete(nexusOperationRegistry, second) + }) + sc := converter.WorkflowSerializationContext{WorkflowID: "target"} + entry := NexusOperationRegistryEntry{ + SerializationContext: func(any) converter.SerializationContext { return sc }, + } + registry := map[NexusOperationKey]NexusOperationRegistryEntry{first: entry} + RegisterNexusOperationRegistry(registry) + clear(registry) + RegisterNexusOperationRegistry(map[NexusOperationKey]NexusOperationRegistryEntry{second: entry}) + require.Equal(t, sc, nexusOperationRegistry[first].SerializationContext(registryModel{})) + require.Equal(t, sc, nexusOperationRegistry[second].SerializationContext(registryModel{})) + require.PanicsWithValue(t, + fmt.Sprintf("Nexus operation registry already contains service %q operation %q", first.Service, first.Operation), + func() { + RegisterNexusOperationRegistry(map[NexusOperationKey]NexusOperationRegistryEntry{first: entry}) + }, + ) + require.Equal(t, sc, nexusOperationRegistry[first].SerializationContext(registryModel{})) +} diff --git a/internal/transfer_types_test.go b/internal/transfer_types_test.go index 04e10967a..726f93c86 100644 --- a/internal/transfer_types_test.go +++ b/internal/transfer_types_test.go @@ -369,6 +369,44 @@ func TestTransferAwareDataConverter_ContextDelegation(t *testing.T) { }) } +func TestTransferAwareDataConverter_TransferWorkflowContext(t *testing.T) { + t.Parallel() + outerCtx := WithValue(Background(), ContextAwareDataConverterContextKey, "value") + innerCtx := WithValue(Background(), ContextAwareDataConverterContextKey, "inner") + innerCtx = WithValue(innerCtx, transferContextKey{}, "inner") + parent := NewContextAwareDataConverter(converter.NewCompositeDataConverter(converter.NewJSONPayloadConverter())) + dc := WithWorkflowContext(outerCtx, converter.MakeTransferAware(parent)) + + for _, tc := range []struct { + name string + dc converter.DataConverter + want string + }{ + { + name: "transfer callbacks only", + dc: converter.WithTransferWorkflowContext(dc, innerCtx), + want: `"wf:inner:?"`, + }, + { + name: "transfer callbacks and parent", + dc: WithWorkflowContext(innerCtx, dc), + want: `"wf:?:value"`, + }, + } { + t.Run(tc.name, func(t *testing.T) { + payload, err := tc.dc.ToPayload(contextualString("value")) + require.NoError(t, err) + require.Equal(t, tc.want, string(payload.GetData())) + + payload, err = converter.GetDefaultDataConverter().ToPayload("value") + require.NoError(t, err) + var got contextualString + require.NoError(t, tc.dc.FromPayload(payload, &got)) + require.Equal(t, contextualString("wf:inner:value"), got) + }) + } +} + func TestTransferAwareDataConverter_ConversionContext(t *testing.T) { t.Parallel() parent := converter.GetDefaultDataConverter() diff --git a/internal/workflow.go b/internal/workflow.go index e4911c89c..861ab3f37 100644 --- a/internal/workflow.go +++ b/internal/workflow.go @@ -3070,6 +3070,16 @@ func (wc *workflowEnvironmentInterceptor) prepareNexusOperationParams(ctx Contex dc := withRootDataConverterSerializationContext(ctx, nsc) fc := converter.WithFailureConverterSerializationContext(getRootFailureConverterFromWorkflowContext(ctx), nsc) + if info, ok := nexusOperationRegistry[NexusOperationKey{Service: nsc.Service, Operation: nsc.Operation}]; ok { + sc := info.SerializationContext(input.Input) + targetDC := withRootDataConverterSerializationContext(ctx, sc) + payloadContext := WithDataConverter(ctx, targetDC) + // Only transfer callbacks get the target context. The parent converter + // must retain the outer Nexus envelope's existing bindings. + dc = converter.WithTransferWorkflowContext(dc, payloadContext) + fc = converter.WithFailureConverterSerializationContext(getRootFailureConverterFromWorkflowContext(ctx), sc) + } + payload, err := dc.ToPayload(input.Input) if err != nil { return ExecuteNexusOperationParams{}, err