@@ -19,40 +19,31 @@ interface TestingZoneType extends ZoneType {
1919}
2020
2121/**
22- * The list of method names for the describe/suite factories.
22+ * The list of method names for the describe/suite and test/it factories
23+ * that are called directly (i.e. same signature as describe/it).
24+ *
2325 * Example: `describe.skip('...', () => { ... });`
2426 * Sourced from https://vitest.dev/api/#describe
2527 */
26- const DESCRIBE_FACTORY_NAMES = [
28+ const DIRECT_MODIFIER_NAMES = [
2729 'skip' ,
28- 'skipIf' ,
29- 'runIf' ,
3030 'only' ,
3131 'concurrent' ,
3232 'sequential' ,
3333 'shuffle' ,
3434 'todo' ,
35- 'each' ,
36- 'for' ,
3735] as const ;
3836
3937/**
40- * The list of method names for the test/it factories.
41- * Example: `test.skip('...', () => { ... });`
42- * Sourced from https://vitest.dev/api/#test
38+ * The list of method names for the describe/suite and test/it modifiers
39+ * that are curried (i.e. called once with a condition/table to get back a chainable fn).
40+ *
41+ * Example: `describe.each([...])('...', () => { ... });`
42+ * Sourced from https://vitest.dev/api/#describe
4343 */
44- const TEST_FACTORY_NAMES = [
45- 'skip' ,
46- 'skipIf' ,
47- 'runIf' ,
48- 'only' ,
49- 'concurrent' ,
50- 'sequential' ,
51- 'shuffle' ,
52- 'todo' ,
53- 'each' ,
54- 'for' ,
55- ] as const ;
44+ const CURRIED_MODIFIER_NAMES = [ 'skipIf' , 'runIf' , 'each' , 'for' ] as const ;
45+
46+ type TEST_MODIFIER_NAME = ( typeof DIRECT_MODIFIER_NAMES | typeof CURRIED_MODIFIER_NAMES ) [ number ] ;
5647
5748export function patchVitest ( Zone : ZoneType ) : void {
5849 Zone . __load_patch ( 'vitest' , ( context : any , Zone : TestingZoneType ) => {
@@ -83,6 +74,10 @@ export function patchVitest(Zone: ZoneType): void {
8374 * synchronous-only zone.
8475 */
8576 function wrapDescribeInZone ( describeBody : Function ) : Function {
77+ // `describe` might be called without a body (e.g. `describe.todo`)
78+ if ( typeof describeBody !== 'function' ) {
79+ return describeBody ;
80+ }
8681 return function ( this : unknown , ...args : unknown [ ] ) {
8782 return syncZone . run ( describeBody , this , args ) ;
8883 } ;
@@ -111,54 +106,53 @@ export function patchVitest(Zone: ZoneType): void {
111106 return wrappedFunc ;
112107 }
113108
114- [ 'suite' , 'describe' ] . forEach ( ( methodName ) => {
115- let originalVitestFn : Function & Record < ( typeof DESCRIBE_FACTORY_NAMES ) [ number ] , Function > =
116- context [ methodName ] ;
109+ /** Patch functions with modifiers (i.e. `describe`/`it`). */
110+ function patchFnWithModifiers ( methodName : string , wrapFn : ( fn : Function ) => Function ) {
111+ const originalVitestFn : Function & Record < TEST_MODIFIER_NAME , Function > = context [ methodName ] ;
117112 // Skip if already patched
118113 if ( context [ Zone . __symbol__ ( methodName ) ] ) {
119114 return ;
120115 }
121-
122116 context [ Zone . __symbol__ ( methodName ) ] = originalVitestFn ;
117+
118+ // Patching the main function
123119 context [ methodName ] = function ( this : unknown , ...args : [ unknown , Function , ...unknown [ ] ] ) {
124- args [ 1 ] = wrapDescribeInZone ( args [ 1 ] ) ;
120+ args [ 1 ] = wrapFn ( args [ 1 ] ) ;
125121 return originalVitestFn . apply ( this , args ) ;
126122 } ;
127123
128- for ( const factoryName of DESCRIBE_FACTORY_NAMES ) {
129- context [ methodName ] [ factoryName ] = function ( this : unknown , ...factoryArgs : unknown [ ] ) {
130- const originalDescribeFn = originalVitestFn . apply ( this , factoryArgs ) ;
131- return function ( this : unknown , ...args : [ unknown , Function , ...unknown [ ] ] ) {
132- args [ 1 ] = wrapDescribeInZone ( args [ 1 ] ) ;
133- return originalDescribeFn . apply ( this , args ) ;
134- } ;
124+ // Patching direct modifier calls
125+ for ( const modifierName of DIRECT_MODIFIER_NAMES ) {
126+ context [ methodName ] [ modifierName ] = function (
127+ this : unknown ,
128+ ...args : [ unknown , Function , ...unknown [ ] ]
129+ ) {
130+ args [ 1 ] = wrapFn ( args [ 1 ] ) ;
131+ return originalVitestFn [ modifierName ] . apply ( this , args ) ;
135132 } ;
136133 }
137- } ) ;
138134
139- [ 'it' , 'test' ] . forEach ( ( methodName ) => {
140- let originalVitestFn : Function & Record < ( typeof TEST_FACTORY_NAMES ) [ number ] , Function > =
141- context [ methodName ] ;
142- // Skip if already patched
143- if ( context [ Zone . __symbol__ ( methodName ) ] ) {
144- return ;
145- }
135+ // Patching curried modifier calls
136+ for ( const modifierName of CURRIED_MODIFIER_NAMES ) {
137+ context [ methodName ] [ modifierName ] = function ( this : unknown , ... modifierArgs : unknown [ ] ) {
138+ // Since we are patching a curried function, we need
139+ // to pass the original context first (`originalVitestFn`).
140+ // Else, the chaining won't be possible (will get an error).
141+ const originalFn = originalVitestFn [ modifierName ] . apply ( originalVitestFn , modifierArgs ) ;
146142
147- context [ Zone . __symbol__ ( methodName ) ] = originalVitestFn ;
148- context [ methodName ] = function ( this : unknown , ...args : [ unknown , Function , ...unknown [ ] ] ) {
149- args [ 1 ] = wrapTestInZone ( args [ 1 ] ) ;
150- return originalVitestFn . apply ( this , args ) ;
151- } ;
152-
153- for ( const factoryName of TEST_FACTORY_NAMES ) {
154- context [ methodName ] [ factoryName ] = function ( this : unknown , ...factoryArgs : unknown [ ] ) {
155143 return function ( this : unknown , ...args : [ unknown , Function , ...unknown [ ] ] ) {
156- args [ 1 ] = wrapTestInZone ( args [ 1 ] ) ;
157- return originalVitestFn . apply ( this , factoryArgs ) . apply ( this , args ) ;
144+ args [ 1 ] = wrapFn ( args [ 1 ] ) ;
145+ return originalFn . apply ( this , args ) ;
158146 } ;
159147 } ;
160148 }
161- } ) ;
149+ }
150+
151+ [ 'suite' , 'describe' ] . forEach ( ( methodName ) =>
152+ patchFnWithModifiers ( methodName , wrapDescribeInZone ) ,
153+ ) ;
154+
155+ [ 'it' , 'test' ] . forEach ( ( methodName ) => patchFnWithModifiers ( methodName , wrapTestInZone ) ) ;
162156
163157 [ 'beforeEach' , 'afterEach' , 'beforeAll' , 'afterAll' ] . forEach ( ( methodName ) => {
164158 const originalVitestFn : Function = context [ methodName ] ;
0 commit comments