package workbench import ( "context" "encoding/json" "net/http" "net/http/httptest" "strings" "testing" "time" ) // fakeOllama 模拟 Ollama 的 /api/embed 与 /api/pull。dim 固定产出向量维数, // requireModel 为 true 时首次 /api/embed 返回 404 触发 /api/pull。 func fakeOllama(t *testing.T, dim int, requireModel bool) *httptest.Server { t.Helper() pulled := false server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { switch r.URL.Path { case "/api/pull": pulled = true _ = json.NewEncoder(w).Encode(map[string]any{"status": "success"}) case "/api/embed": if requireModel && !pulled { w.WriteHeader(http.StatusNotFound) _ = json.NewEncoder(w).Encode(map[string]any{"error": "model 'bge-m3' not found"}) return } var body struct { Input []string `json:"input"` } _ = json.NewDecoder(r.Body).Decode(&body) embeddings := make([][]float32, len(body.Input)) for i := range body.Input { vector := make([]float32, dim) for j := range vector { vector[j] = float32(i + 1) } embeddings[i] = vector } _ = json.NewEncoder(w).Encode(map[string]any{"embeddings": embeddings}) default: http.NotFound(w, r) } })) t.Cleanup(server.Close) return server } func TestOllamaEmbedderBatching(t *testing.T) { server := fakeOllama(t, 4, false) embedder := NewOllamaEmbedder(OllamaEmbedderConfig{BaseURL: server.URL, Model: "bge-m3", Dim: 4, BatchSize: 2, Timeout: 5 * time.Second}) texts := []string{"a", "b", "c", "d", "e"} vectors, err := embedder.Embed(context.Background(), texts) if err != nil { t.Fatal(err) } if len(vectors) != len(texts) { t.Fatalf("got %d vectors want %d", len(vectors), len(texts)) } for i, vector := range vectors { if len(vector) != 4 { t.Fatalf("vector %d has dim %d want 4", i, len(vector)) } } if embedder.Dim() != 4 { t.Fatalf("Dim()=%d want 4", embedder.Dim()) } } func TestOllamaEmbedderDimensionMismatch(t *testing.T) { server := fakeOllama(t, 3, false) embedder := NewOllamaEmbedder(OllamaEmbedderConfig{BaseURL: server.URL, Model: "bge-m3", Dim: 4, BatchSize: 2, Timeout: 5 * time.Second}) _, err := embedder.Embed(context.Background(), []string{"x"}) if err == nil || !strings.Contains(err.Error(), "dimension mismatch") { t.Fatalf("expected dimension mismatch error, got %v", err) } } func TestOllamaEmbedderPullsModelOn404(t *testing.T) { server := fakeOllama(t, 4, true) embedder := NewOllamaEmbedder(OllamaEmbedderConfig{BaseURL: server.URL, Model: "bge-m3", Dim: 4, BatchSize: 8, Timeout: 5 * time.Second}) vectors, err := embedder.Embed(context.Background(), []string{"hello", "world"}) if err != nil { t.Fatalf("expected pull-then-retry to succeed, got %v", err) } if len(vectors) != 2 { t.Fatalf("got %d vectors want 2", len(vectors)) } } func TestOllamaEmbedderEmptyInput(t *testing.T) { embedder := NewOllamaEmbedder(OllamaEmbedderConfig{BaseURL: "http://127.0.0.1:1", Model: "bge-m3", Dim: 4, BatchSize: 2, Timeout: time.Second}) vectors, err := embedder.Embed(context.Background(), nil) if err != nil || len(vectors) != 0 { t.Fatalf("empty input should return nil without HTTP: vectors=%d err=%v", len(vectors), err) } } func TestFormatVector(t *testing.T) { if got := formatVector([]float32{1, 2.5, -0.25}); got != "[1,2.5,-0.25]" { t.Fatalf("formatVector got %q", got) } if got := formatVector([]float32{}); got != "[]" { t.Fatalf("formatVector empty got %q", got) } }