Fragments now carry a memory type in metadata (legacy fragments count as procedural) with SearchByType filtering, and EpisodeCapture summarizes a finished turn with the local LLM and stores it as episodic memory, so the agent can answer "what did we do yesterday?". Includes an E2E test against a live llama.cpp server (gated) and taxonomy unit tests. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
234 lines
7.7 KiB
Go
234 lines
7.7 KiB
Go
package rag_test
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"iter"
|
|
"strings"
|
|
"testing"
|
|
|
|
"github.com/VictorVargas/rony-llm-agent/pkg/llm"
|
|
"github.com/VictorVargas/rony-llm-agent/pkg/rag"
|
|
"github.com/VictorVargas/rony-llm-agent/pkg/rag/embeddings"
|
|
)
|
|
|
|
func TestMemory_Add_StoresMemoryTypeInMetadata(t *testing.T) {
|
|
var gotMeta map[string]string
|
|
m, err := rag.New(rag.Config{
|
|
Backend: &mockBackend{
|
|
upsertFunc: func(_ context.Context, _ string, _ []float32, _ string, metadata map[string]string) error {
|
|
gotMeta = metadata
|
|
return nil
|
|
},
|
|
},
|
|
Embedder: &embeddings.MockEmbedder{},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
|
|
if err := m.Add(context.Background(), rag.Fragment{Content: "an event", Type: rag.MemoryEpisodic}); err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if gotMeta["memory_type"] != "episodic" {
|
|
t.Fatalf("expected memory_type=episodic in metadata, got %q", gotMeta["memory_type"])
|
|
}
|
|
}
|
|
|
|
func TestMemory_Add_DefaultsToProcedural(t *testing.T) {
|
|
var gotMeta map[string]string
|
|
m, err := rag.New(rag.Config{
|
|
Backend: &mockBackend{
|
|
upsertFunc: func(_ context.Context, _ string, _ []float32, _ string, metadata map[string]string) error {
|
|
gotMeta = metadata
|
|
return nil
|
|
},
|
|
},
|
|
Embedder: &embeddings.MockEmbedder{},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
|
|
if err := m.Add(context.Background(), rag.Fragment{Content: "a process"}); err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if gotMeta["memory_type"] != "procedural" {
|
|
t.Fatalf("expected untyped fragments to default to procedural, got %q", gotMeta["memory_type"])
|
|
}
|
|
}
|
|
|
|
func TestMemory_SearchByType_FiltersAndTreatsLegacyAsProcedural(t *testing.T) {
|
|
backendResults := []rag.SearchResult{
|
|
{ID: "1", Content: "episode", Metadata: map[string]string{"memory_type": "episodic"}},
|
|
{ID: "2", Content: "fact", Metadata: map[string]string{"memory_type": "semantic"}},
|
|
{ID: "3", Content: "legacy process", Metadata: map[string]string{}}, // pre-taxonomy fragment
|
|
{ID: "4", Content: "typed process", Metadata: map[string]string{"memory_type": "procedural"}},
|
|
}
|
|
m, err := rag.New(rag.Config{
|
|
Backend: &mockBackend{
|
|
searchFunc: func(_ context.Context, _ string, _ []float32, _ int) ([]rag.SearchResult, error) {
|
|
return backendResults, nil
|
|
},
|
|
},
|
|
Embedder: &embeddings.MockEmbedder{},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
|
|
got, err := m.SearchByType(context.Background(), "q", 10, rag.MemoryProcedural)
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if len(got) != 2 || got[0].ID != "3" || got[1].ID != "4" {
|
|
t.Fatalf("expected legacy + typed procedural fragments, got %+v", got)
|
|
}
|
|
if got[0].Type != rag.MemoryProcedural {
|
|
t.Fatalf("expected legacy fragment to surface as procedural, got %q", got[0].Type)
|
|
}
|
|
|
|
episodes, err := m.SearchByType(context.Background(), "q", 10, rag.MemoryEpisodic)
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if len(episodes) != 1 || episodes[0].ID != "1" {
|
|
t.Fatalf("expected only the episodic fragment, got %+v", episodes)
|
|
}
|
|
|
|
all, err := m.SearchByType(context.Background(), "q", 10)
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if len(all) != 4 {
|
|
t.Fatalf("expected no type restriction to return everything, got %d", len(all))
|
|
}
|
|
}
|
|
|
|
func TestMemory_SearchByType_OverfetchesWhenFiltering(t *testing.T) {
|
|
var gotTopK int
|
|
m, err := rag.New(rag.Config{
|
|
Backend: &mockBackend{
|
|
searchFunc: func(_ context.Context, _ string, _ []float32, topK int) ([]rag.SearchResult, error) {
|
|
gotTopK = topK
|
|
return nil, nil
|
|
},
|
|
},
|
|
Embedder: &embeddings.MockEmbedder{},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
|
|
if _, err := m.SearchByType(context.Background(), "q", 5, rag.MemoryEpisodic); err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if gotTopK <= 5 {
|
|
t.Fatalf("expected the backend to be asked for more than topK candidates when filtering, got %d", gotTopK)
|
|
}
|
|
|
|
if _, err := m.Search(context.Background(), "q", 5); err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if gotTopK != 5 {
|
|
t.Fatalf("expected unfiltered search to request exactly topK, got %d", gotTopK)
|
|
}
|
|
}
|
|
|
|
// captureLLM is a minimal llm.LLMClient stub for capture tests.
|
|
type captureLLM struct {
|
|
response string
|
|
err error
|
|
gotUser string
|
|
}
|
|
|
|
func (c *captureLLM) Generate(_ context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) {
|
|
for _, m := range req.Messages {
|
|
if m.Role == llm.RoleUser {
|
|
c.gotUser = m.Content
|
|
}
|
|
}
|
|
if c.err != nil {
|
|
return llm.CompletionResponse{}, c.err
|
|
}
|
|
return llm.CompletionResponse{Content: c.response}, nil
|
|
}
|
|
|
|
func (c *captureLLM) Stream(_ context.Context, _ llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] {
|
|
return func(func(llm.StreamChunk, error) bool) {}
|
|
}
|
|
func (c *captureLLM) Name() string { return "capture-stub" }
|
|
func (c *captureLLM) Capabilities() llm.ProviderCapabilities { return llm.ProviderCapabilities{} }
|
|
|
|
func TestEpisodeCapture_SavesEpisodicSummary(t *testing.T) {
|
|
var saved rag.Fragment
|
|
backend := &mockBackend{
|
|
upsertFunc: func(_ context.Context, id string, _ []float32, content string, metadata map[string]string) error {
|
|
saved = rag.Fragment{ID: id, Content: content, Metadata: metadata}
|
|
return nil
|
|
},
|
|
}
|
|
mem, err := rag.New(rag.Config{Backend: backend, Embedder: &embeddings.MockEmbedder{}})
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
|
|
stub := &captureLLM{response: " The user asked how to deploy and the assistant explained the release steps. "}
|
|
cap := &rag.EpisodeCapture{Memory: mem, LLM: stub, ProjectID: "proj1"}
|
|
|
|
err = cap.Capture(context.Background(), "how do I deploy?", "You run make release...", "read", "bash")
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
|
|
if saved.Content != "The user asked how to deploy and the assistant explained the release steps." {
|
|
t.Fatalf("expected trimmed summary as content, got %q", saved.Content)
|
|
}
|
|
if saved.Metadata["memory_type"] != "episodic" {
|
|
t.Fatalf("expected episodic type, got %q", saved.Metadata["memory_type"])
|
|
}
|
|
if saved.Metadata["tools"] != "read,bash" {
|
|
t.Fatalf("expected tools metadata, got %q", saved.Metadata["tools"])
|
|
}
|
|
if saved.Metadata["project_id"] != "proj1" {
|
|
t.Fatalf("expected project_id metadata, got %q", saved.Metadata["project_id"])
|
|
}
|
|
if !strings.Contains(stub.gotUser, "how do I deploy?") {
|
|
t.Fatalf("expected the turn transcript to reach the LLM, got %q", stub.gotUser)
|
|
}
|
|
}
|
|
|
|
func TestEpisodeCapture_SkipsEmptyTurnsAndFailures(t *testing.T) {
|
|
upserts := 0
|
|
backend := &mockBackend{
|
|
upsertFunc: func(_ context.Context, _ string, _ []float32, _ string, _ map[string]string) error {
|
|
upserts++
|
|
return nil
|
|
},
|
|
}
|
|
mem, err := rag.New(rag.Config{Backend: backend, Embedder: &embeddings.MockEmbedder{}})
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
|
|
cap := &rag.EpisodeCapture{Memory: mem, LLM: &captureLLM{response: "summary"}, ProjectID: "p"}
|
|
if err := cap.Capture(context.Background(), "", "reply"); err == nil {
|
|
t.Fatal("expected error for empty user input")
|
|
}
|
|
if err := cap.Capture(context.Background(), "input", " "); err == nil {
|
|
t.Fatal("expected error for empty assistant reply")
|
|
}
|
|
|
|
failing := &rag.EpisodeCapture{Memory: mem, LLM: &captureLLM{err: fmt.Errorf("llm down")}, ProjectID: "p"}
|
|
if err := failing.Capture(context.Background(), "input", "reply"); err == nil {
|
|
t.Fatal("expected error when the LLM fails")
|
|
}
|
|
empty := &rag.EpisodeCapture{Memory: mem, LLM: &captureLLM{response: " "}, ProjectID: "p"}
|
|
if err := empty.Capture(context.Background(), "input", "reply"); err == nil {
|
|
t.Fatal("expected error for an empty summary")
|
|
}
|
|
|
|
if upserts != 0 {
|
|
t.Fatalf("expected nothing to be saved on failures, got %d upserts", upserts)
|
|
}
|
|
}
|