Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
119 changes: 119 additions & 0 deletions typescript/tests/agents/BedrockFlowsAgent.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,119 @@
import { BedrockFlowsAgent, BedrockFlowsAgentOptions } from '../../src/agents/bedrockFlowsAgent';
import { BedrockAgentRuntimeClient, InvokeFlowCommand } from '@aws-sdk/client-bedrock-agent-runtime';
import { ConversationMessage, ParticipantRole } from '../../src/types';

jest.mock('@aws-sdk/client-bedrock-agent-runtime');
jest.mock('../../src/common/src/awsSdkUtils', () => ({
addUserAgentMiddleware: jest.fn(),
}));

describe('BedrockFlowsAgent', () => {
let mockClient: jest.Mocked<BedrockAgentRuntimeClient>;

const defaultOptions: BedrockFlowsAgentOptions = {
name: 'TestFlowsAgent',
description: 'Test BedrockFlowsAgent',
flowIdentifier: 'flow-123',
flowAliasIdentifier: 'alias-456',
};

function makeResponseStream(document: string) {
async function* gen() {
yield { flowOutputEvent: { content: { document } } };
}
return { responseStream: gen() };
}

beforeEach(() => {
jest.clearAllMocks();
mockClient = { send: jest.fn(), use: jest.fn(), destroy: jest.fn() } as any;
(BedrockAgentRuntimeClient as jest.Mock).mockImplementation(() => mockClient);
(InvokeFlowCommand as unknown as jest.Mock).mockImplementation((p) => p);
});

describe('constructor', () => {
it('creates a BedrockAgentRuntimeClient with no arguments when no region given', () => {
new BedrockFlowsAgent(defaultOptions);
expect(BedrockAgentRuntimeClient).toHaveBeenCalledWith();
});

it('creates a BedrockAgentRuntimeClient with specified region', () => {
new BedrockFlowsAgent({ ...defaultOptions, region: 'eu-west-1' });
expect(BedrockAgentRuntimeClient).toHaveBeenCalledWith({ region: 'eu-west-1' });
});

it('uses a pre-built client when bedrockAgentClient is provided', () => {
const prebuilt = { send: jest.fn() } as any;
const agent = new BedrockFlowsAgent({ ...defaultOptions, bedrockAgentClient: prebuilt });
expect(agent).toBeDefined();
expect(BedrockAgentRuntimeClient).not.toHaveBeenCalled();
});

it('defaults enableTrace to false', () => {
const agent = new BedrockFlowsAgent(defaultOptions);
expect((agent as any).enableTrace).toBe(false);
});

it('sets enableTrace when provided', () => {
const agent = new BedrockFlowsAgent({ ...defaultOptions, enableTrace: true });
expect((agent as any).enableTrace).toBe(true);
});
});

describe('processRequest', () => {
const userId = 'user1';
const sessionId = 'session1';
const chatHistory: ConversationMessage[] = [];

it('invokes the flow and returns the decoded response', async () => {
mockClient.send.mockResolvedValueOnce(makeResponseStream('Flow answer') as never);

const agent = new BedrockFlowsAgent(defaultOptions);
const result = await agent.processRequest('What is 2+2?', userId, sessionId, chatHistory);

expect(mockClient.send).toHaveBeenCalledTimes(1);
expect(result).toEqual<ConversationMessage>({
role: ParticipantRole.ASSISTANT,
content: [{ text: 'Flow answer' }],
});
});

it('throws when the response stream is missing', async () => {
mockClient.send.mockResolvedValueOnce({ responseStream: null } as never);

const agent = new BedrockFlowsAgent(defaultOptions);
await expect(agent.processRequest('input', userId, sessionId, chatHistory))
.rejects.toThrow(/Error processing request with Bedrock/);
});

it('uses a custom flowInputEncoder when provided', async () => {
const encoder = jest.fn().mockReturnValue('encoded input');
mockClient.send.mockResolvedValueOnce(makeResponseStream('ok') as never);

const agent = new BedrockFlowsAgent({ ...defaultOptions, flowInputEncoder: encoder });
await agent.processRequest('hello', userId, sessionId, chatHistory);

expect(encoder).toHaveBeenCalledWith(agent, 'hello', expect.objectContaining({ userId, sessionId }));
});

it('uses a custom flowOutputDecoder when provided', async () => {
const decoded: ConversationMessage = { role: ParticipantRole.ASSISTANT, content: [{ text: 'custom' }] };
const decoder = jest.fn().mockReturnValue(decoded);
mockClient.send.mockResolvedValueOnce(makeResponseStream('raw') as never);

const agent = new BedrockFlowsAgent({ ...defaultOptions, flowOutputDecoder: decoder });
const result = await agent.processRequest('input', userId, sessionId, chatHistory);

expect(decoder).toHaveBeenCalledWith(agent, 'raw');
expect(result).toEqual(decoded);
});

it('wraps client errors in a descriptive error message', async () => {
mockClient.send.mockRejectedValueOnce(new Error('Network failure') as never);

const agent = new BedrockFlowsAgent(defaultOptions);
await expect(agent.processRequest('input', userId, sessionId, chatHistory))
.rejects.toThrow('Error processing request with Bedrock: Network failure');
});
});
});
136 changes: 136 additions & 0 deletions typescript/tests/agents/ChainAgent.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,136 @@
import { ChainAgent, ChainAgentOptions } from '../../src/agents/chainAgent';
import { Agent, AgentOptions } from '../../src/agents/agent';
import { ConversationMessage, ParticipantRole } from '../../src/types';

