josie / simplegit

package git

import (
	"bufio"
	"bytes"
	"errors"
	"fmt"
	"io"
	"os/exec"
	"strconv"
	"strings"
)

// TreeEntry is one item in a tree listing at a ref.
type TreeEntry struct {
	Mode string
	Type string // "blob", "tree", or "commit" for a submodule
	OID  string
	Path string // relative to the repo root
}

// LsTree lists one directory level of ref. dir is "" for the root;
// entry paths come back relative to the repo root.
func LsTree(repoPath, ref, dir string) ([]TreeEntry, error) {
	args := []string{"ls-tree", "-z", ref}
	if dir != "" {
		args = append(args, "--", dir+"/")
	}
	cmd := gitCommand(repoPath, args...)
	var stderr bytes.Buffer
	cmd.Stderr = &stderr
	out, err := cmd.Output()
	if err != nil {
		return nil, fmt.Errorf("git ls-tree %v in %s: %w: %s", args, repoPath, err, strings.TrimSpace(stderr.String()))
	}

	var entries []TreeEntry
	for _, line := range bytes.Split(out, []byte{0}) {
		if len(line) == 0 {
			continue
		}
		// "<mode> <type> <oid>\t<path>"
		meta, path, ok := bytes.Cut(line, []byte{'\t'})
		if !ok {
			return nil, fmt.Errorf("git ls-tree: malformed entry %q", line)
		}
		fields := strings.Fields(string(meta))
		if len(fields) != 3 {
			return nil, fmt.Errorf("git ls-tree: malformed meta %q", meta)
		}
		entries = append(entries, TreeEntry{
			Mode: fields[0], Type: fields[1], OID: fields[2], Path: string(path),
		})
	}
	return entries, nil
}

// LsTreePaths lists every file path in ref's tree, recursively, in git's
// tree order. Directories are implied by path segments and not listed.
func LsTreePaths(repoPath, ref string) ([]string, error) {
	cmd := gitCommand(repoPath, "ls-tree", "-r", "-z", "--name-only", ref)
	var stderr bytes.Buffer
	cmd.Stderr = &stderr
	out, err := cmd.Output()
	if err != nil {
		return nil, fmt.Errorf("git ls-tree -r %s in %s: %w: %s", ref, repoPath, err, strings.TrimSpace(stderr.String()))
	}
	var paths []string
	for _, p := range bytes.Split(out, []byte{0}) {
		if len(p) > 0 {
			paths = append(paths, string(p))
		}
	}
	return paths, nil
}

// ShowFile returns the contents of path at ref, reading at most limit+1
// bytes; callers detect oversized content with len(source) > limit.
func ShowFile(repoPath, ref, path string, limit int) ([]byte, error) {
	cmd := gitCommand(repoPath, "show", ref+":"+path)
	var stderr bytes.Buffer
	cmd.Stderr = &stderr
	out, oversized, err := runBounded(cmd, limit)
	if oversized {
		return out, nil
	}
	if err != nil {
		var ee *exec.ExitError
		if errors.As(err, &ee) && ee.ExitCode() == 128 {
			return nil, fmt.Errorf("git show %s:%s: %w: %s", ref, path, ErrNotFound, strings.TrimSpace(stderr.String()))
		}
		return nil, fmt.Errorf("git show %s:%s: %w: %s", ref, path, err, strings.TrimSpace(stderr.String()))
	}
	return out, nil
}

var ErrNotFound = errors.New("not found")

// RefExists reports whether ref resolves in the repository (an empty
// bare repo resolves nothing).
func RefExists(repoPath, ref string) (bool, error) {
	cmd := gitCommand(repoPath, "rev-parse", "--verify", "--quiet", ref)
	if err := cmd.Run(); err != nil {
		var ee *exec.ExitError
		if errors.As(err, &ee) && ee.ExitCode() == 1 {
			return false, nil
		}
		return false, fmt.Errorf("git rev-parse %s in %s: %w", ref, repoPath, err)
	}
	return true, nil
}

