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 + + +
+
+ +
+ Home Page | + Documentation | + Examples | + Discord | + Blog +
+
+ +
+ +
+ 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 7aed83d1f..5423428e7 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, @@ -80,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, @@ -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,37 @@ export class Agent { const preparedStaticTools = this.toolManager.prepareToolsForExecution(createToolExecuteFunction); - return { ...preparedStaticTools, ...preparedDynamicTools }; + 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 }; + } + + const exposedNames = this.getToolRoutingExposedNames(toolRouting); + const filteredStaticTools = Object.fromEntries( + Object.entries(preparedStaticTools).filter(([name]) => exposedNames.has(name)), + ); + + return { ...filteredStaticTools, ...preparedDynamicTools }; } /** @@ -4614,6 +4674,7 @@ export class Agent { abortSignal: abortSignal, }, }; + executionOptions.hooks = hooks; // Event tracking now handled by OpenTelemetry spans const toolTags = (tool as { tags?: string[] | undefined }).tags; @@ -4652,7 +4713,6 @@ export class Agent { args, options: executionOptions, }); - await hooks.onToolStart?.({ agent: this, tool, @@ -4746,6 +4806,14 @@ export class Agent { options: executionOptions, }); + await tool.hooks?.onEnd?.({ + tool, + args, + output: undefined, + error: voltAgentError, + options: executionOptions, + }); + await hooks.onToolEnd?.({ agent: this, tool, @@ -4857,6 +4925,522 @@ 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 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; + 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) { + return { + toolName: tool.name, + 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.", + approvedArgsInstruction, + ] + .filter(Boolean) + .join("\n"), + }, + { 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) { + 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 */ @@ -5177,8 +5761,15 @@ 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) => { await options.hooks?.onStepFinish?.(...args); @@ -5478,6 +6069,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, @@ -5488,10 +6106,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, @@ -5543,6 +6173,9 @@ export class Agent { added: (Tool | Toolkit | VercelTool)[]; } { this.toolManager.addItems(tools); + if (this.toolRouting && !this.toolRoutingPoolExplicit) { + this.toolPoolManager.addItems(tools); + } return { added: tools }; } @@ -5556,6 +6189,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); + } } } @@ -5574,6 +6210,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}`); @@ -5596,6 +6235,9 @@ export class Agent { sourceAgent: this as any, }); this.toolManager.addStandaloneTool(delegateTool); + if (this.toolRouting && !this.toolRoutingPoolExplicit) { + this.toolPoolManager.addStandaloneTool(delegateTool); + } } } @@ -5608,6 +6250,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"); + } } } @@ -5622,7 +6267,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()); } /** @@ -5693,6 +6346,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 f14e64bf9..1e28d5c31 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..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,6 @@ 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"; // Mock the AI SDK @@ -9,9 +10,13 @@ vi.mock("ai", () => ({ })); describe("AiSdkEmbeddingAdapter", () => { - let mockModel: EmbeddingModel; + let mockModel: Exclude; let adapter: AiSdkEmbeddingAdapter; + afterEach(() => { + vi.restoreAllMocks(); + }); + beforeEach(() => { vi.clearAllMocks(); @@ -19,7 +24,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 +182,37 @@ 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"); + }); + }); }); diff --git a/packages/core/src/memory/adapters/embedding/ai-sdk.ts b/packages/core/src/memory/adapters/embedding/ai-sdk.ts index 4d2ce387d..82ed807c3 100644 --- a/packages/core/src/memory/adapters/embedding/ai-sdk.ts +++ b/packages/core/src/memory/adapters/embedding/ai-sdk.ts @@ -1,5 +1,11 @@ -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"; + +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 @@ -10,11 +16,20 @@ 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; - // EmbeddingModel can be either a string or an object with modelId - this.modelName = typeof model === "string" ? model : model.modelId; + constructor(model: EmbeddingModelReference, options: EmbeddingOptions = {}) { + 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, @@ -23,10 +38,48 @@ 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) { + 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; + } + + 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 +108,7 @@ export class AiSdkEmbeddingAdapter implements EmbeddingAdapter { return []; } + const model = await this.resolveModel(); const maxBatchSize = this.options.maxBatchSize ?? 100; const embeddings: number[][] = []; @@ -64,7 +118,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 a84337e4d..0daa95b84 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..178fcb49f --- /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" ? zodSchemaToJsonUI(tool.parameters) : 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..70a1bcb9a --- /dev/null +++ b/packages/core/src/tool/routing/index.ts @@ -0,0 +1,98 @@ +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, + ToolRouterResult, + 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 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, + 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; + routerToolRef.current = routerTool; + + 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/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/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..a08506db6 --- /dev/null +++ b/website/docs/tools/tool-routing.md @@ -0,0 +1,311 @@ +--- +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"` + +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..d3f7460d8 --- /dev/null +++ b/website/recipes/tool-routing.md @@ -0,0 +1,143 @@ +--- +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 +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(), + ...mcpTools, + ], +} +``` 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",