Skip to content
Closed
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 @@ -62,7 +62,10 @@
from data_designer.engine.models.errors import (
RETRYABLE_MODEL_ERRORS,
GenerationValidationFailureError,
ModelAPIConnectionError,
ModelInternalServerError,
ModelRateLimitError,
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 @@ -269,6 +272,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 @@ -528,6 +535,20 @@ 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())),
"deferred_tasks": len(self._deferred),
}

def _scheduler_job_diagnostics(self) -> dict[str, object]:
Expand Down Expand Up @@ -1508,7 +1529,36 @@ 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, 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
Expand All @@ -1518,9 +1568,28 @@ def _record_retryable_outcome(self, *, retryable: bool) -> None:
rate. Only retryable errors (rate-limit, timeout, 5xx, connection) count
toward the rate; non-retryable failures register as 0.
"""
kind = self._retryable_error_kind(exc) if retryable else "success"

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Non-retryable failures misclassified as "success"

kind is always "success" when retryable=False, but the exception handler at line 1751 also calls _record_retryable_outcome(retryable=False, exc=exc) for non-retryable LLM failures (auth errors, schema errors, code bugs). Those failures increment _retryable_outcome_counts["success"] and are appended to _recent_retryable_kinds as "success", so the diagnostics inflate the success count and understate actual failure rates. An operator seeing cumulative_counts: {success: 80, rate_limit: 20} cannot tell whether the 80 represent genuine successes or a mix of successes and non-retryable failures.

Prompt To Fix With AI
This is a comment left during a code review.
Path: packages/data-designer-engine/src/data_designer/engine/dataset_builders/async_scheduler.py
Line: 1571

Comment:
**Non-retryable failures misclassified as "success"**

`kind` is always `"success"` when `retryable=False`, but the exception handler at line 1751 also calls `_record_retryable_outcome(retryable=False, exc=exc)` for non-retryable LLM failures (auth errors, schema errors, code bugs). Those failures increment `_retryable_outcome_counts["success"]` and are appended to `_recent_retryable_kinds` as `"success"`, so the diagnostics inflate the success count and understate actual failure rates. An operator seeing `cumulative_counts: {success: 80, rate_limit: 20}` cannot tell whether the 80 represent genuine successes or a mix of successes and non-retryable failures.

How can I resolve this? If you propose a fix, please make it concise.

self._retryable_outcome_counts[kind] += 1
status_code, retry_after = self._provider_error_details(exc)
if retryable:
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 @@ -1531,10 +1600,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 @@ -1669,7 +1748,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 @@ -47,6 +47,7 @@
from data_designer.engine.dataset_builders.utils.row_group_buffer import RowGroupBufferManager
from data_designer.engine.models.errors import (
RETRYABLE_MODEL_ERRORS,
ModelAPIConnectionError,
ModelInternalServerError,
ModelRateLimitError,
ModelTimeoutError,
Expand Down Expand Up @@ -1620,6 +1621,27 @@ 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"),
(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 "sensitive provider text" not in str(metrics)


@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
5 changes: 5 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,11 @@ members = [
package = false
required-version = ">=0.7.10"

[tool.uv.sources]
data-designer = { workspace = true }
data-designer-config = { workspace = true }
data-designer-engine = { workspace = true }

[dependency-groups]
dev = [
"jsonpath-ng>=1.5.3,<2",
Expand Down
Loading