memoryweave/go/internal/storage/memvector.go

308 lines
7.3 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// 织忆 MemoryWeave — 内存向量存储LanceDB 零外部依赖实现)
package storage
import (
"fmt"
"math"
"sort"
"sync"
"time"
"github.com/xiaoxue/memoryweave/internal/models"
)
// MemLanceClient 内存实现 LanceDB 接口,生产部署时替换为 Rust LanceDB
type MemLanceClient struct {
mu sync.RWMutex
memories map[string]*memEntry
episodes []models.EpisodeRecord
auditLog []map[string]interface{}
embed *Embedder
}
type memEntry struct {
ID string
Content string
Vector []float32
Category string
Namespace string
AgentID string
Tier string
QualityScore float64
UsefulCount int
NotUsefulCount int
RecallCount int
Version int
VersionHistory []map[string]interface{}
Source string
IsDeleted bool
LastRecalledAt string
CreatedAt time.Time
}
func NewMemLanceClient(embedder *Embedder) *MemLanceClient {
return &MemLanceClient{
memories: make(map[string]*memEntry),
embed: embedder,
}
}
// ─── LanceDB 接口实现 ─────────────────────────────────
func (mlc *MemLanceClient) InsertEpisode(agentID, namespace, content, category string) (string, error) {
id := fmt.Sprintf("ep_%d", time.Now().UnixNano())
mlc.mu.Lock()
mlc.episodes = append(mlc.episodes, models.EpisodeRecord{
ID: id,
AgentID: agentID,
Namespace: namespace,
Content: content,
Category: category,
CreatedAt: time.Now(),
})
mlc.mu.Unlock()
return id, nil
}
func (mlc *MemLanceClient) InsertMemory(m models.MemoryRecord) error {
vec, _ := mlc.embed.EncodeSingle(m.Content)
if vec == nil {
vec = make([]float32, 1024)
}
mlc.mu.Lock()
mlc.memories[m.ID] = &memEntry{
ID: m.ID,
Content: m.Content,
Vector: vec,
Category: m.Category,
Namespace: m.Namespace,
AgentID: m.AgentID,
Tier: m.Tier,
Version: 1,
CreatedAt: time.Now(),
}
mlc.mu.Unlock()
return nil
}
func (mlc *MemLanceClient) Search(table string, vector []float32, topK int, namespaceFilter string) ([]models.MemoryRecord, error) {
mlc.mu.RLock()
defer mlc.mu.RUnlock()
type scored struct {
entry models.MemoryRecord
score float64
}
var candidates []scored
for _, e := range mlc.memories {
if e.IsDeleted {
continue
}
if namespaceFilter != "" && e.Namespace != namespaceFilter {
continue
}
sim := cosineSim(vector, e.Vector)
candidates = append(candidates, scored{
entry: models.MemoryRecord{
ID: e.ID,
Content: e.Content,
Category: e.Category,
},
score: sim,
})
}
sort.Slice(candidates, func(i, j int) bool {
return candidates[i].score > candidates[j].score
})
if topK > len(candidates) {
topK = len(candidates)
}
results := make([]models.MemoryRecord, topK)
for i := 0; i < topK; i++ {
results[i] = candidates[i].entry
}
return results, nil
}
func (mlc *MemLanceClient) Insert(table string, record any) error {
if r, ok := record.(models.MemoryRecord); ok {
return mlc.InsertMemory(r)
}
if r, ok := record.(models.EpisodeRecord); ok {
mlc.mu.Lock()
mlc.episodes = append(mlc.episodes, r)
mlc.mu.Unlock()
return nil
}
return fmt.Errorf("unknown record type")
}
func (mlc *MemLanceClient) GetTopByQuality(agentID string, limit int) ([]models.MemoryRecord, error) {
mlc.mu.RLock()
defer mlc.mu.RUnlock()
var results []models.MemoryRecord
for _, e := range mlc.memories {
if e.IsDeleted {
continue
}
results = append(results, models.MemoryRecord{
ID: e.ID,
Content: e.Content,
Category: e.Category,
Namespace: e.Namespace,
})
if len(results) >= limit {
break
}
}
return results, nil
}
func (mlc *MemLanceClient) Stats() (map[string]interface{}, error) {
mlc.mu.RLock()
defer mlc.mu.RUnlock()
return map[string]interface{}{
"memory_count": len(mlc.memories),
"episode_count": len(mlc.episodes),
"backend": "in-memory (zero-deps)",
}, nil
}
func (mlc *MemLanceClient) SoftDelete(id, reason string) error {
mlc.mu.Lock()
defer mlc.mu.Unlock()
if e, ok := mlc.memories[id]; ok {
e.IsDeleted = true
}
return nil
}
func (mlc *MemLanceClient) GetVersionHistory(id string) ([]map[string]interface{}, error) {
mlc.mu.RLock()
defer mlc.mu.RUnlock()
if e, ok := mlc.memories[id]; ok {
return e.VersionHistory, nil
}
return nil, nil
}
func (mlc *MemLanceClient) GetCandidatesForForgetting() ([]map[string]interface{}, error) {
mlc.mu.RLock()
defer mlc.mu.RUnlock()
var results []map[string]interface{}
for _, e := range mlc.memories {
if !e.IsDeleted && e.Tier != "core" {
results = append(results, map[string]interface{}{
"id": e.ID,
})
}
}
return results, nil
}
func (mlc *MemLanceClient) Backup(path string) error {
return nil // 内存实现无需备份
}
func (mlc *MemLanceClient) GetAuditLog(limit int) ([]map[string]interface{}, error) {
mlc.mu.RLock()
defer mlc.mu.RUnlock()
if limit > len(mlc.auditLog) {
limit = len(mlc.auditLog)
}
return mlc.auditLog[:limit], nil
}
func (mlc *MemLanceClient) IncrementUseful(id string) {
mlc.mu.Lock()
defer mlc.mu.Unlock()
if e, ok := mlc.memories[id]; ok {
e.UsefulCount++
e.RecallCount++
}
}
func (mlc *MemLanceClient) IncrementNotUseful(id string) {
mlc.mu.Lock()
defer mlc.mu.Unlock()
if e, ok := mlc.memories[id]; ok {
e.NotUsefulCount++
}
}
func (mlc *MemLanceClient) UpdateMemoryContent(id, newContent, source string) error {
mlc.mu.Lock()
defer mlc.mu.Unlock()
if e, ok := mlc.memories[id]; ok {
e.Version++
e.VersionHistory = append(e.VersionHistory, map[string]interface{}{
"version": e.Version,
"content": newContent,
"source": source,
"reason": "corrected",
"timestamp": time.Now().Format(time.RFC3339),
})
e.Content = newContent
e.Source = source
// 重新编码
if vec, err := mlc.embed.EncodeSingle(newContent); err == nil {
e.Vector = vec
}
}
return nil
}
// Update 更新记录字段
func (mlc *MemLanceClient) Update(table, id string, fields map[string]any) error {
mlc.mu.Lock()
defer mlc.mu.Unlock()
if table == "memories" {
if e, ok := mlc.memories[id]; ok {
if v, ok := fields["recall_count"]; ok {
// handle $inc
if incMap, ok := v.(map[string]string); ok {
if incStr, ok := incMap["$inc"]; ok {
if incStr == "1" {
e.RecallCount++
}
}
}
}
}
}
return nil
}
// ─── SearchCache 兼容 ──────────────────────────────────
func (mlc *MemLanceClient) SearchCompat(queryVec []float32, topK int, namespace string) ([]interface{}, error) {
records, err := mlc.Search("memories", queryVec, topK, namespace)
if err != nil {
return nil, err
}
results := make([]interface{}, len(records))
for i, r := range records {
results[i] = r
}
return results, nil
}
// ─── 工具 ─────────────────────────────────────────────
func cosineSim(a, b []float32) float64 {
if len(a) != len(b) || len(a) == 0 {
return 0
}
var dot, normA, normB float64
for i := range a {
dot += float64(a[i]) * float64(b[i])
normA += float64(a[i]) * float64(a[i])
normB += float64(b[i]) * float64(b[i])
}
if normA == 0 || normB == 0 {
return 0
}
return dot / (math.Sqrt(normA) * math.Sqrt(normB))
}