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) } }