diff --git a/packages/react-dom/src/__tests__/refs-test.js b/packages/react-dom/src/__tests__/refs-test.js
index d68c6f9334..9dcaed18c2 100644
--- a/packages/react-dom/src/__tests__/refs-test.js
+++ b/packages/react-dom/src/__tests__/refs-test.js
@@ -9,9 +9,9 @@
'use strict';
-let React = require('react');
-let ReactDOMClient = require('react-dom/client');
-let act = require('internal-test-utils').act;
+const React = require('react');
+const ReactDOMClient = require('react-dom/client');
+const act = require('internal-test-utils').act;
// This is testing if string refs are deleted from `instance.refs`
// Once support for string refs is removed, this test can be removed.
@@ -19,13 +19,6 @@ let act = require('internal-test-utils').act;
describe('reactiverefs', () => {
let container;
- beforeEach(() => {
- jest.resetModules();
- React = require('react');
- ReactDOMClient = require('react-dom/client');
- act = require('internal-test-utils').act;
- });
-
afterEach(() => {
if (container) {
document.body.removeChild(container);
@@ -199,11 +192,6 @@ describe('reactiverefs', () => {
describe('ref swapping', () => {
let RefHopsAround;
beforeEach(() => {
- jest.resetModules();
- React = require('react');
- ReactDOMClient = require('react-dom/client');
- act = require('internal-test-utils').act;
-
RefHopsAround = class extends React.Component {
container = null;
state = {count: 0};
@@ -804,3 +792,96 @@ describe('refs return clean up function', () => {
expect(nullHandler).toHaveBeenCalledTimes(0);
});
});
+
+describe('useImerativeHandle refs', () => {
+ const ImperativeHandleComponent = React.forwardRef(({name}, ref) => {
+ React.useImperativeHandle(
+ ref,
+ () => ({
+ greet() {
+ return `Hello ${name}`;
+ },
+ }),
+ [name],
+ );
+ return null;
+ });
+
+ it('should work with object style refs', async () => {
+ const container = document.createElement('div');
+ const root = ReactDOMClient.createRoot(container);
+ const ref = React.createRef();
+
+ await act(async () => {
+ root.render();
+ });
+ expect(ref.current.greet()).toBe('Hello Alice');
+ await act(() => {
+ root.render(null);
+ });
+ expect(ref.current).toBe(null);
+ });
+
+ it('should work with callback style refs', async () => {
+ const container = document.createElement('div');
+ const root = ReactDOMClient.createRoot(container);
+ let current = null;
+
+ await act(async () => {
+ root.render(
+ {
+ current = r;
+ }}
+ />,
+ );
+ });
+ expect(current.greet()).toBe('Hello Alice');
+ await act(() => {
+ root.render(null);
+ });
+ expect(current).toBe(null);
+ });
+
+ it('should work with callback style refs with cleanup function', async () => {
+ const container = document.createElement('div');
+ const root = ReactDOMClient.createRoot(container);
+
+ let cleanupCalls = 0;
+ let createCalls = 0;
+ let current = null;
+
+ const ref = r => {
+ current = r;
+ createCalls++;
+ return () => {
+ current = null;
+ cleanupCalls++;
+ };
+ };
+
+ await act(async () => {
+ root.render();
+ });
+ expect(current.greet()).toBe('Hello Alice');
+ expect(createCalls).toBe(1);
+ expect(cleanupCalls).toBe(0);
+
+ // update a dep should recreate the ref
+ await act(async () => {
+ root.render();
+ });
+ expect(current.greet()).toBe('Hello Bob');
+ expect(createCalls).toBe(2);
+ expect(cleanupCalls).toBe(1);
+
+ // unmounting should call cleanup
+ await act(() => {
+ root.render(null);
+ });
+ expect(current).toBe(null);
+ expect(createCalls).toBe(2);
+ expect(cleanupCalls).toBe(2);
+ });
+});
diff --git a/packages/react-reconciler/src/ReactFiberHooks.js b/packages/react-reconciler/src/ReactFiberHooks.js
index 59e7ae780f..8a65e6170f 100644
--- a/packages/react-reconciler/src/ReactFiberHooks.js
+++ b/packages/react-reconciler/src/ReactFiberHooks.js
@@ -2564,9 +2564,14 @@ function imperativeHandleEffect(
if (typeof ref === 'function') {
const refCallback = ref;
const inst = create();
- refCallback(inst);
+ const refCleanup = refCallback(inst);
return () => {
- refCallback(null);
+ if (typeof refCleanup === 'function') {
+ // $FlowFixMe[incompatible-use] we need to assume no parameters
+ refCleanup();
+ } else {
+ refCallback(null);
+ }
};
} else if (ref !== null && ref !== undefined) {
const refObject = ref;