Files

120 lines
2.4 KiB
Go
Raw Permalink Normal View History

package index
import (
"math"
"sort"
"sync"
)
type SearchHit struct {
ID string
Score float64
}
type InMemoryVectorIndex struct {
mu sync.RWMutex
ids []string
vectors [][]float64
idToIdx map[string]int
}
func NewInMemoryVectorIndex() *InMemoryVectorIndex {
return &InMemoryVectorIndex{
idToIdx: make(map[string]int),
}
}
func (idx *InMemoryVectorIndex) Add(id string, vector []float64) {
idx.mu.Lock()
defer idx.mu.Unlock()
if i, exists := idx.idToIdx[id]; exists {
idx.vectors[i] = vector
return
}
idx.idToIdx[id] = len(idx.ids)
idx.ids = append(idx.ids, id)
idx.vectors = append(idx.vectors, vector)
}
func (idx *InMemoryVectorIndex) Remove(id string) {
idx.mu.Lock()
defer idx.mu.Unlock()
i, exists := idx.idToIdx[id]
if !exists {
return
}
delete(idx.idToIdx, id)
last := len(idx.ids) - 1
if i < last {
idx.ids[i] = idx.ids[last]
idx.vectors[i] = idx.vectors[last]
idx.idToIdx[idx.ids[i]] = i
}
idx.ids = idx.ids[:last]
idx.vectors = idx.vectors[:last]
}
func (idx *InMemoryVectorIndex) RemoveByPrefix(prefix string) {
idx.mu.Lock()
defer idx.mu.Unlock()
n := 0
for i, id := range idx.ids {
if len(id) >= len(prefix) && id[:len(prefix)] == prefix {
delete(idx.idToIdx, id)
continue
}
if i != n {
idx.ids[n] = idx.ids[i]
idx.vectors[n] = idx.vectors[i]
idx.idToIdx[idx.ids[n]] = n
}
n++
}
idx.ids = idx.ids[:n]
idx.vectors = idx.vectors[:n]
}
func (idx *InMemoryVectorIndex) Search(query []float64, topK int) []SearchHit {
idx.mu.RLock()
defer idx.mu.RUnlock()
scores := make([]SearchHit, len(idx.ids))
for i, vec := range idx.vectors {
scores[i] = SearchHit{ID: idx.ids[i], Score: cosine(query, vec)}
}
sort.Slice(scores, func(i, j int) bool { return scores[i].Score > scores[j].Score })
if topK > len(scores) {
topK = len(scores)
}
return scores[:topK]
}
func (idx *InMemoryVectorIndex) GetVector(id string) ([]float64, bool) {
idx.mu.RLock()
defer idx.mu.RUnlock()
i, exists := idx.idToIdx[id]
if !exists {
return nil, false
}
return idx.vectors[i], true
}
func (idx *InMemoryVectorIndex) Len() int {
idx.mu.RLock()
defer idx.mu.RUnlock()
return len(idx.ids)
}
func cosine(a, b []float64) float64 {
var dot, normA, normB float64
for i := range a {
dot += a[i] * b[i]
normA += a[i] * a[i]
normB += b[i] * b[i]
}
if normA == 0 || normB == 0 {
return 0
}
return dot / (math.Sqrt(normA) * math.Sqrt(normB))
}