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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 0 additions & 16 deletions src/api/catalog.ts
Original file line number Diff line number Diff line change
Expand Up @@ -7,22 +7,6 @@ import type {
GatewayTestResponse,
} from "@/generated/types";

export interface OAuthUserTokenStatus {
status: "valid" | "near_expiry" | "expired" | "missing";
authorized: boolean;
scopes?: string[];
expires_at?: string | null;
updated_at?: string | null;
}

export interface OAuthGatewayStatus {
oauth_enabled: boolean;
grant_type?: string;
user_token_status?: OAuthUserTokenStatus;
}

export type OAuthGatewayStatusMap = Record<string, OAuthGatewayStatus>;

export interface GatewayImpactPreview {
gatewayId: string;
servers: Array<{ id: string; name: string }>;
Expand Down
120 changes: 120 additions & 0 deletions src/api/oauth.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,120 @@
import { describe, expect, it, vi } from "vitest";
import { http, HttpResponse } from "msw";

import { server } from "@/test/mocks/server";
import { getOAuthStatuses, normalizeGatewayIds } from "./oauth";

describe("normalizeGatewayIds", () => {
it("trims, removes empty values, deduplicates, and preserves order", () => {
expect(normalizeGatewayIds([" second ", "", "first", "second", " "])).toEqual([
"second",
"first",
]);
});
});

describe("getOAuthStatuses", () => {
it("returns immediately for empty input", async () => {
const request = vi.fn();
server.use(
http.get("*/api/oauth/status", () => {
request();
return HttpResponse.json({});
}),
);

await expect(getOAuthStatuses([])).resolves.toEqual({ statuses: {}, failures: {} });
expect(request).not.toHaveBeenCalled();
});

it("deduplicates IDs and keeps 100 IDs in one request", async () => {
const seen: string[][] = [];
const ids = Array.from({ length: 100 }, (_, index) => `gateway-${index}`);

server.use(
http.get("*/api/oauth/status", ({ request }) => {
seen.push(new URL(request.url).searchParams.getAll("gateway_ids"));
return HttpResponse.json({});
}),
);

await getOAuthStatuses([...ids, "gateway-99", " "]);

expect(seen).toHaveLength(1);
expect(seen[0]).toHaveLength(100);
expect(seen[0]).toContain("gateway-99");
});

it("rejects invalid gateway IDs before making a request", async () => {
const request = vi.fn();
server.use(
http.get("*/api/oauth/status", () => {
request();
return HttpResponse.json({});
}),
);

await expect(getOAuthStatuses(["gateway/../secret"])).rejects.toThrow(
"Invalid server ID format",
);
expect(request).not.toHaveBeenCalled();
});

it("chunks 201 IDs and preserves successful batches when one batch fails", async () => {
const batchSizes: number[] = [];
server.use(
http.get("*/api/oauth/status", ({ request }) => {
const ids = new URL(request.url).searchParams.getAll("gateway_ids");
batchSizes.push(ids.length);
if (ids[0] === "gateway-100") {
return HttpResponse.json({ detail: "failed" }, { status: 500 });
}
return HttpResponse.json(
Object.fromEntries(
ids.map((id) => [
id,
{
oauth_enabled: true,
grant_type: "authorization_code",
user_token_status: { status: "valid", authorized: true },
},
]),
),
);
}),
);

const result = await getOAuthStatuses(
Array.from({ length: 201 }, (_, index) => `gateway-${index}`),
);

expect(batchSizes).toEqual([100, 100, 1]);
expect(Object.keys(result.statuses)).toHaveLength(101);
expect(result.statuses["gateway-0"]?.user_token_status?.status).toBe("valid");
expect(result.statuses["gateway-200"]?.user_token_status?.status).toBe("valid");
expect(Object.keys(result.failures)).toHaveLength(100);
expect(result.failures["gateway-100"]).toEqual({ retryable: true, status: 500 });
});

it("marks a forbidden batch unavailable and non-retryable", async () => {
server.use(
http.get("*/api/oauth/status", () =>
HttpResponse.json({ detail: "forbidden" }, { status: 403 }),
),
);

await expect(getOAuthStatuses(["gateway-1"])).resolves.toEqual({
statuses: {},
failures: { "gateway-1": { retryable: false, status: 403 } },
});
});

it("propagates abort instead of converting it to a batch failure", async () => {
const controller = new AbortController();
controller.abort();

await expect(getOAuthStatuses(["gateway-1"], controller.signal)).rejects.toMatchObject({
name: "AbortError",
});
});
});
91 changes: 91 additions & 0 deletions src/api/oauth.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,91 @@
import { api, ApiError } from "./client";
import { validateServerId } from "@/utils/serverId";

export type OAuthTokenStatus = "valid" | "near_expiry" | "expired" | "missing" | "unknown";

export interface OAuthUserTokenStatus {
status: OAuthTokenStatus;
authorized: boolean;
scopes?: string[];
expires_at?: string | null;
updated_at?: string | null;
}

export interface OAuthGatewayStatus {
oauth_enabled: boolean;
grant_type?: string;
authorization_url?: string;
message?: string;
user_token_status?: OAuthUserTokenStatus;
}

export type OAuthGatewayStatusMap = Record<string, OAuthGatewayStatus>;

export interface OAuthStatusFailure {
retryable: boolean;
status?: number;
}

export interface OAuthStatusBatchResult {
statuses: OAuthGatewayStatusMap;
failures: Record<string, OAuthStatusFailure>;
}

const OAUTH_STATUS_MAX_IDS = 100;

export function normalizeGatewayIds(gatewayIds: string[]): string[] {
return [
...new Set(
gatewayIds
.map((id) => id.trim())
.filter(Boolean)
.map(validateServerId),
),
];
}

/** Fetch caller-scoped OAuth state, respecting the backend's 100-id batch cap. */
export async function getOAuthStatuses(
gatewayIds: string[],
signal?: AbortSignal,
): Promise<OAuthStatusBatchResult> {
const ids = normalizeGatewayIds(gatewayIds);
if (ids.length === 0) return { statuses: {}, failures: {} };

const batches: string[][] = [];
for (let start = 0; start < ids.length; start += OAUTH_STATUS_MAX_IDS) {
batches.push(ids.slice(start, start + OAUTH_STATUS_MAX_IDS));
}

const settled = await Promise.allSettled(
batches.map(async (batch) => {
const params = new URLSearchParams();
batch.forEach((id) => params.append("gateway_ids", id));
const statuses = await api.get<OAuthGatewayStatusMap>(
`/oauth/status?${params.toString()}`,
undefined,
signal,
);
return { batch, statuses };
}),
);

if (signal?.aborted) throw new DOMException("Aborted", "AbortError");

const result: OAuthStatusBatchResult = { statuses: {}, failures: {} };
settled.forEach((entry, index) => {
const batch = batches[index];
if (entry.status === "fulfilled") {
Object.assign(result.statuses, entry.value.statuses);
return;
}

const status = entry.reason instanceof ApiError ? entry.reason.status : undefined;
const failure = { retryable: status !== 403, ...(status === undefined ? {} : { status }) };
batch.forEach((id) => {
result.failures[id] = failure;
});
});

return result;
}
55 changes: 0 additions & 55 deletions src/api/servers.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -499,61 +499,6 @@ describe("serversApi", () => {
});
});

