package storage import ( "context" "crypto/sha256" "encoding/hex" "errors" "fmt" "io" "net" "net/http" "os" "strconv" "strings" "time" "github.com/aws/aws-sdk-go-v2/aws" awshttp "github.com/aws/aws-sdk-go-v2/aws/transport/http" "github.com/aws/aws-sdk-go-v2/config" "github.com/aws/aws-sdk-go-v2/credentials" "github.com/aws/aws-sdk-go-v2/feature/s3/manager" "github.com/aws/aws-sdk-go-v2/service/s3" "github.com/aws/aws-sdk-go-v2/service/s3/types" "github.com/aws/smithy-go" ) // S3Storage 基于 aws-sdk-go-v2 的 S3 兼容对象存储引擎(AWS / MinIO / R2 / OSS 等)。 // // 相比参考实现(S3FileStorage,aioboto3)的改进: // - 单例客户端 + 自定义连接池 Transport(参考实现每次操作新建 session); // - SaveFile 走 manager.Uploader 分片并发上传(未知长度也可流式,内存占用 ≤ partSize); // - SaveChunk 落本地临时文件获取精确 Content-Length(参考实现整块读入内存); // - MergeChunks 用 S3 原生 multipart 流式合并,边读边校验哈希(不落盘、不整块进内存); // - 5xx/网络错误由 SDK 内置指数退避重试器处理(可配次数)。 type S3Storage struct { client *s3.Client presigner *s3.PresignClient uploader *manager.Uploader bucket string } // NewS3Storage 构造 S3 引擎。 func NewS3Storage(opts S3Options) (*S3Storage, error) { if strings.TrimSpace(opts.Bucket) == "" { return nil, fmt.Errorf("storage/s3: 缺少 bucket 配置(s3_bucket_name)") } region := strings.TrimSpace(opts.Region) if region == "" { region = "us-east-1" } loadOpts := []func(*config.LoadOptions) error{ config.WithRegion(region), config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider( opts.AccessKeyID, opts.SecretAccessKey, opts.SessionToken, )), // SDK 内置重试器:标准模式,指数退避 + 抖动,覆盖 5xx 与网络错误。 config.WithRetryMaxAttempts(3), // 兼容性:仅协议要求时才计算校验和。默认的 trailing CRC32 需要可重放流 // 或 TLS,MinIO/R2 等自建端点通常不需要,关闭后 MergeChunks 的 // GET→UploadPart 纯流式转发才能工作。 config.WithRequestChecksumCalculation(aws.RequestChecksumCalculationWhenRequired), config.WithResponseChecksumValidation(aws.ResponseChecksumValidationWhenRequired), } if ep := strings.TrimSpace(opts.Endpoint); ep != "" { loadOpts = append(loadOpts, config.WithBaseEndpoint(ep)) } awsCfg, err := config.LoadDefaultConfig(context.Background(), loadOpts...) if err != nil { return nil, fmt.Errorf("storage/s3: 初始化 SDK 配置失败: %w", err) } client := s3.NewFromConfig(awsCfg, func(o *s3.Options) { // 寻址风格:path 显式启用;auto 时自定义端点(自建 MinIO 等)默认 path-style。 switch strings.ToLower(strings.TrimSpace(opts.AddressingStyle)) { case "path": o.UsePathStyle = true case "virtual": o.UsePathStyle = false default: // auto o.UsePathStyle = strings.TrimSpace(opts.Endpoint) != "" } // 连接复用:自定义 Transport 连接池。 o.HTTPClient = newPooledHTTPClient() }) st := &S3Storage{ client: client, presigner: s3.NewPresignClient(client), bucket: opts.Bucket, } st.uploader = manager.NewUploader(client, func(u *manager.Uploader) { u.PartSize = 5 * 1024 * 1024 // 5MB,S3 multipart 最小分片 u.Concurrency = 4 u.LeavePartsOnError = false }) return st, nil } func init() { RegisterEngine("s3", func(ctx context.Context) (Storage, error) { return NewS3Storage(engineOptions.S3) }) } // newPooledHTTPClient 供 SDK 使用的连接池化 HTTP 客户端。 func newPooledHTTPClient() *http.Client { return &http.Client{ Transport: &http.Transport{ Proxy: http.ProxyFromEnvironment, DialContext: (&net.Dialer{ Timeout: 10 * time.Second, KeepAlive: 30 * time.Second, }).DialContext, ForceAttemptHTTP2: true, MaxIdleConns: 100, MaxIdleConnsPerHost: 16, IdleConnTimeout: 90 * time.Second, TLSHandshakeTimeout: 10 * time.Second, ExpectContinueTimeout: time.Second, ResponseHeaderTimeout: 60 * time.Second, }, } } // key 校验并规范化对象键(拒绝穿越,统一斜杠)。 func (s *S3Storage) key(savePath string) (string, error) { cleaned, ok := SanitizePath(savePath) if !ok { return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath) } return cleaned, nil } // SaveFile 流式保存:manager.Uploader 按需分片并发上传,内存占用恒定。 func (s *S3Storage) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) { key, err := s.key(savePath) if err != nil { return 0, err } src := &countingReader{r: r} _, err = s.uploader.Upload(ctx, &s3.PutObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), Body: src, ContentType: aws.String("application/octet-stream"), }) if err != nil { return src.count(), mapS3Error(err, "PutObject") } return src.count(), nil } // DeleteFile 删除对象;S3 对不存在的键也返回成功。 func (s *S3Storage) DeleteFile(ctx context.Context, savePath string) error { key, err := s.key(savePath) if err != nil { return err } _, err = s.client.DeleteObject(ctx, &s3.DeleteObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), }) return mapS3Error(err, "DeleteObject") } // Open 获取下载流:Range 直接透传为 GetObject Range 头(流式,不落盘)。 func (s *S3Storage) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) { key, err := s.key(savePath) if err != nil { return nil, err } input := &s3.GetObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), } if rng != nil { input.Range = aws.String(rangeHeaderValue(rng)) } out, err := s.client.GetObject(ctx, input) if err != nil { return nil, mapS3Error(err, "GetObject") } size := aws.ToInt64(out.ContentLength) start, end := int64(0), size-1 if cr := aws.ToString(out.ContentRange); cr != "" { // 服务端按 206 返回了区间 if sr, e, total, ok := parseContentRange(cr); ok { start, end = sr, e if total >= 0 { size = total } } } return &Download{ ReadCloser: out.Body, Start: start, End: end, Total: size, Meta: FileMeta{ Size: size, ContentType: aws.ToString(out.ContentType), AcceptRanges: true, }, }, nil } // rangeHeaderValue 将 Range 结构转为 HTTP Range 头值。 func rangeHeaderValue(rng *Range) string { if rng.End < 0 { return fmt.Sprintf("bytes=%d-", rng.Start) } return fmt.Sprintf("bytes=%d-%d", rng.Start, rng.End) } // parseContentRange 解析 "bytes 0-99/1000"(total 可能为 "*")。 func parseContentRange(v string) (start, end, total int64, ok bool) { v = strings.TrimSpace(v) if !strings.HasPrefix(v, "bytes ") { return 0, 0, -1, false } parts := strings.SplitN(strings.TrimPrefix(v, "bytes "), "/", 2) if len(parts) != 2 { return 0, 0, -1, false } total = -1 if parts[1] != "*" { t, err := strconv.ParseInt(parts[1], 10, 64) if err != nil { return 0, 0, -1, false } total = t } se := strings.SplitN(parts[0], "-", 2) if len(se) != 2 { return 0, 0, -1, false } s0, err1 := strconv.ParseInt(se[0], 10, 64) e0, err2 := strconv.ParseInt(se[1], 10, 64) if err1 != nil || err2 != nil { return 0, 0, -1, false } return s0, e0, total, true } // Stat 获取对象元信息(HeadObject)。 func (s *S3Storage) Stat(ctx context.Context, savePath string) (*FileMeta, error) { key, err := s.key(savePath) if err != nil { return nil, err } out, err := s.client.HeadObject(ctx, &s3.HeadObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), }) if err != nil { return nil, mapS3Error(err, "HeadObject") } return &FileMeta{ Size: aws.ToInt64(out.ContentLength), ContentType: aws.ToString(out.ContentType), AcceptRanges: true, }, nil } // SaveChunk 保存分片对象:落临时文件获取精确长度后 PutObject(可重试)。 func (s *S3Storage) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) { if _, err := s.key(savePath); err != nil { return 0, err } key, err := s.key(ChunkPartPath(savePath, uploadID, chunkIndex)) if err != nil { return 0, err } // 落临时文件:获得精确 Content-Length 与可重放 Body(网络失败可安全重试)。 tmp, err := os.CreateTemp("", "fcb-s3-chunk-*") if err != nil { return 0, fmt.Errorf("storage/s3: 创建临时文件失败: %w", err) } tmpName := tmp.Name() defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }() size, err := io.Copy(tmp, r) if err != nil { return 0, fmt.Errorf("storage/s3: 缓存分片失败: %w", err) } if _, err := tmp.Seek(0, io.SeekStart); err != nil { return 0, fmt.Errorf("storage/s3: 回卷分片失败: %w", err) } _, err = s.client.PutObject(ctx, &s3.PutObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), Body: tmp, ContentLength: aws.Int64(size), ContentType: aws.String("application/octet-stream"), }) if err != nil { return size, mapS3Error(err, "PutObject(分片)") } return size, nil } // S3 multipart 最小分片限制:除最后一片外每片 ≥5MB,否则 Complete 返回 EntityTooSmall。 // 分片上传的分块大小由服务端配置保证(建议 ≥5MB)。 const s3MinPartSize = 5 * 1024 * 1024 // MergeChunks 用 S3 原生 multipart 流式合并: // 逐分片 GET → 边流边算哈希 → UploadPart(带精确 Content-Length)→ Complete。 // 任一步失败即 Abort 并返回错误;成功后清理分片对象。 func (s *S3Storage) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) { if total <= 0 { return 0, "", fmt.Errorf("storage/s3: 非法分片总数 %d", total) } key, err := s.key(savePath) if err != nil { return 0, "", err } chunkPrefix, err := s.key(chunkDirOf(savePath, uploadID)) if err != nil { return 0, "", err } // 单分片快速路径:直接流式 PutObject,绕过 multipart 的 5MB 限制。 if total == 1 { return s.mergeSingle(ctx, chunkPrefix+"/0.part", key, verifyHash, 0) } mpu, err := s.client.CreateMultipartUpload(ctx, &s3.CreateMultipartUploadInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), ContentType: aws.String("application/octet-stream"), }) if err != nil { return 0, "", mapS3Error(err, "CreateMultipartUpload") } _ = aws.ToString(mpu.UploadId) // S3 侧 multipart 会话 ID(Abort 时复用 mpu.UploadId) size := int64(0) totalHash := sha256.New() parts := make([]types.CompletedPart, 0, total) defer func() { // 出错时取消 multipart(避免残留分片产生存储费用)。 if len(parts) < total { _, _ = s.client.AbortMultipartUpload(ctx, &s3.AbortMultipartUploadInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), UploadId: mpu.UploadId, }) } }() for i := 0; i < total; i++ { if err := ctx.Err(); err != nil { return 0, "", err } var expected string if verifyHash != nil { expected, err = verifyHash(i) if err != nil { return 0, "", err } } getOut, err := s.client.GetObject(ctx, &s3.GetObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(fmt.Sprintf("%s/%d.part", chunkPrefix, i)), }) if err != nil { return 0, "", mapS3Error(err, fmt.Sprintf("GetObject(分片 %d)", i)) } // 分片先流式落临时文件:计算哈希 + 获得可回卷 body(SDK 签名哈希需要 seekable 流, // 同时为 UploadPart 失败重试保留数据)。 chunkHash := sha256.New() tmp, err := os.CreateTemp("", "fcb-s3-part-*") if err != nil { _ = getOut.Body.Close() return 0, "", fmt.Errorf("storage/s3: 创建分片临时文件失败: %w", err) } partLen, err := io.Copy(io.MultiWriter(tmp, totalHash, chunkHash), getOut.Body) _ = getOut.Body.Close() if err != nil { _ = tmp.Close() _ = os.Remove(tmp.Name()) return 0, "", fmt.Errorf("storage/s3: 读取分片 %d 失败: %w", i, err) } if expected != "" && expected != hex.EncodeToString(chunkHash.Sum(nil)) { _ = tmp.Close() _ = os.Remove(tmp.Name()) return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, i, expected) } if _, err := tmp.Seek(0, io.SeekStart); err != nil { _ = tmp.Close() _ = os.Remove(tmp.Name()) return 0, "", fmt.Errorf("storage/s3: 回卷分片 %d 失败: %w", i, err) } up, err := s.client.UploadPart(ctx, &s3.UploadPartInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), UploadId: mpu.UploadId, PartNumber: aws.Int32(int32(i + 1)), Body: tmp, ContentLength: aws.Int64(partLen), }) _ = tmp.Close() _ = os.Remove(tmp.Name()) if err != nil { return 0, "", mapS3Error(err, fmt.Sprintf("UploadPart(分片 %d)", i)) } parts = append(parts, types.CompletedPart{ PartNumber: aws.Int32(int32(i + 1)), ETag: up.ETag, }) size += partLen } if _, err := s.client.CompleteMultipartUpload(ctx, &s3.CompleteMultipartUploadInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), UploadId: mpu.UploadId, MultipartUpload: &types.CompletedMultipartUpload{Parts: parts}, }); err != nil { return 0, "", mapS3Error(err, "CompleteMultipartUpload") } // 合并成功后清理分片对象(静默容错)。 _ = s.CleanChunks(ctx, uploadID, savePath) return size, hex.EncodeToString(totalHash.Sum(nil)), nil } // mergeSingle 单分片合并快速路径:GET 分片 → 落临时文件校验 → PutObject 正式键。 func (s *S3Storage) mergeSingle(ctx context.Context, chunkKey, dstKey string, verifyHash func(index int) (string, error), index int) (int64, string, error) { getOut, err := s.client.GetObject(ctx, &s3.GetObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(chunkKey), }) if err != nil { return 0, "", mapS3Error(err, "GetObject(分片)") } defer func() { _ = getOut.Body.Close() }() tmp, err := os.CreateTemp("", "fcb-s3-merge-*") if err != nil { return 0, "", fmt.Errorf("storage/s3: 创建临时文件失败: %w", err) } tmpName := tmp.Name() defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }() fileHash := sha256.New() size, err := io.Copy(io.MultiWriter(tmp, fileHash), getOut.Body) if err != nil { return 0, "", fmt.Errorf("storage/s3: 读取分片失败: %w", err) } if verifyHash != nil { expected, err := verifyHash(index) if err != nil { return 0, "", err } if expected != "" && expected != hex.EncodeToString(fileHash.Sum(nil)) { return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, index, expected) } } if _, err := tmp.Seek(0, io.SeekStart); err != nil { return 0, "", fmt.Errorf("storage/s3: 回卷临时文件失败: %w", err) } if _, err := s.client.PutObject(ctx, &s3.PutObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(dstKey), Body: tmp, ContentLength: aws.Int64(size), ContentType: aws.String("application/octet-stream"), }); err != nil { return 0, "", mapS3Error(err, "PutObject(合并)") } return size, hex.EncodeToString(fileHash.Sum(nil)), nil } // CleanChunks 列举并批量删除分片对象;前缀不存在时静默成功。 func (s *S3Storage) CleanChunks(ctx context.Context, uploadID string, savePath string) error { prefix, err := s.key(chunkDirOf(savePath, uploadID)) if err != nil { return err } prefix += "/" paginator := s3.NewListObjectsV2Paginator(s.client, &s3.ListObjectsV2Input{ Bucket: aws.String(s.bucket), Prefix: aws.String(prefix), }) for paginator.HasMorePages() { page, err := paginator.NextPage(ctx) if err != nil { return mapS3Error(err, "ListObjectsV2(分片)") } if len(page.Contents) == 0 { return nil } objs := make([]types.ObjectIdentifier, 0, len(page.Contents)) for _, obj := range page.Contents { objs = append(objs, types.ObjectIdentifier{Key: obj.Key}) } if _, err := s.client.DeleteObjects(ctx, &s3.DeleteObjectsInput{ Bucket: aws.String(s.bucket), Delete: &types.Delete{Objects: objs, Quiet: aws.Bool(true)}, }); err != nil { return mapS3Error(err, "DeleteObjects(分片)") } } return nil } // FileExists HeadObject 探测存在性。 func (s *S3Storage) FileExists(ctx context.Context, savePath string) (bool, error) { key, err := s.key(savePath) if err != nil { return false, err } _, err = s.client.HeadObject(ctx, &s3.HeadObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), }) if err != nil { if isS3NotFound(err) { return false, nil } return false, mapS3Error(err, "HeadObject") } return true, nil } // PresignGetURL 生成限时下载直链。 func (s *S3Storage) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) { key, err := s.key(savePath) if err != nil { return "", err } if expires <= 0 { expires = 3600 } out, err := s.presigner.PresignGetObject(ctx, &s3.GetObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), }, s3.WithPresignExpires(time.Duration(expires)*time.Second)) if err != nil { return "", mapS3Error(err, "PresignGetObject") } return out.URL, nil } // PresignPutURL 生成限时直传 URL。 func (s *S3Storage) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) { key, err := s.key(savePath) if err != nil { return "", err } if expires <= 0 { expires = 900 } out, err := s.presigner.PresignPutObject(ctx, &s3.PutObjectInput{ Bucket: aws.String(s.bucket), Key: aws.String(key), }, s3.WithPresignExpires(time.Duration(expires)*time.Second)) if err != nil { return "", mapS3Error(err, "PresignPutObject") } return out.URL, nil } // HealthCheck 健康检查:列举 bucket(MaxKeys=1),同时校验连通性、凭据与 bucket 存在。 func (s *S3Storage) HealthCheck(ctx context.Context) error { maxKeys := int32(1) _, err := s.client.ListObjectsV2(ctx, &s3.ListObjectsV2Input{ Bucket: aws.String(s.bucket), MaxKeys: aws.Int32(maxKeys), Prefix: aws.String(""), }) if err != nil { return fmt.Errorf("%w: S3 健康检查失败: %v", ErrUnavailable, err) } return nil } // ReadHead 读取对象前 n 字节(保留的便捷封装:HeadMeta 的仅头部形态)。 func (s *S3Storage) ReadHead(ctx context.Context, savePath string, n int64) ([]byte, error) { _, head, err := s.HeadMeta(ctx, savePath, n) return head, err } // HeadMeta 读取对象元信息与头部字节(S3 引擎实现:HeadObject + Range GET)。 func (s *S3Storage) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) { key, err := s.key(savePath) if err != nil { return nil, nil, err } head, err := s.headBytes(ctx, s.bucket, key, headBytes) if err != nil { return nil, nil, err } meta, err := s.Stat(ctx, savePath) if err != nil { return nil, nil, err } return meta, head, nil } // headBytes 通过 Range GET 读取对象前 n 字节。 func (s *S3Storage) headBytes(ctx context.Context, bucket, key string, n int64) ([]byte, error) { if n <= 0 { return nil, nil } out, err := s.client.GetObject(ctx, &s3.GetObjectInput{ Bucket: aws.String(bucket), Key: aws.String(key), Range: aws.String(fmt.Sprintf("bytes=0-%d", n-1)), }) if err != nil { return nil, mapS3Error(err, "GetObject(head)") } defer func() { _ = out.Body.Close() }() return io.ReadAll(io.LimitReader(out.Body, n)) } // isS3NotFound 判断错误是否为对象不存在。 func isS3NotFound(err error) bool { var nf *types.NotFound if errors.As(err, &nf) { return true } var ae smithy.APIError if errors.As(err, &ae) { switch ae.ErrorCode() { case "NotFound", "NoSuchKey": return true } } var re *awshttp.ResponseError if errors.As(err, &re) && re.HTTPStatusCode() == http.StatusNotFound { return true } return false } // mapS3Error 将 SDK 错误映射为包内哨兵错误。 func mapS3Error(err error, op string) error { if err == nil { return nil } if isS3NotFound(err) { return fmt.Errorf("%w(%s)", ErrNotFound, op) } var re *awshttp.ResponseError if errors.As(err, &re) && re.HTTPStatusCode() == http.StatusRequestedRangeNotSatisfiable { return fmt.Errorf("%w(%s)", ErrRangeNotSatisfiable, op) } var ae smithy.APIError if errors.As(err, &ae) && ae.ErrorCode() == "InvalidRange" { return fmt.Errorf("%w(%s)", ErrRangeNotSatisfiable, op) } return fmt.Errorf("storage/s3: %s 失败: %w", op, err) } // 接口编译期断言。 var _ Storage = (*S3Storage)(nil)