package assets import ( "context" "errors" "fmt" "io" "mime" "net/url" "os" "path/filepath" "strings" ) type LocalFS struct{ root, baseURL string } func NewLocalFS(root, baseURL string) (*LocalFS, error) { abs, err := filepath.Abs(root) if err != nil { return nil, err } if err = os.MkdirAll(abs, 0750); err != nil { return nil, err } info, err := os.Lstat(abs) if err != nil { return nil, err } if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() { return nil, fmt.Errorf("%w: root", ErrUnsafeBlobKey) } return &LocalFS{root: abs, baseURL: strings.TrimRight(baseURL, "/")}, nil } func (s *LocalFS) Put(_ context.Context, key string, body io.Reader, _ int64, contentType string) (StoredObject, error) { full, parts, err := s.resolve(key) if err != nil { return StoredObject{}, err } if err = s.ensureParents(parts[:len(parts)-1]); err != nil { return StoredObject{}, err } if info, e := os.Lstat(full); e == nil && info.Mode()&os.ModeSymlink != 0 { return StoredObject{}, ErrUnsafeBlobKey } else if e != nil && !errors.Is(e, os.ErrNotExist) { return StoredObject{}, e } tmp, err := os.CreateTemp(filepath.Dir(full), ".asset-*") if err != nil { return StoredObject{}, err } tmpName := tmp.Name() defer os.Remove(tmpName) if _, err = io.Copy(tmp, body); err == nil { err = tmp.Sync() } if closeErr := tmp.Close(); err == nil { err = closeErr } if err != nil { return StoredObject{}, err } if err = os.Rename(tmpName, full); err != nil { return StoredObject{}, err } _ = contentType return StoredObject{Key: strings.Join(parts, "/"), URL: s.baseURL + "/" + escapeParts(parts)}, nil } func (s *LocalFS) Read(_ context.Context, key string) (Blob, error) { full, parts, err := s.resolve(key) if err != nil { return Blob{}, err } if err = s.rejectSymlinks(parts); err != nil { if errors.Is(err, os.ErrNotExist) { return Blob{}, ErrBlobNotFound } return Blob{}, err } f, err := os.Open(full) if errors.Is(err, os.ErrNotExist) { return Blob{}, ErrBlobNotFound } if err != nil { return Blob{}, err } info, err := f.Stat() if err != nil { f.Close() return Blob{}, err } if !info.Mode().IsRegular() { f.Close() return Blob{}, ErrUnsafeBlobKey } contentType := mime.TypeByExtension(filepath.Ext(full)) if contentType == "" { contentType = "application/octet-stream" } return Blob{Body: f, ContentType: contentType, Size: info.Size()}, nil } func (s *LocalFS) Delete(_ context.Context, key string) error { full, parts, err := s.resolve(key) if err != nil { return err } if err = s.rejectSymlinks(parts); errors.Is(err, os.ErrNotExist) { return nil } else if err != nil { return err } err = os.Remove(full) if errors.Is(err, os.ErrNotExist) { return nil } return err } func (s *LocalFS) resolve(key string) (string, []string, error) { if key == "" || filepath.IsAbs(key) || strings.Contains(key, "\\") { return "", nil, ErrUnsafeBlobKey } clean := filepath.ToSlash(filepath.Clean(key)) if clean == "." || clean == ".." || strings.HasPrefix(clean, "../") { return "", nil, ErrUnsafeBlobKey } parts := strings.Split(clean, "/") for _, p := range parts { if p == "" || p == "." || p == ".." { return "", nil, ErrUnsafeBlobKey } } full := filepath.Join(append([]string{s.root}, parts...)...) rel, err := filepath.Rel(s.root, full) if err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) { return "", nil, ErrUnsafeBlobKey } return full, parts, nil } func (s *LocalFS) ensureParents(parts []string) error { current := s.root for _, part := range parts { current = filepath.Join(current, part) info, err := os.Lstat(current) if errors.Is(err, os.ErrNotExist) { if err = os.Mkdir(current, 0750); err != nil { return err } continue } if err != nil { return err } if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() { return ErrUnsafeBlobKey } } return nil } func (s *LocalFS) rejectSymlinks(parts []string) error { current := s.root for _, part := range parts { current = filepath.Join(current, part) info, err := os.Lstat(current) if err != nil { return err } if info.Mode()&os.ModeSymlink != 0 { return ErrUnsafeBlobKey } } return nil } func escapeParts(parts []string) string { out := make([]string, len(parts)) for i, p := range parts { out[i] = url.PathEscape(p) } return strings.Join(out, "/") }