cachefs/pkg/provider/crypto/fs.go
2022-04-06 01:08:19 +02:00

129 lines
2.3 KiB
Go

package crypto
import (
"cachefs/pkg/provider"
"encoding/base64"
"errors"
"io/fs"
"os"
"path/filepath"
"runtime"
"strings"
)
var _ provider.FS = (*FS)(nil)
type FS struct {
fs provider.FS
cbc *aesCbc
key []byte
}
func NewFS(fs provider.FS, key []byte) (*FS, error) {
if len(key) != 32 {
return nil, errors.New("encryption key length must be 32 bytes")
}
cbc, err := newAesCbc(key)
if err != nil {
return nil, err
}
return &FS{
fs: fs,
cbc: cbc,
key: key,
}, nil
}
func (fs *FS) encPath(path string) (n string, err error) {
parts := strings.Split(path, "/")
cparts := make([]string, len(parts))
for i, p := range parts {
if p == "" {
continue
}
cparts[i] = base64.RawURLEncoding.EncodeToString(
fs.cbc.encrypt([]byte(p)),
)
}
n = filepath.Join(cparts...)
return
}
func (fs *FS) decPath(path string) (n string, err error) {
parts := strings.Split(path, "/")
pparts := make([]string, len(parts))
for i, p := range parts {
pparts[i] = base64.RawURLEncoding.EncodeToString(
fs.cbc.decrypt([]byte(p)),
)
}
n = filepath.Join(pparts...)
return
}
func (fs *FS) name(p string) string {
return filepath.Join(fs.Root(), p)
}
func (fs *FS) Stat(p string) (fs.FileInfo, error) {
c, err := fs.encPath(p)
if err != nil {
return nil, err
}
fi, err := fs.fs.Stat(c)
return stat{FileInfo: fi, name: fs.name(p)}, err
}
func (fs *FS) Remove(p string) error {
c, err := fs.encPath(p)
if err != nil {
return err
}
return fs.fs.Remove(c)
}
func (fs *FS) Open(p string) (provider.File, error) {
c, err := fs.encPath(p)
if err != nil {
return nil, err
}
f, err := fs.fs.Open(c)
if err != nil {
return nil, err
}
return newFile(fs.key, p, fs, f)
}
func (fs *FS) OpenFile(p string, flags int, mode os.FileMode) (provider.File, error) {
c, err := fs.encPath(p)
if err != nil {
return nil, err
}
f, err := fs.fs.OpenFile(c, flags, mode)
if err != nil {
return nil, err
}
return newFile(fs.key, p, fs, f)
}
func (fs *FS) MkdirAll(p string, mode os.FileMode) error {
c, err := fs.encPath(p)
if err != nil {
return err
}
return fs.fs.MkdirAll(c, mode)
}
func (fs *FS) Root() string {
return fs.fs.Root()
}
func (fs *FS) Fstat(_ uintptr) (int64, error) {
return 0, provider.ErrFstat
}
func (fs *FS) Close() error {
fs.key = nil
runtime.GC()
return nil
}