Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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).
Expand Down Expand Up @@ -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())),
Comment thread
danecor marked this conversation as resolved.
}

def _scheduler_job_diagnostics(self) -> dict[str, object]:
Expand Down Expand Up @@ -1778,19 +1798,75 @@ 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
or failure) feeds this window so the rate reflects the provider's overall
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
Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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).
Expand Down
Loading