1515# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
1616# See the License for the specific language governing permissions and
1717# limitations under the License.
18-
19-
18+ import dataclasses
2019import functools
2120import sys
2221import threading
9796WrappedFn = t .TypeVar ("WrappedFn" , bound = t .Callable [..., t .Any ])
9897
9998
99+ dataclass_kwargs = {}
100+ if sys .version_info >= (3 , 10 ):
101+ dataclass_kwargs .update ({"slots" : True })
102+
103+
104+ @dataclasses .dataclass (** dataclass_kwargs )
105+ class IterState :
106+ actions : t .List [t .Callable [["RetryCallState" ], t .Any ]] = dataclasses .field (
107+ default_factory = list
108+ )
109+ retry_run_result : bool = False
110+ delay_since_first_attempt : int = 0
111+ stop_run_result : bool = False
112+ is_explicit_retry : bool = False
113+
114+ def reset (self ) -> None :
115+ self .actions = []
116+ self .retry_run_result = False
117+ self .delay_since_first_attempt = 0
118+ self .stop_run_result = False
119+ self .is_explicit_retry = False
120+
121+
100122class TryAgain (Exception ):
101123 """Always retry the executed function when raised."""
102124
@@ -287,6 +309,14 @@ def statistics(self) -> t.Dict[str, t.Any]:
287309 self ._local .statistics = t .cast (t .Dict [str , t .Any ], {})
288310 return self ._local .statistics
289311
312+ @property
313+ def iter_state (self ) -> IterState :
314+ try :
315+ return self ._local .iter_state # type: ignore[no-any-return]
316+ except AttributeError :
317+ self ._local .iter_state = IterState ()
318+ return self ._local .iter_state
319+
290320 def wraps (self , f : WrappedFn ) -> WrappedFn :
291321 """Wrap a function for retrying.
292322
@@ -313,45 +343,89 @@ def begin(self) -> None:
313343 self .statistics ["attempt_number" ] = 1
314344 self .statistics ["idle_for" ] = 0
315345
316- def iter (self , retry_state : "RetryCallState" ) -> t .Union [DoAttempt , DoSleep , t .Any ]: # noqa
317- fut = retry_state .outcome
318- if fut is None :
319- if self .before is not None :
320- self .before (retry_state )
321- return DoAttempt ()
322-
323- is_explicit_retry = fut .failed and isinstance (fut .exception (), TryAgain )
324- if not (is_explicit_retry or self .retry (retry_state )):
325- return fut .result ()
346+ def _add_action_func (self , fn : t .Callable [..., t .Any ]) -> None :
347+ self .iter_state .actions .append (fn )
326348
327- if self . after is not None :
328- self .after (retry_state )
349+ def _run_retry ( self , retry_state : "RetryCallState" ) -> None :
350+ self . iter_state . retry_run_result = self .retry (retry_state )
329351
352+ def _run_wait (self , retry_state : "RetryCallState" ) -> None :
330353 if self .wait :
331354 sleep = self .wait (retry_state )
332355 else :
333356 sleep = 0.0
334357
335358 retry_state .upcoming_sleep = sleep
336359
360+ def _run_stop (self , retry_state : "RetryCallState" ) -> None :
337361 self .statistics ["delay_since_first_attempt" ] = retry_state .seconds_since_start
338- if self .stop (retry_state ):
362+ self .iter_state .stop_run_result = self .stop (retry_state )
363+
364+ def iter (self , retry_state : "RetryCallState" ) -> t .Union [DoAttempt , DoSleep , t .Any ]: # noqa
365+ self ._begin_iter (retry_state )
366+ result = None
367+ for action in self .iter_state .actions :
368+ result = action (retry_state )
369+ return result
370+
371+ def _begin_iter (self , retry_state : "RetryCallState" ) -> None : # noqa
372+ self .iter_state .reset ()
373+
374+ fut = retry_state .outcome
375+ if fut is None :
376+ if self .before is not None :
377+ self ._add_action_func (self .before )
378+ self ._add_action_func (lambda rs : DoAttempt ())
379+ return
380+
381+ self .iter_state .is_explicit_retry = fut .failed and isinstance (
382+ fut .exception (), TryAgain
383+ )
384+ if not self .iter_state .is_explicit_retry :
385+ self ._add_action_func (self ._run_retry )
386+ self ._add_action_func (self ._post_retry_check_actions )
387+
388+ def _post_retry_check_actions (self , retry_state : "RetryCallState" ) -> None :
389+ if not (self .iter_state .is_explicit_retry or self .iter_state .retry_run_result ):
390+ self ._add_action_func (lambda rs : rs .outcome .result ())
391+ return
392+
393+ if self .after is not None :
394+ self ._add_action_func (self .after )
395+
396+ self ._add_action_func (self ._run_wait )
397+ self ._add_action_func (self ._run_stop )
398+ self ._add_action_func (self ._post_stop_check_actions )
399+
400+ def _post_stop_check_actions (self , retry_state : "RetryCallState" ) -> None :
401+ if self .iter_state .stop_run_result :
339402 if self .retry_error_callback :
340- return self .retry_error_callback (retry_state )
341- retry_exc = self .retry_error_cls (fut )
342- if self .reraise :
343- raise retry_exc .reraise ()
344- raise retry_exc from fut .exception ()
403+ self ._add_action_func (self .retry_error_callback )
404+ return
405+
406+ def exc_check (rs : "RetryCallState" ) -> None :
407+ fut = t .cast (Future , rs .outcome )
408+ retry_exc = self .retry_error_cls (fut )
409+ if self .reraise :
410+ raise retry_exc .reraise ()
411+ raise retry_exc from fut .exception ()
412+
413+ self ._add_action_func (exc_check )
414+ return
415+
416+ def next_action (rs : "RetryCallState" ) -> None :
417+ sleep = rs .upcoming_sleep
418+ rs .next_action = RetryAction (sleep )
419+ rs .idle_for += sleep
420+ self .statistics ["idle_for" ] += sleep
421+ self .statistics ["attempt_number" ] += 1
345422
346- retry_state .next_action = RetryAction (sleep )
347- retry_state .idle_for += sleep
348- self .statistics ["idle_for" ] += sleep
349- self .statistics ["attempt_number" ] += 1
423+ self ._add_action_func (next_action )
350424
351425 if self .before_sleep is not None :
352- self .before_sleep ( retry_state )
426+ self ._add_action_func ( self . before_sleep )
353427
354- return DoSleep (sleep )
428+ self . _add_action_func ( lambda rs : DoSleep (rs . upcoming_sleep ) )
355429
356430 def __iter__ (self ) -> t .Generator [AttemptManager , None , None ]:
357431 self .begin ()
0 commit comments