Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 29 additions & 14 deletions scripts/stubsabot.py
Original file line number Diff line number Diff line change
Expand Up @@ -179,10 +179,14 @@ def __str__(self) -> str:
@dataclass
class NoUpdate:
distribution: str
reason: Literal["obsolete", "no longer updated", "up to date"]
reason: Literal["obsolete", "no longer updated", "up to date", "pr open"]
pr_number: int = 0

def __str__(self) -> str:
return f"{colored('skipping', 'green')} ({self.reason})"
if self.reason == "pr open":
return f"{colored('skipping', 'green')} ({self.reason} {self.pr_number})"
else:
return f"{colored('skipping', 'green')} ({self.reason})"


@dataclass
Expand Down Expand Up @@ -749,10 +753,9 @@ async def create_or_update_pull_request(*, title: str, body: str, branch_name: s
await update_pull_request_label(pr_number=pr_number, session=session)


async def update_existing_pull_request(*, title: str, body: str, branch_name: str, session: aiohttp.ClientSession) -> int:
async def find_existing_pr(*, branch_name: str, session: aiohttp.ClientSession) -> int:
fork_owner = get_origin_owner()

# Find the existing PR
async with session.get(
f"{TYPESHED_API_URL}/pulls",
params={"state": "open", "head": f"{fork_owner}:{branch_name}", "base": "main"},
Expand All @@ -763,6 +766,13 @@ async def update_existing_pull_request(*, title: str, body: str, branch_name: st
assert len(resp_json) >= 1
pr_number = resp_json[0]["number"]
assert isinstance(pr_number, int)

return pr_number


async def update_existing_pull_request(*, title: str, body: str, branch_name: str, session: aiohttp.ClientSession) -> int:
# Find the existing PR
pr_number = await find_existing_pr(branch_name=branch_name, session=session)
# Update the PR's title and body
async with session.patch(
f"{TYPESHED_API_URL}/pulls/{pr_number}", json={"title": title, "body": body}, headers=get_github_api_headers()
Expand Down Expand Up @@ -914,11 +924,12 @@ def run_action(action: Update | Obsolete | Remove) -> None:
remove_stubs(action.distribution)


async def process_typeshed_change(
action: Update | Obsolete | Remove, session: aiohttp.ClientSession, action_level: ActionLevel
) -> None:
_A = TypeVar("_A", bound=Update | Obsolete | Remove)


async def process_typeshed_change(action: _A, session: aiohttp.ClientSession, action_level: ActionLevel) -> _A | NoUpdate:
if action_level <= ActionLevel.nothing:
return
return action

async with _repo_lock:
branch_name = f"{BRANCH_PREFIX}/{normalize(action.distribution)}"
Expand All @@ -927,15 +938,16 @@ async def process_typeshed_change(
title, body = get_commit_message(action)
subprocess.check_call(["git", "commit", "--quiet", "--all", "-m", f"{title}\n\n{body}"])
if action_level <= ActionLevel.local:
return
return action
if not latest_commit_is_different_to_last_commit_on_origin(branch_name):
print(f"No pushing to origin required: 'origin/{branch_name}' exists and requires no changes!")
return
pr_number = await find_existing_pr(branch_name=branch_name, session=session)
return NoUpdate(action.distribution, "pr open", pr_number)
somewhat_safe_force_push(branch_name)
if action_level <= ActionLevel.fork:
return
return action

await create_or_update_pull_request(title=title, body=body, branch_name=branch_name, session=session)
return action


async def main() -> int:
Expand Down Expand Up @@ -1001,21 +1013,24 @@ async def main() -> int:
for task in asyncio.as_completed(tasks):
action = await task
print(f"{action.distribution}... ", end="")
print(action)

if isinstance(action, NoUpdate):
print(action)
continue
if isinstance(action, Error):
print(action)
error = True
continue

if args.action_count_limit is not None and action_count >= args.action_count_limit:
print(action)
print(colored("... but we've reached action count limit", "red"))
continue
action_count += 1

try:
await process_typeshed_change(action, session, action_level=args.action_level)
action_result = await process_typeshed_change(action, session, action_level=args.action_level)
print(action_result)
continue
except RemoteConflictError as e:
print(colored(f"... but ran into {type(e).__qualname__}: {e}", "red"))
Expand Down