Merge pull request #142 from Dvorinka/fix/phase-c-bugs
Phase C bug fixes: SMTP auth selection, trash-walk OOM, WebDAV ranged PUTpull/3582/head
commit
a82c9f5277
@ -0,0 +1,51 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/cloudreve/Cloudreve/v4/application/dependency"
|
||||
"github.com/cloudreve/Cloudreve/v4/pkg/serializer"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const rateLimitPrefix = "rate_limit:"
|
||||
|
||||
// RateLimit applies a fixed-window rate limit to requests sharing the same
|
||||
// bucket key. Once `limit` requests arrive within `window`, further requests
|
||||
// are rejected until the window expires. State lives in the shared KV store,
|
||||
// so limits hold across cluster nodes when Redis is configured.
|
||||
//
|
||||
// The counter is approximate (read-then-write), which is acceptable for
|
||||
// abuse prevention — it bounds attempts to roughly `limit` per window.
|
||||
func RateLimit(limit int, window time.Duration, key func(*gin.Context) string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
kv := dependency.FromContext(c).KV()
|
||||
bucket := rateLimitPrefix + key(c)
|
||||
|
||||
count := 0
|
||||
if raw, ok := kv.Get(bucket); ok {
|
||||
if v, ok := raw.(int); ok {
|
||||
count = v
|
||||
}
|
||||
}
|
||||
|
||||
if count >= limit {
|
||||
c.Header("Retry-After", strconv.Itoa(int(window.Seconds())))
|
||||
c.JSON(200, serializer.NewError(serializer.CodeRateLimited, "Too many requests, please try again later.", nil))
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
_ = kv.Set(bucket, count+1, int(window.Seconds()))
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// RateLimitByIP keys the rate limit on the client IP and a static bucket
|
||||
// name, e.g. RateLimitByIP("login", 10, time.Minute).
|
||||
func RateLimitByIP(bucket string, limit int, window time.Duration) gin.HandlerFunc {
|
||||
return RateLimit(limit, window, func(c *gin.Context) string {
|
||||
return bucket + ":" + c.ClientIP()
|
||||
})
|
||||
}
|
||||
@ -0,0 +1,123 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/cloudreve/Cloudreve/v4/application/dependency"
|
||||
"github.com/cloudreve/Cloudreve/v4/pkg/cache"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
var testEngine = func() *gin.Engine {
|
||||
e := gin.New()
|
||||
e.ContextWithFallback = true
|
||||
return e
|
||||
}()
|
||||
|
||||
func newRateLimitContext(t *testing.T, dep dependency.Dep, remoteAddr string) (*gin.Context, *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
w := httptest.NewRecorder()
|
||||
c := gin.CreateTestContextOnly(w, testEngine)
|
||||
req := httptest.NewRequest(http.MethodPost, "/session/token", nil)
|
||||
req.RemoteAddr = remoteAddr + ":12345"
|
||||
c.Request = req.WithContext(context.WithValue(req.Context(), dependency.DepCtx{}, dep))
|
||||
return c, w
|
||||
}
|
||||
|
||||
func TestRateLimitByIP(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
dep := dependency.NewDependency(dependency.WithKV(cache.NewMemoStore("", nil)))
|
||||
|
||||
handler := RateLimitByIP("login", 3, time.Minute)
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
c, _ := newRateLimitContext(t, dep, "192.0.2.1")
|
||||
handler(c)
|
||||
if c.IsAborted() {
|
||||
t.Fatalf("request %d unexpectedly rejected", i+1)
|
||||
}
|
||||
}
|
||||
|
||||
c, w := newRateLimitContext(t, dep, "192.0.2.1")
|
||||
handler(c)
|
||||
if !c.IsAborted() {
|
||||
t.Fatal("request over limit was not rejected")
|
||||
}
|
||||
if w.Header().Get("Retry-After") != "60" {
|
||||
t.Fatalf("missing/incorrect Retry-After: %q", w.Header().Get("Retry-After"))
|
||||
}
|
||||
|
||||
// A different IP has its own bucket.
|
||||
c, _ = newRateLimitContext(t, dep, "192.0.2.2")
|
||||
handler(c)
|
||||
if c.IsAborted() {
|
||||
t.Fatal("unrelated IP shared the bucket")
|
||||
}
|
||||
|
||||
// A different bucket name on the same IP has its own counter.
|
||||
handler2 := RateLimitByIP("register", 1, time.Minute)
|
||||
c, _ = newRateLimitContext(t, dep, "192.0.2.1")
|
||||
handler2(c)
|
||||
if c.IsAborted() {
|
||||
t.Fatal("unrelated bucket shared the counter")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRateLimitExpiry(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
dep := dependency.NewDependency(dependency.WithKV(cache.NewMemoStore("", nil)))
|
||||
|
||||
handler := RateLimitByIP("login", 1, time.Second)
|
||||
|
||||
c, _ := newRateLimitContext(t, dep, "192.0.2.1")
|
||||
handler(c)
|
||||
if c.IsAborted() {
|
||||
t.Fatal("first request rejected")
|
||||
}
|
||||
|
||||
c, _ = newRateLimitContext(t, dep, "192.0.2.1")
|
||||
handler(c)
|
||||
if !c.IsAborted() {
|
||||
t.Fatal("second request not rejected")
|
||||
}
|
||||
|
||||
// Window expiry frees the bucket. MemoStore TTL is Unix-second granular
|
||||
// (item valid while Expires >= now), so a 1s window can live ~2s.
|
||||
time.Sleep(2100 * time.Millisecond)
|
||||
|
||||
c, _ = newRateLimitContext(t, dep, "192.0.2.1")
|
||||
handler(c)
|
||||
if c.IsAborted() {
|
||||
t.Fatal("request after window expiry still rejected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRateLimitCountOverflow(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
kv := cache.NewMemoStore("", nil)
|
||||
dep := dependency.NewDependency(dependency.WithKV(kv))
|
||||
|
||||
// A non-int value in the bucket is treated as a fresh window.
|
||||
if err := kv.Set(rateLimitPrefix+"login:192.0.2.1", "garbage", 60); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
handler := RateLimitByIP("login", 1, time.Minute)
|
||||
c, _ := newRateLimitContext(t, dep, "192.0.2.1")
|
||||
handler(c)
|
||||
if c.IsAborted() {
|
||||
t.Fatal("corrupt bucket state rejected request")
|
||||
}
|
||||
|
||||
raw, ok := kv.Get(rateLimitPrefix + "login:192.0.2.1")
|
||||
if !ok {
|
||||
t.Fatal("bucket not persisted")
|
||||
}
|
||||
if _, isInt := raw.(int); !isInt {
|
||||
t.Fatalf("bucket value type drifted: %T", raw)
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,27 @@
|
||||
package email
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
mail "github.com/wneessen/go-mail"
|
||||
)
|
||||
|
||||
func TestSMTPAuthType(t *testing.T) {
|
||||
assert.Equal(t, mail.SMTPAuthAutoDiscover, SMTPAuthType(""))
|
||||
assert.Equal(t, mail.SMTPAuthAutoDiscover, SMTPAuthType("autodiscover"))
|
||||
assert.Equal(t, mail.SMTPAuthAutoDiscover, SMTPAuthType("bogus"))
|
||||
|
||||
assert.Equal(t, mail.SMTPAuthPlain, SMTPAuthType("plain"))
|
||||
assert.Equal(t, mail.SMTPAuthPlainNoEnc, SMTPAuthType("plain-noenc"))
|
||||
assert.Equal(t, mail.SMTPAuthLogin, SMTPAuthType("login"))
|
||||
assert.Equal(t, mail.SMTPAuthLoginNoEnc, SMTPAuthType("login-noenc"))
|
||||
assert.Equal(t, mail.SMTPAuthCramMD5, SMTPAuthType("cram-md5"))
|
||||
assert.Equal(t, mail.SMTPAuthSCRAMSHA1, SMTPAuthType("scram-sha-1"))
|
||||
assert.Equal(t, mail.SMTPAuthSCRAMSHA256, SMTPAuthType("scram-sha-256"))
|
||||
assert.Equal(t, mail.SMTPAuthXOAUTH2, SMTPAuthType("xoauth2"))
|
||||
assert.Equal(t, mail.SMTPAuthNoAuth, SMTPAuthType("noauth"))
|
||||
|
||||
// Values are normalized before matching.
|
||||
assert.Equal(t, mail.SMTPAuthPlainNoEnc, SMTPAuthType(" Plain-NoEnc "))
|
||||
}
|
||||
@ -0,0 +1,195 @@
|
||||
package dbfs
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/cloudreve/Cloudreve/v4/ent"
|
||||
"github.com/cloudreve/Cloudreve/v4/ent/enttest"
|
||||
entfile "github.com/cloudreve/Cloudreve/v4/ent/file"
|
||||
entuser "github.com/cloudreve/Cloudreve/v4/ent/user"
|
||||
"github.com/cloudreve/Cloudreve/v4/inventory"
|
||||
"github.com/cloudreve/Cloudreve/v4/inventory/types"
|
||||
"github.com/cloudreve/Cloudreve/v4/pkg/boolset"
|
||||
"github.com/cloudreve/Cloudreve/v4/pkg/conf"
|
||||
"github.com/cloudreve/Cloudreve/v4/pkg/hashid"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// buildWalkTree creates root -> folder children. Each folder gets
|
||||
// filesPerFolder file children. Returns the root model wrapped as *File.
|
||||
func buildWalkTree(t *testing.T, client *ent.Client, folders, filesPerFolder int) (*ent.User, *File) {
|
||||
ctx := context.Background()
|
||||
group := client.Group.Create().SetName("walkers").SetPermissions(&boolset.BooleanSet{}).SaveX(ctx)
|
||||
user := client.User.Create().SetEmail("walk@example.com").SetNick("walk").SetGroup(group).SaveX(ctx)
|
||||
root := client.File.Create().SetName(inventory.RootFolderName).SetType(int(types.FileTypeFolder)).SetOwner(user).SaveX(ctx)
|
||||
|
||||
folderModels := make([]*ent.File, 0, folders)
|
||||
for i := 0; i < folders; i++ {
|
||||
folderModels = append(folderModels,
|
||||
client.File.Create().SetName(fmt.Sprintf("dir-%04d", i)).SetType(int(types.FileTypeFolder)).SetOwner(user).SetParent(root).SaveX(ctx))
|
||||
}
|
||||
for i, folder := range folderModels {
|
||||
creates := make([]*ent.FileCreate, 0, filesPerFolder)
|
||||
for j := 0; j < filesPerFolder; j++ {
|
||||
creates = append(creates,
|
||||
client.File.Create().SetName(fmt.Sprintf("f-%04d-%04d", i, j)).SetType(int(types.FileTypeFile)).SetOwner(user).SetParent(folder))
|
||||
}
|
||||
client.File.CreateBulk(creates...).SaveX(ctx)
|
||||
}
|
||||
|
||||
rootFile := newFile(nil, root)
|
||||
rootFile.OwnerModel = user
|
||||
return user, rootFile
|
||||
}
|
||||
|
||||
func walkTestNavigator(t *testing.T, client *ent.Client, user *ent.User) *baseNavigator {
|
||||
t.Helper()
|
||||
hasher, err := hashid.New("walk-test-salt")
|
||||
require.NoError(t, err)
|
||||
return newBaseNavigator(inventory.NewFileClient(client, conf.SQLiteDB, hasher), defaultFilter, user, nil, nil)
|
||||
}
|
||||
|
||||
func TestWalkEmitsWholeTree(t *testing.T) {
|
||||
client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared")
|
||||
t.Cleanup(func() { require.NoError(t, client.Close()) })
|
||||
user, rootFile := buildWalkTree(t, client, 3, 5)
|
||||
|
||||
nav := walkTestNavigator(t, client, user)
|
||||
|
||||
emitted := make(map[int]int) // level -> count
|
||||
seen := make(map[string]bool) // dedup guard
|
||||
err := nav.walk(context.Background(), []*File{rootFile}, 1000, 10, func(files []*File, level int) error {
|
||||
emitted[level] += len(files)
|
||||
for _, f := range files {
|
||||
key := fmt.Sprintf("%d:%d", level, f.ID())
|
||||
require.False(t, seen[key], "file %d emitted twice", f.ID())
|
||||
seen[key] = true
|
||||
}
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
// level 0: root. level 1: 3 dirs. level 2: 15 files.
|
||||
require.Equal(t, 1, emitted[0])
|
||||
require.Equal(t, 3, emitted[1])
|
||||
require.Equal(t, 15, emitted[2])
|
||||
require.Equal(t, 19, len(seen))
|
||||
}
|
||||
|
||||
func TestWalkBatchesWideLevels(t *testing.T) {
|
||||
client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared")
|
||||
t.Cleanup(func() { require.NoError(t, client.Close()) })
|
||||
// One folder with more children than walkFetchPageSize forces multiple
|
||||
// callback invocations for the same level — the pre-fix implementation
|
||||
// fetched and emitted the entire level at once.
|
||||
user, rootFile := buildWalkTree(t, client, 1, walkFetchPageSize+500)
|
||||
|
||||
nav := walkTestNavigator(t, client, user)
|
||||
|
||||
levelCalls := make(map[int]int)
|
||||
total := 0
|
||||
err := nav.walk(context.Background(), []*File{rootFile}, 10000, 10, func(files []*File, level int) error {
|
||||
levelCalls[level]++
|
||||
total += len(files)
|
||||
require.LessOrEqual(t, len(files), walkFetchPageSize, "callback batch exceeded page bound")
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
// root + 1 dir + (walkFetchPageSize+500) files
|
||||
require.Equal(t, 2+walkFetchPageSize+500, total)
|
||||
require.Greater(t, levelCalls[2], 1, "wide level must be emitted in multiple batches")
|
||||
}
|
||||
|
||||
func TestWalkLimit(t *testing.T) {
|
||||
client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared")
|
||||
t.Cleanup(func() { require.NoError(t, client.Close()) })
|
||||
user, rootFile := buildWalkTree(t, client, 2, 10)
|
||||
|
||||
nav := walkTestNavigator(t, client, user)
|
||||
|
||||
emitted := 0
|
||||
err := nav.walk(context.Background(), []*File{rootFile}, 5, 10, func(files []*File, level int) error {
|
||||
emitted += len(files)
|
||||
return nil
|
||||
})
|
||||
require.ErrorIs(t, err, ErrFileCountLimitedReached)
|
||||
require.LessOrEqual(t, emitted, 5)
|
||||
}
|
||||
|
||||
// TestDeleteFilesBatched walks a nested tree inside a tx and asserts the
|
||||
// batched delete removes every row — including grandchildren of folders
|
||||
// whose deletion is deferred until after the walk (HasParentWith relies
|
||||
// on parent rows staying alive mid-walk).
|
||||
func TestDeleteFilesBatched(t *testing.T) {
|
||||
client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared")
|
||||
t.Cleanup(func() { require.NoError(t, client.Close()) })
|
||||
ctx := context.Background()
|
||||
|
||||
group := client.Group.Create().SetName("deleters").SetPermissions(&boolset.BooleanSet{}).SaveX(ctx)
|
||||
user := client.User.Create().SetEmail("del@example.com").SetNick("del").SetGroup(group).SaveX(ctx)
|
||||
policy := client.StoragePolicy.Create().SetName("local").SetType("local").SaveX(ctx)
|
||||
root := client.File.Create().SetName(inventory.RootFolderName).SetType(int(types.FileTypeFolder)).SetOwner(user).SaveX(ctx)
|
||||
|
||||
// dir -> sub -> leaf; dir also holds two files.
|
||||
dir := client.File.Create().SetName("dir").SetType(int(types.FileTypeFolder)).SetOwner(user).SetParent(root).SaveX(ctx)
|
||||
sub := client.File.Create().SetName("sub").SetType(int(types.FileTypeFolder)).SetOwner(user).SetParent(dir).SaveX(ctx)
|
||||
leaf := client.File.Create().SetName("leaf.txt").SetType(int(types.FileTypeFile)).SetOwner(user).SetParent(sub).SaveX(ctx)
|
||||
f1 := client.File.Create().SetName("a.txt").SetType(int(types.FileTypeFile)).SetOwner(user).SetParent(dir).SaveX(ctx)
|
||||
f2 := client.File.Create().SetName("b.txt").SetType(int(types.FileTypeFile)).SetOwner(user).SetParent(dir).SaveX(ctx)
|
||||
|
||||
for _, fm := range []*ent.File{leaf, f1, f2} {
|
||||
client.Entity.Create().
|
||||
SetType(1).SetSource("src").SetSize(10).
|
||||
SetStoragePolicyEntities(policy.ID).
|
||||
AddFileIDs(fm.ID).
|
||||
SaveX(ctx)
|
||||
}
|
||||
|
||||
fc := inventory.NewFileClient(client, conf.SQLiteDB, nil)
|
||||
uc := inventory.NewUserClient(client)
|
||||
txFc, tx, txCtx, err := inventory.WithTx(ctx, fc)
|
||||
require.NoError(t, err)
|
||||
txCtx = context.WithValue(txCtx, inventory.LoadFileEntity{}, true)
|
||||
|
||||
// Targets: dir wrapped as *File with entities edge loaded.
|
||||
dirModel := client.File.Query().WithEntities().Where(entfile.IDEQ(dir.ID)).OnlyX(txCtx)
|
||||
dirFile := newFile(nil, dirModel)
|
||||
dirFile.OwnerModel = user
|
||||
|
||||
user = client.User.Query().WithGroup().Where(entuser.IDEQ(user.ID)).OnlyX(txCtx)
|
||||
f := &DBFS{user: user}
|
||||
targets := map[Navigator][]*File{
|
||||
&myNavigator{baseNavigator: newBaseNavigator(txFc, defaultFilter, user, nil, nil), user: user, fileClient: txFc, userClient: uc}: {dirFile},
|
||||
}
|
||||
|
||||
stale, diff, indexToDelete, err := f.deleteFiles(txCtx, targets, txFc, nil)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, inventory.Commit(tx))
|
||||
|
||||
// dir, sub, leaf, f1, f2 all deleted; root survives.
|
||||
remaining := client.File.Query().AllX(ctx)
|
||||
require.Len(t, remaining, 1)
|
||||
require.Equal(t, root.ID, remaining[0].ID)
|
||||
require.ElementsMatch(t, []int{dir.ID, sub.ID, leaf.ID, f1.ID, f2.ID}, indexToDelete)
|
||||
require.Len(t, stale, 3)
|
||||
require.Equal(t, int64(-30), diff[user.ID])
|
||||
}
|
||||
|
||||
func TestWalkDepthLimit(t *testing.T) {
|
||||
client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared")
|
||||
t.Cleanup(func() { require.NoError(t, client.Close()) })
|
||||
user, rootFile := buildWalkTree(t, client, 2, 4)
|
||||
|
||||
nav := walkTestNavigator(t, client, user)
|
||||
|
||||
maxLevel := -1
|
||||
err := nav.walk(context.Background(), []*File{rootFile}, 1000, 1, func(files []*File, level int) error {
|
||||
if level > maxLevel {
|
||||
maxLevel = level
|
||||
}
|
||||
return nil
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, maxLevel)
|
||||
}
|
||||
@ -0,0 +1,139 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/cloudreve/Cloudreve/v4/pkg/cache"
|
||||
"github.com/cloudreve/Cloudreve/v4/pkg/filemanager/fs"
|
||||
)
|
||||
|
||||
func TestMergeByteRanges(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
ranges [][2]int64
|
||||
start int64
|
||||
end int64
|
||||
want [][2]int64
|
||||
}{
|
||||
{"empty", nil, 0, 10, [][2]int64{{0, 10}}},
|
||||
{"append gap", [][2]int64{{0, 10}}, 20, 30, [][2]int64{{0, 10}, {20, 30}}},
|
||||
{"prepend gap", [][2]int64{{20, 30}}, 0, 10, [][2]int64{{0, 10}, {20, 30}}},
|
||||
{"adjacent merge", [][2]int64{{0, 10}}, 10, 20, [][2]int64{{0, 20}}},
|
||||
{"bridge two", [][2]int64{{0, 10}, {20, 30}}, 10, 20, [][2]int64{{0, 30}}},
|
||||
{"contained", [][2]int64{{0, 30}}, 5, 15, [][2]int64{{0, 30}}},
|
||||
{"extend right", [][2]int64{{0, 10}}, 5, 20, [][2]int64{{0, 20}}},
|
||||
{"extend left", [][2]int64{{10, 20}}, 0, 15, [][2]int64{{0, 20}}},
|
||||
{"middle gap", [][2]int64{{0, 5}, {30, 40}}, 10, 20, [][2]int64{{0, 5}, {10, 20}, {30, 40}}},
|
||||
{"merge all", [][2]int64{{0, 5}, {10, 15}, {20, 25}}, 4, 21, [][2]int64{{0, 25}}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := mergeByteRanges(tt.ranges, tt.start, tt.end)
|
||||
if len(got) != len(tt.want) {
|
||||
t.Fatalf("got %v, want %v", got, tt.want)
|
||||
}
|
||||
for i := range got {
|
||||
if got[i] != tt.want[i] {
|
||||
t.Fatalf("got %v, want %v", got, tt.want)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRangesCoverFull(t *testing.T) {
|
||||
if !rangesCoverFull([][2]int64{{0, 100}}, 100) {
|
||||
t.Fatal("expected full coverage")
|
||||
}
|
||||
if rangesCoverFull([][2]int64{{0, 99}}, 100) {
|
||||
t.Fatal("partial coverage reported as full")
|
||||
}
|
||||
if rangesCoverFull([][2]int64{{0, 50}, {50, 100}}, 100) {
|
||||
t.Fatal("unmerged intervals should not report full coverage")
|
||||
}
|
||||
if rangesCoverFull(nil, 100) {
|
||||
t.Fatal("empty intervals reported as full")
|
||||
}
|
||||
}
|
||||
|
||||
func newRangedSession(id string, size int64) fs.UploadSession {
|
||||
return fs.UploadSession{
|
||||
Props: &fs.UploadProps{
|
||||
UploadSessionID: id,
|
||||
Size: size,
|
||||
ExpireAt: time.Now().Add(time.Hour),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestMarkRangeUploaded(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
m := &manager{kv: cache.NewMemoStore("", nil)}
|
||||
session := newRangedSession("test-range-session", 100)
|
||||
if err := m.kv.Set(UploadSessionCachePrefix+"test-range-session", session, 60); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Out-of-order arrival: second half first.
|
||||
all, err := m.MarkRangeUploaded(ctx, &session, 50, 50)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if all {
|
||||
t.Fatal("reported complete after only second half")
|
||||
}
|
||||
|
||||
all, err = m.MarkRangeUploaded(ctx, &session, 0, 50)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !all {
|
||||
t.Fatal("expected complete after full coverage")
|
||||
}
|
||||
|
||||
// Session record carries merged coverage.
|
||||
raw, ok := m.kv.Get(UploadSessionCachePrefix + "test-range-session")
|
||||
if !ok {
|
||||
t.Fatal("session missing from KV")
|
||||
}
|
||||
stored := raw.(fs.UploadSession)
|
||||
if len(stored.RangesReceived) != 1 || stored.RangesReceived[0] != [2]int64{0, 100} {
|
||||
t.Fatalf("unexpected ranges: %v", stored.RangesReceived)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMarkRangeUploadedGap(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
m := &manager{kv: cache.NewMemoStore("", nil)}
|
||||
session := newRangedSession("test-gap-session", 100)
|
||||
if err := m.kv.Set(UploadSessionCachePrefix+"test-gap-session", session, 60); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
for _, r := range [][2]int64{{0, 40}, {60, 100}} {
|
||||
all, err := m.MarkRangeUploaded(ctx, &session, r[0], r[1]-r[0])
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if all {
|
||||
t.Fatal("reported complete with a gap in coverage")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMarkRangeUploadedMissingSession(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
m := &manager{kv: cache.NewMemoStore("", nil)}
|
||||
session := newRangedSession("gone", 100)
|
||||
|
||||
all, err := m.MarkRangeUploaded(ctx, &session, 0, 100)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if all {
|
||||
t.Fatal("missing session reported complete")
|
||||
}
|
||||
}
|
||||
Loading…
Reference in new issue