// RefNames lists the short names of every branch and tag, so callers can
// test ref membership without one subprocess per candidate.
func RefNames(repoPath string) ([]string, error) {
	return forEachRefNames(repoPath, "refs/heads", "refs/tags")
}

// forEachRefNames lists %(refname:short) for each matching pattern.
func forEachRefNames(repoPath string, patterns ...string) ([]string, error) {
	args := append([]string{"for-each-ref", "--format=%(refname:short)"}, patterns...)
	cmd := gitCommand(repoPath, args...)
	var stderr bytes.Buffer
	cmd.Stderr = &stderr
	out, err := cmd.Output()
	if err != nil {
		return nil, fmt.Errorf("git for-each-ref in %s: %w: %s", repoPath, err, strings.TrimSpace(stderr.String()))
	}
	var names []string
	for _, line := range strings.Split(string(out), "\n") {
		if name := strings.TrimSpace(line); name != "" {
			names = append(names, name)
		}
	}
	return names, nil
}

// CatFileInfo inspects an object (typically "ref:path") in a single
// batched cat-file read: it reports the object type, its size, whether the
// first max bytes contain a NUL (the same binary heuristic as the blob
// view), and whether the object exceeds max. At most max+1 bytes are read
// whatever the object's size; unknown objects return ErrNotFound.
func CatFileInfo(repoPath, object string, max int) (objectType string, size int, isBinary, truncated bool, err error) {
	if object == "" || strings.HasPrefix(object, "-") {
		return "", 0, false, false, fmt.Errorf("git cat-file: invalid object %q", object)
	}
	cmd := gitCommand(repoPath, "cat-file", "--batch", "--buffer")
	stdin, err := cmd.StdinPipe()
	if err != nil {
		return "", 0, false, false, fmt.Errorf("cat-file stdin: %w", err)
	}
	stdout, err := cmd.StdoutPipe()
	if err != nil {
		return "", 0, false, false, fmt.Errorf("cat-file stdout: %w", err)
	}
	if err := cmd.Start(); err != nil {
		return "", 0, false, false, fmt.Errorf("start cat-file: %w", err)
	}
	// done stops the (read-only) cat-file in every path; Wait reaps it.
	done := func() {
		stdout.Close()
		_ = cmd.Process.Kill()
		_ = cmd.Wait()
	}

	if _, err := io.WriteString(stdin, object+"\n"); err != nil {
		done()
		return "", 0, false, false, fmt.Errorf("write cat-file request: %w", err)
	}
	stdin.Close()

	line, err := bufio.NewReader(stdout).ReadString('\n')
	if err != nil {
		done()
		return "", 0, false, false, fmt.Errorf("read cat-file header: %w", err)
	}
	fields := strings.Fields(line)
	if len(fields) < 2 || fields[1] == "missing" {
		done()
		return "", 0, false, false, fmt.Errorf("git cat-file %s: %w", object, ErrNotFound)
	}
	if len(fields) < 3 {
		done()
		return "", 0, false, false, fmt.Errorf("git cat-file %s: malformed header %q", object, line)
	}
	size, err = strconv.Atoi(fields[2])
	if err != nil {
		done()
		return "", 0, false, false, fmt.Errorf("git cat-file %s: bad size %q: %w", object, fields[2], err)
	}
	objectType = fields[1]

	if objectType == "blob" {
		br := bufio.NewReader(stdout)
		buf := make([]byte, max+1)
		n, err := io.ReadFull(br, buf)
		if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, io.ErrUnexpectedEOF) {
			done()
			return "", 0, false, false, fmt.Errorf("read cat-file body: %w", err)
		}
		truncated = n > max
		for _, b := range buf[:min(n, max)] {
			if b == 0 {
				isBinary = true
				break
			}
		}
	}
	done()
	return objectType, size, isBinary, truncated, nil
}

// Commit is the metadata of one revision.
type Commit struct {
	SHA     string
	Author  string
	Email   string
	Date    string // author date, ISO 8601
	Subject string
	Body    string
	Parents []string
}

const (
	logFormat = "%x1e%H%x1f%an%x1f%ae%x1f%aI%x1f%s%x1f%b"
	recordSep = "\x1e"
	fieldSep  = "\x1f"
)

func parseLog(out []byte) []Commit {
	var commits []Commit
	for _, rec := range strings.Split(string(out), recordSep) {
		rec = strings.TrimPrefix(rec, "\n")
		if rec == "" {
			continue
		}
		fields := strings.Split(rec, fieldSep)
		if len(fields) < 6 {
			continue
		}
		c := Commit{
			SHA: fields[0], Author: fields[1], Email: fields[2],
			Date: fields[3], Subject: fields[4], Body: strings.TrimSpace(fields[5]),
		}
		commits = append(commits, c)
	}
	return commits
}

// Log returns up to limit commits reachable from ref, newest first.
func Log(repoPath, ref string, limit int) ([]Commit, error) {
	cmd := gitCommand(repoPath, "log", "-n", strconv.Itoa(limit), "--format="+logFormat, ref, "--")
	// trailing `--` disambiguates the rev from paths (ref is regex-validated upstream).
	var stderr bytes.Buffer
	cmd.Stderr = &stderr
	out, err := cmd.Output()
	if err != nil {
		return nil, fmt.Errorf("git log %s in %s: %w: %s", ref, repoPath, err, strings.TrimSpace(stderr.String()))
	}
	return parseLog(out), nil
}

// CommitAt returns the single commit addressed by rev (full or short sha).
func CommitAt(repoPath, rev string) (Commit, error) {
	commits, err := Log(repoPath, rev, 1)
	if err != nil {
		return Commit{}, err
	}
	if len(commits) == 0 {
		return Commit{}, fmt.Errorf("commit %s: %w", rev, ErrNotFound)
	}
	parents, err := parentsOf(repoPath, commits[0].SHA)
	if err != nil {
		return Commit{}, err
	}
	commits[0].Parents = parents
	return commits[0], nil
}

func parentsOf(repoPath, sha string) ([]string, error) {
	cmd := gitCommand(repoPath, "rev-list", "--parents", "-n", "1", sha)
	var stderr bytes.Buffer
	cmd.Stderr = &stderr
	out, err := cmd.Output()
	if err != nil {
		return nil, fmt.Errorf("git rev-list %s: %w: %s", sha, err, strings.TrimSpace(stderr.String()))
	}
	fields := strings.Fields(string(out))
	if len(fields) < 2 {
		return nil, nil // root commit
	}
	return fields[1:], nil
}

// ShowPatch returns the diff a commit introduces (format suppressed, so
// only the patch body), reading at most limit+1 bytes.
func ShowPatch(repoPath, sha string, limit int) ([]byte, error) {
	cmd := gitCommand(repoPath, "show", "--format=", "--patch", "--find-renames", sha)
	var stderr bytes.Buffer
	cmd.Stderr = &stderr
	out, oversized, err := runBounded(cmd, limit)
	if oversized {
		return out, nil
	}
	if err != nil {
		return nil, fmt.Errorf("git show %s: %w: %s", sha, err, strings.TrimSpace(stderr.String()))
	}
	return out, nil
}

// CommitCount returns the number of commits reachable from ref.
func CommitCount(repoPath, ref string) (int, error) {
	cmd := gitCommand(repoPath, "rev-list", "--count", ref)
	var stderr bytes.Buffer
	cmd.Stderr = &stderr
	out, err := cmd.Output()
	if err != nil {
		return 0, fmt.Errorf("git rev-list --count %s: %w: %s", ref, err, strings.TrimSpace(stderr.String()))
	}
	n, err := strconv.Atoi(strings.TrimSpace(string(out)))
	if err != nil {
		return 0, fmt.Errorf("git rev-list --count %s: parse %q: %w", ref, out, err)
	}
	return n, nil
}