diff --git a/server/internal/api/api.go b/server/internal/api/api.go index e8c6d07..35cfedf 100644 --- a/server/internal/api/api.go +++ b/server/internal/api/api.go @@ -65,7 +65,7 @@ type MemoryService interface { Restore(context.Context, memory.LifecycleCommand) (*memory.MutationResult, error) Forget(context.Context, memory.ForgetCommand) (*memory.ForgetResult, error) CreateRelation(context.Context, memory.CreateRelationCommand) (*memory.CreateRelationResult, error) - ListRelations(context.Context, memory.ListRelationsQuery) ([]memory.Relation, error) + ListRelations(context.Context, memory.ListRelationsQuery) (*memory.ListRelationsResult, error) } // DurableContextService is the scoped durable-context port (mem#70). Handlers diff --git a/server/internal/api/handlers_memory.go b/server/internal/api/handlers_memory.go index 879e1d4..f43b0c5 100644 --- a/server/internal/api/handlers_memory.go +++ b/server/internal/api/handlers_memory.go @@ -763,13 +763,15 @@ func (s *Server) handleListMemoryRelations(w http.ResponseWriter, r *http.Reques } tok := r.Context().Value(ctxToken).(*auth.Token) - relations, err := s.Memory.ListRelations(r.Context(), memory.ListRelationsQuery{ + cursor := strings.TrimSpace(r.URL.Query().Get("cursor")) + result, err := s.Memory.ListRelations(r.Context(), memory.ListRelationsQuery{ WorkspaceID: currentWorkspace(r).ID, MemoryID: id, Direction: direction, RelationType: relationType, AllowedPaths: tok.Paths, Limit: limit, + Cursor: cursor, }) if err != nil { switch { @@ -789,8 +791,12 @@ func (s *Server) handleListMemoryRelations(w http.ResponseWriter, r *http.Reques } return } - if relations == nil { - relations = []memory.Relation{} + if result.Relations == nil { + result.Relations = []memory.Relation{} } - writeJSON(w, http.StatusOK, map[string]any{"relations": relations}) + response := map[string]any{"relations": result.Relations} + if result.NextCursor != "" { + response["next_cursor"] = result.NextCursor + } + writeJSON(w, http.StatusOK, response) } diff --git a/server/internal/api/handlers_memory_test.go b/server/internal/api/handlers_memory_test.go index 54887a4..0fc2e6d 100644 --- a/server/internal/api/handlers_memory_test.go +++ b/server/internal/api/handlers_memory_test.go @@ -107,10 +107,10 @@ func (s *memoryServiceStub) CreateRelation( func (s *memoryServiceStub) ListRelations( _ context.Context, q memory.ListRelationsQuery, -) ([]memory.Relation, error) { +) (*memory.ListRelationsResult, error) { s.calls++ s.listRelationsQuery = q - return s.relations, s.controlErr + return &memory.ListRelationsResult{Relations: s.relations}, s.controlErr } func memoryHandlerContext(req *http.Request, paths []string) (*http.Request, uuid.UUID, uuid.UUID, uuid.UUID) { diff --git a/server/internal/db/migrations/0024_data_plane_hygiene.sql b/server/internal/db/migrations/0024_data_plane_hygiene.sql new file mode 100644 index 0000000..41ecabb --- /dev/null +++ b/server/internal/db/migrations/0024_data_plane_hygiene.sql @@ -0,0 +1,91 @@ +-- +goose Up +-- Data-plane hygiene: uniqueness, missing FK indexes, and cascade alignment. +-- See https://github.com/bytefolk/mem/issues/178 + +-- 1. embeddings_text: enforce one row per (file_id, chunk_index). +-- The write path (indexer.go) DELETEs by file_id then batch-inserts; this +-- constraint gives the invariant database-level teeth, matching the +-- per-file guarantee that embeddings_visual already has via its PK. +-- +goose StatementBegin +ALTER TABLE embeddings_text + ADD CONSTRAINT uq_embeddings_text_file_chunk UNIQUE (file_id, chunk_index); +-- +goose StatementEnd + +-- 2. memories: index the unindexed ON DELETE SET NULL foreign keys. +-- Without these, every file or user delete takes a RowExclusiveLock on +-- memories and performs a sequential scan to find rows to null out. +-- +goose StatementBegin +CREATE INDEX IF NOT EXISTS idx_memories_source_file_id + ON memories (source_file_id) WHERE source_file_id IS NOT NULL; +-- +goose StatementEnd + +-- +goose StatementBegin +CREATE INDEX IF NOT EXISTS idx_memories_created_by_user_id + ON memories (created_by_user_id) WHERE created_by_user_id IS NOT NULL; +-- +goose StatementEnd + +-- 3. memory_relations: align FK referential actions with memories CASCADE. +-- memories.workspace_id is ON DELETE CASCADE, but memory_relations FKs +-- defaulted to NO ACTION, so a cascading workspace delete would fail on +-- any edge touching the deleted memories. +-- +goose StatementBegin +ALTER TABLE memory_relations + DROP CONSTRAINT IF EXISTS memory_relations_workspace_id_fkey; +ALTER TABLE memory_relations + ADD CONSTRAINT memory_relations_workspace_id_fkey + FOREIGN KEY (workspace_id) REFERENCES workspaces(id) ON DELETE CASCADE; +-- +goose StatementEnd + +-- +goose StatementBegin +ALTER TABLE memory_relations + DROP CONSTRAINT IF EXISTS memory_relations_source_id_fkey; +ALTER TABLE memory_relations + ADD CONSTRAINT memory_relations_source_id_fkey + FOREIGN KEY (source_id) REFERENCES memories(id) ON DELETE CASCADE; +-- +goose StatementEnd + +-- +goose StatementBegin +ALTER TABLE memory_relations + DROP CONSTRAINT IF EXISTS memory_relations_target_id_fkey; +ALTER TABLE memory_relations + ADD CONSTRAINT memory_relations_target_id_fkey + FOREIGN KEY (target_id) REFERENCES memories(id) ON DELETE CASCADE; +-- +goose StatementEnd + +-- +goose Down +-- +goose StatementBegin +ALTER TABLE memory_relations + DROP CONSTRAINT IF EXISTS memory_relations_target_id_fkey; +ALTER TABLE memory_relations + ADD CONSTRAINT memory_relations_target_id_fkey + FOREIGN KEY (target_id) REFERENCES memories(id); +-- +goose StatementEnd + +-- +goose StatementBegin +ALTER TABLE memory_relations + DROP CONSTRAINT IF EXISTS memory_relations_source_id_fkey; +ALTER TABLE memory_relations + ADD CONSTRAINT memory_relations_source_id_fkey + FOREIGN KEY (source_id) REFERENCES memories(id); +-- +goose StatementEnd + +-- +goose StatementBegin +ALTER TABLE memory_relations + DROP CONSTRAINT IF EXISTS memory_relations_workspace_id_fkey; +ALTER TABLE memory_relations + ADD CONSTRAINT memory_relations_workspace_id_fkey + FOREIGN KEY (workspace_id) REFERENCES workspaces(id); +-- +goose StatementEnd + +-- +goose StatementBegin +DROP INDEX IF EXISTS idx_memories_created_by_user_id; +-- +goose StatementEnd + +-- +goose StatementBegin +DROP INDEX IF EXISTS idx_memories_source_file_id; +-- +goose StatementEnd + +-- +goose StatementBegin +ALTER TABLE embeddings_text + DROP CONSTRAINT IF EXISTS uq_embeddings_text_file_chunk; +-- +goose StatementEnd diff --git a/server/internal/indexer/enrichment_integration_test.go b/server/internal/indexer/enrichment_integration_test.go index d7acca3..517c572 100644 --- a/server/internal/indexer/enrichment_integration_test.go +++ b/server/internal/indexer/enrichment_integration_test.go @@ -640,6 +640,80 @@ func TestIndexerEnrichmentIntegration(t *testing.T) { ) } +func TestEmbeddingsTextChunkUniqueness(t *testing.T) { + dsn := os.Getenv("MEM_TEST_DB") + if dsn == "" { + t.Skip("MEM_TEST_DB not set; skipping DB integration test") + } + config, err := pgxpool.ParseConfig(dsn) + if err != nil { + t.Fatalf("parse MEM_TEST_DB: %v", err) + } + if !strings.HasSuffix(config.ConnConfig.Database, "_test") { + t.Fatalf( + "refusing to modify non-test database %q; MEM_TEST_DB must end in _test", + config.ConnConfig.Database, + ) + } + + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + database, err := memdb.Open(ctx, dsn) + if err != nil { + t.Fatalf("open test database: %v", err) + } + t.Cleanup(database.Close) + if err := database.Migrate(ctx); err != nil { + t.Fatalf("migrate test database: %v", err) + } + + var userID uuid.UUID + if err := database.Pool.QueryRow(ctx, ` + INSERT INTO users (email, password_hash) + VALUES ($1, 'integration-test') + RETURNING id + `, "embeddings-unique-"+uuid.NewString()+"@example.com").Scan(&userID); err != nil { + t.Fatalf("create user: %v", err) + } + t.Cleanup(func() { + cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cleanupCancel() + _, _ = database.Pool.Exec(cleanupCtx, `DELETE FROM users WHERE id = $1`, userID) + }) + + fileID := uuid.New() + if _, err := database.Pool.Exec(ctx, ` + INSERT INTO files ( + id, user_id, name, path, size, sha256, mime, storage_key, index_status + ) VALUES ($1,$2,'chunks.txt','/',1,$3,'text/plain',$4,'ready') + `, fileID, userID, strings.Repeat("b", 64), "embeddings-unique/"+fileID.String()); err != nil { + t.Fatalf("insert file: %v", err) + } + + vec := vectorLiteral(make([]float32, 768)) + if _, err := database.Pool.Exec(ctx, ` + INSERT INTO embeddings_text (file_id, chunk_index, chunk_text, embedding, provider) + VALUES ($1, 0, 'first', $2::vector, 'test:embed') + `, fileID, vec); err != nil { + t.Fatalf("insert first chunk: %v", err) + } + + _, err = database.Pool.Exec(ctx, ` + INSERT INTO embeddings_text (file_id, chunk_index, chunk_text, embedding, provider) + VALUES ($1, 0, 'duplicate', $2::vector, 'test:embed') + `, fileID, vec) + if err == nil { + t.Fatal("expected duplicate (file_id, chunk_index) to be rejected, got nil") + } + + if _, err := database.Pool.Exec(ctx, ` + INSERT INTO embeddings_text (file_id, chunk_index, chunk_text, embedding, provider) + VALUES ($1, 1, 'second chunk', $2::vector, 'test:embed') + `, fileID, vec); err != nil { + t.Fatalf("different chunk_index for same file should succeed: %v", err) + } +} + func assertIndexerFileProjection( t *testing.T, ctx context.Context, diff --git a/server/internal/indexgeneration/service.go b/server/internal/indexgeneration/service.go index a8bb35b..0188f95 100644 --- a/server/internal/indexgeneration/service.go +++ b/server/internal/indexgeneration/service.go @@ -287,19 +287,29 @@ func (s *Service) List(ctx context.Context, workspaceID uuid.UUID, limit int) ([ } defer rows.Close() out := make([]Build, 0, limit) + buildIDs := make([]uuid.UUID, 0, limit) for rows.Next() { build, err := scanBuild(rows) if err != nil { return nil, fmt.Errorf("%w: scan build: %v", ErrUnavailable, err) } - generations, err := listGenerations(ctx, s.pool, workspaceID, build.ID) - if err != nil { - return nil, err - } - build.Generations = generations + buildIDs = append(buildIDs, build.ID) out = append(out, *build) } - return out, rows.Err() + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("%w: list builds: %v", ErrUnavailable, err) + } + if len(buildIDs) == 0 { + return out, nil + } + generations, err := listGenerationsForBuilds(ctx, s.pool, workspaceID, buildIDs) + if err != nil { + return nil, err + } + for i := range out { + out[i].Generations = generations[out[i].ID] + } + return out, nil } func (s *Service) Cancel( diff --git a/server/internal/indexgeneration/store.go b/server/internal/indexgeneration/store.go index 245236e..8657473 100644 --- a/server/internal/indexgeneration/store.go +++ b/server/internal/indexgeneration/store.go @@ -106,6 +106,45 @@ func listGenerations( return out, rows.Err() } +func listGenerationsForBuilds( + ctx context.Context, + q queryer, + workspaceID uuid.UUID, + buildIDs []uuid.UUID, +) (map[uuid.UUID][]Generation, error) { + rows, err := q.Query(ctx, ` + SELECT id, build_id, workspace_id, route_kind, provider, + model_revision, output_dimension, pipeline_revision, + profile_id, profile_revision, state, created_at, updated_at + FROM index_generations + WHERE workspace_id = $1 AND build_id = ANY($2::uuid[]) + ORDER BY route_kind + `, workspaceID, buildIDs) + if err != nil { + return nil, fmt.Errorf("%w: list route generations: %v", ErrUnavailable, err) + } + defer rows.Close() + out := make(map[uuid.UUID][]Generation) + for rows.Next() { + var generation Generation + if err := rows.Scan( + &generation.ID, &generation.BuildID, &generation.WorkspaceID, + &generation.RouteKind, &generation.Provider, + &generation.ModelRevision, &generation.OutputDimension, + &generation.PipelineRevision, &generation.ProfileID, + &generation.ProfileRevision, &generation.State, + &generation.CreatedAt, &generation.UpdatedAt, + ); err != nil { + return nil, fmt.Errorf("%w: scan route generation: %v", ErrUnavailable, err) + } + out[generation.BuildID] = append(out[generation.BuildID], generation) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("%w: list route generations: %v", ErrUnavailable, err) + } + return out, nil +} + func getInflightIdentity( ctx context.Context, q queryer, diff --git a/server/internal/memory/memory_integration_test.go b/server/internal/memory/memory_integration_test.go index fbe694a..df171e7 100644 --- a/server/internal/memory/memory_integration_test.go +++ b/server/internal/memory/memory_integration_test.go @@ -1224,37 +1224,37 @@ func TestMemoryPostgres(t *testing.T) { t.Fatalf("superseded markers = %+v", superseded) } - outbound, err := service.ListRelations(ctx, ListRelationsQuery{ + outboundResult, err := service.ListRelations(ctx, ListRelationsQuery{ WorkspaceID: workspaceA, MemoryID: fresh.Memory.ID, Direction: "source", AllowedPaths: []string{scope}, }) - if err != nil || len(outbound) != 1 || - outbound[0].TargetID != old.Memory.ID || - outbound[0].RelationType != RelSupersedes || - outbound[0].Reason != "decision updated" { - t.Fatalf("outbound relations = %+v err=%v", outbound, err) + if err != nil || len(outboundResult.Relations) != 1 || + outboundResult.Relations[0].TargetID != old.Memory.ID || + outboundResult.Relations[0].RelationType != RelSupersedes || + outboundResult.Relations[0].Reason != "decision updated" { + t.Fatalf("outbound relations = %+v err=%v", outboundResult.Relations, err) } - inbound, err := service.ListRelations(ctx, ListRelationsQuery{ + inboundResult, err := service.ListRelations(ctx, ListRelationsQuery{ WorkspaceID: workspaceA, MemoryID: old.Memory.ID, Direction: "target", RelationType: RelSupersedes, AllowedPaths: []string{scope}, }) - if err != nil || len(inbound) != 1 || inbound[0].SourceID != fresh.Memory.ID { - t.Fatalf("inbound relations = %+v err=%v", inbound, err) + if err != nil || len(inboundResult.Relations) != 1 || inboundResult.Relations[0].SourceID != fresh.Memory.ID { + t.Fatalf("inbound relations = %+v err=%v", inboundResult.Relations, err) } - typed, err := service.ListRelations(ctx, ListRelationsQuery{ + typedResult, err := service.ListRelations(ctx, ListRelationsQuery{ WorkspaceID: workspaceA, MemoryID: old.Memory.ID, Direction: "target", RelationType: RelCorrects, AllowedPaths: []string{scope}, }) - if err != nil || len(typed) != 0 { - t.Fatalf("type-filtered relations = %+v err=%v", typed, err) + if err != nil || len(typedResult.Relations) != 0 { + t.Fatalf("type-filtered relations = %+v err=%v", typedResult.Relations, err) } // Anchors the caller cannot read are hidden as not found. diff --git a/server/internal/memory/relation.go b/server/internal/memory/relation.go index b9c238f..7645d9c 100644 --- a/server/internal/memory/relation.go +++ b/server/internal/memory/relation.go @@ -2,6 +2,9 @@ package memory import ( "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" "errors" "fmt" "strings" @@ -79,6 +82,13 @@ type ListRelationsQuery struct { RelationType string // optional filter AllowedPaths []string Limit int + Cursor string // opaque keyset cursor from a previous ListRelationsResult +} + +// ListRelationsResult is the paginated response for ListRelations. +type ListRelationsResult struct { + Relations []Relation `json:"relations"` + NextCursor string `json:"next_cursor,omitempty"` } // validateCreateRelationCommand checks the command fields without touching the database. @@ -236,7 +246,7 @@ func validateListRelationsQuery(q ListRelationsQuery) error { // ListRelations returns relations for a given memory. Direction controls // whether MemoryID is matched as source or target. -func (s *Service) ListRelations(ctx context.Context, q ListRelationsQuery) ([]Relation, error) { +func (s *Service) ListRelations(ctx context.Context, q ListRelationsQuery) (*ListRelationsResult, error) { if s == nil || s.pool == nil { return nil, fmt.Errorf("memory service is not configured") } @@ -280,6 +290,19 @@ func (s *Service) ListRelations(ctx context.Context, q ListRelationsQuery) ([]Re return nil, ErrForgotten } + filterHash, err := relationFilterHash(q.WorkspaceID, q.MemoryID, direction, q.RelationType) + if err != nil { + return nil, err + } + var cursor *decodedListCursor + if q.Cursor != "" { + decoded, err := decodeListCursor(q.Cursor, filterHash) + if err != nil { + return nil, err + } + cursor = &decoded + } + args = []any{q.WorkspaceID, q.MemoryID} var dirColumn string if direction == "source" { @@ -296,7 +319,15 @@ func (s *Service) ListRelations(ctx context.Context, q ListRelationsQuery) ([]Re args = append(args, relType) where = append(where, fmt.Sprintf("r.relation_type = $%d", len(args))) } - args = append(args, q.Limit) + if cursor != nil { + args = append(args, cursor.createdAt, cursor.id) + timeArg, idArg := len(args)-1, len(args) + where = append(where, fmt.Sprintf( + "(r.created_at < $%d OR (r.created_at = $%d AND r.id < $%d))", + timeArg, timeArg, idArg, + )) + } + args = append(args, q.Limit+1) limitIdx := len(args) sql := fmt.Sprintf(` @@ -313,16 +344,50 @@ func (s *Service) ListRelations(ctx context.Context, q ListRelationsQuery) ([]Re } defer rows.Close() - out := make([]Relation, 0, q.Limit) + relations := make([]Relation, 0, q.Limit+1) for rows.Next() { var rel Relation if err := rows.Scan(&rel.ID, &rel.WorkspaceID, &rel.SourceID, &rel.TargetID, &rel.RelationType, &rel.ActorUserID, &rel.ActorTokenID, &rel.Reason, &rel.CreatedAt); err != nil { return nil, fmt.Errorf("scan relation: %w", err) } - out = append(out, rel) + relations = append(relations, rel) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("list relations: %w", err) + } + + result := &ListRelationsResult{Relations: relations} + if len(relations) <= q.Limit { + return result, nil + } + result.Relations = relations[:q.Limit] + last := result.Relations[len(result.Relations)-1] + result.NextCursor, err = encodeListCursor(last.CreatedAt, last.ID, filterHash) + if err != nil { + return nil, err + } + return result, nil +} + +func relationFilterHash(workspaceID, memoryID uuid.UUID, direction, relationType string) (string, error) { + payload := struct { + WorkspaceID string `json:"workspace_id"` + MemoryID string `json:"memory_id"` + Direction string `json:"direction"` + RelationType string `json:"relation_type"` + }{ + WorkspaceID: workspaceID.String(), + MemoryID: memoryID.String(), + Direction: direction, + RelationType: strings.ToLower(strings.TrimSpace(relationType)), + } + encoded, err := json.Marshal(payload) + if err != nil { + return "", fmt.Errorf("encode relation filter hash: %w", err) } - return out, rows.Err() + sum := sha256.Sum256(encoded) + return hex.EncodeToString(sum[:]), nil } // IsSuperseded returns true if the given memory has been superseded or corrected