Skip to main content

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​