ollie/embedding/index.go

84 lines
1.9 KiB
Go

package embedding
import (
"fmt"
"sort"
)
// Item represents something that can be embedded and matched.
type Item struct {
Name string
Description string // text used for embedding
}
// Index holds precomputed embeddings for semantic matching.
// Index is immutable after construction — no synchronization needed.
type Index struct {
model *Model
items []Item
vecs []Vector
}
// NewIndex creates an index from items using the given embedding model.
// The model is borrowed, not owned — caller is responsible for its lifecycle.
func NewIndex(model *Model, items []Item) (*Index, error) {
if model == nil {
return nil, fmt.Errorf("model is nil")
}
vecs := make([]Vector, len(items))
for i, item := range items {
vec, err := model.Embed(item.Description)
if err != nil {
return nil, fmt.Errorf("embed %q: %w", item.Name, err)
}
vecs[i] = vec
}
return &Index{
model: model,
items: items,
vecs: vecs,
}, nil
}
// Match returns items matching the query, sorted by relevance.
// threshold is the minimum cosine similarity (0-1) to include.
// limit is the maximum number of results (0 for no limit).
func (idx *Index) Match(query string, threshold float32, limit int) ([]MatchResult, error) {
if len(idx.items) == 0 {
return nil, nil
}
qvec, err := idx.model.Embed(query)
if err != nil {
return nil, fmt.Errorf("embed query: %w", err)
}
var results []MatchResult
for i, item := range idx.items {
sim := CosineSimilarity(qvec, idx.vecs[i])
if sim >= threshold {
results = append(results, MatchResult{
Item: item,
Score: sim,
})
}
}
sort.Slice(results, func(i, j int) bool {
return results[i].Score > results[j].Score
})
if limit > 0 && len(results) > limit {
results = results[:limit]
}
return results, nil
}
// MatchResult holds a matched item and its similarity score.
type MatchResult struct {
Item Item
Score float32
}