diff --git a/packages/sdk-go/internal/extensionassets/stagehand-extension.zip b/packages/sdk-go/internal/extensionassets/stagehand-extension.zip index fa42ed748..f43334303 100644 Binary files a/packages/sdk-go/internal/extensionassets/stagehand-extension.zip and b/packages/sdk-go/internal/extensionassets/stagehand-extension.zip differ diff --git a/packages/sdk-ts/tests/integration/timeouts.test.ts b/packages/sdk-ts/tests/integration/timeouts.test.ts new file mode 100644 index 000000000..acd198b0d --- /dev/null +++ b/packages/sdk-ts/tests/integration/timeouts.test.ts @@ -0,0 +1,36 @@ +import { z } from "zod/v4"; +import { afterEach, beforeEach, describe, expect, it } from "vitest"; +import type { ClientLLM, Stagehand } from "../../src/index.js"; +import { closeStagehand, createStagehand, firstPage } from "./_support.js"; + +const hangingModel: ClientLLM = { + generate: () => new Promise(() => {}), +}; + +describe("Stagehand operation timeouts", () => { + let stagehand: Stagehand; + + beforeEach(async () => { + stagehand = await createStagehand({ model: hangingModel }); + const page = await firstPage(stagehand); + await page.goto("data:text/html,

Timeout fixture

"); + }); + + afterEach(async () => { + await closeStagehand(stagehand); + }); + + it("observe() enforces timeout", async () => { + await expect(stagehand.observe("find something", { timeout: 5 })).rejects.toThrow(/timed out/i); + }, 5_000); + + it("extract() enforces timeout", async () => { + await expect( + stagehand.extract("Extract title", z.object({ title: z.string() }), { timeout: 5 }), + ).rejects.toThrow(/timed out/i); + }, 5_000); + + it("act() enforces timeout", async () => { + await expect(stagehand.act("click Continue", { timeout: 5 })).rejects.toThrow(/timed out/i); + }, 5_000); +}); diff --git a/packages/server/controllers/stagehandController.ts b/packages/server/controllers/stagehandController.ts index 42988de09..f60968235 100644 --- a/packages/server/controllers/stagehandController.ts +++ b/packages/server/controllers/stagehandController.ts @@ -13,6 +13,7 @@ import * as cacheService from "../services/cacheService.js"; import * as extractService from "../services/extractService.js"; import { buildGatewayContext } from "../llm/gatewayClient.js"; import * as observeService from "../services/observeService.js"; +import { withTimeout } from "../timeoutConfig.js"; export type StagehandControllerOptions = { initialize?: (params: StagehandInitParams) => Promise; @@ -51,18 +52,22 @@ export function createStagehandController( throw new Error("An LLM was not configured during Stagehand initialization"); } - const result = await actService.act({ - params, - page: runtime.resolveUnderstudyPage(params.pageId), - model, - clientLLMGenerate: runtime.adapters.clientLLMGenerate, - logger, - systemPrompt: state.initParams.systemPrompt, - selfHeal: state.initParams.selfHeal, - domSettleTimeoutMs: state.initParams.domSettleTimeoutMs, - cache: cacheService.buildCacheContext(state.initParams), - gateway, - }); + const result = await withTimeout( + actService.act({ + params, + page: runtime.resolveUnderstudyPage(params.pageId), + model, + clientLLMGenerate: runtime.adapters.clientLLMGenerate, + logger, + systemPrompt: state.initParams.systemPrompt, + selfHeal: state.initParams.selfHeal, + domSettleTimeoutMs: state.initParams.domSettleTimeoutMs, + cache: cacheService.buildCacheContext(state.initParams), + gateway, + }), + params.options?.timeout, + "act()", + ); runtime.metrics.record("act", result.metadata.usage); return result; } @@ -80,16 +85,20 @@ export function createStagehandController( throw new Error("An LLM was not configured during Stagehand initialization"); } - const result = await observeService.observe({ - params, - page: runtime.resolvePage(params.pageId), - model, - clientLLMGenerate: runtime.adapters.clientLLMGenerate, - logger, - systemPrompt: state.initParams.systemPrompt, - cache: cacheService.buildCacheContext(state.initParams), - gateway, - }); + const result = await withTimeout( + observeService.observe({ + params, + page: runtime.resolvePage(params.pageId), + model, + clientLLMGenerate: runtime.adapters.clientLLMGenerate, + logger, + systemPrompt: state.initParams.systemPrompt, + cache: cacheService.buildCacheContext(state.initParams), + gateway, + }), + params.options?.timeout, + "observe()", + ); runtime.metrics.record("observe", result.metadata.usage); return result; } @@ -107,16 +116,20 @@ export function createStagehandController( throw new Error("An LLM was not configured during Stagehand initialization"); } - const result = await extractService.extract({ - params, - page: runtime.resolvePage(params.pageId), - model, - clientLLMGenerate: runtime.adapters.clientLLMGenerate, - logger, - systemPrompt: state.initParams.systemPrompt, - cache: cacheService.buildCacheContext(state.initParams), - gateway, - }); + const result = await withTimeout( + extractService.extract({ + params, + page: runtime.resolvePage(params.pageId), + model, + clientLLMGenerate: runtime.adapters.clientLLMGenerate, + logger, + systemPrompt: state.initParams.systemPrompt, + cache: cacheService.buildCacheContext(state.initParams), + gateway, + }), + params.options?.timeout, + "extract()", + ); runtime.metrics.record("extract", result.metadata.usage); return result; } diff --git a/packages/server/tests/stagehand-controller-timeout.test.ts b/packages/server/tests/stagehand-controller-timeout.test.ts new file mode 100644 index 000000000..e688d6d65 --- /dev/null +++ b/packages/server/tests/stagehand-controller-timeout.test.ts @@ -0,0 +1,110 @@ +import { afterEach, describe, expect, it, vi } from "vitest"; +import { StagehandInitParamsSchema, STAGEHAND_PROTOCOL_VERSION } from "../../protocol/schemas.js"; +import { createStagehandController } from "../controllers/stagehandController.js"; +import type { HandlerContext } from "../rpcRouter.js"; +import { createStagehandRuntime } from "../runtime.js"; +import * as actService from "../services/actService.js"; +import * as extractService from "../services/extractService.js"; +import * as observeService from "../services/observeService.js"; +import { zeroStagehandResultUsage } from "../services/resultUsage.js"; + +const TIMEOUT_MS = 5; +type Operation = "act" | "observe" | "extract"; + +function createHarness() { + const runtime = createStagehandRuntime(); + runtime.state.setState( + { + status: "initialized", + initParams: StagehandInitParamsSchema.parse({ + protocolVersion: STAGEHAND_PROTOCOL_VERSION, + clientInfo: { name: "stagehand-controller-test", version: "1.0.0" }, + model: { source: "client" }, + logLevel: "off", + }), + }, + true, + ); + vi.spyOn(runtime, "resolvePage").mockReturnValue({} as never); + vi.spyOn(runtime, "resolveUnderstudyPage").mockReturnValue({} as never); + return { + controller: createStagehandController(runtime), + context: { logger: runtime.logger }, + }; +} + +function callOperation( + controller: ReturnType, + context: HandlerContext, + operation: Operation, + timeout?: number, +) { + const options = timeout === undefined ? undefined : { timeout }; + switch (operation) { + case "act": + return controller.act( + { + pageId: "page-1", + instruction: "Click the button", + ...(options ? { options } : {}), + }, + context, + ); + case "observe": + return controller.observe({ pageId: "page-1", ...(options ? { options } : {}) }, context); + case "extract": + return controller.extract( + { + pageId: "page-1", + instruction: "Extract the page", + schema: { type: "object" }, + ...(options ? { options } : {}), + }, + context, + ); + } +} + +function mockOperation(operation: Operation, result: Promise) { + switch (operation) { + case "act": + return vi.spyOn(actService, "act").mockReturnValue(result as never); + case "observe": + return vi.spyOn(observeService, "observe").mockReturnValue(result as never); + case "extract": + return vi.spyOn(extractService, "extract").mockReturnValue(result as never); + } +} + +describe("Stagehand controller timeouts", () => { + afterEach(() => { + vi.restoreAllMocks(); + }); + + it.each(["act", "observe", "extract"] as const)( + "applies the full-operation timeout to %s()", + async (operation) => { + const { controller, context } = createHarness(); + const pending = new Promise(() => {}); + mockOperation(operation, pending); + + await expect(callOperation(controller, context, operation, TIMEOUT_MS)).rejects.toThrow( + `${operation}() timed out after ${TIMEOUT_MS}ms`, + ); + }, + ); + + it.each(["act", "observe", "extract"] as const)( + "passes through a successful %s() when timeout is omitted", + async (operation) => { + const { controller, context } = createHarness(); + const result = { + data: { operation }, + metadata: { usage: zeroStagehandResultUsage() }, + }; + mockOperation(operation, Promise.resolve(result)); + + await expect(callOperation(controller, context, operation)).resolves.toBe(result); + }, + ); +}); diff --git a/packages/server/timeoutConfig.ts b/packages/server/timeoutConfig.ts index 0f43b4449..91f30db4e 100644 --- a/packages/server/timeoutConfig.ts +++ b/packages/server/timeoutConfig.ts @@ -1,5 +1,10 @@ import { TimeoutError } from "./errors.js"; +/** + * Enforce a caller-facing deadline. The supplied promise has no cancellation contract, so timing + * out does not abort underlying service or model work; callers that need cancellation must pass an + * abort signal through that operation's own API. + */ export async function withTimeout( promise: Promise, timeout: number | null | undefined, diff --git a/scripts/test-integration.ts b/scripts/test-integration.ts index 12c8f3a89..892db549f 100644 --- a/scripts/test-integration.ts +++ b/scripts/test-integration.ts @@ -49,7 +49,7 @@ export const integrationTestGroups = { "local/page-navigation": ["page-addInitScript", "page-extra-http-headers", "page-goto-response"], "local/page-interactions": ["click-count", "page-drag-and-drop", "page-hover", "page-scroll"], "local/snapshots-ai": ["observe-element-id-format"], - "local/waits-timeouts": ["wait-for-selector", "wait-for-timeout"], + "local/waits-timeouts": ["timeouts", "wait-for-selector", "wait-for-timeout"], } as const; const repoRoot = path.resolve(path.dirname(fileURLToPath(import.meta.url)), "..");