Middleware Interface
Structs for implementing custom middleware to intercept and modify model behavior. Each field is an optional hook; a middleware only needs to set the hooks it uses, and leave the rest nil.
LanguageModelMiddleware
type LanguageModelMiddleware struct {
// SpecificationVersion should be "v3" for the current version
SpecificationVersion string
// OverrideProvider allows overriding the provider name
OverrideProvider func(model provider.LanguageModel) string
// OverrideModelID allows overriding the model ID
OverrideModelID func(model provider.LanguageModel) string
// OverrideSupportedURLs allows overriding the URL patterns (regular
// expressions keyed by media type) the wrapped model reports supporting.
OverrideSupportedURLs func(model provider.LanguageModel) map[string][]string
// TransformParams transforms the parameters before they are passed to the language model
TransformParams func(ctx context.Context, callType string, params *provider.GenerateOptions, model provider.LanguageModel) (*provider.GenerateOptions, error)
// WrapGenerate wraps the generate operation of the language model
WrapGenerate func(ctx context.Context, doGenerate func() (*types.GenerateResult, error), doStream func() (provider.TextStream, error), params *provider.GenerateOptions, model provider.LanguageModel) (*types.GenerateResult, error)
// WrapStream wraps the stream operation of the language model
WrapStream func(ctx context.Context, doGenerate func() (*types.GenerateResult, error), doStream func() (provider.TextStream, error), params *provider.GenerateOptions, model provider.LanguageModel) (provider.TextStream, error)
}
Middleware for language models that can intercept both generation and
streaming calls. WrapGenerate and WrapStream each receive both
doGenerate and doStream closures — a WrapGenerate hook calls
doGenerate() to invoke the wrapped model (or the next middleware in the
chain), and a WrapStream hook calls doStream().
EmbeddingModelMiddleware
type EmbeddingModelMiddleware struct {
SpecificationVersion string
OverrideProvider func(model provider.EmbeddingModel) string
OverrideModelID func(model provider.EmbeddingModel) string
OverrideMaxEmbeddingsPerCall func(model provider.EmbeddingModel) int
OverrideSupportsParallelCalls func(model provider.EmbeddingModel) bool
// TransformInput transforms the input before it is passed to the embedding model
TransformInput func(ctx context.Context, input string, model provider.EmbeddingModel) (string, error)
// WrapEmbed wraps a single embed call
WrapEmbed func(ctx context.Context, doEmbed func() (*types.EmbeddingResult, error), input string, model provider.EmbeddingModel) (*types.EmbeddingResult, error)
// WrapEmbedMany wraps a batch embed call
WrapEmbedMany func(ctx context.Context, doEmbedMany func() (*types.EmbeddingsResult, error), inputs []string, model provider.EmbeddingModel) (*types.EmbeddingsResult, error)
}
Middleware for embedding models.
ImageModelMiddleware
type ImageModelMiddleware struct {
// SpecificationVersion should be "v3" or "v4" for the current version
SpecificationVersion string
OverrideProvider func(model provider.ImageModel) string
OverrideModelID func(model provider.ImageModel) string
OverrideMaxImagesPerCall func(model provider.ImageModel) int
// TransformParams transforms the parameters before they are passed to the image model
TransformParams func(ctx context.Context, params *provider.ImageGenerateOptions, model provider.ImageModel) (*provider.ImageGenerateOptions, error)
// WrapGenerate wraps the generate operation of the image model
WrapGenerate func(ctx context.Context, doGenerate func() (*types.ImageResult, error), params *provider.ImageGenerateOptions, model provider.ImageModel) (*types.ImageResult, error)
}
Middleware for image models.
Examples
Logging Middleware
package main
import (
"context"
"log"
"time"
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)
func LoggingMiddleware() *middleware.LanguageModelMiddleware {
return &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()
log.Printf("Starting generation...")
result, err := doGenerate()
duration := time.Since(start)
if err != nil {
log.Printf("Generation failed after %v: %v", duration, err)
} else {
log.Printf("Generation completed in %v (%d tokens)",
duration, result.Usage.GetTotalTokens())
}
return result, err
},
}
}
Retry Middleware
func RetryMiddleware(maxRetries int) *middleware.LanguageModelMiddleware {
return &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) {
var result *types.GenerateResult
var err error
for attempt := 0; attempt < maxRetries; attempt++ {
result, err = doGenerate()
if err == nil {
return result, nil
}
log.Printf("Retry attempt %d/%d", attempt+1, maxRetries)
time.Sleep(time.Duration(attempt+1) * time.Second)
}
return nil, fmt.Errorf("max retries exceeded: %w", err)
},
}
}
Token Counter Middleware
type TokenCounter struct {
totalTokens int
mu sync.Mutex
}
func (tc *TokenCounter) Middleware() *middleware.LanguageModelMiddleware {
return &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) {
result, err := doGenerate()
if err != nil {
return nil, err
}
tc.mu.Lock()
tc.totalTokens += int(result.Usage.GetTotalTokens())
tc.mu.Unlock()
return result, nil
},
}
}
func (tc *TokenCounter) Total() int {
tc.mu.Lock()
defer tc.mu.Unlock()
return tc.totalTokens
}
See Also
- WrapLanguageModel - Apply middleware to models
- Built-in Middleware - Available middleware catalog