diff --git a/temporal-sdk/src/jackson3Tests/java/io/temporal/common/converter/Jackson3JsonPayloadConverterTest.java b/temporal-sdk/src/jackson3Tests/java/io/temporal/common/converter/Jackson3JsonPayloadConverterTest.java index 5d2485a7a9..84af8cabf8 100644 --- a/temporal-sdk/src/jackson3Tests/java/io/temporal/common/converter/Jackson3JsonPayloadConverterTest.java +++ b/temporal-sdk/src/jackson3Tests/java/io/temporal/common/converter/Jackson3JsonPayloadConverterTest.java @@ -3,9 +3,15 @@ import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertTrue; +import com.fasterxml.jackson.annotation.JsonSubTypes; +import com.fasterxml.jackson.annotation.JsonTypeInfo; import com.fasterxml.jackson.databind.ObjectMapper; +import com.google.common.reflect.TypeToken; import io.temporal.api.common.v1.Payload; +import java.lang.reflect.Type; import java.time.Instant; +import java.util.Collections; +import java.util.List; import java.util.Objects; import java.util.Optional; import org.junit.After; @@ -76,6 +82,18 @@ public void testEncodingType() { assertEquals("json/plain", converter.getEncodingType()); } + @Test + public void serializationUsesTypeHint() { + Jackson3JsonPayloadConverter converter = new Jackson3JsonPayloadConverter(); + Type type = new TypeToken>() {}.getType(); + + Payload payload = converter.toData(Collections.singletonList(new Cat("Milo")), type).get(); + List converted = converter.fromData(payload, List.class, type); + + assertTrue(converted.get(0) instanceof Cat); + assertEquals("Milo", ((Cat) converted.get(0)).getName()); + } + @Test public void testWireCompatibilityBetweenJackson2AndJackson3() { JacksonJsonPayloadConverter jackson2 = new JacksonJsonPayloadConverter(); @@ -213,4 +231,22 @@ public String toString() { + '}'; } } + + @JsonTypeInfo(use = JsonTypeInfo.Id.NAME, property = "type") + @JsonSubTypes(@JsonSubTypes.Type(value = Cat.class, name = "cat")) + private interface Animal {} + + private static class Cat implements Animal { + private String name; + + public Cat() {} + + Cat(String name) { + this.name = name; + } + + public String getName() { + return name; + } + } } diff --git a/temporal-sdk/src/main/java/io/temporal/client/ActivityClientImpl.java b/temporal-sdk/src/main/java/io/temporal/client/ActivityClientImpl.java index 77bf2bd794..958ae4f22e 100644 --- a/temporal-sdk/src/main/java/io/temporal/client/ActivityClientImpl.java +++ b/temporal-sdk/src/main/java/io/temporal/client/ActivityClientImpl.java @@ -67,10 +67,10 @@ public ActivityClientCallsInterceptor getInvoker() { @Override public ActivityHandle start( Class activityInterface, Functions.Proc1 activity, StartActivityOptions options) { - String activityType = - MethodExtractor.activityTypeName( - activityInterface, MethodExtractor.extract(activityInterface, activity)); - UntypedActivityHandle untyped = start(activityType, options, new Object[0]); + Method method = MethodExtractor.extract(activityInterface, activity); + String activityType = MethodExtractor.activityTypeName(activityInterface, method); + UntypedActivityHandle untyped = + start(activityType, options, method.getGenericParameterTypes(), new Object[0]); return ActivityHandle.fromUntyped(untyped, Void.class, null); } @@ -80,10 +80,10 @@ public ActivityHandle start( Functions.Proc2 activity, StartActivityOptions options, A1 arg1) { - String activityType = - MethodExtractor.activityTypeName( - activityInterface, MethodExtractor.extract(activityInterface, activity)); - UntypedActivityHandle untyped = start(activityType, options, arg1); + Method method = MethodExtractor.extract(activityInterface, activity); + String activityType = MethodExtractor.activityTypeName(activityInterface, method); + UntypedActivityHandle untyped = + start(activityType, options, method.getGenericParameterTypes(), arg1); return ActivityHandle.fromUntyped(untyped, Void.class, null); } @@ -94,10 +94,10 @@ public ActivityHandle start( StartActivityOptions options, A1 arg1, A2 arg2) { - String activityType = - MethodExtractor.activityTypeName( - activityInterface, MethodExtractor.extract(activityInterface, activity)); - UntypedActivityHandle untyped = start(activityType, options, arg1, arg2); + Method method = MethodExtractor.extract(activityInterface, activity); + String activityType = MethodExtractor.activityTypeName(activityInterface, method); + UntypedActivityHandle untyped = + start(activityType, options, method.getGenericParameterTypes(), arg1, arg2); return ActivityHandle.fromUntyped(untyped, Void.class, null); } @@ -109,10 +109,10 @@ public ActivityHandle start( A1 arg1, A2 arg2, A3 arg3) { - String activityType = - MethodExtractor.activityTypeName( - activityInterface, MethodExtractor.extract(activityInterface, activity)); - UntypedActivityHandle untyped = start(activityType, options, arg1, arg2, arg3); + Method method = MethodExtractor.extract(activityInterface, activity); + String activityType = MethodExtractor.activityTypeName(activityInterface, method); + UntypedActivityHandle untyped = + start(activityType, options, method.getGenericParameterTypes(), arg1, arg2, arg3); return ActivityHandle.fromUntyped(untyped, Void.class, null); } @@ -125,10 +125,10 @@ public ActivityHandle start( A2 arg2, A3 arg3, A4 arg4) { - String activityType = - MethodExtractor.activityTypeName( - activityInterface, MethodExtractor.extract(activityInterface, activity)); - UntypedActivityHandle untyped = start(activityType, options, arg1, arg2, arg3, arg4); + Method method = MethodExtractor.extract(activityInterface, activity); + String activityType = MethodExtractor.activityTypeName(activityInterface, method); + UntypedActivityHandle untyped = + start(activityType, options, method.getGenericParameterTypes(), arg1, arg2, arg3, arg4); return ActivityHandle.fromUntyped(untyped, Void.class, null); } @@ -142,10 +142,11 @@ public ActivityHandle start( A3 arg3, A4 arg4, A5 arg5) { - String activityType = - MethodExtractor.activityTypeName( - activityInterface, MethodExtractor.extract(activityInterface, activity)); - UntypedActivityHandle untyped = start(activityType, options, arg1, arg2, arg3, arg4, arg5); + Method method = MethodExtractor.extract(activityInterface, activity); + String activityType = MethodExtractor.activityTypeName(activityInterface, method); + UntypedActivityHandle untyped = + start( + activityType, options, method.getGenericParameterTypes(), arg1, arg2, arg3, arg4, arg5); return ActivityHandle.fromUntyped(untyped, Void.class, null); } @@ -160,11 +161,19 @@ public ActivityHandle start( A4 arg4, A5 arg5, A6 arg6) { - String activityType = - MethodExtractor.activityTypeName( - activityInterface, MethodExtractor.extract(activityInterface, activity)); + Method method = MethodExtractor.extract(activityInterface, activity); + String activityType = MethodExtractor.activityTypeName(activityInterface, method); UntypedActivityHandle untyped = - start(activityType, options, arg1, arg2, arg3, arg4, arg5, arg6); + start( + activityType, + options, + method.getGenericParameterTypes(), + arg1, + arg2, + arg3, + arg4, + arg5, + arg6); return ActivityHandle.fromUntyped(untyped, Void.class, null); } @@ -178,7 +187,8 @@ public ActivityHandle start( @SuppressWarnings("unchecked") Class resultClass = (Class) method.getReturnType(); Type resultType = method.getGenericReturnType(); - UntypedActivityHandle untyped = start(activityType, options, new Object[0]); + UntypedActivityHandle untyped = + start(activityType, options, method.getGenericParameterTypes(), new Object[0]); return ActivityHandle.fromUntyped(untyped, resultClass, resultType); } @@ -193,7 +203,8 @@ public ActivityHandle start( @SuppressWarnings("unchecked") Class resultClass = (Class) method.getReturnType(); Type resultType = method.getGenericReturnType(); - UntypedActivityHandle untyped = start(activityType, options, arg1); + UntypedActivityHandle untyped = + start(activityType, options, method.getGenericParameterTypes(), arg1); return ActivityHandle.fromUntyped(untyped, resultClass, resultType); } @@ -209,7 +220,8 @@ public ActivityHandle start( @SuppressWarnings("unchecked") Class resultClass = (Class) method.getReturnType(); Type resultType = method.getGenericReturnType(); - UntypedActivityHandle untyped = start(activityType, options, arg1, arg2); + UntypedActivityHandle untyped = + start(activityType, options, method.getGenericParameterTypes(), arg1, arg2); return ActivityHandle.fromUntyped(untyped, resultClass, resultType); } @@ -226,7 +238,8 @@ public ActivityHandle start( @SuppressWarnings("unchecked") Class resultClass = (Class) method.getReturnType(); Type resultType = method.getGenericReturnType(); - UntypedActivityHandle untyped = start(activityType, options, arg1, arg2, arg3); + UntypedActivityHandle untyped = + start(activityType, options, method.getGenericParameterTypes(), arg1, arg2, arg3); return ActivityHandle.fromUntyped(untyped, resultClass, resultType); } @@ -244,7 +257,8 @@ public ActivityHandle start( @SuppressWarnings("unchecked") Class resultClass = (Class) method.getReturnType(); Type resultType = method.getGenericReturnType(); - UntypedActivityHandle untyped = start(activityType, options, arg1, arg2, arg3, arg4); + UntypedActivityHandle untyped = + start(activityType, options, method.getGenericParameterTypes(), arg1, arg2, arg3, arg4); return ActivityHandle.fromUntyped(untyped, resultClass, resultType); } @@ -263,7 +277,9 @@ public ActivityHandle start( @SuppressWarnings("unchecked") Class resultClass = (Class) method.getReturnType(); Type resultType = method.getGenericReturnType(); - UntypedActivityHandle untyped = start(activityType, options, arg1, arg2, arg3, arg4, arg5); + UntypedActivityHandle untyped = + start( + activityType, options, method.getGenericParameterTypes(), arg1, arg2, arg3, arg4, arg5); return ActivityHandle.fromUntyped(untyped, resultClass, resultType); } @@ -284,7 +300,16 @@ public ActivityHandle start( Class resultClass = (Class) method.getReturnType(); Type resultType = method.getGenericReturnType(); UntypedActivityHandle untyped = - start(activityType, options, arg1, arg2, arg3, arg4, arg5, arg6); + start( + activityType, + options, + method.getGenericParameterTypes(), + arg1, + arg2, + arg3, + arg4, + arg5, + arg6); return ActivityHandle.fromUntyped(untyped, resultClass, resultType); } @@ -293,11 +318,20 @@ public ActivityHandle start( @Override public UntypedActivityHandle start( String activityType, StartActivityOptions options, @Nullable Object... args) { + return start(activityType, options, null, args); + } + + private UntypedActivityHandle start( + String activityType, + StartActivityOptions options, + @Nullable Type[] argTypes, + @Nullable Object... args) { ActivityClientCallsInterceptor.StartActivityOutput output = invoker.startActivity( new ActivityClientCallsInterceptor.StartActivityInput( activityType, Arrays.asList(args != null ? args : new Object[0]), + argTypes, options, propagatedHeader())); return new ActivityHandleImpl(output.getActivityId(), output.getActivityRunId(), invoker); diff --git a/temporal-sdk/src/main/java/io/temporal/client/SignalWithStartBatchRequest.java b/temporal-sdk/src/main/java/io/temporal/client/SignalWithStartBatchRequest.java index 8daf486ed7..5b7df9a518 100644 --- a/temporal-sdk/src/main/java/io/temporal/client/SignalWithStartBatchRequest.java +++ b/temporal-sdk/src/main/java/io/temporal/client/SignalWithStartBatchRequest.java @@ -2,6 +2,7 @@ import io.temporal.api.common.v1.WorkflowExecution; import io.temporal.workflow.Functions; +import java.lang.reflect.Type; import java.util.ArrayList; import java.util.List; import java.util.concurrent.atomic.AtomicBoolean; @@ -13,7 +14,9 @@ final class SignalWithStartBatchRequest implements BatchRequest { private WorkflowStub stub; private String signalName; private Object[] signalArgs; + private Type[] signalArgTypes; private Object[] startArgs; + private Type[] startArgTypes; private final AtomicBoolean invoked = new AtomicBoolean(); WorkflowExecution invoke() { @@ -34,18 +37,21 @@ WorkflowExecution invoke() { } private WorkflowExecution signalWithStart() { - return stub.signalWithStart(signalName, signalArgs, startArgs); + return stub.signalWithStartWithTypeHints( + signalName, signalArgs, signalArgTypes, startArgs, startArgTypes); } - void signal(WorkflowStub stub, String signalName, Object[] args) { + void signal(WorkflowStub stub, String signalName, Object[] args, Type[] argTypes) { setStub(stub); this.signalName = signalName; this.signalArgs = args; + this.signalArgTypes = argTypes; } - void start(WorkflowStub stub, Object[] args) { + void start(WorkflowStub stub, Object[] args, Type[] argTypes) { setStub(stub); this.startArgs = args; + this.startArgTypes = argTypes; } private void setStub(WorkflowStub stub) { diff --git a/temporal-sdk/src/main/java/io/temporal/client/WorkflowInvocationHandler.java b/temporal-sdk/src/main/java/io/temporal/client/WorkflowInvocationHandler.java index cd14942a66..bfba2f748f 100644 --- a/temporal-sdk/src/main/java/io/temporal/client/WorkflowInvocationHandler.java +++ b/temporal-sdk/src/main/java/io/temporal/client/WorkflowInvocationHandler.java @@ -19,6 +19,7 @@ import io.temporal.workflow.*; import java.lang.reflect.InvocationHandler; import java.lang.reflect.Method; +import java.lang.reflect.Type; import java.util.*; import javax.annotation.Nullable; @@ -184,14 +185,14 @@ public Object invoke(Object proxy, Method method, Object[] args) throws Throwabl return Defaults.defaultValue(method.getReturnType()); } - private static void startWorkflow(WorkflowStub untyped, Object[] args) { + private static void startWorkflow(WorkflowStub untyped, Method method, Object[] args) { Optional options = untyped.getOptions(); if (untyped.getExecution() == null || (options.isPresent() && options.get().getWorkflowIdReusePolicy() == WorkflowIdReusePolicy.WORKFLOW_ID_REUSE_POLICY_ALLOW_DUPLICATE)) { try { - untyped.start(args); + untyped.startWithTypeHints(method.getGenericParameterTypes(), args); } catch (WorkflowExecutionAlreadyStarted e) { // We do allow duplicated calls if policy is not AllowDuplicate. Semantic is to wait for // result. @@ -241,7 +242,7 @@ public void invoke( throw new IllegalArgumentException( "WorkflowClient.start can be called only on a method annotated with @WorkflowMethod"); } - result = untyped.start(args); + result = untyped.startWithTypeHints(method.getGenericParameterTypes(), args); } @Override @@ -298,7 +299,7 @@ private void signalWorkflow( throw new IllegalArgumentException("Signal method must have void return type: " + method); } String signalName = methodMetadata.getName(); - untyped.signal(signalName, args); + untyped.signalWithTypeHints(signalName, method.getGenericParameterTypes(), args); } private Object queryWorkflow( @@ -310,7 +311,12 @@ private Object queryWorkflow( throw new IllegalArgumentException("Query method cannot have void return type: " + method); } String queryType = methodMetadata.getName(); - return untyped.query(queryType, method.getReturnType(), method.getGenericReturnType(), args); + return untyped.queryWithTypeHints( + queryType, + method.getReturnType(), + method.getGenericReturnType(), + method.getGenericParameterTypes(), + args); } private Object updateWorkflow( @@ -319,12 +325,17 @@ private Object updateWorkflow( Method method, Object[] args) { String updateType = methodMetadata.getName(); - return untyped.update(updateType, method.getReturnType(), args); + return untyped.updateWithTypeHints( + updateType, + method.getReturnType(), + method.getGenericReturnType(), + method.getGenericParameterTypes(), + args); } @SuppressWarnings("FutureReturnValueIgnored") private Object startWorkflow(WorkflowStub untyped, Method method, Object[] args) { - WorkflowInvocationHandler.startWorkflow(untyped, args); + WorkflowInvocationHandler.startWorkflow(untyped, method, args); return untyped.getResult(method.getReturnType(), method.getGenericReturnType()); } } @@ -349,7 +360,7 @@ public void invoke( throw new IllegalArgumentException( "WorkflowClient.execute can be called only on a method annotated with @WorkflowMethod"); } - WorkflowInvocationHandler.startWorkflow(untyped, args); + WorkflowInvocationHandler.startWorkflow(untyped, method, args); result = untyped.getResultAsync(method.getReturnType(), method.getGenericReturnType()); } @@ -389,10 +400,10 @@ public void invoke( throw new IllegalArgumentException( "SignalWithStart batch doesn't accept methods annotated with @UpdateMethod"); case WORKFLOW: - batch.start(untyped, args); + batch.start(untyped, args, method.getGenericParameterTypes()); break; case SIGNAL: - batch.signal(untyped, methodMetadata.getName(), args); + batch.signal(untyped, methodMetadata.getName(), args, method.getGenericParameterTypes()); break; } } @@ -428,7 +439,9 @@ public void invoke( "Only on a method annotated with @WorkflowMethod can be used to start a Nexus operation."); } - result = createNexusBoundStub(untyped, request).start(args); + result = + createNexusBoundStub(untyped, request) + .startWithTypeHints(method.getGenericParameterTypes(), args); } @Override @@ -463,7 +476,11 @@ public void invoke( throw new IllegalArgumentException( "Only a method annotated with @UpdateMethod can be used to start an Update."); } - result = untyped.startUpdate(mergeUpdateOptions(options, workflowMetadata, method), args); + result = + untyped.startUpdateWithTypeHints( + mergeUpdateOptions(options, workflowMetadata, method), + method.getGenericParameterTypes(), + args); } @Override @@ -506,7 +523,9 @@ enum State { private final UpdateOptions userProvidedUpdateOptions; private Object[] updateArgs; + private Type[] updateArgTypes; private Object[] startArgs; + private Type[] startArgTypes; private UpdateOptions updateOptions; private final WithStartWorkflowOperation startOp; private State state = State.INIT; @@ -541,6 +560,7 @@ public void invoke( } this.setStub(untyped); this.updateArgs = args; + this.updateArgTypes = method.getGenericParameterTypes(); this.updateOptions = UpdateInvocationHandler.mergeUpdateOptions( userProvidedUpdateOptions, workflowMetadata, method); @@ -556,11 +576,14 @@ public void invoke( } this.setStub(untyped); this.startArgs = args; + this.startArgTypes = method.getGenericParameterTypes(); this.startOp.setStub(untyped); this.startOp.setResultClass(method.getReturnType()); state = State.START_RECEIVED; - this.result = untyped.startUpdateWithStart(updateOptions, updateArgs, this.startArgs); + this.result = + untyped.startUpdateWithStartWithTypeHints( + updateOptions, updateArgs, updateArgTypes, this.startArgs, startArgTypes); } else { throw new IllegalArgumentException( "UpdateWithStartInvocationHandler called too many times"); diff --git a/temporal-sdk/src/main/java/io/temporal/client/WorkflowStub.java b/temporal-sdk/src/main/java/io/temporal/client/WorkflowStub.java index 6bc93b5e9d..390ea170db 100644 --- a/temporal-sdk/src/main/java/io/temporal/client/WorkflowStub.java +++ b/temporal-sdk/src/main/java/io/temporal/client/WorkflowStub.java @@ -54,6 +54,15 @@ static WorkflowStub fromTyped(T typed) { */ void signal(String signalName, Object... args); + /** + * Signals a workflow using declared argument types as serialization hints. + * + *

The default implementation delegates to {@link #signal(String, Object...)}. + */ + default void signalWithTypeHints(String signalName, Type[] argTypes, Object... args) { + signal(signalName, args); + } + /** * Synchronously update a workflow execution by invoking its update handler. Usually a update * handler is a method annotated with {@link io.temporal.workflow.UpdateMethod}. @@ -71,6 +80,16 @@ static WorkflowStub fromTyped(T typed) { */ R update(String updateName, Class resultClass, Object... args); + /** + * Updates a workflow using declared argument types as serialization hints. + * + *

The default implementation delegates to {@link #update(String, Class, Object...)}. + */ + default R updateWithTypeHints( + String updateName, Class resultClass, Type resultType, Type[] argTypes, Object... args) { + return update(updateName, resultClass, args); + } + /** * Asynchronously update a workflow execution by invoking its update handler and returning a * handle to the update request. Usually an update handler is a method annotated with {@link @@ -104,6 +123,16 @@ WorkflowUpdateHandle startUpdate( */ WorkflowUpdateHandle startUpdate(UpdateOptions options, Object... args); + /** + * Starts an update using declared argument types as serialization hints. + * + *

The default implementation delegates to {@link #startUpdate(UpdateOptions, Object...)}. + */ + default WorkflowUpdateHandle startUpdateWithTypeHints( + UpdateOptions options, Type[] argTypes, Object... args) { + return startUpdate(options, args); + } + /** * Get an update handle to a previously started update request. Getting an update handle does not * guarantee the update ID exists. @@ -131,6 +160,15 @@ WorkflowUpdateHandle getUpdateHandle( WorkflowExecution start(Object... args); + /** + * Starts a workflow using declared argument types as serialization hints. + * + *

The default implementation delegates to {@link #start(Object...)}. + */ + default WorkflowExecution startWithTypeHints(Type[] argTypes, Object... args) { + return start(args); + } + /** * Asynchronously update a workflow execution by invoking its update handler, and start the * workflow according to the option's {@link WorkflowIdConflictPolicy}. It returns a handle to the @@ -146,6 +184,21 @@ WorkflowUpdateHandle getUpdateHandle( WorkflowUpdateHandle startUpdateWithStart( UpdateOptions updateOptions, Object[] updateArgs, Object[] startArgs); + /** + * Starts an update and workflow using declared argument types as serialization hints. + * + *

The default implementation delegates to {@link #startUpdateWithStart(UpdateOptions, + * Object[], Object[])}. + */ + default WorkflowUpdateHandle startUpdateWithStartWithTypeHints( + UpdateOptions updateOptions, + Object[] updateArgs, + Type[] updateArgTypes, + Object[] startArgs, + Type[] startArgTypes) { + return startUpdateWithStart(updateOptions, updateArgs, startArgs); + } + /** * Synchronously update a workflow execution by invoking its update handler, and start the * workflow according to the option's {@link WorkflowIdConflictPolicy}. It returns the update @@ -170,6 +223,21 @@ R executeUpdateWithStart( */ WorkflowExecution signalWithStart(String signalName, Object[] signalArgs, Object[] startArgs); + /** + * Signals and starts a workflow using declared argument types as serialization hints. + * + *

The default implementation delegates to {@link #signalWithStart(String, Object[], + * Object[])}. + */ + default WorkflowExecution signalWithStartWithTypeHints( + String signalName, + Object[] signalArgs, + Type[] signalArgTypes, + Object[] startArgs, + Type[] startArgTypes) { + return signalWithStart(signalName, signalArgs, startArgs); + } + /** * @return workflow type name if it was provided when the stub was created. */ @@ -371,6 +439,16 @@ CompletableFuture getResultAsync( */ R query(String queryType, Class resultClass, Type resultType, Object... args); + /** + * Queries a workflow using declared argument types as serialization hints. + * + *

The default implementation delegates to {@link #query(String, Class, Type, Object...)}. + */ + default R queryWithTypeHints( + String queryType, Class resultClass, Type resultType, Type[] argTypes, Object... args) { + return query(queryType, resultClass, resultType, args); + } + /** * Request cancellation of a workflow execution. * diff --git a/temporal-sdk/src/main/java/io/temporal/client/WorkflowStubImpl.java b/temporal-sdk/src/main/java/io/temporal/client/WorkflowStubImpl.java index 69018ce65f..d0dfeba35a 100644 --- a/temporal-sdk/src/main/java/io/temporal/client/WorkflowStubImpl.java +++ b/temporal-sdk/src/main/java/io/temporal/client/WorkflowStubImpl.java @@ -75,19 +75,25 @@ class WorkflowStubImpl implements WorkflowStub { @Override public void signal(String signalName, Object... args) { + signalWithTypeHints(signalName, null, args); + } + + @Override + public void signalWithTypeHints(String signalName, Type[] argTypes, Object... args) { checkStarted(); WorkflowExecution targetExecution = currentExecutionCheckLegacy(); try { workflowClientInvoker.signal( new WorkflowClientCallsInterceptor.WorkflowSignalInput( - targetExecution, signalName, Header.empty(), args)); + targetExecution, signalName, Header.empty(), args, argTypes)); } catch (Exception e) { Throwable throwable = throwAsWorkflowFailureException(e, targetExecution); throw new WorkflowServiceException(targetExecution, workflowType.orElse(null), throwable); } } - private WorkflowExecution startWithOptions(WorkflowOptions options, Object... args) { + private WorkflowExecution startWithOptions( + WorkflowOptions options, @Nullable Type[] argTypes, Object... args) { checkExecutionIsNotStarted(); String workflowId = getWorkflowIdForStart(options); WorkflowExecution workflowExecution = null; @@ -95,7 +101,7 @@ private WorkflowExecution startWithOptions(WorkflowOptions options, Object... ar WorkflowClientCallsInterceptor.WorkflowStartOutput workflowStartOutput = workflowClientInvoker.start( new WorkflowClientCallsInterceptor.WorkflowStartInput( - workflowId, workflowType.get(), Header.empty(), args, options)); + workflowId, workflowType.get(), Header.empty(), args, argTypes, options)); workflowExecution = workflowStartOutput.getWorkflowExecution(); populateExecutionAfterStart(workflowExecution); return workflowExecution; @@ -115,15 +121,30 @@ private WorkflowExecution startWithOptions(WorkflowOptions options, Object... ar @Override public WorkflowExecution start(Object... args) { + return startWithTypeHints(null, args); + } + + @Override + public WorkflowExecution startWithTypeHints(Type[] argTypes, Object... args) { if (options == null) { throw new IllegalStateException("Required parameter WorkflowOptions is missing"); } - return startWithOptions(WorkflowOptions.merge(null, null, options), args); + return startWithOptions(WorkflowOptions.merge(null, null, options), argTypes, args); } @Override public WorkflowUpdateHandle startUpdateWithStart( UpdateOptions updateOptions, Object[] updateArgs, Object[] startArgs) { + return startUpdateWithStartWithTypeHints(updateOptions, updateArgs, null, startArgs, null); + } + + @Override + public WorkflowUpdateHandle startUpdateWithStartWithTypeHints( + UpdateOptions updateOptions, + Object[] updateArgs, + Type[] updateArgTypes, + Object[] startArgs, + Type[] startArgTypes) { if (options == null) { throw new IllegalStateException( "Required parameter WorkflowOptions is missing in WorkflowStub"); @@ -140,11 +161,12 @@ public WorkflowUpdateHandle startUpdateWithStart( // gather inputs WorkflowClientCallsInterceptor.WorkflowStartInput startInput = new WorkflowClientCallsInterceptor.WorkflowStartInput( - workflowId, workflowType.get(), Header.empty(), startArgs, options); + workflowId, workflowType.get(), Header.empty(), startArgs, startArgTypes, options); WorkflowClientCallsInterceptor.StartUpdateInput updateInput = startUpdateInput( updateOptions, updateArgs, + updateArgTypes, WorkflowExecution.newBuilder().setWorkflowId(workflowId).build()); WorkflowClientCallsInterceptor.WorkflowUpdateWithStartInput input = new WorkflowClientCallsInterceptor.WorkflowUpdateWithStartInput<>( @@ -180,7 +202,12 @@ public R executeUpdateWithStart( } private WorkflowExecution signalWithStartWithOptions( - WorkflowOptions options, String signalName, Object[] signalArgs, Object[] startArgs) { + WorkflowOptions options, + String signalName, + Object[] signalArgs, + @Nullable Type[] signalArgTypes, + Object[] startArgs, + @Nullable Type[] startArgTypes) { checkExecutionIsNotStarted(); String workflowId = getWorkflowIdForStart(options); WorkflowExecution workflowExecution = null; @@ -189,9 +216,15 @@ private WorkflowExecution signalWithStartWithOptions( workflowClientInvoker.signalWithStart( new WorkflowClientCallsInterceptor.WorkflowSignalWithStartInput( new WorkflowClientCallsInterceptor.WorkflowStartInput( - workflowId, workflowType.get(), Header.empty(), startArgs, options), + workflowId, + workflowType.get(), + Header.empty(), + startArgs, + startArgTypes, + options), signalName, - signalArgs)); + signalArgs, + signalArgTypes)); workflowExecution = workflowStartOutput.getWorkflowStartOutput().getWorkflowExecution(); populateExecutionAfterStart(workflowExecution); return workflowExecution; @@ -220,11 +253,26 @@ private static String getWorkflowIdForStart(WorkflowOptions options) { @Override public WorkflowExecution signalWithStart( String signalName, Object[] signalArgs, Object[] startArgs) { + return signalWithStartWithTypeHints(signalName, signalArgs, null, startArgs, null); + } + + @Override + public WorkflowExecution signalWithStartWithTypeHints( + String signalName, + Object[] signalArgs, + Type[] signalArgTypes, + Object[] startArgs, + Type[] startArgTypes) { if (options == null) { throw new IllegalStateException("Required parameter WorkflowOptions is missing"); } return signalWithStartWithOptions( - WorkflowOptions.merge(null, null, options), signalName, signalArgs, startArgs); + WorkflowOptions.merge(null, null, options), + signalName, + signalArgs, + signalArgTypes, + startArgs, + startArgTypes); } @Override @@ -318,6 +366,12 @@ public R query(String queryType, Class resultClass, Object... args) { @Override public R query(String queryType, Class resultClass, Type resultType, Object... args) { + return queryWithTypeHints(queryType, resultClass, resultType, null, args); + } + + @Override + public R queryWithTypeHints( + String queryType, Class resultClass, Type resultType, Type[] argTypes, Object... args) { checkStarted(); WorkflowClientCallsInterceptor.QueryOutput result; WorkflowExecution targetExecution = execution.get(); @@ -325,7 +379,13 @@ public R query(String queryType, Class resultClass, Type resultType, Obje result = workflowClientInvoker.query( new WorkflowClientCallsInterceptor.QueryInput<>( - targetExecution, queryType, Header.empty(), args, resultClass, resultType)); + targetExecution, + queryType, + Header.empty(), + args, + argTypes, + resultClass, + resultType)); } catch (Exception e) { return throwAsWorkflowFailureExceptionForQuery(e, resultClass, targetExecution); } @@ -342,6 +402,12 @@ public R query(String queryType, Class resultClass, Type resultType, Obje @Override public R update(String updateName, Class resultClass, Object... args) { + return updateWithTypeHints(updateName, resultClass, resultClass, null, args); + } + + @Override + public R updateWithTypeHints( + String updateName, Class resultClass, Type resultType, Type[] argTypes, Object... args) { checkStarted(); try { UpdateOptions options = @@ -349,9 +415,10 @@ public R update(String updateName, Class resultClass, Object... args) { .setUpdateName(updateName) .setWaitForStage(WorkflowUpdateStage.COMPLETED) .setResultClass(resultClass) + .setResultType(resultType) .setFirstExecutionRunId(firstExecutionRunId) .build(); - return startUpdate(options, args).getResultAsync().get(); + return startUpdateWithTypeHints(options, argTypes, args).getResultAsync().get(); } catch (InterruptedException e) { throw new RuntimeException(e); } catch (ExecutionException e) { @@ -379,12 +446,18 @@ public WorkflowUpdateHandle startUpdate( @Override public WorkflowUpdateHandle startUpdate(UpdateOptions options, Object... args) { + return startUpdateWithTypeHints(options, null, args); + } + + @Override + public WorkflowUpdateHandle startUpdateWithTypeHints( + UpdateOptions options, Type[] argTypes, Object... args) { checkStarted(); options.validate(); WorkflowExecution targetExecution = execution.get(); try { WorkflowClientCallsInterceptor.StartUpdateInput input = - startUpdateInput(options, args, targetExecution); + startUpdateInput(options, args, argTypes, targetExecution); return workflowClientInvoker.startUpdate(input); } catch (Exception e) { Throwable throwable = throwAsWorkflowFailureException(e, targetExecution); @@ -393,7 +466,10 @@ public WorkflowUpdateHandle startUpdate(UpdateOptions options, Object. } private WorkflowClientCallsInterceptor.StartUpdateInput startUpdateInput( - UpdateOptions options, Object[] args, WorkflowExecution targetExecution) { + UpdateOptions options, + Object[] args, + @Nullable Type[] argTypes, + WorkflowExecution targetExecution) { String updateId = Strings.isNullOrEmpty(options.getUpdateId()) ? UUID.randomUUID().toString() @@ -405,6 +481,7 @@ private WorkflowClientCallsInterceptor.StartUpdateInput startUpdateInput( Header.empty(), updateId, args, + argTypes, options.getResultClass(), options.getResultType(), options.getFirstExecutionRunId(), diff --git a/temporal-sdk/src/main/java/io/temporal/common/converter/CodecDataConverter.java b/temporal-sdk/src/main/java/io/temporal/common/converter/CodecDataConverter.java index a82172348b..135279477a 100644 --- a/temporal-sdk/src/main/java/io/temporal/common/converter/CodecDataConverter.java +++ b/temporal-sdk/src/main/java/io/temporal/common/converter/CodecDataConverter.java @@ -113,6 +113,17 @@ public CodecDataConverter( public Optional toPayload(T value) { Optional payload = ConverterUtils.withContext(dataConverter, serializationContext).toPayload(value); + return encodeSingle(payload); + } + + @Override + public Optional toPayload(T value, Type valueType) { + Optional payload = + ConverterUtils.withContext(dataConverter, serializationContext).toPayload(value, valueType); + return encodeSingle(payload); + } + + private Optional encodeSingle(Optional payload) { List encodedPayloads = ConverterUtils.withContext(chainCodec, serializationContext) .encode(Collections.singletonList(payload.get())); @@ -134,6 +145,19 @@ public T fromPayload(Payload payload, Class valueClass, Type valueType) { public Optional toPayloads(Object... values) throws DataConverterException { Optional payloads = ConverterUtils.withContext(dataConverter, serializationContext).toPayloads(values); + return encodePayloads(payloads); + } + + @Override + public Optional toPayloads(Object[] values, Type[] valueTypes) + throws DataConverterException { + Optional payloads = + ConverterUtils.withContext(dataConverter, serializationContext) + .toPayloads(values, valueTypes); + return encodePayloads(payloads); + } + + private Optional encodePayloads(Optional payloads) { if (payloads.isPresent()) { List encodedPayloads = ConverterUtils.withContext(chainCodec, serializationContext) diff --git a/temporal-sdk/src/main/java/io/temporal/common/converter/DataConverter.java b/temporal-sdk/src/main/java/io/temporal/common/converter/DataConverter.java index decf6181b0..732bc2b654 100644 --- a/temporal-sdk/src/main/java/io/temporal/common/converter/DataConverter.java +++ b/temporal-sdk/src/main/java/io/temporal/common/converter/DataConverter.java @@ -71,6 +71,22 @@ static DataConverter getDefaultInstance() { */ Optional toPayload(T value) throws DataConverterException; + /** + * Serializes a value using the supplied type hint. + * + *

The type hint may be used by converters whose serialization format depends on the declared + * type of the value. The default implementation preserves compatibility with existing data + * converters by delegating to {@link #toPayload(Object)}. + * + * @param value value to convert + * @param valueType declared type of {@code value} + * @return a {@link Payload} containing the serialized representation of {@code value} + * @throws DataConverterException if conversion fails + */ + default Optional toPayload(T value, Type valueType) throws DataConverterException { + return toPayload(value); + } + T fromPayload(Payload payload, Class valueClass, Type valueType) throws DataConverterException; @@ -84,6 +100,22 @@ T fromPayload(Payload payload, Class valueClass, Type valueType) */ Optional toPayloads(Object... values) throws DataConverterException; + /** + * Serializes a list of values using the supplied type hints. + * + *

The default implementation preserves compatibility with existing data converters by + * delegating to {@link #toPayloads(Object...)}. + * + * @param values Java values to convert + * @param valueTypes declared types of {@code values} + * @return converted values, or an empty Optional if {@code values} is empty + * @throws DataConverterException if conversion fails + */ + default Optional toPayloads(Object[] values, Type[] valueTypes) + throws DataConverterException { + return toPayloads(values); + } + /** * Implements conversion of a single {@link Payload} from the serialized {@link Payloads}. * diff --git a/temporal-sdk/src/main/java/io/temporal/common/converter/GsonJsonPayloadConverter.java b/temporal-sdk/src/main/java/io/temporal/common/converter/GsonJsonPayloadConverter.java index c2d361b7e4..f89a1ace00 100644 --- a/temporal-sdk/src/main/java/io/temporal/common/converter/GsonJsonPayloadConverter.java +++ b/temporal-sdk/src/main/java/io/temporal/common/converter/GsonJsonPayloadConverter.java @@ -51,8 +51,18 @@ public String getEncodingType() { */ @Override public Optional toData(Object value) throws DataConverterException { + return toData(value, null, false); + } + + @Override + public Optional toData(Object value, Type valueType) throws DataConverterException { + return toData(value, valueType, true); + } + + private Optional toData(Object value, Type valueType, boolean useTypeHint) + throws DataConverterException { try { - String json = gson.toJson(value); + String json = useTypeHint ? gson.toJson(value, valueType) : gson.toJson(value); return Optional.of( Payload.newBuilder() .putMetadata(EncodingKeys.METADATA_ENCODING_KEY, EncodingKeys.METADATA_ENCODING_JSON) diff --git a/temporal-sdk/src/main/java/io/temporal/common/converter/JacksonJsonPayloadConverter.java b/temporal-sdk/src/main/java/io/temporal/common/converter/JacksonJsonPayloadConverter.java index 7022d71f53..b8cd2662ee 100644 --- a/temporal-sdk/src/main/java/io/temporal/common/converter/JacksonJsonPayloadConverter.java +++ b/temporal-sdk/src/main/java/io/temporal/common/converter/JacksonJsonPayloadConverter.java @@ -105,18 +105,39 @@ public Optional toData(Object value) throws DataConverterException { } try { - byte[] serialized = mapper.writeValueAsBytes(value); - return Optional.of( - Payload.newBuilder() - .putMetadata(EncodingKeys.METADATA_ENCODING_KEY, EncodingKeys.METADATA_ENCODING_JSON) - .setData(ByteString.copyFrom(serialized)) - .build()); + return toPayload(mapper.writeValueAsBytes(value)); + } catch (JsonProcessingException e) { + throw new DataConverterException(e); + } + } + @Override + public Optional toData(Object value, Type valueType) throws DataConverterException { + // Delegate to Jackson 3 converter if globally opted in via setDefaultAsJackson3 + PayloadConverter delegate = jackson3Delegate; + if (delegate != null && useDefaultJackson3Delegate) { + return delegate.toData(value, valueType); + } + + try { + byte[] serialized = + mapper + .writerFor(mapper.getTypeFactory().constructType(valueType)) + .writeValueAsBytes(value); + return toPayload(serialized); } catch (JsonProcessingException e) { throw new DataConverterException(e); } } + private Optional toPayload(byte[] serialized) { + return Optional.of( + Payload.newBuilder() + .putMetadata(EncodingKeys.METADATA_ENCODING_KEY, EncodingKeys.METADATA_ENCODING_JSON) + .setData(ByteString.copyFrom(serialized)) + .build()); + } + @Override public T fromData(Payload content, Class valueClass, Type valueType) throws DataConverterException { diff --git a/temporal-sdk/src/main/java/io/temporal/common/converter/PayloadAndFailureDataConverter.java b/temporal-sdk/src/main/java/io/temporal/common/converter/PayloadAndFailureDataConverter.java index 935fd8462c..85bd2045c6 100644 --- a/temporal-sdk/src/main/java/io/temporal/common/converter/PayloadAndFailureDataConverter.java +++ b/temporal-sdk/src/main/java/io/temporal/common/converter/PayloadAndFailureDataConverter.java @@ -11,6 +11,7 @@ import io.temporal.payload.context.SerializationContext; import java.lang.reflect.Type; import java.util.*; +import java.util.function.Function; import javax.annotation.Nonnull; import javax.annotation.Nullable; @@ -45,6 +46,17 @@ public PayloadAndFailureDataConverter(@Nonnull List converters @Override public Optional toPayload(T value) throws DataConverterException { + return toPayload(value, converter -> converter.toData(value)); + } + + @Override + public Optional toPayload(T value, Type valueType) throws DataConverterException { + return toPayload(value, converter -> converter.toData(value, valueType)); + } + + private Optional toPayload( + T value, Function> conversion) + throws DataConverterException { // Raw values payload should be passed through without conversion if (value instanceof RawValue) { RawValue rv = (RawValue) value; @@ -52,9 +64,9 @@ public Optional toPayload(T value) throws DataConverterException { } for (PayloadConverter converter : converters) { - Optional result = - (serializationContext != null ? converter.withContext(serializationContext) : converter) - .toData(value); + PayloadConverter contextAwareConverter = + serializationContext != null ? converter.withContext(serializationContext) : converter; + Optional result = conversion.apply(contextAwareConverter); if (result.isPresent()) { return result; } @@ -92,13 +104,36 @@ public T fromPayload(Payload payload, Class valueClass, Type valueType) @Override public Optional toPayloads(Object... values) throws DataConverterException { + return toPayloads(values, null, false); + } + + @Override + public Optional toPayloads(Object[] values, Type[] valueTypes) + throws DataConverterException { + if (valueTypes == null) { + return toPayloads(values); + } + int valuesLength = values == null ? 0 : values.length; + if (valuesLength != valueTypes.length) { + throw new IllegalArgumentException( + "values don't match length of valueTypes: " + + Arrays.toString(values) + + "<>" + + Arrays.toString(valueTypes)); + } + return toPayloads(values, valueTypes, true); + } + + private Optional toPayloads(Object[] values, Type[] valueTypes, boolean useTypeHints) + throws DataConverterException { if (values == null || values.length == 0) { return Optional.empty(); } try { Payloads.Builder result = Payloads.newBuilder(); - for (Object value : values) { - result.addPayloads(toPayload(value).get()); + for (int i = 0; i < values.length; i++) { + result.addPayloads( + (useTypeHints ? toPayload(values[i], valueTypes[i]) : toPayload(values[i])).get()); } return Optional.of(result.build()); } catch (DataConverterException e) { diff --git a/temporal-sdk/src/main/java/io/temporal/common/converter/PayloadConverter.java b/temporal-sdk/src/main/java/io/temporal/common/converter/PayloadConverter.java index bf14a17c07..fbfb7665aa 100644 --- a/temporal-sdk/src/main/java/io/temporal/common/converter/PayloadConverter.java +++ b/temporal-sdk/src/main/java/io/temporal/common/converter/PayloadConverter.java @@ -41,6 +41,22 @@ public interface PayloadConverter { */ Optional toData(Object value) throws DataConverterException; + /** + * Serializes a value using the supplied type hint. + * + *

The type hint may be used by converters whose serialization format depends on the declared + * type of the value. The default implementation preserves compatibility with existing payload + * converters by delegating to {@link #toData(Object)}. + * + * @param value Java value to convert + * @param valueType declared type of {@code value} + * @return converted value + * @throws DataConverterException if conversion fails + */ + default Optional toData(Object value, Type valueType) throws DataConverterException { + return toData(value); + } + /** * Implements conversion of a single value. * diff --git a/temporal-sdk/src/main/java/io/temporal/common/interceptors/ActivityClientCallsInterceptor.java b/temporal-sdk/src/main/java/io/temporal/common/interceptors/ActivityClientCallsInterceptor.java index e2e5174c55..c3486d7f6d 100644 --- a/temporal-sdk/src/main/java/io/temporal/common/interceptors/ActivityClientCallsInterceptor.java +++ b/temporal-sdk/src/main/java/io/temporal/common/interceptors/ActivityClientCallsInterceptor.java @@ -154,13 +154,24 @@ CompletableFuture> getActivityResultAsync( final class StartActivityInput { private final String activityType; private final List args; + private final @Nullable Type[] argTypes; private final StartActivityOptions options; private final Header header; public StartActivityInput( String activityType, List args, StartActivityOptions options, Header header) { + this(activityType, args, null, options, header); + } + + public StartActivityInput( + String activityType, + List args, + @Nullable Type[] argTypes, + StartActivityOptions options, + Header header) { this.activityType = activityType; this.args = args; + this.argTypes = argTypes; this.options = options; this.header = header; } @@ -173,6 +184,10 @@ public List getArgs() { return args; } + public @Nullable Type[] getArgTypes() { + return argTypes; + } + public StartActivityOptions getOptions() { return options; } diff --git a/temporal-sdk/src/main/java/io/temporal/common/interceptors/WorkflowClientCallsInterceptor.java b/temporal-sdk/src/main/java/io/temporal/common/interceptors/WorkflowClientCallsInterceptor.java index e3593e9c2c..6b4885bd90 100644 --- a/temporal-sdk/src/main/java/io/temporal/common/interceptors/WorkflowClientCallsInterceptor.java +++ b/temporal-sdk/src/main/java/io/temporal/common/interceptors/WorkflowClientCallsInterceptor.java @@ -117,6 +117,7 @@ final class WorkflowStartInput { private final String workflowType; private final Header header; private final Object[] arguments; + private final @Nullable Type[] argumentTypes; private final WorkflowOptions options; /** @@ -133,10 +134,21 @@ public WorkflowStartInput( @Nonnull Header header, @Nonnull Object[] arguments, @Nonnull WorkflowOptions options) { + this(workflowId, workflowType, header, arguments, null, options); + } + + public WorkflowStartInput( + @Nonnull String workflowId, + @Nonnull String workflowType, + @Nonnull Header header, + @Nonnull Object[] arguments, + @Nullable Type[] argumentTypes, + @Nonnull WorkflowOptions options) { this.workflowId = workflowId; this.workflowType = workflowType; this.header = header; this.arguments = arguments; + this.argumentTypes = argumentTypes; this.options = options; } @@ -156,6 +168,10 @@ public Object[] getArguments() { return arguments; } + public @Nullable Type[] getArgumentTypes() { + return argumentTypes; + } + public WorkflowOptions getOptions() { return options; } @@ -179,16 +195,27 @@ final class WorkflowSignalInput { private final String signalName; private final Header header; private final Object[] arguments; + private final @Nullable Type[] argumentTypes; public WorkflowSignalInput( WorkflowExecution workflowExecution, String signalName, Header header, Object[] signalArguments) { + this(workflowExecution, signalName, header, signalArguments, null); + } + + public WorkflowSignalInput( + WorkflowExecution workflowExecution, + String signalName, + Header header, + Object[] signalArguments, + @Nullable Type[] argumentTypes) { this.workflowExecution = workflowExecution; this.signalName = signalName; this.header = header; this.arguments = signalArguments; + this.argumentTypes = argumentTypes; } public WorkflowExecution getWorkflowExecution() { @@ -206,6 +233,10 @@ public Header getHeader() { public Object[] getArguments() { return arguments; } + + public @Nullable Type[] getArgumentTypes() { + return argumentTypes; + } } final class WorkflowSignalOutput {} @@ -214,12 +245,22 @@ final class WorkflowSignalWithStartInput { private final WorkflowStartInput workflowStartInput; private final String signalName; private final Object[] signalArguments; + private final @Nullable Type[] signalArgumentTypes; public WorkflowSignalWithStartInput( WorkflowStartInput workflowStartInput, String signalName, Object[] signalArguments) { + this(workflowStartInput, signalName, signalArguments, null); + } + + public WorkflowSignalWithStartInput( + WorkflowStartInput workflowStartInput, + String signalName, + Object[] signalArguments, + @Nullable Type[] signalArgumentTypes) { this.workflowStartInput = workflowStartInput; this.signalName = signalName; this.signalArguments = signalArguments; + this.signalArgumentTypes = signalArgumentTypes; } public WorkflowStartInput getWorkflowStartInput() { @@ -233,6 +274,10 @@ public String getSignalName() { public Object[] getSignalArguments() { return signalArguments; } + + public @Nullable Type[] getSignalArgumentTypes() { + return signalArgumentTypes; + } } final class WorkflowSignalWithStartOutput { @@ -362,6 +407,7 @@ final class QueryInput { private final String queryType; private final Header header; private final Object[] arguments; + private final @Nullable Type[] argumentTypes; private final Class resultClass; private final Type resultType; @@ -372,10 +418,22 @@ public QueryInput( Object[] arguments, Class resultClass, Type resultType) { + this(workflowExecution, queryType, header, arguments, null, resultClass, resultType); + } + + public QueryInput( + WorkflowExecution workflowExecution, + String queryType, + Header header, + Object[] arguments, + @Nullable Type[] argumentTypes, + Class resultClass, + Type resultType) { this.workflowExecution = workflowExecution; this.queryType = queryType; this.header = header; this.arguments = arguments; + this.argumentTypes = argumentTypes; this.resultClass = resultClass; this.resultType = resultType; } @@ -396,6 +454,10 @@ public Object[] getArguments() { return arguments; } + public @Nullable Type[] getArgumentTypes() { + return argumentTypes; + } + public Class getResultClass() { return resultClass; } @@ -479,6 +541,7 @@ final class StartUpdateInput { private final String updateName; private final Header header; private final Object[] arguments; + private final @Nullable Type[] argumentTypes; private final Class resultClass; private final Type resultType; private final String updateId; @@ -496,12 +559,39 @@ public StartUpdateInput( Type resultType, String firstExecutionRunId, WaitPolicy waitPolicy) { + this( + workflowExecution, + workflowType, + updateName, + header, + updateId, + arguments, + null, + resultClass, + resultType, + firstExecutionRunId, + waitPolicy); + } + + public StartUpdateInput( + WorkflowExecution workflowExecution, + Optional workflowType, + String updateName, + Header header, + String updateId, + Object[] arguments, + @Nullable Type[] argumentTypes, + Class resultClass, + Type resultType, + String firstExecutionRunId, + WaitPolicy waitPolicy) { this.workflowExecution = workflowExecution; this.workflowType = workflowType; this.header = header; this.updateId = updateId; this.updateName = updateName; this.arguments = arguments; + this.argumentTypes = argumentTypes; this.resultClass = resultClass; this.resultType = resultType; this.firstExecutionRunId = firstExecutionRunId; @@ -532,6 +622,10 @@ public Object[] getArguments() { return arguments; } + public @Nullable Type[] getArgumentTypes() { + return argumentTypes; + } + public Class getResultClass() { return resultClass; } diff --git a/temporal-sdk/src/main/java/io/temporal/common/interceptors/WorkflowOutboundCallsInterceptor.java b/temporal-sdk/src/main/java/io/temporal/common/interceptors/WorkflowOutboundCallsInterceptor.java index 357df1c4da..0f7ff05fe9 100644 --- a/temporal-sdk/src/main/java/io/temporal/common/interceptors/WorkflowOutboundCallsInterceptor.java +++ b/temporal-sdk/src/main/java/io/temporal/common/interceptors/WorkflowOutboundCallsInterceptor.java @@ -44,6 +44,7 @@ final class ActivityInput { private final Class resultClass; private final Type resultType; private final Object[] args; + private final @Nullable Type[] argTypes; private final ActivityOptions options; private final Header header; @@ -54,10 +55,22 @@ public ActivityInput( Object[] args, ActivityOptions options, Header header) { + this(activityName, resultClass, resultType, args, null, options, header); + } + + public ActivityInput( + String activityName, + Class resultClass, + Type resultType, + Object[] args, + @Nullable Type[] argTypes, + ActivityOptions options, + Header header) { this.activityName = activityName; this.resultClass = resultClass; this.resultType = resultType; this.args = args; + this.argTypes = argTypes; this.options = options; this.header = header; } @@ -78,6 +91,10 @@ public Object[] getArgs() { return args; } + public @Nullable Type[] getArgTypes() { + return argTypes; + } + public ActivityOptions getOptions() { return options; } @@ -110,6 +127,7 @@ final class LocalActivityInput { private final Class resultClass; private final Type resultType; private final Object[] args; + private final @Nullable Type[] argTypes; private final LocalActivityOptions options; private final Header header; @@ -120,10 +138,22 @@ public LocalActivityInput( Object[] args, LocalActivityOptions options, Header header) { + this(activityName, resultClass, resultType, args, null, options, header); + } + + public LocalActivityInput( + String activityName, + Class resultClass, + Type resultType, + Object[] args, + @Nullable Type[] argTypes, + LocalActivityOptions options, + Header header) { this.activityName = activityName; this.resultClass = resultClass; this.resultType = resultType; this.args = args; + this.argTypes = argTypes; this.options = options; this.header = header; } @@ -144,6 +174,10 @@ public Object[] getArgs() { return args; } + public @Nullable Type[] getArgTypes() { + return argTypes; + } + public LocalActivityOptions getOptions() { return options; } @@ -171,6 +205,7 @@ final class ChildWorkflowInput { private final Class resultClass; private final Type resultType; private final Object[] args; + private final @Nullable Type[] argTypes; private final ChildWorkflowOptions options; private final Header header; @@ -182,11 +217,24 @@ public ChildWorkflowInput( Object[] args, ChildWorkflowOptions options, Header header) { + this(workflowId, workflowType, resultClass, resultType, args, null, options, header); + } + + public ChildWorkflowInput( + String workflowId, + String workflowType, + Class resultClass, + Type resultType, + Object[] args, + @Nullable Type[] argTypes, + ChildWorkflowOptions options, + Header header) { this.workflowId = workflowId; this.workflowType = workflowType; this.resultClass = resultClass; this.resultType = resultType; this.args = args; + this.argTypes = argTypes; this.options = options; this.header = header; } @@ -211,6 +259,10 @@ public Object[] getArgs() { return args; } + public @Nullable Type[] getArgTypes() { + return argTypes; + } + public ChildWorkflowOptions getOptions() { return options; } @@ -338,13 +390,24 @@ final class SignalExternalInput { private final String signalName; private final Header header; private final Object[] args; + private final @Nullable Type[] argTypes; public SignalExternalInput( WorkflowExecution execution, String signalName, Header header, Object[] args) { + this(execution, signalName, header, args, null); + } + + public SignalExternalInput( + WorkflowExecution execution, + String signalName, + Header header, + Object[] args, + @Nullable Type[] argTypes) { this.execution = execution; this.signalName = signalName; this.header = header; this.args = args; + this.argTypes = argTypes; } public WorkflowExecution getExecution() { @@ -362,6 +425,10 @@ public Header getHeader() { public Object[] getArgs() { return args; } + + public @Nullable Type[] getArgTypes() { + return argTypes; + } } final class SignalExternalOutput { @@ -417,6 +484,7 @@ final class ContinueAsNewInput { private final @Nullable String workflowType; private final @Nullable ContinueAsNewOptions options; private final Object[] args; + private final @Nullable Type[] argTypes; private final Header header; public ContinueAsNewInput( @@ -424,9 +492,19 @@ public ContinueAsNewInput( @Nullable ContinueAsNewOptions options, Object[] args, Header header) { + this(workflowType, options, args, null, header); + } + + public ContinueAsNewInput( + @Nullable String workflowType, + @Nullable ContinueAsNewOptions options, + Object[] args, + @Nullable Type[] argTypes, + Header header) { this.workflowType = workflowType; this.options = options; this.args = args; + this.argTypes = argTypes; this.header = header; } @@ -450,6 +528,10 @@ public Object[] getArgs() { return args; } + public @Nullable Type[] getArgTypes() { + return argTypes; + } + public Header getHeader() { return header; } @@ -551,6 +633,7 @@ final class UpdateRegistrationRequest { private final HandlerUnfinishedPolicy unfinishedPolicy; private final Class[] argTypes; private final Type[] genericArgTypes; + private final @Nullable Type resultType; private final Functions.Func1 executeCallback; private final Functions.Proc1 validateCallback; @@ -562,13 +645,34 @@ public UpdateRegistrationRequest( Type[] genericArgTypes, Functions.Proc1 validateCallback, Functions.Func1 executeCallback) { - this.updateName = updateName; - this.description = ""; - this.unfinishedPolicy = unfinishedPolicy; - this.argTypes = argTypes; - this.genericArgTypes = genericArgTypes; - this.validateCallback = validateCallback; - this.executeCallback = executeCallback; + this( + updateName, + "", + unfinishedPolicy, + argTypes, + genericArgTypes, + null, + validateCallback, + executeCallback); + } + + public UpdateRegistrationRequest( + String updateName, + String description, + HandlerUnfinishedPolicy unfinishedPolicy, + Class[] argTypes, + Type[] genericArgTypes, + Functions.Proc1 validateCallback, + Functions.Func1 executeCallback) { + this( + updateName, + description, + unfinishedPolicy, + argTypes, + genericArgTypes, + null, + validateCallback, + executeCallback); } public UpdateRegistrationRequest( @@ -577,6 +681,7 @@ public UpdateRegistrationRequest( HandlerUnfinishedPolicy unfinishedPolicy, Class[] argTypes, Type[] genericArgTypes, + @Nullable Type resultType, Functions.Proc1 validateCallback, Functions.Func1 executeCallback) { this.updateName = updateName; @@ -584,6 +689,7 @@ public UpdateRegistrationRequest( this.unfinishedPolicy = unfinishedPolicy; this.argTypes = argTypes; this.genericArgTypes = genericArgTypes; + this.resultType = resultType; this.validateCallback = validateCallback; this.executeCallback = executeCallback; } @@ -609,6 +715,10 @@ public Type[] getGenericArgTypes() { return genericArgTypes; } + public @Nullable Type getResultType() { + return resultType; + } + public Functions.Proc1 getValidateCallback() { return validateCallback; } @@ -635,6 +745,7 @@ final class RegisterQueryInput { private final String description; private final Class[] argTypes; private final Type[] genericArgTypes; + private final @Nullable Type resultType; private final Functions.Func1 callback; // Kept for backward compatibility @@ -643,11 +754,7 @@ public RegisterQueryInput( Class[] argTypes, Type[] genericArgTypes, Functions.Func1 callback) { - this.queryType = queryType; - this.description = ""; - this.argTypes = argTypes; - this.genericArgTypes = genericArgTypes; - this.callback = callback; + this(queryType, "", argTypes, genericArgTypes, null, callback); } public RegisterQueryInput( @@ -656,10 +763,21 @@ public RegisterQueryInput( Class[] argTypes, Type[] genericArgTypes, Functions.Func1 callback) { + this(queryType, description, argTypes, genericArgTypes, null, callback); + } + + public RegisterQueryInput( + String queryType, + String description, + Class[] argTypes, + Type[] genericArgTypes, + @Nullable Type resultType, + Functions.Func1 callback) { this.queryType = queryType; this.description = description; this.argTypes = argTypes; this.genericArgTypes = genericArgTypes; + this.resultType = resultType; this.callback = callback; } @@ -680,6 +798,10 @@ public Type[] getGenericArgTypes() { return genericArgTypes; } + public @Nullable Type getResultType() { + return resultType; + } + public Functions.Func1 getCallback() { return callback; } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/activity/ActivityTaskExecutors.java b/temporal-sdk/src/main/java/io/temporal/internal/activity/ActivityTaskExecutors.java index 7789bd2716..d69ced9f7e 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/activity/ActivityTaskExecutors.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/activity/ActivityTaskExecutors.java @@ -21,6 +21,7 @@ import io.temporal.payload.context.ActivitySerializationContext; import io.temporal.serviceclient.CheckedExceptionWrapper; import java.lang.reflect.Method; +import java.lang.reflect.Type; import java.util.HashMap; import java.util.List; import java.util.Map; @@ -158,11 +159,22 @@ ActivityTaskHandler.Result constructResultValue( ActivityInfoInternal info, @Nullable ActivityOutput result, DataConverter dataConverterWithActivityContext) { + return constructResultValue(info, result, null, dataConverterWithActivityContext); + } + + ActivityTaskHandler.Result constructResultValue( + ActivityInfoInternal info, + @Nullable ActivityOutput result, + @Nullable Type resultType, + DataConverter dataConverterWithActivityContext) { RespondActivityTaskCompletedRequest.Builder request = RespondActivityTaskCompletedRequest.newBuilder(); if (result != null) { Optional serialized = - dataConverterWithActivityContext.toPayloads(result.getResult()); + resultType == null + ? dataConverterWithActivityContext.toPayloads(result.getResult()) + : dataConverterWithActivityContext.toPayloads( + new Object[] {result.getResult()}, new Type[] {resultType}); serialized.ifPresent(request::setResult); } return new ActivityTaskHandler.Result( @@ -225,6 +237,7 @@ protected ActivityTaskHandler.Result constructSuccessfulResultValue( info, // if the expected result of the method is null, we don't publish result at all method.getReturnType() != Void.TYPE ? result : null, + method.getGenericReturnType(), dataConverterWithActivityContext); } } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/client/RootActivityClientInvoker.java b/temporal-sdk/src/main/java/io/temporal/internal/client/RootActivityClientInvoker.java index fc85f0039f..95e4d8127a 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/client/RootActivityClientInvoker.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/client/RootActivityClientInvoker.java @@ -83,7 +83,8 @@ public StartActivityOutput startActivity(StartActivityInput input) { .setIdReusePolicy(options.getIdReusePolicy()) .setIdConflictPolicy(options.getIdConflictPolicy()); - Optional activityInput = dc.toPayloads(input.getArgs().toArray()); + Optional activityInput = + dc.toPayloads(input.getArgs().toArray(), input.getArgTypes()); activityInput.ifPresent(request::setInput); if (options.getScheduleToCloseTimeout() != null) { diff --git a/temporal-sdk/src/main/java/io/temporal/internal/client/RootWorkflowClientInvoker.java b/temporal-sdk/src/main/java/io/temporal/internal/client/RootWorkflowClientInvoker.java index d3b3321ebe..58e3659c2f 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/client/RootWorkflowClientInvoker.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/client/RootWorkflowClientInvoker.java @@ -190,7 +190,8 @@ public WorkflowSignalOutput signal(WorkflowSignalInput input) { DataConverter dataConverterWitSignalContext = workflowConverter(input.getWorkflowExecution()); - Optional inputArgs = dataConverterWitSignalContext.toPayloads(input.getArguments()); + Optional inputArgs = + dataConverterWitSignalContext.toPayloads(input.getArguments(), input.getArgumentTypes()); inputArgs.ifPresent(request::setInput); storeHeader( request.getHeaderBuilder(), @@ -217,7 +218,8 @@ public WorkflowSignalWithStartOutput signalWithStart(WorkflowSignalWithStartInpu toStartRequest(dataConverterWithWorkflowContext, workflowStartInput); Optional signalInput = - dataConverterWithWorkflowContext.toPayloads(input.getSignalArguments()); + dataConverterWithWorkflowContext.toPayloads( + input.getSignalArguments(), input.getSignalArgumentTypes()); SignalWithStartWorkflowExecutionRequest.Builder requestBuilder = requestsHelper.newSignalWithStartWorkflowExecutionRequest( startRequest, input.getSignalName(), signalInput.orElse(null)); @@ -365,7 +367,8 @@ public WorkflowUpdateWithStartOutput updateWithStart( private StartWorkflowExecutionRequest.Builder toStartRequest( DataConverter dataConverterWithWorkflowContext, WorkflowStartInput workflowStartInput) { Optional workflowInput = - dataConverterWithWorkflowContext.toPayloads(workflowStartInput.getArguments()); + dataConverterWithWorkflowContext.toPayloads( + workflowStartInput.getArguments(), workflowStartInput.getArgumentTypes()); @Nullable Memo memo = @@ -456,7 +459,7 @@ public QueryOutput query(QueryInput input) { workflowConverter(input.getWorkflowExecution()); Optional inputArgs = - dataConverterWithWorkflowContext.toPayloads(input.getArguments()); + dataConverterWithWorkflowContext.toPayloads(input.getArguments(), input.getArgumentTypes()); inputArgs.ifPresent(query::setQueryArgs); storeHeader( query.getHeaderBuilder(), @@ -552,7 +555,7 @@ private boolean updateNotYetDurable( private UpdateWorkflowExecutionRequest toUpdateWorkflowExecutionRequest( StartUpdateInput input, DataConverter dataConverterWithWorkflowContext) { Optional inputArgs = - dataConverterWithWorkflowContext.toPayloads(input.getArguments()); + dataConverterWithWorkflowContext.toPayloads(input.getArguments(), input.getArgumentTypes()); Input.Builder updateInput = Input.newBuilder() .setHeader(HeaderUtils.toHeaderGrpc(input.getHeader(), null)) diff --git a/temporal-sdk/src/main/java/io/temporal/internal/common/NexusWorkflowStarter.java b/temporal-sdk/src/main/java/io/temporal/internal/common/NexusWorkflowStarter.java index be0f1f0080..c4b5977c48 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/common/NexusWorkflowStarter.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/common/NexusWorkflowStarter.java @@ -3,6 +3,7 @@ import io.temporal.api.common.v1.WorkflowExecution; import io.temporal.client.WorkflowStub; import io.temporal.internal.client.NexusStartWorkflowResponse; +import java.lang.reflect.Type; public class NexusWorkflowStarter { private final WorkflowStub workflowStub; @@ -17,4 +18,9 @@ public NexusStartWorkflowResponse start(Object... args) { WorkflowExecution workflowExecution = workflowStub.start(args); return new NexusStartWorkflowResponse(workflowExecution, operationToken); } + + public NexusStartWorkflowResponse startWithTypeHints(Type[] argTypes, Object... args) { + WorkflowExecution workflowExecution = workflowStub.startWithTypeHints(argTypes, args); + return new NexusStartWorkflowResponse(workflowExecution, operationToken); + } } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/payload/storage/ExternalStorageDataConverter.java b/temporal-sdk/src/main/java/io/temporal/internal/payload/storage/ExternalStorageDataConverter.java index 1f8c1198c7..aedc7c42df 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/payload/storage/ExternalStorageDataConverter.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/payload/storage/ExternalStorageDataConverter.java @@ -46,7 +46,15 @@ public ExternalStorageDataConverter withStorageTarget( @Override public Optional toPayload(T value) throws DataConverterException { - Optional converted = delegate.toPayload(value); + return storePayload(delegate.toPayload(value)); + } + + @Override + public Optional toPayload(T value, Type valueType) throws DataConverterException { + return storePayload(delegate.toPayload(value, valueType)); + } + + private Optional storePayload(Optional converted) { if (!converted.isPresent()) { return converted; } @@ -56,7 +64,16 @@ public Optional toPayload(T value) throws DataConverterException { @Override public Optional toPayloads(Object... values) throws DataConverterException { - Optional converted = delegate.toPayloads(values); + return storePayloads(delegate.toPayloads(values)); + } + + @Override + public Optional toPayloads(Object[] values, Type[] valueTypes) + throws DataConverterException { + return storePayloads(delegate.toPayloads(values, valueTypes)); + } + + private Optional storePayloads(Optional converted) { if (!converted.isPresent()) { return converted; } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityInvocationHandler.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityInvocationHandler.java index e46408ca06..8bc5b8200b 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityInvocationHandler.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityInvocationHandler.java @@ -4,7 +4,6 @@ import io.temporal.activity.ActivityOptions; import io.temporal.common.MethodRetry; import io.temporal.common.interceptors.WorkflowOutboundCallsInterceptor; -import io.temporal.workflow.ActivityStub; import io.temporal.workflow.Functions; import java.lang.reflect.InvocationHandler; import java.lang.reflect.Method; @@ -58,9 +57,15 @@ protected Function getActivityFunc( + activityName + " activity. Please set at least one of the above through the ActivityStub or WorkflowImplementationOptions."); } - ActivityStub stub = ActivityStubImpl.newInstance(merged, activityExecutor, assertReadOnly); + ActivityStubBase stub = ActivityStubImpl.newInstance(merged, activityExecutor, assertReadOnly); function = - (a) -> stub.execute(activityName, method.getReturnType(), method.getGenericReturnType(), a); + (a) -> + stub.execute( + activityName, + method.getReturnType(), + method.getGenericReturnType(), + method.getGenericParameterTypes(), + a); return function; } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityStubBase.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityStubBase.java index 95698f6ef8..d5db0a41d3 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityStubBase.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityStubBase.java @@ -6,7 +6,7 @@ import io.temporal.workflow.Promise; import java.lang.reflect.Type; -/** Supports calling activity by name and arguments without its strongly typed interface. */ +/** Supports calling an activity by name and arguments without its strongly typed interface. */ abstract class ActivityStubBase implements ActivityStub { @Override @@ -16,7 +16,12 @@ public T execute(String activityName, Class resultClass, Object... args) @Override public T execute(String activityName, Class resultClass, Type resultType, Object... args) { - Promise result = executeAsync(activityName, resultClass, resultType, args); + return execute(activityName, resultClass, resultType, null, args); + } + + T execute( + String activityName, Class resultClass, Type resultType, Type[] argTypes, Object... args) { + Promise result = executeAsync(activityName, resultClass, resultType, argTypes, args); if (AsyncInternal.isAsync()) { AsyncInternal.setAsyncResult(result); return Defaults.defaultValue(resultClass); @@ -40,4 +45,7 @@ public Promise executeAsync(String activityName, Class resultClass, Ob @Override public abstract Promise executeAsync( String activityName, Class resultClass, Type resultType, Object... args); + + abstract Promise executeAsync( + String activityName, Class resultClass, Type resultType, Type[] argTypes, Object... args); } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityStubImpl.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityStubImpl.java index 8ed9e62be3..5a2eb0fb58 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityStubImpl.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/ActivityStubImpl.java @@ -3,7 +3,6 @@ import io.temporal.activity.ActivityOptions; import io.temporal.common.interceptors.Header; import io.temporal.common.interceptors.WorkflowOutboundCallsInterceptor; -import io.temporal.workflow.ActivityStub; import io.temporal.workflow.Functions; import io.temporal.workflow.Promise; import java.lang.reflect.Type; @@ -13,7 +12,7 @@ final class ActivityStubImpl extends ActivityStubBase { private final WorkflowOutboundCallsInterceptor activityExecutor; private final Functions.Proc assertReadOnly; - static ActivityStub newInstance( + static ActivityStubBase newInstance( ActivityOptions options, WorkflowOutboundCallsInterceptor activityExecutor, Functions.Proc assertReadOnly) { @@ -34,11 +33,17 @@ static ActivityStub newInstance( @Override public Promise executeAsync( String activityName, Class resultClass, Type resultType, Object... args) { + return executeAsync(activityName, resultClass, resultType, null, args); + } + + @Override + Promise executeAsync( + String activityName, Class resultClass, Type resultType, Type[] argTypes, Object... args) { this.assertReadOnly.apply(); return activityExecutor .executeActivity( new WorkflowOutboundCallsInterceptor.ActivityInput<>( - activityName, resultClass, resultType, args, options, Header.empty())) + activityName, resultClass, resultType, args, argTypes, options, Header.empty())) .getResult(); } } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/ChildWorkflowInvocationHandler.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/ChildWorkflowInvocationHandler.java index 7268859259..78ba9d5145 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/ChildWorkflowInvocationHandler.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/ChildWorkflowInvocationHandler.java @@ -9,7 +9,6 @@ import io.temporal.common.metadata.POJOWorkflowMethodMetadata; import io.temporal.common.metadata.WorkflowMethodType; import io.temporal.workflow.ChildWorkflowOptions; -import io.temporal.workflow.ChildWorkflowStub; import io.temporal.workflow.Functions; import java.lang.reflect.InvocationHandler; import java.lang.reflect.Method; @@ -18,7 +17,7 @@ /** Dynamic implementation of a strongly typed child workflow interface. */ class ChildWorkflowInvocationHandler implements InvocationHandler { - private final ChildWorkflowStub stub; + private final ChildWorkflowStubImpl stub; private final POJOWorkflowInterfaceMetadata workflowMetadata; ChildWorkflowInvocationHandler( @@ -69,11 +68,15 @@ public Object invoke(Object proxy, Method method, Object[] args) { if (type == WorkflowMethodType.WORKFLOW) { return getValueOrDefault( - stub.execute(method.getReturnType(), method.getGenericReturnType(), args), + stub.execute( + method.getReturnType(), + method.getGenericReturnType(), + method.getGenericParameterTypes(), + args), method.getReturnType()); } if (type == WorkflowMethodType.SIGNAL) { - stub.signal(methodMetadata.getName(), args); + stub.signal(methodMetadata.getName(), method.getGenericParameterTypes(), args); return null; } if (type == WorkflowMethodType.QUERY) { diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/ChildWorkflowStubImpl.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/ChildWorkflowStubImpl.java index 8d38d7cab1..f385fde4da 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/ChildWorkflowStubImpl.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/ChildWorkflowStubImpl.java @@ -60,8 +60,12 @@ public R execute(Class resultClass, Object... args) { @Override public R execute(Class resultClass, Type resultType, Object... args) { + return execute(resultClass, resultType, null, args); + } + + R execute(Class resultClass, Type resultType, Type[] argTypes, Object... args) { assertReadOnly.apply("schedule child workflow"); - Promise result = executeAsync(resultClass, resultType, args); + Promise result = executeAsync(resultClass, resultType, argTypes, args); if (AsyncInternal.isAsync()) { AsyncInternal.setAsyncResult(result); return Defaults.defaultValue(resultClass); @@ -83,6 +87,11 @@ public Promise executeAsync(Class resultClass, Object... args) { @Override public Promise executeAsync(Class resultClass, Type resultType, Object... args) { + return executeAsync(resultClass, resultType, null, args); + } + + Promise executeAsync( + Class resultClass, Type resultType, Type[] argTypes, Object... args) { assertReadOnly.apply("schedule child workflow"); ChildWorkflowOutput result = outboundCallsInterceptor.executeChildWorkflow( @@ -92,6 +101,7 @@ public Promise executeAsync(Class resultClass, Type resultType, Object resultClass, resultType, args, + argTypes, options, Header.empty())); execution.completeFrom(result.getWorkflowExecution()); @@ -100,12 +110,16 @@ public Promise executeAsync(Class resultClass, Type resultType, Object @Override public void signal(String signalName, Object... args) { + signal(signalName, null, args); + } + + void signal(String signalName, Type[] argTypes, Object... args) { assertReadOnly.apply("signal workflow"); Promise signaled = outboundCallsInterceptor .signalExternalWorkflow( new WorkflowOutboundCallsInterceptor.SignalExternalInput( - execution.get(), signalName, Header.empty(), args)) + execution.get(), signalName, Header.empty(), args, argTypes)) .getResult(); if (AsyncInternal.isAsync()) { AsyncInternal.setAsyncResult(signaled); diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/ContinueAsNewWorkflowInvocationHandler.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/ContinueAsNewWorkflowInvocationHandler.java index 6bf8947e95..c60663617c 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/ContinueAsNewWorkflowInvocationHandler.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/ContinueAsNewWorkflowInvocationHandler.java @@ -30,7 +30,8 @@ class ContinueAsNewWorkflowInvocationHandler implements InvocationHandler { @Override public Object invoke(Object proxy, Method method, Object[] args) { String workflowType = workflowMetadata.getMethodMetadata(method).getName(); - WorkflowInternal.continueAsNew(workflowType, options, args, outboundCallsInterceptor); + WorkflowInternal.continueAsNew( + workflowType, options, args, method.getGenericParameterTypes(), outboundCallsInterceptor); return getValueOrDefault(null, method.getReturnType()); } } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/ExternalWorkflowInvocationHandler.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/ExternalWorkflowInvocationHandler.java index 0f7def87ea..e977598fee 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/ExternalWorkflowInvocationHandler.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/ExternalWorkflowInvocationHandler.java @@ -4,7 +4,6 @@ import io.temporal.common.interceptors.WorkflowOutboundCallsInterceptor; import io.temporal.common.metadata.POJOWorkflowInterfaceMetadata; import io.temporal.common.metadata.POJOWorkflowMethodMetadata; -import io.temporal.workflow.ExternalWorkflowStub; import io.temporal.workflow.Functions; import java.lang.reflect.InvocationHandler; import java.lang.reflect.Method; @@ -12,7 +11,7 @@ /** Dynamic implementation of a strongly typed child workflow interface. */ class ExternalWorkflowInvocationHandler implements InvocationHandler { - private final ExternalWorkflowStub stub; + private final ExternalWorkflowStubImpl stub; private final POJOWorkflowInterfaceMetadata workflowMetadata; public ExternalWorkflowInvocationHandler( @@ -50,7 +49,7 @@ public Object invoke(Object proxy, Method method, Object[] args) { "Cannot start a workflow with an external workflow stub " + "created through Workflow.newExternalWorkflowStub"); case SIGNAL: - stub.signal(methodMetadata.getName(), args); + stub.signal(methodMetadata.getName(), method.getGenericParameterTypes(), args); break; case UPDATE: throw new UnsupportedOperationException( diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/ExternalWorkflowStubImpl.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/ExternalWorkflowStubImpl.java index 5234f74ceb..397abbe2b7 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/ExternalWorkflowStubImpl.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/ExternalWorkflowStubImpl.java @@ -4,6 +4,7 @@ import io.temporal.common.interceptors.Header; import io.temporal.common.interceptors.WorkflowOutboundCallsInterceptor; import io.temporal.workflow.*; +import java.lang.reflect.Type; import java.util.Objects; import javax.annotation.Nullable; @@ -30,12 +31,16 @@ public WorkflowExecution getExecution() { @Override public void signal(String signalName, Object... args) { + signal(signalName, null, args); + } + + void signal(String signalName, Type[] argTypes, Object... args) { assertReadOnly.apply("signal external workflow"); Promise signaled = outboundCallsInterceptor .signalExternalWorkflow( new WorkflowOutboundCallsInterceptor.SignalExternalInput( - execution, signalName, Header.empty(), args)) + execution, signalName, Header.empty(), args, argTypes)) .getResult(); if (AsyncInternal.isAsync()) { AsyncInternal.setAsyncResult(signaled); diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/LocalActivityInvocationHandler.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/LocalActivityInvocationHandler.java index 5b173d33f3..652e2892ac 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/LocalActivityInvocationHandler.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/LocalActivityInvocationHandler.java @@ -4,7 +4,6 @@ import io.temporal.activity.LocalActivityOptions; import io.temporal.common.MethodRetry; import io.temporal.common.interceptors.WorkflowOutboundCallsInterceptor; -import io.temporal.workflow.ActivityStub; import io.temporal.workflow.Functions; import java.lang.reflect.InvocationHandler; import java.lang.reflect.Method; @@ -53,10 +52,16 @@ public Function getActivityFunc( .mergeActivityOptions(activityMethodOptions.get(activityName)) .setMethodRetry(methodRetry) .build(); - ActivityStub stub = + ActivityStubBase stub = LocalActivityStubImpl.newInstance(mergedOptions, activityExecutor, assertReadOnly); function = - (a) -> stub.execute(activityName, method.getReturnType(), method.getGenericReturnType(), a); + (a) -> + stub.execute( + activityName, + method.getReturnType(), + method.getGenericReturnType(), + method.getGenericParameterTypes(), + a); return function; } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/LocalActivityStubImpl.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/LocalActivityStubImpl.java index 6744c26cde..84277cc1bd 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/LocalActivityStubImpl.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/LocalActivityStubImpl.java @@ -3,7 +3,6 @@ import io.temporal.activity.LocalActivityOptions; import io.temporal.common.interceptors.Header; import io.temporal.common.interceptors.WorkflowOutboundCallsInterceptor; -import io.temporal.workflow.ActivityStub; import io.temporal.workflow.Functions; import io.temporal.workflow.Promise; import java.lang.reflect.Type; @@ -13,7 +12,7 @@ class LocalActivityStubImpl extends ActivityStubBase { private final WorkflowOutboundCallsInterceptor activityExecutor; private final Functions.Proc assertReadOnly; - static ActivityStub newInstance( + static ActivityStubBase newInstance( LocalActivityOptions options, WorkflowOutboundCallsInterceptor activityExecutor, Functions.Proc assertReadOnly) { @@ -34,11 +33,17 @@ private LocalActivityStubImpl( @Override public Promise executeAsync( String activityName, Class resultClass, Type resultType, Object... args) { + return executeAsync(activityName, resultClass, resultType, null, args); + } + + @Override + Promise executeAsync( + String activityName, Class resultClass, Type resultType, Type[] argTypes, Object... args) { this.assertReadOnly.apply(); return activityExecutor .executeLocalActivity( new WorkflowOutboundCallsInterceptor.LocalActivityInput<>( - activityName, resultClass, resultType, args, options, Header.empty())) + activityName, resultClass, resultType, args, argTypes, options, Header.empty())) .getResult(); } } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/POJOWorkflowImplementationFactory.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/POJOWorkflowImplementationFactory.java index fb20638296..f879853fd5 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/POJOWorkflowImplementationFactory.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/POJOWorkflowImplementationFactory.java @@ -33,6 +33,7 @@ import java.lang.reflect.Constructor; import java.lang.reflect.InvocationTargetException; import java.lang.reflect.Method; +import java.lang.reflect.Type; import java.util.Collections; import java.util.HashMap; import java.util.List; @@ -362,7 +363,8 @@ public Optional execute(Header header, Optional input) if (workflowMethod.getReturnType() == Void.TYPE) { return Optional.empty(); } - return dataConverterWithWorkflowContext.toPayloads(result.getResult()); + return dataConverterWithWorkflowContext.toPayloads( + new Object[] {result.getResult()}, new Type[] {workflowMethod.getGenericReturnType()}); } @Nullable diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/QueryDispatcher.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/QueryDispatcher.java index b92ac3b282..cb738ad3a4 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/QueryDispatcher.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/QueryDispatcher.java @@ -104,7 +104,10 @@ public Optional handleQuery( inboundCallsInterceptor .handleQuery(new WorkflowInboundCallsInterceptor.QueryInput(queryName, header, args)) .getResult(); - return dataConverterWithWorkflowContext.toPayloads(result); + return handler == null || handler.getResultType() == null + ? dataConverterWithWorkflowContext.toPayloads(result) + : dataConverterWithWorkflowContext.toPayloads( + new Object[] {result}, new java.lang.reflect.Type[] {handler.getResultType()}); } finally { replayContext.setReadOnly(false); queryHandlerWorkflowContext.set(null); diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/SyncWorkflow.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/SyncWorkflow.java index 9351f0f34e..38a3c01f56 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/SyncWorkflow.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/SyncWorkflow.java @@ -5,6 +5,7 @@ import io.temporal.api.history.v1.HistoryEvent; import io.temporal.api.history.v1.WorkflowExecutionStartedEventAttributes; import io.temporal.api.query.v1.WorkflowQuery; +import io.temporal.api.sdk.v1.WorkflowMetadata; import io.temporal.client.WorkflowClient; import io.temporal.common.context.ContextPropagator; import io.temporal.common.converter.DataConverter; @@ -229,7 +230,7 @@ public Optional query(WorkflowQuery query) { // metadata should be readable independent of user DataConverter settings Payload payload = DefaultDataConverter.STANDARD_INSTANCE - .toPayload(workflowContext.getWorkflowMetadata()) + .toPayload(workflowContext.getWorkflowMetadata(), WorkflowMetadata.class) .orElseThrow(() -> new IllegalStateException("Failed to serialize metadata")); return dataConverterWithWorkflowContext.toPayloads(new RawValue(payload)); } diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/SyncWorkflowContext.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/SyncWorkflowContext.java index 065ce71428..114cf7f85c 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/SyncWorkflowContext.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/SyncWorkflowContext.java @@ -279,7 +279,8 @@ public ActivityOutput executeActivity(ActivityInput input) { false); DataConverter dataConverterWithActivityContext = dataConverter.withContext(serializationContext); - Optional args = dataConverterWithActivityContext.toPayloads(input.getArgs()); + Optional args = + dataConverterWithActivityContext.toPayloads(input.getArgs(), input.getArgTypes()); ActivityOutput> output = executeActivityOnce(input.getActivityName(), input.getOptions(), input.getHeader(), args); @@ -440,7 +441,8 @@ public LocalActivityOutput executeLocalActivity(LocalActivityInput inp true); DataConverter dataConverterWithActivityContext = dataConverter.withContext(serializationContext); - Optional payloads = dataConverterWithActivityContext.toPayloads(input.getArgs()); + Optional payloads = + dataConverterWithActivityContext.toPayloads(input.getArgs(), input.getArgTypes()); long originalScheduledTime = System.currentTimeMillis(); CompletablePromise> serializedResult = @@ -700,7 +702,8 @@ public ChildWorkflowOutput executeChildWorkflow(ChildWorkflowInput inp DataConverter dataConverterWithChildWorkflowContext = dataConverter.withContext( new WorkflowSerializationContext(replayContext.getNamespace(), input.getWorkflowId())); - Optional payloads = dataConverterWithChildWorkflowContext.toPayloads(input.getArgs()); + Optional payloads = + dataConverterWithChildWorkflowContext.toPayloads(input.getArgs(), input.getArgTypes()); @Nullable Memo memo = @@ -1060,7 +1063,8 @@ public R sideEffect( try { readOnly = true; R r = func.apply(); - return dataConverterWithCurrentWorkflowContext.toPayloads(r); + return dataConverterWithCurrentWorkflowContext.toPayloads( + new Object[] {r}, new Type[] {resultType}); } finally { readOnly = false; } @@ -1130,7 +1134,8 @@ private R mutableSideEffectImpl( func.apply(), "mutableSideEffect function " + "returned null"); if (!stored.isPresent() || updated.test(stored.get(), funcResult)) { unserializedResult.set(funcResult); - return dataConverterWithCurrentWorkflowContext.toPayloads(funcResult); + return dataConverterWithCurrentWorkflowContext.toPayloads( + new Object[] {funcResult}, new Type[] {resultType}); } return Optional.empty(); // returned only when value doesn't need to be updated } finally { @@ -1304,7 +1309,8 @@ public SignalExternalOutput signalExternalWorkflow(SignalExternalInput input) { attributes.setSignalName(input.getSignalName()); attributes.setExecution(childExecution); attributes.setHeader(HeaderUtils.toHeaderGrpc(input.getHeader(), null)); - Optional payloads = dataConverterWithChildWorkflowContext.toPayloads(input.getArgs()); + Optional payloads = + dataConverterWithChildWorkflowContext.toPayloads(input.getArgs(), input.getArgTypes()); payloads.ifPresent(attributes::setInput); CompletablePromise result = Workflow.newPromise(); Functions.Proc1 cancellationCallback = @@ -1475,7 +1481,7 @@ public void continueAsNew(ContinueAsNewInput input) { attributes.setHeader(grpcHeader); Optional payloads = - dataConverterWithCurrentWorkflowContext.toPayloads(input.getArgs()); + dataConverterWithCurrentWorkflowContext.toPayloads(input.getArgs(), input.getArgTypes()); payloads.ifPresent(attributes::setInput); replayContext.continueAsNewOnCompletion(attributes.build()); diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/UpdateDispatcher.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/UpdateDispatcher.java index 8ce53e87a5..1e39c7760b 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/UpdateDispatcher.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/UpdateDispatcher.java @@ -99,7 +99,10 @@ public Optional handleExecuteUpdate( .executeUpdate( new WorkflowInboundCallsInterceptor.UpdateInput(updateName, header, args)) .getResult(); - return dataConverterWithWorkflowContext.toPayloads(result); + return handler == null || handler.getResultType() == null + ? dataConverterWithWorkflowContext.toPayloads(result) + : dataConverterWithWorkflowContext.toPayloads( + new Object[] {result}, new java.lang.reflect.Type[] {handler.getResultType()}); } catch (DestroyWorkflowThreadError e) { threadDestroyed = true; throw e; diff --git a/temporal-sdk/src/main/java/io/temporal/internal/sync/WorkflowInternal.java b/temporal-sdk/src/main/java/io/temporal/internal/sync/WorkflowInternal.java index 84b1e91fd3..825b82720d 100644 --- a/temporal-sdk/src/main/java/io/temporal/internal/sync/WorkflowInternal.java +++ b/temporal-sdk/src/main/java/io/temporal/internal/sync/WorkflowInternal.java @@ -159,6 +159,7 @@ public static void registerListener(Object implementation) { methodMetadata.getDescription(), method.getParameterTypes(), method.getGenericParameterTypes(), + method.getGenericReturnType(), (args) -> { try { return method.invoke(implementation, args); @@ -237,6 +238,7 @@ public static void registerListener(Object implementation) { updateMethod.unfinishedPolicy(), method.getParameterTypes(), method.getGenericParameterTypes(), + method.getGenericReturnType(), (args) -> { try { if (validatorMethod != null) { @@ -670,11 +672,20 @@ public static void continueAsNew( @Nullable ContinueAsNewOptions options, Object[] args, WorkflowOutboundCallsInterceptor outboundCallsInterceptor) { + continueAsNew(workflowType, options, args, null, outboundCallsInterceptor); + } + + static void continueAsNew( + @Nullable String workflowType, + @Nullable ContinueAsNewOptions options, + Object[] args, + @Nullable Type[] argTypes, + WorkflowOutboundCallsInterceptor outboundCallsInterceptor) { assertNotReadOnly("continue as new"); assertNotInUpdateHandler("ContinueAsNew is not supported in an update handler"); outboundCallsInterceptor.continueAsNew( new WorkflowOutboundCallsInterceptor.ContinueAsNewInput( - workflowType, options, args, Header.empty())); + workflowType, options, args, argTypes, Header.empty())); } public static Promise cancelWorkflow(WorkflowExecution execution) { diff --git a/temporal-sdk/src/main/java/io/temporal/nexus/TemporalNexusClientImpl.java b/temporal-sdk/src/main/java/io/temporal/nexus/TemporalNexusClientImpl.java index d962d38770..26c4a195b9 100644 --- a/temporal-sdk/src/main/java/io/temporal/nexus/TemporalNexusClientImpl.java +++ b/temporal-sdk/src/main/java/io/temporal/nexus/TemporalNexusClientImpl.java @@ -45,6 +45,7 @@ import java.util.Map; import java.util.Objects; import java.util.concurrent.atomic.AtomicBoolean; +import javax.annotation.Nullable; /** Package-private implementation of {@link TemporalNexusClient}. */ @Experimental @@ -902,7 +903,8 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Collections.emptyList(), options); + return startActivityImpl( + activityType, Collections.emptyList(), method.getGenericParameterTypes(), options); } @Override @@ -913,7 +915,8 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Collections.singletonList(arg1), options); + return startActivityImpl( + activityType, Collections.singletonList(arg1), method.getGenericParameterTypes(), options); } @Override @@ -925,7 +928,8 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Arrays.asList(arg1, arg2), options); + return startActivityImpl( + activityType, Arrays.asList(arg1, arg2), method.getGenericParameterTypes(), options); } @Override @@ -938,7 +942,8 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Arrays.asList(arg1, arg2, arg3), options); + return startActivityImpl( + activityType, Arrays.asList(arg1, arg2, arg3), method.getGenericParameterTypes(), options); } @Override @@ -952,7 +957,11 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Arrays.asList(arg1, arg2, arg3, arg4), options); + return startActivityImpl( + activityType, + Arrays.asList(arg1, arg2, arg3, arg4), + method.getGenericParameterTypes(), + options); } @Override @@ -967,7 +976,11 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Arrays.asList(arg1, arg2, arg3, arg4, arg5), options); + return startActivityImpl( + activityType, + Arrays.asList(arg1, arg2, arg3, arg4, arg5), + method.getGenericParameterTypes(), + options); } @Override @@ -984,7 +997,10 @@ public TemporalOperationResult startActivity( Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); return startActivityImpl( - activityType, Arrays.asList(arg1, arg2, arg3, arg4, arg5, arg6), options); + activityType, + Arrays.asList(arg1, arg2, arg3, arg4, arg5, arg6), + method.getGenericParameterTypes(), + options); } // ---------- Activity overloads (Proc void) ---------- @@ -994,7 +1010,8 @@ public TemporalOperationResult startActivity( Class activityInterface, Functions.Proc1 activityMethod, StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Collections.emptyList(), options); + return startActivityImpl( + activityType, Collections.emptyList(), method.getGenericParameterTypes(), options); } @Override @@ -1005,7 +1022,8 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Collections.singletonList(arg1), options); + return startActivityImpl( + activityType, Collections.singletonList(arg1), method.getGenericParameterTypes(), options); } @Override @@ -1017,7 +1035,8 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Arrays.asList(arg1, arg2), options); + return startActivityImpl( + activityType, Arrays.asList(arg1, arg2), method.getGenericParameterTypes(), options); } @Override @@ -1030,7 +1049,8 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Arrays.asList(arg1, arg2, arg3), options); + return startActivityImpl( + activityType, Arrays.asList(arg1, arg2, arg3), method.getGenericParameterTypes(), options); } @Override @@ -1044,7 +1064,11 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Arrays.asList(arg1, arg2, arg3, arg4), options); + return startActivityImpl( + activityType, + Arrays.asList(arg1, arg2, arg3, arg4), + method.getGenericParameterTypes(), + options); } @Override @@ -1059,7 +1083,11 @@ public TemporalOperationResult startActivity( StartActivityOptions options) { Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); - return startActivityImpl(activityType, Arrays.asList(arg1, arg2, arg3, arg4, arg5), options); + return startActivityImpl( + activityType, + Arrays.asList(arg1, arg2, arg3, arg4, arg5), + method.getGenericParameterTypes(), + options); } @Override @@ -1076,7 +1104,10 @@ public TemporalOperationResult startActivity( Method method = MethodExtractor.extract(activityInterface, activityMethod); String activityType = MethodExtractor.activityTypeName(activityInterface, method); return startActivityImpl( - activityType, Arrays.asList(arg1, arg2, arg3, arg4, arg5, arg6), options); + activityType, + Arrays.asList(arg1, arg2, arg3, arg4, arg5, arg6), + method.getGenericParameterTypes(), + options); } // ---------- Activity untyped ---------- @@ -1085,11 +1116,14 @@ public TemporalOperationResult startActivity( public TemporalOperationResult startActivity( String activityType, Class resultClass, StartActivityOptions options, Object... args) { List argList = args == null ? Collections.emptyList() : Arrays.asList(args); - return startActivityImpl(activityType, argList, options); + return startActivityImpl(activityType, argList, null, options); } private TemporalOperationResult startActivityImpl( - String activityType, List args, StartActivityOptions options) { + String activityType, + List args, + @Nullable Type[] argTypes, + StartActivityOptions options) { markAsyncOperationStarted(); InternalNexusOperationContext nexusContext = CurrentNexusOperationContext.get(); try { @@ -1112,6 +1146,7 @@ private TemporalOperationResult startActivityImpl( new ActivityClientCallsInterceptor.StartActivityInput( request.getActivityType(), request.getArgs(), + argTypes, request.getOptions(), request.getHeader()); // Build an internal ActivityClient aligned with the surrounding WorkflowClient. diff --git a/temporal-sdk/src/main/java17/io/temporal/common/converter/Jackson3JsonPayloadConverter.java b/temporal-sdk/src/main/java17/io/temporal/common/converter/Jackson3JsonPayloadConverter.java index 46820e6b21..3c84393cc0 100644 --- a/temporal-sdk/src/main/java17/io/temporal/common/converter/Jackson3JsonPayloadConverter.java +++ b/temporal-sdk/src/main/java17/io/temporal/common/converter/Jackson3JsonPayloadConverter.java @@ -106,17 +106,31 @@ public String getEncodingType() { @Override public Optional toData(Object value) throws DataConverterException { try { - byte[] serialized = mapper.writeValueAsBytes(value); - return Optional.of( - Payload.newBuilder() - .putMetadata(EncodingKeys.METADATA_ENCODING_KEY, EncodingKeys.METADATA_ENCODING_JSON) - .setData(ByteString.copyFrom(serialized)) - .build()); + return toPayload(mapper.writeValueAsBytes(value)); } catch (JacksonException e) { throw new DataConverterException(e); } } + @Override + public Optional toData(Object value, Type valueType) throws DataConverterException { + try { + byte[] serialized = + mapper.writerFor(mapper.getTypeFactory().constructType(valueType)).writeValueAsBytes(value); + return toPayload(serialized); + } catch (JacksonException e) { + throw new DataConverterException(e); + } + } + + private Optional toPayload(byte[] serialized) { + return Optional.of( + Payload.newBuilder() + .putMetadata(EncodingKeys.METADATA_ENCODING_KEY, EncodingKeys.METADATA_ENCODING_JSON) + .setData(ByteString.copyFrom(serialized)) + .build()); + } + @Override public T fromData(Payload content, Class valueClass, Type valueType) throws DataConverterException { diff --git a/temporal-sdk/src/test/java/io/temporal/client/ActivityClientImplTest.java b/temporal-sdk/src/test/java/io/temporal/client/ActivityClientImplTest.java index 539fa33e6b..0524d9a183 100644 --- a/temporal-sdk/src/test/java/io/temporal/client/ActivityClientImplTest.java +++ b/temporal-sdk/src/test/java/io/temporal/client/ActivityClientImplTest.java @@ -16,8 +16,10 @@ import io.temporal.serviceclient.WorkflowServiceStubs; import io.temporal.serviceclient.WorkflowServiceStubsOptions; import io.temporal.workflow.Functions; +import java.lang.reflect.Type; import java.time.Duration; import java.util.Collections; +import java.util.List; import java.util.concurrent.atomic.AtomicReference; import org.junit.Before; import org.junit.Test; @@ -36,6 +38,12 @@ public interface WrongActivity { void doWrong(); } + @ActivityInterface + public interface GenericActivity { + @ActivityMethod + void doIt(List values); + } + private WorkflowServiceStubs stubs; private ActivityClient client; private StartActivityOptions options; @@ -91,6 +99,34 @@ public ActivityClientCallsInterceptor.StartActivityOutput startActivity( assertEquals(payload, capturedHeader.get().getValues().get("my-key")); } + @Test + public void testTypedStartIncludesArgumentTypes() throws NoSuchMethodException { + AtomicReference capturedTypes = new AtomicReference<>(); + ActivityClientInterceptor capturingInterceptor = + next -> + new ActivityClientCallsInterceptorBase(next) { + @Override + public ActivityClientCallsInterceptor.StartActivityOutput startActivity( + ActivityClientCallsInterceptor.StartActivityInput input) { + capturedTypes.set(input.getArgTypes()); + return new ActivityClientCallsInterceptor.StartActivityOutput("fake-id", null); + } + }; + ActivityClient client = + ActivityClient.newInstance( + stubs, + ActivityClientOptions.newBuilder() + .setInterceptors(Collections.singletonList(capturingInterceptor)) + .build()); + + client.start( + GenericActivity.class, GenericActivity::doIt, options, Collections.singletonList("value")); + + assertArrayEquals( + GenericActivity.class.getMethod("doIt", List.class).getGenericParameterTypes(), + capturedTypes.get()); + } + @Test(expected = NoSuchMethodError.class) @SuppressWarnings("unchecked") public void testStartWithMethodFromWrongClass() { diff --git a/temporal-sdk/src/test/java/io/temporal/common/converter/CodecDataConverterTest.java b/temporal-sdk/src/test/java/io/temporal/common/converter/CodecDataConverterTest.java index f3957ada40..530790b094 100644 --- a/temporal-sdk/src/test/java/io/temporal/common/converter/CodecDataConverterTest.java +++ b/temporal-sdk/src/test/java/io/temporal/common/converter/CodecDataConverterTest.java @@ -3,6 +3,9 @@ import static org.junit.Assert.assertArrayEquals; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; import com.google.protobuf.ByteString; import io.temporal.api.common.v1.Payload; @@ -12,6 +15,7 @@ import io.temporal.failure.TemporalFailure; import io.temporal.payload.codec.PayloadCodec; import io.temporal.payload.codec.PayloadCodecException; +import java.lang.reflect.Type; import java.util.Collections; import java.util.List; import java.util.Optional; @@ -119,6 +123,41 @@ public void testRawValuePassThrough() { assertEquals(p, converted.getPayload()); } + @Test + public void testSerializationTypeHintIsPassedToDataConverter() { + DataConverter delegate = mock(DataConverter.class); + Type typeHint = new com.google.common.reflect.TypeToken>() {}.getType(); + List value = Collections.singletonList("value"); + Payload payload = Payload.newBuilder().setData(ByteString.copyFromUtf8("value")).build(); + when(delegate.toPayload(value, typeHint)).thenReturn(Optional.of(payload)); + CodecDataConverter converter = + new CodecDataConverter( + delegate, Collections.singletonList(new PrefixPayloadCodec()), false); + + Optional encoded = converter.toPayload(value, typeHint); + + verify(delegate).toPayload(value, typeHint); + assertTrue(isEncoded(encoded.get())); + } + + @Test + public void testSerializationTypeHintsArePassedToDataConverter() { + DataConverter delegate = mock(DataConverter.class); + Type[] typeHints = {new com.google.common.reflect.TypeToken>() {}.getType()}; + Object[] values = {Collections.singletonList("value")}; + Payload payload = Payload.newBuilder().setData(ByteString.copyFromUtf8("value")).build(); + when(delegate.toPayloads(values, typeHints)) + .thenReturn(Optional.of(Payloads.newBuilder().addPayloads(payload).build())); + CodecDataConverter converter = + new CodecDataConverter( + delegate, Collections.singletonList(new PrefixPayloadCodec()), false); + + Optional encoded = converter.toPayloads(values, typeHints); + + verify(delegate).toPayloads(values, typeHints); + assertTrue(isEncoded(encoded.get().getPayloads(0))); + } + static boolean isEncoded(Payload payload) { return payload.getData().startsWith(PrefixPayloadCodec.PREFIX); } diff --git a/temporal-sdk/src/test/java/io/temporal/common/converter/DataConverterTest.java b/temporal-sdk/src/test/java/io/temporal/common/converter/DataConverterTest.java index 01a300b4b8..cbf1098f6b 100644 --- a/temporal-sdk/src/test/java/io/temporal/common/converter/DataConverterTest.java +++ b/temporal-sdk/src/test/java/io/temporal/common/converter/DataConverterTest.java @@ -1,7 +1,9 @@ package io.temporal.common.converter; +import io.temporal.api.common.v1.Payload; import io.temporal.api.common.v1.Payloads; import java.lang.reflect.Method; +import java.lang.reflect.Type; import java.util.List; import java.util.Optional; import org.junit.Assert; @@ -61,4 +63,171 @@ public void addGenericArrayParameter() throws NoSuchMethodException { Assert.assertEquals("test", result[0]); Assert.assertNull(result[1]); } + + @Test + public void passesSerializationTypeHintToPayloadConverter() throws NoSuchMethodException { + Method method = + this.getClass().getMethod("testMethodGenericParameter", String.class, List.class); + Type expectedType = method.getGenericParameterTypes()[1]; + Type[] receivedType = new Type[1]; + Payload expectedPayload = Payload.getDefaultInstance(); + PayloadConverter payloadConverter = + new PayloadConverter() { + @Override + public String getEncodingType() { + return "test/type-hint"; + } + + @Override + public Optional toData(Object value) { + throw new AssertionError("The type-aware overload should be used"); + } + + @Override + public Optional toData(Object value, Type valueType) { + receivedType[0] = valueType; + return Optional.of(expectedPayload); + } + + @Override + public T fromData(Payload content, Class valueType, Type valueGenericType) { + throw new UnsupportedOperationException(); + } + }; + + Optional result = + new DefaultDataConverter(payloadConverter) + .toPayload(java.util.Collections.emptyList(), expectedType); + + Assert.assertSame(expectedPayload, result.get()); + Assert.assertSame(expectedType, receivedType[0]); + } + + @Test + public void passesSerializationTypeHintsToPayloadConverter() throws NoSuchMethodException { + Method method = + this.getClass().getMethod("testMethodGenericParameter", String.class, List.class); + Type[] expectedTypes = method.getGenericParameterTypes(); + List receivedTypes = new java.util.ArrayList<>(); + PayloadConverter payloadConverter = + new PayloadConverter() { + @Override + public String getEncodingType() { + return "test/type-hints"; + } + + @Override + public Optional toData(Object value) { + throw new AssertionError("The type-aware overload should be used"); + } + + @Override + public Optional toData(Object value, Type valueType) { + receivedTypes.add(valueType); + return Optional.of(Payload.getDefaultInstance()); + } + + @Override + public T fromData(Payload content, Class valueType, Type valueGenericType) { + throw new UnsupportedOperationException(); + } + }; + + new DefaultDataConverter(payloadConverter) + .toPayloads(new Object[] {"value", java.util.Collections.emptyList()}, expectedTypes); + + Assert.assertArrayEquals(expectedTypes, receivedTypes.toArray(new Type[0])); + } + + @Test + public void typeHintFallsBackToLegacyPayloadConverter() { + Payload expectedPayload = Payload.getDefaultInstance(); + PayloadConverter payloadConverter = + new PayloadConverter() { + @Override + public String getEncodingType() { + return "test/legacy"; + } + + @Override + public Optional toData(Object value) { + return Optional.of(expectedPayload); + } + + @Override + public T fromData(Payload content, Class valueType, Type valueGenericType) { + throw new UnsupportedOperationException(); + } + }; + + Optional result = + new DefaultDataConverter(payloadConverter).toPayload("value", String.class); + + Assert.assertSame(expectedPayload, result.get()); + } + + @Test + public void typeHintFallsBackToLegacyDataConverter() { + Payload expectedPayload = Payload.getDefaultInstance(); + DataConverter dataConverter = + new DataConverter() { + @Override + public Optional toPayload(T value) { + return Optional.of(expectedPayload); + } + + @Override + public T fromPayload(Payload payload, Class valueClass, Type valueType) { + throw new UnsupportedOperationException(); + } + + @Override + public Optional toPayloads(Object... values) { + throw new UnsupportedOperationException(); + } + + @Override + public T fromPayloads( + int index, Optional content, Class valueType, Type valueGenericType) { + throw new UnsupportedOperationException(); + } + }; + + Optional result = dataConverter.toPayload("value", String.class); + + Assert.assertSame(expectedPayload, result.get()); + } + + @Test + public void typeHintsFallBackToLegacyDataConverter() { + Payloads expectedPayloads = Payloads.getDefaultInstance(); + DataConverter dataConverter = + new DataConverter() { + @Override + public Optional toPayload(T value) { + throw new UnsupportedOperationException(); + } + + @Override + public T fromPayload(Payload payload, Class valueClass, Type valueType) { + throw new UnsupportedOperationException(); + } + + @Override + public Optional toPayloads(Object... values) { + return Optional.of(expectedPayloads); + } + + @Override + public T fromPayloads( + int index, Optional content, Class valueType, Type valueGenericType) { + throw new UnsupportedOperationException(); + } + }; + + Optional result = + dataConverter.toPayloads(new Object[] {"value"}, new Type[] {String.class}); + + Assert.assertSame(expectedPayloads, result.get()); + } } diff --git a/temporal-sdk/src/test/java/io/temporal/common/converter/GsonJsonPayloadConverterTest.java b/temporal-sdk/src/test/java/io/temporal/common/converter/GsonJsonPayloadConverterTest.java new file mode 100644 index 0000000000..d0c5afb517 --- /dev/null +++ b/temporal-sdk/src/test/java/io/temporal/common/converter/GsonJsonPayloadConverterTest.java @@ -0,0 +1,28 @@ +package io.temporal.common.converter; + +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertTrue; + +import io.temporal.api.common.v1.Payload; +import org.junit.Test; + +public class GsonJsonPayloadConverterTest { + + @Test + public void serializationUsesTypeHint() { + GsonJsonPayloadConverter converter = new GsonJsonPayloadConverter(); + Payload payload = converter.toData(new Child(), Parent.class).get(); + String json = payload.getData().toStringUtf8(); + + assertTrue(json.contains("parent")); + assertFalse(json.contains("child")); + } + + private static class Parent { + private final String parent = "parent"; + } + + private static class Child extends Parent { + private final String child = "child"; + } +} diff --git a/temporal-sdk/src/test/java/io/temporal/common/converter/JacksonJsonPayloadConverterTest.java b/temporal-sdk/src/test/java/io/temporal/common/converter/JacksonJsonPayloadConverterTest.java index 5154bdc0ad..66e0a90677 100644 --- a/temporal-sdk/src/test/java/io/temporal/common/converter/JacksonJsonPayloadConverterTest.java +++ b/temporal-sdk/src/test/java/io/temporal/common/converter/JacksonJsonPayloadConverterTest.java @@ -4,9 +4,16 @@ import static org.junit.Assert.assertTrue; import static org.junit.Assert.fail; +import com.fasterxml.jackson.annotation.JsonSubTypes; +import com.fasterxml.jackson.annotation.JsonTypeInfo; +import com.google.common.reflect.TypeToken; +import io.temporal.api.common.v1.Payload; import io.temporal.api.common.v1.Payloads; import java.lang.reflect.InvocationTargetException; +import java.lang.reflect.Type; import java.time.Instant; +import java.util.Collections; +import java.util.List; import java.util.Objects; import java.util.Optional; import org.junit.After; @@ -86,6 +93,36 @@ public void testJsonWithOptional() { assertEquals("myPayload", converted.getName().get()); } + @Test + public void serializationUsesTypeHint() { + JacksonJsonPayloadConverter converter = new JacksonJsonPayloadConverter(); + Type type = new TypeToken>() {}.getType(); + + Payload payload = converter.toData(Collections.singletonList(new Cat("Milo")), type).get(); + List converted = converter.fromData(payload, List.class, type); + + assertTrue(converted.get(0) instanceof Cat); + assertEquals("Milo", ((Cat) converted.get(0)).getName()); + } + + @JsonTypeInfo(use = JsonTypeInfo.Id.NAME, property = "type") + @JsonSubTypes(@JsonSubTypes.Type(value = Cat.class, name = "cat")) + private interface Animal {} + + private static class Cat implements Animal { + private String name; + + public Cat() {} + + Cat(String name) { + this.name = name; + } + + public String getName() { + return name; + } + } + static class TestOptionalPayload { private Optional id; private Optional timestamp; diff --git a/temporal-sdk/src/test/java/io/temporal/functional/serialization/SerializationTypeHintTest.java b/temporal-sdk/src/test/java/io/temporal/functional/serialization/SerializationTypeHintTest.java new file mode 100644 index 0000000000..d0e309a114 --- /dev/null +++ b/temporal-sdk/src/test/java/io/temporal/functional/serialization/SerializationTypeHintTest.java @@ -0,0 +1,214 @@ +package io.temporal.functional.serialization; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertTrue; + +import com.fasterxml.jackson.annotation.JsonSubTypes; +import com.fasterxml.jackson.annotation.JsonTypeInfo; +import com.google.common.reflect.TypeToken; +import io.temporal.activity.ActivityInterface; +import io.temporal.activity.ActivityMethod; +import io.temporal.activity.ActivityOptions; +import io.temporal.api.common.v1.Payload; +import io.temporal.client.WorkflowClient; +import io.temporal.client.WorkflowClientOptions; +import io.temporal.client.WorkflowOptions; +import io.temporal.client.WorkflowStub; +import io.temporal.common.converter.DataConverter; +import io.temporal.common.converter.DataConverterException; +import io.temporal.common.converter.DefaultDataConverter; +import io.temporal.common.converter.JacksonJsonPayloadConverter; +import io.temporal.common.converter.PayloadConverter; +import io.temporal.testing.internal.SDKTestWorkflowRule; +import io.temporal.workflow.ChildWorkflowOptions; +import io.temporal.workflow.QueryMethod; +import io.temporal.workflow.SignalMethod; +import io.temporal.workflow.UpdateMethod; +import io.temporal.workflow.Workflow; +import io.temporal.workflow.WorkflowInterface; +import io.temporal.workflow.WorkflowMethod; +import java.lang.reflect.Type; +import java.time.Duration; +import java.util.Collections; +import java.util.List; +import java.util.Optional; +import java.util.concurrent.ConcurrentLinkedQueue; +import org.junit.Rule; +import org.junit.Test; + +public class SerializationTypeHintTest { + private static final Type ANIMALS_TYPE = new TypeToken>() {}.getType(); + + private final TypeRecordingPayloadConverter recordingConverter = + new TypeRecordingPayloadConverter(); + private final DataConverter dataConverter = + DefaultDataConverter.newDefaultInstance().withPayloadConverterOverrides(recordingConverter); + + @Rule + public SDKTestWorkflowRule testWorkflowRule = + SDKTestWorkflowRule.newBuilder() + .setWorkflowClientOptions( + WorkflowClientOptions.newBuilder().setDataConverter(dataConverter).build()) + .setWorkflowTypes(TestWorkflowImpl.class, ChildWorkflowImpl.class) + .setActivityImplementations(new TestActivitiesImpl()) + .build(); + + @Test + public void typedCallersPassSerializationTypeHints() { + TestWorkflow workflow = + testWorkflowRule + .getWorkflowClient() + .newWorkflowStub( + TestWorkflow.class, + WorkflowOptions.newBuilder().setTaskQueue(testWorkflowRule.getTaskQueue()).build()); + List initial = Collections.singletonList(new Cat("initial")); + List updated = Collections.singletonList(new Cat("updated")); + + WorkflowClient.start(workflow::run, initial); + SDKTestWorkflowRule.waitForOKQuery(WorkflowStub.fromTyped(workflow)); + + assertEquals(initial, workflow.echo(initial)); + assertEquals(updated, workflow.update(updated)); + workflow.finish(updated); + assertEquals(updated, WorkflowStub.fromTyped(workflow).getResult(List.class, ANIMALS_TYPE)); + + assertTrue(recordingConverter.receivedTypes.size() >= 11); + assertTrue(recordingConverter.receivedTypes.stream().allMatch(ANIMALS_TYPE::equals)); + } + + @WorkflowInterface + public interface TestWorkflow { + @WorkflowMethod + List run(List animals); + + @QueryMethod + List echo(List animals); + + @UpdateMethod + List update(List animals); + + @SignalMethod + void finish(List animals); + } + + public static class TestWorkflowImpl implements TestWorkflow { + private final TestActivities activities = + Workflow.newActivityStub( + TestActivities.class, + ActivityOptions.newBuilder().setStartToCloseTimeout(Duration.ofSeconds(10)).build()); + private final ChildWorkflow child = + Workflow.newChildWorkflowStub( + ChildWorkflow.class, ChildWorkflowOptions.newBuilder().build()); + private List animals; + private boolean finished; + + @Override + public List run(List animals) { + this.animals = animals; + Workflow.await(() -> finished); + return activities.echo(child.run(this.animals)); + } + + @Override + public List echo(List animals) { + return animals; + } + + @Override + public List update(List animals) { + this.animals = animals; + return animals; + } + + @Override + public void finish(List animals) { + this.animals = animals; + this.finished = true; + } + } + + @ActivityInterface + public interface TestActivities { + @ActivityMethod + List echo(List animals); + } + + public static class TestActivitiesImpl implements TestActivities { + @Override + public List echo(List animals) { + return animals; + } + } + + @WorkflowInterface + public interface ChildWorkflow { + @WorkflowMethod + List run(List animals); + } + + public static class ChildWorkflowImpl implements ChildWorkflow { + @Override + public List run(List animals) { + return animals; + } + } + + @JsonTypeInfo(use = JsonTypeInfo.Id.NAME, property = "type") + @JsonSubTypes(@JsonSubTypes.Type(value = Cat.class, name = "cat")) + public interface Animal {} + + public static class Cat implements Animal { + private String name; + + public Cat() {} + + public Cat(String name) { + this.name = name; + } + + public String getName() { + return name; + } + + public void setName(String name) { + this.name = name; + } + + @Override + public boolean equals(Object o) { + return o instanceof Cat && name.equals(((Cat) o).name); + } + + @Override + public int hashCode() { + return name.hashCode(); + } + } + + private static final class TypeRecordingPayloadConverter implements PayloadConverter { + private final PayloadConverter delegate = new JacksonJsonPayloadConverter(); + private final ConcurrentLinkedQueue receivedTypes = new ConcurrentLinkedQueue<>(); + + @Override + public String getEncodingType() { + return delegate.getEncodingType(); + } + + @Override + public Optional toData(Object value) throws DataConverterException { + return delegate.toData(value); + } + + @Override + public Optional toData(Object value, Type valueType) throws DataConverterException { + receivedTypes.add(valueType); + return delegate.toData(value, valueType); + } + + @Override + public T fromData(Payload content, Class valueType, Type valueGenericType) + throws DataConverterException { + return delegate.fromData(content, valueType, valueGenericType); + } + } +} diff --git a/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageDataConverterTest.java b/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageDataConverterTest.java index 3cc35284bf..b960c54c28 100644 --- a/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageDataConverterTest.java +++ b/temporal-sdk/src/test/java/io/temporal/internal/payload/storage/ExternalStorageDataConverterTest.java @@ -6,6 +6,9 @@ import static org.junit.Assert.assertNull; import static org.junit.Assert.assertThrows; import static org.junit.Assert.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; import com.google.protobuf.ByteString; import io.temporal.api.common.v1.Payload; @@ -80,6 +83,36 @@ public void singlePayloadRoundTrips() { assertEquals("value", converter.fromPayload(stored.get(), String.class, String.class)); } + @Test + public void singlePayloadPassesSerializationTypeHintToDelegate() { + DataConverter delegate = mock(DataConverter.class); + Type typeHint = new com.google.common.reflect.TypeToken>() {}.getType(); + List value = Collections.singletonList("value"); + Payload payload = Payload.getDefaultInstance(); + when(delegate.toPayload(value, typeHint)).thenReturn(Optional.of(payload)); + DataConverter converter = new ExternalStorageDataConverter(delegate, null); + + Optional converted = converter.toPayload(value, typeHint); + + verify(delegate).toPayload(value, typeHint); + assertEquals(payload, converted.get()); + } + + @Test + public void payloadsPassSerializationTypeHintsToDelegate() { + DataConverter delegate = mock(DataConverter.class); + Type[] typeHints = {new com.google.common.reflect.TypeToken>() {}.getType()}; + Object[] values = {Collections.singletonList("value")}; + Payloads payloads = Payloads.newBuilder().addPayloads(Payload.getDefaultInstance()).build(); + when(delegate.toPayloads(values, typeHints)).thenReturn(Optional.of(payloads)); + DataConverter converter = new ExternalStorageDataConverter(delegate, null); + + Optional converted = converter.toPayloads(values, typeHints); + + verify(delegate).toPayloads(values, typeHints); + assertEquals(payloads, converted.get()); + } + @Test public void failureDetailsRoundTrip() { DataConverter converter = resolving(TestStorageDriver.create(), 0); diff --git a/temporal-spring-boot-autoconfigure/src/test/java/io/temporal/spring/boot/autoconfigure/CustomDataConverterTest.java b/temporal-spring-boot-autoconfigure/src/test/java/io/temporal/spring/boot/autoconfigure/CustomDataConverterTest.java index b3685a8cdd..204d7ea1f3 100644 --- a/temporal-spring-boot-autoconfigure/src/test/java/io/temporal/spring/boot/autoconfigure/CustomDataConverterTest.java +++ b/temporal-spring-boot-autoconfigure/src/test/java/io/temporal/spring/boot/autoconfigure/CustomDataConverterTest.java @@ -9,6 +9,7 @@ import io.temporal.common.converter.DefaultDataConverter; import io.temporal.spring.boot.autoconfigure.bytaskqueue.TestWorkflow; import io.temporal.testing.TestWorkflowEnvironment; +import java.lang.reflect.Type; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.TestInstance; @@ -45,7 +46,7 @@ public void testCustomDataConverterBeanIsPickedUpByTestWorkflowEnvironment() { workflowClient.newWorkflowStub( TestWorkflow.class, WorkflowOptions.newBuilder().setTaskQueue("UnitTest").build()); testWorkflow.execute("input"); - verify(spyDataConverter, atLeastOnce()).toPayloads(any()); + verify(spyDataConverter, atLeastOnce()).toPayloads(any(Object[].class), any(Type[].class)); } @ComponentScan( diff --git a/temporal-testing/src/main/java/io/temporal/testing/TestActivityEnvironmentInternal.java b/temporal-testing/src/main/java/io/temporal/testing/TestActivityEnvironmentInternal.java index dc2a0d3d3c..b90b979927 100644 --- a/temporal-testing/src/main/java/io/temporal/testing/TestActivityEnvironmentInternal.java +++ b/temporal-testing/src/main/java/io/temporal/testing/TestActivityEnvironmentInternal.java @@ -276,7 +276,7 @@ public ActivityOutput executeActivity(ActivityInput i) { testEnvironmentOptions .getWorkflowClientOptions() .getDataConverter() - .toPayloads(i.getArgs()); + .toPayloads(i.getArgs(), i.getArgTypes()); Optional heartbeatPayload = Optional.ofNullable(heartbeatDetails.getAndSet(null)) .flatMap( @@ -319,7 +319,7 @@ public LocalActivityOutput executeLocalActivity(LocalActivityInput i) testEnvironmentOptions .getWorkflowClientOptions() .getDataConverter() - .toPayloads(i.getArgs()); + .toPayloads(i.getArgs(), i.getArgTypes()); LocalActivityOptions options = i.getOptions(); PollActivityTaskQueueResponse.Builder taskBuilder = PollActivityTaskQueueResponse.newBuilder() diff --git a/temporal-testing/src/main/java/io/temporal/testing/TimeLockingInterceptor.java b/temporal-testing/src/main/java/io/temporal/testing/TimeLockingInterceptor.java index f338ca4e0a..c1739bd700 100644 --- a/temporal-testing/src/main/java/io/temporal/testing/TimeLockingInterceptor.java +++ b/temporal-testing/src/main/java/io/temporal/testing/TimeLockingInterceptor.java @@ -49,17 +49,38 @@ public void signal(String signalName, Object... args) { next.signal(signalName, args); } + @Override + public void signalWithTypeHints(String signalName, Type[] argTypes, Object... args) { + next.signalWithTypeHints(signalName, argTypes, args); + } + @Override public WorkflowExecution start(Object... args) { return next.start(args); } + @Override + public WorkflowExecution startWithTypeHints(Type[] argTypes, Object... args) { + return next.startWithTypeHints(argTypes, args); + } + @Override public WorkflowUpdateHandle startUpdateWithStart( UpdateOptions options, Object[] updateArgs, Object[] startArgs) { return next.startUpdateWithStart(options, updateArgs, startArgs); } + @Override + public WorkflowUpdateHandle startUpdateWithStartWithTypeHints( + UpdateOptions options, + Object[] updateArgs, + Type[] updateArgTypes, + Object[] startArgs, + Type[] startArgTypes) { + return next.startUpdateWithStartWithTypeHints( + options, updateArgs, updateArgTypes, startArgs, startArgTypes); + } + @Override public R executeUpdateWithStart( UpdateOptions updateOptions, Object[] updateArgs, Object[] startArgs) { @@ -72,6 +93,17 @@ public WorkflowExecution signalWithStart( return next.signalWithStart(signalName, signalArgs, startArgs); } + @Override + public WorkflowExecution signalWithStartWithTypeHints( + String signalName, + Object[] signalArgs, + Type[] signalArgTypes, + Object[] startArgs, + Type[] startArgTypes) { + return next.signalWithStartWithTypeHints( + signalName, signalArgs, signalArgTypes, startArgs, startArgTypes); + } + @Override public Optional getWorkflowType() { return next.getWorkflowType(); @@ -159,6 +191,12 @@ public R query(String queryType, Class resultClass, Type resultType, Obje return next.query(queryType, resultClass, resultType, args); } + @Override + public R queryWithTypeHints( + String queryType, Class resultClass, Type resultType, Type[] argTypes, Object... args) { + return next.queryWithTypeHints(queryType, resultClass, resultType, argTypes, args); + } + @Override public void cancel() { next.cancel(); @@ -242,6 +280,12 @@ public R update(String updateName, Class resultClass, Object... args) { return next.update(updateName, resultClass, args); } + @Override + public R updateWithTypeHints( + String updateName, Class resultClass, Type resultType, Type[] argTypes, Object... args) { + return next.updateWithTypeHints(updateName, resultClass, resultType, argTypes, args); + } + @Override public WorkflowUpdateHandle startUpdate( String updateName, WorkflowUpdateStage waitForStage, Class resultClass, Object... args) { @@ -253,6 +297,12 @@ public WorkflowUpdateHandle startUpdate(UpdateOptions options, Object. return next.startUpdate(options, args); } + @Override + public WorkflowUpdateHandle startUpdateWithTypeHints( + UpdateOptions options, Type[] argTypes, Object... args) { + return next.startUpdateWithTypeHints(options, argTypes, args); + } + @Override public WorkflowUpdateHandle getUpdateHandle(String updateId, Class resultClass) { return next.getUpdateHandle(updateId, resultClass); diff --git a/temporal-testing/src/main/java/io/temporal/testing/internal/TracingWorkerInterceptor.java b/temporal-testing/src/main/java/io/temporal/testing/internal/TracingWorkerInterceptor.java index 77d82bdaee..9401794855 100644 --- a/temporal-testing/src/main/java/io/temporal/testing/internal/TracingWorkerInterceptor.java +++ b/temporal-testing/src/main/java/io/temporal/testing/internal/TracingWorkerInterceptor.java @@ -321,6 +321,7 @@ public void registerQuery(RegisterQueryInput input) { input.getDescription(), input.getArgTypes(), input.getGenericArgTypes(), + input.getResultType(), (args) -> { Object result = input.getCallback().apply(args); if (!WorkflowUnsafe.isReplaying()) {