From fbd5cd243d1c3c1630c430953c4ae4d13e03f5cb Mon Sep 17 00:00:00 2001 From: XingfenD Date: Mon, 14 Sep 2026 22:00:25 +0800 Subject: [PATCH] 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) --- backend/cmd/webui/api/router.go | 5 +- backend/cmd/webui/api/router_test.go | 2 +- backend/cmd/webui/handlers/auth_test.go | 6 +- backend/cmd/webui/handlers/handlers.go | 6 +- backend/cmd/webui/handlers/libraries.go | 29 +- backend/cmd/webui/handlers/uploads.go | 287 +++++-------------- backend/cmd/webui/main.go | 6 +- backend/internal/scanner/scanner.go | 21 +- backend/internal/upload/upload.go | 357 ++++++++++++++++++++++++ backend/internal/upload/upload_test.go | 252 +++++++++++++++++ 10 files changed, 708 insertions(+), 263 deletions(-) create mode 100644 backend/internal/upload/upload.go create mode 100644 backend/internal/upload/upload_test.go diff --git a/backend/cmd/webui/api/router.go b/backend/cmd/webui/api/router.go index 98fe02a..33e3ecb 100644 --- a/backend/cmd/webui/api/router.go +++ b/backend/cmd/webui/api/router.go @@ -10,11 +10,12 @@ import ( "booklib/internal/redispkg" "booklib/internal/scanner" "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) - h := handlers.New(cfg, st, rdb, sc) + h := handlers.New(cfg, st, rdb, sc, up) r := gin.New() if e := r.SetTrustedProxies(cfg.TrustedProxies); e != nil { panic(e) diff --git a/backend/cmd/webui/api/router_test.go b/backend/cmd/webui/api/router_test.go index 26557cb..3b3a324 100644 --- a/backend/cmd/webui/api/router_test.go +++ b/backend/cmd/webui/api/router_test.go @@ -16,7 +16,7 @@ func testCfg() *config.Config { } 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) w := httptest.NewRecorder() r.ServeHTTP(w, req) diff --git a/backend/cmd/webui/handlers/auth_test.go b/backend/cmd/webui/handlers/auth_test.go index 829df80..33d3dae 100644 --- a/backend/cmd/webui/handlers/auth_test.go +++ b/backend/cmd/webui/handlers/auth_test.go @@ -21,6 +21,7 @@ import ( "booklib/internal/redispkg" "booklib/internal/scanner" "booklib/internal/store" + "booklib/internal/upload" ) func testCfg() *config.Config { @@ -57,8 +58,9 @@ func setupAPI(t *testing.T) (*store.Store, *scanner.Scanner, http.Handler, strin cfg.BooksDir = booksDir cfg.CacheDir = t.TempDir() rdb := redispkg.New(os.Getenv("REDIS_URL")) - sc := scanner.New(st, cfg, rdb) - r := api.NewRouter(cfg, st, rdb, sc) + up := upload.New(cfg.BooksDir, cfg.UploadMaxMB) + 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 opt, e := redis.ParseURL(u); e == nil { rc := redis.NewClient(opt) diff --git a/backend/cmd/webui/handlers/handlers.go b/backend/cmd/webui/handlers/handlers.go index e5a0633..3e46aab 100644 --- a/backend/cmd/webui/handlers/handlers.go +++ b/backend/cmd/webui/handlers/handlers.go @@ -18,6 +18,7 @@ import ( "booklib/internal/redispkg" "booklib/internal/scanner" "booklib/internal/store" + "booklib/internal/upload" ) type H struct { @@ -25,10 +26,11 @@ type H struct { st *store.Store rdb *redispkg.R sc *scanner.Scanner + up *upload.U } -func New(cfg *config.Config, st *store.Store, rdb *redispkg.R, sc *scanner.Scanner) *H { - return &H{cfg: cfg, st: st, rdb: rdb, sc: sc} +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, up: up} } func err(c *gin.Context, status int, code, msg string) { diff --git a/backend/cmd/webui/handlers/libraries.go b/backend/cmd/webui/handlers/libraries.go index 6c2e6bd..e7c30c1 100644 --- a/backend/cmd/webui/handlers/libraries.go +++ b/backend/cmd/webui/handlers/libraries.go @@ -136,12 +136,12 @@ func (h *H) Upload(c *gin.Context) { return } // 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 out *os.File for attempt := 0; attempt < 5; attempt++ { var e error - dst, e = h.uniquePath(root, name) + dst, e = h.up.UniquePath(root, name) if e != nil { err(c, http.StatusForbidden, "forbidden", e.Error()) 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))}) } - -// 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 - } - } -} diff --git a/backend/cmd/webui/handlers/uploads.go b/backend/cmd/webui/handlers/uploads.go index fbc9cb1..4709a81 100644 --- a/backend/cmd/webui/handlers/uploads.go +++ b/backend/cmd/webui/handlers/uploads.go @@ -1,86 +1,22 @@ package handlers import ( - "crypto/sha256" - "encoding/hex" - "encoding/json" "errors" - "fmt" - "io" "net/http" "os" - "path/filepath" - "sort" "strconv" - "strings" - "time" "github.com/gin-gonic/gin" - "booklib/internal/bookfile" + "booklib/internal/ports" + "booklib/internal/upload" ) -// 分片上传:init(指纹→确定性 uploadId,天然支持续传)→ PUT parts → complete 拼接原子落盘。 -// 会话即 /.uploads//(meta.json + parts/N),无独立状态存储;24h 未 complete opportunistic 清扫。 +// 分片上传:HTTP 层只做参数绑定与错误映射,域逻辑(指纹续传/分片落盘/拼接/清扫) +// 全部在 internal/upload;会话清扫由 scanner ticker 接管(B16),不在请求路径。 -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())) - } - } -} +// maxChunkBytes 与 upload 包内常量同值,作为 PutPart 的防御性读取上限。 +const maxChunkBytes = 32 << 20 func (h *H) UploadInit(c *gin.Context) { 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") return } - name := bookfile.SafeName(req.Name) - 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") + uid, e := h.up.Init(c, lib.ID, req.Name, req.Size, req.ChunkSize) + if e != nil { + h.mapUploadErr(c, e) 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 { - err(c, http.StatusNotFound, "not_found", "no such upload") - return uploadMeta{}, "", false - } - if json.Unmarshal(b, &m) != nil { - err(c, http.StatusInternalServerError, "internal", "corrupt session") - return uploadMeta{}, "", false - } - return m, dir, true -} - func (h *H) UploadStatus(c *gin.Context) { - _, dir, ok := h.loadMeta(c) - if !ok { + recv, e := h.up.Status(c, c.Param("uid")) + if e != nil { + h.mapUploadErr(c, e) 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}) } func (h *H) UploadPart(c *gin.Context) { - m, dir, ok := h.loadMeta(c) - if !ok { - return - } + uid := c.Param("uid") 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") return } - lo, hi := chunkRange(m, idx) - c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, hi-lo) - // 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) + // maxSize 防御性上限:分片声明大小由会话 meta 决定,这里再垫一层 32MB 全局上限 + e = h.up.PutPart(c, uid, idx, c.Request.Body, maxChunkBytes) if e != nil { - err(c, http.StatusInternalServerError, "internal", "create part") - 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") + h.mapUploadErr(c, e) return } c.JSON(http.StatusAccepted, gin.H{"accepted": true}) } func (h *H) UploadComplete(c *gin.Context) { - m, dir, ok := h.loadMeta(c) - if !ok { + uid := c.Param("uid") + libID, e := h.up.LibraryID(c, uid) + if e != nil { + h.mapUploadErr(c, e) return } - var total int64 - 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) + lib, e := h.st.GetLibrary(c, libID) if e != nil { dbErr(c, e) return @@ -255,39 +84,53 @@ func (h *H) UploadComplete(c *gin.Context) { if !ok { return } - dst, e := h.uniquePath(root, m.Name) + rel, e := h.up.Complete(c, uid, root) if e != nil { - err(c, http.StatusForbidden, "forbidden", e.Error()) + h.mapUploadErr(c, e) return } - tmp := filepath.Join(dir, "assembled") - 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") + 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()) + 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 } - pf.Close() - } - out.Close() - if e := os.Rename(tmp, dst); e != nil { // 原子落盘,scanner 自动收编 - os.Remove(tmp) - err(c, http.StatusInternalServerError, "internal", "rename") - return + err(c, http.StatusInternalServerError, "internal", "upload failed") } - os.RemoveAll(dir) - c.JSON(http.StatusAccepted, gin.H{"accepted": true, "path": strings.TrimPrefix(dst, root+string(os.PathSeparator))}) } diff --git a/backend/cmd/webui/main.go b/backend/cmd/webui/main.go index 2a01e5e..b880708 100644 --- a/backend/cmd/webui/main.go +++ b/backend/cmd/webui/main.go @@ -16,6 +16,7 @@ import ( "booklib/internal/scanner" "booklib/internal/seed" "booklib/internal/store" + "booklib/internal/upload" ) func main() { @@ -39,11 +40,12 @@ func main() { log.Fatalf("seed: %v", err) } 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) 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} go func() { log.Printf("listening on %s", cfg.Addr) diff --git a/backend/internal/scanner/scanner.go b/backend/internal/scanner/scanner.go index 3543f9c..3703af8 100644 --- a/backend/internal/scanner/scanner.go +++ b/backend/internal/scanner/scanner.go @@ -19,17 +19,23 @@ import ( "booklib/internal/store" ) +// Sweeper 是 scanner 每轮顺手调用的清理钩子;upload.U 满足它(B16)。 +type Sweeper interface { + Sweep(ctx context.Context) error +} + type Scanner struct { - st *store.Store - cfg *config.Config - rdb *redispkg.R + st *store.Store + cfg *config.Config + rdb *redispkg.R + sweepers []Sweeper // B9-②: per-library single-flight — concurrent scan triggers for the same // library are merged into one execution, even without redis. flights sync.Map // map[int64]*sync.WaitGroup } -func New(st *store.Store, cfg *config.Config, rdb *redispkg.R) *Scanner { - return &Scanner{st: st, cfg: cfg, rdb: rdb} +func New(st *store.Store, cfg *config.Config, rdb *redispkg.R, sweepers ...Sweeper) *Scanner { + return &Scanner{st: st, cfg: cfg, rdb: rdb, sweepers: sweepers} } func (s *Scanner) Run(ctx context.Context) { @@ -40,6 +46,11 @@ func (s *Scanner) Run(ctx context.Context) { case <-ctx.Done(): return 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) if err != nil { log.Printf("scan: list libraries: %v", err) diff --git a/backend/internal/upload/upload.go b/backend/internal/upload/upload.go new file mode 100644 index 0000000..17abb21 --- /dev/null +++ b/backend/internal/upload/upload.go @@ -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 /.uploads// (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 + } + } +} diff --git a/backend/internal/upload/upload_test.go b/backend/internal/upload/upload_test.go new file mode 100644 index 0000000..e8fbe99 --- /dev/null +++ b/backend/internal/upload/upload_test.go @@ -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) + } +}