amaniquery-agent / internal /memory /memory_test.go
Deployment
Automated deployment update
4b1daed
Raw
History Blame Contribute Delete
11 kB
package memory_test
import (
"context"
"testing"
"time"
"github.com/AmaniQuery/amaniquery/internal/memory"
)
// MockEmbeddingClient for testing
type MockEmbeddingClient struct{}
func (m *MockEmbeddingClient) Generate(ctx context.Context, text string) ([]float32, error) {
// Return a simple embedding based on text length
embedding := make([]float32, 128)
for i := range embedding {
embedding[i] = float32(len(text)%10) / 10.0
}
return embedding, nil
}
func (m *MockEmbeddingClient) GenerateBatch(ctx context.Context, texts []string) ([][]float32, error) {
embeddings := make([][]float32, len(texts))
for i, text := range texts {
embedding := make([]float32, 128)
for j := range embedding {
embedding[j] = float32(len(text)%10) / 10.0
}
embeddings[i] = embedding
}
return embeddings, nil
}
// MockLLMClient for testing
type MockLLMClient struct{}
func (m *MockLLMClient) Generate(ctx context.Context, prompt string) (string, error) {
return "Mock LLM response", nil
}
func TestLocalBackend_StoreAndRetrieve(t *testing.T) {
config := memory.DefaultMemoryConfig()
backend := memory.NewLocalMemoryBackend(config)
defer backend.Close()
ctx := context.Background()
// Store an entry
entry := &memory.MemoryEntry{
ID: "test-1",
Type: memory.EpisodicMemory,
Content: "This is a test memory entry",
Timestamp: time.Now(),
Confidence: 0.9,
UserID: "user-1",
SessionID: "session-1",
Source: "test",
}
err := backend.Store(ctx, entry)
if err != nil {
t.Fatalf("Store failed: %v", err)
}
// Retrieve the entry
query := &memory.MemoryQuery{
UserID: "user-1",
TopK: 10,
}
results, err := backend.Retrieve(ctx, query)
if err != nil {
t.Fatalf("Retrieve failed: %v", err)
}
if len(results) != 1 {
t.Fatalf("Expected 1 result, got %d", len(results))
}
if results[0].ID != "test-1" {
t.Errorf("Expected ID 'test-1', got '%s'", results[0].ID)
}
}
func TestLocalBackend_DeleteUserData(t *testing.T) {
config := memory.DefaultMemoryConfig()
backend := memory.NewLocalMemoryBackend(config)
defer backend.Close()
ctx := context.Background()
// Store multiple entries for a user
for i := 0; i < 5; i++ {
entry := &memory.MemoryEntry{
ID: "test-" + string(rune('0'+i)),
Type: memory.EpisodicMemory,
Content: "Test content",
Timestamp: time.Now(),
UserID: "user-gdpr",
SessionID: "session-1",
}
backend.Store(ctx, entry)
}
// Verify entries exist
query := &memory.MemoryQuery{UserID: "user-gdpr", TopK: 10}
results, _ := backend.Retrieve(ctx, query)
if len(results) != 5 {
t.Fatalf("Expected 5 entries before delete, got %d", len(results))
}
// Delete user data
err := backend.DeleteUserData(ctx, "user-gdpr")
if err != nil {
t.Fatalf("DeleteUserData failed: %v", err)
}
// Verify entries are gone
results, _ = backend.Retrieve(ctx, query)
if len(results) != 0 {
t.Fatalf("Expected 0 entries after delete, got %d", len(results))
}
}
func TestWorkingMemory_AddAndPrune(t *testing.T) {
wm := memory.NewWorkingMemory("session-1", "user-1", 1024) // 1KB limit
// Add entries until we exceed the limit
for i := 0; i < 10; i++ {
entry := &memory.MemoryEntry{
ID: "test-" + string(rune('0'+i)),
Content: "This is a test entry with some content to take up space",
}
wm.Add(entry)
}
// Check that pruning occurred
stats := wm.Stats()
if stats.SizeBytes > 1024 {
t.Errorf("Expected size <= 1024, got %d", stats.SizeBytes)
}
}
func TestWorkingMemory_GetRecent(t *testing.T) {
wm := memory.NewWorkingMemory("session-1", "user-1", 10*1024)
// Add entries
for i := 0; i < 5; i++ {
entry := &memory.MemoryEntry{
ID: "test-" + string(rune('0'+i)),
Content: "Entry content",
}
wm.Add(entry)
}
// Get recent 3
recent := wm.GetRecent(3)
if len(recent) != 3 {
t.Fatalf("Expected 3 recent entries, got %d", len(recent))
}
// Should be the last 3 added
if recent[0].ID != "test-2" {
t.Errorf("Expected first recent to be test-2, got %s", recent[0].ID)
}
}
func TestTemporalContext_RecencyScoring(t *testing.T) {
tc := memory.NewTemporalContext("user-1", "session-1")
// Create entries with different timestamps
now := time.Now()
recentEntry := &memory.MemoryEntry{
ID: "recent",
Timestamp: now,
Confidence: 1.0,
}
oldEntry := &memory.MemoryEntry{
ID: "old",
Timestamp: now.Add(-24 * time.Hour),
Confidence: 1.0,
}
recentScore := tc.CalculateRecencyScore(recentEntry)
oldScore := tc.CalculateRecencyScore(oldEntry)
if recentScore <= oldScore {
t.Errorf("Recent entry should have higher score: recent=%f, old=%f", recentScore, oldScore)
}
}
func TestContextWindowBuilder_Build(t *testing.T) {
entries := []*memory.MemoryEntry{
{
ID: "1",
Type: memory.SemanticMemory,
Content: "First entry",
Timestamp: time.Now(),
Source: "test",
},
{
ID: "2",
Type: memory.EpisodicMemory,
Content: "Second entry",
Timestamp: time.Now(),
Source: "test",
},
}
builder := memory.NewContextWindowBuilder().
WithMaxTokens(1000).
WithFormatStyle(memory.FormatMarkdown)
window := builder.Build(entries)
if window.TotalTokens == 0 {
t.Error("Expected non-zero token count")
}
if len(window.Entries) != 2 {
t.Errorf("Expected 2 entries in window, got %d", len(window.Entries))
}
if window.FormattedContext == "" {
t.Error("Expected non-empty formatted context")
}
}
func TestMemoryOrchestrator_StoreAndQuery(t *testing.T) {
config := memory.DefaultMemoryConfig()
backend := memory.NewLocalMemoryBackend(config)
embedder := &MockEmbeddingClient{}
llmClient := &MockLLMClient{}
orchestrator := memory.NewMemoryOrchestrator(backend, embedder, llmClient, config)
defer orchestrator.Close()
ctx := context.Background()
// Store an entry
entry := &memory.MemoryEntry{
Type: memory.SemanticMemory,
Content: "The capital of France is Paris",
UserID: "user-1",
SessionID: "session-1",
}
err := orchestrator.Store(ctx, entry)
if err != nil {
t.Fatalf("Store failed: %v", err)
}
// Query
memCtx, err := orchestrator.ProcessQuery(ctx, &memory.QueryRequest{
Query: "What is the capital of France?",
UserID: "user-1",
SessionID: "session-1",
MaxTokens: 1000,
})
if err != nil {
t.Fatalf("ProcessQuery failed: %v", err)
}
if memCtx == nil {
t.Fatal("Expected non-nil memory context")
}
}
func TestAgentMemoryIntegration(t *testing.T) {
config := memory.DefaultMemoryConfig()
embedder := &MockEmbeddingClient{}
llmClient := &MockLLMClient{}
integration, err := memory.NewAgentMemoryIntegration(config, embedder, llmClient)
if err != nil {
t.Fatalf("Failed to create integration: %v", err)
}
defer integration.Stop()
ctx := context.Background()
// Store a conversation turn
err = integration.StoreConversationTurn(ctx, "user-1", "session-1",
"What is the weather?",
"I don't have access to real-time weather data.")
if err != nil {
t.Fatalf("StoreConversationTurn failed: %v", err)
}
// Get context for a follow-up query
memCtx, err := integration.GetContextForQuery(ctx, "user-1", "session-1",
"Tell me more about the weather",
1000)
if err != nil {
t.Fatalf("GetContextForQuery failed: %v", err)
}
if memCtx == nil {
t.Fatal("Expected non-nil memory context")
}
// Check metrics
metrics := integration.GetMetrics()
if metrics.TotalStores == 0 {
t.Error("Expected some stores to be recorded")
}
}
func TestMemoryWorker(t *testing.T) {
config := memory.DefaultMemoryConfig()
embedder := &MockEmbeddingClient{}
llmClient := &MockLLMClient{}
integration, err := memory.NewAgentMemoryIntegration(config, embedder, llmClient)
if err != nil {
t.Fatalf("Failed to create integration: %v", err)
}
worker := memory.NewMemoryWorker(integration, 2)
worker.Start()
defer worker.Stop()
ctx := context.Background()
// Submit async work
entry := &memory.MemoryEntry{
Type: memory.EpisodicMemory,
Content: "Async stored entry",
UserID: "user-1",
SessionID: "session-1",
}
result, err := worker.SubmitWorkWithResult(ctx, memory.MemoryWorkItem{
Type: memory.WorkStoreEntry,
Data: entry,
})
if err != nil {
t.Fatalf("SubmitWorkWithResult failed: %v", err)
}
if !result.Success {
t.Errorf("Work item failed: %v", result.Error)
}
}
func TestGDPRManager_ExportUserData(t *testing.T) {
config := memory.DefaultMemoryConfig()
backend := memory.NewLocalMemoryBackend(config)
gdprManager := memory.NewGDPRManager(backend, nil)
ctx := context.Background()
// Store some data
for i := 0; i < 3; i++ {
entry := &memory.MemoryEntry{
ID: "export-test-" + string(rune('0'+i)),
Type: memory.SemanticMemory,
Content: "Test content for export",
Timestamp: time.Now(),
UserID: "export-user",
SessionID: "session-1",
}
backend.Store(ctx, entry)
}
// Export user data
export, err := gdprManager.ExportUserData(ctx, "export-user", "admin")
if err != nil {
t.Fatalf("ExportUserData failed: %v", err)
}
if export == nil {
t.Fatal("Expected non-nil export")
}
if export.EntryCount != 3 {
t.Errorf("Expected 3 entries in export, got %d", export.EntryCount)
}
}
// Benchmark tests
func BenchmarkLocalBackend_Store(b *testing.B) {
config := memory.DefaultMemoryConfig()
backend := memory.NewLocalMemoryBackend(config)
defer backend.Close()
ctx := context.Background()
b.ResetTimer()
for i := 0; i < b.N; i++ {
entry := &memory.MemoryEntry{
ID: "bench-" + string(rune(i%256)),
Type: memory.EpisodicMemory,
Content: "Benchmark test content",
Timestamp: time.Now(),
UserID: "bench-user",
SessionID: "bench-session",
}
backend.Store(ctx, entry)
}
}
func BenchmarkLocalBackend_Retrieve(b *testing.B) {
config := memory.DefaultMemoryConfig()
backend := memory.NewLocalMemoryBackend(config)
defer backend.Close()
ctx := context.Background()
// Pre-populate
for i := 0; i < 1000; i++ {
entry := &memory.MemoryEntry{
ID: "bench-" + string(rune(i%256)) + string(rune(i/256)),
Type: memory.EpisodicMemory,
Content: "Benchmark test content for retrieval",
Timestamp: time.Now(),
UserID: "bench-user",
SessionID: "bench-session",
}
backend.Store(ctx, entry)
}
query := &memory.MemoryQuery{
UserID: "bench-user",
TopK: 10,
}
b.ResetTimer()
for i := 0; i < b.N; i++ {
backend.Retrieve(ctx, query)
}
}
func BenchmarkContextWindowBuilder(b *testing.B) {
entries := make([]*memory.MemoryEntry, 100)
for i := range entries {
entries[i] = &memory.MemoryEntry{
ID: "bench-" + string(rune(i%256)),
Type: memory.SemanticMemory,
Content: "This is benchmark content for context window building",
Timestamp: time.Now(),
Source: "benchmark",
}
}
builder := memory.NewContextWindowBuilder().WithMaxTokens(4000)
b.ResetTimer()
for i := 0; i < b.N; i++ {
builder.Build(entries)
}
}