Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 29 additions & 17 deletions internal/search/graph.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,6 @@ package search

import (
"context"
"database/sql"
"errors"
"fmt"
"log/slog"
"sort"
Expand All @@ -18,6 +16,9 @@ const (
pprAlpha = 0.15 // teleport probability (jump back to seed)
pprMaxIter = 20 // max power iterations
pprEpsilon = 1e-6 // convergence threshold

// Minimum hydration batch size avoids many tiny IN queries when limit is small.
minHydrationBatchSize = 64
)

// searchGraph expands seed entity IDs via BFS on the knowledge graph,
Expand Down Expand Up @@ -184,26 +185,37 @@ func hydrateEntityResults(
scores map[string]float64,
limit int,
) ([]*domain.SearchResult, error) {
if limit <= 0 || len(ids) == 0 {
return nil, nil
}

results := make([]*domain.SearchResult, 0, min(len(ids), limit))
for _, id := range ids {
if len(results) >= limit {
break
}
ent, err := store.GetEntity(ctx, kbID, id)
batchSize := min(len(ids), max(limit*2, minHydrationBatchSize))
for start := 0; start < len(ids) && len(results) < limit; start += batchSize {
end := min(start+batchSize, len(ids))
batchIDs := ids[start:end]

entitiesByID, err := store.GetEntitiesByIDs(ctx, kbID, batchIDs)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
slog.Warn("graph entity batch lookup failed", "count", len(batchIDs), "error", err)
return nil, fmt.Errorf("get entities by ids: %w", err)
}
for _, id := range batchIDs {
if len(results) >= limit {
break
}
ent, ok := entitiesByID[id]
if !ok {
continue
}
slog.Warn("graph entity lookup failed", "id", id, "error", err)
return nil, fmt.Errorf("get entity %s: %w", id, err)
results = append(results, &domain.SearchResult{
ID: ent.ID,
KBID: ent.KBID,
Type: domain.ItemEntity,
Content: ent.Name + ": " + ent.Summary,
Score: scores[id],
})
}
results = append(results, &domain.SearchResult{
ID: ent.ID,
KBID: ent.KBID,
Type: domain.ItemEntity,
Content: ent.Name + ": " + ent.Summary,
Score: scores[id],
})
}
return results, nil
}
Expand Down
130 changes: 130 additions & 0 deletions internal/search/graph_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,130 @@
package search

import (
"context"
"fmt"
"testing"
"time"

memex "github.com/vndee/memex"
"github.com/vndee/memex/internal/domain"
"github.com/vndee/memex/internal/storage"
)

func init() {
storage.MigrationSQL = memex.MigrationSQL()
}

type countingSearchStore struct {
storage.Store
getEntityCalls int
getEntitiesByIDsCalls int
getEntitiesByIDsBatchSizes []int
}

func (s *countingSearchStore) GetEntity(ctx context.Context, kbID, id string) (*domain.Entity, error) {
s.getEntityCalls++
return s.Store.GetEntity(ctx, kbID, id)
}

func (s *countingSearchStore) GetEntitiesByIDs(ctx context.Context, kbID string, ids []string) (map[string]*domain.Entity, error) {
s.getEntitiesByIDsCalls++
s.getEntitiesByIDsBatchSizes = append(s.getEntitiesByIDsBatchSizes, len(ids))
return s.Store.GetEntitiesByIDs(ctx, kbID, ids)
}

func TestHydrateEntityResults_UsesBatchLookup(t *testing.T) {
sqliteStore, err := storage.NewSQLiteStore(":memory:")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = sqliteStore.Close() })

now := time.Now().UTC()
if err := sqliteStore.CreateKB(context.Background(), &domain.KnowledgeBase{
ID: "kb1", Name: "KB 1",
EmbedConfig: domain.EmbedConfig{Provider: "ollama", Model: "nomic-embed-text"},
LLMConfig: domain.LLMConfig{Provider: "ollama", Model: "llama3.2"},
CreatedAt: now,
}); err != nil {
t.Fatal(err)
}
for _, e := range []*domain.Entity{
{ID: "e1", KBID: "kb1", Name: "Alice", Type: "person", Summary: "Engineer", CreatedAt: now, UpdatedAt: now},
{ID: "e2", KBID: "kb1", Name: "Bob", Type: "person", Summary: "Manager", CreatedAt: now, UpdatedAt: now},
} {
if err := sqliteStore.CreateEntity(context.Background(), e); err != nil {
t.Fatal(err)
}
}

