Test code that calls a model
Write fast, deterministic tests for code that calls a language model.
package main
import (
"context"
"fmt"
"log"
"os"
"strings"
"github.com/digitallysavvy/go-ai/pkg/ai"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/providers/openai"
)
// Summarize is the code under test. It takes the model as a parameter, so a
// test can pass a mock.
func Summarize(ctx context.Context, model provider.LanguageModel, text string) (string, error) {
result, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: model,
System: "Summarize the text in one sentence.",
Prompt: text,
})
if err != nil {
return "", err
}
return strings.TrimSpace(result.Text), nil
}
func main() {
model, err := openai.New(openai.Config{APIKey: os.Getenv("OPENAI_API_KEY")}).
LanguageModel(openai.ModelGPT6Astra)
if err != nil {
log.Fatal(err)
}
summary, err := Summarize(context.Background(), model, "Go is a statically typed, compiled language designed at Google.")
if err != nil {
log.Fatal(err)
}
fmt.Println(summary)
}
The tests, in main_test.go next to it:
// main_test.go
package main
import (
"context"
"errors"
"strings"
"testing"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
"github.com/digitallysavvy/go-ai/pkg/testutil"
)
func TestSummarize(t *testing.T) {
model := &testutil.MockLanguageModel{
DoGenerateFunc: func(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
return &types.GenerateResult{
Text: " Go is a compiled language. ",
FinishReason: types.FinishReasonStop,
}, nil
},
}
got, err := Summarize(context.Background(), model, "long text")
if err != nil {
t.Fatal(err)
}
if got != "Go is a compiled language." {
t.Errorf("got %q", got)
}
// The mock records every call, so you can check what the model was sent.
if len(model.GenerateCalls) != 1 {
t.Fatalf("calls = %d, want 1", len(model.GenerateCalls))
}
if !strings.Contains(promptText(model.GenerateCalls[0]), "long text") {
t.Error("prompt did not contain the input text")
}
}
func TestSummarizeError(t *testing.T) {
boom := errors.New("provider unavailable")
model := &testutil.MockLanguageModel{
DoGenerateFunc: func(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
return nil, boom
},
}
if _, err := Summarize(context.Background(), model, "x"); !errors.Is(err, boom) {
t.Errorf("err = %v, want %v", err, boom)
}
}
func promptText(opts *provider.GenerateOptions) string {
var b strings.Builder
for _, m := range opts.Prompt.Messages {
for _, p := range m.Content {
if tc, ok := p.(types.TextContent); ok {
b.WriteString(tc.Text)
}
}
}
return b.String()
}
Run it:
go test ./examples/recipes/mock-model-tests
Notes
- Take
provider.LanguageModelas a parameter. Production code passes a real model and tests pass a mock. testutil.MockLanguageModellets you setDoGenerateFuncandDoStreamFunc, and records every call inGenerateCallsandStreamCalls.- Return an error from the mock to test your failure paths.
- Keep a few integration tests against real providers, and skip them when the API key is not set.
Go deeper
The full program is at examples/recipes/mock-model-tests/main.go.