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
12 changes: 10 additions & 2 deletions src/auth/oauth.ts
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ type TokenResponse = {
};

const CALLBACK_HOST = "127.0.0.1";
const CALLBACK_PATH = "/callback";
const DEFAULT_CALLBACK_PORT = 53682;
const OAUTH_CALLBACK_PORT_ENV_KEY = "OPENWIKI_OAUTH_CALLBACK_PORT";
const HTTPS_OAUTH_REDIRECT_URI_ENV_KEY = "OPENWIKI_HTTPS_OAUTH_REDIRECT_URI";
Expand Down Expand Up @@ -404,6 +405,13 @@ export async function createCallbackServer(
request.url ?? "/",
`http://${CALLBACK_HOST}:${callbackPort}`,
);

if (requestUrl.pathname !== CALLBACK_PATH) {
response.writeHead(404, getCallbackResponseHeaders());
response.end("OpenWiki OAuth callback server is waiting for /callback.");
return;
}

const code = requestUrl.searchParams.get("code");
const state = requestUrl.searchParams.get("state");
const error = requestUrl.searchParams.get("error");
Expand Down Expand Up @@ -438,7 +446,7 @@ export async function createCallbackServer(
if (!address || typeof address === "string") {
throw new Error("Could not start OAuth callback server.");
}
const localRedirectUri = `http://${CALLBACK_HOST}:${address.port}/callback`;
const localRedirectUri = `http://${CALLBACK_HOST}:${address.port}${CALLBACK_PATH}`;

return {
close: () => closeCallbackServer(server),
Expand Down Expand Up @@ -524,7 +532,7 @@ function getProviderRedirectUri(

const url = new URL(override);

if (url.pathname !== "/callback") {
if (url.pathname !== CALLBACK_PATH) {
throw new Error(
`${HTTPS_OAUTH_REDIRECT_URI_ENV_KEY} must end with /callback.`,
);
Expand Down
61 changes: 59 additions & 2 deletions test/oauth-callback-server.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,8 @@ describe("createCallbackServer", () => {
const callback = await createCallbackServer(getAuthProvider("gmail"));

try {
expect(callback.redirectUri).toBe(`http://127.0.0.1:${port}/callback`);

const codePromise = callback.waitForCode("expected-state");
const redirect = await fetch(
`http://127.0.0.1:${port}/callback?code=test-code&state=expected-state`,
Expand All @@ -58,6 +60,61 @@ describe("createCallbackServer", () => {
}
});

test("ignores non-callback requests without rejecting the pending flow", async () => {
const callback = await createCallbackServer(getAuthProvider("gmail"));

try {
const codePromise = callback.waitForCode("expected-state");
const strayRequest = await fetch(`http://127.0.0.1:${port}/favicon.ico`);

expect(strayRequest.status).toBe(404);
await expect(
fetch(
`http://127.0.0.1:${port}/callback?code=test-code&state=expected-state`,
),
).resolves.toMatchObject({ status: 200 });
await expect(codePromise).resolves.toBe("test-code");
} finally {
await callback.close();
}
});

test("rejects a callback request that is missing code or state", async () => {
const callback = await createCallbackServer(getAuthProvider("gmail"));

try {
const codePromise = callback.waitForCode("expected-state");
const codeRejection = expect(codePromise).rejects.toThrow(
"OAuth callback was missing code or state.",
);
const redirect = await fetch(`http://127.0.0.1:${port}/callback`);

expect(redirect.status).toBe(400);
await codeRejection;
} finally {
await callback.close();
}
});

test("rejects a callback request with a provider error", async () => {
const callback = await createCallbackServer(getAuthProvider("gmail"));

try {
const codePromise = callback.waitForCode("expected-state");
const codeRejection = expect(codePromise).rejects.toThrow(
"OAuth provider returned error: access_denied",
);
const redirect = await fetch(
`http://127.0.0.1:${port}/callback?error=access_denied`,
);

expect(redirect.status).toBe(400);
await codeRejection;
} finally {
await callback.close();
}
});

test("answers a trailing request that arrives while the server is closing", async () => {
const callback = await createCallbackServer(getAuthProvider("gmail"));
const codePromise = callback.waitForCode("expected-state");
Expand All @@ -83,7 +140,7 @@ describe("createCallbackServer", () => {
await closePromise;

const response = Buffer.concat(chunks).toString();
expect(response).toMatch(/^HTTP\/1\.1 400 /);
expect(response).toContain("missing required data");
expect(response).toMatch(/^HTTP\/1\.1 404 /);
expect(response).toContain("waiting for /callback");
});
});