compiler/packages/babel-plugin-react-compiler/src/Inference/DropManualMemoization.ts TYPESCRIPT 600 lines View on github.com → Search inside
1/**2 * Copyright (c) Meta Platforms, Inc. and affiliates.3 *4 * This source code is licensed under the MIT license found in the5 * LICENSE file in the root directory of this source tree.6 */78import {CompilerDiagnostic, CompilerError, SourceLocation} from '..';9import {ErrorCategory} from '../CompilerError';10import {11  CallExpression,12  Effect,13  Environment,14  FinishMemoize,15  FunctionExpression,16  HIRFunction,17  IdentifierId,18  Instruction,19  InstructionId,20  InstructionValue,21  LoadGlobal,22  LoadLocal,23  ManualMemoDependency,24  MethodCall,25  Place,26  PropertyLoad,27  SpreadPattern,28  StartMemoize,29  TInstruction,30  getHookKindForType,31  makeInstructionId,32} from '../HIR';33import {createTemporaryPlace, markInstructionIds} from '../HIR/HIRBuilder';3435type ManualMemoCallee = {36  kind: 'useMemo' | 'useCallback';37  loadInstr: TInstruction<LoadGlobal> | TInstruction<PropertyLoad>;38};3940type IdentifierSidemap = {41  functions: Map<IdentifierId, TInstruction<FunctionExpression>>;42  manualMemos: Map<IdentifierId, ManualMemoCallee>;43  react: Set<IdentifierId>;44  maybeDepsLists: Map<IdentifierId, {loc: SourceLocation; deps: Array<Place>}>;45  maybeDeps: Map<IdentifierId, ManualMemoDependency>;46  optionals: Set<IdentifierId>;47};4849/**50 * Collect loads from named variables and property reads from @value51 * into `maybeDeps`52 * Returns the variable + property reads represented by @instr53 */54export function collectMaybeMemoDependencies(55  value: InstructionValue,56  maybeDeps: Map<IdentifierId, ManualMemoDependency>,57  optional: boolean,58): ManualMemoDependency | null {59  switch (value.kind) {60    case 'LoadGlobal': {61      return {62        root: {63          kind: 'Global',64          identifierName: value.binding.name,65        },66        path: [],67        loc: value.loc,68      };69    }70    case 'PropertyLoad': {71      const object = maybeDeps.get(value.object.identifier.id);72      if (object != null) {73        return {74          root: object.root,75          // TODO: determine if the access is optional76          path: [77            ...object.path,78            {property: value.property, optional, loc: value.loc},79          ],80          loc: value.loc,81        };82      }83      break;84    }8586    case 'LoadLocal':87    case 'LoadContext': {88      const source = maybeDeps.get(value.place.identifier.id);89      if (source != null) {90        return source;91      } else if (92        value.place.identifier.name != null &&93        value.place.identifier.name.kind === 'named'94      ) {95        return {96          root: {97            kind: 'NamedLocal',98            value: {...value.place},99            constant: false,100          },101          path: [],102          loc: value.place.loc,103        };104      }105      break;106    }107    case 'StoreLocal': {108      /*109       * Value blocks rely on StoreLocal to populate their return value.110       * We need to track these as optional property chains are valid in111       * source depslists112       */113      const lvalue = value.lvalue.place.identifier;114      const rvalue = value.value.identifier.id;115      const aliased = maybeDeps.get(rvalue);116      if (aliased != null && lvalue.name?.kind !== 'named') {117        maybeDeps.set(lvalue.id, aliased);118        return aliased;119      }120      break;121    }122  }123  return null;124}125126function collectTemporaries(127  instr: Instruction,128  env: Environment,129  sidemap: IdentifierSidemap,130): void {131  const {value, lvalue} = instr;132  switch (value.kind) {133    case 'FunctionExpression': {134      sidemap.functions.set(135        instr.lvalue.identifier.id,136        instr as TInstruction<FunctionExpression>,137      );138      break;139    }140    case 'LoadGlobal': {141      const global = env.getGlobalDeclaration(value.binding, value.loc);142      const hookKind = global !== null ? getHookKindForType(env, global) : null;143      const lvalId = instr.lvalue.identifier.id;144      if (hookKind === 'useMemo' || hookKind === 'useCallback') {145        sidemap.manualMemos.set(lvalId, {146          kind: hookKind,147          loadInstr: instr as TInstruction<LoadGlobal>,148        });149      } else if (value.binding.name === 'React') {150        sidemap.react.add(lvalId);151      }152      break;153    }154    case 'PropertyLoad': {155      if (sidemap.react.has(value.object.identifier.id)) {156        const property = value.property;157        if (property === 'useMemo' || property === 'useCallback') {158          sidemap.manualMemos.set(instr.lvalue.identifier.id, {159            kind: property as 'useMemo' | 'useCallback',160            loadInstr: instr as TInstruction<PropertyLoad>,161          });162        }163      }164      break;165    }166    case 'ArrayExpression': {167      if (value.elements.every(e => e.kind === 'Identifier')) {168        sidemap.maybeDepsLists.set(instr.lvalue.identifier.id, {169          loc: value.loc,170          deps: value.elements as Array<Place>,171        });172      }173      break;174    }175  }176  const maybeDep = collectMaybeMemoDependencies(177    value,178    sidemap.maybeDeps,179    sidemap.optionals.has(lvalue.identifier.id),180  );181  // We don't expect named lvalues during this pass (unlike ValidatePreservingManualMemo)182  if (maybeDep != null) {183    sidemap.maybeDeps.set(lvalue.identifier.id, maybeDep);184  }185}186187function makeManualMemoizationMarkers(188  fnExpr: Place,189  env: Environment,190  depsList: Array<ManualMemoDependency> | null,191  depsLoc: SourceLocation | null,192  memoDecl: Place,193  manualMemoId: number,194): [TInstruction<StartMemoize>, TInstruction<FinishMemoize>] {195  return [196    {197      id: makeInstructionId(0),198      lvalue: createTemporaryPlace(env, fnExpr.loc),199      value: {200        kind: 'StartMemoize',201        manualMemoId,202        /*203         * Use deps list from source instead of inferred deps204         * as dependencies205         */206        deps: depsList,207        depsLoc,208        loc: fnExpr.loc,209      },210      effects: null,211      loc: fnExpr.loc,212    },213    {214      id: makeInstructionId(0),215      lvalue: createTemporaryPlace(env, fnExpr.loc),216      value: {217        kind: 'FinishMemoize',218        manualMemoId,219        decl: {...memoDecl},220        loc: fnExpr.loc,221      },222      effects: null,223      loc: fnExpr.loc,224    },225  ];226}227228function getManualMemoizationReplacement(229  fn: Place,230  loc: SourceLocation,231  kind: 'useMemo' | 'useCallback',232): LoadLocal | CallExpression {233  if (kind === 'useMemo') {234    /*235     * Replace the hook callee with the fn arg.236     *237     * before:238     *   $1 = LoadGlobal useMemo       // load the useMemo global239     *   $2 = FunctionExpression ...   // memo function240     *   $3 = ArrayExpression [ ... ]  // deps array241     *   $4 = Call $1 ($2, $3 )        // invoke useMemo w fn and deps242     *243     * after:244     *   $1 = LoadGlobal useMemo       // load the useMemo global (dead code)245     *   $2 = FunctionExpression ...   // memo function246     *   $3 = ArrayExpression [ ... ]  // deps array (dead code)247     *   $4 = Call $2 ()               // invoke the memo function itself248     *249     * Note that a later pass (InlineImmediatelyInvokedFunctionExpressions) will250     * inline the useMemo callback along with any other immediately invoked IIFEs.251     */252    return {253      kind: 'CallExpression',254      callee: fn,255      /*256       * Drop the args, including the deps array which DCE will remove257       * later.258       */259      args: [],260      loc,261    };262  } else {263    /*264     * Instead of a Call, just alias the callback directly.265     *266     * before:267     *   $1 = LoadGlobal useCallback268     *   $2 = FunctionExpression ...   // the callback being memoized269     *   $3 = ArrayExpression ...      // deps array270     *   $4 = Call $1 ( $2, $3 )       // invoke useCallback271     *272     * after:273     *   $1 = LoadGlobal useCallback   // dead code274     *   $2 = FunctionExpression ...   // the callback being memoized275     *   $3 = ArrayExpression ...      // deps array (dead code)276     *   $4 = LoadLocal $2             // reference the function277     */278    return {279      kind: 'LoadLocal',280      place: {281        kind: 'Identifier',282        identifier: fn.identifier,283        effect: Effect.Unknown,284        reactive: false,285        loc,286      },287      loc,288    };289  }290}291292function extractManualMemoizationArgs(293  instr: TInstruction<CallExpression> | TInstruction<MethodCall>,294  kind: 'useCallback' | 'useMemo',295  sidemap: IdentifierSidemap,296  env: Environment,297): {298  fnPlace: Place;299  depsList: Array<ManualMemoDependency> | null;300  depsLoc: SourceLocation | null;301} | null {302  const [fnPlace, depsListPlace] = instr.value.args as Array<303    Place | SpreadPattern | undefined304  >;305  if (fnPlace == null || fnPlace.kind !== 'Identifier') {306    env.recordError(307      CompilerDiagnostic.create({308        category: ErrorCategory.UseMemo,309        reason: `Expected a callback function to be passed to ${kind}`,310        description:311          kind === 'useCallback'312            ? 'The first argument to useCallback() must be a function to cache'313            : 'The first argument to useMemo() must be a function that calculates a result to cache',314        suggestions: null,315      }).withDetails({316        kind: 'error',317        loc: instr.value.loc,318        message:319          kind === 'useCallback'320            ? `Expected a callback function`321            : `Expected a memoization function`,322      }),323    );324    return null;325  }326  if (depsListPlace == null) {327    return {328      fnPlace,329      depsList: null,330      depsLoc: null,331    };332  }333  const maybeDepsList =334    depsListPlace.kind === 'Identifier'335      ? sidemap.maybeDepsLists.get(depsListPlace.identifier.id)336      : null;337  if (maybeDepsList == null) {338    env.recordError(339      CompilerDiagnostic.create({340        category: ErrorCategory.UseMemo,341        reason: `Expected the dependency list for ${kind} to be an array literal`,342        description: `Expected the dependency list for ${kind} to be an array literal`,343        suggestions: null,344      }).withDetails({345        kind: 'error',346        loc:347          depsListPlace?.kind === 'Identifier' ? depsListPlace.loc : instr.loc,348        message: `Expected the dependency list for ${kind} to be an array literal`,349      }),350    );351    return null;352  }353  const depsList: Array<ManualMemoDependency> = [];354  for (const dep of maybeDepsList.deps) {355    const maybeDep = sidemap.maybeDeps.get(dep.identifier.id);356    if (maybeDep == null) {357      env.recordError(358        CompilerDiagnostic.create({359          category: ErrorCategory.UseMemo,360          reason: `Expected the dependency list to be an array of simple expressions (e.g. \`x\`, \`x.y.z\`, \`x?.y?.z\`)`,361          description: `Expected the dependency list to be an array of simple expressions (e.g. \`x\`, \`x.y.z\`, \`x?.y?.z\`)`,362          suggestions: null,363        }).withDetails({364          kind: 'error',365          loc: dep.loc,366          message: `Expected the dependency list to be an array of simple expressions (e.g. \`x\`, \`x.y.z\`, \`x?.y?.z\`)`,367        }),368      );369    } else {370      depsList.push(maybeDep);371    }372  }373  return {374    fnPlace,375    depsList,376    depsLoc: maybeDepsList.loc,377  };378}379380/*381 * Removes manual memoization using the `useMemo` and `useCallback` APIs. This pass is designed382 * to compose with InlineImmediatelyInvokedFunctionExpressions, and needs to run prior to entering383 * SSA form (alternatively we could refactor and re-EnterSSA after inlining). Therefore it cannot384 * rely on type inference to find useMemo/useCallback invocations, and instead does basic tracking385 * of globals and property loads to find both direct calls as well as usage via the React namespace,386 * eg `React.useMemo()`.387 *388 * This pass also validates that useMemo callbacks return a value (not void), ensuring that useMemo389 * is only used for memoizing values and not for running arbitrary side effects.390 */391export function dropManualMemoization(func: HIRFunction): void {392  const isValidationEnabled =393    func.env.config.validatePreserveExistingMemoizationGuarantees ||394    func.env.config.validateNoSetStateInRender ||395    func.env.config.enablePreserveExistingMemoizationGuarantees;396  const optionals = findOptionalPlaces(func);397  const sidemap: IdentifierSidemap = {398    functions: new Map(),399    manualMemos: new Map(),400    react: new Set(),401    maybeDeps: new Map(),402    maybeDepsLists: new Map(),403    optionals,404  };405  let nextManualMemoId = 0;406407  /**408   * Phase 1:409   * - Overwrite manual memoization from410   *   CallExpression callee="useMemo/Callback", args=[fnArg, depslist])411   *   to either412   *   CallExpression callee=fnArg413   *   LoadLocal fnArg414   * - (if validation is enabled) collect manual memoization markers415   */416  const queuedInserts: Map<417    InstructionId,418    TInstruction<StartMemoize> | TInstruction<FinishMemoize>419  > = new Map();420  for (const [_, block] of func.body.blocks) {421    for (let i = 0; i < block.instructions.length; i++) {422      const instr = block.instructions[i]!;423      if (424        instr.value.kind === 'CallExpression' ||425        instr.value.kind === 'MethodCall'426      ) {427        const id =428          instr.value.kind === 'CallExpression'429            ? instr.value.callee.identifier.id430            : instr.value.property.identifier.id;431432        const manualMemo = sidemap.manualMemos.get(id);433        if (manualMemo != null) {434          const memoDetails = extractManualMemoizationArgs(435            instr as TInstruction<CallExpression> | TInstruction<MethodCall>,436            manualMemo.kind,437            sidemap,438            func.env,439          );440441          if (memoDetails == null) {442            continue;443          }444          const {fnPlace, depsList, depsLoc} = memoDetails;445446          instr.value = getManualMemoizationReplacement(447            fnPlace,448            instr.value.loc,449            manualMemo.kind,450          );451          if (isValidationEnabled) {452            /**453             * Explicitly bail out when we encounter manual memoization454             * without inline instructions, as our current validation455             * assumes that source depslists closely match inferred deps456             * due to the `exhaustive-deps` lint rule (which only provides457             * diagnostics for inline memo functions)458             * ```js459             * useMemo(opaqueFn, [dep1, dep2]);460             * ```461             * While we could handle this by diffing reactive scope deps462             * of the opaque arg against the source depslist, this pattern463             * is rare and likely sketchy.464             */465            if (!sidemap.functions.has(fnPlace.identifier.id)) {466              func.env.recordError(467                CompilerDiagnostic.create({468                  category: ErrorCategory.UseMemo,469                  reason: `Expected the first argument to be an inline function expression`,470                  description: `Expected the first argument to be an inline function expression`,471                  suggestions: [],472                }).withDetails({473                  kind: 'error',474                  loc: fnPlace.loc,475                  message: `Expected the first argument to be an inline function expression`,476                }),477              );478              continue;479            }480            const memoDecl: Place =481              manualMemo.kind === 'useMemo'482                ? instr.lvalue483                : {484                    kind: 'Identifier',485                    identifier: fnPlace.identifier,486                    effect: Effect.Unknown,487                    reactive: false,488                    loc: fnPlace.loc,489                  };490491            const [startMarker, finishMarker] = makeManualMemoizationMarkers(492              fnPlace,493              func.env,494              depsList,495              depsLoc,496              memoDecl,497              nextManualMemoId++,498            );499500            /**501             * Insert StartMarker right after the `useMemo`/`useCallback` load to502             * ensure all temporaries created when lowering the inline fn expression503             * are included.504             * e.g.505             * ```506             * 0: LoadGlobal useMemo507             * 1: StartMarker deps=[var]508             * 2: t0 = LoadContext [var]509             * 3: function deps=t0510             * ...511             * ```512             */513            queuedInserts.set(manualMemo.loadInstr.id, startMarker);514            queuedInserts.set(instr.id, finishMarker);515          }516        }517      } else {518        collectTemporaries(instr, func.env, sidemap);519      }520    }521  }522523  /**524   * Phase 2: Insert manual memoization markers as needed525   */526  if (queuedInserts.size > 0) {527    let hasChanges = false;528    for (const [_, block] of func.body.blocks) {529      let nextInstructions: Array<Instruction> | null = null;530      for (let i = 0; i < block.instructions.length; i++) {531        const instr = block.instructions[i];532        const insertInstr = queuedInserts.get(instr.id);533        if (insertInstr != null) {534          nextInstructions = nextInstructions ?? block.instructions.slice(0, i);535          nextInstructions.push(instr);536          nextInstructions.push(insertInstr);537        } else if (nextInstructions != null) {538          nextInstructions.push(instr);539        }540      }541      if (nextInstructions !== null) {542        block.instructions = nextInstructions;543        hasChanges = true;544      }545    }546547    if (hasChanges) {548      markInstructionIds(func.body);549    }550  }551}552553function findOptionalPlaces(fn: HIRFunction): Set<IdentifierId> {554  const optionals = new Set<IdentifierId>();555  for (const [, block] of fn.body.blocks) {556    if (block.terminal.kind === 'optional' && block.terminal.optional) {557      const optionalTerminal = block.terminal;558      let testBlock = fn.body.blocks.get(block.terminal.test)!;559      loop: while (true) {560        const terminal = testBlock.terminal;561        switch (terminal.kind) {562          case 'branch': {563            if (terminal.fallthrough === optionalTerminal.fallthrough) {564              // found it565              const consequent = fn.body.blocks.get(terminal.consequent)!;566              const last = consequent.instructions.at(-1);567              if (last !== undefined && last.value.kind === 'StoreLocal') {568                optionals.add(last.value.value.identifier.id);569              }570              break loop;571            } else {572              testBlock = fn.body.blocks.get(terminal.fallthrough)!;573            }574            break;575          }576          case 'optional':577          case 'logical':578          case 'sequence':579          case 'ternary': {580            testBlock = fn.body.blocks.get(terminal.fallthrough)!;581            break;582          }583          case 'maybe-throw': {584            testBlock = fn.body.blocks.get(terminal.continuation)!;585            break;586          }587          default: {588            CompilerError.invariant(false, {589              reason: `Unexpected terminal in optional`,590              message: `Unexpected ${terminal.kind} in optional`,591              loc: terminal.loc,592            });593          }594        }595      }596    }597  }598  return optionals;599}

Findings

✓ No findings reported for this file.

Get this view in your editor

Same data, no extra tab — call code_get_file + code_get_findings over MCP from Claude/Cursor/Copilot.