From a4dbe94ab9ef53af4704ee9722c0d7e9020fad32 Mon Sep 17 00:00:00 2001 From: amoussa1229 <166072409+amoussa1229@users.noreply.github.com> Date: Tue, 30 Jun 2026 22:18:46 +0000 Subject: [PATCH] Restore reject backstop for autofix dispatch A burst of near-simultaneous CI events for one head SHA can slip past the busy-check before the dedupe SHA is recorded, so dispatch the autofix path with multitask_strategy=reject (dev's prior platform default) to drop duplicate concurrent creates instead of letting them interrupt each other. Also make the completion failure-reply dedup claim-then-post and drop the unreachable interrupted branch. --- agent/ci_autofix.py | 6 ++++ agent/completion.py | 53 +++++++++++++++++++++++--------- agent/dispatch.py | 7 +++-- tests/test_ci_autofix.py | 9 ++++++ tests/test_completion_webhook.py | 30 ++++++++++++++++++ 5 files changed, 89 insertions(+), 16 deletions(-) diff --git a/agent/ci_autofix.py b/agent/ci_autofix.py index 721aa99b..13571f58 100644 --- a/agent/ci_autofix.py +++ b/agent/ci_autofix.py @@ -280,11 +280,17 @@ async def _dispatch_or_batch( logger.info("Agent thread %s busy; batching auto-fix event %s", thread_id, reason) await _mark_pending_autofix_event(thread_id, reason, detail) return "batched" + # The busy-check above has a TOCTOU window (the dedupe SHA is only recorded + # after dispatch), so a burst of near-simultaneous CI events for one head SHA + # can all pass the gate. Dispatch with ``reject`` — matching ``dev``'s prior + # platform default — so the platform drops the duplicate concurrent creates + # instead of letting them interrupt each other. await dispatch_agent_run( thread_id, prompt, configurable, source=str(configurable.get("source") or "github_autofix"), + multitask_strategy="reject", ) logger.info( "Created auto-fix run for thread %s (source=%s)", thread_id, configurable.get("source") diff --git a/agent/completion.py b/agent/completion.py index f2e030d5..7b5ddda1 100644 --- a/agent/completion.py +++ b/agent/completion.py @@ -16,6 +16,7 @@ from __future__ import annotations import hmac import logging import os +from collections.abc import Awaitable, Callable from typing import Any from .utils.github_app import get_github_app_installation_token @@ -33,6 +34,11 @@ logger = logging.getLogger(__name__) _TERMINAL_FAILURE_STATUSES = frozenset({"error", "timeout"}) _FAILURE_REPLY_FLAG = "failure_reply_posted" + +class _ClaimFailed(Exception): + """Raised when the dedup flag couldn't be claimed, so we skip the post.""" + + # Shared-secret bearer token proving a /webhooks/run-complete call came from our # own dispatch (which appends ?token= when this is set) rather than from an # attacker hitting the public route. Fail closed when unset: the route rejects @@ -58,20 +64,26 @@ def verify_run_complete_token(token: str | None) -> bool: def _failure_text(status: str) -> str: - if status == "timeout": - reason = "timed out" - elif status == "interrupted": - reason = "was interrupted before it could finish" - else: - reason = "hit an unexpected error" + reason = "timed out" if status == "timeout" else "hit an unexpected error" return ( f"⚠️ I wasn't able to finish that — the run {reason}. " "Send another message and I'll pick it back up." ) -async def _post_failure_reply(thread_id: str, metadata: dict[str, Any], status: str) -> bool: - """Post a failure reply to the run's originating channel. Best-effort.""" +async def _post_failure_reply( + thread_id: str, + metadata: dict[str, Any], + status: str, + *, + claim: Callable[[], Awaitable[None]], +) -> bool: + """Post a failure reply to the run's originating channel. Best-effort. + + ``claim`` is awaited immediately before the network post (claim-then-post), + only on a branch that actually delivers, so a retried/concurrent webhook + can't double-post and threads with no channel never burn the flag. + """ source = metadata.get("source") ctx = metadata.get("source_context") ctx = ctx if isinstance(ctx, dict) else {} @@ -83,6 +95,7 @@ async def _post_failure_reply(thread_id: str, metadata: dict[str, Any], status: channel_id = slack_thread.get("channel_id") thread_ts = slack_thread.get("thread_ts") if channel_id and thread_ts: + await claim() return await post_slack_thread_reply(channel_id, thread_ts, text) return False @@ -91,6 +104,7 @@ async def _post_failure_reply(thread_id: str, metadata: dict[str, Any], status: if isinstance(linear_issue, dict): issue_id = linear_issue.get("id") if issue_id: + await claim() return await comment_on_linear_issue(issue_id, text) return False @@ -104,6 +118,7 @@ async def _post_failure_reply(thread_id: str, metadata: dict[str, Any], status: if isinstance(repo_config, dict) and isinstance(number, int): token = await get_github_app_installation_token() if token: + await claim() return await post_github_comment(repo_config, number, text, token=token) return False @@ -136,13 +151,23 @@ async def handle_run_completion(payload: dict[str, Any]) -> dict[str, str]: if metadata.get(_FAILURE_REPLY_FLAG): return {"status": "ignored", "reason": "failure reply already posted"} - posted = await _post_failure_reply(thread_id, metadata, status) - if not posted: - return {"status": "ignored", "reason": "no reply posted"} + # Claim-then-post: set the dedup flag immediately before the actual post (via + # the claim callback) so a retried/concurrent completion webhook can't + # double-post. The flag is only claimed on a branch that delivers, so a + # thread with no reply channel never burns it. If the claim itself fails we + # skip the post, leaving the flag unset so a later retry can try again. + async def _claim() -> None: + try: + await client.threads.update(thread_id=thread_id, metadata={_FAILURE_REPLY_FLAG: True}) + except Exception as exc: # noqa: BLE001 + logger.warning("run-complete: could not flag thread %s", thread_id, exc_info=True) + raise _ClaimFailed from exc try: - await client.threads.update(thread_id=thread_id, metadata={_FAILURE_REPLY_FLAG: True}) - except Exception: # noqa: BLE001 - logger.warning("run-complete: could not flag thread %s", thread_id, exc_info=True) + posted = await _post_failure_reply(thread_id, metadata, status, claim=_claim) + except _ClaimFailed: + return {"status": "error", "reason": "could not claim failure reply"} + if not posted: + return {"status": "ignored", "reason": "no reply posted"} logger.info("Posted failure reply for thread %s (status=%s)", thread_id, status) return {"status": "ok", "reason": "failure reply posted"} diff --git a/agent/dispatch.py b/agent/dispatch.py index cd6dd554..aa72225e 100644 --- a/agent/dispatch.py +++ b/agent/dispatch.py @@ -62,12 +62,15 @@ async def dispatch_agent_run( assistant_id: str = "agent", metadata: dict[str, Any] | None = None, client: LangGraphClient | None = None, + multitask_strategy: str = "interrupt", ) -> dict[str, Any]: """Create (or interrupt-and-resume) a run for ``thread_id``. Routes every Slack / Linear / GitHub / dashboard trigger through one contract. ``source`` is for logging/metadata only; ``assistant_id`` selects - the graph (``"agent"`` or ``"reviewer"``). + the graph (``"agent"`` or ``"reviewer"``). ``multitask_strategy`` defaults to + ``"interrupt"`` (human follow-ups halt + resume); autofix passes ``"reject"`` + so a burst of concurrent CI events for one head SHA can't interrupt each other. """ client = client or dispatch_client() run = await client.runs.create( @@ -75,7 +78,7 @@ async def dispatch_agent_run( assistant_id, input={"messages": [{"role": "user", "content": content}]}, config={"configurable": configurable, "metadata": metadata or {}}, - multitask_strategy="interrupt", + multitask_strategy=multitask_strategy, durability="sync", webhook=COMPLETION_WEBHOOK_URL, if_not_exists="create", diff --git a/tests/test_ci_autofix.py b/tests/test_ci_autofix.py index fd204161..22684e48 100644 --- a/tests/test_ci_autofix.py +++ b/tests/test_ci_autofix.py @@ -89,6 +89,15 @@ async def test_dispatch_happy_path(happy: dict[str, Any]) -> None: happy["status_check"].assert_awaited() +@pytest.mark.asyncio +async def test_autofix_dispatch_uses_reject_strategy(happy: dict[str, Any]) -> None: + # A burst of concurrent CI events for one head SHA can slip past the busy-check + # before the dedupe SHA is recorded; dispatching with "reject" lets the platform + # drop the duplicate concurrent creates instead of interrupting each other. + await _run() + assert happy["runs_create"].await_args.kwargs["multitask_strategy"] == "reject" + + @pytest.mark.asyncio async def test_batches_when_thread_busy(happy: dict[str, Any], monkeypatch) -> None: monkeypatch.setattr(ci_autofix, "get_thread_active_status", AsyncMock(return_value=True)) diff --git a/tests/test_completion_webhook.py b/tests/test_completion_webhook.py index 4ee233ac..e909d91d 100644 --- a/tests/test_completion_webhook.py +++ b/tests/test_completion_webhook.py @@ -98,6 +98,36 @@ async def test_missing_thread_id_is_ignored() -> None: assert result["status"] == "ignored" +@pytest.mark.asyncio +async def test_claims_flag_before_posting(monkeypatch: pytest.MonkeyPatch) -> None: + # Claim-then-post: the dedup flag must be set before the reply is posted so a + # retried/concurrent webhook can't double-post the canned failure message. + client = _FakeClient(_slack_metadata()) + monkeypatch.setattr(completion, "langgraph_client", lambda: client) + + async def _reply(*_args: Any, **_kwargs: Any) -> bool: + assert client.threads.updates == [{"failure_reply_posted": True}] + return True + + monkeypatch.setattr(completion, "post_slack_thread_reply", AsyncMock(side_effect=_reply)) + + result = await completion.handle_run_completion({"thread_id": "t1", "status": "error"}) + assert result["status"] == "ok" + + +@pytest.mark.asyncio +async def test_does_not_post_when_claim_fails(monkeypatch: pytest.MonkeyPatch) -> None: + client = _FakeClient(_slack_metadata()) + client.threads.update = AsyncMock(side_effect=RuntimeError("boom")) + monkeypatch.setattr(completion, "langgraph_client", lambda: client) + reply = AsyncMock(return_value=True) + monkeypatch.setattr(completion, "post_slack_thread_reply", reply) + + result = await completion.handle_run_completion({"thread_id": "t1", "status": "error"}) + assert result["status"] == "error" + reply.assert_not_called() + + @pytest.mark.asyncio async def test_no_reply_channel_does_not_flag(monkeypatch: pytest.MonkeyPatch) -> None: client = _FakeClient({"source": "schedule"})