diff --git a/context.go b/context.go index 5174033eb3..04d683ea17 100644 --- a/context.go +++ b/context.go @@ -717,6 +717,37 @@ func (c *Context) MultipartForm() (*multipart.Form, error) { // SaveUploadedFile uploads the form file to specific dst. func (c *Context) SaveUploadedFile(file *multipart.FileHeader, dst string, perm ...fs.FileMode) error { + if c != nil && c.engine != nil && c.engine.fileRoot != nil { + if err := validateRootDestination(dst); err != nil { + return err + } + return saveUploadedFileWithRoot(file, c.engine.fileRoot, dst, perm...) + } + return saveUploadedFile(file, dst, perm...) +} + +func validateRootDestination(dst string) error { + if dst == "" { + return &fs.PathError{Op: "save", Path: dst, Err: fs.ErrInvalid} + } + if filepath.IsAbs(dst) || filepath.VolumeName(dst) != "" { + return &fs.PathError{Op: "save", Path: dst, Err: fs.ErrInvalid} + } + clean := filepath.Clean(dst) + if clean == ".." || strings.HasPrefix(clean, ".."+string(filepath.Separator)) { + return &fs.PathError{Op: "save", Path: dst, Err: fs.ErrInvalid} + } + return nil +} + +func sanitizeRootMode(mode os.FileMode) os.FileMode { + if mode&0o777 != mode { + return mode & 0o777 + } + return mode +} + +func saveUploadedFile(file *multipart.FileHeader, dst string, perm ...fs.FileMode) error { src, err := file.Open() if err != nil { return err @@ -745,6 +776,38 @@ func (c *Context) SaveUploadedFile(file *multipart.FileHeader, dst string, perm return err } +func saveUploadedFileWithRoot(file *multipart.FileHeader, root *os.Root, dst string, perm ...fs.FileMode) error { + src, err := file.Open() + if err != nil { + return err + } + defer src.Close() + + var mode os.FileMode = 0o750 + if len(perm) > 0 { + mode = perm[0] + } + mode = sanitizeRootMode(mode) + + if dir := filepath.Dir(dst); dir != "." { + if err = root.MkdirAll(dir, mode); err != nil { + return err + } + if err = root.Chmod(dir, mode); err != nil { + return err + } + } + + out, err := root.OpenFile(dst, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, mode) + if err != nil { + return err + } + defer out.Close() + + _, err = io.Copy(out, src) + return err +} + // Bind checks the Method and Content-Type to select a binding engine automatically, // Depending on the "Content-Type" header different bindings are used, for example: // diff --git a/context_test.go b/context_test.go index ef60379d77..1354929e05 100644 --- a/context_test.go +++ b/context_test.go @@ -20,6 +20,7 @@ import ( "os" "path/filepath" "reflect" + "runtime" "strconv" "strings" "sync" @@ -233,6 +234,10 @@ func TestSaveUploadedCreateFailed(t *testing.T) { } func TestSaveUploadedFileWithPermission(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("chmod semantics differ on windows") + } + buf := new(bytes.Buffer) mw := multipart.NewWriter(buf) w, err := mw.CreateFormFile("file", "permission_test") @@ -274,6 +279,94 @@ func TestSaveUploadedFileWithPermissionFailed(t *testing.T) { require.Error(t, c.SaveUploadedFile(f, "test/permission_test", mode)) } +func TestSaveUploadedFileWithRoot(t *testing.T) { + rootDir := t.TempDir() + root, err := os.OpenRoot(rootDir) + require.NoError(t, err) + t.Cleanup(func() { + assert.NoError(t, root.Close()) + }) + + buf := new(bytes.Buffer) + mw := multipart.NewWriter(buf) + w, err := mw.CreateFormFile("file", "test") + require.NoError(t, err) + _, err = w.Write([]byte("test")) + require.NoError(t, err) + mw.Close() + + c, _ := CreateTestContext(httptest.NewRecorder()) + c.engine.SetFileRoot(root) + c.Request, _ = http.NewRequest(http.MethodPost, "/", buf) + c.Request.Header.Set("Content-Type", mw.FormDataContentType()) + f, err := c.FormFile("file") + require.NoError(t, err) + + dst := filepath.Join("subdir", "test") + require.NoError(t, c.SaveUploadedFile(f, dst)) + _, err = os.Stat(filepath.Join(rootDir, dst)) + require.NoError(t, err) +} + +func TestSaveUploadedFileWithRootRejectsEscape(t *testing.T) { + rootDir := t.TempDir() + root, err := os.OpenRoot(rootDir) + require.NoError(t, err) + t.Cleanup(func() { + assert.NoError(t, root.Close()) + }) + + buf := new(bytes.Buffer) + mw := multipart.NewWriter(buf) + w, err := mw.CreateFormFile("file", "test") + require.NoError(t, err) + _, err = w.Write([]byte("test")) + require.NoError(t, err) + mw.Close() + + c, _ := CreateTestContext(httptest.NewRecorder()) + c.engine.SetFileRoot(root) + c.Request, _ = http.NewRequest(http.MethodPost, "/", buf) + c.Request.Header.Set("Content-Type", mw.FormDataContentType()) + f, err := c.FormFile("file") + require.NoError(t, err) + + dst := filepath.Join("..", "escape") + require.Error(t, c.SaveUploadedFile(f, dst)) +} + +func TestSaveUploadedFileWithRootRejectsAbsPath(t *testing.T) { + rootDir := t.TempDir() + root, err := os.OpenRoot(rootDir) + require.NoError(t, err) + t.Cleanup(func() { + assert.NoError(t, root.Close()) + }) + + buf := new(bytes.Buffer) + mw := multipart.NewWriter(buf) + w, err := mw.CreateFormFile("file", "test") + require.NoError(t, err) + _, err = w.Write([]byte("test")) + require.NoError(t, err) + mw.Close() + + c, _ := CreateTestContext(httptest.NewRecorder()) + c.engine.SetFileRoot(root) + c.Request, _ = http.NewRequest(http.MethodPost, "/", buf) + c.Request.Header.Set("Content-Type", mw.FormDataContentType()) + f, err := c.FormFile("file") + require.NoError(t, err) + + var dst string + if runtime.GOOS == "windows" { + dst = filepath.Join(rootDir, "abs") + } else { + dst = filepath.Join(string(os.PathSeparator), "abs") + } + require.Error(t, c.SaveUploadedFile(f, dst)) +} + func TestContextReset(t *testing.T) { router := New() c := router.allocateContext(0) diff --git a/gin.go b/gin.go index 2e033bf347..d4a102551a 100644 --- a/gin.go +++ b/gin.go @@ -166,6 +166,11 @@ type Engine struct { // method call. MaxMultipartMemory int64 + // fileRoot constrains file operations performed by Context.SaveUploadedFile. + // If nil, SaveUploadedFile uses regular OS operations. + // Caller is responsible for closing the root. + fileRoot *os.Root + // UseH2C enable h2c support. UseH2C bool @@ -322,6 +327,12 @@ func (engine *Engine) SetFuncMap(funcMap template.FuncMap) { engine.FuncMap = funcMap } +// SetFileRoot sets the root directory used by Context.SaveUploadedFile. +// If root is nil, SaveUploadedFile uses regular OS operations. +func (engine *Engine) SetFileRoot(root *os.Root) { + engine.fileRoot = root +} + // NoRoute adds handlers for NoRoute. It returns a 404 code by default. func (engine *Engine) NoRoute(handlers ...HandlerFunc) { engine.noRoute = handlers diff --git a/gin_integration_test.go b/gin_integration_test.go index 720b140fee..ec01c34601 100644 --- a/gin_integration_test.go +++ b/gin_integration_test.go @@ -304,8 +304,10 @@ func TestFileDescriptor(t *testing.T) { require.NoError(t, err) socketFile, err := listener.File() if isWindows() { - // not supported by windows, it is unimplemented now - require.Error(t, err) + // On some Windows/Go versions this may be unsupported; skip if so. + if err != nil { + t.Skip("listener.File not supported on Windows") + } } else { require.NoError(t, err) }