Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions embeddings/cohere.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,9 @@ import (
"time"
)

// cohereEndpoint is overridable in tests.
var cohereEndpoint = "https://api.cohere.com/v2/embed"

// cohere implements the Cohere embedding API.
// Supported models: embed-english-v3.0, embed-multilingual-v3.0.
type cohere struct {
Expand Down Expand Up @@ -39,7 +42,7 @@ func (p *cohere) Embed(ctx context.Context, text string) ([]float32, error) {
if err != nil {
return nil, fmt.Errorf("cohere: marshal: %w", err)
}
req, err := http.NewRequestWithContext(ctx, "POST", "https://api.cohere.com/v2/embed", bytes.NewReader(body))
req, err := http.NewRequestWithContext(ctx, "POST", cohereEndpoint, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("cohere: create request: %w", err)
}
Expand Down Expand Up @@ -98,7 +101,7 @@ func (p *cohere) EmbedBatch(ctx context.Context, texts []string) ([][]float32, e
if err != nil {
return nil, fmt.Errorf("cohere: marshal batch: %w", err)
}
req, err := http.NewRequestWithContext(ctx, "POST", "https://api.cohere.com/v2/embed", bytes.NewReader(body))
req, err := http.NewRequestWithContext(ctx, "POST", cohereEndpoint, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("cohere: create request: %w", err)
}
Expand Down Expand Up @@ -159,7 +162,7 @@ func (p *cohere) EmbedWithMode(ctx context.Context, text string, mode EmbedMode)
if err != nil {
return nil, fmt.Errorf("cohere: marshal: %w", err)
}
req, err := http.NewRequestWithContext(ctx, "POST", "https://api.cohere.com/v2/embed", bytes.NewReader(body))
req, err := http.NewRequestWithContext(ctx, "POST", cohereEndpoint, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("cohere: create request: %w", err)
}
Expand Down
54 changes: 54 additions & 0 deletions embeddings/memo_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,60 @@ func (p *namedProvider) Name() string { return p.name }
// TestMemoizedProvider_ModelSwapInvalidates pins the core #1 fix: a memo built
// for one model must not serve its vectors to a different model. Since the memo
// is namespaced by Name(), each model has its own key space.
func TestEmbeddingMemoLen(t *testing.T) {
t.Parallel()
m := NewEmbeddingMemo(10)
if m.Len() != 0 {
t.Fatalf("expected empty memo, got %d entries", m.Len())
}
m.Put("a", []float32{1})
if m.Len() != 1 {
t.Fatalf("expected 1 entry, got %d", m.Len())
}
m.Put("b", []float32{2})
if m.Len() != 2 {
t.Fatalf("expected 2 entries, got %d", m.Len())
}
}

func TestEmbeddingMemoPutModeUpdate(t *testing.T) {
t.Parallel()
m := NewEmbeddingMemo(10)
m.Put("a", []float32{1, 2, 3})
if m.Len() != 1 {
t.Fatalf("expected 1 entry after put, got %d", m.Len())
}
// Update existing entry
m.Put("a", []float32{4, 5, 6})
if m.Len() != 1 {
t.Fatalf("expected still 1 entry after update, got %d", m.Len())
}
vec, ok := m.Get("a")
if !ok || vec[0] != 4 {
t.Errorf("expected updated vector [4 5 6], got %v", vec)
}
}

func TestEmbeddingMemoNewWithZeroEntries(t *testing.T) {
t.Parallel()
m := NewEmbeddingMemoNS("test", 0)
if m.max != 1024 {
t.Errorf("expected default 1024 max for zero input, got %d", m.max)
}
}

func TestMemoizedProviderNameAndDims(t *testing.T) {
t.Parallel()
inner := NewLocal()
p := NewMemoizedProvider(inner, 10)
if p.Name() != inner.Name() {
t.Errorf("expected Name=%q, got %q", inner.Name(), p.Name())
}
if p.Dims() != inner.Dims() {
t.Errorf("expected Dims=%d, got %d", inner.Dims(), p.Dims())
}
}

func TestMemoizedProvider_ModelSwapInvalidates(t *testing.T) {
t.Parallel()
ctx := context.Background()
Expand Down
14 changes: 9 additions & 5 deletions embeddings/provider.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,10 @@ import (
// between OpenAI and Voyage is ever needed.
var httpClient = &http.Client{Timeout: 30 * time.Second}

// Endpoint URLs are package-level vars so tests can override them with httptest servers.
var openAIEndpoint = "https://api.openai.com/v1/embeddings"
var voyageEndpoint = "https://api.voyageai.com/v1/embeddings"

// maxRetries is the maximum number of rate-limit retries before giving up.
const maxRetries = 5

Expand Down Expand Up @@ -71,7 +75,7 @@ func (p *openAI) Embed(ctx context.Context, text string) ([]float32, error) {
if err != nil {
return nil, fmt.Errorf("openai: marshal: %w", err)
}
req, err := http.NewRequestWithContext(ctx, "POST", "https://api.openai.com/v1/embeddings", bytes.NewReader(body))
req, err := http.NewRequestWithContext(ctx, "POST", openAIEndpoint, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("openai: create request: %w", err)
}
Expand Down Expand Up @@ -147,7 +151,7 @@ func (p *openAI) EmbedBatch(ctx context.Context, texts []string) ([][]float32, e
if err != nil {
return nil, fmt.Errorf("openai: marshal batch: %w", err)
}
req, err := http.NewRequestWithContext(ctx, "POST", "https://api.openai.com/v1/embeddings", bytes.NewReader(body))
req, err := http.NewRequestWithContext(ctx, "POST", openAIEndpoint, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("openai: create request: %w", err)
}
Expand Down Expand Up @@ -252,7 +256,7 @@ func (p *voyage) Embed(ctx context.Context, text string) ([]float32, error) {
if err != nil {
return nil, fmt.Errorf("voyage: marshal: %w", err)
}
req, err := http.NewRequestWithContext(ctx, "POST", "https://api.voyageai.com/v1/embeddings", bytes.NewReader(body))
req, err := http.NewRequestWithContext(ctx, "POST", voyageEndpoint, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("voyage: create request: %w", err)
}
Expand Down Expand Up @@ -312,7 +316,7 @@ func (p *voyage) EmbedBatch(ctx context.Context, texts []string) ([][]float32, e
if err != nil {
return nil, fmt.Errorf("voyage: marshal batch: %w", err)
}
req, err := http.NewRequestWithContext(ctx, "POST", "https://api.voyageai.com/v1/embeddings", bytes.NewReader(body))
req, err := http.NewRequestWithContext(ctx, "POST", voyageEndpoint, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("voyage: create request: %w", err)
}
Expand Down Expand Up @@ -376,7 +380,7 @@ func (p *voyage) EmbedWithMode(ctx context.Context, text string, mode EmbedMode)
if err != nil {
return nil, fmt.Errorf("voyage: marshal: %w", err)
}
req, err := http.NewRequestWithContext(ctx, "POST", "https://api.voyageai.com/v1/embeddings", bytes.NewReader(body))
req, err := http.NewRequestWithContext(ctx, "POST", voyageEndpoint, bytes.NewReader(body))
if err != nil {
return nil, fmt.Errorf("voyage: create request: %w", err)
}
Expand Down
Loading
Loading