From d92b82f127e1ff2f368286b14761f2d79c93af8f Mon Sep 17 00:00:00 2001 From: Max Brooks <45081361+MaxBrooks114@users.noreply.github.com> Date: Mon, 28 Apr 2025 13:05:05 -0400 Subject: [PATCH] Refactor ActivityRegistry to populate all activities in workflow editor page (#6605) - Modify ListAll to return distinct activity descriptors. - Update RegisterAsync to call Add with correct parameters. - Refactor RefreshDescriptorsAsync for better collection usage. - Split Add method into two overloads for clarity. - Improve logging for replacing existing activity descriptors. Co-authored-by: Max Brooks --- .../Services/ActivityRegistry.cs | 262 +++++++++--------- 1 file changed, 132 insertions(+), 130 deletions(-) diff --git a/src/modules/Elsa.Workflows.Core/Services/ActivityRegistry.cs b/src/modules/Elsa.Workflows.Core/Services/ActivityRegistry.cs index 66b94a110..b0bceac32 100644 --- a/src/modules/Elsa.Workflows.Core/Services/ActivityRegistry.cs +++ b/src/modules/Elsa.Workflows.Core/Services/ActivityRegistry.cs @@ -1,138 +1,140 @@ -using System.Collections.Concurrent; -using System.Diagnostics.CodeAnalysis; -using Elsa.Workflows.Helpers; -using Elsa.Workflows.Models; -using Microsoft.Extensions.Logging; - -namespace Elsa.Workflows; - -/// -public class ActivityRegistry(IActivityDescriber activityDescriber, IEnumerable modifiers, ILogger logger) : IActivityRegistry -{ - private readonly ISet _manualActivityDescriptors = new HashSet(); - private ConcurrentDictionary> _providedActivityDescriptors = new(); - private ConcurrentDictionary<(string Type, int Version), ActivityDescriptor> _activityDescriptors = new(); - - /// - public void Add(Type providerType, ActivityDescriptor descriptor) => Add(descriptor, GetOrCreateDescriptors(providerType)); - - /// - public void Remove(Type providerType, ActivityDescriptor descriptor) - { - _providedActivityDescriptors[providerType].Remove(descriptor); - _activityDescriptors.Remove((descriptor.TypeName, descriptor.Version), out _); - } - - /// - public IEnumerable ListAll() => _activityDescriptors.Values; - - /// - public IEnumerable ListByProvider(Type providerType) => _providedActivityDescriptors.TryGetValue(providerType, out var descriptors) ? descriptors : ArraySegment.Empty; - - /// - public ActivityDescriptor? Find(string type) => _activityDescriptors.Values.Where(x => x.TypeName == type).MaxBy(x => x.Version); - - /// - public ActivityDescriptor? Find(string type, int version) => _activityDescriptors.TryGetValue((type, version), out var descriptor) ? descriptor : null; - - /// - public ActivityDescriptor? Find(Func predicate) => _activityDescriptors.Values.FirstOrDefault(predicate); - - /// - public IEnumerable FindMany(Func predicate) => _activityDescriptors.Values.Where(predicate); - - /// - public void Register(ActivityDescriptor descriptor) - { - Add(GetType(), descriptor); - } - - /// - public async Task RegisterAsync([DynamicallyAccessedMembers(DynamicallyAccessedMemberTypes.PublicProperties)] Type activityType, CancellationToken cancellationToken) - { - var activityTypeName = ActivityTypeNameHelper.GenerateTypeName(activityType); - - if (_activityDescriptors.Values.Any(x => x.TypeName == activityTypeName)) - return; - - var activityDescriptor = await activityDescriber.DescribeActivityAsync(activityType, cancellationToken); - Add(GetType(), activityDescriptor); - _manualActivityDescriptors.Add(activityDescriptor); - } - - /// - public async Task RegisterAsync(IEnumerable activityTypes, CancellationToken cancellationToken = default) - { - foreach (var activityType in activityTypes) - await RegisterAsync(activityType, cancellationToken); +using System.Collections.Concurrent; +using System.Diagnostics.CodeAnalysis; +using Elsa.Workflows.Helpers; +using Elsa.Workflows.Models; +using Microsoft.Extensions.Logging; + +namespace Elsa.Workflows; + +/// +public class ActivityRegistry(IActivityDescriber activityDescriber, IEnumerable modifiers, ILogger logger) : IActivityRegistry +{ + private readonly ISet _manualActivityDescriptors = new HashSet(); + private ConcurrentDictionary> _providedActivityDescriptors = new(); + private ConcurrentDictionary<(string Type, int Version), ActivityDescriptor> _activityDescriptors = new(); + + /// + public void Add(Type providerType, ActivityDescriptor descriptor) => Add(descriptor, GetOrCreateDescriptors(providerType)); + + /// + public void Remove(Type providerType, ActivityDescriptor descriptor) + { + _providedActivityDescriptors[providerType].Remove(descriptor); + _activityDescriptors.Remove((descriptor.TypeName, descriptor.Version), out _); } - /// - public ValueTask> GetDescriptorsAsync(CancellationToken cancellationToken = default) => new(_manualActivityDescriptors); - - /// - public async Task RefreshDescriptorsAsync(IEnumerable activityProviders, CancellationToken cancellationToken = default) - { - var providersDictionary = new ConcurrentDictionary>(); - var activityDescriptors = new ConcurrentDictionary<(string Type, int Version), ActivityDescriptor>(); - foreach (var activityProvider in activityProviders) - { - var descriptors = (await activityProvider.GetDescriptorsAsync(cancellationToken)).ToList(); - var providerDescriptors = new List(); - providersDictionary[activityProvider.GetType()] = providerDescriptors; - foreach (var descriptor in descriptors) - { - Add(descriptor, activityDescriptors, providerDescriptors); - } + /// + public IEnumerable ListAll() => _activityDescriptors.Values.DistinctBy(x => x.TypeName); + + /// + public IEnumerable ListByProvider(Type providerType) => _providedActivityDescriptors.TryGetValue(providerType, out var descriptors) ? descriptors : ArraySegment.Empty; + + /// + public ActivityDescriptor? Find(string type) => _activityDescriptors.Values.Where(x => x.TypeName == type).MaxBy(x => x.Version); + + /// + public ActivityDescriptor? Find(string type, int version) => _activityDescriptors.TryGetValue((type, version), out var descriptor) ? descriptor : null; + + /// + public ActivityDescriptor? Find(Func predicate) => _activityDescriptors.Values.FirstOrDefault(predicate); + + /// + public IEnumerable FindMany(Func predicate) => _activityDescriptors.Values.Where(predicate); + + /// + public void Register(ActivityDescriptor descriptor) + { + Add(GetType(), descriptor); + } + + /// + public async Task RegisterAsync([DynamicallyAccessedMembers(DynamicallyAccessedMemberTypes.PublicProperties)] Type activityType, CancellationToken cancellationToken) + { + var activityTypeName = ActivityTypeNameHelper.GenerateTypeName(activityType); + + if (_activityDescriptors.Values.Any(x => x.TypeName == activityTypeName)) + return; + + var activityDescriptor = await activityDescriber.DescribeActivityAsync(activityType, cancellationToken); + + + Add(activityDescriptor, _activityDescriptors, _manualActivityDescriptors); + _manualActivityDescriptors.Add(activityDescriptor); + } + + /// + public async Task RegisterAsync(IEnumerable activityTypes, CancellationToken cancellationToken = default) + { + foreach (var activityType in activityTypes) + await RegisterAsync(activityType, cancellationToken); + } + + /// + public ValueTask> GetDescriptorsAsync(CancellationToken cancellationToken = default) => new(_manualActivityDescriptors); + + /// + public async Task RefreshDescriptorsAsync(IEnumerable activityProviders, CancellationToken cancellationToken = default) + { + var providersDictionary = new ConcurrentDictionary>(); + var activityDescriptors = new ConcurrentDictionary<(string Type, int Version), ActivityDescriptor>(_activityDescriptors); + foreach (var activityProvider in activityProviders) + { + var descriptors = (await activityProvider.GetDescriptorsAsync(cancellationToken)).ToList(); + var providerDescriptors = new List(); + providersDictionary[activityProvider.GetType()] = providerDescriptors; + foreach (var descriptor in descriptors) + { + Add(descriptor, activityDescriptors, providerDescriptors); + } } - Interlocked.Exchange(ref _activityDescriptors, activityDescriptors); - Interlocked.Exchange(ref _providedActivityDescriptors, providersDictionary); + Interlocked.Exchange(ref _activityDescriptors, activityDescriptors); + Interlocked.Exchange(ref _providedActivityDescriptors, providersDictionary); } - public async Task RefreshDescriptorsAsync(IActivityProvider activityProvider, CancellationToken cancellationToken = default) - { - var providersDictionary = new ConcurrentDictionary>(_providedActivityDescriptors); - var activityDescriptors = new ConcurrentDictionary<(string Type, int Version), ActivityDescriptor>(_activityDescriptors); - var descriptors = (await activityProvider.GetDescriptorsAsync(cancellationToken)).ToList(); - var providerDescriptors = new List(); + public async Task RefreshDescriptorsAsync(IActivityProvider activityProvider, CancellationToken cancellationToken = default) + { + var providersDictionary = new ConcurrentDictionary>(_providedActivityDescriptors); + var activityDescriptors = new ConcurrentDictionary<(string Type, int Version), ActivityDescriptor>(_activityDescriptors); + var descriptors = (await activityProvider.GetDescriptorsAsync(cancellationToken)).ToList(); + var providerDescriptors = new List(); providersDictionary[activityProvider.GetType()] = providerDescriptors; foreach (var descriptor in descriptors) Add(descriptor, activityDescriptors, providerDescriptors); - Interlocked.Exchange(ref _activityDescriptors, activityDescriptors); - Interlocked.Exchange(ref _providedActivityDescriptors, providersDictionary); + Interlocked.Exchange(ref _activityDescriptors, activityDescriptors); + Interlocked.Exchange(ref _providedActivityDescriptors, providersDictionary); } - private void Add(ActivityDescriptor descriptor, ICollection target) - { - Add(descriptor, _activityDescriptors, target); - } - - private void Add(ActivityDescriptor? descriptor, ConcurrentDictionary<(string Type, int Version), ActivityDescriptor> activityDescriptors, ICollection providerDescriptors) - { - if (descriptor is null) - { - logger.LogError("Unable to add a null descriptor"); - return; + private void Add(ActivityDescriptor descriptor, ICollection target) + { + Add(descriptor, _activityDescriptors, target); + } + + private void Add(ActivityDescriptor? descriptor, ConcurrentDictionary<(string Type, int Version), ActivityDescriptor> activityDescriptors, ICollection providerDescriptors) + { + if (descriptor is null) + { + logger.LogError("Unable to add a null descriptor"); + return; } - foreach (var modifier in modifiers) - modifier.Modify(descriptor); - + foreach (var modifier in modifiers) + modifier.Modify(descriptor); + // If the descriptor already exists, replace it. But log a warning. - if (activityDescriptors.TryGetValue((descriptor.TypeName, descriptor.Version), out var existingDescriptor)) - { + if (activityDescriptors.TryGetValue((descriptor.TypeName, descriptor.Version), out var existingDescriptor)) + { // Remove the existing descriptor from the providerDescriptors collection. - providerDescriptors.Remove(existingDescriptor); - + providerDescriptors.Remove(existingDescriptor); + // Log a warning. - logger.LogWarning("Activity descriptor {ActivityType} v{ActivityVersion} was already registered. Replacing with new descriptor", descriptor.TypeName, descriptor.Version); - } - - activityDescriptors[(descriptor.TypeName, descriptor.Version)] = descriptor; - providerDescriptors.Add(descriptor); + logger.LogWarning("Activity descriptor {ActivityType} v{ActivityVersion} was already registered. Replacing with new descriptor", descriptor.TypeName, descriptor.Version); + } + + activityDescriptors[(descriptor.TypeName, descriptor.Version)] = descriptor; + providerDescriptors.Add(descriptor); } /// @@ -151,16 +153,16 @@ public class ActivityRegistry(IActivityDescriber activityDescriber, IEnumerable< _activityDescriptors.Remove((descriptor.TypeName, descriptor.Version), out _); _providedActivityDescriptors.Remove(providerType, out _); - } - - private ICollection GetOrCreateDescriptors(Type provider) - { - if (_providedActivityDescriptors.TryGetValue(provider, out var descriptors)) - return descriptors; - - descriptors = new List(); - _providedActivityDescriptors[provider] = descriptors; - - return descriptors; - } -} \ No newline at end of file + } + + private ICollection GetOrCreateDescriptors(Type provider) + { + if (_providedActivityDescriptors.TryGetValue(provider, out var descriptors)) + return descriptors; + + descriptors = new List(); + _providedActivityDescriptors[provider] = descriptors; + + return descriptors; + } +}