Skip to content
163 changes: 163 additions & 0 deletions tr_sys/tests/unit/test_merge_retry.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,163 @@
"""
Unit tests for the retry semantics of merge_and_post_process.

The DB, redis token gate and merge internals are all mocked out; these tests
exercise only the retry/backoff decision logic. Celery's eager mode re-executes
retries synchronously (Task.apply resubmits on Retry), so contention scenarios
run their full retry budget inline and we can assert exact attempt counts.
"""
from contextlib import nullcontext
from unittest.mock import patch, MagicMock

import pytest
from celery.exceptions import Retry, MaxRetriesExceededError

from tr_ars import utils
from tr_sys.celery_gates.expensive_gate import TASK_MAX_RETRIES


TASK_ARGS = ("parent-pk-1", {"knowledge_graph": {}}, "ara-test")


@pytest.fixture
def mock_parent():
parent = MagicMock()
parent.merge_semaphore = False
parent.merged_versions_list = []
return parent


@pytest.fixture
def merge_env(mock_parent):
"""Patch the DB, redis gate and merge internals; yield the mocks for per-test tuning."""
merged = MagicMock()
merged.id = "merged-pk-1"
merged.pk = "merged-pk-1"
env = {"parent": mock_parent, "merged": merged}
with patch.object(utils, "transaction") as txn, \
patch.object(utils, "get_object_or_404", return_value=mock_parent), \
patch.object(utils, "Message") as message_cls, \
patch.object(utils, "try_lock_merge", return_value=True) as try_lock, \
patch.object(utils, "unlock_merge"), \
patch.object(utils, "merge_received", return_value=(merged, mock_parent, {})) as merge_received, \
patch.object(utils, "post_process", return_value=(merged, 200, "D")) as post_process, \
patch.object(utils, "record_error"), \
patch.object(utils, "exp_backoff_with_jitter", return_value=0), \
patch.object(utils, "constant_backoff_with_jitter", return_value=0), \
patch("tr_sys.celery_gates.context.try_acquire", return_value=True) as try_acquire, \
patch("tr_sys.celery_gates.context.release") as release, \
patch("tr_sys.celery_gates.context.LeaseRenewer"), \
patch("tr_sys.celery_gates.context.constant_backoff_with_jitter", return_value=0):
txn.atomic.return_value = nullcontext()
message_cls.objects.filter.return_value.first.return_value = mock_parent
message_cls.DoesNotExist = type("DoesNotExist", (Exception,), {}) # never raised here
env.update(try_lock=try_lock, try_acquire=try_acquire, release=release,
merge_received=merge_received, post_process=post_process)
yield env


def _failed_events(mock_parent):
return [c.args[0]["event_type"] for c in mock_parent.notify_subscribers.call_args_list
if c.args[0].get("event_type") == "merged_version_failed"]


def test_happy_path_merges(merge_env):
result = utils.merge_and_post_process.apply(args=TASK_ARGS)
assert result.state == "SUCCESS"
merge_env["merge_received"].assert_called_once()
merge_env["post_process"].assert_called_once()


def test_token_contention_retries_then_succeeds(merge_env):
"""A task that can't get a token keeps retrying and completes once one frees up."""
merge_env["try_acquire"].side_effect = [False, False, False, True]
result = utils.merge_and_post_process.apply(args=TASK_ARGS)
assert result.state == "SUCCESS"
assert merge_env["try_acquire"].call_count == 4
merge_env["merge_received"].assert_called_once()


def test_token_exhaustion_fails_loudly_after_full_budget(merge_env, mock_parent):
"""Budget exhaustion must fail (not silently return), after exactly the budgeted attempts."""
merge_env["try_acquire"].return_value = False
result = utils.merge_and_post_process.apply(args=TASK_ARGS)
assert result.state == "FAILURE"
assert isinstance(result.result, MaxRetriesExceededError)
# initial attempt + TASK_MAX_RETRIES retries
assert merge_env["try_acquire"].call_count == TASK_MAX_RETRIES + 1
merge_env["merge_received"].assert_not_called()
# the token was never held, so it must never be released
merge_env["release"].assert_not_called()
# exhaustion is surfaced to API consumers, not just logs
assert _failed_events(mock_parent)


