From ddc71b7cca74b0eee627bbcbf9de1167f053c1e3 Mon Sep 17 00:00:00 2001 From: kugouming Date: Thu, 10 Sep 2026 08:46:14 +0800 Subject: [PATCH] =?UTF-8?q?test:=20=E8=A1=A5=E9=BD=90=E6=A0=B8=E5=BF=83?= =?UTF-8?q?=E5=8C=85=E4=B8=8E=20cmd=20=E7=BA=AF=E5=87=BD=E6=95=B0=E5=8D=95?= =?UTF-8?q?=E6=B5=8B,=E6=96=B0=E5=A2=9E=E6=97=A0=E4=BE=9D=E8=B5=96=20test?= =?UTF-8?q?=20CI?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 新增 13 个测试文件、扩充 3 个,测试代码 794 → 2387 行,零外部依赖 (纯函数 + 接口 mock + t.TempDir),为 go test ./... 建立可信回归网: - internal/s3path 达 100% 覆盖;config 补 Resolve 全分支(默认 profile 回退/env 展开与覆盖/path-style 强制 true/缺 key 报错) - internal/view 覆盖格式探测、suffix Range、内容嗅探不丢字节、readBounded - internal/uploader 覆盖上传路径(s3 字段改 manager.UploadAPIClient 接口 以便 mock,抽 buildKey 纯函数),连同 client.xmlTimeNormalizer(aws.HTTPClient stub)覆盖全部规范化分支 - cmd 覆盖 syncNeed 决策、matchFind 过滤、deriveDstKey、buildTree、 headLinesFromReader、parseS3、humanBytes 等 修复 uploader 中 partSize 常量未被引用等既有警告问题一并纳入。 新增 .github/workflows/test.yml:每次 push/PR 自动 go vet + go test(stable/ oldstable),此前唯一 CI 仅 release.yml 发布用、不跑任何测试。 --- .github/workflows/test.yml | 30 +++ cmd/checksum_test.go | 20 ++ cmd/cp_test.go | 30 +++ cmd/du_test.go | 19 ++ cmd/find_test.go | 110 +++++++++++ cmd/head_test.go | 52 +++++ cmd/ls_test.go | 32 +++ cmd/root_test.go | 76 ++++++++ cmd/stat_test.go | 25 +++ cmd/sync_more_test.go | 165 ++++++++++++++++ cmd/tree_test.go | 78 ++++++++ cmd/wc_test.go | 18 ++ internal/client/client_test.go | 145 ++++++++++++++ internal/config/config_test.go | 224 +++++++++++++++++++++ internal/s3path/s3path_test.go | 109 +++++++++++ internal/uploader/uploader.go | 22 ++- internal/uploader/uploader_test.go | 187 ++++++++++++++++++ internal/view/view_test.go | 303 +++++++++++++++++++++++++++++ 18 files changed, 1638 insertions(+), 7 deletions(-) create mode 100644 .github/workflows/test.yml create mode 100644 cmd/checksum_test.go create mode 100644 cmd/cp_test.go create mode 100644 cmd/du_test.go create mode 100644 cmd/head_test.go create mode 100644 cmd/ls_test.go create mode 100644 cmd/root_test.go create mode 100644 cmd/stat_test.go create mode 100644 cmd/sync_more_test.go create mode 100644 cmd/tree_test.go create mode 100644 cmd/wc_test.go create mode 100644 internal/client/client_test.go create mode 100644 internal/s3path/s3path_test.go create mode 100644 internal/uploader/uploader_test.go create mode 100644 internal/view/view_test.go diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml new file mode 100644 index 0000000..e616fa4 --- /dev/null +++ b/.github/workflows/test.yml @@ -0,0 +1,30 @@ +name: test + +# 每次 push / PR 自动跑静态检查与单元测试,保证回归网始终被执行。 +# 单测零外部依赖(纯函数 + mock + t.TempDir),无需凭证与容器。 +on: + push: + branches: ['*'] + pull_request: + +permissions: + contents: read + +jobs: + test: + runs-on: ubuntu-latest + strategy: + matrix: + go-version: [stable, oldstable] + steps: + - uses: actions/checkout@v7 + + - uses: actions/setup-go@v7 + with: + go-version: ${{ matrix.go-version }} + + - name: Vet + run: go vet ./... + + - name: Unit tests + run: go test ./... diff --git a/cmd/checksum_test.go b/cmd/checksum_test.go new file mode 100644 index 0000000..a267b31 --- /dev/null +++ b/cmd/checksum_test.go @@ -0,0 +1,20 @@ +package cmd + +import "testing" + +func TestNewHasher(t *testing.T) { + md5h, err := newHasher("md5") + if err != nil || md5h == nil { + t.Errorf("md5 应可用,err=%v", err) + } + sha, err := newHasher("sha256") + if err != nil || sha == nil { + t.Errorf("sha256 应可用,err=%v", err) + } + if _, err := newHasher("crc32"); err == nil { + t.Errorf("非法算法应报错") + } + if _, err := newHasher(""); err == nil { + t.Errorf("空算法应报错") + } +} diff --git a/cmd/cp_test.go b/cmd/cp_test.go new file mode 100644 index 0000000..07bbf84 --- /dev/null +++ b/cmd/cp_test.go @@ -0,0 +1,30 @@ +package cmd + +import ( + "testing" + + "github.com/BeCrafter/sail/internal/s3path" +) + +// TestDeriveDstKey 校验 cp/mv 目标 S3 key 推导: +// dst.Key 为空或尾 / 表示"进目录",追加 srcBase;否则用 dst.Key。 +func TestDeriveDstKey(t *testing.T) { + cases := []struct { + name string + srcBase string + dst s3path.S3Path + want string + }{ + {name: "dst 空 key 进目录", srcBase: "a.txt", dst: s3path.S3Path{Bucket: "b"}, want: "a.txt"}, + {name: "dst 尾斜杠进目录", srcBase: "a.txt", dst: s3path.S3Path{Bucket: "b", Key: "dir/"}, want: "dir/a.txt"}, + {name: "dst 指定完整 key", srcBase: "a.txt", dst: s3path.S3Path{Bucket: "b", Key: "x/y.txt"}, want: "x/y.txt"}, + {name: "dst 尾斜杠进子目录", srcBase: "sub/b.log", dst: s3path.S3Path{Bucket: "b", Key: "logs/"}, want: "logs/sub/b.log"}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + if got := deriveDstKey(c.srcBase, &c.dst); got != c.want { + t.Errorf("deriveDstKey(%q,%v) = %q,期望 %q", c.srcBase, c.dst, got, c.want) + } + }) + } +} diff --git a/cmd/du_test.go b/cmd/du_test.go new file mode 100644 index 0000000..f60d79c --- /dev/null +++ b/cmd/du_test.go @@ -0,0 +1,19 @@ +package cmd + +import "testing" + +func TestDisplayRoot(t *testing.T) { + cases := []struct { + bucket, base string + want string + }{ + {"bucket", "", "s3://bucket"}, + {"bucket", "logs", "s3://bucket/logs"}, + {"bucket", "logs/2026", "s3://bucket/logs/2026"}, + } + for _, c := range cases { + if got := displayRoot(c.bucket, c.base); got != c.want { + t.Errorf("displayRoot(%q,%q) = %q,期望 %q", c.bucket, c.base, got, c.want) + } + } +} diff --git a/cmd/find_test.go b/cmd/find_test.go index c64a338..3c8cff1 100644 --- a/cmd/find_test.go +++ b/cmd/find_test.go @@ -3,6 +3,8 @@ package cmd import ( "testing" "time" + + "github.com/aws/aws-sdk-go-v2/service/s3/types" ) func TestParseSizeSpec(t *testing.T) { @@ -67,3 +69,111 @@ func TestParseTimeArg(t *testing.T) { t.Errorf("parseTimeArg(2026/01/02) 期望报错") } } + +func strPtr(s string) *string { return &s } + +func sizePtr(n int64) *int64 { return &n } + +func timePtr(t time.Time) *time.Time { return &t } + +// findObj 构造一个 types.Object 用于 matchFind。 +func findObj(size int64, mod time.Time) types.Object { + o := types.Object{Key: strPtr("unused")} + if size >= 0 { + o.Size = sizePtr(size) + } + if !mod.IsZero() { + o.LastModified = timePtr(mod) + } + return o +} + +// TestMatchFind 校验 find 四类过滤的组合匹配。 +func TestMatchFind(t *testing.T) { + base := time.Date(2026, 6, 1, 0, 0, 0, 0, time.UTC) + old := findObj(100, base.Add(-24*time.Hour)) // 对象本身尺寸/时间 + new := findObj(500, base.Add(24*time.Hour)) + + reset := func() { + findMaxDepth = 0 + findNames = nil + findSize = "" + findNewer = "" + } + reset() + defer reset() + + t.Run("无过滤全命中", func(t *testing.T) { + if !matchFind(old, "logs/a.log", "", nil, time.Time{}) { + t.Errorf("无过滤应命中") + } + }) + + t.Run("name 通配按 basename", func(t *testing.T) { + findNames = []string{"*.log"} + defer reset() + if !matchFind(old, "logs/a.log", "", nil, time.Time{}) { + t.Errorf("*.log 应命中 a.log") + } + if matchFind(old, "logs/a.txt", "", nil, time.Time{}) { + t.Errorf("*.log 不应命中 a.txt") + } + }) + + t.Run("size 精确/大于/小于", func(t *testing.T) { + cases := []struct { + spec string + obj types.Object + want bool + }{ + {"100", old, true}, + {"101", old, false}, + {"+200", new, true}, + {"+200", old, false}, + {"-200", old, true}, + {"-200", new, false}, + } + for _, c := range cases { + spec, err := parseSizeSpec(c.spec) + if err != nil { + t.Fatalf("parseSizeSpec(%q): %v", c.spec, err) + } + if got := matchFind(c.obj, "a.log", "", spec, time.Time{}); got != c.want { + t.Errorf("size %q 对 obj(size=%d) = %v,期望 %v", c.spec, *c.obj.Size, got, c.want) + } + } + }) + + t.Run("size 过滤中 Size 为 nil 视为 0", func(t *testing.T) { + noSize := types.Object{Key: strPtr("x")} + if !matchFind(noSize, "x", "", []int64{'=', 0}, time.Time{}) { + t.Errorf("Size nil 应视为 0,=0 应命中") + } + }) + + t.Run("newer 过滤", func(t *testing.T) { + if !matchFind(new, "a.log", "", nil, base) { + t.Errorf("对象晚于 newer 应命中") + } + if matchFind(old, "a.log", "", nil, base) { + t.Errorf("对象早于 newer 应被拒") + } + noMod := types.Object{Key: strPtr("x")} // LastModified nil + if matchFind(noMod, "x", "", nil, base) { + t.Errorf("LastModified nil 且 newer 非零应被拒") + } + }) + + t.Run("max-depth 过滤", func(t *testing.T) { + findMaxDepth = 2 + defer reset() + // base="logs",rel 层级数 + obj := findObj(1, base) + if matchFind(obj, "logs/a/b/c.txt", "logs", nil, time.Time{}) { + t.Errorf("深度 ≥2 应被拒") + } + if !matchFind(obj, "logs/a.txt", "logs", nil, time.Time{}) { + t.Errorf("深度 0 应命中") + } + }) +} diff --git a/cmd/head_test.go b/cmd/head_test.go new file mode 100644 index 0000000..a99fea6 --- /dev/null +++ b/cmd/head_test.go @@ -0,0 +1,52 @@ +package cmd + +import ( + "io" + "os" + "strings" + "testing" +) + +// captureStdout 在 fn 执行期间重定向 os.Stdout,返回其输出。 +func captureStdout(t *testing.T, fn func()) string { + t.Helper() + old := os.Stdout + r, w, err := os.Pipe() + if err != nil { + t.Fatalf("os.Pipe: %v", err) + } + os.Stdout = w + defer func() { os.Stdout = old }() + fn() + w.Close() + out, _ := io.ReadAll(r) + return string(out) +} + +func TestHeadLinesFromReader(t *testing.T) { + cases := []struct { + name string + in string + n int64 + want string + }{ + {name: "n=0 无输出", in: "a\nb\n", n: 0, want: ""}, + {name: "前 2 行", in: "1\n2\n3\n", n: 2, want: "1\n2\n"}, + {name: "n 超实际行数", in: "1\n2\n", n: 10, want: "1\n2\n"}, + {name: "末行无换行也输出", in: "1\n2", n: 10, want: "1\n2"}, + {name: "空输入", in: "", n: 10, want: ""}, + {name: "恰一行无换行 n=1", in: "solo", n: 1, want: "solo"}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + got := captureStdout(t, func() { + if err := headLinesFromReader(strings.NewReader(c.in), c.n); err != nil { + t.Errorf("headLinesFromReader 报错: %v", err) + } + }) + if got != c.want { + t.Errorf("headLinesFromReader(%q,%d) = %q,期望 %q", c.in, c.n, got, c.want) + } + }) + } +} diff --git a/cmd/ls_test.go b/cmd/ls_test.go new file mode 100644 index 0000000..9257853 --- /dev/null +++ b/cmd/ls_test.go @@ -0,0 +1,32 @@ +package cmd + +import ( + "testing" + "time" + + "github.com/aws/aws-sdk-go-v2/service/s3/types" +) + +func TestObjSize(t *testing.T) { + var o types.Object + if got := objSize(o); got != 0 { + t.Errorf("Size nil 应返回 0,got %d", got) + } + n := int64(42) + o.Size = &n + if got := objSize(o); got != 42 { + t.Errorf("Size=42 应返回 42,got %d", got) + } +} + +func TestObjTime(t *testing.T) { + var o types.Object + if got := objTime(o); !got.IsZero() { + t.Errorf("LastModified nil 应返回零值,got %v", got) + } + ts := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC) + o.LastModified = &ts + if got := objTime(o); !got.Equal(ts) { + t.Errorf("objTime = %v,期望 %v", got, ts) + } +} diff --git a/cmd/root_test.go b/cmd/root_test.go new file mode 100644 index 0000000..64c99c5 --- /dev/null +++ b/cmd/root_test.go @@ -0,0 +1,76 @@ +package cmd + +import ( + "testing" + + "github.com/BeCrafter/sail/internal/config" +) + +func TestParseS3DefaultBucket(t *testing.T) { + r := &config.Resolved{Bucket: "defbucket"} + + // 空桶段填充默认桶 + p, err := parseS3("s3:///logs/a.txt", r) + if err != nil { + t.Fatalf("parseS3 报错: %v", err) + } + if p.Bucket != "defbucket" || p.Key != "logs/a.txt" { + t.Errorf("空桶段应填默认桶,got %+v", p) + } + + // 显式桶不覆盖 + p, err = parseS3("s3://real/a.txt", r) + if err != nil { + t.Fatalf("parseS3 报错: %v", err) + } + if p.Bucket != "real" { + t.Errorf("显式桶被覆盖,got %q", p.Bucket) + } + + // r 为 nil 且空桶段 -> 报错 + if _, err := parseS3("s3:///a.txt", nil); err == nil { + t.Errorf("无默认桶时应报错") + } + + // r.Bucket 为空且空桶段 -> 报错 + if _, err := parseS3("s3:///a.txt", &config.Resolved{}); err == nil { + t.Errorf("默认桶为空时应报错") + } +} + +func TestLangFlagFromArgs(t *testing.T) { + cases := []struct { + args []string + want string + }{ + {[]string{"--lang", "zh", "ls"}, "zh"}, + {[]string{"--lang=zh"}, "zh"}, + {[]string{"ls", "--lang", "en"}, "en"}, + {[]string{"ls"}, ""}, + {[]string{"--lang"}, ""}, // 缺值 + {[]string{"x", "--lang", "zh", "--help"}, "zh"}, + } + for _, c := range cases { + if got := langFlagFromArgs(c.args); got != c.want { + t.Errorf("langFlagFromArgs(%v) = %q,期望 %q", c.args, got, c.want) + } + } +} + +func TestConfigFlagFromArgs(t *testing.T) { + cases := []struct { + args []string + want string + }{ + {[]string{"-c", "/tmp/x.yaml", "ls"}, "/tmp/x.yaml"}, + {[]string{"--config", "/a/b.yaml"}, "/a/b.yaml"}, + {[]string{"--config=/a/b.yaml"}, "/a/b.yaml"}, + {[]string{"ls"}, ""}, + {[]string{"--config"}, ""}, + } + for _, c := range cases { + if got := configFlagFromArgs(c.args); got != c.want { + t.Errorf("configFlagFromArgs(%v) = %q,期望 %q", c.args, got, c.want) + } + } +} diff --git a/cmd/stat_test.go b/cmd/stat_test.go new file mode 100644 index 0000000..a4e6c7e --- /dev/null +++ b/cmd/stat_test.go @@ -0,0 +1,25 @@ +package cmd + +import "testing" + +// TestHumanBytes 覆盖 stat/ls/du/tree 共用的字节格式化(注意与 internal 包内 +// 同名函数是各自独立实现;此处针对 cmd/humanBytes)。 +func TestHumanBytes(t *testing.T) { + cases := []struct { + in int64 + want string + }{ + {0, "0 B"}, + {1023, "1023 B"}, + {1024, "1.0 KiB"}, + {1536, "1.5 KiB"}, + {5 * 1024 * 1024, "5.0 MiB"}, + {2 << 30, "2.0 GiB"}, + {1024 * 1024 * 1024 * 1024, "1.0 TiB"}, + } + for _, c := range cases { + if got := humanBytes(c.in); got != c.want { + t.Errorf("humanBytes(%d) = %q,期望 %q", c.in, got, c.want) + } + } +} diff --git a/cmd/sync_more_test.go b/cmd/sync_more_test.go new file mode 100644 index 0000000..a151d95 --- /dev/null +++ b/cmd/sync_more_test.go @@ -0,0 +1,165 @@ +package cmd + +import ( + "context" + "os" + "path/filepath" + "testing" + "time" +) + +// TestSyncNeed 覆盖同步决策核心:目标缺失 / 大小不同 / update / checksum 幂等 / mtime 容差。 +// 非 checksum 路径不触网络,s3c 传 nil 即可。 +func TestSyncNeed(t *testing.T) { + ctx := context.Background() + base := time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + se := syncEntry{size: 100, mtime: base} + src := &syncPath{isS3: true, bucket: "b", keyBase: "src"} + + reset := func() { + syncUpdate = false + syncChecksum = false + } + reset() + defer reset() + + // 目标不存在 -> 需要传输 + if need, err := syncNeed(ctx, nil, nil, src, src, "a.txt", se, syncEntry{}, false); err != nil || !need { + t.Errorf("目标缺失应需要传输,got need=%v err=%v", need, err) + } + + // 大小不同(默认):需要传输 + de := syncEntry{size: 200, mtime: base} + if need, _ := syncNeed(ctx, nil, nil, src, src, "a.txt", se, de, true); !need { + t.Errorf("大小不同应需要传输") + } + + // 大小不同 + --update:源较新(>1s)才传输 + syncUpdate = true + newer := syncEntry{size: 200, mtime: base.Add(5 * time.Second)} // 目标更新 + older := syncEntry{size: 200, mtime: base.Add(-5 * time.Second)} // 目标更旧 + if need, _ := syncNeed(ctx, nil, nil, src, src, "a.txt", se, newer, true); need { + t.Errorf("--update 目标较新应跳过") + } + if need, _ := syncNeed(ctx, nil, nil, src, src, "a.txt", se, older, true); !need { + t.Errorf("--update 源较新应传输") + } + syncUpdate = false + + // 大小相同 + 默认:mtime 差 >1s 才传输(容差吸收 S3 秒级截断) + sameSize := syncEntry{size: 100, mtime: base} + if need, _ := syncNeed(ctx, nil, nil, src, src, "a.txt", se, sameSize, true); need { + t.Errorf("mtime 相同应跳过(幂等)") + } + tinyDiff := syncEntry{size: 100, mtime: base.Add(500 * time.Millisecond)} + if need, _ := syncNeed(ctx, nil, nil, src, src, "a.txt", se, tinyDiff, true); need { + t.Errorf("mtime 差 ≤1s 应视为已同步而跳过") + } + bigDiff := syncEntry{size: 100, mtime: base.Add(-5 * time.Second)} + if need, _ := syncNeed(ctx, nil, nil, src, src, "a.txt", se, bigDiff, true); !need { + t.Errorf("源较目标新 >1s 应传输") + } + + // --checksum 大小相同且 mtime 差 ≤1s:幂等快路径直接跳过(不触网络) + syncChecksum = true + if need, _ := syncNeed(ctx, nil, nil, src, src, "a.txt", se, sameSize, true); need { + t.Errorf("checksum 且 mtime 差 ≤1s 应跳过(快路径)") + } + syncChecksum = false +} + +// TestChecksumDifferFastPath 覆盖 checksum 快路径:两侧 ETag 均为单分片直接比较。 +func TestChecksumDifferFastPath(t *testing.T) { + ctx := context.Background() + src := &syncPath{isS3: true, bucket: "a", keyBase: "s"} + dst := &syncPath{isS3: true, bucket: "a", keyBase: "d"} + + // 相同 ETag -> 无差异 + diff, err := checksumDiffer(ctx, nil, nil, src, dst, "x", + syncEntry{etag: "df31ab9d4881a1a91ab8be84e6186d6a"}, + syncEntry{etag: "df31ab9d4881a1a91ab8be84e6186d6a"}) + if err != nil || diff { + t.Errorf("相同 ETag 应无差异,got diff=%v err=%v", diff, err) + } + + // 不同 ETag -> 有差异 + diff, err = checksumDiffer(ctx, nil, nil, src, dst, "x", + syncEntry{etag: "df31ab9d4881a1a91ab8be84e6186d6a"}, + syncEntry{etag: "11111111111111111111111111111111"}) + if err != nil || !diff { + t.Errorf("不同 ETag 应有差异,got diff=%v err=%v", diff, err) + } + + // ETag 大小写不敏感 + diff, err = checksumDiffer(ctx, nil, nil, src, dst, "x", + syncEntry{etag: "DF31AB9D4881A1A91AB8BE84E6186D6A"}, + syncEntry{etag: "df31ab9d4881a1a91ab8be84e6186d6a"}) + if err != nil || diff { + t.Errorf("ETag 应大小写不敏感比较") + } +} + +func TestNewSyncPathLocal(t *testing.T) { + sp, err := newSyncPath(nil, "./mydir") + if err != nil { + t.Fatalf("newSyncPath 报错: %v", err) + } + if sp.isS3 || sp.localDir != "./mydir" { + t.Errorf("本地分支解析错误: %+v", sp) + } +} + +func TestSyncPathURI(t *testing.T) { + // s3 分支 + s3sp := &syncPath{isS3: true, bucket: "bucket", keyBase: "base"} + if got := s3sp.uri("a/b.txt"); got != "s3://bucket/base/a/b.txt" { + t.Errorf("s3 uri = %q,期望 s3://bucket/base/a/b.txt", got) + } + // keyBase 空 + root := &syncPath{isS3: true, bucket: "bucket"} + if got := root.uri("a.txt"); got != "s3://bucket/a.txt" { + t.Errorf("空 keyBase uri = %q", got) + } + // 本地分支 + local := &syncPath{localDir: "/tmp/d"} + if got := local.uri("x/y.txt"); got != filepath.Join("/tmp/d", "x", "y.txt") { + t.Errorf("local uri = %q", got) + } + if got := local.display("x"); got != local.uri("x") { + t.Errorf("display 应等于 uri") + } +} + +// TestIndexLocal 校验本地目录索引:relKey 用 / 分隔、目录跳过。 +func TestIndexLocal(t *testing.T) { + dir := t.TempDir() + write := func(rel, content string) { + p := filepath.Join(dir, filepath.FromSlash(rel)) + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(p, []byte(content), 0o600); err != nil { + t.Fatal(err) + } + } + write("a/b.txt", "hello") + write("c.log", "x") + + entries, err := indexLocal(dir) + if err != nil { + t.Fatalf("indexLocal 报错: %v", err) + } + if len(entries) != 2 { + t.Fatalf("索引条目 = %d,期望 2", len(entries)) + } + e, ok := entries["a/b.txt"] + if !ok { + t.Fatalf("缺少 a/b.txt,got %v", entries) + } + if e.size != 5 { + t.Errorf("a/b.txt size = %d,期望 5", e.size) + } + if _, ok := entries["c.log"]; !ok { + t.Errorf("缺少 c.log") + } +} diff --git a/cmd/tree_test.go b/cmd/tree_test.go new file mode 100644 index 0000000..c0bb3bd --- /dev/null +++ b/cmd/tree_test.go @@ -0,0 +1,78 @@ +package cmd + +import ( + "testing" +) + +// TestBuildTree 校验目录树构建:隐式中间目录、叶子写 size/isDir、空段过滤。 +func TestBuildTree(t *testing.T) { + entries := []tentry{ + {path: "a/b/c.txt", size: 3, isDir: false}, + {path: "a/d.log", size: 1, isDir: false}, + {path: "top.log", size: 5, isDir: false}, + {path: "a/", size: 0, isDir: true}, // 显式目录占位 + } + root := buildTree(entries) + + a, ok := root.children["a"] + if !ok { + t.Fatalf("缺隐式目录 a") + } + if !a.isDir { + t.Errorf("a 应为目录") + } + if b, ok := a.children["b"]; !ok { + t.Errorf("缺隐式目录 a/b") + } else if !b.isDir { + t.Errorf("a/b 应为目录") + } else if c := b.children["c.txt"]; c == nil || c.size != 3 || c.isDir { + t.Errorf("a/b/c.txt 应为叶子 size=3,got %+v", c) + } + if d := a.children["d.log"]; d == nil || d.size != 1 { + t.Errorf("a/d.log 叶子错误") + } + if top := root.children["top.log"]; top == nil || top.size != 5 { + t.Errorf("top.log 叶子错误") + } +} + +// TestVisibleKids 校验排序与 -d 过滤。 +func TestVisibleKids(t *testing.T) { + root := buildTree([]tentry{ + {path: "z.txt", size: 1}, + {path: "a.txt", size: 1}, + {path: "dir/", size: 0, isDir: true}, + }) + + // 不过滤:按名排序 + treeDirsOnly = false + defer func() { treeDirsOnly = false }() + kids := visibleKids(root) + names := make([]string, len(kids)) + for i, k := range kids { + names[i] = k.name + } + want := []string{"a.txt", "dir", "z.txt"} + if !equalStrings(names, want) { + t.Errorf("visibleKids 排序 = %v,期望 %v", names, want) + } + + // -d:只留目录 + treeDirsOnly = true + kids = visibleKids(root) + if len(kids) != 1 || kids[0].name != "dir" { + t.Errorf("-d 过滤后应只剩 dir,got %+v", kids) + } +} + +func equalStrings(a, b []string) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} diff --git a/cmd/wc_test.go b/cmd/wc_test.go new file mode 100644 index 0000000..6c8b5b1 --- /dev/null +++ b/cmd/wc_test.go @@ -0,0 +1,18 @@ +package cmd + +import "testing" + +func TestIsSpace(t *testing.T) { + cases := []struct { + ch byte + want bool + }{ + {' ', true}, {'\t', true}, {'\n', true}, {'\r', true}, {'\v', true}, {'\f', true}, + {'a', false}, {'0', false}, {255, false}, + } + for _, c := range cases { + if got := isSpace(c.ch); got != c.want { + t.Errorf("isSpace(%q) = %v,期望 %v", c.ch, got, c.want) + } + } +} diff --git a/internal/client/client_test.go b/internal/client/client_test.go new file mode 100644 index 0000000..88e0c21 --- /dev/null +++ b/internal/client/client_test.go @@ -0,0 +1,145 @@ +package client + +import ( + "errors" + "io" + "net/http" + "strings" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" +) + +// stubHTTPClient 实现 aws.HTTPClient,作为 xmlTimeNormalizer 的可控 base。 +type stubHTTPClient struct { + status int + header http.Header + body string + err error +} + +func (s *stubHTTPClient) Do(req *http.Request) (*http.Response, error) { + if s.err != nil { + return nil, s.err + } + return &http.Response{ + StatusCode: s.status, + Header: s.header, + Body: io.NopCloser(strings.NewReader(s.body)), + }, nil +} + +func TestXMLTimeNormalizer(t *testing.T) { + broken := `b2022-08-12 16:15:54` + fixed := `b2022-08-12T16:15:54Z` + + t.Run("非 200 不改写", func(t *testing.T) { + n := &xmlTimeNormalizer{base: &stubHTTPClient{status: 404, header: func() http.Header { + h := http.Header{} + h.Set("Content-Type", "application/xml") + return h + }(), body: broken}} + r, err := n.Do(nil) // stub 忽略 req,直接返回配置的响应 + if err != nil { + t.Fatalf("Do 报错: %v", err) + } + if r.StatusCode != 404 { + t.Errorf("非 200 状态码被改动") + } + body, _ := io.ReadAll(r.Body) + if string(body) != broken { + t.Errorf("非 200 响应体被改写") + } + }) + + t.Run("非 XML Content-Type 不改写", func(t *testing.T) { + n := &xmlTimeNormalizer{base: &stubHTTPClient{status: 200, header: func() http.Header { + h := http.Header{} + h.Set("Content-Type", "application/json") + return h + }(), body: broken}} + r, err := n.Do(nil) + if err != nil { + t.Fatalf("Do 报错: %v", err) + } + body, _ := io.ReadAll(r.Body) + if string(body) != broken { + t.Errorf("非 XML 响应体不应被改写") + } + }) + + t.Run("XML 且 200 规范化时间并更新 ContentLength", func(t *testing.T) { + n := &xmlTimeNormalizer{base: &stubHTTPClient{status: 200, header: func() http.Header { + h := http.Header{} + h.Set("Content-Type", "application/xml") + return h + }(), body: broken}} + r, err := n.Do(nil) + if err != nil { + t.Fatalf("Do 报错: %v", err) + } + body, _ := io.ReadAll(r.Body) + if string(body) != fixed { + t.Errorf("XML 时间未规范化:\n%s", string(body)) + } + if r.ContentLength != int64(len(fixed)) { + t.Errorf("ContentLength = %d,期望 %d", r.ContentLength, len(fixed)) + } + }) + + t.Run("空 Body 原样返回", func(t *testing.T) { + n := &xmlTimeNormalizer{base: &stubHTTPClient{status: 200, header: func() http.Header { + h := http.Header{} + h.Set("Content-Type", "application/xml") + return h + }(), body: ""}} + r, err := n.Do(nil) + if err != nil { + t.Fatalf("Do 报错: %v", err) + } + body, _ := io.ReadAll(r.Body) + if len(body) != 0 { + t.Errorf("空 Body 被改动为 %q", string(body)) + } + }) + + t.Run("base 报错透传", func(t *testing.T) { + baseErr := errors.New("boom") + n := &xmlTimeNormalizer{base: &stubHTTPClient{err: baseErr}} + if _, err := n.Do(nil); err != baseErr { + t.Errorf("base 错误未透传,got %v", err) + } + }) +} + +func TestBrokenXMLTimeRe(t *testing.T) { + got := brokenXMLTimeRe.ReplaceAllString("2022-08-12 16:15:54", "${1}T${2}Z") + if got != "2022-08-12T16:15:54Z" { + t.Errorf("正则替换 = %q,期望 2022-08-12T16:15:54Z", got) + } + // 不匹配已标准化的 ISO 时间 + iso := "2022-08-12T16:15:54Z" + if out := brokenXMLTimeRe.ReplaceAllString(iso, "${1}T${2}Z"); out != iso { + t.Errorf("已标准化的时间不应被误改: %q", out) + } +} + +// TestCoalesce 覆盖兜底选择逻辑(当前无调用点,保留以固化行为;若删除函数则删除本测试)。 +func TestCoalesce(t *testing.T) { + cases := []struct { + a, b string + want string + }{ + {"x", "y", "x"}, + {"", "y", "y"}, + {"x", "", "x"}, + {"", "", ""}, + } + for _, c := range cases { + if got := coalesce(c.a, c.b); got != c.want { + t.Errorf("coalesce(%q,%q) = %q,期望 %q", c.a, c.b, got, c.want) + } + } +} + +var _ aws.HTTPClient = (*stubHTTPClient)(nil) diff --git a/internal/config/config_test.go b/internal/config/config_test.go index eae4bec..3bf27ae 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -106,3 +106,227 @@ profiles: t.Errorf("cdn-bucket-path: false 应为 *false,got %v", r.CDNBucketPath) } } + +// TestResolveDefaultProfile 校验 profile 为空时的回退链:default-profile → "prod"。 +func TestResolveDefaultProfile(t *testing.T) { + // 配置了 default-profile:用默认 + cfg, err := Load(writeConfig(t, `default-profile: test +profiles: + prod: + endpoint: https://p.example.com + access-key: ak + secret-key: sk + test: + endpoint: https://t.example.com + access-key: ak + secret-key: sk +`)) + if err != nil { + t.Fatalf("Load: %v", err) + } + r, err := cfg.Resolve("") // 空 profile → default-profile + if err != nil { + t.Fatalf("Resolve: %v", err) + } + if r.ProfileName != "test" || r.Endpoint != "https://t.example.com" { + t.Errorf("应回退到 default-profile test,got %q %q", r.ProfileName, r.Endpoint) + } + + // 无 default-profile:回退 prod + cfg, err = Load(writeConfig(t, `profiles: + prod: + endpoint: https://p.example.com + access-key: ak + secret-key: sk +`)) + if err != nil { + t.Fatalf("Load: %v", err) + } + r, err = cfg.Resolve("") + if err != nil { + t.Fatalf("Resolve: %v", err) + } + if r.ProfileName != "prod" { + t.Errorf("无 default 应回退 prod,got %q", r.ProfileName) + } +} + +func TestResolveProfileNotFound(t *testing.T) { + cfg, err := Load(writeConfig(t, `profiles: + prod: + endpoint: https://p.example.com + access-key: ak + secret-key: sk +`)) + if err != nil { + t.Fatalf("Load: %v", err) + } + if _, err := cfg.Resolve("ghost"); err == nil { + t.Errorf("Resolve(ghost) 期望报错") + } +} + +func TestResolveMissingKeys(t *testing.T) { + // 缺 endpoint + cfg, err := Load(writeConfig(t, `profiles: + prod: + access-key: ak + secret-key: sk +`)) + if err != nil { + t.Fatalf("Load: %v", err) + } + if _, err := cfg.Resolve("prod"); err == nil { + t.Errorf("缺 endpoint 应报错") + } + + // 缺 access-key/secret-key + cfg, err = Load(writeConfig(t, `profiles: + prod: + endpoint: https://p.example.com +`)) + if err != nil { + t.Fatalf("Load: %v", err) + } + if _, err := cfg.Resolve("prod"); err == nil { + t.Errorf("缺密钥应报错") + } +} + +// TestResolvePathStyleDefaultTrue 校验 path-style 缺省及显式 false 都被强制为 true。 +func TestResolvePathStyleDefaultTrue(t *testing.T) { + // 缺省:应为 true + r := resolveFrom(t, `default-profile: prod +profiles: + prod: + endpoint: https://p.example.com + access-key: ak + secret-key: sk +`) + if !r.PathStyle { + t.Errorf("缺省 path-style 应为 true") + } + + // 显式 false:仍强制为 true + r = resolveFrom(t, `default-profile: prod +profiles: + prod: + endpoint: https://p.example.com + access-key: ak + secret-key: sk + path-style: false +`) + if !r.PathStyle { + t.Errorf("path-style: false 应被强制为 true(自建 S3 兼容服务默认)") + } + + // 显式 true:保持 true + r = resolveFrom(t, `default-profile: prod +profiles: + prod: + endpoint: https://p.example.com + access-key: ak + secret-key: sk + path-style: true +`) + if !r.PathStyle { + t.Errorf("path-style: true 应保持 true") + } +} + +// TestResolveExpandEnv 校验 ${VAR} 在配置值中被展开;未设置则置空。 +func TestResolveExpandEnv(t *testing.T) { + t.Setenv("SAIL_TEST_AK", "from-env-ak") + t.Setenv("SAIL_TEST_SK", "from-env-sk") + // endpoint 混用文本 + 占位符 + t.Setenv("SAIL_TEST_HOST", "s3.example.com") + r := resolveFrom(t, `default-profile: prod +profiles: + prod: + endpoint: https://${SAIL_TEST_HOST}:9000 + access-key: ${SAIL_TEST_AK} + secret-key: ${SAIL_TEST_SK} +`) + if r.Endpoint != "https://s3.example.com:9000" { + t.Errorf("endpoint 展开 = %q,期望 https://s3.example.com:9000", r.Endpoint) + } + if r.AccessKey != "from-env-ak" || r.SecretKey != "from-env-sk" { + t.Errorf("密钥展开错误: %q/%q", r.AccessKey, r.SecretKey) + } + + // 占位符未设置 → 空串 → 触发缺密钥报错 + cfg, err := Load(writeConfig(t, `default-profile: prod +profiles: + prod: + endpoint: https://p.example.com + access-key: ${SAIL_UNSET_AK_XYZ} + secret-key: ${SAIL_UNSET_SK_XYZ} +`)) + if err != nil { + t.Fatalf("Load: %v", err) + } + if _, err := cfg.Resolve("prod"); err == nil { + t.Errorf("占位符未设置导致密钥为空,应报错") + } +} + +// TestResolveEnvOverride 校验 SAIL_* 环境变量优先于配置文件。 +func TestResolveEnvOverride(t *testing.T) { + t.Setenv("SAIL_ENDPOINT", "https://env.example.com") + t.Setenv("SAIL_ACCESS_KEY", "env-ak") + t.Setenv("SAIL_SECRET_KEY", "env-sk") + t.Setenv("SAIL_BUCKET", "env-bucket") + t.Setenv("SAIL_CDN_DOMAIN", "https://env-cdn.example.com") + r := resolveFrom(t, `default-profile: prod +profiles: + prod: + endpoint: https://cfg.example.com + access-key: cfg-ak + secret-key: cfg-sk + bucket: cfg-bucket + cdn-domain: https://cfg-cdn.example.com +`) + if r.Endpoint != "https://env.example.com" || r.AccessKey != "env-ak" || + r.SecretKey != "env-sk" || r.Bucket != "env-bucket" || r.CDNDomain != "https://env-cdn.example.com" { + t.Errorf("环境变量应覆盖配置文件: %+v", r) + } +} + +func TestResolveFieldPassthrough(t *testing.T) { + r := resolveFrom(t, `default-profile: prod +profiles: + prod: + endpoint: https://p.example.com + access-key: ak + secret-key: sk + bucket: mybucket + region: us-west-2 + path-style: true + cdn-domain: https://cdn.example.com +`) + if r.Bucket != "mybucket" || r.Region != "us-west-2" || r.CDNDomain != "https://cdn.example.com" { + t.Errorf("字段透传错误: %+v", r) + } +} + +// TestLoadErrors 校验 Load 对缺失文件 / 非法 YAML 的报错。 +func TestLoadErrors(t *testing.T) { + if _, err := Load(filepath.Join(t.TempDir(), "nope.yaml")); err == nil { + t.Errorf("缺失文件应报错") + } + if _, err := Load(writeConfig(t, ":\n bad: [yaml")); err == nil { + t.Errorf("非法 YAML 应报错") + } +} + +// TestConfigPath 校验 ConfigPath 拼接 ~/.sail/config.yaml。 +func TestConfigPath(t *testing.T) { + t.Setenv("HOME", "/tmp/fake-home") + p, err := ConfigPath() + if err != nil { + t.Fatalf("ConfigPath: %v", err) + } + if p != "/tmp/fake-home/.sail/config.yaml" { + t.Errorf("ConfigPath = %q,期望 /tmp/fake-home/.sail/config.yaml", p) + } +} diff --git a/internal/s3path/s3path_test.go b/internal/s3path/s3path_test.go new file mode 100644 index 0000000..2caeb86 --- /dev/null +++ b/internal/s3path/s3path_test.go @@ -0,0 +1,109 @@ +package s3path + +import "testing" + +func TestParse(t *testing.T) { + cases := []struct { + name string + in string + wantBucket string + wantKey string + wantErr bool + }{ + {name: "空串报错", in: "", wantErr: true}, + {name: "无 s3:// 前缀报错", in: "bucket/key", wantErr: true}, + {name: "仅前缀无 bucket 报错", in: "s3://", wantErr: true}, + {name: "仅 bucket(key 空)", in: "s3://bucket", wantBucket: "bucket"}, + {name: "bucket + key", in: "s3://bucket/key", wantBucket: "bucket", wantKey: "key"}, + {name: "多级 key", in: "s3://bucket/a/b/c.txt", wantBucket: "bucket", wantKey: "a/b/c.txt"}, + {name: "尾斜杠 key 为空", in: "s3://bucket/", wantBucket: "bucket", wantKey: ""}, + {name: "bucket 内含连字符", in: "s3://my-bucket.1/x", wantBucket: "my-bucket.1", wantKey: "x"}, + {name: "空 key 段保留点", in: "s3://b//c", wantBucket: "b", wantKey: "/c"}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + p, err := Parse(c.in) + if c.wantErr { + if err == nil { + t.Errorf("Parse(%q) 期望报错,实际 %v", c.in, p) + } + return + } + if err != nil { + t.Fatalf("Parse(%q) 报错: %v", c.in, err) + } + if p.Bucket != c.wantBucket || p.Key != c.wantKey { + t.Errorf("Parse(%q) = {%q,%q},期望 {%q,%q}", c.in, p.Bucket, p.Key, c.wantBucket, c.wantKey) + } + }) + } +} + +func TestFormat(t *testing.T) { + cases := []struct { + name string + in S3Path + want string + }{ + {name: "仅 bucket", in: S3Path{Bucket: "bucket"}, want: "s3://bucket"}, + {name: "bucket + key", in: S3Path{Bucket: "bucket", Key: "a/b.txt"}, want: "s3://bucket/a/b.txt"}, + {name: "空 key 段", in: S3Path{Bucket: "b", Key: "/c"}, want: "s3://b//c"}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + if got := c.in.Format(); got != c.want { + t.Errorf("(%v).Format() = %q,期望 %q", c.in, got, c.want) + } + }) + } +} + +// TestParseFormatRoundTrip 校验 Parse 与 Format 往返一致。 +func TestParseFormatRoundTrip(t *testing.T) { + for _, s := range []string{"s3://bucket", "s3://bucket/key", "s3://bucket/a/b/c.txt"} { + p, err := Parse(s) + if err != nil { + t.Fatalf("Parse(%q) 报错: %v", s, err) + } + if got := p.Format(); got != s { + t.Errorf("Parse(%q).Format() = %q,往返不一致", s, got) + } + } +} + +func TestBaseName(t *testing.T) { + cases := []struct { + in string + want string + }{ + {"", ""}, + {"a", "a"}, + {"a/b/c", "c"}, + {"a/b/", ""}, // 尾斜杠:末段为空 + {"/a", "a"}, // 前导斜杠 + {"a.txt", "a.txt"}, + } + for _, c := range cases { + if got := BaseName(c.in); got != c.want { + t.Errorf("BaseName(%q) = %q,期望 %q", c.in, got, c.want) + } + } +} + +func TestJoinKey(t *testing.T) { + cases := []struct { + base, name string + want string + }{ + {"", "b", "b"}, // base 空:直接用 name + {"a", "b", "a/b"}, // base 无尾斜杠:补 / + {"a/", "b", "a/b"}, // base 已带尾斜杠:不重复 + {"a/b", "c/d", "a/b/c/d"}, // name 可含多级 + {"a/", "", "a/"}, // name 空:保留 base + } + for _, c := range cases { + if got := JoinKey(c.base, c.name); got != c.want { + t.Errorf("JoinKey(%q, %q) = %q,期望 %q", c.base, c.name, got, c.want) + } + } +} diff --git a/internal/uploader/uploader.go b/internal/uploader/uploader.go index e9d1717..8ec47ec 100644 --- a/internal/uploader/uploader.go +++ b/internal/uploader/uploader.go @@ -20,8 +20,10 @@ import ( const partSize = 5 * 1024 * 1024 // 5MB,S3 multipart 最小分片大小 // Uploader 包装 s3manager,提供单文件/目录上传。 +// s3 字段声明为 manager.UploadAPIClient 接口(而非具体 *s3.Client), +// 使上传路径可在测试中用 mock 覆盖;真实客户端天然满足该接口。 type Uploader struct { - s3 *s3.Client + s3 manager.UploadAPIClient uploader *manager.Uploader } @@ -87,7 +89,6 @@ func (u *Uploader) UploadStream(ctx context.Context, r io.Reader, bucket, key st // UploadDir 递归上传本地目录到 bucket 下的 prefix。 func (u *Uploader) UploadDir(ctx context.Context, localDir, bucket, prefix string) error { - prefix = strings.Trim(prefix, "/") err := filepath.Walk(localDir, func(path string, info os.FileInfo, err error) error { if err != nil { return err @@ -99,17 +100,24 @@ func (u *Uploader) UploadDir(ctx context.Context, localDir, bucket, prefix strin if err != nil { return err } - rel = filepath.ToSlash(rel) - key := rel - if prefix != "" { - key = prefix + "/" + rel - } + key := buildKey(prefix, filepath.ToSlash(rel)) fmt.Printf(i18n.T("uploading %s -> s3://%s/%s\n"), path, bucket, key) return u.UploadFile(ctx, path, bucket, key) }) return err } +// buildKey 在 S3 key 的 base prefix 之上追加子路径 rel。 +// prefix 首尾的 / 会被裁掉(空 prefix 直接返回 rel),避免拼出形如 +// "prefix//rel" 或 "/rel" 的脏 key。 +func buildKey(prefix, rel string) string { + prefix = strings.Trim(prefix, "/") + if prefix == "" { + return rel + } + return prefix + "/" + rel +} + // progressReader 跟踪读取字节数并周期性打印进度。 type progressReader struct { r io.Reader diff --git a/internal/uploader/uploader_test.go b/internal/uploader/uploader_test.go new file mode 100644 index 0000000..fb4e793 --- /dev/null +++ b/internal/uploader/uploader_test.go @@ -0,0 +1,187 @@ +package uploader + +import ( + "bytes" + "context" + "io" + "os" + "path/filepath" + "sort" + "testing" + + "github.com/aws/aws-sdk-go-v2/feature/s3/manager" + "github.com/aws/aws-sdk-go-v2/service/s3" +) + +func TestHumanBytes(t *testing.T) { + cases := []struct { + in int64 + want string + }{ + {0, "0 B"}, + {1023, "1023 B"}, + {1024, "1.0 KiB"}, + {1536, "1.5 KiB"}, + {5 * 1024 * 1024, "5.0 MiB"}, + {2 << 30, "2.0 GiB"}, + {1024 * 1024 * 1024 * 1024, "1.0 TiB"}, + } + for _, c := range cases { + if got := humanBytes(c.in); got != c.want { + t.Errorf("humanBytes(%d) = %q,期望 %q", c.in, got, c.want) + } + } +} + +func TestBuildKey(t *testing.T) { + cases := []struct { + prefix, rel string + want string + }{ + {"", "a/b.txt", "a/b.txt"}, // 空 prefix:直接用 rel + {"mirror", "a/b.txt", "mirror/a/b.txt"}, + {"mirror/", "a.txt", "mirror/a.txt"}, // prefix 尾斜杠被裁 + {"/mirror", "a.txt", "mirror/a.txt"}, // prefix 首斜杠被裁 + {"/mirror/", "a.txt", "mirror/a.txt"}, // 首尾都裁 + {"a/b", "c.txt", "a/b/c.txt"}, + } + for _, c := range cases { + if got := buildKey(c.prefix, c.rel); got != c.want { + t.Errorf("buildKey(%q,%q) = %q,期望 %q", c.prefix, c.rel, got, c.want) + } + } +} + +func TestProgressReaderRead(t *testing.T) { + pr := newProgressReader(bytes.NewReader([]byte("hello")), 5) + buf := make([]byte, 3) + n, err := pr.Read(buf) + if n != 3 || err != nil { + t.Fatalf("首次 Read = (%d,%v),期望 (3,nil)", n, err) + } + pr.mu.Lock() + if pr.read != 3 { + t.Errorf("read 计数 = %d,期望 3", pr.read) + } + pr.mu.Unlock() + pr.Read(buf) // 消费剩余 + pr.mu.Lock() + if pr.read != 5 { + t.Errorf("读尽后 read = %d,期望 5", pr.read) + } + pr.mu.Unlock() +} + +// uploadClientRecorder 实现 manager.UploadAPIClient,把每次 PutObject 的 +// key 与 body 记下来,使 manager.Uploader 的上传路径可在无网络下走通。 +type uploadClientRecorder struct { + keys []string + body string +} + +func (r *uploadClientRecorder) PutObject(ctx context.Context, in *s3.PutObjectInput, opts ...func(*s3.Options)) (*s3.PutObjectOutput, error) { + if in.Key != nil { + r.keys = append(r.keys, *in.Key) + } + if in.Body != nil { + b, err := io.ReadAll(in.Body) + if err != nil { + return nil, err + } + r.body += string(b) + } + return &s3.PutObjectOutput{}, nil +} +func (r *uploadClientRecorder) UploadPart(context.Context, *s3.UploadPartInput, ...func(*s3.Options)) (*s3.UploadPartOutput, error) { + return nil, nil +} +func (r *uploadClientRecorder) CreateMultipartUpload(context.Context, *s3.CreateMultipartUploadInput, ...func(*s3.Options)) (*s3.CreateMultipartUploadOutput, error) { + return nil, nil +} +func (r *uploadClientRecorder) CompleteMultipartUpload(context.Context, *s3.CompleteMultipartUploadInput, ...func(*s3.Options)) (*s3.CompleteMultipartUploadOutput, error) { + return nil, nil +} +func (r *uploadClientRecorder) AbortMultipartUpload(context.Context, *s3.AbortMultipartUploadInput, ...func(*s3.Options)) (*s3.AbortMultipartUploadOutput, error) { + return nil, nil +} + +// newTestUploader 构造指向 recorder 的 Uploader。 +func newTestUploader(rec *uploadClientRecorder) *Uploader { + return &Uploader{ + s3: rec, + uploader: manager.NewUploader(rec, func(o *manager.Uploader) { + o.PartSize = partSize + o.Concurrency = 1 + }), + } +} + +func TestUploadFileCapturesKeyAndBody(t *testing.T) { + path := filepath.Join(t.TempDir(), "hello.txt") + if err := os.WriteFile(path, []byte("file-content"), 0o600); err != nil { + t.Fatalf("写文件: %v", err) + } + rec := &uploadClientRecorder{} + u := newTestUploader(rec) + + if err := u.UploadFile(context.Background(), path, "mybucket", "dir/file.txt"); err != nil { + t.Fatalf("UploadFile 报错: %v", err) + } + if len(rec.keys) != 1 || rec.keys[0] != "dir/file.txt" { + t.Errorf("上传 key = %v,期望 [dir/file.txt]", rec.keys) + } + if rec.body != "file-content" { + t.Errorf("上传内容 = %q,期望 file-content", rec.body) + } +} + +func TestUploadFileMissingPath(t *testing.T) { + u := newTestUploader(&uploadClientRecorder{}) + if err := u.UploadFile(context.Background(), filepath.Join(t.TempDir(), "nope"), "b", "k"); err == nil { + t.Errorf("缺失本地文件应报错") + } +} + +func TestUploadStream(t *testing.T) { + rec := &uploadClientRecorder{} + u := newTestUploader(rec) + if err := u.UploadStream(context.Background(), bytes.NewReader([]byte("stream-data")), "b", "k"); err != nil { + t.Fatalf("UploadStream 报错: %v", err) + } + if rec.body != "stream-data" { + t.Errorf("上传内容 = %q,期望 stream-data", rec.body) + } +} + +// TestUploadDir 用临时目录验证:目录跳过、rel 转 / 分隔、prefix 拼接。 +func TestUploadDir(t *testing.T) { + dir := t.TempDir() + mustWrite := func(rel, content string) { + p := filepath.Join(dir, filepath.FromSlash(rel)) + if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(p, []byte(content), 0o600); err != nil { + t.Fatal(err) + } + } + mustWrite("sub/a.txt", "A") + mustWrite("top.log", "T") + + rec := &uploadClientRecorder{} + u := newTestUploader(rec) + if err := u.UploadDir(context.Background(), dir, "bucket", "/mirror/"); err != nil { + t.Fatalf("UploadDir 报错: %v", err) + } + got := append([]string(nil), rec.keys...) + sort.Strings(got) + want := []string{"mirror/sub/a.txt", "mirror/top.log"} + if len(got) != len(want) { + t.Fatalf("上传对象数 = %d,期望 %d(keys=%v)", len(got), len(want), got) + } + for i := range want { + if got[i] != want[i] { + t.Errorf("key[%d] = %q,期望 %q", i, got[i], want[i]) + } + } +} diff --git a/internal/view/view_test.go b/internal/view/view_test.go new file mode 100644 index 0000000..3adeca8 --- /dev/null +++ b/internal/view/view_test.go @@ -0,0 +1,303 @@ +package view + +import ( + "bytes" + "image" + "image/png" + "io" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestParseSuffixRange(t *testing.T) { + cases := []struct { + name string + in string + want int64 + wantErr bool + }{ + {name: "正常后缀窗口", in: "bytes=-10", want: 10}, + {name: "零字节窗口", in: "bytes=-0", want: 0}, + {name: "非后缀 Range 报错", in: "bytes=10-20", wantErr: true}, + {name: "非数字报错", in: "bytes=-abc", wantErr: true}, + {name: "空串报错", in: "", wantErr: true}, + {name: "负数报错", in: "bytes=--5", wantErr: true}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + got, err := parseSuffixRange(c.in) + if c.wantErr { + if err == nil { + t.Errorf("parseSuffixRange(%q) 期望报错,实际 %d", c.in, got) + } + return + } + if err != nil { + t.Fatalf("parseSuffixRange(%q) 报错: %v", c.in, err) + } + if got != c.want { + t.Errorf("parseSuffixRange(%q) = %d,期望 %d", c.in, got, c.want) + } + }) + } +} + +func TestParseFormat(t *testing.T) { + cases := []struct { + name string + in string + want Format + ok bool + }{ + {"空串=auto", "", FormatAuto, true}, + {"auto", "auto", FormatAuto, true}, + {"大小写不敏感", "JSON", FormatJSON, true}, + {"txt 别名", "txt", FormatText, true}, + {"yml 别名", "yml", FormatYAML, true}, + {"tsv 别名", "tsv", FormatCSV, true}, + {"svg 别名", "svg", FormatXML, true}, + {"img 别名", "img", FormatImage, true}, + {"bin 别名", "bin", FormatBinary, true}, + {"未知值", "pdf", 0, false}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + got, ok := ParseFormat(c.in) + if ok != c.ok || (c.ok && got != c.want) { + t.Errorf("ParseFormat(%q) = %v,%v,期望 %v,%v", c.in, got, ok, c.want, c.ok) + } + }) + } +} + +func TestMimeFormat(t *testing.T) { + cases := []struct { + name string + in string + want Format + ok bool + }{ + {"csv 带 charset", "text/csv; charset=utf-8", FormatCSV, true}, + {"tsv 变体", "text/tab-separated-values", FormatCSV, true}, + {"纯文本", "text/plain", FormatText, true}, + {"json", "application/json", FormatJSON, true}, + {"yaml", "application/x-yaml", FormatYAML, true}, + {"xml", "text/xml", FormatXML, true}, + {"image 前缀", "image/png", FormatImage, true}, + {"svg 当前命中 image(固化现状)", "image/svg+xml", FormatImage, true}, + {"未知类型", "application/octet-stream", 0, false}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + got, ok := mimeFormat(c.in) + if ok != c.ok || (c.ok && got != c.want) { + t.Errorf("mimeFormat(%q) = %v,%v,期望 %v,%v", c.in, got, ok, c.want, c.ok) + } + }) + } +} + +func TestDetectFormat(t *testing.T) { + // 扩展名优先 + if got := DetectFormat(&Source{Name: "a.json"}); got != FormatJSON { + t.Errorf("a.json = %v,期望 JSON", got) + } + // 扩展名未命中,ContentType 命中(含参数剥离) + if got := DetectFormat(&Source{Name: "noext", ContentType: "text/csv; charset=utf-8"}); got != FormatCSV { + t.Errorf("text/csv = %v,期望 CSV", got) + } + // ContentType 空则嗅探(用 PNG 魔数:DetectContentType 可靠判为 image/png) + pngBytes := encodeTestPNG() + s := &Source{Name: "noext", Reader: io.NopCloser(bytes.NewReader(pngBytes))} + if got := DetectFormat(s); got != FormatImage { + t.Errorf("嗅探 PNG = %v,期望 FormatImage(嗅探得 image/png)", got) + } + // 扩展名优先于 ContentType(同文件扩展名胜出) + if got := DetectFormat(&Source{Name: "x.txt", ContentType: "image/png"}); got != FormatText { + t.Errorf("x.txt + image/png = %v,期望 Text(扩展名优先)", got) + } + // 全部未知 -> Binary + if got := DetectFormat(&Source{Name: "noext", ContentType: "application/octet-stream"}); got != FormatBinary { + t.Errorf("未知 = %v,期望 Binary", got) + } +} + +func TestHumanBytes(t *testing.T) { + cases := []struct { + in int64 + want string + }{ + {0, "0 B"}, + {1023, "1023 B"}, + {1024, "1.0 KiB"}, + {1536, "1.5 KiB"}, + {5 * 1024 * 1024, "5.0 MiB"}, + {2 << 30, "2.0 GiB"}, + } + for _, c := range cases { + if got := humanBytes(c.in); got != c.want { + t.Errorf("humanBytes(%d) = %q,期望 %q", c.in, got, c.want) + } + } +} + +func TestAnsiFG(t *testing.T) { + if got := string(ansiFG(1, 2, 3)); got != "\x1b[38;2;1;2;3m" { + t.Errorf("ansiFG = %q", got) + } + if got := string(ansiBG(10, 20, 30)); got != "\x1b[48;2;10;20;30m" { + t.Errorf("ansiBG = %q", got) + } +} + +func TestRGBA(t *testing.T) { + img := newTestRGBA(1, 1, 0xAB, 0xCD, 0xEF) + r, g, b := rgba(img, 0, 0) + if r != 0xAB || g != 0xCD || b != 0xEF { + t.Errorf("rgba = (%d,%d,%d),期望 (171,205,239)", r, g, b) + } +} + +func TestSniffContentType(t *testing.T) { + // 已有 ContentType:短路不读 Reader + body := "hello" + s := &Source{Reader: io.NopCloser(strings.NewReader(body)), ContentType: "text/plain"} + if err := SniffContentType(s); err != nil { + t.Fatalf("已声明 ContentType 短路仍报错: %v", err) + } + if s.ContentType != "text/plain" { + t.Errorf("ContentType 被改写成 %q", s.ContentType) + } + + // 嗅探后数据不丢失(关键不变量) + payload := `x` + s = &Source{Reader: io.NopCloser(strings.NewReader(payload))} + if err := SniffContentType(s); err != nil { + t.Fatalf("SniffContentType 报错: %v", err) + } + if s.ContentType != "text/html; charset=utf-8" { + t.Errorf("嗅探 ContentType = %q,期望 text/html", s.ContentType) + } + all, err := io.ReadAll(s.Reader) + if err != nil { + t.Fatalf("读取拼接 Reader 报错: %v", err) + } + if string(all) != payload { + t.Errorf("嗅探后数据丢失: got %d 字节,期望 %d", len(all), len(payload)) + } +} + +func TestReadBounded(t *testing.T) { + // 未超限通过 + s := &Source{Reader: io.NopCloser(strings.NewReader("12345")), Size: 5} + b, err := readBounded(s, 10, false) + if err != nil || string(b) != "12345" { + t.Errorf("未超限应通过,got %q err=%v", b, err) + } + + // Size 已知且超限:提前报错 + s = &Source{Reader: io.NopCloser(strings.NewReader(strings.Repeat("a", 100))), Size: 100} + if _, err := readBounded(s, 10, false); err == nil { + t.Errorf("Size 超限应报错") + } + + // force 放行超限 + s = &Source{Reader: io.NopCloser(strings.NewReader(strings.Repeat("a", 100))), Size: 100} + if _, err := readBounded(s, 10, true); err != nil { + t.Errorf("force 应放行超限,报错: %v", err) + } +} + +func TestOpenLocalRejectsDir(t *testing.T) { + dir := t.TempDir() + if _, err := openLocal(dir); err == nil { + t.Errorf("目录应被拒绝") + } +} + +func TestOpenLocalRangeWindow(t *testing.T) { + path := filepath.Join(t.TempDir(), "data.txt") + content := "0123456789" // 10 字节 + if err := os.WriteFile(path, []byte(content), 0o600); err != nil { + t.Fatalf("写文件: %v", err) + } + + // bytes=-4:尾部 4 字节,Size 为窗口长度 + s, err := openLocalRange(path, "bytes=-4") + if err != nil { + t.Fatalf("openLocalRange 报错: %v", err) + } + defer s.Close() + if s.Size != 4 { + t.Errorf("窗口 Size = %d,期望 4", s.Size) + } + b, err := io.ReadAll(s.Reader) + if err != nil { + t.Fatalf("读取报错: %v", err) + } + if string(b) != "6789" { + t.Errorf("窗口内容 = %q,期望 6789", string(b)) + } + + // bytes=-100:超过文件大小,截断为全量 + s2, err := openLocalRange(path, "bytes=-100") + if err != nil { + t.Fatalf("openLocalRange 报错: %v", err) + } + defer s2.Close() + if s2.Size != 10 { + t.Errorf("超大窗口 Size = %d,期望 10(全量)", s2.Size) + } +} + +func TestOpenLocalRangeInvalidRange(t *testing.T) { + path := filepath.Join(t.TempDir(), "data.txt") + if err := os.WriteFile(path, []byte("12345"), 0o600); err != nil { + t.Fatalf("写文件: %v", err) + } + if _, err := openLocalRange(path, "bytes=0-5"); err == nil { + t.Errorf("非后缀 Range 应报错") + } +} + +func TestOpenLocal(t *testing.T) { + path := filepath.Join(t.TempDir(), "hello.txt") + if err := os.WriteFile(path, []byte("hello"), 0o600); err != nil { + t.Fatalf("写文件: %v", err) + } + s, err := openLocal(path) + if err != nil { + t.Fatalf("openLocal 报错: %v", err) + } + defer s.Close() + if s.Size != 5 || s.IsS3 { + t.Errorf("openLocal 元信息错误: Size=%d IsS3=%v", s.Size, s.IsS3) + } + if s.Name != "hello.txt" { + t.Errorf("Name = %q,期望 hello.txt", s.Name) + } +} + +// newTestRGBA 构造一个像素可预期的 image.RGBA 便于 rgba() 断言。 +func newTestRGBA(w, h, r, g, b int) *image.RGBA { + img := image.NewRGBA(image.Rect(0, 0, w, h)) + for y := 0; y < h; y++ { + for x := 0; x < w; x++ { + i := (y*img.Stride + x*4) + img.Pix[i] = uint8(r) + img.Pix[i+1] = uint8(g) + img.Pix[i+2] = uint8(b) + img.Pix[i+3] = 255 + } + } + return img +} + +// encodeTestPNG 编码一张 1x1 PNG 字节流,供嗅探路径构造 Source。 +func encodeTestPNG() []byte { + var buf bytes.Buffer + _ = png.Encode(&buf, newTestRGBA(1, 1, 1, 2, 3)) + return buf.Bytes() +}