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()) } } }