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:
2026-09-05 04:22:41 +08:00
commit 7f060dd0e4
173 changed files with 32455 additions and 0 deletions
+10
View File
@@ -0,0 +1,10 @@
package storage
import "errors"
// ErrRangeNotSatisfiable 请求的字节范围超出文件大小(对应 HTTP 416)。
// 属于增量错误定义,不改动 interface.go 的既有签名。
var ErrRangeNotSatisfiable = errors.New("storage: 请求范围超出文件大小")
// ErrHashMismatch 分片哈希校验失败(合并分片时检测到数据损坏)。
var ErrHashMismatch = errors.New("storage: 分片哈希校验失败")
+26
View File
@@ -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)
}
+111
View File
@@ -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=0End=Total-1);
// - rng 非 nil:返回 [Start, End] 区间流。
// 引擎应尽量透传 RangeWebDAV/S3)或按块 seeklocal)。
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
}
+449
View File
@@ -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)
+346
View File
@@ -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("未知引擎应报错")
}
}
+187
View File
@@ -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)
}
+197
View File
@@ -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 构造测试用 Managerlocal 健康引擎起步;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)
}
}
+70
View File
@@ -0,0 +1,70 @@
package storage
// EngineOptions 引擎构造选项:由 main.goAPI 层任务)从 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_urlMinIO 等;空则 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
}
}
+103
View File
@@ -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:])
}
+74
View File
@@ -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
}
+650
View File
@@ -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 等)。
//
// 相比参考实现(S3FileStorageaioboto3)的改进:
// - 单例客户端 + 自定义连接池 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 需要可重放流
// 或 TLSMinIO/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 // 5MBS3 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 会话 IDAbort 时复用 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 健康检查:列举 bucketMaxKeys=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)
+495
View File
@@ -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
}
+889
View File
@@ -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 引擎(本次重写的重点优化对象)。
//
// 相比参考实现(WebDAVFileStorageaiohttp)的改进:
// - 单例 http.Client + 连接池化 Transport(参考实现每个操作新建 ClientSession,无复用);
// - Basic 与 DigestRFC 2617qop=authMD5/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 连接池化 TransportKeep-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 将远端路径转为完整 URLURL.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)
+244
View File
@@ -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=authMD5 / 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)
}
+824
View File
@@ -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|digestdigest 配合 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 DigestMD5)认证协商。
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)
}
// 认证后的 PROPFINDStat 已有目录)应得到 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 DigestSHA-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 可重放 bodyseekable)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)
}
// 非 seekableio.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 响应慢于 Timeout1s)→ 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)
}
}