From c795166cdd2fda4b1d31f08c4bc7778d8abf6ea8 Mon Sep 17 00:00:00 2001 From: Ryan Wang Date: Mon, 24 Aug 2026 12:17:38 +0800 Subject: [PATCH 1/3] Restrict MCP access keys by IP address --- README.md | 11 +++ api-docs/openapi/v3_0/mcpV1alpha1Api.json | 33 ++++++++- .../java/run/halo/mcpserver/McpAccessKey.java | 2 + .../halo/mcpserver/McpAccessKeyEndpoint.java | 19 ++++- .../halo/mcpserver/McpAccessKeyService.java | 12 +++- .../run/halo/mcpserver/McpIpAllowlist.java | 54 ++++++++++++++ .../mcpserver/McpKeyAuthenticationFilter.java | 2 +- .../mcpserver/McpAccessKeyEndpointTest.java | 31 ++++++++ .../mcpserver/McpAccessKeyServiceTest.java | 50 ++++++++++--- .../halo/mcpserver/McpIpAllowlistTest.java | 71 ++++++++++++++++++ .../McpKeyAuthenticationFilterTest.java | 22 +++--- .../models/create-mcp-access-key-request.ts | 6 ++ ui/src/api/generated/models/mcp-access-key.ts | 6 ++ .../models/update-mcp-access-key-request.ts | 6 ++ ui/src/components/AccessKeyCreationModal.vue | 1 + ui/src/components/AccessKeyForm.vue | 18 +++++ ui/src/components/AccessKeyListItem.vue | 8 +++ .../__tests__/AccessKeyForm.test.ts | 72 +++++++++++++++++++ .../__tests__/AccessKeyListItem.test.ts | 8 +++ .../__tests__/McpRecentCalls.test.ts | 1 + .../__tests__/useAccessKeys.test.ts | 1 + 21 files changed, 411 insertions(+), 23 deletions(-) create mode 100644 src/main/java/run/halo/mcpserver/McpIpAllowlist.java create mode 100644 src/test/java/run/halo/mcpserver/McpIpAllowlistTest.java create mode 100644 ui/src/components/__tests__/AccessKeyForm.test.ts diff --git a/README.md b/README.md index 351098d..153bfbf 100644 --- a/README.md +++ b/README.md @@ -40,6 +40,17 @@ Tool access is independent of Halo content RBAC: the key's exact tool allowlist is the authorization boundary. Newly installed tools are denied until an administrator explicitly adds them to a key. Disabled and expired keys are rejected, and rotating a key invalidates its previous secret immediately. +Each key can optionally restrict access to exact IPv4 or IPv6 addresses and CIDR +ranges. An empty IP allowlist means unrestricted access. Requests from an +unmatched or unknown source are rejected as unauthorized and do not update the +key's last-used time. + +IP restrictions use the remote address normalized by Halo's HTTP stack. When +Halo is behind a reverse proxy, configure the proxy and Halo so that untrusted +clients cannot supply or preserve `Forwarded` or `X-Forwarded-*` headers, and +prevent direct access that bypasses the trusted proxy. An IP allowlist is an +additional control, not a replacement for TLS and least-privilege tool access. + Requests carrying an MCP Bearer token are limited to 600 per minute per observed network source before key validation. This is an overall source-level ceiling and includes successful requests. Tool calls are additionally limited to 120 per diff --git a/api-docs/openapi/v3_0/mcpV1alpha1Api.json b/api-docs/openapi/v3_0/mcpV1alpha1Api.json index a99f853..3322ce9 100644 --- a/api-docs/openapi/v3_0/mcpV1alpha1Api.json +++ b/api-docs/openapi/v3_0/mcpV1alpha1Api.json @@ -272,9 +272,18 @@ } }, "CreateMcpAccessKeyRequest" : { - "required" : [ "allowedTools", "displayName" ], + "required" : [ "allowedIpRanges", "allowedTools", "displayName" ], "type" : "object", "properties" : { + "allowedIpRanges" : { + "uniqueItems" : true, + "type" : "array", + "description" : "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.", + "items" : { + "type" : "string", + "description" : "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted." + } + }, "allowedTools" : { "uniqueItems" : true, "type" : "array", @@ -325,9 +334,18 @@ } }, "McpAccessKey" : { - "required" : [ "allowedTools", "displayName", "enabled", "keyPrefix", "name", "ownerName" ], + "required" : [ "allowedIpRanges", "allowedTools", "displayName", "enabled", "keyPrefix", "name", "ownerName" ], "type" : "object", "properties" : { + "allowedIpRanges" : { + "uniqueItems" : true, + "type" : "array", + "description" : "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.", + "items" : { + "type" : "string", + "description" : "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted." + } + }, "allowedTools" : { "uniqueItems" : true, "type" : "array", @@ -583,9 +601,18 @@ } }, "UpdateMcpAccessKeyRequest" : { - "required" : [ "allowedTools", "displayName", "enabled" ], + "required" : [ "allowedIpRanges", "allowedTools", "displayName", "enabled" ], "type" : "object", "properties" : { + "allowedIpRanges" : { + "uniqueItems" : true, + "type" : "array", + "description" : "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.", + "items" : { + "type" : "string", + "description" : "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted." + } + }, "allowedTools" : { "uniqueItems" : true, "type" : "array", diff --git a/src/main/java/run/halo/mcpserver/McpAccessKey.java b/src/main/java/run/halo/mcpserver/McpAccessKey.java index 2135e32..bfc6059 100644 --- a/src/main/java/run/halo/mcpserver/McpAccessKey.java +++ b/src/main/java/run/halo/mcpserver/McpAccessKey.java @@ -34,6 +34,8 @@ public static class Spec { private boolean enabled = true; private Instant expiresAt; private Set allowedTools = new LinkedHashSet<>(); + @Schema(description = "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.") + private Set allowedIpRanges = new LinkedHashSet<>(); } @Data diff --git a/src/main/java/run/halo/mcpserver/McpAccessKeyEndpoint.java b/src/main/java/run/halo/mcpserver/McpAccessKeyEndpoint.java index c30e0b6..65b429d 100644 --- a/src/main/java/run/halo/mcpserver/McpAccessKeyEndpoint.java +++ b/src/main/java/run/halo/mcpserver/McpAccessKeyEndpoint.java @@ -148,6 +148,7 @@ private Mono create(ServerRequest request) { tuple.getT1().displayName(), tuple.getT2(), tools(tuple.getT1().allowedTools()), + tuple.getT1().allowedIpRanges(), tuple.getT1().expiresAt())) .flatMap(created -> ServerResponse.created(URI.create("keys/" + created.accessKey() .getMetadata() @@ -166,11 +167,14 @@ private Mono update(ServerRequest request) { name, body.displayName(), tools(body.allowedTools()), + body.allowedIpRanges(), body.expiresAt(), body.enabled())) .flatMap(key -> ServerResponse.ok().bodyValue(view(key))) .onErrorMap(McpAccessKeyService.AccessKeyNotFoundException.class, error -> - new ResponseStatusException(HttpStatus.NOT_FOUND, error.getMessage(), error)); + new ResponseStatusException(HttpStatus.NOT_FOUND, error.getMessage(), error)) + .onErrorMap(IllegalArgumentException.class, error -> + new ServerWebInputException(error.getMessage(), null, error)); } private Mono rotate(ServerRequest request) { @@ -277,6 +281,7 @@ private static AccessKeyView view(McpAccessKey key) { spec.isEnabled(), spec.getExpiresAt(), spec.getAllowedTools() == null ? Set.of() : Set.copyOf(spec.getAllowedTools()), + spec.getAllowedIpRanges() == null ? Set.of() : Set.copyOf(spec.getAllowedIpRanges()), status == null ? null : status.getLastUsedAt(), metadata.getCreationTimestamp(), metadata.getDeletionTimestamp()); @@ -291,12 +296,20 @@ public GroupVersion groupVersion() { record CreateKeyRequest( @Schema(requiredMode = Schema.RequiredMode.REQUIRED) String displayName, @Schema(requiredMode = Schema.RequiredMode.REQUIRED) Set allowedTools, + @Schema( + requiredMode = Schema.RequiredMode.REQUIRED, + description = "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.") + Set allowedIpRanges, Instant expiresAt) {} @Schema(name = "UpdateMcpAccessKeyRequest") record UpdateKeyRequest( @Schema(requiredMode = Schema.RequiredMode.REQUIRED) String displayName, @Schema(requiredMode = Schema.RequiredMode.REQUIRED) Set allowedTools, + @Schema( + requiredMode = Schema.RequiredMode.REQUIRED, + description = "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.") + Set allowedIpRanges, Instant expiresAt, @Schema(requiredMode = Schema.RequiredMode.REQUIRED) boolean enabled) {} @@ -309,6 +322,10 @@ record AccessKeyView( @Schema(requiredMode = Schema.RequiredMode.REQUIRED) boolean enabled, Instant expiresAt, @Schema(requiredMode = Schema.RequiredMode.REQUIRED) Set allowedTools, + @Schema( + requiredMode = Schema.RequiredMode.REQUIRED, + description = "Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted.") + Set allowedIpRanges, Instant lastUsedAt, Instant creationTimestamp, Instant deletionTimestamp) {} diff --git a/src/main/java/run/halo/mcpserver/McpAccessKeyService.java b/src/main/java/run/halo/mcpserver/McpAccessKeyService.java index c2879be..9d90c59 100644 --- a/src/main/java/run/halo/mcpserver/McpAccessKeyService.java +++ b/src/main/java/run/halo/mcpserver/McpAccessKeyService.java @@ -1,5 +1,6 @@ package run.halo.mcpserver; +import java.net.InetSocketAddress; import java.security.SecureRandom; import java.time.Instant; import java.util.Base64; @@ -51,7 +52,9 @@ Mono create( String displayName, String ownerName, Set allowedTools, + Set allowedIpRanges, Instant expiresAt) { + var normalizedIpRanges = McpIpAllowlist.normalize(allowedIpRanges); var id = UUID.randomUUID().toString(); var secret = randomSecret(); var token = token(id, secret); @@ -68,6 +71,7 @@ Mono create( spec.setEnabled(true); spec.setExpiresAt(expiresAt); spec.setAllowedTools(copyTools(allowedTools)); + spec.setAllowedIpRanges(normalizedIpRanges); accessKey.setSpec(spec); return client.create(accessKey).map(created -> new CreatedKey(created, token)); }); @@ -77,12 +81,15 @@ Mono update( String id, String displayName, Set allowedTools, + Set allowedIpRanges, Instant expiresAt, boolean enabled) { + var normalizedIpRanges = McpIpAllowlist.normalize(allowedIpRanges); return get(id).flatMap(accessKey -> { var spec = accessKey.getSpec(); spec.setDisplayName(requireDisplayName(displayName)); spec.setAllowedTools(copyTools(allowedTools)); + spec.setAllowedIpRanges(normalizedIpRanges); spec.setExpiresAt(expiresAt); spec.setEnabled(enabled); return client.update(accessKey); @@ -108,7 +115,8 @@ Mono delete(String id) { return get(id).flatMap(client::delete).then(); } - Mono authenticate(String rawToken) { + Mono authenticate( + String rawToken, InetSocketAddress remoteAddress) { var parsed = parse(rawToken); if (parsed == null) { return Mono.empty(); @@ -117,6 +125,8 @@ Mono authenticate(String rawToken) { .filter(this::active) .flatMap(accessKey -> matches(parsed.secret(), accessKey.getSpec().getKeyHash()) .filter(Boolean::booleanValue) + .filter(ignored -> McpIpAllowlist.allows( + accessKey.getSpec().getAllowedIpRanges(), remoteAddress)) .flatMap(ignored -> touch(accessKey).thenReturn(new McpKeyAuthenticationToken( parsed.id(), accessKey.getSpec().getDisplayName(), diff --git a/src/main/java/run/halo/mcpserver/McpIpAllowlist.java b/src/main/java/run/halo/mcpserver/McpIpAllowlist.java new file mode 100644 index 0000000..c14e673 --- /dev/null +++ b/src/main/java/run/halo/mcpserver/McpIpAllowlist.java @@ -0,0 +1,54 @@ +package run.halo.mcpserver; + +import java.net.InetSocketAddress; +import java.util.LinkedHashSet; +import java.util.Set; +import org.springframework.security.util.matcher.InetAddressMatchers; +import org.springframework.util.StringUtils; + +final class McpIpAllowlist { + + private McpIpAllowlist() {} + + static Set normalize(Set ranges) { + var normalized = new LinkedHashSet(); + if (ranges == null) { + return normalized; + } + for (var range : ranges) { + if (!StringUtils.hasText(range)) { + continue; + } + var value = range.trim(); + try { + InetAddressMatchers.fromIpAddress(value); + } catch (IllegalArgumentException error) { + throw new IllegalArgumentException("Invalid IP address or CIDR: " + value, error); + } + normalized.add(value); + } + return normalized; + } + + static boolean allows(Set ranges, InetSocketAddress remoteAddress) { + if (ranges == null || ranges.isEmpty()) { + return true; + } + if (remoteAddress == null || remoteAddress.getAddress() == null) { + return false; + } + try { + var matchers = ranges.stream() + .map(InetAddressMatchers::fromIpAddress) + .toList(); + for (var matcher : matchers) { + if (matcher.matches(remoteAddress.getAddress())) { + return true; + } + } + } catch (IllegalArgumentException error) { + return false; + } + return false; + } +} diff --git a/src/main/java/run/halo/mcpserver/McpKeyAuthenticationFilter.java b/src/main/java/run/halo/mcpserver/McpKeyAuthenticationFilter.java index 9d9e162..7803e18 100644 --- a/src/main/java/run/halo/mcpserver/McpKeyAuthenticationFilter.java +++ b/src/main/java/run/halo/mcpserver/McpKeyAuthenticationFilter.java @@ -51,7 +51,7 @@ public Mono filter(ServerWebExchange exchange, WebFilterChain chain) { return tooManyRequests(exchange); } var rawToken = authorization.substring(BEARER_SCHEME.length()); - return accessKeyService.authenticate(rawToken) + return accessKeyService.authenticate(rawToken, exchange.getRequest().getRemoteAddress()) .flatMap(authentication -> { if (!hasSupportedProtocolVersion(exchange)) { return badRequest(exchange).thenReturn(true); diff --git a/src/test/java/run/halo/mcpserver/McpAccessKeyEndpointTest.java b/src/test/java/run/halo/mcpserver/McpAccessKeyEndpointTest.java index 71d4b58..b34b36b 100644 --- a/src/test/java/run/halo/mcpserver/McpAccessKeyEndpointTest.java +++ b/src/test/java/run/halo/mcpserver/McpAccessKeyEndpointTest.java @@ -1,14 +1,17 @@ package run.halo.mcpserver; import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyBoolean; import static org.mockito.Mockito.when; import java.time.Instant; +import java.util.Set; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.InjectMocks; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.http.MediaType; import org.springframework.test.web.reactive.server.WebTestClient; import reactor.core.publisher.Flux; import run.halo.app.extension.Metadata; @@ -39,6 +42,7 @@ void listsDeletionTimestamp() { key.getSpec().setDisplayName("Test key"); key.getSpec().setKeyPrefix("hmcp_test"); key.getSpec().setOwnerName("admin"); + key.getSpec().setAllowedIpRanges(Set.of("203.0.113.0/24")); when(accessKeyService.list()).thenReturn(Flux.just(key)); WebTestClient.bindToRouterFunction(endpoint.endpoint()) @@ -49,10 +53,37 @@ void listsDeletionTimestamp() { .expectStatus() .isOk() .expectBody() + .jsonPath("$[0].allowedIpRanges[0]") + .isEqualTo("203.0.113.0/24") .jsonPath("$[0].deletionTimestamp") .isEqualTo(deletionTimestamp.toString()); } + @Test + void mapsInvalidIpRangesToBadRequestWhenUpdating() { + when(toolCatalog.availableNames()).thenReturn(reactor.core.publisher.Mono.just(Set.of())); + when(accessKeyService.update(any(), any(), any(), any(), any(), anyBoolean())) + .thenReturn(reactor.core.publisher.Mono.error( + new IllegalArgumentException("Invalid IP address or CIDR: invalid"))); + + WebTestClient.bindToRouterFunction(endpoint.endpoint()) + .build() + .put() + .uri("/keys/test-key") + .contentType(MediaType.APPLICATION_JSON) + .bodyValue(""" + { + "displayName": "Test key", + "allowedTools": [], + "allowedIpRanges": ["invalid"], + "enabled": true + } + """) + .exchange() + .expectStatus() + .isBadRequest(); + } + @Test void listsRecentCallsWithFilters() { var call = new McpRecentCall( diff --git a/src/test/java/run/halo/mcpserver/McpAccessKeyServiceTest.java b/src/test/java/run/halo/mcpserver/McpAccessKeyServiceTest.java index d6aedc9..db2dd17 100644 --- a/src/test/java/run/halo/mcpserver/McpAccessKeyServiceTest.java +++ b/src/test/java/run/halo/mcpserver/McpAccessKeyServiceTest.java @@ -6,6 +6,7 @@ import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; +import java.net.InetSocketAddress; import java.time.Instant; import java.util.Set; import org.junit.jupiter.api.BeforeEach; @@ -39,6 +40,7 @@ void createsAHashedKeyAndAuthenticatesIt() { "Automation", "admin", Set.of("halo_search_content"), + Set.of(), Instant.now().plusSeconds(3600)) .block(); @@ -51,7 +53,7 @@ void createsAHashedKeyAndAuthenticatesIt() { when(client.fetch(McpAccessKey.class, created.accessKey().getMetadata().getName())) .thenReturn(Mono.just(created.accessKey())); - var authentication = service.authenticate(created.token()).block(); + var authentication = service.authenticate(created.token(), null).block(); assertThat(authentication).isNotNull(); assertThat(authentication.getName()).isEqualTo("admin"); @@ -64,38 +66,68 @@ void createsAHashedKeyAndAuthenticatesIt() { @Test void rejectsDisabledAndExpiredKeys() { - var created = service.create("Expired", "admin", Set.of(), Instant.now().minusSeconds(1)) + var created = service.create( + "Expired", "admin", Set.of(), Set.of(), Instant.now().minusSeconds(1)) .block(); when(client.fetch(McpAccessKey.class, created.accessKey().getMetadata().getName())) .thenReturn(Mono.just(created.accessKey())); - assertThat(service.authenticate(created.token()).block()).isNull(); + assertThat(service.authenticate(created.token(), null).block()).isNull(); created.accessKey().getSpec().setExpiresAt(null); created.accessKey().getSpec().setEnabled(false); - assertThat(service.authenticate(created.token()).block()).isNull(); + assertThat(service.authenticate(created.token(), null).block()).isNull(); } @Test void rejectsKeysBeingDeleted() { - var created = service.create("Deleting", "admin", Set.of(), null).block(); + var created = service.create("Deleting", "admin", Set.of(), Set.of(), null) + .block(); created.accessKey().getMetadata().setDeletionTimestamp(Instant.now()); when(client.fetch(McpAccessKey.class, created.accessKey().getMetadata().getName())) .thenReturn(Mono.just(created.accessKey())); - assertThat(service.authenticate(created.token()).block()).isNull(); + assertThat(service.authenticate(created.token(), null).block()).isNull(); } @Test void rotationInvalidatesThePreviousSecret() { - var created = service.create("Automation", "admin", Set.of(), null).block(); + var created = service.create("Automation", "admin", Set.of(), Set.of(), null) + .block(); var id = created.accessKey().getMetadata().getName(); when(client.fetch(McpAccessKey.class, id)).thenReturn(Mono.just(created.accessKey())); var rotated = service.rotate(id).block(); assertThat(rotated.token()).isNotEqualTo(created.token()); - assertThat(service.authenticate(created.token()).block()).isNull(); - assertThat(service.authenticate(rotated.token()).block()).isNotNull(); + assertThat(service.authenticate(created.token(), null).block()).isNull(); + assertThat(service.authenticate(rotated.token(), null).block()).isNotNull(); + } + + @Test + void authenticatesOnlyFromAnAllowedIpRange() { + var created = service.create( + "Restricted", + "admin", + Set.of(), + Set.of(" 203.0.113.0/24 "), + null) + .block(); + var id = created.accessKey().getMetadata().getName(); + when(client.fetch(McpAccessKey.class, id)).thenReturn(Mono.just(created.accessKey())); + + assertThat(created.accessKey().getSpec().getAllowedIpRanges()) + .containsExactly("203.0.113.0/24"); + assertThat(service.authenticate( + created.token(), new InetSocketAddress("198.51.100.10", 443)) + .block()) + .isNull(); + assertThat(created.accessKey().getStatus().getLastUsedAt()).isNull(); + + assertThat(service.authenticate( + created.token(), new InetSocketAddress("203.0.113.42", 443)) + .block()) + .isNotNull(); + assertThat(created.accessKey().getStatus().getLastUsedAt()).isNotNull(); } } diff --git a/src/test/java/run/halo/mcpserver/McpIpAllowlistTest.java b/src/test/java/run/halo/mcpserver/McpIpAllowlistTest.java new file mode 100644 index 0000000..f7d6203 --- /dev/null +++ b/src/test/java/run/halo/mcpserver/McpIpAllowlistTest.java @@ -0,0 +1,71 @@ +package run.halo.mcpserver; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.net.InetSocketAddress; +import java.util.LinkedHashSet; +import java.util.Set; +import org.junit.jupiter.api.Test; + +class McpIpAllowlistTest { + + @Test + void normalizesAndValidatesRanges() { + var ranges = new LinkedHashSet<>( + java.util.List.of(" 203.0.113.10 ", "203.0.113.0/24", "2001:db8::/32")); + + assertThat(McpIpAllowlist.normalize(ranges)) + .containsExactly("203.0.113.10", "203.0.113.0/24", "2001:db8::/32"); + } + + @Test + void rejectsHostnamesAndInvalidMasks() { + assertThatThrownBy(() -> McpIpAllowlist.normalize(Set.of("example.com"))) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("Invalid IP address or CIDR: example.com"); + assertThatThrownBy(() -> McpIpAllowlist.normalize(Set.of("203.0.113.0/33"))) + .isInstanceOf(IllegalArgumentException.class) + .hasMessage("Invalid IP address or CIDR: 203.0.113.0/33"); + } + + @Test + void allowsExactAddressesAndCidrsForBothAddressFamilies() { + assertThat(McpIpAllowlist.allows( + Set.of("203.0.113.10"), new InetSocketAddress("203.0.113.10", 443))) + .isTrue(); + assertThat(McpIpAllowlist.allows( + Set.of("203.0.113.0/24"), new InetSocketAddress("203.0.113.42", 443))) + .isTrue(); + assertThat(McpIpAllowlist.allows( + Set.of("2001:db8::/32"), new InetSocketAddress("2001:db8::42", 443))) + .isTrue(); + } + + @Test + void rejectsMismatchesAndUnknownAddressesWhenConfigured() { + assertThat(McpIpAllowlist.allows( + Set.of("203.0.113.0/24"), new InetSocketAddress("198.51.100.10", 443))) + .isFalse(); + assertThat(McpIpAllowlist.allows(Set.of("203.0.113.0/24"), null)) + .isFalse(); + assertThat(McpIpAllowlist.allows( + Set.of("203.0.113.0/24"), + InetSocketAddress.createUnresolved("client.example.com", 443))) + .isFalse(); + } + + @Test + void allowsUnknownAddressesWhenNotConfigured() { + assertThat(McpIpAllowlist.allows(Set.of(), null)).isTrue(); + } + + @Test + void failsClosedWhenStoredConfigurationIsInvalid() { + var ranges = new LinkedHashSet<>(java.util.List.of("203.0.113.10", "invalid")); + + assertThat(McpIpAllowlist.allows( + ranges, new InetSocketAddress("203.0.113.10", 443))) + .isFalse(); + } +} diff --git a/src/test/java/run/halo/mcpserver/McpKeyAuthenticationFilterTest.java b/src/test/java/run/halo/mcpserver/McpKeyAuthenticationFilterTest.java index d9d283b..321c618 100644 --- a/src/test/java/run/halo/mcpserver/McpKeyAuthenticationFilterTest.java +++ b/src/test/java/run/halo/mcpserver/McpKeyAuthenticationFilterTest.java @@ -5,6 +5,7 @@ import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; +import java.net.InetSocketAddress; import java.util.Set; import java.util.concurrent.atomic.AtomicReference; import org.junit.jupiter.api.BeforeEach; @@ -67,21 +68,26 @@ void authenticatesAndStripsTheMcpBearerTokenBeforeTheHaloJwtFilter() { "hmcp_00000000", "admin", Set.of("halo_search_content")); - when(accessKeyService.authenticate(rawToken)).thenReturn(Mono.just(authentication)); + var remoteAddress = new InetSocketAddress("203.0.113.8", 41321); + when(accessKeyService.authenticate(rawToken, remoteAddress)) + .thenReturn(Mono.just(authentication)); var exchange = MockServerWebExchange.from(MockServerHttpRequest.post(McpKeyAuthenticationFilter.MCP_PATH) - .header(HttpHeaders.AUTHORIZATION, "bEaReR " + rawToken)); + .remoteAddress(remoteAddress) + .header(HttpHeaders.AUTHORIZATION, "bEaReR " + rawToken) + .header("X-Forwarded-For", "198.51.100.9")); filter.filter(exchange, ignored -> Mono.error(new AssertionError("Halo chain must not continue"))) .block(); assertThat(handledAuthorization.get()).isNull(); assertThat(handledPath.get()).isEqualTo(McpKeyAuthenticationFilter.MCP_PATH); assertThat(currentAuthentication.get()).isSameAs(authentication); + verify(accessKeyService).authenticate(rawToken, remoteAddress); } @Test void rejectsAnInvalidMcpKey() { var rawToken = "hmcp_00000000-0000-0000-0000-000000000000_invalid"; - when(accessKeyService.authenticate(rawToken)).thenReturn(Mono.empty()); + when(accessKeyService.authenticate(rawToken, null)).thenReturn(Mono.empty()); var exchange = MockServerWebExchange.from(MockServerHttpRequest.post(McpKeyAuthenticationFilter.MCP_PATH) .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken)); @@ -119,7 +125,7 @@ void neverUsesAnMcpKeyOutsideTheMcpEndpoint() { }) .block(); - verify(accessKeyService, never()).authenticate(rawToken); + verify(accessKeyService, never()).authenticate(rawToken, null); assertThat(forwarded.get().getRequest().getHeaders().getFirst(HttpHeaders.AUTHORIZATION)) .isEqualTo("Bearer " + rawToken); } @@ -133,14 +139,14 @@ void authenticatesATrailingSlashAsTheSameMcpEndpoint() { "hmcp_00000000", "admin", Set.of()); - when(accessKeyService.authenticate(rawToken)).thenReturn(Mono.just(authentication)); + when(accessKeyService.authenticate(rawToken, null)).thenReturn(Mono.just(authentication)); var exchange = MockServerWebExchange.from(MockServerHttpRequest.post("/mcp/") .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken)); filter.filter(exchange, ignored -> Mono.error(new AssertionError("Halo chain must not continue"))) .block(); - verify(accessKeyService).authenticate(rawToken); + verify(accessKeyService).authenticate(rawToken, null); assertThat(handledPath.get()).isEqualTo("/mcp/"); } @@ -153,7 +159,7 @@ void rejectsAnUnsupportedProtocolVersionAfterAuthentication() { "hmcp_00000000", "admin", Set.of("halo_search_content")); - when(accessKeyService.authenticate(rawToken)).thenReturn(Mono.just(authentication)); + when(accessKeyService.authenticate(rawToken, null)).thenReturn(Mono.just(authentication)); var exchange = MockServerWebExchange.from(MockServerHttpRequest.post(McpKeyAuthenticationFilter.MCP_PATH) .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken) .header(io.modelcontextprotocol.spec.HttpHeaders.PROTOCOL_VERSION, "2099-01-01")); @@ -179,6 +185,6 @@ void rateLimitsRequestsBeforeAccessKeyLookup() { assertThat(exchange.getResponse().getStatusCode()).isEqualTo(HttpStatus.TOO_MANY_REQUESTS); assertThat(exchange.getResponse().getHeaders().getFirst(HttpHeaders.RETRY_AFTER)) .isEqualTo("60"); - verify(accessKeyService, never()).authenticate(rawToken); + verify(accessKeyService, never()).authenticate(rawToken, null); } } diff --git a/ui/src/api/generated/models/create-mcp-access-key-request.ts b/ui/src/api/generated/models/create-mcp-access-key-request.ts index cca296c..24bc83c 100644 --- a/ui/src/api/generated/models/create-mcp-access-key-request.ts +++ b/ui/src/api/generated/models/create-mcp-access-key-request.ts @@ -20,6 +20,12 @@ * @interface CreateMcpAccessKeyRequest */ export interface CreateMcpAccessKeyRequest { + /** + * Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted. + * @type {Array} + * @memberof CreateMcpAccessKeyRequest + */ + 'allowedIpRanges': Array; /** * * @type {Array} diff --git a/ui/src/api/generated/models/mcp-access-key.ts b/ui/src/api/generated/models/mcp-access-key.ts index e302c3c..e5e4bf5 100644 --- a/ui/src/api/generated/models/mcp-access-key.ts +++ b/ui/src/api/generated/models/mcp-access-key.ts @@ -20,6 +20,12 @@ * @interface McpAccessKey */ export interface McpAccessKey { + /** + * Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted. + * @type {Array} + * @memberof McpAccessKey + */ + 'allowedIpRanges': Array; /** * * @type {Array} diff --git a/ui/src/api/generated/models/update-mcp-access-key-request.ts b/ui/src/api/generated/models/update-mcp-access-key-request.ts index c7bd640..788c0b2 100644 --- a/ui/src/api/generated/models/update-mcp-access-key-request.ts +++ b/ui/src/api/generated/models/update-mcp-access-key-request.ts @@ -20,6 +20,12 @@ * @interface UpdateMcpAccessKeyRequest */ export interface UpdateMcpAccessKeyRequest { + /** + * Allowed IPv4/IPv6 addresses or CIDR ranges. Empty means unrestricted. + * @type {Array} + * @memberof UpdateMcpAccessKeyRequest + */ + 'allowedIpRanges': Array; /** * * @type {Array} diff --git a/ui/src/components/AccessKeyCreationModal.vue b/ui/src/components/AccessKeyCreationModal.vue index d57d21c..90c5265 100644 --- a/ui/src/components/AccessKeyCreationModal.vue +++ b/ui/src/components/AccessKeyCreationModal.vue @@ -24,6 +24,7 @@ const { mutate, isLoading: submitting } = useMutation({ const { data } = await mcpConsoleApiClient.createMcpAccessKey({ createMcpAccessKeyRequest: { displayName: input.displayName, + allowedIpRanges: input.allowedIpRanges, allowedTools: input.allowedTools, expiresAt: input.expiresAt, }, diff --git a/ui/src/components/AccessKeyForm.vue b/ui/src/components/AccessKeyForm.vue index debd4d9..43a2183 100644 --- a/ui/src/components/AccessKeyForm.vue +++ b/ui/src/components/AccessKeyForm.vue @@ -22,6 +22,7 @@ const expiresAt = shallowRef( props.accessKey?.expiresAt ? utils.date.toDatetimeLocal(props.accessKey.expiresAt) : '', ) const enabled = shallowRef(props.accessKey?.enabled ?? true) +const allowedIpRanges = shallowRef((props.accessKey?.allowedIpRanges ?? []).join('\n')) const selected = reactive>( Object.fromEntries( props.tools.map((tool) => [ @@ -47,6 +48,14 @@ function applyPreset(preset: 'read' | 'content' | 'all' | 'none') { function onSubmit() { emit('submit', { displayName: displayName.value, + allowedIpRanges: [ + ...new Set( + allowedIpRanges.value + .split(/\r?\n/) + .map((range) => range.trim()) + .filter(Boolean), + ), + ], allowedTools: props.tools.filter((tool) => selected[tool.name]).map((tool) => tool.name), expiresAt: expiresAt.value ? utils.date.toISOString(expiresAt.value) : undefined, enabled: enabled.value, @@ -76,6 +85,15 @@ defineExpose({ help="留空表示永不过期" /> +
diff --git a/ui/src/components/AccessKeyListItem.vue b/ui/src/components/AccessKeyListItem.vue index 3fd41e8..9f5a4ce 100644 --- a/ui/src/components/AccessKeyListItem.vue +++ b/ui/src/components/AccessKeyListItem.vue @@ -42,6 +42,7 @@ const { mutate: changeEnabled, isLoading: changingEnabled } = useMutation({ name: props.mcpAccessKey.name, updateMcpAccessKeyRequest: { displayName: props.mcpAccessKey.displayName, + allowedIpRanges: props.mcpAccessKey.allowedIpRanges, allowedTools: props.mcpAccessKey.allowedTools, expiresAt: props.mcpAccessKey.expiresAt, enabled, @@ -112,6 +113,13 @@ function handleDelete() {