Make the ssh layer injectable and cover the destructive server commands

The remote commands are the only code in mgsh that can destroy data, and
they were the least verifiable: each one is a single string handed to a
login shell, so a missing quote silently changes which paths it touches.
Every one of them now goes through the sshExec variable, and yesno is a
variable too, so a test can record what would have been sent and answer
the confirmations without a terminal.

The tests pin down what the previous commit fixed by reasoning alone:
that a declined or unreachable `init` sends no rm -rf, that a project
named "my 'weird' project" reaches the server fully quoted, and that
archive's cp/tar/rm sequence is named, ordered and quoted correctly.
push is driven end to end against a real local bare repository.

archive also gained the server-side existence check that init, clone and
show already had, so a project that was never pushed reports that instead
of failing inside cp -r.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-26 11:38:22 +02:00
co-authored by Claude Opus 5
parent ae8c7a3ec0
commit 689a3fe594
4 changed files with 369 additions and 6 deletions
+11
View File
@@ -304,6 +304,17 @@ func runCommandDepth(line string, depth int) bool {
if !requireProject() { if !requireProject() {
break break
} }
// like `init`/`clone`/`show`, ask the server before acting on it —
// otherwise a missing repository surfaces as a raw `cp -r` failure
exists, err := serverEntryExists(".", PRJ+".git")
if err != nil {
errorln("could not reach the git server: " + err.Error())
break
}
if !exists {
errorln("no repository " + PRJ + " on the git server — run 'init' first")
break
}
comment := sanitizeComment(strings.Join(fields[1:], " ")) comment := sanitizeComment(strings.Join(fields[1:], " "))
z := archiveStamp() z := archiveStamp()
name := PRJ + "_" + z name := PRJ + "_" + z
+22 -4
View File
@@ -108,9 +108,29 @@ func shq(s string) string {
return "'" + strings.ReplaceAll(s, "'", `'\''`) + "'" return "'" + strings.ReplaceAll(s, "'", `'\''`) + "'"
} }
// sshExec runs one command on the git server. Every remote command mgsh issues
// — including the `rm -rf` in `init` — funnels through here, which is why it is
// a variable: tests replace it with a recorder and assert on exactly what would
// reach the server. With capture set the output is returned, otherwise the
// command inherits mgsh's stdio.
var sshExec = func(remote string, capture bool) ([]byte, error) {
args := sshArgs(cfg.GitUser+"@"+cfg.GitHost, remote)
c := exec.Command("ssh", args...)
if !capture {
c.Stdin = os.Stdin
c.Stdout = os.Stdout
c.Stderr = os.Stderr
return nil, c.Run()
}
c.Stderr = os.Stderr
out, err := c.Output()
return out, err
}
// ssh runs a single remote command over ssh with inherited stdio. // ssh runs a single remote command over ssh with inherited stdio.
func ssh(remote string) error { func ssh(remote string) error {
return runInDir("", "ssh", sshArgs(cfg.GitUser+"@"+cfg.GitHost, remote)...) _, err := sshExec(remote, false)
return err
} }
// sshOK runs a remote command and, on failure, prints a red summary. // sshOK runs a remote command and, on failure, prints a red summary.
@@ -124,9 +144,7 @@ func sshOK(remote string) bool {
// sshOut runs a remote command and returns its stdout split into lines. // sshOut runs a remote command and returns its stdout split into lines.
func sshOut(remote string) ([]string, error) { func sshOut(remote string) ([]string, error) {
c := exec.Command("ssh", sshArgs(cfg.GitUser+"@"+cfg.GitHost, remote)...) out, err := sshExec(remote, true)
c.Stderr = os.Stderr
out, err := c.Output()
lines := strings.Split(string(out), "\n") lines := strings.Split(string(out), "\n")
for len(lines) > 0 && lines[len(lines)-1] == "" { for len(lines) > 0 && lines[len(lines)-1] == "" {
lines = lines[:len(lines)-1] lines = lines[:len(lines)-1]
+5 -2
View File
@@ -9,8 +9,11 @@ import (
"github.com/chzyer/readline" "github.com/chzyer/readline"
) )
// yesno asks a y/n question with a default. Reads a single keypress. // yesno asks a y/n question with a default and reads a single keypress. It is a
func yesno(prompt string, def bool) bool { // variable because it guards the destructive operations: tests replace it to
// drive those paths without a terminal, and to assert that the question was
// asked at all.
var yesno = func(prompt string, def bool) bool {
suffix := " y/N ? " suffix := " y/N ? "
if def { if def {
suffix = " Y/n ? " suffix = " Y/n ? "
+331
View File
@@ -0,0 +1,331 @@
package main
// remotecmd_test.go — what mgsh actually sends to the git server.
//
// These are the only commands that can destroy data (`rm -rf` on the server),
// and they are the hardest to check by reading: the remote command is a single
// string handed to a login shell, so a missing quote silently changes which
// paths it touches. sshExec and yesno are swapped for recorders, and the tests
// assert on the exact command strings that would have been sent.
import (
"errors"
"os"
"os/exec"
"path/filepath"
"strings"
"testing"
)
// fakeServer replaces the ssh transport for the duration of the test. reply
// returns the stdout for a remote command (and optionally an error); a nil
// reply answers every command with empty output and success.
func fakeServer(t *testing.T, reply func(cmd string) (string, error)) *[]string {
t.Helper()
var sent []string
old := sshExec
sshExec = func(remote string, capture bool) ([]byte, error) {
sent = append(sent, remote)
if reply == nil {
return nil, nil
}
out, err := reply(remote)
return []byte(out), err
}
t.Cleanup(func() { sshExec = old })
return &sent
}
// fakeAnswers replaces the y/n prompt, recording the questions asked and
// answering each one with answer.
func fakeAnswers(t *testing.T, answer bool) *[]string {
t.Helper()
var asked []string
old := yesno
yesno = func(prompt string, def bool) bool {
asked = append(asked, prompt)
return answer
}
t.Cleanup(func() { yesno = old })
return &asked
}
// useProject points the globals at a fresh project directory under a temp base.
func useProject(t *testing.T, name string) string {
t.Helper()
base := t.TempDir()
dir := filepath.Join(base, name)
if err := os.MkdirAll(dir, 0755); err != nil {
t.Fatal(err)
}
oldBase, oldPrj, oldDir, oldCfg := BASE, PRJ, DIR, cfg
t.Cleanup(func() { BASE, PRJ, DIR, cfg = oldBase, oldPrj, oldDir, oldCfg })
BASE, PRJ, DIR = base, name, dir
cfg = Config{
Base: base, GitHost: "git.example", GitPort: "22",
GitUser: "git", GitPath: "/home/git",
}
return dir
}
// findCmd returns the first recorded command containing sub, or "".
func findCmd(sent []string, sub string) string {
for _, c := range sent {
if strings.Contains(c, sub) {
return c
}
}
return ""
}
// TestInitKeepsServerRepoWhenDeclined is the regression test for the worst bug
// in this command set: `init` used the *local* remote.origin.url to decide
// whether a server repository existed, so an unlinked project directory made it
// skip the prompt and wipe the remote history.
func TestInitKeepsServerRepoWhenDeclined(t *testing.T) {
useProject(t, "notes") // deliberately no .git, so there is no local origin
sent := fakeServer(t, func(cmd string) (string, error) {
if strings.HasPrefix(cmd, "/bin/ls") {
return "notes.git\nother.git\n", nil // the repo DOES exist there
}
return "", nil
})
asked := fakeAnswers(t, false) // ... and the user declines
runCommand("init")
if len(*asked) == 0 {
t.Fatal("init did not ask before overwriting an existing server repository")
}
if c := findCmd(*sent, "rm -rf"); c != "" {
t.Fatalf("init destroyed the server repository after the user declined: %q", c)
}
if c := findCmd(*sent, "git --bare init"); c != "" {
t.Fatalf("init re-created the repository after the user declined: %q", c)
}
}
// TestInitAbortsWhenServerUnreachable: an unreachable server must not be read
// as "no repository there" — that would turn a network glitch into data loss.
func TestInitAbortsWhenServerUnreachable(t *testing.T) {
useProject(t, "notes")
sent := fakeServer(t, func(cmd string) (string, error) {
return "", errors.New("ssh: connect to host git.example port 22: Network is unreachable")
})
asked := fakeAnswers(t, true) // even a "yes" must not get that far
runCommand("init")
if c := findCmd(*sent, "rm -rf"); c != "" {
t.Fatalf("init sent %q although the server could not be listed", c)
}
if len(*asked) != 0 {
t.Fatalf("init asked %q despite not knowing the server state", (*asked)[0])
}
}
// TestInitOnFreshProjectCreatesRepo: with nothing on the server, `init` goes
// ahead without asking and builds the bare repository at the configured path.
func TestInitOnFreshProjectCreatesRepo(t *testing.T) {
useProject(t, "notes")
sent := fakeServer(t, func(cmd string) (string, error) {
if strings.HasPrefix(cmd, "/bin/ls") {
return "other.git\n", nil // notes.git is not there
}
return "", nil
})
asked := fakeAnswers(t, false)
runCommand("init")
for _, prompt := range *asked {
if strings.Contains(prompt, "overwrite existing repository") {
t.Errorf("init asked about overwriting although nothing was there: %q", prompt)
}
}
if findCmd(*sent, "rm -rf "+shq("/home/git/notes.git")) == "" {
t.Errorf("init did not clear the target path, sent: %q", *sent)
}
if c := findCmd(*sent, "git --bare init"); !strings.Contains(c, shq("/home/git/notes.git")) {
t.Errorf("init did not create the bare repository at the configured path: %q", c)
}
}
// TestRemoteCommandsQuoteProjectNames covers the injection the review found: a
// project name reaches the server inside a single shell string, so a space
// turns one path into two arguments and a backtick runs on the server.
func TestRemoteCommandsQuoteProjectNames(t *testing.T) {
const evil = "my 'weird' project"
useProject(t, evil)
sent := fakeServer(t, func(cmd string) (string, error) {
if strings.HasPrefix(cmd, "/bin/ls") {
return evil + ".git\n", nil
}
return "", nil
})
fakeAnswers(t, true) // confirm the overwrite
runCommand("init")
runCommand("push")
runCommand("archive")
if len(*sent) == 0 {
t.Fatal("no remote commands recorded")
}
for _, c := range *sent {
// the bare name must never appear outside single quotes: every
// occurrence has to be preceded by the quote shq() adds
for _, idx := range indexesOf(c, evil) {
if idx == 0 || c[idx-1] != '\'' {
t.Errorf("unquoted project name in remote command: %q", c)
break
}
}
}
// spot-check the destructive one in full
want := "rm -rf " + shq("/home/git/"+evil+".git")
if findCmd(*sent, want) == "" {
t.Errorf("expected %q among the sent commands, got %q", want, *sent)
}
}
// TestArchiveRequiresServerRepo: `archive` copies the server-side repository,
// so a project that was never pushed must produce a clear message rather than a
// raw `cp -r` failure.
func TestArchiveRequiresServerRepo(t *testing.T) {
useProject(t, "notes")
sent := fakeServer(t, func(cmd string) (string, error) {
return "", nil // empty server
})
runCommand("archive")
if c := findCmd(*sent, "cp -r"); c != "" {
t.Errorf("archive copied although the server has no such repository: %q", c)
}
}
// TestArchiveNamesAndQuotesSnapshot checks the snapshot pipeline end to end:
// the three remote commands, their order, and their quoting.
func TestArchiveNamesAndQuotesSnapshot(t *testing.T) {
useProject(t, "notes")
sent := fakeServer(t, func(cmd string) (string, error) {
if strings.HasPrefix(cmd, "/bin/ls") {
return "notes.git\n", nil
}
return "", nil
})
runCommand("archive before rewrite")
if len(*sent) != 4 { // ls, cp, tar, rm
t.Fatalf("expected 4 remote commands, got %d: %q", len(*sent), *sent)
}
cp, tar, rm := (*sent)[1], (*sent)[2], (*sent)[3]
// name is <project>_<stamp>_<sanitised comment>
stamp := archiveStamp()
name := "notes_" + stamp + "_before_rewrite"
if !strings.HasPrefix(cp, "cp -r "+shq("notes.git")+" ") {
t.Errorf("cp command = %q", cp)
}
if !strings.Contains(cp, shq("archive/"+name+".git")) {
t.Errorf("cp target = %q, want it to contain %q", cp, name)
}
// `cd archive && tar`, not `cd archive;tar`: a failed cd must not let tar
// run in the login directory
if !strings.HasPrefix(tar, "cd archive && tar ") || !strings.Contains(tar, shq(name+".git.tar.gz")) {
t.Errorf("tar command = %q", tar)
}
if rm != "rm -rf "+shq("archive/"+name+".git") {
t.Errorf("cleanup command = %q", rm)
}
}
// TestCloneRefusesUnknownRepository: `clone` must not start a git clone for a
// repository the server does not list.
func TestCloneRefusesUnknownRepository(t *testing.T) {
useProject(t, "notes")
PRJ = ""
sent := fakeServer(t, func(cmd string) (string, error) {
return "other.git\n", nil
})
runCommand("clone notes")
if len(*sent) != 1 || !strings.HasPrefix((*sent)[0], "/bin/ls") {
t.Errorf("clone did more than look: %q", *sent)
}
if PRJ != "" {
t.Errorf("clone selected %q although nothing was cloned", PRJ)
}
}
// TestPushCommitsAndTouchesServerRepo drives the whole `push` sequence against
// a real local bare repository: stage, commit, push, then bump the mtime of the
// server-side repository so `list` can sort by last activity. The project name
// contains a space, so the touch also has to be quoted.
func TestPushCommitsAndTouchesServerRepo(t *testing.T) {
dir := useProject(t, "my project")
bare := filepath.Join(t.TempDir(), "origin.git")
mustGit(t, "", "init", "--bare", "-q", bare)
mustGit(t, dir, "init", "-q")
mustGit(t, dir, "config", "user.name", "t")
mustGit(t, dir, "config", "user.email", "t@e")
if err := os.WriteFile(filepath.Join(dir, "a.txt"), []byte("hi\n"), 0644); err != nil {
t.Fatal(err)
}
mustGit(t, dir, "add", ".")
mustGit(t, dir, "commit", "-q", "-m", "first")
mustGit(t, dir, "remote", "add", "origin", bare)
mustGit(t, dir, "push", "-q", "-u", "origin", "HEAD")
if err := os.WriteFile(filepath.Join(dir, "b.txt"), []byte("more\n"), 0644); err != nil {
t.Fatal(err)
}
sent := fakeServer(t, nil)
runCommand("push a second file")
want := "touch " + shq("/home/git/my project.git")
if findCmd(*sent, want) == "" {
t.Fatalf("expected %q among the sent commands, got %q", want, *sent)
}
// the new file really made it into the bare repository, with the
// "[user@host] comment" message mgsh builds
out, err := gitCapture(bare, "log", "-1", "--format=%s")
if err != nil {
t.Fatal(err)
}
if subj := strings.TrimSpace(out); !strings.HasSuffix(subj, "] a second file") {
t.Errorf("commit subject = %q, want it to end in %q", subj, "] a second file")
}
if files, err := gitCapture(bare, "ls-tree", "--name-only", "HEAD"); err != nil {
t.Fatal(err)
} else if !strings.Contains(files, "b.txt") {
t.Errorf("pushed tree = %q, want it to contain b.txt", files)
}
}
// indexesOf returns every start index of sub in s.
func indexesOf(s, sub string) []int {
var out []int
for i := 0; ; {
j := strings.Index(s[i:], sub)
if j < 0 {
return out
}
out = append(out, i+j)
i += j + 1
}
}
// mustGit runs a git command in dir, failing the test on error.
func mustGit(t *testing.T, dir string, args ...string) {
t.Helper()
out, err := exec.Command("git", append([]string{"-C", dir}, args...)...).CombinedOutput()
if err != nil {
t.Fatalf("git %v: %v\n%s", args, err, out)
}
}