diff --git a/agent/dashboard/thread_api.py b/agent/dashboard/thread_api.py index a6ef69c7..5954c3cd 100644 --- a/agent/dashboard/thread_api.py +++ b/agent/dashboard/thread_api.py @@ -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"] = [] diff --git a/agent/tools/open_pull_request.py b/agent/tools/open_pull_request.py index c41c0c3b..371fdd72 100644 --- a/agent/tools/open_pull_request.py +++ b/agent/tools/open_pull_request.py @@ -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": { diff --git a/agent/utils/github_comments.py b/agent/utils/github_comments.py index ce57d710..7b771810 100644 --- a/agent/utils/github_comments.py +++ b/agent/utils/github_comments.py @@ -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}", diff --git a/agent/webapp.py b/agent/webapp.py index 9d43272e..be42640a 100644 --- a/agent/webapp.py +++ b/agent/webapp.py @@ -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"} diff --git a/tests/test_agent_thread_pr_state.py b/tests/test_agent_thread_pr_state.py new file mode 100644 index 00000000..61d5686c --- /dev/null +++ b/tests/test_agent_thread_pr_state.py @@ -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() diff --git a/tests/test_dashboard_thread_api.py b/tests/test_dashboard_thread_api.py index 5e8696ac..3dbb3c9c 100644 --- a/tests/test_dashboard_thread_api.py +++ b/tests/test_dashboard_thread_api.py @@ -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: diff --git a/tests/test_open_pull_request.py b/tests/test_open_pull_request.py index e0270ecf..06edb0b5 100644 --- a/tests/test_open_pull_request.py +++ b/tests/test_open_pull_request.py @@ -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" diff --git a/ui/src/components/agents/AgentRunCard.tsx b/ui/src/components/agents/AgentRunCard.tsx index 90390114..2d4d1da5 100644 --- a/ui/src/components/agents/AgentRunCard.tsx +++ b/ui/src/components/agents/AgentRunCard.tsx @@ -61,7 +61,7 @@ export function AgentRunCard({ thread }: AgentRunCardProps) { {hasPr ? ( <> - Draft + {thread.pr?.state ?? "open"} ) : ( <> diff --git a/ui/src/components/agents/AgentsSidebar.tsx b/ui/src/components/agents/AgentsSidebar.tsx index ce66f9e3..f0b96c2a 100644 --- a/ui/src/components/agents/AgentsSidebar.tsx +++ b/ui/src/components/agents/AgentsSidebar.tsx @@ -8,6 +8,8 @@ import { ChartLineUpIcon, ChatCircleIcon, CircleNotchIcon, + GitMergeIcon, + GitPullRequestIcon, LightningIcon, PlusIcon, TrashIcon, @@ -46,6 +48,34 @@ const SOURCE_META: Record = { schedule: { icon: CalendarBlankIcon, label: "Triggered from a schedule" }, } +type PrState = NonNullable["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({ {thread.title} + {prMeta && PrIcon && ( + + {prMeta.label} + + )} {badge && ( {badge} diff --git a/ui/src/lib/agents/types.ts b/ui/src/lib/agents/types.ts index 34f828ad..1ba1c5f4 100644 --- a/ui/src/lib/agents/types.ts +++ b/ui/src/lib/agents/types.ts @@ -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