feat: add vector search for memory and lifecycle hooks
Replace unbounded MEMORY.md dump with opt-in vector search that injects
only the top-K relevant chunks into each prompt. When disabled (default),
behavior is unchanged — zero overhead.
New packages:
- pkg/vecstore: in-memory vector store with gob persistence, brute-force
cosine search, markdown chunker (split by ## headers + paragraphs),
and OpenAI-compatible HTTP embedder with retry/backoff
- pkg/agent/hooks.go: lifecycle hook struct (OnContextBuild, OnPreTool,
OnPostTool, OnPreLLM, OnPostMessage) — nil checks only, no interface
- pkg/tools/memory_search.go: memory_search tool for explicit semantic
search over indexed memory
Modified:
- pkg/config: add memory.vector_search config section with sensible
defaults (enabled: false, model: text-embedding-3-small, top-5)
- pkg/agent/loop.go: integrate hooks at context build, tool execution,
and post-message points
- pkg/agent/context.go: accept optional enriched context in BuildMessages
- cmd/picoclaw/main.go: wire vector search setup into agent and gateway
commands with automatic initial indexing and incremental re-indexing
Config example:
{ "memory": { "vector_search": { "enabled": true } } }
API key and base URL fall back to providers.openai if not specified.
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
d83fb6e081
commit
f984a11291
12 changed files with 1148 additions and 2 deletions
|
|
@ -31,6 +31,7 @@ import (
|
|||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
"github.com/sipeed/picoclaw/pkg/skills"
|
||||
"github.com/sipeed/picoclaw/pkg/tools"
|
||||
"github.com/sipeed/picoclaw/pkg/vecstore"
|
||||
"github.com/sipeed/picoclaw/pkg/voice"
|
||||
)
|
||||
|
||||
|
|
@ -496,6 +497,10 @@ func agentCmd() {
|
|||
msgBus := bus.NewMessageBus()
|
||||
agentLoop := agent.NewAgentLoop(cfg, msgBus, provider)
|
||||
|
||||
// Setup vector search hooks (no-op when disabled)
|
||||
hooks := setupVectorSearch(cfg, agentLoop)
|
||||
agentLoop.SetHooks(hooks)
|
||||
|
||||
// Print agent startup info (only for interactive mode)
|
||||
startupInfo := agentLoop.GetStartupInfo()
|
||||
logger.InfoCF("agent", "Agent initialized",
|
||||
|
|
@ -631,6 +636,10 @@ func gatewayCmd() {
|
|||
msgBus := bus.NewMessageBus()
|
||||
agentLoop := agent.NewAgentLoop(cfg, msgBus, provider)
|
||||
|
||||
// Setup vector search hooks (no-op when disabled)
|
||||
hooks := setupVectorSearch(cfg, agentLoop)
|
||||
agentLoop.SetHooks(hooks)
|
||||
|
||||
// Print agent startup info
|
||||
fmt.Println("\n📦 Agent Status:")
|
||||
startupInfo := agentLoop.GetStartupInfo()
|
||||
|
|
@ -1050,6 +1059,193 @@ func loadConfig() (*config.Config, error) {
|
|||
return config.LoadConfig(getConfigPath())
|
||||
}
|
||||
|
||||
// setupVectorSearch initializes vector search hooks and tools when enabled.
|
||||
// Returns the Hooks struct (always non-nil) and a save function to call on shutdown.
|
||||
func setupVectorSearch(cfg *config.Config, agentLoop *agent.AgentLoop) *agent.Hooks {
|
||||
hooks := &agent.Hooks{}
|
||||
|
||||
vsCfg := cfg.Memory.VectorSearch
|
||||
if !vsCfg.Enabled {
|
||||
return hooks
|
||||
}
|
||||
|
||||
// Resolve API key and base URL, falling back to OpenAI provider config
|
||||
apiKey := vsCfg.APIKey
|
||||
if apiKey == "" {
|
||||
apiKey = cfg.Providers.OpenAI.APIKey
|
||||
}
|
||||
apiBase := vsCfg.APIBase
|
||||
if apiBase == "" {
|
||||
apiBase = cfg.Providers.OpenAI.APIBase
|
||||
}
|
||||
if apiBase == "" {
|
||||
apiBase = "https://api.openai.com/v1"
|
||||
}
|
||||
|
||||
model := vsCfg.Model
|
||||
if model == "" {
|
||||
model = "text-embedding-3-small"
|
||||
}
|
||||
|
||||
maxResults := vsCfg.MaxResults
|
||||
if maxResults <= 0 {
|
||||
maxResults = 5
|
||||
}
|
||||
|
||||
chunkSize := vsCfg.ChunkSize
|
||||
if chunkSize <= 0 {
|
||||
chunkSize = 800
|
||||
}
|
||||
|
||||
workspace := agentLoop.Workspace()
|
||||
storePath := filepath.Join(workspace, "memory", ".vecstore.gob")
|
||||
|
||||
embedder := vecstore.NewHTTPEmbedder(apiBase, apiKey, model)
|
||||
store := vecstore.NewVectorStore(storePath)
|
||||
|
||||
// Load existing store
|
||||
if err := store.Load(); err != nil {
|
||||
logger.ErrorCF("vecstore", "Failed to load vector store",
|
||||
map[string]interface{}{"error": err.Error()})
|
||||
}
|
||||
|
||||
// Initial indexing if store is empty
|
||||
if store.Len() == 0 {
|
||||
go indexMemoryFiles(context.Background(), workspace, embedder, store, chunkSize)
|
||||
}
|
||||
|
||||
// OnContextBuild: embed query and return top-K relevant chunks
|
||||
hooks.OnContextBuild = func(ctx context.Context, query string) (string, error) {
|
||||
embeddings, err := embedder.Embed(ctx, []string{query})
|
||||
if err != nil || len(embeddings) == 0 || len(embeddings[0]) == 0 {
|
||||
return "", err
|
||||
}
|
||||
|
||||
results := store.Search(embeddings[0], maxResults)
|
||||
if len(results) == 0 {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
var sb strings.Builder
|
||||
for _, r := range results {
|
||||
snippet := r.Text
|
||||
if len(snippet) > 700 {
|
||||
snippet = snippet[:700] + "..."
|
||||
}
|
||||
sb.WriteString(fmt.Sprintf("[%s | score: %.2f]\n%s\n\n", r.Source, r.Score, snippet))
|
||||
}
|
||||
return sb.String(), nil
|
||||
}
|
||||
|
||||
// Track which files were touched by tool calls in this message
|
||||
var touchedFiles []string
|
||||
hooks.OnPostTool = func(_ context.Context, name string, result string, _ time.Duration) {
|
||||
if name == "write_file" || name == "append_file" || name == "edit_file" {
|
||||
// Check if the result mentions a memory/ path
|
||||
if strings.Contains(result, "memory/") || strings.Contains(result, "memory\\") {
|
||||
touchedFiles = append(touchedFiles, result)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// OnPostMessage: re-index memory files that were modified
|
||||
hooks.OnPostMessage = func(ctx context.Context, _, _, _ string) {
|
||||
if len(touchedFiles) > 0 {
|
||||
go indexMemoryFiles(ctx, workspace, embedder, store, chunkSize)
|
||||
touchedFiles = nil
|
||||
}
|
||||
}
|
||||
|
||||
// Register memory_search tool
|
||||
agentLoop.RegisterTool(tools.NewMemorySearchTool(embedder, store, maxResults))
|
||||
|
||||
logger.InfoCF("vecstore", "Vector search enabled",
|
||||
map[string]interface{}{
|
||||
"model": model,
|
||||
"max_results": maxResults,
|
||||
"chunk_size": chunkSize,
|
||||
"store_path": storePath,
|
||||
})
|
||||
|
||||
return hooks
|
||||
}
|
||||
|
||||
// indexMemoryFiles chunks and embeds all markdown files in the memory directory.
|
||||
func indexMemoryFiles(ctx context.Context, workspace string, embedder vecstore.Embedder, store *vecstore.VectorStore, chunkSize int) {
|
||||
memoryDir := filepath.Join(workspace, "memory")
|
||||
|
||||
var allChunks []vecstore.Chunk
|
||||
filepath.Walk(memoryDir, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil || info.IsDir() {
|
||||
return nil
|
||||
}
|
||||
// Skip non-markdown and the store file itself
|
||||
if !strings.HasSuffix(path, ".md") {
|
||||
return nil
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
relPath, _ := filepath.Rel(workspace, path)
|
||||
chunks := vecstore.ChunkMarkdown(relPath, string(data), chunkSize)
|
||||
allChunks = append(allChunks, chunks...)
|
||||
return nil
|
||||
})
|
||||
|
||||
if len(allChunks) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
// Collect texts that need embedding (skip chunks already in store with same ID)
|
||||
texts := make([]string, len(allChunks))
|
||||
for i, c := range allChunks {
|
||||
texts[i] = c.Text
|
||||
}
|
||||
|
||||
// Embed in batches of 100
|
||||
const batchSize = 100
|
||||
for i := 0; i < len(texts); i += batchSize {
|
||||
end := i + batchSize
|
||||
if end > len(texts) {
|
||||
end = len(texts)
|
||||
}
|
||||
|
||||
embeddings, err := embedder.Embed(ctx, texts[i:end])
|
||||
if err != nil {
|
||||
logger.ErrorCF("vecstore", "Embedding batch failed",
|
||||
map[string]interface{}{"error": err.Error(), "batch": i / batchSize})
|
||||
continue
|
||||
}
|
||||
|
||||
for j, emb := range embeddings {
|
||||
allChunks[i+j].Embedding = emb
|
||||
}
|
||||
}
|
||||
|
||||
// Filter out chunks that failed to embed
|
||||
var valid []vecstore.Chunk
|
||||
for _, c := range allChunks {
|
||||
if len(c.Embedding) > 0 {
|
||||
valid = append(valid, c)
|
||||
}
|
||||
}
|
||||
|
||||
store.Upsert(valid)
|
||||
if err := store.Save(); err != nil {
|
||||
logger.ErrorCF("vecstore", "Failed to save vector store",
|
||||
map[string]interface{}{"error": err.Error()})
|
||||
}
|
||||
|
||||
logger.InfoCF("vecstore", "Memory indexed",
|
||||
map[string]interface{}{
|
||||
"chunks_total": len(allChunks),
|
||||
"chunks_embedded": len(valid),
|
||||
})
|
||||
}
|
||||
|
||||
func cronCmd() {
|
||||
if len(os.Args) < 3 {
|
||||
cronHelp()
|
||||
|
|
|
|||
|
|
@ -157,11 +157,16 @@ func (cb *ContextBuilder) LoadBootstrapFiles() string {
|
|||
return result
|
||||
}
|
||||
|
||||
func (cb *ContextBuilder) BuildMessages(history []providers.Message, summary string, currentMessage string, media []string, channel, chatID string) []providers.Message {
|
||||
func (cb *ContextBuilder) BuildMessages(history []providers.Message, summary string, currentMessage string, media []string, channel, chatID string, enrichedContext ...string) []providers.Message {
|
||||
messages := []providers.Message{}
|
||||
|
||||
systemPrompt := cb.BuildSystemPrompt()
|
||||
|
||||
// Inject enriched context (e.g. vector search results) before session info
|
||||
if len(enrichedContext) > 0 && enrichedContext[0] != "" {
|
||||
systemPrompt += "\n\n## Relevant Memory Context\n\n" + enrichedContext[0]
|
||||
}
|
||||
|
||||
// Add Current Session info if provided
|
||||
if channel != "" && chatID != "" {
|
||||
systemPrompt += fmt.Sprintf("\n\n## Current Session\nChannel: %s\nChat ID: %s", channel, chatID)
|
||||
|
|
|
|||
28
pkg/agent/hooks.go
Normal file
28
pkg/agent/hooks.go
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/providers"
|
||||
)
|
||||
|
||||
// Hooks provides lifecycle integration points for the agent loop.
|
||||
// All fields are optional — nil checks only, zero cost when unused.
|
||||
type Hooks struct {
|
||||
// OnContextBuild is called before building messages to inject extra context.
|
||||
// Returns additional context string to include in the system prompt.
|
||||
OnContextBuild func(ctx context.Context, query string) (string, error)
|
||||
|
||||
// OnPreTool is called before each tool execution.
|
||||
OnPreTool func(ctx context.Context, name string, args map[string]interface{}) error
|
||||
|
||||
// OnPostTool is called after each tool execution with the result and duration.
|
||||
OnPostTool func(ctx context.Context, name string, result string, dur time.Duration)
|
||||
|
||||
// OnPreLLM is called before each LLM call, allowing message mutation.
|
||||
OnPreLLM func(ctx context.Context, messages []providers.Message) []providers.Message
|
||||
|
||||
// OnPostMessage is called after a complete message exchange is saved.
|
||||
OnPostMessage func(ctx context.Context, sessionKey, userMsg, response string)
|
||||
}
|
||||
|
|
@ -36,6 +36,7 @@ type AgentLoop struct {
|
|||
sessions *session.SessionManager
|
||||
contextBuilder *ContextBuilder
|
||||
tools *tools.ToolRegistry
|
||||
hooks *Hooks
|
||||
running atomic.Bool
|
||||
summarizing sync.Map // Tracks which sessions are currently being summarized
|
||||
}
|
||||
|
|
@ -148,6 +149,16 @@ func (al *AgentLoop) RegisterTool(tool tools.Tool) {
|
|||
al.tools.Register(tool)
|
||||
}
|
||||
|
||||
// SetHooks sets lifecycle hooks on the agent loop.
|
||||
func (al *AgentLoop) SetHooks(h *Hooks) {
|
||||
al.hooks = h
|
||||
}
|
||||
|
||||
// Workspace returns the agent's workspace path (for wiring hooks).
|
||||
func (al *AgentLoop) Workspace() string {
|
||||
return al.workspace
|
||||
}
|
||||
|
||||
func (al *AgentLoop) ProcessDirect(ctx context.Context, content, sessionKey string) (string, error) {
|
||||
return al.ProcessDirectWithChannel(ctx, content, sessionKey, "cli", "direct")
|
||||
}
|
||||
|
|
@ -236,7 +247,14 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, opts processOptions) (str
|
|||
// 1. Update tool contexts
|
||||
al.updateToolContexts(opts.Channel, opts.ChatID)
|
||||
|
||||
// 2. Build messages
|
||||
// 2. Build messages (with optional enriched context from hooks)
|
||||
var enrichedCtx string
|
||||
if al.hooks != nil && al.hooks.OnContextBuild != nil {
|
||||
if extra, err := al.hooks.OnContextBuild(ctx, opts.UserMessage); err == nil && extra != "" {
|
||||
enrichedCtx = extra
|
||||
}
|
||||
}
|
||||
|
||||
history := al.sessions.GetHistory(opts.SessionKey)
|
||||
summary := al.sessions.GetSummary(opts.SessionKey)
|
||||
messages := al.contextBuilder.BuildMessages(
|
||||
|
|
@ -246,6 +264,7 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, opts processOptions) (str
|
|||
nil,
|
||||
opts.Channel,
|
||||
opts.ChatID,
|
||||
enrichedCtx,
|
||||
)
|
||||
|
||||
// 3. Save user message to session
|
||||
|
|
@ -266,6 +285,11 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, opts processOptions) (str
|
|||
al.sessions.AddMessage(opts.SessionKey, "assistant", finalContent)
|
||||
al.sessions.Save(al.sessions.GetOrCreate(opts.SessionKey))
|
||||
|
||||
// 6b. OnPostMessage hook
|
||||
if al.hooks != nil && al.hooks.OnPostMessage != nil {
|
||||
al.hooks.OnPostMessage(ctx, opts.SessionKey, opts.UserMessage, finalContent)
|
||||
}
|
||||
|
||||
// 7. Optional: summarization
|
||||
if opts.EnableSummary {
|
||||
al.maybeSummarize(opts.SessionKey)
|
||||
|
|
@ -411,10 +435,25 @@ func (al *AgentLoop) runLLMIteration(ctx context.Context, messages []providers.M
|
|||
"iteration": iteration,
|
||||
})
|
||||
|
||||
// OnPreTool hook
|
||||
if al.hooks != nil && al.hooks.OnPreTool != nil {
|
||||
if err := al.hooks.OnPreTool(ctx, tc.Name, tc.Arguments); err != nil {
|
||||
logger.DebugCF("agent", "OnPreTool hook error",
|
||||
map[string]interface{}{"tool": tc.Name, "error": err.Error()})
|
||||
}
|
||||
}
|
||||
|
||||
toolStart := time.Now()
|
||||
result, err := al.tools.ExecuteWithContext(ctx, tc.Name, tc.Arguments, opts.Channel, opts.ChatID)
|
||||
if err != nil {
|
||||
result = fmt.Sprintf("Error: %v", err)
|
||||
}
|
||||
toolDur := time.Since(toolStart)
|
||||
|
||||
// OnPostTool hook
|
||||
if al.hooks != nil && al.hooks.OnPostTool != nil {
|
||||
al.hooks.OnPostTool(ctx, tc.Name, result, toolDur)
|
||||
}
|
||||
|
||||
toolResultMsg := providers.Message{
|
||||
Role: "tool",
|
||||
|
|
|
|||
|
|
@ -49,9 +49,24 @@ type Config struct {
|
|||
Providers ProvidersConfig `json:"providers"`
|
||||
Gateway GatewayConfig `json:"gateway"`
|
||||
Tools ToolsConfig `json:"tools"`
|
||||
Memory MemoryConfig `json:"memory"`
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
type MemoryConfig struct {
|
||||
VectorSearch VectorSearchConfig `json:"vector_search"`
|
||||
}
|
||||
|
||||
type VectorSearchConfig struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
Provider string `json:"provider"` // default "openai"
|
||||
Model string `json:"model"` // default "text-embedding-3-small"
|
||||
APIKey string `json:"api_key"` // falls back to providers config
|
||||
APIBase string `json:"api_base"` // falls back to providers config
|
||||
MaxResults int `json:"max_results"` // default 5
|
||||
ChunkSize int `json:"chunk_size"` // default 800
|
||||
}
|
||||
|
||||
type AgentsConfig struct {
|
||||
Defaults AgentDefaults `json:"defaults"`
|
||||
}
|
||||
|
|
@ -255,6 +270,15 @@ func DefaultConfig() *Config {
|
|||
},
|
||||
},
|
||||
},
|
||||
Memory: MemoryConfig{
|
||||
VectorSearch: VectorSearchConfig{
|
||||
Enabled: false,
|
||||
Provider: "openai",
|
||||
Model: "text-embedding-3-small",
|
||||
MaxResults: 5,
|
||||
ChunkSize: 800,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
81
pkg/tools/memory_search.go
Normal file
81
pkg/tools/memory_search.go
Normal file
|
|
@ -0,0 +1,81 @@
|
|||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/vecstore"
|
||||
)
|
||||
|
||||
// MemorySearchTool searches memory using vector similarity.
|
||||
type MemorySearchTool struct {
|
||||
embedder vecstore.Embedder
|
||||
store *vecstore.VectorStore
|
||||
maxResults int
|
||||
}
|
||||
|
||||
// NewMemorySearchTool creates a memory search tool.
|
||||
func NewMemorySearchTool(embedder vecstore.Embedder, store *vecstore.VectorStore, maxResults int) *MemorySearchTool {
|
||||
if maxResults <= 0 {
|
||||
maxResults = 5
|
||||
}
|
||||
return &MemorySearchTool{
|
||||
embedder: embedder,
|
||||
store: store,
|
||||
maxResults: maxResults,
|
||||
}
|
||||
}
|
||||
|
||||
func (t *MemorySearchTool) Name() string { return "memory_search" }
|
||||
|
||||
func (t *MemorySearchTool) Description() string {
|
||||
return "Search long-term memory for relevant information using semantic similarity. Use this to find specific memories, notes, or facts."
|
||||
}
|
||||
|
||||
func (t *MemorySearchTool) Parameters() map[string]interface{} {
|
||||
return map[string]interface{}{
|
||||
"type": "object",
|
||||
"properties": map[string]interface{}{
|
||||
"query": map[string]interface{}{
|
||||
"type": "string",
|
||||
"description": "The search query to find relevant memories",
|
||||
},
|
||||
},
|
||||
"required": []string{"query"},
|
||||
}
|
||||
}
|
||||
|
||||
func (t *MemorySearchTool) Execute(ctx context.Context, args map[string]interface{}) (string, error) {
|
||||
query, _ := args["query"].(string)
|
||||
if query == "" {
|
||||
return "", fmt.Errorf("query is required")
|
||||
}
|
||||
|
||||
// Embed the query
|
||||
embeddings, err := t.embedder.Embed(ctx, []string{query})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("embed query: %w", err)
|
||||
}
|
||||
if len(embeddings) == 0 || len(embeddings[0]) == 0 {
|
||||
return "No results found.", nil
|
||||
}
|
||||
|
||||
// Search
|
||||
results := t.store.Search(embeddings[0], t.maxResults)
|
||||
if len(results) == 0 {
|
||||
return "No relevant memories found.", nil
|
||||
}
|
||||
|
||||
// Format results
|
||||
var sb strings.Builder
|
||||
sb.WriteString(fmt.Sprintf("Found %d relevant memories:\n\n", len(results)))
|
||||
for i, r := range results {
|
||||
snippet := r.Text
|
||||
if len(snippet) > 700 {
|
||||
snippet = snippet[:700] + "..."
|
||||
}
|
||||
sb.WriteString(fmt.Sprintf("--- Result %d (score: %.2f, source: %s) ---\n%s\n\n", i+1, r.Score, r.Source, snippet))
|
||||
}
|
||||
return sb.String(), nil
|
||||
}
|
||||
109
pkg/tools/memory_search_test.go
Normal file
109
pkg/tools/memory_search_test.go
Normal file
|
|
@ -0,0 +1,109 @@
|
|||
package tools
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/sipeed/picoclaw/pkg/vecstore"
|
||||
)
|
||||
|
||||
// mockEmbedder returns a fixed embedding for any input.
|
||||
type mockEmbedder struct {
|
||||
embedding []float32
|
||||
err error
|
||||
}
|
||||
|
||||
func (m *mockEmbedder) Embed(_ context.Context, texts []string) ([][]float32, error) {
|
||||
if m.err != nil {
|
||||
return nil, m.err
|
||||
}
|
||||
result := make([][]float32, len(texts))
|
||||
for i := range texts {
|
||||
result[i] = m.embedding
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func TestMemorySearchExecute(t *testing.T) {
|
||||
store := vecstore.NewVectorStore("")
|
||||
now := time.Now()
|
||||
store.Upsert([]vecstore.Chunk{
|
||||
{ID: "a", Text: "The user prefers dark mode", Source: "memory/MEMORY.md", Embedding: []float32{1, 0, 0}, UpdatedAt: now},
|
||||
{ID: "b", Text: "Meeting notes from Monday", Source: "memory/202601/20260112.md", Embedding: []float32{0, 1, 0}, UpdatedAt: now},
|
||||
{ID: "c", Text: "User timezone is PST", Source: "memory/MEMORY.md", Embedding: []float32{0.9, 0.1, 0}, UpdatedAt: now},
|
||||
})
|
||||
|
||||
embedder := &mockEmbedder{embedding: []float32{1, 0, 0}}
|
||||
tool := NewMemorySearchTool(embedder, store, 2)
|
||||
|
||||
result, err := tool.Execute(context.Background(), map[string]interface{}{
|
||||
"query": "user preferences",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if !strings.Contains(result, "dark mode") {
|
||||
t.Error("expected top result to contain 'dark mode'")
|
||||
}
|
||||
if !strings.Contains(result, "timezone") {
|
||||
t.Error("expected second result to contain 'timezone'")
|
||||
}
|
||||
if strings.Contains(result, "Meeting notes") {
|
||||
t.Error("should not contain third result (maxResults=2)")
|
||||
}
|
||||
if !strings.Contains(result, "score:") {
|
||||
t.Error("result should include scores")
|
||||
}
|
||||
if !strings.Contains(result, "memory/MEMORY.md") {
|
||||
t.Error("result should include source path")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMemorySearchEmptyQuery(t *testing.T) {
|
||||
store := vecstore.NewVectorStore("")
|
||||
embedder := &mockEmbedder{embedding: []float32{1, 0}}
|
||||
tool := NewMemorySearchTool(embedder, store, 5)
|
||||
|
||||
_, err := tool.Execute(context.Background(), map[string]interface{}{
|
||||
"query": "",
|
||||
})
|
||||
if err == nil {
|
||||
t.Error("expected error for empty query")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMemorySearchEmbedError(t *testing.T) {
|
||||
store := vecstore.NewVectorStore("")
|
||||
embedder := &mockEmbedder{err: fmt.Errorf("API unavailable")}
|
||||
tool := NewMemorySearchTool(embedder, store, 5)
|
||||
|
||||
_, err := tool.Execute(context.Background(), map[string]interface{}{
|
||||
"query": "test",
|
||||
})
|
||||
if err == nil {
|
||||
t.Error("expected error when embedder fails")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "embed query") {
|
||||
t.Errorf("error should wrap embed failure, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMemorySearchNoResults(t *testing.T) {
|
||||
store := vecstore.NewVectorStore("") // empty store
|
||||
embedder := &mockEmbedder{embedding: []float32{1, 0}}
|
||||
tool := NewMemorySearchTool(embedder, store, 5)
|
||||
|
||||
result, err := tool.Execute(context.Background(), map[string]interface{}{
|
||||
"query": "anything",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !strings.Contains(result, "No relevant memories") {
|
||||
t.Errorf("expected no-results message, got: %s", result)
|
||||
}
|
||||
}
|
||||
112
pkg/vecstore/chunker.go
Normal file
112
pkg/vecstore/chunker.go
Normal file
|
|
@ -0,0 +1,112 @@
|
|||
package vecstore
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ChunkMarkdown splits markdown text into chunks at semantic boundaries.
|
||||
// Splits first by ## headers, then sub-splits long sections by paragraphs.
|
||||
// Each chunk gets a deterministic ID: sha256(source + ":" + text)[:12].
|
||||
func ChunkMarkdown(source, text string, maxChars int) []Chunk {
|
||||
if maxChars <= 0 {
|
||||
maxChars = 800
|
||||
}
|
||||
|
||||
sections := splitByHeaders(text)
|
||||
now := time.Now()
|
||||
|
||||
var chunks []Chunk
|
||||
for _, section := range sections {
|
||||
section = strings.TrimSpace(section)
|
||||
if section == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
if len(section) <= maxChars {
|
||||
chunks = append(chunks, makeChunk(source, section, now))
|
||||
continue
|
||||
}
|
||||
|
||||
// Sub-split long sections by paragraphs
|
||||
for _, part := range splitByParagraphs(section, maxChars) {
|
||||
part = strings.TrimSpace(part)
|
||||
if part == "" {
|
||||
continue
|
||||
}
|
||||
chunks = append(chunks, makeChunk(source, part, now))
|
||||
}
|
||||
}
|
||||
return chunks
|
||||
}
|
||||
|
||||
// splitByHeaders splits text at ## header boundaries, keeping the header with its content.
|
||||
func splitByHeaders(text string) []string {
|
||||
lines := strings.Split(text, "\n")
|
||||
var sections []string
|
||||
var current strings.Builder
|
||||
|
||||
for _, line := range lines {
|
||||
if strings.HasPrefix(line, "## ") && current.Len() > 0 {
|
||||
sections = append(sections, current.String())
|
||||
current.Reset()
|
||||
}
|
||||
current.WriteString(line)
|
||||
current.WriteByte('\n')
|
||||
}
|
||||
if current.Len() > 0 {
|
||||
sections = append(sections, current.String())
|
||||
}
|
||||
return sections
|
||||
}
|
||||
|
||||
// splitByParagraphs splits text at double-newline boundaries, respecting maxChars.
|
||||
func splitByParagraphs(text string, maxChars int) []string {
|
||||
paragraphs := strings.Split(text, "\n\n")
|
||||
var parts []string
|
||||
var current strings.Builder
|
||||
|
||||
for _, p := range paragraphs {
|
||||
p = strings.TrimSpace(p)
|
||||
if p == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
// If adding this paragraph would exceed max, flush current
|
||||
if current.Len() > 0 && current.Len()+len(p)+2 > maxChars {
|
||||
parts = append(parts, current.String())
|
||||
current.Reset()
|
||||
}
|
||||
|
||||
// If a single paragraph exceeds max, just add it as-is
|
||||
if current.Len() == 0 && len(p) > maxChars {
|
||||
parts = append(parts, p)
|
||||
continue
|
||||
}
|
||||
|
||||
if current.Len() > 0 {
|
||||
current.WriteString("\n\n")
|
||||
}
|
||||
current.WriteString(p)
|
||||
}
|
||||
if current.Len() > 0 {
|
||||
parts = append(parts, current.String())
|
||||
}
|
||||
return parts
|
||||
}
|
||||
|
||||
func makeChunk(source, text string, now time.Time) Chunk {
|
||||
return Chunk{
|
||||
ID: chunkID(source, text),
|
||||
Text: text,
|
||||
Source: source,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
}
|
||||
|
||||
func chunkID(source, text string) string {
|
||||
h := sha256.Sum256([]byte(source + ":" + text))
|
||||
return fmt.Sprintf("%x", h[:6]) // 12 hex chars
|
||||
}
|
||||
104
pkg/vecstore/chunker_test.go
Normal file
104
pkg/vecstore/chunker_test.go
Normal file
|
|
@ -0,0 +1,104 @@
|
|||
package vecstore
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestChunkMarkdownByHeaders(t *testing.T) {
|
||||
md := `# Title
|
||||
|
||||
Intro paragraph.
|
||||
|
||||
## Section A
|
||||
|
||||
Content A here.
|
||||
|
||||
## Section B
|
||||
|
||||
Content B here.
|
||||
`
|
||||
chunks := ChunkMarkdown("test.md", md, 800)
|
||||
if len(chunks) < 2 {
|
||||
t.Fatalf("expected at least 2 chunks, got %d", len(chunks))
|
||||
}
|
||||
|
||||
// First chunk should contain "Title" and "Intro"
|
||||
if !strings.Contains(chunks[0].Text, "Title") {
|
||||
t.Errorf("first chunk should contain Title: %q", chunks[0].Text)
|
||||
}
|
||||
|
||||
// Should have sections A and B as separate chunks
|
||||
foundA, foundB := false, false
|
||||
for _, c := range chunks {
|
||||
if strings.Contains(c.Text, "Section A") {
|
||||
foundA = true
|
||||
}
|
||||
if strings.Contains(c.Text, "Section B") {
|
||||
foundB = true
|
||||
}
|
||||
}
|
||||
if !foundA || !foundB {
|
||||
t.Errorf("expected sections A and B in separate chunks, foundA=%v foundB=%v", foundA, foundB)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChunkMarkdownLongSection(t *testing.T) {
|
||||
// Create a long section that exceeds maxChars
|
||||
long := "## Big Section\n\n"
|
||||
for i := 0; i < 20; i++ {
|
||||
long += "This is paragraph number " + string(rune('A'+i)) + ". It has some content.\n\n"
|
||||
}
|
||||
|
||||
chunks := ChunkMarkdown("test.md", long, 200)
|
||||
if len(chunks) < 2 {
|
||||
t.Fatalf("expected multiple chunks for long section, got %d", len(chunks))
|
||||
}
|
||||
|
||||
for _, c := range chunks {
|
||||
if c.Source != "test.md" {
|
||||
t.Errorf("expected source 'test.md', got %q", c.Source)
|
||||
}
|
||||
if c.ID == "" {
|
||||
t.Error("chunk ID should not be empty")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestChunkMarkdownDeterministicIDs(t *testing.T) {
|
||||
md := "## Hello\n\nWorld"
|
||||
c1 := ChunkMarkdown("src.md", md, 800)
|
||||
c2 := ChunkMarkdown("src.md", md, 800)
|
||||
|
||||
if len(c1) != len(c2) {
|
||||
t.Fatalf("chunk counts differ: %d vs %d", len(c1), len(c2))
|
||||
}
|
||||
for i := range c1 {
|
||||
if c1[i].ID != c2[i].ID {
|
||||
t.Errorf("chunk %d: IDs differ %q vs %q", i, c1[i].ID, c2[i].ID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestChunkMarkdownEmpty(t *testing.T) {
|
||||
chunks := ChunkMarkdown("test.md", "", 800)
|
||||
if len(chunks) != 0 {
|
||||
t.Errorf("expected 0 chunks for empty text, got %d", len(chunks))
|
||||
}
|
||||
}
|
||||
|
||||
func TestChunkIDUniqueness(t *testing.T) {
|
||||
// Same text, different source → different ID
|
||||
id1 := chunkID("a.md", "hello")
|
||||
id2 := chunkID("b.md", "hello")
|
||||
if id1 == id2 {
|
||||
t.Error("IDs should differ for different sources")
|
||||
}
|
||||
|
||||
// Same source, different text → different ID
|
||||
id3 := chunkID("a.md", "hello")
|
||||
id4 := chunkID("a.md", "world")
|
||||
if id3 == id4 {
|
||||
t.Error("IDs should differ for different text")
|
||||
}
|
||||
}
|
||||
131
pkg/vecstore/embed.go
Normal file
131
pkg/vecstore/embed.go
Normal file
|
|
@ -0,0 +1,131 @@
|
|||
package vecstore
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Embedder generates embedding vectors from text.
|
||||
type Embedder interface {
|
||||
Embed(ctx context.Context, texts []string) ([][]float32, error)
|
||||
}
|
||||
|
||||
// HTTPEmbedder calls an OpenAI-compatible /v1/embeddings endpoint.
|
||||
type HTTPEmbedder struct {
|
||||
apiBase string
|
||||
apiKey string
|
||||
model string
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
// NewHTTPEmbedder creates an embedder targeting an OpenAI-compatible API.
|
||||
func NewHTTPEmbedder(apiBase, apiKey, model string) *HTTPEmbedder {
|
||||
return &HTTPEmbedder{
|
||||
apiBase: apiBase,
|
||||
apiKey: apiKey,
|
||||
model: model,
|
||||
client: &http.Client{
|
||||
Timeout: 60 * time.Second,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
type embeddingRequest struct {
|
||||
Input []string `json:"input"`
|
||||
Model string `json:"model"`
|
||||
}
|
||||
|
||||
type embeddingResponse struct {
|
||||
Data []struct {
|
||||
Embedding []float32 `json:"embedding"`
|
||||
Index int `json:"index"`
|
||||
} `json:"data"`
|
||||
Error *struct {
|
||||
Message string `json:"message"`
|
||||
} `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// Embed sends all texts in one batch request and returns their embeddings.
|
||||
// Retries up to 3 times with exponential backoff on transient errors.
|
||||
func (e *HTTPEmbedder) Embed(ctx context.Context, texts []string) ([][]float32, error) {
|
||||
if len(texts) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
body, err := json.Marshal(embeddingRequest{
|
||||
Input: texts,
|
||||
Model: e.model,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("marshal embedding request: %w", err)
|
||||
}
|
||||
|
||||
url := e.apiBase + "/embeddings"
|
||||
|
||||
const maxRetries = 3
|
||||
var lastErr error
|
||||
for attempt := range maxRetries {
|
||||
result, err := e.doRequest(ctx, url, body, len(texts))
|
||||
if err == nil {
|
||||
return result, nil
|
||||
}
|
||||
lastErr = err
|
||||
// Exponential backoff: 500ms, 2s, 8s
|
||||
backoff := time.Duration(math.Pow(4, float64(attempt))) * 500 * time.Millisecond
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-time.After(backoff):
|
||||
}
|
||||
}
|
||||
return nil, fmt.Errorf("embedding failed after %d retries: %w", maxRetries, lastErr)
|
||||
}
|
||||
|
||||
func (e *HTTPEmbedder) doRequest(ctx context.Context, url string, body []byte, n int) ([][]float32, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if e.apiKey != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+e.apiKey)
|
||||
}
|
||||
|
||||
resp, err := e.client.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
respBody, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("embedding API %d: %s", resp.StatusCode, string(respBody))
|
||||
}
|
||||
|
||||
var result embeddingResponse
|
||||
if err := json.Unmarshal(respBody, &result); err != nil {
|
||||
return nil, fmt.Errorf("decode response: %w", err)
|
||||
}
|
||||
if result.Error != nil {
|
||||
return nil, fmt.Errorf("embedding API error: %s", result.Error.Message)
|
||||
}
|
||||
|
||||
// Order by index
|
||||
embeddings := make([][]float32, n)
|
||||
for _, d := range result.Data {
|
||||
if d.Index < len(embeddings) {
|
||||
embeddings[d.Index] = d.Embedding
|
||||
}
|
||||
}
|
||||
return embeddings, nil
|
||||
}
|
||||
167
pkg/vecstore/store.go
Normal file
167
pkg/vecstore/store.go
Normal file
|
|
@ -0,0 +1,167 @@
|
|||
package vecstore
|
||||
|
||||
import (
|
||||
"encoding/gob"
|
||||
"math"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Chunk represents a text chunk with its embedding vector.
|
||||
type Chunk struct {
|
||||
ID string
|
||||
Text string
|
||||
Source string // file path the chunk came from
|
||||
Embedding []float32
|
||||
UpdatedAt time.Time
|
||||
}
|
||||
|
||||
// Result is a search result with similarity score.
|
||||
type Result struct {
|
||||
Chunk
|
||||
Score float32
|
||||
}
|
||||
|
||||
// VectorStore is an in-memory vector store with gob persistence.
|
||||
type VectorStore struct {
|
||||
path string
|
||||
chunks []Chunk
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
// NewVectorStore creates a store that persists to the given path.
|
||||
func NewVectorStore(path string) *VectorStore {
|
||||
return &VectorStore{path: path}
|
||||
}
|
||||
|
||||
// Load reads the store from disk. Returns nil if file doesn't exist.
|
||||
func (vs *VectorStore) Load() error {
|
||||
vs.mu.Lock()
|
||||
defer vs.mu.Unlock()
|
||||
|
||||
f, err := os.Open(vs.path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
vs.chunks = nil
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
var chunks []Chunk
|
||||
if err := gob.NewDecoder(f).Decode(&chunks); err != nil {
|
||||
// Corrupt file — start fresh
|
||||
vs.chunks = nil
|
||||
return nil
|
||||
}
|
||||
vs.chunks = chunks
|
||||
return nil
|
||||
}
|
||||
|
||||
// Save writes the store to disk.
|
||||
func (vs *VectorStore) Save() error {
|
||||
vs.mu.RLock()
|
||||
defer vs.mu.RUnlock()
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(vs.path), 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
f, err := os.Create(vs.path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
return gob.NewEncoder(f).Encode(vs.chunks)
|
||||
}
|
||||
|
||||
// Search returns the top-K chunks most similar to the query embedding.
|
||||
func (vs *VectorStore) Search(query []float32, topK int) []Result {
|
||||
vs.mu.RLock()
|
||||
defer vs.mu.RUnlock()
|
||||
|
||||
if len(vs.chunks) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
results := make([]Result, 0, len(vs.chunks))
|
||||
for _, c := range vs.chunks {
|
||||
if len(c.Embedding) == 0 {
|
||||
continue
|
||||
}
|
||||
score := cosine(query, c.Embedding)
|
||||
results = append(results, Result{Chunk: c, Score: score})
|
||||
}
|
||||
|
||||
sort.Slice(results, func(i, j int) bool {
|
||||
return results[i].Score > results[j].Score
|
||||
})
|
||||
|
||||
if topK > len(results) {
|
||||
topK = len(results)
|
||||
}
|
||||
return results[:topK]
|
||||
}
|
||||
|
||||
// Upsert adds or replaces chunks by ID.
|
||||
func (vs *VectorStore) Upsert(chunks []Chunk) {
|
||||
vs.mu.Lock()
|
||||
defer vs.mu.Unlock()
|
||||
|
||||
idx := make(map[string]int, len(vs.chunks))
|
||||
for i, c := range vs.chunks {
|
||||
idx[c.ID] = i
|
||||
}
|
||||
|
||||
for _, c := range chunks {
|
||||
if i, ok := idx[c.ID]; ok {
|
||||
vs.chunks[i] = c
|
||||
} else {
|
||||
vs.chunks = append(vs.chunks, c)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// DeleteBySource removes all chunks from a given source.
|
||||
func (vs *VectorStore) DeleteBySource(source string) {
|
||||
vs.mu.Lock()
|
||||
defer vs.mu.Unlock()
|
||||
|
||||
filtered := vs.chunks[:0]
|
||||
for _, c := range vs.chunks {
|
||||
if c.Source != source {
|
||||
filtered = append(filtered, c)
|
||||
}
|
||||
}
|
||||
vs.chunks = filtered
|
||||
}
|
||||
|
||||
// Len returns the number of chunks in the store.
|
||||
func (vs *VectorStore) Len() int {
|
||||
vs.mu.RLock()
|
||||
defer vs.mu.RUnlock()
|
||||
return len(vs.chunks)
|
||||
}
|
||||
|
||||
// cosine computes cosine similarity between two vectors.
|
||||
func cosine(a, b []float32) float32 {
|
||||
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])
|
||||
}
|
||||
denom := math.Sqrt(normA) * math.Sqrt(normB)
|
||||
if denom == 0 {
|
||||
return 0
|
||||
}
|
||||
return float32(dot / denom)
|
||||
}
|
||||
150
pkg/vecstore/store_test.go
Normal file
150
pkg/vecstore/store_test.go
Normal file
|
|
@ -0,0 +1,150 @@
|
|||
package vecstore
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCosine(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
a, b []float32
|
||||
want float32
|
||||
tol float32
|
||||
}{
|
||||
{"identical", []float32{1, 0, 0}, []float32{1, 0, 0}, 1.0, 0.001},
|
||||
{"orthogonal", []float32{1, 0, 0}, []float32{0, 1, 0}, 0.0, 0.001},
|
||||
{"opposite", []float32{1, 0}, []float32{-1, 0}, -1.0, 0.001},
|
||||
{"similar", []float32{1, 1}, []float32{1, 0.9}, 0.998, 0.01},
|
||||
{"empty", []float32{}, []float32{}, 0.0, 0.001},
|
||||
{"mismatched", []float32{1, 2}, []float32{1, 2, 3}, 0.0, 0.001},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := cosine(tt.a, tt.b)
|
||||
if diff := got - tt.want; diff > tt.tol || diff < -tt.tol {
|
||||
t.Errorf("cosine(%v, %v) = %f, want %f (tol %f)", tt.a, tt.b, got, tt.want, tt.tol)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSearchReturnsTopK(t *testing.T) {
|
||||
store := NewVectorStore("")
|
||||
now := time.Now()
|
||||
|
||||
store.Upsert([]Chunk{
|
||||
{ID: "a", Text: "alpha", Embedding: []float32{1, 0, 0}, UpdatedAt: now},
|
||||
{ID: "b", Text: "beta", Embedding: []float32{0, 1, 0}, UpdatedAt: now},
|
||||
{ID: "c", Text: "gamma", Embedding: []float32{0.9, 0.1, 0}, UpdatedAt: now},
|
||||
})
|
||||
|
||||
results := store.Search([]float32{1, 0, 0}, 2)
|
||||
if len(results) != 2 {
|
||||
t.Fatalf("expected 2 results, got %d", len(results))
|
||||
}
|
||||
if results[0].ID != "a" {
|
||||
t.Errorf("expected first result 'a', got %q", results[0].ID)
|
||||
}
|
||||
if results[1].ID != "c" {
|
||||
t.Errorf("expected second result 'c', got %q", results[1].ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpsertReplacesExisting(t *testing.T) {
|
||||
store := NewVectorStore("")
|
||||
now := time.Now()
|
||||
|
||||
store.Upsert([]Chunk{
|
||||
{ID: "a", Text: "original", Embedding: []float32{1, 0}, UpdatedAt: now},
|
||||
})
|
||||
store.Upsert([]Chunk{
|
||||
{ID: "a", Text: "replaced", Embedding: []float32{0, 1}, UpdatedAt: now},
|
||||
})
|
||||
|
||||
if store.Len() != 1 {
|
||||
t.Fatalf("expected 1 chunk, got %d", store.Len())
|
||||
}
|
||||
|
||||
results := store.Search([]float32{0, 1}, 1)
|
||||
if results[0].Text != "replaced" {
|
||||
t.Errorf("expected replaced text, got %q", results[0].Text)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteBySource(t *testing.T) {
|
||||
store := NewVectorStore("")
|
||||
now := time.Now()
|
||||
|
||||
store.Upsert([]Chunk{
|
||||
{ID: "a", Text: "a", Source: "file1.md", Embedding: []float32{1, 0}, UpdatedAt: now},
|
||||
{ID: "b", Text: "b", Source: "file2.md", Embedding: []float32{0, 1}, UpdatedAt: now},
|
||||
{ID: "c", Text: "c", Source: "file1.md", Embedding: []float32{1, 1}, UpdatedAt: now},
|
||||
})
|
||||
|
||||
store.DeleteBySource("file1.md")
|
||||
if store.Len() != 1 {
|
||||
t.Fatalf("expected 1 chunk after delete, got %d", store.Len())
|
||||
}
|
||||
|
||||
results := store.Search([]float32{0, 1}, 10)
|
||||
if results[0].Source != "file2.md" {
|
||||
t.Errorf("expected remaining chunk from file2.md, got %q", results[0].Source)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadSave(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "test.gob")
|
||||
now := time.Now()
|
||||
|
||||
// Save
|
||||
store1 := NewVectorStore(path)
|
||||
store1.Upsert([]Chunk{
|
||||
{ID: "x", Text: "hello", Source: "src", Embedding: []float32{0.5, 0.5}, UpdatedAt: now},
|
||||
})
|
||||
if err := store1.Save(); err != nil {
|
||||
t.Fatalf("save: %v", err)
|
||||
}
|
||||
|
||||
// Load into new store
|
||||
store2 := NewVectorStore(path)
|
||||
if err := store2.Load(); err != nil {
|
||||
t.Fatalf("load: %v", err)
|
||||
}
|
||||
if store2.Len() != 1 {
|
||||
t.Fatalf("expected 1 chunk after load, got %d", store2.Len())
|
||||
}
|
||||
|
||||
results := store2.Search([]float32{0.5, 0.5}, 1)
|
||||
if results[0].Text != "hello" {
|
||||
t.Errorf("expected 'hello', got %q", results[0].Text)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadMissingFile(t *testing.T) {
|
||||
store := NewVectorStore(filepath.Join(t.TempDir(), "nonexistent.gob"))
|
||||
if err := store.Load(); err != nil {
|
||||
t.Fatalf("load missing file should not error: %v", err)
|
||||
}
|
||||
if store.Len() != 0 {
|
||||
t.Fatalf("expected 0 chunks, got %d", store.Len())
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadCorruptFile(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "corrupt.gob")
|
||||
os.WriteFile(path, []byte("not valid gob"), 0644)
|
||||
|
||||
store := NewVectorStore(path)
|
||||
if err := store.Load(); err != nil {
|
||||
t.Fatalf("load corrupt file should not error: %v", err)
|
||||
}
|
||||
if store.Len() != 0 {
|
||||
t.Fatalf("expected 0 chunks after corrupt load, got %d", store.Len())
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue