diff --git a/src/computed.test.ts b/src/computed.test.ts index 9e1729b..ff6a702 100644 --- a/src/computed.test.ts +++ b/src/computed.test.ts @@ -1,36 +1,36 @@ -import { describe, expect, test, beforeEach, mock } from "bun:test"; -import { type StateCreator, create } from "zustand"; -import { type ComputedStateOpts, createComputed } from "./computed"; +import { describe, expect, test, beforeEach, mock } from "bun:test" +import { type StateCreator, create } from "zustand" +import { type ComputedStateOpts, createComputed } from "./computed" type Store = { - count: number; - x: number; - y: number; - inc: () => void; - dec: () => void; -}; + count: number + x: number + y: number + inc: () => void + dec: () => void +} type ComputedStore = { - countSq: number; + countSq: number nestedResult: { - stringified: string; - }; -}; + stringified: string + } +} function computeState(state: Store): ComputedStore { const nestedResult = { stringified: JSON.stringify(state.count), - }; + } return { countSq: state.count ** 2, nestedResult, - }; + } } describe("default config", () => { - const computeStateMock = mock(computeState); - const computed = createComputed(computeStateMock); + const computeStateMock = mock(computeState) + const computed = createComputed(computeStateMock) const makeStore = () => create( computed((set) => ({ @@ -40,54 +40,54 @@ describe("default config", () => { inc: () => set((state) => ({ count: state.count + 1 })), dec: () => set((state) => ({ count: state.count - 1 })), })), - ); + ) - let useStore: ReturnType; + let useStore: ReturnType beforeEach(() => { - computeStateMock.mockClear(); - useStore = makeStore(); - }); + computeStateMock.mockClear() + useStore = makeStore() + }) test("computed works on simple counter example", () => { // note: this function should have been called once on store creation - expect(computeStateMock).toHaveBeenCalledTimes(1); - expect(useStore.getState().count).toEqual(1); - expect(useStore.getState().countSq).toEqual(1); - useStore.getState().inc(); - expect(useStore.getState().count).toEqual(2); - expect(useStore.getState().countSq).toEqual(4); - useStore.getState().dec(); - expect(useStore.getState().count).toEqual(1); - expect(useStore.getState().countSq).toEqual(1); - useStore.setState({ count: 4 }); - expect(useStore.getState().countSq).toEqual(16); - expect(computeStateMock).toHaveBeenCalledTimes(4); - }); + expect(computeStateMock).toHaveBeenCalledTimes(1) + expect(useStore.getState().count).toEqual(1) + expect(useStore.getState().countSq).toEqual(1) + useStore.getState().inc() + expect(useStore.getState().count).toEqual(2) + expect(useStore.getState().countSq).toEqual(4) + useStore.getState().dec() + expect(useStore.getState().count).toEqual(1) + expect(useStore.getState().countSq).toEqual(1) + useStore.setState({ count: 4 }) + expect(useStore.getState().countSq).toEqual(16) + expect(computeStateMock).toHaveBeenCalledTimes(4) + }) test("computed does not modify object ref even after change", () => { - useStore.setState({ count: 4 }); - expect(useStore.getState().count).toEqual(4); - const obj = useStore.getState().nestedResult; - useStore.setState({ count: 4 }); - const toCompare = useStore.getState().nestedResult; - expect(obj).toEqual(toCompare); - }); + useStore.setState({ count: 4 }) + expect(useStore.getState().count).toEqual(4) + const obj = useStore.getState().nestedResult + useStore.setState({ count: 4 }) + const toCompare = useStore.getState().nestedResult + 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", () => { - expect(computeStateMock).toHaveBeenCalledTimes(1); - useStore.setState({ x: 2 }); - expect(computeStateMock).toHaveBeenCalledTimes(2); - useStore.setState({ x: 3 }); - expect(computeStateMock).toHaveBeenCalledTimes(2); - useStore.setState({ y: 2 }); - expect(computeStateMock).toHaveBeenCalledTimes(2); - }); -}); + expect(computeStateMock).toHaveBeenCalledTimes(1) + useStore.setState({ x: 2 }) + expect(computeStateMock).toHaveBeenCalledTimes(2) + useStore.setState({ x: 3 }) + expect(computeStateMock).toHaveBeenCalledTimes(2) + useStore.setState({ y: 2 }) + expect(computeStateMock).toHaveBeenCalledTimes(2) + }) +}) describe("custom config", () => { - const computeStateMock = mock(computeState); + const computeStateMock = mock(computeState) const makeStore = (opts?: ComputedStateOpts) => { - const computed = createComputed(computeStateMock, opts); + const computed = createComputed(computeStateMock, opts) return create( computed((set) => ({ count: 1, @@ -96,56 +96,56 @@ describe("custom config", () => { inc: () => set((state) => ({ count: state.count + 1 })), dec: () => set((state) => ({ count: state.count - 1 })), })), - ); - }; + ) + } beforeEach(() => { - computeStateMock.mockClear(); - }); + computeStateMock.mockClear() + }) test("computed does not update when a custom key selector is given", () => { - const useStore = makeStore({ keys: ["x", "y"] }); + const useStore = makeStore({ keys: ["x", "y"] }) // because we only care about x and y, the compute function should not be called when count changes - expect(computeStateMock).toHaveBeenCalledTimes(1); - expect(useStore.getState().count).toEqual(1); - expect(useStore.getState().countSq).toEqual(1); - useStore.getState().inc(); - expect(useStore.getState().count).toEqual(2); - expect(useStore.getState().countSq).toEqual(1); - useStore.getState().dec(); - expect(useStore.getState().count).toEqual(1); - expect(useStore.getState().countSq).toEqual(1); - expect(computeStateMock).toHaveBeenCalledTimes(1); - }); + expect(computeStateMock).toHaveBeenCalledTimes(1) + expect(useStore.getState().count).toEqual(1) + expect(useStore.getState().countSq).toEqual(1) + useStore.getState().inc() + expect(useStore.getState().count).toEqual(2) + expect(useStore.getState().countSq).toEqual(1) + useStore.getState().dec() + expect(useStore.getState().count).toEqual(1) + expect(useStore.getState().countSq).toEqual(1) + expect(computeStateMock).toHaveBeenCalledTimes(1) + }) test("disabling proxy causes compute to run every time", () => { - const useStore = makeStore({ disableProxy: true }); - expect(computeStateMock).toHaveBeenCalledTimes(1); - useStore.setState({ count: 4 }); - useStore.setState({ x: 2 }); - useStore.setState({ y: 3 }); - expect(useStore.getState().count).toEqual(4); - expect(useStore.getState().countSq).toEqual(16); - expect(computeStateMock).toHaveBeenCalledTimes(4); - }); -}); - -type CountSlice = Pick; -type XYSlice = Pick; + const useStore = makeStore({ disableProxy: true }) + expect(computeStateMock).toHaveBeenCalledTimes(1) + useStore.setState({ count: 4 }) + useStore.setState({ x: 2 }) + useStore.setState({ y: 3 }) + expect(useStore.getState().count).toEqual(4) + expect(useStore.getState().countSq).toEqual(16) + expect(computeStateMock).toHaveBeenCalledTimes(4) + }) +}) + +type CountSlice = Pick +type XYSlice = Pick function computeSlice(state: CountSlice): ComputedStore { const nestedResult = { stringified: JSON.stringify(state.count), - }; + } return { countSq: state.count ** 2, nestedResult, - }; + } } describe("slices pattern", () => { - const computeSliceMock = mock(computeSlice); - const computed = createComputed(computeSliceMock); + const computeSliceMock = mock(computeSlice) + const computed = createComputed(computeSliceMock) const makeStore = () => { const createCountSlice: StateCreator< Store, @@ -155,40 +155,40 @@ describe("slices pattern", () => { > = computed((set) => ({ count: 1, dec: () => set((state) => ({ count: state.count - 1 })), - })); + })) const createXySlice: StateCreator = (set) => ({ x: 1, y: 1, // this should not trigger compute function inc: () => set((state) => ({ count: state.count + 2 })), - }); + }) return create()((...a) => ({ ...createCountSlice(...a), ...createXySlice(...a), - })); - }; + })) + } beforeEach(() => { - computeSliceMock.mockClear(); - }); + computeSliceMock.mockClear() + }) test("computed works on slices pattern example", () => { - const useStore = makeStore(); - expect(computeSliceMock).toHaveBeenCalledTimes(1); - expect(useStore.getState().count).toEqual(1); - expect(useStore.getState().countSq).toEqual(1); - useStore.getState().inc(); - expect(useStore.getState().count).toEqual(3); - expect(useStore.getState().countSq).toEqual(1); - expect(computeSliceMock).toHaveBeenCalledTimes(1); - useStore.getState().dec(); - expect(useStore.getState().count).toEqual(2); - expect(useStore.getState().countSq).toEqual(4); - expect(computeSliceMock).toHaveBeenCalledTimes(2); - useStore.setState({ count: 4 }); - expect(useStore.getState().countSq).toEqual(16); - expect(computeSliceMock).toHaveBeenCalledTimes(3); - }); -}); + const useStore = makeStore() + expect(computeSliceMock).toHaveBeenCalledTimes(1) + expect(useStore.getState().count).toEqual(1) + expect(useStore.getState().countSq).toEqual(1) + useStore.getState().inc() + expect(useStore.getState().count).toEqual(3) + expect(useStore.getState().countSq).toEqual(1) + expect(computeSliceMock).toHaveBeenCalledTimes(1) + useStore.getState().dec() + expect(useStore.getState().count).toEqual(2) + expect(useStore.getState().countSq).toEqual(4) + expect(computeSliceMock).toHaveBeenCalledTimes(2) + useStore.setState({ count: 4 }) + expect(useStore.getState().countSq).toEqual(16) + expect(computeSliceMock).toHaveBeenCalledTimes(3) + }) +}) diff --git a/src/computed.ts b/src/computed.ts index 9675070..3ef7272 100644 --- a/src/computed.ts +++ b/src/computed.ts @@ -1,16 +1,24 @@ -import type { - Mutate, - StateCreator, - StoreApi, - StoreMutatorIdentifier, -} from "zustand"; -import { shallow } from "zustand/shallow"; +import type { Mutate, StateCreator, StoreApi, StoreMutatorIdentifier } from "zustand" +import { shallow } from "zustand/shallow" +/** + * Options for when and how your compute function is called. + */ export type ComputedStateOpts = { - keys?: (keyof T)[]; - disableProxy?: boolean; - equalityFn?: (a: Y, b: Y) => boolean; -}; + /** + * An explicit list of keys to track for recomputation. + */ + keys?: (keyof T)[] + /** + * Disable the use of Proxy for tracking. + */ + disableProxy?: boolean + /** + * Custom equality function for comparing computed values. By default, we use + * `zustand/shallow` to compare the stor + */ + equalityFn?: (a: Y, b: Y) => boolean +} export type ComputedStateCreator = ( compute: (state: T) => A, @@ -21,114 +29,108 @@ export type ComputedStateCreator = ( U = T, >( f: StateCreator, -) => StateCreator; +) => StateCreator -type Cast = T extends U ? T : U; -type Write = Omit & U; +type Cast = T extends U ? T : U +type Write = Omit & U type StoreCompute = S extends { - getState: () => infer T; + getState: () => infer T } ? Omit, "setState"> - : never; -type WithCompute = Write>; + : never +type WithCompute = Write> declare module "zustand/vanilla" { interface StoreMutators { - "chrisvander/zustand-computed": WithCompute, A>; + "chrisvander/zustand-computed": WithCompute, A> } } type ComputedStateImpl = ( compute: (state: T) => A, opts?: ComputedStateOpts, -) => (f: StateCreator) => StateCreator; +) => (f: StateCreator) => StateCreator -type SetStateWithArgs = Parameters< - ReturnType> ->[0] extends (...args: infer U) => void +type SetStateWithArgs = Parameters>>[0] extends (...args: infer U) => void ? (...args: [...U, ...unknown[]]) => void - : never; + : never const computedImpl: ComputedStateImpl = (compute, opts) => (f) => { - // set of keys that have been accessed in any compute call - const trackedSelectors = new Set(); + // 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 equalityFn = opts?.equalityFn ?? shallow if (opts?.keys) { - const selectorKeys = opts.keys; + const selectorKeys = opts.keys for (const key of selectorKeys) { - trackedSelectors.add(key); + trackedSelectors.add(key) } } - // we track which selectors are accessed - const useSelectors = opts?.disableProxy !== true || !!opts?.keys; - const useProxy = opts?.disableProxy !== true && !opts?.keys; + // Determine if selectors or proxy should be used. + const useSelectors = opts?.disableProxy !== true || !!opts?.keys + const useProxy = opts?.disableProxy !== true && !opts?.keys + const computeAndMerge = (state: T | (T & A)): T & A => { - // create a Proxy to track which selectors are accessed - const stateProxy = new Proxy( - { ...state }, - { - get: (_, prop) => { - trackedSelectors.add(prop); - return state[prop as keyof T]; + // 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] + }, }, - }, - ); + ) - // calculate the new computed state - const computedState: A = compute(useProxy ? stateProxy : { ...state }); + // Calculate the new computed state. + const computedState: A = compute(useProxy ? createStateProxy() : { ...state }) - // if part of the computed state did not change according to the equalityFn - // then we use the object ref from the previous state. This is to prevent - // unnecessary re-renders. + // If part of the computed state did not change according to the equalityFn, + // then delete that key from the newly calculated computed state. for (const k of Object.keys(computedState) as (keyof A)[]) { if (equalityFn(computedState[k], (state as T & A)[k])) { - computedState[k] = (state as T & A)[k]; + delete computedState[k] } } - return { ...state, ...computedState }; - }; + return { ...state, ...computedState } + } - // higher level function to handle compute & compare overhead - const setWithComputed = ( - update: T | ((state: T) => T), - replace?: boolean, - ...args: unknown[] - ) => { - (set as SetStateWithArgs)( + /** + * Higher level function to handle compute & compare overhead. + */ + const setWithComputed = (update: T | ((state: T) => T), replace?: boolean, ...args: unknown[]) => { + ;(set as SetStateWithArgs)( (state: T): T & A => { - const updated = typeof update === "object" ? update : update(state); + 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 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 } - return computeAndMerge({ ...state, ...updated }); + return computeAndMerge({ ...state, ...updated }) }, replace, ...args, - ); - }; - - const _api = api as Mutate< - StoreApi, - [["chrisvander/zustand-computed", A]] - >; - _api.setState = setWithComputed; - const st = f(setWithComputed, get, _api) as T & A; - return Object.assign({}, st, compute(st)); - }; -}; - -export const createComputed = computedImpl as unknown as ComputedStateCreator; + ) + } + + const _api = api as Mutate, [["chrisvander/zustand-computed", A]]> + _api.setState = setWithComputed + const st = f(setWithComputed, get, _api) as T & A + return Object.assign({}, st, compute(st)) + } +} + +export const createComputed = computedImpl as unknown as ComputedStateCreator diff --git a/src/index.ts b/src/index.ts index 6c6e96f..e0fe5e4 100644 --- a/src/index.ts +++ b/src/index.ts @@ -1 +1 @@ -export * from "./computed"; +export * from "./computed"