diff --git a/src/modules/Elsa.Common/Multitenancy/Implementations/DefaultTenantService.cs b/src/modules/Elsa.Common/Multitenancy/Implementations/DefaultTenantService.cs index 5dc7f0c05..95a336aa9 100644 --- a/src/modules/Elsa.Common/Multitenancy/Implementations/DefaultTenantService.cs +++ b/src/modules/Elsa.Common/Multitenancy/Implementations/DefaultTenantService.cs @@ -7,14 +7,13 @@ public class DefaultTenantService(IServiceScopeFactory scopeFactory, ITenantScop { private readonly AsyncServiceScope _serviceScope = scopeFactory.CreateAsyncScope(); private readonly SemaphoreSlim _initializationLock = new(1, 1); - private readonly SemaphoreSlim _refreshLock = new(1, 1); + private readonly SemaphoreSlim _tenantMutationLock = new(1, 1); private IDictionary? _tenantsDictionary; private IDictionary? _tenantScopesDictionary; public async ValueTask DisposeAsync() { await _serviceScope.DisposeAsync(); - _initializationLock.Dispose(); } public async Task FindAsync(string id, CancellationToken cancellationToken = default) @@ -60,22 +59,33 @@ public class DefaultTenantService(IServiceScopeFactory scopeFactory, ITenantScop public async Task DeactivateTenantsAsync(CancellationToken cancellationToken = default) { - var dictionary = await GetTenantsDictionaryAsync(cancellationToken); - var tenants = dictionary.Values.ToArray(); + var dictionary = await GetTenantsDictionaryForMutationAsync(cancellationToken); + await _tenantMutationLock.WaitAsync(cancellationToken); - foreach (var tenant in tenants) - await UnregisterTenantAsync(tenant, false, cancellationToken); + try + { + var tenants = dictionary.Values.ToArray(); + + foreach (var tenant in tenants) + { + await UnregisterTenantAsync(tenant, false, cancellationToken); + } + } + finally + { + _tenantMutationLock.Release(); + } } public async Task RefreshAsync(CancellationToken cancellationToken = default) { - await _refreshLock.WaitAsync(cancellationToken); + var currentTenants = await GetTenantsDictionaryForMutationAsync(cancellationToken); + await _tenantMutationLock.WaitAsync(cancellationToken); try { await using var scope = scopeFactory.CreateAsyncScope(); var tenantsProvider = scope.ServiceProvider.GetRequiredService(); - var currentTenants = await GetTenantsDictionaryAsync(cancellationToken); var currentTenantIds = currentTenants.Keys; var tenantsFromProvider = (await tenantsProvider.ListAsync(cancellationToken)).ToList(); var newTenants = tenantsFromProvider.Count == 0 @@ -99,10 +109,22 @@ public class DefaultTenantService(IServiceScopeFactory scopeFactory, ITenantScop } finally { - _refreshLock.Release(); + _tenantMutationLock.Release(); } } + private async Task> GetTenantsDictionaryForMutationAsync(CancellationToken cancellationToken) + { + var dictionary = await GetTenantsDictionaryAsync(cancellationToken); + + // The dictionary is published before initialization completes so lifecycle event handlers can read it. + // Wait for any concurrent initializer before allowing a mutation to proceed. + await _initializationLock.WaitAsync(cancellationToken); + _initializationLock.Release(); + + return dictionary; + } + private async Task> GetTenantsDictionaryAsync(CancellationToken cancellationToken) { if (_tenantsDictionary == null) @@ -157,4 +179,4 @@ public class DefaultTenantService(IServiceScopeFactory scopeFactory, ITenantScop } } } -} \ No newline at end of file +} diff --git a/test/unit/Elsa.Common.UnitTests/Multitenancy/DefaultTenantServiceTests.cs b/test/unit/Elsa.Common.UnitTests/Multitenancy/DefaultTenantServiceTests.cs index 89685a0d5..13566e714 100644 --- a/test/unit/Elsa.Common.UnitTests/Multitenancy/DefaultTenantServiceTests.cs +++ b/test/unit/Elsa.Common.UnitTests/Multitenancy/DefaultTenantServiceTests.cs @@ -119,6 +119,147 @@ public class DefaultTenantServiceTests } } + [Fact] + public async Task DeactivateTenantsAsync_WhenRefreshIsInProgress_WaitsForRefreshBeforeDeactivating() + { + var tenant = new Tenant { Id = "tenant-1", Name = "Tenant 1" }; + var timeout = TimeSpan.FromSeconds(5); + var refreshStarted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var completeRefresh = new TaskCompletionSource>(TaskCreationOptions.RunContinuationsAsynchronously); + var listRequests = 0; + var tenantsProvider = Substitute.For(); + tenantsProvider.ListAsync(Arg.Any()).Returns(async _ => + { + if (Interlocked.Increment(ref listRequests) == 1) + { + return new[] { tenant }; + } + + refreshStarted.TrySetResult(); + return await completeRefresh.Task; + }); + + var (tenantService, serviceProvider) = await CreateTenantServiceAsync(tenantsProvider); + Task? refreshTask = null; + Task? deactivateTask = null; + + try + { + await tenantService.ListAsync(); + + refreshTask = tenantService.RefreshAsync(); + await refreshStarted.Task.WaitAsync(timeout); + + deactivateTask = tenantService.DeactivateTenantsAsync(); + + Assert.False(deactivateTask.IsCompleted); + + completeRefresh.SetResult([tenant]); + await refreshTask.WaitAsync(timeout); + await deactivateTask.WaitAsync(timeout); + + Assert.Empty(await tenantService.ListAsync()); + } + finally + { + completeRefresh.TrySetResult([tenant]); + + if (refreshTask != null) + { + await refreshTask.WaitAsync(timeout); + } + + if (deactivateTask != null) + { + await deactivateTask.WaitAsync(timeout); + } + + if (tenantService is IAsyncDisposable disposable) + { + await disposable.DisposeAsync(); + } + + await serviceProvider.DisposeAsync(); + } + } + + [Fact] + public async Task DeactivateTenantsAsync_WhenInitializationIsInProgress_WaitsForInitializationBeforeDeactivating() + { + var tenant = new Tenant { Id = "tenant-1", Name = "Tenant 1" }; + var timeout = TimeSpan.FromSeconds(5); + var initializationStarted = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var completeInitialization = new TaskCompletionSource>(TaskCreationOptions.RunContinuationsAsynchronously); + var tenantsProvider = Substitute.For(); + tenantsProvider.ListAsync(Arg.Any()).Returns(async _ => + { + initializationStarted.TrySetResult(); + return await completeInitialization.Task; + }); + + var (tenantService, serviceProvider) = await CreateTenantServiceAsync(tenantsProvider); + Task>? initializationTask = null; + Task? deactivateTask = null; + + try + { + initializationTask = tenantService.ListAsync(); + await initializationStarted.Task.WaitAsync(timeout); + + deactivateTask = tenantService.DeactivateTenantsAsync(); + + Assert.False(deactivateTask.IsCompleted); + + completeInitialization.SetResult([tenant]); + await initializationTask.WaitAsync(timeout); + await deactivateTask.WaitAsync(timeout); + + Assert.Empty(await tenantService.ListAsync()); + } + finally + { + completeInitialization.TrySetResult([tenant]); + + if (initializationTask != null) + { + await initializationTask.WaitAsync(timeout); + } + + if (deactivateTask != null) + { + await deactivateTask.WaitAsync(timeout); + } + + if (tenantService is IAsyncDisposable disposable) + { + await disposable.DisposeAsync(); + } + + await serviceProvider.DisposeAsync(); + } + } + + [Fact] + public async Task DeactivateTenantsAsync_AfterDisposeAsync_DoesNotUseDisposedSynchronizationPrimitives() + { + var tenant = new Tenant { Id = "tenant-1", Name = "Tenant 1" }; + var (tenantService, serviceProvider) = await CreateTenantServiceAsync([tenant]); + + try + { + await tenantService.ListAsync(); + await ((IAsyncDisposable)tenantService).DisposeAsync(); + + await tenantService.DeactivateTenantsAsync(); + + Assert.Empty(await tenantService.ListAsync()); + } + finally + { + await serviceProvider.DisposeAsync(); + } + } + private static Task<(ITenantService TenantService, ServiceProvider ServiceProvider)> CreateTenantServiceAsync(IEnumerable tenants, Func>? tenantsFactory = null) { var tenantList = tenants.ToList(); @@ -127,6 +268,11 @@ public class DefaultTenantServiceTests var tenantsProvider = Substitute.For(); tenantsProvider.ListAsync(Arg.Any()).Returns(_ => getTenants()); + return CreateTenantServiceAsync(tenantsProvider); + } + + private static Task<(ITenantService TenantService, ServiceProvider ServiceProvider)> CreateTenantServiceAsync(ITenantsProvider tenantsProvider) + { var services = new ServiceCollection(); services.AddSingleton(_ => tenantsProvider); services.AddSingleton();