fix: gate retries by error category with backoff

Co-authored-by: Copilot App <[email protected]>
This commit is contained in:
Jim Lancaster
2026-08-23 18:30:54 -05:00
co-authored by Copilot App
parent f193b2800b
commit 86cdb4035c
7 changed files with 112 additions and 2 deletions
+1
View File
@@ -58,6 +58,7 @@ DATABASE_BACKUP_DIR=./data/backups
WORKER_MAX_RETRIES=0 WORKER_MAX_RETRIES=0
WORKER_PROVIDER_TIMEOUT_SECONDS=30.0 WORKER_PROVIDER_TIMEOUT_SECONDS=30.0
WORKER_STALE_JOB_SECONDS=30.0 WORKER_STALE_JOB_SECONDS=30.0
WORKER_RETRY_BACKOFF_SECONDS=1.0
WORKER_MIN_TRANSCRIPTION_CHARS=0 WORKER_MIN_TRANSCRIPTION_CHARS=0
WORKER_MIN_TRANSCRIPTION_LINES=0 WORKER_MIN_TRANSCRIPTION_LINES=0
WORKER_FAIL_ON_FINISH_REASON_LENGTH=false WORKER_FAIL_ON_FINISH_REASON_LENGTH=false
+1
View File
@@ -115,6 +115,7 @@ class Settings(BaseSettings):
# well past twenty seconds, so an upper cap here would silently fail real work. # well past twenty seconds, so an upper cap here would silently fail real work.
worker_provider_timeout_seconds: float = Field(default=30.0, gt=0.0) worker_provider_timeout_seconds: float = Field(default=30.0, gt=0.0)
worker_stale_job_seconds: float = Field(default=30.0, gt=0.0) worker_stale_job_seconds: float = Field(default=30.0, gt=0.0)
worker_retry_backoff_seconds: float = Field(default=1.0, ge=0.0)
worker_min_transcription_chars: int = Field(default=0, ge=0) worker_min_transcription_chars: int = Field(default=0, ge=0)
worker_min_transcription_lines: int = Field(default=0, ge=0) worker_min_transcription_lines: int = Field(default=0, ge=0)
worker_fail_on_finish_reason_length: bool = False worker_fail_on_finish_reason_length: bool = False
+20
View File
@@ -43,6 +43,26 @@ class LatestExecutionAttempt:
class EvidenceService(ServiceBase): class EvidenceService(ServiceBase):
"""Read, project, and export execution attempt evidence.""" """Read, project, and export execution attempt evidence."""
async def read_latest_job_error_category(
self,
*,
job_id: UUID,
session: AsyncSession | None = None,
) -> str | None:
"""Read the latest persisted execution-attempt error category for a job."""
async with self._session_scope(session) as _session:
query = (
select(ExecutionAttempt.error_category)
.where(ExecutionAttempt.job_id == job_id)
.where(col(ExecutionAttempt.error_category).is_not(None))
.order_by(
col(ExecutionAttempt.created_at).desc(),
col(ExecutionAttempt.id).desc(),
)
.limit(1)
)
return (await _session.exec(query)).first()
async def read_latest_execution_attempt( async def read_latest_execution_attempt(
self, self,
*, *,
+24 -1
View File
@@ -40,6 +40,12 @@ from .sources import transcribe_document_image
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_RETRIABLE_FAILED_JOB_ERROR_CATEGORIES = {
ErrorCategory.EXTERNAL_PROVIDER.value,
ErrorCategory.EXTERNAL_TIMEOUT.value,
ErrorCategory.INFRA_TRANSIENT.value,
}
async def create_document_with_people( async def create_document_with_people(
*, *,
@@ -182,13 +188,30 @@ async def advance_job(
# Recover mid-flight jobs by continuing the queued processing path. # Recover mid-flight jobs by continuing the queued processing path.
return await process_queued_job(job=job, services=services, settings=settings, session=session) return await process_queued_job(job=job, services=services, settings=settings, session=session)
case JobStatus.FAILED: case JobStatus.FAILED:
if job.retry_count < settings.worker_max_retries: latest_error_category = await services.evidence.read_latest_job_error_category(
job_id=job.id,
session=session,
)
can_retry = (
job.retry_count < settings.worker_max_retries
and latest_error_category in _RETRIABLE_FAILED_JOB_ERROR_CATEGORIES
)
if can_retry:
if settings.worker_retry_backoff_seconds > 0:
await asyncio.sleep(settings.worker_retry_backoff_seconds)
return await services.jobs.update_job_state( return await services.jobs.update_job_state(
job_id=job.id, job_id=job.id,
status=JobStatus.QUEUED, status=JobStatus.QUEUED,
retry_count_increment=1, retry_count_increment=1,
session=session, session=session,
) )
if job.retry_count < settings.worker_max_retries:
logger.warning(
"Job %s failed with non-retriable category %s; skipping retry.",
job.id,
latest_error_category or "unknown",
)
else: else:
logger.error("Job %s has failed and reached max retries.", job.id) logger.error("Job %s has failed and reached max retries.", job.id)
return return
@@ -2,6 +2,8 @@
import asyncio import asyncio
import time import time
from datetime import UTC
from datetime import datetime
from pathlib import Path from pathlib import Path
from uuid import uuid4 from uuid import uuid4
@@ -20,6 +22,7 @@ from transcription.db.models import Source
from transcription.providers.base import TranscriptionResult from transcription.providers.base import TranscriptionResult
from transcription.services import ServiceBundle from transcription.services import ServiceBundle
from transcription.services import workflows as workflows_module from transcription.services import workflows as workflows_module
from transcription.services.workflows import advance_job
from transcription.services.workflows import process_queued_job from transcription.services.workflows import process_queued_job
@@ -366,3 +369,63 @@ class TestWorkflowReliability:
result = await task result = await task
assert result is not None assert result is not None
assert result.status == JobStatus.TRANSCRIBED assert result.status == JobStatus.TRANSCRIBED
@pytest.mark.asyncio
async def test_failed_job_with_validation_category_is_not_requeued(self, default_session_factory):
services = ServiceBundle.from_session_factory(default_session_factory)
async with services.jobs._session_scope() as session:
document = Document(id=uuid4(), name="validation-failure-doc")
session.add(document)
await session.flush()
source = Source(
document_id=document.id,
page_number=1,
upload_name="validation.jpg",
filename="validation.jpg",
file_path=str(Path("tests/fixtures/images/real/Book Two - page 02.jpg")),
file_hash="9" * 64,
file_size_bytes=1,
)
session.add(source)
await session.flush()
job = Job(document_id=document.id, status=JobStatus.FAILED, retry_count=0)
session.add(job)
await session.flush()
job_source = JobSource(job_id=job.id, source_id=source.id, status=JobSourceStatus.FAILED)
session.add(job_source)
await session.flush()
now = datetime.now(UTC)
session.add(
ExecutionAttempt(
job_source_id=job_source.id,
job_id=job.id,
source_id=source.id,
attempt_number=1,
status=JobSourceStatus.FAILED,
provider="fixture",
started_at=now,
finished_at=now,
duration_ms=0,
error_category="validation_error",
error_detail="invalid payload",
)
)
await session.commit()
failed_job = await services.jobs.read_job(job_id=job.id, session=session)
result = await advance_job(
failed_job,
services=services,
settings=Settings(openrouter_api_key="test-key", worker_max_retries=1),
)
assert result is None
async with services.jobs._session_scope() as session:
persisted = await session.get(Job, failed_job.id)
assert persisted is not None
assert persisted.status == JobStatus.FAILED
assert persisted.retry_count == 0
+1
View File
@@ -168,6 +168,7 @@ class TestWorkerReliabilitySettings:
settings = _make_settings() settings = _make_settings()
assert settings.worker_max_retries == 0 assert settings.worker_max_retries == 0
assert settings.worker_stale_job_seconds == 30.0 assert settings.worker_stale_job_seconds == 30.0
assert settings.worker_retry_backoff_seconds == 1.0
def test_provider_timeout_is_not_capped_at_twenty_seconds(): def test_provider_timeout_is_not_capped_at_twenty_seconds():
+1
View File
@@ -168,6 +168,7 @@ def test_env_example_default_values_match_settings_defaults():
"WORKER_MAX_RETRIES": str(defaults.worker_max_retries), "WORKER_MAX_RETRIES": str(defaults.worker_max_retries),
"WORKER_PROVIDER_TIMEOUT_SECONDS": str(defaults.worker_provider_timeout_seconds), "WORKER_PROVIDER_TIMEOUT_SECONDS": str(defaults.worker_provider_timeout_seconds),
"WORKER_STALE_JOB_SECONDS": str(defaults.worker_stale_job_seconds), "WORKER_STALE_JOB_SECONDS": str(defaults.worker_stale_job_seconds),
"WORKER_RETRY_BACKOFF_SECONDS": str(defaults.worker_retry_backoff_seconds),
"WORKER_MIN_TRANSCRIPTION_CHARS": str(defaults.worker_min_transcription_chars), "WORKER_MIN_TRANSCRIPTION_CHARS": str(defaults.worker_min_transcription_chars),
"WORKER_MIN_TRANSCRIPTION_LINES": str(defaults.worker_min_transcription_lines), "WORKER_MIN_TRANSCRIPTION_LINES": str(defaults.worker_min_transcription_lines),
"WORKER_FAIL_ON_FINISH_REASON_LENGTH": str(defaults.worker_fail_on_finish_reason_length).lower(), "WORKER_FAIL_ON_FINISH_REASON_LENGTH": str(defaults.worker_fail_on_finish_reason_length).lower(),