26.9(安全审计修复版)
Go 1.27.1 (Gin+GORM) + Vue 3 文件快传服务: - 安全审计全部修复(docs/security-audit-2026-09-05.md): bcrypt 密码哈希与自动升级、presign 直传服务端大小/内容校验、 全局请求体上限、依赖升级(govulncheck 0 命中)、janitor 后台清理、 管理端审计动作落库、/admin CORS 收紧、通知内容白名单净化、 会话默认 7 天、限流缓存故障降级、robots.txt 端点等 - 前端:取件链接复制修复(不再重复拼接提取码)、markdown 净化器加固 - Redis 支持库号(FCB_REDIS_DB / redis://…/db URL) - 文档:docs/api/* 与 openapi.yaml 同步最新行为(robots.txt、 提码 5 位起、chunk 32MiB 上限、admin 审计动作等) 验证:gofmt/go vet/go test 全绿;二进制端到端冒烟通过
This commit is contained in:
@@ -0,0 +1,10 @@
|
||||
package storage
|
||||
|
||||
import "errors"
|
||||
|
||||
// ErrRangeNotSatisfiable 请求的字节范围超出文件大小(对应 HTTP 416)。
|
||||
// 属于增量错误定义,不改动 interface.go 的既有签名。
|
||||
var ErrRangeNotSatisfiable = errors.New("storage: 请求范围超出文件大小")
|
||||
|
||||
// ErrHashMismatch 分片哈希校验失败(合并分片时检测到数据损坏)。
|
||||
var ErrHashMismatch = errors.New("storage: 分片哈希校验失败")
|
||||
@@ -0,0 +1,26 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// Factory 按配置构造存储引擎。由 go-storage 提供 New* 实现后接入。
|
||||
// 这里提供注册表模式:各引擎实现注册自己的构造函数,main.go 按名称选择。
|
||||
type Factory func(ctx context.Context) (Storage, error)
|
||||
|
||||
var registry = map[string]Factory{}
|
||||
|
||||
// RegisterEngine 注册引擎构造函数(init 时调用,名称:local|s3|webdav)。
|
||||
func RegisterEngine(name string, f Factory) {
|
||||
registry[name] = f
|
||||
}
|
||||
|
||||
// NewEngine 按名称构造引擎。
|
||||
func NewEngine(ctx context.Context, name string) (Storage, error) {
|
||||
f, ok := registry[name]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("storage: 未知存储引擎 %q(仅支持 local|s3|webdav)", name)
|
||||
}
|
||||
return f(ctx)
|
||||
}
|
||||
@@ -0,0 +1,111 @@
|
||||
// Package storage 定义存储引擎统一契约。
|
||||
//
|
||||
// 本文件是 go-storage 并行开发的接口契约:签名一经定义不再改动。
|
||||
// 三种引擎(local/s3/webdav)都要实现该接口;工厂按 FCB_STORAGE_ENGINE 选择。
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
)
|
||||
|
||||
// 错误定义:实现方应返回这些哨兵错误(可用 %w 包装),便于 API 层映射 HTTP 状态码。
|
||||
var (
|
||||
// ErrNotFound 文件不存在(HTTP 404)。
|
||||
ErrNotFound = errors.New("storage: 文件不存在")
|
||||
// ErrInvalidPath 非法路径(路径穿越等,HTTP 400)。
|
||||
ErrInvalidPath = errors.New("storage: 非法文件路径")
|
||||
// ErrUnavailable 存储服务不可用(连接失败等,HTTP 503)。
|
||||
ErrUnavailable = errors.New("storage: 存储服务不可用")
|
||||
)
|
||||
|
||||
// FileMeta 文件元信息(大小等)。
|
||||
type FileMeta struct {
|
||||
Size int64 // 字节数
|
||||
ContentType string // MIME 类型,可为空
|
||||
AcceptRanges bool // 是否支持 Range 请求
|
||||
}
|
||||
|
||||
// Download 流式下载句柄。调用方负责 Close。
|
||||
type Download struct {
|
||||
// ReadCloser 文件内容流(已按 Range 重定位)。
|
||||
io.ReadCloser
|
||||
// Meta 文件元信息。
|
||||
Meta FileMeta
|
||||
// Start 当前流的起始字节偏移(Range 请求时为 rangeStart)。
|
||||
Start int64
|
||||
// End 流的结束字节偏移(含);未知为 -1。
|
||||
End int64
|
||||
// Total 文件总大小(字节);未知为 -1。
|
||||
Total int64
|
||||
}
|
||||
|
||||
// Range 字节范围(对齐 HTTP Range 语义)。
|
||||
// nil 指针表示完整文件。
|
||||
type Range struct {
|
||||
Start int64 // 起始字节(含)
|
||||
End int64 // 结束字节(含);-1 表示到文件末尾
|
||||
}
|
||||
|
||||
// Storage 存储引擎统一接口。
|
||||
//
|
||||
// 约定:
|
||||
// - savePath 为存储侧相对路径(引擎内部负责安全解析,拒绝 .. 穿越);
|
||||
// - 所有方法必须是并发安全的;
|
||||
// - 实现方遇到不可恢复错误时返回本包哨兵错误(或用 %w 包装)。
|
||||
type Storage interface {
|
||||
// SaveFile 流式保存文件:r 读取到 EOF 即完成,返回实际写入字节数。
|
||||
// 引擎必须按 256KB 级别分块读取,不得将整个文件读入内存。
|
||||
SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error)
|
||||
|
||||
// DeleteFile 删除文件;文件不存在时返回 ErrNotFound 或 nil 均可接受。
|
||||
DeleteFile(ctx context.Context, savePath string) error
|
||||
|
||||
// Open 以下载模式打开文件,支持 HTTP Range 请求语义:
|
||||
// - rng 为 nil:返回完整文件流(Start=0,End=Total-1);
|
||||
// - rng 非 nil:返回 [Start, End] 区间流。
|
||||
// 引擎应尽量透传 Range(WebDAV/S3)或按块 seek(local)。
|
||||
Open(ctx context.Context, savePath string, rng *Range) (*Download, error)
|
||||
|
||||
// Stat 获取文件元信息;不存在返回 ErrNotFound。
|
||||
Stat(ctx context.Context, savePath string) (*FileMeta, error)
|
||||
|
||||
// SaveChunk 保存一个分片到临时区(upload_id 隔离),返回分片字节数。
|
||||
SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error)
|
||||
|
||||
// MergeChunks 按索引 0..total-1 有序合并分片并落为正式文件。
|
||||
// verifyHash 为nil 时不校验;否则为分片 SHA256 校验函数(输入索引,输出期望哈希,空串表示跳过)。
|
||||
// 返回 (最终文件大小, 整个文件 SHA256)。
|
||||
MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error)
|
||||
|
||||
// CleanChunks 清理分片临时区;不存在时静默成功。
|
||||
CleanChunks(ctx context.Context, uploadID string, savePath string) error
|
||||
|
||||
// FileExists 检查文件是否存在。
|
||||
FileExists(ctx context.Context, savePath string) (bool, error)
|
||||
|
||||
// HeadMeta 读取对象元信息与头部字节(可选能力,供直传 confirm 校验实际
|
||||
// 大小与内容;不支持时返回 ErrNotSupported)。
|
||||
// meta 允许为 nil(仅取头部);head 为对象前 headBytes 字节(不足时取实际长度)。
|
||||
HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error)
|
||||
|
||||
// PresignGetURL 生成限时直链(下载);不支持直链的引擎返回 ErrNotSupported。
|
||||
PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error)
|
||||
|
||||
// PresignPutURL 生成限时直传(上传)URL;不支持直传的引擎返回 ErrNotSupported。
|
||||
PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error)
|
||||
|
||||
// HealthCheck 引擎健康检查(启动时与 /health 使用)。
|
||||
HealthCheck(ctx context.Context) error
|
||||
}
|
||||
|
||||
// ErrNotSupported 当前引擎不支持该能力(如本地引擎不支持预签名)。
|
||||
var ErrNotSupported = errors.New("storage: 当前引擎不支持该操作")
|
||||
|
||||
// ChunkPath 返回分片临时路径(约定统一为 <dir>/chunks/<upload_id>/<index>.part)。
|
||||
// 引擎可使用 ChunkDir 拼接自身路径。
|
||||
type PathBuilder interface {
|
||||
// ChunkDir 分片临时目录(相对 savePath 所在目录)。
|
||||
ChunkDir(savePath, uploadID string) string
|
||||
}
|
||||
@@ -0,0 +1,449 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
)
|
||||
|
||||
// 每次读写使用的缓冲大小:256KB,对齐参考实现 SystemFileStorage.chunk_size。
|
||||
const localChunkSize = 256 * 1024
|
||||
|
||||
// LocalStorage 本地文件系统引擎。
|
||||
//
|
||||
// 相比参考实现(SystemFileStorage)的改进:
|
||||
// - 双重路径防护:清洗相对路径 + 根目录前缀校验 + 符号链接逃逸校验;
|
||||
// - 全部落盘走「临时文件 + fsync + 原子重命名」,断电/中断不产生半截文件;
|
||||
// - 下载使用 io.NewSectionReader 支持任意 Range,无需整文件读入内存。
|
||||
type LocalStorage struct {
|
||||
// root 存储根目录(绝对路径)。
|
||||
root string
|
||||
// rootReal 经符号链接解析后的真实根目录,用于逃逸校验。
|
||||
rootReal string
|
||||
}
|
||||
|
||||
// NewLocalStorage 构造本地引擎。root 为空时使用系统临时目录下的 filecodebox_storage。
|
||||
func NewLocalStorage(root string) (*LocalStorage, error) {
|
||||
if strings.TrimSpace(root) == "" {
|
||||
root = filepath.Join(os.TempDir(), "filecodebox_storage")
|
||||
}
|
||||
abs, err := filepath.Abs(root)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("storage/local: 解析根目录失败: %w", err)
|
||||
}
|
||||
if err := os.MkdirAll(abs, 0o755); err != nil {
|
||||
return nil, fmt.Errorf("storage/local: 创建根目录失败: %w", err)
|
||||
}
|
||||
real := abs
|
||||
if resolved, err := filepath.EvalSymlinks(abs); err == nil {
|
||||
real = resolved
|
||||
}
|
||||
return &LocalStorage{root: abs, rootReal: real}, nil
|
||||
}
|
||||
|
||||
func init() {
|
||||
RegisterEngine("local", func(ctx context.Context) (Storage, error) {
|
||||
return NewLocalStorage(engineOptions.Local.Root)
|
||||
})
|
||||
}
|
||||
|
||||
// withinRoot 判断路径 p 是否位于 root 内(含 root 本身)。
|
||||
func withinRoot(p, root string) bool {
|
||||
p = filepath.Clean(p)
|
||||
root = filepath.Clean(root)
|
||||
if p == root {
|
||||
return true
|
||||
}
|
||||
return strings.HasPrefix(p, root+string(os.PathSeparator))
|
||||
}
|
||||
|
||||
// absPath 将存储侧相对路径解析为根目录内的绝对路径。
|
||||
// 任何路径穿越或符号链接逃逸都会返回 ErrInvalidPath。
|
||||
func (l *LocalStorage) absPath(savePath string) (string, error) {
|
||||
cleaned, ok := SanitizePath(savePath)
|
||||
if !ok {
|
||||
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
|
||||
}
|
||||
full := filepath.Join(l.root, filepath.FromSlash(cleaned))
|
||||
if !withinRoot(full, l.root) {
|
||||
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
|
||||
}
|
||||
// 符号链接逃逸校验:文件已存在时解析真实路径;不存在时校验最深已存在的父目录。
|
||||
if real, err := filepath.EvalSymlinks(full); err == nil {
|
||||
if !withinRoot(real, l.rootReal) {
|
||||
return "", fmt.Errorf("%w: 符号链接逃逸 %q", ErrInvalidPath, savePath)
|
||||
}
|
||||
} else {
|
||||
dir := filepath.Dir(full)
|
||||
if realDir, err := filepath.EvalSymlinks(dir); err == nil && !withinRoot(realDir, l.rootReal) {
|
||||
return "", fmt.Errorf("%w: 符号链接逃逸 %q", ErrInvalidPath, savePath)
|
||||
}
|
||||
}
|
||||
return full, nil
|
||||
}
|
||||
|
||||
// SaveFile 流式保存:256KB 分块读取写入临时文件,fsync 后原子重命名。
|
||||
func (l *LocalStorage) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
|
||||
full, err := l.absPath(savePath)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
src := &countingReader{r: r}
|
||||
if err := writeFileAtomic(full, src); err != nil {
|
||||
return src.count(), err
|
||||
}
|
||||
return src.count(), nil
|
||||
}
|
||||
|
||||
// DeleteFile 删除文件;文件不存在时静默成功(对齐契约)。
|
||||
func (l *LocalStorage) DeleteFile(ctx context.Context, savePath string) error {
|
||||
full, err := l.absPath(savePath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Remove(full); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
return fmt.Errorf("storage/local: 删除失败: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Open 打开文件下载流;rng 非 nil 时用 SectionReader 实现 Range 语义。
|
||||
func (l *LocalStorage) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
|
||||
full, err := l.absPath(savePath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
f, err := os.Open(full)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, fmt.Errorf("storage/local: 打开文件失败: %w", err)
|
||||
}
|
||||
info, err := f.Stat()
|
||||
if err != nil {
|
||||
_ = f.Close()
|
||||
return nil, fmt.Errorf("storage/local: 获取文件信息失败: %w", err)
|
||||
}
|
||||
if info.IsDir() {
|
||||
_ = f.Close()
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
size := info.Size()
|
||||
start, end := int64(0), size-1
|
||||
if rng != nil {
|
||||
if rng.Start < 0 || (rng.End != -1 && rng.End < rng.Start) {
|
||||
_ = f.Close()
|
||||
return nil, fmt.Errorf("%w: 非法 Range", ErrRangeNotSatisfiable)
|
||||
}
|
||||
if rng.Start >= size {
|
||||
_ = f.Close()
|
||||
return nil, ErrRangeNotSatisfiable
|
||||
}
|
||||
start = rng.Start
|
||||
end = size - 1
|
||||
if rng.End != -1 && rng.End < end {
|
||||
end = rng.End
|
||||
}
|
||||
}
|
||||
section := io.NewSectionReader(f, start, end-start+1)
|
||||
dl := &Download{
|
||||
ReadCloser: &fileSection{Reader: section, closer: f},
|
||||
Start: start,
|
||||
End: end,
|
||||
Total: size,
|
||||
Meta: FileMeta{
|
||||
Size: size,
|
||||
ContentType: mime.TypeByExtension(strings.ToLower(filepath.Ext(full))),
|
||||
AcceptRanges: true,
|
||||
},
|
||||
}
|
||||
if end < 0 { // 空文件:End 语义上等于 -1(未知),Total=0 已表达大小
|
||||
dl.End = -1
|
||||
}
|
||||
return dl, nil
|
||||
}
|
||||
|
||||
// fileSection 组合 SectionReader 与文件关闭器。
|
||||
type fileSection struct {
|
||||
io.Reader
|
||||
closer io.Closer
|
||||
}
|
||||
|
||||
func (f *fileSection) Close() error { return f.closer.Close() }
|
||||
|
||||
// Stat 获取文件元信息。
|
||||
func (l *LocalStorage) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
|
||||
full, err := l.absPath(savePath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
info, err := os.Stat(full)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, fmt.Errorf("storage/local: Stat 失败: %w", err)
|
||||
}
|
||||
if info.IsDir() {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return &FileMeta{
|
||||
Size: info.Size(),
|
||||
ContentType: mime.TypeByExtension(strings.ToLower(filepath.Ext(full))),
|
||||
AcceptRanges: true,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// SaveChunk 保存分片到 <父目录>/chunks/<uploadID>/<index>.part,原子写入。
|
||||
func (l *LocalStorage) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
|
||||
// 先校验目标路径合法性,分片目录随合法路径派生。
|
||||
if _, err := l.absPath(savePath); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
chunkRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, chunkIndex))
|
||||
if !ok {
|
||||
return 0, fmt.Errorf("%w: 分片路径非法 %q", ErrInvalidPath, savePath)
|
||||
}
|
||||
full := filepath.Join(l.root, filepath.FromSlash(chunkRel))
|
||||
if !withinRoot(full, l.root) {
|
||||
return 0, fmt.Errorf("%w: 分片路径非法 %q", ErrInvalidPath, chunkRel)
|
||||
}
|
||||
src := &countingReader{r: r}
|
||||
if err := writeFileAtomic(full, src); err != nil {
|
||||
return src.count(), err
|
||||
}
|
||||
return src.count(), nil
|
||||
}
|
||||
|
||||
// MergeChunks 按索引 0..total-1 有序合并分片:
|
||||
// - 逐分片流式拷贝到临时输出(边拷贝边计算整文件与分片 SHA256);
|
||||
// - verifyHash 非 nil 时校验分片哈希(空串跳过);
|
||||
// - 全部通过后 fsync + 原子重命名,并清理分片临时目录。
|
||||
func (l *LocalStorage) 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/local: 非法分片总数 %d", total)
|
||||
}
|
||||
full, err := l.absPath(savePath)
|
||||
if err != nil {
|
||||
return 0, "", err
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(full), 0o755); err != nil {
|
||||
return 0, "", fmt.Errorf("storage/local: 创建目标目录失败: %w", err)
|
||||
}
|
||||
tmp, err := os.CreateTemp(filepath.Dir(full), "."+filepath.Base(full)+".merging-*")
|
||||
if err != nil {
|
||||
return 0, "", fmt.Errorf("storage/local: 创建临时文件失败: %w", err)
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
defer func() {
|
||||
_ = tmp.Close()
|
||||
_ = os.Remove(tmpName) // 成功时已被重命名,删除静默失败
|
||||
}()
|
||||
|
||||
totalHash := sha256.New()
|
||||
buf := make([]byte, localChunkSize)
|
||||
var size int64
|
||||
for i := 0; i < total; i++ {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return 0, "", err
|
||||
}
|
||||
partRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, i))
|
||||
if !ok {
|
||||
return 0, "", fmt.Errorf("%w: 分片 %d 路径非法", ErrInvalidPath, i)
|
||||
}
|
||||
partPath := filepath.Join(l.root, filepath.FromSlash(partRel))
|
||||
in, err := os.Open(partPath)
|
||||
if err != nil {
|
||||
return 0, "", fmt.Errorf("storage/local: 分片 %d 不存在: %w", i, err)
|
||||
}
|
||||
chunkHash := sha256.New()
|
||||
n, err := io.CopyBuffer(io.MultiWriter(tmp, totalHash, chunkHash), in, buf)
|
||||
_ = in.Close()
|
||||
if err != nil {
|
||||
return 0, "", fmt.Errorf("storage/local: 读取分片 %d 失败: %w", i, err)
|
||||
}
|
||||
if verifyHash != nil {
|
||||
expected, err := verifyHash(i)
|
||||
if err != nil {
|
||||
return 0, "", err
|
||||
}
|
||||
if expected != "" && expected != hex.EncodeToString(chunkHash.Sum(nil)) {
|
||||
return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, i, expected)
|
||||
}
|
||||
}
|
||||
size += n
|
||||
}
|
||||
if err := tmp.Sync(); err != nil {
|
||||
return 0, "", fmt.Errorf("storage/local: 落盘失败: %w", err)
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
return 0, "", fmt.Errorf("storage/local: 关闭临时文件失败: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmpName, full); err != nil {
|
||||
return 0, "", fmt.Errorf("storage/local: 原子重命名失败: %w", err)
|
||||
}
|
||||
// 合并成功后清理分片临时目录(静默容错,不掩盖成功结果)。
|
||||
_ = l.CleanChunks(ctx, uploadID, savePath)
|
||||
return size, hex.EncodeToString(totalHash.Sum(nil)), nil
|
||||
}
|
||||
|
||||
// CleanChunks 清理分片临时目录;不存在时静默成功,并尝试移除空 chunks 父目录。
|
||||
func (l *LocalStorage) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
|
||||
dirRel, ok := SanitizePath(chunkDirOf(savePath, uploadID))
|
||||
if !ok {
|
||||
return fmt.Errorf("%w: 分片目录非法 %q", ErrInvalidPath, savePath)
|
||||
}
|
||||
dir := filepath.Join(l.root, filepath.FromSlash(dirRel))
|
||||
if !withinRoot(dir, l.root) {
|
||||
return fmt.Errorf("%w: 分片目录非法 %q", ErrInvalidPath, dirRel)
|
||||
}
|
||||
if err := os.RemoveAll(dir); err != nil {
|
||||
return fmt.Errorf("storage/local: 清理分片目录失败: %w", err)
|
||||
}
|
||||
// 父级 chunks 目录为空则一并清理(对齐参考实现)。
|
||||
chunksParent := filepath.Dir(dir)
|
||||
if entries, err := os.ReadDir(chunksParent); err == nil && len(entries) == 0 {
|
||||
_ = os.Remove(chunksParent)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// chunkDirOf 返回分片目录(去掉文件名部分):<父目录>/chunks/<uploadID>。
|
||||
func chunkDirOf(savePath, uploadID string) string {
|
||||
cd := ChunkDir(savePath, uploadID)
|
||||
// ChunkDir 返回 "<dir>/chunks/<uploadID>/<name>",去掉末段文件名即目录。
|
||||
if idx := strings.LastIndex(cd, "/"); idx > 0 {
|
||||
return cd[:idx]
|
||||
}
|
||||
return cd
|
||||
}
|
||||
|
||||
// FileExists 检查文件是否存在;非法路径按不存在处理(对齐参考实现)。
|
||||
func (l *LocalStorage) FileExists(ctx context.Context, savePath string) (bool, error) {
|
||||
full, err := l.absPath(savePath)
|
||||
if err != nil {
|
||||
return false, nil
|
||||
}
|
||||
info, err := os.Stat(full)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return false, nil
|
||||
}
|
||||
return false, fmt.Errorf("storage/local: Stat 失败: %w", err)
|
||||
}
|
||||
return !info.IsDir(), nil
|
||||
}
|
||||
|
||||
// HeadMeta 读取文件元信息与前 n 字节(本地引擎实现)。
|
||||
func (l *LocalStorage) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
|
||||
full, err := l.absPath(savePath)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
f, err := os.Open(full)
|
||||
if err != nil {
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil, nil, ErrNotFound
|
||||
}
|
||||
return nil, nil, fmt.Errorf("storage/local: 打开文件失败: %w", err)
|
||||
}
|
||||
defer func() { _ = f.Close() }()
|
||||
info, err := f.Stat()
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("storage/local: Stat 失败: %w", err)
|
||||
}
|
||||
if info.IsDir() {
|
||||
return nil, nil, ErrNotFound
|
||||
}
|
||||
head := make([]byte, headBytes)
|
||||
n, _ := io.ReadFull(f, head)
|
||||
return &FileMeta{
|
||||
Size: info.Size(),
|
||||
ContentType: mime.TypeByExtension(strings.ToLower(filepath.Ext(full))),
|
||||
AcceptRanges: true,
|
||||
}, head[:n], nil
|
||||
}
|
||||
|
||||
// PresignGetURL 本地引擎不支持直链。
|
||||
func (l *LocalStorage) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
|
||||
return "", ErrNotSupported
|
||||
}
|
||||
|
||||
// PresignPutURL 本地引擎不支持直传。
|
||||
func (l *LocalStorage) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
|
||||
return "", ErrNotSupported
|
||||
}
|
||||
|
||||
// HealthCheck 健康检查:根目录可写(写入并删除探针文件)。
|
||||
func (l *LocalStorage) HealthCheck(ctx context.Context) error {
|
||||
if err := os.MkdirAll(l.root, 0o755); err != nil {
|
||||
return fmt.Errorf("%w: 本地存储根目录不可创建: %v", ErrUnavailable, err)
|
||||
}
|
||||
probe := filepath.Join(l.root, ".health-probe")
|
||||
if err := os.WriteFile(probe, []byte("ok"), 0o644); err != nil {
|
||||
return fmt.Errorf("%w: 本地存储不可写: %v", ErrUnavailable, err)
|
||||
}
|
||||
_ = os.Remove(probe)
|
||||
return nil
|
||||
}
|
||||
|
||||
// writeFileAtomic 临时文件 + fsync + rename 的原子落盘。
|
||||
func writeFileAtomic(dst string, src io.Reader) error {
|
||||
dir := filepath.Dir(dst)
|
||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
||||
return fmt.Errorf("storage/local: 创建目录失败: %w", err)
|
||||
}
|
||||
tmp, err := os.CreateTemp(dir, "."+filepath.Base(dst)+".tmp-*")
|
||||
if err != nil {
|
||||
return fmt.Errorf("storage/local: 创建临时文件失败: %w", err)
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
cleanup := func() { _ = tmp.Close(); _ = os.Remove(tmpName) }
|
||||
if _, err := io.CopyBuffer(tmp, src, make([]byte, localChunkSize)); err != nil {
|
||||
cleanup()
|
||||
return fmt.Errorf("storage/local: 写入失败: %w", err)
|
||||
}
|
||||
if err := tmp.Sync(); err != nil {
|
||||
cleanup()
|
||||
return fmt.Errorf("storage/local: fsync 失败: %w", err)
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
_ = os.Remove(tmpName)
|
||||
return fmt.Errorf("storage/local: 关闭临时文件失败: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmpName, dst); err != nil {
|
||||
_ = os.Remove(tmpName)
|
||||
return fmt.Errorf("storage/local: 原子重命名失败: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// countingReader 统计累计读取字节数(并发安全)。
|
||||
type countingReader struct {
|
||||
r io.Reader
|
||||
n atomic.Int64
|
||||
}
|
||||
|
||||
func (c *countingReader) Read(p []byte) (int, error) {
|
||||
n, err := c.r.Read(p)
|
||||
c.n.Add(int64(n))
|
||||
return n, err
|
||||
}
|
||||
|
||||
// count 返回累计字节数。
|
||||
func (c *countingReader) count() int64 { return c.n.Load() }
|
||||
|
||||
// reset 归零计数(请求体重放时使用)。
|
||||
func (c *countingReader) reset() { c.n.Store(0) }
|
||||
|
||||
// 接口编译期断言。
|
||||
var _ Storage = (*LocalStorage)(nil)
|
||||
@@ -0,0 +1,346 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// newTestLocal 构造以临时目录为根的本地引擎。
|
||||
func newTestLocal(t *testing.T) *LocalStorage {
|
||||
t.Helper()
|
||||
st, err := NewLocalStorage(t.TempDir())
|
||||
if err != nil {
|
||||
t.Fatalf("NewLocalStorage: %v", err)
|
||||
}
|
||||
return st
|
||||
}
|
||||
|
||||
func sha256Hex(b []byte) string {
|
||||
sum := sha256.Sum256(b)
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// TestLocalSaveOpenRange 保存/Stat/完整与 Range 下载。
|
||||
func TestLocalSaveOpenRange(t *testing.T) {
|
||||
st := newTestLocal(t)
|
||||
ctx := context.Background()
|
||||
data := []byte("hello filecodebox 本地引擎 0123456789")
|
||||
|
||||
n, err := st.SaveFile(ctx, bytes.NewReader(data), "2025/08/测试文件.bin")
|
||||
if err != nil {
|
||||
t.Fatalf("SaveFile: %v", err)
|
||||
}
|
||||
if n != int64(len(data)) {
|
||||
t.Fatalf("SaveFile n = %d, want %d", n, len(data))
|
||||
}
|
||||
|
||||
// Stat
|
||||
meta, err := st.Stat(ctx, "2025/08/测试文件.bin")
|
||||
if err != nil {
|
||||
t.Fatalf("Stat: %v", err)
|
||||
}
|
||||
if meta.Size != int64(len(data)) || !meta.AcceptRanges {
|
||||
t.Fatalf("Stat = %+v", meta)
|
||||
}
|
||||
|
||||
// 完整下载:对齐 go-api 约定 Start=0、End=Total-1
|
||||
dl, err := st.Open(ctx, "2025/08/测试文件.bin", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Open: %v", err)
|
||||
}
|
||||
got, err := io.ReadAll(dl)
|
||||
_ = dl.Close()
|
||||
if err != nil {
|
||||
t.Fatalf("read full: %v", err)
|
||||
}
|
||||
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 download offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
|
||||
}
|
||||
|
||||
// Range 下载 [2, 7]
|
||||
dl, err = st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: 2, End: 7})
|
||||
if err != nil {
|
||||
t.Fatalf("Open range: %v", err)
|
||||
}
|
||||
got, err = io.ReadAll(dl)
|
||||
_ = dl.Close()
|
||||
if err != nil {
|
||||
t.Fatalf("read range: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, data[2:8]) {
|
||||
t.Fatalf("range content = %q, want %q", got, data[2:8])
|
||||
}
|
||||
|
||||
// Range end 越界自动钳制
|
||||
dl, err = st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: 5, End: 99999})
|
||||
if err != nil {
|
||||
t.Fatalf("Open clamp range: %v", err)
|
||||
}
|
||||
got, _ = io.ReadAll(dl)
|
||||
_ = dl.Close()
|
||||
if !bytes.Equal(got, data[5:]) {
|
||||
t.Fatalf("clamp range mismatch")
|
||||
}
|
||||
|
||||
// 起点越界 → 416
|
||||
if _, err := st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: int64(len(data)) + 1, End: -1}); err == nil {
|
||||
t.Fatalf("out-of-range start should fail")
|
||||
} else if !strings.Contains(err.Error(), ErrRangeNotSatisfiable.Error()) {
|
||||
t.Fatalf("want ErrRangeNotSatisfiable, got %v", err)
|
||||
}
|
||||
|
||||
// 不存在 → 404
|
||||
if _, err := st.Open(ctx, "no/such/file.bin", nil); err == nil || !strings.Contains(err.Error(), ErrNotFound.Error()) {
|
||||
t.Fatalf("want ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLocalSaveFileAtomic 原子写:目录自动创建,无临时残留。
|
||||
func TestLocalSaveFileAtomic(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
dir := t.TempDir()
|
||||
st2, _ := NewLocalStorage(dir)
|
||||
|
||||
big := bytes.Repeat([]byte("abc123"), 128*1024) // 768KB,跨多个 256KB 缓冲
|
||||
if _, err := st2.SaveFile(ctx, bytes.NewReader(big), "a/b/c/big.bin"); err != nil {
|
||||
t.Fatalf("SaveFile: %v", err)
|
||||
}
|
||||
// 无临时残留
|
||||
entries, _ := os.ReadDir(filepath.Join(dir, "a", "b", "c"))
|
||||
for _, e := range entries {
|
||||
if strings.HasPrefix(e.Name(), ".") {
|
||||
t.Fatalf("临时文件残留: %s", e.Name())
|
||||
}
|
||||
}
|
||||
got, _ := os.ReadFile(filepath.Join(dir, "a", "b", "c", "big.bin"))
|
||||
if !bytes.Equal(got, big) {
|
||||
t.Fatalf("content mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
// TestLocalTraversal 防路径穿越(含符号链接逃逸)。
|
||||
func TestLocalTraversal(t *testing.T) {
|
||||
st := newTestLocal(t)
|
||||
ctx := context.Background()
|
||||
|
||||
for _, p := range []string{"../escape.txt", "a/../../escape", "..", "/../x"} {
|
||||
if _, err := st.SaveFile(ctx, strings.NewReader("x"), p); err == nil ||
|
||||
!strings.Contains(err.Error(), ErrInvalidPath.Error()) {
|
||||
t.Fatalf("SaveFile(%q) 应拒绝: %v", p, err)
|
||||
}
|
||||
if _, err := st.Open(ctx, p, nil); err == nil || !strings.Contains(err.Error(), ErrInvalidPath.Error()) {
|
||||
t.Fatalf("Open(%q) 应拒绝: %v", p, err)
|
||||
}
|
||||
}
|
||||
|
||||
// 符号链接逃逸
|
||||
outside := filepath.Join(t.TempDir(), "outside.txt")
|
||||
if err := os.WriteFile(outside, []byte("secret"), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
link := filepath.Join(st.root, "link.txt")
|
||||
if err := os.Symlink(outside, link); err != nil {
|
||||
t.Skipf("symlink 不可用: %v", err)
|
||||
}
|
||||
if _, err := st.Open(ctx, "link.txt", nil); err == nil ||
|
||||
!strings.Contains(err.Error(), ErrInvalidPath.Error()) {
|
||||
t.Fatalf("symlink escape 应拒绝: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLocalChunkLifecycle 分片保存/合并(索引有序 + SHA256 校验 + 清理临时目录)。
|
||||
func TestLocalChunkLifecycle(t *testing.T) {
|
||||
st := newTestLocal(t)
|
||||
ctx := context.Background()
|
||||
savePath := "2025/09/chunked.bin"
|
||||
uploadID := "upload-abc"
|
||||
|
||||
chunks := [][]byte{[]byte("AAAA"), []byte("BB"), []byte("CCCCCC")}
|
||||
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)
|
||||
}
|
||||
if n != int64(len(c)) {
|
||||
t.Fatalf("SaveChunk %d n = %d", i, n)
|
||||
}
|
||||
hashes[i] = sha256Hex(c)
|
||||
total += int64(len(c))
|
||||
}
|
||||
|
||||
// 分片文件确实存在于临时目录
|
||||
for i := range chunks {
|
||||
exists, err := st.FileExists(ctx, ChunkPartPath(savePath, uploadID, i))
|
||||
if err != nil || !exists {
|
||||
t.Fatalf("分片 %d 应存在: %v %v", i, exists, err)
|
||||
}
|
||||
}
|
||||
|
||||
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("merged size = %d, want %d", size, total)
|
||||
}
|
||||
want := sha256Hex(bytes.Join(chunks, nil))
|
||||
if fileHash != want {
|
||||
t.Fatalf("file hash = %s, want %s", fileHash, want)
|
||||
}
|
||||
// 合并内容 = 按索引有序拼接
|
||||
got, err := os.ReadFile(filepath.Join(st.root, filepath.FromSlash(savePath)))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !bytes.Equal(got, bytes.Join(chunks, nil)) {
|
||||
t.Fatalf("merged content mismatch")
|
||||
}
|
||||
// 合并后分片目录已清理
|
||||
exists, _ := st.FileExists(ctx, ChunkPartPath(savePath, uploadID, 0))
|
||||
if exists {
|
||||
t.Fatalf("合并后分片应已清理")
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(st.root, "2025/09/chunks")); !os.IsNotExist(err) {
|
||||
t.Fatalf("chunks 父目录应已清理: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLocalChunkHashMismatch 哈希不匹配 → 合并失败且不留输出文件。
|
||||
func TestLocalChunkHashMismatch(t *testing.T) {
|
||||
st := newTestLocal(t)
|
||||
ctx := context.Background()
|
||||
savePath := "mismatch.bin"
|
||||
if _, err := st.SaveChunk(ctx, "uid", 0, strings.NewReader("data"), savePath); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, _, err := st.MergeChunks(ctx, "uid", 1, func(i int) (string, error) {
|
||||
return sha256Hex([]byte("WRONG")), nil
|
||||
}, savePath)
|
||||
if err == nil || !strings.Contains(err.Error(), ErrHashMismatch.Error()) {
|
||||
t.Fatalf("want ErrHashMismatch, got %v", err)
|
||||
}
|
||||
exists, _ := st.FileExists(ctx, savePath)
|
||||
if exists {
|
||||
t.Fatalf("校验失败不应产出正式文件")
|
||||
}
|
||||
}
|
||||
|
||||
// TestLocalVerifyHashAbort verifyHash 返回 error → 合并中止。
|
||||
func TestLocalVerifyHashAbort(t *testing.T) {
|
||||
st := newTestLocal(t)
|
||||
ctx := context.Background()
|
||||
if _, err := st.SaveChunk(ctx, "uid2", 0, strings.NewReader("data"), "v.bin"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
boom := context.Canceled
|
||||
if _, _, err := st.MergeChunks(ctx, "uid2", 1, func(i int) (string, error) {
|
||||
return "", boom
|
||||
}, "v.bin"); err == nil {
|
||||
t.Fatalf("verifyHash 错误应向上传播")
|
||||
}
|
||||
}
|
||||
|
||||
// TestLocalDeleteExistsClean 删除/存在性/清理。
|
||||
func TestLocalDeleteExistsClean(t *testing.T) {
|
||||
st := newTestLocal(t)
|
||||
ctx := context.Background()
|
||||
if _, err := st.SaveFile(ctx, strings.NewReader("x"), "d/f.bin"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if ok, _ := st.FileExists(ctx, "d/f.bin"); !ok {
|
||||
t.Fatalf("文件应存在")
|
||||
}
|
||||
if err := st.DeleteFile(ctx, "d/f.bin"); err != nil {
|
||||
t.Fatalf("DeleteFile: %v", err)
|
||||
}
|
||||
if ok, _ := st.FileExists(ctx, "d/f.bin"); ok {
|
||||
t.Fatalf("文件应已删除")
|
||||
}
|
||||
// 删除不存在 → nil
|
||||
if err := st.DeleteFile(ctx, "d/f.bin"); err != nil {
|
||||
t.Fatalf("删除不存在应静默: %v", err)
|
||||
}
|
||||
// CleanChunks 幂等
|
||||
if err := st.CleanChunks(ctx, "uid-x", "y.bin"); err != nil {
|
||||
t.Fatalf("CleanChunks: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLocalPresignNotSupported 预签名 → ErrNotSupported。
|
||||
func TestLocalPresignNotSupported(t *testing.T) {
|
||||
st := newTestLocal(t)
|
||||
ctx := context.Background()
|
||||
if _, err := st.PresignGetURL(ctx, "a.bin", 60); err == nil || !strings.Contains(err.Error(), ErrNotSupported.Error()) {
|
||||
t.Fatalf("PresignGetURL want ErrNotSupported, got %v", err)
|
||||
}
|
||||
if _, err := st.PresignPutURL(ctx, "a.bin", 60); err == nil || !strings.Contains(err.Error(), ErrNotSupported.Error()) {
|
||||
t.Fatalf("PresignPutURL want ErrNotSupported, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLocalHealthCheck 健康检查。
|
||||
func TestLocalHealthCheck(t *testing.T) {
|
||||
st := newTestLocal(t)
|
||||
if err := st.HealthCheck(context.Background()); err != nil {
|
||||
t.Fatalf("HealthCheck: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestLocalConcurrent 并发安全冒烟。
|
||||
func TestLocalConcurrent(t *testing.T) {
|
||||
st := newTestLocal(t)
|
||||
ctx := context.Background()
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 16; i++ {
|
||||
wg.Add(1)
|
||||
go func(i int) {
|
||||
defer wg.Done()
|
||||
p := "conc/" + itoa(i) + ".bin"
|
||||
if _, err := st.SaveFile(ctx, strings.NewReader(strings.Repeat("x", i+1)), p); err != nil {
|
||||
t.Errorf("save %d: %v", i, err)
|
||||
return
|
||||
}
|
||||
dl, err := st.Open(ctx, p, nil)
|
||||
if err != nil {
|
||||
t.Errorf("open %d: %v", i, err)
|
||||
return
|
||||
}
|
||||
_, _ = io.ReadAll(dl)
|
||||
_ = dl.Close()
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
// TestLocalFactoryRegistry 工厂注册与构造。
|
||||
func TestLocalFactoryRegistry(t *testing.T) {
|
||||
prev := engineOptions.Local.Root
|
||||
engineOptions.Local.Root = t.TempDir()
|
||||
defer func() { engineOptions.Local.Root = prev }()
|
||||
st, err := NewEngine(context.Background(), "local")
|
||||
if err != nil {
|
||||
t.Fatalf("NewEngine(local): %v", err)
|
||||
}
|
||||
if err := st.HealthCheck(context.Background()); err != nil {
|
||||
t.Fatalf("HealthCheck: %v", err)
|
||||
}
|
||||
if _, err := NewEngine(context.Background(), "unknown"); err == nil {
|
||||
t.Fatalf("未知引擎应报错")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,187 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// Manager 存储引擎管理器:实现 Storage 全接口并支持运行时热切换。
|
||||
//
|
||||
// v3 需求:管理后台可设置存储类型(local|s3|webdav)与各引擎参数,
|
||||
// 保存后无需重启即生效。设计要点:
|
||||
// - 读写/保存类操作全部委托到"当前引擎"(原子指针,无锁热路径);
|
||||
// - Switch 先构建并健康检查新引擎,成功才替换指针,失败保持原引擎;
|
||||
// - EngineOf 按名字取引擎实例(带缓存),供"按文件归属引擎取回旧文件"使用;
|
||||
// - 管理端修改引擎参数后调用 Invalidate 使对应实例缓存失效,下次构建生效。
|
||||
type Manager struct {
|
||||
// build 构建指定引擎实例(由装配方注入:内部刷新全局 EngineOptions 后走工厂)。
|
||||
build func(name string) (Storage, error)
|
||||
|
||||
mu sync.RWMutex
|
||||
current Storage
|
||||
curName string
|
||||
cache map[string]Storage
|
||||
}
|
||||
|
||||
// validEngines 合法引擎名(与 FCB_STORAGE_ENGINE 枚举一致)。
|
||||
var validEngines = map[string]bool{"local": true, "s3": true, "webdav": true}
|
||||
|
||||
// ValidEngine 校验引擎名是否合法。
|
||||
func ValidEngine(name string) bool { return validEngines[name] }
|
||||
|
||||
// NewManager 创建管理器:current 为启动时已构建的引擎(主装配流已做过健康检查)。
|
||||
// build 注入构建函数(管理端切换/参数变更时使用,内部须串行——Manager 已加锁)。
|
||||
func NewManager(name string, current Storage, build func(name string) (Storage, error)) *Manager {
|
||||
return &Manager{
|
||||
build: build,
|
||||
current: current,
|
||||
curName: name,
|
||||
cache: map[string]Storage{name: current},
|
||||
}
|
||||
}
|
||||
|
||||
// —— Storage 接口委托(全部走当前引擎)——
|
||||
|
||||
// SaveFile 流式保存文件(委托当前引擎)。
|
||||
func (m *Manager) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
|
||||
return m.current.SaveFile(ctx, r, savePath)
|
||||
}
|
||||
|
||||
// DeleteFile 删除文件(委托当前引擎)。
|
||||
func (m *Manager) DeleteFile(ctx context.Context, savePath string) error {
|
||||
return m.current.DeleteFile(ctx, savePath)
|
||||
}
|
||||
|
||||
// Open 打开文件流(委托当前引擎;旧文件由 API 层先经 EngineOf 按归属引擎取)。
|
||||
func (m *Manager) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
|
||||
return m.current.Open(ctx, savePath, rng)
|
||||
}
|
||||
|
||||
// Stat 文件元信息(委托当前引擎)。
|
||||
func (m *Manager) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
|
||||
return m.current.Stat(ctx, savePath)
|
||||
}
|
||||
|
||||
// SaveChunk 保存分片(委托当前引擎)。
|
||||
func (m *Manager) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
|
||||
return m.current.SaveChunk(ctx, uploadID, chunkIndex, r, savePath)
|
||||
}
|
||||
|
||||
// MergeChunks 合并分片(委托当前引擎)。
|
||||
func (m *Manager) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
|
||||
return m.current.MergeChunks(ctx, uploadID, total, verifyHash, savePath)
|
||||
}
|
||||
|
||||
// CleanChunks 清理分片临时区(委托当前引擎)。
|
||||
func (m *Manager) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
|
||||
return m.current.CleanChunks(ctx, uploadID, savePath)
|
||||
}
|
||||
|
||||
// FileExists 文件存在性(委托当前引擎)。
|
||||
func (m *Manager) FileExists(ctx context.Context, savePath string) (bool, error) {
|
||||
return m.current.FileExists(ctx, savePath)
|
||||
}
|
||||
|
||||
// HeadMeta 元信息与头部字节(委托当前引擎;供直传 confirm 校验)。
|
||||
func (m *Manager) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
|
||||
return m.current.HeadMeta(ctx, savePath, headBytes)
|
||||
}
|
||||
|
||||
// PresignGetURL 限时直链下载(委托当前引擎)。
|
||||
func (m *Manager) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
|
||||
return m.current.PresignGetURL(ctx, savePath, expires)
|
||||
}
|
||||
|
||||
// PresignPutURL 限时直传(委托当前引擎)。
|
||||
func (m *Manager) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
|
||||
return m.current.PresignPutURL(ctx, savePath, expires)
|
||||
}
|
||||
|
||||
// HealthCheck 健康检查(委托当前引擎)。
|
||||
func (m *Manager) HealthCheck(ctx context.Context) error {
|
||||
return m.current.HealthCheck(ctx)
|
||||
}
|
||||
|
||||
// —— 管理面:当前引擎名 / 按名取实例 / 热切换 / 缓存失效 ——
|
||||
|
||||
// CurrentName 当前引擎名(管理端展示与文件归属戳用;并发安全)。
|
||||
func (m *Manager) CurrentName() string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.curName
|
||||
}
|
||||
|
||||
// Current 当前引擎实例。
|
||||
func (m *Manager) Current() Storage {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
return m.current
|
||||
}
|
||||
|
||||
// EngineOf 按名字取引擎实例(带缓存;用于按文件归属引擎取回旧文件)。
|
||||
// 实例不存在时现场构建(不健康检查——读旧文件尽力而为,构建失败即报错)。
|
||||
func (m *Manager) EngineOf(name string) (Storage, error) {
|
||||
if !validEngines[name] {
|
||||
return nil, fmt.Errorf("storage: 未知存储引擎 %q(仅支持 local|s3|webdav)", name)
|
||||
}
|
||||
m.mu.RLock()
|
||||
if s, ok := m.cache[name]; ok {
|
||||
m.mu.RUnlock()
|
||||
return s, nil
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
// 双检:拿写锁期间可能已被并发构建
|
||||
if s, ok := m.cache[name]; ok {
|
||||
return s, nil
|
||||
}
|
||||
s, err := m.build(name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
m.cache[name] = s
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// Switch 热切换当前引擎:构建新实例 → 健康检查 → 成功才替换指针。
|
||||
// 任一步失败返回错误且当前引擎保持不变(管理端 503 上报)。
|
||||
func (m *Manager) Switch(name string) (Storage, error) {
|
||||
if !validEngines[name] {
|
||||
return nil, fmt.Errorf("storage: 未知存储引擎 %q(仅支持 local|s3|webdav)", name)
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.curName == name {
|
||||
return m.current, nil
|
||||
}
|
||||
s, ok := m.cache[name]
|
||||
if !ok {
|
||||
var err error
|
||||
s, err = m.build(name)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("storage: 构建 %s 引擎失败: %w", name, err)
|
||||
}
|
||||
}
|
||||
if err := s.HealthCheck(context.Background()); err != nil {
|
||||
return nil, fmt.Errorf("storage: %s 引擎健康检查未通过: %w", name, err)
|
||||
}
|
||||
m.current = s
|
||||
m.curName = name
|
||||
m.cache[name] = s
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// Invalidate 引擎参数变更后使对应实例缓存失效(下次 EngineOf/Switch 重建生效)。
|
||||
// 当前引擎不受影响(运行中实例继续服务,直到显式 Switch)。
|
||||
func (m *Manager) Invalidate(name string) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if name == m.curName {
|
||||
return // 当前引擎实例仍被热路径使用,不重建;参数生效由下一次 Switch 完成
|
||||
}
|
||||
delete(m.cache, name)
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fakeEngine 可配置健康检查结果的桩引擎。
|
||||
type fakeEngine struct{ failHealth bool }
|
||||
|
||||
func (f *fakeEngine) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
|
||||
return 0, nil
|
||||
}
|
||||
func (f *fakeEngine) DeleteFile(ctx context.Context, savePath string) error { return nil }
|
||||
func (f *fakeEngine) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
func (f *fakeEngine) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
func (f *fakeEngine) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
|
||||
return 0, nil
|
||||
}
|
||||
func (f *fakeEngine) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
|
||||
return 0, "", nil
|
||||
}
|
||||
func (f *fakeEngine) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
|
||||
return nil
|
||||
}
|
||||
func (f *fakeEngine) FileExists(ctx context.Context, savePath string) (bool, error) {
|
||||
return false, nil
|
||||
}
|
||||
func (f *fakeEngine) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
|
||||
return nil, nil, ErrNotSupported
|
||||
}
|
||||
func (f *fakeEngine) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
|
||||
return "", ErrNotSupported
|
||||
}
|
||||
func (f *fakeEngine) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
|
||||
return "", ErrNotSupported
|
||||
}
|
||||
func (f *fakeEngine) HealthCheck(ctx context.Context) error {
|
||||
if f.failHealth {
|
||||
return ErrUnavailable
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// newTestManager 构造测试用 Manager:local 健康引擎起步;s3/webdav 由计数器控制健康。
|
||||
func newTestManager(s3Fail *atomic.Bool) *Manager {
|
||||
build := func(name string) (Storage, error) {
|
||||
switch name {
|
||||
case "local":
|
||||
return &fakeEngine{}, nil
|
||||
case "s3":
|
||||
return &fakeEngine{failHealth: s3Fail.Load()}, nil
|
||||
case "webdav":
|
||||
return &fakeEngine{}, nil
|
||||
}
|
||||
return nil, errors.New("unknown")
|
||||
}
|
||||
return NewManager("local", &fakeEngine{}, build)
|
||||
}
|
||||
|
||||
// TestSwitchSuccessAndCurrentName 切换成功后当前引擎名与实例更新。
|
||||
func TestSwitchSuccessAndCurrentName(t *testing.T) {
|
||||
var s3Fail atomic.Bool
|
||||
m := newTestManager(&s3Fail)
|
||||
if m.CurrentName() != "local" {
|
||||
t.Fatalf("初始引擎应为 local,得到 %s", m.CurrentName())
|
||||
}
|
||||
if _, err := m.Switch("s3"); err != nil {
|
||||
t.Fatalf("Switch(s3) 失败: %v", err)
|
||||
}
|
||||
if m.CurrentName() != "s3" {
|
||||
t.Fatalf("切换后引擎应为 s3,得到 %s", m.CurrentName())
|
||||
}
|
||||
if _, err := m.Switch("webdav"); err != nil {
|
||||
t.Fatalf("Switch(webdav) 失败: %v", err)
|
||||
}
|
||||
if m.CurrentName() != "webdav" {
|
||||
t.Fatalf("切换后引擎应为 webdav,得到 %s", m.CurrentName())
|
||||
}
|
||||
}
|
||||
|
||||
// TestSwitchFailureKeepsCurrent 健康检查失败时保持原引擎(v3 核心语义)。
|
||||
func TestSwitchFailureKeepsCurrent(t *testing.T) {
|
||||
var s3Fail atomic.Bool
|
||||
s3Fail.Store(true) // s3 不健康
|
||||
m := newTestManager(&s3Fail)
|
||||
if _, err := m.Switch("s3"); err == nil {
|
||||
t.Fatal("s3 不健康时 Switch 应失败")
|
||||
}
|
||||
if m.CurrentName() != "local" {
|
||||
t.Fatalf("切换失败后应保持 local,得到 %s", m.CurrentName())
|
||||
}
|
||||
// 恢复健康后可切换成功
|
||||
s3Fail.Store(false)
|
||||
if _, err := m.Switch("s3"); err != nil {
|
||||
t.Fatalf("恢复健康后 Switch(s3) 应成功: %v", err)
|
||||
}
|
||||
if m.CurrentName() != "s3" {
|
||||
t.Fatalf("恢复后引擎应为 s3,得到 %s", m.CurrentName())
|
||||
}
|
||||
}
|
||||
|
||||
// TestSwitchInvalidName 非法引擎名拒绝。
|
||||
func TestSwitchInvalidName(t *testing.T) {
|
||||
m := newTestManager(&atomic.Bool{})
|
||||
if _, err := m.Switch("ftp"); err == nil || !strings.Contains(err.Error(), "未知存储引擎") {
|
||||
t.Fatalf("非法引擎名应报未知存储引擎,得到 %v", err)
|
||||
}
|
||||
if !ValidEngine("local") || ValidEngine("ftp") {
|
||||
t.Fatal("ValidEngine 判定错误")
|
||||
}
|
||||
}
|
||||
|
||||
// TestEngineOfCacheAndInvalidate EngineOf 缓存命中 + Invalidate 后重建(参数生效路径)。
|
||||
func TestEngineOfCacheAndInvalidate(t *testing.T) {
|
||||
var builds atomic.Int64
|
||||
build := func(name string) (Storage, error) {
|
||||
builds.Add(1)
|
||||
return &fakeEngine{failHealth: false}, nil
|
||||
}
|
||||
m := NewManager("local", &fakeEngine{}, build)
|
||||
|
||||
s1, err := m.EngineOf("s3")
|
||||
if err != nil {
|
||||
t.Fatalf("EngineOf(s3): %v", err)
|
||||
}
|
||||
s2, err := m.EngineOf("s3")
|
||||
if err != nil {
|
||||
t.Fatalf("EngineOf(s3) second: %v", err)
|
||||
}
|
||||
if s1 != s2 {
|
||||
t.Fatal("EngineOf 应命中缓存返回同一实例")
|
||||
}
|
||||
if n := builds.Load(); n != 1 {
|
||||
t.Fatalf("应只构建 1 次,实际 %d", n)
|
||||
}
|
||||
|
||||
// Invalidate 后下次取重建新实例
|
||||
m.Invalidate("s3")
|
||||
s3, err := m.EngineOf("s3")
|
||||
if err != nil {
|
||||
t.Fatalf("EngineOf(s3) after invalidate: %v", err)
|
||||
}
|
||||
if s3 == s1 {
|
||||
t.Fatal("Invalidate 后应返回重建的新实例")
|
||||
}
|
||||
if n := builds.Load(); n != 2 {
|
||||
t.Fatalf("Invalidate 后应再构建 1 次,实际累计 %d", n)
|
||||
}
|
||||
}
|
||||
|
||||
// TestInvalidateCurrentNoop Invalidate 当前引擎不生效(热路径实例保持)。
|
||||
func TestInvalidateCurrentNoop(t *testing.T) {
|
||||
m := newTestManager(&atomic.Bool{})
|
||||
cur := m.Current()
|
||||
m.Invalidate("local") // 当前引擎:应为 no-op
|
||||
if m.Current() != cur {
|
||||
t.Fatal("Invalidate 当前引擎不应替换实例")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSwitchSameNameNoop 同名 Switch 幂等。
|
||||
func TestSwitchSameNameNoop(t *testing.T) {
|
||||
m := newTestManager(&atomic.Bool{})
|
||||
s, err := m.Switch("local")
|
||||
if err != nil {
|
||||
t.Fatalf("Switch(local) 同名应成功: %v", err)
|
||||
}
|
||||
if s != m.Current() {
|
||||
t.Fatal("同名 Switch 应返回当前实例")
|
||||
}
|
||||
}
|
||||
|
||||
// TestDelegateToCurrent 保存/读取类操作委托当前引擎(切换后指向新引擎)。
|
||||
func TestDelegateToCurrent(t *testing.T) {
|
||||
var s3Fail atomic.Bool
|
||||
m := newTestManager(&s3Fail)
|
||||
ctx := context.Background()
|
||||
// local 引擎 HealthCheck 健康
|
||||
if err := m.HealthCheck(ctx); err != nil {
|
||||
t.Fatalf("委托 HealthCheck(local): %v", err)
|
||||
}
|
||||
if _, err := m.Switch("s3"); err != nil {
|
||||
t.Fatalf("Switch(s3): %v", err)
|
||||
}
|
||||
if err := m.HealthCheck(ctx); err != nil {
|
||||
t.Fatalf("委托 HealthCheck(s3): %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package storage
|
||||
|
||||
// EngineOptions 引擎构造选项:由 main.go(API 层任务)从 config KV 填充。
|
||||
// 各引擎的 RegisterEngine 工厂读取本结构;零值即安全默认。
|
||||
type EngineOptions struct {
|
||||
// Local 本地引擎选项。
|
||||
Local LocalOptions
|
||||
// S3 S3 引擎选项。
|
||||
S3 S3Options
|
||||
// WebDAV WebDAV 引擎选项。
|
||||
WebDAV WebDAVOptions
|
||||
}
|
||||
|
||||
// LocalOptions 本地引擎配置(对齐 local_storage_path)。
|
||||
type LocalOptions struct {
|
||||
// Root 存储根目录;空则使用系统临时目录。
|
||||
Root string
|
||||
}
|
||||
|
||||
// S3Options S3 引擎配置(对齐 s3_* 配置键)。
|
||||
type S3Options struct {
|
||||
AccessKeyID string // s3_access_key_id
|
||||
SecretAccessKey string // s3_secret_access_key
|
||||
SessionToken string // aws_session_token
|
||||
Bucket string // s3_bucket_name
|
||||
Endpoint string // s3_endpoint_url(MinIO 等;空则 AWS 默认端点)
|
||||
Region string // s3_region_name,默认 auto
|
||||
AddressingStyle string // s3_addressing_style: auto|path|virtual
|
||||
}
|
||||
|
||||
// WebDAVOptions WebDAV 引擎配置(对齐 webdav_* 配置键 + 本次优化项)。
|
||||
type WebDAVOptions struct {
|
||||
// BaseURL 服务地址,如 https://dav.example.com/dav/。
|
||||
BaseURL string
|
||||
// Username/Password 凭据(Basic 与 Digest 共用)。
|
||||
Username string
|
||||
Password string
|
||||
// RootPath 远端根目录(webdav_root_path),会自动逐级创建。
|
||||
RootPath string
|
||||
// MaxRetries 5xx/网络错误最大重试次数(指数退避),0 取默认 3。
|
||||
MaxRetries int
|
||||
// BaseBackoff 重试基础退避时长,0 取默认 200ms。
|
||||
BaseBackoff int64
|
||||
// Timeout 单请求超时秒数,0 取默认 30s。
|
||||
Timeout int64
|
||||
// MaxIdleConnsPerHost 连接池每主机最大空闲连接,0 取默认 16(连接复用优化)。
|
||||
MaxIdleConnsPerHost int
|
||||
}
|
||||
|
||||
// engineOptions 全局引擎选项(由 main.go 注入;默认零值)。
|
||||
var engineOptions EngineOptions
|
||||
|
||||
// SetEngineOptions 注入引擎构造选项(在 RegisterEngine 工厂执行前调用)。
|
||||
func SetEngineOptions(opts EngineOptions) { engineOptions = opts }
|
||||
|
||||
// applyDefaults 填充零值默认项。
|
||||
func (o *WebDAVOptions) applyDefaults() {
|
||||
if o.MaxRetries <= 0 {
|
||||
o.MaxRetries = 3
|
||||
}
|
||||
if o.BaseBackoff <= 0 {
|
||||
o.BaseBackoff = 200
|
||||
}
|
||||
if o.Timeout <= 0 {
|
||||
o.Timeout = 30
|
||||
}
|
||||
if o.MaxIdleConnsPerHost <= 0 {
|
||||
o.MaxIdleConnsPerHost = 16
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"path"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ChunkDir 实现默认分片目录约定:<父目录>/chunks/<uploadID>。
|
||||
// local/s3/webdav 三引擎共用,保持分片路径一致。
|
||||
func ChunkDir(savePath, uploadID string) string {
|
||||
dir := path.Dir(savePath)
|
||||
name := path.Base(savePath)
|
||||
// 防御:savePath 非法时仍返回明确结构,具体引擎再做安全校验
|
||||
if name == "." || name == "/" {
|
||||
name = "file"
|
||||
}
|
||||
return path.Join(dir, "chunks", uploadID) + "/" + name
|
||||
}
|
||||
|
||||
// ChunkPartPath 分片对象完整路径(相对存储根)。
|
||||
func ChunkPartPath(savePath, uploadID string, index int) string {
|
||||
dir := path.Dir(savePath)
|
||||
return path.Join(dir, "chunks", uploadID, itoa(index)+".part")
|
||||
}
|
||||
|
||||
// SanitizePath 清理相对路径:统一斜杠、去首尾斜杠、拒绝 .. 穿越。
|
||||
// 返回清理后的相对路径与是否合法。
|
||||
func SanitizePath(p string) (string, bool) {
|
||||
raw := strings.ReplaceAll(strings.TrimSpace(p), "\\", "/")
|
||||
raw = strings.TrimPrefix(raw, "/")
|
||||
if raw == "" {
|
||||
return "", false
|
||||
}
|
||||
cleaned := path.Clean(raw)
|
||||
if cleaned == ".." || strings.HasPrefix(cleaned, "../") || path.IsAbs(cleaned) {
|
||||
return "", false
|
||||
}
|
||||
// 拒绝任何单独的 .. 段
|
||||
for _, seg := range strings.Split(cleaned, "/") {
|
||||
if seg == ".." {
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
return cleaned, true
|
||||
}
|
||||
|
||||
// SanitizeFileName 清理文件名:剥离路径、替换非法字符、限制长度。
|
||||
// 对齐参考 core/utils.py 的 sanitize_filename。
|
||||
func SanitizeFileName(name string) string {
|
||||
// 剥离路径
|
||||
if idx := strings.LastIndexAny(name, "/\\"); idx >= 0 {
|
||||
name = name[idx+1:]
|
||||
}
|
||||
var b strings.Builder
|
||||
for _, r := range name {
|
||||
switch {
|
||||
case r < 0x20 || r == 0x7f:
|
||||
b.WriteByte('_')
|
||||
case strings.ContainsRune(`\*?:"<>|`, r):
|
||||
b.WriteByte('_')
|
||||
case r == ' ':
|
||||
b.WriteByte('_')
|
||||
default:
|
||||
b.WriteRune(r)
|
||||
}
|
||||
}
|
||||
cleaned := b.String()
|
||||
// 压缩连续下划线
|
||||
for strings.Contains(cleaned, "__") {
|
||||
cleaned = strings.ReplaceAll(cleaned, "__", "_")
|
||||
}
|
||||
cleaned = strings.Trim(cleaned, "._")
|
||||
if cleaned == "" {
|
||||
return "unnamed_file"
|
||||
}
|
||||
if len(cleaned) > 255 {
|
||||
cleaned = cleaned[:255]
|
||||
}
|
||||
return cleaned
|
||||
}
|
||||
|
||||
// itoa 小整数转字符串。
|
||||
func itoa(n int) string {
|
||||
if n == 0 {
|
||||
return "0"
|
||||
}
|
||||
neg := n < 0
|
||||
if neg {
|
||||
n = -n
|
||||
}
|
||||
var buf [21]byte
|
||||
i := len(buf)
|
||||
for n > 0 {
|
||||
i--
|
||||
buf[i] = byte('0' + n%10)
|
||||
n /= 10
|
||||
}
|
||||
if neg {
|
||||
i--
|
||||
buf[i] = '-'
|
||||
}
|
||||
return string(buf[i:])
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"io"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestSanitizePath 校验路径穿越防护。
|
||||
func TestSanitizePath(t *testing.T) {
|
||||
cases := []struct {
|
||||
in string
|
||||
ok bool
|
||||
out string
|
||||
}{
|
||||
{"2025/08/uuid.zip", true, "2025/08/uuid.zip"},
|
||||
{"/2025/08/uuid.zip", true, "2025/08/uuid.zip"},
|
||||
{"a\\b\\c.txt", true, "a/b/c.txt"},
|
||||
{"../etc/passwd", false, ""},
|
||||
{"a/../../b", false, ""},
|
||||
{"..", false, ""},
|
||||
{"", false, ""},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
got, ok := SanitizePath(tc.in)
|
||||
if ok != tc.ok || (ok && got != tc.out) {
|
||||
t.Errorf("SanitizePath(%q) = (%q, %v), want (%q, %v)", tc.in, got, ok, tc.out, tc.ok)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestSanitizeFileName 校验文件名清理。
|
||||
func TestSanitizeFileName(t *testing.T) {
|
||||
cases := []struct{ in, want string }{
|
||||
{"hello world.zip", "hello_world.zip"},
|
||||
{"/path/to/file.txt", "file.txt"},
|
||||
{"a<b>:c?.mp4", "a_b_c_.mp4"}, // 连续下划线压缩,对齐参考 re.sub(r"_+", "_")
|
||||
{"", "unnamed_file"},
|
||||
{"__..__", "unnamed_file"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
if got := SanitizeFileName(tc.in); got != tc.want {
|
||||
t.Errorf("SanitizeFileName(%q) = %q, want %q", tc.in, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestChunkPartPath 校验分片路径约定。
|
||||
func TestChunkPartPath(t *testing.T) {
|
||||
got := ChunkPartPath("2025/08/uuid.zip", "upload-1", 3)
|
||||
want := "2025/08/chunks/upload-1/3.part"
|
||||
if got != want {
|
||||
t.Errorf("ChunkPartPath = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestChunkDir 校验分片目录约定。
|
||||
func TestChunkDir(t *testing.T) {
|
||||
got := ChunkDir("2025/08/uuid.zip", "upload-1")
|
||||
want := "2025/08/chunks/upload-1/uuid.zip"
|
||||
if got != want {
|
||||
t.Errorf("ChunkDir = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSHA256Helper 辅助:确认 sha256 用法一致(合并校验依赖)。
|
||||
func TestSHA256Helper(t *testing.T) {
|
||||
h := sha256.Sum256([]byte("abc"))
|
||||
if got := hex.EncodeToString(h[:]); got != "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad" {
|
||||
t.Errorf("sha256(abc) = %s", got)
|
||||
}
|
||||
_ = io.EOF
|
||||
}
|
||||
@@ -0,0 +1,650 @@
|
||||
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)
|
||||
@@ -0,0 +1,495 @@
|
||||
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(
|
||||
`<?xml version="1.0" encoding="UTF-8"?><Error><Code>%s</Code><Message>%s</Message></Error>`, 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(
|
||||
`<?xml version="1.0" encoding="UTF-8"?><InitiateMultipartUploadResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"><Bucket>%s</Bucket><Key>%s</Key><UploadId>%s</UploadId></InitiateMultipartUploadResult>`,
|
||||
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(
|
||||
`<?xml version="1.0" encoding="UTF-8"?><CompleteMultipartUploadResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"><Location>http://%s/%s/%s</Location><Bucket>%s</Bucket><Key>%s</Key><ETag>"merged"</ETag></CompleteMultipartUploadResult>`,
|
||||
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(`<?xml version="1.0" encoding="UTF-8"?><DeleteResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"></DeleteResult>`))
|
||||
// ListObjectsV2
|
||||
case r.Method == http.MethodGet && q.Get("list-type") == "2":
|
||||
f.listCount++
|
||||
prefix := q.Get("prefix")
|
||||
var body strings.Builder
|
||||
body.WriteString(`<?xml version="1.0" encoding="UTF-8"?><ListBucketResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"><Name>` + bucket + `</Name><Prefix>` + prefix + `</Prefix><IsTruncated>false</IsTruncated>`)
|
||||
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("<Contents><Key>%s</Key><Size>%d</Size></Contents>", k, len(f.objects[k])))
|
||||
}
|
||||
body.WriteString("</ListBucketResult>")
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,889 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/xml"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math/rand"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// WebDAVStorage 基于 net/http 的 WebDAV 引擎(本次重写的重点优化对象)。
|
||||
//
|
||||
// 相比参考实现(WebDAVFileStorage,aiohttp)的改进:
|
||||
// - 单例 http.Client + 连接池化 Transport(参考实现每个操作新建 ClientSession,无复用);
|
||||
// - Basic 与 Digest(RFC 2617,qop=auth,MD5/SHA-256)双认证自动协商(参考实现仅 Basic);
|
||||
// - GET 下载透传 Range 头(参考实现全量 GET,无法断点/分段);
|
||||
// - 5xx/429/网络错误指数退避重试,可配次数(参考实现无重试);
|
||||
// - 下载经 io.Pipe 流式转发,全程不落盘;
|
||||
// - 目录存在性内存缓存,按需逐级 MKCOL,避免每次保存都发 PROPFIND;
|
||||
// - 非流式操作带可配超时;流式传输由调用方 ctx 管控(可取消)。
|
||||
type WebDAVStorage struct {
|
||||
base *url.URL // 服务基址(含可能的路径前缀),以 / 结尾
|
||||
root string // 远端根目录(webdav_root_path)
|
||||
username string
|
||||
password string
|
||||
client *http.Client
|
||||
transport *http.Transport
|
||||
auth *authState
|
||||
|
||||
maxRetries int // 5xx/网络错误最大重试次数
|
||||
baseBackoff time.Duration // 退避基数
|
||||
opTimeout time.Duration // 非流式操作超时
|
||||
|
||||
dirMu sync.RWMutex
|
||||
knownDirs map[string]struct{} // 已确认存在的远端目录(含根前缀)
|
||||
spacesPool sync.Pool // 256KB 复用缓冲
|
||||
}
|
||||
|
||||
// NewWebDAVStorage 构造 WebDAV 引擎。
|
||||
func NewWebDAVStorage(opts WebDAVOptions) (*WebDAVStorage, error) {
|
||||
opts.applyDefaults()
|
||||
raw := strings.TrimSpace(opts.BaseURL)
|
||||
if raw == "" {
|
||||
return nil, fmt.Errorf("storage/webdav: 缺少 webdav_url 配置")
|
||||
}
|
||||
if !strings.Contains(raw, "://") {
|
||||
raw = "http://" + raw
|
||||
}
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("storage/webdav: webdav_url 非法: %w", err)
|
||||
}
|
||||
if u.Scheme != "http" && u.Scheme != "https" {
|
||||
return nil, fmt.Errorf("storage/webdav: webdav_url 仅支持 http/https,收到 %q", u.Scheme)
|
||||
}
|
||||
if !strings.HasSuffix(u.Path, "/") {
|
||||
u.Path += "/"
|
||||
}
|
||||
root := strings.Trim(opts.RootPath, "/")
|
||||
if root == "" {
|
||||
root = "filebox_storage"
|
||||
}
|
||||
root = strings.ReplaceAll(root, "\\", "/")
|
||||
transport := newPooledTransport(opts.MaxIdleConnsPerHost)
|
||||
return &WebDAVStorage{
|
||||
base: u,
|
||||
root: root,
|
||||
username: opts.Username,
|
||||
password: opts.Password,
|
||||
client: &http.Client{Transport: transport},
|
||||
transport: transport,
|
||||
auth: newAuthState(opts.Username, opts.Password),
|
||||
maxRetries: opts.MaxRetries,
|
||||
baseBackoff: time.Duration(opts.BaseBackoff) * time.Millisecond,
|
||||
opTimeout: time.Duration(opts.Timeout) * time.Second,
|
||||
knownDirs: map[string]struct{}{},
|
||||
spacesPool: sync.Pool{New: func() any {
|
||||
b := make([]byte, localChunkSize)
|
||||
return &b
|
||||
}},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func init() {
|
||||
RegisterEngine("webdav", func(ctx context.Context) (Storage, error) {
|
||||
return NewWebDAVStorage(engineOptions.WebDAV)
|
||||
})
|
||||
}
|
||||
|
||||
// newPooledTransport 连接池化 Transport:Keep-Alive 连接复用是 WebDAV 优化的核心。
|
||||
func newPooledTransport(maxIdlePerHost int) *http.Transport {
|
||||
return &http.Transport{
|
||||
Proxy: http.ProxyFromEnvironment,
|
||||
DialContext: (&net.Dialer{
|
||||
Timeout: 10 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}).DialContext,
|
||||
ForceAttemptHTTP2: true,
|
||||
MaxIdleConns: 100,
|
||||
MaxIdleConnsPerHost: maxIdlePerHost,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: time.Second,
|
||||
ResponseHeaderTimeout: 60 * time.Second,
|
||||
}
|
||||
}
|
||||
|
||||
// requestOpts 单次 WebDAV 请求参数。
|
||||
type requestOpts struct {
|
||||
// body 请求体工厂:每次尝试调用一次(重试时重新获取,可重放)。
|
||||
body func() (io.Reader, int64, error)
|
||||
// retryBody 请求体是否可重放(seekable);false 时 PUT 类请求失败不重试。
|
||||
retryBody bool
|
||||
// headers 附加请求头。
|
||||
headers map[string]string
|
||||
// streaming 流式传输(GET/PUT 大 body):不套 opTimeout,由调用方 ctx 管控。
|
||||
streaming bool
|
||||
}
|
||||
|
||||
// do 执行一次 WebDAV 请求:认证自动协商 + 指数退避重试。
|
||||
// 返回的响应由调用方负责关闭(drainClose / readErrorBody)。
|
||||
//
|
||||
// 重要:非流式操作的可配超时通过 ctx 实现,cancel 不随 do() 返回而调用,
|
||||
// 而是挂在 davResponse 上、待响应体读完后再触发——否则取消会提前杀掉
|
||||
// Keep-Alive 连接,破坏连接复用。
|
||||
func (w *WebDAVStorage) do(ctx context.Context, method, rawURL string, opts requestOpts) (*davResponse, error) {
|
||||
// 非流式操作套可配超时(流式由调用方 ctx 管控)。
|
||||
var cancel context.CancelFunc
|
||||
if !opts.streaming {
|
||||
if _, hasDeadline := ctx.Deadline(); !hasDeadline {
|
||||
ctx, cancel = context.WithTimeout(ctx, w.opTimeout)
|
||||
}
|
||||
}
|
||||
fail := func(err error) (*davResponse, error) {
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
// 幂等方法或可重放 body 才允许整体重试。
|
||||
idempotent := method == http.MethodGet || method == http.MethodHead ||
|
||||
method == "PROPFIND" || method == "MKCOL" || method == http.MethodDelete ||
|
||||
method == http.MethodOptions
|
||||
retryable := idempotent || opts.retryBody
|
||||
|
||||
const maxAuthRetries = 2
|
||||
budget := w.maxRetries + maxAuthRetries // 认证挑战重试不消耗退避预算
|
||||
authRetries := 0
|
||||
for attempt := 0; attempt < budget; attempt++ {
|
||||
var body io.Reader
|
||||
var length int64 = -1
|
||||
if opts.body != nil {
|
||||
var err error
|
||||
body, length, err = opts.body()
|
||||
if err != nil {
|
||||
return fail(fmt.Errorf("%w: 构造请求体失败: %v", ErrUnavailable, err))
|
||||
}
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, method, rawURL, body)
|
||||
if err != nil {
|
||||
return fail(fmt.Errorf("%w: 构造请求失败: %v", ErrInvalidPath, err))
|
||||
}
|
||||
if length >= 0 {
|
||||
req.ContentLength = length
|
||||
}
|
||||
for k, v := range opts.headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
w.auth.apply(req)
|
||||
resp, err := w.client.Do(req)
|
||||
if err != nil {
|
||||
if ctx.Err() != nil { // 调用方取消/超时优先
|
||||
return fail(ctx.Err())
|
||||
}
|
||||
if retryable && attempt+1 < budget {
|
||||
if sleepErr := w.backoff(ctx, attempt, 0); sleepErr != nil {
|
||||
return fail(sleepErr)
|
||||
}
|
||||
continue
|
||||
}
|
||||
return fail(fmt.Errorf("%w: %s %s: %v", ErrUnavailable, method, rawURL, err))
|
||||
}
|
||||
// 401 认证挑战:切换 Basic/Digest 后立即重试(不退避、不额外计数)。
|
||||
if resp.StatusCode == http.StatusUnauthorized && authRetries < maxAuthRetries {
|
||||
challenge := resp.Header.Get("WWW-Authenticate")
|
||||
drainClose(&davResponse{Response: resp})
|
||||
if challenge != "" && w.auth.challenge(challenge) {
|
||||
authRetries++
|
||||
continue
|
||||
}
|
||||
return fail(fmt.Errorf("%w: WebDAV 认证失败(401,%s)", ErrUnavailable, rawURL))
|
||||
}
|
||||
// 5xx/429/408:幂等或可重放 body 时指数退避重试。
|
||||
if retryable && isRetryStatus(resp.StatusCode) && attempt+1 < budget {
|
||||
retryAfter := retryAfterSeconds(resp.Header.Get("Retry-After"))
|
||||
drainClose(&davResponse{Response: resp})
|
||||
if sleepErr := w.backoff(ctx, attempt, retryAfter); sleepErr != nil {
|
||||
return fail(sleepErr)
|
||||
}
|
||||
continue
|
||||
}
|
||||
return &davResponse{Response: resp, cancel: cancel}, nil
|
||||
}
|
||||
return fail(fmt.Errorf("%w: WebDAV 重试耗尽(%s %s)", ErrUnavailable, method, rawURL))
|
||||
}
|
||||
|
||||
// davResponse WebDAV 响应 + 关联的超时取消函数。
|
||||
// 非流式操作读完响应体后必须经 drainClose/readErrorBody 释放(触发 cancel)。
|
||||
type davResponse struct {
|
||||
*http.Response
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
// isRetryStatus 判断状态码是否值得重试。
|
||||
func isRetryStatus(code int) bool {
|
||||
switch code {
|
||||
case http.StatusRequestTimeout, http.StatusTooManyRequests,
|
||||
http.StatusInternalServerError, http.StatusBadGateway,
|
||||
http.StatusServiceUnavailable, http.StatusGatewayTimeout:
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// retryAfterSeconds 解析 Retry-After(秒);非法或负值返回 0。
|
||||
func retryAfterSeconds(v string) time.Duration {
|
||||
if v == "" {
|
||||
return 0
|
||||
}
|
||||
n, err := strconv.Atoi(strings.TrimSpace(v))
|
||||
if err != nil || n <= 0 {
|
||||
return 0
|
||||
}
|
||||
if n > 5 {
|
||||
n = 5 // 上限 5s,避免异常服务端拖死请求
|
||||
}
|
||||
return time.Duration(n) * time.Second
|
||||
}
|
||||
|
||||
// backoff 指数退避:base * 2^attempt,封顶 2s,带 ±20% 抖动;retryAfter 优先。
|
||||
func (w *WebDAVStorage) backoff(ctx context.Context, attempt int, retryAfter time.Duration) error {
|
||||
d := retryAfter
|
||||
if d <= 0 {
|
||||
d = w.baseBackoff << attempt
|
||||
if d > 2*time.Second {
|
||||
d = 2 * time.Second
|
||||
}
|
||||
// ±20% 抖动
|
||||
jitter := time.Duration(int64(d) / 5)
|
||||
if jitter > 0 {
|
||||
d -= time.Duration(rand.Int63n(int64(jitter)))
|
||||
}
|
||||
}
|
||||
if d <= 0 {
|
||||
d = time.Millisecond
|
||||
}
|
||||
timer := time.NewTimer(d)
|
||||
defer timer.Stop()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-timer.C:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// drainClose 读取少量残余并关闭响应体,保证连接可复用;随后触发超时清理。
|
||||
func drainClose(resp *davResponse) {
|
||||
if resp == nil || resp.Body == nil {
|
||||
return
|
||||
}
|
||||
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 8<<10))
|
||||
_ = resp.Body.Close()
|
||||
if resp.cancel != nil {
|
||||
resp.cancel()
|
||||
}
|
||||
}
|
||||
|
||||
// joinRemote 校验 savePath 并拼接远端完整路径(含根目录前缀)。
|
||||
func (w *WebDAVStorage) joinRemote(savePath string) (string, error) {
|
||||
cleaned, ok := SanitizePath(savePath)
|
||||
if !ok {
|
||||
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
|
||||
}
|
||||
return path.Join(w.root, cleaned), nil
|
||||
}
|
||||
|
||||
// urlFor 将远端路径转为完整 URL(URL.String 自动按段转义)。
|
||||
func (w *WebDAVStorage) urlFor(remotePath string) string {
|
||||
u := *w.base
|
||||
p := strings.TrimSuffix(u.Path, "/")
|
||||
remotePath = strings.Trim(remotePath, "/")
|
||||
if remotePath != "" && remotePath != "." {
|
||||
p += "/" + remotePath
|
||||
}
|
||||
u.Path = p
|
||||
return u.String()
|
||||
}
|
||||
|
||||
// propfindBody PROPFIND 请求体:只取需要的属性。
|
||||
const propfindBody = `<?xml version="1.0" encoding="utf-8"?>` +
|
||||
`<D:propfind xmlns:D="DAV:"><D:prop>` +
|
||||
`<D:resourcetype/><D:getcontentlength/><D:getcontenttype/>` +
|
||||
`</D:prop></D:propfind>`
|
||||
|
||||
// davMultistatus 207 Multi-Status XML 解析结构(标签名与命名空间无关匹配)。
|
||||
type davMultistatus struct {
|
||||
Responses []struct {
|
||||
Href string `xml:"href"`
|
||||
Propstat []struct {
|
||||
Status string `xml:"status"`
|
||||
Prop struct {
|
||||
ContentLength int64 `xml:"getcontentlength"`
|
||||
ContentType string `xml:"getcontenttype"`
|
||||
ResourceType struct {
|
||||
Collection *struct{} `xml:"collection"`
|
||||
} `xml:"resourcetype"`
|
||||
} `xml:"prop"`
|
||||
} `xml:"propstat"`
|
||||
} `xml:"response"`
|
||||
}
|
||||
|
||||
// firstProp 取第一个 HTTP 2xx 状态的属性块。
|
||||
func (m *davMultistatus) firstProp() (length int64, ctype string, isDir bool, ok bool) {
|
||||
for _, r := range m.Responses {
|
||||
for _, ps := range r.Propstat {
|
||||
if !strings.Contains(ps.Status, " 200 ") {
|
||||
continue
|
||||
}
|
||||
return ps.Prop.ContentLength, ps.Prop.ContentType, ps.Prop.ResourceType.Collection != nil, true
|
||||
}
|
||||
}
|
||||
return 0, "", false, false
|
||||
}
|
||||
|
||||
// propfind 执行 PROPFIND 并解析 207 响应;404 时返回 (nil, nil)。
|
||||
func (w *WebDAVStorage) propfind(ctx context.Context, rawURL string, depth string) (*davMultistatus, error) {
|
||||
resp, err := w.do(ctx, "PROPFIND", rawURL, requestOpts{
|
||||
body: func() (io.Reader, int64, error) {
|
||||
return strings.NewReader(propfindBody), int64(len(propfindBody)), nil
|
||||
},
|
||||
headers: map[string]string{"Depth": depth, "Content-Type": "application/xml"},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer drainClose(resp)
|
||||
switch resp.StatusCode {
|
||||
case http.StatusMultiStatus, http.StatusOK:
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: PROPFIND 读取失败: %v", ErrUnavailable, err)
|
||||
}
|
||||
var ms davMultistatus
|
||||
if err := xml.Unmarshal(body, &ms); err != nil {
|
||||
return nil, fmt.Errorf("%w: PROPFIND XML 解析失败: %v", ErrUnavailable, err)
|
||||
}
|
||||
return &ms, nil
|
||||
case http.StatusNotFound:
|
||||
return nil, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("%w: PROPFIND %s → %d", ErrUnavailable, rawURL, resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
// remoteExists PROPFIND 探测远端路径存在性。
|
||||
func (w *WebDAVStorage) remoteExists(ctx context.Context, remotePath string) (bool, error) {
|
||||
ms, err := w.propfind(ctx, w.urlFor(remotePath), "0")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return ms != nil, nil
|
||||
}
|
||||
|
||||
// markDir 记录已确认存在的目录(避免重复 PROPFIND/MKCOL 往返)。
|
||||
func (w *WebDAVStorage) markDir(remotePath string) {
|
||||
w.dirMu.Lock()
|
||||
defer w.dirMu.Unlock()
|
||||
w.knownDirs[remotePath] = struct{}{}
|
||||
}
|
||||
|
||||
// unmarkDir 目录被删除时移除缓存。
|
||||
func (w *WebDAVStorage) unmarkDir(remotePath string) {
|
||||
w.dirMu.Lock()
|
||||
defer w.dirMu.Unlock()
|
||||
delete(w.knownDirs, remotePath)
|
||||
}
|
||||
|
||||
// isMarkedDir 查询目录缓存。
|
||||
func (w *WebDAVStorage) isMarkedDir(remotePath string) bool {
|
||||
w.dirMu.RLock()
|
||||
defer w.dirMu.RUnlock()
|
||||
_, ok := w.knownDirs[remotePath]
|
||||
return ok
|
||||
}
|
||||
|
||||
// ensureDirs 按需逐级创建远端目录(含根前缀;MKCOL 级联,成功后写缓存)。
|
||||
func (w *WebDAVStorage) ensureDirs(ctx context.Context, remotePath string) error {
|
||||
segments := splitRemoteSegments(remotePath)
|
||||
cur := ""
|
||||
for _, seg := range segments {
|
||||
cur = path.Join(cur, seg)
|
||||
if w.isMarkedDir(cur) {
|
||||
continue
|
||||
}
|
||||
exists, err := w.remoteExists(ctx, cur)
|
||||
if err == nil && exists {
|
||||
w.markDir(cur)
|
||||
continue
|
||||
}
|
||||
if err != nil && !errors.Is(err, ErrNotFound) {
|
||||
return err
|
||||
}
|
||||
resp, err := w.do(ctx, "MKCOL", w.urlFor(cur), requestOpts{})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
status := resp.StatusCode
|
||||
drainClose(resp)
|
||||
// 201 创建成功;405 已存在;其余视为失败(409 通常因父目录缺失,理论上不会出现)。
|
||||
if status == http.StatusCreated || status == http.StatusOK ||
|
||||
status == http.StatusNoContent || status == http.StatusMethodNotAllowed {
|
||||
w.markDir(cur)
|
||||
continue
|
||||
}
|
||||
return fmt.Errorf("%w: MKCOL %s → %d", ErrUnavailable, cur, status)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// splitRemoteSegments 拆分远端路径段。
|
||||
func splitRemoteSegments(p string) []string {
|
||||
p = strings.Trim(strings.ReplaceAll(p, "\\", "/"), "/")
|
||||
if p == "" {
|
||||
return nil
|
||||
}
|
||||
return strings.Split(p, "/")
|
||||
}
|
||||
|
||||
// deleteEmptyParents 删除空父目录(含根前缀,但不删根目录本身);尽力而为。
|
||||
func (w *WebDAVStorage) deleteEmptyParents(ctx context.Context, remotePath string) {
|
||||
dir := path.Dir(remotePath)
|
||||
for dir != "" && dir != "." && dir != w.root && strings.HasPrefix(dir+"/", w.root+"/") {
|
||||
ms, err := w.propfind(ctx, w.urlFor(dir), "1")
|
||||
if err != nil || ms == nil {
|
||||
return
|
||||
}
|
||||
if len(ms.Responses) > 1 { // 非空(自身 + 子项)
|
||||
return
|
||||
}
|
||||
resp, err := w.do(ctx, http.MethodDelete, w.urlFor(dir), requestOpts{})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
ok := resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusNoContent
|
||||
drainClose(resp)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
w.unmarkDir(dir)
|
||||
dir = path.Dir(dir)
|
||||
}
|
||||
}
|
||||
|
||||
// putFile PUT 上传:body 工厂每次尝试返回可重放的读取器。
|
||||
func (w *WebDAVStorage) putFile(ctx context.Context, rawURL string, body func() (io.Reader, int64, error), retryBody bool) (*davResponse, error) {
|
||||
return w.do(ctx, http.MethodPut, rawURL, requestOpts{
|
||||
body: body,
|
||||
retryBody: retryBody,
|
||||
headers: map[string]string{"Content-Type": "application/octet-stream"},
|
||||
streaming: true,
|
||||
})
|
||||
}
|
||||
|
||||
// checkPutStatus 校验 PUT 响应状态。
|
||||
func checkPutStatus(resp *davResponse, op string) error {
|
||||
switch resp.StatusCode {
|
||||
case http.StatusOK, http.StatusCreated, http.StatusNoContent:
|
||||
drainClose(resp)
|
||||
return nil
|
||||
default:
|
||||
msg := readErrorBody(resp)
|
||||
return fmt.Errorf("%w: %s → %d %s", ErrUnavailable, op, resp.StatusCode, msg)
|
||||
}
|
||||
}
|
||||
|
||||
// readErrorBody 读取错误响应前 200 字节并释放连接。
|
||||
func readErrorBody(resp *davResponse) string {
|
||||
if resp == nil || resp.Body == nil {
|
||||
return ""
|
||||
}
|
||||
b, _ := io.ReadAll(io.LimitReader(resp.Body, 200))
|
||||
_ = resp.Body.Close()
|
||||
if resp.cancel != nil {
|
||||
resp.cancel()
|
||||
}
|
||||
return strings.TrimSpace(string(b))
|
||||
}
|
||||
|
||||
// SaveFile 流式保存(PUT):按需建目录,seekable 源可安全重试。
|
||||
func (w *WebDAVStorage) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
|
||||
remote, err := w.joinRemote(savePath)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
if err := w.ensureDirs(ctx, path.Dir(remote)); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
// 可重放判定:seekable 源失败后可从头重传(PUT 覆盖语义保证最终一致)。
|
||||
seeker, seekable := r.(io.Seeker)
|
||||
var knownLen int64 = -1
|
||||
if seekable {
|
||||
if cur, err := seeker.Seek(0, io.SeekCurrent); err == nil {
|
||||
if end, err := seeker.Seek(0, io.SeekEnd); err == nil {
|
||||
knownLen = end - cur
|
||||
_, _ = seeker.Seek(cur, io.SeekStart)
|
||||
}
|
||||
}
|
||||
}
|
||||
src := &countingReader{r: r}
|
||||
body := func() (io.Reader, int64, error) {
|
||||
if seekable {
|
||||
if _, err := seeker.Seek(0, io.SeekStart); err != nil {
|
||||
return nil, -1, err
|
||||
}
|
||||
src.reset()
|
||||
}
|
||||
return src, knownLen, nil
|
||||
}
|
||||
resp, err := w.putFile(ctx, w.urlFor(remote), body, seekable)
|
||||
if err != nil {
|
||||
return src.count(), err
|
||||
}
|
||||
if err := checkPutStatus(resp, fmt.Sprintf("PUT %s", remote)); err != nil {
|
||||
return src.count(), err
|
||||
}
|
||||
return src.count(), nil
|
||||
}
|
||||
|
||||
// DeleteFile DELETE 文件 + 尽力清理空父目录。
|
||||
func (w *WebDAVStorage) DeleteFile(ctx context.Context, savePath string) error {
|
||||
remote, err := w.joinRemote(savePath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
resp, err := w.do(ctx, http.MethodDelete, w.urlFor(remote), requestOpts{})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
switch resp.StatusCode {
|
||||
case http.StatusOK, http.StatusNoContent, http.StatusNotFound:
|
||||
drainClose(resp)
|
||||
default:
|
||||
msg := readErrorBody(resp)
|
||||
return fmt.Errorf("%w: DELETE %s → %d %s", ErrUnavailable, remote, resp.StatusCode, msg)
|
||||
}
|
||||
w.deleteEmptyParents(ctx, remote)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Open 打开下载流:Range 透传,io.Pipe 流式转发不落盘,ctx 可取消。
|
||||
func (w *WebDAVStorage) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
|
||||
remote, err := w.joinRemote(savePath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
opts := requestOpts{streaming: true}
|
||||
if rng != nil {
|
||||
opts.headers = map[string]string{"Range": rangeHeaderValue(rng)}
|
||||
}
|
||||
resp, err := w.do(ctx, http.MethodGet, w.urlFor(remote), opts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch resp.StatusCode {
|
||||
case http.StatusOK, http.StatusPartialContent:
|
||||
// 正常,继续
|
||||
case http.StatusNotFound:
|
||||
drainClose(resp)
|
||||
return nil, ErrNotFound
|
||||
case http.StatusRequestedRangeNotSatisfiable:
|
||||
drainClose(resp)
|
||||
return nil, ErrRangeNotSatisfiable
|
||||
default:
|
||||
msg := readErrorBody(resp)
|
||||
return nil, fmt.Errorf("%w: GET %s → %d %s", ErrUnavailable, remote, resp.StatusCode, msg)
|
||||
}
|
||||
|
||||
total := resp.ContentLength
|
||||
start, end := int64(0), total-1
|
||||
if resp.StatusCode == http.StatusPartialContent {
|
||||
if cr := resp.Header.Get("Content-Range"); cr != "" {
|
||||
if s0, e0, t0, ok := parseContentRange(cr); ok {
|
||||
start, end = s0, e0
|
||||
if t0 >= 0 {
|
||||
total = t0
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if total < 0 { // 服务端未给出长度(chunked):按未知大小处理
|
||||
start, end, total = 0, -1, -1
|
||||
}
|
||||
if rng == nil { // 对齐契约:完整文件 Start=0、End=Total-1
|
||||
start, end = 0, total-1
|
||||
}
|
||||
if end < 0 { // 空文件或未知大小:End 未知语义
|
||||
end = -1
|
||||
}
|
||||
|
||||
// io.Pipe 流式桥接:HTTP 响应体 → 管道 → 调用方,全程不落盘;
|
||||
// 调用方提前 Close 或 ctx 取消都会终止拷贝并释放连接。
|
||||
body := resp.Body
|
||||
pr, pw := io.Pipe()
|
||||
go func() {
|
||||
bufp, _ := w.spacesPool.Get().(*[]byte)
|
||||
_, copyErr := io.CopyBuffer(pw, body, *bufp)
|
||||
w.spacesPool.Put(bufp)
|
||||
_ = body.Close()
|
||||
pw.CloseWithError(copyErr) // copyErr 为 nil 时写入 EOF
|
||||
}()
|
||||
context.AfterFunc(ctx, func() {
|
||||
_ = pw.CloseWithError(ctx.Err())
|
||||
})
|
||||
|
||||
contentType := resp.Header.Get("Content-Type")
|
||||
return &Download{
|
||||
ReadCloser: pr,
|
||||
Start: start,
|
||||
End: end,
|
||||
Total: total,
|
||||
Meta: FileMeta{
|
||||
Size: total,
|
||||
ContentType: contentType,
|
||||
AcceptRanges: true,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Stat PROPFIND Depth 0 获取元信息。
|
||||
func (w *WebDAVStorage) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
|
||||
remote, err := w.joinRemote(savePath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ms, err := w.propfind(ctx, w.urlFor(remote), "0")
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrNotFound) {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if ms == nil {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
length, ctype, _, ok := ms.firstProp()
|
||||
if !ok {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return &FileMeta{Size: length, ContentType: ctype, AcceptRanges: true}, nil
|
||||
}
|
||||
|
||||
// HeadMeta 读取文件元信息与前 n 字节(WebDAV 实现:PROPFIND + Range GET)。
|
||||
func (w *WebDAVStorage) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
|
||||
meta, err := w.Stat(ctx, savePath)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if headBytes <= 0 {
|
||||
return meta, nil, nil
|
||||
}
|
||||
dl, err := w.Open(ctx, savePath, &Range{Start: 0, End: headBytes - 1})
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrRangeNotSatisfiable) { // 空文件等边界:返回空头
|
||||
return meta, nil, nil
|
||||
}
|
||||
return nil, nil, err
|
||||
}
|
||||
defer func() { _ = dl.Close() }()
|
||||
head := make([]byte, headBytes)
|
||||
n, _ := io.ReadFull(dl.ReadCloser, head)
|
||||
return meta, head[:n], nil
|
||||
}
|
||||
|
||||
// SaveChunk 保存分片:落临时文件获得精确长度与可重放 body,PUT 到分片路径。
|
||||
func (w *WebDAVStorage) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
|
||||
if _, err := w.joinRemote(savePath); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
chunkRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, chunkIndex))
|
||||
if !ok {
|
||||
return 0, fmt.Errorf("%w: 分片路径非法 %q", ErrInvalidPath, savePath)
|
||||
}
|
||||
remote := path.Join(w.root, chunkRel)
|
||||
if err := w.ensureDirs(ctx, path.Dir(remote)); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
// 分片体积有限(默认 ≤8MB):落临时文件换取精确 Content-Length 与可重试性。
|
||||
tmp, err := os.CreateTemp("", "fcb-webdav-chunk-*")
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("storage/webdav: 创建临时文件失败: %w", err)
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }()
|
||||
size, err := io.CopyBuffer(tmp, r, make([]byte, localChunkSize))
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("storage/webdav: 缓存分片失败: %w", err)
|
||||
}
|
||||
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
|
||||
return 0, fmt.Errorf("storage/webdav: 回卷分片失败: %w", err)
|
||||
}
|
||||
body := func() (io.Reader, int64, error) {
|
||||
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
|
||||
return nil, -1, err
|
||||
}
|
||||
return tmp, size, nil
|
||||
}
|
||||
resp, err := w.putFile(ctx, w.urlFor(remote), body, true)
|
||||
if err != nil {
|
||||
return size, err
|
||||
}
|
||||
if err := checkPutStatus(resp, fmt.Sprintf("PUT 分片 %s", remote)); err != nil {
|
||||
return size, err
|
||||
}
|
||||
return size, nil
|
||||
}
|
||||
|
||||
// MergeChunks 合并 WebDAV 分片:
|
||||
// 逐分片 GET 流式拼入本地临时文件(边拷贝边校验哈希)→ PUT 上传目标 → 清理远端分片与本地临时文件。
|
||||
// 说明:WebDAV 无服务端聚合能力,合并必须经服务端中转;临时文件仅用于拼接与重试,最终 PUT 可重放。
|
||||
func (w *WebDAVStorage) 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/webdav: 非法分片总数 %d", total)
|
||||
}
|
||||
remote, err := w.joinRemote(savePath)
|
||||
if err != nil {
|
||||
return 0, "", err
|
||||
}
|
||||
if err := w.ensureDirs(ctx, path.Dir(remote)); err != nil {
|
||||
return 0, "", err
|
||||
}
|
||||
tmp, err := os.CreateTemp("", "fcb-webdav-merge-*")
|
||||
if err != nil {
|
||||
return 0, "", fmt.Errorf("storage/webdav: 创建合并临时文件失败: %w", err)
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }()
|
||||
|
||||
totalHash := sha256.New()
|
||||
var size int64
|
||||
for i := 0; i < total; i++ {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return 0, "", err
|
||||
}
|
||||
chunkRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, i))
|
||||
if !ok {
|
||||
return 0, "", fmt.Errorf("%w: 分片 %d 路径非法", ErrInvalidPath, i)
|
||||
}
|
||||
resp, err := w.do(ctx, http.MethodGet, w.urlFor(path.Join(w.root, chunkRel)), requestOpts{streaming: true})
|
||||
if err != nil {
|
||||
return 0, "", fmt.Errorf("storage/webdav: 读取分片 %d 失败: %w", i, err)
|
||||
}
|
||||
if resp.StatusCode == http.StatusNotFound {
|
||||
drainClose(resp)
|
||||
return 0, "", fmt.Errorf("storage/webdav: 分片 %d 不存在", i)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusPartialContent {
|
||||
msg := readErrorBody(resp)
|
||||
return 0, "", fmt.Errorf("storage/webdav: 读取分片 %d → %d %s", i, resp.StatusCode, msg)
|
||||
}
|
||||
chunkHash := sha256.New()
|
||||
n, err := io.CopyBuffer(io.MultiWriter(tmp, totalHash, chunkHash), resp.Body, make([]byte, localChunkSize))
|
||||
_ = resp.Body.Close()
|
||||
if err != nil {
|
||||
return 0, "", fmt.Errorf("storage/webdav: 拼接分片 %d 失败: %w", i, err)
|
||||
}
|
||||
if verifyHash != nil {
|
||||
expected, err := verifyHash(i)
|
||||
if err != nil {
|
||||
return 0, "", err
|
||||
}
|
||||
if expected != "" && expected != hex.EncodeToString(chunkHash.Sum(nil)) {
|
||||
return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, i, expected)
|
||||
}
|
||||
}
|
||||
size += n
|
||||
}
|
||||
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
|
||||
return 0, "", fmt.Errorf("storage/webdav: 回卷合并文件失败: %w", err)
|
||||
}
|
||||
body := func() (io.Reader, int64, error) {
|
||||
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
|
||||
return nil, -1, err
|
||||
}
|
||||
return tmp, size, nil
|
||||
}
|
||||
resp, err := w.putFile(ctx, w.urlFor(remote), body, true)
|
||||
if err != nil {
|
||||
return size, "", err
|
||||
}
|
||||
if err := checkPutStatus(resp, fmt.Sprintf("PUT 合并 %s", remote)); err != nil {
|
||||
return size, "", err
|
||||
}
|
||||
// 合并成功后清理远端分片目录与本地临时文件(defer 兜底删除本地文件)。
|
||||
_ = w.CleanChunks(ctx, uploadID, savePath)
|
||||
return size, hex.EncodeToString(totalHash.Sum(nil)), nil
|
||||
}
|
||||
|
||||
// CleanChunks 递归删除远端分片目录(RFC 4918 DELETE 对 collection 递归)。
|
||||
func (w *WebDAVStorage) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
|
||||
dirRel, ok := SanitizePath(chunkDirOf(savePath, uploadID))
|
||||
if !ok {
|
||||
return fmt.Errorf("%w: 分片目录非法 %q", ErrInvalidPath, savePath)
|
||||
}
|
||||
remote := path.Join(w.root, dirRel)
|
||||
resp, err := w.do(ctx, http.MethodDelete, w.urlFor(remote), requestOpts{})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
switch resp.StatusCode {
|
||||
case http.StatusOK, http.StatusNoContent, http.StatusNotFound:
|
||||
drainClose(resp)
|
||||
w.unmarkDir(remote)
|
||||
default:
|
||||
msg := readErrorBody(resp)
|
||||
return fmt.Errorf("%w: 清理分片目录 %s → %d %s", ErrUnavailable, remote, resp.StatusCode, msg)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// FileExists PROPFIND 探测存在性;非法路径按不存在处理。
|
||||
func (w *WebDAVStorage) FileExists(ctx context.Context, savePath string) (bool, error) {
|
||||
remote, err := w.joinRemote(savePath)
|
||||
if err != nil {
|
||||
return false, nil
|
||||
}
|
||||
return w.remoteExists(ctx, remote)
|
||||
}
|
||||
|
||||
// PresignGetURL WebDAV 无预签名直链能力。
|
||||
func (w *WebDAVStorage) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
|
||||
return "", ErrNotSupported
|
||||
}
|
||||
|
||||
// PresignPutURL WebDAV 无预签名直传能力。
|
||||
func (w *WebDAVStorage) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
|
||||
return "", ErrNotSupported
|
||||
}
|
||||
|
||||
// HealthCheck 健康检查:PROPFIND 根目录;不存在时 MKCOL 创建(启动自愈)。
|
||||
// 同时完成凭据与连通性验证(do 内 401 协商)。
|
||||
func (w *WebDAVStorage) HealthCheck(ctx context.Context) error {
|
||||
exists, err := w.remoteExists(ctx, w.root)
|
||||
if err == nil && exists {
|
||||
w.markDir(w.root)
|
||||
return nil
|
||||
}
|
||||
if err != nil && !errors.Is(err, ErrNotFound) {
|
||||
return fmt.Errorf("%w: WebDAV 健康检查失败: %v", ErrUnavailable, err)
|
||||
}
|
||||
resp, err := w.do(ctx, "MKCOL", w.urlFor(w.root), requestOpts{})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
switch resp.StatusCode {
|
||||
case http.StatusCreated, http.StatusOK, http.StatusNoContent, http.StatusMethodNotAllowed:
|
||||
drainClose(resp)
|
||||
w.markDir(w.root)
|
||||
return nil
|
||||
default:
|
||||
msg := readErrorBody(resp)
|
||||
return fmt.Errorf("%w: WebDAV 根目录创建失败 → %d %s", ErrUnavailable, resp.StatusCode, msg)
|
||||
}
|
||||
}
|
||||
|
||||
// 接口编译期断言。
|
||||
var _ Storage = (*WebDAVStorage)(nil)
|
||||
@@ -0,0 +1,244 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"crypto/md5"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// authMode 认证模式(WebDAV 服务端挑战后自动协商)。
|
||||
type authMode int
|
||||
|
||||
const (
|
||||
authModeUnknown authMode = iota // 未定:先发 Basic 探测
|
||||
authModeBasic
|
||||
authModeDigest
|
||||
)
|
||||
|
||||
// authState WebDAV Basic/Digest 认证状态。
|
||||
//
|
||||
// 策略:
|
||||
// - 首个请求预置 Basic;若服务端 401 且挑战为 Digest,则解析挑战参数切换为 Digest;
|
||||
// - Digest 按 RFC 2617/7616 实现 qop=auth(MD5 / SHA-256,含 -sess 变体);
|
||||
// qop 缺失时回退 RFC 2069 旧式响应;
|
||||
// - nonce 变更时重置 nc 计数;nc/cnonce 在互斥锁内生成保证并发唯一。
|
||||
type authState struct {
|
||||
mu sync.Mutex
|
||||
username string
|
||||
password string
|
||||
mode authMode
|
||||
realm string
|
||||
nonce string
|
||||
qop string // 选定的 qop("auth" 或空 = RFC2069)
|
||||
opaque string
|
||||
algorithm string // MD5 | MD5-sess | SHA-256 | SHA-256-sess
|
||||
nc uint32
|
||||
knownBasicOK bool // 已确认 Basic 可用
|
||||
}
|
||||
|
||||
// newAuthState 构造认证状态(默认以 Basic 起步)。
|
||||
func newAuthState(username, password string) *authState {
|
||||
return &authState{username: username, password: password}
|
||||
}
|
||||
|
||||
// apply 为请求设置 Authorization 头(每次请求调用,Digest 时消耗一个 nc)。
|
||||
func (a *authState) apply(req *http.Request) {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
switch {
|
||||
case a.mode == authModeDigest && a.nonce != "":
|
||||
req.Header.Set("Authorization", a.digestHeader(req))
|
||||
default:
|
||||
req.SetBasicAuth(a.username, a.password)
|
||||
}
|
||||
}
|
||||
|
||||
// digestHeader 依据缓存的挑战参数计算 Digest Authorization 头(调用方需持锁)。
|
||||
func (a *authState) digestHeader(req *http.Request) string {
|
||||
uri := req.URL.RequestURI()
|
||||
method := strings.ToUpper(req.Method)
|
||||
ncStr := fmt.Sprintf("%08x", a.nc+1)
|
||||
a.nc++
|
||||
cnonce := randomHex(8)
|
||||
|
||||
var ha1 string
|
||||
switch strings.ToLower(a.algorithm) {
|
||||
case "md5-sess":
|
||||
ha1 = hashHex("md5", hashHex("md5", a.username+":"+a.realm+":"+a.password)+":"+a.nonce+":"+cnonce)
|
||||
case "sha-256-sess":
|
||||
ha1 = hashHex("sha256", hashHex("sha256", a.username+":"+a.realm+":"+a.password)+":"+a.nonce+":"+cnonce)
|
||||
case "sha-256":
|
||||
ha1 = hashHex("sha256", a.username+":"+a.realm+":"+a.password)
|
||||
default: // md5
|
||||
ha1 = hashHex("md5", a.username+":"+a.realm+":"+a.password)
|
||||
}
|
||||
ha2 := hashHex(algoName(a.algorithm), method+":"+uri)
|
||||
|
||||
var response string
|
||||
var fields []string
|
||||
esc := escapeDigestValue(a.username)
|
||||
if a.qop == "" { // RFC 2069
|
||||
response = hashHex(algoName(a.algorithm), ha1+":"+a.nonce+":"+ha2)
|
||||
fields = append(fields,
|
||||
`Digest username="`+esc+`"`,
|
||||
`realm="`+escapeDigestValue(a.realm)+`"`,
|
||||
`nonce="`+escapeDigestValue(a.nonce)+`"`,
|
||||
`uri="`+escapeDigestValue(uri)+`"`,
|
||||
`response="`+response+`"`)
|
||||
} else {
|
||||
response = hashHex(algoName(a.algorithm), ha1+":"+a.nonce+":"+ncStr+":"+cnonce+":"+a.qop+":"+ha2)
|
||||
fields = append(fields,
|
||||
`Digest username="`+esc+`"`,
|
||||
`realm="`+escapeDigestValue(a.realm)+`"`,
|
||||
`nonce="`+escapeDigestValue(a.nonce)+`"`,
|
||||
`uri="`+escapeDigestValue(uri)+`"`,
|
||||
`cnonce="`+cnonce+`"`,
|
||||
`nc=`+ncStr,
|
||||
`qop=`+a.qop,
|
||||
`response="`+response+`"`,
|
||||
`algorithm=`+a.algorithm)
|
||||
}
|
||||
if a.opaque != "" {
|
||||
fields = append(fields, `opaque="`+escapeDigestValue(a.opaque)+`"`)
|
||||
}
|
||||
return strings.Join(fields, ", ")
|
||||
}
|
||||
|
||||
// challenge 处理 401 的 WWW-Authenticate 挑战;返回是否已切换认证方式可重试。
|
||||
// 返回 false 表示凭据错误或算法不受支持,调用方应直接报错。
|
||||
func (a *authState) challenge(header string) bool {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
h := strings.TrimSpace(header)
|
||||
lower := strings.ToLower(h)
|
||||
switch {
|
||||
case strings.HasPrefix(lower, "digest"):
|
||||
params := parseChallengeParams(strings.TrimPrefix(h[len("Digest"):], " "))
|
||||
algo := strings.ToUpper(strings.TrimSpace(params["algorithm"]))
|
||||
if algo == "" {
|
||||
algo = "MD5"
|
||||
}
|
||||
switch algo {
|
||||
case "MD5", "MD5-SESS", "SHA-256", "SHA-256-SESS":
|
||||
default:
|
||||
return false // 不支持的摘要算法
|
||||
}
|
||||
if params["nonce"] == "" || params["realm"] == "" {
|
||||
return false
|
||||
}
|
||||
qop := ""
|
||||
if raw := strings.TrimSpace(params["qop"]); raw != "" {
|
||||
for _, candidate := range strings.Split(raw, ",") {
|
||||
if strings.EqualFold(strings.TrimSpace(candidate), "auth") {
|
||||
qop = "auth"
|
||||
break
|
||||
}
|
||||
}
|
||||
if qop == "" {
|
||||
return false // 仅支持 auth-int 等需要 body 哈希的模式
|
||||
}
|
||||
}
|
||||
if a.nonce != params["nonce"] {
|
||||
a.nc = 0
|
||||
}
|
||||
a.realm, a.nonce, a.qop = params["realm"], params["nonce"], qop
|
||||
a.opaque, a.algorithm = params["opaque"], strings.ToLower(algo)
|
||||
a.mode = authModeDigest
|
||||
a.knownBasicOK = false
|
||||
return true
|
||||
case strings.HasPrefix(lower, "basic"):
|
||||
if a.knownBasicOK || a.mode == authModeBasic {
|
||||
return false // 已用 Basic 仍 401:凭据错误
|
||||
}
|
||||
a.mode = authModeBasic
|
||||
a.knownBasicOK = true
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// parseChallengeParams 解析 "realm=\"x\", nonce=\"y\"" 形式的挑战参数(引号内逗号不切分)。
|
||||
func parseChallengeParams(s string) map[string]string {
|
||||
out := map[string]string{}
|
||||
for _, item := range splitAuthParams(s) {
|
||||
kv := strings.SplitN(item, "=", 2)
|
||||
if len(kv) != 2 {
|
||||
continue
|
||||
}
|
||||
k := strings.ToLower(strings.TrimSpace(kv[0]))
|
||||
v := strings.TrimSpace(kv[1])
|
||||
if len(v) >= 2 && strings.HasPrefix(v, `"`) && strings.HasSuffix(v, `"`) {
|
||||
v = v[1 : len(v)-1]
|
||||
}
|
||||
out[k] = v
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// splitAuthParams 逗号切分但忽略引号内的逗号。
|
||||
func splitAuthParams(s string) []string {
|
||||
var parts []string
|
||||
var b strings.Builder
|
||||
inQuote := false
|
||||
for i := 0; i < len(s); i++ {
|
||||
c := s[i]
|
||||
switch {
|
||||
case c == '"':
|
||||
inQuote = !inQuote
|
||||
b.WriteByte(c)
|
||||
case c == ',' && !inQuote:
|
||||
if t := strings.TrimSpace(b.String()); t != "" {
|
||||
parts = append(parts, t)
|
||||
}
|
||||
b.Reset()
|
||||
default:
|
||||
b.WriteByte(c)
|
||||
}
|
||||
}
|
||||
if t := strings.TrimSpace(b.String()); t != "" {
|
||||
parts = append(parts, t)
|
||||
}
|
||||
return parts
|
||||
}
|
||||
|
||||
// algoName 映射哈希函数名。
|
||||
func algoName(algorithm string) string {
|
||||
switch strings.ToLower(algorithm) {
|
||||
case "sha-256", "sha-256-sess":
|
||||
return "sha256"
|
||||
default:
|
||||
return "md5"
|
||||
}
|
||||
}
|
||||
|
||||
// hashHex 通用哈希摘要(algo: md5|sha256)。
|
||||
func hashHex(algo, s string) string {
|
||||
if algo == "sha256" {
|
||||
sum := sha256.Sum256([]byte(s))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
sum := md5.Sum([]byte(s))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// escapeDigestValue 转义引号。
|
||||
func escapeDigestValue(s string) string {
|
||||
return strings.ReplaceAll(s, `"`, `\"`)
|
||||
}
|
||||
|
||||
// randomHex 生成 n 字节随机 hex。
|
||||
func randomHex(n int) string {
|
||||
b := make([]byte, n)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
// crypto/rand 失败极其罕见;退化为全零仍保持协议可用。
|
||||
for i := range b {
|
||||
b[i] = 0
|
||||
}
|
||||
}
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
@@ -0,0 +1,824 @@
|
||||
package storage
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ---- 最小 WebDAV 假服务:PUT/GET/HEAD/PROPFIND/MKCOL/DELETE + Basic/Digest 认证 ----
|
||||
|
||||
type davLog struct {
|
||||
Method string
|
||||
Path string
|
||||
Status int
|
||||
}
|
||||
|
||||
type fakeDav struct {
|
||||
mu sync.Mutex
|
||||
dirs map[string]bool
|
||||
files map[string][]byte
|
||||
|
||||
// 认证配置:mode = none|basic|digest;digest 配合 algo = MD5|SHA-256。
|
||||
mode string
|
||||
username string
|
||||
password string
|
||||
realm string
|
||||
nonce string
|
||||
opaque string
|
||||
algo string
|
||||
|
||||
failNext map[string]int // method → 剩余 503 次数
|
||||
logs []davLog
|
||||
}
|
||||
|
||||
func newFakeDav(mode string) *fakeDav {
|
||||
return &fakeDav{
|
||||
dirs: map[string]bool{},
|
||||
files: map[string][]byte{},
|
||||
mode: mode,
|
||||
username: "fcb",
|
||||
password: "fcb-pass",
|
||||
realm: "test-realm",
|
||||
nonce: "dcd98b7102dd2f0e8b11d0f600bfb0c0",
|
||||
opaque: "5ccc069c403ebaf9f0171e9517f40e41",
|
||||
algo: "MD5",
|
||||
failNext: map[string]int{},
|
||||
}
|
||||
}
|
||||
|
||||
// auth 校验请求凭据;失败时写出 401 与对应挑战。
|
||||
func (f *fakeDav) auth(w http.ResponseWriter, r *http.Request) bool {
|
||||
if f.mode == "none" {
|
||||
return true
|
||||
}
|
||||
h := r.Header.Get("Authorization")
|
||||
ok := false
|
||||
switch f.mode {
|
||||
case "basic":
|
||||
ok = h == "Basic "+basicAuth(f.username, f.password)
|
||||
case "digest":
|
||||
ok = f.checkDigest(r)
|
||||
}
|
||||
if ok {
|
||||
return true
|
||||
}
|
||||
switch f.mode {
|
||||
case "basic":
|
||||
w.Header().Set("WWW-Authenticate", `Basic realm="`+f.realm+`"`)
|
||||
case "digest":
|
||||
w.Header().Set("WWW-Authenticate", fmt.Sprintf(
|
||||
`Digest realm="%s", qop="auth", nonce="%s", opaque="%s", algorithm=%s, stale=false`,
|
||||
f.realm, f.nonce, f.opaque, f.algo))
|
||||
}
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
return false
|
||||
}
|
||||
|
||||
// checkDigest 服务端重算 RFC 2617 摘要响应。
|
||||
func (f *fakeDav) checkDigest(r *http.Request) bool {
|
||||
h := r.Header.Get("Authorization")
|
||||
if !strings.HasPrefix(h, "Digest ") {
|
||||
return false
|
||||
}
|
||||
p := parseChallengeParams(strings.TrimSpace(h[len("Digest "):]))
|
||||
ha1 := hashHex(algoName(f.algo), f.username+":"+f.realm+":"+f.password)
|
||||
ha2 := hashHex(algoName(f.algo), strings.ToUpper(r.Method)+":"+r.URL.RequestURI())
|
||||
got := hashHex(algoName(f.algo), ha1+":"+f.nonce+":"+p["nc"]+":"+p["cnonce"]+":"+p["qop"]+":"+ha2)
|
||||
return p["username"] == f.username && p["response"] == got
|
||||
}
|
||||
|
||||
func basicAuth(user, pass string) string {
|
||||
return base64.StdEncoding.EncodeToString([]byte(user + ":" + pass))
|
||||
}
|
||||
|
||||
func (f *fakeDav) record(method, path string, status int) {
|
||||
f.logs = append(f.logs, davLog{Method: method, Path: path, Status: status})
|
||||
}
|
||||
|
||||
// maybeFail 命中失败注入时返回 true(已写出 503)。
|
||||
func (f *fakeDav) maybeFail(w http.ResponseWriter, method string) bool {
|
||||
if f.failNext[method] > 0 {
|
||||
f.failNext[method]--
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (f *fakeDav) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
if !f.auth(w, r) {
|
||||
f.record(r.Method, r.URL.Path, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
p := strings.Trim(r.URL.Path, "/")
|
||||
|
||||
switch r.Method {
|
||||
case http.MethodPut:
|
||||
if f.maybeFail(w, "PUT") {
|
||||
f.record(r.Method, p, 503)
|
||||
return
|
||||
}
|
||||
if parent := parentOf(p); parent != "" && !f.dirs[parent] {
|
||||
w.WriteHeader(http.StatusConflict) // 强制客户端先建目录
|
||||
f.record(r.Method, p, 409)
|
||||
return
|
||||
}
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
f.files[p] = body
|
||||
f.record(r.Method, p, 201)
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
case http.MethodGet:
|
||||
if f.maybeFail(w, "GET") {
|
||||
f.record(r.Method, p, 503)
|
||||
return
|
||||
}
|
||||
data, ok := f.files[p]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
f.record(r.Method, p, 404)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/octet-stream")
|
||||
w.Header().Set("Accept-Ranges", "bytes")
|
||||
if rng := r.Header.Get("Range"); rng != "" {
|
||||
start, end := int64(0), int64(len(data))-1
|
||||
spec := strings.TrimPrefix(rng, "bytes=")
|
||||
if strings.HasSuffix(spec, "-") { // bytes=N- → 到文件尾
|
||||
if s, err := strconv.ParseInt(strings.TrimSuffix(spec, "-"), 10, 64); err == nil {
|
||||
start = s
|
||||
}
|
||||
} else if _, err := fmt.Sscanf(spec, "%d-%d", &start, &end); err != nil {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
f.record(r.Method, p, 400)
|
||||
return
|
||||
}
|
||||
if start < 0 || start >= int64(len(data)) {
|
||||
w.WriteHeader(http.StatusRequestedRangeNotSatisfiable)
|
||||
f.record(r.Method, p, 416)
|
||||
return
|
||||
}
|
||||
if end >= int64(len(data)) {
|
||||
end = int64(len(data)) - 1
|
||||
}
|
||||
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, end, len(data)))
|
||||
w.WriteHeader(http.StatusPartialContent)
|
||||
_, _ = w.Write(data[start : end+1])
|
||||
f.record(r.Method, p, 206)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(data)
|
||||
f.record(r.Method, p, 200)
|
||||
case http.MethodHead:
|
||||
data, ok := f.files[p]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
f.record(r.Method, p, 404)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Length", strconv.Itoa(len(data)))
|
||||
w.WriteHeader(http.StatusOK)
|
||||
f.record(r.Method, p, 200)
|
||||
case "PROPFIND":
|
||||
depth := r.Header.Get("Depth")
|
||||
self, isDirSelf := f.stat(p)
|
||||
if !isDirSelf {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
f.record(r.Method, p, 404)
|
||||
return
|
||||
}
|
||||
var b strings.Builder
|
||||
b.WriteString(`<?xml version="1.0" encoding="utf-8"?>` +
|
||||
`<D:multistatus xmlns:D="DAV:">`)
|
||||
f.writeResponse(&b, p, self)
|
||||
if depth == "1" && self.isDir {
|
||||
for _, name := range f.children(p) {
|
||||
child := name
|
||||
cs, cd := f.stat(child)
|
||||
f.writeResponse(&b, child, davStat{isDir: cd, size: cs.size})
|
||||
}
|
||||
}
|
||||
b.WriteString(`</D:multistatus>`)
|
||||
w.Header().Set("Content-Type", "application/xml; charset=utf-8")
|
||||
w.WriteHeader(http.StatusMultiStatus)
|
||||
_, _ = w.Write([]byte(b.String()))
|
||||
f.record(r.Method, p, 207)
|
||||
case "MKCOL":
|
||||
if f.dirs[p] || f.files[p] != nil {
|
||||
w.WriteHeader(http.StatusMethodNotAllowed) // 已存在
|
||||
f.record(r.Method, p, 405)
|
||||
return
|
||||
}
|
||||
if parent := parentOf(p); parent != "" && !f.dirs[parent] {
|
||||
w.WriteHeader(http.StatusConflict)
|
||||
f.record(r.Method, p, 409)
|
||||
return
|
||||
}
|
||||
f.dirs[p] = true
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
f.record(r.Method, p, 201)
|
||||
case http.MethodDelete:
|
||||
if _, ok := f.files[p]; ok {
|
||||
delete(f.files, p)
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
f.record(r.Method, p, 204)
|
||||
return
|
||||
}
|
||||
if f.dirs[p] {
|
||||
// 递归删除目录
|
||||
prefix := p + "/"
|
||||
for name := range f.files {
|
||||
if strings.HasPrefix(name, prefix) {
|
||||
delete(f.files, name)
|
||||
}
|
||||
}
|
||||
for name := range f.dirs {
|
||||
if name == p || strings.HasPrefix(name+"/", prefix) {
|
||||
delete(f.dirs, name)
|
||||
}
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
f.record(r.Method, p, 204)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
f.record(r.Method, p, 404)
|
||||
default:
|
||||
w.WriteHeader(http.StatusMethodNotAllowed)
|
||||
f.record(r.Method, p, 405)
|
||||
}
|
||||
}
|
||||
|
||||
type davStat struct {
|
||||
isDir bool
|
||||
size int
|
||||
}
|
||||
|
||||
func (f *fakeDav) stat(p string) (davStat, bool) {
|
||||
if data, ok := f.files[p]; ok {
|
||||
return davStat{size: len(data)}, true
|
||||
}
|
||||
if f.dirs[p] {
|
||||
return davStat{isDir: true}, true
|
||||
}
|
||||
return davStat{}, false
|
||||
}
|
||||
|
||||
func (f *fakeDav) children(p string) []string {
|
||||
var out []string
|
||||
prefix := p + "/"
|
||||
for name := range f.files {
|
||||
if strings.HasPrefix(name, prefix) && !strings.Contains(strings.TrimPrefix(name, prefix), "/") {
|
||||
out = append(out, name)
|
||||
}
|
||||
}
|
||||
for name := range f.dirs {
|
||||
if strings.HasPrefix(name, prefix) && !strings.Contains(strings.TrimPrefix(name, prefix), "/") {
|
||||
out = append(out, name)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func (f *fakeDav) writeResponse(b *strings.Builder, href string, st davStat) {
|
||||
b.WriteString(`<D:response><D:href>/` + href + `</D:href><D:propstat><D:prop><D:resourcetype>`)
|
||||
if st.isDir {
|
||||
b.WriteString(`<D:collection/>`)
|
||||
}
|
||||
b.WriteString(`</D:resourcetype><D:getcontentlength>` + strconv.Itoa(st.size) +
|
||||
`</D:getcontentlength><D:getcontenttype>application/octet-stream</D:getcontenttype>` +
|
||||
`</D:prop><D:status>HTTP/1.1 200 OK</D:status></D:propstat></D:response>`)
|
||||
}
|
||||
|
||||
func parentOf(p string) string {
|
||||
if i := strings.LastIndex(p, "/"); i > 0 {
|
||||
return p[:i]
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// newTestDav 构造 WebDAV 引擎 + 假服务。
|
||||
func newTestDav(t *testing.T, mode string, tweak func(o *WebDAVOptions)) (*WebDAVStorage, *fakeDav, *int32) {
|
||||
t.Helper()
|
||||
f := newFakeDav(mode)
|
||||
var conns int32
|
||||
srv := httptest.NewUnstartedServer(f)
|
||||
srv.Config.ConnState = func(c net.Conn, cs http.ConnState) {
|
||||
if cs == http.StateNew {
|
||||
atomic.AddInt32(&conns, 1)
|
||||
}
|
||||
}
|
||||
srv.Start()
|
||||
t.Cleanup(srv.Close)
|
||||
opts := WebDAVOptions{
|
||||
BaseURL: srv.URL,
|
||||
Username: f.username,
|
||||
Password: f.password,
|
||||
RootPath: "fcb_root",
|
||||
MaxRetries: 3,
|
||||
}
|
||||
if tweak != nil {
|
||||
tweak(&opts)
|
||||
}
|
||||
st, err := NewWebDAVStorage(opts)
|
||||
if err != nil {
|
||||
t.Fatalf("NewWebDAVStorage: %v", err)
|
||||
}
|
||||
return st, f, &conns
|
||||
}
|
||||
|
||||
// TestWebDAVBasicCRUD Basic 认证下的完整 CRUD 与 Range。
|
||||
func TestWebDAVBasicCRUD(t *testing.T) {
|
||||
st, f, _ := newTestDav(t, "basic", nil)
|
||||
ctx := context.Background()
|
||||
|
||||
// 健康检查:根目录 404 → MKCOL 自建
|
||||
if err := st.HealthCheck(ctx); err != nil {
|
||||
t.Fatalf("HealthCheck: %v", err)
|
||||
}
|
||||
if !f.dirs["fcb_root"] {
|
||||
t.Fatalf("根目录应被自动创建")
|
||||
}
|
||||
|
||||
data := []byte("WebDAV 引擎数据 0123456789 ABCDEF")
|
||||
n, err := st.SaveFile(ctx, bytes.NewReader(data), "2025/08/w.bin")
|
||||
if err != nil {
|
||||
t.Fatalf("SaveFile: %v", err)
|
||||
}
|
||||
if n != int64(len(data)) {
|
||||
t.Fatalf("n = %d", n)
|
||||
}
|
||||
if string(f.files["fcb_root/2025/08/w.bin"]) != string(data) {
|
||||
t.Fatalf("PUT 内容不匹配")
|
||||
}
|
||||
// 按需建目录:两级目录都应已创建
|
||||
if !f.dirs["fcb_root/2025"] || !f.dirs["fcb_root/2025/08"] {
|
||||
t.Fatalf("目录未按需创建: %v %v", f.dirs["fcb_root/2025"], f.dirs["fcb_root/2025/08"])
|
||||
}
|
||||
|
||||
meta, err := st.Stat(ctx, "2025/08/w.bin")
|
||||
if err != nil {
|
||||
t.Fatalf("Stat: %v", err)
|
||||
}
|
||||
if meta.Size != int64(len(data)) || !meta.AcceptRanges {
|
||||
t.Fatalf("Stat = %+v", meta)
|
||||
}
|
||||
if ok, _ := st.FileExists(ctx, "2025/08/w.bin"); !ok {
|
||||
t.Fatalf("FileExists 应为 true")
|
||||
}
|
||||
|
||||
// 完整下载(对齐 go-api 约定)
|
||||
dl, err := st.Open(ctx, "2025/08/w.bin", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Open: %v", err)
|
||||
}
|
||||
got, err := io.ReadAll(dl)
|
||||
_ = dl.Close()
|
||||
if err != nil {
|
||||
t.Fatalf("read: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, data) {
|
||||
t.Fatalf("full 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)
|
||||
}
|
||||
|
||||
// Range 下载
|
||||
dl, err = st.Open(ctx, "2025/08/w.bin", &Range{Start: 2, End: 7})
|
||||
if err != nil {
|
||||
t.Fatalf("Open range: %v", err)
|
||||
}
|
||||
got, err = io.ReadAll(dl)
|
||||
_ = dl.Close()
|
||||
if err != nil {
|
||||
t.Fatalf("read range: %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, data[2:8]) {
|
||||
t.Fatalf("range mismatch")
|
||||
}
|
||||
if dl.Start != 2 || dl.End != 7 || dl.Total != int64(len(data)) {
|
||||
t.Fatalf("range offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
|
||||
}
|
||||
|
||||
// 416(起点越界)/ 404
|
||||
if _, err := st.Open(ctx, "2025/08/w.bin", &Range{Start: int64(len(data)) + 9, End: -1}); err == nil ||
|
||||
!strings.Contains(err.Error(), ErrRangeNotSatisfiable.Error()) {
|
||||
t.Fatalf("want ErrRangeNotSatisfiable, got %v", err)
|
||||
}
|
||||
if _, err := st.Open(ctx, "no/such.bin", nil); err == nil ||
|
||||
!strings.Contains(err.Error(), ErrNotFound.Error()) {
|
||||
t.Fatalf("want ErrNotFound, got %v", err)
|
||||
}
|
||||
|
||||
// 删除 + 空父目录清理
|
||||
if err := st.DeleteFile(ctx, "2025/08/w.bin"); err != nil {
|
||||
t.Fatalf("DeleteFile: %v", err)
|
||||
}
|
||||
if ok, _ := st.FileExists(ctx, "2025/08/w.bin"); ok {
|
||||
t.Fatalf("删除后仍存在")
|
||||
}
|
||||
if _, err := st.Stat(ctx, "2025/08/w.bin"); err == nil ||
|
||||
!strings.Contains(err.Error(), ErrNotFound.Error()) {
|
||||
t.Fatalf("Stat 应 ErrNotFound, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVDigestAuth Digest(MD5)认证协商。
|
||||
func TestWebDAVDigestAuth(t *testing.T) {
|
||||
st, f, _ := newTestDav(t, "digest", nil)
|
||||
ctx := context.Background()
|
||||
if err := st.HealthCheck(ctx); err != nil {
|
||||
t.Fatalf("HealthCheck(digest): %v", err)
|
||||
}
|
||||
// HealthCheck 流程应观察到 401 挑战(客户端先 Basic 探测 → 401 → Digest 重试)
|
||||
saw401 := false
|
||||
for _, l := range f.logs {
|
||||
if l.Status == 401 {
|
||||
saw401 = true
|
||||
}
|
||||
}
|
||||
if !saw401 {
|
||||
t.Fatalf("未观察到 401 挑战: %+v", f.logs)
|
||||
}
|
||||
|
||||
// 认证后的 PROPFIND(Stat 已有目录)应得到 207
|
||||
if _, err := st.Stat(ctx, ""); err == nil {
|
||||
// Stat("") 非法路径属预期;这里换用 FileExists 对已有根目录探测
|
||||
_ = err
|
||||
}
|
||||
data := []byte("digest 内容")
|
||||
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "d.bin"); err != nil {
|
||||
t.Fatalf("SaveFile(digest): %v", err)
|
||||
}
|
||||
dl, err := st.Open(ctx, "d.bin", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Open(digest): %v", err)
|
||||
}
|
||||
got, _ := io.ReadAll(dl)
|
||||
_ = dl.Close()
|
||||
if !bytes.Equal(got, data) {
|
||||
t.Fatalf("digest 下载内容不匹配")
|
||||
}
|
||||
// 全链路完成:确认存在成功的 2xx/207 请求
|
||||
saw2xx := false
|
||||
for _, l := range f.logs {
|
||||
if l.Status == 207 || l.Status == 201 || l.Status == 200 {
|
||||
saw2xx = true
|
||||
}
|
||||
}
|
||||
if !saw2xx {
|
||||
t.Fatalf("认证后应有成功请求: %+v", f.logs)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVDigestSHA256 Digest(SHA-256)算法。
|
||||
func TestWebDAVDigestSHA256(t *testing.T) {
|
||||
st, _, _ := newTestDav(t, "digest", nil)
|
||||
st.auth.mu.Lock()
|
||||
st.auth.algorithm = "sha-256"
|
||||
st.auth.mu.Unlock()
|
||||
// 服务端也切换到 SHA-256 重算摘要
|
||||
st2, f, _ := newTestDav(t, "digest", nil)
|
||||
f.algo = "SHA-256"
|
||||
// 先让客户端完成一次 MD5 协商拿到挑战参数,再切 SHA-256 会 401 失败——
|
||||
// 因此这里直接对 SHA-256 服务端做完整链路(client 首次探测 Basic→401→Digest)。
|
||||
_ = st
|
||||
ctx := context.Background()
|
||||
if err := st2.HealthCheck(ctx); err != nil {
|
||||
t.Fatalf("HealthCheck(SHA-256): %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVDigestWrongPassword 凭据错误 → 明确报错而非重试风暴。
|
||||
func TestWebDAVDigestWrongPassword(t *testing.T) {
|
||||
f := newFakeDav("digest")
|
||||
srv := httptest.NewServer(f)
|
||||
defer srv.Close()
|
||||
st, err := NewWebDAVStorage(WebDAVOptions{
|
||||
BaseURL: srv.URL, Username: f.username, Password: "WRONG",
|
||||
RootPath: "r", MaxRetries: 1, BaseBackoff: 1,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := st.HealthCheck(context.Background()); err == nil ||
|
||||
!strings.Contains(err.Error(), "401") {
|
||||
t.Fatalf("错误凭据应报 401 相关错误, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVRetryGet 5xx 指数退避重试(GET 幂等)。
|
||||
func TestWebDAVRetryGet(t *testing.T) {
|
||||
st, f, _ := newTestDav(t, "basic", func(o *WebDAVOptions) { o.BaseBackoff = 5 })
|
||||
ctx := context.Background()
|
||||
data := []byte("retry target")
|
||||
if err := st.HealthCheck(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "r.bin"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
f.mu.Lock()
|
||||
f.failNext["GET"] = 2
|
||||
f.mu.Unlock()
|
||||
dl, err := st.Open(ctx, "r.bin", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("503×2 后应重试成功: %v", err)
|
||||
}
|
||||
got, _ := io.ReadAll(dl)
|
||||
_ = dl.Close()
|
||||
if !bytes.Equal(got, data) {
|
||||
t.Fatalf("content mismatch")
|
||||
}
|
||||
// 验证确实发了 3 次 GET
|
||||
gets := 0
|
||||
for _, l := range f.logs {
|
||||
if l.Method == "GET" && strings.HasSuffix(l.Path, "r.bin") {
|
||||
gets++
|
||||
}
|
||||
}
|
||||
if gets != 3 {
|
||||
t.Fatalf("GET 次数 = %d, want 3", gets)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVRetryPut 可重放 body(seekable)PUT 失败重试;不可重放不重试。
|
||||
func TestWebDAVRetryPut(t *testing.T) {
|
||||
st, f, _ := newTestDav(t, "basic", func(o *WebDAVOptions) { o.BaseBackoff = 5 })
|
||||
ctx := context.Background()
|
||||
if err := st.HealthCheck(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// seekable:重试成功
|
||||
f.mu.Lock()
|
||||
f.failNext["PUT"] = 1
|
||||
f.mu.Unlock()
|
||||
data := []byte("put with retry")
|
||||
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "pr.bin"); err != nil {
|
||||
t.Fatalf("PUT 重试应成功: %v", err)
|
||||
}
|
||||
puts := 0
|
||||
for _, l := range f.logs {
|
||||
if l.Method == "PUT" && strings.HasSuffix(l.Path, "pr.bin") {
|
||||
puts++
|
||||
}
|
||||
}
|
||||
if puts != 2 {
|
||||
t.Fatalf("PUT 次数 = %d, want 2", puts)
|
||||
}
|
||||
// 非 seekable(io.Pipe):不重试,直接失败
|
||||
f.mu.Lock()
|
||||
f.failNext["PUT"] = 1
|
||||
f.mu.Unlock()
|
||||
pr, pw := io.Pipe()
|
||||
go func() {
|
||||
_, _ = pw.Write([]byte("non-seekable"))
|
||||
_ = pw.Close()
|
||||
}()
|
||||
if _, err := st.SaveFile(ctx, pr, "ns.bin"); err == nil {
|
||||
t.Fatalf("非重放 PUT 注入 503 应失败")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVConnectionReuse 连接复用:多次请求不应各建一条 TCP 连接。
|
||||
func TestWebDAVConnectionReuse(t *testing.T) {
|
||||
st, _, conns := newTestDav(t, "basic", nil)
|
||||
ctx := context.Background()
|
||||
if err := st.HealthCheck(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := 0; i < 12; i++ {
|
||||
if _, err := st.SaveFile(ctx, bytes.NewReader([]byte("x")), fmt.Sprintf("reuse/%d.bin", i)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := st.Stat(ctx, fmt.Sprintf("reuse/%d.bin", i)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
// 25 次请求(12 PUT + 12 PROPFIND + 1 HealthCheck 的 PROPFIND/MKCOL)只允许极少量新连接
|
||||
if got := atomic.LoadInt32(conns); got > 4 {
|
||||
t.Fatalf("新建 TCP 连接数 = %d,连接复用失效(应 ≤4)", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVPipeStreaming io.Pipe 流式转发:完整读取 + 提前关闭。
|
||||
func TestWebDAVPipeStreaming(t *testing.T) {
|
||||
st, _, _ := newTestDav(t, "basic", nil)
|
||||
ctx := context.Background()
|
||||
if err := st.HealthCheck(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
big := bytes.Repeat([]byte("0123456789abcdef"), 4096) // 64KB
|
||||
if _, err := st.SaveFile(ctx, bytes.NewReader(big), "big.bin"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
dl, err := st.Open(ctx, "big.bin", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, err := io.ReadAll(dl)
|
||||
if err != nil {
|
||||
t.Fatalf("read pipe: %v", err)
|
||||
}
|
||||
_ = dl.Close()
|
||||
if !bytes.Equal(got, big) {
|
||||
t.Fatalf("pipe content mismatch")
|
||||
}
|
||||
|
||||
// 提前关闭:后续读取返回错误且不挂死
|
||||
dl2, err := st.Open(ctx, "big.bin", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
buf := make([]byte, 10)
|
||||
if _, err := io.ReadFull(dl2, buf); err != nil {
|
||||
t.Fatalf("read head: %v", err)
|
||||
}
|
||||
if err := dl2.Close(); err != nil {
|
||||
t.Fatalf("early close: %v", err)
|
||||
}
|
||||
// ctx 取消同样会终止流
|
||||
cctx, cancel := context.WithCancel(context.Background())
|
||||
dl3, err := st.Open(cctx, "big.bin", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cancel()
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
_, err = dl3.Read(buf)
|
||||
if err == nil {
|
||||
_ = dl3.Close()
|
||||
t.Fatalf("ctx 取消后读取应报错")
|
||||
}
|
||||
_ = dl3.Close()
|
||||
}
|
||||
|
||||
// TestWebDAVChunkMerge 分片保存/合并/清理。
|
||||
func TestWebDAVChunkMerge(t *testing.T) {
|
||||
st, f, _ := newTestDav(t, "basic", nil)
|
||||
ctx := context.Background()
|
||||
if err := st.HealthCheck(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
savePath := "2025/09/merged.bin"
|
||||
uploadID := "uid-webdav"
|
||||
chunks := [][]byte{[]byte("AAA"), []byte("BB"), []byte("CCCC")}
|
||||
hashes := make([]string, len(chunks))
|
||||
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)
|
||||
}
|
||||
if n != int64(len(c)) {
|
||||
t.Fatalf("chunk %d size = %d", i, n)
|
||||
}
|
||||
hashes[i] = sha256Hex(c)
|
||||
}
|
||||
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 != 9 || fileHash != sha256Hex(bytes.Join(chunks, nil)) {
|
||||
t.Fatalf("merge result = %d %s", size, fileHash)
|
||||
}
|
||||
if string(f.files["fcb_root/"+savePath]) != "AAABBCCCC" {
|
||||
t.Fatalf("合并内容错误: %q", f.files["fcb_root/"+savePath])
|
||||
}
|
||||
// 分片目录已清理
|
||||
for k := range f.files {
|
||||
if strings.Contains(k, "chunks/"+uploadID) {
|
||||
t.Fatalf("分片残留: %s", k)
|
||||
}
|
||||
}
|
||||
if f.dirs["fcb_root/2025/09/chunks/"+uploadID] {
|
||||
t.Fatalf("分片目录残留")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVCleanChunks 清理与哈希失败路径。
|
||||
func TestWebDAVCleanChunks(t *testing.T) {
|
||||
st, f, _ := newTestDav(t, "basic", nil)
|
||||
ctx := context.Background()
|
||||
if err := st.HealthCheck(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for i := 0; i < 2; i++ {
|
||||
if _, err := st.SaveChunk(ctx, "uidc", i, strings.NewReader("z"), "c.bin"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if err := st.CleanChunks(ctx, "uidc", "c.bin"); err != nil {
|
||||
t.Fatalf("CleanChunks: %v", err)
|
||||
}
|
||||
if len(f.files) != 0 {
|
||||
t.Fatalf("分片未清理: %v", f.files)
|
||||
}
|
||||
// 幂等
|
||||
if err := st.CleanChunks(ctx, "uidc", "c.bin"); err != nil {
|
||||
t.Fatalf("CleanChunks idempotent: %v", err)
|
||||
}
|
||||
// 哈希不匹配
|
||||
if _, err := st.SaveChunk(ctx, "uidm", 0, strings.NewReader("real"), "m.bin"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, _, err := st.MergeChunks(ctx, "uidm", 1, func(i int) (string, error) {
|
||||
return sha256Hex([]byte("wrong")), nil
|
||||
}, "m.bin"); err == nil || !strings.Contains(err.Error(), ErrHashMismatch.Error()) {
|
||||
t.Fatalf("want ErrHashMismatch, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVTimeout 非流式操作超时:PROPFIND 响应慢于 Timeout(1s)→ context deadline exceeded。
|
||||
func TestWebDAVTimeout(t *testing.T) {
|
||||
f := newFakeDav("basic")
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == "PROPFIND" {
|
||||
time.Sleep(1500 * time.Millisecond) // > Timeout 1s
|
||||
}
|
||||
f.ServeHTTP(w, r)
|
||||
}))
|
||||
defer srv.Close()
|
||||
st, err := NewWebDAVStorage(WebDAVOptions{
|
||||
BaseURL: srv.URL, Username: f.username, Password: f.password,
|
||||
RootPath: "r", MaxRetries: 0, Timeout: 1, BaseBackoff: 5,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
start := time.Now()
|
||||
err = st.HealthCheck(context.Background())
|
||||
if err == nil {
|
||||
t.Fatalf("超时应报错")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "context deadline exceeded") {
|
||||
t.Fatalf("应为超时错误, got %v", err)
|
||||
}
|
||||
// 单次尝试 1s 超时 + 一次重试 ≈ 2s;若超时未生效会拖满 2×1.5s
|
||||
if elapsed := time.Since(start); elapsed > 3500*time.Millisecond {
|
||||
t.Fatalf("超时未生效(耗时 %v)", elapsed)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVPresignNotSupported 预签名 → ErrNotSupported。
|
||||
func TestWebDAVPresignNotSupported(t *testing.T) {
|
||||
st, _, _ := newTestDav(t, "basic", nil)
|
||||
ctx := context.Background()
|
||||
if _, err := st.PresignGetURL(ctx, "x.bin", 60); err == nil ||
|
||||
!strings.Contains(err.Error(), ErrNotSupported.Error()) {
|
||||
t.Fatalf("want ErrNotSupported, got %v", err)
|
||||
}
|
||||
if _, err := st.PresignPutURL(ctx, "x.bin", 60); err == nil ||
|
||||
!strings.Contains(err.Error(), ErrNotSupported.Error()) {
|
||||
t.Fatalf("want ErrNotSupported, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebDAVFactoryRegistry 工厂构造 + Digest 全链路。
|
||||
func TestWebDAVFactoryRegistry(t *testing.T) {
|
||||
f := newFakeDav("digest")
|
||||
srv := httptest.NewServer(f)
|
||||
defer srv.Close()
|
||||
prev := engineOptions.WebDAV
|
||||
engineOptions.WebDAV = WebDAVOptions{
|
||||
BaseURL: srv.URL, Username: f.username, Password: f.password,
|
||||
RootPath: "factory_root", MaxRetries: 3, BaseBackoff: 5,
|
||||
}
|
||||
defer func() { engineOptions.WebDAV = prev }()
|
||||
st, err := NewEngine(context.Background(), "webdav")
|
||||
if err != nil {
|
||||
t.Fatalf("NewEngine(webdav): %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
if err := st.HealthCheck(ctx); err != nil {
|
||||
t.Fatalf("HealthCheck: %v", err)
|
||||
}
|
||||
if _, err := st.SaveFile(ctx, strings.NewReader("factory"), "f.txt"); err != nil {
|
||||
t.Fatalf("SaveFile: %v", err)
|
||||
}
|
||||
dl, err := st.Open(ctx, "f.txt", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("Open: %v", err)
|
||||
}
|
||||
got, _ := io.ReadAll(dl)
|
||||
_ = dl.Close()
|
||||
if string(got) != "factory" {
|
||||
t.Fatalf("content = %q", got)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user