diff --git a/src/computed.test.ts b/src/computed.test.ts index ff6a702..2ae98c0 100644 --- a/src/computed.test.ts +++ b/src/computed.test.ts @@ -1,4 +1,4 @@ -import { describe, expect, test, beforeEach, mock } from "bun:test" +import { beforeEach, describe, expect, mock, test } from "bun:test" import { type StateCreator, create } from "zustand" import { type ComputedStateOpts, createComputed } from "./computed" @@ -28,11 +28,11 @@ function computeState(state: Store): ComputedStore { } } -describe("default config", () => { +describe("single store", () => { const computeStateMock = mock(computeState) - const computed = createComputed(computeStateMock) - const makeStore = () => - create( + const makeStore = (opts?: ComputedStateOpts) => { + const computed = createComputed(computeStateMock, opts) + return create( computed((set) => ({ count: 1, x: 1, @@ -41,14 +41,14 @@ describe("default config", () => { dec: () => set((state) => ({ count: state.count - 1 })), })), ) + } - let useStore: ReturnType beforeEach(() => { computeStateMock.mockClear() - useStore = makeStore() }) test("computed works on simple counter example", () => { + const useStore = makeStore() // note: this function should have been called once on store creation expect(computeStateMock).toHaveBeenCalledTimes(1) expect(useStore.getState().count).toEqual(1) @@ -65,6 +65,7 @@ describe("default config", () => { }) test("computed does not modify object ref even after change", () => { + const useStore = makeStore() useStore.setState({ count: 4 }) expect(useStore.getState().count).toEqual(4) const obj = useStore.getState().nestedResult @@ -73,34 +74,37 @@ describe("default config", () => { expect(obj).toEqual(toCompare) }) - test("modifying variables x and y do not trigger compute function more than once, as they are not used in compute function", () => { + test("any store change, by default, triggers compute function", () => { + const useStore = makeStore() expect(computeStateMock).toHaveBeenCalledTimes(1) useStore.setState({ x: 2 }) expect(computeStateMock).toHaveBeenCalledTimes(2) useStore.setState({ x: 3 }) - expect(computeStateMock).toHaveBeenCalledTimes(2) + expect(computeStateMock).toHaveBeenCalledTimes(3) useStore.setState({ y: 2 }) - expect(computeStateMock).toHaveBeenCalledTimes(2) + expect(computeStateMock).toHaveBeenCalledTimes(4) }) -}) -describe("custom config", () => { - const computeStateMock = mock(computeState) - const makeStore = (opts?: ComputedStateOpts) => { - const computed = createComputed(computeStateMock, opts) - return create( - computed((set) => ({ - count: 1, - x: 1, - y: 1, - inc: () => set((state) => ({ count: state.count + 1 })), - dec: () => set((state) => ({ count: state.count - 1 })), - })), - ) - } + test("modifying variables x and y do not trigger compute function when `keys` are specified", () => { + const useStore = makeStore({ keys: ["count"] }) + expect(computeStateMock).toHaveBeenCalledTimes(1) + useStore.setState({ x: 2 }) + expect(computeStateMock).toHaveBeenCalledTimes(1) + useStore.setState({ x: 3 }) + expect(computeStateMock).toHaveBeenCalledTimes(1) + useStore.setState({ y: 2 }) + expect(computeStateMock).toHaveBeenCalledTimes(1) + }) - beforeEach(() => { - computeStateMock.mockClear() + test("modifying variables x and y do not trigger compute function when `shouldRecompute` is defined", () => { + const useStore = makeStore({ shouldRecompute: (_, nextState) => "count" in nextState }) + expect(computeStateMock).toHaveBeenCalledTimes(1) + useStore.setState({ x: 2 }) + expect(computeStateMock).toHaveBeenCalledTimes(1) + useStore.setState({ x: 3 }) + expect(computeStateMock).toHaveBeenCalledTimes(1) + useStore.setState({ y: 2 }) + expect(computeStateMock).toHaveBeenCalledTimes(1) }) test("computed does not update when a custom key selector is given", () => { diff --git a/src/computed.ts b/src/computed.ts index 3ef7272..8c280bd 100644 --- a/src/computed.ts +++ b/src/computed.ts @@ -4,18 +4,43 @@ import { shallow } from "zustand/shallow" /** * Options for when and how your compute function is called. */ -export type ComputedStateOpts = { - /** - * An explicit list of keys to track for recomputation. - */ - keys?: (keyof T)[] +export type ComputedStateOpts = ( + | { + /** + * An explicit list of keys to track for recomputation. By default, + * `zustand-computed` will run your compute function on any change. + * This lets you filter those keys out. It's better to use the + * `compareFn` to be more explicit about how comparison is determined. + */ + keys?: (keyof T)[] + } + | { + /** + * Custom comparison function to determine whether to recompute. + * Receives the previous and next store state, should return true if + * compute should run, false to skip recomputation. This function + * should be *fast* - it determines whether or not you need to + * recompute. + */ + shouldRecompute?: (state: T, nextState: T) => boolean + } +) & { /** - * Disable the use of Proxy for tracking. + * @deprecated removed proxy; this does nothing and will be removed. */ disableProxy?: boolean /** * Custom equality function for comparing computed values. By default, we use - * `zustand/shallow` to compare the stor + * `zustand/shallow` to compare each key in your store against the newly- + * computed values. This is likely the desired method of comparison. + * + * The motivation for this function is to ensure that, in the case your + * computed function returns a value that is identical structurally, it + * should not cause a re-render despite the reference being different. + * + * You can disable comparison, so that the most recent result of the + * compute function always triggers downstream re-renders, by simply + * returning false. */ equalityFn?: (a: Y, b: Y) => boolean } @@ -56,40 +81,27 @@ type SetStateWithArgs = Parameters>>[0] : never const computedImpl: ComputedStateImpl = (compute, opts) => (f) => { - // Set of keys that have been accessed in any compute call. - const trackedSelectors = new Set() - return (set, get, api) => { - type T = ReturnType - type A = ReturnType + type T = ReturnType + type A = ReturnType - const equalityFn = opts?.equalityFn ?? shallow + const optsKeys = !opts || !("keys" in opts) || opts.keys == null ? undefined : opts.keys + const keysSet = optsKeys ? new Set(optsKeys as string[]) : undefined - if (opts?.keys) { - const selectorKeys = opts.keys - for (const key of selectorKeys) { - trackedSelectors.add(key) - } - } + function defaultShouldRecomputeFn(_: T, nextState: T): boolean { + if (!keysSet || nextState == null) return true + return Object.keys(nextState).some((k) => keysSet.has(k)) + } - // Determine if selectors or proxy should be used. - const useSelectors = opts?.disableProxy !== true || !!opts?.keys - const useProxy = opts?.disableProxy !== true && !opts?.keys + const shouldRecomputeFn = + opts && "shouldRecompute" in opts ? (opts.shouldRecompute ?? defaultShouldRecomputeFn) : defaultShouldRecomputeFn - const computeAndMerge = (state: T | (T & A)): T & A => { - // Create a Proxy to track which selectors are accessed. - const createStateProxy = () => - new Proxy( - { ...state }, - { - get: (_, prop) => { - trackedSelectors.add(prop) - return state[prop as keyof T] - }, - }, - ) + // Set of keys that have been accessed in any compute call. + return (set, get, api) => { + const equalityFn = opts?.equalityFn ?? shallow + const computeAndMerge = (state: T | (T & A)): T & A => { // Calculate the new computed state. - const computedState: A = compute(useProxy ? createStateProxy() : { ...state }) + const computedState: A = compute({ ...state }) // If part of the computed state did not change according to the equalityFn, // then delete that key from the newly calculated computed state. @@ -109,16 +121,7 @@ const computedImpl: ComputedStateImpl = (compute, opts) => (f) => { ;(set as SetStateWithArgs)( (state: T): T & A => { const updated = typeof update === "object" ? update : update(state) - - if ( - useSelectors && - trackedSelectors.size !== 0 && - !Object.keys(updated).some((k) => trackedSelectors.has(k)) - ) { - // If we have a selector set, but none of the updated keys are in the selector set, then we can skip the compute. - return { ...state, ...updated } as T & A - } - + if (!shouldRecomputeFn?.(state, updated)) return { ...state, ...updated } as T & A return computeAndMerge({ ...state, ...updated }) }, replace,