Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 9 additions & 3 deletions api-docs/openapi/v3_0/mcpV1alpha1Api.json
Original file line number Diff line number Diff line change
Expand Up @@ -287,8 +287,10 @@
"allowedTools" : {
"uniqueItems" : true,
"type" : "array",
"description" : "Allowed MCP tool names. Use '*' to automatically allow all current and future tools.",
"items" : {
"type" : "string"
"type" : "string",
"description" : "Allowed MCP tool names. Use '*' to automatically allow all current and future tools."
}
},
"displayName" : {
Expand Down Expand Up @@ -349,8 +351,10 @@
"allowedTools" : {
"uniqueItems" : true,
"type" : "array",
"description" : "Allowed MCP tool names. Use '*' to automatically allow all current and future tools.",
"items" : {
"type" : "string"
"type" : "string",
"description" : "Allowed MCP tool names. Use '*' to automatically allow all current and future tools."
}
},
"creationTimestamp" : {
Expand Down Expand Up @@ -616,8 +620,10 @@
"allowedTools" : {
"uniqueItems" : true,
"type" : "array",
"description" : "Allowed MCP tool names. Use '*' to automatically allow all current and future tools.",
"items" : {
"type" : "string"
"type" : "string",
"description" : "Allowed MCP tool names. Use '*' to automatically allow all current and future tools."
}
},
"displayName" : {
Expand Down
16 changes: 9 additions & 7 deletions src/main/java/run/halo/mcpserver/AuthorizedMcpTransport.java
Original file line number Diff line number Diff line change
Expand Up @@ -69,10 +69,10 @@ public Mono<Void> handleNotification(
}

private Mono<McpSchema.JSONRPCResponse> listTools(McpSchema.JSONRPCRequest request) {
return authorization.allowedTools()
return authorization.authentication()
.zipWith(catalog.protocolTools())
.map(tuple -> tuple.getT2().stream()
.filter(tool -> tuple.getT1().contains(tool.name()))
.filter(tool -> tuple.getT1().allows(tool.name()))
.toList())
.map(tools -> McpSchema.JSONRPCResponse.result(
request.id(), McpSchema.ListToolsResult.builder(tools).build()))
Expand All @@ -94,14 +94,16 @@ private Mono<McpSchema.JSONRPCResponse> callTool(
() -> Mono.just(protocolError(request, -32602, "Invalid tools/call parameters")));
}
var toolName = call.name() == null ? "" : call.name();
return recentCallHistory.observe(
var rateLimitTool = authentication.allowsAllTools()
? catalog.hasProtocolTool(toolName)
.map(available -> available ? toolName : "<unauthorized>")
: Mono.just(authentication.allows(toolName) ? toolName : "<unauthorized>");
Comment thread
ruibaby marked this conversation as resolved.
return rateLimitTool.flatMap(name -> recentCallHistory.observe(
authentication,
toolName,
() -> rateLimiter.allowTool(
authentication.keyId(),
authentication.allows(toolName) ? toolName : "<unauthorized>")
() -> rateLimiter.allowTool(authentication.keyId(), name)
? executeTool(context, request, handler, call, toolName)
: Mono.just(rateLimited(request)));
: Mono.just(rateLimited(request))));
});
}

