diff --git a/go/core/internal/database/client_postgres.go b/go/core/internal/database/client_postgres.go index 2090d0e7c..4e0bdc66d 100644 --- a/go/core/internal/database/client_postgres.go +++ b/go/core/internal/database/client_postgres.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "fmt" + "log" "strings" "time" @@ -712,13 +713,14 @@ func (c *postgresClient) SearchAgentMemory(ctx context.Context, agentName, userI } } + // Access-count bookkeeping is best-effort: a failure must not fail the search. if len(results) > 0 { ids := make([]string, len(results)) for i, r := range results { ids[i] = r.ID } if err := c.q.IncrementMemoryAccessCount(ctx, ids); err != nil { - return nil, fmt.Errorf("failed to increment access count: %w", err) + log.Printf("failed to increment memory access count: %v", err) } } diff --git a/go/core/internal/database/client_test.go b/go/core/internal/database/client_test.go index 0768cddb9..a8c263d7d 100644 --- a/go/core/internal/database/client_test.go +++ b/go/core/internal/database/client_test.go @@ -726,3 +726,58 @@ func TestPruneExpiredMemories(t *testing.T) { assert.Contains(t, ids, hotMem.ID, "Expired popular memory should have TTL extended and be retained") assert.Contains(t, ids, liveMem.ID, "Non-expired memory should be retained") } + +// TestSearchAgentMemoryConcurrentAccessCount verifies concurrent searches over +// overlapping rows do not deadlock when incrementing access_count and still +// return results. +func TestSearchAgentMemoryConcurrentAccessCount(t *testing.T) { + db := setupTestDB(t) + client := NewClient(db) + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + t.Cleanup(cancel) + + agentName := "concurrent-agent" + userID := "concurrent-user" + + // Small store so every search hits the same top rows (max overlap). + for i := range 5 { + err := client.StoreAgentMemory(ctx, &dbpkg.Memory{ + AgentName: agentName, + UserID: userID, + Content: fmt.Sprintf("shared memory %d", i), + Embedding: makeEmbedding(float32(i+1) * 0.15), + }) + require.NoError(t, err) + } + + const numGoroutines = 20 + const searchesPerGoroutine = 10 + + var wg sync.WaitGroup + errs := make(chan error, numGoroutines*searchesPerGoroutine) + wg.Add(numGoroutines) + + for range numGoroutines { + go func() { + defer wg.Done() + for range searchesPerGoroutine { + results, err := client.SearchAgentMemory(ctx, agentName, userID, makeEmbedding(0.5), 5) + if err != nil { + errs <- err + return + } + if len(results) == 0 { + errs <- fmt.Errorf("expected search results, got none") + return + } + } + }() + } + + wg.Wait() + close(errs) + + for err := range errs { + require.NoError(t, err, "concurrent memory search must not fail") + } +} diff --git a/go/core/internal/database/gen/memory.sql.go b/go/core/internal/database/gen/memory.sql.go index f2cd5f22b..655d88fd8 100644 --- a/go/core/internal/database/gen/memory.sql.go +++ b/go/core/internal/database/gen/memory.sql.go @@ -48,10 +48,17 @@ func (q *Queries) ExtendMemoryTTL(ctx context.Context) error { } const incrementMemoryAccessCount = `-- name: IncrementMemoryAccessCount :exec -UPDATE memory SET access_count = access_count + 1 -WHERE id = ANY($1::text[]) +UPDATE memory +SET access_count = access_count + 1 +WHERE id IN ( + SELECT id FROM memory + WHERE id = ANY($1::text[]) + ORDER BY id + FOR UPDATE +) ` +// Lock rows in id order to avoid deadlocks between concurrent overlapping increments. func (q *Queries) IncrementMemoryAccessCount(ctx context.Context, dollar_1 []string) error { _, err := q.db.Exec(ctx, incrementMemoryAccessCount, dollar_1) return err diff --git a/go/core/internal/database/gen/querier.go b/go/core/internal/database/gen/querier.go index 7ef2b8a8c..078a06502 100644 --- a/go/core/internal/database/gen/querier.go +++ b/go/core/internal/database/gen/querier.go @@ -25,6 +25,7 @@ type Querier interface { GetTool(ctx context.Context, id string) (Tool, error) GetToolServer(ctx context.Context, name string) (Toolserver, error) HardDeleteCrewAIMemory(ctx context.Context, arg HardDeleteCrewAIMemoryParams) error + // Lock rows in id order to avoid deadlocks between concurrent overlapping increments. IncrementMemoryAccessCount(ctx context.Context, dollar_1 []string) error InsertEvent(ctx context.Context, arg InsertEventParams) error InsertFeedback(ctx context.Context, arg InsertFeedbackParams) error diff --git a/go/core/internal/database/queries/memory.sql b/go/core/internal/database/queries/memory.sql index 461636450..b390fd9c5 100644 --- a/go/core/internal/database/queries/memory.sql +++ b/go/core/internal/database/queries/memory.sql @@ -13,8 +13,15 @@ ORDER BY embedding <=> $1 ASC LIMIT $4; -- name: IncrementMemoryAccessCount :exec -UPDATE memory SET access_count = access_count + 1 -WHERE id = ANY($1::text[]); +-- Lock rows in id order to avoid deadlocks between concurrent overlapping increments. +UPDATE memory +SET access_count = access_count + 1 +WHERE id IN ( + SELECT id FROM memory + WHERE id = ANY($1::text[]) + ORDER BY id + FOR UPDATE +); -- name: ListAgentMemories :many SELECT * FROM memory