Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 10 additions & 7 deletions router/downloader/downloader.go
Original file line number Diff line number Diff line change
Expand Up @@ -252,10 +252,6 @@ func (dl *Download) Execute() error {
return errors.New("downloader: got bad response status from endpoint: " + res.Status)
}

if res.ContentLength < 1 {
return errors.New("downloader: request is missing ContentLength")
}

if dl.req.UseHeader {
if contentDisposition := res.Header.Get("Content-Disposition"); contentDisposition != "" {
_, params, err := mime.ParseMediaType(contentDisposition)
Expand Down Expand Up @@ -285,8 +281,14 @@ func (dl *Download) Execute() error {
dl.server.Log().WithField("path", p).Debug("writing remote file to disk")

// Write the file while tracking the progress, Write will check that the
// size of the file won't exceed the disk limit.
r := io.TeeReader(res.Body, dl.counter(res.ContentLength))
// size of the file won't exceed the disk limit. The length is unknown (-1)
// for chunked responses and for responses that the HTTP client transparently
// decompressed, in which case Write enforces the limit as the data comes in
// and the progress stays at zero since there is nothing to measure it against.
var r io.Reader = res.Body
if res.ContentLength > 0 {
r = io.TeeReader(res.Body, dl.counter(res.ContentLength))
}
if err := dl.server.Filesystem().Write(p, r, res.ContentLength, 0o644); err != nil {
return errors.WrapIf(err, "downloader: failed to write file to server directory")
}
Expand All @@ -308,7 +310,8 @@ func (dl *Download) BelongsTo(s *server.Server) bool {
}

// Progress returns the current progress of the download as a float value between 0 and 1 where
// 1 indicates that the download is completed.
// 1 indicates that the download is completed. If the remote server did not report the size of
// the file, the progress remains at 0 until the download finishes.
func (dl *Download) Progress() float64 {
dl.mu.RLock()
defer dl.mu.RUnlock()
Expand Down
263 changes: 263 additions & 0 deletions router/downloader/downloader_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,263 @@
package downloader

import (
"bytes"
"compress/gzip"
"encoding/json"
"net/http"
"net/http/httptest"
"net/url"
"os"
"path/filepath"
"strconv"
"strings"
"testing"

"github.com/google/uuid"

"github.com/pelican/wings/config"
"github.com/pelican/wings/remote"
"github.com/pelican/wings/server"
"github.com/pelican/wings/server/filesystem"
)

// newTestServer returns a server whose data directory lives in a temporary
// directory and whose disk is limited to diskSpaceMiB.
func newTestServer(t *testing.T, diskSpaceMiB int64) *server.Server {
t.Helper()

cfg := &config.Configuration{AuthenticationToken: "test-token"}
cfg.System.Data = t.TempDir()
cfg.System.User.Uid = os.Getuid()
cfg.System.User.Gid = os.Getgid()
config.Set(cfg)

settings, err := json.Marshal(map[string]any{
"uuid": uuid.NewString(),
"build": map[string]any{"disk_space": diskSpaceMiB},
"allocations": map[string]any{
"default": map[string]any{"ip": "127.0.0.1", "port": 25565},
},
})
if err != nil {
t.Fatal(err)
}

s, err := server.NewEmptyManager(nil).InitServer(remote.ServerConfigurationResponse{Settings: settings})
if err != nil {
t.Fatal(err)
}
t.Cleanup(s.CtxCancel)
return s
}

// allowLoopback swaps out the transport of the package level client so that
// the downloader can reach an httptest server listening on 127.0.0.1.
func allowLoopback(t *testing.T) {
t.Helper()

previous := client.Transport
client.Transport = http.DefaultTransport.(*http.Transport).Clone()
t.Cleanup(func() {
client.Transport = previous
})
}

func download(t *testing.T, s *server.Server, rawURL string) (*Download, error) {
t.Helper()

u, err := url.Parse(rawURL)
if err != nil {
t.Fatal(err)
}
dl := New(s, DownloadRequest{Directory: "/", URL: u})
return dl, dl.Execute()
}

func readServerFile(t *testing.T, s *server.Server, name string) []byte {
t.Helper()

contents, err := os.ReadFile(filepath.Join(s.Filesystem().Path(), name))
if err != nil {
t.Fatal(err)
}
return contents
}

func TestDownloadGzipEncodedResponse(t *testing.T) {
allowLoopback(t)
s := newTestServer(t, 1)

payload := []byte(strings.Repeat("server.properties\n", 512))
var compressed bytes.Buffer
gz := gzip.NewWriter(&compressed)
if _, err := gz.Write(payload); err != nil {
t.Fatal(err)
}
if err := gz.Close(); err != nil {
t.Fatal(err)
}

// Mirror raw.githubusercontent.com: the body is gzip encoded whenever the
// client asks for it, which Go's transport does on its own and then hides
// the Content-Length of the response from us after decompressing it.
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if !strings.Contains(r.Header.Get("Accept-Encoding"), "gzip") {
w.Header().Set("Content-Length", strconv.Itoa(len(payload)))
_, _ = w.Write(payload)
return
}
w.Header().Set("Content-Encoding", "gzip")
w.Header().Set("Content-Length", strconv.Itoa(compressed.Len()))
_, _ = w.Write(compressed.Bytes())
}))
defer srv.Close()

if _, err := download(t, s, srv.URL+"/gzip.txt"); err != nil {
t.Fatalf("expected download to succeed, got %v", err)
}
if got := readServerFile(t, s, "gzip.txt"); !bytes.Equal(got, payload) {
t.Fatalf("expected %d decompressed bytes on disk, got %d", len(payload), len(got))
}
if got := s.Filesystem().CachedUsage(); got != int64(len(payload)) {
t.Fatalf("expected disk usage of %d, got %d", len(payload), got)
}
}

func TestDownloadChunkedResponse(t *testing.T) {
allowLoopback(t)
s := newTestServer(t, 1)

chunk := []byte(strings.Repeat("a", 4096))
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
for i := 0; i < 8; i++ {
_, _ = w.Write(chunk)
w.(http.Flusher).Flush()
}
}))
defer srv.Close()