Expand Down
4 changes: 4 additions & 0 deletions src/main/java/run/halo/mcpserver/McpAccessKey.java
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,9 @@
singular = "mcpaccesskey")
public class McpAccessKey extends AbstractExtension {

static final String ALLOWED_TOOLS_DESCRIPTION =
"Allowed MCP tool names. Use '*' to automatically allow all current and future tools.";

@Schema(requiredMode = Schema.RequiredMode.REQUIRED)
private Spec spec = new Spec();

Expand All @@ -33,6 +36,7 @@ public static class Spec {
private String ownerName;
private boolean enabled = true;
private Instant expiresAt;
@Schema(description = ALLOWED_TOOLS_DESCRIPTION)
private Set<String> allowedTools = new LinkedHashSet<>();
@Schema(description = "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.")
private Set<String> allowedIpRanges = new LinkedHashSet<>();
Expand Down
16 changes: 13 additions & 3 deletions src/main/java/run/halo/mcpserver/McpAccessKeyEndpoint.java
Original file line number Diff line number Diff line change
Expand Up @@ -255,6 +255,7 @@ private Mono<Void> validateTools(Set<String> requestedTools, Set<String> allowed
var requested = tools(requestedTools);
return toolCatalog.availableNames().flatMap(available -> {
var unknown = new LinkedHashSet<>(requested);
unknown.remove(McpKeyAuthenticationToken.ALL_TOOLS);
unknown.removeAll(available);
unknown.removeAll(allowedUnavailableTools);
if (!unknown.isEmpty()) {
Expand Down Expand Up @@ -302,7 +303,10 @@ public GroupVersion groupVersion() {
@Schema(name = "CreateMcpAccessKeyRequest")
record CreateKeyRequest(
@Schema(requiredMode = Schema.RequiredMode.REQUIRED) String displayName,
@Schema(requiredMode = Schema.RequiredMode.REQUIRED) Set<String> allowedTools,
@Schema(
requiredMode = Schema.RequiredMode.REQUIRED,
description = McpAccessKey.ALLOWED_TOOLS_DESCRIPTION)
Set<String> allowedTools,
@Schema(
requiredMode = Schema.RequiredMode.REQUIRED,
description = "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.")
Expand All @@ -312,7 +316,10 @@ record CreateKeyRequest(
@Schema(name = "UpdateMcpAccessKeyRequest")
record UpdateKeyRequest(
@Schema(requiredMode = Schema.RequiredMode.REQUIRED) String displayName,
@Schema(requiredMode = Schema.RequiredMode.REQUIRED) Set<String> allowedTools,
@Schema(
requiredMode = Schema.RequiredMode.REQUIRED,
description = McpAccessKey.ALLOWED_TOOLS_DESCRIPTION)
Set<String> allowedTools,
@Schema(
requiredMode = Schema.RequiredMode.REQUIRED,
description = "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.")
Expand All @@ -328,7 +335,10 @@ record AccessKeyView(
@Schema(requiredMode = Schema.RequiredMode.REQUIRED) String ownerName,
@Schema(requiredMode = Schema.RequiredMode.REQUIRED) boolean enabled,
Instant expiresAt,
@Schema(requiredMode = Schema.RequiredMode.REQUIRED) Set<String> allowedTools,
@Schema(
requiredMode = Schema.RequiredMode.REQUIRED,
description = McpAccessKey.ALLOWED_TOOLS_DESCRIPTION)
Set<String> allowedTools,
@Schema(
requiredMode = Schema.RequiredMode.REQUIRED,
description = "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.")
Expand Down
5 changes: 0 additions & 5 deletions src/main/java/run/halo/mcpserver/McpAuthorization.java
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
package run.halo.mcpserver;

import java.util.Set;
import java.util.function.Supplier;
import org.springframework.security.core.context.ReactiveSecurityContextHolder;
import org.springframework.security.core.context.SecurityContext;
Expand Down Expand Up @@ -48,8 +47,4 @@ Mono<McpKeyAuthenticationToken> authentication() {
public Mono<String> username() {
return authentication().map(McpKeyAuthenticationToken::getName);
}

Mono<Set<String>> allowedTools() {
return authentication().map(McpKeyAuthenticationToken::allowedTools);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@

final class McpKeyAuthenticationToken extends AbstractAuthenticationToken {

static final String ALL_TOOLS = "*";

private final String keyId;
private final String keyDisplayName;
private final String keyPrefix;
Expand Down Expand Up @@ -40,12 +42,12 @@ String keyPrefix() {
return keyPrefix;
}

Set<String> allowedTools() {
return allowedTools;
boolean allowsAllTools() {
return allowedTools.contains(ALL_TOOLS);
}

boolean allows(String toolName) {
return allowedTools.contains(toolName);
return allowsAllTools() || allowedTools.contains(toolName);
}

@Override
Expand Down
8 changes: 8 additions & 0 deletions src/main/java/run/halo/mcpserver/McpToolCatalog.java
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,14 @@ Mono<List<McpSchema.Tool>> protocolTools() {
});
}

Mono<Boolean> hasProtocolTool(String name) {
if (builtInTools.tools().stream().anyMatch(tool -> tool.protocolTool().name().equals(name))) {
return Mono.just(true);
}
return contributedTools().map(tools ->
tools.stream().anyMatch(tool -> tool.protocolName().equals(name)));
}

private Mono<List<RegisteredTool>> contributedTools() {
return toolRegistry.registeredTools().onErrorResume(error -> {
log.warn("Ignoring unavailable contributed MCP tools", error);
Expand Down
58 changes: 55 additions & 3 deletions src/test/java/run/halo/mcpserver/AuthorizedMcpTransportTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -4,13 +4,15 @@
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;

import io.modelcontextprotocol.common.McpTransportContext;
import io.modelcontextprotocol.json.jackson3.JacksonMcpJsonMapper;
import io.modelcontextprotocol.server.McpStatelessServerHandler;
import io.modelcontextprotocol.spec.McpSchema;
import java.util.List;
import java.util.Map;
import java.util.Set;
import org.junit.jupiter.api.Test;
Expand Down Expand Up @@ -51,6 +53,50 @@ void sanitizesInternalErrorsFromBuiltInToolCalls() {
assertInternalErrorIsSanitized(response, 2);
}

@Test
void wildcardListsEveryCurrentlyAvailableTool() {
var fixture = fixture();
var first = McpSchema.Tool.builder("halo_first", Map.of("type", "object")).build();
var addedLater = McpSchema.Tool.builder("PluginExample__added_later", Map.of("type", "object"))
.build();
when(fixture.catalog().protocolTools()).thenReturn(Mono.just(List.of(first, addedLater)));
var request = new McpSchema.JSONRPCRequest("tools/list", 3, Map.of());
var authentication = new McpKeyAuthenticationToken(
"key-id", "Automation", "hmcp_key", "admin", Set.of("*"));

var response = fixture.handler()
.handleRequest(McpTransportContext.EMPTY, request)
.contextWrite(ReactiveSecurityContextHolder.withAuthentication(authentication))
.block();

assertThat(response.result()).isInstanceOf(McpSchema.ListToolsResult.class);
var result = (McpSchema.ListToolsResult) response.result();
assertThat(result.tools()).extracting(McpSchema.Tool::name)
.containsExactly("halo_first", "PluginExample__added_later");
}

@Test
void wildcardCallsKeepAvailableToolBucketsAndShareTheUnauthorizedBucket() {
var fixture = fixture();
when(fixture.catalog().hasProtocolTool(any())).thenAnswer(invocation -> Mono.just(
Set.of("known_one", "known_two").contains(invocation.getArgument(0))));
var authentication = new McpKeyAuthenticationToken(
"key-id", "Automation", "hmcp_key", "admin", Set.of("*"));

for (var toolName : List.of("known_one", "known_two", "missing_one", "missing_two")) {
var request = new McpSchema.JSONRPCRequest(
"tools/call", 4, Map.of("name", toolName, "arguments", Map.of()));
fixture.handler()
.handleRequest(McpTransportContext.EMPTY, request)
.contextWrite(ReactiveSecurityContextHolder.withAuthentication(authentication))
.block();
}

verify(fixture.rateLimiter()).allowTool("key-id", "known_one");
verify(fixture.rateLimiter()).allowTool("key-id", "known_two");
verify(fixture.rateLimiter(), times(2)).allowTool("key-id", "<unauthorized>");
}

@Test
void acceptsInitializedNotificationWithoutDelegatingToMissingSdkHandler() {
var fixture = fixture();
Expand Down Expand Up @@ -79,14 +125,16 @@ private static Fixture fixture() {
var delegate = mock(WebFluxStatelessServerTransport.class);
var catalog = mock(McpToolCatalog.class);
var registry = mock(McpToolRegistry.class);
var rateLimiter = mock(McpRequestRateLimiter.class);
when(rateLimiter.allowTool(any(), any())).thenReturn(true);
var authorization = new McpAuthorization();
var transport = new AuthorizedMcpTransport(
delegate,
new JacksonMcpJsonMapper(JsonMapper.shared()),
catalog,
registry,
authorization,
new McpRequestRateLimiter(),
rateLimiter,
new McpRecentCallHistory());
var sdkHandler = mock(McpStatelessServerHandler.class);
when(sdkHandler.handleNotification(any(), any())).thenReturn(Mono.empty());
Expand All @@ -99,7 +147,7 @@ private static Fixture fixture() {
transport.setMcpHandler(sdkHandler);
var captor = ArgumentCaptor.forClass(McpStatelessServerHandler.class);
verify(delegate).setMcpHandler(captor.capture());
return new Fixture(captor.getValue(), sdkHandler);
return new Fixture(captor.getValue(), sdkHandler, catalog, rateLimiter);
}

private static void assertInternalErrorIsSanitized(
Expand All @@ -112,5 +160,9 @@ private static void assertInternalErrorIsSanitized(
assertThat(response.toString()).doesNotContain("database password");
}

private record Fixture(McpStatelessServerHandler handler, McpStatelessServerHandler sdkHandler) {}
private record Fixture(
McpStatelessServerHandler handler,
McpStatelessServerHandler sdkHandler,
McpToolCatalog catalog,
McpRequestRateLimiter rateLimiter) {}
}
Loading