fix: restart abandoned provider sign-ins with shared OAuth handling (#160574)

* fix(auth): allow retrying abandoned Model Setup sign-ins

* style(auth): remove redundant retry cleanup branch

* refactor(auth): share OAuth callbacks and fence queued retries
This commit is contained in:
stevenlee-oai 2026-09-28 19:40:37 -07:00 • committed by GitHub
parent 58b1860232
commit 2079ed92a5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
19 changed files with 1418 additions and 446 deletions

View file

@ -158,6 +158,17 @@ and [Z.AI / GLM Coding Plan](/providers/zai).
OpenClaw's OAuth registry and adapters live in `src/llm/utils/oauth/`. Shared provider helpers live in `src/plugin-sdk/provider-oauth-runtime.ts` and `src/plugin-sdk/provider-auth-runtime.ts`. The auth commands in `src/commands/models/auth.ts` run the selected provider method and persist the returned profiles.
### Restarting sign-in in Model Setup
In **Model Setup**, starting the same sign-in again replaces your unfinished
attempt for the same agent and workspace. This includes **Sign in with ChatGPT**
and other provider sign-in flows. Use the newest browser link; callbacks from
the previous attempt are no longer accepted.
Another user's sign-in, a different setup flow, or an attempt already saving
credentials or configuration stays protected. Let that operation finish before
retrying.
### Anthropic setup-token
Flow shape:

View file

@ -53,6 +53,25 @@ Unavailable storage or an unusable matching OAuth profile continues to interacti
sign-in. A matching account identity alone does not make expired credentials usable.
A failed selected import stops the operation instead of silently starting a different login.
## Loopback OAuth callbacks
Bundled providers use `startProviderOAuthLoopbackCallbackServer` from
`openclaw/plugin-sdk/provider-auth-runtime` to bind their callback before opening
the browser. `waitForCallback()` returns either an OAuth error or a validated
code/state pair with `parameters: URLSearchParams` for provider-specific fields.
Repeated parameters remain available for the provider to validate.
The default response acknowledges the callback and closes the listener. Set
`deferResponse: true` to finish token exchange and identity checks before calling
`complete({ status, body, contentType })`. Abort, optional `timeoutMs`, and browser
disconnection still close a deferred response; a late `complete()` is a no-op.
Always call `close()` in `finally`. The caller's signal and authority checks own
token requests and persistence; the listener deadline does not cancel that work.
By default, the listener binds every loopback address resolved for the redirect
hostname. `bindHostname` adds a loopback host. Use `bindOnlyHostname` instead to
preserve a provider's exact Node bind host (`localhost`, `127.0.0.1`, or `::1`).
## Handle model access after sign-in
Existing consumers of `runModelsAuthLoginFlow` from

View file

@ -310,9 +310,17 @@ describe("loginOpenAICodexDeviceCode", () => {
const oauthTokenRequest = fetchCall(fetchMock, 3);
expect(oauthTokenRequest[0]).toBe("https://auth.openai.com/oauth/token");
expect(oauthTokenRequest[1]?.method).toBe("POST");
expect(await new Response(oauthTokenRequest[1]?.body).text()).toBe(
"grant_type=authorization_code&code=authorization-code-123&redirect_uri=https%3A%2F%2Fauth.openai.com%2Fdeviceauth%2Fcallback&client_id=app_EMoamEEZ73f0CkXaXp7hrann&code_verifier=code-verifier-123",
);
expect(
Object.fromEntries(
new URLSearchParams(await new Response(oauthTokenRequest[1]?.body).text()),
),
).toEqual({
grant_type: "authorization_code",
code: "authorization-code-123",
redirect_uri: "https://auth.openai.com/deviceauth/callback",
client_id: "app_EMoamEEZ73f0CkXaXp7hrann",
code_verifier: "code-verifier-123",
});
expect(oauthTokenRequest[1]?.signal).toBeInstanceOf(AbortSignal);
expect(oauthTokenRequest[1]?.headers).toEqual({
"Content-Type": "application/x-www-form-urlencoded",

View file

@ -10,11 +10,14 @@ import {
import { resolveOpenAICodexAccessTokenExpiry } from "openclaw/plugin-sdk/provider-auth";
import { readResponseTextLimited } from "openclaw/plugin-sdk/provider-http";
import { classifyTransientNetworkErrorCode } from "openclaw/plugin-sdk/retry-runtime";
import { fetchWithSsrFGuard } from "openclaw/plugin-sdk/ssrf-runtime";
import {
asNullableObjectRecord,
normalizeOptionalString,
} from "openclaw/plugin-sdk/string-coerce-runtime";
import {
createOpenAIAuthorizationCodeForm,
withOpenAIOAuthResponse,
} from "./openai-oauth-http.runtime.js";
const OPENAI_AUTH_BASE_URL = "https://auth.openai.com";
const OPENAI_CODEX_CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann";
@ -151,13 +154,11 @@ async function runOpenAICodexDeviceRequest(params: {
requireHttps: true,
auditContext: "openai-chatgpt-device-code",
};
const { response, release } = await fetchWithSsrFGuard(
return await withOpenAIOAuthResponse(
shouldUseEnvHttpProxyForUrl(params.url)
? withTrustedEnvProxyGuardedFetchMode(guardedOptions)
: guardedOptions,
);
try {
return {
async (response) => ({
ok: response.ok,
status: response.status,
bodyText: await readResponseTextLimited(
@ -167,10 +168,8 @@ async function runOpenAICodexDeviceRequest(params: {
: OPENAI_CODEX_DEVICE_ERROR_BODY_LIMIT_BYTES,
{ chunkTimeoutMs: params.timeoutMs },
),
};
} finally {
await release();
}
}),
);
}
async function fetchOpenAICodexDeviceCode(params: {
@ -353,12 +352,11 @@ async function exchangeOpenAICodexDeviceCode(params: {
init: {
method: "POST",
headers: resolveOpenAICodexDeviceCodeHeaders("application/x-www-form-urlencoded"),
body: new URLSearchParams({
grant_type: "authorization_code",
body: createOpenAIAuthorizationCodeForm({
code: params.authorizationCode,
redirect_uri: OPENAI_CODEX_DEVICE_CALLBACK_URL,
client_id: OPENAI_CODEX_CLIENT_ID,
code_verifier: params.codeVerifier,
redirectUri: OPENAI_CODEX_DEVICE_CALLBACK_URL,
clientId: OPENAI_CODEX_CLIENT_ID,
verifier: params.codeVerifier,
}),
},
timeoutOperation: "token exchange",

View file

@ -251,6 +251,44 @@ describe("OpenAI Codex OAuth flow", () => {
expect(onAuth).not.toHaveBeenCalled();
});
it.each(["prompt", "manual input"] as const)(
"uses %s when the callback port is already occupied",
async (entry) => {
const server = createServer();
await new Promise<void>((resolve, reject) => {
server.once("error", reject);
server.listen(1455, resolveOpenAICallbackHost(), resolve);
});
const onAuth = vi.fn();
const onPrompt = vi.fn(async () => "manual-code");
mockTokenResponse({
access_token: fakeJwt({
"https://api.openai.com/auth": { chatgpt_account_id: "manual-account" },
}),
refresh_token: "test-refresh-token",
expires_in: 3600,
});
try {
await expect(
loginOpenAICodex({
onAuth,
onPrompt,
...(entry === "manual input" ? { onManualCodeInput: async () => "manual-code" } : {}),
}),
).resolves.toMatchObject({ accountId: "manual-account" });
expect(onAuth).toHaveBeenCalledOnce();
if (entry === "manual input") {
expect(onPrompt).not.toHaveBeenCalled();
}
expect(ssrfMocks.fetchWithSsrFGuard.mock.calls[0]?.[0].init.body.get("code")).toBe(
"manual-code",
);
} finally {
await closeServer(server);
}
},
);
it.each(["callback", "manual input", "transport preparation"] as const)(
"revalidates live authority after held %s before exchanging the code",
async (boundary) => {
@ -379,6 +417,56 @@ describe("OpenAI Codex OAuth flow", () => {
}
});
it("restarts browser login while a cancelled token exchange still releases its transport", async () => {
const releasing = createDeferred<void>();
const released = createDeferred<void>();
const firstController = new AbortController();
const nextController = new AbortController();
const agent = new Agent();
const token = (accountId: string) => ({
access_token: fakeJwt({ "https://api.openai.com/auth": { chatgpt_account_id: accountId } }),
refresh_token: "test-refresh-token",
expires_in: 3600,
});
ssrfMocks.fetchWithSsrFGuard.mockResolvedValueOnce({
response: new Response(JSON.stringify(token("old-account"))),
release: async () => {
releasing.resolve();
await released.promise;
},
});
mockTokenResponse(token("new-account"));
const onAuth = async ({ url }: { url: string }) => {
const authorization = new URL(url);
const callback = new URL(authorization.searchParams.get("redirect_uri")!);
callback.searchParams.set("state", authorization.searchParams.get("state")!);
callback.searchParams.set("code", "callback-code");
const response = await requestCallback(callback.toString(), agent);
expect(response.body).toContain("OpenAI authentication completed");
};
const onPrompt = vi.fn(async () => {
throw new Error("Browser login must not fall back to manual input");
});
const first = loginOpenAICodex({ onAuth, onPrompt, signal: firstController.signal }).catch(
() => undefined,
);
try {
await releasing.promise;
firstController.abort();
await expect(
loginOpenAICodex({ onAuth, onPrompt, signal: nextController.signal }),
).resolves.toMatchObject({ accountId: "new-account" });
expect(onPrompt).not.toHaveBeenCalled();
expect(ssrfMocks.fetchWithSsrFGuard).toHaveBeenCalledTimes(2);
} finally {
firstController.abort();
nextController.abort();
released.resolve();
await first;
agent.destroy();
}
});
it("waits for Node OAuth runtime before creating an authorization flow", async () => {
const callbackHost = resolveOpenAICallbackHost();
const flow = await createOpenAIAuthorizationFlow(

View file

@ -20,12 +20,15 @@ import {
refreshOpenAIAccessToken,
} from "./openai-chatgpt-oauth-token.runtime.js";
const CALLBACK_PORT = 1455;
const CALLBACK_HOST = resolveOpenAICallbackHost();
const REDIRECT_URI = resolveOpenAIRedirectUri(CALLBACK_HOST);
const MANUAL_PROMPT_FALLBACK_MS = 15_000;
const loadNodeOAuthHttp = createLazyRuntimeModule(() => import("node:http"));
const loadOAuthCallbackServer = createLazyRuntimeModule(() =>
import("openclaw/plugin-sdk/provider-auth-runtime").then(
({ startProviderOAuthLoopbackCallbackServer }) => startProviderOAuthLoopbackCallbackServer,
),
);
function waitForManualPromptFallback(signal?: AbortSignal): Promise<null> {
return new Promise((resolve, reject) => {
@ -70,99 +73,6 @@ async function promptForAuthorizationCode(
);
}
type OAuthServerInfo = {
close: () => void;
cancelWait: () => void;
waitForCode: () => Promise<{ code: string } | null>;
};
function sendOAuthHtmlResponse(
res: import("node:http").ServerResponse,
statusCode: number,
html: string,
): void {
res.statusCode = statusCode;
// Callback browsers may reuse HTTP/1.1 connections. Force disconnect after
// the response so an accepted socket cannot keep the auth process alive.
res.setHeader("Connection", "close");
res.setHeader("Content-Type", "text/html; charset=utf-8");
res.end(html);
}
async function startLocalOAuthServer(
state: string,
assertCurrent?: () => void,
): Promise<OAuthServerInfo> {
const http = await loadNodeOAuthHttp();
assertCurrent?.();
let settleWait: ((value: { code: string } | null) => void) | undefined;
const waitForCodePromise = new Promise<{ code: string } | null>((resolve) => {
settleWait = resolve;
});
const server = http.createServer((req, res) => {
try {
const url = new URL(req.url || "", "http://localhost");
if (url.pathname !== "/auth/callback") {
sendOAuthHtmlResponse(res, 404, oauthErrorHtml("Callback route not found."));
return;
}
if (url.searchParams.get("state") !== state) {
sendOAuthHtmlResponse(res, 400, oauthErrorHtml("State mismatch."));
return;
}
const code = url.searchParams.get("code");
if (!code) {
sendOAuthHtmlResponse(res, 400, oauthErrorHtml("Missing authorization code."));
return;
}
sendOAuthHtmlResponse(
res,
200,
oauthSuccessHtml("OpenAI authentication completed. You can close this window."),
);
settleWait?.({ code });
} catch {
sendOAuthHtmlResponse(
res,
500,
oauthErrorHtml("Internal error while processing OAuth callback."),
);
}
});
return new Promise((resolve) => {
server
.listen(CALLBACK_PORT, CALLBACK_HOST, () => {
resolve({
close: () => {
server.close();
// Force-close preconnected sockets so they cannot pin the CLI process.
server.closeAllConnections();
},
cancelWait: () => {
settleWait?.(null);
},
waitForCode: () => waitForCodePromise,
});
})
.on("error", () => {
settleWait?.(null);
resolve({
close: () => {
try {
server.close();
} catch {
// ignore
}
},
cancelWait: () => {},
waitForCode: async () => null,
});
});
});
}
function resolveOpenAICredentials(
result: Awaited<ReturnType<typeof refreshOpenAIAccessToken>>,
): OAuthCredentials {
@ -216,8 +126,34 @@ export async function loginOpenAICodex(options: {
options.originator ?? "openclaw",
REDIRECT_URI,
);
const server = await startLocalOAuthServer(state, options.assertCurrent);
const startCallbackServer = await loadOAuthCallbackServer();
options.assertCurrent?.();
throwIfOAuthLoginAborted(options.signal);
let server: Awaited<ReturnType<typeof startCallbackServer>> | undefined;
try {
server = await startCallbackServer({
redirectUrl: REDIRECT_URI,
expectedState: state,
bindOnlyHostname: CALLBACK_HOST,
signal: options.signal,
renderSuccess: () => ({
body: oauthSuccessHtml("OpenAI authentication completed. You can close this window."),
contentType: "text/html; charset=utf-8",
}),
renderError: (message) => ({
body: oauthErrorHtml(message),
contentType: "text/html; charset=utf-8",
}),
});
} catch {
// An unavailable callback port still permits manual entry; retired owners do not.
options.assertCurrent?.();
throwIfOAuthLoginAborted(options.signal);
}
let cancelWait!: () => void;
const cancelledWait = new Promise<null>((resolve) => {
cancelWait = () => resolve(null);
});
let code: string | undefined;
try {
options.assertCurrent?.();
@ -230,9 +166,21 @@ export async function loginOpenAICodex(options: {
}),
),
options.signal,
server.cancelWait,
cancelWait,
);
throwIfOAuthLoginAborted(options.signal);
const callbackPromise = Promise.race([
server
? server.waitForCallback().then((result) => {
if (result.type === "oauth_error") {
throw new Error("OpenAI authorization was not completed.");
}
return { code: result.code };
})
: Promise.resolve(null),
cancelledWait,
]);
void callbackPromise.catch(() => undefined);
if (options.onManualCodeInput) {
let manualCode: string | undefined;
@ -241,21 +189,17 @@ export async function loginOpenAICodex(options: {
.onManualCodeInput()
.then((input) => {
manualCode = input;
server.cancelWait();
cancelWait();
})
.catch((err: unknown) => {
manualError = err instanceof Error ? err : new Error(String(err));
server.cancelWait();
cancelWait();
});
const result = await withOAuthLoginAbort(
server.waitForCode(),
options.signal,
server.cancelWait,
);
const result = await withOAuthLoginAbort(callbackPromise, options.signal, cancelWait);
if (!result?.code && !manualCode && !manualError) {
await withOAuthLoginAbort(manualPromise, options.signal, server.cancelWait);
await withOAuthLoginAbort(manualPromise, options.signal, cancelWait);
}
if (manualError) {
throw manualError;
@ -266,25 +210,24 @@ export async function loginOpenAICodex(options: {
code = parseAuthorizationCode(manualCode, state);
}
} else {
const callbackPromise = server.waitForCode();
const result = await withOAuthLoginAbort(
Promise.race([callbackPromise, waitForManualPromptFallback(options.signal)]),
options.signal,
server.cancelWait,
cancelWait,
);
if (result?.code) {
code = result.code;
} else {
const promptCodePromise = promptForAuthorizationCode(options.onPrompt, state).then(
(promptCode) => {
server.cancelWait();
cancelWait();
return promptCode;
},
);
code = await withOAuthLoginAbort(
Promise.race([callbackPromise.then((callback) => callback?.code), promptCodePromise]),
options.signal,
server.cancelWait,
cancelWait,
);
}
}
@ -293,7 +236,7 @@ export async function loginOpenAICodex(options: {
code = await withOAuthLoginAbort(
promptForAuthorizationCode(options.onPrompt, state),
options.signal,
server.cancelWait,
cancelWait,
);
}
@ -308,7 +251,7 @@ export async function loginOpenAICodex(options: {
}),
);
} finally {
server.close();
await server?.close();
}
}

View file

@ -5,13 +5,17 @@ import {
} from "openclaw/plugin-sdk/provider-oauth-runtime";
import { readResponseWithLimit } from "openclaw/plugin-sdk/response-limit-runtime";
import { redactSensitiveText } from "openclaw/plugin-sdk/security-runtime";
import { fetchWithSsrFGuard, type SsrFPolicy } from "openclaw/plugin-sdk/ssrf-runtime";
import type { SsrFPolicy } from "openclaw/plugin-sdk/ssrf-runtime";
import {
asOptionalRecord,
isRecord,
normalizeBoundedOptionalString,
normalizeOptionalString,
} from "openclaw/plugin-sdk/string-coerce-runtime";
import {
createOpenAIAuthorizationCodeForm,
withOpenAIOAuthResponse,
} from "./openai-oauth-http.runtime.js";
const CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann";
const TOKEN_URL = "https://auth.openai.com/oauth/token";
@ -154,61 +158,64 @@ function formatTokenRequestError(
);
}
async function postTokenForm(
async function requestOpenAIToken(
body: URLSearchParams,
operation: "exchange" | "refresh",
options: TokenRequestOptions = {},
): Promise<Response> {
existingRefreshToken?: string,
): Promise<TokenResult> {
const timeoutMs = options.timeoutMs ?? TOKEN_REQUEST_TIMEOUT_MS;
throwIfOAuthLoginAborted(options.signal);
const { response, release } = await fetchWithSsrFGuard({
url: TOKEN_URL,
// Match device-code login's operator proxy policy. The guard keeps direct DNS
// pinning when no proxy applies; the exact-host policy also permits fake-IP DNS.
mode: "trusted_env_proxy",
policy: OAUTH_TOKEN_SSRF_POLICY,
init: {
method: "POST",
headers: { "Content-Type": "application/x-www-form-urlencoded" },
body,
},
timeoutMs,
signal: options.signal,
beforeRequest: options.assertCurrent,
auditContext: "openai-chatgpt-oauth-token",
});
try {
const responseBody = await readResponseWithLimit(
response,
OAUTH_TOKEN_RESPONSE_BODY_LIMIT_BYTES,
throwIfOAuthLoginAborted(options.signal);
const { response: tokenResponse, text } = await withOpenAIOAuthResponse(
{
onOverflow: ({ size, maxBytes }) =>
new Error(
`OpenAI Codex OAuth token response body too large: ${size} bytes (limit: ${maxBytes} bytes)`,
),
url: TOKEN_URL,
// Keep direct DNS pinning without a proxy; the exact-host policy permits fake-IP DNS.
mode: "trusted_env_proxy",
policy: OAUTH_TOKEN_SSRF_POLICY,
init: {
method: "POST",
headers: { "Content-Type": "application/x-www-form-urlencoded" },
body,
},
timeoutMs,
signal: options.signal,
beforeRequest: options.assertCurrent,
auditContext: "openai-chatgpt-oauth-token",
},
async (response) => {
const bytes = await readResponseWithLimit(response, OAUTH_TOKEN_RESPONSE_BODY_LIMIT_BYTES, {
onOverflow: ({ size, maxBytes }) =>
new Error(
`OpenAI Codex OAuth token response body too large: ${size} bytes (limit: ${maxBytes} bytes)`,
),
});
return { response, text: new TextDecoder().decode(bytes) };
},
);
return new Response(new Uint8Array(responseBody), {
status: response.status,
statusText: response.statusText,
headers: response.headers,
});
} finally {
await release();
return readOpenAITokenResponse(tokenResponse, text, operation, existingRefreshToken);
} catch (error) {
return {
type: "failed",
operation,
...(options.signal?.aborted ? { cancelled: true } : {}),
summary: formatTokenRequestError(operation, error, timeoutMs, options.signal),
};
}
}
async function readOpenAITokenResponse(
function readOpenAITokenResponse(
response: Response,
text: string,
operation: "exchange" | "refresh",
existingRefreshToken?: string,
): Promise<TokenResult> {
): TokenResult {
if (!response.ok) {
const text = await response.text().catch(() => "");
return buildTokenResponseFailure({ response, operation, text });
}
let json: TokenResponseJson;
try {
json = (await response.json()) as TokenResponseJson;
json = JSON.parse(text) as TokenResponseJson;
} catch {
return {
type: "failed",
@ -246,51 +253,25 @@ export async function exchangeOpenAIAuthorizationCode(
redirectUri: string,
options: TokenRequestOptions = {},
): Promise<TokenResult> {
const timeoutMs = options.timeoutMs ?? TOKEN_REQUEST_TIMEOUT_MS;
let response: Response;
try {
response = await postTokenForm(
new URLSearchParams({
grant_type: "authorization_code",
client_id: CLIENT_ID,
code,
code_verifier: verifier,
redirect_uri: redirectUri,
}),
{ ...options, timeoutMs },
);
} catch (error) {
return {
type: "failed",
operation: "exchange",
...(options.signal?.aborted ? { cancelled: true } : {}),
summary: formatTokenRequestError("exchange", error, timeoutMs, options.signal),
};
}
return await readOpenAITokenResponse(response, "exchange");
return await requestOpenAIToken(
createOpenAIAuthorizationCodeForm({ clientId: CLIENT_ID, code, verifier, redirectUri }),
"exchange",
options,
);
}
export async function refreshOpenAIAccessToken(
refreshToken: string,
options: TokenRequestOptions = {},
): Promise<TokenResult> {
const timeoutMs = options.timeoutMs ?? TOKEN_REQUEST_TIMEOUT_MS;
try {
const response = await postTokenForm(
new URLSearchParams({
grant_type: "refresh_token",
refresh_token: refreshToken,
client_id: CLIENT_ID,
}),
{ ...options, timeoutMs },
);
return await readOpenAITokenResponse(response, "refresh", refreshToken);
} catch (error) {
return {
type: "failed",
operation: "refresh",
...(options.signal?.aborted ? { cancelled: true } : {}),
summary: formatTokenRequestError("refresh", error, timeoutMs, options.signal),
};
}
return await requestOpenAIToken(
new URLSearchParams({
grant_type: "refresh_token",
refresh_token: refreshToken,
client_id: CLIENT_ID,
}),
"refresh",
options,
refreshToken,
);
}

View file

@ -0,0 +1,31 @@
import { fetchWithSsrFGuard } from "openclaw/plugin-sdk/ssrf-runtime";
export async function withOpenAIOAuthResponse<T>(
request: Parameters<typeof fetchWithSsrFGuard>[0],
consume: (response: Response) => Promise<T>,
): Promise<T> {
const { response, release } = await fetchWithSsrFGuard(request);
try {
// Keep the guarded transport alive through bounded reads and owner validation.
return await consume(response);
} finally {
await release();
}
}
export function createOpenAIAuthorizationCodeForm(params: {
clientId: string;
code: string;
verifier: string;
redirectUri: string;
resource?: string;
}): URLSearchParams {
return new URLSearchParams({
grant_type: "authorization_code",
client_id: params.clientId,
code: params.code,
code_verifier: params.verifier,
redirect_uri: params.redirectUri,
...(params.resource ? { resource: params.resource } : {}),
});
}

View file

@ -1,6 +1,7 @@
import { createHash } from "node:crypto";
import { request as httpRequest } from "node:http";
import { exportJWK, generateKeyPair, SignJWT } from "jose";
import { createDeferred } from "openclaw/plugin-sdk/extension-shared";
import type { ProviderAuthContext } from "openclaw/plugin-sdk/plugin-entry";
import type { OAuthCredential } from "openclaw/plugin-sdk/provider-auth";
import { afterEach, beforeAll, beforeEach, describe, expect, it, vi } from "vitest";
@ -188,6 +189,101 @@ describe("ChatGPT token-sharing authorization", () => {
},
);
it("restarts a cancelled login without accepting its stale callback", async () => {
const controller = new AbortController();
const opened = createDeferred<URL>();
const ctx = context();
ctx.signal = AbortSignal.any([ctx.signal!, controller.signal]);
ctx.openUrl = async (url) => {
opened.resolve(new URL(url));
};
const login = loginTokenSharing(ctx);
void login.catch(() => undefined);
try {
const previousAuthorization = await opened.promise;
controller.abort();
await expect(login).rejects.toThrow();
expect(request).not.toHaveBeenCalled();
const nextContext = context();
const completeNextCallback = nextContext.openUrl;
nextContext.openUrl = async (url) => {
const nextAuthorization = new URL(url);
const previousState = previousAuthorization.searchParams.get("state")!;
expect(nextAuthorization.searchParams.get("state")).not.toBe(previousState);
const staleCallback = new URL(nextAuthorization.searchParams.get("redirect_uri")!);
staleCallback.hostname = "127.0.0.1";
staleCallback.search = new URLSearchParams({
code: "cancelled-code",
state: previousState,
}).toString();
const staleResponse = await fetch(staleCallback);
expect(staleResponse.status).toBe(400);
await staleResponse.text();
expect(request).not.toHaveBeenCalled();
await completeNextCallback(url);
};
const restarted = await loginTokenSharing(nextContext);
expect(restarted.profiles).toHaveLength(1);
expect(restarted.profiles[0]?.credential).toMatchObject({
access: "opaque-test-access",
authFlow: TOKEN_SHARING_AUTH_FLOW,
});
expect((await callbackResponse!).status).toBe(200);
const exchanges = request.mock.calls.filter(([params]) => params.init?.method === "POST");
expect(exchanges).toHaveLength(1);
expect(exchanges[0]![0].init.body.get("code")).toBe("test-code");
} finally {
controller.abort();
await login.catch(() => undefined);
}
});
it("releases the callback listener on cancellation while a token request is still cleaning up", async () => {
const releaseEntered = createDeferred<void>();
const allowRelease = createDeferred<void>();
const fetchResponse = request.getMockImplementation()!;
request.mockImplementationOnce(async (params) => {
const result = await fetchResponse(params);
return {
...result,
release: async () => {
releaseEntered.resolve();
await allowRelease.promise;
await result.release();
},
};
});
const controller = new AbortController();
const ctx = context();
ctx.signal = AbortSignal.any([ctx.signal!, controller.signal]);
const login = loginTokenSharing(ctx);
const settled = vi.fn();
void login.then(settled, settled);
let originalCallback: Promise<Response> | undefined;
try {
await releaseEntered.promise;
originalCallback = callbackResponse!;
controller.abort();
expect(settled).not.toHaveBeenCalled();
const replacement = await loginTokenSharing(context());
expect(replacement.profiles).toHaveLength(1);
expect(replacement.profiles[0]?.credential).toMatchObject({
access: "opaque-test-access",
authFlow: TOKEN_SHARING_AUTH_FLOW,
});
expect((await callbackResponse!).status).toBe(200);
await expect(originalCallback).rejects.toThrow();
expect(settled).not.toHaveBeenCalled();
} finally {
controller.abort();
allowRelease.resolve();
await expect(login).rejects.toThrow();
await originalCallback?.then((response) => response.text()).catch(() => undefined);
}
});
it.each(["without-id-token", "legacy"] as const)(
"reuses registered client for %s reconnect",
async (state) => {

View file

@ -1,9 +1,9 @@
import { createHash } from "node:crypto";
import { createServer, type ServerResponse } from "node:http";
import { createLocalJWKSet, decodeJwt, jwtVerify } from "jose";
import type { ProviderAuthContext, ProviderAuthResult } from "openclaw/plugin-sdk/plugin-entry";
import type { OAuthCredential } from "openclaw/plugin-sdk/provider-auth";
import { buildOauthProviderAuthResult } from "openclaw/plugin-sdk/provider-auth-result";
import { startProviderOAuthLoopbackCallbackServer } from "openclaw/plugin-sdk/provider-auth-runtime";
import {
generateOAuthState,
generatePKCE,
@ -13,13 +13,16 @@ import {
withOAuthLoginAbort,
} from "openclaw/plugin-sdk/provider-oauth-runtime";
import { readResponseWithLimit } from "openclaw/plugin-sdk/response-limit-runtime";
import { fetchWithSsrFGuard } from "openclaw/plugin-sdk/ssrf-runtime";
import {
asOptionalRecord,
isRecord,
normalizeOptionalString,
} from "openclaw/plugin-sdk/string-coerce-runtime";
import { OPENAI_DEFAULT_MODEL } from "./default-models.js";
import {
createOpenAIAuthorizationCodeForm,
withOpenAIOAuthResponse,
} from "./openai-oauth-http.runtime.js";
import {
IDENTITY_AUTH_FLOW,
isSIWCAuthFlow,
@ -40,60 +43,60 @@ type LoginOwner = Pick<ProviderAuthContext, "signal" | "assertCurrent">;
async function requestJson(url: string, owner: LoginOwner, body?: URLSearchParams) {
owner.signal?.throwIfAborted();
owner.assertCurrent?.();
const { response, release } = await fetchWithSsrFGuard({
url,
policy: { hostnameAllowlist: ["auth.openai.com"] },
mode: "trusted_env_proxy",
requireHttps: true,
maxRedirects: 0,
capture: false,
timeoutMs: 30_000,
signal: owner.signal,
beforeRequest: owner.assertCurrent,
auditContext: "openai-token-sharing-oauth",
...(body
? {
init: {
method: "POST",
headers: { "Content-Type": "application/x-www-form-urlencoded" },
body,
return await withOpenAIOAuthResponse(
{
url,
policy: { hostnameAllowlist: ["auth.openai.com"] },
mode: "trusted_env_proxy",
requireHttps: true,
maxRedirects: 0,
capture: false,
timeoutMs: 30_000,
signal: owner.signal,
beforeRequest: owner.assertCurrent,
auditContext: "openai-token-sharing-oauth",
...(body
? {
init: {
method: "POST",
headers: { "Content-Type": "application/x-www-form-urlencoded" },
body,
},
}
: {}),
},
async (response) => {
const bytes = await readResponseWithLimit(response, MAX_RESPONSE_BYTES);
owner.signal?.throwIfAborted();
owner.assertCurrent?.();
let json: Record<string, unknown> | undefined;
try {
json = asOptionalRecord(JSON.parse(Buffer.from(bytes).toString("utf8")));
} catch {
// Provider bodies can contain credentials: report bounded status, never the body.
}
if (!response.ok) {
const invalidGrant = json?.error === "invalid_grant";
throw Object.assign(
new Error(
invalidGrant
? "ChatGPT connection expired or was revoked. Sign in again to reconnect."
: `ChatGPT authentication request failed (HTTP ${response.status}). Retry sign-in later.`,
),
{
oauthRefreshFailure: {
status: response.status,
...(invalidGrant ? { reason: "invalid_grant", errorType: "invalid_grant" } : {}),
},
},
}
: {}),
});
try {
const bytes = await readResponseWithLimit(response, MAX_RESPONSE_BYTES);
owner.signal?.throwIfAborted();
owner.assertCurrent?.();
let json: Record<string, unknown> | undefined;
try {
json = asOptionalRecord(JSON.parse(Buffer.from(bytes).toString("utf8")));
} catch {
// Provider bodies can contain credentials: report bounded status, never the body.
}
if (!response.ok) {
const invalidGrant = json?.error === "invalid_grant";
throw Object.assign(
new Error(
invalidGrant
? "ChatGPT connection expired or was revoked. Sign in again to reconnect."
: `ChatGPT authentication request failed (HTTP ${response.status}). Retry sign-in later.`,
),
{
oauthRefreshFailure: {
status: response.status,
...(invalidGrant ? { reason: "invalid_grant", errorType: "invalid_grant" } : {}),
},
},
);
}
if (!json) {
throw new Error("ChatGPT authentication returned an invalid response.");
}
return json;
} finally {
await release();
}
);
}
if (!json) {
throw new Error("ChatGPT authentication returned an invalid response.");
}
return json;
},
);
}
async function verifyIdentity(
@ -298,87 +301,15 @@ export async function loginTokenSharing(ctx: ProviderAuthContext): Promise<Provi
nonce,
...(registering ? { agent_name_hint: "OpenClaw" } : {}),
}).toString();
let resolveCode!: (authorization: { code: string; clientId: string }) => void;
let rejectCode!: (error: Error) => void;
let callbackConsumed = false;
let browserResponse: ServerResponse | undefined;
const callback = new Promise<{ code: string; clientId: string }>((resolve, reject) => {
resolveCode = resolve;
rejectCode = reject;
});
// Register a rejection handler before browser I/O, which can outlive the callback.
void callback.catch(() => undefined);
const server = createServer((request, response) => {
response.setHeader("Content-Type", "text/html; charset=utf-8");
response.setHeader("Connection", "close");
response.setHeader("Cache-Control", "no-store");
response.setHeader("Referrer-Policy", "no-referrer");
let callbackUrl: URL;
try {
callbackUrl = new URL(request.url ?? "/", TOKEN_SHARING_REDIRECT_URI);
} catch {
response.writeHead(400).end(oauthErrorHtml("Invalid sign-in callback."));
return;
}
if (
request.method !== "GET" ||
callbackUrl.pathname !== "/auth/callback" ||
callbackUrl.searchParams.getAll("state").length !== 1 ||
callbackUrl.searchParams.get("state") !== state ||
callbackConsumed
) {
response
.writeHead(400)
.end(oauthErrorHtml("Invalid or expired sign-in callback. Return to OpenClaw to retry."));
return;
}
callbackConsumed = true;
try {
owner.assertCurrent?.();
owner.signal.throwIfAborted();
if (callbackUrl.searchParams.has("error")) {
throw new Error(
callbackUrl.searchParams.get("error") === "access_denied"
? "ChatGPT authorization was declined. Start sign-in again when ready."
: "ChatGPT authorization failed. Start sign-in again.",
);
}
const code = callbackUrl.searchParams.get("code");
if (!code || callbackUrl.searchParams.getAll("code").length !== 1) {
throw new Error(
"ChatGPT callback did not contain an authorization code. Start sign-in again.",
);
}
const returnedIds = callbackUrl.searchParams.getAll("client_id");
const returnedId = returnedIds[0];
// Registration changes the client ID mid-flow. Ordinary reauthorization
// may omit it, but must never replace the selected registration.
if (
returnedIds.length > 1 ||
(registering
? !returnedId || !/^oaiapp_[A-Za-z0-9_-]+$/u.test(returnedId)
: returnedId !== undefined && returnedId !== clientId)
) {
throw new Error("ChatGPT returned an invalid OAuth client ID. Start sign-in again.");
}
browserResponse = response;
resolveCode({ code, clientId: returnedId ?? clientId });
} catch (error) {
response
.writeHead(400)
.end(oauthErrorHtml("Authorization did not complete. Return to OpenClaw to retry."));
rejectCode(error instanceof Error ? error : new Error("ChatGPT authorization failed."));
}
const callback = await startProviderOAuthLoopbackCallbackServer({
redirectUrl: TOKEN_SHARING_REDIRECT_URI,
expectedState: state,
signal: owner.signal,
// SSH forwards target IPv4 loopback; keep the registered localhost redirect unchanged.
bindOnlyHostname: "127.0.0.1",
deferResponse: true,
});
try {
await withOAuthLoginAbort(
new Promise<void>((resolve, reject) => {
server.once("error", reject);
// SSH forwards target IPv4 loopback; keep the registered localhost redirect unchanged.
server.listen(8080, "127.0.0.1", resolve);
}),
owner.signal,
);
owner.assertCurrent?.();
// Gateway wizards attach the browser URL to the next note they publish.
await withOAuthLoginAbort(ctx.openUrl(url.toString()), owner.signal);
@ -407,22 +338,42 @@ export async function loginTokenSharing(ctx: ProviderAuthContext): Promise<Provi
);
owner.assertCurrent?.();
owner.signal.throwIfAborted();
const authorization = await withOAuthLoginAbort(callback, owner.signal);
const authorization = await withOAuthLoginAbort(callback.waitForCallback(), owner.signal);
owner.assertCurrent?.();
owner.signal.throwIfAborted();
if (authorization.type === "oauth_error") {
throw new Error(
authorization.error === "access_denied"
? "ChatGPT authorization was declined. Start sign-in again when ready."
: "ChatGPT authorization failed. Start sign-in again.",
);
}
const returnedIds = authorization.parameters.getAll("client_id");
const returnedId = returnedIds[0];
// Registration changes the client ID mid-flow. Reauthorization cannot replace it.
if (
returnedIds.length > 1 ||
(registering
? !returnedId || !/^oaiapp_[A-Za-z0-9_-]+$/u.test(returnedId)
: returnedId !== undefined && returnedId !== clientId)
) {
throw new Error("ChatGPT returned an invalid OAuth client ID. Start sign-in again.");
}
const authorizedClientId = returnedId ?? clientId;
const json = await requestJson(
TOKEN_ENDPOINT,
owner,
new URLSearchParams({
grant_type: "authorization_code",
client_id: authorization.clientId,
createOpenAIAuthorizationCodeForm({
clientId: authorizedClientId,
code: authorization.code,
code_verifier: verifier,
redirect_uri: TOKEN_SHARING_REDIRECT_URI,
verifier,
redirectUri: TOKEN_SHARING_REDIRECT_URI,
resource: TOKEN_SHARING_RESOURCE,
}),
);
const { credential, subject } = await readCredential({
json,
clientId: authorization.clientId,
clientId: authorizedClientId,
nonce,
owner,
});
@ -443,15 +394,17 @@ export async function loginTokenSharing(ctx: ProviderAuthContext): Promise<Provi
}
const profileName = credential.accountId.slice(0, 24);
const sharing = credential.authFlow === TOKEN_SHARING_AUTH_FLOW;
browserResponse
?.writeHead(200)
.end(
oauthSuccessHtml(
sharing
? "ChatGPT token sharing is connected. You can return to OpenClaw."
: "ChatGPT sign-in succeeded. Token sharing is disabled; return to OpenClaw to choose inference access.",
),
);
await callback.complete({
status: 200,
contentType: "text/html; charset=utf-8",
body: oauthSuccessHtml(
sharing
? "ChatGPT token sharing is connected. You can return to OpenClaw."
: "ChatGPT sign-in succeeded. Token sharing is disabled; return to OpenClaw to choose inference access.",
),
});
owner.signal.throwIfAborted();
owner.assertCurrent?.();
const result = buildOauthProviderAuthResult({
providerId: "openai",
profilePrefix: "openai:token-sharing",
@ -479,16 +432,13 @@ export async function loginTokenSharing(ctx: ProviderAuthContext): Promise<Provi
}
return result;
} catch (error) {
browserResponse
?.writeHead(400)
.end(oauthErrorHtml("Sign-in did not complete. Return to OpenClaw for details and retry."));
await callback.complete({
status: 400,
contentType: "text/html; charset=utf-8",
body: oauthErrorHtml("Sign-in did not complete. Return to OpenClaw for details and retry."),
});
throw error;
} finally {
server.close();
if (browserResponse && !browserResponse.writableFinished) {
browserResponse.once("finish", () => server.closeAllConnections());
} else {
server.closeAllConnections();
}
await callback.close();
}
}

View file

@ -473,9 +473,14 @@ describe("OpenRouter OAuth", () => {
type: "authorization_code" as const,
code: "AUTHCODE",
state: "state-1",
parameters: new URLSearchParams({ code: "AUTHCODE", state: "state-1" }),
}));
const close = vi.fn(async () => undefined);
const startCallback = vi.fn(async () => ({ waitForCallback, close }));
const startCallback = vi.fn(async () => ({
waitForCallback,
complete: async () => undefined,
close,
}));
const { ctx, openUrl, text } = createOpenRouterOAuthContext({ isRemote: false });
await loginOpenRouterOAuth(ctx, {
@ -514,7 +519,11 @@ describe("OpenRouter OAuth", () => {
errorDescription: "Denied",
}));
const close = vi.fn(async () => undefined);
const startCallback = vi.fn(async () => ({ waitForCallback, close }));
const startCallback = vi.fn(async () => ({
waitForCallback,
complete: async () => undefined,
close,
}));
const { ctx, text } = createOpenRouterOAuthContext({ isRemote: false });
await expect(

View file

@ -0,0 +1,480 @@
import { expectDefined } from "@openclaw/normalization-core";
import { Compile } from "typebox/compile";
import { afterEach, describe, expect, it, vi } from "vitest";
import type {
SystemAgentSetupAuthStartParams,
WizardNextParams,
WizardNextResult,
} from "../../../packages/gateway-protocol/src/index.js";
import { WizardNextResultSchema } from "../../../packages/gateway-protocol/src/schema/wizard.js";
import { resetCommandQueueStateForTest } from "../../process/command-queue.test-support.js";
import { createDeferredCore } from "../../shared/deferred.js";
import { WizardSession } from "../../wizard/session.js";
import { whenAdmittedWizardSessionSettled } from "./setup-admission.js";
import { systemAgentHandlers } from "./system-agent.js";
import type {
GatewayClient,
GatewayRequestContext,
GatewayRequestHandlerOptions,
} from "./types.js";
import { wizardHandlers } from "./wizard.js";
const setupInferenceMocks = vi.hoisted(() => ({ activateSetupInference: vi.fn() }));
vi.mock("../../system-agent/setup-inference.js", () => ({
activateSetupInference: setupInferenceMocks.activateSetupInference,
}));
const validateWizardResult = Compile(WizardNextResultSchema);
function makeContext() {
const wizardSessions = new Map<string, WizardSession>();
return {
wizardSessions,
context: {
wizardSessions,
findRunningWizard: () => undefined,
purgeWizardSession: (id: string) => wizardSessions.delete(id),
} as unknown as GatewayRequestContext,
};
}
function makeRespond() {
const calls: Array<{ ok: boolean; payload?: unknown; error?: unknown }> = [];
return {
calls,
respond: (ok: boolean, payload?: unknown, error?: unknown) => {
calls.push({ ok, payload, error });
},
};
}
function systemAgentHandler(method: keyof typeof systemAgentHandlers) {
return expectDefined(systemAgentHandlers[method], `systemAgentHandlers["${method}"] invariant`);
}
const authClient = {
connId: "auth-connection",
connect: { device: { id: "auth-device" } },
authenticatedUserProfile: { profileId: "auth-owner" },
} as GatewayClient;
const authParams = {
authChoice: "github-copilot",
agentId: "research",
workspace: "/tmp/auth-workspace",
};
function startAuthRequest(
context: GatewayRequestContext,
sessionId: string,
overrides: Partial<SystemAgentSetupAuthStartParams> = {},
client: GatewayClient = authClient,
authority: Pick<GatewayRequestHandlerOptions, "sessionMutationCommitGuard"> = {},
) {
const { calls, respond } = makeRespond();
const pending = Promise.resolve(
systemAgentHandler("openclaw.setup.auth.start")({
params: { ...authParams, ...overrides, sessionId },
client,
context,
respond,
...authority,
} as never),
);
return { calls, pending };
}
async function settleAuthRequests(
wizardSessions: Map<string, WizardSession>,
pending: Array<Promise<unknown>>,
release: () => void,
) {
for (const session of wizardSessions.values()) {
session.cancel();
}
release();
await Promise.all(pending);
for (const session of wizardSessions.values()) {
session.cancel();
await whenAdmittedWizardSessionSettled(session);
}
}
async function callWizardNext(
context: GatewayRequestContext,
params: WizardNextParams,
): Promise<WizardNextResult> {
const { calls, respond } = makeRespond();
await expectDefined(
wizardHandlers["wizard.next"],
"wizard.next handler",
)({
params,
respond,
context,
} as never);
expect(calls).toHaveLength(1);
expect(calls[0]?.ok).toBe(true);
const payload = calls[0]?.payload;
if (!validateWizardResult.Check(payload)) {
throw new Error("wizard.next returned an invalid result");
}
return payload;
}
describe("openclaw.setup auth retries", () => {
afterEach(() => {
vi.resetAllMocks();
resetCommandQueueStateForTest();
});
it.each(["running", "cancelled"] as const)(
"replaces the owner's %s sign-in after provider cleanup settles",
async (status) => {
const { wizardSessions, context } = makeContext();
const cleanupStarted = createDeferredCore();
const cleanupReleased = createDeferredCore();
setupInferenceMocks.activateSetupInference
.mockImplementationOnce(async (params) => {
try {
await params.prompter.note("Complete browser sign-in");
} finally {
cleanupStarted.resolve();
await cleanupReleased.promise;
}
})
.mockImplementationOnce(async (params) => {
await params.prompter.note("Complete the replacement sign-in");
return { ok: true, modelRef: "github-copilot/test", latencyMs: 1, lines: [] };
});
const first = startAuthRequest(context, "auth-first");
const requests = [first.pending];
try {
await first.pending;
const session = expectDefined(wizardSessions.get("auth-first"), "first auth session");
await callWizardNext(context, { sessionId: "auth-first" });
if (status === "cancelled") {
await expectDefined(
wizardHandlers["wizard.cancel"],
"wizard.cancel",
)({
params: { sessionId: "auth-first" },
context,
respond: () => undefined,
} as never);
}
const replacement = startAuthRequest(context, "auth-replacement");
requests.push(replacement.pending);
await Promise.race([cleanupStarted.promise, replacement.pending]);
expect(session.signal.aborted).toBe(true);
expect(replacement.calls).toEqual([]);
expect(setupInferenceMocks.activateSetupInference).toHaveBeenCalledOnce();
cleanupReleased.resolve();
await replacement.pending;
expect(wizardSessions.has("auth-first")).toBe(false);
expect(replacement.calls).toEqual([
{
ok: true,
payload: { sessionId: "auth-replacement", done: false, status: "running" },
error: undefined,
},
]);
const step = await callWizardNext(context, { sessionId: "auth-replacement" });
expect(step.step?.message).toBe("Complete the replacement sign-in");
expect(setupInferenceMocks.activateSetupInference).toHaveBeenCalledTimes(2);
} finally {
await settleAuthRequests(wizardSessions, requests, () => cleanupReleased.resolve());
}
},
);
it("runs only the latest sign-in when retries overlap provider cleanup", async () => {
const { wizardSessions, context } = makeContext();
const cleanupStarted = createDeferredCore();
const cleanupReleased = createDeferredCore();
setupInferenceMocks.activateSetupInference
.mockImplementationOnce(async (params) => {
try {
await params.prompter.note("Complete browser sign-in");
} finally {
cleanupStarted.resolve();
await cleanupReleased.promise;
}
})
.mockImplementation(async (params) => {
await params.prompter.note("Latest sign-in");
return { ok: true, modelRef: "github-copilot/test", latencyMs: 1, lines: [] };
});
const first = startAuthRequest(context, "auth-first");
const requests = [first.pending];
try {
await first.pending;
const session = expectDefined(wizardSessions.get("auth-first"), "first auth session");
await callWizardNext(context, { sessionId: "auth-first" });
const second = startAuthRequest(context, "auth-second");
requests.push(second.pending);
await Promise.race([cleanupStarted.promise, second.pending]);
expect(session.signal.aborted).toBe(true);
const duplicate = startAuthRequest(context, "auth-second");
requests.push(duplicate.pending);
await duplicate.pending;
expect(duplicate.calls[0]).toMatchObject({
ok: false,
error: { message: "wizard session already exists" },
});
expect(second.calls).toEqual([]);
const third = startAuthRequest(context, "auth-third");
requests.push(third.pending);
cleanupReleased.resolve();
await Promise.all([second.pending, third.pending]);
expect(second.calls).toEqual([
{
ok: true,
payload: { sessionId: "auth-second", done: true, status: "cancelled" },
error: undefined,
},
]);
expect(third.calls[0]).toMatchObject({
ok: true,
payload: { sessionId: "auth-third", done: false, status: "running" },
});
expect(wizardSessions.has("auth-second")).toBe(false);
expect((await callWizardNext(context, { sessionId: "auth-third" })).step?.message).toBe(
"Latest sign-in",
);
expect(setupInferenceMocks.activateSetupInference).toHaveBeenCalledTimes(2);
} finally {
await settleAuthRequests(wizardSessions, requests, () => cleanupReleased.resolve());
}
});
it.each([
{ when: "before cancellation", revocation: "client invalidation" },
{ when: "before cancellation", revocation: "request guard" },
{ when: "during cleanup", revocation: "client invalidation" },
{ when: "during cleanup", revocation: "request guard" },
] as const)(
"rejects a queued retry on $revocation $when and preserves a later live retry",
async ({ when, revocation }) => {
const { wizardSessions, context } = makeContext();
const cleanupStarted = createDeferredCore();
const cleanupReleased = createDeferredCore();
const authorityError = new Error("Queued request authority was revoked");
const retryClient = { ...authClient, connId: "queued-connection", invalidated: false };
let guardRevoked = false;
setupInferenceMocks.activateSetupInference
.mockImplementationOnce(async (params) => {
try {
await params.prompter.note("Complete the original sign-in");
} finally {
cleanupStarted.resolve();
if (when === "during cleanup") {
await cleanupReleased.promise;
}
}
})
.mockImplementation(async (params) => {
await params.prompter.note("Complete the live replacement sign-in");
return { ok: true, modelRef: "github-copilot/test", latencyMs: 1, lines: [] };
});
const first = startAuthRequest(context, "auth-first");
const requests: Array<Promise<unknown>> = [first.pending];
try {
await first.pending;
const original = expectDefined(wizardSessions.get("auth-first"), "original sign-in");
await callWizardNext(context, { sessionId: "auth-first" });
const denied = startAuthRequest(context, "auth-denied", {}, retryClient, {
sessionMutationCommitGuard: () => {
if (guardRevoked) {
throw authorityError;
}
},
});
const deniedResult = denied.pending.then(
() => undefined,
(error: unknown) => error,
);
requests.push(deniedResult);
if (when === "during cleanup") {
await Promise.race([cleanupStarted.promise, deniedResult]);
expect(original.signal.aborted).toBe(true);
}
if (revocation === "client invalidation") {
retryClient.invalidated = true;
} else {
guardRevoked = true;
}
cleanupReleased.resolve();
const error = await deniedResult;
if (when === "before cancellation") {
expect(original.signal.aborted).toBe(false);
expect(wizardSessions.get("auth-first")).toBe(original);
}
expect(setupInferenceMocks.activateSetupInference).toHaveBeenCalledOnce();
expect(wizardSessions.has("auth-denied")).toBe(false);
expect(denied.calls).toEqual([]);
if (revocation === "request guard") {
expect(error).toBe(authorityError);
} else {
expect(error).toMatchObject({ message: "Gateway requester authority changed" });
}
const live = startAuthRequest(context, "auth-live");
requests.push(live.pending);
await live.pending;
expect(live.calls[0]).toMatchObject({
ok: true,
payload: { sessionId: "auth-live", done: false, status: "running" },
});
expect(original.signal.aborted).toBe(true);
expect(wizardSessions.has("auth-first")).toBe(false);
expect((await callWizardNext(context, { sessionId: "auth-live" })).step?.message).toBe(
"Complete the live replacement sign-in",
);
expect(setupInferenceMocks.activateSetupInference).toHaveBeenCalledTimes(2);
} finally {
await settleAuthRequests(wizardSessions, requests, () => cleanupReleased.resolve());
}
},
);
it("retains sign-in ownership after overlapping retries during preparation", async () => {
const { wizardSessions, context } = makeContext();
const preparationStarted = createDeferredCore();
const releasePreparation = createDeferredCore();
setupInferenceMocks.activateSetupInference
.mockImplementationOnce(async (params) => {
await params.beforePersistentEffect();
preparationStarted.resolve();
await releasePreparation.promise;
params.onPreparationComplete();
await params.prompter.note("Complete the original sign-in");
return { ok: true, modelRef: "github-copilot/test", latencyMs: 1, lines: [] };
})
.mockImplementationOnce(async (params) => {
await params.prompter.note("Complete the replacement sign-in");
return { ok: true, modelRef: "github-copilot/test", latencyMs: 1, lines: [] };
});
const first = startAuthRequest(context, "auth-first");
const requests = [first.pending];
try {
await first.pending;
await preparationStarted.promise;
const session = expectDefined(wizardSessions.get("auth-first"), "preparing auth session");
const second = startAuthRequest(context, "auth-second");
const third = startAuthRequest(context, "auth-third");
requests.push(second.pending, third.pending);
await Promise.all([second.pending, third.pending]);
for (const retry of [second, third]) {
expect(retry.calls).toEqual([
{
ok: false,
payload: undefined,
error: expect.objectContaining({ details: { code: "SETUP_ADMISSION_BUSY" } }),
},
]);
}
expect(session.signal.aborted).toBe(false);
expect(setupInferenceMocks.activateSetupInference).toHaveBeenCalledOnce();
releasePreparation.resolve();
expect((await callWizardNext(context, { sessionId: "auth-first" })).step?.message).toBe(
"Complete the original sign-in",
);
const replacement = startAuthRequest(context, "auth-replacement");
requests.push(replacement.pending);
await replacement.pending;
expect(replacement.calls).toEqual([
{
ok: true,
payload: { sessionId: "auth-replacement", done: false, status: "running" },
error: undefined,
},
]);
expect(session.signal.aborted).toBe(true);
expect(wizardSessions.has("auth-first")).toBe(false);
expect((await callWizardNext(context, { sessionId: "auth-replacement" })).step?.message).toBe(
"Complete the replacement sign-in",
);
expect(setupInferenceMocks.activateSetupInference).toHaveBeenCalledTimes(2);
} finally {
await settleAuthRequests(wizardSessions, requests, () => releasePreparation.resolve());
}
});
it.each([
"owner",
"choice",
"agent",
"workspace",
"modelTarget",
"nativeSessionCatalogsEnabled",
"commit",
] as const)("keeps another sign-in busy when its %s prevents replacement", async (difference) => {
const { wizardSessions, context } = makeContext();
const released = createDeferredCore();
setupInferenceMocks.activateSetupInference.mockImplementationOnce(async (params) => {
if (difference === "commit") {
await params.onCommitStarted();
}
await params.prompter.note("Sign-in still owns setup");
await released.promise;
return { ok: true, modelRef: "github-copilot/test", latencyMs: 1, lines: [] };
});
const first = startAuthRequest(context, "auth-first");
const requests = [first.pending];
try {
await first.pending;
const session = expectDefined(wizardSessions.get("auth-first"), "first auth session");
const note = await callWizardNext(context, { sessionId: "auth-first" });
const replacement = startAuthRequest(
context,
"auth-replacement",
difference === "choice"
? { authChoice: "xai" }
: difference === "agent"
? { agentId: "other-agent" }
: difference === "workspace"
? { workspace: "/tmp/other-auth-workspace" }
: difference === "modelTarget"
? { modelTarget: "utility" }
: difference === "nativeSessionCatalogsEnabled"
? { nativeSessionCatalogsEnabled: true }
: {},
difference === "owner"
? {
...authClient,
authenticatedUserProfile: {
...authClient.authenticatedUserProfile!,
profileId: "other-owner",
},
}
: authClient,
);
requests.push(replacement.pending);
await replacement.pending;
expect(replacement.calls).toEqual([
{
ok: false,
payload: undefined,
error: expect.objectContaining({
message: "OpenClaw setup is already in progress; try again when it finishes.",
}),
},
]);
expect(session.signal.aborted).toBe(false);
expect(setupInferenceMocks.activateSetupInference).toHaveBeenCalledOnce();
if (difference === "commit") {
await session.answer(expectDefined(note.step, "locked sign-in step").id, null);
}
} finally {
released.resolve();
const session = wizardSessions.get("auth-first");
const step = session?.getCurrentStep();
if (step) {
await session?.answer(step.id, null);
}
await settleAuthRequests(wizardSessions, requests, () => released.resolve());
}
});
});

View file

@ -1,16 +1,149 @@
import { ErrorCodes, errorShape } from "../../../packages/gateway-protocol/src/index.js";
import { defaultRuntime } from "../../runtime.js";
import { WizardSession } from "../../wizard/session.js";
import { createAdmittedWizardSession, respondSetupAdmissionBusy } from "./setup-admission.js";
import {
createAdmittedWizardSession,
respondSetupAdmissionBusy,
whenAdmittedWizardSessionSettled,
} from "./setup-admission.js";
import { activateGatewaySetupInference } from "./system-agent-execution.js";
import type { GatewayRequestContext, RespondFn } from "./types.js";
type SetupActivation = Pick<
Parameters<typeof activateGatewaySetupInference>[0],
| "kind"
| "agentId"
| "modelRef"
| "modelTarget"
| "authChoice"
| "apiKey"
| "workspace"
| "nativeSessionCatalogsEnabled"
>;
type AuthWizardRequest = {
ownerKey: string;
activation: SetupActivation;
pendingSessionIds: Set<string>;
session: Promise<{ sessionId: string; session: WizardSession } | "superseded" | undefined>;
};
const authWizardRequests = new WeakMap<
GatewayRequestContext["wizardSessions"],
AuthWizardRequest
>();
async function createSetupActivationSession(
params: {
sessionId: string;
ownerKey?: string;
assertCurrent?: () => void;
activation: SetupActivation;
context: GatewayRequestContext;
},
createSession: () => WizardSession,
): Promise<WizardSession | "superseded" | undefined> {
const { ownerKey, activation } = params;
if (!ownerKey || activation.kind !== "provider-auth") {
return createAdmittedWizardSession(createSession);
}
const sessions = params.context.wizardSessions;
const previous = authWizardRequests.get(sessions);
if (
previous &&
(previous.ownerKey !== ownerKey ||
previous.activation.authChoice !== activation.authChoice ||
previous.activation.agentId !== activation.agentId ||
previous.activation.workspace !== activation.workspace ||
previous.activation.modelTarget !== activation.modelTarget ||
previous.activation.nativeSessionCatalogsEnabled !== activation.nativeSessionCatalogsEnabled)
) {
return undefined;
}
let authorityFailure: { error: unknown } | undefined;
const assertCurrent = () => {
try {
params.assertCurrent?.();
} catch (error) {
authorityFailure = { error };
throw error;
}
};
const request: AuthWizardRequest = {
ownerKey,
activation,
pendingSessionIds: previous?.pendingSessionIds ?? new Set(),
session: Promise.resolve().then(async () => {
const predecessor = previous ? await previous.session : undefined;
try {
assertCurrent();
if (predecessor && predecessor !== "superseded") {
if (predecessor.session.getStatus() === "running" && !predecessor.session.cancel()) {
// Preparation locks can lift later; rejected retries must retain that owner.
return predecessor;
}
// Cancellation retires prompts before provider sockets and the setup lock.
// Every queued replacement inherits this barrier, even if superseded.
await whenAdmittedWizardSessionSettled(predecessor.session);
if (sessions.get(predecessor.sessionId) === predecessor.session) {
params.context.purgeWizardSession(predecessor.sessionId);
}
}
if (authWizardRequests.get(sessions) !== request) {
return "superseded";
}
const session = await createAdmittedWizardSession(() => {
assertCurrent();
return createSession();
});
return session ? { sessionId: params.sessionId, session } : undefined;
} catch (error) {
if (!authorityFailure) {
throw error;
}
// Denied callers receive their error, but later retries still inherit the live owner.
return predecessor;
}
}),
};
// Reserve before awaiting admission so only the newest request can start login.
authWizardRequests.set(sessions, request);
request.pendingSessionIds.add(params.sessionId);
const release = () => {
if (authWizardRequests.get(sessions) === request) {
authWizardRequests.delete(sessions);
}
};
try {
const session = await request.session.catch((error: unknown) => {
release();
throw error;
});
let result: WizardSession | "superseded" | undefined;
if (session && session !== "superseded") {
void whenAdmittedWizardSessionSettled(session.session).then(release, release);
result = session.sessionId === params.sessionId ? session.session : undefined;
} else {
release();
result = session;
}
if (authorityFailure) {
throw authorityFailure.error;
}
return result;
} finally {
request.pendingSessionIds.delete(params.sessionId);
}
}
export function rejectExistingSetupWizardSession(params: {
sessionId: string;
context: GatewayRequestContext;
respond: RespondFn;
}): boolean {
if (!params.context.wizardSessions.has(params.sessionId)) {
const sessions = params.context.wizardSessions;
if (
!sessions.has(params.sessionId) &&
!authWizardRequests.get(sessions)?.pendingSessionIds.has(params.sessionId)
) {
return false;
}
params.respond(
@ -23,17 +156,9 @@ export function rejectExistingSetupWizardSession(params: {
export async function startSetupActivationWizard(params: {
sessionId: string;
activation: Pick<
Parameters<typeof activateGatewaySetupInference>[0],
| "kind"
| "agentId"
| "modelRef"
| "modelTarget"
| "authChoice"
| "apiKey"
| "workspace"
| "nativeSessionCatalogsEnabled"
>;
ownerKey?: string;
assertCurrent?: () => void;
activation: SetupActivation;
isLocalClient?: boolean;
timeoutMs: number;
context: GatewayRequestContext;
@ -42,7 +167,8 @@ export async function startSetupActivationWizard(params: {
if (rejectExistingSetupWizardSession(params)) {
return;
}
const session = await createAdmittedWizardSession(
const session = await createSetupActivationSession(
params,
() =>
new WizardSession(
async (prompter, signal, runnerSession) => {
@ -86,6 +212,14 @@ export async function startSetupActivationWizard(params: {
respondSetupAdmissionBusy(params.respond);
return;
}
if (session === "superseded") {
params.respond(
true,
{ sessionId: params.sessionId, done: true, status: "cancelled" },
undefined,
);
return;
}
params.context.wizardSessions.set(params.sessionId, session);
// Return ownership before any prompt so cancellation survives a lost start reply.
params.respond(true, { sessionId: params.sessionId, done: false, status: "running" }, undefined);

View file

@ -22,6 +22,7 @@ import * as setupAdmission from "./setup-admission.js";
import type { SystemAgentChatSession } from "./system-agent.js";
import {
callChat,
defaultClient,
inferenceFallbackMocks,
makeContext,
makeRespond,
@ -116,7 +117,7 @@ describe("openclaw.setup", () => {
sessionId,
authChoice: "custom-api-key",
},
client: { internal: { isLocalClient } },
client: { ...defaultClient, internal: { isLocalClient } },
context,
respond,
} as never);

View file

@ -37,6 +37,7 @@ import {
authenticatedProfileUnavailableError,
isGatewayClientProfilePending,
} from "./gateway-client-identity.js";
import { readGatewayRequestMutationAuthority } from "./session-mutation-guards.js";
import {
createAdmittedWizardSession,
runExclusiveSystemAgentSetupActivation,
@ -191,7 +192,8 @@ export const systemAgentHandlers: GatewayRequestHandlers = {
});
},
/** Start one provider-owned OAuth/device-code login over the shared wizard transport. */
"openclaw.setup.auth.start": async ({ params, respond, context, client }) => {
"openclaw.setup.auth.start": async (options) => {
const { params, respond, context, client } = options;
if (
!assertValidParams(
params,
@ -205,6 +207,8 @@ export const systemAgentHandlers: GatewayRequestHandlers = {
const { sessionId, ...activation } = params;
await startSetupActivationWizard({
sessionId,
ownerKey: resolveSystemAgentSessionOwnerKey({ client }),
assertCurrent: readGatewayRequestMutationAuthority(options).assertCurrent,
activation: { ...activation, kind: "provider-auth" },
timeoutMs: PROVIDER_AUTH_SESSION_TIMEOUT_MS,
context,

View file

@ -18,6 +18,7 @@ const portClaims: TestPortClaim[] = [];
afterEach(async () => {
await Promise.all(openCallbacks.splice(0).map((callback) => callback.close()));
await Promise.all(portClaims.splice(0).map((claim) => claim.release()));
vi.useRealTimers();
vi.restoreAllMocks();
});
@ -55,7 +56,7 @@ async function getClaimedIpv6Port(): Promise<number | undefined> {
async function start(
hostname = "127.0.0.1",
renderSuccess?: Parameters<typeof startOAuthLoopbackCallbackServer>[0]["renderSuccess"],
options: Partial<Parameters<typeof startOAuthLoopbackCallbackServer>[0]> = {},
) {
const port = hostname === "::1" ? await getClaimedIpv6Port() : await getClaimedPort();
if (!port) {
@ -65,7 +66,7 @@ async function start(
redirectUrl: callbackUrl(hostname, port),
expectedState: "state-1234567890",
timeoutMs: 5_000,
renderSuccess,
...options,
});
openCallbacks.push(callback);
return { callback, port };
@ -75,17 +76,17 @@ describe("OAuth loopback callback server", () => {
it.each(["default", "provider"] as const)(
"serves a styled %s response permitted by CSP before closing",
async (renderer) => {
const started = await start(
"127.0.0.1",
renderer === "provider"
? () => ({
body: oauthSuccessHtml(
"Authorization received; return to the terminal while OpenClaw finishes.",
),
contentType: "text/html; charset=utf-8",
})
: undefined,
);
const started = await start("127.0.0.1", {
renderSuccess:
renderer === "provider"
? () => ({
body: oauthSuccessHtml(
"Authorization received; return to the terminal while OpenClaw finishes.",
),
contentType: "text/html; charset=utf-8",
})
: undefined,
});
if (!started) {
throw new Error("IPv4 loopback unavailable");
}
@ -101,6 +102,7 @@ describe("OAuth loopback callback server", () => {
type: "authorization_code",
code: "authorization-code",
state: "state-1234567890",
parameters: new URLSearchParams("code=authorization-code&state=state-1234567890"),
});
const response = await responsePromise;
expect(response.status).toBe(200);
@ -133,7 +135,7 @@ describe("OAuth loopback callback server", () => {
},
);
it("keeps waiting after wrong path, method, missing state, and wrong state", async () => {
it("keeps waiting after wrong path, method, missing state, and ambiguous or wrong callback fields", async () => {
const started = await start();
if (!started) {
throw new Error("IPv4 loopback unavailable");
@ -143,6 +145,10 @@ describe("OAuth loopback callback server", () => {
expect((await fetch(base, { method: "POST" })).status).toBe(405);
expect((await fetch(`${base}?code=code`)).status).toBe(400);
expect((await fetch(`${base}?code=code&state=wrong`)).status).toBe(400);
expect(
(await fetch(`${base}?code=code&state=state-1234567890&state=state-1234567890`)).status,
).toBe(400);
expect((await fetch(`${base}?code=first&code=second&state=state-1234567890`)).status).toBe(400);
const response = await fetch(`${base}?code=right&state=state-1234567890`);
expect(response.status).toBe(200);
@ -223,6 +229,109 @@ describe("OAuth loopback callback server", () => {
await expect(aborted.waitForCallback()).rejects.toThrow("cancelled");
});
it("retains provider callback parameters and waits for the verified browser outcome", async () => {
const started = await start("127.0.0.1", { deferResponse: true });
if (!started) {
throw new Error("IPv4 loopback unavailable");
}
const responsePromise = fetch(
callbackUrl(
"127.0.0.1",
started.port,
"?code=code&state=state-1234567890&client_id=first&client_id=second",
),
);
const received = vi.fn();
void responsePromise.then(received, received);
const result = await started.callback.waitForCallback();
expect(result.type).toBe("authorization_code");
if (result.type !== "authorization_code") {
throw new Error("Expected authorization code");
}
expect(result.parameters.getAll("client_id")).toEqual(["first", "second"]);
expect(
(await fetch(callbackUrl("127.0.0.1", started.port, "?code=other&state=state-1234567890")))
.status,
).toBe(409);
expect(received).not.toHaveBeenCalled();
await started.callback.complete({
status: 400,
body: "Provider rejected the registration",
contentType: "text/plain",
});
const response = await responsePromise;
expect(response.status).toBe(400);
expect(await response.text()).toBe("Provider rejected the registration");
});
it("closes an admitted callback when the browser disconnects before provider completion", async () => {
const started = await start("127.0.0.1", { deferResponse: true });
if (!started) {
throw new Error("IPv4 loopback unavailable");
}
const browser = new AbortController();
const response = fetch(
callbackUrl("127.0.0.1", started.port, "?code=code&state=state-1234567890"),
{ signal: browser.signal },
);
void response.catch(() => undefined);
await started.callback.waitForCallback();
browser.abort();
await expect(response).rejects.toThrow();
await vi.waitFor(async () => {
await expect(fetch(callbackUrl("127.0.0.1", started.port))).rejects.toThrow();
});
await started.callback.complete({ status: 200, body: "Too late", contentType: "text/plain" });
const replacement = await startOAuthLoopbackCallbackServer({
redirectUrl: callbackUrl("127.0.0.1", started.port),
expectedState: "replacement-state",
});
openCallbacks.push(replacement);
});
it.each(["abort", "timeout"] as const)(
"releases an admitted callback on %s before provider work completes",
async (terminal) => {
const controller = new AbortController();
if (terminal === "timeout") {
vi.useFakeTimers({ toFake: ["setTimeout", "clearTimeout"] });
}
const started = await start("127.0.0.1", {
deferResponse: true,
signal: controller.signal,
});
if (!started) {
throw new Error("IPv4 loopback unavailable");
}
const responsePromise = fetch(
callbackUrl("127.0.0.1", started.port, "?code=code&state=state-1234567890"),
);
void responsePromise.catch(() => undefined);
await started.callback.waitForCallback();
if (terminal === "timeout") {
await vi.advanceTimersByTimeAsync(5_000);
vi.useRealTimers();
} else {
controller.abort();
}
await expect(responsePromise).rejects.toThrow();
await started.callback.complete({ status: 200, body: "Too late", contentType: "text/plain" });
const replacement = await startOAuthLoopbackCallbackServer({
redirectUrl: callbackUrl("127.0.0.1", started.port),
expectedState: "replacement-state",
timeoutMs: 5_000,
});
openCallbacks.push(replacement);
const response = await fetch(
callbackUrl("127.0.0.1", started.port, "?code=new-code&state=replacement-state"),
);
expect(response.status).toBe(200);
await response.text();
await expect(replacement.waitForCallback()).resolves.toMatchObject({ code: "new-code" });
},
);
it("observes aborts that arrive while localhost resolution is pending", async () => {
let releaseLookup!: () => void;
const pendingLookup = new Promise<LookupAddress[]>((resolve) => {
@ -296,35 +405,45 @@ describe("OAuth loopback callback server", () => {
await expect(started.callback.waitForCallback()).resolves.toMatchObject({ code: "ipv6" });
});
it("uses HTTP port 80 when the redirect omits a port and rejects port zero", async () => {
let observedPort: number | undefined;
const fakeServer = {
listening: false,
once: () => fakeServer,
listen: (port: number, _hostname: string, callback: () => void) => {
observedPort = port;
fakeServer.listening = true;
callback();
return fakeServer;
},
removeAllListeners: () => fakeServer,
on: () => fakeServer,
close: (callback: () => void) => {
fakeServer.listening = false;
callback();
return fakeServer;
},
closeAllConnections: () => undefined,
};
const callback = await startOAuthLoopbackCallbackServer({
redirectUrl: "http://127.0.0.1/oauth/callback",
expectedState: "state-1234567890",
timeoutMs: 5_000,
createServer: (() =>
fakeServer as unknown as Server) as typeof import("node:http").createServer,
});
expect(observedPort).toBe(80);
await callback.close();
it.each([undefined, "localhost", "127.0.0.1", "::1"])(
"uses HTTP port 80 and the exact bind host %s",
async (bindOnlyHostname) => {
let observedPort: number | undefined;
let observedHostname: string | undefined;
const fakeServer = {
listening: false,
once: () => fakeServer,
listen: (port: number, hostname: string, callback: () => void) => {
observedPort = port;
observedHostname = hostname;
fakeServer.listening = true;
callback();
return fakeServer;
},
removeAllListeners: () => fakeServer,
on: () => fakeServer,
close: (callback: () => void) => {
fakeServer.listening = false;
callback();
return fakeServer;
},
closeAllConnections: () => undefined,
};
const callback = await startOAuthLoopbackCallbackServer({
redirectUrl: "http://127.0.0.1/oauth/callback",
bindOnlyHostname,
expectedState: "state-1234567890",
timeoutMs: 5_000,
createServer: (() =>
fakeServer as unknown as Server) as typeof import("node:http").createServer,
});
expect(observedPort).toBe(80);
expect(observedHostname).toBe(bindOnlyHostname ?? "127.0.0.1");
await callback.close();
},
);
it("rejects non-loopback bind hosts and port zero", async () => {
await expect(
startOAuthLoopbackCallbackServer({
redirectUrl: "http://127.0.0.1:0/oauth/callback",
@ -332,5 +451,12 @@ describe("OAuth loopback callback server", () => {
timeoutMs: 5_000,
}),
).rejects.toThrow("valid TCP port");
await expect(
startOAuthLoopbackCallbackServer({
redirectUrl: "http://localhost:8080/oauth/callback",
bindOnlyHostname: "0.0.0.0",
expectedState: "state-1234567890",
}),
).rejects.toThrow("OAuth callback bind must use");
});
});

View file

@ -5,11 +5,12 @@ import { oauthErrorHtml, renderOAuthPage } from "../shared/oauth-page.js";
import { OAUTH_PAGE_CSP } from "./oauth-page-csp.js";
type OAuthLoopbackCallbackResult =
| { type: "authorization_code"; code: string; state: string }
| { type: "authorization_code"; code: string; state: string; parameters: URLSearchParams }
| { type: "oauth_error"; error: string; errorDescription?: string };
export type OAuthLoopbackCallbackServer = {
waitForCallback: () => Promise<OAuthLoopbackCallbackResult>;
complete: (response: RenderedResponse & { status: number }) => Promise<void>;
close: () => Promise<void>;
};
@ -64,7 +65,15 @@ function resolveBindAddresses(
redirectUrl: URL,
bindHostname?: string,
lookup?: LoopbackLookup,
bindOnlyHostname?: string,
): string[] | Promise<string[]> {
if (bindOnlyHostname !== undefined) {
const hostname = unbracket(bindOnlyHostname);
if (!["localhost", "127.0.0.1", "::1"].includes(hostname)) {
throw new Error("OAuth callback bind must use localhost, 127.0.0.1, or ::1");
}
return [hostname];
}
const redirectHostname = unbracket(redirectUrl.hostname);
const redirectAddresses = resolveLoopbackHostname(redirectHostname, lookup);
const requestedHostname = bindHostname ? unbracket(bindHostname) : redirectHostname;
@ -104,6 +113,7 @@ function prepareResponse(
response: ServerResponse,
resolveCorsOrigin?: CorsOriginResolver,
): void {
response.setHeader("Connection", "close");
response.setHeader("Cache-Control", "no-store");
response.setHeader("Content-Security-Policy", OAUTH_PAGE_CSP);
response.setHeader("Referrer-Policy", "no-referrer");
@ -148,9 +158,11 @@ async function closeServers(servers: readonly Server[]): Promise<void> {
export async function startOAuthLoopbackCallbackServer(params: {
redirectUrl: string | URL;
expectedState: string;
timeoutMs: number;
timeoutMs?: number;
signal?: AbortSignal;
bindHostname?: string;
bindOnlyHostname?: string;
deferResponse?: boolean;
lookup?: LoopbackLookup;
createServer?: typeof import("node:http").createServer;
resolveCorsOrigin?: CorsOriginResolver;
@ -165,14 +177,26 @@ export async function startOAuthLoopbackCallbackServer(params: {
) {
throw new Error("OAuth callback redirect must use HTTP on a loopback address");
}
if (!params.expectedState || !Number.isFinite(params.timeoutMs) || params.timeoutMs <= 0) {
if (
!params.expectedState ||
(params.timeoutMs !== undefined &&
(!Number.isFinite(params.timeoutMs) || params.timeoutMs <= 0))
) {
throw new Error("OAuth callback requires state and a positive timeout");
}
if (params.signal?.aborted) {
throw new Error("OAuth callback cancelled");
}
const resolvedAddresses = resolveBindAddresses(redirectUrl, params.bindHostname, params.lookup);
if (params.bindHostname !== undefined && params.bindOnlyHostname !== undefined) {
throw new Error("Choose either an additional or an exact OAuth callback bind host");
}
const resolvedAddresses = resolveBindAddresses(
redirectUrl,
params.bindHostname,
params.lookup,
params.bindOnlyHostname,
);
const addresses = Array.isArray(resolvedAddresses)
? resolvedAddresses
: await waitForAbortable(resolvedAddresses, params.signal);
@ -180,7 +204,9 @@ export async function startOAuthLoopbackCallbackServer(params: {
const callbackPath = redirectUrl.pathname || "/";
const createServer = params.createServer ?? (await import("node:http")).createServer;
const servers: Server[] = [];
let received = false;
let settled = false;
let pendingResponse: ServerResponse | undefined;
let binding = true;
const timeoutRef: { current?: NodeJS.Timeout } = {};
let closePromise: Promise<void> | undefined;
@ -198,13 +224,23 @@ export async function startOAuthLoopbackCallbackServer(params: {
return;
}
settled = true;
pendingResponse = undefined;
cleanup();
callback.reject(error instanceof Error ? error : new Error("OAuth callback failed"));
void close();
};
const onAbort = () => settleError(new Error("OAuth callback cancelled"));
const settleResult = (result: OAuthLoopbackCallbackResult, response: ServerResponse) => {
if (settled) {
received = true;
if (params.deferResponse) {
// Admission consumes the callback, but cancellation owns the socket until verification ends.
pendingResponse = response;
response.once("close", () => {
if (!response.writableFinished) {
settleError(new Error("OAuth callback disconnected"));
}
});
callback.resolve(result);
return;
}
settled = true;
@ -241,21 +277,51 @@ export async function startOAuthLoopbackCallbackServer(params: {
response.writeHead(status, { "Content-Type": rendered.contentType });
response.end(rendered.body);
};
const complete = async (rendered: RenderedResponse & { status: number }) => {
if (settled || !pendingResponse) {
return;
}
const response = pendingResponse;
if (response.destroyed || response.writableFinished) {
settleError(new Error("OAuth callback disconnected"));
await close();
return;
}
pendingResponse = undefined;
const finished = new Promise<void>((resolve) => {
response.once("finish", resolve);
response.once("close", resolve);
});
respond(response, rendered.status, rendered);
await finished;
settled = true;
cleanup();
await close();
};
const handleRequest = (request: IncomingMessage, response: ServerResponse) => {
try {
prepareResponse(request, response, params.resolveCorsOrigin);
if (settled) {
if (received || settled) {
respond(response, 409, renderError("OAuth callback was already received."));
} else if (request.method === "OPTIONS") {
response.writeHead(204).end();
} else {
const url = new URL(request.url ?? "/", redirectUrl.origin);
let url: URL;
try {
url = new URL(request.url ?? "/", redirectUrl.origin);
} catch {
respond(response, 400, renderError("Invalid OAuth callback."));
return;
}
if (url.pathname !== callbackPath) {
respond(response, 404, renderError("Callback route not found."));
} else if (request.method !== "GET") {
response.setHeader("Allow", "GET, OPTIONS");
respond(response, 405, renderError("Method not allowed."));
} else if (url.searchParams.get("state") !== params.expectedState) {
} else if (
url.searchParams.getAll("state").length !== 1 ||
url.searchParams.get("state") !== params.expectedState
) {
respond(response, 400, renderError("Invalid OAuth state."));
} else if (url.searchParams.has("error")) {
const error = url.searchParams.get("error")!;
@ -264,17 +330,26 @@ export async function startOAuthLoopbackCallbackServer(params: {
{ type: "oauth_error", error, ...(errorDescription ? { errorDescription } : {}) },
response,
);
respond(response, 400, renderError("Authorization was not completed."));
if (!params.deferResponse) {
respond(response, 400, renderError("Authorization was not completed."));
}
} else {
const code = url.searchParams.get("code")?.trim();
if (!code) {
if (!code || url.searchParams.getAll("code").length !== 1) {
respond(response, 400, renderError("Missing OAuth authorization code."));
} else {
settleResult(
{ type: "authorization_code", code, state: params.expectedState },
{
type: "authorization_code",
code,
state: params.expectedState,
parameters: url.searchParams,
},
response,
);
respond(response, 200, renderSuccess());
if (!params.deferResponse) {
respond(response, 200, renderSuccess());
}
}
}
}
@ -313,12 +388,15 @@ export async function startOAuthLoopbackCallbackServer(params: {
throw error;
}
binding = false;
timeoutRef.current = setTimeout(
() => settleError(new Error("OAuth callback timeout")),
params.timeoutMs,
);
if (params.timeoutMs !== undefined) {
timeoutRef.current = setTimeout(
() => settleError(new Error("OAuth callback timeout")),
params.timeoutMs,
);
}
return {
waitForCallback: () => callback.promise,
complete,
close: async () => {
if (!settled) {
settleError(new Error("OAuth callback cancelled"));

View file

@ -50,6 +50,7 @@ describe("Anthropic OAuth token responses", () => {
const close = vi.fn(async () => undefined);
startOAuthLoopbackCallbackServer.mockResolvedValueOnce({
waitForCallback: vi.fn(),
complete: vi.fn(async () => undefined),
close,
});
const loginPromise = anthropicOAuthProvider.login({
@ -152,7 +153,12 @@ describe("Anthropic OAuth callback host", () => {
type: "authorization_code" as const,
code: "authorization-code",
state: params.expectedState,
parameters: new URLSearchParams({
code: "authorization-code",
state: params.expectedState,
}),
}),
complete: async () => undefined,
close: async () => undefined,
}));
const tokenExchange = vi.fn(async (_url: string | URL | Request, init?: RequestInit) => {
@ -193,6 +199,7 @@ describe("Anthropic OAuth callback host", () => {
vi.stubEnv("OPENCLAW_OAUTH_CALLBACK_HOST", "127.0.0.1");
startOAuthLoopbackCallbackServer.mockResolvedValueOnce({
waitForCallback: async () => ({ type: "oauth_error", error: "access_denied" }),
complete: async () => undefined,
close: async () => undefined,
});
const login = anthropicOAuthProvider.login({

View file

@ -39,11 +39,13 @@ export type OAuthCallbackResult = {
};
type ProviderOAuthLoopbackCallbackResult =
| { type: "authorization_code"; code: string; state: string }
| { type: "authorization_code"; code: string; state: string; parameters: URLSearchParams }
| { type: "oauth_error"; error: string; errorDescription?: string };
type ProviderOAuthLoopbackCallbackServer = {
waitForCallback: () => Promise<ProviderOAuthLoopbackCallbackResult>;
/** Flushes a deferred browser result, then closes; closed listeners ignore late completion. */
complete: (response: ProviderOAuthLoopbackRenderedResponse & { status: number }) => Promise<void>;
close: () => Promise<void>;
};
@ -59,9 +61,15 @@ type ProviderOAuthLoopbackCorsOriginResolver = (
export async function startProviderOAuthLoopbackCallbackServer(params: {
redirectUrl: string | URL;
expectedState: string;
timeoutMs: number;
/** Optional listener deadline; the caller signal continues to own provider work. */
timeoutMs?: number;
signal?: AbortSignal;
/** Additional loopback host; all addresses of the redirect hostname remain bound. */
bindHostname?: string;
/** Exact Node bind host for providers whose existing redirect uses one address family. */
bindOnlyHostname?: string;
/** Admit callback parameters now, then render the browser outcome through complete(). */
deferResponse?: boolean;
resolveCorsOrigin?: ProviderOAuthLoopbackCorsOriginResolver;
renderSuccess?: () => ProviderOAuthLoopbackRenderedResponse;
renderError?: (message: string) => ProviderOAuthLoopbackRenderedResponse;