package storage import ( "bytes" "context" "crypto/sha256" "encoding/hex" "io" "os" "path/filepath" "strings" "sync" "testing" ) // newTestLocal 构造以临时目录为根的本地引擎。 func newTestLocal(t *testing.T) *LocalStorage { t.Helper() st, err := NewLocalStorage(t.TempDir()) if err != nil { t.Fatalf("NewLocalStorage: %v", err) } return st } func sha256Hex(b []byte) string { sum := sha256.Sum256(b) return hex.EncodeToString(sum[:]) } // TestLocalSaveOpenRange 保存/Stat/完整与 Range 下载。 func TestLocalSaveOpenRange(t *testing.T) { st := newTestLocal(t) ctx := context.Background() data := []byte("hello fileshare 本地引擎 0123456789") n, err := st.SaveFile(ctx, bytes.NewReader(data), "2025/08/测试文件.bin") if err != nil { t.Fatalf("SaveFile: %v", err) } if n != int64(len(data)) { t.Fatalf("SaveFile n = %d, want %d", n, len(data)) } // Stat meta, err := st.Stat(ctx, "2025/08/测试文件.bin") if err != nil { t.Fatalf("Stat: %v", err) } if meta.Size != int64(len(data)) || !meta.AcceptRanges { t.Fatalf("Stat = %+v", meta) } // 完整下载:对齐 go-api 约定 Start=0、End=Total-1 dl, err := st.Open(ctx, "2025/08/测试文件.bin", nil) if err != nil { t.Fatalf("Open: %v", err) } got, err := io.ReadAll(dl) _ = dl.Close() if err != nil { t.Fatalf("read full: %v", err) } if !bytes.Equal(got, data) { t.Fatalf("full content mismatch") } if dl.Start != 0 || dl.End != int64(len(data))-1 || dl.Total != int64(len(data)) { t.Fatalf("full download offsets = %d,%d,%d", dl.Start, dl.End, dl.Total) } // Range 下载 [2, 7] dl, err = st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: 2, End: 7}) if err != nil { t.Fatalf("Open range: %v", err) } got, err = io.ReadAll(dl) _ = dl.Close() if err != nil { t.Fatalf("read range: %v", err) } if !bytes.Equal(got, data[2:8]) { t.Fatalf("range content = %q, want %q", got, data[2:8]) } // Range end 越界自动钳制 dl, err = st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: 5, End: 99999}) if err != nil { t.Fatalf("Open clamp range: %v", err) } got, _ = io.ReadAll(dl) _ = dl.Close() if !bytes.Equal(got, data[5:]) { t.Fatalf("clamp range mismatch") } // 起点越界 → 416 if _, err := st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: int64(len(data)) + 1, End: -1}); err == nil { t.Fatalf("out-of-range start should fail") } else if !strings.Contains(err.Error(), ErrRangeNotSatisfiable.Error()) { t.Fatalf("want ErrRangeNotSatisfiable, got %v", err) } // 不存在 → 404 if _, err := st.Open(ctx, "no/such/file.bin", nil); err == nil || !strings.Contains(err.Error(), ErrNotFound.Error()) { t.Fatalf("want ErrNotFound, got %v", err) } } // TestLocalSaveFileAtomic 原子写:目录自动创建,无临时残留。 func TestLocalSaveFileAtomic(t *testing.T) { ctx := context.Background() dir := t.TempDir() st2, _ := NewLocalStorage(dir) big := bytes.Repeat([]byte("abc123"), 128*1024) // 768KB,跨多个 256KB 缓冲 if _, err := st2.SaveFile(ctx, bytes.NewReader(big), "a/b/c/big.bin"); err != nil { t.Fatalf("SaveFile: %v", err) } // 无临时残留 entries, _ := os.ReadDir(filepath.Join(dir, "a", "b", "c")) for _, e := range entries { if strings.HasPrefix(e.Name(), ".") { t.Fatalf("临时文件残留: %s", e.Name()) } } got, _ := os.ReadFile(filepath.Join(dir, "a", "b", "c", "big.bin")) if !bytes.Equal(got, big) { t.Fatalf("content mismatch") } } // TestLocalTraversal 防路径穿越(含符号链接逃逸)。 func TestLocalTraversal(t *testing.T) { st := newTestLocal(t) ctx := context.Background() for _, p := range []string{"../escape.txt", "a/../../escape", "..", "/../x"} { if _, err := st.SaveFile(ctx, strings.NewReader("x"), p); err == nil || !strings.Contains(err.Error(), ErrInvalidPath.Error()) { t.Fatalf("SaveFile(%q) 应拒绝: %v", p, err) } if _, err := st.Open(ctx, p, nil); err == nil || !strings.Contains(err.Error(), ErrInvalidPath.Error()) { t.Fatalf("Open(%q) 应拒绝: %v", p, err) } } // 符号链接逃逸 outside := filepath.Join(t.TempDir(), "outside.txt") if err := os.WriteFile(outside, []byte("secret"), 0o644); err != nil { t.Fatal(err) } link := filepath.Join(st.root, "link.txt") if err := os.Symlink(outside, link); err != nil { t.Skipf("symlink 不可用: %v", err) } if _, err := st.Open(ctx, "link.txt", nil); err == nil || !strings.Contains(err.Error(), ErrInvalidPath.Error()) { t.Fatalf("symlink escape 应拒绝: %v", err) } } // TestLocalChunkLifecycle 分片保存/合并(索引有序 + SHA256 校验 + 清理临时目录)。 func TestLocalChunkLifecycle(t *testing.T) { st := newTestLocal(t) ctx := context.Background() savePath := "2025/09/chunked.bin" uploadID := "upload-abc" chunks := [][]byte{[]byte("AAAA"), []byte("BB"), []byte("CCCCCC")} hashes := make([]string, len(chunks)) var total int64 for i, c := range chunks { n, err := st.SaveChunk(ctx, uploadID, i, bytes.NewReader(c), savePath) if err != nil { t.Fatalf("SaveChunk %d: %v", i, err) } if n != int64(len(c)) { t.Fatalf("SaveChunk %d n = %d", i, n) } hashes[i] = sha256Hex(c) total += int64(len(c)) } // 分片文件确实存在于临时目录 for i := range chunks { exists, err := st.FileExists(ctx, ChunkPartPath(savePath, uploadID, i)) if err != nil || !exists { t.Fatalf("分片 %d 应存在: %v %v", i, exists, err) } } size, fileHash, err := st.MergeChunks(ctx, uploadID, len(chunks), func(i int) (string, error) { return hashes[i], nil }, savePath) if err != nil { t.Fatalf("MergeChunks: %v", err) } if size != total { t.Fatalf("merged size = %d, want %d", size, total) } want := sha256Hex(bytes.Join(chunks, nil)) if fileHash != want { t.Fatalf("file hash = %s, want %s", fileHash, want) } // 合并内容 = 按索引有序拼接 got, err := os.ReadFile(filepath.Join(st.root, filepath.FromSlash(savePath))) if err != nil { t.Fatal(err) } if !bytes.Equal(got, bytes.Join(chunks, nil)) { t.Fatalf("merged content mismatch") } // 合并后分片目录已清理 exists, _ := st.FileExists(ctx, ChunkPartPath(savePath, uploadID, 0)) if exists { t.Fatalf("合并后分片应已清理") } if _, err := os.Stat(filepath.Join(st.root, "2025/09/chunks")); !os.IsNotExist(err) { t.Fatalf("chunks 父目录应已清理: %v", err) } } // TestLocalChunkHashMismatch 哈希不匹配 → 合并失败且不留输出文件。 func TestLocalChunkHashMismatch(t *testing.T) { st := newTestLocal(t) ctx := context.Background() savePath := "mismatch.bin" if _, err := st.SaveChunk(ctx, "uid", 0, strings.NewReader("data"), savePath); err != nil { t.Fatal(err) } _, _, err := st.MergeChunks(ctx, "uid", 1, func(i int) (string, error) { return sha256Hex([]byte("WRONG")), nil }, savePath) if err == nil || !strings.Contains(err.Error(), ErrHashMismatch.Error()) { t.Fatalf("want ErrHashMismatch, got %v", err) } exists, _ := st.FileExists(ctx, savePath) if exists { t.Fatalf("校验失败不应产出正式文件") } } // TestLocalVerifyHashAbort verifyHash 返回 error → 合并中止。 func TestLocalVerifyHashAbort(t *testing.T) { st := newTestLocal(t) ctx := context.Background() if _, err := st.SaveChunk(ctx, "uid2", 0, strings.NewReader("data"), "v.bin"); err != nil { t.Fatal(err) } boom := context.Canceled if _, _, err := st.MergeChunks(ctx, "uid2", 1, func(i int) (string, error) { return "", boom }, "v.bin"); err == nil { t.Fatalf("verifyHash 错误应向上传播") } } // TestLocalDeleteExistsClean 删除/存在性/清理。 func TestLocalDeleteExistsClean(t *testing.T) { st := newTestLocal(t) ctx := context.Background() if _, err := st.SaveFile(ctx, strings.NewReader("x"), "d/f.bin"); err != nil { t.Fatal(err) } if ok, _ := st.FileExists(ctx, "d/f.bin"); !ok { t.Fatalf("文件应存在") } if err := st.DeleteFile(ctx, "d/f.bin"); err != nil { t.Fatalf("DeleteFile: %v", err) } if ok, _ := st.FileExists(ctx, "d/f.bin"); ok { t.Fatalf("文件应已删除") } // 删除不存在 → nil if err := st.DeleteFile(ctx, "d/f.bin"); err != nil { t.Fatalf("删除不存在应静默: %v", err) } // CleanChunks 幂等 if err := st.CleanChunks(ctx, "uid-x", "y.bin"); err != nil { t.Fatalf("CleanChunks: %v", err) } } // TestLocalPresignNotSupported 预签名 → ErrNotSupported。 func TestLocalPresignNotSupported(t *testing.T) { st := newTestLocal(t) ctx := context.Background() if _, err := st.PresignGetURL(ctx, "a.bin", 60); err == nil || !strings.Contains(err.Error(), ErrNotSupported.Error()) { t.Fatalf("PresignGetURL want ErrNotSupported, got %v", err) } if _, err := st.PresignPutURL(ctx, "a.bin", 60); err == nil || !strings.Contains(err.Error(), ErrNotSupported.Error()) { t.Fatalf("PresignPutURL want ErrNotSupported, got %v", err) } } // TestLocalHealthCheck 健康检查。 func TestLocalHealthCheck(t *testing.T) { st := newTestLocal(t) if err := st.HealthCheck(context.Background()); err != nil { t.Fatalf("HealthCheck: %v", err) } } // TestLocalConcurrent 并发安全冒烟。 func TestLocalConcurrent(t *testing.T) { st := newTestLocal(t) ctx := context.Background() var wg sync.WaitGroup for i := 0; i < 16; i++ { wg.Add(1) go func(i int) { defer wg.Done() p := "conc/" + itoa(i) + ".bin" if _, err := st.SaveFile(ctx, strings.NewReader(strings.Repeat("x", i+1)), p); err != nil { t.Errorf("save %d: %v", i, err) return } dl, err := st.Open(ctx, p, nil) if err != nil { t.Errorf("open %d: %v", i, err) return } _, _ = io.ReadAll(dl) _ = dl.Close() }(i) } wg.Wait() } // TestLocalFactoryRegistry 工厂注册与构造。 func TestLocalFactoryRegistry(t *testing.T) { prev := engineOptions.Local.Root engineOptions.Local.Root = t.TempDir() defer func() { engineOptions.Local.Root = prev }() st, err := NewEngine(context.Background(), "local") if err != nil { t.Fatalf("NewEngine(local): %v", err) } if err := st.HealthCheck(context.Background()); err != nil { t.Fatalf("HealthCheck: %v", err) } if _, err := NewEngine(context.Background(), "unknown"); err == nil { t.Fatalf("未知引擎应报错") } }