aboutsummaryrefslogtreecommitdiff
path: root/main_test.go
diff options
context:
space:
mode:
Diffstat (limited to 'main_test.go')
-rw-r--r--main_test.go246
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())
+ }
+ }
+}