mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 23:56:37 +08:00
feat(storage): add dynamic storage config and migration
Move storage backend configuration from startup YAML to system_config-backed runtime configuration. Add local, S3-compatible, R2, MinIO, OSS, and WebDAV backend support. Add a storage migration async task using the existing task dispatch framework. Migration target config is carried in task payload, and maintenance mode is derived from task execution state. Split upload file management and storage operations, add the admin storage configuration tab, and update migrations and Swagger docs.
This commit is contained in:
@@ -1,205 +0,0 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package storage
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/logger"
|
||||
"github.com/Rain-kl/Wavelet/internal/otel_trace"
|
||||
"go.opentelemetry.io/otel/codes"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
"golang.org/x/sync/singleflight"
|
||||
)
|
||||
|
||||
var localCacheEnabled = false
|
||||
var localCacheDir = ""
|
||||
var cacheFilePath = "%s/%s"
|
||||
var cacheMetaFilePath = "%s/%s.meta"
|
||||
var group singleflight.Group
|
||||
|
||||
type metaInfo struct {
|
||||
ContentType string `json:"content_type"`
|
||||
ContentLength int64 `json:"content_length"`
|
||||
}
|
||||
|
||||
// cacheDirPerm 缓存目录权限
|
||||
const cacheDirPerm = 0755
|
||||
|
||||
func init() {
|
||||
cfg := config.Config.S3.LocalCache
|
||||
localCacheEnabled = cfg.Enabled && cfg.CacheDir != ""
|
||||
localCacheDir = strings.TrimSuffix(cfg.CacheDir, "/")
|
||||
if localCacheEnabled {
|
||||
if err := os.MkdirAll(cfg.CacheDir, cacheDirPerm); err != nil {
|
||||
log.Fatalf("[Storage] failed to create local cache directory: %v\n", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// GetObjectViaCache 通过本地缓存获取对象,缓存未命中时从 S3/CDN 拉取
|
||||
func GetObjectViaCache(ctx context.Context, key string) (*ObjectInfo, error) {
|
||||
// 没有开启本地缓存
|
||||
if !localCacheEnabled {
|
||||
return GetObjectViaProxy(ctx, key)
|
||||
}
|
||||
|
||||
// 初始化 Trace
|
||||
ctx, span := otel_trace.Start(ctx, "S3.GetObjectViaCache", trace.WithSpanKind(trace.SpanKindClient))
|
||||
defer span.End()
|
||||
|
||||
// 检查本地缓存
|
||||
key = strings.TrimPrefix(key, "/")
|
||||
localPath := fmt.Sprintf(cacheFilePath, localCacheDir, key)
|
||||
metaPath := fmt.Sprintf(cacheMetaFilePath, localCacheDir, key)
|
||||
objInfo, err := getLocalCacheFile(ctx, localPath, metaPath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if objInfo != nil {
|
||||
return objInfo, nil
|
||||
}
|
||||
|
||||
// 使用 singleflight 确保同一时间只有一个请求会触发 CDN 获取和本地缓存保存
|
||||
_, err, _ = group.Do(key, func() (interface{}, error) {
|
||||
ctx := context.WithoutCancel(ctx)
|
||||
|
||||
// 没有缓存,通过 CDN 获取
|
||||
objInfo, err := GetObjectViaProxy(ctx, key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 保存到本地
|
||||
if err := saveToLocalCache(ctx, localPath, metaPath, objInfo); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return nil, nil
|
||||
})
|
||||
if err != nil {
|
||||
logger.ErrorF(ctx, "Failed to get object via singleflight for key %s: %v", key, err)
|
||||
return nil, LocalCacheError{}
|
||||
}
|
||||
|
||||
return GetObjectViaCache(ctx, key)
|
||||
}
|
||||
|
||||
func getLocalCacheFile(ctx context.Context, localPath, metaPath string) (*ObjectInfo, error) {
|
||||
_, span := otel_trace.Start(ctx, "S3.GetLocalCacheFile", trace.WithSpanKind(trace.SpanKindClient))
|
||||
defer span.End()
|
||||
|
||||
// 尝试打开本地缓存文件
|
||||
file, err := os.Open(localPath) //nolint:gosec // localPath is internally managed cache path
|
||||
if err == nil {
|
||||
defer func() { _ = file.Close() }()
|
||||
}
|
||||
|
||||
// 文件不存在
|
||||
if err != nil && os.IsNotExist(err) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// 判断是否为其他异常
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 读取元信息
|
||||
metaData, err := os.ReadFile(metaPath) //nolint:gosec // metaPath is internally managed cache path
|
||||
|
||||
// 文件不存在
|
||||
if err != nil && os.IsNotExist(err) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// 判断是否为其他异常
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 解析元信息
|
||||
meta := &metaInfo{}
|
||||
if err := json.Unmarshal(metaData, meta); err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &ObjectInfo{CachePath: localPath, ContentLength: meta.ContentLength, ContentType: meta.ContentType}, nil
|
||||
}
|
||||
|
||||
func saveToLocalCache(ctx context.Context, localPath, metaPath string, objInfo *ObjectInfo) error {
|
||||
_, span := otel_trace.Start(ctx, "S3.SaveToLocalCache", trace.WithSpanKind(trace.SpanKindClient))
|
||||
defer span.End()
|
||||
|
||||
// 创建目录
|
||||
localDir := filepath.Dir(localPath)
|
||||
if err := os.MkdirAll(localDir, cacheDirPerm); err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return err
|
||||
}
|
||||
|
||||
// 创建文件
|
||||
if err := saveFile(localPath, objInfo.Body); err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return err
|
||||
}
|
||||
|
||||
// 创建元信息文件
|
||||
meta := &metaInfo{ContentType: objInfo.ContentType, ContentLength: objInfo.ContentLength}
|
||||
metaData, err := json.Marshal(meta)
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return err
|
||||
}
|
||||
if err := saveFile(metaPath, bytes.NewReader(metaData)); err != nil {
|
||||
span.SetStatus(codes.Error, err.Error())
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func saveFile(localPath string, data io.Reader) error {
|
||||
// 创建临时文件
|
||||
tempFile, err := os.CreateTemp(filepath.Dir(localPath), "cache_temp_*")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = os.Remove(tempFile.Name()) }()
|
||||
|
||||
// 将内容写入临时文件
|
||||
if _, err := tempFile.ReadFrom(data); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 确保数据写入磁盘
|
||||
if err := tempFile.Sync(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 关闭临时文件
|
||||
if err := tempFile.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 将临时文件重命名为最终文件
|
||||
if err := os.Rename(tempFile.Name(), localPath); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,234 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package storage provides dynamically configured file storage backends.
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/db"
|
||||
"github.com/Rain-kl/Wavelet/internal/model"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// Driver identifies a supported storage backend.
|
||||
type Driver string
|
||||
|
||||
const (
|
||||
// DriverLocal stores files on the local filesystem.
|
||||
DriverLocal Driver = "local"
|
||||
// DriverS3 stores files in an S3-compatible object store.
|
||||
DriverS3 Driver = "s3"
|
||||
// DriverR2 stores files in Cloudflare R2.
|
||||
DriverR2 Driver = "r2"
|
||||
// DriverMinIO stores files in MinIO.
|
||||
DriverMinIO Driver = "minio"
|
||||
// DriverOSS stores files in Aliyun OSS.
|
||||
DriverOSS Driver = "oss"
|
||||
// DriverWebDAV stores files through WebDAV.
|
||||
DriverWebDAV Driver = "webdav"
|
||||
|
||||
// ConfigMask replaces secrets returned to the frontend.
|
||||
ConfigMask = "******"
|
||||
)
|
||||
|
||||
// LocalConfig configures local filesystem storage.
|
||||
type LocalConfig struct {
|
||||
Root string `json:"root"`
|
||||
}
|
||||
|
||||
// ObjectConfig configures S3-compatible or OSS object storage.
|
||||
type ObjectConfig struct {
|
||||
Endpoint string `json:"endpoint"`
|
||||
Region string `json:"region"`
|
||||
Bucket string `json:"bucket"`
|
||||
AccessKeyID string `json:"access_key_id"`
|
||||
SecretAccessKey string `json:"secret_access_key"`
|
||||
AccountID string `json:"account_id,omitempty"`
|
||||
PathStyle bool `json:"path_style"`
|
||||
KeyPrefix string `json:"key_prefix"`
|
||||
CDNURL string `json:"cdn_url"`
|
||||
}
|
||||
|
||||
// WebDAVConfig configures WebDAV storage.
|
||||
type WebDAVConfig struct {
|
||||
Endpoint string `json:"endpoint"`
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
BasePath string `json:"base_path"`
|
||||
}
|
||||
|
||||
// Config contains all storage backends and the currently active driver.
|
||||
type Config struct {
|
||||
Driver Driver `json:"driver"`
|
||||
Local LocalConfig `json:"local"`
|
||||
S3 ObjectConfig `json:"s3"`
|
||||
R2 ObjectConfig `json:"r2"`
|
||||
MinIO ObjectConfig `json:"minio"`
|
||||
OSS ObjectConfig `json:"oss"`
|
||||
WebDAV WebDAVConfig `json:"webdav"`
|
||||
}
|
||||
|
||||
// DefaultConfig returns the local-storage default configuration.
|
||||
func DefaultConfig() Config {
|
||||
return Config{
|
||||
Driver: DriverLocal,
|
||||
Local: LocalConfig{Root: "."},
|
||||
S3: ObjectConfig{Region: "us-east-1"},
|
||||
R2: ObjectConfig{Region: "auto"},
|
||||
MinIO: ObjectConfig{Region: "us-east-1", PathStyle: true},
|
||||
}
|
||||
}
|
||||
|
||||
// LoadConfig loads the active storage configuration.
|
||||
func LoadConfig(ctx context.Context) (Config, error) {
|
||||
return loadConfigByKey(ctx, model.ConfigKeyStorageConfig, DefaultConfig())
|
||||
}
|
||||
|
||||
func loadConfigByKey(ctx context.Context, key string, fallback Config) (Config, error) {
|
||||
var sc model.SystemConfig
|
||||
if err := sc.GetByKey(ctx, key); err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return fallback, nil
|
||||
}
|
||||
return Config{}, err
|
||||
}
|
||||
if strings.TrimSpace(sc.Value) == "" {
|
||||
return fallback, nil
|
||||
}
|
||||
if err := json.Unmarshal([]byte(sc.Value), &fallback); err != nil {
|
||||
return Config{}, fmt.Errorf("parse %s: %w", key, err)
|
||||
}
|
||||
return fallback, nil
|
||||
}
|
||||
|
||||
// ValidateConfig validates the selected backend configuration.
|
||||
func ValidateConfig(cfg Config) error {
|
||||
switch cfg.Driver {
|
||||
case DriverLocal:
|
||||
if strings.TrimSpace(cfg.Local.Root) == "" {
|
||||
return errors.New("local root is required")
|
||||
}
|
||||
case DriverS3:
|
||||
return validateObjectConfig(cfg.S3, false)
|
||||
case DriverR2:
|
||||
if strings.TrimSpace(cfg.R2.AccountID) == "" && strings.TrimSpace(cfg.R2.Endpoint) == "" {
|
||||
return errors.New("R2 account ID or endpoint is required")
|
||||
}
|
||||
return validateObjectConfig(cfg.R2, false)
|
||||
case DriverMinIO:
|
||||
if strings.TrimSpace(cfg.MinIO.Endpoint) == "" {
|
||||
return errors.New("MinIO endpoint is required")
|
||||
}
|
||||
return validateObjectConfig(cfg.MinIO, true)
|
||||
case DriverOSS:
|
||||
if strings.TrimSpace(cfg.OSS.Endpoint) == "" {
|
||||
return errors.New("OSS endpoint is required")
|
||||
}
|
||||
return validateObjectConfig(cfg.OSS, true)
|
||||
case DriverWebDAV:
|
||||
if strings.TrimSpace(cfg.WebDAV.Endpoint) == "" {
|
||||
return errors.New("WebDAV endpoint is required")
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("unsupported storage driver %q", cfg.Driver)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateObjectConfig(cfg ObjectConfig, endpointRequired bool) error {
|
||||
if endpointRequired && strings.TrimSpace(cfg.Endpoint) == "" {
|
||||
return errors.New("endpoint is required")
|
||||
}
|
||||
if strings.TrimSpace(cfg.Region) == "" {
|
||||
return errors.New("region is required")
|
||||
}
|
||||
if strings.TrimSpace(cfg.Bucket) == "" {
|
||||
return errors.New("bucket is required")
|
||||
}
|
||||
if strings.TrimSpace(cfg.AccessKeyID) == "" || strings.TrimSpace(cfg.SecretAccessKey) == "" {
|
||||
return errors.New("access key ID and secret access key are required")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SaveActiveConfig persists the active storage configuration.
|
||||
func SaveActiveConfig(ctx context.Context, cfg Config) error {
|
||||
return saveSystemConfig(ctx, model.ConfigKeyStorageConfig, cfg, "文件存储驱动及连接配置(JSON)")
|
||||
}
|
||||
|
||||
func saveSystemConfig(ctx context.Context, key string, value any, description string) error {
|
||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return upsertSystemConfig(ctx, tx, key, value, description)
|
||||
})
|
||||
}
|
||||
|
||||
func upsertSystemConfig(ctx context.Context, tx *gorm.DB, key string, value any, description string) error {
|
||||
data, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal %s: %w", key, err)
|
||||
}
|
||||
sc := model.SystemConfig{
|
||||
Key: key,
|
||||
Value: string(data),
|
||||
Type: "system",
|
||||
Visibility: model.ConfigVisibilityHidden,
|
||||
Description: description,
|
||||
}
|
||||
if err := tx.Where("key = ?", key).
|
||||
Assign(map[string]any{"value": sc.Value, "description": description, "visibility": model.ConfigVisibilityHidden}).
|
||||
FirstOrCreate(&sc).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
if db.Redis != nil {
|
||||
if err := db.HSetJSON(ctx, model.SystemConfigRedisHashKey, key, &sc); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// MergeMaskedSecrets restores unchanged secrets from the current configuration.
|
||||
func MergeMaskedSecrets(next, current Config) Config {
|
||||
mergeObjectSecret := func(dst *ObjectConfig, src ObjectConfig) {
|
||||
if dst.AccessKeyID == ConfigMask {
|
||||
dst.AccessKeyID = src.AccessKeyID
|
||||
}
|
||||
if dst.SecretAccessKey == ConfigMask {
|
||||
dst.SecretAccessKey = src.SecretAccessKey
|
||||
}
|
||||
}
|
||||
mergeObjectSecret(&next.S3, current.S3)
|
||||
mergeObjectSecret(&next.R2, current.R2)
|
||||
mergeObjectSecret(&next.MinIO, current.MinIO)
|
||||
mergeObjectSecret(&next.OSS, current.OSS)
|
||||
if next.WebDAV.Password == ConfigMask {
|
||||
next.WebDAV.Password = current.WebDAV.Password
|
||||
}
|
||||
return next
|
||||
}
|
||||
|
||||
// MaskSecrets replaces stored credentials with placeholders for API responses.
|
||||
func MaskSecrets(cfg Config) Config {
|
||||
maskObject := func(value *ObjectConfig) {
|
||||
if value.AccessKeyID != "" {
|
||||
value.AccessKeyID = ConfigMask
|
||||
}
|
||||
if value.SecretAccessKey != "" {
|
||||
value.SecretAccessKey = ConfigMask
|
||||
}
|
||||
}
|
||||
maskObject(&cfg.S3)
|
||||
maskObject(&cfg.R2)
|
||||
maskObject(&cfg.MinIO)
|
||||
maskObject(&cfg.OSS)
|
||||
if cfg.WebDAV.Password != "" {
|
||||
cfg.WebDAV.Password = ConfigMask
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
@@ -1,29 +0,0 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
// Package storage 提供文件存储抽象层,包括 S3 兼容存储和本地缓存。
|
||||
package storage
|
||||
|
||||
// ErrS3InitializationFailed S3 存储初始化失败错误
|
||||
type ErrS3InitializationFailed struct{}
|
||||
|
||||
func (e ErrS3InitializationFailed) Error() string {
|
||||
return errS3InitializationFailed
|
||||
}
|
||||
|
||||
// LocalCacheError 本地缓存错误
|
||||
type LocalCacheError struct{}
|
||||
|
||||
func (e LocalCacheError) Error() string {
|
||||
return errLocalCache
|
||||
}
|
||||
|
||||
const (
|
||||
errS3InitializationFailed = "S3存储初始化失败"
|
||||
errLocalCache = "本地缓存错误"
|
||||
errS3PutObjectFailed = "s3 put object failed: %w"
|
||||
errS3GetObjectFailed = "s3 get object failed: %w"
|
||||
errCDNRequestFailed = "cdn request failed: %w"
|
||||
errCDNStatusFailed = "cdn returned status %d"
|
||||
errS3DeleteObjectFailed = "s3 delete object failed: %w"
|
||||
)
|
||||
@@ -0,0 +1,39 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
)
|
||||
|
||||
func getHTTPObject(ctx context.Context, baseURL, key string) (*Object, error) {
|
||||
objectURL, err := url.JoinPath(baseURL, key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("build CDN object URL: %w", err)
|
||||
}
|
||||
request, err := http.NewRequestWithContext(ctx, http.MethodGet, objectURL, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create CDN request: %w", err)
|
||||
}
|
||||
response, err := http.DefaultClient.Do(request)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get CDN object: %w", err)
|
||||
}
|
||||
if response.StatusCode != http.StatusOK {
|
||||
_ = response.Body.Close()
|
||||
return nil, fmt.Errorf("get CDN object: unexpected status %d", response.StatusCode)
|
||||
}
|
||||
contentType := response.Header.Get("Content-Type")
|
||||
if contentType == "" {
|
||||
contentType = defaultContentType
|
||||
}
|
||||
return &Object{
|
||||
Body: response.Body,
|
||||
ContentLength: response.ContentLength,
|
||||
ContentType: contentType,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type localBackend struct {
|
||||
root string
|
||||
}
|
||||
|
||||
func newLocalBackend(cfg LocalConfig) (*localBackend, error) {
|
||||
root := filepath.Clean(cfg.Root)
|
||||
if root == "" {
|
||||
return nil, fmt.Errorf("local root is required")
|
||||
}
|
||||
return &localBackend{root: root}, nil
|
||||
}
|
||||
|
||||
func (b *localBackend) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (string, error) {
|
||||
path, err := b.path(key)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(path), storageDirPerm); err != nil {
|
||||
return "", err
|
||||
}
|
||||
file, err := os.OpenFile( //nolint:gosec // path is constrained to the configured storage root.
|
||||
path,
|
||||
os.O_CREATE|os.O_TRUNC|os.O_WRONLY,
|
||||
storageFilePerm,
|
||||
)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if _, err := io.Copy(file, body); err != nil {
|
||||
_ = file.Close()
|
||||
_ = os.Remove(path)
|
||||
return "", err
|
||||
}
|
||||
if err := file.Close(); err != nil {
|
||||
_ = os.Remove(path)
|
||||
return "", err
|
||||
}
|
||||
return filepath.ToSlash(key), nil
|
||||
}
|
||||
|
||||
func (b *localBackend) Get(_ context.Context, key string) (*Object, error) {
|
||||
path, err := b.path(key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
file, err := os.Open(path) //nolint:gosec // path is constrained to the configured storage root.
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
info, err := file.Stat()
|
||||
if err != nil {
|
||||
_ = file.Close()
|
||||
return nil, err
|
||||
}
|
||||
contentType := mime.TypeByExtension(filepath.Ext(path))
|
||||
if contentType == "" {
|
||||
contentType = defaultContentType
|
||||
}
|
||||
return &Object{Body: file, ContentLength: info.Size(), ContentType: contentType}, nil
|
||||
}
|
||||
|
||||
func (b *localBackend) Delete(_ context.Context, key string) error {
|
||||
path, err := b.path(key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = os.Remove(path)
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (b *localBackend) Test(_ context.Context) error {
|
||||
return os.MkdirAll(b.root, storageDirPerm)
|
||||
}
|
||||
|
||||
func (b *localBackend) path(key string) (string, error) {
|
||||
if filepath.IsAbs(key) {
|
||||
cleanPath := filepath.Clean(key)
|
||||
absRoot, err := filepath.Abs(b.root)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
absPath, err := filepath.Abs(cleanPath)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
rel, err := filepath.Rel(absRoot, absPath)
|
||||
if err != nil || strings.HasPrefix(rel, "..") {
|
||||
return "", fmt.Errorf("storage key escapes local root")
|
||||
}
|
||||
return cleanPath, nil
|
||||
}
|
||||
cleanKey := filepath.Clean(filepath.FromSlash(strings.TrimPrefix(key, "/")))
|
||||
if cleanKey == "." || cleanKey == "" || strings.HasPrefix(cleanKey, "..") {
|
||||
return "", fmt.Errorf("invalid local storage key %q", key)
|
||||
}
|
||||
path := filepath.Join(b.root, cleanKey)
|
||||
rel, err := filepath.Rel(b.root, path)
|
||||
if err != nil || strings.HasPrefix(rel, "..") {
|
||||
return "", fmt.Errorf("storage key escapes local root")
|
||||
}
|
||||
return path, nil
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package storage
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestLocalBackendRoundTrip(t *testing.T) {
|
||||
backend, err := newLocalBackend(LocalConfig{Root: t.TempDir()})
|
||||
if err != nil {
|
||||
t.Fatalf("newLocalBackend() returned error: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
const key = "uploads/2026/06/13/test.txt"
|
||||
const content = "wavelet storage"
|
||||
|
||||
storedKey, err := backend.Put(ctx, key, bytes.NewBufferString(content), int64(len(content)), "text/plain")
|
||||
if err != nil {
|
||||
t.Fatalf("Put(%q) returned error: %v", key, err)
|
||||
}
|
||||
if storedKey != key {
|
||||
t.Errorf("Put(%q) key = %q, want %q", key, storedKey, key)
|
||||
}
|
||||
|
||||
object, err := backend.Get(ctx, key)
|
||||
if err != nil {
|
||||
t.Fatalf("Get(%q) returned error: %v", key, err)
|
||||
}
|
||||
got, err := io.ReadAll(object.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadAll(Get(%q)) returned error: %v", key, err)
|
||||
}
|
||||
if err := object.Body.Close(); err != nil {
|
||||
t.Fatalf("Close(Get(%q)) returned error: %v", key, err)
|
||||
}
|
||||
if string(got) != content {
|
||||
t.Errorf("Get(%q) content = %q, want %q", key, got, content)
|
||||
}
|
||||
|
||||
if err := backend.Delete(ctx, key); err != nil {
|
||||
t.Fatalf("Delete(%q) returned error: %v", key, err)
|
||||
}
|
||||
if _, err := backend.Get(ctx, key); err == nil {
|
||||
t.Errorf("Get(%q) after Delete() returned nil error", key)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"github.com/aliyun/alibabacloud-oss-go-sdk-v2/oss"
|
||||
"github.com/aliyun/alibabacloud-oss-go-sdk-v2/oss/credentials"
|
||||
)
|
||||
|
||||
type ossBackend struct {
|
||||
client *oss.Client
|
||||
bucket string
|
||||
keyPrefix string
|
||||
cdnURL string
|
||||
}
|
||||
|
||||
func newOSSBackend(cfg ObjectConfig) (*ossBackend, error) {
|
||||
options := oss.LoadDefaultConfig().
|
||||
WithCredentialsProvider(credentials.NewStaticCredentialsProvider(cfg.AccessKeyID, cfg.SecretAccessKey)).
|
||||
WithRegion(cfg.Region)
|
||||
if cfg.Endpoint != "" {
|
||||
options.WithEndpoint(cfg.Endpoint)
|
||||
}
|
||||
return &ossBackend{
|
||||
client: oss.NewClient(options),
|
||||
bucket: cfg.Bucket,
|
||||
keyPrefix: strings.Trim(cfg.KeyPrefix, "/"),
|
||||
cdnURL: strings.TrimRight(cfg.CDNURL, "/"),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *ossBackend) Put(ctx context.Context, key string, body io.Reader, _ int64, _ string) (string, error) {
|
||||
key = b.key(key)
|
||||
_, err := b.client.PutObject(ctx, &oss.PutObjectRequest{
|
||||
Bucket: oss.Ptr(b.bucket),
|
||||
Key: oss.Ptr(key),
|
||||
Body: body,
|
||||
})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("put OSS object: %w", err)
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func (b *ossBackend) Get(ctx context.Context, key string) (*Object, error) {
|
||||
key = b.key(key)
|
||||
if b.cdnURL != "" {
|
||||
return getHTTPObject(ctx, b.cdnURL, key)
|
||||
}
|
||||
output, err := b.client.GetObject(ctx, &oss.GetObjectRequest{
|
||||
Bucket: oss.Ptr(b.bucket),
|
||||
Key: oss.Ptr(key),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get OSS object: %w", err)
|
||||
}
|
||||
contentType := defaultContentType
|
||||
if output.ContentType != nil {
|
||||
contentType = *output.ContentType
|
||||
}
|
||||
return &Object{Body: output.Body, ContentLength: output.ContentLength, ContentType: contentType}, nil
|
||||
}
|
||||
|
||||
func (b *ossBackend) Delete(ctx context.Context, key string) error {
|
||||
_, err := b.client.DeleteObject(ctx, &oss.DeleteObjectRequest{
|
||||
Bucket: oss.Ptr(b.bucket),
|
||||
Key: oss.Ptr(b.key(key)),
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("delete OSS object: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *ossBackend) Test(ctx context.Context) error {
|
||||
ok, err := b.client.IsBucketExist(ctx, b.bucket)
|
||||
if err != nil {
|
||||
return fmt.Errorf("access OSS bucket: %w", err)
|
||||
}
|
||||
if !ok {
|
||||
return fmt.Errorf("OSS bucket %q does not exist", b.bucket)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *ossBackend) key(key string) string {
|
||||
key = strings.TrimLeft(key, "/")
|
||||
if b.keyPrefix == "" || strings.HasPrefix(key, b.keyPrefix+"/") {
|
||||
return key
|
||||
}
|
||||
return b.keyPrefix + "/" + key
|
||||
}
|
||||
+63
-201
@@ -1,4 +1,3 @@
|
||||
// Copyright 2025 linux.do
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
@@ -8,253 +7,116 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/Rain-kl/Wavelet/internal/config"
|
||||
"github.com/Rain-kl/Wavelet/internal/otel_trace"
|
||||
"github.com/Rain-kl/Wavelet/internal/util"
|
||||
"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"
|
||||
"go.opentelemetry.io/otel/attribute"
|
||||
"go.opentelemetry.io/otel/codes"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
)
|
||||
|
||||
var (
|
||||
type s3Backend struct {
|
||||
client *s3.Client
|
||||
bucket string
|
||||
keyPrefix string
|
||||
cdnURL string
|
||||
)
|
||||
}
|
||||
|
||||
func init() {
|
||||
cfg := config.Config.S3
|
||||
if !cfg.Enabled {
|
||||
log.Println("[Storage] S3 storage disabled")
|
||||
return
|
||||
}
|
||||
|
||||
bucket = cfg.Bucket
|
||||
keyPrefix = cfg.KeyPrefix
|
||||
cdnURL = strings.TrimRight(cfg.CdnURL, "/")
|
||||
|
||||
awsCfg, err := awsconfig.LoadDefaultConfig(context.Background(),
|
||||
func newS3Backend(ctx context.Context, cfg ObjectConfig) (*s3Backend, error) {
|
||||
awsCfg, err := awsconfig.LoadDefaultConfig(ctx,
|
||||
awsconfig.WithRegion(cfg.Region),
|
||||
awsconfig.WithCredentialsProvider(
|
||||
credentials.NewStaticCredentialsProvider(cfg.AccessKeyID, cfg.SecretAccessKey, ""),
|
||||
),
|
||||
awsconfig.WithCredentialsProvider(credentials.NewStaticCredentialsProvider(
|
||||
cfg.AccessKeyID,
|
||||
cfg.SecretAccessKey,
|
||||
"",
|
||||
)),
|
||||
)
|
||||
if err != nil {
|
||||
log.Fatalf("[Storage] failed to load AWS config: %v\n", err)
|
||||
return nil, fmt.Errorf("load S3 config: %w", err)
|
||||
}
|
||||
|
||||
client = s3.NewFromConfig(awsCfg, func(o *s3.Options) {
|
||||
client := s3.NewFromConfig(awsCfg, func(options *s3.Options) {
|
||||
if cfg.Endpoint != "" {
|
||||
o.BaseEndpoint = aws.String(cfg.Endpoint)
|
||||
options.BaseEndpoint = aws.String(strings.TrimRight(cfg.Endpoint, "/"))
|
||||
}
|
||||
o.UsePathStyle = cfg.PathStyle
|
||||
options.UsePathStyle = cfg.PathStyle
|
||||
})
|
||||
|
||||
log.Printf("[Storage] S3 storage initialized (bucket: %s, prefix: %s, cdn: %s)\n", bucket, keyPrefix, cdnURL)
|
||||
return &s3Backend{
|
||||
client: client,
|
||||
bucket: cfg.Bucket,
|
||||
keyPrefix: strings.Trim(cfg.KeyPrefix, "/"),
|
||||
cdnURL: strings.TrimRight(cfg.CDNURL, "/"),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// IsEnabledFunc 检查 S3 存储是否已初始化(可替换用于测试)
|
||||
var IsEnabledFunc = func() bool {
|
||||
return client != nil
|
||||
}
|
||||
|
||||
// IsEnabled 检查 S3 存储是否可用
|
||||
func IsEnabled() bool {
|
||||
return IsEnabledFunc()
|
||||
}
|
||||
|
||||
// BuildKey constructs a full S3 object key with the configured prefix.
|
||||
func BuildKey(path string) string {
|
||||
return keyPrefix + path
|
||||
}
|
||||
|
||||
var (
|
||||
// PutObjectFunc enables mocking S3 uploads in tests.
|
||||
PutObjectFunc = putObjectDefault
|
||||
// GetObjectFunc enables mocking S3 downloads in tests.
|
||||
GetObjectFunc = getObjectDefault
|
||||
// DeleteObjectFunc enables mocking S3 deletion in tests.
|
||||
DeleteObjectFunc = deleteObjectDefault
|
||||
)
|
||||
|
||||
// MockStorage is a test helper to mock S3 storage operations.
|
||||
// It returns a function that restores original implementations.
|
||||
func MockStorage(
|
||||
mockPut func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error,
|
||||
mockGet func(ctx context.Context, key string) (*ObjectInfo, error),
|
||||
mockDelete func(ctx context.Context, key string) error,
|
||||
) func() {
|
||||
origPut, origGet, origDelete := PutObjectFunc, GetObjectFunc, DeleteObjectFunc
|
||||
PutObjectFunc = mockPut
|
||||
GetObjectFunc = mockGet
|
||||
DeleteObjectFunc = mockDelete
|
||||
return func() {
|
||||
PutObjectFunc = origPut
|
||||
GetObjectFunc = origGet
|
||||
DeleteObjectFunc = origDelete
|
||||
func newR2Backend(ctx context.Context, cfg ObjectConfig) (*s3Backend, error) {
|
||||
if cfg.Endpoint == "" {
|
||||
cfg.Endpoint = fmt.Sprintf("https://%s.r2.cloudflarestorage.com", cfg.AccountID)
|
||||
}
|
||||
cfg.Region = "auto"
|
||||
return newS3Backend(ctx, cfg)
|
||||
}
|
||||
|
||||
// PutObject uploads a file to S3.
|
||||
func PutObject(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
||||
return PutObjectFunc(ctx, key, body, size, contentType)
|
||||
}
|
||||
|
||||
func putObjectDefault(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
|
||||
ctx, span := otel_trace.Start(ctx, "S3.PutObject", trace.WithSpanKind(trace.SpanKindClient))
|
||||
defer span.End()
|
||||
|
||||
span.SetAttributes(
|
||||
attribute.String("s3.key", key),
|
||||
attribute.Int64("s3.content_length", size),
|
||||
attribute.String("s3.content_type", contentType),
|
||||
)
|
||||
|
||||
if !IsEnabled() {
|
||||
span.SetStatus(codes.Error, "S3 not initialized")
|
||||
return ErrS3InitializationFailed{}
|
||||
}
|
||||
|
||||
input := &s3.PutObjectInput{
|
||||
Bucket: aws.String(bucket),
|
||||
func (b *s3Backend) Put(ctx context.Context, key string, body io.Reader, size int64, contentType string) (string, error) {
|
||||
key = b.key(key)
|
||||
_, err := b.client.PutObject(ctx, &s3.PutObjectInput{
|
||||
Bucket: aws.String(b.bucket),
|
||||
Key: aws.String(key),
|
||||
Body: body,
|
||||
ContentLength: aws.Int64(size),
|
||||
ContentType: aws.String(contentType),
|
||||
}
|
||||
|
||||
_, err := client.PutObject(ctx, input)
|
||||
})
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, fmt.Sprintf("S3 put object failed: %v", err))
|
||||
return fmt.Errorf(errS3PutObjectFailed, err)
|
||||
return "", fmt.Errorf("put S3 object: %w", err)
|
||||
}
|
||||
return nil
|
||||
return key, nil
|
||||
}
|
||||
|
||||
// ObjectInfo holds metadata about a retrieved object.
|
||||
type ObjectInfo struct {
|
||||
CachePath string
|
||||
Body io.ReadCloser
|
||||
ContentLength int64
|
||||
ContentType string
|
||||
}
|
||||
|
||||
// GetObject retrieves a file directly from S3.
|
||||
func GetObject(ctx context.Context, key string) (*ObjectInfo, error) {
|
||||
return GetObjectFunc(ctx, key)
|
||||
}
|
||||
|
||||
func getObjectDefault(ctx context.Context, key string) (*ObjectInfo, error) {
|
||||
ctx, span := otel_trace.Start(ctx, "S3.GetObject", trace.WithSpanKind(trace.SpanKindClient))
|
||||
defer span.End()
|
||||
|
||||
span.SetAttributes(attribute.String("s3.key", key))
|
||||
|
||||
if !IsEnabled() {
|
||||
span.SetStatus(codes.Error, "S3 not initialized")
|
||||
return nil, ErrS3InitializationFailed{}
|
||||
func (b *s3Backend) Get(ctx context.Context, key string) (*Object, error) {
|
||||
key = b.key(key)
|
||||
if b.cdnURL != "" {
|
||||
return getHTTPObject(ctx, b.cdnURL, key)
|
||||
}
|
||||
|
||||
output, err := client.GetObject(ctx, &s3.GetObjectInput{
|
||||
Bucket: aws.String(bucket),
|
||||
output, err := b.client.GetObject(ctx, &s3.GetObjectInput{
|
||||
Bucket: aws.String(b.bucket),
|
||||
Key: aws.String(key),
|
||||
})
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, fmt.Sprintf("S3 get object failed: %v", err))
|
||||
return nil, fmt.Errorf(errS3GetObjectFailed, err)
|
||||
return nil, fmt.Errorf("get S3 object: %w", err)
|
||||
}
|
||||
|
||||
contentType := "application/octet-stream"
|
||||
contentType := defaultContentType
|
||||
if output.ContentType != nil {
|
||||
contentType = *output.ContentType
|
||||
}
|
||||
|
||||
var contentLength int64
|
||||
var size int64
|
||||
if output.ContentLength != nil {
|
||||
contentLength = *output.ContentLength
|
||||
size = *output.ContentLength
|
||||
}
|
||||
|
||||
return &ObjectInfo{
|
||||
Body: output.Body,
|
||||
ContentLength: contentLength,
|
||||
ContentType: contentType,
|
||||
}, nil
|
||||
return &Object{Body: output.Body, ContentLength: size, ContentType: contentType}, nil
|
||||
}
|
||||
|
||||
// GetObjectViaProxy retrieves a file via CDN if configured, otherwise falls back to S3.
|
||||
func GetObjectViaProxy(ctx context.Context, key string) (*ObjectInfo, error) {
|
||||
ctx, span := otel_trace.Start(ctx, "S3.GetObjectViaProxy", trace.WithSpanKind(trace.SpanKindClient))
|
||||
defer span.End()
|
||||
|
||||
span.SetAttributes(attribute.String("s3.key", key))
|
||||
|
||||
if !IsEnabled() {
|
||||
span.SetStatus(codes.Error, "S3 not initialized")
|
||||
return nil, ErrS3InitializationFailed{}
|
||||
}
|
||||
|
||||
if cdnURL == "" {
|
||||
return GetObject(ctx, key)
|
||||
}
|
||||
|
||||
url := cdnURL + "/" + key
|
||||
span.SetAttributes(attribute.Bool("s3.use_cdn", true))
|
||||
|
||||
resp, err := util.Request(ctx, http.MethodGet, url, nil, nil, nil)
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, fmt.Sprintf("cdn request failed: %v", err))
|
||||
return nil, fmt.Errorf(errCDNRequestFailed, err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
_ = resp.Body.Close()
|
||||
span.SetStatus(codes.Error, fmt.Sprintf("cdn returned status %d", resp.StatusCode))
|
||||
return nil, fmt.Errorf(errCDNStatusFailed, resp.StatusCode)
|
||||
}
|
||||
|
||||
contentType := resp.Header.Get("Content-Type")
|
||||
if contentType == "" {
|
||||
contentType = "application/octet-stream"
|
||||
}
|
||||
|
||||
return &ObjectInfo{
|
||||
Body: resp.Body,
|
||||
ContentLength: resp.ContentLength,
|
||||
ContentType: contentType,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// DeleteObject deletes a file from S3.
|
||||
func DeleteObject(ctx context.Context, key string) error {
|
||||
return DeleteObjectFunc(ctx, key)
|
||||
}
|
||||
|
||||
func deleteObjectDefault(ctx context.Context, key string) error {
|
||||
ctx, span := otel_trace.Start(ctx, "S3.DeleteObject", trace.WithSpanKind(trace.SpanKindClient))
|
||||
defer span.End()
|
||||
|
||||
span.SetAttributes(attribute.String("s3.key", key))
|
||||
|
||||
if !IsEnabled() {
|
||||
return ErrS3InitializationFailed{}
|
||||
}
|
||||
|
||||
_, err := client.DeleteObject(ctx, &s3.DeleteObjectInput{
|
||||
Bucket: aws.String(bucket),
|
||||
Key: aws.String(key),
|
||||
func (b *s3Backend) Delete(ctx context.Context, key string) error {
|
||||
_, err := b.client.DeleteObject(ctx, &s3.DeleteObjectInput{
|
||||
Bucket: aws.String(b.bucket),
|
||||
Key: aws.String(b.key(key)),
|
||||
})
|
||||
if err != nil {
|
||||
span.SetStatus(codes.Error, fmt.Sprintf("S3 delete object failed: %v", err))
|
||||
return fmt.Errorf(errS3DeleteObjectFailed, err)
|
||||
return fmt.Errorf("delete S3 object: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *s3Backend) Test(ctx context.Context) error {
|
||||
_, err := b.client.HeadBucket(ctx, &s3.HeadBucketInput{Bucket: aws.String(b.bucket)})
|
||||
if err != nil {
|
||||
return fmt.Errorf("access S3 bucket: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *s3Backend) key(key string) string {
|
||||
key = strings.TrimLeft(key, "/")
|
||||
if b.keyPrefix == "" || strings.HasPrefix(key, b.keyPrefix+"/") {
|
||||
return key
|
||||
}
|
||||
return b.keyPrefix + "/" + key
|
||||
}
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultContentType = "application/octet-stream"
|
||||
storageDirPerm = 0o750
|
||||
storageFilePerm = 0o600
|
||||
)
|
||||
|
||||
// Object describes a readable stored object.
|
||||
type Object struct {
|
||||
CachePath string
|
||||
Body io.ReadCloser
|
||||
ContentLength int64
|
||||
ContentType string
|
||||
}
|
||||
|
||||
// Backend defines storage operations used by the upload domain.
|
||||
type Backend interface {
|
||||
Put(ctx context.Context, key string, body io.Reader, size int64, contentType string) (string, error)
|
||||
Get(ctx context.Context, key string) (*Object, error)
|
||||
Delete(ctx context.Context, key string) error
|
||||
Test(ctx context.Context) error
|
||||
}
|
||||
|
||||
var (
|
||||
// IsEnabledFunc preserves the legacy S3 test hook while tests migrate to backend injection.
|
||||
IsEnabledFunc = func() bool { return false }
|
||||
mockBackend Backend
|
||||
)
|
||||
|
||||
// Active returns the configured active driver and backend.
|
||||
func Active(ctx context.Context) (Driver, Backend, error) {
|
||||
if IsEnabledFunc() && mockBackend != nil {
|
||||
return DriverS3, mockBackend, nil
|
||||
}
|
||||
cfg, err := LoadConfig(ctx)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
backend, err := NewBackend(ctx, cfg, cfg.Driver)
|
||||
return cfg.Driver, backend, err
|
||||
}
|
||||
|
||||
// ForDriver returns the active or pending backend for an upload record.
|
||||
func ForDriver(ctx context.Context, driver Driver) (Backend, error) {
|
||||
if driver == DriverS3 && mockBackend != nil {
|
||||
return mockBackend, nil
|
||||
}
|
||||
cfg, err := LoadConfig(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cfg.Driver == driver {
|
||||
return NewBackend(ctx, cfg, driver)
|
||||
}
|
||||
return nil, fmt.Errorf("storage configuration for driver %q is unavailable", driver)
|
||||
}
|
||||
|
||||
type functionBackend struct {
|
||||
put func(context.Context, string, io.Reader, int64, string) error
|
||||
get func(context.Context, string) (*Object, error)
|
||||
delete func(context.Context, string) error
|
||||
}
|
||||
|
||||
func (b *functionBackend) Put(ctx context.Context, key string, body io.Reader, size int64, contentType string) (string, error) {
|
||||
if err := b.put(ctx, key, body, size, contentType); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func (b *functionBackend) Get(ctx context.Context, key string) (*Object, error) {
|
||||
return b.get(ctx, key)
|
||||
}
|
||||
|
||||
func (b *functionBackend) Delete(ctx context.Context, key string) error {
|
||||
return b.delete(ctx, key)
|
||||
}
|
||||
|
||||
func (b *functionBackend) Test(context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// MockStorage replaces object operations for package tests and returns a restore function.
|
||||
func MockStorage(
|
||||
put func(context.Context, string, io.Reader, int64, string) error,
|
||||
get func(context.Context, string) (*Object, error),
|
||||
deleteObject func(context.Context, string) error,
|
||||
) func() {
|
||||
previous := mockBackend
|
||||
mockBackend = &functionBackend{put: put, get: get, delete: deleteObject}
|
||||
return func() {
|
||||
mockBackend = previous
|
||||
}
|
||||
}
|
||||
|
||||
// NewBackend constructs a concrete backend from configuration.
|
||||
func NewBackend(ctx context.Context, cfg Config, driver Driver) (Backend, error) {
|
||||
if driver == DriverS3 && mockBackend != nil {
|
||||
return mockBackend, nil
|
||||
}
|
||||
switch driver {
|
||||
case DriverLocal:
|
||||
return newLocalBackend(cfg.Local)
|
||||
case DriverS3:
|
||||
return newS3Backend(ctx, cfg.S3)
|
||||
case DriverR2:
|
||||
return newR2Backend(ctx, cfg.R2)
|
||||
case DriverMinIO:
|
||||
return newS3Backend(ctx, cfg.MinIO)
|
||||
case DriverOSS:
|
||||
return newOSSBackend(cfg.OSS)
|
||||
case DriverWebDAV:
|
||||
return newWebDAVBackend(cfg.WebDAV)
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported storage driver %q", driver)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
// Copyright 2026 Arctel.net
|
||||
// SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"path"
|
||||
"strings"
|
||||
|
||||
"github.com/studio-b12/gowebdav"
|
||||
)
|
||||
|
||||
type webDAVBackend struct {
|
||||
client *gowebdav.Client
|
||||
basePath string
|
||||
}
|
||||
|
||||
func newWebDAVBackend(cfg WebDAVConfig) (*webDAVBackend, error) {
|
||||
return &webDAVBackend{
|
||||
client: gowebdav.NewClient(strings.TrimRight(cfg.Endpoint, "/"), cfg.Username, cfg.Password),
|
||||
basePath: strings.Trim(cfg.BasePath, "/"),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (b *webDAVBackend) Put(_ context.Context, key string, body io.Reader, size int64, _ string) (string, error) {
|
||||
key = b.key(key)
|
||||
if dir := path.Dir(key); dir != "." && dir != "/" {
|
||||
if err := b.client.MkdirAll(dir, storageDirPerm); err != nil {
|
||||
return "", fmt.Errorf("create WebDAV directory: %w", err)
|
||||
}
|
||||
}
|
||||
if err := b.client.WriteStreamWithLength(key, body, size, storageFilePerm); err != nil {
|
||||
return "", fmt.Errorf("put WebDAV object: %w", err)
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func (b *webDAVBackend) Get(_ context.Context, key string) (*Object, error) {
|
||||
key = b.key(key)
|
||||
info, err := b.client.Stat(key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("stat WebDAV object: %w", err)
|
||||
}
|
||||
body, err := b.client.ReadStream(key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get WebDAV object: %w", err)
|
||||
}
|
||||
contentType := defaultContentType
|
||||
if typed, ok := info.(interface{ ContentType() string }); ok && typed.ContentType() != "" {
|
||||
contentType = typed.ContentType()
|
||||
}
|
||||
return &Object{Body: body, ContentLength: info.Size(), ContentType: contentType}, nil
|
||||
}
|
||||
|
||||
func (b *webDAVBackend) Delete(_ context.Context, key string) error {
|
||||
if err := b.client.Remove(b.key(key)); err != nil {
|
||||
return fmt.Errorf("delete WebDAV object: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *webDAVBackend) Test(_ context.Context) error {
|
||||
if err := b.client.Connect(); err != nil {
|
||||
return fmt.Errorf("connect WebDAV: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (b *webDAVBackend) key(key string) string {
|
||||
return "/" + path.Join(b.basePath, strings.TrimLeft(key, "/"))
|
||||
}
|
||||
Reference in New Issue
Block a user