diff --git a/packages/protocol/schemas.ts b/packages/protocol/schemas.ts index 6a8e6627d..0a2f7e74f 100644 --- a/packages/protocol/schemas.ts +++ b/packages/protocol/schemas.ts @@ -1735,35 +1735,21 @@ export const ContextActivePageResultSchema = PageRefSchema.nullable().meta({ id: "ContextActivePageResult", }); -export const ContextGetDomainPolicyResultSchema = z - .strictObject({ - policy: DomainPolicySchema.nullable(), - }) - .meta({ id: "ContextGetDomainPolicyResult" }); +export const ContextGetDomainPolicyResultSchema = DomainPolicySchema.nullable().meta({ + id: "ContextGetDomainPolicyResult", +}); export const ContextCookiesResultSchema = z - .strictObject({ - cookies: z.array(CookieSchema), - }) + .array(CookieSchema) .meta({ id: "ContextCookiesResult" }); export const ContextClipboardReadTextResultSchema = z - .strictObject({ - text: z.string(), - }) + .string() .meta({ id: "ContextClipboardReadTextResult" }); -export const PageUrlResultSchema = z - .strictObject({ - url: z.string(), - }) - .meta({ id: "PageUrlResult" }); +export const PageUrlResultSchema = z.string().meta({ id: "PageUrlResult" }); -export const PageTitleResultSchema = z - .strictObject({ - title: z.string(), - }) - .meta({ id: "PageTitleResult" }); +export const PageTitleResultSchema = z.string().meta({ id: "PageTitleResult" }); export const PageCloseResultSchema = z .strictObject({ @@ -1816,46 +1802,24 @@ export const LocatorHoverResultSchema = z .meta({ id: "LocatorHoverResult" }); export const LocatorCountResultSchema = z - .strictObject({ - count: z.number().int().nonnegative(), - }) + .number() + .int() + .nonnegative() .meta({ id: "LocatorCountResult" }); -export const LocatorIsCheckedResultSchema = z - .strictObject({ - checked: z.boolean(), - }) - .meta({ id: "LocatorIsCheckedResult" }); +export const LocatorIsCheckedResultSchema = z.boolean().meta({ id: "LocatorIsCheckedResult" }); -export const LocatorInputValueResultSchema = z - .strictObject({ - value: z.string(), - }) - .meta({ id: "LocatorInputValueResult" }); +export const LocatorInputValueResultSchema = z.string().meta({ id: "LocatorInputValueResult" }); -export const LocatorIsVisibleResultSchema = z - .strictObject({ - visible: z.boolean(), - }) - .meta({ id: "LocatorIsVisibleResult" }); +export const LocatorIsVisibleResultSchema = z.boolean().meta({ id: "LocatorIsVisibleResult" }); -export const LocatorInnerTextResultSchema = z - .strictObject({ - text: z.string(), - }) - .meta({ id: "LocatorInnerTextResult" }); +export const LocatorInnerTextResultSchema = z.string().meta({ id: "LocatorInnerTextResult" }); -export const LocatorInnerHtmlResultSchema = z - .strictObject({ - html: z.string(), - }) - .meta({ id: "LocatorInnerHtmlResult" }); +export const LocatorInnerHtmlResultSchema = z.string().meta({ id: "LocatorInnerHtmlResult" }); -export const LocatorTextContentResultSchema = z - .strictObject({ - textContent: z.string(), - }) - .meta({ id: "LocatorTextContentResult" }); +export const LocatorTextContentResultSchema = z.string().meta({ + id: "LocatorTextContentResult", +}); export const LocatorScrollToResultSchema = z .strictObject({ @@ -1889,9 +1853,7 @@ export const LocatorTypeResultSchema = z .meta({ id: "LocatorTypeResult" }); export const LocatorSelectOptionResultSchema = z - .strictObject({ - values: z.array(z.string()), - }) + .array(z.string()) .meta({ id: "LocatorSelectOptionResult" }); export const StagehandLogLevelSchema = z diff --git a/packages/protocol/stagehand.v4.json b/packages/protocol/stagehand.v4.json index 4bf25c266..daa5cbde8 100644 --- a/packages/protocol/stagehand.v4.json +++ b/packages/protocol/stagehand.v4.json @@ -2886,21 +2886,14 @@ "additionalProperties": false }, "ContextGetDomainPolicyResult": { - "type": "object", - "properties": { - "policy": { - "anyOf": [ - { - "$ref": "#/$defs/DomainPolicy" - }, - { - "type": "null" - } - ] + "anyOf": [ + { + "$ref": "#/$defs/DomainPolicy" + }, + { + "type": "null" } - }, - "required": ["policy"], - "additionalProperties": false + ] }, "DomainPolicy": { "type": "object", @@ -2957,17 +2950,10 @@ "additionalProperties": false }, "ContextCookiesResult": { - "type": "object", - "properties": { - "cookies": { - "type": "array", - "items": { - "$ref": "#/$defs/Cookie" - } - } - }, - "required": ["cookies"], - "additionalProperties": false + "type": "array", + "items": { + "$ref": "#/$defs/Cookie" + } }, "Cookie": { "type": "object", @@ -3116,14 +3102,7 @@ "additionalProperties": false }, "ContextClipboardReadTextResult": { - "type": "object", - "properties": { - "text": { - "type": "string" - } - }, - "required": ["text"], - "additionalProperties": false + "type": "string" }, "ContextClipboardWriteTextParams": { "type": "object", @@ -3197,24 +3176,10 @@ "additionalProperties": false }, "PageUrlResult": { - "type": "object", - "properties": { - "url": { - "type": "string" - } - }, - "required": ["url"], - "additionalProperties": false + "type": "string" }, "PageTitleResult": { - "type": "object", - "properties": { - "title": { - "type": "string" - } - }, - "required": ["title"], - "additionalProperties": false + "type": "string" }, "PageCloseResult": { "type": "object", @@ -3982,76 +3947,27 @@ "additionalProperties": false }, "LocatorCountResult": { - "type": "object", - "properties": { - "count": { - "type": "integer", - "minimum": 0, - "maximum": 9007199254740991 - } - }, - "required": ["count"], - "additionalProperties": false + "type": "integer", + "minimum": 0, + "maximum": 9007199254740991 }, "LocatorIsCheckedResult": { - "type": "object", - "properties": { - "checked": { - "type": "boolean" - } - }, - "required": ["checked"], - "additionalProperties": false + "type": "boolean" }, "LocatorInputValueResult": { - "type": "object", - "properties": { - "value": { - "type": "string" - } - }, - "required": ["value"], - "additionalProperties": false + "type": "string" }, "LocatorIsVisibleResult": { - "type": "object", - "properties": { - "visible": { - "type": "boolean" - } - }, - "required": ["visible"], - "additionalProperties": false + "type": "boolean" }, "LocatorInnerTextResult": { - "type": "object", - "properties": { - "text": { - "type": "string" - } - }, - "required": ["text"], - "additionalProperties": false + "type": "string" }, "LocatorInnerHtmlResult": { - "type": "object", - "properties": { - "html": { - "type": "string" - } - }, - "required": ["html"], - "additionalProperties": false + "type": "string" }, "LocatorTextContentResult": { - "type": "object", - "properties": { - "text_content": { - "type": "string" - } - }, - "required": ["text_content"], - "additionalProperties": false + "type": "string" }, "LocatorScrollToParams": { "type": "object", @@ -4305,17 +4221,10 @@ "additionalProperties": false }, "LocatorSelectOptionResult": { - "type": "object", - "properties": { - "values": { - "type": "array", - "items": { - "type": "string" - } - } - }, - "required": ["values"], - "additionalProperties": false + "type": "array", + "items": { + "type": "string" + } }, "StagehandLog": { "type": "object", diff --git a/packages/protocol/tests/browser-runtime/rpc-client-smoke.test.ts b/packages/protocol/tests/browser-runtime/rpc-client-smoke.test.ts index b2871e5b3..0f41c21b0 100644 --- a/packages/protocol/tests/browser-runtime/rpc-client-smoke.test.ts +++ b/packages/protocol/tests/browser-runtime/rpc-client-smoke.test.ts @@ -154,7 +154,7 @@ describe("Stagehand service worker RPC client smoke", () => { pageId: page.pageId, selector: "#blank-page-button", }), - ).resolves.toStrictEqual({ textContent: "Clicked" }); + ).resolves.toBe("Clicked"); } finally { await activeRpcClient.send(StagehandMethods.pageClose, { pageId: page.pageId }); } @@ -174,7 +174,7 @@ describe("Stagehand service worker RPC client smoke", () => { pageId: page.pageId, selector: "#locator-message", }), - ).resolves.toStrictEqual({ textContent: "locator text" }); + ).resolves.toBe("locator text"); } finally { await activeRpcClient.send(StagehandMethods.pageClose, { pageId: page.pageId }); } @@ -201,7 +201,7 @@ describe("Stagehand service worker RPC client smoke", () => { pageId: page.pageId, selector: "#data-button", }), - ).resolves.toStrictEqual({ textContent: "Clicked" }); + ).resolves.toBe("Clicked"); await activeRpcClient.send(StagehandMethods.pageEvaluate, { pageId: page.pageId, @@ -217,7 +217,7 @@ describe("Stagehand service worker RPC client smoke", () => { pageId: page.pageId, selector: "#data-shadow-text", }), - ).resolves.toStrictEqual({ textContent: "Open shadow" }); + ).resolves.toBe("Open shadow"); await expect( activeRpcClient.send(StagehandMethods.pageEvaluate, { pageId: page.pageId, @@ -243,7 +243,7 @@ describe("Stagehand service worker RPC client smoke", () => { pageId: page.pageId, selector: "#locator-message", }), - ).resolves.toStrictEqual({ textContent: "locator text" }); + ).resolves.toBe("locator text"); } finally { await activeRpcClient.send(StagehandMethods.pageClose, { pageId: page.pageId }); } @@ -273,12 +273,10 @@ describe("Stagehand service worker RPC client smoke", () => { await expect( activeRpcClient.send(StagehandMethods.pageUrl, { pageId: page.pageId }), - ).resolves.toStrictEqual({ url: activeFixtureServer.url }); + ).resolves.toBe(activeFixtureServer.url); await expect( activeRpcClient.send(StagehandMethods.pageTitle, { pageId: page.pageId }), - ).resolves.toStrictEqual({ - title: "Stagehand Smoke", - }); + ).resolves.toBe("Stagehand Smoke"); }); it("closes a throwaway PageRef in a browser session", async () => { @@ -310,18 +308,14 @@ describe("Stagehand service worker RPC client smoke", () => { pageId: page.pageId, selector: "#locator-message", }), - ).resolves.toStrictEqual({ - visible: true, - }); + ).resolves.toBe(true); await expect( activeRpcClient.send(StagehandMethods.locatorTextContent, { pageId: page.pageId, selector: "#locator-message", }), - ).resolves.toStrictEqual({ - textContent: "locator text", - }); + ).resolves.toBe("locator text"); await expect( activeRpcClient.send(StagehandMethods.locatorFill, { @@ -347,9 +341,7 @@ describe("Stagehand service worker RPC client smoke", () => { pageId: page.pageId, selector: "#locator-output", }), - ).resolves.toStrictEqual({ - textContent: "clicked:user@example.com", - }); + ).resolves.toBe("clicked:user@example.com"); await expect( activeRpcClient.send(StagehandMethods.locatorFill, { @@ -363,26 +355,26 @@ describe("Stagehand service worker RPC client smoke", () => { pageId: page.pageId, selector: "#locator-date", }), - ).resolves.toStrictEqual({ value: "2026-07-21" }); + ).resolves.toBe("2026-07-21"); await expect( activeRpcClient.send(StagehandMethods.locatorIsVisible, { pageId: page.pageId, selector: "#closed-message", }), - ).resolves.toStrictEqual({ visible: true }); + ).resolves.toBe(true); await expect( activeRpcClient.send(StagehandMethods.locatorTextContent, { pageId: page.pageId, selector: "xpath=//div[@id='closed-host']//p[1]", }), - ).resolves.toStrictEqual({ textContent: "closed root text" }); + ).resolves.toBe("closed root text"); await expect( activeRpcClient.send(StagehandMethods.locatorTextContent, { pageId: page.pageId, selector: "#shadow-frame >> #frame-closed-message", }), - ).resolves.toStrictEqual({ textContent: "closed root iframe text" }); + ).resolves.toBe("closed root iframe text"); await activeRpcClient.send(StagehandMethods.locatorFill, { pageId: page.pageId, @@ -398,21 +390,21 @@ describe("Stagehand service worker RPC client smoke", () => { pageId: page.pageId, selector: "#closed-output", }), - ).resolves.toStrictEqual({ textContent: "clicked:inside closed root" }); + ).resolves.toBe("clicked:inside closed root"); await expect( activeRpcClient.send(StagehandMethods.locatorCount, { pageId: page.pageId, selector: ".mixed-shadow", }), - ).resolves.toStrictEqual({ count: 3 }); + ).resolves.toBe(3); await expect( activeRpcClient.send(StagehandMethods.locatorTextContent, { pageId: page.pageId, selector: ".mixed-shadow", nth: 1, }), - ).resolves.toStrictEqual({ textContent: "closed" }); + ).resolves.toBe("closed"); await activeRpcClient.send(StagehandMethods.pageHover, { pageId: page.pageId, diff --git a/packages/protocol/tests/protocol/context-command-schemas.test.ts b/packages/protocol/tests/protocol/context-command-schemas.test.ts index 00baa637f..56d2af3eb 100644 --- a/packages/protocol/tests/protocol/context-command-schemas.test.ts +++ b/packages/protocol/tests/protocol/context-command-schemas.test.ts @@ -87,18 +87,21 @@ describe("context lifecycle and configuration command schemas", () => { expect(ContextSetDomainPolicyParamsSchema.parse({ policy: null })).toStrictEqual({ policy: null, }); - expect(ContextGetDomainPolicyResultSchema.parse({ policy: null })).toStrictEqual({ - policy: null, - }); + expect(ContextGetDomainPolicyResultSchema.parse(null)).toBeNull(); expect( ContextGetDomainPolicyResultSchema.parse({ - policy: { blockedDomains: ["ads.example.com"] }, + blockedDomains: ["ads.example.com"], }), - ).toStrictEqual({ policy: { blockedDomains: ["ads.example.com"] } }); + ).toStrictEqual({ blockedDomains: ["ads.example.com"] }); expect(() => ContextSetDomainPolicyParamsSchema.parse({})).toThrow(); expect(() => ContextSetDomainPolicyParamsSchema.parse({ policy: undefined })).toThrow(); - expect(() => ContextGetDomainPolicyResultSchema.parse({ policy: null, extra: true })).toThrow(); + expect(() => + ContextGetDomainPolicyResultSchema.parse({ + blockedDomains: ["ads.example.com"], + extra: true, + }), + ).toThrow(); }); it("keeps context mutation and close results strict", () => { @@ -139,15 +142,8 @@ describe("context cookie command schemas", () => { }); it("parses cookie results and rejects invalid browser values", () => { - expect(ContextCookiesResultSchema.parse({ cookies: [cookie] })).toStrictEqual({ - cookies: [cookie], - }); - expect(() => - ContextCookiesResultSchema.parse({ - cookies: [{ ...cookie, sameSite: "Invalid" }], - }), - ).toThrow(); - expect(() => ContextCookiesResultSchema.parse({ cookies: [], extra: true })).toThrow(); + expect(ContextCookiesResultSchema.parse([cookie])).toStrictEqual([cookie]); + expect(() => ContextCookiesResultSchema.parse([{ ...cookie, sameSite: "Invalid" }])).toThrow(); }); it("validates cookies before adding them", () => { @@ -252,12 +248,8 @@ describe("context clipboard command schemas", () => { expect(() => ContextClipboardPasteParamsSchema.parse({ shortcut: "Shift+V" })).toThrow(); }); - it("parses strict clipboard read results", () => { - expect(ContextClipboardReadTextResultSchema.parse({ text: "copied text" })).toStrictEqual({ - text: "copied text", - }); - expect(() => - ContextClipboardReadTextResultSchema.parse({ text: "copied text", extra: true }), - ).toThrow(); + it("parses clipboard read results", () => { + expect(ContextClipboardReadTextResultSchema.parse("copied text")).toBe("copied text"); + expect(() => ContextClipboardReadTextResultSchema.parse(1)).toThrow(); }); }); diff --git a/packages/protocol/tests/protocol/object-model-protocol.test.ts b/packages/protocol/tests/protocol/object-model-protocol.test.ts index 5baf0994e..859b4ff0c 100644 --- a/packages/protocol/tests/protocol/object-model-protocol.test.ts +++ b/packages/protocol/tests/protocol/object-model-protocol.test.ts @@ -399,25 +399,15 @@ describe("Stagehand object-model protocol", () => { hovered: true, }); - expect(StagehandMethods.locatorCount.result.parse({ count: 2 })).toStrictEqual({ - count: 2, - }); - expect(() => StagehandMethods.locatorCount.result.parse({ count: -1 })).toThrow(); + expect(StagehandMethods.locatorCount.result.parse(2)).toBe(2); + expect(() => StagehandMethods.locatorCount.result.parse(-1)).toThrow(); - expect(StagehandMethods.locatorIsChecked.result.parse({ checked: true })).toStrictEqual({ - checked: true, - }); - expect( - StagehandMethods.locatorInputValue.result.parse({ value: "user@example.com" }), - ).toStrictEqual({ value: "user@example.com" }); - expect(StagehandMethods.locatorInnerText.result.parse({ text: "Submit" })).toStrictEqual({ - text: "Submit", - }); - expect(StagehandMethods.locatorInnerHtml.result.parse({ html: "Submit" })).toStrictEqual( - { - html: "Submit", - }, + expect(StagehandMethods.locatorIsChecked.result.parse(true)).toBe(true); + expect(StagehandMethods.locatorInputValue.result.parse("user@example.com")).toBe( + "user@example.com", ); + expect(StagehandMethods.locatorInnerText.result.parse("Submit")).toBe("Submit"); + expect(StagehandMethods.locatorInnerHtml.result.parse("Submit")).toBe("Submit"); expect( StagehandMethods.locatorScrollTo.params.parse({ @@ -480,9 +470,7 @@ describe("Stagehand object-model protocol", () => { ...locatorDescriptor(), values: ["a", "b"], }); - expect(StagehandMethods.locatorSelectOption.result.parse({ values: ["a"] })).toStrictEqual({ - values: ["a"], - }); + expect(StagehandMethods.locatorSelectOption.result.parse(["a"])).toStrictEqual(["a"]); }); it("exports a JSON-RPC request schema for generated clients", () => { diff --git a/packages/protocol/tests/protocol/schema-registry.test-d.ts b/packages/protocol/tests/protocol/schema-registry.test-d.ts index 82c8efc86..ef68af824 100644 --- a/packages/protocol/tests/protocol/schema-registry.test-d.ts +++ b/packages/protocol/tests/protocol/schema-registry.test-d.ts @@ -87,9 +87,9 @@ expectTypeOf>().toEq nth?: number; values: string | string[]; }>(); -expectTypeOf>().toEqualTypeOf<{ - values: string[]; -}>(); +expectTypeOf>().toEqualTypeOf< + string[] +>(); expectTypeOf().toEqualTypeOf>(); expectTypeOf>().toEqualTypeOf(); expectTypeOf(StagehandNotifications.log.name).toEqualTypeOf<"stagehand.log">(); diff --git a/packages/protocol/tests/protocol/wire-casing.test.ts b/packages/protocol/tests/protocol/wire-casing.test.ts index cff89bd76..2b1de1025 100644 --- a/packages/protocol/tests/protocol/wire-casing.test.ts +++ b/packages/protocol/tests/protocol/wire-casing.test.ts @@ -248,34 +248,30 @@ describe("JSON-RPC wire casing", () => { expect(wireSchema(clipboard.params).parse(clipboardWireParams)).toStrictEqual(clipboardParams); const cookies = StagehandMethods.contextCookies; - const cookiesResult = { - cookies: [ - { - name: "session", - value: "abc123", - domain: "example.com", - path: "/", - expires: -1, - httpOnly: true, - secure: true, - sameSite: "Lax" as const, - }, - ], - }; - const cookiesWireResult = { - cookies: [ - { - name: "session", - value: "abc123", - domain: "example.com", - path: "/", - expires: -1, - http_only: true, - secure: true, - same_site: "Lax" as const, - }, - ], - }; + const cookiesResult = [ + { + name: "session", + value: "abc123", + domain: "example.com", + path: "/", + expires: -1, + httpOnly: true, + secure: true, + sameSite: "Lax" as const, + }, + ]; + const cookiesWireResult = [ + { + name: "session", + value: "abc123", + domain: "example.com", + path: "/", + expires: -1, + http_only: true, + secure: true, + same_site: "Lax" as const, + }, + ]; expect(encodeWireValue(cookiesResult)).toStrictEqual(cookiesWireResult); expect(wireSchema(cookies.result).parse(cookiesWireResult)).toStrictEqual(cookiesResult); }); diff --git a/packages/sdk-go/browser_clipboard.go b/packages/sdk-go/browser_clipboard.go index 6670c4ad3..c329d6e8d 100644 --- a/packages/sdk-go/browser_clipboard.go +++ b/packages/sdk-go/browser_clipboard.go @@ -25,7 +25,7 @@ func (c *BrowserClipboard) ReadText(ctx context.Context, options *ClipboardOptio if err := c.rpc.call(ctx, "context.clipboard_read_text", params, &result); err != nil { return "", err } - return result.Text, nil + return string(result), nil } // WriteText replaces the current clipboard text. diff --git a/packages/sdk-go/browser_clipboard_test.go b/packages/sdk-go/browser_clipboard_test.go index fdef116b0..04875d7b6 100644 --- a/packages/sdk-go/browser_clipboard_test.go +++ b/packages/sdk-go/browser_clipboard_test.go @@ -10,7 +10,7 @@ func TestBrowserClipboardMapsScopedAndUnscopedCalls(t *testing.T) { shortcut := ContextClipboardPasteParamsShortcutControlOrMetaV rpc := &recordingProtocolClient{responses: map[string]any{ - "context.clipboard_read_text": ContextClipboardReadTextResult{Text: "copied"}, + "context.clipboard_read_text": ContextClipboardReadTextResult("copied"), }} clipboard := &BrowserClipboard{rpc: rpc} page := &Page{rpc: rpc, ref: PageRef{PageID: "page-1"}} diff --git a/packages/sdk-go/browser_context.go b/packages/sdk-go/browser_context.go index 68a0eec8f..eb076fec9 100644 --- a/packages/sdk-go/browser_context.go +++ b/packages/sdk-go/browser_context.go @@ -90,7 +90,7 @@ func (c *BrowserContext) GetDomainPolicy(ctx context.Context) (*DomainPolicy, er if err := c.rpc.call(ctx, "context.get_domain_policy", EmptyParams{}, &result); err != nil { return nil, err } - return result.Policy, nil + return result, nil } // SetDomainPolicy changes the current domain policy. @@ -107,7 +107,7 @@ func (c *BrowserContext) Cookies(ctx context.Context, urls *StringList) ([]Cooki if err := c.rpc.call(ctx, "context.cookies", params, &result); err != nil { return nil, err } - return result.Cookies, nil + return []Cookie(result), nil } // AddCookies adds cookies to the context. diff --git a/packages/sdk-go/browser_context_test.go b/packages/sdk-go/browser_context_test.go index 9b509c652..cd10563e7 100644 --- a/packages/sdk-go/browser_context_test.go +++ b/packages/sdk-go/browser_context_test.go @@ -14,10 +14,10 @@ func TestBrowserContextMapsPagesAndCookies(t *testing.T) { rpc := &recordingProtocolClient{responses: map[string]any{ "context.pages": ContextPagesResult{{PageID: "page-1"}}, "context.new_page": PageRef{PageID: "page-2", URL: &pageURL}, - "context.cookies": ContextCookiesResult{Cookies: []Cookie{{ + "context.cookies": ContextCookiesResult{{ Name: "session", Value: "value", Domain: "example.com", Path: "/", SameSite: CookieSameSiteLax, - }}}, + }}, }} browserContext := &BrowserContext{rpc: rpc} diff --git a/packages/sdk-go/internal/extensionassets/stagehand-extension.zip b/packages/sdk-go/internal/extensionassets/stagehand-extension.zip index d043c3584..230179684 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-go/internal/generator/main.go b/packages/sdk-go/internal/generator/main.go index 91fa26a62..3dfbe8f9e 100644 --- a/packages/sdk-go/internal/generator/main.go +++ b/packages/sdk-go/internal/generator/main.go @@ -26,31 +26,32 @@ const ( ) var customDefinitions = map[string]string{ - "Caching": "Caching", - "CookieFilter": "CookieFilter", - "ContextActivePageResult": "ContextActivePageResult", - "EmptyParams": "EmptyParams", - "LLMGenerateParams": "LLMGenerateParams", - "LLMGenerateResult": "LLMGenerateResult", - "LLMMessageContentBlock": "LLMMessageContentBlock", - "LLMMessageGenerateResult": "LLMMessageGenerateResult", - "LLMStructuredGenerateResult": "LLMStructuredGenerateResult", - "LLMToolResultContentBlock": "LLMToolResultContentBlock", - "LoadState": "LoadState", - "ModelConfig": "ModelConfig", - "ModelName": "ModelName", - "ProxyConfig": "ProxyConfig", - "VariablePrimitive": "VariablePrimitive", - "VariableValue": "VariableValue", - "__schema0": "json.RawMessage", - "__schema1": "json.RawMessage", - "__schema2": "json.RawMessage", - "__schema3": "json.RawMessage", - "__schema4": "json.RawMessage", - "__schema5": "json.RawMessage", - "__schema6": "json.RawMessage", - "__schema7": "json.RawMessage", - "__schema8": "json.RawMessage", + "Caching": "Caching", + "CookieFilter": "CookieFilter", + "ContextActivePageResult": "ContextActivePageResult", + "ContextGetDomainPolicyResult": "ContextGetDomainPolicyResult", + "EmptyParams": "EmptyParams", + "LLMGenerateParams": "LLMGenerateParams", + "LLMGenerateResult": "LLMGenerateResult", + "LLMMessageContentBlock": "LLMMessageContentBlock", + "LLMMessageGenerateResult": "LLMMessageGenerateResult", + "LLMStructuredGenerateResult": "LLMStructuredGenerateResult", + "LLMToolResultContentBlock": "LLMToolResultContentBlock", + "LoadState": "LoadState", + "ModelConfig": "ModelConfig", + "ModelName": "ModelName", + "ProxyConfig": "ProxyConfig", + "VariablePrimitive": "VariablePrimitive", + "VariableValue": "VariableValue", + "__schema0": "json.RawMessage", + "__schema1": "json.RawMessage", + "__schema2": "json.RawMessage", + "__schema3": "json.RawMessage", + "__schema4": "json.RawMessage", + "__schema5": "json.RawMessage", + "__schema6": "json.RawMessage", + "__schema7": "json.RawMessage", + "__schema8": "json.RawMessage", } var customProperties = map[string]string{ diff --git a/packages/sdk-go/locator.go b/packages/sdk-go/locator.go index da22ae65d..0d8d3975f 100644 --- a/packages/sdk-go/locator.go +++ b/packages/sdk-go/locator.go @@ -51,7 +51,7 @@ func (l *PageLocator) Count(ctx context.Context) (int, error) { if err := l.rpc.call(ctx, "locator.count", params, &result); err != nil { return 0, err } - return result.Count, nil + return int(result), nil } // IsChecked reports whether the matching control is checked. @@ -61,7 +61,7 @@ func (l *PageLocator) IsChecked(ctx context.Context) (bool, error) { if err := l.rpc.call(ctx, "locator.is_checked", params, &result); err != nil { return false, err } - return result.Checked, nil + return bool(result), nil } // InputValue returns the matching input's value. @@ -71,7 +71,7 @@ func (l *PageLocator) InputValue(ctx context.Context) (string, error) { if err := l.rpc.call(ctx, "locator.input_value", params, &result); err != nil { return "", err } - return result.Value, nil + return string(result), nil } // IsVisible reports whether the matching element is visible. @@ -81,7 +81,7 @@ func (l *PageLocator) IsVisible(ctx context.Context) (bool, error) { if err := l.rpc.call(ctx, "locator.is_visible", params, &result); err != nil { return false, err } - return result.Visible, nil + return bool(result), nil } // InnerText returns the matching element's rendered text. @@ -91,7 +91,7 @@ func (l *PageLocator) InnerText(ctx context.Context) (string, error) { if err := l.rpc.call(ctx, "locator.inner_text", params, &result); err != nil { return "", err } - return result.Text, nil + return string(result), nil } // InnerHTML returns the matching element's HTML. @@ -101,7 +101,7 @@ func (l *PageLocator) InnerHTML(ctx context.Context) (string, error) { if err := l.rpc.call(ctx, "locator.inner_html", params, &result); err != nil { return "", err } - return result.HTML, nil + return string(result), nil } // TextContent returns the matching element's text content. @@ -111,7 +111,7 @@ func (l *PageLocator) TextContent(ctx context.Context) (string, error) { if err := l.rpc.call(ctx, "locator.text_content", params, &result); err != nil { return "", err } - return result.TextContent, nil + return string(result), nil } // ScrollTo scrolls the matching element to a generated percentage value. @@ -175,7 +175,7 @@ func (l *PageLocator) SelectOption(ctx context.Context, values StringList) ([]st if err := l.rpc.call(ctx, "locator.select_option", params, &result); err != nil { return nil, err } - return result.Values, nil + return []string(result), nil } // First returns a locator restricted to the first match. diff --git a/packages/sdk-go/locator_test.go b/packages/sdk-go/locator_test.go index 69695ec9a..055c6c237 100644 --- a/packages/sdk-go/locator_test.go +++ b/packages/sdk-go/locator_test.go @@ -9,8 +9,8 @@ func TestPageLocatorPropagatesDescriptorAndMapsResults(t *testing.T) { t.Parallel() rpc := &recordingProtocolClient{responses: map[string]any{ - "locator.count": LocatorCountResult{Count: 3}, - "locator.select_option": LocatorSelectOptionResult{Values: []string{"one"}}, + "locator.count": LocatorCountResult(3), + "locator.select_option": LocatorSelectOptionResult{"one"}, }} locator := (&Page{ rpc: rpc, diff --git a/packages/sdk-go/models.gen.go b/packages/sdk-go/models.gen.go index 13343442e..b69714e6a 100644 --- a/packages/sdk-go/models.gen.go +++ b/packages/sdk-go/models.gen.go @@ -316,10 +316,7 @@ const ContextClipboardPasteParamsShortcutControlOrMetaV ContextClipboardPastePar const ContextClipboardPasteParamsShortcutControlV ContextClipboardPasteParamsShortcut = "Control+V" const ContextClipboardPasteParamsShortcutMetaV ContextClipboardPasteParamsShortcut = "Meta+V" -type ContextClipboardReadTextResult struct { - // Text corresponds to the JSON schema field "text". - Text string `json:"text"` -} +type ContextClipboardReadTextResult string type ContextClipboardTarget struct { // PageID corresponds to the JSON schema field "page_id". @@ -344,15 +341,7 @@ type ContextCookiesParams struct { Urls *StringList `json:"urls,omitempty,omitzero"` } -type ContextCookiesResult struct { - // Cookies corresponds to the JSON schema field "cookies". - Cookies []Cookie `json:"cookies"` -} - -type ContextGetDomainPolicyResult struct { - // Policy corresponds to the JSON schema field "policy". - Policy *DomainPolicy `json:"policy"` -} +type ContextCookiesResult []Cookie type ContextNewPageParams struct { // URL corresponds to the JSON schema field "url". @@ -869,10 +858,7 @@ type LocatorClickResult struct { Clicked bool `json:"clicked"` } -type LocatorCountResult struct { - // Count corresponds to the JSON schema field "count". - Count int `json:"count"` -} +type LocatorCountResult int type LocatorDescriptor struct { // Nth corresponds to the JSON schema field "nth". @@ -939,30 +925,15 @@ type LocatorHoverResult struct { Hovered bool `json:"hovered"` } -type LocatorInnerHTMLResult struct { - // HTML corresponds to the JSON schema field "html". - HTML string `json:"html"` -} +type LocatorInnerHTMLResult string -type LocatorInnerTextResult struct { - // Text corresponds to the JSON schema field "text". - Text string `json:"text"` -} +type LocatorInnerTextResult string -type LocatorInputValueResult struct { - // Value corresponds to the JSON schema field "value". - Value string `json:"value"` -} +type LocatorInputValueResult string -type LocatorIsCheckedResult struct { - // Checked corresponds to the JSON schema field "checked". - Checked bool `json:"checked"` -} +type LocatorIsCheckedResult bool -type LocatorIsVisibleResult struct { - // Visible corresponds to the JSON schema field "visible". - Visible bool `json:"visible"` -} +type LocatorIsVisibleResult bool type LocatorScrollToParams struct { // Nth corresponds to the JSON schema field "nth". @@ -997,10 +968,7 @@ type LocatorSelectOptionParams struct { Values StringList `json:"values"` } -type LocatorSelectOptionResult struct { - // Values corresponds to the JSON schema field "values". - Values []string `json:"values"` -} +type LocatorSelectOptionResult []string type LocatorSendClickEventOptions struct { // Bubbles corresponds to the JSON schema field "bubbles". @@ -1035,10 +1003,7 @@ type LocatorSendClickEventResult struct { Clicked bool `json:"clicked"` } -type LocatorTextContentResult struct { - // TextContent corresponds to the JSON schema field "text_content". - TextContent string `json:"text_content"` -} +type LocatorTextContentResult string type LocatorTypeOptions struct { // Delay corresponds to the JSON schema field "delay". @@ -1471,10 +1436,7 @@ type PageSnapshotParams struct { PageID string `json:"page_id"` } -type PageTitleResult struct { - // Title corresponds to the JSON schema field "title". - Title string `json:"title"` -} +type PageTitleResult string type PageTypeOptions struct { // Delay corresponds to the JSON schema field "delay". @@ -1495,10 +1457,7 @@ type PageTypeParams struct { Text string `json:"text"` } -type PageURLResult struct { - // URL corresponds to the JSON schema field "url". - URL string `json:"url"` -} +type PageURLResult string type PageVoidResult struct { // Ok corresponds to the JSON schema field "ok". @@ -1943,7 +1902,7 @@ type generatedModelCatalog struct { // ContextCookiesResult corresponds to the JSON schema field // "ContextCookiesResult". - ContextCookiesResult *ContextCookiesResult `json:"ContextCookiesResult,omitempty,omitzero"` + ContextCookiesResult ContextCookiesResult `json:"ContextCookiesResult,omitempty,omitzero"` // ContextGetDomainPolicyResult corresponds to the JSON schema field // "ContextGetDomainPolicyResult". @@ -2178,7 +2137,7 @@ type generatedModelCatalog struct { // LocatorSelectOptionResult corresponds to the JSON schema field // "LocatorSelectOptionResult". - LocatorSelectOptionResult *LocatorSelectOptionResult `json:"LocatorSelectOptionResult,omitempty,omitzero"` + LocatorSelectOptionResult LocatorSelectOptionResult `json:"LocatorSelectOptionResult,omitempty,omitzero"` // LocatorSendClickEventOptions corresponds to the JSON schema field // "LocatorSendClickEventOptions". diff --git a/packages/sdk-go/page.go b/packages/sdk-go/page.go index aa4878326..c398f2d65 100644 --- a/packages/sdk-go/page.go +++ b/packages/sdk-go/page.go @@ -270,7 +270,7 @@ func (p *Page) URL(ctx context.Context) (string, error) { if err := p.rpc.call(ctx, "page.url", params, &result); err != nil { return "", err } - return result.URL, nil + return string(result), nil } // Title returns the page title. @@ -280,7 +280,7 @@ func (p *Page) Title(ctx context.Context) (string, error) { if err := p.rpc.call(ctx, "page.title", params, &result); err != nil { return "", err } - return result.Title, nil + return string(result), nil } // Close closes the page. diff --git a/packages/sdk-go/scalar_unions.go b/packages/sdk-go/scalar_unions.go index 49efc1ac0..dc08261a0 100644 --- a/packages/sdk-go/scalar_unions.go +++ b/packages/sdk-go/scalar_unions.go @@ -117,6 +117,9 @@ const ( // ContextActivePageResult is either the active page or JSON null. type ContextActivePageResult = *PageRef +// ContextGetDomainPolicyResult is either the current domain policy or JSON null. +type ContextGetDomainPolicyResult = *DomainPolicy + // StringList accepts either a single JSON string or an array of strings and // always marshals as an array. type StringList []string diff --git a/packages/sdk-python/src/stagehand/_generated/models.py b/packages/sdk-python/src/stagehand/_generated/models.py index 42203cbb5..64c53963c 100644 --- a/packages/sdk-python/src/stagehand/_generated/models.py +++ b/packages/sdk-python/src/stagehand/_generated/models.py @@ -348,12 +348,8 @@ class ContextClipboardPasteParams(WireModel): shortcut: Optional[Shortcut] = None -class ContextClipboardReadTextResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - text: StrictStr +class ContextClipboardReadTextResult(RootModel[StrictStr]): + root: StrictStr class ContextClipboardTarget(WireModel): @@ -389,22 +385,6 @@ class ContextCookiesParams(WireModel): urls: Optional[Union[StrictStr, list[StrictStr]]] = None -class ContextCookiesResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - cookies: list[Cookie] - - -class ContextGetDomainPolicyResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - policy: Optional[DomainPolicy] - - class ContextNewPageParams(WireModel): model_config = ConfigDict( extra="forbid", @@ -460,6 +440,10 @@ class Cookie(WireModel): same_site: SameSite +class ContextCookiesResult(RootModel[list[Cookie]]): + root: list[Cookie] + + class CookieParam(WireModel): model_config = ConfigDict( extra="forbid", @@ -542,6 +526,10 @@ class DomainPolicy(WireModel): blocked_domains: Optional[list[StrictStr]] = None +class ContextGetDomainPolicyResult(RootModel[Optional[DomainPolicy]]): + root: Optional[DomainPolicy] + + class EmptyParams(WireModel): model_config = ConfigDict( extra="forbid", @@ -1048,12 +1036,8 @@ class LocatorClickResult(WireModel): clicked: Literal[True] -class LocatorCountResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - count: Annotated[StrictInt, Field(ge=0, le=9007199254740991)] +class LocatorCountResult(RootModel[StrictInt]): + root: Annotated[StrictInt, Field(ge=0, le=9007199254740991)] class LocatorDescriptor(WireModel): @@ -1122,44 +1106,24 @@ class LocatorHoverResult(WireModel): hovered: Literal[True] -class LocatorInnerHtmlResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - html: StrictStr +class LocatorInnerHtmlResult(RootModel[StrictStr]): + root: StrictStr -class LocatorInnerTextResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - text: StrictStr +class LocatorInnerTextResult(RootModel[StrictStr]): + root: StrictStr -class LocatorInputValueResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - value: StrictStr +class LocatorInputValueResult(RootModel[StrictStr]): + root: StrictStr -class LocatorIsCheckedResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - checked: StrictBool +class LocatorIsCheckedResult(RootModel[StrictBool]): + root: StrictBool -class LocatorIsVisibleResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - visible: StrictBool +class LocatorIsVisibleResult(RootModel[StrictBool]): + root: StrictBool class LocatorScrollToParams(WireModel): @@ -1192,12 +1156,8 @@ class LocatorSelectOptionParams(WireModel): values: Union[StrictStr, list[StrictStr]] -class LocatorSelectOptionResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - values: list[StrictStr] +class LocatorSelectOptionResult(RootModel[list[StrictStr]]): + root: list[StrictStr] class LocatorSendClickEventOptions(WireModel): @@ -1230,12 +1190,8 @@ class LocatorSendClickEventResult(WireModel): clicked: Literal[True] -class LocatorTextContentResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - text_content: StrictStr +class LocatorTextContentResult(RootModel[StrictStr]): + root: StrictStr class LocatorTypeOptions(WireModel): @@ -1711,12 +1667,8 @@ class PageSnapshotParams(WireModel): options: Optional[PageSnapshotOptions] = None -class PageTitleResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - title: StrictStr +class PageTitleResult(RootModel[StrictStr]): + root: StrictStr class PageTypeOptions(WireModel): @@ -1738,12 +1690,8 @@ class PageTypeParams(WireModel): options: Optional[PageTypeOptions] = None -class PageUrlResult(WireModel): - model_config = ConfigDict( - extra="forbid", - validate_by_name=True, - ) - url: StrictStr +class PageUrlResult(RootModel[StrictStr]): + root: StrictStr class PageVoidResult(WireModel): diff --git a/packages/sdk-python/src/stagehand/browser_clipboard.py b/packages/sdk-python/src/stagehand/browser_clipboard.py index 90f0d8819..cef2a42d5 100644 --- a/packages/sdk-python/src/stagehand/browser_clipboard.py +++ b/packages/sdk-python/src/stagehand/browser_clipboard.py @@ -21,12 +21,11 @@ def __init__(self, rpc_client: RPCClient) -> None: self._rpc_client = rpc_client async def read_text(self, *, page: Page | None = None) -> str: - result = await self._rpc_client.send( + return await self._rpc_client.send( "context.clipboard_read_text", _clipboard_target(page), ContextClipboardReadTextResult, ) - return result.text async def write_text(self, text: str, *, page: Page | None = None) -> None: params = ContextClipboardWriteTextParams(text=text) diff --git a/packages/sdk-python/src/stagehand/browser_context.py b/packages/sdk-python/src/stagehand/browser_context.py index e6ef7f3a7..4978d0973 100644 --- a/packages/sdk-python/src/stagehand/browser_context.py +++ b/packages/sdk-python/src/stagehand/browser_context.py @@ -50,7 +50,7 @@ async def pages(self) -> list[Page]: EmptyParams(), ContextPagesResult, ) - return [Page(self._rpc_client, page_ref) for page_ref in result.root] + return [Page(self._rpc_client, page_ref) for page_ref in result] async def new_page(self, *, url: str | None = None) -> Page: params = ContextNewPageParams() @@ -65,7 +65,7 @@ async def active_page(self) -> Page | None: EmptyParams(), ContextActivePageResult, ) - return None if result.root is None else Page(self._rpc_client, result.root) + return None if result is None else Page(self._rpc_client, result) async def set_active_page(self, page: Page) -> None: await self._rpc_client.send( @@ -102,12 +102,11 @@ async def set_extra_http_headers(self, headers: Mapping[str, str]) -> None: ) async def get_domain_policy(self) -> DomainPolicy | None: - result = await self._rpc_client.send( + return await self._rpc_client.send( "context.get_domain_policy", EmptyParams(), ContextGetDomainPolicyResult, ) - return result.policy async def set_domain_policy(self, policy: DomainPolicy | None) -> None: await self._rpc_client.send( @@ -120,12 +119,11 @@ async def cookies(self, urls: str | Sequence[str] | None = None) -> list[Cookie] params = ContextCookiesParams() if urls is not None: params.urls = list(urls) if not isinstance(urls, str) else urls - result = await self._rpc_client.send( + return await self._rpc_client.send( "context.cookies", params, ContextCookiesResult, ) - return result.cookies async def add_cookies(self, cookies: Sequence[CookieParam]) -> None: await self._rpc_client.send( diff --git a/packages/sdk-python/src/stagehand/locator.py b/packages/sdk-python/src/stagehand/locator.py index 90938ca66..df8b41d8b 100644 --- a/packages/sdk-python/src/stagehand/locator.py +++ b/packages/sdk-python/src/stagehand/locator.py @@ -106,60 +106,53 @@ async def fill(self, value: str) -> None: ) async def count(self) -> int: - result = await self._rpc_client.send( + return await self._rpc_client.send( "locator.count", self._descriptor, LocatorCountResult, ) - return result.count async def is_checked(self) -> bool: - result = await self._rpc_client.send( + return await self._rpc_client.send( "locator.is_checked", self._descriptor, LocatorIsCheckedResult, ) - return result.checked async def input_value(self) -> str: - result = await self._rpc_client.send( + return await self._rpc_client.send( "locator.input_value", self._descriptor, LocatorInputValueResult, ) - return result.value async def is_visible(self) -> bool: - result = await self._rpc_client.send( + return await self._rpc_client.send( "locator.is_visible", self._descriptor, LocatorIsVisibleResult, ) - return result.visible async def inner_text(self) -> str: - result = await self._rpc_client.send( + return await self._rpc_client.send( "locator.inner_text", self._descriptor, LocatorInnerTextResult, ) - return result.text async def inner_html(self) -> str: - result = await self._rpc_client.send( + return await self._rpc_client.send( "locator.inner_html", self._descriptor, LocatorInnerHtmlResult, ) - return result.html async def text_content(self) -> str: - result = await self._rpc_client.send( + return await self._rpc_client.send( "locator.text_content", self._descriptor, LocatorTextContentResult, ) - return result.text_content async def scroll_to(self, percent: float | str) -> None: await self._rpc_client.send( @@ -241,7 +234,7 @@ async def type(self, text: str, *, delay: float | None = None) -> None: ) async def select_option(self, values: str | Sequence[str]) -> list[str]: - result = await self._rpc_client.send( + return await self._rpc_client.send( "locator.select_option", LocatorSelectOptionParams.model_validate({ **self._descriptor.model_dump(exclude_unset=True), @@ -249,7 +242,6 @@ async def select_option(self, values: str | Sequence[str]) -> list[str]: }), LocatorSelectOptionResult, ) - return result.values def first(self) -> Self: return self.nth(0) diff --git a/packages/sdk-python/src/stagehand/page.py b/packages/sdk-python/src/stagehand/page.py index 4c820becf..642136dde 100644 --- a/packages/sdk-python/src/stagehand/page.py +++ b/packages/sdk-python/src/stagehand/page.py @@ -440,20 +440,18 @@ async def snapshot(self, *, include_iframes: bool | None = None) -> SnapshotResu return await self._rpc_client.send("page.snapshot", params, SnapshotResult) async def url(self) -> str: - result = await self._rpc_client.send( + return await self._rpc_client.send( "page.url", PageIdParams(page_id=self.page_id), PageUrlResult, ) - return result.url async def title(self) -> str: - result = await self._rpc_client.send( + return await self._rpc_client.send( "page.title", PageIdParams(page_id=self.page_id), PageTitleResult, ) - return result.title async def close(self) -> None: await self._rpc_client.send( diff --git a/packages/sdk-python/src/stagehand/rpc_client.py b/packages/sdk-python/src/stagehand/rpc_client.py index 6f658c9a6..ae15ca74b 100644 --- a/packages/sdk-python/src/stagehand/rpc_client.py +++ b/packages/sdk-python/src/stagehand/rpc_client.py @@ -6,9 +6,17 @@ from collections.abc import Awaitable, Callable, Coroutine from contextlib import suppress from importlib.metadata import version -from typing import Annotated, Literal, Protocol, TypeVar, cast - -from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter, ValidationError +from typing import Annotated, Literal, Protocol, TypeVar, cast, overload + +from pydantic import ( + BaseModel, + ConfigDict, + Field, + JsonValue, + RootModel, + TypeAdapter, + ValidationError, +) from ._generated import models @@ -21,6 +29,7 @@ ParamsT = TypeVar("ParamsT", bound=BaseModel) ResultT = TypeVar("ResultT", bound=BaseModel) +RootResultT = TypeVar("RootResultT") SchemaT = TypeVar("SchemaT", bound=BaseModel) _RequestId = Annotated[int, Field(ge=0, le=_MAX_REQUEST_ID, strict=True)] @@ -111,12 +120,28 @@ def __init__(self, transport: _Transport, *, request_timeout_ms: int = 10_000) - self._close_reason: BaseException | None = None self._reader = asyncio.create_task(self._read(), name="stagehand-rpc-reader") + @overload + async def send( + self, + method: str, + params: BaseModel, + result_model: type[RootModel[RootResultT]], + ) -> RootResultT: ... + + @overload async def send( self, method: str, params: BaseModel, result_model: type[ResultT], - ) -> ResultT: + ) -> ResultT: ... + + async def send( + self, + method: str, + params: BaseModel, + result_model: type[BaseModel], + ) -> object: if self._closed: raise RuntimeError("RPC client is closed") from self._close_reason @@ -132,7 +157,7 @@ async def send( response: asyncio.Future[object] = asyncio.get_running_loop().create_future() self._pending[request_id] = ( method, - cast(type[BaseModel], result_model), + result_model, response, ) request = _JSONRPCRequest( @@ -155,7 +180,7 @@ async def send( ) ) result = await response - return cast(ResultT, result) + return result except TimeoutError as error: raise TimeoutError(f"RPC request timed out: {method}") from error finally: @@ -345,10 +370,11 @@ def _receive_response( return try: - result = result_model.model_validate_json( + parsed_result = result_model.model_validate_json( json.dumps(response.result, separators=(",", ":")), strict=True, ) + result = parsed_result.root if isinstance(parsed_result, RootModel) else parsed_result except (TypeError, ValueError, ValidationError) as error: future.set_exception(error) else: diff --git a/packages/sdk-python/tests/_support.py b/packages/sdk-python/tests/_support.py index 353cd3096..fcd87be97 100644 --- a/packages/sdk-python/tests/_support.py +++ b/packages/sdk-python/tests/_support.py @@ -1,12 +1,13 @@ from __future__ import annotations from collections.abc import Awaitable, Callable -from typing import TypeVar +from typing import TypeVar, overload -from pydantic import BaseModel +from pydantic import BaseModel, RootModel ParamsT = TypeVar("ParamsT", bound=BaseModel) ResultT = TypeVar("ResultT", bound=BaseModel) +RootResultT = TypeVar("RootResultT") class RecordingRPCClient: @@ -17,17 +18,34 @@ def __init__(self, responses: dict[str, object] | None = None) -> None: self.notifications: dict[str, tuple[object, object]] = {} self.closed = False + @overload + async def send( + self, + method: str, + params: BaseModel, + result_model: type[RootModel[RootResultT]], + ) -> RootResultT: ... + + @overload async def send( self, method: str, params: BaseModel, result_model: type[ResultT], - ) -> ResultT: + ) -> ResultT: ... + + async def send( + self, + method: str, + params: BaseModel, + result_model: type[BaseModel], + ) -> object: self.calls.append((method, params, result_model)) response = self.responses[method] if isinstance(response, BaseException): raise response - return result_model.model_validate(response, strict=True) + parsed_result = result_model.model_validate(response, strict=True) + return parsed_result.root if isinstance(parsed_result, RootModel) else parsed_result def on_request( self, diff --git a/packages/sdk-python/tests/test_browser_clipboard.py b/packages/sdk-python/tests/test_browser_clipboard.py index d3969a991..5c0d35770 100644 --- a/packages/sdk-python/tests/test_browser_clipboard.py +++ b/packages/sdk-python/tests/test_browser_clipboard.py @@ -5,7 +5,6 @@ import pytest from stagehand._generated.models import ( - ContextClipboardReadTextResult, ContextVoidResult, PageRef, ) @@ -19,7 +18,7 @@ @pytest.mark.asyncio async def test_browser_clipboard_uses_the_optional_page_as_its_wire_target() -> None: recording = RecordingRPCClient({ - "context.clipboard_read_text": ContextClipboardReadTextResult(text="hello"), + "context.clipboard_read_text": "hello", "context.clipboard_write_text": ContextVoidResult(ok=True), }) rpc_client = cast(RPCClient, recording) diff --git a/packages/sdk-python/tests/test_locator.py b/packages/sdk-python/tests/test_locator.py index bada5c7ea..8cf755df3 100644 --- a/packages/sdk-python/tests/test_locator.py +++ b/packages/sdk-python/tests/test_locator.py @@ -9,7 +9,8 @@ LocatorClickResult, LocatorCountResult, LocatorDescriptor, - LocatorSelectOptionResult, + LocatorInputValueResult, + LocatorIsCheckedResult, ) from stagehand.locator import Locator from stagehand.rpc_client import RPCClient @@ -21,8 +22,8 @@ async def test_locator_methods_use_generated_models_and_keep_the_descriptor_internal() -> None: recording = RecordingRPCClient({ "locator.click": LocatorClickResult(clicked=True), - "locator.count": LocatorCountResult(count=2), - "locator.select_option": LocatorSelectOptionResult(values=["one"]), + "locator.count": 2, + "locator.select_option": ["one"], }) locator = Locator( cast(RPCClient, recording), @@ -52,6 +53,27 @@ async def test_locator_methods_use_generated_models_and_keep_the_descriptor_inte ) +@pytest.mark.asyncio +async def test_locator_boolean_and_string_getters_return_scalars() -> None: + recording = RecordingRPCClient({ + "locator.is_checked": True, + "locator.input_value": "selected", + }) + locator = Locator( + cast(RPCClient, recording), + page_id="page-1", + selector="select", + ) + + assert await locator.is_checked() is True + assert await locator.input_value() == "selected" + descriptor = LocatorDescriptor(page_id="page-1", selector="select") + assert recording.calls == [ + ("locator.is_checked", descriptor, LocatorIsCheckedResult), + ("locator.input_value", descriptor, LocatorInputValueResult), + ] + + def test_locator_first_and_nth_validate_the_generated_descriptor() -> None: recording = RecordingRPCClient() locator = Locator(cast(RPCClient, recording), page_id="page-1", selector="button") diff --git a/packages/sdk-python/tests/test_page.py b/packages/sdk-python/tests/test_page.py index 64634cc7a..0be295f9b 100644 --- a/packages/sdk-python/tests/test_page.py +++ b/packages/sdk-python/tests/test_page.py @@ -8,8 +8,9 @@ from stagehand._generated.models import ( PageEvaluateResult, PageGotoParams, + PageIdParams, PageRef, - PageTitleResult, + PageUrlResult, ) from stagehand.page import Page from stagehand.rpc_client import RPCClient @@ -25,7 +26,7 @@ class EvaluationResult(BaseModel): async def test_page_navigation_uses_generated_wire_models_and_updates_the_page_reference() -> None: recording = RecordingRPCClient({ "page.goto": PageRef(page_id="page-2", url="https://example.com"), - "page.title": PageTitleResult(title="Example Domain"), + "page.title": "Example Domain", }) page = Page(cast(RPCClient, recording), PageRef(page_id="page-1")) @@ -49,6 +50,17 @@ async def test_page_navigation_uses_generated_wire_models_and_updates_the_page_r assert result_model is PageRef +@pytest.mark.asyncio +async def test_page_url_returns_a_scalar_string() -> None: + recording = RecordingRPCClient({"page.url": "https://example.com/path"}) + page = Page(cast(RPCClient, recording), PageRef(page_id="page-1")) + + assert await page.url() == "https://example.com/path" + assert recording.calls == [ + ("page.url", PageIdParams(page_id="page-1"), PageUrlResult), + ] + + def test_page_locator_keeps_the_page_identifier_internal() -> None: recording = RecordingRPCClient() page = Page(cast(RPCClient, recording), PageRef(page_id="page-1")) diff --git a/packages/sdk-python/tests/test_rpc_client.py b/packages/sdk-python/tests/test_rpc_client.py index cc905553b..293ed31f6 100644 --- a/packages/sdk-python/tests/test_rpc_client.py +++ b/packages/sdk-python/tests/test_rpc_client.py @@ -87,7 +87,7 @@ async def test_send_strictly_validates_root_model_results() -> None: "id": request["id"], "result": [{"page_id": "page-1"}], }) - assert (await asyncio.wait_for(call, timeout=1)).root == [models.PageRef(page_id="page-1")] + assert await asyncio.wait_for(call, timeout=1) == [models.PageRef(page_id="page-1")] invalid_call = asyncio.create_task( client.send("context.pages", models.EmptyParams(), models.ContextPagesResult) @@ -131,7 +131,7 @@ async def test_send_revalidates_mutated_params_and_strictly_validates_results() await transport.incoming.put({ "jsonrpc": "2.0", "id": request["id"], - "result": {"count": "1"}, + "result": "1", }) with pytest.raises(ValidationError): await call diff --git a/packages/sdk-ts/src/browserClipboard.ts b/packages/sdk-ts/src/browserClipboard.ts index d3f93fecd..865a58755 100644 --- a/packages/sdk-ts/src/browserClipboard.ts +++ b/packages/sdk-ts/src/browserClipboard.ts @@ -15,11 +15,10 @@ export class BrowserClipboard { constructor(readonly rpcClient: RPCClient) {} async readText(options?: ClipboardOptions): Promise { - const { text } = await this.rpcClient.send( + return await this.rpcClient.send( StagehandMethods.contextClipboardReadText, clipboardTarget(options), ); - return text; } async writeText(text: string, options?: ClipboardOptions): Promise { diff --git a/packages/sdk-ts/src/browserContext.ts b/packages/sdk-ts/src/browserContext.ts index db4a54537..b3a0d509e 100644 --- a/packages/sdk-ts/src/browserContext.ts +++ b/packages/sdk-ts/src/browserContext.ts @@ -65,8 +65,7 @@ export class BrowserContext { } async getDomainPolicy(): Promise { - const { policy } = await this.rpcClient.send(StagehandMethods.contextGetDomainPolicy, {}); - return policy; + return await this.rpcClient.send(StagehandMethods.contextGetDomainPolicy, {}); } async setDomainPolicy(policy: DomainPolicy | null): Promise { @@ -75,8 +74,7 @@ export class BrowserContext { async cookies(urls?: string | string[]): Promise { const params = urls === undefined ? {} : { urls }; - const { cookies } = await this.rpcClient.send(StagehandMethods.contextCookies, params); - return cookies; + return await this.rpcClient.send(StagehandMethods.contextCookies, params); } async addCookies(cookies: CookieParam[]): Promise { diff --git a/packages/sdk-ts/src/locator.ts b/packages/sdk-ts/src/locator.ts index 534a967b4..46b5e1a68 100644 --- a/packages/sdk-ts/src/locator.ts +++ b/packages/sdk-ts/src/locator.ts @@ -37,38 +37,31 @@ export class Locator { } async count(): Promise { - const result = await this.rpcClient.send(StagehandMethods.locatorCount, this.descriptor); - return result.count; + return await this.rpcClient.send(StagehandMethods.locatorCount, this.descriptor); } async isChecked(): Promise { - const result = await this.rpcClient.send(StagehandMethods.locatorIsChecked, this.descriptor); - return result.checked; + return await this.rpcClient.send(StagehandMethods.locatorIsChecked, this.descriptor); } async inputValue(): Promise { - const result = await this.rpcClient.send(StagehandMethods.locatorInputValue, this.descriptor); - return result.value; + return await this.rpcClient.send(StagehandMethods.locatorInputValue, this.descriptor); } async isVisible(): Promise { - const result = await this.rpcClient.send(StagehandMethods.locatorIsVisible, this.descriptor); - return result.visible; + return await this.rpcClient.send(StagehandMethods.locatorIsVisible, this.descriptor); } async innerText(): Promise { - const result = await this.rpcClient.send(StagehandMethods.locatorInnerText, this.descriptor); - return result.text; + return await this.rpcClient.send(StagehandMethods.locatorInnerText, this.descriptor); } async innerHtml(): Promise { - const result = await this.rpcClient.send(StagehandMethods.locatorInnerHtml, this.descriptor); - return result.html; + return await this.rpcClient.send(StagehandMethods.locatorInnerHtml, this.descriptor); } async textContent(): Promise { - const result = await this.rpcClient.send(StagehandMethods.locatorTextContent, this.descriptor); - return result.textContent; + return await this.rpcClient.send(StagehandMethods.locatorTextContent, this.descriptor); } async scrollTo(percent: LocatorScrollToParams["percent"]): Promise { @@ -105,11 +98,10 @@ export class Locator { } async selectOption(values: LocatorSelectOptionParams["values"]): Promise { - const result = await this.rpcClient.send(StagehandMethods.locatorSelectOption, { + return await this.rpcClient.send(StagehandMethods.locatorSelectOption, { ...this.descriptor, values, }); - return result.values; } first(): Locator { diff --git a/packages/sdk-ts/src/page.ts b/packages/sdk-ts/src/page.ts index 3cfd3e79c..91848bb8d 100644 --- a/packages/sdk-ts/src/page.ts +++ b/packages/sdk-ts/src/page.ts @@ -243,17 +243,15 @@ export class Page { } async url(): Promise { - const result = await this.rpcClient.send(StagehandMethods.pageUrl, { + return await this.rpcClient.send(StagehandMethods.pageUrl, { pageId: this.pageId, }); - return result.url; } async title(): Promise { - const result = await this.rpcClient.send(StagehandMethods.pageTitle, { + return await this.rpcClient.send(StagehandMethods.pageTitle, { pageId: this.pageId, }); - return result.title; } async close(): Promise { diff --git a/packages/sdk-ts/tests/object-wrapper.test.ts b/packages/sdk-ts/tests/object-wrapper.test.ts index ed25f9b5b..125365e31 100644 --- a/packages/sdk-ts/tests/object-wrapper.test.ts +++ b/packages/sdk-ts/tests/object-wrapper.test.ts @@ -225,12 +225,10 @@ describe("Stagehand TS object wrapper", () => { const client = new FakeProtocolClient(); client.queueResponse(StagehandMethods.contextSetExtraHTTPHeaders, { ok: true }); client.queueResponse(StagehandMethods.contextGetDomainPolicy, { - policy: { - allowedDomains: ["example.com"], - blockedDomains: ["blocked.example.com"], - }, + allowedDomains: ["example.com"], + blockedDomains: ["blocked.example.com"], }); - client.queueResponse(StagehandMethods.contextGetDomainPolicy, { policy: null }); + client.queueResponse(StagehandMethods.contextGetDomainPolicy, null); client.queueResponse(StagehandMethods.contextSetDomainPolicy, { ok: true }); client.queueResponse(StagehandMethods.contextSetDomainPolicy, { ok: true }); const stagehand = createStagehandWithClientForTest(client); @@ -274,8 +272,8 @@ describe("Stagehand TS object wrapper", () => { secure: true, sameSite: "Lax" as const, }; - client.queueResponse(StagehandMethods.contextCookies, { cookies: [cookie] }); - client.queueResponse(StagehandMethods.contextCookies, { cookies: [] }); + client.queueResponse(StagehandMethods.contextCookies, [cookie]); + client.queueResponse(StagehandMethods.contextCookies, []); client.queueResponse(StagehandMethods.contextAddCookies, { ok: true }); client.queueResponse(StagehandMethods.contextClearCookies, { ok: true }); client.queueResponse(StagehandMethods.contextClearCookies, { ok: true }); @@ -321,7 +319,7 @@ describe("Stagehand TS object wrapper", () => { it("lazily exposes a clipboard facade and routes all clipboard operations", async () => { const client = new FakeProtocolClient(); - client.queueResponse(StagehandMethods.contextClipboardReadText, { text: "clipboard text" }); + client.queueResponse(StagehandMethods.contextClipboardReadText, "clipboard text"); client.queueResponse(StagehandMethods.contextClipboardWriteText, { ok: true }); client.queueResponse(StagehandMethods.contextClipboardClear, { ok: true }); client.queueResponse(StagehandMethods.contextClipboardPaste, { ok: true }); @@ -663,9 +661,9 @@ describe("Stagehand TS object wrapper", () => { ]); }); - it("routes page.url and unwraps the result", async () => { + it("routes page.url", async () => { const client = new FakeProtocolClient(); - client.queueResponse(StagehandMethods.pageUrl, { url: "https://example.com" }); + client.queueResponse(StagehandMethods.pageUrl, "https://example.com"); const page = new Page(client, { pageId: "page-1" }); await expect(page.url()).resolves.toBe("https://example.com"); @@ -674,9 +672,9 @@ describe("Stagehand TS object wrapper", () => { ]); }); - it("routes page.title and unwraps the result", async () => { + it("routes page.title", async () => { const client = new FakeProtocolClient(); - client.queueResponse(StagehandMethods.pageTitle, { title: "Example" }); + client.queueResponse(StagehandMethods.pageTitle, "Example"); const page = new Page(client, { pageId: "page-1" }); await expect(page.title()).resolves.toBe("Example"); @@ -943,9 +941,9 @@ describe("Stagehand TS object wrapper", () => { ]); }); - it("routes locator.isVisible and unwraps the result", async () => { + it("routes locator.isVisible", async () => { const client = new FakeProtocolClient(); - client.queueResponse(StagehandMethods.locatorIsVisible, { visible: true }); + client.queueResponse(StagehandMethods.locatorIsVisible, true); const page = new Page(client, { pageId: "page-1" }); await expect(page.locator("#message").isVisible()).resolves.toBe(true); @@ -957,9 +955,9 @@ describe("Stagehand TS object wrapper", () => { ]); }); - it("routes locator.textContent and unwraps the result", async () => { + it("routes locator.textContent", async () => { const client = new FakeProtocolClient(); - client.queueResponse(StagehandMethods.locatorTextContent, { textContent: "hello" }); + client.queueResponse(StagehandMethods.locatorTextContent, "hello"); const page = new Page(client, { pageId: "page-1" }); await expect(page.locator("#message").textContent()).resolves.toBe("hello"); @@ -971,13 +969,13 @@ describe("Stagehand TS object wrapper", () => { ]); }); - it("routes read locator methods and unwraps their results", async () => { + it("routes read locator methods", async () => { const client = new FakeProtocolClient(); - client.queueResponse(StagehandMethods.locatorCount, { count: 2 }); - client.queueResponse(StagehandMethods.locatorIsChecked, { checked: true }); - client.queueResponse(StagehandMethods.locatorInputValue, { value: "user@example.com" }); - client.queueResponse(StagehandMethods.locatorInnerText, { text: "visible text" }); - client.queueResponse(StagehandMethods.locatorInnerHtml, { html: "visible text" }); + client.queueResponse(StagehandMethods.locatorCount, 2); + client.queueResponse(StagehandMethods.locatorIsChecked, true); + client.queueResponse(StagehandMethods.locatorInputValue, "user@example.com"); + client.queueResponse(StagehandMethods.locatorInnerText, "visible text"); + client.queueResponse(StagehandMethods.locatorInnerHtml, "visible text"); client.queueResponse(StagehandMethods.locatorCentroid, { x: 12, y: 34 }); const page = new Page(client, { pageId: "page-1" }); const locator = page.locator("#field"); @@ -1006,7 +1004,7 @@ describe("Stagehand TS object wrapper", () => { client.queueResponse(StagehandMethods.locatorHighlight, { highlighted: true }); client.queueResponse(StagehandMethods.locatorSendClickEvent, { clicked: true }); client.queueResponse(StagehandMethods.locatorType, { typed: true }); - client.queueResponse(StagehandMethods.locatorSelectOption, { values: ["pro"] }); + client.queueResponse(StagehandMethods.locatorSelectOption, ["pro"]); const page = new Page(client, { pageId: "page-1" }); const locator = page.locator("#field"); diff --git a/packages/sdk-ts/tests/rpcClient.test.ts b/packages/sdk-ts/tests/rpcClient.test.ts index af335d290..24ae75cad 100644 --- a/packages/sdk-ts/tests/rpcClient.test.ts +++ b/packages/sdk-ts/tests/rpcClient.test.ts @@ -114,20 +114,18 @@ describe("RPCClient", () => { }); it("accepts context methods without SDK wrapper methods", async () => { - const cdp = new FakeCDPTransport({ - cookies: [ - { - name: "session", - value: "abc123", - domain: "example.com", - path: "/", - expires: -1, - http_only: true, - secure: true, - same_site: "Lax", - }, - ], - }); + const cdp = new FakeCDPTransport([ + { + name: "session", + value: "abc123", + domain: "example.com", + path: "/", + expires: -1, + http_only: true, + secure: true, + same_site: "Lax", + }, + ]); const client = new RPCClient(cdp, 1_000); const request = client.send(StagehandMethods.contextCookies, { @@ -135,8 +133,8 @@ describe("RPCClient", () => { }); expectTypeOf(request).toEqualTypeOf< - Promise<{ - cookies: Array<{ + Promise< + Array<{ name: string; value: string; domain: string; @@ -145,23 +143,21 @@ describe("RPCClient", () => { httpOnly: boolean; secure: boolean; sameSite: "Strict" | "Lax" | "None"; - }>; - }> + }> + > >(); - await expect(request).resolves.toStrictEqual({ - cookies: [ - { - name: "session", - value: "abc123", - domain: "example.com", - path: "/", - expires: -1, - httpOnly: true, - secure: true, - sameSite: "Lax", - }, - ], - }); + await expect(request).resolves.toStrictEqual([ + { + name: "session", + value: "abc123", + domain: "example.com", + path: "/", + expires: -1, + httpOnly: true, + secure: true, + sameSite: "Lax", + }, + ]); expect(cdp.sent).toContainEqual({ jsonrpc: "2.0", id: 1, diff --git a/packages/server/runtime.ts b/packages/server/runtime.ts index ac8a6e0cd..e35ca21e9 100644 --- a/packages/server/runtime.ts +++ b/packages/server/runtime.ts @@ -368,7 +368,7 @@ export class StagehandRuntime { } contextGetDomainPolicy(): ContextGetDomainPolicyResult { - return { policy: this.requireBrowserSession().getDomainPolicy() }; + return this.requireBrowserSession().getDomainPolicy(); } async contextSetDomainPolicy(params: ContextSetDomainPolicyParams): Promise { @@ -377,7 +377,7 @@ export class StagehandRuntime { } async contextCookies(params: ContextCookiesParams): Promise { - return { cookies: await this.requireBrowserSession().cookies(params.urls) }; + return await this.requireBrowserSession().cookies(params.urls); } async contextAddCookies(params: ContextAddCookiesParams): Promise { @@ -394,7 +394,7 @@ export class StagehandRuntime { params: ContextClipboardReadTextParams, ): Promise { const clipboard = this.requireBrowserSession().clipboard; - return { text: await clipboard.readText(this.clipboardOptions(params.pageId)) }; + return await clipboard.readText(this.clipboardOptions(params.pageId)); } async contextClipboardWriteText( @@ -574,15 +574,11 @@ export class StagehandRuntime { } pageUrl(params: PageIdParams): PageUrlResult { - return { - url: this.resolvePage(params.pageId).url(), - }; + return this.resolvePage(params.pageId).url(); } async pageTitle(params: PageIdParams): Promise { - return { - title: await this.resolvePage(params.pageId).title(), - }; + return await this.resolvePage(params.pageId).title(); } async pageClose(params: PageIdParams): Promise { @@ -608,45 +604,31 @@ export class StagehandRuntime { } async locatorCount(params: LocatorDescriptor): Promise { - return { - count: await this.resolveLocator(params).count(), - }; + return await this.resolveLocator(params).count(); } async locatorIsChecked(params: LocatorDescriptor): Promise { - return { - checked: await this.resolveLocator(params).isChecked(), - }; + return await this.resolveLocator(params).isChecked(); } async locatorInputValue(params: LocatorDescriptor): Promise { - return { - value: await this.resolveLocator(params).inputValue(), - }; + return await this.resolveLocator(params).inputValue(); } async locatorIsVisible(params: LocatorDescriptor): Promise { - return { - visible: await this.resolveLocator(params).isVisible(), - }; + return await this.resolveLocator(params).isVisible(); } async locatorInnerText(params: LocatorDescriptor): Promise { - return { - text: await this.resolveLocator(params).innerText(), - }; + return await this.resolveLocator(params).innerText(); } async locatorInnerHtml(params: LocatorDescriptor): Promise { - return { - html: await this.resolveLocator(params).innerHtml(), - }; + return await this.resolveLocator(params).innerHtml(); } async locatorTextContent(params: LocatorDescriptor): Promise { - return { - textContent: await this.resolveLocator(params).textContent(), - }; + return await this.resolveLocator(params).textContent(); } async locatorScrollTo(params: LocatorScrollToParams): Promise { @@ -676,9 +658,7 @@ export class StagehandRuntime { } async locatorSelectOption(params: LocatorSelectOptionParams): Promise { - return { - values: await this.resolveLocator(params).selectOption(params.values), - }; + return await this.resolveLocator(params).selectOption(params.values); } async close(): Promise { diff --git a/packages/server/tests/stagehand-clients.test.ts b/packages/server/tests/stagehand-clients.test.ts index 5783dc7d3..f20f46eb8 100644 --- a/packages/server/tests/stagehand-clients.test.ts +++ b/packages/server/tests/stagehand-clients.test.ts @@ -424,7 +424,7 @@ class FakeUnderstudyRuntimeLocator implements UnderstudyRuntimeLocator { innerHtml?: string; count?: number; centroid?: LocatorCentroidResult; - selectedValues?: LocatorSelectOptionResult["values"]; + selectedValues?: LocatorSelectOptionResult; } = {}, ) {} @@ -1197,10 +1197,8 @@ describe("Stagehand worker clients", () => { jsonrpc: "2.0", id: 15, result: { - policy: { - allowed_domains: ["example.test"], - blocked_domains: ["blocked.example.test"], - }, + allowed_domains: ["example.test"], + blocked_domains: ["blocked.example.test"], }, }); @@ -1238,7 +1236,7 @@ describe("Stagehand worker clients", () => { ).resolves.toStrictEqual({ jsonrpc: "2.0", id: 17, - result: { policy: null }, + result: null, }); expect(context.setDomainPolicyCalls).toStrictEqual([null]); @@ -1270,20 +1268,18 @@ describe("Stagehand worker clients", () => { ).resolves.toStrictEqual({ jsonrpc: "2.0", id: 18, - result: { - cookies: [ - { - name: "session-id", - value: "abc123", - domain: "example.test", - path: "/", - expires: -1, - http_only: true, - secure: true, - same_site: "Lax", - }, - ], - }, + result: [ + { + name: "session-id", + value: "abc123", + domain: "example.test", + path: "/", + expires: -1, + http_only: true, + secure: true, + same_site: "Lax", + }, + ], }); await expect( handle({ @@ -1360,7 +1356,7 @@ describe("Stagehand worker clients", () => { ).resolves.toStrictEqual({ jsonrpc: "2.0", id: 22, - result: { text: "clipboard text" }, + result: "clipboard text", }); await expect( handle({ @@ -1870,9 +1866,7 @@ describe("Stagehand worker clients", () => { ).resolves.toStrictEqual({ jsonrpc: "2.0", id: 10, - result: { - url: "https://example.test/current", - }, + result: "https://example.test/current", }); }); @@ -1895,9 +1889,7 @@ describe("Stagehand worker clients", () => { ).resolves.toStrictEqual({ jsonrpc: "2.0", id: 11, - result: { - title: "Current Title", - }, + result: "Current Title", }); }); @@ -2030,9 +2022,7 @@ describe("Stagehand worker clients", () => { pageId: "page-a", selector: "section.visible", }), - ).resolves.toStrictEqual({ - visible: true, - }); + ).resolves.toBe(true); expect(page.locatorRefs).toHaveLength(1); expect(page.locatorRefs[0]?.selector).toBe("section.visible"); @@ -2051,9 +2041,7 @@ describe("Stagehand worker clients", () => { pageId: "page-a", selector: "p.message", }), - ).resolves.toStrictEqual({ - textContent: "hello from locator", - }); + ).resolves.toBe("hello from locator"); expect(page.locatorRefs).toHaveLength(1); expect(page.locatorRefs[0]?.selector).toBe("p.message"); @@ -2071,9 +2059,7 @@ describe("Stagehand worker clients", () => { selector: "li.item", nth: 2, }), - ).resolves.toStrictEqual({ - count: 1, - }); + ).resolves.toBe(1); expect(locator.nthCalls).toStrictEqual([2]); }); @@ -2097,17 +2083,11 @@ describe("Stagehand worker clients", () => { selector: "input.email", }; - await expect(runtime.locatorCount(descriptor)).resolves.toStrictEqual({ count: 3 }); - await expect(runtime.locatorIsChecked(descriptor)).resolves.toStrictEqual({ checked: true }); - await expect(runtime.locatorInputValue(descriptor)).resolves.toStrictEqual({ - value: "user@example.com", - }); - await expect(runtime.locatorInnerText(descriptor)).resolves.toStrictEqual({ - text: "visible text", - }); - await expect(runtime.locatorInnerHtml(descriptor)).resolves.toStrictEqual({ - html: "visible text", - }); + await expect(runtime.locatorCount(descriptor)).resolves.toBe(3); + await expect(runtime.locatorIsChecked(descriptor)).resolves.toBe(true); + await expect(runtime.locatorInputValue(descriptor)).resolves.toBe("user@example.com"); + await expect(runtime.locatorInnerText(descriptor)).resolves.toBe("visible text"); + await expect(runtime.locatorInnerHtml(descriptor)).resolves.toBe("visible text"); await expect(runtime.locatorCentroid(descriptor)).resolves.toStrictEqual({ x: 12, y: 34 }); }); @@ -2141,7 +2121,7 @@ describe("Stagehand worker clients", () => { ).resolves.toStrictEqual({ typed: true }); await expect( runtime.locatorSelectOption({ ...descriptor, values: ["a", "b"] }), - ).resolves.toStrictEqual({ values: ["b"] }); + ).resolves.toStrictEqual(["b"]); expect(locator.scrollToCalls).toStrictEqual([50]); expect(locator.highlightCalls).toStrictEqual([ @@ -2259,9 +2239,7 @@ describe("Stagehand worker clients", () => { ).resolves.toStrictEqual({ jsonrpc: "2.0", id: 15, - result: { - visible: true, - }, + result: true, }); }); @@ -2286,9 +2264,7 @@ describe("Stagehand worker clients", () => { ).resolves.toStrictEqual({ jsonrpc: "2.0", id: 16, - result: { - text_content: "hello from locator", - }, + result: "hello from locator", }); }); @@ -2315,9 +2291,7 @@ describe("Stagehand worker clients", () => { ).resolves.toStrictEqual({ jsonrpc: "2.0", id: 17, - result: { - count: 2, - }, + result: 2, }); await expect( @@ -2334,9 +2308,7 @@ describe("Stagehand worker clients", () => { ).resolves.toStrictEqual({ jsonrpc: "2.0", id: 18, - result: { - values: ["pro"], - }, + result: ["pro"], }); expect(locator.nthCalls).toStrictEqual([0]);