diff options
| author | Lena <lena@omega> | 2026-07-01 00:00:00 +0000 |
|---|---|---|
| committer | Lena <lena@omega> | 2026-07-01 00:00:00 +0000 |
| commit | 6aca57a7cdcc2cd0e431146e61f3321100638baa (patch) | |
| tree | b4aaaa5528adf8127a9a7d805cbe722ebbdb86e9 /main_test.go | |
| download | xf-master.tar.gz | |
Serve a directory tree over HTTP through one embedded page: a listing
with download links and a multi-file upload form.
Filesystem access is confined to the served root with os.Root, and
only regular files are served. Uploads are bounded by size and part
count, and every file is staged in the target directory before any is
renamed into place, so a refused or interrupted request leaves the
directory untouched.
Standard library only, no third-party dependencies.
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()) + } + } +} |