1use std::path::Path;
2use std::process::Command;
3
4use anyhow::{Context, Result, bail};
5use serde_json::Value;
6
7use crate::git::{RangeSpec, Side, git};
8
9#[derive(Debug, Clone, PartialEq, Eq)]
11pub enum Repo {
12 GitHub {
13 owner: String,
14 name: String,
15 },
16 AzureDevOps {
17 org: String,
18 project: String,
19 name: String,
20 },
21}
22
23impl Repo {
24 fn same_as(&self, other: &Repo) -> bool {
25 let eq = |a: &str, b: &str| a.eq_ignore_ascii_case(b);
26 match (self, other) {
27 (Repo::GitHub { owner: a, name: b }, Repo::GitHub { owner: c, name: d }) => {
28 eq(a, c) && eq(b, d)
29 }
30 (
31 Repo::AzureDevOps {
32 org: a,
33 project: b,
34 name: c,
35 },
36 Repo::AzureDevOps {
37 org: d,
38 project: e,
39 name: f,
40 },
41 ) => eq(a, d) && eq(b, e) && eq(c, f),
42 _ => false,
43 }
44 }
45
46 pub fn display(&self) -> String {
47 match self {
48 Repo::GitHub { owner, name } => format!("github.com/{owner}/{name}"),
49 Repo::AzureDevOps { org, project, name } => {
50 format!("dev.azure.com/{org}/{project}/{name}")
51 }
52 }
53 }
54}
55
56#[derive(Debug, Clone, PartialEq, Eq)]
57pub struct PullRequest {
58 pub repo: Repo,
59 pub number: u64,
60}
61
62impl PullRequest {
63 pub fn review_key(&self) -> String {
65 format!("{} PR {}", self.repo.display(), self.number)
66 }
67
68 pub fn head_ref(&self) -> String {
70 format!("refs/declutter/pr/{}/head", self.number)
71 }
72}
73
74#[derive(Debug, Clone, PartialEq, Eq)]
76pub struct PrRefs {
77 pub base: String,
78 pub head: String,
79}
80
81pub trait PrLookup {
83 fn refs(&self, pr: &PullRequest) -> Result<PrRefs>;
84}
85
86pub struct CliLookup;
89
90impl PrLookup for CliLookup {
91 fn refs(&self, pr: &PullRequest) -> Result<PrRefs> {
92 match &pr.repo {
93 Repo::GitHub { owner, name } => {
94 let json = run_json(
95 "gh",
96 &[
97 "pr",
98 "view",
99 &pr.number.to_string(),
100 "--repo",
101 &format!("{owner}/{name}"),
102 "--json",
103 "baseRefName",
104 ],
105 )?;
106 let base = json["baseRefName"]
107 .as_str()
108 .context("`gh pr view` returned no baseRefName")?;
109 Ok(PrRefs {
110 base: format!("refs/heads/{base}"),
111 head: format!("refs/pull/{}/head", pr.number),
113 })
114 }
115 Repo::AzureDevOps { org, .. } => {
116 let json = run_json(
117 "az",
118 &[
119 "repos",
120 "pr",
121 "show",
122 "--id",
123 &pr.number.to_string(),
124 "--org",
125 &format!("https://dev.azure.com/{org}"),
126 "--detect",
127 "false",
128 "--output",
129 "json",
130 ],
131 )?;
132 let field = |key: &str| {
133 json[key]
134 .as_str()
135 .map(str::to_string)
136 .with_context(|| format!("`az repos pr show` returned no {key}"))
137 };
138 Ok(PrRefs {
139 base: field("targetRefName")?,
140 head: field("sourceRefName")?,
141 })
142 }
143 }
144 }
145}
146
147fn run_json(program: &str, args: &[&str]) -> Result<Value> {
148 let output = Command::new(program)
149 .args(args)
150 .output()
151 .with_context(|| format!("failed to run `{program}`; is it installed and logged in?"))?;
152 if !output.status.success() {
153 let login = if program == "gh" {
154 "gh auth login"
155 } else {
156 "az login"
157 };
158 bail!(
159 "`{program} {}` failed: {}\n(if this is a sign-in problem, run `{login}` and try again)",
160 args.join(" "),
161 String::from_utf8_lossy(&output.stderr).trim()
162 );
163 }
164 serde_json::from_slice(&output.stdout)
165 .with_context(|| format!("`{program}` returned invalid JSON"))
166}
167
168pub fn parse_pr_url(url: &str) -> Option<PullRequest> {
170 let url = url.split(['?', '#']).next()?.trim_end_matches('/');
171 let rest = url
172 .strip_prefix("https://")
173 .or_else(|| url.strip_prefix("http://"))?;
174 let parts: Vec<String> = rest.split('/').map(percent_decode).collect();
175 let parts: Vec<&str> = parts.iter().map(String::as_str).collect();
176 match parts.as_slice() {
177 ["github.com", owner, name, "pull", number, ..] => Some(PullRequest {
178 repo: Repo::GitHub {
179 owner: owner.to_string(),
180 name: name.to_string(),
181 },
182 number: number.parse().ok()?,
183 }),
184 [
185 "dev.azure.com",
186 org,
187 project,
188 "_git",
189 name,
190 "pullrequest",
191 number,
192 ..,
193 ] => Some(PullRequest {
194 repo: azure(org, project, name),
195 number: number.parse().ok()?,
196 }),
197 [host, project, "_git", name, "pullrequest", number, ..] => Some(PullRequest {
198 repo: azure(host.strip_suffix(".visualstudio.com")?, project, name),
199 number: number.parse().ok()?,
200 }),
201 _ => None,
202 }
203}
204
205pub fn parse_remote_url(url: &str) -> Option<Repo> {
207 let url = url.trim().trim_end_matches('/');
208 let url = url.strip_suffix(".git").unwrap_or(url);
209 let path = if let Some(rest) = url.strip_prefix("git@github.com:") {
210 format!("github.com/{rest}")
211 } else if let Some(rest) = url.strip_prefix("git@ssh.dev.azure.com:v3/") {
212 format!("ssh.dev.azure.com/{rest}")
213 } else if let Some((_, rest)) = url.split_once("@vs-ssh.visualstudio.com:v3/") {
214 format!("ssh.dev.azure.com/{rest}")
215 } else {
216 let rest = url.split_once("://")?.1;
217 let rest = match rest.split_once('@') {
219 Some((_, host_and_path)) if !host_and_path.contains('@') => host_and_path,
220 _ => rest,
221 };
222 rest.to_string()
223 };
224 let parts: Vec<String> = path.split('/').map(percent_decode).collect();
225 let parts: Vec<&str> = parts.iter().map(String::as_str).collect();
226 match parts.as_slice() {
227 ["github.com", owner, name] => Some(Repo::GitHub {
228 owner: owner.to_string(),
229 name: name.to_string(),
230 }),
231 ["dev.azure.com", org, project, "_git", name]
232 | ["ssh.dev.azure.com", org, project, name] => Some(azure(org, project, name)),
233 [host, "DefaultCollection", project, "_git", name] | [host, project, "_git", name] => Some(
234 azure(host.strip_suffix(".visualstudio.com")?, project, name),
235 ),
236 _ => None,
237 }
238}
239
240fn azure(org: &str, project: &str, name: &str) -> Repo {
241 Repo::AzureDevOps {
242 org: org.to_string(),
243 project: project.to_string(),
244 name: name.to_string(),
245 }
246}
247
248fn percent_decode(segment: &str) -> String {
249 let bytes = segment.as_bytes();
250 let mut out = Vec::with_capacity(bytes.len());
251 let mut i = 0;
252 while i < bytes.len() {
253 let hex = |b: u8| (b as char).to_digit(16);
254 if bytes[i] == b'%'
255 && let (Some(hi), Some(lo)) = (
256 bytes.get(i + 1).copied().and_then(hex),
257 bytes.get(i + 2).copied().and_then(hex),
258 )
259 {
260 out.push((hi * 16 + lo) as u8);
261 i += 3;
262 } else {
263 out.push(bytes[i]);
264 i += 1;
265 }
266 }
267 String::from_utf8_lossy(&out).into_owned()
268}
269
270pub fn resolve(
275 dir: &Path,
276 target: &str,
277 lookup: &dyn PrLookup,
278) -> Result<(RangeSpec, PullRequest)> {
279 let remotes = remotes(dir)?;
280 let pr = match target.parse::<u64>() {
281 Ok(number) => {
282 let (_, repo) = remotes
283 .iter()
284 .find(|(name, _)| name == "origin")
285 .or_else(|| remotes.first())
286 .context("no GitHub or Azure DevOps remote found; pass the full PR URL")?;
287 PullRequest {
288 repo: repo.clone(),
289 number,
290 }
291 }
292 Err(_) => parse_pr_url(target).with_context(|| {
293 format!("`{target}` is not a PR number or a GitHub / Azure DevOps PR URL")
294 })?,
295 };
296 let remote = remotes
297 .iter()
298 .filter(|(_, repo)| repo.same_as(&pr.repo))
299 .min_by_key(|(name, _)| name != "origin")
300 .map(|(name, _)| name.clone())
301 .with_context(|| {
302 format!(
303 "this PR is in {}, but no remote of this repository points there",
304 pr.repo.display()
305 )
306 })?;
307
308 let local = |side: &str| format!("refs/declutter/pr/{}/{side}", pr.number);
309 let spec = RangeSpec {
310 old: local("base"),
311 new: Side::Rev(local("head")),
312 merge_base: true,
313 };
314
315 let merge = local("merge");
319 let merge_ref = format!("+refs/pull/{}/merge:{merge}", pr.number);
320 let fetched = git(dir, &["fetch", "--quiet", "--no-tags", &remote, &merge_ref]).is_ok()
321 && git(dir, &["update-ref", &local("base"), &format!("{merge}^1")]).is_ok()
322 && git(dir, &["update-ref", &local("head"), &format!("{merge}^2")]).is_ok();
323 if fetched {
324 return Ok((spec, pr));
325 }
326
327 let refs = lookup.refs(&pr)?;
329 git(
330 dir,
331 &[
332 "fetch",
333 "--quiet",
334 "--no-tags",
335 &remote,
336 &format!("+{}:{}", refs.head, local("head")),
337 &format!("+{}:{}", refs.base, local("base")),
338 ],
339 )
340 .with_context(|| format!("fetching PR {} from `{remote}`", pr.number))?;
341 Ok((spec, pr))
342}
343
344pub fn pr_number(target: &str) -> Option<u64> {
346 target
347 .parse()
348 .ok()
349 .or_else(|| parse_pr_url(target).map(|pr| pr.number))
350}
351
352fn remotes(dir: &Path) -> Result<Vec<(String, Repo)>> {
355 let output = git(dir, &["config", "--get-regexp", r"^remote\..*\.url$"]).unwrap_or_default();
356 Ok(String::from_utf8_lossy(&output)
357 .lines()
358 .filter_map(|line| {
359 let (key, url) = line.split_once(' ')?;
360 let name = key.strip_prefix("remote.")?.strip_suffix(".url")?;
361 Some((name.to_string(), parse_remote_url(url)?))
362 })
363 .collect())
364}