Skip to content
34 changes: 24 additions & 10 deletions src/client/connect.ts
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,8 @@ export interface LinkClientCredential {
}

export interface ClientConnectDeps {
/** Abort enrollment network work and refuse subsequent writes; rollback still drains. */
signal?: AbortSignal;
fetchImpl?: typeof fetch;
now?: () => Date;
lifecycleLockDeps?: ClientLifecycleLockDeps;
Expand Down Expand Up @@ -539,6 +541,18 @@ export async function connectClient(
options: ConnectOptions,
deps: ClientConnectDeps = {},
): Promise<OcxClientConnectionConfig> {
deps.signal?.throwIfAborted();
const rawFetch = deps.fetchImpl ?? fetch;
const fetchImpl: typeof fetch = deps.signal ? Object.assign(async (...[input, init = {}]: Parameters<typeof fetch>) => {
deps.signal!.throwIfAborted();
const signals = [deps.signal, init.signal, input instanceof Request ? input.signal : undefined]
.filter((signal): signal is AbortSignal => signal != null);
return rawFetch(input, { ...init, signal: AbortSignal.any(signals), redirect: "manual" });
}, { preconnect: rawFetch.preconnect }) : rawFetch;
const assertActiveConnectingState = (fingerprint?: string) => {
deps.signal?.throwIfAborted();
assertConnectingState(fingerprint);
};
let serverUrl = "";
let managementUrl = "";
let linkAdmissionToken: string | null = null;
Expand Down Expand Up @@ -586,7 +600,7 @@ export async function connectClient(
throw new Error("link mode requires a valid local config port");
}
withClientLifecycleSync(() => withConfigMutationLockSync(() => {
assertConnectingState();
assertActiveConnectingState();
const externalProvider = currentExternalCodexModelProvider();
if (externalProvider) throw new Error("connect refused: an external Codex provider owns config.toml");
if (linkMode) {
Expand All @@ -603,7 +617,7 @@ export async function connectClient(
}
}), deps.lifecycleLockDeps);

const upstreamFetch = deps.fetchImpl ?? fetch;
const upstreamFetch = fetchImpl;
const readinessFetch = linkMode
? (async (input, init = {}) => {
const url = input instanceof Request ? input.url : String(input);
Expand All @@ -622,18 +636,18 @@ export async function connectClient(
managementUrl,
localGuiOrigin(),
options.credential.value,
{ fetchImpl: deps.fetchImpl },
{ fetchImpl },
);
cleanupCredential = { kind: "gui-session", value: session };
} else if (!linkMode && options.credential.kind === "admin") {
cleanupCredential = { kind: "admin", value: options.credential.value };
}
if (!linkMode) {
issued = await issueClientKey(managementUrl, cleanupCredential!, clientKeyName(), { fetchImpl: deps.fetchImpl });
issued = await issueClientKey(managementUrl, cleanupCredential!, clientKeyName(), { fetchImpl });
}

const initialFiles = withClientLifecycleSync(() => withConfigMutationLockSync(() => {
assertConnectingState(linkMode ? pendingConnectFingerprint! : undefined);
assertActiveConnectingState(linkMode ? pendingConnectFingerprint! : undefined);
const persisted = linkMode
? (() => {
const current = readServiceApiTokenState();
Expand All @@ -656,7 +670,7 @@ export async function connectClient(
if (!admissionToken) throw new Error("client admission credential unavailable");
const apiKeyId = linkMode ? (options.credential as LinkClientCredential).apiKeyId : issued!.id;
const catalog = await downloadClientCatalog(serverUrl, admissionToken, {
fetchImpl: deps.fetchImpl,
fetchImpl,
timeoutMs: options.catalogTimeoutMs,
});
// Fail closed BEFORE the write (#4207). The hub being reachable and the credential working
Expand All @@ -666,7 +680,7 @@ export async function connectClient(
// than writing one and restoring it afterwards.
assertClientCatalogCompatible(catalog.body, deps.catalogCompatibility);
writtenCatalogFingerprint = withClientLifecycleSync(() => withConfigMutationLockSync(() => {
assertConnectingState(persisted.fingerprint);
assertActiveConnectingState(persisted.fingerprint);
atomicWriteFile(DEFAULT_CATALOG_PATH, catalog.body);
return sha256(catalog.body);
}), deps.lifecycleLockDeps);
Expand All @@ -679,7 +693,7 @@ export async function connectClient(
routingTarget: target,
catalogPath: DEFAULT_CATALOG_PATH,
journalOwner: { kind: "client", apiKeyId },
beforeClientWrite: () => assertConnectingState(persisted.fingerprint),
beforeClientWrite: () => assertActiveConnectingState(persisted.fingerprint),
});
if (!preflight.success) throw new Error(preflight.message);

Expand All @@ -688,7 +702,7 @@ export async function connectClient(
routingTarget: target,
catalogPath: DEFAULT_CATALOG_PATH,
journalOwner: { kind: "client", apiKeyId },
beforeClientWrite: () => assertConnectingState(persisted.fingerprint),
beforeClientWrite: () => assertActiveConnectingState(persisted.fingerprint),
});
if (!injected.success || injected.status === "skipped") throw new Error(injected.message);
injectionCommitted = true;
Expand All @@ -715,7 +729,7 @@ export async function connectClient(
...(linkMode ? { transport: "link" as const, link: linkMetadata } : {}),
};
withClientLifecycleSync(() => withConfigMutationLockSync(() => {
assertConnectingState(persisted.fingerprint);
assertActiveConnectingState(persisted.fingerprint);
clearClientConnectPending(persisted.fingerprint);
commitClientConnection(connection);
committed = true;
Expand Down
105 changes: 86 additions & 19 deletions src/client/link-join.ts
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import { randomBytes } from "node:crypto";
import { hostname } from "node:os";
import { isPortAvailable } from "../server/ports";
import { scanListenPidsForAddress, type ListenPidScan } from "../server/port-reclaim";
import { isLinkPort, JOIN_TUNNEL_PORT_MAX, JOIN_TUNNEL_PORT_MIN } from "../link/ports";
import { buildExecArgv, REMOTE_COMMAND_NOT_FOUND, remoteOcxArgv } from "../link/ssh-argv";
import { sshFailureHint, sshRunnerErrorHint, type SshRunner, type SshRunResult } from "../link/ssh-runner";
Expand All @@ -22,6 +23,7 @@ import type { OcxConnectedClientId } from "../types";

const JOIN_TUNNEL_READY_TIMEOUT_MS = 15_000;
const JOIN_TUNNEL_POLL_MS = 100;
const JOIN_TUNNEL_SPAWN_GRACE_MS = 100;
const JOIN_REVOKE_TIMEOUT_MS = 30_000;
const JOIN_CONFIRM_TTL_MS = 5 * 60_000;
const JOIN_PORT_ATTEMPTS = 32;
Expand Down Expand Up @@ -74,6 +76,12 @@ export interface ClientLinkJoinDeps {
hostname?: () => string;
randomBytes?: (size: number) => Uint8Array;
fetchImpl?: typeof fetch;
/**
* LISTEN-owner probe for the tunnel port; defaults to the netstat/lsof/ss scan.
* Receives the loopback address the tunnel binds so listeners on unrelated
* addresses do not confuse the readiness check.
*/
scanListenPids?: (port: number, address?: string) => ListenPidScan;
spawnTunnel?: (spec: {
linkId: string;
alias: string;
Expand Down Expand Up @@ -222,22 +230,58 @@ async function compensateStaleSidecar(deps: ClientLinkJoinDeps): Promise<void> {
}
}

/** Wait for the live tunnel's authenticated readiness, retaining the deadline after failed ownership rechecks. */
async function waitForReady(
deps: ClientLinkJoinDeps,
tunnel: ClientLinkTunnelHandle,
port: number,
key: string,
): Promise<void> {
const fetchImpl = deps.fetchImpl ?? fetch;
const now = deps.now ?? Date.now;
const sleep = deps.sleep ?? ((ms: number) => new Promise<void>(resolve => setTimeout(resolve, ms)));
const deadline = now() + JOIN_TUNNEL_READY_TIMEOUT_MS;
const tunnelExited = tunnel.exited.then(() => { throw new ClientLinkJoinError("join_tunnel_failed"); });
await Promise.race([
tunnelExited,
new Promise<void>(resolve => setTimeout(resolve, JOIN_TUNNEL_SPAWN_GRACE_MS)),
]);
const listenPids = deps.scanListenPids ?? scanListenPidsForAddress;
// The tunnel binds 127.0.0.1; a listener on a different loopback or interface address
// never receives our requests, so ownership is only judged among sockets that serve it.
const tunnelAddress = "127.0.0.1";
for (;;) {
try {
const response = await fetchImpl(`http://127.0.0.1:${port}/readyz`, {
headers: { "x-opencodex-api-key": key },
});
if (response.status === 200) return;
if (response.status === 401) throw new ClientLinkJoinError("admission_failed");
// A squatter answering the 401 challenge would otherwise collect the issued key:
// the only listener allowed a keyed request is the ssh process we spawned — it owns
// the port only after a successful bind, and ExitOnForwardFailure makes it exit when
// it cannot take the port. An unverifiable scan stays "not ready", never a pass.
const ownership = listenPids(port, tunnelAddress);
if (ownership.ok && ownership.pids.length === 1 && ownership.pids[0] === tunnel.pid) {
// Never follow redirects: a port occupant must not reroute the challenge, and a
// redirected keyed request would carry the issued key to an unrelated listener.
const probe = await Promise.race([
tunnelExited,
fetchImpl(`http://127.0.0.1:${port}/readyz`, { redirect: "manual" }),
]);
if (probe.status === 401) {
// Ownership can flip between the probe and the keyed request (a squatter
// takes the port after the tunnel dies). Re-scan in the same iteration and
// skip only the keyed request on failure, not the deadline check and sleep.
const recheck = listenPids(port, tunnelAddress);
if (recheck.ok && recheck.pids.length === 1 && recheck.pids[0] === tunnel.pid) {
const response = await Promise.race([
tunnelExited,
fetchImpl(`http://127.0.0.1:${port}/readyz`, {
headers: { "x-opencodex-api-key": key },
redirect: "manual",
}),
]);
if (response.status === 200) return;
if (response.status === 401) throw new ClientLinkJoinError("admission_failed");
}
}
}
} catch (error) {
if (error instanceof ClientLinkJoinError) throw error;
}
Expand All @@ -256,6 +300,7 @@ function requireConfirmedHost(deps: ClientLinkJoinDeps, alias: string): JoinConf
return confirmed;
}

/** Issue and enroll a confirmed Home link, compensating failures before committing the connection. */
export async function joinHome(deps: ClientLinkJoinDeps, input: { alias: string }): Promise<{ linkId: string; apiKeyId: string }> {
const confirmed = requireConfirmedHost(deps, input.alias);
await compensateStaleSidecar(deps);
Expand Down Expand Up @@ -308,30 +353,52 @@ export async function joinHome(deps: ClientLinkJoinDeps, input: { alias: string
configDir: deps.configDir,
knownHostsFile: deps.knownHostsFile,
});
await waitForReady(deps, tunnelPort, issued.key);
await waitForReady(deps, tunnel, tunnelPort, issued.key);
} catch (error) {
const code = error instanceof ClientLinkJoinError ? error.code : "join_tunnel_failed";
await rollback(deps, issued.linkId, tunnel);
throw new ClientLinkJoinError(code);
}

const enrollmentAbort = new AbortController();
let enrollment: Promise<unknown> | undefined;
let enrollmentFinished = false;
try {
if (!tunnel) throw new ClientLinkJoinError("join_tunnel_failed");
const connect = deps.connect ?? connectClient;
await connect({
serverUrl: `http://127.0.0.1:${tunnelPort}`,
managementUrl: `http://127.0.0.1:${tunnelPort}`,
credential: { kind: "link", apiKeyId: issued.apiKeyId, key: issued.key },
transport: "link",
link: { tunnelPort, linkId: issued.linkId },
selectedClients: deps.selectedClients ?? ["codex", "claude"],
managementTransport: "direct",
}, {
fetchImpl: deps.fetchImpl,
...deps.connectDeps,
// Keep watching the tunnel until the connection commits: an exited tunnel
// must not let the issued key ride out to whatever next holds the port.
const tunnelFailure = tunnel.exited.then(() => {
const error = new ClientLinkJoinError("join_tunnel_failed");
if (!enrollmentFinished) enrollmentAbort.abort(error);
throw error;
});
} catch {
const signal = deps.connectDeps?.signal
? AbortSignal.any([enrollmentAbort.signal, deps.connectDeps.signal]) : enrollmentAbort.signal;
enrollment = Promise.resolve().then(() => connect({
serverUrl: `http://127.0.0.1:${tunnelPort}`,
managementUrl: `http://127.0.0.1:${tunnelPort}`,
credential: { kind: "link", apiKeyId: issued.apiKeyId, key: issued.key },
transport: "link",
link: { tunnelPort, linkId: issued.linkId },
selectedClients: deps.selectedClients ?? ["codex", "claude"],
managementTransport: "direct",
}, {
fetchImpl: deps.fetchImpl,
...deps.connectDeps,
signal,
}));
await Promise.race([tunnelFailure, enrollment]);
enrollmentFinished = true;
} catch (error) {
enrollmentAbort.abort(error);
// An observed tunnel exit is not completion of the losing exchange. Drain its
// cancellation and local rollback before stopping/revoking the shared link.
if (enrollment) await Promise.allSettled([enrollment]);
await rollback(deps, issued.linkId, tunnel);
throw new ClientLinkJoinError("join_connect_failed");
throw new ClientLinkJoinError(error instanceof ClientLinkJoinError ? error.code : "join_connect_failed");
} finally {
enrollmentFinished = true;
}

await stopTunnel(tunnel);
Expand Down
Loading
Loading