From f2df1cc829df8045ebd06850c4561e235d17143c Mon Sep 17 00:00:00 2001 From: Omer Aplak Date: Thu, 22 Jan 2026 07:01:52 -0800 Subject: [PATCH 1/7] feat: allow tool-specific hooks and let `onToolEnd` override tool output #975 --- .changeset/tool-hooks-and-overrides.md | 43 ++++++++ packages/core/src/agent/agent.spec-d.ts | 1 + packages/core/src/agent/agent.spec.ts | 92 ++++++++++++++++ packages/core/src/agent/agent.ts | 122 ++++++++++++++++------ packages/core/src/agent/hooks/index.ts | 8 +- packages/core/src/planagent/plan-agent.ts | 34 +++++- packages/core/src/tool/index.spec.ts | 17 +++ packages/core/src/tool/index.ts | 41 ++++++++ website/docs/agents/hooks.md | 47 ++++++++- website/docs/agents/tools.md | 33 ++++++ 10 files changed, 402 insertions(+), 36 deletions(-) create mode 100644 .changeset/tool-hooks-and-overrides.md diff --git a/.changeset/tool-hooks-and-overrides.md b/.changeset/tool-hooks-and-overrides.md new file mode 100644 index 000000000..20211e7d8 --- /dev/null +++ b/.changeset/tool-hooks-and-overrides.md @@ -0,0 +1,43 @@ +--- +"@voltagent/core": minor +--- + +feat: allow tool-specific hooks and let `onToolEnd` override tool output #975 + +Tool hooks run alongside agent hooks. `onToolEnd` can now return `{ output }` to replace the tool result (validated again if an output schema exists). + +```ts +import { Agent, createTool } from "@voltagent/core"; +import { z } from "zod"; + +const normalizeTool = createTool({ + name: "normalize_text", + description: "Normalizes and truncates text", + parameters: z.object({ text: z.string() }), + execute: async ({ text }) => text, + hooks: { + onStart: ({ tool }) => { + console.log(`[tool] ${tool.name} starting`); + }, + onEnd: ({ output }) => { + if (typeof output === "string") { + return { output: output.slice(0, 1000) }; + } + }, + }, +}); + +const agent = new Agent({ + name: "ToolHooksAgent", + instructions: "Use tools as needed.", + model: myModel, + tools: [normalizeTool], + hooks: { + onToolEnd: ({ output }) => { + if (typeof output === "string") { + return { output: output.trim() }; + } + }, + }, +}); +``` diff --git a/packages/core/src/agent/agent.spec-d.ts b/packages/core/src/agent/agent.spec-d.ts index c40a8f909..336789595 100644 --- a/packages/core/src/agent/agent.spec-d.ts +++ b/packages/core/src/agent/agent.spec-d.ts @@ -551,6 +551,7 @@ describe("Agent Type System", () => { onToolEnd: async ({ context, tool: _tool, output: _output, error: _error }) => { expectTypeOf(context).toMatchTypeOf(); // tool/output/error types are intentionally flexible + return { output: "override" }; }, }; diff --git a/packages/core/src/agent/agent.spec.ts b/packages/core/src/agent/agent.spec.ts index 4695a5001..a4acb4a67 100644 --- a/packages/core/src/agent/agent.spec.ts +++ b/packages/core/src/agent/agent.spec.ts @@ -892,6 +892,98 @@ describe("Agent", () => { operationContext.traceContext.end("completed"); }); + + it("allows onToolEnd to override tool output", async () => { + const onToolEnd = vi.fn().mockResolvedValue({ output: "trimmed" }); + const agent = new Agent({ + name: "TestAgent", + instructions: "Test", + model: mockModel as any, + hooks: createHooks({ onToolEnd }), + }); + + const tool = new Tool({ + name: "text-tool", + description: "Returns text", + parameters: z.object({}), + execute: async () => "original", + }); + + const operationContext = (agent as any).createOperationContext("input"); + const executeFactory = (agent as any).createToolExecutionFactory( + operationContext, + agent.hooks, + ); + + const execute = executeFactory(tool); + const result = await execute({}); + + expect(result).toBe("trimmed"); + expect(onToolEnd).toHaveBeenCalledWith( + expect.objectContaining({ + tool, + output: "original", + error: undefined, + }), + ); + + operationContext.traceContext.end("completed"); + }); + + it("supports tool-level hooks for start and end", async () => { + const toolOnStart = vi.fn(); + const toolOnEnd = vi.fn().mockResolvedValue({ output: "tool-hook" }); + const onToolEnd = vi.fn().mockResolvedValue({ output: "agent-hook" }); + const agent = new Agent({ + name: "TestAgent", + instructions: "Test", + model: mockModel as any, + hooks: createHooks({ onToolEnd }), + }); + + const tool = new Tool({ + name: "hooked-tool", + description: "Returns text", + parameters: z.object({}), + execute: async () => "original", + hooks: { + onStart: toolOnStart, + onEnd: toolOnEnd, + }, + }); + + const operationContext = (agent as any).createOperationContext("input"); + const executeFactory = (agent as any).createToolExecutionFactory( + operationContext, + agent.hooks, + ); + + const execute = executeFactory(tool); + const result = await execute({}); + + expect(result).toBe("agent-hook"); + expect(toolOnStart).toHaveBeenCalledWith( + expect.objectContaining({ + tool, + }), + ); + expect(toolOnEnd).toHaveBeenCalledWith( + expect.objectContaining({ + tool, + output: "original", + error: undefined, + }), + ); + expect(onToolEnd).toHaveBeenCalledWith( + expect.objectContaining({ + tool, + output: "tool-hook", + error: undefined, + }), + ); + + operationContext.traceContext.end("completed"); + }); }); describe("Agent as Tool (toTool)", () => { diff --git a/packages/core/src/agent/agent.ts b/packages/core/src/agent/agent.ts index 47fabe64d..e10af0a0f 100644 --- a/packages/core/src/agent/agent.ts +++ b/packages/core/src/agent/agent.ts @@ -4635,19 +4635,74 @@ export class Agent { oc.systemContext.set("historyEntryId", oc.operationId); oc.systemContext.set("parentToolSpan", toolSpan); - const handleToolSuccess = async (result: any, validatedResult: any) => { - toolSpan.setAttribute("output", safeStringify(result)); - toolSpan.setStatus({ code: SpanStatusCode.OK }); - toolSpan.end(); + const hasOutputOverride = ( + value: unknown, + ): value is { + output?: unknown; + } => { + if (!value || typeof value !== "object") { + return false; + } + return Object.prototype.hasOwnProperty.call(value, "output"); + }; - await hooks.onToolEnd?.({ + const runToolStartHooks = async () => { + await hooks.onToolStart?.({ + agent: this, + tool, + context: oc, + args, + options: executionOptions, + }); + + await tool.hooks?.onStart?.({ + tool, + args, + options: executionOptions, + }); + }; + + const resolveToolEndOutput = async (currentOutput: any) => { + let output = currentOutput; + + const toolHookResult = await tool.hooks?.onEnd?.({ + tool, + args, + output, + error: undefined, + options: executionOptions, + }); + if (hasOutputOverride(toolHookResult)) { + output = toolHookResult.output; + } + + const agentHookResult = await hooks.onToolEnd?.({ agent: this, tool, - output: validatedResult, + output, error: undefined, context: oc, options: executionOptions, }); + if (hasOutputOverride(agentHookResult)) { + output = agentHookResult.output; + } + + if (output !== currentOutput) { + output = await this.validateToolOutput(output, tool); + } + + return output; + }; + + const handleToolSuccess = async (_result: any, validatedResult: any) => { + const finalOutput = await resolveToolEndOutput(validatedResult); + + toolSpan.setAttribute("output", safeStringify(finalOutput)); + toolSpan.setStatus({ code: SpanStatusCode.OK }); + toolSpan.end(); + + return finalOutput; }; const handleToolError = async (errorValue: unknown) => { @@ -4667,6 +4722,14 @@ export class Agent { toolSpan.recordException(error); toolSpan.end(); + await tool.hooks?.onEnd?.({ + tool, + args, + output: undefined, + error: voltAgentError, + options: executionOptions, + }); + await hooks.onToolEnd?.({ agent: this, tool, @@ -4688,13 +4751,7 @@ export class Agent { return async function* (this: Agent): AsyncGenerator { try { await oc.traceContext.withSpan(toolSpan, async () => { - await hooks.onToolStart?.({ - agent: this, - tool, - context: oc, - args, - options: executionOptions, - }); + await runToolStartHooks(); }); const result = execute(args, executionOptions); @@ -4702,15 +4759,15 @@ export class Agent { if (!isAsyncIterable(result)) { const resolved = await result; const validatedResult = await this.validateToolOutput(resolved, tool); - await oc.traceContext.withSpan(toolSpan, async () => { - await handleToolSuccess(resolved, validatedResult); + const finalOutput = await oc.traceContext.withSpan(toolSpan, async () => { + return await handleToolSuccess(resolved, validatedResult); }); - yield resolved; + yield finalOutput; return; } const iterator = result[Symbol.asyncIterator](); - let finalOutput: any = undefined; + let pendingOutput: any = undefined; let validatedResult: any = undefined; let hasOutput = false; @@ -4720,19 +4777,26 @@ export class Agent { break; } - finalOutput = next.value; + if (hasOutput) { + yield pendingOutput; + } + + pendingOutput = next.value; hasOutput = true; - validatedResult = await this.validateToolOutput(finalOutput, tool); - yield finalOutput; + validatedResult = await this.validateToolOutput(pendingOutput, tool); } if (!hasOutput) { - validatedResult = await this.validateToolOutput(finalOutput, tool); + validatedResult = await this.validateToolOutput(pendingOutput, tool); } - await oc.traceContext.withSpan(toolSpan, async () => { - await handleToolSuccess(finalOutput, validatedResult); + const finalOutput = await oc.traceContext.withSpan(toolSpan, async () => { + return await handleToolSuccess(pendingOutput, validatedResult); }); + + if (hasOutput || finalOutput !== undefined) { + yield finalOutput; + } } catch (e) { const errorResult = await oc.traceContext.withSpan(toolSpan, async () => { return await handleToolError(e); @@ -4745,13 +4809,7 @@ export class Agent { return oc.traceContext.withSpan(toolSpan, async () => { try { // Call tool start hook - can throw ToolDeniedError - await hooks.onToolStart?.({ - agent: this, - tool, - context: oc, - args, - options: executionOptions, - }); + await runToolStartHooks(); // Execute tool with merged options if (!tool.execute) { @@ -4769,9 +4827,9 @@ export class Agent { const validatedResult = await this.validateToolOutput(result, tool); - await handleToolSuccess(result, validatedResult); + const finalOutput = await handleToolSuccess(result, validatedResult); - return result; + return finalOutput; } catch (e) { return await handleToolError(e); } finally { diff --git a/packages/core/src/agent/hooks/index.ts b/packages/core/src/agent/hooks/index.ts index 6896e0e35..2d7af227c 100644 --- a/packages/core/src/agent/hooks/index.ts +++ b/packages/core/src/agent/hooks/index.ts @@ -70,6 +70,10 @@ export interface OnToolEndHookArgs { options?: ToolExecuteOptions; } +export interface OnToolEndHookResult { + output?: unknown; +} + export interface OnPrepareMessagesHookArgs { /** The messages that will be sent to the LLM (AI SDK UIMessage). */ messages: UIMessage[]; @@ -168,7 +172,9 @@ export type AgentHookOnEnd = (args: OnEndHookArgs) => Promise | void; export type AgentHookOnHandoff = (args: OnHandoffHookArgs) => Promise | void; export type AgentHookOnHandoffComplete = (args: OnHandoffCompleteHookArgs) => Promise | void; export type AgentHookOnToolStart = (args: OnToolStartHookArgs) => Promise | void; -export type AgentHookOnToolEnd = (args: OnToolEndHookArgs) => Promise | void; +export type AgentHookOnToolEnd = ( + args: OnToolEndHookArgs, +) => Promise | OnToolEndHookResult | undefined; export type AgentHookOnPrepareMessages = ( args: OnPrepareMessagesHookArgs, ) => Promise | OnPrepareMessagesHookResult; diff --git a/packages/core/src/planagent/plan-agent.ts b/packages/core/src/planagent/plan-agent.ts index 856e19dee..b2f416b67 100644 --- a/packages/core/src/planagent/plan-agent.ts +++ b/packages/core/src/planagent/plan-agent.ts @@ -551,6 +551,38 @@ function chainPrepareModelMessagesHooks( }; } +function chainToolEndHooks( + hooks: Array, +): AgentHooks["onToolEnd"] | undefined { + const sequence = compactHooks(hooks); + if (sequence.length === 0) { + return undefined; + } + return async (args) => { + if (args.error) { + for (const hook of sequence) { + await hook(args); + } + return; + } + + let currentOutput = args.output; + let hasOverride = false; + for (const hook of sequence) { + const result = await hook({ ...args, output: currentOutput }); + if (result && Object.prototype.hasOwnProperty.call(result, "output")) { + currentOutput = result.output; + hasOverride = true; + } + } + + if (hasOverride) { + return { output: currentOutput }; + } + return; + }; +} + function composeAgentHooks(sequences: { onStart?: AgentHooks["onStart"][]; onEnd?: AgentHooks["onEnd"][]; @@ -569,7 +601,7 @@ function composeAgentHooks(sequences: { onHandoff: chainVoidHooks(sequences.onHandoff ?? []), onHandoffComplete: chainVoidHooks(sequences.onHandoffComplete ?? []), onToolStart: chainVoidHooks(sequences.onToolStart ?? []), - onToolEnd: chainVoidHooks(sequences.onToolEnd ?? []), + onToolEnd: chainToolEndHooks(sequences.onToolEnd ?? []), onPrepareMessages: chainPrepareMessagesHooks(sequences.onPrepareMessages ?? []), onPrepareModelMessages: chainPrepareModelMessagesHooks(sequences.onPrepareModelMessages ?? []), onError: chainVoidHooks(sequences.onError ?? []), diff --git a/packages/core/src/tool/index.spec.ts b/packages/core/src/tool/index.spec.ts index 676ac572b..76d3d00f5 100644 --- a/packages/core/src/tool/index.spec.ts +++ b/packages/core/src/tool/index.spec.ts @@ -74,6 +74,23 @@ describe("Tool", () => { expect(tool.needsApproval).toBe(true); }); + it("should keep hooks when provided", () => { + const onStart = vi.fn(); + const onEnd = vi.fn(); + const options = { + name: "hookedTool", + description: "Hooked tool", + parameters: z.object({}), + execute: vi.fn(), + hooks: { onStart, onEnd }, + }; + + const tool = new Tool(options); + + expect(tool.hooks?.onStart).toBe(onStart); + expect(tool.hooks?.onEnd).toBe(onEnd); + }); + it("should throw error if name is missing", () => { const options = { parameters: z.object({}), diff --git a/packages/core/src/tool/index.ts b/packages/core/src/tool/index.ts index e5acb93f5..31c5c0a57 100644 --- a/packages/core/src/tool/index.ts +++ b/packages/core/src/tool/index.ts @@ -11,6 +11,36 @@ type JSONValue = string | number | boolean | null | { [key: string]: JSONValue } export type ToolExecutionResult = PromiseLike | AsyncIterable | T; +export interface ToolHookOnStartArgs { + tool: Tool; + args: unknown; + options?: ToolExecuteOptions; +} + +export interface ToolHookOnEndArgs { + tool: Tool; + args: unknown; + /** The successful output from the tool. Undefined on error. */ + output: unknown | undefined; + /** The error if the tool execution failed. */ + error: unknown | undefined; + options?: ToolExecuteOptions; +} + +export interface ToolHookOnEndResult { + output?: unknown; +} + +export type ToolHookOnStart = (args: ToolHookOnStartArgs) => Promise | void; +export type ToolHookOnEnd = ( + args: ToolHookOnEndArgs, +) => Promise | ToolHookOnEndResult | undefined; + +export type ToolHooks = { + onStart?: ToolHookOnStart; + onEnd?: ToolHookOnEnd; +}; + /** * Tool result output format for multi-modal content. * Matches AI SDK's LanguageModelV2ToolResultOutput type. @@ -142,6 +172,11 @@ export type ToolOptions< args: z.infer, options?: ToolExecuteOptions, ) => ToolExecutionResult : unknown>; + + /** + * Optional tool-specific hooks for lifecycle events. + */ + hooks?: ToolHooks; }; /** @@ -198,6 +233,11 @@ export class Tool : unknown; }) => ToolResultOutput; + /** + * Optional tool-specific hooks for lifecycle events. + */ + readonly hooks?: ToolHooks; + /** * Internal discriminator to make runtime/type checks simpler across module boundaries. * Marking our Tool instances with a stable string avoids instanceof issues. @@ -248,6 +288,7 @@ export class Tool { ### `onToolStart` - **Triggered:** Before an agent executes a tool. -- **Argument Object (`OnToolStartHookArgs`):** `{ agent: Agent, tool: AgentTool, args: any, context: OperationContext }` +- **Argument Object (`OnToolStartHookArgs`):** + - `agent`: Agent instance running the tool + - `tool`: The tool being executed + - `args`: Tool input arguments + - `context`: Operation context for the current call + - `options`: ToolExecuteOptions (includes `toolContext`, `abortController`, etc.) - **Use Cases:** Logging tool usage, inspecting tool arguments, or validating inputs before execution. ```ts @@ -396,7 +401,14 @@ onToolStart: async ({ agent, tool, args, context }) => { ### `onToolEnd` - **Triggered:** After a tool execution completes or throws an error. -- **Argument Object (`OnToolEndHookArgs`):** `{ agent: Agent, tool: AgentTool, output: unknown | undefined, error: VoltAgentError | undefined, context: OperationContext }` +- **Argument Object (`OnToolEndHookArgs`):** + - `agent`: Agent instance running the tool + - `tool`: The tool that completed + - `output`: Tool output (undefined on error) + - `error`: VoltAgentError when tool throws (undefined on success) + - `context`: Operation context for the current call + - `options`: ToolExecuteOptions (includes `toolContext`, `abortController`, etc.) +- **Return:** `{ output }` to replace the tool result. The replacement is validated again if the tool has an `outputSchema`. - **Use Cases:** Logging tool results or errors, post-processing output, triggering actions based on success or failure. ```ts @@ -415,6 +427,37 @@ onToolEnd: async ({ agent, tool, output, error, context }) => { }; ``` +### Tool-level hooks (per tool) + +Tool hooks run for a specific tool instance and are called before/after execution. Tool-level `onEnd` runs before agent-level `onToolEnd`. If both return `{ output }`, the agent hook wins. Any override is re-validated when `outputSchema` is present. + +**Tool hook parameters:** + +- `onStart`: `{ tool, args, options }` +- `onEnd`: `{ tool, args, output, error, options }` (return `{ output }` to override) + +```ts +import { createTool } from "@voltagent/core"; +import { z } from "zod"; + +const normalizeTool = createTool({ + name: "normalize_text", + description: "Normalize and trim text", + parameters: z.object({ text: z.string() }), + execute: async ({ text }) => text, + hooks: { + onStart: ({ tool }) => { + console.log(`[tool] ${tool.name} starting`); + }, + onEnd: ({ output }) => { + if (typeof output === "string") { + return { output: output.trim() }; + } + }, + }, +}); +``` + ### `onHandoff` - **Triggered:** When one agent delegates a task to another agent via the `delegate_task` tool. diff --git a/website/docs/agents/tools.md b/website/docs/agents/tools.md index f860de813..f438b5678 100644 --- a/website/docs/agents/tools.md +++ b/website/docs/agents/tools.md @@ -46,6 +46,39 @@ Each tool has: The `execute` function's parameter types are automatically inferred from the Zod schema, providing full IntelliSense support. +## Tool Hooks + +You can attach tool-specific hooks to observe or post-process a tool result. Tool hooks run before agent-level hooks, and agent `onToolEnd` can still override the output afterward. + +**Hook parameters:** + +- `onStart`: `{ tool, args, options }` +- `onEnd`: `{ tool, args, output, error, options }` (return `{ output }` to override the result) + +> Overrides are re-validated if the tool has an `outputSchema`. For streaming tools (AsyncIterable), overrides apply only to the final output. + +```ts +import { createTool } from "@voltagent/core"; +import { z } from "zod"; + +const summarizeTool = createTool({ + name: "summarize_text", + description: "Summarize text with a hard cap", + parameters: z.object({ text: z.string() }), + execute: async ({ text }) => text, + hooks: { + onStart: ({ tool }) => { + console.log(`[tool] ${tool.name} starting`); + }, + onEnd: ({ output }) => { + if (typeof output === "string") { + return { output: output.slice(0, 500) }; + } + }, + }, +}); +``` + ## Streaming Tool Results (Preliminary) If your tool can provide progress or intermediate status, return an `AsyncIterable` from `execute`. From 939d7ba42e05d3338c53574cbeb68b5c3d932046 Mon Sep 17 00:00:00 2001 From: Omer Aplak Date: Thu, 22 Jan 2026 07:04:13 -0800 Subject: [PATCH 2/7] chore: add recipes --- website/recipes/tool-hooks.md | 65 +++++++++++++++++++++++++++++++++++ website/sidebarsRecipes.ts | 5 +++ 2 files changed, 70 insertions(+) create mode 100644 website/recipes/tool-hooks.md diff --git a/website/recipes/tool-hooks.md b/website/recipes/tool-hooks.md new file mode 100644 index 000000000..877a00dbf --- /dev/null +++ b/website/recipes/tool-hooks.md @@ -0,0 +1,65 @@ +--- +id: tool-hooks +title: Tool Hooks +slug: tool-hooks +description: Override or post-process tool results with tool-level and agent-level hooks. +--- + +# Tool Hooks + +Use tool hooks to observe execution and optionally replace tool outputs. Tool-level hooks run first, then agent `onToolEnd` runs and can override the result again. + +## Quick Setup + +```typescript +import { openai } from "@ai-sdk/openai"; +import { Agent, createTool, VoltAgent } from "@voltagent/core"; +import { honoServer } from "@voltagent/server-hono"; +import { z } from "zod"; + +const normalizeTool = createTool({ + name: "normalize_text", + description: "Normalize and cap text length", + parameters: z.object({ text: z.string() }), + execute: async ({ text }) => text, + hooks: { + onStart: ({ tool }) => { + console.log(`[tool] ${tool.name} starting`); + }, + onEnd: ({ output }) => { + if (typeof output === "string") { + return { output: output.slice(0, 1000) }; + } + }, + }, +}); + +const agent = new Agent({ + name: "ToolHooksAgent", + instructions: "Use tools when needed.", + model: openai("gpt-4o-mini"), + tools: [normalizeTool], + hooks: { + onToolEnd: ({ output }) => { + if (typeof output === "string") { + return { output: output.trim() }; + } + }, + }, +}); + +new VoltAgent({ + agents: { agent }, + server: honoServer({ port: 3141 }), +}); +``` + +## Notes + +- Tool hook parameters: + - `onStart`: `{ tool, args, options }` + - `onEnd`: `{ tool, args, output, error, options }` (return `{ output }` to override) +- Agent `onToolEnd` receives `{ agent, tool, output, error, context, options }` and can also return `{ output }`. +- Overrides are re-validated if the tool has an `outputSchema`. +- For streaming tools (AsyncIterable), overrides apply to the final output only. +- These hooks also work with `PlanAgent`. diff --git a/website/sidebarsRecipes.ts b/website/sidebarsRecipes.ts index eae8d0627..85a7d4e22 100644 --- a/website/sidebarsRecipes.ts +++ b/website/sidebarsRecipes.ts @@ -86,6 +86,11 @@ const sidebars: SidebarsConfig = { id: "hooks", label: "Hooks", }, + { + type: "doc", + id: "tool-hooks", + label: "Tool Hooks", + }, { type: "doc", id: "retrying", From 86f94d158859edc55282b8cac37840b257276722 Mon Sep 17 00:00:00 2001 From: Omer Aplak Date: Thu, 22 Jan 2026 18:24:27 -0800 Subject: [PATCH 3/7] feat: add tool routing for agents with router tools, pool/expose controls, and embedding routing --- .changeset/yellow-carpets-admire.md | 67 ++ docs/tool-routing-plan.md | 99 +++ examples/README.md | 1 + examples/base/src/index.ts | 8 +- examples/github-repo-analyzer/src/index.ts | 11 +- examples/with-auth/src/index.ts | 8 +- examples/with-cloudflare-workers/README.md | 2 +- examples/with-lancedb/src/retriever/index.ts | 3 +- examples/with-subagents/src/index.ts | 4 +- examples/with-tool-routing/README.md | 53 ++ examples/with-tool-routing/package.json | 38 + examples/with-tool-routing/src/index.ts | 95 +++ examples/with-tool-routing/tsconfig.json | 14 + examples/with-tools/src/index.ts | 4 +- examples/with-vector-search/README.md | 5 +- examples/with-vector-search/src/index.ts | 8 +- .../src/index.ts | 5 +- examples/with-working-memory/README.md | 5 +- examples/with-working-memory/src/index.ts | 8 +- packages/core/CHANGELOG.md | 5 +- .../generate-model-provider-registry.js | 47 +- packages/core/src/agent/agent.ts | 692 +++++++++++++++++- packages/core/src/agent/context-keys.ts | 1 + packages/core/src/agent/hooks/index.ts | 2 +- packages/core/src/agent/types.ts | 15 +- packages/core/src/index.ts | 3 + .../memory/adapters/embedding/ai-sdk.spec.ts | 40 +- .../src/memory/adapters/embedding/ai-sdk.ts | 55 +- .../src/memory/adapters/embedding/types.ts | 7 + packages/core/src/memory/index.spec-d.ts | 22 + packages/core/src/memory/index.ts | 39 +- packages/core/src/memory/types.ts | 20 +- packages/core/src/planagent/plan-agent.ts | 25 +- .../core/src/registries/agent-registry.ts | 16 + .../embedding-model-router-types.generated.ts | 62 ++ .../embedding-model-router-types.ts | 1 + .../src/registries/model-provider-registry.ts | 93 ++- packages/core/src/tool/index.ts | 17 + .../core/src/tool/manager/BaseToolManager.ts | 17 + packages/core/src/tool/routing/constants.ts | 1 + packages/core/src/tool/routing/embedding.ts | 195 +++++ packages/core/src/tool/routing/index.ts | 86 +++ packages/core/src/tool/routing/types.ts | 103 +++ packages/core/src/types.ts | 5 + packages/core/src/voltagent.ts | 4 + packages/postgres/CHANGELOG.md | 4 +- packages/voltagent-memory/CHANGELOG.md | 5 +- pnpm-lock.yaml | 34 + website/deployment-docs/cloudflare-workers.md | 2 +- website/docs/agents/memory/cloudflare-d1.md | 2 +- website/docs/agents/memory/in-memory.md | 10 +- website/docs/agents/memory/libsql.md | 5 +- website/docs/agents/memory/managed-memory.md | 5 +- website/docs/agents/memory/overview.md | 5 +- website/docs/agents/memory/postgres.md | 5 +- website/docs/agents/memory/semantic-search.md | 28 +- website/docs/agents/memory/supabase.md | 5 +- .../docs/getting-started/migration-guide.md | 7 +- .../docs/getting-started/providers-models.md | 2 +- website/docs/rag/lancedb.md | 6 +- website/docs/tools/tool-routing.md | 312 ++++++++ website/evaluation-docs/prebuilt-scorers.md | 4 +- website/recipes/memory.md | 4 +- website/recipes/tool-routing.md | 133 ++++ website/sidebars.ts | 2 +- website/sidebarsRecipes.ts | 5 + 66 files changed, 2453 insertions(+), 143 deletions(-) create mode 100644 .changeset/yellow-carpets-admire.md create mode 100644 docs/tool-routing-plan.md create mode 100644 examples/with-tool-routing/README.md create mode 100644 examples/with-tool-routing/package.json create mode 100644 examples/with-tool-routing/src/index.ts create mode 100644 examples/with-tool-routing/tsconfig.json create mode 100644 packages/core/src/registries/embedding-model-router-types.generated.ts create mode 100644 packages/core/src/registries/embedding-model-router-types.ts create mode 100644 packages/core/src/tool/routing/constants.ts create mode 100644 packages/core/src/tool/routing/embedding.ts create mode 100644 packages/core/src/tool/routing/index.ts create mode 100644 packages/core/src/tool/routing/types.ts create mode 100644 website/docs/tools/tool-routing.md create mode 100644 website/recipes/tool-routing.md diff --git a/.changeset/yellow-carpets-admire.md b/.changeset/yellow-carpets-admire.md new file mode 100644 index 000000000..d5093e448 --- /dev/null +++ b/.changeset/yellow-carpets-admire.md @@ -0,0 +1,67 @@ +--- +"@voltagent/core": patch +--- + +feat: add tool routing for agents with router tools, pool/expose controls, and embedding routing. + +Embedding model strings also accept provider-qualified IDs like `openai/text-embedding-3-small` using the same model registry as agent model strings. + +Basic embedding router: + +```ts +import { openai } from "@ai-sdk/openai"; +import { Agent, createTool } from "@voltagent/core"; +import { z } from "zod"; + +const getWeather = createTool({ + name: "get_weather", + description: "Get the current weather for a city", + parameters: z.object({ location: z.string() }), + execute: async ({ location }) => ({ location, temperatureC: 22 }), +}); + +const agent = new Agent({ + name: "Tool Routing Agent", + instructions: "Use tool_router for tools. Pass the user request as the query.", + model: "openai/gpt-4o-mini", + tools: [getWeather], + toolRouting: { + embedding: openai.embedding("text-embedding-3-small"), + topK: 2, + }, +}); +``` + +Pool and expose: + +```ts +const agent = new Agent({ + name: "Support Agent", + instructions: "Use tool_router for tools.", + model: "openai/gpt-4o-mini", + toolRouting: { + embedding: "text-embedding-3-small", + pool: [getWeather], + expose: [getStatus], + }, +}); +``` + +Custom router strategy + resolver mode: + +```ts +import { createToolRouter, type ToolArgumentResolver } from "@voltagent/core"; + +const resolver: ToolArgumentResolver = async ({ query, tool }) => { + if (tool.name === "get_weather") return { location: query }; + return {}; +}; + +const router = createToolRouter({ + name: "tool_router", + description: "Route requests with a resolver", + embedding: "text-embedding-3-small", + mode: "resolver", + resolver, +}); +``` diff --git a/docs/tool-routing-plan.md b/docs/tool-routing-plan.md new file mode 100644 index 000000000..1dcf355e0 --- /dev/null +++ b/docs/tool-routing-plan.md @@ -0,0 +1,99 @@ +# Tool Routing Implementation Plan + +This document captures the agreed implementation plan for VoltAgent tool routing with router tools, tool pools, and optional embedding-based selection. We will track execution with Markdown checkboxes. + +## Decisions (Locked) + +- Tool routing config is supported at both Agent and VoltAgent levels (global default + per-agent override). +- Pool includes user-defined tools, provider-defined tools, and MCP tools. +- Default router execution mode is "agent"; users can override. +- Agent mode uses the same model by default; can be overridden with executionModel. +- Args are generated via generateText with structured output (schema-based output). +- Provider tool selection triggers agent-mode fallback and emits an info log. +- Embedding selection auto-activates when an embedding model or adapter is provided. +- Embedding index uses in-memory cache (extensible later). +- Router executes multiple selected tools in parallel. +- API visibility includes pool tools (not hidden); observability also includes pool tools. +- Tool approvals and hooks (tool hooks + agent onToolStart/onToolEnd) still run for pool tools. + +## Scope + +- Core API: ToolRoutingConfig, ToolRouterStrategy, ToolRouter, execution modes. +- Agent runtime: tool pool, router execution, tool execution path reuse. +- Embedding strategy: optional selector using EmbeddingAdapter / AiSdkEmbeddingAdapter. +- Documentation: recipe + usage examples. + +## Checklist + +### 1) API + Types + +- [x] Add ToolRoutingConfig (global + per-agent) to types. +- [x] Define ToolRouterStrategy interface and ToolRouter types. +- [x] Define execution mode enums and router result types. +- [x] Define embedding strategy config (embedding model/adapter, topK, cache). + +### 2) Registry + Defaults + +- [x] Add global toolRouting defaults to AgentRegistry. +- [x] Wire VoltAgentOptions.toolRouting to registry defaults. +- [x] Add agent internal setter to apply default tool routing when unset. + +### 3) Tool Pool Manager + +- [x] Introduce ToolPoolManager (or extend ToolManager) to hold pool tools. +- [x] Add lookup by name (for executing pool tools). +- [x] Ensure pool supports user-defined, provider-defined, and MCP tools. + +### 4) Router Tool Runtime + +- [x] Implement createToolRouter (router tool factory). +- [x] Agent.prepareTools uses routers + exposed tools; pool tools are not added to LLM tools by default. +- [x] Router execution path: + - [x] Select tools via strategy. + - [x] Execute selected tools in parallel. + - [x] Return structured router output. +- [x] Ensure tool hooks + approvals run (no bypass). + +### 5) Agent Mode Execution + +- [x] Implement agent-mode arg generation via generateText with structured output. +- [x] Default to agent model; allow executionModel override. +- [x] Provider tool fallback: + - [x] Force toolChoice to the selected provider tool. + - [x] Log info for fallback. + +### 6) Embedding Strategy (Optional) + +- [x] Add embedding-based ToolRouterStrategy. +- [x] Auto-enable when embedding model/adapter is provided. +- [x] Implement tool-to-text serialization and in-memory embedding cache. +- [x] Invalidate cache when tool pool changes. + +### 7) Observability + API + +- [x] Include pool tools in API responses (getToolsForApi / /agents). +- [x] Add router + selection metadata to spans/logs (safeStringify). +- [x] Ensure pool tools appear in observability with correct tool names. + +### 8) Tests + +- [ ] Unit tests for strategy selection and router output shape. +- [ ] Agent-mode arg generation tests (schema output). +- [ ] Provider tool fallback tests. +- [ ] Embedding strategy tests (cache + selection order). + +### 9) Docs + Recipes + +- [x] New recipe: tool routing with router + pool. +- [x] Embedding-based routing example. +- [x] Update sidebars (if needed). + +## Open Questions + +- None. + +## Notes + +- Use safeStringify for logs and span attributes. +- Keep output schemas for router results explicit. +- Maintain compatibility with PlanAgent (router tools should work there too). diff --git a/examples/README.md b/examples/README.md index fa48fd1fb..b4d3840db 100644 --- a/examples/README.md +++ b/examples/README.md @@ -132,6 +132,7 @@ Create a multi-agent research workflow where different AI agents collaborate to - [Supabase](./with-supabase) — Use Supabase auth/database in tools and server endpoints. - [Tavily Search](./with-tavily-search) — Augment answers with web results from Tavily. - [Thinking Tool](./with-thinking-tool) — Structured reasoning via a dedicated “thinking” tool and schema. +- [Tool Routing](./with-tool-routing) — Route large tool pools through a small set of router tools. - [Tools](./with-tools) — Author Zod‑typed tools with cancellation and streaming support. - [VoltOps Actions + Airtable](./with-voltagent-actions) — Call VoltOps Actions as tools to create and list Airtable records. - [Turso](./with-turso) — Persist memory on LibSQL/Turso with simple setup. diff --git a/examples/base/src/index.ts b/examples/base/src/index.ts index 75766d29e..ff58a0c20 100644 --- a/examples/base/src/index.ts +++ b/examples/base/src/index.ts @@ -1,12 +1,8 @@ -import { openai } from "@ai-sdk/openai"; import { Agent, Memory, VoltAgent } from "@voltagent/core"; +import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; import { createPinoLogger } from "@voltagent/logger"; import { honoServer } from "@voltagent/server-hono"; -// Import Memory and TelemetryStore from core -import { AiSdkEmbeddingAdapter, InMemoryVectorAdapter } from "@voltagent/core"; -import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; - // Create logger const logger = createPinoLogger({ name: "base", @@ -16,7 +12,7 @@ const logger = createPinoLogger({ // Create Memory instance with vector support for semantic search and working memory const memory = new Memory({ storage: new LibSQLMemoryAdapter(), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new LibSQLVectorAdapter(), }); diff --git a/examples/github-repo-analyzer/src/index.ts b/examples/github-repo-analyzer/src/index.ts index 89181d084..dd54ca152 100644 --- a/examples/github-repo-analyzer/src/index.ts +++ b/examples/github-repo-analyzer/src/index.ts @@ -1,11 +1,4 @@ -import { openai } from "@ai-sdk/openai"; -import { - Agent, - AiSdkEmbeddingAdapter, - InMemoryVectorAdapter, - Memory, - VoltAgent, -} from "@voltagent/core"; +import { Agent, InMemoryVectorAdapter, Memory, VoltAgent } from "@voltagent/core"; import { LibSQLMemoryAdapter } from "@voltagent/libsql"; import { createPinoLogger } from "@voltagent/logger"; import { honoServer } from "@voltagent/server-hono"; @@ -20,7 +13,7 @@ const logger = createPinoLogger({ const memory = new Memory({ storage: new LibSQLMemoryAdapter({}), - embedding: new AiSdkEmbeddingAdapter(openai.embeddingModel("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new InMemoryVectorAdapter(), }); diff --git a/examples/with-auth/src/index.ts b/examples/with-auth/src/index.ts index 6ae859e7b..b8d7a2196 100644 --- a/examples/with-auth/src/index.ts +++ b/examples/with-auth/src/index.ts @@ -1,12 +1,8 @@ -import { openai } from "@ai-sdk/openai"; import { Agent, Memory, VoltAgent } from "@voltagent/core"; +import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; import { createPinoLogger } from "@voltagent/logger"; import { authNext, honoServer, jwtAuth } from "@voltagent/server-hono"; -// Import Memory and TelemetryStore from core -import { AiSdkEmbeddingAdapter, InMemoryVectorAdapter } from "@voltagent/core"; -import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; - // Import tools import { weatherTool } from "./tools/index.js"; @@ -19,7 +15,7 @@ const logger = createPinoLogger({ // Create Memory instance with vector support for semantic search and working memory const memory = new Memory({ storage: new LibSQLMemoryAdapter(), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new LibSQLVectorAdapter(), }); diff --git a/examples/with-cloudflare-workers/README.md b/examples/with-cloudflare-workers/README.md index 165594fc6..89d7b848a 100644 --- a/examples/with-cloudflare-workers/README.md +++ b/examples/with-cloudflare-workers/README.md @@ -149,7 +149,7 @@ This example uses in-memory storage adapters: ```typescript const memory = new Memory({ storage: new InMemoryStorageAdapter(), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new InMemoryVectorAdapter(), }); ``` diff --git a/examples/with-lancedb/src/retriever/index.ts b/examples/with-lancedb/src/retriever/index.ts index 65d5382d7..12b12e334 100644 --- a/examples/with-lancedb/src/retriever/index.ts +++ b/examples/with-lancedb/src/retriever/index.ts @@ -1,6 +1,5 @@ import fs from "node:fs/promises"; import path from "node:path"; -import { openai } from "@ai-sdk/openai"; import { type Connection, type Table, connect } from "@lancedb/lancedb"; import { type BaseMessage, BaseRetriever, type RetrieveOptions } from "@voltagent/core"; import { embed } from "ai"; @@ -40,7 +39,7 @@ let table: Table | null = null; async function getEmbedding(text: string): Promise { const { embedding } = await embed({ - model: openai.embedding("text-embedding-3-small"), + model: "openai/text-embedding-3-small", value: text, }); return embedding; diff --git a/examples/with-subagents/src/index.ts b/examples/with-subagents/src/index.ts index da0304ba4..15856156e 100644 --- a/examples/with-subagents/src/index.ts +++ b/examples/with-subagents/src/index.ts @@ -1,7 +1,5 @@ -import { openai } from "@ai-sdk/openai"; import { Agent, - AiSdkEmbeddingAdapter, InMemoryVectorAdapter, Memory, VoltAgent, @@ -21,7 +19,7 @@ const logger = createPinoLogger({ const memory = new Memory({ storage: new LibSQLMemoryAdapter(), - embedding: new AiSdkEmbeddingAdapter(openai.embeddingModel("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new InMemoryVectorAdapter(), }); diff --git a/examples/with-tool-routing/README.md b/examples/with-tool-routing/README.md new file mode 100644 index 000000000..bf0d7ed18 --- /dev/null +++ b/examples/with-tool-routing/README.md @@ -0,0 +1,53 @@ +
+ +435380213-b6253409-8741-462b-a346-834cd18565a9 + + +
+
+ + +
+ +
+ +
+ VoltAgent is an open source TypeScript framework for building and orchestrating AI agents.
+Escape the limitations of no-code builders and the complexity of starting from scratch. +
+
+
+ +
+ +[![npm version](https://img.shields.io/npm/v/@voltagent/core.svg)](https://www.npmjs.com/package/@voltagent/core) +[![Contributor Covenant](https://img.shields.io/badge/Contributor%20Covenant-2.0-4baaaa.svg)](CODE_OF_CONDUCT.md) +[![Discord](https://img.shields.io/discord/1361559153780195478.svg?label=&logo=discord&logoColor=ffffff&color=7389D8&labelColor=6A7EC2)](https://s.voltagent.dev/discord) +[![Twitter Follow](https://img.shields.io/twitter/follow/voltagent_dev?style=social)](https://twitter.com/voltagent_dev) + +
+ +
+ +
+ +VoltAgent Schema + + +
+ +## VoltAgent: Build AI Agents Fast and Flexibly + +VoltAgent is an open-source TypeScript framework for creating and managing AI agents. It provides modular components to build, customize, and scale agents with ease. From connecting to APIs and memory management to supporting multiple LLMs, VoltAgent simplifies the process of creating sophisticated AI systems. It enables fast development, maintains clean code, and offers flexibility to switch between models and tools without vendor lock-in. + +## Try Example + +```bash +npm create voltagent-app@latest -- --example with-tool-routing +``` diff --git a/examples/with-tool-routing/package.json b/examples/with-tool-routing/package.json new file mode 100644 index 000000000..50f959fd2 --- /dev/null +++ b/examples/with-tool-routing/package.json @@ -0,0 +1,38 @@ +{ + "name": "voltagent-example-with-tool-routing", + "author": "", + "dependencies": { + "@ai-sdk/openai": "^3.0.0", + "@voltagent/cli": "^0.1.21", + "@voltagent/core": "^2.1.6", + "@voltagent/logger": "^2.0.2", + "@voltagent/server-hono": "^2.0.4", + "ai": "^6.0.0", + "zod": "^3.25.76" + }, + "devDependencies": { + "@types/node": "^24.2.1", + "tsx": "^4.19.3", + "typescript": "^5.8.2" + }, + "keywords": [ + "agent", + "ai", + "tool-routing", + "voltagent" + ], + "license": "MIT", + "private": true, + "repository": { + "type": "git", + "url": "https://github.com/VoltAgent/voltagent.git", + "directory": "examples/with-tool-routing" + }, + "scripts": { + "build": "tsc", + "dev": "tsx watch --env-file=.env ./src", + "start": "node dist/index.js", + "volt": "volt" + }, + "type": "module" +} diff --git a/examples/with-tool-routing/src/index.ts b/examples/with-tool-routing/src/index.ts new file mode 100644 index 000000000..c1cd73893 --- /dev/null +++ b/examples/with-tool-routing/src/index.ts @@ -0,0 +1,95 @@ +import { Agent, VoltAgent, createTool } from "@voltagent/core"; +import { createPinoLogger } from "@voltagent/logger"; +import { honoServer } from "@voltagent/server-hono"; +import { z } from "zod"; + +const weatherTool = createTool({ + name: "get_weather", + description: "Get the current weather for a city", + parameters: z.object({ + location: z.string().describe("City name, e.g. Berlin"), + }), + tags: ["weather", "forecast"], + execute: async ({ location }) => { + return { + location, + temperatureC: 22, + condition: "sunny", + humidityPercent: 45, + }; + }, +}); + +const convertCurrencyTool = createTool({ + name: "convert_currency", + description: "Convert money between currencies using a sample rate table", + parameters: z.object({ + amount: z.number().describe("Amount to convert"), + from: z.string().describe("Source currency code, e.g. USD"), + to: z.string().describe("Target currency code, e.g. EUR"), + }), + tags: ["finance", "currency"], + execute: async ({ amount, from, to }) => { + const rates: Record = { + USD: 1, + EUR: 0.92, + GBP: 0.79, + TRY: 32.5, + }; + const fromCode = from.toUpperCase(); + const toCode = to.toUpperCase(); + const fromRate = rates[fromCode] ?? 1; + const toRate = rates[toCode] ?? 1; + const rate = toRate / fromRate; + + return { + amount, + from: fromCode, + to: toCode, + rate, + converted: Math.round(amount * rate * 100) / 100, + }; + }, +}); + +const timeZoneTool = createTool({ + name: "get_time_zone", + description: "Get the time zone offset for a city", + parameters: z.object({ + location: z.string().describe("City name"), + }), + tags: ["time", "timezone"], + execute: async ({ location }) => { + return { + location, + timeZone: "UTC+1", + }; + }, +}); + +const weatherPool = [weatherTool, timeZoneTool]; +const financePool = [convertCurrencyTool]; +const toolPool = [...weatherPool, ...financePool]; + +const logger = createPinoLogger({ + name: "with-tool-routing", + level: "info", +}); + +const agent = new Agent({ + name: "Tool Routing Agent", + instructions: + "You are a helpful assistant. Use tool_router when you need tools, and pass the user request as the query.", + model: "openai/gpt-4o-mini", + toolRouting: { + embedding: "openai/text-embedding-3-small", + pool: toolPool, + topK: 2, + }, +}); + +new VoltAgent({ + agents: { agent }, + server: honoServer(), + logger, +}); diff --git a/examples/with-tool-routing/tsconfig.json b/examples/with-tool-routing/tsconfig.json new file mode 100644 index 000000000..cee90c6f3 --- /dev/null +++ b/examples/with-tool-routing/tsconfig.json @@ -0,0 +1,14 @@ +{ + "compilerOptions": { + "target": "ES2022", + "module": "NodeNext", + "moduleResolution": "NodeNext", + "esModuleInterop": true, + "forceConsistentCasingInFileNames": true, + "strict": true, + "outDir": "dist", + "skipLibCheck": true + }, + "include": ["src"], + "exclude": ["node_modules", "dist"] +} diff --git a/examples/with-tools/src/index.ts b/examples/with-tools/src/index.ts index 1d3c8fd1b..36d08f695 100644 --- a/examples/with-tools/src/index.ts +++ b/examples/with-tools/src/index.ts @@ -1,5 +1,5 @@ import { openai } from "@ai-sdk/openai"; -import { Agent, AiSdkEmbeddingAdapter, Memory, VoltAgent } from "@voltagent/core"; +import { Agent, Memory, VoltAgent } from "@voltagent/core"; import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; import { createPinoLogger } from "@voltagent/logger"; import { honoServer } from "@voltagent/server-hono"; @@ -15,7 +15,7 @@ const logger = createPinoLogger({ const memory = new Memory({ storage: new LibSQLMemoryAdapter({}), - embedding: new AiSdkEmbeddingAdapter(openai.embeddingModel("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new LibSQLVectorAdapter(), }); diff --git a/examples/with-vector-search/README.md b/examples/with-vector-search/README.md index 671cc8d91..eaf5b855e 100644 --- a/examples/with-vector-search/README.md +++ b/examples/with-vector-search/README.md @@ -62,14 +62,13 @@ npm create voltagent-app@latest -- --example with-vector-search ## Snippet ```ts -import { openai } from "@ai-sdk/openai"; -import { Agent, AiSdkEmbeddingAdapter, Memory, VoltAgent } from "@voltagent/core"; +import { Agent, Memory, VoltAgent } from "@voltagent/core"; import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; import { honoServer } from "@voltagent/server-hono"; const memory = new Memory({ storage: new LibSQLMemoryAdapter(), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new LibSQLVectorAdapter(), }); diff --git a/examples/with-vector-search/src/index.ts b/examples/with-vector-search/src/index.ts index 15e79bbbb..64efdd8fd 100644 --- a/examples/with-vector-search/src/index.ts +++ b/examples/with-vector-search/src/index.ts @@ -1,5 +1,4 @@ -import { openai } from "@ai-sdk/openai"; -import { Agent, AiSdkEmbeddingAdapter, Memory, VoltAgent } from "@voltagent/core"; +import { Agent, Memory, VoltAgent } from "@voltagent/core"; import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; import { createPinoLogger } from "@voltagent/logger"; import { honoServer } from "@voltagent/server-hono"; @@ -10,10 +9,11 @@ const logger = createPinoLogger({ name: "with-vector-search", level: "info" }); // Memory configured with embeddings + vector DB const memory = new Memory({ storage: new LibSQLMemoryAdapter(), // default: file:./.voltagent/memory.db - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small"), { + embedding: { + model: "openai/text-embedding-3-small", // Optional caching/normalization settings normalize: false, - }), + }, vector: new LibSQLVectorAdapter(), enableCache: true, cacheSize: 1000, diff --git a/examples/with-voltagent-managed-memory/src/index.ts b/examples/with-voltagent-managed-memory/src/index.ts index 9c9a25d14..c7f29f0c8 100644 --- a/examples/with-voltagent-managed-memory/src/index.ts +++ b/examples/with-voltagent-managed-memory/src/index.ts @@ -1,5 +1,4 @@ -import { openai } from "@ai-sdk/openai"; -import { Agent, AiSdkEmbeddingAdapter, Memory, VoltAgent, VoltOpsClient } from "@voltagent/core"; +import { Agent, Memory, VoltAgent, VoltOpsClient } from "@voltagent/core"; import { createPinoLogger } from "@voltagent/logger"; import { honoServer } from "@voltagent/server-hono"; import { ManagedMemoryAdapter, ManagedMemoryVectorAdapter } from "@voltagent/voltagent-memory"; @@ -26,7 +25,7 @@ const agent = new Agent({ memory: new Memory({ storage: managedMemory, vector: managedVector, - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", }), }); diff --git a/examples/with-working-memory/README.md b/examples/with-working-memory/README.md index fa286ef27..859e7ed1f 100644 --- a/examples/with-working-memory/README.md +++ b/examples/with-working-memory/README.md @@ -62,8 +62,7 @@ npm create voltagent-app@latest -- --example with-working-memory ## Snippet ```ts -import { openai } from "@ai-sdk/openai"; -import { Agent, AiSdkEmbeddingAdapter, Memory, VoltAgent } from "@voltagent/core"; +import { Agent, Memory, VoltAgent } from "@voltagent/core"; import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; import { honoServer } from "@voltagent/server-hono"; import { z } from "zod"; @@ -75,7 +74,7 @@ const workingMemorySchema = z.object({ const memory = new Memory({ storage: new LibSQLMemoryAdapter(), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new LibSQLVectorAdapter(), workingMemory: { enabled: true, scope: "conversation", schema: workingMemorySchema }, }); diff --git a/examples/with-working-memory/src/index.ts b/examples/with-working-memory/src/index.ts index 08c160614..9181efcb1 100644 --- a/examples/with-working-memory/src/index.ts +++ b/examples/with-working-memory/src/index.ts @@ -1,10 +1,4 @@ -import { - Agent, - AiSdkEmbeddingAdapter, - Memory, - VoltAgent, - VoltAgentObservability, -} from "@voltagent/core"; +import { Agent, Memory, VoltAgent, VoltAgentObservability } from "@voltagent/core"; import { LibSQLMemoryAdapter, LibSQLObservabilityAdapter, diff --git a/packages/core/CHANGELOG.md b/packages/core/CHANGELOG.md index 1d5501e12..547dc2c56 100644 --- a/packages/core/CHANGELOG.md +++ b/packages/core/CHANGELOG.md @@ -3326,14 +3326,13 @@ ```typescript import { ManagedMemoryAdapter, ManagedMemoryVectorAdapter } from "@voltagent/voltagent-memory"; - import { AiSdkEmbeddingAdapter, Memory } from "@voltagent/core"; - import { openai } from "@ai-sdk/openai"; + import { Memory } from "@voltagent/core"; const memory = new Memory({ storage: new ManagedMemoryAdapter({ databaseName: "production-memory", }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new ManagedMemoryVectorAdapter({ databaseName: "production-memory", }), diff --git a/packages/core/scripts/generate-model-provider-registry.js b/packages/core/scripts/generate-model-provider-registry.js index 4d964c30b..6be3a0853 100644 --- a/packages/core/scripts/generate-model-provider-registry.js +++ b/packages/core/scripts/generate-model-provider-registry.js @@ -7,6 +7,7 @@ const API_URL = "https://models.dev/api.json"; const OUTPUT_DIR = path.resolve(__dirname, "../src/registries"); const REGISTRY_PATH = path.join(OUTPUT_DIR, "model-provider-registry.generated.ts"); const TYPES_PATH = path.join(OUTPUT_DIR, "model-provider-types.generated.ts"); +const EMBEDDING_TYPES_PATH = path.join(OUTPUT_DIR, "embedding-model-router-types.generated.ts"); const HEADER = `/** * THIS FILE IS AUTO-GENERATED - DO NOT EDIT @@ -20,6 +21,15 @@ const normalizeModelId = (id) => id.trim(); const isDeprecatedModel = (modelInfo) => Boolean(modelInfo && typeof modelInfo === "object" && modelInfo.status === "deprecated"); +const isEmbeddingModel = (modelId, modelInfo) => { + const id = normalizeModelId(modelId).toLowerCase(); + const family = + modelInfo && typeof modelInfo === "object" && typeof modelInfo.family === "string" + ? modelInfo.family.toLowerCase() + : ""; + return id.includes("embed") || id.includes("embedding") || family.includes("embed"); +}; + const formatStringLiteral = (value) => `'${String(value).replace(/\\/g, "\\\\").replace(/'/g, "\\'")}'`; @@ -39,6 +49,7 @@ async function run() { const registry = {}; const providerModels = {}; + const providerEmbeddingModels = {}; for (const [providerId, info] of providers) { const normalizedId = normalizeProviderId(info.id || providerId); @@ -58,6 +69,18 @@ async function run() { .sort(); providerModels[normalizedId] = models; + + const embeddingModels = Object.entries(info.models) + .filter( + ([modelId, modelInfo]) => + !isDeprecatedModel(modelInfo) && isEmbeddingModel(modelId, modelInfo), + ) + .map(([modelId]) => normalizeModelId(modelId)) + .sort(); + + if (embeddingModels.length) { + providerEmbeddingModels[normalizedId] = embeddingModels; + } } const registryContent = `${HEADER} @@ -97,15 +120,37 @@ export type ModelRouterModelId = export type ModelForProvider

= ProviderModelsMap[P][number]; `; + const embeddingModelLines = Object.entries(providerEmbeddingModels).map( + ([providerId, models]) => { + const modelLines = models.map((modelId) => ` ${formatStringLiteral(modelId)},`).join("\n"); + return ` readonly ${formatStringLiteral(providerId)}: readonly [\n${modelLines}\n ];`; + }, + ); + + const embeddingTypesContent = `${HEADER} +export type EmbeddingModelsMap = { +${embeddingModelLines.join("\n")} +}; + +export type EmbeddingProviderId = keyof EmbeddingModelsMap; + +export type EmbeddingRouterModelId = + | { + [P in EmbeddingProviderId]: \`\${P}/\${EmbeddingModelsMap[P][number]}\`; + }[EmbeddingProviderId] + | (string & {}); +`; + fs.mkdirSync(OUTPUT_DIR, { recursive: true }); fs.writeFileSync(REGISTRY_PATH, registryContent, "utf8"); fs.writeFileSync(TYPES_PATH, typesContent, "utf8"); + fs.writeFileSync(EMBEDDING_TYPES_PATH, embeddingTypesContent, "utf8"); console.info( `Generated ${path.relative(process.cwd(), REGISTRY_PATH)} and ${path.relative( process.cwd(), TYPES_PATH, - )}`, + )} and ${path.relative(process.cwd(), EMBEDDING_TYPES_PATH)}`, ); } diff --git a/packages/core/src/agent/agent.ts b/packages/core/src/agent/agent.ts index e10af0a0f..4afd0bebd 100644 --- a/packages/core/src/agent/agent.ts +++ b/packages/core/src/agent/agent.ts @@ -26,7 +26,7 @@ import { type FinishReason, type InferGenerateOutput, type LanguageModelUsage, - type Output, + Output, type Warning, consumeStream, convertToModelMessages, @@ -52,12 +52,24 @@ import { type ObservabilityFlushState, flushObservability } from "../observabili import { AgentRegistry } from "../registries/agent-registry"; import { ModelProviderRegistry } from "../registries/model-provider-registry"; import type { BaseRetriever } from "../retriever/retriever"; -import type { Tool, ToolExecutionResult, Toolkit, VercelTool } from "../tool"; +import type { ProviderTool, Tool, ToolExecutionResult, Toolkit, VercelTool } from "../tool"; import { createTool } from "../tool"; +import { isProviderTool } from "../tool/manager"; import { ToolManager } from "../tool/manager"; +import { createToolRouter, isToolRouter } from "../tool/routing"; +import { TOOL_ROUTER_SYMBOL } from "../tool/routing/constants"; +import type { + ToolArgumentResolver, + ToolRouter, + ToolRouterCandidate, + ToolRouterInput, + ToolRouterResult, + ToolRouterSelection, +} from "../tool/routing/types"; import { randomUUID } from "../utils/id"; import { convertModelMessagesToUIMessages } from "../utils/message-converter"; import { NodeType, createNodeId } from "../utils/node-utils"; +import { zodSchemaToJsonUI } from "../utils/toolParser"; import { convertUsage } from "../utils/usage-converter"; import { normalizeFinishUsageStream, resolveFinishUsage } from "../utils/usage-normalizer"; import type { Voice } from "../voice"; @@ -66,6 +78,7 @@ import type { VoltOpsClient } from "../voltops/client"; import type { PromptContent, PromptHelper } from "../voltops/types"; import { buildToolErrorResult } from "./error-utils"; import { + ToolDeniedError, createAbortError, createBailError, createVoltAgentError, @@ -94,8 +107,9 @@ import { P, match } from "ts-pattern"; import type { StopWhen } from "../ai-types"; import type { SamplingPolicy } from "../eval/runtime"; import type { ConversationStepRecord } from "../memory/types"; +import type { ToolRoutingConfig } from "../tool/routing/types"; import { applySummarization } from "./apply-summarization"; -import { FORCED_TOOL_CHOICE_CONTEXT_KEY } from "./context-keys"; +import { AGENT_REF_CONTEXT_KEY, FORCED_TOOL_CHOICE_CONTEXT_KEY } from "./context-keys"; import { ConversationBuffer } from "./conversation-buffer"; import { type NormalizedInputGuardrail, @@ -138,6 +152,8 @@ import type { AgentModelValue, AgentOptions, AgentSummarizationOptions, + AgentToolRoutingState, + ApiToolInfo, DynamicValue, DynamicValueOptions, InputGuardrail, @@ -397,6 +413,10 @@ export interface BaseGenerationOptions extends Partial { // Tools (can provide additional tools dynamically) tools?: (Tool | Toolkit)[]; + /** + * Optional per-call tool routing override. + */ + toolRouting?: ToolRoutingConfig | false; // Hooks (can override agent hooks) hooks?: AgentHooks; @@ -481,6 +501,7 @@ export class Agent { private readonly summarization?: AgentSummarizationOptions | false; private defaultObservability?: VoltAgentObservability; private readonly toolManager: ToolManager; + private readonly toolPoolManager: ToolManager; private readonly subAgentManager: SubAgentManager; private readonly voltOpsClient?: VoltOpsClient; private readonly prompts?: PromptHelper; @@ -494,6 +515,10 @@ export class Agent { private readonly observabilityAuthWarningState: ObservabilityFlushState = { authWarningLogged: false, }; + private toolRouting?: ToolRoutingConfig | false; + private toolRoutingConfigured: boolean; + private toolRoutingExposedNames: Set = new Set(); + private toolRoutingPoolExplicit = false; constructor(options: AgentOptions) { this.id = options.id || options.name; @@ -521,6 +546,8 @@ export class Agent { this.inputMiddlewares = normalizeInputMiddlewareList(options.inputMiddlewares || []); this.outputMiddlewares = normalizeOutputMiddlewareList(options.outputMiddlewares || []); this.maxMiddlewareRetries = options.maxMiddlewareRetries ?? 0; + this.toolRoutingConfigured = options.toolRouting !== undefined; + this.toolRouting = options.toolRouting ?? AgentRegistry.getInstance().getGlobalToolRouting(); // Initialize logger - always use LoggerProxy for consistency // If external logger is provided, it will be used by LoggerProxy @@ -560,6 +587,8 @@ export class Agent { if (options.toolkits) { this.toolManager.addItems(options.toolkits); } + this.toolPoolManager = new ToolManager([], this.logger); + this.applyToolRoutingConfig(this.toolRouting); // Initialize sub-agent manager this.subAgentManager = new SubAgentManager( @@ -2993,6 +3022,7 @@ export class Agent { agentId: this.id, agentName: this.name, }); + systemContext.set(AGENT_REF_CONTEXT_KEY, this); const elicitationHandler = options?.elicitation ?? options?.parentOperationContext?.elicitation; @@ -4568,7 +4598,17 @@ export class Agent { const preparedStaticTools = this.toolManager.prepareToolsForExecution(createToolExecuteFunction); - return { ...preparedStaticTools, ...preparedDynamicTools }; + const toolRouting = this.resolveToolRouting(options); + if (!toolRouting) { + return { ...preparedStaticTools, ...preparedDynamicTools }; + } + + const exposedNames = this.getToolRoutingExposedNames(toolRouting); + const filteredStaticTools = Object.fromEntries( + Object.entries(preparedStaticTools).filter(([name]) => exposedNames.has(name)), + ); + + return { ...filteredStaticTools, ...preparedDynamicTools }; } /** @@ -4614,6 +4654,7 @@ export class Agent { abortSignal: abortSignal, }, }; + executionOptions.hooks = hooks; // Event tracking now handled by OpenTelemetry spans const toolTags = (tool as { tags?: string[] | undefined }).tags; @@ -4840,6 +4881,472 @@ export class Agent { }; } + /** + * Internal: execute a tool router with access to the agent runtime. + */ + public async __executeToolRouter(params: { + router: ToolRouter; + input: ToolRouterInput; + options?: ToolExecuteOptions; + }): Promise { + const { router, input, options } = params; + if (!options) { + throw new Error("Tool router execution requires tool options."); + } + + const metadata = (router as ToolRouter)[TOOL_ROUTER_SYMBOL]; + if (!metadata) { + throw new Error("Tool router metadata is missing."); + } + + const oc = options as OperationContext; + const toolRouting = this.resolveToolRouting(); + const routingConfig = toolRouting && typeof toolRouting === "object" ? toolRouting : undefined; + const mode = metadata.mode ?? routingConfig?.mode; + const effectiveMode = mode ?? "agent"; + const executionModel = metadata.executionModel ?? routingConfig?.executionModel; + const topK = Math.max(1, input.topK ?? metadata.topK ?? routingConfig?.topK ?? 1); + const parallel = metadata.parallel ?? routingConfig?.parallel ?? true; + + const candidates = this.buildToolRouterCandidates(); + const parentToolSpan = oc.systemContext.get("parentToolSpan") as Span | undefined; + const selectionSpanAttributes = { + "tool.name": router.name, + "tool.description": router.description, + "tool.router.name": router.name, + "tool.router.query": input.query, + "tool.router.candidates": candidates.length, + "tool.router.top_k": topK, + input: input.query, + }; + const selectionSpan = parentToolSpan + ? oc.traceContext.createChildSpanWithParent( + parentToolSpan, + `tool.router.selection:${router.name}`, + "tool", + { + label: `Tool Router Selection: ${router.name}`, + attributes: selectionSpanAttributes, + }, + ) + : oc.traceContext.createChildSpan(`tool.router.selection:${router.name}`, "tool", { + label: `Tool Router Selection: ${router.name}`, + attributes: selectionSpanAttributes, + }); + + const context = { + agentId: this.id, + agentName: this.name, + operationContext: oc, + routerName: router.name, + parentSpan: selectionSpan, + }; + + let selections: ToolRouterSelection[] = []; + try { + selections = await oc.traceContext.withSpan(selectionSpan, () => + metadata.strategy.select({ + query: input.query, + tools: candidates, + topK, + context, + }), + ); + oc.traceContext.endChildSpan(selectionSpan, "completed", { + output: selections, + attributes: { + "tool.router.selection.count": selections.length, + "tool.router.selection.names": safeStringify( + selections.map((selection) => selection.name), + ), + }, + }); + oc.logger.debug("Tool router selections computed", { + router: router.name, + query: input.query, + selections: safeStringify(selections), + }); + } catch (error) { + oc.traceContext.endChildSpan(selectionSpan, "error", { + output: { error: error instanceof Error ? error.message : String(error) }, + }); + throw error; + } + + const hooks = + ((options as { hooks?: AgentHooks }).hooks as AgentHooks | undefined) ?? + this.getMergedHooks(); + + const executeSelection = async ( + selection: ToolRouterSelection, + ): Promise<{ toolName: string; toolCallId?: string; output?: unknown; error?: string }> => { + const tool = this.toolPoolManager.getToolByName(selection.name); + if (!tool) { + return { + toolName: selection.name, + error: "Tool not found in pool.", + }; + } + + if (isProviderTool(tool)) { + return this.executeProviderToolViaRouter({ + tool, + query: input.query, + oc, + hooks, + executionModel, + }); + } + + const toolCallId = randomUUID(); + const executionOptions: ToolExecuteOptions = { + ...oc, + toolContext: { + name: tool.name, + callId: toolCallId, + messages: [], + abortSignal: oc.abortController.signal, + }, + }; + executionOptions.hooks = hooks; + + try { + const args = await this.resolveRoutedToolArgs({ + tool, + query: input.query, + mode: effectiveMode, + resolver: metadata.resolver, + executionModel, + oc, + }); + await this.ensureToolApproval(tool, args, executionOptions, toolCallId); + + const execute = this.createToolExecutionFactory(oc, hooks)(tool); + const output = await execute(args, { + toolCallId, + messages: executionOptions.toolContext?.messages ?? [], + abortSignal: executionOptions.toolContext?.abortSignal, + }); + return { + toolName: tool.name, + toolCallId, + output, + }; + } catch (error) { + return { + toolName: tool.name, + toolCallId, + error: error instanceof Error ? error.message : String(error), + }; + } + }; + + const results: ToolRouterResult["results"] = []; + if (parallel) { + const resolved = await Promise.all( + selections.map((selection) => executeSelection(selection)), + ); + results.push(...resolved); + } else { + for (const selection of selections) { + results.push(await executeSelection(selection)); + } + } + + return { + query: input.query, + selections, + results, + }; + } + + private buildToolRouterCandidates(): ToolRouterCandidate[] { + return this.toolPoolManager + .getAllTools() + .filter((tool) => !isToolRouter(tool)) + .map((tool) => { + const parameters = isProviderTool(tool) + ? tool.args + : zodSchemaToJsonUI((tool as Tool).parameters); + const tags = "tags" in tool ? (tool as { tags?: string[] }).tags : undefined; + return { + name: tool.name, + description: tool.description || "", + tags, + parameters, + tool, + }; + }); + } + + private async resolveRoutedToolArgs(params: { + tool: Tool; + query: string; + mode: "agent" | "resolver"; + resolver?: ToolArgumentResolver; + executionModel?: AgentModelValue; + oc: OperationContext; + }): Promise> { + const { tool, query, mode, resolver, executionModel, oc } = params; + if (mode === "resolver") { + if (!resolver) { + throw new Error("Tool router resolver mode requires a resolver function."); + } + return await resolver({ + query, + tool, + context: { agentId: this.id, agentName: this.name, operationContext: oc }, + }); + } + + const schema = zodSchemaToJsonUI(tool.parameters); + const prompt = [ + "Generate JSON arguments for the tool based on the user request.", + `Tool name: ${tool.name}`, + tool.description ? `Tool description: ${tool.description}` : "", + schema ? `Tool schema: ${safeStringify(schema)}` : "", + `User request: ${query}`, + "Return only a JSON object that matches the tool schema.", + ] + .filter(Boolean) + .join("\n"); + + const result = await this.runInternalGenerateText({ + oc, + modelValue: executionModel, + messages: [ + { + role: "system", + content: "You generate tool arguments that strictly match the provided schema.", + }, + { role: "user", content: prompt }, + ], + output: Output.object({ schema: tool.parameters }), + toolChoice: "none", + temperature: 0, + }); + + return (result.output ?? {}) as Record; + } + + private async executeProviderToolViaRouter(params: { + tool: ProviderTool; + query: string; + oc: OperationContext; + hooks: AgentHooks; + executionModel?: AgentModelValue; + }): Promise<{ toolName: string; toolCallId?: string; output?: unknown; error?: string }> { + const { tool, query, oc, hooks, executionModel } = params; + oc.logger.info("Tool router using provider tool fallback", { + toolName: tool.name, + query, + }); + const toolCallId = randomUUID(); + const executionOptions: ToolExecuteOptions = { + ...oc, + toolContext: { + name: tool.name, + callId: toolCallId, + messages: [], + abortSignal: oc.abortController.signal, + }, + }; + executionOptions.hooks = hooks; + + const needsApproval = (tool as { needsApproval?: Tool["needsApproval"] }) + .needsApproval; + if (needsApproval === true) { + throw new ToolDeniedError({ + toolName: tool.name, + message: `Tool ${tool.name} requires approval.`, + code: "TOOL_FORBIDDEN", + httpStatus: 403, + }); + } + + const tools: Record = { + [tool.name]: tool, + }; + + const result = await this.runInternalGenerateText({ + oc, + modelValue: executionModel, + messages: [ + { + role: "system", + content: "Call the required tool with appropriate arguments to satisfy the request.", + }, + { role: "user", content: query }, + ], + tools, + toolChoice: { type: "tool", toolName: tool.name }, + temperature: this.temperature ?? 0, + }); + + const { toolCalls, toolResults } = this.collectToolDataFromResult(result); + const toolCall = toolCalls.find((call) => call.toolName === tool.name); + const toolResult = toolResults.find( + (res) => res.toolName === tool.name && (!toolCall || res.toolCallId === toolCall.toolCallId), + ); + + if (toolCall?.toolCallId && executionOptions.toolContext) { + executionOptions.toolContext.callId = toolCall.toolCallId; + } + + if (toolCall) { + try { + await this.ensureToolApproval( + tool, + toolCall.input as Record, + executionOptions, + toolCall.toolCallId ?? toolCallId, + ); + } catch (error) { + if (isToolDeniedError(error)) { + return { + toolName: tool.name, + toolCallId: toolCall.toolCallId ?? toolCallId, + error: error.message, + }; + } + throw error; + } + } + + if (toolCall) { + await hooks.onToolStart?.({ + agent: this, + tool: tool as any, + args: toolCall.input, + context: oc, + options: executionOptions, + }); + } + + if (!toolResult) { + return { + toolName: tool.name, + toolCallId: toolCall?.toolCallId ?? toolCallId, + error: "Provider tool did not return a result.", + }; + } + + const toolError = + toolResult.output && typeof toolResult.output === "object" && "error" in toolResult.output + ? String((toolResult.output as { error?: unknown }).error ?? "Tool error") + : undefined; + const hookError = toolError + ? createVoltAgentError(toolError, { stage: "tool_execution" }) + : undefined; + + await hooks.onToolEnd?.({ + agent: this, + tool: tool as any, + output: toolError ? undefined : toolResult.output, + error: hookError, + context: oc, + options: executionOptions, + }); + + if (toolError) { + return { + toolName: tool.name, + toolCallId: toolResult.toolCallId ?? toolCallId, + error: toolError, + }; + } + + return { + toolName: tool.name, + toolCallId: toolResult.toolCallId ?? toolCallId, + output: toolResult.output, + }; + } + + private async ensureToolApproval( + tool: Tool | ProviderTool, + args: Record, + options: ToolExecuteOptions, + toolCallId: string, + ): Promise { + const needsApproval = (tool as { needsApproval?: Tool["needsApproval"] }) + .needsApproval; + if (!needsApproval) { + return; + } + + const requiresApproval = + typeof needsApproval === "function" + ? await needsApproval(args as any, { + toolCallId, + messages: (options.toolContext?.messages ?? []) as ModelMessage[], + experimental_context: undefined, + }) + : needsApproval; + + if (requiresApproval) { + throw new ToolDeniedError({ + toolName: tool.name, + message: `Tool ${tool.name} requires approval.`, + code: "TOOL_FORBIDDEN", + httpStatus: 403, + }); + } + } + + private async runInternalGenerateText(params: { + oc: OperationContext; + modelValue?: AgentModelValue; + messages: ModelMessage[]; + tools?: ToolSet; + output?: OutputSpec; + toolChoice?: ToolChoice>; + temperature?: number; + }): Promise> { + const { oc, modelValue, messages, tools, output, toolChoice, temperature } = params; + const model = await this.resolveModel(modelValue ?? this.model, oc); + const modelName = this.getModelName(model); + + const llmSpan = this.createLLMSpan(oc, { + operation: "generateText", + modelName, + isStreaming: false, + messages: messages.map((msg) => ({ role: msg.role, content: msg.content })), + tools, + callOptions: { + temperature, + }, + }); + const finalizeLLMSpan = this.createLLMSpanFinalizer(llmSpan); + + try { + const response = await oc.traceContext.withSpan(llmSpan, () => + generateText({ + model, + messages, + tools, + output, + toolChoice, + temperature, + maxRetries: 0, + stopWhen: stepCountIs(1), + abortSignal: oc.abortController.signal, + }), + ); + + const resolvedUsage = response.usage ? await Promise.resolve(response.usage) : undefined; + finalizeLLMSpan(SpanStatusCode.OK, { + usage: resolvedUsage, + finishReason: response.finishReason, + }); + + return response; + } catch (error) { + finalizeLLMSpan(SpanStatusCode.ERROR, { message: (error as Error).message }); + throw error; + } + } + /** * Create step handler for memory and hooks */ @@ -5162,6 +5669,7 @@ export class Agent { onToolEnd: async (...args) => { await options.hooks?.onToolEnd?.(...args); await this.hooks.onToolEnd?.(...args); + return undefined; }, onStepFinish: async (...args) => { await options.hooks?.onStepFinish?.(...args); @@ -5461,6 +5969,33 @@ export class Agent { const activeMemory = this.getMemory(); const memoryInstance: Memory | undefined = activeMemory || undefined; + const toolRoutingConfig = + this.toolRouting && typeof this.toolRouting === "object" ? this.toolRouting : undefined; + const toolRoutingState: AgentToolRoutingState | undefined = toolRoutingConfig + ? (() => { + const routerTools = this.toolManager + .getAllBaseTools() + .filter((tool) => isToolRouter(tool)); + const routerApiTools = + routerTools.length > 0 + ? new ToolManager(routerTools, this.logger).getToolsForApi() + : []; + const routerNames = new Set(routerApiTools.map((tool) => tool.name)); + const poolApiTools = this.toolPoolManager + .getToolsForApi() + .filter((tool) => !routerNames.has(tool.name)); + const exposeApiTools = + toolRoutingConfig.expose && toolRoutingConfig.expose.length > 0 + ? new ToolManager(toolRoutingConfig.expose, this.logger).getToolsForApi() + : []; + + return { + routers: routerApiTools.length > 0 ? routerApiTools : undefined, + expose: exposeApiTools.length > 0 ? exposeApiTools : undefined, + pool: poolApiTools.length > 0 ? poolApiTools : undefined, + }; + })() + : undefined; return { id: this.id, @@ -5471,10 +6006,22 @@ export class Agent { model: this.getModelName(), node_id: createNodeId(NodeType.AGENT, this.id), - tools: this.toolManager.getAllBaseTools().map((tool) => ({ - ...tool, - node_id: createNodeId(NodeType.TOOL, tool.name, this.id), - })), + tools: (() => { + const merged = new Map(); + for (const tool of [ + ...this.toolManager.getAllTools(), + ...this.toolPoolManager.getAllTools(), + ]) { + if (!merged.has(tool.name)) { + merged.set(tool.name, tool); + } + } + return Array.from(merged.values()).map((tool) => ({ + ...tool, + node_id: createNodeId(NodeType.TOOL, tool.name, this.id), + })); + })(), + toolRouting: toolRoutingState, subAgents: this.subAgentManager.getSubAgentDetails().map((subAgent) => ({ ...subAgent, @@ -5526,6 +6073,9 @@ export class Agent { added: (Tool | Toolkit | VercelTool)[]; } { this.toolManager.addItems(tools); + if (this.toolRouting && !this.toolRoutingPoolExplicit) { + this.toolPoolManager.addItems(tools); + } return { added: tools }; } @@ -5539,6 +6089,9 @@ export class Agent { for (const name of toolNames) { if (this.toolManager.removeTool(name)) { removed.push(name); + if (this.toolRouting && !this.toolRoutingPoolExplicit) { + this.toolPoolManager.removeTool(name); + } } } @@ -5557,6 +6110,9 @@ export class Agent { */ public removeToolkit(toolkitName: string): boolean { const result = this.toolManager.removeToolkit(toolkitName); + if (result && this.toolRouting && !this.toolRoutingPoolExplicit) { + this.toolPoolManager.removeToolkit(toolkitName); + } if (result) { this.logger.debug(`Removed toolkit: ${toolkitName}`); @@ -5579,6 +6135,9 @@ export class Agent { sourceAgent: this as any, }); this.toolManager.addStandaloneTool(delegateTool); + if (this.toolRouting && !this.toolRoutingPoolExplicit) { + this.toolPoolManager.addStandaloneTool(delegateTool); + } } } @@ -5591,6 +6150,9 @@ export class Agent { // Remove delegate tool if no sub-agents left if (this.subAgentManager.getSubAgents().length === 0) { this.toolManager.removeTool("delegate_task"); + if (this.toolRouting && !this.toolRoutingPoolExplicit) { + this.toolPoolManager.removeTool("delegate_task"); + } } } @@ -5605,7 +6167,15 @@ export class Agent { * Get tools for API */ public getToolsForApi() { - return this.toolManager.getToolsForApi(); + const exposed = this.toolManager.getToolsForApi(); + const pooled = this.toolPoolManager.getToolsForApi(); + const merged = new Map(); + for (const tool of [...exposed, ...pooled]) { + if (!merged.has(tool.name)) { + merged.set(tool.name, tool); + } + } + return Array.from(merged.values()); } /** @@ -5676,6 +6246,110 @@ export class Agent { this.memoryManager.setMemory(memory); } + /** + * Internal: apply a default tool routing config when none was configured explicitly. + */ + public __setDefaultToolRouting(toolRouting?: ToolRoutingConfig): void { + if (this.toolRoutingConfigured) { + return; + } + this.toolRouting = toolRouting; + this.toolRoutingConfigured = true; + this.applyToolRoutingConfig(this.toolRouting); + } + + private applyToolRoutingConfig(toolRouting?: ToolRoutingConfig | false): void { + if (!toolRouting) { + this.toolRoutingPoolExplicit = false; + return; + } + + this.toolRoutingPoolExplicit = Object.prototype.hasOwnProperty.call(toolRouting, "pool"); + this.toolRoutingExposedNames = new Set(); + + const routers = this.resolveToolRoutingRouters(toolRouting); + if (routers.length > 0) { + this.toolManager.addItems(routers); + } + routers.forEach((router) => this.toolRoutingExposedNames.add(router.name)); + + const existingRouters = this.toolManager + .getAllBaseTools() + .filter((tool) => isToolRouter(tool)) as ToolRouter[]; + existingRouters.forEach((router) => this.toolRoutingExposedNames.add(router.name)); + + if (toolRouting.expose && toolRouting.expose.length > 0) { + this.toolManager.addItems(toolRouting.expose); + const exposedManager = new ToolManager(toolRouting.expose, this.logger); + exposedManager.getAllToolNames().forEach((name) => this.toolRoutingExposedNames.add(name)); + } + if (toolRouting.pool && toolRouting.pool.length > 0) { + this.toolPoolManager.addItems(toolRouting.pool); + } else if (!this.toolRoutingPoolExplicit) { + const autoPool = this.toolManager.getAllTools().filter((tool) => !isToolRouter(tool)); + if (autoPool.length > 0) { + this.toolPoolManager.addItems(autoPool); + } + } + } + + private resolveToolRoutingRouters(toolRouting: ToolRoutingConfig): ToolRouter[] { + if (toolRouting.routers && toolRouting.routers.length > 0) { + return toolRouting.routers; + } + + if (toolRouting.embedding) { + const embeddingConfig = + typeof toolRouting.embedding === "object" && + toolRouting.embedding !== null && + "model" in toolRouting.embedding + ? (toolRouting.embedding as { topK?: number }) + : undefined; + const router = createToolRouter({ + name: "tool_router", + description: "Routes requests to the most relevant tool and executes it.", + embedding: toolRouting.embedding, + mode: toolRouting.mode, + executionModel: toolRouting.executionModel, + topK: toolRouting.topK ?? embeddingConfig?.topK, + parallel: toolRouting.parallel, + }); + return [router]; + } + + return []; + } + + private resolveToolRouting( + options?: BaseGenerationOptions, + ): ToolRoutingConfig | false | undefined { + if (options?.toolRouting !== undefined) { + return options.toolRouting; + } + return this.toolRouting; + } + + private getToolRoutingExposedNames(toolRouting: ToolRoutingConfig): Set { + if (toolRouting === this.toolRouting && this.toolRoutingExposedNames.size > 0) { + return this.toolRoutingExposedNames; + } + + const exposedNames = new Set(); + toolRouting.routers?.forEach((router) => exposedNames.add(router.name)); + + const existingRouters = this.toolManager + .getAllBaseTools() + .filter((tool) => isToolRouter(tool)) as ToolRouter[]; + existingRouters.forEach((router) => exposedNames.add(router.name)); + + if (toolRouting.expose && toolRouting.expose.length > 0) { + const exposedManager = new ToolManager(toolRouting.expose, this.logger); + exposedManager.getAllToolNames().forEach((name) => exposedNames.add(name)); + } + + return exposedNames; + } + /** * Convert this agent into a tool that can be used by other agents. * This enables supervisor/coordinator patterns where one agent can delegate diff --git a/packages/core/src/agent/context-keys.ts b/packages/core/src/agent/context-keys.ts index 684027f33..6360ad690 100644 --- a/packages/core/src/agent/context-keys.ts +++ b/packages/core/src/agent/context-keys.ts @@ -1 +1,2 @@ export const FORCED_TOOL_CHOICE_CONTEXT_KEY = Symbol("forcedToolChoice"); +export const AGENT_REF_CONTEXT_KEY = Symbol("agentRef"); diff --git a/packages/core/src/agent/hooks/index.ts b/packages/core/src/agent/hooks/index.ts index 2d7af227c..f4f47a8e2 100644 --- a/packages/core/src/agent/hooks/index.ts +++ b/packages/core/src/agent/hooks/index.ts @@ -214,7 +214,7 @@ const defaultHooks: Required = { onHandoff: async (_args: OnHandoffHookArgs) => {}, onHandoffComplete: async (_args: OnHandoffCompleteHookArgs) => {}, onToolStart: async (_args: OnToolStartHookArgs) => {}, - onToolEnd: async (_args: OnToolEndHookArgs) => {}, + onToolEnd: async (_args: OnToolEndHookArgs) => undefined, onPrepareMessages: async (_args: OnPrepareMessagesHookArgs) => ({}), onPrepareModelMessages: async (_args: OnPrepareModelMessagesHookArgs) => ({}), onError: async (_args: OnErrorHookArgs) => {}, diff --git a/packages/core/src/agent/types.ts b/packages/core/src/agent/types.ts index 520c89fba..1e2a4aa32 100644 --- a/packages/core/src/agent/types.ts +++ b/packages/core/src/agent/types.ts @@ -12,7 +12,8 @@ import type { StopWhen } from "../ai-types"; import type { LanguageModel, TextStreamPart, UIMessage } from "ai"; import type { Memory } from "../memory"; import type { BaseRetriever } from "../retriever/retriever"; -import type { Tool, Toolkit, VercelTool } from "../tool"; +import type { ProviderTool, Tool, Toolkit, VercelTool } from "../tool"; +import type { ToolRoutingConfig } from "../tool/routing/types"; import type { StreamEvent } from "../utils/streams"; import type { Voice } from "../voice/types"; import type { VoltOpsClient } from "../voltops/client"; @@ -56,6 +57,12 @@ export interface ApiToolInfo { parameters?: any; } +export interface AgentToolRoutingState { + routers?: ApiToolInfo[]; + expose?: ApiToolInfo[]; + pool?: ApiToolInfo[]; +} + export type AgentFeedbackOptions = { key?: string; feedbackConfig?: VoltOpsFeedbackConfig | null; @@ -75,9 +82,9 @@ export type AgentFeedbackMetadata = { /** * Tool with node_id for agent state */ -export interface ToolWithNodeId extends BaseTool { +export type ToolWithNodeId = (BaseTool | ProviderTool) & { node_id: string; -} +}; export interface AgentScorerState { key: string; @@ -151,6 +158,7 @@ export interface AgentFullState { model: string; node_id: string; tools: ToolWithNodeId[]; + toolRouting?: AgentToolRoutingState; subAgents: SubAgentStateData[]; memory: AgentMemoryState; scorers?: AgentScorerState[]; @@ -564,6 +572,7 @@ export type AgentOptions = { // Tools & Memory tools?: (Tool | Toolkit | VercelTool)[] | DynamicValue<(Tool | Toolkit)[]>; toolkits?: Toolkit[]; + toolRouting?: ToolRoutingConfig | false; memory?: Memory | false; summarization?: AgentSummarizationOptions | false; diff --git a/packages/core/src/index.ts b/packages/core/src/index.ts index 2ed88b2b5..81b916756 100644 --- a/packages/core/src/index.ts +++ b/packages/core/src/index.ts @@ -147,6 +147,7 @@ export { export { InMemoryStorageAdapter } from "./memory/adapters/storage/in-memory"; export { InMemoryVectorAdapter } from "./memory/adapters/vector/in-memory"; export { AiSdkEmbeddingAdapter } from "./memory/adapters/embedding/ai-sdk"; +export type { EmbeddingModelReference } from "./memory/adapters/embedding/types"; export type { WorkingMemoryScope, WorkingMemoryConfig, @@ -155,6 +156,7 @@ export type { export * from "./agent/providers"; export { ModelProviderRegistry, + type EmbeddingModelFactory, type LanguageModelFactory, type ModelProvider, type ModelProviderEntry, @@ -166,6 +168,7 @@ export type { ProviderId, ProviderModelsMap, } from "./registries/model-provider-types.generated"; +export type { EmbeddingRouterModelId } from "./registries/embedding-model-router-types"; export * from "./events/types"; export type { AgentOptions, diff --git a/packages/core/src/memory/adapters/embedding/ai-sdk.spec.ts b/packages/core/src/memory/adapters/embedding/ai-sdk.spec.ts index 9ac11e621..52782aa35 100644 --- a/packages/core/src/memory/adapters/embedding/ai-sdk.spec.ts +++ b/packages/core/src/memory/adapters/embedding/ai-sdk.spec.ts @@ -1,5 +1,6 @@ import type { EmbedManyResult, EmbedResult, EmbeddingModel } from "ai"; import { beforeEach, describe, expect, it, vi } from "vitest"; +import { ModelProviderRegistry } from "../../../registries/model-provider-registry"; import { AiSdkEmbeddingAdapter } from "./ai-sdk"; // Mock the AI SDK @@ -9,7 +10,7 @@ vi.mock("ai", () => ({ })); describe("AiSdkEmbeddingAdapter", () => { - let mockModel: EmbeddingModel; + let mockModel: Exclude; let adapter: AiSdkEmbeddingAdapter; beforeEach(() => { @@ -19,7 +20,7 @@ describe("AiSdkEmbeddingAdapter", () => { modelId: "test-model", provider: "test-provider", doEmbed: vi.fn(), - } as unknown as EmbeddingModel; + } as unknown as Exclude; adapter = new AiSdkEmbeddingAdapter(mockModel, { maxBatchSize: 2, @@ -177,4 +178,39 @@ describe("AiSdkEmbeddingAdapter", () => { expect(adapter.getModelName()).toBe("test-model"); }); }); + + describe("model resolution", () => { + it("should resolve provider-qualified model strings via the registry", async () => { + const mockEmbedding = [0.1, 0.2]; + const { embed } = await import("ai"); + vi.mocked(embed).mockResolvedValue({ + value: "test text", + embedding: mockEmbedding, + usage: { tokens: 10 }, + } as EmbedResult); + + const resolvedModel = { + modelId: "text-embedding-3-small", + provider: "openai", + doEmbed: vi.fn(), + } as unknown as Exclude; + + const registry = ModelProviderRegistry.getInstance(); + const resolveSpy = vi + .spyOn(registry, "resolveEmbeddingModel") + .mockResolvedValue(resolvedModel); + + const stringAdapter = new AiSdkEmbeddingAdapter("openai/text-embedding-3-small"); + await stringAdapter.embed("test text"); + + expect(resolveSpy).toHaveBeenCalledWith("openai/text-embedding-3-small"); + expect(embed).toHaveBeenCalledWith({ + model: resolvedModel, + value: "test text", + }); + expect(stringAdapter.getModelName()).toBe("openai/text-embedding-3-small"); + + resolveSpy.mockRestore(); + }); + }); }); diff --git a/packages/core/src/memory/adapters/embedding/ai-sdk.ts b/packages/core/src/memory/adapters/embedding/ai-sdk.ts index 4d2ce387d..523862039 100644 --- a/packages/core/src/memory/adapters/embedding/ai-sdk.ts +++ b/packages/core/src/memory/adapters/embedding/ai-sdk.ts @@ -1,5 +1,7 @@ -import { type EmbeddingModel, embed, embedMany } from "ai"; -import type { EmbeddingAdapter, EmbeddingOptions } from "./types"; +import { embed, embedMany } from "ai"; +import type { EmbeddingModel } from "ai"; +import { ModelProviderRegistry } from "../../../registries/model-provider-registry"; +import type { EmbeddingAdapter, EmbeddingModelReference, EmbeddingOptions } from "./types"; /** * AI SDK Embedding Adapter @@ -10,11 +12,14 @@ export class AiSdkEmbeddingAdapter implements EmbeddingAdapter { private dimensions: number; private modelName: string; private options: EmbeddingOptions; + private modelResolvePromise?: Promise; - constructor(model: EmbeddingModel, options: EmbeddingOptions = {}) { - this.model = model; + constructor(model: EmbeddingModelReference, options: EmbeddingOptions = {}) { + const normalizedModel = typeof model === "string" ? model.trim() : model; + this.model = normalizedModel; // EmbeddingModel can be either a string or an object with modelId - this.modelName = typeof model === "string" ? model : model.modelId; + this.modelName = + typeof normalizedModel === "string" ? normalizedModel : normalizedModel.modelId; this.dimensions = 0; // Will be set after first embedding this.options = { maxBatchSize: options.maxBatchSize ?? 100, @@ -23,10 +28,45 @@ export class AiSdkEmbeddingAdapter implements EmbeddingAdapter { }; } + private async resolveModel(): Promise { + if (typeof this.model !== "string") { + return this.model; + } + if (this.modelResolvePromise) { + return this.modelResolvePromise; + } + + const trimmed = this.model.trim(); + if (!trimmed) { + return trimmed; + } + + const hasProviderPrefix = trimmed.includes("/") || trimmed.includes(":"); + if (!hasProviderPrefix) { + this.model = trimmed; + this.modelName = trimmed; + return trimmed; + } + + this.modelResolvePromise = ModelProviderRegistry.getInstance() + .resolveEmbeddingModel(trimmed) + .then((resolved) => { + this.model = resolved; + this.modelName = trimmed; + return resolved; + }) + .finally(() => { + this.modelResolvePromise = undefined; + }); + + return this.modelResolvePromise; + } + async embed(text: string): Promise { try { + const model = await this.resolveModel(); const result = await embed({ - model: this.model, + model, value: text, }); @@ -55,6 +95,7 @@ export class AiSdkEmbeddingAdapter implements EmbeddingAdapter { return []; } + const model = await this.resolveModel(); const maxBatchSize = this.options.maxBatchSize ?? 100; const embeddings: number[][] = []; @@ -64,7 +105,7 @@ export class AiSdkEmbeddingAdapter implements EmbeddingAdapter { try { const result = await embedMany({ - model: this.model, + model, values: batch, }); diff --git a/packages/core/src/memory/adapters/embedding/types.ts b/packages/core/src/memory/adapters/embedding/types.ts index 835de0922..4ecc669a1 100644 --- a/packages/core/src/memory/adapters/embedding/types.ts +++ b/packages/core/src/memory/adapters/embedding/types.ts @@ -1,3 +1,10 @@ +import type { EmbeddingModel } from "ai"; +import type { EmbeddingRouterModelId } from "../../../registries/embedding-model-router-types"; + +type EmbeddingModelInstance = Exclude; + +export type EmbeddingModelReference = EmbeddingRouterModelId | EmbeddingModelInstance; + /** * Embedding adapter interface for converting text to vectors */ diff --git a/packages/core/src/memory/index.spec-d.ts b/packages/core/src/memory/index.spec-d.ts index 58ec7f869..22d92d8be 100644 --- a/packages/core/src/memory/index.spec-d.ts +++ b/packages/core/src/memory/index.spec-d.ts @@ -96,6 +96,28 @@ describe("Memory V2 Type System", () => { expectTypeOf(memory).toMatchTypeOf(); }); + it("should accept embedding model string", () => { + const memory = new Memory({ + storage: mockStorageAdapter, + embedding: "openai/text-embedding-3-small", + }); + + expectTypeOf(memory).toMatchTypeOf(); + }); + + it("should accept embedding config object", () => { + const memory = new Memory({ + storage: mockStorageAdapter, + embedding: { + model: "openai/text-embedding-3-small", + normalize: true, + maxBatchSize: 50, + }, + }); + + expectTypeOf(memory).toMatchTypeOf(); + }); + it("should accept optional VectorAdapter", () => { const memory = new Memory({ storage: mockStorageAdapter, diff --git a/packages/core/src/memory/index.ts b/packages/core/src/memory/index.ts index 6bee01a7b..9c76978fa 100644 --- a/packages/core/src/memory/index.ts +++ b/packages/core/src/memory/index.ts @@ -6,6 +6,7 @@ import { type Logger, safeStringify } from "@voltagent/internal"; import type { UIMessage } from "ai"; import type { z } from "zod"; import type { OperationContext } from "../agent/types"; +import { AiSdkEmbeddingAdapter } from "./adapters/embedding/ai-sdk"; import { EmbeddingAdapterNotConfiguredError, VectorAdapterNotConfiguredError } from "./errors"; import type { Conversation, @@ -14,6 +15,8 @@ import type { CreateConversationInput, Document, EmbeddingAdapter, + EmbeddingAdapterConfig, + EmbeddingAdapterInput, GetConversationStepsOptions, GetMessagesOptions, MemoryConfig, @@ -30,6 +33,40 @@ import type { } from "./types"; import { BatchEmbeddingCache } from "./utils/cache"; +const isEmbeddingAdapter = (value: EmbeddingAdapterInput): value is EmbeddingAdapter => + typeof value === "object" && + value !== null && + "embed" in value && + typeof (value as EmbeddingAdapter).embed === "function" && + "embedBatch" in value && + typeof (value as EmbeddingAdapter).embedBatch === "function"; + +const isEmbeddingAdapterConfig = (value: EmbeddingAdapterInput): value is EmbeddingAdapterConfig => + typeof value === "object" && value !== null && "model" in value && !isEmbeddingAdapter(value); + +const resolveEmbeddingAdapter = ( + embedding?: EmbeddingAdapterInput, +): EmbeddingAdapter | undefined => { + if (!embedding) { + return undefined; + } + + if (isEmbeddingAdapter(embedding)) { + return embedding; + } + + if (typeof embedding === "string") { + return new AiSdkEmbeddingAdapter(embedding); + } + + if (isEmbeddingAdapterConfig(embedding)) { + const { model, ...options } = embedding; + return new AiSdkEmbeddingAdapter(model, options); + } + + return new AiSdkEmbeddingAdapter(embedding); +}; + /** * Memory Class * Handles conversation memory with optional vector search capabilities @@ -47,7 +84,7 @@ export class Memory { constructor(options: MemoryConfig) { this.storage = options.storage; - this.embedding = options.embedding; + this.embedding = resolveEmbeddingAdapter(options.embedding); this.vector = options.vector; this.workingMemoryConfig = options.workingMemory; diff --git a/packages/core/src/memory/types.ts b/packages/core/src/memory/types.ts index 99cd4d465..b043a4c60 100644 --- a/packages/core/src/memory/types.ts +++ b/packages/core/src/memory/types.ts @@ -7,6 +7,7 @@ import type { UIMessage } from "ai"; import type { z } from "zod"; import type { MessageRole, UsageInfo } from "../agent/providers/base/types"; import type { OperationContext } from "../agent/types"; +import type { EmbeddingModelReference, EmbeddingOptions } from "./adapters/embedding/types"; // ============================================================================ // Core Types (Re-exported from existing memory system) @@ -226,9 +227,9 @@ export interface MemoryConfig { storage: StorageAdapter; /** - * Optional embedding adapter for semantic operations + * Optional embedding adapter or model reference for semantic operations */ - embedding?: EmbeddingAdapter; + embedding?: EmbeddingAdapterInput; /** * Optional vector adapter for similarity search @@ -260,6 +261,21 @@ export interface MemoryConfig { workingMemory?: WorkingMemoryConfig; } +/** + * Embedding adapter config for Memory + */ +export type EmbeddingAdapterConfig = EmbeddingOptions & { + model: EmbeddingModelReference; +}; + +/** + * Embedding input options for Memory + */ +export type EmbeddingAdapterInput = + | EmbeddingAdapter + | EmbeddingModelReference + | EmbeddingAdapterConfig; + /** * Metadata about the underlying storage adapter */ diff --git a/packages/core/src/planagent/plan-agent.ts b/packages/core/src/planagent/plan-agent.ts index b2f416b67..2e9bedd3f 100644 --- a/packages/core/src/planagent/plan-agent.ts +++ b/packages/core/src/planagent/plan-agent.ts @@ -16,6 +16,7 @@ import type { } from "../agent/types"; import type { Tool, VercelTool } from "../tool"; import { createTool } from "../tool"; +import type { ToolRoutingConfig } from "../tool/routing/types"; import type { Toolkit } from "../tool/toolkit"; import { createToolkit } from "../tool/toolkit"; import { randomUUID } from "../utils/id"; @@ -102,6 +103,7 @@ type PlanAgentCustomSubagentRuntimeDefinition = { model?: unknown; tools?: (Tool | Toolkit | VercelTool)[]; toolkits?: Toolkit[]; + toolRouting?: ToolRoutingConfig | false; memory?: AgentOptions["memory"]; logger?: Logger; } & Record; @@ -128,6 +130,7 @@ export type PlanAgentOptions = Omit< systemPrompt?: InstructionsDynamicValue; tools?: (Tool | Toolkit | VercelTool)[]; toolkits?: Toolkit[]; + toolRouting?: ToolRoutingConfig | false; subagents?: PlanAgentSubagentDefinition[]; generalPurposeAgent?: boolean; planning?: PlanningToolkitOptions | false; @@ -563,7 +566,7 @@ function chainToolEndHooks( for (const hook of sequence) { await hook(args); } - return; + return undefined; } let currentOutput = args.output; @@ -579,7 +582,7 @@ function chainToolEndHooks( if (hasOverride) { return { output: currentOutput }; } - return; + return undefined; }; } @@ -644,8 +647,16 @@ function normalizeSubagentDefinitions(options: { defaultToolkits: Toolkit[]; defaultMemory: AgentOptions["memory"]; defaultLogger?: Logger; + defaultToolRouting?: ToolRoutingConfig | false; }): Array<{ name: string; description: string; config: SubAgentConfig }> { - const { defaultModel, defaultTools, defaultToolkits, defaultMemory, defaultLogger } = options; + const { + defaultModel, + defaultTools, + defaultToolkits, + defaultMemory, + defaultLogger, + defaultToolRouting, + } = options; const normalized: Array<{ name: string; description: string; config: SubAgentConfig }> = []; const rawDefinitions = options.definitions as unknown[]; @@ -683,6 +694,7 @@ function normalizeSubagentDefinitions(options: { instructions: custom.systemPrompt, tools, toolkits, + toolRouting: custom.toolRouting ?? defaultToolRouting, memory: custom.memory ?? defaultMemory, logger: custom.logger ?? defaultLogger, } as AgentOptions); @@ -824,16 +836,17 @@ function createPlanningExtension(options: { }, onToolEnd: async (args) => { if (args.error) { - return; + return undefined; } if (args.tool.name === WRITE_TODOS_TOOL_NAME) { args.context.systemContext.set(PLAN_WRITTEN_CONTEXT_KEY, true); args.context.systemContext.delete(FORCED_TOOL_CHOICE_CONTEXT_KEY); - return; + return undefined; } markPlanProgress(args.context); + return undefined; }, onStepFinish: async (args) => { const step = args.step as StepResult | undefined; @@ -1084,6 +1097,7 @@ export class PlanAgent extends Agent { defaultToolkits: subagentToolkits, defaultMemory: options.memory, defaultLogger: options.logger, + defaultToolRouting: options.toolRouting, }); if (generalPurposeAgent) { @@ -1098,6 +1112,7 @@ export class PlanAgent extends Agent { instructions: DEFAULT_SUBAGENT_PROMPT, tools: wrappedTools, toolkits: subagentToolkits, + toolRouting: options.toolRouting, memory: options.memory, logger: options.logger, }); diff --git a/packages/core/src/registries/agent-registry.ts b/packages/core/src/registries/agent-registry.ts index 2fcc75ee4..bb24e6d80 100644 --- a/packages/core/src/registries/agent-registry.ts +++ b/packages/core/src/registries/agent-registry.ts @@ -2,6 +2,7 @@ import type { Logger } from "@voltagent/internal"; import type { Agent } from "../agent/agent"; import type { Memory } from "../memory"; import type { VoltAgentObservability } from "../observability"; +import type { ToolRoutingConfig } from "../tool/routing/types"; import type { VoltOpsClient } from "../voltops/client"; /** @@ -25,6 +26,7 @@ export class AgentRegistry { private globalMemory?: Memory; private globalAgentMemory?: Memory; private globalWorkflowMemory?: Memory; + private globalToolRouting?: ToolRoutingConfig; /** * Track parent-child relationships between agents (child -> parents) @@ -276,4 +278,18 @@ export class AgentRegistry { public getGlobalWorkflowMemory(): Memory | undefined { return this.globalWorkflowMemory ?? this.globalMemory; } + + /** + * Set the global default tool routing configuration. + */ + public setGlobalToolRouting(toolRouting: ToolRoutingConfig | undefined): void { + this.globalToolRouting = toolRouting; + } + + /** + * Get the global default tool routing configuration. + */ + public getGlobalToolRouting(): ToolRoutingConfig | undefined { + return this.globalToolRouting; + } } diff --git a/packages/core/src/registries/embedding-model-router-types.generated.ts b/packages/core/src/registries/embedding-model-router-types.generated.ts new file mode 100644 index 000000000..8d9f570fd --- /dev/null +++ b/packages/core/src/registries/embedding-model-router-types.generated.ts @@ -0,0 +1,62 @@ +/** + * THIS FILE IS AUTO-GENERATED - DO NOT EDIT + * Generated from https://models.dev/api.json + */ + +export type EmbeddingModelsMap = { + readonly azure: readonly [ + "cohere-embed-v-4-0", + "cohere-embed-v3-english", + "cohere-embed-v3-multilingual", + "text-embedding-3-large", + "text-embedding-3-small", + "text-embedding-ada-002", + ]; + readonly "azure-cognitive-services": readonly [ + "cohere-embed-v-4-0", + "cohere-embed-v3-english", + "cohere-embed-v3-multilingual", + "text-embedding-3-large", + "text-embedding-3-small", + "text-embedding-ada-002", + ]; + readonly "cloudflare-ai-gateway": readonly [ + "workers-ai/@cf/pfnet/plamo-embedding-1b", + "workers-ai/@cf/qwen/qwen3-embedding-0.6b", + ]; + readonly google: readonly ["gemini-embedding-001"]; + readonly "google-vertex": readonly ["gemini-embedding-001"]; + readonly huggingface: readonly ["Qwen/Qwen3-Embedding-4B", "Qwen/Qwen3-Embedding-8B"]; + readonly inference: readonly ["qwen/qwen3-embedding-4b"]; + readonly mistral: readonly ["mistral-embed"]; + readonly nvidia: readonly ["nvidia/llama-embed-nemotron-8b"]; + readonly openai: readonly [ + "text-embedding-3-large", + "text-embedding-3-small", + "text-embedding-ada-002", + ]; + readonly "privatemode-ai": readonly ["qwen3-embedding-4b"]; + readonly vercel: readonly [ + "alibaba/qwen3-embedding-0.6b", + "alibaba/qwen3-embedding-4b", + "alibaba/qwen3-embedding-8b", + "amazon/titan-embed-text-v2", + "cohere/embed-v4.0", + "google/gemini-embedding-001", + "google/text-embedding-005", + "google/text-multilingual-embedding-002", + "mistral/codestral-embed", + "mistral/mistral-embed", + "openai/text-embedding-3-large", + "openai/text-embedding-3-small", + "openai/text-embedding-ada-002", + ]; +}; + +export type EmbeddingProviderId = keyof EmbeddingModelsMap; + +export type EmbeddingRouterModelId = + | { + [P in EmbeddingProviderId]: `${P}/${EmbeddingModelsMap[P][number]}`; + }[EmbeddingProviderId] + | (string & {}); diff --git a/packages/core/src/registries/embedding-model-router-types.ts b/packages/core/src/registries/embedding-model-router-types.ts new file mode 100644 index 000000000..e60a7da50 --- /dev/null +++ b/packages/core/src/registries/embedding-model-router-types.ts @@ -0,0 +1 @@ +export type { EmbeddingRouterModelId } from "./embedding-model-router-types.generated"; diff --git a/packages/core/src/registries/model-provider-registry.ts b/packages/core/src/registries/model-provider-registry.ts index c68b3994d..cdad24cdf 100644 --- a/packages/core/src/registries/model-provider-registry.ts +++ b/packages/core/src/registries/model-provider-registry.ts @@ -1,14 +1,20 @@ import { safeStringify } from "@voltagent/internal"; -import type { LanguageModel } from "ai"; +import type { EmbeddingModel, LanguageModel } from "ai"; import { MODEL_PROVIDER_REGISTRY, type ModelProviderRegistryEntry, } from "./model-provider-registry.generated"; export type LanguageModelFactory = (modelId: string) => LanguageModel; +type EmbeddingModelInstance = Exclude; +export type EmbeddingModelFactory = (modelId: string) => EmbeddingModelInstance; export type ModelProvider = { languageModel: LanguageModelFactory; + embeddingModel?: EmbeddingModelFactory; + embedding?: EmbeddingModelFactory; + textEmbeddingModel?: EmbeddingModelFactory; + textEmbedding?: EmbeddingModelFactory; }; export type ModelProviderEntry = ModelProvider | LanguageModelFactory; @@ -506,6 +512,25 @@ const isModelProvider = (value: unknown): value is ModelProvider => const isLanguageModelFactory = (value: unknown): value is LanguageModelFactory => typeof value === "function"; +const resolveEmbeddingFactory = ( + provider: ModelProviderEntry, +): EmbeddingModelFactory | undefined => { + const candidate = provider as { + embeddingModel?: EmbeddingModelFactory; + embedding?: EmbeddingModelFactory; + textEmbeddingModel?: EmbeddingModelFactory; + textEmbedding?: EmbeddingModelFactory; + }; + + const factory = + candidate.embeddingModel ?? + candidate.embedding ?? + candidate.textEmbeddingModel ?? + candidate.textEmbedding; + + return typeof factory === "function" ? factory.bind(provider as object) : undefined; +}; + const resolveProviderExport = ( moduleExports: Record, exportName: string, @@ -755,8 +780,9 @@ const splitModelId = (value: string): { providerId: string; modelId: string } => export class ModelProviderRegistry { private providers = new Map(); + private providerEntries = new Map(); private loaders = new Map(); - private loading = new Map>(); + private entryLoading = new Map>(); private dynamicRegistry = new Map(); private refreshInterval: ReturnType | null = null; private lastRefreshTime: Date | null = null; @@ -804,7 +830,11 @@ export class ModelProviderRegistry { private registerProviderConfig(config: ModelProviderRegistryEntry): void { const providerId = normalizeProviderId(config.id); - if (this.providers.has(providerId) || this.loaders.has(providerId)) { + if ( + this.providers.has(providerId) || + this.providerEntries.has(providerId) || + this.loaders.has(providerId) + ) { return; } this.registerProviderLoader(providerId, this.createProviderLoader(config)); @@ -944,6 +974,7 @@ export class ModelProviderRegistry { public registerProvider(providerId: string, provider: ModelProviderEntry): void { const normalizedId = normalizeProviderId(providerId); + this.providerEntries.set(normalizedId, provider); this.providers.set(normalizedId, this.normalizeProvider(provider, normalizedId)); } @@ -955,10 +986,16 @@ export class ModelProviderRegistry { public unregisterProvider(providerId: string): void { const normalizedId = normalizeProviderId(providerId); this.providers.delete(normalizedId); + this.providerEntries.delete(normalizedId); + this.entryLoading.delete(normalizedId); } public listProviders(): string[] { - const providers = new Set([...this.providers.keys(), ...this.loaders.keys()]); + const providers = new Set([ + ...this.providers.keys(), + ...this.providerEntries.keys(), + ...this.loaders.keys(), + ]); return [...providers].sort(); } @@ -975,34 +1012,70 @@ export class ModelProviderRegistry { return provider(resolvedModelId); } + public async resolveEmbeddingModel(modelId: string): Promise { + const { providerId, modelId: resolvedModelId } = splitModelId(modelId); + const providerEntry = await this.getProviderEntry(providerId); + if (!providerEntry) { + const available = this.listProviders(); + const availableMessage = available.length + ? `Available providers: ${available.join(", ")}.` + : "No providers are registered."; + throw new Error(`No provider registered for "${providerId}". ${availableMessage}`); + } + + const embeddingFactory = resolveEmbeddingFactory(providerEntry); + if (!embeddingFactory) { + throw new Error(`Provider "${providerId}" does not support embedding models.`); + } + + return embeddingFactory(resolvedModelId); + } + private async getProvider(providerId: string): Promise { const normalizedId = normalizeProviderId(providerId); const existing = this.providers.get(normalizedId); if (existing) { return existing; } + const entry = await this.getProviderEntry(normalizedId); + if (!entry) { + return undefined; + } + const normalizedProvider = this.normalizeProvider(entry, normalizedId); + this.providers.set(normalizedId, normalizedProvider); + return normalizedProvider; + } + + private async getProviderEntry(providerId: string): Promise { + const normalizedId = normalizeProviderId(providerId); + const existing = this.providerEntries.get(normalizedId); + if (existing) { + return existing; + } const loader = this.loaders.get(normalizedId); if (!loader) { return undefined; } - const pending = this.loading.get(normalizedId); + const pending = this.entryLoading.get(normalizedId); if (pending) { return pending; } const loadPromise = loader() .then((provider) => { - const normalizedProvider = this.normalizeProvider(provider, normalizedId); - this.providers.set(normalizedId, normalizedProvider); - return normalizedProvider; + this.providerEntries.set(normalizedId, provider); + if (!this.providers.has(normalizedId)) { + this.providers.set(normalizedId, this.normalizeProvider(provider, normalizedId)); + } + return provider; }) .finally(() => { - this.loading.delete(normalizedId); + this.entryLoading.delete(normalizedId); }); - this.loading.set(normalizedId, loadPromise); + this.entryLoading.set(normalizedId, loadPromise); return loadPromise; } diff --git a/packages/core/src/tool/index.ts b/packages/core/src/tool/index.ts index 31c5c0a57..135a0e7d7 100644 --- a/packages/core/src/tool/index.ts +++ b/packages/core/src/tool/index.ts @@ -64,6 +64,23 @@ export type { ProviderOptions } from "@ai-sdk/provider-utils"; export { ToolManager, ToolStatus, ToolStatusInfo } from "./manager"; // Export Toolkit type and createToolkit function export { type Toolkit, createToolkit } from "./toolkit"; +// Export tool routing helpers +export { createToolRouter, createEmbeddingToolRouterStrategy, isToolRouter } from "./routing"; +export type { + ToolArgumentResolver, + ToolRouter, + ToolRouterCandidate, + ToolRouterContext, + ToolRouterInput, + ToolRouterMode, + ToolRouterResult, + ToolRouterResultItem, + ToolRouterSelection, + ToolRouterStrategy, + ToolRoutingConfig, + ToolRoutingEmbeddingConfig, + ToolRoutingEmbeddingInput, +} from "./routing/types"; /** * Tool definition compatible with Vercel AI SDK diff --git a/packages/core/src/tool/manager/BaseToolManager.ts b/packages/core/src/tool/manager/BaseToolManager.ts index e4f8d5924..7ddb110d7 100644 --- a/packages/core/src/tool/manager/BaseToolManager.ts +++ b/packages/core/src/tool/manager/BaseToolManager.ts @@ -197,6 +197,23 @@ export abstract class BaseToolManager< ]; } + /** + * Get a tool by name across standalone tools and toolkits. + */ + getToolByName(toolName: string): BaseTool | ProviderTool | undefined { + const standalone = this.baseTools.get(toolName) ?? this.providerTools.get(toolName); + if (standalone) { + return standalone; + } + for (const toolkit of this.toolkits.values()) { + const tool = toolkit.getToolByName(toolName); + if (tool) { + return tool; + } + } + return undefined; + } + /** * Get names of all tools (standalone and inside toolkits), deduplicated. */ diff --git a/packages/core/src/tool/routing/constants.ts b/packages/core/src/tool/routing/constants.ts new file mode 100644 index 000000000..9c26bd9c9 --- /dev/null +++ b/packages/core/src/tool/routing/constants.ts @@ -0,0 +1 @@ +export const TOOL_ROUTER_SYMBOL = Symbol("voltagent.toolRouter"); diff --git a/packages/core/src/tool/routing/embedding.ts b/packages/core/src/tool/routing/embedding.ts new file mode 100644 index 000000000..a0e8eefb9 --- /dev/null +++ b/packages/core/src/tool/routing/embedding.ts @@ -0,0 +1,195 @@ +import { safeStringify } from "@voltagent/internal/utils"; +import { AiSdkEmbeddingAdapter } from "../../memory/adapters/embedding/ai-sdk"; +import type { + EmbeddingAdapter, + EmbeddingModelReference, +} from "../../memory/adapters/embedding/types"; +import { cosineSimilarity } from "../../memory/utils/vector-math"; +import { zodSchemaToJsonUI } from "../../utils/toolParser"; +import type { + ToolRouterCandidate, + ToolRouterStrategy, + ToolRoutingEmbeddingConfig, + ToolRoutingEmbeddingInput, +} from "./types"; + +const isEmbeddingAdapter = (value: unknown): value is EmbeddingAdapter => { + if (!value || typeof value !== "object") { + return false; + } + const candidate = value as EmbeddingAdapter; + return ( + typeof candidate.embed === "function" && + typeof candidate.embedBatch === "function" && + typeof candidate.getModelName === "function" + ); +}; + +const normalizeEmbeddingConfig = (input: ToolRoutingEmbeddingInput): ToolRoutingEmbeddingConfig => { + if (isEmbeddingAdapter(input)) { + return { model: input }; + } + if (typeof input === "object" && input !== null && "model" in input) { + return input as ToolRoutingEmbeddingConfig; + } + return { + model: input as EmbeddingAdapter | EmbeddingModelReference, + }; +}; + +const defaultToolText = (tool: ToolRouterCandidate): string => { + const parts: string[] = []; + parts.push(`Tool: ${tool.name}`); + if (tool.description) { + parts.push(`Description: ${tool.description}`); + } + if (tool.tags && tool.tags.length > 0) { + parts.push(`Tags: ${tool.tags.join(", ")}`); + } + if (tool.parameters) { + const normalized = + typeof tool.parameters === "object" ? tool.parameters : zodSchemaToJsonUI(tool.parameters); + parts.push(`Parameters: ${safeStringify(normalized)}`); + } + return parts.join("\n"); +}; + +export const createEmbeddingToolRouterStrategy = ( + input: ToolRoutingEmbeddingInput, +): ToolRouterStrategy => { + const config = normalizeEmbeddingConfig(input); + const adapter = isEmbeddingAdapter(config.model) + ? config.model + : new AiSdkEmbeddingAdapter(config.model as EmbeddingModelReference, { + normalize: config.normalize ?? true, + maxBatchSize: config.maxBatchSize, + }); + const toolText = config.toolText ?? defaultToolText; + const cache = new Map(); + + const getToolEmbeddings = async ( + tools: ToolRouterCandidate[], + ): Promise<{ embeddings: number[][]; stats: { cached: number; computed: number } }> => { + const texts: { name: string; text: string }[] = tools.map((tool) => ({ + name: tool.name, + text: toolText(tool), + })); + + const currentNames = new Set(texts.map((entry) => entry.name)); + for (const cachedName of cache.keys()) { + if (!currentNames.has(cachedName)) { + cache.delete(cachedName); + } + } + + const pending: Array<{ name: string; text: string }> = []; + for (const entry of texts) { + const cached = cache.get(entry.name); + if (!cached || cached.text !== entry.text) { + pending.push(entry); + } + } + + if (pending.length > 0) { + const embeddings = await adapter.embedBatch(pending.map((entry) => entry.text)); + pending.forEach((entry, index) => { + cache.set(entry.name, { text: entry.text, embedding: embeddings[index] }); + }); + } + + return { + embeddings: texts.map((entry) => { + const cached = cache.get(entry.name); + return cached ? cached.embedding : []; + }), + stats: { + cached: texts.length - pending.length, + computed: pending.length, + }, + }; + }; + + return { + select: async ({ query, tools, topK, context }) => { + if (tools.length === 0) { + return []; + } + + const oc = context?.operationContext; + const routerName = context?.routerName; + const parentSpan = context?.parentSpan; + const dimensions = adapter.getDimensions(); + const embeddingSpanAttributes = { + "embedding.model": adapter.getModelName(), + ...(dimensions ? { "embedding.dimensions": dimensions } : {}), + ...(routerName ? { "tool.router.name": routerName } : {}), + "tool.router.query": query, + "tool.router.tool_count": tools.length, + "tool.router.top_k": topK, + "tool.router.strategy": "embedding", + input: query, + }; + const embeddingSpan = oc?.traceContext + ? parentSpan + ? oc.traceContext.createChildSpanWithParent( + parentSpan, + `tool.router.embedding:${routerName ?? "router"}`, + "embedding", + { + label: routerName + ? `Tool Router Embedding: ${routerName}` + : "Tool Router Embedding", + attributes: embeddingSpanAttributes, + }, + ) + : oc.traceContext.createChildSpan( + `tool.router.embedding:${routerName ?? "router"}`, + "embedding", + { + label: routerName + ? `Tool Router Embedding: ${routerName}` + : "Tool Router Embedding", + attributes: embeddingSpanAttributes, + }, + ) + : null; + + const runSelection = async () => { + const queryEmbedding = await adapter.embed(query); + const { embeddings: toolEmbeddings, stats } = await getToolEmbeddings(tools); + + const scored = tools.map((tool, index) => { + const embedding = toolEmbeddings[index] ?? []; + return { + name: tool.name, + score: embedding.length > 0 ? cosineSimilarity(queryEmbedding, embedding) : 0, + }; + }); + + return { + scored, + stats, + }; + }; + + if (!embeddingSpan || !oc) { + const { scored } = await runSelection(); + return scored.sort((a, b) => (b.score ?? 0) - (a.score ?? 0)).slice(0, Math.max(0, topK)); + } + + try { + const { scored, stats } = await oc.traceContext.withSpan(embeddingSpan, runSelection); + oc.traceContext.endChildSpan(embeddingSpan, "completed", { + attributes: { + "tool.router.embedding.cache_hits": stats.cached, + "tool.router.embedding.cache_misses": stats.computed, + }, + }); + return scored.sort((a, b) => (b.score ?? 0) - (a.score ?? 0)).slice(0, Math.max(0, topK)); + } catch (error) { + oc.traceContext.endChildSpan(embeddingSpan, "error", { error }); + throw error; + } + }, + }; +}; diff --git a/packages/core/src/tool/routing/index.ts b/packages/core/src/tool/routing/index.ts new file mode 100644 index 000000000..ddee5d4c1 --- /dev/null +++ b/packages/core/src/tool/routing/index.ts @@ -0,0 +1,86 @@ +import { z } from "zod"; +import type { Agent } from "../../agent/agent"; +import { AGENT_REF_CONTEXT_KEY } from "../../agent/context-keys"; +import type { ToolExecuteOptions, ToolSchema } from "../../agent/providers/base/types"; +import { createTool } from "../index"; +import { TOOL_ROUTER_SYMBOL } from "./constants"; +import { createEmbeddingToolRouterStrategy } from "./embedding"; +import type { + ToolArgumentResolver, + ToolRouter, + ToolRouterInput, + ToolRouterMetadata, + ToolRouterMode, + ToolRouterStrategy, + ToolRoutingEmbeddingInput, +} from "./types"; + +export type CreateToolRouterOptions = { + name: string; + description: string; + strategy?: ToolRouterStrategy; + embedding?: ToolRoutingEmbeddingInput; + mode?: ToolRouterMode; + executionModel?: ToolRouterMetadata["executionModel"]; + resolver?: ToolArgumentResolver; + topK?: number; + parallel?: boolean; + parameters?: ToolSchema; +}; + +const defaultRouterParameters = z.object({ + query: z.string().describe("The user request or query to route"), + topK: z.number().int().positive().optional().describe("Number of tools to select"), +}); + +export const createToolRouter = (options: CreateToolRouterOptions): ToolRouter => { + const strategy = + options.strategy ?? + (options.embedding ? createEmbeddingToolRouterStrategy(options.embedding) : undefined); + + if (!strategy) { + throw new Error("Tool router requires a strategy or embedding configuration."); + } + + const metadata: ToolRouterMetadata = { + strategy, + mode: options.mode, + executionModel: options.executionModel, + resolver: options.resolver, + topK: options.topK, + parallel: options.parallel, + }; + + const execute = async (input: ToolRouterInput, execOptions?: ToolExecuteOptions) => { + const agent = execOptions?.systemContext?.get(AGENT_REF_CONTEXT_KEY) as Agent | undefined; + const executor = agent && (agent as any).__executeToolRouter; + if (typeof executor !== "function") { + throw new Error("Tool router requires an Agent execution context."); + } + + return await executor.call(agent, { + router: tool as ToolRouter, + input, + options: execOptions, + }); + }; + + const tool = createTool({ + name: options.name, + description: options.description, + parameters: options.parameters ?? defaultRouterParameters, + execute, + }); + + const routerTool = tool as ToolRouter; + routerTool[TOOL_ROUTER_SYMBOL] = metadata; + + return routerTool; +}; + +export const isToolRouter = (tool: unknown): tool is ToolRouter => { + return Boolean(tool && typeof tool === "object" && TOOL_ROUTER_SYMBOL in tool); +}; + +export { createEmbeddingToolRouterStrategy } from "./embedding"; +export type { ToolRouterStrategy }; diff --git a/packages/core/src/tool/routing/types.ts b/packages/core/src/tool/routing/types.ts new file mode 100644 index 000000000..c49338ef6 --- /dev/null +++ b/packages/core/src/tool/routing/types.ts @@ -0,0 +1,103 @@ +import type { Span } from "@opentelemetry/api"; +import type { AgentModelValue, OperationContext } from "../../agent/types"; +import type { + EmbeddingAdapter, + EmbeddingModelReference, +} from "../../memory/adapters/embedding/types"; +import type { ProviderTool, Tool, VercelTool } from "../index"; +import type { Toolkit } from "../toolkit"; +import { TOOL_ROUTER_SYMBOL } from "./constants"; + +export type ToolRouterMode = "agent" | "resolver"; + +export type ToolRouterSelection = { + name: string; + score?: number; + reason?: string; +}; + +export type ToolRouterInput = { + query: string; + topK?: number; +}; + +export type ToolRouterResultItem = { + toolName: string; + toolCallId?: string; + output?: unknown; + error?: string; +}; + +export type ToolRouterResult = { + query: string; + selections: ToolRouterSelection[]; + results: ToolRouterResultItem[]; +}; + +export type ToolRouterCandidate = { + name: string; + description?: string; + tags?: string[]; + parameters?: unknown; + tool: Tool | ProviderTool; +}; + +export type ToolRouterContext = { + agentId: string; + agentName: string; + operationContext: OperationContext; + routerName?: string; + parentSpan?: Span; +}; + +export type ToolRouterStrategy = { + select: (params: { + query: string; + tools: ToolRouterCandidate[]; + topK: number; + context: ToolRouterContext; + }) => Promise; +}; + +export type ToolArgumentResolver = (params: { + query: string; + tool: Tool; + context: ToolRouterContext; +}) => Promise>; + +export type ToolRouterMetadata = { + strategy: ToolRouterStrategy; + mode?: ToolRouterMode; + executionModel?: AgentModelValue; + resolver?: ToolArgumentResolver; + topK?: number; + parallel?: boolean; +}; + +export type ToolRouter = Tool & { + [TOOL_ROUTER_SYMBOL]: ToolRouterMetadata; +}; + +export type ToolRoutingEmbeddingConfig = { + model: EmbeddingAdapter | EmbeddingModelReference; + normalize?: boolean; + maxBatchSize?: number; + topK?: number; + toolText?: (tool: ToolRouterCandidate) => string; +}; + +export type ToolRoutingEmbeddingInput = + | ToolRoutingEmbeddingConfig + | EmbeddingAdapter + | EmbeddingModelReference; + +export type ToolRoutingConfig = { + routers?: ToolRouter[]; + pool?: (Tool | Toolkit | VercelTool)[]; + expose?: (Tool | Toolkit | VercelTool)[]; + mode?: ToolRouterMode; + executionModel?: AgentModelValue; + embedding?: ToolRoutingEmbeddingInput; + topK?: number; + parallel?: boolean; +}; diff --git a/packages/core/src/types.ts b/packages/core/src/types.ts index ffb18a2c7..c4e507165 100644 --- a/packages/core/src/types.ts +++ b/packages/core/src/types.ts @@ -13,6 +13,7 @@ import type { MCPServerRegistry } from "./mcp"; import type { Memory } from "./memory"; import type { VoltAgentObservability } from "./observability"; import type { ToolStatusInfo } from "./tool"; +import type { ToolRoutingConfig } from "./tool/routing/types"; import type { TriggerRegistry } from "./triggers/registry"; import type { VoltAgentTriggersConfig } from "./triggers/types"; import type { VoltOpsClient } from "./voltops/client"; @@ -227,6 +228,10 @@ export type VoltAgentOptions = { * Falls back to `memory` when not provided. */ workflowMemory?: Memory; + /** + * Global tool routing defaults (applied to agents without explicit toolRouting config). + */ + toolRouting?: ToolRoutingConfig; /** Optional VoltOps trigger handlers */ triggers?: VoltAgentTriggersConfig; /** diff --git a/packages/core/src/voltagent.ts b/packages/core/src/voltagent.ts index 03709664d..97584c8c5 100644 --- a/packages/core/src/voltagent.ts +++ b/packages/core/src/voltagent.ts @@ -61,6 +61,9 @@ export class VoltAgent { if (options.workflowMemory) { this.registry.setGlobalWorkflowMemory(options.workflowMemory); } + if (options.toolRouting) { + this.registry.setGlobalToolRouting(options.toolRouting); + } // Initialize logger this.logger = (options.logger || getGlobalLogger()).child({ component: "voltagent" }); @@ -378,6 +381,7 @@ export class VoltAgent { public registerAgent(agent: Agent): void { // Register the agent this.applyDefaultMemoryToAgent(agent); + agent.__setDefaultToolRouting?.(this.registry.getGlobalToolRouting()); this.registry.registerAgent(agent); } diff --git a/packages/postgres/CHANGELOG.md b/packages/postgres/CHANGELOG.md index 5b55e094e..619661619 100644 --- a/packages/postgres/CHANGELOG.md +++ b/packages/postgres/CHANGELOG.md @@ -372,7 +372,7 @@ ## New: PostgresVectorAdapter ```typescript - import { Agent, Memory, AiSdkEmbeddingAdapter } from "@voltagent/core"; + import { Agent, Memory } from "@voltagent/core"; import { PostgresMemoryAdapter, PostgresVectorAdapter } from "@voltagent/postgres"; import { openai } from "@ai-sdk/openai"; @@ -380,7 +380,7 @@ storage: new PostgresMemoryAdapter({ connectionString: process.env.DATABASE_URL, }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new PostgresVectorAdapter({ connectionString: process.env.DATABASE_URL, }), diff --git a/packages/voltagent-memory/CHANGELOG.md b/packages/voltagent-memory/CHANGELOG.md index 7626f696f..75ea76e9c 100644 --- a/packages/voltagent-memory/CHANGELOG.md +++ b/packages/voltagent-memory/CHANGELOG.md @@ -228,14 +228,13 @@ ```typescript import { ManagedMemoryAdapter, ManagedMemoryVectorAdapter } from "@voltagent/voltagent-memory"; - import { AiSdkEmbeddingAdapter, Memory } from "@voltagent/core"; - import { openai } from "@ai-sdk/openai"; + import { Memory } from "@voltagent/core"; const memory = new Memory({ storage: new ManagedMemoryAdapter({ databaseName: "production-memory", }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new ManagedMemoryVectorAdapter({ databaseName: "production-memory", }), diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 68b0a6501..fffa72593 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -2758,6 +2758,40 @@ importers: specifier: ^5.8.2 version: 5.9.2 + examples/with-tool-routing: + dependencies: + '@ai-sdk/openai': + specifier: ^3.0.0 + version: 3.0.12(zod@3.25.76) + '@voltagent/cli': + specifier: ^0.1.21 + version: link:../../packages/cli + '@voltagent/core': + specifier: ^2.1.6 + version: link:../../packages/core + '@voltagent/logger': + specifier: ^2.0.2 + version: link:../../packages/logger + '@voltagent/server-hono': + specifier: ^2.0.4 + version: link:../../packages/server-hono + ai: + specifier: ^6.0.0 + version: 6.0.3(zod@3.25.76) + zod: + specifier: ^3.25.76 + version: 3.25.76 + devDependencies: + '@types/node': + specifier: ^24.2.1 + version: 24.6.2 + tsx: + specifier: ^4.19.3 + version: 4.20.4 + typescript: + specifier: ^5.8.2 + version: 5.9.3 + examples/with-tools: dependencies: '@ai-sdk/openai': diff --git a/website/deployment-docs/cloudflare-workers.md b/website/deployment-docs/cloudflare-workers.md index bb1bcb8e9..4bc126261 100644 --- a/website/deployment-docs/cloudflare-workers.md +++ b/website/deployment-docs/cloudflare-workers.md @@ -311,7 +311,7 @@ const memory = new Memory({ vector: new PostgresVectorAdapter({ connectionString: env.POSTGRES_URL, }), - // embedding adapter (e.g. AiSdkEmbeddingAdapter) stays the same + embedding: "openai/text-embedding-3-small", }); const agent = new Agent({ diff --git a/website/docs/agents/memory/cloudflare-d1.md b/website/docs/agents/memory/cloudflare-d1.md index a9b53e217..441007547 100644 --- a/website/docs/agents/memory/cloudflare-d1.md +++ b/website/docs/agents/memory/cloudflare-d1.md @@ -150,7 +150,7 @@ See [Working Memory](./working-memory.md) for configuration details. ### Semantic Search (Optional) -D1 provides storage only. To enable semantic search, pair it with an embedding + vector adapter (for example, `AiSdkEmbeddingAdapter` and `InMemoryVectorAdapter`). +D1 provides storage only. To enable semantic search, pair it with an embedding model string (for example, `openai/text-embedding-3-small`) and a vector adapter such as `InMemoryVectorAdapter`. See [Semantic Search](./semantic-search.md) for usage. diff --git a/website/docs/agents/memory/in-memory.md b/website/docs/agents/memory/in-memory.md index 12d5fea07..3e8aebf09 100644 --- a/website/docs/agents/memory/in-memory.md +++ b/website/docs/agents/memory/in-memory.md @@ -69,17 +69,11 @@ See [Working Memory](./working-memory.md) for configuration details. Combine with `InMemoryVectorAdapter` for semantic search during development: ```ts -import { - Memory, - AiSdkEmbeddingAdapter, - InMemoryVectorAdapter, - InMemoryStorageAdapter, -} from "@voltagent/core"; -import { openai } from "@ai-sdk/openai"; +import { Memory, InMemoryVectorAdapter, InMemoryStorageAdapter } from "@voltagent/core"; const memory = new Memory({ storage: new InMemoryStorageAdapter(), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new InMemoryVectorAdapter(), }); ``` diff --git a/website/docs/agents/memory/libsql.md b/website/docs/agents/memory/libsql.md index 2855a2cec..241d541fb 100644 --- a/website/docs/agents/memory/libsql.md +++ b/website/docs/agents/memory/libsql.md @@ -98,13 +98,12 @@ See [Working Memory](./working-memory.md) for configuration details. Use `LibSQLVectorAdapter` for persistent vector storage: ```ts -import { Memory, AiSdkEmbeddingAdapter } from "@voltagent/core"; +import { Memory } from "@voltagent/core"; import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; -import { openai } from "@ai-sdk/openai"; const memory = new Memory({ storage: new LibSQLMemoryAdapter({ url: "file:./.voltagent/memory.db" }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new LibSQLVectorAdapter({ url: "file:./.voltagent/memory.db" }), }); ``` diff --git a/website/docs/agents/memory/managed-memory.md b/website/docs/agents/memory/managed-memory.md index 72fc4e05c..d22869fc6 100644 --- a/website/docs/agents/memory/managed-memory.md +++ b/website/docs/agents/memory/managed-memory.md @@ -126,14 +126,13 @@ Enable semantic search with `ManagedMemoryVectorAdapter`: ```ts import { ManagedMemoryAdapter, ManagedMemoryVectorAdapter } from "@voltagent/voltagent-memory"; -import { AiSdkEmbeddingAdapter, Memory } from "@voltagent/core"; -import { openai } from "@ai-sdk/openai"; +import { Memory } from "@voltagent/core"; const memory = new Memory({ storage: new ManagedMemoryAdapter({ databaseName: "production-memory", }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new ManagedMemoryVectorAdapter({ databaseName: "production-memory", }), diff --git a/website/docs/agents/memory/overview.md b/website/docs/agents/memory/overview.md index 41623bbe3..9cae3f049 100644 --- a/website/docs/agents/memory/overview.md +++ b/website/docs/agents/memory/overview.md @@ -161,14 +161,13 @@ const agent = new Agent({ ### Semantic Search + Working Memory ```ts -import { Agent, Memory, AiSdkEmbeddingAdapter, InMemoryVectorAdapter } from "@voltagent/core"; +import { Agent, Memory, InMemoryVectorAdapter } from "@voltagent/core"; import { LibSQLMemoryAdapter } from "@voltagent/libsql"; -import { openai } from "@ai-sdk/openai"; import { z } from "zod"; const memory = new Memory({ storage: new LibSQLMemoryAdapter({ url: "file:./.voltagent/memory.db" }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new InMemoryVectorAdapter(), workingMemory: { enabled: true, diff --git a/website/docs/agents/memory/postgres.md b/website/docs/agents/memory/postgres.md index 881de26ba..6e05093d4 100644 --- a/website/docs/agents/memory/postgres.md +++ b/website/docs/agents/memory/postgres.md @@ -131,15 +131,14 @@ See [Working Memory](./working-memory.md) for configuration details. Store vector embeddings directly in PostgreSQL for semantic search (no extensions required): ```ts -import { Memory, AiSdkEmbeddingAdapter } from "@voltagent/core"; +import { Memory } from "@voltagent/core"; import { PostgreSQLMemoryAdapter, PostgresVectorAdapter } from "@voltagent/postgres"; -import { openai } from "@ai-sdk/openai"; const memory = new Memory({ storage: new PostgreSQLMemoryAdapter({ connection: process.env.DATABASE_URL!, }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new PostgresVectorAdapter({ connection: process.env.DATABASE_URL!, }), diff --git a/website/docs/agents/memory/semantic-search.md b/website/docs/agents/memory/semantic-search.md index f0867a2e5..4a417559d 100644 --- a/website/docs/agents/memory/semantic-search.md +++ b/website/docs/agents/memory/semantic-search.md @@ -10,13 +10,12 @@ Semantic search retrieves past messages by similarity rather than recency. It re ## Configuration ```ts -import { Agent, Memory, AiSdkEmbeddingAdapter, InMemoryVectorAdapter } from "@voltagent/core"; +import { Agent, Memory, InMemoryVectorAdapter } from "@voltagent/core"; import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; -import { openai } from "@ai-sdk/openai"; const memory = new Memory({ storage: new LibSQLMemoryAdapter({ url: "file:./.voltagent/memory.db" }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new LibSQLVectorAdapter({ url: "file:./.voltagent/memory.db" }), // or InMemoryVectorAdapter() for dev enableCache: true, // optional embedding cache }); @@ -28,10 +27,26 @@ const agent = new Agent({ }); ``` +You can pass an embedding config object to set adapter options: + +```ts +const memory = new Memory({ + storage: new LibSQLMemoryAdapter({ url: "file:./.voltagent/memory.db" }), + embedding: { + model: "openai/text-embedding-3-small", + normalize: true, + }, + vector: new LibSQLVectorAdapter({ url: "file:./.voltagent/memory.db" }), +}); +``` + ### Available Adapters **Embedding:** +Memory accepts an embedding adapter or a provider-qualified model string such as +`"openai/text-embedding-3-small"`. + - `AiSdkEmbeddingAdapter` - Wraps any AI SDK embedding model **Vector Storage:** @@ -122,7 +137,7 @@ Enable caching to avoid re-embedding identical text: ```ts const memory = new Memory({ storage: new LibSQLMemoryAdapter({ url: "file:./.voltagent/memory.db" }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new LibSQLVectorAdapter({ url: "file:./.voltagent/memory.db" }), enableCache: true, // enable cache cacheSize: 1000, // max entries (default: 1000) @@ -185,13 +200,12 @@ const vector = new ManagedMemoryVectorAdapter({ ## Example: Full Semantic Search Setup ```ts -import { Agent, Memory, AiSdkEmbeddingAdapter } from "@voltagent/core"; +import { Agent, Memory } from "@voltagent/core"; import { LibSQLMemoryAdapter, LibSQLVectorAdapter } from "@voltagent/libsql"; -import { openai } from "@ai-sdk/openai"; const memory = new Memory({ storage: new LibSQLMemoryAdapter({ url: "file:./.voltagent/memory.db" }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new LibSQLVectorAdapter({ url: "file:./.voltagent/memory.db" }), enableCache: true, }); diff --git a/website/docs/agents/memory/supabase.md b/website/docs/agents/memory/supabase.md index 4fe2efa72..0dda57382 100644 --- a/website/docs/agents/memory/supabase.md +++ b/website/docs/agents/memory/supabase.md @@ -224,16 +224,15 @@ See [Working Memory](./working-memory.md). ### Semantic Search ```ts -import { Memory, AiSdkEmbeddingAdapter, InMemoryVectorAdapter } from "@voltagent/core"; +import { Memory, InMemoryVectorAdapter } from "@voltagent/core"; import { SupabaseMemoryAdapter } from "@voltagent/supabase"; -import { openai } from "@ai-sdk/openai"; const memory = new Memory({ storage: new SupabaseMemoryAdapter({ supabaseUrl: process.env.SUPABASE_URL!, supabaseKey: process.env.SUPABASE_KEY!, }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new InMemoryVectorAdapter(), // or pgvector adapter }); ``` diff --git a/website/docs/getting-started/migration-guide.md b/website/docs/getting-started/migration-guide.md index 063d197db..4eea7b9b4 100644 --- a/website/docs/getting-started/migration-guide.md +++ b/website/docs/getting-started/migration-guide.md @@ -492,15 +492,14 @@ const agent = new Agent({ ### Optional: Vector search and working memory -To enable semantic search and working-memory features, add an embedding adapter and a vector adapter. For example, using ai-sdk embeddings and the in-memory vector store: +To enable semantic search and working-memory features, add an embedding model string and a vector adapter. For example, using ai-sdk embeddings and the in-memory vector store: ```ts -import { Memory, AiSdkEmbeddingAdapter, InMemoryVectorAdapter } from "@voltagent/core"; -import { openai } from "@ai-sdk/openai"; // or any ai-sdk embedding model +import { Memory, InMemoryVectorAdapter } from "@voltagent/core"; const memory = new Memory({ storage: new LibSQLMemoryAdapter({ url: "file:./.voltagent/memory.db" }), - embedding: new AiSdkEmbeddingAdapter(openai.embedding("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new InMemoryVectorAdapter(), // optional working-memory config workingMemory: { diff --git a/website/docs/getting-started/providers-models.md b/website/docs/getting-started/providers-models.md index 9d1263481..fcc6bf06a 100644 --- a/website/docs/getting-started/providers-models.md +++ b/website/docs/getting-started/providers-models.md @@ -74,7 +74,7 @@ Install the AI SDK base package (required): -If you plan to import ai-sdk providers directly (for embeddings or provider-specific helpers like `openai.embedding(...)`), install those packages too. If you only use model strings, you can skip them: +If you plan to import ai-sdk providers directly (for embeddings or provider-specific helpers), install those packages too. If you only use model strings such as `openai/text-embedding-3-small`, you can skip them: diff --git a/website/docs/rag/lancedb.md b/website/docs/rag/lancedb.md index 349931d96..7bb3ed186 100644 --- a/website/docs/rag/lancedb.md +++ b/website/docs/rag/lancedb.md @@ -173,7 +173,7 @@ export class LanceDBRetriever extends BaseRetriever { // 2. Generate Embedding const { embedding } = await embed({ - model: openai.embedding("text-embedding-3-small"), + model: "openai/text-embedding-3-small", value: searchText, }); @@ -199,7 +199,7 @@ You can swap OpenAI for other providers: ```typescript // Using a larger model const { embedding } = await embed({ - model: openai.embedding("text-embedding-3-large"), + model: "openai/text-embedding-3-large", value: query, }); ``` @@ -212,7 +212,7 @@ async function addDocument(text: string, metadata: Record) { const table = await db.openTable(tableName); const { embedding } = await embed({ - model: openai.embedding("text-embedding-3-small"), + model: "openai/text-embedding-3-small", value: text, }); diff --git a/website/docs/tools/tool-routing.md b/website/docs/tools/tool-routing.md new file mode 100644 index 000000000..77f8d085d --- /dev/null +++ b/website/docs/tools/tool-routing.md @@ -0,0 +1,312 @@ +--- +title: Tool Routing +--- + +# Tool Routing + +Tool routing exposes a small set of router tools to the model and keeps the full tool pool hidden. A router receives a query, selects tools from the pool, and runs them. The router itself is a tool, so it appears in the agent tool list and can be called by the model. + +## How Tool Routing Works + +- A router is created with `createToolRouter` or via `toolRouting.embedding`. +- The router selects tools from the pool based on a strategy. +- The router executes selected tools and returns their results as its own tool output. +- Only routers and any tools listed in `toolRouting.expose` are visible to the model. + +## Quick Setup With Embedding Routing + +This configuration creates a default router named `tool_router` and uses embeddings to rank tools in the pool. + +```ts +import { openai } from "@ai-sdk/openai"; +import { Agent, createTool } from "@voltagent/core"; +import { z } from "zod"; + +const getWeather = createTool({ + name: "get_weather", + description: "Get the current weather for a city", + parameters: z.object({ + location: z.string(), + }), + execute: async ({ location }) => ({ + location, + temperatureC: 22, + condition: "sunny", + }), +}); + +const getTimeZone = createTool({ + name: "get_time_zone", + description: "Get the time zone offset for a city", + parameters: z.object({ + location: z.string(), + }), + execute: async ({ location }) => ({ + location, + timeZone: "UTC+1", + }), +}); + +const agent = new Agent({ + name: "Tool Routing Agent", + instructions: "Use tool_router when you need a tool. Pass the user request as the query.", + model: "openai/gpt-4o-mini", + tools: [getWeather, getTimeZone], + toolRouting: { + embedding: "openai/text-embedding-3-small", + topK: 2, + }, +}); +``` + +Notes: + +- If `toolRouting.pool` is not set, the pool defaults to all non-router tools registered in the agent. +- If `toolRouting.expose` is not set, only routers are visible to the model. + +## Tool Pool and Exposed Tools + +Use `pool` to define the hidden tool set and `expose` to keep a subset visible. + +```ts +import { Agent, createTool } from "@voltagent/core"; +import { z } from "zod"; + +const getWeather = createTool({ + name: "get_weather", + description: "Get the current weather for a city", + parameters: z.object({ location: z.string() }), + execute: async ({ location }) => ({ location, temperatureC: 22 }), +}); + +const getStatus = createTool({ + name: "get_status", + description: "Return the service status", + parameters: z.object({}), + execute: async () => ({ status: "ok" }), +}); + +const agent = new Agent({ + name: "Support Agent", + instructions: "Use tool_router for tool lookups.", + model: "openai/gpt-4o-mini", + toolRouting: { + embedding: "text-embedding-3-small", + pool: [getWeather], + expose: [getStatus], + }, +}); +``` + +## Multiple Routers With Custom Strategy + +You can register multiple routers and define a custom strategy per router. A router strategy receives the query and the full list of tool candidates and returns tool names. + +```ts +import { Agent, createToolRouter, type ToolRouterStrategy } from "@voltagent/core"; + +const weatherStrategy: ToolRouterStrategy = { + select: async ({ tools, topK, query }) => { + const candidates = tools.filter( + (tool) => tool.tags?.includes("weather") || tool.name.startsWith("weather_") + ); + return candidates.slice(0, topK).map((tool) => ({ + name: tool.name, + reason: `matched weather tags for query: ${query}`, + })); + }, +}; + +const financeStrategy: ToolRouterStrategy = { + select: async ({ tools, topK }) => { + const candidates = tools.filter((tool) => tool.tags?.includes("finance")); + return candidates.slice(0, topK).map((tool) => ({ name: tool.name })); + }, +}; + +const weatherRouter = createToolRouter({ + name: "weather_router", + description: "Route weather-related requests", + strategy: weatherStrategy, +}); + +const financeRouter = createToolRouter({ + name: "finance_router", + description: "Route finance-related requests", + strategy: financeStrategy, +}); + +const agent = new Agent({ + name: "Multi Router Agent", + instructions: "Use weather_router or finance_router based on the request.", + model: "openai/gpt-4o-mini", + tools: [weatherRouter, financeRouter], + toolRouting: { + pool: [ + /* tools and toolkits */ + ], + topK: 1, + }, +}); +``` + +## Resolver Mode + +Resolver mode bypasses LLM argument generation and uses a function to produce tool arguments. + +```ts +import { Agent, createToolRouter, type ToolArgumentResolver } from "@voltagent/core"; + +const resolver: ToolArgumentResolver = async ({ query, tool }) => { + if (tool.name === "get_weather") { + return { location: query }; + } + return {}; +}; + +const router = createToolRouter({ + name: "tool_router", + description: "Route requests with a resolver", + embedding: "text-embedding-3-small", + mode: "resolver", + resolver, +}); + +const agent = new Agent({ + name: "Resolver Agent", + instructions: "Use tool_router for tools.", + model: "openai/gpt-4o-mini", + tools: [router], + toolRouting: { + pool: [ + /* tools and toolkits */ + ], + }, +}); +``` + +## Execution Model Override + +Router argument generation (agent mode) and provider-tool fallback use the router execution model if provided. + +```ts +const agent = new Agent({ + name: "Model Override Agent", + instructions: "Use tool_router for tools.", + model: "openai/gpt-4o-mini", + tools: [ + /* tools */ + ], + toolRouting: { + embedding: "text-embedding-3-small", + executionModel: "openai/gpt-4o", + }, +}); +``` + +## Provider Tools and MCP Tools + +Pool and expose lists accept Vercel AI SDK tools, including provider-defined tools. MCP tools can be added if they are represented as Vercel tools. + +```ts +import { openai } from "@ai-sdk/openai"; +import { Agent, createTool } from "@voltagent/core"; +import { z } from "zod"; + +const localTool = createTool({ + name: "get_weather", + description: "Get the current weather for a city", + parameters: z.object({ location: z.string() }), + execute: async ({ location }) => ({ location, temperatureC: 22 }), +}); + +const webSearch = openai.tools.webSearch(); + +const agent = new Agent({ + name: "Mixed Pool Agent", + instructions: "Use tool_router for tools.", + model: "openai/gpt-4o-mini", + toolRouting: { + embedding: "openai/text-embedding-3-small", + pool: [localTool, webSearch], + }, +}); +``` + +## Embedding Configuration + +The embedding router accepts several model forms and supports a custom tool text format. + +- `"openai/text-embedding-3-small"` +- `"text-embedding-3-small"` +- `"openai/text-embedding-3-small"` + +Provider-qualified strings use the same model registry and type list as agent model strings. + +```ts +import { createToolRouter } from "@voltagent/core"; + +const router = createToolRouter({ + name: "tool_router", + description: "Route with custom embedding text", + embedding: { + model: "openai/text-embedding-3-small", + topK: 3, + toolText: (tool) => { + const tags = tool.tags?.join(", ") ?? ""; + return [tool.name, tool.description, tags].filter(Boolean).join("\n"); + }, + }, +}); +``` + +The embedding strategy caches vectors in memory. The cache is reset when the process restarts. A tool text change for a given tool name triggers a new embedding for that tool. + +## Per-Call Overrides + +You can disable tool routing for a single call or provide a temporary routing config. + +```ts +const result = await agent.generateText("List tools", { + toolRouting: false, +}); + +const routed = await agent.generateText("What time is it in Berlin?", { + toolRouting: { + embedding: "text-embedding-3-small", + topK: 1, + }, +}); +``` + +## Hooks, Approval, and Error Handling + +- `needsApproval`, input/output guardrails, and tool hooks still run when tools are executed via a router. +- If a selected tool is not in the pool, the router returns an error for that selection. +- If a provider tool requires approval, the router returns a tool error. +- When `parallel` is true (default), selected tools execute concurrently. + +## Observability + +Tool routing adds spans for selection and embedding steps: + +- `tool.router.selection:*` spans include `tool.router.candidates`, `tool.router.selection.count`, and selected names. +- `tool.router.embedding:*` spans include `embedding.model`, `embedding.dimensions`, and cache hit/miss counts. + +These spans appear under the router tool span in the execution trace. + +## PlanAgent + +`PlanAgent` accepts `toolRouting` in its options. It passes routing config to its internal agent. + +```ts +import { PlanAgent } from "@voltagent/core"; + +const agent = new PlanAgent({ + name: "Planning Agent", + model: "openai/gpt-4o-mini", + toolRouting: { + embedding: "text-embedding-3-small", + }, +}); +``` diff --git a/website/evaluation-docs/prebuilt-scorers.md b/website/evaluation-docs/prebuilt-scorers.md index 7669e4d60..fa2dc641a 100644 --- a/website/evaluation-docs/prebuilt-scorers.md +++ b/website/evaluation-docs/prebuilt-scorers.md @@ -268,7 +268,7 @@ import { openai } from "@ai-sdk/openai"; const scorer = createAnswerRelevancyScorer({ model: openai("gpt-4o-mini"), - embeddingModel: openai.embedding("text-embedding-3-small"), + embeddingModel: "openai/text-embedding-3-small", strictness: 3, buildPayload: ({ payload, params }) => ({ input: String(payload.input), @@ -848,7 +848,7 @@ const allScorers = [ }), createAnswerRelevancyScorer({ model: openai("gpt-4o-mini"), - embeddingModel: openai.embedding("text-embedding-3-small"), + embeddingModel: "openai/text-embedding-3-small", }), ]; diff --git a/website/recipes/memory.md b/website/recipes/memory.md index 37229689e..d518e3bb7 100644 --- a/website/recipes/memory.md +++ b/website/recipes/memory.md @@ -91,11 +91,11 @@ const memory = new Memory({ ## With Vector Search ```typescript -import { AiSdkEmbeddingAdapter, InMemoryVectorAdapter } from "@voltagent/core"; +import { InMemoryVectorAdapter } from "@voltagent/core"; const memory = new Memory({ storage: new LibSQLMemoryAdapter(), - embedding: new AiSdkEmbeddingAdapter(openai.embeddingModel("text-embedding-3-small")), + embedding: "openai/text-embedding-3-small", vector: new InMemoryVectorAdapter(), }); ``` diff --git a/website/recipes/tool-routing.md b/website/recipes/tool-routing.md new file mode 100644 index 000000000..4bad78d0d --- /dev/null +++ b/website/recipes/tool-routing.md @@ -0,0 +1,133 @@ +--- +id: tool-routing +title: Tool Routing +slug: tool-routing +description: Route a large tool pool through a small set of router tools. +--- + +# Tool Routing + +Tool routing keeps prompts small by exposing only a few router tools. Routers select and execute tools from a larger pool on demand. + +## Quick Setup (Embedding Router) + +```typescript +import { Agent, createTool, VoltAgent } from "@voltagent/core"; +import { openai } from "@ai-sdk/openai"; +import { z } from "zod"; + +const weatherTool = createTool({ + name: "get_weather", + description: "Get the current weather for a location", + parameters: z.object({ + location: z.string().describe("City name"), + }), + execute: async ({ location }) => ({ location, temperature: 22 }), +}); + +const convertCurrencyTool = createTool({ + name: "convert_currency", + description: "Convert money between currencies", + parameters: z.object({ + amount: z.number().describe("Amount to convert"), + from: z.string().describe("Source currency code"), + to: z.string().describe("Target currency code"), + }), + execute: async ({ amount, from, to }) => ({ amount, from, to, rate: 0.92 }), +}); + +const agent = new Agent({ + name: "Assistant", + instructions: "Use tool_router for tool access.", + model: openai("gpt-4o-mini"), + tools: [weatherTool, convertCurrencyTool], + toolRouting: { + embedding: "openai/text-embedding-3-small", + topK: 3, + }, +}); + +new VoltAgent({ agents: { agent } }); +``` + +If `toolRouting.pool` is not provided, VoltAgent uses the agent's registered tools as the pool (router tools are excluded). + +## Explicit Pools (Two Categories) + +```typescript +const weatherPool = [weatherTool]; +const financePool = [convertCurrencyTool]; + +toolRouting: { + embedding: "openai/text-embedding-3-small", + pool: [...weatherPool, ...financePool], +} +``` + +## Expose Tools Directly + +Expose specific tools to the model alongside routers: + +```typescript +toolRouting: { + embedding: "openai/text-embedding-3-small", + expose: [healthCheckTool], +} +``` + +## Custom Router Strategy + +```typescript +import { createToolRouter } from "@voltagent/core"; + +const router = createToolRouter({ + name: "tool_router", + description: "Route to the best tool and execute it.", + strategy: { + select: async ({ query, tools, topK }) => { + const matches = tools + .filter((tool) => tool.description?.toLowerCase().includes(query.toLowerCase())) + .slice(0, topK); + return matches.map((tool, index) => ({ name: tool.name, score: 1 - index * 0.01 })); + }, + }, + mode: "resolver", + resolver: async ({ query, tool }) => { + return { query, tool: tool.name }; + }, +}); + +const agent = new Agent({ + name: "Custom Router Agent", + instructions: "Route requests with tool_router.", + model: openai("gpt-4o-mini"), + toolRouting: { + routers: [router], + pool: [weatherTool, convertCurrencyTool], + }, +}); +``` + +`mode: "agent"` is the default and uses the agent model to generate tool arguments. Use `resolver` when you want to build arguments deterministically. + +## Global Defaults + +```typescript +new VoltAgent({ + toolRouting: { + embedding: "openai/text-embedding-3-small", + }, + agents: { agent }, +}); +``` + +## Pooling Provider and MCP Tools + +```typescript +toolRouting: { + pool: [ + openai.tools.webSearch(), + mcpToolkit, + ], +} +``` diff --git a/website/sidebars.ts b/website/sidebars.ts index 3da76beca..7771c8287 100644 --- a/website/sidebars.ts +++ b/website/sidebars.ts @@ -206,7 +206,7 @@ const sidebars: SidebarsConfig = { type: "category", label: "Tools", collapsed: true, - items: ["tools/overview", "tools/reasoning-tool"], + items: ["tools/overview", "tools/tool-routing", "tools/reasoning-tool"], }, { type: "category", diff --git a/website/sidebarsRecipes.ts b/website/sidebarsRecipes.ts index 85a7d4e22..c5a04c78c 100644 --- a/website/sidebarsRecipes.ts +++ b/website/sidebarsRecipes.ts @@ -76,6 +76,11 @@ const sidebars: SidebarsConfig = { id: "tools", label: "Tools", }, + { + type: "doc", + id: "tool-routing", + label: "Tool Routing", + }, { type: "doc", id: "authentication", From 2e55636cc802611402515b37523f5ac5857d993a Mon Sep 17 00:00:00 2001 From: Omer Aplak Date: Thu, 22 Jan 2026 18:30:03 -0800 Subject: [PATCH 4/7] chore: fix code reviews --- packages/core/src/agent/agent.ts | 5 ++++- packages/core/src/agent/hooks/index.ts | 2 +- packages/core/src/tool/index.ts | 2 +- website/docs/agents/tools.md | 4 ++-- 4 files changed, 8 insertions(+), 5 deletions(-) diff --git a/packages/core/src/agent/agent.ts b/packages/core/src/agent/agent.ts index e10af0a0f..a357338e3 100644 --- a/packages/core/src/agent/agent.ts +++ b/packages/core/src/agent/agent.ts @@ -4664,6 +4664,7 @@ export class Agent { const resolveToolEndOutput = async (currentOutput: any) => { let output = currentOutput; + let overrideProvided = false; const toolHookResult = await tool.hooks?.onEnd?.({ tool, @@ -4674,6 +4675,7 @@ export class Agent { }); if (hasOutputOverride(toolHookResult)) { output = toolHookResult.output; + overrideProvided = true; } const agentHookResult = await hooks.onToolEnd?.({ @@ -4686,9 +4688,10 @@ export class Agent { }); if (hasOutputOverride(agentHookResult)) { output = agentHookResult.output; + overrideProvided = true; } - if (output !== currentOutput) { + if (overrideProvided) { output = await this.validateToolOutput(output, tool); } diff --git a/packages/core/src/agent/hooks/index.ts b/packages/core/src/agent/hooks/index.ts index 2d7af227c..f14e64bf9 100644 --- a/packages/core/src/agent/hooks/index.ts +++ b/packages/core/src/agent/hooks/index.ts @@ -174,7 +174,7 @@ export type AgentHookOnHandoffComplete = (args: OnHandoffCompleteHookArgs) => Pr export type AgentHookOnToolStart = (args: OnToolStartHookArgs) => Promise | void; export type AgentHookOnToolEnd = ( args: OnToolEndHookArgs, -) => Promise | OnToolEndHookResult | undefined; +) => Promise | Promise | OnToolEndHookResult | undefined; export type AgentHookOnPrepareMessages = ( args: OnPrepareMessagesHookArgs, ) => Promise | OnPrepareMessagesHookResult; diff --git a/packages/core/src/tool/index.ts b/packages/core/src/tool/index.ts index 31c5c0a57..a84337e4d 100644 --- a/packages/core/src/tool/index.ts +++ b/packages/core/src/tool/index.ts @@ -34,7 +34,7 @@ export interface ToolHookOnEndResult { export type ToolHookOnStart = (args: ToolHookOnStartArgs) => Promise | void; export type ToolHookOnEnd = ( args: ToolHookOnEndArgs, -) => Promise | ToolHookOnEndResult | undefined; +) => Promise | Promise | ToolHookOnEndResult | undefined; export type ToolHooks = { onStart?: ToolHookOnStart; diff --git a/website/docs/agents/tools.md b/website/docs/agents/tools.md index f438b5678..90fa6335a 100644 --- a/website/docs/agents/tools.md +++ b/website/docs/agents/tools.md @@ -688,9 +688,9 @@ function ReadClipboardTool({ **Important**: You must call `addToolResult` to send the tool result back to the model. Without this, the model considers the tool call a failure. -## Tool Hooks +## Agent Tool Hooks -Hooks let you respond to tool execution events for logging, UI updates, or additional actions. +Hooks let you respond to tool execution events for logging, UI updates, or additional actions. For tool-level hooks (per tool), see the **Tool Hooks** section above. ```ts import { Agent, createHooks, isAbortError } from "@voltagent/core"; From 3a41a310c2c2e1c0f7fe3b49d83a6508dce34300 Mon Sep 17 00:00:00 2001 From: Omer Aplak Date: Thu, 22 Jan 2026 18:43:44 -0800 Subject: [PATCH 5/7] fix: build error --- packages/core/src/tool/routing/index.ts | 16 ++++++++++++++-- 1 file changed, 14 insertions(+), 2 deletions(-) diff --git a/packages/core/src/tool/routing/index.ts b/packages/core/src/tool/routing/index.ts index ddee5d4c1..70a1bcb9a 100644 --- a/packages/core/src/tool/routing/index.ts +++ b/packages/core/src/tool/routing/index.ts @@ -11,6 +11,7 @@ import type { ToolRouterInput, ToolRouterMetadata, ToolRouterMode, + ToolRouterResult, ToolRouterStrategy, ToolRoutingEmbeddingInput, } from "./types"; @@ -51,15 +52,25 @@ export const createToolRouter = (options: CreateToolRouterOptions): ToolRouter = parallel: options.parallel, }; - const execute = async (input: ToolRouterInput, execOptions?: ToolExecuteOptions) => { + const routerToolRef: { current?: ToolRouter } = {}; + + const execute = async ( + input: ToolRouterInput, + execOptions?: ToolExecuteOptions, + ): Promise => { const agent = execOptions?.systemContext?.get(AGENT_REF_CONTEXT_KEY) as Agent | undefined; const executor = agent && (agent as any).__executeToolRouter; if (typeof executor !== "function") { throw new Error("Tool router requires an Agent execution context."); } + const router = routerToolRef.current; + if (!router) { + throw new Error("Tool router is not initialized."); + } + return await executor.call(agent, { - router: tool as ToolRouter, + router, input, options: execOptions, }); @@ -74,6 +85,7 @@ export const createToolRouter = (options: CreateToolRouterOptions): ToolRouter = const routerTool = tool as ToolRouter; routerTool[TOOL_ROUTER_SYMBOL] = metadata; + routerToolRef.current = routerTool; return routerTool; }; From 068933e591f17f3626ebaf0ca5a394d042cf6b72 Mon Sep 17 00:00:00 2001 From: Omer Aplak Date: Thu, 22 Jan 2026 19:14:09 -0800 Subject: [PATCH 6/7] chore: fix code reviews --- packages/core/src/agent/agent.ts | 122 ++++++++++++++++++++++++------- 1 file changed, 96 insertions(+), 26 deletions(-) diff --git a/packages/core/src/agent/agent.ts b/packages/core/src/agent/agent.ts index 413dd44c6..2eb0c05fc 100644 --- a/packages/core/src/agent/agent.ts +++ b/packages/core/src/agent/agent.ts @@ -4599,6 +4599,26 @@ export class Agent { this.toolManager.prepareToolsForExecution(createToolExecuteFunction); const toolRouting = this.resolveToolRouting(options); + if (toolRouting === false) { + const routerNames = new Set(); + this.toolManager + .getAllBaseTools() + .filter((tool) => isToolRouter(tool)) + .forEach((tool) => routerNames.add(tool.name)); + runtimeTools + .filter((tool) => isToolRouter(tool)) + .forEach((tool) => routerNames.add(tool.name)); + + const filteredStaticTools = Object.fromEntries( + Object.entries(preparedStaticTools).filter(([name]) => !routerNames.has(name)), + ); + const filteredDynamicTools = Object.fromEntries( + Object.entries(preparedDynamicTools).filter(([name]) => !routerNames.has(name)), + ); + + return { ...filteredStaticTools, ...filteredDynamicTools }; + } + if (!toolRouting) { return { ...preparedStaticTools, ...preparedDynamicTools }; } @@ -5132,6 +5152,43 @@ export class Agent { return (result.output ?? {}) as Record; } + private async resolveProviderToolArgs(params: { + tool: ProviderTool; + query: string; + oc: OperationContext; + executionModel?: AgentModelValue; + }): Promise> { + const { tool, query, oc, executionModel } = params; + const schema = tool.args as unknown; + const prompt = [ + "Generate JSON arguments for the tool based on the user request.", + `Tool name: ${tool.name}`, + tool.description ? `Tool description: ${tool.description}` : "", + schema ? `Tool schema: ${safeStringify(schema)}` : "", + `User request: ${query}`, + "Return only a JSON object that matches the tool schema.", + ] + .filter(Boolean) + .join("\n"); + + const result = await this.runInternalGenerateText({ + oc, + modelValue: executionModel, + messages: [ + { + role: "system", + content: "You generate tool arguments that strictly match the provided schema.", + }, + { role: "user", content: prompt }, + ], + output: Output.object({ schema: schema as any }), + toolChoice: "none", + temperature: 0, + }); + + return (result.output ?? {}) as Record; + } + private async executeProviderToolViaRouter(params: { tool: ProviderTool; query: string; @@ -5159,25 +5216,58 @@ export class Agent { const needsApproval = (tool as { needsApproval?: Tool["needsApproval"] }) .needsApproval; if (needsApproval === true) { - throw new ToolDeniedError({ + return { toolName: tool.name, - message: `Tool ${tool.name} requires approval.`, - code: "TOOL_FORBIDDEN", - httpStatus: 403, - }); + toolCallId, + error: `Tool ${tool.name} requires approval.`, + }; + } + + let approvedArgs: Record | undefined; + if (needsApproval) { + try { + approvedArgs = await this.resolveProviderToolArgs({ + tool, + query, + oc, + executionModel, + }); + await this.ensureToolApproval(tool, approvedArgs, executionOptions, toolCallId); + } catch (error) { + if (isToolDeniedError(error)) { + return { + toolName: tool.name, + toolCallId, + error: error.message, + }; + } + return { + toolName: tool.name, + toolCallId, + error: error instanceof Error ? error.message : String(error), + }; + } } const tools: Record = { [tool.name]: tool, }; + const approvedArgsInstruction = approvedArgs + ? `Use these tool arguments: ${safeStringify(approvedArgs)}` + : ""; const result = await this.runInternalGenerateText({ oc, modelValue: executionModel, messages: [ { role: "system", - content: "Call the required tool with appropriate arguments to satisfy the request.", + content: [ + "Call the required tool with appropriate arguments to satisfy the request.", + approvedArgsInstruction, + ] + .filter(Boolean) + .join("\n"), }, { role: "user", content: query }, ], @@ -5196,26 +5286,6 @@ export class Agent { executionOptions.toolContext.callId = toolCall.toolCallId; } - if (toolCall) { - try { - await this.ensureToolApproval( - tool, - toolCall.input as Record, - executionOptions, - toolCall.toolCallId ?? toolCallId, - ); - } catch (error) { - if (isToolDeniedError(error)) { - return { - toolName: tool.name, - toolCallId: toolCall.toolCallId ?? toolCallId, - error: error.message, - }; - } - throw error; - } - } - if (toolCall) { await hooks.onToolStart?.({ agent: this, From efc8815e8afa6e15c9a5ca8c0fbbb745ae90ff92 Mon Sep 17 00:00:00 2001 From: Omer Aplak Date: Thu, 22 Jan 2026 19:48:48 -0800 Subject: [PATCH 7/7] fix: code reviews --- packages/core/src/agent/agent.ts | 12 ++++++--- .../memory/adapters/embedding/ai-sdk.spec.ts | 8 +++--- .../src/memory/adapters/embedding/ai-sdk.ts | 25 ++++++++++++++----- packages/core/src/tool/routing/embedding.ts | 2 +- website/docs/getting-started/comparison.mdx | 8 ++++++ website/docs/tools/tool-routing.md | 1 - website/recipes/tool-routing.md | 12 ++++++++- 7 files changed, 53 insertions(+), 15 deletions(-) diff --git a/packages/core/src/agent/agent.ts b/packages/core/src/agent/agent.ts index 243d7622c..5423428e7 100644 --- a/packages/core/src/agent/agent.ts +++ b/packages/core/src/agent/agent.ts @@ -93,7 +93,7 @@ import { type EnqueueEvalScoringArgs, enqueueEvalScoring as enqueueEvalScoringHelper, } from "./eval"; -import type { AgentHooks } from "./hooks"; +import type { AgentHooks, OnToolEndHookResult } from "./hooks"; import { AgentTraceContext, addModelAttributesToSpan } from "./open-telemetry/trace-context"; import type { BaseMessage, @@ -5761,8 +5761,14 @@ export class Agent { await this.hooks.onToolStart?.(...args); }, onToolEnd: async (...args) => { - await options.hooks?.onToolEnd?.(...args); - await this.hooks.onToolEnd?.(...args); + const resOptions = await options.hooks?.onToolEnd?.(...args); + const resThis = await this.hooks.onToolEnd?.(...args); + if (resThis && typeof resThis === "object") { + return resThis as OnToolEndHookResult; + } + if (resOptions && typeof resOptions === "object") { + return resOptions as OnToolEndHookResult; + } return undefined; }, onStepFinish: async (...args) => { diff --git a/packages/core/src/memory/adapters/embedding/ai-sdk.spec.ts b/packages/core/src/memory/adapters/embedding/ai-sdk.spec.ts index 52782aa35..bf5947910 100644 --- a/packages/core/src/memory/adapters/embedding/ai-sdk.spec.ts +++ b/packages/core/src/memory/adapters/embedding/ai-sdk.spec.ts @@ -1,5 +1,5 @@ import type { EmbedManyResult, EmbedResult, EmbeddingModel } from "ai"; -import { beforeEach, describe, expect, it, vi } from "vitest"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; import { ModelProviderRegistry } from "../../../registries/model-provider-registry"; import { AiSdkEmbeddingAdapter } from "./ai-sdk"; @@ -13,6 +13,10 @@ describe("AiSdkEmbeddingAdapter", () => { let mockModel: Exclude; let adapter: AiSdkEmbeddingAdapter; + afterEach(() => { + vi.restoreAllMocks(); + }); + beforeEach(() => { vi.clearAllMocks(); @@ -209,8 +213,6 @@ describe("AiSdkEmbeddingAdapter", () => { value: "test text", }); expect(stringAdapter.getModelName()).toBe("openai/text-embedding-3-small"); - - resolveSpy.mockRestore(); }); }); }); diff --git a/packages/core/src/memory/adapters/embedding/ai-sdk.ts b/packages/core/src/memory/adapters/embedding/ai-sdk.ts index 523862039..82ed807c3 100644 --- a/packages/core/src/memory/adapters/embedding/ai-sdk.ts +++ b/packages/core/src/memory/adapters/embedding/ai-sdk.ts @@ -3,6 +3,10 @@ import type { EmbeddingModel } from "ai"; import { ModelProviderRegistry } from "../../../registries/model-provider-registry"; import type { EmbeddingAdapter, EmbeddingModelReference, EmbeddingOptions } from "./types"; +const BARE_MODEL_NAME_REGEX = /^[A-Za-z0-9][A-Za-z0-9_.-]*$/; + +const isValidBareModelName = (value: string): boolean => BARE_MODEL_NAME_REGEX.test(value); + /** * AI SDK Embedding Adapter * Wraps Vercel AI SDK embedding models for use with Memory V2 @@ -15,11 +19,17 @@ export class AiSdkEmbeddingAdapter implements EmbeddingAdapter { private modelResolvePromise?: Promise; constructor(model: EmbeddingModelReference, options: EmbeddingOptions = {}) { - const normalizedModel = typeof model === "string" ? model.trim() : model; - this.model = normalizedModel; - // EmbeddingModel can be either a string or an object with modelId - this.modelName = - typeof normalizedModel === "string" ? normalizedModel : normalizedModel.modelId; + if (typeof model === "string") { + const trimmed = model.trim(); + if (!trimmed) { + throw new Error("Embedding model is required."); + } + this.model = trimmed; + this.modelName = trimmed; + } else { + this.model = model; + this.modelName = model.modelId; + } this.dimensions = 0; // Will be set after first embedding this.options = { maxBatchSize: options.maxBatchSize ?? 100, @@ -38,11 +48,14 @@ export class AiSdkEmbeddingAdapter implements EmbeddingAdapter { const trimmed = this.model.trim(); if (!trimmed) { - return trimmed; + throw new Error("Embedding model is required."); } const hasProviderPrefix = trimmed.includes("/") || trimmed.includes(":"); if (!hasProviderPrefix) { + if (!isValidBareModelName(trimmed)) { + throw new Error(`Invalid embedding model id "${trimmed}".`); + } this.model = trimmed; this.modelName = trimmed; return trimmed; diff --git a/packages/core/src/tool/routing/embedding.ts b/packages/core/src/tool/routing/embedding.ts index a0e8eefb9..178fcb49f 100644 --- a/packages/core/src/tool/routing/embedding.ts +++ b/packages/core/src/tool/routing/embedding.ts @@ -48,7 +48,7 @@ const defaultToolText = (tool: ToolRouterCandidate): string => { } if (tool.parameters) { const normalized = - typeof tool.parameters === "object" ? tool.parameters : zodSchemaToJsonUI(tool.parameters); + typeof tool.parameters === "object" ? zodSchemaToJsonUI(tool.parameters) : tool.parameters; parts.push(`Parameters: ${safeStringify(normalized)}`); } return parts.join("\n"); diff --git a/website/docs/getting-started/comparison.mdx b/website/docs/getting-started/comparison.mdx index e4ad8ced6..481f283a1 100644 --- a/website/docs/getting-started/comparison.mdx +++ b/website/docs/getting-started/comparison.mdx @@ -57,6 +57,14 @@ This comparison table strives to be as accurate and as unbiased as possible. If aiSdk: { status: "ready" }, aiSdkTools: { status: "ready" }, }, + { + feature: "Tool Routing", + link: "/docs/tools/tool-routing/", + voltagent: { status: "ready" }, + mastra: { status: "not-supported" }, + aiSdk: { status: "not-supported" }, + aiSdkTools: { status: "not-supported" }, + }, { feature: "Working Memory", voltagent: { status: "ready" }, diff --git a/website/docs/tools/tool-routing.md b/website/docs/tools/tool-routing.md index 77f8d085d..a08506db6 100644 --- a/website/docs/tools/tool-routing.md +++ b/website/docs/tools/tool-routing.md @@ -239,7 +239,6 @@ The embedding router accepts several model forms and supports a custom tool text - `"openai/text-embedding-3-small"` - `"text-embedding-3-small"` -- `"openai/text-embedding-3-small"` Provider-qualified strings use the same model registry and type list as agent model strings. diff --git a/website/recipes/tool-routing.md b/website/recipes/tool-routing.md index 4bad78d0d..d3f7460d8 100644 --- a/website/recipes/tool-routing.md +++ b/website/recipes/tool-routing.md @@ -124,10 +124,20 @@ new VoltAgent({ ## Pooling Provider and MCP Tools ```typescript +import { MCPConfiguration } from "@voltagent/core"; +import { openai } from "@ai-sdk/openai"; + +const mcp = new MCPConfiguration({ + servers: [{ name: "zapier", url: process.env.ZAPIER_MCP_URL ?? "" }], +}); + +// MCP docs: /docs/agents/mcp/mcp/ +const mcpTools = await mcp.getTools(); + toolRouting: { pool: [ openai.tools.webSearch(), - mcpToolkit, + ...mcpTools, ], } ```