dl, err := download(t, s, srv.URL+"/chunked.txt")
if err != nil {
t.Fatalf("expected download to succeed, got %v", err)
}
if got := dl.Progress(); got != 0 {
t.Fatalf("expected progress to stay at 0 without a known length, got %v", got)
}
if got := readServerFile(t, s, "chunked.txt"); !bytes.Equal(got, bytes.Repeat(chunk, 8)) {
t.Fatalf("expected %d bytes on disk, got %d", len(chunk)*8, len(got))
}
if got := s.Filesystem().CachedUsage(); got != int64(len(chunk)*8) {
t.Fatalf("expected disk usage of %d, got %d", len(chunk)*8, got)
}
}

func TestDownloadKnownLengthResponse(t *testing.T) {
allowLoopback(t)
s := newTestServer(t, 1)

payload := bytes.Repeat([]byte("a"), 64*1024)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Length", strconv.Itoa(len(payload)))
_, _ = w.Write(payload)
}))
defer srv.Close()

dl, err := download(t, s, srv.URL+"/known.bin")
if err != nil {
t.Fatalf("expected download to succeed, got %v", err)
}
if got := readServerFile(t, s, "known.bin"); !bytes.Equal(got, payload) {
t.Fatalf("expected %d bytes on disk, got %d", len(payload), len(got))
}
if got := dl.Progress(); got != 1 {
t.Fatalf("expected progress of 1, got %v", got)
}
}

func TestDownloadChunkedResponseExceedingDiskLimit(t *testing.T) {
allowLoopback(t)
s := newTestServer(t, 1)

chunk := []byte(strings.Repeat("a", 64*1024))
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
for i := 0; i < 32; i++ {
if _, err := w.Write(chunk); err != nil {
return
}
w.(http.Flusher).Flush()
}
}))
defer srv.Close()

_, err := download(t, s, srv.URL+"/too-big.bin")
if !filesystem.IsErrorCode(err, filesystem.ErrCodeDiskSpace) {
t.Fatalf("expected a disk space error, got %v", err)
}
if _, err := os.Stat(filepath.Join(s.Filesystem().Path(), "too-big.bin")); !os.IsNotExist(err) {
t.Fatalf("expected the partial download to be removed, got %v", err)
}
if got := s.Filesystem().CachedUsage(); got != 0 {
t.Fatalf("expected disk usage to be released, got %d", got)
}
}

func TestDownloadChunkedResponseExceedingDiskLimitKeepsExistingFile(t *testing.T) {
allowLoopback(t)
s := newTestServer(t, 1)

original := []byte("motd=A Minecraft Server\n")
if err := s.Filesystem().Write("server.properties", bytes.NewReader(original), int64(len(original)), 0o644); err != nil {
t.Fatal(err)
}
usage := s.Filesystem().CachedUsage()

chunk := []byte(strings.Repeat("a", 64*1024))
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
for i := 0; i < 32; i++ {
if _, err := w.Write(chunk); err != nil {
return
}
w.(http.Flusher).Flush()
}
}))
defer srv.Close()

_, err := download(t, s, srv.URL+"/server.properties")
if !filesystem.IsErrorCode(err, filesystem.ErrCodeDiskSpace) {
t.Fatalf("expected a disk space error, got %v", err)
}
if got := readServerFile(t, s, "server.properties"); !bytes.Equal(got, original) {
t.Fatalf("expected the existing file to be left untouched, got %d bytes", len(got))
}
if got := s.Filesystem().CachedUsage(); got != usage {
t.Fatalf("expected disk usage to stay at %d, got %d", usage, got)
}
entries, err := os.ReadDir(s.Filesystem().Path())
if err != nil {
t.Fatal(err)
}
if len(entries) != 1 {
t.Fatalf("expected only the existing file to remain, got %d entries", len(entries))
}
}

func TestDownloadKnownLengthExceedingDiskLimit(t *testing.T) {
allowLoopback(t)
s := newTestServer(t, 1)

payload := bytes.Repeat([]byte("a"), 2*1024*1024)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Length", strconv.Itoa(len(payload)))
_, _ = w.Write(payload)
}))
defer srv.Close()

_, err := download(t, s, srv.URL+"/too-big.bin")
if !filesystem.IsErrorCode(err, filesystem.ErrCodeDiskSpace) {
t.Fatalf("expected a disk space error, got %v", err)
}
if _, err := os.Stat(filepath.Join(s.Filesystem().Path(), "too-big.bin")); !os.IsNotExist(err) {
t.Fatalf("expected no file to be written, got %v", err)
}
}
Loading