parent
88409cc1f0
commit
1d52acfda2
Binary file not shown.
@ -0,0 +1,3 @@
|
|||||||
|
package constant
|
||||||
|
|
||||||
|
// var HashIDTable = []int{0, 1, 2, 3, 4, 5}
|
||||||
@ -1,605 +0,0 @@
|
|||||||
package middleware
|
|
||||||
|
|
||||||
import (
|
|
||||||
"database/sql"
|
|
||||||
"errors"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/filesystem"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mq"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/serializer"
|
|
||||||
"github.com/qiniu/go-sdk/v7/auth/qbox"
|
|
||||||
"io/ioutil"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/auth"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
var mock sqlmock.Sqlmock
|
|
||||||
|
|
||||||
// TestMain 初始化数据库Mock
|
|
||||||
func TestMain(m *testing.M) {
|
|
||||||
var db *sql.DB
|
|
||||||
var err error
|
|
||||||
db, mock, err = sqlmock.New()
|
|
||||||
if err != nil {
|
|
||||||
panic("An error was not expected when opening a stub database connection")
|
|
||||||
}
|
|
||||||
model.DB, _ = gorm.Open("mysql", db)
|
|
||||||
defer db.Close()
|
|
||||||
m.Run()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCurrentUser(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/test", nil)
|
|
||||||
|
|
||||||
//session为空
|
|
||||||
sessionFunc := Session("233")
|
|
||||||
sessionFunc(c)
|
|
||||||
CurrentUser()(c)
|
|
||||||
user, _ := c.Get("user")
|
|
||||||
asserts.Nil(user)
|
|
||||||
|
|
||||||
//session正确
|
|
||||||
c, _ = gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/test", nil)
|
|
||||||
sessionFunc(c)
|
|
||||||
util.SetSession(c, map[string]interface{}{"user_id": 1})
|
|
||||||
rows := sqlmock.NewRows([]string{"id", "deleted_at", "email", "options"}).
|
|
||||||
AddRow(1, nil, "admin@cloudreve.org", "{}")
|
|
||||||
mock.ExpectQuery("^SELECT (.+)").WillReturnRows(rows)
|
|
||||||
CurrentUser()(c)
|
|
||||||
user, _ = c.Get("user")
|
|
||||||
asserts.NotNil(user)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAuthRequired(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/test", nil)
|
|
||||||
AuthRequiredFunc := AuthRequired()
|
|
||||||
|
|
||||||
// 未登录
|
|
||||||
AuthRequiredFunc(c)
|
|
||||||
asserts.NotNil(c)
|
|
||||||
|
|
||||||
// 类型错误
|
|
||||||
c.Set("user", 123)
|
|
||||||
AuthRequiredFunc(c)
|
|
||||||
asserts.NotNil(c)
|
|
||||||
|
|
||||||
// 正常
|
|
||||||
c.Set("user", &model.User{})
|
|
||||||
AuthRequiredFunc(c)
|
|
||||||
asserts.NotNil(c)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSignRequired(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/test", nil)
|
|
||||||
authInstance := auth.HMACAuth{SecretKey: []byte(util.RandStringRunes(256))}
|
|
||||||
SignRequiredFunc := SignRequired(authInstance)
|
|
||||||
|
|
||||||
// 鉴权失败
|
|
||||||
SignRequiredFunc(c)
|
|
||||||
asserts.NotNil(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
|
|
||||||
c, _ = gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("PUT", "/test", nil)
|
|
||||||
SignRequiredFunc(c)
|
|
||||||
asserts.NotNil(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
|
|
||||||
// Sign verify success
|
|
||||||
c, _ = gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("PUT", "/test", nil)
|
|
||||||
c.Request = auth.SignRequest(authInstance, c.Request, 0)
|
|
||||||
SignRequiredFunc(c)
|
|
||||||
asserts.NotNil(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWebDAVAuth(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
AuthFunc := WebDAVAuth()
|
|
||||||
|
|
||||||
// options请求跳过验证
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("OPTIONS", "/test", nil)
|
|
||||||
AuthFunc(c)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 请求HTTP Basic Auth
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/test", nil)
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.NotEmpty(c.Writer.Header()["WWW-Authenticate"])
|
|
||||||
}
|
|
||||||
|
|
||||||
// 用户名不存在
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/test", nil)
|
|
||||||
c.Request.Header = map[string][]string{
|
|
||||||
"Authorization": {"Basic d2hvQGNsb3VkcmV2ZS5vcmc6YWRtaW4="},
|
|
||||||
}
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "password", "email"}),
|
|
||||||
)
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal(c.Writer.Status(), http.StatusUnauthorized)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 密码错误
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/test", nil)
|
|
||||||
c.Request.Header = map[string][]string{
|
|
||||||
"Authorization": {"Basic d2hvQGNsb3VkcmV2ZS5vcmc6YWRtaW4="},
|
|
||||||
}
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "password", "email", "options"}).AddRow(1, "123", "who@cloudreve.org", "{}"),
|
|
||||||
)
|
|
||||||
// 查找密码
|
|
||||||
mock.ExpectQuery("SELECT(.+)webdav(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal(c.Writer.Status(), http.StatusUnauthorized)
|
|
||||||
}
|
|
||||||
|
|
||||||
//未启用 WebDAV
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/test", nil)
|
|
||||||
c.Request.Header = map[string][]string{
|
|
||||||
"Authorization": {"Basic d2hvQGNsb3VkcmV2ZS5vcmc6YWRtaW4="},
|
|
||||||
}
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows(
|
|
||||||
[]string{"id", "password", "email", "group_id", "options"}).
|
|
||||||
AddRow(1,
|
|
||||||
"rfBd67ti3SMtYvSg:ce6dc7bca4f17f2660e18e7608686673eae0fdf3",
|
|
||||||
"who@cloudreve.org",
|
|
||||||
1,
|
|
||||||
"{}",
|
|
||||||
),
|
|
||||||
)
|
|
||||||
mock.ExpectQuery("SELECT(.+)groups(.+)").WillReturnRows(sqlmock.NewRows([]string{"id", "web_dav_enabled"}).AddRow(1, false))
|
|
||||||
// 查找密码
|
|
||||||
mock.ExpectQuery("SELECT(.+)webdav(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal(c.Writer.Status(), http.StatusForbidden)
|
|
||||||
}
|
|
||||||
|
|
||||||
//正常
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/test", nil)
|
|
||||||
c.Request.Header = map[string][]string{
|
|
||||||
"Authorization": {"Basic d2hvQGNsb3VkcmV2ZS5vcmc6YWRtaW4="},
|
|
||||||
}
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows(
|
|
||||||
[]string{"id", "password", "email", "group_id", "options"}).
|
|
||||||
AddRow(1,
|
|
||||||
"rfBd67ti3SMtYvSg:ce6dc7bca4f17f2660e18e7608686673eae0fdf3",
|
|
||||||
"who@cloudreve.org",
|
|
||||||
1,
|
|
||||||
"{}",
|
|
||||||
),
|
|
||||||
)
|
|
||||||
mock.ExpectQuery("SELECT(.+)groups(.+)").WillReturnRows(sqlmock.NewRows([]string{"id", "web_dav_enabled"}).AddRow(1, true))
|
|
||||||
// 查找密码
|
|
||||||
mock.ExpectQuery("SELECT(.+)webdav(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal(c.Writer.Status(), 200)
|
|
||||||
_, ok := c.Get("user")
|
|
||||||
asserts.True(ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUseUploadSession(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
AuthFunc := UseUploadSession("local")
|
|
||||||
|
|
||||||
// sessionID 为空
|
|
||||||
{
|
|
||||||
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/remote/sessionID", nil)
|
|
||||||
authInstance := auth.HMACAuth{SecretKey: []byte("123")}
|
|
||||||
auth.SignRequest(authInstance, c.Request, 0)
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
cache.Set(
|
|
||||||
filesystem.UploadSessionCachePrefix+"testCallBackRemote",
|
|
||||||
serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{Type: "local"},
|
|
||||||
},
|
|
||||||
0,
|
|
||||||
)
|
|
||||||
cache.Deletes([]string{"1"}, "policy_")
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "group_id"}).AddRow(1, 1))
|
|
||||||
mock.ExpectQuery("SELECT(.+)groups(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "policies"}).AddRow(1, "[513]"))
|
|
||||||
mock.ExpectQuery("SELECT(.+)policies(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "secret_key"}).AddRow(2, "123"))
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{
|
|
||||||
{"sessionID", "testCallBackRemote"},
|
|
||||||
}
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/remote/testCallBackRemote", nil)
|
|
||||||
authInstance := auth.HMACAuth{SecretKey: []byte("123")}
|
|
||||||
auth.SignRequest(authInstance, c.Request, 0)
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUploadCallbackCheck(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
|
|
||||||
// 上传会话不存在
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{
|
|
||||||
{"sessionID", "testSessionNotExist"},
|
|
||||||
}
|
|
||||||
res := uploadCallbackCheck(c, "local")
|
|
||||||
a.Contains("上传会话不存在或已过期", res.Msg)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 上传策略不一致
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{
|
|
||||||
{"sessionID", "testPolicyNotMatch"},
|
|
||||||
}
|
|
||||||
cache.Set(
|
|
||||||
filesystem.UploadSessionCachePrefix+"testPolicyNotMatch",
|
|
||||||
serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{Type: "remote"},
|
|
||||||
},
|
|
||||||
0,
|
|
||||||
)
|
|
||||||
res := uploadCallbackCheck(c, "local")
|
|
||||||
a.Contains("Policy not supported", res.Msg)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 用户不存在
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{
|
|
||||||
{"sessionID", "testUserNotExist"},
|
|
||||||
}
|
|
||||||
cache.Set(
|
|
||||||
filesystem.UploadSessionCachePrefix+"testUserNotExist",
|
|
||||||
serializer.UploadSession{
|
|
||||||
UID: 313,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{Type: "remote"},
|
|
||||||
},
|
|
||||||
0,
|
|
||||||
)
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "group_id"}))
|
|
||||||
res := uploadCallbackCheck(c, "remote")
|
|
||||||
a.Contains("找不到用户", res.Msg)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
_, ok := cache.Get(filesystem.UploadSessionCachePrefix + "testUserNotExist")
|
|
||||||
a.False(ok)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoteCallbackAuth(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
AuthFunc := RemoteCallbackAuth()
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{SecretKey: "123"},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/remote/testCallBackRemote", nil)
|
|
||||||
authInstance := auth.HMACAuth{SecretKey: []byte("123")}
|
|
||||||
auth.SignRequest(authInstance, c.Request, 0)
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 签名错误
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{SecretKey: "123"},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/remote/testCallBackRemote", nil)
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestQiniuCallbackAuth(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
AuthFunc := QiniuCallbackAuth()
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{
|
|
||||||
SecretKey: "123",
|
|
||||||
AccessKey: "123",
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/qiniu/testCallBackQiniu", nil)
|
|
||||||
mac := qbox.NewMac("123", "123")
|
|
||||||
token, err := mac.SignRequest(c.Request)
|
|
||||||
asserts.NoError(err)
|
|
||||||
c.Request.Header["Authorization"] = []string{"QBox " + token}
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 验证失败
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{
|
|
||||||
SecretKey: "123",
|
|
||||||
AccessKey: "123",
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/qiniu/testCallBackQiniu", nil)
|
|
||||||
mac := qbox.NewMac("123", "1213")
|
|
||||||
token, err := mac.SignRequest(c.Request)
|
|
||||||
asserts.NoError(err)
|
|
||||||
c.Request.Header["Authorization"] = []string{"QBox " + token}
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOSSCallbackAuth(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
AuthFunc := OSSCallbackAuth()
|
|
||||||
|
|
||||||
// 签名验证失败
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{
|
|
||||||
SecretKey: "123",
|
|
||||||
AccessKey: "123",
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/oss/testCallBackOSS", nil)
|
|
||||||
mac := qbox.NewMac("123", "123")
|
|
||||||
token, err := mac.SignRequest(c.Request)
|
|
||||||
asserts.NoError(err)
|
|
||||||
c.Request.Header["Authorization"] = []string{"QBox " + token}
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{
|
|
||||||
SecretKey: "123",
|
|
||||||
AccessKey: "123",
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/oss/TnXx5E5VyfJUyM1UdkdDu1rtnJ34EbmH", ioutil.NopCloser(strings.NewReader(`{"name":"2f7b2ccf30e9270ea920f1ab8a4037a546a2f0d5.jpg","source_name":"1/1_hFRtDLgM_2f7b2ccf30e9270ea920f1ab8a4037a546a2f0d5.jpg","size":114020,"pic_info":"810,539"}`)))
|
|
||||||
c.Request.Header["Authorization"] = []string{"e5LwzwTkP9AFAItT4YzvdJOHd0Y0wqTMWhsV/h5SG90JYGAmMd+8LQyj96R+9qUfJWjMt6suuUh7LaOryR87Dw=="}
|
|
||||||
c.Request.Header["X-Oss-Pub-Key-Url"] = []string{"aHR0cHM6Ly9nb3NzcHVibGljLmFsaWNkbi5jb20vY2FsbGJhY2tfcHViX2tleV92MS5wZW0="}
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
type fakeRead string
|
|
||||||
|
|
||||||
func (r fakeRead) Read(p []byte) (int, error) {
|
|
||||||
return 0, errors.New("error")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUpyunCallbackAuth(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
AuthFunc := UpyunCallbackAuth()
|
|
||||||
|
|
||||||
// 无法获取请求正文
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{
|
|
||||||
SecretKey: "123",
|
|
||||||
AccessKey: "123",
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/upyun/testCallBackUpyun", ioutil.NopCloser(fakeRead("")))
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 正文MD5不一致
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{
|
|
||||||
SecretKey: "123",
|
|
||||||
AccessKey: "123",
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/upyun/testCallBackUpyun", ioutil.NopCloser(strings.NewReader("1")))
|
|
||||||
c.Request.Header["Content-Md5"] = []string{"123"}
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 签名不一致
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{
|
|
||||||
SecretKey: "123",
|
|
||||||
AccessKey: "123",
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/upyun/testCallBackUpyun", ioutil.NopCloser(strings.NewReader("1")))
|
|
||||||
c.Request.Header["Content-Md5"] = []string{"c4ca4238a0b923820dcc509a6f75849b"}
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{
|
|
||||||
SecretKey: "123",
|
|
||||||
AccessKey: "123",
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/upyun/testCallBackUpyun", ioutil.NopCloser(strings.NewReader("1")))
|
|
||||||
c.Request.Header["Content-Md5"] = []string{"c4ca4238a0b923820dcc509a6f75849b"}
|
|
||||||
c.Request.Header["Authorization"] = []string{"UPYUN 123:GWueK9x493BKFFk5gmfdO2Mn6EM="}
|
|
||||||
AuthFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestOneDriveCallbackAuth(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
AuthFunc := OneDriveCallbackAuth()
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{
|
|
||||||
{"sessionID", "TestOneDriveCallbackAuth"},
|
|
||||||
}
|
|
||||||
c.Set(filesystem.UploadSessionCtx, &serializer.UploadSession{
|
|
||||||
UID: 1,
|
|
||||||
VirtualPath: "/",
|
|
||||||
Policy: model.Policy{
|
|
||||||
SecretKey: "123",
|
|
||||||
AccessKey: "123",
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c.Request, _ = http.NewRequest("POST", "/api/v3/callback/upyun/TestOneDriveCallbackAuth", ioutil.NopCloser(strings.NewReader("1")))
|
|
||||||
res := mq.GlobalMQ.Subscribe("TestOneDriveCallbackAuth", 1)
|
|
||||||
AuthFunc(c)
|
|
||||||
select {
|
|
||||||
case <-res:
|
|
||||||
case <-time.After(time.Millisecond * 500):
|
|
||||||
asserts.Fail("mq message should be published")
|
|
||||||
}
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestIsAdmin(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
testFunc := IsAdmin()
|
|
||||||
|
|
||||||
// 非管理员
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set("user", &model.User{})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 是管理员
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
user := &model.User{}
|
|
||||||
user.Group.ID = 1
|
|
||||||
c.Set("user", user)
|
|
||||||
testFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 初始用户,非管理组
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
user := &model.User{}
|
|
||||||
user.Group.ID = 2
|
|
||||||
user.ID = 1
|
|
||||||
c.Set("user", user)
|
|
||||||
testFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,177 +0,0 @@
|
|||||||
package middleware
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"errors"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
type errReader int
|
|
||||||
|
|
||||||
func (errReader) Read(p []byte) (n int, err error) {
|
|
||||||
return 0, errors.New("test error")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCaptchaRequired_General(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
|
|
||||||
// 未启用验证码
|
|
||||||
{
|
|
||||||
cache.SetSettings(map[string]string{
|
|
||||||
"login_captcha": "0",
|
|
||||||
"captcha_type": "1",
|
|
||||||
"captcha_ReCaptchaSecret": "1",
|
|
||||||
"captcha_TCaptcha_SecretId": "1",
|
|
||||||
"captcha_TCaptcha_SecretKey": "1",
|
|
||||||
"captcha_TCaptcha_CaptchaAppId": "1",
|
|
||||||
"captcha_TCaptcha_AppSecretKey": "1",
|
|
||||||
}, "setting_")
|
|
||||||
TestFunc := CaptchaRequired("login_captcha")
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", nil)
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// body 无法读取
|
|
||||||
{
|
|
||||||
cache.SetSettings(map[string]string{
|
|
||||||
"login_captcha": "1",
|
|
||||||
"captcha_type": "1",
|
|
||||||
"captcha_ReCaptchaSecret": "1",
|
|
||||||
"captcha_TCaptcha_SecretId": "1",
|
|
||||||
"captcha_TCaptcha_SecretKey": "1",
|
|
||||||
"captcha_TCaptcha_CaptchaAppId": "1",
|
|
||||||
"captcha_TCaptcha_AppSecretKey": "1",
|
|
||||||
}, "setting_")
|
|
||||||
TestFunc := CaptchaRequired("login_captcha")
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", errReader(1))
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// body JSON 解析失败
|
|
||||||
{
|
|
||||||
cache.SetSettings(map[string]string{
|
|
||||||
"login_captcha": "1",
|
|
||||||
"captcha_type": "1",
|
|
||||||
"captcha_ReCaptchaSecret": "1",
|
|
||||||
"captcha_TCaptcha_SecretId": "1",
|
|
||||||
"captcha_TCaptcha_SecretKey": "1",
|
|
||||||
"captcha_TCaptcha_CaptchaAppId": "1",
|
|
||||||
"captcha_TCaptcha_AppSecretKey": "1",
|
|
||||||
}, "setting_")
|
|
||||||
TestFunc := CaptchaRequired("login_captcha")
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
r := bytes.NewReader([]byte("123"))
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", r)
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCaptchaRequired_Normal(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
|
|
||||||
// 验证码错误
|
|
||||||
{
|
|
||||||
cache.SetSettings(map[string]string{
|
|
||||||
"login_captcha": "1",
|
|
||||||
"captcha_type": "normal",
|
|
||||||
"captcha_ReCaptchaSecret": "1",
|
|
||||||
"captcha_TCaptcha_SecretId": "1",
|
|
||||||
"captcha_TCaptcha_SecretKey": "1",
|
|
||||||
"captcha_TCaptcha_CaptchaAppId": "1",
|
|
||||||
"captcha_TCaptcha_AppSecretKey": "1",
|
|
||||||
}, "setting_")
|
|
||||||
TestFunc := CaptchaRequired("login_captcha")
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
r := bytes.NewReader([]byte("{}"))
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", r)
|
|
||||||
Session("233")(c)
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCaptchaRequired_Recaptcha(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
|
|
||||||
// 无法初始化reCaptcha实例
|
|
||||||
{
|
|
||||||
cache.SetSettings(map[string]string{
|
|
||||||
"login_captcha": "1",
|
|
||||||
"captcha_type": "recaptcha",
|
|
||||||
"captcha_ReCaptchaSecret": "",
|
|
||||||
"captcha_TCaptcha_SecretId": "1",
|
|
||||||
"captcha_TCaptcha_SecretKey": "1",
|
|
||||||
"captcha_TCaptcha_CaptchaAppId": "1",
|
|
||||||
"captcha_TCaptcha_AppSecretKey": "1",
|
|
||||||
}, "setting_")
|
|
||||||
TestFunc := CaptchaRequired("login_captcha")
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
r := bytes.NewReader([]byte("{}"))
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", r)
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 验证码错误
|
|
||||||
{
|
|
||||||
cache.SetSettings(map[string]string{
|
|
||||||
"login_captcha": "1",
|
|
||||||
"captcha_type": "recaptcha",
|
|
||||||
"captcha_ReCaptchaSecret": "233",
|
|
||||||
"captcha_TCaptcha_SecretId": "1",
|
|
||||||
"captcha_TCaptcha_SecretKey": "1",
|
|
||||||
"captcha_TCaptcha_CaptchaAppId": "1",
|
|
||||||
"captcha_TCaptcha_AppSecretKey": "1",
|
|
||||||
}, "setting_")
|
|
||||||
TestFunc := CaptchaRequired("login_captcha")
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
r := bytes.NewReader([]byte("{}"))
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", r)
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCaptchaRequired_Tcaptcha(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
|
|
||||||
// 验证出错
|
|
||||||
{
|
|
||||||
cache.SetSettings(map[string]string{
|
|
||||||
"login_captcha": "1",
|
|
||||||
"captcha_type": "tcaptcha",
|
|
||||||
"captcha_ReCaptchaSecret": "",
|
|
||||||
"captcha_TCaptcha_SecretId": "1",
|
|
||||||
"captcha_TCaptcha_SecretKey": "1",
|
|
||||||
"captcha_TCaptcha_CaptchaAppId": "1",
|
|
||||||
"captcha_TCaptcha_AppSecretKey": "1",
|
|
||||||
}, "setting_")
|
|
||||||
TestFunc := CaptchaRequired("login_captcha")
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
r := bytes.NewReader([]byte("{}"))
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", r)
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,120 +0,0 @@
|
|||||||
package middleware
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/aria2/common"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/auth"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cluster"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mocks/controllermock"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"net/http/httptest"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestMasterMetadata(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
masterMetaDataFunc := MasterMetadata()
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request = httptest.NewRequest("GET", "/", nil)
|
|
||||||
|
|
||||||
c.Request.Header = map[string][]string{
|
|
||||||
"X-Cr-Site-Id": {"expectedSiteID"},
|
|
||||||
"X-Cr-Site-Url": {"expectedSiteURL"},
|
|
||||||
"X-Cr-Cloudreve-Version": {"expectedMasterVersion"},
|
|
||||||
}
|
|
||||||
masterMetaDataFunc(c)
|
|
||||||
siteID, _ := c.Get("MasterSiteID")
|
|
||||||
siteURL, _ := c.Get("MasterSiteURL")
|
|
||||||
siteVersion, _ := c.Get("MasterVersion")
|
|
||||||
|
|
||||||
a.Equal("expectedSiteID", siteID.(string))
|
|
||||||
a.Equal("expectedSiteURL", siteURL.(string))
|
|
||||||
a.Equal("expectedMasterVersion", siteVersion.(string))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveRPCSignRequired(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
np := &cluster.NodePool{}
|
|
||||||
np.Init()
|
|
||||||
slaveRPCSignRequiredFunc := SlaveRPCSignRequired(np)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
|
|
||||||
// id parse failed
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request = httptest.NewRequest("GET", "/", nil)
|
|
||||||
c.Request.Header.Set("X-Cr-Node-Id", "unknown")
|
|
||||||
slaveRPCSignRequiredFunc(c)
|
|
||||||
a.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// node id not exist
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request = httptest.NewRequest("GET", "/", nil)
|
|
||||||
c.Request.Header.Set("X-Cr-Node-Id", "38")
|
|
||||||
slaveRPCSignRequiredFunc(c)
|
|
||||||
a.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// success
|
|
||||||
{
|
|
||||||
authInstance := auth.HMACAuth{SecretKey: []byte("")}
|
|
||||||
np.Add(&model.Node{Model: gorm.Model{
|
|
||||||
ID: 38,
|
|
||||||
}})
|
|
||||||
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request = httptest.NewRequest("POST", "/", nil)
|
|
||||||
c.Request.Header.Set("X-Cr-Node-Id", "38")
|
|
||||||
c.Request = auth.SignRequest(authInstance, c.Request, 0)
|
|
||||||
slaveRPCSignRequiredFunc(c)
|
|
||||||
a.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUseSlaveAria2Instance(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
// MasterSiteID not set
|
|
||||||
{
|
|
||||||
testController := &controllermock.SlaveControllerMock{}
|
|
||||||
useSlaveAria2InstanceFunc := UseSlaveAria2Instance(testController)
|
|
||||||
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
|
||||||
c.Request = httptest.NewRequest("GET", "/", nil)
|
|
||||||
useSlaveAria2InstanceFunc(c)
|
|
||||||
a.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// Cannot get aria2 instances
|
|
||||||
{
|
|
||||||
testController := &controllermock.SlaveControllerMock{}
|
|
||||||
useSlaveAria2InstanceFunc := UseSlaveAria2Instance(testController)
|
|
||||||
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
|
||||||
c.Request = httptest.NewRequest("GET", "/", nil)
|
|
||||||
c.Set("MasterSiteID", "expectedSiteID")
|
|
||||||
testController.On("GetAria2Instance", "expectedSiteID").Return(&common.DummyAria2{}, errors.New("error"))
|
|
||||||
useSlaveAria2InstanceFunc(c)
|
|
||||||
a.True(c.IsAborted())
|
|
||||||
testController.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Success
|
|
||||||
{
|
|
||||||
testController := &controllermock.SlaveControllerMock{}
|
|
||||||
useSlaveAria2InstanceFunc := UseSlaveAria2Instance(testController)
|
|
||||||
c, _ := gin.CreateTestContext(httptest.NewRecorder())
|
|
||||||
c.Request = httptest.NewRequest("GET", "/", nil)
|
|
||||||
c.Set("MasterSiteID", "expectedSiteID")
|
|
||||||
testController.On("GetAria2Instance", "expectedSiteID").Return(&common.DummyAria2{}, nil)
|
|
||||||
useSlaveAria2InstanceFunc(c)
|
|
||||||
a.False(c.IsAborted())
|
|
||||||
res, _ := c.Get("MasterAria2Instance")
|
|
||||||
a.NotNil(res)
|
|
||||||
testController.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,57 +0,0 @@
|
|||||||
package middleware
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"net/http/httptest"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestValidateSourceLink(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
testFunc := ValidateSourceLink()
|
|
||||||
|
|
||||||
// ID 不存在
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
testFunc(c)
|
|
||||||
a.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// SourceLink 不存在
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set("object_id", 1)
|
|
||||||
mock.ExpectQuery("SELECT(.+)source_links(.+)").WithArgs(1).WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
testFunc(c)
|
|
||||||
a.True(c.IsAborted())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 原文件不存在
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set("object_id", 1)
|
|
||||||
mock.ExpectQuery("SELECT(.+)source_links(.+)").WithArgs(1).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").WithArgs(0).WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
testFunc(c)
|
|
||||||
a.True(c.IsAborted())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set("object_id", 1)
|
|
||||||
mock.ExpectQuery("SELECT(.+)source_links(.+)").WithArgs(1).WillReturnRows(sqlmock.NewRows([]string{"id", "file_id"}).AddRow(1, 2))
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").WithArgs(2).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)source_links").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
testFunc(c)
|
|
||||||
a.False(c.IsAborted())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
@ -1,144 +0,0 @@
|
|||||||
package middleware
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/bootstrap"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
testMock "github.com/stretchr/testify/mock"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"os"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
type StaticMock struct {
|
|
||||||
testMock.Mock
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m StaticMock) Open(name string) (http.File, error) {
|
|
||||||
args := m.Called(name)
|
|
||||||
return args.Get(0).(http.File), args.Error(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m StaticMock) Exists(prefix string, filepath string) bool {
|
|
||||||
args := m.Called(prefix, filepath)
|
|
||||||
return args.Bool(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFrontendFileHandler(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
|
|
||||||
// 静态资源未加载
|
|
||||||
{
|
|
||||||
TestFunc := FrontendFileHandler()
|
|
||||||
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", nil)
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// index.html 不存在
|
|
||||||
{
|
|
||||||
testStatic := &StaticMock{}
|
|
||||||
bootstrap.StaticFS = testStatic
|
|
||||||
testStatic.On("Open", "/index.html").
|
|
||||||
Return(&os.File{}, errors.New("error"))
|
|
||||||
TestFunc := FrontendFileHandler()
|
|
||||||
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", nil)
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// index.html 读取失败
|
|
||||||
{
|
|
||||||
file, _ := util.CreatNestedFile("tests/index.html")
|
|
||||||
file.Close()
|
|
||||||
testStatic := &StaticMock{}
|
|
||||||
bootstrap.StaticFS = testStatic
|
|
||||||
testStatic.On("Open", "/index.html").
|
|
||||||
Return(file, nil)
|
|
||||||
TestFunc := FrontendFileHandler()
|
|
||||||
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", nil)
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功且命中
|
|
||||||
{
|
|
||||||
file, _ := util.CreatNestedFile("tests/index.html")
|
|
||||||
defer file.Close()
|
|
||||||
testStatic := &StaticMock{}
|
|
||||||
bootstrap.StaticFS = testStatic
|
|
||||||
testStatic.On("Open", "/index.html").
|
|
||||||
Return(file, nil)
|
|
||||||
TestFunc := FrontendFileHandler()
|
|
||||||
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/", nil)
|
|
||||||
|
|
||||||
cache.Set("setting_siteName", "cloudreve", 0)
|
|
||||||
cache.Set("setting_siteKeywords", "cloudreve", 0)
|
|
||||||
cache.Set("setting_siteScript", "cloudreve", 0)
|
|
||||||
cache.Set("setting_pwa_small_icon", "cloudreve", 0)
|
|
||||||
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功且命中静态文件
|
|
||||||
{
|
|
||||||
file, _ := util.CreatNestedFile("tests/index.html")
|
|
||||||
defer file.Close()
|
|
||||||
testStatic := &StaticMock{}
|
|
||||||
bootstrap.StaticFS = testStatic
|
|
||||||
testStatic.On("Open", "/index.html").
|
|
||||||
Return(file, nil)
|
|
||||||
testStatic.On("Exists", "/", "/2").
|
|
||||||
Return(true)
|
|
||||||
testStatic.On("Open", "/2").
|
|
||||||
Return(file, nil)
|
|
||||||
TestFunc := FrontendFileHandler()
|
|
||||||
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/2", nil)
|
|
||||||
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
testStatic.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// API 相关跳过
|
|
||||||
{
|
|
||||||
for _, reqPath := range []string{"/api/user", "/manifest.json", "/dav/path"} {
|
|
||||||
file, _ := util.CreatNestedFile("tests/index.html")
|
|
||||||
defer file.Close()
|
|
||||||
testStatic := &StaticMock{}
|
|
||||||
bootstrap.StaticFS = testStatic
|
|
||||||
testStatic.On("Open", "/index.html").
|
|
||||||
Return(file, nil)
|
|
||||||
TestFunc := FrontendFileHandler()
|
|
||||||
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{}
|
|
||||||
c.Request, _ = http.NewRequest("GET", reqPath, nil)
|
|
||||||
|
|
||||||
TestFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
@ -1,37 +0,0 @@
|
|||||||
package middleware
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestMockHelper(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
MockHelperFunc := MockHelper()
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/test", nil)
|
|
||||||
|
|
||||||
// 写入session
|
|
||||||
{
|
|
||||||
SessionMock["test"] = "pass"
|
|
||||||
Session("test")(c)
|
|
||||||
MockHelperFunc(c)
|
|
||||||
asserts.Equal("pass", util.GetSession(c, "test").(string))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 写入context
|
|
||||||
{
|
|
||||||
ContextMock["test"] = "pass"
|
|
||||||
MockHelperFunc(c)
|
|
||||||
test, exist := c.Get("test")
|
|
||||||
asserts.True(exist)
|
|
||||||
asserts.Equal("pass", test.(string))
|
|
||||||
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,64 +0,0 @@
|
|||||||
package middleware
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestSession(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
{
|
|
||||||
handler := Session("2333")
|
|
||||||
asserts.NotNil(handler)
|
|
||||||
asserts.NotNil(Store)
|
|
||||||
asserts.IsType(emptyFunc(), handler)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func emptyFunc() gin.HandlerFunc {
|
|
||||||
return func(c *gin.Context) {}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCSRFInit(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
sessionFunc := Session("233")
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/test", nil)
|
|
||||||
sessionFunc(c)
|
|
||||||
CSRFInit()(c)
|
|
||||||
asserts.True(util.GetSession(c, "CSRF").(bool))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCSRFCheck(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
sessionFunc := Session("233")
|
|
||||||
|
|
||||||
// 通过检查
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/test", nil)
|
|
||||||
sessionFunc(c)
|
|
||||||
CSRFInit()(c)
|
|
||||||
CSRFCheck()(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 未通过检查
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request, _ = http.NewRequest("GET", "/test", nil)
|
|
||||||
sessionFunc(c)
|
|
||||||
CSRFCheck()(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,190 +0,0 @@
|
|||||||
package middleware
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/http/httptest"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/conf"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestShareAvailable(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
testFunc := ShareAvailable()
|
|
||||||
|
|
||||||
// 分享不存在
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{
|
|
||||||
{"id", "empty"},
|
|
||||||
}
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 通过
|
|
||||||
{
|
|
||||||
conf.SystemConfig.HashIDSalt = ""
|
|
||||||
// 用户组
|
|
||||||
mock.ExpectQuery("SELECT(.+)groups(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(3))
|
|
||||||
mock.ExpectQuery("SELECT(.+)shares(.+)").
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows(
|
|
||||||
[]string{"id", "remain_downloads", "source_id"}).
|
|
||||||
AddRow(1, 1, 2),
|
|
||||||
)
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2))
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Params = []gin.Param{
|
|
||||||
{"id", "x9T4"},
|
|
||||||
}
|
|
||||||
testFunc(c)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
asserts.NotNil(c.Get("user"))
|
|
||||||
asserts.NotNil(c.Get("share"))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestShareCanPreview(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
testFunc := ShareCanPreview()
|
|
||||||
|
|
||||||
// 无分享上下文
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 可以预览
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set("share", &model.Share{PreviewEnabled: true})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 未开启预览
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set("share", &model.Share{PreviewEnabled: false})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCheckShareUnlocked(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
testFunc := CheckShareUnlocked()
|
|
||||||
|
|
||||||
// 无分享上下文
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无密码
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set("share", &model.Share{})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBeforeShareDownload(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
testFunc := BeforeShareDownload()
|
|
||||||
|
|
||||||
// 无分享上下文
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
|
|
||||||
c, _ = gin.CreateTestContext(rec)
|
|
||||||
c.Set("share", &model.Share{})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 用户不能下载
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set("share", &model.Share{})
|
|
||||||
c.Set("user", &model.User{
|
|
||||||
Group: model.Group{OptionsSerialized: model.GroupOption{}},
|
|
||||||
})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 可以下载
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set("share", &model.Share{})
|
|
||||||
c.Set("user", &model.User{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
Group: model.Group{OptionsSerialized: model.GroupOption{
|
|
||||||
ShareDownload: true,
|
|
||||||
}},
|
|
||||||
})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestShareOwner(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
testFunc := ShareOwner()
|
|
||||||
|
|
||||||
// 未登录
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
|
|
||||||
c, _ = gin.CreateTestContext(rec)
|
|
||||||
c.Set("share", &model.Share{})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 非用户所创建分享
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
|
|
||||||
c, _ = gin.CreateTestContext(rec)
|
|
||||||
c.Set("share", &model.Share{User: model.User{Model: gorm.Model{ID: 1}}})
|
|
||||||
c.Set("user", &model.User{})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 正常
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
|
|
||||||
c, _ = gin.CreateTestContext(rec)
|
|
||||||
c.Set("share", &model.Share{})
|
|
||||||
c.Set("user", &model.User{})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,112 +0,0 @@
|
|||||||
package middleware
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mocks/wopimock"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/wopi"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"net/http/httptest"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestWopiWriteAccess(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
testFunc := WopiWriteAccess()
|
|
||||||
|
|
||||||
// deny preview only session
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(WopiSessionCtx, &wopi.SessionCache{Action: wopi.ActionPreview})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// pass
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Set(WopiSessionCtx, &wopi.SessionCache{Action: wopi.ActionEdit})
|
|
||||||
testFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWopiAccessValidation(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
mockWopi := &wopimock.WopiClientMock{}
|
|
||||||
mockCache := cache.NewMemoStore()
|
|
||||||
testFunc := WopiAccessValidation(mockWopi, mockCache)
|
|
||||||
|
|
||||||
// malformed access token
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.AddParam(wopi.AccessTokenQuery, "000")
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// session key not exist
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request = httptest.NewRequest("GET", "/wopi/files/1?access_token=", nil)
|
|
||||||
query := c.Request.URL.Query()
|
|
||||||
query.Set(wopi.AccessTokenQuery, "sessionID.key")
|
|
||||||
c.Request.URL.RawQuery = query.Encode()
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
}
|
|
||||||
|
|
||||||
// user key not exist
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request = httptest.NewRequest("GET", "/wopi/files/1?access_token=", nil)
|
|
||||||
query := c.Request.URL.Query()
|
|
||||||
query.Set(wopi.AccessTokenQuery, "sessionID.key")
|
|
||||||
c.Request.URL.RawQuery = query.Encode()
|
|
||||||
mockCache.Set(wopi.SessionCachePrefix+"sessionID", wopi.SessionCache{UserID: 1, FileID: 1}, 0)
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").WillReturnError(errors.New("error"))
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// file not found
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request = httptest.NewRequest("GET", "/wopi/files/1?access_token=", nil)
|
|
||||||
query := c.Request.URL.Query()
|
|
||||||
query.Set(wopi.AccessTokenQuery, "sessionID.key")
|
|
||||||
c.Request.URL.RawQuery = query.Encode()
|
|
||||||
mockCache.Set(wopi.SessionCachePrefix+"sessionID", wopi.SessionCache{UserID: 1, FileID: 1}, 0)
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
c.Set("object_id", uint(0))
|
|
||||||
testFunc(c)
|
|
||||||
asserts.True(c.IsAborted())
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// all pass
|
|
||||||
{
|
|
||||||
c, _ := gin.CreateTestContext(rec)
|
|
||||||
c.Request = httptest.NewRequest("GET", "/wopi/files/1?access_token=", nil)
|
|
||||||
query := c.Request.URL.Query()
|
|
||||||
query.Set(wopi.AccessTokenQuery, "sessionID.key")
|
|
||||||
c.Request.URL.RawQuery = query.Encode()
|
|
||||||
mockCache.Set(wopi.SessionCachePrefix+"sessionID", wopi.SessionCache{UserID: 1, FileID: 1}, 0)
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
c.Set("object_id", uint(1))
|
|
||||||
testFunc(c)
|
|
||||||
asserts.False(c.IsAborted())
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
c.MustGet(WopiSessionCtx)
|
|
||||||
})
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
c.MustGet("user")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,190 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestDownload_Create(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
download := Download{GID: "1"}
|
|
||||||
id, err := download.Create()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.EqualValues(1, id)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
download := Download{GID: "1"}
|
|
||||||
id, err := download.Create()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.EqualValues(0, id)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDownload_AfterFind(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
download := Download{Attrs: `{"gid":"123"}`}
|
|
||||||
err := download.AfterFind()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("123", download.StatusInfo.Gid)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 忽略空值
|
|
||||||
{
|
|
||||||
download := Download{Attrs: ``}
|
|
||||||
err := download.AfterFind()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("", download.StatusInfo.Gid)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 解析失败
|
|
||||||
{
|
|
||||||
download := Download{Attrs: `?`}
|
|
||||||
err := download.BeforeSave()
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Equal("", download.StatusInfo.Gid)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDownload_Save(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
download := Download{
|
|
||||||
Model: gorm.Model{
|
|
||||||
ID: 1,
|
|
||||||
},
|
|
||||||
Attrs: `{"gid":"123"}`,
|
|
||||||
}
|
|
||||||
err := download.Save()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("123", download.StatusInfo.Gid)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
download := Download{
|
|
||||||
Model: gorm.Model{
|
|
||||||
ID: 1,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
err := download.Save()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetDownloadsByStatus(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(0, 1).WillReturnRows(sqlmock.NewRows([]string{"gid"}).AddRow("0").AddRow("1"))
|
|
||||||
res := GetDownloadsByStatus(0, 1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Len(res, 2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetDownloadByGid(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(2, "gid").WillReturnRows(sqlmock.NewRows([]string{"g_id"}).AddRow("1"))
|
|
||||||
res, err := GetDownloadByGid("gid", 2)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(res.GID, "1")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDownload_GetOwner(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 已经有User对象
|
|
||||||
{
|
|
||||||
download := &Download{User: &User{Nick: "nick"}}
|
|
||||||
user := download.GetOwner()
|
|
||||||
asserts.NotNil(user)
|
|
||||||
asserts.Equal("nick", user.Nick)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无User对象
|
|
||||||
{
|
|
||||||
download := &Download{UserID: 3}
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"nick"}).AddRow("nick"))
|
|
||||||
user := download.GetOwner()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NotNil(user)
|
|
||||||
asserts.Equal("nick", user.Nick)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetDownloadsByStatusAndUser(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 列出全部
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 1, 2).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2).AddRow(3))
|
|
||||||
res := GetDownloadsByStatusAndUser(0, 1, 1, 2)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Len(res, 2)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 列出全部,分页
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)DESC(.+)").WithArgs(1, 1, 2).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2).AddRow(3))
|
|
||||||
res := GetDownloadsByStatusAndUser(2, 1, 1, 2)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Len(res, 2)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDownload_Delete(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
share := Download{}
|
|
||||||
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := share.Delete()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDownload_GetNodeID(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
record := Download{}
|
|
||||||
|
|
||||||
// compatible with 3.4
|
|
||||||
a.EqualValues(1, record.GetNodeID())
|
|
||||||
|
|
||||||
record.NodeID = 5
|
|
||||||
a.EqualValues(5, record.GetNodeID())
|
|
||||||
}
|
|
||||||
@ -1,785 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestFile_Create(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
file := File{
|
|
||||||
Name: "123",
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法插入文件记录
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
err := file.Create()
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法更新用户容量
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(5, 1))
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
err := file.Create()
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(5, 1))
|
|
||||||
mock.ExpectExec("UPDATE(.+)storage(.+)").WillReturnResult(sqlmock.NewResult(0, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := file.Create()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(uint(5), file.ID)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_AfterFind(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
// metadata not empty
|
|
||||||
{
|
|
||||||
file := File{
|
|
||||||
Name: "123",
|
|
||||||
Metadata: "{\"name\":\"123\"}",
|
|
||||||
}
|
|
||||||
|
|
||||||
a.NoError(file.AfterFind())
|
|
||||||
a.Equal("123", file.MetadataSerialized["name"])
|
|
||||||
}
|
|
||||||
|
|
||||||
// metadata empty
|
|
||||||
{
|
|
||||||
file := File{
|
|
||||||
Name: "123",
|
|
||||||
Metadata: "",
|
|
||||||
}
|
|
||||||
a.Nil(file.MetadataSerialized)
|
|
||||||
a.NoError(file.AfterFind())
|
|
||||||
a.NotNil(file.MetadataSerialized)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_BeforeSave(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
// metadata not empty
|
|
||||||
{
|
|
||||||
file := File{
|
|
||||||
Name: "123",
|
|
||||||
MetadataSerialized: map[string]string{
|
|
||||||
"name": "123",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
a.NoError(file.BeforeSave())
|
|
||||||
a.Equal("{\"name\":\"123\"}", file.Metadata)
|
|
||||||
}
|
|
||||||
|
|
||||||
// metadata empty
|
|
||||||
{
|
|
||||||
file := File{
|
|
||||||
Name: "123",
|
|
||||||
}
|
|
||||||
a.NoError(file.BeforeSave())
|
|
||||||
a.Equal("", file.Metadata)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFolder_GetChildFile(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
folder := Folder{Model: gorm.Model{ID: 1}, Name: "/"}
|
|
||||||
// 存在
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, "1.txt").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "1.txt"))
|
|
||||||
file, err := folder.GetChildFile("1.txt")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("1.txt", file.Name)
|
|
||||||
asserts.Equal("/", file.Position)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 不存在
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, "1.txt").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}))
|
|
||||||
_, err := folder.GetChildFile("1.txt")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFolder_GetChildFiles(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
folder := &Folder{
|
|
||||||
Model: gorm.Model{
|
|
||||||
ID: 1,
|
|
||||||
},
|
|
||||||
Position: "/123",
|
|
||||||
Name: "456",
|
|
||||||
}
|
|
||||||
|
|
||||||
// 找不到
|
|
||||||
mock.ExpectQuery("SELECT(.+)folder_id(.+)").WithArgs(1).WillReturnError(errors.New("error"))
|
|
||||||
files, err := folder.GetChildFiles()
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Len(files, 0)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
// 找到了
|
|
||||||
mock.ExpectQuery("SELECT(.+)folder_id(.+)").WithArgs(1).WillReturnRows(sqlmock.NewRows([]string{"name", "id"}).AddRow("1.txt", 1).AddRow("2.txt", 2))
|
|
||||||
files, err = folder.GetChildFiles()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(files, 2)
|
|
||||||
asserts.Equal("/123/456", files[0].Position)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetFilesByIDs(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 出错
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 2, 3, 1).
|
|
||||||
WillReturnError(errors.New("error"))
|
|
||||||
folders, err := GetFilesByIDs([]uint{1, 2, 3}, 1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Len(folders, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 部分找到
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 2, 3, 1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "1"))
|
|
||||||
folders, err := GetFilesByIDs([]uint{1, 2, 3}, 1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(folders, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 忽略UID查找
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 2, 3).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "1"))
|
|
||||||
folders, err := GetFilesByIDs([]uint{1, 2, 3}, 0)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(folders, 1)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetChildFilesOfFolders(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
testFolder := []Folder{
|
|
||||||
Folder{
|
|
||||||
Model: gorm.Model{ID: 3},
|
|
||||||
},
|
|
||||||
Folder{
|
|
||||||
Model: gorm.Model{ID: 4},
|
|
||||||
}, Folder{
|
|
||||||
Model: gorm.Model{ID: 5},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// 出错
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)folder_id").WithArgs(3, 4, 5).WillReturnError(errors.New("not found"))
|
|
||||||
files, err := GetChildFilesOfFolders(&testFolder)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Len(files, 0)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 找到2个
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)folder_id").
|
|
||||||
WithArgs(3, 4, 5).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).
|
|
||||||
AddRow(3, "3").
|
|
||||||
AddRow(4, "4"),
|
|
||||||
)
|
|
||||||
files, err := GetChildFilesOfFolders(&testFolder)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(files, 2)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 全部找到
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)folder_id").
|
|
||||||
WithArgs(3, 4, 5).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).
|
|
||||||
AddRow(3, "3").
|
|
||||||
AddRow(4, "4").
|
|
||||||
AddRow(5, "5"),
|
|
||||||
)
|
|
||||||
files, err := GetChildFilesOfFolders(&testFolder)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(files, 3)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetUploadPlaceholderFiles(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)upload_session_id(.+)").
|
|
||||||
WithArgs(1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "1"))
|
|
||||||
files := GetUploadPlaceholderFiles(1)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.Len(files, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_GetPolicy(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 空策略
|
|
||||||
{
|
|
||||||
file := File{
|
|
||||||
PolicyID: 23,
|
|
||||||
}
|
|
||||||
mock.ExpectQuery("SELECT(.+)policies(.+)").
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "name"}).
|
|
||||||
AddRow(23, "name"),
|
|
||||||
)
|
|
||||||
file.GetPolicy()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal(uint(23), file.Policy.ID)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 非空策略
|
|
||||||
{
|
|
||||||
file := File{
|
|
||||||
PolicyID: 23,
|
|
||||||
Policy: Policy{Model: gorm.Model{ID: 24}},
|
|
||||||
}
|
|
||||||
file.GetPolicy()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal(uint(24), file.Policy.ID)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoveFilesWithSoftLinks_EmptyArg(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
// 传入空
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)")
|
|
||||||
file, err := RemoveFilesWithSoftLinks([]File{})
|
|
||||||
asserts.Error(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(len(file), 0)
|
|
||||||
DB.Find(&File{})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoveFilesWithSoftLinks(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
files := []File{
|
|
||||||
File{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
SourceName: "1.txt",
|
|
||||||
PolicyID: 23,
|
|
||||||
},
|
|
||||||
File{
|
|
||||||
Model: gorm.Model{ID: 2},
|
|
||||||
SourceName: "2.txt",
|
|
||||||
PolicyID: 24,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// 传入空文件列表
|
|
||||||
{
|
|
||||||
file, err := RemoveFilesWithSoftLinks([]File{})
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Empty(file)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 全都没有
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs("1.txt", 23, 1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "policy_id", "source_name"}))
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs("2.txt", 24, 2).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "policy_id", "source_name"}))
|
|
||||||
file, err := RemoveFilesWithSoftLinks(files)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(files, file)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 第二个是软链
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs("1.txt", 23, 1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "policy_id", "source_name"}))
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs("2.txt", 24, 2).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "policy_id", "source_name"}).
|
|
||||||
AddRow(3, 24, "2.txt"),
|
|
||||||
)
|
|
||||||
file, err := RemoveFilesWithSoftLinks(files)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(files[:1], file)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 第一个是软链
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs("1.txt", 23, 1).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "policy_id", "source_name"}).
|
|
||||||
AddRow(3, 23, "1.txt"),
|
|
||||||
)
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs("2.txt", 24, 2).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "policy_id", "source_name"}))
|
|
||||||
file, err := RemoveFilesWithSoftLinks(files)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(files[1:], file)
|
|
||||||
}
|
|
||||||
// 全部是软链
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs("1.txt", 23, 1).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "policy_id", "source_name"}).
|
|
||||||
AddRow(3, 23, "1.txt"),
|
|
||||||
)
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs("2.txt", 24, 2).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "policy_id", "source_name"}).
|
|
||||||
AddRow(3, 24, "2.txt"),
|
|
||||||
)
|
|
||||||
file, err := RemoveFilesWithSoftLinks(files)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(file, 0)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDeleteFiles(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
// uid 不一致
|
|
||||||
{
|
|
||||||
err := DeleteFiles([]*File{{UserID: 2}}, 1)
|
|
||||||
a.Contains("user id not consistent", err.Error())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 删除失败
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("DELETE(.+)").
|
|
||||||
WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
err := DeleteFiles([]*File{{UserID: 1}}, 1)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法变更用户容量
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("DELETE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectExec("UPDATE(.+)storage(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
err := DeleteFiles([]*File{{UserID: 1}}, 1)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 文件脏读
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("DELETE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 0))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
err := DeleteFiles([]*File{{Size: 1, UserID: 1}, {Size: 2, UserID: 1}}, 1)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.Error(err)
|
|
||||||
a.Contains("file size is dirty", err.Error())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("DELETE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(2, 1))
|
|
||||||
mock.ExpectExec("DELETE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(2, 1))
|
|
||||||
mock.ExpectExec("UPDATE(.+)storage(.+)").WithArgs(uint64(3), sqlmock.AnyArg(), uint(1)).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := DeleteFiles([]*File{{Size: 1, UserID: 1}, {Size: 2, UserID: 1}}, 1)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功, 关联用户不存在
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("DELETE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(2, 1))
|
|
||||||
mock.ExpectExec("DELETE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(2, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := DeleteFiles([]*File{{Size: 1, UserID: 1}, {Size: 2, UserID: 1}}, 0)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.NoError(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetFilesByParentIDs(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 4, 5, 6).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "name"}).
|
|
||||||
AddRow(4, "4.txt").
|
|
||||||
AddRow(5, "5.txt").
|
|
||||||
AddRow(6, "6.txt"),
|
|
||||||
)
|
|
||||||
files, err := GetFilesByParentIDs([]uint{4, 5, 6}, 1)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Len(files, 3)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetFilesByUploadSession(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, "sessionID").
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "name"}).AddRow(4, "4.txt"))
|
|
||||||
files, err := GetFilesByUploadSession("sessionID", 1)
|
|
||||||
a.NoError(err)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.Equal("4.txt", files.Name)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_Updates(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
file := File{Model: gorm.Model{ID: 1}}
|
|
||||||
|
|
||||||
// rename
|
|
||||||
{
|
|
||||||
// not reset thumb
|
|
||||||
{
|
|
||||||
file := File{Model: gorm.Model{ID: 1}}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)files(.+)SET(.+)").WithArgs("", "newName", sqlmock.AnyArg(), 1).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := file.Rename("newName")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// thumb not available, rename base name only
|
|
||||||
{
|
|
||||||
file := File{Model: gorm.Model{ID: 1}, Name: "1.txt", MetadataSerialized: map[string]string{
|
|
||||||
ThumbStatusMetadataKey: ThumbStatusNotAvailable,
|
|
||||||
},
|
|
||||||
Metadata: "{}"}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)files(.+)SET(.+)").WithArgs("{}", "newName.txt", sqlmock.AnyArg(), 1).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := file.Rename("newName.txt")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(ThumbStatusNotAvailable, file.MetadataSerialized[ThumbStatusMetadataKey])
|
|
||||||
}
|
|
||||||
|
|
||||||
// thumb not available, rename base name only
|
|
||||||
{
|
|
||||||
file := File{Model: gorm.Model{ID: 1}, Name: "1.txt", MetadataSerialized: map[string]string{
|
|
||||||
ThumbStatusMetadataKey: ThumbStatusNotAvailable,
|
|
||||||
}}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)files(.+)SET(.+)").WithArgs("{}", "newName.jpg", sqlmock.AnyArg(), 1).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := file.Rename("newName.jpg")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Empty(file.MetadataSerialized[ThumbStatusMetadataKey])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// UpdatePicInfo
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WithArgs("1,1", 1).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := file.UpdatePicInfo("1,1")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// UpdateSourceName
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WithArgs("", "newName", sqlmock.AnyArg(), 1).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := file.UpdateSourceName("newName")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_UpdateSize(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
// 增加成功
|
|
||||||
{
|
|
||||||
file := File{Size: 10}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)files(.+)").WithArgs("", 11, sqlmock.AnyArg(), 10).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectExec("UPDATE(.+)storage(.+)+(.+)").WithArgs(uint64(1), sqlmock.AnyArg()).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.NoError(file.UpdateSize(11))
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 减少成功
|
|
||||||
{
|
|
||||||
file := File{Size: 10}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)files(.+)").WithArgs("", 8, sqlmock.AnyArg(), 10).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectExec("UPDATE(.+)storage(.+)-(.+)").WithArgs(uint64(2), sqlmock.AnyArg()).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.NoError(file.UpdateSize(8))
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 文件更新失败
|
|
||||||
{
|
|
||||||
file := File{Size: 10}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)files(.+)").WithArgs("", 8, sqlmock.AnyArg(), 10).WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
|
|
||||||
a.Error(file.UpdateSize(8))
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 用户容量更新失败
|
|
||||||
{
|
|
||||||
file := File{Size: 10}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)files(.+)").WithArgs("", 8, sqlmock.AnyArg(), 10).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectExec("UPDATE(.+)storage(.+)-(.+)").WithArgs(uint64(2), sqlmock.AnyArg()).WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
|
|
||||||
a.Error(file.UpdateSize(8))
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_PopChunkToFile(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
timeNow := time.Now()
|
|
||||||
file := File{}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)files(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
a.NoError(file.PopChunkToFile(&timeNow, "1,1"))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_CanCopy(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
file := File{}
|
|
||||||
a.True(file.CanCopy())
|
|
||||||
file.UploadSessionID = &file.Name
|
|
||||||
a.False(file.CanCopy())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_FileInfoInterface(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
file := File{
|
|
||||||
Model: gorm.Model{
|
|
||||||
UpdatedAt: time.Date(2019, 12, 21, 12, 40, 0, 0, time.UTC),
|
|
||||||
},
|
|
||||||
Name: "test_name",
|
|
||||||
SourceName: "",
|
|
||||||
UserID: 0,
|
|
||||||
Size: 10,
|
|
||||||
PicInfo: "",
|
|
||||||
FolderID: 0,
|
|
||||||
PolicyID: 0,
|
|
||||||
Policy: Policy{},
|
|
||||||
Position: "/test",
|
|
||||||
}
|
|
||||||
|
|
||||||
name := file.GetName()
|
|
||||||
asserts.Equal("test_name", name)
|
|
||||||
|
|
||||||
size := file.GetSize()
|
|
||||||
asserts.Equal(uint64(10), size)
|
|
||||||
|
|
||||||
asserts.Equal(time.Date(2019, 12, 21, 12, 40, 0, 0, time.UTC), file.ModTime())
|
|
||||||
asserts.False(file.IsDir())
|
|
||||||
asserts.Equal("/test", file.GetPosition())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetFilesByKeywords(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 未指定用户
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs("k1", "k2").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
res, err := GetFilesByKeywords(0, nil, "k1", "k2")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(res, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 指定用户
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, "k1", "k2").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
res, err := GetFilesByKeywords(1, nil, "k1", "k2")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(res, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 指定父目录
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 12, "k1", "k2").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
res, err := GetFilesByKeywords(1, []uint{12}, "k1", "k2")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(res, 1)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_CreateOrGetSourceLink(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
file := &File{}
|
|
||||||
file.ID = 1
|
|
||||||
|
|
||||||
// 已存在,返回老的 SourceLink
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)source_links(.+)").WithArgs(1).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2))
|
|
||||||
res, err := file.CreateOrGetSourceLink()
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues(2, res.ID)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 不存在,插入失败
|
|
||||||
{
|
|
||||||
expectedErr := errors.New("error")
|
|
||||||
mock.ExpectQuery("SELECT(.+)source_links(.+)").WithArgs(1).WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)source_links(.+)").WillReturnError(expectedErr)
|
|
||||||
mock.ExpectRollback()
|
|
||||||
res, err := file.CreateOrGetSourceLink()
|
|
||||||
a.Nil(res)
|
|
||||||
a.ErrorIs(err, expectedErr)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)source_links(.+)").WithArgs(1).WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)source_links(.+)").WillReturnResult(sqlmock.NewResult(2, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
res, err := file.CreateOrGetSourceLink()
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues(2, res.ID)
|
|
||||||
a.EqualValues(file.ID, res.File.ID)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_UpdateMetadata(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
file := &File{}
|
|
||||||
file.ID = 1
|
|
||||||
|
|
||||||
// 更新失败
|
|
||||||
{
|
|
||||||
expectedErr := errors.New("error")
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)files(.+)").WithArgs(sqlmock.AnyArg(), 1).WillReturnError(expectedErr)
|
|
||||||
mock.ExpectRollback()
|
|
||||||
a.ErrorIs(file.UpdateMetadata(map[string]string{"1": "1"}), expectedErr)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)files(.+)").WithArgs(sqlmock.AnyArg(), 1).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
a.NoError(file.UpdateMetadata(map[string]string{"1": "1"}))
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.Equal("1", file.MetadataSerialized["1"])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_ShouldLoadThumb(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
file := &File{
|
|
||||||
MetadataSerialized: map[string]string{},
|
|
||||||
}
|
|
||||||
file.ID = 1
|
|
||||||
|
|
||||||
// 无缩略图
|
|
||||||
{
|
|
||||||
file.MetadataSerialized[ThumbStatusMetadataKey] = ThumbStatusNotAvailable
|
|
||||||
a.False(file.ShouldLoadThumb())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 有缩略图
|
|
||||||
{
|
|
||||||
file.MetadataSerialized[ThumbStatusMetadataKey] = ThumbStatusExist
|
|
||||||
a.True(file.ShouldLoadThumb())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFile_ThumbFile(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
file := &File{
|
|
||||||
SourceName: "test",
|
|
||||||
MetadataSerialized: map[string]string{},
|
|
||||||
}
|
|
||||||
file.ID = 1
|
|
||||||
|
|
||||||
a.Equal("test._thumb", file.ThumbFile())
|
|
||||||
}
|
|
||||||
@ -1,622 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/conf"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestFolder_Create(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
folder := &Folder{
|
|
||||||
Name: "new folder",
|
|
||||||
}
|
|
||||||
|
|
||||||
// 不存在,插入成功
|
|
||||||
mock.ExpectQuery("SELECT(.+)folders(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(5, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
fid, err := folder.Create()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(uint(5), fid)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
// 插入失败
|
|
||||||
mock.ExpectQuery("SELECT(.+)folders(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
mock.ExpectQuery("SELECT(.+)folders(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
fid, err = folder.Create()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(uint(1), fid)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
// 存在,直接返回
|
|
||||||
mock.ExpectQuery("SELECT(.+)folders(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(5))
|
|
||||||
fid, err = folder.Create()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(uint(5), fid)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFolder_GetChild(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
folder := Folder{
|
|
||||||
Model: gorm.Model{ID: 5},
|
|
||||||
OwnerID: 1,
|
|
||||||
Name: "/",
|
|
||||||
}
|
|
||||||
|
|
||||||
// 目录存在
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(5, 1, "sub").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "sub"))
|
|
||||||
sub, err := folder.GetChild("sub")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(sub.Name, "sub")
|
|
||||||
asserts.Equal("/", sub.Position)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 目录不存在
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(5, 1, "sub").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}))
|
|
||||||
sub, err := folder.GetChild("sub")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Equal(uint(0), sub.ID)
|
|
||||||
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFolder_GetChildFolder(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
folder := &Folder{
|
|
||||||
Model: gorm.Model{
|
|
||||||
ID: 1,
|
|
||||||
},
|
|
||||||
Position: "/123",
|
|
||||||
Name: "456",
|
|
||||||
}
|
|
||||||
|
|
||||||
// 找不到
|
|
||||||
mock.ExpectQuery("SELECT(.+)parent_id(.+)").WithArgs(1).WillReturnError(errors.New("error"))
|
|
||||||
files, err := folder.GetChildFolder()
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Len(files, 0)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
// 找到了
|
|
||||||
mock.ExpectQuery("SELECT(.+)parent_id(.+)").WithArgs(1).WillReturnRows(sqlmock.NewRows([]string{"name", "id"}).AddRow("1.txt", 1).AddRow("2.txt", 2))
|
|
||||||
files, err = folder.GetChildFolder()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(files, 2)
|
|
||||||
asserts.Equal("/123/456", files[0].Position)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetRecursiveChildFolderSQLite(t *testing.T) {
|
|
||||||
conf.DatabaseConfig.Type = "sqlite"
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 测试目录结构
|
|
||||||
// 1
|
|
||||||
// 2 3
|
|
||||||
// 4 5 6
|
|
||||||
|
|
||||||
// 查询第一层
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 1).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "name"}).
|
|
||||||
AddRow(1, "folder1"),
|
|
||||||
)
|
|
||||||
// 查询第二层
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 1).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "name"}).
|
|
||||||
AddRow(2, "folder2").
|
|
||||||
AddRow(3, "folder3"),
|
|
||||||
)
|
|
||||||
// 查询第三层
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 2, 3).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "name"}).
|
|
||||||
AddRow(4, "folder4").
|
|
||||||
AddRow(5, "folder5").
|
|
||||||
AddRow(6, "folder6"),
|
|
||||||
)
|
|
||||||
// 查询第四层
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 4, 5, 6).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "name"}),
|
|
||||||
)
|
|
||||||
|
|
||||||
folders, err := GetRecursiveChildFolder([]uint{1}, 1, true)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Len(folders, 6)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDeleteFolderByIDs(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 出错
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("DELETE(.+)").
|
|
||||||
WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
err := DeleteFolderByIDs([]uint{1, 2, 3})
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("DELETE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(0, 3))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := DeleteFolderByIDs([]uint{1, 2, 3})
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetFoldersByIDs(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 出错
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 2, 3, 1).
|
|
||||||
WillReturnError(errors.New("error"))
|
|
||||||
folders, err := GetFoldersByIDs([]uint{1, 2, 3}, 1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Len(folders, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 部分找到
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 2, 3, 1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "1"))
|
|
||||||
folders, err := GetFoldersByIDs([]uint{1, 2, 3}, 1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(folders, 1)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFolder_MoveOrCopyFileTo(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
// 当前目录
|
|
||||||
folder := Folder{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
OwnerID: 1,
|
|
||||||
Name: "test",
|
|
||||||
}
|
|
||||||
// 目标目录
|
|
||||||
dstFolder := Folder{
|
|
||||||
Model: gorm.Model{ID: 10},
|
|
||||||
Name: "dst",
|
|
||||||
}
|
|
||||||
|
|
||||||
// 复制文件
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(
|
|
||||||
1,
|
|
||||||
2,
|
|
||||||
3,
|
|
||||||
1,
|
|
||||||
1,
|
|
||||||
).WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "size", "upload_session_id"}).
|
|
||||||
AddRow(1, 10, nil).
|
|
||||||
AddRow(2, 20, nil).
|
|
||||||
AddRow(2, 20, &folder.Name),
|
|
||||||
)
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
storage, err := folder.MoveOrCopyFileTo(
|
|
||||||
[]uint{1, 2, 3},
|
|
||||||
&dstFolder,
|
|
||||||
true,
|
|
||||||
)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal(uint64(30), storage)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 复制文件, 检索文件出错
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(
|
|
||||||
1,
|
|
||||||
2,
|
|
||||||
1,
|
|
||||||
1,
|
|
||||||
).WillReturnError(errors.New("error"))
|
|
||||||
|
|
||||||
storage, err := folder.MoveOrCopyFileTo(
|
|
||||||
[]uint{1, 2},
|
|
||||||
&dstFolder,
|
|
||||||
true,
|
|
||||||
)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal(uint64(0), storage)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 复制文件,第二个文件插入出错
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(
|
|
||||||
1,
|
|
||||||
2,
|
|
||||||
1,
|
|
||||||
1,
|
|
||||||
).WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "size"}).
|
|
||||||
AddRow(1, 10).
|
|
||||||
AddRow(2, 20),
|
|
||||||
)
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
storage, err := folder.MoveOrCopyFileTo(
|
|
||||||
[]uint{1, 2},
|
|
||||||
&dstFolder,
|
|
||||||
true,
|
|
||||||
)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal(uint64(10), storage)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 移动文件 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").
|
|
||||||
WithArgs(10, sqlmock.AnyArg(), 1, 2, 1, 1).
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 2))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
storage, err := folder.MoveOrCopyFileTo(
|
|
||||||
[]uint{1, 2},
|
|
||||||
&dstFolder,
|
|
||||||
false,
|
|
||||||
)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(uint64(0), storage)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 移动文件 出错
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").
|
|
||||||
WithArgs(10, sqlmock.AnyArg(), 1, 2, 1, 1).
|
|
||||||
WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
storage, err := folder.MoveOrCopyFileTo(
|
|
||||||
[]uint{1, 2},
|
|
||||||
&dstFolder,
|
|
||||||
false,
|
|
||||||
)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Equal(uint64(0), storage)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFolder_CopyFolderTo(t *testing.T) {
|
|
||||||
conf.DatabaseConfig.Type = "mysql"
|
|
||||||
asserts := assert.New(t)
|
|
||||||
// 父目录
|
|
||||||
parFolder := Folder{
|
|
||||||
Model: gorm.Model{ID: 9},
|
|
||||||
OwnerID: 1,
|
|
||||||
}
|
|
||||||
// 目标目录
|
|
||||||
dstFolder := Folder{
|
|
||||||
Model: gorm.Model{ID: 10},
|
|
||||||
}
|
|
||||||
|
|
||||||
// 测试复制目录结构
|
|
||||||
// test(2)(5)
|
|
||||||
// 1(3)(6) 2.txt
|
|
||||||
// 3(4)(7) 4.txt 5.txt(上传中)
|
|
||||||
|
|
||||||
// 正常情况 成功
|
|
||||||
{
|
|
||||||
// GetRecursiveChildFolder
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 2).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(2, 9))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 2).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(3, 2))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 3).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(4, 3))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 4).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}))
|
|
||||||
|
|
||||||
// 复制目录
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(5, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(6, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(7, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
// 查找子文件
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 2, 3, 4).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "name", "folder_id", "size", "upload_session_id"}).
|
|
||||||
AddRow(1, "2.txt", 2, 10, nil).
|
|
||||||
AddRow(2, "3.txt", 3, 20, nil).
|
|
||||||
AddRow(3, "5.txt", 3, 20, &dstFolder.Name),
|
|
||||||
)
|
|
||||||
|
|
||||||
// 复制子文件
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(5, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(6, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
size, err := parFolder.CopyFolderTo(2, &dstFolder)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(uint64(30), size)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 递归查询失败
|
|
||||||
{
|
|
||||||
// GetRecursiveChildFolder
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 2).WillReturnError(errors.New("error"))
|
|
||||||
|
|
||||||
size, err := parFolder.CopyFolderTo(2, &dstFolder)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Equal(uint64(0), size)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 父目录ID不存在
|
|
||||||
{
|
|
||||||
// GetRecursiveChildFolder
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 2).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(2, 9))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 2).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(3, 99))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 3).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(4, 3))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 4).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}))
|
|
||||||
|
|
||||||
// 复制目录
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(5, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
size, err := parFolder.CopyFolderTo(2, &dstFolder)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Equal(uint64(0), size)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 查询子文件失败
|
|
||||||
{
|
|
||||||
// GetRecursiveChildFolder
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 2).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(2, 9))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 2).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(3, 2))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 3).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(4, 3))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 4).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}))
|
|
||||||
|
|
||||||
// 复制目录
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(5, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(6, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(7, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
// 查找子文件
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 2, 3, 4).
|
|
||||||
WillReturnError(errors.New("error"))
|
|
||||||
|
|
||||||
size, err := parFolder.CopyFolderTo(2, &dstFolder)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Equal(uint64(0), size)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 复制文件 一个失败
|
|
||||||
{
|
|
||||||
// GetRecursiveChildFolder
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 2).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(2, 9))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 2).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(3, 2))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 3).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}).AddRow(4, 3))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 4).WillReturnRows(sqlmock.NewRows([]string{"id", "parent_id"}))
|
|
||||||
|
|
||||||
// 复制目录
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(5, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(6, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(7, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
// 查找子文件
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 2, 3, 4).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "name", "folder_id", "size"}).
|
|
||||||
AddRow(1, "2.txt", 2, 10).
|
|
||||||
AddRow(2, "3.txt", 3, 20),
|
|
||||||
)
|
|
||||||
|
|
||||||
// 复制子文件
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(5, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
|
|
||||||
size, err := parFolder.CopyFolderTo(2, &dstFolder)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Equal(uint64(10), size)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFolder_MoveOrCopyFolderTo_Move(t *testing.T) {
|
|
||||||
conf.DatabaseConfig.Type = "mysql"
|
|
||||||
asserts := assert.New(t)
|
|
||||||
// 父目录
|
|
||||||
parFolder := Folder{
|
|
||||||
Model: gorm.Model{ID: 9},
|
|
||||||
OwnerID: 1,
|
|
||||||
}
|
|
||||||
// 目标目录
|
|
||||||
dstFolder := Folder{
|
|
||||||
Model: gorm.Model{ID: 10},
|
|
||||||
OwnerID: 1,
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").
|
|
||||||
WithArgs(10, sqlmock.AnyArg(), 1, 2, 1, 9).
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := parFolder.MoveFolderTo([]uint{1, 2}, &dstFolder)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 移动自己到自己内部,失败
|
|
||||||
{
|
|
||||||
err := parFolder.MoveFolderTo([]uint{10, 2}, &dstFolder)
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFolder_FileInfoInterface(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
folder := Folder{
|
|
||||||
Model: gorm.Model{
|
|
||||||
UpdatedAt: time.Date(2019, 12, 21, 12, 40, 0, 0, time.UTC),
|
|
||||||
},
|
|
||||||
Name: "test_name",
|
|
||||||
OwnerID: 0,
|
|
||||||
Position: "/test",
|
|
||||||
}
|
|
||||||
|
|
||||||
name := folder.GetName()
|
|
||||||
asserts.Equal("test_name", name)
|
|
||||||
|
|
||||||
size := folder.GetSize()
|
|
||||||
asserts.Equal(uint64(0), size)
|
|
||||||
|
|
||||||
asserts.Equal(time.Date(2019, 12, 21, 12, 40, 0, 0, time.UTC), folder.ModTime())
|
|
||||||
asserts.True(folder.IsDir())
|
|
||||||
asserts.Equal("/test", folder.GetPosition())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTraceRoot(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
var parentId uint
|
|
||||||
parentId = 5
|
|
||||||
folder := Folder{
|
|
||||||
ParentID: &parentId,
|
|
||||||
OwnerID: 1,
|
|
||||||
Name: "test_name",
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(5, 1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name", "parent_id"}).AddRow(5, "parent", 1))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 0).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(5, "/"))
|
|
||||||
asserts.NoError(folder.TraceRoot())
|
|
||||||
asserts.Equal("/parent", folder.Position)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 出现错误
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(5, 1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name", "parent_id"}).AddRow(5, "parent", 1))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs(1, 0).
|
|
||||||
WillReturnError(errors.New("error"))
|
|
||||||
asserts.Error(folder.TraceRoot())
|
|
||||||
asserts.Equal("parent", folder.Position)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFolder_Rename(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
folder := Folder{
|
|
||||||
Model: gorm.Model{
|
|
||||||
ID: 1,
|
|
||||||
},
|
|
||||||
Name: "test_name",
|
|
||||||
OwnerID: 1,
|
|
||||||
Position: "/test",
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)folders(.+)SET(.+)").
|
|
||||||
WithArgs("test_name_new", 1).
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := folder.Rename("test_name_new")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 出现错误
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)folders(.+)SET(.+)").
|
|
||||||
WithArgs("test_name_new", 1).
|
|
||||||
WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
err := folder.Rename("test_name_new")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,77 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/pkg/errors"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestGetGroupByID(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
//找到用户组时
|
|
||||||
groupRows := sqlmock.NewRows([]string{"id", "name", "policies"}).
|
|
||||||
AddRow(1, "管理员", "[1]")
|
|
||||||
mock.ExpectQuery("^SELECT (.+)").WillReturnRows(groupRows)
|
|
||||||
|
|
||||||
group, err := GetGroupByID(1)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(Group{
|
|
||||||
Model: gorm.Model{
|
|
||||||
ID: 1,
|
|
||||||
},
|
|
||||||
Name: "管理员",
|
|
||||||
Policies: "[1]",
|
|
||||||
PolicyList: []uint{1},
|
|
||||||
}, group)
|
|
||||||
|
|
||||||
//未找到用户时
|
|
||||||
mock.ExpectQuery("^SELECT (.+)").WillReturnError(errors.New("not found"))
|
|
||||||
group, err = GetGroupByID(1)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Equal(Group{}, group)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGroup_AfterFind(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
testCase := Group{
|
|
||||||
Model: gorm.Model{
|
|
||||||
ID: 1,
|
|
||||||
},
|
|
||||||
Name: "管理员",
|
|
||||||
Policies: "[1]",
|
|
||||||
}
|
|
||||||
err := testCase.AfterFind()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(testCase.PolicyList, []uint{1})
|
|
||||||
|
|
||||||
testCase.Policies = "[1,2,3,4,5]"
|
|
||||||
err = testCase.AfterFind()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(testCase.PolicyList, []uint{1, 2, 3, 4, 5})
|
|
||||||
|
|
||||||
testCase.Policies = "[1,2,3,4,5"
|
|
||||||
err = testCase.AfterFind()
|
|
||||||
asserts.Error(err)
|
|
||||||
|
|
||||||
testCase.Policies = "[]"
|
|
||||||
err = testCase.AfterFind()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(testCase.PolicyList, []uint{})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGroup_BeforeSave(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
group := Group{
|
|
||||||
PolicyList: []uint{1, 2, 3},
|
|
||||||
}
|
|
||||||
{
|
|
||||||
err := group.BeforeSave()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("[1,2,3]", group.Policies)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
@ -1,21 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/conf"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestMigration(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
conf.DatabaseConfig.Type = "sqlite"
|
|
||||||
DB, _ = gorm.Open("sqlite", ":memory:")
|
|
||||||
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
migration()
|
|
||||||
})
|
|
||||||
conf.DatabaseConfig.Type = "mysql"
|
|
||||||
DB = mockDB
|
|
||||||
}
|
|
||||||
@ -1,64 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestGetNodeByID(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)nodes").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
res, err := GetNodeByID(1)
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues(1, res.ID)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetNodesByStatus(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)nodes").WillReturnRows(sqlmock.NewRows([]string{"status"}).AddRow(NodeActive))
|
|
||||||
res, err := GetNodesByStatus(NodeActive)
|
|
||||||
a.NoError(err)
|
|
||||||
a.Len(res, 1)
|
|
||||||
a.EqualValues(NodeActive, res[0].Status)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNode_AfterFind(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
node := &Node{}
|
|
||||||
|
|
||||||
// No aria2 options
|
|
||||||
{
|
|
||||||
a.NoError(node.AfterFind())
|
|
||||||
}
|
|
||||||
|
|
||||||
// with aria2 options
|
|
||||||
{
|
|
||||||
node.Aria2Options = `{"timeout":1}`
|
|
||||||
a.NoError(node.AfterFind())
|
|
||||||
a.Equal(1, node.Aria2OptionsSerialized.Timeout)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNode_BeforeSave(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
node := &Node{}
|
|
||||||
|
|
||||||
node.Aria2OptionsSerialized.Timeout = 1
|
|
||||||
a.NoError(node.BeforeSave())
|
|
||||||
a.Contains(node.Aria2Options, "1")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNode_SetStatus(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
node := &Node{}
|
|
||||||
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)nodes").WithArgs(NodeActive, sqlmock.AnyArg()).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
a.NoError(node.SetStatus(NodeActive))
|
|
||||||
a.Equal(NodeActive, node.Status)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
@ -0,0 +1,59 @@
|
|||||||
|
package model
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
||||||
|
"github.com/jinzhu/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// PackOrderType 容量包订单
|
||||||
|
PackOrderType = iota
|
||||||
|
// GroupOrderType 用户组订单
|
||||||
|
GroupOrderType
|
||||||
|
// ScoreOrderType 积分充值订单
|
||||||
|
ScoreOrderType
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// OrderUnpaid 未支付
|
||||||
|
OrderUnpaid = iota
|
||||||
|
// OrderPaid 已支付
|
||||||
|
OrderPaid
|
||||||
|
// OrderCanceled 已取消
|
||||||
|
OrderCanceled
|
||||||
|
)
|
||||||
|
|
||||||
|
// Order 交易订单
|
||||||
|
type Order struct {
|
||||||
|
gorm.Model
|
||||||
|
UserID uint // 创建者ID
|
||||||
|
OrderNo string `gorm:"index:order_number"` // 商户自定义订单编号
|
||||||
|
Type int // 订单类型
|
||||||
|
Method string // 支付类型
|
||||||
|
ProductID int64 // 商品ID
|
||||||
|
Num int // 商品数量
|
||||||
|
Name string // 订单标题
|
||||||
|
Price int // 商品单价
|
||||||
|
Status int // 订单状态
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create 创建订单记录
|
||||||
|
func (order *Order) Create() (uint, error) {
|
||||||
|
if err := DB.Create(order).Error; err != nil {
|
||||||
|
util.Log().Warning("Failed to insert order record: %s", err)
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return order.ID, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// UpdateStatus 更新订单状态
|
||||||
|
func (order *Order) UpdateStatus(status int) {
|
||||||
|
DB.Model(order).Update("status", status)
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetOrderByNo 根据商户订单号查询订单
|
||||||
|
func GetOrderByNo(id string) (*Order, error) {
|
||||||
|
var order Order
|
||||||
|
err := DB.Where("order_no = ?", id).First(&order).Error
|
||||||
|
return &order, err
|
||||||
|
}
|
||||||
@ -1,269 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"strconv"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestGetPolicyByID(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
cache.Deletes([]string{"22", "23"}, "policy_")
|
|
||||||
// 缓存未命中
|
|
||||||
{
|
|
||||||
rows := sqlmock.NewRows([]string{"name", "type", "options"}).
|
|
||||||
AddRow("默认存储策略", "local", "{\"od_redirect\":\"123\"}")
|
|
||||||
mock.ExpectQuery("^SELECT(.+)").WillReturnRows(rows)
|
|
||||||
policy, err := GetPolicyByID(uint(22))
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal("默认存储策略", policy.Name)
|
|
||||||
asserts.Equal("123", policy.OptionsSerialized.OauthRedirect)
|
|
||||||
|
|
||||||
rows = sqlmock.NewRows([]string{"name", "type", "options"})
|
|
||||||
mock.ExpectQuery("^SELECT(.+)").WillReturnRows(rows)
|
|
||||||
policy, err = GetPolicyByID(uint(23))
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 命中
|
|
||||||
{
|
|
||||||
policy, err := GetPolicyByID(uint(22))
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("默认存储策略", policy.Name)
|
|
||||||
asserts.Equal("123", policy.OptionsSerialized.OauthRedirect)
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPolicy_BeforeSave(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
testPolicy := Policy{
|
|
||||||
OptionsSerialized: PolicyOption{
|
|
||||||
OauthRedirect: "123",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
expected, _ := json.Marshal(testPolicy.OptionsSerialized)
|
|
||||||
err := testPolicy.BeforeSave()
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal(string(expected), testPolicy.Options)
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPolicy_GeneratePath(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
testPolicy := Policy{}
|
|
||||||
|
|
||||||
testPolicy.DirNameRule = "{randomkey16}"
|
|
||||||
asserts.Len(testPolicy.GeneratePath(1, "/"), 16)
|
|
||||||
|
|
||||||
testPolicy.DirNameRule = "{randomkey8}"
|
|
||||||
asserts.Len(testPolicy.GeneratePath(1, "/"), 8)
|
|
||||||
|
|
||||||
testPolicy.DirNameRule = "{timestamp}"
|
|
||||||
asserts.Equal(testPolicy.GeneratePath(1, "/"), strconv.FormatInt(time.Now().Unix(), 10))
|
|
||||||
|
|
||||||
testPolicy.DirNameRule = "{uid}"
|
|
||||||
asserts.Equal(testPolicy.GeneratePath(1, "/"), strconv.Itoa(int(1)))
|
|
||||||
|
|
||||||
testPolicy.DirNameRule = "{datetime}"
|
|
||||||
asserts.Len(testPolicy.GeneratePath(1, "/"), 14)
|
|
||||||
|
|
||||||
testPolicy.DirNameRule = "{date}"
|
|
||||||
asserts.Len(testPolicy.GeneratePath(1, "/"), 8)
|
|
||||||
|
|
||||||
testPolicy.DirNameRule = "123{date}ss{datetime}"
|
|
||||||
asserts.Len(testPolicy.GeneratePath(1, "/"), 27)
|
|
||||||
|
|
||||||
testPolicy.DirNameRule = "/1/{path}/456"
|
|
||||||
asserts.Condition(func() (success bool) {
|
|
||||||
res := testPolicy.GeneratePath(1, "/23")
|
|
||||||
return res == "/1/23/456" || res == "\\1\\23\\456"
|
|
||||||
})
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPolicy_GenerateFileName(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
// 重命名关闭
|
|
||||||
{
|
|
||||||
testPolicy := Policy{
|
|
||||||
AutoRename: false,
|
|
||||||
}
|
|
||||||
testPolicy.FileNameRule = "{randomkey16}"
|
|
||||||
asserts.Equal("123.txt", testPolicy.GenerateFileName(1, "123.txt"))
|
|
||||||
|
|
||||||
testPolicy.Type = "oss"
|
|
||||||
asserts.Equal("origin", testPolicy.GenerateFileName(1, "origin"))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 重命名开启
|
|
||||||
{
|
|
||||||
testPolicy := Policy{
|
|
||||||
AutoRename: true,
|
|
||||||
}
|
|
||||||
|
|
||||||
testPolicy.FileNameRule = "{randomkey16}"
|
|
||||||
asserts.Len(testPolicy.GenerateFileName(1, "123.txt"), 16)
|
|
||||||
|
|
||||||
testPolicy.FileNameRule = "{randomkey8}"
|
|
||||||
asserts.Len(testPolicy.GenerateFileName(1, "123.txt"), 8)
|
|
||||||
|
|
||||||
testPolicy.FileNameRule = "{timestamp}"
|
|
||||||
asserts.Equal(testPolicy.GenerateFileName(1, "123.txt"), strconv.FormatInt(time.Now().Unix(), 10))
|
|
||||||
|
|
||||||
testPolicy.FileNameRule = "{uid}"
|
|
||||||
asserts.Equal(testPolicy.GenerateFileName(1, "123.txt"), strconv.Itoa(int(1)))
|
|
||||||
|
|
||||||
testPolicy.FileNameRule = "{datetime}"
|
|
||||||
asserts.Len(testPolicy.GenerateFileName(1, "123.txt"), 14)
|
|
||||||
|
|
||||||
testPolicy.FileNameRule = "{date}"
|
|
||||||
asserts.Len(testPolicy.GenerateFileName(1, "123.txt"), 8)
|
|
||||||
|
|
||||||
testPolicy.FileNameRule = "123{date}ss{datetime}"
|
|
||||||
asserts.Len(testPolicy.GenerateFileName(1, "123.txt"), 27)
|
|
||||||
|
|
||||||
testPolicy.FileNameRule = "{originname_without_ext}"
|
|
||||||
asserts.Len(testPolicy.GenerateFileName(1, "123.txt"), 3)
|
|
||||||
|
|
||||||
testPolicy.FileNameRule = "{originname_without_ext}_{randomkey8}{ext}"
|
|
||||||
asserts.Len(testPolicy.GenerateFileName(1, "123.txt"), 16)
|
|
||||||
|
|
||||||
// 支持{originname}的策略
|
|
||||||
testPolicy.Type = "local"
|
|
||||||
testPolicy.FileNameRule = "123{originname}"
|
|
||||||
asserts.Equal("123123.txt", testPolicy.GenerateFileName(1, "123.txt"))
|
|
||||||
|
|
||||||
testPolicy.Type = "qiniu"
|
|
||||||
testPolicy.FileNameRule = "{uid}123{originname}"
|
|
||||||
asserts.Equal("1123123.txt", testPolicy.GenerateFileName(1, "123.txt"))
|
|
||||||
|
|
||||||
testPolicy.Type = "oss"
|
|
||||||
testPolicy.FileNameRule = "{uid}123{originname}"
|
|
||||||
asserts.Equal("1123123321", testPolicy.GenerateFileName(1, "123321"))
|
|
||||||
|
|
||||||
testPolicy.Type = "upyun"
|
|
||||||
testPolicy.FileNameRule = "{uid}123{originname}"
|
|
||||||
asserts.Equal("1123123321", testPolicy.GenerateFileName(1, "123321"))
|
|
||||||
|
|
||||||
testPolicy.Type = "qiniu"
|
|
||||||
testPolicy.FileNameRule = "{uid}123{originname}"
|
|
||||||
asserts.Equal("1123123321", testPolicy.GenerateFileName(1, "123321"))
|
|
||||||
|
|
||||||
testPolicy.Type = "local"
|
|
||||||
testPolicy.FileNameRule = "{uid}123{originname}"
|
|
||||||
asserts.Equal("1123", testPolicy.GenerateFileName(1, ""))
|
|
||||||
|
|
||||||
testPolicy.Type = "local"
|
|
||||||
testPolicy.FileNameRule = "{ext}123{uuid}"
|
|
||||||
asserts.Contains(testPolicy.GenerateFileName(1, "123.txt"), ".txt123")
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPolicy_IsDirectlyPreview(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
policy := Policy{Type: "local"}
|
|
||||||
asserts.True(policy.IsDirectlyPreview())
|
|
||||||
policy.Type = "remote"
|
|
||||||
asserts.False(policy.IsDirectlyPreview())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPolicy_ClearCache(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
cache.Set("policy_202", 1, 0)
|
|
||||||
policy := Policy{Model: gorm.Model{ID: 202}}
|
|
||||||
policy.ClearCache()
|
|
||||||
_, ok := cache.Get("policy_202")
|
|
||||||
asserts.False(ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPolicy_UpdateAccessKey(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
policy := Policy{Model: gorm.Model{ID: 202}}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
policy.AccessKey = "123"
|
|
||||||
err := policy.SaveAndClearCache()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPolicy_Props(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
policy := Policy{Type: "onedrive"}
|
|
||||||
policy.OptionsSerialized.PlaceholderWithSize = true
|
|
||||||
asserts.False(policy.IsThumbGenerateNeeded())
|
|
||||||
asserts.False(policy.IsTransitUpload(4))
|
|
||||||
asserts.False(policy.IsTransitUpload(5 * 1024 * 1024))
|
|
||||||
asserts.True(policy.CanStructureBeListed())
|
|
||||||
asserts.True(policy.IsUploadPlaceholderWithSize())
|
|
||||||
policy.Type = "local"
|
|
||||||
asserts.True(policy.IsThumbGenerateNeeded())
|
|
||||||
asserts.False(policy.CanStructureBeListed())
|
|
||||||
asserts.False(policy.IsUploadPlaceholderWithSize())
|
|
||||||
policy.Type = "remote"
|
|
||||||
asserts.True(policy.IsUploadPlaceholderWithSize())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPolicy_UpdateAccessKeyAndClearCache(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
cache.Set("policy_1331", Policy{}, 3600)
|
|
||||||
p := &Policy{}
|
|
||||||
p.ID = 1331
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WithArgs("ak", sqlmock.AnyArg()).WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.NoError(p.UpdateAccessKeyAndClearCache("ak"))
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
_, ok := cache.Get("policy_1331")
|
|
||||||
a.False(ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPolicy_CouldProxyThumb(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
p := &Policy{Type: "local"}
|
|
||||||
|
|
||||||
// local policy
|
|
||||||
{
|
|
||||||
a.False(p.CouldProxyThumb())
|
|
||||||
}
|
|
||||||
|
|
||||||
// feature not enabled
|
|
||||||
{
|
|
||||||
p.Type = "remote"
|
|
||||||
cache.Set("setting_thumb_proxy_enabled", "0", 0)
|
|
||||||
a.False(p.CouldProxyThumb())
|
|
||||||
}
|
|
||||||
|
|
||||||
// list not contain current policy
|
|
||||||
{
|
|
||||||
p.ID = 2
|
|
||||||
cache.Set("setting_thumb_proxy_enabled", "1", 0)
|
|
||||||
cache.Set("setting_thumb_proxy_policy", "[1]", 0)
|
|
||||||
a.False(p.CouldProxyThumb())
|
|
||||||
}
|
|
||||||
|
|
||||||
// enabled
|
|
||||||
{
|
|
||||||
p.ID = 2
|
|
||||||
cache.Set("setting_thumb_proxy_enabled", "1", 0)
|
|
||||||
cache.Set("setting_thumb_proxy_policy", "[2]", 0)
|
|
||||||
a.True(p.CouldProxyThumb())
|
|
||||||
}
|
|
||||||
|
|
||||||
cache.Deletes([]string{"thumb_proxy_enabled", "thumb_proxy_policy"}, "setting_")
|
|
||||||
}
|
|
||||||
@ -0,0 +1,27 @@
|
|||||||
|
package model
|
||||||
|
|
||||||
|
import "github.com/jinzhu/gorm"
|
||||||
|
|
||||||
|
// Redeem 兑换码
|
||||||
|
type Redeem struct {
|
||||||
|
gorm.Model
|
||||||
|
Type int // 订单类型
|
||||||
|
ProductID int64 // 商品ID
|
||||||
|
Num int // 商品数量
|
||||||
|
Code string `gorm:"size:64,index:redeem_code"` // 兑换码
|
||||||
|
Used bool // 是否已被使用
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetAvailableRedeem 根据code查找可用兑换码
|
||||||
|
func GetAvailableRedeem(code string) (*Redeem, error) {
|
||||||
|
redeem := &Redeem{}
|
||||||
|
result := DB.Where("code = ? and used = ?", code, false).First(redeem)
|
||||||
|
return redeem, result.Error
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use 设定为已使用状态
|
||||||
|
func (redeem *Redeem) Use() {
|
||||||
|
DB.Model(redeem).Updates(map[string]interface{}{
|
||||||
|
"used": true,
|
||||||
|
})
|
||||||
|
}
|
||||||
@ -0,0 +1,21 @@
|
|||||||
|
package model
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/jinzhu/gorm"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Report 举报模型
|
||||||
|
type Report struct {
|
||||||
|
gorm.Model
|
||||||
|
ShareID uint `gorm:"index:share_id"` // 对应分享ID
|
||||||
|
Reason int // 举报原因
|
||||||
|
Description string // 补充描述
|
||||||
|
|
||||||
|
// 关联模型
|
||||||
|
Share Share `gorm:"save_associations:false:false"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create 创建举报
|
||||||
|
func (report *Report) Create() error {
|
||||||
|
return DB.Create(report).Error
|
||||||
|
}
|
||||||
@ -1,39 +0,0 @@
|
|||||||
package invoker
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
type TestScript int
|
|
||||||
|
|
||||||
func (script TestScript) Run(ctx context.Context) {
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRunDBScript(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
Register("test", TestScript(0))
|
|
||||||
|
|
||||||
// 不存在
|
|
||||||
{
|
|
||||||
asserts.Error(RunDBScript("else", context.Background()))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 存在
|
|
||||||
{
|
|
||||||
asserts.NoError(RunDBScript("test", context.Background()))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestListPrefix(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
Register("U1", TestScript(0))
|
|
||||||
Register("U2", TestScript(0))
|
|
||||||
Register("U3", TestScript(0))
|
|
||||||
Register("P1", TestScript(0))
|
|
||||||
|
|
||||||
res := ListPrefix("U")
|
|
||||||
asserts.Len(res, 3)
|
|
||||||
}
|
|
||||||
@ -1,50 +0,0 @@
|
|||||||
package scripts
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestResetAdminPassword_Run(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
script := ResetAdminPassword(0)
|
|
||||||
|
|
||||||
// 初始用户不存在
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "email", "storage"}))
|
|
||||||
asserts.Panics(func() {
|
|
||||||
script.Run(context.Background())
|
|
||||||
})
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 密码更新失败
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "email", "storage"}).AddRow(1, "a@a.com", 10))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
asserts.Panics(func() {
|
|
||||||
script.Run(context.Background())
|
|
||||||
})
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "email", "storage"}).AddRow(1, "a@a.com", 10))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
script.Run(context.Background())
|
|
||||||
})
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,61 +0,0 @@
|
|||||||
package scripts
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"database/sql"
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
var mock sqlmock.Sqlmock
|
|
||||||
var mockDB *gorm.DB
|
|
||||||
|
|
||||||
// TestMain 初始化数据库Mock
|
|
||||||
func TestMain(m *testing.M) {
|
|
||||||
var db *sql.DB
|
|
||||||
var err error
|
|
||||||
db, mock, err = sqlmock.New()
|
|
||||||
if err != nil {
|
|
||||||
panic("An error was not expected when opening a stub database connection")
|
|
||||||
}
|
|
||||||
model.DB, _ = gorm.Open("mysql", db)
|
|
||||||
mockDB = model.DB
|
|
||||||
defer db.Close()
|
|
||||||
m.Run()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUserStorageCalibration_Run(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
script := UserStorageCalibration(0)
|
|
||||||
|
|
||||||
// 容量异常
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "email", "storage"}).AddRow(1, "a@a.com", 10))
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs(1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"total"}).AddRow(11))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
script.Run(context.Background())
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 容量正常
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "email", "storage"}).AddRow(1, "a@a.com", 10))
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs(1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"total"}).AddRow(10))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
script.Run(context.Background())
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -0,0 +1,22 @@
|
|||||||
|
package scripts
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
model "github.com/cloudreve/Cloudreve/v3/models"
|
||||||
|
)
|
||||||
|
|
||||||
|
type UpgradeToPro int
|
||||||
|
|
||||||
|
// Run 运行脚本从社区版升级至 Pro 版
|
||||||
|
func (script UpgradeToPro) Run(ctx context.Context) {
|
||||||
|
// folder.PolicyID 字段设为 0
|
||||||
|
model.DB.Model(model.Folder{}).UpdateColumn("policy_id", 0)
|
||||||
|
// shares.Score 字段设为0
|
||||||
|
model.DB.Model(model.Share{}).UpdateColumn("score", 0)
|
||||||
|
// user 表相关初始字段
|
||||||
|
model.DB.Model(model.User{}).Updates(map[string]interface{}{
|
||||||
|
"score": 0,
|
||||||
|
"previous_group_id": 0,
|
||||||
|
"open_id": "",
|
||||||
|
})
|
||||||
|
}
|
||||||
@ -1,66 +0,0 @@
|
|||||||
package scripts
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestUpgradeTo340_Run(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
script := UpgradeTo340(0)
|
|
||||||
|
|
||||||
// skip
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)settings").WillReturnRows(sqlmock.NewRows([]string{"name"}))
|
|
||||||
script.Run(context.Background())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// node not found
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)settings").WillReturnRows(sqlmock.NewRows([]string{"name"}).AddRow("1"))
|
|
||||||
mock.ExpectQuery("SELECT(.+)nodes").WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
script.Run(context.Background())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// success
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)settings").WillReturnRows(sqlmock.NewRows([]string{"name", "value"}).
|
|
||||||
AddRow("aria2_rpcurl", "expected_aria2_rpcurl").
|
|
||||||
AddRow("aria2_interval", "expected_aria2_interval").
|
|
||||||
AddRow("aria2_temp_path", "expected_aria2_temp_path").
|
|
||||||
AddRow("aria2_token", "expected_aria2_token").
|
|
||||||
AddRow("aria2_options", "{}"))
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)nodes").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
script.Run(context.Background())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// failed
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)settings").WillReturnRows(sqlmock.NewRows([]string{"name", "value"}).
|
|
||||||
AddRow("aria2_rpcurl", "expected_aria2_rpcurl").
|
|
||||||
AddRow("aria2_interval", "expected_aria2_interval").
|
|
||||||
AddRow("aria2_temp_path", "expected_aria2_temp_path").
|
|
||||||
AddRow("aria2_token", "expected_aria2_token").
|
|
||||||
AddRow("aria2_options", "{}"))
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)nodes").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
script.Run(context.Background())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,196 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"database/sql"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
var mock sqlmock.Sqlmock
|
|
||||||
var mockDB *gorm.DB
|
|
||||||
|
|
||||||
// TestMain 初始化数据库Mock
|
|
||||||
func TestMain(m *testing.M) {
|
|
||||||
var db *sql.DB
|
|
||||||
var err error
|
|
||||||
db, mock, err = sqlmock.New()
|
|
||||||
if err != nil {
|
|
||||||
panic("An error was not expected when opening a stub database connection")
|
|
||||||
}
|
|
||||||
DB, _ = gorm.Open("mysql", db)
|
|
||||||
mockDB = DB
|
|
||||||
defer db.Close()
|
|
||||||
m.Run()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetSettingByType(t *testing.T) {
|
|
||||||
cache.Store = cache.NewMemoStore()
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
//找到设置时
|
|
||||||
rows := sqlmock.NewRows([]string{"name", "value", "type"}).
|
|
||||||
AddRow("siteName", "Cloudreve", "basic").
|
|
||||||
AddRow("siteDes", "Something wonderful", "basic")
|
|
||||||
mock.ExpectQuery("^SELECT \\* FROM `(.+)` WHERE `(.+)`\\.`deleted_at` IS NULL AND(.+)$").WillReturnRows(rows)
|
|
||||||
settings := GetSettingByType([]string{"basic"})
|
|
||||||
asserts.Equal(map[string]string{
|
|
||||||
"siteName": "Cloudreve",
|
|
||||||
"siteDes": "Something wonderful",
|
|
||||||
}, settings)
|
|
||||||
|
|
||||||
rows = sqlmock.NewRows([]string{"name", "value", "type"}).
|
|
||||||
AddRow("siteName", "Cloudreve", "basic").
|
|
||||||
AddRow("siteDes", "Something wonderful", "basic2")
|
|
||||||
mock.ExpectQuery("^SELECT \\* FROM `(.+)` WHERE `(.+)`\\.`deleted_at` IS NULL AND(.+)$").WillReturnRows(rows)
|
|
||||||
settings = GetSettingByType([]string{"basic", "basic2"})
|
|
||||||
asserts.Equal(map[string]string{
|
|
||||||
"siteName": "Cloudreve",
|
|
||||||
"siteDes": "Something wonderful",
|
|
||||||
}, settings)
|
|
||||||
|
|
||||||
//找不到
|
|
||||||
rows = sqlmock.NewRows([]string{"name", "value", "type"})
|
|
||||||
mock.ExpectQuery("^SELECT \\* FROM `(.+)` WHERE `(.+)`\\.`deleted_at` IS NULL AND(.+)$").WillReturnRows(rows)
|
|
||||||
settings = GetSettingByType([]string{"basic233"})
|
|
||||||
asserts.Equal(map[string]string{}, settings)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetSettingByNameWithDefault(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
rows := sqlmock.NewRows([]string{"name", "value", "type"})
|
|
||||||
mock.ExpectQuery("^SELECT \\* FROM `(.+)` WHERE `(.+)`\\.`deleted_at` IS NULL AND(.+)$").WillReturnRows(rows)
|
|
||||||
settings := GetSettingByNameWithDefault("123", "123321")
|
|
||||||
a.Equal("123321", settings)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetSettingByNames(t *testing.T) {
|
|
||||||
cache.Store = cache.NewMemoStore()
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
//找到设置时
|
|
||||||
rows := sqlmock.NewRows([]string{"name", "value", "type"}).
|
|
||||||
AddRow("siteName", "Cloudreve", "basic").
|
|
||||||
AddRow("siteDes", "Something wonderful", "basic")
|
|
||||||
mock.ExpectQuery("^SELECT \\* FROM `(.+)` WHERE `(.+)`\\.`deleted_at` IS NULL AND(.+)$").WillReturnRows(rows)
|
|
||||||
settings := GetSettingByNames("siteName", "siteDes")
|
|
||||||
asserts.Equal(map[string]string{
|
|
||||||
"siteName": "Cloudreve",
|
|
||||||
"siteDes": "Something wonderful",
|
|
||||||
}, settings)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
//找到其中一个设置时
|
|
||||||
rows = sqlmock.NewRows([]string{"name", "value", "type"}).
|
|
||||||
AddRow("siteName2", "Cloudreve", "basic")
|
|
||||||
mock.ExpectQuery("^SELECT \\* FROM `(.+)` WHERE `(.+)`\\.`deleted_at` IS NULL AND(.+)$").WillReturnRows(rows)
|
|
||||||
settings = GetSettingByNames("siteName2", "siteDes2333")
|
|
||||||
asserts.Equal(map[string]string{
|
|
||||||
"siteName2": "Cloudreve",
|
|
||||||
}, settings)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
//找不到设置时
|
|
||||||
rows = sqlmock.NewRows([]string{"name", "value", "type"})
|
|
||||||
mock.ExpectQuery("^SELECT \\* FROM `(.+)` WHERE `(.+)`\\.`deleted_at` IS NULL AND(.+)$").WillReturnRows(rows)
|
|
||||||
settings = GetSettingByNames("siteName2333", "siteDes2333")
|
|
||||||
asserts.Equal(map[string]string{}, settings)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
// 一个设置命中缓存
|
|
||||||
mock.ExpectQuery("^SELECT \\* FROM `(.+)` WHERE `(.+)`\\.`deleted_at` IS NULL AND(.+)$").WithArgs("siteDes2").WillReturnRows(sqlmock.NewRows([]string{"name", "value", "type"}).
|
|
||||||
AddRow("siteDes2", "Cloudreve2", "basic"))
|
|
||||||
settings = GetSettingByNames("siteName", "siteDes2")
|
|
||||||
asserts.Equal(map[string]string{
|
|
||||||
"siteName": "Cloudreve",
|
|
||||||
"siteDes2": "Cloudreve2",
|
|
||||||
}, settings)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestGetSettingByName 测试GetSettingByName
|
|
||||||
func TestGetSettingByName(t *testing.T) {
|
|
||||||
cache.Store = cache.NewMemoStore()
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
//找到设置时
|
|
||||||
rows := sqlmock.NewRows([]string{"name", "value", "type"}).
|
|
||||||
AddRow("siteName", "Cloudreve", "basic")
|
|
||||||
mock.ExpectQuery("^SELECT \\* FROM `(.+)` WHERE `(.+)`\\.`deleted_at` IS NULL AND(.+)$").WillReturnRows(rows)
|
|
||||||
|
|
||||||
siteName := GetSettingByName("siteName")
|
|
||||||
asserts.Equal("Cloudreve", siteName)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
// 第二次查询应返回缓存内容
|
|
||||||
siteNameCache := GetSettingByName("siteName")
|
|
||||||
asserts.Equal("Cloudreve", siteNameCache)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
// 找不到设置
|
|
||||||
rows = sqlmock.NewRows([]string{"name", "value", "type"})
|
|
||||||
mock.ExpectQuery("^SELECT \\* FROM `(.+)` WHERE `(.+)`\\.`deleted_at` IS NULL AND(.+)$").WillReturnRows(rows)
|
|
||||||
|
|
||||||
siteName = GetSettingByName("siteName not exist")
|
|
||||||
asserts.Equal("", siteName)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestIsTrueVal(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
asserts.True(IsTrueVal("1"))
|
|
||||||
asserts.True(IsTrueVal("true"))
|
|
||||||
asserts.False(IsTrueVal("0"))
|
|
||||||
asserts.False(IsTrueVal("false"))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetSiteURL(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 正常
|
|
||||||
{
|
|
||||||
err := cache.Deletes([]string{"siteURL"}, "setting_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs("siteURL").WillReturnRows(sqlmock.NewRows([]string{"id", "value"}).AddRow(1, "https://drive.cloudreve.org"))
|
|
||||||
siteURL := GetSiteURL()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal("https://drive.cloudreve.org", siteURL.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 失败 返回默认值
|
|
||||||
{
|
|
||||||
err := cache.Deletes([]string{"siteURL"}, "setting_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WithArgs("siteURL").WillReturnRows(sqlmock.NewRows([]string{"id", "value"}).AddRow(1, ":][\\/\\]sdf"))
|
|
||||||
siteURL := GetSiteURL()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Equal("https://cloudreve.org", siteURL.String())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetIntSetting(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 正常
|
|
||||||
{
|
|
||||||
cache.Set("setting_TestGetIntSetting", "10", 0)
|
|
||||||
res := GetIntSetting("TestGetIntSetting", 20)
|
|
||||||
asserts.Equal(10, res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 使用默认值
|
|
||||||
{
|
|
||||||
res := GetIntSetting("TestGetIntSetting_2", 20)
|
|
||||||
asserts.Equal(20, res)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
@ -1,321 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"net/http/httptest"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/conf"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestShare_Create(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
share := Share{UserID: 1}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(2, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
id, err := share.Create()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.EqualValues(2, id)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
id, err := share.Create()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.EqualValues(0, id)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetShareByHashID(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
conf.SystemConfig.HashIDSalt = ""
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
res := GetShareByHashID("x9T4")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NotNil(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 查询失败
|
|
||||||
{
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WillReturnError(errors.New("error"))
|
|
||||||
res := GetShareByHashID("x9T4")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Nil(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ID解码失败
|
|
||||||
{
|
|
||||||
res := GetShareByHashID("empty")
|
|
||||||
asserts.Nil(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestShare_IsAvailable(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 下载剩余次数为0
|
|
||||||
{
|
|
||||||
share := Share{}
|
|
||||||
asserts.False(share.IsAvailable())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 时效过期
|
|
||||||
{
|
|
||||||
expires := time.Unix(10, 10)
|
|
||||||
share := Share{
|
|
||||||
RemainDownloads: -1,
|
|
||||||
Expires: &expires,
|
|
||||||
}
|
|
||||||
asserts.False(share.IsAvailable())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 源对象为目录,但不存在
|
|
||||||
{
|
|
||||||
share := Share{
|
|
||||||
RemainDownloads: -1,
|
|
||||||
SourceID: 2,
|
|
||||||
IsDir: true,
|
|
||||||
}
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
asserts.False(share.IsAvailable())
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 源对象为目录,存在
|
|
||||||
{
|
|
||||||
share := Share{
|
|
||||||
RemainDownloads: -1,
|
|
||||||
SourceID: 2,
|
|
||||||
IsDir: false,
|
|
||||||
}
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(13))
|
|
||||||
asserts.True(share.IsAvailable())
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 用户被封禁
|
|
||||||
{
|
|
||||||
share := Share{
|
|
||||||
RemainDownloads: -1,
|
|
||||||
SourceID: 2,
|
|
||||||
IsDir: true,
|
|
||||||
User: User{Status: Baned},
|
|
||||||
}
|
|
||||||
asserts.False(share.IsAvailable())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestShare_GetCreator(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
share := Share{UserID: 1}
|
|
||||||
res := share.Creator()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.EqualValues(1, res.ID)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestShare_Source(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 目录
|
|
||||||
{
|
|
||||||
share := Share{IsDir: true, SourceID: 3}
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(3))
|
|
||||||
asserts.EqualValues(3, share.Source().(*Folder).ID)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 文件
|
|
||||||
{
|
|
||||||
share := Share{IsDir: false, SourceID: 3}
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(3))
|
|
||||||
asserts.EqualValues(3, share.Source().(*File).ID)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestShare_CanBeDownloadBy(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
share := Share{}
|
|
||||||
|
|
||||||
// 未登录,无权
|
|
||||||
{
|
|
||||||
user := &User{
|
|
||||||
Group: Group{
|
|
||||||
OptionsSerialized: GroupOption{
|
|
||||||
ShareDownload: false,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
asserts.Error(share.CanBeDownloadBy(user))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 已登录,无权
|
|
||||||
{
|
|
||||||
user := &User{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
Group: Group{
|
|
||||||
OptionsSerialized: GroupOption{
|
|
||||||
ShareDownload: false,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
asserts.Error(share.CanBeDownloadBy(user))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
user := &User{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
Group: Group{
|
|
||||||
OptionsSerialized: GroupOption{
|
|
||||||
ShareDownload: true,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
asserts.NoError(share.CanBeDownloadBy(user))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestShare_WasDownloadedBy(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
share := Share{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
}
|
|
||||||
|
|
||||||
// 已登录,已下载
|
|
||||||
{
|
|
||||||
user := User{
|
|
||||||
Model: gorm.Model{
|
|
||||||
ID: 1,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
r := httptest.NewRecorder()
|
|
||||||
c, _ := gin.CreateTestContext(r)
|
|
||||||
cache.Set("share_1_1", true, 0)
|
|
||||||
asserts.True(share.WasDownloadedBy(&user, c))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestShare_DownloadBy(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
share := Share{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
}
|
|
||||||
user := User{
|
|
||||||
Model: gorm.Model{
|
|
||||||
ID: 1,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
cache.Deletes([]string{"1_1"}, "share_")
|
|
||||||
r := httptest.NewRecorder()
|
|
||||||
c, _ := gin.CreateTestContext(r)
|
|
||||||
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
err := share.DownloadBy(&user, c)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
_, ok := cache.Get("share_1_1")
|
|
||||||
asserts.True(ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestShare_Viewed(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
share := Share{}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
share.Viewed()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.EqualValues(1, share.Views)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestShare_UpdateAndDelete(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
share := Share{}
|
|
||||||
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := share.Update(map[string]interface{}{"id": 1})
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := share.Delete()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := DeleteShareBySourceIDs([]uint{1}, true)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestListShares(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2).AddRow(2))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1).AddRow(2))
|
|
||||||
|
|
||||||
res, total := ListShares(1, 1, 10, "desc", true)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Len(res, 2)
|
|
||||||
asserts.Equal(2, total)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSearchShares(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs("", sqlmock.AnyArg(), "%1%2%").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
res, total := SearchShares(1, 10, "id", "1 2")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Len(res, 1)
|
|
||||||
asserts.Equal(1, total)
|
|
||||||
}
|
|
||||||
@ -1,52 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestSourceLink_Link(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
s := &SourceLink{}
|
|
||||||
s.ID = 1
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
s.File.Name = string([]byte{0x7f})
|
|
||||||
res, err := s.Link()
|
|
||||||
a.Error(err)
|
|
||||||
a.Empty(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
s.File.Name = "filename"
|
|
||||||
res, err := s.Link()
|
|
||||||
a.NoError(err)
|
|
||||||
a.Contains(res, s.Name)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetSourceLinkByID(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)source_links(.+)").WithArgs(1).WillReturnRows(sqlmock.NewRows([]string{"id", "file_id"}).AddRow(1, 2))
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").WithArgs(2).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2))
|
|
||||||
|
|
||||||
res, err := GetSourceLinkByID(1)
|
|
||||||
a.NoError(err)
|
|
||||||
a.NotNil(res)
|
|
||||||
a.EqualValues(2, res.File.ID)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSourceLink_Downloaded(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
s := &SourceLink{}
|
|
||||||
s.ID = 1
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)source_links(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
s.Downloaded()
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
@ -1,63 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestTag_Create(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
tag := Tag{}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
id, err := tag.Create()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.EqualValues(1, id)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
id, err := tag.Create()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.EqualValues(0, id)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDeleteTagByID(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
err := DeleteTagByID(1, 2)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetTagsByUID(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
res, err := GetTagsByUID(1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(res, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetTagsByID(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"name"}).AddRow("tag"))
|
|
||||||
res, err := GetTagsByID(1, 1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.EqualValues("tag", res.Name)
|
|
||||||
}
|
|
||||||
@ -1,104 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestTask_Create(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
task := Task{Props: "1"}
|
|
||||||
id, err := task.Create()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.EqualValues(1, id)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
task := Task{Props: "1"}
|
|
||||||
id, err := task.Create()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.EqualValues(0, id)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTask_SetError(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
task := Task{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
asserts.NoError(task.SetError("error"))
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTask_SetStatus(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
task := Task{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
asserts.NoError(task.SetStatus(1))
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTask_SetProgress(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
task := Task{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
asserts.NoError(task.SetProgress(1))
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetTasksByID(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
res, err := GetTasksByID(1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.EqualValues(1, res.ID)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestListTasks(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(5))
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(5))
|
|
||||||
|
|
||||||
res, total := ListTasks(1, 1, 10, "")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.EqualValues(5, total)
|
|
||||||
asserts.Len(res, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetTasksByStatus(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").
|
|
||||||
WithArgs(1, 2).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
res := GetTasksByStatus(1, 2)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.Len(res, 1)
|
|
||||||
}
|
|
||||||
@ -1,100 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/duo-labs/webauthn/webauthn"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestUser_RegisterAuthn(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
credential := webauthn.Credential{}
|
|
||||||
user := User{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
}
|
|
||||||
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
user.RegisterAuthn(&credential)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUser_WebAuthnCredentials(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
user := User{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
Authn: `[{"ID":"123","PublicKey":"+4sg1vYcjg/+=","AttestationType":"packed","Authenticator":{"AAGUID":"+lg==","SignCount":0,"CloneWarning":false}}]`,
|
|
||||||
}
|
|
||||||
{
|
|
||||||
credentials := user.WebAuthnCredentials()
|
|
||||||
asserts.Len(credentials, 1)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUser_WebAuthnDisplayName(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
user := User{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
Nick: "123",
|
|
||||||
}
|
|
||||||
{
|
|
||||||
nick := user.WebAuthnDisplayName()
|
|
||||||
asserts.Equal("123", nick)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUser_WebAuthnIcon(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
user := User{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
}
|
|
||||||
{
|
|
||||||
icon := user.WebAuthnIcon()
|
|
||||||
asserts.NotEmpty(icon)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUser_WebAuthnID(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
user := User{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
}
|
|
||||||
{
|
|
||||||
id := user.WebAuthnID()
|
|
||||||
asserts.Len(id, 8)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUser_WebAuthnName(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
user := User{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
Email: "abslant@foxmail.com",
|
|
||||||
}
|
|
||||||
{
|
|
||||||
name := user.WebAuthnName()
|
|
||||||
asserts.Equal("abslant@foxmail.com", name)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestUser_RemoveAuthn(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
user := User{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
Authn: `[{"ID":"123","PublicKey":"+4sg1vYcjg/+=","AttestationType":"packed","Authenticator":{"AAGUID":"+lg==","SignCount":0,"CloneWarning":false}}]`,
|
|
||||||
}
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").
|
|
||||||
WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
user.RemoveAuthn("123")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,60 +0,0 @@
|
|||||||
package model
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestWebdav_Create(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
task := Webdav{}
|
|
||||||
id, err := task.Create()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.EqualValues(1, id)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
task := Webdav{}
|
|
||||||
id, err := task.Create()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.EqualValues(0, id)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetWebdavByPassword(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
_, err := GetWebdavByPassword("e", 1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestListWebDAVAccounts(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}))
|
|
||||||
res := ListWebDAVAccounts(1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Len(res, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDeleteWebDAVAccountByID(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
DeleteWebDAVAccountByID(1, 1)
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
@ -1,66 +0,0 @@
|
|||||||
package aria2
|
|
||||||
|
|
||||||
import (
|
|
||||||
"database/sql"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mocks"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mq"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
testMock "github.com/stretchr/testify/mock"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
)
|
|
||||||
|
|
||||||
var mock sqlmock.Sqlmock
|
|
||||||
|
|
||||||
// TestMain 初始化数据库Mock
|
|
||||||
func TestMain(m *testing.M) {
|
|
||||||
var db *sql.DB
|
|
||||||
var err error
|
|
||||||
db, mock, err = sqlmock.New()
|
|
||||||
if err != nil {
|
|
||||||
panic("An error was not expected when opening a stub database connection")
|
|
||||||
}
|
|
||||||
model.DB, _ = gorm.Open("mysql", db)
|
|
||||||
defer db.Close()
|
|
||||||
m.Run()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestInit(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockPool := &mocks.NodePoolMock{}
|
|
||||||
mockPool.On("GetNodeByID", testMock.Anything).Return(nil)
|
|
||||||
mockQueue := mq.NewMQ()
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
|
|
||||||
Init(false, mockPool, mockQueue)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
mockPool.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTestRPCConnection(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
// url not legal
|
|
||||||
{
|
|
||||||
res, err := TestRPCConnection(string([]byte{0x7f}), "", 10)
|
|
||||||
a.Error(err)
|
|
||||||
a.Empty(res.Version)
|
|
||||||
}
|
|
||||||
|
|
||||||
// rpc failed
|
|
||||||
{
|
|
||||||
res, err := TestRPCConnection("ws://0.0.0.0", "", 0)
|
|
||||||
a.Error(err)
|
|
||||||
a.Empty(res.Version)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetLoadBalancer(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
a.NotPanics(func() {
|
|
||||||
GetLoadBalancer()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@ -1,54 +0,0 @@
|
|||||||
package common
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/aria2/rpc"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestDummyAria2(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
d := &DummyAria2{}
|
|
||||||
|
|
||||||
a.NoError(d.Init())
|
|
||||||
|
|
||||||
res, err := d.CreateTask(&model.Download{}, map[string]interface{}{})
|
|
||||||
a.Empty(res)
|
|
||||||
a.Error(err)
|
|
||||||
|
|
||||||
_, err = d.Status(&model.Download{})
|
|
||||||
a.Error(err)
|
|
||||||
|
|
||||||
err = d.Cancel(&model.Download{})
|
|
||||||
a.Error(err)
|
|
||||||
|
|
||||||
err = d.Select(&model.Download{}, []int{})
|
|
||||||
a.Error(err)
|
|
||||||
|
|
||||||
configRes := d.GetConfig()
|
|
||||||
a.NotNil(configRes)
|
|
||||||
|
|
||||||
err = d.DeleteTempFile(&model.Download{})
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetStatus(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
a.Equal(GetStatus(rpc.StatusInfo{Status: "complete"}), Complete)
|
|
||||||
a.Equal(GetStatus(rpc.StatusInfo{Status: "active",
|
|
||||||
BitTorrent: rpc.BitTorrentInfo{Mode: ""}}), Downloading)
|
|
||||||
a.Equal(GetStatus(rpc.StatusInfo{Status: "active",
|
|
||||||
BitTorrent: rpc.BitTorrentInfo{Mode: "single"},
|
|
||||||
TotalLength: "100", CompletedLength: "50"}), Downloading)
|
|
||||||
a.Equal(GetStatus(rpc.StatusInfo{Status: "active",
|
|
||||||
BitTorrent: rpc.BitTorrentInfo{Mode: "multi"},
|
|
||||||
TotalLength: "100", CompletedLength: "100"}), Seeding)
|
|
||||||
a.Equal(GetStatus(rpc.StatusInfo{Status: "waiting"}), Ready)
|
|
||||||
a.Equal(GetStatus(rpc.StatusInfo{Status: "paused"}), Paused)
|
|
||||||
a.Equal(GetStatus(rpc.StatusInfo{Status: "error"}), Error)
|
|
||||||
a.Equal(GetStatus(rpc.StatusInfo{Status: "removed"}), Canceled)
|
|
||||||
a.Equal(GetStatus(rpc.StatusInfo{Status: "unknown"}), Unknown)
|
|
||||||
}
|
|
||||||
@ -1,447 +0,0 @@
|
|||||||
package monitor
|
|
||||||
|
|
||||||
import (
|
|
||||||
"database/sql"
|
|
||||||
"errors"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/aria2/common"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/aria2/rpc"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/filesystem"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mocks"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mq"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
testMock "github.com/stretchr/testify/mock"
|
|
||||||
)
|
|
||||||
|
|
||||||
var mock sqlmock.Sqlmock
|
|
||||||
|
|
||||||
// TestMain 初始化数据库Mock
|
|
||||||
func TestMain(m *testing.M) {
|
|
||||||
var db *sql.DB
|
|
||||||
var err error
|
|
||||||
db, mock, err = sqlmock.New()
|
|
||||||
if err != nil {
|
|
||||||
panic("An error was not expected when opening a stub database connection")
|
|
||||||
}
|
|
||||||
model.DB, _ = gorm.Open("mysql", db)
|
|
||||||
defer db.Close()
|
|
||||||
m.Run()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewMonitor(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockMQ := mq.NewMQ()
|
|
||||||
|
|
||||||
// node not available
|
|
||||||
{
|
|
||||||
mockPool := &mocks.NodePoolMock{}
|
|
||||||
mockPool.On("GetNodeByID", uint(1)).Return(nil)
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
task := &model.Download{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
}
|
|
||||||
NewMonitor(task, mockPool, mockMQ)
|
|
||||||
mockPool.AssertExpectations(t)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.NotEmpty(task.Error)
|
|
||||||
}
|
|
||||||
|
|
||||||
// success
|
|
||||||
{
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(&common.DummyAria2{})
|
|
||||||
mockPool := &mocks.NodePoolMock{}
|
|
||||||
mockPool.On("GetNodeByID", uint(1)).Return(mockNode)
|
|
||||||
|
|
||||||
task := &model.Download{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
}
|
|
||||||
NewMonitor(task, mockPool, mockMQ)
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
mockPool.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_Loop(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockMQ := mq.NewMQ()
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(&common.DummyAria2{})
|
|
||||||
m := &Monitor{
|
|
||||||
retried: MAX_RETRY,
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{Model: gorm.Model{ID: 1}},
|
|
||||||
notifier: mockMQ.Subscribe("test", 1),
|
|
||||||
}
|
|
||||||
|
|
||||||
// into interval loop
|
|
||||||
{
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
m.Loop(mockMQ)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.NotEmpty(m.Task.Error)
|
|
||||||
}
|
|
||||||
|
|
||||||
// into notifier loop
|
|
||||||
{
|
|
||||||
m.Task.Error = ""
|
|
||||||
mockMQ.Publish("test", mq.Message{})
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
m.Loop(mockMQ)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.NotEmpty(m.Task.Error)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_UpdateFailedAfterRetry(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(&common.DummyAria2{})
|
|
||||||
m := &Monitor{
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
for i := 0; i < MAX_RETRY; i++ {
|
|
||||||
a.False(m.Update())
|
|
||||||
}
|
|
||||||
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
a.True(m.Update())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.NotEmpty(m.Task.Error)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_UpdateMagentoFollow(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockAria2 := &mocks.Aria2Mock{}
|
|
||||||
mockAria2.On("Status", testMock.Anything).Return(rpc.StatusInfo{
|
|
||||||
FollowedBy: []string{"next"},
|
|
||||||
}, nil)
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(mockAria2)
|
|
||||||
m := &Monitor{
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.False(m.Update())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
a.Equal("next", m.Task.GID)
|
|
||||||
mockAria2.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_UpdateFailedToUpdateInfo(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockAria2 := &mocks.Aria2Mock{}
|
|
||||||
mockAria2.On("Status", testMock.Anything).Return(rpc.StatusInfo{}, nil)
|
|
||||||
mockAria2.On("DeleteTempFile", testMock.Anything).Return(nil)
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(mockAria2)
|
|
||||||
m := &Monitor{
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectRollback()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.True(m.Update())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
mockAria2.AssertExpectations(t)
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
a.NotEmpty(m.Task.Error)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_UpdateCompleted(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockAria2 := &mocks.Aria2Mock{}
|
|
||||||
mockAria2.On("Status", testMock.Anything).Return(rpc.StatusInfo{
|
|
||||||
Status: "complete",
|
|
||||||
}, nil)
|
|
||||||
mockAria2.On("DeleteTempFile", testMock.Anything).Return(nil)
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(mockAria2)
|
|
||||||
mockNode.On("ID").Return(uint(1))
|
|
||||||
m := &Monitor{
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectQuery("SELECT(.+)users(.+)").WillReturnError(errors.New("error"))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.True(m.Update())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
mockAria2.AssertExpectations(t)
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
a.NotEmpty(m.Task.Error)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_UpdateError(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockAria2 := &mocks.Aria2Mock{}
|
|
||||||
mockAria2.On("Status", testMock.Anything).Return(rpc.StatusInfo{
|
|
||||||
Status: "error",
|
|
||||||
ErrorMessage: "error",
|
|
||||||
}, nil)
|
|
||||||
mockAria2.On("DeleteTempFile", testMock.Anything).Return(nil)
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(mockAria2)
|
|
||||||
m := &Monitor{
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.True(m.Update())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
mockAria2.AssertExpectations(t)
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
a.NotEmpty(m.Task.Error)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_UpdateActive(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockAria2 := &mocks.Aria2Mock{}
|
|
||||||
mockAria2.On("Status", testMock.Anything).Return(rpc.StatusInfo{
|
|
||||||
Status: "active",
|
|
||||||
}, nil)
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(mockAria2)
|
|
||||||
m := &Monitor{
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.False(m.Update())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
mockAria2.AssertExpectations(t)
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_UpdateRemoved(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockAria2 := &mocks.Aria2Mock{}
|
|
||||||
mockAria2.On("Status", testMock.Anything).Return(rpc.StatusInfo{
|
|
||||||
Status: "removed",
|
|
||||||
}, nil)
|
|
||||||
mockAria2.On("DeleteTempFile", testMock.Anything).Return(nil)
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(mockAria2)
|
|
||||||
m := &Monitor{
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.True(m.Update())
|
|
||||||
a.Equal(common.Canceled, m.Task.Status)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
mockAria2.AssertExpectations(t)
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_UpdateUnknown(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockAria2 := &mocks.Aria2Mock{}
|
|
||||||
mockAria2.On("Status", testMock.Anything).Return(rpc.StatusInfo{
|
|
||||||
Status: "unknown",
|
|
||||||
}, nil)
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(mockAria2)
|
|
||||||
m := &Monitor{
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.True(m.Update())
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
mockAria2.AssertExpectations(t)
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_UpdateTaskInfoValidateFailed(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
status := rpc.StatusInfo{
|
|
||||||
Status: "completed",
|
|
||||||
TotalLength: "100",
|
|
||||||
CompletedLength: "50",
|
|
||||||
DownloadSpeed: "20",
|
|
||||||
}
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(&common.DummyAria2{})
|
|
||||||
m := &Monitor{
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
err := m.UpdateTaskInfo(status)
|
|
||||||
a.Error(err)
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_ValidateFile(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &Monitor{
|
|
||||||
Task: &model.Download{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
TotalSize: 100,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// failed to create filesystem
|
|
||||||
{
|
|
||||||
m.Task.User = &model.User{
|
|
||||||
Policy: model.Policy{
|
|
||||||
Type: "random",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
a.Equal(filesystem.ErrUnknownPolicyType, m.ValidateFile())
|
|
||||||
}
|
|
||||||
|
|
||||||
// User capacity not enough
|
|
||||||
{
|
|
||||||
m.Task.User = &model.User{
|
|
||||||
Group: model.Group{
|
|
||||||
MaxStorage: 99,
|
|
||||||
},
|
|
||||||
Policy: model.Policy{
|
|
||||||
Type: "local",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
a.Equal(filesystem.ErrInsufficientCapacity, m.ValidateFile())
|
|
||||||
}
|
|
||||||
|
|
||||||
// single file too big
|
|
||||||
{
|
|
||||||
m.Task.StatusInfo.Files = []rpc.FileInfo{
|
|
||||||
{
|
|
||||||
Length: "100",
|
|
||||||
Selected: "true",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
m.Task.User = &model.User{
|
|
||||||
Group: model.Group{
|
|
||||||
MaxStorage: 100,
|
|
||||||
},
|
|
||||||
Policy: model.Policy{
|
|
||||||
Type: "local",
|
|
||||||
MaxSize: 99,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
a.Equal(filesystem.ErrFileSizeTooBig, m.ValidateFile())
|
|
||||||
}
|
|
||||||
|
|
||||||
// all pass
|
|
||||||
{
|
|
||||||
m.Task.StatusInfo.Files = []rpc.FileInfo{
|
|
||||||
{
|
|
||||||
Length: "100",
|
|
||||||
Selected: "true",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
m.Task.User = &model.User{
|
|
||||||
Group: model.Group{
|
|
||||||
MaxStorage: 100,
|
|
||||||
},
|
|
||||||
Policy: model.Policy{
|
|
||||||
Type: "local",
|
|
||||||
MaxSize: 100,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
a.NoError(m.ValidateFile())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMonitor_Complete(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockNode := &mocks.NodeMock{}
|
|
||||||
mockNode.On("ID").Return(uint(1))
|
|
||||||
mockPool := &mocks.TaskPoolMock{}
|
|
||||||
mockPool.On("Submit", testMock.Anything)
|
|
||||||
m := &Monitor{
|
|
||||||
node: mockNode,
|
|
||||||
Task: &model.Download{
|
|
||||||
Model: gorm.Model{ID: 1},
|
|
||||||
TotalSize: 100,
|
|
||||||
UserID: 9414,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
m.Task.StatusInfo.Files = []rpc.FileInfo{
|
|
||||||
{
|
|
||||||
Length: "100",
|
|
||||||
Selected: "true",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)users").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(9414))
|
|
||||||
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)tasks").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("UPDATE(.+)downloads").WillReturnResult(sqlmock.NewResult(1, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
mock.ExpectQuery("SELECT(.+)tasks").WillReturnRows(sqlmock.NewRows([]string{"id", "type", "status"}).AddRow(1, 2, 4))
|
|
||||||
mock.ExpectQuery("SELECT(.+)users").WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(9414))
|
|
||||||
mock.ExpectBegin()
|
|
||||||
mock.ExpectExec("INSERT(.+)tasks").WillReturnResult(sqlmock.NewResult(2, 1))
|
|
||||||
mock.ExpectCommit()
|
|
||||||
|
|
||||||
a.False(m.Complete(mockPool))
|
|
||||||
m.Task.StatusInfo.Status = "complete"
|
|
||||||
a.True(m.Complete(mockPool))
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
mockPool.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
@ -1,136 +0,0 @@
|
|||||||
package auth
|
|
||||||
|
|
||||||
import (
|
|
||||||
"io/ioutil"
|
|
||||||
"net/http"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestSignURI(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
General = HMACAuth{SecretKey: []byte(util.RandStringRunes(256))}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
sign, err := SignURI(General, "/api/v3/something?id=1", 0)
|
|
||||||
asserts.NoError(err)
|
|
||||||
queries := sign.Query()
|
|
||||||
asserts.Equal("1", queries.Get("id"))
|
|
||||||
asserts.NotEmpty(queries.Get("sign"))
|
|
||||||
}
|
|
||||||
|
|
||||||
// URI解码失败
|
|
||||||
{
|
|
||||||
sign, err := SignURI(General, "://dg.;'f]gh./'", 0)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Nil(sign)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCheckURI(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
General = HMACAuth{SecretKey: []byte(util.RandStringRunes(256))}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
sign, err := SignURI(General, "/api/ok?if=sdf&fd=go", 10)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.NoError(CheckURI(General, sign))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 过期
|
|
||||||
{
|
|
||||||
sign, err := SignURI(General, "/api/ok?if=sdf&fd=go", -1)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Error(CheckURI(General, sign))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSignRequest(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
General = HMACAuth{SecretKey: []byte(util.RandStringRunes(256))}
|
|
||||||
|
|
||||||
// 非上传请求
|
|
||||||
{
|
|
||||||
req, err := http.NewRequest("POST", "http://127.0.0.1/api/v3/slave/upload", strings.NewReader("I am body."))
|
|
||||||
asserts.NoError(err)
|
|
||||||
req = SignRequest(General, req, 0)
|
|
||||||
asserts.NotEmpty(req.Header["Authorization"])
|
|
||||||
}
|
|
||||||
|
|
||||||
// 上传请求
|
|
||||||
{
|
|
||||||
req, err := http.NewRequest(
|
|
||||||
"POST",
|
|
||||||
"http://127.0.0.1/api/v3/slave/upload",
|
|
||||||
strings.NewReader("I am body."),
|
|
||||||
)
|
|
||||||
asserts.NoError(err)
|
|
||||||
req.Header["X-Cr-Policy"] = []string{"I am Policy"}
|
|
||||||
req = SignRequest(General, req, 10)
|
|
||||||
asserts.NotEmpty(req.Header["Authorization"])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestCheckRequest(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
General = HMACAuth{SecretKey: []byte(util.RandStringRunes(256))}
|
|
||||||
|
|
||||||
// 缺少请求头
|
|
||||||
{
|
|
||||||
req, err := http.NewRequest(
|
|
||||||
"POST",
|
|
||||||
"http://127.0.0.1/api/v3/upload",
|
|
||||||
strings.NewReader("I am body."),
|
|
||||||
)
|
|
||||||
asserts.NoError(err)
|
|
||||||
err = CheckRequest(General, req)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Equal(ErrAuthHeaderMissing, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 非上传请求 验证成功
|
|
||||||
{
|
|
||||||
req, err := http.NewRequest(
|
|
||||||
"POST",
|
|
||||||
"http://127.0.0.1/api/v3/upload",
|
|
||||||
strings.NewReader("I am body."),
|
|
||||||
)
|
|
||||||
asserts.NoError(err)
|
|
||||||
req = SignRequest(General, req, 0)
|
|
||||||
err = CheckRequest(General, req)
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 上传请求 验证成功
|
|
||||||
{
|
|
||||||
req, err := http.NewRequest(
|
|
||||||
"POST",
|
|
||||||
"http://127.0.0.1/api/v3/upload",
|
|
||||||
strings.NewReader("I am body."),
|
|
||||||
)
|
|
||||||
asserts.NoError(err)
|
|
||||||
req.Header["X-Cr-Policy"] = []string{"I am Policy"}
|
|
||||||
req = SignRequest(General, req, 0)
|
|
||||||
err = CheckRequest(General, req)
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 非上传请求 失败
|
|
||||||
{
|
|
||||||
req, err := http.NewRequest(
|
|
||||||
"POST",
|
|
||||||
"http://127.0.0.1/api/v3/upload",
|
|
||||||
strings.NewReader("I am body."),
|
|
||||||
)
|
|
||||||
asserts.NoError(err)
|
|
||||||
req = SignRequest(General, req, 0)
|
|
||||||
req.Body = ioutil.NopCloser(strings.NewReader("2333"))
|
|
||||||
err = CheckRequest(General, req)
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,94 +0,0 @@
|
|||||||
package auth
|
|
||||||
|
|
||||||
import (
|
|
||||||
"database/sql"
|
|
||||||
"fmt"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/conf"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
|
||||||
"github.com/gin-gonic/gin"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
var mock sqlmock.Sqlmock
|
|
||||||
|
|
||||||
func TestMain(m *testing.M) {
|
|
||||||
// 设置gin为测试模式
|
|
||||||
gin.SetMode(gin.TestMode)
|
|
||||||
|
|
||||||
// 初始化sqlmock
|
|
||||||
var db *sql.DB
|
|
||||||
var err error
|
|
||||||
db, mock, err = sqlmock.New()
|
|
||||||
if err != nil {
|
|
||||||
panic("An error was not expected when opening a stub database connection")
|
|
||||||
}
|
|
||||||
|
|
||||||
mockDB, _ := gorm.Open("mysql", db)
|
|
||||||
model.DB = mockDB
|
|
||||||
defer db.Close()
|
|
||||||
|
|
||||||
m.Run()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHMACAuth_Sign(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
auth := HMACAuth{
|
|
||||||
SecretKey: []byte(util.RandStringRunes(256)),
|
|
||||||
}
|
|
||||||
|
|
||||||
asserts.NotEmpty(auth.Sign("content", 0))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHMACAuth_Check(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
auth := HMACAuth{
|
|
||||||
SecretKey: []byte(util.RandStringRunes(256)),
|
|
||||||
}
|
|
||||||
|
|
||||||
// 正常,永不过期
|
|
||||||
{
|
|
||||||
sign := auth.Sign("content", 0)
|
|
||||||
asserts.NoError(auth.Check("content", sign))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 过期
|
|
||||||
{
|
|
||||||
sign := auth.Sign("content", 1)
|
|
||||||
asserts.Error(auth.Check("content", sign))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 签名格式错误
|
|
||||||
{
|
|
||||||
sign := auth.Sign("content", 1)
|
|
||||||
asserts.Error(auth.Check("content", sign+":"))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 过期日期格式错误
|
|
||||||
{
|
|
||||||
asserts.Error(auth.Check("content", "ErrAuthFailed:ErrAuthFailed"))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 签名有误
|
|
||||||
{
|
|
||||||
asserts.Error(auth.Check("content", fmt.Sprintf("sign:%d", time.Now().Unix()+10)))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestInit(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id", "value"}).AddRow(1, "12312312312312"))
|
|
||||||
Init()
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
|
|
||||||
// slave模式
|
|
||||||
conf.SystemConfig.Mode = "slave"
|
|
||||||
asserts.Panics(func() {
|
|
||||||
Init()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@ -1,17 +0,0 @@
|
|||||||
package authn
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestInit(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
cache.Set("setting_siteURL", "http://cloudreve.org", 0)
|
|
||||||
cache.Set("setting_siteName", "Cloudreve", 0)
|
|
||||||
res, err := NewAuthnInstance()
|
|
||||||
asserts.NotNil(res)
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
@ -1,12 +0,0 @@
|
|||||||
package balancer
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestNewBalancer(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
a.NotNil(NewBalancer(""))
|
|
||||||
a.IsType(&RoundRobin{}, NewBalancer("RoundRobin"))
|
|
||||||
}
|
|
||||||
@ -1,42 +0,0 @@
|
|||||||
package balancer
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestRoundRobin_NextIndex(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
r := &RoundRobin{}
|
|
||||||
total := 5
|
|
||||||
for i := 1; i < total; i++ {
|
|
||||||
a.Equal(i, r.NextIndex(total))
|
|
||||||
}
|
|
||||||
for i := 0; i < total; i++ {
|
|
||||||
a.Equal(i, r.NextIndex(total))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRoundRobin_NextPeer(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
r := &RoundRobin{}
|
|
||||||
|
|
||||||
// not slice
|
|
||||||
{
|
|
||||||
err, _ := r.NextPeer("s")
|
|
||||||
a.Equal(ErrInputNotSlice, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// no nodes
|
|
||||||
{
|
|
||||||
err, _ := r.NextPeer([]string{})
|
|
||||||
a.Equal(ErrNoAvaliableNode, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// pass
|
|
||||||
{
|
|
||||||
err, res := r.NextPeer([]string{"a"})
|
|
||||||
a.NoError(err)
|
|
||||||
a.Equal("a", res.(string))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,69 +0,0 @@
|
|||||||
package cache
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestSet(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
asserts.NoError(Set("123", "321", -1))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGet(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
asserts.NoError(Set("123", "321", -1))
|
|
||||||
|
|
||||||
value, ok := Get("123")
|
|
||||||
asserts.True(ok)
|
|
||||||
asserts.Equal("321", value)
|
|
||||||
|
|
||||||
value, ok = Get("not_exist")
|
|
||||||
asserts.False(ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDeletes(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
asserts.NoError(Set("123", "321", -1))
|
|
||||||
err := Deletes([]string{"123"}, "")
|
|
||||||
asserts.NoError(err)
|
|
||||||
_, exist := Get("123")
|
|
||||||
asserts.False(exist)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestGetSettings(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
asserts.NoError(Set("test_1", "1", -1))
|
|
||||||
|
|
||||||
values, missed := GetSettings([]string{"1", "2"}, "test_")
|
|
||||||
asserts.Equal(map[string]string{"1": "1"}, values)
|
|
||||||
asserts.Equal([]string{"2"}, missed)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSetSettings(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
err := SetSettings(map[string]string{"3": "3", "4": "4"}, "test_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
value1, _ := Get("test_3")
|
|
||||||
value2, _ := Get("test_4")
|
|
||||||
asserts.Equal("3", value1)
|
|
||||||
asserts.Equal("4", value2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestInit(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
Init()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestInitSlaveOverwrites(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
InitSlaveOverwrites()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
@ -1,191 +0,0 @@
|
|||||||
package cache
|
|
||||||
|
|
||||||
import (
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"path/filepath"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestNewMemoStore(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
store := NewMemoStore()
|
|
||||||
asserts.NotNil(store)
|
|
||||||
asserts.NotNil(store.Store)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMemoStore_Set(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
store := NewMemoStore()
|
|
||||||
err := store.Set("KEY", "vAL", -1)
|
|
||||||
asserts.NoError(err)
|
|
||||||
|
|
||||||
val, ok := store.Store.Load("KEY")
|
|
||||||
asserts.True(ok)
|
|
||||||
asserts.Equal("vAL", val.(itemWithTTL).Value)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMemoStore_Get(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
store := NewMemoStore()
|
|
||||||
|
|
||||||
// 正常情况
|
|
||||||
{
|
|
||||||
_ = store.Set("string", "string_val", -1)
|
|
||||||
val, ok := store.Get("string")
|
|
||||||
asserts.Equal("string_val", val)
|
|
||||||
asserts.True(ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Key不存在
|
|
||||||
{
|
|
||||||
val, ok := store.Get("something")
|
|
||||||
asserts.Equal(nil, val)
|
|
||||||
asserts.False(ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 存储struct
|
|
||||||
{
|
|
||||||
type testStruct struct {
|
|
||||||
key int
|
|
||||||
}
|
|
||||||
test := testStruct{key: 233}
|
|
||||||
_ = store.Set("struct", test, -1)
|
|
||||||
val, ok := store.Get("struct")
|
|
||||||
asserts.True(ok)
|
|
||||||
res, ok := val.(testStruct)
|
|
||||||
asserts.True(ok)
|
|
||||||
asserts.Equal(test, res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 过期
|
|
||||||
{
|
|
||||||
_ = store.Set("string", "string_val", 1)
|
|
||||||
time.Sleep(time.Duration(2) * time.Second)
|
|
||||||
val, ok := store.Get("string")
|
|
||||||
asserts.Nil(val)
|
|
||||||
asserts.False(ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMemoStore_Gets(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
store := NewMemoStore()
|
|
||||||
|
|
||||||
err := store.Set("1", "1,val", -1)
|
|
||||||
err = store.Set("2", "2,val", -1)
|
|
||||||
err = store.Set("3", "3,val", -1)
|
|
||||||
err = store.Set("4", "4,val", -1)
|
|
||||||
asserts.NoError(err)
|
|
||||||
|
|
||||||
// 全部命中
|
|
||||||
{
|
|
||||||
values, miss := store.Gets([]string{"1", "2", "3", "4"}, "")
|
|
||||||
asserts.Len(values, 4)
|
|
||||||
asserts.Len(miss, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 命中一半
|
|
||||||
{
|
|
||||||
values, miss := store.Gets([]string{"1", "2", "9", "10"}, "")
|
|
||||||
asserts.Len(values, 2)
|
|
||||||
asserts.Equal([]string{"9", "10"}, miss)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMemoStore_Sets(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
store := NewMemoStore()
|
|
||||||
|
|
||||||
err := store.Sets(map[string]interface{}{
|
|
||||||
"1": "1.val",
|
|
||||||
"2": "2.val",
|
|
||||||
"3": "3.val",
|
|
||||||
"4": "4.val",
|
|
||||||
}, "test_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
|
|
||||||
vals, miss := store.Gets([]string{"1", "2", "3", "4"}, "test_")
|
|
||||||
asserts.Len(miss, 0)
|
|
||||||
asserts.Equal(map[string]interface{}{
|
|
||||||
"1": "1.val",
|
|
||||||
"2": "2.val",
|
|
||||||
"3": "3.val",
|
|
||||||
"4": "4.val",
|
|
||||||
}, vals)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMemoStore_Delete(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
store := NewMemoStore()
|
|
||||||
|
|
||||||
err := store.Sets(map[string]interface{}{
|
|
||||||
"1": "1.val",
|
|
||||||
"2": "2.val",
|
|
||||||
"3": "3.val",
|
|
||||||
"4": "4.val",
|
|
||||||
}, "test_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
|
|
||||||
err = store.Delete([]string{"1", "2"}, "test_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
values, miss := store.Gets([]string{"1", "2", "3", "4"}, "test_")
|
|
||||||
asserts.Equal([]string{"1", "2"}, miss)
|
|
||||||
asserts.Equal(map[string]interface{}{"3": "3.val", "4": "4.val"}, values)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMemoStore_GarbageCollect(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
store := NewMemoStore()
|
|
||||||
store.Set("test", 1, 1)
|
|
||||||
time.Sleep(time.Duration(2000) * time.Millisecond)
|
|
||||||
store.GarbageCollect()
|
|
||||||
_, ok := store.Get("test")
|
|
||||||
asserts.False(ok)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMemoStore_PersistFailed(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
store := NewMemoStore()
|
|
||||||
type testStruct struct{ v string }
|
|
||||||
store.Set("test", 1, 0)
|
|
||||||
store.Set("test2", testStruct{v: "test"}, 0)
|
|
||||||
err := store.Persist(filepath.Join(t.TempDir(), "TestMemoStore_PersistFailed"))
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMemoStore_PersistAndRestore(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
store := NewMemoStore()
|
|
||||||
store.Set("test", 1, 0)
|
|
||||||
// already expired
|
|
||||||
store.Store.Store("test2", itemWithTTL{Value: "test", Expires: 1})
|
|
||||||
// expired after persist
|
|
||||||
store.Set("test3", 1, 1)
|
|
||||||
temp := filepath.Join(t.TempDir(), "TestMemoStore_PersistFailed")
|
|
||||||
|
|
||||||
// Persist
|
|
||||||
err := store.Persist(temp)
|
|
||||||
a.NoError(err)
|
|
||||||
a.FileExists(temp)
|
|
||||||
|
|
||||||
time.Sleep(2 * time.Second)
|
|
||||||
// Restore
|
|
||||||
store2 := NewMemoStore()
|
|
||||||
err = store2.Restore(temp)
|
|
||||||
a.NoError(err)
|
|
||||||
test, testOk := store2.Get("test")
|
|
||||||
a.EqualValues(1, test)
|
|
||||||
a.True(testOk)
|
|
||||||
test2, test2Ok := store2.Get("test2")
|
|
||||||
a.Nil(test2)
|
|
||||||
a.False(test2Ok)
|
|
||||||
test3, test3Ok := store2.Get("test3")
|
|
||||||
a.Nil(test3)
|
|
||||||
a.False(test3Ok)
|
|
||||||
|
|
||||||
a.NoFileExists(temp)
|
|
||||||
}
|
|
||||||
@ -1,324 +0,0 @@
|
|||||||
package cache
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"github.com/gomodule/redigo/redis"
|
|
||||||
"github.com/rafaeljusto/redigomock"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestNewRedisStore(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
store := NewRedisStore(10, "tcp", "", "", "", "0")
|
|
||||||
asserts.NotNil(store)
|
|
||||||
|
|
||||||
asserts.Panics(func() {
|
|
||||||
store.pool.Dial()
|
|
||||||
})
|
|
||||||
|
|
||||||
testConn := redigomock.NewConn()
|
|
||||||
cmd := testConn.Command("PING").Expect("PONG")
|
|
||||||
err := store.pool.TestOnBorrow(testConn, time.Now())
|
|
||||||
if testConn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRedisStore_Set(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
conn := redigomock.NewConn()
|
|
||||||
pool := &redis.Pool{
|
|
||||||
Dial: func() (redis.Conn, error) { return conn, nil },
|
|
||||||
MaxIdle: 10,
|
|
||||||
}
|
|
||||||
store := &RedisStore{pool: pool}
|
|
||||||
|
|
||||||
// 正常情况
|
|
||||||
{
|
|
||||||
cmd := conn.Command("SET", "test", redigomock.NewAnyData()).ExpectStringSlice("OK")
|
|
||||||
err := store.Set("test", "test val", -1)
|
|
||||||
asserts.NoError(err)
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 带有TTL
|
|
||||||
// 正常情况
|
|
||||||
{
|
|
||||||
cmd := conn.Command("SETEX", "test", 10, redigomock.NewAnyData()).ExpectStringSlice("OK")
|
|
||||||
err := store.Set("test", "test val", 10)
|
|
||||||
asserts.NoError(err)
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 序列化出错
|
|
||||||
{
|
|
||||||
value := struct {
|
|
||||||
Key string
|
|
||||||
}{
|
|
||||||
Key: "123",
|
|
||||||
}
|
|
||||||
err := store.Set("test", value, -1)
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 命令执行失败
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
cmd := conn.Command("SET", "test", redigomock.NewAnyData()).ExpectError(errors.New("error"))
|
|
||||||
err := store.Set("test", "test val", -1)
|
|
||||||
asserts.Error(err)
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
// 获取连接失败
|
|
||||||
{
|
|
||||||
store.pool = &redis.Pool{
|
|
||||||
Dial: func() (redis.Conn, error) { return nil, errors.New("error") },
|
|
||||||
MaxIdle: 10,
|
|
||||||
}
|
|
||||||
err := store.Set("test", "123", -1)
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRedisStore_Get(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
conn := redigomock.NewConn()
|
|
||||||
pool := &redis.Pool{
|
|
||||||
Dial: func() (redis.Conn, error) { return conn, nil },
|
|
||||||
MaxIdle: 10,
|
|
||||||
}
|
|
||||||
store := &RedisStore{pool: pool}
|
|
||||||
|
|
||||||
// 正常情况
|
|
||||||
{
|
|
||||||
expectVal, _ := serializer("test val")
|
|
||||||
cmd := conn.Command("GET", "test").Expect(expectVal)
|
|
||||||
val, ok := store.Get("test")
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
asserts.True(ok)
|
|
||||||
asserts.Equal("test val", val.(string))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Key不存在
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
cmd := conn.Command("GET", "test").Expect(nil)
|
|
||||||
val, ok := store.Get("test")
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
asserts.False(ok)
|
|
||||||
asserts.Nil(val)
|
|
||||||
}
|
|
||||||
// 解码错误
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
cmd := conn.Command("GET", "test").Expect([]byte{0x20})
|
|
||||||
val, ok := store.Get("test")
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
asserts.False(ok)
|
|
||||||
asserts.Nil(val)
|
|
||||||
}
|
|
||||||
// 获取连接失败
|
|
||||||
{
|
|
||||||
store.pool = &redis.Pool{
|
|
||||||
Dial: func() (redis.Conn, error) { return nil, errors.New("error") },
|
|
||||||
MaxIdle: 10,
|
|
||||||
}
|
|
||||||
val, ok := store.Get("test")
|
|
||||||
asserts.False(ok)
|
|
||||||
asserts.Nil(val)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRedisStore_Gets(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
conn := redigomock.NewConn()
|
|
||||||
pool := &redis.Pool{
|
|
||||||
Dial: func() (redis.Conn, error) { return conn, nil },
|
|
||||||
MaxIdle: 10,
|
|
||||||
}
|
|
||||||
store := &RedisStore{pool: pool}
|
|
||||||
|
|
||||||
// 全部命中
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
value1, _ := serializer("1")
|
|
||||||
value2, _ := serializer("2")
|
|
||||||
cmd := conn.Command("MGET", "test_1", "test_2").ExpectSlice(
|
|
||||||
value1, value2)
|
|
||||||
res, missed := store.Gets([]string{"1", "2"}, "test_")
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
asserts.Len(missed, 0)
|
|
||||||
asserts.Len(res, 2)
|
|
||||||
asserts.Equal("1", res["1"].(string))
|
|
||||||
asserts.Equal("2", res["2"].(string))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 命中一个
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
value2, _ := serializer("2")
|
|
||||||
cmd := conn.Command("MGET", "test_1", "test_2").ExpectSlice(
|
|
||||||
nil, value2)
|
|
||||||
res, missed := store.Gets([]string{"1", "2"}, "test_")
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
asserts.Len(missed, 1)
|
|
||||||
asserts.Len(res, 1)
|
|
||||||
asserts.Equal("1", missed[0])
|
|
||||||
asserts.Equal("2", res["2"].(string))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 命令出错
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
cmd := conn.Command("MGET", "test_1", "test_2").ExpectError(errors.New("error"))
|
|
||||||
res, missed := store.Gets([]string{"1", "2"}, "test_")
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
asserts.Len(missed, 2)
|
|
||||||
asserts.Len(res, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 连接出错
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
store.pool = &redis.Pool{
|
|
||||||
Dial: func() (redis.Conn, error) { return nil, errors.New("error") },
|
|
||||||
MaxIdle: 10,
|
|
||||||
}
|
|
||||||
res, missed := store.Gets([]string{"1", "2"}, "test_")
|
|
||||||
asserts.Len(missed, 2)
|
|
||||||
asserts.Len(res, 0)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRedisStore_Sets(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
conn := redigomock.NewConn()
|
|
||||||
pool := &redis.Pool{
|
|
||||||
Dial: func() (redis.Conn, error) { return conn, nil },
|
|
||||||
MaxIdle: 10,
|
|
||||||
}
|
|
||||||
store := &RedisStore{pool: pool}
|
|
||||||
|
|
||||||
// 正常
|
|
||||||
{
|
|
||||||
cmd := conn.Command("MSET", redigomock.NewAnyData(), redigomock.NewAnyData(), redigomock.NewAnyData(), redigomock.NewAnyData()).ExpectSlice("OK")
|
|
||||||
err := store.Sets(map[string]interface{}{"1": "1", "2": "2"}, "test_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 序列化失败
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
value := struct {
|
|
||||||
Key string
|
|
||||||
}{
|
|
||||||
Key: "123",
|
|
||||||
}
|
|
||||||
err := store.Sets(map[string]interface{}{"1": value, "2": "2"}, "test_")
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 执行失败
|
|
||||||
{
|
|
||||||
cmd := conn.Command("MSET", redigomock.NewAnyData(), redigomock.NewAnyData(), redigomock.NewAnyData(), redigomock.NewAnyData()).ExpectError(errors.New("error"))
|
|
||||||
err := store.Sets(map[string]interface{}{"1": "1", "2": "2"}, "test_")
|
|
||||||
asserts.Error(err)
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 连接失败
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
store.pool = &redis.Pool{
|
|
||||||
Dial: func() (redis.Conn, error) { return nil, errors.New("error") },
|
|
||||||
MaxIdle: 10,
|
|
||||||
}
|
|
||||||
err := store.Sets(map[string]interface{}{"1": "1", "2": "2"}, "test_")
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRedisStore_Delete(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
conn := redigomock.NewConn()
|
|
||||||
pool := &redis.Pool{
|
|
||||||
Dial: func() (redis.Conn, error) { return conn, nil },
|
|
||||||
MaxIdle: 10,
|
|
||||||
}
|
|
||||||
store := &RedisStore{pool: pool}
|
|
||||||
|
|
||||||
// 正常
|
|
||||||
{
|
|
||||||
cmd := conn.Command("DEL", redigomock.NewAnyData(), redigomock.NewAnyData(), redigomock.NewAnyData(), redigomock.NewAnyData()).ExpectSlice("OK")
|
|
||||||
err := store.Delete([]string{"1", "2", "3", "4"}, "test_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 命令执行失败
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
cmd := conn.Command("DEL", redigomock.NewAnyData(), redigomock.NewAnyData(), redigomock.NewAnyData(), redigomock.NewAnyData()).ExpectError(errors.New("error"))
|
|
||||||
err := store.Delete([]string{"1", "2", "3", "4"}, "test_")
|
|
||||||
asserts.Error(err)
|
|
||||||
if conn.Stats(cmd) != 1 {
|
|
||||||
fmt.Println("Command was not used")
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 连接失败
|
|
||||||
{
|
|
||||||
conn.Clear()
|
|
||||||
store.pool = &redis.Pool{
|
|
||||||
Dial: func() (redis.Conn, error) { return nil, errors.New("error") },
|
|
||||||
MaxIdle: 10,
|
|
||||||
}
|
|
||||||
err := store.Delete([]string{"1", "2", "3", "4"}, "test_")
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,385 +0,0 @@
|
|||||||
package cluster
|
|
||||||
|
|
||||||
import (
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/aria2/common"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/auth"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mq"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/request"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/serializer"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
testMock "github.com/stretchr/testify/mock"
|
|
||||||
"io"
|
|
||||||
"io/ioutil"
|
|
||||||
"net/http"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestInitController(t *testing.T) {
|
|
||||||
assert.NotPanics(t, func() {
|
|
||||||
InitController()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveController_HandleHeartBeat(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
c := &slaveController{
|
|
||||||
masters: make(map[string]MasterInfo),
|
|
||||||
}
|
|
||||||
|
|
||||||
// first heart beat
|
|
||||||
{
|
|
||||||
_, err := c.HandleHeartBeat(&serializer.NodePingReq{
|
|
||||||
SiteID: "1",
|
|
||||||
Node: &model.Node{},
|
|
||||||
})
|
|
||||||
a.NoError(err)
|
|
||||||
|
|
||||||
_, err = c.HandleHeartBeat(&serializer.NodePingReq{
|
|
||||||
SiteID: "2",
|
|
||||||
Node: &model.Node{},
|
|
||||||
})
|
|
||||||
a.NoError(err)
|
|
||||||
|
|
||||||
a.Len(c.masters, 2)
|
|
||||||
}
|
|
||||||
|
|
||||||
// second heart beat, no fresh
|
|
||||||
{
|
|
||||||
_, err := c.HandleHeartBeat(&serializer.NodePingReq{
|
|
||||||
SiteID: "1",
|
|
||||||
SiteURL: "http://127.0.0.1",
|
|
||||||
Node: &model.Node{},
|
|
||||||
})
|
|
||||||
a.NoError(err)
|
|
||||||
a.Len(c.masters, 2)
|
|
||||||
a.Empty(c.masters["1"].URL)
|
|
||||||
}
|
|
||||||
|
|
||||||
// second heart beat, fresh
|
|
||||||
{
|
|
||||||
_, err := c.HandleHeartBeat(&serializer.NodePingReq{
|
|
||||||
SiteID: "1",
|
|
||||||
IsUpdate: true,
|
|
||||||
SiteURL: "http://127.0.0.1",
|
|
||||||
Node: &model.Node{},
|
|
||||||
})
|
|
||||||
a.NoError(err)
|
|
||||||
a.Len(c.masters, 2)
|
|
||||||
a.Equal("http://127.0.0.1", c.masters["1"].URL.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
// second heart beat, fresh, url illegal
|
|
||||||
{
|
|
||||||
_, err := c.HandleHeartBeat(&serializer.NodePingReq{
|
|
||||||
SiteID: "1",
|
|
||||||
IsUpdate: true,
|
|
||||||
SiteURL: string([]byte{0x7f}),
|
|
||||||
Node: &model.Node{},
|
|
||||||
})
|
|
||||||
a.Error(err)
|
|
||||||
a.Len(c.masters, 2)
|
|
||||||
a.Equal("http://127.0.0.1", c.masters["1"].URL.String())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type nodeMock struct {
|
|
||||||
testMock.Mock
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) Init(node *model.Node) {
|
|
||||||
n.Called(node)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) IsFeatureEnabled(feature string) bool {
|
|
||||||
args := n.Called(feature)
|
|
||||||
return args.Bool(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) SubscribeStatusChange(callback func(isActive bool, id uint)) {
|
|
||||||
n.Called(callback)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) Ping(req *serializer.NodePingReq) (*serializer.NodePingResp, error) {
|
|
||||||
args := n.Called(req)
|
|
||||||
return args.Get(0).(*serializer.NodePingResp), args.Error(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) IsActive() bool {
|
|
||||||
args := n.Called()
|
|
||||||
return args.Bool(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) GetAria2Instance() common.Aria2 {
|
|
||||||
args := n.Called()
|
|
||||||
return args.Get(0).(common.Aria2)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) ID() uint {
|
|
||||||
args := n.Called()
|
|
||||||
return args.Get(0).(uint)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) Kill() {
|
|
||||||
n.Called()
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) IsMater() bool {
|
|
||||||
args := n.Called()
|
|
||||||
return args.Bool(0)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) MasterAuthInstance() auth.Auth {
|
|
||||||
args := n.Called()
|
|
||||||
return args.Get(0).(auth.Auth)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) SlaveAuthInstance() auth.Auth {
|
|
||||||
args := n.Called()
|
|
||||||
return args.Get(0).(auth.Auth)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (n nodeMock) DBModel() *model.Node {
|
|
||||||
args := n.Called()
|
|
||||||
return args.Get(0).(*model.Node)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveController_GetAria2Instance(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mockNode := &nodeMock{}
|
|
||||||
mockNode.On("GetAria2Instance").Return(&common.DummyAria2{})
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {Instance: mockNode},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// node node found
|
|
||||||
{
|
|
||||||
res, err := c.GetAria2Instance("2")
|
|
||||||
a.Nil(res)
|
|
||||||
a.Equal(ErrMasterNotFound, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// node found
|
|
||||||
{
|
|
||||||
res, err := c.GetAria2Instance("1")
|
|
||||||
a.NotNil(res)
|
|
||||||
a.NoError(err)
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
type requestMock struct {
|
|
||||||
testMock.Mock
|
|
||||||
}
|
|
||||||
|
|
||||||
func (r requestMock) Request(method, target string, body io.Reader, opts ...request.Option) *request.Response {
|
|
||||||
return r.Called(method, target, body, opts).Get(0).(*request.Response)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveController_SendNotification(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// node not exit
|
|
||||||
{
|
|
||||||
a.Equal(ErrMasterNotFound, c.SendNotification("2", "", mq.Message{}))
|
|
||||||
}
|
|
||||||
|
|
||||||
// gob encode error
|
|
||||||
{
|
|
||||||
type randomType struct{}
|
|
||||||
a.Error(c.SendNotification("1", "", mq.Message{
|
|
||||||
Content: randomType{},
|
|
||||||
}))
|
|
||||||
}
|
|
||||||
|
|
||||||
// return none 200
|
|
||||||
{
|
|
||||||
mockRequest := &requestMock{}
|
|
||||||
mockRequest.On("Request", "PUT", "/api/v3/slave/notification/s1", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{StatusCode: http.StatusConflict},
|
|
||||||
})
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {Client: mockRequest},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
a.Error(c.SendNotification("1", "s1", mq.Message{}))
|
|
||||||
mockRequest.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return error
|
|
||||||
{
|
|
||||||
mockRequest := &requestMock{}
|
|
||||||
mockRequest.On("Request", "PUT", "/api/v3/slave/notification/s2", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"code\":1}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {Client: mockRequest},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
a.Equal(1, c.SendNotification("1", "s2", mq.Message{}).(serializer.AppError).Code)
|
|
||||||
mockRequest.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// success
|
|
||||||
{
|
|
||||||
mockRequest := &requestMock{}
|
|
||||||
mockRequest.On("Request", "PUT", "/api/v3/slave/notification/s3", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"code\":0}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {Client: mockRequest},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
a.NoError(c.SendNotification("1", "s3", mq.Message{}))
|
|
||||||
mockRequest.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveController_SubmitTask(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {
|
|
||||||
jobTracker: map[string]bool{},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// node not exit
|
|
||||||
{
|
|
||||||
a.Equal(ErrMasterNotFound, c.SubmitTask("2", "", "", nil))
|
|
||||||
}
|
|
||||||
|
|
||||||
// success
|
|
||||||
{
|
|
||||||
submitted := false
|
|
||||||
a.NoError(c.SubmitTask("1", "", "hash", func(i interface{}) {
|
|
||||||
submitted = true
|
|
||||||
}))
|
|
||||||
a.True(submitted)
|
|
||||||
}
|
|
||||||
|
|
||||||
// job already submitted
|
|
||||||
{
|
|
||||||
submitted := false
|
|
||||||
a.NoError(c.SubmitTask("1", "", "hash", func(i interface{}) {
|
|
||||||
submitted = true
|
|
||||||
}))
|
|
||||||
a.False(submitted)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveController_GetMasterInfo(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// node not exit
|
|
||||||
{
|
|
||||||
res, err := c.GetMasterInfo("2")
|
|
||||||
a.Equal(ErrMasterNotFound, err)
|
|
||||||
a.Nil(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// success
|
|
||||||
{
|
|
||||||
res, err := c.GetMasterInfo("1")
|
|
||||||
a.NoError(err)
|
|
||||||
a.NotNil(res)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveController_GetOneDriveToken(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// node not exit
|
|
||||||
{
|
|
||||||
res, err := c.GetPolicyOauthToken("2", 1)
|
|
||||||
a.Equal(ErrMasterNotFound, err)
|
|
||||||
a.Empty(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// return none 200
|
|
||||||
{
|
|
||||||
mockRequest := &requestMock{}
|
|
||||||
mockRequest.On("Request", "GET", "/api/v3/slave/credential/1", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{StatusCode: http.StatusConflict},
|
|
||||||
})
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {Client: mockRequest},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
res, err := c.GetPolicyOauthToken("1", 1)
|
|
||||||
a.Error(err)
|
|
||||||
a.Empty(res)
|
|
||||||
mockRequest.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return error
|
|
||||||
{
|
|
||||||
mockRequest := &requestMock{}
|
|
||||||
mockRequest.On("Request", "GET", "/api/v3/slave/credential/1", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"code\":1}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {Client: mockRequest},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
res, err := c.GetPolicyOauthToken("1", 1)
|
|
||||||
a.Equal(1, err.(serializer.AppError).Code)
|
|
||||||
a.Empty(res)
|
|
||||||
mockRequest.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// success
|
|
||||||
{
|
|
||||||
mockRequest := &requestMock{}
|
|
||||||
mockRequest.On("Request", "GET", "/api/v3/slave/credential/1", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"data\":\"expected\"}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
c := &slaveController{
|
|
||||||
masters: map[string]MasterInfo{
|
|
||||||
"1": {Client: mockRequest},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
res, err := c.GetPolicyOauthToken("1", 1)
|
|
||||||
a.NoError(err)
|
|
||||||
a.Equal("expected", res)
|
|
||||||
mockRequest.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
@ -1,186 +0,0 @@
|
|||||||
package cluster
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/aria2/rpc"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/serializer"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"os"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestMasterNode_Init(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &MasterNode{}
|
|
||||||
m.Init(&model.Node{Status: model.NodeSuspend})
|
|
||||||
a.Equal(model.NodeSuspend, m.DBModel().Status)
|
|
||||||
m.Init(&model.Node{Aria2Enabled: true})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMasterNode_DummyMethods(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &MasterNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
}
|
|
||||||
|
|
||||||
m.Model.ID = 5
|
|
||||||
a.Equal(m.Model.ID, m.ID())
|
|
||||||
|
|
||||||
res, err := m.Ping(&serializer.NodePingReq{})
|
|
||||||
a.NoError(err)
|
|
||||||
a.NotNil(res)
|
|
||||||
|
|
||||||
a.True(m.IsActive())
|
|
||||||
a.True(m.IsMater())
|
|
||||||
|
|
||||||
m.SubscribeStatusChange(func(isActive bool, id uint) {})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMasterNode_IsFeatureEnabled(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &MasterNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
}
|
|
||||||
|
|
||||||
a.False(m.IsFeatureEnabled("aria2"))
|
|
||||||
a.False(m.IsFeatureEnabled("random"))
|
|
||||||
m.Model.Aria2Enabled = true
|
|
||||||
a.True(m.IsFeatureEnabled("aria2"))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMasterNode_AuthInstance(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &MasterNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
}
|
|
||||||
|
|
||||||
a.NotNil(m.MasterAuthInstance())
|
|
||||||
a.NotNil(m.SlaveAuthInstance())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMasterNode_Kill(t *testing.T) {
|
|
||||||
m := &MasterNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
}
|
|
||||||
|
|
||||||
m.Kill()
|
|
||||||
|
|
||||||
caller, _ := rpc.New(context.Background(), "http://", "", 0, nil)
|
|
||||||
m.aria2RPC.Caller = caller
|
|
||||||
m.Kill()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMasterNode_GetAria2Instance(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &MasterNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
aria2RPC: rpcService{},
|
|
||||||
}
|
|
||||||
|
|
||||||
m.aria2RPC.parent = m
|
|
||||||
|
|
||||||
a.NotNil(m.GetAria2Instance())
|
|
||||||
m.Model.Aria2Enabled = true
|
|
||||||
a.NotNil(m.GetAria2Instance())
|
|
||||||
m.aria2RPC.Initialized = true
|
|
||||||
a.NotNil(m.GetAria2Instance())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRpcService_Init(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &MasterNode{
|
|
||||||
Model: &model.Node{
|
|
||||||
Aria2OptionsSerialized: model.Aria2Option{
|
|
||||||
Options: "{",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
aria2RPC: rpcService{},
|
|
||||||
}
|
|
||||||
m.aria2RPC.parent = m
|
|
||||||
|
|
||||||
// failed to decode address
|
|
||||||
{
|
|
||||||
m.Model.Aria2OptionsSerialized.Server = string([]byte{0x7f})
|
|
||||||
a.Error(m.aria2RPC.Init())
|
|
||||||
}
|
|
||||||
|
|
||||||
// failed to decode options
|
|
||||||
{
|
|
||||||
m.Model.Aria2OptionsSerialized.Server = ""
|
|
||||||
a.Error(m.aria2RPC.Init())
|
|
||||||
}
|
|
||||||
|
|
||||||
// failed to initialized
|
|
||||||
{
|
|
||||||
m.Model.Aria2OptionsSerialized.Server = ""
|
|
||||||
m.Model.Aria2OptionsSerialized.Options = "{}"
|
|
||||||
caller, _ := rpc.New(context.Background(), "http://", "", 0, nil)
|
|
||||||
m.aria2RPC.Caller = caller
|
|
||||||
a.Error(m.aria2RPC.Init())
|
|
||||||
a.False(m.aria2RPC.Initialized)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func getTestRPCNode() *MasterNode {
|
|
||||||
m := &MasterNode{
|
|
||||||
Model: &model.Node{
|
|
||||||
Aria2OptionsSerialized: model.Aria2Option{},
|
|
||||||
},
|
|
||||||
aria2RPC: rpcService{
|
|
||||||
options: &clientOptions{
|
|
||||||
Options: map[string]interface{}{"1": "1"},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
m.aria2RPC.parent = m
|
|
||||||
caller, _ := rpc.New(context.Background(), "http://", "", 0, nil)
|
|
||||||
m.aria2RPC.Caller = caller
|
|
||||||
return m
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRpcService_CreateTask(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := getTestRPCNode()
|
|
||||||
|
|
||||||
res, err := m.aria2RPC.CreateTask(&model.Download{}, map[string]interface{}{"1": "1"})
|
|
||||||
a.Error(err)
|
|
||||||
a.Empty(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRpcService_Status(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := getTestRPCNode()
|
|
||||||
|
|
||||||
res, err := m.aria2RPC.Status(&model.Download{})
|
|
||||||
a.Error(err)
|
|
||||||
a.Empty(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRpcService_Cancel(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := getTestRPCNode()
|
|
||||||
|
|
||||||
a.Error(m.aria2RPC.Cancel(&model.Download{}))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRpcService_Select(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := getTestRPCNode()
|
|
||||||
|
|
||||||
a.NotNil(m.aria2RPC.GetConfig())
|
|
||||||
a.Error(m.aria2RPC.Select(&model.Download{}, []int{1, 2, 3}))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRpcService_DeleteTempFile(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := getTestRPCNode()
|
|
||||||
fdName := "TestRpcService_DeleteTempFile"
|
|
||||||
a.NoError(os.Mkdir(fdName, 0644))
|
|
||||||
|
|
||||||
a.NoError(m.aria2RPC.DeleteTempFile(&model.Download{Parent: fdName}))
|
|
||||||
time.Sleep(500 * time.Millisecond)
|
|
||||||
a.False(util.Exists(fdName))
|
|
||||||
}
|
|
||||||
@ -1,17 +0,0 @@
|
|||||||
package cluster
|
|
||||||
|
|
||||||
import (
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestNewNodeFromDBModel(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
a.IsType(&SlaveNode{}, NewNodeFromDBModel(&model.Node{
|
|
||||||
Type: model.SlaveNodeType,
|
|
||||||
}))
|
|
||||||
a.IsType(&MasterNode{}, NewNodeFromDBModel(&model.Node{
|
|
||||||
Type: model.MasterNodeType,
|
|
||||||
}))
|
|
||||||
}
|
|
||||||
@ -1,161 +0,0 @@
|
|||||||
package cluster
|
|
||||||
|
|
||||||
import (
|
|
||||||
"database/sql"
|
|
||||||
"errors"
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/balancer"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
var mock sqlmock.Sqlmock
|
|
||||||
|
|
||||||
// TestMain 初始化数据库Mock
|
|
||||||
func TestMain(m *testing.M) {
|
|
||||||
var db *sql.DB
|
|
||||||
var err error
|
|
||||||
db, mock, err = sqlmock.New()
|
|
||||||
if err != nil {
|
|
||||||
panic("An error was not expected when opening a stub database connection")
|
|
||||||
}
|
|
||||||
model.DB, _ = gorm.Open("mysql", db)
|
|
||||||
defer db.Close()
|
|
||||||
m.Run()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestInitFailed(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnError(errors.New("error"))
|
|
||||||
Init()
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestInitSuccess(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
mock.ExpectQuery("SELECT(.+)").WillReturnRows(sqlmock.NewRows([]string{"id", "aria2_enabled", "type"}).AddRow(1, true, model.MasterNodeType))
|
|
||||||
Init()
|
|
||||||
a.NoError(mock.ExpectationsWereMet())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNodePool_GetNodeByID(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
p := &NodePool{}
|
|
||||||
p.Init()
|
|
||||||
mockNode := &nodeMock{}
|
|
||||||
|
|
||||||
// inactive
|
|
||||||
{
|
|
||||||
p.inactive[1] = mockNode
|
|
||||||
a.Equal(mockNode, p.GetNodeByID(1))
|
|
||||||
}
|
|
||||||
|
|
||||||
// active
|
|
||||||
{
|
|
||||||
delete(p.inactive, 1)
|
|
||||||
p.active[1] = mockNode
|
|
||||||
a.Equal(mockNode, p.GetNodeByID(1))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNodePool_NodeStatusChange(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
p := &NodePool{}
|
|
||||||
n := &MasterNode{Model: &model.Node{}}
|
|
||||||
p.Init()
|
|
||||||
p.inactive[1] = n
|
|
||||||
|
|
||||||
p.nodeStatusChange(true, 1)
|
|
||||||
a.Len(p.inactive, 0)
|
|
||||||
a.Equal(n, p.active[1])
|
|
||||||
|
|
||||||
p.nodeStatusChange(false, 1)
|
|
||||||
a.Len(p.active, 0)
|
|
||||||
a.Equal(n, p.inactive[1])
|
|
||||||
|
|
||||||
p.nodeStatusChange(false, 1)
|
|
||||||
a.Len(p.active, 0)
|
|
||||||
a.Equal(n, p.inactive[1])
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNodePool_Add(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
p := &NodePool{}
|
|
||||||
p.Init()
|
|
||||||
|
|
||||||
// new node
|
|
||||||
{
|
|
||||||
p.Add(&model.Node{})
|
|
||||||
a.Len(p.active, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// old node
|
|
||||||
{
|
|
||||||
p.inactive[0] = p.active[0]
|
|
||||||
delete(p.active, 0)
|
|
||||||
p.Add(&model.Node{})
|
|
||||||
a.Len(p.active, 0)
|
|
||||||
a.Len(p.inactive, 1)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNodePool_Delete(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
p := &NodePool{}
|
|
||||||
p.Init()
|
|
||||||
|
|
||||||
// active
|
|
||||||
{
|
|
||||||
mockNode := &nodeMock{}
|
|
||||||
mockNode.On("Kill")
|
|
||||||
p.active[0] = mockNode
|
|
||||||
p.Delete(0)
|
|
||||||
a.Len(p.active, 0)
|
|
||||||
a.Len(p.inactive, 0)
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
p.Init()
|
|
||||||
|
|
||||||
// inactive
|
|
||||||
{
|
|
||||||
mockNode := &nodeMock{}
|
|
||||||
mockNode.On("Kill")
|
|
||||||
p.inactive[0] = mockNode
|
|
||||||
p.Delete(0)
|
|
||||||
a.Len(p.active, 0)
|
|
||||||
a.Len(p.inactive, 0)
|
|
||||||
mockNode.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNodePool_BalanceNodeByFeature(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
p := &NodePool{}
|
|
||||||
p.Init()
|
|
||||||
|
|
||||||
// success
|
|
||||||
{
|
|
||||||
p.featureMap["test"] = []Node{&MasterNode{}}
|
|
||||||
err, res := p.BalanceNodeByFeature("test", balancer.NewBalancer("round-robin"))
|
|
||||||
a.NoError(err)
|
|
||||||
a.Equal(p.featureMap["test"][0], res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// NoNodes
|
|
||||||
{
|
|
||||||
p.featureMap["test"] = []Node{}
|
|
||||||
err, res := p.BalanceNodeByFeature("test", balancer.NewBalancer("round-robin"))
|
|
||||||
a.Error(err)
|
|
||||||
a.Nil(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// No match feature
|
|
||||||
{
|
|
||||||
err, res := p.BalanceNodeByFeature("test2", balancer.NewBalancer("round-robin"))
|
|
||||||
a.Error(err)
|
|
||||||
a.Nil(res)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,559 +0,0 @@
|
|||||||
package cluster
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"encoding/json"
|
|
||||||
"errors"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mocks/requestmock"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/request"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/serializer"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
testMock "github.com/stretchr/testify/mock"
|
|
||||||
"io/ioutil"
|
|
||||||
"net/http"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestSlaveNode_InitAndKill(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
n := &SlaveNode{
|
|
||||||
callback: func(b bool, u uint) {
|
|
||||||
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
a.NotPanics(func() {
|
|
||||||
n.Init(&model.Node{})
|
|
||||||
time.Sleep(time.Millisecond * 500)
|
|
||||||
n.Init(&model.Node{})
|
|
||||||
n.Kill()
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveNode_DummyMethods(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &SlaveNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
}
|
|
||||||
|
|
||||||
m.Model.ID = 5
|
|
||||||
a.Equal(m.Model.ID, m.ID())
|
|
||||||
a.Equal(m.Model.ID, m.DBModel().ID)
|
|
||||||
|
|
||||||
a.False(m.IsActive())
|
|
||||||
a.False(m.IsMater())
|
|
||||||
|
|
||||||
m.SubscribeStatusChange(func(isActive bool, id uint) {})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveNode_IsFeatureEnabled(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &SlaveNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
}
|
|
||||||
|
|
||||||
a.False(m.IsFeatureEnabled("aria2"))
|
|
||||||
a.False(m.IsFeatureEnabled("random"))
|
|
||||||
m.Model.Aria2Enabled = true
|
|
||||||
a.True(m.IsFeatureEnabled("aria2"))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveNode_Ping(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &SlaveNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return error code
|
|
||||||
{
|
|
||||||
mockRequest := &requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "heartbeat", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"code\":1}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
res, err := m.Ping(&serializer.NodePingReq{})
|
|
||||||
a.Error(err)
|
|
||||||
a.Nil(res)
|
|
||||||
a.Equal(1, err.(serializer.AppError).Code)
|
|
||||||
}
|
|
||||||
|
|
||||||
// return unexpected json
|
|
||||||
{
|
|
||||||
mockRequest := &requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "heartbeat", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"data\":\"233\"}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
res, err := m.Ping(&serializer.NodePingReq{})
|
|
||||||
a.Error(err)
|
|
||||||
a.Nil(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// return success
|
|
||||||
{
|
|
||||||
mockRequest := &requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "heartbeat", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"data\":\"{}\"}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
res, err := m.Ping(&serializer.NodePingReq{})
|
|
||||||
a.NoError(err)
|
|
||||||
a.NotNil(res)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveNode_GetAria2Instance(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &SlaveNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
}
|
|
||||||
|
|
||||||
a.NotNil(m.GetAria2Instance())
|
|
||||||
m.Model.Aria2Enabled = true
|
|
||||||
a.NotNil(m.GetAria2Instance())
|
|
||||||
a.NotNil(m.GetAria2Instance())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveNode_StartPingLoop(t *testing.T) {
|
|
||||||
callbackCount := 0
|
|
||||||
finishedChan := make(chan struct{})
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "heartbeat", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 404,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m := &SlaveNode{
|
|
||||||
Active: true,
|
|
||||||
Model: &model.Node{},
|
|
||||||
callback: func(b bool, u uint) {
|
|
||||||
callbackCount++
|
|
||||||
if callbackCount == 2 {
|
|
||||||
close(finishedChan)
|
|
||||||
}
|
|
||||||
if callbackCount == 1 {
|
|
||||||
mockRequest.AssertExpectations(t)
|
|
||||||
mockRequest = requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "heartbeat", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"data\":\"{}\"}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
},
|
|
||||||
}
|
|
||||||
cache.Set("setting_slave_ping_interval", "0", 0)
|
|
||||||
cache.Set("setting_slave_recover_interval", "0", 0)
|
|
||||||
cache.Set("setting_slave_node_retry", "1", 0)
|
|
||||||
|
|
||||||
m.caller.Client = &mockRequest
|
|
||||||
go func() {
|
|
||||||
select {
|
|
||||||
case <-finishedChan:
|
|
||||||
m.Kill()
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
m.StartPingLoop()
|
|
||||||
mockRequest.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveNode_AuthInstance(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := &SlaveNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
}
|
|
||||||
|
|
||||||
a.NotNil(m.MasterAuthInstance())
|
|
||||||
a.NotNil(m.SlaveAuthInstance())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveNode_ChangeStatus(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
isActive := false
|
|
||||||
m := &SlaveNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
callback: func(b bool, u uint) {
|
|
||||||
isActive = b
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
a.NotPanics(func() {
|
|
||||||
m.changeStatus(false)
|
|
||||||
})
|
|
||||||
m.changeStatus(true)
|
|
||||||
a.True(isActive)
|
|
||||||
}
|
|
||||||
|
|
||||||
func getTestRPCNodeSlave() *SlaveNode {
|
|
||||||
m := &SlaveNode{
|
|
||||||
Model: &model.Node{},
|
|
||||||
}
|
|
||||||
m.caller.parent = m
|
|
||||||
return m
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveCaller_CreateTask(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := getTestRPCNodeSlave()
|
|
||||||
|
|
||||||
// master return 404
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/task", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 404,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
res, err := m.caller.CreateTask(&model.Download{}, nil)
|
|
||||||
a.Empty(res)
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return error
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/task", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"code\":1}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
res, err := m.caller.CreateTask(&model.Download{}, nil)
|
|
||||||
a.Empty(res)
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return success
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/task", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"data\":\"res\"}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
res, err := m.caller.CreateTask(&model.Download{}, nil)
|
|
||||||
a.Equal("res", res)
|
|
||||||
a.NoError(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveCaller_Status(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := getTestRPCNodeSlave()
|
|
||||||
|
|
||||||
// master return 404
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/status", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 404,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
res, err := m.caller.Status(&model.Download{})
|
|
||||||
a.Empty(res.Status)
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return error
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/status", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"code\":1}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
res, err := m.caller.Status(&model.Download{})
|
|
||||||
a.Empty(res.Status)
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return success
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/status", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"data\":\"re456456s\"}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
res, err := m.caller.Status(&model.Download{})
|
|
||||||
a.Empty(res.Status)
|
|
||||||
a.NoError(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveCaller_Cancel(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := getTestRPCNodeSlave()
|
|
||||||
|
|
||||||
// master return 404
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/cancel", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 404,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
err := m.caller.Cancel(&model.Download{})
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return error
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/cancel", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"code\":1}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
err := m.caller.Cancel(&model.Download{})
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return success
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/cancel", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"data\":\"res\"}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
err := m.caller.Cancel(&model.Download{})
|
|
||||||
a.NoError(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveCaller_Select(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := getTestRPCNodeSlave()
|
|
||||||
m.caller.Init()
|
|
||||||
m.caller.GetConfig()
|
|
||||||
|
|
||||||
// master return 404
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/select", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 404,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
err := m.caller.Select(&model.Download{}, nil)
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return error
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/select", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"code\":1}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
err := m.caller.Select(&model.Download{}, nil)
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return success
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/select", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"data\":\"res\"}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
err := m.caller.Select(&model.Download{}, nil)
|
|
||||||
a.NoError(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSlaveCaller_DeleteTempFile(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
m := getTestRPCNodeSlave()
|
|
||||||
m.caller.Init()
|
|
||||||
m.caller.GetConfig()
|
|
||||||
|
|
||||||
// master return 404
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/delete", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 404,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
err := m.caller.DeleteTempFile(&model.Download{})
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return error
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/delete", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"code\":1}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
err := m.caller.DeleteTempFile(&model.Download{})
|
|
||||||
a.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// master return success
|
|
||||||
{
|
|
||||||
mockRequest := requestMock{}
|
|
||||||
mockRequest.On("Request", "POST", "aria2/delete", testMock.Anything, testMock.Anything).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("{\"data\":\"res\"}")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
m.caller.Client = mockRequest
|
|
||||||
err := m.caller.DeleteTempFile(&model.Download{})
|
|
||||||
a.NoError(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoteCallback(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 回调成功
|
|
||||||
{
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
mockResp, _ := json.Marshal(serializer.Response{Code: 0})
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test/test/url",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(bytes.NewReader(mockResp)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
request.GeneralClient = clientMock
|
|
||||||
resp := RemoteCallback("http://test/test/url", serializer.UploadCallback{})
|
|
||||||
asserts.NoError(resp)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 服务端返回业务错误
|
|
||||||
{
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
mockResp, _ := json.Marshal(serializer.Response{Code: 401})
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test/test/url",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(bytes.NewReader(mockResp)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
request.GeneralClient = clientMock
|
|
||||||
resp := RemoteCallback("http://test/test/url", serializer.UploadCallback{})
|
|
||||||
asserts.EqualValues(401, resp.(serializer.AppError).Code)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法解析回调响应
|
|
||||||
{
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test/test/url",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("mockResp")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
request.GeneralClient = clientMock
|
|
||||||
resp := RemoteCallback("http://test/test/url", serializer.UploadCallback{})
|
|
||||||
asserts.Error(resp)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// HTTP状态码非200
|
|
||||||
{
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test/test/url",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 404,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader("mockResp")),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
request.GeneralClient = clientMock
|
|
||||||
resp := RemoteCallback("http://test/test/url", serializer.UploadCallback{})
|
|
||||||
asserts.Error(resp)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法发起回调
|
|
||||||
{
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test/test/url",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: errors.New("error"),
|
|
||||||
})
|
|
||||||
request.GeneralClient = clientMock
|
|
||||||
resp := RemoteCallback("http://test/test/url", serializer.UploadCallback{})
|
|
||||||
asserts.Error(resp)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,100 +0,0 @@
|
|||||||
package conf
|
|
||||||
|
|
||||||
import (
|
|
||||||
"io/ioutil"
|
|
||||||
"os"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
// 测试Init日志路径错误
|
|
||||||
func TestInitPanic(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
// 日志路径不存在时
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
Init("not/exist/path/conf.ini")
|
|
||||||
})
|
|
||||||
|
|
||||||
asserts.True(util.Exists("not/exist/path/conf.ini"))
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestInitDelimiterNotFound 日志路径存在但 Key 格式错误时
|
|
||||||
func TestInitDelimiterNotFound(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
testCase := `[Database]
|
|
||||||
Type = mysql
|
|
||||||
User = root
|
|
||||||
Password233root
|
|
||||||
Host = 127.0.0.1:3306
|
|
||||||
Name = v3
|
|
||||||
TablePrefix = v3_`
|
|
||||||
err := ioutil.WriteFile("testConf.ini", []byte(testCase), 0644)
|
|
||||||
defer func() { err = os.Remove("testConf.ini") }()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
asserts.Panics(func() {
|
|
||||||
Init("testConf.ini")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// TestInitNoPanic 日志路径存在且合法时
|
|
||||||
func TestInitNoPanic(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
testCase := `
|
|
||||||
[System]
|
|
||||||
Listen = 3000
|
|
||||||
HashIDSalt = 1
|
|
||||||
|
|
||||||
[Database]
|
|
||||||
Type = mysql
|
|
||||||
User = root
|
|
||||||
Password = root
|
|
||||||
Host = 127.0.0.1:3306
|
|
||||||
Name = v3
|
|
||||||
TablePrefix = v3_
|
|
||||||
|
|
||||||
[OptionOverwrite]
|
|
||||||
key=value
|
|
||||||
`
|
|
||||||
err := ioutil.WriteFile("testConf.ini", []byte(testCase), 0644)
|
|
||||||
defer func() { err = os.Remove("testConf.ini") }()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
asserts.NotPanics(func() {
|
|
||||||
Init("testConf.ini")
|
|
||||||
})
|
|
||||||
asserts.Equal(OptionOverwrite["key"], "value")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMapSection(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
|
|
||||||
//正常情况
|
|
||||||
testCase := `
|
|
||||||
[System]
|
|
||||||
Listen = 3000
|
|
||||||
HashIDSalt = 1
|
|
||||||
|
|
||||||
[Database]
|
|
||||||
Type = mysql
|
|
||||||
User = root
|
|
||||||
Password:root
|
|
||||||
Host = 127.0.0.1:3306
|
|
||||||
Name = v3
|
|
||||||
TablePrefix = v3_`
|
|
||||||
err := ioutil.WriteFile("testConf.ini", []byte(testCase), 0644)
|
|
||||||
defer func() { err = os.Remove("testConf.ini") }()
|
|
||||||
if err != nil {
|
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
Init("testConf.ini")
|
|
||||||
err = mapSection("Database", DatabaseConfig)
|
|
||||||
asserts.NoError(err)
|
|
||||||
|
|
||||||
}
|
|
||||||
@ -1,16 +1,22 @@
|
|||||||
package conf
|
package conf
|
||||||
|
|
||||||
|
// plusVersion 增强版版本号
|
||||||
|
const plusVersion = "+1.1"
|
||||||
|
|
||||||
// BackendVersion 当前后端版本号
|
// BackendVersion 当前后端版本号
|
||||||
var BackendVersion = "3.8.3"
|
const BackendVersion = "3.8.3" + plusVersion
|
||||||
|
|
||||||
|
// KeyVersion 授权版本号
|
||||||
|
const KeyVersion = "3.3.1"
|
||||||
|
|
||||||
// RequiredDBVersion 与当前版本匹配的数据库版本
|
// RequiredDBVersion 与当前版本匹配的数据库版本
|
||||||
var RequiredDBVersion = "3.8.1"
|
const RequiredDBVersion = "3.8.1+1.0-plus"
|
||||||
|
|
||||||
// RequiredStaticVersion 与当前版本匹配的静态资源版本
|
// RequiredStaticVersion 与当前版本匹配的静态资源版本
|
||||||
var RequiredStaticVersion = "3.8.3"
|
const RequiredStaticVersion = "3.8.3" + plusVersion
|
||||||
|
|
||||||
// IsPro 是否为Pro版本
|
// IsPlus 是否为Plus版本
|
||||||
var IsPro = "false"
|
const IsPlus = "true"
|
||||||
|
|
||||||
// LastCommit 最后commit id
|
// LastCommit 最后commit id
|
||||||
var LastCommit = "a11f819"
|
const LastCommit = "88409cc"
|
||||||
|
|||||||
@ -0,0 +1,83 @@
|
|||||||
|
package crontab
|
||||||
|
|
||||||
|
import (
|
||||||
|
model "github.com/cloudreve/Cloudreve/v3/models"
|
||||||
|
"github.com/cloudreve/Cloudreve/v3/pkg/email"
|
||||||
|
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
||||||
|
)
|
||||||
|
|
||||||
|
func notifyExpiredVAS() {
|
||||||
|
checkStoragePack()
|
||||||
|
checkUserGroup()
|
||||||
|
util.Log().Info("Crontab job \"cron_notify_user\" complete.")
|
||||||
|
}
|
||||||
|
|
||||||
|
// banOverusedUser 封禁超出宽容期的用户
|
||||||
|
func banOverusedUser() {
|
||||||
|
users := model.GetTolerantExpiredUser()
|
||||||
|
for _, user := range users {
|
||||||
|
|
||||||
|
// 清除最后通知日期标记
|
||||||
|
user.ClearNotified()
|
||||||
|
|
||||||
|
// 检查容量是否超额
|
||||||
|
if user.Storage > user.Group.MaxStorage+user.GetAvailablePackSize() {
|
||||||
|
// 封禁用户
|
||||||
|
user.SetStatus(model.OveruseBaned)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkUserGroup 检查已过期用户组
|
||||||
|
func checkUserGroup() {
|
||||||
|
users := model.GetGroupExpiredUsers()
|
||||||
|
for _, user := range users {
|
||||||
|
|
||||||
|
// 将用户回退到初始用户组
|
||||||
|
user.GroupFallback()
|
||||||
|
|
||||||
|
// 重新加载用户
|
||||||
|
user, _ = model.GetUserByID(user.ID)
|
||||||
|
|
||||||
|
// 检查容量是否超额
|
||||||
|
if user.Storage > user.Group.MaxStorage+user.GetAvailablePackSize() {
|
||||||
|
// 如果超额,则通知用户
|
||||||
|
sendNotification(&user, "用户组过期")
|
||||||
|
// 更新最后通知日期
|
||||||
|
user.Notified()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// checkStoragePack 检查已过期的容量包
|
||||||
|
func checkStoragePack() {
|
||||||
|
packs := model.GetExpiredStoragePack()
|
||||||
|
for _, pack := range packs {
|
||||||
|
// 删除过期的容量包
|
||||||
|
pack.Delete()
|
||||||
|
|
||||||
|
//找到所属用户
|
||||||
|
user, err := model.GetUserByID(pack.UserID)
|
||||||
|
if err != nil {
|
||||||
|
util.Log().Warning("Crontab job failed to get user info of [UID=%d]: %s", pack.UserID, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// 检查容量是否超额
|
||||||
|
if user.Storage > user.Group.MaxStorage+user.GetAvailablePackSize() {
|
||||||
|
// 如果超额,则通知用户
|
||||||
|
sendNotification(&user, "容量包过期")
|
||||||
|
|
||||||
|
// 更新最后通知日期
|
||||||
|
user.Notified()
|
||||||
|
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sendNotification(user *model.User, reason string) {
|
||||||
|
title, body := email.NewOveruseNotification(user.Nick, reason)
|
||||||
|
if err := email.Send(user.Email, title, body); err != nil {
|
||||||
|
util.Log().Warning("Failed to send notification email: %s", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@ -1,256 +0,0 @@
|
|||||||
package filesystem
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/util"
|
|
||||||
testMock "github.com/stretchr/testify/mock"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"runtime"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/DATA-DOG/go-sqlmock"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/filesystem/fsctx"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestFileSystem_Compress(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
fs := FileSystem{
|
|
||||||
User: &model.User{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
// 查找压缩父目录
|
|
||||||
mock.ExpectQuery("SELECT(.+)folders(.+)").
|
|
||||||
WithArgs(1, 1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "parent"))
|
|
||||||
// 查找顶级待压缩文件
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs(1, 1).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows(
|
|
||||||
[]string{"id", "name", "source_name", "policy_id"}).
|
|
||||||
AddRow(1, "1.txt", "tests/file1.txt", 1),
|
|
||||||
)
|
|
||||||
asserts.NoError(cache.Set("setting_temp_path", "tests", -1))
|
|
||||||
// 查找父目录子文件
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs(1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name", "source_name", "policy_id"}))
|
|
||||||
// 查找子目录
|
|
||||||
mock.ExpectQuery("SELECT(.+)folders(.+)").
|
|
||||||
WithArgs(1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(2, "sub"))
|
|
||||||
// 查找子目录子文件
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs(2).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows([]string{"id", "name", "source_name", "policy_id"}).
|
|
||||||
AddRow(2, "2.txt", "tests/file2.txt", 1),
|
|
||||||
)
|
|
||||||
// 查找上传策略
|
|
||||||
asserts.NoError(cache.Set("policy_1", model.Policy{Type: "local"}, -1))
|
|
||||||
w := &bytes.Buffer{}
|
|
||||||
|
|
||||||
err := fs.Compress(ctx, w, []uint{1}, []uint{1}, true)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.NotEmpty(w.Len())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 上下文取消
|
|
||||||
{
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
|
||||||
cancel()
|
|
||||||
// 查找压缩父目录
|
|
||||||
mock.ExpectQuery("SELECT(.+)folders(.+)").
|
|
||||||
WithArgs(1, 1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "parent"))
|
|
||||||
// 查找顶级待压缩文件
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs(1, 1).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows(
|
|
||||||
[]string{"id", "name", "source_name", "policy_id"}).
|
|
||||||
AddRow(1, "1.txt", "tests/file1.txt", 1),
|
|
||||||
)
|
|
||||||
asserts.NoError(cache.Set("setting_temp_path", "tests", -1))
|
|
||||||
|
|
||||||
w := &bytes.Buffer{}
|
|
||||||
err := fs.Compress(ctx, w, []uint{1}, []uint{1}, true)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.NotEmpty(w.Len())
|
|
||||||
}
|
|
||||||
|
|
||||||
// 限制父目录
|
|
||||||
{
|
|
||||||
ctx := context.WithValue(context.Background(), fsctx.LimitParentCtx, &model.Folder{
|
|
||||||
Model: gorm.Model{ID: 3},
|
|
||||||
})
|
|
||||||
// 查找压缩父目录
|
|
||||||
mock.ExpectQuery("SELECT(.+)folders(.+)").
|
|
||||||
WithArgs(1, 1).
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name", "parent_id"}).AddRow(1, "parent", 3))
|
|
||||||
// 查找顶级待压缩文件
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WithArgs(1, 1).
|
|
||||||
WillReturnRows(
|
|
||||||
sqlmock.NewRows(
|
|
||||||
[]string{"id", "name", "source_name", "policy_id"}).
|
|
||||||
AddRow(1, "1.txt", "tests/file1.txt", 1),
|
|
||||||
)
|
|
||||||
asserts.NoError(cache.Set("setting_temp_path", "tests", -1))
|
|
||||||
|
|
||||||
w := &bytes.Buffer{}
|
|
||||||
err := fs.Compress(ctx, w, []uint{1}, []uint{1}, true)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Equal(ErrObjectNotExist, err)
|
|
||||||
asserts.Empty(w.Len())
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
type MockNopRSC string
|
|
||||||
|
|
||||||
func (m MockNopRSC) Read(b []byte) (int, error) {
|
|
||||||
return 0, errors.New("read error")
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m MockNopRSC) Seek(n int64, offset int) (int64, error) {
|
|
||||||
return 0, errors.New("read error")
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m MockNopRSC) Close() error {
|
|
||||||
return errors.New("read error")
|
|
||||||
}
|
|
||||||
|
|
||||||
type MockRSC struct {
|
|
||||||
rs io.ReadSeeker
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m MockRSC) Read(b []byte) (int, error) {
|
|
||||||
return m.rs.Read(b)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m MockRSC) Seek(n int64, offset int) (int64, error) {
|
|
||||||
return m.rs.Seek(n, offset)
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m MockRSC) Close() error {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
var basepath string
|
|
||||||
|
|
||||||
func init() {
|
|
||||||
_, currentFile, _, _ := runtime.Caller(0)
|
|
||||||
basepath = filepath.Dir(currentFile)
|
|
||||||
}
|
|
||||||
|
|
||||||
func Path(rel string) string {
|
|
||||||
return filepath.Join(basepath, rel)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFileSystem_Decompress(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
ctx := context.Background()
|
|
||||||
fs := FileSystem{
|
|
||||||
User: &model.User{Model: gorm.Model{ID: 1}},
|
|
||||||
}
|
|
||||||
os.RemoveAll(util.RelativePath("tests/decompress"))
|
|
||||||
|
|
||||||
// 压缩文件不存在
|
|
||||||
{
|
|
||||||
// 查找根目录
|
|
||||||
mock.ExpectQuery("SELECT(.+)folders(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "/"))
|
|
||||||
// 查找压缩文件,未找到
|
|
||||||
mock.ExpectQuery("SELECT(.+)files(.+)").
|
|
||||||
WillReturnRows(sqlmock.NewRows([]string{"id", "name"}))
|
|
||||||
err := fs.Decompress(ctx, "/1.zip", "/", "")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法下载压缩文件
|
|
||||||
{
|
|
||||||
fs.FileTarget = []model.File{{SourceName: "1.zip", Policy: model.Policy{Type: "mock"}}}
|
|
||||||
fs.FileTarget[0].Policy.ID = 1
|
|
||||||
testHandler := new(FileHeaderMock)
|
|
||||||
testHandler.On("Get", testMock.Anything, "1.zip").Return(MockRSC{}, errors.New("error"))
|
|
||||||
fs.Handler = testHandler
|
|
||||||
err := fs.Decompress(ctx, "/1.zip", "/", "")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.EqualError(err, "error")
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法创建临时压缩文件
|
|
||||||
{
|
|
||||||
cache.Set("setting_temp_path", "/tests:", 0)
|
|
||||||
fs.FileTarget = []model.File{{SourceName: "1.zip", Policy: model.Policy{Type: "mock"}}}
|
|
||||||
fs.FileTarget[0].Policy.ID = 1
|
|
||||||
testHandler := new(FileHeaderMock)
|
|
||||||
testHandler.On("Get", testMock.Anything, "1.zip").Return(MockRSC{}, nil)
|
|
||||||
fs.Handler = testHandler
|
|
||||||
err := fs.Decompress(ctx, "/1.zip", "/", "")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法写入压缩文件
|
|
||||||
{
|
|
||||||
cache.Set("setting_temp_path", "tests", 0)
|
|
||||||
fs.FileTarget = []model.File{{SourceName: "1.zip", Policy: model.Policy{Type: "mock"}}}
|
|
||||||
fs.FileTarget[0].Policy.ID = 1
|
|
||||||
testHandler := new(FileHeaderMock)
|
|
||||||
testHandler.On("Get", testMock.Anything, "1.zip").Return(MockNopRSC("1"), nil)
|
|
||||||
fs.Handler = testHandler
|
|
||||||
err := fs.Decompress(ctx, "/1.zip", "/", "")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Contains(err.Error(), "read error")
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法重设上传策略
|
|
||||||
{
|
|
||||||
cache.Set("setting_temp_path", "tests", 0)
|
|
||||||
fs.FileTarget = []model.File{{SourceName: "1.zip", Policy: model.Policy{Type: "mock"}}}
|
|
||||||
fs.FileTarget[0].Policy.ID = 1
|
|
||||||
testHandler := new(FileHeaderMock)
|
|
||||||
testHandler.On("Get", testMock.Anything, "1.zip").Return(MockRSC{rs: strings.NewReader("read")}, nil)
|
|
||||||
fs.Handler = testHandler
|
|
||||||
err := fs.Decompress(ctx, "/1.zip", "/", "")
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.True(util.IsEmpty(util.RelativePath("tests/decompress")))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法上传,容量不足
|
|
||||||
{
|
|
||||||
cache.Set("setting_max_parallel_transfer", "1", 0)
|
|
||||||
zipFile, _ := os.Open(Path("tests/test.zip"))
|
|
||||||
fs.FileTarget = []model.File{{SourceName: "1.zip", Policy: model.Policy{Type: "mock"}}}
|
|
||||||
fs.FileTarget[0].Policy.ID = 1
|
|
||||||
fs.User.Policy.Type = "mock"
|
|
||||||
testHandler := new(FileHeaderMock)
|
|
||||||
testHandler.On("Get", testMock.Anything, "1.zip").Return(zipFile, nil)
|
|
||||||
fs.Handler = testHandler
|
|
||||||
|
|
||||||
fs.Decompress(ctx, "/1.zip", "/", "")
|
|
||||||
|
|
||||||
zipFile.Close()
|
|
||||||
|
|
||||||
asserts.NoError(mock.ExpectationsWereMet())
|
|
||||||
testHandler.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,61 +0,0 @@
|
|||||||
package backoff
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"net/http"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestConstantBackoff_Next(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
// General error
|
|
||||||
{
|
|
||||||
err := errors.New("error")
|
|
||||||
b := &ConstantBackoff{Sleep: time.Duration(0), Max: 3}
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.False(b.Next(err))
|
|
||||||
b.Reset()
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.False(b.Next(err))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Retryable error
|
|
||||||
{
|
|
||||||
err := &RetryableError{RetryAfter: time.Duration(1)}
|
|
||||||
b := &ConstantBackoff{Sleep: time.Duration(0), Max: 3}
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.False(b.Next(err))
|
|
||||||
b.Reset()
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.True(b.Next(err))
|
|
||||||
a.False(b.Next(err))
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNewRetryableErrorFromHeader(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
// no retry-after header
|
|
||||||
{
|
|
||||||
err := NewRetryableErrorFromHeader(nil, http.Header{})
|
|
||||||
a.Empty(err.RetryAfter)
|
|
||||||
}
|
|
||||||
|
|
||||||
// with retry-after header
|
|
||||||
{
|
|
||||||
header := http.Header{}
|
|
||||||
header.Add("retry-after", "120")
|
|
||||||
err := NewRetryableErrorFromHeader(nil, header)
|
|
||||||
a.EqualValues(time.Duration(120)*time.Second, err.RetryAfter)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,250 +0,0 @@
|
|||||||
package chunk
|
|
||||||
|
|
||||||
import (
|
|
||||||
"errors"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/filesystem/chunk/backoff"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/filesystem/fsctx"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
"io"
|
|
||||||
"os"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestNewChunkGroup(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
testCases := []struct {
|
|
||||||
fileSize uint64
|
|
||||||
chunkSize uint64
|
|
||||||
expectedInnerChunkSize uint64
|
|
||||||
expectedChunkNum uint64
|
|
||||||
expectedInfo [][2]int //Start, Index,Length
|
|
||||||
}{
|
|
||||||
{10, 0, 10, 1, [][2]int{{0, 10}}},
|
|
||||||
{0, 0, 0, 1, [][2]int{{0, 0}}},
|
|
||||||
{0, 10, 10, 1, [][2]int{{0, 0}}},
|
|
||||||
{50, 10, 10, 5, [][2]int{
|
|
||||||
{0, 10},
|
|
||||||
{10, 10},
|
|
||||||
{20, 10},
|
|
||||||
{30, 10},
|
|
||||||
{40, 10},
|
|
||||||
}},
|
|
||||||
{50, 50, 50, 1, [][2]int{
|
|
||||||
{0, 50},
|
|
||||||
}},
|
|
||||||
|
|
||||||
{50, 15, 15, 4, [][2]int{
|
|
||||||
{0, 15},
|
|
||||||
{15, 15},
|
|
||||||
{30, 15},
|
|
||||||
{45, 5},
|
|
||||||
}},
|
|
||||||
}
|
|
||||||
|
|
||||||
for index, testCase := range testCases {
|
|
||||||
file := &fsctx.FileStream{Size: testCase.fileSize}
|
|
||||||
chunkGroup := NewChunkGroup(file, testCase.chunkSize, &backoff.ConstantBackoff{}, true)
|
|
||||||
a.EqualValues(testCase.expectedChunkNum, chunkGroup.Num(),
|
|
||||||
"TestCase:%d,ChunkNum()", index)
|
|
||||||
a.EqualValues(testCase.expectedInnerChunkSize, chunkGroup.chunkSize,
|
|
||||||
"TestCase:%d,InnerChunkSize()", index)
|
|
||||||
a.EqualValues(testCase.expectedChunkNum, chunkGroup.Num(),
|
|
||||||
"TestCase:%d,len(Chunks)", index)
|
|
||||||
a.EqualValues(testCase.fileSize, chunkGroup.Total())
|
|
||||||
|
|
||||||
for cIndex, info := range testCase.expectedInfo {
|
|
||||||
a.True(chunkGroup.Next())
|
|
||||||
a.EqualValues(info[1], chunkGroup.Length(),
|
|
||||||
"TestCase:%d,Chunks[%d].Length()", index, cIndex)
|
|
||||||
a.EqualValues(info[0], chunkGroup.Start(),
|
|
||||||
"TestCase:%d,Chunks[%d].Start()", index, cIndex)
|
|
||||||
|
|
||||||
a.Equal(cIndex == len(testCase.expectedInfo)-1, chunkGroup.IsLast(),
|
|
||||||
"TestCase:%d,Chunks[%d].IsLast()", index, cIndex)
|
|
||||||
|
|
||||||
a.NotEmpty(chunkGroup.RangeHeader())
|
|
||||||
}
|
|
||||||
a.False(chunkGroup.Next())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestChunkGroup_TempAvailablet(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
file := &fsctx.FileStream{Size: 1}
|
|
||||||
c := NewChunkGroup(file, 0, &backoff.ConstantBackoff{}, true)
|
|
||||||
a.False(c.TempAvailable())
|
|
||||||
|
|
||||||
f, err := os.CreateTemp("", "TestChunkGroup_TempAvailablet.*")
|
|
||||||
defer func() {
|
|
||||||
f.Close()
|
|
||||||
os.Remove(f.Name())
|
|
||||||
}()
|
|
||||||
a.NoError(err)
|
|
||||||
c.bufferTemp = f
|
|
||||||
|
|
||||||
a.False(c.TempAvailable())
|
|
||||||
f.Write([]byte("1"))
|
|
||||||
a.True(c.TempAvailable())
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestChunkGroup_Process(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
file := &fsctx.FileStream{Size: 10}
|
|
||||||
|
|
||||||
// success
|
|
||||||
{
|
|
||||||
file.File = io.NopCloser(strings.NewReader("1234567890"))
|
|
||||||
c := NewChunkGroup(file, 5, &backoff.ConstantBackoff{}, true)
|
|
||||||
count := 0
|
|
||||||
a.True(c.Next())
|
|
||||||
a.NoError(c.Process(func(c *ChunkGroup, chunk io.Reader) error {
|
|
||||||
count++
|
|
||||||
res, err := io.ReadAll(chunk)
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues("12345", string(res))
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
a.True(c.Next())
|
|
||||||
a.NoError(c.Process(func(c *ChunkGroup, chunk io.Reader) error {
|
|
||||||
count++
|
|
||||||
res, err := io.ReadAll(chunk)
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues("67890", string(res))
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
a.False(c.Next())
|
|
||||||
a.Equal(2, count)
|
|
||||||
}
|
|
||||||
|
|
||||||
// retry, read from buffer file
|
|
||||||
{
|
|
||||||
file.File = io.NopCloser(strings.NewReader("1234567890"))
|
|
||||||
c := NewChunkGroup(file, 5, &backoff.ConstantBackoff{Max: 2}, true)
|
|
||||||
count := 0
|
|
||||||
a.True(c.Next())
|
|
||||||
a.NoError(c.Process(func(c *ChunkGroup, chunk io.Reader) error {
|
|
||||||
count++
|
|
||||||
res, err := io.ReadAll(chunk)
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues("12345", string(res))
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
a.True(c.Next())
|
|
||||||
a.NoError(c.Process(func(c *ChunkGroup, chunk io.Reader) error {
|
|
||||||
count++
|
|
||||||
res, err := io.ReadAll(chunk)
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues("67890", string(res))
|
|
||||||
if count == 2 {
|
|
||||||
return errors.New("error")
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
a.False(c.Next())
|
|
||||||
a.Equal(3, count)
|
|
||||||
}
|
|
||||||
|
|
||||||
// retry, read from seeker
|
|
||||||
{
|
|
||||||
f, _ := os.CreateTemp("", "TestChunkGroup_Process.*")
|
|
||||||
f.Write([]byte("1234567890"))
|
|
||||||
f.Seek(0, 0)
|
|
||||||
defer func() {
|
|
||||||
f.Close()
|
|
||||||
os.Remove(f.Name())
|
|
||||||
}()
|
|
||||||
file.File = f
|
|
||||||
file.Seeker = f
|
|
||||||
c := NewChunkGroup(file, 5, &backoff.ConstantBackoff{Max: 2}, false)
|
|
||||||
count := 0
|
|
||||||
a.True(c.Next())
|
|
||||||
a.NoError(c.Process(func(c *ChunkGroup, chunk io.Reader) error {
|
|
||||||
count++
|
|
||||||
res, err := io.ReadAll(chunk)
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues("12345", string(res))
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
a.True(c.Next())
|
|
||||||
a.NoError(c.Process(func(c *ChunkGroup, chunk io.Reader) error {
|
|
||||||
count++
|
|
||||||
res, err := io.ReadAll(chunk)
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues("67890", string(res))
|
|
||||||
if count == 2 {
|
|
||||||
return errors.New("error")
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
a.False(c.Next())
|
|
||||||
a.Equal(3, count)
|
|
||||||
}
|
|
||||||
|
|
||||||
// retry, seek error
|
|
||||||
{
|
|
||||||
f, _ := os.CreateTemp("", "TestChunkGroup_Process.*")
|
|
||||||
f.Write([]byte("1234567890"))
|
|
||||||
f.Seek(0, 0)
|
|
||||||
defer func() {
|
|
||||||
f.Close()
|
|
||||||
os.Remove(f.Name())
|
|
||||||
}()
|
|
||||||
file.File = f
|
|
||||||
file.Seeker = f
|
|
||||||
c := NewChunkGroup(file, 5, &backoff.ConstantBackoff{Max: 2}, false)
|
|
||||||
count := 0
|
|
||||||
a.True(c.Next())
|
|
||||||
a.NoError(c.Process(func(c *ChunkGroup, chunk io.Reader) error {
|
|
||||||
count++
|
|
||||||
res, err := io.ReadAll(chunk)
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues("12345", string(res))
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
a.True(c.Next())
|
|
||||||
f.Close()
|
|
||||||
a.Error(c.Process(func(c *ChunkGroup, chunk io.Reader) error {
|
|
||||||
count++
|
|
||||||
if count == 2 {
|
|
||||||
return errors.New("error")
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
a.False(c.Next())
|
|
||||||
a.Equal(2, count)
|
|
||||||
}
|
|
||||||
|
|
||||||
// retry, finally error
|
|
||||||
{
|
|
||||||
f, _ := os.CreateTemp("", "TestChunkGroup_Process.*")
|
|
||||||
f.Write([]byte("1234567890"))
|
|
||||||
f.Seek(0, 0)
|
|
||||||
defer func() {
|
|
||||||
f.Close()
|
|
||||||
os.Remove(f.Name())
|
|
||||||
}()
|
|
||||||
file.File = f
|
|
||||||
file.Seeker = f
|
|
||||||
c := NewChunkGroup(file, 5, &backoff.ConstantBackoff{Max: 2}, false)
|
|
||||||
count := 0
|
|
||||||
a.True(c.Next())
|
|
||||||
a.NoError(c.Process(func(c *ChunkGroup, chunk io.Reader) error {
|
|
||||||
count++
|
|
||||||
res, err := io.ReadAll(chunk)
|
|
||||||
a.NoError(err)
|
|
||||||
a.EqualValues("12345", string(res))
|
|
||||||
return nil
|
|
||||||
}))
|
|
||||||
a.True(c.Next())
|
|
||||||
a.Error(c.Process(func(c *ChunkGroup, chunk io.Reader) error {
|
|
||||||
count++
|
|
||||||
return errors.New("error")
|
|
||||||
}))
|
|
||||||
a.False(c.Next())
|
|
||||||
a.Equal(4, count)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
File diff suppressed because it is too large
Load Diff
@ -1,32 +0,0 @@
|
|||||||
package onedrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestNewClient(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
// getOAuthEndpoint失败
|
|
||||||
{
|
|
||||||
policy := model.Policy{
|
|
||||||
BaseURL: string([]byte{0x7f}),
|
|
||||||
}
|
|
||||||
res, err := NewClient(&policy)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Nil(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
policy := model.Policy{}
|
|
||||||
res, err := NewClient(&policy)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.NotNil(res)
|
|
||||||
asserts.NotNil(res.Credential)
|
|
||||||
asserts.NotNil(res.Endpoints)
|
|
||||||
asserts.NotNil(res.Endpoints.OAuthEndpoints)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,420 +0,0 @@
|
|||||||
package onedrive
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"fmt"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mq"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/serializer"
|
|
||||||
"github.com/jinzhu/gorm"
|
|
||||||
"io"
|
|
||||||
"io/ioutil"
|
|
||||||
"net/http"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/filesystem/fsctx"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/request"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
testMock "github.com/stretchr/testify/mock"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestDriver_Token(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
h, _ := NewDriver(&model.Policy{
|
|
||||||
AccessKey: "ak",
|
|
||||||
SecretKey: "sk",
|
|
||||||
BucketName: "test",
|
|
||||||
Server: "test.com",
|
|
||||||
})
|
|
||||||
handler := h.(Driver)
|
|
||||||
|
|
||||||
// 分片上传 失败
|
|
||||||
{
|
|
||||||
cache.Set("setting_siteURL", "http://test.cloudreve.org", 0)
|
|
||||||
handler.Client, _ = NewClient(&model.Policy{})
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 400,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"uploadUrl":"123321"}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client.Request = clientMock
|
|
||||||
res, err := handler.Token(context.Background(), 10, &serializer.UploadSession{}, &fsctx.FileStream{})
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Nil(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 分片上传 成功
|
|
||||||
{
|
|
||||||
cache.Set("setting_siteURL", "http://test.cloudreve.org", 0)
|
|
||||||
cache.Set("setting_onedrive_monitor_timeout", "600", 0)
|
|
||||||
cache.Set("setting_onedrive_callback_check", "20", 0)
|
|
||||||
handler.Client, _ = NewClient(&model.Policy{})
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
handler.Client.Credential.AccessToken = "1"
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"uploadUrl":"123321"}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client.Request = clientMock
|
|
||||||
go func() {
|
|
||||||
time.Sleep(time.Duration(1) * time.Second)
|
|
||||||
mq.GlobalMQ.Publish("TestDriver_Token", mq.Message{})
|
|
||||||
}()
|
|
||||||
res, err := handler.Token(context.Background(), 10, &serializer.UploadSession{Key: "TestDriver_Token"}, &fsctx.FileStream{})
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("123321", res.UploadURLs[0])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriver_Source(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
AccessKey: "ak",
|
|
||||||
SecretKey: "sk",
|
|
||||||
BucketName: "test",
|
|
||||||
Server: "test.com",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
handler.Client, _ = NewClient(&model.Policy{})
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
cache.Set("setting_onedrive_source_timeout", "1800", 0)
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
res, err := handler.Source(context.Background(), "123.jpg", 1, true, 0)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Empty(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 命中缓存 成功
|
|
||||||
{
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
handler.Client.Credential.AccessToken = "1"
|
|
||||||
cache.Set("onedrive_source_0_123.jpg", "res", 1)
|
|
||||||
res, err := handler.Source(context.Background(), "123.jpg", 0, true, 0)
|
|
||||||
cache.Deletes([]string{"0_123.jpg"}, "onedrive_source_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("res", res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 命中缓存 上下文存在文件 成功
|
|
||||||
{
|
|
||||||
file := model.File{}
|
|
||||||
file.ID = 1
|
|
||||||
file.UpdatedAt = time.Now()
|
|
||||||
ctx := context.WithValue(context.Background(), fsctx.FileModelCtx, file)
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
handler.Client.Credential.AccessToken = "1"
|
|
||||||
cache.Set(fmt.Sprintf("onedrive_source_file_%d_1", file.UpdatedAt.Unix()), "res", 0)
|
|
||||||
res, err := handler.Source(ctx, "123.jpg", 1, true, 0)
|
|
||||||
cache.Deletes([]string{"0_123.jpg"}, "onedrive_source_")
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("res", res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"GET",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"@microsoft.graph.downloadUrl":"123321"}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client.Request = clientMock
|
|
||||||
handler.Client.Credential.AccessToken = "1"
|
|
||||||
res, err := handler.Source(context.Background(), "123.jpg", 1, true, 0)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("123321", res)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriver_List(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
AccessKey: "ak",
|
|
||||||
SecretKey: "sk",
|
|
||||||
BucketName: "test",
|
|
||||||
Server: "test.com",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
handler.Client, _ = NewClient(&model.Policy{})
|
|
||||||
handler.Client.Credential.AccessToken = "AccessToken"
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
|
|
||||||
// 非递归
|
|
||||||
{
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"GET",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"value":[{}]}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client.Request = clientMock
|
|
||||||
res, err := handler.List(context.Background(), "/", false)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(res, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 递归一次
|
|
||||||
{
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"GET",
|
|
||||||
"me/drive/root/children?$top=999999999",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"value":[{"name":"1","folder":{}}]}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"GET",
|
|
||||||
"me/drive/root:/1:/children?$top=999999999",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"value":[{"name":"2"}]}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client.Request = clientMock
|
|
||||||
res, err := handler.List(context.Background(), "/", true)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(res, 2)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriver_Thumb(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
AccessKey: "ak",
|
|
||||||
SecretKey: "sk",
|
|
||||||
BucketName: "test",
|
|
||||||
Server: "test.com",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
handler.Client, _ = NewClient(&model.Policy{})
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
file := &model.File{PicInfo: "1,1", Model: gorm.Model{ID: 1}}
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
ctx := context.WithValue(context.Background(), fsctx.ThumbSizeCtx, [2]uint{10, 20})
|
|
||||||
res, err := handler.Thumb(ctx, file)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Empty(res.URL)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 上下文错误
|
|
||||||
{
|
|
||||||
_, err := handler.Thumb(context.Background(), file)
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriver_Delete(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
AccessKey: "ak",
|
|
||||||
SecretKey: "sk",
|
|
||||||
BucketName: "test",
|
|
||||||
Server: "test.com",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
handler.Client, _ = NewClient(&model.Policy{})
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
_, err := handler.Delete(context.Background(), []string{"1"})
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriver_Put(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
AccessKey: "ak",
|
|
||||||
SecretKey: "sk",
|
|
||||||
BucketName: "test",
|
|
||||||
Server: "test.com",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
handler.Client, _ = NewClient(&model.Policy{})
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
err := handler.Put(context.Background(), &fsctx.FileStream{})
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriver_Get(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
AccessKey: "ak",
|
|
||||||
SecretKey: "sk",
|
|
||||||
BucketName: "test",
|
|
||||||
Server: "test.com",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
handler.Client, _ = NewClient(&model.Policy{})
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
|
|
||||||
// 无法获取source
|
|
||||||
{
|
|
||||||
res, err := handler.Get(context.Background(), "123.txt")
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Nil(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"GET",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"@microsoft.graph.downloadUrl":"123321"}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client.Request = clientMock
|
|
||||||
handler.Client.Credential.AccessToken = "1"
|
|
||||||
|
|
||||||
driverClientMock := ClientMock{}
|
|
||||||
driverClientMock.On(
|
|
||||||
"Request",
|
|
||||||
"GET",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`123`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.HTTPClient = driverClientMock
|
|
||||||
res, err := handler.Get(context.Background(), "123.txt")
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
asserts.NoError(err)
|
|
||||||
_, err = res.Seek(0, io.SeekEnd)
|
|
||||||
asserts.NoError(err)
|
|
||||||
content, err := ioutil.ReadAll(res)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Equal("123", string(content))
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriver_replaceSourceHost(t *testing.T) {
|
|
||||||
tests := []struct {
|
|
||||||
name string
|
|
||||||
origin string
|
|
||||||
cdn string
|
|
||||||
want string
|
|
||||||
wantErr bool
|
|
||||||
}{
|
|
||||||
{"TestNoReplace", "http://1dr.ms/download.aspx?123456", "", "http://1dr.ms/download.aspx?123456", false},
|
|
||||||
{"TestReplaceCorrect", "http://1dr.ms/download.aspx?123456", "https://test.com:8080", "https://test.com:8080/download.aspx?123456", false},
|
|
||||||
{"TestCdnFormatError", "http://1dr.ms/download.aspx?123456", string([]byte{0x7f}), "", true},
|
|
||||||
{"TestSrcFormatError", string([]byte{0x7f}), "https://test.com:8080", "", true},
|
|
||||||
}
|
|
||||||
for _, tt := range tests {
|
|
||||||
t.Run(tt.name, func(t *testing.T) {
|
|
||||||
policy := &model.Policy{}
|
|
||||||
policy.OptionsSerialized.OdProxy = tt.cdn
|
|
||||||
handler := Driver{
|
|
||||||
Policy: policy,
|
|
||||||
}
|
|
||||||
got, err := handler.replaceSourceHost(tt.origin)
|
|
||||||
if (err != nil) != tt.wantErr {
|
|
||||||
t.Errorf("replaceSourceHost() error = %v, wantErr %v", err, tt.wantErr)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if got != tt.want {
|
|
||||||
t.Errorf("replaceSourceHost() got = %v, want %v", got, tt.want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriver_CancelToken(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
AccessKey: "ak",
|
|
||||||
SecretKey: "sk",
|
|
||||||
BucketName: "test",
|
|
||||||
Server: "test.com",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
handler.Client, _ = NewClient(&model.Policy{})
|
|
||||||
handler.Client.Credential.ExpiresIn = time.Now().Add(time.Duration(100) * time.Hour).Unix()
|
|
||||||
|
|
||||||
// 失败
|
|
||||||
{
|
|
||||||
err := handler.CancelToken(context.Background(), &serializer.UploadSession{})
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -0,0 +1,25 @@
|
|||||||
|
package onedrive
|
||||||
|
|
||||||
|
import "sync"
|
||||||
|
|
||||||
|
// CredentialLock 针对存储策略凭证的锁
|
||||||
|
type CredentialLock interface {
|
||||||
|
Lock(uint)
|
||||||
|
Unlock(uint)
|
||||||
|
}
|
||||||
|
|
||||||
|
var GlobalMutex = mutexMap{}
|
||||||
|
|
||||||
|
type mutexMap struct {
|
||||||
|
locks sync.Map
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mutexMap) Lock(id uint) {
|
||||||
|
lock, _ := m.locks.LoadOrStore(id, &sync.Mutex{})
|
||||||
|
lock.(*sync.Mutex).Lock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mutexMap) Unlock(id uint) {
|
||||||
|
lock, _ := m.locks.LoadOrStore(id, &sync.Mutex{})
|
||||||
|
lock.(*sync.Mutex).Unlock()
|
||||||
|
}
|
||||||
@ -1,262 +0,0 @@
|
|||||||
package remote
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/filesystem/fsctx"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mocks/requestmock"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/request"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
testMock "github.com/stretchr/testify/mock"
|
|
||||||
"io/ioutil"
|
|
||||||
"net/http"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestNewClient(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
policy := &model.Policy{}
|
|
||||||
|
|
||||||
// 无法解析服务端url
|
|
||||||
{
|
|
||||||
policy.Server = string([]byte{0x7f})
|
|
||||||
c, err := NewClient(policy)
|
|
||||||
a.Error(err)
|
|
||||||
a.Nil(c)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
policy.Server = ""
|
|
||||||
c, err := NewClient(policy)
|
|
||||||
a.NoError(err)
|
|
||||||
a.NotNil(c)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoteClient_Upload(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
c, _ := NewClient(&model.Policy{})
|
|
||||||
|
|
||||||
// 无法创建上传会话
|
|
||||||
{
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
c.(*remoteClient).httpClient = &clientMock
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"PUT",
|
|
||||||
"upload",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: errors.New("error"),
|
|
||||||
})
|
|
||||||
err := c.Upload(context.Background(), &fsctx.FileStream{})
|
|
||||||
a.Error(err)
|
|
||||||
a.Contains(err.Error(), "error")
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 分片上传失败,成功删除上传会话
|
|
||||||
{
|
|
||||||
cache.Set("setting_chunk_retries", "1", 0)
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
c.(*remoteClient).httpClient = &clientMock
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"PUT",
|
|
||||||
"upload",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: errors.New("error"),
|
|
||||||
})
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"DELETE",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
err := c.Upload(context.Background(), &fsctx.FileStream{})
|
|
||||||
a.Error(err)
|
|
||||||
a.Contains(err.Error(), "error")
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 分片上传失败,无法删除上传会话
|
|
||||||
{
|
|
||||||
cache.Set("setting_chunk_retries", "1", 0)
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
c.(*remoteClient).httpClient = &clientMock
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"PUT",
|
|
||||||
"upload",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: errors.New("error"),
|
|
||||||
})
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"DELETE",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: errors.New("error2"),
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
err := c.Upload(context.Background(), &fsctx.FileStream{})
|
|
||||||
a.Error(err)
|
|
||||||
a.Contains(err.Error(), "error")
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
cache.Set("setting_chunk_retries", "1", 0)
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
c.(*remoteClient).httpClient = &clientMock
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"PUT",
|
|
||||||
"upload",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
err := c.Upload(context.Background(), &fsctx.FileStream{})
|
|
||||||
a.NoError(err)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoteClient_CreateUploadSessionFailed(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
c, _ := NewClient(&model.Policy{})
|
|
||||||
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
c.(*remoteClient).httpClient = &clientMock
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"PUT",
|
|
||||||
"upload",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":500,"msg":"error"}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
err := c.Upload(context.Background(), &fsctx.FileStream{})
|
|
||||||
a.Error(err)
|
|
||||||
a.Contains(err.Error(), "error")
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoteClient_UploadChunkFailed(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
c, _ := NewClient(&model.Policy{})
|
|
||||||
|
|
||||||
clientMock := requestmock.RequestMock{}
|
|
||||||
c.(*remoteClient).httpClient = &clientMock
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":500,"msg":"error"}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
err := c.(*remoteClient).uploadChunk(context.Background(), "", 0, strings.NewReader(""), false, 0)
|
|
||||||
a.Error(err)
|
|
||||||
a.Contains(err.Error(), "error")
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRemoteClient_GetUploadURL(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
c, _ := NewClient(&model.Policy{})
|
|
||||||
|
|
||||||
// url 解析失败
|
|
||||||
{
|
|
||||||
c.(*remoteClient).policy.Server = string([]byte{0x7f})
|
|
||||||
res, sign, err := c.GetUploadURL(0, "")
|
|
||||||
a.Error(err)
|
|
||||||
a.Empty(res)
|
|
||||||
a.Empty(sign)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
c.(*remoteClient).policy.Server = ""
|
|
||||||
res, sign, err := c.GetUploadURL(0, "")
|
|
||||||
a.NoError(err)
|
|
||||||
a.NotEmpty(res)
|
|
||||||
a.NotEmpty(sign)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@ -1,460 +0,0 @@
|
|||||||
package remote
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"errors"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/filesystem/driver"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/mocks/remoteclientmock"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/serializer"
|
|
||||||
"io"
|
|
||||||
"io/ioutil"
|
|
||||||
"net/http"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
model "github.com/cloudreve/Cloudreve/v3/models"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/auth"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/cache"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/filesystem/fsctx"
|
|
||||||
"github.com/cloudreve/Cloudreve/v3/pkg/request"
|
|
||||||
"github.com/stretchr/testify/assert"
|
|
||||||
testMock "github.com/stretchr/testify/mock"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestNewDriver(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
|
|
||||||
// remoteClient 初始化失败
|
|
||||||
{
|
|
||||||
d, err := NewDriver(&model.Policy{Server: string([]byte{0x7f})})
|
|
||||||
a.Error(err)
|
|
||||||
a.Nil(d)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
d, err := NewDriver(&model.Policy{})
|
|
||||||
a.NoError(err)
|
|
||||||
a.NotNil(d)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandler_Source(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
auth.General = auth.HMACAuth{SecretKey: []byte("test")}
|
|
||||||
|
|
||||||
// 无法获取上下文
|
|
||||||
{
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{Server: "/"},
|
|
||||||
AuthInstance: auth.HMACAuth{},
|
|
||||||
}
|
|
||||||
ctx := context.Background()
|
|
||||||
res, err := handler.Source(ctx, "", 0, true, 0)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.NotEmpty(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{Server: "/"},
|
|
||||||
AuthInstance: auth.HMACAuth{},
|
|
||||||
}
|
|
||||||
file := model.File{
|
|
||||||
SourceName: "1.txt",
|
|
||||||
}
|
|
||||||
ctx := context.WithValue(context.Background(), fsctx.FileModelCtx, file)
|
|
||||||
res, err := handler.Source(ctx, "", 10, true, 0)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Contains(res, "api/v3/slave/download/0")
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功 自定义CDN
|
|
||||||
{
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{Server: "/", BaseURL: "https://cqu.edu.cn"},
|
|
||||||
AuthInstance: auth.HMACAuth{},
|
|
||||||
}
|
|
||||||
file := model.File{
|
|
||||||
SourceName: "1.txt",
|
|
||||||
}
|
|
||||||
ctx := context.WithValue(context.Background(), fsctx.FileModelCtx, file)
|
|
||||||
res, err := handler.Source(ctx, "", 10, true, 0)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Contains(res, "api/v3/slave/download/0")
|
|
||||||
asserts.Contains(res, "https://cqu.edu.cn")
|
|
||||||
}
|
|
||||||
|
|
||||||
// 解析失败 自定义CDN
|
|
||||||
{
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{Server: "/", BaseURL: string([]byte{0x7f})},
|
|
||||||
AuthInstance: auth.HMACAuth{},
|
|
||||||
}
|
|
||||||
file := model.File{
|
|
||||||
SourceName: "1.txt",
|
|
||||||
}
|
|
||||||
ctx := context.WithValue(context.Background(), fsctx.FileModelCtx, file)
|
|
||||||
res, err := handler.Source(ctx, "", 10, true, 0)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Empty(res)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功 预览
|
|
||||||
{
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{Server: "/"},
|
|
||||||
AuthInstance: auth.HMACAuth{},
|
|
||||||
}
|
|
||||||
file := model.File{
|
|
||||||
SourceName: "1.txt",
|
|
||||||
}
|
|
||||||
ctx := context.WithValue(context.Background(), fsctx.FileModelCtx, file)
|
|
||||||
res, err := handler.Source(ctx, "", 10, false, 0)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Contains(res, "api/v3/slave/source/0")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
type ClientMock struct {
|
|
||||||
testMock.Mock
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m ClientMock) Request(method, target string, body io.Reader, opts ...request.Option) *request.Response {
|
|
||||||
args := m.Called(method, target, body, opts)
|
|
||||||
return args.Get(0).(*request.Response)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandler_Delete(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
SecretKey: "test",
|
|
||||||
Server: "http://test.com",
|
|
||||||
},
|
|
||||||
AuthInstance: auth.HMACAuth{},
|
|
||||||
}
|
|
||||||
ctx := context.Background()
|
|
||||||
cache.Set("setting_slave_api_timeout", "60", 0)
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test.com/api/v3/slave/delete",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client = clientMock
|
|
||||||
failed, err := handler.Delete(ctx, []string{"/test1.txt", "test2.txt"})
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(failed, 0)
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
// 结果解析失败
|
|
||||||
{
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test.com/api/v3/slave/delete",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":203}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client = clientMock
|
|
||||||
failed, err := handler.Delete(ctx, []string{"/test1.txt", "test2.txt"})
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Len(failed, 2)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 一个失败
|
|
||||||
{
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test.com/api/v3/slave/delete",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":203,"data":"{\"files\":[\"1\"]}"}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client = clientMock
|
|
||||||
failed, err := handler.Delete(ctx, []string{"/test1.txt", "test2.txt"})
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Len(failed, 1)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriver_List(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
SecretKey: "test",
|
|
||||||
Server: "http://test.com",
|
|
||||||
},
|
|
||||||
AuthInstance: auth.HMACAuth{},
|
|
||||||
}
|
|
||||||
ctx := context.Background()
|
|
||||||
cache.Set("setting_slave_api_timeout", "60", 0)
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test.com/api/v3/slave/list",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0,"data":"[{}]"}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client = clientMock
|
|
||||||
res, err := handler.List(ctx, "/", true)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.Len(res, 1)
|
|
||||||
|
|
||||||
}
|
|
||||||
|
|
||||||
// 响应解析失败
|
|
||||||
{
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test.com/api/v3/slave/list",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0,"data":"233"}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client = clientMock
|
|
||||||
res, err := handler.List(ctx, "/", true)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Len(res, 0)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 从机返回错误
|
|
||||||
{
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"POST",
|
|
||||||
"http://test.com/api/v3/slave/list",
|
|
||||||
testMock.Anything,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":203}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client = clientMock
|
|
||||||
res, err := handler.List(ctx, "/", true)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
asserts.Error(err)
|
|
||||||
asserts.Len(res, 0)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandler_Get(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
SecretKey: "test",
|
|
||||||
Server: "http://test.com",
|
|
||||||
},
|
|
||||||
AuthInstance: auth.HMACAuth{},
|
|
||||||
}
|
|
||||||
ctx := context.Background()
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
ctx = context.WithValue(ctx, fsctx.UserCtx, model.User{})
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"GET",
|
|
||||||
testMock.Anything,
|
|
||||||
nil,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 200,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client = clientMock
|
|
||||||
resp, err := handler.Get(ctx, "/test.txt")
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
asserts.NotNil(resp)
|
|
||||||
asserts.NoError(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 请求失败
|
|
||||||
{
|
|
||||||
ctx = context.WithValue(ctx, fsctx.UserCtx, model.User{})
|
|
||||||
clientMock := ClientMock{}
|
|
||||||
clientMock.On(
|
|
||||||
"Request",
|
|
||||||
"GET",
|
|
||||||
testMock.Anything,
|
|
||||||
nil,
|
|
||||||
testMock.Anything,
|
|
||||||
).Return(&request.Response{
|
|
||||||
Err: nil,
|
|
||||||
Response: &http.Response{
|
|
||||||
StatusCode: 404,
|
|
||||||
Body: ioutil.NopCloser(strings.NewReader(`{"code":0}`)),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
handler.Client = clientMock
|
|
||||||
resp, err := handler.Get(ctx, "/test.txt")
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
asserts.Nil(resp)
|
|
||||||
asserts.Error(err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandler_Put(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
handler, _ := NewDriver(&model.Policy{
|
|
||||||
Type: "remote",
|
|
||||||
SecretKey: "test",
|
|
||||||
Server: "http://test.com",
|
|
||||||
})
|
|
||||||
clientMock := &remoteclientmock.RemoteClientMock{}
|
|
||||||
handler.uploadClient = clientMock
|
|
||||||
clientMock.On("Upload", testMock.Anything, testMock.Anything).Return(errors.New("error"))
|
|
||||||
a.Error(handler.Put(context.Background(), &fsctx.FileStream{}))
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandler_Thumb(t *testing.T) {
|
|
||||||
asserts := assert.New(t)
|
|
||||||
handler := Driver{
|
|
||||||
Policy: &model.Policy{
|
|
||||||
Type: "remote",
|
|
||||||
SecretKey: "test",
|
|
||||||
Server: "http://test.com",
|
|
||||||
OptionsSerialized: model.PolicyOption{
|
|
||||||
ThumbExts: []string{"txt"},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
AuthInstance: auth.HMACAuth{},
|
|
||||||
}
|
|
||||||
file := &model.File{
|
|
||||||
Name: "1.txt",
|
|
||||||
SourceName: "1.txt",
|
|
||||||
}
|
|
||||||
ctx := context.Background()
|
|
||||||
asserts.NoError(cache.Set("setting_preview_timeout", "60", 0))
|
|
||||||
|
|
||||||
// no error
|
|
||||||
{
|
|
||||||
resp, err := handler.Thumb(ctx, file)
|
|
||||||
asserts.NoError(err)
|
|
||||||
asserts.True(resp.Redirect)
|
|
||||||
}
|
|
||||||
|
|
||||||
// ext not support
|
|
||||||
{
|
|
||||||
file.Name = "1.jpg"
|
|
||||||
resp, err := handler.Thumb(ctx, file)
|
|
||||||
asserts.ErrorIs(err, driver.ErrorThumbNotSupported)
|
|
||||||
asserts.Nil(resp)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestHandler_Token(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
handler, _ := NewDriver(&model.Policy{})
|
|
||||||
|
|
||||||
// 无法创建上传会话
|
|
||||||
{
|
|
||||||
clientMock := &remoteclientmock.RemoteClientMock{}
|
|
||||||
handler.uploadClient = clientMock
|
|
||||||
clientMock.On("CreateUploadSession", testMock.Anything, testMock.Anything, int64(10), false).Return(errors.New("error"))
|
|
||||||
res, err := handler.Token(context.Background(), 10, &serializer.UploadSession{}, &fsctx.FileStream{})
|
|
||||||
a.Error(err)
|
|
||||||
a.Contains(err.Error(), "error")
|
|
||||||
a.Nil(res)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 无法创建上传地址
|
|
||||||
{
|
|
||||||
clientMock := &remoteclientmock.RemoteClientMock{}
|
|
||||||
handler.uploadClient = clientMock
|
|
||||||
clientMock.On("CreateUploadSession", testMock.Anything, testMock.Anything, int64(10), false).Return(nil)
|
|
||||||
clientMock.On("GetUploadURL", int64(10), "").Return("", "", errors.New("error"))
|
|
||||||
res, err := handler.Token(context.Background(), 10, &serializer.UploadSession{}, &fsctx.FileStream{})
|
|
||||||
a.Error(err)
|
|
||||||
a.Contains(err.Error(), "error")
|
|
||||||
a.Nil(res)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 成功
|
|
||||||
{
|
|
||||||
clientMock := &remoteclientmock.RemoteClientMock{}
|
|
||||||
handler.uploadClient = clientMock
|
|
||||||
clientMock.On("CreateUploadSession", testMock.Anything, testMock.Anything, int64(10), false).Return(nil)
|
|
||||||
clientMock.On("GetUploadURL", int64(10), "").Return("1", "2", nil)
|
|
||||||
res, err := handler.Token(context.Background(), 10, &serializer.UploadSession{}, &fsctx.FileStream{})
|
|
||||||
a.NoError(err)
|
|
||||||
a.NotNil(res)
|
|
||||||
a.Equal("1", res.UploadURLs[0])
|
|
||||||
a.Equal("2", res.Credential)
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDriver_CancelToken(t *testing.T) {
|
|
||||||
a := assert.New(t)
|
|
||||||
handler, _ := NewDriver(&model.Policy{})
|
|
||||||
|
|
||||||
clientMock := &remoteclientmock.RemoteClientMock{}
|
|
||||||
handler.uploadClient = clientMock
|
|
||||||
clientMock.On("DeleteUploadSession", testMock.Anything, "key").Return(errors.New("error"))
|
|
||||||
err := handler.CancelToken(context.Background(), &serializer.UploadSession{Key: "key"})
|
|
||||||
a.Error(err)
|
|
||||||
a.Contains(err.Error(), "error")
|
|
||||||
clientMock.AssertExpectations(t)
|
|
||||||
}
|
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Loading…
Reference in new issue