store := &countingSearchStore{Store: sqliteStore}
ids := []string{"e1", "missing", "e2"}
scores := map[string]float64{"e1": 1, "e2": 0.5, "missing": 0.1}

results, err := hydrateEntityResults(context.Background(), store, "kb1", ids, scores, 10)
if err != nil {
t.Fatal(err)
}
if len(results) != 2 {
t.Fatalf("got %d results, want 2", len(results))
}
if store.getEntitiesByIDsCalls != 1 {
t.Fatalf("GetEntitiesByIDs calls = %d, want 1", store.getEntitiesByIDsCalls)
}
if store.getEntityCalls != 0 {
t.Fatalf("GetEntity calls = %d, want 0", store.getEntityCalls)
}
}

func TestHydrateEntityResults_DoesNotFetchAllIDsWhenLimitSmall(t *testing.T) {
sqliteStore, err := storage.NewSQLiteStore(":memory:")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = sqliteStore.Close() })

now := time.Now().UTC()
if err := sqliteStore.CreateKB(context.Background(), &domain.KnowledgeBase{
ID: "kb1", Name: "KB 1",
EmbedConfig: domain.EmbedConfig{Provider: "ollama", Model: "nomic-embed-text"},
LLMConfig: domain.LLMConfig{Provider: "ollama", Model: "llama3.2"},
CreatedAt: now,
}); err != nil {
t.Fatal(err)
}
if err := sqliteStore.CreateEntity(context.Background(), &domain.Entity{
ID: "e001", KBID: "kb1", Name: "Alice", Type: "person", Summary: "Engineer", CreatedAt: now, UpdatedAt: now,
}); err != nil {
t.Fatal(err)
}

store := &countingSearchStore{Store: sqliteStore}
ids := make([]string, 0, 200)
scores := make(map[string]float64, 200)
for i := 1; i <= 200; i++ {
id := fmt.Sprintf("missing-%03d", i)
if i == 1 {
id = "e001"
}
ids = append(ids, id)
scores[id] = 1
}

results, err := hydrateEntityResults(context.Background(), store, "kb1", ids, scores, 1)
if err != nil {
t.Fatal(err)
}
if len(results) != 1 {
t.Fatalf("got %d results, want 1", len(results))
}
if store.getEntitiesByIDsCalls != 1 {
t.Fatalf("GetEntitiesByIDs calls = %d, want 1", store.getEntitiesByIDsCalls)
}
if len(store.getEntitiesByIDsBatchSizes) != 1 {
t.Fatalf("batch size captures = %d, want 1", len(store.getEntitiesByIDsBatchSizes))
}
if store.getEntitiesByIDsBatchSizes[0] >= len(ids) {
t.Fatalf("batch fetched %d IDs, expected less than total %d", store.getEntitiesByIDsBatchSizes[0], len(ids))
}
}
115 changes: 85 additions & 30 deletions internal/server/helpers.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,19 @@ package server

import (
"context"
"log/slog"
"time"

"github.com/vndee/memex/internal/domain"
"github.com/vndee/memex/internal/graph"
"github.com/vndee/memex/internal/storage"
)

type subgraphMetadataLoader interface {
GetSubgraphEntitiesByIDs(ctx context.Context, kbID string, ids []string) (map[string]storage.SubgraphEntityMetadata, error)
GetSubgraphRelationsByIDs(ctx context.Context, kbID string, ids []string) (map[string]storage.SubgraphRelationMetadata, error)
}

