Parent directory

mirror_test.go

3377 bytes
  1package jobs
  2
  3import (
  4	"errors"
  5	"os/exec"
  6	"path/filepath"
  7	"slices"
  8	"testing"
  9
 10	"github.com/charmbracelet/soft-serve/git"
 11	"github.com/charmbracelet/soft-serve/pkg/ssrf"
 12)
 13
 14// newRepoWithRemotes creates a bare repo with the given name=url remotes.
 15func newRepoWithRemotes(t *testing.T, remotes map[string]string) *git.Repository {
 16	t.Helper()
 17
 18	path := filepath.Join(t.TempDir(), "repo.git")
 19	if _, err := git.Init(path, true); err != nil {
 20		t.Fatalf("init: %v", err)
 21	}
 22
 23	for name, url := range remotes {
 24		cmd := exec.CommandContext(t.Context(), "git", "remote", "add", name, url)
 25		cmd.Dir = path
 26		if out, err := cmd.CombinedOutput(); err != nil {
 27			t.Fatalf("adding remote %s: %v: %s", name, err, out)
 28		}
 29	}
 30
 31	r, err := git.Open(path)
 32	if err != nil {
 33		t.Fatalf("open: %v", err)
 34	}
 35	return r
 36}
 37
 38func TestValidateMirrorRemotes(t *testing.T) {
 39	tests := []struct {
 40		name    string
 41		remotes map[string]string
 42		wantErr error
 43	}{
 44		{
 45			name:    "public origin",
 46			remotes: map[string]string{"origin": "https://1.1.1.1/x.git"},
 47		},
 48		{
 49			name:    "ssh origin",
 50			remotes: map[string]string{"origin": "ssh://git@10.0.0.1/x.git"},
 51		},
 52		{
 53			name:    "private origin",
 54			remotes: map[string]string{"origin": "http://127.0.0.1:8080/x.git"},
 55			wantErr: ssrf.ErrPrivateIP,
 56		},
 57		{
 58			name:    "metadata origin",
 59			remotes: map[string]string{"origin": "http://169.254.169.254/x.git"},
 60			wantErr: ssrf.ErrPrivateIP,
 61		},
 62		{
 63			// `git remote update` fetches from every remote, not just
 64			// origin. A guard that only reads origin misses this entirely.
 65			name: "private non-origin remote",
 66			remotes: map[string]string{
 67				"origin": "https://1.1.1.1/x.git",
 68				"backup": "http://192.168.1.1/x.git",
 69			},
 70			wantErr: ssrf.ErrPrivateIP,
 71		},
 72		{
 73			// git:// is a raw TCP connect and is just as usable for SSRF.
 74			name:    "private git scheme",
 75			remotes: map[string]string{"origin": "git://10.0.0.1/x.git"},
 76			wantErr: ssrf.ErrPrivateIP,
 77		},
 78		{
 79			// Go and libcurl disagree on non-canonical IPv4 literals.
 80			name:    "octal loopback",
 81			remotes: map[string]string{"origin": "http://0177.0.0.1/x.git"},
 82			wantErr: ssrf.ErrAmbiguousHost,
 83		},
 84		{
 85			// A remote that cannot be parsed must not be fetched from.
 86			// Failing open here was how unvalidated remotes kept firing.
 87			name:    "unparseable remote",
 88			remotes: map[string]string{"origin": "http://[::1/x.git"},
 89			wantErr: ssrf.ErrInvalidURL,
 90		},
 91	}
 92
 93	for _, tt := range tests {
 94		t.Run(tt.name, func(t *testing.T) {
 95			r := newRepoWithRemotes(t, tt.remotes)
 96			env, err := validateMirrorRemotes(r)
 97
 98			if tt.wantErr != nil {
 99				if !errors.Is(err, tt.wantErr) {
100					t.Fatalf("validateMirrorRemotes() error = %v, want %v", err, tt.wantErr)
101				}
102				return
103			}
104
105			if err != nil {
106				t.Fatalf("validateMirrorRemotes() unexpected error: %v", err)
107			}
108			if !slices.Contains(env, "GIT_CONFIG_KEY_0=http.followRedirects") ||
109				!slices.Contains(env, "GIT_CONFIG_VALUE_0=false") {
110				t.Errorf("sync env did not disable redirects: %v", env)
111			}
112		})
113	}
114}
115
116// TestValidateMirrorRemotesNoRemote verifies a mirror with no remote URL is
117// skipped rather than treated as valid.
118func TestValidateMirrorRemotesNoRemote(t *testing.T) {
119	r := newRepoWithRemotes(t, nil)
120	if _, err := validateMirrorRemotes(r); err == nil {
121		t.Error("validateMirrorRemotes() accepted a repo with no remote")
122	}
123}