def test_lock_contention_releases_token_and_retries(merge_env):
merge_env["try_lock"].side_effect = [False, False, True]
result = utils.merge_and_post_process.apply(args=TASK_ARGS)
assert result.state == "SUCCESS"
assert merge_env["try_lock"].call_count == 3
# the expensive token must be released on each lock-contention requeue (and on success)
assert merge_env["release"].call_count == 3


def test_lock_exhaustion_fails_loudly_after_full_budget(merge_env, mock_parent):
merge_env["try_lock"].return_value = False
result = utils.merge_and_post_process.apply(args=TASK_ARGS)
assert result.state == "FAILURE"
assert isinstance(result.result, MaxRetriesExceededError)
assert merge_env["try_lock"].call_count == TASK_MAX_RETRIES + 1
merge_env["merge_received"].assert_not_called()
assert _failed_events(mock_parent)


def test_error_retries_stop_at_cap(merge_env):
"""A persistently failing merge is attempted exactly MERGE_ERROR_MAX_RETRIES + 1 times."""
merge_env["merge_received"].side_effect = RuntimeError("boom")
result = utils.merge_and_post_process.apply(args=TASK_ARGS)
assert result.state == "FAILURE"
assert isinstance(result.result, RuntimeError)
assert merge_env["merge_received"].call_count == utils.MERGE_ERROR_MAX_RETRIES + 1


def test_error_retry_recovers_after_transient_failure(merge_env, mock_parent):
merged = merge_env["merged"]
merge_env["merge_received"].side_effect = [RuntimeError("boom"), RuntimeError("boom"),
(merged, mock_parent, {})]
result = utils.merge_and_post_process.apply(args=TASK_ARGS)
assert result.state == "SUCCESS"
assert merge_env["merge_received"].call_count == 3


def test_error_retry_preserves_caller_kwargs(merge_env):
"""The error retry merges into request.kwargs instead of replacing them wholesale."""
merge_env["merge_received"].side_effect = RuntimeError("boom")
retry_spy = MagicMock(side_effect=Retry("requested retry"))
kwargs = {"parent_pk": "parent-pk-1", "message_to_merge": {}, "agent_name": "ara-test"}
with patch.object(utils.merge_and_post_process, "retry", retry_spy):
result = utils.merge_and_post_process.apply(kwargs=kwargs)
assert result.state == "RETRY"
retry_kwargs = retry_spy.call_args.kwargs["kwargs"]
assert retry_kwargs["error_retries"] == 1
for key, value in kwargs.items():
assert retry_kwargs[key] == value


def test_error_retry_increments_existing_counter(merge_env):
merge_env["merge_received"].side_effect = RuntimeError("boom")
retry_spy = MagicMock(side_effect=Retry("requested retry"))
with patch.object(utils.merge_and_post_process, "retry", retry_spy):
utils.merge_and_post_process.apply(
kwargs={"parent_pk": "p", "message_to_merge": {}, "agent_name": "a",
"error_retries": 3})
assert retry_spy.call_args.kwargs["kwargs"]["error_retries"] == 4


def test_direct_call_error_retry_survives_missing_request_kwargs(merge_env):
"""Called directly (not via a worker), request.kwargs is None; the retry must not TypeError."""
merge_env["merge_received"].side_effect = RuntimeError("boom")
retry_spy = MagicMock(side_effect=Retry("requested retry"))
with patch.object(utils.merge_and_post_process, "retry", retry_spy):
with pytest.raises(Retry):
utils.merge_and_post_process(*TASK_ARGS)
assert retry_spy.call_args.kwargs["kwargs"]["error_retries"] == 1
96 changes: 70 additions & 26 deletions tr_sys/tr_ars/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,14 +37,16 @@
from opentelemetry import trace
from opentelemetry.trace import Status, StatusCode

from tr_sys.celery_gates.expensive_gate import exp_backoff_with_jitter
from tr_sys.celery_gates.expensive_gate import (exp_backoff_with_jitter, constant_backoff_with_jitter,
TASK_MAX_RETRIES)

