Files

93 lines
1.6 KiB
Go

package snapshot
import (
"archive/zip"
"fmt"
"io"
"io/fs"
"os"
"path/filepath"
"strings"
)
func ZipDir(src, dest string) error {
f, err := os.Create(dest)
if err != nil {
return err
}
defer f.Close()
w := zip.NewWriter(f)
defer w.Close()
return filepath.WalkDir(src, func(path string, d fs.DirEntry, err error) error {
if err != nil {
return err
}
rel, err := filepath.Rel(src, path)
if err != nil {
return err
}
rel = filepath.ToSlash(rel)
if d.IsDir() {
if rel == "." {
return nil
}
_, err = w.Create(rel + "/")
return err
}
fw, err := w.Create(rel)
if err != nil {
return err
}
file, err := os.Open(path)
if err != nil {
return err
}
defer file.Close()
_, err = io.Copy(fw, file)
return err
})
}
func UnzipTo(src, dest string) error {
r, err := zip.OpenReader(src)
if err != nil {
return err
}
defer r.Close()
for _, f := range r.File {
target := filepath.Join(dest, filepath.FromSlash(f.Name))
if !strings.HasPrefix(filepath.Clean(target), filepath.Clean(dest)+string(os.PathSeparator)) &&
filepath.Clean(target) != filepath.Clean(dest) {
return fmt.Errorf("非法路径: %s", f.Name)
}
if f.FileInfo().IsDir() {
if err := os.MkdirAll(target, 0o755); err != nil {
return err
}
continue
}
if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil {
return err
}
out, err := os.Create(target)
if err != nil {
return err
}
rc, err := f.Open()
if err != nil {
out.Close()
return err
}
_, copyErr := io.Copy(out, rc)
rc.Close()
out.Close()
if copyErr != nil {
return copyErr
}
}
return nil
}