Files
JuYou/API/internal/storage/cos.go
T
2026-08-25 17:59:42 +08:00

214 lines
7.0 KiB
Go

// Package storage 封装腾讯云 COS 对象存储访问,统一提供上传、读取、删除和公开地址生成能力。
package storage
import (
"context"
"fmt"
"io"
"net/http"
"os"
"path"
"strings"
"juhe-factory/api/internal/config"
"github.com/aws/aws-sdk-go-v2/aws"
awsconfig "github.com/aws/aws-sdk-go-v2/config"
"github.com/aws/aws-sdk-go-v2/credentials"
"github.com/aws/aws-sdk-go-v2/service/s3"
)
const publicMediaCacheControl = "public, max-age=31536000, immutable"
// COS 保存腾讯云 COS 客户端、存储桶、公开域名和媒体体积限制。
type COS struct {
client *s3.Client
bucket string
publicBaseURL string
maxImageBytes int64
maxVideoBytes int64
maxAudioBytes int64
}
// NewCOS 根据服务端配置初始化腾讯云 COS;完全未配置时返回 nil,部分配置时返回错误。
func NewCOS(ctx context.Context, cfg config.Config) (*COS, error) {
configured := cfg.COSEndpoint != "" || cfg.COSSecretID != "" || cfg.COSSecretKey != "" || cfg.COSBucket != "" || cfg.COSRegion != ""
if !configured {
return nil, nil
}
if cfg.COSEndpoint == "" || cfg.COSSecretID == "" || cfg.COSSecretKey == "" || cfg.COSBucket == "" || cfg.COSRegion == "" {
return nil, fmt.Errorf("COS 配置不完整,必须同时配置 Region、Endpoint、SecretId、SecretKey 和 Bucket")
}
if cfg.COSPublicBaseURL == "" {
return nil, fmt.Errorf("视频转绘需要可供模型读取的 COS_PUBLIC_BASE_URL")
}
endpoint := normalizeHTTPURL(cfg.COSEndpoint)
publicBaseURL := normalizeHTTPURL(cfg.COSPublicBaseURL)
awsCfg, err := awsconfig.LoadDefaultConfig(ctx,
awsconfig.WithRegion(cfg.COSRegion),
awsconfig.WithCredentialsProvider(credentials.NewStaticCredentialsProvider(cfg.COSSecretID, cfg.COSSecretKey, "")),
)
if err != nil {
return nil, err
}
client := s3.NewFromConfig(awsCfg, func(options *s3.Options) {
options.BaseEndpoint = aws.String(endpoint)
options.UsePathStyle = false
})
return &COS{
client: client, bucket: cfg.COSBucket, publicBaseURL: publicBaseURL,
maxImageBytes: int64(cfg.COSMaxImageSizeMB) * 1024 * 1024,
maxVideoBytes: int64(cfg.COSMaxVideoSizeMB) * 1024 * 1024,
maxAudioBytes: int64(cfg.COSMaxAudioSizeMB) * 1024 * 1024,
}, nil
}
// normalizeHTTPURL 补全缺少协议的 HTTP 地址,并保留显式配置的协议。
func normalizeHTTPURL(value string) string {
value = strings.TrimSpace(value)
if value != "" && !strings.Contains(value, "://") {
return "https://" + value
}
return value
}
// MaxVideoBytes 返回允许处理的视频最大字节数。
func (storage *COS) MaxVideoBytes() int64 {
if storage == nil {
return 0
}
return storage.maxVideoBytes
}
// MaxAudioBytes 返回允许处理的音频最大字节数。
func (storage *COS) MaxAudioBytes() int64 {
if storage == nil {
return 0
}
return storage.maxAudioBytes
}
// PublicURL 根据对象键生成无需签名的公开访问地址。
func (storage *COS) PublicURL(key string) string {
if storage == nil || storage.publicBaseURL == "" {
return ""
}
return strings.TrimRight(storage.publicBaseURL, "/") + "/" + strings.TrimLeft(key, "/")
}
// MaxImageBytes 返回允许处理的图片最大字节数。
func (storage *COS) MaxImageBytes() int64 {
if storage == nil {
return 0
}
return storage.maxImageBytes
}
// Put 上传对象并返回其公开访问地址。
func (storage *COS) Put(ctx context.Context, key, contentType string, body io.Reader, size int64) (string, error) {
if storage == nil {
return "", fmt.Errorf("COS 未配置")
}
input := &s3.PutObjectInput{
Bucket: &storage.bucket,
Key: &key,
Body: body,
ContentType: &contentType,
CacheControl: aws.String(publicMediaCacheControl),
}
if size >= 0 {
input.ContentLength = aws.Int64(size)
}
if _, err := storage.client.PutObject(ctx, input); err != nil {
return "", err
}
return storage.PublicURL(key), nil
}
// Open 打开对象读取流,并返回内容类型和内容长度;调用方负责关闭读取流。
func (storage *COS) Open(ctx context.Context, key string) (io.ReadCloser, string, int64, error) {
if storage == nil || strings.TrimSpace(key) == "" {
return nil, "", 0, fmt.Errorf("COS 未配置或对象不存在")
}
result, err := storage.client.GetObject(ctx, &s3.GetObjectInput{Bucket: &storage.bucket, Key: &key})
if err != nil {
return nil, "", 0, err
}
return result.Body, aws.ToString(result.ContentType), aws.ToInt64(result.ContentLength), nil
}
// Delete 删除指定对象;空对象键视为无需处理。
func (storage *COS) Delete(ctx context.Context, key string) error {
if storage == nil || strings.TrimSpace(key) == "" {
return nil
}
_, err := storage.client.DeleteObject(ctx, &s3.DeleteObjectInput{Bucket: &storage.bucket, Key: &key})
return err
}
// PutFromURL 下载远端结果到临时文件,校验体积后上传到 COS。
func (storage *COS) PutFromURL(ctx context.Context, client *http.Client, sourceURL, key, fallbackContentType string, maxBytes int64) (string, string, string, int64, error) {
request, err := http.NewRequestWithContext(ctx, http.MethodGet, sourceURL, nil)
if err != nil {
return "", "", "", 0, err
}
response, err := client.Do(request)
if err != nil {
return "", "", "", 0, err
}
defer response.Body.Close()
if response.StatusCode < 200 || response.StatusCode >= 300 {
return "", "", "", 0, fmt.Errorf("下载生成结果失败: HTTP %d", response.StatusCode)
}
if maxBytes > 0 && response.ContentLength > maxBytes {
return "", "", "", 0, fmt.Errorf("生成结果超过允许大小")
}
contentType := strings.TrimSpace(strings.Split(response.Header.Get("Content-Type"), ";")[0])
if contentType == "" || contentType == "application/octet-stream" {
contentType = fallbackContentType
}
key = keyWithContentType(key, contentType)
temporary, err := os.CreateTemp("", "jcf-generated-*")
if err != nil {
return "", "", "", 0, err
}
temporaryName := temporary.Name()
defer os.Remove(temporaryName)
reader := io.Reader(response.Body)
if maxBytes > 0 {
reader = io.LimitReader(response.Body, maxBytes+1)
}
written, copyErr := io.Copy(temporary, reader)
if closeErr := temporary.Close(); copyErr == nil {
copyErr = closeErr
}
if copyErr != nil {
return "", "", "", 0, copyErr
}
if maxBytes > 0 && written > maxBytes {
return "", "", "", 0, fmt.Errorf("生成结果超过允许大小")
}
file, err := os.Open(temporaryName)
if err != nil {
return "", "", "", 0, err
}
defer file.Close()
url, err := storage.Put(ctx, key, contentType, file, written)
if err != nil {
return "", "", "", 0, err
}
return url, key, contentType, written, nil
}
// keyWithContentType 根据真实媒体类型修正对象扩展名。
func keyWithContentType(key, contentType string) string {
extension := map[string]string{
"image/jpeg": ".jpg", "image/png": ".png", "image/webp": ".webp", "image/gif": ".gif",
"video/mp4": ".mp4", "video/webm": ".webm", "video/quicktime": ".mov",
}[strings.ToLower(contentType)]
if extension == "" {
return key
}
return strings.TrimSuffix(key, path.Ext(key)) + extension
}