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.
This commit is contained in:
amoussa1229 2026-06-30 22:18:46 +00:00
parent 9af5e36ace
commit a4dbe94ab9
5 changed files with 89 additions and 16 deletions

View file

@ -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")

View file

@ -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"}

View file

@ -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",

View file

@ -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))

View file

@ -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"})