mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
* draft checkpoint * checkpoint * address some comments * address more comments * fix conflicts * fix event loop bug * remove mcp agent * use bedrock model for first example * add uuid package instead of crypto * address final review * final fixes (hopefully)
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
dist
|
||||
node_modules
|
||||
package-lock.json
|
||||
@@ -0,0 +1,21 @@
|
||||
{
|
||||
"name": "first-agent",
|
||||
"private": true,
|
||||
"main": "dist/index.js",
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
"clean": "rm -rf dist node_modules package-lock.json",
|
||||
"build": "tsc",
|
||||
"start": "tsc && node dist/index.js"
|
||||
},
|
||||
"workspaces": [
|
||||
"../../"
|
||||
],
|
||||
"dependencies": {
|
||||
"@strands-agents/sdk": "*"
|
||||
},
|
||||
"devDependencies": {
|
||||
"@types/node": "^20.0.0",
|
||||
"typescript": "^5.5.0"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
import { type Tool, type ToolResult, type ToolContext, Agent, BedrockModel } from '@strands-agents/sdk'
|
||||
|
||||
// Define the shape of the expected input
|
||||
type WeatherToolInput = {
|
||||
location: string
|
||||
}
|
||||
|
||||
// Type Guard: A function that performs a runtime check and informs the TS compiler.
|
||||
function isValidInput(input: any): input is WeatherToolInput {
|
||||
return input && typeof input.location === 'string'
|
||||
}
|
||||
|
||||
class WeatherTool implements Tool {
|
||||
name = 'get_weather'
|
||||
description = 'Get the current weather for a specific location.'
|
||||
|
||||
toolSpec = {
|
||||
name: this.name,
|
||||
description: this.description,
|
||||
inputSchema: {
|
||||
type: 'object' as const,
|
||||
properties: {
|
||||
location: {
|
||||
type: 'string' as const,
|
||||
description: 'The city and state, e.g., San Francisco, CA',
|
||||
},
|
||||
},
|
||||
required: ['location'],
|
||||
},
|
||||
}
|
||||
|
||||
async *stream(context: ToolContext): AsyncGenerator<never, ToolResult, unknown> {
|
||||
const input = context.toolUse.input
|
||||
|
||||
// Use the type guard for validation
|
||||
if (!isValidInput(input)) {
|
||||
throw new Error('Tool input must be an object with a string "location" property.')
|
||||
}
|
||||
|
||||
// After this check, TypeScript knows `input` is `WeatherToolInput`
|
||||
const location = input.location
|
||||
|
||||
console.log(`\n[WeatherTool] Getting weather for ${location}...`)
|
||||
|
||||
const fakeWeatherData = {
|
||||
temperature: '72°F',
|
||||
conditions: 'sunny',
|
||||
}
|
||||
|
||||
const resultText = `The weather in ${location} is ${fakeWeatherData.temperature} and ${fakeWeatherData.conditions}.`
|
||||
|
||||
return {
|
||||
toolUseId: context.toolUse.toolUseId,
|
||||
status: 'success' as const,
|
||||
content: [{ type: 'textBlock', text: resultText }],
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* A helper function to run an agent scenario and handle its output stream.
|
||||
* This avoids repeating the for-await loop and logging logic.
|
||||
* @param title The title of the scenario to be logged.
|
||||
* @param agent The agent instance to use.
|
||||
* @param prompt The user prompt to invoke the agent with.
|
||||
*/
|
||||
async function run(title: string, agent: Agent, prompt: string) {
|
||||
console.log(`--- ${title} ---`)
|
||||
console.log(`User: ${prompt}`)
|
||||
|
||||
const responseStream = agent.invoke(prompt)
|
||||
|
||||
console.log('Agent response stream:')
|
||||
let result = await responseStream.next()
|
||||
while (!result.done) {
|
||||
const event = result.value
|
||||
console.log('[Event]', event)
|
||||
result = await responseStream.next()
|
||||
}
|
||||
|
||||
// Clean up logging for the next scenario
|
||||
console.log('\nInvocation complete.\n')
|
||||
}
|
||||
|
||||
async function main() {
|
||||
// 1. Initialize the components
|
||||
const model = new BedrockModel()
|
||||
const weatherTool = new WeatherTool()
|
||||
|
||||
// 2. Create agents
|
||||
const defaultAgent = new Agent()
|
||||
const agentWithoutTools = new Agent({ model })
|
||||
const agentWithTools = new Agent({
|
||||
systemPrompt: 'You are a helpful assistant that provides weather information using the get_weather tool.',
|
||||
model,
|
||||
tools: [weatherTool],
|
||||
})
|
||||
|
||||
await run('0: Invocation with default agent (no model or tools)', defaultAgent, 'Hello!')
|
||||
await run('1: Invocation with a model but no tools', agentWithoutTools, 'Hello!')
|
||||
await run('2: Invocation that uses a tool', agentWithTools, 'What is the weather in Toronto? Use the weather tool.')
|
||||
}
|
||||
|
||||
main().catch(console.error)
|
||||
@@ -0,0 +1,19 @@
|
||||
{
|
||||
"compilerOptions": {
|
||||
"target": "ES2022",
|
||||
"lib": ["ES2022"],
|
||||
"esModuleInterop": true,
|
||||
"forceConsistentCasingInFileNames": true,
|
||||
"strict": true,
|
||||
"skipLibCheck": true,
|
||||
"module": "NodeNext",
|
||||
"moduleResolution": "NodeNext",
|
||||
"outDir": "./dist",
|
||||
"rootDir": "./src",
|
||||
"declaration": true,
|
||||
"declarationMap": true,
|
||||
"sourceMap": true
|
||||
},
|
||||
"include": ["src/**/*"],
|
||||
"exclude": ["node_modules", "dist", "tests*"]
|
||||
}
|
||||
@@ -84,6 +84,8 @@
|
||||
"homepage": "https://github.com/strands-agents/sdk-typescript#readme",
|
||||
"dependencies": {
|
||||
"@aws-sdk/client-bedrock-runtime": "^3.911.0",
|
||||
"@modelcontextprotocol/sdk": "^1.20.2",
|
||||
"uuid": "^13.0.0",
|
||||
"zod": "^4.1.12"
|
||||
},
|
||||
"optionalDependencies": {
|
||||
|
||||
@@ -38,7 +38,7 @@ export function createMockTool(
|
||||
resultFn: () => ToolResult | AsyncGenerator<never, ToolResult, never>
|
||||
): Tool {
|
||||
return {
|
||||
toolName: name,
|
||||
name,
|
||||
description: `Mock tool ${name}`,
|
||||
toolSpec: {
|
||||
name,
|
||||
|
||||
@@ -1,526 +0,0 @@
|
||||
import { describe, it, expect } from 'vitest'
|
||||
import { runAgentLoop } from '../agent-loop.js'
|
||||
import { TestModelProvider, collectGenerator } from '../../__fixtures__/model-test-helpers.js'
|
||||
import { MockMessageModel } from '../../__fixtures__/mock-message-model.js'
|
||||
import { createMockTool } from '../../__fixtures__/tool-helpers.js'
|
||||
import { ToolRegistry } from '../../tools/registry.js'
|
||||
import { Message, TextBlock } from '../../types/messages.js'
|
||||
import { MaxTokensError } from '../../errors.js'
|
||||
|
||||
describe('runAgentLoop', () => {
|
||||
describe('when handling simple completion without tools', () => {
|
||||
it('yields events and returns final messages array', async () => {
|
||||
const provider = new TestModelProvider(async function* () {
|
||||
yield { type: 'modelMessageStartEvent', role: 'assistant' }
|
||||
yield { type: 'modelContentBlockStartEvent', contentBlockIndex: 0 }
|
||||
yield {
|
||||
type: 'modelContentBlockDeltaEvent',
|
||||
delta: { type: 'textDelta', text: 'Hello, how can I help?' },
|
||||
contentBlockIndex: 0,
|
||||
}
|
||||
yield { type: 'modelContentBlockStopEvent', contentBlockIndex: 0 }
|
||||
yield { type: 'modelMessageStopEvent', stopReason: 'endTurn' }
|
||||
})
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Hi')],
|
||||
}),
|
||||
]
|
||||
|
||||
const { items } = await collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
|
||||
// Verify agent events are present
|
||||
expect(items).toContainEqual({ type: 'beforeInvocationEvent' })
|
||||
expect(items).toContainEqual({ type: 'beforeModelEvent', messages: expect.any(Array) })
|
||||
expect(items).toContainEqual({
|
||||
type: 'afterModelEvent',
|
||||
message: expect.objectContaining({ role: 'assistant' }),
|
||||
stopReason: expect.any(String),
|
||||
})
|
||||
expect(items).toContainEqual({ type: 'afterInvocationEvent' })
|
||||
|
||||
// Verify model events are passed through
|
||||
expect(items).toContainEqual({ type: 'modelMessageStartEvent', role: 'assistant' })
|
||||
|
||||
// Verify final messages array contains assistant response
|
||||
expect(messages).toHaveLength(2)
|
||||
expect(messages[1]).toEqual({
|
||||
type: 'message',
|
||||
role: 'assistant',
|
||||
content: [{ type: 'textBlock', text: 'Hello, how can I help?' }],
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe('when handling single tool use cycle', () => {
|
||||
it('executes tool and continues loop until completion', async () => {
|
||||
const provider = new MockMessageModel()
|
||||
.addTurn({
|
||||
type: 'toolUseBlock',
|
||||
name: 'calculator',
|
||||
toolUseId: 'tool-1',
|
||||
input: { operation: 'add', a: 5, b: 3 },
|
||||
})
|
||||
.addTurn({ type: 'textBlock', text: 'The result is 8' })
|
||||
|
||||
const mockTool = createMockTool('calculator', () => ({
|
||||
toolUseId: 'tool-1',
|
||||
status: 'success',
|
||||
content: [{ type: 'textBlock', text: '8' }],
|
||||
}))
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
registry.register(mockTool)
|
||||
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [{ type: 'textBlock', text: 'What is 5+3?' }],
|
||||
}),
|
||||
]
|
||||
|
||||
const { items } = await collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
|
||||
// Verify tool execution events
|
||||
expect(items).toContainEqual({
|
||||
type: 'beforeToolsEvent',
|
||||
message: expect.objectContaining({ role: 'assistant' }),
|
||||
})
|
||||
expect(items).toContainEqual({
|
||||
type: 'afterToolsEvent',
|
||||
message: expect.objectContaining({ role: 'user' }),
|
||||
})
|
||||
|
||||
// Verify only one beforeInvocationEvent
|
||||
const beforeEvents = items.filter((e) => e.type === 'beforeInvocationEvent')
|
||||
expect(beforeEvents).toHaveLength(1)
|
||||
|
||||
// Verify two iterations using callCount
|
||||
expect(provider.callCount).toBe(2)
|
||||
|
||||
// Verify final messages include tool use and result
|
||||
expect(messages).toHaveLength(4) // user, assistant with tool use, user with tool result, assistant with final response
|
||||
if (!messages[1] || !messages[1].content[0]) {
|
||||
throw new Error('Expected content at index 1')
|
||||
}
|
||||
expect(messages[1].content[0]).toMatchObject({
|
||||
type: 'toolUseBlock',
|
||||
name: 'calculator',
|
||||
toolUseId: 'tool-1',
|
||||
})
|
||||
if (!messages[2] || !messages[2].content[0]) {
|
||||
throw new Error('Expected content at index 2')
|
||||
}
|
||||
expect(messages[2].content[0]).toMatchObject({
|
||||
type: 'toolResultBlock',
|
||||
toolUseId: 'tool-1',
|
||||
status: 'success',
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe('when handling multiple tool uses in sequence', () => {
|
||||
it('executes all tools sequentially', async () => {
|
||||
const provider = new MockMessageModel()
|
||||
.addTurn([
|
||||
{
|
||||
type: 'toolUseBlock',
|
||||
name: 'tool1',
|
||||
toolUseId: 'id-1',
|
||||
input: {},
|
||||
},
|
||||
{
|
||||
type: 'toolUseBlock',
|
||||
name: 'tool2',
|
||||
toolUseId: 'id-2',
|
||||
input: {},
|
||||
},
|
||||
])
|
||||
.addTurn({ type: 'textBlock', text: 'Done' })
|
||||
|
||||
const tool1 = createMockTool('tool1', () => ({
|
||||
toolUseId: 'id-1',
|
||||
status: 'success',
|
||||
content: [{ type: 'textBlock', text: 'result1' }],
|
||||
}))
|
||||
|
||||
const tool2 = createMockTool('tool2', () => ({
|
||||
toolUseId: 'id-2',
|
||||
status: 'success',
|
||||
content: [{ type: 'textBlock', text: 'result2' }],
|
||||
}))
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
registry.register([tool1, tool2])
|
||||
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Test')],
|
||||
}),
|
||||
]
|
||||
|
||||
await collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
|
||||
// Verify both tool results are present
|
||||
const toolResultMessage = messages[2]
|
||||
if (!toolResultMessage) {
|
||||
throw new Error('Expected tool result message at index 2')
|
||||
}
|
||||
expect(toolResultMessage.content).toHaveLength(2)
|
||||
expect(toolResultMessage.content[0]).toMatchObject({
|
||||
type: 'toolResultBlock',
|
||||
toolUseId: 'id-1',
|
||||
})
|
||||
expect(toolResultMessage.content[1]).toMatchObject({
|
||||
type: 'toolResultBlock',
|
||||
toolUseId: 'id-2',
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe('when handling multiple agentic loop iterations', () => {
|
||||
it('continues through multiple tool-use cycles', async () => {
|
||||
const provider = new MockMessageModel()
|
||||
.addTurn({
|
||||
type: 'toolUseBlock',
|
||||
name: 'tool1',
|
||||
toolUseId: 'id-1',
|
||||
input: {},
|
||||
})
|
||||
.addTurn({
|
||||
type: 'toolUseBlock',
|
||||
name: 'tool2',
|
||||
toolUseId: 'id-2',
|
||||
input: {},
|
||||
})
|
||||
.addTurn({ type: 'textBlock', text: 'Complete' })
|
||||
|
||||
const tool1 = createMockTool('tool1', () => ({
|
||||
toolUseId: 'id-1',
|
||||
status: 'success',
|
||||
content: [{ type: 'textBlock', text: 'r1' }],
|
||||
}))
|
||||
|
||||
const tool2 = createMockTool('tool2', () => ({
|
||||
toolUseId: 'id-2',
|
||||
status: 'success',
|
||||
content: [{ type: 'textBlock', text: 'r2' }],
|
||||
}))
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
registry.register([tool1, tool2])
|
||||
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Test')],
|
||||
}),
|
||||
]
|
||||
|
||||
const { items } = await collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
|
||||
// Verify only one beforeInvocationEvent
|
||||
const beforeEvents = items.filter((e) => e.type === 'beforeInvocationEvent')
|
||||
expect(beforeEvents).toHaveLength(1)
|
||||
|
||||
// Verify three iterations using callCount
|
||||
expect(provider.callCount).toBe(3)
|
||||
|
||||
// Verify final message count (1 user + 2 assistant tool use + 2 user tool results + 1 assistant final)
|
||||
expect(messages).toHaveLength(6)
|
||||
})
|
||||
})
|
||||
|
||||
describe('when handling transactional message success', () => {
|
||||
it('adds assistant message to array after first model event', async () => {
|
||||
const provider = new MockMessageModel().addTurn({ type: 'textBlock', text: 'Response' })
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Test')],
|
||||
}),
|
||||
]
|
||||
|
||||
await collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
|
||||
// Verify assistant message was added
|
||||
expect(messages).toHaveLength(2)
|
||||
if (!messages[1]) {
|
||||
throw new Error('Expected assistant message at index 1')
|
||||
}
|
||||
expect(messages[1].role).toBe('assistant')
|
||||
})
|
||||
})
|
||||
|
||||
describe('when handling transactional message with early error', () => {
|
||||
it('throws error without adding message to array', async () => {
|
||||
const provider = new MockMessageModel().addTurn(new Error('Model error before any events'))
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Test')],
|
||||
}),
|
||||
]
|
||||
|
||||
// Verify error is thrown
|
||||
await expect(
|
||||
collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
).rejects.toThrow('Model error before any events')
|
||||
})
|
||||
})
|
||||
|
||||
describe('when model throws error after first event', () => {
|
||||
it('propagates error with messages array preserved', async () => {
|
||||
// For error after first event, we need to use TestModelProvider since TestMessageModelProvider
|
||||
// throws errors before any events are generated
|
||||
const provider = new TestModelProvider(async function* () {
|
||||
yield { type: 'modelMessageStartEvent', role: 'assistant' }
|
||||
throw new Error('Error after first event')
|
||||
})
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Test')],
|
||||
}),
|
||||
]
|
||||
|
||||
// Verify error is thrown
|
||||
await expect(
|
||||
collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
).rejects.toThrow('Error after first event')
|
||||
})
|
||||
})
|
||||
|
||||
describe('when tool throws exception', () => {
|
||||
it('propagates the error from the tool', async () => {
|
||||
const provider = new MockMessageModel().addTurn({
|
||||
type: 'toolUseBlock',
|
||||
name: 'badTool',
|
||||
toolUseId: 'id-1',
|
||||
input: {},
|
||||
})
|
||||
|
||||
// eslint-disable-next-line require-yield
|
||||
const badTool = createMockTool('badTool', async function* () {
|
||||
throw new Error('Tool execution failed')
|
||||
})
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
registry.register(badTool)
|
||||
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Test')],
|
||||
}),
|
||||
]
|
||||
|
||||
await expect(
|
||||
collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
).rejects.toThrow('Tool execution failed')
|
||||
})
|
||||
|
||||
it('does not add assistant message with tool uses when tool execution fails', async () => {
|
||||
const provider = new MockMessageModel().addTurn({
|
||||
type: 'toolUseBlock',
|
||||
name: 'badTool',
|
||||
toolUseId: 'id-1',
|
||||
input: {},
|
||||
})
|
||||
|
||||
// eslint-disable-next-line require-yield
|
||||
const badTool = createMockTool('badTool', async function* () {
|
||||
throw new Error('Tool execution failed')
|
||||
})
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
registry.register(badTool)
|
||||
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Test')],
|
||||
}),
|
||||
]
|
||||
|
||||
try {
|
||||
await collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
throw new Error('Expected error to be thrown')
|
||||
} catch (error) {
|
||||
if (error instanceof Error && error.message === 'Tool execution failed') {
|
||||
// Verify that messages array only contains the initial user message
|
||||
// The assistant message with tool uses should NOT be present since tool execution failed
|
||||
expect(messages).toHaveLength(1)
|
||||
expect(messages[0]).toEqual({
|
||||
type: 'message',
|
||||
role: 'user',
|
||||
content: [{ type: 'textBlock', text: 'Test' }],
|
||||
})
|
||||
} else {
|
||||
throw error
|
||||
}
|
||||
}
|
||||
})
|
||||
})
|
||||
|
||||
describe('when tool is not found in registry', () => {
|
||||
it('returns error tool result and continues loop', async () => {
|
||||
const provider = new MockMessageModel()
|
||||
.addTurn({
|
||||
type: 'toolUseBlock',
|
||||
name: 'nonexistent',
|
||||
toolUseId: 'id-1',
|
||||
input: {},
|
||||
})
|
||||
.addTurn({ type: 'textBlock', text: 'Tool not available' })
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Test')],
|
||||
}),
|
||||
]
|
||||
|
||||
await collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
|
||||
// Verify error tool result was returned
|
||||
const toolResultMessage = messages[2]
|
||||
if (!toolResultMessage || !toolResultMessage.content[0]) {
|
||||
throw new Error('Expected tool result message at index 2')
|
||||
}
|
||||
expect(toolResultMessage.content[0]).toMatchObject({
|
||||
type: 'toolResultBlock',
|
||||
toolUseId: 'id-1',
|
||||
status: 'error',
|
||||
})
|
||||
|
||||
// Verify loop continued and completed
|
||||
expect(messages).toHaveLength(4) // user, assistant tool use, user error result, assistant final
|
||||
})
|
||||
})
|
||||
|
||||
describe('when maxTokens stop reason occurs', () => {
|
||||
it('throws MaxTokensError', async () => {
|
||||
const provider = new MockMessageModel().addTurn({ type: 'textBlock', text: 'Partial' }, 'maxTokens')
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Test')],
|
||||
}),
|
||||
]
|
||||
|
||||
await expect(
|
||||
collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
).rejects.toThrow(MaxTokensError)
|
||||
})
|
||||
})
|
||||
|
||||
describe('when constructing ContentBlocks via streamAggregated', () => {
|
||||
it('handles TextBlock correctly', async () => {
|
||||
const provider = new MockMessageModel().addTurn({ type: 'textBlock', text: 'Hello' })
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Hi')],
|
||||
}),
|
||||
]
|
||||
|
||||
await collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
|
||||
if (!messages[1] || !messages[1].content[0]) {
|
||||
throw new Error('Expected content at index 1')
|
||||
}
|
||||
expect(messages[1].content[0]).toEqual({
|
||||
type: 'textBlock',
|
||||
text: 'Hello',
|
||||
})
|
||||
})
|
||||
|
||||
it('handles ToolUseBlock correctly', async () => {
|
||||
const provider = new MockMessageModel()
|
||||
.addTurn({
|
||||
type: 'toolUseBlock',
|
||||
name: 'test',
|
||||
toolUseId: 'id-1',
|
||||
input: { key: 'value' },
|
||||
})
|
||||
.addTurn({ type: 'textBlock', text: 'Done' })
|
||||
|
||||
const tool = createMockTool('test', () => ({
|
||||
toolUseId: 'id-1',
|
||||
status: 'success',
|
||||
content: [{ type: 'textBlock', text: 'ok' }],
|
||||
}))
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
registry.register(tool)
|
||||
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Hi')],
|
||||
}),
|
||||
]
|
||||
|
||||
await collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
|
||||
const toolUseBlock = messages[1]?.content[0]
|
||||
if (!toolUseBlock || toolUseBlock.type !== 'toolUseBlock') {
|
||||
throw new Error('Expected tool use block at messages[1].content[0]')
|
||||
}
|
||||
expect(toolUseBlock).toEqual({
|
||||
type: 'toolUseBlock',
|
||||
name: 'test',
|
||||
toolUseId: 'id-1',
|
||||
input: { key: 'value' },
|
||||
})
|
||||
})
|
||||
|
||||
it('handles ReasoningBlock correctly', async () => {
|
||||
const provider = new MockMessageModel().addTurn([
|
||||
{ type: 'reasoningBlock', text: 'thinking...' },
|
||||
{ type: 'textBlock', text: 'Response' },
|
||||
])
|
||||
|
||||
const registry = new ToolRegistry()
|
||||
const messages: Message[] = [
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [new TextBlock('Hi')],
|
||||
}),
|
||||
]
|
||||
await collectGenerator(runAgentLoop({ model: provider, messages, toolRegistry: registry }))
|
||||
|
||||
if (!messages[1] || !messages[1].content[0]) {
|
||||
throw new Error('Expected content blocks at index 1')
|
||||
}
|
||||
expect(messages[1].content[0]).toEqual({
|
||||
type: 'reasoningBlock',
|
||||
text: 'thinking...',
|
||||
})
|
||||
if (!messages[1].content[1]) {
|
||||
throw new Error('Expected second content block at index 1')
|
||||
}
|
||||
expect(messages[1].content[1]).toEqual({
|
||||
type: 'textBlock',
|
||||
text: 'Response',
|
||||
})
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -1,227 +0,0 @@
|
||||
import { Message, TextBlock, ToolResultBlock, type SystemPrompt, type ToolUseBlock } from '../types/messages.js'
|
||||
import type { BaseModelConfig, Model, StreamOptions } from '../models/model.js'
|
||||
import type { ToolRegistry } from '../tools/registry.js'
|
||||
import type { AgentStreamEvent } from './streaming.js'
|
||||
import { MaxTokensError } from '../errors.js'
|
||||
import type { AgentResult } from '../types/agent.js'
|
||||
|
||||
/**
|
||||
* Internal configuration for the agent loop.
|
||||
* @internal
|
||||
*/
|
||||
interface AgentLike {
|
||||
/**
|
||||
* Model provider instance for generating responses.
|
||||
*/
|
||||
model: Model<BaseModelConfig>
|
||||
|
||||
/**
|
||||
* Array of conversation messages (will be mutated as the loop progresses).
|
||||
*/
|
||||
messages: Message[]
|
||||
|
||||
/**
|
||||
* Registry containing available tools.
|
||||
*/
|
||||
toolRegistry: ToolRegistry
|
||||
|
||||
/**
|
||||
* Optional system prompt to guide model behavior.
|
||||
*/
|
||||
systemPrompt?: SystemPrompt
|
||||
}
|
||||
|
||||
/**
|
||||
* Async generator that coordinates execution between model providers and tools.
|
||||
*
|
||||
* The agent loop manages the conversation flow by:
|
||||
* 1. Streaming model responses and yielding all events
|
||||
* 2. Executing tools when the model requests them
|
||||
* 3. Continuing the loop until the model completes without tool use
|
||||
*
|
||||
* An explicit goal of this method is to always leave the message array in a way that
|
||||
* the agent can be reinvoked with a user prompt after this method completes. To that end
|
||||
* assistant messages containing tool uses are only added after tool execution succeeds
|
||||
* with valid toolResponses
|
||||
*
|
||||
* @param agent - Configuration including model, messages, toolRegistry, and systemPrompt
|
||||
* @returns Async generator that yields AgentStreamEvent objects and returns AgentResult
|
||||
*
|
||||
* @example
|
||||
* ```typescript
|
||||
* const messages = [{ type: 'message', role: 'user', content: [{ type: 'textBlock', text: 'Hello' }] }]
|
||||
* const registry = new ToolRegistry()
|
||||
* const provider = new BedrockModel(config)
|
||||
*
|
||||
* for await (const event of runAgentLoop({ model: provider, messages, toolRegistry: registry })) {
|
||||
* console.log('Event:', event.type)
|
||||
* }
|
||||
* // Messages array is mutated in place and contains the full conversation
|
||||
* ```
|
||||
*/
|
||||
export async function* runAgentLoop(agent: AgentLike): AsyncGenerator<AgentStreamEvent, AgentResult, never> {
|
||||
// Emit event before the loop starts
|
||||
yield { type: 'beforeInvocationEvent' }
|
||||
|
||||
try {
|
||||
// Main agent loop - continues until model stops without requesting tools
|
||||
while (true) {
|
||||
const modelResult = yield* invokeModel(agent)
|
||||
|
||||
// Handle stop reason
|
||||
if (modelResult.stopReason === 'maxTokens') {
|
||||
throw new MaxTokensError(
|
||||
'Model reached maximum token limit. This is an unrecoverable state that requires intervention.',
|
||||
modelResult.message
|
||||
)
|
||||
}
|
||||
|
||||
if (modelResult.stopReason !== 'toolUse') {
|
||||
// Loop terminates - no tool use requested
|
||||
// Add assistant message now that we're returning
|
||||
agent.messages.push(modelResult.message)
|
||||
return {
|
||||
stopReason: modelResult.stopReason,
|
||||
lastMessage: modelResult.message,
|
||||
}
|
||||
}
|
||||
|
||||
// Execute tools sequentially
|
||||
const toolResultMessage = yield* executeTools(modelResult.message, agent.toolRegistry)
|
||||
|
||||
// Add assistant message with tool uses right before adding tool results
|
||||
// This ensures we don't have dangling tool use messages if tool execution fails
|
||||
agent.messages.push(modelResult.message)
|
||||
agent.messages.push(toolResultMessage)
|
||||
|
||||
// Continue loop
|
||||
}
|
||||
} finally {
|
||||
// Always emit final event
|
||||
yield { type: 'afterInvocationEvent' }
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Invokes the model provider and streams all events.
|
||||
*
|
||||
* @param agent - Agent configuration containing model, messages, toolRegistry, and systemPrompt
|
||||
* @returns Object containing the assistant message and stop reason
|
||||
*/
|
||||
async function* invokeModel(
|
||||
agent: AgentLike
|
||||
): AsyncGenerator<AgentStreamEvent, { message: Message; stopReason: string }, never> {
|
||||
// Emit event before invoking model
|
||||
yield { type: 'beforeModelEvent', messages: [...agent.messages] }
|
||||
|
||||
const toolSpecs = agent.toolRegistry.list().map((tool) => tool.toolSpec)
|
||||
const streamOptions: StreamOptions = { toolSpecs }
|
||||
if (agent.systemPrompt !== undefined) {
|
||||
streamOptions.systemPrompt = agent.systemPrompt
|
||||
}
|
||||
|
||||
const { message, stopReason } = yield* agent.model.streamAggregated(agent.messages, streamOptions)
|
||||
|
||||
yield { type: 'afterModelEvent', message, stopReason }
|
||||
|
||||
return { message, stopReason }
|
||||
}
|
||||
|
||||
/**
|
||||
* Executes tools sequentially and streams all tool events.
|
||||
*
|
||||
* @param assistantMessage - The assistant message containing tool use blocks
|
||||
* @param toolRegistry - Registry containing available tools
|
||||
* @returns User message containing tool results
|
||||
*/
|
||||
async function* executeTools(
|
||||
assistantMessage: Message,
|
||||
toolRegistry: ToolRegistry
|
||||
): AsyncGenerator<AgentStreamEvent, Message, never> {
|
||||
yield { type: 'beforeToolsEvent', message: assistantMessage }
|
||||
|
||||
// Extract tool use blocks from assistant message
|
||||
const toolUseBlocks = assistantMessage.content.filter((block): block is ToolUseBlock => block.type === 'toolUseBlock')
|
||||
|
||||
if (toolUseBlocks.length === 0) {
|
||||
// No tool use blocks found even though stopReason is toolUse
|
||||
throw new Error('Model indicated toolUse but no tool use blocks found in message')
|
||||
}
|
||||
|
||||
const toolResultBlocks: ToolResultBlock[] = []
|
||||
|
||||
for (const toolUseBlock of toolUseBlocks) {
|
||||
const toolResultBlock = yield* executeTool(toolUseBlock, toolRegistry)
|
||||
toolResultBlocks.push(toolResultBlock)
|
||||
|
||||
// Yield the tool result block as it's created
|
||||
yield toolResultBlock as AgentStreamEvent
|
||||
}
|
||||
|
||||
// Create user message with tool results
|
||||
const toolResultMessage: Message = new Message({
|
||||
role: 'user',
|
||||
content: toolResultBlocks,
|
||||
})
|
||||
|
||||
yield { type: 'afterToolsEvent', message: toolResultMessage }
|
||||
|
||||
return toolResultMessage
|
||||
}
|
||||
|
||||
/**
|
||||
* Executes a single tool and returns the result.
|
||||
* If the tool is not found or fails to return a result, returns an error ToolResult
|
||||
* instead of throwing an exception. This allows the agent loop to continue and
|
||||
* let the model handle the error gracefully.
|
||||
*
|
||||
* @param toolUseBlock - Tool use block to execute
|
||||
* @param toolRegistry - Registry containing available tools
|
||||
* @returns Tool result block
|
||||
*/
|
||||
async function* executeTool(
|
||||
toolUseBlock: ToolUseBlock,
|
||||
toolRegistry: ToolRegistry
|
||||
): AsyncGenerator<AgentStreamEvent, ToolResultBlock, never> {
|
||||
const tool = toolRegistry.get(toolUseBlock.name)
|
||||
|
||||
if (!tool) {
|
||||
// Tool not found - return error result instead of throwing
|
||||
return new ToolResultBlock({
|
||||
toolUseId: toolUseBlock.toolUseId,
|
||||
status: 'error',
|
||||
content: [new TextBlock(`Tool '${toolUseBlock.name}' not found in registry`)],
|
||||
})
|
||||
}
|
||||
|
||||
// Execute tool and collect result
|
||||
const toolContext = {
|
||||
toolUse: {
|
||||
name: toolUseBlock.name,
|
||||
toolUseId: toolUseBlock.toolUseId,
|
||||
input: toolUseBlock.input,
|
||||
},
|
||||
invocationState: {},
|
||||
}
|
||||
|
||||
const toolGenerator = tool.stream(toolContext)
|
||||
|
||||
// Use yield* to delegate to the tool generator and capture the return value
|
||||
const toolResult = yield* toolGenerator
|
||||
|
||||
if (!toolResult) {
|
||||
// Tool didn't return a result - return error result instead of throwing
|
||||
return new ToolResultBlock({
|
||||
toolUseId: toolUseBlock.toolUseId,
|
||||
status: 'error',
|
||||
content: [new TextBlock(`Tool '${toolUseBlock.name}' did not return a result`)],
|
||||
})
|
||||
}
|
||||
|
||||
// Create ToolResultBlock from ToolResult
|
||||
return new ToolResultBlock({
|
||||
toolUseId: toolResult.toolUseId,
|
||||
status: toolResult.status,
|
||||
content: toolResult.content,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,297 @@
|
||||
import {
|
||||
BedrockModel,
|
||||
MaxTokensError,
|
||||
type AgentResult,
|
||||
type AgentStreamEvent,
|
||||
Message,
|
||||
ToolResultBlock,
|
||||
type SystemPrompt,
|
||||
type Tool,
|
||||
type ToolUseBlock,
|
||||
type MessageData,
|
||||
TextBlock,
|
||||
} from '../index.js'
|
||||
import type { BaseModelConfig, Model, StreamOptions } from '../models/model.js'
|
||||
import { ToolRegistry } from '../registry/tool-registry.js'
|
||||
|
||||
/**
|
||||
* Configuration object for creating a new Agent.
|
||||
*/
|
||||
export type AgentConfig = {
|
||||
/**
|
||||
* The model instance that the agent will use to make decisions.
|
||||
*/
|
||||
model?: Model<BaseModelConfig>
|
||||
/**
|
||||
* An initial set of messages to seed the agent's conversation history.
|
||||
*/
|
||||
messages?: Message[] | MessageData[]
|
||||
/**
|
||||
* An initial set of tools to register with the agent.
|
||||
*/
|
||||
tools?: Tool[]
|
||||
/**
|
||||
* A system prompt which guides model behavior.
|
||||
*/
|
||||
systemPrompt?: SystemPrompt
|
||||
}
|
||||
|
||||
/**
|
||||
* Arguments for invoking an agent.
|
||||
*
|
||||
* A plain string represents user input to an agent.
|
||||
*/
|
||||
export type InvokeArgs = string
|
||||
|
||||
/**
|
||||
* Orchestrates the interaction between a model, a set of tools, and MCP clients.
|
||||
* The Agent is responsible for managing the lifecycle of tools and clients
|
||||
* and invoking the core decision-making loop.
|
||||
*/
|
||||
export class Agent {
|
||||
private _model: Model<BaseModelConfig>
|
||||
private _toolRegistry: ToolRegistry
|
||||
private _systemPrompt?: SystemPrompt
|
||||
private _messages: Message[]
|
||||
|
||||
/**
|
||||
* Creates an instance of the Agent.
|
||||
* @param config - The configuration for the agent.
|
||||
*/
|
||||
constructor(config?: AgentConfig) {
|
||||
this._model = config?.model ?? new BedrockModel()
|
||||
this._toolRegistry = new ToolRegistry(config?.tools)
|
||||
|
||||
if (config?.systemPrompt !== undefined) {
|
||||
this._systemPrompt = config.systemPrompt
|
||||
}
|
||||
|
||||
this._messages = (config?.messages ?? []).map((msg) =>
|
||||
msg instanceof Message ? msg : Message.fromMessageData(msg)
|
||||
)
|
||||
}
|
||||
|
||||
/**
|
||||
* The tools this agent can use.
|
||||
*/
|
||||
get tools(): Tool[] {
|
||||
return this._toolRegistry.values()
|
||||
}
|
||||
|
||||
/**
|
||||
* The tool registry for managing the agent's tools.
|
||||
*/
|
||||
get toolRegistry(): ToolRegistry {
|
||||
return this._toolRegistry
|
||||
}
|
||||
|
||||
/**
|
||||
* Async generator that coordinates execution between model providers and tools.
|
||||
*
|
||||
* The agent loop manages the conversation flow by:
|
||||
* 1. Streaming model responses and yielding all events
|
||||
* 2. Executing tools when the model requests them
|
||||
* 3. Continuing the loop until the model completes without tool use
|
||||
*
|
||||
* An explicit goal of this method is to always leave the message array in a way that
|
||||
* the agent can be reinvoked with a user prompt after this method completes. To that end
|
||||
* assistant messages containing tool uses are only added after tool execution succeeds
|
||||
* with valid toolResponses
|
||||
*
|
||||
* @param agent - Configuration including model, messages, toolRegistry, and systemPrompt
|
||||
* @returns Async generator that yields AgentStreamEvent objects and returns AgentResult
|
||||
*
|
||||
* @example
|
||||
* ```typescript
|
||||
* const messages = [{ type: 'message', role: 'user', content: [{ type: 'textBlock', text: 'Hello' }] }]
|
||||
* const registry = new ToolRegistry()
|
||||
* const provider = new BedrockModel(config)
|
||||
*
|
||||
* for await (const event of runAgentLoop({ model: provider, messages, toolRegistry: registry })) {
|
||||
* console.log('Event:', event.type)
|
||||
* }
|
||||
* // Messages array is mutated in place and contains the full conversation
|
||||
* ```
|
||||
*/
|
||||
public async *invoke(args: InvokeArgs): AsyncGenerator<AgentStreamEvent, AgentResult, never> {
|
||||
let currentArgs: InvokeArgs | undefined = args
|
||||
|
||||
// Emit event before the loop starts
|
||||
yield { type: 'beforeInvocationEvent' }
|
||||
|
||||
try {
|
||||
// Main agent loop - continues until model stops without requesting tools
|
||||
while (true) {
|
||||
const modelResult = yield* this.invokeModel(currentArgs)
|
||||
currentArgs = undefined // Only pass args on first invocation
|
||||
|
||||
// Handle stop reason
|
||||
if (modelResult.stopReason === 'maxTokens') {
|
||||
throw new MaxTokensError(
|
||||
'Model reached maximum token limit. This is an unrecoverable state that requires intervention.',
|
||||
modelResult.message
|
||||
)
|
||||
}
|
||||
|
||||
if (modelResult.stopReason !== 'toolUse') {
|
||||
// Loop terminates - no tool use requested
|
||||
// Add assistant message now that we're returning
|
||||
this._messages.push(modelResult.message)
|
||||
return {
|
||||
stopReason: modelResult.stopReason,
|
||||
lastMessage: modelResult.message,
|
||||
}
|
||||
}
|
||||
|
||||
// Execute tools sequentially
|
||||
const toolResultMessage = yield* this.executeTools(modelResult.message, this._toolRegistry)
|
||||
|
||||
// Add assistant message with tool uses right before adding tool results
|
||||
// This ensures we don't have dangling tool use messages if tool execution fails
|
||||
this._messages.push(modelResult.message)
|
||||
this._messages.push(toolResultMessage)
|
||||
|
||||
// Continue loop
|
||||
}
|
||||
} finally {
|
||||
// Always emit final event
|
||||
yield { type: 'afterInvocationEvent' }
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Invokes the model provider and streams all events.
|
||||
*
|
||||
* @param args - Optional arguments for invoking the model
|
||||
* @returns Object containing the assistant message and stop reason
|
||||
*/
|
||||
private async *invokeModel(
|
||||
args?: InvokeArgs
|
||||
): AsyncGenerator<AgentStreamEvent, { message: Message; stopReason: string }, never> {
|
||||
// Emit event before invoking model
|
||||
yield { type: 'beforeModelEvent', messages: [...this._messages] }
|
||||
|
||||
const toolSpecs = this._toolRegistry.values().map((tool) => tool.toolSpec)
|
||||
const streamOptions: StreamOptions = { toolSpecs }
|
||||
if (this._systemPrompt !== undefined) {
|
||||
streamOptions.systemPrompt = this._systemPrompt
|
||||
}
|
||||
|
||||
if (args !== undefined && typeof args === 'string') {
|
||||
// Add user message from args
|
||||
this._messages.push(
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [{ type: 'textBlock', text: args }],
|
||||
})
|
||||
)
|
||||
}
|
||||
|
||||
const { message, stopReason } = yield* this._model.streamAggregated(this._messages, streamOptions)
|
||||
|
||||
yield { type: 'afterModelEvent', message, stopReason }
|
||||
|
||||
return { message, stopReason }
|
||||
}
|
||||
|
||||
/**
|
||||
* Executes tools sequentially and streams all tool events.
|
||||
*
|
||||
* @param assistantMessage - The assistant message containing tool use blocks
|
||||
* @param toolRegistry - Registry containing available tools
|
||||
* @returns User message containing tool results
|
||||
*/
|
||||
private async *executeTools(
|
||||
assistantMessage: Message,
|
||||
toolRegistry: ToolRegistry
|
||||
): AsyncGenerator<AgentStreamEvent, Message, never> {
|
||||
yield { type: 'beforeToolsEvent', message: assistantMessage }
|
||||
|
||||
// Extract tool use blocks from assistant message
|
||||
const toolUseBlocks = assistantMessage.content.filter(
|
||||
(block): block is ToolUseBlock => block.type === 'toolUseBlock'
|
||||
)
|
||||
|
||||
if (toolUseBlocks.length === 0) {
|
||||
// No tool use blocks found even though stopReason is toolUse
|
||||
throw new Error('Model indicated toolUse but no tool use blocks found in message')
|
||||
}
|
||||
|
||||
const toolResultBlocks: ToolResultBlock[] = []
|
||||
|
||||
for (const toolUseBlock of toolUseBlocks) {
|
||||
const toolResultBlock = yield* this.executeTool(toolUseBlock, toolRegistry)
|
||||
toolResultBlocks.push(toolResultBlock)
|
||||
|
||||
// Yield the tool result block as it's created
|
||||
yield toolResultBlock as AgentStreamEvent
|
||||
}
|
||||
|
||||
// Create user message with tool results
|
||||
const toolResultMessage: Message = new Message({
|
||||
role: 'user',
|
||||
content: toolResultBlocks,
|
||||
})
|
||||
|
||||
yield { type: 'afterToolsEvent', message: toolResultMessage }
|
||||
|
||||
return toolResultMessage
|
||||
}
|
||||
|
||||
/**
|
||||
* Executes a single tool and returns the result.
|
||||
* If the tool is not found or fails to return a result, returns an error ToolResult
|
||||
* instead of throwing an exception. This allows the agent loop to continue and
|
||||
* let the model handle the error gracefully.
|
||||
*
|
||||
* @param toolUseBlock - Tool use block to execute
|
||||
* @param toolRegistry - Registry containing available tools
|
||||
* @returns Tool result block
|
||||
*/
|
||||
private async *executeTool(
|
||||
toolUseBlock: ToolUseBlock,
|
||||
toolRegistry: ToolRegistry
|
||||
): AsyncGenerator<AgentStreamEvent, ToolResultBlock, never> {
|
||||
const tool = toolRegistry.find((t) => t.name === toolUseBlock.name)
|
||||
|
||||
if (!tool) {
|
||||
// Tool not found - return error result instead of throwing
|
||||
return new ToolResultBlock({
|
||||
toolUseId: toolUseBlock.toolUseId,
|
||||
status: 'error',
|
||||
content: [new TextBlock(`Tool '${toolUseBlock.name}' not found in registry`)],
|
||||
})
|
||||
}
|
||||
|
||||
// Execute tool and collect result
|
||||
const toolContext = {
|
||||
toolUse: {
|
||||
name: toolUseBlock.name,
|
||||
toolUseId: toolUseBlock.toolUseId,
|
||||
input: toolUseBlock.input,
|
||||
},
|
||||
invocationState: {},
|
||||
}
|
||||
|
||||
const toolGenerator = tool.stream(toolContext)
|
||||
|
||||
// Use yield* to delegate to the tool generator and capture the return value
|
||||
const toolResult = yield* toolGenerator
|
||||
|
||||
if (!toolResult) {
|
||||
// Tool didn't return a result - return error result instead of throwing
|
||||
return new ToolResultBlock({
|
||||
toolUseId: toolUseBlock.toolUseId,
|
||||
status: 'error',
|
||||
content: [new TextBlock(`Tool '${toolUseBlock.name}' did not return a result`)],
|
||||
})
|
||||
}
|
||||
|
||||
// Create ToolResultBlock from ToolResult
|
||||
return new ToolResultBlock({
|
||||
toolUseId: toolResult.toolUseId,
|
||||
status: toolResult.status,
|
||||
content: toolResult.content,
|
||||
})
|
||||
}
|
||||
}
|
||||
+3
-3
@@ -5,6 +5,9 @@
|
||||
* public APIs and functionality.
|
||||
*/
|
||||
|
||||
// Agent class
|
||||
export { Agent } from './agent/agent.js'
|
||||
|
||||
// Error types
|
||||
export { ContextWindowOverflowError, MaxTokensError } from './errors.js'
|
||||
|
||||
@@ -44,9 +47,6 @@ export { FunctionTool } from './tools/function-tool.js'
|
||||
// Tool factory function
|
||||
export { tool } from './tools/zod-tool.js'
|
||||
|
||||
// ToolRegistry implementation
|
||||
export { ToolRegistry } from './tools/registry.js'
|
||||
|
||||
// Streaming event types
|
||||
export type {
|
||||
Usage,
|
||||
|
||||
@@ -14,6 +14,8 @@ import type { Message } from '../types/messages.js'
|
||||
import type { ModelStreamEvent } from '../models/streaming.js'
|
||||
import { ContextWindowOverflowError } from '../errors.js'
|
||||
|
||||
const DEFAULT_OPENAI_MODEL_ID = 'gpt-4o'
|
||||
|
||||
/**
|
||||
* Error message patterns that indicate context window overflow.
|
||||
* Used to detect when input exceeds the model's context window.
|
||||
@@ -68,7 +70,7 @@ export interface OpenAIModelConfig extends BaseModelConfig {
|
||||
/**
|
||||
* OpenAI model identifier (e.g., gpt-4o, gpt-3.5-turbo).
|
||||
*/
|
||||
modelId: string
|
||||
modelId?: string
|
||||
|
||||
/**
|
||||
* Controls randomness in generation (0 to 2).
|
||||
@@ -396,7 +398,7 @@ export class OpenAIModel extends Model<OpenAIModelConfig> {
|
||||
): OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming {
|
||||
// Start with required fields
|
||||
const request: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = {
|
||||
model: this._config.modelId,
|
||||
model: this._config.modelId ?? DEFAULT_OPENAI_MODEL_ID,
|
||||
messages: [] as OpenAI.Chat.Completions.ChatCompletionMessageParam[],
|
||||
stream: true,
|
||||
stream_options: { include_usage: true },
|
||||
|
||||
@@ -0,0 +1,359 @@
|
||||
/**
|
||||
* A generic, polymorphic resource registry for managing runtime resources.
|
||||
*
|
||||
* This abstract class provides methods to register, deregister, retrieve,
|
||||
* and find items based on unique identifiers. Subclasses must implement
|
||||
* methods for generating unique IDs and validating items before insertion.
|
||||
*
|
||||
* @typeParam T - The type of the items being stored.
|
||||
* @typeParam I - The type of the identifier for the items.
|
||||
*/
|
||||
|
||||
/**
|
||||
* Thrown when an item with a specific ID cannot be found.
|
||||
* @typeParam I - The type of the item's identifier.
|
||||
*/
|
||||
export class ItemNotFoundError<I> extends Error {
|
||||
constructor(id: I) {
|
||||
super(`Item with id '${id}' not found`)
|
||||
this.name = 'ItemNotFoundError'
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Thrown when attempting to add an item with an ID that already exists.
|
||||
* @typeParam I - The type of the item's identifier.
|
||||
*/
|
||||
export class DuplicateItemError<I> extends Error {
|
||||
constructor(id: I) {
|
||||
super(`An item with the ID '${id}' already exists.`)
|
||||
this.name = 'DuplicateItemError'
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Thrown when an item fails a validation check.
|
||||
*/
|
||||
export class ValidationError extends Error {
|
||||
constructor(message: string) {
|
||||
super(message)
|
||||
this.name = 'ValidationError'
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* A generic, polymorphic registry for managing runtime resources.
|
||||
* @typeParam T - The type of the items being stored.
|
||||
* @typeParam I - The type of the identifier for the items.
|
||||
*/
|
||||
export abstract class Registry<T, I> {
|
||||
protected _items: Map<I, T>
|
||||
|
||||
/**
|
||||
* Abstract method for generating a new, unique identifier.
|
||||
* Subclasses must provide their own implementation (e.g., UUID, auto-increment).
|
||||
* @returns A new, unique identifier.
|
||||
*/
|
||||
protected abstract generateId(item: T): I
|
||||
|
||||
/**
|
||||
* Abstract validation hook called before an item is added.
|
||||
* Subclasses must implement this to provide custom insertion logic.
|
||||
* @param item - The item to be validated.
|
||||
* @throws ValidationError If the item is invalid.
|
||||
*/
|
||||
protected abstract validate(item: T): void
|
||||
|
||||
constructor(items?: T[]) {
|
||||
this._items = new Map()
|
||||
if (items) {
|
||||
this.addAll(items)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Retrieves an item by its ID.
|
||||
* @param id - The identifier of the item to retrieve.
|
||||
* @returns The item if found, otherwise undefined.
|
||||
*/
|
||||
public get(id: I): T | undefined {
|
||||
return this._items.get(id)
|
||||
}
|
||||
|
||||
/**
|
||||
* Finds the first item that satisfies the provided predicate function.
|
||||
* @param predicate - A function to test each item.
|
||||
* @returns The first item that passes the predicate test, otherwise undefined.
|
||||
*/
|
||||
public find(predicate: (item: T) => boolean): T | undefined {
|
||||
for (const item of this._items.values()) {
|
||||
if (predicate(item)) {
|
||||
return item
|
||||
}
|
||||
}
|
||||
|
||||
return undefined
|
||||
}
|
||||
|
||||
/**
|
||||
* Returns an array of all keys (identifiers) in the registry.
|
||||
* @returns An array of all keys.
|
||||
*/
|
||||
public keys(): I[] {
|
||||
return Array.from(this._items.keys())
|
||||
}
|
||||
|
||||
/**
|
||||
* Returns an array of all values (items) in the registry.
|
||||
* @returns An array of all values.
|
||||
*/
|
||||
public values(): T[] {
|
||||
return Array.from(this._items.values())
|
||||
}
|
||||
|
||||
/**
|
||||
* Returns an array of all key-value pairs in the registry.
|
||||
* @returns An array of [id, item] pairs.
|
||||
*/
|
||||
public pairs(): Array<[I, T]> {
|
||||
return Array.from(this._items.entries())
|
||||
}
|
||||
|
||||
/**
|
||||
* Clears all items from the registry.
|
||||
*/
|
||||
public clear(): void {
|
||||
this._items.clear()
|
||||
}
|
||||
|
||||
/**
|
||||
* Validates and adds a new item, assigning it a generated ID.
|
||||
* @param item - The item to add.
|
||||
* @returns The newly generated ID for the item.
|
||||
* @throws DuplicateItemError If the generated ID already exists.
|
||||
* @throws ValidationError If the item fails the validation check.
|
||||
*/
|
||||
public add(item: T): I {
|
||||
this.validate(item)
|
||||
|
||||
const id = this.generateId(item)
|
||||
if (this._items.has(id)) {
|
||||
throw new DuplicateItemError(id)
|
||||
}
|
||||
|
||||
this._items.set(id, item)
|
||||
return id
|
||||
}
|
||||
|
||||
/**
|
||||
* Adds an array of items.
|
||||
* @param items - An array of items to add.
|
||||
* @returns An array of the new IDs for the added items.
|
||||
*/
|
||||
public addAll(items: T[]): I[] {
|
||||
return items.map((item) => this.add(item))
|
||||
}
|
||||
|
||||
/**
|
||||
* Removes an item from the registry by its ID.
|
||||
* @param id - The ID of the item to remove.
|
||||
* @returns The removed item.
|
||||
* @throws ItemNotFoundError If no item with the given ID is found.
|
||||
*/
|
||||
public remove(id: I): T {
|
||||
const item = this._items.get(id)
|
||||
if (item === undefined) {
|
||||
throw new ItemNotFoundError(id)
|
||||
}
|
||||
this._items.delete(id)
|
||||
return item
|
||||
}
|
||||
|
||||
/**
|
||||
* Removes multiple items from the registry by their IDs.
|
||||
* @param ids - An array of IDs of the items to remove.
|
||||
* @returns An array of the removed items.
|
||||
*/
|
||||
public removeAll(ids: I[]): T[] {
|
||||
return ids.map((id) => this.remove(id))
|
||||
}
|
||||
|
||||
/**
|
||||
* Finds the first item matching the predicate, removes it, and returns it.
|
||||
* @param predicate - A function to test each item.
|
||||
* @returns The removed item if found, otherwise undefined.
|
||||
*/
|
||||
public findRemove(predicate: (item: T) => boolean): T | undefined {
|
||||
for (const [id, item] of this._items.entries()) {
|
||||
if (predicate(item)) {
|
||||
this._items.delete(id)
|
||||
return item
|
||||
}
|
||||
}
|
||||
|
||||
return undefined
|
||||
}
|
||||
}
|
||||
|
||||
// Unit tests
|
||||
if (import.meta.vitest) {
|
||||
const { describe, it, expect, beforeEach, vi } = import.meta.vitest
|
||||
|
||||
// A concrete implementation of the abstract Registry for testing purposes
|
||||
class TestRegistry extends Registry<string, number> {
|
||||
private nextId = 1
|
||||
|
||||
protected generateId(): number {
|
||||
return this.nextId++
|
||||
}
|
||||
|
||||
protected validate(item: string): void {
|
||||
if (item.length === 0) {
|
||||
throw new ValidationError('Item cannot be an empty string.')
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
describe('Error Classes', () => {
|
||||
it('ItemNotFoundError should have the correct name and message', () => {
|
||||
const error = new ItemNotFoundError(123)
|
||||
expect(error.name).toBe('ItemNotFoundError')
|
||||
expect(error.message).toBe("Item with id '123' not found")
|
||||
})
|
||||
|
||||
it('DuplicateItemError should have the correct name and message', () => {
|
||||
const error = new DuplicateItemError('abc')
|
||||
expect(error.name).toBe('DuplicateItemError')
|
||||
expect(error.message).toBe("An item with the ID 'abc' already exists.")
|
||||
})
|
||||
|
||||
it('ValidationError should have the correct name and message', () => {
|
||||
const error = new ValidationError('Invalid item')
|
||||
expect(error.name).toBe('ValidationError')
|
||||
expect(error.message).toBe('Invalid item')
|
||||
})
|
||||
})
|
||||
|
||||
describe('Registry', () => {
|
||||
let registry: TestRegistry
|
||||
|
||||
beforeEach(() => {
|
||||
registry = new TestRegistry()
|
||||
})
|
||||
|
||||
it('should register an item and return a new ID', () => {
|
||||
const id = registry.add('test-item')
|
||||
expect(id).toBe(1)
|
||||
expect(registry.get(1)).toBe('test-item')
|
||||
})
|
||||
|
||||
it('should throw DuplicateItemError when registering with an existing ID', () => {
|
||||
// @ts-expect-error - Spying on protected 'generateId' to test duplicate handling.
|
||||
const generateIdSpy = vi.spyOn(registry, 'generateId').mockReturnValue(1)
|
||||
registry.add('test-item') // This will register with ID 1.
|
||||
expect(() => registry.add('another-item')).toThrow(DuplicateItemError)
|
||||
generateIdSpy.mockRestore()
|
||||
})
|
||||
|
||||
it('should deregister an item and return it', () => {
|
||||
const id = registry.add('test-item')
|
||||
const deregisteredItem = registry.remove(id)
|
||||
expect(deregisteredItem).toBe('test-item')
|
||||
expect(registry.get(id)).toBeUndefined()
|
||||
})
|
||||
|
||||
it('should throw ItemNotFoundError when deregistering a non-existent item', () => {
|
||||
expect(() => registry.remove(999)).toThrow(ItemNotFoundError)
|
||||
})
|
||||
|
||||
it('should get an item by its ID', () => {
|
||||
const id = registry.add('test-item')
|
||||
const foundItem = registry.get(id)
|
||||
expect(foundItem).toBe('test-item')
|
||||
})
|
||||
|
||||
it('should return undefined when getting a non-existent item', () => {
|
||||
const foundItem = registry.get(999)
|
||||
expect(foundItem).toBeUndefined()
|
||||
})
|
||||
|
||||
it('should find an item using a predicate', () => {
|
||||
registry.add('item-a')
|
||||
registry.add('item-b')
|
||||
const foundItem = registry.find((item) => item.includes('b'))
|
||||
expect(foundItem).toBe('item-b')
|
||||
})
|
||||
|
||||
it('should return undefined when no item matches the predicate', () => {
|
||||
registry.add('item-a')
|
||||
const foundItem = registry.find((item) => item.includes('c'))
|
||||
expect(foundItem).toBeUndefined()
|
||||
})
|
||||
|
||||
it('should return all keys', () => {
|
||||
registry.add('item-1')
|
||||
registry.add('item-2')
|
||||
expect(registry.keys()).toEqual([1, 2])
|
||||
})
|
||||
|
||||
it('should return all values', () => {
|
||||
registry.add('item-1')
|
||||
registry.add('item-2')
|
||||
expect(registry.values()).toEqual(['item-1', 'item-2'])
|
||||
})
|
||||
|
||||
it('should return all key-value pairs', () => {
|
||||
registry.add('item-1')
|
||||
registry.add('item-2')
|
||||
expect(registry.pairs()).toEqual([
|
||||
[1, 'item-1'],
|
||||
[2, 'item-2'],
|
||||
])
|
||||
})
|
||||
|
||||
it('should clear all items from the registry', () => {
|
||||
registry.add('item-1')
|
||||
registry.clear()
|
||||
expect(registry.keys()).toEqual([])
|
||||
expect(registry.values()).toEqual([])
|
||||
})
|
||||
|
||||
it('should register multiple items', () => {
|
||||
const ids = registry.addAll(['item-a', 'item-b'])
|
||||
expect(ids).toEqual([1, 2])
|
||||
expect(registry.values()).toEqual(['item-a', 'item-b'])
|
||||
})
|
||||
|
||||
it('should deregister multiple items', () => {
|
||||
const ids = registry.addAll(['item-a', 'item-b', 'item-c'])
|
||||
const deregisteredItems = registry.removeAll([ids[0]!, ids[2]!])
|
||||
expect(deregisteredItems).toEqual(['item-a', 'item-c'])
|
||||
expect(registry.values()).toEqual(['item-b'])
|
||||
})
|
||||
|
||||
it('should find and deregister an item', () => {
|
||||
registry.add('item-a')
|
||||
registry.add('item-b')
|
||||
const deregisteredItem = registry.findRemove((item) => item.includes('a'))
|
||||
expect(deregisteredItem).toBe('item-a')
|
||||
expect(registry.values()).toEqual(['item-b'])
|
||||
})
|
||||
|
||||
it('should return undefined from findRemove if no item matches', () => {
|
||||
const removedItem = registry.findRemove((item) => item.includes('c'))
|
||||
expect(removedItem).toBeUndefined()
|
||||
})
|
||||
|
||||
it('should call the validate method on register', () => {
|
||||
// @ts-expect-error - Spying on protected 'validate' to confirm it is called.
|
||||
const validateSpy = vi.spyOn(registry, 'validate')
|
||||
registry.add('a-valid-item')
|
||||
expect(validateSpy).toHaveBeenCalledWith('a-valid-item')
|
||||
validateSpy.mockRestore()
|
||||
})
|
||||
|
||||
it('should throw a validation error for an invalid item', () => {
|
||||
expect(() => registry.add('')).toThrow(ValidationError)
|
||||
})
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,208 @@
|
||||
import { Registry, ValidationError } from './registry.js'
|
||||
import type { Tool, ToolStreamGenerator } from '../tools/tool.js'
|
||||
|
||||
/**
|
||||
* A concrete implementation of the Registry for managing Tool instances.
|
||||
* It adds validation for tool properties and ensures unique tool names.
|
||||
*/
|
||||
export class ToolRegistry extends Registry<Tool, Tool> {
|
||||
/**
|
||||
* Generates a unique identifier for a Tool.
|
||||
* @override
|
||||
* @returns The tool itself as the identifier.
|
||||
*/
|
||||
protected generateId(tool: Tool): Tool {
|
||||
return tool
|
||||
}
|
||||
|
||||
/**
|
||||
* Validates a tool before it is registered.
|
||||
* @override
|
||||
* @param tool - The tool to be validated.
|
||||
* @throws ValidationError If the tool's properties are invalid or its name is already registered.
|
||||
*/
|
||||
protected validate(tool: Tool): void {
|
||||
// Validate tool name is a string
|
||||
if (typeof tool.name !== 'string') {
|
||||
throw new ValidationError('Tool name must be a string')
|
||||
}
|
||||
|
||||
// Validate tool name length (1-64 characters)
|
||||
if (tool.name.length < 1 || tool.name.length > 64) {
|
||||
throw new ValidationError('Tool name must be between 1 and 64 characters')
|
||||
}
|
||||
|
||||
// Validate tool name pattern
|
||||
const validNamePattern = /^[a-zA-Z0-9_-]+$/
|
||||
if (!validNamePattern.test(tool.name)) {
|
||||
throw new ValidationError('Tool name must contain only alphanumeric characters, hyphens, and underscores')
|
||||
}
|
||||
|
||||
// Validate tool description if present
|
||||
if (tool.description !== undefined && tool.description !== null) {
|
||||
if (typeof tool.description !== 'string' || tool.description.length < 1) {
|
||||
throw new ValidationError('Tool description must be a non-empty string')
|
||||
}
|
||||
}
|
||||
|
||||
// Check for duplicate names
|
||||
if (this.values().some((t) => t.name === tool.name)) {
|
||||
throw new ValidationError(`Tool with name '${tool.name}' already registered`)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Retrieves the first tool that matches the given name.
|
||||
* @param name - The name of the tool to retrieve.
|
||||
* @returns The tool if found, otherwise undefined.
|
||||
*/
|
||||
public getByName(name: string): Tool | undefined {
|
||||
return this.values().find((tool) => tool.name === name)
|
||||
}
|
||||
|
||||
/**
|
||||
* Finds and removes the first tool that matches the given name.
|
||||
* If multiple tools have the same name, only the first one found is removed.
|
||||
* @param name - The name of the tool to remove.
|
||||
*/
|
||||
public removeByName(name: string): void {
|
||||
this.findRemove((tool) => tool.name === name)
|
||||
}
|
||||
}
|
||||
|
||||
// Unit tests
|
||||
if (import.meta.vitest) {
|
||||
const { describe, it, expect, beforeEach } = import.meta.vitest
|
||||
|
||||
// Mock Tool definition for testing purposes
|
||||
const createMockTool = (overrides: Partial<Tool> = {}): Tool => ({
|
||||
name: 'valid-tool',
|
||||
description: 'A valid tool description.',
|
||||
toolSpec: {
|
||||
name: 'valid-tool',
|
||||
description: 'A valid tool description.',
|
||||
inputSchema: { type: 'object', properties: {} },
|
||||
},
|
||||
stream: async function* (): ToolStreamGenerator {
|
||||
// Mock stream implementation
|
||||
yield { type: 'toolStreamEvent' as const, data: 'mock data' }
|
||||
return { toolUseId: '', status: 'success' as const, content: [] }
|
||||
},
|
||||
...overrides,
|
||||
})
|
||||
|
||||
describe('ToolRegistry', () => {
|
||||
let registry: ToolRegistry
|
||||
|
||||
beforeEach(() => {
|
||||
registry = new ToolRegistry()
|
||||
})
|
||||
|
||||
it('should register a valid tool successfully', () => {
|
||||
const tool = createMockTool()
|
||||
expect(() => registry.add(tool)).not.toThrow()
|
||||
expect(registry.values()).toHaveLength(1)
|
||||
expect(registry.values()[0]?.name).toBe('valid-tool')
|
||||
})
|
||||
|
||||
it('should throw ValidationError for a duplicate tool name', () => {
|
||||
const tool1 = createMockTool({ name: 'duplicate-name' })
|
||||
const tool2 = createMockTool({ name: 'duplicate-name' })
|
||||
registry.add(tool1)
|
||||
|
||||
expect(() => registry.add(tool2)).toThrow(ValidationError)
|
||||
expect(() => registry.add(tool2)).toThrow("Tool with name 'duplicate-name' already registered")
|
||||
})
|
||||
|
||||
it('should throw ValidationError for an invalid tool name pattern', () => {
|
||||
const tool = createMockTool({ name: 'invalid name!' })
|
||||
expect(() => registry.add(tool)).toThrow(ValidationError)
|
||||
expect(() => registry.add(tool)).toThrow(
|
||||
'Tool name must contain only alphanumeric characters, hyphens, and underscores'
|
||||
)
|
||||
})
|
||||
|
||||
it('should throw ValidationError for a tool name that is too long', () => {
|
||||
const longName = 'a'.repeat(65)
|
||||
const tool = createMockTool({ name: longName })
|
||||
expect(() => registry.add(tool)).toThrow(ValidationError)
|
||||
expect(() => registry.add(tool)).toThrow('Tool name must be between 1 and 64 characters')
|
||||
})
|
||||
|
||||
it('should throw ValidationError for a tool name that is too short', () => {
|
||||
const tool = createMockTool({ name: '' })
|
||||
expect(() => registry.add(tool)).toThrow(ValidationError)
|
||||
expect(() => registry.add(tool)).toThrow('Tool name must be between 1 and 64 characters')
|
||||
})
|
||||
|
||||
it('should throw ValidationError for an invalid description', () => {
|
||||
// @ts-expect-error - Testing invalid type for description
|
||||
const tool = createMockTool({ description: 123 })
|
||||
expect(() => registry.add(tool)).toThrow(ValidationError)
|
||||
expect(() => registry.add(tool)).toThrow('Tool description must be a non-empty string')
|
||||
})
|
||||
|
||||
it('should throw ValidationError for an empty string description', () => {
|
||||
const tool = createMockTool({ description: '' })
|
||||
expect(() => registry.add(tool)).toThrow(ValidationError)
|
||||
expect(() => registry.add(tool)).toThrow('Tool description must be a non-empty string')
|
||||
})
|
||||
|
||||
it('should allow a tool with a null or undefined description', () => {
|
||||
const tool1 = createMockTool()
|
||||
// @ts-expect-error - Testing explicit undefined description
|
||||
tool1.description = undefined
|
||||
|
||||
const tool2 = createMockTool()
|
||||
tool2.name = 'another-valid-tool'
|
||||
// @ts-expect-error - Testing explicit null description
|
||||
tool2.description = null
|
||||
|
||||
expect(() => registry.add(tool1)).not.toThrow()
|
||||
expect(() => registry.add(tool2)).not.toThrow()
|
||||
})
|
||||
|
||||
it('should retrieve a tool by its name', () => {
|
||||
const tool = createMockTool({ name: 'find-me' })
|
||||
registry.add(tool)
|
||||
const foundTool = registry.getByName('find-me')
|
||||
expect(foundTool).toBe(tool)
|
||||
})
|
||||
|
||||
it('should return undefined when getting a tool by a name that does not exist', () => {
|
||||
const foundTool = registry.getByName('non-existent')
|
||||
expect(foundTool).toBeUndefined()
|
||||
})
|
||||
|
||||
it('should remove a tool by its name', () => {
|
||||
const tool = createMockTool({ name: 'remove-me' })
|
||||
registry.add(tool)
|
||||
expect(registry.getByName('remove-me')).toBeDefined()
|
||||
registry.removeByName('remove-me')
|
||||
expect(registry.getByName('remove-me')).toBeUndefined()
|
||||
})
|
||||
|
||||
it('should not throw when removing a tool by a name that does not exist', () => {
|
||||
expect(() => registry.removeByName('non-existent')).not.toThrow()
|
||||
})
|
||||
|
||||
it('should generate a valid ToolIdentifier', () => {
|
||||
const tool = createMockTool()
|
||||
const id = registry['generateId'](tool)
|
||||
expect(id).toBe(tool)
|
||||
})
|
||||
|
||||
it('should register a tool with a name at the maximum length', () => {
|
||||
const longName = 'a'.repeat(64)
|
||||
const tool = createMockTool({ name: longName })
|
||||
expect(() => registry.add(tool)).not.toThrow()
|
||||
})
|
||||
|
||||
it('should throw ValidationError for a non-string tool name', () => {
|
||||
// @ts-expect-error - Testing invalid type for name
|
||||
const tool = createMockTool({ name: 123 })
|
||||
expect(() => registry.add(tool)).toThrow(ValidationError)
|
||||
expect(() => registry.add(tool)).toThrow('Tool name must be a string')
|
||||
})
|
||||
})
|
||||
}
|
||||
@@ -1,311 +0,0 @@
|
||||
import { describe, it, expect } from 'vitest'
|
||||
import { ToolRegistry } from '../registry.js'
|
||||
import type { Tool, ToolStreamEvent } from '../tool.js'
|
||||
import type { ToolResult, ToolSpec } from '../types.js'
|
||||
import { TextBlock } from '../../types/messages.js'
|
||||
|
||||
/**
|
||||
* Helper function to create a mock Tool for testing.
|
||||
* Creates a minimal Tool implementation with configurable name and description.
|
||||
*/
|
||||
function createMockTool(name: string, description = 'Test tool description'): Tool {
|
||||
const toolSpec: ToolSpec = {
|
||||
name,
|
||||
description,
|
||||
inputSchema: { type: 'object' },
|
||||
}
|
||||
|
||||
return {
|
||||
toolName: name,
|
||||
description,
|
||||
toolSpec,
|
||||
// eslint-disable-next-line require-yield
|
||||
async *stream(): AsyncGenerator<ToolStreamEvent, ToolResult, unknown> {
|
||||
return {
|
||||
toolUseId: 'test-id',
|
||||
status: 'success',
|
||||
content: [new TextBlock('test result')],
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
describe('ToolRegistry', () => {
|
||||
describe('constructor', () => {
|
||||
it('creates an empty registry', () => {
|
||||
const registry = new ToolRegistry()
|
||||
expect(registry).toBeDefined()
|
||||
expect(registry.list()).toEqual([])
|
||||
})
|
||||
})
|
||||
|
||||
describe('register', () => {
|
||||
describe('when registering a single tool', () => {
|
||||
it('adds the tool to the registry', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool = createMockTool('testTool')
|
||||
registry.register(tool)
|
||||
|
||||
const retrieved = registry.get('testTool')
|
||||
expect(retrieved).toBe(tool)
|
||||
})
|
||||
|
||||
it('allows retrieval with get()', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool = createMockTool('calculator')
|
||||
registry.register(tool)
|
||||
|
||||
const retrieved = registry.get('calculator')
|
||||
expect(retrieved?.toolName).toBe('calculator')
|
||||
})
|
||||
})
|
||||
|
||||
describe('when registering multiple tools', () => {
|
||||
it('adds all tools to the registry', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool1 = createMockTool('multiTool1')
|
||||
const tool2 = createMockTool('multiTool2')
|
||||
const tool3 = createMockTool('multiTool3')
|
||||
|
||||
registry.register([tool1, tool2, tool3])
|
||||
|
||||
expect(registry.get('multiTool1')).toBe(tool1)
|
||||
expect(registry.get('multiTool2')).toBe(tool2)
|
||||
expect(registry.get('multiTool3')).toBe(tool3)
|
||||
})
|
||||
|
||||
it('allows retrieval of each tool', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tools = [createMockTool('alpha'), createMockTool('beta'), createMockTool('gamma')]
|
||||
|
||||
registry.register(tools)
|
||||
|
||||
const allTools = registry.list()
|
||||
expect(allTools).toHaveLength(3)
|
||||
expect(allTools[0]?.toolName).toBe('alpha')
|
||||
expect(allTools[1]?.toolName).toBe('beta')
|
||||
expect(allTools[2]?.toolName).toBe('gamma')
|
||||
})
|
||||
})
|
||||
|
||||
describe('when registering a duplicate tool name', () => {
|
||||
it('throws an error with descriptive message', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool1 = createMockTool('duplicateTool')
|
||||
const tool2 = createMockTool('duplicateTool')
|
||||
|
||||
registry.register(tool1)
|
||||
|
||||
expect(() => registry.register(tool2)).toThrow("Tool with name 'duplicateTool' already registered")
|
||||
})
|
||||
})
|
||||
|
||||
describe('when registering a tool with empty name', () => {
|
||||
it('throws an error with descriptive message', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool = createMockTool('')
|
||||
|
||||
expect(() => registry.register(tool)).toThrow('Tool name must be between 1 and 64 characters')
|
||||
})
|
||||
})
|
||||
|
||||
describe('when registering a tool with name too long', () => {
|
||||
it('throws an error with descriptive message', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const longName = 'a'.repeat(65)
|
||||
const tool = createMockTool(longName)
|
||||
|
||||
expect(() => registry.register(tool)).toThrow('Tool name must be between 1 and 64 characters')
|
||||
})
|
||||
})
|
||||
|
||||
describe('when registering a tool with invalid name characters', () => {
|
||||
it('throws an error for spaces', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool = createMockTool('invalid name')
|
||||
|
||||
expect(() => registry.register(tool)).toThrow(
|
||||
'Tool name must contain only alphanumeric characters, hyphens, and underscores'
|
||||
)
|
||||
})
|
||||
|
||||
it('throws an error for special characters', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool = createMockTool('invalid@name!')
|
||||
|
||||
expect(() => registry.register(tool)).toThrow(
|
||||
'Tool name must contain only alphanumeric characters, hyphens, and underscores'
|
||||
)
|
||||
})
|
||||
|
||||
it('allows valid characters', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool1 = createMockTool('valid_name')
|
||||
const tool2 = createMockTool('valid-name')
|
||||
const tool3 = createMockTool('ValidName123')
|
||||
|
||||
expect(() => {
|
||||
registry.register([tool1, tool2, tool3])
|
||||
}).not.toThrow()
|
||||
|
||||
expect(registry.list()).toHaveLength(3)
|
||||
})
|
||||
})
|
||||
|
||||
describe('when registering a tool with empty description', () => {
|
||||
it('throws an error with descriptive message', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool = createMockTool('validName', '')
|
||||
|
||||
expect(() => registry.register(tool)).toThrow('Tool description must be a non-empty string')
|
||||
})
|
||||
})
|
||||
|
||||
describe('when registering a tool with valid name at boundary', () => {
|
||||
it('accepts name with 1 character', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool = createMockTool('a')
|
||||
|
||||
expect(() => registry.register(tool)).not.toThrow()
|
||||
expect(registry.get('a')).toBe(tool)
|
||||
})
|
||||
|
||||
it('accepts name with 64 characters', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const name64 = 'a'.repeat(64)
|
||||
const tool = createMockTool(name64)
|
||||
|
||||
expect(() => registry.register(tool)).not.toThrow()
|
||||
expect(registry.get(name64)).toBe(tool)
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe('get', () => {
|
||||
describe('when tool exists', () => {
|
||||
it('returns the tool instance', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool = createMockTool('existingTool')
|
||||
registry.register(tool)
|
||||
|
||||
const retrieved = registry.get('existingTool')
|
||||
expect(retrieved).toBe(tool)
|
||||
})
|
||||
})
|
||||
|
||||
describe('when tool does not exist', () => {
|
||||
it('returns undefined', () => {
|
||||
const registry = new ToolRegistry()
|
||||
expect(registry.get('nonExistentTool')).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
describe('when registry is empty', () => {
|
||||
it('returns undefined', () => {
|
||||
const registry = new ToolRegistry()
|
||||
expect(registry.get('anyTool')).toBeUndefined()
|
||||
})
|
||||
})
|
||||
})
|
||||
describe('remove', () => {
|
||||
describe('when removing an existing tool', () => {
|
||||
it('removes the tool from registry', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool = createMockTool('removableTool')
|
||||
registry.register(tool)
|
||||
|
||||
registry.remove('removableTool')
|
||||
|
||||
expect(registry.list()).toEqual([])
|
||||
})
|
||||
|
||||
it('get() returns undefined after removal', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool = createMockTool('temporaryTool')
|
||||
registry.register(tool)
|
||||
|
||||
registry.remove('temporaryTool')
|
||||
|
||||
expect(registry.get('temporaryTool')).toBeUndefined()
|
||||
})
|
||||
})
|
||||
|
||||
describe('when tool does not exist', () => {
|
||||
it('throws an error with descriptive message', () => {
|
||||
const registry = new ToolRegistry()
|
||||
expect(() => registry.remove('nonExistent')).toThrow("Tool with name 'nonExistent' not found")
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe('list', () => {
|
||||
describe('when registry has tools', () => {
|
||||
it('returns all registered tools', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool1 = createMockTool('listTool1')
|
||||
const tool2 = createMockTool('listTool2')
|
||||
const tool3 = createMockTool('listTool3')
|
||||
|
||||
registry.register([tool1, tool2, tool3])
|
||||
|
||||
const tools = registry.list()
|
||||
expect(tools).toEqual([tool1, tool2, tool3])
|
||||
})
|
||||
|
||||
it('returns a copy (mutation does not affect registry)', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool1 = createMockTool('copyTool1')
|
||||
const tool2 = createMockTool('copyTool2')
|
||||
|
||||
registry.register([tool1, tool2])
|
||||
|
||||
const tools = registry.list()
|
||||
tools.pop() // Mutate the returned array
|
||||
|
||||
// Verify registry still has both tools
|
||||
expect(registry.list()).toEqual([tool1, tool2])
|
||||
})
|
||||
|
||||
it('returns tools in insertion order', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const toolA = createMockTool('orderA')
|
||||
const toolB = createMockTool('orderB')
|
||||
const toolC = createMockTool('orderC')
|
||||
|
||||
registry.register(toolA)
|
||||
registry.register(toolB)
|
||||
registry.register(toolC)
|
||||
|
||||
const tools = registry.list()
|
||||
expect(tools).toHaveLength(3)
|
||||
expect(tools[0]?.toolName).toBe('orderA')
|
||||
expect(tools[1]?.toolName).toBe('orderB')
|
||||
expect(tools[2]?.toolName).toBe('orderC')
|
||||
})
|
||||
})
|
||||
|
||||
describe('when registry is empty', () => {
|
||||
it('returns an empty array', () => {
|
||||
const registry = new ToolRegistry()
|
||||
expect(registry.list()).toEqual([])
|
||||
})
|
||||
})
|
||||
|
||||
describe('after adding and removing tools', () => {
|
||||
it('reflects current state', () => {
|
||||
const registry = new ToolRegistry()
|
||||
const tool1 = createMockTool('stateTool1')
|
||||
const tool2 = createMockTool('stateTool2')
|
||||
const tool3 = createMockTool('stateTool3')
|
||||
|
||||
registry.register([tool1, tool2, tool3])
|
||||
expect(registry.list()).toHaveLength(3)
|
||||
|
||||
registry.remove('stateTool2')
|
||||
const tools = registry.list()
|
||||
expect(tools).toHaveLength(2)
|
||||
expect(tools).toEqual([tool1, tool3])
|
||||
})
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -14,10 +14,10 @@ describe('FunctionTool', () => {
|
||||
inputSchema: { type: 'object' },
|
||||
callback: (): string => 'result',
|
||||
})
|
||||
expect(tool.toolName).toBeTruthy()
|
||||
expect(typeof tool.toolName).toBe('string')
|
||||
expect(tool.toolName.length).toBeGreaterThan(0)
|
||||
expect(tool.toolName).toBe('testTool')
|
||||
expect(tool.name).toBeTruthy()
|
||||
expect(typeof tool.name).toBe('string')
|
||||
expect(tool.name.length).toBeGreaterThan(0)
|
||||
expect(tool.name).toBe('testTool')
|
||||
})
|
||||
|
||||
it('has a non-empty description', () => {
|
||||
@@ -62,7 +62,7 @@ describe('FunctionTool', () => {
|
||||
inputSchema: { type: 'object' },
|
||||
callback: (): string => 'result',
|
||||
})
|
||||
expect(tool.toolName).toBe(tool.toolSpec.name)
|
||||
expect(tool.name).toBe(tool.toolSpec.name)
|
||||
})
|
||||
|
||||
it('has matching description and toolSpec.description', () => {
|
||||
@@ -943,7 +943,7 @@ describe('Tool interface backwards compatibility', () => {
|
||||
})
|
||||
|
||||
// Verify interface properties exist
|
||||
expect(tool).toHaveProperty('toolName')
|
||||
expect(tool).toHaveProperty('name')
|
||||
expect(tool).toHaveProperty('description')
|
||||
expect(tool).toHaveProperty('toolSpec')
|
||||
expect(tool).toHaveProperty('stream')
|
||||
|
||||
@@ -14,7 +14,7 @@ describe('tool', () => {
|
||||
callback: (input) => input.value,
|
||||
})
|
||||
|
||||
expect(myTool.toolName).toBe('testTool')
|
||||
expect(myTool.name).toBe('testTool')
|
||||
expect(myTool.description).toBe('Test description')
|
||||
expect(myTool.toolSpec).toEqual({
|
||||
name: 'testTool',
|
||||
@@ -37,7 +37,7 @@ describe('tool', () => {
|
||||
callback: (input) => input.value,
|
||||
})
|
||||
|
||||
expect(myTool.toolName).toBe('testTool')
|
||||
expect(myTool.name).toBe('testTool')
|
||||
expect(myTool.description).toBe('')
|
||||
})
|
||||
})
|
||||
|
||||
@@ -90,7 +90,7 @@ export class FunctionTool implements Tool {
|
||||
/**
|
||||
* The unique name of the tool.
|
||||
*/
|
||||
readonly toolName: string
|
||||
readonly name: string
|
||||
|
||||
/**
|
||||
* Human-readable description of what the tool does.
|
||||
@@ -127,7 +127,7 @@ export class FunctionTool implements Tool {
|
||||
* ```
|
||||
*/
|
||||
constructor(config: FunctionToolConfig) {
|
||||
this.toolName = config.name
|
||||
this.name = config.name
|
||||
this.description = config.description
|
||||
this.toolSpec = {
|
||||
name: config.name,
|
||||
|
||||
@@ -1,95 +0,0 @@
|
||||
import type { Tool } from './tool.js'
|
||||
|
||||
/**
|
||||
* Registry for managing Tool instances.
|
||||
*/
|
||||
export class ToolRegistry {
|
||||
private readonly _tools: Map<string, Tool>
|
||||
|
||||
/**
|
||||
* Creates a new ToolRegistry instance with an empty registry.
|
||||
*/
|
||||
constructor() {
|
||||
this._tools = new Map<string, Tool>()
|
||||
}
|
||||
|
||||
/**
|
||||
* Registers one or more tools with the registry.
|
||||
* Accepts single Tool or array of Tools for convenience.
|
||||
*
|
||||
* @param tool - Single Tool instance or array of Tool instances to register
|
||||
* @throws If a tool with duplicate name already exists
|
||||
* @throws If tool name is invalid (must be 1-64 chars, alphanumeric with hyphens/underscores)
|
||||
* @throws If tool description is empty
|
||||
*/
|
||||
public register(tool: Tool | Tool[]): void {
|
||||
const tools = Array.isArray(tool) ? tool : [tool]
|
||||
|
||||
for (const t of tools) {
|
||||
// Validate tool name is a string
|
||||
if (typeof t.toolName !== 'string') {
|
||||
throw new Error('Tool name must be a string')
|
||||
}
|
||||
|
||||
// Validate tool name length (1-64 characters)
|
||||
if (t.toolName.length < 1 || t.toolName.length > 64) {
|
||||
throw new Error('Tool name must be between 1 and 64 characters')
|
||||
}
|
||||
|
||||
// Validate tool name pattern (alphanumeric with hyphens and underscores)
|
||||
const validNamePattern = /^[a-zA-Z0-9_-]+$/
|
||||
if (!validNamePattern.test(t.toolName)) {
|
||||
throw new Error('Tool name must contain only alphanumeric characters, hyphens, and underscores')
|
||||
}
|
||||
|
||||
// Validate tool description if present
|
||||
if (t.description !== undefined && t.description !== null) {
|
||||
if (typeof t.description !== 'string' || t.description.length < 1) {
|
||||
throw new Error('Tool description must be a non-empty string')
|
||||
}
|
||||
}
|
||||
|
||||
// Check for duplicate names
|
||||
if (this._tools.has(t.toolName)) {
|
||||
throw new Error(`Tool with name '${t.toolName}' already registered`)
|
||||
}
|
||||
|
||||
this._tools.set(t.toolName, t)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Retrieves a tool by its unique name.
|
||||
*
|
||||
* @param name - The unique name of the tool to retrieve
|
||||
* @returns The Tool instance, or undefined if not found
|
||||
*/
|
||||
public get(name: string): Tool | undefined {
|
||||
return this._tools.get(name)
|
||||
}
|
||||
|
||||
/**
|
||||
* Removes a tool from the registry.
|
||||
*
|
||||
* @param name - The name of the tool to remove
|
||||
* @throws If tool with given name doesn't exist
|
||||
*/
|
||||
public remove(name: string): void {
|
||||
// Check if tool exists
|
||||
if (!this._tools.has(name)) {
|
||||
throw new Error(`Tool with name '${name}' not found`)
|
||||
}
|
||||
|
||||
this._tools.delete(name)
|
||||
}
|
||||
|
||||
/**
|
||||
* Returns all registered tools as an array.
|
||||
* Returns a copy of the internal array to prevent external mutation.
|
||||
*
|
||||
* @returns Array of all registered Tool instances, or empty array if no tools registered
|
||||
*/
|
||||
public list(): Tool[] {
|
||||
return Array.from(this._tools.values())
|
||||
}
|
||||
}
|
||||
+1
-1
@@ -90,7 +90,7 @@ export interface Tool {
|
||||
* The unique name of the tool.
|
||||
* This MUST match the name in the toolSpec.
|
||||
*/
|
||||
toolName: string
|
||||
name: string
|
||||
|
||||
/**
|
||||
* Human-readable description of what the tool does.
|
||||
|
||||
+1
-1
@@ -85,6 +85,6 @@ export interface ToolUse {
|
||||
*
|
||||
* - `{ auto: {} }` - Let the model decide whether to use a tool
|
||||
* - `{ any: {} }` - Force the model to use one of the available tools
|
||||
* - `{ tool: { name: 'toolName' } }` - Force the model to use a specific tool
|
||||
* - `{ tool: { name: 'name' } }` - Force the model to use a specific tool
|
||||
*/
|
||||
export type ToolChoice = { auto: Record<string, never> } | { any: Record<string, never> } | { tool: { name: string } }
|
||||
|
||||
@@ -106,7 +106,7 @@ export function tool<TInput extends z.ZodType, TReturn extends JSONValue = JSONV
|
||||
|
||||
// Create an invokable tool that extends the FunctionTool
|
||||
const invokableTool: InvokableTool<z.infer<TInput>, TReturn> = {
|
||||
toolName: functionTool.toolName,
|
||||
name: functionTool.name,
|
||||
description: functionTool.description,
|
||||
toolSpec: functionTool.toolSpec,
|
||||
|
||||
|
||||
@@ -47,6 +47,44 @@ export class Message {
|
||||
this.role = data.role
|
||||
this.content = data.content
|
||||
}
|
||||
|
||||
/**
|
||||
* Creates a Message instance from MessageData.
|
||||
*/
|
||||
public static fromMessageData(data: MessageData): Message {
|
||||
const contentBlocks: ContentBlock[] = data.content.map((block) => {
|
||||
if ('text' in block) {
|
||||
return new TextBlock(block.text)
|
||||
} else if ('toolUse' in block) {
|
||||
return new ToolUseBlock(block.toolUse)
|
||||
} else if ('toolResult' in block) {
|
||||
return new ToolResultBlock({
|
||||
toolUseId: block.toolResult.toolUseId,
|
||||
status: block.toolResult.status,
|
||||
content: block.toolResult.content.map((contentItem) => {
|
||||
if ('text' in contentItem) {
|
||||
return new TextBlock(contentItem.text)
|
||||
} else if ('json' in contentItem) {
|
||||
return new JsonBlock(contentItem)
|
||||
} else {
|
||||
throw new Error('Unknown ToolResultContentData type')
|
||||
}
|
||||
}),
|
||||
})
|
||||
} else if ('reasoning' in block) {
|
||||
return new ReasoningBlock(block.reasoning)
|
||||
} else if ('cachePoint' in block) {
|
||||
return new CachePointBlock(block.cachePoint)
|
||||
} else {
|
||||
throw new Error('Unknown ContentBlockData type')
|
||||
}
|
||||
})
|
||||
|
||||
return new Message({
|
||||
role: data.role,
|
||||
content: contentBlocks,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import { describe, it, expect } from 'vitest'
|
||||
import { OpenAIModel } from '@strands-agents/sdk/openai'
|
||||
import { ContextWindowOverflowError } from '@strands-agents/sdk'
|
||||
import { ContextWindowOverflowError, ToolResultBlock } from '@strands-agents/sdk'
|
||||
import { Message } from '@strands-agents/sdk'
|
||||
import type { ToolSpec } from '@strands-agents/sdk'
|
||||
|
||||
@@ -234,12 +234,11 @@ describe.skipIf(!hasApiKey)('OpenAIModel Integration Tests', () => {
|
||||
new Message({
|
||||
role: 'user',
|
||||
content: [
|
||||
{
|
||||
type: 'toolResultBlock',
|
||||
new ToolResultBlock({
|
||||
toolUseId: toolUseId!,
|
||||
content: [{ type: 'textBlock', text: '42' }],
|
||||
status: 'success',
|
||||
},
|
||||
}),
|
||||
],
|
||||
}),
|
||||
]
|
||||
|
||||
+2
-1
@@ -25,7 +25,8 @@
|
||||
"isolatedModules": true,
|
||||
"verbatimModuleSyntax": true,
|
||||
"sourceMap": true,
|
||||
"removeComments": false
|
||||
"removeComments": false,
|
||||
"types": ["vitest/importMeta"]
|
||||
},
|
||||
"include": ["src/**/*"],
|
||||
"exclude": ["node_modules", "dist"]
|
||||
|
||||
+7
-3
@@ -6,6 +6,7 @@ export default defineConfig({
|
||||
{
|
||||
test: {
|
||||
include: ['src/**/__tests__/**/*.test.ts'],
|
||||
includeSource: ['src/**/*.{js,ts}'],
|
||||
name: { label: 'unit-node', color: 'green' },
|
||||
typecheck: {
|
||||
enabled: true,
|
||||
@@ -34,12 +35,12 @@ export default defineConfig({
|
||||
name: { label: 'integ', color: 'magenta' },
|
||||
testTimeout: 30000,
|
||||
globalSetup: './tests_integ/integ-setup.ts',
|
||||
sequence: {
|
||||
concurrent: true,
|
||||
},
|
||||
},
|
||||
},
|
||||
],
|
||||
sequence: {
|
||||
concurrent: true,
|
||||
},
|
||||
typecheck: {
|
||||
enabled: true,
|
||||
},
|
||||
@@ -57,4 +58,7 @@ export default defineConfig({
|
||||
},
|
||||
environment: 'node',
|
||||
},
|
||||
define: {
|
||||
'import.meta.vitest': 'undefined',
|
||||
},
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user