mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-07 16:19:09 +00:00
Adopt upstream modular webhook skeleton (#1621)
Apply the durable-interrupt-dispatch refactor: split the monolithic
webapp.py into a thin routing layer plus per-source handlers in
webhooks/{github,slack,linear}.py, and add completion.py, dispatch.py,
and reconcile.py. Reconcile fork divergence by keeping the Bedrock/
Fireworks cross-provider fallback, the no-agent-attribution prompt
policy, the dashboard-handoff re-export, and the Slack channel-info
cache. ci_autofix is restored on the new dispatch model in a later
commit.
Refs: #80
This commit is contained in:
parent
1f060f2a1d
commit
0cac1ad363
24 changed files with 2406 additions and 3266 deletions
|
|
@ -1,607 +0,0 @@
|
|||
"""Auto-fix CI failures and review feedback on agent-authored pull requests.
|
||||
|
||||
This is the shared core for "PR babysitting": when a CI check fails (or a
|
||||
reviewer leaves actionable feedback) on a PR that Open SWE opened, locate the
|
||||
originating agent thread and dispatch a confidence-gated fix run on it.
|
||||
|
||||
Both the GitHub webhook path (:mod:`agent.webapp`) and the polling fallback
|
||||
(:mod:`agent.ci_monitor`) call into here, so all the skip-rules, dedupe, and
|
||||
loop-capping live in one place. Skip-rules mirror Cursor/Claude Code:
|
||||
|
||||
* Only PRs Open SWE authored (an agent thread with this ``pr_url`` exists).
|
||||
* Skip failures inherited from the base branch.
|
||||
* Skip when the latest commit was authored by a human (don't fight pushes).
|
||||
* Dedupe per head SHA; cap total attempts.
|
||||
* Honor the per-user ``auto_fix_ci`` profile flag and the per-PR opt-out.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from langgraph_sdk import get_client
|
||||
|
||||
from .dashboard.agent_overrides import load_profile, resolve_login_from_email_async
|
||||
from .dashboard.autofix_state import is_pr_autofix_disabled
|
||||
from .dashboard.enabled_repos import is_review_repo_enabled
|
||||
from .reviewer_findings import REVIEWER_THREAD_KIND
|
||||
from .utils.dashboard_links import dashboard_thread_url
|
||||
from .utils.github_app import get_github_app_installation_token
|
||||
from .utils.github_checks import post_autofix_status_check
|
||||
from .utils.github_ci import (
|
||||
fetch_open_pr_for_branch,
|
||||
fetch_pr,
|
||||
has_repo_write_permission,
|
||||
head_commit_author_login,
|
||||
list_failing_check_runs,
|
||||
list_failing_statuses,
|
||||
names_failing_on_base,
|
||||
)
|
||||
from .utils.github_org_membership import INTERNAL_BOT_LOGINS
|
||||
from .utils.thread_ops import (
|
||||
is_thread_active,
|
||||
langgraph_client,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Hard cap on auto-fix follow-ups per PR so a failure the agent can't resolve
|
||||
# doesn't loop forever (Cursor caps at 10).
|
||||
MAX_AUTOFIX_ATTEMPTS = 10
|
||||
# Keep the dedupe list bounded on thread metadata.
|
||||
_MAX_HANDLED_KEYS = 30
|
||||
# Store location for batched auto-fix events, consumed by the message-queue middleware.
|
||||
_PENDING_AUTOFIX_NS = "autofix"
|
||||
_PENDING_AUTOFIX_KEY = "pending_event"
|
||||
|
||||
|
||||
def _dedupe_key(head_sha: str) -> str:
|
||||
return head_sha
|
||||
|
||||
|
||||
async def _user_autofix_enabled(github_login: str, user_email: str = "") -> bool:
|
||||
"""Check the per-user ``auto_fix_ci`` profile flag (defaults to True)."""
|
||||
login = github_login.strip() if isinstance(github_login, str) else ""
|
||||
if not login and user_email:
|
||||
login = await resolve_login_from_email_async(user_email)
|
||||
if not login:
|
||||
return True
|
||||
profile = await load_profile(login)
|
||||
if not isinstance(profile, dict):
|
||||
return True
|
||||
value = profile.get("auto_fix_ci")
|
||||
return value if isinstance(value, bool) else True
|
||||
|
||||
|
||||
async def find_agent_thread_for_pr(pr_url: str) -> tuple[str, dict[str, Any]] | None:
|
||||
"""Return ``(thread_id, metadata)`` of the agent thread that opened ``pr_url``.
|
||||
|
||||
Reviewer threads are skipped — only the coding-agent thread can push fixes.
|
||||
"""
|
||||
if not pr_url:
|
||||
return None
|
||||
client = get_client()
|
||||
try:
|
||||
threads = await client.threads.search(metadata={"pr_url": pr_url}, limit=10)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug("Could not search threads for PR %s", pr_url, exc_info=True)
|
||||
return None
|
||||
for thread in threads or []:
|
||||
metadata = thread.get("metadata") if isinstance(thread, dict) else None
|
||||
if not isinstance(metadata, dict):
|
||||
continue
|
||||
if metadata.get("kind") == REVIEWER_THREAD_KIND:
|
||||
continue
|
||||
if metadata.get("agent_kind") != "agent":
|
||||
continue
|
||||
thread_id = thread.get("thread_id") or thread.get("id")
|
||||
if isinstance(thread_id, str) and thread_id:
|
||||
return thread_id, metadata
|
||||
return None
|
||||
|
||||
|
||||
def _build_ci_fix_prompt(
|
||||
*,
|
||||
owner: str,
|
||||
repo: str,
|
||||
pr_number: int,
|
||||
pr_url: str,
|
||||
branch: str,
|
||||
head_sha: str,
|
||||
failing_checks: list[dict[str, Any]],
|
||||
) -> str:
|
||||
lines = []
|
||||
for check in failing_checks:
|
||||
name = check.get("name", "check")
|
||||
conclusion = check.get("conclusion", "failure")
|
||||
details = check.get("details_url") or ""
|
||||
suffix = f" — {details}" if details else ""
|
||||
lines.append(f"- {name} ({conclusion}){suffix}")
|
||||
failing_block = "\n".join(lines)
|
||||
return (
|
||||
"An automated CI check failed on a pull request you opened. Please "
|
||||
"investigate and fix it.\n\n"
|
||||
f"## Repository: {owner}/{repo}\n\n"
|
||||
f"## Pull Request: {pr_url} (#{pr_number})\n\n"
|
||||
f"## Branch: {branch}\n\n"
|
||||
f"## Head commit: {head_sha}\n\n"
|
||||
f"## Failing checks:\n{failing_block}\n\n"
|
||||
"Instructions:\n"
|
||||
"1. Make sure you are on the PR branch, then read the failing logs "
|
||||
"(e.g. `GH_TOKEN=dummy gh pr checks` and `GH_TOKEN=dummy gh run view "
|
||||
"<run-id> --log-failed`).\n"
|
||||
"2. Confidence gating — fix autonomously ONLY when the cause is clear "
|
||||
"and deterministic (lint/format, type errors, missing imports, failed "
|
||||
"assertions, snapshot updates, build errors). Commit and push to the "
|
||||
"existing branch; do NOT open a new PR.\n"
|
||||
"3. If the failure is ambiguous, flaky, infrastructure-related, appears "
|
||||
"pre-existing, or needs an architectural/design decision, do NOT guess. "
|
||||
"Post a short PR comment explaining what you found and what input you "
|
||||
"need, then stop.\n"
|
||||
"4. Never force-push. Never weaken or delete test assertions just to go "
|
||||
"green unless the behavior change is intentional and correct.\n"
|
||||
"5. Before finishing, re-check the PR's latest CI status and review "
|
||||
"comments. Address any newly failed checks or unhandled actionable "
|
||||
"comments that arrived while you were working.\n"
|
||||
"6. After you push, CI re-runs automatically — you don't need to merge."
|
||||
)
|
||||
|
||||
|
||||
def _build_review_feedback_prompt(
|
||||
*,
|
||||
owner: str,
|
||||
repo: str,
|
||||
pr_number: int,
|
||||
pr_url: str,
|
||||
reviewer: str,
|
||||
body: str,
|
||||
) -> str:
|
||||
return (
|
||||
"A reviewer left feedback on a pull request you opened. Please respond.\n\n"
|
||||
f"## Repository: {owner}/{repo}\n\n"
|
||||
f"## Pull Request: {pr_url} (#{pr_number})\n\n"
|
||||
f"## Reviewer: {reviewer}\n\n"
|
||||
f"## Feedback:\n{body}\n\n"
|
||||
"Instructions:\n"
|
||||
"1. If the requested change is unambiguous (rename, typo, missing null "
|
||||
"check, small refactor, add a test), make it, commit, and push to the "
|
||||
"existing branch.\n"
|
||||
"2. If the comment is ambiguous, opinion-based, or needs a design "
|
||||
"decision, reply on the PR asking for clarification instead of guessing.\n"
|
||||
"3. Before finishing, re-check the PR's latest review comments and CI "
|
||||
"status. Address any newly arrived actionable comments or failed checks "
|
||||
"that are clear and deterministic.\n"
|
||||
"4. Never force-push. Reply to the reviewer on GitHub to explain what "
|
||||
"you changed."
|
||||
)
|
||||
|
||||
|
||||
async def _thread_autofix_state(metadata: dict[str, Any]) -> tuple[int, list[str], str, str]:
|
||||
attempts = metadata.get("autofix_attempts")
|
||||
attempts = attempts if isinstance(attempts, int) and attempts >= 0 else 0
|
||||
handled = metadata.get("autofix_handled")
|
||||
handled = [h for h in handled if isinstance(h, str)] if isinstance(handled, list) else []
|
||||
github_login = metadata.get("github_login")
|
||||
github_login = github_login if isinstance(github_login, str) else ""
|
||||
user_email = metadata.get("triggering_user_email")
|
||||
user_email = user_email if isinstance(user_email, str) else ""
|
||||
return attempts, handled, github_login, user_email
|
||||
|
||||
|
||||
async def _record_attempt(
|
||||
thread_id: str, *, attempts: int, handled: list[str], dedupe_key: str, head_sha: str
|
||||
) -> None:
|
||||
new_handled = [*handled, dedupe_key][-_MAX_HANDLED_KEYS:]
|
||||
try:
|
||||
await get_client().threads.update(
|
||||
thread_id=thread_id,
|
||||
metadata={
|
||||
"autofix_attempts": attempts + 1,
|
||||
"autofix_handled": new_handled,
|
||||
"autofix_last_head_sha": head_sha,
|
||||
},
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug("Failed to record auto-fix attempt for thread %s", thread_id, exc_info=True)
|
||||
|
||||
|
||||
# Run sources the agent's GitHub-token resolver knows how to authenticate.
|
||||
_AUTH_RESOLVABLE_SOURCES = frozenset(["github", "slack", "dashboard", "linear", "schedule"])
|
||||
|
||||
|
||||
def _run_configurable(
|
||||
metadata: dict[str, Any], *, repo_config: dict[str, str], pr_number: int
|
||||
) -> dict[str, Any]:
|
||||
"""Build the run config for a fix run by reusing the PR thread's identity.
|
||||
|
||||
The agent's GitHub-token resolver only authenticates known sources, so a
|
||||
bespoke ``github_ci`` source would fail in non-bot-token deployments. Reuse
|
||||
the originating thread's ``source`` + login/email so auth resolves exactly
|
||||
as it did for the run that opened the PR.
|
||||
"""
|
||||
source = metadata.get("source")
|
||||
if source not in _AUTH_RESOLVABLE_SOURCES:
|
||||
source = "github"
|
||||
configurable: dict[str, Any] = {
|
||||
"source": source,
|
||||
"repo": repo_config,
|
||||
"pr_number": pr_number,
|
||||
}
|
||||
login = metadata.get("github_login")
|
||||
if isinstance(login, str) and login:
|
||||
configurable["github_login"] = login
|
||||
email = metadata.get("triggering_user_email")
|
||||
if isinstance(email, str) and email:
|
||||
configurable["user_email"] = email
|
||||
return configurable
|
||||
|
||||
|
||||
async def _mark_pending_autofix_event(thread_id: str, reason: str, detail: str = "") -> None:
|
||||
"""Record a batched auto-fix event in the store the message-queue middleware reads.
|
||||
|
||||
Uses the same store namespace mechanism as ``queue_message_for_thread`` so the
|
||||
in-flight run picks it up in-process at its next ``before_model`` step — no
|
||||
per-step thread fetch. ``detail`` (e.g. a reviewer's comment) is accumulated so
|
||||
specifics aren't lost when several events batch against one busy run.
|
||||
"""
|
||||
client = langgraph_client()
|
||||
namespace = (_PENDING_AUTOFIX_NS, thread_id)
|
||||
try:
|
||||
details: list[str] = []
|
||||
try:
|
||||
existing = await client.store.get_item(namespace, _PENDING_AUTOFIX_KEY)
|
||||
if existing and existing.get("value"):
|
||||
prior = existing["value"].get("details")
|
||||
if isinstance(prior, list):
|
||||
details = [d for d in prior if isinstance(d, str)]
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug("No existing pending auto-fix event for thread %s", thread_id)
|
||||
if detail and detail not in details:
|
||||
details.append(detail)
|
||||
await client.store.put_item(
|
||||
namespace, _PENDING_AUTOFIX_KEY, {"reason": reason, "details": details}
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug(
|
||||
"Failed to record pending auto-fix event for thread %s", thread_id, exc_info=True
|
||||
)
|
||||
|
||||
|
||||
async def _dispatch_or_batch(
|
||||
thread_id: str, prompt: str, *, configurable: dict[str, Any], reason: str, detail: str = ""
|
||||
) -> str:
|
||||
if await is_thread_active(thread_id):
|
||||
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"
|
||||
client = langgraph_client()
|
||||
await client.runs.create(
|
||||
thread_id,
|
||||
"agent",
|
||||
input={"messages": [{"role": "user", "content": prompt}]},
|
||||
config={"configurable": configurable},
|
||||
if_not_exists="create",
|
||||
)
|
||||
logger.info(
|
||||
"Created auto-fix run for thread %s (source=%s)", thread_id, configurable.get("source")
|
||||
)
|
||||
return "dispatched"
|
||||
|
||||
|
||||
async def handle_ci_failure(
|
||||
*,
|
||||
repo_config: dict[str, str],
|
||||
branch: str,
|
||||
head_sha: str,
|
||||
token: str | None = None,
|
||||
source: str = "github_ci",
|
||||
failing_checks: list[dict[str, Any]] | None = None,
|
||||
pr: dict[str, Any] | None = None,
|
||||
) -> str:
|
||||
"""Auto-fix failing CI on an agent-authored PR. Returns a status string."""
|
||||
owner = repo_config.get("owner", "")
|
||||
repo = repo_config.get("name", "")
|
||||
if not owner or not repo:
|
||||
return "missing_repo"
|
||||
|
||||
if not await is_review_repo_enabled(owner, repo):
|
||||
return "repo_not_enabled"
|
||||
|
||||
if token is None:
|
||||
token = await get_github_app_installation_token()
|
||||
if not token:
|
||||
logger.warning("No GitHub App token for CI auto-fix on %s/%s", owner, repo)
|
||||
return "no_token"
|
||||
|
||||
if pr is None:
|
||||
if not branch:
|
||||
return "no_branch"
|
||||
pr = await fetch_open_pr_for_branch(owner=owner, repo=repo, branch=branch, token=token)
|
||||
if not pr:
|
||||
return "no_open_pr"
|
||||
|
||||
pr_number = pr.get("number")
|
||||
if not isinstance(pr_number, int):
|
||||
return "no_pr_number"
|
||||
pr_url = pr.get("html_url") or pr.get("url") or ""
|
||||
base_sha = (pr.get("base") or {}).get("sha", "")
|
||||
branch = branch or (pr.get("head") or {}).get("ref", "")
|
||||
head_sha = head_sha or (pr.get("head") or {}).get("sha", "")
|
||||
if not head_sha:
|
||||
return "no_head_sha"
|
||||
|
||||
if await is_pr_autofix_disabled(owner, repo, pr_number):
|
||||
return "pr_disabled"
|
||||
|
||||
found = await find_agent_thread_for_pr(pr_url)
|
||||
if found is None:
|
||||
return "no_agent_thread"
|
||||
thread_id, metadata = found
|
||||
|
||||
attempts, handled, github_login, user_email = await _thread_autofix_state(metadata)
|
||||
|
||||
if not await _user_autofix_enabled(github_login, user_email):
|
||||
return "autofix_disabled_user"
|
||||
|
||||
if attempts >= MAX_AUTOFIX_ATTEMPTS:
|
||||
await post_autofix_status_check(
|
||||
owner=owner,
|
||||
repo=repo,
|
||||
head_sha=head_sha,
|
||||
token=token,
|
||||
title="Auto-fix limit reached",
|
||||
summary=(
|
||||
f"Open SWE has attempted {attempts} auto-fixes on this PR and "
|
||||
"stopped to avoid a loop. Push a commit or comment to continue."
|
||||
),
|
||||
details_url=dashboard_thread_url(thread_id),
|
||||
)
|
||||
return "max_attempts"
|
||||
|
||||
if failing_checks is None:
|
||||
runs = await list_failing_check_runs(owner=owner, repo=repo, ref=head_sha, token=token)
|
||||
statuses = await list_failing_statuses(owner=owner, repo=repo, ref=head_sha, token=token)
|
||||
if runs is None and statuses is None:
|
||||
return "ci_read_failed"
|
||||
failing_checks = (runs or []) + (statuses or [])
|
||||
if not failing_checks:
|
||||
return "no_failing_checks"
|
||||
|
||||
base_failing = await names_failing_on_base(
|
||||
owner=owner, repo=repo, base_sha=base_sha, token=token
|
||||
)
|
||||
actionable = [c for c in failing_checks if c.get("name") not in base_failing]
|
||||
if not actionable:
|
||||
return "all_failing_on_base"
|
||||
|
||||
dedupe_key = _dedupe_key(head_sha)
|
||||
if dedupe_key in handled:
|
||||
return "already_handled"
|
||||
|
||||
author_login = await head_commit_author_login(owner=owner, repo=repo, sha=head_sha, token=token)
|
||||
if (
|
||||
author_login is not None
|
||||
and author_login not in INTERNAL_BOT_LOGINS
|
||||
and (not github_login or author_login.lower() != github_login.lower())
|
||||
):
|
||||
logger.info(
|
||||
"Skipping CI auto-fix on %s/%s#%s: head commit %s authored by human %s",
|
||||
owner,
|
||||
repo,
|
||||
pr_number,
|
||||
head_sha,
|
||||
author_login,
|
||||
)
|
||||
return "human_commit"
|
||||
|
||||
prompt = _build_ci_fix_prompt(
|
||||
owner=owner,
|
||||
repo=repo,
|
||||
pr_number=pr_number,
|
||||
pr_url=pr_url,
|
||||
branch=branch,
|
||||
head_sha=head_sha,
|
||||
failing_checks=actionable,
|
||||
)
|
||||
result = await _dispatch_or_batch(
|
||||
thread_id,
|
||||
prompt,
|
||||
configurable=_run_configurable(
|
||||
metadata, repo_config={"owner": owner, "name": repo}, pr_number=pr_number
|
||||
),
|
||||
reason="ci_failure",
|
||||
)
|
||||
if result == "dispatched":
|
||||
# Only burn an attempt / mark the SHA handled on a real dispatch. A batched
|
||||
# event is just a nudge to the in-flight run; if that run ends before
|
||||
# consuming it, leaving the SHA un-handled lets a later webhook or the sweep
|
||||
# re-dispatch instead of silently dropping the failure.
|
||||
await _record_attempt(
|
||||
thread_id, attempts=attempts, handled=handled, dedupe_key=dedupe_key, head_sha=head_sha
|
||||
)
|
||||
await post_autofix_status_check(
|
||||
owner=owner,
|
||||
repo=repo,
|
||||
head_sha=head_sha,
|
||||
token=token,
|
||||
title=f"Auto-fixing {len(actionable)} failing check(s)",
|
||||
summary=(
|
||||
"Open SWE is investigating the failing checks and will push a fix if "
|
||||
"the cause is clear. Track progress in the linked run."
|
||||
),
|
||||
details_url=dashboard_thread_url(thread_id),
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
async def handle_review_feedback(
|
||||
*,
|
||||
repo_config: dict[str, str],
|
||||
pr_number: int,
|
||||
pr_url: str,
|
||||
reviewer: str,
|
||||
body: str,
|
||||
token: str | None = None,
|
||||
source: str = "github_review",
|
||||
) -> str:
|
||||
"""Auto-respond to a human review comment on an agent-authored PR."""
|
||||
owner = repo_config.get("owner", "")
|
||||
repo = repo_config.get("name", "")
|
||||
if not owner or not repo or not pr_url:
|
||||
return "missing_repo"
|
||||
|
||||
if not await is_review_repo_enabled(owner, repo):
|
||||
return "repo_not_enabled"
|
||||
if await is_pr_autofix_disabled(owner, repo, pr_number):
|
||||
return "pr_disabled"
|
||||
|
||||
found = await find_agent_thread_for_pr(pr_url)
|
||||
if found is None:
|
||||
return "no_agent_thread"
|
||||
thread_id, metadata = found
|
||||
|
||||
_, _, github_login, user_email = await _thread_autofix_state(metadata)
|
||||
if not await _user_autofix_enabled(github_login, user_email):
|
||||
return "autofix_disabled_user"
|
||||
|
||||
if token is None:
|
||||
token = await get_github_app_installation_token()
|
||||
if not token:
|
||||
return "no_token"
|
||||
if not await has_repo_write_permission(owner=owner, repo=repo, username=reviewer, token=token):
|
||||
logger.info(
|
||||
"Skipping auto-fix review feedback on %s/%s#%s: %s lacks write access",
|
||||
owner,
|
||||
repo,
|
||||
pr_number,
|
||||
reviewer or "<unknown>",
|
||||
)
|
||||
return "reviewer_no_write_permission"
|
||||
|
||||
prompt = _build_review_feedback_prompt(
|
||||
owner=owner,
|
||||
repo=repo,
|
||||
pr_number=pr_number,
|
||||
pr_url=pr_url,
|
||||
reviewer=reviewer,
|
||||
body=body,
|
||||
)
|
||||
return await _dispatch_or_batch(
|
||||
thread_id,
|
||||
prompt,
|
||||
configurable=_run_configurable(
|
||||
metadata, repo_config={"owner": owner, "name": repo}, pr_number=pr_number
|
||||
),
|
||||
reason="review_feedback",
|
||||
detail=f"Reviewer {reviewer or 'unknown'} commented: {body.strip()}"
|
||||
if body.strip()
|
||||
else "",
|
||||
)
|
||||
|
||||
|
||||
async def sweep_open_prs() -> dict[str, int]:
|
||||
"""Poll open agent-authored PRs and auto-fix failing CI / flag conflicts.
|
||||
|
||||
The polling fallback for deployments without reliable CI webhooks, and the
|
||||
only path that can react to base-branch merge conflicts (GitHub emits no
|
||||
webhook for those).
|
||||
"""
|
||||
counts = {"scanned": 0, "dispatched": 0, "batched": 0, "conflicts": 0}
|
||||
token = await get_github_app_installation_token()
|
||||
if not token:
|
||||
logger.warning("CI monitor sweep: no GitHub App token")
|
||||
return counts
|
||||
client = get_client()
|
||||
try:
|
||||
threads = await client.threads.search(
|
||||
metadata={"agent_kind": "agent", "pr_state": "open"}, limit=100
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("CI monitor sweep: thread search failed", exc_info=True)
|
||||
return counts
|
||||
|
||||
for thread in threads or []:
|
||||
metadata = thread.get("metadata") if isinstance(thread, dict) else None
|
||||
if not isinstance(metadata, dict):
|
||||
continue
|
||||
repo = metadata.get("repo")
|
||||
pr_number = metadata.get("pr_number")
|
||||
branch = metadata.get("branch_name")
|
||||
if not isinstance(repo, dict) or not isinstance(pr_number, int):
|
||||
continue
|
||||
owner = repo.get("owner", "")
|
||||
name = repo.get("name", "")
|
||||
if not owner or not name:
|
||||
continue
|
||||
counts["scanned"] += 1
|
||||
pr = await fetch_pr(owner=owner, repo=name, pr_number=pr_number, token=token)
|
||||
if not pr:
|
||||
continue
|
||||
head_sha = (pr.get("head") or {}).get("sha", "")
|
||||
branch = (pr.get("head") or {}).get("ref", "") or (
|
||||
branch if isinstance(branch, str) else ""
|
||||
)
|
||||
if pr.get("mergeable_state") == "dirty":
|
||||
counts["conflicts"] += 1
|
||||
await _flag_merge_conflict(
|
||||
owner=owner,
|
||||
repo=name,
|
||||
pr_number=pr_number,
|
||||
pr_url=pr.get("html_url") or "",
|
||||
head_sha=head_sha,
|
||||
token=token,
|
||||
)
|
||||
continue
|
||||
result = await handle_ci_failure(
|
||||
repo_config={"owner": owner, "name": name},
|
||||
branch=branch,
|
||||
head_sha=head_sha,
|
||||
token=token,
|
||||
source="ci_monitor",
|
||||
pr=pr,
|
||||
)
|
||||
if result == "dispatched":
|
||||
counts["dispatched"] += 1
|
||||
elif result == "batched":
|
||||
counts["batched"] += 1
|
||||
logger.info("CI monitor sweep complete: %s", counts)
|
||||
return counts
|
||||
|
||||
|
||||
async def _flag_merge_conflict(
|
||||
*, owner: str, repo: str, pr_number: int, pr_url: str, head_sha: str, token: str
|
||||
) -> None:
|
||||
"""Ask the agent to rebase a PR that has merge conflicts with its base."""
|
||||
if await is_pr_autofix_disabled(owner, repo, pr_number):
|
||||
return
|
||||
found = await find_agent_thread_for_pr(pr_url)
|
||||
if found is None:
|
||||
return
|
||||
thread_id, metadata = found
|
||||
_, _, github_login, user_email = await _thread_autofix_state(metadata)
|
||||
if not await _user_autofix_enabled(github_login, user_email):
|
||||
return
|
||||
if metadata.get("autofix_conflict_head") == head_sha:
|
||||
return
|
||||
prompt = (
|
||||
f"The pull request you opened (#{pr_number}, {pr_url}) now has merge "
|
||||
"conflicts with its base branch. Rebase or merge the base branch into "
|
||||
"the PR branch, resolve the conflicts carefully, and push. If a "
|
||||
"conflict resolution is ambiguous, comment on the PR and ask before "
|
||||
"guessing. Never force-push over commits already on the remote."
|
||||
)
|
||||
await _dispatch_or_batch(
|
||||
thread_id,
|
||||
prompt,
|
||||
configurable=_run_configurable(
|
||||
metadata, repo_config={"owner": owner, "name": repo}, pr_number=pr_number
|
||||
),
|
||||
reason="merge_conflict",
|
||||
)
|
||||
try:
|
||||
await get_client().threads.update(
|
||||
thread_id=thread_id, metadata={"autofix_conflict_head": head_sha}
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug("Failed to record conflict head for thread %s", thread_id, exc_info=True)
|
||||
|
|
@ -1,35 +0,0 @@
|
|||
"""LangGraph entrypoint that polls open agent PRs for CI failures / conflicts.
|
||||
|
||||
A fallback for deployments where CI webhooks (``check_run`` / ``workflow_run``)
|
||||
aren't reliably delivered, and the only path that can react to base-branch
|
||||
merge conflicts (GitHub emits no webhook for those). Register it on a cron to
|
||||
sweep periodically; each tick calls :func:`agent.ci_autofix.sweep_open_prs`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any, TypedDict
|
||||
|
||||
from langgraph.graph import END, START, StateGraph
|
||||
from langgraph.graph.state import RunnableConfig
|
||||
|
||||
from .ci_autofix import sweep_open_prs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CIMonitorState(TypedDict, total=False):
|
||||
result: dict[str, Any]
|
||||
|
||||
|
||||
async def _sweep(_state: CIMonitorState, _config: RunnableConfig) -> dict[str, Any]:
|
||||
return {"result": await sweep_open_prs()}
|
||||
|
||||
|
||||
def get_ci_monitor(config: RunnableConfig | None = None):
|
||||
builder = StateGraph(CIMonitorState)
|
||||
builder.add_node("sweep", _sweep)
|
||||
builder.add_edge(START, "sweep")
|
||||
builder.add_edge("sweep", END)
|
||||
return builder.compile().with_config(config or {})
|
||||
148
agent/completion.py
Normal file
148
agent/completion.py
Normal file
|
|
@ -0,0 +1,148 @@
|
|||
"""Run-completion webhook handler — guarantees every run ends with a signal.
|
||||
|
||||
The platform POSTs a run-completion payload to ``/webhooks/run-complete`` (wired
|
||||
as the ``webhook`` on every dispatched run, see ``agent.dispatch``). When a run
|
||||
ends in a failure state (``error`` / ``timeout`` / ``interrupted``) we post a
|
||||
short failure reply to the originating channel, so a run that died on a server
|
||||
recycle or hit a limit never leaves the user in silence.
|
||||
|
||||
This decouples "the user gets an answer" from "the agent remembered to reply."
|
||||
The reply is idempotent: a per-thread metadata flag prevents double-posting when
|
||||
the platform retries the webhook or a checkpoint replays.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hmac
|
||||
import logging
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from .utils.github_app import get_github_app_installation_token
|
||||
from .utils.github_comments import post_github_comment
|
||||
from .utils.linear import comment_on_linear_issue
|
||||
from .utils.slack import post_slack_thread_reply
|
||||
from .utils.thread_ops import langgraph_client
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Run statuses that mean the user will otherwise get nothing back. "interrupted"
|
||||
# is intentionally excluded: with multitask_strategy="interrupt", a normal
|
||||
# follow-up halts the prior run (status "interrupted") while its replacement
|
||||
# carries on — that's healthy, not a failure worth a "couldn't finish" reply.
|
||||
_TERMINAL_FAILURE_STATUSES = frozenset({"error", "timeout"})
|
||||
_FAILURE_REPLY_FLAG = "failure_reply_posted"
|
||||
|
||||
# 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
|
||||
# every call, so completion replies stay off until the secret is configured.
|
||||
RUN_COMPLETE_WEBHOOK_SECRET = os.environ.get("RUN_COMPLETE_WEBHOOK_SECRET")
|
||||
if not RUN_COMPLETE_WEBHOOK_SECRET:
|
||||
logger.warning(
|
||||
"RUN_COMPLETE_WEBHOOK_SECRET is not set; /webhooks/run-complete is fail-closed "
|
||||
"(all calls rejected) and run-failure replies are disabled. Set it to enable them."
|
||||
)
|
||||
|
||||
|
||||
def verify_run_complete_token(token: str | None) -> bool:
|
||||
"""Return whether a run-completion webhook token is acceptable.
|
||||
|
||||
Fail closed: with no secret configured, reject every call rather than accept
|
||||
unauthenticated requests on a publicly reachable route.
|
||||
"""
|
||||
secret = RUN_COMPLETE_WEBHOOK_SECRET
|
||||
if not secret:
|
||||
return False
|
||||
return token is not None and hmac.compare_digest(token, secret)
|
||||
|
||||
|
||||
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"
|
||||
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."""
|
||||
source = metadata.get("source")
|
||||
ctx = metadata.get("source_context")
|
||||
ctx = ctx if isinstance(ctx, dict) else {}
|
||||
text = _failure_text(status)
|
||||
|
||||
if source == "slack":
|
||||
slack_thread = ctx.get("slack_thread")
|
||||
if isinstance(slack_thread, dict):
|
||||
channel_id = slack_thread.get("channel_id")
|
||||
thread_ts = slack_thread.get("thread_ts")
|
||||
if channel_id and thread_ts:
|
||||
return await post_slack_thread_reply(channel_id, thread_ts, text)
|
||||
return False
|
||||
|
||||
if source == "linear":
|
||||
linear_issue = ctx.get("linear_issue")
|
||||
if isinstance(linear_issue, dict):
|
||||
issue_id = linear_issue.get("id")
|
||||
if issue_id:
|
||||
return await comment_on_linear_issue(issue_id, text)
|
||||
return False
|
||||
|
||||
if source in ("github", "github_issue"):
|
||||
repo_config = metadata.get("repo")
|
||||
number = ctx.get("pr_number")
|
||||
if number is None:
|
||||
github_issue = ctx.get("github_issue")
|
||||
if isinstance(github_issue, dict):
|
||||
number = github_issue.get("number")
|
||||
if isinstance(repo_config, dict) and isinstance(number, int):
|
||||
token = await get_github_app_installation_token()
|
||||
if token:
|
||||
return await post_github_comment(repo_config, number, text, token=token)
|
||||
return False
|
||||
|
||||
logger.info("No failure-reply channel for thread %s (source=%s)", thread_id, source)
|
||||
return False
|
||||
|
||||
|
||||
async def handle_run_completion(payload: dict[str, Any]) -> dict[str, str]:
|
||||
"""Handle a platform run-completion webhook POST.
|
||||
|
||||
Posts a failure reply only when the run ended in a failure state and we
|
||||
haven't already replied for this thread.
|
||||
"""
|
||||
status = payload.get("status")
|
||||
thread_id = payload.get("thread_id")
|
||||
if not isinstance(thread_id, str) or not thread_id:
|
||||
return {"status": "ignored", "reason": "missing thread_id"}
|
||||
if status not in _TERMINAL_FAILURE_STATUSES:
|
||||
return {"status": "ignored", "reason": f"non-failure status: {status}"}
|
||||
|
||||
client = langgraph_client()
|
||||
try:
|
||||
thread = await client.threads.get(thread_id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("run-complete: could not load thread %s", thread_id, exc_info=True)
|
||||
return {"status": "error", "reason": "thread fetch failed"}
|
||||
|
||||
metadata = thread.get("metadata") if isinstance(thread, dict) else None
|
||||
metadata = metadata if isinstance(metadata, dict) else {}
|
||||
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"}
|
||||
|
||||
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)
|
||||
logger.info("Posted failure reply for thread %s (status=%s)", thread_id, status)
|
||||
return {"status": "ok", "reason": "failure reply posted"}
|
||||
|
|
@ -21,6 +21,7 @@ from fastapi import APIRouter, Depends, HTTPException
|
|||
from langgraph_sdk import get_client
|
||||
from pydantic import BaseModel
|
||||
|
||||
from ..dispatch import dispatch_agent_run
|
||||
from .oauth import require_same_origin_for_mutations, require_session
|
||||
from .plan_store import (
|
||||
PLAN_STATUS_APPROVED,
|
||||
|
|
@ -212,11 +213,9 @@ async def _dispatch_followup(
|
|||
# mode (implement), reject stays in plan mode (revise the plan).
|
||||
configurable["plan_mode"] = plan_mode
|
||||
|
||||
client = get_client()
|
||||
await client.runs.create(
|
||||
await dispatch_agent_run(
|
||||
thread_id,
|
||||
"agent",
|
||||
input={"messages": [{"role": "user", "content": text}]},
|
||||
config={"configurable": configurable},
|
||||
if_not_exists="create",
|
||||
text,
|
||||
configurable,
|
||||
source=configurable["source"],
|
||||
)
|
||||
|
|
|
|||
90
agent/dispatch.py
Normal file
90
agent/dispatch.py
Normal file
|
|
@ -0,0 +1,90 @@
|
|||
"""Single durable dispatch contract behind every agent/reviewer run trigger.
|
||||
|
||||
Replaces the per-site ``runs.create`` calls (plus the ``is_thread_active``
|
||||
busy-check and the custom store-queue) with one function that always uses:
|
||||
|
||||
- ``multitask_strategy="interrupt"`` — a follow-up halts the active run
|
||||
(progress preserved by the sync checkpoint) and resumes the agent with full
|
||||
history + the new message; on an idle thread it just starts. This is the
|
||||
platform-native, cross-process replacement for the racy busy-check + queue.
|
||||
- ``durability="sync"`` — checkpoint before each step so a crash/recycle
|
||||
resumes from the last checkpoint instead of losing all work.
|
||||
- ``webhook=COMPLETION_WEBHOOK_URL`` — the platform calls us on completion or
|
||||
failure so every run ends with a signal even if the agent died.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
from langgraph_sdk import get_client
|
||||
from langgraph_sdk.client import LangGraphClient
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
ContentBlocks = str | list[dict[str, Any]]
|
||||
|
||||
# Same-server FastAPI route the platform POSTs run completion/failure to. A
|
||||
# relative URL loopback-posts into this app (no SSRF/loopback config needed);
|
||||
# override with an absolute URL via env for split deployments. The route is
|
||||
# fail-closed on RUN_COMPLETE_WEBHOOK_SECRET, so only register the webhook when
|
||||
# the secret is set, appending it as ?token= so the route can verify the call
|
||||
# came from us (completion.verify_run_complete_token). Unset → no webhook.
|
||||
_COMPLETION_WEBHOOK_BASE = os.environ.get("COMPLETION_WEBHOOK_URL") or "/webhooks/run-complete"
|
||||
_RUN_COMPLETE_SECRET = os.environ.get("RUN_COMPLETE_WEBHOOK_SECRET")
|
||||
COMPLETION_WEBHOOK_URL: str | None
|
||||
if not _RUN_COMPLETE_SECRET:
|
||||
COMPLETION_WEBHOOK_URL = None
|
||||
elif "?" in _COMPLETION_WEBHOOK_BASE:
|
||||
COMPLETION_WEBHOOK_URL = _COMPLETION_WEBHOOK_BASE
|
||||
else:
|
||||
COMPLETION_WEBHOOK_URL = f"{_COMPLETION_WEBHOOK_BASE}?token={_RUN_COMPLETE_SECRET}"
|
||||
|
||||
|
||||
def _langgraph_url() -> str:
|
||||
return os.environ.get("LANGGRAPH_URL") or os.environ.get(
|
||||
"LANGGRAPH_URL_PROD", "http://localhost:2024"
|
||||
)
|
||||
|
||||
|
||||
def dispatch_client() -> LangGraphClient:
|
||||
return get_client(url=_langgraph_url())
|
||||
|
||||
|
||||
async def dispatch_agent_run(
|
||||
thread_id: str,
|
||||
content: ContentBlocks,
|
||||
configurable: dict[str, Any],
|
||||
*,
|
||||
source: str,
|
||||
assistant_id: str = "agent",
|
||||
metadata: dict[str, Any] | None = None,
|
||||
client: LangGraphClient | None = None,
|
||||
) -> 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"``).
|
||||
"""
|
||||
client = client or dispatch_client()
|
||||
run = await client.runs.create(
|
||||
thread_id,
|
||||
assistant_id,
|
||||
input={"messages": [{"role": "user", "content": content}]},
|
||||
config={"configurable": configurable, "metadata": metadata or {}},
|
||||
multitask_strategy="interrupt",
|
||||
durability="sync",
|
||||
webhook=COMPLETION_WEBHOOK_URL,
|
||||
if_not_exists="create",
|
||||
)
|
||||
logger.info(
|
||||
"Dispatched %s run on thread %s (source=%s, run=%s)",
|
||||
assistant_id,
|
||||
thread_id,
|
||||
source,
|
||||
run.get("run_id") if isinstance(run, dict) else None,
|
||||
)
|
||||
return run
|
||||
121
agent/reconcile.py
Normal file
121
agent/reconcile.py
Normal file
|
|
@ -0,0 +1,121 @@
|
|||
"""Reconciliation sweep: cancel runs stuck in ``pending`` past their deadline.
|
||||
|
||||
The durable-dispatch contract relies on the platform's completion webhook to
|
||||
end every run. When that webhook never fires (crash, lost delivery), a run can
|
||||
sit in ``pending`` forever and hold its thread ``busy``. This sweep is the
|
||||
safety net: find busy threads, look for stale ``pending`` runs on them, and
|
||||
cancel the ones older than ``max_age_seconds`` so the thread frees up.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from .utils.thread_ops import langgraph_client
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SEARCH_PAGE_SIZE = 100
|
||||
|
||||
|
||||
def _parse_created_at(value: Any) -> datetime | None:
|
||||
"""Parse a run's ``created_at`` into an aware UTC datetime, or None."""
|
||||
if isinstance(value, datetime):
|
||||
return value if value.tzinfo else value.replace(tzinfo=UTC)
|
||||
if not isinstance(value, str) or not value:
|
||||
return None
|
||||
text = value.strip()
|
||||
if text.endswith("Z"):
|
||||
text = f"{text[:-1]}+00:00"
|
||||
try:
|
||||
parsed = datetime.fromisoformat(text)
|
||||
except ValueError:
|
||||
return None
|
||||
return parsed if parsed.tzinfo else parsed.replace(tzinfo=UTC)
|
||||
|
||||
|
||||
async def reconcile_stale_runs(*, max_age_seconds: int = 1800) -> dict[str, int]:
|
||||
"""Cancel ``pending`` runs older than ``max_age_seconds`` on busy threads.
|
||||
|
||||
Walks every ``busy`` thread (paginated), lists its ``pending`` runs, and
|
||||
cancels those whose ``created_at`` is older than the cutoff. Per-thread work
|
||||
is wrapped in try/except so one bad thread never aborts the sweep.
|
||||
|
||||
Returns counts: ``{"threads_checked", "stale_runs", "cancelled"}``.
|
||||
"""
|
||||
client = langgraph_client()
|
||||
now = datetime.now(UTC)
|
||||
|
||||
threads_checked = 0
|
||||
stale_runs = 0
|
||||
cancelled = 0
|
||||
|
||||
offset = 0
|
||||
while True:
|
||||
try:
|
||||
threads = await client.threads.search(
|
||||
metadata=None,
|
||||
status="busy",
|
||||
limit=_SEARCH_PAGE_SIZE,
|
||||
offset=offset,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Reconcile sweep: thread search failed at offset %d", offset)
|
||||
break
|
||||
if not threads:
|
||||
break
|
||||
|
||||
for thread in threads:
|
||||
thread_id = thread.get("thread_id") if isinstance(thread, dict) else None
|
||||
if not thread_id:
|
||||
continue
|
||||
threads_checked += 1
|
||||
try:
|
||||
runs = await client.runs.list(thread_id, status="pending")
|
||||
stale_run_ids: list[str] = []
|
||||
for run in runs:
|
||||
created = _parse_created_at(run.get("created_at"))
|
||||
if created is None:
|
||||
logger.warning(
|
||||
"Reconcile sweep: unparseable created_at on run %s (thread %s)",
|
||||
run.get("run_id"),
|
||||
thread_id,
|
||||
)
|
||||
continue
|
||||
if (now - created).total_seconds() <= max_age_seconds:
|
||||
continue
|
||||
run_id = run.get("run_id")
|
||||
if run_id:
|
||||
stale_run_ids.append(run_id)
|
||||
|
||||
if not stale_run_ids:
|
||||
continue
|
||||
stale_runs += len(stale_run_ids)
|
||||
await client.runs.cancel_many(
|
||||
thread_id=thread_id,
|
||||
run_ids=stale_run_ids,
|
||||
action="interrupt",
|
||||
)
|
||||
cancelled += len(stale_run_ids)
|
||||
logger.info(
|
||||
"Reconcile sweep: cancelled %d stale pending run(s) on thread %s",
|
||||
len(stale_run_ids),
|
||||
thread_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Reconcile sweep: failed to reconcile thread %s", thread_id)
|
||||
continue
|
||||
|
||||
if len(threads) < _SEARCH_PAGE_SIZE:
|
||||
break
|
||||
offset += _SEARCH_PAGE_SIZE
|
||||
|
||||
counts = {
|
||||
"threads_checked": threads_checked,
|
||||
"stale_runs": stale_runs,
|
||||
"cancelled": cancelled,
|
||||
}
|
||||
logger.info("Reconcile sweep complete: %s", counts)
|
||||
return counts
|
||||
|
|
@ -9,17 +9,22 @@ from langgraph.graph import END, START, StateGraph
|
|||
from langgraph.graph.state import RunnableConfig
|
||||
|
||||
from .dashboard.schedules import launch_scheduled_agent_run
|
||||
from .reconcile import reconcile_stale_runs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SchedulerState(TypedDict, total=False):
|
||||
schedule_id: str
|
||||
task: str
|
||||
result: dict[str, Any]
|
||||
|
||||
|
||||
async def _launch(state: SchedulerState, config: RunnableConfig) -> dict[str, Any]:
|
||||
configurable = config.get("configurable") or {}
|
||||
task = state.get("task") or configurable.get("task")
|
||||
if task == "reconcile":
|
||||
return {"result": await reconcile_stale_runs()}
|
||||
schedule_id = state.get("schedule_id") or configurable.get("schedule_id")
|
||||
if not isinstance(schedule_id, str) or not schedule_id:
|
||||
logger.warning("Scheduled agent tick missing schedule_id")
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Annotated
|
||||
|
||||
|
|
@ -23,7 +22,7 @@ _ENTERED_MESSAGE = (
|
|||
)
|
||||
|
||||
|
||||
def enter_plan_mode(tool_call_id: Annotated[str, InjectedToolCallId]) -> Command:
|
||||
async def enter_plan_mode(tool_call_id: Annotated[str, InjectedToolCallId]) -> Command:
|
||||
"""Activate plan mode mid-run.
|
||||
|
||||
Call this when you believe the task would benefit from a structured
|
||||
|
|
@ -41,7 +40,7 @@ def enter_plan_mode(tool_call_id: Annotated[str, InjectedToolCallId]) -> Command
|
|||
thread_id = _thread_id_from_config()
|
||||
if thread_id:
|
||||
try:
|
||||
asyncio.run(set_plan_status(thread_id, PLAN_STATUS_PLANNING, plan_mode=True))
|
||||
await set_plan_status(thread_id, PLAN_STATUS_PLANNING, plan_mode=True)
|
||||
except Exception:
|
||||
logger.warning("Failed to persist plan-mode entry for %s", thread_id, exc_info=True)
|
||||
return Command(
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
import asyncio
|
||||
from typing import Any
|
||||
|
||||
from langgraph.config import get_config
|
||||
|
|
@ -7,7 +6,7 @@ from agent.utils.slack import parse_github_pr_url
|
|||
from agent.webapp import trigger_pr_review_from_ref
|
||||
|
||||
|
||||
def request_pr_review(pr_url: str) -> dict[str, Any]:
|
||||
async def request_pr_review(pr_url: str) -> dict[str, Any]:
|
||||
"""Start the reviewer agent for a GitHub pull request URL."""
|
||||
pr_ref = parse_github_pr_url(pr_url)
|
||||
if not pr_ref:
|
||||
|
|
@ -19,13 +18,11 @@ def request_pr_review(pr_url: str) -> dict[str, Any]:
|
|||
configurable = get_config().get("configurable", {})
|
||||
source = configurable.get("source") or "agent"
|
||||
slack_thread = configurable.get("slack_thread") or {}
|
||||
return asyncio.run(
|
||||
trigger_pr_review_from_ref(
|
||||
pr_ref,
|
||||
source=source,
|
||||
github_login=configurable.get("github_login", ""),
|
||||
github_user_id=configurable.get("github_user_id"),
|
||||
slack_channel_id=slack_thread.get("channel_id", ""),
|
||||
slack_thread_ts=slack_thread.get("thread_ts", ""),
|
||||
)
|
||||
return await trigger_pr_review_from_ref(
|
||||
pr_ref,
|
||||
source=source,
|
||||
github_login=configurable.get("github_login", ""),
|
||||
github_user_id=configurable.get("github_user_id"),
|
||||
slack_channel_id=slack_thread.get("channel_id", ""),
|
||||
slack_thread_ts=slack_thread.get("thread_ts", ""),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8,7 +8,6 @@ changes. Available in plan mode (it does not modify the repository under review)
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
|
|
@ -21,7 +20,7 @@ logger = logging.getLogger(__name__)
|
|||
PLAN_FILE_PATH = "plan.md"
|
||||
|
||||
|
||||
def save_plan(plan_markdown: str) -> dict[str, Any]:
|
||||
async def save_plan(plan_markdown: str) -> dict[str, Any]:
|
||||
"""Write your implementation plan as a markdown file and publish it for review.
|
||||
|
||||
Use this in plan mode once your plan is ready. The plan is saved as
|
||||
|
|
@ -54,7 +53,7 @@ def save_plan(plan_markdown: str) -> dict[str, Any]:
|
|||
return {"success": False, "error": "no thread_id in run config"}
|
||||
|
||||
try:
|
||||
path = asyncio.run(_save(str(thread_id), content))
|
||||
path = await _save(str(thread_id), content)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.exception("save_plan failed for thread %s", thread_id)
|
||||
return {"success": False, "error": f"failed to save plan: {exc}"}
|
||||
|
|
|
|||
|
|
@ -1,12 +1,16 @@
|
|||
"""Shared LangGraph thread helpers for webhooks and the dashboard."""
|
||||
"""Shared LangGraph thread helpers for the dashboard.
|
||||
|
||||
The webhook triggers (Slack / Linear / GitHub) dispatch through
|
||||
``agent.dispatch.dispatch_agent_run`` with ``multitask_strategy="interrupt"``,
|
||||
so they no longer need a busy-check or an in-process lock. The store-queue
|
||||
below is retained for the dashboard's deliberate "inject a follow-up into a
|
||||
run that's already in flight" path (``thread_api.send_dashboard_message``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Any
|
||||
|
||||
from langgraph_sdk import get_client
|
||||
|
|
@ -15,25 +19,6 @@ logger = logging.getLogger(__name__)
|
|||
|
||||
MAX_QUEUED_MESSAGES = 100
|
||||
|
||||
_THREAD_RUN_LOCKS: dict[str, asyncio.Lock] = {}
|
||||
|
||||
|
||||
def get_thread_run_lock(thread_id: str) -> asyncio.Lock:
|
||||
"""Return a per-thread-id asyncio.Lock, creating one lazily if needed."""
|
||||
lock = _THREAD_RUN_LOCKS.get(thread_id)
|
||||
if lock is None:
|
||||
lock = asyncio.Lock()
|
||||
_THREAD_RUN_LOCKS[thread_id] = lock
|
||||
return lock
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def thread_run_lock(thread_id: str) -> AsyncIterator[None]:
|
||||
"""Serialize run dispatch for a thread."""
|
||||
lock = get_thread_run_lock(thread_id)
|
||||
async with lock:
|
||||
yield
|
||||
|
||||
|
||||
def langgraph_url() -> str:
|
||||
return os.environ.get("LANGGRAPH_URL") or os.environ.get(
|
||||
|
|
@ -57,15 +42,14 @@ async def get_thread_active_status(thread_id: str) -> bool | None:
|
|||
return None
|
||||
|
||||
|
||||
async def is_thread_active(thread_id: str) -> bool:
|
||||
"""Return whether the thread currently has a running run."""
|
||||
return await get_thread_active_status(thread_id) is True
|
||||
|
||||
|
||||
async def queue_message_for_thread(
|
||||
thread_id: str, message_content: str | list[dict[str, Any]] | dict[str, Any]
|
||||
) -> bool:
|
||||
"""Queue a follow-up message for a busy thread (FIFO store namespace)."""
|
||||
"""Queue a follow-up message for a busy thread (FIFO store namespace).
|
||||
|
||||
Used by the dashboard to inject a follow-up into a run that's already in
|
||||
flight; webhook triggers use ``multitask_strategy="interrupt"`` instead.
|
||||
"""
|
||||
client = langgraph_client()
|
||||
try:
|
||||
namespace = ("queue", thread_id)
|
||||
|
|
|
|||
2048
agent/webapp.py
2048
agent/webapp.py
File diff suppressed because it is too large
Load diff
0
agent/webhooks/__init__.py
Normal file
0
agent/webhooks/__init__.py
Normal file
1013
agent/webhooks/github.py
Normal file
1013
agent/webhooks/github.py
Normal file
File diff suppressed because it is too large
Load diff
235
agent/webhooks/linear.py
Normal file
235
agent/webhooks/linear.py
Normal file
|
|
@ -0,0 +1,235 @@
|
|||
"""Linear webhook handler — moved out of webapp.py (behavior-identical).
|
||||
|
||||
Helpers and constants stay in webapp.py; they are accessed through the module
|
||||
object (``webapp.X``) so tests that monkeypatch them keep working.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from langchain_core.messages.content import create_text_block
|
||||
|
||||
from agent import webapp
|
||||
|
||||
|
||||
async def process_linear_issue( # noqa: PLR0912, PLR0915
|
||||
issue_data: dict[str, Any], repo_config: dict[str, str]
|
||||
) -> None:
|
||||
"""Process a Linear issue by creating a new LangGraph thread and run.
|
||||
|
||||
Args:
|
||||
issue_data: The Linear issue data from webhook (basic info only).
|
||||
repo_config: The repo configuration with owner and name.
|
||||
"""
|
||||
issue_id = issue_data.get("id", "")
|
||||
webapp.logger.info(
|
||||
"Processing Linear issue %s for repo %s/%s",
|
||||
issue_id,
|
||||
repo_config.get("owner"),
|
||||
repo_config.get("name"),
|
||||
)
|
||||
|
||||
triggering_comment_id = issue_data.get("triggering_comment_id", "")
|
||||
if triggering_comment_id:
|
||||
await webapp.react_to_linear_comment(triggering_comment_id, "👀")
|
||||
|
||||
thread_id = webapp.generate_thread_id_from_issue(issue_id)
|
||||
|
||||
full_issue = await webapp.fetch_linear_issue_details(issue_id)
|
||||
if not full_issue:
|
||||
full_issue = issue_data
|
||||
|
||||
user_email = None
|
||||
user_name = None
|
||||
comment_author = issue_data.get("comment_author", {})
|
||||
if comment_author:
|
||||
user_email = comment_author.get("email")
|
||||
user_name = comment_author.get("name")
|
||||
if not user_email:
|
||||
creator = full_issue.get("creator", {})
|
||||
if creator:
|
||||
user_email = creator.get("email")
|
||||
user_name = user_name or creator.get("name")
|
||||
if not user_email:
|
||||
assignee = full_issue.get("assignee", {})
|
||||
if assignee:
|
||||
user_email = assignee.get("email")
|
||||
user_name = user_name or assignee.get("name")
|
||||
|
||||
webapp.logger.info("User email for issue %s: %s", issue_id, user_email)
|
||||
|
||||
title = full_issue.get("title", "No title")
|
||||
description = full_issue.get("description") or "No description"
|
||||
image_urls: list[str] = []
|
||||
description_image_urls = webapp.extract_image_urls(description)
|
||||
if description_image_urls:
|
||||
image_urls.extend(description_image_urls)
|
||||
webapp.logger.debug(
|
||||
"Found %d image URL(s) in issue description",
|
||||
len(description_image_urls),
|
||||
)
|
||||
|
||||
comments = full_issue.get("comments", {}).get("nodes", [])
|
||||
comments_text = ""
|
||||
triggering_comment = issue_data.get("triggering_comment", "")
|
||||
triggering_comment_id = issue_data.get("triggering_comment_id", "")
|
||||
|
||||
bot_message_prefixes = (
|
||||
"🔐 **GitHub Authentication Required**",
|
||||
"✅ **Pull Request Created**",
|
||||
"✅ **Pull Request Updated**",
|
||||
"**Pull Request Created**",
|
||||
"**Pull Request Updated**",
|
||||
"🤖 **Agent Response**",
|
||||
"❌ **Agent Error**",
|
||||
)
|
||||
|
||||
comment_ids: set[str] = set()
|
||||
comment_id_to_index: dict[str, int] = {}
|
||||
if comments:
|
||||
for i, comment in enumerate(comments):
|
||||
comment_id = comment.get("id", "")
|
||||
if comment_id:
|
||||
comment_ids.add(comment_id)
|
||||
comment_id_to_index[comment_id] = i
|
||||
|
||||
relevant_comments = []
|
||||
trigger_index = None
|
||||
if triggering_comment_id:
|
||||
trigger_index = comment_id_to_index.get(triggering_comment_id)
|
||||
if trigger_index is not None:
|
||||
relevant_comments = comments[trigger_index:]
|
||||
webapp.logger.debug(
|
||||
"Using triggering comment index %d to build relevant comments",
|
||||
trigger_index,
|
||||
)
|
||||
else:
|
||||
relevant_comments = webapp.get_recent_comments(comments, bot_message_prefixes)
|
||||
|
||||
if relevant_comments:
|
||||
comments_text = "\n\n## Comments:\n"
|
||||
for comment in relevant_comments:
|
||||
user = comment.get("user") or {}
|
||||
author = user.get("name", "User")
|
||||
body = comment.get("body", "")
|
||||
body_image_urls = webapp.extract_image_urls(body)
|
||||
if body_image_urls:
|
||||
image_urls.extend(body_image_urls)
|
||||
webapp.logger.debug(
|
||||
"Found %d image URL(s) in comment by %s",
|
||||
len(body_image_urls),
|
||||
author,
|
||||
)
|
||||
if any(body.startswith(prefix) for prefix in bot_message_prefixes):
|
||||
continue
|
||||
comments_text += f"\n**{author}:** {body}\n"
|
||||
|
||||
if triggering_comment and triggering_comment_id not in comment_ids:
|
||||
if not comments_text:
|
||||
comments_text = "\n\n## Comments:\n"
|
||||
trigger_author = comment_author.get("name", "Unknown")
|
||||
trigger_body = triggering_comment
|
||||
trigger_image_urls = webapp.extract_image_urls(trigger_body)
|
||||
if trigger_image_urls:
|
||||
image_urls.extend(trigger_image_urls)
|
||||
webapp.logger.debug(
|
||||
"Found %d image URL(s) in triggering comment by %s",
|
||||
len(trigger_image_urls),
|
||||
trigger_author,
|
||||
)
|
||||
comments_text += f"\n**{trigger_author}:** {trigger_body}\n"
|
||||
webapp.logger.debug(
|
||||
"Appended triggering comment %s not present in issue comments list",
|
||||
triggering_comment_id or "<missing-id>",
|
||||
)
|
||||
|
||||
identifier = full_issue.get("identifier", "") or issue_data.get("identifier", "")
|
||||
|
||||
triggered_by_line = f"## Triggered by: {user_name}\n\n" if user_name else ""
|
||||
tag_instruction = (
|
||||
f"When calling linear_comment, tag @{user_name} if you are asking them a question, need their input, or are notifying them of something important (e.g. a completed PR). For simple answers, tagging is not required."
|
||||
if user_name
|
||||
else ""
|
||||
)
|
||||
prompt = (
|
||||
f"Please work on the following issue:\n\n"
|
||||
f"## Repository: {repo_config.get('owner')}/{repo_config.get('name')}\n\n"
|
||||
f"## Title: {title}\n\n"
|
||||
f"{triggered_by_line}"
|
||||
f"## Linear Ticket: {identifier} - Ticket ID: {issue_id}\n\n"
|
||||
f"## Description:\n{description}\n"
|
||||
f"{comments_text}\n\n"
|
||||
f"Please analyze this issue and implement the necessary changes. "
|
||||
f"When you're done, commit and push your changes. {tag_instruction}"
|
||||
)
|
||||
content_blocks: list[dict[str, Any]] = [create_text_block(prompt)]
|
||||
if image_urls:
|
||||
image_urls = webapp.dedupe_urls(image_urls)
|
||||
linear_login = (
|
||||
await webapp.resolve_login_from_email_async(user_email) if user_email else None
|
||||
)
|
||||
resolved_model_id = await webapp.resolve_agent_model_id(linear_login)
|
||||
if webapp.model_supports_images(resolved_model_id):
|
||||
webapp.logger.info("Preparing %d image(s) for multimodal content", len(image_urls))
|
||||
webapp.logger.debug("Image URLs: %s", image_urls)
|
||||
|
||||
async with httpx.AsyncClient(timeout=webapp.DEFAULT_HTTP_TIMEOUT) as client:
|
||||
for image_url in image_urls:
|
||||
image_block = await webapp.fetch_image_block(image_url, client)
|
||||
if image_block:
|
||||
content_blocks.append(image_block)
|
||||
webapp.logger.info("Built %d content block(s) for prompt", len(content_blocks))
|
||||
else:
|
||||
webapp.logger.warning(
|
||||
"Skipping %d image(s) for Linear issue: model %s does not support images",
|
||||
len(image_urls),
|
||||
resolved_model_id,
|
||||
)
|
||||
prompt += webapp.vision_not_supported_warning(resolved_model_id, len(image_urls))
|
||||
content_blocks[0] = create_text_block(prompt)
|
||||
image_urls = []
|
||||
|
||||
linear_project_id = ""
|
||||
linear_issue_number = ""
|
||||
if identifier and "-" in identifier:
|
||||
parts = identifier.split("-", 1)
|
||||
linear_project_id = parts[0]
|
||||
linear_issue_number = parts[1]
|
||||
|
||||
configurable: dict[str, Any] = {
|
||||
"repo": repo_config,
|
||||
"linear_issue": {
|
||||
"id": issue_id,
|
||||
"title": title,
|
||||
"url": full_issue.get("url", "") or issue_data.get("url", ""),
|
||||
"identifier": identifier,
|
||||
"linear_project_id": linear_project_id,
|
||||
"linear_issue_number": linear_issue_number,
|
||||
"triggering_user_name": user_name or "",
|
||||
},
|
||||
"user_email": user_email,
|
||||
"source": "linear",
|
||||
}
|
||||
|
||||
await webapp.upsert_agent_thread_owner_metadata(
|
||||
thread_id,
|
||||
source="linear",
|
||||
repo_config=repo_config,
|
||||
user_email=user_email or "",
|
||||
title=title or identifier or "Linear issue",
|
||||
source_context={"linear_issue": configurable["linear_issue"]},
|
||||
)
|
||||
|
||||
run = await webapp.dispatch_agent_run(
|
||||
thread_id,
|
||||
content_blocks,
|
||||
configurable,
|
||||
source="linear",
|
||||
metadata=webapp._AGENT_VERSION_METADATA,
|
||||
)
|
||||
webapp.logger.info(
|
||||
"LangGraph run dispatched for thread %s (run=%s)",
|
||||
thread_id,
|
||||
run.get("run_id") if isinstance(run, dict) else None,
|
||||
)
|
||||
await webapp.post_linear_trace_comment(issue_id, thread_id, triggering_comment_id)
|
||||
269
agent/webhooks/slack.py
Normal file
269
agent/webhooks/slack.py
Normal file
|
|
@ -0,0 +1,269 @@
|
|||
"""Slack webhook handler — moved out of webapp.py (behavior-identical).
|
||||
|
||||
Helpers and constants stay in webapp.py; they are accessed through the module
|
||||
object (``webapp.X``) so tests that monkeypatch them keep working.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from langchain_core.messages.content import create_text_block
|
||||
|
||||
from agent import webapp
|
||||
|
||||
|
||||
async def process_slack_mention(event_data: dict[str, Any], repo_config: dict[str, str]) -> None:
|
||||
"""Process a Slack app mention by creating a run or queuing a mid-run message."""
|
||||
channel_id = event_data.get("channel_id", "")
|
||||
thread_ts = event_data.get("thread_ts", "")
|
||||
event_ts = event_data.get("event_ts", "")
|
||||
user_id = event_data.get("user_id", "")
|
||||
text = event_data.get("text", "")
|
||||
bot_user_id = event_data.get("bot_user_id", "")
|
||||
|
||||
if not channel_id or not thread_ts or not event_ts:
|
||||
webapp.logger.warning(
|
||||
"Missing Slack event fields (channel_id=%s, thread_ts=%s, event_ts=%s)",
|
||||
channel_id,
|
||||
thread_ts,
|
||||
event_ts,
|
||||
)
|
||||
return
|
||||
|
||||
await webapp.set_slack_assistant_status(channel_id, thread_ts)
|
||||
|
||||
thread_id = webapp.generate_thread_id_from_slack_thread(channel_id, thread_ts)
|
||||
|
||||
# Prime the user-mapping cache so login/email/slack-id lookups below are warm.
|
||||
try:
|
||||
await webapp.refresh_user_mapping_cache()
|
||||
except Exception: # noqa: BLE001
|
||||
webapp.logger.debug("Could not refresh user mapping cache for Slack mention", exc_info=True)
|
||||
|
||||
user_email = None
|
||||
user_name = ""
|
||||
if user_id:
|
||||
slack_user = await webapp.get_slack_user_info(user_id)
|
||||
if slack_user:
|
||||
profile = slack_user.get("profile", {})
|
||||
if isinstance(profile, dict):
|
||||
user_email = profile.get("email")
|
||||
user_name = (
|
||||
profile.get("display_name")
|
||||
or profile.get("real_name")
|
||||
or slack_user.get("real_name")
|
||||
or slack_user.get("name")
|
||||
or ""
|
||||
)
|
||||
|
||||
thread_messages = await webapp.fetch_slack_thread_messages(channel_id, thread_ts)
|
||||
if not any(str(message.get("ts")) == str(event_ts) for message in thread_messages):
|
||||
thread_messages.append({"ts": event_ts, "text": text, "user": user_id})
|
||||
|
||||
context_messages, context_mode = webapp.select_slack_context_messages(
|
||||
thread_messages, event_ts, bot_user_id, webapp.SLACK_BOT_USERNAME
|
||||
)
|
||||
context_user_ids = [
|
||||
value
|
||||
for value in (message.get("user") for message in context_messages)
|
||||
if isinstance(value, str) and value
|
||||
]
|
||||
user_names_by_id = await webapp.get_slack_user_names(context_user_ids)
|
||||
if user_id and user_name and user_id not in user_names_by_id:
|
||||
user_names_by_id[user_id] = user_name
|
||||
context_text = webapp.format_slack_messages_for_prompt(
|
||||
context_messages,
|
||||
user_names_by_id,
|
||||
bot_user_id=bot_user_id,
|
||||
bot_username=webapp.SLACK_BOT_USERNAME,
|
||||
)
|
||||
context_source = (
|
||||
"the previous message where I was tagged"
|
||||
if context_mode == "last_mention"
|
||||
else "the beginning of the thread"
|
||||
)
|
||||
clean_text = (
|
||||
webapp.strip_bot_mention(text, bot_user_id, bot_username=webapp.SLACK_BOT_USERNAME)
|
||||
or "(no text in mention)"
|
||||
)
|
||||
trigger_user = user_name or (f"<@{user_id}>" if user_id else "Unknown user")
|
||||
|
||||
# Auto-resolve cross-posted Slack message links in context
|
||||
resolved_links_section, image_urls_from_links = await webapp.resolve_slack_links_in_context(
|
||||
context_messages, user_names_by_id
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"You were mentioned in Slack.\n\n"
|
||||
"## Default Repository Hint\n"
|
||||
f"{repo_config.get('owner')}/{repo_config.get('name')}\n"
|
||||
"Use this only if the Slack conversation does not identify a different repository.\n\n"
|
||||
f"## Triggered by\n{trigger_user}\n\n"
|
||||
f"## Slack Thread\n- Channel: {channel_id}\n- Thread TS: {thread_ts}\n"
|
||||
f"- Context starts at: {context_source}\n\n"
|
||||
f"## Conversation Context\n{context_text}\n\n"
|
||||
f"## Latest Mention Request\n{clean_text}\n\n"
|
||||
+ (f"{resolved_links_section}\n\n" if resolved_links_section else "")
|
||||
+ "Use `slack_thread_reply` to communicate in this Slack thread for clarifications, "
|
||||
"status updates, and final summaries. Use `slack_read_thread_messages` to read any "
|
||||
"Slack messages by providing channel_id and message_ts."
|
||||
)
|
||||
content_blocks: list[dict[str, Any]] = [create_text_block(prompt)]
|
||||
|
||||
image_urls = webapp.dedupe_urls(
|
||||
[url for msg in context_messages for url in webapp.extract_image_urls(msg.get("text", ""))]
|
||||
+ [
|
||||
f["url_private"]
|
||||
for msg in context_messages
|
||||
for f in msg.get("files", [])
|
||||
if isinstance(f, dict)
|
||||
and f.get("mimetype", "").startswith("image/")
|
||||
and f.get("url_private")
|
||||
]
|
||||
+ image_urls_from_links
|
||||
)
|
||||
|
||||
mapped_login = await webapp.login_for_slack_id(user_id)
|
||||
if not mapped_login and user_email:
|
||||
mapped_login = await webapp.login_for_email(user_email)
|
||||
|
||||
if image_urls:
|
||||
resolved_model_id = await webapp.resolve_agent_model_id(mapped_login)
|
||||
if webapp.model_supports_images(resolved_model_id):
|
||||
webapp.logger.info("Preparing %d image(s) for Slack mention", len(image_urls))
|
||||
async with httpx.AsyncClient(timeout=webapp.DEFAULT_HTTP_TIMEOUT) as http_client:
|
||||
for image_url in image_urls:
|
||||
image_block = await webapp.fetch_image_block(image_url, http_client)
|
||||
if image_block:
|
||||
content_blocks.append(image_block)
|
||||
else:
|
||||
webapp.logger.warning(
|
||||
"Skipping %d image(s) for Slack mention: model %s does not support images",
|
||||
len(image_urls),
|
||||
resolved_model_id,
|
||||
)
|
||||
prompt += webapp.vision_not_supported_warning(resolved_model_id, len(image_urls))
|
||||
content_blocks[0] = create_text_block(prompt)
|
||||
image_urls = []
|
||||
|
||||
# Open SWE opens PRs as the triggering user, so a run only proceeds when we
|
||||
# have a valid user GitHub token. Users who have never signed in with
|
||||
# GitHub, and users whose stored authorization is no longer usable, are
|
||||
# blocked and prompted to set up via the dashboard. Bot-token-only
|
||||
# deployments are exempt — they run on the installation token.
|
||||
user_token: str | None = None
|
||||
if mapped_login:
|
||||
try:
|
||||
user_token = await webapp.get_valid_access_token(mapped_login)
|
||||
except Exception: # noqa: BLE001
|
||||
webapp.logger.debug(
|
||||
"Failed to resolve GitHub token for %s; treating as unauthenticated",
|
||||
mapped_login,
|
||||
exc_info=True,
|
||||
)
|
||||
user_token = None
|
||||
has_valid_user_token = bool(user_token)
|
||||
|
||||
if not has_valid_user_token and not webapp.is_bot_token_only_mode():
|
||||
# A stored-but-unusable token means "sign in again"; no record at all
|
||||
# means the user has never connected GitHub + Slack via the dashboard.
|
||||
# Guard the store read like token resolution above so a transient
|
||||
# failure still yields an actionable prompt and clears the status.
|
||||
has_token_record = False
|
||||
if mapped_login:
|
||||
try:
|
||||
has_token_record = await webapp.has_access_token_record(mapped_login)
|
||||
except Exception: # noqa: BLE001
|
||||
webapp.logger.debug(
|
||||
"Failed to check GitHub token record for %s; prompting sign-in",
|
||||
mapped_login,
|
||||
exc_info=True,
|
||||
)
|
||||
reason = "revoked" if has_token_record else "unlinked"
|
||||
webapp.logger.info(
|
||||
"Blocking Slack run for thread %s: no valid user GitHub token (%s)",
|
||||
thread_id,
|
||||
reason,
|
||||
)
|
||||
if user_id:
|
||||
await webapp._post_account_link_prompt(
|
||||
channel_id, thread_ts, user_id, user_email, reason=reason
|
||||
)
|
||||
await webapp.set_slack_assistant_status(channel_id, thread_ts, status="")
|
||||
return
|
||||
|
||||
configurable: dict[str, Any] = {
|
||||
"repo": repo_config,
|
||||
"slack_thread": {
|
||||
"channel_id": channel_id,
|
||||
"thread_ts": thread_ts,
|
||||
"triggering_user_id": user_id,
|
||||
"triggering_user_name": user_name,
|
||||
"triggering_user_email": user_email,
|
||||
"triggering_event_ts": event_ts,
|
||||
},
|
||||
"user_email": user_email,
|
||||
"source": "slack",
|
||||
}
|
||||
if mapped_login:
|
||||
configurable["github_login"] = mapped_login
|
||||
|
||||
thread_plan_mode = await webapp._get_thread_plan_mode(thread_id)
|
||||
if thread_plan_mode is not None:
|
||||
configurable["plan_mode"] = thread_plan_mode
|
||||
|
||||
langgraph_client = webapp.get_client(url=webapp.LANGGRAPH_URL)
|
||||
is_first_mention = not await webapp._thread_exists(thread_id)
|
||||
await webapp._upsert_slack_thread_repo_metadata(thread_id, repo_config, langgraph_client)
|
||||
# Pass the login resolved above (from the stable Slack user id) so the thread is
|
||||
# always tagged with github_login — the key the dashboard searches by. Without
|
||||
# it, upsert re-resolves from the Slack profile email, which can miss.
|
||||
await webapp.upsert_agent_thread_owner_metadata(
|
||||
thread_id,
|
||||
source="slack",
|
||||
repo_config=repo_config,
|
||||
github_login=mapped_login or "",
|
||||
user_email=user_email or "",
|
||||
title=clean_text if is_first_mention else "",
|
||||
source_context={"slack_thread": configurable["slack_thread"]},
|
||||
)
|
||||
|
||||
run = await webapp.dispatch_agent_run(
|
||||
thread_id,
|
||||
content_blocks,
|
||||
configurable,
|
||||
source="slack",
|
||||
metadata=webapp._AGENT_VERSION_METADATA,
|
||||
client=langgraph_client,
|
||||
)
|
||||
webapp.logger.info(
|
||||
"Slack LangGraph run %s dispatched for thread %s",
|
||||
webapp._run_id_for_logging(run),
|
||||
thread_id,
|
||||
)
|
||||
run_id = run.get("run_id")
|
||||
if is_first_mention:
|
||||
trace_message_ts = await webapp.post_slack_trace_reply(channel_id, thread_ts, thread_id)
|
||||
await webapp.set_slack_assistant_status(channel_id, thread_ts)
|
||||
if isinstance(run_id, str) and run_id:
|
||||
await webapp.store_slack_run_mapping(
|
||||
langgraph_client,
|
||||
channel_id,
|
||||
thread_ts,
|
||||
run_id,
|
||||
message_ts=trace_message_ts,
|
||||
triggering_user_id=user_id,
|
||||
)
|
||||
else:
|
||||
webapp.logger.info(
|
||||
"Skipping Slack trace reply for thread %s — agent will reply when run completes",
|
||||
thread_id,
|
||||
)
|
||||
if isinstance(run_id, str) and run_id:
|
||||
await webapp.store_slack_run_mapping(
|
||||
langgraph_client,
|
||||
channel_id,
|
||||
thread_ts,
|
||||
run_id,
|
||||
triggering_user_id=user_id,
|
||||
)
|
||||
|
|
@ -1,188 +0,0 @@
|
|||
"""Unit tests for the auto-fix webhook helpers in agent.webapp."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from agent import webapp
|
||||
|
||||
|
||||
def test_parse_autofix_command() -> None:
|
||||
assert webapp._parse_autofix_command("@open-swe autofix off") is True
|
||||
assert webapp._parse_autofix_command("@open-swe autofix on") is False
|
||||
assert webapp._parse_autofix_command("@openswe please autofix off now") is True
|
||||
# Missing the mention -> not a command.
|
||||
assert webapp._parse_autofix_command("autofix off") is None
|
||||
# Mention but no command keyword.
|
||||
assert webapp._parse_autofix_command("@open-swe fix this") is None
|
||||
|
||||
|
||||
def test_pr_ref_from_issue_comment() -> None:
|
||||
payload = {
|
||||
"repository": {"owner": {"login": "o"}, "name": "r"},
|
||||
"issue": {
|
||||
"number": 7,
|
||||
"pull_request": {"html_url": "https://github.com/o/r/pull/7"},
|
||||
},
|
||||
}
|
||||
ref = webapp._pr_ref_from_comment_payload(payload, "issue_comment")
|
||||
assert ref == {"owner": "o", "name": "r", "number": 7, "url": "https://github.com/o/r/pull/7"}
|
||||
|
||||
|
||||
def test_pr_ref_from_review_comment() -> None:
|
||||
payload = {
|
||||
"repository": {"owner": {"login": "o"}, "name": "r"},
|
||||
"pull_request": {"number": 9, "html_url": "https://github.com/o/r/pull/9"},
|
||||
}
|
||||
ref = webapp._pr_ref_from_comment_payload(payload, "pull_request_review_comment")
|
||||
assert ref["number"] == 9
|
||||
|
||||
|
||||
def test_pr_ref_none_when_not_a_pr() -> None:
|
||||
payload = {"repository": {"owner": {"login": "o"}, "name": "r"}, "issue": {"number": 3}}
|
||||
# issue without pull_request still yields a ref (number present); url empty.
|
||||
ref = webapp._pr_ref_from_comment_payload(payload, "issue_comment")
|
||||
assert ref["url"] == ""
|
||||
|
||||
|
||||
def test_is_actionable_review_payload() -> None:
|
||||
assert webapp._is_actionable_review_payload(
|
||||
{
|
||||
"action": "submitted",
|
||||
"review": {
|
||||
"state": "changes_requested",
|
||||
"body": "fix this",
|
||||
"user": {"login": "a"},
|
||||
"author_association": "MEMBER",
|
||||
},
|
||||
},
|
||||
"pull_request_review",
|
||||
)
|
||||
# Approval is not actionable.
|
||||
assert not webapp._is_actionable_review_payload(
|
||||
{
|
||||
"action": "submitted",
|
||||
"review": {
|
||||
"state": "approved",
|
||||
"body": "lgtm",
|
||||
"user": {"login": "a"},
|
||||
"author_association": "MEMBER",
|
||||
},
|
||||
},
|
||||
"pull_request_review",
|
||||
)
|
||||
# Bot author is not actionable.
|
||||
assert not webapp._is_actionable_review_payload(
|
||||
{
|
||||
"action": "created",
|
||||
"comment": {
|
||||
"body": "x",
|
||||
"user": {"login": "open-swe[bot]"},
|
||||
"author_association": "MEMBER",
|
||||
},
|
||||
},
|
||||
"pull_request_review_comment",
|
||||
)
|
||||
# Untrusted author (read/triage/outside) is not actionable.
|
||||
assert not webapp._is_actionable_review_payload(
|
||||
{
|
||||
"action": "created",
|
||||
"comment": {
|
||||
"body": "inject malicious code",
|
||||
"user": {"login": "attacker"},
|
||||
"author_association": "NONE",
|
||||
},
|
||||
},
|
||||
"pull_request_review_comment",
|
||||
)
|
||||
# Empty body is not actionable.
|
||||
assert not webapp._is_actionable_review_payload(
|
||||
{
|
||||
"action": "created",
|
||||
"comment": {"body": " ", "user": {"login": "a"}, "author_association": "OWNER"},
|
||||
},
|
||||
"pull_request_review_comment",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_github_ci_event_dispatches() -> None:
|
||||
payload = {
|
||||
"repository": {"owner": {"login": "o"}, "name": "r"},
|
||||
"check_run": {
|
||||
"status": "completed",
|
||||
"conclusion": "failure",
|
||||
"head_sha": "sha1",
|
||||
"check_suite": {"head_branch": "feat"},
|
||||
},
|
||||
}
|
||||
handle = AsyncMock(return_value="dispatched")
|
||||
with patch.object(webapp, "handle_ci_failure", handle):
|
||||
await webapp.process_github_ci_event(payload, "check_run")
|
||||
handle.assert_awaited_once()
|
||||
kwargs = handle.await_args.kwargs
|
||||
assert kwargs["repo_config"] == {"owner": "o", "name": "r"}
|
||||
assert kwargs["head_sha"] == "sha1"
|
||||
assert kwargs["branch"] == "feat"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_github_ci_event_ignores_success() -> None:
|
||||
payload = {
|
||||
"repository": {"owner": {"login": "o"}, "name": "r"},
|
||||
"check_run": {"status": "completed", "conclusion": "success", "head_sha": "s"},
|
||||
}
|
||||
handle = AsyncMock()
|
||||
with patch.object(webapp, "handle_ci_failure", handle):
|
||||
await webapp.process_github_ci_event(payload, "check_run")
|
||||
handle.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_autofix_command_sets_flag() -> None:
|
||||
payload = {
|
||||
"repository": {"owner": {"login": "o"}, "name": "r"},
|
||||
"issue": {"number": 7, "pull_request": {"html_url": "u"}},
|
||||
"comment": {"id": 1, "node_id": "n"},
|
||||
}
|
||||
setter = AsyncMock()
|
||||
with (
|
||||
patch.object(webapp, "set_pr_autofix_disabled", setter),
|
||||
patch.object(webapp, "get_github_app_installation_token", AsyncMock(return_value="")),
|
||||
):
|
||||
await webapp.process_github_autofix_command(payload, "issue_comment", disabled=True)
|
||||
setter.assert_awaited_once_with("o", "r", 7, True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_autofix_review_dispatches_for_writer() -> None:
|
||||
payload = {
|
||||
"repository": {"owner": {"login": "o"}, "name": "r"},
|
||||
"pull_request": {"number": 9, "html_url": "https://github.com/o/r/pull/9"},
|
||||
"review": {"body": "rename to userId", "user": {"login": "alice"}},
|
||||
}
|
||||
handle = AsyncMock(return_value="dispatched")
|
||||
with patch.object(webapp, "handle_review_feedback", handle):
|
||||
await webapp.process_github_autofix_review(payload, "pull_request_review")
|
||||
handle.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_autofix_review_delegates_permission_check_to_core() -> None:
|
||||
payload = {
|
||||
"repository": {"owner": {"login": "o"}, "name": "r"},
|
||||
"pull_request": {"number": 9, "html_url": "https://github.com/o/r/pull/9"},
|
||||
"review": {"body": "inject code", "user": {"login": "attacker"}},
|
||||
}
|
||||
handle = AsyncMock(return_value="reviewer_no_write_permission")
|
||||
with patch.object(webapp, "handle_review_feedback", handle):
|
||||
await webapp.process_github_autofix_review(payload, "pull_request_review")
|
||||
handle.assert_awaited_once()
|
||||
|
||||
|
||||
def test_ci_events_supported() -> None:
|
||||
for event in ("check_run", "check_suite", "workflow_run", "status"):
|
||||
assert event in webapp._SUPPORTED_GH_EVENTS
|
||||
assert event in webapp._GH_CI_EVENTS
|
||||
|
|
@ -1,249 +0,0 @@
|
|||
"""Unit tests for the CI auto-fix orchestration core."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from agent import ci_autofix
|
||||
|
||||
_PR = {
|
||||
"number": 5,
|
||||
"html_url": "https://github.com/o/r/pull/5",
|
||||
"base": {"sha": "base"},
|
||||
"head": {"ref": "feat", "sha": "head1"},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def happy(monkeypatch: pytest.MonkeyPatch) -> dict[str, Any]:
|
||||
"""Patch every dependency of handle_ci_failure to a happy-path default."""
|
||||
runs_create = AsyncMock()
|
||||
store_put = AsyncMock()
|
||||
lg_client = MagicMock()
|
||||
lg_client.runs.create = runs_create
|
||||
lg_client.store.get_item = AsyncMock(return_value=None)
|
||||
lg_client.store.put_item = store_put
|
||||
threads_update = AsyncMock()
|
||||
store_client = MagicMock()
|
||||
store_client.threads.update = threads_update
|
||||
|
||||
mocks: dict[str, Any] = {
|
||||
"runs_create": runs_create,
|
||||
"threads_update": threads_update,
|
||||
"status_check": AsyncMock(return_value=True),
|
||||
"store_put": store_put,
|
||||
}
|
||||
|
||||
monkeypatch.setattr(ci_autofix, "_user_autofix_enabled", AsyncMock(return_value=True))
|
||||
monkeypatch.setattr(ci_autofix, "is_review_repo_enabled", AsyncMock(return_value=True))
|
||||
monkeypatch.setattr(
|
||||
ci_autofix, "get_github_app_installation_token", AsyncMock(return_value="tok")
|
||||
)
|
||||
monkeypatch.setattr(ci_autofix, "is_pr_autofix_disabled", AsyncMock(return_value=False))
|
||||
monkeypatch.setattr(
|
||||
ci_autofix,
|
||||
"find_agent_thread_for_pr",
|
||||
AsyncMock(return_value=("t1", {"github_login": "alice", "autofix_attempts": 0})),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ci_autofix,
|
||||
"list_failing_check_runs",
|
||||
AsyncMock(return_value=[{"name": "lint", "conclusion": "failure", "details_url": ""}]),
|
||||
)
|
||||
monkeypatch.setattr(ci_autofix, "list_failing_statuses", AsyncMock(return_value=[]))
|
||||
monkeypatch.setattr(ci_autofix, "names_failing_on_base", AsyncMock(return_value=set()))
|
||||
monkeypatch.setattr(
|
||||
ci_autofix, "head_commit_author_login", AsyncMock(return_value="open-swe[bot]")
|
||||
)
|
||||
monkeypatch.setattr(ci_autofix, "is_thread_active", AsyncMock(return_value=False))
|
||||
monkeypatch.setattr(ci_autofix, "post_autofix_status_check", mocks["status_check"])
|
||||
monkeypatch.setattr(ci_autofix, "langgraph_client", lambda: lg_client)
|
||||
monkeypatch.setattr(ci_autofix, "get_client", lambda: store_client)
|
||||
return mocks
|
||||
|
||||
|
||||
async def _run(**overrides: Any) -> str:
|
||||
kwargs: dict[str, Any] = {
|
||||
"repo_config": {"owner": "o", "name": "r"},
|
||||
"branch": "feat",
|
||||
"head_sha": "head1",
|
||||
"pr": _PR,
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return await ci_autofix.handle_ci_failure(**kwargs)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_happy_path(happy: dict[str, Any]) -> None:
|
||||
result = await _run()
|
||||
assert result == "dispatched"
|
||||
happy["runs_create"].assert_awaited_once()
|
||||
happy["threads_update"].assert_awaited()
|
||||
happy["status_check"].assert_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batches_when_thread_busy(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "is_thread_active", AsyncMock(return_value=True))
|
||||
result = await _run()
|
||||
assert result == "batched"
|
||||
happy["store_put"].assert_awaited()
|
||||
happy["runs_create"].assert_not_called()
|
||||
# A batched event must not burn an attempt or mark the SHA handled, so a later
|
||||
# webhook/sweep can still dispatch if the in-flight run never consumes it.
|
||||
happy["threads_update"].assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_user_disabled(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "_user_autofix_enabled", AsyncMock(return_value=False))
|
||||
assert await _run() == "autofix_disabled_user"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_repo_not_enabled(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "is_review_repo_enabled", AsyncMock(return_value=False))
|
||||
assert await _run() == "repo_not_enabled"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_pr_disabled(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "is_pr_autofix_disabled", AsyncMock(return_value=True))
|
||||
assert await _run() == "pr_disabled"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_no_agent_thread(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "find_agent_thread_for_pr", AsyncMock(return_value=None))
|
||||
assert await _run() == "no_agent_thread"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_max_attempts(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(
|
||||
ci_autofix,
|
||||
"find_agent_thread_for_pr",
|
||||
AsyncMock(
|
||||
return_value=(
|
||||
"t1",
|
||||
{"github_login": "alice", "autofix_attempts": ci_autofix.MAX_AUTOFIX_ATTEMPTS},
|
||||
)
|
||||
),
|
||||
)
|
||||
assert await _run() == "max_attempts"
|
||||
happy["status_check"].assert_awaited()
|
||||
happy["runs_create"].assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_all_failing_on_base(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "names_failing_on_base", AsyncMock(return_value={"lint"}))
|
||||
assert await _run() == "all_failing_on_base"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_already_handled(happy: dict[str, Any], monkeypatch) -> None:
|
||||
key = ci_autofix._dedupe_key("head1")
|
||||
monkeypatch.setattr(
|
||||
ci_autofix,
|
||||
"find_agent_thread_for_pr",
|
||||
AsyncMock(return_value=("t1", {"github_login": "alice", "autofix_handled": [key]})),
|
||||
)
|
||||
assert await _run() == "already_handled"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skip_human_commit(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "head_commit_author_login", AsyncMock(return_value="mallory"))
|
||||
assert await _run() == "human_commit"
|
||||
happy["runs_create"].assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_failing_checks(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "list_failing_check_runs", AsyncMock(return_value=[]))
|
||||
assert await _run(failing_checks=None) == "no_failing_checks"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ci_read_failed(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "list_failing_check_runs", AsyncMock(return_value=None))
|
||||
monkeypatch.setattr(ci_autofix, "list_failing_statuses", AsyncMock(return_value=None))
|
||||
assert await _run(failing_checks=None) == "ci_read_failed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_review_feedback_skips_user_disabled(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "_user_autofix_enabled", AsyncMock(return_value=False))
|
||||
assert (
|
||||
await ci_autofix.handle_review_feedback(
|
||||
repo_config={"owner": "o", "name": "r"},
|
||||
pr_number=5,
|
||||
pr_url="https://github.com/o/r/pull/5",
|
||||
reviewer="alice",
|
||||
body="fix this",
|
||||
)
|
||||
== "autofix_disabled_user"
|
||||
)
|
||||
happy["runs_create"].assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_review_feedback_batches_when_thread_busy(happy: dict[str, Any], monkeypatch) -> None:
|
||||
monkeypatch.setattr(ci_autofix, "has_repo_write_permission", AsyncMock(return_value=True))
|
||||
monkeypatch.setattr(ci_autofix, "is_thread_active", AsyncMock(return_value=True))
|
||||
result = await ci_autofix.handle_review_feedback(
|
||||
repo_config={"owner": "o", "name": "r"},
|
||||
pr_number=5,
|
||||
pr_url="https://github.com/o/r/pull/5",
|
||||
reviewer="alice",
|
||||
body="fix this",
|
||||
)
|
||||
assert result == "batched"
|
||||
happy["runs_create"].assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_review_feedback_checks_write_permission_after_user_gate(
|
||||
happy: dict[str, Any], monkeypatch
|
||||
) -> None:
|
||||
permission = AsyncMock(return_value=False)
|
||||
monkeypatch.setattr(ci_autofix, "has_repo_write_permission", permission)
|
||||
result = await ci_autofix.handle_review_feedback(
|
||||
repo_config={"owner": "o", "name": "r"},
|
||||
pr_number=5,
|
||||
pr_url="https://github.com/o/r/pull/5",
|
||||
reviewer="alice",
|
||||
body="fix this",
|
||||
)
|
||||
assert result == "reviewer_no_write_permission"
|
||||
permission.assert_awaited_once()
|
||||
happy["runs_create"].assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_agent_thread_picks_agent_skips_reviewer(monkeypatch) -> None:
|
||||
client = MagicMock()
|
||||
client.threads.search = AsyncMock(
|
||||
return_value=[
|
||||
{"thread_id": "rev", "metadata": {"kind": "reviewer", "agent_kind": "agent"}},
|
||||
{"thread_id": "ag", "metadata": {"agent_kind": "agent"}},
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(ci_autofix, "get_client", lambda: client)
|
||||
found = await ci_autofix.find_agent_thread_for_pr("https://github.com/o/r/pull/5")
|
||||
assert found is not None
|
||||
assert found[0] == "ag"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_find_agent_thread_none_when_only_reviewer(monkeypatch) -> None:
|
||||
client = MagicMock()
|
||||
client.threads.search = AsyncMock(
|
||||
return_value=[{"thread_id": "rev", "metadata": {"kind": "reviewer"}}]
|
||||
)
|
||||
monkeypatch.setattr(ci_autofix, "get_client", lambda: client)
|
||||
assert await ci_autofix.find_agent_thread_for_pr("u") is None
|
||||
138
tests/test_completion_webhook.py
Normal file
138
tests/test_completion_webhook.py
Normal file
|
|
@ -0,0 +1,138 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
from agent import completion
|
||||
|
||||
|
||||
class _FakeThreads:
|
||||
def __init__(self, metadata: dict[str, Any]) -> None:
|
||||
self._metadata = metadata
|
||||
self.updates: list[dict[str, Any]] = []
|
||||
|
||||
async def get(self, thread_id: str) -> dict[str, Any]:
|
||||
return {"thread_id": thread_id, "metadata": self._metadata}
|
||||
|
||||
async def update(self, *, thread_id: str, metadata: dict[str, Any]) -> None:
|
||||
self.updates.append(metadata)
|
||||
|
||||
|
||||
class _FakeClient:
|
||||
def __init__(self, metadata: dict[str, Any]) -> None:
|
||||
self.threads = _FakeThreads(metadata)
|
||||
|
||||
|
||||
def _slack_metadata() -> dict[str, Any]:
|
||||
return {
|
||||
"source": "slack",
|
||||
"source_context": {"slack_thread": {"channel_id": "C1", "thread_ts": "123.45"}},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_error_status_posts_slack_failure_reply(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
client = _FakeClient(_slack_metadata())
|
||||
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"] == "ok"
|
||||
reply.assert_awaited_once()
|
||||
args = reply.await_args.args
|
||||
assert args[0] == "C1"
|
||||
assert args[1] == "123.45"
|
||||
assert client.threads.updates == [{"failure_reply_posted": True}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_status_is_ignored(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
client = _FakeClient(_slack_metadata())
|
||||
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": "success"})
|
||||
|
||||
assert result["status"] == "ignored"
|
||||
reply.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_idempotent_when_already_replied(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
metadata = _slack_metadata()
|
||||
metadata["failure_reply_posted"] = True
|
||||
client = _FakeClient(metadata)
|
||||
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": "timeout"})
|
||||
|
||||
assert result["status"] == "ignored"
|
||||
reply.assert_not_called()
|
||||
assert client.threads.updates == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_linear_source_comments_on_issue(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
client = _FakeClient({"source": "linear", "source_context": {"linear_issue": {"id": "iss_1"}}})
|
||||
monkeypatch.setattr(completion, "langgraph_client", lambda: client)
|
||||
comment = AsyncMock(return_value=True)
|
||||
monkeypatch.setattr(completion, "comment_on_linear_issue", comment)
|
||||
|
||||
result = await completion.handle_run_completion({"thread_id": "t1", "status": "timeout"})
|
||||
|
||||
assert result["status"] == "ok"
|
||||
comment.assert_awaited_once()
|
||||
assert comment.await_args.args[0] == "iss_1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_thread_id_is_ignored() -> None:
|
||||
result = await completion.handle_run_completion({"status": "error"})
|
||||
assert result["status"] == "ignored"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_reply_channel_does_not_flag(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
client = _FakeClient({"source": "schedule"})
|
||||
monkeypatch.setattr(completion, "langgraph_client", lambda: client)
|
||||
|
||||
result = await completion.handle_run_completion({"thread_id": "t1", "status": "error"})
|
||||
|
||||
assert result["status"] == "ignored"
|
||||
assert client.threads.updates == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_interrupted_status_is_ignored(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
# Follow-ups use multitask_strategy="interrupt", so an interrupted run is a
|
||||
# healthy hand-off, not a failure to report.
|
||||
client = _FakeClient(_slack_metadata())
|
||||
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": "interrupted"})
|
||||
|
||||
assert result["status"] == "ignored"
|
||||
reply.assert_not_called()
|
||||
assert client.threads.updates == []
|
||||
|
||||
|
||||
def test_verify_run_complete_token(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
# No secret configured: fail closed (reject everything).
|
||||
monkeypatch.setattr(completion, "RUN_COMPLETE_WEBHOOK_SECRET", None)
|
||||
assert completion.verify_run_complete_token(None) is False
|
||||
assert completion.verify_run_complete_token("whatever") is False
|
||||
|
||||
# Secret configured: require an exact match.
|
||||
monkeypatch.setattr(completion, "RUN_COMPLETE_WEBHOOK_SECRET", "s3cret")
|
||||
assert completion.verify_run_complete_token("s3cret") is True
|
||||
assert completion.verify_run_complete_token("wrong") is False
|
||||
assert completion.verify_run_complete_token(None) is False
|
||||
|
|
@ -323,9 +323,6 @@ def test_process_github_review_finding_reply_uses_rereview_config(monkeypatch) -
|
|||
captured["interaction"] = (finding_id, interaction)
|
||||
return {}
|
||||
|
||||
async def fake_is_thread_active(_thread_id: str) -> bool:
|
||||
return False
|
||||
|
||||
async def fake_store_current_run_id(_thread_id: str, _run: object) -> None:
|
||||
return None
|
||||
|
||||
|
|
@ -348,7 +345,6 @@ def test_process_github_review_finding_reply_uses_rereview_config(monkeypatch) -
|
|||
monkeypatch.setattr(webapp, "reconcile_findings_with_review_threads", fake_reconcile)
|
||||
monkeypatch.setattr(webapp, "list_reviewer_findings", fake_list_findings)
|
||||
monkeypatch.setattr(webapp, "append_finding_interaction", fake_append_interaction)
|
||||
monkeypatch.setattr(webapp, "is_thread_active", fake_is_thread_active)
|
||||
monkeypatch.setattr(webapp, "_store_current_reviewer_run_id", fake_store_current_run_id)
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeLangGraphClient())
|
||||
|
||||
|
|
@ -381,7 +377,7 @@ def test_process_github_review_finding_reply_uses_rereview_config(monkeypatch) -
|
|||
assert config["finding_reply_id"] == "f_1"
|
||||
|
||||
|
||||
def test_process_github_review_finding_reply_queues_reply_body_when_active(monkeypatch) -> None:
|
||||
def test_process_github_review_finding_reply_dispatches_sanitized_reply_body(monkeypatch) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
async def fake_get_thread_metadata_safe(_thread_id: str) -> dict[str, object]:
|
||||
|
|
@ -407,15 +403,16 @@ def test_process_github_review_finding_reply_queues_reply_body_when_active(monke
|
|||
) -> dict[str, object]:
|
||||
return {}
|
||||
|
||||
async def fake_is_thread_active(_thread_id: str) -> bool:
|
||||
return True
|
||||
async def fake_store_current_run_id(_thread_id: str, _run: object) -> None:
|
||||
return None
|
||||
|
||||
async def fake_queue_message_for_thread(thread_id: str, message_content: object) -> bool:
|
||||
captured["queued"] = {"thread_id": thread_id, "message_content": message_content}
|
||||
return True
|
||||
class _FakeRunsClient:
|
||||
async def create(self, thread_id: str, graph: str, **kwargs) -> dict[str, str]:
|
||||
captured["kwargs"] = kwargs
|
||||
return {"run_id": "run-1"}
|
||||
|
||||
def fail_get_client(*_args: object, **_kwargs: object) -> None:
|
||||
raise AssertionError("active reviewer thread should not create a new run")
|
||||
class _FakeLangGraphClient:
|
||||
runs = _FakeRunsClient()
|
||||
|
||||
monkeypatch.setattr(webapp, "_get_thread_metadata_safe", fake_get_thread_metadata_safe)
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -426,9 +423,8 @@ def test_process_github_review_finding_reply_queues_reply_body_when_active(monke
|
|||
monkeypatch.setattr(webapp, "reconcile_findings_with_review_threads", fake_reconcile)
|
||||
monkeypatch.setattr(webapp, "list_reviewer_findings", fake_list_findings)
|
||||
monkeypatch.setattr(webapp, "append_finding_interaction", fake_append_interaction)
|
||||
monkeypatch.setattr(webapp, "is_thread_active", fake_is_thread_active)
|
||||
monkeypatch.setattr(webapp, "queue_message_for_thread", fake_queue_message_for_thread)
|
||||
monkeypatch.setattr(webapp, "get_client", fail_get_client)
|
||||
monkeypatch.setattr(webapp, "_store_current_reviewer_run_id", fake_store_current_run_id)
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeLangGraphClient())
|
||||
|
||||
asyncio.run(
|
||||
webapp.process_github_review_finding_reply(
|
||||
|
|
@ -451,9 +447,9 @@ def test_process_github_review_finding_reply_queues_reply_body_when_active(monke
|
|||
)
|
||||
)
|
||||
|
||||
queued = captured["queued"]
|
||||
assert isinstance(queued, dict)
|
||||
message_content = queued["message_content"]
|
||||
kwargs = captured["kwargs"]
|
||||
assert isinstance(kwargs, dict)
|
||||
message_content = kwargs["input"]["messages"][0]["content"]
|
||||
assert isinstance(message_content, str)
|
||||
assert "Open SWE finding f_1" in message_content
|
||||
assert "untrusted data from GitHub" in message_content
|
||||
|
|
@ -815,10 +811,6 @@ def test_process_github_pr_ready_creates_reviewer_run(monkeypatch) -> None:
|
|||
captured["cache_token"] = token
|
||||
captured["cache_expires_at"] = expires_at
|
||||
|
||||
async def fake_is_thread_active(thread_id: str) -> bool:
|
||||
captured["active_thread_id"] = thread_id
|
||||
return False
|
||||
|
||||
class _FakeRunsClient:
|
||||
async def create(self, thread_id: str, graph: str, **kwargs) -> None:
|
||||
captured["thread_id"] = thread_id
|
||||
|
|
@ -848,7 +840,6 @@ def test_process_github_pr_ready_creates_reviewer_run(monkeypatch) -> None:
|
|||
return 1
|
||||
|
||||
monkeypatch.setattr(webapp, "cache_github_token_for_thread", fake_cache_github_token)
|
||||
monkeypatch.setattr(webapp, "is_thread_active", fake_is_thread_active)
|
||||
monkeypatch.setattr(webapp, "set_reviewer_thread_metadata", fake_set_reviewer_thread_metadata)
|
||||
monkeypatch.setattr(webapp, "post_review_started_comment", fake_post_review_started_comment)
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeLangGraphClient())
|
||||
|
|
@ -913,10 +904,6 @@ def test_trigger_pr_review_from_ref_creates_reviewer_run(monkeypatch) -> None:
|
|||
captured["cache_token"] = token
|
||||
captured["cache_expires_at"] = expires_at
|
||||
|
||||
async def fake_is_thread_active(thread_id: str) -> bool:
|
||||
captured["active_thread_id"] = thread_id
|
||||
return False
|
||||
|
||||
class _FakeRunsClient:
|
||||
async def create(self, thread_id: str, graph: str, **kwargs) -> None:
|
||||
captured["thread_id"] = thread_id
|
||||
|
|
@ -950,7 +937,6 @@ def test_trigger_pr_review_from_ref_creates_reviewer_run(monkeypatch) -> None:
|
|||
|
||||
monkeypatch.setattr(webapp, "fetch_github_pr_metadata", fake_fetch_github_pr_metadata)
|
||||
monkeypatch.setattr(webapp, "cache_github_token_for_thread", fake_cache_github_token)
|
||||
monkeypatch.setattr(webapp, "is_thread_active", fake_is_thread_active)
|
||||
monkeypatch.setattr(webapp, "set_reviewer_thread_metadata", fake_set_reviewer_thread_metadata)
|
||||
monkeypatch.setattr(webapp, "post_review_started_comment", fake_post_review_started_comment)
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeLangGraphClient())
|
||||
|
|
@ -1029,7 +1015,7 @@ def test_trigger_pr_review_from_ref_respects_dashboard_opt_in(monkeypatch) -> No
|
|||
assert called is False
|
||||
|
||||
|
||||
def test_request_pr_review_tool_uses_shared_trigger(monkeypatch) -> None:
|
||||
async def test_request_pr_review_tool_uses_shared_trigger(monkeypatch) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
async def fake_trigger_pr_review_from_ref(
|
||||
|
|
@ -1065,7 +1051,7 @@ def test_request_pr_review_tool_uses_shared_trigger(monkeypatch) -> None:
|
|||
},
|
||||
)
|
||||
|
||||
result = request_pr_review_tool("https://github.com/langchain-ai/open-swe/pull/1244")
|
||||
result = await request_pr_review_tool("https://github.com/langchain-ai/open-swe/pull/1244")
|
||||
|
||||
pr_ref = captured["pr_ref"]
|
||||
assert isinstance(pr_ref, GitHubPrRef)
|
||||
|
|
@ -1156,9 +1142,6 @@ def test_process_github_issue_uses_resolved_user_token_for_reaction(monkeypatch)
|
|||
captured["fetch_token"] = token
|
||||
return []
|
||||
|
||||
async def fake_is_thread_active(thread_id: str) -> bool:
|
||||
return False
|
||||
|
||||
class _FakeRunsClient:
|
||||
async def create(self, *args, **kwargs) -> None:
|
||||
captured["run_created"] = True
|
||||
|
|
@ -1175,7 +1158,6 @@ def test_process_github_issue_uses_resolved_user_token_for_reaction(monkeypatch)
|
|||
monkeypatch.setattr(webapp, "_thread_exists", lambda thread_id: asyncio.sleep(0, result=False))
|
||||
monkeypatch.setattr(webapp, "react_to_github_comment", fake_react_to_github_comment)
|
||||
monkeypatch.setattr(webapp, "fetch_issue_comments", fake_fetch_issue_comments)
|
||||
monkeypatch.setattr(webapp, "is_thread_active", fake_is_thread_active)
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeLangGraphClient())
|
||||
monkeypatch.setattr(
|
||||
webapp,
|
||||
|
|
@ -1239,9 +1221,6 @@ def test_process_github_issue_existing_thread_uses_followup_prompt(monkeypatch)
|
|||
async def fake_thread_exists(thread_id: str) -> bool:
|
||||
return True
|
||||
|
||||
async def fake_is_thread_active(thread_id: str) -> bool:
|
||||
return False
|
||||
|
||||
class _FakeRunsClient:
|
||||
async def create(self, *args, **kwargs) -> None:
|
||||
captured["prompt"] = kwargs["input"]["messages"][0]["content"]
|
||||
|
|
@ -1258,7 +1237,6 @@ def test_process_github_issue_existing_thread_uses_followup_prompt(monkeypatch)
|
|||
monkeypatch.setattr(webapp, "_thread_exists", fake_thread_exists)
|
||||
monkeypatch.setattr(webapp, "react_to_github_comment", fake_react_to_github_comment)
|
||||
monkeypatch.setattr(webapp, "fetch_issue_comments", fake_fetch_issue_comments)
|
||||
monkeypatch.setattr(webapp, "is_thread_active", fake_is_thread_active)
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeLangGraphClient())
|
||||
monkeypatch.setattr(
|
||||
webapp,
|
||||
|
|
|
|||
|
|
@ -169,7 +169,7 @@ def test_plan_mode_guidance_section_present_when_enabled() -> None:
|
|||
assert "Plan Mode (ACTIVE)" in prompt
|
||||
|
||||
|
||||
def test_enter_plan_mode_tool_returns_command() -> None:
|
||||
async def test_enter_plan_mode_tool_returns_command() -> None:
|
||||
from langchain_core.messages import ToolMessage
|
||||
from langchain_core.tools import tool as as_tool
|
||||
from langgraph.types import Command
|
||||
|
|
@ -178,7 +178,7 @@ def test_enter_plan_mode_tool_returns_command() -> None:
|
|||
|
||||
# Wrap as the agent does so the InjectedToolCallId is supplied from the call.
|
||||
wrapped = as_tool(enter_plan_mode)
|
||||
result = wrapped.invoke(
|
||||
result = await wrapped.ainvoke(
|
||||
{"name": "enter_plan_mode", "args": {}, "id": "call-1", "type": "tool_call"}
|
||||
)
|
||||
assert isinstance(result, Command)
|
||||
|
|
|
|||
|
|
@ -95,19 +95,19 @@ async def test_clear_plan_comments_deletes_each(monkeypatch: pytest.MonkeyPatch)
|
|||
assert deleted == ["a", "b"]
|
||||
|
||||
|
||||
def test_save_plan_requires_run_context() -> None:
|
||||
async def test_save_plan_requires_run_context() -> None:
|
||||
from agent.tools.save_plan import save_plan
|
||||
|
||||
# No LangGraph run context → no thread_id → graceful error, not a crash.
|
||||
result = save_plan("## Plan")
|
||||
result = await save_plan("## Plan")
|
||||
assert result["success"] is False
|
||||
assert "thread_id" in result["error"]
|
||||
|
||||
|
||||
def test_save_plan_rejects_empty_markdown() -> None:
|
||||
async def test_save_plan_rejects_empty_markdown() -> None:
|
||||
from agent.tools.save_plan import save_plan
|
||||
|
||||
result = save_plan(" ")
|
||||
result = await save_plan(" ")
|
||||
assert result["success"] is False
|
||||
assert "empty" in result["error"]
|
||||
|
||||
|
|
|
|||
158
tests/test_reconcile_sweep.py
Normal file
158
tests/test_reconcile_sweep.py
Normal file
|
|
@ -0,0 +1,158 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
from agent import reconcile
|
||||
|
||||
|
||||
def _run(run_id: str, thread_id: str, age_seconds: float) -> dict[str, Any]:
|
||||
created = datetime.now(UTC) - timedelta(seconds=age_seconds)
|
||||
return {
|
||||
"run_id": run_id,
|
||||
"thread_id": thread_id,
|
||||
"status": "pending",
|
||||
"created_at": created.isoformat(),
|
||||
}
|
||||
|
||||
|
||||
class _FakeThreads:
|
||||
def __init__(self, pages: list[list[dict[str, Any]]]) -> None:
|
||||
self._pages = pages
|
||||
self.search_calls: list[dict[str, Any]] = []
|
||||
|
||||
async def search(self, **kwargs: Any) -> list[dict[str, Any]]:
|
||||
self.search_calls.append(kwargs)
|
||||
offset = kwargs.get("offset", 0)
|
||||
limit = kwargs.get("limit", 100)
|
||||
index = offset // limit if limit else 0
|
||||
if index < len(self._pages):
|
||||
return self._pages[index]
|
||||
return []
|
||||
|
||||
|
||||
class _FakeRuns:
|
||||
def __init__(self, runs_by_thread: dict[str, Any]) -> None:
|
||||
self._runs_by_thread = runs_by_thread
|
||||
self.cancel_many = AsyncMock(return_value=None)
|
||||
self.list_calls: list[tuple[str, dict[str, Any]]] = []
|
||||
|
||||
async def list(self, thread_id: str, **kwargs: Any) -> list[dict[str, Any]]:
|
||||
self.list_calls.append((thread_id, kwargs))
|
||||
value = self._runs_by_thread.get(thread_id, [])
|
||||
if isinstance(value, Exception):
|
||||
raise value
|
||||
return value
|
||||
|
||||
|
||||
class _FakeClient:
|
||||
def __init__(self, threads: _FakeThreads, runs: _FakeRuns) -> None:
|
||||
self.threads = threads
|
||||
self.runs = runs
|
||||
|
||||
|
||||
def _patch(monkeypatch: pytest.MonkeyPatch, client: _FakeClient) -> None:
|
||||
monkeypatch.setattr(reconcile, "langgraph_client", lambda: client)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancels_only_stale_pending_runs(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
threads = _FakeThreads([[{"thread_id": "t1"}]])
|
||||
runs = _FakeRuns(
|
||||
{
|
||||
"t1": [
|
||||
_run("old1", "t1", age_seconds=4000),
|
||||
_run("fresh1", "t1", age_seconds=60),
|
||||
_run("old2", "t1", age_seconds=10000),
|
||||
]
|
||||
}
|
||||
)
|
||||
_patch(monkeypatch, _FakeClient(threads, runs))
|
||||
|
||||
counts = await reconcile.reconcile_stale_runs(max_age_seconds=1800)
|
||||
|
||||
assert counts == {"threads_checked": 1, "stale_runs": 2, "cancelled": 2}
|
||||
runs.cancel_many.assert_awaited_once()
|
||||
kwargs = runs.cancel_many.await_args.kwargs
|
||||
assert kwargs["thread_id"] == "t1"
|
||||
assert sorted(kwargs["run_ids"]) == ["old1", "old2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_stale_runs_means_no_cancel(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
threads = _FakeThreads([[{"thread_id": "t1"}]])
|
||||
runs = _FakeRuns({"t1": [_run("fresh1", "t1", age_seconds=30)]})
|
||||
_patch(monkeypatch, _FakeClient(threads, runs))
|
||||
|
||||
counts = await reconcile.reconcile_stale_runs(max_age_seconds=1800)
|
||||
|
||||
assert counts == {"threads_checked": 1, "stale_runs": 0, "cancelled": 0}
|
||||
runs.cancel_many.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bad_thread_does_not_abort_sweep(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
threads = _FakeThreads([[{"thread_id": "bad"}, {"thread_id": "good"}]])
|
||||
runs = _FakeRuns(
|
||||
{
|
||||
"bad": RuntimeError("runs.list exploded"),
|
||||
"good": [_run("old1", "good", age_seconds=5000)],
|
||||
}
|
||||
)
|
||||
_patch(monkeypatch, _FakeClient(threads, runs))
|
||||
|
||||
counts = await reconcile.reconcile_stale_runs(max_age_seconds=1800)
|
||||
|
||||
# Both threads counted; the good thread is still reconciled despite the bad one.
|
||||
assert counts == {"threads_checked": 2, "stale_runs": 1, "cancelled": 1}
|
||||
runs.cancel_many.assert_awaited_once()
|
||||
assert runs.cancel_many.await_args.kwargs["thread_id"] == "good"
|
||||
assert runs.cancel_many.await_args.kwargs["run_ids"] == ["old1"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_paginates_busy_threads(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
full_page = [{"thread_id": f"t{i}"} for i in range(reconcile._SEARCH_PAGE_SIZE)]
|
||||
second_page = [{"thread_id": "tail"}]
|
||||
threads = _FakeThreads([full_page, second_page])
|
||||
runs_by_thread: dict[str, Any] = {t["thread_id"]: [] for t in full_page}
|
||||
runs_by_thread["tail"] = [_run("old", "tail", age_seconds=9000)]
|
||||
runs = _FakeRuns(runs_by_thread)
|
||||
_patch(monkeypatch, _FakeClient(threads, runs))
|
||||
|
||||
counts = await reconcile.reconcile_stale_runs(max_age_seconds=1800)
|
||||
|
||||
assert counts["threads_checked"] == reconcile._SEARCH_PAGE_SIZE + 1
|
||||
assert counts["cancelled"] == 1
|
||||
# Two search calls: first full page triggers a second page fetch.
|
||||
assert len(threads.search_calls) == 2
|
||||
assert threads.search_calls[0]["offset"] == 0
|
||||
assert threads.search_calls[1]["offset"] == reconcile._SEARCH_PAGE_SIZE
|
||||
assert threads.search_calls[0]["status"] == "busy"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unparseable_created_at_is_skipped(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
threads = _FakeThreads([[{"thread_id": "t1"}]])
|
||||
runs = _FakeRuns(
|
||||
{
|
||||
"t1": [
|
||||
{
|
||||
"run_id": "bad",
|
||||
"thread_id": "t1",
|
||||
"status": "pending",
|
||||
"created_at": "not-a-date",
|
||||
},
|
||||
_run("old", "t1", age_seconds=5000),
|
||||
]
|
||||
}
|
||||
)
|
||||
_patch(monkeypatch, _FakeClient(threads, runs))
|
||||
|
||||
counts = await reconcile.reconcile_stale_runs(max_age_seconds=1800)
|
||||
|
||||
assert counts == {"threads_checked": 1, "stale_runs": 1, "cancelled": 1}
|
||||
assert runs.cancel_many.await_args.kwargs["run_ids"] == ["old"]
|
||||
|
|
@ -454,10 +454,6 @@ def _setup_slack_mention_fakes(
|
|||
captured["user_names_by_id"] = user_names_by_id
|
||||
return "", []
|
||||
|
||||
async def fake_is_thread_active(thread_id: str) -> bool:
|
||||
captured["active_thread_id"] = thread_id
|
||||
return False
|
||||
|
||||
async def fake_post_slack_trace_reply(channel_id: str, thread_ts: str, thread_id: str) -> None:
|
||||
captured["trace_reply"] = {
|
||||
"channel_id": channel_id,
|
||||
|
|
@ -505,7 +501,6 @@ def _setup_slack_mention_fakes(
|
|||
async def fake_post_prompt(*args, **kwargs) -> None:
|
||||
captured["prompt"] = {"args": args, "kwargs": kwargs}
|
||||
|
||||
monkeypatch.setattr(webapp, "is_thread_active", fake_is_thread_active)
|
||||
monkeypatch.setattr(webapp, "post_slack_trace_reply", fake_post_slack_trace_reply)
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeLangGraphClientForProcess())
|
||||
monkeypatch.setattr(webapp, "login_for_slack_id", fake_login_for_slack_id)
|
||||
|
|
@ -547,7 +542,6 @@ def test_process_slack_mention_creates_thread_first_run_with_trace_reply(
|
|||
|
||||
assert captured["thread_exists_check"] == expected_thread_id
|
||||
assert captured["fetch_thread"] == {"channel_id": "C123", "thread_ts": thread_ts}
|
||||
assert captured["active_thread_id"] == expected_thread_id
|
||||
assert captured["metadata_update"] == {
|
||||
"thread_id": expected_thread_id,
|
||||
"metadata": {"repo": {"owner": "langchain-ai", "name": "open-swe"}},
|
||||
|
|
@ -564,7 +558,8 @@ def test_process_slack_mention_creates_thread_first_run_with_trace_reply(
|
|||
assert run_create["graph"] == "agent"
|
||||
kwargs = run_create["kwargs"]
|
||||
assert kwargs["if_not_exists"] == "create"
|
||||
assert "multitask_strategy" not in kwargs
|
||||
assert kwargs["multitask_strategy"] == "interrupt"
|
||||
assert kwargs["durability"] == "sync"
|
||||
assert kwargs["config"]["configurable"]["slack_thread"]["thread_ts"] == thread_ts
|
||||
prompt_block = kwargs["input"]["messages"][0]["content"][0]
|
||||
assert "## Default Repository Hint\nlangchain-ai/open-swe" in prompt_block["text"]
|
||||
|
|
@ -615,217 +610,6 @@ def test_process_slack_mention_skips_trace_reply_on_followup_mention(
|
|||
assert run_create["thread_id"] == expected_thread_id
|
||||
|
||||
|
||||
def test_process_slack_mention_queues_active_thread_message(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
async def fake_get_slack_user_info(user_id: str) -> dict:
|
||||
return {
|
||||
"profile": {
|
||||
"email": "mason@example.com",
|
||||
"display_name": "Mason",
|
||||
}
|
||||
}
|
||||
|
||||
async def fake_fetch_slack_thread_messages(channel_id: str, thread_ts: str) -> list[dict]:
|
||||
return [
|
||||
{"ts": "1700000000.000100", "text": "<@UBOT> first request", "user": "U123"},
|
||||
{
|
||||
"ts": "1700000000.000200",
|
||||
"text": "<@UBOT> include this screenshot https://example.com/image.png",
|
||||
"user": "U123",
|
||||
},
|
||||
]
|
||||
|
||||
async def fake_get_slack_user_names(user_ids: list[str]) -> dict[str, str]:
|
||||
captured["user_ids"] = user_ids
|
||||
return {"U123": "Mason"}
|
||||
|
||||
async def fake_resolve_slack_links_in_context(
|
||||
context_messages: list[dict], user_names_by_id: dict[str, str]
|
||||
) -> tuple[str, list[str]]:
|
||||
captured["context_messages"] = context_messages
|
||||
return "", []
|
||||
|
||||
async def fake_fetch_image_block(image_url: str, http_client: object) -> None:
|
||||
captured["image_url"] = image_url
|
||||
return None
|
||||
|
||||
async def fake_is_thread_active(thread_id: str) -> bool:
|
||||
captured["active_thread_id"] = thread_id
|
||||
return True
|
||||
|
||||
async def fake_queue_message_for_thread(thread_id: str, message_content: object) -> bool:
|
||||
captured["queued"] = {"thread_id": thread_id, "message_content": message_content}
|
||||
return True
|
||||
|
||||
async def fake_post_slack_trace_reply(*args, **kwargs) -> None:
|
||||
raise AssertionError("trace reply should not be posted for queued mid-run Slack messages")
|
||||
|
||||
async def fake_thread_exists(thread_id: str) -> bool:
|
||||
return True
|
||||
|
||||
class _FakeRunsClient:
|
||||
async def create(self, *args, **kwargs) -> None:
|
||||
raise AssertionError("run should not be created for active Slack threads")
|
||||
|
||||
class _FakeThreadsClientForProcess:
|
||||
async def update(self, *, thread_id: str, metadata: dict) -> None:
|
||||
captured["metadata_update"] = {"thread_id": thread_id, "metadata": metadata}
|
||||
|
||||
class _FakeLangGraphClientForProcess:
|
||||
runs = _FakeRunsClient()
|
||||
threads = _FakeThreadsClientForProcess()
|
||||
|
||||
monkeypatch.setattr(webapp, "SLACK_BOT_USERNAME", "open-swe")
|
||||
monkeypatch.setattr(webapp, "get_slack_user_info", fake_get_slack_user_info)
|
||||
monkeypatch.setattr(webapp, "fetch_slack_thread_messages", fake_fetch_slack_thread_messages)
|
||||
monkeypatch.setattr(webapp, "get_slack_user_names", fake_get_slack_user_names)
|
||||
monkeypatch.setattr(
|
||||
webapp, "resolve_slack_links_in_context", fake_resolve_slack_links_in_context
|
||||
)
|
||||
monkeypatch.setattr(webapp, "fetch_image_block", fake_fetch_image_block)
|
||||
monkeypatch.setattr(webapp, "is_thread_active", fake_is_thread_active)
|
||||
monkeypatch.setattr(webapp, "queue_message_for_thread", fake_queue_message_for_thread)
|
||||
monkeypatch.setattr(webapp, "post_slack_trace_reply", fake_post_slack_trace_reply)
|
||||
monkeypatch.setattr(webapp, "_thread_exists", fake_thread_exists)
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeLangGraphClientForProcess())
|
||||
|
||||
async def fake_login_for_slack_id(slack_user_id):
|
||||
return "mason-gh"
|
||||
|
||||
async def fake_login_for_email(email):
|
||||
return None
|
||||
|
||||
async def fake_refresh_cache() -> list:
|
||||
return []
|
||||
|
||||
async def fake_get_valid_access_token(login):
|
||||
return "user-token"
|
||||
|
||||
monkeypatch.setattr(webapp, "login_for_slack_id", fake_login_for_slack_id)
|
||||
monkeypatch.setattr(webapp, "login_for_email", fake_login_for_email)
|
||||
monkeypatch.setattr(webapp, "refresh_user_mapping_cache", fake_refresh_cache)
|
||||
monkeypatch.setattr(webapp, "get_valid_access_token", fake_get_valid_access_token)
|
||||
|
||||
async def fake_resolve_agent_model_id(github_login, per_thread_model_id=None):
|
||||
return "bedrock_converse:us.anthropic.claude-opus-4-8"
|
||||
|
||||
monkeypatch.setattr(webapp, "resolve_agent_model_id", fake_resolve_agent_model_id)
|
||||
|
||||
thread_ts = "1700000000.000100"
|
||||
event_ts = "1700000000.000200"
|
||||
expected_thread_id = generate_thread_id_from_slack_thread("C123", thread_ts)
|
||||
|
||||
asyncio.run(
|
||||
webapp.process_slack_mention(
|
||||
{
|
||||
"channel_id": "C123",
|
||||
"thread_ts": thread_ts,
|
||||
"event_ts": event_ts,
|
||||
"user_id": "U123",
|
||||
"text": "<@UBOT> include this screenshot https://example.com/image.png",
|
||||
"bot_user_id": "UBOT",
|
||||
},
|
||||
{"owner": "langchain-ai", "name": "open-swe"},
|
||||
)
|
||||
)
|
||||
|
||||
assert captured["active_thread_id"] == expected_thread_id
|
||||
assert captured["queued"]["thread_id"] == expected_thread_id
|
||||
queued_payload = captured["queued"]["message_content"]
|
||||
assert queued_payload["image_urls"] == ["https://example.com/image.png"]
|
||||
assert "## Latest Mention Request\ninclude this screenshot" in queued_payload["text"]
|
||||
|
||||
|
||||
def test_process_slack_mention_serializes_concurrent_run_dispatch(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
_setup_slack_mention_fakes(monkeypatch, captured)
|
||||
|
||||
thread_ts = "1700000001.000100"
|
||||
expected_thread_id = generate_thread_id_from_slack_thread("C123", thread_ts)
|
||||
first_active_started = asyncio.Event()
|
||||
finish_first_active = asyncio.Event()
|
||||
active_calls: list[str] = []
|
||||
run_creates: list[dict[str, object]] = []
|
||||
queued_messages: list[dict[str, object]] = []
|
||||
|
||||
async def fake_thread_exists(thread_id: str) -> bool:
|
||||
return False
|
||||
|
||||
async def fake_is_thread_active(thread_id: str) -> bool:
|
||||
active_calls.append(thread_id)
|
||||
if len(active_calls) == 1:
|
||||
first_active_started.set()
|
||||
await finish_first_active.wait()
|
||||
return bool(run_creates)
|
||||
|
||||
async def fake_queue_message_for_thread(thread_id: str, message_content: object) -> bool:
|
||||
queued_messages.append({"thread_id": thread_id, "message_content": message_content})
|
||||
return True
|
||||
|
||||
class _FakeRunsClient:
|
||||
async def create(self, thread_id: str, graph: str, **kwargs) -> dict[str, str]:
|
||||
run_creates.append({"thread_id": thread_id, "graph": graph, "kwargs": kwargs})
|
||||
return {"run_id": f"run-{len(run_creates)}"}
|
||||
|
||||
class _FakeThreadsClientForProcess:
|
||||
async def update(self, *, thread_id: str, metadata: dict) -> None:
|
||||
captured["metadata_update"] = {"thread_id": thread_id, "metadata": metadata}
|
||||
|
||||
class _FakeLangGraphClientForProcess:
|
||||
runs = _FakeRunsClient()
|
||||
threads = _FakeThreadsClientForProcess()
|
||||
|
||||
monkeypatch.setattr(webapp, "_thread_exists", fake_thread_exists)
|
||||
monkeypatch.setattr(webapp, "is_thread_active", fake_is_thread_active)
|
||||
monkeypatch.setattr(webapp, "queue_message_for_thread", fake_queue_message_for_thread)
|
||||
monkeypatch.setattr(webapp, "get_client", lambda url: _FakeLangGraphClientForProcess())
|
||||
|
||||
async def run_concurrent_mentions() -> None:
|
||||
first = asyncio.create_task(
|
||||
webapp.process_slack_mention(
|
||||
{
|
||||
"channel_id": "C123",
|
||||
"thread_ts": thread_ts,
|
||||
"event_ts": "1700000000.000200",
|
||||
"user_id": "U123",
|
||||
"text": "<@UBOT> first request",
|
||||
"bot_user_id": "UBOT",
|
||||
},
|
||||
{"owner": "langchain-ai", "name": "open-swe"},
|
||||
)
|
||||
)
|
||||
await first_active_started.wait()
|
||||
second = asyncio.create_task(
|
||||
webapp.process_slack_mention(
|
||||
{
|
||||
"channel_id": "C123",
|
||||
"thread_ts": thread_ts,
|
||||
"event_ts": "1700000000.000300",
|
||||
"user_id": "U123",
|
||||
"text": "<@UBOT> second request",
|
||||
"bot_user_id": "UBOT",
|
||||
},
|
||||
{"owner": "langchain-ai", "name": "open-swe"},
|
||||
)
|
||||
)
|
||||
await asyncio.sleep(0.05)
|
||||
assert active_calls == [expected_thread_id]
|
||||
finish_first_active.set()
|
||||
await asyncio.gather(first, second)
|
||||
|
||||
asyncio.run(run_concurrent_mentions())
|
||||
|
||||
assert active_calls == [expected_thread_id, expected_thread_id]
|
||||
assert len(run_creates) == 1
|
||||
assert run_creates[0]["thread_id"] == expected_thread_id
|
||||
assert queued_messages[0]["thread_id"] == expected_thread_id
|
||||
|
||||
|
||||
def test_process_slack_mention_unmapped_user_blocked_and_prompted(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue