From 244d8bc4dc0954ca9d94813d8113592f974d3eca Mon Sep 17 00:00:00 2001 From: Brian Yin Date: Fri, 29 May 2026 05:43:58 +0800 Subject: [PATCH 1/6] Refactor ToolContext to parity class taking a list of Tool | Toolset (#1517) --- .changeset/list-syntax-toolcontext.md | 5 + agents/src/beta/workflows/task_group.ts | 6 +- agents/src/inference/llm.ts | 57 +++--- agents/src/llm/chat_context.test.ts | 29 +++ agents/src/llm/chat_context.ts | 2 +- agents/src/llm/fallback_adapter.test.ts | 6 +- agents/src/llm/fallback_adapter.ts | 6 +- agents/src/llm/index.ts | 7 +- agents/src/llm/llm.ts | 18 +- agents/src/llm/tool_context.test.ts | 177 ++++++++++++++++- agents/src/llm/tool_context.ts | 185 ++++++++++++++++-- agents/src/llm/tool_context.type.test.ts | 4 + agents/src/llm/utils.ts | 4 +- agents/src/voice/agent.test.ts | 29 ++- agents/src/voice/agent.ts | 19 +- agents/src/voice/agent_activity.test.ts | 6 +- agents/src/voice/agent_activity.ts | 45 +++-- agents/src/voice/amd.test.ts | 8 +- agents/src/voice/amd.ts | 12 +- agents/src/voice/generation.ts | 7 +- agents/src/voice/generation_tools.test.ts | 18 +- agents/src/voice/remote_session.ts | 2 +- agents/src/voice/testing/fake_llm.ts | 6 +- agents/src/voice/testing/run_result.ts | 3 +- examples/src/background_audio.ts | 5 +- examples/src/basic_agent.ts | 7 +- examples/src/basic_agent_task.ts | 27 +-- examples/src/basic_task_group.ts | 21 +- examples/src/basic_tool_call_agent.ts | 19 +- examples/src/comprehensive_test.ts | 17 +- examples/src/drive-thru/drivethru_agent.ts | 19 +- examples/src/frontdesk/frontdesk_agent.ts | 10 +- examples/src/gemini_realtime_agent.ts | 11 +- examples/src/instructions_per_modality.ts | 7 +- examples/src/llm_fallback_adapter.ts | 7 +- examples/src/manual_shutdown.ts | 10 +- examples/src/multi_agent.ts | 9 +- examples/src/phonic_realtime_agent.ts | 5 +- examples/src/raw_function_description.ts | 7 +- examples/src/realtime_agent.ts | 7 +- examples/src/realtime_with_tts.ts | 5 +- examples/src/restaurant_agent.ts | 78 +++++--- examples/src/survey_agent.ts | 49 +++-- examples/src/testing/agent_task.test.ts | 28 +-- examples/src/testing/basic_task_group.test.ts | 14 +- examples/src/testing/run_result.test.ts | 19 +- examples/src/testing/task_group.test.ts | 28 +-- examples/src/tool_call_disfluency.ts | 5 +- examples/src/xai-realtime.ts | 7 +- plugins/baseten/src/llm.ts | 11 +- plugins/cerebras/src/llm.test.ts | 7 +- .../google/src/beta/realtime/realtime_api.ts | 4 +- plugins/google/src/llm.ts | 7 +- plugins/google/src/utils.ts | 2 +- plugins/mistralai/src/llm.ts | 8 +- plugins/openai/src/llm.ts | 11 +- plugins/openai/src/realtime/realtime_model.ts | 23 +-- plugins/openai/src/responses/llm.ts | 15 +- plugins/openai/src/ws/llm.ts | 13 +- plugins/phonic/src/realtime/realtime_model.ts | 72 ++++--- plugins/test/src/llm.ts | 46 +++-- 61 files changed, 901 insertions(+), 400 deletions(-) create mode 100644 .changeset/list-syntax-toolcontext.md diff --git a/.changeset/list-syntax-toolcontext.md b/.changeset/list-syntax-toolcontext.md new file mode 100644 index 000000000..0a4e0bfa0 --- /dev/null +++ b/.changeset/list-syntax-toolcontext.md @@ -0,0 +1,5 @@ +--- +"@livekit/agents": minor +--- + +**BREAKING**: `Agent({ tools })` and `agent.updateTools()` now accept a flat list `(FunctionTool | ProviderDefinedTool)[]` instead of a `Record` map, and `llm.tool({ ... })` requires a `name` field. `ToolContext` is now a Python-parity class with `functionTools` / `providerTools` / `toolsets` accessors, plus `flatten()`, `hasTool(name)`, `getFunctionTool(name)`, `updateTools()`, `copy()`, and `equals()`. To match the Python reference, registering two **different** function-tool instances under the same `name` now throws `duplicate function name: ` instead of silently overriding the earlier entry; passing the **same instance** twice is a no-op. `agent.toolCtx` returns a defensive copy so callers can no longer mutate the agent's internal state. `LLM.chat({ toolCtx })` accepts either a `ToolContext` instance or a raw `(FunctionTool | ProviderDefinedTool)[]` array (`ToolCtxInput`) and normalizes it internally, so callers don't have to construct a `ToolContext` themselves. Stateful `Toolset` containers are not part of this release — the `toolsets` accessor currently returns an empty list and `TODO`s in `tool_context.ts` mark every site where Python's Toolset support will plug in later. diff --git a/agents/src/beta/workflows/task_group.ts b/agents/src/beta/workflows/task_group.ts index 8c96790dd..add6655fb 100644 --- a/agents/src/beta/workflows/task_group.ts +++ b/agents/src/beta/workflows/task_group.ts @@ -84,10 +84,7 @@ export class TaskGroup extends AgentTask { const outOfScopeTool = this.buildOutOfScopeTool(taskId); if (outOfScopeTool) { - await this._currentTask.updateTools({ - ...this._currentTask.toolCtx, - out_of_scope: outOfScopeTool, - }); + await this._currentTask.updateTools([...this._currentTask.toolCtx.tools, outOfScopeTool]); } try { @@ -190,6 +187,7 @@ export class TaskGroup extends AgentTask { const visitedTasks = this._visitedTasks; return tool({ + name: 'out_of_scope', description, flags: ToolFlag.IGNORE_ON_ENTER, parameters: z.object({ diff --git a/agents/src/inference/llm.ts b/agents/src/inference/llm.ts index 3434e496c..87679389f 100644 --- a/agents/src/inference/llm.ts +++ b/agents/src/inference/llm.ts @@ -249,7 +249,7 @@ export class LLM extends llm.LLM { chat({ chatCtx, - toolCtx, + toolCtx: toolCtxInput, connOptions = DEFAULT_API_CONNECT_OPTIONS, parallelToolCalls, toolChoice, @@ -258,7 +258,7 @@ export class LLM extends llm.LLM { extraKwargs, }: { chatCtx: llm.ChatContext; - toolCtx?: llm.ToolContext; + toolCtx?: llm.ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: llm.ToolChoice; @@ -266,6 +266,7 @@ export class LLM extends llm.LLM { // TODO(AJS-270): Add responseFormat parameter extraKwargs?: Record; }): LLMStream { + const toolCtx = llm.toToolContext(toolCtxInput); let modelOptions: Record = { ...(extraKwargs || {}) }; parallelToolCalls = @@ -273,7 +274,11 @@ export class LLM extends llm.LLM { ? parallelToolCalls : this.opts.modelOptions.parallel_tool_calls; - if (toolCtx && Object.keys(toolCtx).length > 0 && parallelToolCalls !== undefined) { + if ( + toolCtx && + Object.keys(toolCtx.functionTools).length > 0 && + parallelToolCalls !== undefined + ) { modelOptions.parallel_tool_calls = parallelToolCalls; } @@ -379,26 +384,32 @@ export class LLMStream extends llm.LLMStream { )) as OpenAI.ChatCompletionMessageParam[]; const tools = this.toolCtx - ? Object.entries(this.toolCtx).map(([name, func]) => { - const oaiParams = { - type: 'function' as const, - function: { - name, - description: func.description, - parameters: llm.toJsonSchema( - func.parameters, - true, - this.strictToolSchema, - ) as unknown as OpenAI.Chat.Completions.ChatCompletionFunctionTool['function']['parameters'], - } as OpenAI.Chat.Completions.ChatCompletionFunctionTool['function'], - }; - - if (this.strictToolSchema) { - oaiParams.function.strict = true; - } - - return oaiParams; - }) + ? this.toolCtx + .flatten() + .map((t) => { + if (llm.isFunctionTool(t)) { + const oaiParams = { + type: 'function' as const, + function: { + name: t.name, + description: t.description, + parameters: llm.toJsonSchema( + t.parameters, + true, + this.strictToolSchema, + ) as unknown as OpenAI.Chat.Completions.ChatCompletionFunctionTool['function']['parameters'], + } as OpenAI.Chat.Completions.ChatCompletionFunctionTool['function'], + }; + if (this.strictToolSchema) { + oaiParams.function.strict = true; + } + return oaiParams; + } + // Provider-defined tools are not yet supported by the inference adapter; skip them + // here rather than emitting a malformed tool definition. See AJS-112. + return undefined; + }) + .filter((t): t is NonNullable => t !== undefined) : undefined; const requestOptions: Record = dropUnsupportedParams( diff --git a/agents/src/llm/chat_context.test.ts b/agents/src/llm/chat_context.test.ts index 849469e55..350dcd1cd 100644 --- a/agents/src/llm/chat_context.test.ts +++ b/agents/src/llm/chat_context.test.ts @@ -19,6 +19,7 @@ import { isInstructions, renderInstructions, } from './chat_context.js'; +import { ToolContext, tool } from './tool_context.js'; initializeLogger({ pretty: false, level: 'error' }); @@ -1479,3 +1480,31 @@ extra`; expect((baseCtx.items[0]! as ChatMessage).content[0]).toBe(instr); }); }); + +describe('ChatContext.copy with toolCtx filter', () => { + it('drops function calls / outputs whose tool is not in the supplied ToolContext', () => { + const known = tool({ name: 'known', description: 'k', execute: async () => 'ok' }); + const ctx = new ChatContext([ + ChatMessage.create({ role: 'user', content: ['hello'] }), + FunctionCall.create({ callId: 'c1', name: 'known', args: '{}' }), + FunctionCallOutput.create({ callId: 'c1', name: 'known', output: 'done', isError: false }), + FunctionCall.create({ callId: 'c2', name: 'removed', args: '{}' }), + FunctionCallOutput.create({ callId: 'c2', name: 'removed', output: 'x', isError: false }), + ]); + + const filtered = ctx.copy({ toolCtx: new ToolContext([known]) }); + const types = filtered.items.map((i) => `${i.type}:${'name' in i ? i.name : ''}`); + expect(types).toEqual(['message:', 'function_call:known', 'function_call_output:known']); + }); + + it('keeps provider-tool calls when the ToolContext holds a matching provider tool id', () => { + const provider = tool({ id: 'code_runner', config: {} }); + const ctx = new ChatContext([ + FunctionCall.create({ callId: 'p1', name: 'code_runner', args: '{}' }), + FunctionCall.create({ callId: 'p2', name: 'other', args: '{}' }), + ]); + + const filtered = ctx.copy({ toolCtx: new ToolContext([provider]) }); + expect(filtered.items.map((i) => ('name' in i ? i.name : ''))).toEqual(['code_runner']); + }); +}); diff --git a/agents/src/llm/chat_context.ts b/agents/src/llm/chat_context.ts index 743e7efb8..34cf39200 100644 --- a/agents/src/llm/chat_context.ts +++ b/agents/src/llm/chat_context.ts @@ -835,7 +835,7 @@ export class ChatContext { continue; } - if (toolCtx !== undefined && isToolCallOrOutput(item) && toolCtx[item.name] === undefined) { + if (toolCtx !== undefined && isToolCallOrOutput(item) && !toolCtx.hasTool(item.name)) { continue; } diff --git a/agents/src/llm/fallback_adapter.test.ts b/agents/src/llm/fallback_adapter.test.ts index a9747c885..d466c9306 100644 --- a/agents/src/llm/fallback_adapter.test.ts +++ b/agents/src/llm/fallback_adapter.test.ts @@ -9,7 +9,7 @@ import { delay } from '../utils.js'; import type { ChatContext } from './chat_context.js'; import { FallbackAdapter } from './fallback_adapter.js'; import { type ChatChunk, LLM, LLMStream } from './llm.js'; -import type { ToolChoice, ToolContext } from './tool_context.js'; +import type { ToolChoice, ToolCtxInput } from './tool_context.js'; class MockLLMStream extends LLMStream { public myLLM: LLM; @@ -18,7 +18,7 @@ class MockLLMStream extends LLMStream { llm: LLM, opts: { chatCtx: ChatContext; - toolCtx?: ToolContext; + toolCtx?: ToolCtxInput; connOptions: APIConnectOptions; }, private shouldFail: boolean = false, @@ -64,7 +64,7 @@ class MockLLM extends LLM { chat(opts: { chatCtx: ChatContext; - toolCtx?: ToolContext; + toolCtx?: ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: ToolChoice; diff --git a/agents/src/llm/fallback_adapter.ts b/agents/src/llm/fallback_adapter.ts index 128c2392c..27d87d2a0 100644 --- a/agents/src/llm/fallback_adapter.ts +++ b/agents/src/llm/fallback_adapter.ts @@ -8,7 +8,7 @@ import { type APIConnectOptions, DEFAULT_API_CONNECT_OPTIONS } from '../types.js import type { ChatContext } from './chat_context.js'; import type { ChatChunk } from './llm.js'; import { LLM, LLMStream } from './llm.js'; -import type { ToolChoice, ToolContext } from './tool_context.js'; +import type { ToolChoice, ToolCtxInput } from './tool_context.js'; /** * Default connection options for FallbackAdapter. @@ -113,7 +113,7 @@ export class FallbackAdapter extends LLM { chat(opts: { chatCtx: ChatContext; - toolCtx?: ToolContext; + toolCtx?: ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: ToolChoice; @@ -159,7 +159,7 @@ class FallbackLLMStream extends LLMStream { adapter: FallbackAdapter, opts: { chatCtx: ChatContext; - toolCtx?: ToolContext; + toolCtx?: ToolCtxInput; connOptions: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: ToolChoice; diff --git a/agents/src/llm/index.ts b/agents/src/llm/index.ts index 950296c9c..4837f2cb9 100644 --- a/agents/src/llm/index.ts +++ b/agents/src/llm/index.ts @@ -4,15 +4,20 @@ export { handoff, isFunctionTool, + isProviderDefinedTool, + isTool, tool, + ToolContext, ToolError, ToolFlag, + toToolContext, type AgentHandoff, type FunctionTool, type ProviderDefinedTool, type Tool, type ToolChoice, - type ToolContext, + type ToolContextEntry, + type ToolCtxInput, type ToolOptions, type ToolType, } from './tool_context.js'; diff --git a/agents/src/llm/llm.ts b/agents/src/llm/llm.ts index 553541939..0c05bbb2d 100644 --- a/agents/src/llm/llm.ts +++ b/agents/src/llm/llm.ts @@ -11,7 +11,12 @@ import { recordException, traceTypes, tracer } from '../telemetry/index.js'; import { type APIConnectOptions, intervalForRetry } from '../types.js'; import { AsyncIterableQueue, delay, startSoon, toError } from '../utils.js'; import { type ChatContext, type ChatRole, type FunctionCall } from './chat_context.js'; -import type { ToolChoice, ToolContext } from './tool_context.js'; +import { + type ToolChoice, + type ToolContext, + type ToolCtxInput, + toToolContext, +} from './tool_context.js'; export interface ChoiceDelta { role: ChatRole; @@ -91,7 +96,12 @@ export abstract class LLM extends (EventEmitter as new () => TypedEmitter { connOptions, }: { chatCtx: ChatContext; - toolCtx?: ToolContext; + toolCtx?: ToolCtxInput; connOptions: APIConnectOptions; }, ) { this.#llm = llm; this.#chatCtx = chatCtx; - this.#toolCtx = toolCtx; + this.#toolCtx = toToolContext(toolCtx); this._connOptions = connOptions; this.monitorMetrics(); this.abortController.signal.addEventListener('abort', () => { diff --git a/agents/src/llm/tool_context.test.ts b/agents/src/llm/tool_context.test.ts index 183c7ff2a..4d9cdb83d 100644 --- a/agents/src/llm/tool_context.test.ts +++ b/agents/src/llm/tool_context.test.ts @@ -5,7 +5,7 @@ import { describe, expect, it } from 'vitest'; import { z } from 'zod'; import * as z3 from 'zod/v3'; import * as z4 from 'zod/v4'; -import { type ToolOptions, tool } from './tool_context.js'; +import { ToolContext, type ToolOptions, tool } from './tool_context.js'; import { createToolOptions, oaiParams } from './utils.js'; describe('Tool Context', () => { @@ -77,6 +77,7 @@ describe('Tool Context', () => { describe('tool', () => { it('should create and execute a basic core tool', async () => { const getWeather = tool({ + name: 'getWeather', description: 'Get the weather for a given location', parameters: z.object({ location: z.string(), @@ -95,6 +96,7 @@ describe('Tool Context', () => { it('should properly type a callable function', async () => { const testFunction = tool({ + name: 'testFunction', description: 'Test function', parameters: z.object({ name: z.string().describe('The user name'), @@ -114,6 +116,7 @@ describe('Tool Context', () => { it('should handle async execution', async () => { const testFunction = tool({ + name: 'asyncTestFunction', description: 'Async test function', parameters: z.object({ delay: z.number().describe('Delay in milliseconds'), @@ -157,6 +160,7 @@ describe('Tool Context', () => { describe('optional parameters', () => { it('should create a tool without parameters', async () => { const simpleAction = tool({ + name: 'simpleAction', description: 'Perform a simple action', execute: async () => { return 'Action performed'; @@ -175,6 +179,7 @@ describe('Tool Context', () => { it('should support .optional() fields in tool parameters', async () => { const weatherTool = tool({ + name: 'weatherTool', description: 'Get weather information', parameters: z.object({ location: z.string().describe('The city or location').optional(), @@ -205,6 +210,7 @@ describe('Tool Context', () => { it('should handle tools with context but no parameters', async () => { const greetUser = tool({ + name: 'greetUser', description: 'Greet the current user', execute: async (_, { ctx }: ToolOptions<{ username: string }>) => { return `Hello, ${ctx.userData.username}!`; @@ -217,6 +223,7 @@ describe('Tool Context', () => { it('should create a tool that accesses tool call id without parameters', async () => { const getCallId = tool({ + name: 'getCallId', description: 'Get the current tool call ID', execute: async (_, { toolCallId }) => { return `Tool call ID: ${toolCallId}`; @@ -231,6 +238,7 @@ describe('Tool Context', () => { describe('Zod v3 and v4 compatibility', () => { it('should work with Zod v3 schemas', async () => { const v3Tool = tool({ + name: 'v3Tool', description: 'A tool using Zod v3 schema', parameters: z3.object({ name: z3.string(), @@ -250,6 +258,7 @@ describe('Tool Context', () => { it('should work with Zod v4 schemas', async () => { const v4Tool = tool({ + name: 'v4Tool', description: 'A tool using Zod v4 schema', parameters: z4.object({ name: z4.string(), @@ -269,6 +278,7 @@ describe('Tool Context', () => { it('should handle v4 schemas with optional fields', async () => { const v4Tool = tool({ + name: 'v4OptionalTool', description: 'Tool with optional field using v4', parameters: z4.object({ required: z4.string(), @@ -291,6 +301,7 @@ describe('Tool Context', () => { it('should handle v4 enum schemas', async () => { const v4Tool = tool({ + name: 'v4EnumTool', description: 'Tool with enum using v4', parameters: z4.object({ color: z4.enum(['red', 'blue', 'green']), @@ -306,6 +317,7 @@ describe('Tool Context', () => { it('should handle v4 array schemas', async () => { const v4Tool = tool({ + name: 'v4ArrayTool', description: 'Tool with array using v4', parameters: z4.object({ tags: z4.array(z4.string()), @@ -324,6 +336,7 @@ describe('Tool Context', () => { it('should handle v4 nested object schemas', async () => { const v4Tool = tool({ + name: 'v4NestedTool', description: 'Tool with nested object using v4', parameters: z4.object({ user: z4.object({ @@ -405,3 +418,165 @@ describe('Tool Context', () => { }); }); }); + +describe('tool() name requirement', () => { + it('throws when name is missing', () => { + expect(() => + // @ts-expect-error - name is required + tool({ + description: 'no name', + execute: async () => 'x', + }), + ).toThrow('requires a non-empty name'); + }); + + it('throws when name is empty', () => { + expect(() => + tool({ + name: '', + description: 'empty name', + execute: async () => 'x', + }), + ).toThrow('requires a non-empty name'); + }); + + it('stores the name on the returned function tool', () => { + const t = tool({ + name: 'doStuff', + description: 'd', + execute: async () => 'x', + }); + expect(t.name).toBe('doStuff'); + }); +}); + +describe('ToolContext', () => { + const makeFn = (name: string) => + tool({ + name, + description: `${name} tool`, + execute: async () => name, + }); + + it('empty() returns an empty context', () => { + const ctx = ToolContext.empty(); + expect(ctx.functionTools).toEqual({}); + expect(ctx.providerTools).toEqual([]); + expect(ctx.toolsets).toEqual([]); + expect(ctx.flatten()).toEqual([]); + }); + + it('indexes function tools by name and supports lookup', () => { + const a = makeFn('a'); + const b = makeFn('b'); + const ctx = new ToolContext([a, b]); + + expect(ctx.functionTools).toEqual({ a, b }); + expect(ctx.getFunctionTool('a')).toBe(a); + expect(ctx.getFunctionTool('b')).toBe(b); + expect(ctx.getFunctionTool('missing')).toBeUndefined(); + }); + + it('throws on duplicate function names with different instances', () => { + // Matches Python's `if existing is not tool: raise ValueError(...)` — silently overriding + // a registered tool would mask a real bug at the caller (two distinct functions colliding + // on a single advertised name). + const a1 = makeFn('a'); + const a2 = makeFn('a'); + expect(() => new ToolContext([a1, a2])).toThrow('duplicate function name: a'); + }); + + it('silently skips the same function tool instance listed multiple times', () => { + // Matches Python's `return # same instance, skip` branch. Useful when a tool gets + // included both directly and via a future Toolset that re-exports it. + const a = makeFn('a'); + const ctx = new ToolContext([a, a]); + expect(ctx.getFunctionTool('a')).toBe(a); + expect(Object.keys(ctx.functionTools)).toEqual(['a']); + }); + + it('separates provider tools from function tools', () => { + const fnA = makeFn('a'); + const provider = tool({ id: 'code', config: { language: 'python' } }); + const ctx = new ToolContext([fnA, provider]); + + expect(ctx.functionTools).toEqual({ a: fnA }); + expect(ctx.providerTools).toEqual([provider]); + expect(ctx.flatten()).toEqual([fnA, provider]); + }); + + it('updateTools replaces the entire context', () => { + const a = makeFn('a'); + const b = makeFn('b'); + const ctx = new ToolContext([a]); + ctx.updateTools([b]); + expect(ctx.getFunctionTool('a')).toBeUndefined(); + expect(ctx.getFunctionTool('b')).toBe(b); + }); + + it('copy() yields an independent context with the same tools', () => { + const a = makeFn('a'); + const ctx = new ToolContext([a]); + const dup = ctx.copy(); + + expect(dup.getFunctionTool('a')).toBe(a); + dup.updateTools([]); + expect(ctx.getFunctionTool('a')).toBe(a); + expect(dup.getFunctionTool('a')).toBeUndefined(); + }); + + it('equals() compares function tool maps and provider lists by identity', () => { + const a = makeFn('a'); + const b = makeFn('b'); + const c = makeFn('c'); + + expect(new ToolContext([a, b]).equals(new ToolContext([a, b]))).toBe(true); + expect(new ToolContext([a, b]).equals(new ToolContext([a]))).toBe(false); + expect(new ToolContext([a, b]).equals(new ToolContext([a, c]))).toBe(false); + }); + + it('equals() is reflexive', () => { + const a = makeFn('a'); + const provider = tool({ id: 'code', config: { language: 'python' } }); + const ctx = new ToolContext([a, provider]); + expect(ctx.equals(ctx)).toBe(true); + }); + + it('equals() treats provider tool order as insignificant', () => { + // Matches Python's `set(id(t) for t in self._provider_tools)` comparison: two contexts + // that hold the same provider-tool identities in different order are still equal so + // realtime-session / preemptive-generation reuse fast paths are not invalidated. + const a = makeFn('a'); + const p1 = tool({ id: 'code', config: { language: 'python' } }); + const p2 = tool({ id: 'browser', config: {} }); + expect(new ToolContext([a, p1, p2]).equals(new ToolContext([a, p2, p1]))).toBe(true); + }); + + it('equals() supports contexts with only provider tools', () => { + const p1 = tool({ id: 'code', config: {} }); + const p2 = tool({ id: 'browser', config: {} }); + expect(new ToolContext([p1, p2]).equals(new ToolContext([p1, p2]))).toBe(true); + const p3 = tool({ id: 'code', config: {} }); // distinct identity, same id + expect(new ToolContext([p1]).equals(new ToolContext([p3]))).toBe(false); + }); + + it('hasTool() matches function tools by name and provider tools by id', () => { + const a = makeFn('a'); + const provider = tool({ id: 'code_runner', config: {} }); + const ctx = new ToolContext([a, provider]); + + expect(ctx.hasTool('a')).toBe(true); + expect(ctx.hasTool('code_runner')).toBe(true); + expect(ctx.hasTool('missing')).toBe(false); + }); + + it('flatten() returns function tools in insertion order followed by provider tools', () => { + // Matches Python's `flatten()`: list(self._fnc_tools_map.values()) + self._provider_tools. + const a = makeFn('a'); + const b = makeFn('b'); + const provider = tool({ id: 'code', config: {} }); + const ctx = new ToolContext([b, provider, a]); + + expect(ctx.flatten()).toEqual([b, a, provider]); + }); +}); diff --git a/agents/src/llm/tool_context.ts b/agents/src/llm/tool_context.ts index ca2888167..df714d57a 100644 --- a/agents/src/llm/tool_context.ts +++ b/agents/src/llm/tool_context.ts @@ -167,6 +167,12 @@ export interface FunctionTool< > extends Tool { type: 'function'; + /** + * The name of the tool. Used to identify it inside a `ToolContext` and exposed to the LLM + * as the function name to call. + */ + name: string; + /** * The description of the tool. Will be used by the language model to decide whether to use the tool. */ @@ -190,38 +196,168 @@ export interface FunctionTool< [FUNCTION_TOOL_SYMBOL]: true; } -// TODO(AJS-112): support provider-defined tools in the future) -export type ToolContext = { - // eslint-disable-next-line @typescript-eslint/no-explicit-any -- Generic tool registry needs to accept any parameter/result types - [name: string]: FunctionTool; -}; +/** + * Convenience input shape accepted by APIs that want to take a list of tools directly without + * forcing callers to wrap them in `new ToolContext(...)`. + */ +export type ToolCtxInput = + | ToolContext + | readonly ToolContextEntry[]; + +export function toToolContext( + input: ToolCtxInput, +): ToolContext; +export function toToolContext( + input: ToolCtxInput | undefined, +): ToolContext | undefined; +export function toToolContext( + input: ToolCtxInput | undefined, +): ToolContext | undefined { + if (input === undefined) return undefined; + return input instanceof ToolContext ? input : new ToolContext(input); +} -export function isSameToolContext(ctx1: ToolContext, ctx2: ToolContext): boolean { - const toolNames = new Set(Object.keys(ctx1)); - const toolNames2 = new Set(Object.keys(ctx2)); +//TODO: toolset - accept stateful `Toolset` containers alongside `FunctionTool` / +// eslint-disable-next-line @typescript-eslint/no-explicit-any -- ToolContext entries accept any function-tool parameter/result types +export type ToolContextEntry = + // eslint-disable-next-line @typescript-eslint/no-explicit-any + FunctionTool | ProviderDefinedTool; + +export class ToolContext { + // TODO: toolset - widen entries to `FunctionTool | ProviderDefinedTool | Toolset` once Toolset + // lands so this stays heterogeneous like Python's `Sequence[Tool | Toolset]`. + private _tools: ToolContextEntry[] = []; + // eslint-disable-next-line @typescript-eslint/no-explicit-any -- ToolContext stores generic function tools + private _functionToolsMap: Map> = new Map(); + private _providerTools: ProviderDefinedTool[] = []; + // TODO: toolset - populate when Toolset support is supported. + // so the `toolsets` getter and `equals` toolset-identity check stay byte-compatible with the + private _toolSets: unknown[] = []; + + // TODO: toolset - widen `tools` to `Sequence` once Toolset lands. + constructor(tools: readonly ToolContextEntry[] = []) { + this.updateTools(tools); + } - if (toolNames.size !== toolNames2.size) { - return false; + static empty(): ToolContext { + return new ToolContext([]); } - for (const name of toolNames) { - if (!toolNames2.has(name)) { - return false; + /** A copy of all function tools in the tool context, including those in tool sets. */ + // eslint-disable-next-line @typescript-eslint/no-explicit-any + get functionTools(): Record> { + return Object.fromEntries(this._functionToolsMap); + } + + /** A copy of all provider tools in the tool context, including those in tool sets. */ + get providerTools(): ProviderDefinedTool[] { + return this._providerTools; + } + + /** + * A copy of all tool sets in the tool context. + * + * TODO: toolset - wire up once Toolset is ported. + */ + get toolsets(): unknown[] { + return this._toolSets; + } + + /** + * A copy of the raw tool list this context was constructed with. + */ + get tools(): readonly ToolContextEntry[] { + return [...this._tools]; + } + + /** Flatten the tool context to a list of tools. */ + flatten(): Tool[] { + return [...this._functionToolsMap.values(), ...this._providerTools]; + } + + // eslint-disable-next-line @typescript-eslint/no-explicit-any -- Generic registry over any parameter/result types + getFunctionTool(name: string): FunctionTool | undefined { + return this._functionToolsMap.get(name); + } + + hasTool(name: string): boolean { + if (this._functionToolsMap.has(name)) { + return true; } + return this._providerTools.some((tool) => tool.id === name); + } - const tool1 = ctx1[name]; - const tool2 = ctx2[name]; + // TODO: toolset - widen `tools` to `Sequence` once Toolset lands. + updateTools(tools: readonly ToolContextEntry[]): void { + this._tools = [...tools]; + this._functionToolsMap = new Map(); + this._providerTools = []; + this._toolSets = []; + + // Mirrors Python's recursive `add_tool` (minus Toolset flattening, which is TODO). + // eslint-disable-next-line @typescript-eslint/no-explicit-any -- accepts any tool shape + const addTool = (tool: any): void => { + if (isProviderDefinedTool(tool)) { + this._providerTools.push(tool); + return; + } + + if (isFunctionTool(tool)) { + const existing = this._functionToolsMap.get(tool.name); + if (existing !== undefined) { + if (existing !== tool) { + throw new Error(`duplicate function name: ${tool.name}`); + } + return; // same instance, skip + } + this._functionToolsMap.set(tool.name, tool); + return; + } + + // TODO: toolset - if (tool instanceof Toolset) { for (const t of tool.tools) addTool(t); + // this._toolSets.push(tool); return; } + + throw new Error(`unknown tool type: ${typeof tool}`); + }; - if (!tool1 || !tool2) { - return false; + // TODO: toolset - Python also chains `find_function_tools(self)` here so subclasses can + // declare tools as class members. JS doesn't use that decorator pattern, so we only walk + // the explicit input list. + for (const tool of tools) { + addTool(tool); } + } - if (tool1.description !== tool2.description) { + copy(): ToolContext { + return new ToolContext([...this._tools]); + } + + equals(other: ToolContext): boolean { + if (this._functionToolsMap.size !== other._functionToolsMap.size) { + return false; + } + for (const [name, tool] of this._functionToolsMap) { + if (other._functionToolsMap.get(name) !== tool) { + return false; + } + } + if (this._providerTools.length !== other._providerTools.length) { return false; } + // Provider tools compare as identity sets to match Python's `set(id(t) for t in ...)` + // semantics — order is not significant. + const otherProviderIds = new Set(other._providerTools); + for (const tool of this._providerTools) { + if (!otherProviderIds.has(tool)) { + return false; + } + } + // TODO: toolset - once Toolset lands, also compare `_toolSets` as identity sets per Python + // self_tool_set_ids = {id(ts) for ts in self._tool_sets} + // other_tool_set_ids = {id(ts) for ts in other._tool_sets} + // if self_tool_set_ids != other_tool_set_ids: return False + return true; } - - return true; } export function isSameToolChoice(choice1: ToolChoice | null, choice2: ToolChoice | null): boolean { @@ -248,11 +384,13 @@ export function tool< UserData = UnknownUserData, Result = unknown, >({ + name, description, parameters, execute, flags, }: { + name: string; description: string; parameters: Schema; execute: ToolExecuteFunction, UserData, Result>; @@ -263,10 +401,12 @@ export function tool< * Create a function tool without parameters. */ export function tool({ + name, description, execute, flags, }: { + name: string; description: string; parameters?: never; execute: ToolExecuteFunction, UserData, Result>; @@ -290,6 +430,10 @@ export function tool({ // eslint-disable-next-line @typescript-eslint/no-explicit-any export function tool(tool: any): any { if (tool.execute !== undefined) { + if (typeof tool.name !== 'string' || tool.name.length === 0) { + throw new Error('tool({ name, ... }) requires a non-empty name'); + } + // Default parameters to z.object({}) if not provided const parameters = tool.parameters ?? z.object({}); @@ -305,6 +449,7 @@ export function tool(tool: any): any { return { type: 'function', + name: tool.name, description: tool.description, parameters, execute: tool.execute, diff --git a/agents/src/llm/tool_context.type.test.ts b/agents/src/llm/tool_context.type.test.ts index 27cf9fe55..187f95e7b 100644 --- a/agents/src/llm/tool_context.type.test.ts +++ b/agents/src/llm/tool_context.type.test.ts @@ -8,6 +8,7 @@ import { type FunctionTool, type ProviderDefinedTool, type ToolOptions, tool } f describe('tool type inference', () => { it('should infer argument type from zod schema', () => { const toolType = tool({ + name: 'test', description: 'test', parameters: z.object({ number: z.number() }), execute: async () => 'test' as const, @@ -29,6 +30,7 @@ describe('tool type inference', () => { it('should infer run context type', () => { const toolType = tool({ + name: 'test', description: 'test', parameters: z.object({ number: z.number() }), execute: async ({ number }, { ctx }: ToolOptions<{ name: string }>) => { @@ -91,6 +93,7 @@ describe('tool type inference', () => { it('should infer empty object type when parameters are omitted', () => { const toolType = tool({ + name: 'simpleAction', description: 'Simple action without parameters', execute: async () => 'done' as const, }); @@ -100,6 +103,7 @@ describe('tool type inference', () => { it('should infer correct types with context but no parameters', () => { const toolType = tool({ + name: 'actionWithCtx', description: 'Action with context', execute: async (args, { ctx }: ToolOptions<{ userId: number }>) => { expectTypeOf(args).toEqualTypeOf>(); diff --git a/agents/src/llm/utils.ts b/agents/src/llm/utils.ts index 0271deb2a..e8e2c4cab 100644 --- a/agents/src/llm/utils.ts +++ b/agents/src/llm/utils.ts @@ -171,7 +171,7 @@ export const oaiBuildFunctionInfo = ( toolName: string, rawArgs: string, ): FunctionCall => { - const tool = toolCtx[toolName]; + const tool = toolCtx.getFunctionTool(toolName); if (!tool) { throw new Error(`AI tool ${toolName} not found`); } @@ -187,7 +187,7 @@ export async function executeToolCall( toolCall: FunctionCall, toolCtx: ToolContext, ): Promise { - const tool = toolCtx[toolCall.name]!; + const tool = toolCtx.getFunctionTool(toolCall.name)!; let args: object | undefined; let params: object | undefined; diff --git a/agents/src/voice/agent.test.ts b/agents/src/voice/agent.test.ts index 8dd83fee3..f2afa7c42 100644 --- a/agents/src/voice/agent.test.ts +++ b/agents/src/voice/agent.test.ts @@ -29,12 +29,14 @@ describe('Agent', () => { // Create mock tools using the tool function const mockTool1 = tool({ + name: 'getTool1', description: 'First test tool', parameters: z.object({}), execute: async () => 'tool1 result', }); const mockTool2 = tool({ + name: 'getTool2', description: 'Second test tool', parameters: z.object({ input: z.string().describe('Input parameter'), @@ -44,17 +46,14 @@ describe('Agent', () => { const agent = new Agent({ instructions, - tools: { - getTool1: mockTool1, - getTool2: mockTool2, - }, + tools: [mockTool1, mockTool2], }); expect(agent).toBeDefined(); expect(agent.instructions).toBe(instructions); // Assert tools are set correctly - const agentTools = agent.toolCtx; + const agentTools = agent.toolCtx.functionTools; expect(Object.keys(agentTools)).toHaveLength(2); expect(agentTools).toHaveProperty('getTool1'); expect(agentTools).toHaveProperty('getTool2'); @@ -64,27 +63,21 @@ describe('Agent', () => { expect(agentTools.getTool2?.description).toBe('Second test tool'); }); - it('should return a copy of tools, not the original reference', () => { + it('toolCtx returns a defensive copy that exposes the same tools', () => { const instructions = 'You are a helpful assistant'; const mockTool = tool({ + name: 'testTool', description: 'Test tool', parameters: z.object({}), execute: async () => 'result', }); - const tools = { testTool: mockTool }; - const agent = new Agent({ instructions, tools }); - - const tools1 = agent.toolCtx; - const tools2 = agent.toolCtx; - - // Should return different object references - expect(tools1).not.toBe(tools2); - expect(tools1).not.toBe(tools); + const agent = new Agent({ instructions, tools: [mockTool] }); - // Should contain the same set of tools - expect(tools1).toEqual(tools2); - expect(tools1).toEqual(tools); + // Each call returns a fresh ToolContext so external mutation can't escape into the agent's + // internal state. + expect(agent.toolCtx).not.toBe(agent.toolCtx); + expect(agent.toolCtx.getFunctionTool('testTool')).toBe(mockTool); }); it('should require AgentTask to run inside task context', async () => { diff --git a/agents/src/voice/agent.ts b/agents/src/voice/agent.ts index 890d7ea7d..3145de2b1 100644 --- a/agents/src/voice/agent.ts +++ b/agents/src/voice/agent.ts @@ -20,7 +20,8 @@ import { LLM, RealtimeModel, type ToolChoice, - type ToolContext, + ToolContext, + type ToolContextEntry, } from '../llm/index.js'; import { log } from '../log.js'; import type { STT, SpeechEvent } from '../stt/index.js'; @@ -119,7 +120,7 @@ export interface AgentOptions { id?: string; instructions: string | Instructions; chatCtx?: ChatContext; - tools?: ToolContext; + tools?: readonly ToolContextEntry[]; stt?: STT | STTModelString; vad?: VAD; llm?: LLM | RealtimeModel | LLMModels; @@ -157,7 +158,7 @@ export class Agent { _instructions: string | Instructions; /** @internal */ - _tools?: ToolContext; + _toolCtx: ToolContext; constructor({ id, @@ -190,10 +191,10 @@ export class Agent { } this._instructions = instructions; - this._tools = { ...tools }; + this._toolCtx = new ToolContext(tools ?? []); this._chatCtx = chatCtx ? chatCtx.copy({ - toolCtx: this._tools, + toolCtx: this._toolCtx, }) : ChatContext.empty(); @@ -269,7 +270,7 @@ export class Agent { } get toolCtx(): ToolContext { - return { ...this._tools }; + return this._toolCtx.copy(); } get session(): AgentSession { @@ -345,10 +346,10 @@ export class Agent { } // TODO(parity): Add when AgentConfigUpdate is ported to ChatContext. - async updateTools(tools: ToolContext): Promise { + async updateTools(tools: readonly ToolContextEntry[]): Promise { if (!this._agentActivity) { - this._tools = { ...tools }; - this._chatCtx = this._chatCtx.copy({ toolCtx: this._tools }); + this._toolCtx = new ToolContext(tools); + this._chatCtx = this._chatCtx.copy({ toolCtx: this._toolCtx }); return; } diff --git a/agents/src/voice/agent_activity.test.ts b/agents/src/voice/agent_activity.test.ts index 03ddd2dd4..186b4c904 100644 --- a/agents/src/voice/agent_activity.test.ts +++ b/agents/src/voice/agent_activity.test.ts @@ -18,6 +18,7 @@ import { Heap } from 'heap-js'; import { describe, expect, it, vi } from 'vitest'; import type { ChatContext } from '../llm/chat_context.js'; import { LLM, type LLMStream } from '../llm/llm.js'; +import { ToolContext } from '../llm/tool_context.js'; import { Future } from '../utils.js'; import { AgentActivity } from './agent_activity.js'; import type { PreemptiveGenerationInfo } from './audio_recognition.js'; @@ -270,15 +271,16 @@ function buildPreemptiveRunner(opts: Partial = {}) { const fakeChatCtx = { copy: () => fakeChatCtx } as unknown as ChatContext; + const emptyToolCtx = ToolContext.empty(); const fakeActivity = { _preemptiveGenerationCount: 0, _preemptiveGeneration: undefined, _currentSpeech: undefined as SpeechHandle | undefined, schedulingPaused: false, llm: new FakePreemptiveLLM(), - tools: {}, + tools: emptyToolCtx, toolChoice: null, - agent: { chatCtx: fakeChatCtx }, + agent: { chatCtx: fakeChatCtx, _toolCtx: emptyToolCtx }, agentSession: { sessionOptions: { turnHandling: { preemptiveGeneration: preemptiveOpts }, diff --git a/agents/src/voice/agent_activity.ts b/agents/src/voice/agent_activity.ts index e53f39b14..3068b408d 100644 --- a/agents/src/voice/agent_activity.ts +++ b/agents/src/voice/agent_activity.ts @@ -36,11 +36,12 @@ import { type RealtimeModelError, type RealtimeSession, type ToolChoice, - type ToolContext, + ToolContext, + type ToolContextEntry, ToolFlag, } from '../llm/index.js'; import type { LLMError } from '../llm/llm.js'; -import { isSameToolChoice, isSameToolContext } from '../llm/tool_context.js'; +import { isSameToolChoice } from '../llm/tool_context.js'; import { log } from '../log.js'; import type { EOUMetrics, @@ -484,7 +485,12 @@ export class AgentActivity implements RecognitionHooks { } } - const initialTools = Object.keys(this.tools); + // Surface every tool the agent advertises at start — function tools by name and provider + // tools by id. + const initialTools = [ + ...Object.keys(this.agent._toolCtx.functionTools), + ...this.agent._toolCtx.providerTools.map((t) => t.id), + ]; if (runOnEnter && (this.agent.instructions || initialTools.length > 0)) { const initialConfig = new AgentConfigUpdate({ instructions: this.agent.instructions, @@ -609,7 +615,8 @@ export class AgentActivity implements RecognitionHooks { // tools update is supported or tools are the same reusable = reusable && - (capabilities.midSessionToolsUpdate || isSameToolContext(this.tools, newActivity.tools)); + (capabilities.midSessionToolsUpdate || + this.agent._toolCtx.equals(newActivity.agent._toolCtx)); if (reusable) { // detach: remove event listeners but don't close the session @@ -759,13 +766,14 @@ export class AgentActivity implements RecognitionHooks { } } - async updateTools(tools: ToolContext): Promise { - const oldToolNames = new Set(Object.keys(this.tools)); - const newToolNames = new Set(Object.keys(tools)); + async updateTools(tools: readonly ToolContextEntry[]): Promise { + const oldToolNames = new Set(Object.keys(this.agent._toolCtx.functionTools)); + const newToolCtx = new ToolContext(tools); + const newToolNames = new Set(Object.keys(newToolCtx.functionTools)); const toolsAdded = [...newToolNames].filter((name) => !oldToolNames.has(name)); const toolsRemoved = [...oldToolNames].filter((name) => !newToolNames.has(name)); - this.agent._tools = { ...tools }; + this.agent._toolCtx = newToolCtx; if (toolsAdded.length > 0 || toolsRemoved.length > 0) { const configUpdate = new AgentConfigUpdate({ @@ -777,12 +785,12 @@ export class AgentActivity implements RecognitionHooks { } if (this.realtimeSession) { - await this.realtimeSession.updateTools(tools); + await this.realtimeSession.updateTools(newToolCtx); } if (this.llm instanceof LLM) { // for realtime LLM, we assume the server will remove unvalid tool messages - await this.updateChatCtx(this.agent._chatCtx.copy({ toolCtx: tools })); + await this.updateChatCtx(this.agent._chatCtx.copy({ toolCtx: newToolCtx })); } } @@ -1421,7 +1429,7 @@ export class AgentActivity implements RecognitionHooks { userMessage, info, chatCtx: chatCtx.copy(), - tools: { ...this.tools }, + tools: this.agent._toolCtx.copy(), toolChoice: this.toolChoice, createdAt: Date.now(), }; @@ -1725,11 +1733,14 @@ export class AgentActivity implements RecognitionHooks { const shouldFilterTools = onEnterData?.agent === this.agent && onEnterData?.session === this.agentSession; - const tools = shouldFilterTools - ? Object.fromEntries( - Object.entries(this.agent.toolCtx).filter( - ([, fnTool]) => !(fnTool.flags & ToolFlag.IGNORE_ON_ENTER), - ), + const tools: ToolContext = shouldFilterTools + ? new ToolContext( + this.agent.toolCtx.tools.filter((t) => { + if (t.type === 'function') { + return !(t.flags & ToolFlag.IGNORE_ON_ENTER); + } + return true; + }), ) : this.agent.toolCtx; @@ -1912,7 +1923,7 @@ export class AgentActivity implements RecognitionHooks { if ( preemptive.info.newTranscript === userMessage?.textContent && preemptive.chatCtx.isEquivalent(chatCtx) && - isSameToolContext(preemptive.tools, this.tools) && + preemptive.tools.equals(this.agent._toolCtx) && isSameToolChoice(preemptive.toolChoice, this.toolChoice) ) { speechHandle = preemptive.speechHandle; diff --git a/agents/src/voice/amd.test.ts b/agents/src/voice/amd.test.ts index 8f8acd822..197bee11a 100644 --- a/agents/src/voice/amd.test.ts +++ b/agents/src/voice/amd.test.ts @@ -7,7 +7,7 @@ import type { ChatContext } from '../llm/chat_context.js'; import { FunctionCall } from '../llm/chat_context.js'; import type { ChatChunk } from '../llm/llm.js'; import { LLM, type LLMStream } from '../llm/llm.js'; -import type { ToolChoice, ToolContext } from '../llm/tool_context.js'; +import type { ToolChoice, ToolCtxInput } from '../llm/tool_context.js'; import type { SpeechEvent, SpeechStream } from '../stt/stt.js'; import { STT } from '../stt/stt.js'; import type { APIConnectOptions } from '../types.js'; @@ -30,7 +30,7 @@ class StaticLLM extends LLM { connOptions: _connOptions, }: { chatCtx: ChatContext; - toolCtx?: ToolContext; + toolCtx?: ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: ToolChoice; @@ -194,7 +194,7 @@ describe('AMD', () => { } chat({}: { chatCtx: ChatContext; - toolCtx?: ToolContext; + toolCtx?: ToolCtxInput; connOptions?: APIConnectOptions; }): LLMStream { return { @@ -538,7 +538,7 @@ describe('AMD', () => { label(): string { return 'postpone-llm'; } - chat({}: { chatCtx: ChatContext; toolCtx?: ToolContext }): LLMStream { + chat({}: { chatCtx: ChatContext; toolCtx?: ToolCtxInput }): LLMStream { callCount += 1; const isFirst = callCount === 1; return { diff --git a/agents/src/voice/amd.ts b/agents/src/voice/amd.ts index c1af28457..eeb1fd053 100644 --- a/agents/src/voice/amd.ts +++ b/agents/src/voice/amd.ts @@ -13,8 +13,7 @@ import type { LLMModels, STTModels } from '../inference/index.js'; import { ChatContext } from '../llm/chat_context.js'; import type { FunctionCall } from '../llm/chat_context.js'; import { LLM, type LLMStream } from '../llm/llm.js'; -import { isFunctionTool, tool } from '../llm/tool_context.js'; -import type { ToolContext } from '../llm/tool_context.js'; +import { ToolContext, type ToolContextEntry, isFunctionTool, tool } from '../llm/tool_context.js'; import { log } from '../log.js'; import { STT, SpeechEventType, type SpeechStream } from '../stt/stt.js'; import { traceTypes, tracer } from '../telemetry/index.js'; @@ -934,6 +933,7 @@ export class AMD extends (EventEmitter as new () => TypedEmitter) const isStale = (): boolean => generation !== this.detectGeneration || this.settled; const savePrediction = tool({ + name: 'save_prediction', description: 'Save the AMD prediction to the verdict.', parameters: z.object({ label: z.enum([ @@ -966,6 +966,7 @@ export class AMD extends (EventEmitter as new () => TypedEmitter) }); const postponeTermination = tool({ + name: 'postpone_termination', description: 'Postpone the termination of the classification task. ' + 'Use when the transcript is ambiguous and more audio is expected.', @@ -996,10 +997,11 @@ export class AMD extends (EventEmitter as new () => TypedEmitter) }, }); - const toolCtx: ToolContext = { save_prediction: savePrediction }; + const toolList: ToolContextEntry[] = [savePrediction]; if (this.extensionCount < MAX_EXTENSIONS) { - toolCtx.postpone_termination = postponeTermination; + toolList.push(postponeTermination); } + const toolCtx = new ToolContext(toolList); const chatCtx = new ChatContext(); chatCtx.addMessage({ role: 'system', content: this.prompt }); @@ -1035,7 +1037,7 @@ export class AMD extends (EventEmitter as new () => TypedEmitter) // Execute tool calls (save_prediction populates `savedResult`, // postpone_termination mutates the silence timer and returns). for (const tc of toolCalls) { - const fnTool = toolCtx[tc.name]; + const fnTool = toolCtx.getFunctionTool(tc.name); if (!fnTool || !isFunctionTool(fnTool)) continue; let parsedArgs: unknown = {}; try { diff --git a/agents/src/voice/generation.ts b/agents/src/voice/generation.ts index 3d938abdc..8bce7c198 100644 --- a/agents/src/voice/generation.ts +++ b/agents/src/voice/generation.ts @@ -485,7 +485,10 @@ export function performLLMInference( traceTypes.ATTR_CHAT_CTX, JSON.stringify(chatCtx.toJSON({ excludeTimestamp: false })), ); - span.setAttribute(traceTypes.ATTR_FUNCTION_TOOLS, JSON.stringify(Object.keys(toolCtx))); + span.setAttribute( + traceTypes.ATTR_FUNCTION_TOOLS, + JSON.stringify(Object.keys(toolCtx.functionTools)), + ); if (model) { span.setAttribute(traceTypes.ATTR_GEN_AI_REQUEST_MODEL, model); @@ -992,7 +995,7 @@ export function performToolExecutions({ // TODO(brian): assert other toolChoice values - const tool = toolCtx[toolCall.name]; + const tool = toolCtx.getFunctionTool(toolCall.name); if (!tool) { logger.warn( { diff --git a/agents/src/voice/generation_tools.test.ts b/agents/src/voice/generation_tools.test.ts index d53e12196..9b4f7d1df 100644 --- a/agents/src/voice/generation_tools.test.ts +++ b/agents/src/voice/generation_tools.test.ts @@ -4,7 +4,7 @@ import { ReadableStream as NodeReadableStream } from 'stream/web'; import { describe, expect, it } from 'vitest'; import { z } from 'zod'; -import { FunctionCall, tool } from '../llm/index.js'; +import { FunctionCall, ToolContext, tool } from '../llm/index.js'; import { initializeLogger } from '../log.js'; import type { Task } from '../utils.js'; import { cancelAndWait, delay } from '../utils.js'; @@ -63,6 +63,7 @@ describe('Generation + Tool Execution', () => { // Tool that takes > 5 seconds let toolAborted = false; const getWeather = tool({ + name: 'getWeather', description: 'weather', parameters: z.object({ location: z.string() }), execute: async ({ location }, { abortSignal }) => { @@ -87,7 +88,7 @@ describe('Generation + Tool Execution', () => { const [execTask, toolOutput] = performToolExecutions({ session: {} as any, speechHandle: { id: 'speech_test', _itemAdded: () => {} } as any, - toolCtx: { getWeather } as any, + toolCtx: new ToolContext([getWeather]) as any, toolCallStream, controller: replyAbortController, onToolExecutionStarted: () => {}, @@ -115,6 +116,7 @@ describe('Generation + Tool Execution', () => { const replyAbortController = new AbortController(); const echo = tool({ + name: 'echo', description: 'echo', parameters: z.object({ msg: z.string() }), execute: async ({ msg }) => `echo: ${msg}`, @@ -130,7 +132,7 @@ describe('Generation + Tool Execution', () => { const [execTask, toolOutput] = performToolExecutions({ session: {} as any, speechHandle: { id: 'speech_test2', _itemAdded: () => {} } as any, - toolCtx: { echo } as any, + toolCtx: new ToolContext([echo]) as any, toolCallStream, controller: replyAbortController, }); @@ -147,6 +149,7 @@ describe('Generation + Tool Execution', () => { let aborted = false; const longOp = tool({ + name: 'longOp', description: 'longOp', parameters: z.object({ ms: z.number() }), execute: async ({ ms }, { abortSignal }) => { @@ -170,7 +173,7 @@ describe('Generation + Tool Execution', () => { const [execTask, toolOutput] = performToolExecutions({ session: {} as any, speechHandle: { id: 'speech_abort', _itemAdded: () => {} } as any, - toolCtx: { longOp } as any, + toolCtx: new ToolContext([longOp]) as any, toolCallStream, controller: replyAbortController, }); @@ -189,6 +192,7 @@ describe('Generation + Tool Execution', () => { const replyAbortController = new AbortController(); const echo = tool({ + name: 'echo', description: 'echo', parameters: z.object({ msg: z.string() }), execute: async ({ msg }) => `echo: ${msg}`, @@ -205,7 +209,7 @@ describe('Generation + Tool Execution', () => { const [execTask, toolOutput] = performToolExecutions({ session: {} as any, speechHandle: { id: 'speech_invalid', _itemAdded: () => {} } as any, - toolCtx: { echo } as any, + toolCtx: new ToolContext([echo]) as any, toolCallStream, controller: replyAbortController, }); @@ -220,11 +224,13 @@ describe('Generation + Tool Execution', () => { const replyAbortController = new AbortController(); const sum = tool({ + name: 'sum', description: 'sum', parameters: z.object({ a: z.number(), b: z.number() }), execute: async ({ a, b }) => a + b, }); const upper = tool({ + name: 'upper', description: 'upper', parameters: z.object({ s: z.string() }), execute: async ({ s }) => s.toUpperCase(), @@ -245,7 +251,7 @@ describe('Generation + Tool Execution', () => { const [execTask, toolOutput] = performToolExecutions({ session: {} as any, speechHandle: { id: 'speech_multi', _itemAdded: () => {} } as any, - toolCtx: { sum, upper } as any, + toolCtx: new ToolContext([sum, upper]) as any, toolCallStream, controller: replyAbortController, }); diff --git a/agents/src/voice/remote_session.ts b/agents/src/voice/remote_session.ts index f970b064e..ee1a3dab8 100644 --- a/agents/src/voice/remote_session.ts +++ b/agents/src/voice/remote_session.ts @@ -471,7 +471,7 @@ function sessionUsageToProto(usage: AgentSessionUsage): pb.AgentSessionUsage { function toolNames(toolCtx: ToolContext | undefined): string[] { if (!toolCtx) return []; - return Object.keys(toolCtx); + return Object.keys(toolCtx.functionTools); } function protoSerializeOptions(opts: { diff --git a/agents/src/voice/testing/fake_llm.ts b/agents/src/voice/testing/fake_llm.ts index ad3a1bf16..b6ba9b08a 100644 --- a/agents/src/voice/testing/fake_llm.ts +++ b/agents/src/voice/testing/fake_llm.ts @@ -4,7 +4,7 @@ import type { ChatContext } from '../../llm/chat_context.js'; import { FunctionCall } from '../../llm/chat_context.js'; import { LLMStream as BaseLLMStream, LLM, type LLMStream } from '../../llm/llm.js'; -import type { ToolChoice, ToolContext } from '../../llm/tool_context.js'; +import type { ToolChoice, ToolCtxInput } from '../../llm/tool_context.js'; import { type APIConnectOptions, DEFAULT_API_CONNECT_OPTIONS } from '../../types.js'; import { delay } from '../../utils.js'; @@ -42,7 +42,7 @@ export class FakeLLM extends LLM { connOptions = DEFAULT_API_CONNECT_OPTIONS, }: { chatCtx: ChatContext; - toolCtx?: ToolContext; + toolCtx?: ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: ToolChoice; @@ -65,7 +65,7 @@ class FakeLLMStream extends BaseLLMStream { constructor( fake: FakeLLM, - params: { chatCtx: ChatContext; toolCtx?: ToolContext; connOptions: APIConnectOptions }, + params: { chatCtx: ChatContext; toolCtx?: ToolCtxInput; connOptions: APIConnectOptions }, ) { super(fake, params); this.fake = fake; diff --git a/agents/src/voice/testing/run_result.ts b/agents/src/voice/testing/run_result.ts index 4ee0ccc56..0e60d03ce 100644 --- a/agents/src/voice/testing/run_result.ts +++ b/agents/src/voice/testing/run_result.ts @@ -817,6 +817,7 @@ export class MessageAssert extends EventAssert { // Create the check_intent tool const checkIntentTool = tool({ + name: 'check_intent', description: 'Determines whether the message correctly fulfills the given intent. ' + 'Returns success=true if the message satisfies the intent, false otherwise. ' + @@ -853,7 +854,7 @@ export class MessageAssert extends EventAssert { const stream = llm.chat({ chatCtx, - toolCtx: { check_intent: checkIntentTool }, + toolCtx: [checkIntentTool], toolChoice: { type: 'function', function: { name: 'check_intent' } }, extraKwargs: { temperature: 0 }, }); diff --git a/examples/src/background_audio.ts b/examples/src/background_audio.ts index fc718ac0c..8d5884be9 100644 --- a/examples/src/background_audio.ts +++ b/examples/src/background_audio.ts @@ -35,6 +35,7 @@ export default defineAgent({ logger.info('Connected to room'); const searchWeb = llm.tool({ + name: 'searchWeb', description: 'Search the web for information based on the given query. Always use this function whenever the user requests a web search', parameters: z.object({ @@ -49,9 +50,7 @@ export default defineAgent({ const agent = new voice.Agent({ instructions: 'You are a helpful assistant', - tools: { - searchWeb, - }, + tools: [searchWeb], }); const session = new voice.AgentSession({ diff --git a/examples/src/basic_agent.ts b/examples/src/basic_agent.ts index 95ecddb9a..79e79808a 100644 --- a/examples/src/basic_agent.ts +++ b/examples/src/basic_agent.ts @@ -27,8 +27,9 @@ export default defineAgent({ const agent = new voice.Agent({ instructions: "You are a helpful assistant, you can hear the user's message and respond to it.", - tools: { - getWeather: llm.tool({ + tools: [ + llm.tool({ + name: 'getWeather', description: 'Get the weather for a given location.', parameters: z.object({ location: z.string().describe('The location to get the weather for'), @@ -37,7 +38,7 @@ export default defineAgent({ return `The weather in ${location} is sunny.`; }, }), - }, + ], }); const logger = log(); diff --git a/examples/src/basic_agent_task.ts b/examples/src/basic_agent_task.ts index aacbeee5c..c450a81bb 100644 --- a/examples/src/basic_agent_task.ts +++ b/examples/src/basic_agent_task.ts @@ -21,8 +21,9 @@ class InfoTask extends voice.AgentTask { super({ instructions: `Collect the user's information. around ${info}. Once you have the information, call the saveUserInfo tool to save the information to the database IMMEDIATELY. DO NOT have chitchat with the user, just collect the information and call the saveUserInfo tool.`, tts: 'elevenlabs/eleven_turbo_v2_5', - tools: { - saveUserInfo: llm.tool({ + tools: [ + llm.tool({ + name: 'saveUserInfo', description: `Save the user's ${info} to database`, parameters: z.object({ [info]: z.string(), @@ -32,7 +33,7 @@ class InfoTask extends voice.AgentTask { return `Thanks, collected ${info} successfully: ${args[info]}`; }, }), - }, + ], }); } @@ -48,8 +49,9 @@ class SurveyAgent extends voice.Agent { super({ instructions: 'You orchestrate a short intro survey. Speak naturally and keep the interaction brief.', - tools: { - collectUserInfo: llm.tool({ + tools: [ + llm.tool({ + name: 'collectUserInfo', description: 'Call this when user want to provide some information to you', parameters: z.object({ key: z @@ -63,15 +65,17 @@ class SurveyAgent extends voice.Agent { return `Collected ${key} successfully: ${value}`; }, }), - transferToWeatherAgent: llm.tool({ + llm.tool({ + name: 'transferToWeatherAgent', description: 'Call this immediately after user want to know the weather', execute: async () => { const agent = new voice.Agent({ instructions: 'You are a weather agent. You are responsible for providing the weather information to the user.', tts: 'deepgram/aura-2', - tools: { - getWeather: llm.tool({ + tools: [ + llm.tool({ + name: 'getWeather', description: 'Get the weather for a given location', parameters: z.object({ location: z.string().describe('The location to get the weather for'), @@ -80,7 +84,8 @@ class SurveyAgent extends voice.Agent { return `The weather in ${location} is sunny today.`; }, }), - finishWeatherConversation: llm.tool({ + llm.tool({ + name: 'finishWeatherConversation', description: 'Call this when you want to finish the weather conversation', execute: async () => { return llm.handoff({ @@ -89,13 +94,13 @@ class SurveyAgent extends voice.Agent { }); }, }), - }, + ], }); return llm.handoff({ agent, returns: "Let's start the weather conversation!" }); }, }), - }, + ], }); } diff --git a/examples/src/basic_task_group.ts b/examples/src/basic_task_group.ts index d40befe2a..3d8e04464 100644 --- a/examples/src/basic_task_group.ts +++ b/examples/src/basic_task_group.ts @@ -25,8 +25,9 @@ class CollectNameTask extends voice.AgentTask { instructions: 'Collect the user name from the latest user message. As soon as you have it, call save_name.', tts: taskTts, - tools: { - save_name: llm.tool({ + tools: [ + llm.tool({ + name: 'save_name', description: 'Save the user name.', parameters: z.object({ name: z.string().describe('The user name'), @@ -36,7 +37,7 @@ class CollectNameTask extends voice.AgentTask { return `Saved name: ${name}`; }, }), - }, + ], }); } @@ -54,8 +55,9 @@ class CollectEmailTask extends voice.AgentTask { instructions: 'Collect the user email from the latest user message. As soon as you have it, call save_email.', tts: taskTts, - tools: { - save_email: llm.tool({ + tools: [ + llm.tool({ + name: 'save_email', description: 'Save the user email.', parameters: z.object({ email: z.string().describe('The user email'), @@ -65,7 +67,7 @@ class CollectEmailTask extends voice.AgentTask { return `Saved email: ${email}`; }, }), - }, + ], }); } @@ -82,8 +84,9 @@ class TaskGroupDemoAgent extends voice.Agent { super({ instructions: 'You are onboarding assistant. When user asks to begin onboarding, call startOnboarding exactly once.', - tools: { - startOnboarding: llm.tool({ + tools: [ + llm.tool({ + name: 'startOnboarding', description: 'Start a two-step onboarding flow (name then email).', parameters: z.object({}), execute: async () => { @@ -107,7 +110,7 @@ class TaskGroupDemoAgent extends voice.Agent { return JSON.stringify(result.taskResults); }, }), - }, + ], }); } diff --git a/examples/src/basic_tool_call_agent.ts b/examples/src/basic_tool_call_agent.ts index 5642ef488..9a18de362 100644 --- a/examples/src/basic_tool_call_agent.ts +++ b/examples/src/basic_tool_call_agent.ts @@ -44,6 +44,7 @@ export default defineAgent({ }, entry: async (ctx: JobContext) => { const getWeather = llm.tool({ + name: 'getWeather', description: ' Called when the user asks about the weather.', parameters: z.object({ location: z.string().describe('The location to get the weather for'), @@ -55,6 +56,7 @@ export default defineAgent({ }); const toggleLight = llm.tool({ + name: 'toggleLight', description: 'Called when the user asks to turn on or off the light.', parameters: z.object({ room: roomNameSchema.describe('The room to turn the light in'), @@ -70,6 +72,7 @@ export default defineAgent({ }); const getNumber = llm.tool({ + name: 'getNumber', description: 'Called when the user wants to get a number value, None if user want a random value', parameters: z.object({ @@ -87,6 +90,7 @@ export default defineAgent({ }); const checkStoredNumber = llm.tool({ + name: 'checkStoredNumber', description: 'Called when the user wants to check the stored number.', execute: async (_, { ctx }: llm.ToolOptions) => { return `The stored number is ${ctx.userData.number}.`; @@ -94,6 +98,7 @@ export default defineAgent({ }); const updateStoredNumber = llm.tool({ + name: 'updateStoredNumber', description: 'Called when the user wants to update the stored number.', parameters: z.object({ number: z.number().describe('The number to update the stored number to'), @@ -106,31 +111,33 @@ export default defineAgent({ const routerAgent = new RouterAgent({ instructions: 'You are a helpful assistant.', - tools: { + tools: [ getWeather, toggleLight, - playGame: llm.tool({ + llm.tool({ + name: 'playGame', description: 'Called when the user wants to play a game (transfer user to a game agent).', execute: async (): Promise => { return llm.handoff({ agent: gameAgent, returns: 'The game is now playing.' }); }, }), - }, + ], }); const gameAgent = new GameAgent({ instructions: 'You are a game agent. You are playing a game with the user.', - tools: { + tools: [ getNumber, checkStoredNumber, updateStoredNumber, - finishGame: llm.tool({ + llm.tool({ + name: 'finishGame', description: 'Called when the user wants to finish the game.', execute: async () => { return llm.handoff({ agent: routerAgent, returns: 'The game is now finished.' }); }, }), - }, + ], }); const vad = ctx.proc.userData.vad! as silero.VAD; diff --git a/examples/src/comprehensive_test.ts b/examples/src/comprehensive_test.ts index ddebfc4a7..f3b4974f0 100644 --- a/examples/src/comprehensive_test.ts +++ b/examples/src/comprehensive_test.ts @@ -76,8 +76,9 @@ class MainAgent extends voice.Agent { tts: ttsOptions['elevenlabs'](), llm: llmOptions['openai'](), turnDetection: eouOptions['multilingual'](), - tools: { - testAgent: llm.tool({ + tools: [ + llm.tool({ + name: 'testAgent', description: 'Called when user want to test an agent with STT, TTS, EOU, LLM, and optionally realtime LLM configuration', parameters: z.object({ @@ -102,7 +103,7 @@ class MainAgent extends voice.Agent { }); }, }), - }, + ], }); } @@ -159,8 +160,9 @@ class TestAgent extends voice.Agent { tts: tts, llm: realtimeModel ?? model, turnDetection: eou, - tools: { - testTool: llm.tool({ + tools: [ + llm.tool({ + name: 'testTool', description: "Testing agent's tool calling ability", parameters: z .object({ @@ -173,7 +175,8 @@ class TestAgent extends voice.Agent { }; }, }), - nextAgent: llm.tool({ + llm.tool({ + name: 'nextAgent', description: 'Called when user confirm current agent is working and want to proceed to next agent', parameters: z.object({ @@ -204,7 +207,7 @@ class TestAgent extends voice.Agent { }); }, }), - }, + ], }); this.sttChoice = sttChoice; diff --git a/examples/src/drive-thru/drivethru_agent.ts b/examples/src/drive-thru/drivethru_agent.ts index 9882f6fcd..e684f5c00 100644 --- a/examples/src/drive-thru/drivethru_agent.ts +++ b/examples/src/drive-thru/drivethru_agent.ts @@ -57,23 +57,24 @@ export class DriveThruAgent extends voice.Agent { super({ instructions, - tools: { - orderComboMeal: DriveThruAgent.buildComboOrderTool( + tools: [ + DriveThruAgent.buildComboOrderTool( userdata.comboItems, userdata.drinkItems, userdata.sauceItems, ), - orderHappyMeal: DriveThruAgent.buildHappyOrderTool( + DriveThruAgent.buildHappyOrderTool( userdata.happyItems, userdata.drinkItems, userdata.sauceItems, ), - orderRegularItem: DriveThruAgent.buildRegularOrderTool( + DriveThruAgent.buildRegularOrderTool( userdata.regularItems, userdata.drinkItems, userdata.sauceItems, ), - removeOrderItem: llm.tool({ + llm.tool({ + name: 'removeOrderItem', description: `Removes one or more items from the user's order using their \`orderId\`s. Useful when the user asks to cancel or delete existing items (e.g., "Remove the cheeseburger"). @@ -100,7 +101,8 @@ If the \`orderId\`s are unknown, call \`listOrderItems\` first to retrieve them. return 'Removed items:\n' + removedItems.map((item) => JSON.stringify(item)).join('\n'); }, }), - listOrderItems: llm.tool({ + llm.tool({ + name: 'listOrderItems', description: `Retrieves the current list of items in the user's order, including each item's internal \`orderId\`. Helpful when: @@ -120,7 +122,7 @@ Examples: return items.map((item) => JSON.stringify(item)).join('\n'); }, }), - }, + ], }); } @@ -134,6 +136,7 @@ Examples: const availableSauceIds = [...new Set(sauceItems.map((item) => item.id))]; return llm.tool({ + name: 'orderComboMeal', description: `Call this when the user orders a **Combo Meal**, like: "Number 4b with a large Sprite" or "I'll do a medium meal." Do not call this tool unless the user clearly refers to a known combo meal by name or number. @@ -222,6 +225,7 @@ If the user says just "a large meal," assume both drink and fries are that size. const availableSauceIds = [...new Set(sauceItems.map((item) => item.id))]; return llm.tool({ + name: 'orderHappyMeal', description: `Call this when the user orders a **Happy Meal**, typically for children. These meals come with a main item, a drink, and a sauce. The user must clearly specify a valid Happy Meal option (e.g., "Can I get a Happy Meal?"). @@ -299,6 +303,7 @@ Assume Small as default only if the user says "Happy Meal" and gives no size pre const availableIds = [...new Set(allItems.map((item) => item.id))]; return llm.tool({ + name: 'orderRegularItem', description: `Call this when the user orders **a single item on its own**, not as part of a Combo Meal or Happy Meal. The customer must provide clear and specific input. For example, item variants such as flavor must **always** be explicitly stated. diff --git a/examples/src/frontdesk/frontdesk_agent.ts b/examples/src/frontdesk/frontdesk_agent.ts index d5d2e1ab1..2fa60ba57 100644 --- a/examples/src/frontdesk/frontdesk_agent.ts +++ b/examples/src/frontdesk/frontdesk_agent.ts @@ -59,8 +59,9 @@ export class FrontDeskAgent extends voice.Agent { super({ instructions, - tools: { - scheduleAppointment: llm.tool({ + tools: [ + llm.tool({ + name: 'scheduleAppointment', description: 'Schedule an appointment at the given slot.', parameters: z.object({ slotId: z @@ -110,7 +111,8 @@ export class FrontDeskAgent extends voice.Agent { return `The appointment was successfully scheduled for ${formatted}.`; }, }), - listAvailableSlots: llm.tool({ + llm.tool({ + name: 'listAvailableSlots', description: `Return a plain-text list of available slots, one per line. - , , at () @@ -188,7 +190,7 @@ You must infer the appropriate range implicitly from the conversational context return lines.join('\n') || 'No slots available at the moment.'; }, }), - }, + ], }); this.tz = options.timezone; diff --git a/examples/src/gemini_realtime_agent.ts b/examples/src/gemini_realtime_agent.ts index dd22db51f..d70368352 100644 --- a/examples/src/gemini_realtime_agent.ts +++ b/examples/src/gemini_realtime_agent.ts @@ -47,6 +47,7 @@ type StoryData = { const roomNameSchema = z.enum(['bedroom', 'living room', 'kitchen', 'bathroom', 'office']); const getWeather = llm.tool({ + name: 'getWeather', description: 'Called when the user asks about the weather.', parameters: z.object({ location: z.string().describe('The location to get the weather for'), @@ -59,6 +60,7 @@ const getWeather = llm.tool({ }); const toggleLight = llm.tool({ + name: 'toggleLight', description: 'Called when the user asks to turn on or off the light.', parameters: z.object({ room: roomNameSchema.describe('The room to turn the light in'), @@ -79,15 +81,16 @@ class IntroAgent extends voice.Agent { static create() { return new IntroAgent({ instructions: `You are a story teller. Your goal is to gather a few pieces of information from the user to make the story personalized and engaging. Ask the user for their name and where they are from.`, - tools: { - informationGathered: llm.tool({ + tools: [ + llm.tool({ + name: 'informationGathered', description: 'Called when the user has provided the information needed to make the story personalized and engaging.', parameters: z.object({ name: z.string().describe('The name of the user'), location: z.string().describe('The location of the user'), }), - execute: async ({ name, location }, { ctx }) => { + execute: async ({ name, location }, { ctx }: llm.ToolOptions) => { ctx.userData.name = name; ctx.userData.location = location; @@ -97,7 +100,7 @@ class IntroAgent extends voice.Agent { }), getWeather, toggleLight, - }, + ], }); } } diff --git a/examples/src/instructions_per_modality.ts b/examples/src/instructions_per_modality.ts index 71f2f3f04..643e3d433 100644 --- a/examples/src/instructions_per_modality.ts +++ b/examples/src/instructions_per_modality.ts @@ -57,8 +57,9 @@ class SchedulingAgent extends voice.Agent { super({ instructions, - tools: { - bookAppointment: llm.tool({ + tools: [ + llm.tool({ + name: 'bookAppointment', description: 'Book an appointment.', parameters: z.object({ date: z.string().describe('The date of the appointment in the format YYYY-MM-DD'), @@ -69,7 +70,7 @@ class SchedulingAgent extends voice.Agent { return `Appointment booked for ${date} at ${time}`; }, }), - }, + ], }); } diff --git a/examples/src/llm_fallback_adapter.ts b/examples/src/llm_fallback_adapter.ts index d053464dc..d86708548 100644 --- a/examples/src/llm_fallback_adapter.ts +++ b/examples/src/llm_fallback_adapter.ts @@ -71,8 +71,9 @@ export default defineAgent({ const agent = new voice.Agent({ instructions: 'You are a helpful assistant. Demonstrate that you are working by responding to user queries.', - tools: { - getWeather: llm.tool({ + tools: [ + llm.tool({ + name: 'getWeather', description: 'Get the weather for a given location.', parameters: z.object({ location: z.string().describe('The location to get the weather for'), @@ -81,7 +82,7 @@ export default defineAgent({ return `The weather in ${location} is sunny with a temperature of 72°F.`; }, }), - }, + ], }); const session = new voice.AgentSession({ diff --git a/examples/src/manual_shutdown.ts b/examples/src/manual_shutdown.ts index 96bedb901..edb21b87c 100644 --- a/examples/src/manual_shutdown.ts +++ b/examples/src/manual_shutdown.ts @@ -25,8 +25,9 @@ export default defineAgent({ const agent = new voice.Agent({ instructions: "You are a helpful assistant, you can hear the user's message and respond to it, end the call when the user asks you to.", - tools: { - getWeather: llm.tool({ + tools: [ + llm.tool({ + name: 'getWeather', description: 'Get the weather for a given location.', parameters: z.object({ location: z.string().describe('The location to get the weather for'), @@ -35,7 +36,8 @@ export default defineAgent({ return `The weather in ${location} is sunny.`; }, }), - endCall: llm.tool({ + llm.tool({ + name: 'endCall', description: 'End the call.', parameters: z.object({ reason: z @@ -56,7 +58,7 @@ export default defineAgent({ session.shutdown({ reason }); }, }), - }, + ], }); const session = new voice.AgentSession({ diff --git a/examples/src/multi_agent.ts b/examples/src/multi_agent.ts index 7f4819bed..12d87377d 100644 --- a/examples/src/multi_agent.ts +++ b/examples/src/multi_agent.ts @@ -35,15 +35,16 @@ class IntroAgent extends voice.Agent { static create() { return new IntroAgent({ instructions: `You are a story teller. Your goal is to gather a few pieces of information from the user to make the story personalized and engaging. Ask the user for their name and where they are from.`, - tools: { - informationGathered: llm.tool({ + tools: [ + llm.tool({ + name: 'informationGathered', description: 'Called when the user has provided the information needed to make the story personalized and engaging.', parameters: z.object({ name: z.string().describe('The name of the user'), location: z.string().describe('The location of the user'), }), - execute: async ({ name, location }, { ctx }) => { + execute: async ({ name, location }, { ctx }: llm.ToolOptions) => { ctx.userData.name = name; ctx.userData.location = location; @@ -51,7 +52,7 @@ class IntroAgent extends voice.Agent { return llm.handoff({ agent: storyAgent, returns: "Let's start the story!" }); }, }), - }, + ], }); } } diff --git a/examples/src/phonic_realtime_agent.ts b/examples/src/phonic_realtime_agent.ts index 35a038450..934485f5b 100644 --- a/examples/src/phonic_realtime_agent.ts +++ b/examples/src/phonic_realtime_agent.ts @@ -7,6 +7,7 @@ import { fileURLToPath } from 'node:url'; import { z } from 'zod'; const toggleLight = llm.tool({ + name: 'toggle_light', description: 'Toggle a light on or off. Available lights are A05, A06, A07, and A08.', parameters: z.object({ light_id: z.string().describe('The ID of the light to toggle'), @@ -23,9 +24,7 @@ export default defineAgent({ entry: async (ctx: JobContext) => { const agent = new voice.Agent({ instructions: 'You are a helpful voice AI assistant named Alex.', - tools: { - toggle_light: toggleLight, - }, + tools: [toggleLight], }); const session = new voice.AgentSession({ diff --git a/examples/src/raw_function_description.ts b/examples/src/raw_function_description.ts index 6548fd011..1e81c65d9 100644 --- a/examples/src/raw_function_description.ts +++ b/examples/src/raw_function_description.ts @@ -18,8 +18,9 @@ import { fileURLToPath } from 'node:url'; function createRawFunctionAgent() { return new voice.Agent({ instructions: 'You are a helpful assistant.', - tools: { - openGate: llm.tool({ + tools: [ + llm.tool({ + name: 'openGate', description: 'Opens a specified gate from a predefined set of access points.', parameters: { type: 'object', @@ -43,7 +44,7 @@ function createRawFunctionAgent() { return `The gate ${gateId} is now open.`; }, }), - }, + ], }); } diff --git a/examples/src/realtime_agent.ts b/examples/src/realtime_agent.ts index b30171776..647d5759b 100644 --- a/examples/src/realtime_agent.ts +++ b/examples/src/realtime_agent.ts @@ -24,6 +24,7 @@ export default defineAgent({ }, entry: async (ctx: JobContext) => { const getWeather = llm.tool({ + name: 'getWeather', description: ' Called when the user asks about the weather.', parameters: z.object({ location: z.string().describe('The location to get the weather for'), @@ -34,6 +35,7 @@ export default defineAgent({ }); const toggleLight = llm.tool({ + name: 'toggleLight', description: 'Called when the user asks to turn on or off the light.', parameters: z.object({ room: roomNameSchema.describe('The room to turn the light in'), @@ -65,10 +67,7 @@ export default defineAgent({ instructions: "You are a helpful assistant created by LiveKit, always speaking English, you can hear the user's message and respond to it.", chatCtx, - tools: { - getWeather, - toggleLight, - }, + tools: [getWeather, toggleLight], }); const session = new voice.AgentSession({ diff --git a/examples/src/realtime_with_tts.ts b/examples/src/realtime_with_tts.ts index d87db7853..88c450f06 100644 --- a/examples/src/realtime_with_tts.ts +++ b/examples/src/realtime_with_tts.ts @@ -26,6 +26,7 @@ export default defineAgent({ const logger = log(); const getWeather = llm.tool({ + name: 'getWeather', description: 'Called when the user asks about the weather.', parameters: z.object({ location: z.string().describe('The location to get the weather for'), @@ -38,9 +39,7 @@ export default defineAgent({ const agent = new voice.Agent({ instructions: 'You are a helpful assistant. Always speak in English.', - tools: { - getWeather, - }, + tools: [getWeather], }); const session = new voice.AgentSession({ diff --git a/examples/src/restaurant_agent.ts b/examples/src/restaurant_agent.ts index d9faaf9a5..4fb5281ba 100644 --- a/examples/src/restaurant_agent.ts +++ b/examples/src/restaurant_agent.ts @@ -79,6 +79,7 @@ function summarize({ } const updateName = llm.tool({ + name: 'updateName', description: 'Called when the user provides their name. Confirm the spelling with the user before calling the function.', parameters: z.object({ @@ -91,6 +92,7 @@ const updateName = llm.tool({ }); const updatePhone = llm.tool({ + name: 'updatePhone', description: 'Called when the user provides their phone number. Confirm the spelling with the user before calling the function.', parameters: z.object({ @@ -103,6 +105,7 @@ const updatePhone = llm.tool({ }); const toGreeter = llm.tool({ + name: 'toGreeter', description: 'Called when user asks any unrelated questions or requests any other services not in your job description.', execute: async (_, { ctx }: llm.ToolOptions) => { @@ -173,35 +176,37 @@ function createGreeterAgent(menu: string) { instructions: `You are a friendly restaurant receptionist. The menu is: ${menu}\nYour jobs are to greet the caller and understand if they want to make a reservation or order takeaway. Guide them to the right agent using tools.`, llm: new inference.LLM({ model: 'openai/gpt-4.1-mini' }), tts: new inference.TTS({ model: 'cartesia/sonic-3', voice: voices.greeter }), - tools: { - toReservation: llm.tool({ + tools: [ + llm.tool({ + name: 'toReservation', description: dedent` Called when user wants to make or update a reservation. This function handles transitioning to the reservation agent who will collect the necessary details like reservation time, customer name and phone number. `, - execute: async (_, { ctx }): Promise => { + execute: async (_, { ctx }: llm.ToolOptions): Promise => { return await greeter.transferToAgent({ name: 'reservation', ctx, }); }, }), - toTakeaway: llm.tool({ + llm.tool({ + name: 'toTakeaway', description: dedent` Called when the user wants to place a takeaway order. This includes handling orders for pickup, delivery, or when the user wants to proceed to checkout with their existing order. `, - execute: async (_, { ctx }): Promise => { + execute: async (_, { ctx }: llm.ToolOptions): Promise => { return await greeter.transferToAgent({ name: 'takeaway', ctx, }); }, }), - }, + ], }); return greeter; @@ -212,11 +217,12 @@ function createReservationAgent() { name: 'reservation', instructions: `You are a reservation agent at a restaurant. Your jobs are to ask for the reservation time, then customer's name, and phone number. Then confirm the reservation details with the customer.`, tts: new inference.TTS({ model: 'cartesia/sonic-3', voice: voices.reservation }), - tools: { + tools: [ updateName, updatePhone, toGreeter, - updateReservationTime: llm.tool({ + llm.tool({ + name: 'updateReservationTime', description: dedent` Called when the user provides their reservation time. Confirm the time with the user before calling the function. @@ -224,14 +230,18 @@ function createReservationAgent() { parameters: z.object({ time: z.string().describe('The reservation time'), }), - execute: async ({ time }, { ctx }) => { + execute: async ({ time }, { ctx }: llm.ToolOptions) => { ctx.userData.reservationTime = time; return `The reservation time is updated to ${time}`; }, }), - confirmReservation: llm.tool({ + llm.tool({ + name: 'confirmReservation', description: `Called when the user confirms the reservation.`, - execute: async (_, { ctx }): Promise => { + execute: async ( + _, + { ctx }: llm.ToolOptions, + ): Promise => { const userdata = ctx.userData; if (!userdata.customer.name || !userdata.customer.phone) { return 'Please provide your name and phone number first.'; @@ -245,7 +255,7 @@ function createReservationAgent() { }); }, }), - }, + ], }); return reservation; @@ -256,21 +266,26 @@ function createTakeawayAgent(menu: string) { name: 'takeaway', instructions: `Your are a takeaway agent that takes orders from the customer. Our menu is: ${menu}\nClarify special requests and confirm the order with the customer.`, tts: new inference.TTS({ model: 'cartesia/sonic-3', voice: voices.takeaway }), - tools: { + tools: [ toGreeter, - updateOrder: llm.tool({ + llm.tool({ + name: 'updateOrder', description: `Called when the user provides their order.`, parameters: z.object({ items: z.array(z.string()).describe('The items of the full order'), }), - execute: async ({ items }, { ctx }) => { + execute: async ({ items }, { ctx }: llm.ToolOptions) => { ctx.userData.order = items; return `The order is updated to ${items}`; }, }), - toCheckout: llm.tool({ + llm.tool({ + name: 'toCheckout', description: `Called when the user confirms the order.`, - execute: async (_, { ctx }): Promise => { + execute: async ( + _, + { ctx }: llm.ToolOptions, + ): Promise => { const userdata = ctx.userData; if (!userdata.order) { return 'No takeaway order found. Please make an order first.'; @@ -281,7 +296,7 @@ function createTakeawayAgent(menu: string) { }); }, }), - }, + ], }); return takeaway; @@ -292,21 +307,23 @@ function createCheckoutAgent(menu: string) { name: 'checkout', instructions: `You are a checkout agent at a restaurant. The menu is: ${menu}\nYour are responsible for confirming the expense of the order and then collecting customer's name, phone number and credit card information, including the card number, expiry date, and CVV step by step.`, tts: new inference.TTS({ model: 'cartesia/sonic-3', voice: voices.checkout }), - tools: { + tools: [ updateName, updatePhone, toGreeter, - confirmExpense: llm.tool({ + llm.tool({ + name: 'confirmExpense', description: `Called when the user confirms the expense.`, parameters: z.object({ expense: z.number().describe('The expense of the order'), }), - execute: async ({ expense }, { ctx }) => { + execute: async ({ expense }, { ctx }: llm.ToolOptions) => { ctx.userData.expense = expense; return `The expense is confirmed to be ${expense}`; }, }), - updateCreditCard: llm.tool({ + llm.tool({ + name: 'updateCreditCard', description: dedent` Called when the user provides their credit card number, expiry date, and CVV. Confirm the spelling with the user before calling the function. @@ -316,14 +333,18 @@ function createCheckoutAgent(menu: string) { expiry: z.string().describe('The expiry date of the credit card'), cvv: z.string().describe('The CVV of the credit card'), }), - execute: async ({ number, expiry, cvv }, { ctx }) => { + execute: async ({ number, expiry, cvv }, { ctx }: llm.ToolOptions) => { ctx.userData.creditCard = { number, expiry, cvv }; return `The credit card number is updated to ${number}`; }, }), - confirmCheckout: llm.tool({ + llm.tool({ + name: 'confirmCheckout', description: `Called when the user confirms the checkout.`, - execute: async (_, { ctx }): Promise => { + execute: async ( + _, + { ctx }: llm.ToolOptions, + ): Promise => { const userdata = ctx.userData; if (!userdata.expense) { return 'Please confirm the expense first.'; @@ -342,16 +363,17 @@ function createCheckoutAgent(menu: string) { }); }, }), - toTakeaway: llm.tool({ + llm.tool({ + name: 'toTakeaway', description: `Called when the user wants to update their order.`, - execute: async (_, { ctx }): Promise => { + execute: async (_, { ctx }: llm.ToolOptions): Promise => { return await checkout.transferToAgent({ name: 'takeaway', ctx, }); }, }), - }, + ], }); return checkout; diff --git a/examples/src/survey_agent.ts b/examples/src/survey_agent.ts index 8504fa1a7..918fd9859 100644 --- a/examples/src/survey_agent.ts +++ b/examples/src/survey_agent.ts @@ -76,6 +76,7 @@ async function writeCsvRow(path: string, data: Record): Promise function disqualifyTool() { return llm.tool({ + name: 'disqualify', description: 'End the interview if the candidate refuses to cooperate, provides inappropriate answers, or is not a fit.', parameters: z.object({ @@ -101,8 +102,9 @@ export class IntroTask extends voice.AgentTask { super({ instructions: 'You are Alex, an interviewer screening a software engineer candidate. Gather the candidate name and short self-introduction.', - tools: { - saveIntro: llm.tool({ + tools: [ + llm.tool({ + name: 'saveIntro', description: 'Save candidate name and intro notes.', parameters: z.object({ name: z.string().describe('Candidate name'), @@ -114,7 +116,7 @@ export class IntroTask extends voice.AgentTask { return `Saved intro for ${name}.`; }, }), - }, + ], }); } @@ -132,9 +134,10 @@ export class EmailTask extends voice.AgentTask { super({ instructions: 'Collect a valid email address. If the candidate refuses, call disqualify immediately.', - tools: { + tools: [ disqualify, - saveEmail: llm.tool({ + llm.tool({ + name: 'saveEmail', description: 'Save candidate email address.', parameters: z.object({ email: z.string().describe('Candidate email'), @@ -144,7 +147,7 @@ export class EmailTask extends voice.AgentTask { return `Saved email: ${email}`; }, }), - }, + ], }); } @@ -161,9 +164,10 @@ export class CommuteTask extends voice.AgentTask super({ instructions: 'Collect commute flexibility. The role expects office attendance three days per week.', - tools: { + tools: [ disqualify, - saveCommute: llm.tool({ + llm.tool({ + name: 'saveCommute', description: 'Save candidate commute information.', parameters: z.object({ canCommute: z.boolean().describe('Whether the candidate can commute to office'), @@ -176,7 +180,7 @@ export class CommuteTask extends voice.AgentTask return 'Saved commute flexibility.'; }, }), - }, + ], }); } @@ -194,9 +198,10 @@ export class ExperienceTask extends voice.AgentTask { super({ instructions: 'You are a survey interviewer for a software engineer screening. Be concise, professional, and natural. Call endScreening when the process is complete.', - tools: { - endScreening: llm.tool({ + tools: [ + llm.tool({ + name: 'endScreening', description: 'End interview and hang up.', execute: async (_, { ctx }: llm.ToolOptions) => { ctx.session.shutdown(); return 'Interview concluded.'; }, }), - }, + ], }); } diff --git a/examples/src/testing/agent_task.test.ts b/examples/src/testing/agent_task.test.ts index 163d232a1..38d2a85d9 100644 --- a/examples/src/testing/agent_task.test.ts +++ b/examples/src/testing/agent_task.test.ts @@ -148,8 +148,9 @@ describe('AgentTask examples', { timeout: 120_000 }, () => { super({ instructions: 'You are collecting a name and role. Extract both from user input and call recordIntro.', - tools: { - recordIntro: llm.tool({ + tools: [ + llm.tool({ + name: 'recordIntro', description: 'Record the name and role', parameters: z.object({ name: z.string().describe('User name'), @@ -160,7 +161,7 @@ describe('AgentTask examples', { timeout: 120_000 }, () => { return 'recorded'; }, }), - }, + ], }); } @@ -222,8 +223,9 @@ describe('AgentTask examples', { timeout: 120_000 }, () => { super({ instructions: 'When asked to capture email, ALWAYS call captureEmail exactly once, then respond briefly.', - tools: { - captureEmail: llm.tool({ + tools: [ + llm.tool({ + name: 'captureEmail', description: 'Capture an email by running a nested AgentTask.', parameters: z.object({}), execute: async () => { @@ -236,7 +238,7 @@ describe('AgentTask examples', { timeout: 120_000 }, () => { } }, }), - }, + ], }); } @@ -275,8 +277,9 @@ describe('AgentTask examples', { timeout: 120_000 }, () => { instructions: 'You are Alex, an interviewer. Extract the candidate name and a short intro from the latest user input. ' + 'Use the tool recordIntro exactly once when both are available.', - tools: { - recordIntro: llm.tool({ + tools: [ + llm.tool({ + name: 'recordIntro', description: 'Record candidate name and intro summary.', parameters: z.object({ name: z.string().describe('Candidate name'), @@ -288,7 +291,7 @@ describe('AgentTask examples', { timeout: 120_000 }, () => { return 'Intro recorded.'; }, }), - }, + ], }); } @@ -305,8 +308,9 @@ describe('AgentTask examples', { timeout: 120_000 }, () => { super({ instructions: 'When the user asks to run the intro task, ALWAYS call collectIntroWithTask exactly once.', - tools: { - collectIntroWithTask: llm.tool({ + tools: [ + llm.tool({ + name: 'collectIntroWithTask', description: 'Launch the IntroTask and return the captured intro details.', parameters: z.object({}), execute: async () => { @@ -316,7 +320,7 @@ describe('AgentTask examples', { timeout: 120_000 }, () => { return JSON.stringify(result); }, }), - }, + ], }); } } diff --git a/examples/src/testing/basic_task_group.test.ts b/examples/src/testing/basic_task_group.test.ts index afbc7400e..5201a6c7c 100644 --- a/examples/src/testing/basic_task_group.test.ts +++ b/examples/src/testing/basic_task_group.test.ts @@ -82,8 +82,9 @@ class CollectNameTask extends voice.AgentTask { super({ instructions: 'Collect the user name from the latest user message. As soon as you have it, call save_name.', - tools: { - save_name: llm.tool({ + tools: [ + llm.tool({ + name: 'save_name', description: 'Save the user name.', parameters: z.object({ name: z.string().describe('The user name') }), execute: async ({ name }) => { @@ -91,7 +92,7 @@ class CollectNameTask extends voice.AgentTask { return `Saved name: ${name}`; }, }), - }, + ], }); this.ready = ready; } @@ -108,8 +109,9 @@ class CollectEmailTask extends voice.AgentTask { super({ instructions: 'Collect the user email from the latest user message. As soon as you have it, call save_email.', - tools: { - save_email: llm.tool({ + tools: [ + llm.tool({ + name: 'save_email', description: 'Save the user email.', parameters: z.object({ email: z.string().describe('The user email') }), execute: async ({ email }) => { @@ -117,7 +119,7 @@ class CollectEmailTask extends voice.AgentTask { return `Saved email: ${email}`; }, }), - }, + ], }); this.ready = ready; } diff --git a/examples/src/testing/run_result.test.ts b/examples/src/testing/run_result.test.ts index 583cbaffa..4169dea93 100644 --- a/examples/src/testing/run_result.test.ts +++ b/examples/src/testing/run_result.test.ts @@ -45,8 +45,9 @@ Response rules: - After ordering, confirm what was added (e.g., "I've added the burger to your order"). - When asked about sizes, always ask for clarification if not specified. - Be friendly and proactive in suggesting next steps.`, - tools: { - getWeather: llm.tool({ + tools: [ + llm.tool({ + name: 'getWeather', description: 'Get the current weather for a location', parameters: z.object({ location: z.string().describe('The city name'), @@ -59,14 +60,16 @@ Response rules: }); }, }), - getCurrentTime: llm.tool({ + llm.tool({ + name: 'getCurrentTime', description: 'Get the current time', parameters: z.object({}), execute: async () => { return '3:00 PM'; }, }), - orderItem: llm.tool({ + llm.tool({ + name: 'orderItem', description: 'Add an item to the order', parameters: z.object({ itemId: z.string().describe('The menu item ID'), @@ -84,7 +87,8 @@ Response rules: }); }, }), - getOrderStatus: llm.tool({ + llm.tool({ + name: 'getOrderStatus', description: 'Get the current order status', parameters: z.object({}), execute: async () => { @@ -98,7 +102,8 @@ Response rules: }); }, }), - getMenuItems: llm.tool({ + llm.tool({ + name: 'getMenuItems', description: 'Get available menu items and prices', parameters: z.object({ category: z @@ -128,7 +133,7 @@ Response rules: return JSON.stringify(menu); }, }), - }, + ], }); } } diff --git a/examples/src/testing/task_group.test.ts b/examples/src/testing/task_group.test.ts index 3d5afff06..188d46e59 100644 --- a/examples/src/testing/task_group.test.ts +++ b/examples/src/testing/task_group.test.ts @@ -260,8 +260,9 @@ describe('TaskGroup', { timeout: 120_000 }, () => { super({ instructions: 'Extract the user name from the latest user message. Call recordName immediately.', - tools: { - recordName: llm.tool({ + tools: [ + llm.tool({ + name: 'recordName', description: 'Record the user name', parameters: z.object({ name: z.string().describe('The user name') }), execute: async ({ name }) => { @@ -269,7 +270,7 @@ describe('TaskGroup', { timeout: 120_000 }, () => { return 'recorded'; }, }), - }, + ], }); } @@ -283,8 +284,9 @@ describe('TaskGroup', { timeout: 120_000 }, () => { super({ instructions: 'Extract an email address from the latest user message. Call recordEmail immediately.', - tools: { - recordEmail: llm.tool({ + tools: [ + llm.tool({ + name: 'recordEmail', description: 'Record the user email', parameters: z.object({ email: z.string().describe('The email address') }), execute: async ({ email }) => { @@ -292,7 +294,7 @@ describe('TaskGroup', { timeout: 120_000 }, () => { return 'recorded'; }, }), - }, + ], }); } @@ -400,8 +402,9 @@ describe('TaskGroup', { timeout: 120_000 }, () => { super({ instructions: 'Extract the user favorite color from the latest message. Call recordColor immediately.', - tools: { - recordColor: llm.tool({ + tools: [ + llm.tool({ + name: 'recordColor', description: 'Record favorite color', parameters: z.object({ color: z.string() }), execute: async ({ color }) => { @@ -409,7 +412,7 @@ describe('TaskGroup', { timeout: 120_000 }, () => { return 'recorded'; }, }), - }, + ], }); } @@ -423,8 +426,9 @@ describe('TaskGroup', { timeout: 120_000 }, () => { super({ instructions: 'Extract the user favorite food from the latest message. Call recordFood immediately.', - tools: { - recordFood: llm.tool({ + tools: [ + llm.tool({ + name: 'recordFood', description: 'Record favorite food', parameters: z.object({ food: z.string() }), execute: async ({ food }) => { @@ -432,7 +436,7 @@ describe('TaskGroup', { timeout: 120_000 }, () => { return 'recorded'; }, }), - }, + ], }); } diff --git a/examples/src/tool_call_disfluency.ts b/examples/src/tool_call_disfluency.ts index 8f92183a8..7e8018c99 100644 --- a/examples/src/tool_call_disfluency.ts +++ b/examples/src/tool_call_disfluency.ts @@ -39,6 +39,7 @@ export default defineAgent({ const vad = ctx.proc.userData.vad! as silero.VAD; const getWeather = llm.tool({ + name: 'getWeather', description: ' Called when the user asks about the weather.', parameters: z.object({ location: z.string().describe('The location to get the weather for'), @@ -55,9 +56,7 @@ export default defineAgent({ const agent = new VoiceAgent({ instructions: "You are a helpful assistant, you can hear the user's message and respond to it.", - tools: { - getWeather, - }, + tools: [getWeather], }); const session = new voice.AgentSession({ diff --git a/examples/src/xai-realtime.ts b/examples/src/xai-realtime.ts index ad383a0a6..a5130a640 100644 --- a/examples/src/xai-realtime.ts +++ b/examples/src/xai-realtime.ts @@ -10,8 +10,9 @@ export default defineAgent({ entry: async (ctx: JobContext) => { const agent = new voice.Agent({ instructions: 'You are a helpful assistant. Keep your responses short and concise.', - tools: { - getWeather: llm.tool({ + tools: [ + llm.tool({ + name: 'getWeather', description: 'Get the weather for a given location.', parameters: z.object({ location: z.string().describe('The location to get the weather for'), @@ -20,7 +21,7 @@ export default defineAgent({ return `The weather in ${location} is sunny.`; }, }), - }, + ], }); const session = new voice.AgentSession({ diff --git a/plugins/baseten/src/llm.ts b/plugins/baseten/src/llm.ts index 4039b7cac..179a09344 100644 --- a/plugins/baseten/src/llm.ts +++ b/plugins/baseten/src/llm.ts @@ -72,19 +72,20 @@ export class OpenAILLM extends llm.LLM { chat({ chatCtx, - toolCtx, + toolCtx: toolCtxInput, connOptions = DEFAULT_API_CONNECT_OPTIONS, parallelToolCalls, toolChoice, extraKwargs, }: { chatCtx: llm.ChatContext; - toolCtx?: llm.ToolContext; + toolCtx?: llm.ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: llm.ToolChoice; extraKwargs?: Record; }): inference.LLMStream { + const toolCtx = llm.toToolContext(toolCtxInput); const extras: Record = { ...extraKwargs }; if (this.#opts.metadata) { @@ -125,7 +126,11 @@ export class OpenAILLM extends llm.LLM { parallelToolCalls = parallelToolCalls !== undefined ? parallelToolCalls : this.#opts.parallelToolCalls; - if (toolCtx && Object.keys(toolCtx).length > 0 && parallelToolCalls !== undefined) { + if ( + toolCtx && + Object.keys(toolCtx.functionTools).length > 0 && + parallelToolCalls !== undefined + ) { extras.parallel_tool_calls = parallelToolCalls; } diff --git a/plugins/cerebras/src/llm.test.ts b/plugins/cerebras/src/llm.test.ts index 3a0a3ca8b..dda5d7b3c 100644 --- a/plugins/cerebras/src/llm.test.ts +++ b/plugins/cerebras/src/llm.test.ts @@ -88,8 +88,9 @@ class WeatherAgent extends voice.Agent { constructor() { super({ instructions: 'You are a helpful assistant.', - tools: { - get_weather: llm.tool({ + tools: [ + llm.tool({ + name: 'get_weather', description: 'Get the current weather for a location.', parameters: z.object({ location: z.string().describe('The city name'), @@ -98,7 +99,7 @@ class WeatherAgent extends voice.Agent { return `The weather in ${location} is sunny, 72°F.`; }, }), - }, + ], }); } } diff --git a/plugins/google/src/beta/realtime/realtime_api.ts b/plugins/google/src/beta/realtime/realtime_api.ts index 8b9ada6eb..66bc6a7f9 100644 --- a/plugins/google/src/beta/realtime/realtime_api.ts +++ b/plugins/google/src/beta/realtime/realtime_api.ts @@ -451,7 +451,7 @@ export class RealtimeModel extends llm.RealtimeModel { * supporting both text and audio modalities with function calling capabilities. */ export class RealtimeSession extends llm.RealtimeSession { - private _tools: llm.ToolContext = {}; + private _tools: llm.ToolContext = llm.ToolContext.empty(); private _chatCtx = llm.ChatContext.empty(); private options: RealtimeOptions; @@ -780,7 +780,7 @@ export class RealtimeSession extends llm.RealtimeSession { } get tools(): llm.ToolContext { - return { ...this._tools }; + return this._tools.copy(); } get manualActivityDetection(): boolean { diff --git a/plugins/google/src/llm.ts b/plugins/google/src/llm.ts index 302b679af..e452b70d2 100644 --- a/plugins/google/src/llm.ts +++ b/plugins/google/src/llm.ts @@ -189,20 +189,21 @@ export class LLM extends llm.LLM { chat({ chatCtx, - toolCtx, + toolCtx: toolCtxInput, connOptions = DEFAULT_API_CONNECT_OPTIONS, toolChoice, extraKwargs, geminiTools, }: { chatCtx: llm.ChatContext; - toolCtx?: llm.ToolContext; + toolCtx?: llm.ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: llm.ToolChoice; extraKwargs?: Record; geminiTools?: LLMTools; }): LLMStream { + const toolCtx = llm.toToolContext(toolCtxInput); const extras: GenerateContentConfig = { ...extraKwargs } as GenerateContentConfig; toolChoice = toolChoice !== undefined ? toolChoice : this.#opts.toolChoice; @@ -218,7 +219,7 @@ export class LLM extends llm.LLM { }, }; } else if (toolChoice === 'required') { - const toolNames = Object.entries(toolCtx || {}).map(([name]) => name); + const toolNames = Object.keys(toolCtx?.functionTools ?? {}); geminiToolConfig = { functionCallingConfig: { mode: FunctionCallingConfigMode.ANY, diff --git a/plugins/google/src/utils.ts b/plugins/google/src/utils.ts index 732ae0c3d..5548c076e 100644 --- a/plugins/google/src/utils.ts +++ b/plugins/google/src/utils.ts @@ -139,7 +139,7 @@ function isEmptyObjectSchema(jsonSchema: JSONSchema7Definition): boolean { export function toFunctionDeclarations(toolCtx: llm.ToolContext): FunctionDeclaration[] { const functionDeclarations: FunctionDeclaration[] = []; - for (const [name, tool] of Object.entries(toolCtx)) { + for (const [name, tool] of Object.entries(toolCtx.functionTools)) { const { description, parameters } = tool; const jsonSchema = llm.toJsonSchema(parameters, false); diff --git a/plugins/mistralai/src/llm.ts b/plugins/mistralai/src/llm.ts index f0c07f7a2..f6685b042 100644 --- a/plugins/mistralai/src/llm.ts +++ b/plugins/mistralai/src/llm.ts @@ -123,7 +123,7 @@ export class LLM extends llm.LLM { extraKwargs, }: { chatCtx: llm.ChatContext; - toolCtx?: llm.ToolContext; + toolCtx?: llm.ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: llm.ToolChoice; @@ -187,7 +187,7 @@ export class LLMStream extends llm.LLMStream { client: Mistral; opts: LLMOpts; chatCtx: llm.ChatContext; - toolCtx?: llm.ToolContext; + toolCtx?: llm.ToolCtxInput; connOptions: APIConnectOptions; extraKwargs: Record; }, @@ -211,8 +211,8 @@ export class LLMStream extends llm.LLMStream { // eslint-disable-next-line @typescript-eslint/no-explicit-any const toolsList: any[] = []; - if (this.toolCtx && Object.keys(this.toolCtx).length > 0) { - for (const [name, func] of Object.entries(this.toolCtx)) { + if (this.toolCtx && Object.keys(this.toolCtx.functionTools).length > 0) { + for (const [name, func] of Object.entries(this.toolCtx.functionTools)) { toolsList.push({ type: 'function' as const, function: { diff --git a/plugins/openai/src/llm.ts b/plugins/openai/src/llm.ts index e96551011..a358f5abd 100644 --- a/plugins/openai/src/llm.ts +++ b/plugins/openai/src/llm.ts @@ -468,19 +468,20 @@ export class LLM extends llm.LLM { chat({ chatCtx, - toolCtx, + toolCtx: toolCtxInput, connOptions = DEFAULT_API_CONNECT_OPTIONS, parallelToolCalls, toolChoice, extraKwargs, }: { chatCtx: llm.ChatContext; - toolCtx?: llm.ToolContext; + toolCtx?: llm.ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: llm.ToolChoice; extraKwargs?: Record; }): LLMStream { + const toolCtx = llm.toToolContext(toolCtxInput); const extras: Record = { ...extraKwargs }; if (this.#opts.metadata) { @@ -509,7 +510,11 @@ export class LLM extends llm.LLM { parallelToolCalls = parallelToolCalls !== undefined ? parallelToolCalls : this.#opts.parallelToolCalls; - if (toolCtx && Object.keys(toolCtx).length > 0 && parallelToolCalls !== undefined) { + if ( + toolCtx && + Object.keys(toolCtx.functionTools).length > 0 && + parallelToolCalls !== undefined + ) { extras.parallel_tool_calls = parallelToolCalls; } diff --git a/plugins/openai/src/realtime/realtime_model.ts b/plugins/openai/src/realtime/realtime_model.ts index 3a6a82095..8481e0f47 100644 --- a/plugins/openai/src/realtime/realtime_model.ts +++ b/plugins/openai/src/realtime/realtime_model.ts @@ -413,7 +413,7 @@ function processBaseURL({ * - openai_client_event_queued: expose the raw client events sent to the OpenAI Realtime API */ export class RealtimeSession extends llm.RealtimeSession { - private _tools: llm.ToolContext = {}; + private _tools: llm.ToolContext = llm.ToolContext.empty(); private remoteChatCtx: llm.RemoteChatContext = new llm.RemoteChatContext(); private messageChannel = new Queue(); private inputResampler?: AudioResampler; @@ -536,7 +536,7 @@ export class RealtimeSession extends llm.RealtimeSession { } get tools() { - return { ...this._tools } as llm.ToolContext; + return this._tools.copy(); } async updateChatCtx(_chatCtx: llm.ChatContext): Promise { @@ -698,13 +698,11 @@ export class RealtimeSession extends llm.RealtimeSession { // TODO(brian): these logics below are noops I think, leaving it here to keep // parity with the python but we should remove them later const retainedToolNames = new Set(ev.session.tools.map((tool) => tool.name)); - const retainedTools = Object.fromEntries( - Object.entries(_tools).filter( - ([name, tool]) => llm.isFunctionTool(tool) && retainedToolNames.has(name), - ), - ); + const retainedTools = Object.entries(_tools.functionTools) + .filter(([name]) => retainedToolNames.has(name)) + .map(([, tool]) => tool); - this._tools = retainedTools as llm.ToolContext; + this._tools = new llm.ToolContext(retainedTools); unlock(); } @@ -712,12 +710,7 @@ export class RealtimeSession extends llm.RealtimeSession { private createToolsUpdateEvent(_tools: llm.ToolContext): api_proto.SessionUpdateEvent { const oaiTools: api_proto.Tool[] = []; - for (const [name, tool] of Object.entries(_tools)) { - if (!llm.isFunctionTool(tool)) { - this.#logger.error({ name, tool }, "OpenAI Realtime API doesn't support this tool type"); - continue; - } - + for (const [name, tool] of Object.entries(_tools.functionTools)) { const { parameters: toolParameters, description } = tool; try { const parameters = llm.toJsonSchema( @@ -998,7 +991,7 @@ export class RealtimeSession extends llm.RealtimeSession { events.push(this.createSessionUpdateEvent()); // tools - if (Object.keys(this._tools).length > 0) { + if (Object.keys(this._tools.functionTools).length > 0) { events.push(this.createToolsUpdateEvent(this._tools)); } diff --git a/plugins/openai/src/responses/llm.ts b/plugins/openai/src/responses/llm.ts index 494a05c87..9a255d046 100644 --- a/plugins/openai/src/responses/llm.ts +++ b/plugins/openai/src/responses/llm.ts @@ -77,25 +77,30 @@ class ResponsesHttpLLM extends llm.LLM { override chat({ chatCtx, - toolCtx, + toolCtx: toolCtxInput, connOptions = DEFAULT_API_CONNECT_OPTIONS, parallelToolCalls, toolChoice, extraKwargs, }: { chatCtx: llm.ChatContext; - toolCtx?: llm.ToolContext; + toolCtx?: llm.ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: llm.ToolChoice; extraKwargs?: Record; }): ResponsesHttpLLMStream { + const toolCtx = llm.toToolContext(toolCtxInput); const modelOptions: Record = { ...(extraKwargs || {}) }; parallelToolCalls = parallelToolCalls !== undefined ? parallelToolCalls : this.#opts.parallelToolCalls; - if (toolCtx && Object.keys(toolCtx).length > 0 && parallelToolCalls !== undefined) { + if ( + toolCtx && + Object.keys(toolCtx.functionTools).length > 0 && + parallelToolCalls !== undefined + ) { modelOptions.parallel_tool_calls = parallelToolCalls; } @@ -182,7 +187,7 @@ class ResponsesHttpLLMStream extends llm.LLMStream { )) as OpenAI.Responses.ResponseInputItem[]; const tools = this.toolCtx - ? Object.entries(this.toolCtx).map(([name, func]) => { + ? Object.entries(this.toolCtx.functionTools).map(([name, func]) => { const oaiParams = { type: 'function' as const, name: name, @@ -417,7 +422,7 @@ export class LLM extends llm.LLM { extraKwargs, }: { chatCtx: llm.ChatContext; - toolCtx?: llm.ToolContext; + toolCtx?: llm.ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: llm.ToolChoice; diff --git a/plugins/openai/src/ws/llm.ts b/plugins/openai/src/ws/llm.ts index 1dd6483c6..d22d7a753 100644 --- a/plugins/openai/src/ws/llm.ts +++ b/plugins/openai/src/ws/llm.ts @@ -231,24 +231,29 @@ export class WSLLM extends llm.LLM { chat({ chatCtx, - toolCtx, + toolCtx: toolCtxInput, connOptions = DEFAULT_API_CONNECT_OPTIONS, parallelToolCalls, toolChoice, extraKwargs, }: { chatCtx: llm.ChatContext; - toolCtx?: llm.ToolContext; + toolCtx?: llm.ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: llm.ToolChoice; extraKwargs?: Record; }): WSLLMStream { + const toolCtx = llm.toToolContext(toolCtxInput); const modelOptions: Record = { ...(extraKwargs ?? {}) }; parallelToolCalls = parallelToolCalls !== undefined ? parallelToolCalls : this.#opts.parallelToolCalls; - if (toolCtx && Object.keys(toolCtx).length > 0 && parallelToolCalls !== undefined) { + if ( + toolCtx && + Object.keys(toolCtx.functionTools).length > 0 && + parallelToolCalls !== undefined + ) { modelOptions.parallel_tool_calls = parallelToolCalls; } @@ -425,7 +430,7 @@ export class WSLLMStream extends llm.LLMStream { )) as OpenAI.Responses.ResponseInputItem[]; const tools = this.toolCtx - ? Object.entries(this.toolCtx).map(([name, func]) => { + ? Object.entries(this.toolCtx.functionTools).map(([name, func]) => { const oaiParams = { type: 'function' as const, name, diff --git a/plugins/phonic/src/realtime/realtime_model.ts b/plugins/phonic/src/realtime/realtime_model.ts index 665c5c5ba..09933b580 100644 --- a/plugins/phonic/src/realtime/realtime_model.ts +++ b/plugins/phonic/src/realtime/realtime_model.ts @@ -239,7 +239,7 @@ interface GenerationState { * Realtime session for Phonic (https://docs.phonic.co/) */ export class RealtimeSession extends llm.RealtimeSession { - private _tools: llm.ToolContext = {}; + private _tools: llm.ToolContext = llm.ToolContext.empty(); private _chatCtx = llm.ChatContext.empty(); private options: RealtimeModelOptions; @@ -290,7 +290,7 @@ export class RealtimeSession extends llm.RealtimeSession { } get tools(): llm.ToolContext { - return { ...this._tools }; + return this._tools.copy(); } async updateInstructions(instructions: string): Promise { @@ -367,26 +367,24 @@ export class RealtimeSession extends llm.RealtimeSession { return; } - this._tools = { ...tools }; - this.toolDefinitions = Object.entries(tools) - .filter(([_, tool]) => llm.isFunctionTool(tool)) - .map(([name, tool]) => ({ - type: 'custom_websocket', - tool_schema: { - type: 'function', - function: { - name, - description: tool.description, - parameters: llm.toJsonSchema(tool.parameters), - strict: true, - }, + this._tools = tools.copy(); + this.toolDefinitions = Object.entries(tools.functionTools).map(([name, tool]) => ({ + type: 'custom_websocket', + tool_schema: { + type: 'function', + function: { + name, + description: tool.description, + parameters: llm.toJsonSchema(tool.parameters), + strict: true, }, - tool_call_output_timeout_ms: TOOL_CALL_OUTPUT_TIMEOUT_MS, - // Tool chaining and tool calls during speech are not supported at this time - // for ease of implementation within the RealtimeSession generations framework - wait_for_speech_before_tool_call: true, - allow_tool_chaining: false, - })); + }, + tool_call_output_timeout_ms: TOOL_CALL_OUTPUT_TIMEOUT_MS, + // Tool chaining and tool calls during speech are not supported at this time + // for ease of implementation within the RealtimeSession generations framework + wait_for_speech_before_tool_call: true, + allow_tool_chaining: false, + })); this.toolsReady.resolve(); } @@ -405,24 +403,22 @@ export class RealtimeSession extends llm.RealtimeSession { this.options.instructions = instructions; } if (tools !== undefined) { - this._tools = { ...tools }; - this.toolDefinitions = Object.entries(tools) - .filter(([, tool]) => llm.isFunctionTool(tool)) - .map(([name, tool]) => ({ - type: 'custom_websocket', - tool_schema: { - type: 'function', - function: { - name, - description: tool.description, - parameters: llm.toJsonSchema(tool.parameters), - strict: true, - }, + this._tools = tools.copy(); + this.toolDefinitions = Object.entries(tools.functionTools).map(([name, tool]) => ({ + type: 'custom_websocket', + tool_schema: { + type: 'function', + function: { + name, + description: tool.description, + parameters: llm.toJsonSchema(tool.parameters), + strict: true, }, - tool_call_output_timeout_ms: TOOL_CALL_OUTPUT_TIMEOUT_MS, - wait_for_speech_before_tool_call: true, - allow_tool_chaining: false, - })); + }, + tool_call_output_timeout_ms: TOOL_CALL_OUTPUT_TIMEOUT_MS, + wait_for_speech_before_tool_call: true, + allow_tool_chaining: false, + })); } if (chatCtx !== undefined) { this._chatCtx = chatCtx.copy(); diff --git a/plugins/test/src/llm.ts b/plugins/test/src/llm.ts index 8d85c75c3..534dce2df 100644 --- a/plugins/test/src/llm.ts +++ b/plugins/test/src/llm.ts @@ -5,8 +5,9 @@ import { initializeLogger, llm as llmlib } from '@livekit/agents'; import { describe, expect, it } from 'vitest'; import { z } from 'zod/v4'; -const toolCtx: llmlib.ToolContext = { - getWeather: llmlib.tool({ +const toolCtx = new llmlib.ToolContext([ + llmlib.tool({ + name: 'getWeather', description: 'Get the current weather in a given location', parameters: z.object({ location: z.string().describe('The city and state, e.g. San Francisco, CA'), @@ -14,14 +15,16 @@ const toolCtx: llmlib.ToolContext = { }), execute: async () => {}, }), - playMusic: llmlib.tool({ + llmlib.tool({ + name: 'playMusic', description: 'Play music', parameters: z.object({ name: z.string().describe('The artist and name of the song'), }), execute: async () => {}, }), - toggleLight: llmlib.tool({ + llmlib.tool({ + name: 'toggleLight', description: 'Turn on/off the lights in a room', parameters: z.object({ name: z.string().describe('The room to control'), @@ -31,7 +34,8 @@ const toolCtx: llmlib.ToolContext = { await new Promise((resolve) => setTimeout(resolve, 60_000)); }, }), - selectCurrencies: llmlib.tool({ + llmlib.tool({ + name: 'selectCurrencies', description: 'Currencies of a specific area', parameters: z.object({ currencies: z @@ -40,7 +44,8 @@ const toolCtx: llmlib.ToolContext = { }), execute: async () => {}, }), - updateUserInfo: llmlib.tool({ + llmlib.tool({ + name: 'updateUserInfo', description: 'Update user info.', parameters: z.object({ email: z.string().optional().describe("User's email address"), @@ -49,18 +54,20 @@ const toolCtx: llmlib.ToolContext = { }), execute: async () => {}, }), - simulateFailure: llmlib.tool({ + llmlib.tool({ + name: 'simulateFailure', description: 'Simulate a failure', parameters: z.object({}), execute: async () => { throw new Error('Simulated failure'); }, }), -}; +]); // Tool context for strict mode - uses nullable() instead of optional() -const toolCtxStrict: llmlib.ToolContext = { - getWeather: llmlib.tool({ +const toolCtxStrict = new llmlib.ToolContext([ + llmlib.tool({ + name: 'getWeather', description: 'Get the current weather in a given location', parameters: z.object({ location: z.string().describe('The city and state, e.g. San Francisco, CA'), @@ -68,14 +75,16 @@ const toolCtxStrict: llmlib.ToolContext = { }), execute: async () => {}, }), - playMusic: llmlib.tool({ + llmlib.tool({ + name: 'playMusic', description: 'Play music', parameters: z.object({ name: z.string().describe('The artist and name of the song'), }), execute: async () => {}, }), - toggleLight: llmlib.tool({ + llmlib.tool({ + name: 'toggleLight', description: 'Turn on/off the lights in a room', parameters: z.object({ name: z.string().describe('The room to control'), @@ -85,7 +94,8 @@ const toolCtxStrict: llmlib.ToolContext = { await new Promise((resolve) => setTimeout(resolve, 60_000)); }, }), - selectCurrencies: llmlib.tool({ + llmlib.tool({ + name: 'selectCurrencies', description: 'Currencies of a specific area', parameters: z.object({ currencies: z @@ -94,7 +104,8 @@ const toolCtxStrict: llmlib.ToolContext = { }), execute: async () => {}, }), - updateUserInfo: llmlib.tool({ + llmlib.tool({ + name: 'updateUserInfo', description: 'Update user info.', parameters: z.object({ email: z.string().nullable().describe("User's email address"), @@ -103,14 +114,15 @@ const toolCtxStrict: llmlib.ToolContext = { }), execute: async () => {}, }), - simulateFailure: llmlib.tool({ + llmlib.tool({ + name: 'simulateFailure', description: 'Simulate a failure', parameters: z.object({}), execute: async () => { throw new Error('Simulated failure'); }, }), -}; +]); export const llm = async (llm: llmlib.LLM, skipOptionalArgs: boolean) => { initializeLogger({ pretty: false }); @@ -315,7 +327,7 @@ const executeCalls = async (calls: llmlib.FunctionCall[]) => { const results: llmlib.FunctionCallOutput[] = []; for (const call of calls) { - const tool = toolCtx[call.name]; + const tool = toolCtx.getFunctionTool(call.name); if (!tool) { throw new Error(`Tool ${call.name} not found`); } From 6b85bfea91af00a17a80860a674aff456ba5c9df Mon Sep 17 00:00:00 2001 From: Brian Yin Date: Tue, 2 Jun 2026 01:23:21 +0800 Subject: [PATCH 2/6] feat(agents): add Toolset support to ToolContext and AgentActivity (#1525) Co-authored-by: rosetta-livekit-bot[bot] <282703043+rosetta-livekit-bot[bot]@users.noreply.github.com> Co-authored-by: u9g --- .changeset/cold-avocados-behave.md | 5 + .changeset/gemini-provider-tools.md | 5 + .changeset/list-syntax-toolcontext.md | 15 +- .changeset/openai-provider-tools.md | 5 + .changeset/quick-meals-breathe.md | 5 + agents/src/generator.test.ts | 19 + agents/src/generator.ts | 23 +- agents/src/index.test.ts | 23 + agents/src/index.ts | 6 +- agents/src/llm/chat_context.test.ts | 5 +- agents/src/llm/index.ts | 9 +- agents/src/llm/llm.ts | 2 +- agents/src/llm/tool_context.test.ts | 163 +++++- agents/src/llm/tool_context.ts | 320 ++++++++---- agents/src/llm/tool_context.type.test.ts | 21 +- agents/src/utils.test.ts | 79 +++ agents/src/utils.ts | 52 +- agents/src/voice/agent.test.ts | 356 ++++++++++++- agents/src/voice/agent.ts | 25 + agents/src/voice/agent_activity.ts | 59 ++- agents/src/voice/agent_v2.ts | 484 ++++++++++++++++++ agents/src/voice/index.ts | 8 + examples/src/basic_agent.ts | 20 +- examples/src/basic_agent_task.ts | 180 +++---- examples/src/basic_toolsets.ts | 180 +++++++ examples/src/gemini_realtime_agent.ts | 8 +- examples/src/multi_agent.ts | 8 +- .../google/src/beta/realtime/realtime_api.ts | 43 +- plugins/google/src/index.ts | 1 + plugins/google/src/llm.ts | 12 +- plugins/google/src/tools.ts | 98 +++- plugins/google/src/utils.ts | 62 ++- plugins/mistralai/src/llm.ts | 12 +- plugins/openai/src/index.ts | 1 + plugins/openai/src/realtime/realtime_model.ts | 26 +- plugins/openai/src/responses/llm.ts | 20 +- plugins/openai/src/tool_utils.test.ts | 84 +++ plugins/openai/src/tool_utils.ts | 43 ++ plugins/openai/src/tools.ts | 166 ++++++ plugins/openai/src/ws/llm.ts | 22 +- plugins/phonic/src/realtime/realtime_model.ts | 66 +-- plugins/test/src/llm.ts | 51 ++ 42 files changed, 2399 insertions(+), 393 deletions(-) create mode 100644 .changeset/cold-avocados-behave.md create mode 100644 .changeset/gemini-provider-tools.md create mode 100644 .changeset/openai-provider-tools.md create mode 100644 .changeset/quick-meals-breathe.md create mode 100644 agents/src/generator.test.ts create mode 100644 agents/src/index.test.ts create mode 100644 agents/src/voice/agent_v2.ts create mode 100644 examples/src/basic_toolsets.ts create mode 100644 plugins/openai/src/tool_utils.test.ts create mode 100644 plugins/openai/src/tool_utils.ts create mode 100644 plugins/openai/src/tools.ts diff --git a/.changeset/cold-avocados-behave.md b/.changeset/cold-avocados-behave.md new file mode 100644 index 000000000..a2eef8a7b --- /dev/null +++ b/.changeset/cold-avocados-behave.md @@ -0,0 +1,5 @@ +--- +"@livekit/agents": patch +--- + +Add Agent.create method diff --git a/.changeset/gemini-provider-tools.md b/.changeset/gemini-provider-tools.md new file mode 100644 index 000000000..3b1093432 --- /dev/null +++ b/.changeset/gemini-provider-tools.md @@ -0,0 +1,5 @@ +--- +'@livekit/agents-plugin-google': minor +--- + +Add Gemini provider tools for Google Search, Google Maps, URL context, File Search, code execution, and Vertex RAG retrieval, and serialize them from `ToolContext` for Google LLM and realtime sessions. diff --git a/.changeset/list-syntax-toolcontext.md b/.changeset/list-syntax-toolcontext.md index 0a4e0bfa0..5e854a813 100644 --- a/.changeset/list-syntax-toolcontext.md +++ b/.changeset/list-syntax-toolcontext.md @@ -1,5 +1,16 @@ --- -"@livekit/agents": minor +'@livekit/agents': minor --- -**BREAKING**: `Agent({ tools })` and `agent.updateTools()` now accept a flat list `(FunctionTool | ProviderDefinedTool)[]` instead of a `Record` map, and `llm.tool({ ... })` requires a `name` field. `ToolContext` is now a Python-parity class with `functionTools` / `providerTools` / `toolsets` accessors, plus `flatten()`, `hasTool(name)`, `getFunctionTool(name)`, `updateTools()`, `copy()`, and `equals()`. To match the Python reference, registering two **different** function-tool instances under the same `name` now throws `duplicate function name: ` instead of silently overriding the earlier entry; passing the **same instance** twice is a no-op. `agent.toolCtx` returns a defensive copy so callers can no longer mutate the agent's internal state. `LLM.chat({ toolCtx })` accepts either a `ToolContext` instance or a raw `(FunctionTool | ProviderDefinedTool)[]` array (`ToolCtxInput`) and normalizes it internally, so callers don't have to construct a `ToolContext` themselves. Stateful `Toolset` containers are not part of this release — the `toolsets` accessor currently returns an empty list and `TODO`s in `tool_context.ts` mark every site where Python's Toolset support will plug in later. +**BREAKING**: `Agent({ tools })` and `agent.updateTools()` now accept a flat list `(FunctionTool | ProviderTool | Toolset)[]` instead of a `Record` map, and `llm.tool({ ... })` requires a `name` field. `ToolContext` is now a Python-parity class with `functionTools` / `providerTools` / `toolsets` accessors, plus `flatten()`, `hasTool(id)`, `getFunctionTool(id)`, `updateTools()`, `copy()`, and `equals()`. To match the Python reference, registering two **different** function-tool instances under the same `name` now throws `duplicate function name: ` instead of silently overriding the earlier entry; passing the **same instance** twice is a no-op. `agent.toolCtx` returns a defensive copy so callers can no longer mutate the agent's internal state. `LLM.chat({ toolCtx })` accepts either a `ToolContext` instance or a raw `(FunctionTool | ProviderTool | Toolset)[]` array (`ToolCtxInput`) and normalizes it internally, so callers don't have to construct a `ToolContext` themselves. + +Tools also expose an `id: string` field on the base `Tool` interface (parity with Python's `Tool.id` property): for `FunctionTool` it mirrors `name`, for `ProviderTool` it is the provider tool id. `ToolContext` keys and equality now use `tool.id` consistently. + +**BREAKING**: Provider tools are now modeled to match Python's `ProviderTool`: + +- `ProviderDefinedTool` is renamed to `ProviderTool`, and `isProviderDefinedTool` is renamed to `isProviderTool`. +- `ProviderTool` is now an **abstract class** (Python parity). Plugins must subclass it (`class WebSearch extends ProviderTool { ... }`) to attach provider-specific fields and serializers; bare `new ProviderTool(...)` is rejected at compile time. +- The `tool({ id })` factory overload is removed; `tool({ ... })` only creates function tools now. Construct provider tools by instantiating a `ProviderTool` subclass. +- The `ToolType` literal for provider tools is renamed from `'provider-defined'` to `'provider'`. + +`Toolset` now carries a `TOOLSET_SYMBOL` marker and is detected via a new `isToolset()` guard (consistent with `isFunctionTool` / `isProviderTool`). Existing `instanceof Toolset` checks still work, but symbol-based detection is preferred for cross-realm safety. diff --git a/.changeset/openai-provider-tools.md b/.changeset/openai-provider-tools.md new file mode 100644 index 000000000..8e793a935 --- /dev/null +++ b/.changeset/openai-provider-tools.md @@ -0,0 +1,5 @@ +--- +'@livekit/agents-plugin-openai': minor +--- + +Add OpenAI Responses provider tools for web search, file search, and code interpreter. diff --git a/.changeset/quick-meals-breathe.md b/.changeset/quick-meals-breathe.md new file mode 100644 index 000000000..d32233b93 --- /dev/null +++ b/.changeset/quick-meals-breathe.md @@ -0,0 +1,5 @@ +--- +'@livekit/agents': patch +--- + +Adds base `Toolset` support: a stateful container for a group of tools with `setup()` / `aclose()` lifecycle hooks. Toolsets can be passed directly into `Agent({ tools: [...] })` alongside individual function tools; their tools are flattened into the agent's `ToolContext` and the runtime drives `setup()` on activity start, `aclose()` on close, and a setup/close diff when `agent.updateTools()` adds or removes Toolsets mid-session. Per-toolset `setup()` errors are logged but do not abort the activity. The `IGNORE_ON_ENTER` flag is also respected for function tools nested inside a Toolset. Every LLM and realtime plugin tool builder iterates `ToolContext.flatten()` so toolset-contributed tools are correctly advertised. Also exports `ToolCalledEvent` / `ToolCompletedEvent` payload types. diff --git a/agents/src/generator.test.ts b/agents/src/generator.test.ts new file mode 100644 index 000000000..a92ff9e58 --- /dev/null +++ b/agents/src/generator.test.ts @@ -0,0 +1,19 @@ +// SPDX-FileCopyrightText: 2026 LiveKit, Inc. +// +// SPDX-License-Identifier: Apache-2.0 +import { describe, expect, it } from 'vitest'; +import { defineAgent, isAgent } from './generator.js'; + +describe('generator', () => { + it('marks definitions created with defineAgent as agents', () => { + const agent = defineAgent({ + entry: async () => {}, + }); + + expect(isAgent(agent)).toBe(true); + }); + + it('does not treat unmarked structural objects as agents', () => { + expect(isAgent({ entry: async () => {} })).toBe(false); + }); +}); diff --git a/agents/src/generator.ts b/agents/src/generator.ts index b7e7ad93f..e9bb09075 100644 --- a/agents/src/generator.ts +++ b/agents/src/generator.ts @@ -3,24 +3,22 @@ // SPDX-License-Identifier: Apache-2.0 import type { JobContext, JobProcess } from './job.js'; +export const AGENT_DEFINITION_SYMBOL = Symbol.for('livekit.agents.AgentDefinition'); + /** @see {@link defineAgent} */ -export interface Agent> { +export interface AgentDefinition> { entry: (ctx: JobContext) => Promise; prewarm?: (proc: JobProcess) => unknown; } +export type Agent> = AgentDefinition; + /** Helper to check if an object is an agent before running it. * * @internal */ -export function isAgent(obj: unknown): obj is Agent { - return ( - typeof obj === 'object' && - obj !== null && - 'entry' in obj && - typeof (obj as Agent).entry === 'function' && - (('prewarm' in obj && typeof (obj as Agent).prewarm === 'function') || !('prewarm' in obj)) - ); +export function isAgent(obj: unknown): obj is AgentDefinition { + return typeof obj === 'object' && obj !== null && AGENT_DEFINITION_SYMBOL in obj; } /** @@ -34,7 +32,10 @@ export function isAgent(obj: unknown): obj is Agent { * ``` */ export function defineAgent>( - agent: Agent, -): Agent { + agent: AgentDefinition, +): AgentDefinition { + Object.defineProperty(agent, AGENT_DEFINITION_SYMBOL, { + value: true, + }); return agent; } diff --git a/agents/src/index.test.ts b/agents/src/index.test.ts new file mode 100644 index 000000000..40b7ade44 --- /dev/null +++ b/agents/src/index.test.ts @@ -0,0 +1,23 @@ +// SPDX-FileCopyrightText: 2026 LiveKit, Inc. +// +// SPDX-License-Identifier: Apache-2.0 +import { describe, expect, it } from 'vitest'; +import { + Agent, + AgentSession, + ChatContext, + ModelUsageCollector, + logMetrics, + tool, +} from './index.js'; + +describe('index exports', () => { + it('exports voice, llm, and metrics APIs directly from the package root', () => { + expect(Agent).toBeDefined(); + expect(AgentSession).toBeDefined(); + expect(ChatContext).toBeDefined(); + expect(tool).toBeDefined(); + expect(ModelUsageCollector).toBeDefined(); + expect(logMetrics).toBeDefined(); + }); +}); diff --git a/agents/src/index.ts b/agents/src/index.ts index 07e4b45da..9f1519bf3 100644 --- a/agents/src/index.ts +++ b/agents/src/index.ts @@ -14,14 +14,16 @@ export * from './audio.js'; export * as beta from './beta/index.js'; export * as cli from './cli.js'; export * from './connection_pool.js'; -export * from './generator.js'; +export { defineAgent, isAgent, type AgentDefinition } from './generator.js'; export * as inference from './inference/index.js'; export * from './inference_runner.js'; export * as ipc from './ipc/index.js'; export * from './job.js'; export * from './language.js'; +export * from './llm/index.js'; export * as llm from './llm/index.js'; export * from './log.js'; +export * from './metrics/index.js'; export * as metrics from './metrics/index.js'; export * from './plugin.js'; export * as stream from './stream/index.js'; @@ -34,6 +36,6 @@ export * from './types.js'; export * from './utils.js'; export * from './vad.js'; export * from './version.js'; +export * from './voice/index.js'; export * as voice from './voice/index.js'; -export { createTimedString, isTimedString, type TimedString } from './voice/io.js'; export * from './worker.js'; diff --git a/agents/src/llm/chat_context.test.ts b/agents/src/llm/chat_context.test.ts index 350dcd1cd..101a9fe46 100644 --- a/agents/src/llm/chat_context.test.ts +++ b/agents/src/llm/chat_context.test.ts @@ -19,7 +19,7 @@ import { isInstructions, renderInstructions, } from './chat_context.js'; -import { ToolContext, tool } from './tool_context.js'; +import { ProviderTool, ToolContext, tool } from './tool_context.js'; initializeLogger({ pretty: false, level: 'error' }); @@ -1498,7 +1498,8 @@ describe('ChatContext.copy with toolCtx filter', () => { }); it('keeps provider-tool calls when the ToolContext holds a matching provider tool id', () => { - const provider = tool({ id: 'code_runner', config: {} }); + class CodeRunner extends ProviderTool {} + const provider = new CodeRunner({ id: 'code_runner' }); const ctx = new ChatContext([ FunctionCall.create({ callId: 'p1', name: 'code_runner', args: '{}' }), FunctionCall.create({ callId: 'p2', name: 'other', args: '{}' }), diff --git a/agents/src/llm/index.ts b/agents/src/llm/index.ts index 4837f2cb9..098a84fca 100644 --- a/agents/src/llm/index.ts +++ b/agents/src/llm/index.ts @@ -4,21 +4,26 @@ export { handoff, isFunctionTool, - isProviderDefinedTool, + isProviderTool, isTool, + isToolset, + ProviderTool, tool, ToolContext, ToolError, ToolFlag, + Toolset, toToolContext, type AgentHandoff, type FunctionTool, - type ProviderDefinedTool, type Tool, + type ToolCalledEvent, type ToolChoice, + type ToolCompletedEvent, type ToolContextEntry, type ToolCtxInput, type ToolOptions, + type ToolsetCreateOptions, type ToolType, } from './tool_context.js'; diff --git a/agents/src/llm/llm.ts b/agents/src/llm/llm.ts index 0c05bbb2d..828e134e9 100644 --- a/agents/src/llm/llm.ts +++ b/agents/src/llm/llm.ts @@ -98,7 +98,7 @@ export abstract class LLM extends (EventEmitter as new () => TypedEmitter { @@ -448,8 +455,20 @@ describe('tool() name requirement', () => { }); expect(t.name).toBe('doStuff'); }); + + it('exposes id mirroring the function tool name', () => { + const t = tool({ + name: 'doStuff', + description: 'd', + execute: async () => 'x', + }); + expect(t.id).toBe('doStuff'); + expect(t.id).toBe(t.name); + }); }); +class TestProviderTool extends ProviderTool {} + describe('ToolContext', () => { const makeFn = (name: string) => tool({ @@ -497,7 +516,7 @@ describe('ToolContext', () => { it('separates provider tools from function tools', () => { const fnA = makeFn('a'); - const provider = tool({ id: 'code', config: { language: 'python' } }); + const provider = new TestProviderTool({ id: 'code' }); const ctx = new ToolContext([fnA, provider]); expect(ctx.functionTools).toEqual({ a: fnA }); @@ -537,7 +556,7 @@ describe('ToolContext', () => { it('equals() is reflexive', () => { const a = makeFn('a'); - const provider = tool({ id: 'code', config: { language: 'python' } }); + const provider = new TestProviderTool({ id: 'code' }); const ctx = new ToolContext([a, provider]); expect(ctx.equals(ctx)).toBe(true); }); @@ -547,22 +566,22 @@ describe('ToolContext', () => { // that hold the same provider-tool identities in different order are still equal so // realtime-session / preemptive-generation reuse fast paths are not invalidated. const a = makeFn('a'); - const p1 = tool({ id: 'code', config: { language: 'python' } }); - const p2 = tool({ id: 'browser', config: {} }); + const p1 = new TestProviderTool({ id: 'code' }); + const p2 = new TestProviderTool({ id: 'browser' }); expect(new ToolContext([a, p1, p2]).equals(new ToolContext([a, p2, p1]))).toBe(true); }); it('equals() supports contexts with only provider tools', () => { - const p1 = tool({ id: 'code', config: {} }); - const p2 = tool({ id: 'browser', config: {} }); + const p1 = new TestProviderTool({ id: 'code' }); + const p2 = new TestProviderTool({ id: 'browser' }); expect(new ToolContext([p1, p2]).equals(new ToolContext([p1, p2]))).toBe(true); - const p3 = tool({ id: 'code', config: {} }); // distinct identity, same id + const p3 = new TestProviderTool({ id: 'code' }); // distinct identity, same id expect(new ToolContext([p1]).equals(new ToolContext([p3]))).toBe(false); }); it('hasTool() matches function tools by name and provider tools by id', () => { const a = makeFn('a'); - const provider = tool({ id: 'code_runner', config: {} }); + const provider = new TestProviderTool({ id: 'code_runner' }); const ctx = new ToolContext([a, provider]); expect(ctx.hasTool('a')).toBe(true); @@ -574,9 +593,133 @@ describe('ToolContext', () => { // Matches Python's `flatten()`: list(self._fnc_tools_map.values()) + self._provider_tools. const a = makeFn('a'); const b = makeFn('b'); - const provider = tool({ id: 'code', config: {} }); + const provider = new TestProviderTool({ id: 'code' }); const ctx = new ToolContext([b, provider, a]); expect(ctx.flatten()).toEqual([b, a, provider]); }); }); + +describe('Toolset', () => { + const makeFn = (name: string) => + tool({ + name, + description: `${name} tool`, + execute: async () => name, + }); + + it('exposes its id and the tools it was constructed with', () => { + const a = makeFn('a'); + const b = makeFn('b'); + const ts = new Toolset({ id: 'set1', tools: [a, b] }); + + expect(ts.id).toBe('set1'); + expect(ts.tools).toEqual([a, b]); + }); + + it('default setup and aclose are no-ops', async () => { + const ts = new Toolset({ id: 'noop', tools: [] }); + await expect(ts.setup()).resolves.toBeUndefined(); + await expect(ts.aclose()).resolves.toBeUndefined(); + }); + + it('lets subclasses override lifecycle hooks', async () => { + const events: string[] = []; + class Recording extends Toolset { + override async setup(): Promise { + events.push(`setup:${this.id}`); + } + override async aclose(): Promise { + events.push(`close:${this.id}`); + } + } + + const ts = new Recording({ id: 'rec', tools: [] }); + await ts.setup(); + await ts.aclose(); + expect(events).toEqual(['setup:rec', 'close:rec']); + }); + + it('Toolset.create() composes lifecycle callbacks without subclassing', async () => { + const a = makeFn('a'); + const events: string[] = []; + const ts = Toolset.create({ + id: 'composed', + tools: [a], + setup: async () => { + events.push('setup'); + }, + aclose: async () => { + events.push('close'); + }, + }); + + expect(ts).toBeInstanceOf(Toolset); + expect(ts.id).toBe('composed'); + expect(ts.tools).toEqual([a]); + + await ts.setup(); + await ts.aclose(); + expect(events).toEqual(['setup', 'close']); + }); + + it('Toolset.create() defaults setup and aclose to no-ops when callbacks are omitted', async () => { + const ts = Toolset.create({ id: 'bare', tools: [] }); + await expect(ts.setup()).resolves.toBeUndefined(); + await expect(ts.aclose()).resolves.toBeUndefined(); + }); + + it('Toolset.create() accepts a tools thunk, re-evaluated on every access (dynamic)', () => { + const a = makeFn('a'); + const b = makeFn('b'); + const current: Tool[] = [a]; + let calls = 0; + const ts = Toolset.create({ + id: 'dynamic', + tools: () => { + calls += 1; + return current; + }, + }); + // Each access re-invokes the thunk so the toolset reflects the current source-of-truth. + expect(ts.tools).toEqual([a]); + expect(calls).toBe(1); + current.push(b); + expect(ts.tools).toEqual([a, b]); + expect(calls).toBe(2); + }); + + it('is flattened into a ToolContext: function tools merged, toolset tracked', () => { + const a = makeFn('a'); + const b = makeFn('b'); + const ts = new Toolset({ id: 'set', tools: [a, b] }); + const direct = makeFn('direct'); + + const ctx = new ToolContext([direct, ts]); + + expect(Object.keys(ctx.functionTools).sort()).toEqual(['a', 'b', 'direct']); + expect(ctx.toolsets).toEqual([ts]); + }); + + it('throws when a Toolset contributes a duplicate function name', () => { + // Mirrors Python's `add_tool`: a name collision between top-level and toolset-contributed + // tools is an error, not silent overwrite. + const a1 = makeFn('a'); + const a2 = makeFn('a'); + const ts = new Toolset({ id: 'collides', tools: [a2] }); + + expect(() => new ToolContext([a1, ts])).toThrow(/duplicate function name: a/); + }); + + it('equals() compares toolsets as identity sets, not by order', () => { + // Matches Python's `{id(ts) for ts in self._tool_sets}` semantics. + const ts1 = new Toolset({ id: 'one', tools: [] }); + const ts2 = new Toolset({ id: 'two', tools: [] }); + + expect(new ToolContext([ts1, ts2]).equals(new ToolContext([ts2, ts1]))).toBe(true); + + const ts3 = new Toolset({ id: 'three', tools: [] }); + expect(new ToolContext([ts1, ts2]).equals(new ToolContext([ts1, ts3]))).toBe(false); + expect(new ToolContext([ts1]).equals(new ToolContext([ts1, ts2]))).toBe(false); + }); +}); diff --git a/agents/src/llm/tool_context.ts b/agents/src/llm/tool_context.ts index df714d57a..284d5d7cf 100644 --- a/agents/src/llm/tool_context.ts +++ b/agents/src/llm/tool_context.ts @@ -12,7 +12,8 @@ import { isZodObjectSchema, isZodSchema } from './zod-utils.js'; const TOOL_SYMBOL = Symbol('tool'); const FUNCTION_TOOL_SYMBOL = Symbol('function_tool'); -const PROVIDER_DEFINED_TOOL_SYMBOL = Symbol('provider_defined_tool'); +const PROVIDER_TOOL_SYMBOL = Symbol('provider_tool'); +const TOOLSET_SYMBOL = Symbol('toolset'); const TOOL_ERROR_SYMBOL = Symbol('tool_error'); const HANDOFF_SYMBOL = Symbol('handoff'); @@ -57,7 +58,7 @@ export type InferToolInput = T extends { _output: infer O } ? O : any; // eslint-disable-line @typescript-eslint/no-explicit-any -- Fallback type for JSON Schema objects without type inference -export type ToolType = 'function' | 'provider-defined'; +export type ToolType = 'function' | 'provider'; export type ToolChoice = | 'auto' @@ -136,28 +137,32 @@ export type ToolExecuteFunction< export interface Tool { /** * The type of the tool. - * @internal Either user-defined core tool or provider-defined tool. + * @internal Either user-defined function tool or provider-side tool. */ type: ToolType; + /** + * Stable identifier used to key the tool inside a `ToolContext`. For function tools this + * mirrors `name`; for provider tools this is the provider tool id. + */ + id: string; + [TOOL_SYMBOL]: true; } -// TODO(AJS-112): support provider-defined tools -export interface ProviderDefinedTool extends Tool { - type: 'provider-defined'; +// TODO(AJS-112): support provider tools +export abstract class ProviderTool implements Tool { + readonly type = 'provider' as const; - /** - * The ID of the tool. - */ - id: string; + readonly id: string; - /** - * The configuration of the tool. - */ - config: Record; + readonly [TOOL_SYMBOL] = true as const; - [PROVIDER_DEFINED_TOOL_SYMBOL]: true; + readonly [PROVIDER_TOOL_SYMBOL] = true as const; + + constructor({ id }: { id: string }) { + this.id = id; + } } export interface FunctionTool< @@ -169,7 +174,7 @@ export interface FunctionTool< /** * The name of the tool. Used to identify it inside a `ToolContext` and exposed to the LLM - * as the function name to call. + * as the function name to call. Also surfaced as the inherited `Tool.id`. */ name: string; @@ -196,6 +201,125 @@ export interface FunctionTool< [FUNCTION_TOOL_SYMBOL]: true; } +export interface ToolCalledEvent { + ctx: RunContext; + arguments: Record; +} + +export interface ToolCompletedEvent { + ctx: RunContext; + output?: { type: 'output'; value: unknown } | { type: 'error'; value: Error }; +} + +/** + * A stateful collection of tools sharing a lifecycle. Tools registered through a `Toolset` are + * flattened into the surrounding `ToolContext`, while the `Toolset` itself is tracked so its + * `setup()` / `aclose()` hooks can be invoked by the agent runtime. + */ +export class Toolset { + readonly #id: string; + + readonly #tools: Tool[]; + + readonly [TOOLSET_SYMBOL] = true as const; + + constructor({ id, tools }: { id: string; tools: readonly Tool[] }) { + this.#id = id; + this.#tools = [...tools]; + } + + /** + * Compose a `Toolset` with inline `setup` / `aclose` hooks instead of subclassing. `tools` + * may also be a thunk that is re-evaluated on every `.tools` access, so the toolset can + * expose a dynamic list that changes after `setup()` runs. + * + * @example Static tool list with a shared backing resource + * ```ts + * function createPostgresToolset(connectionUrl: string): Toolset { + * const pool = new pg.Pool({ connectionString: connectionUrl }); + * return Toolset.create({ + * id: 'postgres', + * tools: [queryOrders, queryCustomers], + * setup: () => pool.connect(), + * aclose: () => pool.end(), + * }); + * } + * ``` + * + * @example Dynamic tool list + * ```ts + * function createMcpToolset(url: string): Toolset { + * const client = new MCPClient({ url }); + * return Toolset.create({ + * id: 'mcp_remote', + * tools: () => client.getTools(), + * setup: () => client.connect(), + * aclose: () => client.disconnect(), + * }); + * } + * ``` + */ + static create(options: ToolsetCreateOptions): Toolset { + return new ToolsetFactory(options); + } + + get id(): string { + return this.#id; + } + + get tools(): readonly Tool[] { + return this.#tools; + } + + async setup(): Promise {} + + async aclose(): Promise {} +} + +/** Options accepted by `Toolset.create()` — id + tools plus optional lifecycle hooks. */ +export interface ToolsetCreateOptions { + id: string; + /** + * Either a static list of tools, or a thunk re-evaluated on every `tools` access — useful + * when the underlying source (e.g. an MCP discovery loop) can produce a dynamic tool list. + */ + tools: readonly Tool[] | (() => readonly Tool[]); + /** Invoked when the toolset becomes active in an `AgentActivity`. */ + setup?: () => Promise; + /** Invoked when the toolset is being torn down. */ + aclose?: () => Promise; +} + +/** Backing implementation of `Toolset.create()`. Kept private so callers go through the factory. */ +class ToolsetFactory extends Toolset { + readonly #toolsSource: readonly Tool[] | (() => readonly Tool[]); + + readonly #setupFn?: () => Promise; + + readonly #acloseFn?: () => Promise; + + constructor({ id, tools, setup, aclose }: ToolsetCreateOptions) { + // Pass [] to super and override the `tools` getter so a thunk can be re-evaluated on + // every access (lets callers expose a dynamic tool list). + super({ id, tools: [] }); + this.#toolsSource = tools; + this.#setupFn = setup; + this.#acloseFn = aclose; + } + + override get tools(): readonly Tool[] { + return typeof this.#toolsSource === 'function' ? this.#toolsSource() : this.#toolsSource; + } + + override async setup(): Promise { + if (this.#setupFn) await this.#setupFn(); + } + + override async aclose(): Promise { + if (this.#acloseFn) await this.#acloseFn(); + } +} + /** * Convenience input shape accepted by APIs that want to take a list of tools directly without * forcing callers to wrap them in `new ToolContext(...)`. @@ -217,24 +341,18 @@ export function toToolContext( return input instanceof ToolContext ? input : new ToolContext(input); } -//TODO: toolset - accept stateful `Toolset` containers alongside `FunctionTool` / // eslint-disable-next-line @typescript-eslint/no-explicit-any -- ToolContext entries accept any function-tool parameter/result types export type ToolContextEntry = // eslint-disable-next-line @typescript-eslint/no-explicit-any - FunctionTool | ProviderDefinedTool; + FunctionTool | ProviderTool | Toolset; export class ToolContext { - // TODO: toolset - widen entries to `FunctionTool | ProviderDefinedTool | Toolset` once Toolset - // lands so this stays heterogeneous like Python's `Sequence[Tool | Toolset]`. private _tools: ToolContextEntry[] = []; // eslint-disable-next-line @typescript-eslint/no-explicit-any -- ToolContext stores generic function tools private _functionToolsMap: Map> = new Map(); - private _providerTools: ProviderDefinedTool[] = []; - // TODO: toolset - populate when Toolset support is supported. - // so the `toolsets` getter and `equals` toolset-identity check stay byte-compatible with the - private _toolSets: unknown[] = []; + private _providerTools: ProviderTool[] = []; + private _toolsets: Toolset[] = []; - // TODO: toolset - widen `tools` to `Sequence` once Toolset lands. constructor(tools: readonly ToolContextEntry[] = []) { this.updateTools(tools); } @@ -250,17 +368,13 @@ export class ToolContext { } /** A copy of all provider tools in the tool context, including those in tool sets. */ - get providerTools(): ProviderDefinedTool[] { + get providerTools(): ProviderTool[] { return this._providerTools; } - /** - * A copy of all tool sets in the tool context. - * - * TODO: toolset - wire up once Toolset is ported. - */ - get toolsets(): unknown[] { - return this._toolSets; + /** A copy of all toolsets registered in the context. */ + get toolsets(): readonly Toolset[] { + return [...this._toolsets]; } /** @@ -276,53 +390,53 @@ export class ToolContext { } // eslint-disable-next-line @typescript-eslint/no-explicit-any -- Generic registry over any parameter/result types - getFunctionTool(name: string): FunctionTool | undefined { - return this._functionToolsMap.get(name); + getFunctionTool(id: string): FunctionTool | undefined { + return this._functionToolsMap.get(id); } - hasTool(name: string): boolean { - if (this._functionToolsMap.has(name)) { + hasTool(id: string): boolean { + if (this._functionToolsMap.has(id)) { return true; } - return this._providerTools.some((tool) => tool.id === name); + return this._providerTools.some((tool) => tool.id === id); } - // TODO: toolset - widen `tools` to `Sequence` once Toolset lands. updateTools(tools: readonly ToolContextEntry[]): void { this._tools = [...tools]; this._functionToolsMap = new Map(); this._providerTools = []; - this._toolSets = []; + this._toolsets = []; - // Mirrors Python's recursive `add_tool` (minus Toolset flattening, which is TODO). // eslint-disable-next-line @typescript-eslint/no-explicit-any -- accepts any tool shape const addTool = (tool: any): void => { - if (isProviderDefinedTool(tool)) { + if (isToolset(tool)) { + for (const inner of tool.tools) { + addTool(inner); + } + this._toolsets.push(tool); + return; + } + + if (isProviderTool(tool)) { this._providerTools.push(tool); return; } if (isFunctionTool(tool)) { - const existing = this._functionToolsMap.get(tool.name); + const existing = this._functionToolsMap.get(tool.id); if (existing !== undefined) { if (existing !== tool) { - throw new Error(`duplicate function name: ${tool.name}`); + throw new Error(`duplicate function name: ${tool.id}`); } return; // same instance, skip } - this._functionToolsMap.set(tool.name, tool); + this._functionToolsMap.set(tool.id, tool); return; } - // TODO: toolset - if (tool instanceof Toolset) { for (const t of tool.tools) addTool(t); - // this._toolSets.push(tool); return; } - throw new Error(`unknown tool type: ${typeof tool}`); }; - // TODO: toolset - Python also chains `find_function_tools(self)` here so subclasses can - // declare tools as class members. JS doesn't use that decorator pattern, so we only walk - // the explicit input list. for (const tool of tools) { addTool(tool); } @@ -336,14 +450,17 @@ export class ToolContext { if (this._functionToolsMap.size !== other._functionToolsMap.size) { return false; } - for (const [name, tool] of this._functionToolsMap) { - if (other._functionToolsMap.get(name) !== tool) { + + for (const [id, tool] of this._functionToolsMap) { + if (other._functionToolsMap.get(id) !== tool) { return false; } } + if (this._providerTools.length !== other._providerTools.length) { return false; } + // Provider tools compare as identity sets to match Python's `set(id(t) for t in ...)` // semantics — order is not significant. const otherProviderIds = new Set(other._providerTools); @@ -352,10 +469,17 @@ export class ToolContext { return false; } } - // TODO: toolset - once Toolset lands, also compare `_toolSets` as identity sets per Python - // self_tool_set_ids = {id(ts) for ts in self._tool_sets} - // other_tool_set_ids = {id(ts) for ts in other._tool_sets} - // if self_tool_set_ids != other_tool_set_ids: return False + + if (this._toolsets.length !== other._toolsets.length) { + return false; + } + + const otherToolsets = new Set(other._toolsets); + for (const ts of this._toolsets) { + if (!otherToolsets.has(ts)) { + return false; + } + } return true; } } @@ -413,63 +537,36 @@ export function tool({ flags?: number; }): FunctionTool, UserData, Result>; -/** - * Create a provider-defined tool. - * - * @param id - The ID of the tool. - * @param config - The configuration of the tool. - */ -export function tool({ - id, - config, -}: { - id: string; - config: Record; -}): ProviderDefinedTool; - // eslint-disable-next-line @typescript-eslint/no-explicit-any export function tool(tool: any): any { - if (tool.execute !== undefined) { - if (typeof tool.name !== 'string' || tool.name.length === 0) { - throw new Error('tool({ name, ... }) requires a non-empty name'); - } - - // Default parameters to z.object({}) if not provided - const parameters = tool.parameters ?? z.object({}); - - // if parameters is a Zod schema, ensure it's an object schema - if (isZodSchema(parameters) && !isZodObjectSchema(parameters)) { - throw new Error('Tool parameters must be a Zod object schema (z.object(...))'); - } + if (typeof tool.name !== 'string' || tool.name.length === 0) { + throw new Error('tool({ name, ... }) requires a non-empty name'); + } - // Ensure parameters is either a Zod schema or a plain object (JSON schema) - if (!isZodSchema(parameters) && !(typeof parameters === 'object')) { - throw new Error('Tool parameters must be a Zod object schema or a raw JSON schema'); - } + // Default parameters to z.object({}) if not provided + const parameters = tool.parameters ?? z.object({}); - return { - type: 'function', - name: tool.name, - description: tool.description, - parameters, - execute: tool.execute, - flags: tool.flags ?? ToolFlag.NONE, - [TOOL_SYMBOL]: true, - [FUNCTION_TOOL_SYMBOL]: true, - }; + // if parameters is a Zod schema, ensure it's an object schema + if (isZodSchema(parameters) && !isZodObjectSchema(parameters)) { + throw new Error('Tool parameters must be a Zod object schema (z.object(...))'); } - if (tool.config !== undefined && tool.id !== undefined) { - return { - type: 'provider-defined', - id: tool.id, - config: tool.config, - [TOOL_SYMBOL]: true, - [PROVIDER_DEFINED_TOOL_SYMBOL]: true, - }; + // Ensure parameters is either a Zod schema or a plain object (JSON schema) + if (!isZodSchema(parameters) && !(typeof parameters === 'object')) { + throw new Error('Tool parameters must be a Zod object schema or a raw JSON schema'); } - throw new Error('Invalid tool'); + return { + type: 'function', + id: tool.name, + name: tool.name, + description: tool.description, + parameters, + execute: tool.execute, + flags: tool.flags ?? ToolFlag.NONE, + [TOOL_SYMBOL]: true, + [FUNCTION_TOOL_SYMBOL]: true, + }; } // eslint-disable-next-line @typescript-eslint/no-explicit-any @@ -485,10 +582,15 @@ export function isFunctionTool(tool: any): tool is FunctionTool { } // eslint-disable-next-line @typescript-eslint/no-explicit-any -export function isProviderDefinedTool(tool: any): tool is ProviderDefinedTool { +export function isProviderTool(tool: any): tool is ProviderTool { const isTool = tool && tool[TOOL_SYMBOL] === true; - const isProviderDefinedTool = tool[PROVIDER_DEFINED_TOOL_SYMBOL] === true; - return isTool && isProviderDefinedTool; + const isProviderTool = tool[PROVIDER_TOOL_SYMBOL] === true; + return isTool && isProviderTool; +} + +// eslint-disable-next-line @typescript-eslint/no-explicit-any +export function isToolset(value: any): value is Toolset { + return value && value[TOOLSET_SYMBOL] === true; } // eslint-disable-next-line @typescript-eslint/no-explicit-any diff --git a/agents/src/llm/tool_context.type.test.ts b/agents/src/llm/tool_context.type.test.ts index 187f95e7b..5d33124ad 100644 --- a/agents/src/llm/tool_context.type.test.ts +++ b/agents/src/llm/tool_context.type.test.ts @@ -3,7 +3,7 @@ // SPDX-License-Identifier: Apache-2.0 import { describe, expect, expectTypeOf, it } from 'vitest'; import { z } from 'zod'; -import { type FunctionTool, type ProviderDefinedTool, type ToolOptions, tool } from './index.js'; +import { type FunctionTool, ProviderTool, type ToolOptions, tool } from './index.js'; describe('tool type inference', () => { it('should infer argument type from zod schema', () => { @@ -17,15 +17,15 @@ describe('tool type inference', () => { expectTypeOf(toolType).toEqualTypeOf>(); }); - it('should infer provider defined tool type', () => { - const toolType = tool({ - id: 'code-interpreter', - config: { - language: 'python', - }, - }); + it('rejects direct instantiation of the abstract ProviderTool base', () => { + // @ts-expect-error - ProviderTool is abstract; plugins must subclass it. + new ProviderTool({ id: 'code-interpreter' }); - expectTypeOf(toolType).toEqualTypeOf(); + class CodeInterpreter extends ProviderTool {} + const providerTool = new CodeInterpreter({ id: 'code-interpreter' }); + expectTypeOf(providerTool).toMatchTypeOf(); + expect(providerTool.id).toBe('code-interpreter'); + expect(providerTool.type).toBe('provider'); }); it('should infer run context type', () => { @@ -45,7 +45,6 @@ describe('tool type inference', () => { it('should not accept primitive zod schemas', () => { expect(() => { - // @ts-expect-error - Testing that non-object schemas are rejected tool({ name: 'test', description: 'test', @@ -57,7 +56,6 @@ describe('tool type inference', () => { it('should not accept array schemas', () => { expect(() => { - // @ts-expect-error - Testing that array schemas are rejected tool({ name: 'test', description: 'test', @@ -69,7 +67,6 @@ describe('tool type inference', () => { it('should not accept union schemas', () => { expect(() => { - // @ts-expect-error - Testing that union schemas are rejected tool({ name: 'test', description: 'test', diff --git a/agents/src/utils.test.ts b/agents/src/utils.test.ts index df2fc91de..ddfe230cd 100644 --- a/agents/src/utils.test.ts +++ b/agents/src/utils.test.ts @@ -14,6 +14,7 @@ import { delay, isPending, resampleStream, + toStream, } from '../src/utils.js'; describe('utils', () => { @@ -32,6 +33,84 @@ describe('utils', () => { }); }); + describe('toStream', () => { + it('converts an async iterable into a ReadableStream', async () => { + async function* source() { + yield 1; + yield 2; + yield 3; + } + + const reader = toStream(source()).getReader(); + + await expect(reader.read()).resolves.toEqual({ done: false, value: 1 }); + await expect(reader.read()).resolves.toEqual({ done: false, value: 2 }); + await expect(reader.read()).resolves.toEqual({ done: false, value: 3 }); + await expect(reader.read()).resolves.toEqual({ done: true, value: undefined }); + }); + + it('propagates errors from the async iterable', async () => { + const expectedError = new Error('source failed'); + async function* source() { + yield 1; + throw expectedError; + } + + const reader = toStream(source()).getReader(); + + await expect(reader.read()).resolves.toEqual({ done: false, value: 1 }); + await expect(reader.read()).rejects.toBe(expectedError); + }); + + it('runs async iterable cleanup when the stream is canceled mid-stream', async () => { + let cleanupRan = false; + async function* source() { + try { + yield 1; + yield 2; + } finally { + cleanupRan = true; + } + } + + const reader = toStream(source()).getReader(); + + await expect(reader.read()).resolves.toEqual({ done: false, value: 1 }); + await reader.cancel('stop early'); + + expect(cleanupRan).toBe(true); + }); + + it('does not wait for a pending next value when canceled mid-stream', async () => { + let releaseNextValue: (() => void) | undefined; + let cleanupRan = false; + async function* source() { + try { + yield 1; + await new Promise((resolve) => { + releaseNextValue = resolve; + }); + yield 2; + } finally { + cleanupRan = true; + } + } + + const reader = toStream(source()).getReader(); + + await expect(reader.read()).resolves.toEqual({ done: false, value: 1 }); + const pendingRead = reader.read(); + await delay(1); + await expect( + Promise.race([reader.cancel('stop early'), delay(10).then(() => 'timeout')]), + ).resolves.not.toBe('timeout'); + releaseNextValue?.(); + await expect(pendingRead).resolves.toEqual({ done: true, value: undefined }); + + expect(cleanupRan).toBe(true); + }); + }); + describe('Task', () => { it('should execute task successfully and return result', async () => { const expectedResult = 'task completed'; diff --git a/agents/src/utils.ts b/agents/src/utils.ts index 6343c932f..2b3d651b2 100644 --- a/agents/src/utils.ts +++ b/agents/src/utils.ts @@ -13,8 +13,11 @@ import { type Throws, ThrowsPromise } from '@livekit/throws-transformer/throws'; import { AsyncLocalStorage } from 'node:async_hooks'; import { randomUUID } from 'node:crypto'; import { EventEmitter, once } from 'node:events'; -import type { ReadableStream } from 'node:stream/web'; -import { TransformStream, type TransformStreamDefaultController } from 'node:stream/web'; +import { + ReadableStream, + TransformStream, + type TransformStreamDefaultController, +} from 'node:stream/web'; import { log } from './log.js'; /** @@ -1127,15 +1130,21 @@ export async function* readStream( const abortPromise = waitForAbort(signal); while (true) { const result = await ThrowsPromise.race([reader.read(), abortPromise]); - if (!result) break; + if (!result) { + break; + } const { done, value } = result; - if (done) break; + if (done) { + break; + } yield value; } } else { while (true) { const { done, value } = await reader.read(); - if (done) break; + if (done) { + break; + } yield value; } } @@ -1148,6 +1157,39 @@ export async function* readStream( } } +export function toStream(iterable: AsyncIterable): ReadableStream { + let iterator: AsyncIterator | undefined; + let cancelled = false; + + return new ReadableStream({ + async start(controller) { + iterator = iterable[Symbol.asyncIterator](); + + try { + while (true) { + const { done, value } = await iterator.next(); + if (done || cancelled) { + break; + } + controller.enqueue(value); + } + + if (!cancelled) { + controller.close(); + } + } catch (error) { + if (!cancelled) { + controller.error(error); + } + } + }, + cancel(reason) { + cancelled = true; + void iterator?.return?.(reason).catch(() => {}); + }, + }); +} + export async function waitForAbort(signal: AbortSignal) { if (signal.aborted) { return; diff --git a/agents/src/voice/agent.test.ts b/agents/src/voice/agent.test.ts index f2afa7c42..0bfbf7df3 100644 --- a/agents/src/voice/agent.test.ts +++ b/agents/src/voice/agent.test.ts @@ -1,9 +1,11 @@ // SPDX-FileCopyrightText: 2025 LiveKit, Inc. // // SPDX-License-Identifier: Apache-2.0 +import type { AudioFrame } from '@livekit/rtc-node'; +import { ReadableStream } from 'node:stream/web'; import { describe, expect, it, vi } from 'vitest'; import { z } from 'zod'; -import { tool } from '../llm/index.js'; +import { ChatContext, ChatMessage, tool } from '../llm/index.js'; import { initializeLogger } from '../log.js'; import { Task } from '../utils.js'; import { Agent, AgentTask, _setActivityTaskInfo } from './agent.js'; @@ -15,6 +17,14 @@ vi.mock('ofetch', () => ({ ofetch: vi.fn() })); initializeLogger({ pretty: false, level: 'error' }); +async function collectReadableStream(stream: ReadableStream): Promise { + const chunks: T[] = []; + for await (const chunk of stream) { + chunks.push(chunk); + } + return chunks; +} + describe('Agent', () => { it('should create agent with basic instructions', () => { const instructions = 'You are a helpful assistant'; @@ -80,6 +90,258 @@ describe('Agent', () => { expect(agent.toolCtx.getFunctionTool('testTool')).toBe(mockTool); }); + describe('create', () => { + it('preserves constructor options and base Agent default id', () => { + const mockTool = tool({ + name: 'testTool', + description: 'Test tool', + parameters: z.object({}), + execute: async () => 'result', + }); + + const agent = Agent.create({ + instructions: 'factory instructions', + tools: [mockTool], + }); + + expect(agent).toBeInstanceOf(Agent); + expect(agent.instructions).toBe('factory instructions'); + expect(agent.id).toBe('default_agent'); + expect(agent.toolCtx.getFunctionTool('testTool')).toBe(mockTool); + }); + + it('passes AgentContext to lifecycle hooks', async () => { + const calls: string[] = []; + const chatCtx = ChatContext.empty(); + const newMessage = ChatMessage.create({ role: 'user', content: ['hello'] }); + const agent = Agent.create({ + id: 'factory_agent', + instructions: 'factory instructions', + minConsecutiveSpeechDelay: 12, + ttsPronunciationMap: { LiveKit: 'live kit' }, + onEnter: (ctx) => { + expect(ctx.agent).toBe(agent); + expect(ctx.id).toBe(agent.id); + expect(ctx.instructions).toBe(agent.instructions); + expect(ctx.toolCtx.functionTools).toEqual(agent.toolCtx.functionTools); + expect(ctx.chatCtx.items).toEqual(agent.chatCtx.items); + expect(ctx.minConsecutiveSpeechDelay).toBe(agent.minConsecutiveSpeechDelay); + expect(ctx.ttsPronunciationMap).toBe(agent.ttsPronunciationMap); + calls.push('enter'); + }, + onExit: async (ctx) => { + expect(ctx.agent).toBe(agent); + calls.push('exit'); + }, + onUserTurnCompleted: (ctx, receivedChatCtx, receivedMessage) => { + expect(ctx.agent).toBe(agent); + expect(receivedChatCtx).toBe(chatCtx); + expect(receivedMessage).toBe(newMessage); + calls.push('turn'); + }, + }); + + await agent.onEnter(); + await agent.onExit(); + await agent.onUserTurnCompleted(chatCtx, newMessage); + + expect(calls).toEqual(['enter', 'exit', 'turn']); + }); + + it('adapts stream node hooks between ReadableStream and AsyncIterable', async () => { + const audioFrame = 'audio' as unknown as AudioFrame; + const agent = Agent.create({ + instructions: 'factory instructions', + async sttNode(ctx, audio) { + async function* stream() { + expect(ctx.agent).toBe(agent); + const frames: AudioFrame[] = []; + for await (const frame of audio) { + frames.push(frame); + } + expect(frames).toEqual([audioFrame]); + yield 'transcript'; + } + + return stream(); + }, + }); + const audio = new ReadableStream({ + start(controller) { + controller.enqueue(audioFrame); + controller.close(); + }, + }); + + const result = await agent.sttNode(audio, {}); + + expect(result).not.toBeNull(); + await expect(collectReadableStream(result!)).resolves.toEqual(['transcript']); + }); + + it('supports async generator stream node hooks', async () => { + const audioFrame = 'audio' as unknown as AudioFrame; + const outputFrame = 'output-audio' as unknown as AudioFrame; + const agent = Agent.create({ + instructions: 'factory instructions', + async *sttNode(ctx, audio) { + expect(ctx.agent).toBe(agent); + const frames: AudioFrame[] = []; + for await (const frame of audio) { + frames.push(frame); + } + expect(frames).toEqual([audioFrame]); + yield 'transcript'; + }, + async *llmNode(ctx, chatCtx, toolCtx) { + expect(ctx.agent).toBe(agent); + expect(chatCtx).toBeInstanceOf(ChatContext); + expect(toolCtx.equals(agent.toolCtx)).toBe(true); + yield 'llm-output'; + }, + async *ttsNode(ctx, text) { + expect(ctx.agent).toBe(agent); + const chunks: string[] = []; + for await (const chunk of text) { + chunks.push(chunk); + } + expect(chunks).toEqual(['hello']); + yield outputFrame; + }, + async *realtimeAudioOutputNode(ctx, audio) { + expect(ctx.agent).toBe(agent); + const frames: AudioFrame[] = []; + for await (const frame of audio) { + frames.push(frame); + } + expect(frames).toEqual([audioFrame]); + yield outputFrame; + }, + }); + const audio = new ReadableStream({ + start(controller) { + controller.enqueue(audioFrame); + controller.close(); + }, + }); + const text = new ReadableStream({ + start(controller) { + controller.enqueue('hello'); + controller.close(); + }, + }); + + const sttResult = await agent.sttNode(audio, {}); + const llmResult = await agent.llmNode(ChatContext.empty(), agent.toolCtx, {}); + const ttsResult = await agent.ttsNode(text, {}); + const realtimeResult = await agent.realtimeAudioOutputNode( + new ReadableStream({ + start(controller) { + controller.enqueue(audioFrame); + controller.close(); + }, + }), + {}, + ); + + expect(sttResult).not.toBeNull(); + expect(llmResult).not.toBeNull(); + expect(ttsResult).not.toBeNull(); + expect(realtimeResult).not.toBeNull(); + await expect(collectReadableStream(sttResult!)).resolves.toEqual(['transcript']); + await expect(collectReadableStream(llmResult!)).resolves.toEqual(['llm-output']); + await expect(collectReadableStream(ttsResult!)).resolves.toEqual([outputFrame]); + await expect(collectReadableStream(realtimeResult!)).resolves.toEqual([outputFrame]); + }); + + it('supports stream node hooks that return async iterables', async () => { + function asyncIterableOf(...items: T[]): AsyncIterable { + return { + async *[Symbol.asyncIterator]() { + for (const item of items) { + yield item; + } + }, + }; + } + + const audioFrame = 'audio' as unknown as AudioFrame; + const outputFrame = 'output-audio' as unknown as AudioFrame; + const agent = Agent.create({ + instructions: 'factory instructions', + sttNode(ctx) { + expect(ctx.agent).toBe(agent); + return asyncIterableOf('transcript'); + }, + llmNode(ctx) { + expect(ctx.agent).toBe(agent); + return asyncIterableOf('llm-output'); + }, + ttsNode(ctx) { + expect(ctx.agent).toBe(agent); + return asyncIterableOf(outputFrame); + }, + realtimeAudioOutputNode(ctx) { + expect(ctx.agent).toBe(agent); + return asyncIterableOf(outputFrame); + }, + }); + + const sttResult = await agent.sttNode( + new ReadableStream({ + start(controller) { + controller.enqueue(audioFrame); + controller.close(); + }, + }), + {}, + ); + const llmResult = await agent.llmNode(ChatContext.empty(), agent.toolCtx, {}); + const ttsResult = await agent.ttsNode( + new ReadableStream({ + start(controller) { + controller.enqueue('hello'); + controller.close(); + }, + }), + {}, + ); + const realtimeResult = await agent.realtimeAudioOutputNode( + new ReadableStream({ + start(controller) { + controller.enqueue(audioFrame); + controller.close(); + }, + }), + {}, + ); + + expect(sttResult).not.toBeNull(); + expect(llmResult).not.toBeNull(); + expect(ttsResult).not.toBeNull(); + expect(realtimeResult).not.toBeNull(); + await expect(collectReadableStream(sttResult!)).resolves.toEqual(['transcript']); + await expect(collectReadableStream(llmResult!)).resolves.toEqual(['llm-output']); + await expect(collectReadableStream(ttsResult!)).resolves.toEqual([outputFrame]); + await expect(collectReadableStream(realtimeResult!)).resolves.toEqual([outputFrame]); + }); + + it('falls back to existing defaults for missing hooks', async () => { + const audioFrame = 'audio' as unknown as AudioFrame; + const audio = new ReadableStream({ + start(controller) { + controller.enqueue(audioFrame); + controller.close(); + }, + }); + const agent = Agent.create({ instructions: 'factory instructions' }); + + const result = await agent.realtimeAudioOutputNode(audio, {}); + + expect(result).toBe(audio); + }); + }); + it('should require AgentTask to run inside task context', async () => { class TestTask extends AgentTask { constructor() { @@ -147,6 +409,94 @@ describe('Agent', () => { await expect(wrapper.result).resolves.toBe('ok'); }); + describe('AgentTask.create', () => { + it('exposes complete on hook context', async () => { + const task = AgentTask.create({ + instructions: 'factory task', + onEnter: (ctx) => { + expect(ctx.agent).toBe(task); + expect(ctx.id).toBe('default_agent'); + expect(ctx.instructions).toBe('factory task'); + ctx.complete('ok'); + }, + }); + const oldAgent = new Agent({ instructions: 'old agent' }); + const mockSession = { + currentAgent: oldAgent, + _globalRunState: undefined, + _updateActivity: async (agent: Agent) => { + if (agent === task) { + await agent.onEnter(); + } + }, + }; + const mockActivity = { + agent: oldAgent, + agentSession: mockSession, + _onEnterTask: undefined, + llm: undefined, + close: async () => {}, + }; + + const wrapper = Task.from(async () => { + const currentTask = Task.current(); + if (!currentTask) { + throw new Error('expected task context'); + } + _setActivityTaskInfo(currentTask, { inlineTask: true }); + return await agentActivityStorage.run(mockActivity as any, () => task.run()); + }); + + await expect(wrapper.result).resolves.toBe('ok'); + }); + + it('adapts stream node hooks between ReadableStream and AsyncIterable', async () => { + const audioFrame = 'audio' as unknown as AudioFrame; + const task = AgentTask.create({ + instructions: 'factory task', + async sttNode(ctx, audio) { + async function* stream() { + expect(ctx.agent).toBe(task); + const frames: AudioFrame[] = []; + for await (const frame of audio) { + frames.push(frame); + } + expect(frames).toEqual([audioFrame]); + yield 'transcript'; + } + + return stream(); + }, + }); + const audio = new ReadableStream({ + start(controller) { + controller.enqueue(audioFrame); + controller.close(); + }, + }); + + const result = await task.sttNode(audio, {}); + + expect(result).not.toBeNull(); + await expect(collectReadableStream(result!)).resolves.toEqual(['transcript']); + }); + + it('falls back to existing defaults for missing hooks', async () => { + const audioFrame = 'audio' as unknown as AudioFrame; + const audio = new ReadableStream({ + start(controller) { + controller.enqueue(audioFrame); + controller.close(); + }, + }); + const task = AgentTask.create({ instructions: 'factory task' }); + + const result = await task.realtimeAudioOutputNode(audio, {}); + + expect(result).toBe(audio); + }); + }); + it('should require AgentTask to run inside AgentActivity context', async () => { class TestTask extends AgentTask { constructor() { @@ -235,6 +585,7 @@ describe('Agent', () => { turnHandling: { endpointing: { minDelay: 999 }, interruption: {}, + preemptiveGeneration: {}, turnDetection: 'vad', }, allowInterruptions: false, @@ -251,6 +602,7 @@ describe('Agent', () => { turnHandling: { endpointing: { minDelay: 999, maxDelay: 4000 }, interruption: { enabled: true }, + preemptiveGeneration: {}, turnDetection: 'vad', }, allowInterruptions: false, @@ -268,6 +620,7 @@ describe('Agent', () => { turnHandling: { interruption: { mode: 'adaptive' }, endpointing: {}, + preemptiveGeneration: {}, turnDetection: undefined, }, }); @@ -280,6 +633,7 @@ describe('Agent', () => { turnHandling: { endpointing: { minDelay: 111, maxDelay: 222 }, interruption: { enabled: false }, + preemptiveGeneration: {}, turnDetection: 'manual', }, }); diff --git a/agents/src/voice/agent.ts b/agents/src/voice/agent.ts index 3145de2b1..4072543ef 100644 --- a/agents/src/voice/agent.ts +++ b/agents/src/voice/agent.ts @@ -34,11 +34,26 @@ import { Future, Task } from '../utils.js'; import type { VAD } from '../vad.js'; import { type AgentActivity, agentActivityStorage } from './agent_activity.js'; import type { AgentSession, TurnDetectionMode } from './agent_session.js'; +import { + type AgentCreateOptions, + type AgentTaskCreateOptions, + createAgentTaskV2, + createAgentV2, +} from './agent_v2.js'; import type { TimedString } from './io.js'; import type { SpeechHandle } from './speech_handle.js'; import type { TurnHandlingOptions } from './turn_config/turn_handling.js'; import { migrateTurnHandling } from './turn_config/utils.js'; +export type { + AgentContext, + AgentCreateOptions, + AgentHookNodeResult, + AgentHooks, + AgentTaskContext, + AgentTaskCreateOptions, +} from './agent_v2.js'; + // speechHandle identifies which SpeechHandle owns the current tool call, enabling // SpeechHandle.waitForPlayout() to distinguish self-wait (deadlock) from waiting // on a different handle scheduled inside the tool. @@ -160,6 +175,10 @@ export class Agent { /** @internal */ _toolCtx: ToolContext; + static create(options: AgentCreateOptions): Agent { + return createAgentV2(Agent, options); + } + constructor({ id, instructions, @@ -660,6 +679,12 @@ export class AgentTask extends Agent( + options: AgentTaskCreateOptions, + ): AgentTask { + return createAgentTaskV2(AgentTask, options); + } + constructor(options: AgentTaskOptions) { const { preserveFunctionCallHistory = false, ...rest } = options; super(rest); diff --git a/agents/src/voice/agent_activity.ts b/agents/src/voice/agent_activity.ts index 3068b408d..793809a3e 100644 --- a/agents/src/voice/agent_activity.ts +++ b/agents/src/voice/agent_activity.ts @@ -35,10 +35,14 @@ import { RealtimeModel, type RealtimeModelError, type RealtimeSession, + type Tool, type ToolChoice, ToolContext, type ToolContextEntry, ToolFlag, + Toolset, + isFunctionTool, + isToolset, } from '../llm/index.js'; import type { LLMError } from '../llm/llm.js'; import { isSameToolChoice } from '../llm/tool_context.js'; @@ -215,6 +219,7 @@ export class AgentActivity implements RecognitionHooks { private toolChoice: ToolChoice | null = null; private _preemptiveGeneration?: PreemptiveGeneration; private _preemptiveGenerationCount = 0; + private _toolsetsSetup = false; private interruptionDetector?: AdaptiveInterruptionDetector; private isInterruptionDetectionEnabled: boolean; private isInterruptionByAudioActivityEnabled: boolean; @@ -421,6 +426,8 @@ export class AgentActivity implements RecognitionHooks { this.agent._agentActivity = this; + await this.setupToolsets(); + if (this.llm instanceof RealtimeModel) { const rtReused = reuseResources?.rtSession !== undefined; @@ -767,13 +774,20 @@ export class AgentActivity implements RecognitionHooks { } async updateTools(tools: readonly ToolContextEntry[]): Promise { - const oldToolNames = new Set(Object.keys(this.agent._toolCtx.functionTools)); + const oldToolCtx = this.agent._toolCtx; + const oldToolNames = new Set(Object.keys(oldToolCtx.functionTools)); + const oldToolsets = oldToolCtx.toolsets; const newToolCtx = new ToolContext(tools); const newToolNames = new Set(Object.keys(newToolCtx.functionTools)); + const newToolsets = newToolCtx.toolsets; const toolsAdded = [...newToolNames].filter((name) => !oldToolNames.has(name)); const toolsRemoved = [...oldToolNames].filter((name) => !newToolNames.has(name)); + const addedToolsets = newToolsets.filter((ts) => !oldToolsets.includes(ts)); + const removedToolsets = oldToolsets.filter((ts) => !newToolsets.includes(ts)); + await this.setupToolsetList(addedToolsets); this.agent._toolCtx = newToolCtx; + await this.closeToolsetList(removedToolsets); if (toolsAdded.length > 0 || toolsRemoved.length > 0) { const configUpdate = new AgentConfigUpdate({ @@ -1735,11 +1749,13 @@ export class AgentActivity implements RecognitionHooks { const tools: ToolContext = shouldFilterTools ? new ToolContext( - this.agent.toolCtx.tools.filter((t) => { - if (t.type === 'function') { - return !(t.flags & ToolFlag.IGNORE_ON_ENTER); + this.agent.toolCtx.tools.flatMap((t): ToolContextEntry[] => { + const keepFn = (fn: Tool): boolean => + !isFunctionTool(fn) || !(fn.flags & ToolFlag.IGNORE_ON_ENTER); + if (isToolset(t)) { + return t.tools.filter(keepFn) as ToolContextEntry[]; } - return true; + return keepFn(t) ? [t] : []; }), ) : this.agent.toolCtx; @@ -3728,9 +3744,42 @@ export class AgentActivity implements RecognitionHooks { this.realtimeSpans?.clear(); await this.realtimeSession?.close(); await this.audioRecognition?.close(); + await this.closeToolsets(); this.realtimeSession = undefined; this.audioRecognition = undefined; } + + private async setupToolsets(): Promise { + // Guard against resume() re-entering _startSession on an activity whose toolsets are + // already initialized. + if (this._toolsetsSetup) return; + this._toolsetsSetup = true; + await this.setupToolsetList(this.agent.toolCtx.toolsets); + } + + private async closeToolsets(): Promise { + if (!this._toolsetsSetup) return; + this._toolsetsSetup = false; + await this.closeToolsetList(this.agent.toolCtx.toolsets); + } + + private async setupToolsetList(toolsets: readonly Toolset[]): Promise { + const outputs = await Promise.allSettled(toolsets.map((ts) => ts.setup())); + for (const output of outputs) { + if (output.status === 'rejected') { + this.logger.error({ error: output.reason }, 'error setting up toolset'); + } + } + } + + private async closeToolsetList(toolsets: readonly Toolset[]): Promise { + const outputs = await Promise.allSettled(toolsets.map((ts) => ts.aclose())); + for (const output of outputs) { + if (output.status === 'rejected') { + this.logger.error({ error: output.reason }, 'error closing toolset'); + } + } + } } function toOaiToolChoice(toolChoice: ToolChoice | null): ToolChoice | undefined { diff --git a/agents/src/voice/agent_v2.ts b/agents/src/voice/agent_v2.ts new file mode 100644 index 000000000..6e368e2d3 --- /dev/null +++ b/agents/src/voice/agent_v2.ts @@ -0,0 +1,484 @@ +// SPDX-FileCopyrightText: 2026 LiveKit, Inc. +// +// SPDX-License-Identifier: Apache-2.0 +import type { AudioFrame } from '@livekit/rtc-node'; +import type { ReadableStream } from 'node:stream/web'; +import type { Instructions, ReadonlyChatContext } from '../llm/chat_context.js'; +import type { + ChatChunk, + ChatContext, + ChatMessage, + LLM, + RealtimeModel, + ToolContext, +} from '../llm/index.js'; +import type { STT, SpeechEvent } from '../stt/index.js'; +import type { TTS } from '../tts/index.js'; +import { readStream, toStream } from '../utils.js'; +import type { VAD } from '../vad.js'; +import type { + Agent, + AgentOptions, + AgentTask, + AgentTaskOptions, + ModelSettings, + TTSPronunciationMap, +} from './agent.js'; +import type { AgentSession } from './agent_session.js'; +import type { TurnHandlingOptions } from './turn_config/turn_handling.js'; + +/** Context passed to hooks created with `Agent.create()`. */ +export interface AgentContext { + /** The agent instance currently running the hook. */ + agent: Agent; + /** Voice activity detector configured for the agent. */ + vad: VAD | undefined; + /** Speech-to-text model configured for the agent. */ + stt: STT | undefined; + /** LLM or realtime model configured for the agent. */ + llm: LLM | RealtimeModel | undefined; + /** Text-to-speech model configured for the agent. */ + tts: TTS | undefined; + /** Whether TTS-aligned transcripts are enabled for the agent. */ + useTtsAlignedTranscript: boolean | undefined; + /** Pronunciation replacements applied before TTS synthesis. */ + ttsPronunciationMap: TTSPronunciationMap | undefined; + /** Readonly view of the agent's current chat context. */ + chatCtx: ReadonlyChatContext; + /** Agent identifier. */ + id: string; + /** Agent instructions. */ + instructions: string | Instructions; + /** Copy of the agent tool context. */ + toolCtx: ToolContext; + /** Current session for the agent. */ + session: AgentSession; + /** Agent-level turn handling configuration. */ + turnHandling: Partial | undefined; + /** Minimum delay between consecutive speech. */ + minConsecutiveSpeechDelay: number | undefined; +} + +/** Return type for stream hooks. Returning `null` stops that pipeline node. */ +export type AgentHookNodeResult = AsyncIterable | Promise | null> | null; + +export interface AgentHooks< + UserData, + ContextT extends AgentContext = AgentContext, +> { + /** Called when the agent becomes active in a session. */ + onEnter?: (ctx: ContextT) => Promise | void; + /** Called when the agent is leaving the active session. */ + onExit?: (ctx: ContextT) => Promise | void; + /** Called after the user's turn has been committed to the chat context. */ + onUserTurnCompleted?: ( + ctx: ContextT, + chatCtx: ChatContext, + newMessage: ChatMessage, + ) => Promise | void; + /** Transforms incoming audio into speech events or transcript text for the agent. */ + sttNode?: ( + ctx: ContextT, + audio: AsyncIterable, + modelSettings: ModelSettings, + ) => AgentHookNodeResult; + /** Produces LLM chunks or text from the current chat and tool context. */ + llmNode?: ( + ctx: ContextT, + chatCtx: ChatContext, + toolCtx: ToolContext, + modelSettings: ModelSettings, + ) => AgentHookNodeResult; + /** Synthesizes agent text into audio frames for playout. */ + ttsNode?: ( + ctx: ContextT, + text: AsyncIterable, + modelSettings: ModelSettings, + ) => AgentHookNodeResult; + /** Processes realtime model audio before it is sent to the agent output. */ + realtimeAudioOutputNode?: ( + ctx: ContextT, + audio: AsyncIterable, + modelSettings: ModelSettings, + ) => AgentHookNodeResult; +} + +export interface AgentCreateOptions + extends AgentOptions, + AgentHooks {} + +/** Context passed to hooks created with `AgentTask.create()`. */ +export interface AgentTaskContext + extends AgentContext { + /** The task instance currently running the hook. */ + agent: AgentTask; + /** Complete the task with either a result or an error. */ + complete(result: ResultT | Error): void; +} + +export interface AgentTaskCreateOptions + extends AgentTaskOptions, + AgentHooks> {} + +// agent.ts passes these runtime base classes in to avoid a circular runtime import. +type AgentCtor = new (options: AgentOptions) => Agent; + +type AgentTaskCtor = new ( + options: AgentTaskOptions, +) => AgentTask; + +export function createAgentV2( + AgentBase: AgentCtor, + options: AgentCreateOptions, +): Agent { + class AgentV2 extends AgentBase { + private readonly hookAdapter: AgentHookAdapter>; + + constructor({ + onEnter, + onExit, + onUserTurnCompleted, + sttNode, + llmNode, + ttsNode, + realtimeAudioOutputNode, + ...agentOptions + }: AgentCreateOptions) { + super({ + ...agentOptions, + id: agentOptions.id ?? 'default_agent', + }); + + this.hookAdapter = new AgentHookAdapter( + { + onEnter, + onExit, + onUserTurnCompleted, + sttNode, + llmNode, + ttsNode, + realtimeAudioOutputNode, + }, + new AgentHookContext(this), + ); + } + + override async onEnter(): Promise { + return this.hookAdapter.onEnter(() => super.onEnter()); + } + + override async onExit(): Promise { + return this.hookAdapter.onExit(() => super.onExit()); + } + + override async onUserTurnCompleted( + chatCtx: ChatContext, + newMessage: ChatMessage, + ): Promise { + return this.hookAdapter.onUserTurnCompleted(chatCtx, newMessage, () => + super.onUserTurnCompleted(chatCtx, newMessage), + ); + } + + override async sttNode( + audio: ReadableStream, + modelSettings: ModelSettings, + ): Promise | null> { + return this.hookAdapter.sttNode(audio, modelSettings, () => + super.sttNode(audio, modelSettings), + ); + } + + override async llmNode( + chatCtx: ChatContext, + toolCtx: ToolContext, + modelSettings: ModelSettings, + ): Promise | null> { + return this.hookAdapter.llmNode(chatCtx, toolCtx, modelSettings, () => + super.llmNode(chatCtx, toolCtx, modelSettings), + ); + } + + override async ttsNode( + text: ReadableStream, + modelSettings: ModelSettings, + ): Promise | null> { + return this.hookAdapter.ttsNode(text, modelSettings, () => + super.ttsNode(text, modelSettings), + ); + } + + override async realtimeAudioOutputNode( + audio: ReadableStream, + modelSettings: ModelSettings, + ): Promise | null> { + return this.hookAdapter.realtimeAudioOutputNode(audio, modelSettings, () => + super.realtimeAudioOutputNode(audio, modelSettings), + ); + } + } + + return new AgentV2(options); +} + +export function createAgentTaskV2( + AgentTaskBase: AgentTaskCtor, + options: AgentTaskCreateOptions, +): AgentTask { + class AgentTaskV2 extends AgentTaskBase { + private readonly hookAdapter: AgentHookAdapter>; + + constructor({ + onEnter, + onExit, + onUserTurnCompleted, + sttNode, + llmNode, + ttsNode, + realtimeAudioOutputNode, + ...taskOptions + }: AgentTaskCreateOptions) { + super({ + ...taskOptions, + id: taskOptions.id ?? 'default_agent', + }); + + this.hookAdapter = new AgentHookAdapter( + { + onEnter, + onExit, + onUserTurnCompleted, + sttNode, + llmNode, + ttsNode, + realtimeAudioOutputNode, + }, + new AgentTaskHookContext(this), + ); + } + + override async onEnter(): Promise { + return this.hookAdapter.onEnter(() => super.onEnter()); + } + + override async onExit(): Promise { + return this.hookAdapter.onExit(() => super.onExit()); + } + + override async onUserTurnCompleted( + chatCtx: ChatContext, + newMessage: ChatMessage, + ): Promise { + return this.hookAdapter.onUserTurnCompleted(chatCtx, newMessage, () => + super.onUserTurnCompleted(chatCtx, newMessage), + ); + } + + override async sttNode( + audio: ReadableStream, + modelSettings: ModelSettings, + ): Promise | null> { + return this.hookAdapter.sttNode(audio, modelSettings, () => + super.sttNode(audio, modelSettings), + ); + } + + override async llmNode( + chatCtx: ChatContext, + toolCtx: ToolContext, + modelSettings: ModelSettings, + ): Promise | null> { + return this.hookAdapter.llmNode(chatCtx, toolCtx, modelSettings, () => + super.llmNode(chatCtx, toolCtx, modelSettings), + ); + } + + override async ttsNode( + text: ReadableStream, + modelSettings: ModelSettings, + ): Promise | null> { + return this.hookAdapter.ttsNode(text, modelSettings, () => + super.ttsNode(text, modelSettings), + ); + } + + override async realtimeAudioOutputNode( + audio: ReadableStream, + modelSettings: ModelSettings, + ): Promise | null> { + return this.hookAdapter.realtimeAudioOutputNode(audio, modelSettings, () => + super.realtimeAudioOutputNode(audio, modelSettings), + ); + } + } + + return new AgentTaskV2(options); +} + +class AgentHookAdapter> { + constructor( + private readonly hooks: AgentHooks, + private readonly context: ContextT, + ) {} + + async onEnter(fallback: () => Promise): Promise { + if (!this.hooks.onEnter) { + return fallback(); + } + + return this.hooks.onEnter(this.context); + } + + async onExit(fallback: () => Promise): Promise { + if (!this.hooks.onExit) { + return fallback(); + } + + return this.hooks.onExit(this.context); + } + + async onUserTurnCompleted( + chatCtx: ChatContext, + newMessage: ChatMessage, + fallback: () => Promise, + ): Promise { + if (!this.hooks.onUserTurnCompleted) { + return fallback(); + } + + return this.hooks.onUserTurnCompleted(this.context, chatCtx, newMessage); + } + + async sttNode( + audio: ReadableStream, + modelSettings: ModelSettings, + fallback: () => Promise | null>, + ): Promise | null> { + if (!this.hooks.sttNode) { + return fallback(); + } + + const result = await this.hooks.sttNode(this.context, readStream(audio), modelSettings); + return result === null ? null : toStream(result); + } + + async llmNode( + chatCtx: ChatContext, + toolCtx: ToolContext, + modelSettings: ModelSettings, + fallback: () => Promise | null>, + ): Promise | null> { + if (!this.hooks.llmNode) { + return fallback(); + } + + const result = await this.hooks.llmNode( + this.context, + chatCtx, + toolCtx as ToolContext, + modelSettings, + ); + return result === null ? null : toStream(result); + } + + async ttsNode( + text: ReadableStream, + modelSettings: ModelSettings, + fallback: () => Promise | null>, + ): Promise | null> { + if (!this.hooks.ttsNode) { + return fallback(); + } + + const result = await this.hooks.ttsNode(this.context, readStream(text), modelSettings); + return result === null ? null : toStream(result); + } + + async realtimeAudioOutputNode( + audio: ReadableStream, + modelSettings: ModelSettings, + fallback: () => Promise | null>, + ): Promise | null> { + if (!this.hooks.realtimeAudioOutputNode) { + return fallback(); + } + + const result = await this.hooks.realtimeAudioOutputNode( + this.context, + readStream(audio), + modelSettings, + ); + return result === null ? null : toStream(result); + } +} + +class AgentHookContext implements AgentContext { + constructor(readonly agent: Agent) {} + + get vad(): VAD | undefined { + return this.agent.vad; + } + + get stt(): STT | undefined { + return this.agent.stt; + } + + get llm(): LLM | RealtimeModel | undefined { + return this.agent.llm; + } + + get tts(): TTS | undefined { + return this.agent.tts; + } + + get useTtsAlignedTranscript(): boolean | undefined { + return this.agent.useTtsAlignedTranscript; + } + + get ttsPronunciationMap(): TTSPronunciationMap | undefined { + return this.agent.ttsPronunciationMap; + } + + get chatCtx(): ReadonlyChatContext { + return this.agent.chatCtx; + } + + get id(): string { + return this.agent.id; + } + + get instructions(): string | Instructions { + return this.agent.instructions; + } + + get toolCtx(): ToolContext { + return this.agent.toolCtx; + } + + get session(): AgentSession { + return this.agent.session; + } + + get turnHandling(): Partial | undefined { + return this.agent.turnHandling; + } + + get minConsecutiveSpeechDelay(): number | undefined { + return this.agent.minConsecutiveSpeechDelay; + } +} + +class AgentTaskHookContext + extends AgentHookContext + implements AgentTaskContext +{ + declare readonly agent: AgentTask; + + constructor(agent: AgentTask) { + super(agent); + } + + complete(result: ResultT | Error): void { + this.agent.complete(result); + } +} diff --git a/agents/src/voice/index.ts b/agents/src/voice/index.ts index e8813e460..5833f95f2 100644 --- a/agents/src/voice/index.ts +++ b/agents/src/voice/index.ts @@ -5,7 +5,13 @@ export { Agent, AgentTask, StopResponse, + type AgentContext, + type AgentCreateOptions, + type AgentHookNodeResult, + type AgentHooks, type AgentOptions, + type AgentTaskContext, + type AgentTaskCreateOptions, type ModelSettings, type TTSPronunciationMap, } from './agent.js'; @@ -35,6 +41,8 @@ export { type PlaybackFinishedEvent, type PlaybackStartedEvent, type TimedString, + createTimedString, + isTimedString, } from './io.js'; export * from './report.js'; export * from './room_io/index.js'; diff --git a/examples/src/basic_agent.ts b/examples/src/basic_agent.ts index 79e79808a..21463556e 100644 --- a/examples/src/basic_agent.ts +++ b/examples/src/basic_agent.ts @@ -2,16 +2,18 @@ // // SPDX-License-Identifier: Apache-2.0 import { + Agent, + AgentSession, + AgentSessionEventTypes, type JobContext, type JobProcess, ServerOptions, cli, defineAgent, inference, - llm, log, - metrics, - voice, + logMetrics, + tool, } from '@livekit/agents'; import * as livekit from '@livekit/agents-plugin-livekit'; import * as silero from '@livekit/agents-plugin-silero'; @@ -24,11 +26,11 @@ export default defineAgent({ proc.userData.vad = await silero.VAD.load(); }, entry: async (ctx: JobContext) => { - const agent = new voice.Agent({ + const agent = Agent.create({ instructions: "You are a helpful assistant, you can hear the user's message and respond to it.", tools: [ - llm.tool({ + tool({ name: 'getWeather', description: 'Get the weather for a given location.', parameters: z.object({ @@ -43,7 +45,7 @@ export default defineAgent({ const logger = log(); - const session = new voice.AgentSession({ + const session = new AgentSession({ // VAD and turn detection are used to determine when the user is speaking and when the agent should respond // See more at https://docs.livekit.io/agents/build/turns vad: ctx.proc.userData.vad! as silero.VAD, @@ -104,8 +106,8 @@ export default defineAgent({ }); // Log metrics as they are emitted - session.on(voice.AgentSessionEventTypes.MetricsCollected, (ev) => { - metrics.logMetrics(ev.metrics); + session.on(AgentSessionEventTypes.MetricsCollected, (ev) => { + logMetrics(ev.metrics); }); // Log usage summary when job shuts down @@ -118,7 +120,7 @@ export default defineAgent({ ); }); - session.on(voice.AgentSessionEventTypes.OverlappingSpeech, (ev) => { + session.on(AgentSessionEventTypes.OverlappingSpeech, (ev) => { logger.warn({ type: ev.type, isInterruption: ev.isInterruption }, 'user overlapping speech'); }); diff --git a/examples/src/basic_agent_task.ts b/examples/src/basic_agent_task.ts index c450a81bb..68a049694 100644 --- a/examples/src/basic_agent_task.ts +++ b/examples/src/basic_agent_task.ts @@ -2,11 +2,14 @@ // // SPDX-License-Identifier: Apache-2.0 import { + Agent, + AgentTask, type JobContext, type JobProcess, ServerOptions, cli, defineAgent, + handoff, inference, llm, voice, @@ -16,102 +19,101 @@ import * as silero from '@livekit/agents-plugin-silero'; import { fileURLToPath } from 'node:url'; import { z } from 'zod'; -class InfoTask extends voice.AgentTask { - constructor(private info: string) { - super({ - instructions: `Collect the user's information. around ${info}. Once you have the information, call the saveUserInfo tool to save the information to the database IMMEDIATELY. DO NOT have chitchat with the user, just collect the information and call the saveUserInfo tool.`, - tts: 'elevenlabs/eleven_turbo_v2_5', - tools: [ - llm.tool({ - name: 'saveUserInfo', - description: `Save the user's ${info} to database`, - parameters: z.object({ - [info]: z.string(), - }), - execute: async (args) => { - this.complete(args[info] as string); - return `Thanks, collected ${info} successfully: ${args[info]}`; - }, +function createInfoTask(info: string): AgentTask { + const task = AgentTask.create({ + instructions: `Collect the user's information. around ${info}. Once you have the information, call the saveUserInfo tool to save the information to the database IMMEDIATELY. DO NOT have chitchat with the user, just collect the information and call the saveUserInfo tool.`, + tts: 'elevenlabs/eleven_turbo_v2_5', + tools: [ + llm.tool({ + name: 'saveUserInfo', + description: `Save the user's ${info} to database`, + parameters: z.object({ + [info]: z.string(), }), - ], - }); - } + execute: async (args) => { + task.complete(args[info] as string); + return `Thanks, collected ${info} successfully: ${args[info]}`; + }, + }), + ], + onEnter: (ctx) => { + ctx.session.generateReply({ + userInput: `Ask the user for their ${info}`, + }); + }, + }); - async onEnter() { - this.session.generateReply({ - userInput: `Ask the user for their ${this.info}`, - }); - } + return task; } -class SurveyAgent extends voice.Agent { - constructor() { - super({ - instructions: - 'You orchestrate a short intro survey. Speak naturally and keep the interaction brief.', - tools: [ - llm.tool({ - name: 'collectUserInfo', - description: 'Call this when user want to provide some information to you', - parameters: z.object({ - key: z - .string() - .describe( - 'The key of the information to collect, e.g. "name" or "role" should be no space and underscore separated', - ), - }), - execute: async ({ key }) => { - const value = await new InfoTask(key).run(); - return `Collected ${key} successfully: ${value}`; - }, +function createWeatherAgent() { + return Agent.create({ + instructions: + 'You are a weather agent. You are responsible for providing the weather information to the user.', + tts: 'deepgram/aura-2', + tools: [ + llm.tool({ + name: 'getWeather', + description: 'Get the weather for a given location', + parameters: z.object({ + location: z.string().describe('The location to get the weather for'), }), - llm.tool({ - name: 'transferToWeatherAgent', - description: 'Call this immediately after user want to know the weather', - execute: async () => { - const agent = new voice.Agent({ - instructions: - 'You are a weather agent. You are responsible for providing the weather information to the user.', - tts: 'deepgram/aura-2', - tools: [ - llm.tool({ - name: 'getWeather', - description: 'Get the weather for a given location', - parameters: z.object({ - location: z.string().describe('The location to get the weather for'), - }), - execute: async ({ location }) => { - return `The weather in ${location} is sunny today.`; - }, - }), - llm.tool({ - name: 'finishWeatherConversation', - description: 'Call this when you want to finish the weather conversation', - execute: async () => { - return llm.handoff({ - agent: new SurveyAgent(), - returns: 'Transfer to survey agent successfully!', - }); - }, - }), - ], - }); + execute: async ({ location }) => { + return `The weather in ${location} is sunny today.`; + }, + }), + llm.tool({ + name: 'finishWeatherConversation', + description: 'Call this when you want to finish the weather conversation', + execute: async () => { + return llm.handoff({ + agent: createSurveyAgent(), + returns: 'Transfer to survey agent successfully!', + }); + }, + }), + ], + }); +} - return llm.handoff({ agent, returns: "Let's start the weather conversation!" }); - }, +function createSurveyAgent(): Agent { + return Agent.create({ + instructions: + 'You orchestrate a short intro survey. Speak naturally and keep the interaction brief.', + tools: [ + llm.tool({ + name: 'collectUserInfo', + description: 'Call this when user want to provide some information to you', + parameters: z.object({ + key: z + .string() + .describe( + 'The key of the information to collect, e.g. "name" or "role" should be no space and underscore separated', + ), }), - ], - }); - } - - async onEnter() { - const name = await new InfoTask('name').run(); - const role = await new InfoTask('role').run(); + execute: async ({ key }) => { + const value = await createInfoTask(key).run(); + return `Collected ${key} successfully: ${value}`; + }, + }), + llm.tool({ + name: 'transferToWeatherAgent', + description: 'Call this immediately after user want to know the weather', + execute: async () => { + const agent = createWeatherAgent(); + return handoff({ agent, returns: "Let's start the weather conversation!" }); + }, + }), + ], + onEnter: async (ctx) => { + const name = await createInfoTask('name').run(); + const role = await createInfoTask('role').run(); - await this.session.say( - `Great to meet you ${name}. I noted your role as ${role}. We can continue now.`, - ); - } + await ctx.session.say( + `Great to meet you ${name}. I noted your role as ${role}. We can continue now.`, + ); + }, + }); } export default defineAgent({ @@ -131,7 +133,7 @@ export default defineAgent({ await session.start({ room: ctx.room, - agent: new SurveyAgent(), + agent: createSurveyAgent(), }); }, }); diff --git a/examples/src/basic_toolsets.ts b/examples/src/basic_toolsets.ts new file mode 100644 index 000000000..cb0cf2b6e --- /dev/null +++ b/examples/src/basic_toolsets.ts @@ -0,0 +1,180 @@ +// SPDX-FileCopyrightText: 2026 LiveKit, Inc. +// +// SPDX-License-Identifier: Apache-2.0 +import { + type JobContext, + type JobProcess, + ServerOptions, + cli, + defineAgent, + inference, + llm, + voice, +} from '@livekit/agents'; +import * as livekit from '@livekit/agents-plugin-livekit'; +import * as silero from '@livekit/agents-plugin-silero'; +import { BackgroundVoiceCancellation } from '@livekit/noise-cancellation-node'; +import { fileURLToPath } from 'node:url'; +import { z } from 'zod'; + +class InfoTask extends voice.AgentTask { + private key: string; + + constructor(key: string, sharedToolset: llm.Toolset) { + super({ + instructions: `Collect the user's ${key}. Once you have it, call saveUserInfo IMMEDIATELY. No chitchat.`, + tools: [ + sharedToolset, + llm.tool({ + name: 'saveUserInfo', + description: `Save the user's ${key} to the database`, + parameters: z.object({ + [key]: z.string(), + }), + execute: async (args) => { + this.complete(args[key] as string); + return `Thanks, collected ${key} successfully: ${args[key]}`; + }, + }), + ], + }); + this.key = key; + } + + async onEnter() { + this.session.generateReply({ userInput: `Ask the user for their ${this.key}` }); + } +} + +function makeWeatherAgent(returnHome: () => voice.Agent) { + const weatherToolset = new llm.Toolset({ + id: 'weather_tools', + tools: [ + llm.tool({ + name: 'getWeather', + description: 'Get the weather for a given location', + parameters: z.object({ location: z.string() }), + execute: async ({ location }) => `The weather in ${location} is sunny today.`, + }), + ], + }); + + return new voice.Agent({ + instructions: 'You are a weather agent. Provide weather information then hand back when done.', + tools: [ + weatherToolset, + llm.tool({ + name: 'finishWeatherConversation', + description: 'Call this when you want to finish the weather conversation', + execute: async () => { + return llm.handoff({ agent: returnHome(), returns: 'Transfer back to main agent.' }); + }, + }), + ], + }); +} + +class MainAgent extends voice.Agent { + private locationToolset: llm.Toolset; + + constructor(locationToolset: llm.Toolset) { + super({ + instructions: + 'You are a helpful assistant. Use the location toolset for weather/timezone. Use transferToWeather when the user asks about weather. Use swapToolset / reapplyTools to exercise updateTools.', + tools: [ + locationToolset, + llm.tool({ + name: 'transferToWeather', + description: 'Call this when the user wants to know the weather', + execute: async () => { + return llm.handoff({ + agent: makeWeatherAgent(() => new MainAgent(locationToolset)), + returns: "Let's switch to the weather agent.", + }); + }, + }), + llm.tool({ + name: 'swapToolset', + description: 'Replace the active toolset with a brand-new toolset (tests updateTools).', + execute: async () => { + const replacement = new llm.Toolset({ + id: 'location_tools_v2', + tools: [ + llm.tool({ + name: 'getWeather', + description: 'v2 weather', + parameters: z.object({ location: z.string() }), + execute: async ({ location }) => `v2: ${location} -> sunny`, + }), + ], + }); + await this.updateTools([replacement]); + return 'Swapped toolset.'; + }, + }), + llm.tool({ + name: 'reapplyTools', + description: 'Re-apply the current tool list unchanged (idempotent updateTools).', + execute: async () => { + await this.updateTools([...this.toolCtx.tools]); + return 'Re-applied the same tool list.'; + }, + }), + ], + }); + this.locationToolset = locationToolset; + } + + async onEnter() { + const name = await new InfoTask('name', this.locationToolset).run(); + await this.session.say( + `Got it, ${name}. Ask me about weather, or say "swap" / "reapply" to exercise updateTools.`, + ); + } +} + +export default defineAgent({ + prewarm: async (proc: JobProcess) => { + proc.userData.vad = await silero.VAD.load(); + }, + entry: async (ctx: JobContext) => { + const locationToolset = new llm.Toolset({ + id: 'location_tools', + tools: [ + llm.tool({ + name: 'getWeather', + description: 'Get the weather for a given location.', + parameters: z.object({ location: z.string() }), + execute: async ({ location }) => `The weather in ${location} is sunny.`, + }), + llm.tool({ + name: 'lookupTimezone', + description: 'Look up the timezone for a city or region.', + parameters: z.object({ location: z.string() }), + execute: async ({ location }) => `${location} is in the America/Los_Angeles timezone.`, + }), + ], + }); + + const session = new voice.AgentSession({ + vad: ctx.proc.userData.vad! as silero.VAD, + stt: new inference.STT({ model: 'deepgram/nova-3', language: 'en' }), + llm: new inference.LLM({ model: 'openai/gpt-4.1-mini' }), + tts: new inference.TTS({ + model: 'cartesia/sonic-3', + voice: '9626c31c-bec5-4cca-baa8-f8ba9e84c8bc', + }), + turnDetection: new livekit.turnDetector.MultilingualModel(), + }); + + await session.start({ + agent: new MainAgent(locationToolset), + room: ctx.room, + inputOptions: { noiseCancellation: BackgroundVoiceCancellation() }, + }); + + session.say('Hello! I will ask you a quick question, then we can chat.'); + }, +}); + +cli.runApp(new ServerOptions({ agent: fileURLToPath(import.meta.url) })); diff --git a/examples/src/gemini_realtime_agent.ts b/examples/src/gemini_realtime_agent.ts index d70368352..93e1bb890 100644 --- a/examples/src/gemini_realtime_agent.ts +++ b/examples/src/gemini_realtime_agent.ts @@ -78,7 +78,7 @@ class IntroAgent extends voice.Agent { }); } - static create() { + static createIntroAgent() { return new IntroAgent({ instructions: `You are a story teller. Your goal is to gather a few pieces of information from the user to make the story personalized and engaging. Ask the user for their name and where they are from.`, tools: [ @@ -94,7 +94,7 @@ class IntroAgent extends voice.Agent { ctx.userData.name = name; ctx.userData.location = location; - const storyAgent = StoryAgent.create(name, location); + const storyAgent = StoryAgent.createStoryAgent(name, location); return llm.handoff({ agent: storyAgent, returns: "Let's start the story!" }); }, }), @@ -110,7 +110,7 @@ class StoryAgent extends voice.Agent { this.session.generateReply(); } - static create(name: string, location: string) { + static createStoryAgent(name: string, location: string) { return new StoryAgent({ instructions: dedent` You are a storyteller. Use the user's information in order to make the story personalized. @@ -142,7 +142,7 @@ export default defineAgent({ }); await session.start({ - agent: IntroAgent.create(), + agent: IntroAgent.createIntroAgent(), room: ctx.room, }); diff --git a/examples/src/multi_agent.ts b/examples/src/multi_agent.ts index 12d87377d..cd22d299d 100644 --- a/examples/src/multi_agent.ts +++ b/examples/src/multi_agent.ts @@ -32,7 +32,7 @@ class IntroAgent extends voice.Agent { }); } - static create() { + static createIntroAgent() { return new IntroAgent({ instructions: `You are a story teller. Your goal is to gather a few pieces of information from the user to make the story personalized and engaging. Ask the user for their name and where they are from.`, tools: [ @@ -48,7 +48,7 @@ class IntroAgent extends voice.Agent { ctx.userData.name = name; ctx.userData.location = location; - const storyAgent = StoryAgent.create(name, location); + const storyAgent = StoryAgent.createStoryAgent(name, location); return llm.handoff({ agent: storyAgent, returns: "Let's start the story!" }); }, }), @@ -62,7 +62,7 @@ class StoryAgent extends voice.Agent { this.session.generateReply(); } - static create(name: string, location: string) { + static createStoryAgent(name: string, location: string) { return new StoryAgent({ instructions: dedent` You are a storyteller. Use the user's information in order to make the story personalized. @@ -94,7 +94,7 @@ export default defineAgent({ }); await session.start({ - agent: IntroAgent.create(), + agent: IntroAgent.createIntroAgent(), room: ctx.room, }); diff --git a/plugins/google/src/beta/realtime/realtime_api.ts b/plugins/google/src/beta/realtime/realtime_api.ts index 66bc6a7f9..1d8365584 100644 --- a/plugins/google/src/beta/realtime/realtime_api.ts +++ b/plugins/google/src/beta/realtime/realtime_api.ts @@ -33,7 +33,7 @@ import { import { Mutex } from '@livekit/mutex'; import { AudioFrame, AudioResampler, type VideoFrame } from '@livekit/rtc-node'; import { type LLMTools } from '../../tools.js'; -import { toFunctionDeclarations } from '../../utils.js'; +import { toToolsConfig } from '../../utils.js'; import type * as api_proto from './api_proto.js'; import type { LiveAPIModels, Voice } from './api_proto.js'; @@ -70,13 +70,6 @@ export interface InputTranscription { transcript: string; } -/** - * Helper function to check if two sets are equal - */ -function setsEqual(a: Set, b: Set): boolean { - return a.size === b.size && [...a].every((x) => b.has(x)); -} - /** * Internal realtime options for Google Realtime API */ @@ -455,7 +448,6 @@ export class RealtimeSession extends llm.RealtimeSession { private _chatCtx = llm.ChatContext.empty(); private options: RealtimeOptions; - private geminiDeclarations: types.FunctionDeclaration[] = []; private messageChannel = new Queue(); private inputResampler?: AudioResampler; private inputResamplerInputRate?: number; @@ -764,15 +756,12 @@ export class RealtimeSession extends llm.RealtimeSession { } async updateTools(tools: llm.ToolContext): Promise { - const newDeclarations = toFunctionDeclarations(tools); - const currentToolNames = new Set(this.geminiDeclarations.map((f) => f.name)); - const newToolNames = new Set(newDeclarations.map((f) => f.name)); - - if (!setsEqual(currentToolNames, newToolNames)) { - this.geminiDeclarations = newDeclarations; - this._tools = tools; - this.markRestartNeeded(); + if (this._tools.equals(tools)) { + return; } + + this._tools = tools; + this.markRestartNeeded(); } get chatCtx(): llm.ChatContext { @@ -1424,21 +1413,11 @@ export class RealtimeSession extends llm.RealtimeSession { }, languageCode: opts.language, }, - tools: - this.geminiDeclarations.length > 0 || this.options.geminiTools - ? [ - { - functionDeclarations: - this.options.toolBehavior !== undefined - ? this.geminiDeclarations.map((d) => ({ - ...d, - behavior: this.options.toolBehavior, - })) - : this.geminiDeclarations, - ...this.options.geminiTools, - }, - ] - : undefined, + tools: toToolsConfig({ + toolCtx: this._tools, + geminiTools: this.options.geminiTools, + toolBehavior: this.options.toolBehavior, + }), inputAudioTranscription: opts.inputAudioTranscription, outputAudioTranscription: opts.outputAudioTranscription, sessionResumption: this.sessionResumptionHandle diff --git a/plugins/google/src/index.ts b/plugins/google/src/index.ts index fbafc1d66..326ca6270 100644 --- a/plugins/google/src/index.ts +++ b/plugins/google/src/index.ts @@ -6,6 +6,7 @@ import { Plugin } from '@livekit/agents'; export * as beta from './beta/index.js'; export { LLM, LLMStream, type LLMOptions } from './llm.js'; export * from './models.js'; +export * from './tools.js'; class GooglePlugin extends Plugin { constructor() { diff --git a/plugins/google/src/llm.ts b/plugins/google/src/llm.ts index e452b70d2..4958ccc6c 100644 --- a/plugins/google/src/llm.ts +++ b/plugins/google/src/llm.ts @@ -13,7 +13,7 @@ import { } from '@livekit/agents'; import type { ChatModels } from './models.js'; import type { LLMTools } from './tools.js'; -import { toFunctionDeclarations } from './utils.js'; +import { toToolsConfig } from './utils.js'; interface GoogleFormatData { systemMessages: string[] | null; @@ -355,11 +355,11 @@ export class LLMStream extends llm.LLMStream { parts: turn.parts as types.Part[], })); - const functionDeclarations = this.toolCtx ? toFunctionDeclarations(this.toolCtx) : undefined; - const tools = - functionDeclarations && functionDeclarations.length > 0 - ? [{ functionDeclarations }] - : undefined; + const tools = toToolsConfig({ + toolCtx: this.toolCtx, + geminiTools: this.#geminiTools, + onlySingleType: true, + }); let systemInstruction: types.Content | undefined = undefined; if (extraData.systemMessages && extraData.systemMessages.length > 0) { diff --git a/plugins/google/src/tools.ts b/plugins/google/src/tools.ts index 90cd9cc7c..b864e98a6 100644 --- a/plugins/google/src/tools.ts +++ b/plugins/google/src/tools.ts @@ -1,6 +1,100 @@ // SPDX-FileCopyrightText: 2025 LiveKit, Inc. // // SPDX-License-Identifier: Apache-2.0 -import type { Tool } from '@google/genai'; +import type * as types from '@google/genai'; +import { llm } from '@livekit/agents'; -export type LLMTools = Omit; +export type LLMTools = Omit; + +export abstract class GeminiTool extends llm.ProviderTool { + abstract toToolConfig(): types.Tool; +} + +export class GoogleSearch extends GeminiTool { + constructor(public readonly options: types.GoogleSearch = {}) { + super({ id: 'gemini_google_search' }); + } + + toToolConfig(): types.Tool { + return { googleSearch: this.options }; + } +} + +export class GoogleMaps extends GeminiTool { + constructor(public readonly options: types.GoogleMaps = {}) { + super({ id: 'gemini_google_maps' }); + } + + toToolConfig(): types.Tool { + return { googleMaps: this.options }; + } +} + +export class URLContext extends GeminiTool { + constructor() { + super({ id: 'gemini_url_context' }); + } + + toToolConfig(): types.Tool { + return { urlContext: {} }; + } +} + +export interface FileSearchOptions extends types.FileSearch { + fileSearchStoreNames: string[]; +} + +export class FileSearch extends GeminiTool { + constructor(public readonly options: FileSearchOptions) { + super({ id: 'gemini_file_search' }); + } + + toToolConfig(): types.Tool { + return { fileSearch: this.options }; + } +} + +export class ToolCodeExecution extends GeminiTool { + constructor() { + super({ id: 'gemini_code_execution' }); + } + + toToolConfig(): types.Tool { + return { codeExecution: {} }; + } +} + +export interface VertexRAGRetrievalOptions { + ragResources: string[]; + similarityTopK?: number; + vectorDistanceThreshold?: number; +} + +export class VertexRAGRetrieval extends GeminiTool { + readonly ragResources: string[]; + readonly similarityTopK: number; + readonly vectorDistanceThreshold?: number; + + constructor({ + ragResources, + similarityTopK = 3, + vectorDistanceThreshold, + }: VertexRAGRetrievalOptions) { + super({ id: 'gemini_vertex_rag_retrieval' }); + this.ragResources = ragResources; + this.similarityTopK = similarityTopK; + this.vectorDistanceThreshold = vectorDistanceThreshold; + } + + toToolConfig(): types.Tool { + return { + retrieval: { + vertexRagStore: { + ragResources: this.ragResources.map((ragCorpus) => ({ ragCorpus })), + similarityTopK: this.similarityTopK, + vectorDistanceThreshold: this.vectorDistanceThreshold, + }, + }, + }; + } +} diff --git a/plugins/google/src/utils.ts b/plugins/google/src/utils.ts index 5548c076e..30b5fcc89 100644 --- a/plugins/google/src/utils.ts +++ b/plugins/google/src/utils.ts @@ -1,9 +1,11 @@ // SPDX-FileCopyrightText: 2025 LiveKit, Inc. // // SPDX-License-Identifier: Apache-2.0 +import type * as types from '@google/genai'; import type { FunctionDeclaration, Schema } from '@google/genai'; import { llm } from '@livekit/agents'; import type { JSONSchema7 } from 'json-schema'; +import { GeminiTool, type LLMTools } from './tools.js'; /** * JSON Schema v7 @@ -139,8 +141,10 @@ function isEmptyObjectSchema(jsonSchema: JSONSchema7Definition): boolean { export function toFunctionDeclarations(toolCtx: llm.ToolContext): FunctionDeclaration[] { const functionDeclarations: FunctionDeclaration[] = []; - for (const [name, tool] of Object.entries(toolCtx.functionTools)) { - const { description, parameters } = tool; + for (const tool of toolCtx.flatten()) { + // TODO: support provider tools in the Gemini schema. + if (!llm.isFunctionTool(tool)) continue; + const { name, description, parameters } = tool; const jsonSchema = llm.toJsonSchema(parameters, false); // Create a deep copy to prevent the Google GenAI library from mutating the schema @@ -155,3 +159,57 @@ export function toFunctionDeclarations(toolCtx: llm.ToolContext): FunctionDeclar return functionDeclarations; } + +export function toToolsConfig({ + toolCtx, + geminiTools, + toolBehavior, + onlySingleType = false, +}: { + toolCtx?: llm.ToolContext; + geminiTools?: LLMTools; + toolBehavior?: types.Behavior; + onlySingleType?: boolean; +}): types.Tool[] | undefined { + const tools: types.Tool[] = []; + const providerTools: types.Tool[] = []; + + if (toolCtx) { + const functionDeclarations = toFunctionDeclarations(toolCtx); + if (functionDeclarations.length > 0) { + tools.push({ + functionDeclarations: + toolBehavior !== undefined + ? functionDeclarations.map((declaration) => ({ + ...declaration, + behavior: toolBehavior, + })) + : functionDeclarations, + }); + } + } + + if (geminiTools !== undefined) { + providerTools.push(geminiTools); + } + + if (toolCtx) { + for (const tool of toolCtx.providerTools) { + if (tool instanceof GeminiTool) { + providerTools.push(tool.toToolConfig()); + } + } + } + + if (tools.length > 0 && providerTools.length > 0) { + throw new Error('Gemini does not support mixing function tools and provider tools'); + } + + if (onlySingleType && tools.length > 0) { + return tools; + } + + tools.push(...providerTools); + + return tools.length > 0 ? tools : undefined; +} diff --git a/plugins/mistralai/src/llm.ts b/plugins/mistralai/src/llm.ts index f6685b042..901325f0a 100644 --- a/plugins/mistralai/src/llm.ts +++ b/plugins/mistralai/src/llm.ts @@ -211,14 +211,16 @@ export class LLMStream extends llm.LLMStream { // eslint-disable-next-line @typescript-eslint/no-explicit-any const toolsList: any[] = []; - if (this.toolCtx && Object.keys(this.toolCtx.functionTools).length > 0) { - for (const [name, func] of Object.entries(this.toolCtx.functionTools)) { + if (this.toolCtx) { + for (const t of this.toolCtx.flatten()) { + // TODO: support provider tools in the Mistral schema. + if (!llm.isFunctionTool(t)) continue; toolsList.push({ type: 'function' as const, function: { - name, - description: func.description, - parameters: llm.toJsonSchema(func.parameters, true, false), + name: t.name, + description: t.description, + parameters: llm.toJsonSchema(t.parameters, true, false), }, }); } diff --git a/plugins/openai/src/index.ts b/plugins/openai/src/index.ts index ccffdcb3f..6a5d9cb7c 100644 --- a/plugins/openai/src/index.ts +++ b/plugins/openai/src/index.ts @@ -5,6 +5,7 @@ import { Plugin } from '@livekit/agents'; export { LLM, LLMStream, type LLMOptions } from './llm.js'; export * from './models.js'; +export * from './tools.js'; export * as realtime from './realtime/index.js'; export * as responses from './responses/index.js'; export { STT, type STTOptions } from './stt.js'; diff --git a/plugins/openai/src/realtime/realtime_model.ts b/plugins/openai/src/realtime/realtime_model.ts index 8481e0f47..49a0cb506 100644 --- a/plugins/openai/src/realtime/realtime_model.ts +++ b/plugins/openai/src/realtime/realtime_model.ts @@ -698,11 +698,12 @@ export class RealtimeSession extends llm.RealtimeSession { // TODO(brian): these logics below are noops I think, leaving it here to keep // parity with the python but we should remove them later const retainedToolNames = new Set(ev.session.tools.map((tool) => tool.name)); - const retainedTools = Object.entries(_tools.functionTools) - .filter(([name]) => retainedToolNames.has(name)) - .map(([, tool]) => tool); + // Keep provider tools and Toolsets as-is; only drop function tools the server didn't accept. + const retainedEntries = _tools.tools.filter( + (entry) => !llm.isFunctionTool(entry) || retainedToolNames.has(entry.name), + ); - this._tools = new llm.ToolContext(retainedTools); + this._tools = new llm.ToolContext(retainedEntries); unlock(); } @@ -710,21 +711,26 @@ export class RealtimeSession extends llm.RealtimeSession { private createToolsUpdateEvent(_tools: llm.ToolContext): api_proto.SessionUpdateEvent { const oaiTools: api_proto.Tool[] = []; - for (const [name, tool] of Object.entries(_tools.functionTools)) { - const { parameters: toolParameters, description } = tool; + for (const t of _tools.flatten()) { + // TODO: support provider tools in the Realtime session-update schema. + if (!llm.isFunctionTool(t)) continue; + try { const parameters = llm.toJsonSchema( - toolParameters, + t.parameters, ) as unknown as api_proto.Tool['parameters']; oaiTools.push({ - name, - description, + name: t.name, + description: t.description, parameters: parameters, type: 'function', }); } catch (e) { - this.#logger.error({ name, tool }, "OpenAI Realtime API doesn't support this tool type"); + this.#logger.error( + { name: t.name, tool: t }, + "OpenAI Realtime API doesn't support this tool type", + ); continue; } } diff --git a/plugins/openai/src/responses/llm.ts b/plugins/openai/src/responses/llm.ts index 9a255d046..4a1dc9d92 100644 --- a/plugins/openai/src/responses/llm.ts +++ b/plugins/openai/src/responses/llm.ts @@ -13,6 +13,7 @@ import { } from '@livekit/agents'; import OpenAI from 'openai'; import type { ChatModels } from '../models.js'; +import { toResponsesTools } from '../tool_utils.js'; import { WSLLM } from '../ws/llm.js'; export interface LLMOptions { @@ -187,24 +188,7 @@ class ResponsesHttpLLMStream extends llm.LLMStream { )) as OpenAI.Responses.ResponseInputItem[]; const tools = this.toolCtx - ? Object.entries(this.toolCtx.functionTools).map(([name, func]) => { - const oaiParams = { - type: 'function' as const, - name: name, - description: func.description, - parameters: llm.toJsonSchema( - func.parameters, - true, - this.strictToolSchema, - ) as unknown as OpenAI.Responses.FunctionTool['parameters'], - } as OpenAI.Responses.FunctionTool; - - if (this.strictToolSchema) { - oaiParams.strict = true; - } - - return oaiParams; - }) + ? toResponsesTools(this.toolCtx, this.strictToolSchema) : undefined; const requestOptions: Record = { ...this.modelOptions }; diff --git a/plugins/openai/src/tool_utils.test.ts b/plugins/openai/src/tool_utils.test.ts new file mode 100644 index 000000000..ce1922f52 --- /dev/null +++ b/plugins/openai/src/tool_utils.test.ts @@ -0,0 +1,84 @@ +// SPDX-FileCopyrightText: 2026 LiveKit, Inc. +// +// SPDX-License-Identifier: Apache-2.0 +import { llm } from '@livekit/agents'; +import { describe, expect, it } from 'vitest'; +import { z } from 'zod'; +import { toResponsesTools } from './tool_utils.js'; +import { CodeInterpreter, FileSearch, WebSearch } from './tools.js'; + +describe('toResponsesTools', () => { + it('serializes function tools', () => { + const fn = llm.tool({ + name: 'lookup_weather', + description: 'Look up weather', + parameters: z.object({ city: z.string() }), + execute: async () => 'sunny', + }); + + expect(toResponsesTools(new llm.ToolContext([fn]), true)).toEqual([ + { + type: 'function', + name: 'lookup_weather', + description: 'Look up weather', + parameters: { + $schema: 'http://json-schema.org/draft-07/schema#', + type: 'object', + properties: { city: { type: 'string' } }, + required: ['city'], + additionalProperties: false, + }, + strict: true, + }, + ]); + }); + + it('serializes OpenAI provider tools', () => { + const tools = toResponsesTools( + new llm.ToolContext([ + new WebSearch({ + filters: { allowed_domains: ['docs.livekit.io'] }, + searchContextSize: 'low', + userLocation: { type: 'approximate', country: 'US' }, + }), + new FileSearch({ + vectorStoreIds: ['vs_123'], + maxNumResults: 3, + rankingOptions: { ranker: 'auto' }, + }), + new CodeInterpreter({ container: { type: 'auto', file_ids: ['file_123'] } }), + ]), + false, + ); + + expect(tools).toEqual([ + { + type: 'web_search', + search_context_size: 'low', + filters: { allowed_domains: ['docs.livekit.io'] }, + user_location: { type: 'approximate', country: 'US' }, + }, + { + type: 'file_search', + vector_store_ids: ['vs_123'], + max_num_results: 3, + ranking_options: { ranker: 'auto' }, + }, + { type: 'code_interpreter', container: { type: 'auto', file_ids: ['file_123'] } }, + ]); + }); + + it('omits the code interpreter container when unset', () => { + expect(toResponsesTools(new llm.ToolContext([new CodeInterpreter()]), false)).toEqual([ + { type: 'code_interpreter' }, + ]); + }); + + it('ignores non-OpenAI provider tools', () => { + class OtherProviderTool extends llm.ProviderTool {} + + expect( + toResponsesTools(new llm.ToolContext([new OtherProviderTool({ id: 'other' })]), false), + ).toBeUndefined(); + }); +}); diff --git a/plugins/openai/src/tool_utils.ts b/plugins/openai/src/tool_utils.ts new file mode 100644 index 000000000..1e2e6c709 --- /dev/null +++ b/plugins/openai/src/tool_utils.ts @@ -0,0 +1,43 @@ +// SPDX-FileCopyrightText: 2026 LiveKit, Inc. +// +// SPDX-License-Identifier: Apache-2.0 +import { llm } from '@livekit/agents'; +import type OpenAI from 'openai'; +import { OpenAITool } from './tools.js'; + +export function toResponsesTools( + toolCtx: llm.ToolContext, + strictToolSchema: boolean, +): OpenAI.Responses.Tool[] | undefined { + const tools = toolCtx + .flatten() + .map((tool) => { + if (llm.isFunctionTool(tool)) { + const oaiParams = { + type: 'function' as const, + name: tool.name, + description: tool.description, + parameters: llm.toJsonSchema( + tool.parameters, + true, + strictToolSchema, + ) as unknown as OpenAI.Responses.FunctionTool['parameters'], + } as OpenAI.Responses.FunctionTool; + + if (strictToolSchema) { + oaiParams.strict = true; + } + + return oaiParams; + } + + if (tool instanceof OpenAITool) { + return tool.toToolConfig() as unknown as OpenAI.Responses.Tool; + } + + return undefined; + }) + .filter((tool): tool is OpenAI.Responses.Tool => tool !== undefined); + + return tools.length > 0 ? tools : undefined; +} diff --git a/plugins/openai/src/tools.ts b/plugins/openai/src/tools.ts new file mode 100644 index 000000000..1ce779376 --- /dev/null +++ b/plugins/openai/src/tools.ts @@ -0,0 +1,166 @@ +// SPDX-FileCopyrightText: 2026 LiveKit, Inc. +// +// SPDX-License-Identifier: Apache-2.0 +import { llm } from '@livekit/agents'; +import type OpenAI from 'openai'; + +/** Base class for OpenAI Responses API provider tools. */ +export abstract class OpenAITool extends llm.ProviderTool { + /** Convert this provider tool to the OpenAI Responses API tool configuration. */ + abstract toToolConfig(): Record; +} + +/** + * High-level guidance for the amount of context window space to use for web search. + * OpenAI defaults this to `medium`. + */ +export type WebSearchContextSize = 'low' | 'medium' | 'high'; + +/** Options for the OpenAI web search tool. */ +export interface WebSearchOptions { + /** + * Filters for the search, such as allowed domains. If not provided, all domains are allowed. + */ + filters?: OpenAI.Responses.WebSearchTool['filters']; + + /** + * Amount of context window space to use for the search. Defaults to `medium`. + */ + searchContextSize?: WebSearchContextSize | null; + + /** Approximate location of the user, such as city, region, country, or timezone. */ + userLocation?: OpenAI.Responses.WebSearchTool['user_location']; +} + +/** + * Search the Internet for sources related to the prompt. + * + * @see https://platform.openai.com/docs/guides/tools-web-search + */ +export class WebSearch extends OpenAITool { + /** Filters for the search, such as allowed domains. */ + readonly filters: OpenAI.Responses.WebSearchTool['filters'] | undefined; + + /** Amount of context window space to use for the search. */ + readonly searchContextSize: WebSearchContextSize | null; + + /** Approximate location of the user. */ + readonly userLocation: OpenAI.Responses.WebSearchTool['user_location'] | undefined; + + constructor({ filters, searchContextSize = 'medium', userLocation }: WebSearchOptions = {}) { + super({ id: 'openai_web_search' }); + this.filters = filters; + this.searchContextSize = searchContextSize; + this.userLocation = userLocation; + } + + toToolConfig(): Record { + const result: Record = { + type: 'web_search', + search_context_size: this.searchContextSize, + }; + if (this.userLocation !== undefined) { + result.user_location = this.userLocation; + } + if (this.filters !== undefined) { + result.filters = this.filters; + } + return result; + } +} + +/** Options for the OpenAI file search tool. */ +export interface FileSearchOptions { + /** IDs of the vector stores to search. */ + vectorStoreIds?: string[]; + + /** Filter to apply to file search results. */ + filters?: OpenAI.Responses.FileSearchTool['filters']; + + /** Maximum number of results to return. This should be between 1 and 50 inclusive. */ + maxNumResults?: number; + + /** Ranking options for search, including ranker and score threshold. */ + rankingOptions?: OpenAI.Responses.FileSearchTool.RankingOptions; +} + +/** + * Search for relevant content from uploaded files. + * + * @see https://platform.openai.com/docs/guides/tools-file-search + */ +export class FileSearch extends OpenAITool { + /** IDs of the vector stores to search. */ + readonly vectorStoreIds: string[]; + + /** Filter to apply to file search results. */ + readonly filters: OpenAI.Responses.FileSearchTool['filters'] | undefined; + + /** Maximum number of results to return. */ + readonly maxNumResults: number | undefined; + + /** Ranking options for search. */ + readonly rankingOptions: OpenAI.Responses.FileSearchTool.RankingOptions | undefined; + + constructor({ + vectorStoreIds = [], + filters, + maxNumResults, + rankingOptions, + }: FileSearchOptions = {}) { + super({ id: 'openai_file_search' }); + this.vectorStoreIds = [...vectorStoreIds]; + this.filters = filters; + this.maxNumResults = maxNumResults; + this.rankingOptions = rankingOptions; + } + + toToolConfig(): Record { + const result: Record = { + type: 'file_search', + vector_store_ids: this.vectorStoreIds, + }; + if (this.filters !== undefined) { + result.filters = this.filters; + } + if (this.maxNumResults !== undefined) { + result.max_num_results = this.maxNumResults; + } + if (this.rankingOptions !== undefined) { + result.ranking_options = this.rankingOptions; + } + return result; + } +} + +/** Options for the OpenAI code interpreter tool. */ +export interface CodeInterpreterOptions { + /** + * Code interpreter container. Can be a container ID or an object that specifies uploaded file IDs + * to make available to the code. + */ + container?: OpenAI.Responses.Tool.CodeInterpreter['container'] | null; +} + +/** + * Run Python code to help generate a response to a prompt. + * + * @see https://platform.openai.com/docs/guides/tools-code-interpreter + */ +export class CodeInterpreter extends OpenAITool { + /** Code interpreter container ID or configuration. */ + readonly container: OpenAI.Responses.Tool.CodeInterpreter['container'] | null; + + constructor({ container = null }: CodeInterpreterOptions = {}) { + super({ id: 'openai_code_interpreter' }); + this.container = container; + } + + toToolConfig(): Record { + const result: Record = { type: 'code_interpreter' }; + if (this.container !== null) { + result.container = this.container; + } + return result; + } +} diff --git a/plugins/openai/src/ws/llm.ts b/plugins/openai/src/ws/llm.ts index d22d7a753..f75054387 100644 --- a/plugins/openai/src/ws/llm.ts +++ b/plugins/openai/src/ws/llm.ts @@ -15,6 +15,7 @@ import { import type OpenAI from 'openai'; import { WebSocket } from 'ws'; import type { ChatModels } from '../models.js'; +import { toResponsesTools } from '../tool_utils.js'; import type { WsOutputItemDoneEvent, WsOutputTextDeltaEvent, @@ -429,26 +430,7 @@ export class WSLLMStream extends llm.LLMStream { 'openai.responses', )) as OpenAI.Responses.ResponseInputItem[]; - const tools = this.toolCtx - ? Object.entries(this.toolCtx.functionTools).map(([name, func]) => { - const oaiParams = { - type: 'function' as const, - name, - description: func.description, - parameters: llm.toJsonSchema( - func.parameters, - true, - this.#strictToolSchema, - ) as unknown as OpenAI.Responses.FunctionTool['parameters'], - } as OpenAI.Responses.FunctionTool; - - if (this.#strictToolSchema) { - oaiParams.strict = true; - } - - return oaiParams; - }) - : undefined; + const tools = this.toolCtx ? toResponsesTools(this.toolCtx, this.#strictToolSchema) : undefined; const requestOptions: Record = { ...this.#modelOptions }; if (!tools) { diff --git a/plugins/phonic/src/realtime/realtime_model.ts b/plugins/phonic/src/realtime/realtime_model.ts index 09933b580..bec3906dd 100644 --- a/plugins/phonic/src/realtime/realtime_model.ts +++ b/plugins/phonic/src/realtime/realtime_model.ts @@ -368,23 +368,25 @@ export class RealtimeSession extends llm.RealtimeSession { } this._tools = tools.copy(); - this.toolDefinitions = Object.entries(tools.functionTools).map(([name, tool]) => ({ - type: 'custom_websocket', - tool_schema: { - type: 'function', - function: { - name, - description: tool.description, - parameters: llm.toJsonSchema(tool.parameters), - strict: true, + // TODO: support provider tools in the Phonic schema. + this.toolDefinitions = tools + .flatten() + .filter(llm.isFunctionTool) + .map((t) => ({ + type: 'custom_websocket' as const, + tool_schema: { + type: 'function' as const, + function: { + name: t.name, + description: t.description, + parameters: llm.toJsonSchema(t.parameters), + strict: true, + }, }, - }, - tool_call_output_timeout_ms: TOOL_CALL_OUTPUT_TIMEOUT_MS, - // Tool chaining and tool calls during speech are not supported at this time - // for ease of implementation within the RealtimeSession generations framework - wait_for_speech_before_tool_call: true, - allow_tool_chaining: false, - })); + tool_call_output_timeout_ms: TOOL_CALL_OUTPUT_TIMEOUT_MS, + wait_for_speech_before_tool_call: true, + allow_tool_chaining: false, + })); this.toolsReady.resolve(); } @@ -404,21 +406,25 @@ export class RealtimeSession extends llm.RealtimeSession { } if (tools !== undefined) { this._tools = tools.copy(); - this.toolDefinitions = Object.entries(tools.functionTools).map(([name, tool]) => ({ - type: 'custom_websocket', - tool_schema: { - type: 'function', - function: { - name, - description: tool.description, - parameters: llm.toJsonSchema(tool.parameters), - strict: true, + // TODO: support provider tools in the Phonic schema. + this.toolDefinitions = tools + .flatten() + .filter(llm.isFunctionTool) + .map((t) => ({ + type: 'custom_websocket' as const, + tool_schema: { + type: 'function' as const, + function: { + name: t.name, + description: t.description, + parameters: llm.toJsonSchema(t.parameters), + strict: true, + }, }, - }, - tool_call_output_timeout_ms: TOOL_CALL_OUTPUT_TIMEOUT_MS, - wait_for_speech_before_tool_call: true, - allow_tool_chaining: false, - })); + tool_call_output_timeout_ms: TOOL_CALL_OUTPUT_TIMEOUT_MS, + wait_for_speech_before_tool_call: true, + allow_tool_chaining: false, + })); } if (chatCtx !== undefined) { this._chatCtx = chatCtx.copy(); diff --git a/plugins/test/src/llm.ts b/plugins/test/src/llm.ts index 534dce2df..1f912654a 100644 --- a/plugins/test/src/llm.ts +++ b/plugins/test/src/llm.ts @@ -200,6 +200,57 @@ export const llm = async (llm: llmlib.LLM, skipOptionalArgs: boolean) => { expect(JSON.parse(calls[0]!.args).address).toBeUndefined(); }); }); + + describe('toolset', async () => { + const buildToolsetContext = () => { + const weatherToolset = new llmlib.Toolset({ + id: 'weather_toolset', + tools: [ + llmlib.tool({ + name: 'getWeather', + description: 'Get the current weather in a given location', + parameters: z.object({ + location: z.string().describe('The city and state, e.g. San Francisco, CA'), + unit: z.enum(['celsius', 'fahrenheit']).describe('The temperature unit to use'), + }), + execute: async () => {}, + }), + ], + }); + + const directTool = llmlib.tool({ + name: 'playMusic', + description: 'Play music', + parameters: z.object({ + name: z.string().describe('The artist and name of the song'), + }), + execute: async () => {}, + }); + + return new llmlib.ToolContext([weatherToolset, directTool]); + }; + + it('should call a function tool that lives inside a Toolset', async () => { + const ctx = buildToolsetContext(); + const calls = await requestFncCall( + llm, + "What's the weather in San Francisco, in Celsius?", + ctx, + ); + + expect(calls.length).toStrictEqual(1); + expect(calls[0]!.name).toStrictEqual('getWeather'); + expect(JSON.parse(calls[0]!.args).unit).toStrictEqual('celsius'); + }); + + it('should expose direct tools alongside Toolset tools', async () => { + const ctx = buildToolsetContext(); + const calls = await requestFncCall(llm, 'Play the song "Bohemian Rhapsody" by Queen.', ctx); + + expect(calls.length).toStrictEqual(1); + expect(calls[0]!.name).toStrictEqual('playMusic'); + }); + }); }); }; From a24fb0be51c96f1f8cee06e4cf04143ca774d232 Mon Sep 17 00:00:00 2001 From: Toubat Date: Mon, 8 Jun 2026 14:35:49 -0700 Subject: [PATCH 3/6] fix(agents): thread FlushSentinel through Agent.create llmNode types Agent.llmNode now returns ReadableStream, but the agent_v2 hook overrides and AgentHookAdapter still declared the narrower ChatChunk | string union, so passing super.llmNode as the fallback failed to type-check. Widen the override return types and the adapter's fallback/return signatures to include FlushSentinel. Co-authored-by: Cursor --- agents/src/voice/agent_v2.ts | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/agents/src/voice/agent_v2.ts b/agents/src/voice/agent_v2.ts index 47df78b85..2f00402ce 100644 --- a/agents/src/voice/agent_v2.ts +++ b/agents/src/voice/agent_v2.ts @@ -14,6 +14,7 @@ import type { } from '../llm/index.js'; import type { STT, SpeechEvent } from '../stt/index.js'; import type { TTS } from '../tts/index.js'; +import type { FlushSentinel } from '../types.js'; import { readStream, toStream } from '../utils.js'; import type { VAD } from '../vad.js'; import type { Agent, AgentOptions, AgentTask, AgentTaskOptions, ModelSettings } from './agent.js'; @@ -184,7 +185,7 @@ export function createAgentV2( chatCtx: ChatContext, toolCtx: ToolContext, modelSettings: ModelSettings, - ): Promise | null> { + ): Promise | null> { return this.hookAdapter.llmNode(chatCtx, toolCtx, modelSettings, () => super.llmNode(chatCtx, toolCtx, modelSettings), ); @@ -278,7 +279,7 @@ export function createAgentTaskV2( chatCtx: ChatContext, toolCtx: ToolContext, modelSettings: ModelSettings, - ): Promise | null> { + ): Promise | null> { return this.hookAdapter.llmNode(chatCtx, toolCtx, modelSettings, () => super.llmNode(chatCtx, toolCtx, modelSettings), ); @@ -357,8 +358,8 @@ class AgentHookAdapter> { chatCtx: ChatContext, toolCtx: ToolContext, modelSettings: ModelSettings, - fallback: () => Promise | null>, - ): Promise | null> { + fallback: () => Promise | null>, + ): Promise | null> { if (!this.hooks.llmNode) { return fallback(); } From c8b56bd5d43e6f532ba40b69cb9de2ab31df07fb Mon Sep 17 00:00:00 2001 From: "rosetta-livekit-bot[bot]" <282703043+rosetta-livekit-bot[bot]@users.noreply.github.com> Date: Mon, 8 Jun 2026 14:59:00 -0700 Subject: [PATCH 4/6] feat(agents): add beta end call tool (#1474) Co-authored-by: Brian Yin Co-authored-by: rosetta-livekit-bot[bot] <282703043+rosetta-livekit-bot[bot]@users.noreply.github.com> Co-authored-by: u9g --- .changeset/port-end-call-tool.md | 5 + agents/src/beta/index.ts | 7 + agents/src/beta/tools/end_call.ts | 188 ++++++++++++++++++++++++ agents/src/beta/tools/index.ts | 10 ++ agents/src/llm/index.ts | 1 + agents/src/llm/tool_context.test.ts | 49 +++--- agents/src/llm/tool_context.ts | 79 ++++++---- agents/src/voice/agent_activity.test.ts | 97 +++++++++++- agents/src/voice/agent_activity.ts | 41 +++++- plugins/google/src/aiplatform_llm.ts | 5 +- plugins/openai/src/responses/llm.ts | 1 + 11 files changed, 420 insertions(+), 63 deletions(-) create mode 100644 .changeset/port-end-call-tool.md create mode 100644 agents/src/beta/tools/end_call.ts create mode 100644 agents/src/beta/tools/index.ts diff --git a/.changeset/port-end-call-tool.md b/.changeset/port-end-call-tool.md new file mode 100644 index 000000000..f1646f943 --- /dev/null +++ b/.changeset/port-end-call-tool.md @@ -0,0 +1,5 @@ +--- +'@livekit/agents': minor +--- + +Add beta EndCallTool for ending calls from agent tools diff --git a/agents/src/beta/index.ts b/agents/src/beta/index.ts index 98ac382e9..d5141be55 100644 --- a/agents/src/beta/index.ts +++ b/agents/src/beta/index.ts @@ -12,3 +12,10 @@ export { type WarmTransferTaskOptions, } from './workflows/index.js'; export { Instructions } from '../llm/index.js'; +export { + END_CALL_DESCRIPTION, + createEndCallTool, + type EndCallToolCalledEvent, + type EndCallToolCompletedEvent, + type EndCallToolOptions, +} from './tools/index.js'; diff --git a/agents/src/beta/tools/end_call.ts b/agents/src/beta/tools/end_call.ts new file mode 100644 index 000000000..34eb63f11 --- /dev/null +++ b/agents/src/beta/tools/end_call.ts @@ -0,0 +1,188 @@ +// SPDX-FileCopyrightText: 2026 LiveKit, Inc. +// +// SPDX-License-Identifier: Apache-2.0 +import { type EventEmitter, once } from 'node:events'; +import { setTimeout as waitFor } from 'node:timers/promises'; +import { getJobContext } from '../../job.js'; +import { + RealtimeModel, + type ToolCalledEvent, + type ToolCompletedEvent, + Toolset, + tool, +} from '../../llm/index.js'; +import { log } from '../../log.js'; +import type { AgentSession, AgentSessionCallbacks } from '../../voice/agent_session.js'; +import { AgentSessionEventTypes } from '../../voice/events.js'; +import type { UnknownUserData } from '../../voice/run_context.js'; + +/** How long to wait for the agent's goodbye reply to play out before forcing shutdown. */ +const END_CALL_REPLY_TIMEOUT = 5000; + +/** + * `events.once` typed against {@link AgentSessionCallbacks}, resolving to the event payload (e.g. + * `CloseEvent`) — every session event is single-arg — instead of the raw `any[]` tuple. The cast + * `events.once` forces is confined here. + * + * If `signal` aborts before the event fires, `events.once` rejects with an AbortError. Since these + * waits are fire-and-forget (or race losers), an uncaught rejection would surface as an unhandled + * rejection — so we absorb the abort and resolve to `undefined`, letting callers treat "aborted" + * as a normal (typed) outcome. Other rejections (e.g. an emitter `error`) still propagate. + */ +function onceEvent( + // eslint-disable-next-line @typescript-eslint/no-explicit-any -- callbacks don't depend on UserData + session: AgentSession, + event: E, + options?: { signal?: AbortSignal }, +): Promise[0] | undefined> { + return ( + once(session as unknown as EventEmitter, event, options) as Promise< + Parameters + > + ).then( + ([payload]) => payload, + (err) => { + if (options?.signal?.aborted) return undefined; + throw err; + }, + ); +} + +export const END_CALL_DESCRIPTION = ` +Ends the current call and disconnects immediately. + +Call when: +- The user clearly indicates they are done (e.g., "that's all, bye"). + +Do not call when: +- The user asks to pause, hold, or transfer. +- Intent is unclear. + +This is the final action the agent can take. +Once called, no further interaction is possible with the user. +Don't generate any other text or response when the tool is called. +`; + +export type EndCallToolCalledEvent = ToolCalledEvent; + +export type EndCallToolCompletedEvent = ToolCompletedEvent; + +export type EndCallToolOptions = { + /** Additional description to add to the end call tool. */ + extraDescription?: string; + /** + * Whether to delete the room when the user ends the call. + * Deleting the room disconnects all remote users, including SIP callers. + */ + deleteRoom?: boolean; + /** Tool output to the LLM for generating the tool response. */ + endInstructions?: string | null; + /** Callback to call when the tool is called. */ + onToolCalled?: (event: EndCallToolCalledEvent) => Promise | void; + /** Callback to call when the tool is completed. */ + onToolCompleted?: (event: EndCallToolCompletedEvent) => Promise | void; +}; + +/** + * Allows the agent to end the call and disconnect from the room. + */ +export function createEndCallTool({ + extraDescription = '', + deleteRoom = true, + endInstructions = 'say goodbye to the user', + onToolCalled, + onToolCompleted, +}: EndCallToolOptions = {}): Toolset { + // For a realtime LLM that generates the goodbye reply itself, wait for that reply to play out + // (bounded by END_CALL_REPLY_TIMEOUT) before shutting down. `signal` is aborted when the call + // ends or the toolset is torn down, which cancels whichever of the two races is still pending. + const delayedSessionShutdown = async ( + session: AgentSession, + signal: AbortSignal, + ): Promise => { + const speech = onceEvent(session, AgentSessionEventTypes.SpeechCreated, { signal }).then( + (event) => event?.speechHandle, + ); + const timeout = waitFor(END_CALL_REPLY_TIMEOUT, 'timeout' as const, { signal }).catch( + () => undefined, + ); + + const winner = await Promise.race([speech, timeout]); + if (signal.aborted) return; // session already closed or toolset torn down + + if (winner === 'timeout') { + log().warn('tool reply timed out, shutting down session'); + session.shutdown(); + } else if (winner) { + await winner.waitForPlayout(); + session.shutdown(); + } + }; + + return Toolset.create({ + id: 'end_call', + tools: [ + tool({ + name: 'end_call', + description: `${END_CALL_DESCRIPTION}\n${extraDescription}`, + execute: async (_args, { ctx, abortSignal }) => { + log().debug('end_call tool called'); + const session = ctx.session; + const llm = session.currentAgent.getActivityOrThrow().llm; + + // Lifetime of this invocation: aborts when the session closes, and also when the tool + // call itself is aborted. All listeners/timers below are scoped to it. + const controller = new AbortController(); + const signal = abortSignal + ? AbortSignal.any([abortSignal, controller.signal]) + : controller.signal; + + void onceEvent(session, AgentSessionEventTypes.Close, { signal }).then((event) => { + if (!event) return; // signal aborted before close fired + controller.abort(); // stop the delayed-shutdown race + + const jobCtx = getJobContext(false); + if (!jobCtx) return; + + if (deleteRoom) { + jobCtx.addShutdownCallback(async () => { + log().info('deleting the room because the user ended the call'); + await jobCtx.deleteRoom(); + }); + } + + jobCtx.shutdown(String(event.reason)); + }); + + ctx.speechHandle.addDoneCallback(() => { + if (!(llm instanceof RealtimeModel) || !llm.capabilities.autoToolReplyGeneration) { + session.shutdown(); + return; + } + + void delayedSessionShutdown(session, signal).catch((error) => + log().error({ error }, 'error during delayed session shutdown'), + ); + }); + + if (onToolCalled) { + await onToolCalled({ ctx, arguments: {} }); + } + + const completedEvent = { + ctx, + output: + endInstructions === null + ? undefined + : ({ type: 'output', value: endInstructions } as const), + }; + if (onToolCompleted) { + await onToolCompleted(completedEvent); + } + + return endInstructions ?? undefined; + }, + }), + ], + }); +} diff --git a/agents/src/beta/tools/index.ts b/agents/src/beta/tools/index.ts new file mode 100644 index 000000000..8ef18b993 --- /dev/null +++ b/agents/src/beta/tools/index.ts @@ -0,0 +1,10 @@ +// SPDX-FileCopyrightText: 2026 LiveKit, Inc. +// +// SPDX-License-Identifier: Apache-2.0 +export { + END_CALL_DESCRIPTION, + createEndCallTool, + type EndCallToolCalledEvent, + type EndCallToolCompletedEvent, + type EndCallToolOptions, +} from './end_call.js'; diff --git a/agents/src/llm/index.ts b/agents/src/llm/index.ts index 03515c5ce..400849be4 100644 --- a/agents/src/llm/index.ts +++ b/agents/src/llm/index.ts @@ -25,6 +25,7 @@ export { type ToolContextEntry, type ToolCtxInput, type ToolOptions, + type ToolsetContext, type ToolsetCreateOptions, type ToolType, } from './tool_context.js'; diff --git a/agents/src/llm/tool_context.test.ts b/agents/src/llm/tool_context.test.ts index c42380478..a18946eb2 100644 --- a/agents/src/llm/tool_context.test.ts +++ b/agents/src/llm/tool_context.test.ts @@ -11,6 +11,7 @@ import { ToolContext, type ToolOptions, Toolset, + type ToolsetContext, tool, } from './tool_context.js'; import { createToolOptions, oaiParams } from './utils.js'; @@ -617,9 +618,13 @@ describe('Toolset', () => { expect(ts.tools).toEqual([a, b]); }); + const fakeToolsetContext = ( + updateTools: (tools: readonly Tool[]) => void = () => {}, + ): ToolsetContext => ({ updateTools }); + it('default setup and aclose are no-ops', async () => { const ts = new Toolset({ id: 'noop', tools: [] }); - await expect(ts.setup()).resolves.toBeUndefined(); + await expect(ts.setup(fakeToolsetContext())).resolves.toBeUndefined(); await expect(ts.aclose()).resolves.toBeUndefined(); }); @@ -635,20 +640,17 @@ describe('Toolset', () => { } const ts = new Recording({ id: 'rec', tools: [] }); - await ts.setup(); + await ts.setup(fakeToolsetContext()); await ts.aclose(); expect(events).toEqual(['setup:rec', 'close:rec']); }); - it('Toolset.create() composes lifecycle callbacks without subclassing', async () => { + it('Toolset.create() resolves a static tools list eagerly and composes aclose', async () => { const a = makeFn('a'); const events: string[] = []; const ts = Toolset.create({ id: 'composed', tools: [a], - setup: async () => { - events.push('setup'); - }, aclose: async () => { events.push('close'); }, @@ -656,37 +658,38 @@ describe('Toolset', () => { expect(ts).toBeInstanceOf(Toolset); expect(ts.id).toBe('composed'); - expect(ts.tools).toEqual([a]); + expect(ts.tools).toEqual([a]); // static tools available before activation - await ts.setup(); + await ts.setup(fakeToolsetContext()); await ts.aclose(); - expect(events).toEqual(['setup', 'close']); + expect(events).toEqual(['close']); }); - it('Toolset.create() defaults setup and aclose to no-ops when callbacks are omitted', async () => { + it('Toolset.create() defaults aclose to a no-op when omitted', async () => { const ts = Toolset.create({ id: 'bare', tools: [] }); - await expect(ts.setup()).resolves.toBeUndefined(); + await expect(ts.setup(fakeToolsetContext())).resolves.toBeUndefined(); await expect(ts.aclose()).resolves.toBeUndefined(); }); - it('Toolset.create() accepts a tools thunk, re-evaluated on every access (dynamic)', () => { + it('lets setup push tools after activation via ctx.updateTools', async () => { const a = makeFn('a'); const b = makeFn('b'); - const current: Tool[] = [a]; - let calls = 0; + let push!: (tools: readonly Tool[]) => void; const ts = Toolset.create({ - id: 'dynamic', - tools: () => { - calls += 1; - return current; + id: 'mcp', + setup: async ({ updateTools }) => { + push = updateTools; }, + tools: [], }); - // Each access re-invokes the thunk so the toolset reflects the current source-of-truth. - expect(ts.tools).toEqual([a]); - expect(calls).toBe(1); - current.push(b); + + // Mimic the runtime: ctx.updateTools writes the toolset's current tools. + await ts.setup(fakeToolsetContext((tools) => ts._setTools(tools))); + expect(ts.tools).toEqual([]); + + // A dynamic source (e.g. an MCP server) pushes its tools after connecting. + push([a, b]); expect(ts.tools).toEqual([a, b]); - expect(calls).toBe(2); }); it('is flattened into a ToolContext: function tools merged, toolset tracked', () => { diff --git a/agents/src/llm/tool_context.ts b/agents/src/llm/tool_context.ts index 85ed81223..4cc087352 100644 --- a/agents/src/llm/tool_context.ts +++ b/agents/src/llm/tool_context.ts @@ -211,6 +211,15 @@ export interface ToolCompletedEvent { output?: { type: 'output'; value: unknown } | { type: 'error'; value: Error }; } +/** Context passed to a {@link Toolset}'s `setup` hook when it activates. */ +export interface ToolsetContext { + /** + * Replace the toolset's tools. Useful for dynamic sources + * (e.g. an MCP server) whose tools are discovered after `setup` or change at runtime. + */ + updateTools(tools: readonly Tool[]): void; +} + /** * Function tools of a `ToolContext`, sorted by name for deterministic provider payloads. * Provider tools are intentionally excluded — callers that need them iterate `flatten()`. @@ -239,7 +248,7 @@ export function sortedToolNames(toolCtx: ToolContext | undefined): string[] { export class Toolset { readonly #id: string; - readonly #tools: Tool[]; + #tools: readonly Tool[]; readonly [TOOLSET_SYMBOL] = true as const; @@ -249,9 +258,10 @@ export class Toolset { } /** - * Compose a `Toolset` with inline `setup` / `aclose` hooks instead of subclassing. `tools` - * may also be a thunk that is re-evaluated on every `.tools` access, so the toolset can - * expose a dynamic list that changes after `setup()` runs. + * For when your tools share something that needs setup or cleanup, like a DB pool, an open MCP + * client, or listeners on a shared bus. `setup` runs once at activation, `aclose` once at + * teardown. If the tool list itself is dynamic (e.g. an MCP server), push it from `setup` via + * {@link ToolsetContext.updateTools}. * * @example Static tool list with a shared backing resource * ```ts @@ -260,20 +270,26 @@ export class Toolset { * return Toolset.create({ * id: 'postgres', * tools: [queryOrders, queryCustomers], - * setup: () => pool.connect(), * aclose: () => pool.end(), * }); * } * ``` * - * @example Dynamic tool list + * @example Dynamic tool list bound to an external source * ```ts * function createMcpToolset(url: string): Toolset { * const client = new MCPClient({ url }); * return Toolset.create({ * id: 'mcp_remote', - * tools: () => client.getTools(), - * setup: () => client.connect(), + * // setup connects and wires listeners that push the server's tools whenever they change; + * // the runtime re-advertises without re-running anything. + * setup: async ({ updateTools }) => { + * const sync = async () => updateTools(await client.listTools()); + * client.on('connect', sync); + * client.on('tool_list_changed', sync); + * await client.connect(); + * }, + * tools: [], * aclose: () => client.disconnect(), * }); * } @@ -291,48 +307,49 @@ export class Toolset { return this.#tools; } - async setup(): Promise {} + /** + * Replace the toolset's current tools. Backs {@link ToolsetContext.updateTools}; the runtime + * re-flattens and re-advertises after calling it. + * + * @internal + */ + _setTools(tools: readonly Tool[]): void { + this.#tools = [...tools]; + } + + async setup(_ctx: ToolsetContext): Promise {} async aclose(): Promise {} } -/** Options accepted by `Toolset.create()` — id + tools plus optional lifecycle hooks. */ +/** Options accepted by `Toolset.create()` — id + tools plus optional setup/teardown hooks. */ export interface ToolsetCreateOptions { id: string; /** - * Either a static list of tools, or a thunk re-evaluated on every `tools` access — useful - * when the underlying source (e.g. an MCP discovery loop) can produce a dynamic tool list. + * One-time async initialization run when the toolset activates — e.g. connecting to a server + * and wiring listeners. Push a changed tool list via {@link ToolsetContext.updateTools}. */ - tools: readonly Tool[] | (() => readonly Tool[]); - /** Invoked when the toolset becomes active in an `AgentActivity`. */ - setup?: () => Promise; - /** Invoked when the toolset is being torn down. */ + setup?: (ctx: ToolsetContext) => Promise; + /** The toolset's initial tools. */ + tools: readonly Tool[]; + /** Invoked when the toolset is being torn down. Release awaitable resources here. */ aclose?: () => Promise; } /** Backing implementation of `Toolset.create()`. Kept private so callers go through the factory. */ class ToolsetFactory extends Toolset { - readonly #toolsSource: readonly Tool[] | (() => readonly Tool[]); - - readonly #setupFn?: () => Promise; + readonly #setupFn?: (ctx: ToolsetContext) => Promise; readonly #acloseFn?: () => Promise; - constructor({ id, tools, setup, aclose }: ToolsetCreateOptions) { - // Pass [] to super and override the `tools` getter so a thunk can be re-evaluated on - // every access (lets callers expose a dynamic tool list). - super({ id, tools: [] }); - this.#toolsSource = tools; + constructor({ id, setup, tools, aclose }: ToolsetCreateOptions) { + super({ id, tools }); this.#setupFn = setup; this.#acloseFn = aclose; } - override get tools(): readonly Tool[] { - return typeof this.#toolsSource === 'function' ? this.#toolsSource() : this.#toolsSource; - } - - override async setup(): Promise { - if (this.#setupFn) await this.#setupFn(); + override async setup(ctx: ToolsetContext): Promise { + if (this.#setupFn) await this.#setupFn(ctx); } override async aclose(): Promise { @@ -389,7 +406,7 @@ export class ToolContext { /** A copy of all provider tools in the tool context, including those in tool sets. */ get providerTools(): ProviderTool[] { - return this._providerTools; + return [...this._providerTools]; } /** A copy of all toolsets registered in the context. */ diff --git a/agents/src/voice/agent_activity.test.ts b/agents/src/voice/agent_activity.test.ts index 2e76d9d71..da8bab4ba 100644 --- a/agents/src/voice/agent_activity.test.ts +++ b/agents/src/voice/agent_activity.test.ts @@ -16,9 +16,9 @@ */ import { Heap } from 'heap-js'; import { describe, expect, it, vi } from 'vitest'; -import { ChatContext } from '../llm/chat_context.js'; +import { AgentConfigUpdate, ChatContext } from '../llm/chat_context.js'; import { LLM, type LLMStream } from '../llm/llm.js'; -import { ToolContext } from '../llm/tool_context.js'; +import { type Tool, ToolContext, Toolset, tool } from '../llm/tool_context.js'; import { Future } from '../utils.js'; import { AgentActivity } from './agent_activity.js'; import type { PreemptiveGenerationInfo } from './audio_recognition.js'; @@ -409,3 +409,96 @@ describe('AgentActivity - onPreemptiveGeneration guards', () => { expect(cancelPreemptiveGeneration).not.toHaveBeenCalled(); }); }); + +/** + * Regression test for the dynamic-toolset push path. + * + * When an already-activated toolset swaps its tools at runtime (e.g. an MCP server pushes a new + * tool list via `ToolsetContext.updateTools`), `setupToolsetList`'s wiring must (1) invoke + * `onToolsetToolsChanged`, which now funnels through `updateTools`, and (2) record an + * `AgentConfigUpdate` in the agent chat context + session history so a non-realtime pipeline's + * chat context reflects the new tool set on the next turn. + */ +class FakeToolsetLLM extends LLM { + label(): string { + return 'fake.toolset.LLM'; + } + chat(): LLMStream { + throw new Error('not used in these tests'); + } +} + +describe('AgentActivity - onToolsetToolsChanged (dynamic toolset push)', () => { + const makeFn = (name: string) => + tool({ name, description: `${name} tool`, execute: async () => name }); + + function buildToolsetActivity(toolset: Toolset) { + const history = new ChatContext(); + const fakeActivity = { + _toolsetsSetup: true, + realtimeSession: undefined, + llm: new FakeToolsetLLM(), + agent: { + _toolCtx: new ToolContext([toolset]), + _chatCtx: new ChatContext(), + }, + agentSession: { history }, + updateChatCtx: vi.fn(async () => {}), + logger: { info() {}, debug() {}, warn() {}, error() {} }, + }; + Object.setPrototypeOf(fakeActivity, AgentActivity.prototype); + return { fakeActivity, history }; + } + + it('fires onToolsetToolsChanged on a dynamic push and records an AgentConfigUpdate', async () => { + const toolA = makeFn('toolA'); + const toolB = makeFn('toolB'); + + // Capture the wired ctx.updateTools the framework hands the toolset during setup. + let pushTools!: (tools: readonly Tool[]) => void; + const toolset = Toolset.create({ + id: 'dynamic', + tools: [toolA], + setup: async ({ updateTools }) => { + pushTools = updateTools; + }, + }); + + const { fakeActivity, history } = buildToolsetActivity(toolset); + + const changedSpy = vi.spyOn( + AgentActivity.prototype as unknown as Record<'onToolsetToolsChanged', () => Promise>, + 'onToolsetToolsChanged', + ); + + // Activate the toolset through the real path so it captures the push channel. + const setupToolsetList = (AgentActivity.prototype as Record) + .setupToolsetList as (this: unknown, toolsets: readonly Toolset[]) => Promise; + await setupToolsetList.call(fakeActivity, [toolset]); + + expect(changedSpy).not.toHaveBeenCalled(); + + // (1) A dynamic push swaps the toolset's tools — the wiring must invoke onToolsetToolsChanged. + pushTools([toolA, toolB]); + expect(changedSpy).toHaveBeenCalledTimes(1); + await changedSpy.mock.results[0]!.value; + + // (2) An AgentConfigUpdate naming the added tool lands in the session history. + const updates = history.items.filter( + (i): i is AgentConfigUpdate => i instanceof AgentConfigUpdate, + ); + expect(updates).toHaveLength(1); + expect(updates[0]!.toolsAdded).toContain('toolB'); + expect(updates[0]!.toolsRemoved ?? []).not.toContain('toolA'); + + // The refreshed tool context advertises the new tool to the next turn, and the non-realtime + // pipeline's chat context was refreshed via updateChatCtx. + expect(Object.keys(fakeActivity.agent._toolCtx.functionTools).sort()).toEqual([ + 'toolA', + 'toolB', + ]); + expect(fakeActivity.updateChatCtx).toHaveBeenCalledTimes(1); + + changedSpy.mockRestore(); + }); +}); diff --git a/agents/src/voice/agent_activity.ts b/agents/src/voice/agent_activity.ts index fd7f5c806..66c55a747 100644 --- a/agents/src/voice/agent_activity.ts +++ b/agents/src/voice/agent_activity.ts @@ -23,6 +23,7 @@ import { instructionsEqual, renderInstructions, } from '../llm/chat_context.js'; +import type { Toolset } from '../llm/index.js'; import { type ChatItem, type FunctionCall, @@ -41,7 +42,6 @@ import { ToolContext, type ToolContextEntry, ToolFlag, - Toolset, isFunctionTool, isToolset, } from '../llm/index.js'; @@ -786,14 +786,18 @@ export class AgentActivity implements RecognitionHooks { const oldToolNames = new Set(Object.keys(oldToolCtx.functionTools)); const oldToolsets = oldToolCtx.toolsets; const newToolCtx = new ToolContext(tools); - const newToolNames = new Set(Object.keys(newToolCtx.functionTools)); const newToolsets = newToolCtx.toolsets; - const toolsAdded = [...newToolNames].filter((name) => !oldToolNames.has(name)); - const toolsRemoved = [...oldToolNames].filter((name) => !newToolNames.has(name)); const addedToolsets = newToolsets.filter((ts) => !oldToolsets.includes(ts)); const removedToolsets = oldToolsets.filter((ts) => !newToolsets.includes(ts)); + // Resolve added factory toolsets before re-flattening, so their tools are included in the + // advertised set (newToolNames is computed below, after resolution). await this.setupToolsetList(addedToolsets); + newToolCtx.updateTools(newToolCtx.tools); + const newToolNames = new Set(Object.keys(newToolCtx.functionTools)); + const toolsAdded = [...newToolNames].filter((name) => !oldToolNames.has(name)); + const toolsRemoved = [...oldToolNames].filter((name) => !newToolNames.has(name)); + this.agent._toolCtx = newToolCtx; await this.closeToolsetList(removedToolsets); @@ -4185,6 +4189,8 @@ export class AgentActivity implements RecognitionHooks { if (this._toolsetsSetup) return; this._toolsetsSetup = true; await this.setupToolsetList(this.agent.toolCtx.toolsets); + // Re-flatten now that any factory toolsets have resolved their tools, so they're advertised. + this.agent._toolCtx.updateTools(this.agent._toolCtx.tools); } private async closeToolsets(): Promise { @@ -4193,8 +4199,33 @@ export class AgentActivity implements RecognitionHooks { await this.closeToolsetList(this.agent.toolCtx.toolsets); } + /** + * Refresh the agent's tool context after a dynamic toolset pushed a new tool list (via + * `ToolsetContext.updateTools`), routing through updateTools so history, chat context, and the + * realtime session stay in sync. The next LLM turn picks up the new tools automatically. + */ + private async onToolsetToolsChanged(): Promise { + if (!this._toolsetsSetup) return; + const current = this.agent._toolCtx; + if (new ToolContext(current.tools).equals(current)) return; + // Same toolset entries, so updateTools' setup/close steps are no-ops (no re-entrancy here). + await this.updateTools(current.tools); + } + private async setupToolsetList(toolsets: readonly Toolset[]): Promise { - const outputs = await Promise.allSettled(toolsets.map((ts) => ts.setup())); + const outputs = await Promise.allSettled( + toolsets.map((ts) => + ts.setup({ + // A dynamic toolset pushes a changed tool list here; re-flatten and re-advertise it. + updateTools: (tools) => { + ts._setTools(tools); + void this.onToolsetToolsChanged().catch((error) => + this.logger.error({ error }, 'error re-advertising toolset tools'), + ); + }, + }), + ), + ); for (const output of outputs) { if (output.status === 'rejected') { this.logger.error({ error: output.reason }, 'error setting up toolset'); diff --git a/plugins/google/src/aiplatform_llm.ts b/plugins/google/src/aiplatform_llm.ts index 8ac75d99c..19e530eed 100644 --- a/plugins/google/src/aiplatform_llm.ts +++ b/plugins/google/src/aiplatform_llm.ts @@ -165,19 +165,20 @@ export class AIPlatformLLM extends llm.LLM { chat({ chatCtx, - toolCtx, + toolCtx: toolCtxInput, connOptions = DEFAULT_API_CONNECT_OPTIONS, parallelToolCalls, toolChoice, extraKwargs, }: { chatCtx: llm.ChatContext; - toolCtx?: llm.ToolContext; + toolCtx?: llm.ToolCtxInput; connOptions?: APIConnectOptions; parallelToolCalls?: boolean; toolChoice?: llm.ToolChoice; extraKwargs?: Record; }): inference.LLMStream { + const toolCtx = llm.toToolContext(toolCtxInput); const extras: Record = { ...extraKwargs }; if (this.#opts.temperature !== undefined) { diff --git a/plugins/openai/src/responses/llm.ts b/plugins/openai/src/responses/llm.ts index 67fde0e91..aaf68a5ed 100644 --- a/plugins/openai/src/responses/llm.ts +++ b/plugins/openai/src/responses/llm.ts @@ -187,6 +187,7 @@ class ResponsesHttpLLMStream extends llm.LLMStream { 'openai.responses', )) as OpenAI.Responses.ResponseInputItem[]; + // TODO: support provider tools in the Responses schema. const tools = this.toolCtx ? toResponsesTools(this.toolCtx, this.strictToolSchema) : undefined; From f089da479135e750cb44e1fb4f4a288789e63f4e Mon Sep 17 00:00:00 2001 From: Toubat Date: Mon, 8 Jun 2026 16:20:39 -0700 Subject: [PATCH 5/6] arden end-call shutdown and tool guards Catch end-call close listener errors to avoid unhandled rejections during shutdown, and make public tool type guards return false for null inputs. --- agents/src/beta/tools/end_call.ts | 39 ++++++++++++----------------- agents/src/llm/tool_context.test.ts | 2 +- agents/src/llm/tool_context.ts | 16 +++++------- 3 files changed, 23 insertions(+), 34 deletions(-) diff --git a/agents/src/beta/tools/end_call.ts b/agents/src/beta/tools/end_call.ts index 34eb63f11..962deda36 100644 --- a/agents/src/beta/tools/end_call.ts +++ b/agents/src/beta/tools/end_call.ts @@ -19,16 +19,7 @@ import type { UnknownUserData } from '../../voice/run_context.js'; /** How long to wait for the agent's goodbye reply to play out before forcing shutdown. */ const END_CALL_REPLY_TIMEOUT = 5000; -/** - * `events.once` typed against {@link AgentSessionCallbacks}, resolving to the event payload (e.g. - * `CloseEvent`) — every session event is single-arg — instead of the raw `any[]` tuple. The cast - * `events.once` forces is confined here. - * - * If `signal` aborts before the event fires, `events.once` rejects with an AbortError. Since these - * waits are fire-and-forget (or race losers), an uncaught rejection would surface as an unhandled - * rejection — so we absorb the abort and resolve to `undefined`, letting callers treat "aborted" - * as a normal (typed) outcome. Other rejections (e.g. an emitter `error`) still propagate. - */ +/** Typed wrapper around `events.once`; abort resolves to `undefined`, other errors propagate. */ function onceEvent( // eslint-disable-next-line @typescript-eslint/no-explicit-any -- callbacks don't depend on UserData session: AgentSession, @@ -137,22 +128,24 @@ export function createEndCallTool({ ? AbortSignal.any([abortSignal, controller.signal]) : controller.signal; - void onceEvent(session, AgentSessionEventTypes.Close, { signal }).then((event) => { - if (!event) return; // signal aborted before close fired - controller.abort(); // stop the delayed-shutdown race + void onceEvent(session, AgentSessionEventTypes.Close, { signal }) + .then((event) => { + if (!event) return; // signal aborted before close fired + controller.abort(); // stop the delayed-shutdown race - const jobCtx = getJobContext(false); - if (!jobCtx) return; + const jobCtx = getJobContext(false); + if (!jobCtx) return; - if (deleteRoom) { - jobCtx.addShutdownCallback(async () => { - log().info('deleting the room because the user ended the call'); - await jobCtx.deleteRoom(); - }); - } + if (deleteRoom) { + jobCtx.addShutdownCallback(async () => { + log().info('deleting the room because the user ended the call'); + await jobCtx.deleteRoom(); + }); + } - jobCtx.shutdown(String(event.reason)); - }); + jobCtx.shutdown(String(event.reason)); + }) + .catch((error) => log().error({ error }, 'error during end call shutdown')); ctx.speechHandle.addDoneCallback(() => { if (!(llm instanceof RealtimeModel) || !llm.capabilities.autoToolReplyGeneration) { diff --git a/agents/src/llm/tool_context.test.ts b/agents/src/llm/tool_context.test.ts index a18946eb2..5f2f80cca 100644 --- a/agents/src/llm/tool_context.test.ts +++ b/agents/src/llm/tool_context.test.ts @@ -631,7 +631,7 @@ describe('Toolset', () => { it('lets subclasses override lifecycle hooks', async () => { const events: string[] = []; class Recording extends Toolset { - override async setup(): Promise { + override async setup(_ctx: ToolsetContext): Promise { events.push(`setup:${this.id}`); } override async aclose(): Promise { diff --git a/agents/src/llm/tool_context.ts b/agents/src/llm/tool_context.ts index 4cc087352..50c4220f2 100644 --- a/agents/src/llm/tool_context.ts +++ b/agents/src/llm/tool_context.ts @@ -608,34 +608,30 @@ export function tool(tool: any): any { // eslint-disable-next-line @typescript-eslint/no-explicit-any export function isTool(tool: any): tool is Tool { - return tool && tool[TOOL_SYMBOL] === true; + return !!tool && tool[TOOL_SYMBOL] === true; } // eslint-disable-next-line @typescript-eslint/no-explicit-any export function isFunctionTool(tool: any): tool is FunctionTool { - const isTool = tool && tool[TOOL_SYMBOL] === true; - const isFunctionTool = tool[FUNCTION_TOOL_SYMBOL] === true; - return isTool && isFunctionTool; + return isTool(tool) && (tool as FunctionTool)[FUNCTION_TOOL_SYMBOL] === true; } // eslint-disable-next-line @typescript-eslint/no-explicit-any export function isProviderTool(tool: any): tool is ProviderTool { - const isTool = tool && tool[TOOL_SYMBOL] === true; - const isProviderTool = tool[PROVIDER_TOOL_SYMBOL] === true; - return isTool && isProviderTool; + return isTool(tool) && (tool as ProviderTool)[PROVIDER_TOOL_SYMBOL] === true; } // eslint-disable-next-line @typescript-eslint/no-explicit-any export function isToolset(value: any): value is Toolset { - return value && value[TOOLSET_SYMBOL] === true; + return !!value && value[TOOLSET_SYMBOL] === true; } // eslint-disable-next-line @typescript-eslint/no-explicit-any export function isToolError(error: any): error is ToolError { - return error && error[TOOL_ERROR_SYMBOL] === true; + return !!error && error[TOOL_ERROR_SYMBOL] === true; } // eslint-disable-next-line @typescript-eslint/no-explicit-any export function isAgentHandoff(handoff: any): handoff is AgentHandoff { - return handoff && handoff[HANDOFF_SYMBOL] === true; + return !!handoff && handoff[HANDOFF_SYMBOL] === true; } From c201f8b3455e0d7bf5b466b33c882fdff826d03e Mon Sep 17 00:00:00 2001 From: "rosetta-livekit-bot[bot]" <282703043+rosetta-livekit-bot[bot]@users.noreply.github.com> Date: Wed, 10 Jun 2026 05:11:53 +0000 Subject: [PATCH 6/6] fix(sarvam): emit speech timing for STT metrics --- .changeset/sarvam-stt-speech-timing.md | 5 + plugins/sarvam/src/stt.ts | 215 ++++++++++++++++++++++--- 2 files changed, 197 insertions(+), 23 deletions(-) create mode 100644 .changeset/sarvam-stt-speech-timing.md diff --git a/.changeset/sarvam-stt-speech-timing.md b/.changeset/sarvam-stt-speech-timing.md new file mode 100644 index 000000000..76907e748 --- /dev/null +++ b/.changeset/sarvam-stt-speech-timing.md @@ -0,0 +1,5 @@ +--- +'@livekit/agents-plugin-sarvam': patch +--- + +Emit Sarvam STT speech timing for streaming metrics. diff --git a/plugins/sarvam/src/stt.ts b/plugins/sarvam/src/stt.ts index 48a6d7abc..62e974b3a 100644 --- a/plugins/sarvam/src/stt.ts +++ b/plugins/sarvam/src/stt.ts @@ -35,6 +35,7 @@ const SARVAM_STT_TRANSLATE_WS_URL = 'wss://api.sarvam.ai/speech-to-text-translat const SAMPLE_RATE = 16000; const NUM_CHANNELS = 1; +const EOS_FALLBACK_TIMEOUT = 1000; // --------------------------------------------------------------------------- // Model-specific option types @@ -409,6 +410,8 @@ interface SarvamWSTranscriptData { transcript?: string; language_code?: string | null; language_probability?: number | null; + speech_start?: number | null; + speech_end?: number | null; timestamps?: Record | null; diarized_transcript?: Record | null; metrics?: { @@ -423,6 +426,7 @@ interface SarvamWSEventData { timestamp?: string; signal_type?: 'START_SPEECH' | 'END_SPEECH'; occured_at?: number; + request_id?: string; } /** type: "error" — server sends data with message and code fields */ @@ -554,6 +558,15 @@ export class SpeechStream extends stt.SpeechStream { #speaking = false; #resetWS = new Future(); #requestId = ''; + #audioPosition = 0; + #utteranceStartAudioPos = 0; + #utteranceSpeechEndAudioPos?: number; + #utteranceSpeechEndWallTime?: number; + #pendingFinalData?: SarvamWSTranscriptData; + #pendingEos = false; + #eosFallbackTimer?: ReturnType; + #finalReceivedForUtterance = false; + #eosEmittedForUtterance = false; label = 'sarvam.SpeechStream'; constructor(sttInstance: STT, opts: ResolvedSTTOptions, connOptions?: APIConnectOptions) { @@ -583,6 +596,136 @@ export class SpeechStream extends stt.SpeechStream { this.#resetWS.resolve(); } + #maybeSetRequestId(requestId: string | undefined) { + if (requestId) { + this.#requestId = requestId; + } + } + + #positiveTime(value: unknown): number | undefined { + if (typeof value !== 'number' || !Number.isFinite(value) || value <= 0) { + return undefined; + } + return value; + } + + #resetUtteranceState() { + this.#cancelEosFallback(); + this.#pendingFinalData = undefined; + this.#pendingEos = false; + this.#utteranceStartAudioPos = this.#audioPosition; + this.#utteranceSpeechEndAudioPos = undefined; + this.#utteranceSpeechEndWallTime = undefined; + this.#finalReceivedForUtterance = false; + this.#eosEmittedForUtterance = false; + } + + #cancelEosFallback() { + if (this.#eosFallbackTimer !== undefined) { + clearTimeout(this.#eosFallbackTimer); + this.#eosFallbackTimer = undefined; + } + } + + #sendFinalTranscript( + transcriptData: SarvamWSTranscriptData, + opts: { requireEndTime?: boolean } = {}, + ): boolean { + const transcript = transcriptData.transcript ?? ''; + if (!transcript) return false; + + const startTime = + this.#positiveTime(transcriptData.speech_start) ?? this.#utteranceStartAudioPos; + let endTime = + this.#positiveTime(transcriptData.speech_end) ?? + this.#utteranceSpeechEndAudioPos ?? + this.#positiveTime(transcriptData.metrics?.audio_duration); + if (endTime === undefined) { + if (opts.requireEndTime) return false; + endTime = 0; + } + + if (!this.queue.closed) { + this.queue.put({ + type: stt.SpeechEventType.FINAL_TRANSCRIPT, + requestId: transcriptData.request_id ?? this.#requestId, + alternatives: [ + { + text: transcript, + language: normalizeLanguage( + transcriptData.language_code ?? this.#opts.languageCode ?? 'unknown', + ), + startTime, + endTime, + confidence: extractConfidence(transcriptData, this.#logger), + }, + ], + }); + } + return true; + } + + #tryCommitUtterance() { + if ( + this.#pendingFinalData === undefined || + this.#utteranceSpeechEndAudioPos === undefined || + this.#eosEmittedForUtterance + ) { + return; + } + + const committedData = this.#pendingFinalData; + if (this.#sendFinalTranscript(committedData, { requireEndTime: true })) { + this.#logger.debug( + `Sarvam STT utterance committed: end_time=${this.#utteranceSpeechEndAudioPos}`, + ); + this.#emitEndOfSpeech(); + this.#pendingFinalData = undefined; + } + } + + #emitEndOfSpeech() { + if (this.#eosEmittedForUtterance) return; + + this.#cancelEosFallback(); + const alternatives: [stt.SpeechData, ...stt.SpeechData[]] | undefined = + this.#utteranceSpeechEndAudioPos !== undefined + ? [ + { + text: '', + language: normalizeLanguage(this.#opts.languageCode ?? 'unknown'), + startTime: this.#utteranceStartAudioPos, + endTime: this.#utteranceSpeechEndAudioPos, + confidence: 0, + metadata: + this.#utteranceSpeechEndWallTime !== undefined + ? { speechEndWallTime: this.#utteranceSpeechEndWallTime } + : undefined, + }, + ] + : undefined; + + if (!this.queue.closed) { + this.queue.put({ + type: stt.SpeechEventType.END_OF_SPEECH, + requestId: this.#requestId, + alternatives, + }); + } + this.#eosEmittedForUtterance = true; + this.#pendingEos = false; + } + + #scheduleEosFallback() { + this.#cancelEosFallback(); + this.#eosFallbackTimer = setTimeout(() => { + this.#eosFallbackTimer = undefined; + if (this.#pendingEos && !this.#eosEmittedForUtterance) { + this.#emitEndOfSpeech(); + } + }, EOS_FALLBACK_TIMEOUT); + } + protected async run() { const maxRetry = 32; let retries = 0; @@ -717,6 +860,7 @@ export class SpeechStream extends stt.SpeechStream { }, }), ); + this.#audioPosition += frame.samplesPerChannel / SAMPLE_RATE; } } @@ -765,55 +909,79 @@ export class SpeechStream extends stt.SpeechStream { if (msgType === 'events') { const eventData = (json['data'] as SarvamWSEventData | undefined) ?? {}; + this.#maybeSetRequestId(eventData.request_id); const signalType = eventData.signal_type; if (signalType === 'START_SPEECH') { if (!this.#speaking) { + this.#resetUtteranceState(); this.#speaking = true; - putMessage({ type: stt.SpeechEventType.START_OF_SPEECH }); + putMessage({ + type: stt.SpeechEventType.START_OF_SPEECH, + requestId: this.#requestId, + }); } } else if (signalType === 'END_SPEECH') { if (this.#speaking) { this.#speaking = false; - putMessage({ type: stt.SpeechEventType.END_OF_SPEECH }); + this.#utteranceSpeechEndAudioPos = this.#audioPosition; + this.#utteranceSpeechEndWallTime = Date.now(); + this.#pendingEos = true; + this.#tryCommitUtterance(); + if (!this.#eosEmittedForUtterance && this.#pendingFinalData === undefined) { + if (this.#finalReceivedForUtterance) { + this.#emitEndOfSpeech(); + } else { + this.#scheduleEosFallback(); + } + } } } } else if (msgType === 'data') { const td = (json['data'] as SarvamWSTranscriptData | undefined) ?? {}; const transcript = td.transcript ?? ''; - const language = normalizeLanguage( - td.language_code ?? this.#opts.languageCode ?? 'unknown', - ); - const requestId = td.request_id ?? ''; - const confidence = extractConfidence(td, this.#logger); - this.#requestId = requestId; + this.#maybeSetRequestId(td.request_id); // Log metrics when available if (td.metrics) { this.#logger.debug( `Sarvam STT metrics: audio_duration=${td.metrics.audio_duration}s, latency=${td.metrics.processing_latency}s`, ); + if (typeof td.metrics.audio_duration === 'number' && !this.queue.closed) { + putMessage({ + type: stt.SpeechEventType.RECOGNITION_USAGE, + requestId: this.#requestId, + recognitionUsage: { audioDuration: td.metrics.audio_duration }, + }); + } } if (transcript) { - if (!this.#speaking) { + if (!this.#speaking && !this.#pendingEos && !this.#eosEmittedForUtterance) { + this.#resetUtteranceState(); this.#speaking = true; - putMessage({ type: stt.SpeechEventType.START_OF_SPEECH }); + putMessage({ + type: stt.SpeechEventType.START_OF_SPEECH, + requestId: this.#requestId, + }); + } + + const transcriptEndTime = this.#positiveTime(td.speech_end); + if ( + transcriptEndTime !== undefined && + this.#utteranceSpeechEndAudioPos === undefined + ) { + this.#utteranceSpeechEndAudioPos = transcriptEndTime; + this.#utteranceSpeechEndWallTime = Date.now(); } - putMessage({ - type: stt.SpeechEventType.FINAL_TRANSCRIPT, - requestId, - alternatives: [ - { - text: transcript, - language, - startTime: 0, - endTime: td.metrics?.audio_duration ?? 0, - confidence, - }, - ], - }); + if (this.#pendingEos) { + this.#pendingFinalData = td; + this.#finalReceivedForUtterance = true; + this.#tryCommitUtterance(); + } else if (this.#sendFinalTranscript(td)) { + this.#finalReceivedForUtterance = true; + } } } else if (msgType === 'error') { // Server format: { type: "error", data: { message: "...", code: "..." } } @@ -857,6 +1025,7 @@ export class SpeechStream extends stt.SpeechStream { // triggers the ws.once('close') handler inside listenMessage, letting listenTask // exit naturally. On close(), the parent abort signal handles it directly. wsMonitor.cancel(); + this.#cancelEosFallback(); ws.close(); // Suppress unhandled rejection from orphaned listenTask on reconnect listenTask.result.catch(() => {});