diff --git a/app/mcp/route.ts b/app/mcp/route.ts index aa26953a..bc09f40c 100644 --- a/app/mcp/route.ts +++ b/app/mcp/route.ts @@ -35,8 +35,7 @@ import { McpLibraryItemSchema, } from "@/lib/integrations/mcp/protocol"; import { - incrementMcpRateCounter, - isOverLimit, + checkMcpRateLimit, MCP_RATE_BUCKETS, } from "@/lib/integrations/mcp/rate-limit"; import { @@ -123,24 +122,46 @@ async function authorizeToolCall( required === "library:write" ? MCP_RATE_BUCKETS.write : MCP_RATE_BUCKETS.read; - const decision = await incrementMcpRateCounter(auth.userId, bucket); - if (decision !== null && isOverLimit(decision, bucket)) { - return { result: rateLimitResult(bucket, decision) }; + const rateLimit = await checkMcpRateLimit(auth.userId, bucket); + if (rateLimit.status === "limited") { + return { + result: rateLimitResult(bucket, rateLimit.retryAfterSeconds), + }; + } + if (rateLimit.status === "unavailable") { + return { result: rateLimitUnavailableResult(auth.userId, bucket) }; } return { userId: auth.userId }; } function rateLimitResult( - bucket: { limit: number; name: string }, - decision: { retryAfterSeconds: number } + bucket: { name: string }, + retryAfterSeconds: number +): CallToolResult { + log.warn(`rate limit hit (${bucket.name})`, { retryAfterSeconds }); + return { + content: [ + { + text: `Rate limit reached for \`${bucket.name}\` operations. Retry in about ${retryAfterSeconds} seconds.`, + type: "text", + }, + ], + isError: true, + }; +} + +function rateLimitUnavailableResult( + userId: string, + bucket: { name: string } ): CallToolResult { - log.warn(`rate limit hit (${bucket.name})`, { - retryAfterSeconds: decision.retryAfterSeconds, + log.warn("rate limit unavailable; rejecting request", { + bucket: bucket.name, + userId, }); return { content: [ { - text: `Rate limit reached for \`${bucket.name}\` operations. Retry in about ${decision.retryAfterSeconds} seconds.`, + text: "Rate limiting is unavailable right now. Retry shortly.", type: "text", }, ], diff --git a/lib/collections/link-reachability.test.ts b/lib/collections/link-reachability.test.ts new file mode 100644 index 00000000..7fce1313 --- /dev/null +++ b/lib/collections/link-reachability.test.ts @@ -0,0 +1,47 @@ +import { afterEach, beforeEach, describe, expect, mock, test } from "bun:test"; + +let redisConfigured = false; + +mock.module("server-only", () => ({})); +mock.module("@/lib/common/redis", () => ({ + getReadyRedisClient: () => Promise.resolve(null), + isRedisConfigured: () => redisConfigured, +})); + +const { consumeProbeBudget } = await import( + "@/lib/collections/link-reachability" +); + +beforeEach(() => { + redisConfigured = false; +}); + +afterEach(() => { + mock.restore(); +}); + +describe("consumeProbeBudget", () => { + test("refuses probes when Redis is configured but unreachable", async () => { + redisConfigured = true; + expect(await consumeProbeBudget("user-1", 1)).toEqual({ + allowed: false, + retryAfterMs: 60_000, + }); + }); + + test("uses the local fallback when Redis is not configured", async () => { + redisConfigured = false; + expect(await consumeProbeBudget("user-1", 1)).toEqual({ + allowed: true, + retryAfterMs: 0, + }); + }); + + test("allows a zero-probe batch without touching the fallback", async () => { + redisConfigured = true; + expect(await consumeProbeBudget("user-1", 0)).toEqual({ + allowed: true, + retryAfterMs: 0, + }); + }); +}); diff --git a/lib/collections/link-reachability.ts b/lib/collections/link-reachability.ts index 8d31437b..d7bc8985 100644 --- a/lib/collections/link-reachability.ts +++ b/lib/collections/link-reachability.ts @@ -9,7 +9,7 @@ import { import { isAbortError } from "@/lib/common/abort"; import { mapConcurrent } from "@/lib/common/array"; import { createLogger } from "@/lib/common/logs/console/logger"; -import { getRedisClient } from "@/lib/common/redis"; +import { getReadyRedisClient, isRedisConfigured } from "@/lib/common/redis"; import { type FetchHttpRedirectResult, fetchPublicRedirect, @@ -92,10 +92,16 @@ function tryConsumeLocalProbeBudget( } /** - * Shared fixed-window budget via Redis when available; falls back to an - * in-process Map so local/dev still rate-limits a single isolate. + * Shared fixed-window budget for outbound probes. + * + * Redis holds the shared budget. When it is configured but unreachable the + * budget cannot be enforced across isolates, so probes are refused rather + * than fanned out; the client retries after the window. `getReadyRedisClient` + * waits out the connect window of a freshly created client, so a cold start + * is not read as an outage. A deployment with no `REDIS_URL` uses the + * in-process Map, which still bounds a single local isolate. */ -async function tryConsumeProbeBudget( +export async function consumeProbeBudget( userId: string, amount: number ): Promise<{ allowed: boolean; retryAfterMs: number }> { @@ -103,8 +109,15 @@ async function tryConsumeProbeBudget( return { allowed: true, retryAfterMs: 0 }; } - const redis = getRedisClient(); + const redis = await getReadyRedisClient(); if (!redis) { + if (isRedisConfigured()) { + log.warn("Link probe budget unavailable; refusing probes", { + amount, + userId, + }); + return { allowed: false, retryAfterMs: PROBE_BUDGET_WINDOW_MS }; + } return tryConsumeLocalProbeBudget(userId, amount); } @@ -128,6 +141,14 @@ async function tryConsumeProbeBudget( } return { allowed: true, retryAfterMs: 0 }; } catch (error) { + if (isRedisConfigured()) { + log.warn("Link probe Redis budget failed; refusing probes", { + amount, + error, + userId, + }); + return { allowed: false, retryAfterMs: PROBE_BUDGET_WINDOW_MS }; + } log.warn("Link probe Redis budget failed; using local fallback", { error, userId, @@ -293,7 +314,7 @@ export async function probeLibraryItemsReachability({ }); const probeCount = work.filter((entry) => entry.kind === "probe").length; - const budget = await tryConsumeProbeBudget(userId, probeCount); + const budget = await consumeProbeBudget(userId, probeCount); if (!budget.allowed) { return { didPersist: false, diff --git a/lib/common/redis.test.ts b/lib/common/redis.test.ts new file mode 100644 index 00000000..613dbe7f --- /dev/null +++ b/lib/common/redis.test.ts @@ -0,0 +1,62 @@ +import { afterAll, describe, expect, mock, test } from "bun:test"; + +type RedisHandler = (...args: unknown[]) => void; + +const handlers = new Map(); +let resolveInitialConnect: (() => void) | null = null; + +const redisClient = { + connect: () => + new Promise((resolve) => { + resolveInitialConnect = resolve; + }), + destroy: () => undefined, + isReady: false, + on(event: string, handler: RedisHandler) { + handlers.set(event, handler); + return redisClient; + }, + ping: () => Promise.resolve("PONG"), +}; + +mock.module("redis", () => ({ + createClient: () => redisClient, +})); + +process.env.REDIS_URL = "redis://localhost:6379"; + +// `mock.module` is process-wide and sibling test files mock this module's +// path, so load a private copy to test the real implementation. +const redis: typeof import("./redis") = await import( + `${import.meta.dir}/redis.ts?isolation` +); + +afterAll(() => { + mock.restore(); + delete process.env.REDIS_URL; +}); + +describe("getReadyRedisClient", () => { + test("waits for the cold-start connect instead of reporting an outage", async () => { + const pending = redis.getReadyRedisClient(); + + let didSettle = false; + pending.then(() => { + didSettle = true; + }); + await Promise.resolve(); + expect(didSettle).toBe(false); + + redisClient.isReady = true; + handlers.get("ready")?.(); + resolveInitialConnect?.(); + + expect((await pending)?.isReady).toBe(true); + }); + + test("reports unavailable once an established connection is lost", async () => { + redisClient.isReady = false; + + expect(await redis.getReadyRedisClient()).toBeNull(); + }); +}); diff --git a/lib/common/redis.ts b/lib/common/redis.ts index 14d471ca..490332ec 100644 --- a/lib/common/redis.ts +++ b/lib/common/redis.ts @@ -15,18 +15,77 @@ const RedisConnectionError = NamedError.create( ); export type RedisConnectionError = InstanceType; +/** + * How long an abuse-bounding caller waits for an in-flight connection before + * it treats a configured Redis as unavailable. + */ +const REDIS_READY_WAIT_TIMEOUT_MS = 1000; + let globalRedisClient: RedisClientType | null = null; +let redisConnectPromise: Promise | null = null; +let didWarnRedisUnavailable = false; +let hasRedisConnected = false; + +/** + * Whether Redis is configured through `REDIS_URL`. + * + * Separates "the operator never configured Redis" from "Redis is configured + * but its socket is not ready yet". Abuse-bounding callers must fail closed + * for the second state without breaking a deliberately Redis-less setup. + */ +export function isRedisConfigured(): boolean { + return Boolean(process.env.REDIS_URL); +} + +function warnRedisUnavailableOnce(): void { + if (didWarnRedisUnavailable) { + return; + } + didWarnRedisUnavailable = true; + log.warn( + "Redis client not ready (disconnected or reconnecting); callers degrade or fail closed" + ); +} + +/** + * Wait for the connect started by {@link getRedisClient} to settle, up to a + * bound. A client that is still connecting on a cold start resolves here in + * milliseconds; a client whose socket cannot connect keeps its connect + * pending across reconnects, so the timeout keeps the caller from hanging. + */ +function waitForRedisReady(timeoutMs: number): Promise { + const connecting = redisConnectPromise; + if (!connecting) { + return Promise.resolve(); + } + + return new Promise((resolve) => { + const timer = setTimeout(() => resolve(), timeoutMs); + connecting.then( + () => { + clearTimeout(timer); + resolve(); + }, + () => { + clearTimeout(timer); + resolve(); + } + ); + }); +} /** * Get a Redis client instance. * Returns null in browser environments, when Redis is not configured, or - * when the underlying socket has not connected yet (including during an - * automatic reconnection). Callers already handle null throughout the - * codebase via their existing degraded-path branches. + * when the underlying socket is not ready — which includes both the + * cold-start connect window and a reconnect after a dropped connection. * * The client auto-reconnects on disconnection. Once the socket is ready * again the returned value flips from null back to the client — no * instance is lost or re-created. + * + * Abuse-bounding callers use {@link getReadyRedisClient} so the transient + * cold-start window is not mistaken for an outage. */ export function getRedisClient(): RedisClientType | null { if (typeof window !== "undefined") { @@ -37,11 +96,20 @@ export function getRedisClient(): RedisClientType | null { // The client auto-reconnects after disconnection, but commands sent before // reconnection completes queue indefinitely (the offline queue is enabled by // default). Rather than returning a client that will hang callers, return null - // so every caller's existing null-handling branch gracefully degrades. + // so every caller can degrade or fail closed. // // Once the underlying socket is ready again the client will be returned on // the next call — no client is lost or re-created. - return globalRedisClient.isReady ? globalRedisClient : null; + if (globalRedisClient.isReady) { + return globalRedisClient; + } + // A dropped connection and a not-yet-established one both leave + // `isReady` false. Only the first is an outage, so warn only after an + // established connection has been lost. + if (hasRedisConnected) { + warnRedisUnavailableOnce(); + } + return null; } const url = @@ -57,12 +125,16 @@ export function getRedisClient(): RedisClientType | null { try { globalRedisClient = createClient({ url }); + redisConnectPromise = null; + hasRedisConnected = false; globalRedisClient.on("error", (error) => { log.error("Redis client error", { error }); }); - globalRedisClient.on("connect", () => { + globalRedisClient.on("ready", () => { + hasRedisConnected = true; + didWarnRedisUnavailable = false; log.info("Redis connection established"); }); @@ -75,17 +147,49 @@ export function getRedisClient(): RedisClientType | null { }); // Kick off connect eagerly so the client is ready by the first data request. - globalRedisClient.connect().catch((error) => { - log.error("Redis initial connect failed", { error }); - }); + redisConnectPromise = globalRedisClient.connect().then( + () => undefined, + (error) => { + log.error("Redis initial connect failed", { error }); + } + ); - return globalRedisClient.isReady ? globalRedisClient : null; + // The first call creates the client before its socket is ready. Return + // null now; getReadyRedisClient waits for the connect to settle. + return null; } catch (error) { log.error("Failed to initialize Redis client", { error }); return null; } } +/** + * Get a Redis client that has finished connecting, for callers that must fail + * closed when Redis is configured but unreachable. + * + * The transient state where a freshly created client is still connecting is + * waited out, so a cold start is not reported as an outage. A client that is + * genuinely down (or reconnecting) still reports null after the bounded wait. + * A deployment with no `REDIS_URL` returns null immediately. + */ +export async function getReadyRedisClient(): Promise { + const client = getRedisClient(); + if (client) { + return client; + } + if (!isRedisConfigured()) { + return null; + } + + await waitForRedisReady(REDIS_READY_WAIT_TIMEOUT_MS); + + const readyClient = getRedisClient(); + if (!readyClient) { + warnRedisUnavailableOnce(); + } + return readyClient; +} + /** * Close the Redis connection gracefully. * Important for proper cleanup in serverless environments. @@ -93,6 +197,8 @@ export function getRedisClient(): RedisClientType | null { export async function closeRedisConnection(): Promise { const client = globalRedisClient; globalRedisClient = null; + redisConnectPromise = null; + hasRedisConnected = false; if (!client) { return; } diff --git a/lib/integrations/mcp/rate-limit.test.ts b/lib/integrations/mcp/rate-limit.test.ts new file mode 100644 index 00000000..497d6859 --- /dev/null +++ b/lib/integrations/mcp/rate-limit.test.ts @@ -0,0 +1,75 @@ +import { afterEach, beforeEach, describe, expect, mock, test } from "bun:test"; + +interface RedisStub { + incr: (key: string) => Promise; + pExpire: (key: string, ms: number) => Promise; +} + +let redisClient: RedisStub | null = null; +let redisConfigured = false; + +beforeEach(() => { + redisClient = null; + redisConfigured = false; + mock.module("@/lib/common/redis", () => ({ + getReadyRedisClient: () => Promise.resolve(redisClient), + isRedisConfigured: () => redisConfigured, + })); +}); + +afterEach(() => { + mock.restore(); +}); + +const { checkMcpRateLimit, MCP_RATE_BUCKETS } = await import( + "@/lib/integrations/mcp/rate-limit" +); + +describe("checkMcpRateLimit", () => { + test("fails closed when Redis is configured but unreachable", async () => { + redisConfigured = true; + expect( + await checkMcpRateLimit("user-1", MCP_RATE_BUCKETS.read) + ).toEqual({ status: "unavailable" }); + }); + + test("stays open when Redis was never configured", async () => { + redisConfigured = false; + expect( + await checkMcpRateLimit("user-1", MCP_RATE_BUCKETS.read) + ).toEqual({ status: "allowed" }); + }); + + test("allows requests within the bucket", async () => { + redisConfigured = true; + redisClient = { + incr: () => Promise.resolve(5), + pExpire: () => Promise.resolve(1), + }; + expect( + await checkMcpRateLimit("user-1", MCP_RATE_BUCKETS.read) + ).toEqual({ status: "allowed" }); + }); + + test("limits requests over the bucket and reports the window", async () => { + redisConfigured = true; + redisClient = { + incr: () => Promise.resolve(121), + pExpire: () => Promise.resolve(1), + }; + expect( + await checkMcpRateLimit("user-1", MCP_RATE_BUCKETS.read) + ).toEqual({ retryAfterSeconds: 60, status: "limited" }); + }); + + test("fails closed when the counter command throws", async () => { + redisConfigured = true; + redisClient = { + incr: () => Promise.reject(new Error("connection reset")), + pExpire: () => Promise.resolve(1), + }; + expect( + await checkMcpRateLimit("user-1", MCP_RATE_BUCKETS.write) + ).toEqual({ status: "unavailable" }); + }); +}); diff --git a/lib/integrations/mcp/rate-limit.ts b/lib/integrations/mcp/rate-limit.ts index ddeefa2b..f1e4bae4 100644 --- a/lib/integrations/mcp/rate-limit.ts +++ b/lib/integrations/mcp/rate-limit.ts @@ -12,11 +12,15 @@ * constant). The window is short enough that an attacker can't amortize * the burst and long enough not to feel like a quota. * - * When Redis is unavailable, `incrementMcpRateCounter` returns `null` - * (fail-open) because we'd rather degrade UX than block legitimate user - * actions during an infrastructure blip. The danger is bounded: stealing - * a token still gives the user their existing library access; rate-limits - * are a defense in depth, not the only line. + * Redis is the counter's only home, so the limit cannot be enforced without + * it. When Redis is configured but unreachable this fails closed + * (`unavailable`) rather than letting the request through: the counter is + * the blast-radius control for a stolen token, and letting the request + * through would remove the only throttle. `getReadyRedisClient` waits out the + * connect window of a freshly created client, so a cold start is not read as + * an outage. A deployment with no `REDIS_URL` at all is a deliberate + * Redis-less setup, so that case stays fail-open; the Redis client logs it + * separately from an outage. * * No `import "server-only"` here on purpose: this module's only callers are * the MCP route handler (a Next.js server route); pulling in the client @@ -24,7 +28,10 @@ * marker would also keep us from unit-testing the decision helpers without a * preload hack. */ -import { getRedisClient } from "@/lib/common/redis"; +import { createLogger } from "@/lib/common/logs/console/logger"; +import { getReadyRedisClient, isRedisConfigured } from "@/lib/common/redis"; + +const log = createLogger("mcp.rate-limit"); const WINDOW_SECONDS = 60; @@ -40,52 +47,46 @@ export const MCP_RATE_BUCKETS = { write: { limit: 30, name: "write" }, } as const satisfies Record; -export interface RateLimitDecision { - count: number; - limit: number; - retryAfterSeconds: number; -} +export type McpRateLimitOutcome = + | { status: "allowed" } + | { status: "limited"; retryAfterSeconds: number } + | { status: "unavailable" }; /** - * Atomically increment the counter for the user's bucket and return the new - * count. Returns `null` when Redis isn't reachable so callers can decide - * whether to fail-open or surface the error. + * Count the request against the user's bucket and decide whether it is over + * the limit. Counts the request atomically and returns `unavailable` when the + * counter cannot be read, instead of throwing. */ -export async function incrementMcpRateCounter( +export async function checkMcpRateLimit( userId: string, bucket: Bucket -): Promise { - const redis = getRedisClient(); +): Promise { + const redis = await getReadyRedisClient(); if (!redis) { - return null; + return isRedisConfigured() + ? { status: "unavailable" } + : { status: "allowed" }; } - const key = `mcp:rate:${bucket.name}:${userId}`; - const count = await redis.incr(key); - if (count === 1) { - // First request in a fresh window — establish the TTL atomically. - // `pexpire` is preferred so a partial-second drift doesn't cut the - // window short; the worst case is we let a request through on the - // 60.0001-second boundary, which is fine for a defense-in-depth cap. - await redis.pExpire(key, WINDOW_SECONDS * 1000); - } - const retryAfterSeconds = count > bucket.limit ? WINDOW_SECONDS : 0; - return { - count, - limit: bucket.limit, - retryAfterSeconds, - }; -} -/** - * Centralized gate. `null` means Redis is absent — fall through and let the - * tool run. A `decision` with `count > limit` means fail-closed at this call. - */ -export function isOverLimit( - decision: RateLimitDecision | null, - bucket: Bucket -): boolean { - if (!decision) { - return false; + try { + const key = `mcp:rate:${bucket.name}:${userId}`; + const count = await redis.incr(key); + if (count === 1) { + // First request in a fresh window — establish the TTL atomically. + // `pexpire` is preferred so a partial-second drift doesn't cut the + // window short; the worst case is we let a request through on the + // 60.0001-second boundary, which is fine for a defense-in-depth cap. + await redis.pExpire(key, WINDOW_SECONDS * 1000); + } + return count > bucket.limit + ? { retryAfterSeconds: WINDOW_SECONDS, status: "limited" } + : { status: "allowed" }; + } catch (error) { + log.warn("MCP rate limit counter failed; failing closed", { + bucket: bucket.name, + error, + userId, + }); + return { status: "unavailable" }; } - return decision.count > bucket.limit; }