ollie/embedding/embedding_test.go

134 lines
3.3 KiB
Go

package embedding
import (
"os"
"path/filepath"
"testing"
)
func TestEmbedAndSimilarity(t *testing.T) {
// Find model directory
modelDir := os.Getenv("OLLIE_MODEL_DIR")
if modelDir == "" {
// Try common locations
home, _ := os.UserHomeDir()
candidates := []string{
"../data/models",
filepath.Join(home, ".local/share/ollie/models"),
}
for _, c := range candidates {
if _, err := os.Stat(filepath.Join(c, "model.onnx")); err == nil {
modelDir = c
break
}
}
}
if modelDir == "" {
t.Skip("model not found, set OLLIE_MODEL_DIR")
}
model, err := LoadModel(modelDir)
if err != nil {
t.Fatalf("LoadModel: %v", err)
}
defer model.Close()
// Test embedding
vec, err := model.Embed("hello world")
if err != nil {
t.Fatalf("Embed: %v", err)
}
if len(vec) != 384 {
t.Fatalf("expected 384 dimensions, got %d", len(vec))
}
// Test similarity - similar sentences should have high similarity
v1, _ := model.Embed("I love programming in Go")
v2, _ := model.Embed("Go is my favorite programming language")
v3, _ := model.Embed("The weather is nice today")
sim12 := CosineSimilarity(v1, v2)
sim13 := CosineSimilarity(v1, v3)
t.Logf("Similar sentences: %.4f", sim12)
t.Logf("Dissimilar sentences: %.4f", sim13)
if sim12 < 0.5 {
t.Errorf("expected similar sentences to have similarity > 0.5, got %.4f", sim12)
}
if sim13 > sim12 {
t.Errorf("expected similar sentences to have higher similarity than dissimilar")
}
}
func TestSkillMatching(t *testing.T) {
modelDir := os.Getenv("OLLIE_MODEL_DIR")
if modelDir == "" {
home, _ := os.UserHomeDir()
candidates := []string{
"../data/models",
filepath.Join(home, ".local/share/ollie/models"),
}
for _, c := range candidates {
if _, err := os.Stat(filepath.Join(c, "model.onnx")); err == nil {
modelDir = c
break
}
}
}
if modelDir == "" {
t.Skip("model not found")
}
model, err := LoadModel(modelDir)
if err != nil {
t.Fatalf("LoadModel: %v", err)
}
defer model.Close()
// Simulate skill descriptions
skills := []string{
"Interact with pascom Atlassian (Jira/Confluence). Use for issues, pages, and project management.",
"Interact with GitHub repositories, issues, PRs, workflows using the gh CLI.",
"Execute bash commands on remote dev server with synced mobydick workspace.",
"Write ast-grep rules for AST-based structural code search and analysis.",
}
skillVecs := make([]Vector, len(skills))
for i, s := range skills {
vec, err := model.Embed(s)
if err != nil {
t.Fatalf("Embed skill %d: %v", i, err)
}
skillVecs[i] = vec
}
// Test queries
queries := []struct {
query string
expected int // index of expected best match
}{
{"read jira ticket PR-12345", 0}, // should match Atlassian
{"create a github issue", 1}, // should match GitHub
{"run make on the remote server", 2}, // should match remote-bash
{"find all function calls in Go", 3}, // should match ast-grep
}
for _, tc := range queries {
qvec, _ := model.Embed(tc.query)
best := -1
bestSim := float32(-1)
for i, sv := range skillVecs {
sim := CosineSimilarity(qvec, sv)
t.Logf("%q vs skill[%d]: %.4f", tc.query, i, sim)
if sim > bestSim {
bestSim = sim
best = i
}
}
if best != tc.expected {
t.Errorf("%q: expected skill %d, got %d (sim=%.4f)", tc.query, tc.expected, best, bestSim)
}
}
}