Testing
The best way to test agents is to pass a MockLanguageModel to the agent instead of a real language model.
import { LanguageModelError } from "@hoangvvo/llm-sdk";import { MockLanguageModel } from "@hoangvvo/llm-sdk/test";import test, { suite, type TestContext } from "node:test";import { Agent } from "./agent.ts";import type { Toolkit } from "./toolkit.ts";
function createCloseTrackingToolkit<TContext>(): { toolkit: Toolkit<TContext>; getCloseCalls: () => number;} { let closeCalls = 0; return { toolkit: { createSession: () => Promise.resolve({ getSystemPrompt: () => undefined, getTools: () => [], close: () => { closeCalls++; return Promise.resolve(); }, }), }, getCloseCalls: () => closeCalls, };}
suite("Agent#run", () => { test("creates session, runs, and closes", async (t: TestContext) => { const closeTracker = createCloseTrackingToolkit<object>(); const model = new MockLanguageModel(); model.enqueueGenerateResult({ response: { content: [{ type: "text", text: "Mock response" }] }, }); const agent = new Agent({ name: "test-agent", model, toolkits: [closeTracker.toolkit], });
const response = await agent.run({ context: {}, input: [ { type: "message", role: "user", content: [{ type: "text", text: "Hello" }], }, ], });
t.assert.deepStrictEqual(response, { content: [{ type: "text", text: "Mock response" }], output: [ { type: "model", content: [{ type: "text", text: "Mock response" }], }, ], }); t.assert.strictEqual(closeTracker.getCloseCalls(), 1); });
test("closes the session when generation fails", async (t: TestContext) => { const closeTracker = createCloseTrackingToolkit<object>(); const model = new MockLanguageModel(); model.enqueueGenerateResult({ error: new LanguageModelError("generation failed"), }); const agent = new Agent({ name: "test-agent", model, toolkits: [closeTracker.toolkit], });
await t.assert.rejects(() => agent.run({ context: {}, input: [ { type: "message", role: "user", content: [{ type: "text", text: "Hello" }], }, ], }), ); t.assert.strictEqual(closeTracker.getCloseCalls(), 1); });});
suite("Agent#runStream", () => { test("creates session, streams, and closes", async (t: TestContext) => { const closeTracker = createCloseTrackingToolkit<object>(); const model = new MockLanguageModel(); model.enqueueStreamResult({ partials: [ { delta: { index: 0, part: { type: "text", text: "Mock" } }, }, ], }); const agent = new Agent({ name: "test-agent", model, toolkits: [closeTracker.toolkit], });
const generator = agent.runStream({ context: {}, input: [ { type: "message", role: "user", content: [{ type: "text", text: "Hello" }], }, ], });
const events = []; let current = await generator.next(); while (!current.done) { events.push(current.value); current = await generator.next(); }
t.assert.deepStrictEqual(events, [ { event: "partial", delta: { index: 0, part: { type: "text", text: "Mock" } }, }, { event: "item", index: 0, item: { type: "model", content: [{ type: "text", text: "Mock" }], }, }, { event: "response", content: [{ type: "text", text: "Mock" }], output: [ { type: "model", content: [{ type: "text", text: "Mock" }], }, ], }, ]); t.assert.deepStrictEqual(current.value, { content: [{ type: "text", text: "Mock" }], output: [ { type: "model", content: [{ type: "text", text: "Mock" }], }, ], }); t.assert.strictEqual(closeTracker.getCloseCalls(), 1); });
test("closes the session when streaming fails", async (t: TestContext) => { const closeTracker = createCloseTrackingToolkit<object>(); const model = new MockLanguageModel(); model.enqueueStreamResult({ error: new LanguageModelError("stream failed"), }); const agent = new Agent({ name: "test-agent", model, toolkits: [closeTracker.toolkit], });
await t.assert.rejects(async () => { for await (const event of agent.runStream({ context: {}, input: [ { type: "message", role: "user", content: [{ type: "text", text: "Hello" }], }, ], })) { // No events are expected before the mocked stream fails. t.assert.fail(`unexpected event: ${JSON.stringify(event)}`); } }); t.assert.strictEqual(closeTracker.getCloseCalls(), 1); });});use futures::{future::BoxFuture, TryStreamExt};use llm_agent::{ Agent, AgentError, AgentItem, AgentParams, AgentRequest, AgentResponse, AgentStreamEvent, AgentStreamItemEvent, AgentTool, BoxedError, Toolkit, ToolkitSession,};use llm_sdk::{ llm_sdk_test::{MockGenerateResult, MockLanguageModel, MockStreamResult}, ContentDelta, LanguageModelError, Message, ModelResponse, Part, PartDelta, PartialModelResponse, TextPartDelta,};use std::sync::{ atomic::{AtomicUsize, Ordering}, Arc,};
struct CloseTrackingToolkit { close_calls: Arc<AtomicUsize>,}
impl Toolkit<()> for CloseTrackingToolkit { fn create_session<'a>( &'a self, _context: &'a (), ) -> BoxFuture<'a, Result<Box<dyn ToolkitSession<()> + Send + Sync>, BoxedError>> { let close_calls = self.close_calls.clone(); Box::pin(async move { Ok(Box::new(CloseTrackingToolkitSession { close_calls }) as Box<dyn ToolkitSession<()> + Send + Sync>) }) }}
struct CloseTrackingToolkitSession { close_calls: Arc<AtomicUsize>,}
impl ToolkitSession<()> for CloseTrackingToolkitSession { fn system_prompt(&self) -> Option<String> { None }
fn tools(&self) -> Vec<AgentTool<()>> { Vec::new() }
fn close(self: Box<Self>) -> BoxFuture<'static, Result<(), BoxedError>> { Box::pin(async move { self.close_calls.fetch_add(1, Ordering::SeqCst); Ok(()) }) }}
#[tokio::test]async fn agent_run_creates_session_runs_and_finishes() { let close_calls = Arc::new(AtomicUsize::new(0)); let model = Arc::new(MockLanguageModel::new()); model.enqueue_generate(ModelResponse { content: vec![Part::text("Mock response")], ..Default::default() });
let agent = Agent::new(AgentParams::new("test-agent", model.clone()).add_toolkit( CloseTrackingToolkit { close_calls: close_calls.clone(), }, ));
let response = agent .run(AgentRequest { context: (), input: vec![AgentItem::Message(Message::user(vec![Part::text("Hello")]))], }) .await .expect("agent run succeeds");
let expected = AgentResponse { content: vec![Part::text("Mock response")], output: vec![AgentItem::Model(ModelResponse { content: vec![Part::text("Mock response")], ..Default::default() })], };
assert_eq!(response, expected); assert_eq!(close_calls.load(Ordering::SeqCst), 1);}
#[tokio::test]async fn agent_run_closes_session_when_generation_fails() { let close_calls = Arc::new(AtomicUsize::new(0)); let model = Arc::new(MockLanguageModel::new()); model.enqueue_generate(MockGenerateResult::error(LanguageModelError::InvalidInput( "generation failed".to_string(), ))); let agent = Agent::new(AgentParams::new("test-agent", model).add_toolkit( CloseTrackingToolkit { close_calls: close_calls.clone(), }, ));
let result = agent .run(AgentRequest { context: (), input: vec![AgentItem::Message(Message::user(vec![Part::text("Hello")]))], }) .await;
assert!(matches!(result, Err(AgentError::LanguageModel(_)))); assert_eq!(close_calls.load(Ordering::SeqCst), 1);}
#[tokio::test]async fn agent_run_stream_creates_session_streams_and_finishes() { let close_calls = Arc::new(AtomicUsize::new(0)); let model = Arc::new(MockLanguageModel::new()); model.enqueue_stream(MockStreamResult::partials(vec![PartialModelResponse { delta: Some(ContentDelta { index: 0, part: PartDelta::Text(TextPartDelta::new("Mock")), }), ..Default::default() }]));
let agent = Agent::new(AgentParams::new("test-agent", model.clone()).add_toolkit( CloseTrackingToolkit { close_calls: close_calls.clone(), }, ));
let stream = agent .run_stream(AgentRequest { context: (), input: vec![AgentItem::Message(Message::user(vec![Part::text("Hello")]))], }) .await .expect("agent run_stream succeeds");
let events = stream .map_err(|err| err.to_string()) .try_collect::<Vec<_>>() .await .expect("collect stream");
let expected = vec![ AgentStreamEvent::Partial(PartialModelResponse { delta: Some(ContentDelta { index: 0, part: PartDelta::Text(TextPartDelta::new("Mock")), }), ..Default::default() }), AgentStreamEvent::Item(AgentStreamItemEvent { index: 0, item: AgentItem::Model(ModelResponse { content: vec![Part::text("Mock")], ..Default::default() }), }), AgentStreamEvent::Response(AgentResponse { content: vec![Part::text("Mock")], output: vec![AgentItem::Model(ModelResponse { content: vec![Part::text("Mock")], ..Default::default() })], }), ];
assert_eq!(events, expected); assert_eq!(close_calls.load(Ordering::SeqCst), 1);}
#[tokio::test]async fn agent_run_stream_closes_session_when_streaming_fails() { let close_calls = Arc::new(AtomicUsize::new(0)); let model = Arc::new(MockLanguageModel::new()); model.enqueue_stream(MockStreamResult::error(LanguageModelError::InvalidInput( "stream failed".to_string(), ))); let agent = Agent::new(AgentParams::new("test-agent", model).add_toolkit( CloseTrackingToolkit { close_calls: close_calls.clone(), }, ));
let stream = agent .run_stream(AgentRequest { context: (), input: vec![AgentItem::Message(Message::user(vec![Part::text("Hello")]))], }) .await .expect("create agent stream"); let result = stream.try_collect::<Vec<_>>().await;
assert!(matches!(result, Err(AgentError::LanguageModel(_)))); assert_eq!(close_calls.load(Ordering::SeqCst), 1);}package llmagent_test
import ( "context" "errors" "testing"
"github.com/google/go-cmp/cmp" llmagent "github.com/hoangvvo/llm-sdk/agent-go" llmsdk "github.com/hoangvvo/llm-sdk/sdk-go" "github.com/hoangvvo/llm-sdk/sdk-go/llmsdktest")
func TestAgent_Run(t *testing.T) { t.Run("creates session, runs, and closes", func(t *testing.T) { toolkitSession := &mockToolkitSession[map[string]interface{}]{} model := llmsdktest.NewMockLanguageModel() model.EnqueueGenerateResult( llmsdktest.NewMockGenerateResultResponse(llmsdk.ModelResponse{ Content: []llmsdk.Part{ llmsdk.NewTextPart("Mock response"), }, }), ) agent := llmagent.NewAgent( "test-agent", model, llmagent.WithToolkits(&mockToolkit[map[string]interface{}]{ createFn: func(context.Context, map[string]interface{}) (llmagent.ToolkitSession[map[string]interface{}], error) { return toolkitSession, nil }, }), )
response, err := agent.Run(context.Background(), llmagent.AgentRequest[map[string]interface{}]{ Context: map[string]interface{}{}, Input: []llmagent.AgentItem{ llmagent.NewAgentItemMessage(llmsdk.NewUserMessage(llmsdk.NewTextPart("Hello"))), }, })
if err != nil { t.Fatalf("expected no error, got %v", err) }
expectedResponse := &llmagent.AgentResponse{ Content: []llmsdk.Part{ llmsdk.NewTextPart("Mock response"), }, Output: []llmagent.AgentItem{ llmagent.NewAgentItemModelResponse(llmsdk.ModelResponse{ Content: []llmsdk.Part{ llmsdk.NewTextPart("Mock response"), }, }), }, }
if diff := cmp.Diff(expectedResponse, response); diff != "" { t.Errorf("response mismatch (-want +got): %s", diff) } if toolkitSession.closeCalls != 1 { t.Fatalf("expected toolkit session to close once, got %d", toolkitSession.closeCalls) } })
t.Run("closes session when generation fails", func(t *testing.T) { toolkitSession := &mockToolkitSession[map[string]interface{}]{} model := llmsdktest.NewMockLanguageModel() modelErr := llmsdk.NewInvalidInputError("generation failed") model.EnqueueGenerateResult(llmsdktest.NewMockGenerateResultError(modelErr)) agent := llmagent.NewAgent( "test-agent", model, llmagent.WithToolkits(&mockToolkit[map[string]interface{}]{ createFn: func(context.Context, map[string]interface{}) (llmagent.ToolkitSession[map[string]interface{}], error) { return toolkitSession, nil }, }), )
_, err := agent.Run(t.Context(), llmagent.AgentRequest[map[string]interface{}]{ Context: map[string]interface{}{}, Input: []llmagent.AgentItem{ llmagent.NewAgentItemMessage(llmsdk.NewUserMessage(llmsdk.NewTextPart("Hello"))), }, }) if !errors.Is(err, modelErr) { t.Fatalf("expected wrapped model error, got %v", err) } if toolkitSession.closeCalls != 1 { t.Fatalf("expected toolkit session to close once, got %d", toolkitSession.closeCalls) } })}
func TestAgent_RunStream(t *testing.T) { t.Run("creates session, streams, and closes", func(t *testing.T) { toolkitSession := &mockToolkitSession[map[string]interface{}]{} model := llmsdktest.NewMockLanguageModel() model.EnqueueStreamResult( llmsdktest.NewMockStreamResultPartials([]llmsdk.PartialModelResponse{ { Delta: &llmsdk.ContentDelta{ Index: 0, Part: llmsdk.NewTextPartDelta("Mock"), }, }, }), ) agent := llmagent.NewAgent( "test-agent", model, llmagent.WithToolkits(&mockToolkit[map[string]interface{}]{ createFn: func(context.Context, map[string]interface{}) (llmagent.ToolkitSession[map[string]interface{}], error) { return toolkitSession, nil }, }), )
stream, err := agent.RunStream(context.Background(), llmagent.AgentRequest[map[string]interface{}]{ Context: map[string]interface{}{}, Input: []llmagent.AgentItem{ llmagent.NewAgentItemMessage(llmsdk.NewUserMessage(llmsdk.NewTextPart("Hello"))), }, })
if err != nil { t.Fatalf("expected no error, got %v", err) }
events := []*llmagent.AgentStreamEvent{} for stream.Next() { events = append(events, stream.Current()) }
if err := stream.Err(); err != nil { t.Fatalf("expected no error, got %v", err) }
expectedEvents := []*llmagent.AgentStreamEvent{ { Partial: &llmsdk.PartialModelResponse{ Delta: &llmsdk.ContentDelta{ Index: 0, Part: llmsdk.NewTextPartDelta("Mock"), }, }, }, llmagent.NewAgentStreamItemEvent( 0, llmagent.NewAgentItemModelResponse(llmsdk.ModelResponse{ Content: []llmsdk.Part{ llmsdk.NewTextPart("Mock"), }, }), ), { Response: &llmagent.AgentResponse{ Content: []llmsdk.Part{ llmsdk.NewTextPart("Mock"), }, Output: []llmagent.AgentItem{ llmagent.NewAgentItemModelResponse(llmsdk.ModelResponse{ Content: []llmsdk.Part{ llmsdk.NewTextPart("Mock"), }, }), }, }, }, }
if diff := cmp.Diff(expectedEvents, events); diff != "" { t.Errorf("stream events mismatch (-want +got):\n%s", diff) } if toolkitSession.closeCalls != 1 { t.Fatalf("expected toolkit session to close once, got %d", toolkitSession.closeCalls) } })
t.Run("closes session when streaming fails", func(t *testing.T) { toolkitSession := &mockToolkitSession[map[string]interface{}]{} model := llmsdktest.NewMockLanguageModel() modelErr := llmsdk.NewInvalidInputError("stream failed") model.EnqueueStreamResult(llmsdktest.NewMockStreamResultError(modelErr)) agent := llmagent.NewAgent( "test-agent", model, llmagent.WithToolkits(&mockToolkit[map[string]interface{}]{ createFn: func(context.Context, map[string]interface{}) (llmagent.ToolkitSession[map[string]interface{}], error) { return toolkitSession, nil }, }), )
stream, err := agent.RunStream(t.Context(), llmagent.AgentRequest[map[string]interface{}]{ Context: map[string]interface{}{}, Input: []llmagent.AgentItem{ llmagent.NewAgentItemMessage(llmsdk.NewUserMessage(llmsdk.NewTextPart("Hello"))), }, }) if err != nil { t.Fatalf("create stream: %v", err) } for stream.Next() { } if !errors.Is(stream.Err(), modelErr) { t.Fatalf("expected wrapped model error, got %v", stream.Err()) } if toolkitSession.closeCalls != 1 { t.Fatalf("expected toolkit session to close once, got %d", toolkitSession.closeCalls) } })}