diff --git a/src/forge/integrations/jira/client.py b/src/forge/integrations/jira/client.py index b519e58c..dd335f94 100644 --- a/src/forge/integrations/jira/client.py +++ b/src/forge/integrations/jira/client.py @@ -697,10 +697,28 @@ async def get_comments(self, issue_key: str) -> list[JiraComment]: List of JiraComment objects. """ client = await self._get_client() - response = await client.get(f"/issue/{issue_key}/comment") - response.raise_for_status() - data = response.json() - return [JiraComment.from_api_response(c) for c in data.get("comments", [])] + comments: list[JiraComment] = [] + start_at = 0 + max_results = 100 + + while True: + response = await client.get( + f"/issue/{issue_key}/comment", + params={"startAt": start_at, "maxResults": max_results}, + ) + response.raise_for_status() + data = response.json() + page = data.get("comments", []) + comments.extend(JiraComment.from_api_response(comment) for comment in page) + + page_start = int(data.get("startAt", start_at)) + total = int(data.get("total", page_start + len(page))) + next_start = page_start + len(page) + if not page or next_start >= total: + break + start_at = next_start + + return comments async def get_labels(self, issue_key: str) -> list[str]: """Get labels for a Jira issue. diff --git a/src/forge/orchestrator/worker.py b/src/forge/orchestrator/worker.py index 2a425870..1c326598 100644 --- a/src/forge/orchestrator/worker.py +++ b/src/forge/orchestrator/worker.py @@ -119,7 +119,10 @@ def __init__( """ self.settings = get_settings() self.consumer_name = consumer_name or f"worker-{uuid.uuid4().hex[:8]}" - self.consumer = QueueConsumer(self.consumer_name) + self.consumer = QueueConsumer( + self.consumer_name, + terminal_failure_handler=self._handle_terminal_failure, + ) self.router = router or create_default_router() self._shutdown_event = asyncio.Event() self._checkpointer = None @@ -140,6 +143,37 @@ async def _get_forge_github_login(self) -> str: self._forge_github_login = login return login + async def _handle_terminal_failure(self, message: QueueMessage, error: str) -> None: + """Post one Jira comment after queue retries are exhausted.""" + jira = JiraClient() + event_marker = f"Event/correlation ID: {message.event_id}" + try: + comments = await jira.get_comments(message.ticket_key) + if any(event_marker in comment.body for comment in comments): + logger.info( + f"Terminal failure notification already exists for event {message.event_id}" + ) + return + + safe_error = redact_secrets(error) + if len(safe_error) > 500: + safe_error = f"{safe_error[:500]}..." + details = ( + f"{safe_error}\n\n" + f"Ticket: {message.ticket_key}\n" + f"{event_marker}\n" + "Recovery: inspect the dead-letter entry, resolve the root cause, " + "then requeue the event." + ) + await jira.add_error_comment( + issue_key=message.ticket_key, + error_message=details, + node_name="queue execution (retries exhausted)", + ) + logger.info(f"Posted terminal queue failure notification to {message.ticket_key}") + finally: + await jira.close() + async def _handle_jira_event(self, message: QueueMessage) -> None: """Handle a Jira webhook event. diff --git a/src/forge/queue/consumer.py b/src/forge/queue/consumer.py index 2cff8219..1a95fd9f 100644 --- a/src/forge/queue/consumer.py +++ b/src/forge/queue/consumer.py @@ -3,7 +3,7 @@ import asyncio import logging from collections import defaultdict -from collections.abc import Callable, Coroutine +from collections.abc import AsyncIterator, Callable, Coroutine from contextlib import asynccontextmanager, suppress from typing import Any @@ -16,7 +16,7 @@ from forge.orchestrator.checkpointer import get_redis_client from forge.queue.models import QueueMessage from forge.queue.producer import GITHUB_STREAM, JIRA_STREAM -from forge.queue.retry import RetryQueue +from forge.queue.retry import RETRY_CLAIM_RENEW_SECONDS, RetryEntry, RetryQueue logger = logging.getLogger(__name__) @@ -35,12 +35,17 @@ # Handler type for message processing MessageHandler = Callable[[QueueMessage], Coroutine[Any, Any, None]] +TerminalFailureHandler = Callable[[QueueMessage, str], Coroutine[Any, Any, None]] class TicketLockLostError(RuntimeError): """Raised when a worker loses its distributed per-ticket lease.""" +class RetryLeaseLostError(RuntimeError): + """Raised when a worker loses ownership of a claimed retry entry.""" + + class QueueConsumer: """Consumes webhook events from Redis Streams with FIFO ordering per ticket. @@ -54,6 +59,7 @@ def __init__( redis_client: redis.Redis | None = None, jira_client: JiraClient | None = None, max_concurrent_tasks: int | None = None, + terminal_failure_handler: TerminalFailureHandler | None = None, ): """Initialize the queue consumer. @@ -63,6 +69,8 @@ def __init__( jira_client: Optional Jira client for freshness checks. max_concurrent_tasks: Maximum concurrent in-flight tasks. Defaults to ``settings.queue_max_concurrent_tasks`` when not provided. + terminal_failure_handler: Optional callback invoked after a message + is moved to the dead-letter queue. """ self.consumer_name = consumer_name self._redis = redis_client @@ -78,9 +86,10 @@ def __init__( self._semaphore = asyncio.Semaphore(concurrency) self._active_tasks: set[asyncio.Task[None]] = set() self._retry_queue = RetryQueue() + self._terminal_failure_handler = terminal_failure_handler @asynccontextmanager - async def _distributed_ticket_lock(self, ticket_key: str): + async def _distributed_ticket_lock(self, ticket_key: str) -> AsyncIterator[None]: """Hold a renewable, deployment-wide lease for one ticket.""" redis_client = await self._get_redis() lock = redis_client.lock( @@ -199,6 +208,16 @@ async def _ack(self, stream: str, message_id: str) -> None: redis_client = await self._get_redis() await redis_client.xack(stream, CONSUMER_GROUP, message_id) + async def _ack_terminal_message(self, message: QueueMessage, stream: str) -> None: + """Best-effort acknowledge a message after it has reached the DLQ.""" + try: + await self._ack(stream, message.message_id) + except Exception as ack_error: + logger.warning( + f"xack failed for terminal event {message.event_id}; " + f"PEL entry may linger: {ack_error}" + ) + async def _process_message( self, message: QueueMessage, @@ -251,7 +270,7 @@ async def _process_message( try: moved_to_dlq = not await self._retry_queue.enqueue_for_retry(message, str(exc)) if moved_to_dlq: - await self._ack(stream, message.message_id) + await self._ack_terminal_message(message, stream) except Exception as retry_err: logger.error( f"Failed to enqueue {message.event_id} for retry after ticket lock loss: " @@ -290,7 +309,7 @@ async def _handle_locked_message( try: moved_to_dlq = not await self._retry_queue.enqueue_for_retry(message, str(e)) if moved_to_dlq: - await self._ack(stream, message.message_id) + await self._ack_terminal_message(message, stream) except Exception as retry_err: logger.error( f"Failed to enqueue {message.event_id} for retry: {retry_err}. " @@ -337,21 +356,70 @@ async def _consume_stream(self, stream: str, _source: EventSource) -> None: if self._active_tasks: await asyncio.gather(*self._active_tasks, return_exceptions=True) + @asynccontextmanager + async def _renew_retry_claim(self, entry: RetryEntry) -> AsyncIterator[None]: + """Renew a retry lease and abort processing if ownership is lost.""" + if entry.lease_until is None: + yield + return + + owner_task = asyncio.current_task() + lease_lost = asyncio.Event() + + async def renew() -> None: + while True: + await asyncio.sleep(RETRY_CLAIM_RENEW_SECONDS) + try: + if not await self._retry_queue.renew_retry_claim(entry): + lease_lost.set() + if owner_task is not None: + owner_task.cancel() + return + except asyncio.CancelledError: + raise + except Exception: + logger.exception("Failed to renew retry lease for %s", entry.message.event_id) + lease_lost.set() + if owner_task is not None: + owner_task.cancel() + return + + renewal_task = asyncio.create_task(renew(), name=f"renew-retry-{entry.message.event_id}") + try: + yield + except asyncio.CancelledError: + if lease_lost.is_set(): + raise RetryLeaseLostError( + f"Lost retry lease for {entry.message.event_id}" + ) from None + raise + finally: + renewal_task.cancel() + with suppress(asyncio.CancelledError): + await renewal_task + async def _process_due_retries_once(self) -> None: """Dispatch one batch of due retries. Kept separate from the polling loop so recovery and terminal DLQ behavior can be exercised deterministically in integration tests. """ - entries = await self._retry_queue.get_due_messages() + entries = await self._retry_queue.claim_due_messages() for entry in entries: retry_stream = ( JIRA_STREAM if entry.message.source == EventSource.JIRA else GITHUB_STREAM ) try: - await self._process_message( - entry.message, retry_stream, raise_on_error=True, skip_ack=True + async with self._renew_retry_claim(entry): + await self._process_message( + entry.message, retry_stream, raise_on_error=True, skip_ack=True + ) + except RetryLeaseLostError: + logger.error( + "Stopped retry processing after losing ownership of %s", + entry.message.event_id, ) + continue except Exception as e: logger.warning( f"Retry attempt {entry.attempt} failed for " @@ -364,7 +432,7 @@ async def _process_due_retries_once(self) -> None: if not queued: # Terminal failure: the DLQ now owns the message, so clear # its original stream PEL entry. - await self._ack(retry_stream, entry.message.message_id) + await self._ack_terminal_message(entry.message, retry_stream) continue # Success: clear retry state and acknowledge the original entry. @@ -381,11 +449,30 @@ async def _process_due_retries_once(self) -> None: f"retry (PEL entry may linger): {xack_err}" ) + async def _process_terminal_notifications_once(self) -> None: + """Deliver one leased batch from the terminal-notification outbox.""" + if self._terminal_failure_handler is None: + return + + entries = await self._retry_queue.claim_due_terminal_notifications() + for entry in entries: + try: + await self._terminal_failure_handler(entry.message, entry.error) + except Exception as exc: + logger.error( + "Terminal failure callback failed for %s: %s", + entry.message.event_id, + exc, + ) + continue + await self._retry_queue.remove_terminal_notification(entry) + async def _process_retry_queue(self) -> None: """Poll the retry queue and re-dispatch due messages.""" while self._running: try: await self._process_due_retries_once() + await self._process_terminal_notifications_once() except asyncio.CancelledError: break except Exception as e: diff --git a/src/forge/queue/retry.py b/src/forge/queue/retry.py index 6733f443..eb537809 100644 --- a/src/forge/queue/retry.py +++ b/src/forge/queue/retry.py @@ -16,12 +16,48 @@ RETRY_QUEUE_KEY = "forge:retry:queue" DEAD_LETTER_KEY = "forge:retry:dlq" RETRY_ATTEMPTS_KEY = "forge:retry:attempts" +TERMINAL_NOTIFICATION_QUEUE_KEY = "forge:retry:terminal-notifications" # Retry configuration MAX_RETRY_ATTEMPTS = 3 RETRY_BACKOFF_MULTIPLIER = 2 INITIAL_RETRY_DELAY_SECONDS = 30 MAX_RETRY_DELAY_SECONDS = 3600 # 1 hour +RETRY_CLAIM_LEASE_SECONDS = 900 # 15 minutes +RETRY_CLAIM_RENEW_SECONDS = RETRY_CLAIM_LEASE_SECONDS // 3 +TERMINAL_NOTIFICATION_LEASE_SECONDS = 300 + +# Atomically move due entries beyond the visibility window before returning +# them. Other pollers cannot claim the same entries while this lease is active; +# if a worker exits before cleanup, the entries become visible again. +_CLAIM_DUE_MESSAGES_SCRIPT = """ +local entries = redis.call( + "ZRANGEBYSCORE", KEYS[1], "-inf", ARGV[1], "LIMIT", 0, ARGV[3] +) +for _, entry in ipairs(entries) do + redis.call("ZADD", KEYS[1], ARGV[2], entry) +end +return entries +""" + +_RENEW_CLAIM_SCRIPT = """ +local score = redis.call("ZSCORE", KEYS[1], ARGV[1]) +if score and tonumber(score) == tonumber(ARGV[2]) then + redis.call("ZADD", KEYS[1], ARGV[3], ARGV[1]) + return 1 +end +return 0 +""" + +_MOVE_TO_DEAD_LETTER_SCRIPT = """ +redis.call("RPUSH", KEYS[1], ARGV[1]) +redis.call("ZADD", KEYS[2], ARGV[2], ARGV[1]) +return 1 +""" + + +def _now_timestamp() -> float: + return datetime.utcnow().timestamp() @dataclass @@ -32,6 +68,7 @@ class RetryEntry: attempt: int next_retry: datetime last_error: str + lease_until: float | None = None def to_dict(self) -> dict[str, Any]: """Convert to dictionary for Redis storage.""" @@ -58,14 +95,24 @@ def from_dict(cls, data: dict[str, Any]) -> "RetryEntry": ) +@dataclass +class TerminalNotification: + """A leased terminal-failure notification from the durable outbox.""" + + message: QueueMessage + error: str + serialized: str | bytes + lease_until: float + + class RetryQueue: """Manages webhook retry queue with exponential backoff and dead-letter.""" - def __init__(self): + def __init__(self) -> None: """Initialize retry queue.""" - self._redis = None + self._redis: Any = None - async def _get_redis(self): + async def _get_redis(self) -> Any: """Get Redis client.""" if self._redis is None: self._redis = await get_redis_client() @@ -147,43 +194,111 @@ async def _move_to_dead_letter( "failed_at": datetime.utcnow().isoformat(), } - await redis.rpush(DEAD_LETTER_KEY, json.dumps(entry)) + serialized = json.dumps(entry) + await redis.eval( + _MOVE_TO_DEAD_LETTER_SCRIPT, + 2, + DEAD_LETTER_KEY, + TERMINAL_NOTIFICATION_QUEUE_KEY, + serialized, + _now_timestamp(), + ) logger.warning( f"Message {message_id} moved to dead-letter queue after " f"{attempt} attempts. Error: {error}" ) - async def get_due_messages(self, limit: int = 10) -> list[RetryEntry]: - """Get messages that are due for retry. + async def claim_due_messages(self, limit: int = 10) -> list[RetryEntry]: + """Atomically lease and return messages that are due for retry. Args: limit: Maximum number of messages to return. Returns: - List of retry entries ready to be processed. + List of retry entries exclusively leased to this poller for the + visibility window. """ redis = await self._get_redis() - now = datetime.utcnow().timestamp() + now = _now_timestamp() + lease_until = now + RETRY_CLAIM_LEASE_SECONDS - # Get messages with score <= now - entries = await redis.zrangebyscore( + entries = await redis.eval( + _CLAIM_DUE_MESSAGES_SCRIPT, + 1, RETRY_QUEUE_KEY, - "-inf", now, - start=0, - num=limit, + lease_until, + limit, ) results = [] for entry_json in entries: try: data = json.loads(entry_json) - results.append(RetryEntry.from_dict(data)) + entry = RetryEntry.from_dict(data) + entry.lease_until = lease_until + results.append(entry) except (json.JSONDecodeError, KeyError) as e: logger.error(f"Failed to parse retry entry: {e}") return results + async def renew_retry_claim(self, entry: RetryEntry) -> bool: + """Extend a retry lease only while this worker still owns it.""" + if entry.lease_until is None: + return False + + redis = await self._get_redis() + lease_until = _now_timestamp() + RETRY_CLAIM_LEASE_SECONDS + renewed = await redis.eval( + _RENEW_CLAIM_SCRIPT, + 1, + RETRY_QUEUE_KEY, + json.dumps(entry.to_dict()), + entry.lease_until, + lease_until, + ) + if renewed: + entry.lease_until = lease_until + return bool(renewed) + + async def claim_due_terminal_notifications(self, limit: int = 10) -> list[TerminalNotification]: + """Lease due terminal notifications for exclusive delivery.""" + redis = await self._get_redis() + now = _now_timestamp() + lease_until = now + TERMINAL_NOTIFICATION_LEASE_SECONDS + entries = await redis.eval( + _CLAIM_DUE_MESSAGES_SCRIPT, + 1, + TERMINAL_NOTIFICATION_QUEUE_KEY, + now, + lease_until, + limit, + ) + + results = [] + for serialized in entries: + try: + data = json.loads(serialized) + msg_data = dict(data["message"]) + message_id = msg_data.pop("message_id", "") + results.append( + TerminalNotification( + message=QueueMessage.from_redis(message_id, msg_data), + error=data["error"], + serialized=serialized, + lease_until=lease_until, + ) + ) + except (json.JSONDecodeError, KeyError) as exc: + logger.error(f"Failed to parse terminal notification: {exc}") + return results + + async def remove_terminal_notification(self, entry: TerminalNotification) -> None: + """Remove an outbox entry after successful delivery.""" + redis = await self._get_redis() + await redis.zrem(TERMINAL_NOTIFICATION_QUEUE_KEY, entry.serialized) + async def remove_from_retry(self, entry: RetryEntry) -> None: """Remove a message from the retry queue after successful processing. @@ -261,6 +376,7 @@ async def requeue_dead_letter(self, index: int) -> bool: # Reset attempt counter message_id = f"{message.source}:{message.ticket_key}:{message.event_id}" await redis.delete(f"{RETRY_ATTEMPTS_KEY}:{message_id}") + await redis.zrem(TERMINAL_NOTIFICATION_QUEUE_KEY, entries[0]) # Add back to retry queue entry = RetryEntry( diff --git a/tests/integration/redis/test_queue_integration.py b/tests/integration/redis/test_queue_integration.py index d6ccdd2a..826aa650 100644 --- a/tests/integration/redis/test_queue_integration.py +++ b/tests/integration/redis/test_queue_integration.py @@ -5,6 +5,8 @@ """ import asyncio +import json +from datetime import datetime, timedelta import pytest @@ -12,6 +14,7 @@ from forge.queue.consumer import CONSUMER_GROUP, QueueConsumer from forge.queue.models import QueueMessage from forge.queue.producer import GITHUB_STREAM, JIRA_STREAM, QueueProducer +from forge.queue.retry import RETRY_QUEUE_KEY, RetryEntry, RetryQueue @pytest.mark.integration @@ -262,6 +265,48 @@ async def handler(message: QueueMessage): assert len(messages) == 0 +@pytest.mark.integration +class TestRetryQueue: + """Test retry leasing against a real shared Redis instance.""" + + async def test_due_entry_is_claimed_by_only_one_worker(self, redis_client): + """Concurrent pollers must not receive the same due retry entry.""" + message = QueueMessage( + message_id="1-0", + event_id="lease-test-1", + source=EventSource.JIRA, + event_type="issue_updated", + ticket_key="TEST-LEASE", + ) + entry = RetryEntry( + message=message, + attempt=1, + next_retry=datetime.utcnow() - timedelta(seconds=1), + last_error="temporary failure", + ) + serialized_entry = json.dumps(entry.to_dict()) + await redis_client.zadd( + RETRY_QUEUE_KEY, + {serialized_entry: entry.next_retry.timestamp()}, + ) + + first = RetryQueue() + second = RetryQueue() + first._redis = redis_client + second._redis = redis_client + + claims = await asyncio.gather( + first.claim_due_messages(), + second.claim_due_messages(), + ) + + assert sorted(len(claim) for claim in claims) == [0, 1] + assert sum(claim[0].message.event_id == message.event_id for claim in claims if claim) == 1 + lease_score = await redis_client.zscore(RETRY_QUEUE_KEY, serialized_entry) + assert lease_score is not None + assert lease_score > datetime.utcnow().timestamp() + + @pytest.mark.integration class TestQueueMessageSerialization: """Test message serialization through the queue.""" diff --git a/tests/unit/integrations/jira/test_client.py b/tests/unit/integrations/jira/test_client.py index 321668b7..872d2336 100644 --- a/tests/unit/integrations/jira/test_client.py +++ b/tests/unit/integrations/jira/test_client.py @@ -129,6 +129,44 @@ async def test_add_structured_comment_includes_interaction_options_outside_marke assert "## 🤖 Forge interaction options" not in body[:marker_end] +class TestJiraClientComments: + """Tests for paginated comment retrieval.""" + + @pytest.mark.asyncio + async def test_get_comments_reads_marker_from_later_page(self): + with patch("forge.integrations.jira.client.get_settings") as mock_settings: + mock_settings.return_value.jira_base_url = "https://test.atlassian.net" + mock_settings.return_value.jira_api_token = MagicMock() + mock_settings.return_value.jira_api_token.get_secret_value.return_value = "token" + mock_settings.return_value.jira_user_email = "test@example.com" + jira = JiraClient() + + first_response = MagicMock() + first_response.json.return_value = { + "startAt": 0, + "maxResults": 1, + "total": 2, + "comments": [{"id": "1", "body": "Older comment"}], + } + second_response = MagicMock() + second_response.json.return_value = { + "startAt": 1, + "maxResults": 1, + "total": 2, + "comments": [{"id": "2", "body": "Event/correlation ID: evt-terminal-1"}], + } + http = AsyncMock() + http.get = AsyncMock(side_effect=[first_response, second_response]) + + with patch.object(jira, "_get_client", return_value=http): + comments = await jira.get_comments("TEST-123") + + assert [comment.id for comment in comments] == ["1", "2"] + assert "evt-terminal-1" in comments[1].body + assert http.get.await_args_list[0].kwargs["params"]["startAt"] == 0 + assert http.get.await_args_list[1].kwargs["params"]["startAt"] == 1 + + class TestJiraClientLabels: """Tests for label operations.""" diff --git a/tests/unit/orchestrator/test_worker.py b/tests/unit/orchestrator/test_worker.py index ac68b0b3..2bf4f7a7 100644 --- a/tests/unit/orchestrator/test_worker.py +++ b/tests/unit/orchestrator/test_worker.py @@ -74,6 +74,57 @@ async def test_report_new_workflow_error_skips_non_reportable_errors( notify.assert_not_awaited() +@pytest.mark.asyncio +async def test_terminal_failure_posts_sanitized_recovery_comment(): + worker = OrchestratorWorker(consumer_name="test-worker") + message = QueueMessage( + message_id="1-0", + event_id="evt-terminal-1", + source=EventSource.JIRA, + event_type="issue_updated", + ticket_key="TEST-123", + ) + jira = AsyncMock() + jira.get_comments = AsyncMock(return_value=[]) + + with patch("forge.orchestrator.worker.JiraClient", return_value=jira): + await worker._handle_terminal_failure( + message, + "clone https://ghp_abcdefghijklmnopqrstuvwxyz123456@github.com/acme/repo failed", + ) + + jira.add_error_comment.assert_awaited_once() + kwargs = jira.add_error_comment.await_args.kwargs + assert kwargs["issue_key"] == "TEST-123" + assert "[REDACTED]" in kwargs["error_message"] + assert "ghp_" not in kwargs["error_message"] + assert "Event/correlation ID: evt-terminal-1" in kwargs["error_message"] + assert "Recovery:" in kwargs["error_message"] + jira.close.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_terminal_failure_skips_existing_event_comment(): + worker = OrchestratorWorker(consumer_name="test-worker") + message = QueueMessage( + message_id="1-0", + event_id="evt-terminal-1", + source=EventSource.JIRA, + event_type="issue_updated", + ticket_key="TEST-123", + ) + jira = AsyncMock() + jira.get_comments = AsyncMock( + return_value=[MagicMock(body="Event/correlation ID: evt-terminal-1")] + ) + + with patch("forge.orchestrator.worker.JiraClient", return_value=jira): + await worker._handle_terminal_failure(message, "failed") + + jira.add_error_comment.assert_not_awaited() + jira.close.assert_awaited_once() + + class TestQuestionDetection: """Tests for Q&A mode question detection.""" diff --git a/tests/unit/queue/test_consumer.py b/tests/unit/queue/test_consumer.py index 4ce3fcc3..7ba21fcd 100644 --- a/tests/unit/queue/test_consumer.py +++ b/tests/unit/queue/test_consumer.py @@ -81,7 +81,7 @@ def _make_consumer(redis_mock: MagicMock, max_tasks: int = 20) -> QueueConsumer: # the message for retry. retry_mock = MagicMock(spec=RetryQueue) retry_mock.enqueue_for_retry = AsyncMock(return_value=True) # queued, not DLQ - retry_mock.get_due_messages = AsyncMock(return_value=[]) + retry_mock.claim_due_messages = AsyncMock(return_value=[]) retry_mock.remove_from_retry = AsyncMock() retry_mock.remove_from_retry_without_counter_reset = AsyncMock() consumer._retry_queue = retry_mock diff --git a/tests/unit/queue/test_consumer_retry.py b/tests/unit/queue/test_consumer_retry.py index b5d82323..991c0092 100644 --- a/tests/unit/queue/test_consumer_retry.py +++ b/tests/unit/queue/test_consumer_retry.py @@ -8,7 +8,7 @@ from forge.models.events import EventSource from forge.queue.consumer import CONSUMER_GROUP, QueueConsumer from forge.queue.models import QueueMessage -from forge.queue.retry import RetryEntry, RetryQueue +from forge.queue.retry import RetryEntry, RetryQueue, TerminalNotification # --------------------------------------------------------------------------- # Helpers @@ -55,9 +55,12 @@ def make_consumer() -> QueueConsumer: # Replace the real RetryQueue with a mock consumer._retry_queue = MagicMock(spec=RetryQueue) consumer._retry_queue.enqueue_for_retry = AsyncMock(return_value=True) - consumer._retry_queue.get_due_messages = AsyncMock(return_value=[]) + consumer._retry_queue.claim_due_messages = AsyncMock(return_value=[]) + consumer._retry_queue.renew_retry_claim = AsyncMock(return_value=True) consumer._retry_queue.remove_from_retry = AsyncMock() consumer._retry_queue.remove_from_retry_without_counter_reset = AsyncMock() + consumer._retry_queue.claim_due_terminal_notifications = AsyncMock(return_value=[]) + consumer._retry_queue.remove_terminal_notification = AsyncMock() return consumer @@ -182,6 +185,52 @@ async def xreadgroup_side_effect(*_args, **_kwargs): # xack must be called to clear PEL after DLQ move redis_mock.xack.assert_called_once_with("stream:jira", CONSUMER_GROUP, "1-0") + @pytest.mark.asyncio + async def test_dlq_ack_does_not_wait_for_terminal_callback(self): + """Durable notification delivery is independent from stream acknowledgement.""" + callback = AsyncMock(side_effect=RuntimeError("Jira unavailable")) + consumer = QueueConsumer("test-worker", terminal_failure_handler=callback) + consumer._retry_queue = MagicMock(spec=RetryQueue) + consumer._retry_queue.enqueue_for_retry = AsyncMock(return_value=False) + + redis_mock = AsyncMock() + configure_redis_mock(redis_mock) + redis_mock.xack = AsyncMock() + consumer._redis = redis_mock + consumer.register_handler( + EventSource.JIRA, + AsyncMock(side_effect=RuntimeError("permanent failure")), + ) + + await consumer._process_message(make_message(), "stream:jira") + + callback.assert_not_awaited() + redis_mock.xack.assert_awaited_once_with("stream:jira", CONSUMER_GROUP, "1-0") + + @pytest.mark.asyncio + async def test_terminal_callback_failure_is_retried_until_success(self): + """An outbox entry remains after failure and is removed after later success.""" + callback = AsyncMock(side_effect=[RuntimeError("Jira unavailable"), None]) + consumer = make_consumer() + consumer._terminal_failure_handler = callback + message = make_message() + notification = TerminalNotification( + message=message, + error="failed", + serialized=b'{"event":"evt-001"}', + lease_until=0, + ) + consumer._retry_queue.claim_due_terminal_notifications = AsyncMock( + side_effect=[[notification], [notification]] + ) + + await consumer._process_terminal_notifications_once() + consumer._retry_queue.remove_terminal_notification.assert_not_awaited() + + await consumer._process_terminal_notifications_once() + assert callback.await_count == 2 + consumer._retry_queue.remove_terminal_notification.assert_awaited_once_with(notification) + # --------------------------------------------------------------------------- # _process_retry_queue — background poller @@ -203,7 +252,7 @@ async def test_successful_retry_removes_entry(self): call_count = 0 - async def get_due_side_effect(*_args, **_kwargs): + async def claim_due_side_effect(*_args, **_kwargs): nonlocal call_count call_count += 1 if call_count == 1: @@ -211,7 +260,7 @@ async def get_due_side_effect(*_args, **_kwargs): consumer._running = False return [] - consumer._retry_queue.get_due_messages = AsyncMock(side_effect=get_due_side_effect) + consumer._retry_queue.claim_due_messages = AsyncMock(side_effect=claim_due_side_effect) handler = AsyncMock() consumer.register_handler(EventSource.JIRA, handler) @@ -235,7 +284,7 @@ async def test_failed_retry_reenqueues(self): call_count = 0 - async def get_due_side_effect(*_args, **_kwargs): + async def claim_due_side_effect(*_args, **_kwargs): nonlocal call_count call_count += 1 if call_count == 1: @@ -243,7 +292,7 @@ async def get_due_side_effect(*_args, **_kwargs): consumer._running = False return [] - consumer._retry_queue.get_due_messages = AsyncMock(side_effect=get_due_side_effect) + consumer._retry_queue.claim_due_messages = AsyncMock(side_effect=claim_due_side_effect) failing_handler = AsyncMock(side_effect=RuntimeError("still broken")) consumer.register_handler(EventSource.JIRA, failing_handler) @@ -261,6 +310,83 @@ async def get_due_side_effect(*_args, **_kwargs): assert enqueue_args[0].event_id == "evt-001" assert "still broken" in enqueue_args[1] + @pytest.mark.asyncio + async def test_exhausted_retry_invokes_terminal_callback(self): + """The outbox poller delivers after the retry moves the event to the DLQ.""" + callback = AsyncMock() + consumer = make_consumer() + consumer._terminal_failure_handler = callback + message = make_message() + entry = RetryEntry( + message=message, + attempt=3, + next_retry=datetime.utcnow(), + last_error="still broken", + ) + notification = TerminalNotification( + message=message, + error="retries exhausted", + serialized=b'{"event":"evt-001"}', + lease_until=0, + ) + consumer._retry_queue.claim_due_messages = AsyncMock(return_value=[entry]) + consumer._retry_queue.claim_due_terminal_notifications = AsyncMock( + return_value=[notification] + ) + consumer._retry_queue.enqueue_for_retry = AsyncMock(return_value=False) + consumer.register_handler( + EventSource.JIRA, + AsyncMock(side_effect=RuntimeError("retries exhausted")), + ) + + await consumer._process_due_retries_once() + await consumer._process_terminal_notifications_once() + + callback.assert_awaited_once_with(message, "retries exhausted") + consumer._redis.xack.assert_awaited_once_with( + "forge:events:jira", CONSUMER_GROUP, message.message_id + ) + consumer._retry_queue.remove_terminal_notification.assert_awaited_once_with(notification) + + @pytest.mark.asyncio + async def test_terminal_callback_failure_still_acknowledges_retry(self): + """A callback failure cannot leave a terminal retry in the PEL.""" + callback = AsyncMock(side_effect=RuntimeError("Jira unavailable")) + consumer = make_consumer() + consumer._terminal_failure_handler = callback + message = make_message() + entry = RetryEntry( + message=message, + attempt=3, + next_retry=datetime.utcnow(), + last_error="still broken", + ) + + notification = TerminalNotification( + message=message, + error="retries exhausted", + serialized=b'{"event":"evt-001"}', + lease_until=0, + ) + consumer._retry_queue.claim_due_messages = AsyncMock(return_value=[entry]) + consumer._retry_queue.claim_due_terminal_notifications = AsyncMock( + return_value=[notification] + ) + consumer._retry_queue.enqueue_for_retry = AsyncMock(return_value=False) + consumer.register_handler( + EventSource.JIRA, + AsyncMock(side_effect=RuntimeError("retries exhausted")), + ) + + await consumer._process_due_retries_once() + await consumer._process_terminal_notifications_once() + + callback.assert_awaited_once_with(message, "retries exhausted") + consumer._redis.xack.assert_awaited_once_with( + "forge:events:jira", CONSUMER_GROUP, message.message_id + ) + consumer._retry_queue.remove_terminal_notification.assert_not_awaited() + @pytest.mark.asyncio async def test_empty_retry_queue_no_action(self): """Poller with no due messages does nothing.""" @@ -268,14 +394,14 @@ async def test_empty_retry_queue_no_action(self): call_count = 0 - async def get_due_side_effect(*_args, **_kwargs): + async def claim_due_side_effect(*_args, **_kwargs): nonlocal call_count call_count += 1 if call_count >= 2: consumer._running = False return [] - consumer._retry_queue.get_due_messages = AsyncMock(side_effect=get_due_side_effect) + consumer._retry_queue.claim_due_messages = AsyncMock(side_effect=claim_due_side_effect) with patch("forge.queue.consumer.asyncio.sleep", new_callable=AsyncMock): await consumer._process_retry_queue() diff --git a/tests/unit/queue/test_retry_queue.py b/tests/unit/queue/test_retry_queue.py index 1ee18088..962ab626 100644 --- a/tests/unit/queue/test_retry_queue.py +++ b/tests/unit/queue/test_retry_queue.py @@ -1,5 +1,6 @@ """Unit tests for RetryQueue class.""" +import asyncio import json from datetime import datetime from unittest.mock import AsyncMock, patch @@ -11,7 +12,9 @@ from forge.queue.retry import ( DEAD_LETTER_KEY, MAX_RETRY_ATTEMPTS, + RETRY_CLAIM_LEASE_SECONDS, RETRY_QUEUE_KEY, + TERMINAL_NOTIFICATION_QUEUE_KEY, RetryEntry, RetryQueue, ) @@ -38,6 +41,7 @@ def make_redis_mock() -> AsyncMock: mock.expire = AsyncMock() mock.zadd = AsyncMock() mock.zrangebyscore = AsyncMock(return_value=[]) + mock.eval = AsyncMock(return_value=[]) mock.zrem = AsyncMock() mock.delete = AsyncMock() mock.rpush = AsyncMock() @@ -77,9 +81,10 @@ async def test_exceeds_max_attempts_moves_to_dlq(self): result = await rq.enqueue_for_retry(message, "final failure") assert result is False - redis.rpush.assert_called_once() - dlq_call_args = redis.rpush.call_args[0] - assert dlq_call_args[0] == DEAD_LETTER_KEY + eval_args = redis.eval.await_args.args + assert eval_args[1] == 2 + assert eval_args[2:4] == (DEAD_LETTER_KEY, TERMINAL_NOTIFICATION_QUEUE_KEY) + assert json.loads(eval_args[4])["message"]["event_id"] == message.event_id @pytest.mark.asyncio async def test_backoff_increases_with_attempt(self): @@ -111,11 +116,11 @@ async def capture_zadd(_key, mapping): # --------------------------------------------------------------------------- -# get_due_messages +# claim_due_messages # --------------------------------------------------------------------------- -class TestGetDueMessages: +class TestClaimDueMessages: @pytest.mark.asyncio async def test_returns_parsed_entries(self): message = make_message() @@ -130,25 +135,154 @@ async def test_returns_parsed_entries(self): rq = RetryQueue() redis = make_redis_mock() - redis.zrangebyscore = AsyncMock(return_value=[json.dumps(raw_entry).encode()]) + redis.eval = AsyncMock(return_value=[json.dumps(raw_entry).encode()]) rq._redis = redis - results = await rq.get_due_messages() + results = await rq.claim_due_messages() assert len(results) == 1 assert results[0].message.event_id == "evt-001" assert results[0].attempt == 1 + eval_args = redis.eval.await_args.args + assert eval_args[1] == 1 + assert eval_args[2] == RETRY_QUEUE_KEY + assert eval_args[4] - eval_args[3] == RETRY_CLAIM_LEASE_SECONDS + assert eval_args[5] == 10 @pytest.mark.asyncio async def test_empty_queue_returns_empty_list(self): rq = RetryQueue() redis = make_redis_mock() - redis.zrangebyscore = AsyncMock(return_value=[]) + redis.eval = AsyncMock(return_value=[]) rq._redis = redis - results = await rq.get_due_messages() + results = await rq.claim_due_messages() assert results == [] + @pytest.mark.asyncio + async def test_concurrent_workers_receive_entry_once(self): + """Two workers sharing Redis cannot claim the same due entry.""" + entry = RetryEntry( + message=make_message(), + attempt=1, + next_retry=datetime(2024, 1, 1), + last_error="oops", + ) + serialized_entry = json.dumps(entry.to_dict()).encode() + claim_lock = asyncio.Lock() + entry_is_due = True + + async def atomic_eval(*_args): + nonlocal entry_is_due + async with claim_lock: + if not entry_is_due: + return [] + entry_is_due = False + return [serialized_entry] + + redis = make_redis_mock() + redis.eval = AsyncMock(side_effect=atomic_eval) + first = RetryQueue() + second = RetryQueue() + first._redis = redis + second._redis = redis + + claims = await asyncio.gather( + first.claim_due_messages(), + second.claim_due_messages(), + ) + + assert sorted(len(claim) for claim in claims) == [0, 1] + assert redis.eval.await_count == 2 + + @pytest.mark.asyncio + async def test_renewal_keeps_entry_hidden_past_initial_lease(self): + entry = RetryEntry( + message=make_message(), + attempt=1, + next_retry=datetime(2024, 1, 1), + last_error="oops", + ) + serialized = json.dumps(entry.to_dict()).encode() + score = 0.0 + + async def eval_script(script, _keys, _key, *args): + nonlocal score + if "ZRANGEBYSCORE" in script: + now, lease_until, _limit = args + if score <= now: + score = float(lease_until) + return [serialized] + return [] + member, expected_score, lease_until = args + if member == serialized.decode() and score == float(expected_score): + score = float(lease_until) + return 1 + return 0 + + redis = make_redis_mock() + redis.eval = AsyncMock(side_effect=eval_script) + first = RetryQueue() + second = RetryQueue() + first._redis = redis + second._redis = redis + + with patch( + "forge.queue.retry._now_timestamp", + side_effect=[0, 600, 901], + ): + claimed = (await first.claim_due_messages())[0] + initial_lease = claimed.lease_until + + assert await first.renew_retry_claim(claimed) + assert claimed.lease_until is not None + assert initial_lease is not None + assert claimed.lease_until > initial_lease + + assert initial_lease == 900 + assert await second.claim_due_messages() == [] + + +class TestTerminalNotificationQueue: + @pytest.mark.asyncio + async def test_concurrent_workers_claim_notification_once(self): + message = make_message() + stored = json.dumps( + { + "message": {**message.to_dict(), "message_id": message.message_id}, + "error": "final failure", + "attempts": 4, + "failed_at": datetime.utcnow().isoformat(), + } + ).encode() + claim_lock = asyncio.Lock() + entry_is_due = True + + async def atomic_eval(*_args): + nonlocal entry_is_due + async with claim_lock: + if not entry_is_due: + return [] + entry_is_due = False + return [stored] + + redis = make_redis_mock() + redis.eval = AsyncMock(side_effect=atomic_eval) + first = RetryQueue() + second = RetryQueue() + first._redis = redis + second._redis = redis + + claims = await asyncio.gather( + first.claim_due_terminal_notifications(), + second.claim_due_terminal_notifications(), + ) + + assert sorted(len(claim) for claim in claims) == [0, 1] + notification = next(claim[0] for claim in claims if claim) + assert notification.message.event_id == message.event_id + assert notification.error == "final failure" + # --------------------------------------------------------------------------- # remove_from_retry @@ -197,3 +331,25 @@ async def test_removes_entry_without_clearing_counter(self): redis.zrem.assert_called_once_with(RETRY_QUEUE_KEY, json.dumps({"stub": True})) redis.delete.assert_not_called() + + +class TestRequeueDeadLetter: + @pytest.mark.asyncio + async def test_removes_stale_terminal_notification(self): + message = make_message() + stored = json.dumps( + { + "message": {**message.to_dict(), "message_id": message.message_id}, + "error": "final failure", + "attempts": 4, + "failed_at": datetime.utcnow().isoformat(), + } + ).encode() + redis = make_redis_mock() + redis.lrange = AsyncMock(return_value=[stored]) + rq = RetryQueue() + rq._redis = redis + + assert await rq.requeue_dead_letter(0) + + redis.zrem.assert_awaited_once_with(TERMINAL_NOTIFICATION_QUEUE_KEY, stored)