mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 06:53:14 +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.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
|
||||
|
|
|
|||
|
|
@ -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
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,
|
||||
*,
|
||||
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
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