1
0
Эх сурвалжийг харах

feat(core): route subagent models by role

Filip Hejmowski 23 цаг өмнө
parent
commit
c71ebe46df

+ 90 - 0
packages/core/src/model-routing.ts

@@ -0,0 +1,90 @@
+export * as ModelRouting from "./model-routing.js"
+
+import { Model } from "./model.js"
+import { Provider } from "./provider.js"
+
+export const roles = ["fast", "smart", "vision", "long-context"] as const
+export type Role = (typeof roles)[number]
+
+export function resolve(selection: string, available: readonly Model.Info[]) {
+  if (!isRole(selection)) return exact(selection, available)
+  return select(selection, available)
+}
+
+export function select(role: Role, available: readonly Model.Info[]) {
+  const candidates = available.filter(
+    (model) =>
+      model.status === "active" &&
+      model.capabilities.tools &&
+      model.capabilities.input.includes("text") &&
+      model.capabilities.output.includes("text"),
+  )
+  const eligible =
+    role === "vision"
+      ? candidates.filter((model) => model.capabilities.input.includes("image"))
+      : role === "fast"
+        ? candidates.filter((model) => !SLOW_MODEL_RE.test(identity(model)))
+        : candidates
+  if (eligible.length === 0) return
+
+  const sorted = eligible.toSorted((a, b) => {
+    if (role === "fast") {
+      const tagged = Number(fast(b)) - Number(fast(a))
+      if (tagged !== 0) return tagged
+      const price = cost(a) - cost(b)
+      if (price !== 0) return price
+    }
+    if (role === "smart" || role === "vision") {
+      const tagged = Number(smart(b)) - Number(smart(a))
+      if (tagged !== 0) return tagged
+    }
+    if (role === "long-context") {
+      const context = b.limit.context - a.limit.context
+      if (context !== 0) return context
+    }
+    const released = b.time.released - a.time.released
+    if (released !== 0) return released
+    return `${a.providerID}/${a.id}`.localeCompare(`${b.providerID}/${b.id}`)
+  })
+  const selected = sorted[0]
+  return Model.Ref.make({ providerID: selected.providerID, id: selected.id })
+}
+
+function isRole(selection: string): selection is Role {
+  return roles.includes(selection as Role)
+}
+
+function exact(selection: string, available: readonly Model.Info[]) {
+  const providerEnd = selection.indexOf("/")
+  if (providerEnd <= 0) return
+  const variantStart = selection.indexOf("#", providerEnd + 1)
+  const providerID = Provider.ID.make(selection.slice(0, providerEnd))
+  const id = Model.ID.make(selection.slice(providerEnd + 1, variantStart === -1 ? undefined : variantStart))
+  const variant = variantStart === -1 ? undefined : Model.VariantID.make(selection.slice(variantStart + 1))
+  if (!id || !providerID || (variantStart !== -1 && !variant)) return
+  const model = available.find((item) => item.providerID === providerID && item.id === id)
+  if (!model) return
+  if (variant && !model.variants.some((item) => item.id === variant)) return
+  return Model.Ref.make({ providerID, id, variant })
+}
+
+function cost(model: Model.Info) {
+  const price = model.cost[0]
+  return price ? price.input + price.output : Number.MAX_SAFE_INTEGER
+}
+
+function fast(model: Model.Info) {
+  return FAST_MODEL_RE.test(identity(model))
+}
+
+function smart(model: Model.Info) {
+  return SMART_MODEL_RE.test(identity(model))
+}
+
+function identity(model: Model.Info) {
+  return `${model.id} ${model.family ?? ""} ${model.name}`.toLowerCase()
+}
+
+const FAST_MODEL_RE = /\b(nano|flash|lite|mini|small|fast)\b/
+const SLOW_MODEL_RE = /\b(haiku)\b/
+const SMART_MODEL_RE = /\b(opus|pro|max|ultra|reasoner|reasoning)\b|\b(gpt-5|grok-4|deepseek-v4|kimi-k2)\b/

+ 15 - 2
packages/core/src/tool/plugin/subagent.ts

@@ -4,7 +4,9 @@ import { ToolFailure } from "@opencode-ai/ai"
 import type { Context as PluginContext } from "@opencode-ai/plugin/effect/plugin"
 import { Effect, Schema, Scope } from "effect"
 import { Agent } from "../../agent.js"
+import { Catalog } from "../../catalog.js"
 import { Config } from "../../config.js"
+import { ModelRouting } from "../../model-routing.js"
 import { PluginRuntime } from "../../plugin/runtime.js"
 import { Permission } from "../../permission.js"
 import { SessionSchema } from "../../session/schema.js"
