Testing
Testing language models can be challenging, because they are non-deterministic and calling them is slow and expensive.
The Go AI SDK is designed with testability in mind. All models are interfaces, which makes it easy to create mock implementations for testing. This allows you to test your code in a repeatable and deterministic way without actually calling a language model provider.
Testing Strategies
1. Interface-Based Mocking
Since all models in the Go AI SDK are interfaces, you can create mock implementations for testing:
package myapp
import (
"context"
"testing"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)
func int64Ptr(v int64) *int64 {
return &v
}
// MockLanguageModel implements provider.LanguageModel for testing
type MockLanguageModel struct {
GenerateFunc func(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error)
StreamFunc func(ctx context.Context, opts *provider.GenerateOptions) (provider.TextStream, error)
}
func (m *MockLanguageModel) SpecificationVersion() string {
return "v3"
}
func (m *MockLanguageModel) Provider() string {
return "mock"
}
func (m *MockLanguageModel) ModelID() string {
return "mock-model"
}
func (m *MockLanguageModel) SupportsTools() bool {
return true
}
func (m *MockLanguageModel) SupportsStructuredOutput() bool {
return true
}
func (m *MockLanguageModel) SupportsImageInput() bool {
return false
}
func (m *MockLanguageModel) DoGenerate(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
if m.GenerateFunc != nil {
return m.GenerateFunc(ctx, opts)
}
return &types.GenerateResult{
Text: "Mock response",
FinishReason: types.FinishReasonStop,
Usage: types.Usage{
InputTokens: int64Ptr(10),
OutputTokens: int64Ptr(20),
TotalTokens: int64Ptr(30),
},
}, nil
}
func (m *MockLanguageModel) DoStream(ctx context.Context, opts *provider.GenerateOptions) (provider.TextStream, error) {
if m.StreamFunc != nil {
return m.StreamFunc(ctx, opts)
}
return NewMockTextStream("Mock streaming response"), nil
}
2. Mock Text Stream
Create a mock stream for testing streaming operations:
package myapp
import (
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
"github.com/digitallysavvy/go-ai/pkg/testutil"
)
func NewMockTextStream(text string) provider.TextStream {
return testutil.NewMockTextStream([]provider.StreamChunk{
{
Type: provider.ChunkTypeText,
Text: text,
},
{
Type: provider.ChunkTypeFinish,
FinishReason: types.FinishReasonStop,
Usage: &types.Usage{
InputTokens: int64Ptr(10),
OutputTokens: int64Ptr(20),
TotalTokens: int64Ptr(30),
},
},
})
}
Testing Examples
Testing GenerateText
package myapp
import (
"context"
"testing"
"github.com/digitallysavvy/go-ai/pkg/ai"
)
func TestGenerateText(t *testing.T) {
ctx := context.Background()
tests := []struct {
name string
prompt string
mockResponse string
expectedResult string
expectError bool
}{
{
name: "successful generation",
prompt: "Hello, test!",
mockResponse: "Hello, world!",
expectedResult: "Hello, world!",
expectError: false,
},
{
name: "empty response",
prompt: "Test",
mockResponse: "",
expectedResult: "",
expectError: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
// Create mock model
mockModel := &MockLanguageModel{
GenerateFunc: func(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
return &types.GenerateResult{
Text: tt.mockResponse,
FinishReason: types.FinishReasonStop,
Usage: types.Usage{
InputTokens: int64Ptr(10),
OutputTokens: int64Ptr(20),
TotalTokens: int64Ptr(30),
},
}, nil
},
}
// Call function under test
result, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: mockModel,
Prompt: tt.prompt,
})
// Assertions
if tt.expectError {
if err == nil {
t.Errorf("expected error, got nil")
}
return
}
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if result.Text != tt.expectedResult {
t.Errorf("expected %q, got %q", tt.expectedResult, result.Text)
}
})
}
}
Testing StreamText
package myapp
import (
"context"
"strings"
"testing"
"github.com/digitallysavvy/go-ai/pkg/ai"
)
func TestStreamText(t *testing.T) {
ctx := context.Background()
tests := []struct {
name string
prompt string
mockChunks []string
expectedResult string
}{
{
name: "streaming response",
prompt: "Hello, test!",
mockChunks: []string{"Hello", ", ", "world", "!"},
expectedResult: "Hello, world!",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
// Create mock model
mockModel := &MockLanguageModel{
StreamFunc: func(ctx context.Context, opts *provider.GenerateOptions) (provider.TextStream, error) {
return NewMockStreamWithChunks(tt.mockChunks), nil
},
}
// Call function under test
stream, err := ai.StreamText(ctx, ai.StreamTextOptions{
Model: mockModel,
Prompt: tt.prompt,
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
// Collect streamed text
var builder strings.Builder
for chunk := range stream.Chunks() {
builder.WriteString(chunk.Text)
}
if err := stream.Err(); err != nil {
t.Fatalf("stream error: %v", err)
}
result := builder.String()
if result != tt.expectedResult {
t.Errorf("expected %q, got %q", tt.expectedResult, result)
}
})
}
}
// NewMockStreamWithChunks creates a mock stream with multiple chunks
func NewMockStreamWithChunks(chunks []string) provider.TextStream {
streamChunks := make([]provider.StreamChunk, 0, len(chunks)+1)
for _, chunk := range chunks {
streamChunks = append(streamChunks, provider.StreamChunk{
Type: provider.ChunkTypeText,
Text: chunk,
})
}
streamChunks = append(streamChunks, provider.StreamChunk{
Type: provider.ChunkTypeFinish,
FinishReason: types.FinishReasonStop,
})
return testutil.NewMockTextStream(streamChunks)
}
Testing GenerateObject
package myapp
import (
"context"
"encoding/json"
"testing"
"github.com/digitallysavvy/go-ai/pkg/ai"
"github.com/digitallysavvy/go-ai/pkg/schema"
)
func TestGenerateObject(t *testing.T) {
ctx := context.Background()
objSchema := map[string]interface{}{
"type": "object",
"properties": map[string]interface{}{
"name": map[string]interface{}{"type": "string"},
"age": map[string]interface{}{"type": "number"},
},
"required": []string{"name", "age"},
}
tests := []struct {
name string
mockResponse string
expectedName string
expectedAge float64
expectError bool
}{
{
name: "valid object",
mockResponse: `{"name":"John","age":30}`,
expectedName: "John",
expectedAge: 30,
expectError: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
mockModel := &MockLanguageModel{
GenerateFunc: func(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
return &types.GenerateResult{
Text: tt.mockResponse,
FinishReason: types.FinishReasonStop,
}, nil
},
}
result, err := ai.GenerateObject(ctx, ai.GenerateObjectOptions{
Model: mockModel,
Schema: schema.NewSimpleJSONSchema(objSchema),
Prompt: "Generate a person",
})
if tt.expectError {
if err == nil {
t.Errorf("expected error, got nil")
}
return
}
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
// Parse the object (result.Object is already unmarshaled into a
// generic interface{}; re-marshal it to decode into a typed struct)
var person struct {
Name string `json:"name"`
Age float64 `json:"age"`
}
objBytes, err := json.Marshal(result.Object)
if err != nil {
t.Fatalf("failed to marshal object: %v", err)
}
if err := json.Unmarshal(objBytes, &person); err != nil {
t.Fatalf("failed to unmarshal object: %v", err)
}
if person.Name != tt.expectedName {
t.Errorf("expected name %q, got %q", tt.expectedName, person.Name)
}
if person.Age != tt.expectedAge {
t.Errorf("expected age %f, got %f", tt.expectedAge, person.Age)
}
})
}
}
Testing Tool Calling
package myapp
import (
"context"
"testing"
"github.com/digitallysavvy/go-ai/pkg/ai"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)
func TestToolCalling(t *testing.T) {
ctx := context.Background()
weatherTool := types.Tool{
Name: "weather",
Description: "Get the weather",
Parameters: map[string]interface{}{
"type": "object",
"properties": map[string]interface{}{
"location": map[string]interface{}{"type": "string"},
},
"required": []string{"location"},
},
Execute: func(ctx context.Context, input map[string]interface{}, opts types.ToolExecutionOptions) (interface{}, error) {
return map[string]interface{}{
"temperature": 72,
"condition": "sunny",
}, nil
},
}
mockModel := &MockLanguageModel{
GenerateFunc: func(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
// Simulate tool call
return &types.GenerateResult{
ToolCalls: []types.ToolCall{
{
ID: "call-1",
ToolName: "weather",
Arguments: map[string]interface{}{
"location": "San Francisco",
},
},
},
FinishReason: types.FinishReasonToolCalls,
}, nil
},
}
result, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: mockModel,
Prompt: "What's the weather?",
Tools: []types.Tool{weatherTool},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(result.ToolResults) == 0 {
t.Errorf("expected tool results, got none")
}
}
Testing with Embedding Models
package myapp
import (
"context"
"testing"
"github.com/digitallysavvy/go-ai/pkg/ai"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)
// MockEmbeddingModel implements provider.EmbeddingModel for testing
type MockEmbeddingModel struct {
EmbedFunc func(ctx context.Context, input string, opts *provider.EmbedModelOptions) (*types.EmbeddingResult, error)
EmbedManyFunc func(ctx context.Context, inputs []string, opts *provider.EmbedModelOptions) (*types.EmbeddingsResult, error)
}
func (m *MockEmbeddingModel) SpecificationVersion() string {
return "v4"
}
func (m *MockEmbeddingModel) Provider() string {
return "mock"
}
func (m *MockEmbeddingModel) ModelID() string {
return "mock-embedding"
}
func (m *MockEmbeddingModel) MaxEmbeddingsPerCall() int {
return 100
}
func (m *MockEmbeddingModel) SupportsParallelCalls() bool {
return true
}
func (m *MockEmbeddingModel) DoEmbed(ctx context.Context, input string, opts *provider.EmbedModelOptions) (*types.EmbeddingResult, error) {
if m.EmbedFunc != nil {
return m.EmbedFunc(ctx, input, opts)
}
return &types.EmbeddingResult{
Embedding: []float64{0.1, 0.2, 0.3},
Usage: types.EmbeddingUsage{
TotalTokens: 10,
},
}, nil
}
func (m *MockEmbeddingModel) DoEmbedMany(ctx context.Context, inputs []string, opts *provider.EmbedModelOptions) (*types.EmbeddingsResult, error) {
if m.EmbedManyFunc != nil {
return m.EmbedManyFunc(ctx, inputs, opts)
}
embeddings := make([][]float64, len(inputs))
for i := range inputs {
embeddings[i] = []float64{0.1, 0.2, 0.3}
}
return &types.EmbeddingsResult{
Embeddings: embeddings,
Usage: types.EmbeddingUsage{
TotalTokens: 10 * len(inputs),
},
}, nil
}
func TestEmbed(t *testing.T) {
ctx := context.Background()
mockModel := &MockEmbeddingModel{}
result, err := ai.Embed(ctx, ai.EmbedOptions{
Model: mockModel,
Input: "test text",
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(result.Embedding) != 3 {
t.Errorf("expected 3-dimensional embedding, got %d", len(result.Embedding))
}
}
Testing Best Practices
1. Use Table-Driven Tests
Go's table-driven tests are ideal for testing multiple scenarios:
func TestMultipleScenarios(t *testing.T) {
tests := []struct {
name string
input string
expected string
}{
{"scenario 1", "input1", "output1"},
{"scenario 2", "input2", "output2"},
{"scenario 3", "input3", "output3"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
// Test logic here
})
}
}
2. Test Error Conditions
Always test error paths:
func TestErrorHandling(t *testing.T) {
ctx := context.Background()
mockModel := &MockLanguageModel{
GenerateFunc: func(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
return nil, errors.New("mock error")
},
}
_, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: mockModel,
Prompt: "test",
})
if err == nil {
t.Error("expected error, got nil")
}
}
3. Test Context Cancellation
Test that your code respects context cancellation:
func TestContextCancellation(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel() // Cancel immediately
mockModel := &MockLanguageModel{
GenerateFunc: func(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
// A real model call would fail on a canceled context; the mock
// must check it explicitly since it never makes an HTTP request.
if err := ctx.Err(); err != nil {
return nil, err
}
return &types.GenerateResult{Text: "unreachable"}, nil
},
}
_, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: mockModel,
Prompt: "test",
})
if err == nil {
t.Error("expected context cancellation error")
}
}
4. Use Test Helpers
Create reusable test helpers:
func newTestModel(response string) *MockLanguageModel {
return &MockLanguageModel{
GenerateFunc: func(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
return &types.GenerateResult{
Text: response,
FinishReason: types.FinishReasonStop,
}, nil
},
}
}
func TestWithHelper(t *testing.T) {
model := newTestModel("test response")
// Use model in test...
}
5. Test Callbacks
Test that callbacks are called correctly:
func TestOnFinishCallback(t *testing.T) {
ctx := context.Background()
callbackCalled := false
mockModel := &MockLanguageModel{}
_, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: mockModel,
Prompt: "test",
OnFinish: func(ctx context.Context, result *ai.GenerateTextResult, userContext interface{}) {
callbackCalled = true
},
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !callbackCalled {
t.Error("OnFinish callback was not called")
}
}
6. Use Subtests for Complex Scenarios
Break down complex tests into subtests:
func TestComplexScenario(t *testing.T) {
t.Run("setup", func(t *testing.T) {
// Setup test
})
t.Run("execution", func(t *testing.T) {
// Execute test
})
t.Run("validation", func(t *testing.T) {
// Validate results
})
}
7. Mock External Dependencies
When testing code that uses AI SDK, mock the model interface:
type MyService struct {
model provider.LanguageModel
}
func (s *MyService) ProcessText(ctx context.Context, text string) (string, error) {
result, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: s.model,
Prompt: text,
})
if err != nil {
return "", err
}
return result.Text, nil
}
func TestMyService(t *testing.T) {
mockModel := &MockLanguageModel{
GenerateFunc: func(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
return &types.GenerateResult{
Text: "processed: " + opts.Prompt.Messages[0].Content[0].(types.TextContent).Text,
}, nil
},
}
service := &MyService{model: mockModel}
result, err := service.ProcessText(context.Background(), "test")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if result != "processed: test" {
t.Errorf("unexpected result: %s", result)
}
}
Integration Testing
For integration tests with real providers, use build tags to separate them from unit tests:
//go:build integration
// +build integration
package myapp
import (
"context"
"os"
"testing"
"github.com/digitallysavvy/go-ai/pkg/ai"
"github.com/digitallysavvy/go-ai/pkg/providers/openai"
)
func TestRealProvider(t *testing.T) {
if os.Getenv("OPENAI_API_KEY") == "" {
t.Skip("OPENAI_API_KEY not set")
}
ctx := context.Background()
provider := openai.New(openai.Config{APIKey: os.Getenv("OPENAI_API_KEY")})
model, _ := provider.LanguageModel("gpt-4")
result, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: model,
Prompt: "Say hello",
})
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if result.Text == "" {
t.Error("expected non-empty response")
}
}
Run integration tests with:
go test -tags=integration ./...