129 lines
3.1 KiB
Go
129 lines
3.1 KiB
Go
package rag
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"time"
|
|
|
|
"github.com/google/uuid"
|
|
)
|
|
|
|
// Fragment represents a piece of content stored in the RAG system.
|
|
type Fragment struct {
|
|
ID string
|
|
Content string
|
|
Vector []float32
|
|
Metadata map[string]string
|
|
Timestamp time.Time
|
|
ProjectID string
|
|
}
|
|
|
|
// Memory provides persistent memory and semantic search over agent content.
|
|
type Memory interface {
|
|
Add(ctx context.Context, fragment Fragment) error
|
|
Search(ctx context.Context, query string, topK int) ([]Fragment, error)
|
|
Forget(ctx context.Context, id string) error
|
|
ForgetAll(ctx context.Context) error
|
|
}
|
|
|
|
// Config holds the settings for creating a Memory.
|
|
type Config struct {
|
|
Backend Backend
|
|
Embedder Embedder
|
|
}
|
|
|
|
// Backend is the interface for vector database backends.
|
|
type Backend interface {
|
|
Upsert(ctx context.Context, id string, vector []float32, metadata map[string]string) error
|
|
Search(ctx context.Context, queryVector []float32, topK int) ([]SearchResult, error)
|
|
Forget(ctx context.Context, id string) error
|
|
ForgetAll(ctx context.Context) error
|
|
}
|
|
|
|
// SearchResult represents a matched fragment from a search.
|
|
type SearchResult struct {
|
|
ID string
|
|
Content string
|
|
Score float32
|
|
Metadata map[string]string
|
|
}
|
|
|
|
// Embedder generates embeddings for text.
|
|
type Embedder interface {
|
|
Embed(ctx context.Context, text string) ([]float32, error)
|
|
Dimensions() int
|
|
}
|
|
|
|
// memory implements Memory using a Backend and Embedder.
|
|
type memory struct {
|
|
backend Backend
|
|
embedder Embedder
|
|
}
|
|
|
|
// New creates a new Memory with the given config.
|
|
func New(cfg Config) (Memory, error) {
|
|
if cfg.Backend == nil {
|
|
return nil, fmt.Errorf("backend is required")
|
|
}
|
|
if cfg.Embedder == nil {
|
|
return nil, fmt.Errorf("embedder is required")
|
|
}
|
|
return &memory{
|
|
backend: cfg.Backend,
|
|
embedder: cfg.Embedder,
|
|
}, nil
|
|
}
|
|
|
|
func (m *memory) Add(ctx context.Context, fragment Fragment) error {
|
|
if fragment.ID == "" {
|
|
fragment.ID = uuid.New().String()
|
|
}
|
|
if fragment.Metadata == nil {
|
|
fragment.Metadata = make(map[string]string)
|
|
}
|
|
fragment.Metadata["project_id"] = fragment.ProjectID
|
|
fragment.Timestamp = time.Now()
|
|
|
|
vector, err := m.embedder.Embed(ctx, fragment.Content)
|
|
if err != nil {
|
|
return fmt.Errorf("embedding: %w", err)
|
|
}
|
|
fragment.Vector = vector
|
|
|
|
return m.backend.Upsert(ctx, fragment.ID, fragment.Vector, fragment.Metadata)
|
|
}
|
|
|
|
func (m *memory) Search(ctx context.Context, query string, topK int) ([]Fragment, error) {
|
|
if topK <= 0 {
|
|
topK = 5
|
|
}
|
|
|
|
queryVector, err := m.embedder.Embed(ctx, query)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("embedding query: %w", err)
|
|
}
|
|
|
|
results, err := m.backend.Search(ctx, queryVector, topK)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("search: %w", err)
|
|
}
|
|
|
|
fragments := make([]Fragment, len(results))
|
|
for i, r := range results {
|
|
fragments[i] = Fragment{
|
|
ID: r.ID,
|
|
Content: r.Content,
|
|
Metadata: r.Metadata,
|
|
ProjectID: r.Metadata["project_id"],
|
|
}
|
|
}
|
|
return fragments, nil
|
|
}
|
|
|
|
func (m *memory) Forget(ctx context.Context, id string) error {
|
|
return m.backend.Forget(ctx, id)
|
|
}
|
|
|
|
func (m *memory) ForgetAll(ctx context.Context) error {
|
|
return m.backend.ForgetAll(ctx)
|
|
}
|