package storage import ( "bytes" "context" "encoding/xml" "fmt" "io" "net/http" "net/http/httptest" "sort" "strings" "sync" "testing" "time" ) // ---- 最小 S3 兼容假服务(仅覆盖本引擎用到的 API)---- type fakeS3Upload struct { key string parts map[int][]byte } type fakeS3 struct { mu sync.Mutex objects map[string][]byte uploads map[string]*fakeS3Upload nextID int putCount int getCount int headCount int deleteCount int listCount int completeN int // failNextGet:让接下来 N 次 GET 返回 503(重试测试用)。 failNextGet int } func newFakeS3() *fakeS3 { return &fakeS3{objects: map[string][]byte{}, uploads: map[string]*fakeS3Upload{}} } // s3Key 从 path-style 路径剥离 bucket 前缀得到对象键。 func s3Key(r *http.Request) (bucket, key string) { p := strings.TrimPrefix(r.URL.Path, "/") if i := strings.Index(p, "/"); i >= 0 { return p[:i], p[i+1:] } return p, "" } func s3ErrorXML(w http.ResponseWriter, status int, code, msg string) { w.Header().Set("Content-Type", "application/xml") w.WriteHeader(status) _, _ = w.Write([]byte(fmt.Sprintf( `%s%s`, code, msg))) } func (f *fakeS3) ServeHTTP(w http.ResponseWriter, r *http.Request) { f.mu.Lock() defer f.mu.Unlock() q := r.URL.Query() bucket, key := s3Key(r) switch { // UploadPart case r.Method == http.MethodPut && q.Get("partNumber") != "" && q.Get("uploadId") != "": up, ok := f.uploads[q.Get("uploadId")] if !ok { s3ErrorXML(w, 404, "NoSuchUpload", "upload not found") return } body, _ := io.ReadAll(r.Body) var n int _, _ = fmt.Sscanf(q.Get("partNumber"), "%d", &n) up.parts[n] = body w.Header().Set("ETag", fmt.Sprintf(`"part-%d"`, n)) w.WriteHeader(http.StatusOK) // CreateMultipartUpload case r.Method == http.MethodPost && q.Has("uploads"): f.nextID++ id := fmt.Sprintf("mpu-%d", f.nextID) f.uploads[id] = &fakeS3Upload{key: key, parts: map[int][]byte{}} w.Header().Set("Content-Type", "application/xml") _, _ = w.Write([]byte(fmt.Sprintf( `%s%s%s`, bucket, key, id))) // CompleteMultipartUpload case r.Method == http.MethodPost && q.Get("uploadId") != "": up, ok := f.uploads[q.Get("uploadId")] if !ok { s3ErrorXML(w, 404, "NoSuchUpload", "upload not found") return } // 按 partNumber 有序拼接 nums := make([]int, 0, len(up.parts)) for n := range up.parts { nums = append(nums, n) } sort.Ints(nums) var merged bytes.Buffer for _, n := range nums { merged.Write(up.parts[n]) } f.objects[up.key] = merged.Bytes() delete(f.uploads, q.Get("uploadId")) f.completeN++ w.Header().Set("Content-Type", "application/xml") _, _ = w.Write([]byte(fmt.Sprintf( `http://%s/%s/%s%s%s"merged"`, r.Host, bucket, up.key, bucket, up.key))) // AbortMultipartUpload case r.Method == http.MethodDelete && q.Get("uploadId") != "": delete(f.uploads, q.Get("uploadId")) w.WriteHeader(http.StatusNoContent) // DeleteObjects(批量) case r.Method == http.MethodPost && q.Has("delete"): var req struct { Objects []struct { Key string `xml:"Key"` } `xml:"Object"` } _ = xml.NewDecoder(r.Body).Decode(&req) for _, o := range req.Objects { delete(f.objects, o.Key) } w.Header().Set("Content-Type", "application/xml") _, _ = w.Write([]byte(``)) // ListObjectsV2 case r.Method == http.MethodGet && q.Get("list-type") == "2": f.listCount++ prefix := q.Get("prefix") var body strings.Builder body.WriteString(`` + bucket + `` + prefix + `false`) keys := make([]string, 0, len(f.objects)) for k := range f.objects { if strings.HasPrefix(k, prefix) { keys = append(keys, k) } } sort.Strings(keys) for _, k := range keys { body.WriteString(fmt.Sprintf("%s%d", k, len(f.objects[k]))) } body.WriteString("") w.Header().Set("Content-Type", "application/xml") _, _ = w.Write([]byte(body.String())) // PutObject case r.Method == http.MethodPut: f.putCount++ body, _ := io.ReadAll(r.Body) f.objects[key] = body w.Header().Set("ETag", `"put"`) w.WriteHeader(http.StatusOK) // GetObject case r.Method == http.MethodGet: if f.failNextGet > 0 { f.failNextGet-- s3ErrorXML(w, 503, "ServiceUnavailable", "flaky") return } f.getCount++ data, ok := f.objects[key] if !ok { s3ErrorXML(w, 404, "NoSuchKey", "not found") return } w.Header().Set("Content-Type", "application/octet-stream") if rng := r.Header.Get("Range"); rng != "" { start, end := int64(0), int64(len(data))-1 if _, err := fmt.Sscanf(rng, "bytes=%d-%d", &start, &end); err != nil { var s int64 if _, err := fmt.Sscanf(rng, "bytes=%d-", &s); err == nil { start, end = s, int64(len(data))-1 } } if start < 0 || start >= int64(len(data)) { s3ErrorXML(w, 416, "InvalidRange", "range not satisfiable") return } if end >= int64(len(data)) { end = int64(len(data)) - 1 } slice := data[start : end+1] w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, end, len(data))) w.WriteHeader(http.StatusPartialContent) _, _ = w.Write(slice) return } w.WriteHeader(http.StatusOK) _, _ = w.Write(data) // HeadObject case r.Method == http.MethodHead: f.headCount++ data, ok := f.objects[key] if !ok { w.WriteHeader(http.StatusNotFound) return } w.Header().Set("Content-Type", "application/octet-stream") w.Header().Set("Content-Length", fmt.Sprintf("%d", len(data))) w.WriteHeader(http.StatusOK) // DeleteObject case r.Method == http.MethodDelete: f.deleteCount++ delete(f.objects, key) w.WriteHeader(http.StatusNoContent) default: s3ErrorXML(w, 400, "NotImplemented", "unsupported") } } // newTestS3 构造对接假服务的 S3 引擎。 func newTestS3(t *testing.T) (*S3Storage, *fakeS3) { t.Helper() f := newFakeS3() srv := httptest.NewServer(f) t.Cleanup(srv.Close) st, err := NewS3Storage(S3Options{ AccessKeyID: "test-ak", SecretAccessKey: "test-sk", Bucket: "test-bucket", Endpoint: srv.URL, Region: "us-east-1", AddressingStyle: "path", }) if err != nil { t.Fatalf("NewS3Storage: %v", err) } return st, f } // TestS3SaveStatOpenRange 保存/元信息/完整与 Range 下载。 func TestS3SaveStatOpenRange(t *testing.T) { st, f := newTestS3(t) ctx := context.Background() data := []byte("0123456789abcdef S3 引擎测试数据") n, err := st.SaveFile(ctx, bytes.NewReader(data), "2025/08/s3.bin") if err != nil { t.Fatalf("SaveFile: %v", err) } if n != int64(len(data)) { t.Fatalf("n = %d", n) } if got := f.objects["2025/08/s3.bin"]; !bytes.Equal(got, data) { t.Fatalf("stored mismatch") } meta, err := st.Stat(ctx, "2025/08/s3.bin") if err != nil { t.Fatalf("Stat: %v", err) } if meta.Size != int64(len(data)) || !meta.AcceptRanges { t.Fatalf("Stat = %+v", meta) } dl, err := st.Open(ctx, "2025/08/s3.bin", nil) if err != nil { t.Fatalf("Open: %v", err) } got, _ := io.ReadAll(dl) _ = dl.Close() 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 offsets = %d,%d,%d", dl.Start, dl.End, dl.Total) } dl, err = st.Open(ctx, "2025/08/s3.bin", &Range{Start: 4, End: 9}) if err != nil { t.Fatalf("Open range: %v", err) } got, _ = io.ReadAll(dl) _ = dl.Close() if !bytes.Equal(got, data[4:10]) { t.Fatalf("range mismatch: %q", got) } if dl.Start != 4 || dl.End != 9 || dl.Total != int64(len(data)) { t.Fatalf("range offsets = %d,%d,%d", dl.Start, dl.End, dl.Total) } // 404 / 416 if _, err := st.Open(ctx, "missing.bin", nil); err == nil || !strings.Contains(err.Error(), ErrNotFound.Error()) { t.Fatalf("want ErrNotFound, got %v", err) } if _, err := st.Open(ctx, "2025/08/s3.bin", &Range{Start: int64(len(data)) + 5, End: -1}); err == nil || !strings.Contains(err.Error(), ErrRangeNotSatisfiable.Error()) { t.Fatalf("want ErrRangeNotSatisfiable, got %v", err) } } // TestS3DeleteExists 删除与存在性。 func TestS3DeleteExists(t *testing.T) { st, _ := newTestS3(t) ctx := context.Background() if _, err := st.SaveFile(ctx, strings.NewReader("x"), "del.bin"); err != nil { t.Fatal(err) } ok, err := st.FileExists(ctx, "del.bin") if err != nil || !ok { t.Fatalf("exists = %v %v", ok, err) } if err := st.DeleteFile(ctx, "del.bin"); err != nil { t.Fatalf("DeleteFile: %v", err) } ok, err = st.FileExists(ctx, "del.bin") if err != nil || ok { t.Fatalf("after delete exists = %v %v", ok, err) } } // TestS3ChunkMergeMulti 多分片合并:原生 multipart + 哈希校验 + 分片清理。 func TestS3ChunkMergeMulti(t *testing.T) { st, f := newTestS3(t) ctx := context.Background() savePath := "2025/09/s3-chunked.bin" uploadID := "uid-s3" chunks := [][]byte{bytes.Repeat([]byte("A"), 6*1024*1024/3), []byte("BBBB"), []byte("CC")} // 注意:multipart 除最后一片需 ≥5MB;此处只验证代码路径,真实约束由部署配置保证。 // 为避免 EntityTooSmall,将第一片放大: chunks[0] = bytes.Repeat([]byte("A"), 5*1024*1024) 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) } hashes[i] = sha256Hex(c) total += int64(len(c)) _ = n } 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("size = %d want %d", size, total) } if fileHash != sha256Hex(bytes.Join(chunks, nil)) { t.Fatalf("file hash mismatch") } merged := f.objects[savePath] if !bytes.Equal(merged, bytes.Join(chunks, nil)) { t.Fatalf("merged object mismatch (len=%d)", len(merged)) } if f.completeN != 1 { t.Fatalf("CompleteMultipartUpload 次数 = %d", f.completeN) } // 分片对象已清理 for k := range f.objects { if strings.Contains(k, "chunks/"+uploadID) { t.Fatalf("分片对象残留: %s", k) } } } // TestS3ChunkMergeSingle 单分片快速路径。 func TestS3ChunkMergeSingle(t *testing.T) { st, f := newTestS3(t) ctx := context.Background() data := []byte("single-chunk") if _, err := st.SaveChunk(ctx, "uid1", 0, bytes.NewReader(data), "one.bin"); err != nil { t.Fatal(err) } size, fileHash, err := st.MergeChunks(ctx, "uid1", 1, func(i int) (string, error) { return sha256Hex(data), nil }, "one.bin") if err != nil { t.Fatalf("MergeChunks: %v", err) } if size != int64(len(data)) || fileHash != sha256Hex(data) { t.Fatalf("size/hash mismatch") } if !bytes.Equal(f.objects["one.bin"], data) { t.Fatalf("object mismatch") } } // TestS3CleanChunks 清理残留分片。 func TestS3CleanChunks(t *testing.T) { st, f := newTestS3(t) ctx := context.Background() savePath := "clean.bin" for i := 0; i < 3; i++ { if _, err := st.SaveChunk(ctx, "uidc", i, strings.NewReader("zz"), savePath); err != nil { t.Fatal(err) } } prefix := chunkDirOf(savePath, "uidc") + "/" count := 0 for k := range f.objects { if strings.HasPrefix(k, prefix) { count++ } } if count != 3 { t.Fatalf("期望 3 个分片对象,实际 %d", count) } if err := st.CleanChunks(ctx, "uidc", savePath); err != nil { t.Fatalf("CleanChunks: %v", err) } for k := range f.objects { if strings.HasPrefix(k, prefix) { t.Fatalf("分片未清理: %s", k) } } // 幂等:再清理一次不报错 if err := st.CleanChunks(ctx, "uidc", savePath); err != nil { t.Fatalf("CleanChunks idempotent: %v", err) } } // TestS3Presign 预签名 URL 生成。 func TestS3Presign(t *testing.T) { st, _ := newTestS3(t) ctx := context.Background() getURL, err := st.PresignGetURL(ctx, "presign.bin", 600) if err != nil { t.Fatalf("PresignGetURL: %v", err) } if !strings.Contains(getURL, "X-Amz-Signature") || !strings.Contains(getURL, "X-Amz-Expires=600") { t.Fatalf("GET 直链缺少签名参数: %s", getURL) } putURL, err := st.PresignPutURL(ctx, "presign.bin", 300) if err != nil { t.Fatalf("PresignPutURL: %v", err) } if !strings.Contains(putURL, "X-Amz-Signature") { t.Fatalf("PUT 直链缺少签名参数: %s", putURL) } } // TestS3HealthCheck 健康检查(ListObjectsV2)。 func TestS3HealthCheck(t *testing.T) { st, _ := newTestS3(t) if err := st.HealthCheck(context.Background()); err != nil { t.Fatalf("HealthCheck: %v", err) } } // TestS3RetryOn503 SDK 内置重试器:503 后成功。 func TestS3RetryOn503(t *testing.T) { st, f := newTestS3(t) ctx := context.Background() data := []byte("retry-me") if _, err := st.SaveFile(ctx, bytes.NewReader(data), "retry.bin"); err != nil { t.Fatal(err) } f.mu.Lock() f.failNextGet = 1 f.mu.Unlock() dl, err := st.Open(ctx, "retry.bin", nil) if err != nil { t.Fatalf("503 后应重试成功: %v", err) } got, _ := io.ReadAll(dl) _ = dl.Close() if !bytes.Equal(got, data) { t.Fatalf("content mismatch") } } // TestS3FactoryRegistry 工厂构造。 func TestS3FactoryRegistry(t *testing.T) { f := newFakeS3() srv := httptest.NewServer(f) defer srv.Close() prev := engineOptions.S3 engineOptions.S3 = S3Options{ AccessKeyID: "ak", SecretAccessKey: "sk", Bucket: "b", Endpoint: srv.URL, Region: "us-east-1", AddressingStyle: "path", } defer func() { engineOptions.S3 = prev }() st, err := NewEngine(context.Background(), "s3") if err != nil { t.Fatalf("NewEngine(s3): %v", err) } if err := st.HealthCheck(context.Background()); err != nil { t.Fatalf("HealthCheck: %v", err) } if _, err := st.PresignGetURL(context.Background(), "x.bin", 60); err != nil { t.Fatalf("Presign: %v", err) } _ = time.Now }