diff --git a/src/ModelContextProtocol.AspNetCore/AuthorizationFilterSetup.cs b/src/ModelContextProtocol.AspNetCore/AuthorizationFilterSetup.cs index 5b371d9a1..6290dd5c6 100644 --- a/src/ModelContextProtocol.AspNetCore/AuthorizationFilterSetup.cs +++ b/src/ModelContextProtocol.AspNetCore/AuthorizationFilterSetup.cs @@ -20,6 +20,7 @@ internal sealed class AuthorizationFilterSetup( public void Configure(McpServerOptions options) { ConfigureListToolsFilter(options); + ConfigureOrdinaryCallToolFilter(options); ConfigureListResourcesFilter(options); ConfigureListResourceTemplatesFilter(options); @@ -85,6 +86,22 @@ private static void CheckListToolsFilter(McpServerOptions options) }); } + private void ConfigureOrdinaryCallToolFilter(McpServerOptions options) + { + options.Filters.Request.CallToolFilters.Add(next => async (context, cancellationToken) => + { + var authResult = await GetAuthorizationResultAsync(context.User, context.MatchedPrimitive, context.Services, context); + if (!authResult.Succeeded) + { + throw new McpProtocolException("Access forbidden: This tool requires authorization.", McpErrorCode.InvalidRequest); + } + + context.Items[AuthorizationFilterInvokedKey] = true; + + return await next(context, cancellationToken); + }); + } + private void ConfigureCallToolFilter(McpServerOptions options) { #pragma warning disable MCPEXP002 // Authorization must run in the alternate-result pipeline before task dispatch. diff --git a/src/ModelContextProtocol.AspNetCore/HttpMcpServerBuilderExtensions.cs b/src/ModelContextProtocol.AspNetCore/HttpMcpServerBuilderExtensions.cs index 801da5f6e..a52268341 100644 --- a/src/ModelContextProtocol.AspNetCore/HttpMcpServerBuilderExtensions.cs +++ b/src/ModelContextProtocol.AspNetCore/HttpMcpServerBuilderExtensions.cs @@ -58,6 +58,9 @@ public static IMcpServerBuilder WithHttpTransport(this IMcpServerBuilder builder /// authorization attributes such as /// and . Tool authorization runs in the alternate-result pipeline before /// the Tasks extension dispatches background execution, so an unauthorized tool call does not create a task. + /// Each call to this method also adds an ordinary call-tool authorization checkpoint at that point in the filter + /// pipeline. Call this method again after any call-tool filter that changes the matched tool or user to authorize + /// the replacement using the updated context. /// public static IMcpServerBuilder AddAuthorizationFilters(this IMcpServerBuilder builder) { diff --git a/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs b/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs index 9c46b0a2c..207ffebd7 100644 --- a/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs +++ b/src/ModelContextProtocol.Extensions.Tasks/Server/McpTasksBuilderExtensions.cs @@ -23,7 +23,6 @@ public static class McpTasksBuilderExtensions /// Tasks are implemented as an alternate-result call-tool filter. Alternate-result filters registered before /// the Tasks filter run before task creation. Filters registered after it, along with all ordinary call-tool /// filters, run in the background before the tool. - /// Register Tasks before configuring ordinary call-tool filters. /// /// The server builder. /// The task store. @@ -79,13 +78,6 @@ public void Configure(McpServerOptions options) options.RequestHandlers.Add(new McpServerRequestHandler { Method = TasksProtocol.MethodTasksUpdate, Handler = HandleUpdateTask }); options.RequestHandlers.Add(new McpServerRequestHandler { Method = TasksProtocol.MethodTasksCancel, Handler = HandleCancelTask }); - if (options.Filters.Request.CallToolFilters.Count > 0) - { - throw new InvalidOperationException( - $"{nameof(WithTasks)} must be configured before ordinary call-tool filters because " + - "the Tasks filter must execute outside the ordinary call-tool pipeline."); - } - // Use a filter rather than a handler so it wraps around Core's tool dispatch. // This ensures it intercepts tool calls BEFORE the tool is invoked, allowing // it to spawn background execution and return the task alternate immediately. diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/AuthorizeAttributeTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/AuthorizeAttributeTests.cs index 76a7201d8..6ad643ef5 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/AuthorizeAttributeTests.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/AuthorizeAttributeTests.cs @@ -93,6 +93,91 @@ public async Task Authorize_Tool_AllowsAuthenticatedUser() Assert.Equal("Authorized: test", content.Text); } + [Fact] + public async Task Authorize_Tool_ReauthorizesPrimitiveChangedByInterveningFilter() + { + await using var app = await StartServerWithAuth(builder => + { + builder.WithTools(); + builder.Services.Configure(options => + { + if (options.ToolCollection is null || + !options.ToolCollection.TryGetPrimitive("authorized_tool", out var authorizedTool)) + { + throw new InvalidOperationException("The replacement tool was not registered."); + } + + options.Filters.Request.CallToolFilters.Add(next => async (context, cancellationToken) => + { + context.MatchedPrimitive = authorizedTool; + return await next(context, cancellationToken); + }); + }); + builder.AddAuthorizationFilters(); + }); + + var client = await ConnectAsync(); + + var exception = await Assert.ThrowsAsync(async () => + await client.CallToolAsync( + "anonymous_tool", + new Dictionary { ["message"] = "test" }, + cancellationToken: TestContext.Current.CancellationToken)); + + Assert.Equal("Request failed (remote): Access forbidden: This tool requires authorization.", exception.Message); + Assert.Equal(McpErrorCode.InvalidRequest, exception.ErrorCode); + } + + [Fact] + public async Task Authorize_Tool_ReauthorizesAfterEachInterveningFilter() + { + await using var app = await StartServerWithAuth(builder => + { + builder.WithTools(); + builder.Services.Configure(options => + { + if (options.ToolCollection is null || + !options.ToolCollection.TryGetPrimitive("authorized_tool", out var authorizedTool)) + { + throw new InvalidOperationException("The first replacement tool was not registered."); + } + + options.Filters.Request.CallToolFilters.Add(next => async (context, cancellationToken) => + { + context.MatchedPrimitive = authorizedTool; + return await next(context, cancellationToken); + }); + }); + builder.AddAuthorizationFilters(); + builder.Services.Configure(options => + { + if (options.ToolCollection is null || + !options.ToolCollection.TryGetPrimitive("admin_tool", out var adminTool)) + { + throw new InvalidOperationException("The second replacement tool was not registered."); + } + + options.Filters.Request.CallToolFilters.Add(next => async (context, cancellationToken) => + { + context.MatchedPrimitive = adminTool; + return await next(context, cancellationToken); + }); + }); + builder.AddAuthorizationFilters(); + }, "TestUser"); + + var client = await ConnectAsync(); + + var exception = await Assert.ThrowsAsync(async () => + await client.CallToolAsync( + "anonymous_tool", + new Dictionary { ["message"] = "test" }, + cancellationToken: TestContext.Current.CancellationToken)); + + Assert.Equal("Request failed (remote): Access forbidden: This tool requires authorization.", exception.Message); + Assert.Equal(McpErrorCode.InvalidRequest, exception.ErrorCode); + } + [Fact] public async Task AuthorizeWithRoles_Tool_RequiresAdminRole() { diff --git a/tests/ModelContextProtocol.AspNetCore.Tests/HttpTaskIntegrationTests.cs b/tests/ModelContextProtocol.AspNetCore.Tests/HttpTaskIntegrationTests.cs index e7b2a5578..d64a55987 100644 --- a/tests/ModelContextProtocol.AspNetCore.Tests/HttpTaskIntegrationTests.cs +++ b/tests/ModelContextProtocol.AspNetCore.Tests/HttpTaskIntegrationTests.cs @@ -1,7 +1,6 @@ using Microsoft.AspNetCore.Builder; using Microsoft.AspNetCore.Authorization; using Microsoft.Extensions.DependencyInjection; -using Microsoft.Extensions.Options; using ModelContextProtocol.AspNetCore.Tests.Utils; using ModelContextProtocol.Client; using ModelContextProtocol.Extensions.Tasks; @@ -44,23 +43,41 @@ public async Task WithTasks_CanCallToolOverHttp() } [Fact] - public async Task WithTasks_AfterOrdinaryFilter_ThrowsActionableError() + public async Task WithTasks_AfterOrdinaryFilter_RunsFilter() { + var filterInvocationCount = 0; Builder.Services .AddMcpServer(options => { - options.Filters.Request.CallToolFilters.Add(next => next); + options.Filters.Request.CallToolFilters.Add(next => async (context, cancellationToken) => + { + Interlocked.Increment(ref filterInvocationCount); + return await next(context, cancellationToken); + }); }) .WithHttpTransport() .WithTasks(new InMemoryMcpTaskStore { DefaultPollIntervalMs = 10 }) .WithTools(); await using var app = Builder.Build(); + app.MapMcp(); + await app.StartAsync(TestContext.Current.CancellationToken); + + await using var transport = new HttpClientTransport( + new HttpClientTransportOptions { Endpoint = new("http://localhost:5000") }, + HttpClient, + LoggerFactory); + await using var client = await McpClient.CreateAsync( + transport, + loggerFactory: LoggerFactory, + cancellationToken: TestContext.Current.CancellationToken); - var exception = Assert.Throws( - () => app.Services.GetRequiredService>().Value); - Assert.Contains(nameof(McpTasksBuilderExtensions.WithTasks), exception.Message); - Assert.Contains("before ordinary call-tool filters", exception.Message); + var result = await client.CallToolWithPollingAsync( + new CallToolRequestParams { Name = "test" }, + cancellationToken: TestContext.Current.CancellationToken); + + Assert.Equal("Hello World!", Assert.IsType(Assert.Single(result.Content)).Text); + Assert.Equal(1, filterInvocationCount); } [Theory] @@ -165,6 +182,54 @@ public async Task WithTasks_UnauthorizedTool_DoesNotCreateTask(bool registerTask Times.Never); } + [Fact] + public async Task WithTasks_ReauthorizesToolChangedByOrdinaryFilter() + { + var serverBuilder = Builder.Services + .AddMcpServer() + .WithHttpTransport() + .WithTasks(new InMemoryMcpTaskStore { DefaultPollIntervalMs = 10 }) + .WithTools() + .AddAuthorizationFilters(); + + serverBuilder.Services.Configure(options => + { + if (options.ToolCollection is null || + !options.ToolCollection.TryGetPrimitive("authorized-test", out var authorizedTool)) + { + throw new InvalidOperationException("The replacement tool was not registered."); + } + + options.Filters.Request.CallToolFilters.Add(next => async (context, cancellationToken) => + { + context.MatchedPrimitive = authorizedTool; + return await next(context, cancellationToken); + }); + }); + serverBuilder.AddAuthorizationFilters(); + Builder.Services.AddAuthorization(); + + await using var app = Builder.Build(); + app.MapMcp(); + await app.StartAsync(TestContext.Current.CancellationToken); + + await using var transport = new HttpClientTransport( + new HttpClientTransportOptions { Endpoint = new("http://localhost:5000") }, + HttpClient, + LoggerFactory); + await using var client = await McpClient.CreateAsync( + transport, + loggerFactory: LoggerFactory, + cancellationToken: TestContext.Current.CancellationToken); + + var exception = await Assert.ThrowsAsync(() => + client.CallToolWithPollingAsync( + new CallToolRequestParams { Name = "test" }, + cancellationToken: TestContext.Current.CancellationToken).AsTask()); + + Assert.Contains("Access forbidden: This tool requires authorization.", exception.Message); + } + [McpServerToolType] private sealed class TestTools {