fix: Add scripts for getting usage, fix dual committing (#1131)

This commit is contained in:
Brace Sproul 2026-03-25 13:39:33 -07:00 • committed by GitHub
parent 0d8647c2cd
commit 87968ab813
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 707 additions and 9 deletions

View file

@ -16,6 +16,13 @@ from langchain.agents.middleware import AgentState, after_agent
from langgraph.config import get_config
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 (
create_github_pr,
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_body = pr_payload.get("body", "Automated PR created by Open SWE agent.")
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:
raise ValueError("No thread_id found in config")
@ -135,14 +148,12 @@ async def open_pr_if_needed(
git_config_user,
sandbox_backend,
repo_dir,
"open-swe[bot]",
"open-swe@users.noreply.github.com",
OPEN_SWE_BOT_NAME,
OPEN_SWE_BOT_EMAIL,
)
await asyncio.to_thread(git_add_all, sandbox_backend, repo_dir)
await asyncio.to_thread(git_commit, sandbox_backend, repo_dir, commit_message)
github_token = get_github_token()
if github_token:
await asyncio.to_thread(
git_push, sandbox_backend, repo_dir, target_branch, github_token

View file

@ -4,6 +4,13 @@ from typing import Any
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 (
create_github_pr,
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}
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)
git_fetch_origin(sandbox_backend, repo_dir)
@ -163,12 +173,12 @@ def commit_and_open_pr(
git_config_user(
sandbox_backend,
repo_dir,
"open-swe[bot]",
"open-swe@users.noreply.github.com",
OPEN_SWE_BOT_NAME,
OPEN_SWE_BOT_EMAIL,
)
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:
commit_result = git_commit(sandbox_backend, repo_dir, commit_msg)
if commit_result.exit_code != 0:
@ -178,7 +188,6 @@ def commit_and_open_pr(
"pr_url": None,
}
github_token = get_github_token()
if not github_token:
logger.error("commit_and_open_pr missing GitHub token for thread %s", thread_id)
return {
@ -204,7 +213,7 @@ def commit_and_open_pr(
title=title,
head_branch=target_branch,
base_branch=base_branch,
body=body,
body=pr_body,
)
)

165
agent/utils/authorship.py Normal file
View 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}"

View file

@ -1168,6 +1168,7 @@ async def _trigger_or_queue_run(
prompt: str,
*,
github_login: str,
github_user_id: int | None,
repo_config: dict[str, str],
pr_number: int,
) -> None:
@ -1188,6 +1189,7 @@ async def _trigger_or_queue_run(
"configurable": {
"source": "github",
"github_login": github_login,
"github_user_id": github_user_id,
"repo": repo_config,
"pr_number": pr_number,
},
@ -1251,6 +1253,7 @@ async def process_github_pr_comment(payload: dict[str, Any], event_type: str) ->
comment_id,
node_id,
) = await extract_pr_context(payload, event_type)
github_user_id = payload.get("sender", {}).get("id")
logger.info(
"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,
prompt,
github_login=github_login,
github_user_id=github_user_id,
repo_config=repo_config,
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_number = issue.get("number")
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", "")
title = issue.get("title", "No title")
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] = {
"source": "github",
"github_login": github_login,
"github_user_id": github_user_id,
"repo": repo_config,
"github_issue": {
"id": issue_id,

1
scripts/__init__.py Normal file
View file

@ -0,0 +1 @@
"""Utility scripts for working with Open SWE thread and PR data."""

View 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
View 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()