diff --git a/src/modules/Elsa.Tenants.AspNetCore/Middleware/TenantResolutionMiddleware.cs b/src/modules/Elsa.Tenants.AspNetCore/Middleware/TenantResolutionMiddleware.cs index c4d35194f..b8babc087 100644 --- a/src/modules/Elsa.Tenants.AspNetCore/Middleware/TenantResolutionMiddleware.cs +++ b/src/modules/Elsa.Tenants.AspNetCore/Middleware/TenantResolutionMiddleware.cs @@ -37,7 +37,14 @@ public class TenantResolutionMiddleware(RequestDelegate next, ITenantScopeFactor await using var tenantScope = tenantScopeFactory.CreateScope(tenant); var originalServiceProvider = context.RequestServices; context.RequestServices = tenantScope.ServiceProvider; - await next(context); - context.RequestServices = originalServiceProvider; + + try + { + await next(context); + } + finally + { + context.RequestServices = originalServiceProvider; + } } -} \ No newline at end of file +} diff --git a/test/unit/Elsa.Tenants.UnitTests/Elsa.Tenants.UnitTests.csproj b/test/unit/Elsa.Tenants.UnitTests/Elsa.Tenants.UnitTests.csproj index 057369bbd..0a2af3fe0 100644 --- a/test/unit/Elsa.Tenants.UnitTests/Elsa.Tenants.UnitTests.csproj +++ b/test/unit/Elsa.Tenants.UnitTests/Elsa.Tenants.UnitTests.csproj @@ -7,6 +7,7 @@ + diff --git a/test/unit/Elsa.Tenants.UnitTests/Middleware/TenantResolutionMiddlewareTests.cs b/test/unit/Elsa.Tenants.UnitTests/Middleware/TenantResolutionMiddlewareTests.cs new file mode 100644 index 000000000..a1813367d --- /dev/null +++ b/test/unit/Elsa.Tenants.UnitTests/Middleware/TenantResolutionMiddlewareTests.cs @@ -0,0 +1,43 @@ +using Elsa.Common.Multitenancy; +using Elsa.Tenants.AspNetCore.Middleware; +using Microsoft.AspNetCore.Http; +using Microsoft.Extensions.DependencyInjection; +using NSubstitute; + +namespace Elsa.Tenants.UnitTests.Middleware; + +public class TenantResolutionMiddlewareTests +{ + [Fact] + public async Task InvokeAsync_WhenNextThrows_RestoresOriginalRequestServices() + { + await using var rootProvider = new ServiceCollection() + .AddScoped(_ => new ScopedProbe()) + .BuildServiceProvider(); + await using var originalRequestScope = rootProvider.CreateAsyncScope(); + var originalRequestServices = originalRequestScope.ServiceProvider; + var context = new DefaultHttpContext { RequestServices = originalRequestServices }; + var expectedException = new InvalidOperationException("Downstream failure"); + var tenantScopeFactory = new DefaultTenantScopeFactory( + new DefaultTenantAccessor(), + rootProvider.GetRequiredService()); + var middleware = new TenantResolutionMiddleware( + _ => Task.FromException(expectedException), + tenantScopeFactory); + var tenantResolverPipelineInvoker = Substitute.For(); + tenantResolverPipelineInvoker + .InvokePipelineAsync(Arg.Any()) + .Returns(Task.FromResult(null)); + + var exception = await Assert.ThrowsAsync( + () => middleware.InvokeAsync(context, tenantResolverPipelineInvoker)); + + Assert.Same(expectedException, exception); + Assert.Same(originalRequestServices, context.RequestServices); + Assert.NotNull(context.RequestServices.GetRequiredService()); + } + + private sealed class ScopedProbe + { + } +}