feat(upload): extract upload subsystem + move sweep to scanner ticker (B16)

- internal/upload owns chunked-upload domain logic (fingerprint resume,
  part tmp+rename, assemble, session sweep) and the shared UniquePath
  helper used by both single-file and chunked completion paths
- handlers/uploads.go is now HTTP-only: bind params, call upload.U, map
  sentinel errors to the unchanged status/code/message contract
- B16: session sweep moved out of the UploadInit request path onto the
  scanner ticker via a Sweeper hook (upload.U satisfies it)
- unit tests for the package without PG/Redis; contract pinned by the
  existing handler integration tests (all green)
This commit is contained in:
2026-09-14 22:00:25 +08:00
parent 443f4acfa9
commit fbd5cd243d
10 changed files with 708 additions and 263 deletions
+3 -2
View File
@@ -10,11 +10,12 @@ import (
"booklib/internal/redispkg" "booklib/internal/redispkg"
"booklib/internal/scanner" "booklib/internal/scanner"
"booklib/internal/store" "booklib/internal/store"
"booklib/internal/upload"
) )
func NewRouter(cfg *config.Config, st *store.Store, rdb *redispkg.R, sc *scanner.Scanner) *gin.Engine { func NewRouter(cfg *config.Config, st *store.Store, rdb *redispkg.R, sc *scanner.Scanner, up *upload.U) *gin.Engine {
gin.SetMode(gin.ReleaseMode) gin.SetMode(gin.ReleaseMode)
h := handlers.New(cfg, st, rdb, sc) h := handlers.New(cfg, st, rdb, sc, up)
r := gin.New() r := gin.New()
if e := r.SetTrustedProxies(cfg.TrustedProxies); e != nil { if e := r.SetTrustedProxies(cfg.TrustedProxies); e != nil {
panic(e) panic(e)
+1 -1
View File
@@ -16,7 +16,7 @@ func testCfg() *config.Config {
} }
func TestHealthz(t *testing.T) { func TestHealthz(t *testing.T) {
r := NewRouter(testCfg(), nil, redispkg.New(""), nil) r := NewRouter(testCfg(), nil, redispkg.New(""), nil, nil)
req := httptest.NewRequest(http.MethodGet, "/api/healthz", nil) req := httptest.NewRequest(http.MethodGet, "/api/healthz", nil)
w := httptest.NewRecorder() w := httptest.NewRecorder()
r.ServeHTTP(w, req) r.ServeHTTP(w, req)
+4 -2
View File
@@ -21,6 +21,7 @@ import (
"booklib/internal/redispkg" "booklib/internal/redispkg"
"booklib/internal/scanner" "booklib/internal/scanner"
"booklib/internal/store" "booklib/internal/store"
"booklib/internal/upload"
) )
func testCfg() *config.Config { func testCfg() *config.Config {
@@ -57,8 +58,9 @@ func setupAPI(t *testing.T) (*store.Store, *scanner.Scanner, http.Handler, strin
cfg.BooksDir = booksDir cfg.BooksDir = booksDir
cfg.CacheDir = t.TempDir() cfg.CacheDir = t.TempDir()
rdb := redispkg.New(os.Getenv("REDIS_URL")) rdb := redispkg.New(os.Getenv("REDIS_URL"))
sc := scanner.New(st, cfg, rdb) up := upload.New(cfg.BooksDir, cfg.UploadMaxMB)
r := api.NewRouter(cfg, st, rdb, sc) sc := scanner.New(st, cfg, rdb, up)
r := api.NewRouter(cfg, st, rdb, sc, up)
if u := os.Getenv("REDIS_URL"); u != "" { // 测试卫生: 共享 redis 上重置登录限流桶, 防跨测试累计 429 if u := os.Getenv("REDIS_URL"); u != "" { // 测试卫生: 共享 redis 上重置登录限流桶, 防跨测试累计 429
if opt, e := redis.ParseURL(u); e == nil { if opt, e := redis.ParseURL(u); e == nil {
rc := redis.NewClient(opt) rc := redis.NewClient(opt)
+4 -2
View File
@@ -18,6 +18,7 @@ import (
"booklib/internal/redispkg" "booklib/internal/redispkg"
"booklib/internal/scanner" "booklib/internal/scanner"
"booklib/internal/store" "booklib/internal/store"
"booklib/internal/upload"
) )
type H struct { type H struct {
@@ -25,10 +26,11 @@ type H struct {
st *store.Store st *store.Store
rdb *redispkg.R rdb *redispkg.R
sc *scanner.Scanner sc *scanner.Scanner
up *upload.U
} }
func New(cfg *config.Config, st *store.Store, rdb *redispkg.R, sc *scanner.Scanner) *H { func New(cfg *config.Config, st *store.Store, rdb *redispkg.R, sc *scanner.Scanner, up *upload.U) *H {
return &H{cfg: cfg, st: st, rdb: rdb, sc: sc} return &H{cfg: cfg, st: st, rdb: rdb, sc: sc, up: up}
} }
func err(c *gin.Context, status int, code, msg string) { func err(c *gin.Context, status int, code, msg string) {
+2 -27
View File
@@ -136,12 +136,12 @@ func (h *H) Upload(c *gin.Context) {
return return
} }
// B12: retry on O_EXCL collision — concurrent uploads with the same name // B12: retry on O_EXCL collision — concurrent uploads with the same name
// can both get the same candidate from uniquePath (stat-then-create race). // can both get the same candidate from UniquePath (stat-then-create race).
var dst, tmp string var dst, tmp string
var out *os.File var out *os.File
for attempt := 0; attempt < 5; attempt++ { for attempt := 0; attempt < 5; attempt++ {
var e error var e error
dst, e = h.uniquePath(root, name) dst, e = h.up.UniquePath(root, name)
if e != nil { if e != nil {
err(c, http.StatusForbidden, "forbidden", e.Error()) err(c, http.StatusForbidden, "forbidden", e.Error())
return return
@@ -190,28 +190,3 @@ func (h *H) Upload(c *gin.Context) {
} }
c.JSON(http.StatusAccepted, gin.H{"accepted": true, "path": strings.TrimPrefix(dst, root+string(os.PathSeparator))}) c.JSON(http.StatusAccepted, gin.H{"accepted": true, "path": strings.TrimPrefix(dst, root+string(os.PathSeparator))})
} }
// uniquePath 清洗后的 name 必须仍在 root 内;重名加 " (n)" 后缀
func (h *H) uniquePath(root, name string) (string, error) {
ext := filepath.Ext(name)
base := strings.TrimSuffix(name, ext)
for i := 0; ; i++ {
cand := base + ext
if i > 0 {
cand = base + " (" + strconv.Itoa(i) + ")" + ext
}
p := filepath.Join(root, cand)
if filepath.Clean(p) != filepath.Join(root, filepath.Clean(cand)) ||
!strings.HasPrefix(filepath.Clean(p), root+string(os.PathSeparator)) {
return "", os.ErrInvalid
}
if _, e := os.Stat(p); os.IsNotExist(e) {
return p, nil
} else if e != nil {
return "", e
}
if i > 999 {
return "", os.ErrExist
}
}
}
+67 -224
View File
@@ -1,86 +1,22 @@
package handlers package handlers
import ( import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors" "errors"
"fmt"
"io"
"net/http" "net/http"
"os" "os"
"path/filepath"
"sort"
"strconv" "strconv"
"strings"
"time"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"booklib/internal/bookfile" "booklib/internal/ports"
"booklib/internal/upload"
) )
// 分片上传:init(指纹→确定性 uploadId,天然支持续传)→ PUT parts → complete 拼接原子落盘。 // 分片上传:HTTP 层只做参数绑定与错误映射,域逻辑(指纹续传/分片落盘/拼接/清扫)
// 会话即 <BooksDir>/.uploads/<uid>/(meta.json + parts/N),无独立状态存储;24h 未 complete opportunistic 清扫。 // 全部在 internal/upload;会话清扫由 scanner ticker 接管(B16),不在请求路径。
const ( // maxChunkBytes 与 upload 包内常量同值,作为 PutPart 的防御性读取上限。
maxChunkBytes = 32 << 20 const maxChunkBytes = 32 << 20
defaultChunk = 8 << 20
uploadSessTTL = 24 * time.Hour
uploadSessionIn = ".uploads"
)
type uploadMeta struct {
Name string `json:"name"`
Size int64 `json:"size"`
ChunkSize int64 `json:"chunkSize"`
LibraryID int64 `json:"libraryId"`
}
func validUploadID(s string) bool {
if len(s) != 32 {
return false
}
for _, r := range s {
if !((r >= '0' && r <= '9') || (r >= 'a' && r <= 'f')) {
return false
}
}
return true
}
func uploadIDFor(libID int64, name string, size, chunk int64) string {
h := sha256.Sum256([]byte(fmt.Sprintf("%d|%s|%d|%d", libID, name, size, chunk)))
return hex.EncodeToString(h[:16])
}
func (h *H) uploadDir(uid string) string {
return filepath.Join(filepath.Clean(h.cfg.BooksDir), uploadSessionIn, uid)
}
func chunkRange(m uploadMeta, i int64) (int64, int64) {
lo := i * m.ChunkSize
hi := min(lo+m.ChunkSize, m.Size)
return lo, hi
}
func (h *H) numParts(m uploadMeta) int64 {
return (m.Size + m.ChunkSize - 1) / m.ChunkSize
}
// sweepUploads 删除过期会话目录;尽力而为,失败不影响主流程
func (h *H) sweepUploads() {
base := filepath.Join(filepath.Clean(h.cfg.BooksDir), uploadSessionIn)
es, e := os.ReadDir(base)
if e != nil {
return
}
for _, en := range es {
if fi, e := en.Info(); e == nil && time.Since(fi.ModTime()) > uploadSessTTL {
os.RemoveAll(filepath.Join(base, en.Name()))
}
}
}
func (h *H) UploadInit(c *gin.Context) { func (h *H) UploadInit(c *gin.Context) {
lib, ok := h.getLibrary(c) lib, ok := h.getLibrary(c)
@@ -99,154 +35,47 @@ func (h *H) UploadInit(c *gin.Context) {
err(c, http.StatusBadRequest, "bad_request", "name, size required") err(c, http.StatusBadRequest, "bad_request", "name, size required")
return return
} }
name := bookfile.SafeName(req.Name) uid, e := h.up.Init(c, lib.ID, req.Name, req.Size, req.ChunkSize)
if name == "" {
err(c, http.StatusBadRequest, "bad_request", "bad name")
return
}
if bookfile.FormatFromExt(name) == "" {
err(c, http.StatusBadRequest, "bad_format", "extension must be cbz/pdf/epub/txt/md")
return
}
if req.Size <= 0 {
err(c, http.StatusBadRequest, "bad_request", "bad size")
return
}
if req.Size > h.cfg.UploadMaxMB<<20 {
err(c, http.StatusRequestEntityTooLarge, "too_large", "file exceeds upload limit of "+strconv.FormatInt(h.cfg.UploadMaxMB, 10)+"MB")
return
}
if req.ChunkSize == 0 {
req.ChunkSize = defaultChunk
}
if req.ChunkSize > maxChunkBytes {
err(c, http.StatusBadRequest, "bad_request", "chunkSize must be <= 33554432")
return
}
uid := uploadIDFor(lib.ID, name, req.Size, req.ChunkSize)
dir := h.uploadDir(uid)
meta := uploadMeta{Name: name, Size: req.Size, ChunkSize: req.ChunkSize, LibraryID: lib.ID}
if b, e := os.ReadFile(filepath.Join(dir, "meta.json")); e == nil {
var old uploadMeta
if json.Unmarshal(b, &old) == nil && old == meta { // 同指纹 → 复用会话(续传)
c.JSON(http.StatusOK, gin.H{"uploadId": uid})
return
}
os.RemoveAll(dir) // 指纹撞上但内容不同 → 从头再来
}
h.sweepUploads()
if e := os.MkdirAll(filepath.Join(dir, "parts"), 0o755); e != nil {
err(c, http.StatusInternalServerError, "internal", "create session")
return
}
b, _ := json.Marshal(meta)
if e := os.WriteFile(filepath.Join(dir, "meta.json"), b, 0o644); e != nil {
err(c, http.StatusInternalServerError, "internal", "write meta")
return
}
c.JSON(http.StatusOK, gin.H{"uploadId": uid})
}
// loadMeta 校验 uid 与路径,404/400 已回复
func (h *H) loadMeta(c *gin.Context) (uploadMeta, string, bool) {
uid := c.Param("uid")
if !validUploadID(uid) {
err(c, http.StatusBadRequest, "bad_request", "bad upload id")
return uploadMeta{}, "", false
}
dir := h.uploadDir(uid)
var m uploadMeta
b, e := os.ReadFile(filepath.Join(dir, "meta.json"))
if e != nil { if e != nil {
err(c, http.StatusNotFound, "not_found", "no such upload") h.mapUploadErr(c, e)
return uploadMeta{}, "", false return
} }
if json.Unmarshal(b, &m) != nil { c.JSON(http.StatusOK, gin.H{"uploadId": uid})
err(c, http.StatusInternalServerError, "internal", "corrupt session")
return uploadMeta{}, "", false
}
return m, dir, true
} }
func (h *H) UploadStatus(c *gin.Context) { func (h *H) UploadStatus(c *gin.Context) {
_, dir, ok := h.loadMeta(c) recv, e := h.up.Status(c, c.Param("uid"))
if !ok { if e != nil {
h.mapUploadErr(c, e)
return return
} }
recv := []int64{}
es, e := os.ReadDir(filepath.Join(dir, "parts"))
if e == nil {
for _, en := range es {
if i, e := strconv.ParseInt(en.Name(), 10, 64); e == nil {
recv = append(recv, i)
}
}
}
sort.Slice(recv, func(i, j int) bool { return recv[i] < recv[j] })
c.JSON(http.StatusOK, gin.H{"received": recv}) c.JSON(http.StatusOK, gin.H{"received": recv})
} }
func (h *H) UploadPart(c *gin.Context) { func (h *H) UploadPart(c *gin.Context) {
m, dir, ok := h.loadMeta(c) uid := c.Param("uid")
if !ok {
return
}
idx, e := strconv.ParseInt(c.Param("index"), 10, 64) idx, e := strconv.ParseInt(c.Param("index"), 10, 64)
if e != nil || idx < 0 || idx >= h.numParts(m) { if e != nil {
err(c, http.StatusBadRequest, "bad_request", "bad part index") err(c, http.StatusBadRequest, "bad_request", "bad part index")
return return
} }
lo, hi := chunkRange(m, idx) // maxSize 防御性上限:分片声明大小由会话 meta 决定,这里再垫一层 32MB 全局上限
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, hi-lo) e = h.up.PutPart(c, uid, idx, c.Request.Body, maxChunkBytes)
// B4: write to .tmp then rename — prevents truncated parts from being
// reported as "received" by UploadStatus if the process crashes mid-write.
p := filepath.Join(dir, "parts", strconv.FormatInt(idx, 10))
tmp := p + ".tmp"
f, e := os.OpenFile(tmp, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o644)
if e != nil { if e != nil {
err(c, http.StatusInternalServerError, "internal", "create part") h.mapUploadErr(c, e)
return
}
n, e := io.Copy(f, c.Request.Body)
f.Close()
if e != nil || n != hi-lo {
os.Remove(tmp)
var mbe *http.MaxBytesError
code, msg := "too_large", "part size mismatch"
if errors.As(e, &mbe) {
msg = "part exceeds declared size"
}
err(c, http.StatusRequestEntityTooLarge, code, msg)
return
}
if e := os.Rename(tmp, p); e != nil {
os.Remove(tmp)
err(c, http.StatusInternalServerError, "internal", "rename part")
return return
} }
c.JSON(http.StatusAccepted, gin.H{"accepted": true}) c.JSON(http.StatusAccepted, gin.H{"accepted": true})
} }
func (h *H) UploadComplete(c *gin.Context) { func (h *H) UploadComplete(c *gin.Context) {
m, dir, ok := h.loadMeta(c) uid := c.Param("uid")
if !ok { libID, e := h.up.LibraryID(c, uid)
if e != nil {
h.mapUploadErr(c, e)
return return
} }
var total int64 lib, e := h.st.GetLibrary(c, libID)
for i := int64(0); i < h.numParts(m); i++ {
lo, hi := chunkRange(m, i)
fi, e := os.Stat(filepath.Join(dir, "parts", strconv.FormatInt(i, 10)))
if e != nil || fi.Size() != hi-lo {
err(c, http.StatusBadRequest, "bad_request", "upload incomplete; missing or corrupt parts, re-upload them")
return
}
total += fi.Size()
}
if total != m.Size {
err(c, http.StatusBadRequest, "bad_request", "total size mismatch")
return
}
lib, e := h.st.GetLibrary(c, m.LibraryID)
if e != nil { if e != nil {
dbErr(c, e) dbErr(c, e)
return return
@@ -255,39 +84,53 @@ func (h *H) UploadComplete(c *gin.Context) {
if !ok { if !ok {
return return
} }
dst, e := h.uniquePath(root, m.Name) rel, e := h.up.Complete(c, uid, root)
if e != nil { if e != nil {
h.mapUploadErr(c, e)
return
}
c.JSON(http.StatusAccepted, gin.H{"accepted": true, "path": rel})
}
// mapUploadErr 把 upload 包的 sentinel 错误映射为原契约的 status/code/message。
func (h *H) mapUploadErr(c *gin.Context, e error) {
switch {
case errors.Is(e, ports.ErrTooLarge):
err(c, http.StatusRequestEntityTooLarge, "too_large", "file exceeds upload limit of "+strconv.FormatInt(h.cfg.UploadMaxMB, 10)+"MB")
case errors.Is(e, upload.ErrBadName):
err(c, http.StatusBadRequest, "bad_request", "bad name")
case errors.Is(e, upload.ErrBadFormat):
err(c, http.StatusBadRequest, "bad_format", "extension must be cbz/pdf/epub/txt/md")
case errors.Is(e, upload.ErrBadSize):
err(c, http.StatusBadRequest, "bad_request", "bad size")
case errors.Is(e, upload.ErrBadChunk):
err(c, http.StatusBadRequest, "bad_request", "chunkSize must be <= 33554432")
case errors.Is(e, upload.ErrBadUploadID):
err(c, http.StatusBadRequest, "bad_request", "bad upload id")
case errors.Is(e, upload.ErrBadIndex):
err(c, http.StatusBadRequest, "bad_request", "bad part index")
case errors.Is(e, upload.ErrPartTooBig):
err(c, http.StatusRequestEntityTooLarge, "too_large", "part exceeds declared size")
case errors.Is(e, upload.ErrPartSizeMismatch):
err(c, http.StatusRequestEntityTooLarge, "too_large", "part size mismatch")
case errors.Is(e, ports.ErrNotFound):
err(c, http.StatusNotFound, "not_found", "no such upload")
case errors.Is(e, upload.ErrCorrupt):
err(c, http.StatusInternalServerError, "internal", "corrupt session")
case errors.Is(e, ports.ErrIncomplete):
err(c, http.StatusBadRequest, "bad_request", "upload incomplete; missing or corrupt parts, re-upload them")
case errors.Is(e, ports.ErrSizeMismatch):
err(c, http.StatusBadRequest, "bad_request", "total size mismatch")
case errors.Is(e, os.ErrInvalid), errors.Is(e, os.ErrExist):
err(c, http.StatusForbidden, "forbidden", e.Error()) err(c, http.StatusForbidden, "forbidden", e.Error())
case errors.Is(e, os.ErrPermission), errors.Is(e, os.ErrClosed):
err(c, http.StatusInternalServerError, "internal", "io error")
default:
var oe *upload.OpError
if errors.As(e, &oe) {
err(c, http.StatusInternalServerError, "internal", oe.Op)
return return
} }
tmp := filepath.Join(dir, "assembled") err(c, http.StatusInternalServerError, "internal", "upload failed")
out, e := os.OpenFile(tmp, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o644)
if e != nil {
err(c, http.StatusInternalServerError, "internal", "create tmp")
return
} }
for i := int64(0); i < h.numParts(m); i++ {
pf, e := os.Open(filepath.Join(dir, "parts", strconv.FormatInt(i, 10)))
if e != nil {
out.Close()
err(c, http.StatusInternalServerError, "internal", "open part")
return
}
if _, e := io.Copy(out, pf); e != nil {
pf.Close()
out.Close()
os.Remove(tmp)
err(c, http.StatusInternalServerError, "internal", "assemble")
return
}
pf.Close()
}
out.Close()
if e := os.Rename(tmp, dst); e != nil { // 原子落盘,scanner 自动收编
os.Remove(tmp)
err(c, http.StatusInternalServerError, "internal", "rename")
return
}
os.RemoveAll(dir)
c.JSON(http.StatusAccepted, gin.H{"accepted": true, "path": strings.TrimPrefix(dst, root+string(os.PathSeparator))})
} }
+4 -2
View File
@@ -16,6 +16,7 @@ import (
"booklib/internal/scanner" "booklib/internal/scanner"
"booklib/internal/seed" "booklib/internal/seed"
"booklib/internal/store" "booklib/internal/store"
"booklib/internal/upload"
) )
func main() { func main() {
@@ -39,11 +40,12 @@ func main() {
log.Fatalf("seed: %v", err) log.Fatalf("seed: %v", err)
} }
rdb := redispkg.New(cfg.RedisURL) rdb := redispkg.New(cfg.RedisURL)
sc := scanner.New(st, cfg, rdb) up := upload.New(cfg.BooksDir, cfg.UploadMaxMB)
sc := scanner.New(st, cfg, rdb, up) // B16: sweep rides the scan ticker
go sc.Run(ctx) go sc.Run(ctx)
serveErr := make(chan error, 1) serveErr := make(chan error, 1)
srv := &http.Server{Addr: cfg.Addr, Handler: api.NewRouter(cfg, st, rdb, sc), srv := &http.Server{Addr: cfg.Addr, Handler: api.NewRouter(cfg, st, rdb, sc, up),
ReadHeaderTimeout: 10 * time.Second} ReadHeaderTimeout: 10 * time.Second}
go func() { go func() {
log.Printf("listening on %s", cfg.Addr) log.Printf("listening on %s", cfg.Addr)
+13 -2
View File
@@ -19,17 +19,23 @@ import (
"booklib/internal/store" "booklib/internal/store"
) )
// Sweeper 是 scanner 每轮顺手调用的清理钩子;upload.U 满足它(B16)。
type Sweeper interface {
Sweep(ctx context.Context) error
}
type Scanner struct { type Scanner struct {
st *store.Store st *store.Store
cfg *config.Config cfg *config.Config
rdb *redispkg.R rdb *redispkg.R
sweepers []Sweeper
// B9-②: per-library single-flight — concurrent scan triggers for the same // B9-②: per-library single-flight — concurrent scan triggers for the same
// library are merged into one execution, even without redis. // library are merged into one execution, even without redis.
flights sync.Map // map[int64]*sync.WaitGroup flights sync.Map // map[int64]*sync.WaitGroup
} }
func New(st *store.Store, cfg *config.Config, rdb *redispkg.R) *Scanner { func New(st *store.Store, cfg *config.Config, rdb *redispkg.R, sweepers ...Sweeper) *Scanner {
return &Scanner{st: st, cfg: cfg, rdb: rdb} return &Scanner{st: st, cfg: cfg, rdb: rdb, sweepers: sweepers}
} }
func (s *Scanner) Run(ctx context.Context) { func (s *Scanner) Run(ctx context.Context) {
@@ -40,6 +46,11 @@ func (s *Scanner) Run(ctx context.Context) {
case <-ctx.Done(): case <-ctx.Done():
return return
case <-t.C: case <-t.C:
for _, sw := range s.sweepers { // B16: 上传会话清扫随扫描周期跑
if err := sw.Sweep(ctx); err != nil {
log.Printf("scan: sweep: %v", err)
}
}
libs, err := s.st.ListLibraries(ctx) libs, err := s.st.ListLibraries(ctx)
if err != nil { if err != nil {
log.Printf("scan: list libraries: %v", err) log.Printf("scan: list libraries: %v", err)
+357
View File
@@ -0,0 +1,357 @@
// Package upload owns the resumable chunked-upload subsystem and the shared
// unique-path placement helper used by both chunked and single-file uploads.
//
// A session lives at <BooksDir>/.uploads/<uid>/ (meta.json + parts/N). The uid
// is a fingerprint of (libID, name, size, chunkSize), so re-initialising the
// same file resumes the existing session instead of restarting it. Expired
// sessions are swept by the scanner ticker (B16), not on the request path.
//
// This package is HTTP-free: it returns sentinel errors that the handlers map
// onto status/code/message tuples. The prior client-facing contract is
// preserved exactly.
package upload
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
"time"
"booklib/internal/bookfile"
"booklib/internal/ports"
)
const (
maxChunkBytes = 32 << 20
defaultChunk = 8 << 20
uploadSessTTL = 24 * time.Hour
uploadSessionIn = ".uploads"
)
// Sentinel errors returned to handlers for status/code mapping. The four shared
// with ports (ErrTooLarge/ErrNotFound/ErrIncomplete/ErrSizeMismatch) live in the
// ports package so the interface contract and the implementation agree.
var (
ErrBadName = errors.New("bad name")
ErrBadFormat = errors.New("bad format")
ErrBadSize = errors.New("bad size")
ErrBadChunk = errors.New("bad chunk size")
ErrBadUploadID = errors.New("bad upload id")
ErrBadIndex = errors.New("bad part index")
ErrPartTooBig = errors.New("part exceeds declared size")
ErrPartSizeMismatch = errors.New("part size mismatch")
ErrCorrupt = errors.New("corrupt session")
)
// OpError wraps an internal filesystem/IO failure with the short operation label
// that the client-facing 500 message uses, preserving the prior contract strings
// ("create session", "write meta", "create part", ...).
type OpError struct {
Op string
Err error
}
func (e *OpError) Error() string { return e.Op + ": " + e.Err.Error() }
func (e *OpError) Unwrap() error { return e.Err }
func opErr(op string, err error) error { return &OpError{Op: op, Err: err} }
type uploadMeta struct {
Name string `json:"name"`
Size int64 `json:"size"`
ChunkSize int64 `json:"chunkSize"`
LibraryID int64 `json:"libraryId"`
}
// U is the upload subsystem. It is stateless beyond the filesystem session dir.
type U struct {
booksDir string
uploadMaxMB int64
}
// New builds the upload subsystem. booksDir is the storage root (sessions live
// under booksDir/.uploads); uploadMaxMB caps the total declared file size.
func New(booksDir string, uploadMaxMB int64) *U {
return &U{booksDir: booksDir, uploadMaxMB: uploadMaxMB}
}
// compile-time proof that *U satisfies the consumer-side interface.
var _ ports.UploadSessions = (*U)(nil)
// ---------- pure helpers ----------
func validUploadID(s string) bool {
if len(s) != 32 {
return false
}
for _, r := range s {
if !((r >= '0' && r <= '9') || (r >= 'a' && r <= 'f')) {
return false
}
}
return true
}
func uploadIDFor(libID int64, name string, size, chunk int64) string {
h := sha256.Sum256([]byte(fmt.Sprintf("%d|%s|%d|%d", libID, name, size, chunk)))
return hex.EncodeToString(h[:16])
}
func (u *U) uploadDir(uid string) string {
return filepath.Join(filepath.Clean(u.booksDir), uploadSessionIn, uid)
}
func chunkRange(m uploadMeta, i int64) (int64, int64) {
lo := i * m.ChunkSize
hi := min(lo+m.ChunkSize, m.Size)
return lo, hi
}
func numParts(m uploadMeta) int64 {
return (m.Size + m.ChunkSize - 1) / m.ChunkSize
}
// ---------- session meta ----------
// loadMeta validates uid + reads meta.json. Returns ErrBadUploadID,
// ports.ErrNotFound, or ErrCorrupt on failure.
func (u *U) loadMeta(uid string) (uploadMeta, string, error) {
if !validUploadID(uid) {
return uploadMeta{}, "", ErrBadUploadID
}
dir := u.uploadDir(uid)
b, e := os.ReadFile(filepath.Join(dir, "meta.json"))
if e != nil {
return uploadMeta{}, "", ports.ErrNotFound
}
var m uploadMeta
if json.Unmarshal(b, &m) != nil {
return uploadMeta{}, "", ErrCorrupt
}
return m, dir, nil
}
// LibraryID returns the target library id recorded in the session, so the
// handler can resolve+validate the library root before calling Complete.
func (u *U) LibraryID(ctx context.Context, uid string) (int64, error) {
m, _, e := u.loadMeta(uid)
if e != nil {
return 0, e
}
return m.LibraryID, nil
}
// ---------- public API (satisfies ports.UploadSessions) ----------
// Init validates the declared upload, then creates or resumes a session. The
// returned uid is deterministic for a given (libID, safeName, size, chunkSize),
// so a re-init of the same file resumes; a fingerprint collision with different
// content restarts the session.
func (u *U) Init(ctx context.Context, libID int64, name string, size, chunkSize int64) (string, error) {
safe := bookfile.SafeName(name)
if safe == "" {
return "", ErrBadName
}
if bookfile.FormatFromExt(safe) == "" {
return "", ErrBadFormat
}
if size <= 0 {
return "", ErrBadSize
}
if size > u.uploadMaxMB<<20 {
return "", ports.ErrTooLarge
}
if chunkSize == 0 {
chunkSize = defaultChunk
}
if chunkSize > maxChunkBytes {
return "", ErrBadChunk
}
uid := uploadIDFor(libID, safe, size, chunkSize)
dir := u.uploadDir(uid)
meta := uploadMeta{Name: safe, Size: size, ChunkSize: chunkSize, LibraryID: libID}
if b, e := os.ReadFile(filepath.Join(dir, "meta.json")); e == nil {
var old uploadMeta
if json.Unmarshal(b, &old) == nil && old == meta { // same fingerprint → resume
return uid, nil
}
os.RemoveAll(dir) // fingerprint collided but content differs → restart
}
if e := os.MkdirAll(filepath.Join(dir, "parts"), 0o755); e != nil {
return "", opErr("create session", e)
}
b, _ := json.Marshal(meta)
if e := os.WriteFile(filepath.Join(dir, "meta.json"), b, 0o644); e != nil {
return "", opErr("write meta", e)
}
return uid, nil
}
// Status returns the sorted indices of parts already received.
func (u *U) Status(ctx context.Context, uid string) ([]int64, error) {
_, dir, e := u.loadMeta(uid)
if e != nil {
return nil, e
}
recv := []int64{}
es, e := os.ReadDir(filepath.Join(dir, "parts"))
if e == nil {
for _, en := range es {
if i, e := strconv.ParseInt(en.Name(), 10, 64); e == nil {
recv = append(recv, i)
}
}
}
sort.Slice(recv, func(i, j int) bool { return recv[i] < recv[j] })
return recv, nil
}
// PutPart writes one part via tmp+rename (B4: a truncated part is never reported
// as received). body is read up to the declared part size; reading past it yields
// ErrPartTooBig, a short read yields ErrPartSizeMismatch. maxSize, when > 0, is a
// defensive ceiling on bytes read.
func (u *U) PutPart(ctx context.Context, uid string, index int64, body io.Reader, maxSize int64) error {
m, dir, e := u.loadMeta(uid)
if e != nil {
return e
}
if index < 0 || index >= numParts(m) {
return ErrBadIndex
}
lo, hi := chunkRange(m, index)
want := hi - lo
limit := want + 1
if maxSize > 0 && maxSize+1 < limit {
limit = maxSize + 1
}
p := filepath.Join(dir, "parts", strconv.FormatInt(index, 10))
tmp := p + ".tmp"
f, e := os.OpenFile(tmp, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o644)
if e != nil {
return opErr("create part", e)
}
n, copyErr := io.Copy(f, io.LimitReader(body, limit))
f.Close()
if copyErr != nil {
os.Remove(tmp)
return ErrPartSizeMismatch
}
if n > want {
os.Remove(tmp)
return ErrPartTooBig
}
if n != want {
os.Remove(tmp)
return ErrPartSizeMismatch
}
if e := os.Rename(tmp, p); e != nil {
os.Remove(tmp)
return opErr("rename part", e)
}
return nil
}
// Complete verifies every part is present and correctly sized, assembles them
// into a tmp file, then atomically renames it into root under a unique name.
// It returns the path relative to root. The session dir is removed on success.
func (u *U) Complete(ctx context.Context, uid, root string) (string, error) {
m, dir, e := u.loadMeta(uid)
if e != nil {
return "", e
}
var total int64
for i := int64(0); i < numParts(m); i++ {
lo, hi := chunkRange(m, i)
fi, e := os.Stat(filepath.Join(dir, "parts", strconv.FormatInt(i, 10)))
if e != nil || fi.Size() != hi-lo {
return "", ports.ErrIncomplete
}
total += fi.Size()
}
if total != m.Size {
return "", ports.ErrSizeMismatch
}
dst, e := u.UniquePath(root, m.Name)
if e != nil {
return "", e // os.ErrInvalid / os.ErrExist → handler maps to 403
}
tmp := filepath.Join(dir, "assembled")
out, e := os.OpenFile(tmp, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o644)
if e != nil {
return "", opErr("create tmp", e)
}
for i := int64(0); i < numParts(m); i++ {
pf, e := os.Open(filepath.Join(dir, "parts", strconv.FormatInt(i, 10)))
if e != nil {
out.Close()
return "", opErr("open part", e)
}
_, copyErr := io.Copy(out, pf)
pf.Close()
if copyErr != nil {
out.Close()
os.Remove(tmp)
return "", opErr("assemble", copyErr)
}
}
out.Close()
if e := os.Rename(tmp, dst); e != nil { // atomic placement; scanner picks it up
os.Remove(tmp)
return "", opErr("rename", e)
}
os.RemoveAll(dir)
return strings.TrimPrefix(dst, root+string(os.PathSeparator)), nil
}
// Sweep removes session dirs untouched for longer than the TTL. Best-effort:
// a missing/unreadable base dir is not an error. Run from the scanner ticker
// (B16) instead of the request path.
func (u *U) Sweep(ctx context.Context) error {
base := filepath.Join(filepath.Clean(u.booksDir), uploadSessionIn)
es, e := os.ReadDir(base)
if e != nil {
return nil // no sessions yet
}
for _, en := range es {
if fi, e := en.Info(); e == nil && time.Since(fi.ModTime()) > uploadSessTTL {
os.RemoveAll(filepath.Join(base, en.Name()))
}
}
return nil
}
// UniquePath returns a path under root for name that does not yet exist,
// appending " (n)" on collision. The cleaned name must stay inside root.
// Shared by single-file and chunked upload completion.
func (u *U) UniquePath(root, name string) (string, error) {
ext := filepath.Ext(name)
base := strings.TrimSuffix(name, ext)
for i := 0; ; i++ {
cand := base + ext
if i > 0 {
cand = base + " (" + strconv.Itoa(i) + ")" + ext
}
p := filepath.Join(root, cand)
if filepath.Clean(p) != filepath.Join(root, filepath.Clean(cand)) ||
!strings.HasPrefix(filepath.Clean(p), root+string(os.PathSeparator)) {
return "", os.ErrInvalid
}
if _, e := os.Stat(p); os.IsNotExist(e) {
return p, nil
} else if e != nil {
return "", e
}
if i > 999 {
return "", os.ErrExist
}
}
}
+252
View File
@@ -0,0 +1,252 @@
package upload
import (
"bytes"
"context"
"errors"
"os"
"path/filepath"
"strings"
"testing"
"time"
"booklib/internal/ports"
)
func newU(t *testing.T) *U {
t.Helper()
dir := t.TempDir()
resolved, err := filepath.EvalSymlinks(dir)
if err != nil {
t.Fatal(err)
}
return New(resolved, 1) // 1MB cap, mirrors handler test cfg
}
func mustInit(t *testing.T, u *U, libID int64, name string, size, chunk int64) string {
t.Helper()
uid, err := u.Init(context.Background(), libID, name, size, chunk)
if err != nil {
t.Fatalf("Init(%q): %v", name, err)
}
return uid
}
func TestUploadIDForDeterministic(t *testing.T) {
a := uploadIDFor(1, "x.cbz", 100, 50)
b := uploadIDFor(1, "x.cbz", 100, 50)
c := uploadIDFor(2, "x.cbz", 100, 50)
if a != b {
t.Fatalf("same fingerprint must yield same uid: %s vs %s", a, b)
}
if a == c {
t.Fatal("different libID must yield different uid")
}
if !validUploadID(a) || validUploadID("zzz") || validUploadID(strings.Repeat("a", 31)) {
t.Fatal("validUploadID broken")
}
}
func TestInitValidation(t *testing.T) {
u := newU(t)
ctx := context.Background()
if _, err := u.Init(ctx, 1, "", 100, 0); !errors.Is(err, ErrBadName) {
t.Fatalf("empty name: want ErrBadName got %v", err)
}
if _, err := u.Init(ctx, 1, "virus.exe", 100, 0); !errors.Is(err, ErrBadFormat) {
t.Fatalf("bad ext: want ErrBadFormat got %v", err)
}
if _, err := u.Init(ctx, 1, "ok.cbz", 0, 0); !errors.Is(err, ErrBadSize) {
t.Fatalf("zero size: want ErrBadSize got %v", err)
}
if _, err := u.Init(ctx, 1, "big.cbz", 2<<20, 0); !errors.Is(err, ports.ErrTooLarge) {
t.Fatalf("oversize: want ErrTooLarge got %v", err)
}
if _, err := u.Init(ctx, 1, "ok.cbz", 100, 40<<20); !errors.Is(err, ErrBadChunk) {
t.Fatalf("huge chunk: want ErrBadChunk got %v", err)
}
// 默认 chunk 生效且 uid 稳定
uid := mustInit(t, u, 1, "ok.cbz", 100, 0)
if uid != uploadIDFor(1, "ok.cbz", 100, defaultChunk) {
t.Fatal("chunkSize=0 must default to defaultChunk in fingerprint")
}
// SafeName 清洗:path 形式取 base
if uid2 := mustInit(t, u, 1, `C:\dir\book.cbz`, 100, 0); uid2 != uploadIDFor(1, "book.cbz", 100, defaultChunk) {
t.Fatalf("windows path must be sanitized to base name, got uid %s", uid2)
}
}
func TestInitResumeAndRestart(t *testing.T) {
u := newU(t)
ctx := context.Background()
uid := mustInit(t, u, 1, "r.cbz", 900000, 400000)
if err := u.PutPart(ctx, uid, 0, bytes.NewReader(bytes.Repeat([]byte("x"), 400000)), 0); err != nil {
t.Fatalf("PutPart: %v", err)
}
// 同指纹 re-init → 复用会话,分片保留
uid2, err := u.Init(ctx, 1, "r.cbz", 900000, 400000)
if err != nil || uid2 != uid {
t.Fatalf("resume: want same uid, got %s err=%v", uid2, err)
}
recv, err := u.Status(ctx, uid)
if err != nil || len(recv) != 1 || recv[0] != 0 {
t.Fatalf("resume must keep parts: %v err=%v", recv, err)
}
// 同指纹位但内容不同(size 变)→ 新会话,旧目录被换掉是安全的(uid 不同)
uid3 := mustInit(t, u, 1, "r.cbz", 800000, 400000)
if uid3 == uid {
t.Fatal("different size must yield different uid")
}
}
func TestPutPartErrors(t *testing.T) {
u := newU(t)
ctx := context.Background()
uid := mustInit(t, u, 1, "p.cbz", 1000, 400)
if err := u.PutPart(ctx, uid, 9, bytes.NewReader(bytes.Repeat([]byte("y"), 400)), 0); !errors.Is(err, ErrBadIndex) {
t.Fatalf("index oob: want ErrBadIndex got %v", err)
}
if err := u.PutPart(ctx, uid, -1, bytes.NewReader(nil), 0); !errors.Is(err, ErrBadIndex) {
t.Fatalf("negative index: want ErrBadIndex got %v", err)
}
// 超期望体积 → ErrPartTooBig
if err := u.PutPart(ctx, uid, 0, bytes.NewReader(bytes.Repeat([]byte("y"), 500)), 0); !errors.Is(err, ErrPartTooBig) {
t.Fatalf("oversize part: want ErrPartTooBig got %v", err)
}
// 短读 → ErrPartSizeMismatch
if err := u.PutPart(ctx, uid, 0, bytes.NewReader(bytes.Repeat([]byte("y"), 300)), 0); !errors.Is(err, ErrPartSizeMismatch) {
t.Fatalf("short part: want ErrPartSizeMismatch got %v", err)
}
// 失败不留 parts/B4: 截断分片不被 Status 报告为已接收
recv, err := u.Status(ctx, uid)
if err != nil || len(recv) != 0 {
t.Fatalf("failed parts must not be received: %v err=%v", recv, err)
}
// 未知/非法 uid
if _, err := u.Status(ctx, "deadbeefdeadbeefdeadbeefdeadbeef"); !errors.Is(err, ports.ErrNotFound) {
t.Fatalf("unknown uid: want ErrNotFound got %v", err)
}
if _, err := u.Status(ctx, "zzz"); !errors.Is(err, ErrBadUploadID) {
t.Fatalf("bad uid: want ErrBadUploadID got %v", err)
}
}
func TestCompleteHappyPathAndCleanup(t *testing.T) {
u := newU(t)
ctx := context.Background()
root := filepath.Join(u.booksDir, "lib")
if err := os.MkdirAll(root, 0o755); err != nil {
t.Fatal(err)
}
content := bytes.Repeat([]byte("調教開關第二季!"), 300) // ~6.3KB multibyte
size := int64(len(content))
const chunk = int64(2000)
uid := mustInit(t, u, 7, "調教開關:第二季.zip", size, chunk)
n := int((size + chunk - 1) / chunk)
// 乱序上传:索引顺序 1,2,...,n-1,0 —— 覆盖所有分片且非递增
for step := 1; step <= n; step++ {
idx := int64(step % n)
lo := idx * chunk
hi := min(lo+chunk, size)
if err := u.PutPart(ctx, uid, idx, bytes.NewReader(content[lo:hi]), 0); err != nil {
t.Fatalf("PutPart %d: %v", idx, err)
}
}
// 缺片 → ErrIncomplete(先删一片验证)
part0 := filepath.Join(u.uploadDir(uid), "parts", "0")
saved, _ := os.ReadFile(part0)
os.Remove(part0)
if _, err := u.Complete(ctx, uid, root); !errors.Is(err, ports.ErrIncomplete) {
t.Fatalf("missing part: want ErrIncomplete got %v", err)
}
os.WriteFile(part0, saved, 0o644)
rel, err := u.Complete(ctx, uid, root)
if err != nil {
t.Fatalf("Complete: %v", err)
}
if rel != "調教開關:第二季.zip" {
t.Fatalf("rel path: %q", rel)
}
got, err := os.ReadFile(filepath.Join(root, rel))
if err != nil || !bytes.Equal(got, content) {
t.Fatalf("assembled content wrong: err=%v", err)
}
if _, err := os.Stat(u.uploadDir(uid)); !os.IsNotExist(err) {
t.Fatal("session dir must be removed after complete")
}
}
func TestCompleteUniquePathCollision(t *testing.T) {
u := newU(t)
ctx := context.Background()
root := filepath.Join(u.booksDir, "lib")
os.MkdirAll(root, 0o755)
os.WriteFile(filepath.Join(root, "dup.cbz"), []byte("existing"), 0o644)
uid := mustInit(t, u, 1, "dup.cbz", 4, 4)
if err := u.PutPart(ctx, uid, 0, bytes.NewReader([]byte("new!")), 0); err != nil {
t.Fatal(err)
}
rel, err := u.Complete(ctx, uid, root)
if err != nil {
t.Fatalf("Complete: %v", err)
}
if rel != "dup (1).cbz" {
t.Fatalf("collision must add suffix, got %q", rel)
}
}
func TestUniquePathTraversalRejected(t *testing.T) {
u := newU(t)
root := filepath.Join(u.booksDir, "lib")
if _, err := u.UniquePath(root, "../escape.cbz"); !errors.Is(err, os.ErrInvalid) {
t.Fatalf("traversal: want os.ErrInvalid got %v", err)
}
p, err := u.UniquePath(root, "ok.cbz")
if err != nil || p != filepath.Join(root, "ok.cbz") {
t.Fatalf("valid name: %q %v", p, err)
}
}
func TestSweepRemovesOnlyExpired(t *testing.T) {
u := newU(t)
ctx := context.Background()
fresh := mustInit(t, u, 1, "fresh.cbz", 10, 10)
stale := mustInit(t, u, 1, "stale.cbz", 10, 10)
// 把 stale 会话 mtime 拨到 TTL 之前
old := time.Now().Add(-uploadSessTTL - time.Hour)
dir := u.uploadDir(stale)
if err := os.Chtimes(filepath.Join(dir, "meta.json"), old, old); err != nil {
t.Fatal(err)
}
if err := os.Chtimes(dir, old, old); err != nil {
t.Fatal(err)
}
if err := u.Sweep(ctx); err != nil {
t.Fatalf("Sweep: %v", err)
}
if _, err := os.Stat(dir); !os.IsNotExist(err) {
t.Fatal("stale session must be swept")
}
if _, err := os.Stat(u.uploadDir(fresh)); err != nil {
t.Fatalf("fresh session must survive: %v", err)
}
// 空目录/不存在 base 都不报错
os.RemoveAll(filepath.Join(u.booksDir, uploadSessionIn))
if err := u.Sweep(ctx); err != nil {
t.Fatalf("Sweep on missing base: %v", err)
}
}
func TestLibraryID(t *testing.T) {
u := newU(t)
ctx := context.Background()
uid := mustInit(t, u, 42, "lib.cbz", 10, 10)
id, err := u.LibraryID(ctx, uid)
if err != nil || id != 42 {
t.Fatalf("LibraryID: %d %v", id, err)
}
if _, err := u.LibraryID(ctx, "deadbeefdeadbeefdeadbeefdeadbeef"); !errors.Is(err, ports.ErrNotFound) {
t.Fatalf("unknown uid: want ErrNotFound got %v", err)
}
}