diff --git a/src/useScrolling.ts b/src/useScrolling.ts index 2e2ddbef9a..dc7452e45b 100644 --- a/src/useScrolling.ts +++ b/src/useScrolling.ts @@ -5,7 +5,8 @@ const useScrolling = (ref: RefObject): boolean => { const [scrolling, setScrolling] = useState(false); useEffect(() => { - if (ref.current) { + const element = ref.current; + if (element) { let scrollingTimeout; const handleScrollEnd = () => { @@ -18,11 +19,9 @@ const useScrolling = (ref: RefObject): boolean => { scrollingTimeout = setTimeout(() => handleScrollEnd(), 150); }; - on(ref.current, 'scroll', handleScroll, false); + on(element, 'scroll', handleScroll, false); return () => { - if (ref.current) { - off(ref.current, 'scroll', handleScroll, false); - } + off(element, 'scroll', handleScroll, false); }; } return () => {}; diff --git a/tests/useScrolling.test.ts b/tests/useScrolling.test.ts new file mode 100644 index 0000000000..92b57986ce --- /dev/null +++ b/tests/useScrolling.test.ts @@ -0,0 +1,32 @@ +import { renderHook } from '@testing-library/react-hooks'; +import useScrolling from '../src/useScrolling'; + +describe('useScrolling', () => { + it('removes the listener from the element it originally observed', () => { + const element = document.createElement('div'); + const replacement = document.createElement('div'); + const removeListener = jest.spyOn(element, 'removeEventListener'); + const removeReplacementListener = jest.spyOn(replacement, 'removeEventListener'); + const ref = { current: element }; + + const { unmount } = renderHook(() => useScrolling(ref)); + + ref.current = replacement; + unmount(); + + expect(removeListener).toHaveBeenCalledWith('scroll', expect.any(Function), false); + expect(removeReplacementListener).not.toHaveBeenCalled(); + }); + + it('removes the listener when the ref has been cleared', () => { + const element = document.createElement('div'); + const removeListener = jest.spyOn(element, 'removeEventListener'); + const ref: { current: HTMLElement | null } = { current: element }; + const { unmount } = renderHook(() => useScrolling(ref)); + + ref.current = null; + unmount(); + + expect(removeListener).toHaveBeenCalledWith('scroll', expect.any(Function), false); + }); +});