diff --git a/compiler/forget/packages/playground/components/Editor/index.tsx b/compiler/forget/packages/playground/components/Editor/index.tsx index 1d19de18f3..5714bd8720 100644 --- a/compiler/forget/packages/playground/components/Editor/index.tsx +++ b/compiler/forget/packages/playground/components/Editor/index.tsx @@ -74,8 +74,6 @@ const COMMON_HOOKS: Array<[string, Hook]> = [ [ "useFragment", { - name: "useFragment", - kind: "Custom", valueKind: ValueKind.Frozen, effectKind: Effect.Freeze, }, @@ -83,8 +81,6 @@ const COMMON_HOOKS: Array<[string, Hook]> = [ [ "usePaginationFragment", { - name: "usePaginationFragment", - kind: "Custom", valueKind: ValueKind.Frozen, effectKind: Effect.Freeze, }, @@ -92,8 +88,6 @@ const COMMON_HOOKS: Array<[string, Hook]> = [ [ "useRefetchableFragment", { - name: "useRefetchableFragment", - kind: "Custom", valueKind: ValueKind.Frozen, effectKind: Effect.Freeze, }, @@ -101,8 +95,6 @@ const COMMON_HOOKS: Array<[string, Hook]> = [ [ "useLazyLoadQuery", { - name: "useLazyLoadQuery", - kind: "Custom", valueKind: ValueKind.Frozen, effectKind: Effect.Freeze, }, @@ -110,8 +102,6 @@ const COMMON_HOOKS: Array<[string, Hook]> = [ [ "usePreloadedQuery", { - name: "usePreloadedQuery", - kind: "Custom", valueKind: ValueKind.Frozen, effectKind: Effect.Freeze, }, diff --git a/compiler/forget/packages/snap/src/compiler-worker.ts b/compiler/forget/packages/snap/src/compiler-worker.ts index 7f3bb45bf8..b70f784693 100644 --- a/compiler/forget/packages/snap/src/compiler-worker.ts +++ b/compiler/forget/packages/snap/src/compiler-worker.ts @@ -134,8 +134,6 @@ export async function compile( [ "useFreeze", { - name: "useFreeze", - kind: "Custom", valueKind: "frozen", effectKind: "freeze", }, diff --git a/compiler/forget/src/HIR/Environment.ts b/compiler/forget/src/HIR/Environment.ts index e60e9e4b66..1a78bde8d0 100644 --- a/compiler/forget/src/HIR/Environment.ts +++ b/compiler/forget/src/HIR/Environment.ts @@ -17,15 +17,27 @@ import { import { BlockId, BuiltInType, + Effect, FunctionType, IdentifierId, ObjectType, PolyType, + ValueKind, makeBlockId, makeIdentifierId, } from "./HIR"; -import { Hook } from "./Hooks"; -import { FunctionSignature, ShapeRegistry } from "./ObjectShape"; +import { + DefaultMutatingHook, + DefaultNonmutatingHook, + FunctionSignature, + ShapeRegistry, + addHook, +} from "./ObjectShape"; + +export type Hook = { + effectKind: Effect; + valueKind: ValueKind; +}; // TODO(mofeiZ): User defined global types (with corresponding shapes). // User defined global types should have inline ObjectShapes instead of directly @@ -107,7 +119,7 @@ export class Environment { config: EnvironmentConfig | null, contextIdentifiers: Set ) { - this.#shapes = DEFAULT_SHAPES; + this.#shapes = new Map(DEFAULT_SHAPES); if (config?.customHooks) { this.#globals = new Map(DEFAULT_GLOBALS); @@ -116,10 +128,17 @@ export class Environment { !this.#globals.has(hookName), `[Globals] Found existing definition in global registry for custom hook ${hookName}` ); - this.#globals.set(hookName, { - kind: "Hook", - definition: hook, - }); + this.#globals.set( + hookName, + addHook(this.#shapes, [], { + positionalParams: [], + restParam: hook.effectKind, + returnType: { kind: "Poly" }, + returnValueKind: hook.valueKind, + calleeEffect: Effect.Read, + hookKind: "Custom", + }) + ); } } else { this.#globals = DEFAULT_GLOBALS; @@ -150,10 +169,11 @@ export class Environment { if (resolvedGlobal === null) { // Hack, since we don't track module level declarations and imports if (isHookName(name)) { - return { - kind: "Hook", - definition: null, - }; + if (this.enableAssumeHooksFollowRulesOfReact) { + return DefaultNonmutatingHook; + } else { + return DefaultMutatingHook; + } } else { log(() => `Undefined global '${name}'`); } diff --git a/compiler/forget/src/HIR/Globals.ts b/compiler/forget/src/HIR/Globals.ts index 12e17b0f56..16b39c4769 100644 --- a/compiler/forget/src/HIR/Globals.ts +++ b/compiler/forget/src/HIR/Globals.ts @@ -6,15 +6,15 @@ */ import { Effect, ValueKind } from "./HIR"; -import { Hook } from "./Hooks"; import { BUILTIN_SHAPES, BuiltInArrayId, ShapeRegistry, addFunction, + addHook, addObject, } from "./ObjectShape"; -import { BuiltInType, HookType, PolyType } from "./Types"; +import { BuiltInType, FunctionType, PolyType } from "./Types"; /** * This file exports types and defaults for JavaScript global objects. @@ -229,85 +229,91 @@ const TYPED_GLOBALS: Array<[string, BuiltInType]> = [ // TODO: rest of Global objects ]; -const BUILTIN_HOOKS: Array<[string, Hook]> = [ +// TODO(mofeiZ): We currently only store rest param effects for hooks +// until FeatureFlag `enableTreatHooksAsFunctions` is removed +const BUILTIN_HOOKS: Array<[string, FunctionType]> = [ [ "useContext", - { - kind: "State", - name: "useContext", - effectKind: Effect.Read, - valueKind: ValueKind.Mutable, - }, + addHook(DEFAULT_SHAPES, [], { + positionalParams: [], + restParam: Effect.Read, + returnType: { kind: "Poly" }, + calleeEffect: Effect.Read, + hookKind: "useContext", + returnValueKind: ValueKind.Mutable, + }), ], [ "useState", - { - kind: "State", - name: "useState", - effectKind: Effect.Freeze, - valueKind: ValueKind.Frozen, - }, + addHook(DEFAULT_SHAPES, [], { + positionalParams: [], + restParam: Effect.Freeze, + returnType: { kind: "Poly" }, + calleeEffect: Effect.Read, + hookKind: "useState", + returnValueKind: ValueKind.Frozen, + }), ], [ "useRef", - { - kind: "Ref", - name: "useRef", - effectKind: Effect.Capture, - valueKind: ValueKind.Mutable, - }, + addHook(DEFAULT_SHAPES, [], { + positionalParams: [], + restParam: Effect.Capture, + returnType: { kind: "Poly" }, + calleeEffect: Effect.Read, + hookKind: "useRef", + returnValueKind: ValueKind.Mutable, + }), ], [ "useMemo", - { - kind: "Memo", - name: "useMemo", - effectKind: Effect.Freeze, - valueKind: ValueKind.Frozen, - }, + addHook(DEFAULT_SHAPES, [], { + positionalParams: [], + restParam: Effect.Freeze, + returnType: { kind: "Poly" }, + calleeEffect: Effect.Read, + hookKind: "useMemo", + returnValueKind: ValueKind.Frozen, + }), ], [ "useCallback", - { - kind: "Memo", - name: "useCallback", - effectKind: Effect.Freeze, - valueKind: ValueKind.Frozen, - }, + addHook(DEFAULT_SHAPES, [], { + positionalParams: [], + restParam: Effect.Freeze, + returnType: { kind: "Poly" }, + calleeEffect: Effect.Read, + hookKind: "useCallback", + returnValueKind: ValueKind.Frozen, + }), ], [ "useEffect", - { - kind: "Memo", - name: "useEffect", - effectKind: Effect.Freeze, - valueKind: ValueKind.Frozen, - }, + addHook(DEFAULT_SHAPES, [], { + positionalParams: [], + restParam: Effect.Freeze, + returnType: { kind: "Poly" }, + calleeEffect: Effect.Read, + hookKind: "useEffect", + returnValueKind: ValueKind.Frozen, + }), ], [ "useLayoutEffect", - { - kind: "Memo", - name: "useLayoutEffect", - effectKind: Effect.Freeze, - valueKind: ValueKind.Frozen, - }, + addHook(DEFAULT_SHAPES, [], { + positionalParams: [], + restParam: Effect.Freeze, + returnType: { kind: "Poly" }, + calleeEffect: Effect.Read, + hookKind: "useLayoutEffect", + returnValueKind: ValueKind.Frozen, + }), ], ]; -export type Global = BuiltInType | HookType | PolyType; +export type Global = BuiltInType | PolyType; export type GlobalRegistry = Map; -export const DEFAULT_GLOBALS: GlobalRegistry = new Map( - BUILTIN_HOOKS.map(([hookName, hook]) => { - return [ - hookName, - { - kind: "Hook", - definition: hook, - }, - ]; - }) -); +export const DEFAULT_GLOBALS: GlobalRegistry = new Map(BUILTIN_HOOKS); // Hack until we add ObjectShapes for all globals for (const name of UNTYPED_GLOBALS) { diff --git a/compiler/forget/src/HIR/HIR.ts b/compiler/forget/src/HIR/HIR.ts index 59a3988932..ec23776565 100644 --- a/compiler/forget/src/HIR/HIR.ts +++ b/compiler/forget/src/HIR/HIR.ts @@ -8,6 +8,7 @@ import * as t from "@babel/types"; import invariant from "invariant"; import { Environment } from "./Environment"; +import { HookKind } from "./ObjectShape"; import { Type } from "./Types"; // ******************************************************************************************* @@ -951,8 +952,13 @@ export function isPrimitiveType(id: Identifier): boolean { return id.type.kind === "Primitive"; } -export function isHookType(id: Identifier): boolean { - return id.type.kind === "Hook"; +export function getHookKind(env: Environment, id: Identifier): HookKind | null { + const idType = id.type; + if (idType.kind === "Function") { + const signature = env.getFunctionSignature(idType); + return signature?.hookKind ?? null; + } + return null; } export * from "./Types"; diff --git a/compiler/forget/src/HIR/Hooks.ts b/compiler/forget/src/HIR/Hooks.ts deleted file mode 100644 index b8bb0936c2..0000000000 --- a/compiler/forget/src/HIR/Hooks.ts +++ /dev/null @@ -1,16 +0,0 @@ -/** - * Copyright (c) Meta Platforms, Inc. and affiliates. - * - * This source code is licensed under the MIT license found in the - * LICENSE file in the root directory of this source tree. - */ - -import { Effect, ValueKind } from "./HIR"; - -export type HookKind = "State" | "Ref" | "Custom" | "Memo"; -export type Hook = { - kind: HookKind; - name: string; - effectKind: Effect; - valueKind: ValueKind; -}; diff --git a/compiler/forget/src/HIR/ObjectShape.ts b/compiler/forget/src/HIR/ObjectShape.ts index 873140f75c..0ba10a973a 100644 --- a/compiler/forget/src/HIR/ObjectShape.ts +++ b/compiler/forget/src/HIR/ObjectShape.ts @@ -33,14 +33,36 @@ function createAnonId(): string { } /** - * Add a function to an existing ShapeRegistry. + * Add a non-hook function to an existing ShapeRegistry. * * @returns a {@link FunctionType} representing the added function. */ export function addFunction( registry: ShapeRegistry, properties: Iterable<[string, BuiltInType | PolyType]>, - fn: FunctionSignature + fn: Omit +): FunctionType { + const shapeId = createAnonId(); + addShape(registry, shapeId, properties, { + ...fn, + hookKind: null, + }); + return { + kind: "Function", + return: fn.returnType, + shapeId, + }; +} + +/** + * Add a hook to an existing ShapeRegistry. + * + * @returns a {@link FunctionType} representing the added hook function. + */ +export function addHook( + registry: ShapeRegistry, + properties: Iterable<[string, BuiltInType | PolyType]>, + fn: FunctionSignature & { hookKind: HookKind } ): FunctionType { const shapeId = createAnonId(); addShape(registry, shapeId, properties, fn); @@ -88,6 +110,16 @@ function addShape( return shape; } +export type HookKind = + | "useContext" + | "useState" + | "useRef" + | "useEffect" + | "useLayoutEffect" + | "useMemo" + | "useCallback" + | "Custom"; + /** * Call signature of a function, used for type and effect inference. * @@ -102,6 +134,7 @@ export type FunctionSignature = { returnType: BuiltInType | PolyType; returnValueKind: ValueKind; calleeEffect: Effect; + hookKind: HookKind | null; }; /** @@ -184,3 +217,21 @@ addObject(BUILTIN_SHAPES, BuiltInObjectId, [ // TODO: // hasOwnProperty, isPrototypeOf, propertyIsEnumerable, toLocaleString, valueOf ]); + +export const DefaultMutatingHook = addHook(BUILTIN_SHAPES, [], { + positionalParams: [], + restParam: Effect.Mutate, + returnType: { kind: "Poly" }, + calleeEffect: Effect.Read, + hookKind: "Custom", + returnValueKind: ValueKind.Mutable, +}); + +export const DefaultNonmutatingHook = addHook(BUILTIN_SHAPES, [], { + positionalParams: [], + restParam: Effect.Freeze, + returnType: { kind: "Poly" }, + calleeEffect: Effect.Read, + hookKind: "Custom", + returnValueKind: ValueKind.Frozen, +}); diff --git a/compiler/forget/src/HIR/Types.ts b/compiler/forget/src/HIR/Types.ts index 14083d107f..8d26868bdc 100644 --- a/compiler/forget/src/HIR/Types.ts +++ b/compiler/forget/src/HIR/Types.ts @@ -6,22 +6,11 @@ */ import invariant from "invariant"; -import { Hook } from "./Hooks"; export type BuiltInType = PrimitiveType | FunctionType | ObjectType; -export type Type = - | BuiltInType - | HookType - | PhiType - | TypeVar - | PolyType - | PropType; +export type Type = BuiltInType | PhiType | TypeVar | PolyType | PropType; export type PrimitiveType = { kind: "Primitive" }; -export type HookType = { - kind: "Hook"; - definition: Hook | null; -}; /** * An {@link FunctionType} or {@link ObjectType} (also a JS object) may be associated with an @@ -93,7 +82,6 @@ export function typeEquals(tA: Type, tB: Type): boolean { return ( typeVarEquals(tA, tB) || funcTypeEquals(tA, tB) || - hookTypeEquals(tA, tB) || objectTypeEquals(tA, tB) || primitiveTypeEquals(tA, tB) || polyTypeEquals(tA, tB) || @@ -131,12 +119,6 @@ function funcTypeEquals(tA: Type, tB: Type): boolean { return typeEquals(tA.return, tB.return); } -function hookTypeEquals(tA: Type, tB: Type): boolean { - return ( - tA.kind === "Hook" && tB.kind === "Hook" && tA.definition === tB.definition - ); -} - function phiTypeEquals(tA: Type, tB: Type): boolean { if (tA.kind === "Phi" && tB.kind === "Phi") { if (tA.operands.length !== tB.operands.length) { diff --git a/compiler/forget/src/HIR/ValidateHooksUsage.ts b/compiler/forget/src/HIR/ValidateHooksUsage.ts index f1134e9edc..c1a588d126 100644 --- a/compiler/forget/src/HIR/ValidateHooksUsage.ts +++ b/compiler/forget/src/HIR/ValidateHooksUsage.ts @@ -11,7 +11,7 @@ import { ErrorSeverity, } from "../CompilerError"; import { hasBackEdge } from "../Optimization/DeadCodeElimination"; -import { HIRFunction, IdentifierId, Place, isHookType } from "./HIR"; +import { HIRFunction, IdentifierId, Place, getHookKind } from "./HIR"; import { eachInstructionValueOperand, eachTerminalOperand } from "./visitors"; /** @@ -56,7 +56,7 @@ export function validateHooksUsage(fn: HIRFunction): void { for (const instr of block.instructions) { if ( instr.value.kind === "LoadGlobal" && - isHookType(instr.lvalue.identifier) + getHookKind(fn.env, instr.lvalue.identifier) != null ) { hooks.add(instr.lvalue.identifier.id); } else if (instr.value.kind === "CallExpression") { diff --git a/compiler/forget/src/HIR/ValidateUnconditionalHooks.ts b/compiler/forget/src/HIR/ValidateUnconditionalHooks.ts index a80eebed11..77d5656b87 100644 --- a/compiler/forget/src/HIR/ValidateUnconditionalHooks.ts +++ b/compiler/forget/src/HIR/ValidateUnconditionalHooks.ts @@ -13,7 +13,7 @@ import { import { findBlocksWithBackEdges } from "../Optimization/DeadCodeElimination"; import { Err, Ok, Result } from "../Utils/Result"; import { PostDominator, computePostDominatorTree } from "./Dominator"; -import { BlockId, HIRFunction, isHookType } from "./HIR"; +import { BlockId, HIRFunction, getHookKind } from "./HIR"; /** * Validates that the function honors the [Rules of Hooks](https://react.dev/warnings/invalid-hook-call-warning) @@ -85,7 +85,7 @@ export function validateUnconditionalHooks( for (const instr of block.instructions) { if ( instr.value.kind === "CallExpression" && - isHookType(instr.value.callee.identifier) + getHookKind(fn.env, instr.value.callee.identifier) != null ) { const loc = instr.loc; // TODO: the current ESLint rule has different error messages for code that is called conditionally, in a loop, etc. diff --git a/compiler/forget/src/HIR/index.ts b/compiler/forget/src/HIR/index.ts index 0a06c43cfd..47e296ed83 100644 --- a/compiler/forget/src/HIR/index.ts +++ b/compiler/forget/src/HIR/index.ts @@ -7,7 +7,7 @@ export { lower } from "./BuildHIR"; export { computeDominatorTree, computePostDominatorTree } from "./Dominator"; -export { Environment } from "./Environment"; +export { Environment, Hook } from "./Environment"; export * from "./HIR"; export { markInstructionIds, @@ -15,7 +15,6 @@ export { removeUnreachableFallthroughs, reversePostorderBlocks, } from "./HIRBuilder"; -export { Hook } from "./Hooks"; export { mergeConsecutiveBlocks } from "./MergeConsecutiveBlocks"; export { printFunction, printHIR } from "./PrintHIR"; export { validateConsistentIdentifiers } from "./ValidateConsistentIdentifiers"; diff --git a/compiler/forget/src/Inference/DropMemoCalls.ts b/compiler/forget/src/Inference/DropMemoCalls.ts index 4b1669f7d3..6267b7c385 100644 --- a/compiler/forget/src/Inference/DropMemoCalls.ts +++ b/compiler/forget/src/Inference/DropMemoCalls.ts @@ -5,17 +5,16 @@ * LICENSE file in the root directory of this source tree. */ -import { Effect, HIRFunction, HookType, isHookType } from "../HIR"; +import { Effect, HIRFunction, getHookKind } from "../HIR"; export default function (func: HIRFunction): void { for (const [_, block] of func.body.blocks) { for (const instr of block.instructions) { switch (instr.value.kind) { case "CallExpression": { - if (isHookType(instr.value.callee.identifier)) { - const name = (instr.value.callee.identifier.type as HookType) - .definition?.name; - if (name === "useMemo") { + const hookKind = getHookKind(func.env, instr.value.callee.identifier); + if (hookKind != null) { + if (hookKind === "useMemo") { const [fn] = instr.value.args; // TODO(gsn): Consider inlining the function passed to useMemo, @@ -38,7 +37,7 @@ export default function (func: HIRFunction): void { loc: instr.value.loc, }; } - } else if (name === "useCallback") { + } else if (hookKind === "useCallback") { const [fn] = instr.value.args; // Instead of a Call, just alias the callback directly. diff --git a/compiler/forget/src/Inference/InferReferenceEffects.ts b/compiler/forget/src/Inference/InferReferenceEffects.ts index edfc37dbc9..a76cd1bdce 100644 --- a/compiler/forget/src/Inference/InferReferenceEffects.ts +++ b/compiler/forget/src/Inference/InferReferenceEffects.ts @@ -24,7 +24,11 @@ import { Type, ValueKind, } from "../HIR/HIR"; -import { FunctionSignature } from "../HIR/ObjectShape"; +import { + DefaultMutatingHook, + DefaultNonmutatingHook, + FunctionSignature, +} from "../HIR/ObjectShape"; import { printIdentifier, printMixedHIR, @@ -708,25 +712,33 @@ function inferBlock( continue; } case "CallExpression": { - if (instrValue.callee.identifier.type.kind === "Hook") { - const definition = instrValue.callee.identifier.type.definition; - if (definition !== null) { - effectKind = definition.effectKind; - valueKind = definition.valueKind; - } else if (env.enableAssumeHooksFollowRulesOfReact) { - effectKind = Effect.Freeze; - valueKind = ValueKind.Frozen; - } else { - effectKind = Effect.Mutate; - valueKind = ValueKind.Mutable; - } + let signature = getFunctionCallSignature( + env, + instrValue.callee.identifier.type + ); + signature = + env.enableFunctionCallSignatureOptimizations || + signature?.hookKind != null + ? signature + : null; + + if ( + signature && + signature.hookKind != null && + !env.enableTreatHooksAsFunctions + ) { + effectKind = signature.restParam; + valueKind = signature.returnValueKind; break; } - const signature = env.enableFunctionCallSignatureOptimizations - ? getFunctionCallSignature(env, instrValue.callee.identifier.type) - : null; - + // We currently always check reference effects of typed functions + // (i.e. call `referenceAndCheckError`). However, default custom hooks + // should not assert reference effects, since their signatures are only + // assumptions / defaults. + const isDefaultCustomHook = + instrValue.callee.identifier.type === DefaultMutatingHook || + instrValue.callee.identifier.type === DefaultNonmutatingHook; const effects = signature !== null ? getFunctionEffects(instrValue, signature) : null; const returnValueKind = @@ -735,9 +747,13 @@ function inferBlock( const arg = instrValue.args[i]; const place = arg.kind === "Identifier" ? arg : arg.place; if (effects !== null) { - // If effects are inferred for an argument, we should fail invalid - // mutating effects - state.referenceAndCheckError(place, effects[i]); + if (isDefaultCustomHook) { + state.reference(place, effects[i]); + } else { + // If effects are inferred for an argument, we should fail invalid + // mutating effects + state.referenceAndCheckError(place, effects[i]); + } } else { state.reference(place, Effect.Mutate); } diff --git a/compiler/forget/src/ReactiveScopes/FlattenScopesWithHooks.ts b/compiler/forget/src/ReactiveScopes/FlattenScopesWithHooks.ts index 6b6432a799..14fa5fe11b 100644 --- a/compiler/forget/src/ReactiveScopes/FlattenScopesWithHooks.ts +++ b/compiler/forget/src/ReactiveScopes/FlattenScopesWithHooks.ts @@ -6,12 +6,13 @@ */ import { + Environment, InstructionId, - isHookType, ReactiveFunction, ReactiveScopeBlock, ReactiveStatement, ReactiveValue, + getHookKind, } from "../HIR"; import { ReactiveFunctionTransform, @@ -31,17 +32,26 @@ import { * to ensure the hook call does not inadvertently become conditional. */ export function flattenScopesWithHooks(fn: ReactiveFunction): void { - visitReactiveFunction(fn, new Transform(), { hasHook: false }); + visitReactiveFunction(fn, new Transform(), { + env: fn.env, + hasHook: false, + }); } -type State = { hasHook: boolean }; +type State = { + env: Environment; + hasHook: boolean; +}; class Transform extends ReactiveFunctionTransform { override transformScope( scope: ReactiveScopeBlock, outerState: State ): Transformed { - const innerState: State = { hasHook: false }; + const innerState: State = { + env: outerState.env, + hasHook: false, + }; this.visitScope(scope, innerState); outerState.hasHook ||= innerState.hasHook; if (innerState.hasHook) { @@ -59,7 +69,7 @@ class Transform extends ReactiveFunctionTransform { this.traverseValue(id, value, state); if ( value.kind === "CallExpression" && - isHookType(value.callee.identifier) + getHookKind(state.env, value.callee.identifier) != null ) { state.hasHook = true; } diff --git a/compiler/forget/src/ReactiveScopes/InferReactiveIdentifiers.ts b/compiler/forget/src/ReactiveScopes/InferReactiveIdentifiers.ts index da749fd527..6c7c9ec54b 100644 --- a/compiler/forget/src/ReactiveScopes/InferReactiveIdentifiers.ts +++ b/compiler/forget/src/ReactiveScopes/InferReactiveIdentifiers.ts @@ -6,26 +6,32 @@ */ import { CompilerError } from "../CompilerError"; +import { Environment } from "../HIR"; import { Effect, IdentifierId, - isHookType, ReactiveFunction, ReactiveInstruction, + getHookKind, } from "../HIR/HIR"; import { eachInstructionLValue } from "../HIR/visitors"; import { assertExhaustive } from "../Utils/utils"; import { - eachReactiveValueOperand, ReactiveFunctionVisitor, + eachReactiveValueOperand, visitReactiveFunction, } from "./visitors"; type IdentifierReactivity = Map; class State { + env: Environment; reactivityMap: IdentifierReactivity = new Map(); temporaries: Map = new Map(); + + constructor(env: Environment) { + this.env = env; + } } class Visitor extends ReactiveFunctionVisitor { @@ -65,7 +71,7 @@ class Visitor extends ReactiveFunctionVisitor { if ( !hasReactiveInput && instr.value.kind === "CallExpression" && - isHookType(instr.value.callee.identifier) + getHookKind(state.env, instr.value.callee.identifier) != null ) { // Hooks cannot be memoized. Even if they do not accept any reactive inputs, // they are not guaranteed to memoize their return value, and their result @@ -169,7 +175,7 @@ export function inferReactiveIdentifiers( fn: ReactiveFunction ): Set { const visitor = new Visitor(); - const state = new State(); + const state = new State(fn.env); for (const param of fn.params) { state.reactivityMap.set(param.identifier.id, true); } diff --git a/compiler/forget/src/ReactiveScopes/PruneNonEscapingScopes.ts b/compiler/forget/src/ReactiveScopes/PruneNonEscapingScopes.ts index eb0e82ca09..0db85aa470 100644 --- a/compiler/forget/src/ReactiveScopes/PruneNonEscapingScopes.ts +++ b/compiler/forget/src/ReactiveScopes/PruneNonEscapingScopes.ts @@ -9,10 +9,9 @@ import invariant from "invariant"; import prettyFormat from "pretty-format"; import { CompilerError } from "../CompilerError"; import { + Environment, IdentifierId, InstructionId, - isHookType, - isMutableEffect, Pattern, Place, ReactiveFunction, @@ -23,6 +22,8 @@ import { ReactiveTerminalStatement, ReactiveValue, ScopeId, + getHookKind, + isMutableEffect, } from "../HIR"; import { eachInstructionValueOperand } from "../HIR/visitors"; import { log } from "../Utils/logger"; @@ -30,10 +31,10 @@ import { assertExhaustive } from "../Utils/utils"; import { getPlaceScope } from "./BuildReactiveBlocks"; import { printReactiveFunction } from "./PrintReactiveFunction"; import { - eachReactiveValueOperand, ReactiveFunctionTransform, ReactiveFunctionVisitor, Transformed, + eachReactiveValueOperand, visitReactiveFunction, } from "./visitors"; @@ -116,7 +117,7 @@ export function pruneNonEscapingScopes( ): void { // First build up a map of which instructions are involved in creating which values, // and which values are returned. - const state = new State(); + const state = new State(fn.env); if (fn.id !== null) { state.declare(fn.id.id); } @@ -200,6 +201,7 @@ type ScopeNode = { // Stores the identifier and scope graphs, set of returned identifiers, etc class State { + env: Environment; // Maps lvalues for LoadLocal to the identifier being loaded, to resolve indirections // in subsequent lvalues/rvalues definitions: Map = new Map(); @@ -208,6 +210,10 @@ class State { scopes: Map = new Map(); escapingValues: Set = new Set(); + constructor(env: Environment) { + this.env = env; + } + /** * Declare a new identifier, used for function id and params */ @@ -715,7 +721,7 @@ class CollectDependenciesVisitor extends ReactiveFunctionVisitor { ); } else if (instruction.value.kind === "CallExpression") { const callee = instruction.value.callee; - if (isHookType(callee.identifier)) { + if (getHookKind(state.env, callee.identifier)) { for (const operand of eachInstructionValueOperand(instruction.value)) { state.escapingValues.add(operand.identifier.id); }