diff --git a/.github/workflows/cd.yaml b/.github/workflows/cd.yaml index be35136..7ab5153 100644 --- a/.github/workflows/cd.yaml +++ b/.github/workflows/cd.yaml @@ -7,7 +7,7 @@ on: jobs: cd: - uses: halo-sigs/reusable-workflows/.github/workflows/plugin-cd.yaml@v4 + uses: ./.github/workflows/plugin-cd.yaml permissions: contents: write with: diff --git a/.github/workflows/plugin-cd.yaml b/.github/workflows/plugin-cd.yaml new file mode 100644 index 0000000..abf4813 --- /dev/null +++ b/.github/workflows/plugin-cd.yaml @@ -0,0 +1,133 @@ +name: CD + +on: + workflow_call: + inputs: + artifacts-path: + type: string + required: false + default: build/libs + description: Artifacts path, default is build/libs. Must be a folder. + node-version: + description: Node.js version. + type: string + required: false + default: "24" + pnpm-version: + description: pnpm version. + type: string + required: false + default: "10" + java-version: + description: Java version. + type: string + required: false + default: "21" + ui-path: + description: Path of UI project. + type: string + required: false + default: "console" + skip-node-setup: + description: Indicates if the node setup should be skipped. + type: boolean + required: false + default: false + skip-appstore-release: + description: Indicates if the appstore release should be skipped. + type: boolean + required: false + default: false + app-id: + description: Application ID from Halo App Store. + required: false + type: string + default: not-configured-app-id + halo-backend-baseurl: + description: Base URL of Halo App Store. + required: false + default: https://www.halo.run + type: string + npm-registry-url: + type: string + default: "" + description: NPM registry. + required: false + build-args: + type: string + description: Additional build arguments. + required: false + default: "" + secrets: + halo-pat: + description: Personal access token of Halo Appstore. + required: false + npm-auth-token: + description: NPM auth token. + required: false +jobs: + build: + name: Build + runs-on: ubuntu-latest + if: github.event_name == 'release' + steps: + - uses: actions/checkout@v6 + - name: Setup Environment + uses: halo-sigs/reusable-workflows/plugin-setup-env@99c24709710a37e286adca6879d8d5faf46224fc # v4 + with: + cache-dept-path: ${{ inputs.ui-path }}/pnpm-lock.yaml + skip-node-setup: ${{ inputs.skip-node-setup }} + node-version: ${{ inputs.node-version }} + pnpm-version: ${{ inputs.pnpm-version }} + java-version: ${{ inputs.java-version }} + npm-registry-url: ${{ inputs.npm-registry-url }} + - name: Build + run: | + version=${{ github.event.release.tag_name }} + ./gradlew clean build -x check -Pversion=${version#v} ${{ inputs.build-args }} + env: + NODE_AUTH_TOKEN: ${{ secrets.npm-auth-token }} + - name: Upload Artifacts + uses: actions/upload-artifact@v7 + with: + name: artifacts + path: ${{ inputs.artifacts-path }} + retention-days: 1 + + github-release: + name: GitHub Release + runs-on: ubuntu-latest + needs: build + if: github.event_name == 'release' + steps: + - uses: actions/checkout@v6 + - name: Download Artifacts + uses: actions/download-artifact@v8 + with: + name: artifacts + path: ${{ inputs.artifacts-path }} + - name: Upload Release Assets + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + run: gh release upload ${{ github.event.release.tag_name }} ${{ inputs.artifacts-path }}/* + appstore-release: + name: App Store Release + runs-on: ubuntu-latest + needs: build + if: ${{ github.event_name == 'release' && !inputs.skip-appstore-release }} + steps: + - uses: actions/checkout@v6 + - name: Download Artifacts + uses: actions/download-artifact@v8 + with: + name: artifacts + path: ${{ inputs.artifacts-path }} + - name: Release to App Store + uses: halo-sigs/app-store-release-action@v4 + with: + github-token: ${{secrets.GITHUB_TOKEN}} + app-id: ${{ inputs.app-id }} + halo-backend-baseurl: ${{ inputs.halo-backend-baseurl }} + release-id: ${{ github.event.release.id }} + assets-dir: ${{ inputs.artifacts-path }} + halo-pat: ${{ secrets.halo-pat }} diff --git a/README.md b/README.md index 153bfbf..211d17c 100644 --- a/README.md +++ b/README.md @@ -55,7 +55,10 @@ 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 minute for each access-key and tool pair. Limits are process-local and therefore -apply independently to each Halo replica. +apply independently to each Halo replica. Authenticated request processing has a +30-second deadline and is limited to 8 concurrent requests per key and 64 per +plugin instance. A timed-out request returns HTTP 504; a concurrency rejection +returns HTTP 429 with `Retry-After: 1`. Use a dedicated, least-privilege key. Do not put keys in URLs, configuration files committed to source control, shell history, or logs. diff --git a/src/main/java/run/halo/mcpserver/CategoryParentMutationLock.java b/src/main/java/run/halo/mcpserver/CategoryParentMutationLock.java new file mode 100644 index 0000000..15bccb7 --- /dev/null +++ b/src/main/java/run/halo/mcpserver/CategoryParentMutationLock.java @@ -0,0 +1,30 @@ +package run.halo.mcpserver; + +import io.swagger.v3.oas.annotations.media.Schema; +import java.time.Instant; +import lombok.Data; +import lombok.EqualsAndHashCode; +import run.halo.app.extension.AbstractExtension; +import run.halo.app.extension.GVK; + +@Data +@EqualsAndHashCode(callSuper = true) +@GVK( + group = "mcp.halo.run", + version = "v1alpha1", + kind = "CategoryParentMutationLock", + plural = "categoryparentmutationlocks", + singular = "categoryparentmutationlock") +public class CategoryParentMutationLock extends AbstractExtension { + + @Schema(requiredMode = Schema.RequiredMode.REQUIRED) + private Spec spec = new Spec(); + + @Data + @Schema(name = "CategoryParentMutationLockSpec") + public static class Spec { + private String holder; + private Instant expiresAt; + private long generation; + } +} diff --git a/src/main/java/run/halo/mcpserver/HaloMcpServer.java b/src/main/java/run/halo/mcpserver/HaloMcpServer.java index af276c6..62cdc2d 100644 --- a/src/main/java/run/halo/mcpserver/HaloMcpServer.java +++ b/src/main/java/run/halo/mcpserver/HaloMcpServer.java @@ -6,7 +6,6 @@ import io.modelcontextprotocol.server.McpStatelessAsyncServer; import io.modelcontextprotocol.server.transport.DefaultServerTransportSecurityValidator; import io.modelcontextprotocol.spec.McpSchema; -import java.time.Duration; import org.springframework.ai.mcp.server.webflux.transport.WebFluxStatelessServerTransport; import org.springframework.stereotype.Component; import org.springframework.web.reactive.function.server.RouterFunction; @@ -53,7 +52,6 @@ class HaloMcpServer { .capabilities(McpSchema.ServerCapabilities.builder() .tools(false) .build()) - .requestTimeout(Duration.ofSeconds(30)) .tools(builtInTools.specifications()) .build(); } diff --git a/src/main/java/run/halo/mcpserver/McpAccessKeyService.java b/src/main/java/run/halo/mcpserver/McpAccessKeyService.java index 9d90c59..1e6d917 100644 --- a/src/main/java/run/halo/mcpserver/McpAccessKeyService.java +++ b/src/main/java/run/halo/mcpserver/McpAccessKeyService.java @@ -4,6 +4,7 @@ import java.security.SecureRandom; import java.time.Instant; import java.util.Base64; +import java.util.Collections; import java.util.LinkedHashSet; import java.util.List; import java.util.Set; @@ -123,18 +124,32 @@ Mono authenticate( } return client.fetch(McpAccessKey.class, parsed.id()) .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(), - accessKey.getSpec().getKeyPrefix(), - accessKey.getSpec().getOwnerName(), - accessKey.getSpec().getAllowedTools() == null - ? Set.of() - : accessKey.getSpec().getAllowedTools())))); + .flatMap(accessKey -> { + var expected = AuthenticationState.from(accessKey); + return matches(parsed.secret(), expected.keyHash()) + .filter(Boolean::booleanValue) + .filter(ignored -> McpIpAllowlist.allows( + expected.allowedIpRanges(), remoteAddress)) + .flatMap(ignored -> touch(accessKey) + .then(revalidate(parsed.id(), expected, remoteAddress))) + .map(current -> new McpKeyAuthenticationToken( + parsed.id(), + current.getSpec().getDisplayName(), + current.getSpec().getKeyPrefix(), + current.getSpec().getOwnerName(), + current.getSpec().getAllowedTools() == null + ? Set.of() + : current.getSpec().getAllowedTools())); + }); + } + + private Mono revalidate( + String id, AuthenticationState expected, InetSocketAddress remoteAddress) { + return client.fetch(McpAccessKey.class, id) + .filter(this::active) + .filter(accessKey -> expected.equals(AuthenticationState.from(accessKey))) + .filter(accessKey -> McpIpAllowlist.allows( + accessKey.getSpec().getAllowedIpRanges(), remoteAddress)); } private Mono get(String id) { @@ -233,6 +248,32 @@ private static Set copyTools(Set allowedTools) { record CreatedKey(McpAccessKey accessKey, String token) {} + private record AuthenticationState( + String keyHash, + String ownerName, + boolean enabled, + Instant expiresAt, + Set allowedTools, + Set allowedIpRanges) { + + private static AuthenticationState from(McpAccessKey accessKey) { + var spec = accessKey.getSpec(); + return new AuthenticationState( + spec.getKeyHash(), + spec.getOwnerName(), + spec.isEnabled(), + spec.getExpiresAt(), + immutableSet(spec.getAllowedTools()), + immutableSet(spec.getAllowedIpRanges())); + } + + private static Set immutableSet(Set values) { + return values == null + ? Set.of() + : Collections.unmodifiableSet(new LinkedHashSet<>(values)); + } + } + private record ParsedKey(String id, String secret) {} static final class AccessKeyNotFoundException extends RuntimeException { diff --git a/src/main/java/run/halo/mcpserver/McpAuthorization.java b/src/main/java/run/halo/mcpserver/McpAuthorization.java index bf4ca23..adac717 100644 --- a/src/main/java/run/halo/mcpserver/McpAuthorization.java +++ b/src/main/java/run/halo/mcpserver/McpAuthorization.java @@ -45,6 +45,10 @@ public Mono username() { return authentication().map(McpKeyAuthenticationToken::getName); } + public Mono keyId() { + return authentication().map(McpKeyAuthenticationToken::keyId); + } + Mono> allowedTools() { return authentication().map(McpKeyAuthenticationToken::allowedTools); } diff --git a/src/main/java/run/halo/mcpserver/McpIpAllowlist.java b/src/main/java/run/halo/mcpserver/McpIpAllowlist.java index 9f43198..04b622d 100644 --- a/src/main/java/run/halo/mcpserver/McpIpAllowlist.java +++ b/src/main/java/run/halo/mcpserver/McpIpAllowlist.java @@ -37,13 +37,8 @@ static boolean allows(Set ranges, InetSocketAddress remoteAddress) { if (ranges == null || ranges.isEmpty()) { return true; } - if (remoteAddress == null) { - return false; - } try { - var address = remoteAddress.getAddress() == null - ? parseNumericAddress(remoteAddress.getHostString()) - : remoteAddress.getAddress(); + var address = resolve(remoteAddress); var matchers = ranges.stream() .map(McpIpAllowlist::compile) .toList(); @@ -58,6 +53,46 @@ static boolean allows(Set ranges, InetSocketAddress remoteAddress) { return false; } + static InetAddress resolve(InetSocketAddress remoteAddress) { + if (remoteAddress == null) { + throw new IllegalArgumentException("Remote address is unavailable"); + } + return remoteAddress.getAddress() == null + ? parseForwardedAddress(remoteAddress.getHostString()) + : remoteAddress.getAddress(); + } + + private static InetAddress parseForwardedAddress(String address) { + try { + return parseNumericAddress(address); + } catch (IllegalArgumentException directError) { + var value = address.startsWith("[") && address.endsWith("]") + ? address.substring(1, address.length() - 1) + : address; + String host; + String port; + if (value.startsWith("[")) { + var closingBracket = value.indexOf(']'); + if (closingBracket < 0 + || closingBracket + 1 >= value.length() + || value.charAt(closingBracket + 1) != ':') { + throw directError; + } + host = value.substring(1, closingBracket); + port = value.substring(closingBracket + 2); + } else { + var portSeparator = value.lastIndexOf(':'); + if (portSeparator < 0) { + throw directError; + } + host = value.substring(0, portSeparator); + port = value.substring(portSeparator + 1); + } + validatePort(port); + return parseNumericAddress(host); + } + } + private static CompiledRange compile(String range) { var matcher = InetAddressMatchers.fromIpAddress(range); var slashIndex = range.indexOf('/'); @@ -94,6 +129,19 @@ private static void validateMask(String mask, int maxBits) { } } + private static void validatePort(String port) { + if (!port.matches("[0-9]+")) { + throw new IllegalArgumentException("Invalid port"); + } + try { + if (Integer.parseInt(port) > 65_535) { + throw new IllegalArgumentException("Invalid port"); + } + } catch (NumberFormatException error) { + throw new IllegalArgumentException("Invalid port", error); + } + } + private record CompiledRange(int addressLength, InetAddressMatcher matcher) { boolean matches(InetAddress address) { diff --git a/src/main/java/run/halo/mcpserver/McpKeyAuthenticationFilter.java b/src/main/java/run/halo/mcpserver/McpKeyAuthenticationFilter.java index 7803e18..90e1476 100644 --- a/src/main/java/run/halo/mcpserver/McpKeyAuthenticationFilter.java +++ b/src/main/java/run/halo/mcpserver/McpKeyAuthenticationFilter.java @@ -4,6 +4,8 @@ import static org.springframework.http.HttpHeaders.RETRY_AFTER; import static org.springframework.http.HttpHeaders.WWW_AUTHENTICATE; +import java.time.Duration; +import java.util.concurrent.TimeoutException; import org.springframework.core.Ordered; import org.springframework.core.annotation.Order; import org.springframework.http.HttpStatus; @@ -20,19 +22,24 @@ class McpKeyAuthenticationFilter implements BeforeSecurityWebFilter { static final String MCP_PATH = "/mcp"; + static final Duration REQUEST_TIMEOUT = Duration.ofSeconds(30); private static final String BEARER_SCHEME = "Bearer "; + private static final String CONCURRENCY_RETRY_AFTER_SECONDS = "1"; private final McpAccessKeyService accessKeyService; private final McpRequestRateLimiter rateLimiter; + private final McpRequestConcurrencyLimiter concurrencyLimiter; private final WebHandler mcpHandler; private final java.util.Set protocolVersions; McpKeyAuthenticationFilter( McpAccessKeyService accessKeyService, McpRequestRateLimiter rateLimiter, + McpRequestConcurrencyLimiter concurrencyLimiter, HaloMcpServer mcpServer) { this.accessKeyService = accessKeyService; this.rateLimiter = rateLimiter; + this.concurrencyLimiter = concurrencyLimiter; this.mcpHandler = RouterFunctions.toWebHandler(mcpServer.routerFunction()); this.protocolVersions = java.util.Set.copyOf(mcpServer.protocolVersions()); } @@ -56,16 +63,29 @@ public Mono filter(ServerWebExchange exchange, WebFilterChain chain) { if (!hasSupportedProtocolVersion(exchange)) { return badRequest(exchange).thenReturn(true); } - var request = exchange.getRequest().mutate() - .headers(headers -> headers.remove(AUTHORIZATION)) - .build(); - return mcpHandler.handle(exchange.mutate().request(request).build()) - .contextWrite(org.springframework.security.core.context.ReactiveSecurityContextHolder - .withAuthentication(authentication)) + var permit = concurrencyLimiter.tryAcquire(authentication.keyId()); + if (permit.isEmpty()) { + return atCapacity(exchange).thenReturn(true); + } + return Mono.using( + permit::orElseThrow, + ignored -> { + var request = exchange.getRequest().mutate() + .headers(headers -> headers.remove(AUTHORIZATION)) + .build(); + return mcpHandler.handle( + exchange.mutate().request(request).build()) + .contextWrite(org.springframework.security.core.context + .ReactiveSecurityContextHolder + .withAuthentication(authentication)); + }, + McpRequestConcurrencyLimiter.Permit::close) .thenReturn(true); }) .defaultIfEmpty(false) - .flatMap(handled -> handled ? Mono.empty() : unauthorized(exchange)); + .flatMap(handled -> handled ? Mono.empty() : unauthorized(exchange)) + .timeout(REQUEST_TIMEOUT) + .onErrorResume(TimeoutException.class, ignored -> requestTimedOut(exchange)); } private static boolean isMcpPath(String path) { @@ -104,4 +124,16 @@ private static Mono tooManyRequests(ServerWebExchange exchange) { RETRY_AFTER, String.valueOf(McpRequestRateLimiter.RETRY_AFTER_SECONDS)); return exchange.getResponse().setComplete(); } + + private static Mono atCapacity(ServerWebExchange exchange) { + exchange.getResponse().setStatusCode(HttpStatus.TOO_MANY_REQUESTS); + exchange.getResponse().getHeaders().set( + RETRY_AFTER, CONCURRENCY_RETRY_AFTER_SECONDS); + return exchange.getResponse().setComplete(); + } + + private static Mono requestTimedOut(ServerWebExchange exchange) { + exchange.getResponse().setStatusCode(HttpStatus.GATEWAY_TIMEOUT); + return exchange.getResponse().setComplete(); + } } diff --git a/src/main/java/run/halo/mcpserver/McpRequestConcurrencyLimiter.java b/src/main/java/run/halo/mcpserver/McpRequestConcurrencyLimiter.java new file mode 100644 index 0000000..f3823f8 --- /dev/null +++ b/src/main/java/run/halo/mcpserver/McpRequestConcurrencyLimiter.java @@ -0,0 +1,76 @@ +package run.halo.mcpserver; + +import java.util.HashMap; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; +import java.util.concurrent.atomic.AtomicBoolean; +import org.springframework.stereotype.Component; + +/** Bounds authenticated MCP work that is active in this plugin instance. */ +@Component +final class McpRequestConcurrencyLimiter { + + static final int GLOBAL_LIMIT = 64; + static final int PER_KEY_LIMIT = 8; + + private final int globalLimit; + private final int perKeyLimit; + private final Map activeByKey = new HashMap<>(); + private int activeGlobal; + + McpRequestConcurrencyLimiter() { + this(GLOBAL_LIMIT, PER_KEY_LIMIT); + } + + McpRequestConcurrencyLimiter(int globalLimit, int perKeyLimit) { + if (globalLimit < 1 || perKeyLimit < 1 || perKeyLimit > globalLimit) { + throw new IllegalArgumentException("Invalid MCP concurrency limits"); + } + this.globalLimit = globalLimit; + this.perKeyLimit = perKeyLimit; + } + + synchronized Optional tryAcquire(String keyId) { + Objects.requireNonNull(keyId, "keyId must not be null"); + var activeForKey = activeByKey.getOrDefault(keyId, 0); + if (activeGlobal >= globalLimit || activeForKey >= perKeyLimit) { + return Optional.empty(); + } + activeGlobal++; + activeByKey.put(keyId, activeForKey + 1); + return Optional.of(new Permit(this, keyId)); + } + + private synchronized void release(String keyId) { + var activeForKey = activeByKey.get(keyId); + if (activeForKey == null) { + return; + } + activeGlobal--; + if (activeForKey == 1) { + activeByKey.remove(keyId); + } else { + activeByKey.put(keyId, activeForKey - 1); + } + } + + static final class Permit implements AutoCloseable { + + private final McpRequestConcurrencyLimiter limiter; + private final String keyId; + private final AtomicBoolean released = new AtomicBoolean(); + + private Permit(McpRequestConcurrencyLimiter limiter, String keyId) { + this.limiter = limiter; + this.keyId = keyId; + } + + @Override + public void close() { + if (released.compareAndSet(false, true)) { + limiter.release(keyId); + } + } + } +} diff --git a/src/main/java/run/halo/mcpserver/McpRequestRateLimiter.java b/src/main/java/run/halo/mcpserver/McpRequestRateLimiter.java index 225bc95..e49d19c 100644 --- a/src/main/java/run/halo/mcpserver/McpRequestRateLimiter.java +++ b/src/main/java/run/halo/mcpserver/McpRequestRateLimiter.java @@ -31,10 +31,12 @@ class McpRequestRateLimiter { } boolean allowRequest(InetSocketAddress remoteAddress) { - var source = remoteAddress == null || remoteAddress.getAddress() == null - ? "unknown" - : remoteAddress.getAddress().getHostAddress(); - return allow(currentWindow().requestCounts, source, REQUESTS_PER_MINUTE); + try { + var source = McpIpAllowlist.resolve(remoteAddress).getHostAddress(); + return allow(currentWindow().requestCounts, source, REQUESTS_PER_MINUTE); + } catch (IllegalArgumentException error) { + return false; + } } boolean allowTool(String keyId, String toolName) { diff --git a/src/main/java/run/halo/mcpserver/McpServerPlugin.java b/src/main/java/run/halo/mcpserver/McpServerPlugin.java index e5dfd58..f29e75b 100644 --- a/src/main/java/run/halo/mcpserver/McpServerPlugin.java +++ b/src/main/java/run/halo/mcpserver/McpServerPlugin.java @@ -40,6 +40,7 @@ public McpServerPlugin( @Override public void start() { schemeManager.register(McpAccessKey.class); + schemeManager.register(CategoryParentMutationLock.class); log.info("Halo MCP server started"); } @@ -47,6 +48,7 @@ public void start() { public void stop() { mcpServer.closeGracefully().block(Duration.ofSeconds(5)); rateLimiter.clear(); + schemeManager.unregister(Scheme.buildFromType(CategoryParentMutationLock.class)); schemeManager.unregister(Scheme.buildFromType(McpAccessKey.class)); log.info("Halo MCP server stopped"); } diff --git a/src/main/java/run/halo/mcpserver/McpToolRegistry.java b/src/main/java/run/halo/mcpserver/McpToolRegistry.java index e6620ae..614a790 100644 --- a/src/main/java/run/halo/mcpserver/McpToolRegistry.java +++ b/src/main/java/run/halo/mcpserver/McpToolRegistry.java @@ -54,8 +54,8 @@ Mono> executeIfContributed( .filter(tool -> tool.definition().name().equals(name)) .findFirst()) .flatMap(tool -> tool - .map(value -> gateway(() -> authorization.require(name) - .then(execute(value.definition(), arguments))) + .map(value -> gateway(() -> authorization.authorize( + name, () -> execute(value.definition(), arguments))) .map(Optional::of)) .orElseGet(() -> Mono.just(Optional.empty()))); } diff --git a/src/main/java/run/halo/mcpserver/tools/AttachmentTools.java b/src/main/java/run/halo/mcpserver/tools/AttachmentTools.java index eca71cd..33143ab 100644 --- a/src/main/java/run/halo/mcpserver/tools/AttachmentTools.java +++ b/src/main/java/run/halo/mcpserver/tools/AttachmentTools.java @@ -3,14 +3,16 @@ import static org.springframework.data.domain.Sort.Order.asc; import static org.springframework.data.domain.Sort.Order.desc; +import java.io.InputStream; import java.util.Base64; import java.util.List; import java.util.Map; +import java.util.Objects; +import org.springframework.core.io.buffer.DataBufferUtils; import org.springframework.core.io.buffer.DefaultDataBufferFactory; import org.springframework.data.domain.Sort; import org.springframework.http.MediaType; import org.springframework.stereotype.Component; -import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; import run.halo.app.core.extension.attachment.Attachment; import run.halo.app.core.extension.service.AttachmentService; @@ -30,17 +32,22 @@ class AttachmentTools extends ToolSupport implements ToolGroup { static final String DELETE = "halo_delete_attachment"; private static final int MAX_CONTENT_BYTES = 8 * 1024 * 1024; + private static final int MAX_ENCODED_CHARACTERS = 4 * ((MAX_CONTENT_BYTES + 2) / 3); + private static final int DECODE_BUFFER_BYTES = 64 * 1024; private final ReactiveExtensionClient client; private final AttachmentService attachmentService; + private final AttachmentUploadLimiter uploadLimiter; AttachmentTools( ReactiveExtensionClient client, AttachmentService attachmentService, - McpAuthorization authorization) { + McpAuthorization authorization, + AttachmentUploadLimiter uploadLimiter) { super(authorization); this.client = client; this.attachmentService = attachmentService; + this.uploadLimiter = uploadLimiter; } @Override @@ -71,27 +78,43 @@ Mono upload(Map arguments) { var policyName = requiredString(arguments, "policyName"); var groupName = optionalString(arguments, "groupName", null); var encoded = requiredString(arguments, "contentBase64"); - if (encoded.length() > 12_000_000) { + if (encoded.length() > MAX_ENCODED_CHARACTERS) { throw new McpToolException("INVALID_ARGUMENT", "contentBase64 exceeds the 8 MiB limit"); } - final byte[] bytes; - try { - bytes = Base64.getDecoder().decode(encoded); - } catch (IllegalArgumentException error) { - throw new McpToolException("INVALID_ARGUMENT", "contentBase64 is not valid Base64", error); - } - if (bytes.length == 0 || bytes.length > MAX_CONTENT_BYTES) { + var contentBytes = decodedLength(encoded); + if (contentBytes == 0 || contentBytes > MAX_CONTENT_BYTES) { throw new McpToolException( "INVALID_ARGUMENT", "Attachment content must be between 1 byte and 8 MiB"); } - var buffer = DefaultDataBufferFactory.sharedInstance.wrap(bytes); + var contentType = mediaType(arguments.get("mediaType")); + return authorization.keyId().flatMap(keyId -> Mono.using( + () -> uploadLimiter + .tryAcquire(keyId, policyName, contentBytes) + .orElseThrow(() -> new McpToolException( + "RATE_LIMITED", + "Attachment upload capacity is exhausted; retry later")), + ignored -> uploadReserved( + policyName, groupName, filename, encoded, contentType), + AttachmentUploadLimiter.Permit::close)); + } + + private Mono uploadReserved( + String policyName, + String groupName, + String filename, + String encoded, + MediaType contentType) { + var content = DataBufferUtils.readInputStream( + () -> Base64.getDecoder().wrap(new CharSequenceInputStream(encoded)), + DefaultDataBufferFactory.sharedInstance, + DECODE_BUFFER_BYTES); return attachmentService .upload( policyName, groupName, filename, - Flux.just(buffer), - mediaType(arguments.get("mediaType"))) + content, + contentType) .switchIfEmpty(Mono.error(new McpToolException( "ATTACHMENT_UNAVAILABLE", "Halo did not create the attachment"))) .map(attachment -> payload(ContentPayloads.attachment(attachment), "Uploaded attachment " + filename)); @@ -154,8 +177,12 @@ private BuiltInTool uploadTool() { "groupName", stringSchema(), "mediaType", stringSchema(), "contentBase64", - stringSchema( - "Base64-encoded file bytes; do not include a data URL prefix.")), + map( + "type", "string", + "minLength", 1, + "maxLength", MAX_ENCODED_CHARACTERS, + "description", + "Base64-encoded file bytes; do not include a data URL prefix.")), List.of("filename", "policyName", "contentBase64")), ContentPayloads.attachmentSchema(), CREATE, @@ -200,4 +227,76 @@ private static String safeFilename(String filename) { } return filename; } + + private static int decodedLength(String encoded) { + var length = encoded.length(); + var padding = 0; + while (padding < 2 + && padding < length + && encoded.charAt(length - padding - 1) == '=') { + padding++; + } + var dataLength = length - padding; + for (var index = 0; index < dataLength; index++) { + if (!isBase64Character(encoded.charAt(index))) { + throw invalidBase64(); + } + } + for (var index = dataLength; index < length; index++) { + if (encoded.charAt(index) != '=') { + throw invalidBase64(); + } + } + var remainder = dataLength % 4; + if (remainder == 1 + || (padding > 0 + && (length % 4 != 0 + || padding != (remainder == 2 ? 2 : remainder == 3 ? 1 : 0)))) { + throw invalidBase64(); + } + return (dataLength / 4) * 3 + (remainder == 2 ? 1 : remainder == 3 ? 2 : 0); + } + + private static boolean isBase64Character(char character) { + return character >= 'A' && character <= 'Z' + || character >= 'a' && character <= 'z' + || character >= '0' && character <= '9' + || character == '+' + || character == '/'; + } + + private static McpToolException invalidBase64() { + return new McpToolException("INVALID_ARGUMENT", "contentBase64 is not valid Base64"); + } + + private static final class CharSequenceInputStream extends InputStream { + + private final CharSequence source; + private int position; + + private CharSequenceInputStream(CharSequence source) { + this.source = source; + } + + @Override + public int read() { + return position == source.length() ? -1 : source.charAt(position++); + } + + @Override + public int read(byte[] bytes, int offset, int length) { + Objects.checkFromIndexSize(offset, length, bytes.length); + if (length == 0) { + return 0; + } + if (position == source.length()) { + return -1; + } + var count = Math.min(length, source.length() - position); + for (var index = 0; index < count; index++) { + bytes[offset + index] = (byte) source.charAt(position++); + } + return count; + } + } } diff --git a/src/main/java/run/halo/mcpserver/tools/AttachmentUploadLimiter.java b/src/main/java/run/halo/mcpserver/tools/AttachmentUploadLimiter.java new file mode 100644 index 0000000..8b575b5 --- /dev/null +++ b/src/main/java/run/halo/mcpserver/tools/AttachmentUploadLimiter.java @@ -0,0 +1,136 @@ +package run.halo.mcpserver.tools; + +import java.util.HashMap; +import java.util.Map; +import java.util.Objects; +import java.util.Optional; +import java.util.concurrent.atomic.AtomicBoolean; +import org.springframework.stereotype.Component; + +/** Bounds attachment uploads retained by this plugin instance. */ +@Component +final class AttachmentUploadLimiter { + + static final long GLOBAL_BYTE_LIMIT = 32L * 1024 * 1024; + static final long PER_KEY_BYTE_LIMIT = 16L * 1024 * 1024; + static final int GLOBAL_UPLOAD_LIMIT = 4; + static final int PER_KEY_UPLOAD_LIMIT = 2; + static final int PER_POLICY_UPLOAD_LIMIT = 2; + + private final long globalByteLimit; + private final long perKeyByteLimit; + private final int globalUploadLimit; + private final int perKeyUploadLimit; + private final int perPolicyUploadLimit; + private final Map activeByKey = new HashMap<>(); + private final Map activeByPolicy = new HashMap<>(); + private long activeBytes; + private int activeUploads; + + AttachmentUploadLimiter() { + this( + GLOBAL_BYTE_LIMIT, + PER_KEY_BYTE_LIMIT, + GLOBAL_UPLOAD_LIMIT, + PER_KEY_UPLOAD_LIMIT, + PER_POLICY_UPLOAD_LIMIT); + } + + AttachmentUploadLimiter( + long globalByteLimit, + long perKeyByteLimit, + int globalUploadLimit, + int perKeyUploadLimit, + int perPolicyUploadLimit) { + if (globalByteLimit < 1 + || perKeyByteLimit < 1 + || perKeyByteLimit > globalByteLimit + || globalUploadLimit < 1 + || perKeyUploadLimit < 1 + || perKeyUploadLimit > globalUploadLimit + || perPolicyUploadLimit < 1 + || perPolicyUploadLimit > globalUploadLimit) { + throw new IllegalArgumentException("Invalid attachment upload limits"); + } + this.globalByteLimit = globalByteLimit; + this.perKeyByteLimit = perKeyByteLimit; + this.globalUploadLimit = globalUploadLimit; + this.perKeyUploadLimit = perKeyUploadLimit; + this.perPolicyUploadLimit = perPolicyUploadLimit; + } + + synchronized Optional tryAcquire(String keyId, String policyName, int contentBytes) { + Objects.requireNonNull(keyId, "keyId must not be null"); + Objects.requireNonNull(policyName, "policyName must not be null"); + if (contentBytes < 1) { + throw new IllegalArgumentException("contentBytes must be positive"); + } + var keyUsage = activeByKey.getOrDefault(keyId, Usage.EMPTY); + var policyUploads = activeByPolicy.getOrDefault(policyName, 0); + if (activeBytes + contentBytes > globalByteLimit + || keyUsage.bytes + contentBytes > perKeyByteLimit + || activeUploads >= globalUploadLimit + || keyUsage.uploads >= perKeyUploadLimit + || policyUploads >= perPolicyUploadLimit) { + return Optional.empty(); + } + activeBytes += contentBytes; + activeUploads++; + activeByKey.put(keyId, new Usage(keyUsage.bytes + contentBytes, keyUsage.uploads + 1)); + activeByPolicy.put(policyName, policyUploads + 1); + return Optional.of(new Permit(this, keyId, policyName, contentBytes)); + } + + private synchronized void release(String keyId, String policyName, int contentBytes) { + var keyUsage = activeByKey.get(keyId); + var policyUploads = activeByPolicy.get(policyName); + if (keyUsage == null || policyUploads == null) { + return; + } + activeBytes -= contentBytes; + activeUploads--; + if (keyUsage.uploads == 1) { + activeByKey.remove(keyId); + } else { + activeByKey.put( + keyId, new Usage(keyUsage.bytes - contentBytes, keyUsage.uploads - 1)); + } + if (policyUploads == 1) { + activeByPolicy.remove(policyName); + } else { + activeByPolicy.put(policyName, policyUploads - 1); + } + } + + private record Usage(long bytes, int uploads) { + + private static final Usage EMPTY = new Usage(0, 0); + } + + static final class Permit implements AutoCloseable { + + private final AttachmentUploadLimiter limiter; + private final String keyId; + private final String policyName; + private final int contentBytes; + private final AtomicBoolean released = new AtomicBoolean(); + + private Permit( + AttachmentUploadLimiter limiter, + String keyId, + String policyName, + int contentBytes) { + this.limiter = limiter; + this.keyId = keyId; + this.policyName = policyName; + this.contentBytes = contentBytes; + } + + @Override + public void close() { + if (released.compareAndSet(false, true)) { + limiter.release(keyId, policyName, contentBytes); + } + } + } +} diff --git a/src/main/java/run/halo/mcpserver/tools/CategoryParentMutationCoordinator.java b/src/main/java/run/halo/mcpserver/tools/CategoryParentMutationCoordinator.java new file mode 100644 index 0000000..914c461 --- /dev/null +++ b/src/main/java/run/halo/mcpserver/tools/CategoryParentMutationCoordinator.java @@ -0,0 +1,147 @@ +package run.halo.mcpserver.tools; + +import java.time.Duration; +import java.time.Instant; +import java.util.UUID; +import java.util.function.Supplier; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; +import org.springframework.dao.DataIntegrityViolationException; +import org.springframework.dao.OptimisticLockingFailureException; +import org.springframework.stereotype.Component; +import org.springframework.util.StringUtils; +import reactor.core.Disposable; +import reactor.core.publisher.Flux; +import reactor.core.publisher.Mono; +import reactor.util.retry.Retry; +import run.halo.app.extension.ReactiveExtensionClient; +import run.halo.mcpserver.CategoryParentMutationLock; + +@Component +class CategoryParentMutationCoordinator { + + static final String LOCK_NAME = "category-parent-mutations"; + + private static final Logger log = + LoggerFactory.getLogger(CategoryParentMutationCoordinator.class); + private static final Duration LOCK_LEASE = Duration.ofMinutes(5); + private static final Duration LOCK_RENEW_INTERVAL = Duration.ofMinutes(1); + private static final Duration RETRY_DELAY = Duration.ofMillis(100); + + private final ReactiveExtensionClient client; + + CategoryParentMutationCoordinator(ReactiveExtensionClient client) { + this.client = client; + } + + Mono serialize(Supplier> mutation) { + return Mono.defer(() -> acquire(UUID.randomUUID().toString()) + .flatMap(permit -> runToCompletion(permit, mutation))); + } + + private Mono acquire(String holder) { + return getOrCreateLock().flatMap(lock -> { + var now = Instant.now(); + var spec = lock.getSpec(); + if (StringUtils.hasText(spec.getHolder()) + && spec.getExpiresAt() != null + && spec.getExpiresAt().isAfter(now)) { + return retryAcquire(holder); + } + spec.setHolder(holder); + spec.setExpiresAt(now.plus(LOCK_LEASE)); + spec.setGeneration(spec.getGeneration() + 1); + return client.update(lock) + .thenReturn(new Permit(holder)) + .onErrorResume(this::isContention, ignored -> retryAcquire(holder)); + }); + } + + private Mono retryAcquire(String holder) { + return Mono.delay(RETRY_DELAY).then(Mono.defer(() -> acquire(holder))); + } + + private Mono getOrCreateLock() { + return client.fetch(CategoryParentMutationLock.class, LOCK_NAME) + .switchIfEmpty(Mono.defer(() -> client.create(newLock()) + .onErrorResume(this::isContention, ignored -> Mono.empty()) + .then(Mono.defer(() -> client.fetch( + CategoryParentMutationLock.class, LOCK_NAME))) + .switchIfEmpty(Mono.delay(RETRY_DELAY) + .then(Mono.defer(this::getOrCreateLock))))); + } + + private Mono runToCompletion(Permit permit, Supplier> mutation) { + return Mono.usingWhen( + Mono.just(permit), + ignored -> Mono.using( + () -> renewLease(permit), + renewal -> Mono.defer(mutation), + Disposable::dispose), + this::release, + (ignored, error) -> release(permit), + this::release) + .cache(); + } + + private Disposable renewLease(Permit permit) { + return Flux.interval(LOCK_RENEW_INTERVAL) + .concatMap(ignored -> extendLease(permit) + .doOnError(error -> log.warn( + "Failed to renew the category parent mutation lock", + error)) + .onErrorComplete()) + .subscribe(); + } + + private Mono extendLease(Permit permit) { + return Mono.defer(() -> client.fetch(CategoryParentMutationLock.class, LOCK_NAME) + .flatMap(lock -> { + var spec = lock.getSpec(); + if (!permit.holder().equals(spec.getHolder())) { + return Mono.empty(); + } + spec.setExpiresAt(Instant.now().plus(LOCK_LEASE)); + return client.update(lock).then(); + })) + .retryWhen(contentionRetry()); + } + + private Mono release(Permit permit) { + return Mono.defer(() -> client.fetch(CategoryParentMutationLock.class, LOCK_NAME) + .flatMap(lock -> { + var spec = lock.getSpec(); + if (!permit.holder().equals(spec.getHolder())) { + return Mono.empty(); + } + spec.setHolder(null); + spec.setExpiresAt(null); + return client.update(lock).then(); + })) + .retryWhen(contentionRetry()) + .doOnError(error -> log.warn( + "Failed to release the category parent mutation lock; its lease will expire", + error)) + .onErrorComplete(); + } + + private Retry contentionRetry() { + return Retry.backoff(3, Duration.ofMillis(25)) + .maxBackoff(Duration.ofMillis(100)) + .filter(this::isContention) + .onRetryExhaustedThrow((spec, signal) -> signal.failure()); + } + + private boolean isContention(Throwable error) { + return error instanceof OptimisticLockingFailureException + || error instanceof DataIntegrityViolationException; + } + + private static CategoryParentMutationLock newLock() { + var lock = new CategoryParentMutationLock(); + lock.setMetadata(ToolSupport.metadata(LOCK_NAME)); + return lock; + } + + private record Permit(String holder) {} +} diff --git a/src/main/java/run/halo/mcpserver/tools/CategoryTools.java b/src/main/java/run/halo/mcpserver/tools/CategoryTools.java index 9ca486f..41b6934 100644 --- a/src/main/java/run/halo/mcpserver/tools/CategoryTools.java +++ b/src/main/java/run/halo/mcpserver/tools/CategoryTools.java @@ -8,6 +8,7 @@ import java.util.HashSet; import java.util.List; import java.util.Map; +import java.util.function.Supplier; import org.springframework.data.domain.Sort; import org.springframework.stereotype.Component; import org.springframework.util.StringUtils; @@ -28,10 +29,15 @@ class CategoryTools extends ToolSupport implements ToolGroup { static final String UPDATE = "halo_update_category"; private final ReactiveExtensionClient client; + private final CategoryParentMutationCoordinator parentMutationCoordinator; - CategoryTools(ReactiveExtensionClient client, McpAuthorization authorization) { + CategoryTools( + ReactiveExtensionClient client, + McpAuthorization authorization, + CategoryParentMutationCoordinator parentMutationCoordinator) { super(authorization); this.client = client; + this.parentMutationCoordinator = parentMutationCoordinator; } @Override @@ -79,28 +85,35 @@ Mono list(Map arguments) { Mono create(Map arguments) { var name = resourceName(arguments, "name"); var parent = parent(arguments); - return validateParent(name, parent).then(Mono.defer(() -> { - var category = new Category(); - category.setMetadata(metadata(name)); - category.setSpec(new Category.CategorySpec()); - apply(category.getSpec(), arguments, name, true); - return client.create(category); - })).map(category -> payload(ContentPayloads.category(category), "Created category " + name)); + Supplier> create = () -> validateParent(name, parent) + .then(Mono.defer(() -> { + var category = new Category(); + category.setMetadata(metadata(name)); + category.setSpec(new Category.CategorySpec()); + apply(category.getSpec(), arguments, name, true); + return client.create(category); + })) + .map(category -> payload(ContentPayloads.category(category), "Created category " + name)); + return parent == null ? create.get() : parentMutationCoordinator.serialize(create); } Mono update(Map arguments) { var name = resourceName(arguments, "name"); var expectedVersion = optionalLong(arguments, "expectedVersion"); var parent = arguments.containsKey("parent") ? parent(arguments) : null; - var validation = arguments.containsKey("parent") ? validateParent(name, parent) : Mono.empty(); - return validation.then(Mono.defer(() -> client.fetch(Category.class, name) - .switchIfEmpty(notFound("Category", name)) - .flatMap(category -> checkVersion(category.getMetadata().getVersion(), expectedVersion) - .then(Mono.defer(() -> { - apply(category.getSpec(), arguments, name, false); - return client.update(category); - }))))) + Supplier> update = () -> client.fetch(Category.class, name) + .switchIfEmpty(notFound("Category", name)) + .flatMap(category -> checkVersion(category.getMetadata().getVersion(), expectedVersion) + .then(Mono.defer(() -> { + apply(category.getSpec(), arguments, name, false); + return client.update(category); + }))) .map(category -> payload(ContentPayloads.category(category), "Updated category " + name)); + if (!arguments.containsKey("parent")) { + return update.get(); + } + return parentMutationCoordinator.serialize( + () -> validateParent(name, parent).then(Mono.defer(update))); } private BuiltInTool createTool() { diff --git a/src/test/java/run/halo/mcpserver/HaloMcpServerTest.java b/src/test/java/run/halo/mcpserver/HaloMcpServerTest.java index c0343d7..b97f56f 100644 --- a/src/test/java/run/halo/mcpserver/HaloMcpServerTest.java +++ b/src/test/java/run/halo/mcpserver/HaloMcpServerTest.java @@ -3,7 +3,9 @@ import static org.mockito.Mockito.lenient; import static org.mockito.Mockito.when; +import java.net.InetSocketAddress; import java.time.Duration; +import java.util.concurrent.atomic.AtomicBoolean; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; @@ -11,10 +13,14 @@ import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpStatus; import org.springframework.http.MediaType; +import org.springframework.mock.http.server.reactive.MockServerHttpRequest; +import org.springframework.mock.web.server.MockServerWebExchange; import org.springframework.test.web.reactive.server.WebTestClient; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import reactor.test.StepVerifier; import run.halo.app.plugin.extensionpoint.ExtensionGetter; import run.halo.app.extension.ReactiveExtensionClient; import run.halo.app.plugin.PluginContext; @@ -42,6 +48,9 @@ class HaloMcpServerTest { @Mock BuiltInTools builtInTools; + @Mock + McpAccessKeyService accessKeyService; + HaloMcpServer server; McpRecentCallHistory recentCallHistory; WebTestClient client; @@ -346,6 +355,99 @@ void keepsBuiltInToolsWhenProviderDiscoveryFails() { .assertThat(body.toString()).doesNotContain("storage password")); } + @Test + void timesOutAndCancelsAStalledProviderAndReleasesItsPermit() { + var providerCancelled = new AtomicBoolean(); + var definition = McpToolDefinition.builder() + .name("demo/stall") + .title("Stall") + .description("Test timeout cancellation") + .inputSchema(java.util.Map.of("type", "object", "properties", java.util.Map.of())) + .permission(invocation -> Mono.just(true)) + .handler(invocation -> Mono.never() + .doOnCancel(() -> providerCancelled.set(true))) + .build(); + when(extensionGetter.getEnabledExtensions(McpToolProvider.class)) + .thenReturn(Flux.just(provider)); + when(provider.tools()).thenReturn(Flux.just(definition)); + var plugin = new run.halo.app.core.extension.Plugin(); + var metadata = new run.halo.app.extension.Metadata(); + metadata.setName("demo"); + plugin.setMetadata(metadata); + var pluginStatus = new run.halo.app.core.extension.Plugin.PluginStatus(); + try { + pluginStatus.setLoadLocation(provider.getClass() + .getProtectionDomain() + .getCodeSource() + .getLocation() + .toURI() + .normalize()); + } catch (java.net.URISyntaxException error) { + throw new AssertionError(error); + } + plugin.setStatus(pluginStatus); + when(extensionClient.fetch(run.halo.app.core.extension.Plugin.class, "demo")) + .thenReturn(Mono.just(plugin)); + + var rawToken = "hmcp_00000000-0000-0000-0000-000000000000_secret"; + var remoteAddress = new InetSocketAddress("203.0.113.8", 41321); + var authentication = new McpKeyAuthenticationToken( + "key-id", + "Automation", + "hmcp_key", + "admin", + java.util.Set.of("demo/stall")); + when(accessKeyService.authenticate(rawToken, remoteAddress)) + .thenReturn(Mono.just(authentication)); + var concurrencyLimiter = new McpRequestConcurrencyLimiter(1, 1); + var filter = new McpKeyAuthenticationFilter( + accessKeyService, rateLimiter, concurrencyLimiter, server); + var stalledExchange = MockServerWebExchange.from(MockServerHttpRequest.post("/mcp") + .remoteAddress(remoteAddress) + .contentType(MediaType.APPLICATION_JSON) + .header(HttpHeaders.ACCEPT, "application/json, text/event-stream") + .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken) + .body(""" + { + "jsonrpc":"2.0", + "id":11, + "method":"tools/call", + "params":{"name":"demo/stall","arguments":{}} + } + """)); + + StepVerifier.withVirtualTime(() -> filter.filter( + stalledExchange, + ignored -> Mono.error(new AssertionError("Halo chain must not continue")))) + .thenAwait(McpKeyAuthenticationFilter.REQUEST_TIMEOUT.plusMillis(1)) + .verifyComplete(); + + org.assertj.core.api.Assertions.assertThat(stalledExchange.getResponse().getStatusCode()) + .isEqualTo(HttpStatus.GATEWAY_TIMEOUT); + org.assertj.core.api.Assertions.assertThat(providerCancelled).isTrue(); + var cancelledCall = recentCallHistory.list( + new McpRecentCallQuery(1, 20, "key-id", "demo/stall", null)); + org.assertj.core.api.Assertions.assertThat(cancelledCall.total()).isEqualTo(1); + org.assertj.core.api.Assertions.assertThat(cancelledCall.items().getFirst().outcome()) + .isEqualTo(McpCallOutcome.CANCELLED); + + var pingExchange = MockServerWebExchange.from(MockServerHttpRequest.post("/mcp") + .remoteAddress(remoteAddress) + .contentType(MediaType.APPLICATION_JSON) + .header(HttpHeaders.ACCEPT, "application/json, text/event-stream") + .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken) + .body(""" + {"jsonrpc":"2.0","id":12,"method":"ping","params":{}} + """)); + filter.filter( + pingExchange, + ignored -> Mono.error(new AssertionError("Halo chain must not continue"))) + .block(); + + org.assertj.core.api.Assertions.assertThat(pingExchange.getResponse().getStatusCode()) + .isEqualTo(HttpStatus.OK); + } + private static BuiltInTool builtInTool(String name, String title) { var tool = io.modelcontextprotocol.spec.McpSchema.Tool.builder( name, diff --git a/src/test/java/run/halo/mcpserver/McpAccessKeyServiceTest.java b/src/test/java/run/halo/mcpserver/McpAccessKeyServiceTest.java index db2dd17..2c88c84 100644 --- a/src/test/java/run/halo/mcpserver/McpAccessKeyServiceTest.java +++ b/src/test/java/run/halo/mcpserver/McpAccessKeyServiceTest.java @@ -3,17 +3,21 @@ import static org.assertj.core.api.Assertions.assertThat; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.times; 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 java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.dao.OptimisticLockingFailureException; import reactor.core.publisher.Mono; import run.halo.app.extension.ReactiveExtensionClient; @@ -130,4 +134,176 @@ void authenticatesOnlyFromAnAllowedIpRange() { .isNotNull(); assertThat(created.accessKey().getStatus().getLastUsedAt()).isNotNull(); } + + @Test + void rejectsAuthenticationWhenToolScopesChangeAfterInitialFetch() { + var created = service.create( + "Automation", + "admin", + Set.of("halo_delete_post"), + Set.of(), + null) + .block(); + var initial = created.accessKey(); + initial.getMetadata().setVersion(1L); + var restricted = copyOf(initial); + restricted.getMetadata().setVersion(2L); + restricted.getSpec().setAllowedTools(Set.of("halo_search_content")); + var stored = new AtomicReference(initial); + var mutationCommitted = new AtomicBoolean(); + when(client.fetch(McpAccessKey.class, initial.getMetadata().getName())) + .thenAnswer(ignored -> Mono.defer(() -> Mono.justOrEmpty(stored.get())) + .doOnNext(snapshot -> { + if (mutationCommitted.compareAndSet(false, true)) { + stored.set(restricted); + } + })); + + var authentication = service.authenticate(created.token(), null).block(); + + assertThat(authentication).isNull(); + assertThat(mutationCommitted).isTrue(); + } + + @Test + void rejectsAuthenticationWhenOtherSecurityStateChangesAfterInitialFetch() { + assertMutationRejected("rotated hash", key -> key.getSpec().setKeyHash("rotated-hash")); + assertMutationRejected("disabled key", key -> key.getSpec().setEnabled(false)); + assertMutationRejected( + "expired key", key -> key.getSpec().setExpiresAt(Instant.now().minusSeconds(1))); + assertMutationRejected( + "restricted IP", + key -> key.getSpec().setAllowedIpRanges(Set.of("198.51.100.0/24"))); + assertMutationRejected("changed owner", key -> key.getSpec().setOwnerName("other-user")); + assertStoredChangeRejected("deleted key", null); + } + + @Test + void authenticatesWhenOnlyStatusAndResourceVersionChangeAfterInitialFetch() { + var created = newCreatedKey(); + var initial = created.accessKey(); + initial.getMetadata().setVersion(1L); + var current = copyOf(initial); + current.getMetadata().setVersion(2L); + current.getStatus().setLastUsedAt(Instant.now().minusSeconds(30)); + stubMutationAfterInitialFetch(initial, current); + + var authentication = service.authenticate( + created.token(), new InetSocketAddress("203.0.113.42", 443)) + .block(); + + assertThat(authentication).isNotNull(); + assertThat(authentication.allows("halo_delete_post")).isTrue(); + } + + @Test + void revalidatesWhenLastUsedWriteConflictsWithARestriction() { + var created = newCreatedKey(); + var initial = created.accessKey(); + initial.getMetadata().setVersion(1L); + var restricted = copyOf(initial); + restricted.getMetadata().setVersion(2L); + restricted.getSpec().setAllowedTools(Set.of("halo_search_content")); + var id = initial.getMetadata().getName(); + when(client.fetch(McpAccessKey.class, id)) + .thenReturn(Mono.just(initial), Mono.just(restricted)); + when(client.update(initial)) + .thenReturn(Mono.error(new OptimisticLockingFailureException("version changed"))); + + var authentication = service.authenticate(created.token(), null).block(); + + assertThat(authentication).isNull(); + verify(client, times(2)).fetch(McpAccessKey.class, id); + } + + @Test + void keepsLastUsedWritesBestEffortWhenThereIsNoVersionConflict() { + var created = newCreatedKey(); + var initial = created.accessKey(); + var validated = copyOf(initial); + var id = initial.getMetadata().getName(); + when(client.fetch(McpAccessKey.class, id)) + .thenReturn(Mono.just(initial), Mono.just(validated)); + when(client.update(initial)).thenReturn(Mono.error(new IllegalStateException("unavailable"))); + + var authentication = service.authenticate(created.token(), null).block(); + + assertThat(authentication).isNotNull(); + assertThat(authentication.allows("halo_delete_post")).isTrue(); + } + + private void assertMutationRejected( + String description, java.util.function.Consumer mutation) { + var created = newCreatedKey(); + var initial = created.accessKey(); + initial.getMetadata().setVersion(1L); + var changed = copyOf(initial); + changed.getMetadata().setVersion(2L); + mutation.accept(changed); + stubMutationAfterInitialFetch(initial, changed); + + var authentication = service.authenticate( + created.token(), new InetSocketAddress("203.0.113.42", 443)) + .block(); + + assertThat(authentication).as(description).isNull(); + } + + private void assertStoredChangeRejected(String description, McpAccessKey changed) { + var created = newCreatedKey(); + var initial = created.accessKey(); + initial.getMetadata().setVersion(1L); + stubMutationAfterInitialFetch(initial, changed); + + var authentication = service.authenticate( + created.token(), new InetSocketAddress("203.0.113.42", 443)) + .block(); + + assertThat(authentication).as(description).isNull(); + } + + private McpAccessKeyService.CreatedKey newCreatedKey() { + return service.create( + "Automation", + "admin", + Set.of("halo_delete_post"), + Set.of(), + null) + .block(); + } + + private void stubMutationAfterInitialFetch(McpAccessKey initial, McpAccessKey changed) { + var stored = new AtomicReference<>(initial); + var mutationCommitted = new AtomicBoolean(); + when(client.fetch(McpAccessKey.class, initial.getMetadata().getName())) + .thenAnswer(ignored -> Mono.defer(() -> Mono.justOrEmpty(stored.get())) + .doOnNext(snapshot -> { + if (mutationCommitted.compareAndSet(false, true)) { + stored.set(changed); + } + })); + } + + private static McpAccessKey copyOf(McpAccessKey source) { + var copy = new McpAccessKey(); + var metadata = new run.halo.app.extension.Metadata(); + metadata.setName(source.getMetadata().getName()); + metadata.setVersion(source.getMetadata().getVersion()); + metadata.setDeletionTimestamp(source.getMetadata().getDeletionTimestamp()); + copy.setMetadata(metadata); + var spec = new McpAccessKey.Spec(); + spec.setDisplayName(source.getSpec().getDisplayName()); + spec.setKeyHash(source.getSpec().getKeyHash()); + spec.setKeyPrefix(source.getSpec().getKeyPrefix()); + spec.setOwnerName(source.getSpec().getOwnerName()); + spec.setEnabled(source.getSpec().isEnabled()); + spec.setExpiresAt(source.getSpec().getExpiresAt()); + spec.setAllowedTools(Set.copyOf(source.getSpec().getAllowedTools())); + spec.setAllowedIpRanges(Set.copyOf(source.getSpec().getAllowedIpRanges())); + copy.setSpec(spec); + var status = new McpAccessKey.Status(); + status.setLastUsedAt(source.getStatus() == null ? null : source.getStatus().getLastUsedAt()); + copy.setStatus(status); + return copy; + } } diff --git a/src/test/java/run/halo/mcpserver/McpKeyAuthenticationFilterTest.java b/src/test/java/run/halo/mcpserver/McpKeyAuthenticationFilterTest.java index 321c618..532bab2 100644 --- a/src/test/java/run/halo/mcpserver/McpKeyAuthenticationFilterTest.java +++ b/src/test/java/run/halo/mcpserver/McpKeyAuthenticationFilterTest.java @@ -7,12 +7,14 @@ import java.net.InetSocketAddress; import java.util.Set; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicReference; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.core.io.buffer.DataBuffer; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpStatus; import org.springframework.mock.http.server.reactive.MockServerHttpRequest; @@ -21,11 +23,16 @@ import org.springframework.web.reactive.function.server.RouterFunctions; import org.springframework.web.reactive.function.server.ServerResponse; import org.springframework.web.server.ServerWebExchange; +import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import reactor.test.StepVerifier; @ExtendWith(MockitoExtension.class) class McpKeyAuthenticationFilterTest { + private static final InetSocketAddress REMOTE_ADDRESS = + new InetSocketAddress("203.0.113.8", 41321); + @Mock McpAccessKeyService accessKeyService; @@ -34,6 +41,7 @@ class McpKeyAuthenticationFilterTest { McpKeyAuthenticationFilter filter; McpRequestRateLimiter rateLimiter; + McpRequestConcurrencyLimiter concurrencyLimiter; AtomicReference handledPath; AtomicReference handledAuthorization; AtomicReference currentAuthentication; @@ -56,7 +64,9 @@ void setUp() { .build()); when(mcpServer.protocolVersions()).thenReturn(java.util.List.of("2025-11-25")); rateLimiter = new McpRequestRateLimiter(); - filter = new McpKeyAuthenticationFilter(accessKeyService, rateLimiter, mcpServer); + concurrencyLimiter = new McpRequestConcurrencyLimiter(); + filter = new McpKeyAuthenticationFilter( + accessKeyService, rateLimiter, concurrencyLimiter, mcpServer); } @Test @@ -87,8 +97,9 @@ void authenticatesAndStripsTheMcpBearerTokenBeforeTheHaloJwtFilter() { @Test void rejectsAnInvalidMcpKey() { var rawToken = "hmcp_00000000-0000-0000-0000-000000000000_invalid"; - when(accessKeyService.authenticate(rawToken, null)).thenReturn(Mono.empty()); + when(accessKeyService.authenticate(rawToken, REMOTE_ADDRESS)).thenReturn(Mono.empty()); var exchange = MockServerWebExchange.from(MockServerHttpRequest.post(McpKeyAuthenticationFilter.MCP_PATH) + .remoteAddress(REMOTE_ADDRESS) .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken)); filter.filter(exchange, ignored -> Mono.error(new AssertionError("Request must not continue"))) @@ -139,14 +150,16 @@ void authenticatesATrailingSlashAsTheSameMcpEndpoint() { "hmcp_00000000", "admin", Set.of()); - when(accessKeyService.authenticate(rawToken, null)).thenReturn(Mono.just(authentication)); + when(accessKeyService.authenticate(rawToken, REMOTE_ADDRESS)) + .thenReturn(Mono.just(authentication)); var exchange = MockServerWebExchange.from(MockServerHttpRequest.post("/mcp/") + .remoteAddress(REMOTE_ADDRESS) .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken)); filter.filter(exchange, ignored -> Mono.error(new AssertionError("Halo chain must not continue"))) .block(); - verify(accessKeyService).authenticate(rawToken, null); + verify(accessKeyService).authenticate(rawToken, REMOTE_ADDRESS); assertThat(handledPath.get()).isEqualTo("/mcp/"); } @@ -159,8 +172,10 @@ void rejectsAnUnsupportedProtocolVersionAfterAuthentication() { "hmcp_00000000", "admin", Set.of("halo_search_content")); - when(accessKeyService.authenticate(rawToken, null)).thenReturn(Mono.just(authentication)); + when(accessKeyService.authenticate(rawToken, REMOTE_ADDRESS)) + .thenReturn(Mono.just(authentication)); var exchange = MockServerWebExchange.from(MockServerHttpRequest.post(McpKeyAuthenticationFilter.MCP_PATH) + .remoteAddress(REMOTE_ADDRESS) .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken) .header(io.modelcontextprotocol.spec.HttpHeaders.PROTOCOL_VERSION, "2099-01-01")); @@ -173,10 +188,11 @@ void rejectsAnUnsupportedProtocolVersionAfterAuthentication() { @Test void rateLimitsRequestsBeforeAccessKeyLookup() { for (var i = 0; i < McpRequestRateLimiter.REQUESTS_PER_MINUTE; i++) { - assertThat(rateLimiter.allowRequest(null)).isTrue(); + assertThat(rateLimiter.allowRequest(REMOTE_ADDRESS)).isTrue(); } var rawToken = "hmcp_00000000-0000-0000-0000-000000000000_secret"; var exchange = MockServerWebExchange.from(MockServerHttpRequest.post(McpKeyAuthenticationFilter.MCP_PATH) + .remoteAddress(REMOTE_ADDRESS) .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken)); filter.filter(exchange, ignored -> Mono.error(new AssertionError("Request must not continue"))) @@ -185,6 +201,95 @@ 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, null); + verify(accessKeyService, never()).authenticate(rawToken, REMOTE_ADDRESS); + } + + @Test + void timesOutAndCancelsAnIncompleteBodyAndReleasesItsPermit() { + var rawToken = "hmcp_00000000-0000-0000-0000-000000000000_secret"; + var authentication = new McpKeyAuthenticationToken( + "key-id", "Automation", "hmcp_00000000", "admin", Set.of()); + when(accessKeyService.authenticate(rawToken, REMOTE_ADDRESS)) + .thenReturn(Mono.just(authentication)); + var bodyCancelled = new AtomicBoolean(); + org.springframework.web.reactive.function.server.HandlerFunction handler = + request -> request.bodyToMono(String.class) + .then(ServerResponse.noContent().build()); + when(mcpServer.routerFunction()).thenReturn(RouterFunctions.route() + .POST("/mcp", handler) + .build()); + concurrencyLimiter = new McpRequestConcurrencyLimiter(1, 1); + filter = new McpKeyAuthenticationFilter( + accessKeyService, rateLimiter, concurrencyLimiter, mcpServer); + var stalledExchange = MockServerWebExchange.from(MockServerHttpRequest.post( + McpKeyAuthenticationFilter.MCP_PATH) + .remoteAddress(REMOTE_ADDRESS) + .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken) + .body(Flux.never() + .doOnCancel(() -> bodyCancelled.set(true)))); + + StepVerifier.withVirtualTime(() -> filter.filter( + stalledExchange, + ignored -> Mono.error(new AssertionError("Halo chain must not continue")))) + .thenAwait(McpKeyAuthenticationFilter.REQUEST_TIMEOUT.plusMillis(1)) + .verifyComplete(); + + assertThat(stalledExchange.getResponse().getStatusCode()) + .isEqualTo(HttpStatus.GATEWAY_TIMEOUT); + assertThat(bodyCancelled).isTrue(); + + var completedExchange = MockServerWebExchange.from(MockServerHttpRequest.post( + McpKeyAuthenticationFilter.MCP_PATH) + .remoteAddress(REMOTE_ADDRESS) + .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken) + .body("complete")); + filter.filter( + completedExchange, + ignored -> Mono.error(new AssertionError("Halo chain must not continue"))) + .block(); + + assertThat(completedExchange.getResponse().getStatusCode()) + .isEqualTo(HttpStatus.NO_CONTENT); + } + + @Test + void rejectsPerKeyConcurrencyBeforeInvokingTheHandler() { + assertCapacityRejection(new McpRequestConcurrencyLimiter(2, 1), "key-id"); + } + + @Test + void rejectsGlobalConcurrencyBeforeInvokingTheHandler() { + assertCapacityRejection(new McpRequestConcurrencyLimiter(1, 1), "other-key-id"); + } + + private void assertCapacityRejection( + McpRequestConcurrencyLimiter limiter, String occupiedKeyId) { + var rawToken = "hmcp_00000000-0000-0000-0000-000000000000_secret"; + var authentication = new McpKeyAuthenticationToken( + "key-id", "Automation", "hmcp_00000000", "admin", Set.of()); + when(accessKeyService.authenticate(rawToken, REMOTE_ADDRESS)) + .thenReturn(Mono.just(authentication)); + concurrencyLimiter = limiter; + filter = new McpKeyAuthenticationFilter( + accessKeyService, rateLimiter, concurrencyLimiter, mcpServer); + var occupied = concurrencyLimiter.tryAcquire(occupiedKeyId).orElseThrow(); + var exchange = MockServerWebExchange.from(MockServerHttpRequest.post( + McpKeyAuthenticationFilter.MCP_PATH) + .remoteAddress(REMOTE_ADDRESS) + .header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken)); + try { + filter.filter( + exchange, + ignored -> Mono.error(new AssertionError("Halo chain must not continue"))) + .block(); + } finally { + occupied.close(); + } + + assertThat(exchange.getResponse().getStatusCode()) + .isEqualTo(HttpStatus.TOO_MANY_REQUESTS); + assertThat(exchange.getResponse().getHeaders().getFirst(HttpHeaders.RETRY_AFTER)) + .isEqualTo("1"); + assertThat(handledPath).hasNullValue(); } } diff --git a/src/test/java/run/halo/mcpserver/McpRequestConcurrencyLimiterTest.java b/src/test/java/run/halo/mcpserver/McpRequestConcurrencyLimiterTest.java new file mode 100644 index 0000000..ae5dc4e --- /dev/null +++ b/src/test/java/run/halo/mcpserver/McpRequestConcurrencyLimiterTest.java @@ -0,0 +1,28 @@ +package run.halo.mcpserver; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.Test; + +class McpRequestConcurrencyLimiterTest { + + @Test + void enforcesPerKeyAndGlobalLimitsAndReleasesPermits() { + var limiter = new McpRequestConcurrencyLimiter(2, 1); + var first = limiter.tryAcquire("key-one").orElseThrow(); + + assertThat(limiter.tryAcquire("key-one")).isEmpty(); + + var second = limiter.tryAcquire("key-two").orElseThrow(); + assertThat(limiter.tryAcquire("key-three")).isEmpty(); + + first.close(); + first.close(); + var replacement = limiter.tryAcquire("key-three").orElseThrow(); + + second.close(); + replacement.close(); + var finalPermit = limiter.tryAcquire("key-one").orElseThrow(); + finalPermit.close(); + } +} diff --git a/src/test/java/run/halo/mcpserver/McpRequestRateLimiterTest.java b/src/test/java/run/halo/mcpserver/McpRequestRateLimiterTest.java index 6cc3340..3ac4747 100644 --- a/src/test/java/run/halo/mcpserver/McpRequestRateLimiterTest.java +++ b/src/test/java/run/halo/mcpserver/McpRequestRateLimiterTest.java @@ -5,6 +5,8 @@ import java.net.InetSocketAddress; import java.util.concurrent.atomic.AtomicLong; import org.junit.jupiter.api.Test; +import org.springframework.mock.http.server.reactive.MockServerHttpRequest; +import org.springframework.web.server.adapter.ForwardedHeaderTransformer; class McpRequestRateLimiterTest { @@ -23,6 +25,47 @@ void limitsRequestsByRemoteAddressAndResetsAfterAWindow() { assertThat(limiter.allowRequest(remote)).isTrue(); } + @Test + void isolatesNumericAddressesFromForwardedHeaders() { + var limiter = new McpRequestRateLimiter(() -> 0L); + var firstClient = transformedRemoteAddress("X-Forwarded-For", "203.0.113.42"); + var secondClient = transformedRemoteAddress( + "Forwarded", "for=\"[2001:db8::9]\""); + assertThat(firstClient.isUnresolved()).isTrue(); + assertThat(secondClient.isUnresolved()).isTrue(); + + for (var i = 0; i < McpRequestRateLimiter.REQUESTS_PER_MINUTE; i++) { + assertThat(limiter.allowRequest(firstClient)).isTrue(); + } + assertThat(limiter.allowRequest(firstClient)).isFalse(); + assertThat(limiter.allowRequest(secondClient)).isTrue(); + } + + @Test + void acceptsForwardedAddressesWithPorts() { + var limiter = new McpRequestRateLimiter(() -> 0L); + + assertThat(limiter.allowRequest(transformedRemoteAddress( + "X-Forwarded-For", "203.0.113.42:4567"))) + .isTrue(); + assertThat(limiter.allowRequest(transformedRemoteAddress( + "Forwarded", "for=\"[2001:db8::42]:4567\""))) + .isTrue(); + } + + @Test + void rejectsUnavailableAndNonnumericAddresses() { + var limiter = new McpRequestRateLimiter(() -> 0L); + + assertThat(limiter.allowRequest(null)).isFalse(); + assertThat(limiter.allowRequest( + InetSocketAddress.createUnresolved("client.example.com", 443))) + .isFalse(); + assertThat(limiter.allowRequest(transformedRemoteAddress( + "X-Forwarded-For", "203.0.113.42:not-a-port"))) + .isFalse(); + } + @Test void isolatesToolLimitsByKeyAndToolAndCanBeCleared() { var limiter = new McpRequestRateLimiter(() -> 0L); @@ -37,4 +80,12 @@ void isolatesToolLimitsByKeyAndToolAndCanBeCleared() { limiter.clear(); assertThat(limiter.allowTool("key-one", "demo/one")).isTrue(); } + + private static InetSocketAddress transformedRemoteAddress(String header, String value) { + var request = MockServerHttpRequest.get("http://localhost/mcp") + .remoteAddress(new InetSocketAddress("192.0.2.10", 443)) + .header(header, value) + .build(); + return new ForwardedHeaderTransformer().apply(request).getRemoteAddress(); + } } diff --git a/src/test/java/run/halo/mcpserver/McpServerPluginTest.java b/src/test/java/run/halo/mcpserver/McpServerPluginTest.java index 5ede4f6..837d654 100644 --- a/src/test/java/run/halo/mcpserver/McpServerPluginTest.java +++ b/src/test/java/run/halo/mcpserver/McpServerPluginTest.java @@ -1,16 +1,18 @@ package run.halo.mcpserver; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + 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 reactor.core.publisher.Mono; +import run.halo.app.extension.Scheme; import run.halo.app.extension.SchemeManager; import run.halo.app.plugin.PluginContext; -import static org.mockito.Mockito.when; - @ExtendWith(MockitoExtension.class) class McpServerPluginTest { @@ -34,5 +36,10 @@ void contextLoads() { when(mcpServer.closeGracefully()).thenReturn(Mono.empty()); plugin.start(); plugin.stop(); + + verify(schemeManager).register(McpAccessKey.class); + verify(schemeManager).register(CategoryParentMutationLock.class); + verify(schemeManager).unregister(Scheme.buildFromType(CategoryParentMutationLock.class)); + verify(schemeManager).unregister(Scheme.buildFromType(McpAccessKey.class)); } } diff --git a/src/test/java/run/halo/mcpserver/McpToolRegistryTest.java b/src/test/java/run/halo/mcpserver/McpToolRegistryTest.java index c9c62af..d3546c3 100644 --- a/src/test/java/run/halo/mcpserver/McpToolRegistryTest.java +++ b/src/test/java/run/halo/mcpserver/McpToolRegistryTest.java @@ -6,6 +6,7 @@ import java.util.List; import java.util.Map; import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; @@ -170,6 +171,47 @@ void rejectsToolWhenKeyOrProviderPermissionDenies() { .verifyComplete(); } + @Test + void invokesProviderCallbacksOnlyAfterKeyAllowlistSucceeds() { + var permissionConstructed = new AtomicInteger(); + var permissionSubscribed = new AtomicInteger(); + var handlerConstructed = new AtomicInteger(); + var tool = McpToolDefinition.builder() + .name("demo/secret") + .inputSchema(objectSchema(Map.of(), List.of())) + .permission(invocation -> { + permissionConstructed.incrementAndGet(); + return Mono.fromSupplier(() -> { + permissionSubscribed.incrementAndGet(); + return true; + }); + }) + .handler(invocation -> { + handlerConstructed.incrementAndGet(); + return Mono.just(McpToolResult.success(Map.of("ok", true))); + }) + .build(); + providerTools(provider, "demo", tool); + + StepVerifier.create(registry.executeIfContributed("demo/secret", Map.of()) + .contextWrite(context("demo/other"))) + .assertNext(result -> assertThat(result.orElseThrow().structuredContent().toString()) + .contains("FORBIDDEN")) + .verifyComplete(); + assertThat(permissionConstructed).hasValue(0); + assertThat(permissionSubscribed).hasValue(0); + assertThat(handlerConstructed).hasValue(0); + + StepVerifier.create(registry.executeIfContributed("demo/secret", Map.of()) + .contextWrite(context("demo/secret"))) + .assertNext(result -> assertThat(result.orElseThrow().structuredContent().toString()) + .contains("ok=true")) + .verifyComplete(); + assertThat(permissionConstructed).hasValue(1); + assertThat(permissionSubscribed).hasValue(1); + assertThat(handlerConstructed).hasValue(1); + } + @Test void resolvesProviderListForEachRequest() { var current = new java.util.concurrent.atomic.AtomicReference(tool("demo/one")); diff --git a/src/test/java/run/halo/mcpserver/tools/AttachmentToolsTest.java b/src/test/java/run/halo/mcpserver/tools/AttachmentToolsTest.java index ee72f11..74ffebf 100644 --- a/src/test/java/run/halo/mcpserver/tools/AttachmentToolsTest.java +++ b/src/test/java/run/halo/mcpserver/tools/AttachmentToolsTest.java @@ -2,15 +2,28 @@ import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.nullable; 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 java.util.Base64; import java.util.Map; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicLong; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentMatchers; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferUtils; +import org.springframework.http.MediaType; +import reactor.core.Disposable; +import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; import reactor.test.StepVerifier; import run.halo.app.core.extension.attachment.Attachment; @@ -33,7 +46,8 @@ class AttachmentToolsTest { @Test void validatesAttachmentBase64BeforeUpload() { - var tools = new AttachmentTools(client, attachmentService, authorization); + var tools = new AttachmentTools( + client, attachmentService, authorization, new AttachmentUploadLimiter()); assertThatThrownBy(() -> tools.upload(Map.of( "filename", "a.txt", @@ -41,6 +55,177 @@ void validatesAttachmentBase64BeforeUpload() { "contentBase64", "not-base64"))) .isInstanceOf(McpToolException.class) .hasMessageContaining("valid Base64"); + + assertThatThrownBy(() -> tools.upload(Map.of( + "filename", "a.txt", + "policyName", "local", + "contentBase64", "A".repeat(11_184_813)))) + .isInstanceOf(McpToolException.class) + .hasMessageContaining("8 MiB"); + + verify(attachmentService, never()) + .upload( + anyString(), + nullable(String.class), + anyString(), + ArgumentMatchers.>any(), + any(MediaType.class)); + } + + @Test + void limitsUploadsRetainedByAStalledAttachmentBackend() { + when(authorization.keyId()).thenReturn(Mono.just("key-one")); + when(attachmentService.upload( + anyString(), + nullable(String.class), + anyString(), + ArgumentMatchers.>any(), + any(MediaType.class))) + .thenReturn(Mono.never()); + var tools = new AttachmentTools( + client, + attachmentService, + authorization, + new AttachmentUploadLimiter()); + var arguments = Map.of( + "filename", "a.txt", + "policyName", "local", + "contentBase64", "YQ=="); + + var subscriptions = java.util.stream.IntStream.range(0, 2) + .mapToObj(ignored -> tools.upload(arguments).subscribe()) + .toList(); + try { + StepVerifier.create(tools.upload(arguments)) + .expectErrorSatisfies(error -> assertThat(error) + .isInstanceOf(McpToolException.class) + .extracting(value -> ((McpToolException) value).code()) + .isEqualTo("RATE_LIMITED")) + .verify(); + verify(attachmentService, times(2)) + .upload( + anyString(), + nullable(String.class), + anyString(), + ArgumentMatchers.>any(), + any(MediaType.class)); + } finally { + subscriptions.forEach(Disposable::dispose); + } + + var replacement = tools.upload(arguments).subscribe(); + try { + verify(attachmentService, times(3)) + .upload( + anyString(), + nullable(String.class), + anyString(), + ArgumentMatchers.>any(), + any(MediaType.class)); + } finally { + replacement.dispose(); + } + } + + @Test + void streamsTheMaximumLegitimateUploadInBoundedBuffersAndReleasesCapacity() { + when(authorization.keyId()).thenReturn(Mono.just("key-one")); + var attachment = new Attachment(); + attachment.setMetadata(ToolSupport.metadata("attachment-one")); + var receivedBytes = new AtomicLong(); + var largestBuffer = new AtomicInteger(); + when(attachmentService.upload( + anyString(), + nullable(String.class), + anyString(), + ArgumentMatchers.>any(), + any(MediaType.class))) + .thenAnswer(invocation -> { + Flux content = invocation.getArgument(3); + return content.doOnNext(buffer -> { + var readableBytes = buffer.readableByteCount(); + receivedBytes.addAndGet(readableBytes); + largestBuffer.accumulateAndGet(readableBytes, Math::max); + DataBufferUtils.release(buffer); + }) + .then(Mono.just(attachment)); + }); + var tools = new AttachmentTools( + client, + attachmentService, + authorization, + new AttachmentUploadLimiter(8L * 1024 * 1024, 8L * 1024 * 1024, 1, 1, 1)); + var bytes = new byte[8 * 1024 * 1024]; + var arguments = Map.of( + "filename", "maximum.bin", + "policyName", "local", + "contentBase64", Base64.getEncoder().encodeToString(bytes)); + + StepVerifier.create(tools.upload(arguments)) + .assertNext(payload -> assertThat(payload.summary()).contains("maximum.bin")) + .verifyComplete(); + StepVerifier.create(tools.upload(Map.of( + "filename", "small.bin", + "policyName", "local", + "contentBase64", "YQ=="))) + .expectNextCount(1) + .verifyComplete(); + + assertThat(receivedBytes).hasValue(8L * 1024 * 1024 + 1); + assertThat(largestBuffer).hasValueLessThanOrEqualTo(64 * 1024); + } + + @Test + void rejectsDecodedContentOverEightMiBBeforeCallingTheBackend() { + var tools = new AttachmentTools( + client, attachmentService, authorization, new AttachmentUploadLimiter()); + var bytes = new byte[8 * 1024 * 1024 + 1]; + + assertThatThrownBy(() -> tools.upload(Map.of( + "filename", "too-large.bin", + "policyName", "local", + "contentBase64", Base64.getEncoder().encodeToString(bytes)))) + .isInstanceOf(McpToolException.class) + .hasMessageContaining("8 MiB"); + + verify(attachmentService, never()) + .upload( + anyString(), + nullable(String.class), + anyString(), + ArgumentMatchers.>any(), + any(MediaType.class)); + } + + @Test + void releasesUploadCapacityWhenTheBackendFails() { + when(authorization.keyId()).thenReturn(Mono.just("key-one")); + var attachment = new Attachment(); + attachment.setMetadata(ToolSupport.metadata("attachment-one")); + when(attachmentService.upload( + anyString(), + nullable(String.class), + anyString(), + ArgumentMatchers.>any(), + any(MediaType.class))) + .thenReturn(Mono.error(new IllegalStateException("storage failed"))) + .thenReturn(Mono.just(attachment)); + var tools = new AttachmentTools( + client, + attachmentService, + authorization, + new AttachmentUploadLimiter(1, 1, 1, 1, 1)); + var arguments = Map.of( + "filename", "a.txt", + "policyName", "local", + "contentBase64", "YQ=="); + + StepVerifier.create(tools.upload(arguments)) + .expectErrorMessage("storage failed") + .verify(); + StepVerifier.create(tools.upload(arguments)) + .expectNextCount(1) + .verifyComplete(); } @Test @@ -50,7 +235,8 @@ void deletionUsesExtensionLifecycleSoReconcilerCleansStorage() { attachment.getMetadata().setVersion(1L); when(client.fetch(Attachment.class, "attachment-one")).thenReturn(Mono.just(attachment)); when(client.delete(attachment)).thenReturn(Mono.just(attachment)); - var tools = new AttachmentTools(client, attachmentService, authorization); + var tools = new AttachmentTools( + client, attachmentService, authorization, new AttachmentUploadLimiter()); StepVerifier.create(tools.delete(Map.of("name", "attachment-one"))) .assertNext(payload -> assertThat(payload.summary()).contains("attachment-one")) diff --git a/src/test/java/run/halo/mcpserver/tools/AttachmentUploadLimiterTest.java b/src/test/java/run/halo/mcpserver/tools/AttachmentUploadLimiterTest.java new file mode 100644 index 0000000..2d72ef6 --- /dev/null +++ b/src/test/java/run/halo/mcpserver/tools/AttachmentUploadLimiterTest.java @@ -0,0 +1,42 @@ +package run.halo.mcpserver.tools; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.Test; + +class AttachmentUploadLimiterTest { + + @Test + void enforcesGlobalAndPerKeyByteBudgets() { + var limiter = new AttachmentUploadLimiter(10, 6, 4, 3, 4); + var first = limiter.tryAcquire("key-one", "policy-one", 6).orElseThrow(); + + assertThat(limiter.tryAcquire("key-one", "policy-two", 1)).isEmpty(); + var second = limiter.tryAcquire("key-two", "policy-two", 4).orElseThrow(); + assertThat(limiter.tryAcquire("key-three", "policy-three", 1)).isEmpty(); + + first.close(); + var replacement = limiter.tryAcquire("key-three", "policy-three", 6).orElseThrow(); + + second.close(); + replacement.close(); + } + + @Test + void enforcesUploadAndPolicyLimitsAndReleasesIdempotently() { + var limiter = new AttachmentUploadLimiter(100, 100, 2, 1, 1); + var first = limiter.tryAcquire("key-one", "policy-one", 1).orElseThrow(); + + assertThat(limiter.tryAcquire("key-one", "policy-two", 1)).isEmpty(); + assertThat(limiter.tryAcquire("key-two", "policy-one", 1)).isEmpty(); + var second = limiter.tryAcquire("key-two", "policy-two", 1).orElseThrow(); + assertThat(limiter.tryAcquire("key-three", "policy-three", 1)).isEmpty(); + + first.close(); + first.close(); + var replacement = limiter.tryAcquire("key-three", "policy-one", 1).orElseThrow(); + + second.close(); + replacement.close(); + } +} diff --git a/src/test/java/run/halo/mcpserver/tools/BuiltInToolsTest.java b/src/test/java/run/halo/mcpserver/tools/BuiltInToolsTest.java index 2e602ab..ebebb4d 100644 --- a/src/test/java/run/halo/mcpserver/tools/BuiltInToolsTest.java +++ b/src/test/java/run/halo/mcpserver/tools/BuiltInToolsTest.java @@ -24,10 +24,15 @@ void organizesToolsByHaloDomainAndUsesChineseConsoleDescriptions() { new ContentSearchTools(mock(SearchService.class), authorization), new PostTools(client, mock(PostContentService.class), snapshots, authorization), new SinglePageTools(client, snapshots, authorization), - new CategoryTools(client, authorization), + new CategoryTools( + client, authorization, new CategoryParentMutationCoordinator(client)), new TagTools(client, authorization), new CommentTools(client, authorization), - new AttachmentTools(client, mock(AttachmentService.class), authorization)); + new AttachmentTools( + client, + mock(AttachmentService.class), + authorization, + new AttachmentUploadLimiter())); var names = tools.names().toList(); assertThat(tools.tools()).hasSize(29); @@ -67,10 +72,15 @@ void marksOnlyRecycleAndDeleteToolsAsDestructive() { new ContentSearchTools(mock(SearchService.class), authorization), new PostTools(client, mock(PostContentService.class), snapshots, authorization), new SinglePageTools(client, snapshots, authorization), - new CategoryTools(client, authorization), + new CategoryTools( + client, authorization, new CategoryParentMutationCoordinator(client)), new TagTools(client, authorization), new CommentTools(client, authorization), - new AttachmentTools(client, mock(AttachmentService.class), authorization)); + new AttachmentTools( + client, + mock(AttachmentService.class), + authorization, + new AttachmentUploadLimiter())); var destructive = tools.tools().stream() .filter(tool -> Boolean.TRUE.equals(tool.protocolTool().annotations().destructiveHint())) @@ -95,10 +105,15 @@ void publishesDetailedOutputObjectSchemas() { new ContentSearchTools(mock(SearchService.class), authorization), new PostTools(client, mock(PostContentService.class), snapshots, authorization), new SinglePageTools(client, snapshots, authorization), - new CategoryTools(client, authorization), + new CategoryTools( + client, authorization, new CategoryParentMutationCoordinator(client)), new TagTools(client, authorization), new CommentTools(client, authorization), - new AttachmentTools(client, mock(AttachmentService.class), authorization)); + new AttachmentTools( + client, + mock(AttachmentService.class), + authorization, + new AttachmentUploadLimiter())); var schemaValidator = new DefaultJsonSchemaValidator(); assertThat(tools.tools()).allSatisfy(tool -> { diff --git a/src/test/java/run/halo/mcpserver/tools/CategoryToolsTest.java b/src/test/java/run/halo/mcpserver/tools/CategoryToolsTest.java index d6c97f9..80208d6 100644 --- a/src/test/java/run/halo/mcpserver/tools/CategoryToolsTest.java +++ b/src/test/java/run/halo/mcpserver/tools/CategoryToolsTest.java @@ -2,17 +2,32 @@ import static org.assertj.core.api.Assertions.assertThat; import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.lenient; import static org.mockito.Mockito.when; +import java.time.Duration; +import java.util.HashSet; import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.dao.DataIntegrityViolationException; +import org.springframework.dao.OptimisticLockingFailureException; +import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import reactor.core.publisher.Sinks; +import reactor.core.scheduler.Schedulers; import reactor.test.StepVerifier; import run.halo.app.core.extension.content.Category; import run.halo.app.extension.ReactiveExtensionClient; +import run.halo.mcpserver.CategoryParentMutationLock; import run.halo.mcpserver.McpAuthorization; @ExtendWith(MockitoExtension.class) @@ -24,12 +39,58 @@ class CategoryToolsTest { @Mock McpAuthorization authorization; + AtomicReference storedLock; + + @BeforeEach + void setUpLockStore() { + storedLock = new AtomicReference<>(); + lenient() + .when(client.fetch( + eq(CategoryParentMutationLock.class), + eq(CategoryParentMutationCoordinator.LOCK_NAME))) + .thenAnswer(ignored -> Mono.defer(() -> + Mono.justOrEmpty(copyOf(storedLock.get())))); + lenient() + .when(client.create(any(CategoryParentMutationLock.class))) + .thenAnswer(invocation -> Mono.defer(() -> { + synchronized (storedLock) { + if (storedLock.get() != null) { + return Mono.error(new DataIntegrityViolationException("lock exists")); + } + var created = copyOf( + invocation.getArgument(0)); + created.getMetadata().setVersion(1L); + storedLock.set(created); + return Mono.just(copyOf(created)); + } + })); + lenient() + .when(client.update(any(CategoryParentMutationLock.class))) + .thenAnswer(invocation -> Mono.defer(() -> { + synchronized (storedLock) { + var current = storedLock.get(); + var candidate = invocation.getArgument(0); + if (current == null + || !current.getMetadata() + .getVersion() + .equals(candidate.getMetadata().getVersion())) { + return Mono.error(new OptimisticLockingFailureException( + "lock version changed")); + } + var updated = copyOf(candidate); + updated.getMetadata().setVersion(current.getMetadata().getVersion() + 1); + storedLock.set(updated); + return Mono.just(copyOf(updated)); + } + })); + } + @Test void createsChildCategoryWithDisplayMetadata() { var parent = category("parent", null, 1L); when(client.fetch(Category.class, "parent")).thenReturn(Mono.just(parent)); when(client.create(any(Category.class))).thenAnswer(invocation -> Mono.just(invocation.getArgument(0))); - var tools = new CategoryTools(client, authorization); + var tools = newTools(); StepVerifier.create(tools.create(ToolSupport.map( "name", "child", @@ -46,7 +107,7 @@ void createsChildCategoryWithDisplayMetadata() { void rejectsCategoryParentCycle() { var child = category("child", "current", 1L); when(client.fetch(Category.class, "child")).thenReturn(Mono.just(child)); - var tools = new CategoryTools(client, authorization); + var tools = newTools(); StepVerifier.create(tools.update(Map.of("name", "current", "parent", "child"))) .expectErrorSatisfies(error -> assertThat(error) @@ -54,6 +115,121 @@ void rejectsCategoryParentCycle() { .verify(); } + @Test + void preventsCycleFromConcurrentReparenting() { + var stored = new ConcurrentHashMap(); + stored.put("a", category("a", null, 1L)); + stored.put("b", category("b", null, 1L)); + + assertConcurrentReparentingBlocked(stored, "a", "b", "b", "a"); + } + + @Test + void preventsLongerCycleFromConcurrentReparenting() { + var stored = new ConcurrentHashMap(); + stored.put("a", category("a", null, 1L)); + stored.put("b", category("b", "a", 1L)); + stored.put("c", category("c", null, 1L)); + + assertConcurrentReparentingBlocked(stored, "a", "c", "c", "b"); + } + + @Test + void keepsCoordinationLockUntilCancelledWriteFinishes() { + var stored = new ConcurrentHashMap(); + stored.put("a", category("a", null, 1L)); + stored.put("b", category("b", null, 1L)); + var updateDispatched = Sinks.one(); + var allowCommit = Sinks.one(); + when(client.fetch(eq(Category.class), anyString())).thenAnswer(invocation -> + Mono.defer(() -> Mono.just(copyOf(stored.get(invocation.getArgument(1)))))); + when(client.update(any(Category.class))).thenAnswer(invocation -> { + var updated = copyOf(invocation.getArgument(0)); + if (!updated.getMetadata().getName().equals("a")) { + return commit(updated, stored); + } + return Mono.defer(() -> { + updateDispatched.tryEmitEmpty(); + return allowCommit.asMono().then(commit(updated, stored)); + }); + }); + var firstTools = newTools(); + var secondTools = newTools(); + + var first = firstTools.update(Map.of("name", "a", "parent", "b")).subscribe(); + StepVerifier.create(updateDispatched.asMono()).verifyComplete(); + first.dispose(); + + var secondOutcome = secondTools + .update(Map.of("name", "b", "parent", "a")) + .materialize() + .cache(); + var second = secondOutcome.subscribe(); + StepVerifier.create(Mono.delay(Duration.ofMillis(250)) + .doOnNext(ignored -> { + assertThat(stored.get("a").getSpec().getParent()).isNull(); + assertThat(stored.get("b").getSpec().getParent()).isNull(); + allowCommit.tryEmitEmpty(); + }) + .then(secondOutcome)) + .assertNext(signal -> { + assertThat(signal.isOnError()).isTrue(); + assertThat(signal.getThrowable()).hasMessageContaining("category cycle"); + assertAcyclic(stored); + }) + .verifyComplete(); + second.dispose(); + } + + private void assertConcurrentReparentingBlocked( + ConcurrentHashMap stored, + String firstName, + String firstParent, + String secondName, + String secondParent) { + var fetchPair = new FetchPair(); + when(client.fetch(eq(Category.class), anyString())) + .thenAnswer(invocation -> fetchPair.fetch(invocation.getArgument(1), stored)); + when(client.update(any(Category.class))).thenAnswer(invocation -> Mono.fromSupplier(() -> { + var updated = copyOf(invocation.getArgument(0)); + stored.put(updated.getMetadata().getName(), updated); + return copyOf(updated); + })); + var firstTools = newTools(); + var secondTools = newTools(); + + var firstUpdate = firstTools.update(Map.of("name", firstName, "parent", firstParent)) + .thenReturn("updated-" + firstName) + .onErrorResume(error -> Mono.just("error:" + error.getMessage())) + .subscribeOn(Schedulers.parallel()); + var secondUpdate = secondTools.update(Map.of("name", secondName, "parent", secondParent)) + .thenReturn("updated-" + secondName) + .onErrorResume(error -> Mono.just("error:" + error.getMessage())) + .subscribeOn(Schedulers.parallel()); + + StepVerifier.create(Flux.merge(firstUpdate, secondUpdate).collectList()) + .assertNext(outcomes -> { + assertThat(outcomes).filteredOn(value -> value.startsWith("updated-")).hasSize(1); + assertThat(outcomes).anySatisfy(value -> assertThat(value) + .startsWith("error:") + .contains("category cycle")); + assertAcyclic(stored); + }) + .verifyComplete(); + } + + private static void assertAcyclic(Map stored) { + stored.keySet().forEach(name -> { + var visited = new HashSet(); + var current = name; + while (current != null) { + assertThat(visited.add(current)).as("parent chain from %s", name).isTrue(); + var category = stored.get(current); + current = category == null ? null : category.getSpec().getParent(); + } + }); + } + private static Category category(String name, String parent, long version) { var category = new Category(); category.setMetadata(ToolSupport.metadata(name)); @@ -65,4 +241,69 @@ private static Category category(String name, String parent, long version) { category.getSpec().setPriority(0); return category; } + + private static Category copyOf(Category source) { + return category( + source.getMetadata().getName(), + source.getSpec().getParent(), + source.getMetadata().getVersion()); + } + + private CategoryTools newTools() { + return new CategoryTools( + client, authorization, new CategoryParentMutationCoordinator(client)); + } + + private static Mono commit( + Category updated, ConcurrentHashMap stored) { + return Mono.fromSupplier(() -> { + var committed = copyOf(updated); + stored.put(committed.getMetadata().getName(), committed); + return copyOf(committed); + }); + } + + private static CategoryParentMutationLock copyOf(CategoryParentMutationLock source) { + if (source == null) { + return null; + } + var copy = new CategoryParentMutationLock(); + copy.setMetadata(ToolSupport.metadata(source.getMetadata().getName())); + copy.getMetadata().setVersion(source.getMetadata().getVersion()); + var spec = new CategoryParentMutationLock.Spec(); + spec.setHolder(source.getSpec().getHolder()); + spec.setExpiresAt(source.getSpec().getExpiresAt()); + spec.setGeneration(source.getSpec().getGeneration()); + copy.setSpec(spec); + return copy; + } + + private static final class FetchPair { + private static final Duration WAIT_FOR_CONCURRENT_FETCH = Duration.ofMillis(250); + + private final Object monitor = new Object(); + private Sinks.One release = Sinks.one(); + private int arrivals; + + Mono fetch(String name, Map stored) { + return Mono.defer(() -> { + var snapshot = copyOf(stored.get(name)); + Sinks.One currentRelease; + synchronized (monitor) { + currentRelease = release; + arrivals++; + if (arrivals == 2) { + arrivals = 0; + release = Sinks.one(); + currentRelease.tryEmitEmpty(); + } + } + return currentRelease + .asMono() + .timeout(WAIT_FOR_CONCURRENT_FETCH) + .onErrorResume(TimeoutException.class, ignored -> Mono.empty()) + .thenReturn(snapshot); + }); + } + } } diff --git a/ui/src/components/AccessKeySecretModal.vue b/ui/src/components/AccessKeySecretModal.vue index 709dfab..3cefcf0 100644 --- a/ui/src/components/AccessKeySecretModal.vue +++ b/ui/src/components/AccessKeySecretModal.vue @@ -32,7 +32,7 @@ const { copy, copied } = useClipboard({
接入方式
- +
diff --git a/ui/src/components/__tests__/McpConnectionGuide.test.ts b/ui/src/components/__tests__/McpConnectionGuide.test.ts new file mode 100644 index 0000000..eb2ecc3 --- /dev/null +++ b/ui/src/components/__tests__/McpConnectionGuide.test.ts @@ -0,0 +1,63 @@ +import { flushPromises, mount, type VueWrapper } from '@vue/test-utils' +import { beforeEach, describe, expect, it, vi } from 'vitest' + +import McpConnectionGuide from '../McpConnectionGuide.vue' + +const SECRET_MARKER = 'hmcp_component_secret_marker' + +const { clipboardWrite } = vi.hoisted(() => ({ + clipboardWrite: vi.fn<(value: string) => Promise>(), +})) + +vi.mock('@halo-dev/components', () => ({ + Toast: { + success: vi.fn<(message: string) => void>(), + error: vi.fn<(message: string) => void>(), + }, +})) + +function tab(wrapper: VueWrapper, label: string) { + const button = wrapper.findAll('button').find((candidate) => candidate.text().includes(label)) + if (!button) { + throw new Error(`Missing ${label} tab`) + } + return button +} + +describe('McpConnectionGuide', () => { + beforeEach(() => { + clipboardWrite.mockReset().mockResolvedValue() + Object.defineProperty(navigator, 'clipboard', { + configurable: true, + value: { writeText: clipboardWrite }, + }) + }) + + it('keeps a legacy token attribute out of rendered, copied and navigated artifacts', async () => { + const wrapper = mount(McpConnectionGuide, { + attrs: { token: SECRET_MARKER }, + }) + + expect(wrapper.html()).not.toContain(SECRET_MARKER) + expect(wrapper.text()).toContain('“添加端点”不包含认证信息') + + await tab(wrapper, 'Cursor').trigger('click') + expect(wrapper.get('a').text()).toBe('添加端点') + expect(wrapper.get('a').attributes('href')).not.toContain(SECRET_MARKER) + + await tab(wrapper, 'VS Code').trigger('click') + expect(wrapper.get('a').attributes('href')).not.toContain(SECRET_MARKER) + + const copyButton = wrapper + .findAll('button') + .find((candidate) => candidate.text().trim() === '复制') + if (!copyButton) { + throw new Error('Missing copy button') + } + await copyButton.trigger('click') + await flushPromises() + + expect(clipboardWrite).toHaveBeenCalledOnce() + expect(clipboardWrite.mock.calls[0]?.[0]).not.toContain(SECRET_MARKER) + }) +}) diff --git a/ui/src/utils/__tests__/mcp-config.test.ts b/ui/src/utils/__tests__/mcp-config.test.ts index ac6c4e0..b99db40 100644 --- a/ui/src/utils/__tests__/mcp-config.test.ts +++ b/ui/src/utils/__tests__/mcp-config.test.ts @@ -1,12 +1,56 @@ import { describe, expect, it } from 'vitest' -import { mcpClientGuides } from '../mcp-config' +import { mcpClientGuides, mcpEndpoint, mcpHttpConfig, type McpClientGuide } from '../mcp-config' + +const SECRET_MARKER = 'hmcp_regression_secret_marker' + +function guidesFromLegacySecret(secret: string) { + const legacyGenerator = mcpClientGuides as unknown as (secret: string) => McpClientGuide[] + return legacyGenerator(secret) +} describe('mcpClientGuides', () => { - it('writes the Codex token directly to a static authorization header', () => { - const codex = mcpClientGuides('hmcp_secret').find((guide) => guide.id === 'codex') + it('does not expose a supplied token in generated artifacts', () => { + const guides = guidesFromLegacySecret(SECRET_MARKER) + const legacyHttpGenerator = mcpHttpConfig as unknown as (secret: string) => string + + for (const guide of guides) { + const artifacts = [guide.content, guide.installUrl ?? ''] + expect(artifacts.join('\n')).not.toContain(SECRET_MARKER) + expect(artifacts.join('\n')).not.toContain(encodeURIComponent(SECRET_MARKER)) + } + expect(legacyHttpGenerator(SECRET_MARKER)).not.toContain(SECRET_MARKER) + + const cursor = guides.find((guide) => guide.id === 'cursor') + const cursorUrl = new URL(cursor!.installUrl!) + const cursorConfig = JSON.parse(window.atob(cursorUrl.searchParams.get('config')!)) + expect(cursorConfig).toEqual({ type: 'http', url: mcpEndpoint() }) + + const vscode = guides.find((guide) => guide.id === 'vscode') + const encodedVscodeConfig = vscode!.installUrl!.split('?')[1] + expect(encodedVscodeConfig).toBeDefined() + const vscodeConfig = JSON.parse(decodeURIComponent(encodedVscodeConfig!)) + expect(vscodeConfig).toEqual({ name: 'halo', type: 'http', url: mcpEndpoint() }) + }) + + it('keeps working authentication setup through client-native secret indirection', () => { + const guides = mcpClientGuides() + + expect(guides).toHaveLength(4) + expect(guides.every((guide) => guide.content.includes(mcpEndpoint()))).toBe(true) + const claude = guides.find((guide) => guide.id === 'claude-code') + expect(claude?.content).toContain('"Authorization": "Bearer ${HALO_MCP_TOKEN}"') + expect(() => JSON.parse(claude!.content)).not.toThrow() + expect(guides.find((guide) => guide.id === 'codex')?.content).toContain( + 'bearer_token_env_var = "HALO_MCP_TOKEN"', + ) + const cursor = guides.find((guide) => guide.id === 'cursor') + expect(cursor?.content).toContain('"Authorization": "Bearer ${env:HALO_MCP_TOKEN}"') + expect(() => JSON.parse(cursor!.content)).not.toThrow() - expect(codex?.content).toContain('http_headers = { Authorization = "Bearer hmcp_secret" }') - expect(codex?.content).not.toContain('bearer_token_env_var') + const vscode = guides.find((guide) => guide.id === 'vscode') + expect(vscode?.content).toContain('"Authorization": "Bearer ${input:halo-mcp-token}"') + expect(vscode?.content).toContain('"password": true') + expect(() => JSON.parse(vscode!.content)).not.toThrow() }) }) diff --git a/ui/src/utils/mcp-config.ts b/ui/src/utils/mcp-config.ts index 5347802..75179e8 100644 --- a/ui/src/utils/mcp-config.ts +++ b/ui/src/utils/mcp-config.ts @@ -3,9 +3,11 @@ export function mcpEndpoint() { } const SERVER_NAME = 'halo' -const TOKEN_PLACEHOLDER = '$HALO_MCP_TOKEN' +const CLAUDE_TOKEN_REFERENCE = '${HALO_MCP_TOKEN}' +const CURSOR_TOKEN_REFERENCE = '${env:HALO_MCP_TOKEN}' +const VSCODE_TOKEN_INPUT_ID = 'halo-mcp-token' -export function mcpHttpConfig(token = TOKEN_PLACEHOLDER) { +export function mcpHttpConfig() { return JSON.stringify( { mcpServers: { @@ -13,7 +15,7 @@ export function mcpHttpConfig(token = TOKEN_PLACEHOLDER) { type: 'http', url: mcpEndpoint(), headers: { - Authorization: `Bearer ${token}`, + Authorization: `Bearer ${CURSOR_TOKEN_REFERENCE}`, }, }, }, @@ -32,20 +34,30 @@ export interface McpClientGuide { installUrl?: string } -export function mcpClientGuides(token?: string): McpClientGuide[] { +export function mcpClientGuides(): McpClientGuide[] { const endpoint = mcpEndpoint() - const bearer = `Bearer ${token ?? TOKEN_PLACEHOLDER}` - const serverConfig = { + const installConfig = { type: 'http', url: endpoint, - headers: { Authorization: bearer }, } const guides: McpClientGuide[] = [ { id: 'claude-code', label: 'Claude Code', - content: `claude mcp add --transport http ${SERVER_NAME} ${endpoint} --header "Authorization: ${bearer}"`, + content: JSON.stringify( + { + mcpServers: { + [SERVER_NAME]: { + type: 'http', + url: endpoint, + headers: { Authorization: `Bearer ${CLAUDE_TOKEN_REFERENCE}` }, + }, + }, + }, + null, + 2, + ), }, { id: 'codex', @@ -53,32 +65,46 @@ export function mcpClientGuides(token?: string): McpClientGuide[] { content: `# ~/.codex/config.toml [mcp_servers.${SERVER_NAME}] url = "${endpoint}" -http_headers = { Authorization = "${bearer}" }`, +bearer_token_env_var = "HALO_MCP_TOKEN"`, }, { id: 'cursor', label: 'Cursor', - content: mcpHttpConfig(token), + content: mcpHttpConfig(), + installUrl: `cursor://anysphere.cursor-deeplink/mcp/install?name=${SERVER_NAME}&config=${window.btoa( + JSON.stringify(installConfig), + )}`, }, { id: 'vscode', label: 'VS Code', - content: JSON.stringify({ servers: { [SERVER_NAME]: serverConfig } }, null, 2), + content: JSON.stringify( + { + inputs: [ + { + type: 'promptString', + id: VSCODE_TOKEN_INPUT_ID, + description: 'Halo MCP access key', + password: true, + }, + ], + servers: { + [SERVER_NAME]: { + ...installConfig, + headers: { + Authorization: `Bearer \${input:${VSCODE_TOKEN_INPUT_ID}}`, + }, + }, + }, + }, + null, + 2, + ), + installUrl: `vscode:mcp/install?${encodeURIComponent( + JSON.stringify({ name: SERVER_NAME, ...installConfig }), + )}`, }, ] - if (token) { - for (const guide of guides) { - if (guide.id === 'cursor') { - const config = window.btoa(JSON.stringify(serverConfig)) - guide.installUrl = `cursor://anysphere.cursor-deeplink/mcp/install?name=${SERVER_NAME}&config=${config}` - } - if (guide.id === 'vscode') { - const config = encodeURIComponent(JSON.stringify({ name: SERVER_NAME, ...serverConfig })) - guide.installUrl = `vscode:mcp/install?${config}` - } - } - } - return guides }