Skip to main content

Custom Provider Guide

Learn how to implement custom AI providers for the Go-AI SDK.

Provider Interface​

type Provider interface {
Name() string

LanguageModel(modelID string) (LanguageModel, error)
EmbeddingModel(modelID string) (EmbeddingModel, error)
ImageModel(modelID string) (ImageModel, error)
SpeechModel(modelID string) (SpeechModel, error)
TranscriptionModel(modelID string) (TranscriptionModel, error)
RerankingModel(modelID string) (RerankingModel, error)
}

Implementing a Language Model​

package main

import (
"context"
"fmt"

"github.com/digitallysavvy/go-ai/pkg/provider"
"github.com/digitallysavvy/go-ai/pkg/provider/types"
)

type CustomLanguageModel struct {
modelID string
apiKey string
}

func (m *CustomLanguageModel) SpecificationVersion() string {
return "v3"
}

func (m *CustomLanguageModel) Provider() string {
return "custom"
}

func (m *CustomLanguageModel) ModelID() string {
return m.modelID
}

func (m *CustomLanguageModel) SupportsTools() bool {
return false
}

func (m *CustomLanguageModel) SupportsStructuredOutput() bool {
return false
}

func (m *CustomLanguageModel) SupportsImageInput() bool {
return false
}

func int64Ptr(v int64) *int64 {
return &v
}

func (m *CustomLanguageModel) DoGenerate(ctx context.Context, opts *provider.GenerateOptions) (*types.GenerateResult, error) {
// Implement your API call here
return &types.GenerateResult{
Text: "Generated text",
FinishReason: types.FinishReasonStop,
Usage: types.Usage{
InputTokens: int64Ptr(10),
OutputTokens: int64Ptr(20),
TotalTokens: int64Ptr(30),
},
}, nil
}

func (m *CustomLanguageModel) DoStream(ctx context.Context, opts *provider.GenerateOptions) (provider.TextStream, error) {
return nil, fmt.Errorf("streaming not implemented")
}

Implementing a Provider​

type CustomProvider struct {
apiKey string
}

func NewCustomProvider(apiKey string) *CustomProvider {
return &CustomProvider{apiKey: apiKey}
}

func (p *CustomProvider) Name() string {
return "custom"
}

func (p *CustomProvider) LanguageModel(modelID string) (provider.LanguageModel, error) {
return &CustomLanguageModel{
modelID: modelID,
apiKey: p.apiKey,
}, nil
}

func (p *CustomProvider) EmbeddingModel(modelID string) (provider.EmbeddingModel, error) {
return nil, fmt.Errorf("embedding models not supported")
}

func (p *CustomProvider) ImageModel(modelID string) (provider.ImageModel, error) {
return nil, fmt.Errorf("image models not supported")
}

func (p *CustomProvider) SpeechModel(modelID string) (provider.SpeechModel, error) {
return nil, fmt.Errorf("speech models not supported")
}

func (p *CustomProvider) TranscriptionModel(modelID string) (provider.TranscriptionModel, error) {
return nil, fmt.Errorf("transcription models not supported")
}

func (p *CustomProvider) RerankingModel(modelID string) (provider.RerankingModel, error) {
return nil, fmt.Errorf("reranking models not supported")
}

Using Your Custom Provider​

func main() {
// Create provider
provider := NewCustomProvider("your-api-key")

// Get model
model, err := provider.LanguageModel("custom-model")
if err != nil {
log.Fatal(err)
}

// Use with ai package
result, err := ai.GenerateText(context.Background(), ai.GenerateTextOptions{
Model: model,
Prompt: "Hello!",
})
if err != nil {
log.Fatal(err)
}

fmt.Println(result.Text)
}

Implementing Streaming​

provider.TextStream only has Next, Err, and Close — it does not implement io.Reader. Don't add a Read method to a stream wrapper; consume chunks with Next() in a loop, as shown in the provider.TextStream doc comment.

type CustomTextStream struct {
chunks chan *provider.StreamChunk
err error
}

func NewCustomTextStream() *CustomTextStream {
return &CustomTextStream{
chunks: make(chan *provider.StreamChunk, 10),
}
}

func (s *CustomTextStream) Next() (*provider.StreamChunk, error) {
chunk, ok := <-s.chunks
if !ok {
return nil, io.EOF
}
return chunk, nil
}

func (s *CustomTextStream) Close() error {
close(s.chunks)
return nil
}

func (s *CustomTextStream) Err() error {
return s.err
}

See Also​