@@ -23,6 +25,10 @@ export const Input = Schema.Struct({
   agent: Schema.String.annotate({ description: "The type of specialized agent to use for this task" }),
   description: Schema.String.annotate({ description: "A short 3-5 word label for the task, displayed to the user" }),
   prompt: Schema.String.annotate({ description: "The task for the subagent to perform" }),
+  model: Schema.optionalKey(Schema.String).annotate({
+    description:
+      'Optional model route. Use "fast" for cheap bounded work, "smart" for difficult reasoning, "vision" for image input, or "long-context" for very large inputs. Pass an exact provider/model ID only when the user requests one. Omit this field to use the agent model or inherit the parent model.',
+  }),
   background: Schema.optionalKey(Schema.Boolean).annotate({
     description:
       "Run the subagent in the background and return immediately. You will be notified when it completes. DO NOT sleep, poll, or proactively check on its progress.",
@@ -47,6 +53,7 @@ export const Plugin = {
   effect: Effect.fn("SubagentTool.Plugin")(function* (ctx: PluginContext) {
     const runtime = yield* PluginRuntime.Service
     const agents = yield* Agent.Service
+    const catalog = yield* Catalog.Service
     const config = yield* Config.Service
     const permission = yield* Permission.Service
     const scope = yield* Scope.Scope
@@ -163,8 +170,14 @@ export const Plugin = {
                 })
                 .pipe(Effect.mapError((error) => new ToolFailure({ message: `Subagent denied: ${agent.id}`, error })))
 
-              // Model selection is policy/config/session state, not an LLM-facing tool argument.
-              const model = agent.model ?? parent.model
+              const routed = input.model
+                ? ModelRouting.resolve(input.model, yield* catalog.model.available())
+                : undefined
+              if (input.model && !routed)
+                return yield* new ToolFailure({
+                  message: `No available model matches route: ${input.model}`,
+                })
+              const model = routed ?? agent.model ?? parent.model
               const child = yield* runtime.session
                 .create({
                   parentID: context.sessionID,

+ 75 - 0
packages/core/test/model-routing.test.ts

@@ -0,0 +1,75 @@
+import { describe, expect, test } from "bun:test"
+import { Money } from "@opencode-ai/schema/money"
+import { ModelRouting } from "@opencode-ai/core/model-routing"
+import { Model } from "@opencode-ai/core/model"
+import { Provider } from "@opencode-ai/core/provider"
+
+const model = (
+  providerID: string,
+  id: string,
+  input: readonly string[],
+  options: { cost?: number; context?: number; released?: number } = {},
+) =>
+  Model.Info.make({
+    ...Model.Info.default(Provider.ID.make(providerID), Model.ID.make(id)),
+    name: id,
+    capabilities: { tools: true, input: [...input], output: ["text"] },
+    time: { released: options.released ?? 1 },
+    cost: [
+      {
+        input: Money.USDPerMillionTokens.make(options.cost ?? 1),
+        output: Money.USDPerMillionTokens.make(options.cost ?? 1),
+        cache: { read: Money.USDPerMillionTokens.zero, write: Money.USDPerMillionTokens.zero },
+      },
+    ],
+    limit: { context: options.context ?? 100_000, output: 10_000 },
+  })
+
+describe("ModelRouting.select", () => {
+  test("routes fast work to an inexpensive fast family without selecting haiku", () => {
+    const selected = ModelRouting.select("fast", [
+      model("anthropic", "claude-haiku-4", ["text"], { cost: 0.1, released: 4 }),
+      model("google", "gemini-flash", ["text"], { cost: 0.2, released: 3 }),
+      model("openai", "gpt-5", ["text"], { cost: 2, released: 5 }),
+    ])
+
+    expect(selected).toEqual(
+      Model.Ref.make({ providerID: Provider.ID.make("google"), id: Model.ID.make("gemini-flash") }),
+    )
+  })
+
+  test("routes smart work to a high-capability family", () => {
+    const selected = ModelRouting.select("smart", [
+      model("google", "gemini-flash", ["text"], { released: 5 }),
+      model("anthropic", "claude-opus-4", ["text"], { released: 3 }),
+    ])
+
+    expect(selected).toEqual(
+      Model.Ref.make({ providerID: Provider.ID.make("anthropic"), id: Model.ID.make("claude-opus-4") }),
+    )
+  })
+
+  test("enforces role capabilities and provider availability through the candidate set", () => {
+    expect(
+      ModelRouting.select("vision", [
+        model("openai", "gpt-5", ["text"], { released: 5 }),
+        model("google", "gemini-pro-vision", ["text", "image"], { released: 3 }),
+      ]),
+    ).toEqual(Model.Ref.make({ providerID: Provider.ID.make("google"), id: Model.ID.make("gemini-pro-vision") }))
+
+    expect(
+      ModelRouting.select("long-context", [
+        model("openai", "gpt-5", ["text"], { context: 200_000 }),
+        model("google", "gemini-pro", ["text"], { context: 1_000_000 }),
+      ]),
+    ).toEqual(Model.Ref.make({ providerID: Provider.ID.make("google"), id: Model.ID.make("gemini-pro") }))
+  })
+
+  test("resolves only exact models present in the available catalog", () => {
+    const available = [model("xai", "grok-4", ["text"])]
+    expect(ModelRouting.resolve("xai/grok-4", available)).toEqual(
+      Model.Ref.make({ providerID: Provider.ID.make("xai"), id: Model.ID.make("grok-4") }),
+    )
+    expect(ModelRouting.resolve("xai/grok-5", available)).toBeUndefined()
+  })
+})

+ 55 - 1
packages/core/test/tool-subagent.test.ts

@@ -8,6 +8,7 @@ import { Global } from "@opencode-ai/util/global"
 import { makeGlobalNode, makeLocationNode } from "@opencode-ai/util/effect/app-node"
 import { Database } from "@opencode-ai/core/database/database"
 import { Bus } from "@opencode-ai/core/bus"
+import { Catalog } from "@opencode-ai/core/catalog"
 import { Config } from "@opencode-ai/core/config"
 import { Location } from "@opencode-ai/core/location"
 import { Model } from "@opencode-ai/core/model"
@@ -36,6 +37,7 @@ import { executeTool, registerToolPlugin, toolIdentity } from "./lib/tool"
 const childText = "child final response"
 const childModel = Model.Ref.make({ id: Model.ID.make("child"), providerID: Provider.ID.make("test") })
 const parentModel = Model.Ref.make({ id: Model.ID.make("parent"), providerID: Provider.ID.make("test") })
+const fastModel = Model.Ref.make({ id: Model.ID.make("gemini-flash"), providerID: Provider.ID.make("route") })
 const tokens = { input: 0, output: 0, reasoning: 0, cache: { read: 0, write: 0 } }
 
 const outputSessionID = (value: unknown) =>
@@ -100,7 +102,7 @@ const subagentPluginSupervisor = makeLocationNode({
     PluginSupervisor.Service,
     registerToolPlugin(SubagentTool.Plugin).pipe(Effect.as(PluginSupervisor.Service.of({ flush: Effect.void }))),
   ),
-  deps: [Agent.node, Config.node, Permission.node, PluginRuntime.node, Tool.node],
+  deps: [Agent.node, Catalog.node, Config.node, Permission.node, PluginRuntime.node, Tool.node],
 })
 
 const nodes = LayerNode.group([
@@ -142,6 +144,27 @@ const withSubagent = (location: Location.Ref) =>
         })
       }),
     ).pipe(Effect.provide(locations.get(location)))
+    yield* Catalog.Service.use((catalog) =>
+      catalog.transform((draft) => {
+        draft.provider.update(fastModel.providerID, (provider) => {
+          provider.activation = "enabled"
+        })
+        draft.model.update(fastModel.providerID, fastModel.id, (model) => {
+          Object.assign(model, Model.Info.default(fastModel.providerID, fastModel.id), {
+            name: "Gemini Flash",
+            family: Model.Family.make("gemini-flash"),
+            time: { released: Date.now() },
+            cost: [
+              {
+                input: Money.USDPerMillionTokens.make(0.1),
+                output: Money.USDPerMillionTokens.make(0.2),
+                cache: { read: Money.USDPerMillionTokens.zero, write: Money.USDPerMillionTokens.zero },
+              },
+            ],
+          })
+        })
+      }),
+    ).pipe(Effect.provide(locations.get(location)))
   })
 
 describe("SubagentTool", () => {
@@ -320,6 +343,37 @@ describe("SubagentTool", () => {
           })
           const fallbackChild = yield* sessions.get(outputSessionID(fallback.metadata))
           expect(fallbackChild).toMatchObject({ parentID: parent.id, model: parentModel })
+
+          const routed = yield* executeTool(registry, {
+            sessionID: parent.id,
+            ...toolIdentity,
+            call: {
+              type: "tool-call",
+              id: "call-subagent-routed",
+              name: SubagentTool.name,
+              input: { agent: "reviewer", description: "fast", prompt: "quick check", model: "fast" },
+            },
+          })
+          expect(yield* sessions.get(outputSessionID(routed.metadata))).toMatchObject({
+            parentID: parent.id,
+            model: fastModel,
+          })
+
+          expect(
+            yield* executeTool(registry, {
+              sessionID: parent.id,
+              ...toolIdentity,
+              call: {
+                type: "tool-call",
+                id: "call-subagent-missing-model",
+                name: SubagentTool.name,
+                input: { agent: "reviewer", description: "missing", prompt: "check", model: "missing/model" },
+              },
+            }),
+          ).toEqual({
+            status: "error",
+            error: { type: "tool.execution", message: "No available model matches route: missing/model" },
+          })
         }),
       ),
     ),