tracer = trace.get_tracer(__name__)
import asyncio
import zstandard as zstd
from tr_sys.celery_gates.context import (expensive_section)
from celery.exceptions import Retry, MaxRetriesExceededError
from tr_sys.otel_config import count_error, record_error
from celery.exceptions import Retry


ARS_ACTOR = {
'channel': [],
Expand All @@ -59,6 +61,7 @@
NORMALIZER_URL=os.getenv("TR_NORMALIZER") if os.getenv("TR_NORMALIZER") is not None else "https://nodenorm-es.ci.transltr.io/get_normalized_nodes"
ANNOTATOR_URL=os.getenv("TR_ANNOTATOR") if os.getenv("TR_ANNOTATOR") is not None else "https://biothings.ncats.io/curie"
APPRAISER_URL=os.getenv("TR_APPRAISE") if os.getenv("TR_APPRAISE") is not None else "https://answerappraiser.ci.transltr.io/get_appraisal"
MERGE_ERROR_MAX_RETRIES = int(os.getenv("ARS_MERGE_ERROR_MAX_RETRIES", "8"))

class QueryGraph():
def __init__(self,qg):
Expand Down Expand Up @@ -704,15 +707,31 @@ def unlock_merge(message: Message) -> None:
logging.exception("Failed to release merge_semaphore for message %s", getattr(message, "pk", "<unknown>"))
raise

@shared_task(name="merge-and-post-process", bind=True, acks_late=True, max_retries=20)
def merge_and_post_process(self, parent_pk,message_to_merge, agent_name):

@shared_task(name="merge-and-post-process", bind=True, acks_late=True, max_retries=None)
def merge_and_post_process(self, parent_pk, message_to_merge, agent_name, error_retries=0):
"""
Safe merge & post-process task:
- Acquire expensive token (expensive_section) first
- Acquire DB boolean lock (inside short select_for_update atomic)
- Do merge_received(), post_process() and notifications
- Always release DB boolean lock in finally
- Use self.request.retries for retry/backoff

Retry functionality:
- The built-in Celery retry functionality, tracked by the counter self.request.retries and enforced by Celery,
is used for all tasks, contention is the main cause for a retry (no available expensive tokens or the
merge db lock is occupied). This is done by setting max_retries to TASK_MAX_RETRIES in expensive_section()
and in this function after a failed merge lock acquisition. When this kind of retry occurs we use a quick
delay to avoid pointless downtime waiting on ourselves.
- Retries for errors in merging are handled separately by error_retries, an arg based counter
for actual merge_and_post_process errors, capped by MERGE_ERROR_MAX_RETRIES, enforced by the code here
and not celery's internal retry check. These use exponential backoff (1, 2, 4, 8, 16s, then capped at 30s)
to give downstream services time to recover, so the default cap of 8 buys about two minutes of patience —
enough to outlast an annotator/appraiser restart or a rolling deploy.
- Note that max_retries=None in the merge_and_post_process decorator signature means celery imposes no
limit of its own: contention retry sites must pass max_retries explicitly, while the error path
intentionally omits it and relies on the error_retries cap instead. Future changes should preserve that
split — a retry with neither limit would loop forever.
"""
merged=None
stats={}
Expand All @@ -727,7 +746,7 @@ def merge_and_post_process(self, parent_pk,message_to_merge, agent_name):
logging.info(f"🚀Starting merge for %s with parent PK: %s"% (agent_name,parent_pk))
try:
#Acquire an expensive token so we don't hold DB locked while waiting
with expensive_section(self):
with expensive_section(self, max_retries=TASK_MAX_RETRIES):
logging.info("[%s] 🟢 acquired expensive token", self.request.id)

# short critical section: lock row + decide if we can merge
Expand All @@ -748,23 +767,15 @@ def merge_and_post_process(self, parent_pk,message_to_merge, agent_name):
lock_span.set_attribute("merge.lock.acquired", lock_acquired)
task_span.set_attribute("merge.lock.acquired", lock_acquired)
if not lock_acquired:
#someone else already locked it & possibly going through merge, retry ...
retries = self.request.retries
if retries < 10:
delay = 5
outcome = "lock_contended"
task_span.set_attribute("merge.lock.retry_delay_seconds", delay)
if retries >= 7:
logging.warning("⚠️ High retry count (%s) for merge lock on parent=%s agent=%s.Retrying in %ss",retries, parent_pk, agent_name, delay)
else:
logging.info(" 🔄 Merged_version locked for %s. Attempt %s. Retrying in %ss",agent_name, retries, delay)
# raise retry — this will release the expensive token (expensive_section finally), stop the task, requeues it and free the worker
raise self.retry(countdown=delay)
else:
outcome = "lock_budget_exhausted"
task_span.set_status(Status(StatusCode.ERROR, "gave up waiting for the per-parent merge lock"))
logging.info("❌ Merging failed for %s %s after retries", agent_name, parent_pk)
return
# Another task holds the merge lock. Retry after a quick delay.
delay = constant_backoff_with_jitter()
outcome = "lock_contended"
task_span.set_attribute("merge.lock.retry_delay_seconds", delay)
logging.info("🔄 merge lock held for %s (parent=%s); retrying in %.1fs (retries=%s)",
agent_name, parent_pk, delay, self.request.retries)
# raise retry — releases the expensive token (expensive_section finally),
# requeues the task and frees the worker.
raise self.retry(countdown=delay, max_retries=TASK_MAX_RETRIES)
else:
logging.info(" the merge semaphore for agent %s is %s"% (agent_name, parent.merge_semaphore))
locked = True
Expand Down Expand Up @@ -843,7 +854,35 @@ def merge_and_post_process(self, parent_pk,message_to_merge, agent_name):
logging.info("Task retry requested — requeueing (parent=%s, retries=%s)",parent_pk, self.request.retries)
raise

except MaxRetriesExceededError:
# General retry budget exhausted (due to token gate or merge db lock).
# This is not a real error, so keep it out of the catch-all Exception handling below.
outcome = "retry_budget_exhausted"
task_span.set_status(Status(StatusCode.ERROR, "exhausted the retry budget waiting on contention"))
logging.error("❌ Merge gave up after exhausting the retry budget for agent %s pk %s (retries=%s).",
agent_name, parent_pk, self.request.retries)
# surface the failure to API consumers, not just logs/OTEL
try:
if parent is None:
parent = Message.objects.filter(pk=parent_pk).first()
if parent is not None:
parent.notify_subscribers({
"event_type": "merged_version_failed",
"agent_name": agent_name,
"reason": "retry_budget_exhausted",
"merged_versions_list": parent.merged_versions_list if parent.merged_versions_list is not None else []
})
except Exception as notify_error:
logging.exception("Failed to notify subscribers of merge retry exhaustion for parent %s", parent_pk)
record_error(notify_error)
raise

except Exception as e:
# TODO - it would be best to categorize and handle different kinds of errors in different ways here.
# If there is a deterministic error (ie something wrong with the data) retrying is a waste of time and
# hogs resources for no reason. Transient errors such as issues with postgres or service calls to annotator
# or appraiser are legitimate reasons to retry and should go through the retry process. It would also be
# good to log the kinds of errors occurring and/or make them visible to OTEL.
outcome = "error"
logging.info("Problem with merger for agent %s pk: %s " % (agent_name, (parent_pk)))
logging.info(e, exc_info=True)
Expand All @@ -853,9 +892,14 @@ def merge_and_post_process(self, parent_pk,message_to_merge, agent_name):
merged.status='E'
merged.code = 422
merged.save()
# retry on post_process failure
delay = exp_backoff_with_jitter(self.request.retries)
raise self.retry(exc=e, countdown=delay)
if error_retries >= MERGE_ERROR_MAX_RETRIES:
logging.error("Merge for agent %s pk %s failed after %s error retries; giving up.",
agent_name, parent_pk, error_retries)
raise
delay = exp_backoff_with_jitter(error_retries)
raise self.retry(
kwargs={**(self.request.kwargs or {}), "error_retries": error_retries + 1},
exc=e, countdown=delay)

finally:
task_span.set_attribute("merge.outcome", outcome)
Expand Down
Loading