mirror of
https://github.com/OpenListTeam/OpenList.git
synced 2026-10-10 21:13:10 +08:00
Compare commits
28 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 5be5bc7873 | |||
| 966164d3a0 | |||
| f828938095 | |||
| fe575686f1 | |||
| e73e69f05e | |||
| 6d17d37ee7 | |||
| f933d59ec4 | |||
| 79c1d9d721 | |||
| 6f2ba09df9 | |||
| 624fdd24e9 | |||
| 06423083c7 | |||
| 4903bf61a9 | |||
| d29baa0eaf | |||
| d684e45d63 | |||
| ac1192a8d9 | |||
| 172ef17421 | |||
| 85be4214ee | |||
| e10ebd2694 | |||
| 8467a3abe0 | |||
| aa42d9e5fd | |||
| 596e284fb7 | |||
| 2ed55487df | |||
| 2e6dd00d91 | |||
| 09b150a0ef | |||
| 238dbb66ec | |||
| 9c7ad84242 | |||
| a5e5048555 | |||
| b0f6919f86 |
@@ -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=' }}
|
||||
|
||||
@@ -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=' }}
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
package _115
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
|
||||
"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/pkg/http_range"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回 115 支持的秒传哈希算法(SHA1)
|
||||
func (d *Pan115) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*utils.SHA1}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 115 不需要分片哈希
|
||||
func (d *Pan115) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 SHA1 哈希尝试秒传。
|
||||
//
|
||||
// 115 的秒传协议除了整文件 SHA1,还需要文件头部 128KB 的 SHA1(pre_hash),
|
||||
// 因此当内容源可用时会打开它来计算前置哈希;内容不可用时无法秒传。
|
||||
func (d *Pan115) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
fullHash := strings.ToUpper(req.Whole.GetHash(utils.SHA1))
|
||||
if len(fullHash) != utils.SHA1.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
if req.Open == nil {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
src, err := req.Open()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer src.Close()
|
||||
|
||||
const PreHashSize int64 = 128 * utils.KB
|
||||
hashSize := PreHashSize
|
||||
if req.Size < PreHashSize {
|
||||
hashSize = req.Size
|
||||
}
|
||||
reader, err := src.RangeRead(http_range.Range{Start: 0, Length: hashSize})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
preHash, err := utils.HashReader(utils.SHA1, reader)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
preHash = strings.ToUpper(preHash)
|
||||
|
||||
fastInfo, err := d.rapidUpload(req.Size, req.Name, dstDir.GetID(), preHash, fullHash, src)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
matched, err := fastInfo.Ok()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !matched {
|
||||
return nil, errs.ErrRapidUploadFailed
|
||||
}
|
||||
f, err := d.getNewFileByPickCode(fastInfo.PickCode)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return f, nil
|
||||
}
|
||||
@@ -171,11 +171,22 @@ func (d *Open115) Get(ctx context.Context, path string) (model.Obj, error) {
|
||||
path = stdpath.Join(d.parentPath, path)
|
||||
resp, err := d.client.GetFolderInfoByPath(ctx, path)
|
||||
if err != nil {
|
||||
// SDK-level "object not found" (empty array response from API)
|
||||
if errors.Is(err, sdk.ErrObjectNotFound) {
|
||||
return d.getFromParent(ctx, path, "")
|
||||
}
|
||||
// API-level error response (State=false), treat as not found
|
||||
// since this is a path-lookup that can't resolve the target
|
||||
var apiErr *sdk.Error
|
||||
if errors.As(err, &apiErr) {
|
||||
return nil, errs.ObjectNotFound
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
// Handle null/empty response (e.g., API returns null data for non-existent path)
|
||||
if resp.FileID == "" {
|
||||
return nil, errs.ObjectNotFound
|
||||
}
|
||||
obj := &Obj{
|
||||
Fid: resp.FileID,
|
||||
Fn: resp.FileName,
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
package _123
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/base"
|
||||
"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/pkg/utils"
|
||||
"github.com/go-resty/resty/v2"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回 123 云盘支持的秒传哈希算法(MD5)
|
||||
func (d *Pan123) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*utils.MD5}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 123 云盘不需要分片哈希
|
||||
func (d *Pan123) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 MD5(Etag)哈希尝试秒传。
|
||||
//
|
||||
// 123 的秒传即「上传请求返回 reuse=true」,无需真正传输内容。
|
||||
func (d *Pan123) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
etag := req.Whole.GetHash(utils.MD5)
|
||||
if len(etag) < utils.MD5.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
duplicate := 0
|
||||
if overwrite {
|
||||
duplicate = 2
|
||||
}
|
||||
data := base.Json{
|
||||
"driveId": 0,
|
||||
"duplicate": duplicate,
|
||||
"etag": strings.ToLower(etag),
|
||||
"fileName": req.Name,
|
||||
"parentFileId": dstDir.GetID(),
|
||||
"size": req.Size,
|
||||
"type": 0,
|
||||
}
|
||||
|
||||
var resp UploadResp
|
||||
_, err := d.Request(UploadRequest, http.MethodPost, func(r *resty.Request) {
|
||||
r.SetBody(data).SetContext(ctx)
|
||||
}, &resp)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// reuse=true 或未返回上传 Key 均视为秒传成功
|
||||
if !resp.Data.Reuse && resp.Data.Key != "" {
|
||||
return nil, errs.ErrHashMismatch
|
||||
}
|
||||
|
||||
return &model.ObjThumb{
|
||||
Object: model.Object{
|
||||
ID: strconv.FormatInt(resp.Data.FileId, 10),
|
||||
Name: req.Name,
|
||||
Size: req.Size,
|
||||
Modified: time.Now(),
|
||||
IsFolder: false,
|
||||
HashInfo: utils.NewHashInfo(utils.MD5, strings.ToLower(etag)),
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package _123_open
|
||||
|
||||
import (
|
||||
"context"
|
||||
"path"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"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/pkg/utils"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回 123 开放平台支持的秒传哈希算法(SHA1)
|
||||
func (d *Open123) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*utils.SHA1}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 123 开放平台不需要分片哈希
|
||||
func (d *Open123) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 SHA1 哈希尝试秒传
|
||||
func (d *Open123) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
sha1Hash := req.Whole.GetHash(utils.SHA1)
|
||||
if len(sha1Hash) < utils.SHA1.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
parentID, err := strconv.ParseInt(dstDir.GetID(), 10, 64)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resp, err := d.sha1Reuse(parentID, req.Name, sha1Hash, req.Size, 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !resp.Data.Reuse {
|
||||
return nil, errs.ErrHashMismatch
|
||||
}
|
||||
|
||||
return &model.ObjThumb{
|
||||
Object: model.Object{
|
||||
ID: strconv.FormatInt(resp.Data.FileID, 10),
|
||||
Name: req.Name,
|
||||
Size: req.Size,
|
||||
IsFolder: false,
|
||||
Path: path.Join(dstDir.GetPath(), req.Name),
|
||||
Modified: time.Now(),
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
@@ -6,7 +6,6 @@ import (
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/torrent"
|
||||
@@ -15,12 +14,8 @@ import (
|
||||
|
||||
// GenerateTorrent 根据上传过程中收集的哈希信息生成包含 CAS 扩展的 torrent 文件
|
||||
func GenerateTorrent(fileName string, fileSize int64, fileMD5 string, sliceMD5s []string, sliceSize int64, pieceHashes []byte) ([]byte, error) {
|
||||
// 计算 sliceMD5
|
||||
sliceMD5 := fileMD5
|
||||
if len(sliceMD5s) > 1 {
|
||||
joined := strings.Join(sliceMD5s, "\n")
|
||||
sliceMD5 = strings.ToUpper(torrent.GetMD5Str(joined))
|
||||
}
|
||||
// 计算 sliceMD5(统一走规范实现)
|
||||
sliceMD5 := torrent.SliceMD5FromPieces(sliceMD5s, fileMD5)
|
||||
|
||||
t := torrent.NewTorrent(fileName, fileSize, fileMD5)
|
||||
t.Info.PieceLength = sliceSize
|
||||
@@ -30,7 +25,7 @@ func GenerateTorrent(fileName string, fileSize int64, fileMD5 string, sliceMD5s
|
||||
SliceMD5: sliceMD5,
|
||||
SliceMD5s: sliceMD5s,
|
||||
SliceSize: sliceSize,
|
||||
Cloud: "189",
|
||||
Cloud: torrent.Cloud189,
|
||||
})
|
||||
|
||||
return t.Encode()
|
||||
@@ -95,7 +90,7 @@ func ComputeTorrentFromReader(reader io.Reader, fileName string, fileSize int64,
|
||||
sliceSize = torrent.DefaultPieceSize
|
||||
}
|
||||
|
||||
hw := torrent.NewHashWriter(sliceSize, sliceSize)
|
||||
hw := torrent.NewHashWriter(sliceSize, sliceSize, fileSize)
|
||||
|
||||
buf := make([]byte, 32*1024)
|
||||
for {
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
package _189_tv
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"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/pkg/utils"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回 189 电视支持的秒传哈希算法(MD5)
|
||||
func (d *Cloud189TV) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*utils.MD5}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 189 电视不需要分片哈希
|
||||
func (d *Cloud189TV) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 MD5 哈希尝试秒传
|
||||
func (d *Cloud189TV) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
md5Hash := req.Whole.GetHash(utils.MD5)
|
||||
if len(md5Hash) < utils.MD5.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
stream := driver.NewSeedHashStream(req)
|
||||
obj, err := d.RapidUpload(ctx, dstDir, stream, d.isFamily(), overwrite)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return obj, nil
|
||||
}
|
||||
@@ -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"`
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
package _189pc
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"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/pkg/utils"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回 189pc 支持的秒传哈希算法(MD5)
|
||||
func (d *Cloud189PC) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*utils.MD5}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 189pc 的 CAS 秒传依赖分片 MD5
|
||||
func (d *Cloud189PC) RapidHashNeedsPieces() bool {
|
||||
return true
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 MD5 哈希尝试秒传
|
||||
func (d *Cloud189PC) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
md5Hash := req.Whole.GetHash(utils.MD5)
|
||||
if len(md5Hash) < utils.MD5.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
stream := driver.NewSeedHashStream(req)
|
||||
obj, err := d.RapidUpload(ctx, dstDir, stream, d.isFamily(), overwrite)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return obj, nil
|
||||
}
|
||||
+17
-109
@@ -7,7 +7,6 @@ import (
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -28,12 +27,8 @@ import (
|
||||
// fileName: 文件名
|
||||
// fileSize: 文件大小
|
||||
func GenerateTorrent(fileName string, fileSize int64, fileMD5 string, sliceMD5s []string, sliceSize int64, pieceHashes []byte) ([]byte, error) {
|
||||
// 计算 sliceMD5
|
||||
sliceMD5 := fileMD5
|
||||
if len(sliceMD5s) > 1 {
|
||||
joined := strings.Join(sliceMD5s, "\n")
|
||||
sliceMD5 = strings.ToUpper(torrent.GetMD5Str(joined))
|
||||
}
|
||||
// 计算 sliceMD5(统一走规范实现)
|
||||
sliceMD5 := torrent.SliceMD5FromPieces(sliceMD5s, fileMD5)
|
||||
|
||||
t := torrent.NewTorrent(fileName, fileSize, fileMD5)
|
||||
t.Info.PieceLength = sliceSize
|
||||
@@ -43,7 +38,7 @@ func GenerateTorrent(fileName string, fileSize int64, fileMD5 string, sliceMD5s
|
||||
SliceMD5: sliceMD5,
|
||||
SliceMD5s: sliceMD5s,
|
||||
SliceSize: sliceSize,
|
||||
Cloud: "189",
|
||||
Cloud: torrent.Cloud189,
|
||||
})
|
||||
|
||||
return t.Encode()
|
||||
@@ -69,101 +64,18 @@ func (y *Cloud189PC) RapidUploadFromTorrent(ctx context.Context, dstDir model.Ob
|
||||
fileName := t.Info.Name
|
||||
fileSize := t.GetTotalSize()
|
||||
|
||||
// 统一 MD5 为大写(与正常上传保持一致,天翼云盘要求大写)
|
||||
fileMD5Upper := strings.ToUpper(cas.FileMD5)
|
||||
|
||||
// 优先使用 torrent 中嵌入的分片大小,与生成时保持一致
|
||||
sliceSize := cas.SliceSize
|
||||
if sliceSize <= 0 {
|
||||
sliceSize = partSize(fileSize)
|
||||
// 优先使用 torrent 中嵌入的分片 MD5 与大写整文件 MD5
|
||||
fileMD5 := strings.ToUpper(cas.FileMD5)
|
||||
sliceMD5s := make([]string, len(cas.SliceMD5s))
|
||||
for i, s := range cas.SliceMD5s {
|
||||
sliceMD5s[i] = strings.ToUpper(s)
|
||||
}
|
||||
|
||||
// 计算 sliceMd5(与上传时一致的算法)
|
||||
// 优先使用 torrent 中已有的 SliceMD5;仅当有多分片列表时才重新计算
|
||||
sliceMd5Hex := strings.ToUpper(cas.SliceMD5)
|
||||
if sliceMd5Hex == "" {
|
||||
sliceMd5Hex = fileMD5Upper
|
||||
}
|
||||
if len(cas.SliceMD5s) > 1 {
|
||||
// 分片 MD5 也需要统一大写后再拼接计算
|
||||
upperSliceMD5s := make([]string, len(cas.SliceMD5s))
|
||||
for i, s := range cas.SliceMD5s {
|
||||
upperSliceMD5s[i] = strings.ToUpper(s)
|
||||
}
|
||||
sliceMd5Hex = strings.ToUpper(utils.GetMD5EncodeStr(strings.Join(upperSliceMD5s, "\n")))
|
||||
}
|
||||
|
||||
// 使用与 Web 端一致的三步秒传流程
|
||||
fullUrl := "https://upload.cloud.189.cn"
|
||||
if isFamily {
|
||||
fullUrl += "/family"
|
||||
} else {
|
||||
fullUrl += "/person"
|
||||
}
|
||||
|
||||
// Step 1: initMultiUpload(不传 fileMd5/sliceMd5,只传 lazyCheck)
|
||||
initParams := Params{
|
||||
"parentFolderId": dstDir.GetID(),
|
||||
"fileName": url.QueryEscape(fileName),
|
||||
"fileSize": fmt.Sprint(fileSize),
|
||||
"sliceSize": fmt.Sprint(sliceSize),
|
||||
"lazyCheck": "1",
|
||||
}
|
||||
if isFamily {
|
||||
initParams.Set("familyId", y.FamilyID)
|
||||
}
|
||||
|
||||
var uploadInfo InitMultiUploadResp
|
||||
_, err = y.request(fullUrl+"/initMultiUpload", "GET", func(req *resty.Request) {
|
||||
req.SetContext(ctx)
|
||||
}, initParams, &uploadInfo, isFamily)
|
||||
// 复用统一的 CAS 秒传核心实现
|
||||
respObj, err := y.rapidUploadByCAS(ctx, dstDir, fileName, fileSize, fileMD5, sliceMD5s, cas.SliceSize, overwrite)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("initMultiUpload 失败: %w", err)
|
||||
}
|
||||
|
||||
uploadFileId := uploadInfo.Data.UploadFileID
|
||||
|
||||
// Step 2: checkTransSecond(用 fileMd5 + sliceMd5 + uploadFileId 检查秒传)
|
||||
checkParams := Params{
|
||||
"fileMd5": fileMD5Upper,
|
||||
"sliceMd5": sliceMd5Hex,
|
||||
"uploadFileId": uploadFileId,
|
||||
}
|
||||
|
||||
var checkResp struct {
|
||||
Data struct {
|
||||
FileDataExists int `json:"fileDataExists"`
|
||||
} `json:"data"`
|
||||
}
|
||||
_, err = y.request(fullUrl+"/checkTransSecond", "GET", func(req *resty.Request) {
|
||||
req.SetContext(ctx)
|
||||
}, checkParams, &checkResp, isFamily)
|
||||
if err != nil {
|
||||
utils.Log.Errorf("[RapidUpload] checkTransSecond 失败: uploadFileId=%s, err=%v", uploadFileId, err)
|
||||
return nil, fmt.Errorf("秒传检查失败: %w", err)
|
||||
}
|
||||
|
||||
if checkResp.Data.FileDataExists != 1 {
|
||||
return nil, fmt.Errorf("秒传失败:云端不存在该文件(fileMD5=%s, sliceMD5=%s, size=%d)", fileMD5Upper, sliceMd5Hex, fileSize)
|
||||
}
|
||||
|
||||
// Step 3: commitMultiUploadFile(传 fileMd5 + sliceMd5)
|
||||
|
||||
var resp CommitMultiUploadFileResp
|
||||
commitParams := Params{
|
||||
"uploadFileId": uploadFileId,
|
||||
"fileMd5": fileMD5Upper,
|
||||
"sliceMd5": sliceMd5Hex,
|
||||
"lazyCheck": "1",
|
||||
"opertype": IF(overwrite, "3", "1"),
|
||||
}
|
||||
|
||||
_, err = y.request(fullUrl+"/commitMultiUploadFile", "GET", func(req *resty.Request) {
|
||||
req.SetContext(ctx)
|
||||
}, commitParams, &resp, isFamily)
|
||||
if err != nil {
|
||||
utils.Log.Errorf("[RapidUpload] commitMultiUploadFile 失败: uploadFileId=%s, err=%v", uploadFileId, err)
|
||||
return nil, fmt.Errorf("提交上传失败: %w", err)
|
||||
utils.Log.Errorf("[RapidUpload] 秒传失败: fileMD5=%s, err=%v", fileMD5, err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 秒传成功后,将 torrent 文件上传到目标目录(异步,不影响秒传结果)
|
||||
@@ -196,7 +108,7 @@ func (y *Cloud189PC) RapidUploadFromTorrent(ctx context.Context, dstDir model.Ob
|
||||
}()
|
||||
}
|
||||
|
||||
return resp.toFile(), nil
|
||||
return respObj, nil
|
||||
}
|
||||
|
||||
// ComputeTorrentFromReader 从 io.Reader 计算并生成 torrent 文件
|
||||
@@ -206,7 +118,7 @@ func ComputeTorrentFromReader(reader io.Reader, fileName string, fileSize int64,
|
||||
sliceSize = torrent.DefaultPieceSize
|
||||
}
|
||||
|
||||
hw := torrent.NewHashWriter(sliceSize, sliceSize)
|
||||
hw := torrent.NewHashWriter(sliceSize, sliceSize, fileSize)
|
||||
|
||||
buf := make([]byte, 32*1024)
|
||||
for {
|
||||
@@ -259,12 +171,8 @@ func InjectCASIntoTorrent(torrentData []byte, fileMD5 string, sliceMD5s []string
|
||||
return nil, fmt.Errorf("解析 torrent 失败: %w", err)
|
||||
}
|
||||
|
||||
// 计算 sliceMD5
|
||||
sliceMD5 := fileMD5
|
||||
if len(sliceMD5s) > 1 {
|
||||
joined := strings.Join(sliceMD5s, "\n")
|
||||
sliceMD5 = strings.ToUpper(torrent.GetMD5Str(joined))
|
||||
}
|
||||
// 计算 sliceMD5(统一走规范实现)
|
||||
sliceMD5 := torrent.SliceMD5FromPieces(sliceMD5s, fileMD5)
|
||||
|
||||
// 注入 CAS 信息
|
||||
t.SetCASInfo(&torrent.CASInfo{
|
||||
@@ -272,7 +180,7 @@ func InjectCASIntoTorrent(torrentData []byte, fileMD5 string, sliceMD5s []string
|
||||
SliceMD5: sliceMD5,
|
||||
SliceMD5s: sliceMD5s,
|
||||
SliceSize: sliceSize,
|
||||
Cloud: "189",
|
||||
Cloud: torrent.Cloud189,
|
||||
})
|
||||
|
||||
// 同时更新 info 中的 md5sum 字段
|
||||
|
||||
@@ -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"`
|
||||
|
||||
+173
-355
@@ -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)
|
||||
}
|
||||
@@ -1607,6 +1329,102 @@ func (y *Cloud189PC) OldUploadCommit(ctx context.Context, fileCommitUrl string,
|
||||
return resp.toFile(), nil
|
||||
}
|
||||
|
||||
// rapidUploadByCAS 使用 MD5 + 分片 MD5(CAS)执行天翼云盘秒传。
|
||||
//
|
||||
// 流程与 Web 端一致:
|
||||
// 1. initMultiUpload(仅传 lazyCheck=1)
|
||||
// 2. checkTransSecond(用 fileMd5 + sliceMd5 检查云端是否已存在文件数据)
|
||||
// 3. commitMultiUploadFile(提交并返回文件对象)
|
||||
func (y *Cloud189PC) rapidUploadByCAS(ctx context.Context, dstDir model.Obj, fileName string, fileSize int64, fileMD5 string, sliceMD5s []string, sliceSize int64, overwrite bool) (model.Obj, error) {
|
||||
isFamily := y.isFamily()
|
||||
|
||||
// 统一 MD5 为大写(天翼云盘要求大写)
|
||||
fileMD5Upper := strings.ToUpper(fileMD5)
|
||||
|
||||
// 优先使用传入的分片大小,否则按文件大小推导
|
||||
if sliceSize <= 0 {
|
||||
sliceSize = partSize(fileSize)
|
||||
}
|
||||
|
||||
// 计算 sliceMd5(与上传时一致的算法)
|
||||
sliceMd5Hex := fileMD5Upper
|
||||
if len(sliceMD5s) > 1 {
|
||||
upperSliceMD5s := make([]string, len(sliceMD5s))
|
||||
for i, s := range sliceMD5s {
|
||||
upperSliceMD5s[i] = strings.ToUpper(s)
|
||||
}
|
||||
sliceMd5Hex = strings.ToUpper(utils.GetMD5EncodeStr(strings.Join(upperSliceMD5s, "\n")))
|
||||
} else if len(sliceMD5s) == 1 {
|
||||
sliceMd5Hex = strings.ToUpper(sliceMD5s[0])
|
||||
}
|
||||
|
||||
fullUrl := "https://upload.cloud.189.cn"
|
||||
if isFamily {
|
||||
fullUrl += "/family"
|
||||
} else {
|
||||
fullUrl += "/person"
|
||||
}
|
||||
|
||||
// Step 1: initMultiUpload(不传 fileMd5/sliceMd5,只传 lazyCheck)
|
||||
initParams := Params{
|
||||
"parentFolderId": dstDir.GetID(),
|
||||
"fileName": url.QueryEscape(fileName),
|
||||
"fileSize": fmt.Sprint(fileSize),
|
||||
"sliceSize": fmt.Sprint(sliceSize),
|
||||
"lazyCheck": "1",
|
||||
}
|
||||
if isFamily {
|
||||
initParams.Set("familyId", y.FamilyID)
|
||||
}
|
||||
|
||||
var uploadInfo InitMultiUploadResp
|
||||
if _, err := y.request(fullUrl+"/initMultiUpload", "GET", func(req *resty.Request) {
|
||||
req.SetContext(ctx)
|
||||
}, initParams, &uploadInfo, isFamily); err != nil {
|
||||
return nil, fmt.Errorf("initMultiUpload 失败: %w", err)
|
||||
}
|
||||
uploadFileId := uploadInfo.Data.UploadFileID
|
||||
|
||||
// Step 2: checkTransSecond(用 fileMd5 + sliceMd5 + uploadFileId 检查秒传)
|
||||
checkParams := Params{
|
||||
"fileMd5": fileMD5Upper,
|
||||
"sliceMd5": sliceMd5Hex,
|
||||
"uploadFileId": uploadFileId,
|
||||
}
|
||||
|
||||
var checkResp struct {
|
||||
Data struct {
|
||||
FileDataExists int `json:"fileDataExists"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if _, err := y.request(fullUrl+"/checkTransSecond", "GET", func(req *resty.Request) {
|
||||
req.SetContext(ctx)
|
||||
}, checkParams, &checkResp, isFamily); err != nil {
|
||||
return nil, fmt.Errorf("秒传检查失败: %w", err)
|
||||
}
|
||||
if checkResp.Data.FileDataExists != 1 {
|
||||
return nil, fmt.Errorf("秒传失败:云端不存在该文件(fileMD5=%s, sliceMD5=%s, size=%d)", fileMD5Upper, sliceMd5Hex, fileSize)
|
||||
}
|
||||
|
||||
// Step 3: commitMultiUploadFile(传 fileMd5 + sliceMd5)
|
||||
commitParams := Params{
|
||||
"uploadFileId": uploadFileId,
|
||||
"fileMd5": fileMD5Upper,
|
||||
"sliceMd5": sliceMd5Hex,
|
||||
"lazyCheck": "1",
|
||||
"opertype": IF(overwrite, "3", "1"),
|
||||
}
|
||||
|
||||
var resp CommitMultiUploadFileResp
|
||||
if _, err := y.request(fullUrl+"/commitMultiUploadFile", "GET", func(req *resty.Request) {
|
||||
req.SetContext(ctx)
|
||||
}, commitParams, &resp, isFamily); err != nil {
|
||||
return nil, fmt.Errorf("提交上传失败: %w", err)
|
||||
}
|
||||
|
||||
return resp.toFile(), nil
|
||||
}
|
||||
|
||||
func (y *Cloud189PC) isFamily() bool {
|
||||
return y.Type == "family"
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
package aliyundrive_open
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/base"
|
||||
"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/pkg/utils"
|
||||
"github.com/go-resty/resty/v2"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回阿里云盘支持的秒传哈希算法(SHA1)
|
||||
func (d *AliyundriveOpen) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*utils.SHA1}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 阿里云盘不需要分片哈希
|
||||
func (d *AliyundriveOpen) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 SHA1 哈希尝试秒传。
|
||||
//
|
||||
// 阿里云盘的秒传还需要 proof_code(按 proof range 读取的一段内容),
|
||||
// 因此当内容源不可用时无法完成秒传。
|
||||
func (d *AliyundriveOpen) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
sha1Hash := req.Whole.GetHash(utils.SHA1)
|
||||
if len(sha1Hash) < utils.SHA1.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
if req.Open == nil {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
stream, err := req.Open()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer stream.Close()
|
||||
|
||||
proofCode, err := d.calProofCode(stream)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var resp CreateResp
|
||||
_, err = d.request(ctx, limiterOther, "/adrive/v1.0/openFile/create", http.MethodPost, func(r *resty.Request) {
|
||||
r.SetBody(base.Json{
|
||||
"drive_id": d.DriveId,
|
||||
"parent_file_id": dstDir.GetID(),
|
||||
"name": req.Name,
|
||||
"type": "file",
|
||||
"check_name_mode": "auto_rename",
|
||||
"size": req.Size,
|
||||
"content_hash": sha1Hash,
|
||||
"content_hash_name": "sha1",
|
||||
"proof_version": "v1",
|
||||
"proof_code": proofCode,
|
||||
}).SetResult(&resp)
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !resp.RapidUpload {
|
||||
return nil, errs.ErrHashMismatch
|
||||
}
|
||||
|
||||
if resp.FileId != "" {
|
||||
obj, err := d.completeUpload(ctx, resp.FileId, resp.UploadId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return obj, nil
|
||||
}
|
||||
|
||||
return &model.ObjThumb{
|
||||
Object: model.Object{
|
||||
Name: req.Name,
|
||||
Size: req.Size,
|
||||
Modified: time.Now(),
|
||||
IsFolder: false,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package baidu_netdisk
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"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/pkg/utils"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回百度网盘支持的秒传哈希算法(MD5)
|
||||
func (d *BaiduNetdisk) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*utils.MD5}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 百度网盘不需要分片哈希
|
||||
func (d *BaiduNetdisk) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 MD5 哈希尝试秒传
|
||||
func (d *BaiduNetdisk) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
md5Hash := req.Whole.GetHash(utils.MD5)
|
||||
if len(md5Hash) < utils.MD5.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
stream := driver.NewSeedHashStream(req)
|
||||
obj, err := d.PutRapid(ctx, dstDir, stream)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return obj, nil
|
||||
}
|
||||
+8
-43
@@ -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
|
||||
}
|
||||
|
||||
@@ -245,6 +245,7 @@ func mFile2Object(file MFile) *model.ObjThumbURL {
|
||||
Ctime: ctime,
|
||||
IsFolder: false,
|
||||
Size: file.Size,
|
||||
HashInfo: utils.NewHashInfo(utils.MD5, file.MD5),
|
||||
},
|
||||
Thumbnail: model.Thumbnail{
|
||||
Thumbnail: file.ThumbnailURL,
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
package pikpak
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/base"
|
||||
"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/pkg/utils"
|
||||
hash_extend "github.com/OpenListTeam/OpenList/v4/pkg/utils/hash"
|
||||
"github.com/go-resty/resty/v2"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回 PikPak 支持的秒传哈希算法(GCID)
|
||||
func (d *PikPak) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*hash_extend.GCID}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces PikPak 不需要分片哈希
|
||||
func (d *PikPak) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 GCID 哈希尝试秒传
|
||||
func (d *PikPak) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
gcid := req.Whole.GetHash(hash_extend.GCID)
|
||||
if len(gcid) < hash_extend.GCID.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
var resp UploadTaskData
|
||||
_, err := d.request("https://api-drive.mypikpak.net/drive/v1/files", http.MethodPost, func(r *resty.Request) {
|
||||
r.SetContext(ctx).SetBody(base.Json{
|
||||
"kind": "drive#file",
|
||||
"name": req.Name,
|
||||
"size": req.Size,
|
||||
"hash": strings.ToUpper(gcid),
|
||||
"upload_type": "UPLOAD_TYPE_RESUMABLE",
|
||||
"objProvider": base.Json{"provider": "UPLOAD_TYPE_UNKNOWN"},
|
||||
"parent_id": dstDir.GetID(),
|
||||
"folder_type": "NORMAL",
|
||||
})
|
||||
}, &resp)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 秒传成功时不会返回 Resumable
|
||||
if resp.Resumable == nil {
|
||||
file := fileToObj(resp.File)
|
||||
return file, nil
|
||||
}
|
||||
|
||||
return nil, errs.ErrHashMismatch
|
||||
}
|
||||
+10
-30
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
package quark_open
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"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/pkg/utils"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回夸克网盘支持的秒传哈希算法(MD5 + SHA1)
|
||||
func (d *QuarkOpen) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*utils.MD5, *utils.SHA1}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 夸克网盘不需要分片哈希
|
||||
func (d *QuarkOpen) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 MD5/SHA1 哈希尝试秒传。
|
||||
//
|
||||
// 夸克网盘的预上传需要 proof_code(按 proof range 读取的一段内容),
|
||||
// 因此当内容源不可用时无法完成秒传。
|
||||
func (d *QuarkOpen) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
md5Hash := req.Whole.GetHash(utils.MD5)
|
||||
sha1Hash := req.Whole.GetHash(utils.SHA1)
|
||||
if len(md5Hash) < utils.MD5.Width || len(sha1Hash) < utils.SHA1.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
if req.Open == nil {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
stream, err := req.Open()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer stream.Close()
|
||||
|
||||
pre, err := d.upPre(ctx, stream, dstDir.GetID(), md5Hash, sha1Hash)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !pre.Data.Finish {
|
||||
return nil, errs.ErrHashMismatch
|
||||
}
|
||||
|
||||
return &model.ObjThumb{
|
||||
Object: model.Object{
|
||||
ID: pre.Data.Fid,
|
||||
Name: req.Name,
|
||||
Size: req.Size,
|
||||
IsFolder: false,
|
||||
Modified: time.Now(),
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package quark_open
|
||||
|
||||
import (
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"time"
|
||||
)
|
||||
|
||||
@@ -57,6 +58,7 @@ func fileToObj(f File) *model.ObjThumb {
|
||||
Modified: time.UnixMilli(f.UpdatedAt),
|
||||
IsFolder: f.FileType == "0",
|
||||
Ctime: time.UnixMilli(f.CreatedAt),
|
||||
HashInfo: utils.NewHashInfo(utils.SHA1, f.ContentHash),
|
||||
},
|
||||
Thumbnail: model.Thumbnail{Thumbnail: f.ThumbnailURL},
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
package thunder
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/base"
|
||||
"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/pkg/utils"
|
||||
hash_extend "github.com/OpenListTeam/OpenList/v4/pkg/utils/hash"
|
||||
"github.com/go-resty/resty/v2"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回迅雷支持的秒传哈希算法(GCID)
|
||||
func (xc *XunLeiCommon) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*hash_extend.GCID}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 迅雷不需要分片哈希
|
||||
func (xc *XunLeiCommon) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 GCID 哈希尝试秒传
|
||||
func (xc *XunLeiCommon) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
gcid := req.Whole.GetHash(hash_extend.GCID)
|
||||
if len(gcid) < hash_extend.GCID.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
var resp UploadTaskResponse
|
||||
_, err := xc.Request(FILE_API_URL, http.MethodPost, func(r *resty.Request) {
|
||||
r.SetContext(ctx)
|
||||
r.SetBody(&base.Json{
|
||||
"kind": FILE,
|
||||
"parent_id": dstDir.GetID(),
|
||||
"name": req.Name,
|
||||
"size": req.Size,
|
||||
"hash": gcid,
|
||||
"upload_type": UPLOAD_TYPE_RESUMABLE,
|
||||
"space": xc.Space,
|
||||
})
|
||||
}, &resp)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 秒传成功(UploadType != UPLOAD_TYPE_RESUMABLE)
|
||||
if resp.UploadType != UPLOAD_TYPE_RESUMABLE {
|
||||
return &resp.File, nil
|
||||
}
|
||||
|
||||
return nil, errs.ErrHashMismatch
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package thunder_browser
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/base"
|
||||
"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/pkg/utils"
|
||||
hash_extend "github.com/OpenListTeam/OpenList/v4/pkg/utils/hash"
|
||||
"github.com/go-resty/resty/v2"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回迅雷浏览器支持的秒传哈希算法(GCID)
|
||||
func (xc *XunLeiBrowserCommon) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*hash_extend.GCID}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 迅雷浏览器不需要分片哈希
|
||||
func (xc *XunLeiBrowserCommon) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 GCID 哈希尝试秒传
|
||||
func (xc *XunLeiBrowserCommon) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
gcid := req.Whole.GetHash(hash_extend.GCID)
|
||||
if len(gcid) < hash_extend.GCID.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
var resp UploadTaskResponse
|
||||
_, err := xc.Request(FILE_API_URL, http.MethodPost, func(r *resty.Request) {
|
||||
r.SetContext(ctx)
|
||||
r.SetBody(&base.Json{
|
||||
"kind": FILE,
|
||||
"parent_id": dstDir.GetID(),
|
||||
"name": req.Name,
|
||||
"size": req.Size,
|
||||
"hash": gcid,
|
||||
"upload_type": UPLOAD_TYPE_RESUMABLE,
|
||||
})
|
||||
}, &resp)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 秒传成功(UploadType != UPLOAD_TYPE_RESUMABLE)
|
||||
if resp.UploadType != UPLOAD_TYPE_RESUMABLE {
|
||||
return &resp.File, nil
|
||||
}
|
||||
|
||||
return nil, errs.ErrHashMismatch
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package thunderx
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/drivers/base"
|
||||
"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/pkg/utils"
|
||||
hash_extend "github.com/OpenListTeam/OpenList/v4/pkg/utils/hash"
|
||||
"github.com/go-resty/resty/v2"
|
||||
)
|
||||
|
||||
// RapidHashAlgos 返回迅雷X支持的秒传哈希算法(GCID)
|
||||
func (xc *XunLeiXCommon) RapidHashAlgos() []utils.HashType {
|
||||
return []utils.HashType{*hash_extend.GCID}
|
||||
}
|
||||
|
||||
// RapidHashNeedsPieces 迅雷X不需要分片哈希
|
||||
func (xc *XunLeiXCommon) RapidHashNeedsPieces() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// RapidUploadByHashes 使用种子中的 GCID 哈希尝试秒传
|
||||
func (xc *XunLeiXCommon) RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *driver.SeedRapidUploadRequest, overwrite bool) (model.Obj, error) {
|
||||
gcid := req.Whole.GetHash(hash_extend.GCID)
|
||||
if len(gcid) < hash_extend.GCID.Width {
|
||||
return nil, errs.ErrUnavailableHash
|
||||
}
|
||||
|
||||
var resp UploadTaskResponse
|
||||
_, err := xc.Request(FILE_API_URL, http.MethodPost, func(r *resty.Request) {
|
||||
r.SetContext(ctx)
|
||||
r.SetBody(&base.Json{
|
||||
"kind": FILE,
|
||||
"parent_id": dstDir.GetID(),
|
||||
"name": req.Name,
|
||||
"size": req.Size,
|
||||
"hash": gcid,
|
||||
"upload_type": UPLOAD_TYPE_RESUMABLE,
|
||||
})
|
||||
}, &resp)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 秒传成功(UploadType != UPLOAD_TYPE_RESUMABLE)
|
||||
if resp.UploadType != UPLOAD_TYPE_RESUMABLE {
|
||||
return &resp.File, nil
|
||||
}
|
||||
|
||||
return nil, errs.ErrHashMismatch
|
||||
}
|
||||
@@ -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
|
||||
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
|
||||
|
||||
@@ -51,10 +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.2 h1:iR4B8WH0qWqWkzVTNSj2TgWw6kovTh2bV8TGMOFSnVI=
|
||||
github.com/OpenListTeam/gofakes3 v0.8.2/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=
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -187,6 +187,14 @@ func InitialSettings() []model.SettingItem {
|
||||
{Key: conf.HandleHookAfterWriting, Value: "false", Type: conf.TypeBool, Group: model.GLOBAL, Flag: model.PRIVATE},
|
||||
{Key: conf.HandleHookRateLimit, Value: "0", Type: conf.TypeNumber, Group: model.GLOBAL, Flag: model.PRIVATE},
|
||||
{Key: conf.IgnoreSystemFiles, Value: "false", Type: conf.TypeBool, Group: model.GLOBAL, Flag: model.PRIVATE, Help: `When enabled, ignores common system files during upload (.DS_Store, desktop.ini, Thumbs.db, and files starting with ._)`},
|
||||
{Key: conf.SeedSiteURL, Value: "", Type: conf.TypeString, Group: model.GLOBAL, Flag: model.PRIVATE, Help: `Public base URL embedded in generated transfer seed sources when configured`},
|
||||
{Key: conf.SeedDefaultMatrix, Value: `{"md5":{"whole":true,"pieces":false},"sha1":{"whole":true,"pieces":false},"sha256":{"whole":true,"pieces":false}}`, Type: conf.TypeText, Group: model.GLOBAL, Flag: model.PRIVATE, Help: `Default right-click hash matrix for transfer seed generation`},
|
||||
{Key: conf.SeedFormatPolicies, Value: `{"oss":"off","torrent":"off","cas":"off"}`, Type: conf.TypeText, Group: model.GLOBAL, Flag: model.PRIVATE},
|
||||
{Key: conf.SeedDefaultFormat, Value: "oss", Type: conf.TypeSelect, Options: "oss,torrent,cas", Group: model.GLOBAL, Flag: model.PRIVATE},
|
||||
{Key: conf.SeedSingleDirectPreview, Value: "false", Type: conf.TypeBool, Group: model.GLOBAL, Flag: model.PUBLIC},
|
||||
{Key: conf.SeedCASDirectAccess, Value: "false", Type: conf.TypeBool, Group: model.GLOBAL, Flag: model.PUBLIC, Help: `When opening a single-file CAS seed, immediately rapid-upload it into the same folder and preview the restored file`},
|
||||
{Key: conf.SeedAutoGeneratePolicy, Value: "off", Type: conf.TypeSelect, Options: "off,on", Group: model.GLOBAL, Flag: model.PRIVATE, Help: `Global upload sidecar policy; storage-specific inheritance can be layered without changing the safe default`},
|
||||
{Key: conf.SeedDefaultTrackers, Value: "", Type: conf.TypeText, Group: model.GLOBAL, Flag: model.PRIVATE, Help: `Default tracker list offered when generating torrent seeds (one tracker per line)`},
|
||||
|
||||
// single settings
|
||||
{Key: conf.Token, Value: token, Type: conf.TypeString, Group: model.SINGLE, Flag: model.PRIVATE},
|
||||
@@ -211,7 +219,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},
|
||||
|
||||
@@ -49,4 +49,5 @@ func InitTaskManager() {
|
||||
op.RegisterSettingChangingCallback(func() {
|
||||
fs.ArchiveContentUploadTaskManager.SetWorkersNumActive(taskFilterNegative(setting.GetInt(conf.TaskDecompressUploadThreadsNum, conf.Conf.Tasks.DecompressUpload.Workers)))
|
||||
})
|
||||
fs.SeedGenerateTaskManager = tache.NewManager[*fs.SeedGenerateTask](tache.WithWorks(setting.GetInt(conf.TaskUploadThreadsNum, conf.Conf.Tasks.Upload.Workers)), tache.WithMaxRetry(conf.Conf.Tasks.Upload.MaxRetry)) //seed generation will not support persist
|
||||
}
|
||||
|
||||
+10
-1
@@ -60,6 +60,16 @@ const (
|
||||
HandleHookRateLimit = "handle_hook_rate_limit"
|
||||
IgnoreSystemFiles = "ignore_system_files"
|
||||
|
||||
// transfer seeds
|
||||
SeedSiteURL = "seed_site_url"
|
||||
SeedDefaultMatrix = "seed_default_matrix"
|
||||
SeedFormatPolicies = "seed_format_policies"
|
||||
SeedDefaultFormat = "seed_default_format"
|
||||
SeedSingleDirectPreview = "seed_single_direct_preview"
|
||||
SeedCASDirectAccess = "seed_cas_direct_access"
|
||||
SeedAutoGeneratePolicy = "seed_auto_generate_policy"
|
||||
SeedDefaultTrackers = "seed_default_trackers"
|
||||
|
||||
// index
|
||||
SearchIndex = "search_index"
|
||||
AutoUpdateIndex = "auto_update_index"
|
||||
@@ -117,7 +127,6 @@ const (
|
||||
SSODefaultDir = "sso_default_dir"
|
||||
SSODefaultPermission = "sso_default_permission"
|
||||
SSOCompatibilityMode = "sso_compatibility_mode"
|
||||
SSOPostMessageOrigin = "sso_postmessage_origin"
|
||||
|
||||
// ldap
|
||||
LdapLoginEnabled = "ldap_login_enabled"
|
||||
|
||||
@@ -1,8 +0,0 @@
|
||||
package conf
|
||||
|
||||
import "context"
|
||||
|
||||
func GetApiUrl(ctx context.Context) string {
|
||||
api, _ := ctx.Value(ApiUrlKey).(string)
|
||||
return api
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
)
|
||||
|
||||
type Driver interface {
|
||||
@@ -218,3 +219,54 @@ type DirectUploader interface {
|
||||
// return errs.NotImplement if the driver does not support the given direct upload tool
|
||||
GetDirectUploadInfo(ctx context.Context, tool string, dstDir model.Obj, fileName string, fileSize int64) (any, error)
|
||||
}
|
||||
|
||||
// SeedRapidUploadRequest carries everything a driver needs to attempt a
|
||||
// hash-driven rapid upload (秒传/CAS) without transferring the full content.
|
||||
//
|
||||
// It is populated from a parsed transfer seed, so that any cloud drive
|
||||
// supporting hash-based rapid upload can be used as a seed save target.
|
||||
type SeedRapidUploadRequest struct {
|
||||
// Name is the target file name.
|
||||
Name string
|
||||
// Size is the total file size in bytes.
|
||||
Size int64
|
||||
// Whole holds whole-file hashes indexed by algorithm (e.g. utils.MD5,
|
||||
// utils.SHA1). It must never be nil; individual entries may be empty.
|
||||
Whole *utils.HashInfo
|
||||
// SliceSize is the per-slice/piece size in bytes (0 when unknown).
|
||||
SliceSize int64
|
||||
// SliceMD5s is the ordered per-slice MD5 list (used by 189pc-style CAS).
|
||||
SliceMD5s []string
|
||||
// SliceSHA1s is the ordered per-slice SHA1 list (used by SHA1-piece drives).
|
||||
SliceSHA1s []string
|
||||
// Open lazily yields the file content as a model.FileStreamer. Drivers whose
|
||||
// rapid-upload protocol needs partial or full content (e.g. a leading
|
||||
// pre-hash or a proof-code) may call it; hash-only drivers may ignore it.
|
||||
// Open may be nil when no content source is available, in which case
|
||||
// drivers that strictly require content must fail gracefully.
|
||||
Open func() (model.FileStreamer, error)
|
||||
}
|
||||
|
||||
// SeedRapidUploader is an optional capability interface implemented by drivers
|
||||
// that can perform a "rapid upload" (秒传/CAS) driven by precomputed hashes
|
||||
// instead of a full content transfer.
|
||||
//
|
||||
// It generalizes the previous hard-coded 189pc-specific CAS path so that any
|
||||
// cloud drive supporting hash-based rapid upload (e.g. 189pc via MD5+slice MD5,
|
||||
// 115/aliyundrive_open via SHA1) can be used as a transfer-seed save target.
|
||||
type SeedRapidUploader interface {
|
||||
// RapidUploadByHashes attempts a rapid upload of a file into dstDir.
|
||||
//
|
||||
// Implementations should return errs.NotImplement (or a descriptive error)
|
||||
// when the required hash is missing, the file does not exist remotely, or
|
||||
// the rapid upload cannot be confirmed.
|
||||
RapidUploadByHashes(ctx context.Context, dstDir model.Obj, req *SeedRapidUploadRequest, overwrite bool) (model.Obj, error)
|
||||
|
||||
// RapidHashAlgos reports the whole-file hash algorithms accepted by
|
||||
// RapidUploadByHashes (e.g. utils.MD5, utils.SHA1).
|
||||
RapidHashAlgos() []utils.HashType
|
||||
|
||||
// RapidHashNeedsPieces reports whether RapidUploadByHashes relies on
|
||||
// per-slice hashes (CAS slice MD5s / SHA1 pieces) for this driver.
|
||||
RapidHashNeedsPieces() bool
|
||||
}
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
package driver
|
||||
|
||||
import (
|
||||
"io"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
)
|
||||
|
||||
// SeedHashStream is a model.FileStreamer that only carries file metadata and
|
||||
// precomputed hashes; it never yields real content.
|
||||
//
|
||||
// It exists so that hash-driven rapid upload (秒传/CAS) implementations can reuse
|
||||
// the drivers' existing Put/RapidUpload code paths, which expect a
|
||||
// model.FileStreamer but only read the name/size/hash. Drivers that additionally
|
||||
// need the real content (e.g. to compute a leading proof hash) should supply
|
||||
// Source, which is exposed via GetReadCloser-like lazy opening.
|
||||
type SeedHashStream struct {
|
||||
name string
|
||||
size int64
|
||||
hashInfo utils.HashInfo
|
||||
// Source lazily opens the underlying content streamer. May be nil.
|
||||
Source func() (model.FileStreamer, error)
|
||||
}
|
||||
|
||||
var (
|
||||
_ model.FileStreamer = (*SeedHashStream)(nil)
|
||||
_ utils.ClosersIF = (*SeedHashStream)(nil)
|
||||
)
|
||||
|
||||
// NewSeedHashStream builds a hash-only streamer from a rapid-upload request.
|
||||
func NewSeedHashStream(req *SeedRapidUploadRequest) *SeedHashStream {
|
||||
s := &SeedHashStream{name: req.Name, size: req.Size, Source: req.Open}
|
||||
if req.Whole != nil {
|
||||
s.hashInfo = *req.Whole
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func (s *SeedHashStream) GetName() string { return s.name }
|
||||
func (s *SeedHashStream) GetSize() int64 { return s.size }
|
||||
func (s *SeedHashStream) GetHash() utils.HashInfo { return s.hashInfo }
|
||||
func (s *SeedHashStream) GetMimetype() string { return "" }
|
||||
func (s *SeedHashStream) ModTime() time.Time { return time.Now() }
|
||||
func (s *SeedHashStream) CreateTime() time.Time { return time.Now() }
|
||||
func (s *SeedHashStream) IsDir() bool { return false }
|
||||
func (s *SeedHashStream) GetID() string { return "" }
|
||||
func (s *SeedHashStream) GetPath() string { return "" }
|
||||
|
||||
func (s *SeedHashStream) NeedStore() bool { return false }
|
||||
func (s *SeedHashStream) IsForceStreamUpload() bool { return true }
|
||||
func (s *SeedHashStream) GetExist() model.Obj { return nil }
|
||||
func (s *SeedHashStream) SetExist(model.Obj) {}
|
||||
func (s *SeedHashStream) GetFile() model.File { return nil }
|
||||
func (s *SeedHashStream) Add(io.Closer) {}
|
||||
func (s *SeedHashStream) AddIfCloser(any) {}
|
||||
func (s *SeedHashStream) Close() error { return nil }
|
||||
|
||||
// Read returns EOF: the stream carries hashes only, no content.
|
||||
func (s *SeedHashStream) Read([]byte) (int, error) { return 0, io.EOF }
|
||||
|
||||
// RangeRead returns an empty reader, since no content is available.
|
||||
func (s *SeedHashStream) RangeRead(http_range.Range) (io.Reader, error) {
|
||||
return nil, io.EOF
|
||||
}
|
||||
|
||||
// CacheFullAndWriter reports that the content cannot be materialized.
|
||||
func (s *SeedHashStream) CacheFullAndWriter(*model.UpdateProgress, io.Writer) (model.File, error) {
|
||||
return nil, io.EOF
|
||||
}
|
||||
@@ -4,4 +4,16 @@ import "errors"
|
||||
|
||||
var (
|
||||
EmptyToken = errors.New("empty token")
|
||||
|
||||
// ErrUnavailableHash indicates the seed does not carry the hash algorithm
|
||||
// required by the destination driver, so rapid upload cannot be attempted.
|
||||
ErrUnavailableHash = errors.New("required hash is unavailable")
|
||||
// ErrEmptyHash indicates a required hash exists but is empty/too short.
|
||||
ErrEmptyHash = errors.New("empty hash")
|
||||
// ErrHashMismatch indicates the remote side rejected the provided hash, so
|
||||
// a full content transfer is required instead of a rapid upload.
|
||||
ErrHashMismatch = errors.New("hash mismatch")
|
||||
// ErrRapidUploadFailed indicates the driver attempted a rapid upload but
|
||||
// could not confirm success.
|
||||
ErrRapidUploadFailed = errors.New("rapid upload failed")
|
||||
)
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -0,0 +1,603 @@
|
||||
package fs
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
stdpath "path"
|
||||
"slices"
|
||||
"strings"
|
||||
"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/internal/op"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/setting"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/task"
|
||||
"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"
|
||||
)
|
||||
|
||||
// MaxSeedGenerateSyncSize bounds synchronous seed generation (1GB). Larger
|
||||
// requests are turned into an asynchronous task.
|
||||
const MaxSeedGenerateSyncSize = 1 * 1024 * 1024 * 1024
|
||||
|
||||
// SeedGenerateNeedsAsync reports whether the given files must be generated
|
||||
// asynchronously because they exceed the synchronous size limit. It resolves
|
||||
// each path and sums the file sizes, returning the first error encountered.
|
||||
func SeedGenerateNeedsAsync(ctx context.Context, user *model.User, paths []string) (bool, error) {
|
||||
var total int64
|
||||
for _, requestedPath := range paths {
|
||||
fullPath, err := user.JoinPath(requestedPath)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
storage, actualPath, err := op.GetStorageAndActualPath(fullPath)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
obj, err := op.Get(ctx, storage, actualPath)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("seed path must be a readable file: %s", requestedPath)
|
||||
}
|
||||
if obj.IsDir() {
|
||||
return false, fmt.Errorf("seed path must be a readable file: %s", requestedPath)
|
||||
}
|
||||
total += obj.GetSize()
|
||||
if total > MaxSeedGenerateSyncSize {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// SeedHashSelection controls whole-file and piece hash inclusion.
|
||||
type SeedHashSelection struct {
|
||||
Whole bool `json:"whole"`
|
||||
Pieces bool `json:"pieces"`
|
||||
}
|
||||
|
||||
// SeedHashMatrix controls the optional hash metadata stored in a seed.
|
||||
type SeedHashMatrix struct {
|
||||
MD5 SeedHashSelection `json:"md5"`
|
||||
SHA1 SeedHashSelection `json:"sha1"`
|
||||
SHA256 SeedHashSelection `json:"sha256"`
|
||||
}
|
||||
|
||||
// SeedGenerateParams carries a fully-resolved seed generation request, free of
|
||||
// any HTTP transport concerns so it can run synchronously or as a task.
|
||||
type SeedGenerateParams struct {
|
||||
Paths []string
|
||||
Formats []string
|
||||
Name string
|
||||
Comment string
|
||||
FileComments map[string]string
|
||||
HashMatrix SeedHashMatrix
|
||||
PieceSize int64
|
||||
Trackers []string
|
||||
Channels []torrent.SeedChannel
|
||||
OutputPath string
|
||||
IncludeShare bool
|
||||
IncludeDirectSource bool
|
||||
ShareFiles []string
|
||||
DirectFiles []string
|
||||
}
|
||||
|
||||
// SeedArtifact is one generated seed container.
|
||||
type SeedArtifact struct {
|
||||
Format string `json:"format"`
|
||||
Name string `json:"name"`
|
||||
FileName string `json:"file_name"`
|
||||
SeedData string `json:"seed_data"`
|
||||
Size int `json:"size"`
|
||||
Path string `json:"path,omitempty"`
|
||||
}
|
||||
|
||||
// DeriveSeedName derives a sensible default seed name from the source paths:
|
||||
// single selection uses the file name, multi selection uses the common base
|
||||
// name (ignoring extensions) when all files share one, otherwise the folder name.
|
||||
func DeriveSeedName(paths []string) string {
|
||||
if len(paths) == 0 {
|
||||
return "OpenList Seed"
|
||||
}
|
||||
if len(paths) == 1 {
|
||||
return seedBaseName(paths[0])
|
||||
}
|
||||
// Common base name ignoring extensions (e.g. a.docx + a.exe -> "a").
|
||||
commonBase := seedBaseName(paths[0])
|
||||
for _, p := range paths[1:] {
|
||||
if base := seedBaseName(p); base != commonBase {
|
||||
commonBase = ""
|
||||
break
|
||||
}
|
||||
}
|
||||
if commonBase != "" {
|
||||
return commonBase
|
||||
}
|
||||
// Fall back to the common parent directory name.
|
||||
dir := commonParentDir(paths)
|
||||
if base := stdpath.Base(dir); base != "" && base != "/" && base != "." {
|
||||
return base
|
||||
}
|
||||
return "OpenList Seed"
|
||||
}
|
||||
|
||||
// seedBaseName returns the file name without its extension.
|
||||
func seedBaseName(p string) string {
|
||||
base := stdpath.Base(p)
|
||||
return strings.TrimSuffix(base, stdpath.Ext(base))
|
||||
}
|
||||
|
||||
// commonParentDir returns the longest common parent directory of the given paths.
|
||||
func commonParentDir(paths []string) string {
|
||||
if len(paths) == 0 {
|
||||
return "/"
|
||||
}
|
||||
parts := strings.Split(strings.Trim(stdpath.Dir(paths[0]), "/"), "/")
|
||||
for _, p := range paths[1:] {
|
||||
cur := strings.Split(strings.Trim(stdpath.Dir(p), "/"), "/")
|
||||
n := 0
|
||||
for n < len(parts) && n < len(cur) && parts[n] == cur[n] {
|
||||
n++
|
||||
}
|
||||
parts = parts[:n]
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return "/"
|
||||
}
|
||||
return "/" + strings.Join(parts, "/")
|
||||
}
|
||||
|
||||
// NormalizeSeedFormats validates and deduplicates a list of seed format names.
|
||||
func NormalizeSeedFormats(rawFormats []string) ([]string, error) {
|
||||
formats := append([]string(nil), rawFormats...)
|
||||
if len(formats) == 1 && strings.TrimSpace(formats[0]) == "" {
|
||||
formats[0] = setting.GetStr(conf.SeedDefaultFormat, "oss")
|
||||
}
|
||||
if len(formats) > 3 {
|
||||
return nil, fmt.Errorf("at most three seed formats may be generated")
|
||||
}
|
||||
result := make([]string, 0, len(formats))
|
||||
seen := make(map[string]struct{}, len(formats))
|
||||
for _, rawFormat := range formats {
|
||||
format := strings.ToLower(strings.TrimPrefix(strings.TrimSpace(rawFormat), "."))
|
||||
if format == "bt" {
|
||||
format = "torrent"
|
||||
}
|
||||
if format != "oss" && format != "torrent" && format != "cas" {
|
||||
return nil, fmt.Errorf("unsupported seed format %q", rawFormat)
|
||||
}
|
||||
if _, exists := seen[format]; exists {
|
||||
continue
|
||||
}
|
||||
seen[format] = struct{}{}
|
||||
result = append(result, format)
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func seedFormats(params SeedGenerateParams) ([]string, error) {
|
||||
formats, err := NormalizeSeedFormats(params.Formats)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(formats) == 0 {
|
||||
formats = []string{setting.GetStr(conf.SeedDefaultFormat, "oss")}
|
||||
}
|
||||
return formats, nil
|
||||
}
|
||||
|
||||
func seedMatrixEmpty(matrix SeedHashMatrix) bool {
|
||||
return !matrix.MD5.Whole && !matrix.MD5.Pieces && !matrix.SHA1.Whole && !matrix.SHA1.Pieces && !matrix.SHA256.Whole && !matrix.SHA256.Pieces
|
||||
}
|
||||
|
||||
func loadSeedDefaultMatrix() SeedHashMatrix {
|
||||
var matrix SeedHashMatrix
|
||||
if raw := strings.TrimSpace(setting.GetStr(conf.SeedDefaultMatrix)); raw != "" {
|
||||
_ = json.Unmarshal([]byte(raw), &matrix)
|
||||
}
|
||||
return matrix
|
||||
}
|
||||
|
||||
func normalizedSeedMatrix(matrix SeedHashMatrix, formats []string) SeedHashMatrix {
|
||||
if seedMatrixEmpty(matrix) {
|
||||
matrix = loadSeedDefaultMatrix()
|
||||
}
|
||||
if seedMatrixEmpty(matrix) {
|
||||
matrix = SeedHashMatrix{
|
||||
MD5: SeedHashSelection{Whole: true, Pieces: true}, SHA1: SeedHashSelection{Whole: true, Pieces: true},
|
||||
SHA256: SeedHashSelection{Whole: true, Pieces: true},
|
||||
}
|
||||
}
|
||||
for _, format := range formats {
|
||||
switch format {
|
||||
case "torrent":
|
||||
matrix.SHA1 = SeedHashSelection{Whole: true, Pieces: true}
|
||||
case "cas":
|
||||
matrix.MD5 = SeedHashSelection{Whole: true, Pieces: true}
|
||||
}
|
||||
}
|
||||
return matrix
|
||||
}
|
||||
|
||||
func canReuseListedHashes(hashInfo utils.HashInfo, matrix SeedHashMatrix) bool {
|
||||
if matrix.MD5.Pieces || matrix.SHA1.Pieces || matrix.SHA256.Pieces {
|
||||
return false
|
||||
}
|
||||
if matrix.MD5.Whole && hashInfo.GetHash(utils.MD5) == "" {
|
||||
return false
|
||||
}
|
||||
if matrix.SHA1.Whole && hashInfo.GetHash(utils.SHA1) == "" {
|
||||
return false
|
||||
}
|
||||
if matrix.SHA256.Whole && hashInfo.GetHash(utils.SHA256) == "" {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func applySeedMatrix(file *torrent.SeedFile, matrix SeedHashMatrix) {
|
||||
if !matrix.MD5.Whole {
|
||||
file.Hashes.MD5 = ""
|
||||
}
|
||||
if !matrix.SHA1.Whole {
|
||||
file.Hashes.SHA1 = ""
|
||||
}
|
||||
if !matrix.SHA256.Whole {
|
||||
file.Hashes.SHA256 = ""
|
||||
}
|
||||
if file.Hashes.Pieces == nil {
|
||||
return
|
||||
}
|
||||
if !matrix.MD5.Pieces {
|
||||
file.Hashes.Pieces.MD5 = nil
|
||||
}
|
||||
if !matrix.SHA1.Pieces {
|
||||
file.Hashes.Pieces.SHA1 = nil
|
||||
}
|
||||
if !matrix.SHA256.Pieces {
|
||||
file.Hashes.Pieces.SHA256 = nil
|
||||
}
|
||||
if len(file.Hashes.Pieces.MD5) == 0 && len(file.Hashes.Pieces.SHA1) == 0 && len(file.Hashes.Pieces.SHA256) == 0 {
|
||||
file.Hashes.Pieces = nil
|
||||
}
|
||||
}
|
||||
|
||||
// EncodeGeneratedSeed serializes a seed in the requested container format.
|
||||
func EncodeGeneratedSeed(seed *torrent.Seed, format string, standardPieces []byte) ([]byte, error) {
|
||||
if format != "torrent" {
|
||||
return torrent.EncodeSeed(seed, format)
|
||||
}
|
||||
t := &torrent.Torrent{
|
||||
Info: torrent.TorrentInfo{Name: seed.Name, PieceLength: seed.PieceSize, Pieces: standardPieces},
|
||||
Comment: seed.Comment,
|
||||
CreatedBy: seed.CreatedBy,
|
||||
CreationDate: time.Now().Unix(),
|
||||
OpenList: seed,
|
||||
}
|
||||
if len(seed.Trackers) > 0 {
|
||||
t.Announce = seed.Trackers[0]
|
||||
for _, tracker := range seed.Trackers {
|
||||
t.AnnounceList = append(t.AnnounceList, []string{tracker})
|
||||
}
|
||||
}
|
||||
if len(seed.Files) == 1 {
|
||||
file := seed.Files[0]
|
||||
t.Info.Name = stdpath.Base(file.Path)
|
||||
t.Info.Length = file.Size
|
||||
t.Info.MD5Sum = file.Hashes.MD5
|
||||
casCloud := file.CASCloud
|
||||
if casCloud == "" {
|
||||
casCloud = torrent.Cloud189
|
||||
}
|
||||
if file.CASSliceMD5 != "" {
|
||||
t.SetCASInfo(&torrent.CASInfo{
|
||||
FileMD5: strings.ToUpper(file.Hashes.MD5), SliceMD5: strings.ToUpper(file.CASSliceMD5),
|
||||
SliceSize: torrent.DefaultPieceSize, Cloud: casCloud,
|
||||
})
|
||||
} else if seed.PieceSize == torrent.DefaultPieceSize && file.Hashes.Pieces != nil && len(file.Hashes.Pieces.MD5) > 0 {
|
||||
t.SetCASInfo(torrent.BuildCASInfoFromMD5sWithCloud(file.Hashes.MD5, file.Hashes.Pieces.MD5, torrent.DefaultPieceSize, casCloud))
|
||||
}
|
||||
} else {
|
||||
for _, file := range seed.Files {
|
||||
t.Info.Files = append(t.Info.Files, torrent.TorrentFile{Length: file.Size, Path: strings.Split(file.Path, "/"), MD5Sum: file.Hashes.MD5})
|
||||
}
|
||||
}
|
||||
return t.Encode()
|
||||
}
|
||||
|
||||
// GenerateSeedArtifacts reads each file once while computing the requested
|
||||
// hashes, then emits one or more seed containers. It has no HTTP dependency and
|
||||
// can be driven synchronously or from a background task.
|
||||
func GenerateSeedArtifacts(ctx context.Context, user *model.User, params SeedGenerateParams) ([]SeedArtifact, *torrent.Seed, error) {
|
||||
if len(params.Paths) == 0 || len(params.Paths) > torrent.DefaultMaxSeedFiles {
|
||||
return nil, nil, fmt.Errorf("invalid seed file count")
|
||||
}
|
||||
formats, err := seedFormats(params)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
matrix := normalizedSeedMatrix(params.HashMatrix, formats)
|
||||
pieceSize := params.PieceSize
|
||||
if pieceSize <= 0 {
|
||||
pieceSize = torrent.DefaultPieceSize
|
||||
}
|
||||
if slices.Contains(formats, "cas") {
|
||||
pieceSize = torrent.DefaultPieceSize
|
||||
}
|
||||
seedName := strings.TrimSpace(params.Name)
|
||||
if seedName == "" {
|
||||
seedName = DeriveSeedName(params.Paths)
|
||||
}
|
||||
seed := torrent.NewSeed(seedName, "OpenList", pieceSize)
|
||||
seed.Comment = params.Comment
|
||||
seed.Trackers = params.Trackers
|
||||
seed.Channels = params.Channels
|
||||
shareSet := make(map[string]bool, len(params.ShareFiles))
|
||||
for _, p := range params.ShareFiles {
|
||||
if strings.TrimSpace(p) != "" {
|
||||
shareSet[p] = true
|
||||
}
|
||||
}
|
||||
directSet := make(map[string]bool, len(params.DirectFiles))
|
||||
for _, p := range params.DirectFiles {
|
||||
if strings.TrimSpace(p) != "" {
|
||||
directSet[p] = true
|
||||
}
|
||||
}
|
||||
useGlobalShare := len(shareSet) == 0 && params.IncludeShare
|
||||
useGlobalDirect := len(directSet) == 0 && params.IncludeDirectSource
|
||||
hasShare := useGlobalShare || len(shareSet) > 0
|
||||
hasDirect := useGlobalDirect || len(directSet) > 0
|
||||
if hasShare && !user.CanShare() {
|
||||
return nil, nil, errs.PermissionDenied
|
||||
}
|
||||
if hasDirect && setting.GetBool(conf.SignAll) && !hasShare {
|
||||
return nil, nil, fmt.Errorf("direct sources require an automatic share when global signing is enabled")
|
||||
}
|
||||
if (hasShare || hasDirect) && strings.TrimSpace(setting.GetStr(conf.SeedSiteURL)) == "" {
|
||||
return nil, nil, fmt.Errorf("seed_site_url must be configured before embedding download sources")
|
||||
}
|
||||
globalHasher := torrent.NewHashWriter(pieceSize, pieceSize, 0)
|
||||
fullPaths := make([]string, 0, len(params.Paths))
|
||||
var total int64
|
||||
for _, requestedPath := range params.Paths {
|
||||
fullPath, err := user.JoinPath(requestedPath)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
meta, err := op.GetNearestMeta(fullPath)
|
||||
if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) {
|
||||
return nil, nil, err
|
||||
}
|
||||
if !common.CanRead(user, meta, fullPath) {
|
||||
return nil, nil, errs.PermissionDenied
|
||||
}
|
||||
storage, actualPath, err := op.GetStorageAndActualPath(fullPath)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
obj, err := op.Get(ctx, storage, actualPath)
|
||||
if err != nil || obj.IsDir() {
|
||||
return nil, nil, fmt.Errorf("seed path must be a readable file: %s", requestedPath)
|
||||
}
|
||||
total += obj.GetSize()
|
||||
modified := ""
|
||||
if !obj.ModTime().IsZero() {
|
||||
modified = obj.ModTime().UTC().Format(time.RFC3339)
|
||||
}
|
||||
seedPath := stdpath.Base(requestedPath)
|
||||
if len(params.Paths) > 1 {
|
||||
seedPath = strings.TrimPrefix(stdpath.Clean(requestedPath), "/")
|
||||
}
|
||||
|
||||
hashInfo := obj.GetHash()
|
||||
if canReuseListedHashes(hashInfo, matrix) {
|
||||
seedFile := torrent.SeedFile{
|
||||
Path: seedPath,
|
||||
Size: obj.GetSize(),
|
||||
Modified: modified,
|
||||
Hashes: torrent.SeedHashes{
|
||||
MD5: strings.ToLower(hashInfo.GetHash(utils.MD5)),
|
||||
SHA1: strings.ToLower(hashInfo.GetHash(utils.SHA1)),
|
||||
SHA256: strings.ToLower(hashInfo.GetHash(utils.SHA256)),
|
||||
},
|
||||
}
|
||||
applySeedMatrix(&seedFile, matrix)
|
||||
if comment := strings.TrimSpace(params.FileComments[requestedPath]); comment != "" {
|
||||
seedFile.Comment = comment
|
||||
} else if comment := strings.TrimSpace(params.FileComments[obj.GetName()]); comment != "" {
|
||||
seedFile.Comment = comment
|
||||
}
|
||||
if useGlobalDirect || directSet[requestedPath] {
|
||||
baseURL := strings.TrimRight(setting.GetStr(conf.SeedSiteURL), "/")
|
||||
seedFile.Sources = []torrent.SeedSource{{Type: "openlist-direct", URL: baseURL + utils.EncodePath("/d"+fullPath)}}
|
||||
}
|
||||
seed.Files = append(seed.Files, seedFile)
|
||||
fullPaths = append(fullPaths, fullPath)
|
||||
continue
|
||||
}
|
||||
|
||||
link, _, err := op.Link(ctx, storage, actualPath, model.LinkArgs{})
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("storage cannot stream %s: %v", requestedPath, err)
|
||||
}
|
||||
rangeReader, err := stream.GetRangeReaderFromLink(obj.GetSize(), link)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("storage cannot stream %s", requestedPath)
|
||||
}
|
||||
rc, err := rangeReader.RangeRead(ctx, http_range.Range{Length: obj.GetSize()})
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
fileHasher := torrent.NewHashWriter(pieceSize, pieceSize, 0)
|
||||
n, copyErr := io.Copy(io.MultiWriter(globalHasher, fileHasher), rc)
|
||||
_ = rc.Close()
|
||||
if copyErr != nil {
|
||||
return nil, nil, fmt.Errorf("read %s: %w", requestedPath, copyErr)
|
||||
}
|
||||
if n != obj.GetSize() {
|
||||
return nil, nil, fmt.Errorf("read %s: got %d of %d bytes", requestedPath, n, obj.GetSize())
|
||||
}
|
||||
fileHasher.Finish()
|
||||
seedFile := fileHasher.BuildSeedFile(seedPath, modified)
|
||||
applySeedMatrix(&seedFile, matrix)
|
||||
if comment := strings.TrimSpace(params.FileComments[requestedPath]); comment != "" {
|
||||
seedFile.Comment = comment
|
||||
} else if comment := strings.TrimSpace(params.FileComments[obj.GetName()]); comment != "" {
|
||||
seedFile.Comment = comment
|
||||
}
|
||||
if useGlobalDirect || directSet[requestedPath] {
|
||||
baseURL := strings.TrimRight(setting.GetStr(conf.SeedSiteURL), "/")
|
||||
seedFile.Sources = []torrent.SeedSource{{Type: "openlist-direct", URL: baseURL + utils.EncodePath("/d"+fullPath)}}
|
||||
}
|
||||
seed.Files = append(seed.Files, seedFile)
|
||||
fullPaths = append(fullPaths, fullPath)
|
||||
}
|
||||
globalHasher.Finish()
|
||||
createdShares := make([]string, 0, len(seed.Files))
|
||||
keepCreatedShares := false
|
||||
defer func() {
|
||||
if !keepCreatedShares {
|
||||
for _, createdID := range createdShares {
|
||||
_ = op.DeleteSharing(createdID)
|
||||
}
|
||||
}
|
||||
}()
|
||||
if hasShare {
|
||||
for index, fullPath := range fullPaths {
|
||||
if !useGlobalShare && !shareSet[params.Paths[index]] {
|
||||
continue
|
||||
}
|
||||
sharing := &model.Sharing{
|
||||
SharingDB: &model.SharingDB{Remark: "Transfer seed source"},
|
||||
Files: []string{fullPath}, Creator: user,
|
||||
}
|
||||
shareID, createErr := op.CreateSharing(sharing)
|
||||
if createErr != nil {
|
||||
return nil, nil, fmt.Errorf("create seed share: %w", createErr)
|
||||
}
|
||||
createdShares = append(createdShares, shareID)
|
||||
seed.Files[index].Sources = []torrent.SeedSource{{
|
||||
Type: "openlist-share", URL: strings.TrimRight(setting.GetStr(conf.SeedSiteURL), "/") + "/sd/" + shareID,
|
||||
ShareID: shareID,
|
||||
}}
|
||||
}
|
||||
}
|
||||
if err := torrent.ValidateSeed(seed, torrent.DefaultParseLimits()); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
outputPath := strings.TrimSpace(params.OutputPath)
|
||||
artifacts := make([]SeedArtifact, 0, len(formats))
|
||||
seenFormats := make(map[string]struct{}, len(formats))
|
||||
|
||||
// writeArtifact persists the encoded container into the destination folder
|
||||
// when an output path is configured, returning the fully-populated artifact.
|
||||
writeArtifact := func(format, fileName string, data []byte) (SeedArtifact, error) {
|
||||
artifact := SeedArtifact{
|
||||
Format: format,
|
||||
Name: fileName,
|
||||
FileName: fileName,
|
||||
SeedData: base64.StdEncoding.EncodeToString(data),
|
||||
Size: len(data),
|
||||
}
|
||||
if outputPath == "" {
|
||||
return artifact, nil
|
||||
}
|
||||
dstDir, err := user.JoinPath(outputPath)
|
||||
if err != nil {
|
||||
return artifact, err
|
||||
}
|
||||
meta, metaErr := op.GetNearestMeta(dstDir)
|
||||
if metaErr != nil && !errors.Is(errors.Cause(metaErr), errs.MetaNotFound) {
|
||||
return artifact, metaErr
|
||||
}
|
||||
if (!user.CanWriteContent() && !common.CanWriteContentBypassUserPerms(meta, dstDir)) || !common.CanWrite(user, meta, dstDir) {
|
||||
return artifact, errs.PermissionDenied
|
||||
}
|
||||
fileStream := &stream.FileStream{
|
||||
Ctx: ctx,
|
||||
Obj: &model.Object{Name: fileName, Size: int64(len(data)), Modified: time.Now()},
|
||||
Reader: bytes.NewReader(data), Mimetype: "application/octet-stream",
|
||||
}
|
||||
if err = PutDirectly(ctx, dstDir, fileStream); err != nil {
|
||||
return artifact, err
|
||||
}
|
||||
artifact.Path = stdpath.Join(outputPath, fileName)
|
||||
return artifact, nil
|
||||
}
|
||||
|
||||
for _, requestedFormat := range formats {
|
||||
format := strings.ToLower(strings.TrimPrefix(strings.TrimSpace(requestedFormat), "."))
|
||||
if format == "bt" {
|
||||
format = "torrent"
|
||||
}
|
||||
if _, exists := seenFormats[format]; exists {
|
||||
continue
|
||||
}
|
||||
seenFormats[format] = struct{}{}
|
||||
|
||||
data, err := EncodeGeneratedSeed(seed, format, globalHasher.GetPieceHashes())
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("generate %s seed: %w", format, err)
|
||||
}
|
||||
fileName := stdpath.Base(seed.Name) + "." + format
|
||||
artifact, err := writeArtifact(format, fileName, data)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
artifacts = append(artifacts, artifact)
|
||||
}
|
||||
keepCreatedShares = true
|
||||
return artifacts, seed, nil
|
||||
}
|
||||
|
||||
// SeedGenerateTask generates seed containers asynchronously, writing them into
|
||||
// params.OutputPath so they appear in the target folder once done.
|
||||
type SeedGenerateTask struct {
|
||||
task.TaskExtension
|
||||
params SeedGenerateParams
|
||||
}
|
||||
|
||||
func (t *SeedGenerateTask) GetName() string {
|
||||
if len(t.params.Paths) == 1 {
|
||||
return fmt.Sprintf("generate seed for %s", stdpath.Base(t.params.Paths[0]))
|
||||
}
|
||||
return fmt.Sprintf("generate seed for %d files", len(t.params.Paths))
|
||||
}
|
||||
|
||||
func (t *SeedGenerateTask) GetStatus() string {
|
||||
return "generating seed"
|
||||
}
|
||||
|
||||
func (t *SeedGenerateTask) Run() error {
|
||||
t.ClearEndTime()
|
||||
t.SetStartTime(time.Now())
|
||||
defer func() { t.SetEndTime(time.Now()) }()
|
||||
_, _, err := GenerateSeedArtifacts(t.Ctx(), t.Creator, t.params)
|
||||
return err
|
||||
}
|
||||
|
||||
var SeedGenerateTaskManager *tache.Manager[*SeedGenerateTask]
|
||||
|
||||
// AddSeedGenerateTask schedules an asynchronous seed generation task.
|
||||
func AddSeedGenerateTask(ctx context.Context, user *model.User, params SeedGenerateParams) (task.TaskExtensionInfo, error) {
|
||||
t := &SeedGenerateTask{
|
||||
TaskExtension: task.TaskExtension{
|
||||
Creator: user,
|
||||
ApiUrl: common.GetApiUrl(ctx),
|
||||
},
|
||||
params: params,
|
||||
}
|
||||
SeedGenerateTaskManager.Add(t)
|
||||
return t, nil
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -19,6 +19,7 @@ type Storage struct {
|
||||
Disabled bool `json:"disabled"` // if disabled
|
||||
DisableIndex bool `json:"disable_index"`
|
||||
EnableSign bool `json:"enable_sign"`
|
||||
SeedPolicy string `json:"seed_policy" gorm:"default:inherit"`
|
||||
Sort
|
||||
Proxy
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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() }
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -173,6 +173,14 @@ func getMainItems(config driver.Config) []driver.Item {
|
||||
Default: "false",
|
||||
Required: true,
|
||||
})
|
||||
items = append(items, driver.Item{
|
||||
Name: "seed_policy",
|
||||
Type: conf.TypeSelect,
|
||||
Options: "inherit,on,off",
|
||||
Default: "inherit",
|
||||
Required: true,
|
||||
Help: "Override automatic transfer-seed generation for this storage",
|
||||
})
|
||||
return items
|
||||
}
|
||||
func getAdditionalItems(t reflect.Type, defaultRoot string) []driver.Item {
|
||||
|
||||
+1
-4
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+32
-13
@@ -145,15 +145,24 @@ func bencodeEncodeOrderedDict(w io.Writer, d OrderedDict) error {
|
||||
|
||||
// BencodeDecode 从字节数组解码 bencode 数据
|
||||
func BencodeDecode(data []byte) (interface{}, error) {
|
||||
if int64(len(data)) > DefaultMaxSeedSize {
|
||||
return nil, fmt.Errorf("bencode: input exceeds %d bytes", DefaultMaxSeedSize)
|
||||
}
|
||||
reader := bytes.NewReader(data)
|
||||
val, err := bencodeDecodeValue(reader)
|
||||
val, err := bencodeDecodeValue(reader, 0)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if reader.Len() != 0 {
|
||||
return nil, fmt.Errorf("bencode: trailing data")
|
||||
}
|
||||
return val, nil
|
||||
}
|
||||
|
||||
func bencodeDecodeValue(r *bytes.Reader) (interface{}, error) {
|
||||
func bencodeDecodeValue(r *bytes.Reader, depth int) (interface{}, error) {
|
||||
if depth > DefaultParseLimits().MaxDepth {
|
||||
return nil, fmt.Errorf("bencode: nesting depth exceeds limit")
|
||||
}
|
||||
b, err := r.ReadByte()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -163,9 +172,9 @@ func bencodeDecodeValue(r *bytes.Reader) (interface{}, error) {
|
||||
case b == 'i':
|
||||
return bencodeDecodeInt(r)
|
||||
case b == 'l':
|
||||
return bencodeDecodeList(r)
|
||||
return bencodeDecodeList(r, depth+1)
|
||||
case b == 'd':
|
||||
return bencodeDecodeDict(r)
|
||||
return bencodeDecodeDict(r, depth+1)
|
||||
case b >= '0' && b <= '9':
|
||||
r.UnreadByte()
|
||||
return bencodeDecodeString(r)
|
||||
@@ -206,10 +215,14 @@ func bencodeDecodeString(r *bytes.Reader) ([]byte, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bencode: invalid string length: %v", err)
|
||||
}
|
||||
if length < 0 || length > 100*1024*1024 {
|
||||
return nil, fmt.Errorf("bencode: string length out of bounds: %d", length)
|
||||
// A single string can never exceed the whole input, which BencodeDecode
|
||||
// already caps at DefaultMaxSeedSize. Deriving the bound from the same
|
||||
// constant keeps the constraint self-consistent instead of maintaining a
|
||||
// second, unreachable 100MB ceiling.
|
||||
if length < 0 || length > DefaultMaxSeedSize {
|
||||
return nil, fmt.Errorf("bencode: string length out of bounds: %d (limit %d)", length, DefaultMaxSeedSize)
|
||||
}
|
||||
// Safe to convert to int: bounds check above ensures length <= 100MB which fits in int32
|
||||
// Bounded by DefaultMaxSeedSize, so the int conversion cannot truncate.
|
||||
data := make([]byte, int(length))
|
||||
_, err = io.ReadFull(r, data)
|
||||
if err != nil {
|
||||
@@ -218,9 +231,12 @@ func bencodeDecodeString(r *bytes.Reader) ([]byte, error) {
|
||||
return data, nil
|
||||
}
|
||||
|
||||
func bencodeDecodeList(r *bytes.Reader) ([]interface{}, error) {
|
||||
func bencodeDecodeList(r *bytes.Reader, depth int) ([]interface{}, error) {
|
||||
var list []interface{}
|
||||
for {
|
||||
if len(list) >= DefaultMaxSeedFiles*4 {
|
||||
return nil, fmt.Errorf("bencode: list item limit exceeded")
|
||||
}
|
||||
b, err := r.ReadByte()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -228,8 +244,8 @@ func bencodeDecodeList(r *bytes.Reader) ([]interface{}, error) {
|
||||
if b == 'e' {
|
||||
return list, nil
|
||||
}
|
||||
r.UnreadByte()
|
||||
val, err := bencodeDecodeValue(r)
|
||||
_ = r.UnreadByte()
|
||||
val, err := bencodeDecodeValue(r, depth)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -237,9 +253,12 @@ func bencodeDecodeList(r *bytes.Reader) ([]interface{}, error) {
|
||||
}
|
||||
}
|
||||
|
||||
func bencodeDecodeDict(r *bytes.Reader) (map[string]interface{}, error) {
|
||||
func bencodeDecodeDict(r *bytes.Reader, depth int) (map[string]interface{}, error) {
|
||||
dict := make(map[string]interface{})
|
||||
for {
|
||||
if len(dict) >= DefaultMaxSeedFiles*4 {
|
||||
return nil, fmt.Errorf("bencode: dictionary item limit exceeded")
|
||||
}
|
||||
b, err := r.ReadByte()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -247,12 +266,12 @@ func bencodeDecodeDict(r *bytes.Reader) (map[string]interface{}, error) {
|
||||
if b == 'e' {
|
||||
return dict, nil
|
||||
}
|
||||
r.UnreadByte()
|
||||
_ = r.UnreadByte()
|
||||
keyBytes, err := bencodeDecodeString(r)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
val, err := bencodeDecodeValue(r)
|
||||
val, err := bencodeDecodeValue(r, depth)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
+31
-10
@@ -1,9 +1,10 @@
|
||||
package torrent
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
"path"
|
||||
)
|
||||
|
||||
// GenerateFromFile 从文件路径生成通用的 torrent 文件(不含 CAS 扩展)
|
||||
@@ -30,7 +31,7 @@ func GenerateFromReader(reader io.Reader, fileName string, fileSize int64, piece
|
||||
pieceSize = DefaultPieceSize
|
||||
}
|
||||
|
||||
hw := NewHashWriter(pieceSize, pieceSize)
|
||||
hw := NewHashWriter(pieceSize, pieceSize, fileSize)
|
||||
|
||||
buf := make([]byte, 32*1024)
|
||||
for {
|
||||
@@ -64,7 +65,7 @@ func GenerateFromReaderWithCAS(reader io.Reader, fileName string, fileSize int64
|
||||
pieceSize = DefaultPieceSize
|
||||
}
|
||||
|
||||
hw := NewHashWriter(pieceSize, pieceSize)
|
||||
hw := NewHashWriter(pieceSize, pieceSize, fileSize)
|
||||
|
||||
buf := make([]byte, 32*1024)
|
||||
for {
|
||||
@@ -85,12 +86,8 @@ func GenerateFromReaderWithCAS(reader io.Reader, fileName string, fileSize int64
|
||||
sliceMD5s := hw.GetSliceMD5s()
|
||||
pieceHashes := hw.GetPieceHashes()
|
||||
|
||||
// 计算 sliceMD5
|
||||
sliceMD5 := fileMD5
|
||||
if len(sliceMD5s) > 1 {
|
||||
joined := strings.Join(sliceMD5s, "\n")
|
||||
sliceMD5 = strings.ToUpper(GetMD5Str(joined))
|
||||
}
|
||||
// 计算 sliceMD5(统一走规范实现)
|
||||
sliceMD5 := SliceMD5FromPieces(sliceMD5s, fileMD5)
|
||||
|
||||
t := NewTorrent(fileName, fileSize, fileMD5)
|
||||
t.Info.PieceLength = pieceSize
|
||||
@@ -100,7 +97,7 @@ func GenerateFromReaderWithCAS(reader io.Reader, fileName string, fileSize int64
|
||||
SliceMD5: sliceMD5,
|
||||
SliceMD5s: sliceMD5s,
|
||||
SliceSize: pieceSize,
|
||||
Cloud: "189",
|
||||
Cloud: Cloud189,
|
||||
})
|
||||
|
||||
return t.Encode()
|
||||
@@ -121,3 +118,27 @@ func GenerateFromFileWithCAS(filePath string) ([]byte, error) {
|
||||
|
||||
return GenerateFromReaderWithCAS(f, info.Name(), info.Size(), DefaultPieceSize)
|
||||
}
|
||||
|
||||
// GenerateSeedFromReader computes the complete OSS hash matrix in one stream pass.
|
||||
func GenerateSeedFromReader(reader io.Reader, filePath string, expectedSize, pieceSize int64, createdBy string) (*Seed, error) {
|
||||
if pieceSize <= 0 {
|
||||
pieceSize = DefaultPieceSize
|
||||
}
|
||||
if err := validateRelativeSeedPath(filePath); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
hw := NewHashWriter(pieceSize, pieceSize, expectedSize)
|
||||
if _, err := CopyAndHash(nil, reader, hw); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
hw.Finish()
|
||||
if expectedSize >= 0 && hw.GetTotalWritten() != expectedSize {
|
||||
return nil, fmt.Errorf("stream size mismatch: read %d bytes, expected %d", hw.GetTotalWritten(), expectedSize)
|
||||
}
|
||||
seed := NewSeed(path.Base(filePath), createdBy, pieceSize)
|
||||
seed.Files = []SeedFile{hw.BuildSeedFile(filePath, "")}
|
||||
if err := ValidateSeed(seed, DefaultParseLimits()); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return seed, nil
|
||||
}
|
||||
|
||||
+129
-22
@@ -3,11 +3,14 @@ package torrent
|
||||
import (
|
||||
"crypto/md5"
|
||||
"crypto/sha1"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"hash"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
hash_extend "github.com/OpenListTeam/OpenList/v4/pkg/utils/hash"
|
||||
)
|
||||
|
||||
// HashWriter 同时计算文件的 MD5、分片 MD5 和 SHA-1 piece hash
|
||||
@@ -15,15 +18,24 @@ import (
|
||||
type HashWriter struct {
|
||||
// 整文件 MD5
|
||||
fileMD5 hash.Hash
|
||||
// fileSHA1 and fileSHA256 complete the portable full-file matrix.
|
||||
fileSHA1 hash.Hash
|
||||
fileSHA256 hash.Hash
|
||||
// fileGCID 用于迅雷、PikPak 等
|
||||
fileGCID hash.Hash
|
||||
// 当前分片 MD5
|
||||
sliceMD5 hash.Hash
|
||||
// 当前 piece 的 SHA-1
|
||||
pieceSHA1 hash.Hash
|
||||
// Per-piece hashers are updated in the same pass as whole-file hashes.
|
||||
pieceMD5 hash.Hash
|
||||
pieceSHA1 hash.Hash
|
||||
pieceSHA256 hash.Hash
|
||||
|
||||
// 分片大小(默认 10MB)
|
||||
sliceSize int64
|
||||
// piece 大小(与 sliceSize 相同,保持对齐)
|
||||
pieceSize int64
|
||||
// 文件总大小(用于 GCID 初始化)
|
||||
fileSize int64
|
||||
|
||||
// 当前分片已写入字节数
|
||||
sliceWritten int64
|
||||
@@ -34,14 +46,19 @@ type HashWriter struct {
|
||||
|
||||
// 每个分片的 MD5(大写十六进制)
|
||||
sliceMD5Hexs []string
|
||||
// 所有 piece 的 SHA-1 哈希拼接
|
||||
// all standard BitTorrent SHA-1 piece hashes concatenated
|
||||
pieceHashes []byte
|
||||
// portable per-file piece matrix
|
||||
pieceMD5Hexs []string
|
||||
pieceSHA1Hexs []string
|
||||
pieceSHA256Hexs []string
|
||||
}
|
||||
|
||||
// NewHashWriter 创建一个新的 HashWriter
|
||||
// sliceSize: CAS 分片大小(通常 10MB)
|
||||
// pieceSize: BT piece 大小(设为与 sliceSize 相同以保持对齐)
|
||||
func NewHashWriter(sliceSize, pieceSize int64) *HashWriter {
|
||||
// fileSize: 文件总大小(用于 GCID 初始化,0 表示未知)
|
||||
func NewHashWriter(sliceSize, pieceSize, fileSize int64) *HashWriter {
|
||||
if sliceSize <= 0 {
|
||||
sliceSize = DefaultPieceSize
|
||||
}
|
||||
@@ -49,17 +66,23 @@ func NewHashWriter(sliceSize, pieceSize int64) *HashWriter {
|
||||
pieceSize = DefaultPieceSize
|
||||
}
|
||||
return &HashWriter{
|
||||
fileMD5: md5.New(),
|
||||
sliceMD5: md5.New(),
|
||||
pieceSHA1: sha1.New(),
|
||||
sliceSize: sliceSize,
|
||||
pieceSize: pieceSize,
|
||||
fileMD5: md5.New(),
|
||||
fileSHA1: sha1.New(),
|
||||
fileSHA256: sha256.New(),
|
||||
fileGCID: hash_extend.GCID.NewFunc(fileSize),
|
||||
sliceMD5: md5.New(),
|
||||
pieceMD5: md5.New(),
|
||||
pieceSHA1: sha1.New(),
|
||||
pieceSHA256: sha256.New(),
|
||||
sliceSize: sliceSize,
|
||||
pieceSize: pieceSize,
|
||||
fileSize: fileSize,
|
||||
}
|
||||
}
|
||||
|
||||
// NewDefaultHashWriter 创建默认的 HashWriter(10MB 分片)
|
||||
func NewDefaultHashWriter() *HashWriter {
|
||||
return NewHashWriter(DefaultPieceSize, DefaultPieceSize)
|
||||
return NewHashWriter(DefaultPieceSize, DefaultPieceSize, 0)
|
||||
}
|
||||
|
||||
// Write 实现 io.Writer 接口
|
||||
@@ -76,12 +99,15 @@ func (hw *HashWriter) Write(p []byte) (n int, err error) {
|
||||
|
||||
chunk := p[offset : offset+int(canWrite)]
|
||||
|
||||
// 写入整文件 MD5
|
||||
hw.fileMD5.Write(chunk)
|
||||
// 写入当前分片 MD5
|
||||
hw.sliceMD5.Write(chunk)
|
||||
// 写入当前 piece SHA-1
|
||||
hw.pieceSHA1.Write(chunk)
|
||||
// Write all whole-file and boundary-specific hashes in one pass.
|
||||
_, _ = hw.fileMD5.Write(chunk)
|
||||
_, _ = hw.fileSHA1.Write(chunk)
|
||||
_, _ = hw.fileSHA256.Write(chunk)
|
||||
_, _ = hw.fileGCID.Write(chunk)
|
||||
_, _ = hw.sliceMD5.Write(chunk)
|
||||
_, _ = hw.pieceMD5.Write(chunk)
|
||||
_, _ = hw.pieceSHA1.Write(chunk)
|
||||
_, _ = hw.pieceSHA256.Write(chunk)
|
||||
|
||||
hw.sliceWritten += canWrite
|
||||
hw.pieceWritten += canWrite
|
||||
@@ -112,8 +138,16 @@ func (hw *HashWriter) finishSlice() {
|
||||
|
||||
// finishPiece 完成当前 piece 的 SHA-1 计算
|
||||
func (hw *HashWriter) finishPiece() {
|
||||
hw.pieceHashes = append(hw.pieceHashes, hw.pieceSHA1.Sum(nil)...)
|
||||
md5Sum := hw.pieceMD5.Sum(nil)
|
||||
sha1Sum := hw.pieceSHA1.Sum(nil)
|
||||
sha256Sum := hw.pieceSHA256.Sum(nil)
|
||||
hw.pieceMD5Hexs = append(hw.pieceMD5Hexs, hex.EncodeToString(md5Sum))
|
||||
hw.pieceSHA1Hexs = append(hw.pieceSHA1Hexs, hex.EncodeToString(sha1Sum))
|
||||
hw.pieceSHA256Hexs = append(hw.pieceSHA256Hexs, hex.EncodeToString(sha256Sum))
|
||||
hw.pieceHashes = append(hw.pieceHashes, sha1Sum...)
|
||||
hw.pieceMD5.Reset()
|
||||
hw.pieceSHA1.Reset()
|
||||
hw.pieceSHA256.Reset()
|
||||
hw.pieceWritten = 0
|
||||
}
|
||||
|
||||
@@ -134,6 +168,56 @@ func (hw *HashWriter) GetFileMD5() string {
|
||||
return strings.ToUpper(hex.EncodeToString(hw.fileMD5.Sum(nil)))
|
||||
}
|
||||
|
||||
// GetFileSHA1 returns the lowercase whole-file SHA-1 digest.
|
||||
func (hw *HashWriter) GetFileSHA1() string {
|
||||
return hex.EncodeToString(hw.fileSHA1.Sum(nil))
|
||||
}
|
||||
|
||||
// GetFileSHA256 returns the lowercase whole-file SHA-256 digest.
|
||||
func (hw *HashWriter) GetFileSHA256() string {
|
||||
return hex.EncodeToString(hw.fileSHA256.Sum(nil))
|
||||
}
|
||||
|
||||
// GetFileGCID returns the uppercase GCID digest for Thunder/PikPak.
|
||||
func (hw *HashWriter) GetFileGCID() string {
|
||||
return strings.ToUpper(hex.EncodeToString(hw.fileGCID.Sum(nil)))
|
||||
}
|
||||
|
||||
// GetPieceMD5s returns independent per-file MD5 piece hashes.
|
||||
func (hw *HashWriter) GetPieceMD5s() []string {
|
||||
return append([]string(nil), hw.pieceMD5Hexs...)
|
||||
}
|
||||
|
||||
// GetPieceSHA1s returns independent per-file SHA-1 piece hashes.
|
||||
func (hw *HashWriter) GetPieceSHA1s() []string {
|
||||
return append([]string(nil), hw.pieceSHA1Hexs...)
|
||||
}
|
||||
|
||||
// GetPieceSHA256s returns independent per-file SHA-256 piece hashes.
|
||||
func (hw *HashWriter) GetPieceSHA256s() []string {
|
||||
return append([]string(nil), hw.pieceSHA256Hexs...)
|
||||
}
|
||||
|
||||
// BuildSeedFile exports all hashes accumulated during this single stream pass.
|
||||
func (hw *HashWriter) BuildSeedFile(filePath string, modified string) SeedFile {
|
||||
return SeedFile{
|
||||
Path: filePath,
|
||||
Size: hw.totalWritten,
|
||||
Modified: modified,
|
||||
Hashes: SeedHashes{
|
||||
MD5: strings.ToLower(hw.GetFileMD5()),
|
||||
SHA1: hw.GetFileSHA1(),
|
||||
SHA256: hw.GetFileSHA256(),
|
||||
GCID: strings.ToLower(hw.GetFileGCID()),
|
||||
Pieces: &SeedPieceHashes{
|
||||
MD5: hw.GetPieceMD5s(),
|
||||
SHA1: hw.GetPieceSHA1s(),
|
||||
SHA256: hw.GetPieceSHA256s(),
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// GetSliceMD5s 获取所有分片的 MD5 列表
|
||||
func (hw *HashWriter) GetSliceMD5s() []string {
|
||||
return hw.sliceMD5Hexs
|
||||
@@ -141,11 +225,34 @@ func (hw *HashWriter) GetSliceMD5s() []string {
|
||||
|
||||
// GetSliceMD5 获取最终的 sliceMD5(用于秒传)
|
||||
func (hw *HashWriter) GetSliceMD5(fileMD5 string) string {
|
||||
if len(hw.sliceMD5Hexs) <= 1 {
|
||||
return fileMD5
|
||||
return SliceMD5FromPieces(hw.sliceMD5Hexs, fileMD5)
|
||||
}
|
||||
|
||||
// SliceMD5FromPieces is the single canonical implementation of the sliceMd5
|
||||
// rule shared by every CAS producer and consumer:
|
||||
//
|
||||
// - no piece, or a single piece -> the whole-file MD5
|
||||
// - two or more pieces -> MD5 of the piece MD5s joined by "\n"
|
||||
//
|
||||
// Keeping one implementation matters because this value is what the remote
|
||||
// provider compares against: a divergence between the hash-generation side and
|
||||
// the torrent/CAS encoding side silently turns rapid uploads into mismatches.
|
||||
// All comparisons and the returned value are upper-case.
|
||||
func SliceMD5FromPieces(sliceMD5s []string, fileMD5 string) string {
|
||||
switch len(sliceMD5s) {
|
||||
case 0, 1:
|
||||
// A single piece covers the whole file, so the two hashes coincide.
|
||||
if len(sliceMD5s) == 1 && sliceMD5s[0] != "" {
|
||||
return strings.ToUpper(sliceMD5s[0])
|
||||
}
|
||||
return strings.ToUpper(fileMD5)
|
||||
default:
|
||||
upper := make([]string, len(sliceMD5s))
|
||||
for i, piece := range sliceMD5s {
|
||||
upper[i] = strings.ToUpper(piece)
|
||||
}
|
||||
return strings.ToUpper(GetMD5Str(strings.Join(upper, "\n")))
|
||||
}
|
||||
joined := strings.Join(hw.sliceMD5Hexs, "\n")
|
||||
return strings.ToUpper(GetMD5Str(joined))
|
||||
}
|
||||
|
||||
// GetPieceHashes 获取所有 piece 的 SHA-1 哈希拼接
|
||||
@@ -170,7 +277,7 @@ func (hw *HashWriter) BuildTorrent(fileName string, fileSize int64) *Torrent {
|
||||
SliceMD5: sliceMD5,
|
||||
SliceMD5s: hw.GetSliceMD5s(),
|
||||
SliceSize: hw.sliceSize,
|
||||
Cloud: "189",
|
||||
Cloud: Cloud189,
|
||||
})
|
||||
|
||||
return t
|
||||
|
||||
@@ -0,0 +1,156 @@
|
||||
package torrent
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// --- path traversal / malformed path handling ---------------------------------
|
||||
|
||||
func TestValidateSeedRejectsUnsafePaths(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
path string
|
||||
}{
|
||||
{"parent traversal", "../secret"},
|
||||
{"nested traversal", "a/../../secret"},
|
||||
{"absolute unix", "/etc/passwd"},
|
||||
{"empty", ""},
|
||||
{"current dir", "."},
|
||||
{"nul byte", "a\x00b"},
|
||||
{"backslash traversal", `..\secret`},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
seed := testSeed()
|
||||
seed.Files[0].Path = tc.path
|
||||
if err := ValidateSeed(seed, DefaultParseLimits()); err == nil {
|
||||
t.Fatalf("ValidateSeed() accepted unsafe path %q", tc.path)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateSeedRejectsTooManyFiles(t *testing.T) {
|
||||
seed := testSeed()
|
||||
seed.Files = make([]SeedFile, DefaultMaxSeedFiles+1)
|
||||
for i := range seed.Files {
|
||||
seed.Files[i] = SeedFile{
|
||||
Path: "f" + strings.Repeat("0", i%3) + ".bin",
|
||||
Size: 1,
|
||||
Hashes: SeedHashes{MD5: strings.Repeat("1", 32)},
|
||||
}
|
||||
}
|
||||
if err := ValidateSeed(seed, DefaultParseLimits()); err == nil {
|
||||
t.Fatal("ValidateSeed() accepted more files than DefaultMaxSeedFiles")
|
||||
}
|
||||
}
|
||||
|
||||
// --- sliceMd5 canonical rule ---------------------------------------------------
|
||||
|
||||
func TestSliceMD5FromPiecesMatchesSpec(t *testing.T) {
|
||||
fileMD5 := strings.Repeat("a", 32)
|
||||
|
||||
// Zero pieces: fall back to the whole-file MD5.
|
||||
if got := SliceMD5FromPieces(nil, fileMD5); got != strings.ToUpper(fileMD5) {
|
||||
t.Fatalf("SliceMD5FromPieces(nil) = %q, want %q", got, strings.ToUpper(fileMD5))
|
||||
}
|
||||
// A single piece covers the whole file, so it equals the file MD5.
|
||||
single := strings.ToUpper(fileMD5)
|
||||
if got := SliceMD5FromPieces([]string{single}, fileMD5); got != single {
|
||||
t.Fatalf("SliceMD5FromPieces(single) = %q, want %q", got, single)
|
||||
}
|
||||
// Two or more pieces: MD5 of the newline-joined, upper-cased piece list.
|
||||
pieces := []string{strings.Repeat("b", 32), strings.Repeat("c", 32)}
|
||||
want := strings.ToUpper(GetMD5Str(strings.Join([]string{strings.Repeat("B", 32), strings.Repeat("C", 32)}, "\n")))
|
||||
if got := SliceMD5FromPieces(pieces, fileMD5); got != want {
|
||||
t.Fatalf("SliceMD5FromPieces(multi) = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSliceMD5AgreesWithBuildCASInfo locks the rule shared by the hash-generation
|
||||
// side and the CAS-encoding side. A divergence here silently turns rapid uploads
|
||||
// into hash mismatches, so the two entry points must never drift apart.
|
||||
func TestSliceMD5AgreesWithBuildCASInfo(t *testing.T) {
|
||||
fileMD5 := strings.Repeat("a", 32)
|
||||
sets := [][]string{
|
||||
nil,
|
||||
{strings.Repeat("b", 32)},
|
||||
{strings.Repeat("b", 32), strings.Repeat("c", 32)},
|
||||
{strings.Repeat("b", 32), strings.Repeat("c", 32), strings.Repeat("d", 32)},
|
||||
}
|
||||
for i, pieces := range sets {
|
||||
hw := &HashWriter{sliceMD5Hexs: pieces}
|
||||
fromWriter := hw.GetSliceMD5(fileMD5)
|
||||
fromCAS := BuildCASInfoFromMD5s(fileMD5, pieces, DefaultPieceSize).SliceMD5
|
||||
if fromWriter != fromCAS {
|
||||
t.Fatalf("case %d: GetSliceMD5() = %q, BuildCASInfoFromMD5s() = %q", i, fromWriter, fromCAS)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- bencode robustness --------------------------------------------------------
|
||||
|
||||
func TestBencodeDecodeRejectsOversizedStringLength(t *testing.T) {
|
||||
// Declares a 4GiB string while the buffer is empty; parsing must fail on the
|
||||
// declared length instead of attempting a huge allocation.
|
||||
payload := []byte("9999999999:")
|
||||
if _, err := BencodeDecode(payload); err == nil {
|
||||
t.Fatal("BencodeDecode() accepted an out-of-bounds string length")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBencodeDecodeRejectsDeepNesting(t *testing.T) {
|
||||
depth := DefaultParseLimits().MaxDepth + 2
|
||||
payload := strings.Repeat("l", depth) + strings.Repeat("e", depth)
|
||||
if _, err := BencodeDecode([]byte(payload)); err == nil {
|
||||
t.Fatal("BencodeDecode() accepted nesting beyond MaxDepth")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBencodeDecodeRejectsTrailingData(t *testing.T) {
|
||||
if _, err := BencodeDecode([]byte("i1eextra")); err == nil {
|
||||
t.Fatal("BencodeDecode() accepted trailing data")
|
||||
}
|
||||
}
|
||||
|
||||
// --- cross-format conversion consistency --------------------------------------
|
||||
|
||||
// TestConvertConsistencyAcrossFormats ensures a seed survives OSS -> torrent ->
|
||||
// OSS and OSS -> CAS -> OSS without losing whole-file hashes.
|
||||
func TestConvertConsistencyAcrossFormats(t *testing.T) {
|
||||
original := testSeed()
|
||||
|
||||
torrentData, err := EncodeSeed(original, "torrent")
|
||||
if err != nil {
|
||||
t.Fatalf("EncodeSeed(torrent) error = %v", err)
|
||||
}
|
||||
fromTorrent, err := DecodeSeed(torrentData, "torrent", DefaultParseLimits())
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeSeed(torrent) error = %v", err)
|
||||
}
|
||||
casData, err := EncodeCAS(fromTorrent)
|
||||
if err != nil {
|
||||
t.Fatalf("EncodeCAS() error = %v", err)
|
||||
}
|
||||
fromCAS, err := DecodeCAS(casData, DefaultParseLimits())
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeCAS() error = %v", err)
|
||||
}
|
||||
if got := fromCAS.Files[0].Hashes.MD5; got != strings.ToUpper(original.Files[0].Hashes.MD5) {
|
||||
t.Fatalf("MD5 changed across formats: %q", got)
|
||||
}
|
||||
if len(fromCAS.Files) != len(original.Files) {
|
||||
t.Fatalf("file count changed across formats: %d", len(fromCAS.Files))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSeedRejectsUnknownFormat(t *testing.T) {
|
||||
data, err := EncodeOSS(testSeed())
|
||||
if err != nil {
|
||||
t.Fatalf("EncodeOSS() error = %v", err)
|
||||
}
|
||||
if _, err := DecodeSeed(data, "does-not-exist", DefaultParseLimits()); err == nil {
|
||||
t.Fatal("DecodeSeed() accepted an unknown format")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,160 @@
|
||||
package torrent
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func testSeed() *Seed {
|
||||
return &Seed{
|
||||
Format: OSSFormat,
|
||||
Version: OSSVersion,
|
||||
Name: "example.bin",
|
||||
CreatedAt: time.Unix(1, 0).UTC().Format(time.RFC3339),
|
||||
CreatedBy: "OpenList",
|
||||
PieceSize: DefaultPieceSize,
|
||||
Files: []SeedFile{{
|
||||
Path: "example.bin",
|
||||
Size: DefaultPieceSize + 1,
|
||||
Hashes: SeedHashes{
|
||||
MD5: strings.Repeat("1", 32),
|
||||
SHA1: strings.Repeat("2", 40),
|
||||
SHA256: strings.Repeat("3", 64),
|
||||
Pieces: &SeedPieceHashes{
|
||||
MD5: []string{strings.Repeat("4", 32), strings.Repeat("5", 32)},
|
||||
SHA1: []string{strings.Repeat("6", 40), strings.Repeat("7", 40)},
|
||||
SHA256: []string{strings.Repeat("8", 64), strings.Repeat("9", 64)},
|
||||
},
|
||||
},
|
||||
}},
|
||||
}
|
||||
}
|
||||
|
||||
func TestOSSRoundTrip(t *testing.T) {
|
||||
encoded, err := EncodeOSS(testSeed())
|
||||
if err != nil {
|
||||
t.Fatalf("EncodeOSS() error = %v", err)
|
||||
}
|
||||
decoded, err := DecodeOSS(encoded, DefaultParseLimits())
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeOSS() error = %v", err)
|
||||
}
|
||||
if decoded.Name != "example.bin" || len(decoded.Files) != 1 || decoded.Files[0].Hashes.SHA256 == "" {
|
||||
t.Fatalf("DecodeOSS() returned incomplete seed: %#v", decoded)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCASWireFormatIsLegacyCompatible(t *testing.T) {
|
||||
encoded, err := EncodeCAS(testSeed())
|
||||
if err != nil {
|
||||
t.Fatalf("EncodeCAS() error = %v", err)
|
||||
}
|
||||
decodedJSON, err := base64.StdEncoding.DecodeString(string(encoded))
|
||||
if err != nil {
|
||||
t.Fatalf("base64.DecodeString() error = %v", err)
|
||||
}
|
||||
var payload map[string]any
|
||||
if err = json.Unmarshal(decodedJSON, &payload); err != nil {
|
||||
t.Fatalf("json.Unmarshal() error = %v", err)
|
||||
}
|
||||
// The five legacy fields must always be present so the reference client can
|
||||
// parse the payload. slice_md5s / slice_size are optional extensions.
|
||||
for _, key := range []string{"name", "size", "md5", "sliceMd5", "create_time"} {
|
||||
if _, ok := payload[key]; !ok {
|
||||
t.Fatalf("CAS payload missing %q: %#v", key, payload)
|
||||
}
|
||||
}
|
||||
decoded, err := DecodeCAS(encoded, DefaultParseLimits())
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeCAS() error = %v", err)
|
||||
}
|
||||
if decoded.Files[0].Hashes.MD5 != strings.Repeat("1", 32) {
|
||||
t.Fatalf("DecodeCAS() MD5 = %q", decoded.Files[0].Hashes.MD5)
|
||||
}
|
||||
// Per-piece MD5 list must round-trip through the slice_md5s extension.
|
||||
if decoded.Files[0].Hashes.Pieces == nil || len(decoded.Files[0].Hashes.Pieces.MD5) != 2 {
|
||||
t.Fatalf("DecodeCAS() piece MD5 list = %#v", decoded.Files[0].Hashes.Pieces)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTorrentRoundTripPreservesOpenListExtension(t *testing.T) {
|
||||
encoded, err := EncodeSeed(testSeed(), "torrent")
|
||||
if err != nil {
|
||||
t.Fatalf("EncodeSeed(torrent) error = %v", err)
|
||||
}
|
||||
decoded, err := DecodeSeed(encoded, "torrent", DefaultParseLimits())
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeSeed(torrent) error = %v", err)
|
||||
}
|
||||
if decoded.Files[0].Hashes.SHA256 != strings.Repeat("3", 64) {
|
||||
t.Fatalf("torrent extension lost SHA-256: %#v", decoded.Files[0].Hashes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCASCloudGeneralizationRoundTrip(t *testing.T) {
|
||||
seed := testSeed()
|
||||
seed.Files[0].CASSliceMD5 = strings.Repeat("a", 32)
|
||||
seed.Files[0].CASCloud = CloudAliyundriveOpen
|
||||
|
||||
// CAS (base64 JSON) round-trip must preserve the cloud identifier.
|
||||
encoded, err := EncodeCAS(seed)
|
||||
if err != nil {
|
||||
t.Fatalf("EncodeCAS() error = %v", err)
|
||||
}
|
||||
decoded, err := DecodeCAS(encoded, DefaultParseLimits())
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeCAS() error = %v", err)
|
||||
}
|
||||
if got := decoded.Files[0].CASCloud; got != CloudAliyundriveOpen {
|
||||
t.Fatalf("DecodeCAS() cloud = %q, want %q", got, CloudAliyundriveOpen)
|
||||
}
|
||||
|
||||
// Torrent bencode round-trip must preserve the cloud identifier too.
|
||||
torrentData, err := EncodeSeed(seed, "torrent")
|
||||
if err != nil {
|
||||
t.Fatalf("EncodeSeed(torrent) error = %v", err)
|
||||
}
|
||||
decodedTorrent, err := DecodeSeed(torrentData, "torrent", DefaultParseLimits())
|
||||
if err != nil {
|
||||
t.Fatalf("DecodeSeed(torrent) error = %v", err)
|
||||
}
|
||||
if got := decodedTorrent.Files[0].CASCloud; got != CloudAliyundriveOpen {
|
||||
t.Fatalf("DecodeSeed(torrent) cloud = %q, want %q", got, CloudAliyundriveOpen)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildCASInfoFromMD5sDefaultsToCloud189(t *testing.T) {
|
||||
info := BuildCASInfoFromMD5s(strings.Repeat("1", 32), []string{strings.Repeat("4", 32)}, DefaultPieceSize)
|
||||
if info.Cloud != Cloud189 {
|
||||
t.Fatalf("BuildCASInfoFromMD5s() cloud = %q, want %q", info.Cloud, Cloud189)
|
||||
}
|
||||
other := BuildCASInfoFromMD5sWithCloud(strings.Repeat("1", 32), []string{strings.Repeat("4", 32)}, DefaultPieceSize, Cloud115)
|
||||
if other.Cloud != Cloud115 {
|
||||
t.Fatalf("BuildCASInfoFromMD5sWithCloud() cloud = %q, want %q", other.Cloud, Cloud115)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateSeedRejectsTraversalAndInvalidHash(t *testing.T) {
|
||||
seed := testSeed()
|
||||
seed.Files[0].Path = "../secret"
|
||||
if err := ValidateSeed(seed, DefaultParseLimits()); err == nil {
|
||||
t.Fatal("ValidateSeed() accepted path traversal")
|
||||
}
|
||||
seed = testSeed()
|
||||
seed.Files[0].Hashes.MD5 = "not-a-hash"
|
||||
if err := ValidateSeed(seed, DefaultParseLimits()); err == nil {
|
||||
t.Fatal("ValidateSeed() accepted an invalid hash")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeSeedHonorsSizeLimit(t *testing.T) {
|
||||
data := []byte(`{"format":"openlist-sharing-seed"}`)
|
||||
limits := DefaultParseLimits()
|
||||
limits.MaxBytes = int64(len(data) - 1)
|
||||
if _, err := DecodeSeed(data, "oss", limits); err == nil {
|
||||
t.Fatal("DecodeSeed() accepted input above MaxBytes")
|
||||
}
|
||||
}
|
||||
+1036
-9
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
+102
-17
@@ -63,6 +63,7 @@ type MoveCopyReq struct {
|
||||
Overwrite bool `json:"overwrite"`
|
||||
SkipExisting bool `json:"skip_existing"`
|
||||
Merge bool `json:"merge"`
|
||||
FollowSeed bool `json:"follow_seed"`
|
||||
}
|
||||
|
||||
// FsMove performs batch move (individual item permission checks skipped for performance).
|
||||
@@ -113,10 +114,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) {
|
||||
@@ -156,6 +153,14 @@ func FsMove(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if req.FollowSeed {
|
||||
seedTasks, followErr := followSeedTransfer(c, "move", p, dstDir)
|
||||
if followErr != nil {
|
||||
common.ErrorResp(c, followErr, 500)
|
||||
return
|
||||
}
|
||||
addedTasks = append(addedTasks, seedTasks...)
|
||||
}
|
||||
}
|
||||
|
||||
// Return immediately with task information
|
||||
@@ -220,10 +225,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) {
|
||||
@@ -269,6 +270,14 @@ func FsCopy(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if req.FollowSeed {
|
||||
seedTasks, followErr := followSeedTransfer(c, "copy", p, dstDir)
|
||||
if followErr != nil {
|
||||
common.ErrorResp(c, followErr, 500)
|
||||
return
|
||||
}
|
||||
addedTasks = append(addedTasks, seedTasks...)
|
||||
}
|
||||
}
|
||||
|
||||
// Return immediately with task information
|
||||
@@ -285,9 +294,10 @@ func FsCopy(c *gin.Context) {
|
||||
}
|
||||
|
||||
type RenameReq struct {
|
||||
Path string `json:"path"`
|
||||
Name string `json:"name"`
|
||||
Overwrite bool `json:"overwrite"`
|
||||
Path string `json:"path"`
|
||||
Name string `json:"name"`
|
||||
Overwrite bool `json:"overwrite"`
|
||||
FollowSeed bool `json:"follow_seed"`
|
||||
}
|
||||
|
||||
func FsRename(c *gin.Context) {
|
||||
@@ -332,6 +342,12 @@ func FsRename(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if req.FollowSeed {
|
||||
if err := followSeedRename(c, reqPath, req.Name); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
}
|
||||
common.SuccessResp(c)
|
||||
}
|
||||
|
||||
@@ -343,8 +359,9 @@ func checkRelativePath(path string) error {
|
||||
}
|
||||
|
||||
type RemoveReq struct {
|
||||
Dir string `json:"dir"`
|
||||
Names []string `json:"names"`
|
||||
Dir string `json:"dir"`
|
||||
Names []string `json:"names"`
|
||||
FollowSeed bool `json:"follow_seed"`
|
||||
}
|
||||
|
||||
// FsRemove performs batch remove (individual item permission checks skipped for performance).
|
||||
@@ -381,10 +398,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] = ""
|
||||
@@ -396,16 +409,88 @@ func FsRemove(c *gin.Context) {
|
||||
if path == "" {
|
||||
continue
|
||||
}
|
||||
source, _ := fs.Get(c.Request.Context(), path, &fs.GetArgs{NoLog: true})
|
||||
err := fs.Remove(c.Request.Context(), path)
|
||||
if err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if req.FollowSeed && source != nil && !source.IsDir() {
|
||||
if err = followSeedRemove(c, path); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
//fs.ClearCache(req.Dir)
|
||||
common.SuccessResp(c)
|
||||
}
|
||||
|
||||
func seedSidecarPaths(filePath string) []string {
|
||||
return []string{filePath + ".oss", filePath + ".torrent", filePath + ".cas", filePath + ".cas.torrent"}
|
||||
}
|
||||
|
||||
func followSeedTransfer(c *gin.Context, operation, srcPath, dstDir string) ([]task.TaskExtensionInfo, error) {
|
||||
source, err := fs.Get(c.Request.Context(), srcPath, &fs.GetArgs{NoLog: true})
|
||||
if err != nil || source == nil || source.IsDir() {
|
||||
return nil, nil
|
||||
}
|
||||
var tasks []task.TaskExtensionInfo
|
||||
for _, sidecarPath := range seedSidecarPaths(srcPath) {
|
||||
obj, err := fs.Get(c.Request.Context(), sidecarPath, &fs.GetArgs{NoLog: true})
|
||||
if err != nil || obj == nil || obj.IsDir() {
|
||||
continue
|
||||
}
|
||||
var current task.TaskExtensionInfo
|
||||
switch operation {
|
||||
case "copy":
|
||||
current, err = fs.Copy(c.Request.Context(), sidecarPath, dstDir, true)
|
||||
case "move":
|
||||
current, err = fs.Move(c.Request.Context(), sidecarPath, dstDir, true)
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported seed sidecar operation %q", operation)
|
||||
}
|
||||
if err != nil {
|
||||
return tasks, fmt.Errorf("%s seed sidecar %s: %w", operation, sidecarPath, err)
|
||||
}
|
||||
if current != nil {
|
||||
tasks = append(tasks, current)
|
||||
}
|
||||
}
|
||||
return tasks, nil
|
||||
}
|
||||
|
||||
func followSeedRename(c *gin.Context, srcPath, newName string) error {
|
||||
source, err := fs.Get(c.Request.Context(), srcPath, &fs.GetArgs{NoLog: true})
|
||||
if err != nil || source == nil || source.IsDir() {
|
||||
return nil
|
||||
}
|
||||
for _, sidecarPath := range seedSidecarPaths(srcPath) {
|
||||
obj, err := fs.Get(c.Request.Context(), sidecarPath, &fs.GetArgs{NoLog: true})
|
||||
if err != nil || obj == nil || obj.IsDir() {
|
||||
continue
|
||||
}
|
||||
suffix := strings.TrimPrefix(sidecarPath, srcPath)
|
||||
if err = fs.Rename(c.Request.Context(), sidecarPath, newName+suffix, true); err != nil {
|
||||
return fmt.Errorf("rename seed sidecar %s: %w", sidecarPath, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func followSeedRemove(c *gin.Context, srcPath string) error {
|
||||
for _, sidecarPath := range seedSidecarPaths(srcPath) {
|
||||
obj, err := fs.Get(c.Request.Context(), sidecarPath, &fs.GetArgs{NoLog: true})
|
||||
if err != nil || obj == nil || obj.IsDir() {
|
||||
continue
|
||||
}
|
||||
if err = fs.Remove(c.Request.Context(), sidecarPath); err != nil {
|
||||
return fmt.Errorf("remove seed sidecar %s: %w", sidecarPath, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type RemoveEmptyDirectoryReq struct {
|
||||
SrcDir string `json:"src_dir"`
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
},
|
||||
|
||||
+146
-2
@@ -1,19 +1,25 @@
|
||||
package handles
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/url"
|
||||
stdpath "path"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/conf"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/errs"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/fs"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/model"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/op"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/setting"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/stream"
|
||||
"github.com/OpenListTeam/OpenList/v4/internal/task"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/torrent"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
"github.com/OpenListTeam/OpenList/v4/server/common"
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -97,6 +103,17 @@ func FsStream(c *gin.Context) {
|
||||
if len(mimetype) == 0 {
|
||||
mimetype = utils.GetMimeType(name)
|
||||
}
|
||||
generateSeed := shouldGenerateUploadSeed(c, dir)
|
||||
if generateSeed && asTask {
|
||||
common.ErrorStrResp(c, "seed sidecar generation requires synchronous upload", 400)
|
||||
return
|
||||
}
|
||||
var seedHasher *torrent.HashWriter
|
||||
var uploadReader io.Reader = c.Request.Body
|
||||
if generateSeed {
|
||||
seedHasher = torrent.NewHashWriter(seedPieceSize(c), seedPieceSize(c), 0)
|
||||
uploadReader = io.TeeReader(c.Request.Body, seedHasher)
|
||||
}
|
||||
s := &stream.FileStream{
|
||||
Obj: &model.Object{
|
||||
Name: name,
|
||||
@@ -104,7 +121,7 @@ func FsStream(c *gin.Context) {
|
||||
Modified: getLastModified(c),
|
||||
HashInfo: utils.NewHashInfoByMap(h),
|
||||
},
|
||||
Reader: c.Request.Body,
|
||||
Reader: uploadReader,
|
||||
Mimetype: mimetype,
|
||||
WebPutAsTask: asTask,
|
||||
}
|
||||
@@ -118,6 +135,12 @@ func FsStream(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if generateSeed {
|
||||
if err = writeUploadSeedSidecar(c, dir, name, size, seedHasher); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
}
|
||||
if t == nil {
|
||||
common.SuccessResp(c)
|
||||
return
|
||||
@@ -194,6 +217,17 @@ func FsForm(c *gin.Context) {
|
||||
if len(mimetype) == 0 {
|
||||
mimetype = utils.GetMimeType(name)
|
||||
}
|
||||
generateSeed := shouldGenerateUploadSeed(c, dir)
|
||||
if generateSeed && asTask {
|
||||
common.ErrorStrResp(c, "seed sidecar generation requires synchronous upload", 400)
|
||||
return
|
||||
}
|
||||
var seedHasher *torrent.HashWriter
|
||||
var uploadReader io.Reader = f
|
||||
if generateSeed {
|
||||
seedHasher = torrent.NewHashWriter(seedPieceSize(c), seedPieceSize(c), 0)
|
||||
uploadReader = io.TeeReader(f, seedHasher)
|
||||
}
|
||||
s := &stream.FileStream{
|
||||
Obj: &model.Object{
|
||||
Name: name,
|
||||
@@ -201,7 +235,7 @@ func FsForm(c *gin.Context) {
|
||||
Modified: getLastModified(c),
|
||||
HashInfo: utils.NewHashInfoByMap(h),
|
||||
},
|
||||
Reader: f,
|
||||
Reader: uploadReader,
|
||||
Mimetype: mimetype,
|
||||
WebPutAsTask: asTask,
|
||||
}
|
||||
@@ -218,6 +252,12 @@ func FsForm(c *gin.Context) {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
if generateSeed {
|
||||
if err = writeUploadSeedSidecar(c, dir, name, file.Size, seedHasher); err != nil {
|
||||
common.ErrorResp(c, err, 500)
|
||||
return
|
||||
}
|
||||
}
|
||||
if t == nil {
|
||||
common.SuccessResp(c)
|
||||
return
|
||||
@@ -226,3 +266,107 @@ func FsForm(c *gin.Context) {
|
||||
"task": getTaskInfo(t),
|
||||
})
|
||||
}
|
||||
|
||||
func shouldGenerateUploadSeed(c *gin.Context, path string) bool {
|
||||
if strings.TrimSpace(c.GetHeader("X-Seed-Sidecars")) != "" {
|
||||
return true
|
||||
}
|
||||
policy := strings.ToLower(strings.TrimSpace(c.GetHeader("X-Generate-Seed")))
|
||||
if policy != "" && policy != "inherit" {
|
||||
return policy == "on" || policy == "true" || policy == "1"
|
||||
}
|
||||
if storage := op.GetBalancedStorage(path); storage != nil {
|
||||
policy = strings.ToLower(strings.TrimSpace(storage.GetStorage().SeedPolicy))
|
||||
}
|
||||
if policy == "" || policy == "inherit" {
|
||||
policy = strings.ToLower(setting.GetStr(conf.SeedAutoGeneratePolicy, "off"))
|
||||
}
|
||||
return (policy == "on" || policy == "true" || policy == "1") && configuredSeedFormats() != ""
|
||||
}
|
||||
|
||||
func seedPieceSize(c *gin.Context) int64 {
|
||||
for _, format := range strings.Split(strings.ToLower(c.GetHeader("X-Seed-Sidecars")), ",") {
|
||||
if strings.TrimSpace(format) == "cas" {
|
||||
return torrent.DefaultPieceSize
|
||||
}
|
||||
}
|
||||
value := c.GetHeader("X-Seed-Piece-Size")
|
||||
if value == "" {
|
||||
return torrent.DefaultPieceSize
|
||||
}
|
||||
size, err := strconv.ParseInt(value, 10, 64)
|
||||
if err != nil || size <= 0 || size > 1<<30 {
|
||||
return torrent.DefaultPieceSize
|
||||
}
|
||||
return size
|
||||
}
|
||||
|
||||
func configuredSeedFormats() string {
|
||||
policies := make(map[string]string)
|
||||
if err := json.Unmarshal([]byte(setting.GetStr(conf.SeedFormatPolicies)), &policies); err != nil {
|
||||
return ""
|
||||
}
|
||||
formats := make([]string, 0, 3)
|
||||
for _, format := range []string{"oss", "torrent", "cas"} {
|
||||
if strings.EqualFold(strings.TrimSpace(policies[format]), "on") {
|
||||
formats = append(formats, format)
|
||||
}
|
||||
}
|
||||
return strings.Join(formats, ",")
|
||||
}
|
||||
|
||||
func writeUploadSeedSidecar(c *gin.Context, dir, name string, expectedSize int64, hasher *torrent.HashWriter) error {
|
||||
if hasher == nil {
|
||||
return nil
|
||||
}
|
||||
hasher.Finish()
|
||||
if expectedSize >= 0 && hasher.GetTotalWritten() != expectedSize {
|
||||
return fmt.Errorf("seed sidecar requires a complete stream: read %d of %d bytes", hasher.GetTotalWritten(), expectedSize)
|
||||
}
|
||||
seed := torrent.NewSeed(name, "OpenList", seedPieceSize(c))
|
||||
formatsHeader := strings.TrimSpace(c.GetHeader("X-Seed-Sidecars"))
|
||||
if formatsHeader == "" {
|
||||
formatsHeader = strings.TrimSpace(c.GetHeader("X-Seed-Format"))
|
||||
}
|
||||
if formatsHeader == "" {
|
||||
formatsHeader = configuredSeedFormats()
|
||||
}
|
||||
formats, err := fs.NormalizeSeedFormats(strings.Split(formatsHeader, ","))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
matrix := SeedHashMatrix{}
|
||||
if rawMatrix := strings.TrimSpace(c.GetHeader("X-Seed-Hash-Matrix")); rawMatrix != "" {
|
||||
if err = json.Unmarshal([]byte(rawMatrix), &matrix); err != nil {
|
||||
return fmt.Errorf("invalid seed hash matrix: %w", err)
|
||||
}
|
||||
}
|
||||
matrix = normalizedSeedMatrix(matrix, formats)
|
||||
seedFile := hasher.BuildSeedFile(name, getLastModified(c).UTC().Format(time.RFC3339))
|
||||
applySeedMatrix(&seedFile, matrix)
|
||||
seed.Files = []torrent.SeedFile{seedFile}
|
||||
seen := make(map[string]struct{}, 3)
|
||||
for _, rawFormat := range formats {
|
||||
format := strings.ToLower(strings.TrimPrefix(strings.TrimSpace(rawFormat), "."))
|
||||
if format == "bt" {
|
||||
format = "torrent"
|
||||
}
|
||||
if _, exists := seen[format]; exists {
|
||||
continue
|
||||
}
|
||||
seen[format] = struct{}{}
|
||||
data, err := fs.EncodeGeneratedSeed(seed, format, hasher.GetPieceHashes())
|
||||
if err != nil {
|
||||
return fmt.Errorf("generate %s seed sidecar: %w", format, err)
|
||||
}
|
||||
sidecar := &stream.FileStream{
|
||||
Ctx: c.Request.Context(),
|
||||
Obj: &model.Object{Name: name + "." + format, Size: int64(len(data)), Modified: time.Now()},
|
||||
Reader: bytes.NewReader(data), Mimetype: "application/octet-stream",
|
||||
}
|
||||
if err = fs.PutDirectly(c.Request.Context(), dir, sidecar, true); err != nil {
|
||||
return fmt.Errorf("upload %s seed sidecar: %w", format, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+1512
-12
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,259 @@
|
||||
package handles
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/torrent"
|
||||
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
|
||||
)
|
||||
|
||||
// withSeedSite pins the configured seed site for the duration of a test.
|
||||
func withSeedSite(t *testing.T, site string) {
|
||||
t.Helper()
|
||||
previous := seedSiteURLProvider
|
||||
seedSiteURLProvider = func() string { return site }
|
||||
t.Cleanup(func() { seedSiteURLProvider = previous })
|
||||
}
|
||||
|
||||
func mustParse(t *testing.T, raw string) *url.URL {
|
||||
t.Helper()
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("url.Parse(%q) error = %v", raw, err)
|
||||
}
|
||||
return u
|
||||
}
|
||||
|
||||
// TestSameSeedHostNormalizesDefaultPort guards against the regression where a
|
||||
// configured "https://pan.example.com" and an embedded
|
||||
// "https://pan.example.com:443/..." were treated as different hosts, silently
|
||||
// discarding otherwise valid sources.
|
||||
func TestSameSeedHostNormalizesDefaultPort(t *testing.T) {
|
||||
cases := []struct {
|
||||
a, b string
|
||||
want bool
|
||||
}{
|
||||
{"https://pan.example.com", "https://pan.example.com:443/x", true},
|
||||
{"http://pan.example.com", "http://pan.example.com:80/x", true},
|
||||
{"https://pan.example.com:8443", "https://pan.example.com:8443/x", true},
|
||||
{"https://pan.example.com:8443", "https://pan.example.com", false},
|
||||
{"https://pan.example.com", "http://pan.example.com", false},
|
||||
{"https://pan.example.com", "https://evil.example.com", false},
|
||||
{"https://pan.example.com", "https://pan.example.com.evil.com", false},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
if got := sameSeedHost(mustParse(t, tc.a), mustParse(t, tc.b)); got != tc.want {
|
||||
t.Errorf("sameSeedHost(%q, %q) = %v, want %v", tc.a, tc.b, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestValidateSeedHostRejectsForeignHosts is the core SSRF guarantee: a seed
|
||||
// source may only ever point at the operator-configured site.
|
||||
func TestValidateSeedHostRejectsForeignHosts(t *testing.T) {
|
||||
withSeedSite(t, "https://pan.example.com")
|
||||
|
||||
rejected := []string{
|
||||
"http://169.254.169.254/latest/meta-data/", // cloud metadata
|
||||
"http://127.0.0.1:5244/api/fs/list", // local admin API
|
||||
"http://localhost:5244/d/secret",
|
||||
"https://evil.example.com/d/secret",
|
||||
"https://pan.example.com.evil.com/d/x",
|
||||
"file:///etc/passwd",
|
||||
"ftp://pan.example.com/x",
|
||||
"https://user:pass@pan.example.com/d/x", // embedded credentials
|
||||
}
|
||||
for _, raw := range rejected {
|
||||
if err := validateSeedHost(mustParse(t, raw)); err == nil {
|
||||
t.Errorf("validateSeedHost(%q) accepted a disallowed URL", raw)
|
||||
}
|
||||
}
|
||||
|
||||
allowed := []string{
|
||||
"https://pan.example.com/d/some/file",
|
||||
"https://pan.example.com:443/sd/abc123",
|
||||
}
|
||||
for _, raw := range allowed {
|
||||
if err := validateSeedHost(mustParse(t, raw)); err != nil {
|
||||
t.Errorf("validateSeedHost(%q) rejected a valid URL: %v", raw, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateSeedHostWithoutConfiguredSite(t *testing.T) {
|
||||
withSeedSite(t, "")
|
||||
if err := validateSeedHost(mustParse(t, "https://pan.example.com/d/x")); err == nil {
|
||||
t.Fatal("validateSeedHost() accepted a source while seed_site_url is unset")
|
||||
}
|
||||
}
|
||||
|
||||
// TestRedirectGuardBlocksSSRF is the regression test for the bypass: the first
|
||||
// hop passes the host allow-list, but a redirect must not be allowed to escape
|
||||
// to an internal address.
|
||||
func TestRedirectGuardBlocksSSRF(t *testing.T) {
|
||||
withSeedSite(t, "https://pan.example.com")
|
||||
|
||||
origin := mustParse(t, "https://pan.example.com/d/file")
|
||||
check := seedSourceHTTPClient.CheckRedirect
|
||||
|
||||
// A redirect staying on the configured site is fine.
|
||||
sameHost := &http.Request{URL: mustParse(t, "https://pan.example.com/d/file-2")}
|
||||
if err := check(sameHost, []*http.Request{{URL: origin}}); err != nil {
|
||||
t.Fatalf("CheckRedirect() rejected a same-host redirect: %v", err)
|
||||
}
|
||||
|
||||
// A redirect to the cloud metadata endpoint must be refused.
|
||||
metadata := &http.Request{URL: mustParse(t, "http://169.254.169.254/latest/meta-data/")}
|
||||
if err := check(metadata, []*http.Request{{URL: origin}}); err == nil {
|
||||
t.Fatal("CheckRedirect() allowed a redirect to the cloud metadata endpoint")
|
||||
}
|
||||
|
||||
// A redirect to localhost must be refused.
|
||||
local := &http.Request{URL: mustParse(t, "http://127.0.0.1:5244/api/fs/list")}
|
||||
if err := check(local, []*http.Request{{URL: origin}}); err == nil {
|
||||
t.Fatal("CheckRedirect() allowed a redirect to localhost")
|
||||
}
|
||||
|
||||
// A same-host redirect that downgrades https -> http must be refused.
|
||||
downgrade := &http.Request{URL: mustParse(t, "http://pan.example.com/d/file")}
|
||||
if err := check(downgrade, []*http.Request{{URL: origin}}); err == nil {
|
||||
t.Fatal("CheckRedirect() allowed a scheme downgrade")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRedirectGuardLimitsHopCount(t *testing.T) {
|
||||
withSeedSite(t, "https://pan.example.com")
|
||||
|
||||
origin := mustParse(t, "https://pan.example.com/d/file")
|
||||
via := make([]*http.Request, maxSeedSourceRedirects)
|
||||
for i := range via {
|
||||
via[i] = &http.Request{URL: origin}
|
||||
}
|
||||
next := &http.Request{URL: mustParse(t, "https://pan.example.com/d/file-2")}
|
||||
if err := seedSourceHTTPClient.CheckRedirect(next, via); err == nil {
|
||||
t.Fatal("CheckRedirect() accepted more redirects than maxSeedSourceRedirects")
|
||||
}
|
||||
}
|
||||
|
||||
// TestValidateSeedSourceEnforcesPathPrefix documents the per-type path contract.
|
||||
func TestValidateSeedSourceEnforcesPathPrefix(t *testing.T) {
|
||||
withSeedSite(t, "https://pan.example.com")
|
||||
|
||||
if err := validateSeedSource(torrent.SeedSource{
|
||||
Type: "openlist-direct",
|
||||
URL: "https://pan.example.com/sd/abc",
|
||||
}); err == nil {
|
||||
t.Fatal("validateSeedSource() accepted a share path for a direct source")
|
||||
}
|
||||
if err := validateSeedSource(torrent.SeedSource{
|
||||
Type: "openlist-share",
|
||||
URL: "https://pan.example.com/d/file",
|
||||
}); err == nil {
|
||||
t.Fatal("validateSeedSource() accepted a direct path for a share source")
|
||||
}
|
||||
if err := validateSeedSource(torrent.SeedSource{
|
||||
Type: "openlist-direct",
|
||||
URL: "https://evil.example.com/d/file",
|
||||
}); err == nil {
|
||||
t.Fatal("validateSeedSource() accepted a foreign host")
|
||||
}
|
||||
if err := validateSeedSource(torrent.SeedSource{
|
||||
Type: "openlist-direct",
|
||||
URL: "https://pan.example.com/d/file",
|
||||
}); err != nil {
|
||||
t.Fatalf("validateSeedSource() rejected a valid direct source: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFirstUsableSeedSourceSkipsExpiredAndForeign verifies the selection logic
|
||||
// only returns sources that are both on-site and not expired.
|
||||
func TestFirstUsableSeedSourceSkipsExpiredAndForeign(t *testing.T) {
|
||||
withSeedSite(t, "https://pan.example.com")
|
||||
|
||||
file := torrent.SeedFile{
|
||||
Sources: []torrent.SeedSource{
|
||||
{Type: "openlist-direct", URL: "https://evil.example.com/d/a"},
|
||||
{Type: "openlist-direct", URL: "https://pan.example.com/d/b", ExpiresAt: "2000-01-01T00:00:00Z"},
|
||||
{Type: "openlist-share", URL: "https://pan.example.com/sd/good"},
|
||||
},
|
||||
}
|
||||
if got := firstUsableSeedSource(file); got != "https://pan.example.com/sd/good" {
|
||||
t.Fatalf("firstUsableSeedSource() = %q, want the valid share source", got)
|
||||
}
|
||||
|
||||
// No usable source at all.
|
||||
none := torrent.SeedFile{
|
||||
Sources: []torrent.SeedSource{
|
||||
{Type: "openlist-direct", URL: "https://evil.example.com/d/a"},
|
||||
},
|
||||
}
|
||||
if got := firstUsableSeedSource(none); got != "" {
|
||||
t.Fatalf("firstUsableSeedSource() = %q, want empty", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestBuildSeedRapidUploadRequestRejectsMultiFile documents that a multi-file
|
||||
// torrent cannot be described by a single rapid-upload request. Returning nil
|
||||
// (instead of silently using Files[0] with the aggregate size) prevents sending
|
||||
// the destination a size/hash combination that contradicts itself.
|
||||
func TestBuildSeedRapidUploadRequestRejectsMultiFile(t *testing.T) {
|
||||
const md5Hex = "0123456789abcdef0123456789abcdef"
|
||||
|
||||
multi := &torrent.Torrent{
|
||||
OpenList: &torrent.Seed{
|
||||
PieceSize: torrent.DefaultPieceSize,
|
||||
Files: []torrent.SeedFile{
|
||||
{Path: "a.bin", Size: 10, Hashes: torrent.SeedHashes{MD5: md5Hex}},
|
||||
{Path: "b.bin", Size: 20, Hashes: torrent.SeedHashes{MD5: md5Hex}},
|
||||
},
|
||||
},
|
||||
}
|
||||
if req := buildSeedRapidUploadRequest(multi, nil); req != nil {
|
||||
t.Fatalf("buildSeedRapidUploadRequest() accepted a multi-file torrent: %#v", req)
|
||||
}
|
||||
|
||||
// A single file must still work.
|
||||
single := &torrent.Torrent{
|
||||
Info: torrent.TorrentInfo{Name: "a.bin", Length: 10},
|
||||
OpenList: &torrent.Seed{
|
||||
PieceSize: torrent.DefaultPieceSize,
|
||||
Files: []torrent.SeedFile{
|
||||
{Path: "a.bin", Size: 10, Hashes: torrent.SeedHashes{MD5: md5Hex}},
|
||||
},
|
||||
},
|
||||
}
|
||||
req := buildSeedRapidUploadRequest(single, nil)
|
||||
if req == nil {
|
||||
t.Fatal("buildSeedRapidUploadRequest() rejected a valid single-file torrent")
|
||||
}
|
||||
if req.Size != 10 {
|
||||
t.Fatalf("buildSeedRapidUploadRequest() size = %d, want 10", req.Size)
|
||||
}
|
||||
if got := req.Whole.GetHash(utils.MD5); !strings.EqualFold(got, md5Hex) {
|
||||
t.Fatalf("buildSeedRapidUploadRequest() MD5 = %q, want %q", got, md5Hex)
|
||||
}
|
||||
|
||||
// A single file whose metadata size disagrees with the torrent length is
|
||||
// internally inconsistent and must also be refused.
|
||||
mismatched := &torrent.Torrent{
|
||||
Info: torrent.TorrentInfo{Name: "a.bin", Length: 99},
|
||||
OpenList: &torrent.Seed{
|
||||
PieceSize: torrent.DefaultPieceSize,
|
||||
Files: []torrent.SeedFile{
|
||||
{Path: "a.bin", Size: 10, Hashes: torrent.SeedHashes{MD5: md5Hex}},
|
||||
},
|
||||
},
|
||||
}
|
||||
if req := buildSeedRapidUploadRequest(mismatched, nil); req != nil {
|
||||
t.Fatalf("buildSeedRapidUploadRequest() accepted a size/hash mismatch: %#v", req)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildSeedRapidUploadRequestHandlesNil(t *testing.T) {
|
||||
if req := buildSeedRapidUploadRequest(nil, nil); req != nil {
|
||||
t.Fatalf("buildSeedRapidUploadRequest(nil) = %#v, want nil", req)
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
}
|
||||
@@ -234,6 +234,19 @@ func _fs(g *gin.RouterGroup) {
|
||||
g.POST("/torrent/upload_parse", handles.UploadTorrentAndParse)
|
||||
g.POST("/torrent/rapid_upload", handles.TorrentRapidUpload)
|
||||
g.POST("/torrent/generate", handles.GenerateTorrentForPath)
|
||||
// Unified transfer seed APIs. Legacy torrent routes above remain supported.
|
||||
seed := g.Group("/seed")
|
||||
seed.POST("/parse", handles.ParseSeed)
|
||||
seed.POST("/upload_parse", handles.UploadSeedAndParse)
|
||||
seed.POST("/generate", handles.GenerateSeedForPaths)
|
||||
seed.POST("/convert", handles.ConvertSeed)
|
||||
seed.POST("/diagnose", handles.DiagnoseSeed)
|
||||
seed.POST("/capabilities", handles.SeedCapabilities)
|
||||
seed.POST("/rapid_upload", handles.QuickSaveSeed)
|
||||
seed.POST("/offline_download", handles.QuickSaveSeed)
|
||||
seed.POST("/update", handles.UpdateSeed)
|
||||
seed.POST("/quick_save", handles.QuickSaveSeed)
|
||||
seed.POST("/update_channels", handles.UpdateSeedChannels)
|
||||
// Direct upload (client-side upload to storage)
|
||||
g.POST("/get_direct_upload_info", handles.FsGetDirectUploadInfo)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+2
-11
@@ -159,18 +159,9 @@ func s3RequestAuthorized(r *http.Request, authPairs map[string]string) bool {
|
||||
if len(authPairs) == 0 {
|
||||
return true
|
||||
}
|
||||
// Verify against the keys this server was configured with. V4SignVerify and
|
||||
// V2SignVerify read the signature package's process-wide key store, which
|
||||
// gofakes3 never writes (keys are kept per instance), so they always
|
||||
// returned InvalidAccessKeyId and the 302/307 direct-transfer redirects
|
||||
// never ran. Same V4-then-V2 order the auth middleware uses.
|
||||
lookup := func(accessKey string) (string, bool) {
|
||||
secret, ok := authPairs[accessKey]
|
||||
return secret, ok
|
||||
}
|
||||
result := signature.V4SignVerifyWithLookup(r, lookup)
|
||||
result := signature.V4SignVerify(r)
|
||||
if result == signature.ErrUnsupportAlgorithm {
|
||||
result = signature.V2SignVerifyWithLookup(r, lookup)
|
||||
result = signature.V2SignVerify(r)
|
||||
}
|
||||
return result == signature.ErrNone
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user