mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-10-06 18:12:13 +00:00
fix: Add scripts for getting usage, fix dual committing (#1131)
This commit is contained in:
parent
0d8647c2cd
commit
87968ab813
7 changed files with 707 additions and 9 deletions
|
|
@ -16,6 +16,13 @@ from langchain.agents.middleware import AgentState, after_agent
|
||||||
from langgraph.config import get_config
|
from langgraph.config import get_config
|
||||||
from langgraph.runtime import Runtime
|
from langgraph.runtime import Runtime
|
||||||
|
|
||||||
|
from ..utils.authorship import (
|
||||||
|
OPEN_SWE_BOT_EMAIL,
|
||||||
|
OPEN_SWE_BOT_NAME,
|
||||||
|
add_pr_collaboration_note,
|
||||||
|
add_user_coauthor_trailer,
|
||||||
|
resolve_triggering_user_identity,
|
||||||
|
)
|
||||||
from ..utils.github import (
|
from ..utils.github import (
|
||||||
create_github_pr,
|
create_github_pr,
|
||||||
get_github_default_branch,
|
get_github_default_branch,
|
||||||
|
|
@ -84,6 +91,12 @@ async def open_pr_if_needed(
|
||||||
pr_title = pr_payload.get("title", "feat: Open SWE PR")
|
pr_title = pr_payload.get("title", "feat: Open SWE PR")
|
||||||
pr_body = pr_payload.get("body", "Automated PR created by Open SWE agent.")
|
pr_body = pr_payload.get("body", "Automated PR created by Open SWE agent.")
|
||||||
commit_message = pr_payload.get("commit_message", pr_title)
|
commit_message = pr_payload.get("commit_message", pr_title)
|
||||||
|
github_token = get_github_token()
|
||||||
|
user_identity = await asyncio.to_thread(
|
||||||
|
resolve_triggering_user_identity, config, github_token
|
||||||
|
)
|
||||||
|
pr_body = add_pr_collaboration_note(pr_body, user_identity)
|
||||||
|
commit_message = add_user_coauthor_trailer(commit_message, user_identity)
|
||||||
|
|
||||||
if not thread_id:
|
if not thread_id:
|
||||||
raise ValueError("No thread_id found in config")
|
raise ValueError("No thread_id found in config")
|
||||||
|
|
@ -135,14 +148,12 @@ async def open_pr_if_needed(
|
||||||
git_config_user,
|
git_config_user,
|
||||||
sandbox_backend,
|
sandbox_backend,
|
||||||
repo_dir,
|
repo_dir,
|
||||||
"open-swe[bot]",
|
OPEN_SWE_BOT_NAME,
|
||||||
"open-swe@users.noreply.github.com",
|
OPEN_SWE_BOT_EMAIL,
|
||||||
)
|
)
|
||||||
await asyncio.to_thread(git_add_all, sandbox_backend, repo_dir)
|
await asyncio.to_thread(git_add_all, sandbox_backend, repo_dir)
|
||||||
await asyncio.to_thread(git_commit, sandbox_backend, repo_dir, commit_message)
|
await asyncio.to_thread(git_commit, sandbox_backend, repo_dir, commit_message)
|
||||||
|
|
||||||
github_token = get_github_token()
|
|
||||||
|
|
||||||
if github_token:
|
if github_token:
|
||||||
await asyncio.to_thread(
|
await asyncio.to_thread(
|
||||||
git_push, sandbox_backend, repo_dir, target_branch, github_token
|
git_push, sandbox_backend, repo_dir, target_branch, github_token
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,13 @@ from typing import Any
|
||||||
|
|
||||||
from langgraph.config import get_config
|
from langgraph.config import get_config
|
||||||
|
|
||||||
|
from ..utils.authorship import (
|
||||||
|
OPEN_SWE_BOT_EMAIL,
|
||||||
|
OPEN_SWE_BOT_NAME,
|
||||||
|
add_pr_collaboration_note,
|
||||||
|
add_user_coauthor_trailer,
|
||||||
|
resolve_triggering_user_identity,
|
||||||
|
)
|
||||||
from ..utils.github import (
|
from ..utils.github import (
|
||||||
create_github_pr,
|
create_github_pr,
|
||||||
get_github_default_branch,
|
get_github_default_branch,
|
||||||
|
|
@ -131,6 +138,9 @@ def commit_and_open_pr(
|
||||||
return {"success": False, "error": "No sandbox found for thread", "pr_url": None}
|
return {"success": False, "error": "No sandbox found for thread", "pr_url": None}
|
||||||
|
|
||||||
repo_dir = resolve_repo_dir(sandbox_backend, repo_name)
|
repo_dir = resolve_repo_dir(sandbox_backend, repo_name)
|
||||||
|
github_token = get_github_token()
|
||||||
|
user_identity = resolve_triggering_user_identity(config, github_token)
|
||||||
|
pr_body = add_pr_collaboration_note(body, user_identity)
|
||||||
|
|
||||||
has_uncommitted_changes = git_has_uncommitted_changes(sandbox_backend, repo_dir)
|
has_uncommitted_changes = git_has_uncommitted_changes(sandbox_backend, repo_dir)
|
||||||
git_fetch_origin(sandbox_backend, repo_dir)
|
git_fetch_origin(sandbox_backend, repo_dir)
|
||||||
|
|
@ -163,12 +173,12 @@ def commit_and_open_pr(
|
||||||
git_config_user(
|
git_config_user(
|
||||||
sandbox_backend,
|
sandbox_backend,
|
||||||
repo_dir,
|
repo_dir,
|
||||||
"open-swe[bot]",
|
OPEN_SWE_BOT_NAME,
|
||||||
"open-swe@users.noreply.github.com",
|
OPEN_SWE_BOT_EMAIL,
|
||||||
)
|
)
|
||||||
git_add_all(sandbox_backend, repo_dir)
|
git_add_all(sandbox_backend, repo_dir)
|
||||||
|
|
||||||
commit_msg = commit_message or title
|
commit_msg = add_user_coauthor_trailer(commit_message or title, user_identity)
|
||||||
if has_uncommitted_changes:
|
if has_uncommitted_changes:
|
||||||
commit_result = git_commit(sandbox_backend, repo_dir, commit_msg)
|
commit_result = git_commit(sandbox_backend, repo_dir, commit_msg)
|
||||||
if commit_result.exit_code != 0:
|
if commit_result.exit_code != 0:
|
||||||
|
|
@ -178,7 +188,6 @@ def commit_and_open_pr(
|
||||||
"pr_url": None,
|
"pr_url": None,
|
||||||
}
|
}
|
||||||
|
|
||||||
github_token = get_github_token()
|
|
||||||
if not github_token:
|
if not github_token:
|
||||||
logger.error("commit_and_open_pr missing GitHub token for thread %s", thread_id)
|
logger.error("commit_and_open_pr missing GitHub token for thread %s", thread_id)
|
||||||
return {
|
return {
|
||||||
|
|
@ -204,7 +213,7 @@ def commit_and_open_pr(
|
||||||
title=title,
|
title=title,
|
||||||
head_branch=target_branch,
|
head_branch=target_branch,
|
||||||
base_branch=base_branch,
|
base_branch=base_branch,
|
||||||
body=body,
|
body=pr_body,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
|
||||||
165
agent/utils/authorship.py
Normal file
165
agent/utils/authorship.py
Normal file
|
|
@ -0,0 +1,165 @@
|
||||||
|
"""Helpers for collaborative commit and PR attribution."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
from .github_user_email_map import GITHUB_USER_EMAIL_MAP
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
OPEN_SWE_BOT_NAME = "open-swe[bot]"
|
||||||
|
OPEN_SWE_BOT_EMAIL = "open-swe@users.noreply.github.com"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class CollaboratorIdentity:
|
||||||
|
"""Identity used for git trailers and PR attribution."""
|
||||||
|
|
||||||
|
display_name: str
|
||||||
|
commit_name: str
|
||||||
|
commit_email: str
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_text(value: Any) -> str:
|
||||||
|
return value.strip() if isinstance(value, str) else ""
|
||||||
|
|
||||||
|
|
||||||
|
def _github_noreply_email(login: str, user_id: Any = None) -> str:
|
||||||
|
normalized_login = _normalize_text(login)
|
||||||
|
if not normalized_login:
|
||||||
|
return ""
|
||||||
|
|
||||||
|
normalized_user_id = str(user_id).strip() if user_id is not None else ""
|
||||||
|
if normalized_user_id:
|
||||||
|
return f"{normalized_user_id}+{normalized_login}@users.noreply.github.com"
|
||||||
|
return f"{normalized_login}@users.noreply.github.com"
|
||||||
|
|
||||||
|
|
||||||
|
def _identity_from_github_token(github_token: str | None) -> CollaboratorIdentity | None:
|
||||||
|
if not github_token:
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
response = httpx.get(
|
||||||
|
"https://api.github.com/user",
|
||||||
|
headers={
|
||||||
|
"Authorization": f"Bearer {github_token}",
|
||||||
|
"Accept": "application/vnd.github+json",
|
||||||
|
"X-GitHub-Api-Version": "2022-11-28",
|
||||||
|
},
|
||||||
|
timeout=5.0,
|
||||||
|
)
|
||||||
|
if response.status_code != 200: # noqa: PLR2004
|
||||||
|
logger.debug("GitHub user lookup returned %s", response.status_code)
|
||||||
|
return None
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
login = _normalize_text(payload.get("login"))
|
||||||
|
display_name = _normalize_text(payload.get("name")) or login
|
||||||
|
commit_email = _github_noreply_email(login, payload.get("id")) or _normalize_text(
|
||||||
|
payload.get("email")
|
||||||
|
)
|
||||||
|
if not display_name or not commit_email:
|
||||||
|
return None
|
||||||
|
if commit_email == OPEN_SWE_BOT_EMAIL and display_name == OPEN_SWE_BOT_NAME:
|
||||||
|
return None
|
||||||
|
return CollaboratorIdentity(
|
||||||
|
display_name=display_name,
|
||||||
|
commit_name=display_name,
|
||||||
|
commit_email=commit_email,
|
||||||
|
)
|
||||||
|
except httpx.HTTPError:
|
||||||
|
logger.debug("Failed to resolve GitHub user identity from token", exc_info=True)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _identity_from_config(config: dict[str, Any]) -> CollaboratorIdentity | None:
|
||||||
|
configurable = config.get("configurable", {})
|
||||||
|
|
||||||
|
github_login = _normalize_text(configurable.get("github_login"))
|
||||||
|
if github_login:
|
||||||
|
github_user_id = configurable.get("github_user_id")
|
||||||
|
commit_email = _github_noreply_email(github_login, github_user_id) or _normalize_text(
|
||||||
|
GITHUB_USER_EMAIL_MAP.get(github_login)
|
||||||
|
)
|
||||||
|
if commit_email:
|
||||||
|
return CollaboratorIdentity(
|
||||||
|
display_name=github_login,
|
||||||
|
commit_name=github_login,
|
||||||
|
commit_email=commit_email,
|
||||||
|
)
|
||||||
|
|
||||||
|
slack_thread = configurable.get("slack_thread", {})
|
||||||
|
linear_issue = configurable.get("linear_issue", {})
|
||||||
|
|
||||||
|
display_name = (
|
||||||
|
_normalize_text(slack_thread.get("triggering_user_name"))
|
||||||
|
or _normalize_text(linear_issue.get("triggering_user_name"))
|
||||||
|
or _normalize_text(configurable.get("user_email")).split("@", 1)[0]
|
||||||
|
)
|
||||||
|
commit_email = _normalize_text(configurable.get("user_email")) or _normalize_text(
|
||||||
|
slack_thread.get("triggering_user_email")
|
||||||
|
)
|
||||||
|
if display_name and commit_email:
|
||||||
|
return CollaboratorIdentity(
|
||||||
|
display_name=display_name,
|
||||||
|
commit_name=display_name,
|
||||||
|
commit_email=commit_email,
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_triggering_user_identity(
|
||||||
|
config: dict[str, Any],
|
||||||
|
github_token: str | None = None,
|
||||||
|
) -> CollaboratorIdentity | None:
|
||||||
|
"""Resolve the triggering user's git identity.
|
||||||
|
|
||||||
|
Prefer the GitHub account identity derived from the token when available.
|
||||||
|
Fall back to config metadata when the run originated from GitHub or when
|
||||||
|
Slack/Linear supplied an explicit user name and email.
|
||||||
|
"""
|
||||||
|
|
||||||
|
return _identity_from_github_token(github_token) or _identity_from_config(config)
|
||||||
|
|
||||||
|
|
||||||
|
def add_user_coauthor_trailer(
|
||||||
|
commit_message: str,
|
||||||
|
identity: CollaboratorIdentity | None,
|
||||||
|
) -> str:
|
||||||
|
"""Append a Co-authored-by trailer when a user identity is available."""
|
||||||
|
normalized_message = commit_message.rstrip()
|
||||||
|
if not identity:
|
||||||
|
return normalized_message
|
||||||
|
|
||||||
|
trailer = f"Co-authored-by: {identity.commit_name} <{identity.commit_email}>"
|
||||||
|
if trailer in normalized_message:
|
||||||
|
return normalized_message
|
||||||
|
return f"{normalized_message}\n\n{trailer}"
|
||||||
|
|
||||||
|
|
||||||
|
def add_pr_collaboration_note(
|
||||||
|
pr_body: str,
|
||||||
|
identity: CollaboratorIdentity | None,
|
||||||
|
) -> str:
|
||||||
|
"""Append a best-effort PR attribution note.
|
||||||
|
|
||||||
|
GitHub supports commit co-authors, but not PR co-authors. This note makes
|
||||||
|
the collaboration explicit in the automatically-opened PR body.
|
||||||
|
"""
|
||||||
|
|
||||||
|
normalized_body = pr_body.rstrip()
|
||||||
|
if not identity:
|
||||||
|
return normalized_body
|
||||||
|
|
||||||
|
note = f"_Opened collaboratively by {identity.display_name} and open-swe._"
|
||||||
|
if note in normalized_body:
|
||||||
|
return normalized_body
|
||||||
|
if not normalized_body:
|
||||||
|
return note
|
||||||
|
return f"{normalized_body}\n\n{note}"
|
||||||
|
|
@ -1168,6 +1168,7 @@ async def _trigger_or_queue_run(
|
||||||
prompt: str,
|
prompt: str,
|
||||||
*,
|
*,
|
||||||
github_login: str,
|
github_login: str,
|
||||||
|
github_user_id: int | None,
|
||||||
repo_config: dict[str, str],
|
repo_config: dict[str, str],
|
||||||
pr_number: int,
|
pr_number: int,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
|
@ -1188,6 +1189,7 @@ async def _trigger_or_queue_run(
|
||||||
"configurable": {
|
"configurable": {
|
||||||
"source": "github",
|
"source": "github",
|
||||||
"github_login": github_login,
|
"github_login": github_login,
|
||||||
|
"github_user_id": github_user_id,
|
||||||
"repo": repo_config,
|
"repo": repo_config,
|
||||||
"pr_number": pr_number,
|
"pr_number": pr_number,
|
||||||
},
|
},
|
||||||
|
|
@ -1251,6 +1253,7 @@ async def process_github_pr_comment(payload: dict[str, Any], event_type: str) ->
|
||||||
comment_id,
|
comment_id,
|
||||||
node_id,
|
node_id,
|
||||||
) = await extract_pr_context(payload, event_type)
|
) = await extract_pr_context(payload, event_type)
|
||||||
|
github_user_id = payload.get("sender", {}).get("id")
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Processing GitHub PR comment: event=%s, pr=%s, branch=%s",
|
"Processing GitHub PR comment: event=%s, pr=%s, branch=%s",
|
||||||
|
|
@ -1319,6 +1322,7 @@ async def process_github_pr_comment(payload: dict[str, Any], event_type: str) ->
|
||||||
thread_id,
|
thread_id,
|
||||||
prompt,
|
prompt,
|
||||||
github_login=github_login,
|
github_login=github_login,
|
||||||
|
github_user_id=github_user_id,
|
||||||
repo_config=repo_config,
|
repo_config=repo_config,
|
||||||
pr_number=pr_number,
|
pr_number=pr_number,
|
||||||
)
|
)
|
||||||
|
|
@ -1336,6 +1340,7 @@ async def process_github_issue(payload: dict[str, Any], event_type: str) -> None
|
||||||
issue_id = str(issue.get("id", ""))
|
issue_id = str(issue.get("id", ""))
|
||||||
issue_number = issue.get("number")
|
issue_number = issue.get("number")
|
||||||
github_login = payload.get("sender", {}).get("login", "")
|
github_login = payload.get("sender", {}).get("login", "")
|
||||||
|
github_user_id = payload.get("sender", {}).get("id")
|
||||||
issue_url = issue.get("html_url", "") or issue.get("url", "")
|
issue_url = issue.get("html_url", "") or issue.get("url", "")
|
||||||
title = issue.get("title", "No title")
|
title = issue.get("title", "No title")
|
||||||
description = issue.get("body") or "No description"
|
description = issue.get("body") or "No description"
|
||||||
|
|
@ -1414,6 +1419,7 @@ async def process_github_issue(payload: dict[str, Any], event_type: str) -> None
|
||||||
configurable: dict[str, Any] = {
|
configurable: dict[str, Any] = {
|
||||||
"source": "github",
|
"source": "github",
|
||||||
"github_login": github_login,
|
"github_login": github_login,
|
||||||
|
"github_user_id": github_user_id,
|
||||||
"repo": repo_config,
|
"repo": repo_config,
|
||||||
"github_issue": {
|
"github_issue": {
|
||||||
"id": issue_id,
|
"id": issue_id,
|
||||||
|
|
|
||||||
1
scripts/__init__.py
Normal file
1
scripts/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
"""Utility scripts for working with Open SWE thread and PR data."""
|
||||||
185
scripts/check_pr_merge_status.py
Normal file
185
scripts/check_pr_merge_status.py
Normal file
|
|
@ -0,0 +1,185 @@
|
||||||
|
"""Check merge status counts for PR URLs exported from LangGraph threads."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
|
import httpx
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
DEFAULT_INPUT_PATH = "pr_urls.json"
|
||||||
|
DEFAULT_CONCURRENCY = 20
|
||||||
|
GITHUB_API_VERSION = "2022-11-28"
|
||||||
|
|
||||||
|
|
||||||
|
def _load_dotenv_if_available() -> None:
|
||||||
|
try:
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
except ImportError:
|
||||||
|
return
|
||||||
|
load_dotenv()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class PullRequestRef:
|
||||||
|
owner: str
|
||||||
|
repo: str
|
||||||
|
number: int
|
||||||
|
url: str
|
||||||
|
|
||||||
|
|
||||||
|
def parse_github_pr_url(pr_url: str) -> PullRequestRef:
|
||||||
|
parsed_url = urlparse(pr_url)
|
||||||
|
if parsed_url.scheme not in {"http", "https"}:
|
||||||
|
raise ValueError(f"Unsupported PR URL scheme: {pr_url}")
|
||||||
|
if parsed_url.netloc not in {"github.com", "www.github.com"}:
|
||||||
|
raise ValueError(f"Unsupported PR URL host: {pr_url}")
|
||||||
|
|
||||||
|
path_parts = [part for part in parsed_url.path.split("/") if part]
|
||||||
|
if len(path_parts) < 4 or path_parts[2] != "pull":
|
||||||
|
raise ValueError(f"Unsupported GitHub PR URL path: {pr_url}")
|
||||||
|
|
||||||
|
try:
|
||||||
|
number = int(path_parts[3])
|
||||||
|
except ValueError as exc:
|
||||||
|
raise ValueError(f"Invalid GitHub PR number in URL: {pr_url}") from exc
|
||||||
|
|
||||||
|
return PullRequestRef(
|
||||||
|
owner=path_parts[0],
|
||||||
|
repo=path_parts[1],
|
||||||
|
number=number,
|
||||||
|
url=pr_url,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def load_pr_urls(input_path: Path) -> list[str]:
|
||||||
|
payload = json.loads(input_path.read_text(encoding="utf-8"))
|
||||||
|
if not isinstance(payload, list):
|
||||||
|
raise ValueError(f"Expected {input_path} to contain a JSON array of PR URLs")
|
||||||
|
|
||||||
|
unique_urls: list[str] = []
|
||||||
|
seen_urls: set[str] = set()
|
||||||
|
for item in payload:
|
||||||
|
if not isinstance(item, str) or not item:
|
||||||
|
raise ValueError(f"Expected every item in {input_path} to be a non-empty string")
|
||||||
|
if item not in seen_urls:
|
||||||
|
seen_urls.add(item)
|
||||||
|
unique_urls.append(item)
|
||||||
|
return unique_urls
|
||||||
|
|
||||||
|
|
||||||
|
def classify_pr_state(pr_payload: dict[str, Any]) -> str:
|
||||||
|
if pr_payload.get("merged") or pr_payload.get("merged_at"):
|
||||||
|
return "merged"
|
||||||
|
|
||||||
|
state = pr_payload.get("state")
|
||||||
|
if state == "open":
|
||||||
|
return "open_or_draft"
|
||||||
|
if state == "closed":
|
||||||
|
return "closed"
|
||||||
|
|
||||||
|
raise ValueError(f"Unsupported GitHub PR state: {state!r}")
|
||||||
|
|
||||||
|
|
||||||
|
async def _fetch_pr_state(
|
||||||
|
http_client: httpx.AsyncClient,
|
||||||
|
pr_ref: PullRequestRef,
|
||||||
|
github_pat: str,
|
||||||
|
semaphore: asyncio.Semaphore,
|
||||||
|
) -> str:
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Bearer {github_pat}",
|
||||||
|
"Accept": "application/vnd.github+json",
|
||||||
|
"X-GitHub-Api-Version": GITHUB_API_VERSION,
|
||||||
|
}
|
||||||
|
|
||||||
|
async with semaphore:
|
||||||
|
response = await http_client.get(
|
||||||
|
f"https://api.github.com/repos/{pr_ref.owner}/{pr_ref.repo}/pulls/{pr_ref.number}",
|
||||||
|
headers=headers,
|
||||||
|
)
|
||||||
|
|
||||||
|
if response.status_code != 200: # noqa: PLR2004
|
||||||
|
raise RuntimeError(
|
||||||
|
f"GitHub API returned {response.status_code} for {pr_ref.url}: {response.text}"
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
if not isinstance(payload, dict):
|
||||||
|
raise RuntimeError(f"Unexpected GitHub API response for {pr_ref.url}")
|
||||||
|
return classify_pr_state(payload)
|
||||||
|
|
||||||
|
|
||||||
|
async def summarize_pr_statuses(
|
||||||
|
*,
|
||||||
|
pr_urls: list[str],
|
||||||
|
github_pat: str,
|
||||||
|
concurrency: int = DEFAULT_CONCURRENCY,
|
||||||
|
) -> dict[str, int]:
|
||||||
|
semaphore = asyncio.Semaphore(concurrency)
|
||||||
|
async with httpx.AsyncClient(timeout=30.0) as http_client:
|
||||||
|
tasks = [
|
||||||
|
_fetch_pr_state(http_client, parse_github_pr_url(pr_url), github_pat, semaphore)
|
||||||
|
for pr_url in pr_urls
|
||||||
|
]
|
||||||
|
states = await asyncio.gather(*tasks)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"total_prs": len(pr_urls),
|
||||||
|
"total_merged_prs": sum(1 for state in states if state == "merged"),
|
||||||
|
"total_open_draft_prs": sum(1 for state in states if state == "open_or_draft"),
|
||||||
|
"total_closed_prs": sum(1 for state in states if state == "closed"),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser(description="Check merge status for GitHub PR URLs.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--input",
|
||||||
|
default=DEFAULT_INPUT_PATH,
|
||||||
|
help=f"Path to the input JSON file. Defaults to {DEFAULT_INPUT_PATH!r}.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--concurrency",
|
||||||
|
type=int,
|
||||||
|
default=DEFAULT_CONCURRENCY,
|
||||||
|
help=f"Concurrent GitHub API requests. Defaults to {DEFAULT_CONCURRENCY}.",
|
||||||
|
)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
_load_dotenv_if_available()
|
||||||
|
logging.basicConfig(level=logging.INFO, format="%(message)s")
|
||||||
|
|
||||||
|
args = parse_args()
|
||||||
|
github_pat = os.environ.get("GITHUB_PAT")
|
||||||
|
if not github_pat:
|
||||||
|
raise RuntimeError("GITHUB_PAT must be set")
|
||||||
|
|
||||||
|
pr_urls = load_pr_urls(Path(args.input))
|
||||||
|
summary = asyncio.run(
|
||||||
|
summarize_pr_statuses(
|
||||||
|
pr_urls=pr_urls,
|
||||||
|
github_pat=github_pat,
|
||||||
|
concurrency=args.concurrency,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info("total PRs: %d", summary["total_prs"])
|
||||||
|
logger.info("total merged PRs: %d", summary["total_merged_prs"])
|
||||||
|
logger.info("total open/draft PRs: %d", summary["total_open_draft_prs"])
|
||||||
|
logger.info("total closed PRs: %d", summary["total_closed_prs"])
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
321
scripts/export_pr_urls.py
Normal file
321
scripts/export_pr_urls.py
Normal file
|
|
@ -0,0 +1,321 @@
|
||||||
|
"""Export unique PR URLs from commit_and_open_pr tool messages in LangGraph threads."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
from collections.abc import Iterator
|
||||||
|
from datetime import UTC, datetime, timedelta
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from langchain_core.messages import BaseMessage, convert_to_messages
|
||||||
|
from langgraph_sdk import get_client
|
||||||
|
from langgraph_sdk.client import LangGraphClient
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
DEFAULT_OUTPUT_PATH = "pr_urls.json"
|
||||||
|
DEFAULT_PAGE_SIZE = 100
|
||||||
|
DEFAULT_CONCURRENCY = 20
|
||||||
|
DEFAULT_DAYS_BACK = 9
|
||||||
|
|
||||||
|
|
||||||
|
def _load_dotenv_if_available() -> None:
|
||||||
|
try:
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
except ImportError:
|
||||||
|
return
|
||||||
|
load_dotenv()
|
||||||
|
|
||||||
|
|
||||||
|
def get_langgraph_url(explicit_url: str | None = None) -> str:
|
||||||
|
if explicit_url:
|
||||||
|
return explicit_url
|
||||||
|
return os.environ.get("LANGGRAPH_URL") or os.environ.get(
|
||||||
|
"LANGGRAPH_URL_PROD", "http://localhost:2024"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def extract_pr_urls_from_messages(messages: list[BaseMessage]) -> list[str]:
|
||||||
|
pr_urls: list[str] = []
|
||||||
|
|
||||||
|
for message in messages:
|
||||||
|
if getattr(message, "type", None) != "tool":
|
||||||
|
continue
|
||||||
|
if getattr(message, "name", None) != "commit_and_open_pr":
|
||||||
|
continue
|
||||||
|
|
||||||
|
content = getattr(message, "content", None)
|
||||||
|
payload: dict[str, Any] | None = None
|
||||||
|
if isinstance(content, str):
|
||||||
|
try:
|
||||||
|
parsed_content = json.loads(content)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
continue
|
||||||
|
if isinstance(parsed_content, dict):
|
||||||
|
payload = parsed_content
|
||||||
|
elif isinstance(content, dict):
|
||||||
|
payload = content
|
||||||
|
|
||||||
|
if not payload:
|
||||||
|
continue
|
||||||
|
|
||||||
|
pr_url = payload.get("pr_url")
|
||||||
|
if isinstance(pr_url, str) and pr_url:
|
||||||
|
pr_urls.append(pr_url)
|
||||||
|
|
||||||
|
return pr_urls
|
||||||
|
|
||||||
|
|
||||||
|
def extract_pr_urls_from_state_values(state_values: Any) -> list[str]:
|
||||||
|
if not isinstance(state_values, dict):
|
||||||
|
return []
|
||||||
|
|
||||||
|
raw_messages = state_values.get("messages")
|
||||||
|
if not isinstance(raw_messages, list):
|
||||||
|
return []
|
||||||
|
|
||||||
|
try:
|
||||||
|
messages = convert_to_messages(raw_messages)
|
||||||
|
except Exception: # noqa: BLE001
|
||||||
|
logger.exception("Failed to deserialize messages from thread state")
|
||||||
|
raise ValueError("Failed to deserialize messages from thread state") from None
|
||||||
|
|
||||||
|
return extract_pr_urls_from_messages(messages)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_thread_id(thread: Any) -> str | None:
|
||||||
|
if isinstance(thread, dict):
|
||||||
|
thread_id = thread.get("thread_id")
|
||||||
|
else:
|
||||||
|
thread_id = getattr(thread, "thread_id", None)
|
||||||
|
return thread_id if isinstance(thread_id, str) and thread_id else None
|
||||||
|
|
||||||
|
|
||||||
|
def _coerce_datetime(value: Any) -> datetime | None:
|
||||||
|
if isinstance(value, datetime):
|
||||||
|
if value.tzinfo is None:
|
||||||
|
return value.replace(tzinfo=UTC)
|
||||||
|
return value.astimezone(UTC)
|
||||||
|
|
||||||
|
if isinstance(value, str):
|
||||||
|
try:
|
||||||
|
parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||||
|
except ValueError:
|
||||||
|
return None
|
||||||
|
if parsed.tzinfo is None:
|
||||||
|
return parsed.replace(tzinfo=UTC)
|
||||||
|
return parsed.astimezone(UTC)
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _get_thread_created_at(thread: Any) -> datetime | None:
|
||||||
|
if isinstance(thread, dict):
|
||||||
|
created_at = thread.get("created_at")
|
||||||
|
else:
|
||||||
|
created_at = getattr(thread, "created_at", None)
|
||||||
|
return _coerce_datetime(created_at)
|
||||||
|
|
||||||
|
|
||||||
|
def _split_recent_threads(threads: list[Any], cutoff: datetime) -> tuple[list[Any], bool]:
|
||||||
|
recent_threads: list[Any] = []
|
||||||
|
|
||||||
|
for thread in threads:
|
||||||
|
created_at = _get_thread_created_at(thread)
|
||||||
|
if created_at is None:
|
||||||
|
logger.warning(
|
||||||
|
"Skipping thread %s because created_at is missing or invalid",
|
||||||
|
_get_thread_id(thread) or "<unknown>",
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
if created_at >= cutoff:
|
||||||
|
recent_threads.append(thread)
|
||||||
|
continue
|
||||||
|
return recent_threads, True
|
||||||
|
|
||||||
|
return recent_threads, False
|
||||||
|
|
||||||
|
|
||||||
|
def _iter_offset_batches(
|
||||||
|
total_threads: int, page_size: int, batch_size: int
|
||||||
|
) -> Iterator[list[int]]:
|
||||||
|
offsets = range(0, total_threads, page_size)
|
||||||
|
batch: list[int] = []
|
||||||
|
|
||||||
|
for offset in offsets:
|
||||||
|
batch.append(offset)
|
||||||
|
if len(batch) == batch_size:
|
||||||
|
yield batch
|
||||||
|
batch = []
|
||||||
|
|
||||||
|
if batch:
|
||||||
|
yield batch
|
||||||
|
|
||||||
|
|
||||||
|
async def _fetch_thread_page(
|
||||||
|
client: LangGraphClient,
|
||||||
|
*,
|
||||||
|
offset: int,
|
||||||
|
page_size: int,
|
||||||
|
) -> tuple[int, list[Any]]:
|
||||||
|
threads = await client.threads.search(
|
||||||
|
limit=page_size,
|
||||||
|
offset=offset,
|
||||||
|
sort_by="created_at",
|
||||||
|
sort_order="desc",
|
||||||
|
)
|
||||||
|
return offset, threads
|
||||||
|
|
||||||
|
|
||||||
|
async def _fetch_pr_urls_for_thread(
|
||||||
|
client: LangGraphClient,
|
||||||
|
thread_id: str,
|
||||||
|
semaphore: asyncio.Semaphore,
|
||||||
|
) -> list[str]:
|
||||||
|
async with semaphore:
|
||||||
|
try:
|
||||||
|
state = await client.threads.get_state(thread_id)
|
||||||
|
except Exception: # noqa: BLE001
|
||||||
|
logger.exception("Failed to fetch state for thread %s", thread_id)
|
||||||
|
return []
|
||||||
|
|
||||||
|
return extract_pr_urls_from_state_values(state.get("values"))
|
||||||
|
|
||||||
|
|
||||||
|
async def export_pr_urls(
|
||||||
|
*,
|
||||||
|
langgraph_url: str,
|
||||||
|
output_path: Path,
|
||||||
|
page_size: int = DEFAULT_PAGE_SIZE,
|
||||||
|
concurrency: int = DEFAULT_CONCURRENCY,
|
||||||
|
days_back: int = DEFAULT_DAYS_BACK,
|
||||||
|
) -> list[str]:
|
||||||
|
if page_size < 1:
|
||||||
|
raise ValueError("page_size must be greater than 0")
|
||||||
|
if concurrency < 1:
|
||||||
|
raise ValueError("concurrency must be greater than 0")
|
||||||
|
if days_back < 1:
|
||||||
|
raise ValueError("days_back must be greater than 0")
|
||||||
|
|
||||||
|
api_key = os.environ.get("LANGGRAPH_API_KEY")
|
||||||
|
client = get_client(url=langgraph_url, api_key=api_key)
|
||||||
|
try:
|
||||||
|
total_threads = await client.threads.count()
|
||||||
|
cutoff = datetime.now(UTC) - timedelta(days=days_back)
|
||||||
|
logger.info(
|
||||||
|
"Scanning threads from %s created on or after %s",
|
||||||
|
langgraph_url,
|
||||||
|
cutoff.isoformat(),
|
||||||
|
)
|
||||||
|
|
||||||
|
state_semaphore = asyncio.Semaphore(concurrency)
|
||||||
|
unique_pr_urls: set[str] = set()
|
||||||
|
recent_threads_count = 0
|
||||||
|
|
||||||
|
for offset_batch in _iter_offset_batches(total_threads, page_size, concurrency):
|
||||||
|
page_results = await asyncio.gather(
|
||||||
|
*[
|
||||||
|
_fetch_thread_page(client, offset=offset, page_size=page_size)
|
||||||
|
for offset in offset_batch
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
thread_ids: list[str] = []
|
||||||
|
saw_older_thread = False
|
||||||
|
for _offset, threads in sorted(page_results, key=lambda result: result[0]):
|
||||||
|
if not threads:
|
||||||
|
continue
|
||||||
|
|
||||||
|
recent_threads, saw_older_thread = _split_recent_threads(threads, cutoff)
|
||||||
|
recent_threads_count += len(recent_threads)
|
||||||
|
|
||||||
|
for thread in recent_threads:
|
||||||
|
thread_id = _get_thread_id(thread)
|
||||||
|
if thread_id:
|
||||||
|
thread_ids.append(thread_id)
|
||||||
|
|
||||||
|
if saw_older_thread:
|
||||||
|
break
|
||||||
|
|
||||||
|
for pr_urls in await asyncio.gather(
|
||||||
|
*[
|
||||||
|
_fetch_pr_urls_for_thread(client, thread_id, state_semaphore)
|
||||||
|
for thread_id in thread_ids
|
||||||
|
]
|
||||||
|
):
|
||||||
|
unique_pr_urls.update(pr_urls)
|
||||||
|
|
||||||
|
logger.info("Processed %d recent thread(s)", recent_threads_count)
|
||||||
|
|
||||||
|
if saw_older_thread:
|
||||||
|
break
|
||||||
|
|
||||||
|
sorted_pr_urls = sorted(unique_pr_urls)
|
||||||
|
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
output_path.write_text(f"{json.dumps(sorted_pr_urls, indent=2)}\n", encoding="utf-8")
|
||||||
|
logger.info("Total threads in deployment: %d", total_threads)
|
||||||
|
logger.info("Threads in last %d days: %d", days_back, recent_threads_count)
|
||||||
|
logger.info("Wrote %d unique PR URL(s) to %s", len(sorted_pr_urls), output_path)
|
||||||
|
return sorted_pr_urls
|
||||||
|
finally:
|
||||||
|
await client.aclose()
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> argparse.Namespace:
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Export unique PR URLs from commit_and_open_pr tool messages."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--output",
|
||||||
|
default=DEFAULT_OUTPUT_PATH,
|
||||||
|
help=f"Path to the output JSON file. Defaults to {DEFAULT_OUTPUT_PATH!r}.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--langgraph-url",
|
||||||
|
default=None,
|
||||||
|
help="LangGraph deployment URL. Defaults to LANGGRAPH_URL or LANGGRAPH_URL_PROD.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--page-size",
|
||||||
|
type=int,
|
||||||
|
default=DEFAULT_PAGE_SIZE,
|
||||||
|
help=f"Threads to fetch per page. Defaults to {DEFAULT_PAGE_SIZE}.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--concurrency",
|
||||||
|
type=int,
|
||||||
|
default=DEFAULT_CONCURRENCY,
|
||||||
|
help=f"Concurrent LangGraph page/state requests per batch. Defaults to {DEFAULT_CONCURRENCY}.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--days-back",
|
||||||
|
type=int,
|
||||||
|
default=DEFAULT_DAYS_BACK,
|
||||||
|
help=f"Only include threads created in the last N days. Defaults to {DEFAULT_DAYS_BACK}.",
|
||||||
|
)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
_load_dotenv_if_available()
|
||||||
|
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
|
||||||
|
|
||||||
|
args = parse_args()
|
||||||
|
asyncio.run(
|
||||||
|
export_pr_urls(
|
||||||
|
langgraph_url=get_langgraph_url(args.langgraph_url),
|
||||||
|
output_path=Path(args.output),
|
||||||
|
page_size=args.page_size,
|
||||||
|
concurrency=args.concurrency,
|
||||||
|
days_back=args.days_back,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
Loading…
Add table
Reference in a new issue