308 lines
7.3 KiB
Go
308 lines
7.3 KiB
Go
// 织忆 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))
|
||
}
|