Skip to main content

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.LanguageModel as a parameter. Production code passes a real model and tests pass a mock.
  • testutil.MockLanguageModel lets you set DoGenerateFunc and DoStreamFunc, and records every call in GenerateCalls and StreamCalls.
  • 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.