// buildKB constructs a KnowledgeBase with defaults applied for empty fields.
// Shared by both HTTP and MCP handlers.
func buildKB(id, name, desc, embedProvider, embedModel, llmProvider, llmModel string) *domain.KnowledgeBase {
Expand Down Expand Up @@ -47,46 +53,95 @@ func buildKB(id, name, desc, embedProvider, embedModel, llmProvider, llmModel st
// HydrateSubgraph enriches a raw SubgraphResult with entity and relation metadata
// from storage. Used by MCP, HTTP, and CLI graph traversal handlers.
func HydrateSubgraph(ctx context.Context, store storage.Store, kbID string, sg graph.SubgraphResult) (*domain.Subgraph, error) {
nodeIndex := make(map[string]int, len(sg.Nodes))
nodeIDs := make([]string, 0, len(sg.Nodes))
nodes := make([]domain.SubgraphNode, 0, len(sg.Nodes))
for id, dist := range sg.Nodes {
ent, err := store.GetEntity(ctx, kbID, id)
if err != nil {
nodes = append(nodes, domain.SubgraphNode{
ID: id, Distance: dist,
})
continue
}
for id := range sg.Nodes {
nodeIndex[id] = len(nodes)
nodeIDs = append(nodeIDs, id)
nodes = append(nodes, domain.SubgraphNode{
ID: ent.ID,
Name: ent.Name,
Type: ent.Type,
Summary: ent.Summary,
Distance: dist,
ID: id,
Distance: sg.Nodes[id],
})
}

edgeIndex := make(map[string]int, len(sg.Edges))
edgeIDs := make([]string, 0, len(sg.Edges))
edges := make([]domain.SubgraphEdge, 0, len(sg.Edges))
for _, e := range sg.Edges {
rel, err := store.GetRelation(ctx, kbID, e.RelID)
edgeIndex[e.RelID] = len(edges)
edgeIDs = append(edgeIDs, e.RelID)
edges = append(edges, domain.SubgraphEdge{
ID: e.RelID,
SourceID: e.SourceID,
TargetID: e.TargetID,
Type: e.Type,
Weight: e.Weight,
})
}

if loader, ok := store.(subgraphMetadataLoader); ok {
entitiesByID, err := loader.GetSubgraphEntitiesByIDs(ctx, kbID, nodeIDs)
if err != nil {
slog.Warn("subgraph entity hydration failed", "count", len(nodeIDs), "error", err)
} else {
for id, ent := range entitiesByID {
idx, ok := nodeIndex[id]
if !ok {
continue
}
nodes[idx].ID = ent.ID
nodes[idx].Name = ent.Name
nodes[idx].Type = ent.Type
nodes[idx].Summary = ent.Summary
}
}

relationsByID, err := loader.GetSubgraphRelationsByIDs(ctx, kbID, edgeIDs)
if err != nil {
slog.Warn("subgraph relation hydration failed", "count", len(edgeIDs), "error", err)
} else {
for id, rel := range relationsByID {
idx, ok := edgeIndex[id]
if !ok {
continue
}
edges[idx].ID = rel.ID
edges[idx].SourceID = rel.SourceID
edges[idx].TargetID = rel.TargetID
edges[idx].Type = rel.Type
edges[idx].Weight = rel.Weight
edges[idx].ValidAt = rel.ValidAt
edges[idx].InvalidAt = rel.InvalidAt
}
}

return &domain.Subgraph{Nodes: nodes, Edges: edges}, nil
}

for i := range nodes {
ent, err := store.GetEntity(ctx, kbID, nodes[i].ID)
if err != nil {
edges = append(edges, domain.SubgraphEdge{
ID: e.RelID,
SourceID: e.SourceID,
TargetID: e.TargetID,
Type: e.Type,
Weight: e.Weight,
})
continue
}
edges = append(edges, domain.SubgraphEdge{
ID: rel.ID,
SourceID: rel.SourceID,
TargetID: rel.TargetID,
Type: rel.Type,
Weight: rel.Weight,
ValidAt: rel.ValidAt,
InvalidAt: rel.InvalidAt,
})
nodes[i].ID = ent.ID
nodes[i].Name = ent.Name
nodes[i].Type = ent.Type
nodes[i].Summary = ent.Summary
}

for i := range edges {
rel, err := store.GetRelation(ctx, kbID, edges[i].ID)
if err != nil {
continue
}
edges[i].ID = rel.ID
edges[i].SourceID = rel.SourceID
edges[i].TargetID = rel.TargetID
edges[i].Type = rel.Type
edges[i].Weight = rel.Weight
edges[i].ValidAt = rel.ValidAt
edges[i].InvalidAt = rel.InvalidAt
}

return &domain.Subgraph{Nodes: nodes, Edges: edges}, nil
Expand Down
Loading
Loading