mirror of
https://github.com/openclaw/openclaw.git
synced 2026-10-03 09:39:25 +00:00
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:
parent
58b1860232
commit
2079ed92a5
19 changed files with 1418 additions and 446 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
);
|
||||
}
|
||||
|
|
|
|||
31
extensions/openai/openai-oauth-http.runtime.ts
Normal file
31
extensions/openai/openai-oauth-http.runtime.ts
Normal 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 } : {}),
|
||||
});
|
||||
}
|
||||
|
|
@ -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) => {
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
480
src/gateway/server-methods/system-agent-setup-auth-retry.test.ts
Normal file
480
src/gateway/server-methods/system-agent-setup-auth-retry.test.ts
Normal 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());
|
||||
}
|
||||
});
|
||||
});
|
||||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
|
|
|
|||
|
|
@ -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({
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue