Skip to main content

declutter/
pr.rs

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/// A repository on a hosting service.
10#[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    /// Names the PR as a review, for scoping notes to it.
64    pub fn review_key(&self) -> String {
65        format!("{} PR {}", self.repo.display(), self.number)
66    }
67
68    /// Where its fetched head lives locally.
69    pub fn head_ref(&self) -> String {
70        format!("refs/declutter/pr/{}/head", self.number)
71    }
72}
73
74/// The refs to fetch for a pull request, as full ref names on the remote.
75#[derive(Debug, Clone, PartialEq, Eq)]
76pub struct PrRefs {
77    pub base: String,
78    pub head: String,
79}
80
81/// Asks the hosting service which branches a pull request compares.
82pub trait PrLookup {
83    fn refs(&self, pr: &PullRequest) -> Result<PrRefs>;
84}
85
86/// Looks pull requests up with the `gh` and `az` command-line tools, using whatever
87/// account they are logged in with.
88pub 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                    // GitHub publishes every PR head here, forks included.
112                    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
168/// Parses a pull-request URL from GitHub or Azure DevOps.
169pub 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
205/// Parses a git remote URL that points at GitHub or Azure DevOps.
206pub 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        // Drop credentials or a username in front of the host.
218        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
270/// Turns a PR URL or number into a range: fetches the PR's base and head into
271/// `refs/declutter/pr/<n>/` (so no local branch is touched) and compares the head
272/// against its merge base with the base, as the PR page does. The host's CLI is only
273/// asked when the remote publishes no merge ref for the PR.
274pub 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    // Both hosts publish an open PR merged into its target as refs/pull/<n>/merge: the
316    // first parent is the target, the second the PR's head. Reading the PR from it needs
317    // nothing but git's own access to the remote — no `gh` or `az`.
318    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    // No merge ref (a conflicting or closed PR): ask the host for the branches.
328    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
344/// The PR number a `pr` target names, for labels: from a URL or a bare number.
345pub 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
352/// Remotes whose configured URL is a GitHub or Azure DevOps repository. Reads the raw
353/// configured URL, before any `insteadOf` rewriting.
354fn 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}