130 lines
3.2 KiB
Go
130 lines
3.2 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, err := model.EmbedBatch(skills)
|
|
if err != nil {
|
|
t.Fatalf("EmbedBatch: %v", err)
|
|
}
|
|
|
|
// 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)
|
|
}
|
|
}
|
|
}
|