diff --git a/packages/data-designer-engine/src/data_designer/engine/dataset_builders/async_scheduler.py b/packages/data-designer-engine/src/data_designer/engine/dataset_builders/async_scheduler.py index 8b6be7889..5fd3ce6e8 100644 --- a/packages/data-designer-engine/src/data_designer/engine/dataset_builders/async_scheduler.py +++ b/packages/data-designer-engine/src/data_designer/engine/dataset_builders/async_scheduler.py @@ -68,8 +68,11 @@ from data_designer.engine.models.errors import ( RETRYABLE_MODEL_ERRORS, GenerationValidationFailureError, + ModelAPIConnectionError, + ModelInternalServerError, ModelRateLimitError, ModelRequestAdmissionTimeoutError, + ModelTimeoutError, ) from data_designer.engine.models.request_admission.config import RequestAdmissionConfig from data_designer.engine.models.request_admission.resources import RequestResourceKey @@ -308,6 +311,10 @@ def __init__( self._degraded_warn_window = degraded_warn_window self._degraded_warn_interval_s = degraded_warn_interval_s self._recent_retryable: deque[bool] = deque(maxlen=degraded_warn_window) + self._recent_retryable_kinds: deque[str] = deque(maxlen=degraded_warn_window) + self._retryable_outcome_counts: Counter[str] = Counter() + self._retryable_detail_counts: Counter[str] = Counter() + self._logged_retryable_kinds: set[str] = set() # Initialize to -inf so the first WARN is always emitted regardless of # the monotonic clock's absolute value (which can be near-zero on freshly # booted CI runners). @@ -580,6 +587,19 @@ def _scheduler_health_diagnostics(self, *, reason: str) -> dict[str, object]: "request_pressure_advisory_skips": self._request_pressure_advisory_skips, "row_group_admission_blocked_reasons": dict(self._row_group_admission_blocked_reasons), "request_pressure": self._request_pressure_diagnostics(), + "retryable_outcomes": self.retryable_outcome_metrics, + } + + @property + def retryable_outcome_metrics(self) -> dict[str, object]: + """Return sanitized rolling and cumulative model-task outcome counts.""" + rolling = Counter(self._recent_retryable_kinds) + return { + "window_size": self._degraded_warn_window, + "rolling_total": len(self._recent_retryable_kinds), + "rolling_counts": dict(sorted(rolling.items())), + "cumulative_counts": dict(sorted(self._retryable_outcome_counts.items())), + "retryable_details": dict(sorted(self._retryable_detail_counts.items())), } def _scheduler_job_diagnostics(self) -> dict[str, object]: @@ -1778,7 +1798,38 @@ def _check_error_rate(self, *, success: bool) -> None: if errors / self._shutdown_error_window >= self._shutdown_error_rate: self._early_shutdown = True - def _record_retryable_outcome(self, *, retryable: bool) -> None: + @staticmethod + def _retryable_error_kind(exc: Exception | None) -> str: + if isinstance(exc, ModelRequestAdmissionTimeoutError): + return "request_admission_timeout" + if isinstance(exc, ModelRateLimitError): + return "rate_limit" + if isinstance(exc, ModelTimeoutError): + return "timeout" + if isinstance(exc, ModelInternalServerError): + return "internal_server" + if isinstance(exc, ModelAPIConnectionError): + return "connection" + return "other_retryable" + + @staticmethod + def _provider_error_details(exc: Exception | None) -> tuple[int | None, float | None]: + """Extract only safe transport metadata from a wrapped provider error.""" + seen: set[int] = set() + current: BaseException | None = exc + while current is not None and id(current) not in seen: + seen.add(id(current)) + if isinstance(current, ProviderError): + return current.status_code, current.retry_after + current = current.__cause__ or current.__context__ + return None, None + + def _record_retryable_outcome( + self, + *, + retryable: bool, + exc: Exception | None = None, + ) -> None: """Track retryable-error rate and emit a rate-limited WARN under provider degradation. Distinct from ``_check_error_rate``: every LLM-bound task outcome (success @@ -1786,11 +1837,36 @@ def _record_retryable_outcome(self, *, retryable: bool) -> None: health, not just the error mix. The call site filters on ``is_llm`` so non-LLM tasks (samplers, expressions, non-LLM customs) don't dilute the rate. Only retryable errors (rate-limit, timeout, 5xx, connection) count - toward the rate; non-retryable failures register as 0. + toward the rate; non-retryable failures register as 0 and are categorized + separately from successful outcomes. """ + if retryable: + kind = self._retryable_error_kind(exc) + elif exc is not None: + kind = "non_retryable_failure" + else: + kind = "success" + self._retryable_outcome_counts[kind] += 1 + if retryable: + status_code, retry_after = self._provider_error_details(exc) + detail = kind + if status_code is not None: + detail = f"{detail}:http_{status_code}" + self._retryable_detail_counts[detail] += 1 + if kind not in self._logged_retryable_kinds: + self._logged_retryable_kinds.add(kind) + retry_after_text = f", retry_after={retry_after:g}s" if retry_after is not None else "" + status_text = f", http_status={status_code}" if status_code is not None else "" + logger.warning( + "Observed retryable model-task error: kind=%s%s%s; the row task will be deferred.", + kind, + status_text, + retry_after_text, + ) if self._degraded_warn_window <= 0: return self._recent_retryable.append(retryable) + self._recent_retryable_kinds.append(kind) if len(self._recent_retryable) < self._degraded_warn_window: return rate = sum(self._recent_retryable) / self._degraded_warn_window @@ -1801,10 +1877,20 @@ def _record_retryable_outcome(self, *, retryable: bool) -> None: return self._last_degraded_warn_at = now pct = int(round(rate * 100)) + rolling = Counter(self._recent_retryable_kinds) + rolling_summary = ", ".join(f"{key}={rolling[key]}" for key in sorted(rolling)) + cumulative_summary = ", ".join( + f"{key}={self._retryable_outcome_counts[key]}" for key in sorted(self._retryable_outcome_counts) + ) logger.warning( - f"Provider showing degraded performance: {pct}% of last {self._degraded_warn_window} " - "task outcomes were retryable errors (rate-limit, timeout, 5xx, connection). " - "Run may take longer than expected; salvage will retry these." + "Provider showing degraded performance: %s%% of last %s task outcomes were " + "retryable errors (%s); cumulative outcomes: %s; deferred_tasks=%s. " + "Run may take longer than expected; salvage will retry these.", + pct, + self._degraded_warn_window, + rolling_summary, + cumulative_summary, + len(self._deferred), ) async def _dispatch_seeds(self, rg_id: int, rg_size: int) -> None: @@ -1941,7 +2027,7 @@ async def _execute_task_inner_impl(self, task: Task, lease: TaskAdmissionLease, if not retryable: self._check_error_rate(success=False) if uses_model_stage_resource: - self._record_retryable_outcome(retryable=retryable) + self._record_retryable_outcome(retryable=retryable, exc=exc) if not retryable and self._reporter: self._reporter.record_failure(task.column) if self._trace and trace: diff --git a/packages/data-designer-engine/tests/engine/dataset_builders/test_async_scheduler.py b/packages/data-designer-engine/tests/engine/dataset_builders/test_async_scheduler.py index 8de0e4cee..94aeca3e4 100644 --- a/packages/data-designer-engine/tests/engine/dataset_builders/test_async_scheduler.py +++ b/packages/data-designer-engine/tests/engine/dataset_builders/test_async_scheduler.py @@ -45,12 +45,15 @@ from data_designer.engine.dataset_builders.scheduling.task_policies import BoundedBorrowTaskAdmissionPolicyConfig from data_designer.engine.dataset_builders.utils.execution_graph import ExecutionGraph from data_designer.engine.dataset_builders.utils.row_group_buffer import RowGroupBufferManager +from data_designer.engine.models.clients.errors import ProviderError, ProviderErrorKind from data_designer.engine.models.errors import ( RETRYABLE_MODEL_ERRORS, + ModelAPIConnectionError, ModelInternalServerError, ModelRateLimitError, ModelRequestAdmissionTimeoutError, ModelTimeoutError, + handle_llm_exceptions, ) from data_designer.engine.models.request_admission.config import RequestAdmissionConfig from data_designer.engine.models.request_admission.controller import ( @@ -1722,6 +1725,87 @@ async def test_degraded_provider_warn_silent_under_threshold(caplog: pytest.LogC assert _count_degraded_msgs(caplog) == 0 +@pytest.mark.parametrize( + ("exc_cls", "expected_kind"), + [ + (ModelRateLimitError, "rate_limit"), + (ModelRequestAdmissionTimeoutError, "request_admission_timeout"), + (ModelTimeoutError, "timeout"), + (ModelInternalServerError, "internal_server"), + (ModelAPIConnectionError, "connection"), + ], +) +def test_retryable_outcome_metrics_classify_errors( + exc_cls: type[Exception], + expected_kind: str, +) -> None: + scheduler, _tracker = _build_simple_pipeline(num_records=1) + + scheduler._record_retryable_outcome( + retryable=True, + exc=exc_cls("sensitive provider text"), + ) + scheduler._record_retryable_outcome(retryable=False) + + metrics = scheduler.retryable_outcome_metrics + assert metrics["cumulative_counts"] == {expected_kind: 1, "success": 1} + assert metrics["rolling_counts"] == {expected_kind: 1, "success": 1} + assert metrics["retryable_details"] == {expected_kind: 1} + assert "sensitive provider text" not in str(metrics) + + +def test_retryable_outcome_metrics_extract_provider_status_from_suppressed_context( + caplog: pytest.LogCaptureFixture, +) -> None: + scheduler, _tracker = _build_simple_pipeline(num_records=1) + provider_error = ProviderError( + kind=ProviderErrorKind.RATE_LIMIT, + message="sensitive provider text", + status_code=429, + retry_after=2.5, + ) + + try: + raise provider_error + except ProviderError as exc: + with pytest.raises(ModelRateLimitError) as exc_info: + handle_llm_exceptions(exc, MODEL_ALIAS, "stub-provider") + + assert exc_info.value.__cause__ is None + assert exc_info.value.__context__ is provider_error + with caplog.at_level(logging.WARNING): + scheduler._record_retryable_outcome(retryable=True, exc=exc_info.value) + + metrics = scheduler.retryable_outcome_metrics + assert metrics["retryable_details"] == {"rate_limit:http_429": 1} + assert "http_status=429" in caplog.text + assert "retry_after=2.5s" in caplog.text + assert "sensitive provider text" not in str(metrics) + assert "sensitive provider text" not in caplog.text + + +def test_retryable_outcome_metrics_distinguish_non_retryable_failures() -> None: + scheduler, _tracker = _build_simple_pipeline(num_records=1) + + scheduler._record_retryable_outcome( + retryable=False, + exc=RuntimeError("sensitive non-retryable text"), + ) + scheduler._record_retryable_outcome(retryable=False) + + metrics = scheduler.retryable_outcome_metrics + assert metrics["cumulative_counts"] == {"non_retryable_failure": 1, "success": 1} + assert metrics["rolling_counts"] == {"non_retryable_failure": 1, "success": 1} + assert metrics["retryable_details"] == {} + assert "sensitive non-retryable text" not in str(metrics) + + diagnostics = scheduler._scheduler_health_diagnostics(reason="test") + assert diagnostics["deferred_tasks"] == 0 + retryable_outcomes = diagnostics["retryable_outcomes"] + assert isinstance(retryable_outcomes, dict) + assert "deferred_tasks" not in retryable_outcomes + + @pytest.mark.asyncio(loop_scope="session") async def test_degraded_provider_warn_only_counts_llm_tasks() -> None: """The WARN window must ignore non-LLM task outcomes (samplers, expressions, etc).