patch-via-github 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,965 @@
1
+ #!/usr/bin/env python3
2
+
3
+ """
4
+ Basic program to apply a set of patches from various GitHub pull requests
5
+ based on user criteria to a repo sync area
6
+
7
+ PRs are applied in the order requested; PRs found by label are applied in
8
+ the order the labels were given, and in PR number order for each label
9
+ """
10
+
11
+ import argparse
12
+ import configparser
13
+ import contextlib
14
+ import logging
15
+ import os
16
+ import re
17
+ import subprocess
18
+ import sys
19
+ import time
20
+ import xml.etree.ElementTree as EleTree
21
+ from importlib.metadata import version
22
+ from shutil import which
23
+
24
+ import requests
25
+ from requests.adapters import HTTPAdapter
26
+ from requests.auth import AuthBase
27
+ from urllib3.util.retry import Retry
28
+
29
+
30
+ logger = logging.getLogger('patch_via_github')
31
+
32
+ SHA_RE = re.compile(r'^[0-9a-f]{40}$')
33
+ # GitHub caps a PR's commit listing at this many commits
34
+ MAX_LISTED_COMMITS = 250
35
+ # Seconds
36
+ FETCH_TIMEOUT = 10 * 60
37
+ REPO_MANIFEST_TIMEOUT = 5 * 60
38
+
39
+
40
+ class PatchError(RuntimeError):
41
+ """A failure reported to the user before exiting"""
42
+
43
+ def __init__(self, message, exit_code=1):
44
+ super().__init__(message)
45
+ self.exit_code = exit_code
46
+
47
+
48
+ def pr_key(repo_full_name, number):
49
+ """Identity of a PR across repos; GitHub names are case-insensitive"""
50
+ return f'{repo_full_name.lower()}#{int(number)}'
51
+
52
+
53
+ @contextlib.contextmanager
54
+ def api_data(what):
55
+ """Report malformed GitHub API data as a PatchError"""
56
+ try:
57
+ yield
58
+ except (KeyError, TypeError) as exc:
59
+ raise PatchError(
60
+ f'GitHub API error: unexpected {what} '
61
+ f'(missing or malformed field {exc})'
62
+ ) from exc
63
+
64
+
65
+ def run_git(cwd, *args, timeout=None, text=True, input=None):
66
+ """Run git in cwd, capturing its output"""
67
+
68
+ logger.debug(f' git {" ".join(args)} (in {cwd})')
69
+ try:
70
+ return subprocess.run(
71
+ ['git', *args], cwd=cwd, capture_output=True, text=text,
72
+ input=input, timeout=timeout
73
+ )
74
+ except OSError as exc:
75
+ raise PatchError(f'Could not run git in {cwd}: {exc}') from exc
76
+ except subprocess.TimeoutExpired as exc:
77
+ raise PatchError(
78
+ f'`git {" ".join(args)}` in {cwd} timed out after {timeout}s'
79
+ ) from exc
80
+
81
+
82
+ def repo_tool():
83
+ repo = which('repo')
84
+ if repo is None:
85
+ raise PatchError("'repo' tool not found on PATH")
86
+ return repo
87
+
88
+
89
+ class BearerAuth(AuthBase):
90
+ """Token auth; set as the session's auth so ~/.netrc can't override it"""
91
+
92
+ def __init__(self, token):
93
+ self.token = token
94
+
95
+ def __call__(self, request):
96
+ request.headers['Authorization'] = f'Bearer {self.token}'
97
+ return request
98
+
99
+
100
+ class GitHubPR:
101
+ """Encapsulation of relevant information for a given GitHub pull request"""
102
+
103
+ def __init__(self, data, use_ssh=False):
104
+ """
105
+ Initialize with key information for a GitHub PR
106
+ """
107
+ self.number = int(data['number'])
108
+ self.repo_full_name = data['base']['repo']['full_name']
109
+ self.project = data['base']['repo']['name']
110
+ self.branch = data['base']['ref']
111
+ self.head_sha = data['head']['sha']
112
+ self.html_url = data['html_url']
113
+ # Only the single-PR endpoint reports this, not the PR listing
114
+ self.commit_count = data.get('commits')
115
+ self.ref = f'{self.repo_full_name}#{self.number}'
116
+ self.key = pr_key(self.repo_full_name, self.number)
117
+ # Populated by GitHubPatches.load_pr_commits before cherry-picking
118
+ self.commits = []
119
+ self.fetch_url = data['base']['repo'][
120
+ 'ssh_url' if use_ssh else 'clone_url'
121
+ ]
122
+
123
+
124
+ class GitHubPatches:
125
+ """
126
+ Determine all relevant patches to apply to a repo sync based on
127
+ a given set of initial parameters, which can be a set of one of
128
+ the following:
129
+ - PR references (org/repo#number or repo#number)
130
+ - labels (org/repo:label)
131
+
132
+ The resulting data will include the necessary patch commands to
133
+ be applied to the repo sync
134
+ """
135
+
136
+ GITHUB_API_URL = 'https://api.github.com'
137
+
138
+ def __init__(self, token, default_org=None, checkout=False, use_ssh=True):
139
+ """Initial GitHub connection and set base options"""
140
+
141
+ self.default_org = default_org
142
+ self.checkout = checkout
143
+ self.use_ssh = use_ssh
144
+ self.session = requests.Session()
145
+ self.session.headers.update({
146
+ 'Accept': 'application/vnd.github+json',
147
+ 'X-GitHub-Api-Version': '2022-11-28',
148
+ })
149
+ self.session.auth = BearerAuth(token)
150
+ retry = Retry(
151
+ total=3, backoff_factor=1, allowed_methods={'GET'},
152
+ status_forcelist={429, 500, 502, 503, 504},
153
+ retry_after_max=60, raise_on_status=False,
154
+ )
155
+ self.session.mount('https://', HTTPAdapter(max_retries=retry))
156
+
157
+ # pr_key() of each PR applied, in order
158
+ self.applied_prs = []
159
+ # Read from 'repo manifest' on first use
160
+ self.manifest = None
161
+ # proj_path -> (branch or None if detached, HEAD sha) before this
162
+ # run first touched it, for rolling back on failure
163
+ self.originals = {}
164
+
165
+ @classmethod
166
+ def from_config_file(cls, config_path, default_org=None, checkout=False,
167
+ use_ssh=None):
168
+ """
169
+ Factory method: construct a GitHubPatches from the path to a
170
+ config file
171
+ """
172
+ config = configparser.ConfigParser()
173
+ try:
174
+ with open(config_path) as config_file:
175
+ config.read_file(config_file)
176
+ except FileNotFoundError:
177
+ raise PatchError(
178
+ f'Configuration file {config_path} missing!'
179
+ ) from None
180
+ except OSError as exc:
181
+ raise PatchError(
182
+ f'Could not read configuration file {config_path}: '
183
+ f'{exc.strerror}'
184
+ ) from None
185
+ except configparser.Error as exc:
186
+ # Not quoting exc: it includes the offending line, maybe the token
187
+ errors = (getattr(exc, 'errors', None)
188
+ or [(getattr(exc, 'lineno', '?'), None)])
189
+ lines = ', '.join(str(lineno) for lineno, _ in errors)
190
+ raise PatchError(
191
+ f'Invalid config file "{config_path}" (could not parse '
192
+ f'line {lines})'
193
+ ) from None
194
+
195
+ if 'main' not in config.sections():
196
+ raise PatchError(
197
+ f'Invalid config file "{config_path}" '
198
+ '(missing "main" section)'
199
+ )
200
+
201
+ token = config.get('main', 'token', fallback=None)
202
+ if token is None:
203
+ raise PatchError(
204
+ 'Required option "token" is missing from the config '
205
+ 'file. Aborting...'
206
+ )
207
+ if not token:
208
+ raise PatchError(
209
+ f'Option "token" in {config_path} is empty. Aborting...'
210
+ )
211
+
212
+ org = default_org or config.get('main', 'default_org', fallback=None)
213
+
214
+ ssh = use_ssh
215
+ if ssh is None:
216
+ try:
217
+ ssh = config.getboolean('main', 'ssh', fallback=True)
218
+ except ValueError:
219
+ raise PatchError(
220
+ f'Invalid value for "ssh" in {config_path}: '
221
+ f'"{config.get("main", "ssh")}" (expected true or false)'
222
+ ) from None
223
+
224
+ return cls(token, org, checkout, ssh)
225
+
226
+ def patch_repo_sync(self, refs, id_type):
227
+ """
228
+ Patch the repo sync with the list of patch commands. Repo
229
+ sync is presumed to be in current working directory.
230
+ """
231
+
232
+ try:
233
+ prs = self.resolve_prs(refs, id_type)
234
+ self.apply_prs(prs)
235
+ self.check_requested_prs_applied(id_type, refs, prs)
236
+ except (Exception, KeyboardInterrupt) as exc:
237
+ self._abort(exc)
238
+
239
+ def resolve_prs(self, refs, id_type):
240
+ """
241
+ From an initial set of PR references or labels, determine all
242
+ relevant PRs that will need to be applied to a repo sync
243
+ via patching
244
+ """
245
+
246
+ if id_type == 'pr':
247
+ prs = self._resolve_pr_refs(refs)
248
+ else:
249
+ prs = self._resolve_labels(refs)
250
+ logger.info(
251
+ f'Final list of PRs to apply: '
252
+ f'{", ".join(pr.ref for pr in prs.values())}'
253
+ )
254
+ return prs
255
+
256
+ def _resolve_pr_refs(self, pr_refs):
257
+ """Requested PRs are applied whatever their state or target branch"""
258
+
259
+ parsed = [self.parse_pr_reference(ref) for ref in pr_refs]
260
+ prs = {}
261
+ for org, repo, number in parsed:
262
+ try:
263
+ pr = self.get_pr(org, repo, number)
264
+ except PatchError as exc:
265
+ raise PatchError(
266
+ f'Failed to fetch {org}/{repo}#{number}: {exc}'
267
+ ) from exc
268
+ # GitHub follows repo renames, so pr.key is the canonical name
269
+ prs.setdefault(pr.key, pr)
270
+ return prs
271
+
272
+ def parse_pr_reference(self, pr_ref):
273
+ """
274
+ Parse a PR reference into (org, repo, number).
275
+ Supports formats:
276
+ - org/repo#number
277
+ - repo#number (uses default_org)
278
+ """
279
+
280
+ # Full format: org/repo#number
281
+ match = re.match(r'^([^/]+)/([^#]+)#(\d+)$', pr_ref)
282
+ if match:
283
+ return match.groups()
284
+
285
+ # Short format: repo#number
286
+ match = re.match(r'^([^#/]+)#(\d+)$', pr_ref)
287
+ if match:
288
+ if not self.default_org:
289
+ raise PatchError(
290
+ f'PR reference "{pr_ref}" needs an org, but no '
291
+ 'default_org configured. Use org/repo#number format '
292
+ 'or set default_org in config.'
293
+ )
294
+ return self.default_org, *match.groups()
295
+
296
+ raise PatchError(
297
+ f'Invalid PR reference: "{pr_ref}". '
298
+ 'Use format: org/repo#number or repo#number'
299
+ )
300
+
301
+ def _resolve_labels(self, label_refs):
302
+ """Open labelled PRs, if they target their project's manifest branch"""
303
+
304
+ parsed = []
305
+ for label_ref in label_refs:
306
+ match = re.match(r'^([^/]+)/([^:]+):(.+)$', label_ref)
307
+ if not match:
308
+ raise PatchError(
309
+ f'Invalid label reference: "{label_ref}". '
310
+ 'Use format: org/repo:label'
311
+ )
312
+ parsed.append((label_ref, *match.groups()))
313
+
314
+ prs = {}
315
+ for label_ref, org, repo, label in parsed:
316
+ try:
317
+ labelled = self.get_open_prs_by_label(org, repo, label)
318
+ except PatchError as exc:
319
+ raise PatchError(
320
+ f'Failed to fetch PRs for {label_ref}: {exc}'
321
+ ) from exc
322
+ for pr in sorted(labelled.values(), key=lambda pr: pr.number):
323
+ prs.setdefault(pr.key, pr)
324
+
325
+ for key, pr in list(prs.items()):
326
+ _, manifest_branch = (
327
+ self.get_project_path_and_branch_from_manifest(pr.project)
328
+ )
329
+ if (manifest_branch
330
+ and pr.branch != manifest_branch
331
+ and not SHA_RE.match(manifest_branch)):
332
+ logger.info(
333
+ f' Ignoring {pr.ref} because it targets {pr.branch}, '
334
+ f'manifest branch is {manifest_branch}'
335
+ )
336
+ del prs[key]
337
+ return prs
338
+
339
+ def get_pr(self, org, repo, number):
340
+ """Fetch a single PR from GitHub API"""
341
+
342
+ logger.debug(f'Fetching PR {org}/{repo}#{number}')
343
+ data = self._api_get(f'/repos/{org}/{repo}/pulls/{number}')
344
+ return self._make_pr(data, f'{org}/{repo}#{number}')
345
+
346
+ def get_open_prs_by_label(self, org, repo, label):
347
+ """Fetch all open PRs with a given label"""
348
+
349
+ logger.debug(
350
+ f'Fetching open PRs with label "{label}" from {org}/{repo}'
351
+ )
352
+ data = self._api_get_paginated(
353
+ f'/repos/{org}/{repo}/pulls?state=open&per_page=100'
354
+ )
355
+ with api_data(f'PR data from {org}/{repo}'):
356
+ labelled = [
357
+ pr_data for pr_data in data
358
+ if label in (lbl['name'] for lbl in pr_data['labels'])
359
+ ]
360
+ prs = {}
361
+ for pr_data in labelled:
362
+ pr = self._make_pr(pr_data, f'{org}/{repo}')
363
+ prs[pr.key] = pr
364
+ return prs
365
+
366
+ def _make_pr(self, data, source):
367
+ with api_data(f'PR data from {source}'):
368
+ return GitHubPR(data, use_ssh=self.use_ssh)
369
+
370
+ def get_project_path_and_branch_from_manifest(self, project):
371
+ """(path, revision) of project, or (None, None) if not in manifest"""
372
+
373
+ if self.manifest is None:
374
+ self._load_manifest()
375
+
376
+ proj_info = self.manifest.find(f'.//project[@name="{project}"]')
377
+ if proj_info is None:
378
+ return None, None
379
+
380
+ # Same precedence as repo: project, then its remote, then <default>
381
+ default = self.manifest.find('.//default')
382
+ defaults = {} if default is None else default.attrib
383
+ branch = proj_info.get('revision')
384
+ if branch is None:
385
+ remote_name = proj_info.get('remote', defaults.get('remote'))
386
+ remote = self.manifest.find(f'.//remote[@name="{remote_name}"]')
387
+ if remote is not None:
388
+ branch = remote.get('revision')
389
+ if branch is None:
390
+ branch = defaults.get('revision')
391
+ if branch is not None:
392
+ branch = branch.removeprefix('refs/heads/')
393
+ return proj_info.get('path', project), branch
394
+
395
+ def _load_manifest(self):
396
+ try:
397
+ manifest_str = subprocess.check_output(
398
+ [repo_tool(), 'manifest'], timeout=REPO_MANIFEST_TIMEOUT
399
+ )
400
+ except subprocess.CalledProcessError as exc:
401
+ raise PatchError(
402
+ f"'repo manifest' failed with exit code {exc.returncode}"
403
+ ) from exc
404
+ except subprocess.TimeoutExpired as exc:
405
+ raise PatchError(
406
+ f"'repo manifest' timed out after {exc.timeout}s"
407
+ ) from exc
408
+ except OSError as exc:
409
+ raise PatchError(f"Could not run 'repo manifest': {exc}") from exc
410
+ try:
411
+ self.manifest = EleTree.fromstring(manifest_str)
412
+ except EleTree.ParseError as exc:
413
+ raise PatchError(
414
+ f"Could not parse 'repo manifest' output: {exc}"
415
+ ) from exc
416
+
417
+ def apply_prs(self, prs):
418
+ """Applies the given PRs, in order"""
419
+
420
+ for pr in prs.values():
421
+ path, _ = self.get_project_path_and_branch_from_manifest(
422
+ pr.project
423
+ )
424
+ if path is None:
425
+ logger.info(
426
+ f"***** NOTE: ignoring PR {pr.ref} for project "
427
+ f"{pr.project} that is either not part of the "
428
+ f"manifest, or was excluded due to manifest "
429
+ f"group filters."
430
+ )
431
+ continue
432
+
433
+ self.apply_single_pr(pr, path)
434
+
435
+ def apply_single_pr(self, pr, proj_path):
436
+ """
437
+ Given a single PR object and a path, apply the git change to
438
+ that path (using either checkout or cherry-pick as requested).
439
+ """
440
+
441
+ if not os.path.exists(proj_path):
442
+ raise PatchError(
443
+ f'***** Project {pr.project} missing on disk! '
444
+ f'Expected to be in {proj_path}',
445
+ exit_code=5
446
+ )
447
+ self._record_original(proj_path)
448
+
449
+ logger.info(
450
+ f'***** Applying {pr.html_url} to project {pr.project}:'
451
+ )
452
+ if not self.checkout:
453
+ self.load_pr_commits(pr)
454
+
455
+ self._git(pr, proj_path, 'fetch', pr.fetch_url,
456
+ f'pull/{pr.number}/head', timeout=FETCH_TIMEOUT)
457
+ fetched = self._git(pr, proj_path, 'rev-parse', 'FETCH_HEAD').strip()
458
+ if fetched != pr.head_sha:
459
+ raise PatchError(
460
+ f'{pr.ref}: fetched {fetched} but GitHub reported head '
461
+ f'{pr.head_sha}; the PR was updated during this run, '
462
+ 're-run it'
463
+ )
464
+
465
+ if self.checkout:
466
+ self._git(pr, proj_path, 'checkout', 'FETCH_HEAD')
467
+ else:
468
+ self._cherry_pick_commits(pr, proj_path)
469
+ logger.info(
470
+ f'***** Done applying PR {pr.ref} '
471
+ f'to project {pr.project}\n'
472
+ )
473
+ self.applied_prs.append(pr.key)
474
+
475
+ def _record_original(self, proj_path):
476
+ if proj_path in self.originals:
477
+ return
478
+ head = run_git(proj_path, 'rev-parse', 'HEAD')
479
+ if head.returncode != 0:
480
+ raise PatchError(
481
+ f'{proj_path} is not a usable git checkout: '
482
+ f'{head.stderr.strip()}'
483
+ )
484
+ branch = run_git(
485
+ proj_path, 'symbolic-ref', '-q', '--short', 'HEAD'
486
+ ).stdout.strip() or None
487
+ self.originals[proj_path] = (branch, head.stdout.strip())
488
+
489
+ def load_pr_commits(self, pr):
490
+ """
491
+ Record the PR's own commits, oldest first, for cherry-picking.
492
+ Requires linear history, and a commit list that GitHub confirms
493
+ is complete
494
+ """
495
+
496
+ if pr.commit_count is None:
497
+ org, repo = pr.repo_full_name.split('/', 1)
498
+ latest = self.get_pr(org, repo, pr.number)
499
+ if latest.head_sha != pr.head_sha:
500
+ raise PatchError(
501
+ f'{pr.ref}: head moved from {pr.head_sha} to '
502
+ f'{latest.head_sha} during this run; re-run it'
503
+ )
504
+ pr.commit_count = latest.commit_count
505
+
506
+ if not isinstance(pr.commit_count, int):
507
+ raise PatchError(
508
+ f'GitHub API error: no commit count reported for {pr.ref}'
509
+ )
510
+ if pr.commit_count > MAX_LISTED_COMMITS:
511
+ raise PatchError(
512
+ f'{pr.ref} has {pr.commit_count} commits, but GitHub only '
513
+ f'lists {MAX_LISTED_COMMITS}, so they cannot all be '
514
+ 'cherry-picked. Use --checkout, or split the PR'
515
+ )
516
+
517
+ data = self._api_get_paginated(
518
+ f'/repos/{pr.repo_full_name}/pulls/{pr.number}/commits'
519
+ '?per_page=100'
520
+ )
521
+ with api_data(f'commit data for {pr.ref}'):
522
+ parents = {
523
+ commit['sha']: [parent['sha'] for parent in commit['parents']]
524
+ for commit in data
525
+ }
526
+
527
+ if len(parents) != pr.commit_count:
528
+ raise PatchError(
529
+ f'{pr.ref}: GitHub reports {pr.commit_count} commits but '
530
+ f'listed {len(parents)}; the PR may have been updated '
531
+ 'during this run, re-run it'
532
+ )
533
+ if pr.head_sha not in parents:
534
+ raise PatchError(
535
+ f'{pr.ref}: PR head {pr.head_sha} is not in the commit list '
536
+ 'from GitHub; the PR may have been updated during this run, '
537
+ 're-run it'
538
+ )
539
+
540
+ shas = []
541
+ sha = pr.head_sha
542
+ while sha in parents:
543
+ if len(parents[sha]) != 1:
544
+ raise PatchError(
545
+ f'{pr.ref}: commit {sha} is a merge commit; '
546
+ 'cherry-picking needs linear PR history. Rebase the PR '
547
+ 'or use --checkout'
548
+ )
549
+ shas.append(sha)
550
+ sha = parents[sha][0]
551
+
552
+ if len(shas) != len(parents):
553
+ stray = sorted(set(parents) - set(shas))
554
+ raise PatchError(
555
+ f'{pr.ref}: commits {", ".join(stray)} are not ancestors of '
556
+ f'the PR head {pr.head_sha}; the PR may have been updated '
557
+ 'during this run, re-run it'
558
+ )
559
+ pr.commits = shas[::-1]
560
+
561
+ def _cherry_pick_commits(self, pr, proj_path):
562
+ """
563
+ Cherry-pick the PR's commits one at a time, skipping any whose
564
+ changes are already present (e.g. from a stacked PR applied earlier)
565
+ """
566
+
567
+ present = self._present_prefix(pr, proj_path)
568
+ if present:
569
+ logger.info(
570
+ f' Skipped the first {present} commit(s), up to '
571
+ f'{pr.commits[present - 1]}: their changes are already '
572
+ 'present (e.g. from a stacked PR applied earlier)'
573
+ )
574
+
575
+ picked = 0
576
+ for sha in pr.commits[present:]:
577
+ result = run_git(proj_path, 'cherry-pick', sha)
578
+ if result.returncode == 0:
579
+ picked += 1
580
+ continue
581
+ if not self._cherry_pick_is_empty(proj_path):
582
+ self._raise_git_error(
583
+ pr, proj_path, ('cherry-pick', sha), result
584
+ )
585
+ self._git(pr, proj_path, 'cherry-pick', '--skip')
586
+ empty_in_pr = run_git(
587
+ proj_path, 'diff', '--quiet', f'{sha}^', sha
588
+ ).returncode == 0
589
+ reason = ('it is empty in the PR' if empty_in_pr
590
+ else 'its changes are already present')
591
+ logger.info(f' Skipped commit {sha}: {reason}')
592
+
593
+ logger.info(
594
+ f' Cherry-picked {picked} of {len(pr.commits)} commit(s)'
595
+ )
596
+ if not picked:
597
+ logger.warning(
598
+ f' {pr.ref} added no new changes to {pr.project}'
599
+ )
600
+
601
+ def _present_prefix(self, pr, proj_path):
602
+ """
603
+ Length of the longest leading run of the PR's commits whose combined
604
+ changes already exist in the checkout, i.e. their diff
605
+ reverse-applies
606
+ """
607
+
608
+ base = f'{pr.commits[0]}^'
609
+ for count in range(len(pr.commits), 0, -1):
610
+ diff_args = ('diff', '--binary', base, pr.commits[count - 1])
611
+ diff = run_git(proj_path, *diff_args, text=False)
612
+ if diff.returncode != 0:
613
+ self._raise_git_error(pr, proj_path, diff_args, diff)
614
+ if run_git(
615
+ proj_path, 'apply', '--check', '--reverse',
616
+ text=False, input=diff.stdout
617
+ ).returncode == 0:
618
+ return count
619
+ return 0
620
+
621
+ @staticmethod
622
+ def _cherry_pick_is_empty(proj_path):
623
+ """True if a stopped cherry-pick has no conflicts and nothing staged"""
624
+
625
+ return (
626
+ run_git(proj_path, 'rev-parse', '-q', '--verify',
627
+ 'CHERRY_PICK_HEAD').returncode == 0
628
+ and run_git(proj_path, 'ls-files', '--unmerged').stdout == ''
629
+ and run_git(proj_path, 'diff', '--cached',
630
+ '--quiet').returncode == 0
631
+ )
632
+
633
+ def _git(self, pr, proj_path, *args, timeout=None):
634
+ """Run a git command for a PR, returning stdout"""
635
+
636
+ result = run_git(proj_path, *args, timeout=timeout)
637
+ if result.returncode != 0:
638
+ self._raise_git_error(pr, proj_path, args, result)
639
+ return result.stdout
640
+
641
+ @staticmethod
642
+ def _raise_git_error(pr, proj_path, args, result):
643
+ outputs = (
644
+ out.decode(errors='replace') if isinstance(out, bytes) else out
645
+ for out in (result.stdout, result.stderr)
646
+ )
647
+ output = '\n'.join(
648
+ out.strip() for out in outputs if out and out.strip()
649
+ )
650
+ msg = (
651
+ f'Failed to apply {pr.ref} to project {pr.project} '
652
+ f'({proj_path}): `git {" ".join(args)}` exited with code '
653
+ f'{result.returncode}'
654
+ )
655
+ if output:
656
+ msg += f'\n{output}'
657
+ if args[0] == 'cherry-pick':
658
+ msg += (
659
+ f'\n{proj_path} is left mid-cherry-pick; run '
660
+ '`git cherry-pick --abort` there to reset it'
661
+ )
662
+ raise PatchError(msg)
663
+
664
+ def check_requested_prs_applied(self, id_type, refs, requested):
665
+ """
666
+ Verify that every requested PR (keys of those resolved from the
667
+ request) was applied, raising PatchError if not
668
+ """
669
+
670
+ if id_type == 'label':
671
+ if self.applied_prs:
672
+ labels = ', '.join(f"'{ref}'" for ref in refs)
673
+ logger.info(
674
+ f'Applied PRs from label(s) {labels}: {self.applied_prs}'
675
+ )
676
+ return
677
+
678
+ requested = list(requested)
679
+ if any(key not in self.applied_prs for key in requested):
680
+ raise PatchError(
681
+ 'Failed to apply all explicitly-requested PRs! '
682
+ f'Requested: {requested} Applied: {self.applied_prs}'
683
+ )
684
+ if requested:
685
+ logger.info(
686
+ 'All explicitly-requested PRs applied! '
687
+ f'Requested: {requested} Applied: {self.applied_prs}'
688
+ )
689
+
690
+ def _abort(self, exc):
691
+ """Report a failed run, roll back what it changed, and exit"""
692
+
693
+ if isinstance(exc, PatchError):
694
+ logger.critical(exc)
695
+ elif isinstance(exc, KeyboardInterrupt):
696
+ logger.critical('Interrupted')
697
+ else:
698
+ logger.critical(f'Unexpected error: {exc!r}', exc_info=exc)
699
+
700
+ if not self.originals:
701
+ if isinstance(exc, KeyboardInterrupt):
702
+ sys.exit(130)
703
+ sys.exit(getattr(exc, 'exit_code', 1))
704
+
705
+ logger.critical(
706
+ 'Rolling back changes made by this run (PRs applied: '
707
+ f'{", ".join(self.applied_prs) or "none"})'
708
+ )
709
+ if not self.roll_back():
710
+ logger.critical(
711
+ 'Rollback incomplete; see above for projects needing '
712
+ 'manual attention'
713
+ )
714
+ sys.exit(5)
715
+
716
+ def roll_back(self):
717
+ """
718
+ Return every project this run touched to its original HEAD,
719
+ keeping uncommitted local changes. Returns True if all succeeded
720
+ """
721
+
722
+ ok = True
723
+ for proj_path, (branch, sha) in reversed(self.originals.items()):
724
+ target = f'{branch} at {sha}' if branch else sha
725
+ if run_git(proj_path, 'rev-parse', '-q', '--verify',
726
+ 'CHERRY_PICK_HEAD').returncode == 0:
727
+ run_git(proj_path, 'cherry-pick', '--abort')
728
+ if branch:
729
+ steps = [('checkout', '-q', branch),
730
+ ('reset', '-q', '--keep', sha)]
731
+ else:
732
+ steps = [('checkout', '-q', '--detach', sha)]
733
+ for step in steps:
734
+ result = run_git(proj_path, *step)
735
+ if result.returncode != 0:
736
+ logger.critical(
737
+ f' Could not roll back {proj_path}: '
738
+ f'`git {" ".join(step)}` failed: '
739
+ f'{result.stderr.strip()}. Restore it manually to '
740
+ f'{target}'
741
+ )
742
+ ok = False
743
+ break
744
+ else:
745
+ logger.critical(f' Rolled back {proj_path} to {target}')
746
+ return ok
747
+
748
+ def _api_get(self, endpoint):
749
+ """Make a GET request to the GitHub API"""
750
+
751
+ return self._request(f'{self.GITHUB_API_URL}{endpoint}')[1]
752
+
753
+ def _api_get_paginated(self, endpoint):
754
+ """Make paginated GET requests to the GitHub API"""
755
+
756
+ results = []
757
+ url = f'{self.GITHUB_API_URL}{endpoint}'
758
+ while url:
759
+ response, data = self._request(url)
760
+ if not isinstance(data, list):
761
+ raise PatchError(
762
+ f'GitHub API error: expected a list from {url}, got '
763
+ f'{type(data).__name__}: {str(data)[:200]}'
764
+ )
765
+ results.extend(data)
766
+ url = response.links.get('next', {}).get('url')
767
+ # The token goes with every request, so stay on the API host
768
+ if url and not url.startswith(f'{self.GITHUB_API_URL}/'):
769
+ raise PatchError(
770
+ f'GitHub API error: refusing to follow pagination link '
771
+ f'to {url}'
772
+ )
773
+ return results
774
+
775
+ def _request(self, url):
776
+ """GET a GitHub API URL, returning (response, decoded JSON)"""
777
+
778
+ logger.debug(f' API GET: {url}')
779
+ try:
780
+ response = self.session.get(url, timeout=30)
781
+ response.raise_for_status()
782
+ return response, response.json()
783
+ except requests.exceptions.HTTPError as exc:
784
+ raise PatchError(self._http_error_message(exc)) from exc
785
+ except requests.exceptions.JSONDecodeError as exc:
786
+ raise PatchError(
787
+ f'GitHub API error: invalid JSON from {url}: {exc}'
788
+ ) from exc
789
+ except requests.exceptions.RequestException as exc:
790
+ raise PatchError(
791
+ f'GitHub API error: request to {url} failed: {exc}'
792
+ ) from exc
793
+
794
+ @staticmethod
795
+ def _http_error_message(exc):
796
+ response = exc.response
797
+ if response is None:
798
+ return f'GitHub API error: {exc}'
799
+
800
+ msg = (
801
+ f'GitHub API error: {response.status_code} {response.reason} '
802
+ f'for {response.url}'
803
+ )
804
+ try:
805
+ detail = response.json().get('message')
806
+ except (ValueError, AttributeError):
807
+ detail = None
808
+ if detail:
809
+ msg += f' ({detail})'
810
+
811
+ headers = response.headers
812
+ if response.status_code == 401:
813
+ msg += ' - check the token in your config file'
814
+ elif (response.status_code in (403, 429)
815
+ and headers.get('X-RateLimit-Remaining') == '0'):
816
+ msg += ' - API rate limit exhausted'
817
+ reset = headers.get('X-RateLimit-Reset')
818
+ if reset and reset.isdigit():
819
+ reset_at = time.strftime(
820
+ '%H:%M:%S %Z', time.localtime(int(reset))
821
+ )
822
+ msg += f', resets at {reset_at}'
823
+ elif reset:
824
+ msg += f', X-RateLimit-Reset: {reset}'
825
+ elif (response.status_code in (403, 429)
826
+ and headers.get('Retry-After')):
827
+ msg += (
828
+ ' - secondary rate limit hit, retry after '
829
+ f'{headers["Retry-After"]} seconds'
830
+ )
831
+ elif response.status_code == 403:
832
+ msg += (
833
+ ' - check your token has access to the repository, and is '
834
+ 'authorized for the organization if it uses SSO'
835
+ )
836
+ elif response.status_code == 404:
837
+ msg += (
838
+ ' - check the reference is correct and that your token '
839
+ 'can access the repository'
840
+ )
841
+ return msg
842
+
843
+
844
+ class ParseCSVs(argparse.Action):
845
+ """Split each argument on commas, dropping empty values"""
846
+
847
+ def __call__(self, parser, namespace, arguments, option_string=None):
848
+ values = [
849
+ value for arg in arguments for value in arg.split(',') if value
850
+ ]
851
+ if not values:
852
+ parser.error(f'argument {option_string}: expected a value')
853
+ setattr(namespace, self.dest, values)
854
+
855
+
856
+ def default_ini_file():
857
+ """
858
+ Returns a string path to the default patch_via_github.ini
859
+ """
860
+ return os.path.expanduser('~/.ssh/patch_via_github.ini')
861
+
862
+
863
+ def print_divider():
864
+ """Print a visual divider line"""
865
+ print("=" * 80)
866
+
867
+
868
+ def setup_logging():
869
+ """Log to stderr at INFO; returns the handler so --debug can lower it"""
870
+
871
+ logger.setLevel(logging.DEBUG)
872
+ handler = logging.StreamHandler()
873
+ handler.setFormatter(logging.Formatter('%(levelname)s: %(message)s'))
874
+ handler.setLevel(logging.INFO)
875
+ logger.addHandler(handler)
876
+ return handler
877
+
878
+
879
+ def main():
880
+ """
881
+ Parse the arguments, verify the repo sync exists, read and validate
882
+ the configuration file, then determine all the needed GitHub PR
883
+ patches and apply them to the repo sync
884
+ """
885
+
886
+ # PyInstaller binaries get LD_LIBRARY_PATH set for them, and that
887
+ # can have unwanted side-effects for our subprocesses.
888
+ os.environ.pop("LD_LIBRARY_PATH", None)
889
+
890
+ handler = setup_logging()
891
+
892
+ version_string = (
893
+ f"patch_via_github version {version('patch-via-github')}"
894
+ )
895
+ parser = argparse.ArgumentParser(
896
+ description='Patch repo sync with requested GitHub pull requests'
897
+ )
898
+ parser.add_argument('-d', '--debug', action='store_true',
899
+ help='Enable debugging output')
900
+ parser.add_argument('-c', '--config', dest='github_config',
901
+ help='Configuration file for patching via GitHub',
902
+ default=default_ini_file())
903
+ change_group = parser.add_mutually_exclusive_group(required=True)
904
+ change_group.add_argument(
905
+ '-p', '--pull-request', dest='pull_requests', nargs='+',
906
+ action=ParseCSVs,
907
+ help='Pull request references to apply (comma-separated). '
908
+ 'Format: org/repo#number or repo#number (with -o)')
909
+ change_group.add_argument(
910
+ '-l', '--label', dest='labels', nargs='+',
911
+ action=ParseCSVs,
912
+ help='Labels to search for open PRs (comma-separated). '
913
+ 'Format: org/repo:label')
914
+ parser.add_argument('-o', '--default-org', dest='default_org',
915
+ help='Default GitHub organization for short-form '
916
+ 'PR references (repo#number)')
917
+ parser.add_argument('-s', '--source', dest='repo_source',
918
+ help='Location of the repo sync checkout',
919
+ default='.')
920
+ parser.add_argument('-C', '--checkout', action='store_true',
921
+ help='When specified, checkout the PR head '
922
+ 'rather than cherry-picking')
923
+ parser.add_argument('--no-ssh', dest='use_ssh',
924
+ action='store_false', default=None,
925
+ help='Use HTTPS URLs for git fetch instead of '
926
+ 'SSH (also configurable via ini: ssh = false; '
927
+ 'this flag takes precedence)')
928
+ parser.add_argument('-V', '--version', action='version',
929
+ help='Display patch_via_github version information',
930
+ version=version_string)
931
+ args = parser.parse_args()
932
+
933
+ if args.debug:
934
+ handler.setLevel(logging.DEBUG)
935
+
936
+ if not os.path.isdir(args.repo_source):
937
+ logger.error(
938
+ "Path for repo sync checkout doesn't exist. Aborting..."
939
+ )
940
+ sys.exit(1)
941
+ os.chdir(args.repo_source)
942
+
943
+ if args.pull_requests:
944
+ id_type, refs = 'pr', args.pull_requests
945
+ else:
946
+ id_type, refs = 'label', args.labels
947
+
948
+ logger.info(f"******** {version_string} ********")
949
+ print_divider()
950
+ try:
951
+ github_patches = GitHubPatches.from_config_file(
952
+ args.github_config, args.default_org, args.checkout, args.use_ssh
953
+ )
954
+ except PatchError as exc:
955
+ logger.critical(exc)
956
+ sys.exit(exc.exit_code)
957
+
958
+ logger.info(f"Initial request to patch {id_type}s: {', '.join(refs)}")
959
+ github_patches.patch_repo_sync(refs, id_type)
960
+
961
+ print_divider()
962
+
963
+
964
+ if __name__ == '__main__':
965
+ main()