Compare commits

..

1 Commits

Author SHA1 Message Date
renovate[bot] 579ceff18f fix(deps): update module github.com/go-webauthn/webauthn to v0.18.2 2026-09-19 16:39:32 +00:00
71 changed files with 633 additions and 4114 deletions
+2 -2
View File
@@ -124,7 +124,7 @@ jobs:
FRONTEND_REPO: ${{ vars.FRONTEND_REPO }}
- name: Build
uses: OpenListTeam/cgo-actions@d760a8ec1a6be1f8ec181229e11cb2671797f9e0 # v1.3.0
uses: OpenListTeam/cgo-actions@6fcace5934c36d70503dba06e2396bb58b766130 # v1.2.5
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@oplist.org>
github.com/OpenListTeam/OpenList/v4/internal/conf.GitAuthor=The OpenList Projects Contributors <noreply@openlist.team>
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
+2 -2
View File
@@ -42,7 +42,7 @@ jobs:
FRONTEND_REPO: ${{ vars.FRONTEND_REPO }}
- name: Build
uses: OpenListTeam/cgo-actions@d760a8ec1a6be1f8ec181229e11cb2671797f9e0 # v1.3.0
uses: OpenListTeam/cgo-actions@6fcace5934c36d70503dba06e2396bb58b766130 # v1.2.5
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@oplist.org>
github.com/OpenListTeam/OpenList/v4/internal/conf.GitAuthor=The OpenList Projects Contributors <noreply@openlist.team>
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
+3 -3
View File
@@ -1,7 +1,7 @@
set -e
appName="openlist"
builtAt="$(date +'%F %T %z')"
gitAuthor="The OpenList Projects Contributors <noreply@oplist.org>"
gitAuthor="The OpenList Projects Contributors <noreply@openlist.team>"
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.4"
freebsd_version="14.4"
echo "Failed to get FreeBSD version, falling back to 14.3"
freebsd_version="14.3"
fi
echo "Using FreeBSD version: $freebsd_version"
+16 -61
View File
@@ -3,7 +3,6 @@ package _115_open
import (
"context"
"encoding/base64"
"errors"
"io"
"time"
@@ -71,19 +70,6 @@ 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 {
@@ -94,32 +80,7 @@ func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer,
return err
}
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
})
imur, err := bucket.InitiateMultipartUpload(initResp.Object, oss.Sequential())
if err != nil {
return err
}
@@ -148,17 +109,13 @@ func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer,
return err
}
err = retry.Do(func() error {
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
})
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
},
retry.Context(ctx),
retry.Attempts(3),
@@ -177,16 +134,14 @@ func (d *Open115) multpartUpload(ctx context.Context, stream model.FileStreamer,
up(float64(offset) * 100 / float64(fileSize))
}
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
})
// 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),
)
if err != nil {
return err
}
+2 -7
View File
@@ -7,16 +7,11 @@ import (
type Addition struct {
LoginType string `json:"login_type" type:"select" options:"password,qrcode" default:"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"`
Username string `json:"username" required:"true"`
Password string `json:"password" required:"true"`
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"`
-62
View File
@@ -72,8 +72,6 @@ type BaseLoginParam struct {
// 请求头参数
Lt string
ReqId string
// logbox页面地址,作为后续请求的Referer,缺失会被判定为陌生设备
Referer string
// 表单参数
ParamId string
@@ -99,20 +97,10 @@ 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"`
@@ -128,35 +116,6 @@ 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返回
@@ -190,27 +149,6 @@ 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"`
+77 -355
View File
@@ -30,7 +30,6 @@ 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"
@@ -42,13 +41,9 @@ import (
const (
ACCOUNT_TYPE = "02"
// 官方 PC 端(cloud.189.cn 网页/客户端)使用的 appId,
// 登录、生成二维码、换取 session 必须全程使用同一个 appId
APP_ID = "9317140619"
CLIENT_TYPE = "10020"
// 扫码状态轮询使用的 clientType,与密码登录的 10020 不同
QR_CLIENT_TYPE = "1"
VERSION = "7.2.4.0"
APP_ID = "8025431004"
CLIENT_TYPE = "10020"
VERSION = "6.2"
WEB_URL = "https://cloud.189.cn"
AUTH_URL = "https://open.e.189.cn"
@@ -62,18 +57,8 @@ 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 {
@@ -303,72 +288,9 @@ 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 {
@@ -377,16 +299,9 @@ 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
// 遇到错误,重新加载登陆参数(刷新验证码)
@@ -404,18 +319,17 @@ func (y *Cloud189PC) loginByPassword() (err error) {
param := y.loginParam
var loginresp LoginResp
res, err := y.client.R().
_, err = y.client.R().
ForceContentType("application/json;charset=UTF-8").SetResult(&loginresp).
SetHeaders(y.loginHeaders(param.BaseLoginParam)).
SetHeaders(map[string]string{
"REQID": param.ReqId,
"lt": param.Lt,
}).
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,
@@ -431,106 +345,17 @@ 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)
}
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 {
// 获取Session
var erron RespErr
var tokenInfo AppSessionResp
_, err := y.client.R().
_, err = y.client.R().
SetResult(&tokenInfo).SetError(&erron).
SetQueryParams(clientSuffix()).
SetQueryParams(y.deviceParams()).
SetQueryParam("redirectURL", redirectURL).
SetHeader("X-Request-ID", uuid.NewString()).
SetQueryParam("redirectURL", loginresp.ToUrl).
Post(API_URL + "/getSessionForPC.action")
if err != nil {
return err
@@ -540,13 +365,14 @@ func (y *Cloud189PC) getSessionByRedirectURL(redirectURL string) error {
return &erron
}
if tokenInfo.ResCode != 0 {
return errors.New(tokenInfo.ResMessage)
err = fmt.Errorf(tokenInfo.ResMessage)
return err
}
y.Addition.AccessToken = tokenInfo.AccessToken
y.Addition.RefreshToken = tokenInfo.RefreshToken
y.tokenInfo = &tokenInfo
op.MustSaveDriverStorage(y)
return nil
return err
}
func (y *Cloud189PC) loginByQRCode() error {
@@ -557,74 +383,66 @@ func (y *Cloud189PC) loginByQRCode() error {
}
}
// 本地轮询,扫码确认后自动继续,不需要用户反复保存
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)
var state struct {
Status int `json:"status"`
RedirectUrl string `json:"redirectUrl"`
Msg string `json:"msg"`
}
// 轮询超时,把二维码交回前端等待下一次保存
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(y.loginHeaders(y.qrcodeParam.BaseLoginParam)).
SetHeaders(map[string]string{
"Referer": AUTH_URL,
"Reqid": y.qrcodeParam.ReqId,
"lt": y.qrcodeParam.Lt,
}).
SetFormData(map[string]string{
"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),
"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),
}).
ForceContentType("application/json;charset=UTF-8").
SetResult(&state).
Post(AUTH_URL + "/api/logbox/oauth2/qrcodeLoginState.do")
if err != nil {
return nil, err
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 &state, nil
}
func (y *Cloud189PC) genQRCode(text string) error {
@@ -650,9 +468,8 @@ 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().
@@ -667,98 +484,14 @@ 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{
Lt: lt,
ReqId: reqId,
Referer: finalUrl.String(),
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],
}, 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
}
/* 初始化登陆需要的参数
* 如果遇到验证码返回错误
*/
@@ -783,13 +516,12 @@ func (y *Cloud189PC) initLoginParam() error {
}
y.loginParam.jRsaKey = fmt.Sprintf("-----BEGIN PUBLIC KEY-----\n%s\n-----END PUBLIC KEY-----", encryptConf.Data.PubKey)
y.loginParam.rsaPrefix = encryptConf.Data.Pre
y.loginParam.RsaUsername = y.loginParam.encryptSecret(y.Username)
y.loginParam.RsaPassword = y.loginParam.encryptSecret(y.Password)
y.loginParam.RsaUsername = encryptConf.Data.Pre + RsaEncrypt(y.loginParam.jRsaKey, y.Username)
y.loginParam.RsaPassword = encryptConf.Data.Pre + RsaEncrypt(y.loginParam.jRsaKey, y.Password)
// 判断是否需要验证码
resp, err := y.client.R().
SetHeaders(y.loginHeaders(y.loginParam.BaseLoginParam)).
SetHeader("REQID", y.loginParam.ReqId).
SetFormData(map[string]string{
"appKey": APP_ID,
"accountType": ACCOUNT_TYPE,
@@ -844,7 +576,6 @@ 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).
@@ -852,9 +583,6 @@ 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
@@ -875,7 +603,6 @@ 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,
@@ -916,11 +643,12 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
return errors.New("refresh token failed after maximum retries")
}
// 该接口刷新失败时以HTTP 200返回 result/msg,SetError不会触发,必须解析响应体判断
var tokenInfo RefreshTokenResp
var erron RespErr
var tokenInfo AppSessionResp
_, 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,
@@ -933,8 +661,7 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
}
// 如果刷新失败,返回错误给上层处理
if tokenInfo.HasError() {
refreshErr := tokenInfo.Error()
if erron.HasError() {
if y.Addition.RefreshToken != "" {
y.Addition.RefreshToken = ""
op.MustSaveDriverStorage(y)
@@ -942,11 +669,7 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
// 根据登录类型决定下一步行为
if y.LoginType == "qrcode" {
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 errors.New("QR code session has expired, please re-scan the code to log in")
}
// 密码登录模式下,尝试回退到完整登录
return y.login()
@@ -954,8 +677,7 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
y.Addition.AccessToken = tokenInfo.AccessToken
y.Addition.RefreshToken = tokenInfo.RefreshToken
y.tokenInfo.AccessToken = tokenInfo.AccessToken
y.tokenInfo.RefreshToken = tokenInfo.RefreshToken
y.tokenInfo = &tokenInfo
op.MustSaveDriverStorage(y)
return y.refreshSessionWithRetry(retryCount + 1)
}
+1
View File
@@ -328,6 +328,7 @@ 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
}
-295
View File
@@ -1,295 +0,0 @@
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
}
-365
View File
@@ -1,365 +0,0 @@
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())
}
}
+4 -18
View File
@@ -11,7 +11,6 @@ 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"
@@ -23,9 +22,8 @@ type AliyundriveOpen struct {
DriveId string
limiter *limiter
ref *AliyundriveOpen
callback *callbackRegistration
limiter *limiter
ref *AliyundriveOpen
}
func (d *AliyundriveOpen) Config() driver.Config {
@@ -37,7 +35,6 @@ 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"
@@ -55,7 +52,6 @@ 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
}
@@ -69,10 +65,6 @@ 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
@@ -127,16 +119,10 @@ 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
link := &model.Link{
return &model.Link{
URL: url,
Expiration: &exp,
}
if args.Redirect {
return link, nil
}
link.URL = ""
link.RangeReader = stream.RateLimitRangeReaderFunc(d.callbackRangeReader(url, file.GetSize()))
return link, nil
}, nil
}
func (d *AliyundriveOpen) MakeDir(ctx context.Context, parentDir model.Obj, dirName string) (model.Obj, error) {
+13 -14
View File
@@ -8,20 +8,19 @@ 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"`
CallbackConcurrency int `json:"callback_concurrency" type:"number" default:"1" help:"Maximum active proxied downloads shared by Aliyun user"`
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"`
AccessToken string
}
var config = driver.Config{
+8 -43
View File
@@ -315,36 +315,6 @@ 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
@@ -467,23 +437,18 @@ 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 {
// 旧版结构:相对路径 /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 {
if len(matches) < 2 {
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
}
+1 -2
View File
@@ -10,11 +10,10 @@ 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:"false" default:""`
RefreshToken string `json:"refresh_token" required:"true" 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{
+10 -30
View File
@@ -100,13 +100,12 @@ 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"
// Always refresh captcha token before signin (it may be expired)
if err := d.RefreshCaptchaTokenInLogin(GetAction(http.MethodPost, url), d.Username); err != nil {
return err
// 使用 用户填写的 CaptchaToken —————— (验证后的captcha_token)
if d.GetCaptchaToken() == "" {
if err := d.RefreshCaptchaTokenInLogin(GetAction(http.MethodPost, url), d.Username); err != nil {
return err
}
}
var e ErrResp
@@ -126,12 +125,7 @@ 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
}
@@ -165,14 +159,9 @@ 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 = newRefreshToken
d.AccessToken = newAccessToken
d.RefreshToken = jsoniter.Get(data, "refresh_token").ToString()
d.AccessToken = jsoniter.Get(data, "access_token").ToString()
d.Common.SetUserID(jsoniter.Get(data, "sub").ToString())
d.Addition.RefreshToken = d.RefreshToken
op.MustSaveDriverStorage(d)
@@ -208,18 +197,12 @@ func (d *PikPak) request(url string, method string, callback base.ReqCallback, r
case 0:
return res.Body(), nil
case 4122, 4121, 16:
if strings.Contains(url, "/v1/auth/") || strings.Contains(url, "/v1/shield/captcha/") {
return nil, errors.New(e.Error())
}
// access_token expired, refresh and retry
// access_token 过期
if err1 := d.refreshToken(d.RefreshToken); err1 != nil {
return nil, err1
}
return d.request(url, method, callback, resp)
case 9: // captcha token expired
if strings.Contains(url, "/v1/shield/captcha/") {
return nil, errors.New(e.Error())
}
case 9: // 验证码token过期
if err = d.RefreshCaptchaTokenAtLogin(GetAction(method, url), d.GetUserID()); err != nil {
return nil, err
}
@@ -386,9 +369,6 @@ 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)
}
@@ -427,7 +407,7 @@ func (d *PikPak) refreshCaptchaToken(action string, metas map[string]string) err
return errors.New(e.Error())
}
if resp.Url != "" && !d.Addition.SkipVerification {
if resp.Url != "" {
return fmt.Errorf(`need verify: <a target="_blank" href="%s">Click Here</a>`, resp.Url)
}
-720
View File
@@ -1,720 +0,0 @@
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")
}
}
+3 -11
View File
@@ -55,18 +55,10 @@ 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": dirPath,
"path": dir.GetPath(),
"limit": "500",
"page": "1",
})
@@ -95,7 +87,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": dirPath,
"path": dir.GetPath(),
"limit": "500",
"page": strconv.Itoa(page),
})
@@ -122,7 +114,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(dirPath, src.Name),
Path: path.Join(dir.GetPath(), src.Name),
ID: src.ID,
Name: src.Name,
Size: func() int64 {
-43
View File
@@ -4,11 +4,9 @@ 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"
)
@@ -38,44 +36,3 @@ 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)
}
}
+11 -11
View File
@@ -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.20260911142347-cd3c030a83b4
github.com/OpenListTeam/gofakes3 v0.8.1
github.com/OpenListTeam/sftpd-openlist v1.0.1
github.com/OpenListTeam/tache v0.2.2
github.com/OpenListTeam/times v0.1.0
@@ -26,7 +26,7 @@ require (
github.com/blevesearch/bleve/v2 v2.6.1
github.com/bmatcuk/doublestar/v4 v4.10.0
github.com/caarlos0/env/v9 v9.0.0
github.com/charmbracelet/bubbles/v2 v2.2.1
github.com/charmbracelet/bubbles v0.21.1
github.com/charmbracelet/bubbletea v1.3.10
github.com/charmbracelet/lipgloss v1.1.0
github.com/city404/v6-public-rpc-proto/go v0.0.0-20240817070657-90f8e24b653e
@@ -44,7 +44,7 @@ require (
github.com/gin-gonic/gin v1.12.0
github.com/glebarez/sqlite v1.11.0
github.com/go-resty/resty/v2 v2.17.2
github.com/go-webauthn/webauthn v0.18.0
github.com/go-webauthn/webauthn v0.18.2
github.com/golang-jwt/jwt/v4 v4.5.2
github.com/google/uuid v1.6.0
github.com/gorilla/websocket v1.5.3
@@ -75,7 +75,7 @@ require (
github.com/u2takey/ffmpeg-go v0.5.0
github.com/upyun/go-sdk/v3 v3.0.4
github.com/zzzhr1990/go-common-entity v0.0.0-20250202070650-1a200048f0d3
golang.org/x/crypto v0.56.0
golang.org/x/crypto v0.57.0
golang.org/x/image v0.45.0
golang.org/x/net v0.58.0
golang.org/x/oauth2 v0.36.0
@@ -222,7 +222,7 @@ require (
github.com/crackcomm/go-gitignore v0.0.0-20170627025303-887ab5e44cc3 // indirect
github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect
github.com/decred/dcrd/dcrec/secp256k1/v4 v4.1.0 // indirect
github.com/fxamacker/cbor/v2 v2.9.3 // indirect
github.com/fxamacker/cbor/v2 v2.9.4 // indirect
github.com/gabriel-vasile/mimetype v1.4.13 // indirect
github.com/gin-contrib/sse v1.1.0 // indirect
github.com/go-chi/chi/v5 v5.3.1 // indirect
@@ -231,7 +231,7 @@ require (
github.com/go-playground/universal-translator v0.18.1 // indirect
github.com/go-playground/validator/v10 v10.30.3 // indirect
github.com/go-sql-driver/mysql v1.8.1 // indirect
github.com/go-webauthn/x v0.3.0 // indirect
github.com/go-webauthn/x v0.3.1 // indirect
github.com/goccy/go-json v0.10.6 // indirect
github.com/golang-jwt/jwt/v5 v5.3.1 // indirect
github.com/golang/protobuf v1.5.4 // indirect
@@ -309,11 +309,11 @@ require (
github.com/yusufpapurcu/wmi v1.2.4 // indirect
go.etcd.io/bbolt v1.5.0 // indirect
golang.org/x/arch v0.23.0 // indirect
golang.org/x/sync v0.22.0
golang.org/x/sys v0.47.0
golang.org/x/term v0.45.0 // indirect
golang.org/x/text v0.41.0
golang.org/x/tools v0.48.0 // indirect
golang.org/x/sync v0.23.0
golang.org/x/sys v0.48.0
golang.org/x/term v0.46.0 // indirect
golang.org/x/text v0.42.0
golang.org/x/tools v0.49.0 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260715232425-e75dac1f907d // indirect
google.golang.org/grpc v1.85.0-dev
google.golang.org/protobuf v1.36.11 // indirect
+20 -2
View File
@@ -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.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/gofakes3 v0.8.1 h1:uihJ7Zgb4qIafFcXhcm71BzxCyGRIqBVJYg4YOUa6uY=
github.com/OpenListTeam/gofakes3 v0.8.1/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=
@@ -375,6 +375,8 @@ github.com/fxamacker/cbor/v2 v2.9.0 h1:NpKPmjDBgUfBms6tr6JZkTHtfFGcMKsw3eGcmD/sa
github.com/fxamacker/cbor/v2 v2.9.0/go.mod h1:vM4b+DJCtHn+zz7h3FFp/hDAI9WNWCsZj23V5ytsSxQ=
github.com/fxamacker/cbor/v2 v2.9.3 h1:oQBnFATpNdY8gJHTndDDv5Xl4QqNaz51G5LLEPhng3Q=
github.com/fxamacker/cbor/v2 v2.9.3/go.mod h1:vM4b+DJCtHn+zz7h3FFp/hDAI9WNWCsZj23V5ytsSxQ=
github.com/fxamacker/cbor/v2 v2.9.4 h1:xwjVlxEMR3S605oUlgBjKLTTeGFciYPGYCtF/35LKGo=
github.com/fxamacker/cbor/v2 v2.9.4/go.mod h1:vM4b+DJCtHn+zz7h3FFp/hDAI9WNWCsZj23V5ytsSxQ=
github.com/gabriel-vasile/mimetype v1.4.12 h1:e9hWvmLYvtp846tLHam2o++qitpguFiYCKbn0w9jyqw=
github.com/gabriel-vasile/mimetype v1.4.12/go.mod h1:d+9Oxyo1wTzWdyVUPMmXFvp4F9tea18J8ufA774AB3s=
github.com/gabriel-vasile/mimetype v1.4.13 h1:46nXokslUBsAJE/wMsp5gtO500a4F3Nkz9Ufpk2AcUM=
@@ -433,10 +435,14 @@ github.com/go-webauthn/webauthn v0.13.4 h1:q68qusWPcqHbg9STSxBLBHnsKaLxNO0RnVKaA
github.com/go-webauthn/webauthn v0.13.4/go.mod h1:MglN6OH9ECxvhDqoq1wMoF6P6JRYDiQpC9nc5OomQmI=
github.com/go-webauthn/webauthn v0.18.0 h1:PC8R3PNLEmjZf++WwcQlo1Z39S9rf8ma69rlwkypZhA=
github.com/go-webauthn/webauthn v0.18.0/go.mod h1:ymzZQhx3D/PrDjznemBdQJ23gHTaSDxUchM7sH1lUCg=
github.com/go-webauthn/webauthn v0.18.2 h1:0BeftmEHU7i3Dv0VFwBtidy/ba37Vcdjvqst9EYu8Sk=
github.com/go-webauthn/webauthn v0.18.2/go.mod h1:hEXaOuLxvZ3zG9miZe3ehlyeVso9AtklXG+kTn36k+A=
github.com/go-webauthn/x v0.1.23 h1:9lEO0s+g8iTyz5Vszlg/rXTGrx3CjcD0RZQ1GPZCaxI=
github.com/go-webauthn/x v0.1.23/go.mod h1:AJd3hI7NfEp/4fI6T4CHD753u91l510lglU7/NMN6+E=
github.com/go-webauthn/x v0.3.0 h1:Q2X9vbrlP0Ed+QGEzixh1hthGZlDnzVT0XH/9IIQ0kE=
github.com/go-webauthn/x v0.3.0/go.mod h1:5OkdSQdOy7taRXWqvNHggtaPffmW94ybu3rZEER4I+I=
github.com/go-webauthn/x v0.3.1 h1:1ff37z3XfmTTomkhlURgGizLIDyOvPgTt2t9nlzKLRo=
github.com/go-webauthn/x v0.3.1/go.mod h1:ZInxAynYXfBPvvm5gzKZ7geBlL23K71xASMgohHl/Rg=
github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4=
github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M=
github.com/goccy/go-json v0.10.6 h1:p8HrPJzOakx/mn/bQtjgNjdTcN+/S6FcG2CTtQOrHVU=
@@ -950,6 +956,8 @@ golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
golang.org/x/crypto v0.56.0 h1:GUh5Ii4J5jtcseSMiRqr1jXCNHoxjeV9Fmekc2oLy6Y=
golang.org/x/crypto v0.56.0/go.mod h1:OMW5y6CY9l38uPLmxU6l6pwcXp1obtLo3e6gT7gQR2I=
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
golang.org/x/exp v0.0.0-20250606033433-dcc06ee1d476 h1:bsqhLWFR6G6xiQcb+JoGqdKdRU6WzPWmK8E0jxTjzo4=
golang.org/x/exp v0.0.0-20250606033433-dcc06ee1d476/go.mod h1:3//PLf8L/X+8b4vuAfHzxeRUl04Adcb341+IGKfnqS8=
golang.org/x/exp v0.0.0-20260410095643-746e56fc9e2f h1:W3F4c+6OLc6H2lb//N1q4WpJkhzJCK5J6kUi1NTVXfM=
@@ -1001,6 +1009,8 @@ golang.org/x/sync v0.7.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.10.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -1026,6 +1036,8 @@ golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
golang.org/x/telemetry v0.0.0-20240228155512-f48c80bd79b2/go.mod h1:TeRTkGYfJXctD9OcfyVLyj2J3IxLnKwHJR8f4D8a3YE=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/term v0.0.0-20210927222741-03fcf44c2211/go.mod h1:jbD1KX2456YbFQfuXm/mYQcufACuNUgVhRMnK/tPxf8=
@@ -1040,6 +1052,8 @@ golang.org/x/term v0.44.0 h1:0rLvDRCtNj0gZkyIXhCyOb2OAzEhLVqc4B+hrsBhrmc=
golang.org/x/term v0.44.0/go.mod h1:7ze4MdzUzLXpSAoFP1H0bOI9aXDqveSvatT5vKcFh2Y=
golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0=
golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w=
golang.org/x/term v0.46.0 h1:3+OXuTbaKDgwk8jTi3aSLHRlmWqHEUDUtxnbFigO4YE=
golang.org/x/term v0.46.0/go.mod h1:+K02xbkittuwc0Am4abfA3Fc+XRGXkvBXNO88NCXPoc=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.2/go.mod h1:bEr9sfX3Q8Zfm5fL9x+3itogRgK3+ptLWKqgva+5dAk=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
@@ -1056,6 +1070,8 @@ golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs=
golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY=
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI=
golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E=
golang.org/x/time v0.0.0-20190308202827-9d24e82272b4/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ=
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
@@ -1073,6 +1089,8 @@ golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q=
golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA=
golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE=
golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk=
golang.org/x/tools v0.49.0 h1:3NI7VXzL9+1WZD52Dx2ttoPwD5DWrFGpl9mFZDlmisI=
golang.org/x/tools v0.49.0/go.mod h1:SJNXV9DBKT0UbdttsQjbfJlAE/q+y36++zo3uL3N0Oo=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
+2 -1
View File
@@ -6,12 +6,13 @@ 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(conf.GetApiUrl(c.Request.Context()))
siteUrl, err := url.Parse(common.GetApiUrl(c.Request.Context()))
if err != nil {
return nil, err
}
-1
View File
@@ -211,7 +211,6 @@ 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},
-1
View File
@@ -117,7 +117,6 @@ const (
SSODefaultDir = "sso_default_dir"
SSODefaultPermission = "sso_default_permission"
SSOCompatibilityMode = "sso_compatibility_mode"
SSOPostMessageOrigin = "sso_postmessage_origin"
// ldap
LdapLoginEnabled = "ldap_login_enabled"
-8
View File
@@ -1,8 +0,0 @@
package conf
import "context"
func GetApiUrl(ctx context.Context) string {
api, _ := ctx.Value(ApiUrlKey).(string)
return api
}
-29
View File
@@ -1,29 +0,0 @@
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)
}
})
}
}
+1 -1
View File
@@ -35,7 +35,7 @@ func DeleteSearchNodesByParent(path string) error {
if err != nil {
return err
}
dir, name := stdpath.Dir(path), stdpath.Base(path)
dir, name := stdpath.Split(path)
return db.Where(fmt.Sprintf("%s = ? AND %s = ?",
columnName("parent"), columnName("name")),
dir, name).Delete(&model.SearchNode{}).Error
-1
View File
@@ -18,7 +18,6 @@ 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")
+2 -1
View File
@@ -21,6 +21,7 @@ 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"
@@ -414,7 +415,7 @@ func archiveDecompress(ctx context.Context, srcObjPath, dstDirPath string, args
return nil, err
} else {
tsk.Creator, _ = ctx.Value(conf.UserKey).(*model.User)
tsk.ApiUrl = conf.GetApiUrl(ctx)
tsk.ApiUrl = common.GetApiUrl(ctx)
ArchiveDownloadTaskManager.Add(tsk)
return tsk, nil
}
+2 -1
View File
@@ -14,6 +14,7 @@ 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"
)
@@ -165,7 +166,7 @@ func transfer(ctx context.Context, taskType taskType, srcObjPath, dstDirPath str
}
t.Creator, _ = ctx.Value(conf.UserKey).(*model.User)
t.ApiUrl = conf.GetApiUrl(ctx)
t.ApiUrl = common.GetApiUrl(ctx)
if taskType == copy || taskType == merge {
CopyTaskManager.Add(t)
} else {
+2 -2
View File
@@ -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 = conf.GetApiUrl(ctx) + l.URL
l.URL = common.GetApiUrl(ctx) + l.URL
}
return l, obj, nil
}
+2 -1
View File
@@ -7,6 +7,7 @@ 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"
@@ -80,7 +81,7 @@ func putAsTask(ctx context.Context, dstDirPath string, file model.FileStreamer)
t := &UploadTask{
TaskExtension: task.TaskExtension{
Creator: taskCreator,
ApiUrl: conf.GetApiUrl(ctx),
ApiUrl: common.GetApiUrl(ctx),
},
storage: storage,
dstDirActualPath: dstDirActualPath,
+20 -2
View File
@@ -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 expiration; not transferred by Clone
Expiration *time.Duration // local cache expire Duration
//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,3 +118,21 @@ 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
}
-25
View File
@@ -1,25 +0,0 @@
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")
}
}
+5 -7
View File
@@ -206,8 +206,7 @@ func (d *downloader) download() (io.ReadCloser, error) {
if err != nil {
d.cancel(err)
d.cfg.ConcurrencyLimit.Release()
_ = d.interrupt()
return nil, err
return nil, d.interrupt()
}
d.mu.Lock()
@@ -269,6 +268,10 @@ 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,
@@ -283,11 +286,6 @@ 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
}
}
-88
View File
@@ -1,88 +0,0 @@
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)
}
})
}
}
+26 -89
View File
@@ -4,7 +4,6 @@ import (
"compress/gzip"
"context"
"crypto/tls"
stderrors "errors"
"fmt"
"io"
"mime/multipart"
@@ -16,6 +15,7 @@ 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,8 +25,12 @@ 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 range reader. The main benefit of ServeHTTP over io.Copy
// provided RangeReadCloser. 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.
@@ -43,11 +47,13 @@ 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 RangeRead method must return a reader for the requested range.
// The content's RangeReadCloser method must work: ServeHTTP gives a range,
// caller will give the reader for that 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, rangeReader model.RangeReaderIF) (err error) {
func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time.Time, size int64, RangeReadCloser model.RangeReadCloserIF) error {
defer RangeReadCloser.Close()
setLastModified(w, modTime)
done, rangeReq := checkPreconditions(w, r, modTime)
if done {
@@ -107,11 +113,10 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
ctx := r.Context()
switch {
case len(ranges) == 0:
reader, err := openRange(ctx, rangeReader, http_range.Range{Length: -1})
reader, err := RangeReadCloser.RangeRead(ctx, http_range.Range{Length: -1})
if err != nil {
code = http.StatusRequestedRangeNotSatisfiable
var statusCode HttpStatusCodeError
if errors.As(err, &statusCode) {
if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok {
code = int(statusCode)
}
http.Error(w, err.Error(), code)
@@ -131,11 +136,10 @@ 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 = openRange(ctx, rangeReader, ra)
sendContent, err = RangeReadCloser.RangeRead(ctx, ra)
if err != nil {
code = http.StatusRequestedRangeNotSatisfiable
var statusCode HttpStatusCodeError
if errors.As(err, &statusCode) {
if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok {
code = int(statusCode)
}
http.Error(w, err.Error(), code)
@@ -155,6 +159,7 @@ 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))
@@ -162,18 +167,21 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
pw.CloseWithError(err)
return
}
if err := copyRange(ctx, part, rangeReader, ra); err != nil {
reader, err := RangeReadCloser.RangeRead(ctx, ra)
if err != nil {
pw.CloseWithError(err)
return
}
if _, err := utils.CopyWithBufferN(part, reader, ra.Length); err != nil {
pw.CloseWithError(err)
return
}
}
_ = pw.CloseWithError(mw.Close())
mw.Close()
pw.Close()
}()
}
defer func() {
err = closeWithError(err, sendContent)
}()
w.Header().Set("Accept-Ranges", "bytes")
if w.Header().Get("Content-Encoding") == "" {
@@ -193,8 +201,7 @@ 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
var statusCode HttpStatusCodeError
if errors.As(err, &statusCode) {
if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok {
code = int(statusCode)
}
w.WriteHeader(code)
@@ -203,86 +210,16 @@ 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 {
lower := strings.ToLower(h)
if _, unsafe := unsafeProxyHeaders[lower]; unsafe {
continue
}
if utils.SliceContains(conf.SlicesMap[conf.ProxyIgnoreHeaders], lower) {
if utils.SliceContains(conf.SlicesMap[conf.ProxyIgnoreHeaders], strings.ToLower(h)) {
continue
}
result[h] = val
}
// needed header, produced by the storage driver rather than the client
// needed header
for h, val := range override {
result[h] = val
}
-67
View File
@@ -1,67 +0,0 @@
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)
}
}
-236
View File
@@ -1,236 +0,0 @@
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() }
+5 -6
View File
@@ -25,6 +25,7 @@ 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"
)
@@ -183,7 +184,7 @@ func AddURL(ctx context.Context, args *AddURLArgs) (task.TaskExtensionInfo, erro
t := &DownloadTask{
TaskExtension: task.TaskExtension{
Creator: taskCreator,
ApiUrl: conf.GetApiUrl(ctx),
ApiUrl: common.GetApiUrl(ctx),
},
Url: args.URL,
DstDirPath: args.DstDirPath,
@@ -222,17 +223,15 @@ func isEd2kURL(urlStr string) bool {
}
func ed2kToolForStorage(storage driver.Driver) string {
name := NativeToolName(storage)
switch name {
switch toolNameForStorage(storage) {
case "115 Cloud", "115 Open":
return name
return toolNameForStorage(storage)
default:
return ""
}
}
// NativeToolName returns the offline-download tool implemented by storage.
func NativeToolName(storage driver.Driver) string {
func toolNameForStorage(storage driver.Driver) string {
switch storage.(type) {
case *_115.Pan115:
return "115 Cloud"
+3 -3
View File
@@ -58,7 +58,7 @@ func TestEd2kToolForStorage(t *testing.T) {
}
}
func TestNativeToolName(t *testing.T) {
func TestToolNameForStorage(t *testing.T) {
tests := []struct {
name string
storage driver.Driver
@@ -78,8 +78,8 @@ func TestNativeToolName(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := NativeToolName(tt.storage); got != tt.want {
t.Fatalf("NativeToolName(%T) = %q, want %q", tt.storage, got, tt.want)
if got := toolNameForStorage(tt.storage); got != tt.want {
t.Fatalf("toolNameForStorage(%T) = %q, want %q", tt.storage, got, tt.want)
}
})
}
+1 -1
View File
@@ -45,7 +45,7 @@ func (t ToolsManager) NamesForPath(path string) []string {
return names
}
name := NativeToolName(storage)
name := toolNameForStorage(storage)
if name == "" {
return names
}
+3 -2
View File
@@ -20,6 +20,7 @@ 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"
@@ -139,7 +140,7 @@ func transferStd(ctx context.Context, tempDir, dstDirPath string, deletePolicy D
TaskData: fs.TaskData{
TaskExtension: task.TaskExtension{
Creator: taskCreator,
ApiUrl: conf.GetApiUrl(ctx),
ApiUrl: common.GetApiUrl(ctx),
},
SrcActualPath: stdpath.Join(tempDir, entry.Name()),
DstActualPath: dstDirActualPath,
@@ -275,7 +276,7 @@ func transferObj(ctx context.Context, tempDir, dstDirPath string, deletePolicy D
TaskData: fs.TaskData{
TaskExtension: task.TaskExtension{
Creator: taskCreator,
ApiUrl: conf.GetApiUrl(ctx),
ApiUrl: common.GetApiUrl(ctx),
},
SrcActualPath: stdpath.Join(srcObjActualPath, obj.GetName()),
DstActualPath: dstDirActualPath,
+7 -11
View File
@@ -390,9 +390,8 @@ func ArchiveGet(ctx context.Context, storage driver.Driver, path string, args mo
}
type objWithLink struct {
link *model.Link
obj model.Obj
policy linkCachePolicy
link *model.Link
obj model.Obj
}
var (
@@ -406,7 +405,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.acquire() {
if ol.link.Expiration != nil || ol.link.SyncClosers.AcquireReference() || !ol.link.RequireReference {
return ol.link, ol.obj, nil
}
}
@@ -416,8 +415,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.policy.expiration != nil {
extractCache.SetWithTTL(key, ol, *ol.policy.expiration)
if ol.link.Expiration != nil {
extractCache.SetWithTTL(key, ol, *ol.link.Expiration)
} else {
extractCache.SetWithExpirable(key, ol, &ol.link.SyncClosers)
}
@@ -429,7 +428,7 @@ func DriverExtract(ctx context.Context, storage driver.Driver, path string, args
if err != nil {
return nil, nil, err
}
if ol.acquire() {
if ol.link.SyncClosers.AcquireReference() || !ol.link.RequireReference {
return ol.link, ol.obj, nil
}
}
@@ -451,10 +450,7 @@ func driverExtract(ctx context.Context, storage driver.Driver, path string, args
return nil, errors.WithStack(errs.NotFile)
}
link, err := storageAr.Extract(ctx, archiveFile, args)
if err != nil {
return nil, err
}
return admitLink(link, extracted)
return &objWithLink{link: link, obj: extracted}, err
}
type streamWithParent struct {
+7 -12
View File
@@ -233,10 +233,7 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li
if mode == -1 {
mode = storage.(driver.LinkCacheModeResolver).ResolveLinkCacheMode(path)
}
typeKey := "proxy/" + args.Type
if args.Redirect {
typeKey = "redirect/" + args.Type
}
typeKey := args.Type
if mode&driver.LinkCacheIP != 0 {
typeKey += "/" + args.IP
}
@@ -245,7 +242,8 @@ 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.acquire() {
if ol.link.Expiration != nil ||
ol.link.SyncClosers.AcquireReference() || !ol.link.RequireReference {
return ol.link, ol.obj, nil
}
}
@@ -263,12 +261,9 @@ 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, err := admitLink(link, file)
if err != nil {
return nil, err
}
if ol.policy.expiration != nil {
Cache.linkCache.SetTypeWithTTL(key, typeKey, ol, *ol.policy.expiration)
ol := &objWithLink{link: link, obj: file}
if link.Expiration != nil {
Cache.linkCache.SetTypeWithTTL(key, typeKey, ol, *link.Expiration)
} else {
Cache.linkCache.SetTypeWithExpirable(key, typeKey, ol, &link.SyncClosers)
}
@@ -279,7 +274,7 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li
if err != nil {
return nil, nil, err
}
if ol.acquire() {
if ol.link.SyncClosers.AcquireReference() || !ol.link.RequireReference {
return ol.link, ol.obj, nil
}
}
-71
View File
@@ -1,71 +0,0 @@
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)
}
})
}
}
-34
View File
@@ -1,34 +0,0 @@
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
}
-143
View File
@@ -1,143 +0,0 @@
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())
}
})
}
+2 -2
View File
@@ -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 = conf.GetApiUrl(ctx) + l.URL
l.URL = common.GetApiUrl(ctx) + l.URL
}
return sharing, l, obj, nil
}
-179
View File
@@ -1,179 +0,0 @@
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)
}
}
+34 -71
View File
@@ -8,7 +8,6 @@ import (
"io"
"math"
"os"
"sort"
"sync"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
@@ -359,72 +358,10 @@ func (r *ReaderUpdatingProgress) Close() error {
type RangeReadReadAtSeeker struct {
ss *SeekableStream
masterOff int64
readers orderedReaders
readerMap sync.Map
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
@@ -459,7 +396,7 @@ func (r *headCache) Close() error {
func (r *RangeReadReadAtSeeker) InitHeadCache() {
if r.masterOff == 0 {
value, _ := r.readers.takeExact(0)
value, _ := r.readerMap.LoadAndDelete(int64(0))
r.headCache = &headCache{reader: value.(io.Reader)}
r.ss.Closers.Add(r.headCache)
}
@@ -485,9 +422,9 @@ func NewReadAtSeeker(ss *SeekableStream, offset int64, forceRange ...bool) (mode
if err != nil {
return nil, err
}
r.readers.store(offset, reader)
r.readerMap.Store(int64(offset), reader)
} else {
r.readers.store(0, ss)
r.readerMap.Store(int64(offset), ss)
}
return r, nil
}
@@ -505,15 +442,41 @@ func NewMultiReaderAt(ss []*SeekableStream) (readerutil.SizeReaderAt, error) {
}
func (r *RangeReadReadAtSeeker) getReaderAtOffset(off int64) (io.Reader, error) {
if rr, cur, ok := r.readers.takeBest(off); ok {
if cur == off {
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)
return rr, nil
}
n, _ := utils.CopyWithBufferN(io.Discard, rr, off-cur)
if cur+n == off {
cur += n
if cur == off {
// logrus.Debugf("getReaderAtOffset old_%d", 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
@@ -538,7 +501,7 @@ func (r *RangeReadReadAtSeeker) ReadAt(p []byte, off int64) (n int, err error) {
off += int64(n)
switch err {
case nil:
r.readers.store(off, rr)
r.readerMap.Store(int64(off), rr)
case io.ErrUnexpectedEOF:
err = io.EOF
}
-20
View File
@@ -1,20 +0,0 @@
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)
}
}
+2 -1
View File
@@ -31,5 +31,6 @@ func GetApiUrlFromRequest(r *http.Request) string {
}
func GetApiUrl(ctx context.Context) string {
return conf.GetApiUrl(ctx)
api, _ := ctx.Value(conf.ApiUrlKey).(string)
return api
}
+6 -2
View File
@@ -34,7 +34,9 @@ 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, rrf)
return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, &model.RangeReadCloser{
RangeReader: rrf,
})
}
if link.RangeReader != nil {
@@ -43,7 +45,9 @@ 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, link.RangeReader)
return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, &model.RangeReadCloser{
RangeReader: link.RangeReader,
})
}
//transparent proxy
-49
View File
@@ -1,49 +0,0 @@
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())
}
}()
}
}
+2 -3
View File
@@ -94,11 +94,10 @@ func (f *FileUploadProxy) Close() error {
return err
}
arr := make([]byte, 512)
n, err := f.buffer.Read(arr)
if err != nil && err != io.EOF {
if _, err := f.buffer.Read(arr); err != nil {
return err
}
contentType := http.DetectContentType(arr[:n])
contentType := http.DetectContentType(arr)
if _, err := f.buffer.Seek(0, io.SeekStart); err != nil {
return err
}
-12
View File
@@ -113,10 +113,6 @@ 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) {
@@ -220,10 +216,6 @@ 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) {
@@ -381,10 +373,6 @@ 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] = ""
-126
View File
@@ -1,126 +0,0 @@
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)
}
}
+5 -7
View File
@@ -347,13 +347,11 @@ func FsGet(c *gin.Context, req *FsGetReq, user *model.User) {
}
}
}
parentPath := stdpath.Dir(reqPath)
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)
}
parentPath := stdpath.Dir(reqPath)
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)
@@ -368,7 +366,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.GetObjType(obj.GetName(), obj.IsDir()),
Type: utils.GetFileType(obj.GetName()),
Thumb: thumb,
MountDetails: mountDetails,
},
+270 -69
View File
@@ -3,6 +3,15 @@ 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"
@@ -14,44 +23,6 @@ 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"`
@@ -67,8 +38,18 @@ 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},
}
version, ok := saveAndInitOfflineDownloadTool(c, "aria2", items)
if !ok {
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)
return
}
common.SuccessResp(c, version)
@@ -89,7 +70,17 @@ 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 _, ok := saveAndInitOfflineDownloadTool(c, "qBittorrent", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
@@ -110,7 +101,17 @@ 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 _, ok := saveAndInitOfflineDownloadTool(c, "Transmission", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
@@ -126,13 +127,35 @@ func Set115(c *gin.Context) {
common.ErrorResp(c, err, 400)
return
}
if !validateOfflineDownloadStorage(c, req.TempDir, "115 Cloud") {
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
}
}
items := []model.SettingItem{
{Key: conf.Pan115TempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
}
if _, ok := saveAndInitOfflineDownloadTool(c, "115 Cloud", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
@@ -148,13 +171,35 @@ func Set115Open(c *gin.Context) {
common.ErrorResp(c, err, 400)
return
}
if !validateOfflineDownloadStorage(c, req.TempDir, "115 Open") {
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
}
}
items := []model.SettingItem{
{Key: conf.Pan115OpenTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
}
if _, ok := saveAndInitOfflineDownloadTool(c, "115 Open", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
@@ -170,13 +215,35 @@ func Set123Pan(c *gin.Context) {
common.ErrorResp(c, err, 400)
return
}
if !validateOfflineDownloadStorage(c, req.TempDir, "123Pan") {
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
}
}
items := []model.SettingItem{
{Key: conf.Pan123TempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
}
if _, ok := saveAndInitOfflineDownloadTool(c, "123Pan", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
@@ -193,14 +260,36 @@ func Set123Open(c *gin.Context) {
common.ErrorResp(c, err, 400)
return
}
if !validateOfflineDownloadStorage(c, req.TempDir, "123 Open") {
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
}
}
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 _, ok := saveAndInitOfflineDownloadTool(c, "123 Open", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
@@ -216,13 +305,35 @@ func SetPikPak(c *gin.Context) {
common.ErrorResp(c, err, 400)
return
}
if !validateOfflineDownloadStorage(c, req.TempDir, "PikPak") {
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
}
}
items := []model.SettingItem{
{Key: conf.PikPakTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
}
if _, ok := saveAndInitOfflineDownloadTool(c, "PikPak", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
@@ -238,13 +349,35 @@ func SetThunder(c *gin.Context) {
common.ErrorResp(c, err, 400)
return
}
if !validateOfflineDownloadStorage(c, req.TempDir, "Thunder") {
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
}
}
items := []model.SettingItem{
{Key: conf.ThunderTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
}
if _, ok := saveAndInitOfflineDownloadTool(c, "Thunder", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
@@ -260,13 +393,35 @@ func SetThunderX(c *gin.Context) {
common.ErrorResp(c, err, 400)
return
}
if !validateOfflineDownloadStorage(c, req.TempDir, "ThunderX") {
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
}
}
items := []model.SettingItem{
{Key: conf.ThunderXTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
}
if _, ok := saveAndInitOfflineDownloadTool(c, "ThunderX", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
@@ -282,13 +437,37 @@ func SetThunderBrowser(c *gin.Context) {
common.ErrorResp(c, err, 400)
return
}
if !validateOfflineDownloadStorage(c, req.TempDir, "ThunderBrowser") {
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
}
}
items := []model.SettingItem{
{Key: conf.ThunderBrowserTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
}
if _, ok := saveAndInitOfflineDownloadTool(c, "ThunderBrowser", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
@@ -304,13 +483,35 @@ func SetGuangYaPan(c *gin.Context) {
common.ErrorResp(c, err, 400)
return
}
if !validateOfflineDownloadStorage(c, req.TempDir, "GuangYaPan") {
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
}
}
items := []model.SettingItem{
{Key: conf.GuangYaPanTempDir, Value: req.TempDir, Type: conf.TypeString, Group: model.OFFLINE_DOWNLOAD, Flag: model.PRIVATE},
}
if _, ok := saveAndInitOfflineDownloadTool(c, "GuangYaPan", items); !ok {
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)
return
}
common.SuccessResp(c, "ok")
-155
View File
@@ -1,155 +0,0 @@
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)
}
}
+8 -17
View File
@@ -44,7 +44,14 @@ func Search(c *gin.Context) {
return
}
nodes, total, err := search.SearchFiltered(c, req.SearchReq, func(node model.SearchNode) bool {
return isSearchNodeAccessible(user, node, req.Password, op.GetNearestMeta)
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)
})
if err != nil {
common.ErrorResp(c, err, 500)
@@ -56,22 +63,6 @@ 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,
-78
View File
@@ -1,78 +0,0 @@
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)
}
})
}
}
+1 -1
View File
@@ -54,7 +54,7 @@ func SharingGet(c *gin.Context, req *FsGetReq) {
HashInfoStr: obj.GetHash().String(),
HashInfo: obj.GetHash().Export(),
Sign: "",
Type: utils.GetObjType(obj.GetName(), obj.IsDir()),
Type: utils.GetFileType(obj.GetName()),
Thumb: thumb,
},
RawURL: url,
+36 -51
View File
@@ -122,53 +122,6 @@ 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
@@ -385,7 +338,15 @@ func OIDCLoginCallback(c *gin.Context) {
c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof)
return
}
ssoPostMessage(c, map[string]string{"sso_id": bindingProof})
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))
return
}
if method == "sso_get_token" {
@@ -406,7 +367,15 @@ func OIDCLoginCallback(c *gin.Context) {
c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token)
return
}
ssoPostMessage(c, map[string]string{"token": token})
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))
return
}
}
@@ -547,7 +516,15 @@ func SSOLoginCallback(c *gin.Context) {
c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof)
return
}
ssoPostMessage(c, map[string]string{"sso_id": bindingProof})
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))
return
}
username := utils.Json.Get(resp.Body(), usernameField).ToString()
@@ -568,5 +545,13 @@ func SSOLoginCallback(c *gin.Context) {
c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token)
return
}
ssoPostMessage(c, map[string]string{"token": token})
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))
}
-102
View File
@@ -1,102 +0,0 @@
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)
}
}
+4 -6
View File
@@ -68,11 +68,9 @@ func (s *Server) callFSGet(c *gin.Context, raw json.RawMessage) (any, *rpcError)
parentPath := stdpath.Dir(reqPath)
var related []model.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)
}
sameLevelFiles, err := fs.List(ctx, parentPath, &fs.ListArgs{})
if err == nil {
related = filterRelatedObjs(sameLevelFiles, obj)
}
parentMeta, _ := op.GetNearestMeta(parentPath)
@@ -87,7 +85,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.GetObjType(obj.GetName(), obj.IsDir()),
Type: utils.GetFileType(obj.GetName()),
HashInfoStr: obj.GetHash().String(),
HashInfo: obj.GetHash().Export(),
MountDetails: mountDetails,
-62
View File
@@ -1,62 +0,0 @@
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",
)
}
+2 -5
View File
@@ -152,8 +152,6 @@ 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
@@ -195,7 +193,7 @@ func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string
return nil, fmt.Errorf("the remote storage driver need to be enhanced to support s3")
}
var rd io.ReadCloser
var rd io.Reader
if rnge != nil {
rd, err = rrf.RangeRead(ctx, http_range.Range(*rnge))
} else {
@@ -217,7 +215,6 @@ func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string
meta[k] = v
}
}
closers := utils.NewClosers(rd, link)
return &gofakes3.Object{
// Name: gofakes3.URLEncode(objectName),
@@ -226,7 +223,7 @@ func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string
Metadata: meta,
Size: size,
Range: rnge,
Contents: utils.ReadCloser{Reader: rd, Closer: &closers},
Contents: utils.ReadCloser{Reader: rd, Closer: link},
}, nil
}
-66
View File
@@ -1,66 +0,0 @@
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)
}
}
-132
View File
@@ -1,132 +0,0 @@
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)
}
}
-8
View File
@@ -5,7 +5,6 @@ package s3
import (
"context"
"encoding/json"
stderrors "errors"
"strings"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
@@ -22,13 +21,6 @@ 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) {