diff --git a/build.gradle b/build.gradle index 98bf5158a..20b88e924 100644 --- a/build.gradle +++ b/build.gradle @@ -21,7 +21,7 @@ subprojects { ext { otelVersion = '1.30.1' otelVersionAlpha = "${otelVersion}-alpha" - javaSDKVersion = '1.39.0' + javaSDKVersion = '1.40.0' camelVersion = '3.22.1' jarVersion = '1.0.0' } @@ -49,4 +49,4 @@ subprojects { test { useJUnitPlatform() } -} \ No newline at end of file +} diff --git a/core/src/main/java/io/temporal/samples/hello/HelloAccumulator.java b/core/src/main/java/io/temporal/samples/hello/HelloAccumulator.java index 2f8d67457..037b0f723 100644 --- a/core/src/main/java/io/temporal/samples/hello/HelloAccumulator.java +++ b/core/src/main/java/io/temporal/samples/hello/HelloAccumulator.java @@ -15,6 +15,7 @@ import io.temporal.serviceclient.WorkflowServiceStubs; import io.temporal.worker.Worker; import io.temporal.worker.WorkerFactory; +import io.temporal.workflow.Promise; import io.temporal.workflow.SignalMethod; import io.temporal.workflow.Workflow; import io.temporal.workflow.WorkflowInterface; @@ -198,8 +199,10 @@ public String accumulateGreetings( // - if exit signal is received, process any remaining signals and exit do { - boolean timedout = - !Workflow.await(MAX_AWAIT_TIME, () -> !unprocessedGreetings.isEmpty() || exitRequested); + Promise timer = Workflow.newTimer(MAX_AWAIT_TIME); + Workflow.await( + () -> timer.isCompleted() || !unprocessedGreetings.isEmpty() || exitRequested); + boolean timedout = timer.isCompleted() && unprocessedGreetings.isEmpty() && !exitRequested; while (!unprocessedGreetings.isEmpty()) { processGreeting(unprocessedGreetings.removeFirst()); diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/README.md b/core/src/main/java/io/temporal/samples/nexusserializationcontext/README.md new file mode 100644 index 000000000..32fc41bae --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/README.md @@ -0,0 +1,115 @@ +# Nexus serialization context + +This sample calls a synchronous and an asynchronous Nexus operation through each of +two endpoints. Each endpoint routes to its own handler worker, which registers both +`SyncEchoService` and `AsyncEchoService`. The caller's `NexusCodec` uses +`NexusSerializationContext` to select the `PayloadCodec` registered for each endpoint. +The sample uses these keys: + +- Key A: compresses with zlib, then encrypts synchronous and asynchronous payloads for `nexus-serialization-compressed-encrypted`. +- Key B: encrypts synchronous and asynchronous payloads for `nexus-serialization-encrypted`. +- Key C: encrypts the caller workflow's input and final result. + +The caller configures a codec for each endpoint and a separate codec for its own +workflow input and result. Each handler worker uses the fixed key for its endpoint. + +The caller schedules all four operations before waiting for their results, so each +result must be decoded using the context of its own endpoint. The caller's +`NexusCodec` uses Key C for its own workflow payloads, which have no Nexus endpoint. + +For both Nexus endpoints, the outer encrypted payload stores the endpoint name in +`nexus-endpoint-name` metadata alongside `binary/nexus-aes-gcm` and a sample key ID +(`key-a` or `key-b`). A Codec Server can use the endpoint name to select the matching +key and decompression chain without SDK context. + +For the compressed endpoint, the outer payload's decoded metadata looks like: + +```text +encoding: binary/nexus-aes-gcm +encryption-key-id: key-a +nexus-endpoint-name: nexus-serialization-compressed-encrypted +``` + +The asynchronous results have `encryption-key-id: key-a` or `key-b` and their respective +endpoint names in the outer payload metadata. +The caller workflow's input and final result use `encryption-key-id: key-c`; they do not +have endpoint metadata. The starter uses the same converter as the caller worker, so it +can decode the final result before printing it. + +`NexusSerializationContext` works end to end for synchronous Nexus operations. +For an asynchronous operation, the handler's final result is serialized as a workflow +result and does not receive `NexusSerializationContext`. We plan to add +`NexusSerializationContext` support for asynchronous operation results in the Java SDK. +Until then, the sample's `NexusEndpointInterceptor` captures the endpoint when the +handler starts the backing `EchoWorkflow`. `NexusEndpointContextPropagator` saves it in +the workflow headers and restores it on the workflow thread. Each handler's codec uses +its fixed key to encrypt the workflow result and includes the propagated endpoint name +in the outer payload metadata. This propagation is needed for the endpoint metadata; +the handler already knows which key to use. When the result reaches the caller, the SDK +supplies `NexusSerializationContext` to decode it. + +This endpoint propagation covers asynchronous operations backed by workflows. If the +endpoint is not propagated, the result remains encrypted with the handler's fixed key +but lacks the endpoint name in its metadata. + +The hard-coded keys are only for this local example. For production encryption, +use a secure key store, as in the +[AWS Encryption SDK sample](../keymanagementencryption/awsencryptionsdk/README.md). + +Requires Java SDK 1.40.0 or later and Temporal Server 1.30.0 or later with Nexus enabled +so the handler can read the endpoint name. + +## Run locally + +Start a Temporal dev server: + +```bash +temporal server start-dev +``` + +In another terminal, create the namespaces and endpoints: + +```bash +temporal operator namespace create --namespace nexus-serialization-caller +temporal operator namespace create --namespace nexus-serialization-key-a-handler +temporal operator namespace create --namespace nexus-serialization-key-b-handler +temporal operator nexus endpoint create \ + --name nexus-serialization-compressed-encrypted \ + --target-namespace nexus-serialization-key-a-handler \ + --target-task-queue nexus-serialization-key-a-handler +temporal operator nexus endpoint create \ + --name nexus-serialization-encrypted \ + --target-namespace nexus-serialization-key-b-handler \ + --target-task-queue nexus-serialization-key-b-handler +``` + +Run each of the following in its own terminal from the repository root: + +```bash +./gradlew -q :core:execute -PmainClass=io.temporal.samples.nexusserializationcontext.handler.CompressedEncryptedHandlerWorker \ + --args="-namespace nexus-serialization-key-a-handler" +``` + +```bash +./gradlew -q :core:execute -PmainClass=io.temporal.samples.nexusserializationcontext.handler.EncryptedHandlerWorker \ + --args="-namespace nexus-serialization-key-b-handler" +``` + +```bash +./gradlew -q :core:execute -PmainClass=io.temporal.samples.nexusserializationcontext.caller.CallerWorker \ + --args="-namespace nexus-serialization-caller" +``` + +```bash +./gradlew -q :core:execute -PmainClass=io.temporal.samples.nexusserializationcontext.caller.CallerStarter \ + --args="-namespace nexus-serialization-caller" +``` + +The starter will print: + +```text +Compressed and encrypted endpoint sync result: Hello from Nexus +Encrypted endpoint sync result: Hello from Nexus +Compressed and encrypted endpoint async result: Hello from Nexus +Encrypted endpoint async result: Hello from Nexus +``` diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/SampleConfig.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/SampleConfig.java new file mode 100644 index 000000000..1ba965c4c --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/SampleConfig.java @@ -0,0 +1,18 @@ +package io.temporal.samples.nexusserializationcontext; + +public final class SampleConfig { + public static final String COMPRESSED_ENCRYPTED_ENDPOINT = + "nexus-serialization-compressed-encrypted"; + public static final String ENCRYPTED_ENDPOINT = "nexus-serialization-encrypted"; + public static final String ENDPOINT_METADATA_KEY = "nexus-endpoint-name"; + public static final String KEY_A_ID = "key-a"; + public static final String KEY_B_ID = "key-b"; + public static final String KEY_C_ID = "key-c"; + + // Hard-coded keys are only for this local sample. + public static final String KEY_A_VALUE = "sample-key-A-123"; + public static final String KEY_B_VALUE = "sample-key-B-123"; + public static final String KEY_C_VALUE = "sample-key-C-123"; + + private SampleConfig() {} +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerStarter.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerStarter.java new file mode 100644 index 000000000..7defd5223 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerStarter.java @@ -0,0 +1,29 @@ +package io.temporal.samples.nexusserializationcontext.caller; + +import io.temporal.client.WorkflowClient; +import io.temporal.client.WorkflowClientOptions; +import io.temporal.client.WorkflowOptions; +import io.temporal.samples.nexus.options.ClientOptions; + +public class CallerStarter { + public static void main(String[] args) { + WorkflowClient client = + ClientOptions.getWorkflowClient( + args, + WorkflowClientOptions.newBuilder().setDataConverter(CallerWorker.dataConverter())); + CallerWorkflow workflow = + client.newWorkflowStub( + CallerWorkflow.class, + WorkflowOptions.newBuilder().setTaskQueue(CallerWorker.TASK_QUEUE).build()); + + EndpointResults results = workflow.echoThroughEndpoints("Hello from Nexus"); + System.out.println( + "Compressed and encrypted endpoint sync result: " + + results.compressedEncryptedSyncResult()); + System.out.println("Encrypted endpoint sync result: " + results.encryptedSyncResult()); + System.out.println( + "Compressed and encrypted endpoint async result: " + + results.compressedEncryptedAsyncResult()); + System.out.println("Encrypted endpoint async result: " + results.encryptedAsyncResult()); + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerWorker.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerWorker.java new file mode 100644 index 000000000..3c9938091 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerWorker.java @@ -0,0 +1,58 @@ +package io.temporal.samples.nexusserializationcontext.caller; + +import io.temporal.client.WorkflowClient; +import io.temporal.client.WorkflowClientOptions; +import io.temporal.common.converter.CodecDataConverter; +import io.temporal.common.converter.DataConverter; +import io.temporal.common.converter.DefaultDataConverter; +import io.temporal.payload.codec.ChainCodec; +import io.temporal.payload.codec.PayloadCodec; +import io.temporal.samples.nexus.options.ClientOptions; +import io.temporal.samples.nexusserializationcontext.SampleConfig; +import io.temporal.samples.nexusserializationcontext.codec.AesGcmCodec; +import io.temporal.samples.nexusserializationcontext.codec.NexusCodec; +import io.temporal.samples.nexusserializationcontext.codec.ZlibCodec; +import io.temporal.worker.Worker; +import io.temporal.worker.WorkerFactory; +import java.nio.charset.StandardCharsets; +import java.util.List; +import java.util.Map; +import javax.crypto.SecretKey; +import javax.crypto.spec.SecretKeySpec; + +public class CallerWorker { + static final String TASK_QUEUE = "nexus-serialization-caller"; + + public static void main(String[] args) { + WorkflowClient client = + ClientOptions.getWorkflowClient( + args, WorkflowClientOptions.newBuilder().setDataConverter(dataConverter())); + WorkerFactory factory = WorkerFactory.newInstance(client); + Worker worker = factory.newWorker(TASK_QUEUE); + worker.registerWorkflowImplementationTypes(CallerWorkflowImpl.class); + factory.start(); + } + + public static DataConverter dataConverter() { + SecretKey keyA = + new SecretKeySpec(SampleConfig.KEY_A_VALUE.getBytes(StandardCharsets.UTF_8), "AES"); + SecretKey keyB = + new SecretKeySpec(SampleConfig.KEY_B_VALUE.getBytes(StandardCharsets.UTF_8), "AES"); + SecretKey keyC = + new SecretKeySpec(SampleConfig.KEY_C_VALUE.getBytes(StandardCharsets.UTF_8), "AES"); + // ChainCodec encodes last to first: compress, then encrypt. + PayloadCodec compressedEncrypted = + new ChainCodec(List.of(new AesGcmCodec(SampleConfig.KEY_A_ID, keyA), new ZlibCodec())); + PayloadCodec encrypted = new AesGcmCodec(SampleConfig.KEY_B_ID, keyB); + return new CodecDataConverter( + DefaultDataConverter.newDefaultInstance(), + List.of( + new NexusCodec( + Map.of( + SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT, + compressedEncrypted, + SampleConfig.ENCRYPTED_ENDPOINT, + encrypted), + new AesGcmCodec(SampleConfig.KEY_C_ID, keyC)))); + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerWorkflow.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerWorkflow.java new file mode 100644 index 000000000..a7bbe109c --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerWorkflow.java @@ -0,0 +1,10 @@ +package io.temporal.samples.nexusserializationcontext.caller; + +import io.temporal.workflow.WorkflowInterface; +import io.temporal.workflow.WorkflowMethod; + +@WorkflowInterface +public interface CallerWorkflow { + @WorkflowMethod + EndpointResults echoThroughEndpoints(String message); +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerWorkflowImpl.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerWorkflowImpl.java new file mode 100644 index 000000000..9e0505b05 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/CallerWorkflowImpl.java @@ -0,0 +1,51 @@ +package io.temporal.samples.nexusserializationcontext.caller; + +import io.temporal.samples.nexusserializationcontext.SampleConfig; +import io.temporal.samples.nexusserializationcontext.service.AsyncEchoService; +import io.temporal.samples.nexusserializationcontext.service.SyncEchoService; +import io.temporal.workflow.NexusOperationHandle; +import io.temporal.workflow.NexusOperationOptions; +import io.temporal.workflow.NexusServiceOptions; +import io.temporal.workflow.Workflow; +import java.time.Duration; + +public class CallerWorkflowImpl implements CallerWorkflow { + @Override + public EndpointResults echoThroughEndpoints(String message) { + SyncEchoService compressedEncryptedService = + serviceFor(SyncEchoService.class, SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT); + SyncEchoService encryptedService = + serviceFor(SyncEchoService.class, SampleConfig.ENCRYPTED_ENDPOINT); + AsyncEchoService compressedEncryptedAsyncService = + serviceFor(AsyncEchoService.class, SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT); + AsyncEchoService encryptedAsyncService = + serviceFor(AsyncEchoService.class, SampleConfig.ENCRYPTED_ENDPOINT); + + // Start all operations before awaiting results. Each result keeps its endpoint context. + NexusOperationHandle compressedEncryptedSync = + Workflow.startNexusOperation(compressedEncryptedService::echo, message); + NexusOperationHandle encryptedSync = + Workflow.startNexusOperation(encryptedService::echo, message); + NexusOperationHandle compressedEncryptedAsync = + Workflow.startNexusOperation(compressedEncryptedAsyncService::echoAsync, message); + NexusOperationHandle encryptedAsync = + Workflow.startNexusOperation(encryptedAsyncService::echoAsync, message); + return new EndpointResults( + compressedEncryptedSync.getResult().get(), + encryptedSync.getResult().get(), + compressedEncryptedAsync.getResult().get(), + encryptedAsync.getResult().get()); + } + + private static T serviceFor(Class serviceClass, String endpoint) { + return Workflow.newNexusServiceStub( + serviceClass, + NexusServiceOptions.newBuilder() + .setEndpoint(endpoint) + .setOperationOptions( + NexusOperationOptions.newBuilder() + .setScheduleToCloseTimeout(Duration.ofSeconds(30)) + .build()) + .build()); + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/EndpointResults.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/EndpointResults.java new file mode 100644 index 000000000..741f9286e --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/caller/EndpointResults.java @@ -0,0 +1,7 @@ +package io.temporal.samples.nexusserializationcontext.caller; + +public record EndpointResults( + String compressedEncryptedSyncResult, + String encryptedSyncResult, + String compressedEncryptedAsyncResult, + String encryptedAsyncResult) {} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/AesGcmCodec.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/AesGcmCodec.java new file mode 100644 index 000000000..0208849a2 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/AesGcmCodec.java @@ -0,0 +1,140 @@ +package io.temporal.samples.nexusserializationcontext.codec; + +import com.google.protobuf.ByteString; +import com.google.protobuf.InvalidProtocolBufferException; +import io.temporal.api.common.v1.Payload; +import io.temporal.common.converter.EncodingKeys; +import io.temporal.payload.codec.PayloadCodec; +import io.temporal.payload.codec.PayloadCodecException; +import io.temporal.payload.context.NexusSerializationContext; +import io.temporal.payload.context.SerializationContext; +import io.temporal.samples.nexusserializationcontext.SampleConfig; +import io.temporal.samples.nexusserializationcontext.propagation.NexusEndpointContextPropagator; +import java.nio.ByteBuffer; +import java.security.GeneralSecurityException; +import java.security.SecureRandom; +import java.util.ArrayList; +import java.util.List; +import javax.annotation.Nonnull; +import javax.crypto.Cipher; +import javax.crypto.SecretKey; +import javax.crypto.spec.GCMParameterSpec; +import org.apache.commons.lang.StringUtils; + +/** Encrypts payloads with a configured key and includes endpoint metadata when available. */ +public final class AesGcmCodec implements PayloadCodec { + static final String KEY_ID_METADATA_KEY = "encryption-key-id"; + + private static final String CIPHER = "AES/GCM/NoPadding"; + private static final int NONCE_LENGTH = 12; + private static final int TAG_LENGTH_BITS = 128; + private static final SecureRandom RANDOM = new SecureRandom(); + + private final String keyId; + private final SecretKey key; + private final String endpoint; + + public AesGcmCodec(String keyId, SecretKey key) { + this.keyId = keyId; + this.key = key; + this.endpoint = null; + } + + private AesGcmCodec(String keyId, SecretKey key, String endpoint) { + this.keyId = keyId; + this.key = key; + this.endpoint = endpoint; + } + + @Override + @Nonnull + public PayloadCodec withContext(@Nonnull SerializationContext context) { + if (context instanceof NexusSerializationContext nexusSerializationContext) { + return new AesGcmCodec(keyId, key, nexusSerializationContext.getEndpoint()); + } + return this; + } + + @Override + @Nonnull + public List encode(@Nonnull List payloads) { + String endpointName = + StringUtils.defaultIfBlank(endpoint, NexusEndpointContextPropagator.currentEndpoint()); + List encoded = new ArrayList<>(payloads.size()); + for (Payload payload : payloads) { + Payload.Builder encrypted = + Payload.newBuilder() + .putMetadata( + EncodingKeys.METADATA_ENCODING_KEY, + ByteString.copyFromUtf8(NexusEncoding.AES_GCM.encodingName())) + .putMetadata(KEY_ID_METADATA_KEY, ByteString.copyFromUtf8(keyId)) + .setData(ByteString.copyFrom(encrypt(payload.toByteArray()))); + if (StringUtils.isNotBlank(endpointName)) { + encrypted.putMetadata( + SampleConfig.ENDPOINT_METADATA_KEY, ByteString.copyFromUtf8(endpointName)); + } + encoded.add(encrypted.build()); + } + return encoded; + } + + @Override + @Nonnull + public List decode(@Nonnull List payloads) { + List decoded = new ArrayList<>(payloads.size()); + for (Payload payload : payloads) { + String encoding = + payload + .getMetadataOrDefault(EncodingKeys.METADATA_ENCODING_KEY, ByteString.EMPTY) + .toStringUtf8(); + if (!NexusEncoding.AES_GCM.encodingName().equals(encoding)) { + throw new PayloadCodecException("Expected an AES-GCM payload"); + } + String payloadKeyId = + payload.getMetadataOrDefault(KEY_ID_METADATA_KEY, ByteString.EMPTY).toStringUtf8(); + if (!keyId.equals(payloadKeyId)) { + throw new PayloadCodecException("Unexpected encryption key ID: " + payloadKeyId); + } + try { + decoded.add(Payload.parseFrom(decrypt(payload.getData().toByteArray()))); + } catch (InvalidProtocolBufferException e) { + throw new PayloadCodecException(e); + } + } + return decoded; + } + + private byte[] encrypt(byte[] bytes) { + byte[] nonce = new byte[NONCE_LENGTH]; + RANDOM.nextBytes(nonce); + try { + Cipher cipher = Cipher.getInstance(CIPHER); + cipher.init(Cipher.ENCRYPT_MODE, key, new GCMParameterSpec(TAG_LENGTH_BITS, nonce)); + byte[] ciphertext = cipher.doFinal(bytes); + return ByteBuffer.allocate(nonce.length + ciphertext.length) + .put(nonce) + .put(ciphertext) + .array(); + } catch (GeneralSecurityException e) { + throw new PayloadCodecException(e); + } + } + + private byte[] decrypt(byte[] encrypted) { + if (encrypted.length < NONCE_LENGTH + TAG_LENGTH_BITS / Byte.SIZE) { + throw new PayloadCodecException("AES-GCM payload is too short"); + } + ByteBuffer buffer = ByteBuffer.wrap(encrypted); + byte[] nonce = new byte[NONCE_LENGTH]; + buffer.get(nonce); + byte[] ciphertext = new byte[buffer.remaining()]; + buffer.get(ciphertext); + try { + Cipher cipher = Cipher.getInstance(CIPHER); + cipher.init(Cipher.DECRYPT_MODE, key, new GCMParameterSpec(TAG_LENGTH_BITS, nonce)); + return cipher.doFinal(ciphertext); + } catch (GeneralSecurityException e) { + throw new PayloadCodecException(e); + } + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/NexusCodec.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/NexusCodec.java new file mode 100644 index 000000000..3c6284ea9 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/NexusCodec.java @@ -0,0 +1,55 @@ +package io.temporal.samples.nexusserializationcontext.codec; + +import io.temporal.api.common.v1.Payload; +import io.temporal.payload.codec.PayloadCodec; +import io.temporal.payload.codec.PayloadCodecException; +import io.temporal.payload.context.NexusSerializationContext; +import io.temporal.payload.context.SerializationContext; +import java.util.List; +import java.util.Map; +import javax.annotation.Nonnull; + +/** + * Selects an endpoint codec for synchronous and asynchronous Nexus payloads. The caller's own + * workflow payloads use a separate codec. + */ +public final class NexusCodec implements PayloadCodec { + private final Map codecsByEndpoint; + // The caller workflow has no Nexus endpoint, so Key C encrypts its input and final result. + private final PayloadCodec workflowCodec; + + public NexusCodec(Map codecsByEndpoint, PayloadCodec workflowCodec) { + this.codecsByEndpoint = Map.copyOf(codecsByEndpoint); + this.workflowCodec = workflowCodec; + } + + @Override + @Nonnull + public PayloadCodec withContext(@Nonnull SerializationContext context) { + if (context instanceof NexusSerializationContext nexusContext) { + return codecFor(nexusContext.getEndpoint()).withContext(context); + } + // Caller workflow input and result have no Nexus endpoint; use the caller's Key C codec. + return workflowCodec.withContext(context); + } + + @Override + @Nonnull + public List encode(@Nonnull List payloads) { + return workflowCodec.encode(payloads); + } + + @Override + @Nonnull + public List decode(@Nonnull List payloads) { + return workflowCodec.decode(payloads); + } + + private PayloadCodec codecFor(String endpoint) { + PayloadCodec codec = codecsByEndpoint.get(endpoint); + if (codec == null) { + throw new PayloadCodecException("Unknown Nexus endpoint: " + endpoint); + } + return codec; + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/NexusEncoding.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/NexusEncoding.java new file mode 100644 index 000000000..1a7255892 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/NexusEncoding.java @@ -0,0 +1,16 @@ +package io.temporal.samples.nexusserializationcontext.codec; + +enum NexusEncoding { + AES_GCM("binary/nexus-aes-gcm"), + ZLIB("binary/nexus-zlib"); + + private final String encodingName; + + NexusEncoding(String encodingName) { + this.encodingName = encodingName; + } + + String encodingName() { + return encodingName; + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/ZlibCodec.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/ZlibCodec.java new file mode 100644 index 000000000..f7f350ab4 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/codec/ZlibCodec.java @@ -0,0 +1,76 @@ +package io.temporal.samples.nexusserializationcontext.codec; + +import com.google.protobuf.ByteString; +import com.google.protobuf.InvalidProtocolBufferException; +import io.temporal.api.common.v1.Payload; +import io.temporal.common.converter.EncodingKeys; +import io.temporal.payload.codec.PayloadCodec; +import io.temporal.payload.codec.PayloadCodecException; +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.util.ArrayList; +import java.util.List; +import java.util.zip.DeflaterOutputStream; +import java.util.zip.InflaterInputStream; +import javax.annotation.Nonnull; + +/** Compresses Nexus payloads with zlib. */ +public final class ZlibCodec implements PayloadCodec { + @Override + @Nonnull + public List encode(@Nonnull List payloads) { + List encoded = new ArrayList<>(payloads.size()); + for (Payload payload : payloads) { + encoded.add( + Payload.newBuilder() + .putMetadata( + EncodingKeys.METADATA_ENCODING_KEY, + ByteString.copyFromUtf8(NexusEncoding.ZLIB.encodingName())) + .setData(ByteString.copyFrom(compress(payload.toByteArray()))) + .build()); + } + return encoded; + } + + @Override + @Nonnull + public List decode(@Nonnull List payloads) { + List decoded = new ArrayList<>(payloads.size()); + for (Payload payload : payloads) { + String encoding = + payload + .getMetadataOrDefault(EncodingKeys.METADATA_ENCODING_KEY, ByteString.EMPTY) + .toStringUtf8(); + if (!NexusEncoding.ZLIB.encodingName().equals(encoding)) { + throw new PayloadCodecException("Expected a Nexus zlib payload"); + } + try { + decoded.add(Payload.parseFrom(decompress(payload.getData().toByteArray()))); + } catch (InvalidProtocolBufferException e) { + throw new PayloadCodecException(e); + } + } + return decoded; + } + + private static byte[] compress(byte[] bytes) { + try { + ByteArrayOutputStream output = new ByteArrayOutputStream(); + try (DeflaterOutputStream deflater = new DeflaterOutputStream(output)) { + deflater.write(bytes); + } + return output.toByteArray(); + } catch (IOException e) { + throw new PayloadCodecException(e); + } + } + + private static byte[] decompress(byte[] bytes) { + try (InflaterInputStream inflater = new InflaterInputStream(new ByteArrayInputStream(bytes))) { + return inflater.readAllBytes(); + } catch (IOException e) { + throw new PayloadCodecException(e); + } + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/AsyncEchoServiceImpl.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/AsyncEchoServiceImpl.java new file mode 100644 index 000000000..7d85ce9c9 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/AsyncEchoServiceImpl.java @@ -0,0 +1,26 @@ +package io.temporal.samples.nexusserializationcontext.handler; + +import io.nexusrpc.handler.OperationHandler; +import io.nexusrpc.handler.OperationImpl; +import io.nexusrpc.handler.ServiceImpl; +import io.temporal.client.WorkflowOptions; +import io.temporal.nexus.Nexus; +import io.temporal.nexus.WorkflowRunOperation; +import io.temporal.samples.nexusserializationcontext.service.AsyncEchoService; + +@ServiceImpl(service = AsyncEchoService.class) +public class AsyncEchoServiceImpl { + @OperationImpl + public OperationHandler echoAsync() { + return WorkflowRunOperation.fromWorkflowMethod( + (ctx, details, message) -> + Nexus.getOperationContext() + .getWorkflowClient() + .newWorkflowStub( + EchoWorkflow.class, + WorkflowOptions.newBuilder() + .setWorkflowId("nexus-echo-" + details.getRequestId()) + .build()) + ::echo); + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/CompressedEncryptedHandlerWorker.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/CompressedEncryptedHandlerWorker.java new file mode 100644 index 000000000..9640b0f20 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/CompressedEncryptedHandlerWorker.java @@ -0,0 +1,56 @@ +package io.temporal.samples.nexusserializationcontext.handler; + +import io.temporal.client.WorkflowClient; +import io.temporal.client.WorkflowClientOptions; +import io.temporal.common.converter.CodecDataConverter; +import io.temporal.common.converter.DataConverter; +import io.temporal.common.converter.DefaultDataConverter; +import io.temporal.payload.codec.ChainCodec; +import io.temporal.samples.nexus.options.ClientOptions; +import io.temporal.samples.nexusserializationcontext.SampleConfig; +import io.temporal.samples.nexusserializationcontext.codec.AesGcmCodec; +import io.temporal.samples.nexusserializationcontext.codec.ZlibCodec; +import io.temporal.samples.nexusserializationcontext.propagation.NexusEndpointContextPropagator; +import io.temporal.samples.nexusserializationcontext.propagation.NexusEndpointInterceptor; +import io.temporal.worker.Worker; +import io.temporal.worker.WorkerFactory; +import io.temporal.worker.WorkerFactoryOptions; +import java.nio.charset.StandardCharsets; +import java.util.List; +import javax.crypto.SecretKey; +import javax.crypto.spec.SecretKeySpec; + +public class CompressedEncryptedHandlerWorker { + private static final String TASK_QUEUE = "nexus-serialization-key-a-handler"; + + public static void main(String[] args) { + WorkflowClient client = + ClientOptions.getWorkflowClient( + args, + WorkflowClientOptions.newBuilder() + .setDataConverter(dataConverter()) + .setContextPropagators(List.of(new NexusEndpointContextPropagator()))); + WorkerFactory factory = + WorkerFactory.newInstance( + client, + WorkerFactoryOptions.newBuilder() + .setWorkerInterceptors(new NexusEndpointInterceptor()) + .build()); + Worker worker = factory.newWorker(TASK_QUEUE); + worker.registerWorkflowImplementationTypes(EchoWorkflowImpl.class); + worker.registerNexusServiceImplementation(new SyncEchoServiceImpl()); + worker.registerNexusServiceImplementation(new AsyncEchoServiceImpl()); + factory.start(); + } + + public static DataConverter dataConverter() { + SecretKey keyA = + new SecretKeySpec(SampleConfig.KEY_A_VALUE.getBytes(StandardCharsets.UTF_8), "AES"); + // ChainCodec encodes last to first: compress, then encrypt. + return new CodecDataConverter( + DefaultDataConverter.newDefaultInstance(), + List.of( + new ChainCodec( + List.of(new AesGcmCodec(SampleConfig.KEY_A_ID, keyA), new ZlibCodec())))); + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/EchoWorkflow.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/EchoWorkflow.java new file mode 100644 index 000000000..ce6d8f2ea --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/EchoWorkflow.java @@ -0,0 +1,10 @@ +package io.temporal.samples.nexusserializationcontext.handler; + +import io.temporal.workflow.WorkflowInterface; +import io.temporal.workflow.WorkflowMethod; + +@WorkflowInterface +public interface EchoWorkflow { + @WorkflowMethod + String echo(String message); +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/EchoWorkflowImpl.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/EchoWorkflowImpl.java new file mode 100644 index 000000000..03816609a --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/EchoWorkflowImpl.java @@ -0,0 +1,8 @@ +package io.temporal.samples.nexusserializationcontext.handler; + +public final class EchoWorkflowImpl implements EchoWorkflow { + @Override + public String echo(String message) { + return message; + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/EncryptedHandlerWorker.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/EncryptedHandlerWorker.java new file mode 100644 index 000000000..85525aaca --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/EncryptedHandlerWorker.java @@ -0,0 +1,51 @@ +package io.temporal.samples.nexusserializationcontext.handler; + +import io.temporal.client.WorkflowClient; +import io.temporal.client.WorkflowClientOptions; +import io.temporal.common.converter.CodecDataConverter; +import io.temporal.common.converter.DataConverter; +import io.temporal.common.converter.DefaultDataConverter; +import io.temporal.samples.nexus.options.ClientOptions; +import io.temporal.samples.nexusserializationcontext.SampleConfig; +import io.temporal.samples.nexusserializationcontext.codec.AesGcmCodec; +import io.temporal.samples.nexusserializationcontext.propagation.NexusEndpointContextPropagator; +import io.temporal.samples.nexusserializationcontext.propagation.NexusEndpointInterceptor; +import io.temporal.worker.Worker; +import io.temporal.worker.WorkerFactory; +import io.temporal.worker.WorkerFactoryOptions; +import java.nio.charset.StandardCharsets; +import java.util.List; +import javax.crypto.SecretKey; +import javax.crypto.spec.SecretKeySpec; + +public class EncryptedHandlerWorker { + private static final String TASK_QUEUE = "nexus-serialization-key-b-handler"; + + public static void main(String[] args) { + WorkflowClient client = + ClientOptions.getWorkflowClient( + args, + WorkflowClientOptions.newBuilder() + .setDataConverter(dataConverter()) + .setContextPropagators(List.of(new NexusEndpointContextPropagator()))); + WorkerFactory factory = + WorkerFactory.newInstance( + client, + WorkerFactoryOptions.newBuilder() + .setWorkerInterceptors(new NexusEndpointInterceptor()) + .build()); + Worker worker = factory.newWorker(TASK_QUEUE); + worker.registerWorkflowImplementationTypes(EchoWorkflowImpl.class); + worker.registerNexusServiceImplementation(new SyncEchoServiceImpl()); + worker.registerNexusServiceImplementation(new AsyncEchoServiceImpl()); + factory.start(); + } + + public static DataConverter dataConverter() { + SecretKey keyB = + new SecretKeySpec(SampleConfig.KEY_B_VALUE.getBytes(StandardCharsets.UTF_8), "AES"); + return new CodecDataConverter( + DefaultDataConverter.newDefaultInstance(), + List.of(new AesGcmCodec(SampleConfig.KEY_B_ID, keyB))); + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/SyncEchoServiceImpl.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/SyncEchoServiceImpl.java new file mode 100644 index 000000000..9cb1da932 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/handler/SyncEchoServiceImpl.java @@ -0,0 +1,14 @@ +package io.temporal.samples.nexusserializationcontext.handler; + +import io.nexusrpc.handler.OperationHandler; +import io.nexusrpc.handler.OperationImpl; +import io.nexusrpc.handler.ServiceImpl; +import io.temporal.samples.nexusserializationcontext.service.SyncEchoService; + +@ServiceImpl(service = SyncEchoService.class) +public class SyncEchoServiceImpl { + @OperationImpl + public OperationHandler echo() { + return OperationHandler.sync((ctx, details, message) -> message); + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/propagation/NexusEndpointContextPropagator.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/propagation/NexusEndpointContextPropagator.java new file mode 100644 index 000000000..394275f18 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/propagation/NexusEndpointContextPropagator.java @@ -0,0 +1,51 @@ +package io.temporal.samples.nexusserializationcontext.propagation; + +import io.temporal.api.common.v1.Payload; +import io.temporal.common.context.ContextPropagator; +import io.temporal.common.converter.DataConverter; +import io.temporal.samples.nexusserializationcontext.SampleConfig; +import java.util.Map; +import org.slf4j.MDC; + +/** Carries the Nexus endpoint from a handler into the workflow it starts. */ +public final class NexusEndpointContextPropagator implements ContextPropagator { + private static final DataConverter CONVERTER = DataConverter.getDefaultInstance(); + + public static String currentEndpoint() { + return MDC.get(SampleConfig.ENDPOINT_METADATA_KEY); + } + + @Override + public String getName() { + return NexusEndpointContextPropagator.class.getName(); + } + + @Override + public Object getCurrentContext() { + return currentEndpoint(); + } + + @Override + public void setCurrentContext(Object context) { + if (context == null) { + MDC.remove(SampleConfig.ENDPOINT_METADATA_KEY); + } else { + MDC.put(SampleConfig.ENDPOINT_METADATA_KEY, (String) context); + } + } + + @Override + public Map serializeContext(Object context) { + if (context == null) { + return Map.of(); + } + return Map.of( + SampleConfig.ENDPOINT_METADATA_KEY, CONVERTER.toPayload((String) context).orElseThrow()); + } + + @Override + public Object deserializeContext(Map header) { + Payload payload = header.get(SampleConfig.ENDPOINT_METADATA_KEY); + return payload == null ? null : CONVERTER.fromPayload(payload, String.class, String.class); + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/propagation/NexusEndpointInterceptor.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/propagation/NexusEndpointInterceptor.java new file mode 100644 index 000000000..acf05000f --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/propagation/NexusEndpointInterceptor.java @@ -0,0 +1,36 @@ +package io.temporal.samples.nexusserializationcontext.propagation; + +import io.nexusrpc.OperationException; +import io.nexusrpc.handler.OperationContext; +import io.temporal.common.interceptors.NexusOperationInboundCallsInterceptor; +import io.temporal.common.interceptors.NexusOperationInboundCallsInterceptorBase; +import io.temporal.common.interceptors.WorkerInterceptorBase; +import io.temporal.nexus.Nexus; +import io.temporal.samples.nexusserializationcontext.SampleConfig; +import org.slf4j.MDC; + +/** Makes the endpoint available while a Nexus handler starts its backing workflow. */ +public final class NexusEndpointInterceptor extends WorkerInterceptorBase { + @Override + public NexusOperationInboundCallsInterceptor interceptNexusOperation( + OperationContext context, NexusOperationInboundCallsInterceptor next) { + return new NexusOperationInboundCallsInterceptorBase(next) { + @Override + public StartOperationOutput startOperation(StartOperationInput input) + throws OperationException { + String previousEndpoint = NexusEndpointContextPropagator.currentEndpoint(); + String endpoint = Nexus.getOperationContext().getInfo().getEndpoint(); + try { + MDC.put(SampleConfig.ENDPOINT_METADATA_KEY, endpoint); + return super.startOperation(input); + } finally { + if (previousEndpoint == null) { + MDC.remove(SampleConfig.ENDPOINT_METADATA_KEY); + } else { + MDC.put(SampleConfig.ENDPOINT_METADATA_KEY, previousEndpoint); + } + } + } + }; + } +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/service/AsyncEchoService.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/service/AsyncEchoService.java new file mode 100644 index 000000000..b9bee42d1 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/service/AsyncEchoService.java @@ -0,0 +1,13 @@ +package io.temporal.samples.nexusserializationcontext.service; + +import io.nexusrpc.Operation; +import io.nexusrpc.Service; + +@Service(name = AsyncEchoService.SERVICE_NAME) +public interface AsyncEchoService { + String SERVICE_NAME = "AsyncEchoService"; + String ECHO_ASYNC_OPERATION_NAME = "echoAsync"; + + @Operation(name = ECHO_ASYNC_OPERATION_NAME) + String echoAsync(String message); +} diff --git a/core/src/main/java/io/temporal/samples/nexusserializationcontext/service/SyncEchoService.java b/core/src/main/java/io/temporal/samples/nexusserializationcontext/service/SyncEchoService.java new file mode 100644 index 000000000..d3b947279 --- /dev/null +++ b/core/src/main/java/io/temporal/samples/nexusserializationcontext/service/SyncEchoService.java @@ -0,0 +1,13 @@ +package io.temporal.samples.nexusserializationcontext.service; + +import io.nexusrpc.Operation; +import io.nexusrpc.Service; + +@Service(name = SyncEchoService.SERVICE_NAME) +public interface SyncEchoService { + String SERVICE_NAME = "SyncEchoService"; + String ECHO_OPERATION_NAME = "echo"; + + @Operation(name = ECHO_OPERATION_NAME) + String echo(String message); +} diff --git a/core/src/test/java/io/temporal/samples/nexusserializationcontext/codec/NexusCodecTest.java b/core/src/test/java/io/temporal/samples/nexusserializationcontext/codec/NexusCodecTest.java new file mode 100644 index 000000000..b53299522 --- /dev/null +++ b/core/src/test/java/io/temporal/samples/nexusserializationcontext/codec/NexusCodecTest.java @@ -0,0 +1,263 @@ +package io.temporal.samples.nexusserializationcontext.codec; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import com.google.protobuf.ByteString; +import io.temporal.api.common.v1.Payload; +import io.temporal.common.converter.DataConverter; +import io.temporal.common.converter.EncodingKeys; +import io.temporal.payload.codec.PayloadCodecException; +import io.temporal.payload.context.NexusSerializationContext; +import io.temporal.payload.context.WorkflowSerializationContext; +import io.temporal.samples.nexusserializationcontext.SampleConfig; +import io.temporal.samples.nexusserializationcontext.caller.CallerWorker; +import io.temporal.samples.nexusserializationcontext.caller.EndpointResults; +import io.temporal.samples.nexusserializationcontext.handler.CompressedEncryptedHandlerWorker; +import io.temporal.samples.nexusserializationcontext.handler.EncryptedHandlerWorker; +import io.temporal.samples.nexusserializationcontext.propagation.NexusEndpointContextPropagator; +import io.temporal.samples.nexusserializationcontext.service.AsyncEchoService; +import io.temporal.samples.nexusserializationcontext.service.SyncEchoService; +import org.junit.jupiter.api.Test; + +class NexusCodecTest { + @Test + void encryptsSyncAndAsyncOperationsWithTheKeyForTheirEndpoint() { + DataConverter syncA = converterFor(SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT); + DataConverter asyncA = converterFor(SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT, true); + DataConverter syncB = converterFor(SampleConfig.ENCRYPTED_ENDPOINT); + DataConverter asyncB = converterFor(SampleConfig.ENCRYPTED_ENDPOINT, true); + + Payload syncAPayload = syncA.toPayload("hello").orElseThrow(); + Payload asyncAPayload = asyncA.toPayload("hello").orElseThrow(); + Payload syncBPayload = syncB.toPayload("hello").orElseThrow(); + Payload asyncBPayload = asyncB.toPayload("hello").orElseThrow(); + + assertEncryptedFor( + syncAPayload, SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT, SampleConfig.KEY_A_ID); + assertEncryptedFor( + asyncAPayload, SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT, SampleConfig.KEY_A_ID); + assertEncryptedFor(syncBPayload, SampleConfig.ENCRYPTED_ENDPOINT, SampleConfig.KEY_B_ID); + assertEncryptedFor(asyncBPayload, SampleConfig.ENCRYPTED_ENDPOINT, SampleConfig.KEY_B_ID); + assertEquals("hello", syncA.fromPayload(syncAPayload, String.class, String.class)); + assertEquals("hello", asyncA.fromPayload(asyncAPayload, String.class, String.class)); + assertEquals("hello", syncB.fromPayload(syncBPayload, String.class, String.class)); + assertEquals("hello", asyncB.fromPayload(asyncBPayload, String.class, String.class)); + assertThrows( + PayloadCodecException.class, + () -> syncA.fromPayload(syncBPayload, String.class, String.class)); + assertThrows( + PayloadCodecException.class, + () -> asyncB.fromPayload(asyncAPayload, String.class, String.class)); + Payload keyAPayloadMarkedAsB = + syncAPayload.toBuilder() + .putMetadata( + AesGcmCodec.KEY_ID_METADATA_KEY, ByteString.copyFromUtf8(SampleConfig.KEY_B_ID)) + .build(); + assertThrows( + PayloadCodecException.class, + () -> syncB.fromPayload(keyAPayloadMarkedAsB, String.class, String.class)); + } + + @Test + void compressesBeforeEncryptingForBothOperationsOnTheCompressedEndpoint() { + String message = "repeat me ".repeat(100); + Payload syncCompressed = + converterFor(SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT).toPayload(message).orElseThrow(); + Payload syncEncryptedOnly = + converterFor(SampleConfig.ENCRYPTED_ENDPOINT).toPayload(message).orElseThrow(); + Payload asyncCompressed = + converterFor(SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT, true) + .toPayload(message) + .orElseThrow(); + Payload asyncEncrypted = + converterFor(SampleConfig.ENCRYPTED_ENDPOINT, true).toPayload(message).orElseThrow(); + + assertTrue(syncCompressed.getData().size() < syncEncryptedOnly.getData().size()); + assertTrue(asyncCompressed.getData().size() < asyncEncrypted.getData().size()); + } + + @Test + void encryptsCallerWorkflowInputAndResultWithKeyC() { + DataConverter converter = + CallerWorker.dataConverter() + .withContext( + new WorkflowSerializationContext("nexus-serialization-caller", "caller-id")); + EndpointResults results = new EndpointResults("reply A", "reply B", "reply C", "reply D"); + + Payload input = converter.toPayload("Hello from Nexus").orElseThrow(); + Payload result = converter.toPayload(results).orElseThrow(); + + assertEquals( + NexusEncoding.AES_GCM.encodingName(), + input.getMetadataOrThrow(EncodingKeys.METADATA_ENCODING_KEY).toStringUtf8()); + assertEquals( + SampleConfig.KEY_C_ID, + input.getMetadataOrThrow(AesGcmCodec.KEY_ID_METADATA_KEY).toStringUtf8()); + assertEquals( + SampleConfig.KEY_C_ID, + result.getMetadataOrThrow(AesGcmCodec.KEY_ID_METADATA_KEY).toStringUtf8()); + assertFalse(input.getMetadataMap().containsKey(SampleConfig.ENDPOINT_METADATA_KEY)); + assertFalse(input.getData().toStringUtf8().contains("Hello from Nexus")); + assertFalse(result.getData().toStringUtf8().contains("reply A")); + assertEquals("Hello from Nexus", converter.fromPayload(input, String.class, String.class)); + assertEquals( + results, converter.fromPayload(result, EndpointResults.class, EndpointResults.class)); + } + + @Test + void rejectsUnknownEndpoints() { + DataConverter converter = + CallerWorker.dataConverter() + .withContext( + new NexusSerializationContext( + "unknown", SyncEchoService.SERVICE_NAME, SyncEchoService.ECHO_OPERATION_NAME)); + + assertThrows(PayloadCodecException.class, () -> converter.toPayload("hello")); + } + + @Test + void synchronousHandlersUseTheirOwnCodecWithoutEndpointContext() { + DataConverter callerA = converterFor(SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT); + DataConverter callerB = converterFor(SampleConfig.ENCRYPTED_ENDPOINT); + DataConverter handlerA = CompressedEncryptedHandlerWorker.dataConverter(); + DataConverter handlerB = EncryptedHandlerWorker.dataConverter(); + + Payload requestA = callerA.toPayload("hello").orElseThrow(); + Payload requestB = callerB.toPayload("hello").orElseThrow(); + assertEquals("hello", handlerA.fromPayload(requestA, String.class, String.class)); + assertEquals("hello", handlerB.fromPayload(requestB, String.class, String.class)); + + Payload resultA = handlerA.toPayload("reply").orElseThrow(); + Payload resultB = handlerB.toPayload("reply").orElseThrow(); + assertEquals("reply", callerA.fromPayload(resultA, String.class, String.class)); + assertEquals("reply", callerB.fromPayload(resultB, String.class, String.class)); + + Payload contextualResultA = + handlerA + .withContext(contextFor(SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT)) + .toPayload("reply") + .orElseThrow(); + Payload contextualResultB = + handlerB + .withContext(contextFor(SampleConfig.ENCRYPTED_ENDPOINT)) + .toPayload("reply") + .orElseThrow(); + assertEquals( + SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT, + contextualResultA.getMetadataOrThrow(SampleConfig.ENDPOINT_METADATA_KEY).toStringUtf8()); + assertEquals( + SampleConfig.ENCRYPTED_ENDPOINT, + contextualResultB.getMetadataOrThrow(SampleConfig.ENDPOINT_METADATA_KEY).toStringUtf8()); + + assertThrows( + PayloadCodecException.class, + () -> handlerA.fromPayload(requestB, String.class, String.class)); + assertThrows( + PayloadCodecException.class, + () -> handlerB.fromPayload(requestA, String.class, String.class)); + } + + @Test + void handlersDecodeAsyncRequestsWithTheirEndpointKey() { + DataConverter handlerA = + CompressedEncryptedHandlerWorker.dataConverter() + .withContext(contextFor(SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT, true)); + DataConverter handlerB = + EncryptedHandlerWorker.dataConverter() + .withContext(contextFor(SampleConfig.ENCRYPTED_ENDPOINT, true)); + Payload requestA = + converterFor(SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT, true) + .toPayload("hello") + .orElseThrow(); + Payload requestB = + converterFor(SampleConfig.ENCRYPTED_ENDPOINT, true).toPayload("hello").orElseThrow(); + + assertEquals("hello", handlerA.fromPayload(requestA, String.class, String.class)); + assertEquals("hello", handlerB.fromPayload(requestB, String.class, String.class)); + assertThrows( + PayloadCodecException.class, + () -> handlerA.fromPayload(requestB, String.class, String.class)); + } + + @Test + void propagatedEndpointAppearsOnEncryptedAsyncWorkflowResult() { + NexusEndpointContextPropagator propagator = new NexusEndpointContextPropagator(); + try { + propagator.setCurrentContext(null); + DataConverter handlerA = + CompressedEncryptedHandlerWorker.dataConverter() + .withContext(new WorkflowSerializationContext("handler-a", "echo-workflow")); + DataConverter handlerB = + EncryptedHandlerWorker.dataConverter() + .withContext(new WorkflowSerializationContext("handler-b", "echo-workflow")); + + Payload resultA = + asyncResultFor( + propagator, + handlerA, + SampleConfig.COMPRESSED_ENCRYPTED_ENDPOINT, + SampleConfig.KEY_A_ID); + Payload resultB = + asyncResultFor( + propagator, handlerB, SampleConfig.ENCRYPTED_ENDPOINT, SampleConfig.KEY_B_ID); + assertTrue(resultA.getData().size() < resultB.getData().size()); + } finally { + propagator.setCurrentContext(null); + } + } + + private static Payload asyncResultFor( + NexusEndpointContextPropagator propagator, + DataConverter handler, + String endpoint, + String keyId) { + propagator.setCurrentContext(endpoint); + var header = propagator.serializeContext(propagator.getCurrentContext()); + propagator.setCurrentContext(null); + propagator.setCurrentContext(propagator.deserializeContext(header)); + + String message = "reply ".repeat(100); + Payload result = handler.toPayload(message).orElseThrow(); + assertEquals( + endpoint, result.getMetadataOrThrow(SampleConfig.ENDPOINT_METADATA_KEY).toStringUtf8()); + assertEquals(keyId, result.getMetadataOrThrow(AesGcmCodec.KEY_ID_METADATA_KEY).toStringUtf8()); + assertEquals( + message, converterFor(endpoint, true).fromPayload(result, String.class, String.class)); + assertEquals(message, handler.fromPayload(result, String.class, String.class)); + propagator.setCurrentContext(null); + return result; + } + + private static DataConverter converterFor(String endpoint) { + return CallerWorker.dataConverter().withContext(contextFor(endpoint)); + } + + private static DataConverter converterFor(String endpoint, boolean async) { + return CallerWorker.dataConverter().withContext(contextFor(endpoint, async)); + } + + private static NexusSerializationContext contextFor(String endpoint) { + return contextFor(endpoint, false); + } + + private static NexusSerializationContext contextFor(String endpoint, boolean async) { + if (async) { + return new NexusSerializationContext( + endpoint, AsyncEchoService.SERVICE_NAME, AsyncEchoService.ECHO_ASYNC_OPERATION_NAME); + } + return new NexusSerializationContext( + endpoint, SyncEchoService.SERVICE_NAME, SyncEchoService.ECHO_OPERATION_NAME); + } + + private static void assertEncryptedFor(Payload payload, String endpoint, String keyId) { + assertEquals( + NexusEncoding.AES_GCM.encodingName(), + payload.getMetadataOrThrow(EncodingKeys.METADATA_ENCODING_KEY).toStringUtf8()); + assertEquals(keyId, payload.getMetadataOrThrow(AesGcmCodec.KEY_ID_METADATA_KEY).toStringUtf8()); + assertEquals( + endpoint, payload.getMetadataOrThrow(SampleConfig.ENDPOINT_METADATA_KEY).toStringUtf8()); + } +}