Files
FileShare/server/internal/storage/s3_test.go
T
SKYMirror 9686fe887a FileCodeBox Go 重写版 v2.5.6(安全审计修复版)
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 全绿;二进制端到端冒烟通过
2026-09-05 04:22:41 +08:00

496 lines
14 KiB
Go

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
}