@henryqw/pi-pr 3.1.10 → 4.0.3

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}}}}";
@@ -103,6 +105,24 @@ export type PullRequestObservation = {
103
105
 
104
106
  export type PullRequestLoadContext = Pick<ExtensionContext, "cwd" | "signal">;
105
107
 
108
+ export type PullRequestCreationPreflight = {
109
+ head: string;
110
+ base: {
111
+ host: string;
112
+ repository: string;
113
+ fetchSource: string;
114
+ ref: string;
115
+ oid: string;
116
+ mergeBase: string;
117
+ };
118
+ ahead: number;
119
+ };
120
+
121
+ type CreationIdentity = {
122
+ target: PullRequestTarget;
123
+ head: string;
124
+ };
125
+
106
126
  type CommandOutput = {
107
127
  stdout: string;
108
128
  stderr: string;
@@ -149,7 +169,7 @@ type LinkConfiguration = {
149
169
  mirror: string[];
150
170
  };
151
171
 
152
- type SearchPullRequest = {
172
+ export type PullRequestCandidate = {
153
173
  number: number;
154
174
  url: URL;
155
175
  lifecycle: PullRequestLifecycle;
@@ -159,6 +179,37 @@ type SearchPullRequest = {
159
179
  headOid: string;
160
180
  };
161
181
 
182
+ type SearchPullRequest = PullRequestCandidate;
183
+
184
+ export type PullRequestPublication = {
185
+ number: number;
186
+ url: URL;
187
+ lifecycle: PullRequestLifecycle;
188
+ base: { repository: string; ref: string };
189
+ head: PullRequestRef;
190
+ title: string;
191
+ body: string;
192
+ };
193
+
194
+ export type ValidatedRemoteAuthority = {
195
+ fetchSource: string;
196
+ host: string;
197
+ repository: string;
198
+ };
199
+
200
+ export type BranchUpstreamTarget = {
201
+ branch: string;
202
+ remote: string;
203
+ ref: string;
204
+ fetchSource: string;
205
+ remoteOid: string;
206
+ };
207
+
208
+ export type BranchUpstreamConfiguration = {
209
+ remote: string[];
210
+ merge: string[];
211
+ };
212
+
162
213
  type SearchSelection =
163
214
  | { kind: "candidate"; candidate: SearchPullRequest; pullRequest: ListedPullRequest | null }
164
215
  | { kind: "none" }
@@ -179,6 +230,11 @@ type RulesetBranchPolicy = {
179
230
  allowedMergeMethods: MergeMethod[] | null;
180
231
  };
181
232
 
233
+ type CheckState = {
234
+ state: string;
235
+ diagnosableFailure: boolean;
236
+ };
237
+
182
238
  type ListedPullRequest = {
183
239
  id: string;
184
240
  number: number;
@@ -190,7 +246,7 @@ type ListedPullRequest = {
190
246
  mergeable: "MERGEABLE" | "CONFLICTING" | "UNKNOWN";
191
247
  mergeStateStatus: "BEHIND" | "BLOCKED" | "CLEAN" | "DIRTY" | "DRAFT" | "HAS_HOOKS" | "UNKNOWN" | "UNSTABLE";
192
248
  reviewDecision: "APPROVED" | "CHANGES_REQUESTED" | "REVIEW_REQUIRED" | null;
193
- checkStates: string[];
249
+ checkStates: CheckState[];
194
250
  };
195
251
 
196
252
  function fail(action: string, reason: string): never {
@@ -299,6 +355,12 @@ function parseCommandOutput(value: unknown, action: string): CommandOutput {
299
355
  return { stdout, stderr, code, killed };
300
356
  }
301
357
 
358
+ function hasExactKeys(value: Record<string, unknown>, keys: string[]): boolean {
359
+ const actual = Object.keys(value).sort();
360
+ const expected = [...keys].sort();
361
+ return actual.length === expected.length && actual.every((key, index) => key === expected[index]);
362
+ }
363
+
302
364
  async function invoke(
303
365
  pi: Pick<ExtensionAPI, "exec">,
304
366
  context: PullRequestLoadContext,
@@ -536,27 +598,64 @@ function checkOutcome(state: string): "failure" | "success" | "running" {
536
598
  return "running";
537
599
  }
538
600
 
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);
601
+ function canonicalPositiveDecimal(value: string): boolean {
602
+ const parsed = Number(value);
603
+ return Number.isSafeInteger(parsed) && parsed > 0 && String(parsed) === value;
604
+ }
605
+
606
+ function isActionsJobUrl(value: unknown, pullRequestUrl: URL, repository: string): boolean {
607
+ if (typeof value !== "string" || !value) return false;
608
+ let url: URL;
609
+ try {
610
+ url = new URL(value);
611
+ } catch {
612
+ return false;
613
+ }
614
+ if (url.protocol !== "https:" || url.username || url.password || url.port || url.search || url.hash ||
615
+ url.hostname.toLowerCase() !== pullRequestUrl.hostname.toLowerCase()) return false;
616
+ const parts = url.pathname.split("/").filter(Boolean);
617
+ const expected = repository.split("/");
618
+ return parts.length === 7 && expected.length === 2 &&
619
+ parts[0]!.toLowerCase() === expected[0]!.toLowerCase() &&
620
+ parts[1]!.toLowerCase() === expected[1]!.toLowerCase() &&
621
+ parts[2] === "actions" && parts[3] === "runs" && canonicalPositiveDecimal(parts[4]!) &&
622
+ parts[5] === "job" && canonicalPositiveDecimal(parts[6]!);
623
+ }
624
+
625
+ function isDiagnosableActionsCheck(check: Record<string, unknown>, pullRequestUrl: URL, repository: string): boolean {
626
+ const workflowName = check.workflowName;
627
+ return typeof workflowName === "string" && !!workflowName && workflowName.trim() === workflowName &&
628
+ !/\p{Cc}/u.test(workflowName) && isActionsJobUrl(check.detailsUrl, pullRequestUrl, repository);
629
+ }
630
+
631
+ function checkState(value: unknown, pullRequestUrl: URL, repository: string): CheckState {
632
+ if (!isRecord(value) || (value.__typename !== "CheckRun" && value.__typename !== "StatusContext")) {
633
+ return fail("Find pull requests", "invalid statusCheckRollup");
634
+ }
635
+ const conclusion = value.__typename === "CheckRun" ? optionalCheckState(value, "conclusion") : null;
636
+ const state = value.__typename === "StatusContext" ? optionalCheckState(value, "state") : null;
637
+ const status = value.__typename === "CheckRun" ? optionalCheckState(value, "status") : null;
638
+ const states = [conclusion, state, status].filter((candidate): candidate is string => candidate !== null);
545
639
  if (!states.length) fail("Find pull requests", "invalid statusCheckRollup");
546
640
 
547
641
  // COMPLETED describes a check run's lifecycle; its conclusion gives the outcome.
548
- const outcomes = states.filter((value) => value !== "COMPLETED").map(checkOutcome);
642
+ const outcomes = states.filter((candidate) => candidate !== "COMPLETED").map(checkOutcome);
549
643
  if (
550
644
  new Set(outcomes).size > 1 ||
551
645
  (states.includes("COMPLETED") && outcomes.includes("running"))
552
646
  ) fail("Find pull requests", "invalid statusCheckRollup");
553
- return conclusion ?? state ?? status ?? fail("Find pull requests", "invalid statusCheckRollup");
647
+ const selected = conclusion ?? state ?? status ?? fail("Find pull requests", "invalid statusCheckRollup");
648
+ return {
649
+ state: selected,
650
+ diagnosableFailure: FAILED_CHECK_STATES.has(selected) && value.__typename === "CheckRun" &&
651
+ isDiagnosableActionsCheck(value, pullRequestUrl, repository),
652
+ };
554
653
  }
555
654
 
556
- function checkStates(value: unknown): string[] {
655
+ function checkStates(value: unknown, pullRequestUrl: URL, repository: string): CheckState[] {
557
656
  if (value === null) return [];
558
657
  if (!Array.isArray(value)) fail("Find pull requests", "invalid statusCheckRollup");
559
- return value.map(checkState);
658
+ return value.map((check) => checkState(check, pullRequestUrl, repository));
560
659
  }
561
660
 
562
661
  function listedPullRequest(value: unknown): ListedPullRequest | null {
@@ -589,7 +688,7 @@ function listedPullRequest(value: unknown): ListedPullRequest | null {
589
688
  mergeable: mergeable(value.mergeable),
590
689
  mergeStateStatus: mergeStateStatus(value.mergeStateStatus),
591
690
  reviewDecision: reviewDecision(value.reviewDecision),
592
- checkStates: checkStates(value.statusCheckRollup),
691
+ checkStates: checkStates(value.statusCheckRollup, parsedUrl.url, parsedUrl.repository),
593
692
  };
594
693
  }
595
694
 
@@ -634,7 +733,7 @@ function searchPullRequest(value: unknown, host: string): SearchPullRequest {
634
733
  };
635
734
  }
636
735
 
637
- function parseSearchPage(output: string, pushTarget: PushTarget): SearchPage | null {
736
+ function parseSearchPage(output: string, pushTarget: Pick<PushTarget, "repository" | "ref">): SearchPage | null {
638
737
  const page = parseJson(output, "Find pull requests");
639
738
  if (!isRecord(page)) fail("Find pull requests", "invalid GitHub CLI output");
640
739
  if (page.errors !== undefined) {
@@ -718,6 +817,37 @@ function parseLoadedPullRequest(output: string, expectedUrl: URL): ListedPullReq
718
817
  return candidate;
719
818
  }
720
819
 
820
+ function parsePullRequestPublication(output: string, expectedUrl: URL): PullRequestPublication {
821
+ const value = parseJson(output, "Read pull request publication");
822
+ if (!isRecord(value) || !isRecord(value.headRepository)) {
823
+ fail("Read pull request publication", "invalid GitHub CLI output");
824
+ }
825
+ const number = value.number;
826
+ if (typeof number !== "number" || !Number.isSafeInteger(number) || number <= 0) {
827
+ fail("Read pull request publication", "invalid number");
828
+ }
829
+ const parsedUrl = parsePullRequestUrl(value.url, number);
830
+ if (parsedUrl.url.href !== expectedUrl.href) {
831
+ fail("Read pull request publication", "response does not match candidate url");
832
+ }
833
+ return {
834
+ number,
835
+ url: parsedUrl.url,
836
+ lifecycle: lifecycle(value.state),
837
+ base: {
838
+ repository: parsedUrl.repository,
839
+ ref: text(value.baseRefName, "Read pull request publication", "baseRefName"),
840
+ },
841
+ head: {
842
+ repository: repositoryName(value.headRepository.nameWithOwner, "Read pull request publication", "headRepository.nameWithOwner"),
843
+ ref: text(value.headRefName, "Read pull request publication", "headRefName"),
844
+ oid: oid(value.headRefOid, "Read pull request publication", "headRefOid"),
845
+ },
846
+ title: text(value.title, "Read pull request publication", "title"),
847
+ body: typeof value.body === "string" ? value.body : fail("Read pull request publication", "invalid body"),
848
+ };
849
+ }
850
+
721
851
  function selectPullRequest(
722
852
  candidates: ListedPullRequest[],
723
853
  pushTarget: PushTarget,
@@ -745,13 +875,19 @@ function selectPullRequest(
745
875
  return historical[0] ?? null;
746
876
  }
747
877
 
748
- function ciStatus(states: string[]): CiStatus {
749
- if (!states.length) return "none";
878
+ function ciStatus(checks: CheckState[]): CiStatus {
879
+ if (!checks.length) return "none";
880
+ let failed = false;
750
881
  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;
882
+ for (const check of checks) {
883
+ if (FAILED_CHECK_STATES.has(check.state)) {
884
+ if (!check.diagnosableFailure) return "failure-blocked";
885
+ failed = true;
886
+ } else if (!SUCCESSFUL_CHECK_STATES.has(check.state)) {
887
+ running = true;
888
+ }
754
889
  }
890
+ if (failed) return "failure";
755
891
  return running ? "running" : "success";
756
892
  }
757
893
 
@@ -821,7 +957,10 @@ function parseUnresolvedReviewThreads(output: string): number {
821
957
  return total;
822
958
  }
823
959
 
824
- function parseBaseRefOid(output: string, candidate: ListedPullRequest): string {
960
+ function parseBaseRefAuthority(
961
+ output: string,
962
+ expected: { repository: string; ref: string },
963
+ ): string {
825
964
  const value = parseJson(output, "Read base ref");
826
965
  if (!isRecord(value)) fail("Read base ref", "invalid GitHub CLI output");
827
966
  if (value.errors !== undefined) {
@@ -834,12 +973,16 @@ function parseBaseRefOid(output: string, candidate: ListedPullRequest): string {
834
973
  }
835
974
  if (
836
975
  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
976
+ normalizeRepository(expected.repository) ||
977
+ text(repository.ref.name, "Read base ref", "ref") !== expected.ref
839
978
  ) fail("Read base ref", "response does not match pull request base");
840
979
  return oid(repository.ref.target.oid, "Read base ref", "target OID");
841
980
  }
842
981
 
982
+ function parseBaseRefOid(output: string, candidate: ListedPullRequest): string {
983
+ return parseBaseRefAuthority(output, candidate.base);
984
+ }
985
+
843
986
  function parseLegacyBaseBranchPolicy(output: string, candidate: ListedPullRequest): boolean {
844
987
  const value = parseJson(output, "Read base branch policy");
845
988
  if (!isRecord(value)) fail("Read base branch policy", "invalid GitHub CLI output");
@@ -936,27 +1079,232 @@ function parseMergeMethodSettings(output: string, rulesetMethods: MergeMethod[]
936
1079
  return { allowedMergeMethods, viewerDefaultMergeMethod };
937
1080
  }
938
1081
 
939
- export async function hasLocalCommit(
1082
+ function validatedCreationTarget(target: PullRequestTarget): PullRequestTarget {
1083
+ if (!isRecord(target)) fail("Read creation target", "invalid target");
1084
+ if (target.provenance !== "configured" && target.provenance !== "inferred") {
1085
+ fail("Read creation target", "invalid provenance");
1086
+ }
1087
+ return {
1088
+ provenance: target.provenance,
1089
+ branch: text(target.branch, "Read creation target", "branch"),
1090
+ remote: text(target.remote, "Read creation target", "remote"),
1091
+ ref: text(target.ref, "Read creation target", "ref"),
1092
+ repository: repositoryName(target.repository, "Read creation target", "repository"),
1093
+ host: text(target.host, "Read creation target", "host").toLowerCase(),
1094
+ fetchSource: text(target.fetchSource, "Read creation target", "fetch source"),
1095
+ remoteOid: target.remoteOid === null ? null : oid(target.remoteOid, "Read creation target", "remote OID"),
1096
+ };
1097
+ }
1098
+
1099
+ function sameCreationTarget(left: PullRequestTarget, right: PullRequestTarget): boolean {
1100
+ return left.provenance === right.provenance && left.branch === right.branch &&
1101
+ left.remote === right.remote && left.ref === right.ref &&
1102
+ normalizeRepository(left.repository) === normalizeRepository(right.repository) &&
1103
+ left.host === right.host && left.fetchSource === right.fetchSource && left.remoteOid === right.remoteOid;
1104
+ }
1105
+
1106
+ async function validateCreationRef(
940
1107
  pi: Pick<ExtensionAPI, "exec">,
941
1108
  context: PullRequestLoadContext,
942
- ): Promise<boolean> {
1109
+ ref: string,
1110
+ ): Promise<string> {
1111
+ const requested = text(ref, "Validate creation base", "ref");
1112
+ const checked = singleLine(
1113
+ (await execute(pi, context, "Validate creation base", "git", ["check-ref-format", "--branch", requested])).stdout,
1114
+ "Validate creation base",
1115
+ "ref",
1116
+ );
1117
+ if (checked !== requested) fail("Validate creation base", "ref changed");
1118
+ return requested;
1119
+ }
1120
+
1121
+ async function captureCreationIdentity(
1122
+ pi: Pick<ExtensionAPI, "exec">,
1123
+ context: PullRequestLoadContext,
1124
+ target: PullRequestTarget,
1125
+ ): Promise<CreationIdentity> {
1126
+ const validatedTarget = validatedCreationTarget(target);
943
1127
  const branch = singleLine(
944
- (await execute(pi, context, "Read current branch", "git", ["branch", "--show-current"])).stdout,
945
- "Read current branch",
1128
+ (await execute(pi, context, "Read creation branch", "git", ["branch", "--show-current"])).stdout,
1129
+ "Read creation branch",
946
1130
  "branch",
947
1131
  );
948
- const output = (await execute(pi, context, "Read branch history", "git", [
949
- "reflog",
950
- "show",
951
- "--format=%H",
952
- `refs/heads/${branch}`,
953
- ])).stdout.replace(/\r\n/g, "\n");
954
- const entries = output.split("\n");
955
- if (entries.at(-1) === "") entries.pop();
956
- if (!entries.length) fail("Read branch history", "missing branch creation entry");
957
- const commits = entries.map((entry) => oid(entry, "Read branch history", "commit"));
958
- // ponytail: reflog expiry can hide old branch history; resolve the PR base if this becomes observable.
959
- return commits[0] !== commits.at(-1);
1132
+ if (branch !== validatedTarget.branch) fail("Read creation branch", "branch changed");
1133
+ await validateCreationRef(pi, context, branch);
1134
+ const head = oid(singleLine(
1135
+ (await execute(pi, context, "Read creation HEAD", "git", ["rev-parse", "--verify", "HEAD^{commit}"])).stdout,
1136
+ "Read creation HEAD",
1137
+ "OID",
1138
+ ), "Read creation HEAD", "OID");
1139
+ return { target: validatedTarget, head };
1140
+ }
1141
+
1142
+ async function readConfiguredCreationBaseRef(
1143
+ pi: Pick<ExtensionAPI, "exec">,
1144
+ context: PullRequestLoadContext,
1145
+ branch: string,
1146
+ ): Promise<string | null> {
1147
+ const result = await invoke(pi, context, "Read creation base configuration", "git", [
1148
+ "config", "--get-all", `branch.${branch}.gh-merge-base`,
1149
+ ]);
1150
+ if (result.killed) commandFailure("Read creation base configuration", result);
1151
+ if (result.code === 1 && result.stdout === "" && result.stderr === "") return null;
1152
+ if (result.code !== 0) commandFailure("Read creation base configuration", result);
1153
+ if (result.stderr !== "") fail("Read creation base configuration", "unexpected diagnostic");
1154
+ const values = lines(result.stdout, "Read creation base configuration", "base ref");
1155
+ if (values.length !== 1) fail("Read creation base configuration", "multiple base refs");
1156
+ return await validateCreationRef(pi, context, values[0]!);
1157
+ }
1158
+
1159
+ function parseDefaultCreationBaseRef(output: string): string {
1160
+ const value = parseJson(output, "Read creation default branch");
1161
+ if (!isRecord(value) || !hasExactKeys(value, ["defaultBranchRef"]) || !isRecord(value.defaultBranchRef) ||
1162
+ !hasExactKeys(value.defaultBranchRef, ["name"])) {
1163
+ fail("Read creation default branch", "invalid GitHub CLI output");
1164
+ }
1165
+ return text(value.defaultBranchRef.name, "Read creation default branch", "default branch ref");
1166
+ }
1167
+
1168
+ async function readDefaultCreationBaseRef(
1169
+ pi: Pick<ExtensionAPI, "exec">,
1170
+ context: PullRequestLoadContext,
1171
+ origin: { repository: PushRepository },
1172
+ ): Promise<string> {
1173
+ const result = await execute(pi, context, "Read creation default branch", "gh", [
1174
+ "repo", "view", `${origin.repository.host}/${origin.repository.nameWithOwner}`, "--json", "defaultBranchRef",
1175
+ ]);
1176
+ if (result.stderr !== "") fail("Read creation default branch", "unexpected diagnostic");
1177
+ return await validateCreationRef(pi, context, parseDefaultCreationBaseRef(result.stdout));
1178
+ }
1179
+
1180
+ function parseCreationRepositoryLineage(output: string, expected: PushRepository): string {
1181
+ const value = parseJson(output, "Read creation repository");
1182
+ if (!isRecord(value)) fail("Read creation repository", "invalid GitHub CLI output");
1183
+ const fullName = repositoryName(value.full_name, "Read creation repository", "full_name");
1184
+ const url = parseHttpUrl(value.html_url, "Read creation repository", "html_url");
1185
+ if (
1186
+ normalizeRepository(fullName) !== expected.normalizedName || url.protocol !== "https:" || url.port ||
1187
+ url.hostname.toLowerCase() !== expected.host || url.pathname.toLowerCase() !== `/${expected.normalizedName}`
1188
+ ) fail("Read creation repository", "response does not match repository");
1189
+ const source = value.source;
1190
+ if (source !== undefined && source !== null && !isRecord(source)) {
1191
+ fail("Read creation repository", "invalid source");
1192
+ }
1193
+ const sourceName = source === undefined || source === null
1194
+ ? fullName
1195
+ : repositoryName(source.full_name, "Read creation repository", "source.full_name");
1196
+ return normalizeRepository(sourceName);
1197
+ }
1198
+
1199
+ async function assertCreationRepositoryRelation(
1200
+ pi: Pick<ExtensionAPI, "exec">,
1201
+ context: PullRequestLoadContext,
1202
+ target: PullRequestTarget,
1203
+ origin: { repository: PushRepository },
1204
+ baseRef: string,
1205
+ ): Promise<void> {
1206
+ if (target.host !== origin.repository.host) {
1207
+ fail("Read creation repository", "base and head hosts do not match");
1208
+ }
1209
+ if (normalizeRepository(target.repository) === origin.repository.normalizedName) {
1210
+ if (target.ref === baseRef) fail("Read creation repository", "head and base refs match");
1211
+ return;
1212
+ }
1213
+ const read = async (repository: PushRepository): Promise<string> => {
1214
+ const [owner, name] = repository.nameWithOwner.split("/");
1215
+ const result = await execute(pi, context, "Read creation repository", "gh", [
1216
+ "api", "--hostname", origin.repository.host, `repos/${owner}/${name}`,
1217
+ ]);
1218
+ if (result.stderr !== "") fail("Read creation repository", "unexpected diagnostic");
1219
+ return parseCreationRepositoryLineage(result.stdout, repository);
1220
+ };
1221
+ const originSource = await read(origin.repository);
1222
+ const targetSource = await read({
1223
+ nameWithOwner: target.repository,
1224
+ normalizedName: normalizeRepository(target.repository),
1225
+ host: target.host,
1226
+ });
1227
+ if (originSource !== targetSource) fail("Read creation repository", "base and head are unrelated");
1228
+ }
1229
+
1230
+ function parseCreationAhead(output: string): number {
1231
+ const value = singleLine(output, "Count creation commits", "ahead count");
1232
+ if (!/^(?:0|[1-9][0-9]*)$/.test(value)) fail("Count creation commits", "invalid ahead count");
1233
+ const ahead = Number(value);
1234
+ if (!Number.isSafeInteger(ahead) || ahead < 0) fail("Count creation commits", "invalid ahead count");
1235
+ return ahead;
1236
+ }
1237
+
1238
+ async function preflightCreation(
1239
+ pi: Pick<ExtensionAPI, "exec">,
1240
+ context: PullRequestLoadContext,
1241
+ target: PullRequestTarget,
1242
+ explicitBaseRef: string | undefined,
1243
+ identity: CreationIdentity,
1244
+ ): Promise<PullRequestCreationPreflight> {
1245
+ const validatedTarget = validatedCreationTarget(target);
1246
+ if (!sameCreationTarget(identity.target, validatedTarget)) {
1247
+ fail("Read creation target", "target changed");
1248
+ }
1249
+ const origin = await readRemoteAuthority(pi, context, "origin", true);
1250
+ if (!origin) fail("Read creation repository", "origin is unavailable");
1251
+ const configuredBaseRef = explicitBaseRef === undefined
1252
+ ? await readConfiguredCreationBaseRef(pi, context, identity.target.branch)
1253
+ : await validateCreationRef(pi, context, explicitBaseRef);
1254
+ const baseRef = configuredBaseRef ?? await readDefaultCreationBaseRef(pi, context, origin);
1255
+ await assertCreationRepositoryRelation(pi, context, identity.target, origin, baseRef);
1256
+ const trackingRef = `refs/remotes/origin/${baseRef}`;
1257
+ await execute(pi, context, "Fetch creation base", "git", [
1258
+ "fetch", "--no-write-fetch-head", "--no-tags", "--no-recurse-submodules", "--",
1259
+ origin.fetchSource, `+refs/heads/${baseRef}:${trackingRef}`,
1260
+ ]);
1261
+ const baseOid = oid(singleLine(
1262
+ (await execute(pi, context, "Read creation base", "git", ["rev-parse", "--verify", `${trackingRef}^{commit}`])).stdout,
1263
+ "Read creation base",
1264
+ "OID",
1265
+ ), "Read creation base", "OID");
1266
+ const mergeBase = oid(singleLine(
1267
+ (await execute(pi, context, "Find creation merge base", "git", ["merge-base", identity.head, baseOid])).stdout,
1268
+ "Find creation merge base",
1269
+ "OID",
1270
+ ), "Find creation merge base", "OID");
1271
+ const ahead = parseCreationAhead((await execute(pi, context, "Count creation commits", "git", [
1272
+ "rev-list", "--count", `${mergeBase}..${identity.head}`,
1273
+ ])).stdout);
1274
+ return {
1275
+ head: identity.head,
1276
+ base: {
1277
+ host: origin.repository.host,
1278
+ repository: origin.repository.nameWithOwner,
1279
+ fetchSource: origin.fetchSource,
1280
+ ref: baseRef,
1281
+ oid: baseOid,
1282
+ mergeBase,
1283
+ },
1284
+ ahead,
1285
+ };
1286
+ }
1287
+
1288
+ export async function preflightPullRequestCreation(
1289
+ pi: Pick<ExtensionAPI, "exec">,
1290
+ context: PullRequestLoadContext,
1291
+ target: PullRequestTarget,
1292
+ explicitBaseRef?: string,
1293
+ ): Promise<PullRequestCreationPreflight> {
1294
+ const identity = await captureCreationIdentity(pi, context, target);
1295
+ return await preflightCreation(pi, context, target, explicitBaseRef, identity);
1296
+ }
1297
+
1298
+ async function creationDiscovery(
1299
+ pi: Pick<ExtensionAPI, "exec">,
1300
+ context: PullRequestLoadContext,
1301
+ target: PullRequestTarget,
1302
+ explicitBaseRef: string | undefined,
1303
+ identity?: CreationIdentity,
1304
+ ): Promise<CurrentPullRequestDiscovery> {
1305
+ const captured = identity ?? await captureCreationIdentity(pi, context, target);
1306
+ const preflight = await preflightCreation(pi, context, target, explicitBaseRef, captured);
1307
+ return { kind: "none", creationTarget: target, branch: { ahead: preflight.ahead } };
960
1308
  }
961
1309
 
962
1310
  async function readRemoteAuthority(
@@ -1000,6 +1348,20 @@ async function readRemoteAuthority(
1000
1348
  }
1001
1349
  }
1002
1350
 
1351
+ export async function readValidatedRemoteAuthority(
1352
+ pi: Pick<ExtensionAPI, "exec">,
1353
+ context: PullRequestLoadContext,
1354
+ remote: string,
1355
+ ): Promise<ValidatedRemoteAuthority> {
1356
+ const authority = await readRemoteAuthority(pi, context, text(remote, "Read push remotes", "remote"), true);
1357
+ if (!authority) fail("Read push target", "invalid remote authority");
1358
+ return {
1359
+ fetchSource: authority.fetchSource,
1360
+ host: authority.repository.host,
1361
+ repository: authority.repository.nameWithOwner,
1362
+ };
1363
+ }
1364
+
1003
1365
  async function readRemoteHeadOid(
1004
1366
  pi: Pick<ExtensionAPI, "exec">,
1005
1367
  context: PullRequestLoadContext,
@@ -1039,6 +1401,22 @@ async function readConfigValues(
1039
1401
  }
1040
1402
  }
1041
1403
 
1404
+ export async function readBranchUpstreamConfiguration(
1405
+ pi: Pick<ExtensionAPI, "exec">,
1406
+ context: PullRequestLoadContext,
1407
+ branch: string,
1408
+ ): Promise<BranchUpstreamConfiguration> {
1409
+ const checkedBranch = text(branch, "Read branch upstream", "branch");
1410
+ const [remote, merge] = await Promise.all([
1411
+ readConfigValues(pi, context, `branch.${checkedBranch}.remote`),
1412
+ readConfigValues(pi, context, `branch.${checkedBranch}.merge`),
1413
+ ]);
1414
+ if (remote === null || merge === null) {
1415
+ throw new Error("Read branch upstream failed: invalid Git configuration");
1416
+ }
1417
+ return { remote, merge };
1418
+ }
1419
+
1042
1420
  async function readBooleanConfigValues(
1043
1421
  pi: Pick<ExtensionAPI, "exec">,
1044
1422
  context: PullRequestLoadContext,
@@ -1293,6 +1671,25 @@ async function readBaseRefOid(
1293
1671
  return parseBaseRefOid(result.stdout, candidate);
1294
1672
  }
1295
1673
 
1674
+ export async function readPullRequestBaseRefOid(
1675
+ pi: Pick<ExtensionAPI, "exec">,
1676
+ context: PullRequestLoadContext,
1677
+ authority: { host: string; repository: string; ref: string },
1678
+ ): Promise<string> {
1679
+ const host = text(authority.host, "Read base ref", "host").toLowerCase();
1680
+ const repository = repositoryName(authority.repository, "Read base ref", "repository");
1681
+ const ref = text(authority.ref, "Read base ref", "ref");
1682
+ const [owner, name] = repository.split("/");
1683
+ const result = await execute(pi, context, "Read base ref", "gh", [
1684
+ "api", "graphql", "--hostname", host,
1685
+ "-f", `query=${BASE_REF_QUERY}`,
1686
+ "-F", `owner=${owner}`,
1687
+ "-F", `name=${name}`,
1688
+ "-F", `qualifiedName=refs/heads/${ref}`,
1689
+ ]);
1690
+ return parseBaseRefAuthority(result.stdout, { repository, ref });
1691
+ }
1692
+
1296
1693
  async function readLegacyBaseBranchPolicy(
1297
1694
  pi: Pick<ExtensionAPI, "exec">,
1298
1695
  context: PullRequestLoadContext,
@@ -1449,35 +1846,28 @@ async function loadObservedPullRequest(
1449
1846
  return loadPullRequestDetails(pi, context, candidate, pushTarget, inspectedLocal);
1450
1847
  }
1451
1848
 
1452
- async function searchPullRequests(
1849
+ async function enumerateSearchPullRequests(
1453
1850
  pi: Pick<ExtensionAPI, "exec">,
1454
1851
  context: PullRequestLoadContext,
1455
- pushTarget: PushTarget,
1456
- ): Promise<SearchSelection> {
1457
- const [owner, name] = pushTarget.repository.nameWithOwner.split("/");
1852
+ target: Pick<PushTarget, "repository" | "ref">,
1853
+ ): Promise<SearchPullRequest[] | null> {
1854
+ const [owner, name] = target.repository.nameWithOwner.split("/");
1458
1855
  const candidates: SearchPullRequest[] = [];
1459
1856
  const cursors = new Set<string>();
1460
1857
  let totalCount: number | null = null;
1461
1858
  let endCursor: string | null = null;
1462
1859
  for (let pageIndex = 0; pageIndex < PR_DISCOVERY_MAX_PAGES; pageIndex += 1) {
1463
1860
  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}`,
1861
+ "api", "graphql", "--hostname", target.repository.host,
1862
+ "-f", `query=${PR_DISCOVERY_QUERY}`,
1863
+ "-F", `owner=${owner}`,
1864
+ "-F", `name=${name}`,
1865
+ "-F", `qualifiedName=refs/heads/${target.ref}`,
1476
1866
  ];
1477
1867
  if (endCursor !== null) args.push("-F", `endCursor=${endCursor}`);
1478
1868
  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" };
1869
+ const page = parseSearchPage(result.stdout, target);
1870
+ if (page === null) return null;
1481
1871
  if (totalCount !== null && page.totalCount !== totalCount) {
1482
1872
  fail("Find pull requests", "inconsistent search result pages");
1483
1873
  }
@@ -1495,32 +1885,77 @@ async function searchPullRequests(
1495
1885
  }
1496
1886
  const hasMore = candidates.length < totalCount;
1497
1887
  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
- }
1888
+ if (!hasMore) return candidates;
1513
1889
  if (page.endCursor === null) fail("Find pull requests", "invalid search pageInfo");
1514
1890
  endCursor = page.endCursor;
1515
1891
  }
1516
1892
  return fail("Find pull requests", "GitHub pull request result cap reached");
1517
1893
  }
1518
1894
 
1895
+ export async function findExactHeadPullRequests(
1896
+ pi: Pick<ExtensionAPI, "exec">,
1897
+ context: PullRequestLoadContext,
1898
+ target: { host: string; repository: string; ref: string },
1899
+ ): Promise<PullRequestCandidate[]> {
1900
+ const host = text(target.host, "Find pull requests", "host").toLowerCase();
1901
+ const repository = repositoryName(target.repository, "Find pull requests", "head repository");
1902
+ const ref = text(target.ref, "Find pull requests", "head ref");
1903
+ const candidates = await enumerateSearchPullRequests(pi, context, {
1904
+ repository: { host, nameWithOwner: repository, normalizedName: normalizeRepository(repository) },
1905
+ ref,
1906
+ });
1907
+ if (candidates === null) return fail("Find pull requests", "published head ref is unavailable");
1908
+ return candidates.filter((candidate) =>
1909
+ candidate.lifecycle === "open" && candidate.headRepository !== null &&
1910
+ normalizeRepository(candidate.headRepository) === normalizeRepository(repository) && candidate.headRef === ref
1911
+ );
1912
+ }
1913
+
1914
+ export async function loadPullRequestPublication(
1915
+ pi: Pick<ExtensionAPI, "exec">,
1916
+ context: PullRequestLoadContext,
1917
+ url: URL,
1918
+ ): Promise<PullRequestPublication> {
1919
+ const loaded = await execute(pi, context, "Read pull request publication", "gh", [
1920
+ "pr", "view", url.href, "--json", PR_PUBLICATION_FIELDS,
1921
+ ]);
1922
+ return parsePullRequestPublication(loaded.stdout, url);
1923
+ }
1924
+
1925
+ async function searchPullRequests(
1926
+ pi: Pick<ExtensionAPI, "exec">,
1927
+ context: PullRequestLoadContext,
1928
+ pushTarget: PushTarget,
1929
+ ): Promise<SearchSelection> {
1930
+ const candidates = await enumerateSearchPullRequests(pi, context, pushTarget);
1931
+ if (candidates === null) {
1932
+ return pushTarget.remoteHeadOid === null ? { kind: "none" } : { kind: "target-invalid" };
1933
+ }
1934
+ const selected = selectSearchPullRequest(candidates, pushTarget);
1935
+ if (selected.kind !== "candidate") return selected;
1936
+ const loaded = await execute(pi, context, "Find pull requests", "gh", [
1937
+ "pr", "view", selected.candidate.url.href, "--json", PR_FIELDS,
1938
+ ]);
1939
+ return { ...selected, pullRequest: parseLoadedPullRequest(loaded.stdout, selected.candidate.url) };
1940
+ }
1941
+
1519
1942
  export async function loadCurrentPullRequest(
1520
1943
  pi: Pick<ExtensionAPI, "exec">,
1521
1944
  context: PullRequestLoadContext,
1522
1945
  inspectedLocal?: LocalMergeSafety,
1523
1946
  observed?: unknown,
1947
+ explicitCreationBase?: string,
1948
+ ): Promise<CurrentPullRequestDiscovery> {
1949
+ return await loadCurrentPullRequestInternal(pi, context, inspectedLocal, observed, explicitCreationBase);
1950
+ }
1951
+
1952
+ async function loadCurrentPullRequestInternal(
1953
+ pi: Pick<ExtensionAPI, "exec">,
1954
+ context: PullRequestLoadContext,
1955
+ inspectedLocal: LocalMergeSafety | undefined,
1956
+ observed: unknown,
1957
+ explicitCreationBase: string | undefined,
1958
+ creationIdentity?: CreationIdentity,
1524
1959
  ): Promise<CurrentPullRequestDiscovery> {
1525
1960
  const read = await readPushTarget(pi, context);
1526
1961
  if (read.kind === "inactive") return { kind: "inactive" };
@@ -1541,10 +1976,18 @@ export async function loadCurrentPullRequest(
1541
1976
  };
1542
1977
  }
1543
1978
  if (inferred.kind === "none") {
1979
+ const target = publicTarget(inferred.target);
1980
+ if (creationIdentity === undefined) {
1981
+ const captured = await captureCreationIdentity(pi, context, target);
1982
+ return await loadCurrentPullRequestInternal(pi, context, inspectedLocal, observed, explicitCreationBase, captured);
1983
+ }
1984
+ if (!sameCreationTarget(creationIdentity.target, validatedCreationTarget(target))) {
1985
+ fail("Read creation target", "target changed");
1986
+ }
1544
1987
  if (!canLinkTarget(await readLinkConfiguration(pi, context, inferred.target), inferred.target)) {
1545
1988
  return { kind: "blocked", issue: { kind: "link-configuration", remote: inferred.target.remote } };
1546
1989
  }
1547
- return { kind: "none", creationTarget: publicTarget(inferred.target) };
1990
+ return await creationDiscovery(pi, context, target, explicitCreationBase, creationIdentity);
1548
1991
  }
1549
1992
  pushTarget = inferred.target;
1550
1993
  } else {
@@ -1641,7 +2084,10 @@ export async function loadCurrentPullRequest(
1641
2084
  }
1642
2085
  throw error;
1643
2086
  }
1644
- if (candidate === null) return { kind: "none", creationTarget: publicTarget(pushTarget) };
2087
+ if (candidate === null) {
2088
+ if (creationIdentity !== undefined) fail("Read creation target", "target changed");
2089
+ return await creationDiscovery(pi, context, publicTarget(pushTarget), explicitCreationBase);
2090
+ }
1645
2091
  }
1646
2092
 
1647
2093
  return {
@@ -1682,7 +2128,7 @@ function sameLinkedPullRequest(inferred: CurrentPullRequest, configured: Current
1682
2128
  inferred.target.remoteOid === configured.target.remoteOid;
1683
2129
  }
1684
2130
 
1685
- async function readTrackingOid(
2131
+ export async function readTrackingOid(
1686
2132
  pi: Pick<ExtensionAPI, "exec">,
1687
2133
  context: PullRequestLoadContext,
1688
2134
  trackingRef: string,
@@ -1696,6 +2142,46 @@ async function readTrackingOid(
1696
2142
  return oid(singleLine(result.stdout, "Read remote-tracking ref", "OID"), "Read remote-tracking ref", "OID");
1697
2143
  }
1698
2144
 
2145
+ export function branchTrackingRef(target: Pick<BranchUpstreamTarget, "remote" | "ref">): string {
2146
+ return `refs/remotes/${text(target.remote, "Read push remotes", "remote")}/${text(target.ref, "Read push target", "push ref")}`;
2147
+ }
2148
+
2149
+ export async function fetchBranchTrackingRef(
2150
+ pi: Pick<ExtensionAPI, "exec">,
2151
+ context: PullRequestLoadContext,
2152
+ target: BranchUpstreamTarget,
2153
+ ): Promise<void> {
2154
+ await execute(pi, context, "Fetch branch tracking ref", "git", [
2155
+ "fetch", "--no-write-fetch-head", "--no-tags", "--no-recurse-submodules",
2156
+ target.fetchSource, `+${target.remoteOid}:${branchTrackingRef(target)}`,
2157
+ ]);
2158
+ }
2159
+
2160
+ export async function setBranchUpstream(
2161
+ pi: Pick<ExtensionAPI, "exec">,
2162
+ context: PullRequestLoadContext,
2163
+ target: BranchUpstreamTarget,
2164
+ ): Promise<void> {
2165
+ await execute(pi, context, "Set branch upstream", "git", [
2166
+ "branch", `--set-upstream-to=${target.remote}/${target.ref}`, "--", target.branch,
2167
+ ]);
2168
+ }
2169
+
2170
+ export async function verifyBranchUpstream(
2171
+ pi: Pick<ExtensionAPI, "exec">,
2172
+ context: PullRequestLoadContext,
2173
+ target: BranchUpstreamTarget,
2174
+ ): Promise<void> {
2175
+ const trackingOid = await readTrackingOid(pi, context, branchTrackingRef(target));
2176
+ if (trackingOid !== target.remoteOid) throw new Error("Branch tracking ref does not match published OID");
2177
+ const configuredTarget = optionalPushReference((await execute(pi, context, "Verify push target", "git", [
2178
+ "for-each-ref", "--format=%(push:short)", `refs/heads/${target.branch}`,
2179
+ ])).stdout);
2180
+ if (configuredTarget !== `${target.remote}/${target.ref}`) {
2181
+ throw new Error("Configured push target does not match published branch");
2182
+ }
2183
+ }
2184
+
1699
2185
  async function restoreConfigValue(
1700
2186
  pi: Pick<ExtensionAPI, "exec">,
1701
2187
  context: PullRequestLoadContext,
@@ -1725,6 +2211,46 @@ async function restoreConfigValue(
1725
2211
  }
1726
2212
  }
1727
2213
 
2214
+ function sameConfigValues(left: readonly string[], right: readonly string[]): boolean {
2215
+ return left.length === right.length && left.every((value, index) => value === right[index]);
2216
+ }
2217
+
2218
+ export async function restoreBranchUpstreamConfiguration(
2219
+ pi: Pick<ExtensionAPI, "exec">,
2220
+ context: PullRequestLoadContext,
2221
+ target: Pick<BranchUpstreamTarget, "branch" | "remote" | "ref">,
2222
+ original: BranchUpstreamConfiguration,
2223
+ ): Promise<void> {
2224
+ let incomplete = false;
2225
+ for (const [key, expected, values] of [
2226
+ [`branch.${target.branch}.remote`, target.remote, original.remote],
2227
+ [`branch.${target.branch}.merge`, `refs/heads/${target.ref}`, original.merge],
2228
+ ] as const) {
2229
+ try {
2230
+ const current = await readConfigValues(pi, context, key);
2231
+ if (current === null) incomplete = true;
2232
+ else if (sameConfigValues(current, values)) continue;
2233
+ else if (current.length === 1 && current[0] === expected) {
2234
+ await restoreConfigValue(pi, context, key, expected, values);
2235
+ } else incomplete = true;
2236
+ } catch {
2237
+ incomplete = true;
2238
+ }
2239
+ }
2240
+ for (const [key, values] of [
2241
+ [`branch.${target.branch}.remote`, original.remote],
2242
+ [`branch.${target.branch}.merge`, original.merge],
2243
+ ] as const) {
2244
+ try {
2245
+ const current = await readConfigValues(pi, context, key);
2246
+ if (current === null || !sameConfigValues(current, values)) incomplete = true;
2247
+ } catch {
2248
+ incomplete = true;
2249
+ }
2250
+ }
2251
+ if (incomplete) throw new Error("Restore branch upstream failed and rollback was incomplete");
2252
+ }
2253
+
1728
2254
  async function restoreLinkState(
1729
2255
  pi: Pick<ExtensionAPI, "exec">,
1730
2256
  context: PullRequestLoadContext,
@@ -1732,32 +2258,39 @@ async function restoreLinkState(
1732
2258
  configuration: LinkConfiguration,
1733
2259
  trackingRef: string,
1734
2260
  trackingOid: string | null,
1735
- upstreamMutated: boolean,
1736
- fetchedOid: string | undefined,
2261
+ upstreamAttempted: boolean,
2262
+ fetchAttempted: boolean,
1737
2263
  ): Promise<void> {
1738
2264
  let incomplete = false;
1739
- if (upstreamMutated) {
2265
+ if (upstreamAttempted) {
1740
2266
  for (const [key, expected, original] of [
1741
2267
  [`branch.${target.branch}.remote`, target.remote, configuration.upstreamRemote],
1742
2268
  [`branch.${target.branch}.merge`, `refs/heads/${target.ref}`, configuration.upstreamMerge],
1743
2269
  ] as const) {
1744
2270
  try {
1745
- await restoreConfigValue(pi, context, key, expected, original);
2271
+ const current = await readConfigValues(pi, context, key);
2272
+ if (current === null) incomplete = true;
2273
+ else if (sameConfigValues(current, original)) continue;
2274
+ else if (current.length === 1 && current[0] === expected) {
2275
+ await restoreConfigValue(pi, context, key, expected, original);
2276
+ } else incomplete = true;
1746
2277
  } catch {
1747
2278
  incomplete = true;
1748
2279
  }
1749
2280
  }
1750
2281
  }
1751
- if (fetchedOid !== undefined) {
2282
+ if (fetchAttempted) {
1752
2283
  try {
1753
2284
  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);
2285
+ if (currentTrackingOid !== trackingOid) {
2286
+ if (target.remoteHeadOid === null || currentTrackingOid !== target.remoteHeadOid) {
2287
+ incomplete = true;
2288
+ } else {
2289
+ const args = trackingOid === null
2290
+ ? ["update-ref", "-d", trackingRef, target.remoteHeadOid]
2291
+ : ["update-ref", trackingRef, trackingOid, target.remoteHeadOid];
2292
+ await execute(pi, context, "Restore remote-tracking ref", "git", args);
2293
+ }
1761
2294
  }
1762
2295
  } catch {
1763
2296
  incomplete = true;
@@ -1770,90 +2303,84 @@ export async function linkInferredPullRequest(
1770
2303
  pi: Pick<ExtensionAPI, "exec">,
1771
2304
  context: PullRequestLoadContext,
1772
2305
  inferred: CurrentPullRequest,
2306
+ options: { agentDir?: string } = {},
1773
2307
  ): Promise<CurrentPullRequest> {
1774
2308
  if (inferred.target.provenance !== "inferred" || inferred.lifecycle !== "open") {
1775
2309
  throw new Error("Link branch failed: pull request is not an open inferred target");
1776
2310
  }
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");
2311
+ return await withWorktreeLock(context.cwd, async () => {
2312
+ const freshDiscovery = await loadCurrentPullRequest(pi, context);
2313
+ if (
2314
+ freshDiscovery.kind !== "current" ||
2315
+ freshDiscovery.pullRequest.target.provenance !== "inferred" ||
2316
+ !samePullRequestSnapshot(inferred, freshDiscovery.pullRequest)
2317
+ ) throw new Error("Link branch cancelled: inferred pull request context changed");
2318
+ inferred = freshDiscovery.pullRequest;
2319
+ const target: PushTarget = {
2320
+ provenance: "inferred",
2321
+ branch: inferred.target.branch,
2322
+ remote: inferred.target.remote,
2323
+ ref: inferred.target.ref,
2324
+ fetchSource: inferred.target.fetchSource,
2325
+ remoteHeadOid: inferred.target.remoteOid,
2326
+ repository: {
2327
+ nameWithOwner: inferred.target.repository,
2328
+ normalizedName: normalizeRepository(inferred.target.repository),
2329
+ host: inferred.target.host,
2330
+ },
2331
+ };
2332
+ const linkConfiguration = await readLinkConfiguration(pi, context, target);
2333
+ if (target.remoteHeadOid === null || !linkConfiguration || !canLinkTarget(linkConfiguration, target)) {
2334
+ throw new Error("Link branch cancelled: target configuration changed");
1825
2335
  }
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", [
2336
+ const pushReference = optionalPushReference((await execute(pi, context, "Read push target", "git", [
1831
2337
  "for-each-ref", "--format=%(push:short)", `refs/heads/${target.branch}`,
1832
2338
  ])).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 };
2339
+ if (pushReference !== null) throw new Error("Link branch cancelled: push target is no longer empty");
2340
+ const remoteHeadOid = await readRemoteHeadOid(pi, context, target.fetchSource, target.ref);
2341
+ if (remoteHeadOid !== target.remoteHeadOid) throw new Error("Link branch cancelled: remote ref changed");
2342
+
2343
+ const trackingRef = `refs/remotes/${target.remote}/${target.ref}`;
2344
+ const trackingOid = await readTrackingOid(pi, context, trackingRef);
2345
+ let fetchAttempted = false;
2346
+ let upstreamAttempted = false;
1843
2347
  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");
2348
+ const upstreamTarget: BranchUpstreamTarget = {
2349
+ branch: target.branch,
2350
+ remote: target.remote,
2351
+ ref: target.ref,
2352
+ fetchSource: target.fetchSource,
2353
+ remoteOid: target.remoteHeadOid,
2354
+ };
2355
+ fetchAttempted = true;
2356
+ await fetchBranchTrackingRef(pi, context, upstreamTarget);
2357
+ const verifiedFetchedOid = await readTrackingOid(pi, context, trackingRef);
2358
+ if (verifiedFetchedOid !== target.remoteHeadOid) throw new Error("Link branch cancelled: fetched remote ref changed");
2359
+ upstreamAttempted = true;
2360
+ await setBranchUpstream(pi, context, upstreamTarget);
2361
+ await verifyBranchUpstream(pi, context, upstreamTarget);
2362
+ const discovery = await loadCurrentPullRequest(pi, context);
2363
+ if (discovery.kind !== "current" || !sameLinkedPullRequest(inferred, discovery.pullRequest)) {
2364
+ throw new Error("Link branch failed: configured pull request does not match inferred target");
2365
+ }
2366
+ return discovery.pullRequest;
2367
+ } catch (error) {
2368
+ const rollbackContext = { cwd: context.cwd, signal: new AbortController().signal };
2369
+ try {
2370
+ await restoreLinkState(
2371
+ pi,
2372
+ rollbackContext,
2373
+ target,
2374
+ linkConfiguration,
2375
+ trackingRef,
2376
+ trackingOid,
2377
+ upstreamAttempted,
2378
+ fetchAttempted,
2379
+ );
2380
+ } catch {
2381
+ throw new Error("Link branch failed and rollback was incomplete");
2382
+ }
2383
+ throw error;
1856
2384
  }
1857
- throw error;
1858
- }
2385
+ }, { agentDir: options.agentDir, signal: context.signal });
1859
2386
  }