diff --git a/types/jest-when/index.d.ts b/types/jest-when/index.d.ts
index f484502f4d..0579419425 100644
--- a/types/jest-when/index.d.ts
+++ b/types/jest-when/index.d.ts
@@ -6,22 +6,22 @@
///
-export interface PartialMockInstance {
- mockReturnValue: jest.MockInstance['mockReturnValue'];
- mockReturnValueOnce: jest.MockInstance['mockReturnValueOnce'];
- mockResolvedValue: jest.MockInstance['mockResolvedValue'];
- mockResolvedValueOnce: jest.MockInstance['mockResolvedValueOnce'];
- mockRejectedValue: jest.MockInstance['mockRejectedValue'];
- mockRejectedValueOnce: jest.MockInstance['mockRejectedValueOnce'];
-}
-
export interface When {
(fn: jest.Mocked | jest.Mock): When;
+
// due to no-unnecessary-generics lint rule, the generics have been replaced with 'any'
// calledWith(...matchers: any[]): PartialMockInstance;
// expectCalledWith(...matchers: any[]): PartialMockInstance;
- calledWith(...matchers: any[]): PartialMockInstance;
- expectCalledWith(...matchers: any[]): PartialMockInstance;
+ calledWith(...matchers: any[]): When;
+
+ expectCalledWith(...matchers: any[]): When;
+
+ mockReturnValue: (value: any) => jest.MockInstance['mockReturnValue'] & When;
+ mockReturnValueOnce: (value: any) => jest.MockInstance['mockReturnValue'] & When;
+ mockResolvedValue: (value: any) => jest.MockInstance['mockReturnValue'] & When;
+ mockResolvedValueOnce: (value: any) => jest.MockInstance['mockReturnValue'] & When;
+ mockRejectedValue: (value: any) => jest.MockInstance['mockReturnValue'] & When;
+ mockRejectedValueOnce: (value: any) => jest.MockInstance['mockReturnValue'] & When;
}
export const when: When;
diff --git a/types/jest-when/jest-when-tests.ts b/types/jest-when/jest-when-tests.ts
index 13da65928d..179ea2fae1 100644
--- a/types/jest-when/jest-when-tests.ts
+++ b/types/jest-when/jest-when-tests.ts
@@ -43,6 +43,37 @@ describe('mock-when test', () => {
expect(fn(5)).toEqual(undefined);
});
+ it('Supports chained calls:', () => {
+ const fn = jest.fn();
+ when(fn)
+ .calledWith(1).mockReturnValue('no')
+ .calledWith(2).mockReturnValue('way?')
+ .calledWith(3).mockReturnValue('yes')
+ .calledWith(4).mockReturnValue('way!');
+
+ expect(fn(1)).toEqual('no');
+ expect(fn(2)).toEqual('way?');
+ expect(fn(3)).toEqual('yes');
+ expect(fn(4)).toEqual('way!');
+ expect(fn(5)).toEqual(undefined);
+ });
+
+ it('Supports chained calls with defaults:', () => {
+ const fn = jest.fn();
+ when(fn)
+ .mockReturnValue('nice')
+ .calledWith(1).mockReturnValue('no')
+ .calledWith(2).mockReturnValue('way?')
+ .calledWith(3).mockReturnValue('yes')
+ .calledWith(4).mockReturnValue('way!');
+
+ expect(fn(1)).toEqual('no');
+ expect(fn(2)).toEqual('way?');
+ expect(fn(3)).toEqual('yes');
+ expect(fn(4)).toEqual('way!');
+ expect(fn(5)).toEqual('nice');
+ });
+
it('Assert the args:', () => {
const fn = jest.fn();
when(fn).expectCalledWith(1).mockReturnValue('x');