diff --git a/packages/core/src/tool/plugin/websearch.ts b/packages/core/src/tool/plugin/websearch.ts index 86198eaf11f..b605d33daab 100644 --- a/packages/core/src/tool/plugin/websearch.ts +++ b/packages/core/src/tool/plugin/websearch.ts @@ -54,36 +54,62 @@ export const Plugin = { if (!Schema.is(WebSearch.ProviderRequiredError)(error)) return Effect.fail(error) return Effect.gen(function* () { const providers = (yield* ctx.websearch.providers()).data - if (providers.length === 0) return yield* new WebSearch.ProviderRequiredError() + const defaultProvider = providers[0] + if (!defaultProvider) return yield* new WebSearch.ProviderRequiredError() const response = yield* forms.ask({ sessionID: context.sessionID, title: "Web Search", metadata: { kind: "websearch.provider" }, fields: [ { - key: "enabled", - description: "Allow OpenCode to search the web?", + key: "choice", + description: "Allow OpenCode to search the web for up-to-date information?", type: "string", required: true, custom: false, options: [ - { value: "yes", label: "Allow web search" }, - { value: "no", label: "Disable web search" }, + { + value: "allow", + label: `Allow web search via ${defaultProvider.name}`, + }, + { + value: "choose", + label: "Choose another provider", + }, + { value: "disable", label: "Disable web search" }, ], }, ], }) if (response.status === "cancelled") return yield* Effect.fail(new Error("Web search cancelled")) - const answer = response.answer.enabled - if (answer === "no") { + if (response.answer.choice === "disable") { yield* kv.set("websearch:provider", false) return yield* new WebSearch.DisabledError() } - if (answer === "yes") { - yield* kv.set("websearch:provider", providers[Math.floor(Math.random() * providers.length)].id) - return yield* ctx.websearch.query(input) - } - return yield* new WebSearch.ProviderRequiredError() + const selection = + response.answer.choice === "choose" + ? yield* forms.ask({ + sessionID: context.sessionID, + title: "Choose a web search provider", + metadata: { kind: "websearch.provider" }, + fields: [ + { + key: "provider", + description: "Choose a provider for web search.", + type: "string", + required: true, + custom: false, + options: providers.map((provider) => ({ value: provider.id, label: provider.name })), + }, + ], + }) + : undefined + if (selection?.status === "cancelled") return yield* Effect.fail(new Error("Web search cancelled")) + const providerID = selection?.answer.provider ?? defaultProvider.id + if (typeof providerID !== "string" || !providers.some((provider) => provider.id === providerID)) + return yield* new WebSearch.ProviderRequiredError() + yield* kv.set("websearch:provider", providerID) + return yield* ctx.websearch.query(input) }) }), ) diff --git a/packages/core/test/tool-websearch.test.ts b/packages/core/test/tool-websearch.test.ts index 9bd6815f22a..1f1ec9cb5d2 100644 --- a/packages/core/test/tool-websearch.test.ts +++ b/packages/core/test/tool-websearch.test.ts @@ -30,6 +30,15 @@ const webSearchToolNode = makeLocationNode({ const sessionID = Session.ID.make("ses_websearch_test") const assertions: Permission.AssertInput[] = [] const queries: WebSearch.Input[] = [] +const formRequests: Form.CreateInput[] = [] +const values = new Map() +const providers = [ + { id: WebSearch.ID.make("exa"), name: "Exa" }, + { id: WebSearch.ID.make("parallel"), name: "Parallel" }, +] +let providerRequired = false +let formResponse: Form.TerminalState = { status: "cancelled" } +const formResponses: Form.TerminalState[] = [] let result = new WebSearch.Response({ providerID: WebSearch.ID.make("exa"), results: [{ url: "https://example.com", title: "Search results", content: "search results", time: {} }], @@ -38,6 +47,11 @@ let result = new WebSearch.Response({ beforeEach(() => { assertions.length = 0 queries.length = 0 + formRequests.length = 0 + values.clear() + providerRequired = false + formResponse = { status: "cancelled" } + formResponses.length = 0 result = new WebSearch.Response({ providerID: WebSearch.ID.make("exa"), results: [{ url: "https://example.com", title: "Search results", content: "search results", time: {} }], @@ -60,11 +74,15 @@ const websearch = Layer.succeed( WebSearch.Service.of({ transform: () => Effect.die("unused"), reload: () => Effect.die("unused"), - providers: () => Effect.succeed([]), + providers: () => Effect.succeed(providers), default: () => Effect.succeed(undefined), query: (input) => - Effect.sync(() => { + Effect.gen(function* () { queries.push(input) + const stored = values.get("websearch:provider") + if (providerRequired && typeof stored !== "string") return yield* new WebSearch.ProviderRequiredError() + if (typeof stored === "string") + return new WebSearch.Response({ providerID: WebSearch.ID.make(stored), results: result.results }) return result }), }), @@ -73,7 +91,11 @@ const form = Layer.succeed( Form.Service, Form.Service.of({ create: () => Effect.die("unused"), - ask: () => Effect.die("unused"), + ask: (input) => + Effect.sync(() => { + formRequests.push(input) + return formResponses.shift() ?? formResponse + }), get: () => Effect.die("unused"), list: () => Effect.die("unused"), state: () => Effect.die("unused"), @@ -84,22 +106,19 @@ const form = Layer.succeed( const kv = Layer.succeed( KV.Service, KV.Service.of({ - get: () => Effect.succeed(undefined), - set: () => Effect.void, - remove: () => Effect.void, + get: (key) => Effect.succeed(values.get(key)), + set: (key, value) => Effect.sync(() => values.set(key, value)).pipe(Effect.asVoid), + remove: (key) => Effect.sync(() => values.delete(key)).pipe(Effect.asVoid), }), ) const it = testEffect( - AppNodeBuilder.build( - LayerNode.group([Tool.node, WebSearch.node, webSearchToolNode]), - [ - [Permission.node, permission], - [WebSearch.node, websearch], - [Form.node, form], - [KV.node, kv], - [Image.node, imagePassthrough], - ], - ), + AppNodeBuilder.build(LayerNode.group([Tool.node, WebSearch.node, webSearchToolNode]), [ + [Permission.node, permission], + [WebSearch.node, websearch], + [Form.node, form], + [KV.node, kv], + [Image.node, imagePassthrough], + ]), ) describe("WebSearchTool registration", () => { @@ -202,4 +221,116 @@ describe("WebSearchTool registration", () => { }) }), ) + + it.effect("asks once and uses the default provider when web search is first enabled", () => + Effect.gen(function* () { + providerRequired = true + formResponse = { status: "answered", answer: { choice: "allow" } } + const registry = yield* Tool.Service + + expect( + yield* executeTool(registry, { + sessionID, + ...toolIdentity, + call: { type: "tool-call", id: "call-enable", name: "websearch", input: { query: "effect" } }, + }), + ).toMatchObject({ status: "completed", metadata: { provider: "exa" } }) + expect(values.get("websearch:provider")).toBe("exa") + expect(queries).toHaveLength(2) + expect(formRequests).toEqual([ + { + sessionID, + title: "Web Search", + metadata: { kind: "websearch.provider" }, + fields: [ + { + key: "choice", + description: "Allow OpenCode to search the web for up-to-date information?", + type: "string", + required: true, + custom: false, + options: [ + { + value: "allow", + label: "Allow web search via Exa", + }, + { + value: "choose", + label: "Choose another provider", + }, + { value: "disable", label: "Disable web search" }, + ], + }, + ], + }, + ]) + + expect( + yield* executeTool(registry, { + sessionID, + ...toolIdentity, + call: { type: "tool-call", id: "call-enabled", name: "websearch", input: { query: "effect schema" } }, + }), + ).toMatchObject({ status: "completed", metadata: { provider: "exa" } }) + expect(formRequests).toHaveLength(1) + expect(queries).toHaveLength(3) + }), + ) + + it.effect("asks a second form when choosing another provider", () => + Effect.gen(function* () { + providerRequired = true + formResponses.push( + { status: "answered", answer: { choice: "choose" } }, + { status: "answered", answer: { provider: "parallel" } }, + ) + const registry = yield* Tool.Service + + expect( + yield* executeTool(registry, { + sessionID, + ...toolIdentity, + call: { type: "tool-call", id: "call-choose", name: "websearch", input: { query: "effect" } }, + }), + ).toMatchObject({ status: "completed", metadata: { provider: "parallel" } }) + expect(values.get("websearch:provider")).toBe("parallel") + expect(queries).toHaveLength(2) + expect(formRequests[1]).toEqual({ + sessionID, + title: "Choose a web search provider", + metadata: { kind: "websearch.provider" }, + fields: [ + { + key: "provider", + description: "Choose a provider for web search.", + type: "string", + required: true, + custom: false, + options: [ + { value: "exa", label: "Exa" }, + { value: "parallel", label: "Parallel" }, + ], + }, + ], + }) + }), + ) + + it.effect("persists the choice to disable web search", () => + Effect.gen(function* () { + providerRequired = true + formResponse = { status: "answered", answer: { choice: "disable" } } + const registry = yield* Tool.Service + + expect( + yield* executeTool(registry, { + sessionID, + ...toolIdentity, + call: { type: "tool-call", id: "call-disable", name: "websearch", input: { query: "effect" } }, + }), + ).toMatchObject({ status: "error" }) + expect(values.get("websearch:provider")).toBe(false) + expect(queries).toHaveLength(1) + }), + ) })