diff --git a/.changeset/configurable-rpc-timeouts.md b/.changeset/configurable-rpc-timeouts.md new file mode 100644 index 0000000000..7a1fce8f62 --- /dev/null +++ b/.changeset/configurable-rpc-timeouts.md @@ -0,0 +1,5 @@ +--- +"@browserbasehq/stagehand": minor +--- + +Add configurable response timeouts and per-call abort signal propagation for ordinary TypeScript SDK RPC calls. \ No newline at end of file diff --git a/packages/docs/v4/reference/stagehand.mdx b/packages/docs/v4/reference/stagehand.mdx index 63d1738185..476a8736bb 100644 --- a/packages/docs/v4/reference/stagehand.mdx +++ b/packages/docs/v4/reference/stagehand.mdx @@ -77,6 +77,14 @@ const stagehand = await Stagehand.create({ browser }); Client-side log level, output format, and optional log callback. + + Response-wait deadlines for ordinary Stagehand RPC calls. `defaultMs` applies to every ordinary call; `methods` overrides it for a specific RPC method. Initialization and `experimentalBatch()` retain their existing lifecycle timeout behavior. + + + + Called immediately before each ordinary RPC call. Return an `AbortSignal` to stop waiting for that call when the signal aborts. Aborting a call does not close the Stagehand instance or its browser. + + An initialized Stagehand instance. diff --git a/packages/sdk-ts/src/clientSchemas.ts b/packages/sdk-ts/src/clientSchemas.ts index 49c110dfaa..83f11c297c 100644 --- a/packages/sdk-ts/src/clientSchemas.ts +++ b/packages/sdk-ts/src/clientSchemas.ts @@ -20,6 +20,7 @@ import { StagehandInitParamsSchema, StagehandLogLevelSchema, } from "@browserbasehq/stagehand-protocol/schemas"; +import { StagehandMethods } from "@browserbasehq/stagehand-protocol/schema-registry"; import { Page } from "./page.js"; import { Locator } from "./locator.js"; import { isStagehandBrowser, type StagehandBrowser } from "./browser/index.js"; @@ -267,8 +268,40 @@ export const StagehandBrowserSchema = z ) .meta({ id: "StagehandBrowser" }); +export type StagehandRPCMethodName = + (typeof StagehandMethods)[keyof typeof StagehandMethods]["name"]; +export type StagehandRPCTimeouts = { + defaultMs?: number; + methods?: Partial>; +}; + +const MAX_RPC_TIMEOUT_MS = 2_147_473_647; +const RPCTimeoutMsSchema = z.int().positive().max(MAX_RPC_TIMEOUT_MS); + +export const StagehandRPCTimeoutsSchema = z + .strictObject({ + defaultMs: RPCTimeoutMsSchema.optional(), + methods: z.record(z.string(), RPCTimeoutMsSchema).optional(), + }) + .meta({ id: "StagehandRPCTimeouts" }); + +export const StagehandCallOptionsSchema = z + .strictObject({ + signal: z.custom((value) => value instanceof AbortSignal).optional(), + }) + .meta({ id: "StagehandCallOptions" }); + +export const StagehandGetCallOptionsSchema = z + .custom<() => StagehandCallOptions | undefined>( + (value) => typeof value === "function", + "getCallOptions must be a function", + ) + .meta({ id: "StagehandGetCallOptions" }); + export const StagehandCreateOptionsSchema = StagehandClientCreateConfigSchema.extend({ browser: StagehandBrowserSchema, + rpcTimeouts: StagehandRPCTimeoutsSchema.optional(), + getCallOptions: StagehandGetCallOptionsSchema.optional(), }).meta({ id: "StagehandCreateOptions" }); export type ClientLLM = z.infer; @@ -298,6 +331,8 @@ export type ResolvedStagehandClientCreateConfig = z.output< >; export type StagehandCreateOptions = z.input; export type ResolvedStagehandCreateOptions = z.output; +export type StagehandCallOptions = z.input; +export type StagehandGetCallOptions = z.input; export type WebMCPToolsOptions = z.infer; export type WebMCPInvokeOptions = z.infer; export type WebMCPResultOptions = z.infer; diff --git a/packages/sdk-ts/src/index.ts b/packages/sdk-ts/src/index.ts index 93729ed49d..6248563af8 100644 --- a/packages/sdk-ts/src/index.ts +++ b/packages/sdk-ts/src/index.ts @@ -118,19 +118,26 @@ export { StagehandClientLoggingConfigSchema, StagehandClientLogLevelSchema, StagehandClientCreateConfigSchema, + StagehandCallOptionsSchema, StagehandBrowserSchema, StagehandCreateOptionsSchema, + StagehandGetCallOptionsSchema, + StagehandRPCTimeoutsSchema, WebMCPInvokeOptionsSchema, WebMCPResultOptionsSchema, WebMCPToolsOptionsSchema, type ClientLLM, type ResolvedStagehandClientLoggingConfig, type StagehandClientActOptions, + type StagehandCallOptions, type StagehandClientExtractOptions, type StagehandClientLoggingConfig, type StagehandClientObserveOptions, type StagehandClientCreateConfig, type StagehandCreateOptions, + type StagehandGetCallOptions, + type StagehandRPCMethodName, + type StagehandRPCTimeouts, type ResolvedStagehandCreateOptions, type WebMCPInvokeOptions, type WebMCPResultOptions, diff --git a/packages/sdk-ts/src/rpcClient.ts b/packages/sdk-ts/src/rpcClient.ts index dee85b5253..9011df834c 100644 --- a/packages/sdk-ts/src/rpcClient.ts +++ b/packages/sdk-ts/src/rpcClient.ts @@ -38,7 +38,12 @@ import { import type { StagehandRpcNotification } from "@browserbasehq/stagehand-protocol/types"; import { z } from "zod/v4"; import { CDPClient, type ServiceWorkerInfo } from "./cdpClient.js"; -import { abortReason } from "./abort.js"; +import { + StagehandCallOptionsSchema, + type StagehandGetCallOptions, + type StagehandRPCTimeouts, +} from "./clientSchemas.js"; +import { abortReason, throwIfAborted } from "./abort.js"; type PendingRequest = { method: RPCMethod; @@ -144,9 +149,19 @@ export class RPCClient { pendingNotifications: StagehandRpcNotification[] = []; closed = false; readonly cdp: CDPTransport; - - constructor(cdp: CDPTransport) { + readonly rpcTimeouts?: StagehandRPCTimeouts; + readonly getCallOptions?: StagehandGetCallOptions; + + constructor( + cdp: CDPTransport, + options: { + rpcTimeouts?: StagehandRPCTimeouts; + getCallOptions?: StagehandGetCallOptions; + } = {}, + ) { this.cdp = cdp; + this.rpcTimeouts = options.rpcTimeouts; + this.getCallOptions = options.getCallOptions; this.serviceWorker = cdp.serviceWorker; this.browserWebSocketDebuggerUrl = cdp.webSocketDebuggerUrl; this.cdp.onmessage = (message) => this.receive(message); @@ -163,6 +178,16 @@ export class RPCClient { if (method.name === StagehandMethods.stagehandInit.name && !options.signal) { throw new Error("stagehand.init requires an initialization lifecycle signal"); } + const callOptions = + method.name === StagehandMethods.stagehandInit.name || + method.name === StagehandMethods.stagehandCallbackBatch.name + ? undefined + : this.getCallOptions?.(); + const callSignal = combineAbortSignals( + options.signal, + callOptions === undefined ? undefined : StagehandCallOptionsSchema.parse(callOptions).signal, + ); + throwIfAborted(callSignal); const parentContext = context.active(); const span = TRACER.startSpan( @@ -190,13 +215,10 @@ export class RPCClient { ...getTraceContextFields(requestContext), }); span.setAttribute("jsonrpc.request.id", String(request.id)); - const responseTimeoutMs = rpcResponseTimeoutMs(method.name, parsedParams); + const responseTimeoutMs = rpcResponseTimeoutMs(method.name, parsedParams, this.rpcTimeouts); const timeoutController = responseTimeoutMs === undefined ? undefined : new AbortController(); - const signal = - options.signal && timeoutController - ? AbortSignal.any([options.signal, timeoutController.signal]) - : (options.signal ?? timeoutController?.signal); + const signal = combineAbortSignals(callSignal, timeoutController?.signal); const timeoutId = timeoutController && responseTimeoutMs !== undefined ? setTimeout(() => { @@ -213,6 +235,7 @@ export class RPCClient { const [, result] = await Promise.all([ this.cdp.send(request, signal).catch((error: unknown) => { this.rejectPending(request.id, asError(error)); + throw error; }), response, ]); @@ -509,7 +532,25 @@ function asError(error: unknown): Error { return error instanceof Error ? error : new Error(String(error)); } -export function rpcResponseTimeoutMs(method: string, params: unknown): number | undefined { +function combineAbortSignals(...signals: Array): AbortSignal | undefined { + const unique = [...new Set(signals.filter((signal) => signal !== undefined))]; + if (unique.length === 0) return undefined; + if (unique.length === 1) return unique[0]; + return AbortSignal.any(unique); +} + +export function rpcResponseTimeoutMs( + method: string, + params: unknown, + configuredTimeouts?: StagehandRPCTimeouts, +): number | undefined { + const configuredTimeoutMs = + method === StagehandMethods.stagehandInit.name || + method === StagehandMethods.stagehandCallbackBatch.name + ? undefined + : ((configuredTimeouts?.methods as Record | undefined)?.[method] ?? + configuredTimeouts?.defaultMs); + if (configuredTimeoutMs !== undefined) return configuredTimeoutMs; let operationTimeoutMs: number | undefined; switch (method) { case StagehandMethods.stagehandAct.name: diff --git a/packages/sdk-ts/src/stagehand.ts b/packages/sdk-ts/src/stagehand.ts index f5a10d4fe0..9505e7ae33 100644 --- a/packages/sdk-ts/src/stagehand.ts +++ b/packages/sdk-ts/src/stagehand.ts @@ -29,6 +29,8 @@ import { StagehandClientObserveOptionsSchema, type StagehandClientActOptions, type StagehandClientExtractOptions, + type StagehandGetCallOptions, + type StagehandRPCTimeouts, type ResolvedStagehandClientLoggingConfig, type ResolvedStagehandClientCreateConfig, type StagehandCreateOptions, @@ -75,12 +77,17 @@ export class Stagehand { private constructor( private readonly browserHandle: StagehandBrowser, private readonly createConfig: ResolvedStagehandClientCreateConfig, + private readonly rpcClientOptions: { + rpcTimeouts?: StagehandRPCTimeouts; + getCallOptions?: StagehandGetCallOptions; + }, ) {} static async create(input: StagehandCreateOptions): Promise { - const { browser, ...createConfig } = StagehandCreateOptionsSchema.parse(input); + const { browser, rpcTimeouts, getCallOptions, ...createConfig } = + StagehandCreateOptionsSchema.parse(input); const claimedBrowser = claimStagehandBrowser(browser); - const stagehand = new Stagehand(browser, createConfig); + const stagehand = new Stagehand(browser, createConfig, { rpcTimeouts, getCallOptions }); let lifecycleSignal: AbortSignal | undefined; try { await withStagehandInitDeadline((signal) => { @@ -171,7 +178,7 @@ export class Stagehand { private async initialize(browser: ClaimedStagehandBrowser, signal: AbortSignal): Promise { const createConfig = this.createConfig; - const rpcClient = new RPCClient(browser.cdpClient); + const rpcClient = new RPCClient(browser.cdpClient, this.rpcClientOptions); this.rpcClient = rpcClient; try { diff --git a/packages/sdk-ts/tests/packageContract.test.ts b/packages/sdk-ts/tests/packageContract.test.ts index 58f5e1228c..b333ada2c0 100644 --- a/packages/sdk-ts/tests/packageContract.test.ts +++ b/packages/sdk-ts/tests/packageContract.test.ts @@ -54,6 +54,9 @@ describe("published TypeScript SDK", () => { LocalBrowserConnectOptionsSchema, Response, Stagehand, + StagehandCallOptionsSchema, + StagehandGetCallOptionsSchema, + StagehandRPCTimeoutsSchema, WebMCPInvocation, WebMCPTool, WebMCPToolsOptionsSchema, @@ -75,6 +78,12 @@ describe("published TypeScript SDK", () => { } LocalBrowserConnectOptionsSchema.parse({ cdpUrl: "ws://127.0.0.1:9222" }); BrowserbaseConnectOptionsSchema.parse({ apiKey: "bb_key", sessionId: "session_123" }); + StagehandRPCTimeoutsSchema.parse({ + defaultMs: 1000, + methods: { "page.goto": 2000 }, + }); + StagehandCallOptionsSchema.parse({}); + StagehandGetCallOptionsSchema.parse(() => undefined); if (typeof WebMCPTool !== "function") throw new Error("WebMCPTool export is unavailable"); if (typeof WebMCPInvocation !== "function") { throw new Error("WebMCPInvocation export is unavailable"); @@ -116,8 +125,12 @@ describe("published TypeScript SDK", () => { RgbaColor, SnapshotResult, StagehandClientActOptions, + StagehandCallOptions, StagehandClientExtractOptions, StagehandClientObserveOptions, + StagehandGetCallOptions, + StagehandRPCMethodName, + StagehandRPCTimeouts, StagehandResultUsage, Variables, } from "@browserbasehq/stagehand"; @@ -149,6 +162,13 @@ describe("published TypeScript SDK", () => { const actOptions: StagehandClientActOptions = { cache: caching, model, variables }; const observeOptions: StagehandClientObserveOptions = { model, variables }; const extractOptions: StagehandClientExtractOptions = { model }; + const rpcMethod: StagehandRPCMethodName = "page.goto"; + const rpcTimeouts: StagehandRPCTimeouts = { + defaultMs: 1_000, + methods: { [rpcMethod]: 2_000 }, + }; + const callOptions: StagehandCallOptions = {}; + const getCallOptions: StagehandGetCallOptions = () => callOptions; declare const centroid: LocatorCentroidResult; declare const snapshot: SnapshotResult; @@ -172,6 +192,8 @@ describe("published TypeScript SDK", () => { actOptions, observeOptions, extractOptions, + rpcTimeouts, + getCallOptions, centroid, snapshot, usage, diff --git a/packages/sdk-ts/tests/rpcClient.test.ts b/packages/sdk-ts/tests/rpcClient.test.ts index 3f17cdddd1..6f1a4d1885 100644 --- a/packages/sdk-ts/tests/rpcClient.test.ts +++ b/packages/sdk-ts/tests/rpcClient.test.ts @@ -16,6 +16,7 @@ import { rpcResponseTimeoutMs, type CDPTransport, } from "../src/rpcClient.js"; +import { StagehandRPCTimeoutsSchema } from "../src/clientSchemas.js"; const UppercaseMethod = { name: "test.uppercase", @@ -73,6 +74,13 @@ class ManualCDPTransport implements CDPTransport { } } +class FailingCDPTransport extends ManualCDPTransport { + override async send(message: JSONRPCMessage): Promise { + this.sent.push(message); + throw new Error("transport send failed"); + } +} + describe("RPCClientOptionsSchema", () => { it("accepts the preloaded-extension option", () => { const signal = new AbortController().signal; @@ -101,6 +109,24 @@ describe("RPCClientOptionsSchema", () => { }); }); +describe("StagehandRPCTimeoutsSchema", () => { + it("accepts positive integer timeouts within the timer limit", () => { + expect( + StagehandRPCTimeoutsSchema.parse({ + defaultMs: 2_147_473_647, + methods: { "page.goto": 1 }, + }), + ).toStrictEqual({ + defaultMs: 2_147_473_647, + methods: { "page.goto": 1 }, + }); + }); + + it.each([0, -1, 1.5, 2_147_473_648])("rejects an invalid timeout value: %s", (timeout) => { + expect(() => StagehandRPCTimeoutsSchema.parse({ defaultMs: timeout })).toThrow(); + }); +}); + describe("RPCClient", () => { it("keeps callback batches in the ordinary pending RPC path", async () => { const source = "async () => undefined"; @@ -419,6 +445,121 @@ describe("RPCClient", () => { expect(rpcResponseTimeoutMs(method, {})).toBe(timeout); }); + it("prefers configured per-method RPC timeouts over the configured default", () => { + expect( + rpcResponseTimeoutMs( + StagehandMethods.pageGoto.name, + {}, + { + defaultMs: 20_000, + methods: { [StagehandMethods.pageGoto.name]: 5_000 }, + }, + ), + ).toBe(5_000); + expect(rpcResponseTimeoutMs(StagehandMethods.pageReload.name, {}, { defaultMs: 20_000 })).toBe( + 20_000, + ); + }); + + it("does not apply configured RPC timeouts to initialization or callback batches", () => { + const timeouts = { defaultMs: 1 }; + expect(rpcResponseTimeoutMs(StagehandMethods.stagehandInit.name, {}, timeouts)).toBeUndefined(); + expect( + rpcResponseTimeoutMs( + StagehandMethods.stagehandCallbackBatch.name, + { options: { timeout: 2_000 } }, + timeouts, + ), + ).toBe(12_000); + }); + + it("uses the current call signal without closing the client", async () => { + const cdp = new ManualCDPTransport(); + const controller = new AbortController(); + let signal: AbortSignal | undefined = controller.signal; + const client = new RPCClient(cdp, { + getCallOptions: () => (signal ? { signal } : undefined), + }); + const pending = client.send(StagehandMethods.contextPages, {}); + + await vi.waitFor(() => expect(cdp.sent).toHaveLength(1)); + const reason = new Error("cell cancelled"); + controller.abort(reason); + await expect(pending).rejects.toBe(reason); + expect(client.closed).toBe(false); + + signal = undefined; + const followUp = client.send(StagehandMethods.contextPages, {}); + await vi.waitFor(() => expect(cdp.sent).toHaveLength(2)); + await cdp.receive({ jsonrpc: "2.0", id: 2, result: [] }); + await expect(followUp).resolves.toStrictEqual([]); + }); + + it("gets call options once for each ordinary RPC call", async () => { + let calls = 0; + const client = new RPCClient(new FakeCDPTransport([]), { + getCallOptions: () => { + calls += 1; + return undefined; + }, + }); + + await client.send(StagehandMethods.contextPages, {}); + await client.send(StagehandMethods.contextPages, {}); + + expect(calls).toBe(2); + }); + + it("rejects a pre-aborted call signal before sending", async () => { + const cdp = new ManualCDPTransport(); + const controller = new AbortController(); + const reason = new Error("cell cancelled"); + controller.abort(reason); + const client = new RPCClient(cdp, { + getCallOptions: () => ({ signal: controller.signal }), + }); + + await expect(client.send(StagehandMethods.contextPages, {})).rejects.toBe(reason); + expect(cdp.sent).toHaveLength(0); + }); + + it("does not read call options for initialization or callback batches", async () => { + const getCallOptions = () => { + throw new Error("call options must not be read"); + }; + const client = new RPCClient(new FakeCDPTransport({ initialized: true, pages: [] }), { + getCallOptions, + }); + const callbackClient = new RPCClient(new FakeCDPTransport({}), { + getCallOptions, + }); + const signal = new AbortController().signal; + + await expect( + client.sendStagehandInit( + { + protocolVersion: STAGEHAND_PROTOCOL_VERSION, + clientInfo: { name: "test", version: "1.0.0" }, + }, + signal, + ), + ).resolves.toStrictEqual({ initialized: true, pages: [] }); + await expect( + callbackClient.send(StagehandMethods.stagehandCallbackBatch, { + callbackSource: "async () => undefined", + options: { timeout: 1_000 }, + }), + ).resolves.toStrictEqual({}); + }); + + it("rethrows a transport send error after rejecting its pending request", async () => { + const client = new RPCClient(new FailingCDPTransport()); + + await expect(client.send(StagehandMethods.contextPages, {})).rejects.toThrow( + "transport send failed", + ); + }); + it("does not impose response deadlines on operations that were unbounded in v3", () => { const methods = [ StagehandMethods.stagehandInit.name,