# 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

```go
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

```go
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

```go
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

```go
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

```go
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

```go
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](https://goaisdk.com/docs/reference/middleware/wrap-language-model.md) - Apply middleware to models
- [Built-in Middleware](https://goaisdk.com/docs/reference/middleware/built-in-middleware.md) - Available middleware catalog
