@henryqw/pi-pr 3.1.10 → 4.0.2

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.
@@ -5,6 +5,7 @@ import type {
5
5
  import { lstatSync } from "node:fs";
6
6
  import { dirname, join, resolve } from "node:path";
7
7
  import { inspectLocalMergeSafety } from "./pr-merge.ts";
8
+ import { withWorktreeLock } from "./pr-execution.ts";
8
9
  import type {
9
10
  CiStatus,
10
11
  LocalMergeSafety,
@@ -22,6 +23,7 @@ const PR_DISCOVERY_PAGE_SIZE = 100;
22
23
  const PR_DISCOVERY_CAP = 1_000;
23
24
  const PR_DISCOVERY_MAX_PAGES = PR_DISCOVERY_CAP / PR_DISCOVERY_PAGE_SIZE;
24
25
  const PR_FIELDS = "id,number,url,state,isDraft,baseRefName,baseRefOid,headRefName,headRefOid,headRepository,mergeable,mergeStateStatus,reviewDecision,statusCheckRollup";
26
+ const PR_PUBLICATION_FIELDS = "number,url,state,baseRefName,headRefName,headRefOid,headRepository,title,body";
25
27
  const PR_DISCOVERY_QUERY = "query($owner:String!,$name:String!,$qualifiedName:String!,$endCursor:String){repository(owner:$owner,name:$name){nameWithOwner ref(qualifiedName:$qualifiedName){name associatedPullRequests(first:100,after:$endCursor){totalCount edges{cursor node{__typename number url state baseRepository{nameWithOwner}headRepository{nameWithOwner}headRefName headRefOid}}pageInfo{hasNextPage startCursor endCursor}}}}}";
26
28
  const REVIEW_THREADS_QUERY = "query($id:ID!,$endCursor:String){node(id:$id){...on PullRequest{reviewThreads(first:100,after:$endCursor){nodes{isResolved}pageInfo{hasNextPage endCursor}}}}}";
27
29
  const BASE_REF_QUERY = "query($owner:String!,$name:String!,$qualifiedName:String!){repository(owner:$owner,name:$name){nameWithOwner ref(qualifiedName:$qualifiedName){name target{oid}}}}";
@@ -149,7 +151,7 @@ type LinkConfiguration = {
149
151
  mirror: string[];
150
152
  };
151
153
 
152
- type SearchPullRequest = {
154
+ export type PullRequestCandidate = {
153
155
  number: number;
154
156
  url: URL;
155
157
  lifecycle: PullRequestLifecycle;
@@ -159,6 +161,32 @@ type SearchPullRequest = {
159
161
  headOid: string;
160
162
  };
161
163
 
164
+ type SearchPullRequest = PullRequestCandidate;
165
+
166
+ export type PullRequestPublication = {
167
+ number: number;
168
+ url: URL;
169
+ lifecycle: PullRequestLifecycle;
170
+ base: { repository: string; ref: string };
171
+ head: PullRequestRef;
172
+ title: string;
173
+ body: string;
174
+ };
175
+
176
+ export type ValidatedRemoteAuthority = {
177
+ fetchSource: string;
178
+ host: string;
179
+ repository: string;
180
+ };
181
+
182
+ export type BranchUpstreamTarget = {
183
+ branch: string;
184
+ remote: string;
185
+ ref: string;
186
+ fetchSource: string;
187
+ remoteOid: string;
188
+ };
189
+
162
190
  type SearchSelection =
163
191
  | { kind: "candidate"; candidate: SearchPullRequest; pullRequest: ListedPullRequest | null }
164
192
  | { kind: "none" }
@@ -179,6 +207,11 @@ type RulesetBranchPolicy = {
179
207
  allowedMergeMethods: MergeMethod[] | null;
180
208
  };
181
209
 
210
+ type CheckState = {
211
+ state: string;
212
+ diagnosableFailure: boolean;
213
+ };
214
+
182
215
  type ListedPullRequest = {
183
216
  id: string;
184
217
  number: number;
@@ -190,7 +223,7 @@ type ListedPullRequest = {
190
223
  mergeable: "MERGEABLE" | "CONFLICTING" | "UNKNOWN";
191
224
  mergeStateStatus: "BEHIND" | "BLOCKED" | "CLEAN" | "DIRTY" | "DRAFT" | "HAS_HOOKS" | "UNKNOWN" | "UNSTABLE";
192
225
  reviewDecision: "APPROVED" | "CHANGES_REQUESTED" | "REVIEW_REQUIRED" | null;
193
- checkStates: string[];
226
+ checkStates: CheckState[];
194
227
  };
195
228
 
196
229
  function fail(action: string, reason: string): never {
@@ -536,27 +569,64 @@ function checkOutcome(state: string): "failure" | "success" | "running" {
536
569
  return "running";
537
570
  }
538
571
 
539
- function checkState(value: unknown): string {
540
- if (!isRecord(value)) fail("Find pull requests", "invalid statusCheckRollup");
541
- const conclusion = optionalCheckState(value, "conclusion");
542
- const state = optionalCheckState(value, "state");
543
- const status = optionalCheckState(value, "status");
544
- const states = [conclusion, state, status].filter((value): value is string => value !== null);
572
+ function canonicalPositiveDecimal(value: string): boolean {
573
+ const parsed = Number(value);
574
+ return Number.isSafeInteger(parsed) && parsed > 0 && String(parsed) === value;
575
+ }
576
+
577
+ function isActionsJobUrl(value: unknown, pullRequestUrl: URL, repository: string): boolean {
578
+ if (typeof value !== "string" || !value) return false;
579
+ let url: URL;
580
+ try {
581
+ url = new URL(value);
582
+ } catch {
583
+ return false;
584
+ }
585
+ if (url.protocol !== "https:" || url.username || url.password || url.port || url.search || url.hash ||
586
+ url.hostname.toLowerCase() !== pullRequestUrl.hostname.toLowerCase()) return false;
587
+ const parts = url.pathname.split("/").filter(Boolean);
588
+ const expected = repository.split("/");
589
+ return parts.length === 7 && expected.length === 2 &&
590
+ parts[0]!.toLowerCase() === expected[0]!.toLowerCase() &&
591
+ parts[1]!.toLowerCase() === expected[1]!.toLowerCase() &&
592
+ parts[2] === "actions" && parts[3] === "runs" && canonicalPositiveDecimal(parts[4]!) &&
593
+ parts[5] === "job" && canonicalPositiveDecimal(parts[6]!);
594
+ }
595
+
596
+ function isDiagnosableActionsCheck(check: Record<string, unknown>, pullRequestUrl: URL, repository: string): boolean {
597
+ const workflowName = check.workflowName;
598
+ return typeof workflowName === "string" && !!workflowName && workflowName.trim() === workflowName &&
599
+ !/\p{Cc}/u.test(workflowName) && isActionsJobUrl(check.detailsUrl, pullRequestUrl, repository);
600
+ }
601
+
602
+ function checkState(value: unknown, pullRequestUrl: URL, repository: string): CheckState {
603
+ if (!isRecord(value) || (value.__typename !== "CheckRun" && value.__typename !== "StatusContext")) {
604
+ return fail("Find pull requests", "invalid statusCheckRollup");
605
+ }
606
+ const conclusion = value.__typename === "CheckRun" ? optionalCheckState(value, "conclusion") : null;
607
+ const state = value.__typename === "StatusContext" ? optionalCheckState(value, "state") : null;
608
+ const status = value.__typename === "CheckRun" ? optionalCheckState(value, "status") : null;
609
+ const states = [conclusion, state, status].filter((candidate): candidate is string => candidate !== null);
545
610
  if (!states.length) fail("Find pull requests", "invalid statusCheckRollup");
546
611
 
547
612
  // COMPLETED describes a check run's lifecycle; its conclusion gives the outcome.
548
- const outcomes = states.filter((value) => value !== "COMPLETED").map(checkOutcome);
613
+ const outcomes = states.filter((candidate) => candidate !== "COMPLETED").map(checkOutcome);
549
614
  if (
550
615
  new Set(outcomes).size > 1 ||
551
616
  (states.includes("COMPLETED") && outcomes.includes("running"))
552
617
  ) fail("Find pull requests", "invalid statusCheckRollup");
553
- return conclusion ?? state ?? status ?? fail("Find pull requests", "invalid statusCheckRollup");
618
+ const selected = conclusion ?? state ?? status ?? fail("Find pull requests", "invalid statusCheckRollup");
619
+ return {
620
+ state: selected,
621
+ diagnosableFailure: FAILED_CHECK_STATES.has(selected) && value.__typename === "CheckRun" &&
622
+ isDiagnosableActionsCheck(value, pullRequestUrl, repository),
623
+ };
554
624
  }
555
625
 
556
- function checkStates(value: unknown): string[] {
626
+ function checkStates(value: unknown, pullRequestUrl: URL, repository: string): CheckState[] {
557
627
  if (value === null) return [];
558
628
  if (!Array.isArray(value)) fail("Find pull requests", "invalid statusCheckRollup");
559
- return value.map(checkState);
629
+ return value.map((check) => checkState(check, pullRequestUrl, repository));
560
630
  }
561
631
 
562
632
  function listedPullRequest(value: unknown): ListedPullRequest | null {
@@ -589,7 +659,7 @@ function listedPullRequest(value: unknown): ListedPullRequest | null {
589
659
  mergeable: mergeable(value.mergeable),
590
660
  mergeStateStatus: mergeStateStatus(value.mergeStateStatus),
591
661
  reviewDecision: reviewDecision(value.reviewDecision),
592
- checkStates: checkStates(value.statusCheckRollup),
662
+ checkStates: checkStates(value.statusCheckRollup, parsedUrl.url, parsedUrl.repository),
593
663
  };
594
664
  }
595
665
 
@@ -634,7 +704,7 @@ function searchPullRequest(value: unknown, host: string): SearchPullRequest {
634
704
  };
635
705
  }
636
706
 
637
- function parseSearchPage(output: string, pushTarget: PushTarget): SearchPage | null {
707
+ function parseSearchPage(output: string, pushTarget: Pick<PushTarget, "repository" | "ref">): SearchPage | null {
638
708
  const page = parseJson(output, "Find pull requests");
639
709
  if (!isRecord(page)) fail("Find pull requests", "invalid GitHub CLI output");
640
710
  if (page.errors !== undefined) {
@@ -718,6 +788,37 @@ function parseLoadedPullRequest(output: string, expectedUrl: URL): ListedPullReq
718
788
  return candidate;
719
789
  }
720
790
 
791
+ function parsePullRequestPublication(output: string, expectedUrl: URL): PullRequestPublication {
792
+ const value = parseJson(output, "Read pull request publication");
793
+ if (!isRecord(value) || !isRecord(value.headRepository)) {
794
+ fail("Read pull request publication", "invalid GitHub CLI output");
795
+ }
796
+ const number = value.number;
797
+ if (typeof number !== "number" || !Number.isSafeInteger(number) || number <= 0) {
798
+ fail("Read pull request publication", "invalid number");
799
+ }
800
+ const parsedUrl = parsePullRequestUrl(value.url, number);
801
+ if (parsedUrl.url.href !== expectedUrl.href) {
802
+ fail("Read pull request publication", "response does not match candidate url");
803
+ }
804
+ return {
805
+ number,
806
+ url: parsedUrl.url,
807
+ lifecycle: lifecycle(value.state),
808
+ base: {
809
+ repository: parsedUrl.repository,
810
+ ref: text(value.baseRefName, "Read pull request publication", "baseRefName"),
811
+ },
812
+ head: {
813
+ repository: repositoryName(value.headRepository.nameWithOwner, "Read pull request publication", "headRepository.nameWithOwner"),
814
+ ref: text(value.headRefName, "Read pull request publication", "headRefName"),
815
+ oid: oid(value.headRefOid, "Read pull request publication", "headRefOid"),
816
+ },
817
+ title: text(value.title, "Read pull request publication", "title"),
818
+ body: typeof value.body === "string" ? value.body : fail("Read pull request publication", "invalid body"),
819
+ };
820
+ }
821
+
721
822
  function selectPullRequest(
722
823
  candidates: ListedPullRequest[],
723
824
  pushTarget: PushTarget,
@@ -745,13 +846,19 @@ function selectPullRequest(
745
846
  return historical[0] ?? null;
746
847
  }
747
848
 
748
- function ciStatus(states: string[]): CiStatus {
749
- if (!states.length) return "none";
849
+ function ciStatus(checks: CheckState[]): CiStatus {
850
+ if (!checks.length) return "none";
851
+ let failed = false;
750
852
  let running = false;
751
- for (const state of states) {
752
- if (FAILED_CHECK_STATES.has(state)) return "failure";
753
- if (!SUCCESSFUL_CHECK_STATES.has(state)) running = true;
853
+ for (const check of checks) {
854
+ if (FAILED_CHECK_STATES.has(check.state)) {
855
+ if (!check.diagnosableFailure) return "failure-blocked";
856
+ failed = true;
857
+ } else if (!SUCCESSFUL_CHECK_STATES.has(check.state)) {
858
+ running = true;
859
+ }
754
860
  }
861
+ if (failed) return "failure";
755
862
  return running ? "running" : "success";
756
863
  }
757
864
 
@@ -821,7 +928,10 @@ function parseUnresolvedReviewThreads(output: string): number {
821
928
  return total;
822
929
  }
823
930
 
824
- function parseBaseRefOid(output: string, candidate: ListedPullRequest): string {
931
+ function parseBaseRefAuthority(
932
+ output: string,
933
+ expected: { repository: string; ref: string },
934
+ ): string {
825
935
  const value = parseJson(output, "Read base ref");
826
936
  if (!isRecord(value)) fail("Read base ref", "invalid GitHub CLI output");
827
937
  if (value.errors !== undefined) {
@@ -834,12 +944,16 @@ function parseBaseRefOid(output: string, candidate: ListedPullRequest): string {
834
944
  }
835
945
  if (
836
946
  normalizeRepository(repositoryName(repository.nameWithOwner, "Read base ref", "repository")) !==
837
- normalizeRepository(candidate.base.repository) ||
838
- text(repository.ref.name, "Read base ref", "ref") !== candidate.base.ref
947
+ normalizeRepository(expected.repository) ||
948
+ text(repository.ref.name, "Read base ref", "ref") !== expected.ref
839
949
  ) fail("Read base ref", "response does not match pull request base");
840
950
  return oid(repository.ref.target.oid, "Read base ref", "target OID");
841
951
  }
842
952
 
953
+ function parseBaseRefOid(output: string, candidate: ListedPullRequest): string {
954
+ return parseBaseRefAuthority(output, candidate.base);
955
+ }
956
+
843
957
  function parseLegacyBaseBranchPolicy(output: string, candidate: ListedPullRequest): boolean {
844
958
  const value = parseJson(output, "Read base branch policy");
845
959
  if (!isRecord(value)) fail("Read base branch policy", "invalid GitHub CLI output");
@@ -1000,6 +1114,20 @@ async function readRemoteAuthority(
1000
1114
  }
1001
1115
  }
1002
1116
 
1117
+ export async function readValidatedRemoteAuthority(
1118
+ pi: Pick<ExtensionAPI, "exec">,
1119
+ context: PullRequestLoadContext,
1120
+ remote: string,
1121
+ ): Promise<ValidatedRemoteAuthority> {
1122
+ const authority = await readRemoteAuthority(pi, context, text(remote, "Read push remotes", "remote"), true);
1123
+ if (!authority) fail("Read push target", "invalid remote authority");
1124
+ return {
1125
+ fetchSource: authority.fetchSource,
1126
+ host: authority.repository.host,
1127
+ repository: authority.repository.nameWithOwner,
1128
+ };
1129
+ }
1130
+
1003
1131
  async function readRemoteHeadOid(
1004
1132
  pi: Pick<ExtensionAPI, "exec">,
1005
1133
  context: PullRequestLoadContext,
@@ -1293,6 +1421,25 @@ async function readBaseRefOid(
1293
1421
  return parseBaseRefOid(result.stdout, candidate);
1294
1422
  }
1295
1423
 
1424
+ export async function readPullRequestBaseRefOid(
1425
+ pi: Pick<ExtensionAPI, "exec">,
1426
+ context: PullRequestLoadContext,
1427
+ authority: { host: string; repository: string; ref: string },
1428
+ ): Promise<string> {
1429
+ const host = text(authority.host, "Read base ref", "host").toLowerCase();
1430
+ const repository = repositoryName(authority.repository, "Read base ref", "repository");
1431
+ const ref = text(authority.ref, "Read base ref", "ref");
1432
+ const [owner, name] = repository.split("/");
1433
+ const result = await execute(pi, context, "Read base ref", "gh", [
1434
+ "api", "graphql", "--hostname", host,
1435
+ "-f", `query=${BASE_REF_QUERY}`,
1436
+ "-F", `owner=${owner}`,
1437
+ "-F", `name=${name}`,
1438
+ "-F", `qualifiedName=refs/heads/${ref}`,
1439
+ ]);
1440
+ return parseBaseRefAuthority(result.stdout, { repository, ref });
1441
+ }
1442
+
1296
1443
  async function readLegacyBaseBranchPolicy(
1297
1444
  pi: Pick<ExtensionAPI, "exec">,
1298
1445
  context: PullRequestLoadContext,
@@ -1449,35 +1596,28 @@ async function loadObservedPullRequest(
1449
1596
  return loadPullRequestDetails(pi, context, candidate, pushTarget, inspectedLocal);
1450
1597
  }
1451
1598
 
1452
- async function searchPullRequests(
1599
+ async function enumerateSearchPullRequests(
1453
1600
  pi: Pick<ExtensionAPI, "exec">,
1454
1601
  context: PullRequestLoadContext,
1455
- pushTarget: PushTarget,
1456
- ): Promise<SearchSelection> {
1457
- const [owner, name] = pushTarget.repository.nameWithOwner.split("/");
1602
+ target: Pick<PushTarget, "repository" | "ref">,
1603
+ ): Promise<SearchPullRequest[] | null> {
1604
+ const [owner, name] = target.repository.nameWithOwner.split("/");
1458
1605
  const candidates: SearchPullRequest[] = [];
1459
1606
  const cursors = new Set<string>();
1460
1607
  let totalCount: number | null = null;
1461
1608
  let endCursor: string | null = null;
1462
1609
  for (let pageIndex = 0; pageIndex < PR_DISCOVERY_MAX_PAGES; pageIndex += 1) {
1463
1610
  const args = [
1464
- "api",
1465
- "graphql",
1466
- "--hostname",
1467
- pushTarget.repository.host,
1468
- "-f",
1469
- `query=${PR_DISCOVERY_QUERY}`,
1470
- "-F",
1471
- `owner=${owner}`,
1472
- "-F",
1473
- `name=${name}`,
1474
- "-F",
1475
- `qualifiedName=refs/heads/${pushTarget.ref}`,
1611
+ "api", "graphql", "--hostname", target.repository.host,
1612
+ "-f", `query=${PR_DISCOVERY_QUERY}`,
1613
+ "-F", `owner=${owner}`,
1614
+ "-F", `name=${name}`,
1615
+ "-F", `qualifiedName=refs/heads/${target.ref}`,
1476
1616
  ];
1477
1617
  if (endCursor !== null) args.push("-F", `endCursor=${endCursor}`);
1478
1618
  const result = await execute(pi, context, "Find pull requests", "gh", args);
1479
- const page = parseSearchPage(result.stdout, pushTarget);
1480
- if (page === null) return pushTarget.remoteHeadOid === null ? { kind: "none" } : { kind: "target-invalid" };
1619
+ const page = parseSearchPage(result.stdout, target);
1620
+ if (page === null) return null;
1481
1621
  if (totalCount !== null && page.totalCount !== totalCount) {
1482
1622
  fail("Find pull requests", "inconsistent search result pages");
1483
1623
  }
@@ -1495,27 +1635,60 @@ async function searchPullRequests(
1495
1635
  }
1496
1636
  const hasMore = candidates.length < totalCount;
1497
1637
  if (page.hasNextPage !== hasMore) fail("Find pull requests", "incomplete search results");
1498
- if (!hasMore) {
1499
- const selected = selectSearchPullRequest(candidates, pushTarget);
1500
- if (selected.kind !== "candidate") return selected;
1501
- const loaded = await execute(pi, context, "Find pull requests", "gh", [
1502
- "pr",
1503
- "view",
1504
- selected.candidate.url.href,
1505
- "--json",
1506
- PR_FIELDS,
1507
- ]);
1508
- return {
1509
- ...selected,
1510
- pullRequest: parseLoadedPullRequest(loaded.stdout, selected.candidate.url),
1511
- };
1512
- }
1638
+ if (!hasMore) return candidates;
1513
1639
  if (page.endCursor === null) fail("Find pull requests", "invalid search pageInfo");
1514
1640
  endCursor = page.endCursor;
1515
1641
  }
1516
1642
  return fail("Find pull requests", "GitHub pull request result cap reached");
1517
1643
  }
1518
1644
 
1645
+ export async function findExactHeadPullRequests(
1646
+ pi: Pick<ExtensionAPI, "exec">,
1647
+ context: PullRequestLoadContext,
1648
+ target: { host: string; repository: string; ref: string },
1649
+ ): Promise<PullRequestCandidate[]> {
1650
+ const host = text(target.host, "Find pull requests", "host").toLowerCase();
1651
+ const repository = repositoryName(target.repository, "Find pull requests", "head repository");
1652
+ const ref = text(target.ref, "Find pull requests", "head ref");
1653
+ const candidates = await enumerateSearchPullRequests(pi, context, {
1654
+ repository: { host, nameWithOwner: repository, normalizedName: normalizeRepository(repository) },
1655
+ ref,
1656
+ });
1657
+ if (candidates === null) return fail("Find pull requests", "published head ref is unavailable");
1658
+ return candidates.filter((candidate) =>
1659
+ candidate.lifecycle === "open" && candidate.headRepository !== null &&
1660
+ normalizeRepository(candidate.headRepository) === normalizeRepository(repository) && candidate.headRef === ref
1661
+ );
1662
+ }
1663
+
1664
+ export async function loadPullRequestPublication(
1665
+ pi: Pick<ExtensionAPI, "exec">,
1666
+ context: PullRequestLoadContext,
1667
+ url: URL,
1668
+ ): Promise<PullRequestPublication> {
1669
+ const loaded = await execute(pi, context, "Read pull request publication", "gh", [
1670
+ "pr", "view", url.href, "--json", PR_PUBLICATION_FIELDS,
1671
+ ]);
1672
+ return parsePullRequestPublication(loaded.stdout, url);
1673
+ }
1674
+
1675
+ async function searchPullRequests(
1676
+ pi: Pick<ExtensionAPI, "exec">,
1677
+ context: PullRequestLoadContext,
1678
+ pushTarget: PushTarget,
1679
+ ): Promise<SearchSelection> {
1680
+ const candidates = await enumerateSearchPullRequests(pi, context, pushTarget);
1681
+ if (candidates === null) {
1682
+ return pushTarget.remoteHeadOid === null ? { kind: "none" } : { kind: "target-invalid" };
1683
+ }
1684
+ const selected = selectSearchPullRequest(candidates, pushTarget);
1685
+ if (selected.kind !== "candidate") return selected;
1686
+ const loaded = await execute(pi, context, "Find pull requests", "gh", [
1687
+ "pr", "view", selected.candidate.url.href, "--json", PR_FIELDS,
1688
+ ]);
1689
+ return { ...selected, pullRequest: parseLoadedPullRequest(loaded.stdout, selected.candidate.url) };
1690
+ }
1691
+
1519
1692
  export async function loadCurrentPullRequest(
1520
1693
  pi: Pick<ExtensionAPI, "exec">,
1521
1694
  context: PullRequestLoadContext,
@@ -1682,7 +1855,7 @@ function sameLinkedPullRequest(inferred: CurrentPullRequest, configured: Current
1682
1855
  inferred.target.remoteOid === configured.target.remoteOid;
1683
1856
  }
1684
1857
 
1685
- async function readTrackingOid(
1858
+ export async function readTrackingOid(
1686
1859
  pi: Pick<ExtensionAPI, "exec">,
1687
1860
  context: PullRequestLoadContext,
1688
1861
  trackingRef: string,
@@ -1696,6 +1869,46 @@ async function readTrackingOid(
1696
1869
  return oid(singleLine(result.stdout, "Read remote-tracking ref", "OID"), "Read remote-tracking ref", "OID");
1697
1870
  }
1698
1871
 
1872
+ export function branchTrackingRef(target: Pick<BranchUpstreamTarget, "remote" | "ref">): string {
1873
+ return `refs/remotes/${text(target.remote, "Read push remotes", "remote")}/${text(target.ref, "Read push target", "push ref")}`;
1874
+ }
1875
+
1876
+ export async function fetchBranchTrackingRef(
1877
+ pi: Pick<ExtensionAPI, "exec">,
1878
+ context: PullRequestLoadContext,
1879
+ target: BranchUpstreamTarget,
1880
+ ): Promise<void> {
1881
+ await execute(pi, context, "Fetch branch tracking ref", "git", [
1882
+ "fetch", "--no-write-fetch-head", "--no-tags", "--no-recurse-submodules",
1883
+ target.fetchSource, `+${target.remoteOid}:${branchTrackingRef(target)}`,
1884
+ ]);
1885
+ }
1886
+
1887
+ export async function setBranchUpstream(
1888
+ pi: Pick<ExtensionAPI, "exec">,
1889
+ context: PullRequestLoadContext,
1890
+ target: BranchUpstreamTarget,
1891
+ ): Promise<void> {
1892
+ await execute(pi, context, "Set branch upstream", "git", [
1893
+ "branch", `--set-upstream-to=${target.remote}/${target.ref}`, "--", target.branch,
1894
+ ]);
1895
+ }
1896
+
1897
+ export async function verifyBranchUpstream(
1898
+ pi: Pick<ExtensionAPI, "exec">,
1899
+ context: PullRequestLoadContext,
1900
+ target: BranchUpstreamTarget,
1901
+ ): Promise<void> {
1902
+ const trackingOid = await readTrackingOid(pi, context, branchTrackingRef(target));
1903
+ if (trackingOid !== target.remoteOid) throw new Error("Branch tracking ref does not match published OID");
1904
+ const configuredTarget = optionalPushReference((await execute(pi, context, "Verify push target", "git", [
1905
+ "for-each-ref", "--format=%(push:short)", `refs/heads/${target.branch}`,
1906
+ ])).stdout);
1907
+ if (configuredTarget !== `${target.remote}/${target.ref}`) {
1908
+ throw new Error("Configured push target does not match published branch");
1909
+ }
1910
+ }
1911
+
1699
1912
  async function restoreConfigValue(
1700
1913
  pi: Pick<ExtensionAPI, "exec">,
1701
1914
  context: PullRequestLoadContext,
@@ -1725,6 +1938,10 @@ async function restoreConfigValue(
1725
1938
  }
1726
1939
  }
1727
1940
 
1941
+ function sameConfigValues(left: readonly string[], right: readonly string[]): boolean {
1942
+ return left.length === right.length && left.every((value, index) => value === right[index]);
1943
+ }
1944
+
1728
1945
  async function restoreLinkState(
1729
1946
  pi: Pick<ExtensionAPI, "exec">,
1730
1947
  context: PullRequestLoadContext,
@@ -1732,32 +1949,39 @@ async function restoreLinkState(
1732
1949
  configuration: LinkConfiguration,
1733
1950
  trackingRef: string,
1734
1951
  trackingOid: string | null,
1735
- upstreamMutated: boolean,
1736
- fetchedOid: string | undefined,
1952
+ upstreamAttempted: boolean,
1953
+ fetchAttempted: boolean,
1737
1954
  ): Promise<void> {
1738
1955
  let incomplete = false;
1739
- if (upstreamMutated) {
1956
+ if (upstreamAttempted) {
1740
1957
  for (const [key, expected, original] of [
1741
1958
  [`branch.${target.branch}.remote`, target.remote, configuration.upstreamRemote],
1742
1959
  [`branch.${target.branch}.merge`, `refs/heads/${target.ref}`, configuration.upstreamMerge],
1743
1960
  ] as const) {
1744
1961
  try {
1745
- await restoreConfigValue(pi, context, key, expected, original);
1962
+ const current = await readConfigValues(pi, context, key);
1963
+ if (current === null) incomplete = true;
1964
+ else if (sameConfigValues(current, original)) continue;
1965
+ else if (current.length === 1 && current[0] === expected) {
1966
+ await restoreConfigValue(pi, context, key, expected, original);
1967
+ } else incomplete = true;
1746
1968
  } catch {
1747
1969
  incomplete = true;
1748
1970
  }
1749
1971
  }
1750
1972
  }
1751
- if (fetchedOid !== undefined) {
1973
+ if (fetchAttempted) {
1752
1974
  try {
1753
1975
  const currentTrackingOid = await readTrackingOid(pi, context, trackingRef);
1754
- if (currentTrackingOid !== fetchedOid) {
1755
- incomplete = true;
1756
- } else {
1757
- const args = trackingOid === null
1758
- ? ["update-ref", "-d", trackingRef, currentTrackingOid]
1759
- : ["update-ref", trackingRef, trackingOid, currentTrackingOid];
1760
- await execute(pi, context, "Restore remote-tracking ref", "git", args);
1976
+ if (currentTrackingOid !== trackingOid) {
1977
+ if (target.remoteHeadOid === null || currentTrackingOid !== target.remoteHeadOid) {
1978
+ incomplete = true;
1979
+ } else {
1980
+ const args = trackingOid === null
1981
+ ? ["update-ref", "-d", trackingRef, target.remoteHeadOid]
1982
+ : ["update-ref", trackingRef, trackingOid, target.remoteHeadOid];
1983
+ await execute(pi, context, "Restore remote-tracking ref", "git", args);
1984
+ }
1761
1985
  }
1762
1986
  } catch {
1763
1987
  incomplete = true;
@@ -1770,90 +1994,84 @@ export async function linkInferredPullRequest(
1770
1994
  pi: Pick<ExtensionAPI, "exec">,
1771
1995
  context: PullRequestLoadContext,
1772
1996
  inferred: CurrentPullRequest,
1997
+ options: { agentDir?: string } = {},
1773
1998
  ): Promise<CurrentPullRequest> {
1774
1999
  if (inferred.target.provenance !== "inferred" || inferred.lifecycle !== "open") {
1775
2000
  throw new Error("Link branch failed: pull request is not an open inferred target");
1776
2001
  }
1777
- const freshDiscovery = await loadCurrentPullRequest(pi, context);
1778
- if (
1779
- freshDiscovery.kind !== "current" ||
1780
- freshDiscovery.pullRequest.target.provenance !== "inferred" ||
1781
- !samePullRequestSnapshot(inferred, freshDiscovery.pullRequest)
1782
- ) throw new Error("Link branch cancelled: inferred pull request context changed");
1783
- inferred = freshDiscovery.pullRequest;
1784
- const target: PushTarget = {
1785
- provenance: "inferred",
1786
- branch: inferred.target.branch,
1787
- remote: inferred.target.remote,
1788
- ref: inferred.target.ref,
1789
- fetchSource: inferred.target.fetchSource,
1790
- remoteHeadOid: inferred.target.remoteOid,
1791
- repository: {
1792
- nameWithOwner: inferred.target.repository,
1793
- normalizedName: normalizeRepository(inferred.target.repository),
1794
- host: inferred.target.host,
1795
- },
1796
- };
1797
- const linkConfiguration = await readLinkConfiguration(pi, context, target);
1798
- if (target.remoteHeadOid === null || !linkConfiguration || !canLinkTarget(linkConfiguration, target)) {
1799
- throw new Error("Link branch cancelled: target configuration changed");
1800
- }
1801
- const pushReference = optionalPushReference((await execute(pi, context, "Read push target", "git", [
1802
- "for-each-ref", "--format=%(push:short)", `refs/heads/${target.branch}`,
1803
- ])).stdout);
1804
- if (pushReference !== null) throw new Error("Link branch cancelled: push target is no longer empty");
1805
- const remoteHeadOid = await readRemoteHeadOid(pi, context, target.fetchSource, target.ref);
1806
- if (remoteHeadOid !== target.remoteHeadOid) throw new Error("Link branch cancelled: remote ref changed");
1807
-
1808
- const trackingRef = `refs/remotes/${target.remote}/${target.ref}`;
1809
- const trackingOid = await readTrackingOid(pi, context, trackingRef);
1810
- let fetchedOid: string | undefined;
1811
- let upstreamMutated = false;
1812
- try {
1813
- await execute(pi, context, "Fetch inferred branch", "git", [
1814
- "fetch",
1815
- "--no-write-fetch-head",
1816
- "--no-tags",
1817
- "--no-recurse-submodules",
1818
- target.fetchSource,
1819
- `${target.remoteHeadOid}:${trackingRef}`,
1820
- ]);
1821
- fetchedOid = target.remoteHeadOid;
1822
- const verifiedFetchedOid = (await readTrackingOid(pi, context, trackingRef)) ?? undefined;
1823
- if (verifiedFetchedOid !== fetchedOid) {
1824
- throw new Error("Link branch cancelled: fetched remote ref changed");
2002
+ return await withWorktreeLock(context.cwd, async () => {
2003
+ const freshDiscovery = await loadCurrentPullRequest(pi, context);
2004
+ if (
2005
+ freshDiscovery.kind !== "current" ||
2006
+ freshDiscovery.pullRequest.target.provenance !== "inferred" ||
2007
+ !samePullRequestSnapshot(inferred, freshDiscovery.pullRequest)
2008
+ ) throw new Error("Link branch cancelled: inferred pull request context changed");
2009
+ inferred = freshDiscovery.pullRequest;
2010
+ const target: PushTarget = {
2011
+ provenance: "inferred",
2012
+ branch: inferred.target.branch,
2013
+ remote: inferred.target.remote,
2014
+ ref: inferred.target.ref,
2015
+ fetchSource: inferred.target.fetchSource,
2016
+ remoteHeadOid: inferred.target.remoteOid,
2017
+ repository: {
2018
+ nameWithOwner: inferred.target.repository,
2019
+ normalizedName: normalizeRepository(inferred.target.repository),
2020
+ host: inferred.target.host,
2021
+ },
2022
+ };
2023
+ const linkConfiguration = await readLinkConfiguration(pi, context, target);
2024
+ if (target.remoteHeadOid === null || !linkConfiguration || !canLinkTarget(linkConfiguration, target)) {
2025
+ throw new Error("Link branch cancelled: target configuration changed");
1825
2026
  }
1826
- await execute(pi, context, "Set branch upstream", "git", [
1827
- "branch", `--set-upstream-to=${target.remote}/${target.ref}`, "--", target.branch,
1828
- ]);
1829
- upstreamMutated = true;
1830
- const configuredTarget = optionalPushReference((await execute(pi, context, "Verify push target", "git", [
2027
+ const pushReference = optionalPushReference((await execute(pi, context, "Read push target", "git", [
1831
2028
  "for-each-ref", "--format=%(push:short)", `refs/heads/${target.branch}`,
1832
2029
  ])).stdout);
1833
- if (configuredTarget !== `${target.remote}/${target.ref}`) {
1834
- throw new Error("Link branch failed: configured push target does not match inferred target");
1835
- }
1836
- const discovery = await loadCurrentPullRequest(pi, context);
1837
- if (discovery.kind !== "current" || !sameLinkedPullRequest(inferred, discovery.pullRequest)) {
1838
- throw new Error("Link branch failed: configured pull request does not match inferred target");
1839
- }
1840
- return discovery.pullRequest;
1841
- } catch (error) {
1842
- const rollbackContext = { cwd: context.cwd, signal: new AbortController().signal };
2030
+ if (pushReference !== null) throw new Error("Link branch cancelled: push target is no longer empty");
2031
+ const remoteHeadOid = await readRemoteHeadOid(pi, context, target.fetchSource, target.ref);
2032
+ if (remoteHeadOid !== target.remoteHeadOid) throw new Error("Link branch cancelled: remote ref changed");
2033
+
2034
+ const trackingRef = `refs/remotes/${target.remote}/${target.ref}`;
2035
+ const trackingOid = await readTrackingOid(pi, context, trackingRef);
2036
+ let fetchAttempted = false;
2037
+ let upstreamAttempted = false;
1843
2038
  try {
1844
- await restoreLinkState(
1845
- pi,
1846
- rollbackContext,
1847
- target,
1848
- linkConfiguration,
1849
- trackingRef,
1850
- trackingOid,
1851
- upstreamMutated,
1852
- fetchedOid,
1853
- );
1854
- } catch {
1855
- throw new Error("Link branch failed and rollback was incomplete");
2039
+ const upstreamTarget: BranchUpstreamTarget = {
2040
+ branch: target.branch,
2041
+ remote: target.remote,
2042
+ ref: target.ref,
2043
+ fetchSource: target.fetchSource,
2044
+ remoteOid: target.remoteHeadOid,
2045
+ };
2046
+ fetchAttempted = true;
2047
+ await fetchBranchTrackingRef(pi, context, upstreamTarget);
2048
+ const verifiedFetchedOid = await readTrackingOid(pi, context, trackingRef);
2049
+ if (verifiedFetchedOid !== target.remoteHeadOid) throw new Error("Link branch cancelled: fetched remote ref changed");
2050
+ upstreamAttempted = true;
2051
+ await setBranchUpstream(pi, context, upstreamTarget);
2052
+ await verifyBranchUpstream(pi, context, upstreamTarget);
2053
+ const discovery = await loadCurrentPullRequest(pi, context);
2054
+ if (discovery.kind !== "current" || !sameLinkedPullRequest(inferred, discovery.pullRequest)) {
2055
+ throw new Error("Link branch failed: configured pull request does not match inferred target");
2056
+ }
2057
+ return discovery.pullRequest;
2058
+ } catch (error) {
2059
+ const rollbackContext = { cwd: context.cwd, signal: new AbortController().signal };
2060
+ try {
2061
+ await restoreLinkState(
2062
+ pi,
2063
+ rollbackContext,
2064
+ target,
2065
+ linkConfiguration,
2066
+ trackingRef,
2067
+ trackingOid,
2068
+ upstreamAttempted,
2069
+ fetchAttempted,
2070
+ );
2071
+ } catch {
2072
+ throw new Error("Link branch failed and rollback was incomplete");
2073
+ }
2074
+ throw error;
1856
2075
  }
1857
- throw error;
1858
- }
2076
+ }, { agentDir: options.agentDir, signal: context.signal });
1859
2077
  }