diff options
Diffstat (limited to 'main_test.go')
| -rw-r--r-- | main_test.go | 246 |
1 files changed, 246 insertions, 0 deletions
diff --git a/main_test.go b/main_test.go new file mode 100644 index 0000000..b767695 --- /dev/null +++ b/main_test.go @@ -0,0 +1,246 @@ +package main + +import ( + "bytes" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "syscall" + "testing" +) + +type testUpload struct { + name string + body string +} + +func TestRootConfinement(t *testing.T) { + rootDir := t.TempDir() + outsideDir := t.TempDir() + if err := os.WriteFile(filepath.Join(outsideDir, "secret"), []byte("secret"), 0o600); err != nil { + t.Fatal(err) + } + if err := os.Symlink(outsideDir, filepath.Join(rootDir, "outside")); err != nil { + t.Fatal(err) + } + + s := newTestServer(t, rootDir, 1<<20) + rec := httptest.NewRecorder() + s.handleGet(rec, httptest.NewRequest(http.MethodGet, "/outside/secret", nil)) + if rec.Code != http.StatusNotFound || strings.Contains(rec.Body.String(), "secret") { + t.Fatalf("external symlink served: status %d, body %q", rec.Code, rec.Body.String()) + } + + req := newUploadRequest(t, "/outside/", []testUpload{{"written", "data"}}) + rec = httptest.NewRecorder() + s.handleUpload(rec, req) + if rec.Code == http.StatusSeeOther { + t.Fatalf("upload through external symlink succeeded") + } + if _, err := os.Stat(filepath.Join(outsideDir, "written")); !os.IsNotExist(err) { + t.Fatalf("upload escaped root: %v", err) + } +} + +func TestInternalSymlink(t *testing.T) { + rootDir := t.TempDir() + if err := os.Mkdir(filepath.Join(rootDir, "data"), 0o700); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(rootDir, "data", "file"), []byte("content"), 0o600); err != nil { + t.Fatal(err) + } + if err := os.Symlink("data", filepath.Join(rootDir, "internal")); err != nil { + t.Fatal(err) + } + + s := newTestServer(t, rootDir, 1<<20) + rec := httptest.NewRecorder() + s.handleGet(rec, httptest.NewRequest(http.MethodGet, "/internal/file", nil)) + if rec.Code != http.StatusOK || rec.Body.String() != "content" { + t.Fatalf("internal symlink failed: status %d, body %q", rec.Code, rec.Body.String()) + } +} + +func TestSpecialFileIsRejected(t *testing.T) { + rootDir := t.TempDir() + if err := syscall.Mkfifo(filepath.Join(rootDir, "fifo"), 0o600); err != nil { + t.Fatal(err) + } + + s := newTestServer(t, rootDir, 1<<20) + rec := httptest.NewRecorder() + s.handleGet(rec, httptest.NewRequest(http.MethodGet, "/fifo", nil)) + if rec.Code != http.StatusForbidden { + t.Fatalf("status = %d, want %d", rec.Code, http.StatusForbidden) + } +} + +func TestUploadLimitLeavesFilesUntouched(t *testing.T) { + rootDir := t.TempDir() + target := filepath.Join(rootDir, "file") + if err := os.WriteFile(target, []byte("old"), 0o640); err != nil { + t.Fatal(err) + } + if err := os.Chmod(target, 0o640); err != nil { + t.Fatal(err) + } + + s := newTestServer(t, rootDir, 128) + req := newUploadRequest(t, "/", []testUpload{{"file", strings.Repeat("x", 256)}}) + rec := httptest.NewRecorder() + s.handleUpload(rec, req) + if rec.Code != http.StatusRequestEntityTooLarge { + t.Fatalf("status = %d, want %d", rec.Code, http.StatusRequestEntityTooLarge) + } + assertFile(t, target, "old", 0o640) + assertNoTemps(t, rootDir) + + req = newUploadRequest(t, "/", []testUpload{{"file", strings.Repeat("x", 256)}}) + req.ContentLength = -1 + rec = httptest.NewRecorder() + s.handleUpload(rec, req) + if rec.Code != http.StatusRequestEntityTooLarge { + t.Fatalf("chunked status = %d, want %d: %s", + rec.Code, http.StatusRequestEntityTooLarge, rec.Body.String()) + } + assertFile(t, target, "old", 0o640) + assertNoTemps(t, rootDir) +} + +func TestMalformedUploadLeavesFilesUntouched(t *testing.T) { + rootDir := t.TempDir() + target := filepath.Join(rootDir, "first") + if err := os.WriteFile(target, []byte("old"), 0o600); err != nil { + t.Fatal(err) + } + + req := newUploadRequest(t, "/", []testUpload{ + {"first", "replacement"}, + {"second", "new"}, + }) + body, err := io.ReadAll(req.Body) + if err != nil { + t.Fatal(err) + } + req.Body = io.NopCloser(bytes.NewReader(body[:len(body)-10])) + req.ContentLength = int64(len(body) - 10) + + s := newTestServer(t, rootDir, 1<<20) + rec := httptest.NewRecorder() + s.handleUpload(rec, req) + if rec.Code != http.StatusBadRequest { + t.Fatalf("status = %d, want %d", rec.Code, http.StatusBadRequest) + } + assertFile(t, target, "old", 0o600) + if _, err := os.Stat(filepath.Join(rootDir, "second")); !os.IsNotExist(err) { + t.Fatalf("second file committed: %v", err) + } + assertNoTemps(t, rootDir) +} + +func TestUploadModesAndTemporaryFiles(t *testing.T) { + rootDir := t.TempDir() + existing := filepath.Join(rootDir, "existing") + if err := os.WriteFile(existing, []byte("old"), 0o640); err != nil { + t.Fatal(err) + } + if err := os.Chmod(existing, 0o640); err != nil { + t.Fatal(err) + } + + s := newTestServer(t, rootDir, 1<<20) + req := newUploadRequest(t, "/", []testUpload{ + {"existing", "replacement"}, + {"new", "content"}, + }) + rec := httptest.NewRecorder() + s.handleUpload(rec, req) + if rec.Code != http.StatusSeeOther { + t.Fatalf("status = %d, want %d: %s", rec.Code, http.StatusSeeOther, rec.Body.String()) + } + assertFile(t, existing, "replacement", 0o640) + assertFile(t, filepath.Join(rootDir, "new"), "content", 0o600) + assertNoTemps(t, rootDir) + + tempName := tempPrefix + strings.Repeat("0", 32) + if err := os.WriteFile(filepath.Join(rootDir, tempName), []byte("partial"), 0o600); err != nil { + t.Fatal(err) + } + rec = httptest.NewRecorder() + s.handleGet(rec, httptest.NewRequest(http.MethodGet, "/"+tempName, nil)) + if rec.Code != http.StatusNotFound { + t.Fatalf("temporary file status = %d, want %d", rec.Code, http.StatusNotFound) + } + rec = httptest.NewRecorder() + s.handleGet(rec, httptest.NewRequest(http.MethodGet, "/", nil)) + if strings.Contains(rec.Body.String(), tempName) || strings.Contains(rec.Body.String(), rootDir) { + t.Fatalf("listing exposed private path or temporary file: %q", rec.Body.String()) + } +} + +func newTestServer(t *testing.T, dir string, maxUpload int64) *server { + t.Helper() + root, err := os.OpenRoot(dir) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { root.Close() }) + return &server{root: root, maxUpload: maxUpload} +} + +func newUploadRequest(t *testing.T, target string, files []testUpload) *http.Request { + t.Helper() + var body bytes.Buffer + w := multipart.NewWriter(&body) + for _, file := range files { + part, err := w.CreateFormFile("files", file.name) + if err != nil { + t.Fatal(err) + } + if _, err := io.WriteString(part, file.body); err != nil { + t.Fatal(err) + } + } + if err := w.Close(); err != nil { + t.Fatal(err) + } + req := httptest.NewRequest(http.MethodPost, target, bytes.NewReader(body.Bytes())) + req.Header.Set("Content-Type", w.FormDataContentType()) + return req +} + +func assertFile(t *testing.T, name, content string, mode os.FileMode) { + t.Helper() + data, err := os.ReadFile(name) + if err != nil { + t.Fatal(err) + } + if string(data) != content { + t.Fatalf("%s contains %q, want %q", name, data, content) + } + fi, err := os.Stat(name) + if err != nil { + t.Fatal(err) + } + if fi.Mode().Perm() != mode { + t.Fatalf("%s mode = %o, want %o", name, fi.Mode().Perm(), mode) + } +} + +func assertNoTemps(t *testing.T, dir string) { + t.Helper() + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatal(err) + } + for _, entry := range entries { + if isTempName(entry.Name()) { + t.Fatalf("temporary file remains: %s", entry.Name()) + } + } +} |