1import { dequal } from "dequal";
2import { toJsxRuntime, type Props, type Options as JsxOptions } from "hast-util-to-jsx-runtime";
3import type { Components } from "rehype-react";
4import type { ReactNode } from "react";
5import { Fragment, jsx, jsxs } from "react/jsx-runtime";
6import type { ElementContent, Node } from "hast";
7
8const tableElements = new Set(["table", "tbody", "thead", "tfoot", "tr"]);
9const htmlWhitespace = /^[\t\n\f\r ]*$/;
10
11/** Opaque immutable state for {@linkcode memoizedHastToReact} */
12export interface RenderState {
13 readonly children: readonly RenderState[];
14 readonly key: string | undefined;
15 readonly node: Node;
16 readonly react: ReactNode;
17 readonly shell: { props: Props; type: React.ElementType } | null;
18}
19
20/**
21 * Converts `hast` (HTML Ast) into React elements by diffing an opaque 'state'
22 * value provided by the previous invocation. This achieves hyper-memoization.
23 */
24export function memoizedHastToReact(
25 tree: Node,
26 state: RenderState | null | undefined,
27 components: Partial<Components> = {},
28): { react: ReactNode; state: RenderState } {
29 const next = recursiveMemoRender(
30 tree,
31 state ?? null,
32 { Fragment, components, jsx, jsxs },
33 undefined,
34 );
35 return { react: next.react, state: next };
36}
37
38function recursiveMemoRender(
39 node: ElementContent | Node,
40 previous: RenderState | null,
41 options: JsxOptions,
42 key: string | undefined,
43): RenderState {
44 if (previous && previous.key === key && node === previous.node) return previous;
45
46 if (!("children" in node) || !Array.isArray(node.children)) {
47 if (previous && previous.key === key && nodeDeepEquals(node, previous.node)) return previous;
48 return {
49 children: [],
50 key,
51 node,
52 react:
53 node.type === "text" && "value" in node && typeof node.value === "string"
54 ? node.value
55 : undefined,
56 shell: null,
57 };
58 }
59
60 const previousChildren = previous?.children ?? [];
61 const nextChildren: RenderState[] = [];
62 const children: ReactNode[] = [];
63 const countsByName = new Map<string, number>();
64 const usedKeys = new Set(
65 previousChildren.map((child) => child.key).filter((key) => key !== undefined),
66 );
67 let prefixLength = 0;
68 for (
69 ;
70 prefixLength < node.children.length && prefixLength < previousChildren.length;
71 prefixLength += 1
72 ) {
73 if (!nodeDeepEquals(node.children[prefixLength], previousChildren[prefixLength]?.node)) break;
74 }
75 let suffixLength = 0;
76 for (
77 ;
78 suffixLength < node.children.length - prefixLength &&
79 suffixLength < previousChildren.length - prefixLength;
80 suffixLength += 1
81 ) {
82 if (
83 !nodeDeepEquals(
84 node.children[node.children.length - 1 - suffixLength],
85 previousChildren[previousChildren.length - 1 - suffixLength]?.node,
86 )
87 )
88 break;
89 }
90
91 for (let i = 0; i < node.children.length; i += 1) {
92 const child = node.children[i]!;
93 let childKey: string | undefined;
94 if (child.type === "element") {
95 const count = countsByName.get(child.tagName) ?? 0;
96 countsByName.set(child.tagName, count + 1);
97 childKey = `${child.tagName}-${count}`;
98 }
99 let previousIndex = -1;
100 if (i < prefixLength || node.children.length === previousChildren.length) previousIndex = i;
101 if (i >= node.children.length - suffixLength) {
102 previousIndex = previousChildren.length - (node.children.length - i);
103 }
104 if (previousIndex >= 0) childKey = previousChildren[previousIndex]?.key ?? childKey;
105 while (childKey && usedKeys.has(childKey) && previousChildren[previousIndex]?.key !== childKey)
106 childKey = `${childKey}+`;
107 if (childKey) usedKeys.add(childKey);
108 const childState = recursiveMemoRender(
109 child,
110 previousChildren[previousIndex] ?? null,
111 options,
112 childKey,
113 );
114 nextChildren.push(childState);
115 if (childState.react !== undefined) children.push(childState.react);
116 }
117
118 const sameSelf = previous !== null && nodeShallowEqual(node, previous.node);
119 if (
120 previous &&
121 previous.key === key &&
122 sameSelf &&
123 previousChildren.length === nextChildren.length &&
124 nextChildren.every((child, index) => child === previousChildren[index])
125 ) {
126 return previous;
127 }
128
129 const filteredChildren =
130 node.type === "element" && tableElements.has(node.tagName)
131 ? children.filter((child) => typeof child !== "string" || !htmlWhitespace.test(child))
132 : children;
133 const reactChildren =
134 filteredChildren.length > 0
135 ? filteredChildren.length === 1
136 ? filteredChildren[0]
137 : filteredChildren
138 : null;
139
140 let shell = sameSelf ? previous?.shell : null;
141 if (!shell) {
142 const base = toJsxRuntime({ ...node, children: [] }, options);
143 shell = { props: base.props, type: base.type };
144 }
145
146 return {
147 children: nextChildren,
148 key,
149 node,
150 react: jsx(shell.type, { ...shell.props, children: reactChildren }, key),
151 shell,
152 };
153}
154
155function nodeShallowEqual<T extends Node & { children?: unknown }>(left: T, right: T) {
156 if (left === right) return true;
157 const { children: _a, position: _b, ...leftProps } = left;
158 const { children: _c, position: _d, ...rightProps } = right;
159 return dequal(leftProps, rightProps);
160}
161
162export function nodeDeepEquals(left: unknown, right: unknown) {
163 if (left === right) return true;
164 if (Array.isArray(left) || Array.isArray(right)) {
165 if (!Array.isArray(left) || !Array.isArray(right)) return false;
166 if (left.length !== right.length) return false;
167 for (let i = 0; i < left.length; i += 1) {
168 if (!nodeDeepEquals(left[i], right[i])) return false;
169 }
170 return true;
171 }
172 if (!left || !right || typeof left !== "object" || typeof right !== "object") return false;
173
174 const leftRecord = left as Record<string, unknown>;
175 const rightRecord = right as Record<string, unknown>;
176 let leftKeyCount = 0;
177 let rightKeyCount = 0;
178
179 for (const key of Object.keys(leftRecord)) {
180 if (key === "position") continue;
181 leftKeyCount += 1;
182 if (!(key in rightRecord)) return false;
183 if (!nodeDeepEquals(leftRecord[key], rightRecord[key])) return false;
184 }
185
186 for (const key of Object.keys(rightRecord)) {
187 if (key === "position") continue;
188 rightKeyCount += 1;
189 }
190
191 return leftKeyCount === rightKeyCount;
192}