Skip to content
Open
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
2 changes: 1 addition & 1 deletion src/DurableTask.AzureStorage/MessageManager.cs
Original file line number Diff line number Diff line change
Expand Up @@ -119,7 +119,7 @@ public async Task<string> SerializeMessageDataAsync(MessageData messageData, Can
return Utils.SerializeToJson(serializer, wrapperMessageData);
}

return Utils.SerializeToJson(serializer, messageData);
return rawContent;
}

/// <summary>
Expand Down
208 changes: 200 additions & 8 deletions test/DurableTask.AzureStorage.Tests/MessageManagerTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -14,15 +14,20 @@
namespace DurableTask.AzureStorage.Tests
{
using DurableTask.AzureStorage.Storage;
using DurableTask.Core;
using DurableTask.Core.History;
using Microsoft.VisualStudio.TestTools.UnitTesting;
using Newtonsoft.Json;
using System;
using System.Collections.Generic;
using System.Text;
using System.Threading.Tasks;

[TestClass]
public class MessageManagerTests
{
const int MaxStorageQueuePayloadSizeInBytes = 45 * 1024;

[DataTestMethod]
[DataRow("System.Collections.Generic.Dictionary`2[[System.String, System.Private.CoreLib],[System.String, System.Private.CoreLib]]")]
[DataRow("System.Collections.Generic.Dictionary`2[[System.String, mscorlib],[System.String, mscorlib]]")]
Expand Down Expand Up @@ -69,6 +74,102 @@ public void DeserializesCustomTypes()
Assert.AreEqual("tagValue", startedEvent.Tags["tag1"]);
}

[DataTestMethod]
[DataRow(false)]
[DataRow(true)]
public async Task SerializeMessageDataAsync_InlinePayloadMatchesPreviousWirePayload(bool useDataContractSerialization)
{
var binder = new CustomNameTypeBinder();
var settings = new AzureStorageOrchestrationServiceSettings
{
CustomMessageTypeBinder = binder,
UseDataContractSerialization = useDataContractSerialization,
};
var manager = SetupMessageManager(settings, "$root");
MessageData message = CreateMessageData("inline payload");

string actualPayload = await manager.SerializeMessageDataAsync(message);
string previousWirePayload = SerializeUsingMessageSettings(
new AzureStorageOrchestrationServiceSettings
{
CustomMessageTypeBinder = new CustomNameTypeBinder(),
UseDataContractSerialization = useDataContractSerialization,
},
message);

Assert.AreEqual(previousWirePayload, actualPayload);
Assert.AreEqual(Encoding.UTF8.GetByteCount(previousWirePayload), message.TotalMessageSizeBytes);
Assert.AreEqual(MessageFormatFlags.InlineJson, manager.GetMessageFormatFlags(message));
Assert.IsNull(message.CompressedBlobName);
Assert.AreEqual(1, binder.MessageDataSerializationCount);
StringAssert.Contains(actualPayload, CustomNameTypeBinder.CustomAssemblyName);

MessageData roundTrippedMessage = manager.DeserializeMessageData(actualPayload);
Assert.AreEqual(message.ActivityId, roundTrippedMessage.ActivityId);
Assert.AreEqual(
((GenericEvent)message.TaskMessage.Event).Data,
((GenericEvent)roundTrippedMessage.TaskMessage.Event).Data);
}

[DataTestMethod]
[DataRow(false, 0)]
[DataRow(false, 1)]
[DataRow(true, 0)]
[DataRow(true, 1)]
public async Task SerializeMessageDataAsync_UsesExpectedStoragePathAtThreshold(
bool useDataContractSerialization,
int bytesOverThreshold)
{
var settings = new AzureStorageOrchestrationServiceSettings
{
UseDataContractSerialization = useDataContractSerialization,
};
string containerName = $"message-manager-{Guid.NewGuid():N}";
var manager = SetupMessageManager(settings, containerName);
int targetSize = MaxStorageQueuePayloadSizeInBytes + bytesOverThreshold;
MessageData message = CreateMessageDataWithSerializedSize(settings, targetSize);
string fullMessagePayload = SerializeUsingMessageSettings(settings, message);

try
{
string queuePayload = await manager.SerializeMessageDataAsync(message);
MessageFormatFlags expectedFormat = bytesOverThreshold == 0
? MessageFormatFlags.InlineJson
: MessageFormatFlags.StorageBlob;

Assert.AreEqual(targetSize, message.TotalMessageSizeBytes);
Assert.AreEqual(expectedFormat, manager.GetMessageFormatFlags(message));

if (expectedFormat == MessageFormatFlags.InlineJson)
{
Assert.AreEqual(fullMessagePayload, queuePayload);
Assert.IsNull(message.CompressedBlobName);
}
else
{
string expectedWrapperPayload = SerializeUsingMessageSettings(
settings,
new MessageData { CompressedBlobName = message.CompressedBlobName });
Assert.AreEqual(expectedWrapperPayload, queuePayload);

MessageData wrapper = manager.DeserializeMessageData(queuePayload);
Assert.IsNull(wrapper.TaskMessage);
Assert.IsFalse(string.IsNullOrWhiteSpace(wrapper.CompressedBlobName));
Assert.AreEqual(message.CompressedBlobName, wrapper.CompressedBlobName);

string storedPayload = await manager.DownloadAndDecompressAsBytesAsync(wrapper.CompressedBlobName);
Assert.AreEqual(fullMessagePayload, storedPayload);
}
}
finally
{
if (bytesOverThreshold > 0)
{
await manager.DeleteContainerAsync();
}
}
}

[DataTestMethod]
[DataRow("blob.bin", "blob.bin")]
[DataRow("@#$%!", "%40%23%24%25%21")]
Expand All @@ -94,17 +195,108 @@ private string GetMessage(string dictionaryType)

private MessageManager SetupMessageManager(ICustomTypeBinder binder)
{
var azureStorageClient = new AzureStorageClient(
new AzureStorageOrchestrationServiceSettings
{
StorageAccountClientProvider = new StorageAccountClientProvider("UseDevelopmentStorage=true"),
});

return new MessageManager(
return SetupMessageManager(
new AzureStorageOrchestrationServiceSettings { CustomMessageTypeBinder = binder },
azureStorageClient,
"$root");
}

private static MessageManager SetupMessageManager(
AzureStorageOrchestrationServiceSettings settings,
string containerName)
{
settings.StorageAccountClientProvider =
new StorageAccountClientProvider(TestHelpers.GetTestStorageAccountConnectionString());
var azureStorageClient = new AzureStorageClient(settings);
return new MessageManager(settings, azureStorageClient, containerName);
}

private static MessageData CreateMessageData(string payload)
{
var orchestrationInstance = new OrchestrationInstance
{
InstanceId = "message-manager-instance",
ExecutionId = "message-manager-execution",
};
var taskMessage = new TaskMessage
{
Event = new GenericEvent(1, payload),
OrchestrationInstance = orchestrationInstance,
SequenceNumber = 42,
};

return new MessageData(
taskMessage,
Guid.Parse("55f31df8-9abb-4d86-a197-78fd0908efcf"),
"message-manager-queue",
orchestrationEpisode: 3,
sender: orchestrationInstance)
{
SequenceNumber = 43,
SerializableTraceContext = "trace-context",
};
}

private static MessageData CreateMessageDataWithSerializedSize(
AzureStorageOrchestrationServiceSettings settings,
int targetSize)
{
MessageData message = CreateMessageData(string.Empty);
int emptyPayloadSize = Encoding.UTF8.GetByteCount(SerializeUsingMessageSettings(settings, message));
Assert.IsTrue(targetSize >= emptyPayloadSize);

((GenericEvent)message.TaskMessage.Event).Data = new string('x', targetSize - emptyPayloadSize);
int actualSize = Encoding.UTF8.GetByteCount(SerializeUsingMessageSettings(settings, message));
Assert.AreEqual(targetSize, actualSize);
return message;
}

private static string SerializeUsingMessageSettings(
AzureStorageOrchestrationServiceSettings settings,
MessageData message)
{
var serializerSettings = new JsonSerializerSettings
{
TypeNameHandling = TypeNameHandling.Objects,
SerializationBinder = new TypeNameSerializationBinder(settings.CustomMessageTypeBinder),
};

if (settings.UseDataContractSerialization)
{
serializerSettings.Converters.Add(new DataContractJsonConverter());
}

return Utils.SerializeToJson(JsonSerializer.Create(serializerSettings), message);
}
}

internal class CustomNameTypeBinder : ICustomTypeBinder
{
public const string CustomAssemblyName = "MessageManagerTests.CustomAssembly";
readonly Dictionary<string, Type> serializedTypes = new Dictionary<string, Type>();

public int MessageDataSerializationCount { get; private set; }

public void BindToName(Type serializedType, out string assemblyName, out string typeName)
{
assemblyName = CustomAssemblyName;
typeName = serializedType.FullName!;
this.serializedTypes[typeName] = serializedType;

if (serializedType == typeof(MessageData))
{
this.MessageDataSerializationCount++;
}
}

public Type BindToType(string assemblyName, string typeName)
{
if (this.serializedTypes.TryGetValue(typeName, out Type? serializedType))
{
return serializedType;
}

throw new JsonSerializationException($"Unknown serialized type '{typeName}'.");
}
}

internal class KnownTypeBinder : ICustomTypeBinder
Expand Down
Loading