diff --git a/packages/memory-graph/README.md b/packages/memory-graph/README.md index 83710f5a..e18c9be8 100644 --- a/packages/memory-graph/README.md +++ b/packages/memory-graph/README.md @@ -53,7 +53,7 @@ function App() { - **Relationship visualization** - Edges show document similarity and memory version chains - **Space filtering** - Filter by workspace or view all memories - **Two variants** - Full-featured console mode or embedded consumer mode -- **Pagination support** - Load more documents on demand +- **Pagination support** - Keep the initial graph fitted as pages arrive; manual pan, zoom, or selection takes control of the view - **TypeScript support** - Full type definitions included ## Essential Props @@ -67,6 +67,12 @@ function App() { | `loadMoreDocuments` | `() => Promise` | Function to load more data | | `highlightDocumentIds` | `string[]` | IDs of documents to highlight | +Console mode uses the supplied theme colors for its surface, 16px dot grid, document icons, and node fills and strokes. Consumer mode keeps its transparent surface and cluster colors. + +Set `colors.dotColor` or the `--graph-dot` CSS variable to style the dot grid independently of text. When neither is set, the grid uses `textMuted`. + +For paginated initial loading, pass `hasMore` and `isLoadingMore` alongside `documents`. In both variants, new batches gradually warm the force layout from the existing node positions. Initial and appended nodes relax until their movement stays low, then cool automatically; a tick limit bounds settling for layouts that keep drifting. The initial view smoothly follows the changing bounds until loading and settling finish. Manual interaction immediately cancels automatic camera movement. Clicking a node selects it without restarting the forces; dragging warms the layout until release, including release outside the canvas. Changing the document selection starts a new fit; the Fit control remains available at any time. The existing static layout safeguard for more than 6,000 nodes remains in place. + ## Documentation Full documentation available at [docs.supermemory.ai](https://docs.supermemory.ai): diff --git a/packages/memory-graph/src/canvas/input-handler.ts b/packages/memory-graph/src/canvas/input-handler.ts index a26c37ff..d3a82f16 100644 --- a/packages/memory-graph/src/canvas/input-handler.ts +++ b/packages/memory-graph/src/canvas/input-handler.ts @@ -23,6 +23,11 @@ export class InputHandler { private posHistory: Array<{ x: number; y: number; t: number }> = [] private draggingNode: GraphNode | null = null + private pressedNode: GraphNode | null = null + private pressX = 0 + private pressY = 0 + private grabOffset = { x: 0, y: 0 } + private gestureWindow: Window | null = null private didDrag = false private currentHoveredId: string | null = null @@ -45,6 +50,8 @@ export class InputHandler { private boundMouseDown: (e: MouseEvent) => void private boundMouseMove: (e: MouseEvent) => void private boundMouseUp: (e: MouseEvent) => void + private boundWindowMove: (e: MouseEvent) => void + private boundBlur: () => void private boundWheel: (e: WheelEvent) => void private boundClick: (e: MouseEvent) => void private boundDblClick: (e: MouseEvent) => void @@ -67,6 +74,10 @@ export class InputHandler { this.boundMouseDown = this.onMouseDown.bind(this) this.boundMouseMove = this.onMouseMove.bind(this) this.boundMouseUp = this.onMouseUp.bind(this) + this.boundWindowMove = (event) => { + if (event.target !== this.canvas) this.onMouseMove(event) + } + this.boundBlur = () => this.endMouseGesture(false) this.boundWheel = this.onWheel.bind(this) this.boundClick = this.onClick.bind(this) this.boundDblClick = this.onDblClick.bind(this) @@ -77,7 +88,6 @@ export class InputHandler { canvas.addEventListener("mousedown", this.boundMouseDown) canvas.addEventListener("mousemove", this.boundMouseMove) - canvas.addEventListener("mouseup", this.boundMouseUp) canvas.addEventListener("click", this.boundClick) canvas.addEventListener("dblclick", this.boundDblClick) canvas.addEventListener("wheel", this.boundWheel, { passive: false }) @@ -98,10 +108,10 @@ export class InputHandler { } destroy(): void { + this.endMouseGesture(false) const c = this.canvas c.removeEventListener("mousedown", this.boundMouseDown) c.removeEventListener("mousemove", this.boundMouseMove) - c.removeEventListener("mouseup", this.boundMouseUp) c.removeEventListener("click", this.boundClick) c.removeEventListener("dblclick", this.boundDblClick) c.removeEventListener("wheel", this.boundWheel) @@ -117,12 +127,32 @@ export class InputHandler { return this.draggingNode } + syncNodes(nodes: Map): void { + if (this.pressedNode) { + this.pressedNode = nodes.get(this.pressedNode.id) ?? null + } + if (!this.draggingNode) return + const current = nodes.get(this.draggingNode.id) + if (!current) { + this.endMouseGesture(false) + } else if (current !== this.draggingNode) { + current.x = current.fx = this.draggingNode.x + current.y = current.fy = this.draggingNode.y + this.draggingNode.fx = null + this.draggingNode.fy = null + this.draggingNode = current + } + } + private canvasXY(e: MouseEvent): { x: number; y: number } { const rect = this.canvas.getBoundingClientRect() return { x: e.clientX - rect.left, y: e.clientY - rect.top } } private onMouseDown(e: MouseEvent): void { + if (e.button !== 0) return + this.endMouseGesture(false) + this.viewport.cancelAnimation() const { x, y } = this.canvasXY(e) const world = this.viewport.screenToWorld(x, y) const node = this.spatialIndex.queryPoint(world.x, world.y) @@ -131,13 +161,16 @@ export class InputHandler { this.lastMouseY = y this.posHistory = [{ x, y, t: performance.now() }] this.didDrag = false + this.pressX = x + this.pressY = y + this.gestureWindow = this.canvas.ownerDocument.defaultView + this.gestureWindow?.addEventListener("mousemove", this.boundWindowMove) + this.gestureWindow?.addEventListener("mouseup", this.boundMouseUp) + this.gestureWindow?.addEventListener("blur", this.boundBlur) if (node) { - this.draggingNode = node - node.fx = node.x - node.fy = node.y - this.callbacks.onNodeDragStart(node.id, node) - this.canvas.style.cursor = "grabbing" + this.pressedNode = node + this.grabOffset = { x: world.x - node.x, y: world.y - node.y } } else { this.isPanning = true this.canvas.style.cursor = "grabbing" @@ -146,13 +179,18 @@ export class InputHandler { private onMouseMove(e: MouseEvent): void { const { x, y } = this.canvasXY(e) + if (this.pressedNode) { + if (Math.hypot(x - this.pressX, y - this.pressY) < 4) return + this.draggingNode = this.pressedNode + this.pressedNode = null + this.callbacks.onNodeDragStart(this.draggingNode.id, this.draggingNode) + this.canvas.style.cursor = "grabbing" + } if (this.draggingNode) { const world = this.viewport.screenToWorld(x, y) - this.draggingNode.fx = world.x - this.draggingNode.fy = world.y - this.draggingNode.x = world.x - this.draggingNode.y = world.y + this.draggingNode.x = this.draggingNode.fx = world.x - this.grabOffset.x + this.draggingNode.y = this.draggingNode.fy = world.y - this.grabOffset.y this.didDrag = true this.callbacks.onRequestRender() return @@ -186,6 +224,15 @@ export class InputHandler { } private onMouseUp(_e: MouseEvent): void { + this.endMouseGesture(true) + } + + private endMouseGesture(withInertia: boolean): void { + this.gestureWindow?.removeEventListener("mousemove", this.boundWindowMove) + this.gestureWindow?.removeEventListener("mouseup", this.boundMouseUp) + this.gestureWindow?.removeEventListener("blur", this.boundBlur) + this.gestureWindow = null + this.pressedNode = null if (this.draggingNode) { this.draggingNode.fx = null this.draggingNode.fy = null @@ -198,7 +245,7 @@ export class InputHandler { if (this.isPanning) { this.isPanning = false - if (this.posHistory.length >= 2) { + if (withInertia && this.posHistory.length >= 2) { const newest = this.posHistory[this.posHistory.length - 1] const oldest = this.posHistory[0] if (!newest || !oldest) return diff --git a/packages/memory-graph/src/canvas/renderer.ts b/packages/memory-graph/src/canvas/renderer.ts index 291c04dd..0ddfba0f 100644 --- a/packages/memory-graph/src/canvas/renderer.ts +++ b/packages/memory-graph/src/canvas/renderer.ts @@ -728,8 +728,18 @@ function drawDocumentNode( sx + half, sy + half, ) - grad.addColorStop(0, mixHexColors(colors.docFill, clusterColor, 0.1)) - grad.addColorStop(1, mixHexColors(colors.docFill, clusterColor, 0.22)) + grad.addColorStop( + 0, + node.clusterColor + ? mixHexColors(colors.docFill, clusterColor, 0.1) + : colors.docFill, + ) + grad.addColorStop( + 1, + node.clusterColor + ? mixHexColors(colors.docFill, clusterColor, 0.22) + : colors.docFill, + ) ctx.fillStyle = grad ctx.strokeStyle = @@ -750,14 +760,23 @@ function drawDocumentNode( const innerSize = size * 0.72 const innerHalf = innerSize * 0.5 const innerR = 6 * (size / 50) - ctx.fillStyle = mixHexColors(colors.docInnerFill, clusterColor, 0.08) + ctx.fillStyle = node.clusterColor + ? mixHexColors(colors.docInnerFill, clusterColor, 0.08) + : colors.docInnerFill roundRect(ctx, sx - innerHalf, sy - innerHalf, innerSize, innerSize, innerR) ctx.fill() const iconSize = size * 0.35 const docType = node.type === "document" ? (node.data as DocumentNodeData).type : "text" - drawDocIcon(ctx, sx, sy, iconSize, docType || "text", clusterColor) + drawDocIcon( + ctx, + sx, + sy, + iconSize, + docType || "text", + node.clusterColor ? clusterColor : colors.iconColor, + ) } function drawMemoryNode( diff --git a/packages/memory-graph/src/canvas/simulation.ts b/packages/memory-graph/src/canvas/simulation.ts index b04cee27..3d424d88 100644 --- a/packages/memory-graph/src/canvas/simulation.ts +++ b/packages/memory-graph/src/canvas/simulation.ts @@ -6,6 +6,7 @@ export const DENSE_GRAPH_STATIC_THRESHOLD = 6000 export class ForceSimulation { private sim: d3.Simulation | null = null + private holdingHeat = false init(nodes: GraphNode[], edges: GraphEdge[]): void { this.destroy() @@ -72,7 +73,7 @@ export class ForceSimulation { if (nodes.length > DENSE_GRAPH_STATIC_THRESHOLD) { this.stop() } else { - this.sim.alphaTarget(0).restart() + this.settle() } } catch (e) { console.error("ForceSimulation.init failed:", e) @@ -89,22 +90,62 @@ export class ForceSimulation { } reheat(): void { + this.holdingHeat = true + this.sim?.on("tick.settle", null) this.sim?.alphaTarget(FORCE_CONFIG.alphaTarget).restart() } + settle(): void { + const sim = this.sim + if (!sim || this.holdingHeat) return + let ticks = 0 + let stableTicks = 0 + sim + .alphaTarget(FORCE_CONFIG.alphaTarget) + .on("tick.settle", () => { + let squaredVelocity = 0 + let maxSquaredVelocity = 0 + const nodes = sim.nodes() + for (const node of nodes) { + const velocity = (node.vx ?? 0) ** 2 + (node.vy ?? 0) ** 2 + squaredVelocity += velocity + maxSquaredVelocity = Math.max(maxSquaredVelocity, velocity) + } + const settled = + sim.alpha() >= FORCE_CONFIG.alphaTarget * 0.9 && + squaredVelocity / Math.max(1, nodes.length) <= + FORCE_CONFIG.settleMeanVelocity ** 2 && + maxSquaredVelocity <= FORCE_CONFIG.settleMaxVelocity ** 2 + stableTicks = settled ? stableTicks + 1 : 0 + if ( + ++ticks >= FORCE_CONFIG.settleMaxTicks || + stableTicks >= FORCE_CONFIG.settleStableTicks + ) { + sim.alphaTarget(0).on("tick.settle", null) + } + }) + .restart() + } + coolDown(): void { - this.sim?.alphaTarget(0) + this.holdingHeat = false + this.sim?.alphaTarget(0).on("tick.settle", null) } stop(): void { - this.sim?.alpha(0).alphaTarget(0).stop() + this.coolDown() + this.sim?.alpha(0).stop() } isActive(): boolean { - return (this.sim?.alpha() ?? 0) > FORCE_CONFIG.alphaMin + return ( + Math.max(this.sim?.alpha() ?? 0, this.sim?.alphaTarget() ?? 0) > + FORCE_CONFIG.alphaMin + ) } destroy(): void { + this.holdingHeat = false if (this.sim) { this.sim.stop() this.sim = null diff --git a/packages/memory-graph/src/canvas/viewport.ts b/packages/memory-graph/src/canvas/viewport.ts index c33a1681..648b6ce8 100644 --- a/packages/memory-graph/src/canvas/viewport.ts +++ b/packages/memory-graph/src/canvas/viewport.ts @@ -43,12 +43,21 @@ export class ViewportState { } pan(dx: number, dy: number): void { + this.cancelAnimation() this.panX += dx this.panY += dy this.targetPanX = null this.targetPanY = null } + cancelAnimation(): void { + this.velocityX = 0 + this.velocityY = 0 + this.targetZoom = this.zoom + this.targetPanX = null + this.targetPanY = null + } + releaseWithVelocity(vx: number, vy: number): void { this.velocityX = vx this.velocityY = vy @@ -72,6 +81,7 @@ export class ViewportState { nodes: Array<{ x: number; y: number; size: number }>, width: number, height: number, + { animate = true }: { animate?: boolean } = {}, ): void { const fit = computeFit(nodes, width, height) if (!fit) return @@ -83,6 +93,12 @@ export class ViewportState { this.zoomAnchorY = height / 2 this.targetPanX = width / 2 - cx * this.targetZoom this.targetPanY = height / 2 - cy * this.targetZoom + if (!animate) { + this.zoom = this.targetZoom + this.panX = this.targetPanX + this.panY = this.targetPanY + this.cancelAnimation() + } } setMinZoomForNodes( @@ -99,7 +115,6 @@ export class ViewportState { ViewportState.ABSOLUTE_MIN_ZOOM, ViewportState.DEFAULT_MIN_ZOOM, ) - this.zoom = clamp(this.zoom, this.minZoom, ViewportState.MAX_ZOOM) this.targetZoom = clamp( this.targetZoom, this.minZoom, @@ -132,7 +147,7 @@ export class ViewportState { } const zoomDiff = this.targetZoom - this.zoom - if (Math.abs(zoomDiff) > 0.001) { + if (Math.abs(zoomDiff) > this.zoom * 0.001) { const world = this.screenToWorld(this.zoomAnchorX, this.zoomAnchorY) this.zoom += zoomDiff * this.zoomSpring this.panX = this.zoomAnchorX - world.x * this.zoom diff --git a/packages/memory-graph/src/components/graph-canvas.tsx b/packages/memory-graph/src/components/graph-canvas.tsx index c5aea910..ca070b9b 100644 --- a/packages/memory-graph/src/components/graph-canvas.tsx +++ b/packages/memory-graph/src/components/graph-canvas.tsx @@ -99,6 +99,7 @@ export const GraphCanvas = memo(function GraphCanvas({ const map = nodeMapRef.current map.clear() for (const n of nodes) map.set(n.id, n) + inputRef.current?.syncNodes(map) spatialRef.current.rebuild(nodes) renderNeeded.current = true }, [nodes]) diff --git a/packages/memory-graph/src/components/legend.tsx b/packages/memory-graph/src/components/legend.tsx index e9be3a22..a133ce4f 100644 --- a/packages/memory-graph/src/components/legend.tsx +++ b/packages/memory-graph/src/components/legend.tsx @@ -320,6 +320,7 @@ export const Legend = memo(function Legend({ .filter(Boolean), ).size const updateNodeCount = nodes.filter(isUpdateMemoryNode).length + const hasClusterColors = nodes.some((node) => node.clusterColor) const outerStyle: React.CSSProperties = { overflow: "hidden", @@ -571,22 +572,26 @@ export const Legend = memo(function Legend({
- Color -
- -
- Cluster - - Same document or connected memory group - + + {hasClusterColors ? "Color" : "Clusters"} + + {hasClusterColors && ( +
+ +
+ Cluster + + Same document or connected memory group + +
-
+ )}
Visible clusters {clusterCount} diff --git a/packages/memory-graph/src/components/memory-graph.tsx b/packages/memory-graph/src/components/memory-graph.tsx index d165ac71..016d9428 100644 --- a/packages/memory-graph/src/components/memory-graph.tsx +++ b/packages/memory-graph/src/components/memory-graph.tsx @@ -1,4 +1,11 @@ -import { useCallback, useEffect, useMemo, useRef, useState } from "react" +import { + useCallback, + useEffect, + useLayoutEffect, + useMemo, + useRef, + useState, +} from "react" import { DENSE_GRAPH_STATIC_THRESHOLD, ForceSimulation, @@ -8,6 +15,7 @@ import { VersionChainIndex } from "../canvas/version-chain" import type { ViewportState } from "../canvas/viewport" import { useGraphData } from "../hooks/use-graph-data" import { useGraphTheme } from "../hooks/use-graph-theme" +import { useInitialGraphFit } from "../hooks/use-initial-graph-fit" import type { GraphApiDocument, GraphThemeColors, @@ -116,6 +124,7 @@ export function MemoryGraph({ containerSize.width, containerSize.height, colors, + variant === "console" ? "theme" : "cluster", ) const isCompactViewport = containerSize.width > 0 && containerSize.width < 640 const graphFitHeight = isCompactViewport @@ -131,10 +140,8 @@ export function MemoryGraph({ // that makes this a no-op on re-renders where limitedDocuments hasn't changed. chainIndex.current.rebuild(limitedDocuments) - // Initial loads get a full force settle. Append-only pagination keeps - // existing coordinates stable and renders new nodes in nearby open areas. const prevSimIdsRef = useRef>(new Set()) - useEffect(() => { + useLayoutEffect(() => { if (nodes.length === 0) { simulationRef.current?.destroy() simulationRef.current = null @@ -172,7 +179,7 @@ export function MemoryGraph({ } else if (idsChanged && isAppendOnly) { prevSimIdsRef.current = currentIds simulationRef.current.update(nodes, edges) - simulationRef.current.stop() + simulationRef.current.settle() } else { simulationRef.current.update(nodes, edges) } @@ -195,81 +202,17 @@ export function MemoryGraph({ ) }, [nodes, containerSize.width, graphFitHeight]) - // 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 && - hasContainerSize - ) { - const fitDelays = isCompactViewport ? [100, 450, 900] : [100, 300] - const timers = fitDelays.map((delay, index) => - setTimeout(() => { - if (!viewportRef.current || !hasContainerSize) return - viewportRef.current.fitToNodes( - nodes, - containerSize.width, - graphFitHeight, - ) - if (index === fitDelays.length - 1) { - hasAutoFittedRef.current = true - } - }, delay), - ) - return () => { - for (const timer of timers) clearTimeout(timer) - } - } - }, [ + const stopFollowing = useInitialGraphFit({ + documents: limitedDocuments, 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, - containerSize.width, - graphFitHeight, - ) - }, 120) - return () => clearTimeout(timer) - }, [ - isCompactViewport, - nodes, - containerSize.width, - graphFitHeight, - hasContainerSize, - ]) - - useEffect(() => { - if (nodes.length === 0) hasAutoFittedRef.current = false - }, [nodes.length]) - - useEffect(() => { - if (isCompactViewport) { - hasAutoFittedRef.current = false - } - }, [isCompactViewport]) - - useEffect(() => { - if (hasContainerSize && !hadValidContainerSizeRef.current) { - hadValidContainerSizeRef.current = true - hasAutoFittedRef.current = false - } - if (!hasContainerSize) { - hadValidContainerSizeRef.current = false - } - }, [hasContainerSize]) + viewportRef, + simulationRef, + width: containerSize.width, + height: hasContainerSize ? graphFitHeight : 0, + isLoading: externalIsLoading, + isLoadingMore, + hasMore, + }) // Container resize observer useEffect(() => { @@ -375,11 +318,18 @@ export function MemoryGraph({ // Navigation const handleAutoFit = useCallback(() => { if (nodes.length === 0 || !viewportRef.current) return + stopFollowing() + viewportRef.current.setMinZoomForNodes( + nodes, + containerSize.width, + graphFitHeight, + ) viewportRef.current.fitToNodes(nodes, containerSize.width, graphFitHeight) - }, [nodes, containerSize.width, graphFitHeight]) + }, [nodes, containerSize.width, graphFitHeight, stopFollowing]) const handleCenter = useCallback(() => { if (nodes.length === 0 || !viewportRef.current) return + stopFollowing() let sx = 0 let sy = 0 for (const n of nodes) { @@ -392,19 +342,21 @@ export function MemoryGraph({ containerSize.width, graphFitHeight, ) - }, [nodes, containerSize.width, graphFitHeight]) + }, [nodes, containerSize.width, graphFitHeight, stopFollowing]) const handleZoomIn = useCallback(() => { const vp = viewportRef.current if (!vp) return + stopFollowing() vp.zoomTo(vp.zoom * 1.3, containerSize.width / 2, graphFitHeight / 2) - }, [containerSize.width, graphFitHeight]) + }, [containerSize.width, graphFitHeight, stopFollowing]) const handleZoomOut = useCallback(() => { const vp = viewportRef.current if (!vp) return + stopFollowing() vp.zoomTo(vp.zoom / 1.3, containerSize.width / 2, graphFitHeight / 2) - }, [containerSize.width, graphFitHeight]) + }, [containerSize.width, graphFitHeight, stopFollowing]) // Wrap onOpenDocument to dismiss the popover before opening the modal. // Without this, the popover overlay stays mounted on top of the @@ -460,6 +412,7 @@ export function MemoryGraph({ // Arrow key navigation through nodes const selectAndCenter = useCallback( (nodeId: string) => { + stopFollowing() setSelectedNode(nodeId) const n = nodes.find((nd) => nd.id === nodeId) if (n && viewportRef.current) @@ -470,7 +423,7 @@ export function MemoryGraph({ graphFitHeight, ) }, - [nodes, containerSize.width, graphFitHeight], + [nodes, containerSize.width, graphFitHeight, stopFollowing], ) const navigateUp = useCallback(() => { @@ -613,7 +566,6 @@ export function MemoryGraph({ if (!isSlideshowActive || nodes.length === 0) { if (!isSlideshowActive) { setSelectedNode(null) - simulationRef.current?.coolDown() } return } @@ -653,6 +605,7 @@ export function MemoryGraph({ return () => { clearInterval(interval) if (coolDownTimer) clearTimeout(coolDownTimer) + simulationRef.current?.coolDown() } }, [isSlideshowActive, nodes.length]) @@ -692,7 +645,7 @@ export function MemoryGraph({ display: "flex", alignItems: "center", justifyContent: "center", - backgroundColor: "transparent", + backgroundColor: variant === "console" ? colors.bg : "transparent", borderRadius: 12, } @@ -718,7 +671,12 @@ export function MemoryGraph({ height: "100%", borderRadius: 12, overflow: "hidden", - backgroundColor: "transparent", + backgroundColor: variant === "console" ? colors.bg : "transparent", + backgroundImage: + variant === "console" + ? `radial-gradient(circle, ${colors.dotColor ?? colors.textMuted} 0.5px, transparent 0.5px)` + : undefined, + backgroundSize: variant === "console" ? "16px 16px" : undefined, } const canvasContainerStyle: React.CSSProperties = { @@ -767,7 +725,16 @@ export function MemoryGraph({
{children}
)} -
+
{ + if (event.target instanceof HTMLCanvasElement) stopFollowing() + }} + onWheelCapture={(event) => { + if (event.target instanceof HTMLCanvasElement) stopFollowing() + }} + > {hasContainerSize && ( >(new Map()) @@ -510,6 +511,7 @@ export function useGraphData( for (let docIdx = 0; docIdx < docCount; docIdx++) { const doc = documents[docIdx] const docCluster = getDocumentClusterAssignment(doc, clusterAssignments) + const docColor = colorMode === "cluster" ? docCluster.color : null const angle = docIdx * goldenAngle const radius = spiralScale * Math.sqrt((docIdx + 1) / docCount) const initialX = cx + Math.cos(angle) * radius @@ -531,9 +533,9 @@ export function useGraphData( docNode = { ...previousDocNode, data: docData, - borderColor: docCluster.color, + borderColor: docColor ?? colors.docStroke, clusterKey: docCluster.key, - clusterColor: docCluster.color, + clusterColor: docColor, isDragging: draggingNodeId === doc.id, } } else { @@ -556,9 +558,9 @@ export function useGraphData( y: appendPosition?.y ?? initialY, data: docData, size: 50, - borderColor: docCluster.color, + borderColor: docColor ?? colors.docStroke, clusterKey: docCluster.key, - clusterColor: docCluster.color, + clusterColor: docColor, isHovered: false, isDragging: false, } @@ -581,15 +583,16 @@ export function useGraphData( content: mem.memory, } const cluster = clusterAssignments.get(mem.id) + const memoryColor = colorMode === "cluster" ? cluster?.color : undefined let memNode: GraphNode if (previousMemNode) { memNode = { ...previousMemNode, data: memData, - borderColor: getMemoryNodeBorderColor(mem, colors, cluster?.color), + borderColor: getMemoryNodeBorderColor(mem, colors, memoryColor), clusterKey: cluster?.key ?? null, - clusterColor: cluster?.color ?? null, + clusterColor: memoryColor ?? null, isDragging: draggingNodeId === mem.id, } } else { @@ -601,9 +604,9 @@ export function useGraphData( y: docNode.y + memOffset.y, data: memData, size: 36, - borderColor: getMemoryNodeBorderColor(mem, colors, cluster?.color), + borderColor: getMemoryNodeBorderColor(mem, colors, memoryColor), clusterKey: cluster?.key ?? null, - clusterColor: cluster?.color ?? null, + clusterColor: memoryColor ?? null, isHovered: false, isDragging: false, } @@ -618,7 +621,7 @@ export function useGraphData( } return { nodes: result, cache: nextCache } - }, [documents, canvasWidth, canvasHeight, draggingNodeId, colors]) + }, [documents, canvasWidth, canvasHeight, draggingNodeId, colors, colorMode]) useEffect(() => { nodeCache.current = graphData.cache diff --git a/packages/memory-graph/src/hooks/use-graph-theme.ts b/packages/memory-graph/src/hooks/use-graph-theme.ts index f045830a..fc7ae913 100644 --- a/packages/memory-graph/src/hooks/use-graph-theme.ts +++ b/packages/memory-graph/src/hooks/use-graph-theme.ts @@ -13,6 +13,7 @@ function readCssVar(name: string, fallback: string): string { function resolveColors(): GraphThemeColors { return { bg: readCssVar("--graph-bg", DEFAULT_COLORS.bg), + dotColor: readCssVar("--graph-dot", "") || undefined, docFill: readCssVar("--graph-doc-fill", DEFAULT_COLORS.docFill), docStroke: readCssVar("--graph-doc-stroke", DEFAULT_COLORS.docStroke), docInnerFill: readCssVar("--graph-doc-inner", DEFAULT_COLORS.docInnerFill), diff --git a/packages/memory-graph/src/hooks/use-initial-graph-fit.ts b/packages/memory-graph/src/hooks/use-initial-graph-fit.ts new file mode 100644 index 00000000..378c3641 --- /dev/null +++ b/packages/memory-graph/src/hooks/use-initial-graph-fit.ts @@ -0,0 +1,101 @@ +import { useCallback, useLayoutEffect, useRef, type RefObject } from "react" +import type { ViewportState } from "../canvas/viewport" +import type { ForceSimulation } from "../canvas/simulation" +import type { GraphApiDocument, GraphNode } from "../types" + +export function useInitialGraphFit({ + documents, + nodes, + viewportRef, + simulationRef, + width, + height, + isLoading, + isLoadingMore, + hasMore, +}: { + documents: GraphApiDocument[] + nodes: GraphNode[] + viewportRef: RefObject + simulationRef?: RefObject + width: number + height: number + isLoading: boolean + isLoadingMore: boolean + hasMore: boolean +}) { + const session = useRef({ + documentIds: new Set(), + following: true, + userMoved: false, + fitted: false, + width: 0, + height: 0, + }) + + const stopFollowing = useCallback(() => { + session.current.following = false + session.current.userMoved = true + viewportRef.current?.cancelAnimation() + }, [viewportRef]) + + useLayoutEffect(() => { + const current = session.current + const documentIds = new Set(documents.map((document) => document.id)) + const replaced = [...current.documentIds].some((id) => !documentIds.has(id)) + if (nodes.length === 0 || replaced) { + current.following = true + current.userMoved = false + current.fitted = false + } + if ( + (width !== current.width || height !== current.height) && + !current.userMoved + ) { + current.following = true + } + current.documentIds = documentIds + current.width = width + current.height = height + if ( + !current.following || + isLoading || + nodes.length === 0 || + width <= 0 || + height <= 0 + ) + return + + let timer: ReturnType | undefined + const fit = (animate = true) => { + const viewport = viewportRef.current + if (!viewport || !session.current.following) return + viewport.setMinZoomForNodes(nodes, width, height) + viewport.fitToNodes(nodes, width, height, { animate }) + if (simulationRef?.current?.isActive()) { + timer = setTimeout(fit, 200) + } else if (!hasMore && !isLoadingMore) { + session.current.following = false + } + } + if (!current.fitted) { + current.fitted = true + fit(false) + } else { + timer = setTimeout(fit, 100) + } + return () => clearTimeout(timer) + }, [ + documents, + nodes, + viewportRef, + simulationRef, + width, + height, + isLoading, + isLoadingMore, + hasMore, + ]) + + return stopFollowing +} diff --git a/packages/memory-graph/src/types.ts b/packages/memory-graph/src/types.ts index f871b3da..eb89347d 100644 --- a/packages/memory-graph/src/types.ts +++ b/packages/memory-graph/src/types.ts @@ -105,6 +105,7 @@ export interface GraphEdge { export interface GraphThemeColors { bg: string + dotColor?: string docFill: string docStroke: string docInnerFill: string