Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -19,25 +19,25 @@

package io.temporal.samples.encryptedpayloads;

import com.google.common.base.Defaults;
import com.google.protobuf.ByteString;
import io.temporal.api.common.v1.Payload;
import io.temporal.api.common.v1.Payloads;
import io.temporal.common.converter.DataConverter;
import io.temporal.common.converter.DataConverterException;
import java.lang.reflect.Type;
import io.temporal.common.converter.EncodingKeys;
import io.temporal.payload.codec.PayloadCodec;
import io.temporal.payload.codec.PayloadCodecException;
import java.nio.ByteBuffer;
import java.nio.charset.Charset;
import java.nio.charset.StandardCharsets;
import java.security.SecureRandom;
import java.util.Optional;
import java.util.List;
import java.util.stream.Collectors;
import javax.crypto.Cipher;
import javax.crypto.SecretKey;
import javax.crypto.spec.GCMParameterSpec;
import javax.crypto.spec.SecretKeySpec;
import org.jetbrains.annotations.NotNull;

public class CryptDataConverter implements DataConverter {
static final String METADATA_ENCODING_KEY = "encoding";
class CryptCodec implements PayloadCodec {
static final ByteString METADATA_ENCODING =
ByteString.copyFrom("binary/encrypted", StandardCharsets.UTF_8);

Expand All @@ -53,10 +53,61 @@ public class CryptDataConverter implements DataConverter {
private static final int GCM_TAG_LENGTH_BIT = 128;
private static final Charset UTF_8 = StandardCharsets.UTF_8;

private final DataConverter converter;
@NotNull
@Override
public List<Payload> encode(@NotNull List<Payload> payloads) {
return payloads.stream().map(this::encodePayload).collect(Collectors.toList());
}

@NotNull
@Override
public List<Payload> decode(@NotNull List<Payload> payloads) {
return payloads.stream().map(this::decodePayload).collect(Collectors.toList());
}

private Payload encodePayload(Payload payload) {
String keyId = getKeyId();
SecretKey key = getKey(keyId);

byte[] encryptedData;
try {
encryptedData = encrypt(payload.toByteArray(), key);
} catch (Throwable e) {
throw new DataConverterException(e);
}

public CryptDataConverter(DataConverter converter) {
this.converter = converter;
return Payload.newBuilder()
.putMetadata(EncodingKeys.METADATA_ENCODING_KEY, METADATA_ENCODING)
.putMetadata(METADATA_ENCRYPTION_CIPHER_KEY, METADATA_ENCRYPTION_CIPHER)
.putMetadata(METADATA_ENCRYPTION_KEY_ID_KEY, ByteString.copyFromUtf8(keyId))
.setData(ByteString.copyFrom(encryptedData))
.build();
}

private Payload decodePayload(Payload payload) {
if (METADATA_ENCODING.equals(
payload.getMetadataOrDefault(EncodingKeys.METADATA_ENCODING_KEY, null))) {
String keyId;
try {
keyId = payload.getMetadataOrThrow(METADATA_ENCRYPTION_KEY_ID_KEY).toString(UTF_8);
} catch (Exception e) {
throw new PayloadCodecException(e);
}
SecretKey key = getKey(keyId);

byte[] plainData;
Payload decryptedPayload;

try {
plainData = decrypt(payload.getData().toByteArray(), key);
decryptedPayload = Payload.parseFrom(plainData);
return decryptedPayload;
} catch (Throwable e) {
throw new PayloadCodecException(e);
}
} else {
return payload;
}
}

private String getKeyId() {
Expand Down Expand Up @@ -106,105 +157,4 @@ private byte[] decrypt(byte[] encryptedDataWithNonce, SecretKey key) throws Exce

return cipher.doFinal(encryptedData);
}

@Override
public <T> Optional<Payload> toPayload(T value) throws DataConverterException {
return converter.toPayload(value);
}

public <T> Optional<Payload> toEncryptedPayload(T value) throws DataConverterException {
Optional<Payload> optionalPayload = converter.toPayload(value);

if (!optionalPayload.isPresent()) {
return optionalPayload;
}

Payload innerPayload = optionalPayload.get();

String keyId = getKeyId();
SecretKey key = getKey(keyId);

byte[] encryptedData;
try {
encryptedData = encrypt(innerPayload.toByteArray(), key);
} catch (Throwable e) {
throw new DataConverterException(e);
}

Payload encryptedPayload =
Payload.newBuilder()
.putMetadata(METADATA_ENCODING_KEY, METADATA_ENCODING)
.putMetadata(METADATA_ENCRYPTION_CIPHER_KEY, METADATA_ENCRYPTION_CIPHER)
.putMetadata(METADATA_ENCRYPTION_KEY_ID_KEY, ByteString.copyFromUtf8(keyId))
.setData(ByteString.copyFrom(encryptedData))
.build();

return Optional.of(encryptedPayload);
}

@Override
public <T> T fromPayload(Payload payload, Class<T> valueClass, Type valueType) {
ByteString encoding = payload.getMetadataOrDefault(METADATA_ENCODING_KEY, null);
if (!encoding.equals(METADATA_ENCODING)) {
return converter.fromPayload(payload, valueClass, valueType);
}

String keyId;
try {
keyId = payload.getMetadataOrThrow(METADATA_ENCRYPTION_KEY_ID_KEY).toString(UTF_8);
} catch (Exception e) {
throw new DataConverterException(payload, valueClass, e);
}
SecretKey key = getKey(keyId);

byte[] plainData;
Payload decryptedPayload;

try {
plainData = decrypt(payload.getData().toByteArray(), key);
decryptedPayload = Payload.parseFrom(plainData);
} catch (Throwable e) {
throw new DataConverterException(e);
}

return converter.fromPayload(decryptedPayload, valueClass, valueType);
}

@Override
public Optional<Payloads> toPayloads(Object... values) throws DataConverterException {
if (values == null || values.length == 0) {
return Optional.empty();
}
try {
Payloads.Builder result = Payloads.newBuilder();
for (Object value : values) {
Optional<Payload> payload = toEncryptedPayload(value);
if (payload.isPresent()) {
result.addPayloads(payload.get());
} else {
result.addPayloads(Payload.getDefaultInstance());
}
}
return Optional.of(result.build());
} catch (DataConverterException e) {
throw e;
} catch (Throwable e) {
throw new DataConverterException(e);
}
}

@Override
public <T> T fromPayloads(
int index, Optional<Payloads> content, Class<T> parameterType, Type genericParameterType)
throws DataConverterException {
if (!content.isPresent()) {
return (T) Defaults.defaultValue((Class<?>) parameterType);
}
int count = content.get().getPayloadsCount();
// To make adding arguments a backwards compatible change
if (index >= count) {
return (T) Defaults.defaultValue((Class<?>) parameterType);
}
return fromPayload(content.get().getPayloads(index), parameterType, genericParameterType);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -25,14 +25,16 @@
import io.temporal.client.WorkflowClient;
import io.temporal.client.WorkflowClientOptions;
import io.temporal.client.WorkflowOptions;
import io.temporal.common.converter.DataConverter;
import io.temporal.common.converter.CodecDataConverter;
import io.temporal.common.converter.DefaultDataConverter;
import io.temporal.serviceclient.WorkflowServiceStubs;
import io.temporal.worker.Worker;
import io.temporal.worker.WorkerFactory;
import io.temporal.workflow.Workflow;
import io.temporal.workflow.WorkflowInterface;
import io.temporal.workflow.WorkflowMethod;
import java.time.Duration;
import java.util.Collections;

/**
* Hello World Temporal workflow that executes a single activity. Requires a local instance the
Expand Down Expand Up @@ -91,7 +93,10 @@ public static void main(String[] args) {
WorkflowClient.newInstance(
service,
WorkflowClientOptions.newBuilder()
.setDataConverter(new CryptDataConverter(DataConverter.getDefaultInstance()))
.setDataConverter(
new CodecDataConverter(
DefaultDataConverter.newDefaultInstance(),
Collections.singletonList(new CryptCodec())))
.build());

// worker factory that can be used to create workers for specific task queues
Expand Down