// Minimal concrete Agent subclass for testing
class MockAgent extends Agent {
private response: ConversationMessage | AsyncIterable<any>;

constructor(
options: AgentOptions,
response: ConversationMessage | AsyncIterable<any>
) {
super(options);
this.response = response;
}

async processRequest(): Promise<ConversationMessage | AsyncIterable<any>> {
return this.response;
}
}

function makeMessage(text: string): ConversationMessage {
return { role: ParticipantRole.ASSISTANT, content: [{ text }] };
}

function makeAgent(name: string, response: ConversationMessage | AsyncIterable<any>): MockAgent {
return new MockAgent({ name, description: `${name} description` }, response);
}

describe('ChainAgent', () => {
const defaultOptions: ChainAgentOptions = {
name: 'TestChain',
description: 'Test chain agent',
agents: [makeAgent('Agent1', makeMessage('step one output'))],
};

describe('constructor', () => {
it('throws when agents array is empty', () => {
expect(() => new ChainAgent({ ...defaultOptions, agents: [] })).toThrow(
'ChainAgent requires at least one agent in the chain.'
);
});

it('initialises with a single agent', () => {
const chain = new ChainAgent(defaultOptions);
expect(chain.agents).toHaveLength(1);
});

it('uses custom defaultOutput when provided', () => {
const chain = new ChainAgent({ ...defaultOptions, defaultOutput: 'custom fallback' });
expect(chain).toBeDefined();
});
});

describe('processRequest', () => {
const userId = 'user1';
const sessionId = 'session1';
const chatHistory: ConversationMessage[] = [];

it('returns the output of a single-agent chain', async () => {
const chain = new ChainAgent(defaultOptions);
const result = await chain.processRequest('hello', userId, sessionId, chatHistory) as ConversationMessage;

expect(result.role).toBe(ParticipantRole.ASSISTANT);
expect(result.content[0]).toHaveProperty('text', 'step one output');
});

it('pipes output of each agent as input to the next', async () => {
const spy1 = jest.fn().mockResolvedValue(makeMessage('output of agent 1'));
const spy2 = jest.fn().mockResolvedValue(makeMessage('output of agent 2'));

const agent1 = new MockAgent({ name: 'A1', description: 'd' }, makeMessage(''));
const agent2 = new MockAgent({ name: 'A2', description: 'd' }, makeMessage(''));
agent1.processRequest = spy1 as any;
agent2.processRequest = spy2 as any;

const chain = new ChainAgent({ ...defaultOptions, agents: [agent1, agent2] });
await chain.processRequest('initial input', userId, sessionId, chatHistory);

expect(spy1).toHaveBeenCalledWith('initial input', userId, sessionId, chatHistory, undefined);
expect(spy2).toHaveBeenCalledWith('output of agent 1', userId, sessionId, chatHistory, undefined);
});

it('passes additionalParams through to each agent', async () => {
const spy = jest.fn().mockResolvedValue(makeMessage('ok'));
const agent = new MockAgent({ name: 'A', description: 'd' }, makeMessage(''));
agent.processRequest = spy as any;

const chain = new ChainAgent({ ...defaultOptions, agents: [agent] });
const params = { key: 'val' };
await chain.processRequest('input', userId, sessionId, chatHistory, params);

expect(spy).toHaveBeenCalledWith('input', userId, sessionId, chatHistory, params);
});

it('returns default response when an intermediate agent returns no text content', async () => {
const emptyAgent = makeAgent('Empty', { role: ParticipantRole.ASSISTANT, content: [] });
const secondAgent = makeAgent('Second', makeMessage('should not reach'));

const chain = new ChainAgent({ ...defaultOptions, agents: [emptyAgent, secondAgent] });
const result = await chain.processRequest('input', userId, sessionId, chatHistory) as ConversationMessage;

expect(result.content[0]).toHaveProperty('text', 'No output generated from the chain.');
});

it('throws when an agent in the chain errors', async () => {
const errorAgent = new MockAgent({ name: 'ErrorAgent', description: 'd' }, makeMessage(''));
errorAgent.processRequest = jest.fn().mockRejectedValue(new Error('downstream failure')) as any;

const chain = new ChainAgent({ ...defaultOptions, agents: [errorAgent] });
await expect(chain.processRequest('input', userId, sessionId, chatHistory))
.rejects.toMatch(/Error processing request with agent ErrorAgent/);
});

it('allows the last agent to return a streaming response', async () => {
async function* streamGen() { yield 'chunk'; }
const streamingAgent = makeAgent('Streamer', streamGen());

const chain = new ChainAgent({ ...defaultOptions, agents: [streamingAgent] });
const result = await chain.processRequest('input', userId, sessionId, chatHistory);

expect(typeof (result as any)[Symbol.asyncIterator]).toBe('function');
});

it('returns default response when a non-last agent returns a streaming response', async () => {
async function* streamGen() { yield 'chunk'; }
const streamingAgent = makeAgent('Streamer', streamGen());
const nextAgent = makeAgent('Next', makeMessage('next'));

const chain = new ChainAgent({ ...defaultOptions, agents: [streamingAgent, nextAgent] });
const result = await chain.processRequest('input', userId, sessionId, chatHistory) as ConversationMessage;

expect(result.content[0]).toHaveProperty('text', 'No output generated from the chain.');
});
});
});