diff --git a/test/PatternKit.Generators.Tests/SagaGeneratorTests.cs b/test/PatternKit.Generators.Tests/SagaGeneratorTests.cs index e126b15f..bd630b7f 100644 --- a/test/PatternKit.Generators.Tests/SagaGeneratorTests.cs +++ b/test/PatternKit.Generators.Tests/SagaGeneratorTests.cs @@ -99,6 +99,57 @@ private static ValueTask StartAsync(OrderState state, Message message, MessageContext context) + => state with { Started = true }; + + [SagaStep(typeof(Paid), 20)] + private static ValueTask PayAsync(OrderState state, Message message, MessageContext context, CancellationToken cancellationToken) + => ValueTask.FromResult(state with { Paid = true }); + + [SagaCompleteWhen] + private static bool IsComplete(OrderState state) => state.Started && state.Paid; + } + """; + + var comp = CreateCompilation(source, nameof(GeneratesSagaFactoriesForGlobalStructHostWithSyncAndAsyncSteps)); + var gen = new SagaGenerator(); + _ = RoslynTestHelpers.Run(comp, gen, out var run, out var updated); + + ScenarioExpect.All(run.Results, result => ScenarioExpect.Empty(result.Diagnostics)); + + var generated = ScenarioExpect.Single(run.Results.SelectMany(result => result.GeneratedSources)); + var text = generated.SourceText.ToString(); + ScenarioExpect.Equal("OrderSaga.Saga.g.cs", generated.HintName); + ScenarioExpect.DoesNotContain("namespace ", text); + ScenarioExpect.Contains("partial struct OrderSaga", text); + ScenarioExpect.Contains("BuildSync()", text); + ScenarioExpect.Contains("BuildAsync()", text); + ScenarioExpect.Contains(".On().Then(Start)", text); + ScenarioExpect.Contains(".On().Then(PayAsync)", text); + ScenarioExpect.Equal(2, CountOccurrences(text, ".CompleteWhen(IsComplete)")); + + var emit = updated.Emit(Stream.Null); + ScenarioExpect.True(emit.Success, string.Join("\n", emit.Diagnostics)); + } + [Scenario("ReportsDiagnosticForNonPartialSaga")] [Fact] public void ReportsDiagnosticForNonPartialSaga() @@ -177,6 +228,38 @@ public static partial class OrderSaga ScenarioExpect.Equal("PKSG003", ScenarioExpect.Single(run.Results.SelectMany(result => result.Diagnostics)).Id); } + [Scenario("ReportsDiagnosticForInvalidSagaStepShapes")] + [Fact] + public void ReportsDiagnosticForInvalidSagaStepShapes() + { + var source = """ + using PatternKit.Generators.Messaging; + using PatternKit.Messaging; + + namespace MyApp; + + public sealed record OrderState(bool Started); + public sealed record Started(string OrderId); + + [GenerateSaga(typeof(OrderState))] + public static partial class OrderSaga + { + [SagaStep(typeof(Started), 10)] + private static OrderState MissingContext(OrderState state, Message message) => state; + + [SagaStep(typeof(Started), 20)] + private OrderState InstanceStep(OrderState state, Message message, MessageContext context) => state; + } + """; + + var comp = CreateCompilation(source, nameof(ReportsDiagnosticForInvalidSagaStepShapes)); + var gen = new SagaGenerator(); + _ = RoslynTestHelpers.Run(comp, gen, out var run, out _); + + var diagnostics = run.Results.SelectMany(result => result.Diagnostics).ToArray(); + ScenarioExpect.Equal(2, diagnostics.Count(diagnostic => diagnostic.Id == "PKSG003")); + } + [Scenario("ReportsDiagnosticForInvalidCompletionSignature")] [Fact] public void ReportsDiagnosticForInvalidCompletionSignature() @@ -213,4 +296,17 @@ private static CSharpCompilation CreateCompilation(string source, string assembl source, assemblyName, extra: MetadataReference.CreateFromFile(typeof(PatternKit.Messaging.Message<>).Assembly.Location)); + + private static int CountOccurrences(string value, string match) + { + var count = 0; + var index = 0; + while ((index = value.IndexOf(match, index, StringComparison.Ordinal)) >= 0) + { + count++; + index += match.Length; + } + + return count; + } }