diff --git a/apps/web/components/memory-graph/hooks/use-graph-api.ts b/apps/web/components/memory-graph/hooks/use-graph-api.ts
index 7991f600..dd295d45 100644
--- a/apps/web/components/memory-graph/hooks/use-graph-api.ts
+++ b/apps/web/components/memory-graph/hooks/use-graph-api.ts
@@ -1,7 +1,7 @@
"use client"
import { useInfiniteQuery } from "@tanstack/react-query"
-import { useMemo } from "react"
+import { useEffect, useMemo } from "react"
import { $fetch } from "@lib/api"
import type {
GraphApiDocument,
@@ -15,6 +15,7 @@ interface UseGraphApiOptions {
containerTags?: string[]
documentIds?: string[]
enabled?: boolean
+ maxNodes?: number
}
interface ApiMemoryEntry {
@@ -59,6 +60,13 @@ interface ApiDocumentsResponse {
}
}
+function getGraphNodeCount(documents: ApiDocument[]): number {
+ return documents.reduce(
+ (total, doc) => total + 1 + (doc.memoryEntries?.length ?? 0),
+ 0,
+ )
+}
+
function toGraphMemory(mem: ApiMemoryEntry): GraphApiMemory {
return {
id: mem.id,
@@ -108,7 +116,7 @@ function toGraphDocument(
}
export function useGraphApi(options: UseGraphApiOptions = {}) {
- const { containerTags, documentIds, enabled = true } = options
+ const { containerTags, documentIds, enabled = true, maxNodes } = options
const filteredDocumentIds = documentIds?.filter(Boolean)
const hasDocumentIds =
filteredDocumentIds != null && filteredDocumentIds.length > 0
@@ -126,6 +134,7 @@ export function useGraphApi(options: UseGraphApiOptions = {}) {
containerTags,
[],
filteredDocumentIds,
+ maxNodes,
],
initialPageParam: 1,
queryFn: async ({ pageParam }) => {
@@ -155,7 +164,16 @@ export function useGraphApi(options: UseGraphApiOptions = {}) {
return response.data as unknown as ApiDocumentsResponse
},
- getNextPageParam: (lastPage) => {
+ getNextPageParam: (lastPage, allPages) => {
+ if (hasDocumentIds) return undefined
+ if (maxNodes != null) {
+ const loadedNodes = allPages.reduce(
+ (total, page) => total + getGraphNodeCount(page.documents ?? []),
+ 0,
+ )
+ if (loadedNodes >= maxNodes) return undefined
+ }
+
const { currentPage, totalPages } = lastPage.pagination
return currentPage < totalPages ? currentPage + 1 : undefined
},
@@ -163,6 +181,29 @@ export function useGraphApi(options: UseGraphApiOptions = {}) {
enabled,
})
+ const loadedNodeCount = useMemo(() => {
+ if (!data?.pages) return 0
+ return data.pages.reduce(
+ (total, page) => total + getGraphNodeCount(page.documents ?? []),
+ 0,
+ )
+ }, [data])
+
+ useEffect(() => {
+ if (!enabled || hasDocumentIds || maxNodes == null) return
+ if (!hasNextPage || isFetchingNextPage || loadedNodeCount >= maxNodes)
+ return
+ fetchNextPage()
+ }, [
+ enabled,
+ hasDocumentIds,
+ hasNextPage,
+ isFetchingNextPage,
+ loadedNodeCount,
+ maxNodes,
+ fetchNextPage,
+ ])
+
const documents = useMemo(() => {
if (!data?.pages) return []
return data.pages.flatMap((page) =>
diff --git a/apps/web/components/memory-graph/memory-graph-wrapper.tsx b/apps/web/components/memory-graph/memory-graph-wrapper.tsx
index 0b2ca826..f594004f 100644
--- a/apps/web/components/memory-graph/memory-graph-wrapper.tsx
+++ b/apps/web/components/memory-graph/memory-graph-wrapper.tsx
@@ -1,6 +1,5 @@
"use client"
-import { useEffect, useRef, useState } from "react"
import { MemoryGraph as MemoryGraphBase } from "@supermemory/memory-graph"
import type { GraphThemeColors } from "@supermemory/memory-graph"
import { useGraphApi } from "./hooks/use-graph-api"
@@ -34,20 +33,6 @@ export function MemoryGraph({
canvasRef,
...rest
}: MemoryGraphWrapperProps) {
- const [containerSize, setContainerSize] = useState({ width: 0, height: 0 })
- const containerRef = useRef
(null)
-
- useEffect(() => {
- const el = containerRef.current
- if (!el) return
- const ro = new ResizeObserver(() => {
- setContainerSize({ width: el.clientWidth, height: el.clientHeight })
- })
- ro.observe(el)
- setContainerSize({ width: el.clientWidth, height: el.clientHeight })
- return () => ro.disconnect()
- }, [])
-
const {
documents,
isLoading: apiIsLoading,
@@ -59,11 +44,11 @@ export function MemoryGraph({
} = useGraphApi({
containerTags,
documentIds,
- enabled: containerSize.width > 0 && containerSize.height > 0,
+ maxNodes,
})
return (
-
+
{
if (!maxNodes || documents.length === 0) return documents
let totalNodes = 0
- let cutoff = documents.length
+ const limited: GraphApiDocument[] = []
for (let i = 0; i < documents.length; i++) {
- const docNodes = 1 + (documents[i]?.memories?.length ?? 0)
- if (totalNodes + docNodes > maxNodes) {
- cutoff = i
+ const doc = documents[i]
+ if (!doc) continue
+ if (totalNodes >= maxNodes) break
+
+ const remainingNodes = maxNodes - totalNodes
+ const memories = doc.memories ?? []
+ const docNodes = 1 + memories.length
+
+ if (docNodes <= remainingNodes) {
+ limited.push(doc)
+ totalNodes += docNodes
+ continue
+ }
+
+ if (remainingNodes > 1) {
+ limited.push({
+ ...doc,
+ memories: memories.slice(0, remainingNodes - 1),
+ })
+ totalNodes = maxNodes
break
}
- totalNodes += docNodes
+
+ limited.push({ ...doc, memories: [] })
+ totalNodes += 1
}
- return cutoff === documents.length ? documents : documents.slice(0, cutoff)
+ return limited
}, [documents, maxNodes])
+ const hasContainerSize = containerSize.width > 0 && containerSize.height > 0
+
const { nodes, edges } = useGraphData(
- limitedDocuments,
+ hasContainerSize ? limitedDocuments : [],
null,
containerSize.width,
containerSize.height,
@@ -149,17 +174,19 @@ export function MemoryGraph({
// Auto-fit when data first loads. Mobile needs a few passes because the
// force simulation can move nodes after the first layout frame.
const hasAutoFittedRef = useRef(false)
+ const hadValidContainerSizeRef = useRef(false)
useEffect(() => {
if (
!hasAutoFittedRef.current &&
nodes.length > 0 &&
viewportRef.current &&
- containerSize.width > 0
+ hasContainerSize
) {
- const fitDelays = isCompactViewport ? [100, 450, 900] : [100]
+ const fitDelays = isCompactViewport ? [100, 450, 900] : [100, 300]
const timers = fitDelays.map((delay, index) =>
setTimeout(() => {
- viewportRef.current?.fitToNodes(
+ if (!viewportRef.current || !hasContainerSize) return
+ viewportRef.current.fitToNodes(
nodes,
containerSize.width,
graphFitHeight,
@@ -173,10 +200,17 @@ export function MemoryGraph({
for (const timer of timers) clearTimeout(timer)
}
}
- }, [nodes, containerSize.width, graphFitHeight, isCompactViewport])
+ }, [
+ nodes,
+ containerSize.width,
+ graphFitHeight,
+ isCompactViewport,
+ hasContainerSize,
+ ])
useEffect(() => {
if (!isCompactViewport || nodes.length === 0 || !viewportRef.current) return
+ if (!hasContainerSize) return
const timer = setTimeout(() => {
viewportRef.current?.fitToNodes(
nodes,
@@ -185,7 +219,13 @@ export function MemoryGraph({
)
}, 120)
return () => clearTimeout(timer)
- }, [isCompactViewport, nodes, containerSize.width, graphFitHeight])
+ }, [
+ isCompactViewport,
+ nodes,
+ containerSize.width,
+ graphFitHeight,
+ hasContainerSize,
+ ])
useEffect(() => {
if (nodes.length === 0) hasAutoFittedRef.current = false
@@ -197,20 +237,39 @@ export function MemoryGraph({
}
}, [isCompactViewport])
+ useEffect(() => {
+ if (hasContainerSize && !hadValidContainerSizeRef.current) {
+ hadValidContainerSizeRef.current = true
+ hasAutoFittedRef.current = false
+ }
+ if (!hasContainerSize) {
+ hadValidContainerSizeRef.current = false
+ }
+ }, [hasContainerSize])
+
// Container resize observer
useEffect(() => {
const el = containerRef.current
if (!el) return
- const ro = new ResizeObserver(() => {
- setContainerSize({ width: el.clientWidth, height: el.clientHeight })
- setContainerBounds(el.getBoundingClientRect())
- })
- ro.observe(el)
- setContainerSize({ width: el.clientWidth, height: el.clientHeight })
- setContainerBounds(el.getBoundingClientRect())
+ const measure = () => {
+ const rect = el.getBoundingClientRect()
+ const width = Math.round(rect.width) || el.clientWidth
+ const height = Math.round(rect.height) || el.clientHeight
+ setContainerSize({ width, height })
+ setContainerBounds(rect)
+ }
- return () => ro.disconnect()
+ const ro = new ResizeObserver(measure)
+ ro.observe(el)
+ const parent = el.parentElement
+ if (parent) ro.observe(parent)
+ measure()
+ const raf = requestAnimationFrame(measure)
+ return () => {
+ cancelAnimationFrame(raf)
+ ro.disconnect()
+ }
}, [])
// Callbacks for GraphCanvas
@@ -600,7 +659,8 @@ export function MemoryGraph({
return chainIndex.current.getChain(activeNodeData.id)
}, [activeNodeData, limitedDocuments])
- const isLoading = externalIsLoading
+ const isLayoutPending = !hasContainerSize && limitedDocuments.length > 0
+ const isLoading = externalIsLoading || isLayoutPending
if (externalError) {
const errorContainerStyle: React.CSSProperties = {
@@ -682,7 +742,7 @@ export function MemoryGraph({
)}
- {containerSize.width > 0 && containerSize.height > 0 && (
+ {hasContainerSize && (
>
+
export function getAppendPosition(
existingNodes: GraphNode[],
appendIndex: number,
canvasWidth: number,
canvasHeight: number,
+ baseBounds?: NodeBounds | null,
+ spatialGrid?: Map,
) {
- const bounds = getNodeBounds(existingNodes)
+ const bounds = baseBounds ?? getNodeBounds(existingNodes)
if (!bounds) {
return { x: canvasWidth / 2, y: canvasHeight / 2 }
}
- const spatialGrid = buildAppendSpatialGrid(existingNodes)
+ const candidateGrid = spatialGrid ?? buildAppendSpatialGrid(existingNodes)
const boundsWidth = bounds.maxX - bounds.minX
const boundsHeight = bounds.maxY - bounds.minY
const baseRadiusX = boundsWidth / 2 + APPEND_CLUSTER_RADIUS + APPEND_AREA_GAP
@@ -323,7 +327,7 @@ export function getAppendPosition(
y: bounds.centerY + Math.sin(angle) * radiusY,
}
- if (isAppendCandidateOpen(candidate, spatialGrid)) {
+ if (isAppendCandidateOpen(candidate, candidateGrid)) {
return candidate
}
}
@@ -366,20 +370,27 @@ function isAppendCandidateOpen(
function buildAppendSpatialGrid(nodes: GraphNode[]): Map {
const grid = new Map()
for (const node of nodes) {
- const key = getAppendSpatialKey(
- getAppendSpatialCell(node.x),
- getAppendSpatialCell(node.y),
- )
- const bucket = grid.get(key)
- if (bucket) {
- bucket.push(node)
- } else {
- grid.set(key, [node])
- }
+ addAppendSpatialNode(grid, node)
}
return grid
}
+function addAppendSpatialNode(
+ grid: Map,
+ node: GraphNode,
+): void {
+ const key = getAppendSpatialKey(
+ getAppendSpatialCell(node.x),
+ getAppendSpatialCell(node.y),
+ )
+ const bucket = grid.get(key)
+ if (bucket) {
+ bucket.push(node)
+ } else {
+ grid.set(key, [node])
+ }
+}
+
function getAppendSpatialCell(value: number): number {
return Math.floor(value / APPEND_SPATIAL_CELL_SIZE)
}
@@ -472,6 +483,14 @@ export function useGraphData(
const appendPlacementNodes = Array.from(previousCache.values()).filter(
(node) => currentIds.has(node.id),
)
+ const shouldAppendNewNodes = appendPlacementNodes.length > 0
+ const appendBaseBounds = shouldAppendNewNodes
+ ? getNodeBounds(appendPlacementNodes)
+ : null
+ const appendSpatialGrid =
+ shouldAppendNewNodes && appendBaseBounds
+ ? buildAppendSpatialGrid(appendPlacementNodes)
+ : null
let appendIndex = 0
const clusterAssignments = computeClusterAssignments(documents)
@@ -519,12 +538,14 @@ export function useGraphData(
}
} else {
const appendPosition =
- appendPlacementNodes.length > 0
+ shouldAppendNewNodes && appendBaseBounds && appendSpatialGrid
? getAppendPosition(
appendPlacementNodes,
appendIndex++,
canvasWidth,
canvasHeight,
+ appendBaseBounds,
+ appendSpatialGrid,
)
: null
@@ -541,7 +562,10 @@ export function useGraphData(
isHovered: false,
isDragging: false,
}
- appendPlacementNodes.push(docNode)
+ if (appendSpatialGrid) {
+ appendPlacementNodes.push(docNode)
+ addAppendSpatialNode(appendSpatialGrid, docNode)
+ }
}
nextCache.set(doc.id, docNode)
result.push(docNode)
@@ -583,7 +607,10 @@ export function useGraphData(
isHovered: false,
isDragging: false,
}
- appendPlacementNodes.push(memNode)
+ if (appendSpatialGrid) {
+ appendPlacementNodes.push(memNode)
+ addAppendSpatialNode(appendSpatialGrid, memNode)
+ }
}
nextCache.set(mem.id, memNode)
result.push(memNode)