Language Model Middleware
Language model middleware is a way to enhance the behavior of language models by intercepting and modifying the calls to the language model.
It can be used to add features like guardrails, RAG, caching, and logging in a language model agnostic way. Such middleware can be developed and distributed independently from the language models that they are applied to.
Using Language Model Middleware
You can use language model middleware with the middleware.WrapLanguageModel function. It takes a language model and one or more middleware instances and returns a new language model that incorporates the middleware.
For parity-oriented API discovery, equivalent wrappers are also exposed from pkg/ai (for example ai.WrapLanguageModel and ai.DefaultSettingsMiddleware).
import (
"os"
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/providers/openai"
)
provider := openai.New(openai.Config{APIKey: os.Getenv("OPENAI_API_KEY")})
baseModel, _ := provider.LanguageModel("gpt-4")
wrappedModel := middleware.WrapLanguageModel(
baseModel,
[]*middleware.LanguageModelMiddleware{yourMiddleware},
nil, // optional modelID override
nil, // optional providerID override
)
The wrapped language model can be used just like any other language model, e.g. in ai.GenerateText or ai.StreamText:
result, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: wrappedModel,
Prompt: "What cities are in the United States?",
})
Multiple Middlewares
You can provide multiple middlewares to the middleware.WrapLanguageModel function. The middlewares will be applied in the order they are provided.
wrappedModel := middleware.WrapLanguageModel(
baseModel,
[]*middleware.LanguageModelMiddleware{
firstMiddleware,
secondMiddleware,
},
nil,
nil,
)
// Applied as: firstMiddleware(secondMiddleware(baseModel))
Built-in Middleware
The Go AI SDK comes with several built-in middlewares:
Default Settings Middleware
The DefaultSettingsMiddleware applies default settings to a language model. Settings provided in individual calls will override these defaults.
import (
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/provider"
)
temperature := 0.5
maxTokens := 800
model := middleware.WrapLanguageModel(
baseModel,
[]*middleware.LanguageModelMiddleware{
middleware.DefaultSettingsMiddleware(&provider.GenerateOptions{
Temperature: &temperature,
MaxTokens: &maxTokens,
}),
},
nil,
nil,
)
// All calls with this model will use temperature=0.5 and maxTokens=800 by default
result, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: model,
Prompt: "Explain photosynthesis",
})
Implementing Custom Language Model Middleware
Note: Implementing language model middleware is advanced functionality and requires a solid understanding of the language model specification in the provider package.
You can implement any of the following functions to modify the behavior of the language model:
- TransformParams: Transforms the parameters before they are passed to the language model, for both
DoGenerateandDoStream. - WrapGenerate: Wraps the
DoGeneratemethod of the language model. You can modify the parameters, call the language model, and modify the result. - WrapStream: Wraps the
DoStreammethod of the language model. You can modify the parameters, call the language model, and modify the result.
Middleware Structure
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
// TransformParams transforms parameters before passing to the model
TransformParams func(
ctx context.Context,
callType string,
params *provider.GenerateOptions,
model provider.LanguageModel,
) (*provider.GenerateOptions, error)
// WrapGenerate wraps the generate operation
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
WrapStream func(
ctx context.Context,
doGenerate func() (*types.GenerateResult, error),
doStream func() (provider.TextStream, error),
params *provider.GenerateOptions,
model provider.LanguageModel,
) (provider.TextStream, error)
}
Examples
Note: These examples are not meant to be used in production as-is. They demonstrate how to use middleware to enhance language model behavior.
Logging Middleware
This example shows how to log the parameters and generated text of a language model call.
package main
import (
"context"
"encoding/json"
"io"
"log"
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)
func LogMiddleware() *middleware.LanguageModelMiddleware {
return &middleware.LanguageModelMiddleware{
SpecificationVersion: "v3",
WrapGenerate: func(
ctx context.Context,
doGenerate func() (*types.GenerateResult, error),
doStream func() (provider.TextStream, error),
params *provider.GenerateOptions,
model provider.LanguageModel,
) (*types.GenerateResult, error) {
log.Println("DoGenerate called")
paramsJSON, _ := json.MarshalIndent(params, "", " ")
log.Printf("Params: %s\n", paramsJSON)
result, err := doGenerate()
if err != nil {
log.Printf("DoGenerate failed: %v\n", err)
return nil, err
}
log.Println("DoGenerate finished")
log.Printf("Generated text: %s\n", result.Text)
return result, nil
},
WrapStream: func(
ctx context.Context,
doGenerate func() (*types.GenerateResult, error),
doStream func() (provider.TextStream, error),
params *provider.GenerateOptions,
model provider.LanguageModel,
) (provider.TextStream, error) {
log.Println("DoStream called")
paramsJSON, _ := json.MarshalIndent(params, "", " ")
log.Printf("Params: %s\n", paramsJSON)
stream, err := doStream()
if err != nil {
log.Printf("DoStream failed: %v\n", err)
return nil, err
}
return &loggingStream{base: stream}, nil
},
}
}
// loggingStream implements provider.TextStream
type loggingStream struct {
base provider.TextStream
fullText string
}
func (s *loggingStream) Next() (*provider.StreamChunk, error) {
chunk, err := s.base.Next()
if err == io.EOF {
log.Println("DoStream finished")
log.Printf("Generated text: %s\n", s.fullText)
return nil, io.EOF
}
if err != nil {
return nil, err
}
if chunk.Type == provider.ChunkTypeText {
s.fullText += chunk.Text
}
return chunk, nil
}
func (s *loggingStream) Err() error {
return s.base.Err()
}
func (s *loggingStream) Close() error {
return s.base.Close()
}
Caching Middleware
This example shows how to build a simple cache for the generated text of a language model call.
package main
import (
"context"
"crypto/sha256"
"encoding/json"
"fmt"
"sync"
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)
func CacheMiddleware() *middleware.LanguageModelMiddleware {
cache := &sync.Map{}
return &middleware.LanguageModelMiddleware{
SpecificationVersion: "v3",
WrapGenerate: func(
ctx context.Context,
doGenerate func() (*types.GenerateResult, error),
doStream func() (provider.TextStream, error),
params *provider.GenerateOptions,
model provider.LanguageModel,
) (*types.GenerateResult, error) {
// Generate cache key from parameters
cacheKey := generateCacheKey(params)
// Check cache
if cached, ok := cache.Load(cacheKey); ok {
fmt.Println("Cache hit!")
return cached.(*types.GenerateResult), nil
}
// Call model
result, err := doGenerate()
if err != nil {
return nil, err
}
// Store in cache
cache.Store(cacheKey, result)
return result, nil
},
// Note: Implementing caching for streaming is more complex
// as you need to collect the full stream before caching
}
}
func generateCacheKey(params *provider.GenerateOptions) string {
data, _ := json.Marshal(params)
hash := sha256.Sum256(data)
return fmt.Sprintf("%x", hash)
}
Retrieval Augmented Generation (RAG) Middleware
This example shows how to use RAG as middleware.
package main
import (
"context"
"fmt"
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)
func RAGMiddleware(vectorDB VectorDatabase) *middleware.LanguageModelMiddleware {
return &middleware.LanguageModelMiddleware{
SpecificationVersion: "v3",
TransformParams: func(
ctx context.Context,
callType string,
params *provider.GenerateOptions,
model provider.LanguageModel,
) (*provider.GenerateOptions, error) {
// Extract the last user message
lastUserMessage := getLastUserMessage(params)
if lastUserMessage == "" {
return params, nil // No modification needed
}
// Find relevant sources from vector database
sources, err := vectorDB.FindSimilar(ctx, lastUserMessage, 5)
if err != nil {
return nil, fmt.Errorf("failed to find sources: %w", err)
}
// Build instruction with sources
instruction := "Use the following information to answer the question:\n"
for _, source := range sources {
instruction += fmt.Sprintf("\n%s", source.Content)
}
// Add instruction to the last user message
modifiedParams := addToLastUserMessage(params, instruction)
return modifiedParams, nil
},
}
}
// Helper functions (implementation depends on your use case)
func getLastUserMessage(params *provider.GenerateOptions) string {
if len(params.Prompt.Messages) == 0 {
return ""
}
lastMsg := params.Prompt.Messages[len(params.Prompt.Messages)-1]
if lastMsg.Role != types.RoleUser {
return ""
}
for _, part := range lastMsg.Content {
if textPart, ok := part.(types.TextContent); ok {
return textPart.Text
}
}
return ""
}
func addToLastUserMessage(params *provider.GenerateOptions, text string) *provider.GenerateOptions {
modified := *params
if len(modified.Prompt.Messages) == 0 {
return params
}
messages := make([]types.Message, len(modified.Prompt.Messages))
copy(messages, modified.Prompt.Messages)
lastIdx := len(messages) - 1
lastMsg := messages[lastIdx]
// Prepend text to the last user message
newContent := []types.ContentPart{types.TextContent{Text: text + "\n\n"}}
newContent = append(newContent, lastMsg.Content...)
messages[lastIdx] = types.Message{
Role: lastMsg.Role,
Content: newContent,
}
modified.Prompt.Messages = messages
return &modified
}
// VectorDatabase interface (example)
type VectorDatabase interface {
FindSimilar(ctx context.Context, query string, limit int) ([]Document, error)
}
type Document struct {
Content string
Score float64
}
Guardrails Middleware
Guardrails ensure that the generated text of a language model call is safe and appropriate.
package main
import (
"context"
"regexp"
"strings"
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)
func GuardrailMiddleware() *middleware.LanguageModelMiddleware {
// Define patterns to filter
sensitivePatterns := []*regexp.Regexp{
regexp.MustCompile(`\d{3}-\d{2}-\d{4}`), // SSN pattern
regexp.MustCompile(`\d{16}`), // Credit card pattern
regexp.MustCompile(`\b[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Z|a-z]{2,}\b`), // Email
}
return &middleware.LanguageModelMiddleware{
SpecificationVersion: "v3",
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
}
// Filter sensitive information
cleanedText := result.Text
for _, pattern := range sensitivePatterns {
cleanedText = pattern.ReplaceAllString(cleanedText, "<REDACTED>")
}
// Replace inappropriate words
cleanedText = strings.ReplaceAll(cleanedText, "badword", "<FILTERED>")
result.Text = cleanedText
return result, nil
},
// Note: Streaming guardrails are more complex because you don't know
// the full content until the stream is finished
}
}
Retry Middleware
This middleware automatically retries failed requests with exponential backoff.
package main
import (
"context"
"fmt"
"time"
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)
func RetryMiddleware(maxRetries int) *middleware.LanguageModelMiddleware {
return &middleware.LanguageModelMiddleware{
SpecificationVersion: "v3",
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 lastErr error
for attempt := 0; attempt <= maxRetries; attempt++ {
if attempt > 0 {
// Exponential backoff
backoff := time.Duration(1<<uint(attempt-1)) * time.Second
fmt.Printf("Retry attempt %d/%d after %v\n", attempt, maxRetries, backoff)
select {
case <-time.After(backoff):
case <-ctx.Done():
return nil, ctx.Err()
}
}
result, err := doGenerate()
if err == nil {
return result, nil
}
lastErr = err
// Don't retry on context cancellation or timeout
if ctx.Err() != nil {
break
}
}
return nil, fmt.Errorf("failed after %d attempts: %w", maxRetries+1, lastErr)
},
}
}
Metrics Middleware
Track metrics like latency, token usage, and error rates.
package main
import (
"context"
"time"
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)
type Metrics struct {
TotalRequests int64
TotalErrors int64
TotalTokens int64
TotalLatencyMs int64
}
func MetricsMiddleware(metrics *Metrics) *middleware.LanguageModelMiddleware {
return &middleware.LanguageModelMiddleware{
SpecificationVersion: "v3",
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()
metrics.TotalRequests++
result, err := doGenerate()
latency := time.Since(start).Milliseconds()
metrics.TotalLatencyMs += latency
if err != nil {
metrics.TotalErrors++
return nil, err
}
metrics.TotalTokens += int64(result.Usage.GetTotalTokens())
return result, nil
},
}
}
Wrapping Providers with Middleware
Instead of wrapping individual models, you can wrap an entire provider to apply middleware to all models from that provider:
import (
"os"
"github.com/digitallysavvy/go-ai/pkg/middleware"
"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/providers/openai"
)
baseProvider := openai.New(openai.Config{APIKey: os.Getenv("OPENAI_API_KEY")})
temperature := 0.7
wrappedProvider := middleware.WrapProvider(
baseProvider,
// Language model middleware
[]*middleware.LanguageModelMiddleware{
middleware.DefaultSettingsMiddleware(&provider.GenerateOptions{
Temperature: &temperature,
}),
LogMiddleware(),
},
// Embedding model middleware
[]*middleware.EmbeddingModelMiddleware{
// Add embedding-specific middleware here
},
)
// All models from this provider will have the middleware applied
model, _ := wrappedProvider.LanguageModel("gpt-4")
Combining Multiple Middlewares
You can combine multiple middlewares for powerful behavior:
var metrics Metrics
model := middleware.WrapLanguageModel(
baseModel,
[]*middleware.LanguageModelMiddleware{
LogMiddleware(), // Log all requests
MetricsMiddleware(&metrics), // Track metrics
RetryMiddleware(3), // Retry failures
CacheMiddleware(), // Cache results
GuardrailMiddleware(), // Filter sensitive data
middleware.DefaultSettingsMiddleware(&provider.GenerateOptions{
Temperature: &temperature,
}),
},
nil,
nil,
)
Middlewares are applied in order:
- First middleware transforms params first
- Last middleware wraps directly around the model
Configuring Per-Request Custom Metadata
To send and access custom metadata in middleware, you can use ProviderOptions:
func LogMiddlewareWithMetadata() *middleware.LanguageModelMiddleware {
return &middleware.LanguageModelMiddleware{
SpecificationVersion: "v3",
WrapGenerate: func(
ctx context.Context,
doGenerate func() (*types.GenerateResult, error),
doStream func() (provider.TextStream, error),
params *provider.GenerateOptions,
model provider.LanguageModel,
) (*types.GenerateResult, error) {
// Access custom metadata
if metadata, ok := params.ProviderOptions["yourLogMiddleware"].(map[string]interface{}); ok {
fmt.Printf("Metadata: %+v\n", metadata)
}
return doGenerate()
},
}
}
// Usage
result, err := ai.GenerateText(ctx, ai.GenerateTextOptions{
Model: wrappedModel,
Prompt: "Invent a new holiday and describe its traditions.",
ProviderOptions: map[string]interface{}{
"yourLogMiddleware": map[string]interface{}{
"userId": "user-123",
"requestId": "req-456",
"timestamp": time.Now().Unix(),
},
},
})
Best Practices
1. Keep Middleware Focused
Each middleware should have a single responsibility:
// Good: Single responsibility
func LogMiddleware() *middleware.LanguageModelMiddleware { ... }
func CacheMiddleware() *middleware.LanguageModelMiddleware { ... }
func MetricsMiddleware() *middleware.LanguageModelMiddleware { ... }
// Bad: Multiple responsibilities
func LogAndCacheAndMetricsMiddleware() *middleware.LanguageModelMiddleware { ... }
2. Handle Errors Gracefully
Always handle errors and decide whether to fail fast or continue:
TransformParams: func(ctx context.Context, callType string, params *provider.GenerateOptions, model provider.LanguageModel) (*provider.GenerateOptions, error) {
enrichedParams, err := enrichWithContext(params)
if err != nil {
// Log error but don't fail the request
log.Printf("Failed to enrich params: %v", err)
return params, nil // Return original params
}
return enrichedParams, nil
}
3. Be Mindful of Performance
Middleware runs on every request, so optimize for performance:
// Good: Efficient caching
var cache = &sync.Map{}
// Bad: Inefficient locking
var (
cacheMutex sync.Mutex
cache = make(map[string]interface{})
)
4. Document Your Middleware
Provide clear documentation for custom middleware:
// RateLimitMiddleware limits the number of requests per time window.
// It uses a token bucket algorithm with the specified rate and burst size.
//
// Parameters:
// - rate: Maximum requests per second
// - burst: Maximum burst size
//
// Example:
// middleware := RateLimitMiddleware(10, 20) // 10 req/s, burst of 20
func RateLimitMiddleware(rate float64, burst int) *middleware.LanguageModelMiddleware {
// Implementation...
}
5. Consider Streaming
If you implement WrapStream, ensure you handle the streaming channel correctly:
WrapStream: func(ctx context.Context, doGenerate func() (*types.GenerateResult, error), doStream func() (provider.TextStream, error), params *provider.GenerateOptions, model provider.LanguageModel) (provider.TextStream, error) {
stream, err := doStream()
if err != nil {
return nil, err
}
return &processingStream{base: stream}, nil
}