Skip to content

Commit 012c539

Browse files
committed
Bound concurrent MCP authentication
1 parent b82ab0a commit 012c539

3 files changed

Lines changed: 82 additions & 19 deletions

File tree

src/main/java/run/halo/mcpserver/McpInFlightLimiter.java

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,17 +2,24 @@
22

33
import org.springframework.stereotype.Component;
44

5-
/** Process-local bounds on concurrent authenticated MCP request work. */
5+
/** Process-local bounds on concurrent MCP authentication and request work. */
66
@Component
77
class McpInFlightLimiter {
88

9+
static final int AUTHENTICATION_LIMIT = 100;
910
static final int GLOBAL_LIMIT = 100;
1011
static final int PER_KEY_LIMIT = 16;
1112
static final int MAX_TRACKED_KEYS = 10_000;
1213

14+
private final KeyedInFlightLimiter authenticationLimiter =
15+
new KeyedInFlightLimiter(AUTHENTICATION_LIMIT, AUTHENTICATION_LIMIT, 1);
1316
private final KeyedInFlightLimiter limiter =
1417
new KeyedInFlightLimiter(GLOBAL_LIMIT, PER_KEY_LIMIT, MAX_TRACKED_KEYS);
1518

19+
KeyedInFlightLimiter.Lease tryAcquireAuthentication() {
20+
return authenticationLimiter.tryAcquire("authentication", 1);
21+
}
22+
1623
/**
1724
* Reserves one global and per-key slot for the key, or returns null when either budget is
1825
* exhausted. The caller must close the permit exactly once when the request terminates.

src/main/java/run/halo/mcpserver/McpKeyAuthenticationFilter.java

Lines changed: 28 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -68,27 +68,37 @@ public Mono<Void> filter(ServerWebExchange exchange, WebFilterChain chain) {
6868
return tooManyRequests(exchange);
6969
}
7070
var rawToken = authorization.substring(BEARER_SCHEME.length());
71-
return accessKeyService.authenticate(rawToken, exchange.getRequest().getRemoteAddress())
72-
.flatMap(authentication -> {
73-
if (!hasSupportedProtocolVersion(exchange)) {
74-
return badRequest(exchange).thenReturn(true);
75-
}
76-
var permit = inFlightLimiter.tryAcquire(authentication.keyId());
77-
if (permit == null) {
71+
return Mono.defer(() -> {
72+
var authenticationPermit = inFlightLimiter.tryAcquireAuthentication();
73+
if (authenticationPermit == null) {
7874
return tooManyRequests(exchange).thenReturn(true);
7975
}
80-
var request = exchange.getRequest().mutate()
81-
.headers(headers -> headers.remove(AUTHORIZATION))
82-
.build();
83-
// Defer so a synchronously throwing handler assembly still flows through
84-
// the error path and releases the permit.
85-
return Mono.defer(() -> mcpHandler.handle(exchange.mutate().request(request).build()))
86-
.doFinally(ignored -> permit.close())
87-
.contextWrite(org.springframework.security.core.context.ReactiveSecurityContextHolder
88-
.withAuthentication(authentication))
89-
.thenReturn(true);
76+
return Mono.defer(() -> accessKeyService.authenticate(
77+
rawToken, exchange.getRequest().getRemoteAddress()))
78+
.doFinally(ignored -> authenticationPermit.close())
79+
.flatMap(authentication -> {
80+
if (!hasSupportedProtocolVersion(exchange)) {
81+
return badRequest(exchange).thenReturn(true);
82+
}
83+
var permit = inFlightLimiter.tryAcquire(authentication.keyId());
84+
if (permit == null) {
85+
return tooManyRequests(exchange).thenReturn(true);
86+
}
87+
var request = exchange.getRequest().mutate()
88+
.headers(headers -> headers.remove(AUTHORIZATION))
89+
.build();
90+
// Defer so a synchronously throwing handler assembly still flows through
91+
// the error path and releases the permit.
92+
return Mono.defer(() -> mcpHandler.handle(
93+
exchange.mutate().request(request).build()))
94+
.doFinally(ignored -> permit.close())
95+
.contextWrite(org.springframework.security.core.context
96+
.ReactiveSecurityContextHolder
97+
.withAuthentication(authentication))
98+
.thenReturn(true);
99+
})
100+
.defaultIfEmpty(false);
90101
})
91-
.defaultIfEmpty(false)
92102
// The deadline covers the whole authenticate-and-handle chain so stalled
93103
// credential lookups cannot park requests either.
94104
.timeout(requestTimeout)

src/test/java/run/halo/mcpserver/McpKeyAuthenticationFilterTest.java

Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@
77
import static org.mockito.Mockito.when;
88

99
import java.net.InetSocketAddress;
10+
import java.time.Duration;
11+
import java.util.ArrayList;
1012
import java.util.Set;
1113
import java.util.concurrent.atomic.AtomicReference;
1214
import org.junit.jupiter.api.BeforeEach;
@@ -24,6 +26,7 @@
2426
import org.springframework.web.reactive.function.server.ServerResponse;
2527
import org.springframework.web.server.ServerWebExchange;
2628
import reactor.core.publisher.Mono;
29+
import reactor.test.StepVerifier;
2730

2831
@ExtendWith(MockitoExtension.class)
2932
class McpKeyAuthenticationFilterTest {
@@ -249,6 +252,49 @@ accessKeyService, new McpRequestRateLimiter(), new McpInFlightLimiter(), mcpServ
249252
assertThat(exchange.getResponse().getStatusCode()).isEqualTo(HttpStatus.SERVICE_UNAVAILABLE);
250253
}
251254

255+
@Test
256+
void limitsConcurrentAuthenticationBeforeCredentialLookup() {
257+
var rawToken = "hmcp_00000000-0000-0000-0000-000000000000_secret";
258+
when(accessKeyService.authenticate(
259+
org.mockito.ArgumentMatchers.eq(rawToken),
260+
org.mockito.ArgumentMatchers.any(InetSocketAddress.class)))
261+
.thenReturn(Mono.never());
262+
var inFlightLimiter = new McpInFlightLimiter();
263+
var boundedFilter = new McpKeyAuthenticationFilter(
264+
accessKeyService, new McpRequestRateLimiter(), inFlightLimiter, mcpServer);
265+
var requests = new ArrayList<reactor.core.Disposable>();
266+
try {
267+
for (var i = 0; i < McpInFlightLimiter.AUTHENTICATION_LIMIT; i++) {
268+
var exchange = MockServerWebExchange.from(MockServerHttpRequest
269+
.post(McpKeyAuthenticationFilter.MCP_PATH)
270+
.remoteAddress(new InetSocketAddress("192.0.2." + (i + 1), 8080))
271+
.header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken));
272+
requests.add(boundedFilter.filter(exchange,
273+
ignored -> Mono.error(new AssertionError("Request must not continue")))
274+
.subscribe());
275+
}
276+
var overflow = MockServerWebExchange.from(MockServerHttpRequest
277+
.post(McpKeyAuthenticationFilter.MCP_PATH)
278+
.remoteAddress(new InetSocketAddress("198.51.100.1", 8080))
279+
.header(HttpHeaders.AUTHORIZATION, "Bearer " + rawToken));
280+
281+
StepVerifier.create(boundedFilter.filter(overflow,
282+
ignored -> Mono.error(new AssertionError("Request must not continue"))))
283+
.expectComplete()
284+
.verify(Duration.ofMillis(200));
285+
286+
assertThat(overflow.getResponse().getStatusCode()).isEqualTo(HttpStatus.TOO_MANY_REQUESTS);
287+
} finally {
288+
requests.forEach(reactor.core.Disposable::dispose);
289+
}
290+
var released = new ArrayList<KeyedInFlightLimiter.Lease>();
291+
for (var i = 0; i < McpInFlightLimiter.AUTHENTICATION_LIMIT; i++) {
292+
released.add(inFlightLimiter.tryAcquireAuthentication());
293+
}
294+
assertThat(released).doesNotContainNull();
295+
released.forEach(KeyedInFlightLimiter.Lease::close);
296+
}
297+
252298
@Test
253299
void rejectsRequestsWhenTheInFlightBudgetIsExhausted() {
254300
var inFlightLimiter = new McpInFlightLimiter();

0 commit comments

Comments
 (0)