describe("getOAuthStatus", () => {
const statusResponse = (body: unknown) =>
new Response(JSON.stringify(body), {
status: 200,
headers: { "Content-Type": "application/json" },
});

const idsFrom = (call: unknown[]) =>
new URL(String(call[0]), "http://localhost").searchParams.getAll("gateway_ids");

it("sends one request and returns its statuses", async () => {
const body = { "srv-1": { user_token_status: { status: "missing" } } };
mockFetch.mockResolvedValueOnce(statusResponse(body));

const result = await serversApi.getOAuthStatus(["srv-1"]);

expect(result).toEqual(body);
expect(mockFetch).toHaveBeenCalledTimes(1);
expect(idsFrom(mockFetch.mock.calls[0])).toEqual(["srv-1"]);
});

it("splits over 100 ids across requests and merges the responses", async () => {
const ids = Array.from({ length: 101 }, (_, index) => `srv-${index}`);
mockFetch
.mockResolvedValueOnce(
statusResponse({ "srv-0": { user_token_status: { status: "valid" } } }),
)
.mockResolvedValueOnce(
statusResponse({ "srv-100": { user_token_status: { status: "expired" } } }),
);

const result = await serversApi.getOAuthStatus(ids);

expect(mockFetch).toHaveBeenCalledTimes(2);
expect(idsFrom(mockFetch.mock.calls[0])).toHaveLength(100);
expect(idsFrom(mockFetch.mock.calls[1])).toEqual(["srv-100"]);
expect(result).toEqual({
"srv-0": { user_token_status: { status: "valid" } },
"srv-100": { user_token_status: { status: "expired" } },
});
});

it("makes no request for an empty id list", async () => {
await expect(serversApi.getOAuthStatus([])).resolves.toEqual({});
expect(mockFetch).not.toHaveBeenCalled();
});

it("rejects a malformed id before any request", async () => {
await expect(serversApi.getOAuthStatus(["../etc"])).rejects.toThrow(
"Invalid server ID format",
);
expect(mockFetch).not.toHaveBeenCalled();
});
});

