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}