test(llm,embeddings): add unit tests for mock client, types, and Ollama embedder
- Add comprehensive MockLLMClient tests (generate, stream, match variants) - Add TypeRef JSON marshaling tests and StopReason value tests - Add Ollama embedder tests for config defaults and embedding requests
This commit is contained in:
parent
1a8f1557f6
commit
2eed2033f0
3 changed files with 495 additions and 0 deletions
237
pkg/llm/mock/mock_test.go
Normal file
237
pkg/llm/mock/mock_test.go
Normal file
|
|
@ -0,0 +1,237 @@
|
||||||
|
package mock
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/VictorVargas/rony-llm-agent/pkg/llm"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMockLLMClient_Generate_Default(t *testing.T) {
|
||||||
|
client := New()
|
||||||
|
resp, err := client.Generate(context.Background(), llm.CompletionRequest{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if resp.Content != "default response" {
|
||||||
|
t.Errorf("expected 'default response', got %q", resp.Content)
|
||||||
|
}
|
||||||
|
if resp.StopReason != llm.StopReasonEndTurn {
|
||||||
|
t.Errorf("expected StopReasonEndTurn, got %q", resp.StopReason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockLLMClient_Generate_Custom(t *testing.T) {
|
||||||
|
expected := llm.CompletionResponse{
|
||||||
|
ID: "test_123",
|
||||||
|
Model: "gpt-4",
|
||||||
|
Content: "hello world",
|
||||||
|
StopReason: llm.StopReasonToolUse,
|
||||||
|
}
|
||||||
|
client := NewWithGenerate(expected)
|
||||||
|
|
||||||
|
resp, err := client.Generate(context.Background(), llm.CompletionRequest{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if resp.ID != "test_123" {
|
||||||
|
t.Errorf("expected ID 'test_123', got %q", resp.ID)
|
||||||
|
}
|
||||||
|
if resp.Content != "hello world" {
|
||||||
|
t.Errorf("expected 'hello world', got %q", resp.Content)
|
||||||
|
}
|
||||||
|
if resp.StopReason != llm.StopReasonToolUse {
|
||||||
|
t.Errorf("expected StopReasonToolUse, got %q", resp.StopReason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockLLMClient_Generate_Error(t *testing.T) {
|
||||||
|
client := &MockLLMClient{
|
||||||
|
GenerateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
||||||
|
return llm.CompletionResponse{}, fmt.Errorf("test error")
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := client.Generate(context.Background(), llm.CompletionRequest{})
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error, got nil")
|
||||||
|
}
|
||||||
|
if err.Error() != "test error" {
|
||||||
|
t.Errorf("expected 'test error', got %q", err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockLLMClient_Stream_Default(t *testing.T) {
|
||||||
|
client := New()
|
||||||
|
var chunks []string
|
||||||
|
for chunk, err := range client.Stream(context.Background(), llm.CompletionRequest{}) {
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
chunks = append(chunks, chunk.Delta)
|
||||||
|
}
|
||||||
|
if len(chunks) != 1 || chunks[0] != "default response" {
|
||||||
|
t.Errorf("expected single 'default response' chunk, got %v", chunks)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockLLMClient_Stream_Custom(t *testing.T) {
|
||||||
|
expectedChunks := []llm.StreamChunk{
|
||||||
|
{Delta: "chunk1"},
|
||||||
|
{Delta: "chunk2", ToolCalls: []llm.ToolCall{{ID: "call_1"}}},
|
||||||
|
{FinishReason: "stop", Usage: llm.TokenUsage{InputTokens: 5, OutputTokens: 10}},
|
||||||
|
}
|
||||||
|
client := NewWithStream(expectedChunks)
|
||||||
|
|
||||||
|
var actualChunks []llm.StreamChunk
|
||||||
|
for chunk, err := range client.Stream(context.Background(), llm.CompletionRequest{}) {
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
actualChunks = append(actualChunks, chunk)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(actualChunks) != 3 {
|
||||||
|
t.Errorf("expected 3 chunks, got %d", len(actualChunks))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if actualChunks[0].Delta != "chunk1" {
|
||||||
|
t.Errorf("expected 'chunk1', got %q", actualChunks[0].Delta)
|
||||||
|
}
|
||||||
|
if len(actualChunks[1].ToolCalls) != 1 {
|
||||||
|
t.Errorf("expected 1 tool call, got %d", len(actualChunks[1].ToolCalls))
|
||||||
|
}
|
||||||
|
if actualChunks[2].FinishReason != "stop" {
|
||||||
|
t.Errorf("expected 'stop', got %q", actualChunks[2].FinishReason)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockLLMClient_Name_Default(t *testing.T) {
|
||||||
|
client := New()
|
||||||
|
if name := client.Name(); name != "mock" {
|
||||||
|
t.Errorf("expected 'mock', got %q", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockLLMClient_Name_Custom(t *testing.T) {
|
||||||
|
client := &MockLLMClient{
|
||||||
|
NameFunc: func() string { return "custom-model" },
|
||||||
|
}
|
||||||
|
if name := client.Name(); name != "custom-model" {
|
||||||
|
t.Errorf("expected 'custom-model', got %q", name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockLLMClient_Capabilities_Default(t *testing.T) {
|
||||||
|
client := New()
|
||||||
|
caps := client.Capabilities()
|
||||||
|
if !caps.SupportsTools {
|
||||||
|
t.Error("expected SupportsTools to be true")
|
||||||
|
}
|
||||||
|
if caps.SupportsVision {
|
||||||
|
t.Error("expected SupportsVision to be false")
|
||||||
|
}
|
||||||
|
if !caps.SupportsJSON {
|
||||||
|
t.Error("expected SupportsJSON to be true")
|
||||||
|
}
|
||||||
|
if caps.MaxContextWindow != 128000 {
|
||||||
|
t.Errorf("expected MaxContextWindow 128000, got %d", caps.MaxContextWindow)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockLLMClient_Capabilities_Custom(t *testing.T) {
|
||||||
|
client := &MockLLMClient{
|
||||||
|
CapabilitiesFunc: func() llm.ProviderCapabilities {
|
||||||
|
return llm.ProviderCapabilities{
|
||||||
|
SupportsTools: false,
|
||||||
|
SupportsVision: true,
|
||||||
|
SupportsJSON: false,
|
||||||
|
MaxContextWindow: 32000,
|
||||||
|
}
|
||||||
|
},
|
||||||
|
}
|
||||||
|
caps := client.Capabilities()
|
||||||
|
if caps.SupportsTools {
|
||||||
|
t.Error("expected SupportsTools to be false")
|
||||||
|
}
|
||||||
|
if !caps.SupportsVision {
|
||||||
|
t.Error("expected SupportsVision to be true")
|
||||||
|
}
|
||||||
|
if caps.MaxContextWindow != 32000 {
|
||||||
|
t.Errorf("expected MaxContextWindow 32000, got %d", caps.MaxContextWindow)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockLLMClient_MatchResponse(t *testing.T) {
|
||||||
|
client := NewWithMatch([]MatchResponse{
|
||||||
|
{Match: "hello", Response: "hi there!"},
|
||||||
|
{Match: "world", Response: "earth"},
|
||||||
|
{Match: "*", Response: "default"},
|
||||||
|
})
|
||||||
|
|
||||||
|
resp, _ := client.Generate(context.Background(), llm.CompletionRequest{
|
||||||
|
Messages: []llm.Message{{Role: llm.RoleUser, Content: "hello"}},
|
||||||
|
})
|
||||||
|
if resp.Content != "hi there!" {
|
||||||
|
t.Errorf("expected 'hi there!', got %q", resp.Content)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, _ = client.Generate(context.Background(), llm.CompletionRequest{
|
||||||
|
Messages: []llm.Message{{Role: llm.RoleUser, Content: "world"}},
|
||||||
|
})
|
||||||
|
if resp.Content != "earth" {
|
||||||
|
t.Errorf("expected 'earth', got %q", resp.Content)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, _ = client.Generate(context.Background(), llm.CompletionRequest{
|
||||||
|
Messages: []llm.Message{{Role: llm.RoleUser, Content: "anything else"}},
|
||||||
|
})
|
||||||
|
if resp.Content != "default" {
|
||||||
|
t.Errorf("expected 'default', got %q", resp.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockLLMClient_MatchResponse_NoMatch(t *testing.T) {
|
||||||
|
client := NewWithMatch([]MatchResponse{
|
||||||
|
{Match: "exact", Response: "matched"},
|
||||||
|
})
|
||||||
|
|
||||||
|
resp, _ := client.Generate(context.Background(), llm.CompletionRequest{
|
||||||
|
Messages: []llm.Message{{Role: llm.RoleUser, Content: "no match here"}},
|
||||||
|
})
|
||||||
|
if resp.Content != "no match" {
|
||||||
|
t.Errorf("expected 'no match', got %q", resp.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatchesAny(t *testing.T) {
|
||||||
|
if !matchesAny("hello", []llm.Message{
|
||||||
|
{Role: llm.RoleUser, Content: "say hello to me"},
|
||||||
|
}) {
|
||||||
|
t.Error("expected matchesAny to find 'hello' in message")
|
||||||
|
}
|
||||||
|
|
||||||
|
if matchesAny("goodbye", []llm.Message{
|
||||||
|
{Role: llm.RoleUser, Content: "say hello to me"},
|
||||||
|
}) {
|
||||||
|
t.Error("expected matchesAny to not find 'goodbye' in message")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMatchesAny_Wildcard(t *testing.T) {
|
||||||
|
resp, _ := (&MockLLMClient{
|
||||||
|
GenerateFunc: func(ctx context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
||||||
|
for _, r := range []MatchResponse{{Match: "*", Response: "wildcard"}} {
|
||||||
|
if r.Match == "*" || matchesAny(r.Match, req.Messages) {
|
||||||
|
return llm.CompletionResponse{Content: r.Response}, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return llm.CompletionResponse{}, nil
|
||||||
|
},
|
||||||
|
}).Generate(context.Background(), llm.CompletionRequest{})
|
||||||
|
|
||||||
|
if resp.Content != "wildcard" {
|
||||||
|
t.Errorf("expected 'wildcard', got %q", resp.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
145
pkg/llm/types_test.go
Normal file
145
pkg/llm/types_test.go
Normal file
|
|
@ -0,0 +1,145 @@
|
||||||
|
package llm
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestToolRef_MarshalJSON(t *testing.T) {
|
||||||
|
ref := &ToolRef{Name: "my_tool"}
|
||||||
|
data, err := json.Marshal(ref)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var result struct {
|
||||||
|
Type string `json:"type"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(data, &result); err != nil {
|
||||||
|
t.Fatalf("failed to unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
if result.Type != "function" {
|
||||||
|
t.Errorf("expected type 'function', got %q", result.Type)
|
||||||
|
}
|
||||||
|
if result.Name != "my_tool" {
|
||||||
|
t.Errorf("expected name 'my_tool', got %q", result.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStopReason_Values(t *testing.T) {
|
||||||
|
if string(StopReasonEndTurn) != "end_turn" {
|
||||||
|
t.Error("StopReasonEndTurn should be 'end_turn'")
|
||||||
|
}
|
||||||
|
if string(StopReasonToolUse) != "tool_use" {
|
||||||
|
t.Error("StopReasonToolUse should be 'tool_use'")
|
||||||
|
}
|
||||||
|
if string(StopReasonMaxTokens) != "max_tokens" {
|
||||||
|
t.Error("StopReasonMaxTokens should be 'max_tokens'")
|
||||||
|
}
|
||||||
|
if string(StopReasonStopSeq) != "stop_sequence" {
|
||||||
|
t.Error("StopReasonStopSeq should be 'stop_sequence'")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompletionRequest_JSON(t *testing.T) {
|
||||||
|
req := CompletionRequest{
|
||||||
|
Model: "gpt-4",
|
||||||
|
Temperature: float32Ptr(0.7),
|
||||||
|
MaxTokens: intPtr(100),
|
||||||
|
Metadata: map[string]string{"key": "val"},
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(req)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var unmarshaled CompletionRequest
|
||||||
|
if err := json.Unmarshal(data, &unmarshaled); err != nil {
|
||||||
|
t.Fatalf("failed to unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
if unmarshaled.Model != "gpt-4" {
|
||||||
|
t.Errorf("expected model 'gpt-4', got %q", unmarshaled.Model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCompletionResponse_JSON(t *testing.T) {
|
||||||
|
resp := CompletionResponse{
|
||||||
|
ID: "resp_123",
|
||||||
|
Model: "gpt-4",
|
||||||
|
Content: "hello world",
|
||||||
|
StopReason: StopReasonEndTurn,
|
||||||
|
Usage: TokenUsage{
|
||||||
|
InputTokens: 10,
|
||||||
|
OutputTokens: 5,
|
||||||
|
TotalTokens: 15,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(resp)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var unmarshaled CompletionResponse
|
||||||
|
if err := json.Unmarshal(data, &unmarshaled); err != nil {
|
||||||
|
t.Fatalf("failed to unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
if unmarshaled.ID != "resp_123" {
|
||||||
|
t.Errorf("expected ID 'resp_123', got %q", unmarshaled.ID)
|
||||||
|
}
|
||||||
|
if unmarshaled.Content != "hello world" {
|
||||||
|
t.Errorf("expected content 'hello world', got %q", unmarshaled.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolCall_JSON(t *testing.T) {
|
||||||
|
call := ToolCall{
|
||||||
|
ID: "call_1",
|
||||||
|
Name: "my_tool",
|
||||||
|
Arguments: json.RawMessage(`{"key": "val"}`),
|
||||||
|
Thought: "let me think",
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(call)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var unmarshaled ToolCall
|
||||||
|
if err := json.Unmarshal(data, &unmarshaled); err != nil {
|
||||||
|
t.Fatalf("failed to unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
if unmarshaled.ID != "call_1" {
|
||||||
|
t.Errorf("expected ID 'call_1', got %q", unmarshaled.ID)
|
||||||
|
}
|
||||||
|
if string(unmarshaled.Arguments) != `{"key":"val"}` {
|
||||||
|
t.Errorf("unexpected arguments: %s", unmarshaled.Arguments)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStreamChunk_JSON(t *testing.T) {
|
||||||
|
chunk := StreamChunk{
|
||||||
|
Delta: "hello",
|
||||||
|
ToolCalls: []ToolCall{},
|
||||||
|
}
|
||||||
|
|
||||||
|
data, err := json.Marshal(chunk)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var unmarshaled StreamChunk
|
||||||
|
if err := json.Unmarshal(data, &unmarshaled); err != nil {
|
||||||
|
t.Fatalf("failed to unmarshal: %v", err)
|
||||||
|
}
|
||||||
|
if unmarshaled.Delta != "hello" {
|
||||||
|
t.Errorf("expected delta 'hello', got %q", unmarshaled.Delta)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
func float32Ptr(f float32) *float32 { return &f }
|
||||||
|
func intPtr(i int) *int { return &i }
|
||||||
113
pkg/rag/embeddings/embeddings_test.go
Normal file
113
pkg/rag/embeddings/embeddings_test.go
Normal file
|
|
@ -0,0 +1,113 @@
|
||||||
|
package embeddings
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewOllama_Defaults(t *testing.T) {
|
||||||
|
e, err := NewOllama(Config{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if e.baseURL != "http://localhost:11434" {
|
||||||
|
t.Errorf("expected 'http://localhost:11434', got %q", e.baseURL)
|
||||||
|
}
|
||||||
|
if e.model != "nomic-embed-text" {
|
||||||
|
t.Errorf("expected 'nomic-embed-text', got %q", e.model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewOllama_Custom(t *testing.T) {
|
||||||
|
e, err := NewOllama(Config{
|
||||||
|
BaseURL: "http://custom:8080",
|
||||||
|
Model: "my-embedder",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if e.baseURL != "http://custom:8080" {
|
||||||
|
t.Errorf("expected 'http://custom:8080', got %q", e.baseURL)
|
||||||
|
}
|
||||||
|
if e.model != "my-embedder" {
|
||||||
|
t.Errorf("expected 'my-embedder', got %q", e.model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOllama_Dimensions(t *testing.T) {
|
||||||
|
e, _ := NewOllama(Config{})
|
||||||
|
if dims := e.Dimensions(); dims != 768 {
|
||||||
|
t.Errorf("expected 768 dimensions, got %d", dims)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOllama_Embed_Success(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
w.Write([]byte(`{"embedding": [0.1, 0.2, 0.3]}`))
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
e, _ := NewOllama(Config{BaseURL: server.URL})
|
||||||
|
ctx := context.Background()
|
||||||
|
vector, err := e.Embed(ctx, "test input")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(vector) != 3 {
|
||||||
|
t.Errorf("expected 3 dimensions, got %d", len(vector))
|
||||||
|
}
|
||||||
|
if vector[0] != 0.1 || vector[1] != 0.2 || vector[2] != 0.3 {
|
||||||
|
t.Errorf("unexpected vector: %v", vector)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOllama_Embed_Error(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusInternalServerError)
|
||||||
|
w.Write([]byte("internal error"))
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
e, _ := NewOllama(Config{BaseURL: server.URL})
|
||||||
|
ctx := context.Background()
|
||||||
|
_, err := e.Embed(ctx, "test input")
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error, got nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockEmbedder_Default(t *testing.T) {
|
||||||
|
m := &MockEmbedder{}
|
||||||
|
vector, err := m.Embed(context.Background(), "test")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("unexpected error: %v", err)
|
||||||
|
}
|
||||||
|
if len(vector) != 3 {
|
||||||
|
t.Errorf("expected 3 dimensions, got %d", len(vector))
|
||||||
|
}
|
||||||
|
if dims := m.Dimensions(); dims != 3 {
|
||||||
|
t.Errorf("expected 3 dimensions, got %d", dims)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMockEmbedder_Custom(t *testing.T) {
|
||||||
|
m := &MockEmbedder{
|
||||||
|
EmbedFunc: func(ctx context.Context, text string) ([]float32, error) {
|
||||||
|
return []float32{1.0, 2.0}, nil
|
||||||
|
},
|
||||||
|
DimensionsFn: func() int {
|
||||||
|
return 2
|
||||||
|
},
|
||||||
|
}
|
||||||
|
vector, _ := m.Embed(context.Background(), "test")
|
||||||
|
if len(vector) != 2 {
|
||||||
|
t.Errorf("expected 2 dimensions, got %d", len(vector))
|
||||||
|
}
|
||||||
|
if dims := m.Dimensions(); dims != 2 {
|
||||||
|
t.Errorf("expected 2 dimensions, got %d", dims)
|
||||||
|
}
|
||||||
|
}
|
||||||
Loading…
Reference in a new issue