use std::collections::HashMap;
use std::fs;
use std::ops::RangeInclusive;
use std::path::{Path, PathBuf};
use std::process::Command;
use anyhow::{Context, Result, bail};
use crate::paths::absolute_normalized;
const FALLBACK_BASE_REF: &str = "main";
const NEW_SIDE_PREFIX: &str = "b/";
#[derive(clap::Args, Debug, Clone, Default)]
pub struct ChangeScopeArgs {
#[arg(
long = "changed-only",
help = "Only files changed against --base",
help_heading = "Change scope"
)]
pub changed_only: bool,
#[arg(
long = "changed-lines",
help = "Only report comments on lines changed against --base (implies --changed-only)",
help_heading = "Change scope"
)]
pub changed_lines: bool,
#[arg(
long,
conflicts_with = "base",
help = "Compare staged changes against HEAD instead of --base (for pre-commit hooks)",
help_heading = "Change scope"
)]
pub staged: bool,
#[arg(
long,
value_name = "REF",
help = "Base ref for --changed-only (default: origin's default branch)",
help_heading = "Change scope"
)]
pub base: Option<String>,
}
impl ChangeScopeArgs {
pub fn is_active(&self) -> bool {
self.changed_only || self.changed_lines || self.staged
}
fn flag(&self) -> &'static str {
if self.changed_lines {
"--changed-lines"
} else if self.changed_only || !self.staged {
"--changed-only"
} else {
"--staged"
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum DiffSource {
Range(String),
Staged,
}
impl DiffSource {
fn diff_arg(&self, repo_root: &Path) -> Result<String> {
match self {
Self::Range(base) => {
let merge_base = run_git(
repo_root,
&["merge-base", base, "HEAD"],
&format!("merge-base {base} HEAD"),
)?;
Ok(merge_base.trim().to_string())
}
Self::Staged => Ok("--cached".to_string()),
}
}
fn includes_untracked(&self) -> bool {
matches!(self, Self::Range(_))
}
fn describe(&self) -> String {
match self {
Self::Range(base) => base.clone(),
Self::Staged => "the index".to_string(),
}
}
}
#[derive(Debug, Clone)]
pub struct ChangeScope {
files: HashMap<PathBuf, Option<Vec<RangeInclusive<usize>>>>,
per_line: bool,
description: String,
}
impl ChangeScope {
pub fn resolve(args: &ChangeScopeArgs, repo_root: Option<&Path>) -> Result<Option<Self>> {
if !args.is_active() {
return Ok(None);
}
let flag = args.flag();
let root = repo_root
.with_context(|| format!("{flag} needs a git repository, and none encloses the base directory"))?;
let source = if args.staged {
DiffSource::Staged
} else {
DiffSource::Range(args.base.clone().unwrap_or_else(|| default_base_ref(root)))
};
let files = if args.changed_lines {
changed_lines(root, &source)?
.into_iter()
.map(|(path, lines)| (path, Some(lines)))
.collect()
} else {
changed_files(root, &source)?
.into_iter()
.map(|path| (path, None))
.collect()
};
Ok(Some(Self {
files,
per_line: args.changed_lines,
description: format!("{flag} against {}", source.describe()),
}))
}
pub fn contains_file(&self, file: &Path) -> bool {
self.files.contains_key(file)
}
pub fn touches(&self, file: &Path, first: usize, last: usize) -> bool {
match self.files.get(file) {
None => false,
Some(None) => true,
Some(Some(ranges)) => ranges
.iter()
.any(|range| *range.start() <= last && first <= *range.end()),
}
}
pub fn is_per_line(&self) -> bool {
self.per_line
}
pub fn note(&self, kept: usize, total: usize) -> String {
format!("{}: {kept} of {total} file(s) changed", self.description)
}
}
fn default_base_ref(repo_root: &Path) -> String {
let Some(common) = git_common_dir(repo_root) else {
return FALLBACK_BASE_REF.to_string();
};
fs::read_to_string(common.join("refs/remotes/origin/HEAD"))
.ok()
.and_then(|content| {
let reference = content.trim().strip_prefix("ref:")?.trim();
let branch = reference.strip_prefix("refs/remotes/origin/")?;
(!branch.is_empty()).then(|| branch.to_string())
})
.unwrap_or_else(|| FALLBACK_BASE_REF.to_string())
}
fn git_common_dir(repo_root: &Path) -> Option<PathBuf> {
let dot_git = repo_root.join(".git");
let metadata = fs::metadata(&dot_git).ok()?;
let git_dir = if metadata.is_dir() {
dot_git
} else {
let pointer = fs::read_to_string(&dot_git).ok()?;
let pointer = pointer.trim().strip_prefix("gitdir:")?.trim();
if pointer.is_empty() {
return None;
}
absolute_normalized(repo_root, Path::new(pointer))
};
match fs::read_to_string(git_dir.join("commondir")) {
Ok(common) => Some(absolute_normalized(&git_dir, Path::new(common.trim()))),
Err(_) => Some(git_dir),
}
}
fn run_git(repo_root: &Path, args: &[&str], description: &str) -> Result<String> {
let output = Command::new("git")
.arg("-C")
.arg(repo_root)
.args(args)
.output()
.with_context(|| format!("failed to run `git {description}`"))?;
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr);
bail!("`git {description}` failed: {}", stderr.trim());
}
Ok(String::from_utf8_lossy(&output.stdout).into_owned())
}
fn untracked_files(repo_root: &Path) -> Result<Vec<PathBuf>> {
let output = run_git(
repo_root,
&["ls-files", "--others", "--exclude-standard", "-z"],
"ls-files --others --exclude-standard -z",
)?;
Ok(output
.split('\0')
.filter(|name| !name.is_empty())
.map(|name| absolute_normalized(repo_root, Path::new(name)))
.collect())
}
fn line_count(path: &Path) -> usize {
let Ok(content) = fs::read(path) else {
return 0;
};
if content.is_empty() {
return 0;
}
let newlines = content.iter().filter(|&&byte| byte == b'\n').count();
if content.ends_with(b"\n") {
newlines
} else {
newlines + 1
}
}
fn git_diff(repo_root: &Path, source: &DiffSource, args: &[&str]) -> Result<Vec<u8>> {
let diff_arg = source.diff_arg(repo_root)?;
let mut command = Command::new("git");
command
.arg("-C")
.arg(repo_root)
.args(["-c", "core.quotePath=false", "diff", "--no-color", "--no-ext-diff"])
.args(args)
.arg(&diff_arg);
let rendered = format!("git diff {} {diff_arg}", args.join(" "));
let output = command
.output()
.with_context(|| format!("failed to run `{rendered}`"))?;
if !output.status.success() {
let stderr = String::from_utf8_lossy(&output.stderr);
bail!("`{rendered}` failed: {}", stderr.trim());
}
Ok(output.stdout)
}
fn changed_files(repo_root: &Path, source: &DiffSource) -> Result<Vec<PathBuf>> {
let stdout = git_diff(repo_root, source, &["--name-only", "-z"])?;
let mut files: Vec<PathBuf> = stdout
.split(|&byte| byte == 0)
.filter(|name| !name.is_empty())
.map(|name| absolute_normalized(repo_root, Path::new(&*String::from_utf8_lossy(name))))
.collect();
if source.includes_untracked() {
files.extend(untracked_files(repo_root)?);
}
Ok(files)
}
fn changed_lines(repo_root: &Path, source: &DiffSource) -> Result<HashMap<PathBuf, Vec<RangeInclusive<usize>>>> {
let stdout = git_diff(
repo_root,
source,
&["--unified=0", "--src-prefix=a/", "--dst-prefix=b/"],
)?;
let mut files: HashMap<PathBuf, Vec<RangeInclusive<usize>>> = parse_unified_zero(&String::from_utf8_lossy(&stdout))
.into_iter()
.map(|(path, lines)| (absolute_normalized(repo_root, Path::new(&path)), lines))
.collect();
if source.includes_untracked() {
for path in untracked_files(repo_root)? {
let ranges = match line_count(&path) {
0 => Vec::new(),
count => std::iter::once(1..=count).collect(),
};
files.insert(path, ranges);
}
}
Ok(files)
}
fn parse_unified_zero(diff: &str) -> HashMap<String, Vec<RangeInclusive<usize>>> {
let mut files: HashMap<String, Vec<RangeInclusive<usize>>> = HashMap::new();
let mut current: Option<String> = None;
for line in diff.lines() {
if let Some(target) = line.strip_prefix("+++ ") {
current = unquote(target).strip_prefix(NEW_SIDE_PREFIX).map(str::to_string);
if let Some(path) = ¤t {
files.entry(path.clone()).or_default();
}
} else if line.starts_with("@@")
&& let (Some(path), Some(range)) = (¤t, parse_hunk_new_side(line))
{
files.entry(path.clone()).or_default().push(range);
}
}
files
}
fn parse_hunk_new_side(header: &str) -> Option<RangeInclusive<usize>> {
let new_side = header.split_whitespace().nth(2)?.strip_prefix('+')?;
let (start, count) = match new_side.split_once(',') {
Some((start, count)) => (start.parse::<usize>().ok()?, count.parse::<usize>().ok()?),
None => (new_side.parse::<usize>().ok()?, 1),
};
(count > 0).then(|| start..=start + count - 1)
}
fn unquote(raw: &str) -> String {
let Some(inner) = raw.strip_prefix('"').and_then(|rest| rest.strip_suffix('"')) else {
return raw.to_string();
};
let mut bytes = Vec::with_capacity(inner.len());
let mut chars = inner.chars().peekable();
while let Some(ch) = chars.next() {
if ch != '\\' {
let mut buffer = [0u8; 4];
bytes.extend_from_slice(ch.encode_utf8(&mut buffer).as_bytes());
continue;
}
match chars.next() {
Some('n') => bytes.push(b'\n'),
Some('t') => bytes.push(b'\t'),
Some('r') => bytes.push(b'\r'),
Some(digit @ '0'..='7') => {
let mut value = digit.to_digit(8).unwrap_or_default();
for _ in 0..2 {
if let Some(next) = chars.peek().and_then(|c| c.to_digit(8)) {
value = value * 8 + next;
chars.next();
}
}
bytes.push(u8::try_from(value).unwrap_or(u8::MAX));
}
Some(other) => {
let mut buffer = [0u8; 4];
bytes.extend_from_slice(other.encode_utf8(&mut buffer).as_bytes());
}
None => bytes.push(b'\\'),
}
}
String::from_utf8_lossy(&bytes).into_owned()
}
#[cfg(test)]
mod tests {
use super::*;
fn ranges(pairs: &[(usize, usize)]) -> Vec<RangeInclusive<usize>> {
pairs.iter().map(|&(start, end)| start..=end).collect()
}
#[test]
fn the_default_base_ref_is_read_from_origin_head_and_never_master() {
let temp = tempfile::TempDir::new().expect("temp dir");
let root = temp.path();
let refs = root.join(".git/refs/remotes/origin");
fs::create_dir_all(&refs).expect("create refs");
assert_eq!(default_base_ref(root), "main", "no record of origin's default branch");
fs::write(refs.join("HEAD"), "ref: refs/remotes/origin/trunk\n").expect("write HEAD");
assert_eq!(default_base_ref(root), "trunk");
fs::write(refs.join("HEAD"), "ref: refs/remotes/origin/master\n").expect("write HEAD");
assert_eq!(
default_base_ref(root),
"master",
"a repository that really uses master is still honoured"
);
}
#[test]
fn a_worktree_resolves_refs_through_commondir() {
let temp = tempfile::TempDir::new().expect("temp dir");
let main_git = temp.path().join("main/.git");
fs::create_dir_all(main_git.join("refs/remotes/origin")).expect("create refs");
fs::create_dir_all(main_git.join("worktrees/wt")).expect("create worktree dir");
fs::write(
main_git.join("refs/remotes/origin/HEAD"),
"ref: refs/remotes/origin/main\n",
)
.expect("write HEAD");
fs::write(main_git.join("worktrees/wt/commondir"), "../..\n").expect("write commondir");
let worktree = temp.path().join("wt");
fs::create_dir_all(&worktree).expect("create worktree");
fs::write(
worktree.join(".git"),
format!("gitdir: {}\n", main_git.join("worktrees/wt").display()),
)
.expect("write .git");
assert_eq!(default_base_ref(&worktree), "main");
}
#[test]
fn an_inactive_scope_resolves_to_none_even_outside_a_repository() {
let scope = ChangeScope::resolve(&ChangeScopeArgs::default(), None).expect("resolve");
assert!(scope.is_none(), "no flag set must never touch git");
}
#[test]
fn narrowing_outside_a_repository_is_an_error() {
let args = ChangeScopeArgs {
changed_only: true,
..ChangeScopeArgs::default()
};
let error = ChangeScope::resolve(&args, None).expect_err("no repository to diff");
assert!(error.to_string().contains("needs a git repository"), "{error}");
}
#[test]
fn a_hunk_header_yields_its_new_side_lines() {
assert_eq!(parse_hunk_new_side("@@ -3,2 +5,4 @@ fn main() {"), Some(5..=8));
assert_eq!(
parse_hunk_new_side("@@ -3 +7 @@"),
Some(7..=7),
"an omitted count is one line"
);
assert_eq!(
parse_hunk_new_side("@@ -3,2 +2,0 @@"),
None,
"a pure deletion adds no line"
);
assert_eq!(parse_hunk_new_side("@@ garbage @@"), None);
}
#[test]
fn a_unified_zero_diff_is_keyed_by_new_path() {
let diff = "\
diff --git a/src/edited.rs b/src/edited.rs
index 1111111..2222222 100644
--- a/src/edited.rs
+++ b/src/edited.rs
@@ -2,0 +3,2 @@ fn a() {
+ // one
+ // two
@@ -10 +12 @@ fn b() {
-old
+new
diff --git a/src/shrunk.rs b/src/shrunk.rs
--- a/src/shrunk.rs
+++ b/src/shrunk.rs
@@ -4,2 +3,0 @@
-gone
-gone
diff --git a/src/removed.rs b/src/removed.rs
deleted file mode 100644
--- a/src/removed.rs
+++ /dev/null
@@ -1 +0,0 @@
-bye
diff --git a/src/new file.rs b/src/new file.rs
new file mode 100644
--- /dev/null
+++ \"b/src/new\\tfile.rs\"
@@ -0,0 +1,3 @@
+a
+b
+c
";
let parsed = parse_unified_zero(diff);
assert_eq!(parsed.get("src/edited.rs"), Some(&ranges(&[(3, 4), (12, 12)])));
assert_eq!(
parsed.get("src/shrunk.rs"),
Some(&Vec::new()),
"changed, but nothing new"
);
assert!(!parsed.contains_key("src/removed.rs"), "a deleted file has no new side");
assert_eq!(parsed.get("src/new\tfile.rs"), Some(&ranges(&[(1, 3)])), "{parsed:?}");
}
#[test]
fn a_quoted_path_is_unquoted() {
assert_eq!(unquote("b/plain.rs"), "b/plain.rs");
assert_eq!(unquote(r#""b/say \"hi\".rs""#), "b/say \"hi\".rs");
assert_eq!(unquote(r#""b/caf\303\251.rs""#), "b/café.rs");
assert_eq!(unquote(r#""b/back\\slash.rs""#), "b/back\\slash.rs");
}
#[test]
fn line_scope_intersects_a_comment_span_with_the_changed_lines() {
let file = PathBuf::from("/repo/a.rs");
let scope = ChangeScope {
files: HashMap::from([(file.clone(), Some(ranges(&[(5, 7)])))]),
per_line: true,
description: String::new(),
};
assert!(scope.touches(&file, 7, 7));
assert!(scope.touches(&file, 1, 5), "a block comment ending on a changed line");
assert!(scope.touches(&file, 6, 20));
assert!(!scope.touches(&file, 8, 9));
assert!(!scope.touches(&file, 1, 4));
assert!(!scope.touches(Path::new("/repo/other.rs"), 5, 5));
assert!(scope.is_per_line());
}
#[test]
fn file_scope_touches_every_line_of_a_changed_file() {
let file = PathBuf::from("/repo/a.rs");
let scope = ChangeScope {
files: HashMap::from([(file.clone(), None)]),
per_line: false,
description: String::new(),
};
assert!(scope.touches(&file, 1, 1));
assert!(scope.touches(&file, 1000, 1000));
assert!(!scope.is_per_line());
}
}