mirror of
https://github.com/OpenListTeam/OpenList.git
synced 2026-10-10 21:13:10 +08:00
Compare commits
22 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 9e7da25e85 | |||
| 4c39bbe9c2 | |||
| e1d88b071f | |||
| ea10624fb6 | |||
| 54ae9d7451 | |||
| c16701b94b | |||
| 893457cd50 | |||
| 90acfa18e4 | |||
| 1462d63a48 | |||
| cadbf87246 | |||
| 9de3f69b8f | |||
| e73a80c78c | |||
| 6c6009109f | |||
| b51d1c8284 | |||
| f286862c61 | |||
| 40f4f6546f | |||
| 3a31b438a9 | |||
| 56064d1981 | |||
| d894a3983b | |||
| 084c008102 | |||
| 4580c4db33 | |||
| 5447ecb072 |
@@ -124,7 +124,7 @@ jobs:
|
||||
FRONTEND_REPO: ${{ vars.FRONTEND_REPO }}
|
||||
|
||||
- name: Build
|
||||
uses: OpenListTeam/cgo-actions@6fcace5934c36d70503dba06e2396bb58b766130 # v1.2.5
|
||||
uses: OpenListTeam/cgo-actions@d760a8ec1a6be1f8ec181229e11cb2671797f9e0 # v1.3.0
|
||||
with:
|
||||
targets: ${{ matrix.target }}
|
||||
flags: ${{ matrix.flags || '-ldflags=' }}
|
||||
@@ -136,7 +136,7 @@ jobs:
|
||||
musl-base-url: "https://github.com/OpenListTeam/musl-compilers/releases/latest/download/"
|
||||
x-flags: |
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.BuiltAt=$built_at
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.GitAuthor=The OpenList Projects Contributors <noreply@openlist.team>
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.GitAuthor=The OpenList Projects Contributors <noreply@oplist.org>
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.GitCommit=$git_commit
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.Version=$tag
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.WebVersion=rolling
|
||||
|
||||
@@ -42,7 +42,7 @@ jobs:
|
||||
FRONTEND_REPO: ${{ vars.FRONTEND_REPO }}
|
||||
|
||||
- name: Build
|
||||
uses: OpenListTeam/cgo-actions@6fcace5934c36d70503dba06e2396bb58b766130 # v1.2.5
|
||||
uses: OpenListTeam/cgo-actions@d760a8ec1a6be1f8ec181229e11cb2671797f9e0 # v1.3.0
|
||||
with:
|
||||
targets: ${{ matrix.target }}
|
||||
flags: ${{ contains(matrix.target, '-musl') && '-ldflags=-linkmode external -extldflags ''-static -fpic''' || '-ldflags=' }}
|
||||
@@ -52,7 +52,7 @@ jobs:
|
||||
out-dir: build
|
||||
x-flags: |
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.BuiltAt=$built_at
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.GitAuthor=The OpenList Projects Contributors <noreply@openlist.team>
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.GitAuthor=The OpenList Projects Contributors <noreply@oplist.org>
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.GitCommit=$git_commit
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.Version=$tag
|
||||
github.com/OpenListTeam/OpenList/v4/internal/conf.WebVersion=rolling
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
set -e
|
||||
appName="openlist"
|
||||
builtAt="$(date +'%F %T %z')"
|
||||
gitAuthor="The OpenList Projects Contributors <noreply@openlist.team>"
|
||||
gitAuthor="The OpenList Projects Contributors <noreply@oplist.org>"
|
||||
gitCommit=$(git log --pretty=format:"%h" -1)
|
||||
|
||||
# Set frontend repository, default to OpenListTeam/OpenList-Frontend
|
||||
@@ -531,8 +531,8 @@ BuildReleaseFreeBSD() {
|
||||
sed 's/\.0$//')
|
||||
|
||||
if [ -z "$freebsd_version" ]; then
|
||||
echo "Failed to get FreeBSD version, falling back to 14.3"
|
||||
freebsd_version="14.3"
|
||||
echo "Failed to get FreeBSD version, falling back to 14.4"
|
||||
freebsd_version="14.4"
|
||||
fi
|
||||
|
||||
echo "Using FreeBSD version: $freebsd_version"
|
||||
|
||||
+61
-16
@@ -3,6 +3,7 @@ package _115_open
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"io"
|
||||
"time"
|
||||
|
||||
@@ -70,6 +71,19 @@ func (d *Open115) singleUpload(ctx context.Context, tempF model.File, tokenResp
|
||||
// } `json:"data"`
|
||||
// }
|
||||
|
||||
// retryExpiredToken retries only the rejected OSS operation, preserving the upload ID.
|
||||
func retryExpiredToken(refresh func() error, operation func() error) error {
|
||||
err := operation()
|
||||
var serviceErr oss.ServiceError
|
||||
if !errors.As(err, &serviceErr) || serviceErr.Code != "SecurityTokenExpired" {
|
||||
return err
|
||||
}
|
||||
if err := refresh(); err != nil {
|
||||
return err
|
||||
}
|
||||
return operation()
|
||||
}
|
||||
|
||||
func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer, up driver.UpdateProgress, tokenResp *sdk.UploadGetTokenResp, initResp *sdk.UploadInitResp) error {
|
||||
ossClient, err := netutil.NewOSSClient(tokenResp.Endpoint, tokenResp.AccessKeyId, tokenResp.AccessKeySecret, oss.SecurityToken(tokenResp.SecurityToken))
|
||||
if err != nil {
|
||||
@@ -80,7 +94,32 @@ func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer,
|
||||
return err
|
||||
}
|
||||
|
||||
imur, err := bucket.InitiateMultipartUpload(initResp.Object, oss.Sequential())
|
||||
refresh := func() error {
|
||||
if err := d.WaitLimit(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
token, err := d.client.UploadGetToken(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
client, err := netutil.NewOSSClient(token.Endpoint, token.AccessKeyId, token.AccessKeySecret, oss.SecurityToken(token.SecurityToken))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
newBucket, err := client.Bucket(initResp.Bucket)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
bucket = newBucket
|
||||
return nil
|
||||
}
|
||||
|
||||
var imur oss.InitiateMultipartUploadResult
|
||||
err = retryExpiredToken(refresh, func() error {
|
||||
var err error
|
||||
imur, err = bucket.InitiateMultipartUpload(initResp.Object, oss.Sequential(), oss.WithContext(ctx))
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -109,13 +148,17 @@ func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer,
|
||||
return err
|
||||
}
|
||||
err = retry.Do(func() error {
|
||||
rd.Seek(0, io.SeekStart)
|
||||
part, err := bucket.UploadPart(imur, driver.NewLimitedUploadStream(ctx, rd), partSize, int(i))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
parts[i-1] = part
|
||||
return nil
|
||||
return retryExpiredToken(refresh, func() error {
|
||||
if _, err := rd.Seek(0, io.SeekStart); err != nil {
|
||||
return err
|
||||
}
|
||||
part, err := bucket.UploadPart(imur, driver.NewLimitedUploadStream(ctx, rd), partSize, int(i), oss.WithContext(ctx))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
parts[i-1] = part
|
||||
return nil
|
||||
})
|
||||
},
|
||||
retry.Context(ctx),
|
||||
retry.Attempts(3),
|
||||
@@ -134,14 +177,16 @@ func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer,
|
||||
up(float64(offset) * 100 / float64(fileSize))
|
||||
}
|
||||
|
||||
// callbackRespBytes := make([]byte, 1024)
|
||||
_, err = bucket.CompleteMultipartUpload(
|
||||
imur,
|
||||
parts,
|
||||
oss.Callback(base64.StdEncoding.EncodeToString([]byte(initResp.Callback.Value.Callback))),
|
||||
oss.CallbackVar(base64.StdEncoding.EncodeToString([]byte(initResp.Callback.Value.CallbackVar))),
|
||||
// oss.CallbackResult(&callbackRespBytes),
|
||||
)
|
||||
err = retryExpiredToken(refresh, func() error {
|
||||
_, err := bucket.CompleteMultipartUpload(
|
||||
imur,
|
||||
parts,
|
||||
oss.Callback(base64.StdEncoding.EncodeToString([]byte(initResp.Callback.Value.Callback))),
|
||||
oss.CallbackVar(base64.StdEncoding.EncodeToString([]byte(initResp.Callback.Value.CallbackVar))),
|
||||
oss.WithContext(ctx),
|
||||
)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -7,11 +7,16 @@ import (
|
||||
|
||||
type Addition struct {
|
||||
LoginType string `json:"login_type" type:"select" options:"password,qrcode" default:"password" required:"true"`
|
||||
Username string `json:"username" required:"true"`
|
||||
Password string `json:"password" required:"true"`
|
||||
Username string `json:"username" help:"Not needed when an access token or refresh token is provided"`
|
||||
Password string `json:"password" help:"Not needed when an access token or refresh token is provided"`
|
||||
VCode string `json:"validate_code"`
|
||||
SmsCode string `json:"sms_code" help:"SMS code for the second device verification, fill it in and save again when login asks for it"`
|
||||
AccessToken string `json:"access_token" required:"false"`
|
||||
RefreshToken string `json:"refresh_token" help:"To switch accounts, please clear this field"`
|
||||
DeviceID string `json:"device_id" help:"DEVICEID cookie issued after the second device verification, keep it to avoid verifying again"`
|
||||
ClientSn string `json:"client_sn" help:"Device serial number captured from the official client, leave it empty if you do not have one"`
|
||||
JgOpenId string `json:"jg_open_id" help:"Optional push id reported by the official client"`
|
||||
UserFinger string `json:"user_finger" help:"Device fingerprint sent with login requests, generated and kept automatically when empty"`
|
||||
driver.RootID
|
||||
OrderBy string `json:"order_by" type:"select" options:"filename,filesize,lastOpTime" default:"filename"`
|
||||
OrderDirection string `json:"order_direction" type:"select" options:"asc,desc" default:"asc"`
|
||||
|
||||
@@ -72,6 +72,8 @@ type BaseLoginParam struct {
|
||||
// 请求头参数
|
||||
Lt string
|
||||
ReqId string
|
||||
// logbox页面地址,作为后续请求的Referer,缺失会被判定为陌生设备
|
||||
Referer string
|
||||
|
||||
// 表单参数
|
||||
ParamId string
|
||||
@@ -97,10 +99,20 @@ type LoginParam struct {
|
||||
|
||||
// rsa密钥
|
||||
jRsaKey string
|
||||
// 加密字段的前缀,服务端下发(如 {NRP})
|
||||
rsaPrefix string
|
||||
|
||||
// 设备二次校验时服务端返回的加密手机号
|
||||
SecondAuthMobile string
|
||||
|
||||
BaseLoginParam
|
||||
}
|
||||
|
||||
// encryptSecret 用登陆时拿到的公钥加密敏感值,格式与userName/epd一致
|
||||
func (p *LoginParam) encryptSecret(value string) string {
|
||||
return p.rsaPrefix + RsaEncrypt(p.jRsaKey, value)
|
||||
}
|
||||
|
||||
// 登陆加密相关
|
||||
type EncryptConfResp struct {
|
||||
Result int `json:"result"`
|
||||
@@ -116,6 +128,35 @@ type LoginResp struct {
|
||||
Msg string `json:"msg"`
|
||||
Result int `json:"result"`
|
||||
ToUrl string `json:"toUrl"`
|
||||
// 设备二次校验时返回的加密手机号
|
||||
Mobile string `json:"mobile"`
|
||||
}
|
||||
|
||||
// 登陆页配置,新版登陆页的paramId由该接口下发
|
||||
// 该接口的result可能是数字也可能是字符串
|
||||
type AppConfResp struct {
|
||||
Result any `json:"result"`
|
||||
Msg string `json:"msg"`
|
||||
Data struct {
|
||||
ParamId string `json:"paramId"`
|
||||
AccountType string `json:"accountType"`
|
||||
ReturnUrl string `json:"returnUrl"`
|
||||
MailSuffix string `json:"mailSuffix"`
|
||||
} `json:"data"`
|
||||
}
|
||||
|
||||
func (r *AppConfResp) Succeeded() bool {
|
||||
switch v := r.Result.(type) {
|
||||
case nil:
|
||||
return true
|
||||
case string:
|
||||
return v == "0" || v == ""
|
||||
case float64:
|
||||
return v == 0
|
||||
case int:
|
||||
return v == 0
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// 刷新session返回
|
||||
@@ -149,6 +190,27 @@ type AppSessionResp struct {
|
||||
RefreshToken string `json:"refreshToken"`
|
||||
}
|
||||
|
||||
// 刷新token返回,失败时以HTTP 200返回result/msg,需要单独判断
|
||||
type RefreshTokenResp struct {
|
||||
AccessToken string `json:"accessToken"`
|
||||
RefreshToken string `json:"refreshToken"`
|
||||
ExpiresIn int `json:"expiresIn"`
|
||||
|
||||
Result int `json:"result"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
|
||||
func (r *RefreshTokenResp) HasError() bool {
|
||||
return r.Result != 0 || r.AccessToken == ""
|
||||
}
|
||||
|
||||
func (r *RefreshTokenResp) Error() string {
|
||||
if r.Msg != "" {
|
||||
return fmt.Sprintf("refresh token failed, result: %d, msg: %s", r.Result, r.Msg)
|
||||
}
|
||||
return fmt.Sprintf("refresh token failed, result: %d", r.Result)
|
||||
}
|
||||
|
||||
// 家庭云账户
|
||||
type FamilyInfoListResp struct {
|
||||
FamilyInfoResp []FamilyInfoResp `json:"familyInfoResp"`
|
||||
|
||||
+355
-77
@@ -30,6 +30,7 @@ import (
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/errgroup"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils/random"
|
||||
"github.com/skip2/go-qrcode"
|
||||
|
||||
"github.com/avast/retry-go"
|
||||
@@ -41,9 +42,13 @@ import (
|
||||
|
||||
const (
|
||||
ACCOUNT_TYPE = "02"
|
||||
APP_ID = "8025431004"
|
||||
CLIENT_TYPE = "10020"
|
||||
VERSION = "6.2"
|
||||
// 官方 PC 端(cloud.189.cn 网页/客户端)使用的 appId,
|
||||
// 登录、生成二维码、换取 session 必须全程使用同一个 appId
|
||||
APP_ID = "9317140619"
|
||||
CLIENT_TYPE = "10020"
|
||||
// 扫码状态轮询使用的 clientType,与密码登录的 10020 不同
|
||||
QR_CLIENT_TYPE = "1"
|
||||
VERSION = "7.2.4.0"
|
||||
|
||||
WEB_URL = "https://cloud.189.cn"
|
||||
AUTH_URL = "https://open.e.189.cn"
|
||||
@@ -57,8 +62,18 @@ const (
|
||||
|
||||
CHANNEL_ID = "web_cloud.189.cn"
|
||||
|
||||
// 服务端通过短信二次校验后下发的设备标识,复用它可以避免再次触发校验
|
||||
DEVICE_ID_COOKIE = "DEVICEID"
|
||||
|
||||
// 扫码登录本地轮询参数,超时后把二维码交回前端,避免请求被反向代理掐断
|
||||
QRCODE_POLL_INTERVAL = 2 * time.Second
|
||||
QRCODE_POLL_TIMEOUT = 20 * time.Second
|
||||
|
||||
// Error codes
|
||||
UserInvalidOpenTokenError = "UserInvalidOpenToken"
|
||||
|
||||
// 密码登录返回该结果表示需要设备二次校验
|
||||
SecondDeviceAuthResult = -133
|
||||
)
|
||||
|
||||
func (y *Cloud189PC) SignatureHeader(url, method, params string, isFamily bool) map[string]string {
|
||||
@@ -288,9 +303,72 @@ func (y *Cloud189PC) login() error {
|
||||
if y.LoginType == "qrcode" {
|
||||
return y.loginByQRCode()
|
||||
}
|
||||
if y.Username == "" || y.Password == "" {
|
||||
return errors.New("please fill in the username and password, or provide an access token / refresh token")
|
||||
}
|
||||
return y.loginByPassword()
|
||||
}
|
||||
|
||||
// 设备指纹,为空时生成并保存,服务端以此识别是否为同一台设备
|
||||
func (y *Cloud189PC) getUserFinger() string {
|
||||
if y.Addition.UserFinger == "" {
|
||||
y.Addition.UserFinger = fmt.Sprint(random.Rand.Int63n(9e9) + 1e9)
|
||||
op.MustSaveDriverStorage(y)
|
||||
}
|
||||
return y.Addition.UserFinger
|
||||
}
|
||||
|
||||
// 换取会话时携带的设备参数,与官方PC客户端保持一致
|
||||
// clientSn/jgOpenId 只在用户从官方客户端抓到并填写后才发送,避免上报一个服务端不认识的设备号
|
||||
func (y *Cloud189PC) deviceParams() map[string]string {
|
||||
params := map[string]string{"returnType": "JSON"}
|
||||
if y.Addition.ClientSn != "" {
|
||||
params["clientSn"] = y.Addition.ClientSn
|
||||
}
|
||||
if y.Addition.JgOpenId != "" {
|
||||
params["jgOpenId"] = y.Addition.JgOpenId
|
||||
}
|
||||
return params
|
||||
}
|
||||
|
||||
// logbox接口的公共请求头,缺少user-finger和Referer会被判定为陌生设备
|
||||
func (y *Cloud189PC) loginHeaders(param BaseLoginParam) map[string]string {
|
||||
return map[string]string{
|
||||
"REQID": param.ReqId,
|
||||
"lt": param.Lt,
|
||||
"user-finger": y.getUserFinger(),
|
||||
"Referer": IF(param.Referer != "", param.Referer, AUTH_URL),
|
||||
}
|
||||
}
|
||||
|
||||
// 把已保存的设备标识写入cookie,避免重复触发设备二次校验
|
||||
func (y *Cloud189PC) applyDeviceID(jar http.CookieJar) {
|
||||
if y.Addition.DeviceID == "" {
|
||||
return
|
||||
}
|
||||
authUrl, err := url.Parse(AUTH_URL)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
jar.SetCookies(authUrl, []*http.Cookie{{
|
||||
Name: DEVICE_ID_COOKIE,
|
||||
Value: y.Addition.DeviceID,
|
||||
Domain: "e.189.cn",
|
||||
Path: "/",
|
||||
}})
|
||||
}
|
||||
|
||||
// 保存服务端下发的设备标识,下次登陆复用即可跳过设备二次校验
|
||||
func (y *Cloud189PC) saveDeviceID(res *resty.Response) {
|
||||
for _, cookie := range res.Cookies() {
|
||||
if cookie.Name == DEVICE_ID_COOKIE && cookie.Value != "" && cookie.Value != y.Addition.DeviceID {
|
||||
y.Addition.DeviceID = cookie.Value
|
||||
op.MustSaveDriverStorage(y)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (y *Cloud189PC) loginByPassword() (err error) {
|
||||
// 初始化登陆所需参数
|
||||
if y.loginParam == nil {
|
||||
@@ -299,9 +377,16 @@ func (y *Cloud189PC) loginByPassword() (err error) {
|
||||
return err
|
||||
}
|
||||
}
|
||||
// 设备二次校验必须复用同一套登陆参数,此时不能销毁也不能重新初始化
|
||||
keepLoginParam := false
|
||||
defer func() {
|
||||
// 销毁验证码
|
||||
y.VCode = ""
|
||||
if keepLoginParam {
|
||||
y.Status = err.Error()
|
||||
op.MustSaveDriverStorage(y)
|
||||
return
|
||||
}
|
||||
// 销毁登陆参数
|
||||
y.loginParam = nil
|
||||
// 遇到错误,重新加载登陆参数(刷新验证码)
|
||||
@@ -319,17 +404,18 @@ func (y *Cloud189PC) loginByPassword() (err error) {
|
||||
|
||||
param := y.loginParam
|
||||
var loginresp LoginResp
|
||||
_, err = y.client.R().
|
||||
res, err := y.client.R().
|
||||
ForceContentType("application/json;charset=UTF-8").SetResult(&loginresp).
|
||||
SetHeaders(map[string]string{
|
||||
"REQID": param.ReqId,
|
||||
"lt": param.Lt,
|
||||
}).
|
||||
SetHeaders(y.loginHeaders(param.BaseLoginParam)).
|
||||
SetFormData(map[string]string{
|
||||
"version": "v2.0",
|
||||
"apToken": "",
|
||||
"appKey": APP_ID,
|
||||
"pageKey": "normal",
|
||||
"accountType": ACCOUNT_TYPE,
|
||||
"userName": param.RsaUsername,
|
||||
"password": param.RsaPassword,
|
||||
"epd": param.RsaPassword,
|
||||
"validateCode": y.VCode,
|
||||
"captchaToken": param.CaptchaToken,
|
||||
"returnUrl": RETURN_URL,
|
||||
@@ -345,17 +431,106 @@ func (y *Cloud189PC) loginByPassword() (err error) {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
y.saveDeviceID(res)
|
||||
|
||||
// 设备二次校验:服务端要求短信验证,保留登陆参数并引导填写短信验证码
|
||||
if loginresp.Result == SecondDeviceAuthResult {
|
||||
err = y.secondDeviceAuth(loginresp.Mobile)
|
||||
// 校验未完成时保留登陆参数,等待用户回填短信验证码
|
||||
keepLoginParam = err != nil && y.loginParam != nil
|
||||
return err
|
||||
}
|
||||
|
||||
if loginresp.ToUrl == "" {
|
||||
return fmt.Errorf("login failed,No toUrl obtained, msg: %s", loginresp.Msg)
|
||||
}
|
||||
|
||||
// 获取Session
|
||||
return y.getSessionByRedirectURL(loginresp.ToUrl)
|
||||
}
|
||||
|
||||
// 设备二次校验:先发短信,用户回填验证码后再提交
|
||||
func (y *Cloud189PC) secondDeviceAuth(mobile string) error {
|
||||
param := y.loginParam
|
||||
if mobile != "" {
|
||||
param.SecondAuthMobile = mobile
|
||||
}
|
||||
if param.SecondAuthMobile == "" {
|
||||
return errors.New("second device verification is required, but no mobile was returned")
|
||||
}
|
||||
|
||||
// 已填写短信验证码,直接提交校验
|
||||
if y.SmsCode != "" {
|
||||
smsCode := y.SmsCode
|
||||
y.SmsCode = ""
|
||||
op.MustSaveDriverStorage(y)
|
||||
return y.submitSecondDeviceAuth(smsCode)
|
||||
}
|
||||
|
||||
var smsResp LoginResp
|
||||
_, err := y.client.R().
|
||||
ForceContentType("application/json;charset=UTF-8").SetResult(&smsResp).
|
||||
SetHeaders(y.loginHeaders(param.BaseLoginParam)).
|
||||
SetFormData(map[string]string{
|
||||
"mobile": param.SecondAuthMobile,
|
||||
"appKey": APP_ID,
|
||||
}).
|
||||
Post(AUTH_URL + "/api/logbox/oauth2/sendSmsCodeForSecondAuth.do")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if smsResp.Result != 0 {
|
||||
return fmt.Errorf("failed to send the verification SMS: %s", smsResp.Msg)
|
||||
}
|
||||
// 保留登陆参数,等待用户回填短信验证码后重新保存
|
||||
return errors.New("second device verification is required, an SMS code has been sent, please fill it into `sms_code` and save again")
|
||||
}
|
||||
|
||||
// 提交短信验证码完成设备二次校验
|
||||
// 注意:该接口没有独立的短信码字段,短信码要加密后放在epd里(登陆时epd装的是密码)
|
||||
func (y *Cloud189PC) submitSecondDeviceAuth(smsCode string) error {
|
||||
param := y.loginParam
|
||||
var authResp LoginResp
|
||||
res, err := y.client.R().
|
||||
ForceContentType("application/json;charset=UTF-8").SetResult(&authResp).
|
||||
SetHeaders(y.loginHeaders(param.BaseLoginParam)).
|
||||
SetFormData(map[string]string{
|
||||
"mobile": param.SecondAuthMobile,
|
||||
"appKey": APP_ID,
|
||||
"userName": param.RsaUsername,
|
||||
"epd": param.encryptSecret(smsCode),
|
||||
"accountType": ACCOUNT_TYPE,
|
||||
"returnUrl": RETURN_URL,
|
||||
"isOauth2": "false",
|
||||
"cb_SaveName": "1",
|
||||
"state": "",
|
||||
"paramId": param.ParamId,
|
||||
}).
|
||||
Post(AUTH_URL + "/api/logbox/oauth2/submitForSecondAuth.do")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// 校验通过后服务端会下发DEVICEID,保存下来以后就不会再触发二次校验
|
||||
y.saveDeviceID(res)
|
||||
|
||||
if authResp.Result != 0 {
|
||||
return fmt.Errorf("second device verification failed: %s", authResp.Msg)
|
||||
}
|
||||
if authResp.ToUrl == "" {
|
||||
return fmt.Errorf("second device verification failed, no toUrl obtained, msg: %s", authResp.Msg)
|
||||
}
|
||||
return y.getSessionByRedirectURL(authResp.ToUrl)
|
||||
}
|
||||
|
||||
// 用登陆结果的跳转地址换取会话
|
||||
func (y *Cloud189PC) getSessionByRedirectURL(redirectURL string) error {
|
||||
var erron RespErr
|
||||
var tokenInfo AppSessionResp
|
||||
_, err = y.client.R().
|
||||
_, err := y.client.R().
|
||||
SetResult(&tokenInfo).SetError(&erron).
|
||||
SetQueryParams(clientSuffix()).
|
||||
SetQueryParam("redirectURL", loginresp.ToUrl).
|
||||
SetQueryParams(y.deviceParams()).
|
||||
SetQueryParam("redirectURL", redirectURL).
|
||||
SetHeader("X-Request-ID", uuid.NewString()).
|
||||
Post(API_URL + "/getSessionForPC.action")
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -365,14 +540,13 @@ func (y *Cloud189PC) loginByPassword() (err error) {
|
||||
return &erron
|
||||
}
|
||||
if tokenInfo.ResCode != 0 {
|
||||
err = fmt.Errorf(tokenInfo.ResMessage)
|
||||
return err
|
||||
return errors.New(tokenInfo.ResMessage)
|
||||
}
|
||||
y.Addition.AccessToken = tokenInfo.AccessToken
|
||||
y.Addition.RefreshToken = tokenInfo.RefreshToken
|
||||
y.tokenInfo = &tokenInfo
|
||||
op.MustSaveDriverStorage(y)
|
||||
return err
|
||||
return nil
|
||||
}
|
||||
|
||||
func (y *Cloud189PC) loginByQRCode() error {
|
||||
@@ -383,66 +557,74 @@ func (y *Cloud189PC) loginByQRCode() error {
|
||||
}
|
||||
}
|
||||
|
||||
var state struct {
|
||||
Status int `json:"status"`
|
||||
RedirectUrl string `json:"redirectUrl"`
|
||||
Msg string `json:"msg"`
|
||||
// 本地轮询,扫码确认后自动继续,不需要用户反复保存
|
||||
deadline := time.Now().Add(QRCODE_POLL_TIMEOUT)
|
||||
lastStatus := -106
|
||||
for {
|
||||
state, err := y.checkQRCodeState()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check QR code state: %w", err)
|
||||
}
|
||||
lastStatus = state.Status
|
||||
|
||||
switch state.Status {
|
||||
case 0: // 登录成功
|
||||
y.qrcodeParam = nil
|
||||
return y.getSessionByRedirectURL(state.RedirectUrl)
|
||||
case -106, -11002: // -106 等待扫描,-11002 已扫描等待确认
|
||||
case -11001: // 二维码过期
|
||||
y.qrcodeParam = nil
|
||||
return errors.New("QR code expired, please try again")
|
||||
default: // 其他错误
|
||||
y.qrcodeParam = nil
|
||||
return fmt.Errorf("QR code login failed with status %d: %s", state.Status, state.Msg)
|
||||
}
|
||||
|
||||
if time.Now().Add(QRCODE_POLL_INTERVAL).After(deadline) {
|
||||
break
|
||||
}
|
||||
time.Sleep(QRCODE_POLL_INTERVAL)
|
||||
}
|
||||
|
||||
// 轮询超时,把二维码交回前端等待下一次保存
|
||||
if lastStatus == -11002 {
|
||||
return y.genQRCode("QR code has been scanned, please confirm the login on your phone and save again")
|
||||
}
|
||||
return y.genQRCode("QR code has not been scanned yet, please scan and save again")
|
||||
}
|
||||
|
||||
type qrCodeState struct {
|
||||
Status int `json:"status"`
|
||||
RedirectUrl string `json:"redirectUrl"`
|
||||
Msg string `json:"msg"`
|
||||
}
|
||||
|
||||
// 查询扫码状态,参数需与官方PC端一致,否则服务端不会返回授权结果
|
||||
func (y *Cloud189PC) checkQRCodeState() (*qrCodeState, error) {
|
||||
now := time.Now()
|
||||
var state qrCodeState
|
||||
_, err := y.client.R().
|
||||
SetHeaders(map[string]string{
|
||||
"Referer": AUTH_URL,
|
||||
"Reqid": y.qrcodeParam.ReqId,
|
||||
"lt": y.qrcodeParam.Lt,
|
||||
}).
|
||||
SetHeaders(y.loginHeaders(y.qrcodeParam.BaseLoginParam)).
|
||||
SetFormData(map[string]string{
|
||||
"appId": APP_ID,
|
||||
"clientType": CLIENT_TYPE,
|
||||
"returnUrl": RETURN_URL,
|
||||
"paramId": y.qrcodeParam.ParamId,
|
||||
"uuid": y.qrcodeParam.UUID,
|
||||
"encryuuid": y.qrcodeParam.EncryUUID,
|
||||
"date": formatDate(now),
|
||||
"timeStamp": fmt.Sprint(now.UTC().UnixNano() / 1e6),
|
||||
"appId": APP_ID,
|
||||
"clientType": QR_CLIENT_TYPE,
|
||||
"returnUrl": RETURN_URL,
|
||||
"paramId": y.qrcodeParam.ParamId,
|
||||
"uuid": y.qrcodeParam.UUID,
|
||||
"encryuuid": y.qrcodeParam.EncryUUID,
|
||||
"cb_SaveName": "3",
|
||||
"isOauth2": "false",
|
||||
"state": "",
|
||||
"date": formatDate(now),
|
||||
"timeStamp": fmt.Sprint(now.UTC().UnixNano() / 1e6),
|
||||
}).
|
||||
ForceContentType("application/json;charset=UTF-8").
|
||||
SetResult(&state).
|
||||
Post(AUTH_URL + "/api/logbox/oauth2/qrcodeLoginState.do")
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to check QR code state: %w", err)
|
||||
}
|
||||
|
||||
switch state.Status {
|
||||
case 0: // 登录成功
|
||||
var tokenInfo AppSessionResp
|
||||
_, err = y.client.R().
|
||||
SetResult(&tokenInfo).
|
||||
SetQueryParams(clientSuffix()).
|
||||
SetQueryParam("redirectURL", state.RedirectUrl).
|
||||
Post(API_URL + "/getSessionForPC.action")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if tokenInfo.ResCode != 0 {
|
||||
return fmt.Errorf(tokenInfo.ResMessage)
|
||||
}
|
||||
y.Addition.AccessToken = tokenInfo.AccessToken
|
||||
y.Addition.RefreshToken = tokenInfo.RefreshToken
|
||||
y.tokenInfo = &tokenInfo
|
||||
op.MustSaveDriverStorage(y)
|
||||
return nil
|
||||
case -11001: // 二维码过期
|
||||
y.qrcodeParam = nil
|
||||
return errors.New("QR code expired, please try again")
|
||||
case -106: // 等待扫描
|
||||
return y.genQRCode("QR code has not been scanned yet, please scan and save again")
|
||||
case -11002: // 等待确认
|
||||
return y.genQRCode("QR code has been scanned, please confirm the login on your phone and save again")
|
||||
default: // 其他错误
|
||||
y.qrcodeParam = nil
|
||||
return fmt.Errorf("QR code login failed with status %d: %s", state.Status, state.Msg)
|
||||
return nil, err
|
||||
}
|
||||
return &state, nil
|
||||
}
|
||||
|
||||
func (y *Cloud189PC) genQRCode(text string) error {
|
||||
@@ -468,8 +650,9 @@ func (y *Cloud189PC) genQRCode(text string) error {
|
||||
}
|
||||
|
||||
func (y *Cloud189PC) initBaseParams() (*BaseLoginParam, error) {
|
||||
// 清除cookie
|
||||
// 清除cookie,并带上已保存的设备标识
|
||||
jar, _ := cookiejar.New(nil)
|
||||
y.applyDeviceID(jar)
|
||||
y.client.SetCookieJar(jar)
|
||||
|
||||
res, err := y.client.R().
|
||||
@@ -484,14 +667,98 @@ func (y *Cloud189PC) initBaseParams() (*BaseLoginParam, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 当前登陆页把lt/reqId放在跳转地址上,老页面则写在页内变量里,两种都要支持
|
||||
param, err := parseBaseParamFromRedirect(res.RawResponse.Request.URL)
|
||||
if err != nil {
|
||||
param, err = parseBaseParamFromPage(res.String())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return param, nil
|
||||
}
|
||||
|
||||
// 跳转地址上没有paramId,需要再问一次appConf.do
|
||||
var appConf AppConfResp
|
||||
_, err = y.client.R().
|
||||
SetHeaders(y.loginHeaders(*param)).
|
||||
ForceContentType("application/json;charset=UTF-8").
|
||||
SetResult(&appConf).
|
||||
SetFormData(map[string]string{
|
||||
"version": "2.0",
|
||||
"appKey": APP_ID,
|
||||
}).
|
||||
Post(AUTH_URL + "/api/logbox/oauth2/appConf.do")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !appConf.Succeeded() || appConf.Data.ParamId == "" {
|
||||
return nil, fmt.Errorf("failed to get the login paramId: %s", appConf.Msg)
|
||||
}
|
||||
param.ParamId = appConf.Data.ParamId
|
||||
return param, nil
|
||||
}
|
||||
|
||||
// parseBaseParamFromRedirect 从logbox跳转地址提取登陆参数,并以该地址作为后续请求的Referer
|
||||
func parseBaseParamFromRedirect(finalUrl *url.URL) (*BaseLoginParam, error) {
|
||||
if finalUrl == nil {
|
||||
return nil, errors.New("no login page redirect")
|
||||
}
|
||||
query := finalUrl.Query()
|
||||
lt, reqId := query.Get("lt"), query.Get("reqId")
|
||||
if lt == "" || reqId == "" {
|
||||
return nil, errors.New("no lt/reqId in the login page redirect")
|
||||
}
|
||||
return &BaseLoginParam{
|
||||
CaptchaToken: regexp.MustCompile(`'captchaToken' value='(.+?)'`).FindStringSubmatch(res.String())[1],
|
||||
Lt: regexp.MustCompile(`lt = "(.+?)"`).FindStringSubmatch(res.String())[1],
|
||||
ParamId: regexp.MustCompile(`paramId = "(.+?)"`).FindStringSubmatch(res.String())[1],
|
||||
ReqId: regexp.MustCompile(`reqId = "(.+?)"`).FindStringSubmatch(res.String())[1],
|
||||
Lt: lt,
|
||||
ReqId: reqId,
|
||||
Referer: finalUrl.String(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// parseBaseParamFromPage 兼容把参数写在页内变量里的老登陆页
|
||||
func parseBaseParamFromPage(body string) (*BaseLoginParam, error) {
|
||||
lt, err := matchLoginParam(body, `lt = "(.+?)"`, "lt")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reqId, err := matchLoginParam(body, `reqId = "(.+?)"`, "reqId")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
paramId, err := matchLoginParam(body, `paramId = "(.+?)"`, "paramId")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// 老页面才有内嵌的图形验证码token
|
||||
captchaToken, _ := matchLoginParam(body, `'captchaToken' value='(.+?)'`, "captchaToken")
|
||||
|
||||
param := &BaseLoginParam{
|
||||
CaptchaToken: captchaToken,
|
||||
Lt: lt,
|
||||
ParamId: paramId,
|
||||
ReqId: reqId,
|
||||
}
|
||||
encryptUrl, _ := matchLoginParam(body, `encryptUrl = "(.+?)"`, "encryptUrl")
|
||||
param.Referer = AUTH_URL + "/api/logbox/separate/web/index.html?" + strings.Join([]string{
|
||||
"appId=" + url.QueryEscape(APP_ID),
|
||||
"lt=" + url.QueryEscape(param.Lt),
|
||||
"reqId=" + url.QueryEscape(param.ReqId),
|
||||
}, "&")
|
||||
if encryptUrl != "" {
|
||||
param.Referer += "&encryptUrl=" + url.QueryEscape(encryptUrl)
|
||||
}
|
||||
return param, nil
|
||||
}
|
||||
|
||||
// matchLoginParam 从登陆页面提取参数,缺失时返回可读的错误而不是panic
|
||||
func matchLoginParam(body, pattern, name string) (string, error) {
|
||||
matches := regexp.MustCompile(pattern).FindStringSubmatch(body)
|
||||
if len(matches) < 2 {
|
||||
return "", fmt.Errorf("failed to get %s from the login page", name)
|
||||
}
|
||||
return matches[1], nil
|
||||
}
|
||||
|
||||
/* 初始化登陆需要的参数
|
||||
* 如果遇到验证码返回错误
|
||||
*/
|
||||
@@ -516,12 +783,13 @@ func (y *Cloud189PC) initLoginParam() error {
|
||||
}
|
||||
|
||||
y.loginParam.jRsaKey = fmt.Sprintf("-----BEGIN PUBLIC KEY-----\n%s\n-----END PUBLIC KEY-----", encryptConf.Data.PubKey)
|
||||
y.loginParam.RsaUsername = encryptConf.Data.Pre + RsaEncrypt(y.loginParam.jRsaKey, y.Username)
|
||||
y.loginParam.RsaPassword = encryptConf.Data.Pre + RsaEncrypt(y.loginParam.jRsaKey, y.Password)
|
||||
y.loginParam.rsaPrefix = encryptConf.Data.Pre
|
||||
y.loginParam.RsaUsername = y.loginParam.encryptSecret(y.Username)
|
||||
y.loginParam.RsaPassword = y.loginParam.encryptSecret(y.Password)
|
||||
|
||||
// 判断是否需要验证码
|
||||
resp, err := y.client.R().
|
||||
SetHeader("REQID", y.loginParam.ReqId).
|
||||
SetHeaders(y.loginHeaders(y.loginParam.BaseLoginParam)).
|
||||
SetFormData(map[string]string{
|
||||
"appKey": APP_ID,
|
||||
"accountType": ACCOUNT_TYPE,
|
||||
@@ -576,6 +844,7 @@ func (y *Cloud189PC) initQRCodeParam() (err error) {
|
||||
|
||||
var qrcodeParam QRLoginParam
|
||||
_, err = y.client.R().
|
||||
SetHeaders(y.loginHeaders(*baseParam)).
|
||||
SetFormData(map[string]string{"appId": APP_ID}).
|
||||
ForceContentType("application/json;charset=UTF-8").
|
||||
SetResult(&qrcodeParam).
|
||||
@@ -583,6 +852,9 @@ func (y *Cloud189PC) initQRCodeParam() (err error) {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if qrcodeParam.UUID == "" {
|
||||
return errors.New("failed to get the QR code uuid")
|
||||
}
|
||||
qrcodeParam.BaseLoginParam = *baseParam
|
||||
y.qrcodeParam = &qrcodeParam
|
||||
|
||||
@@ -603,6 +875,7 @@ func (y *Cloud189PC) refreshSessionWithRetry(retryCount int) (err error) {
|
||||
_, err = y.client.R().
|
||||
SetResult(&userSessionResp).SetError(&erron).
|
||||
SetQueryParams(clientSuffix()).
|
||||
SetQueryParams(y.deviceParams()).
|
||||
SetQueryParams(map[string]string{
|
||||
"appId": APP_ID,
|
||||
"accessToken": y.tokenInfo.AccessToken,
|
||||
@@ -643,12 +916,11 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
|
||||
return errors.New("refresh token failed after maximum retries")
|
||||
}
|
||||
|
||||
var erron RespErr
|
||||
var tokenInfo AppSessionResp
|
||||
// 该接口刷新失败时以HTTP 200返回 result/msg,SetError不会触发,必须解析响应体判断
|
||||
var tokenInfo RefreshTokenResp
|
||||
_, err = y.client.R().
|
||||
SetResult(&tokenInfo).
|
||||
ForceContentType("application/json;charset=UTF-8").
|
||||
SetError(&erron).
|
||||
SetFormData(map[string]string{
|
||||
"clientId": APP_ID,
|
||||
"refreshToken": y.tokenInfo.RefreshToken,
|
||||
@@ -661,7 +933,8 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
|
||||
}
|
||||
|
||||
// 如果刷新失败,返回错误给上层处理
|
||||
if erron.HasError() {
|
||||
if tokenInfo.HasError() {
|
||||
refreshErr := tokenInfo.Error()
|
||||
if y.Addition.RefreshToken != "" {
|
||||
y.Addition.RefreshToken = ""
|
||||
op.MustSaveDriverStorage(y)
|
||||
@@ -669,7 +942,11 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
|
||||
|
||||
// 根据登录类型决定下一步行为
|
||||
if y.LoginType == "qrcode" {
|
||||
return errors.New("QR code session has expired, please re-scan the code to log in")
|
||||
return fmt.Errorf("QR code session has expired, please re-scan the code to log in: %s", refreshErr)
|
||||
}
|
||||
// 没有账号密码时无法回退到完整登录,直接把刷新失败的原因返回
|
||||
if y.Username == "" || y.Password == "" {
|
||||
return errors.New(refreshErr)
|
||||
}
|
||||
// 密码登录模式下,尝试回退到完整登录
|
||||
return y.login()
|
||||
@@ -677,7 +954,8 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
|
||||
|
||||
y.Addition.AccessToken = tokenInfo.AccessToken
|
||||
y.Addition.RefreshToken = tokenInfo.RefreshToken
|
||||
y.tokenInfo = &tokenInfo
|
||||
y.tokenInfo.AccessToken = tokenInfo.AccessToken
|
||||
y.tokenInfo.RefreshToken = tokenInfo.RefreshToken
|
||||
op.MustSaveDriverStorage(y)
|
||||
return y.refreshSessionWithRetry(retryCount + 1)
|
||||
}
|
||||
|
||||
@@ -328,7 +328,6 @@ func (d *Alias) Link(ctx context.Context, file model.Obj, args model.LinkArgs) (
|
||||
return nil, err
|
||||
}
|
||||
resultLink := link.Clone() // 复制一份,避免修改到原始link
|
||||
resultLink.Expiration = nil
|
||||
if args.Redirect {
|
||||
return resultLink, nil
|
||||
}
|
||||
|
||||
@@ -95,6 +95,7 @@ func (d *AListV3) List(ctx context.Context, dir model.Obj, args model.ListArgs)
|
||||
file := model.ObjThumb{
|
||||
Object: model.Object{
|
||||
Name: f.Name,
|
||||
Path: path.Join(dir.GetPath(), f.Name),
|
||||
Modified: f.Modified,
|
||||
Ctime: f.Created,
|
||||
Size: f.Size,
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
package alist_v3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/base"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/driver"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/go-resty/resty/v2"
|
||||
)
|
||||
|
||||
// TestListSetsChildPaths descends two levels the way op.Get does, feeding an
|
||||
// object from one listing back into List as dir. Without a Path on that object
|
||||
// the driver asks upstream for "", which a real server answers with its own
|
||||
// root -- hence the fake upstream's fallback, and the endless self-similar tree.
|
||||
func TestListSetsChildPaths(t *testing.T) {
|
||||
tree := map[string][]ObjResp{
|
||||
"/": {{Name: "root-marker", IsDir: true}},
|
||||
"/drive": {{Name: "concerts", IsDir: true}},
|
||||
"/drive/concerts": {{Name: "show.mkv"}},
|
||||
}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var req ListReq
|
||||
_ = json.NewDecoder(r.Body).Decode(&req)
|
||||
content, ok := tree[req.Path]
|
||||
if !ok {
|
||||
content = tree["/"]
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json") // resty only unmarshals JSON
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"code": 200, "message": "success",
|
||||
"data": map[string]any{"content": content, "total": len(content)},
|
||||
})
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
// conf.Conf is nil outside a booted server, so base.InitClient() is unusable.
|
||||
prev := base.RestyClient
|
||||
base.RestyClient = resty.New().SetTimeout(5 * time.Second)
|
||||
t.Cleanup(func() { base.RestyClient = prev })
|
||||
|
||||
d := &AListV3{Addition: Addition{
|
||||
RootPath: driver.RootPath{RootFolderPath: "/drive"},
|
||||
Address: srv.URL,
|
||||
}}
|
||||
dir := model.Obj(&model.Object{Path: "/drive", IsFolder: true})
|
||||
for _, want := range []string{"/drive/concerts", "/drive/concerts/show.mkv"} {
|
||||
objs, err := d.List(context.Background(), dir, model.ListArgs{})
|
||||
if err != nil {
|
||||
t.Fatalf("List(%q): %v", dir.GetPath(), err)
|
||||
}
|
||||
if len(objs) != 1 {
|
||||
t.Fatalf("List(%q) returned %d objects, want 1", dir.GetPath(), len(objs))
|
||||
}
|
||||
if got := objs[0].GetPath(); got != want {
|
||||
t.Fatalf("child of %q has path %q, want %q", dir.GetPath(), got, want)
|
||||
}
|
||||
dir = objs[0]
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,295 @@
|
||||
package aliyundrive_open
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"math/rand/v2"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/errs"
|
||||
anet "github.com/OpenListTeam/OpenList/v4/internal/net"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultCallbackConcurrency = 1
|
||||
callbackAcquireTimeout = time.Second
|
||||
callbackRequestAttempts = 3
|
||||
callbackRetryBaseDelay = 200 * time.Millisecond
|
||||
callbackErrorBodyLimit = 64 << 10
|
||||
)
|
||||
|
||||
var callbackLimiters = struct {
|
||||
sync.Mutex
|
||||
byUser map[string]*callbackLimiter
|
||||
}{byUser: make(map[string]*callbackLimiter)}
|
||||
|
||||
type callbackLimiter struct {
|
||||
userID string
|
||||
mu sync.Mutex
|
||||
active int
|
||||
nextID uint64
|
||||
registrations map[uint64]int
|
||||
changed chan struct{}
|
||||
}
|
||||
|
||||
type callbackRegistration struct {
|
||||
limiter *callbackLimiter
|
||||
id uint64
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
type callbackPermit struct {
|
||||
limiter *callbackLimiter
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func normalizeCallbackConcurrency(limit int) int {
|
||||
if limit <= 0 {
|
||||
return defaultCallbackConcurrency
|
||||
}
|
||||
return limit
|
||||
}
|
||||
|
||||
func registerCallbackLimiter(userID string, limit int) *callbackRegistration {
|
||||
callbackLimiters.Lock()
|
||||
defer callbackLimiters.Unlock()
|
||||
|
||||
limiter := callbackLimiters.byUser[userID]
|
||||
if limiter == nil {
|
||||
limiter = &callbackLimiter{
|
||||
userID: userID,
|
||||
registrations: make(map[uint64]int),
|
||||
changed: make(chan struct{}),
|
||||
}
|
||||
callbackLimiters.byUser[userID] = limiter
|
||||
}
|
||||
limiter.mu.Lock()
|
||||
limiter.nextID++
|
||||
id := limiter.nextID
|
||||
limiter.registrations[id] = normalizeCallbackConcurrency(limit)
|
||||
limiter.signalLocked()
|
||||
limiter.mu.Unlock()
|
||||
return &callbackRegistration{limiter: limiter, id: id}
|
||||
}
|
||||
|
||||
func (r *callbackRegistration) unregister() {
|
||||
if r == nil || r.limiter == nil {
|
||||
return
|
||||
}
|
||||
r.once.Do(func() {
|
||||
callbackLimiters.Lock()
|
||||
defer callbackLimiters.Unlock()
|
||||
r.limiter.mu.Lock()
|
||||
delete(r.limiter.registrations, r.id)
|
||||
r.limiter.signalLocked()
|
||||
if len(r.limiter.registrations) == 0 && r.limiter.active == 0 {
|
||||
delete(callbackLimiters.byUser, r.limiter.userID)
|
||||
}
|
||||
r.limiter.mu.Unlock()
|
||||
})
|
||||
}
|
||||
|
||||
func (r *callbackRegistration) acquire(ctx context.Context) (*callbackPermit, error) {
|
||||
if r == nil || r.limiter == nil {
|
||||
return nil, errs.NewErr(errs.TemporaryCapacity, "callback limiter is unavailable")
|
||||
}
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
waitCtx, cancel := context.WithTimeout(ctx, callbackAcquireTimeout)
|
||||
defer cancel()
|
||||
for {
|
||||
r.limiter.mu.Lock()
|
||||
if r.limiter.active < r.limiter.limitLocked() {
|
||||
r.limiter.active++
|
||||
r.limiter.mu.Unlock()
|
||||
return &callbackPermit{limiter: r.limiter}, nil
|
||||
}
|
||||
changed := r.limiter.changed
|
||||
r.limiter.mu.Unlock()
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-waitCtx.Done():
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nil, errs.NewErr(errs.TemporaryCapacity, "timed out waiting for callback admission")
|
||||
case <-changed:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (l *callbackLimiter) limitLocked() int {
|
||||
limit := 0
|
||||
for _, registered := range l.registrations {
|
||||
if limit == 0 || registered < limit {
|
||||
limit = registered
|
||||
}
|
||||
}
|
||||
return limit
|
||||
}
|
||||
|
||||
func (l *callbackLimiter) signalLocked() {
|
||||
close(l.changed)
|
||||
l.changed = make(chan struct{})
|
||||
}
|
||||
|
||||
func (p *callbackPermit) release() {
|
||||
if p == nil || p.limiter == nil {
|
||||
return
|
||||
}
|
||||
p.once.Do(func() {
|
||||
callbackLimiters.Lock()
|
||||
defer callbackLimiters.Unlock()
|
||||
p.limiter.mu.Lock()
|
||||
p.limiter.active--
|
||||
p.limiter.signalLocked()
|
||||
if len(p.limiter.registrations) == 0 && p.limiter.active == 0 {
|
||||
delete(callbackLimiters.byUser, p.limiter.userID)
|
||||
}
|
||||
p.limiter.mu.Unlock()
|
||||
})
|
||||
}
|
||||
|
||||
func (d *AliyundriveOpen) callbackRegistration() *callbackRegistration {
|
||||
if d.callback != nil {
|
||||
return d.callback
|
||||
}
|
||||
if d.ref != nil {
|
||||
return d.ref.callbackRegistration()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *AliyundriveOpen) callbackRangeReader(url string, size int64) stream.RangeReaderFunc {
|
||||
return func(ctx context.Context, requested http_range.Range) (io.ReadCloser, error) {
|
||||
if requested.Length < 0 || requested.Start+requested.Length > size {
|
||||
requested.Length = size - requested.Start
|
||||
}
|
||||
for attempt := 0; attempt < callbackRequestAttempts; attempt++ {
|
||||
permit, err := d.callbackRegistration().acquire(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
body, retry, err := openCallbackRange(ctx, url, size, requested)
|
||||
if !retry && err == nil {
|
||||
return newCallbackBody(ctx, body, permit.release), nil
|
||||
}
|
||||
permit.release()
|
||||
if !retry {
|
||||
return nil, err
|
||||
}
|
||||
if attempt+1 == callbackRequestAttempts {
|
||||
return nil, errs.NewErr(errs.TemporaryCapacity, "Aliyun callback concurrency limit rejected %d attempts", callbackRequestAttempts)
|
||||
}
|
||||
delay := callbackRetryBaseDelay << attempt
|
||||
delay += time.Duration(rand.Int64N(int64(delay / 2)))
|
||||
timer := time.NewTimer(delay)
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
timer.Stop()
|
||||
return nil, ctx.Err()
|
||||
case <-timer.C:
|
||||
}
|
||||
}
|
||||
return nil, errs.NewErr(errs.TemporaryCapacity, "callback attempts exhausted")
|
||||
}
|
||||
}
|
||||
|
||||
func openCallbackRange(ctx context.Context, url string, size int64, requested http_range.Range) (io.ReadCloser, bool, error) {
|
||||
requestHeader, _ := ctx.Value(conf.RequestHeaderKey).(http.Header)
|
||||
header := anet.ProcessHeader(requestHeader, nil)
|
||||
header = http_range.ApplyRangeToHttpHeader(requested, header)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, false, fmt.Errorf("create Aliyun callback request: %w", err)
|
||||
}
|
||||
req.Header = header
|
||||
response, err := anet.HttpClient().Do(req)
|
||||
if err != nil {
|
||||
return nil, false, fmt.Errorf("Aliyun callback request failed: %w", err)
|
||||
}
|
||||
if response.StatusCode >= http.StatusBadRequest {
|
||||
defer response.Body.Close()
|
||||
body, readErr := io.ReadAll(io.LimitReader(response.Body, callbackErrorBodyLimit))
|
||||
if readErr != nil {
|
||||
return nil, false, fmt.Errorf("read Aliyun callback error response: %w", readErr)
|
||||
}
|
||||
if isCallbackCapacityRejection(response.StatusCode, body) {
|
||||
return nil, true, nil
|
||||
}
|
||||
message := strings.ReplaceAll(strings.TrimSpace(string(body)), url, "<redacted>")
|
||||
return nil, false, fmt.Errorf("Aliyun callback request failed: %w; response: %s", anet.HttpStatusCodeError(response.StatusCode), message)
|
||||
}
|
||||
if requested.Start == 0 && requested.Length == size || response.StatusCode == http.StatusPartialContent || callbackContentRangeStartsAt(response.Header, requested.Start) {
|
||||
return response.Body, false, nil
|
||||
}
|
||||
if response.StatusCode == http.StatusOK {
|
||||
body, rangeErr := anet.GetRangedHttpReader(response.Body, requested.Start, requested.Length)
|
||||
if rangeErr != nil {
|
||||
response.Body.Close()
|
||||
return nil, false, rangeErr
|
||||
}
|
||||
return body, false, nil
|
||||
}
|
||||
return response.Body, false, nil
|
||||
}
|
||||
|
||||
func isCallbackCapacityRejection(status int, body []byte) bool {
|
||||
return status == http.StatusForbidden &&
|
||||
strings.Contains(string(body), "RequestDeniedByCallback") &&
|
||||
strings.Contains(string(body), "ExceedMaxConcurrency")
|
||||
}
|
||||
|
||||
func callbackContentRangeStartsAt(header http.Header, offset int64) bool {
|
||||
start, _, err := http_range.ParseContentRange(header.Get("Content-Range"))
|
||||
return err == nil && start == offset
|
||||
}
|
||||
|
||||
type callbackBody struct {
|
||||
body io.ReadCloser
|
||||
release func()
|
||||
once sync.Once
|
||||
mu sync.Mutex
|
||||
stop func() bool
|
||||
}
|
||||
|
||||
func newCallbackBody(ctx context.Context, body io.ReadCloser, release func()) *callbackBody {
|
||||
b := &callbackBody{body: body, release: release}
|
||||
stop := context.AfterFunc(ctx, func() { _ = b.Close() })
|
||||
b.mu.Lock()
|
||||
b.stop = stop
|
||||
b.mu.Unlock()
|
||||
return b
|
||||
}
|
||||
|
||||
func (b *callbackBody) Read(p []byte) (int, error) {
|
||||
n, err := b.body.Read(p)
|
||||
if err != nil {
|
||||
_ = b.Close()
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (b *callbackBody) Close() error {
|
||||
var err error
|
||||
b.once.Do(func() {
|
||||
b.mu.Lock()
|
||||
stop := b.stop
|
||||
b.mu.Unlock()
|
||||
if stop != nil {
|
||||
stop()
|
||||
}
|
||||
err = b.body.Close()
|
||||
b.release()
|
||||
})
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,365 @@
|
||||
package aliyundrive_open
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/base"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/errs"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
)
|
||||
|
||||
func TestLinkSeparatesRedirectAndProxyRepresentations(t *testing.T) {
|
||||
oldConf := conf.Conf
|
||||
conf.Conf = &conf.Config{}
|
||||
t.Cleanup(func() { conf.Conf = oldConf })
|
||||
base.InitClient()
|
||||
|
||||
var server *httptest.Server
|
||||
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
switch r.URL.Path {
|
||||
case "/adrive/v1.0/user/getDriveInfo":
|
||||
_, _ = fmt.Fprint(w, `{"user_id":"user-1","resource_drive_id":"drive-1"}`)
|
||||
case "/adrive/v1.0/openFile/getDownloadUrl":
|
||||
_, _ = fmt.Fprintf(w, `{"url":%q}`, server.URL+"/callback")
|
||||
default:
|
||||
http.NotFound(w, r)
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
oldAPIURL := API_URL
|
||||
API_URL = server.URL
|
||||
defer func() { API_URL = oldAPIURL }()
|
||||
|
||||
d := &AliyundriveOpen{Addition: Addition{AccessToken: "token"}}
|
||||
if err := d.Init(t.Context()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer d.Drop(context.Background())
|
||||
if d.CallbackConcurrency != defaultCallbackConcurrency {
|
||||
t.Fatalf("normalized callback concurrency = %d, want %d", d.CallbackConcurrency, defaultCallbackConcurrency)
|
||||
}
|
||||
|
||||
link, err := d.Link(t.Context(), &model.Object{ID: "file-1", Name: "file", Size: 1}, model.LinkArgs{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if link.RangeReader == nil {
|
||||
t.Fatal("proxy link must own callback acquisition through a range reader")
|
||||
}
|
||||
if _, ok := link.RangeReader.(stream.RateLimitRangeReaderFunc); !ok {
|
||||
t.Fatalf("proxy range reader type = %T, want server-rate-limited reader", link.RangeReader)
|
||||
}
|
||||
direct, err := d.Link(t.Context(), &model.Object{ID: "file-1", Name: "file", Size: 1}, model.LinkArgs{Redirect: true})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if direct.URL == "" || direct.RangeReader != nil {
|
||||
t.Fatal("redirect link must remain URL-only")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackRangeHoldsPermitUntilBodyClose(t *testing.T) {
|
||||
oldConf := conf.Conf
|
||||
conf.Conf = &conf.Config{}
|
||||
t.Cleanup(func() { conf.Conf = oldConf })
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.Header().Set("Content-Length", "1")
|
||||
w.Header().Set("Content-Range", "bytes 0-0/1")
|
||||
w.WriteHeader(http.StatusPartialContent)
|
||||
_, _ = io.WriteString(w, "x")
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
registration := registerCallbackLimiter(t.Name(), 1)
|
||||
t.Cleanup(registration.unregister)
|
||||
d := &AliyundriveOpen{callback: registration}
|
||||
body, err := d.callbackRangeReader(server.URL, 1).RangeRead(t.Context(), http_range.Range{Length: 1})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
registration.limiter.mu.Lock()
|
||||
active := registration.limiter.active
|
||||
registration.limiter.mu.Unlock()
|
||||
if active != 1 {
|
||||
t.Fatalf("active callback bodies = %d, want 1", active)
|
||||
}
|
||||
if err := body.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
registration.limiter.mu.Lock()
|
||||
active = registration.limiter.active
|
||||
registration.limiter.mu.Unlock()
|
||||
if active != 0 {
|
||||
t.Fatalf("active callback bodies after Close = %d, want 0", active)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackLimiterUsesMinimumRegisteredLimit(t *testing.T) {
|
||||
firstRegistration := registerCallbackLimiter(t.Name(), 2)
|
||||
t.Cleanup(firstRegistration.unregister)
|
||||
first, err := firstRegistration.acquire(t.Context())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
second, err := firstRegistration.acquire(t.Context())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer first.release()
|
||||
defer second.release()
|
||||
|
||||
lowerRegistration := registerCallbackLimiter(t.Name(), 1)
|
||||
t.Cleanup(lowerRegistration.unregister)
|
||||
acquired := make(chan *callbackPermit, 1)
|
||||
go func() {
|
||||
permit, acquireErr := lowerRegistration.acquire(t.Context())
|
||||
if acquireErr == nil {
|
||||
acquired <- permit
|
||||
}
|
||||
}()
|
||||
|
||||
first.release()
|
||||
select {
|
||||
case permit := <-acquired:
|
||||
permit.release()
|
||||
t.Fatal("lowering the shared limit must wait for all excess bodies to drain")
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
}
|
||||
second.release()
|
||||
select {
|
||||
case permit := <-acquired:
|
||||
permit.release()
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("admission did not resume after active bodies drained below the new limit")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackLimiterSeparatesUsers(t *testing.T) {
|
||||
firstUser := registerCallbackLimiter(t.Name()+"-first", 1)
|
||||
secondUser := registerCallbackLimiter(t.Name()+"-second", 1)
|
||||
t.Cleanup(firstUser.unregister)
|
||||
t.Cleanup(secondUser.unregister)
|
||||
first, err := firstUser.acquire(t.Context())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer first.release()
|
||||
second, err := secondUser.acquire(t.Context())
|
||||
if err != nil {
|
||||
t.Fatalf("independent user was blocked: %v", err)
|
||||
}
|
||||
second.release()
|
||||
}
|
||||
|
||||
func TestCallbackLimiterReconfigureWaitsForOldBodies(t *testing.T) {
|
||||
userID := t.Name()
|
||||
oldRegistration := registerCallbackLimiter(userID, 2)
|
||||
first, err := oldRegistration.acquire(t.Context())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
second, err := oldRegistration.acquire(t.Context())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
oldRegistration.unregister()
|
||||
|
||||
newRegistration := registerCallbackLimiter(userID, 1)
|
||||
t.Cleanup(newRegistration.unregister)
|
||||
acquired := make(chan *callbackPermit, 1)
|
||||
go func() {
|
||||
permit, acquireErr := newRegistration.acquire(t.Context())
|
||||
if acquireErr == nil {
|
||||
acquired <- permit
|
||||
}
|
||||
}()
|
||||
first.release()
|
||||
select {
|
||||
case permit := <-acquired:
|
||||
permit.release()
|
||||
t.Fatal("reconfigured limiter admitted while an old body still occupied the new limit")
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
}
|
||||
second.release()
|
||||
select {
|
||||
case permit := <-acquired:
|
||||
permit.release()
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("reconfigured limiter did not admit after old bodies drained")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackLimiterDistinguishesTimeoutAndCancellation(t *testing.T) {
|
||||
registration := registerCallbackLimiter(t.Name(), 1)
|
||||
t.Cleanup(registration.unregister)
|
||||
permit, err := registration.acquire(t.Context())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer permit.release()
|
||||
|
||||
started := time.Now()
|
||||
_, err = registration.acquire(t.Context())
|
||||
if !errors.Is(err, errs.TemporaryCapacity) {
|
||||
t.Fatalf("admission timeout error = %v, want TemporaryCapacity", err)
|
||||
}
|
||||
if time.Since(started) < callbackAcquireTimeout {
|
||||
t.Fatal("admission timed out before the configured wait elapsed")
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(t.Context())
|
||||
cancel()
|
||||
_, err = registration.acquire(ctx)
|
||||
if !errors.Is(err, context.Canceled) || errors.Is(err, errs.TemporaryCapacity) {
|
||||
t.Fatalf("canceled admission error = %v, want only context.Canceled", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackCapacityRejectionRequiresBothExactMarkers(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
want bool
|
||||
}{
|
||||
{name: "both", body: `{"code":"RequestDeniedByCallback","message":"ExceedMaxConcurrency"}`, want: true},
|
||||
{name: "code only", body: `{"code":"RequestDeniedByCallback"}`},
|
||||
{name: "message only", body: `{"message":"ExceedMaxConcurrency"}`},
|
||||
{name: "case differs", body: `{"code":"requestdeniedbycallback","message":"ExceedMaxConcurrency"}`},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
if got := isCallbackCapacityRejection(http.StatusForbidden, []byte(test.body)); got != test.want {
|
||||
t.Fatalf("classification = %v, want %v", got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
if isCallbackCapacityRejection(http.StatusTooManyRequests, []byte(`RequestDeniedByCallback ExceedMaxConcurrency`)) {
|
||||
t.Fatal("non-403 response must not be classified as callback capacity")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackRangeRetriesOnlyVerifiedCapacityRejections(t *testing.T) {
|
||||
var requests atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
requests.Add(1)
|
||||
w.WriteHeader(http.StatusForbidden)
|
||||
_, _ = io.WriteString(w, `{"code":"RequestDeniedByCallback","message":"ExceedMaxConcurrency"}`)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
registration := registerCallbackLimiter(t.Name(), 1)
|
||||
t.Cleanup(registration.unregister)
|
||||
d := &AliyundriveOpen{callback: registration}
|
||||
_, err := d.callbackRangeReader(server.URL+"?token=secret", 1).RangeRead(t.Context(), http_range.Range{Length: 1})
|
||||
if !errors.Is(err, errs.TemporaryCapacity) {
|
||||
t.Fatalf("verified rejection error = %v, want TemporaryCapacity", err)
|
||||
}
|
||||
if requests.Load() != callbackRequestAttempts {
|
||||
t.Fatalf("requests = %d, want %d", requests.Load(), callbackRequestAttempts)
|
||||
}
|
||||
if strings.Contains(err.Error(), "secret") {
|
||||
t.Fatal("capacity error leaked the signed callback URL")
|
||||
}
|
||||
permit, acquireErr := registration.acquire(t.Context())
|
||||
if acquireErr != nil {
|
||||
t.Fatalf("capacity retries leaked admission: %v", acquireErr)
|
||||
}
|
||||
permit.release()
|
||||
|
||||
requests.Store(0)
|
||||
permanent := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
requests.Add(1)
|
||||
w.WriteHeader(http.StatusForbidden)
|
||||
_, _ = io.WriteString(w, `{"code":"RequestDeniedByCallback","message":"denied"}`)
|
||||
}))
|
||||
defer permanent.Close()
|
||||
_, err = d.callbackRangeReader(permanent.URL+"?token=secret", 1).RangeRead(t.Context(), http_range.Range{Length: 1})
|
||||
if errors.Is(err, errs.TemporaryCapacity) {
|
||||
t.Fatalf("permanent 403 error = %v, must not be TemporaryCapacity", err)
|
||||
}
|
||||
if requests.Load() != 1 {
|
||||
t.Fatalf("permanent 403 requests = %d, want 1", requests.Load())
|
||||
}
|
||||
if strings.Contains(err.Error(), "secret") {
|
||||
t.Fatal("permanent error leaked the signed callback URL")
|
||||
}
|
||||
}
|
||||
|
||||
type countingReadCloser struct {
|
||||
reader io.Reader
|
||||
closed atomic.Int32
|
||||
}
|
||||
|
||||
func (r *countingReadCloser) Read(p []byte) (int, error) { return r.reader.Read(p) }
|
||||
func (r *countingReadCloser) Close() error {
|
||||
r.closed.Add(1)
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestCallbackBodyReleasesExactlyOnce(t *testing.T) {
|
||||
underlying := &countingReadCloser{reader: strings.NewReader("x")}
|
||||
var released atomic.Int32
|
||||
body := newCallbackBody(t.Context(), underlying, func() { released.Add(1) })
|
||||
_, _ = io.ReadAll(body)
|
||||
if err := body.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := body.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if underlying.closed.Load() != 1 || released.Load() != 1 {
|
||||
t.Fatalf("close count = %d, release count = %d; want 1, 1", underlying.closed.Load(), released.Load())
|
||||
}
|
||||
}
|
||||
|
||||
type failingReadCloser struct {
|
||||
closed atomic.Int32
|
||||
}
|
||||
|
||||
func (*failingReadCloser) Read([]byte) (int, error) { return 0, errors.New("read failed") }
|
||||
func (r *failingReadCloser) Close() error {
|
||||
r.closed.Add(1)
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestCallbackBodyReadFailureReleasesPermit(t *testing.T) {
|
||||
underlying := &failingReadCloser{}
|
||||
var released atomic.Int32
|
||||
body := newCallbackBody(t.Context(), underlying, func() { released.Add(1) })
|
||||
if _, err := body.Read(make([]byte, 1)); err == nil {
|
||||
t.Fatal("read unexpectedly succeeded")
|
||||
}
|
||||
if underlying.closed.Load() != 1 || released.Load() != 1 {
|
||||
t.Fatalf("close count = %d, release count = %d; want 1, 1", underlying.closed.Load(), released.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackBodyCancellationReleasesPermit(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(t.Context())
|
||||
underlying := &countingReadCloser{reader: strings.NewReader("x")}
|
||||
released := make(chan struct{}, 1)
|
||||
_ = newCallbackBody(ctx, underlying, func() { released <- struct{}{} })
|
||||
cancel()
|
||||
select {
|
||||
case <-released:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("context cancellation did not release callback admission")
|
||||
}
|
||||
if underlying.closed.Load() != 1 {
|
||||
t.Fatalf("underlying close count = %d, want 1", underlying.closed.Load())
|
||||
}
|
||||
}
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/driver"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/errs"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"github.com/go-resty/resty/v2"
|
||||
log "github.com/sirupsen/logrus"
|
||||
@@ -22,8 +23,9 @@ type AliyundriveOpen struct {
|
||||
|
||||
DriveId string
|
||||
|
||||
limiter *limiter
|
||||
ref *AliyundriveOpen
|
||||
limiter *limiter
|
||||
ref *AliyundriveOpen
|
||||
callback *callbackRegistration
|
||||
}
|
||||
|
||||
func (d *AliyundriveOpen) Config() driver.Config {
|
||||
@@ -35,6 +37,7 @@ func (d *AliyundriveOpen) GetAddition() driver.Additional {
|
||||
}
|
||||
|
||||
func (d *AliyundriveOpen) Init(ctx context.Context) error {
|
||||
d.CallbackConcurrency = normalizeCallbackConcurrency(d.CallbackConcurrency)
|
||||
d.limiter = getLimiterForUser(globalLimiterUserID) // First create a globally shared limiter to limit the initial requests.
|
||||
if d.LIVPDownloadFormat == "" {
|
||||
d.LIVPDownloadFormat = "jpeg"
|
||||
@@ -52,6 +55,7 @@ func (d *AliyundriveOpen) Init(ctx context.Context) error {
|
||||
userid := utils.Json.Get(res, "user_id").ToString()
|
||||
d.limiter.free()
|
||||
d.limiter = getLimiterForUser(userid) // Allocate a corresponding limiter for each user.
|
||||
d.callback = registerCallbackLimiter(userid, d.CallbackConcurrency)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -65,6 +69,10 @@ func (d *AliyundriveOpen) InitReference(storage driver.Driver) error {
|
||||
}
|
||||
|
||||
func (d *AliyundriveOpen) Drop(ctx context.Context) error {
|
||||
if d.callback != nil {
|
||||
d.callback.unregister()
|
||||
d.callback = nil
|
||||
}
|
||||
d.limiter.free()
|
||||
d.limiter = nil
|
||||
d.ref = nil
|
||||
@@ -119,10 +127,16 @@ func (d *AliyundriveOpen) Link(ctx context.Context, file model.Obj, args model.L
|
||||
url = utils.Json.Get(res, "streamsUrl", d.LIVPDownloadFormat).ToString()
|
||||
}
|
||||
exp := time.Minute
|
||||
return &model.Link{
|
||||
link := &model.Link{
|
||||
URL: url,
|
||||
Expiration: &exp,
|
||||
}, nil
|
||||
}
|
||||
if args.Redirect {
|
||||
return link, nil
|
||||
}
|
||||
link.URL = ""
|
||||
link.RangeReader = stream.RateLimitRangeReaderFunc(d.callbackRangeReader(url, file.GetSize()))
|
||||
return link, nil
|
||||
}
|
||||
|
||||
func (d *AliyundriveOpen) MakeDir(ctx context.Context, parentDir model.Obj, dirName string) (model.Obj, error) {
|
||||
|
||||
@@ -8,19 +8,20 @@ import (
|
||||
type Addition struct {
|
||||
DriveType string `json:"drive_type" type:"select" options:"default,resource,backup" default:"resource"`
|
||||
driver.RootID
|
||||
RefreshToken string `json:"refresh_token" required:"true"`
|
||||
OrderBy string `json:"order_by" type:"select" options:"name,size,updated_at,created_at"`
|
||||
OrderDirection string `json:"order_direction" type:"select" options:"ASC,DESC"`
|
||||
UseOnlineAPI bool `json:"use_online_api" default:"true"`
|
||||
AlipanType string `json:"alipan_type" required:"true" type:"select" default:"default" options:"default,alipanTV"`
|
||||
APIAddress string `json:"api_url_address" default:"https://api.oplist.org/alicloud/renewapi"`
|
||||
ClientID string `json:"client_id" help:"Keep it empty if you don't have one"`
|
||||
ClientSecret string `json:"client_secret" help:"Keep it empty if you don't have one"`
|
||||
RemoveWay string `json:"remove_way" required:"true" type:"select" options:"trash,delete"`
|
||||
RapidUpload bool `json:"rapid_upload" help:"If you enable this option, the file will be uploaded to the server first, so the progress will be incorrect"`
|
||||
InternalUpload bool `json:"internal_upload" help:"If you are using Aliyun ECS is located in Beijing, you can turn it on to boost the upload speed"`
|
||||
LIVPDownloadFormat string `json:"livp_download_format" type:"select" options:"jpeg,mov" default:"jpeg"`
|
||||
AccessToken string
|
||||
RefreshToken string `json:"refresh_token" required:"true"`
|
||||
OrderBy string `json:"order_by" type:"select" options:"name,size,updated_at,created_at"`
|
||||
OrderDirection string `json:"order_direction" type:"select" options:"ASC,DESC"`
|
||||
UseOnlineAPI bool `json:"use_online_api" default:"true"`
|
||||
AlipanType string `json:"alipan_type" required:"true" type:"select" default:"default" options:"default,alipanTV"`
|
||||
APIAddress string `json:"api_url_address" default:"https://api.oplist.org/alicloud/renewapi"`
|
||||
ClientID string `json:"client_id" help:"Keep it empty if you don't have one"`
|
||||
ClientSecret string `json:"client_secret" help:"Keep it empty if you don't have one"`
|
||||
RemoveWay string `json:"remove_way" required:"true" type:"select" options:"trash,delete"`
|
||||
RapidUpload bool `json:"rapid_upload" help:"If you enable this option, the file will be uploaded to the server first, so the progress will be incorrect"`
|
||||
InternalUpload bool `json:"internal_upload" help:"If you are using Aliyun ECS is located in Beijing, you can turn it on to boost the upload speed"`
|
||||
LIVPDownloadFormat string `json:"livp_download_format" type:"select" options:"jpeg,mov" default:"jpeg"`
|
||||
CallbackConcurrency int `json:"callback_concurrency" type:"number" default:"1" help:"Maximum active proxied downloads shared by Aliyun user"`
|
||||
AccessToken string
|
||||
}
|
||||
|
||||
var config = driver.Config{
|
||||
|
||||
+43
-8
@@ -315,6 +315,36 @@ var findDownPageParamReg = regexp.MustCompile(`<iframe.*?src="(.+?)"`)
|
||||
// 获取文件ID
|
||||
var findFileIDReg = regexp.MustCompile(`'/ajax(?:file|m)\.php\?file=(\d+)'`)
|
||||
|
||||
// 2026-10 改版:文件页将下载参数移入 /fn? 内页,接口变为 apifile 绝对地址并携带签名
|
||||
var (
|
||||
fnDomainReg = regexp.MustCompile(`var\s+domain[12]\s*=\s*'([^']*(?:ajaxfile|ajaxm)\.php\?file=(\d+)[^']*)'`)
|
||||
fnSignReg = regexp.MustCompile(`var\s+wp_sign\s*=\s*'([^']*)'`)
|
||||
fnAjaxDataReg = regexp.MustCompile(`var\s+ajaxdata\s*=\s*'([^']*)'`)
|
||||
)
|
||||
|
||||
// parseFnPage 从改版后的 /fn? 内页提取下载接口地址与签名表单
|
||||
// 对应页面 JS:POST domain1 {'action':'downprocess','websignkey':ajaxdata,'signs':ajaxdata,'sign':wp_sign,'websign':'2','kd':kdns,'ves':1}
|
||||
func parseFnPage(pageData string) (string, map[string]string, error) {
|
||||
matches := fnDomainReg.FindStringSubmatch(pageData)
|
||||
if len(matches) < 3 {
|
||||
return "", nil, fmt.Errorf("not find fn ajax url")
|
||||
}
|
||||
sign := fnSignReg.FindStringSubmatch(pageData)
|
||||
ajaxdata := fnAjaxDataReg.FindStringSubmatch(pageData)
|
||||
if len(sign) < 2 || len(ajaxdata) < 2 {
|
||||
return "", nil, fmt.Errorf("not find fn sign")
|
||||
}
|
||||
return matches[1], map[string]string{
|
||||
"action": "downprocess",
|
||||
"websignkey": ajaxdata[1],
|
||||
"signs": ajaxdata[1],
|
||||
"sign": sign[1],
|
||||
"websign": "2",
|
||||
"kd": "1",
|
||||
"ves": "1",
|
||||
}, nil
|
||||
}
|
||||
|
||||
// 获取分享链接主界面
|
||||
func (d *LanZou) getShareUrlHtml(shareID string) (string, error) {
|
||||
var vs string
|
||||
@@ -437,18 +467,23 @@ func (d *LanZou) getFilesByShareUrl(shareID, pwd string, sharePageData string) (
|
||||
return nil, err
|
||||
}
|
||||
nextPageData := RemoveNotes(string(data))
|
||||
param, err = htmlJsonToMap(nextPageData)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var resp FileShareInfoAndUrlResp[int]
|
||||
matches := findFileIDReg.FindStringSubmatch(nextPageData)
|
||||
if len(matches) < 2 {
|
||||
if len(matches) >= 2 {
|
||||
// 旧版结构:相对路径 /ajaxm.php?file=N
|
||||
param, err = htmlJsonToMap(nextPageData)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ajaxUrl := d.ShareUrl + matches[0][1:len(matches[0])-1]
|
||||
_, err = d.post(ajaxUrl, func(req *resty.Request) { req.SetFormData(param) }, &resp)
|
||||
} else if fnUrl, fnForm, ferr := parseFnPage(nextPageData); ferr == nil {
|
||||
// 2026-10 改版结构:/fn? 内页携带 apifile 绝对地址与签名参数
|
||||
_, err = d.post(fnUrl, func(req *resty.Request) { req.SetFormData(fnForm) }, &resp)
|
||||
} else {
|
||||
return nil, fmt.Errorf("not find file id")
|
||||
}
|
||||
ajaxUrl := d.ShareUrl + matches[0][1:len(matches[0])-1]
|
||||
var resp FileShareInfoAndUrlResp[int]
|
||||
_, err = d.post(ajaxUrl, func(req *resty.Request) { req.SetFormData(param) }, &resp)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -10,10 +10,11 @@ type Addition struct {
|
||||
Username string `json:"username" required:"true"`
|
||||
Password string `json:"password" required:"true"`
|
||||
Platform string `json:"platform" required:"true" default:"web" type:"select" options:"android,web,pc"`
|
||||
RefreshToken string `json:"refresh_token" required:"true" default:""`
|
||||
RefreshToken string `json:"refresh_token" required:"false" default:""`
|
||||
CaptchaToken string `json:"captcha_token" default:""`
|
||||
DeviceID string `json:"device_id" required:"false" default:""`
|
||||
DisableMediaLink bool `json:"disable_media_link" default:"true"`
|
||||
SkipVerification bool `json:"skip_verification" default:"false" help:"ignore the human verification URL returned by the captcha API instead of failing; enabling this may trigger PikPak risk control"`
|
||||
}
|
||||
|
||||
var config = driver.Config{
|
||||
|
||||
+30
-10
@@ -100,12 +100,13 @@ func (d *PikPak) login() error {
|
||||
return errors.New("username or password is empty")
|
||||
}
|
||||
|
||||
// Clear expired access token so captcha requests don't carry a stale bearer
|
||||
d.AccessToken = ""
|
||||
|
||||
url := "https://user.mypikpak.net/v1/auth/signin"
|
||||
// 使用 用户填写的 CaptchaToken —————— (验证后的captcha_token)
|
||||
if d.GetCaptchaToken() == "" {
|
||||
if err := d.RefreshCaptchaTokenInLogin(GetAction(http.MethodPost, url), d.Username); err != nil {
|
||||
return err
|
||||
}
|
||||
// Always refresh captcha token before signin (it may be expired)
|
||||
if err := d.RefreshCaptchaTokenInLogin(GetAction(http.MethodPost, url), d.Username); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var e ErrResp
|
||||
@@ -125,7 +126,12 @@ func (d *PikPak) login() error {
|
||||
data := res.Body()
|
||||
d.RefreshToken = jsoniter.Get(data, "refresh_token").ToString()
|
||||
d.AccessToken = jsoniter.Get(data, "access_token").ToString()
|
||||
if d.AccessToken == "" || d.RefreshToken == "" {
|
||||
return errors.New("login failed: server returned empty tokens")
|
||||
}
|
||||
d.Common.SetUserID(jsoniter.Get(data, "sub").ToString())
|
||||
d.Addition.RefreshToken = d.RefreshToken
|
||||
op.MustSaveDriverStorage(d)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -159,9 +165,14 @@ func (d *PikPak) refreshToken(refreshToken string) error {
|
||||
return errors.New(e.Error())
|
||||
}
|
||||
data := res.Body()
|
||||
newAccessToken := jsoniter.Get(data, "access_token").ToString()
|
||||
newRefreshToken := jsoniter.Get(data, "refresh_token").ToString()
|
||||
if newAccessToken == "" || newRefreshToken == "" {
|
||||
return errors.New("refresh failed: server returned empty tokens")
|
||||
}
|
||||
d.Status = "work"
|
||||
d.RefreshToken = jsoniter.Get(data, "refresh_token").ToString()
|
||||
d.AccessToken = jsoniter.Get(data, "access_token").ToString()
|
||||
d.RefreshToken = newRefreshToken
|
||||
d.AccessToken = newAccessToken
|
||||
d.Common.SetUserID(jsoniter.Get(data, "sub").ToString())
|
||||
d.Addition.RefreshToken = d.RefreshToken
|
||||
op.MustSaveDriverStorage(d)
|
||||
@@ -197,12 +208,18 @@ func (d *PikPak) request(url string, method string, callback base.ReqCallback, r
|
||||
case 0:
|
||||
return res.Body(), nil
|
||||
case 4122, 4121, 16:
|
||||
// access_token 过期
|
||||
if strings.Contains(url, "/v1/auth/") || strings.Contains(url, "/v1/shield/captcha/") {
|
||||
return nil, errors.New(e.Error())
|
||||
}
|
||||
// access_token expired, refresh and retry
|
||||
if err1 := d.refreshToken(d.RefreshToken); err1 != nil {
|
||||
return nil, err1
|
||||
}
|
||||
return d.request(url, method, callback, resp)
|
||||
case 9: // 验证码token过期
|
||||
case 9: // captcha token expired
|
||||
if strings.Contains(url, "/v1/shield/captcha/") {
|
||||
return nil, errors.New(e.Error())
|
||||
}
|
||||
if err = d.RefreshCaptchaTokenAtLogin(GetAction(method, url), d.GetUserID()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -369,6 +386,9 @@ func (d *PikPak) RefreshCaptchaTokenInLogin(action, username string) error {
|
||||
} else {
|
||||
metas["username"] = username
|
||||
}
|
||||
metas["client_version"] = d.ClientVersion
|
||||
metas["package_name"] = d.PackageName
|
||||
metas["timestamp"], metas["captcha_sign"] = d.Common.GetCaptchaSign()
|
||||
return d.refreshCaptchaToken(action, metas)
|
||||
}
|
||||
|
||||
@@ -407,7 +427,7 @@ func (d *PikPak) refreshCaptchaToken(action string, metas map[string]string) err
|
||||
return errors.New(e.Error())
|
||||
}
|
||||
|
||||
if resp.Url != "" {
|
||||
if resp.Url != "" && !d.Addition.SkipVerification {
|
||||
return fmt.Errorf(`need verify: <a target="_blank" href="%s">Click Here</a>`, resp.Url)
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,720 @@
|
||||
package pikpak
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/base"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/db"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/go-resty/resty/v2"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// --- Helper function tests ---
|
||||
|
||||
func TestGetAction(t *testing.T) {
|
||||
tests := []struct {
|
||||
method string
|
||||
url string
|
||||
want string
|
||||
}{
|
||||
{"GET", "https://api-drive.mypikpak.net/drive/v1/files", "GET:/drive/v1/files"},
|
||||
{"POST", "https://user.mypikpak.net/v1/auth/signin", "POST:/v1/auth/signin"},
|
||||
{"POST", "https://user.mypikpak.net/v1/shield/captcha/init", "POST:/v1/shield/captcha/init"},
|
||||
{"GET", "https://api-drive.mypikpak.net/drive/v1/files?page_token=abc", "GET:/drive/v1/files"},
|
||||
{"POST", "https://user.mypikpak.net/v1/auth/token", "POST:/v1/auth/token"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.method+":"+tt.url, func(t *testing.T) {
|
||||
got := GetAction(tt.method, tt.url)
|
||||
if got != tt.want {
|
||||
t.Errorf("GetAction(%q, %q) = %q, want %q", tt.method, tt.url, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetCaptchaSign(t *testing.T) {
|
||||
c := &Common{
|
||||
ClientID: "YNxT9w7GMdWvEOKa",
|
||||
ClientVersion: "1.53.2",
|
||||
PackageName: "com.pikcloud.pikpak",
|
||||
DeviceID: "test-device-id",
|
||||
Algorithms: AndroidAlgorithms,
|
||||
}
|
||||
|
||||
timestamp, sign := c.GetCaptchaSign()
|
||||
if timestamp == "" {
|
||||
t.Fatal("timestamp should not be empty")
|
||||
}
|
||||
if len(sign) != 34 {
|
||||
t.Fatalf("sign length should be 34 (\"1.\" + 32 hex), got %d: %q", len(sign), sign)
|
||||
}
|
||||
if sign[:2] != "1." {
|
||||
t.Errorf("sign should start with '1.', got %q", sign[:2])
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateDeviceSign(t *testing.T) {
|
||||
sign := generateDeviceSign("test-device", "com.pikcloud.pikpak")
|
||||
if len(sign) < 7 {
|
||||
t.Fatal("device sign too short")
|
||||
}
|
||||
if sign[:7] != "div101." {
|
||||
t.Errorf("device sign should start with 'div101.', got %q", sign[:7])
|
||||
}
|
||||
// Deterministic
|
||||
if sign != generateDeviceSign("test-device", "com.pikcloud.pikpak") {
|
||||
t.Error("generateDeviceSign should be deterministic")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildCustomUserAgent(t *testing.T) {
|
||||
ua := BuildCustomUserAgent("dev123", AndroidClientID, AndroidPackageName,
|
||||
AndroidSdkVersion, AndroidClientVersion, AndroidPackageName, "user456")
|
||||
for _, want := range []string{"ANDROID-", "clientid/", "deviceid/dev123", "usrno/user456"} {
|
||||
if !strings.Contains(ua, want) {
|
||||
t.Errorf("user agent should contain %q", want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- Auth recovery behavior tests ---
|
||||
|
||||
func TestErrRespErrorClassification(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
resp ErrResp
|
||||
wantError bool
|
||||
wantCode int64
|
||||
}{
|
||||
{"success", ErrResp{ErrorCode: 0}, false, 0},
|
||||
{"access_token_expired_4122", ErrResp{ErrorCode: 4122, ErrorMsg: "access_token expired"}, true, 4122},
|
||||
{"access_token_expired_4121", ErrResp{ErrorCode: 4121, ErrorMsg: "access_token expired"}, true, 4121},
|
||||
{"unauthenticated_16", ErrResp{ErrorCode: 16, ErrorMsg: "unauthenticated"}, true, 16},
|
||||
{"refresh_token_invalid_4126", ErrResp{ErrorCode: 4126, ErrorMsg: "invalid_grant"}, true, 4126},
|
||||
{"captcha_expired_9", ErrResp{ErrorCode: 9, ErrorMsg: "captcha_invalid"}, true, 9},
|
||||
{"rate_limit_10", ErrResp{ErrorCode: 10, ErrorDescription: "too frequent"}, true, 10},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
gotError := tt.resp.IsError()
|
||||
if gotError != tt.wantError {
|
||||
t.Errorf("IsError() = %v, want %v", gotError, tt.wantError)
|
||||
}
|
||||
if tt.resp.ErrorCode != tt.wantCode {
|
||||
t.Errorf("ErrorCode = %d, want %d", tt.resp.ErrorCode, tt.wantCode)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestGuardClauseOnAuthURLDoesNotRefresh verifies that when the auth endpoint
|
||||
// itself reports 4122, request() fails fast instead of calling refreshToken()
|
||||
// (which would recurse). Real behavior, real code path: with the guard
|
||||
// removed from request(), the token endpoint would be hit a second time.
|
||||
func TestGuardClauseOnAuthURLDoesNotRefresh(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
d, _ := newTestDriver(t)
|
||||
m.tokenStatus = http.StatusBadRequest
|
||||
m.tokenBody = map[string]any{"error_code": 4122, "error": "access_token_expired"}
|
||||
|
||||
_, err := d.request("https://user.mypikpak.net/v1/auth/token", http.MethodPost, nil, nil)
|
||||
if err == nil {
|
||||
t.Fatal("request() to an auth URL must fail on 4122 instead of refreshing")
|
||||
}
|
||||
if got := m.count(pathToken); got != 1 {
|
||||
t.Errorf("guard clause violated: token endpoint hit %d times, want exactly 1 (no refreshToken recursion)", got)
|
||||
}
|
||||
if got := m.count(pathSignin); got != 0 {
|
||||
t.Errorf("no re-login expected, got %d signin calls", got)
|
||||
}
|
||||
}
|
||||
|
||||
// --- Integration scaffolding: in-memory DB + mock PikPak endpoints ---
|
||||
|
||||
var (
|
||||
setupDBOnce sync.Once
|
||||
setupDBErr error
|
||||
rowSeq int64
|
||||
)
|
||||
|
||||
// setupTestDB mirrors internal/op/storage_test.go: an in-memory SQLite
|
||||
// database behind internal/db, so op.MustSaveDriverStorage really persists
|
||||
// and tests can assert on the saved row instead of on comments.
|
||||
func setupTestDB(t *testing.T) {
|
||||
t.Helper()
|
||||
setupDBOnce.Do(func() {
|
||||
var gormDB *gorm.DB
|
||||
gormDB, setupDBErr = gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
|
||||
if setupDBErr != nil {
|
||||
return
|
||||
}
|
||||
conf.Conf = conf.DefaultConfig("testdata")
|
||||
db.Init(gormDB)
|
||||
})
|
||||
if setupDBErr != nil {
|
||||
t.Fatalf("failed to set up test database: %v", setupDBErr)
|
||||
}
|
||||
}
|
||||
|
||||
// createStorageRow inserts a fresh storage row and returns it, so that
|
||||
// MustSaveDriverStorage during a test performs an UPDATE that can be read
|
||||
// back afterwards.
|
||||
func createStorageRow(t *testing.T) *model.Storage {
|
||||
t.Helper()
|
||||
rowSeq++
|
||||
st := &model.Storage{
|
||||
Driver: "PikPak",
|
||||
MountPath: fmt.Sprintf("/pikpak-test-%d", rowSeq),
|
||||
Addition: `{"username":"tester@example.com","password":"pw"}`,
|
||||
}
|
||||
if err := db.CreateStorage(st); err != nil {
|
||||
t.Fatalf("failed to create storage row: %v", err)
|
||||
}
|
||||
return st
|
||||
}
|
||||
|
||||
func persistedRefreshToken(t *testing.T, id uint) string {
|
||||
t.Helper()
|
||||
st, err := db.GetStorageById(id)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read storage back: %v", err)
|
||||
}
|
||||
var a Addition
|
||||
if err := json.Unmarshal([]byte(st.Addition), &a); err != nil {
|
||||
t.Fatalf("failed to decode persisted addition %q: %v", st.Addition, err)
|
||||
}
|
||||
return a.RefreshToken
|
||||
}
|
||||
|
||||
// mockCall records one request received by the mock server.
|
||||
type mockCall struct {
|
||||
headers http.Header
|
||||
body map[string]any
|
||||
}
|
||||
|
||||
func (c mockCall) captchaToken() string {
|
||||
s, _ := c.body["captcha_token"].(string)
|
||||
return s
|
||||
}
|
||||
|
||||
// pikpakMock emulates the captcha/auth endpoints used by login() and
|
||||
// refreshToken(), plus one drive endpoint that serves as the entry point of
|
||||
// the recovery chain. The drive endpoint fails exactly once (with the code
|
||||
// configured in driveFirstStatus) and succeeds afterwards, so request() can
|
||||
// only complete if recovery actually ran.
|
||||
type pikpakMock struct {
|
||||
t *testing.T
|
||||
srv *httptest.Server
|
||||
mu sync.Mutex
|
||||
calls map[string][]mockCall
|
||||
|
||||
captchaTokenOut string
|
||||
captchaURL string
|
||||
|
||||
tokenStatus int
|
||||
tokenBody map[string]any
|
||||
|
||||
signinStatus int
|
||||
signinBody map[string]any
|
||||
|
||||
driveFirstStatus int
|
||||
driveFirstBody map[string]any // body served on the first drive call only
|
||||
driveBody map[string]any // body served afterwards
|
||||
driveHits int
|
||||
}
|
||||
|
||||
func newPikpakMock(t *testing.T) *pikpakMock {
|
||||
t.Helper()
|
||||
m := &pikpakMock{
|
||||
t: t,
|
||||
calls: map[string][]mockCall{},
|
||||
captchaTokenOut: "cap-fresh",
|
||||
tokenStatus: http.StatusOK,
|
||||
tokenBody: map[string]any{"access_token": "at-2", "refresh_token": "rt-2", "sub": "user-1"},
|
||||
signinStatus: http.StatusOK,
|
||||
signinBody: map[string]any{"access_token": "at-new", "refresh_token": "rt-new", "sub": "user-1"},
|
||||
driveFirstStatus: http.StatusOK,
|
||||
driveFirstBody: map[string]any{"files": []any{}, "next_page_token": ""},
|
||||
driveBody: map[string]any{"files": []any{}, "next_page_token": ""},
|
||||
}
|
||||
m.srv = httptest.NewServer(http.HandlerFunc(m.serve))
|
||||
return m
|
||||
}
|
||||
|
||||
func (m *pikpakMock) close() { m.srv.Close() }
|
||||
|
||||
func (m *pikpakMock) serve(w http.ResponseWriter, r *http.Request) {
|
||||
body := map[string]any{}
|
||||
if raw, err := io.ReadAll(r.Body); err == nil && len(raw) > 0 {
|
||||
_ = json.Unmarshal(raw, &body)
|
||||
}
|
||||
m.mu.Lock()
|
||||
m.calls[r.URL.Path] = append(m.calls[r.URL.Path], mockCall{headers: r.Header.Clone(), body: body})
|
||||
status := http.StatusOK
|
||||
payload := any(map[string]any{})
|
||||
switch {
|
||||
case strings.HasSuffix(r.URL.Path, "/v1/shield/captcha/init"):
|
||||
payload = map[string]any{"captcha_token": m.captchaTokenOut, "expires_in": 3600, "url": m.captchaURL}
|
||||
case strings.HasSuffix(r.URL.Path, "/v1/auth/signin"):
|
||||
status = m.signinStatus
|
||||
payload = m.signinBody
|
||||
case strings.HasSuffix(r.URL.Path, "/v1/auth/token"):
|
||||
status = m.tokenStatus
|
||||
payload = m.tokenBody
|
||||
case strings.HasSuffix(r.URL.Path, "/drive/v1/files"):
|
||||
m.driveHits++
|
||||
if m.driveHits == 1 {
|
||||
status = m.driveFirstStatus
|
||||
payload = m.driveFirstBody
|
||||
} else {
|
||||
payload = m.driveBody
|
||||
}
|
||||
default:
|
||||
m.mu.Unlock()
|
||||
m.t.Errorf("unexpected request to %s", r.URL.Path)
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
m.mu.Unlock()
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(status)
|
||||
_ = json.NewEncoder(w).Encode(payload)
|
||||
}
|
||||
|
||||
func (m *pikpakMock) count(path string) int {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
return len(m.calls[path])
|
||||
}
|
||||
|
||||
func (m *pikpakMock) reset() {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.calls = map[string][]mockCall{}
|
||||
m.driveHits = 0
|
||||
}
|
||||
|
||||
func (m *pikpakMock) last(path string) mockCall {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
calls := m.calls[path]
|
||||
if len(calls) == 0 {
|
||||
m.t.Fatalf("no recorded call for %s", path)
|
||||
}
|
||||
return calls[len(calls)-1]
|
||||
}
|
||||
|
||||
// installMockClient replaces base.RestyClient with a client whose requests to
|
||||
// the hard-coded PikPak hosts are rewritten onto the mock server, and returns
|
||||
// a restore function. The rewrite happens in OnBeforeRequest, which resty
|
||||
// runs before its internal parseRequestURL/createHTTPRequest middlewares.
|
||||
func installMockClient(m *pikpakMock) func() {
|
||||
old := base.RestyClient
|
||||
client := resty.New()
|
||||
client.OnBeforeRequest(func(_ *resty.Client, req *resty.Request) error {
|
||||
for _, host := range []string{"https://user.mypikpak.net", "https://api-drive.mypikpak.net"} {
|
||||
if strings.HasPrefix(req.URL, host) {
|
||||
req.URL = strings.Replace(req.URL, host, m.srv.URL, 1)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
base.RestyClient = client
|
||||
return func() { base.RestyClient = old }
|
||||
}
|
||||
|
||||
// newTestDriver builds a PikPak with a fully initialized Common (web platform
|
||||
// constants) and a fresh storage row in the DB, ready for auth-flow tests.
|
||||
func newTestDriver(t *testing.T) (*PikPak, uint) {
|
||||
t.Helper()
|
||||
setupTestDB(t)
|
||||
st := createStorageRow(t)
|
||||
d := &PikPak{}
|
||||
d.SetStorage(*st)
|
||||
d.Platform = "web"
|
||||
d.Username = "tester@example.com"
|
||||
d.Password = "pw"
|
||||
d.Common = &Common{
|
||||
ClientID: WebClientID,
|
||||
ClientSecret: WebClientSecret,
|
||||
ClientVersion: WebClientVersion,
|
||||
PackageName: WebPackageName,
|
||||
DeviceID: "test-device",
|
||||
UserAgent: "test-agent",
|
||||
Algorithms: WebAlgorithms,
|
||||
}
|
||||
d.Common.RefreshCTokenCk = func(token string) {
|
||||
d.Common.CaptchaToken = token
|
||||
}
|
||||
return d, st.ID
|
||||
}
|
||||
|
||||
const (
|
||||
pathCaptchaInit = "/v1/shield/captcha/init"
|
||||
pathSignin = "/v1/auth/signin"
|
||||
pathToken = "/v1/auth/token"
|
||||
pathFiles = "/drive/v1/files"
|
||||
)
|
||||
|
||||
// --- Main auth recovery path ---
|
||||
|
||||
// TestMainRecoveryPath exercises the full chain the PR is about: a drive
|
||||
// request fails with 4122, refreshToken fails with 4126, login() runs (fresh
|
||||
// captcha + password signin), the new refresh token is persisted to the DB,
|
||||
// and request() retries the original call successfully with the new tokens.
|
||||
func TestMainRecoveryPath(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
d, id := newTestDriver(t)
|
||||
d.RefreshToken = "rt-old"
|
||||
d.AccessToken = "at-stale"
|
||||
d.SetCaptchaToken("cap-stale")
|
||||
d.Addition.RefreshToken = "rt-old"
|
||||
|
||||
// refresh attempt fails with "refresh token invalid"
|
||||
m.tokenStatus = http.StatusBadRequest
|
||||
m.tokenBody = map[string]any{"error_code": 4126, "error": "invalid_grant"}
|
||||
// the first drive call reports an expired access token; the retry succeeds
|
||||
m.driveFirstStatus = http.StatusBadRequest
|
||||
m.driveFirstBody = map[string]any{"error_code": 4122, "error": "access_token_expired"}
|
||||
m.driveBody = map[string]any{"files": []any{}, "next_page_token": ""}
|
||||
|
||||
var resp Files
|
||||
if _, err := d.request("https://api-drive.mypikpak.net/drive/v1/files", http.MethodGet, nil, &resp); err != nil {
|
||||
t.Fatalf("request() returned error even though recovery should succeed: %v", err)
|
||||
}
|
||||
|
||||
if got := m.count(pathToken); got != 1 {
|
||||
t.Errorf("expected exactly 1 refresh request, got %d", got)
|
||||
}
|
||||
if got := m.count(pathSignin); got != 1 {
|
||||
t.Errorf("expected exactly 1 signin (re-login), got %d", got)
|
||||
}
|
||||
if got := m.count(pathFiles); got != 2 {
|
||||
t.Errorf("expected 2 files requests (failed + retried), got %d", got)
|
||||
}
|
||||
if got := m.count(pathCaptchaInit); got != 1 {
|
||||
t.Errorf("expected exactly 1 captcha/init call during re-login, got %d", got)
|
||||
}
|
||||
|
||||
// The retry must carry the tokens obtained via re-login, not the stale ones.
|
||||
lastFiles := m.last(pathFiles)
|
||||
if got := lastFiles.headers.Get("Authorization"); got != "Bearer at-new" {
|
||||
t.Errorf("retried request Authorization = %q, want %q", got, "Bearer at-new")
|
||||
}
|
||||
if got := lastFiles.headers.Get("X-Captcha-Token"); got != "cap-fresh" {
|
||||
t.Errorf("retried request X-Captcha-Token = %q, want %q", got, "cap-fresh")
|
||||
}
|
||||
|
||||
// Tokens were rotated in memory...
|
||||
if d.AccessToken != "at-new" {
|
||||
t.Errorf("AccessToken = %q, want %q", d.AccessToken, "at-new")
|
||||
}
|
||||
if d.RefreshToken != "rt-new" {
|
||||
t.Errorf("RefreshToken = %q, want %q", d.RefreshToken, "rt-new")
|
||||
}
|
||||
// ...and the rotated refresh token was really persisted.
|
||||
if got := persistedRefreshToken(t, id); got != "rt-new" {
|
||||
t.Errorf("persisted addition refresh_token = %q, want %q", got, "rt-new")
|
||||
}
|
||||
}
|
||||
|
||||
// TestRefreshToken4126WithoutCredentialsDoesNotLogin checks that a 4126 with
|
||||
// empty username/password yields the "re-provide refresh_token" error instead
|
||||
// of attempting a password login.
|
||||
func TestRefreshToken4126WithoutCredentialsDoesNotLogin(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
d, _ := newTestDriver(t)
|
||||
d.Username = ""
|
||||
d.Password = ""
|
||||
|
||||
m.tokenStatus = http.StatusBadRequest
|
||||
m.tokenBody = map[string]any{"error_code": 4126, "error": "invalid_grant"}
|
||||
|
||||
err := d.refreshToken("rt-old")
|
||||
if err == nil {
|
||||
t.Fatal("refreshToken() with invalid refresh token and no credentials must fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "re-provide") {
|
||||
t.Errorf("unexpected error text: %v", err)
|
||||
}
|
||||
if got := m.count(pathSignin); got != 0 {
|
||||
t.Errorf("signin must not be attempted without credentials, got %d calls", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestRefreshTokenOtherErrorDoesNotLogin checks that a non-4126 refresh
|
||||
// failure propagates without triggering a re-login (4126 is the single
|
||||
// documented trigger).
|
||||
func TestRefreshTokenOtherErrorDoesNotLogin(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
d, _ := newTestDriver(t)
|
||||
m.tokenStatus = http.StatusBadRequest
|
||||
m.tokenBody = map[string]any{"error_code": 10, "error_description": "too frequent"}
|
||||
|
||||
if err := d.refreshToken("rt-old"); err == nil {
|
||||
t.Fatal("refreshToken() must propagate a non-4126 error")
|
||||
}
|
||||
if got := m.count(pathSignin); got != 0 {
|
||||
t.Errorf("signin must not be attempted for non-4126 errors, got %d calls", got)
|
||||
}
|
||||
}
|
||||
|
||||
// --- Token validation (replaces TestTokenValidationRejectsEmpty) ---
|
||||
|
||||
// TestTokenValidationRejectsEmpty drives login() and refreshToken() against
|
||||
// 200 responses that carry empty tokens and requires both paths to refuse
|
||||
// them without persisting anything.
|
||||
func TestTokenValidationRejectsEmpty(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
// login(): signin answers 200 but with an empty access_token.
|
||||
d, id := newTestDriver(t)
|
||||
m.signinBody = map[string]any{"access_token": "", "refresh_token": "rt-x", "sub": "user-1"}
|
||||
if err := d.login(); err == nil {
|
||||
t.Fatal("login() must reject empty access_token")
|
||||
}
|
||||
if got := persistedRefreshToken(t, id); got != "" {
|
||||
t.Errorf("login() must not persist tokens when validation fails, persisted %q", got)
|
||||
}
|
||||
|
||||
// login(): symmetric case — empty refresh_token but non-empty access_token.
|
||||
d3, id3 := newTestDriver(t)
|
||||
m.reset()
|
||||
m.signinBody = map[string]any{"access_token": "at-x", "refresh_token": "", "sub": "user-1"}
|
||||
if err := d3.login(); err == nil {
|
||||
t.Fatal("login() must reject empty refresh_token")
|
||||
}
|
||||
if got := persistedRefreshToken(t, id3); got != "" {
|
||||
t.Errorf("login() must not persist tokens when validation fails, persisted %q", got)
|
||||
}
|
||||
|
||||
// refreshToken(): 200 but empty refresh_token.
|
||||
d2, id2 := newTestDriver(t)
|
||||
m.tokenStatus = http.StatusOK
|
||||
m.tokenBody = map[string]any{"access_token": "at-x", "refresh_token": "", "sub": "user-1"}
|
||||
if err := d2.refreshToken("rt-old"); err == nil {
|
||||
t.Fatal("refreshToken() must reject empty refresh_token")
|
||||
}
|
||||
if got := persistedRefreshToken(t, id2); got != "" {
|
||||
t.Errorf("refreshToken() must not persist tokens when validation fails, persisted %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
// --- Captcha refresh (replaces TestCaptchaAlwaysRefreshedBeforeLogin) ---
|
||||
|
||||
// TestCaptchaAlwaysRefreshedBeforeLogin proves login() fetches a fresh captcha
|
||||
// even when a (possibly expired) token is already present, and that signin is
|
||||
// performed with the fresh token rather than the stale one.
|
||||
func TestCaptchaAlwaysRefreshedBeforeLogin(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
d, _ := newTestDriver(t)
|
||||
d.SetCaptchaToken("cap-stale") // non-empty and (conceptually) expired
|
||||
|
||||
if err := d.login(); err != nil {
|
||||
t.Fatalf("login() failed: %v", err)
|
||||
}
|
||||
|
||||
if got := m.count(pathCaptchaInit); got != 1 {
|
||||
t.Fatalf("expected exactly 1 captcha/init call despite a non-empty stale token, got %d", got)
|
||||
}
|
||||
if got := m.last(pathSignin).captchaToken(); got != "cap-fresh" {
|
||||
t.Errorf("signin used captcha_token %q, want the fresh %q", got, "cap-fresh")
|
||||
}
|
||||
if got := d.GetCaptchaToken(); got != "cap-fresh" {
|
||||
t.Errorf("driver CaptchaToken = %q after login, want %q", got, "cap-fresh")
|
||||
}
|
||||
}
|
||||
|
||||
// --- Stale bearer cleared before login ---
|
||||
|
||||
// TestLoginClearsStaleAccessToken checks that the captcha/init and signin
|
||||
// requests issued by login() do not carry the expired bearer token.
|
||||
func TestLoginClearsStaleAccessToken(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
d, _ := newTestDriver(t)
|
||||
d.AccessToken = "at-stale"
|
||||
|
||||
if err := d.login(); err != nil {
|
||||
t.Fatalf("login() failed: %v", err)
|
||||
}
|
||||
|
||||
for _, path := range []string{pathCaptchaInit, pathSignin} {
|
||||
if got := m.last(path).headers.Get("Authorization"); got != "" {
|
||||
t.Errorf("%s request carried Authorization %q, want it cleared before login", path, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- Captcha meta completeness ---
|
||||
|
||||
// TestCaptchaMetaCompleteness asserts captcha/init on the login path carries
|
||||
// the same meta fields RefreshCaptchaTokenAtLogin sends on main.
|
||||
func TestCaptchaMetaCompleteness(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
d, _ := newTestDriver(t)
|
||||
if err := d.login(); err != nil {
|
||||
t.Fatalf("login() failed: %v", err)
|
||||
}
|
||||
|
||||
meta, _ := m.last(pathCaptchaInit).body["meta"].(map[string]any)
|
||||
for _, key := range []string{"email", "client_version", "package_name", "timestamp", "captcha_sign"} {
|
||||
if v, ok := meta[key]; !ok || v == "" {
|
||||
t.Errorf("captcha meta missing or empty %q (got %#v)", key, meta)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- refreshToken success path (highest-frequency production path) ---
|
||||
|
||||
// TestRefreshTokenSuccessRotatesAndPersists covers 4122 -> refreshToken()
|
||||
// succeeding: rotated tokens land in memory, the retry carries the new bearer,
|
||||
// and the new refresh token is persisted to the DB.
|
||||
func TestRefreshTokenSuccessRotatesAndPersists(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
d, id := newTestDriver(t)
|
||||
d.RefreshToken = "rt-old"
|
||||
d.AccessToken = "at-stale"
|
||||
d.Addition.RefreshToken = "rt-old"
|
||||
|
||||
m.tokenStatus = http.StatusOK
|
||||
m.tokenBody = map[string]any{"access_token": "at-2", "refresh_token": "rt-2", "sub": "user-1"}
|
||||
m.driveFirstStatus = http.StatusBadRequest
|
||||
m.driveFirstBody = map[string]any{"error_code": 4122, "error": "access_token_expired"}
|
||||
m.driveBody = map[string]any{"files": []any{}, "next_page_token": ""}
|
||||
|
||||
var resp Files
|
||||
if _, err := d.request("https://api-drive.mypikpak.net/drive/v1/files", http.MethodGet, nil, &resp); err != nil {
|
||||
t.Fatalf("request() failed even though refresh should succeed: %v", err)
|
||||
}
|
||||
|
||||
if got := m.count(pathSignin); got != 0 {
|
||||
t.Errorf("a successful refresh must not fall through to password login, got %d signin calls", got)
|
||||
}
|
||||
if d.AccessToken != "at-2" {
|
||||
t.Errorf("AccessToken = %q, want %q", d.AccessToken, "at-2")
|
||||
}
|
||||
if d.RefreshToken != "rt-2" {
|
||||
t.Errorf("RefreshToken = %q, want %q", d.RefreshToken, "rt-2")
|
||||
}
|
||||
if got := m.last(pathFiles).headers.Get("Authorization"); got != "Bearer at-2" {
|
||||
t.Errorf("retried request Authorization = %q, want %q", got, "Bearer at-2")
|
||||
}
|
||||
if got := persistedRefreshToken(t, id); got != "rt-2" {
|
||||
t.Errorf("persisted addition refresh_token = %q, want %q", got, "rt-2")
|
||||
}
|
||||
}
|
||||
|
||||
// --- captcha expired (case 9) ---
|
||||
|
||||
// TestCaptchaExpiredRefreshesAndRetries covers request() case 9: a captcha
|
||||
// error on a drive call triggers a captcha refresh and one retry.
|
||||
func TestCaptchaExpiredRefreshesAndRetries(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
d, _ := newTestDriver(t)
|
||||
d.AccessToken = "at-ok"
|
||||
d.RefreshToken = "rt-ok"
|
||||
d.SetCaptchaToken("cap-stale")
|
||||
|
||||
m.driveFirstStatus = http.StatusBadRequest
|
||||
m.driveFirstBody = map[string]any{"error_code": 9, "error": "captcha_invalid"}
|
||||
m.driveBody = map[string]any{"files": []any{}, "next_page_token": ""}
|
||||
|
||||
var resp Files
|
||||
if _, err := d.request("https://api-drive.mypikpak.net/drive/v1/files", http.MethodGet, nil, &resp); err != nil {
|
||||
t.Fatalf("request() failed even though captcha refresh should recover: %v", err)
|
||||
}
|
||||
|
||||
if got := m.count(pathCaptchaInit); got == 0 {
|
||||
t.Fatal("expected a captcha refresh after error code 9")
|
||||
}
|
||||
if got := m.count(pathFiles); got != 2 {
|
||||
t.Errorf("expected 2 files requests (failed + retried), got %d", got)
|
||||
}
|
||||
if got := m.count(pathSignin); got != 0 {
|
||||
t.Errorf("captcha recovery must not re-login, got %d signin calls", got)
|
||||
}
|
||||
if got := m.last(pathFiles).headers.Get("X-Captcha-Token"); got != "cap-fresh" {
|
||||
t.Errorf("retried request X-Captcha-Token = %q, want %q", got, "cap-fresh")
|
||||
}
|
||||
}
|
||||
|
||||
// --- SkipVerification (added by this PR) ---
|
||||
|
||||
// TestSkipVerificationControlsVerificationURL covers the new config option:
|
||||
// a captcha/init response carrying a human-verification url is fatal by
|
||||
// default and ignored only when skip_verification is enabled.
|
||||
func TestSkipVerificationControlsVerificationURL(t *testing.T) {
|
||||
m := newPikpakMock(t)
|
||||
defer m.close()
|
||||
restore := installMockClient(m)
|
||||
defer restore()
|
||||
|
||||
m.captchaURL = "https://user.mypikpak.net/forbidden/test"
|
||||
|
||||
d, _ := newTestDriver(t)
|
||||
if err := d.login(); err == nil {
|
||||
t.Fatal("login() must fail on a verification url by default")
|
||||
} else if !strings.Contains(err.Error(), "need verify") {
|
||||
t.Errorf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
d2, _ := newTestDriver(t)
|
||||
d2.SkipVerification = true
|
||||
if err := d2.login(); err != nil {
|
||||
t.Fatalf("login() with skip_verification must ignore the url, got: %v", err)
|
||||
}
|
||||
if d2.AccessToken != "at-new" {
|
||||
t.Errorf("AccessToken = %q after skipped verification, want %q", d2.AccessToken, "at-new")
|
||||
}
|
||||
}
|
||||
@@ -55,10 +55,18 @@ func (d *Teldrive) Drop(ctx context.Context) error {
|
||||
}
|
||||
|
||||
func (d *Teldrive) List(ctx context.Context, dir model.Obj, args model.ListArgs) ([]model.Obj, error) {
|
||||
dirPath := dir.GetPath()
|
||||
if dirPath == "" {
|
||||
dirPath = d.GetRootPath()
|
||||
}
|
||||
if dirPath == "" {
|
||||
dirPath = "/"
|
||||
}
|
||||
|
||||
var firstResp ListResp
|
||||
err := d.request(http.MethodGet, "/api/files", func(req *resty.Request) {
|
||||
req.SetQueryParams(map[string]string{
|
||||
"path": dir.GetPath(),
|
||||
"path": dirPath,
|
||||
"limit": "500",
|
||||
"page": "1",
|
||||
})
|
||||
@@ -87,7 +95,7 @@ func (d *Teldrive) List(ctx context.Context, dir model.Obj, args model.ListArgs)
|
||||
var resp ListResp
|
||||
err := d.request(http.MethodGet, "/api/files", func(req *resty.Request) {
|
||||
req.SetQueryParams(map[string]string{
|
||||
"path": dir.GetPath(),
|
||||
"path": dirPath,
|
||||
"limit": "500",
|
||||
"page": strconv.Itoa(page),
|
||||
})
|
||||
@@ -114,7 +122,7 @@ func (d *Teldrive) List(ctx context.Context, dir model.Obj, args model.ListArgs)
|
||||
|
||||
return utils.SliceConvert(allItems, func(src Object) (model.Obj, error) {
|
||||
return &model.Object{
|
||||
Path: path.Join(dir.GetPath(), src.Name),
|
||||
Path: path.Join(dirPath, src.Name),
|
||||
ID: src.ID,
|
||||
Name: src.Name,
|
||||
Size: func() int64 {
|
||||
|
||||
@@ -4,9 +4,11 @@ import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/base"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/driver"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/go-resty/resty/v2"
|
||||
)
|
||||
@@ -36,3 +38,44 @@ func TestListEmptyDir(t *testing.T) {
|
||||
t.Fatalf("expected no entries for an empty dir, got %d", len(objs))
|
||||
}
|
||||
}
|
||||
|
||||
func TestListRootUsesConfiguredRootPath(t *testing.T) {
|
||||
var (
|
||||
mu sync.Mutex
|
||||
paths []string
|
||||
)
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
mu.Lock()
|
||||
paths = append(paths, r.URL.Query().Get("path"))
|
||||
mu.Unlock()
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"items":[{"id":"child","name":"child","type":"folder"}],"meta":{"count":1,"totalPages":1,"currentPage":1}}`))
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
oldClient := base.RestyClient
|
||||
base.RestyClient = resty.New()
|
||||
defer func() { base.RestyClient = oldClient }()
|
||||
|
||||
d := &Teldrive{
|
||||
Addition: Addition{
|
||||
RootPath: driver.RootPath{RootFolderPath: "/configured-root"},
|
||||
},
|
||||
}
|
||||
d.Address = srv.URL
|
||||
|
||||
objs, err := d.List(context.Background(), &model.Object{}, model.ListArgs{})
|
||||
if err != nil {
|
||||
t.Fatalf("List returned error: %v", err)
|
||||
}
|
||||
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
if len(paths) != 1 || paths[0] != "/configured-root" {
|
||||
t.Fatalf("expected request path %q, got %q", "/configured-root", paths)
|
||||
}
|
||||
if len(objs) != 1 || objs[0].GetPath() != "/configured-root/child" {
|
||||
t.Fatalf("expected child path %q, got %#v", "/configured-root/child", objs)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,7 +10,7 @@ require (
|
||||
github.com/KarpelesLab/reflink v1.0.2
|
||||
github.com/KirCute/zip v1.0.1
|
||||
github.com/OpenListTeam/go-cache v0.1.0
|
||||
github.com/OpenListTeam/gofakes3 v0.8.2-0.20260911144636-404e34e0f5af
|
||||
github.com/OpenListTeam/gofakes3 v0.8.2-0.20260911142347-cd3c030a83b4
|
||||
github.com/OpenListTeam/sftpd-openlist v1.0.1
|
||||
github.com/OpenListTeam/tache v0.2.2
|
||||
github.com/OpenListTeam/times v0.1.0
|
||||
@@ -36,6 +36,7 @@ require (
|
||||
github.com/dhowden/tag v0.0.0-20240417053706-3d75831295e8
|
||||
github.com/disintegration/imaging v1.6.2
|
||||
github.com/dlclark/regexp2 v1.12.0
|
||||
github.com/dlclark/regexp2/v2 v2.8.4
|
||||
github.com/dustinxie/ecc v0.0.0-20210511000915-959544187564
|
||||
github.com/fclairamb/ftpserverlib v0.26.1-0.20250709223522-4a925d79caf6
|
||||
github.com/foxxorcat/mopan-sdk-go v0.1.6
|
||||
|
||||
@@ -51,8 +51,8 @@ github.com/OpenListTeam/115-sdk-go v0.2.6 h1:ehXyStvncvn4qRBuknor3kyGZtUmHc0+stj
|
||||
github.com/OpenListTeam/115-sdk-go v0.2.6/go.mod h1:cfvitk2lwe6036iNi2h+iNxwxWDifKZsSvNtrur5BqU=
|
||||
github.com/OpenListTeam/go-cache v0.1.0 h1:eV2+FCP+rt+E4OCJqLUW7wGccWZNJMV0NNkh+uChbAI=
|
||||
github.com/OpenListTeam/go-cache v0.1.0/go.mod h1:AHWjKhNK3LE4rorVdKyEALDHoeMnP8SjiNyfVlB+Pz4=
|
||||
github.com/OpenListTeam/gofakes3 v0.8.2-0.20260911144636-404e34e0f5af h1:pcORNNaOgbc28g9YliwPc5mpEmGf/dKyZjkXhHZgvVo=
|
||||
github.com/OpenListTeam/gofakes3 v0.8.2-0.20260911144636-404e34e0f5af/go.mod h1:mS9Ywbo6aId6BrRzeYjOIOpK0QDVnMoKOIb0hpaQZ3U=
|
||||
github.com/OpenListTeam/gofakes3 v0.8.2-0.20260911142347-cd3c030a83b4 h1:Zy7/qg6aCS0OF/FPIoJh9/d0IgcIxpWRvn79ACm2R/Y=
|
||||
github.com/OpenListTeam/gofakes3 v0.8.2-0.20260911142347-cd3c030a83b4/go.mod h1:mS9Ywbo6aId6BrRzeYjOIOpK0QDVnMoKOIb0hpaQZ3U=
|
||||
github.com/OpenListTeam/gsync v0.1.0 h1:ywzGybOvA3lW8K1BUjKZ2IUlT2FSlzPO4DOazfYXjcs=
|
||||
github.com/OpenListTeam/gsync v0.1.0/go.mod h1:h/Rvv9aX/6CdW/7B8di3xK3xNV8dUg45Fehrd/ksZ9s=
|
||||
github.com/OpenListTeam/reflink v0.0.0-20260701021214-78760eaeafef h1:67uGHancMF/abMrnkc8abVUWQiG73Wk5d8CKt3RzkFo=
|
||||
@@ -341,6 +341,7 @@ github.com/disintegration/imaging v1.6.2 h1:w1LecBlG2Lnp8B3jk5zSuNqd7b4DXhcjwek1
|
||||
github.com/disintegration/imaging v1.6.2/go.mod h1:44/5580QXChDfwIclfc/PCwrr44amcmDAg8hxG0Ewe4=
|
||||
github.com/dlclark/regexp2 v1.12.0 h1:0j4c5qQmnC6XOWNjP3PIXURXN2gWx76rd3KvgdPkCz8=
|
||||
github.com/dlclark/regexp2 v1.12.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8=
|
||||
github.com/dlclark/regexp2/v2 v2.8.4/go.mod h1:avUrQvPaLz2DrFNHJF0taWAFFX2C1GMSSoeiqFjcBmU=
|
||||
github.com/dsnet/compress v0.0.2-0.20230904184137-39efe44ab707 h1:2tV76y6Q9BB+NEBasnqvs7e49aEBFI8ejC89PSnWH+4=
|
||||
github.com/dsnet/compress v0.0.2-0.20230904184137-39efe44ab707/go.mod h1:qssHWj60/X5sZFNxpG4HBPDHVqxNm4DfnCKgrbZOT+s=
|
||||
github.com/dsnet/golib v0.0.0-20171103203638-1ea166775780/go.mod h1:Lj+Z9rebOhdfkVLjJ8T6VcRQv3SXugXy999NBtR9aFY=
|
||||
|
||||
@@ -6,13 +6,12 @@ import (
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/setting"
|
||||
"github.com/OpenListTeam/OpenList/v4/server/common"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/go-webauthn/webauthn/webauthn"
|
||||
)
|
||||
|
||||
func NewAuthnInstance(c *gin.Context) (*webauthn.WebAuthn, error) {
|
||||
siteUrl, err := url.Parse(common.GetApiUrl(c.Request.Context()))
|
||||
siteUrl, err := url.Parse(conf.GetApiUrl(c.Request.Context()))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -211,6 +211,7 @@ func InitialSettings() []model.SettingItem {
|
||||
{Key: conf.SSODefaultDir, Value: "/", Type: conf.TypeString, Group: model.SSO, Flag: model.PRIVATE},
|
||||
{Key: conf.SSODefaultPermission, Value: "0", Type: conf.TypeNumber, Group: model.SSO, Flag: model.PRIVATE},
|
||||
{Key: conf.SSOCompatibilityMode, Value: "false", Type: conf.TypeBool, Group: model.SSO, Flag: model.PUBLIC},
|
||||
{Key: conf.SSOPostMessageOrigin, Value: "", Type: conf.TypeString, Group: model.SSO, Flag: model.PUBLIC},
|
||||
|
||||
// ldap settings
|
||||
{Key: conf.LdapLoginEnabled, Value: "false", Type: conf.TypeBool, Group: model.LDAP, Flag: model.PUBLIC},
|
||||
|
||||
@@ -117,6 +117,7 @@ const (
|
||||
SSODefaultDir = "sso_default_dir"
|
||||
SSODefaultPermission = "sso_default_permission"
|
||||
SSOCompatibilityMode = "sso_compatibility_mode"
|
||||
SSOPostMessageOrigin = "sso_postmessage_origin"
|
||||
|
||||
// ldap
|
||||
LdapLoginEnabled = "ldap_login_enabled"
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
package conf
|
||||
|
||||
import "context"
|
||||
|
||||
func GetApiUrl(ctx context.Context) string {
|
||||
api, _ := ctx.Value(ApiUrlKey).(string)
|
||||
return api
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
package conf_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
)
|
||||
|
||||
func TestGetApiUrl(t *testing.T) {
|
||||
const want = "https://openlist.example"
|
||||
tests := []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
want string
|
||||
}{
|
||||
{name: "present", ctx: context.WithValue(context.Background(), conf.ApiUrlKey, want), want: want},
|
||||
{name: "absent", ctx: context.Background()},
|
||||
{name: "wrong type", ctx: context.WithValue(context.Background(), conf.ApiUrlKey, 1)},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := conf.GetApiUrl(tt.ctx); got != tt.want {
|
||||
t.Fatalf("origin = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -35,7 +35,7 @@ func DeleteSearchNodesByParent(path string) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dir, name := stdpath.Split(path)
|
||||
dir, name := stdpath.Dir(path), stdpath.Base(path)
|
||||
return db.Where(fmt.Sprintf("%s = ? AND %s = ?",
|
||||
columnName("parent"), columnName("name")),
|
||||
dir, name).Delete(&model.SearchNode{}).Error
|
||||
|
||||
@@ -18,6 +18,7 @@ var (
|
||||
StorageNotInit = errors.New("storage not init")
|
||||
StreamIncomplete = errors.New("upload/download stream incomplete, possible network issue")
|
||||
StreamPeekFail = errors.New("StreamPeekFail")
|
||||
TemporaryCapacity = errors.New("temporary capacity unavailable")
|
||||
|
||||
UnknownArchiveFormat = errors.New("unknown archive format")
|
||||
WrongArchivePassword = errors.New("wrong archive password")
|
||||
|
||||
@@ -21,7 +21,6 @@ import (
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/task"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/task_group"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"github.com/OpenListTeam/OpenList/v4/server/common"
|
||||
"github.com/OpenListTeam/tache"
|
||||
"github.com/pkg/errors"
|
||||
log "github.com/sirupsen/logrus"
|
||||
@@ -415,7 +414,7 @@ func archiveDecompress(ctx context.Context, srcObjPath, dstDirPath string, args
|
||||
return nil, err
|
||||
} else {
|
||||
tsk.Creator, _ = ctx.Value(conf.UserKey).(*model.User)
|
||||
tsk.ApiUrl = common.GetApiUrl(ctx)
|
||||
tsk.ApiUrl = conf.GetApiUrl(ctx)
|
||||
ArchiveDownloadTaskManager.Add(tsk)
|
||||
return tsk, nil
|
||||
}
|
||||
|
||||
@@ -14,7 +14,6 @@ import (
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/task"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/task_group"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"github.com/OpenListTeam/OpenList/v4/server/common"
|
||||
"github.com/OpenListTeam/tache"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
@@ -166,7 +165,7 @@ func transfer(ctx context.Context, taskType taskType, srcObjPath, dstDirPath str
|
||||
}
|
||||
|
||||
t.Creator, _ = ctx.Value(conf.UserKey).(*model.User)
|
||||
t.ApiUrl = common.GetApiUrl(ctx)
|
||||
t.ApiUrl = conf.GetApiUrl(ctx)
|
||||
if taskType == copy || taskType == merge {
|
||||
CopyTaskManager.Add(t)
|
||||
} else {
|
||||
|
||||
+2
-2
@@ -4,9 +4,9 @@ import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/op"
|
||||
"github.com/OpenListTeam/OpenList/v4/server/common"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
@@ -20,7 +20,7 @@ func link(ctx context.Context, path string, args model.LinkArgs) (*model.Link, m
|
||||
return nil, nil, errors.WithMessage(err, "failed link")
|
||||
}
|
||||
if l.URL != "" && !strings.HasPrefix(l.URL, "http://") && !strings.HasPrefix(l.URL, "https://") {
|
||||
l.URL = common.GetApiUrl(ctx) + l.URL
|
||||
l.URL = conf.GetApiUrl(ctx) + l.URL
|
||||
}
|
||||
return l, obj, nil
|
||||
}
|
||||
|
||||
+1
-2
@@ -7,7 +7,6 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"github.com/OpenListTeam/OpenList/v4/server/common"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/driver"
|
||||
@@ -81,7 +80,7 @@ func putAsTask(ctx context.Context, dstDirPath string, file model.FileStreamer)
|
||||
t := &UploadTask{
|
||||
TaskExtension: task.TaskExtension{
|
||||
Creator: taskCreator,
|
||||
ApiUrl: common.GetApiUrl(ctx),
|
||||
ApiUrl: conf.GetApiUrl(ctx),
|
||||
},
|
||||
storage: storage,
|
||||
dstDirActualPath: dstDirActualPath,
|
||||
|
||||
+2
-20
@@ -30,7 +30,7 @@ type Link struct {
|
||||
Header http.Header `json:"header"` // needed header (for url)
|
||||
RangeReader RangeReaderIF `json:"-"` // recommended way if can't use URL
|
||||
|
||||
Expiration *time.Duration // local cache expire Duration
|
||||
Expiration *time.Duration // local cache expiration; not transferred by Clone
|
||||
|
||||
//for accelerating request, use multi-thread downloading
|
||||
Concurrency int `json:"concurrency"`
|
||||
@@ -42,12 +42,12 @@ type Link struct {
|
||||
RequireReference bool `json:"-"`
|
||||
}
|
||||
|
||||
// Clone transfers ownership of l without inheriting its cache expiration.
|
||||
func (l *Link) Clone() *Link {
|
||||
return &Link{
|
||||
URL: l.URL,
|
||||
Header: l.Header,
|
||||
RangeReader: l.RangeReader,
|
||||
Expiration: l.Expiration,
|
||||
Concurrency: l.Concurrency,
|
||||
PartSize: l.PartSize,
|
||||
ContentLength: l.ContentLength,
|
||||
@@ -118,21 +118,3 @@ type SharingLinkArgs struct {
|
||||
type RangeReaderIF interface {
|
||||
RangeRead(ctx context.Context, httpRange http_range.Range) (io.ReadCloser, error)
|
||||
}
|
||||
|
||||
type RangeReadCloserIF interface {
|
||||
RangeReaderIF
|
||||
utils.ClosersIF
|
||||
}
|
||||
|
||||
var _ RangeReadCloserIF = (*RangeReadCloser)(nil)
|
||||
|
||||
type RangeReadCloser struct {
|
||||
RangeReader RangeReaderIF
|
||||
utils.Closers
|
||||
}
|
||||
|
||||
func (r *RangeReadCloser) RangeRead(ctx context.Context, httpRange http_range.Range) (io.ReadCloser, error) {
|
||||
rc, err := r.RangeReader.RangeRead(ctx, httpRange)
|
||||
r.Add(rc)
|
||||
return rc, err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestLinkCloneTransfersOwnershipWithoutCachePolicy(t *testing.T) {
|
||||
ttl := time.Minute
|
||||
source := &Link{URL: "https://example.test/file", Expiration: &ttl}
|
||||
|
||||
clone := source.Clone()
|
||||
if clone.URL != source.URL {
|
||||
t.Fatal("clone did not preserve transport data")
|
||||
}
|
||||
if clone.Expiration != nil {
|
||||
t.Fatal("clone inherited source cache policy")
|
||||
}
|
||||
if err := clone.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !source.Expired() {
|
||||
t.Fatal("closing clone did not release its source")
|
||||
}
|
||||
}
|
||||
@@ -206,7 +206,8 @@ func (d *downloader) download() (io.ReadCloser, error) {
|
||||
if err != nil {
|
||||
d.cancel(err)
|
||||
d.cfg.ConcurrencyLimit.Release()
|
||||
return nil, d.interrupt()
|
||||
_ = d.interrupt()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
d.mu.Lock()
|
||||
@@ -268,10 +269,6 @@ func (d *downloader) sendChunkTask(newConcurrency bool) (err error) {
|
||||
if err != nil {
|
||||
return err // 分片算法错误或者下载中断
|
||||
}
|
||||
if newConcurrency {
|
||||
go d.downloadPart()
|
||||
d.concurrency--
|
||||
}
|
||||
ch := chunk{
|
||||
start: d.pos,
|
||||
size: finalSize,
|
||||
@@ -286,6 +283,11 @@ func (d *downloader) sendChunkTask(newConcurrency bool) (err error) {
|
||||
case <-d.ctx.Done():
|
||||
return context.Cause(d.ctx)
|
||||
case d.chunkCh <- ch:
|
||||
if newConcurrency {
|
||||
// The worker owns the acquired slot only after its chunk is queued.
|
||||
go d.downloadPart()
|
||||
d.concurrency--
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
package net
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
)
|
||||
|
||||
func TestDownloadCancelledAcquisitionReturnsErrorAndReleasesLimit(t *testing.T) {
|
||||
const attempts = 32
|
||||
limits := make([]*ConcurrencyLimit, 0, attempts)
|
||||
for range attempts {
|
||||
limit := &ConcurrencyLimit{Limit: 1}
|
||||
limits = append(limits, limit)
|
||||
d := NewDownloader(func(d *Downloader) {
|
||||
d.Concurrency = 2
|
||||
d.PartSize = 4
|
||||
d.ConcurrencyLimit = limit
|
||||
d.HttpClient = func(ctx context.Context, _ *HttpRequestParams) (*http.Response, error) {
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
})
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
reader, err := d.Download(ctx, &HttpRequestParams{Size: 16, Range: http_range.Range{Length: -1}})
|
||||
if reader == nil && err == nil {
|
||||
t.Error("cancelled download returned a nil reader and nil error")
|
||||
}
|
||||
if reader != nil {
|
||||
_ = reader.Close()
|
||||
} else if !errors.Is(err, context.Canceled) {
|
||||
t.Errorf("cancelled download error = %v, want context.Canceled", err)
|
||||
}
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond) // allow any started workers to release their slots
|
||||
for i, limit := range limits {
|
||||
limit.mu.Lock()
|
||||
got := limit.Limit
|
||||
limit.mu.Unlock()
|
||||
if got != 1 {
|
||||
t.Errorf("attempt %d remaining concurrency = %d, want 1", i, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadSinglePartFailureReleasesLimit(t *testing.T) {
|
||||
upstreamErr := errors.New("upstream failure")
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
want error
|
||||
}{
|
||||
{name: "cancelled", ctx: func() context.Context {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
return ctx
|
||||
}(), want: context.Canceled},
|
||||
{name: "upstream failure", ctx: context.Background(), want: upstreamErr},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
limit := &ConcurrencyLimit{Limit: 1}
|
||||
d := NewDownloader(func(d *Downloader) {
|
||||
d.PartSize = 32
|
||||
d.ConcurrencyLimit = limit
|
||||
d.HttpClient = func(ctx context.Context, _ *HttpRequestParams) (*http.Response, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return nil, upstreamErr
|
||||
}
|
||||
})
|
||||
reader, err := d.Download(tc.ctx, &HttpRequestParams{Size: 16, Range: http_range.Range{Length: -1}})
|
||||
if reader != nil || !errors.Is(err, tc.want) {
|
||||
t.Fatalf("single-part failed download = %v, %v; want nil, %v", reader, err, tc.want)
|
||||
}
|
||||
limit.mu.Lock()
|
||||
got := limit.Limit
|
||||
limit.mu.Unlock()
|
||||
if got != 1 {
|
||||
t.Errorf("remaining concurrency = %d, want 1", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+89
-26
@@ -4,6 +4,7 @@ import (
|
||||
"compress/gzip"
|
||||
"context"
|
||||
"crypto/tls"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
@@ -15,7 +16,6 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/errs"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
@@ -25,12 +25,8 @@ import (
|
||||
|
||||
//this file is inspired by GO_SDK net.http.ServeContent
|
||||
|
||||
//type RangeReadCloser struct {
|
||||
// GetReaderForRange RangeReaderFunc
|
||||
//}
|
||||
|
||||
// ServeHTTP replies to the request using the content in the
|
||||
// provided RangeReadCloser. The main benefit of ServeHTTP over io.Copy
|
||||
// provided range reader. The main benefit of ServeHTTP over io.Copy
|
||||
// is that it handles Range requests properly, sets the MIME type, and
|
||||
// handles If-Match, If-Unmodified-Since, If-None-Match, If-Modified-Since,
|
||||
// and If-Range requests.
|
||||
@@ -47,13 +43,11 @@ import (
|
||||
// request includes an If-Modified-Since header, ServeHTTP uses
|
||||
// modtime to decide whether the content needs to be sent at all.
|
||||
//
|
||||
// The content's RangeReadCloser method must work: ServeHTTP gives a range,
|
||||
// caller will give the reader for that Range.
|
||||
// The content's RangeRead method must return a reader for the requested range.
|
||||
//
|
||||
// If the caller has set w's ETag header formatted per RFC 7232, section 2.3,
|
||||
// ServeHTTP uses it to handle requests using If-Match, If-None-Match, or If-Range.
|
||||
func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time.Time, size int64, RangeReadCloser model.RangeReadCloserIF) error {
|
||||
defer RangeReadCloser.Close()
|
||||
func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time.Time, size int64, rangeReader model.RangeReaderIF) (err error) {
|
||||
setLastModified(w, modTime)
|
||||
done, rangeReq := checkPreconditions(w, r, modTime)
|
||||
if done {
|
||||
@@ -113,10 +107,11 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
|
||||
ctx := r.Context()
|
||||
switch {
|
||||
case len(ranges) == 0:
|
||||
reader, err := RangeReadCloser.RangeRead(ctx, http_range.Range{Length: -1})
|
||||
reader, err := openRange(ctx, rangeReader, http_range.Range{Length: -1})
|
||||
if err != nil {
|
||||
code = http.StatusRequestedRangeNotSatisfiable
|
||||
if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok {
|
||||
var statusCode HttpStatusCodeError
|
||||
if errors.As(err, &statusCode) {
|
||||
code = int(statusCode)
|
||||
}
|
||||
http.Error(w, err.Error(), code)
|
||||
@@ -136,10 +131,11 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
|
||||
// does not request multiple parts might not support
|
||||
// multipart responses."
|
||||
ra := ranges[0]
|
||||
sendContent, err = RangeReadCloser.RangeRead(ctx, ra)
|
||||
sendContent, err = openRange(ctx, rangeReader, ra)
|
||||
if err != nil {
|
||||
code = http.StatusRequestedRangeNotSatisfiable
|
||||
if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok {
|
||||
var statusCode HttpStatusCodeError
|
||||
if errors.As(err, &statusCode) {
|
||||
code = int(statusCode)
|
||||
}
|
||||
http.Error(w, err.Error(), code)
|
||||
@@ -159,7 +155,6 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
|
||||
mw := multipart.NewWriter(pw)
|
||||
w.Header().Set("Content-Type", "multipart/byteranges; boundary="+mw.Boundary())
|
||||
sendContent = pr
|
||||
defer pr.Close() // cause writing goroutine to fail and exit if CopyN doesn't finish.
|
||||
go func() {
|
||||
for _, ra := range ranges {
|
||||
part, err := mw.CreatePart(ra.MimeHeader(contentType, size))
|
||||
@@ -167,21 +162,18 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
|
||||
pw.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
reader, err := RangeReadCloser.RangeRead(ctx, ra)
|
||||
if err != nil {
|
||||
pw.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
if _, err := utils.CopyWithBufferN(part, reader, ra.Length); err != nil {
|
||||
if err := copyRange(ctx, part, rangeReader, ra); err != nil {
|
||||
pw.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
mw.Close()
|
||||
pw.Close()
|
||||
_ = pw.CloseWithError(mw.Close())
|
||||
}()
|
||||
}
|
||||
defer func() {
|
||||
err = closeWithError(err, sendContent)
|
||||
}()
|
||||
|
||||
w.Header().Set("Accept-Ranges", "bytes")
|
||||
if w.Header().Get("Content-Encoding") == "" {
|
||||
@@ -201,7 +193,8 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
|
||||
log.Warnf("Maybe size incorrect or reader not giving correct/full data, or connection closed before finish. written bytes: %d ,sendSize:%d, ", written, sendSize)
|
||||
}
|
||||
code = http.StatusInternalServerError
|
||||
if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok {
|
||||
var statusCode HttpStatusCodeError
|
||||
if errors.As(err, &statusCode) {
|
||||
code = int(statusCode)
|
||||
}
|
||||
w.WriteHeader(code)
|
||||
@@ -210,16 +203,86 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func copyRange(ctx context.Context, dst io.Writer, rangeReader model.RangeReaderIF, requested http_range.Range) (err error) {
|
||||
reader, err := openRange(ctx, rangeReader, requested)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() {
|
||||
err = closeWithError(err, reader)
|
||||
}()
|
||||
_, err = utils.CopyWithBufferN(dst, reader, requested.Length)
|
||||
return err
|
||||
}
|
||||
|
||||
func openRange(ctx context.Context, rangeReader model.RangeReaderIF, requested http_range.Range) (io.ReadCloser, error) {
|
||||
reader, err := rangeReader.RangeRead(ctx, requested)
|
||||
if err != nil {
|
||||
if reader != nil {
|
||||
err = closeWithError(err, reader)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
if reader == nil {
|
||||
return nil, errors.New("range reader returned a nil body")
|
||||
}
|
||||
return reader, nil
|
||||
}
|
||||
|
||||
func closeWithError(err error, closer io.Closer) error {
|
||||
closeErr := closer.Close()
|
||||
if err == nil {
|
||||
return closeErr
|
||||
}
|
||||
if closeErr == nil {
|
||||
return err
|
||||
}
|
||||
return stderrors.Join(err, closeErr)
|
||||
}
|
||||
|
||||
// unsafeProxyHeaders are never forwarded from the client request to the
|
||||
// upstream storage, regardless of the proxy_ignore_headers setting. They either
|
||||
// carry the caller's credentials, describe the hop to this server rather than
|
||||
// the hop to upstream, or let the caller influence how upstream routes and
|
||||
// authenticates the request.
|
||||
var unsafeProxyHeaders = map[string]struct{}{
|
||||
"authorization": {},
|
||||
"cookie": {},
|
||||
"proxy-authorization": {},
|
||||
"www-authenticate": {},
|
||||
"host": {},
|
||||
"referer": {},
|
||||
"origin": {},
|
||||
"connection": {},
|
||||
"keep-alive": {},
|
||||
"proxy-connection": {},
|
||||
"te": {},
|
||||
"trailer": {},
|
||||
"transfer-encoding": {},
|
||||
"upgrade": {},
|
||||
"forwarded": {},
|
||||
"x-forwarded-for": {},
|
||||
"x-forwarded-host": {},
|
||||
"x-forwarded-proto": {},
|
||||
"x-real-ip": {},
|
||||
}
|
||||
|
||||
|
||||
func ProcessHeader(origin, override http.Header) http.Header {
|
||||
result := http.Header{}
|
||||
// client header
|
||||
for h, val := range origin {
|
||||
if utils.SliceContains(conf.SlicesMap[conf.ProxyIgnoreHeaders], strings.ToLower(h)) {
|
||||
lower := strings.ToLower(h)
|
||||
if _, unsafe := unsafeProxyHeaders[lower]; unsafe {
|
||||
continue
|
||||
}
|
||||
if utils.SliceContains(conf.SlicesMap[conf.ProxyIgnoreHeaders], lower) {
|
||||
continue
|
||||
}
|
||||
result[h] = val
|
||||
}
|
||||
// needed header
|
||||
// needed header, produced by the storage driver rather than the client
|
||||
for h, val := range override {
|
||||
result[h] = val
|
||||
}
|
||||
|
||||
@@ -0,0 +1,67 @@
|
||||
package net
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
)
|
||||
|
||||
// The client must not be able to smuggle credential or routing headers into the
|
||||
// request that this server makes to the upstream storage, even when the
|
||||
// proxy_ignore_headers setting has been emptied.
|
||||
func TestProcessHeaderDropsUnsafeClientHeaders(t *testing.T) {
|
||||
conf.SlicesMap[conf.ProxyIgnoreHeaders] = nil
|
||||
|
||||
origin := http.Header{}
|
||||
origin.Set("Authorization", "Bearer victim-token")
|
||||
origin.Set("Cookie", "session=victim")
|
||||
origin.Set("X-Forwarded-For", "127.0.0.1")
|
||||
origin.Set("Host", "internal.example")
|
||||
origin.Set("Range", "bytes=0-1023")
|
||||
|
||||
result := ProcessHeader(origin, nil)
|
||||
|
||||
for _, h := range []string{"Authorization", "Cookie", "X-Forwarded-For", "Host"} {
|
||||
if got := result.Get(h); got != "" {
|
||||
t.Errorf("header %q must not be forwarded upstream, got %q", h, got)
|
||||
}
|
||||
}
|
||||
if got := result.Get("Range"); got != "bytes=0-1023" {
|
||||
t.Errorf("Range must be preserved, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
// Headers supplied by the storage driver still win, since they carry the
|
||||
// credentials needed to reach upstream.
|
||||
func TestProcessHeaderOverrideWins(t *testing.T) {
|
||||
conf.SlicesMap[conf.ProxyIgnoreHeaders] = nil
|
||||
|
||||
origin := http.Header{}
|
||||
origin.Set("Authorization", "Bearer victim-token")
|
||||
|
||||
override := http.Header{}
|
||||
override.Set("Authorization", "Bearer driver-token")
|
||||
|
||||
result := ProcessHeader(origin, override)
|
||||
if got := result.Get("Authorization"); got != "Bearer driver-token" {
|
||||
t.Errorf("driver header must be used, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessHeaderStillHonoursIgnoreSetting(t *testing.T) {
|
||||
conf.SlicesMap[conf.ProxyIgnoreHeaders] = []string{"x-custom"}
|
||||
t.Cleanup(func() { conf.SlicesMap[conf.ProxyIgnoreHeaders] = nil })
|
||||
|
||||
origin := http.Header{}
|
||||
origin.Set("X-Custom", "drop-me")
|
||||
origin.Set("X-Keep", "keep-me")
|
||||
|
||||
result := ProcessHeader(origin, nil)
|
||||
if got := result.Get("X-Custom"); got != "" {
|
||||
t.Errorf("configured ignore header must be dropped, got %q", got)
|
||||
}
|
||||
if got := result.Get("X-Keep"); got != "keep-me" {
|
||||
t.Errorf("unrelated header must be preserved, got %q", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,236 @@
|
||||
package net
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"reflect"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
)
|
||||
|
||||
func TestServeHTTPClosesMultipartRangeBeforeOpeningNext(t *testing.T) {
|
||||
source := newSequentialRangeSource("abc")
|
||||
ctx, cancel := context.WithTimeout(t.Context(), 500*time.Millisecond)
|
||||
defer cancel()
|
||||
request := httptest.NewRequest(http.MethodGet, "/file", nil).WithContext(ctx)
|
||||
request.Header.Set("Range", "bytes=0-0,2-2")
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
if err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source); err != nil {
|
||||
t.Fatalf("ServeHTTP() error = %v", err)
|
||||
}
|
||||
response := recorder.Result()
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode != http.StatusPartialContent {
|
||||
t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusPartialContent)
|
||||
}
|
||||
|
||||
mediaType, params, err := mime.ParseMediaType(response.Header.Get("Content-Type"))
|
||||
if err != nil {
|
||||
t.Fatalf("parse Content-Type: %v", err)
|
||||
}
|
||||
if mediaType != "multipart/byteranges" {
|
||||
t.Fatalf("Content-Type = %q, want multipart/byteranges", mediaType)
|
||||
}
|
||||
multipartReader := multipart.NewReader(response.Body, params["boundary"])
|
||||
var parts []string
|
||||
for {
|
||||
part, err := multipartReader.NextPart()
|
||||
if errors.Is(err, io.EOF) {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("read multipart part: %v", err)
|
||||
}
|
||||
body, err := io.ReadAll(part)
|
||||
if err != nil {
|
||||
t.Fatalf("read multipart body: %v", err)
|
||||
}
|
||||
parts = append(parts, string(body))
|
||||
}
|
||||
if want := []string{"a", "c"}; !reflect.DeepEqual(parts, want) {
|
||||
t.Fatalf("multipart parts = %q, want %q", parts, want)
|
||||
}
|
||||
assertRangeLifecycle(t, source, []string{"open:0", "close:0", "open:2", "close:2"}, []int{1, 1})
|
||||
}
|
||||
|
||||
func TestServeHTTPClosesSelectedRangeBody(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
method string
|
||||
rangeValue string
|
||||
wantStatus int
|
||||
wantEvents []string
|
||||
}{
|
||||
{name: "full", method: http.MethodGet, wantStatus: http.StatusOK, wantEvents: []string{"open:0", "close:0"}},
|
||||
{name: "single range", method: http.MethodGet, rangeValue: "bytes=1-1", wantStatus: http.StatusPartialContent, wantEvents: []string{"open:1", "close:1"}},
|
||||
{name: "head", method: http.MethodHead, wantStatus: http.StatusOK, wantEvents: []string{"open:0", "close:0"}},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
source := newSequentialRangeSource("abc")
|
||||
request := httptest.NewRequest(test.method, "/file", nil)
|
||||
if test.rangeValue != "" {
|
||||
request.Header.Set("Range", test.rangeValue)
|
||||
}
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
if err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source); err != nil {
|
||||
t.Fatalf("ServeHTTP() error = %v", err)
|
||||
}
|
||||
if recorder.Code != test.wantStatus {
|
||||
t.Fatalf("status = %d, want %d", recorder.Code, test.wantStatus)
|
||||
}
|
||||
assertRangeLifecycle(t, source, test.wantEvents, []int{1})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestServeHTTPClosesRangeAfterWriteFailure(t *testing.T) {
|
||||
writeErr := errors.New("write failed")
|
||||
source := newSequentialRangeSource("abc")
|
||||
request := httptest.NewRequest(http.MethodGet, "/file", nil)
|
||||
writer := &failingResponseWriter{header: make(http.Header), err: writeErr}
|
||||
|
||||
err := ServeHTTP(writer, request, "file.txt", time.Time{}, 3, source)
|
||||
if !errors.Is(err, writeErr) {
|
||||
t.Fatalf("ServeHTTP() error = %v, want %v", err, writeErr)
|
||||
}
|
||||
assertRangeLifecycle(t, source, []string{"open:0", "close:0"}, []int{1})
|
||||
}
|
||||
|
||||
func TestServeHTTPClosesBodyReturnedWithOpenError(t *testing.T) {
|
||||
source := newSequentialRangeSource("abc")
|
||||
source.openErr = HttpStatusCodeError(http.StatusServiceUnavailable)
|
||||
source.closeErr = errors.New("close failed")
|
||||
request := httptest.NewRequest(http.MethodGet, "/file", nil)
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
if err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source); err != nil {
|
||||
t.Fatalf("ServeHTTP() error = %v", err)
|
||||
}
|
||||
if recorder.Code != http.StatusServiceUnavailable {
|
||||
t.Fatalf("status = %d, want %d", recorder.Code, http.StatusServiceUnavailable)
|
||||
}
|
||||
assertRangeLifecycle(t, source, []string{"open:0", "close:0"}, []int{1})
|
||||
}
|
||||
|
||||
func TestServeHTTPStopsMultipartAfterRangeCloseFailure(t *testing.T) {
|
||||
closeErr := errors.New("close failed")
|
||||
source := newSequentialRangeSource("abc")
|
||||
source.closeErr = closeErr
|
||||
request := httptest.NewRequest(http.MethodGet, "/file", nil)
|
||||
request.Header.Set("Range", "bytes=0-0,2-2")
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source)
|
||||
if !errors.Is(err, closeErr) {
|
||||
t.Fatalf("ServeHTTP() error = %v, want %v", err, closeErr)
|
||||
}
|
||||
assertRangeLifecycle(t, source, []string{"open:0", "close:0"}, []int{1})
|
||||
}
|
||||
|
||||
type sequentialRangeSource struct {
|
||||
content []byte
|
||||
permit chan struct{}
|
||||
|
||||
mu sync.Mutex
|
||||
events []string
|
||||
closeCounts []int
|
||||
closeErr error
|
||||
openErr error
|
||||
}
|
||||
|
||||
func newSequentialRangeSource(content string) *sequentialRangeSource {
|
||||
return &sequentialRangeSource{
|
||||
content: []byte(content),
|
||||
permit: make(chan struct{}, 1),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *sequentialRangeSource) RangeRead(ctx context.Context, requested http_range.Range) (io.ReadCloser, error) {
|
||||
select {
|
||||
case s.permit <- struct{}{}:
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
|
||||
start := int(requested.Start)
|
||||
length := int(requested.Length)
|
||||
if length < 0 || start+length > len(s.content) {
|
||||
length = len(s.content) - start
|
||||
}
|
||||
end := start + length
|
||||
s.mu.Lock()
|
||||
index := len(s.closeCounts)
|
||||
s.events = append(s.events, fmt.Sprintf("open:%d", requested.Start))
|
||||
s.closeCounts = append(s.closeCounts, 0)
|
||||
s.mu.Unlock()
|
||||
return &testReadCloser{
|
||||
Reader: bytes.NewReader(s.content[start:end]),
|
||||
close: func() error {
|
||||
s.mu.Lock()
|
||||
s.closeCounts[index]++
|
||||
closeCalls := s.closeCounts[index]
|
||||
if closeCalls == 1 {
|
||||
s.events = append(s.events, fmt.Sprintf("close:%d", requested.Start))
|
||||
}
|
||||
s.mu.Unlock()
|
||||
if closeCalls != 1 {
|
||||
return fmt.Errorf("body closed %d times", closeCalls)
|
||||
}
|
||||
<-s.permit
|
||||
return s.closeErr
|
||||
},
|
||||
}, s.openErr
|
||||
}
|
||||
|
||||
func (s *sequentialRangeSource) eventsSnapshot() []string {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return append([]string(nil), s.events...)
|
||||
}
|
||||
|
||||
func (s *sequentialRangeSource) closeCountsSnapshot() []int {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return append([]int(nil), s.closeCounts...)
|
||||
}
|
||||
|
||||
func assertRangeLifecycle(t *testing.T, source *sequentialRangeSource, wantEvents []string, wantCloseCounts []int) {
|
||||
t.Helper()
|
||||
if got := source.eventsSnapshot(); !reflect.DeepEqual(got, wantEvents) {
|
||||
t.Fatalf("range lifecycle = %v, want %v", got, wantEvents)
|
||||
}
|
||||
if got := source.closeCountsSnapshot(); !reflect.DeepEqual(got, wantCloseCounts) {
|
||||
t.Fatalf("close counts = %v, want %v", got, wantCloseCounts)
|
||||
}
|
||||
}
|
||||
|
||||
type failingResponseWriter struct {
|
||||
header http.Header
|
||||
err error
|
||||
}
|
||||
|
||||
func (w *failingResponseWriter) Header() http.Header { return w.header }
|
||||
func (*failingResponseWriter) WriteHeader(int) {}
|
||||
func (w *failingResponseWriter) Write([]byte) (int, error) {
|
||||
return 0, w.err
|
||||
}
|
||||
|
||||
type testReadCloser struct {
|
||||
io.Reader
|
||||
close func() error
|
||||
}
|
||||
|
||||
func (b *testReadCloser) Close() error { return b.close() }
|
||||
@@ -25,7 +25,6 @@ import (
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/op"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/setting"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/task"
|
||||
"github.com/OpenListTeam/OpenList/v4/server/common"
|
||||
"github.com/google/uuid"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
@@ -184,7 +183,7 @@ func AddURL(ctx context.Context, args *AddURLArgs) (task.TaskExtensionInfo, erro
|
||||
t := &DownloadTask{
|
||||
TaskExtension: task.TaskExtension{
|
||||
Creator: taskCreator,
|
||||
ApiUrl: common.GetApiUrl(ctx),
|
||||
ApiUrl: conf.GetApiUrl(ctx),
|
||||
},
|
||||
Url: args.URL,
|
||||
DstDirPath: args.DstDirPath,
|
||||
@@ -223,15 +222,17 @@ func isEd2kURL(urlStr string) bool {
|
||||
}
|
||||
|
||||
func ed2kToolForStorage(storage driver.Driver) string {
|
||||
switch toolNameForStorage(storage) {
|
||||
name := NativeToolName(storage)
|
||||
switch name {
|
||||
case "115 Cloud", "115 Open":
|
||||
return toolNameForStorage(storage)
|
||||
return name
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func toolNameForStorage(storage driver.Driver) string {
|
||||
// NativeToolName returns the offline-download tool implemented by storage.
|
||||
func NativeToolName(storage driver.Driver) string {
|
||||
switch storage.(type) {
|
||||
case *_115.Pan115:
|
||||
return "115 Cloud"
|
||||
|
||||
@@ -58,7 +58,7 @@ func TestEd2kToolForStorage(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolNameForStorage(t *testing.T) {
|
||||
func TestNativeToolName(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
storage driver.Driver
|
||||
@@ -78,8 +78,8 @@ func TestToolNameForStorage(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := toolNameForStorage(tt.storage); got != tt.want {
|
||||
t.Fatalf("toolNameForStorage(%T) = %q, want %q", tt.storage, got, tt.want)
|
||||
if got := NativeToolName(tt.storage); got != tt.want {
|
||||
t.Fatalf("NativeToolName(%T) = %q, want %q", tt.storage, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -45,7 +45,7 @@ func (t ToolsManager) NamesForPath(path string) []string {
|
||||
return names
|
||||
}
|
||||
|
||||
name := toolNameForStorage(storage)
|
||||
name := NativeToolName(storage)
|
||||
if name == "" {
|
||||
return names
|
||||
}
|
||||
|
||||
@@ -20,7 +20,6 @@ import (
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/torrent"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"github.com/OpenListTeam/OpenList/v4/server/common"
|
||||
"github.com/OpenListTeam/tache"
|
||||
"github.com/pkg/errors"
|
||||
log "github.com/sirupsen/logrus"
|
||||
@@ -140,7 +139,7 @@ func transferStd(ctx context.Context, tempDir, dstDirPath string, deletePolicy D
|
||||
TaskData: fs.TaskData{
|
||||
TaskExtension: task.TaskExtension{
|
||||
Creator: taskCreator,
|
||||
ApiUrl: common.GetApiUrl(ctx),
|
||||
ApiUrl: conf.GetApiUrl(ctx),
|
||||
},
|
||||
SrcActualPath: stdpath.Join(tempDir, entry.Name()),
|
||||
DstActualPath: dstDirActualPath,
|
||||
@@ -276,7 +275,7 @@ func transferObj(ctx context.Context, tempDir, dstDirPath string, deletePolicy D
|
||||
TaskData: fs.TaskData{
|
||||
TaskExtension: task.TaskExtension{
|
||||
Creator: taskCreator,
|
||||
ApiUrl: common.GetApiUrl(ctx),
|
||||
ApiUrl: conf.GetApiUrl(ctx),
|
||||
},
|
||||
SrcActualPath: stdpath.Join(srcObjActualPath, obj.GetName()),
|
||||
DstActualPath: dstDirActualPath,
|
||||
|
||||
+11
-7
@@ -390,8 +390,9 @@ func ArchiveGet(ctx context.Context, storage driver.Driver, path string, args mo
|
||||
}
|
||||
|
||||
type objWithLink struct {
|
||||
link *model.Link
|
||||
obj model.Obj
|
||||
link *model.Link
|
||||
obj model.Obj
|
||||
policy linkCachePolicy
|
||||
}
|
||||
|
||||
var (
|
||||
@@ -405,7 +406,7 @@ func DriverExtract(ctx context.Context, storage driver.Driver, path string, args
|
||||
}
|
||||
key := stdpath.Join(Key(storage, path), args.InnerPath)
|
||||
if ol, ok := extractCache.Get(key); ok {
|
||||
if ol.link.Expiration != nil || ol.link.SyncClosers.AcquireReference() || !ol.link.RequireReference {
|
||||
if ol.acquire() {
|
||||
return ol.link, ol.obj, nil
|
||||
}
|
||||
}
|
||||
@@ -415,8 +416,8 @@ func DriverExtract(ctx context.Context, storage driver.Driver, path string, args
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed extract archive")
|
||||
}
|
||||
if ol.link.Expiration != nil {
|
||||
extractCache.SetWithTTL(key, ol, *ol.link.Expiration)
|
||||
if ol.policy.expiration != nil {
|
||||
extractCache.SetWithTTL(key, ol, *ol.policy.expiration)
|
||||
} else {
|
||||
extractCache.SetWithExpirable(key, ol, &ol.link.SyncClosers)
|
||||
}
|
||||
@@ -428,7 +429,7 @@ func DriverExtract(ctx context.Context, storage driver.Driver, path string, args
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if ol.link.SyncClosers.AcquireReference() || !ol.link.RequireReference {
|
||||
if ol.acquire() {
|
||||
return ol.link, ol.obj, nil
|
||||
}
|
||||
}
|
||||
@@ -450,7 +451,10 @@ func driverExtract(ctx context.Context, storage driver.Driver, path string, args
|
||||
return nil, errors.WithStack(errs.NotFile)
|
||||
}
|
||||
link, err := storageAr.Extract(ctx, archiveFile, args)
|
||||
return &objWithLink{link: link, obj: extracted}, err
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return admitLink(link, extracted)
|
||||
}
|
||||
|
||||
type streamWithParent struct {
|
||||
|
||||
+12
-7
@@ -233,7 +233,10 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li
|
||||
if mode == -1 {
|
||||
mode = storage.(driver.LinkCacheModeResolver).ResolveLinkCacheMode(path)
|
||||
}
|
||||
typeKey := args.Type
|
||||
typeKey := "proxy/" + args.Type
|
||||
if args.Redirect {
|
||||
typeKey = "redirect/" + args.Type
|
||||
}
|
||||
if mode&driver.LinkCacheIP != 0 {
|
||||
typeKey += "/" + args.IP
|
||||
}
|
||||
@@ -242,8 +245,7 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li
|
||||
}
|
||||
key := Key(storage, path)
|
||||
if ol, exists := Cache.linkCache.GetType(key, typeKey); exists {
|
||||
if ol.link.Expiration != nil ||
|
||||
ol.link.SyncClosers.AcquireReference() || !ol.link.RequireReference {
|
||||
if ol.acquire() {
|
||||
return ol.link, ol.obj, nil
|
||||
}
|
||||
}
|
||||
@@ -261,9 +263,12 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed get link")
|
||||
}
|
||||
ol := &objWithLink{link: link, obj: file}
|
||||
if link.Expiration != nil {
|
||||
Cache.linkCache.SetTypeWithTTL(key, typeKey, ol, *link.Expiration)
|
||||
ol, err := admitLink(link, file)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if ol.policy.expiration != nil {
|
||||
Cache.linkCache.SetTypeWithTTL(key, typeKey, ol, *ol.policy.expiration)
|
||||
} else {
|
||||
Cache.linkCache.SetTypeWithExpirable(key, typeKey, ol, &link.SyncClosers)
|
||||
}
|
||||
@@ -274,7 +279,7 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if ol.link.SyncClosers.AcquireReference() || !ol.link.RequireReference {
|
||||
if ol.acquire() {
|
||||
return ol.link, ol.obj, nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
package op
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/driver"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
)
|
||||
|
||||
type linkModeDriver struct {
|
||||
driver.Driver
|
||||
storage model.Storage
|
||||
calls int
|
||||
}
|
||||
|
||||
func (d *linkModeDriver) Config() driver.Config { return driver.Config{} }
|
||||
|
||||
func (d *linkModeDriver) GetStorage() *model.Storage { return &d.storage }
|
||||
|
||||
func (d *linkModeDriver) Get(context.Context, string) (model.Obj, error) {
|
||||
return &model.Object{Name: "file"}, nil
|
||||
}
|
||||
|
||||
func (d *linkModeDriver) Link(_ context.Context, _ model.Obj, args model.LinkArgs) (*model.Link, error) {
|
||||
d.calls++
|
||||
expiration := time.Minute
|
||||
if args.Redirect {
|
||||
return &model.Link{URL: "https://example.com/file", Expiration: &expiration}, nil
|
||||
}
|
||||
return &model.Link{
|
||||
RangeReader: stream.RangeReaderFunc(func(context.Context, http_range.Range) (io.ReadCloser, error) {
|
||||
return io.NopCloser(strings.NewReader("file")), nil
|
||||
}),
|
||||
Expiration: &expiration,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func TestLinkCacheSeparatesRedirectAndProxy(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
firstRedirect bool
|
||||
}{
|
||||
{name: "redirect then proxy", firstRedirect: true},
|
||||
{name: "proxy then redirect", firstRedirect: false},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
d := &linkModeDriver{storage: model.Storage{MountPath: "/" + t.Name()}}
|
||||
for _, redirect := range []bool{tc.firstRedirect, !tc.firstRedirect, tc.firstRedirect, !tc.firstRedirect} {
|
||||
link, _, err := Link(context.Background(), d, "/file", model.LinkArgs{Redirect: redirect})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if redirect && (link.URL == "" || link.RangeReader != nil) {
|
||||
t.Fatalf("redirect link has wrong shape: %+v", link)
|
||||
}
|
||||
if !redirect && (link.URL != "" || link.RangeReader == nil) {
|
||||
t.Fatalf("proxy link has wrong shape: %+v", link)
|
||||
}
|
||||
}
|
||||
if d.calls != 2 {
|
||||
t.Fatalf("expected one driver call per mode, got %d", d.calls)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package op
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
)
|
||||
|
||||
var errConflictingLinkLifecycle = errors.New("invalid link lifecycle: expiration cannot be combined with owned closers or RequireReference")
|
||||
|
||||
type linkCachePolicy struct {
|
||||
expiration *time.Duration
|
||||
requireReference bool
|
||||
}
|
||||
|
||||
func admitLink(link *model.Link, obj model.Obj) (*objWithLink, error) {
|
||||
if link.Expiration != nil && (link.RequireReference || link.SyncClosers.Length() > 0) {
|
||||
return nil, errors.Join(errConflictingLinkLifecycle, link.Close())
|
||||
}
|
||||
return &objWithLink{
|
||||
link: link,
|
||||
obj: obj,
|
||||
policy: linkCachePolicy{
|
||||
expiration: link.Expiration,
|
||||
requireReference: link.RequireReference,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (ol *objWithLink) acquire() bool {
|
||||
return ol.policy.expiration != nil ||
|
||||
ol.link.SyncClosers.AcquireReference() || !ol.policy.requireReference
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
package op
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/driver"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/singleflight"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
)
|
||||
|
||||
type linkLifecycleDriver struct {
|
||||
model.Storage
|
||||
links func() *model.Link
|
||||
calls atomic.Int32
|
||||
}
|
||||
|
||||
func (d *linkLifecycleDriver) Config() driver.Config { return driver.Config{} }
|
||||
func (d *linkLifecycleDriver) GetAddition() driver.Additional { return nil }
|
||||
func (d *linkLifecycleDriver) Init(context.Context) error { return nil }
|
||||
func (d *linkLifecycleDriver) Drop(context.Context) error { return nil }
|
||||
func (d *linkLifecycleDriver) List(context.Context, model.Obj, model.ListArgs) ([]model.Obj, error) {
|
||||
return nil, nil
|
||||
}
|
||||
func (d *linkLifecycleDriver) Get(context.Context, string) (model.Obj, error) {
|
||||
return &model.Object{Name: "file", Path: "/file"}, nil
|
||||
}
|
||||
func (d *linkLifecycleDriver) Link(context.Context, model.Obj, model.LinkArgs) (*model.Link, error) {
|
||||
d.calls.Add(1)
|
||||
return d.links(), nil
|
||||
}
|
||||
|
||||
func resetLinkLifecycleState(t *testing.T) {
|
||||
t.Helper()
|
||||
oldCache := Cache
|
||||
Cache, linkG = NewCacheManager(), singleflight.Group[*objWithLink]{}
|
||||
t.Cleanup(func() { Cache, linkG = oldCache, singleflight.Group[*objWithLink]{} })
|
||||
}
|
||||
|
||||
func acquireTestLink(t *testing.T, d *linkLifecycleDriver) *model.Link {
|
||||
t.Helper()
|
||||
link, _, err := Link(context.Background(), d, "/file", model.LinkArgs{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return link
|
||||
}
|
||||
|
||||
func TestLinkLifecycleModes(t *testing.T) {
|
||||
t.Run("TTL descriptor remains reusable after close", func(t *testing.T) {
|
||||
resetLinkLifecycleState(t)
|
||||
ttl := time.Minute
|
||||
d := &linkLifecycleDriver{
|
||||
Storage: model.Storage{MountPath: "/ttl"},
|
||||
links: func() *model.Link { return &model.Link{URL: "https://example.test/file", Expiration: &ttl} },
|
||||
}
|
||||
|
||||
first := acquireTestLink(t, d)
|
||||
_ = first.Close()
|
||||
second := acquireTestLink(t, d)
|
||||
if second.URL != first.URL || d.calls.Load() != 1 {
|
||||
t.Fatalf("TTL link was not reused: calls=%d", d.calls.Load())
|
||||
}
|
||||
_ = second.Close()
|
||||
})
|
||||
|
||||
t.Run("references keep shared resources alive until final close", func(t *testing.T) {
|
||||
resetLinkLifecycleState(t)
|
||||
var closes atomic.Int32
|
||||
d := &linkLifecycleDriver{
|
||||
Storage: model.Storage{MountPath: "/reference"},
|
||||
links: func() *model.Link {
|
||||
return &model.Link{
|
||||
URL: "https://example.test/file",
|
||||
SyncClosers: utils.NewSyncClosers(utils.CloseFunc(func() error { closes.Add(1); return nil })),
|
||||
RequireReference: true,
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
first := acquireTestLink(t, d)
|
||||
second := acquireTestLink(t, d)
|
||||
_ = first.Close()
|
||||
if closes.Load() != 0 {
|
||||
t.Fatal("shared resource closed while another reference was active")
|
||||
}
|
||||
_ = second.Close()
|
||||
if closes.Load() != 1 {
|
||||
t.Fatalf("final close count = %d, want 1", closes.Load())
|
||||
}
|
||||
third := acquireTestLink(t, d)
|
||||
_ = third.Close()
|
||||
if d.calls.Load() != 2 || closes.Load() != 2 {
|
||||
t.Fatalf("stale link was not replaced: calls=%d closes=%d", d.calls.Load(), closes.Load())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("close-invalidated link is reacquired", func(t *testing.T) {
|
||||
resetLinkLifecycleState(t)
|
||||
d := &linkLifecycleDriver{
|
||||
Storage: model.Storage{MountPath: "/close-invalidated"},
|
||||
links: func() *model.Link {
|
||||
return &model.Link{SyncClosers: utils.NewSyncClosers(utils.CloseFunc(func() error { return nil }))}
|
||||
},
|
||||
}
|
||||
|
||||
first := acquireTestLink(t, d)
|
||||
_ = first.Close()
|
||||
second := acquireTestLink(t, d)
|
||||
_ = second.Close()
|
||||
if d.calls.Load() != 2 {
|
||||
t.Fatalf("driver calls = %d, want 2", d.calls.Load())
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("TTL with owned resources is rejected and released", func(t *testing.T) {
|
||||
resetLinkLifecycleState(t)
|
||||
ttl := time.Minute
|
||||
var closes atomic.Int32
|
||||
d := &linkLifecycleDriver{
|
||||
Storage: model.Storage{MountPath: "/conflict"},
|
||||
links: func() *model.Link {
|
||||
return &model.Link{
|
||||
Expiration: &ttl,
|
||||
SyncClosers: utils.NewSyncClosers(utils.CloseFunc(func() error { closes.Add(1); return nil })),
|
||||
RequireReference: true,
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
_, _, err := Link(context.Background(), d, "/file", model.LinkArgs{})
|
||||
if err == nil || !strings.Contains(err.Error(), "expiration cannot be combined") {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if closes.Load() != 1 {
|
||||
t.Fatalf("rejected link close count = %d, want 1", closes.Load())
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -4,11 +4,11 @@ import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/errs"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/op"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"github.com/OpenListTeam/OpenList/v4/server/common"
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
@@ -38,7 +38,7 @@ func link(ctx context.Context, sid, path string, args *LinkArgs) (*model.Sharing
|
||||
return nil, nil, nil, errors.WithMessage(err, "failed get sharing link")
|
||||
}
|
||||
if l.URL != "" && !strings.HasPrefix(l.URL, "http://") && !strings.HasPrefix(l.URL, "https://") {
|
||||
l.URL = common.GetApiUrl(ctx) + l.URL
|
||||
l.URL = conf.GetApiUrl(ctx) + l.URL
|
||||
}
|
||||
return sharing, l, obj, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
package stream_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"math/rand"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
)
|
||||
|
||||
// maxReuseGap mirrors the internal continuation-reuse window (4*utils.MB).
|
||||
const maxReuseGap = 4 * 1024 * 1024
|
||||
|
||||
// newMockSeekableStream builds a SeekableStream whose range reads are served
|
||||
// from data, counting every upstream range request in gets.
|
||||
func newMockSeekableStream(t *testing.T, data []byte, gets *atomic.Int64) *stream.SeekableStream {
|
||||
t.Helper()
|
||||
rr := stream.RangeReaderFunc(func(ctx context.Context, r http_range.Range) (io.ReadCloser, error) {
|
||||
gets.Add(1)
|
||||
if r.Length < 0 || r.Start+r.Length > int64(len(data)) {
|
||||
r.Length = int64(len(data)) - r.Start
|
||||
}
|
||||
return io.NopCloser(io.NewSectionReader(bytes.NewReader(data), r.Start, r.Length)), nil
|
||||
})
|
||||
ss, err := stream.NewSeekableStream(&stream.FileStream{
|
||||
Obj: &model.Object{Size: int64(len(data))},
|
||||
Ctx: context.Background(),
|
||||
}, &model.Link{
|
||||
RangeReader: rr,
|
||||
ContentLength: int64(len(data)),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSeekableStream() error = %v", err)
|
||||
}
|
||||
return ss
|
||||
}
|
||||
|
||||
// readAtFull reads len(p) bytes at off and fails the test on mismatch.
|
||||
func readAtFull(t *testing.T, ra io.ReaderAt, data []byte, off int64, p []byte) {
|
||||
t.Helper()
|
||||
n, err := ra.ReadAt(p, off)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadAt(off=%d) error = %v", off, err)
|
||||
}
|
||||
if !bytes.Equal(p, data[off:off+int64(n)]) {
|
||||
t.Fatalf("ReadAt(off=%d) content mismatch", off)
|
||||
}
|
||||
}
|
||||
|
||||
func randomData(size int) []byte {
|
||||
data := make([]byte, size)
|
||||
x := uint64(42)
|
||||
for i := range data {
|
||||
x = x*6364136223846793005 + 1
|
||||
data[i] = byte(x >> 33)
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
// Sequential reads must reuse a single upstream range request.
|
||||
func TestReadAtSeekerSequentialReuse(t *testing.T) {
|
||||
data := randomData(16 * 1024 * 1024)
|
||||
var gets atomic.Int64
|
||||
ss := newMockSeekableStream(t, data, &gets)
|
||||
ra, err := stream.NewReadAtSeeker(ss, 0, true)
|
||||
if err != nil {
|
||||
t.Fatalf("NewReadAtSeeker() error = %v", err)
|
||||
}
|
||||
buf := make([]byte, 128*1024)
|
||||
for off := 0; off < len(data); off += len(buf) {
|
||||
readAtFull(t, ra, data, int64(off), buf)
|
||||
}
|
||||
if n := gets.Load(); n != 1 {
|
||||
t.Fatalf("sequential read issued %d range requests, want 1", n)
|
||||
}
|
||||
}
|
||||
|
||||
// A read landing up to maxReuseGap bytes past a parked reader must be served
|
||||
// by advancing that reader, without a new range request.
|
||||
func TestReadAtSeekerSkipsAheadWithinWindow(t *testing.T) {
|
||||
data := randomData(16 * 1024 * 1024)
|
||||
var gets atomic.Int64
|
||||
ss := newMockSeekableStream(t, data, &gets)
|
||||
ra, err := stream.NewReadAtSeeker(ss, 0, true)
|
||||
if err != nil {
|
||||
t.Fatalf("NewReadAtSeeker() error = %v", err)
|
||||
}
|
||||
// Park a continuation reader right after reading the first 2 MiB.
|
||||
chunk := make([]byte, 256*1024)
|
||||
for off := 0; off < 2*1024*1024; off += len(chunk) {
|
||||
readAtFull(t, ra, data, int64(off), chunk)
|
||||
}
|
||||
skip := 512 * 1024
|
||||
off := int64(2*1024*1024 + skip)
|
||||
readAtFull(t, ra, data, off, chunk)
|
||||
if n := gets.Load(); n != 1 {
|
||||
t.Fatalf("window skip issued %d range requests, want 1", n)
|
||||
}
|
||||
// A second skip deeper inside the window must also be free.
|
||||
off = int64(4*1024*1024) - 128*1024
|
||||
readAtFull(t, ra, data, off, chunk)
|
||||
if n := gets.Load(); n != 1 {
|
||||
t.Fatalf("second window skip issued %d range requests, want 1", n)
|
||||
}
|
||||
}
|
||||
|
||||
// A forward jump beyond the reuse window must open a new range request but
|
||||
// keep the parked reader available for later window hits.
|
||||
func TestReadAtSeekerFarJumpOpensNewRequest(t *testing.T) {
|
||||
data := randomData(16 * 1024 * 1024)
|
||||
var gets atomic.Int64
|
||||
ss := newMockSeekableStream(t, data, &gets)
|
||||
ra, err := stream.NewReadAtSeeker(ss, 0, true)
|
||||
if err != nil {
|
||||
t.Fatalf("NewReadAtSeeker() error = %v", err)
|
||||
}
|
||||
buf := make([]byte, 256*1024)
|
||||
for off := 0; off < 2*1024*1024; off += len(buf) {
|
||||
readAtFull(t, ra, data, int64(off), buf)
|
||||
}
|
||||
// 2 MiB -> 10 MiB is beyond the 4 MiB reuse window.
|
||||
off := int64(10 * 1024 * 1024)
|
||||
readAtFull(t, ra, data, off, buf)
|
||||
if n := gets.Load(); n != 2 {
|
||||
t.Fatalf("far jump issued %d range requests, want 2", n)
|
||||
}
|
||||
// Back within the window of the 10 MiB chain: free reuse again.
|
||||
readAtFull(t, ra, data, off+maxReuseGap, buf)
|
||||
if n := gets.Load(); n != 2 {
|
||||
t.Fatalf("jump inside new window issued %d range requests, want 2", n)
|
||||
}
|
||||
}
|
||||
|
||||
// Backward reads can never reuse a parked continuation and must open a new
|
||||
// range request.
|
||||
func TestReadAtSeekerBackwardJumpOpensNewRequest(t *testing.T) {
|
||||
data := randomData(8 * 1024 * 1024)
|
||||
var gets atomic.Int64
|
||||
ss := newMockSeekableStream(t, data, &gets)
|
||||
ra, err := stream.NewReadAtSeeker(ss, 0, true)
|
||||
if err != nil {
|
||||
t.Fatalf("NewReadAtSeeker() error = %v", err)
|
||||
}
|
||||
buf := make([]byte, 256*1024)
|
||||
for off := 0; off < 2*1024*1024; off += len(buf) {
|
||||
readAtFull(t, ra, data, int64(off), buf)
|
||||
}
|
||||
readAtFull(t, ra, data, int64(1024*1024), buf)
|
||||
if n := gets.Load(); n != 2 {
|
||||
t.Fatalf("backward jump issued %d range requests, want 2", n)
|
||||
}
|
||||
}
|
||||
|
||||
// Random reads must return correct data and keep upstream requests bounded:
|
||||
// each read is either a window hit or a fresh request, never more than one.
|
||||
func TestReadAtSeekerRandomReads(t *testing.T) {
|
||||
data := randomData(32 * 1024 * 1024)
|
||||
var gets atomic.Int64
|
||||
ss := newMockSeekableStream(t, data, &gets)
|
||||
ra, err := stream.NewReadAtSeeker(ss, 0, true)
|
||||
if err != nil {
|
||||
t.Fatalf("NewReadAtSeeker() error = %v", err)
|
||||
}
|
||||
const chunk = 8 * 1024
|
||||
buf := make([]byte, chunk)
|
||||
r := rand.New(rand.NewSource(7))
|
||||
for i := 0; i < 200; i++ {
|
||||
off := r.Int63n(int64(len(data)) - chunk)
|
||||
readAtFull(t, ra, data, off, buf)
|
||||
}
|
||||
if n := gets.Load(); n > 200 {
|
||||
t.Fatalf("random reads issued %d range requests, want <= 200", n)
|
||||
}
|
||||
}
|
||||
+71
-34
@@ -8,6 +8,7 @@ import (
|
||||
"io"
|
||||
"math"
|
||||
"os"
|
||||
"sort"
|
||||
"sync"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
@@ -358,10 +359,72 @@ func (r *ReaderUpdatingProgress) Close() error {
|
||||
type RangeReadReadAtSeeker struct {
|
||||
ss *SeekableStream
|
||||
masterOff int64
|
||||
readerMap sync.Map
|
||||
readers orderedReaders
|
||||
headCache *headCache
|
||||
}
|
||||
|
||||
type orderedReaders struct {
|
||||
mu sync.Mutex
|
||||
m map[int64]io.Reader
|
||||
keys []int64
|
||||
}
|
||||
|
||||
func (o *orderedReaders) store(off int64, r io.Reader) {
|
||||
o.mu.Lock()
|
||||
defer o.mu.Unlock()
|
||||
if _, ok := o.m[off]; ok {
|
||||
o.m[off] = r
|
||||
return
|
||||
}
|
||||
if o.m == nil {
|
||||
o.m = make(map[int64]io.Reader)
|
||||
}
|
||||
i := sort.Search(len(o.keys), func(i int) bool { return o.keys[i] >= off })
|
||||
o.keys = append(o.keys, 0)
|
||||
copy(o.keys[i+1:], o.keys[i:])
|
||||
o.keys[i] = off
|
||||
o.m[off] = r
|
||||
}
|
||||
|
||||
func (o *orderedReaders) takeExact(off int64) (io.Reader, bool) {
|
||||
o.mu.Lock()
|
||||
defer o.mu.Unlock()
|
||||
r, ok := o.m[off]
|
||||
if ok {
|
||||
delete(o.m, off)
|
||||
o.removeKey(off)
|
||||
}
|
||||
return r, ok
|
||||
}
|
||||
|
||||
func (o *orderedReaders) takeBest(off int64) (io.Reader, int64, bool) {
|
||||
o.mu.Lock()
|
||||
defer o.mu.Unlock()
|
||||
if r, ok := o.m[off]; ok {
|
||||
delete(o.m, off)
|
||||
o.removeKey(off)
|
||||
return r, off, true
|
||||
}
|
||||
i := sort.Search(len(o.keys), func(i int) bool { return o.keys[i] >= off })
|
||||
if i == 0 {
|
||||
return nil, 0, false
|
||||
}
|
||||
k := o.keys[i-1]
|
||||
if off-k > 4*utils.MB {
|
||||
return nil, 0, false
|
||||
}
|
||||
r := o.m[k]
|
||||
delete(o.m, k)
|
||||
o.removeKey(k)
|
||||
return r, k, true
|
||||
}
|
||||
|
||||
func (o *orderedReaders) removeKey(k int64) {
|
||||
i := sort.Search(len(o.keys), func(i int) bool { return o.keys[i] >= k })
|
||||
copy(o.keys[i:], o.keys[i+1:])
|
||||
o.keys = o.keys[:len(o.keys)-1]
|
||||
}
|
||||
|
||||
type headCache struct {
|
||||
reader io.Reader
|
||||
bufs [][]byte
|
||||
@@ -396,7 +459,7 @@ func (r *headCache) Close() error {
|
||||
|
||||
func (r *RangeReadReadAtSeeker) InitHeadCache() {
|
||||
if r.masterOff == 0 {
|
||||
value, _ := r.readerMap.LoadAndDelete(int64(0))
|
||||
value, _ := r.readers.takeExact(0)
|
||||
r.headCache = &headCache{reader: value.(io.Reader)}
|
||||
r.ss.Closers.Add(r.headCache)
|
||||
}
|
||||
@@ -422,9 +485,9 @@ func NewReadAtSeeker(ss *SeekableStream, offset int64, forceRange ...bool) (mode
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r.readerMap.Store(int64(offset), reader)
|
||||
r.readers.store(offset, reader)
|
||||
} else {
|
||||
r.readerMap.Store(int64(offset), ss)
|
||||
r.readers.store(0, ss)
|
||||
}
|
||||
return r, nil
|
||||
}
|
||||
@@ -442,41 +505,15 @@ func NewMultiReaderAt(ss []*SeekableStream) (readerutil.SizeReaderAt, error) {
|
||||
}
|
||||
|
||||
func (r *RangeReadReadAtSeeker) getReaderAtOffset(off int64) (io.Reader, error) {
|
||||
for {
|
||||
var cur int64 = -1
|
||||
r.readerMap.Range(func(key, value any) bool {
|
||||
k := key.(int64)
|
||||
if off == k {
|
||||
cur = k
|
||||
return false
|
||||
}
|
||||
if off > k && off-k <= 4*utils.MB && k > cur {
|
||||
cur = k
|
||||
}
|
||||
return true
|
||||
})
|
||||
if cur < 0 {
|
||||
break
|
||||
}
|
||||
v, ok := r.readerMap.LoadAndDelete(int64(cur))
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
rr := v.(io.Reader)
|
||||
if off == int64(cur) {
|
||||
// logrus.Debugf("getReaderAtOffset match_%d", off)
|
||||
if rr, cur, ok := r.readers.takeBest(off); ok {
|
||||
if cur == off {
|
||||
return rr, nil
|
||||
}
|
||||
n, _ := utils.CopyWithBufferN(io.Discard, rr, off-cur)
|
||||
cur += n
|
||||
if cur == off {
|
||||
// logrus.Debugf("getReaderAtOffset old_%d", off)
|
||||
if cur+n == off {
|
||||
return rr, nil
|
||||
}
|
||||
break
|
||||
}
|
||||
|
||||
// logrus.Debugf("getReaderAtOffset new_%d", off)
|
||||
reader, err := r.ss.RangeRead(http_range.Range{Start: off, Length: -1})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -501,7 +538,7 @@ func (r *RangeReadReadAtSeeker) ReadAt(p []byte, off int64) (n int, err error) {
|
||||
off += int64(n)
|
||||
switch err {
|
||||
case nil:
|
||||
r.readerMap.Store(int64(off), rr)
|
||||
r.readers.store(off, rr)
|
||||
case io.ErrUnexpectedEOF:
|
||||
err = io.EOF
|
||||
}
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
package task_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/task"
|
||||
)
|
||||
|
||||
func TestTaskExtensionRestoresAPIURL(t *testing.T) {
|
||||
const want = "https://openlist.example"
|
||||
extension := task.TaskExtension{ApiUrl: want}
|
||||
|
||||
extension.SetCtx(context.Background())
|
||||
|
||||
if got := conf.GetApiUrl(extension.Ctx()); got != want {
|
||||
t.Fatalf("restored origin = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
@@ -31,6 +31,5 @@ func GetApiUrlFromRequest(r *http.Request) string {
|
||||
}
|
||||
|
||||
func GetApiUrl(ctx context.Context) string {
|
||||
api, _ := ctx.Value(conf.ApiUrlKey).(string)
|
||||
return api
|
||||
return conf.GetApiUrl(ctx)
|
||||
}
|
||||
|
||||
@@ -34,9 +34,7 @@ func Proxy(w http.ResponseWriter, r *http.Request, link *model.Link, file model.
|
||||
if link.RangeReader == nil {
|
||||
r = r.WithContext(context.WithValue(r.Context(), conf.RequestHeaderKey, r.Header))
|
||||
}
|
||||
return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, &model.RangeReadCloser{
|
||||
RangeReader: rrf,
|
||||
})
|
||||
return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, rrf)
|
||||
}
|
||||
|
||||
if link.RangeReader != nil {
|
||||
@@ -45,9 +43,7 @@ func Proxy(w http.ResponseWriter, r *http.Request, link *model.Link, file model.
|
||||
if size <= 0 {
|
||||
size = file.GetSize()
|
||||
}
|
||||
return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, &model.RangeReadCloser{
|
||||
RangeReader: link.RangeReader,
|
||||
})
|
||||
return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, link.RangeReader)
|
||||
}
|
||||
|
||||
//transparent proxy
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
)
|
||||
|
||||
func TestProxyCancelledPartitionedReaderDoesNotPanic(t *testing.T) {
|
||||
oldConf := conf.Conf
|
||||
conf.Conf = conf.DefaultConfig("data")
|
||||
t.Cleanup(func() { conf.Conf = oldConf })
|
||||
link := &model.Link{
|
||||
Concurrency: 2,
|
||||
PartSize: 4,
|
||||
RangeReader: stream.RangeReaderFunc(func(ctx context.Context, requested http_range.Range) (io.ReadCloser, error) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return io.NopCloser(bytes.NewReader([]byte("0123456789abcdef")[requested.Start : requested.Start+requested.Length])), nil
|
||||
}),
|
||||
}
|
||||
file := &model.Object{Name: "fixture.bin", Size: 16}
|
||||
for range 32 {
|
||||
func() {
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
t.Errorf("Proxy panicked on cancelled partitioned read: %v", recovered)
|
||||
}
|
||||
}()
|
||||
r := httptest.NewRequest(http.MethodGet, "/proxy/fixture.bin", nil)
|
||||
ctx, cancel := context.WithCancel(r.Context())
|
||||
cancel()
|
||||
w := httptest.NewRecorder()
|
||||
_ = Proxy(w, r.WithContext(ctx), link, file)
|
||||
if bytes.Contains(w.Body.Bytes(), []byte("0123456789abcdef")) {
|
||||
t.Errorf("cancelled response contained file contents: %q", w.Body.String())
|
||||
}
|
||||
}()
|
||||
}
|
||||
}
|
||||
+3
-2
@@ -94,10 +94,11 @@ func (f *FileUploadProxy) Close() error {
|
||||
return err
|
||||
}
|
||||
arr := make([]byte, 512)
|
||||
if _, err := f.buffer.Read(arr); err != nil {
|
||||
n, err := f.buffer.Read(arr)
|
||||
if err != nil && err != io.EOF {
|
||||
return err
|
||||
}
|
||||
contentType := http.DetectContentType(arr)
|
||||
contentType := http.DetectContentType(arr[:n])
|
||||
if _, err := f.buffer.Seek(0, io.SeekStart); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -113,6 +113,10 @@ func FsMove(c *gin.Context) {
|
||||
srcDir += "/"
|
||||
}
|
||||
for i, name := range req.Names {
|
||||
if err := checkRelativePath(name); err != nil {
|
||||
common.ErrorResp(c, err, 403)
|
||||
return
|
||||
}
|
||||
// ensure req.Names is not a relative path
|
||||
srcPath := stdpath.Join(srcDir, name)
|
||||
if !strings.HasPrefix(srcPath+"/", srcDir) {
|
||||
@@ -216,6 +220,10 @@ func FsCopy(c *gin.Context) {
|
||||
srcDir += "/"
|
||||
}
|
||||
for i, name := range req.Names {
|
||||
if err := checkRelativePath(name); err != nil {
|
||||
common.ErrorResp(c, err, 403)
|
||||
return
|
||||
}
|
||||
// ensure req.Names is not a relative path
|
||||
srcPath := stdpath.Join(srcDir, name)
|
||||
if !strings.HasPrefix(srcPath+"/", srcDir) {
|
||||
@@ -373,6 +381,10 @@ func FsRemove(c *gin.Context) {
|
||||
reqPath += "/"
|
||||
}
|
||||
for i, name := range req.Names {
|
||||
if err := checkRelativePath(name); err != nil {
|
||||
common.ErrorResp(c, err, 403)
|
||||
return
|
||||
}
|
||||
fullPath := stdpath.Join(reqPath, name)
|
||||
if !strings.HasPrefix(fullPath+"/", reqPath) {
|
||||
req.Names[i] = ""
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
package handles
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
_ "github.com/OpenListTeam/OpenList/v4/drivers/local"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/db"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/op"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func setupBackslashTraversalTest(t *testing.T, root string, permission int32) *model.User {
|
||||
t.Helper()
|
||||
database, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
conf.Conf = conf.DefaultConfig(t.TempDir())
|
||||
db.Init(database)
|
||||
addition, err := utils.Json.MarshalToString(map[string]string{"root_folder_path": root})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = op.CreateStorage(context.Background(), model.Storage{
|
||||
Driver: "Local", MountPath: "/", Addition: addition,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return &model.User{
|
||||
Username: "restricted-user", BasePath: "/team/a", Role: model.GENERAL,
|
||||
Permission: permission,
|
||||
}
|
||||
}
|
||||
|
||||
func prepareBackslashTraversalFs(t *testing.T) (root string, secretPath string) {
|
||||
t.Helper()
|
||||
root = t.TempDir()
|
||||
if err := os.MkdirAll(filepath.Join(root, "team", "a", "writable"), 0o700); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Join(root, "team", "ab"), 0o700); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
secretPath = filepath.Join(root, "team", "ab", "secret.txt")
|
||||
if err := os.WriteFile(secretPath, []byte("synthetic-secret"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return root, secretPath
|
||||
}
|
||||
|
||||
func invokeHandler(t *testing.T, user *model.User, payload any, handler gin.HandlerFunc) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
recorder := httptest.NewRecorder()
|
||||
ctx, _ := gin.CreateTestContext(recorder)
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/fs/remove", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req = req.WithContext(context.WithValue(req.Context(), conf.UserKey, user))
|
||||
ctx.Request = req
|
||||
handler(ctx)
|
||||
return recorder
|
||||
}
|
||||
|
||||
func TestFsRemoveRejectsBackslashTraversal(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
root, secretPath := prepareBackslashTraversalFs(t)
|
||||
user := setupBackslashTraversalTest(t, root, 1<<3|1<<7)
|
||||
|
||||
for _, name := range []string{"../../ab/secret.txt", `..\..\ab\secret.txt`} {
|
||||
recorder := invokeHandler(t, user, map[string]any{"dir": "/writable", "names": []string{name}}, FsRemove)
|
||||
if recorder.Code != http.StatusOK || !strings.Contains(recorder.Body.String(), `"code":403`) {
|
||||
t.Fatalf("payload %q: got status=%d body=%s, want 403", name, recorder.Code, recorder.Body.String())
|
||||
}
|
||||
if _, err := os.Stat(secretPath); err != nil {
|
||||
t.Fatalf("payload %q deleted sibling file: %v", name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFsMoveRejectsBackslashTraversal(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
root, secretPath := prepareBackslashTraversalFs(t)
|
||||
user := setupBackslashTraversalTest(t, root, 1<<3|1<<5)
|
||||
|
||||
recorder := invokeHandler(t, user, map[string]any{
|
||||
"src_dir": "/writable", "dst_dir": "/", "names": []string{`..\..\ab\secret.txt`},
|
||||
}, FsMove)
|
||||
if recorder.Code != http.StatusOK || !strings.Contains(recorder.Body.String(), `"code":403`) {
|
||||
t.Fatalf("got status=%d body=%s, want 403", recorder.Code, recorder.Body.String())
|
||||
}
|
||||
if _, err := os.Stat(secretPath); err != nil {
|
||||
t.Fatalf("backslash traversal moved sibling file: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFsCopyRejectsBackslashTraversal(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
root, secretPath := prepareBackslashTraversalFs(t)
|
||||
user := setupBackslashTraversalTest(t, root, 1<<3|1<<6)
|
||||
|
||||
recorder := invokeHandler(t, user, map[string]any{
|
||||
"src_dir": "/writable", "dst_dir": "/", "names": []string{`..\..\ab\secret.txt`},
|
||||
}, FsCopy)
|
||||
if recorder.Code != http.StatusOK || !strings.Contains(recorder.Body.String(), `"code":403`) {
|
||||
t.Fatalf("got status=%d body=%s, want 403", recorder.Code, recorder.Body.String())
|
||||
}
|
||||
if _, err := os.Stat(secretPath); err != nil {
|
||||
t.Fatalf("backslash traversal affected sibling file: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -347,11 +347,13 @@ func FsGet(c *gin.Context, req *FsGetReq, user *model.User) {
|
||||
}
|
||||
}
|
||||
}
|
||||
var related []model.Obj
|
||||
parentPath := stdpath.Dir(reqPath)
|
||||
sameLevelFiles, err := fs.List(c.Request.Context(), parentPath, &fs.ListArgs{})
|
||||
if err == nil {
|
||||
related = filterRelated(sameLevelFiles, obj)
|
||||
var related []model.Obj
|
||||
if !obj.IsDir() && utils.GetFileType(obj.GetName()) == conf.VIDEO {
|
||||
sameLevelFiles, err := fs.List(c.Request.Context(), parentPath, &fs.ListArgs{})
|
||||
if err == nil {
|
||||
related = filterRelated(sameLevelFiles, obj)
|
||||
}
|
||||
}
|
||||
parentMeta, _ := op.GetNearestMeta(parentPath)
|
||||
thumb, _ := model.GetThumb(obj)
|
||||
@@ -366,7 +368,7 @@ func FsGet(c *gin.Context, req *FsGetReq, user *model.User) {
|
||||
HashInfoStr: obj.GetHash().String(),
|
||||
HashInfo: obj.GetHash().Export(),
|
||||
Sign: common.Sign(obj, parentPath, isEncrypt(meta, reqPath)),
|
||||
Type: utils.GetFileType(obj.GetName()),
|
||||
Type: utils.GetObjType(obj.GetName(), obj.IsDir()),
|
||||
Thumb: thumb,
|
||||
MountDetails: mountDetails,
|
||||
},
|
||||
|
||||
@@ -3,15 +3,6 @@ package handles
|
||||
import (
|
||||
"strings"
|
||||
|
||||
_115 "github.com/OpenListTeam/OpenList/v4/drivers/115"
|
||||
_115_open "github.com/OpenListTeam/OpenList/v4/drivers/115_open"
|
||||
_123 "github.com/OpenListTeam/OpenList/v4/drivers/123"
|
||||
_123_open "github.com/OpenListTeam/OpenList/v4/drivers/123_open"
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/guangyapan"
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/pikpak"
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/thunder"
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/thunder_browser"
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/thunderx"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/errs"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
@@ -23,6 +14,44 @@ import (
|
||||
"github.com/pkg/errors"
|
||||
)
|
||||
|
||||
func saveAndInitOfflineDownloadTool(c *gin.Context, name string, items []model.SettingItem) (string, bool) {
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return "", false
|
||||
}
|
||||
downloadTool, err := tool.Tools.Get(name)
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return "", false
|
||||
}
|
||||
version, err := downloadTool.Init()
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return "", false
|
||||
}
|
||||
return version, true
|
||||
}
|
||||
|
||||
func validateOfflineDownloadStorage(c *gin.Context, tempDir, nativeTool string) bool {
|
||||
if tempDir == "" {
|
||||
return true
|
||||
}
|
||||
storage, _, err := op.GetStorageAndActualPath(tempDir)
|
||||
if err != nil {
|
||||
common.ErrorStrResp(c, "storage does not exists", 400)
|
||||
return false
|
||||
}
|
||||
if storage.Config().CheckStatus && storage.GetStorage().Status != op.WORK {
|
||||
common.ErrorStrResp(c, "storage not init: "+storage.GetStorage().Status, 400)
|
||||
return false
|
||||
}
|
||||
if tool.NativeToolName(storage) != nativeTool {
|
||||
common.ErrorStrResp(c, "unsupported storage driver for offline download, only "+nativeTool+" is supported", 400)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
type SetAria2Req struct {
|
||||
Uri string `json:"uri" form:"uri"`
|
||||
Secret string `json:"secret" form:"secret"`
|
||||
@@ -38,18 +67,8 @@ func SetAria2(c *gin.Context) {
|
||||
{Key: conf.Aria2Uri, Value: req.Uri, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
{Key: conf.Aria2Secret, Value: req.Secret, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("aria2")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
version, err := _tool.Init()
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
version, ok := saveAndInitOfflineDownloadTool(c, "aria2", items)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, version)
|
||||
@@ -70,17 +89,7 @@ func SetQbittorrent(c *gin.Context) {
|
||||
{Key: conf.QbittorrentUrl, Value: req.Url, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
{Key: conf.QbittorrentSeedtime, Value: req.Seedtime, Type: conf.TypeNumber, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("qBittorrent")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "qBittorrent", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
@@ -101,17 +110,7 @@ func SetTransmission(c *gin.Context) {
|
||||
{Key: conf.TransmissionUri, Value: req.Uri, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
{Key: conf.TransmissionSeedtime, Value: req.Seedtime, Type: conf.TypeNumber, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("Transmission")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "Transmission", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
@@ -127,35 +126,13 @@ func Set115(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 400)
|
||||
return
|
||||
}
|
||||
if req.TempDir != "" {
|
||||
storage, _, err := op.GetStorageAndActualPath(req.TempDir)
|
||||
if err != nil {
|
||||
common.ErrorStrResp(c, "storage does not exists", 400)
|
||||
return
|
||||
}
|
||||
if storage.Config().CheckStatus && storage.GetStorage().Status != op.WORK {
|
||||
common.ErrorStrResp(c, "storage not init: "+storage.GetStorage().Status, 400)
|
||||
return
|
||||
}
|
||||
if _, ok := storage.(*_115.Pan115); !ok {
|
||||
common.ErrorStrResp(c, "unsupported storage driver for offline download, only 115 Cloud is supported", 400)
|
||||
return
|
||||
}
|
||||
if !validateOfflineDownloadStorage(c, req.TempDir, "115 Cloud") {
|
||||
return
|
||||
}
|
||||
items := []model.SettingItem{
|
||||
{Key: conf.Pan115TempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("115 Cloud")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "115 Cloud", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
@@ -171,35 +148,13 @@ func Set115Open(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 400)
|
||||
return
|
||||
}
|
||||
if req.TempDir != "" {
|
||||
storage, _, err := op.GetStorageAndActualPath(req.TempDir)
|
||||
if err != nil {
|
||||
common.ErrorStrResp(c, "storage does not exists", 400)
|
||||
return
|
||||
}
|
||||
if storage.Config().CheckStatus && storage.GetStorage().Status != op.WORK {
|
||||
common.ErrorStrResp(c, "storage not init: "+storage.GetStorage().Status, 400)
|
||||
return
|
||||
}
|
||||
if _, ok := storage.(*_115_open.Open115); !ok {
|
||||
common.ErrorStrResp(c, "unsupported storage driver for offline download, only 115 Open is supported", 400)
|
||||
return
|
||||
}
|
||||
if !validateOfflineDownloadStorage(c, req.TempDir, "115 Open") {
|
||||
return
|
||||
}
|
||||
items := []model.SettingItem{
|
||||
{Key: conf.Pan115OpenTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("115 Open")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "115 Open", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
@@ -215,35 +170,13 @@ func Set123Pan(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 400)
|
||||
return
|
||||
}
|
||||
if req.TempDir != "" {
|
||||
storage, _, err := op.GetStorageAndActualPath(req.TempDir)
|
||||
if err != nil {
|
||||
common.ErrorStrResp(c, "storage does not exists", 400)
|
||||
return
|
||||
}
|
||||
if storage.Config().CheckStatus && storage.GetStorage().Status != op.WORK {
|
||||
common.ErrorStrResp(c, "storage not init: "+storage.GetStorage().Status, 400)
|
||||
return
|
||||
}
|
||||
if _, ok := storage.(*_123.Pan123); !ok {
|
||||
common.ErrorStrResp(c, "unsupported storage driver for offline download, only 123Pan is supported", 400)
|
||||
return
|
||||
}
|
||||
if !validateOfflineDownloadStorage(c, req.TempDir, "123Pan") {
|
||||
return
|
||||
}
|
||||
items := []model.SettingItem{
|
||||
{Key: conf.Pan123TempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("123Pan")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "123Pan", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
@@ -260,36 +193,14 @@ func Set123Open(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 400)
|
||||
return
|
||||
}
|
||||
if req.TempDir != "" {
|
||||
storage, _, err := op.GetStorageAndActualPath(req.TempDir)
|
||||
if err != nil {
|
||||
common.ErrorStrResp(c, "storage does not exists", 400)
|
||||
return
|
||||
}
|
||||
if storage.Config().CheckStatus && storage.GetStorage().Status != op.WORK {
|
||||
common.ErrorStrResp(c, "storage not init: "+storage.GetStorage().Status, 400)
|
||||
return
|
||||
}
|
||||
if _, ok := storage.(*_123_open.Open123); !ok {
|
||||
common.ErrorStrResp(c, "unsupported storage driver for offline download, only 123 Open is supported", 400)
|
||||
return
|
||||
}
|
||||
if !validateOfflineDownloadStorage(c, req.TempDir, "123 Open") {
|
||||
return
|
||||
}
|
||||
items := []model.SettingItem{
|
||||
{Key: conf.Pan123OpenTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
{Key: conf.Pan123OpenOfflineDownloadCallbackUrl, Value: req.CallbackUrl, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("123 Open")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "123 Open", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
@@ -305,35 +216,13 @@ func SetPikPak(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 400)
|
||||
return
|
||||
}
|
||||
if req.TempDir != "" {
|
||||
storage, _, err := op.GetStorageAndActualPath(req.TempDir)
|
||||
if err != nil {
|
||||
common.ErrorStrResp(c, "storage does not exists", 400)
|
||||
return
|
||||
}
|
||||
if storage.Config().CheckStatus && storage.GetStorage().Status != op.WORK {
|
||||
common.ErrorStrResp(c, "storage not init: "+storage.GetStorage().Status, 400)
|
||||
return
|
||||
}
|
||||
if _, ok := storage.(*pikpak.PikPak); !ok {
|
||||
common.ErrorStrResp(c, "unsupported storage driver for offline download, only PikPak is supported", 400)
|
||||
return
|
||||
}
|
||||
if !validateOfflineDownloadStorage(c, req.TempDir, "PikPak") {
|
||||
return
|
||||
}
|
||||
items := []model.SettingItem{
|
||||
{Key: conf.PikPakTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("PikPak")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "PikPak", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
@@ -349,35 +238,13 @@ func SetThunder(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 400)
|
||||
return
|
||||
}
|
||||
if req.TempDir != "" {
|
||||
storage, _, err := op.GetStorageAndActualPath(req.TempDir)
|
||||
if err != nil {
|
||||
common.ErrorStrResp(c, "storage does not exists", 400)
|
||||
return
|
||||
}
|
||||
if storage.Config().CheckStatus && storage.GetStorage().Status != op.WORK {
|
||||
common.ErrorStrResp(c, "storage not init: "+storage.GetStorage().Status, 400)
|
||||
return
|
||||
}
|
||||
if _, ok := storage.(*thunder.Thunder); !ok {
|
||||
common.ErrorStrResp(c, "unsupported storage driver for offline download, only Thunder is supported", 400)
|
||||
return
|
||||
}
|
||||
if !validateOfflineDownloadStorage(c, req.TempDir, "Thunder") {
|
||||
return
|
||||
}
|
||||
items := []model.SettingItem{
|
||||
{Key: conf.ThunderTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("Thunder")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "Thunder", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
@@ -393,35 +260,13 @@ func SetThunderX(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 400)
|
||||
return
|
||||
}
|
||||
if req.TempDir != "" {
|
||||
storage, _, err := op.GetStorageAndActualPath(req.TempDir)
|
||||
if err != nil {
|
||||
common.ErrorStrResp(c, "storage does not exists", 400)
|
||||
return
|
||||
}
|
||||
if storage.Config().CheckStatus && storage.GetStorage().Status != op.WORK {
|
||||
common.ErrorStrResp(c, "storage not init: "+storage.GetStorage().Status, 400)
|
||||
return
|
||||
}
|
||||
if _, ok := storage.(*thunderx.ThunderX); !ok {
|
||||
common.ErrorStrResp(c, "unsupported storage driver for offline download, only ThunderX is supported", 400)
|
||||
return
|
||||
}
|
||||
if !validateOfflineDownloadStorage(c, req.TempDir, "ThunderX") {
|
||||
return
|
||||
}
|
||||
items := []model.SettingItem{
|
||||
{Key: conf.ThunderXTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("ThunderX")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "ThunderX", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
@@ -437,37 +282,13 @@ func SetThunderBrowser(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 400)
|
||||
return
|
||||
}
|
||||
if req.TempDir != "" {
|
||||
storage, _, err := op.GetStorageAndActualPath(req.TempDir)
|
||||
if err != nil {
|
||||
common.ErrorStrResp(c, "storage does not exists", 400)
|
||||
return
|
||||
}
|
||||
if storage.Config().CheckStatus && storage.GetStorage().Status != op.WORK {
|
||||
common.ErrorStrResp(c, "storage not init: "+storage.GetStorage().Status, 400)
|
||||
return
|
||||
}
|
||||
switch storage.(type) {
|
||||
case *thunder_browser.ThunderBrowser, *thunder_browser.ThunderBrowserExpert:
|
||||
default:
|
||||
common.ErrorStrResp(c, "unsupported storage driver for offline download, only ThunderBrowser is supported", 400)
|
||||
return
|
||||
}
|
||||
if !validateOfflineDownloadStorage(c, req.TempDir, "ThunderBrowser") {
|
||||
return
|
||||
}
|
||||
items := []model.SettingItem{
|
||||
{Key: conf.ThunderBrowserTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("ThunderBrowser")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "ThunderBrowser", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
@@ -483,35 +304,13 @@ func SetGuangYaPan(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 400)
|
||||
return
|
||||
}
|
||||
if req.TempDir != "" {
|
||||
storage, _, err := op.GetStorageAndActualPath(req.TempDir)
|
||||
if err != nil {
|
||||
common.ErrorStrResp(c, "storage does not exists", 400)
|
||||
return
|
||||
}
|
||||
if storage.Config().CheckStatus && storage.GetStorage().Status != op.WORK {
|
||||
common.ErrorStrResp(c, "storage not init: "+storage.GetStorage().Status, 400)
|
||||
return
|
||||
}
|
||||
if _, ok := storage.(*guangyapan.GuangYaPan); !ok {
|
||||
common.ErrorStrResp(c, "unsupported storage driver for offline download, only GuangYaPan is supported", 400)
|
||||
return
|
||||
}
|
||||
if !validateOfflineDownloadStorage(c, req.TempDir, "GuangYaPan") {
|
||||
return
|
||||
}
|
||||
items := []model.SettingItem{
|
||||
{Key: conf.GuangYaPanTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
|
||||
}
|
||||
if err := op.SaveSettingItems(items); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
_tool, err := tool.Tools.Get("GuangYaPan")
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if _, err := _tool.Init(); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
if _, ok := saveAndInitOfflineDownloadTool(c, "GuangYaPan", items); !ok {
|
||||
return
|
||||
}
|
||||
common.SuccessResp(c, "ok")
|
||||
|
||||
@@ -0,0 +1,155 @@
|
||||
package handles
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
_ "github.com/OpenListTeam/OpenList/v4/drivers/local"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/db"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/offline_download/tool"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/op"
|
||||
"github.com/OpenListTeam/OpenList/v4/server/common"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func init() {
|
||||
dataDir, err := os.MkdirTemp("", "openlist-handles-*")
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
conf.Conf = conf.DefaultConfig(dataDir)
|
||||
database, err := gorm.Open(sqlite.Open("file:handles?mode=memory&cache=shared"), &gorm.Config{})
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
db.Init(database)
|
||||
}
|
||||
|
||||
type settingsTool struct {
|
||||
name string
|
||||
version string
|
||||
initCalls int
|
||||
}
|
||||
|
||||
func (t *settingsTool) Name() string { return t.name }
|
||||
func (*settingsTool) Items() []model.SettingItem { return nil }
|
||||
func (t *settingsTool) Init() (string, error) { t.initCalls++; return t.version, nil }
|
||||
func (*settingsTool) IsReady() bool { return true }
|
||||
func (*settingsTool) AddURL(*tool.AddUrlArgs) (string, error) { return "", nil }
|
||||
func (*settingsTool) Remove(*tool.DownloadTask) error { return nil }
|
||||
func (*settingsTool) Status(*tool.DownloadTask) (*tool.Status, error) {
|
||||
return &tool.Status{}, nil
|
||||
}
|
||||
func (*settingsTool) Run(*tool.DownloadTask) error { return nil }
|
||||
|
||||
func TestOfflineDownloadSettingsPreserveSuccessPayloads(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
toolName string
|
||||
version string
|
||||
body string
|
||||
handler gin.HandlerFunc
|
||||
wantData string
|
||||
settingKey string
|
||||
}{
|
||||
{
|
||||
name: "aria2 returns version",
|
||||
toolName: "aria2",
|
||||
version: "v-test",
|
||||
body: `{"uri":"http://aria2","secret":"secret"}`,
|
||||
handler: SetAria2,
|
||||
wantData: "v-test",
|
||||
settingKey: conf.Aria2Uri,
|
||||
},
|
||||
{
|
||||
name: "qBittorrent returns ok",
|
||||
toolName: "qBittorrent",
|
||||
version: "ignored",
|
||||
body: `{"url":"http://qbit","seedtime":"1"}`,
|
||||
handler: SetQbittorrent,
|
||||
wantData: "ok",
|
||||
settingKey: conf.QbittorrentUrl,
|
||||
},
|
||||
}
|
||||
gin.SetMode(gin.TestMode)
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
fake := &settingsTool{name: test.toolName, version: test.version}
|
||||
previous, existed := tool.Tools[test.toolName]
|
||||
tool.Tools[test.toolName] = fake
|
||||
t.Cleanup(func() {
|
||||
if existed {
|
||||
tool.Tools[test.toolName] = previous
|
||||
} else {
|
||||
delete(tool.Tools, test.toolName)
|
||||
}
|
||||
_ = db.DeleteSettingItemByKey(test.settingKey)
|
||||
op.SettingCacheUpdate()
|
||||
})
|
||||
|
||||
response := httptest.NewRecorder()
|
||||
ctx, _ := gin.CreateTestContext(response)
|
||||
ctx.Request = httptest.NewRequest(http.MethodPost, "/", strings.NewReader(test.body))
|
||||
ctx.Request.Header.Set("Content-Type", "application/json")
|
||||
test.handler(ctx)
|
||||
|
||||
var result common.Resp[string]
|
||||
if err := json.Unmarshal(response.Body.Bytes(), &result); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Code != 200 || result.Data != test.wantData {
|
||||
t.Fatalf("response = %#v, want code 200 and data %q", result, test.wantData)
|
||||
}
|
||||
if fake.initCalls != 1 {
|
||||
t.Fatalf("Init calls = %d, want 1", fake.initCalls)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateOfflineDownloadStorageRejectsWrongNativeTool(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
addition, err := json.Marshal(struct {
|
||||
RootFolderPath string `json:"root_folder_path"`
|
||||
}{RootFolderPath: root})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
mount := "/" + strings.ReplaceAll(t.Name(), "/", "_")
|
||||
storageID, err := op.CreateStorage(context.Background(), model.Storage{
|
||||
Driver: "Local",
|
||||
MountPath: mount,
|
||||
Addition: string(addition),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
if err := op.DeleteStorageById(context.Background(), storageID); err != nil {
|
||||
t.Errorf("delete fixture storage: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
response := httptest.NewRecorder()
|
||||
ctx, _ := gin.CreateTestContext(response)
|
||||
if validateOfflineDownloadStorage(ctx, mount, "Thunder") {
|
||||
t.Fatal("Local storage unexpectedly accepted as Thunder")
|
||||
}
|
||||
var result common.Resp[any]
|
||||
if err := json.Unmarshal(response.Body.Bytes(), &result); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := "unsupported storage driver for offline download, only Thunder is supported"
|
||||
if result.Code != 400 || result.Message != want {
|
||||
t.Fatalf("response = %#v, want code 400 and message %q", result, want)
|
||||
}
|
||||
}
|
||||
@@ -44,14 +44,7 @@ func Search(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
nodes, total, err := search.SearchFiltered(c, req.SearchReq, func(node model.SearchNode) bool {
|
||||
if !utils.IsSubPath(user.BasePath, node.Parent) {
|
||||
return false
|
||||
}
|
||||
meta, err := op.GetNearestMeta(node.Parent)
|
||||
if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) {
|
||||
return false
|
||||
}
|
||||
return common.CanAccess(user, meta, path.Join(node.Parent, node.Name), req.Password)
|
||||
return isSearchNodeAccessible(user, node, req.Password, op.GetNearestMeta)
|
||||
})
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
@@ -63,6 +56,22 @@ func Search(c *gin.Context) {
|
||||
})
|
||||
}
|
||||
|
||||
func isSearchNodeAccessible(user *model.User, node model.SearchNode, password string, resolveMeta func(string) (*model.Meta, error)) bool {
|
||||
if !utils.IsSubPath(user.BasePath, node.Parent) {
|
||||
return false
|
||||
}
|
||||
nodePath := path.Join(node.Parent, node.Name)
|
||||
metaPath := node.Parent
|
||||
if node.IsDir {
|
||||
metaPath = nodePath
|
||||
}
|
||||
meta, err := resolveMeta(metaPath)
|
||||
if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) {
|
||||
return false
|
||||
}
|
||||
return common.CanAccess(user, meta, nodePath, password)
|
||||
}
|
||||
|
||||
func nodeToSearchResp(node model.SearchNode) SearchResp {
|
||||
return SearchResp{
|
||||
SearchNode: node,
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
package handles
|
||||
|
||||
import (
|
||||
"path"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/errs"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
)
|
||||
|
||||
func fakeResolveMeta(metas map[string]*model.Meta) func(string) (*model.Meta, error) {
|
||||
return func(p string) (*model.Meta, error) {
|
||||
for {
|
||||
if meta, ok := metas[p]; ok {
|
||||
return meta, nil
|
||||
}
|
||||
if p == "/" {
|
||||
return nil, errs.MetaNotFound
|
||||
}
|
||||
p = path.Dir(p)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsSearchNodeAccessible(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
metas map[string]*model.Meta
|
||||
node model.SearchNode
|
||||
want bool
|
||||
wantMetaPath string
|
||||
}{
|
||||
{
|
||||
name: "restricted directory",
|
||||
metas: map[string]*model.Meta{"/private": {Path: "/private", ReadUsers: []uint{1}}},
|
||||
node: model.SearchNode{Parent: "/", Name: "private", IsDir: true},
|
||||
want: false,
|
||||
wantMetaPath: "/private",
|
||||
},
|
||||
{
|
||||
name: "restricted sub directory",
|
||||
metas: map[string]*model.Meta{"/private": {Path: "/private", ReadUsers: []uint{1}, ReadUsersSub: true}},
|
||||
node: model.SearchNode{Parent: "/private", Name: "sub", IsDir: true},
|
||||
want: false,
|
||||
wantMetaPath: "/private/sub",
|
||||
},
|
||||
{
|
||||
name: "file keeps parent scope",
|
||||
metas: map[string]*model.Meta{"/private": {Path: "/private", ReadUsers: []uint{1}}},
|
||||
node: model.SearchNode{Parent: "/private", Name: "a.txt", IsDir: false},
|
||||
want: true,
|
||||
wantMetaPath: "/private",
|
||||
},
|
||||
{
|
||||
name: "outside base path",
|
||||
node: model.SearchNode{Parent: "/other", Name: "private", IsDir: true},
|
||||
want: false,
|
||||
wantMetaPath: "",
|
||||
},
|
||||
}
|
||||
user := &model.User{ID: 2, BasePath: "/"}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
resolve := fakeResolveMeta(tt.metas)
|
||||
var gotMetaPath string
|
||||
spy := func(p string) (*model.Meta, error) {
|
||||
gotMetaPath = p
|
||||
return resolve(p)
|
||||
}
|
||||
if got := isSearchNodeAccessible(user, tt.node, "", spy); got != tt.want {
|
||||
t.Fatalf("isSearchNodeAccessible() = %v, want %v", got, tt.want)
|
||||
}
|
||||
if tt.wantMetaPath != "" && gotMetaPath != tt.wantMetaPath {
|
||||
t.Fatalf("meta resolved at %q, want %q", gotMetaPath, tt.wantMetaPath)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -54,7 +54,7 @@ func SharingGet(c *gin.Context, req *FsGetReq) {
|
||||
HashInfoStr: obj.GetHash().String(),
|
||||
HashInfo: obj.GetHash().Export(),
|
||||
Sign: "",
|
||||
Type: utils.GetFileType(obj.GetName()),
|
||||
Type: utils.GetObjType(obj.GetName(), obj.IsDir()),
|
||||
Thumb: thumb,
|
||||
},
|
||||
RawURL: url,
|
||||
|
||||
+51
-36
@@ -122,6 +122,53 @@ func generateSSOBindingToken(c *gin.Context, purpose, ssoID string) (string, err
|
||||
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(common.SecretKey)
|
||||
}
|
||||
|
||||
// ssoTargetOrigin returns the origin that is allowed to receive the SSO result
|
||||
// via postMessage. It honours the operator-configured sso_postmessage_origin so
|
||||
// a frontend served from a different origin than the API can still receive the
|
||||
// result; otherwise it falls back to the API origin, or "/" to restrict
|
||||
// delivery to same-origin openers when that cannot be resolved.
|
||||
func ssoTargetOrigin(c *gin.Context) string {
|
||||
if configured := setting.GetStr(conf.SSOPostMessageOrigin); configured != "" {
|
||||
if u, err := url.Parse(configured); err == nil &&
|
||||
(u.Scheme == "http" || u.Scheme == "https") &&
|
||||
u.Host != "" && u.User == nil &&
|
||||
(u.Path == "" || u.Path == "/") &&
|
||||
u.RawQuery == "" && u.Fragment == "" {
|
||||
return u.Scheme + "://" + u.Host
|
||||
}
|
||||
}
|
||||
u, err := url.Parse(common.GetApiUrl(c))
|
||||
if err != nil || u.Scheme == "" || u.Host == "" {
|
||||
return "/"
|
||||
}
|
||||
return u.Scheme + "://" + u.Host
|
||||
}
|
||||
|
||||
// ssoPostMessage hands the SSO result back to the window that started the login.
|
||||
// The target origin is pinned so that an arbitrary page cannot open the SSO
|
||||
// endpoint in a popup and read the payload out of the message event.
|
||||
func ssoPostMessage(c *gin.Context, payload map[string]string) {
|
||||
data, err := utils.Json.MarshalToString(payload)
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
origin, err := utils.Json.MarshalToString(ssoTargetOrigin(c))
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
html := fmt.Sprintf(`<!DOCTYPE html>
|
||||
<head></head>
|
||||
<body>
|
||||
<script>
|
||||
if (window.opener) { window.opener.postMessage(%s, %s) }
|
||||
window.close()
|
||||
</script>
|
||||
</body>`, data, origin)
|
||||
c.Data(200, "text/html; charset=utf-8", []byte(html))
|
||||
}
|
||||
|
||||
func ssoRedirectUri(c *gin.Context, useCompatibility bool, method string) string {
|
||||
if useCompatibility {
|
||||
return common.GetApiUrl(c) + "/api/auth/" + method
|
||||
@@ -338,15 +385,7 @@ func OIDCLoginCallback(c *gin.Context) {
|
||||
c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof)
|
||||
return
|
||||
}
|
||||
html := fmt.Sprintf(`<!DOCTYPE html>
|
||||
<head></head>
|
||||
<body>
|
||||
<script>
|
||||
window.opener.postMessage({"sso_id": "%s"}, "*")
|
||||
window.close()
|
||||
</script>
|
||||
</body>`, bindingProof)
|
||||
c.Data(200, "text/html; charset=utf-8", []byte(html))
|
||||
ssoPostMessage(c, map[string]string{"sso_id": bindingProof})
|
||||
return
|
||||
}
|
||||
if method == "sso_get_token" {
|
||||
@@ -367,15 +406,7 @@ func OIDCLoginCallback(c *gin.Context) {
|
||||
c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token)
|
||||
return
|
||||
}
|
||||
html := fmt.Sprintf(`<!DOCTYPE html>
|
||||
<head></head>
|
||||
<body>
|
||||
<script>
|
||||
window.opener.postMessage({"token":"%s"}, "*")
|
||||
window.close()
|
||||
</script>
|
||||
</body>`, token)
|
||||
c.Data(200, "text/html; charset=utf-8", []byte(html))
|
||||
ssoPostMessage(c, map[string]string{"token": token})
|
||||
return
|
||||
}
|
||||
}
|
||||
@@ -516,15 +547,7 @@ func SSOLoginCallback(c *gin.Context) {
|
||||
c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof)
|
||||
return
|
||||
}
|
||||
html := fmt.Sprintf(`<!DOCTYPE html>
|
||||
<head></head>
|
||||
<body>
|
||||
<script>
|
||||
window.opener.postMessage({"sso_id": "%s"}, "*")
|
||||
window.close()
|
||||
</script>
|
||||
</body>`, bindingProof)
|
||||
c.Data(200, "text/html; charset=utf-8", []byte(html))
|
||||
ssoPostMessage(c, map[string]string{"sso_id": bindingProof})
|
||||
return
|
||||
}
|
||||
username := utils.Json.Get(resp.Body(), usernameField).ToString()
|
||||
@@ -545,13 +568,5 @@ func SSOLoginCallback(c *gin.Context) {
|
||||
c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token)
|
||||
return
|
||||
}
|
||||
html := fmt.Sprintf(`<!DOCTYPE html>
|
||||
<head></head>
|
||||
<body>
|
||||
<script>
|
||||
window.opener.postMessage({"token":"%s"}, "*")
|
||||
window.close()
|
||||
</script>
|
||||
</body>`, token)
|
||||
c.Data(200, "text/html; charset=utf-8", []byte(html))
|
||||
ssoPostMessage(c, map[string]string{"token": token})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
package handles
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/op"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func ssoTestContext(apiUrl string) (*gin.Context, *httptest.ResponseRecorder) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
rec := httptest.NewRecorder()
|
||||
c, engine := gin.CreateTestContext(rec)
|
||||
// Matches server.Init, which is what lets GetApiUrl reach the value the
|
||||
// middleware stored on the request context.
|
||||
engine.ContextWithFallback = true
|
||||
req := httptest.NewRequest(http.MethodGet, "/api/auth/sso?method=sso_get_token", nil)
|
||||
if apiUrl != "" {
|
||||
req = req.WithContext(context.WithValue(req.Context(), conf.ApiUrlKey, apiUrl))
|
||||
}
|
||||
c.Request = req
|
||||
// Keep setting lookups off the (uninitialised) database: ssoTargetOrigin
|
||||
// reads sso_postmessage_origin through the setting cache.
|
||||
op.Cache.SetSetting(conf.SSOPostMessageOrigin, &model.SettingItem{
|
||||
Key: conf.SSOPostMessageOrigin,
|
||||
Value: "",
|
||||
})
|
||||
return c, rec
|
||||
}
|
||||
|
||||
// A page that opens the SSO endpoint in a popup must not be able to read the
|
||||
// token: the postMessage target origin has to name the site, never "*".
|
||||
func TestSSOPostMessagePinsTargetOrigin(t *testing.T) {
|
||||
c, rec := ssoTestContext("https://openlist.example.com/base")
|
||||
ssoPostMessage(c, map[string]string{"token": "secret-token"})
|
||||
|
||||
body := rec.Body.String()
|
||||
if strings.Contains(body, `"*"`) || strings.Contains(body, `, '*'`) {
|
||||
t.Fatalf("wildcard target origin present in response:\n%s", body)
|
||||
}
|
||||
if !strings.Contains(body, `"https://openlist.example.com"`) {
|
||||
t.Errorf("expected the site origin as target, got:\n%s", body)
|
||||
}
|
||||
if !strings.Contains(body, "secret-token") {
|
||||
t.Errorf("payload should still reach a legitimate opener, got:\n%s", body)
|
||||
}
|
||||
}
|
||||
|
||||
// A frontend served from a different origin than the API needs the operator to
|
||||
// be able to point the target at the frontend origin. The configured origin
|
||||
// must win over the API origin.
|
||||
func TestSSOPostMessageUsesConfiguredOrigin(t *testing.T) {
|
||||
c, rec := ssoTestContext("https://api.example.com/base")
|
||||
op.Cache.SetSetting(conf.SSOPostMessageOrigin, &model.SettingItem{
|
||||
Key: conf.SSOPostMessageOrigin,
|
||||
Value: "https://frontend.example.com",
|
||||
})
|
||||
defer op.Cache.ClearAll()
|
||||
|
||||
ssoPostMessage(c, map[string]string{"token": "secret-token"})
|
||||
|
||||
body := rec.Body.String()
|
||||
if !strings.Contains(body, `"https://frontend.example.com"`) {
|
||||
t.Errorf("expected the configured origin as target, got:\n%s", body)
|
||||
}
|
||||
}
|
||||
|
||||
// If the site URL cannot be resolved the fallback must tighten delivery to
|
||||
// same-origin openers, not widen it back to every origin.
|
||||
func TestSSOPostMessageFallsBackToSameOrigin(t *testing.T) {
|
||||
c, rec := ssoTestContext("")
|
||||
ssoPostMessage(c, map[string]string{"token": "secret-token"})
|
||||
|
||||
body := rec.Body.String()
|
||||
if strings.Contains(body, `"*"`) {
|
||||
t.Fatalf("fallback must not be a wildcard origin:\n%s", body)
|
||||
}
|
||||
if !strings.Contains(body, `"/"`) {
|
||||
t.Errorf(`expected "/" fallback origin, got:\n%s`, body)
|
||||
}
|
||||
}
|
||||
|
||||
// userID comes from the identity provider, so it must be encoded rather than
|
||||
// interpolated into the JS string literal it used to land in.
|
||||
func TestSSOPostMessageEscapesProviderControlledValue(t *testing.T) {
|
||||
c, rec := ssoTestContext("https://openlist.example.com")
|
||||
ssoPostMessage(c, map[string]string{"sso_id": `"});alert(document.domain);//`})
|
||||
|
||||
body := rec.Body.String()
|
||||
if strings.Contains(body, `alert(document.domain)`) && !strings.Contains(body, `\"`) {
|
||||
t.Fatalf("provider value was not escaped:\n%s", body)
|
||||
}
|
||||
if !strings.Contains(body, `\"});alert`) {
|
||||
t.Errorf("expected the injected quote to be escaped, got:\n%s", body)
|
||||
}
|
||||
}
|
||||
@@ -68,9 +68,11 @@ func (s *Server) callFSGet(c *gin.Context, raw json.RawMessage) (any, *rpcError)
|
||||
|
||||
parentPath := stdpath.Dir(reqPath)
|
||||
var related []model.Obj
|
||||
sameLevelFiles, err := fs.List(ctx, parentPath, &fs.ListArgs{})
|
||||
if err == nil {
|
||||
related = filterRelatedObjs(sameLevelFiles, obj)
|
||||
if !obj.IsDir() && utils.GetFileType(obj.GetName()) == conf.VIDEO {
|
||||
sameLevelFiles, err := fs.List(ctx, parentPath, &fs.ListArgs{})
|
||||
if err == nil {
|
||||
related = filterRelatedObjs(sameLevelFiles, obj)
|
||||
}
|
||||
}
|
||||
|
||||
parentMeta, _ := op.GetNearestMeta(parentPath)
|
||||
@@ -85,7 +87,7 @@ func (s *Server) callFSGet(c *gin.Context, raw json.RawMessage) (any, *rpcError)
|
||||
Created: obj.CreateTime(),
|
||||
Sign: common.Sign(obj, parentPath, isEncrypt(meta, reqPath)),
|
||||
Thumb: thumb,
|
||||
Type: utils.GetFileType(obj.GetName()),
|
||||
Type: utils.GetObjType(obj.GetName(), obj.IsDir()),
|
||||
HashInfoStr: obj.GetHash().String(),
|
||||
HashInfo: obj.GetHash().Export(),
|
||||
MountDetails: mountDetails,
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
package middlewares
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func TestStoragesLoadedAdmitsRequestOrigin(t *testing.T) {
|
||||
originalMode := gin.Mode()
|
||||
gin.SetMode(gin.TestMode)
|
||||
originalConf := conf.Conf
|
||||
originalLoaded := conf.StoragesLoaded
|
||||
t.Cleanup(func() {
|
||||
gin.SetMode(originalMode)
|
||||
conf.Conf = originalConf
|
||||
conf.StoragesLoaded = originalLoaded
|
||||
})
|
||||
conf.StoragesLoaded = true
|
||||
|
||||
router := gin.New()
|
||||
router.Use(StoragesLoaded)
|
||||
router.GET("/", func(c *gin.Context) {
|
||||
c.String(http.StatusOK, conf.GetApiUrl(c.Request.Context()))
|
||||
})
|
||||
|
||||
assertOrigin := func(name, siteURL, target string, header http.Header, want string) {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
conf.Conf = &conf.Config{SiteURL: siteURL}
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, target, nil)
|
||||
req.Header = header
|
||||
rec := httptest.NewRecorder()
|
||||
router.ServeHTTP(rec, req)
|
||||
|
||||
if got := rec.Body.String(); got != want {
|
||||
t.Fatalf("origin = %q, want %q", got, want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
assertOrigin(
|
||||
"configured site URL",
|
||||
"https://openlist.example/base/",
|
||||
"http://ignored.example/",
|
||||
nil,
|
||||
"https://openlist.example/base",
|
||||
)
|
||||
assertOrigin(
|
||||
"forwarded request",
|
||||
"",
|
||||
"http://internal.example/",
|
||||
http.Header{
|
||||
"X-Forwarded-Proto": {"https"},
|
||||
"X-Forwarded-Host": {"public.example"},
|
||||
},
|
||||
"https://public.example",
|
||||
)
|
||||
}
|
||||
+26
-64
@@ -4,11 +4,9 @@ package s3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/md5"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"path"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -32,7 +30,7 @@ import (
|
||||
|
||||
var (
|
||||
emptyPrefix = &gofakes3.Prefix{}
|
||||
timeFormat = http.TimeFormat
|
||||
timeFormat = "Mon, 2 Jan 2006 15:04:05 GMT"
|
||||
)
|
||||
|
||||
// s3Backend implements the gofakes3.Backend interface to make an S3
|
||||
@@ -119,13 +117,9 @@ func (b *s3Backend) HeadObject(ctx context.Context, bucketName, objectName strin
|
||||
|
||||
fp := path.Join(bucketPath, objectName)
|
||||
fmeta, _ := op.GetNearestMeta(fp)
|
||||
ctx = context.WithValue(ctx, conf.MetaKey, fmeta)
|
||||
node, err := fs.Get(ctx, fp, &fs.GetArgs{})
|
||||
node, err := fs.Get(context.WithValue(ctx, conf.MetaKey, fmeta), fp, &fs.GetArgs{})
|
||||
if err != nil {
|
||||
if errs.IsObjectNotFound(err) {
|
||||
return nil, gofakes3.KeyNotFound(objectName)
|
||||
}
|
||||
return nil, err
|
||||
return nil, gofakes3.KeyNotFound(objectName)
|
||||
}
|
||||
|
||||
if node.IsDir() {
|
||||
@@ -133,27 +127,23 @@ func (b *s3Backend) HeadObject(ctx context.Context, bucketName, objectName strin
|
||||
}
|
||||
|
||||
size := node.GetSize()
|
||||
hash, err := getObjectHash(ctx, fp, node)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// hash := getFileHashByte(fobj)
|
||||
|
||||
meta := map[string]string{
|
||||
"Last-Modified": node.ModTime().UTC().Format(timeFormat),
|
||||
"Last-Modified": node.ModTime().Format(timeFormat),
|
||||
"Content-Type": utils.GetMimeType(fp),
|
||||
}
|
||||
|
||||
stored, etag := b.loadMetadata(fp, hash)
|
||||
for k, v := range stored {
|
||||
if k != "Last-Modified" {
|
||||
if val, ok := b.meta.Load(fp); ok {
|
||||
metaMap := val.(map[string]string)
|
||||
for k, v := range metaMap {
|
||||
meta[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
return &gofakes3.Object{
|
||||
Name: objectName,
|
||||
Hash: hash,
|
||||
ETag: etag,
|
||||
Name: objectName,
|
||||
// Hash: hash,
|
||||
Metadata: meta,
|
||||
Size: size,
|
||||
Contents: noOpReadCloser{},
|
||||
@@ -162,6 +152,8 @@ func (b *s3Backend) HeadObject(ctx context.Context, bucketName, objectName strin
|
||||
|
||||
// GetObject fetchs the object from the filesystem.
|
||||
func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string, rangeRequest *gofakes3.ObjectRangeRequest) (s3Obj *gofakes3.Object, err error) {
|
||||
defer func() { err = mapBackendError(err) }()
|
||||
|
||||
bucket, err := getBucketByName(bucketName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -170,24 +162,15 @@ func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string
|
||||
|
||||
fp := path.Join(bucketPath, objectName)
|
||||
fmeta, _ := op.GetNearestMeta(fp)
|
||||
ctx = context.WithValue(ctx, conf.MetaKey, fmeta)
|
||||
node, err := fs.Get(ctx, fp, &fs.GetArgs{})
|
||||
node, err := fs.Get(context.WithValue(ctx, conf.MetaKey, fmeta), fp, &fs.GetArgs{})
|
||||
if err != nil {
|
||||
if errs.IsObjectNotFound(err) {
|
||||
return nil, gofakes3.KeyNotFound(objectName)
|
||||
}
|
||||
return nil, err
|
||||
return nil, gofakes3.KeyNotFound(objectName)
|
||||
}
|
||||
|
||||
if node.IsDir() {
|
||||
return nil, gofakes3.KeyNotFound(objectName)
|
||||
}
|
||||
|
||||
hash, err := getObjectHash(ctx, fp, node)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
link, file, err := fs.Link(ctx, fp, model.LinkArgs{})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -223,34 +206,27 @@ func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string
|
||||
}
|
||||
|
||||
meta := map[string]string{
|
||||
"Last-Modified": node.ModTime().UTC().Format(timeFormat),
|
||||
"Last-Modified": node.ModTime().Format(timeFormat),
|
||||
"Content-Disposition": utils.GenerateContentDisposition(file.GetName()),
|
||||
"Content-Type": utils.GetMimeType(fp),
|
||||
}
|
||||
|
||||
stored, etag := b.loadMetadata(fp, hash)
|
||||
for k, v := range stored {
|
||||
if k != "Last-Modified" {
|
||||
if val, ok := b.meta.Load(fp); ok {
|
||||
metaMap := val.(map[string]string)
|
||||
for k, v := range metaMap {
|
||||
meta[k] = v
|
||||
}
|
||||
}
|
||||
closers := utils.NewClosers(rd, link)
|
||||
|
||||
return &gofakes3.Object{
|
||||
// Name: gofakes3.URLEncode(objectName),
|
||||
Name: objectName,
|
||||
Hash: hash,
|
||||
ETag: etag,
|
||||
Name: objectName,
|
||||
// Hash: "",
|
||||
Metadata: meta,
|
||||
Size: size,
|
||||
Range: rnge,
|
||||
Contents: utils.NewReadCloser(rd, func() error {
|
||||
readErr := rd.Close()
|
||||
linkErr := link.Close()
|
||||
if readErr != nil {
|
||||
return readErr
|
||||
}
|
||||
return linkErr
|
||||
}),
|
||||
Contents: utils.ReadCloser{Reader: rd, Closer: &closers},
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -266,7 +242,7 @@ func (b *s3Backend) PutObject(
|
||||
meta map[string]string,
|
||||
input io.Reader, size int64,
|
||||
) (result gofakes3.PutObjectResult, err error) {
|
||||
return result, b.putStream(ctx, bucketName, objectName, meta, input, size, "")
|
||||
return result, b.putStream(ctx, bucketName, objectName, meta, input, size)
|
||||
}
|
||||
|
||||
// putStream stores the given object into the underlying storage. It is shared
|
||||
@@ -276,7 +252,6 @@ func (b *s3Backend) putStream(
|
||||
ctx context.Context, bucketName, objectName string,
|
||||
meta map[string]string,
|
||||
input io.Reader, size int64,
|
||||
etag string,
|
||||
) error {
|
||||
bucket, err := getBucketByName(bucketName)
|
||||
if err != nil {
|
||||
@@ -342,10 +317,9 @@ func (b *s3Backend) putStream(
|
||||
if setting.GetBool(conf.IgnoreSystemFiles) && utils.IsSystemFile(obj.Name) {
|
||||
return errs.IgnoredSystemFile
|
||||
}
|
||||
hash := md5.New()
|
||||
stream := &stream.FileStream{
|
||||
Obj: &obj,
|
||||
Reader: io.TeeReader(input, hash),
|
||||
Reader: input,
|
||||
Mimetype: meta["Content-Type"],
|
||||
}
|
||||
if stream.Mimetype == "" {
|
||||
@@ -357,7 +331,7 @@ func (b *s3Backend) putStream(
|
||||
return err
|
||||
}
|
||||
|
||||
b.meta.Store(fp, objectMetadata{headers: meta, hash: hash.Sum(nil), etag: etag})
|
||||
b.meta.Store(fp, meta)
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -403,10 +377,7 @@ func (b *s3Backend) deleteObject(ctx context.Context, bucketName, objectName str
|
||||
return err
|
||||
}
|
||||
|
||||
if err := fs.Remove(ctx, fp); err != nil {
|
||||
return err
|
||||
}
|
||||
b.meta.Delete(fp)
|
||||
fs.Remove(ctx, fp)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -450,18 +421,9 @@ func (b *s3Backend) CopyObject(ctx context.Context, srcBucket, srcKey, dstBucket
|
||||
srcFp := path.Join(srcBucketPath, srcKey)
|
||||
fmeta, _ := op.GetNearestMeta(srcFp)
|
||||
srcNode, err := fs.Get(context.WithValue(ctx, conf.MetaKey, fmeta), srcFp, &fs.GetArgs{})
|
||||
if err != nil {
|
||||
if errs.IsObjectNotFound(err) {
|
||||
return result, gofakes3.KeyNotFound(srcKey)
|
||||
}
|
||||
return result, err
|
||||
}
|
||||
|
||||
c, err := b.GetObject(ctx, srcBucket, srcKey, nil)
|
||||
if err != nil {
|
||||
if errs.IsObjectNotFound(err) {
|
||||
return result, gofakes3.KeyNotFound(srcKey)
|
||||
}
|
||||
return
|
||||
}
|
||||
defer func() {
|
||||
|
||||
@@ -1,17 +0,0 @@
|
||||
package s3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/gofakes3"
|
||||
)
|
||||
|
||||
func TestCopyObjectMissingSourceReturnsNoSuchKey(t *testing.T) {
|
||||
b, _ := setupMultipartBackend(t)
|
||||
|
||||
_, err := b.CopyObject(context.Background(), "mp", "missing.txt", "mp", "copy.txt", nil)
|
||||
if code := s3ErrorCode(err); code != gofakes3.ErrNoSuchKey {
|
||||
t.Fatalf("CopyObject() error = %v, want NoSuchKey", code)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package s3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/xml"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/errs"
|
||||
"github.com/OpenListTeam/gofakes3"
|
||||
"github.com/OpenListTeam/gofakes3/s3mem"
|
||||
)
|
||||
|
||||
func TestMapBackendErrorMapsOnlyTemporaryCapacity(t *testing.T) {
|
||||
capacity := errs.NewErr(errs.TemporaryCapacity, "callback admission timed out")
|
||||
if got := mapBackendError(capacity); got != gofakes3.ErrSlowDown {
|
||||
t.Fatalf("capacity error mapped to %v, want %v", got, gofakes3.ErrSlowDown)
|
||||
}
|
||||
|
||||
permanent := errors.New("permission denied")
|
||||
if got := mapBackendError(permanent); got != permanent {
|
||||
t.Fatalf("permanent error mapped to %v, want original error", got)
|
||||
}
|
||||
if got := mapBackendError(nil); got != nil {
|
||||
t.Fatalf("nil error mapped to %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
type capacityBackend struct {
|
||||
gofakes3.Backend
|
||||
}
|
||||
|
||||
func (b capacityBackend) GetObject(context.Context, string, string, *gofakes3.ObjectRangeRequest) (*gofakes3.Object, error) {
|
||||
return nil, mapBackendError(errs.NewErr(errs.TemporaryCapacity, "callback admission timed out"))
|
||||
}
|
||||
|
||||
func TestTemporaryCapacityProducesS3SlowDownResponse(t *testing.T) {
|
||||
memory := s3mem.New()
|
||||
if err := memory.CreateBucket(t.Context(), "bucket"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server := httptest.NewServer(gofakes3.New(capacityBackend{Backend: memory}).Server())
|
||||
defer server.Close()
|
||||
|
||||
request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, server.URL+"/bucket/object", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
response, err := server.Client().Do(request)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode != http.StatusServiceUnavailable {
|
||||
t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusServiceUnavailable)
|
||||
}
|
||||
var result gofakes3.ErrorResult
|
||||
if err := xml.NewDecoder(response.Body).Decode(&result); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if result.Code != gofakes3.ErrSlowDown || result.Message != gofakes3.ErrSlowDown.Message() {
|
||||
t.Fatalf("S3 error = %#v, want SlowDown with standard message", result)
|
||||
}
|
||||
}
|
||||
@@ -1,73 +0,0 @@
|
||||
package s3
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/md5"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/fs"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
)
|
||||
|
||||
type objectMetadata struct {
|
||||
headers map[string]string
|
||||
hash []byte
|
||||
etag string
|
||||
}
|
||||
|
||||
// Only reuse upload metadata when it still describes the current content.
|
||||
func (b *s3Backend) loadMetadata(name string, hash []byte) (map[string]string, string) {
|
||||
if value, ok := b.meta.Load(name); ok {
|
||||
metadata := value.(objectMetadata)
|
||||
if bytes.Equal(metadata.hash, hash) {
|
||||
return metadata.headers, metadata.etag
|
||||
}
|
||||
}
|
||||
return nil, ""
|
||||
}
|
||||
|
||||
// Metadata alone cannot identify content: clients can preserve both file size
|
||||
// and modification time when overwriting an object. Use the driver's MD5 when
|
||||
// available, otherwise hash the complete content, including for HEAD and ranges.
|
||||
func getObjectHash(ctx context.Context, name string, obj model.Obj) ([]byte, error) {
|
||||
if value := obj.GetHash().GetHash(utils.MD5); value != "" {
|
||||
hash, err := hex.DecodeString(value)
|
||||
if err != nil || len(hash) != md5.Size {
|
||||
return nil, fmt.Errorf("invalid object MD5: %q", value)
|
||||
}
|
||||
return hash, nil
|
||||
}
|
||||
link, file, err := fs.Link(ctx, name, model.LinkArgs{})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer link.Close()
|
||||
size := link.ContentLength
|
||||
if size <= 0 {
|
||||
size = file.GetSize()
|
||||
}
|
||||
ranges, err := stream.GetRangeReaderFromLink(size, link)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
reader, err := ranges.RangeRead(ctx, http_range.Range{Length: -1})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer reader.Close()
|
||||
hash := md5.New()
|
||||
n, err := utils.CopyWithBuffer(hash, reader)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if n != size {
|
||||
return nil, io.ErrUnexpectedEOF
|
||||
}
|
||||
return hash.Sum(nil), nil
|
||||
}
|
||||
@@ -1,64 +0,0 @@
|
||||
package s3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/md5"
|
||||
"encoding/hex"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
)
|
||||
|
||||
func TestObjectHashFromDriver(t *testing.T) {
|
||||
sum := md5.Sum([]byte("content"))
|
||||
obj := &model.Object{HashInfo: utils.NewHashInfo(utils.MD5, hex.EncodeToString(sum[:]))}
|
||||
hash, err := getObjectHash(context.Background(), "object", obj)
|
||||
if err != nil || hex.EncodeToString(hash) != hex.EncodeToString(sum[:]) {
|
||||
t.Fatalf("hash = %x, err = %v", hash, err)
|
||||
}
|
||||
for _, value := range []string{"invalid", "ab"} {
|
||||
obj.HashInfo = utils.NewHashInfo(utils.MD5, value)
|
||||
if _, err := getObjectHash(context.Background(), "object", obj); err == nil {
|
||||
t.Fatalf("accepted invalid MD5 %q", value)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMultipartMetadataMatchesCurrentContent(t *testing.T) {
|
||||
b := &s3Backend{meta: new(sync.Map)}
|
||||
oldHash := md5.Sum([]byte("old"))
|
||||
newHash := md5.Sum([]byte("new"))
|
||||
b.meta.Store("object", objectMetadata{
|
||||
headers: map[string]string{"Content-Type": "text/plain"},
|
||||
hash: oldHash[:],
|
||||
etag: `"multipart-2"`,
|
||||
})
|
||||
meta, etag := b.loadMetadata("object", oldHash[:])
|
||||
if etag != `"multipart-2"` || meta["Content-Type"] != "text/plain" {
|
||||
t.Fatal("lost metadata for unchanged multipart content")
|
||||
}
|
||||
meta, etag = b.loadMetadata("object", newHash[:])
|
||||
if meta != nil || etag != "" {
|
||||
t.Fatal("reused a multipart validator after a same-size content change")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConditionalRequestsDoNotRedirect(t *testing.T) {
|
||||
for _, method := range []string{"GET", "PUT"} {
|
||||
for _, name := range []string{"If-Match", "If-None-Match", "If-Modified-Since", "If-Unmodified-Since", "If-Range"} {
|
||||
for _, value := range []string{"*", ""} {
|
||||
r := httptest.NewRequest(method, "/bucket/object", nil)
|
||||
r.Header.Set(name, value)
|
||||
if url, ok := directObjectURL(r, nil); ok || url != "" {
|
||||
t.Fatalf("redirected %s with %s", method, name)
|
||||
}
|
||||
if url, ok := directUploadURL(r, nil); ok || url != "" {
|
||||
t.Fatalf("redirected %s with %s", method, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
package s3
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/local"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/db"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/driver"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/op"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
const closeTrackingDriverName = "S3CloseTrackingLocal"
|
||||
|
||||
type closeTrackingDriver struct {
|
||||
local.Local
|
||||
closed *[]string
|
||||
}
|
||||
|
||||
func (d *closeTrackingDriver) Config() driver.Config {
|
||||
c := d.Local.Config()
|
||||
c.Name = closeTrackingDriverName
|
||||
return c
|
||||
}
|
||||
|
||||
func (d *closeTrackingDriver) Link(context.Context, model.Obj, model.LinkArgs) (*model.Link, error) {
|
||||
link := &model.Link{
|
||||
ContentLength: 4,
|
||||
RangeReader: stream.RangeReaderFunc(func(context.Context, http_range.Range) (io.ReadCloser, error) {
|
||||
return utils.NewReadCloser(strings.NewReader("body"), func() error {
|
||||
*d.closed = append(*d.closed, "body")
|
||||
return nil
|
||||
}), nil
|
||||
}),
|
||||
RequireReference: true,
|
||||
}
|
||||
link.SyncClosers.Add(utils.CloseFunc(func() error {
|
||||
*d.closed = append(*d.closed, "link")
|
||||
return nil
|
||||
}))
|
||||
return link, nil
|
||||
}
|
||||
|
||||
func TestGetObjectClosesRangeBodyBeforeLink(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
var closed []string
|
||||
op.RegisterDriver(func() driver.Driver {
|
||||
return &closeTrackingDriver{closed: &closed}
|
||||
})
|
||||
|
||||
root := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(root, "fixture.txt"), []byte("body"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
addition, err := json.Marshal(struct {
|
||||
RootFolderPath string `json:"root_folder_path"`
|
||||
}{RootFolderPath: root})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
mount := "/" + sanitizeTestName(t.Name())
|
||||
storageID, err := op.CreateStorage(ctx, model.Storage{
|
||||
Driver: closeTrackingDriverName,
|
||||
MountPath: mount,
|
||||
Addition: string(addition),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
if err := op.DeleteStorageById(ctx, storageID); err != nil {
|
||||
t.Errorf("delete fixture storage: %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
previousBuckets, previousBucketsErr := op.GetSettingItemByKey(conf.S3Buckets)
|
||||
if previousBucketsErr != nil && !errors.Is(previousBucketsErr, gorm.ErrRecordNotFound) {
|
||||
t.Fatal(previousBucketsErr)
|
||||
}
|
||||
if err := op.SaveSettingItem(&model.SettingItem{
|
||||
Key: conf.S3Buckets,
|
||||
Value: `[{"name":"close","path":"` + mount + `"}]`,
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
if previousBucketsErr == nil {
|
||||
if err := op.SaveSettingItem(previousBuckets); err != nil {
|
||||
t.Errorf("restore S3 buckets: %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err := db.DeleteSettingItemByKey(conf.S3Buckets); err != nil {
|
||||
t.Errorf("delete fixture S3 buckets: %v", err)
|
||||
}
|
||||
op.SettingCacheUpdate()
|
||||
})
|
||||
|
||||
object, err := newBackend().(*s3Backend).GetObject(ctx, "close", "fixture.txt", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
contents, err := io.ReadAll(object.Contents)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if string(contents) != "body" {
|
||||
t.Fatalf("contents = %q, want body", contents)
|
||||
}
|
||||
if err := object.Contents.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := object.Contents.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !slices.Equal(closed, []string{"body", "link"}) {
|
||||
t.Fatalf("close order = %v, want [body link] exactly once", closed)
|
||||
}
|
||||
}
|
||||
@@ -234,9 +234,7 @@ func (b *s3Backend) CompleteMultipartUpload(ctx context.Context, bucket, object
|
||||
|
||||
defer combined.Close()
|
||||
|
||||
sum := md5.Sum(concat)
|
||||
etag := fmt.Sprintf("%q", fmt.Sprintf("%s-%d", hex.EncodeToString(sum[:]), len(ordered)))
|
||||
err := b.putStream(ctx, bucket, object, state.meta, combined, total, etag)
|
||||
err := b.putStream(ctx, bucket, object, state.meta, combined, total)
|
||||
if err != nil {
|
||||
// Leave the upload in place so the client may retry completion, per
|
||||
// the gofakes3 MultipartBackend contract.
|
||||
@@ -246,6 +244,8 @@ func (b *s3Backend) CompleteMultipartUpload(ctx context.Context, bucket, object
|
||||
// Success: drop bookkeeping and clean up part files.
|
||||
b.removeUpload(uploadID)
|
||||
|
||||
sum := md5.Sum(concat)
|
||||
etag := fmt.Sprintf("%q", fmt.Sprintf("%s-%d", hex.EncodeToString(sum[:]), len(ordered)))
|
||||
log.Debugf("s3 multipart: completed upload %s -> %s/%s (%d bytes)", uploadID, bucket, object, total)
|
||||
return "", etag, nil
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package s3
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
@@ -67,11 +68,21 @@ func setupMultipartBackend(t *testing.T) (*s3Backend, string) {
|
||||
t.Fatalf("mkdir local root: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = os.RemoveAll(localRoot) })
|
||||
addition, err := json.Marshal(struct {
|
||||
RootFolderPath string `json:"root_folder_path"`
|
||||
Thumbnail bool `json:"thumbnail"`
|
||||
}{
|
||||
RootFolderPath: localRoot,
|
||||
Thumbnail: false,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("marshal local storage addition: %v", err)
|
||||
}
|
||||
|
||||
_, err = op.CreateStorage(ctx, model.Storage{
|
||||
Driver: "Local",
|
||||
MountPath: mount,
|
||||
Addition: `{"root_folder_path":"` + localRoot + `","thumbnail":false}`,
|
||||
Addition: string(addition),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("create local storage: %+v", err)
|
||||
|
||||
@@ -1,12 +0,0 @@
|
||||
package s3
|
||||
|
||||
import "net/http"
|
||||
|
||||
func hasPreconditions(r *http.Request) bool {
|
||||
for _, name := range []string{"If-Match", "If-None-Match", "If-Modified-Since", "If-Unmodified-Since", "If-Range"} {
|
||||
if _, ok := r.Header[name]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -38,7 +38,7 @@ func redirectHandler(next http.Handler, authPairs map[string]string) http.Handle
|
||||
}
|
||||
|
||||
func directObjectURL(r *http.Request, authPairs map[string]string) (string, bool) {
|
||||
if r.Method != http.MethodGet || hasPreconditions(r) {
|
||||
if r.Method != http.MethodGet {
|
||||
return "", false
|
||||
}
|
||||
if hasNonObjectQuery(r) || !s3RequestAuthorized(r, authPairs) {
|
||||
@@ -75,7 +75,7 @@ func directObjectURL(r *http.Request, authPairs map[string]string) (string, bool
|
||||
}
|
||||
|
||||
func directUploadURL(r *http.Request, authPairs map[string]string) (string, bool) {
|
||||
if r.Method != http.MethodPut || r.ContentLength < 0 || hasPreconditions(r) {
|
||||
if r.Method != http.MethodPut || r.ContentLength < 0 {
|
||||
return "", false
|
||||
}
|
||||
if hasNonObjectQuery(r) || !s3RequestAuthorized(r, authPairs) {
|
||||
|
||||
@@ -5,6 +5,7 @@ package s3
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"strings"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
@@ -21,6 +22,13 @@ type Bucket struct {
|
||||
Path string `json:"path"`
|
||||
}
|
||||
|
||||
func mapBackendError(err error) error {
|
||||
if stderrors.Is(err, errs.TemporaryCapacity) {
|
||||
return gofakes3.ErrSlowDown
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
const emptyObjectName = "ThisIsAnEmptyFolderInTheS3Bucket"
|
||||
|
||||
func getAndParseBuckets() ([]Bucket, error) {
|
||||
|
||||
Reference in New Issue
Block a user