feat(backend): zip page index with natural sort, zip-slip rejection, content hash
This commit is contained in:
@@ -0,0 +1,12 @@
|
|||||||
|
package bookfile
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
func Hash(size, modTS int64) string {
|
||||||
|
sum := sha256.Sum256([]byte(fmt.Sprintf("%d:%d", size, modTS)))
|
||||||
|
return hex.EncodeToString(sum[:])[:16]
|
||||||
|
}
|
||||||
@@ -0,0 +1,99 @@
|
|||||||
|
package bookfile
|
||||||
|
|
||||||
|
import (
|
||||||
|
"archive/zip"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"path"
|
||||||
|
"sort"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
var ErrNotZip = errors.New("not a readable zip")
|
||||||
|
var ErrUnsafeZip = errors.New("unsafe zip entry")
|
||||||
|
|
||||||
|
func unsafeEntry(n string) bool {
|
||||||
|
return strings.HasPrefix(n, "/") || strings.Contains(n, "..") || strings.ContainsRune(n, '\\')
|
||||||
|
}
|
||||||
|
|
||||||
|
func isImage(name string) bool {
|
||||||
|
switch strings.ToLower(path.Ext(name)) {
|
||||||
|
case ".jpg", ".jpeg", ".png", ".webp", ".gif", ".avif":
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func PageIndex(f io.ReaderAt, size int64) ([]string, error) {
|
||||||
|
zr, err := zip.NewReader(f, size)
|
||||||
|
if err != nil {
|
||||||
|
return nil, ErrNotZip
|
||||||
|
}
|
||||||
|
var names []string
|
||||||
|
for _, zf := range zr.File {
|
||||||
|
if unsafeEntry(zf.Name) {
|
||||||
|
return nil, fmt.Errorf("%w: %s", ErrUnsafeZip, zf.Name)
|
||||||
|
}
|
||||||
|
if isImage(zf.Name) {
|
||||||
|
names = append(names, zf.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sort.Slice(names, func(i, j int) bool { return NaturalLess(names[i], names[j]) })
|
||||||
|
return names, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func ReadEntry(f io.ReaderAt, size int64, name string) ([]byte, error) {
|
||||||
|
zr, err := zip.NewReader(f, size)
|
||||||
|
if err != nil {
|
||||||
|
return nil, ErrNotZip
|
||||||
|
}
|
||||||
|
for _, zf := range zr.File {
|
||||||
|
if zf.Name == name {
|
||||||
|
rc, err := zf.Open()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rc.Close()
|
||||||
|
return io.ReadAll(io.LimitReader(rc, 64<<20))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil, errors.New("no such entry")
|
||||||
|
}
|
||||||
|
|
||||||
|
func NaturalLess(a, b string) bool {
|
||||||
|
i, j := 0, 0
|
||||||
|
for i < len(a) && j < len(b) {
|
||||||
|
da, db := isDigit(a[i]), isDigit(b[j])
|
||||||
|
switch {
|
||||||
|
case da && db:
|
||||||
|
si, sj := i, j
|
||||||
|
for i < len(a) && isDigit(a[i]) {
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
for j < len(b) && isDigit(b[j]) {
|
||||||
|
j++
|
||||||
|
}
|
||||||
|
na, _ := strconv.Atoi(a[si:i])
|
||||||
|
nb, _ := strconv.Atoi(b[sj:j])
|
||||||
|
if na != nb {
|
||||||
|
return na < nb
|
||||||
|
}
|
||||||
|
if a[si:i] != b[sj:j] {
|
||||||
|
return a[si:i] < b[sj:j]
|
||||||
|
}
|
||||||
|
case !da && !db:
|
||||||
|
if a[i] != b[j] {
|
||||||
|
return a[i] < b[j]
|
||||||
|
}
|
||||||
|
i++
|
||||||
|
j++
|
||||||
|
default:
|
||||||
|
return da // 数字段排在字母前
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return j < len(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func isDigit(c byte) bool { return c >= '0' && c <= '9' }
|
||||||
@@ -0,0 +1,85 @@
|
|||||||
|
package bookfile
|
||||||
|
|
||||||
|
import (
|
||||||
|
"archive/zip"
|
||||||
|
"bytes"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func zipOf(t *testing.T, names ...string) *bytes.Reader {
|
||||||
|
t.Helper()
|
||||||
|
buf := &bytes.Buffer{}
|
||||||
|
zw := zip.NewWriter(buf)
|
||||||
|
for i, n := range names {
|
||||||
|
w, err := zw.Create(n)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
w.Write([]byte{byte(i)})
|
||||||
|
}
|
||||||
|
zw.Close()
|
||||||
|
return bytes.NewReader(buf.Bytes())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPageIndexSortAndFilter(t *testing.T) {
|
||||||
|
r := zipOf(t, "page10.jpg", "page2.jpg", "page1.jpg", "cover.PNG", "meta.xml", "sub/3.webp", "readme.txt")
|
||||||
|
idx, err := PageIndex(r, int64(r.Len()))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
want := []string{"cover.PNG", "page1.jpg", "page2.jpg", "page10.jpg", "sub/3.webp"}
|
||||||
|
if len(idx) != len(want) {
|
||||||
|
t.Fatalf("got %v", idx)
|
||||||
|
}
|
||||||
|
for i := range want {
|
||||||
|
if idx[i] != want[i] {
|
||||||
|
t.Fatalf("got %v want %v", idx, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPageIndexRejectsSlip(t *testing.T) {
|
||||||
|
for _, bad := range []string{"../evil.jpg", "/etc/passwd.jpg", "a\\..\\b.jpg"} {
|
||||||
|
r := zipOf(t, bad)
|
||||||
|
if _, err := PageIndex(r, int64(r.Len())); err == nil {
|
||||||
|
t.Fatalf("entry %q must be rejected", bad)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReadEntry(t *testing.T) {
|
||||||
|
buf := &bytes.Buffer{}
|
||||||
|
zw := zip.NewWriter(buf)
|
||||||
|
w, _ := zw.Create("p1.jpg")
|
||||||
|
w.Write([]byte("jpegbytes"))
|
||||||
|
zw.Close()
|
||||||
|
r := bytes.NewReader(buf.Bytes())
|
||||||
|
got, err := ReadEntry(r, int64(r.Len()), "p1.jpg")
|
||||||
|
if err != nil || string(got) != "jpegbytes" {
|
||||||
|
t.Fatalf("%q %v", got, err)
|
||||||
|
}
|
||||||
|
if _, err := ReadEntry(r, int64(r.Len()), "nope.jpg"); err == nil {
|
||||||
|
t.Fatal("missing entry must error")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNaturalLess(t *testing.T) {
|
||||||
|
pairs := [][2]string{{"2", "10"}, {"a2b", "a10b"}, {"007", "7"}, {"a", "b"}, {"A", "a"}}
|
||||||
|
for _, p := range pairs {
|
||||||
|
if !NaturalLess(p[0], p[1]) {
|
||||||
|
t.Errorf("NaturalLess(%q,%q) want true", p[0], p[1])
|
||||||
|
}
|
||||||
|
if NaturalLess(p[1], p[0]) {
|
||||||
|
t.Errorf("NaturalLess(%q,%q) want false", p[1], p[0])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHash(t *testing.T) {
|
||||||
|
if Hash(100, 200) == Hash(101, 200) || Hash(100, 200) != Hash(100, 200) {
|
||||||
|
t.Fatal("hash broken")
|
||||||
|
}
|
||||||
|
if len(Hash(1, 2)) != 16 {
|
||||||
|
t.Fatal("hash length")
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user