mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 09:13:14 +00:00
feat: track PR lifecycle state per thread for sidebar (#1492)
* feat: track PR lifecycle state per thread for sidebar Persist a PR's draft/open/merged/closed state on the agent thread and keep it in sync as PR webhooks fire, so the Agents UI can show per-thread PR status the way Cursor does. open_pull_request now records the initial draft/open state and pr_title; thread summaries expose diffStats; and the PR webhook refreshes pr_state on close/reopen/draft/ready transitions. Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com> * refactor: consolidate PR state mapping into shared derive_pr_state helper --------- Co-authored-by: open-swe[bot] <open-swe@users.noreply.github.com>
This commit is contained in:
parent
9a2f68e99d
commit
e0678e8c01
10 changed files with 287 additions and 4 deletions
|
|
@ -53,6 +53,8 @@ _PROXY_REQUEST_TIMEOUT = httpx.Timeout(30.0, connect=5.0)
|
|||
_PROXY_STREAM_TIMEOUT = httpx.Timeout(None)
|
||||
# Sources whose threads should surface in the Agents UI (besides "dashboard").
|
||||
_SURFACED_SOURCES: tuple[str, ...] = ("dashboard", "github", "slack", "linear", "schedule")
|
||||
# PR lifecycle states surfaced to the UI for a thread's associated pull request.
|
||||
_PR_STATES: frozenset[str] = frozenset({"draft", "open", "merged", "closed"})
|
||||
|
||||
|
||||
def _agent_version_metadata() -> dict[str, str]:
|
||||
|
|
@ -333,11 +335,18 @@ def _thread_summary(
|
|||
summary["pr"] = {
|
||||
"number": pr_number,
|
||||
"title": pr_title if isinstance(pr_title, str) else title,
|
||||
"state": pr_state if isinstance(pr_state, str) else "open",
|
||||
"state": pr_state if pr_state in _PR_STATES else "open",
|
||||
"headRef": metadata.get("branch_name") or "",
|
||||
"baseRef": metadata.get("base_branch") or "main",
|
||||
"url": pr_url,
|
||||
}
|
||||
diff_stats = metadata.get("diff_stats")
|
||||
if isinstance(diff_stats, dict):
|
||||
summary["diffStats"] = {
|
||||
"files": int(diff_stats.get("files") or 0),
|
||||
"additions": int(diff_stats.get("additions") or 0),
|
||||
"deletions": int(diff_stats.get("deletions") or 0),
|
||||
}
|
||||
# The transcript hydrates client-side from the SDK (`GET …/state` →
|
||||
# `stream.messages`); the summary only carries metadata.
|
||||
summary["messages"] = []
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from langgraph_sdk import get_client
|
|||
|
||||
from ..dashboard.agent_usage import record_agent_pr_usage
|
||||
from ..utils.github_app import get_github_app_installation_token
|
||||
from ..utils.github_comments import derive_pr_state
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
|
@ -118,6 +119,7 @@ async def _record_pr_telemetry(
|
|||
)
|
||||
pr_url = details.get("html_url") or pr.get("html_url")
|
||||
merged = bool(details.get("merged"))
|
||||
is_draft = bool(details.get("draft", pr.get("draft")))
|
||||
state = details.get("state") if isinstance(details.get("state"), str) else "open"
|
||||
additions = details.get("additions") if isinstance(details.get("additions"), int) else 0
|
||||
deletions = details.get("deletions") if isinstance(details.get("deletions"), int) else 0
|
||||
|
|
@ -147,7 +149,8 @@ async def _record_pr_telemetry(
|
|||
"agent_kind": "agent",
|
||||
"pr_url": pr_url if isinstance(pr_url, str) else "",
|
||||
"pr_number": pr_number,
|
||||
"pr_state": "merged" if merged else state,
|
||||
"pr_state": derive_pr_state(state=state, merged=merged, draft=is_draft),
|
||||
"pr_title": details.get("title") or pr.get("title"),
|
||||
"branch_name": head,
|
||||
"base_branch": base,
|
||||
"diff_stats": {
|
||||
|
|
|
|||
|
|
@ -64,6 +64,17 @@ def verify_github_signature(body: bytes, signature: str, *, secret: str) -> bool
|
|||
return hmac.compare_digest(expected, signature)
|
||||
|
||||
|
||||
def derive_pr_state(*, state: str | None, merged: bool, draft: bool) -> str:
|
||||
"""Map GitHub PR fields to the dashboard's pr_state vocabulary."""
|
||||
if merged:
|
||||
return "merged"
|
||||
if state == "closed":
|
||||
return "closed"
|
||||
if draft:
|
||||
return "draft"
|
||||
return "open"
|
||||
|
||||
|
||||
def get_thread_id_from_branch(branch_name: str) -> str | None:
|
||||
match = re.search(
|
||||
r"[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}",
|
||||
|
|
|
|||
|
|
@ -65,6 +65,7 @@ from .utils.github_comments import (
|
|||
OPEN_SWE_TAGS,
|
||||
GitHubAuthError,
|
||||
build_pr_prompt,
|
||||
derive_pr_state,
|
||||
extract_pr_context,
|
||||
fetch_issue_comments,
|
||||
fetch_pr_comments_since_last_tag,
|
||||
|
|
@ -1601,6 +1602,10 @@ _SUPPORTED_GH_PULL_REQUEST_ACTIONS = frozenset(
|
|||
)
|
||||
_GH_PR_WATCH_TOGGLE_ACTIONS = frozenset(["closed", "reopened", "converted_to_draft"])
|
||||
_GH_PR_FIRST_REVIEW_ACTIONS = frozenset(["opened", "ready_for_review"])
|
||||
# PR lifecycle actions that should refresh the agent thread's tracked pr_state.
|
||||
_GH_PR_AGENT_STATE_ACTIONS = frozenset(
|
||||
["closed", "reopened", "converted_to_draft", "ready_for_review"]
|
||||
)
|
||||
_SUPPORTED_GH_COMMENT_ACTIONS = {
|
||||
"issue_comment": frozenset(["created", "edited"]),
|
||||
"pull_request_review_comment": frozenset(["created", "edited"]),
|
||||
|
|
@ -2220,6 +2225,56 @@ async def _get_thread_metadata_safe(thread_id: str) -> dict[str, Any] | None:
|
|||
return metadata if isinstance(metadata, dict) else {}
|
||||
|
||||
|
||||
def _pr_state_from_payload(payload: dict[str, Any]) -> str | None:
|
||||
pull_request = payload.get("pull_request") if isinstance(payload, dict) else None
|
||||
if not isinstance(pull_request, dict):
|
||||
return None
|
||||
state = pull_request.get("state")
|
||||
return derive_pr_state(
|
||||
state=state if isinstance(state, str) else None,
|
||||
merged=bool(pull_request.get("merged")),
|
||||
draft=bool(pull_request.get("draft")),
|
||||
)
|
||||
|
||||
|
||||
async def update_agent_thread_pr_state(payload: dict[str, Any]) -> None:
|
||||
"""Keep an agent thread's tracked PR state in sync with PR lifecycle events.
|
||||
|
||||
The agent thread is located by the PR's html_url persisted in metadata when
|
||||
the PR was opened (``open_pull_request``). Reviewer threads are skipped.
|
||||
"""
|
||||
pull_request = payload.get("pull_request") if isinstance(payload, dict) else None
|
||||
if not isinstance(pull_request, dict):
|
||||
return
|
||||
pr_url = pull_request.get("html_url")
|
||||
new_state = _pr_state_from_payload(payload)
|
||||
if not isinstance(pr_url, str) or not pr_url or new_state is None:
|
||||
return
|
||||
|
||||
langgraph_client = get_client(url=LANGGRAPH_URL)
|
||||
try:
|
||||
threads = await langgraph_client.threads.search(metadata={"pr_url": pr_url}, limit=10)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug("Could not search threads for PR %s state update", pr_url, exc_info=True)
|
||||
return
|
||||
|
||||
for thread in threads or []:
|
||||
metadata = thread.get("metadata") if isinstance(thread, dict) else None
|
||||
if not isinstance(metadata, dict) or metadata.get("kind") == REVIEWER_THREAD_KIND:
|
||||
continue
|
||||
thread_id = thread.get("thread_id") or thread.get("id")
|
||||
if not isinstance(thread_id, str) or not thread_id:
|
||||
continue
|
||||
if metadata.get("pr_state") == new_state:
|
||||
continue
|
||||
try:
|
||||
await langgraph_client.threads.update(
|
||||
thread_id=thread_id, metadata={"pr_state": new_state}
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug("Failed to update pr_state for thread %s", thread_id, exc_info=True)
|
||||
|
||||
|
||||
async def process_github_pr_close(payload: dict[str, Any]) -> None:
|
||||
"""Toggle watch on the canonical reviewer thread on close/reopen/draft transitions.
|
||||
|
||||
|
|
@ -3028,6 +3083,8 @@ async def github_webhook(request: Request, background_tasks: BackgroundTasks) ->
|
|||
"status": "ignored",
|
||||
"reason": f"Unsupported GitHub pull_request action: {action}",
|
||||
}
|
||||
if action in _GH_PR_AGENT_STATE_ACTIONS:
|
||||
background_tasks.add_task(update_agent_thread_pr_state, payload)
|
||||
if action in _GH_PR_WATCH_TOGGLE_ACTIONS:
|
||||
if not await _is_repo_enabled_for_review(webhook_repo_config):
|
||||
return {"status": "ignored", "reason": "Repository not enabled for review"}
|
||||
|
|
|
|||
91
tests/test_agent_thread_pr_state.py
Normal file
91
tests/test_agent_thread_pr_state.py
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
"""Unit tests for agent-thread PR-state tracking from PR webhook events."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from agent import webapp
|
||||
|
||||
|
||||
def _pr_payload(*, state: str, merged: bool = False, draft: bool = False) -> dict[str, Any]:
|
||||
return {
|
||||
"pull_request": {
|
||||
"html_url": "https://github.com/lc/repo/pull/7",
|
||||
"state": state,
|
||||
"merged": merged,
|
||||
"draft": draft,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def test_pr_state_from_payload_merged() -> None:
|
||||
assert webapp._pr_state_from_payload(_pr_payload(state="closed", merged=True)) == "merged"
|
||||
|
||||
|
||||
def test_pr_state_from_payload_closed() -> None:
|
||||
assert webapp._pr_state_from_payload(_pr_payload(state="closed")) == "closed"
|
||||
|
||||
|
||||
def test_pr_state_from_payload_draft() -> None:
|
||||
assert webapp._pr_state_from_payload(_pr_payload(state="open", draft=True)) == "draft"
|
||||
|
||||
|
||||
def test_pr_state_from_payload_open() -> None:
|
||||
assert webapp._pr_state_from_payload(_pr_payload(state="open")) == "open"
|
||||
|
||||
|
||||
def test_pr_state_from_payload_missing_pull_request() -> None:
|
||||
assert webapp._pr_state_from_payload({}) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_agent_thread_pr_state_updates_matching_thread() -> None:
|
||||
fake_client = MagicMock()
|
||||
fake_client.threads.search = AsyncMock(
|
||||
return_value=[
|
||||
{
|
||||
"thread_id": "t1",
|
||||
"metadata": {"kind": "agent", "pr_state": "draft"},
|
||||
}
|
||||
]
|
||||
)
|
||||
fake_client.threads.update = AsyncMock()
|
||||
|
||||
with patch("agent.webapp.get_client", return_value=fake_client):
|
||||
await webapp.update_agent_thread_pr_state(_pr_payload(state="closed"))
|
||||
|
||||
fake_client.threads.search.assert_awaited_once()
|
||||
fake_client.threads.update.assert_awaited_once()
|
||||
assert fake_client.threads.update.await_args.kwargs["thread_id"] == "t1"
|
||||
assert fake_client.threads.update.await_args.kwargs["metadata"] == {"pr_state": "closed"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_agent_thread_pr_state_skips_reviewer_threads() -> None:
|
||||
fake_client = MagicMock()
|
||||
fake_client.threads.search = AsyncMock(
|
||||
return_value=[{"thread_id": "rev", "metadata": {"kind": "reviewer"}}]
|
||||
)
|
||||
fake_client.threads.update = AsyncMock()
|
||||
|
||||
with patch("agent.webapp.get_client", return_value=fake_client):
|
||||
await webapp.update_agent_thread_pr_state(_pr_payload(state="closed"))
|
||||
|
||||
fake_client.threads.update.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_agent_thread_pr_state_noop_when_state_unchanged() -> None:
|
||||
fake_client = MagicMock()
|
||||
fake_client.threads.search = AsyncMock(
|
||||
return_value=[{"thread_id": "t1", "metadata": {"pr_state": "merged"}}]
|
||||
)
|
||||
fake_client.threads.update = AsyncMock()
|
||||
|
||||
with patch("agent.webapp.get_client", return_value=fake_client):
|
||||
await webapp.update_agent_thread_pr_state(_pr_payload(state="closed", merged=True))
|
||||
|
||||
fake_client.threads.update.assert_not_called()
|
||||
|
|
@ -213,6 +213,59 @@ async def test_enrich_run_start_command_rejects_images_for_resolved_text_only_mo
|
|||
assert "does not support image input" in exc_info.value.detail
|
||||
|
||||
|
||||
def _thread_with_metadata(metadata: dict) -> dict:
|
||||
return {"thread_id": "t1", "status": "idle", "metadata": metadata}
|
||||
|
||||
|
||||
def test_thread_summary_includes_pr_and_diff_stats() -> None:
|
||||
summary = thread_api._thread_summary(
|
||||
_thread_with_metadata(
|
||||
{
|
||||
"repo_full_name": "langchain-ai/open-swe",
|
||||
"title": "Add feature",
|
||||
"pr_number": 42,
|
||||
"pr_url": "https://github.com/langchain-ai/open-swe/pull/42",
|
||||
"pr_state": "draft",
|
||||
"pr_title": "feat: add feature",
|
||||
"branch_name": "open-swe/feature",
|
||||
"base_branch": "main",
|
||||
"diff_stats": {"files": 3, "additions": 10, "deletions": 2},
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert summary["pr"] == {
|
||||
"number": 42,
|
||||
"title": "feat: add feature",
|
||||
"state": "draft",
|
||||
"headRef": "open-swe/feature",
|
||||
"baseRef": "main",
|
||||
"url": "https://github.com/langchain-ai/open-swe/pull/42",
|
||||
}
|
||||
assert summary["diffStats"] == {"files": 3, "additions": 10, "deletions": 2}
|
||||
|
||||
|
||||
def test_thread_summary_defaults_unknown_pr_state_to_open() -> None:
|
||||
summary = thread_api._thread_summary(
|
||||
_thread_with_metadata(
|
||||
{
|
||||
"pr_number": 7,
|
||||
"pr_url": "https://example.com/pull/7",
|
||||
"pr_state": "bogus",
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert summary["pr"]["state"] == "open"
|
||||
|
||||
|
||||
def test_thread_summary_omits_pr_when_no_pr_metadata() -> None:
|
||||
summary = thread_api._thread_summary(_thread_with_metadata({"title": "No PR"}))
|
||||
|
||||
assert "pr" not in summary
|
||||
assert "diffStats" not in summary
|
||||
|
||||
|
||||
async def test_proxy_commands_lazily_creates_missing_thread_only_for_run_start(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -206,3 +206,19 @@ def test_error_surfaced_on_failure(monkeypatch: pytest.MonkeyPatch) -> None:
|
|||
|
||||
async def _coro(value: Any) -> Any:
|
||||
return value
|
||||
|
||||
|
||||
def test_derive_pr_state_prefers_merged() -> None:
|
||||
assert opr.derive_pr_state(state="closed", merged=True, draft=True) == "merged"
|
||||
|
||||
|
||||
def test_derive_pr_state_closed_over_draft() -> None:
|
||||
assert opr.derive_pr_state(state="closed", merged=False, draft=True) == "closed"
|
||||
|
||||
|
||||
def test_derive_pr_state_draft() -> None:
|
||||
assert opr.derive_pr_state(state="open", merged=False, draft=True) == "draft"
|
||||
|
||||
|
||||
def test_derive_pr_state_open() -> None:
|
||||
assert opr.derive_pr_state(state="open", merged=False, draft=False) == "open"
|
||||
|
|
|
|||
|
|
@ -61,7 +61,7 @@ export function AgentRunCard({ thread }: AgentRunCardProps) {
|
|||
{hasPr ? (
|
||||
<>
|
||||
<GitPullRequestIcon className="size-3" />
|
||||
Draft
|
||||
<span className="capitalize">{thread.pr?.state ?? "open"}</span>
|
||||
</>
|
||||
) : (
|
||||
<>
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ import {
|
|||
ChartLineUpIcon,
|
||||
ChatCircleIcon,
|
||||
CircleNotchIcon,
|
||||
GitMergeIcon,
|
||||
GitPullRequestIcon,
|
||||
LightningIcon,
|
||||
PlusIcon,
|
||||
TrashIcon,
|
||||
|
|
@ -46,6 +48,34 @@ const SOURCE_META: Record<AgentSource, { icon: SourceIcon; label: string }> = {
|
|||
schedule: { icon: CalendarBlankIcon, label: "Triggered from a schedule" },
|
||||
}
|
||||
|
||||
type PrState = NonNullable<AgentThread["pr"]>["state"]
|
||||
|
||||
const PR_STATE_META: Record<
|
||||
PrState,
|
||||
{ icon: SourceIcon; label: string; className: string }
|
||||
> = {
|
||||
draft: {
|
||||
icon: GitPullRequestIcon,
|
||||
label: "Draft pull request",
|
||||
className: "text-[var(--ui-text-dim)]",
|
||||
},
|
||||
open: {
|
||||
icon: GitPullRequestIcon,
|
||||
label: "Open pull request",
|
||||
className: "text-[var(--ui-success)]",
|
||||
},
|
||||
merged: {
|
||||
icon: GitMergeIcon,
|
||||
label: "Merged pull request",
|
||||
className: "text-[var(--ui-accent)]",
|
||||
},
|
||||
closed: {
|
||||
icon: GitPullRequestIcon,
|
||||
label: "Closed pull request",
|
||||
className: "text-[var(--ui-danger)]",
|
||||
},
|
||||
}
|
||||
|
||||
interface AgentsSidebarProps {
|
||||
user: SessionUser
|
||||
activeThreadId?: string
|
||||
|
|
@ -227,6 +257,8 @@ function ThreadRow({
|
|||
? SOURCE_META[thread.source]
|
||||
: null
|
||||
const SourceIcon = source?.icon
|
||||
const prMeta = thread.pr ? PR_STATE_META[thread.pr.state] : null
|
||||
const PrIcon = prMeta?.icon
|
||||
const showFinishedIndicator = thread.status === "finished" && !thread.viewed
|
||||
|
||||
const openTrace = () => {
|
||||
|
|
@ -282,6 +314,17 @@ function ThreadRow({
|
|||
<span className="min-w-0 flex-1 truncate text-xs">
|
||||
{thread.title}
|
||||
</span>
|
||||
{prMeta && PrIcon && (
|
||||
<PrIcon
|
||||
className={cn(
|
||||
"size-3.5 shrink-0 group-hover:hidden",
|
||||
prMeta.className
|
||||
)}
|
||||
aria-label={prMeta.label}
|
||||
>
|
||||
<title>{prMeta.label}</title>
|
||||
</PrIcon>
|
||||
)}
|
||||
{badge && (
|
||||
<span className="shrink-0 rounded bg-[var(--ui-panel-2)] px-1.5 py-0.5 text-[10px] text-[var(--ui-text-dim)] group-hover:hidden">
|
||||
{badge}
|
||||
|
|
|
|||
|
|
@ -189,7 +189,7 @@ export interface AgentThread {
|
|||
pr?: {
|
||||
number: number
|
||||
title: string
|
||||
state: "draft" | "open" | "merged"
|
||||
state: "draft" | "open" | "merged" | "closed"
|
||||
headRef: string
|
||||
baseRef: string
|
||||
url: string
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue