Skip to main content

WrapLanguageModel

Wraps a language model with middleware to add additional behavior like default settings, logging, or extraction.

Signature​

func WrapLanguageModel(
model provider.LanguageModel,
middleware []*middleware.LanguageModelMiddleware,
modelID *string,
providerID *string,
) provider.LanguageModel

When multiple middleware are provided, the first middleware transforms the input first, and the last middleware is wrapped directly around the model.

Parameters​

ParameterTypeDescription
modelprovider.LanguageModelBase language model to wrap
middleware[]*middleware.LanguageModelMiddlewareMiddleware to apply, in order
modelID*stringOptional override for the reported model ID (nil keeps the wrapped model's ID)
providerID*stringOptional override for the reported provider name (nil keeps the wrapped model's provider)

Examples​

Basic Wrapping​

package main

import (
"context"
"fmt"
"log"

"github.com/digitallysavvy/go-ai/pkg/ai"
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/providers/openai"
)

func main() {
p := openai.New(openai.Config{
APIKey: "your-api-key",
})
baseModel, _ := p.LanguageModel("gpt-4")

temperature := 0.7
wrapped := middleware.WrapLanguageModel(
baseModel,
[]*middleware.LanguageModelMiddleware{
middleware.DefaultSettingsMiddleware(&provider.GenerateOptions{
Temperature: &temperature,
}),
},
nil,
nil,
)

result, err := ai.GenerateText(context.Background(), ai.GenerateTextOptions{
Model: wrapped,
Prompt: "Hello!",
// Temperature will default to 0.7
})
if err != nil {
log.Fatal(err)
}
fmt.Println(result.Text)
}

Multiple Middleware​

temperature := 0.7
maxTokens := 500

wrapped := middleware.WrapLanguageModel(
baseModel,
[]*middleware.LanguageModelMiddleware{
middleware.DefaultSettingsMiddleware(&provider.GenerateOptions{
Temperature: &temperature,
MaxTokens: &maxTokens,
}),
middleware.ExtractReasoningMiddleware(&middleware.ExtractReasoningOptions{
TagName: "think",
}),
},
nil,
nil,
)

Overriding the Reported Model ID​

modelID := "custom-model-alias"

wrapped := middleware.WrapLanguageModel(
baseModel,
nil,
&modelID,
nil,
)

Custom Middleware​

customMiddleware := &middleware.LanguageModelMiddleware{
WrapGenerate: func(
ctx context.Context,
doGenerate func() (*types.GenerateResult, error),
doStream func() (provider.TextStream, error),
params *provider.GenerateOptions,
model provider.LanguageModel,
) (*types.GenerateResult, error) {
start := time.Now()
result, err := doGenerate()
log.Printf("Generation took %v", time.Since(start))
return result, err
},
}

wrapped := middleware.WrapLanguageModel(
baseModel,
[]*middleware.LanguageModelMiddleware{customMiddleware},
nil,
nil,
)

See Also​