瀏覽代碼

fix(core): deduplicate websearch consent prompts (#40869)

Co-authored-by: James Long <17031+jlongster@users.noreply.github.com>
opencode-agent[bot] 1 周之前
父節點
當前提交
727beae2d5
共有 2 個文件被更改,包括 122 次插入57 次删除
  1. 77 55
      packages/core/src/tool/plugin/websearch.ts
  2. 45 2
      packages/core/test/tool-websearch.test.ts

+ 77 - 55
packages/core/src/tool/plugin/websearch.ts

@@ -2,7 +2,7 @@ export * as WebSearchTool from "./websearch"
 
 import type { Context as PluginContext } from "@opencode-ai/plugin/effect/plugin"
 import { ToolFailure } from "@opencode-ai/ai"
-import { Effect, Schema } from "effect"
+import { Effect, Schema, Semaphore } from "effect"
 import { Form } from "../../form"
 import { KV } from "../../kv"
 import { Permission } from "../../permission"
@@ -10,6 +10,7 @@ import { WebSearch } from "../../websearch"
 
 export const name = "websearch"
 export const NO_RESULTS = "No search results found. Please try a different query."
+const providerSelectionLock = Semaphore.makeUnsafe(1)
 
 export const description = `Search the web using the user's selected search integration. Use this for current information beyond knowledge cutoff.
 
@@ -29,6 +30,7 @@ export const Plugin = {
     const permission = yield* Permission.Service
     const forms = yield* Form.Service
     const kv = yield* KV.Service
+    const websearch = yield* WebSearch.Service
 
     yield* ctx.tool
       .transform((draft) =>
@@ -49,70 +51,90 @@ export const Plugin = {
                 agent: context.agent,
                 source: { type: "tool", messageID: context.messageID, id: context.id },
               })
-              const result = yield* ctx.websearch.query(input).pipe(
-                Effect.catch((error) => {
-                  if (!Schema.is(WebSearch.ProviderRequiredError)(error)) return Effect.fail(error)
-                  return Effect.gen(function* () {
-                    const providers = (yield* ctx.websearch.providers()).data
-                    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: "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 ${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"))
-                    if (response.answer.choice === "disable") {
-                      yield* kv.set("websearch:provider", false)
-                      return yield* new WebSearch.DisabledError()
-                    }
-                    const selection =
-                      response.answer.choice === "choose"
-                        ? yield* forms.ask({
+              const search = (): Effect.Effect<Effect.Success<ReturnType<typeof ctx.websearch.query>>, unknown> =>
+                ctx.websearch.query(input).pipe(
+                  Effect.catch((error) => {
+                    if (!Schema.is(WebSearch.ProviderRequiredError)(error)) return Effect.fail(error)
+                    return providerSelectionLock
+                      .withPermit(
+                        Effect.gen(function* () {
+                          if (yield* websearch.default()) return yield* Effect.void
+                          const providers = (yield* ctx.websearch.providers()).data
+                          const defaultProvider = providers[0]
+                          if (!defaultProvider) return yield* new WebSearch.ProviderRequiredError()
+                          const response = yield* forms.ask({
                             sessionID: context.sessionID,
-                            title: "Choose a web search provider",
+                            title: "Web Search",
                             metadata: { kind: "websearch.provider" },
                             fields: [
                               {
-                                key: "provider",
-                                description: "Choose a provider for web search.",
+                                key: "choice",
+                                description: "Allow OpenCode to search the web for up-to-date information?",
                                 type: "string",
                                 required: true,
                                 custom: false,
-                                options: providers.map((provider) => ({ value: provider.id, label: provider.name })),
+                                options: [
+                                  {
+                                    value: "allow",
+                                    label: `Allow web search via ${defaultProvider.name}`,
+                                  },
+                                  {
+                                    value: "choose",
+                                    label: "Choose another provider",
+                                  },
+                                  { value: "disable", label: "Disable web search" },
+                                ],
                               },
                             ],
                           })
-                        : 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)
-                  })
-                }),
-              )
+                          if (response.status === "cancelled")
+                            return yield* Effect.fail(new Error("Web search cancelled"))
+                          if (response.answer.choice === "disable") {
+                            yield* kv.set("websearch:provider", false)
+                            return yield* new WebSearch.DisabledError()
+                          }
+                          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()
+                          return yield* kv.set("websearch:provider", providerID)
+                        }),
+                      )
+                      .pipe(
+                        Effect.timeoutOrElse({
+                          duration: "1 minute",
+                          orElse: () => Effect.fail(new Error("Web search cancelled")),
+                        }),
+                        Effect.andThen(Effect.suspend(search)),
+                      )
+                  }),
+                )
+              const result = yield* search()
               const output = {
                 provider: result.data.providerID,
                 results: result.data.results,

+ 45 - 2
packages/core/test/tool-websearch.test.ts

@@ -1,5 +1,5 @@
 import { beforeEach, describe, expect } from "bun:test"
-import { Effect, Layer } from "effect"
+import { Deferred, Effect, Layer } from "effect"
 import { AppNodeBuilder } from "@opencode-ai/core/effect/app-node-builder"
 import { LayerNode } from "@opencode-ai/util/effect/layer-node"
 import { Permission } from "@opencode-ai/core/permission"
@@ -39,6 +39,8 @@ const providers = [
 let providerRequired = false
 let formResponse: Form.TerminalState = { status: "cancelled" }
 const formResponses: Form.TerminalState[] = []
+let queryBarrier: Deferred.Deferred<void> | undefined
+let synchronizedQueries = 0
 let result = new WebSearch.Response({
   providerID: WebSearch.ID.make("exa"),
   results: [{ url: "https://example.com", title: "Search results", content: "search results", time: {} }],
@@ -52,6 +54,8 @@ beforeEach(() => {
   providerRequired = false
   formResponse = { status: "cancelled" }
   formResponses.length = 0
+  queryBarrier = undefined
+  synchronizedQueries = 0
   result = new WebSearch.Response({
     providerID: WebSearch.ID.make("exa"),
     results: [{ url: "https://example.com", title: "Search results", content: "search results", time: {} }],
@@ -75,11 +79,21 @@ const websearch = Layer.succeed(
     transform: () => Effect.die("unused"),
     reload: () => Effect.die("unused"),
     providers: () => Effect.succeed(providers),
-    default: () => Effect.succeed(undefined),
+    default: () =>
+      Effect.gen(function* () {
+        const stored = values.get("websearch:provider")
+        if (stored === false) return yield* new WebSearch.DisabledError()
+        return typeof stored === "string" ? providers.find((provider) => provider.id === stored) : undefined
+      }),
     query: (input) =>
       Effect.gen(function* () {
         queries.push(input)
         const stored = values.get("websearch:provider")
+        if (queryBarrier && synchronizedQueries < 5) {
+          synchronizedQueries++
+          if (synchronizedQueries === 5) yield* Deferred.succeed(queryBarrier, undefined)
+          yield* Deferred.await(queryBarrier)
+        }
         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 })
@@ -316,6 +330,35 @@ describe("WebSearchTool registration", () => {
     }),
   )
 
+  it.effect("shares provider consent across concurrent searches", () =>
+    Effect.gen(function* () {
+      providerRequired = true
+      formResponse = { status: "answered", answer: { choice: "allow" } }
+      queryBarrier = yield* Deferred.make<void>()
+      const registry = yield* Tool.Service
+
+      const results = yield* Effect.all(
+        Array.from({ length: 5 }, (_, index) =>
+          executeTool(registry, {
+            sessionID,
+            ...toolIdentity,
+            call: {
+              type: "tool-call",
+              id: `call-concurrent-${index}`,
+              name: "websearch",
+              input: { query: `effect ${index}` },
+            },
+          }),
+        ),
+        { concurrency: "unbounded" },
+      )
+
+      expect(results.every((item) => item.status === "completed")).toBe(true)
+      expect(formRequests).toHaveLength(1)
+      expect(values.get("websearch:provider")).toBe("exa")
+    }),
+  )
+
   it.effect("persists the choice to disable web search", () =>
     Effect.gen(function* () {
       providerRequired = true