amaniquery-agent / internal /generator /embedding_test.go
Deployment
Automated deployment update
4b1daed
Raw
History Blame Contribute Delete
3.33 kB
package generator_test
import (
"testing"
"time"
"github.com/AmaniQuery/amaniquery/internal/generator"
)
// TestNewOpenAIEmbeddingClient tests client creation
func TestNewOpenAIEmbeddingClient(t *testing.T) {
cfg := generator.EmbeddingConfig{
Provider: "openai",
APIKey: "test-api-key",
Model: "text-embedding-3-small",
Dimension: 1536,
BatchSize: 100,
Timeout: 30 * time.Second,
}
client := generator.NewOpenAIEmbeddingClient(cfg)
if client == nil {
t.Fatal("Expected non-nil client")
}
}
// TestEmbeddingClient_DefaultValues tests default value assignment
func TestEmbeddingClient_DefaultValues(t *testing.T) {
cfg := generator.EmbeddingConfig{
Provider: "openai",
APIKey: "test-key",
}
client := generator.NewOpenAIEmbeddingClient(cfg)
if client.GetModel() == "" {
t.Error("Expected default model to be set")
}
if client.GetDimension() == 0 {
t.Error("Expected default dimension to be set")
}
}
// TestEmbeddingClient_GetDimension tests dimension getter
func TestEmbeddingClient_GetDimension(t *testing.T) {
cfg := generator.EmbeddingConfig{
Provider: "openai",
APIKey: "test-key",
Dimension: 768,
}
client := generator.NewOpenAIEmbeddingClient(cfg)
if client.GetDimension() != 768 {
t.Errorf("Expected dimension 768, got %d", client.GetDimension())
}
}
// TestEmbeddingClient_GetModel tests model getter
func TestEmbeddingClient_GetModel(t *testing.T) {
cfg := generator.EmbeddingConfig{
Provider: "openai",
APIKey: "test-key",
Model: "text-embedding-ada-002",
}
client := generator.NewOpenAIEmbeddingClient(cfg)
if client.GetModel() != "text-embedding-ada-002" {
t.Errorf("Expected model 'text-embedding-ada-002', got '%s'", client.GetModel())
}
}
// TestEmbeddingClient_ProviderBaseURL tests provider-specific base URL
func TestEmbeddingClient_ProviderBaseURL(t *testing.T) {
testCases := []struct {
provider string
expectValid bool
}{
{"openai", true},
{"ollama", true},
{"unknown", true}, // Falls back to openai
}
for _, tc := range testCases {
t.Run(tc.provider, func(t *testing.T) {
cfg := generator.EmbeddingConfig{
Provider: tc.provider,
APIKey: "test-key",
}
client := generator.NewOpenAIEmbeddingClient(cfg)
if client == nil && tc.expectValid {
t.Error("Expected valid client")
}
})
}
}
// TestEmbeddingClient_CustomBaseURL tests custom base URL
func TestEmbeddingClient_CustomBaseURL(t *testing.T) {
cfg := generator.EmbeddingConfig{
Provider: "openai",
APIKey: "test-key",
BaseURL: "https://custom-api.example.com/v1",
}
client := generator.NewOpenAIEmbeddingClient(cfg)
if client == nil {
t.Fatal("Expected non-nil client with custom base URL")
}
}
// TestEmbeddingConfig_Struct tests config structure
func TestEmbeddingConfig_Struct(t *testing.T) {
cfg := generator.EmbeddingConfig{
Provider: "openai",
APIKey: "api-key",
BaseURL: "https://api.example.com",
Model: "model-name",
Dimension: 512,
BatchSize: 50,
Timeout: time.Minute,
MaxRetries: 3,
}
if cfg.Provider != "openai" {
t.Error("Provider mismatch")
}
if cfg.Dimension != 512 {
t.Error("Dimension mismatch")
}
if cfg.BatchSize != 50 {
t.Error("BatchSize mismatch")
}
if cfg.MaxRetries != 3 {
t.Error("MaxRetries mismatch")
}
}