@@ -810,6 +810,13 @@ def execute(
810810 if output_file :
811811 output_file .metadata ["parcels_kernels" ] = self ._kernel .name
812812
813+ if np .isnat (dt ) or dt is None :
814+ dt = np .timedelta64 (1 , "s" )
815+ self ._data ["dt" ][:] = dt
816+ sign_dt = np .sign (dt ).astype (int )
817+ if sign_dt not in [- 1 , 1 ]:
818+ raise ValueError ("dt must be a positive or negative np.timedelta64 object" )
819+
813820 if self .fieldset .time_interval is None :
814821 start_time = np .timedelta64 (0 , "s" ) # For the execution loop, we need a start time as a timedelta object
815822 if runtime is None :
@@ -831,18 +838,23 @@ def execute(
831838 )
832839 # Ensure that the endtime uses the same type as the start_time
833840 if isinstance (endtime , self .fieldset .time_interval .left .__class__ ):
834- if endtime < self .fieldset .time_interval .left :
835- raise ValueError ("The endtime must be after the start time of the fieldset.time_interval" )
836- end_time = min (endtime , self .fieldset .time_interval .right )
841+ if sign_dt > 0 :
842+ if endtime < self .fieldset .time_interval .left :
843+ raise ValueError ("The endtime must be after the start time of the fieldset.time_interval" )
844+ end_time = min (endtime , self .fieldset .time_interval .right )
845+ else :
846+ if endtime > self .fieldset .time_interval .right :
847+ raise ValueError (
848+ "The endtime must be before the end time of the fieldset.time_interval when dt < 0"
849+ )
850+ end_time = max (endtime , self .fieldset .time_interval .left )
837851 else :
838852 raise TypeError ("The endtime must be of the same type as the fieldset.time_interval start time." )
839853 else :
840- end_time = start_time + runtime
854+ end_time = start_time + runtime * sign_dt
841855
842856 outputdt = output_file .outputdt if output_file else None
843857
844- self ._data ["dt" ][:] = dt
845-
846858 # Set up pbar
847859 if output_file :
848860 logger .info (f"Output files are stored in { output_file .fname } ." )
@@ -853,8 +865,11 @@ def execute(
853865 next_output = outputdt if output_file else None
854866
855867 time = start_time
856- while time < end_time :
857- next_time = min (time + dt , end_time ) # TODO also for time-backward
868+ while sign_dt * (time - end_time ) < 0 :
869+ if sign_dt > 0 :
870+ next_time = min (time + dt , end_time )
871+ else :
872+ next_time = max (time + dt , end_time )
858873 res = self ._kernel .execute (self , endtime = next_time , dt = dt )
859874 if res == StatusCode .StopAllExecution :
860875 return StatusCode .StopAllExecution
0 commit comments