mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-07 16:19:09 +00:00
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:
parent
9af5e36ace
commit
a4dbe94ab9
5 changed files with 89 additions and 16 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue