|
| 1 | +import functools |
1 | 2 | import logging |
2 | 3 | import os |
| 4 | +import random |
| 5 | +import time |
3 | 6 | from collections.abc import AsyncGenerator, Awaitable, Callable |
4 | 7 | from contextlib import asynccontextmanager |
5 | 8 | from dataclasses import dataclass, field |
@@ -157,6 +160,76 @@ def shutdown_runtime_sidecar( |
157 | 160 | return True |
158 | 161 |
|
159 | 162 |
|
| 163 | +# Streams retry the unary transient codes plus INTERNAL/UNKNOWN: the Go |
| 164 | +# controller can surface those during rolling updates (e.g. INTERNAL when |
| 165 | +# the context is cancelled mid-send). |
| 166 | +_RETRYABLE_STREAM_CODES = _TRANSIENT_GRPC_CODES | frozenset({ |
| 167 | + grpc.StatusCode.INTERNAL, |
| 168 | + grpc.StatusCode.UNKNOWN, |
| 169 | +}) |
| 170 | + |
| 171 | + |
| 172 | +class _StreamClosedImmediately(Exception): |
| 173 | + """Stream connected and returned zero items; treated as retryable degradation.""" |
| 174 | + |
| 175 | + |
| 176 | +def _is_retryable(e: Exception) -> bool: |
| 177 | + """Classify whether a streaming error warrants retry or is terminal.""" |
| 178 | + if isinstance(e, _StreamClosedImmediately): |
| 179 | + return True |
| 180 | + if isinstance(e, grpc.aio.AioRpcError): |
| 181 | + return e.code() in _RETRYABLE_STREAM_CODES |
| 182 | + if isinstance(e, ConnectionError): |
| 183 | + return True |
| 184 | + return False |
| 185 | + |
| 186 | + |
| 187 | +@dataclass |
| 188 | +class _GraceWindow: |
| 189 | + """Wall-clock degradation window for stream retries.""" |
| 190 | + |
| 191 | + period: float |
| 192 | + since: float | None = field(default=None, init=False) |
| 193 | + |
| 194 | + def mark_failure(self) -> float: |
| 195 | + now = time.monotonic() |
| 196 | + if self.since is None: |
| 197 | + self.since = now |
| 198 | + return now - self.since |
| 199 | + |
| 200 | + def elapsed(self) -> float: |
| 201 | + if self.since is None: |
| 202 | + return 0.0 |
| 203 | + return time.monotonic() - self.since |
| 204 | + |
| 205 | + def expired(self) -> bool: |
| 206 | + return self.since is not None and time.monotonic() - self.since > self.period |
| 207 | + |
| 208 | + def reset(self): |
| 209 | + self.since = None |
| 210 | + |
| 211 | + |
| 212 | +@dataclass |
| 213 | +class _Backoff: |
| 214 | + """Exponential backoff with jitter for stream retries.""" |
| 215 | + |
| 216 | + max_delay: float |
| 217 | + delay: float = field(default=0.5, init=False) |
| 218 | + _initial: float = field(default=0.5, init=False) |
| 219 | + |
| 220 | + def __post_init__(self): |
| 221 | + self._initial = min(0.5, self.max_delay) |
| 222 | + self.delay = self._initial |
| 223 | + |
| 224 | + def reset(self): |
| 225 | + self.delay = self._initial |
| 226 | + |
| 227 | + async def wait(self): |
| 228 | + jitter = random.uniform(0, self.delay * 0.3) |
| 229 | + await sleep(self.delay + jitter) |
| 230 | + self.delay = min(self.delay * 2, self.max_delay) |
| 231 | + |
| 232 | + |
160 | 233 | class LeaseState(Enum): |
161 | 234 | IDLE = "idle" |
162 | 235 | LEASED = "leased" |
@@ -369,6 +442,9 @@ class Exporter(AsyncContextManagerMixin, Metadata): |
369 | 442 | _status_rpc_event: Event = field(init=False, default_factory=Event) |
370 | 443 | """Signals the drain task that a new status update is pending.""" |
371 | 444 |
|
| 445 | + _fatal_stream_error: tuple[str, Exception] | None = field(init=False, default=None) |
| 446 | + """Set by _cancel_with_fatal_error when a stream hits a terminal error.""" |
| 447 | + |
372 | 448 | @property |
373 | 449 | def _lease_state(self) -> LeaseState: |
374 | 450 | return LeaseState.LEASED if self._lease_context is not None else LeaseState.IDLE |
@@ -423,55 +499,148 @@ async def _controller_stub(self) -> AsyncGenerator[jumpstarter_pb2_grpc.Controll |
423 | 499 | finally: |
424 | 500 | await channel.close() |
425 | 501 |
|
| 502 | + def _cancel_with_fatal_error(self, stream_name: str, error: Exception): |
| 503 | + self._fatal_stream_error = (stream_name, error) |
| 504 | + if self._tg is not None: |
| 505 | + self._tg.cancel_scope.cancel() |
| 506 | + |
| 507 | + def _on_status_exhausted(self, stream_name: str, error: Exception): |
| 508 | + pass |
| 509 | + |
| 510 | + async def _stream_once( |
| 511 | + self, |
| 512 | + stream_name: str, |
| 513 | + stream_factory: Callable[[jumpstarter_pb2_grpc.ControllerServiceStub], AsyncGenerator], |
| 514 | + send_tx: MemoryObjectSendStream[Any], |
| 515 | + window: _GraceWindow, |
| 516 | + backoff: _Backoff, |
| 517 | + ) -> Exception | None: |
| 518 | + """Run one stream connection attempt. |
| 519 | +
|
| 520 | + Returns None if data was yielded (window/backoff reset inline), |
| 521 | + or the failure exception for the caller to handle. |
| 522 | + Raises ClosedResourceError/BrokenResourceError for channel closure. |
| 523 | + """ |
| 524 | + yielded_items = False |
| 525 | + try: |
| 526 | + async with self._controller_stub() as controller: |
| 527 | + logger.debug("%s stream connected to controller", stream_name) |
| 528 | + async for item in stream_factory(controller): |
| 529 | + yielded_items = True |
| 530 | + if window.since is not None: |
| 531 | + logger.info( |
| 532 | + "%s stream recovered after %.1fs", |
| 533 | + stream_name, |
| 534 | + window.elapsed(), |
| 535 | + ) |
| 536 | + window.reset() |
| 537 | + backoff.reset() |
| 538 | + await send_tx.send(item) |
| 539 | + except (anyio.ClosedResourceError, anyio.BrokenResourceError): |
| 540 | + raise |
| 541 | + except Exception as e: |
| 542 | + return e |
| 543 | + else: |
| 544 | + if yielded_items: |
| 545 | + window.reset() |
| 546 | + backoff.reset() |
| 547 | + return None |
| 548 | + return _StreamClosedImmediately( |
| 549 | + f"{stream_name} stream closed immediately" |
| 550 | + ) |
| 551 | + |
426 | 552 | async def _retry_stream( |
427 | 553 | self, |
428 | 554 | stream_name: str, |
429 | 555 | stream_factory: Callable[[jumpstarter_pb2_grpc.ControllerServiceStub], AsyncGenerator], |
430 | | - send_tx, |
431 | | - retries: int = 5, |
432 | | - backoff: float = 1.0, # Reduced from 3.0 for faster recovery from transient errors |
433 | | - ): |
434 | | - """Generic retry wrapper for gRPC streaming calls. |
| 556 | + send_tx: MemoryObjectSendStream[Any], |
| 557 | + grace_period: float = 300.0, |
| 558 | + max_backoff: float = 10.0, |
| 559 | + on_terminal: Callable[[str, Exception], None] | None = None, |
| 560 | + on_exhausted: Callable[[str, Exception], None] | None = None, |
| 561 | + ) -> None: |
| 562 | + """Resilient retry wrapper for gRPC streaming calls. |
435 | 563 |
|
436 | | - Args: |
437 | | - stream_name: Name of the stream for logging purposes |
438 | | - stream_factory: Function that takes a controller stub and returns an async generator |
439 | | - send_tx: Transmission channel to send stream items to |
440 | | - retries: Maximum number of retry attempts |
441 | | - backoff: Seconds to wait between retries |
| 564 | + Retries for up to grace_period seconds after the first failure, with |
| 565 | + exponential backoff and jitter. Data flowing through resets the window. |
| 566 | + Terminal (non-retryable) errors invoke on_terminal immediately. |
| 567 | + When on_exhausted is set, grace window expiry calls it, resets the |
| 568 | + window, and continues retrying instead of stopping. |
442 | 569 | """ |
443 | | - retries_left = retries |
444 | | - while True: |
445 | | - received_data = False |
446 | | - try: |
447 | | - async with self._controller_stub() as controller: |
448 | | - logger.debug("%s stream connected to controller", stream_name) |
449 | | - async for item in stream_factory(controller): |
450 | | - received_data = True |
451 | | - logger.debug("%s stream received item", stream_name) |
452 | | - await send_tx.send(item) |
453 | | - except Exception as e: |
454 | | - if received_data: |
455 | | - logger.debug("%s stream retry counter reset after receiving data", stream_name) |
456 | | - retries_left = retries |
457 | | - if retries_left > 0: |
458 | | - retries_left -= 1 |
459 | | - # Check for common transient errors that warrant faster retry |
460 | | - error_str = str(e) |
461 | | - is_transient = "Stream removed" in error_str or "UNAVAILABLE" in error_str |
462 | | - retry_delay = 0.5 if is_transient else backoff |
463 | | - logger.info( |
464 | | - "%s stream interrupted, restarting in %ss, %s retries left: %s", |
| 570 | + if on_terminal is None: |
| 571 | + on_terminal = self._cancel_with_fatal_error |
| 572 | + window = _GraceWindow(grace_period) |
| 573 | + backoff = _Backoff(max_backoff) |
| 574 | + warned = False |
| 575 | + |
| 576 | + async with send_tx: |
| 577 | + while True: |
| 578 | + try: |
| 579 | + failure = await self._stream_once( |
| 580 | + stream_name, stream_factory, send_tx, window, backoff |
| 581 | + ) |
| 582 | + except (anyio.ClosedResourceError, anyio.BrokenResourceError): |
| 583 | + logger.debug("%s send channel closed, exiting", stream_name) |
| 584 | + return |
| 585 | + |
| 586 | + if failure is None: |
| 587 | + warned = False |
| 588 | + await backoff.wait() |
| 589 | + continue |
| 590 | + |
| 591 | + if not _is_retryable(failure): |
| 592 | + logger.error("%s stream hit terminal error: %s", stream_name, failure) |
| 593 | + on_terminal(stream_name, failure) |
| 594 | + return |
| 595 | + |
| 596 | + fresh = window.since is None |
| 597 | + degraded = window.mark_failure() |
| 598 | + if window.expired(): |
| 599 | + if on_exhausted is not None: |
| 600 | + logger.warning( |
| 601 | + "%s stream unavailable for %.0fs, still retrying: %s", |
| 602 | + stream_name, |
| 603 | + degraded, |
| 604 | + failure, |
| 605 | + ) |
| 606 | + on_exhausted(stream_name, failure) |
| 607 | + window.reset() |
| 608 | + backoff.reset() |
| 609 | + warned = True |
| 610 | + await backoff.wait() |
| 611 | + continue |
| 612 | + else: |
| 613 | + logger.error( |
| 614 | + "%s stream failed after %.1fs grace period: %s", |
| 615 | + stream_name, |
| 616 | + degraded, |
| 617 | + failure, |
| 618 | + ) |
| 619 | + on_terminal(stream_name, failure) |
| 620 | + return |
| 621 | + |
| 622 | + if fresh: |
| 623 | + warned = False |
| 624 | + if not warned: |
| 625 | + warned = True |
| 626 | + logger.warning( |
| 627 | + "%s stream degraded, retrying in %.1fs for %.0fs: %s", |
465 | 628 | stream_name, |
466 | | - retry_delay, |
467 | | - retries_left, |
468 | | - e, |
| 629 | + backoff.delay, |
| 630 | + grace_period, |
| 631 | + failure, |
469 | 632 | ) |
470 | | - await sleep(retry_delay) |
471 | 633 | else: |
472 | | - raise |
473 | | - else: |
474 | | - retries_left = retries |
| 634 | + logger.info( |
| 635 | + "%s stream retrying in %.1fs (degraded %.1fs/%.0fs): %s", |
| 636 | + stream_name, |
| 637 | + backoff.delay, |
| 638 | + degraded, |
| 639 | + grace_period, |
| 640 | + failure, |
| 641 | + ) |
| 642 | + |
| 643 | + await backoff.wait() |
475 | 644 |
|
476 | 645 | def _listen_stream_factory( |
477 | 646 | self, lease_name: str |
@@ -1123,14 +1292,16 @@ async def handle_lease(self, lease_name: str, tg: TaskGroup, lease_scope: LeaseC |
1123 | 1292 | # Type: request is jumpstarter_pb2.ListenResponse with router_endpoint and router_token fields |
1124 | 1293 | try: |
1125 | 1294 | async with create_task_group() as conn_tg: |
1126 | | - # Start listening for connection requests with retry logic |
1127 | | - # This is inside conn_tg so it gets cancelled when the lease ends |
1128 | | - conn_tg.start_soon( |
| 1295 | + conn_tg.start_soon(functools.partial( |
1129 | 1296 | self._retry_stream, |
1130 | | - "Listen", |
1131 | | - self._listen_stream_factory(lease_name), |
1132 | | - listen_tx, |
1133 | | - ) |
| 1297 | + stream_name="Listen", |
| 1298 | + stream_factory=self._listen_stream_factory(lease_name), |
| 1299 | + send_tx=listen_tx, |
| 1300 | + on_terminal=lambda name, err: ( |
| 1301 | + logger.info("Listen stream ended (%s: %s), signaling lease end", name, err), |
| 1302 | + lease_scope.lease_ended.set(), |
| 1303 | + ), |
| 1304 | + )) |
1134 | 1305 |
|
1135 | 1306 | async def wait_for_lease_end(): |
1136 | 1307 | """Wait for lease_ended event and cancel the connection loop.""" |
@@ -1213,13 +1384,21 @@ async def serve(self): |
1213 | 1384 | status_tx, status_rx = create_memory_object_stream[jumpstarter_pb2.StatusResponse](max_buffer_size=5) |
1214 | 1385 | try: |
1215 | 1386 | await self._run_control_plane(status_tx, status_rx) |
| 1387 | + if self._fatal_stream_error: |
| 1388 | + name, err = self._fatal_stream_error |
| 1389 | + logger.warning( |
| 1390 | + "Control plane down (%s: %s)", |
| 1391 | + name, |
| 1392 | + err, |
| 1393 | + ) |
1216 | 1394 | finally: |
1217 | 1395 | if self.exit_on_lease_end: |
1218 | 1396 | # Ensure the runtime container exits whenever this exporter is |
1219 | 1397 | # configured for ExitAndReplace (covers hook on_failure=exit and |
1220 | 1398 | # other stop paths that skip the lease-end branch above). |
1221 | 1399 | await anyio.to_thread.run_sync(shutdown_runtime_sidecar) |
1222 | 1400 | self._tg = None |
| 1401 | + self._fatal_stream_error = None |
1223 | 1402 | self._status_drain_active = False |
1224 | 1403 | clear_log_context() |
1225 | 1404 |
|
@@ -1247,12 +1426,13 @@ async def _run_control_plane( |
1247 | 1426 | tg.start_soon(self._drain_status_reports) |
1248 | 1427 | if self._telemetry_handler is not None: |
1249 | 1428 | tg.start_soon(self._telemetry_handler.flush_loop) |
1250 | | - tg.start_soon( |
| 1429 | + tg.start_soon(functools.partial( |
1251 | 1430 | self._retry_stream, |
1252 | 1431 | "Status", |
1253 | 1432 | self._status_stream_factory(), |
1254 | 1433 | status_tx, |
1255 | | - ) |
| 1434 | + on_exhausted=self._on_status_exhausted, |
| 1435 | + )) |
1256 | 1436 | async for status in status_rx: |
1257 | 1437 | if await self._apply_status(status, tg): |
1258 | 1438 | break |
|
0 commit comments