Deployment
Automated deployment update
4b1daed
Raw
History Blame Contribute Delete
5.45 kB
// Package generator provides embedding generation capabilities
package generator
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"time"
)
// EmbeddingClient interface for embedding generation
type EmbeddingClient interface {
Generate(ctx context.Context, text string) ([]float32, error)
GenerateBatch(ctx context.Context, texts []string) ([][]float32, error)
}
// EmbeddingConfig for embedding client
type EmbeddingConfig struct {
Provider string
APIKey string
BaseURL string
Model string
Dimension int
BatchSize int
Timeout time.Duration
MaxRetries int
}
// OpenAIEmbeddingClient implements EmbeddingClient for OpenAI-compatible APIs
type OpenAIEmbeddingClient struct {
httpClient *http.Client
baseURL string
apiKey string
model string
dimension int
batchSize int
maxRetries int
}
// NewOpenAIEmbeddingClient creates a new embedding client
func NewOpenAIEmbeddingClient(cfg EmbeddingConfig) *OpenAIEmbeddingClient {
baseURL := cfg.BaseURL
if baseURL == "" {
switch cfg.Provider {
case "openai":
baseURL = "https://api.openai.com/v1"
case "ollama":
baseURL = "http://localhost:11434/v1"
default:
baseURL = "https://api.openai.com/v1"
}
}
model := cfg.Model
if model == "" {
model = "text-embedding-3-small"
}
dimension := cfg.Dimension
if dimension == 0 {
dimension = 1536
}
batchSize := cfg.BatchSize
if batchSize == 0 {
batchSize = 100
}
timeout := cfg.Timeout
if timeout == 0 {
timeout = 30 * time.Second
}
return &OpenAIEmbeddingClient{
httpClient: &http.Client{Timeout: timeout},
baseURL: baseURL,
apiKey: cfg.APIKey,
model: model,
dimension: dimension,
batchSize: batchSize,
maxRetries: cfg.MaxRetries,
}
}
// embeddingRequest represents an OpenAI embedding request
type embeddingRequest struct {
Model string `json:"model"`
Input interface{} `json:"input"`
Dimensions int `json:"dimensions,omitempty"`
}
// embeddingResponse represents an OpenAI embedding response
type embeddingResponse struct {
Object string `json:"object"`
Data []struct {
Object string `json:"object"`
Index int `json:"index"`
Embedding []float32 `json:"embedding"`
} `json:"data"`
Model string `json:"model"`
Usage struct {
PromptTokens int `json:"prompt_tokens"`
TotalTokens int `json:"total_tokens"`
} `json:"usage"`
}
// Generate creates an embedding for a single text
func (c *OpenAIEmbeddingClient) Generate(ctx context.Context, text string) ([]float32, error) {
embeddings, err := c.GenerateBatch(ctx, []string{text})
if err != nil {
return nil, err
}
if len(embeddings) == 0 {
return nil, fmt.Errorf("no embeddings returned")
}
return embeddings[0], nil
}
// GenerateBatch creates embeddings for multiple texts
func (c *OpenAIEmbeddingClient) GenerateBatch(ctx context.Context, texts []string) ([][]float32, error) {
if len(texts) == 0 {
return nil, nil
}
// Process in batches
var allEmbeddings [][]float32
for i := 0; i < len(texts); i += c.batchSize {
end := i + c.batchSize
if end > len(texts) {
end = len(texts)
}
batch := texts[i:end]
embeddings, err := c.processEmbeddingBatch(ctx, batch)
if err != nil {
return nil, fmt.Errorf("batch %d failed: %w", i/c.batchSize, err)
}
allEmbeddings = append(allEmbeddings, embeddings...)
}
return allEmbeddings, nil
}
func (c *OpenAIEmbeddingClient) processEmbeddingBatch(ctx context.Context, texts []string) ([][]float32, error) {
reqBody := embeddingRequest{
Model: c.model,
Input: texts,
}
// Only set dimensions for models that support it
if c.model == "text-embedding-3-small" || c.model == "text-embedding-3-large" {
reqBody.Dimensions = c.dimension
}
jsonBody, err := json.Marshal(reqBody)
if err != nil {
return nil, fmt.Errorf("failed to marshal request: %w", err)
}
var resp *embeddingResponse
var lastErr error
for attempt := 0; attempt <= c.maxRetries; attempt++ {
if attempt > 0 {
time.Sleep(time.Duration(attempt) * time.Second)
}
req, err := http.NewRequestWithContext(ctx, "POST", c.baseURL+"/embeddings", bytes.NewReader(jsonBody))
if err != nil {
return nil, fmt.Errorf("failed to create request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+c.apiKey)
httpResp, err := c.httpClient.Do(req)
if err != nil {
lastErr = err
continue
}
body, err := io.ReadAll(httpResp.Body)
httpResp.Body.Close()
if httpResp.StatusCode != http.StatusOK {
lastErr = fmt.Errorf("API error: %d - %s", httpResp.StatusCode, string(body))
if httpResp.StatusCode >= 500 {
continue // Retry on server errors
}
return nil, lastErr
}
if err := json.Unmarshal(body, &resp); err != nil {
lastErr = fmt.Errorf("failed to decode response: %w", err)
continue
}
break
}
if resp == nil {
return nil, lastErr
}
// Sort embeddings by index to ensure correct order
embeddings := make([][]float32, len(texts))
for _, data := range resp.Data {
if data.Index < len(embeddings) {
embeddings[data.Index] = data.Embedding
}
}
return embeddings, nil
}
// GetDimension returns the configured embedding dimension
func (c *OpenAIEmbeddingClient) GetDimension() int {
return c.dimension
}
// GetModel returns the configured model name
func (c *OpenAIEmbeddingClient) GetModel() string {
return c.model
}