diff options
Diffstat (limited to 'hotline')
| -rw-r--r-- | hotline/file_store.go | 127 | ||||
| -rw-r--r-- | hotline/file_transfer.go | 38 | ||||
| -rw-r--r-- | hotline/file_wrapper.go | 8 | ||||
| -rw-r--r-- | hotline/files.go | 22 | ||||
| -rw-r--r-- | hotline/files_test.go | 8 | ||||
| -rw-r--r-- | hotline/server.go | 7 |
6 files changed, 139 insertions, 71 deletions
diff --git a/hotline/file_store.go b/hotline/file_store.go index 8c18256..8cf1403 100644 --- a/hotline/file_store.go +++ b/hotline/file_store.go @@ -1,98 +1,162 @@ package hotline import ( + "io" "io/fs" "os" + "path/filepath" "time" "github.com/stretchr/testify/mock" ) +// FileStore is the storage backend for the file library (the FileRoot that clients browse, +// upload, and download). It is deliberately expressed in terms of io/fs interface types rather +// than *os.File so that non-filesystem backends (e.g. an object store such as S3/R2) can +// implement it. OSFileStore is the default, filesystem-backed implementation. +// +// Symlink/ReadLink exist only to support Hotline aliases and have no analog on object stores; +// such backends may return errors.ErrUnsupported, and callers degrade gracefully. type FileStore interface { - Create(name string) (*os.File, error) - Mkdir(name string, perm os.FileMode) error - Open(name string) (*os.File, error) - OpenFile(name string, flag int, perm fs.FileMode) (*os.File, error) + // Reads + Open(name string) (io.ReadCloser, error) + Stat(name string) (fs.FileInfo, error) + ReadFile(name string) ([]byte, error) + ReadDir(name string) ([]fs.DirEntry, error) + ReadLink(name string) (string, error) + Walk(root string, fn filepath.WalkFunc) error + + // Writes + Create(name string) (io.WriteCloser, error) + OpenFile(name string, flag int, perm fs.FileMode) (io.WriteCloser, error) + WriteFile(name string, data []byte, perm fs.FileMode) error + Mkdir(name string, perm fs.FileMode) error + + // Mutations + Rename(oldpath string, newpath string) error Remove(name string) error RemoveAll(path string) error - Rename(oldpath string, newpath string) error - Stat(name string) (fs.FileInfo, error) Symlink(oldname, newname string) error - WriteFile(name string, data []byte, perm fs.FileMode) error - ReadFile(name string) ([]byte, error) } +// OSFileStore is a FileStore backed by the local filesystem via the os and filepath packages. type OSFileStore struct{} -func (fs *OSFileStore) Mkdir(name string, perm os.FileMode) error { +var _ FileStore = (*OSFileStore)(nil) + +func (*OSFileStore) Mkdir(name string, perm fs.FileMode) error { return os.Mkdir(name, perm) } -func (fs *OSFileStore) Stat(name string) (os.FileInfo, error) { +func (*OSFileStore) Stat(name string) (fs.FileInfo, error) { return os.Stat(name) } -func (fs *OSFileStore) Open(name string) (*os.File, error) { - return os.Open(name) +func (*OSFileStore) Open(name string) (io.ReadCloser, error) { + f, err := os.Open(name) + if err != nil { + return nil, err + } + return f, nil +} + +func (*OSFileStore) ReadDir(name string) ([]fs.DirEntry, error) { + return os.ReadDir(name) +} + +func (*OSFileStore) ReadLink(name string) (string, error) { + return os.Readlink(name) +} + +func (*OSFileStore) Walk(root string, fn filepath.WalkFunc) error { + return filepath.Walk(root, fn) } -func (fs *OSFileStore) Symlink(oldname, newname string) error { +func (*OSFileStore) Symlink(oldname, newname string) error { return os.Symlink(oldname, newname) } -func (fs *OSFileStore) RemoveAll(name string) error { +func (*OSFileStore) RemoveAll(name string) error { return os.RemoveAll(name) } -func (fs *OSFileStore) Remove(name string) error { +func (*OSFileStore) Remove(name string) error { return os.Remove(name) } -func (fs *OSFileStore) Create(name string) (*os.File, error) { - return os.Create(name) +func (*OSFileStore) Create(name string) (io.WriteCloser, error) { + f, err := os.Create(name) + if err != nil { + return nil, err + } + return f, nil } -func (fs *OSFileStore) WriteFile(name string, data []byte, perm fs.FileMode) error { +func (*OSFileStore) WriteFile(name string, data []byte, perm fs.FileMode) error { return os.WriteFile(name, data, perm) } -func (fs *OSFileStore) Rename(oldpath string, newpath string) error { +func (*OSFileStore) Rename(oldpath string, newpath string) error { return os.Rename(oldpath, newpath) } -func (fs *OSFileStore) ReadFile(name string) ([]byte, error) { +func (*OSFileStore) ReadFile(name string) ([]byte, error) { return os.ReadFile(name) } -func (fs *OSFileStore) OpenFile(name string, flag int, perm fs.FileMode) (*os.File, error) { - return os.OpenFile(name, flag, perm) +func (*OSFileStore) OpenFile(name string, flag int, perm fs.FileMode) (io.WriteCloser, error) { + f, err := os.OpenFile(name, flag, perm) + if err != nil { + return nil, err + } + return f, nil } type MockFileStore struct { mock.Mock } -func (mfs *MockFileStore) Mkdir(name string, perm os.FileMode) error { +var _ FileStore = (*MockFileStore)(nil) + +func (mfs *MockFileStore) Mkdir(name string, perm fs.FileMode) error { args := mfs.Called(name, perm) return args.Error(0) } -func (mfs *MockFileStore) Stat(name string) (os.FileInfo, error) { +func (mfs *MockFileStore) Stat(name string) (fs.FileInfo, error) { args := mfs.Called(name) if args.Get(0) == nil { return nil, args.Error(1) } - return args.Get(0).(os.FileInfo), args.Error(1) + return args.Get(0).(fs.FileInfo), args.Error(1) } -func (mfs *MockFileStore) Open(name string) (*os.File, error) { +func (mfs *MockFileStore) Open(name string) (io.ReadCloser, error) { args := mfs.Called(name) - return args.Get(0).(*os.File), args.Error(1) + f, _ := args.Get(0).(io.ReadCloser) + return f, args.Error(1) +} + +func (mfs *MockFileStore) ReadDir(name string) ([]fs.DirEntry, error) { + args := mfs.Called(name) + entries, _ := args.Get(0).([]fs.DirEntry) + return entries, args.Error(1) +} + +func (mfs *MockFileStore) ReadLink(name string) (string, error) { + args := mfs.Called(name) + return args.String(0), args.Error(1) +} + +func (mfs *MockFileStore) Walk(root string, fn filepath.WalkFunc) error { + args := mfs.Called(root, fn) + return args.Error(0) } -func (mfs *MockFileStore) OpenFile(name string, flag int, perm fs.FileMode) (*os.File, error) { +func (mfs *MockFileStore) OpenFile(name string, flag int, perm fs.FileMode) (io.WriteCloser, error) { args := mfs.Called(name, flag, perm) - return args.Get(0).(*os.File), args.Error(1) + f, _ := args.Get(0).(io.WriteCloser) + return f, args.Error(1) } func (mfs *MockFileStore) Symlink(oldname, newname string) error { @@ -110,9 +174,10 @@ func (mfs *MockFileStore) Remove(name string) error { return args.Error(0) } -func (mfs *MockFileStore) Create(name string) (*os.File, error) { +func (mfs *MockFileStore) Create(name string) (io.WriteCloser, error) { args := mfs.Called(name) - return args.Get(0).(*os.File), args.Error(1) + f, _ := args.Get(0).(io.WriteCloser) + return f, args.Error(1) } func (mfs *MockFileStore) WriteFile(name string, data []byte, perm fs.FileMode) error { diff --git a/hotline/file_transfer.go b/hotline/file_transfer.go index 8670439..1526e28 100644 --- a/hotline/file_transfer.go +++ b/hotline/file_transfer.go @@ -13,7 +13,6 @@ import ( "math" "os" "path" - "path/filepath" "slices" "strings" "sync" @@ -289,21 +288,18 @@ func DownloadHandler(w io.Writer, fullPath string, fileTransfer *FileTransfer, f } } - rFile, _ := hlFile.rsrcForkFile() - //if err != nil { - // // return fmt.Errorf("open resource fork file: %v", err) - //} - - _, _ = io.Copy(w, io.TeeReader(rFile, fileTransfer.bytesSentCounter)) - //if err != nil { - // // return fmt.Errorf("send resource fork data: %v", err) - //} + // The resource fork may legitimately not exist. rsrcForkFile returns a nil reader in that + // case; guard against it so backends that return an untyped-nil reader (rather than an + // os.File whose Read tolerates a nil receiver) don't panic. + if rFile, _ := hlFile.rsrcForkFile(); rFile != nil { + _, _ = io.Copy(w, io.TeeReader(rFile, fileTransfer.bytesSentCounter)) + } return nil } func UploadHandler(rwc io.ReadWriter, fullPath string, fileTransfer *FileTransfer, fileStore FileStore, rLogger *slog.Logger, preserveForks bool) error { - var file *os.File + var file io.WriteCloser // A file upload has two possible cases: // 1) Upload a new file @@ -312,7 +308,7 @@ func UploadHandler(rwc io.ReadWriter, fullPath string, fileTransfer *FileTransfe // Check for existing file. If found, do not proceed. This is an invalid scenario, as the file upload transaction // handler should have returned an error to the client indicating there was an existing file present. - _, err := os.Stat(fullPath) + _, err := fileStore.Stat(fullPath) if err == nil { return fmt.Errorf("existing file found: %s", fullPath) } @@ -322,7 +318,7 @@ func UploadHandler(rwc io.ReadWriter, fullPath string, fileTransfer *FileTransfe } // If not found, open or create a new .incomplete file - file, err = os.OpenFile(fullPath+IncompleteFileSuffix, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0644) + file, err = fileStore.OpenFile(fullPath+IncompleteFileSuffix, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0644) if err != nil { return fmt.Errorf("open temp file for uploade: %w", err) } @@ -406,7 +402,7 @@ func DownloadFolderHandler(rwc io.ReadWriter, fullPath string, fileTransfer *Fil } i := 0 - err := filepath.Walk(fullPath+"/", func(path string, info os.FileInfo, err error) error { + err := fileStore.Walk(fullPath+"/", func(path string, info os.FileInfo, err error) error { //s.Stats.DownloadCounter += 1 i += 1 @@ -567,8 +563,8 @@ func UploadFolderHandler(rwc io.ReadWriter, fullPath string, fileTransfer *FileT itemPath := path.Join(fullPath, fu.FormattedPath()) if fu.IsFolder == [2]byte{0, 1} { - if _, err := os.Stat(itemPath); os.IsNotExist(err) { - if err := os.Mkdir(itemPath, 0777); err != nil { + if _, err := fileStore.Stat(itemPath); os.IsNotExist(err) { + if err := fileStore.Mkdir(itemPath, 0777); err != nil { return err } } @@ -581,7 +577,7 @@ func UploadFolderHandler(rwc io.ReadWriter, fullPath string, fileTransfer *FileT nextAction := DlFldrActionSendFile // Check if we have the full file already. If so, send dlFldrAction_NextFile to client to skip. - _, err := os.Stat(itemPath) + _, err := fileStore.Stat(itemPath) if err != nil && !errors.Is(err, fs.ErrNotExist) { return err } @@ -590,7 +586,7 @@ func UploadFolderHandler(rwc io.ReadWriter, fullPath string, fileTransfer *FileT } // Check if we have a partial file already. If so, send dlFldrAction_ResumeFile to client to resume upload. - incompleteFile, err := os.Stat(itemPath + IncompleteFileSuffix) + incompleteFile, err := fileStore.Stat(itemPath + IncompleteFileSuffix) if err != nil && !errors.Is(err, fs.ErrNotExist) { return err } @@ -609,7 +605,7 @@ func UploadFolderHandler(rwc io.ReadWriter, fullPath string, fileTransfer *FileT offset := make([]byte, 4) binary.BigEndian.PutUint32(offset, uint32(incompleteFile.Size())) - file, err := os.OpenFile(itemPath+IncompleteFileSuffix, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644) + file, err := fileStore.OpenFile(itemPath+IncompleteFileSuffix, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644) if err != nil { return err } @@ -633,7 +629,7 @@ func UploadFolderHandler(rwc io.ReadWriter, fullPath string, fileTransfer *FileT rLogger.Error("Error receiving file", "err", err) } - err = os.Rename(itemPath+IncompleteFileSuffix, itemPath) + err = fileStore.Rename(itemPath+IncompleteFileSuffix, itemPath) if err != nil { return err } @@ -680,7 +676,7 @@ func UploadFolderHandler(rwc io.ReadWriter, fullPath string, fileTransfer *FileT } // Rename the temporary upload file to the final file name. - if err := os.Rename(filePath+".incomplete", filePath); err != nil { + if err := fileStore.Rename(filePath+".incomplete", filePath); err != nil { return err } } diff --git a/hotline/file_wrapper.go b/hotline/file_wrapper.go index de3cbe4..4aeaf25 100644 --- a/hotline/file_wrapper.go +++ b/hotline/file_wrapper.go @@ -103,7 +103,7 @@ func (f *File) infoForkName() string { } func (f *File) rsrcForkWriter() (io.WriteCloser, error) { - file, err := os.OpenFile(f.rsrcPath, os.O_CREATE|os.O_WRONLY, 0644) + file, err := f.fs.OpenFile(f.rsrcPath, os.O_CREATE|os.O_WRONLY, 0644) if err != nil { return nil, err } @@ -112,7 +112,7 @@ func (f *File) rsrcForkWriter() (io.WriteCloser, error) { } func (f *File) InfoForkWriter() (io.WriteCloser, error) { - file, err := os.OpenFile(f.infoPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0644) + file, err := f.fs.OpenFile(f.infoPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0644) if err != nil { return nil, err } @@ -121,7 +121,7 @@ func (f *File) InfoForkWriter() (io.WriteCloser, error) { } func (f *File) incFileWriter() (io.WriteCloser, error) { - file, err := os.OpenFile(f.incompletePath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0644) + file, err := f.fs.OpenFile(f.incompletePath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0644) if err != nil { return nil, err } @@ -137,7 +137,7 @@ func (f *File) dataForkReader() (io.Reader, error) { return f.fs.Open(f.incompletePath) } -func (f *File) rsrcForkFile() (*os.File, error) { +func (f *File) rsrcForkFile() (io.ReadCloser, error) { return f.fs.Open(f.rsrcPath) } diff --git a/hotline/files.go b/hotline/files.go index a5b3441..4f9a031 100644 --- a/hotline/files.go +++ b/hotline/files.go @@ -37,8 +37,8 @@ func fileTypeFromInfo(info fs.FileInfo) (ft fileType, err error) { const maxFileSize = 4294967296 -func GetFileNameList(path string, ignoreList []string, encoder *encoding.Encoder, logger *slog.Logger) (fields []Field, err error) { - files, err := os.ReadDir(path) +func GetFileNameList(fileStore FileStore, path string, ignoreList []string, encoder *encoding.Encoder, logger *slog.Logger) (fields []Field, err error) { + files, err := fileStore.ReadDir(path) if err != nil { return fields, fmt.Errorf("error reading path: %s: %w", path, err) } @@ -59,12 +59,12 @@ func GetFileNameList(path string, ignoreList []string, encoder *encoding.Encoder // Check if path is a symlink. If so, follow it. if fileInfo.Mode()&os.ModeSymlink != 0 { - resolvedPath, err := os.Readlink(filepath.Join(path, file.Name())) + resolvedPath, err := fileStore.ReadLink(filepath.Join(path, file.Name())) if err != nil { return fields, fmt.Errorf("error following symlink: %s: %w", resolvedPath, err) } - rFile, err := os.Stat(resolvedPath) + rFile, err := fileStore.Stat(resolvedPath) if errors.Is(err, os.ErrNotExist) { continue } @@ -73,7 +73,7 @@ func GetFileNameList(path string, ignoreList []string, encoder *encoding.Encoder } if rFile.IsDir() { - dir, err := os.ReadDir(filepath.Join(path, file.Name())) + dir, err := fileStore.ReadDir(filepath.Join(path, file.Name())) if err != nil { return fields, err } @@ -94,7 +94,7 @@ func GetFileNameList(path string, ignoreList []string, encoder *encoding.Encoder copy(fnwi.Creator[:], FileTypeFromFilename(rFile.Name()).CreatorCode) } } else if file.IsDir() { - dir, err := os.ReadDir(filepath.Join(path, file.Name())) + dir, err := fileStore.ReadDir(filepath.Join(path, file.Name())) if err != nil { return fields, fmt.Errorf("readDir: %w", err) } @@ -115,7 +115,7 @@ func GetFileNameList(path string, ignoreList []string, encoder *encoding.Encoder continue } - hlFile, err := NewFile(&OSFileStore{}, path+"/"+file.Name(), 0) + hlFile, err := NewFile(fileStore, path+"/"+file.Name(), 0) if err != nil { return nil, fmt.Errorf("NewFile: %w", err) } @@ -148,9 +148,9 @@ func GetFileNameList(path string, ignoreList []string, encoder *encoding.Encoder return fields, nil } -func CalcTotalSize(filePath string) ([]byte, error) { +func CalcTotalSize(fileStore FileStore, filePath string) ([]byte, error) { var totalSize uint32 - err := filepath.Walk(filePath, func(path string, info os.FileInfo, err error) error { + err := fileStore.Walk(filePath, func(path string, info os.FileInfo, err error) error { if err != nil { return err } @@ -174,11 +174,11 @@ func CalcTotalSize(filePath string) ([]byte, error) { } // CalcItemCount recurses through a file path and counts the number of non-hidden files. -func CalcItemCount(filePath string) ([]byte, error) { +func CalcItemCount(fileStore FileStore, filePath string) ([]byte, error) { var itemCount uint16 // Walk the directory and count items - err := filepath.Walk(filePath, func(path string, info os.FileInfo, err error) error { + err := fileStore.Walk(filePath, func(path string, info os.FileInfo, err error) error { if err != nil { return err } diff --git a/hotline/files_test.go b/hotline/files_test.go index 3d9846f..44902d3 100644 --- a/hotline/files_test.go +++ b/hotline/files_test.go @@ -76,7 +76,7 @@ func TestCalcTotalSize(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - got, err := CalcTotalSize(tt.args.filePath) + got, err := CalcTotalSize(&OSFileStore{}, tt.args.filePath) if (err != nil) != tt.wantErr { t.Errorf("CalcTotalSize() error = %v, wantErr %v", err, tt.wantErr) return @@ -160,7 +160,7 @@ func TestCalcItemCount(t *testing.T) { } // Calculate item count - result, err := CalcItemCount(tempDir) + result, err := CalcItemCount(&OSFileStore{}, tempDir) if err != nil { t.Fatalf("CalcItemCount returned an error: %v", err) } @@ -201,7 +201,7 @@ func TestGetFileNameList_Encoding(t *testing.T) { } t.Run("macintosh encoder converts UTF-8 to Mac Roman", func(t *testing.T) { - fields, err := GetFileNameList(tempDir, nil, charmap.Macintosh.NewEncoder(), slog.Default()) + fields, err := GetFileNameList(&OSFileStore{}, tempDir, nil, charmap.Macintosh.NewEncoder(), slog.Default()) assert.NoError(t, err) assert.Len(t, fields, 1) @@ -211,7 +211,7 @@ func TestGetFileNameList_Encoding(t *testing.T) { }) t.Run("nop encoder passes UTF-8 through unchanged", func(t *testing.T) { - fields, err := GetFileNameList(tempDir, nil, encoding.Nop.NewEncoder(), slog.Default()) + fields, err := GetFileNameList(&OSFileStore{}, tempDir, nil, encoding.Nop.NewEncoder(), slog.Default()) assert.NoError(t, err) assert.Len(t, fields, 1) diff --git a/hotline/server.go b/hotline/server.go index 020f878..14cee04 100644 --- a/hotline/server.go +++ b/hotline/server.go @@ -141,6 +141,13 @@ func WithPresenceTracker(p PresenceTracker) func(s *Server) { } } +// WithFileStore sets the storage backend for the file library. Defaults to OSFileStore. +func WithFileStore(fs FileStore) func(s *Server) { + return func(s *Server) { + s.FS = fs + } +} + // WithTLS optionally enables TLS support on the specified port. func WithTLS(tlsConfig *tls.Config, port int) func(s *Server) { return func(s *Server) { |