describe("testConnection", () => {
it("POSTs /v1/mcp-servers/:id/test and returns the result", async () => {
mockFetch.mockResolvedValueOnce(
Expand Down
50 changes: 2 additions & 48 deletions src/api/servers.ts
Original file line number Diff line number Diff line change
Expand Up @@ -6,20 +6,18 @@
*/

import { api } from "./client";
import type { ServersResponse, MCPServer, GatewayOAuthStatus } from "../types/server";
import type { ServersResponse, MCPServer } from "../types/server";
import type {
GatewayHandshakeRequest,
GatewayHandshakeResponse,
GatewayRefreshResponse,
GatewayTestRequest,
GatewayTestResponse,
} from "@/generated/types";
import { validateServerId } from "@/utils/serverId";

const serverByIdRequestCache = new Map<string, Promise<MCPServer>>();

/** Mirrors OAUTH_STATUS_BATCH_MAX_IDS on the backend's /oauth/status route. */
const OAUTH_STATUS_MAX_IDS = 100;

/** The user closed the OAuth popup. Typed so callers can stay quiet about it. */
export class OAuthCancelledError extends Error {
constructor() {
Expand All @@ -28,25 +26,6 @@ export class OAuthCancelledError extends Error {
}
}

/**
* Validates server ID to prevent path traversal and injection attacks
* @param id - The server ID to validate
* @returns The validated ID
* @throws Error if ID is invalid
*/
function validateServerId(id: string): string {
if (!id || typeof id !== "string") {
throw new Error("Invalid server ID");
}

// Ensure ID is alphanumeric with hyphens/underscores only
if (!/^[a-zA-Z0-9_-]+$/.test(id)) {
throw new Error("Invalid server ID format");
}

return id;
}

function openOAuthAuthorizationPopup(): Window {
const width = 600;
const height = 700;
Expand Down Expand Up @@ -214,31 +193,6 @@ export const serversApi = {
return api.post(`/oauth/fetch-tools/${validId}`);
},

/**
* The caller's own OAuth state for each gateway, batched.
*
* Keys stay snake_case, unlike the gateway endpoints. Ids that are missing or
* not visible to the caller are omitted. A paged-through list can exceed the
* backend's id cap, so requests are split and the responses merged.
*/
getOAuthStatus: async (ids: string[]): Promise<Record<string, GatewayOAuthStatus>> => {
const validIds = ids.map(validateServerId);
const batches: string[][] = [];
for (let start = 0; start < validIds.length; start += OAUTH_STATUS_MAX_IDS) {
batches.push(validIds.slice(start, start + OAUTH_STATUS_MAX_IDS));
}

const responses = await Promise.all(
batches.map((batch) => {
const params = new URLSearchParams();
batch.forEach((id) => params.append("gateway_ids", id));
return api.get<Record<string, GatewayOAuthStatus>>(`/oauth/status?${params.toString()}`);
}),
);

return Object.assign({}, ...responses) as Record<string, GatewayOAuthStatus>;
},

/**
* Open blank OAuth popup during an active user gesture.
*
Expand Down
Loading