diff --git a/cmd/serve/http/http.go b/cmd/serve/http/http.go index c22b40d33..51d733115 100644 --- a/cmd/serve/http/http.go +++ b/cmd/serve/http/http.go @@ -199,6 +199,7 @@ func newServer(ctx context.Context, f fs.Fs, opt *Options, vfsOpt *vfscommon.Opt router := s.server.Router() router.Use( + s.provider.HoldVFS, middleware.Compress(5), middleware.SetHeader("Accept-Ranges", "bytes"), middleware.SetHeader("Server", "rclone/"+fs.Version), diff --git a/cmd/serve/http/http_test.go b/cmd/serve/http/http_test.go index 4567fa434..8d6f2a50f 100644 --- a/cmd/serve/http/http_test.go +++ b/cmd/serve/http/http_test.go @@ -22,6 +22,7 @@ import ( "github.com/rclone/rclone/fs/filter" "github.com/rclone/rclone/fs/rc" libhttp "github.com/rclone/rclone/lib/http" + "github.com/rclone/rclone/lib/random" "github.com/rclone/rclone/vfs/vfscommon" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -559,3 +560,40 @@ func TestNewServerError(t *testing.T) { require.Error(t, err) assert.Nil(t, s) } + +// TestAuthProxyDownloadOutlivesCache checks a download in progress +// carries on working when the auth proxy drops its VFS from its cache, +// as it does when a download takes longer than the cache expiry time. +func TestAuthProxyDownloadOutlivesCache(t *testing.T) { + root := t.TempDir() + contents := random.String(32 * 1024 * 1024) + require.NoError(t, os.WriteFile(filepath.Join(root, "download.bin"), []byte(contents), 0666)) + + prog, err := filepath.Abs("../servetest/proxy_code.go") + require.NoError(t, err) + // FIXME this is untidy setting a global variable! + proxy.Opt.AuthProxy = "go run " + prog + " " + root + defer func() { + proxy.Opt.AuthProxy = "" + }() + s, testURL := start(context.Background(), t, nil) + defer func() { assert.NoError(t, s.Shutdown()) }() + + req, err := http.NewRequest("GET", testURL+"download.bin", nil) + require.NoError(t, err) + req.SetBasicAuth(testUser, testPass) + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer func() { _ = resp.Body.Close() }() + require.Equal(t, http.StatusOK, resp.StatusCode) + start := make([]byte, 1024) + _, err = io.ReadFull(resp.Body, start) + require.NoError(t, err) + + // Drop everything from the proxy's cache as if it had expired + s.provider.Proxy().Shutdown() + + rest, err := io.ReadAll(resp.Body) + require.NoError(t, err) + assert.True(t, contents == string(start)+string(rest), "download corrupted") +} diff --git a/cmd/serve/proxy/proxy.go b/cmd/serve/proxy/proxy.go index 5718588b3..ed0bb3395 100644 --- a/cmd/serve/proxy/proxy.go +++ b/cmd/serve/proxy/proxy.go @@ -12,6 +12,7 @@ import ( "encoding/json" "errors" "fmt" + "net/http" "net/netip" "os/exec" "strings" @@ -577,6 +578,26 @@ func (p *Provider) Get(ctx context.Context) (*vfs.VFS, error) { return VFS, nil } +// HoldVFS is HTTP middleware which holds the VFS for the request until +// the request has been served. +// +// An auth proxy shuts down a VFS when it expires from its cache, which +// a long upload or download can outlast as only the start of a request +// counts as a use of the VFS. +func (p *Provider) HoldVFS(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Requests which don't need auth, eg CORS preflight, have no VFS + if VFS, err := p.Get(r.Context()); err == nil { + if !VFS.Hold() { + http.Error(w, "VFS has been shut down", http.StatusServiceUnavailable) + return + } + defer VFS.Shutdown() + } + next.ServeHTTP(w, r) + }) +} + // VFS returns the fixed VFS, or nil if using an auth proxy. func (p *Provider) VFS() *vfs.VFS { if p == nil {