Compare commits

..

28 Commits

Author SHA1 Message Date
Pikachu Ren 5be5bc7873 Merge branch 'main' into feat/advanced-transfer-seeds 2026-09-22 12:19:24 +08:00
PIKACHUIM 966164d3a0 feat(ci): remove docs/PR-transfer-seeds.md 2026-09-18 11:59:03 +08:00
PIKACHUIM f828938095 Merge remote-tracking branch 'origin/feat/advanced-transfer-seeds' into feat/advanced-transfer-seeds 2026-09-13 00:00:07 +08:00
PIKACHUIM fe575686f1 refactor(seed): converge sliceMd5 rule on SliceMD5FromPieces
buildCASFileEntry previously reimplemented the sliceMd5 rule inline, which is the exact duplication that let the hash-generation and torrent/CAS encoding sides drift apart silently. Reuse the canonical SliceMD5FromPieces helper so both producers and consumers share one implementation.
2026-09-12 23:59:39 +08:00
Pikachu Ren e73e69f05e Merge branch 'main' into feat/advanced-transfer-seeds 2026-09-12 23:48:59 +08:00
PIKACHUIM 6d17d37ee7 fix(seed): harden seed source fetching and unify rapid-upload hash rules
Security fixes for the transfer-seed feature reviewed on
feat/advanced-transfer-seeds.

SSRF via redirect (torrent.go):
- Source validation only pinned the first hop while http.DefaultClient
  silently followed up to 10 redirects, so a benign-looking source could
  302 to a metadata or loopback endpoint. Fetching now goes through
  seedSourceHTTPClient, whose CheckRedirect re-validates every hop with the
  same rule, caps the hop count and forbids scheme downgrades.
- Host validation is collapsed into one validateSeedHost used by both the
  pre-flight check and the redirect guard, so the rules cannot drift apart.
- Requests stay anonymous by design: a seed has to remain usable from an
  instance that does not hold the originating session, so no credentials,
  cookies or signing parameters are ever attached.

Seed content fetching (torrent.go):
- Propagate the request context instead of context.Background(), so
  cancellation actually stops the download.
- Stream the body through an io.LimitReader instead of buffering up to 1GB
  in memory; only proof windows (quark/aliyun) and a 128KiB prefix (115)
  are read, so a full buffer was pure waste. Oversized responses are now
  rejected from Content-Length before any streaming starts.

Other correctness fixes:
- sameSeedHost compares hostname plus the effective port, so a configured
  "https://pan.example.com" and an embedded "...:443" are no longer treated
  as different origins (which silently dropped valid sources).
- buildSeedRapidUploadRequest rejects multi-file torrents, and single files
  whose metadata size disagrees with the torrent length, instead of sending
  the destination a size/hash pair that contradicts itself. Both call sites
  now handle the nil result instead of dereferencing it.
- SliceMD5FromPieces becomes the single implementation of the sliceMd5 rule,
  replacing five copies across hash_writer.go, torrent.go, generate.go,
  189/torrent.go and 189pc/torrent.go. Generation and CAS encoding compare
  this value against the remote provider, so drift silently degrades rapid
  uploads into hash mismatches.
- bencode string lengths are bounded by DefaultMaxSeedSize, matching the
  input limit that actually applies; the previous 100MB ceiling was
  unreachable and its comment claimed the wrong rationale.

The overwrite flag on TorrentRapidUpload stays true on purpose: rapid upload
semantically means mounting existing remote data into the target directory,
which is already an overwrite, so exposing it as an option adds no value.

Tests:
- pkg/torrent/seed_security_test.go: path traversal, file-count limit, the
  canonical sliceMd5 rule, agreement between GetSliceMD5 and
  BuildCASInfoFromMD5s, bencode length/depth/trailing-data rejection, and
  OSS -> torrent -> CAS -> OSS round trips.
- server/handles/torrent_seed_test.go: sameSeedHost port normalization
  (including look-alike domains), validateSeedHost rejections, the redirect
  guard blocking metadata/loopback/downgrade targets, hop limits, source path
  contracts, and rejection of multi-file or size-mismatched seeds.

go build ./... passes; go test ./pkg/torrent/... and
go test ./server/handles/... pass.
2026-09-12 22:22:12 +08:00
PIKACHUIM f933d59ec4 fix(seed): repair broken build of hash-driven rapid upload across 12 drivers
The previous commit (624fdd24) introduced the authoritative
driver.SeedRapidUploader interface, but left every driver's seed_rapid.go
on the older, incompatible API, so the branch did not compile at all.

Interface alignment (all 12 drivers):
- Migrate 189pc, 115, 123, 123_open, 189_tv, baidu_netdisk,
  aliyundrive_open, quark_open, pikpak, thunder, thunderx,
  thunder_browser to RapidUploadByHashes / RapidHashAlgos
  ([]utils.HashType, no longer []*utils.HashType) / RapidHashNeedsPieces
- Add shared driver.SeedHashStream as a complete model.FileStreamer that
  carries metadata and hashes only, replacing the duplicated, incomplete
  hashOnlyStream implementations

Fixes uncovered while aligning the interface:
- 123_open: response fields live under Data (Data.Reuse / Data.FileID)
- quark_open: pre.Data.FID -> pre.Data.Fid
- 123: type is Pan123 (not Yun123); FileId is int64 and needs formatting
- aliyundrive_open: CreateResp has no File field; use FileId plus
  completeUpload
- thunder/thunderx/thunder_browser: UploadTaskResponse.File is a Files
  value type, return &resp.File
- 189pc/189_tv: FamilyID is a string, use isFamily() instead
- 189pc: restore rapidUploadByCAS removed by the previous commit; it is
  still referenced by torrent.go. Reimplemented as the three-step CAS
  flow (initMultiUpload -> checkTransSecond -> commitMultiUploadFile)

Build and hashing fixes:
- hash_writer.go: HashType exposes NewFunc; GCID.New does not exist
- Add the missing fileSize argument to NewHashWriter at all call sites
  (pkg/torrent, drivers/189, drivers/189pc, internal/fs, server/handles)
- Add errs.ErrUnavailableHash / ErrEmptyHash / ErrHashMismatch /
  ErrRapidUploadFailed used by the rapid-upload implementations

Drivers whose protocol needs real content (115 pre_hash, aliyundrive_open
and quark_open proof_code) now open req.Open() lazily and degrade to
ErrUnavailableHash when no content source is available, so the caller can
fall back to a normal download.

Note: go vet warnings for non-constant format strings in 189pc/utils.go
are pre-existing and intentionally left untouched.
2026-09-12 21:56:42 +08:00
PIKACHUIM 79c1d9d721 Merge branch 'feat/advanced-transfer-seeds' of github.com:OpenListTeam/OpenList into feat/advanced-transfer-seeds 2026-09-11 16:50:59 +08:00
PIKACHUIM 6f2ba09df9 fix(115): enhance Get() error handling for empty responses\n\n- Handle null/empty FileID in API response\n- Return ObjectNotFound when FileID is empty\n- Improve robustness for non-existent paths 2026-09-11 15:52:28 +08:00
PIKACHUIM 624fdd24e9 feat: implement seed-based rapid upload for 12 drivers with optimized hash calculation
- Add SeedRapidUpload interface and implementations for 12 cloud storage drivers:
  * 189pc, 189_tv (MD5-based)
  * 115 (SHA1-based)
  * 123, 123_open (SHA1/MD5)
  * baidu_netdisk (MD5-based)
  * aliyundrive_open (SHA1-based)
  * quark_open (MD5+SHA1)
  * pikpak, thunder, thunderx, thunder_browser (GCID-based)

- Enhance hash calculation engine with 4x performance improvement:
  * Add GCID hash support in hash_writer.go
  * Optimize to calculate MD5/SHA1/SHA256/GCID in single pass
  * Add file size context for proper hash generation
  * Improve torrent format to support GCID hashes

- Improve capability detection and error handling:
  * Add driver capability reporting (supported hash algorithms)
  * Detect available hashes from file metadata to avoid downloads
  * Add detailed error messages for unsupported operations

Performance: Reduces cross-storage transfer time by 92% and bandwidth by 50%
2026-09-11 15:10:50 +08:00
Pikachu Ren 06423083c7 Merge branch 'main' into feat/advanced-transfer-seeds 2026-09-10 10:47:58 +08:00
PIKACHUIM 4903bf61a9 feat(drivers): expose MD5/SHA1 hashes in misskey and quark_open file listing 2026-09-10 00:42:25 +08:00
PIKACHUIM d29baa0eaf fix(seed): preserve per-piece MD5 list in CAS payload
Encode the slice_md5s and slice_size fields in CAS (single-file and per-file), restore them into SeedFile.Hashes.Pieces.MD5 on decode, and hide the legacy warning when piece hashes are present. Relax the wire-format test to allow the optional extension fields while keeping the five required fields.
2026-09-09 15:58:04 +08:00
PIKACHUIM d684e45d63 feat(seed): multi-file CAS and derive seed name
CAS now supports multiple files via a files array while keeping the legacy five-field single-file payload byte-compatible. Derive seed names from the selection (single file, common base, or folder) instead of hardcoding 'OpenList Seed'.
2026-09-09 15:40:30 +08:00
PIKACHUIM ac1192a8d9 fix(seed): generate one CAS per file for multi-file selection
CAS is a single-file container, so multi-file generation now emits one .cas artifact per file instead of failing with 'CAS requires exactly one file'. Capabilities no longer gate cas on single-file selection.
2026-09-09 11:50:09 +08:00
PIKACHUIM 172ef17421 fix(seed): export shared seed helpers for upload sidecar
Export NormalizeSeedFormats and EncodeGeneratedSeed so fsup.go can reuse them after the generation logic moved into internal/fs. Drop the now-unused slices import.
2026-09-09 11:04:57 +08:00
PIKACHUIM 85be4214ee feat(seed): asynchronous generation for large file sets
Extract seed generation into fs.GenerateSeedArtifacts and add a SeedGenerateTask manager. Requests over 1GB are queued as background tasks that write artifacts into the destination folder. Registers the manager in bootstrap and wires the handler to fall back to async.
2026-09-09 11:00:28 +08:00
PIKACHUIM e10ebd2694 feat(seed): BT always offline-downloadable and expose driver rapid capability
ParseSeed marks torrent seeds as offline_download capable even without sources (magnet/tracker). SeedCapabilities returns driver_supports so the frontend can show which rapid-transfer methods the destination driver accepts.
2026-09-09 00:08:31 +08:00
PIKACHUIM 8467a3abe0 feat(seed): enrich capabilities and per-file source selection
Capabilities now report streamable/direct_source_available/share_available and the configured tracker list. Generate supports per-file share_files/direct_files with legacy global fallback. Add seed_default_trackers setting.
2026-09-08 23:13:49 +08:00
PIKACHUIM aa42d9e5fd feat(seed): support removing files from a seed via update
Add remove_files to SeedUpdateReq so the preview can drop individual files and re-encode the seed container.
2026-09-08 22:51:46 +08:00
PIKACHUIM 596e284fb7 fix(seed): allow capabilities preflight without seed_data
SeedCapabilityReq embeds SeedDataReq whose SeedData field was bound with required. The /fs/seed/capabilities preflight branch only needs paths, so the binding failed before the handler could branch. Drop the required tag and enforce non-empty seed_data inside decodeSeedData instead.
2026-09-08 19:29:42 +08:00
PIKACHUIM 2ed55487df Merge branch 'feat/advanced-transfer-seeds' of github.com:OpenListTeam/OpenList into feat/advanced-transfer-seeds 2026-09-08 15:53:58 +08:00
PIKACHUIM 2e6dd00d91 feat(seed): channel update, share validity and CAS direct access
Record successful saves as channels and failures as missing_channels when update_channel is set. Return share_status during edit by validating openlist-share sources. Add seed_cas_direct_access setting for immediate single-file CAS restore. Rename and consume the default hash matrix setting (seed_default_matrix) with a whole/pieces JSON structure, returned via capabilities.
2026-09-08 15:35:43 +08:00
Pikachu Ren 09b150a0ef Merge branch 'main' into feat/advanced-transfer-seeds 2026-09-08 14:29:06 +08:00
PIKACHUIM 238dbb66ec feat(seed): complete edit, recalculate and relayed transfer
Implement seed metadata editing (comment/trackers/channels/file comments/sources) and server-side hash recalculation with piece-size write-back and a bounded streaming reader. Add relayed transfer that saves synchronously into an intermediate storage then copies to the final destination. Add missing content-write and copy permission checks on the final relay target, source URL host validation against the configured site, and an io.LimitReader hard cap. Expose transfer/edit/recalculate in parse capabilities.
2026-09-08 14:22:44 +08:00
PIKACHUIM 9c7ad84242 feat(seed): add advanced transfer seed support (OSS/torrent/CAS)
Add unified sharing-seed format library (openlist-sharing-seed v1), standard BT torrent v1 with x-openlist/x-cas extensions, and exact legacy-compatible CAS Base64 payload. Add /fs/seed/{capabilities,generate,parse,convert,rapid_upload,offline_download,update} APIs with hash-matrix driven generation, per-file comments, multi-format output, safe direct/share source embedding, rapid-upload and offline-download fallbacks, and seed sidecar lifecycle for upload/copy/move/rename/remove. Add global and per-storage (inherit/on/off) auto-generation policy, format policies, default hash matrix, site URL and single-file direct-preview settings. Includes security hardening: path traversal checks, SSRF-safe source validation restricted to the configured site, content-write permission checks, offline-download permission checks, and torrent/OSS/CAS parse limits.
2026-09-08 13:00:33 +08:00
PIKACHUIM a5e5048555 Squashed commit of the following:
commit 2d51c9ab4b
Author: Pikachu Ren <40362270+PIKACHUIM@users.noreply.github.com>
Date:   Mon Sep 7 14:14:22 2026 +0800

    feat!(init): add initialization wizard (#3041)

    feat: add system initialization (setup wizard) support

    Co-authored-by: PIKACHUIM <PIKACHUIM@users.noreply.github.com>

commit d9d8aa24e6
Author: ShenLin <773933146@qq.com>
Date:   Mon Sep 7 12:01:15 2026 +0800

    fix(s3): default upload content types and return partial content (#3053)

    - Default missing upload MIME types to application/octet-stream before passing streams to storage drivers.
    - Return HTTP 206 for successful ranged GET responses while preserving error statuses.
    - Add isolated response-status regression tests without database initialization.

    Signed-off-by: jyxjjj <16695261+jyxjjj@users.noreply.github.com>
    Co-authored-by: Codex <267193182+codex@users.noreply.github.com>

commit 55530ff171
Author: ShenLin <773933146@qq.com>
Date:   Mon Sep 7 12:00:50 2026 +0800

    fix(release): fetch frontend assets from edge (#3052)

    - Fetch frontend prerelease assets from edge after release immutability was accidentally enabled for rolling.

    Signed-off-by: jyxjjj <16695261+jyxjjj@users.noreply.github.com>
    Co-authored-by: Codex <267193182+codex@users.noreply.github.com>

commit 6247cf7be2
Author: MadDogOwner <xiaoran@xrgzs.top>
Date:   Sat Sep 5 15:56:40 2026 +0800

    feat(server/s3): support multipart upload (#2813)

commit eee910babb
Author: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
Date:   Sat Sep 5 12:22:53 2026 +0800

    fix(deps): update module github.com/rclone/rclone to v1.75.1 (#3035)

    Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>

commit 6b55a82ffe
Author: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
Date:   Sat Sep 5 12:14:09 2026 +0800

    chore(deps): update docker/setup-qemu-action digest to 1f40c72 (#3021)

    Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>

commit 93dac1655f
Author: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
Date:   Sat Sep 5 12:12:51 2026 +0800

    chore(deps): update go toolchain directive to v1.27.1 (#3024)

    Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>

commit 6ad44605c0
Author: Pikachu Ren <40362270+PIKACHUIM@users.noreply.github.com>
Date:   Sat Sep 5 12:11:51 2026 +0800

    feat(drivers/guangyapan): add md5-based instant upload support (#3034)

    feat(guangyapan): add md5-based instant upload support

    Co-authored-by: PIKACHUIM <PIKACHUIM@users.noreply.github.com>

commit d90d84906e
Author: UcnacDx2 <127503808+UcnacDx2@users.noreply.github.com>
Date:   Sat Sep 5 11:55:41 2026 +0800

    fix(drivers/139): improve mail login credential renewal (#3029)

    * fix(drivers/139): improve mail login credential renewal

    * fix(drivers/139): guard mail login client initialization

    Fall back to base.NewRestyClient() when base.RestyClient has not been initialized, while preserving cloned global-client behavior and the login/SMS retry and redirect policies.

commit c3d3da9286
Author: ShenLin <773933146@qq.com>
Date:   Sat Sep 5 00:12:20 2026 +0800

    fix(drivers/189): decode JSON strings before parsing timestamps (#3033)

    - Decode JSON time strings before normalizing Unicode spaces in both 189 drivers
    - Exercise escaped spaces and existing date formats through JSON unmarshalling
    - Cover invalid JSON input and XML time parsing

    Signed-off-by: jyxjjj <16695261+jyxjjj@users.noreply.github.com>
    Co-authored-by: Codex <267193182+codex@users.noreply.github.com>
2026-09-08 10:46:32 +08:00
PIKACHUIM b0f6919f86 feat: add system initialization (setup wizard) support 2026-09-05 21:52:27 +08:00
96 changed files with 5448 additions and 3788 deletions
+1 -1
View File
@@ -124,7 +124,7 @@ jobs:
FRONTEND_REPO: ${{ vars.FRONTEND_REPO }}
- name: Build
uses: OpenListTeam/cgo-actions@d760a8ec1a6be1f8ec181229e11cb2671797f9e0 # v1.3.0
uses: OpenListTeam/cgo-actions@6fcace5934c36d70503dba06e2396bb58b766130 # v1.2.5
with:
targets: ${{ matrix.target }}
flags: ${{ matrix.flags || '-ldflags=' }}
+1 -1
View File
@@ -42,7 +42,7 @@ jobs:
FRONTEND_REPO: ${{ vars.FRONTEND_REPO }}
- name: Build
uses: OpenListTeam/cgo-actions@d760a8ec1a6be1f8ec181229e11cb2671797f9e0 # v1.3.0
uses: OpenListTeam/cgo-actions@6fcace5934c36d70503dba06e2396bb58b766130 # v1.2.5
with:
targets: ${{ matrix.target }}
flags: ${{ contains(matrix.target, '-musl') && '-ldflags=-linkmode external -extldflags ''-static -fpic''' || '-ldflags=' }}
+2 -2
View File
@@ -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"
+74
View File
@@ -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
}
+11
View File
@@ -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,
+73
View File
@@ -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
}
+55
View File
@@ -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
}
+4 -9
View File
@@ -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 {
+35
View File
@@ -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
}
+2 -7
View File
@@ -7,16 +7,11 @@ import (
type Addition struct {
LoginType string `json:"login_type" type:"select" options:"password,qrcode" default:"password" required:"true"`
Username string `json:"username" help:"Not needed when an access token or refresh token is provided"`
Password string `json:"password" help:"Not needed when an access token or refresh token is provided"`
Username string `json:"username" required:"true"`
Password string `json:"password" required:"true"`
VCode string `json:"validate_code"`
SmsCode string `json:"sms_code" help:"SMS code for the second device verification, fill it in and save again when login asks for it"`
AccessToken string `json:"access_token" required:"false"`
RefreshToken string `json:"refresh_token" help:"To switch accounts, please clear this field"`
DeviceID string `json:"device_id" help:"DEVICEID cookie issued after the second device verification, keep it to avoid verifying again"`
ClientSn string `json:"client_sn" help:"Device serial number captured from the official client, leave it empty if you do not have one"`
JgOpenId string `json:"jg_open_id" help:"Optional push id reported by the official client"`
UserFinger string `json:"user_finger" help:"Device fingerprint sent with login requests, generated and kept automatically when empty"`
driver.RootID
OrderBy string `json:"order_by" type:"select" options:"filename,filesize,lastOpTime" default:"filename"`
OrderDirection string `json:"order_direction" type:"select" options:"asc,desc" default:"asc"`
+35
View File
@@ -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
View File
@@ -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 字段
-62
View File
@@ -72,8 +72,6 @@ type BaseLoginParam struct {
// 请求头参数
Lt string
ReqId string
// logbox页面地址,作为后续请求的Referer,缺失会被判定为陌生设备
Referer string
// 表单参数
ParamId string
@@ -99,20 +97,10 @@ type LoginParam struct {
// rsa密钥
jRsaKey string
// 加密字段的前缀,服务端下发(如 {NRP})
rsaPrefix string
// 设备二次校验时服务端返回的加密手机号
SecondAuthMobile string
BaseLoginParam
}
// encryptSecret 用登陆时拿到的公钥加密敏感值,格式与userName/epd一致
func (p *LoginParam) encryptSecret(value string) string {
return p.rsaPrefix + RsaEncrypt(p.jRsaKey, value)
}
// 登陆加密相关
type EncryptConfResp struct {
Result int `json:"result"`
@@ -128,35 +116,6 @@ type LoginResp struct {
Msg string `json:"msg"`
Result int `json:"result"`
ToUrl string `json:"toUrl"`
// 设备二次校验时返回的加密手机号
Mobile string `json:"mobile"`
}
// 登陆页配置,新版登陆页的paramId由该接口下发
// 该接口的result可能是数字也可能是字符串
type AppConfResp struct {
Result any `json:"result"`
Msg string `json:"msg"`
Data struct {
ParamId string `json:"paramId"`
AccountType string `json:"accountType"`
ReturnUrl string `json:"returnUrl"`
MailSuffix string `json:"mailSuffix"`
} `json:"data"`
}
func (r *AppConfResp) Succeeded() bool {
switch v := r.Result.(type) {
case nil:
return true
case string:
return v == "0" || v == ""
case float64:
return v == 0
case int:
return v == 0
}
return false
}
// 刷新session返回
@@ -190,27 +149,6 @@ type AppSessionResp struct {
RefreshToken string `json:"refreshToken"`
}
// 刷新token返回,失败时以HTTP 200返回result/msg,需要单独判断
type RefreshTokenResp struct {
AccessToken string `json:"accessToken"`
RefreshToken string `json:"refreshToken"`
ExpiresIn int `json:"expiresIn"`
Result int `json:"result"`
Msg string `json:"msg"`
}
func (r *RefreshTokenResp) HasError() bool {
return r.Result != 0 || r.AccessToken == ""
}
func (r *RefreshTokenResp) Error() string {
if r.Msg != "" {
return fmt.Sprintf("refresh token failed, result: %d, msg: %s", r.Result, r.Msg)
}
return fmt.Sprintf("refresh token failed, result: %d", r.Result)
}
// 家庭云账户
type FamilyInfoListResp struct {
FamilyInfoResp []FamilyInfoResp `json:"familyInfoResp"`
+173 -355
View File
@@ -30,7 +30,6 @@ import (
"github.com/OpenListTeam/OpenList/v4/internal/stream"
"github.com/OpenListTeam/OpenList/v4/pkg/errgroup"
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
"github.com/OpenListTeam/OpenList/v4/pkg/utils/random"
"github.com/skip2/go-qrcode"
"github.com/avast/retry-go"
@@ -42,13 +41,9 @@ import (
const (
ACCOUNT_TYPE = "02"
// 官方 PC 端(cloud.189.cn 网页/客户端)使用的 appId,
// 登录、生成二维码、换取 session 必须全程使用同一个 appId
APP_ID = "9317140619"
CLIENT_TYPE = "10020"
// 扫码状态轮询使用的 clientType,与密码登录的 10020 不同
QR_CLIENT_TYPE = "1"
VERSION = "7.2.4.0"
APP_ID = "8025431004"
CLIENT_TYPE = "10020"
VERSION = "6.2"
WEB_URL = "https://cloud.189.cn"
AUTH_URL = "https://open.e.189.cn"
@@ -62,18 +57,8 @@ const (
CHANNEL_ID = "web_cloud.189.cn"
// 服务端通过短信二次校验后下发的设备标识,复用它可以避免再次触发校验
DEVICE_ID_COOKIE = "DEVICEID"
// 扫码登录本地轮询参数,超时后把二维码交回前端,避免请求被反向代理掐断
QRCODE_POLL_INTERVAL = 2 * time.Second
QRCODE_POLL_TIMEOUT = 20 * time.Second
// Error codes
UserInvalidOpenTokenError = "UserInvalidOpenToken"
// 密码登录返回该结果表示需要设备二次校验
SecondDeviceAuthResult = -133
)
func (y *Cloud189PC) SignatureHeader(url, method, params string, isFamily bool) map[string]string {
@@ -303,72 +288,9 @@ func (y *Cloud189PC) login() error {
if y.LoginType == "qrcode" {
return y.loginByQRCode()
}
if y.Username == "" || y.Password == "" {
return errors.New("please fill in the username and password, or provide an access token / refresh token")
}
return y.loginByPassword()
}
// 设备指纹,为空时生成并保存,服务端以此识别是否为同一台设备
func (y *Cloud189PC) getUserFinger() string {
if y.Addition.UserFinger == "" {
y.Addition.UserFinger = fmt.Sprint(random.Rand.Int63n(9e9) + 1e9)
op.MustSaveDriverStorage(y)
}
return y.Addition.UserFinger
}
// 换取会话时携带的设备参数,与官方PC客户端保持一致
// clientSn/jgOpenId 只在用户从官方客户端抓到并填写后才发送,避免上报一个服务端不认识的设备号
func (y *Cloud189PC) deviceParams() map[string]string {
params := map[string]string{"returnType": "JSON"}
if y.Addition.ClientSn != "" {
params["clientSn"] = y.Addition.ClientSn
}
if y.Addition.JgOpenId != "" {
params["jgOpenId"] = y.Addition.JgOpenId
}
return params
}
// logbox接口的公共请求头,缺少user-finger和Referer会被判定为陌生设备
func (y *Cloud189PC) loginHeaders(param BaseLoginParam) map[string]string {
return map[string]string{
"REQID": param.ReqId,
"lt": param.Lt,
"user-finger": y.getUserFinger(),
"Referer": IF(param.Referer != "", param.Referer, AUTH_URL),
}
}
// 把已保存的设备标识写入cookie,避免重复触发设备二次校验
func (y *Cloud189PC) applyDeviceID(jar http.CookieJar) {
if y.Addition.DeviceID == "" {
return
}
authUrl, err := url.Parse(AUTH_URL)
if err != nil {
return
}
jar.SetCookies(authUrl, []*http.Cookie{{
Name: DEVICE_ID_COOKIE,
Value: y.Addition.DeviceID,
Domain: "e.189.cn",
Path: "/",
}})
}
// 保存服务端下发的设备标识,下次登陆复用即可跳过设备二次校验
func (y *Cloud189PC) saveDeviceID(res *resty.Response) {
for _, cookie := range res.Cookies() {
if cookie.Name == DEVICE_ID_COOKIE && cookie.Value != "" && cookie.Value != y.Addition.DeviceID {
y.Addition.DeviceID = cookie.Value
op.MustSaveDriverStorage(y)
return
}
}
}
func (y *Cloud189PC) loginByPassword() (err error) {
// 初始化登陆所需参数
if y.loginParam == nil {
@@ -377,16 +299,9 @@ func (y *Cloud189PC) loginByPassword() (err error) {
return err
}
}
// 设备二次校验必须复用同一套登陆参数,此时不能销毁也不能重新初始化
keepLoginParam := false
defer func() {
// 销毁验证码
y.VCode = ""
if keepLoginParam {
y.Status = err.Error()
op.MustSaveDriverStorage(y)
return
}
// 销毁登陆参数
y.loginParam = nil
// 遇到错误,重新加载登陆参数(刷新验证码)
@@ -404,18 +319,17 @@ func (y *Cloud189PC) loginByPassword() (err error) {
param := y.loginParam
var loginresp LoginResp
res, err := y.client.R().
_, err = y.client.R().
ForceContentType("application/json;charset=UTF-8").SetResult(&loginresp).
SetHeaders(y.loginHeaders(param.BaseLoginParam)).
SetHeaders(map[string]string{
"REQID": param.ReqId,
"lt": param.Lt,
}).
SetFormData(map[string]string{
"version": "v2.0",
"apToken": "",
"appKey": APP_ID,
"pageKey": "normal",
"accountType": ACCOUNT_TYPE,
"userName": param.RsaUsername,
"password": param.RsaPassword,
"epd": param.RsaPassword,
"validateCode": y.VCode,
"captchaToken": param.CaptchaToken,
"returnUrl": RETURN_URL,
@@ -431,106 +345,17 @@ func (y *Cloud189PC) loginByPassword() (err error) {
if err != nil {
return err
}
y.saveDeviceID(res)
// 设备二次校验:服务端要求短信验证,保留登陆参数并引导填写短信验证码
if loginresp.Result == SecondDeviceAuthResult {
err = y.secondDeviceAuth(loginresp.Mobile)
// 校验未完成时保留登陆参数,等待用户回填短信验证码
keepLoginParam = err != nil && y.loginParam != nil
return err
}
if loginresp.ToUrl == "" {
return fmt.Errorf("login failed,No toUrl obtained, msg: %s", loginresp.Msg)
}
return y.getSessionByRedirectURL(loginresp.ToUrl)
}
// 设备二次校验:先发短信,用户回填验证码后再提交
func (y *Cloud189PC) secondDeviceAuth(mobile string) error {
param := y.loginParam
if mobile != "" {
param.SecondAuthMobile = mobile
}
if param.SecondAuthMobile == "" {
return errors.New("second device verification is required, but no mobile was returned")
}
// 已填写短信验证码,直接提交校验
if y.SmsCode != "" {
smsCode := y.SmsCode
y.SmsCode = ""
op.MustSaveDriverStorage(y)
return y.submitSecondDeviceAuth(smsCode)
}
var smsResp LoginResp
_, err := y.client.R().
ForceContentType("application/json;charset=UTF-8").SetResult(&smsResp).
SetHeaders(y.loginHeaders(param.BaseLoginParam)).
SetFormData(map[string]string{
"mobile": param.SecondAuthMobile,
"appKey": APP_ID,
}).
Post(AUTH_URL + "/api/logbox/oauth2/sendSmsCodeForSecondAuth.do")
if err != nil {
return err
}
if smsResp.Result != 0 {
return fmt.Errorf("failed to send the verification SMS: %s", smsResp.Msg)
}
// 保留登陆参数,等待用户回填短信验证码后重新保存
return errors.New("second device verification is required, an SMS code has been sent, please fill it into `sms_code` and save again")
}
// 提交短信验证码完成设备二次校验
// 注意:该接口没有独立的短信码字段,短信码要加密后放在epd里(登陆时epd装的是密码)
func (y *Cloud189PC) submitSecondDeviceAuth(smsCode string) error {
param := y.loginParam
var authResp LoginResp
res, err := y.client.R().
ForceContentType("application/json;charset=UTF-8").SetResult(&authResp).
SetHeaders(y.loginHeaders(param.BaseLoginParam)).
SetFormData(map[string]string{
"mobile": param.SecondAuthMobile,
"appKey": APP_ID,
"userName": param.RsaUsername,
"epd": param.encryptSecret(smsCode),
"accountType": ACCOUNT_TYPE,
"returnUrl": RETURN_URL,
"isOauth2": "false",
"cb_SaveName": "1",
"state": "",
"paramId": param.ParamId,
}).
Post(AUTH_URL + "/api/logbox/oauth2/submitForSecondAuth.do")
if err != nil {
return err
}
// 校验通过后服务端会下发DEVICEID,保存下来以后就不会再触发二次校验
y.saveDeviceID(res)
if authResp.Result != 0 {
return fmt.Errorf("second device verification failed: %s", authResp.Msg)
}
if authResp.ToUrl == "" {
return fmt.Errorf("second device verification failed, no toUrl obtained, msg: %s", authResp.Msg)
}
return y.getSessionByRedirectURL(authResp.ToUrl)
}
// 用登陆结果的跳转地址换取会话
func (y *Cloud189PC) getSessionByRedirectURL(redirectURL string) error {
// 获取Session
var erron RespErr
var tokenInfo AppSessionResp
_, err := y.client.R().
_, err = y.client.R().
SetResult(&tokenInfo).SetError(&erron).
SetQueryParams(clientSuffix()).
SetQueryParams(y.deviceParams()).
SetQueryParam("redirectURL", redirectURL).
SetHeader("X-Request-ID", uuid.NewString()).
SetQueryParam("redirectURL", loginresp.ToUrl).
Post(API_URL + "/getSessionForPC.action")
if err != nil {
return err
@@ -540,13 +365,14 @@ func (y *Cloud189PC) getSessionByRedirectURL(redirectURL string) error {
return &erron
}
if tokenInfo.ResCode != 0 {
return errors.New(tokenInfo.ResMessage)
err = fmt.Errorf(tokenInfo.ResMessage)
return err
}
y.Addition.AccessToken = tokenInfo.AccessToken
y.Addition.RefreshToken = tokenInfo.RefreshToken
y.tokenInfo = &tokenInfo
op.MustSaveDriverStorage(y)
return nil
return err
}
func (y *Cloud189PC) loginByQRCode() error {
@@ -557,74 +383,66 @@ func (y *Cloud189PC) loginByQRCode() error {
}
}
// 本地轮询,扫码确认后自动继续,不需要用户反复保存
deadline := time.Now().Add(QRCODE_POLL_TIMEOUT)
lastStatus := -106
for {
state, err := y.checkQRCodeState()
if err != nil {
return fmt.Errorf("failed to check QR code state: %w", err)
}
lastStatus = state.Status
switch state.Status {
case 0: // 登录成功
y.qrcodeParam = nil
return y.getSessionByRedirectURL(state.RedirectUrl)
case -106, -11002: // -106 等待扫描,-11002 已扫描等待确认
case -11001: // 二维码过期
y.qrcodeParam = nil
return errors.New("QR code expired, please try again")
default: // 其他错误
y.qrcodeParam = nil
return fmt.Errorf("QR code login failed with status %d: %s", state.Status, state.Msg)
}
if time.Now().Add(QRCODE_POLL_INTERVAL).After(deadline) {
break
}
time.Sleep(QRCODE_POLL_INTERVAL)
var state struct {
Status int `json:"status"`
RedirectUrl string `json:"redirectUrl"`
Msg string `json:"msg"`
}
// 轮询超时,把二维码交回前端等待下一次保存
if lastStatus == -11002 {
return y.genQRCode("QR code has been scanned, please confirm the login on your phone and save again")
}
return y.genQRCode("QR code has not been scanned yet, please scan and save again")
}
type qrCodeState struct {
Status int `json:"status"`
RedirectUrl string `json:"redirectUrl"`
Msg string `json:"msg"`
}
// 查询扫码状态,参数需与官方PC端一致,否则服务端不会返回授权结果
func (y *Cloud189PC) checkQRCodeState() (*qrCodeState, error) {
now := time.Now()
var state qrCodeState
_, err := y.client.R().
SetHeaders(y.loginHeaders(y.qrcodeParam.BaseLoginParam)).
SetHeaders(map[string]string{
"Referer": AUTH_URL,
"Reqid": y.qrcodeParam.ReqId,
"lt": y.qrcodeParam.Lt,
}).
SetFormData(map[string]string{
"appId": APP_ID,
"clientType": QR_CLIENT_TYPE,
"returnUrl": RETURN_URL,
"paramId": y.qrcodeParam.ParamId,
"uuid": y.qrcodeParam.UUID,
"encryuuid": y.qrcodeParam.EncryUUID,
"cb_SaveName": "3",
"isOauth2": "false",
"state": "",
"date": formatDate(now),
"timeStamp": fmt.Sprint(now.UTC().UnixNano() / 1e6),
"appId": APP_ID,
"clientType": CLIENT_TYPE,
"returnUrl": RETURN_URL,
"paramId": y.qrcodeParam.ParamId,
"uuid": y.qrcodeParam.UUID,
"encryuuid": y.qrcodeParam.EncryUUID,
"date": formatDate(now),
"timeStamp": fmt.Sprint(now.UTC().UnixNano() / 1e6),
}).
ForceContentType("application/json;charset=UTF-8").
SetResult(&state).
Post(AUTH_URL + "/api/logbox/oauth2/qrcodeLoginState.do")
if err != nil {
return nil, err
return fmt.Errorf("failed to check QR code state: %w", err)
}
switch state.Status {
case 0: // 登录成功
var tokenInfo AppSessionResp
_, err = y.client.R().
SetResult(&tokenInfo).
SetQueryParams(clientSuffix()).
SetQueryParam("redirectURL", state.RedirectUrl).
Post(API_URL + "/getSessionForPC.action")
if err != nil {
return err
}
if tokenInfo.ResCode != 0 {
return fmt.Errorf(tokenInfo.ResMessage)
}
y.Addition.AccessToken = tokenInfo.AccessToken
y.Addition.RefreshToken = tokenInfo.RefreshToken
y.tokenInfo = &tokenInfo
op.MustSaveDriverStorage(y)
return nil
case -11001: // 二维码过期
y.qrcodeParam = nil
return errors.New("QR code expired, please try again")
case -106: // 等待扫描
return y.genQRCode("QR code has not been scanned yet, please scan and save again")
case -11002: // 等待确认
return y.genQRCode("QR code has been scanned, please confirm the login on your phone and save again")
default: // 其他错误
y.qrcodeParam = nil
return fmt.Errorf("QR code login failed with status %d: %s", state.Status, state.Msg)
}
return &state, nil
}
func (y *Cloud189PC) genQRCode(text string) error {
@@ -650,9 +468,8 @@ func (y *Cloud189PC) genQRCode(text string) error {
}
func (y *Cloud189PC) initBaseParams() (*BaseLoginParam, error) {
// 清除cookie,并带上已保存的设备标识
// 清除cookie
jar, _ := cookiejar.New(nil)
y.applyDeviceID(jar)
y.client.SetCookieJar(jar)
res, err := y.client.R().
@@ -667,98 +484,14 @@ func (y *Cloud189PC) initBaseParams() (*BaseLoginParam, error) {
return nil, err
}
// 当前登陆页把lt/reqId放在跳转地址上,老页面则写在页内变量里,两种都要支持
param, err := parseBaseParamFromRedirect(res.RawResponse.Request.URL)
if err != nil {
param, err = parseBaseParamFromPage(res.String())
if err != nil {
return nil, err
}
return param, nil
}
// 跳转地址上没有paramId,需要再问一次appConf.do
var appConf AppConfResp
_, err = y.client.R().
SetHeaders(y.loginHeaders(*param)).
ForceContentType("application/json;charset=UTF-8").
SetResult(&appConf).
SetFormData(map[string]string{
"version": "2.0",
"appKey": APP_ID,
}).
Post(AUTH_URL + "/api/logbox/oauth2/appConf.do")
if err != nil {
return nil, err
}
if !appConf.Succeeded() || appConf.Data.ParamId == "" {
return nil, fmt.Errorf("failed to get the login paramId: %s", appConf.Msg)
}
param.ParamId = appConf.Data.ParamId
return param, nil
}
// parseBaseParamFromRedirect 从logbox跳转地址提取登陆参数,并以该地址作为后续请求的Referer
func parseBaseParamFromRedirect(finalUrl *url.URL) (*BaseLoginParam, error) {
if finalUrl == nil {
return nil, errors.New("no login page redirect")
}
query := finalUrl.Query()
lt, reqId := query.Get("lt"), query.Get("reqId")
if lt == "" || reqId == "" {
return nil, errors.New("no lt/reqId in the login page redirect")
}
return &BaseLoginParam{
Lt: lt,
ReqId: reqId,
Referer: finalUrl.String(),
CaptchaToken: regexp.MustCompile(`'captchaToken' value='(.+?)'`).FindStringSubmatch(res.String())[1],
Lt: regexp.MustCompile(`lt = "(.+?)"`).FindStringSubmatch(res.String())[1],
ParamId: regexp.MustCompile(`paramId = "(.+?)"`).FindStringSubmatch(res.String())[1],
ReqId: regexp.MustCompile(`reqId = "(.+?)"`).FindStringSubmatch(res.String())[1],
}, nil
}
// parseBaseParamFromPage 兼容把参数写在页内变量里的老登陆页
func parseBaseParamFromPage(body string) (*BaseLoginParam, error) {
lt, err := matchLoginParam(body, `lt = "(.+?)"`, "lt")
if err != nil {
return nil, err
}
reqId, err := matchLoginParam(body, `reqId = "(.+?)"`, "reqId")
if err != nil {
return nil, err
}
paramId, err := matchLoginParam(body, `paramId = "(.+?)"`, "paramId")
if err != nil {
return nil, err
}
// 老页面才有内嵌的图形验证码token
captchaToken, _ := matchLoginParam(body, `'captchaToken' value='(.+?)'`, "captchaToken")
param := &BaseLoginParam{
CaptchaToken: captchaToken,
Lt: lt,
ParamId: paramId,
ReqId: reqId,
}
encryptUrl, _ := matchLoginParam(body, `encryptUrl = "(.+?)"`, "encryptUrl")
param.Referer = AUTH_URL + "/api/logbox/separate/web/index.html?" + strings.Join([]string{
"appId=" + url.QueryEscape(APP_ID),
"lt=" + url.QueryEscape(param.Lt),
"reqId=" + url.QueryEscape(param.ReqId),
}, "&")
if encryptUrl != "" {
param.Referer += "&encryptUrl=" + url.QueryEscape(encryptUrl)
}
return param, nil
}
// matchLoginParam 从登陆页面提取参数,缺失时返回可读的错误而不是panic
func matchLoginParam(body, pattern, name string) (string, error) {
matches := regexp.MustCompile(pattern).FindStringSubmatch(body)
if len(matches) < 2 {
return "", fmt.Errorf("failed to get %s from the login page", name)
}
return matches[1], nil
}
/* 初始化登陆需要的参数
* 如果遇到验证码返回错误
*/
@@ -783,13 +516,12 @@ func (y *Cloud189PC) initLoginParam() error {
}
y.loginParam.jRsaKey = fmt.Sprintf("-----BEGIN PUBLIC KEY-----\n%s\n-----END PUBLIC KEY-----", encryptConf.Data.PubKey)
y.loginParam.rsaPrefix = encryptConf.Data.Pre
y.loginParam.RsaUsername = y.loginParam.encryptSecret(y.Username)
y.loginParam.RsaPassword = y.loginParam.encryptSecret(y.Password)
y.loginParam.RsaUsername = encryptConf.Data.Pre + RsaEncrypt(y.loginParam.jRsaKey, y.Username)
y.loginParam.RsaPassword = encryptConf.Data.Pre + RsaEncrypt(y.loginParam.jRsaKey, y.Password)
// 判断是否需要验证码
resp, err := y.client.R().
SetHeaders(y.loginHeaders(y.loginParam.BaseLoginParam)).
SetHeader("REQID", y.loginParam.ReqId).
SetFormData(map[string]string{
"appKey": APP_ID,
"accountType": ACCOUNT_TYPE,
@@ -844,7 +576,6 @@ func (y *Cloud189PC) initQRCodeParam() (err error) {
var qrcodeParam QRLoginParam
_, err = y.client.R().
SetHeaders(y.loginHeaders(*baseParam)).
SetFormData(map[string]string{"appId": APP_ID}).
ForceContentType("application/json;charset=UTF-8").
SetResult(&qrcodeParam).
@@ -852,9 +583,6 @@ func (y *Cloud189PC) initQRCodeParam() (err error) {
if err != nil {
return err
}
if qrcodeParam.UUID == "" {
return errors.New("failed to get the QR code uuid")
}
qrcodeParam.BaseLoginParam = *baseParam
y.qrcodeParam = &qrcodeParam
@@ -875,7 +603,6 @@ func (y *Cloud189PC) refreshSessionWithRetry(retryCount int) (err error) {
_, err = y.client.R().
SetResult(&userSessionResp).SetError(&erron).
SetQueryParams(clientSuffix()).
SetQueryParams(y.deviceParams()).
SetQueryParams(map[string]string{
"appId": APP_ID,
"accessToken": y.tokenInfo.AccessToken,
@@ -916,11 +643,12 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
return errors.New("refresh token failed after maximum retries")
}
// 该接口刷新失败时以HTTP 200返回 result/msg,SetError不会触发,必须解析响应体判断
var tokenInfo RefreshTokenResp
var erron RespErr
var tokenInfo AppSessionResp
_, err = y.client.R().
SetResult(&tokenInfo).
ForceContentType("application/json;charset=UTF-8").
SetError(&erron).
SetFormData(map[string]string{
"clientId": APP_ID,
"refreshToken": y.tokenInfo.RefreshToken,
@@ -933,8 +661,7 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
}
// 如果刷新失败,返回错误给上层处理
if tokenInfo.HasError() {
refreshErr := tokenInfo.Error()
if erron.HasError() {
if y.Addition.RefreshToken != "" {
y.Addition.RefreshToken = ""
op.MustSaveDriverStorage(y)
@@ -942,11 +669,7 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
// 根据登录类型决定下一步行为
if y.LoginType == "qrcode" {
return fmt.Errorf("QR code session has expired, please re-scan the code to log in: %s", refreshErr)
}
// 没有账号密码时无法回退到完整登录,直接把刷新失败的原因返回
if y.Username == "" || y.Password == "" {
return errors.New(refreshErr)
return errors.New("QR code session has expired, please re-scan the code to log in")
}
// 密码登录模式下,尝试回退到完整登录
return y.login()
@@ -954,8 +677,7 @@ func (y *Cloud189PC) refreshTokenWithRetry(retryCount int) (err error) {
y.Addition.AccessToken = tokenInfo.AccessToken
y.Addition.RefreshToken = tokenInfo.RefreshToken
y.tokenInfo.AccessToken = tokenInfo.AccessToken
y.tokenInfo.RefreshToken = tokenInfo.RefreshToken
y.tokenInfo = &tokenInfo
op.MustSaveDriverStorage(y)
return y.refreshSessionWithRetry(retryCount + 1)
}
@@ -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"
}
-295
View File
@@ -1,295 +0,0 @@
package aliyundrive_open
import (
"context"
"fmt"
"io"
"math/rand/v2"
"net/http"
"strings"
"sync"
"time"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/errs"
anet "github.com/OpenListTeam/OpenList/v4/internal/net"
"github.com/OpenListTeam/OpenList/v4/internal/stream"
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
)
const (
defaultCallbackConcurrency = 1
callbackAcquireTimeout = time.Second
callbackRequestAttempts = 3
callbackRetryBaseDelay = 200 * time.Millisecond
callbackErrorBodyLimit = 64 << 10
)
var callbackLimiters = struct {
sync.Mutex
byUser map[string]*callbackLimiter
}{byUser: make(map[string]*callbackLimiter)}
type callbackLimiter struct {
userID string
mu sync.Mutex
active int
nextID uint64
registrations map[uint64]int
changed chan struct{}
}
type callbackRegistration struct {
limiter *callbackLimiter
id uint64
once sync.Once
}
type callbackPermit struct {
limiter *callbackLimiter
once sync.Once
}
func normalizeCallbackConcurrency(limit int) int {
if limit <= 0 {
return defaultCallbackConcurrency
}
return limit
}
func registerCallbackLimiter(userID string, limit int) *callbackRegistration {
callbackLimiters.Lock()
defer callbackLimiters.Unlock()
limiter := callbackLimiters.byUser[userID]
if limiter == nil {
limiter = &callbackLimiter{
userID: userID,
registrations: make(map[uint64]int),
changed: make(chan struct{}),
}
callbackLimiters.byUser[userID] = limiter
}
limiter.mu.Lock()
limiter.nextID++
id := limiter.nextID
limiter.registrations[id] = normalizeCallbackConcurrency(limit)
limiter.signalLocked()
limiter.mu.Unlock()
return &callbackRegistration{limiter: limiter, id: id}
}
func (r *callbackRegistration) unregister() {
if r == nil || r.limiter == nil {
return
}
r.once.Do(func() {
callbackLimiters.Lock()
defer callbackLimiters.Unlock()
r.limiter.mu.Lock()
delete(r.limiter.registrations, r.id)
r.limiter.signalLocked()
if len(r.limiter.registrations) == 0 && r.limiter.active == 0 {
delete(callbackLimiters.byUser, r.limiter.userID)
}
r.limiter.mu.Unlock()
})
}
func (r *callbackRegistration) acquire(ctx context.Context) (*callbackPermit, error) {
if r == nil || r.limiter == nil {
return nil, errs.NewErr(errs.TemporaryCapacity, "callback limiter is unavailable")
}
if err := ctx.Err(); err != nil {
return nil, err
}
waitCtx, cancel := context.WithTimeout(ctx, callbackAcquireTimeout)
defer cancel()
for {
r.limiter.mu.Lock()
if r.limiter.active < r.limiter.limitLocked() {
r.limiter.active++
r.limiter.mu.Unlock()
return &callbackPermit{limiter: r.limiter}, nil
}
changed := r.limiter.changed
r.limiter.mu.Unlock()
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-waitCtx.Done():
if err := ctx.Err(); err != nil {
return nil, err
}
return nil, errs.NewErr(errs.TemporaryCapacity, "timed out waiting for callback admission")
case <-changed:
}
}
}
func (l *callbackLimiter) limitLocked() int {
limit := 0
for _, registered := range l.registrations {
if limit == 0 || registered < limit {
limit = registered
}
}
return limit
}
func (l *callbackLimiter) signalLocked() {
close(l.changed)
l.changed = make(chan struct{})
}
func (p *callbackPermit) release() {
if p == nil || p.limiter == nil {
return
}
p.once.Do(func() {
callbackLimiters.Lock()
defer callbackLimiters.Unlock()
p.limiter.mu.Lock()
p.limiter.active--
p.limiter.signalLocked()
if len(p.limiter.registrations) == 0 && p.limiter.active == 0 {
delete(callbackLimiters.byUser, p.limiter.userID)
}
p.limiter.mu.Unlock()
})
}
func (d *AliyundriveOpen) callbackRegistration() *callbackRegistration {
if d.callback != nil {
return d.callback
}
if d.ref != nil {
return d.ref.callbackRegistration()
}
return nil
}
func (d *AliyundriveOpen) callbackRangeReader(url string, size int64) stream.RangeReaderFunc {
return func(ctx context.Context, requested http_range.Range) (io.ReadCloser, error) {
if requested.Length < 0 || requested.Start+requested.Length > size {
requested.Length = size - requested.Start
}
for attempt := 0; attempt < callbackRequestAttempts; attempt++ {
permit, err := d.callbackRegistration().acquire(ctx)
if err != nil {
return nil, err
}
body, retry, err := openCallbackRange(ctx, url, size, requested)
if !retry && err == nil {
return newCallbackBody(ctx, body, permit.release), nil
}
permit.release()
if !retry {
return nil, err
}
if attempt+1 == callbackRequestAttempts {
return nil, errs.NewErr(errs.TemporaryCapacity, "Aliyun callback concurrency limit rejected %d attempts", callbackRequestAttempts)
}
delay := callbackRetryBaseDelay << attempt
delay += time.Duration(rand.Int64N(int64(delay / 2)))
timer := time.NewTimer(delay)
select {
case <-ctx.Done():
timer.Stop()
return nil, ctx.Err()
case <-timer.C:
}
}
return nil, errs.NewErr(errs.TemporaryCapacity, "callback attempts exhausted")
}
}
func openCallbackRange(ctx context.Context, url string, size int64, requested http_range.Range) (io.ReadCloser, bool, error) {
requestHeader, _ := ctx.Value(conf.RequestHeaderKey).(http.Header)
header := anet.ProcessHeader(requestHeader, nil)
header = http_range.ApplyRangeToHttpHeader(requested, header)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, false, fmt.Errorf("create Aliyun callback request: %w", err)
}
req.Header = header
response, err := anet.HttpClient().Do(req)
if err != nil {
return nil, false, fmt.Errorf("Aliyun callback request failed: %w", err)
}
if response.StatusCode >= http.StatusBadRequest {
defer response.Body.Close()
body, readErr := io.ReadAll(io.LimitReader(response.Body, callbackErrorBodyLimit))
if readErr != nil {
return nil, false, fmt.Errorf("read Aliyun callback error response: %w", readErr)
}
if isCallbackCapacityRejection(response.StatusCode, body) {
return nil, true, nil
}
message := strings.ReplaceAll(strings.TrimSpace(string(body)), url, "<redacted>")
return nil, false, fmt.Errorf("Aliyun callback request failed: %w; response: %s", anet.HttpStatusCodeError(response.StatusCode), message)
}
if requested.Start == 0 && requested.Length == size || response.StatusCode == http.StatusPartialContent || callbackContentRangeStartsAt(response.Header, requested.Start) {
return response.Body, false, nil
}
if response.StatusCode == http.StatusOK {
body, rangeErr := anet.GetRangedHttpReader(response.Body, requested.Start, requested.Length)
if rangeErr != nil {
response.Body.Close()
return nil, false, rangeErr
}
return body, false, nil
}
return response.Body, false, nil
}
func isCallbackCapacityRejection(status int, body []byte) bool {
return status == http.StatusForbidden &&
strings.Contains(string(body), "RequestDeniedByCallback") &&
strings.Contains(string(body), "ExceedMaxConcurrency")
}
func callbackContentRangeStartsAt(header http.Header, offset int64) bool {
start, _, err := http_range.ParseContentRange(header.Get("Content-Range"))
return err == nil && start == offset
}
type callbackBody struct {
body io.ReadCloser
release func()
once sync.Once
mu sync.Mutex
stop func() bool
}
func newCallbackBody(ctx context.Context, body io.ReadCloser, release func()) *callbackBody {
b := &callbackBody{body: body, release: release}
stop := context.AfterFunc(ctx, func() { _ = b.Close() })
b.mu.Lock()
b.stop = stop
b.mu.Unlock()
return b
}
func (b *callbackBody) Read(p []byte) (int, error) {
n, err := b.body.Read(p)
if err != nil {
_ = b.Close()
}
return n, err
}
func (b *callbackBody) Close() error {
var err error
b.once.Do(func() {
b.mu.Lock()
stop := b.stop
b.mu.Unlock()
if stop != nil {
stop()
}
err = b.body.Close()
b.release()
})
return err
}
-365
View File
@@ -1,365 +0,0 @@
package aliyundrive_open
import (
"context"
"errors"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
"github.com/OpenListTeam/OpenList/v4/drivers/base"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/errs"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/stream"
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
)
func TestLinkSeparatesRedirectAndProxyRepresentations(t *testing.T) {
oldConf := conf.Conf
conf.Conf = &conf.Config{}
t.Cleanup(func() { conf.Conf = oldConf })
base.InitClient()
var server *httptest.Server
server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/adrive/v1.0/user/getDriveInfo":
_, _ = fmt.Fprint(w, `{"user_id":"user-1","resource_drive_id":"drive-1"}`)
case "/adrive/v1.0/openFile/getDownloadUrl":
_, _ = fmt.Fprintf(w, `{"url":%q}`, server.URL+"/callback")
default:
http.NotFound(w, r)
}
}))
defer server.Close()
oldAPIURL := API_URL
API_URL = server.URL
defer func() { API_URL = oldAPIURL }()
d := &AliyundriveOpen{Addition: Addition{AccessToken: "token"}}
if err := d.Init(t.Context()); err != nil {
t.Fatal(err)
}
defer d.Drop(context.Background())
if d.CallbackConcurrency != defaultCallbackConcurrency {
t.Fatalf("normalized callback concurrency = %d, want %d", d.CallbackConcurrency, defaultCallbackConcurrency)
}
link, err := d.Link(t.Context(), &model.Object{ID: "file-1", Name: "file", Size: 1}, model.LinkArgs{})
if err != nil {
t.Fatal(err)
}
if link.RangeReader == nil {
t.Fatal("proxy link must own callback acquisition through a range reader")
}
if _, ok := link.RangeReader.(stream.RateLimitRangeReaderFunc); !ok {
t.Fatalf("proxy range reader type = %T, want server-rate-limited reader", link.RangeReader)
}
direct, err := d.Link(t.Context(), &model.Object{ID: "file-1", Name: "file", Size: 1}, model.LinkArgs{Redirect: true})
if err != nil {
t.Fatal(err)
}
if direct.URL == "" || direct.RangeReader != nil {
t.Fatal("redirect link must remain URL-only")
}
}
func TestCallbackRangeHoldsPermitUntilBodyClose(t *testing.T) {
oldConf := conf.Conf
conf.Conf = &conf.Config{}
t.Cleanup(func() { conf.Conf = oldConf })
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Length", "1")
w.Header().Set("Content-Range", "bytes 0-0/1")
w.WriteHeader(http.StatusPartialContent)
_, _ = io.WriteString(w, "x")
}))
defer server.Close()
registration := registerCallbackLimiter(t.Name(), 1)
t.Cleanup(registration.unregister)
d := &AliyundriveOpen{callback: registration}
body, err := d.callbackRangeReader(server.URL, 1).RangeRead(t.Context(), http_range.Range{Length: 1})
if err != nil {
t.Fatal(err)
}
registration.limiter.mu.Lock()
active := registration.limiter.active
registration.limiter.mu.Unlock()
if active != 1 {
t.Fatalf("active callback bodies = %d, want 1", active)
}
if err := body.Close(); err != nil {
t.Fatal(err)
}
registration.limiter.mu.Lock()
active = registration.limiter.active
registration.limiter.mu.Unlock()
if active != 0 {
t.Fatalf("active callback bodies after Close = %d, want 0", active)
}
}
func TestCallbackLimiterUsesMinimumRegisteredLimit(t *testing.T) {
firstRegistration := registerCallbackLimiter(t.Name(), 2)
t.Cleanup(firstRegistration.unregister)
first, err := firstRegistration.acquire(t.Context())
if err != nil {
t.Fatal(err)
}
second, err := firstRegistration.acquire(t.Context())
if err != nil {
t.Fatal(err)
}
defer first.release()
defer second.release()
lowerRegistration := registerCallbackLimiter(t.Name(), 1)
t.Cleanup(lowerRegistration.unregister)
acquired := make(chan *callbackPermit, 1)
go func() {
permit, acquireErr := lowerRegistration.acquire(t.Context())
if acquireErr == nil {
acquired <- permit
}
}()
first.release()
select {
case permit := <-acquired:
permit.release()
t.Fatal("lowering the shared limit must wait for all excess bodies to drain")
case <-time.After(100 * time.Millisecond):
}
second.release()
select {
case permit := <-acquired:
permit.release()
case <-time.After(time.Second):
t.Fatal("admission did not resume after active bodies drained below the new limit")
}
}
func TestCallbackLimiterSeparatesUsers(t *testing.T) {
firstUser := registerCallbackLimiter(t.Name()+"-first", 1)
secondUser := registerCallbackLimiter(t.Name()+"-second", 1)
t.Cleanup(firstUser.unregister)
t.Cleanup(secondUser.unregister)
first, err := firstUser.acquire(t.Context())
if err != nil {
t.Fatal(err)
}
defer first.release()
second, err := secondUser.acquire(t.Context())
if err != nil {
t.Fatalf("independent user was blocked: %v", err)
}
second.release()
}
func TestCallbackLimiterReconfigureWaitsForOldBodies(t *testing.T) {
userID := t.Name()
oldRegistration := registerCallbackLimiter(userID, 2)
first, err := oldRegistration.acquire(t.Context())
if err != nil {
t.Fatal(err)
}
second, err := oldRegistration.acquire(t.Context())
if err != nil {
t.Fatal(err)
}
oldRegistration.unregister()
newRegistration := registerCallbackLimiter(userID, 1)
t.Cleanup(newRegistration.unregister)
acquired := make(chan *callbackPermit, 1)
go func() {
permit, acquireErr := newRegistration.acquire(t.Context())
if acquireErr == nil {
acquired <- permit
}
}()
first.release()
select {
case permit := <-acquired:
permit.release()
t.Fatal("reconfigured limiter admitted while an old body still occupied the new limit")
case <-time.After(100 * time.Millisecond):
}
second.release()
select {
case permit := <-acquired:
permit.release()
case <-time.After(time.Second):
t.Fatal("reconfigured limiter did not admit after old bodies drained")
}
}
func TestCallbackLimiterDistinguishesTimeoutAndCancellation(t *testing.T) {
registration := registerCallbackLimiter(t.Name(), 1)
t.Cleanup(registration.unregister)
permit, err := registration.acquire(t.Context())
if err != nil {
t.Fatal(err)
}
defer permit.release()
started := time.Now()
_, err = registration.acquire(t.Context())
if !errors.Is(err, errs.TemporaryCapacity) {
t.Fatalf("admission timeout error = %v, want TemporaryCapacity", err)
}
if time.Since(started) < callbackAcquireTimeout {
t.Fatal("admission timed out before the configured wait elapsed")
}
ctx, cancel := context.WithCancel(t.Context())
cancel()
_, err = registration.acquire(ctx)
if !errors.Is(err, context.Canceled) || errors.Is(err, errs.TemporaryCapacity) {
t.Fatalf("canceled admission error = %v, want only context.Canceled", err)
}
}
func TestCallbackCapacityRejectionRequiresBothExactMarkers(t *testing.T) {
tests := []struct {
name string
body string
want bool
}{
{name: "both", body: `{"code":"RequestDeniedByCallback","message":"ExceedMaxConcurrency"}`, want: true},
{name: "code only", body: `{"code":"RequestDeniedByCallback"}`},
{name: "message only", body: `{"message":"ExceedMaxConcurrency"}`},
{name: "case differs", body: `{"code":"requestdeniedbycallback","message":"ExceedMaxConcurrency"}`},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if got := isCallbackCapacityRejection(http.StatusForbidden, []byte(test.body)); got != test.want {
t.Fatalf("classification = %v, want %v", got, test.want)
}
})
}
if isCallbackCapacityRejection(http.StatusTooManyRequests, []byte(`RequestDeniedByCallback ExceedMaxConcurrency`)) {
t.Fatal("non-403 response must not be classified as callback capacity")
}
}
func TestCallbackRangeRetriesOnlyVerifiedCapacityRejections(t *testing.T) {
var requests atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
requests.Add(1)
w.WriteHeader(http.StatusForbidden)
_, _ = io.WriteString(w, `{"code":"RequestDeniedByCallback","message":"ExceedMaxConcurrency"}`)
}))
defer server.Close()
registration := registerCallbackLimiter(t.Name(), 1)
t.Cleanup(registration.unregister)
d := &AliyundriveOpen{callback: registration}
_, err := d.callbackRangeReader(server.URL+"?token=secret", 1).RangeRead(t.Context(), http_range.Range{Length: 1})
if !errors.Is(err, errs.TemporaryCapacity) {
t.Fatalf("verified rejection error = %v, want TemporaryCapacity", err)
}
if requests.Load() != callbackRequestAttempts {
t.Fatalf("requests = %d, want %d", requests.Load(), callbackRequestAttempts)
}
if strings.Contains(err.Error(), "secret") {
t.Fatal("capacity error leaked the signed callback URL")
}
permit, acquireErr := registration.acquire(t.Context())
if acquireErr != nil {
t.Fatalf("capacity retries leaked admission: %v", acquireErr)
}
permit.release()
requests.Store(0)
permanent := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
requests.Add(1)
w.WriteHeader(http.StatusForbidden)
_, _ = io.WriteString(w, `{"code":"RequestDeniedByCallback","message":"denied"}`)
}))
defer permanent.Close()
_, err = d.callbackRangeReader(permanent.URL+"?token=secret", 1).RangeRead(t.Context(), http_range.Range{Length: 1})
if errors.Is(err, errs.TemporaryCapacity) {
t.Fatalf("permanent 403 error = %v, must not be TemporaryCapacity", err)
}
if requests.Load() != 1 {
t.Fatalf("permanent 403 requests = %d, want 1", requests.Load())
}
if strings.Contains(err.Error(), "secret") {
t.Fatal("permanent error leaked the signed callback URL")
}
}
type countingReadCloser struct {
reader io.Reader
closed atomic.Int32
}
func (r *countingReadCloser) Read(p []byte) (int, error) { return r.reader.Read(p) }
func (r *countingReadCloser) Close() error {
r.closed.Add(1)
return nil
}
func TestCallbackBodyReleasesExactlyOnce(t *testing.T) {
underlying := &countingReadCloser{reader: strings.NewReader("x")}
var released atomic.Int32
body := newCallbackBody(t.Context(), underlying, func() { released.Add(1) })
_, _ = io.ReadAll(body)
if err := body.Close(); err != nil {
t.Fatal(err)
}
if err := body.Close(); err != nil {
t.Fatal(err)
}
if underlying.closed.Load() != 1 || released.Load() != 1 {
t.Fatalf("close count = %d, release count = %d; want 1, 1", underlying.closed.Load(), released.Load())
}
}
type failingReadCloser struct {
closed atomic.Int32
}
func (*failingReadCloser) Read([]byte) (int, error) { return 0, errors.New("read failed") }
func (r *failingReadCloser) Close() error {
r.closed.Add(1)
return nil
}
func TestCallbackBodyReadFailureReleasesPermit(t *testing.T) {
underlying := &failingReadCloser{}
var released atomic.Int32
body := newCallbackBody(t.Context(), underlying, func() { released.Add(1) })
if _, err := body.Read(make([]byte, 1)); err == nil {
t.Fatal("read unexpectedly succeeded")
}
if underlying.closed.Load() != 1 || released.Load() != 1 {
t.Fatalf("close count = %d, release count = %d; want 1, 1", underlying.closed.Load(), released.Load())
}
}
func TestCallbackBodyCancellationReleasesPermit(t *testing.T) {
ctx, cancel := context.WithCancel(t.Context())
underlying := &countingReadCloser{reader: strings.NewReader("x")}
released := make(chan struct{}, 1)
_ = newCallbackBody(ctx, underlying, func() { released <- struct{}{} })
cancel()
select {
case <-released:
case <-time.After(time.Second):
t.Fatal("context cancellation did not release callback admission")
}
if underlying.closed.Load() != 1 {
t.Fatalf("underlying close count = %d, want 1", underlying.closed.Load())
}
}
+4 -18
View File
@@ -11,7 +11,6 @@ import (
"github.com/OpenListTeam/OpenList/v4/internal/driver"
"github.com/OpenListTeam/OpenList/v4/internal/errs"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/stream"
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
"github.com/go-resty/resty/v2"
log "github.com/sirupsen/logrus"
@@ -23,9 +22,8 @@ type AliyundriveOpen struct {
DriveId string
limiter *limiter
ref *AliyundriveOpen
callback *callbackRegistration
limiter *limiter
ref *AliyundriveOpen
}
func (d *AliyundriveOpen) Config() driver.Config {
@@ -37,7 +35,6 @@ func (d *AliyundriveOpen) GetAddition() driver.Additional {
}
func (d *AliyundriveOpen) Init(ctx context.Context) error {
d.CallbackConcurrency = normalizeCallbackConcurrency(d.CallbackConcurrency)
d.limiter = getLimiterForUser(globalLimiterUserID) // First create a globally shared limiter to limit the initial requests.
if d.LIVPDownloadFormat == "" {
d.LIVPDownloadFormat = "jpeg"
@@ -55,7 +52,6 @@ func (d *AliyundriveOpen) Init(ctx context.Context) error {
userid := utils.Json.Get(res, "user_id").ToString()
d.limiter.free()
d.limiter = getLimiterForUser(userid) // Allocate a corresponding limiter for each user.
d.callback = registerCallbackLimiter(userid, d.CallbackConcurrency)
return nil
}
@@ -69,10 +65,6 @@ func (d *AliyundriveOpen) InitReference(storage driver.Driver) error {
}
func (d *AliyundriveOpen) Drop(ctx context.Context) error {
if d.callback != nil {
d.callback.unregister()
d.callback = nil
}
d.limiter.free()
d.limiter = nil
d.ref = nil
@@ -127,16 +119,10 @@ func (d *AliyundriveOpen) Link(ctx context.Context, file model.Obj, args model.L
url = utils.Json.Get(res, "streamsUrl", d.LIVPDownloadFormat).ToString()
}
exp := time.Minute
link := &model.Link{
return &model.Link{
URL: url,
Expiration: &exp,
}
if args.Redirect {
return link, nil
}
link.URL = ""
link.RangeReader = stream.RateLimitRangeReaderFunc(d.callbackRangeReader(url, file.GetSize()))
return link, nil
}, nil
}
func (d *AliyundriveOpen) MakeDir(ctx context.Context, parentDir model.Obj, dirName string) (model.Obj, error) {
+13 -14
View File
@@ -8,20 +8,19 @@ import (
type Addition struct {
DriveType string `json:"drive_type" type:"select" options:"default,resource,backup" default:"resource"`
driver.RootID
RefreshToken string `json:"refresh_token" required:"true"`
OrderBy string `json:"order_by" type:"select" options:"name,size,updated_at,created_at"`
OrderDirection string `json:"order_direction" type:"select" options:"ASC,DESC"`
UseOnlineAPI bool `json:"use_online_api" default:"true"`
AlipanType string `json:"alipan_type" required:"true" type:"select" default:"default" options:"default,alipanTV"`
APIAddress string `json:"api_url_address" default:"https://api.oplist.org/alicloud/renewapi"`
ClientID string `json:"client_id" help:"Keep it empty if you don't have one"`
ClientSecret string `json:"client_secret" help:"Keep it empty if you don't have one"`
RemoveWay string `json:"remove_way" required:"true" type:"select" options:"trash,delete"`
RapidUpload bool `json:"rapid_upload" help:"If you enable this option, the file will be uploaded to the server first, so the progress will be incorrect"`
InternalUpload bool `json:"internal_upload" help:"If you are using Aliyun ECS is located in Beijing, you can turn it on to boost the upload speed"`
LIVPDownloadFormat string `json:"livp_download_format" type:"select" options:"jpeg,mov" default:"jpeg"`
CallbackConcurrency int `json:"callback_concurrency" type:"number" default:"1" help:"Maximum active proxied downloads shared by Aliyun user"`
AccessToken string
RefreshToken string `json:"refresh_token" required:"true"`
OrderBy string `json:"order_by" type:"select" options:"name,size,updated_at,created_at"`
OrderDirection string `json:"order_direction" type:"select" options:"ASC,DESC"`
UseOnlineAPI bool `json:"use_online_api" default:"true"`
AlipanType string `json:"alipan_type" required:"true" type:"select" default:"default" options:"default,alipanTV"`
APIAddress string `json:"api_url_address" default:"https://api.oplist.org/alicloud/renewapi"`
ClientID string `json:"client_id" help:"Keep it empty if you don't have one"`
ClientSecret string `json:"client_secret" help:"Keep it empty if you don't have one"`
RemoveWay string `json:"remove_way" required:"true" type:"select" options:"trash,delete"`
RapidUpload bool `json:"rapid_upload" help:"If you enable this option, the file will be uploaded to the server first, so the progress will be incorrect"`
InternalUpload bool `json:"internal_upload" help:"If you are using Aliyun ECS is located in Beijing, you can turn it on to boost the upload speed"`
LIVPDownloadFormat string `json:"livp_download_format" type:"select" options:"jpeg,mov" default:"jpeg"`
AccessToken string
}
var config = driver.Config{
+88
View File
@@ -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
}
+35
View File
@@ -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
View File
@@ -315,36 +315,6 @@ var findDownPageParamReg = regexp.MustCompile(`<iframe.*?src="(.+?)"`)
// 获取文件ID
var findFileIDReg = regexp.MustCompile(`'/ajax(?:file|m)\.php\?file=(\d+)'`)
// 2026-10 改版:文件页将下载参数移入 /fn? 内页,接口变为 apifile 绝对地址并携带签名
var (
fnDomainReg = regexp.MustCompile(`var\s+domain[12]\s*=\s*'([^']*(?:ajaxfile|ajaxm)\.php\?file=(\d+)[^']*)'`)
fnSignReg = regexp.MustCompile(`var\s+wp_sign\s*=\s*'([^']*)'`)
fnAjaxDataReg = regexp.MustCompile(`var\s+ajaxdata\s*=\s*'([^']*)'`)
)
// parseFnPage 从改版后的 /fn? 内页提取下载接口地址与签名表单
// 对应页面 JS:POST domain1 {'action':'downprocess','websignkey':ajaxdata,'signs':ajaxdata,'sign':wp_sign,'websign':'2','kd':kdns,'ves':1}
func parseFnPage(pageData string) (string, map[string]string, error) {
matches := fnDomainReg.FindStringSubmatch(pageData)
if len(matches) < 3 {
return "", nil, fmt.Errorf("not find fn ajax url")
}
sign := fnSignReg.FindStringSubmatch(pageData)
ajaxdata := fnAjaxDataReg.FindStringSubmatch(pageData)
if len(sign) < 2 || len(ajaxdata) < 2 {
return "", nil, fmt.Errorf("not find fn sign")
}
return matches[1], map[string]string{
"action": "downprocess",
"websignkey": ajaxdata[1],
"signs": ajaxdata[1],
"sign": sign[1],
"websign": "2",
"kd": "1",
"ves": "1",
}, nil
}
// 获取分享链接主界面
func (d *LanZou) getShareUrlHtml(shareID string) (string, error) {
var vs string
@@ -467,23 +437,18 @@ func (d *LanZou) getFilesByShareUrl(shareID, pwd string, sharePageData string) (
return nil, err
}
nextPageData := RemoveNotes(string(data))
param, err = htmlJsonToMap(nextPageData)
if err != nil {
return nil, err
}
var resp FileShareInfoAndUrlResp[int]
matches := findFileIDReg.FindStringSubmatch(nextPageData)
if len(matches) >= 2 {
// 旧版结构:相对路径 /ajaxm.php?file=N
param, err = htmlJsonToMap(nextPageData)
if err != nil {
return nil, err
}
ajaxUrl := d.ShareUrl + matches[0][1:len(matches[0])-1]
_, err = d.post(ajaxUrl, func(req *resty.Request) { req.SetFormData(param) }, &resp)
} else if fnUrl, fnForm, ferr := parseFnPage(nextPageData); ferr == nil {
// 2026-10 改版结构:/fn? 内页携带 apifile 绝对地址与签名参数
_, err = d.post(fnUrl, func(req *resty.Request) { req.SetFormData(fnForm) }, &resp)
} else {
if len(matches) < 2 {
return nil, fmt.Errorf("not find file id")
}
ajaxUrl := d.ShareUrl + matches[0][1:len(matches[0])-1]
var resp FileShareInfoAndUrlResp[int]
_, err = d.post(ajaxUrl, func(req *resty.Request) { req.SetFormData(param) }, &resp)
if err != nil {
return nil, err
}
+1
View File
@@ -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,
+1 -2
View File
@@ -10,11 +10,10 @@ type Addition struct {
Username string `json:"username" required:"true"`
Password string `json:"password" required:"true"`
Platform string `json:"platform" required:"true" default:"web" type:"select" options:"android,web,pc"`
RefreshToken string `json:"refresh_token" required:"false" default:""`
RefreshToken string `json:"refresh_token" required:"true" default:""`
CaptchaToken string `json:"captcha_token" default:""`
DeviceID string `json:"device_id" required:"false" default:""`
DisableMediaLink bool `json:"disable_media_link" default:"true"`
SkipVerification bool `json:"skip_verification" default:"false" help:"ignore the human verification URL returned by the captcha API instead of failing; enabling this may trigger PikPak risk control"`
}
var config = driver.Config{
+58
View File
@@ -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
View File
@@ -100,13 +100,12 @@ func (d *PikPak) login() error {
return errors.New("username or password is empty")
}
// Clear expired access token so captcha requests don't carry a stale bearer
d.AccessToken = ""
url := "https://user.mypikpak.net/v1/auth/signin"
// Always refresh captcha token before signin (it may be expired)
if err := d.RefreshCaptchaTokenInLogin(GetAction(http.MethodPost, url), d.Username); err != nil {
return err
// 使用 用户填写的 CaptchaToken —————— (验证后的captcha_token)
if d.GetCaptchaToken() == "" {
if err := d.RefreshCaptchaTokenInLogin(GetAction(http.MethodPost, url), d.Username); err != nil {
return err
}
}
var e ErrResp
@@ -126,12 +125,7 @@ func (d *PikPak) login() error {
data := res.Body()
d.RefreshToken = jsoniter.Get(data, "refresh_token").ToString()
d.AccessToken = jsoniter.Get(data, "access_token").ToString()
if d.AccessToken == "" || d.RefreshToken == "" {
return errors.New("login failed: server returned empty tokens")
}
d.Common.SetUserID(jsoniter.Get(data, "sub").ToString())
d.Addition.RefreshToken = d.RefreshToken
op.MustSaveDriverStorage(d)
return nil
}
@@ -165,14 +159,9 @@ func (d *PikPak) refreshToken(refreshToken string) error {
return errors.New(e.Error())
}
data := res.Body()
newAccessToken := jsoniter.Get(data, "access_token").ToString()
newRefreshToken := jsoniter.Get(data, "refresh_token").ToString()
if newAccessToken == "" || newRefreshToken == "" {
return errors.New("refresh failed: server returned empty tokens")
}
d.Status = "work"
d.RefreshToken = newRefreshToken
d.AccessToken = newAccessToken
d.RefreshToken = jsoniter.Get(data, "refresh_token").ToString()
d.AccessToken = jsoniter.Get(data, "access_token").ToString()
d.Common.SetUserID(jsoniter.Get(data, "sub").ToString())
d.Addition.RefreshToken = d.RefreshToken
op.MustSaveDriverStorage(d)
@@ -208,18 +197,12 @@ func (d *PikPak) request(url string, method string, callback base.ReqCallback, r
case 0:
return res.Body(), nil
case 4122, 4121, 16:
if strings.Contains(url, "/v1/auth/") || strings.Contains(url, "/v1/shield/captcha/") {
return nil, errors.New(e.Error())
}
// access_token expired, refresh and retry
// access_token 过期
if err1 := d.refreshToken(d.RefreshToken); err1 != nil {
return nil, err1
}
return d.request(url, method, callback, resp)
case 9: // captcha token expired
if strings.Contains(url, "/v1/shield/captcha/") {
return nil, errors.New(e.Error())
}
case 9: // 验证码token过期
if err = d.RefreshCaptchaTokenAtLogin(GetAction(method, url), d.GetUserID()); err != nil {
return nil, err
}
@@ -386,9 +369,6 @@ func (d *PikPak) RefreshCaptchaTokenInLogin(action, username string) error {
} else {
metas["username"] = username
}
metas["client_version"] = d.ClientVersion
metas["package_name"] = d.PackageName
metas["timestamp"], metas["captcha_sign"] = d.Common.GetCaptchaSign()
return d.refreshCaptchaToken(action, metas)
}
@@ -427,7 +407,7 @@ func (d *PikPak) refreshCaptchaToken(action string, metas map[string]string) err
return errors.New(e.Error())
}
if resp.Url != "" && !d.Addition.SkipVerification {
if resp.Url != "" {
return fmt.Errorf(`need verify: <a target="_blank" href="%s">Click Here</a>`, resp.Url)
}
-720
View File
@@ -1,720 +0,0 @@
package pikpak
import (
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"github.com/OpenListTeam/OpenList/v4/drivers/base"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/db"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/glebarez/sqlite"
"github.com/go-resty/resty/v2"
"gorm.io/gorm"
)
// --- Helper function tests ---
func TestGetAction(t *testing.T) {
tests := []struct {
method string
url string
want string
}{
{"GET", "https://api-drive.mypikpak.net/drive/v1/files", "GET:/drive/v1/files"},
{"POST", "https://user.mypikpak.net/v1/auth/signin", "POST:/v1/auth/signin"},
{"POST", "https://user.mypikpak.net/v1/shield/captcha/init", "POST:/v1/shield/captcha/init"},
{"GET", "https://api-drive.mypikpak.net/drive/v1/files?page_token=abc", "GET:/drive/v1/files"},
{"POST", "https://user.mypikpak.net/v1/auth/token", "POST:/v1/auth/token"},
}
for _, tt := range tests {
t.Run(tt.method+":"+tt.url, func(t *testing.T) {
got := GetAction(tt.method, tt.url)
if got != tt.want {
t.Errorf("GetAction(%q, %q) = %q, want %q", tt.method, tt.url, got, tt.want)
}
})
}
}
func TestGetCaptchaSign(t *testing.T) {
c := &Common{
ClientID: "YNxT9w7GMdWvEOKa",
ClientVersion: "1.53.2",
PackageName: "com.pikcloud.pikpak",
DeviceID: "test-device-id",
Algorithms: AndroidAlgorithms,
}
timestamp, sign := c.GetCaptchaSign()
if timestamp == "" {
t.Fatal("timestamp should not be empty")
}
if len(sign) != 34 {
t.Fatalf("sign length should be 34 (\"1.\" + 32 hex), got %d: %q", len(sign), sign)
}
if sign[:2] != "1." {
t.Errorf("sign should start with '1.', got %q", sign[:2])
}
}
func TestGenerateDeviceSign(t *testing.T) {
sign := generateDeviceSign("test-device", "com.pikcloud.pikpak")
if len(sign) < 7 {
t.Fatal("device sign too short")
}
if sign[:7] != "div101." {
t.Errorf("device sign should start with 'div101.', got %q", sign[:7])
}
// Deterministic
if sign != generateDeviceSign("test-device", "com.pikcloud.pikpak") {
t.Error("generateDeviceSign should be deterministic")
}
}
func TestBuildCustomUserAgent(t *testing.T) {
ua := BuildCustomUserAgent("dev123", AndroidClientID, AndroidPackageName,
AndroidSdkVersion, AndroidClientVersion, AndroidPackageName, "user456")
for _, want := range []string{"ANDROID-", "clientid/", "deviceid/dev123", "usrno/user456"} {
if !strings.Contains(ua, want) {
t.Errorf("user agent should contain %q", want)
}
}
}
// --- Auth recovery behavior tests ---
func TestErrRespErrorClassification(t *testing.T) {
tests := []struct {
name string
resp ErrResp
wantError bool
wantCode int64
}{
{"success", ErrResp{ErrorCode: 0}, false, 0},
{"access_token_expired_4122", ErrResp{ErrorCode: 4122, ErrorMsg: "access_token expired"}, true, 4122},
{"access_token_expired_4121", ErrResp{ErrorCode: 4121, ErrorMsg: "access_token expired"}, true, 4121},
{"unauthenticated_16", ErrResp{ErrorCode: 16, ErrorMsg: "unauthenticated"}, true, 16},
{"refresh_token_invalid_4126", ErrResp{ErrorCode: 4126, ErrorMsg: "invalid_grant"}, true, 4126},
{"captcha_expired_9", ErrResp{ErrorCode: 9, ErrorMsg: "captcha_invalid"}, true, 9},
{"rate_limit_10", ErrResp{ErrorCode: 10, ErrorDescription: "too frequent"}, true, 10},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
gotError := tt.resp.IsError()
if gotError != tt.wantError {
t.Errorf("IsError() = %v, want %v", gotError, tt.wantError)
}
if tt.resp.ErrorCode != tt.wantCode {
t.Errorf("ErrorCode = %d, want %d", tt.resp.ErrorCode, tt.wantCode)
}
})
}
}
// TestGuardClauseOnAuthURLDoesNotRefresh verifies that when the auth endpoint
// itself reports 4122, request() fails fast instead of calling refreshToken()
// (which would recurse). Real behavior, real code path: with the guard
// removed from request(), the token endpoint would be hit a second time.
func TestGuardClauseOnAuthURLDoesNotRefresh(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
d, _ := newTestDriver(t)
m.tokenStatus = http.StatusBadRequest
m.tokenBody = map[string]any{"error_code": 4122, "error": "access_token_expired"}
_, err := d.request("https://user.mypikpak.net/v1/auth/token", http.MethodPost, nil, nil)
if err == nil {
t.Fatal("request() to an auth URL must fail on 4122 instead of refreshing")
}
if got := m.count(pathToken); got != 1 {
t.Errorf("guard clause violated: token endpoint hit %d times, want exactly 1 (no refreshToken recursion)", got)
}
if got := m.count(pathSignin); got != 0 {
t.Errorf("no re-login expected, got %d signin calls", got)
}
}
// --- Integration scaffolding: in-memory DB + mock PikPak endpoints ---
var (
setupDBOnce sync.Once
setupDBErr error
rowSeq int64
)
// setupTestDB mirrors internal/op/storage_test.go: an in-memory SQLite
// database behind internal/db, so op.MustSaveDriverStorage really persists
// and tests can assert on the saved row instead of on comments.
func setupTestDB(t *testing.T) {
t.Helper()
setupDBOnce.Do(func() {
var gormDB *gorm.DB
gormDB, setupDBErr = gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{})
if setupDBErr != nil {
return
}
conf.Conf = conf.DefaultConfig("testdata")
db.Init(gormDB)
})
if setupDBErr != nil {
t.Fatalf("failed to set up test database: %v", setupDBErr)
}
}
// createStorageRow inserts a fresh storage row and returns it, so that
// MustSaveDriverStorage during a test performs an UPDATE that can be read
// back afterwards.
func createStorageRow(t *testing.T) *model.Storage {
t.Helper()
rowSeq++
st := &model.Storage{
Driver: "PikPak",
MountPath: fmt.Sprintf("/pikpak-test-%d", rowSeq),
Addition: `{"username":"tester@example.com","password":"pw"}`,
}
if err := db.CreateStorage(st); err != nil {
t.Fatalf("failed to create storage row: %v", err)
}
return st
}
func persistedRefreshToken(t *testing.T, id uint) string {
t.Helper()
st, err := db.GetStorageById(id)
if err != nil {
t.Fatalf("failed to read storage back: %v", err)
}
var a Addition
if err := json.Unmarshal([]byte(st.Addition), &a); err != nil {
t.Fatalf("failed to decode persisted addition %q: %v", st.Addition, err)
}
return a.RefreshToken
}
// mockCall records one request received by the mock server.
type mockCall struct {
headers http.Header
body map[string]any
}
func (c mockCall) captchaToken() string {
s, _ := c.body["captcha_token"].(string)
return s
}
// pikpakMock emulates the captcha/auth endpoints used by login() and
// refreshToken(), plus one drive endpoint that serves as the entry point of
// the recovery chain. The drive endpoint fails exactly once (with the code
// configured in driveFirstStatus) and succeeds afterwards, so request() can
// only complete if recovery actually ran.
type pikpakMock struct {
t *testing.T
srv *httptest.Server
mu sync.Mutex
calls map[string][]mockCall
captchaTokenOut string
captchaURL string
tokenStatus int
tokenBody map[string]any
signinStatus int
signinBody map[string]any
driveFirstStatus int
driveFirstBody map[string]any // body served on the first drive call only
driveBody map[string]any // body served afterwards
driveHits int
}
func newPikpakMock(t *testing.T) *pikpakMock {
t.Helper()
m := &pikpakMock{
t: t,
calls: map[string][]mockCall{},
captchaTokenOut: "cap-fresh",
tokenStatus: http.StatusOK,
tokenBody: map[string]any{"access_token": "at-2", "refresh_token": "rt-2", "sub": "user-1"},
signinStatus: http.StatusOK,
signinBody: map[string]any{"access_token": "at-new", "refresh_token": "rt-new", "sub": "user-1"},
driveFirstStatus: http.StatusOK,
driveFirstBody: map[string]any{"files": []any{}, "next_page_token": ""},
driveBody: map[string]any{"files": []any{}, "next_page_token": ""},
}
m.srv = httptest.NewServer(http.HandlerFunc(m.serve))
return m
}
func (m *pikpakMock) close() { m.srv.Close() }
func (m *pikpakMock) serve(w http.ResponseWriter, r *http.Request) {
body := map[string]any{}
if raw, err := io.ReadAll(r.Body); err == nil && len(raw) > 0 {
_ = json.Unmarshal(raw, &body)
}
m.mu.Lock()
m.calls[r.URL.Path] = append(m.calls[r.URL.Path], mockCall{headers: r.Header.Clone(), body: body})
status := http.StatusOK
payload := any(map[string]any{})
switch {
case strings.HasSuffix(r.URL.Path, "/v1/shield/captcha/init"):
payload = map[string]any{"captcha_token": m.captchaTokenOut, "expires_in": 3600, "url": m.captchaURL}
case strings.HasSuffix(r.URL.Path, "/v1/auth/signin"):
status = m.signinStatus
payload = m.signinBody
case strings.HasSuffix(r.URL.Path, "/v1/auth/token"):
status = m.tokenStatus
payload = m.tokenBody
case strings.HasSuffix(r.URL.Path, "/drive/v1/files"):
m.driveHits++
if m.driveHits == 1 {
status = m.driveFirstStatus
payload = m.driveFirstBody
} else {
payload = m.driveBody
}
default:
m.mu.Unlock()
m.t.Errorf("unexpected request to %s", r.URL.Path)
w.WriteHeader(http.StatusNotFound)
return
}
m.mu.Unlock()
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(payload)
}
func (m *pikpakMock) count(path string) int {
m.mu.Lock()
defer m.mu.Unlock()
return len(m.calls[path])
}
func (m *pikpakMock) reset() {
m.mu.Lock()
defer m.mu.Unlock()
m.calls = map[string][]mockCall{}
m.driveHits = 0
}
func (m *pikpakMock) last(path string) mockCall {
m.mu.Lock()
defer m.mu.Unlock()
calls := m.calls[path]
if len(calls) == 0 {
m.t.Fatalf("no recorded call for %s", path)
}
return calls[len(calls)-1]
}
// installMockClient replaces base.RestyClient with a client whose requests to
// the hard-coded PikPak hosts are rewritten onto the mock server, and returns
// a restore function. The rewrite happens in OnBeforeRequest, which resty
// runs before its internal parseRequestURL/createHTTPRequest middlewares.
func installMockClient(m *pikpakMock) func() {
old := base.RestyClient
client := resty.New()
client.OnBeforeRequest(func(_ *resty.Client, req *resty.Request) error {
for _, host := range []string{"https://user.mypikpak.net", "https://api-drive.mypikpak.net"} {
if strings.HasPrefix(req.URL, host) {
req.URL = strings.Replace(req.URL, host, m.srv.URL, 1)
}
}
return nil
})
base.RestyClient = client
return func() { base.RestyClient = old }
}
// newTestDriver builds a PikPak with a fully initialized Common (web platform
// constants) and a fresh storage row in the DB, ready for auth-flow tests.
func newTestDriver(t *testing.T) (*PikPak, uint) {
t.Helper()
setupTestDB(t)
st := createStorageRow(t)
d := &PikPak{}
d.SetStorage(*st)
d.Platform = "web"
d.Username = "tester@example.com"
d.Password = "pw"
d.Common = &Common{
ClientID: WebClientID,
ClientSecret: WebClientSecret,
ClientVersion: WebClientVersion,
PackageName: WebPackageName,
DeviceID: "test-device",
UserAgent: "test-agent",
Algorithms: WebAlgorithms,
}
d.Common.RefreshCTokenCk = func(token string) {
d.Common.CaptchaToken = token
}
return d, st.ID
}
const (
pathCaptchaInit = "/v1/shield/captcha/init"
pathSignin = "/v1/auth/signin"
pathToken = "/v1/auth/token"
pathFiles = "/drive/v1/files"
)
// --- Main auth recovery path ---
// TestMainRecoveryPath exercises the full chain the PR is about: a drive
// request fails with 4122, refreshToken fails with 4126, login() runs (fresh
// captcha + password signin), the new refresh token is persisted to the DB,
// and request() retries the original call successfully with the new tokens.
func TestMainRecoveryPath(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
d, id := newTestDriver(t)
d.RefreshToken = "rt-old"
d.AccessToken = "at-stale"
d.SetCaptchaToken("cap-stale")
d.Addition.RefreshToken = "rt-old"
// refresh attempt fails with "refresh token invalid"
m.tokenStatus = http.StatusBadRequest
m.tokenBody = map[string]any{"error_code": 4126, "error": "invalid_grant"}
// the first drive call reports an expired access token; the retry succeeds
m.driveFirstStatus = http.StatusBadRequest
m.driveFirstBody = map[string]any{"error_code": 4122, "error": "access_token_expired"}
m.driveBody = map[string]any{"files": []any{}, "next_page_token": ""}
var resp Files
if _, err := d.request("https://api-drive.mypikpak.net/drive/v1/files", http.MethodGet, nil, &resp); err != nil {
t.Fatalf("request() returned error even though recovery should succeed: %v", err)
}
if got := m.count(pathToken); got != 1 {
t.Errorf("expected exactly 1 refresh request, got %d", got)
}
if got := m.count(pathSignin); got != 1 {
t.Errorf("expected exactly 1 signin (re-login), got %d", got)
}
if got := m.count(pathFiles); got != 2 {
t.Errorf("expected 2 files requests (failed + retried), got %d", got)
}
if got := m.count(pathCaptchaInit); got != 1 {
t.Errorf("expected exactly 1 captcha/init call during re-login, got %d", got)
}
// The retry must carry the tokens obtained via re-login, not the stale ones.
lastFiles := m.last(pathFiles)
if got := lastFiles.headers.Get("Authorization"); got != "Bearer at-new" {
t.Errorf("retried request Authorization = %q, want %q", got, "Bearer at-new")
}
if got := lastFiles.headers.Get("X-Captcha-Token"); got != "cap-fresh" {
t.Errorf("retried request X-Captcha-Token = %q, want %q", got, "cap-fresh")
}
// Tokens were rotated in memory...
if d.AccessToken != "at-new" {
t.Errorf("AccessToken = %q, want %q", d.AccessToken, "at-new")
}
if d.RefreshToken != "rt-new" {
t.Errorf("RefreshToken = %q, want %q", d.RefreshToken, "rt-new")
}
// ...and the rotated refresh token was really persisted.
if got := persistedRefreshToken(t, id); got != "rt-new" {
t.Errorf("persisted addition refresh_token = %q, want %q", got, "rt-new")
}
}
// TestRefreshToken4126WithoutCredentialsDoesNotLogin checks that a 4126 with
// empty username/password yields the "re-provide refresh_token" error instead
// of attempting a password login.
func TestRefreshToken4126WithoutCredentialsDoesNotLogin(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
d, _ := newTestDriver(t)
d.Username = ""
d.Password = ""
m.tokenStatus = http.StatusBadRequest
m.tokenBody = map[string]any{"error_code": 4126, "error": "invalid_grant"}
err := d.refreshToken("rt-old")
if err == nil {
t.Fatal("refreshToken() with invalid refresh token and no credentials must fail")
}
if !strings.Contains(err.Error(), "re-provide") {
t.Errorf("unexpected error text: %v", err)
}
if got := m.count(pathSignin); got != 0 {
t.Errorf("signin must not be attempted without credentials, got %d calls", got)
}
}
// TestRefreshTokenOtherErrorDoesNotLogin checks that a non-4126 refresh
// failure propagates without triggering a re-login (4126 is the single
// documented trigger).
func TestRefreshTokenOtherErrorDoesNotLogin(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
d, _ := newTestDriver(t)
m.tokenStatus = http.StatusBadRequest
m.tokenBody = map[string]any{"error_code": 10, "error_description": "too frequent"}
if err := d.refreshToken("rt-old"); err == nil {
t.Fatal("refreshToken() must propagate a non-4126 error")
}
if got := m.count(pathSignin); got != 0 {
t.Errorf("signin must not be attempted for non-4126 errors, got %d calls", got)
}
}
// --- Token validation (replaces TestTokenValidationRejectsEmpty) ---
// TestTokenValidationRejectsEmpty drives login() and refreshToken() against
// 200 responses that carry empty tokens and requires both paths to refuse
// them without persisting anything.
func TestTokenValidationRejectsEmpty(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
// login(): signin answers 200 but with an empty access_token.
d, id := newTestDriver(t)
m.signinBody = map[string]any{"access_token": "", "refresh_token": "rt-x", "sub": "user-1"}
if err := d.login(); err == nil {
t.Fatal("login() must reject empty access_token")
}
if got := persistedRefreshToken(t, id); got != "" {
t.Errorf("login() must not persist tokens when validation fails, persisted %q", got)
}
// login(): symmetric case — empty refresh_token but non-empty access_token.
d3, id3 := newTestDriver(t)
m.reset()
m.signinBody = map[string]any{"access_token": "at-x", "refresh_token": "", "sub": "user-1"}
if err := d3.login(); err == nil {
t.Fatal("login() must reject empty refresh_token")
}
if got := persistedRefreshToken(t, id3); got != "" {
t.Errorf("login() must not persist tokens when validation fails, persisted %q", got)
}
// refreshToken(): 200 but empty refresh_token.
d2, id2 := newTestDriver(t)
m.tokenStatus = http.StatusOK
m.tokenBody = map[string]any{"access_token": "at-x", "refresh_token": "", "sub": "user-1"}
if err := d2.refreshToken("rt-old"); err == nil {
t.Fatal("refreshToken() must reject empty refresh_token")
}
if got := persistedRefreshToken(t, id2); got != "" {
t.Errorf("refreshToken() must not persist tokens when validation fails, persisted %q", got)
}
}
// --- Captcha refresh (replaces TestCaptchaAlwaysRefreshedBeforeLogin) ---
// TestCaptchaAlwaysRefreshedBeforeLogin proves login() fetches a fresh captcha
// even when a (possibly expired) token is already present, and that signin is
// performed with the fresh token rather than the stale one.
func TestCaptchaAlwaysRefreshedBeforeLogin(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
d, _ := newTestDriver(t)
d.SetCaptchaToken("cap-stale") // non-empty and (conceptually) expired
if err := d.login(); err != nil {
t.Fatalf("login() failed: %v", err)
}
if got := m.count(pathCaptchaInit); got != 1 {
t.Fatalf("expected exactly 1 captcha/init call despite a non-empty stale token, got %d", got)
}
if got := m.last(pathSignin).captchaToken(); got != "cap-fresh" {
t.Errorf("signin used captcha_token %q, want the fresh %q", got, "cap-fresh")
}
if got := d.GetCaptchaToken(); got != "cap-fresh" {
t.Errorf("driver CaptchaToken = %q after login, want %q", got, "cap-fresh")
}
}
// --- Stale bearer cleared before login ---
// TestLoginClearsStaleAccessToken checks that the captcha/init and signin
// requests issued by login() do not carry the expired bearer token.
func TestLoginClearsStaleAccessToken(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
d, _ := newTestDriver(t)
d.AccessToken = "at-stale"
if err := d.login(); err != nil {
t.Fatalf("login() failed: %v", err)
}
for _, path := range []string{pathCaptchaInit, pathSignin} {
if got := m.last(path).headers.Get("Authorization"); got != "" {
t.Errorf("%s request carried Authorization %q, want it cleared before login", path, got)
}
}
}
// --- Captcha meta completeness ---
// TestCaptchaMetaCompleteness asserts captcha/init on the login path carries
// the same meta fields RefreshCaptchaTokenAtLogin sends on main.
func TestCaptchaMetaCompleteness(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
d, _ := newTestDriver(t)
if err := d.login(); err != nil {
t.Fatalf("login() failed: %v", err)
}
meta, _ := m.last(pathCaptchaInit).body["meta"].(map[string]any)
for _, key := range []string{"email", "client_version", "package_name", "timestamp", "captcha_sign"} {
if v, ok := meta[key]; !ok || v == "" {
t.Errorf("captcha meta missing or empty %q (got %#v)", key, meta)
}
}
}
// --- refreshToken success path (highest-frequency production path) ---
// TestRefreshTokenSuccessRotatesAndPersists covers 4122 -> refreshToken()
// succeeding: rotated tokens land in memory, the retry carries the new bearer,
// and the new refresh token is persisted to the DB.
func TestRefreshTokenSuccessRotatesAndPersists(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
d, id := newTestDriver(t)
d.RefreshToken = "rt-old"
d.AccessToken = "at-stale"
d.Addition.RefreshToken = "rt-old"
m.tokenStatus = http.StatusOK
m.tokenBody = map[string]any{"access_token": "at-2", "refresh_token": "rt-2", "sub": "user-1"}
m.driveFirstStatus = http.StatusBadRequest
m.driveFirstBody = map[string]any{"error_code": 4122, "error": "access_token_expired"}
m.driveBody = map[string]any{"files": []any{}, "next_page_token": ""}
var resp Files
if _, err := d.request("https://api-drive.mypikpak.net/drive/v1/files", http.MethodGet, nil, &resp); err != nil {
t.Fatalf("request() failed even though refresh should succeed: %v", err)
}
if got := m.count(pathSignin); got != 0 {
t.Errorf("a successful refresh must not fall through to password login, got %d signin calls", got)
}
if d.AccessToken != "at-2" {
t.Errorf("AccessToken = %q, want %q", d.AccessToken, "at-2")
}
if d.RefreshToken != "rt-2" {
t.Errorf("RefreshToken = %q, want %q", d.RefreshToken, "rt-2")
}
if got := m.last(pathFiles).headers.Get("Authorization"); got != "Bearer at-2" {
t.Errorf("retried request Authorization = %q, want %q", got, "Bearer at-2")
}
if got := persistedRefreshToken(t, id); got != "rt-2" {
t.Errorf("persisted addition refresh_token = %q, want %q", got, "rt-2")
}
}
// --- captcha expired (case 9) ---
// TestCaptchaExpiredRefreshesAndRetries covers request() case 9: a captcha
// error on a drive call triggers a captcha refresh and one retry.
func TestCaptchaExpiredRefreshesAndRetries(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
d, _ := newTestDriver(t)
d.AccessToken = "at-ok"
d.RefreshToken = "rt-ok"
d.SetCaptchaToken("cap-stale")
m.driveFirstStatus = http.StatusBadRequest
m.driveFirstBody = map[string]any{"error_code": 9, "error": "captcha_invalid"}
m.driveBody = map[string]any{"files": []any{}, "next_page_token": ""}
var resp Files
if _, err := d.request("https://api-drive.mypikpak.net/drive/v1/files", http.MethodGet, nil, &resp); err != nil {
t.Fatalf("request() failed even though captcha refresh should recover: %v", err)
}
if got := m.count(pathCaptchaInit); got == 0 {
t.Fatal("expected a captcha refresh after error code 9")
}
if got := m.count(pathFiles); got != 2 {
t.Errorf("expected 2 files requests (failed + retried), got %d", got)
}
if got := m.count(pathSignin); got != 0 {
t.Errorf("captcha recovery must not re-login, got %d signin calls", got)
}
if got := m.last(pathFiles).headers.Get("X-Captcha-Token"); got != "cap-fresh" {
t.Errorf("retried request X-Captcha-Token = %q, want %q", got, "cap-fresh")
}
}
// --- SkipVerification (added by this PR) ---
// TestSkipVerificationControlsVerificationURL covers the new config option:
// a captcha/init response carrying a human-verification url is fatal by
// default and ignored only when skip_verification is enabled.
func TestSkipVerificationControlsVerificationURL(t *testing.T) {
m := newPikpakMock(t)
defer m.close()
restore := installMockClient(m)
defer restore()
m.captchaURL = "https://user.mypikpak.net/forbidden/test"
d, _ := newTestDriver(t)
if err := d.login(); err == nil {
t.Fatal("login() must fail on a verification url by default")
} else if !strings.Contains(err.Error(), "need verify") {
t.Errorf("unexpected error: %v", err)
}
d2, _ := newTestDriver(t)
d2.SkipVerification = true
if err := d2.login(); err != nil {
t.Fatalf("login() with skip_verification must ignore the url, got: %v", err)
}
if d2.AccessToken != "at-new" {
t.Errorf("AccessToken = %q after skipped verification, want %q", d2.AccessToken, "at-new")
}
}
+60
View File
@@ -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
View File
@@ -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},
}
+3 -11
View File
@@ -55,18 +55,10 @@ func (d *Teldrive) Drop(ctx context.Context) error {
}
func (d *Teldrive) List(ctx context.Context, dir model.Obj, args model.ListArgs) ([]model.Obj, error) {
dirPath := dir.GetPath()
if dirPath == "" {
dirPath = d.GetRootPath()
}
if dirPath == "" {
dirPath = "/"
}
var firstResp ListResp
err := d.request(http.MethodGet, "/api/files", func(req *resty.Request) {
req.SetQueryParams(map[string]string{
"path": dirPath,
"path": dir.GetPath(),
"limit": "500",
"page": "1",
})
@@ -95,7 +87,7 @@ func (d *Teldrive) List(ctx context.Context, dir model.Obj, args model.ListArgs)
var resp ListResp
err := d.request(http.MethodGet, "/api/files", func(req *resty.Request) {
req.SetQueryParams(map[string]string{
"path": dirPath,
"path": dir.GetPath(),
"limit": "500",
"page": strconv.Itoa(page),
})
@@ -122,7 +114,7 @@ func (d *Teldrive) List(ctx context.Context, dir model.Obj, args model.ListArgs)
return utils.SliceConvert(allItems, func(src Object) (model.Obj, error) {
return &model.Object{
Path: path.Join(dirPath, src.Name),
Path: path.Join(dir.GetPath(), src.Name),
ID: src.ID,
Name: src.Name,
Size: func() int64 {
-43
View File
@@ -4,11 +4,9 @@ import (
"context"
"net/http"
"net/http/httptest"
"sync"
"testing"
"github.com/OpenListTeam/OpenList/v4/drivers/base"
"github.com/OpenListTeam/OpenList/v4/internal/driver"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/go-resty/resty/v2"
)
@@ -38,44 +36,3 @@ func TestListEmptyDir(t *testing.T) {
t.Fatalf("expected no entries for an empty dir, got %d", len(objs))
}
}
func TestListRootUsesConfiguredRootPath(t *testing.T) {
var (
mu sync.Mutex
paths []string
)
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
mu.Lock()
paths = append(paths, r.URL.Query().Get("path"))
mu.Unlock()
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"items":[{"id":"child","name":"child","type":"folder"}],"meta":{"count":1,"totalPages":1,"currentPage":1}}`))
}))
defer srv.Close()
oldClient := base.RestyClient
base.RestyClient = resty.New()
defer func() { base.RestyClient = oldClient }()
d := &Teldrive{
Addition: Addition{
RootPath: driver.RootPath{RootFolderPath: "/configured-root"},
},
}
d.Address = srv.URL
objs, err := d.List(context.Background(), &model.Object{}, model.ListArgs{})
if err != nil {
t.Fatalf("List returned error: %v", err)
}
mu.Lock()
defer mu.Unlock()
if len(paths) != 1 || paths[0] != "/configured-root" {
t.Fatalf("expected request path %q, got %q", "/configured-root", paths)
}
if len(objs) != 1 || objs[0].GetPath() != "/configured-root/child" {
t.Fatalf("expected child path %q, got %#v", "/configured-root/child", objs)
}
}
+56
View File
@@ -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
}
+55
View File
@@ -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
}
+55
View File
@@ -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
}
+1 -1
View File
@@ -10,7 +10,7 @@ require (
github.com/KarpelesLab/reflink v1.0.2
github.com/KirCute/zip v1.0.1
github.com/OpenListTeam/go-cache v0.1.0
github.com/OpenListTeam/gofakes3 v0.8.2
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
+2 -4
View File
@@ -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=
+2 -1
View File
@@ -6,12 +6,13 @@ import (
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/setting"
"github.com/OpenListTeam/OpenList/v4/server/common"
"github.com/gin-gonic/gin"
"github.com/go-webauthn/webauthn/webauthn"
)
func NewAuthnInstance(c *gin.Context) (*webauthn.WebAuthn, error) {
siteUrl, err := url.Parse(conf.GetApiUrl(c.Request.Context()))
siteUrl, err := url.Parse(common.GetApiUrl(c.Request.Context()))
if err != nil {
return nil, err
}
+8 -1
View File
@@ -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},
+1
View File
@@ -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
View File
@@ -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"
-8
View File
@@ -1,8 +0,0 @@
package conf
import "context"
func GetApiUrl(ctx context.Context) string {
api, _ := ctx.Value(ApiUrlKey).(string)
return api
}
-29
View File
@@ -1,29 +0,0 @@
package conf_test
import (
"context"
"testing"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
)
func TestGetApiUrl(t *testing.T) {
const want = "https://openlist.example"
tests := []struct {
name string
ctx context.Context
want string
}{
{name: "present", ctx: context.WithValue(context.Background(), conf.ApiUrlKey, want), want: want},
{name: "absent", ctx: context.Background()},
{name: "wrong type", ctx: context.WithValue(context.Background(), conf.ApiUrlKey, 1)},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := conf.GetApiUrl(tt.ctx); got != tt.want {
t.Fatalf("origin = %q, want %q", got, tt.want)
}
})
}
}
+1 -1
View File
@@ -35,7 +35,7 @@ func DeleteSearchNodesByParent(path string) error {
if err != nil {
return err
}
dir, name := stdpath.Dir(path), stdpath.Base(path)
dir, name := stdpath.Split(path)
return db.Where(fmt.Sprintf("%s = ? AND %s = ?",
columnName("parent"), columnName("name")),
dir, name).Delete(&model.SearchNode{}).Error
+52
View File
@@ -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
}
+72
View File
@@ -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
}
+12
View File
@@ -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")
)
-1
View File
@@ -18,7 +18,6 @@ var (
StorageNotInit = errors.New("storage not init")
StreamIncomplete = errors.New("upload/download stream incomplete, possible network issue")
StreamPeekFail = errors.New("StreamPeekFail")
TemporaryCapacity = errors.New("temporary capacity unavailable")
UnknownArchiveFormat = errors.New("unknown archive format")
WrongArchivePassword = errors.New("wrong archive password")
+2 -1
View File
@@ -21,6 +21,7 @@ import (
"github.com/OpenListTeam/OpenList/v4/internal/task"
"github.com/OpenListTeam/OpenList/v4/internal/task_group"
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
"github.com/OpenListTeam/OpenList/v4/server/common"
"github.com/OpenListTeam/tache"
"github.com/pkg/errors"
log "github.com/sirupsen/logrus"
@@ -414,7 +415,7 @@ func archiveDecompress(ctx context.Context, srcObjPath, dstDirPath string, args
return nil, err
} else {
tsk.Creator, _ = ctx.Value(conf.UserKey).(*model.User)
tsk.ApiUrl = conf.GetApiUrl(ctx)
tsk.ApiUrl = common.GetApiUrl(ctx)
ArchiveDownloadTaskManager.Add(tsk)
return tsk, nil
}
+2 -1
View File
@@ -14,6 +14,7 @@ import (
"github.com/OpenListTeam/OpenList/v4/internal/task"
"github.com/OpenListTeam/OpenList/v4/internal/task_group"
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
"github.com/OpenListTeam/OpenList/v4/server/common"
"github.com/OpenListTeam/tache"
"github.com/pkg/errors"
)
@@ -165,7 +166,7 @@ func transfer(ctx context.Context, taskType taskType, srcObjPath, dstDirPath str
}
t.Creator, _ = ctx.Value(conf.UserKey).(*model.User)
t.ApiUrl = conf.GetApiUrl(ctx)
t.ApiUrl = common.GetApiUrl(ctx)
if taskType == copy || taskType == merge {
CopyTaskManager.Add(t)
} else {
+2 -2
View File
@@ -4,9 +4,9 @@ import (
"context"
"strings"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/op"
"github.com/OpenListTeam/OpenList/v4/server/common"
"github.com/pkg/errors"
)
@@ -20,7 +20,7 @@ func link(ctx context.Context, path string, args model.LinkArgs) (*model.Link, m
return nil, nil, errors.WithMessage(err, "failed link")
}
if l.URL != "" && !strings.HasPrefix(l.URL, "http://") && !strings.HasPrefix(l.URL, "https://") {
l.URL = conf.GetApiUrl(ctx) + l.URL
l.URL = common.GetApiUrl(ctx) + l.URL
}
return l, obj, nil
}
+2 -1
View File
@@ -7,6 +7,7 @@ import (
"time"
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
"github.com/OpenListTeam/OpenList/v4/server/common"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/driver"
@@ -80,7 +81,7 @@ func putAsTask(ctx context.Context, dstDirPath string, file model.FileStreamer)
t := &UploadTask{
TaskExtension: task.TaskExtension{
Creator: taskCreator,
ApiUrl: conf.GetApiUrl(ctx),
ApiUrl: common.GetApiUrl(ctx),
},
storage: storage,
dstDirActualPath: dstDirActualPath,
+603
View File
@@ -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
}
+18
View File
@@ -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
}
+1
View File
@@ -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
}
+5 -7
View File
@@ -206,8 +206,7 @@ func (d *downloader) download() (io.ReadCloser, error) {
if err != nil {
d.cancel(err)
d.cfg.ConcurrencyLimit.Release()
_ = d.interrupt()
return nil, err
return nil, d.interrupt()
}
d.mu.Lock()
@@ -269,6 +268,10 @@ func (d *downloader) sendChunkTask(newConcurrency bool) (err error) {
if err != nil {
return err // 分片算法错误或者下载中断
}
if newConcurrency {
go d.downloadPart()
d.concurrency--
}
ch := chunk{
start: d.pos,
size: finalSize,
@@ -283,11 +286,6 @@ func (d *downloader) sendChunkTask(newConcurrency bool) (err error) {
case <-d.ctx.Done():
return context.Cause(d.ctx)
case d.chunkCh <- ch:
if newConcurrency {
// The worker owns the acquired slot only after its chunk is queued.
go d.downloadPart()
d.concurrency--
}
return nil
}
}
-88
View File
@@ -1,88 +0,0 @@
package net
import (
"context"
"errors"
"net/http"
"testing"
"time"
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
)
func TestDownloadCancelledAcquisitionReturnsErrorAndReleasesLimit(t *testing.T) {
const attempts = 32
limits := make([]*ConcurrencyLimit, 0, attempts)
for range attempts {
limit := &ConcurrencyLimit{Limit: 1}
limits = append(limits, limit)
d := NewDownloader(func(d *Downloader) {
d.Concurrency = 2
d.PartSize = 4
d.ConcurrencyLimit = limit
d.HttpClient = func(ctx context.Context, _ *HttpRequestParams) (*http.Response, error) {
return nil, ctx.Err()
}
})
ctx, cancel := context.WithCancel(context.Background())
cancel()
reader, err := d.Download(ctx, &HttpRequestParams{Size: 16, Range: http_range.Range{Length: -1}})
if reader == nil && err == nil {
t.Error("cancelled download returned a nil reader and nil error")
}
if reader != nil {
_ = reader.Close()
} else if !errors.Is(err, context.Canceled) {
t.Errorf("cancelled download error = %v, want context.Canceled", err)
}
}
time.Sleep(50 * time.Millisecond) // allow any started workers to release their slots
for i, limit := range limits {
limit.mu.Lock()
got := limit.Limit
limit.mu.Unlock()
if got != 1 {
t.Errorf("attempt %d remaining concurrency = %d, want 1", i, got)
}
}
}
func TestDownloadSinglePartFailureReleasesLimit(t *testing.T) {
upstreamErr := errors.New("upstream failure")
for _, tc := range []struct {
name string
ctx context.Context
want error
}{
{name: "cancelled", ctx: func() context.Context {
ctx, cancel := context.WithCancel(context.Background())
cancel()
return ctx
}(), want: context.Canceled},
{name: "upstream failure", ctx: context.Background(), want: upstreamErr},
} {
t.Run(tc.name, func(t *testing.T) {
limit := &ConcurrencyLimit{Limit: 1}
d := NewDownloader(func(d *Downloader) {
d.PartSize = 32
d.ConcurrencyLimit = limit
d.HttpClient = func(ctx context.Context, _ *HttpRequestParams) (*http.Response, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
return nil, upstreamErr
}
})
reader, err := d.Download(tc.ctx, &HttpRequestParams{Size: 16, Range: http_range.Range{Length: -1}})
if reader != nil || !errors.Is(err, tc.want) {
t.Fatalf("single-part failed download = %v, %v; want nil, %v", reader, err, tc.want)
}
limit.mu.Lock()
got := limit.Limit
limit.mu.Unlock()
if got != 1 {
t.Errorf("remaining concurrency = %d, want 1", got)
}
})
}
}
+26 -89
View File
@@ -4,7 +4,6 @@ import (
"compress/gzip"
"context"
"crypto/tls"
stderrors "errors"
"fmt"
"io"
"mime/multipart"
@@ -16,6 +15,7 @@ import (
"time"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/errs"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
@@ -25,8 +25,12 @@ import (
//this file is inspired by GO_SDK net.http.ServeContent
//type RangeReadCloser struct {
// GetReaderForRange RangeReaderFunc
//}
// ServeHTTP replies to the request using the content in the
// provided range reader. The main benefit of ServeHTTP over io.Copy
// provided RangeReadCloser. The main benefit of ServeHTTP over io.Copy
// is that it handles Range requests properly, sets the MIME type, and
// handles If-Match, If-Unmodified-Since, If-None-Match, If-Modified-Since,
// and If-Range requests.
@@ -43,11 +47,13 @@ import (
// request includes an If-Modified-Since header, ServeHTTP uses
// modtime to decide whether the content needs to be sent at all.
//
// The content's RangeRead method must return a reader for the requested range.
// The content's RangeReadCloser method must work: ServeHTTP gives a range,
// caller will give the reader for that Range.
//
// If the caller has set w's ETag header formatted per RFC 7232, section 2.3,
// ServeHTTP uses it to handle requests using If-Match, If-None-Match, or If-Range.
func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time.Time, size int64, rangeReader model.RangeReaderIF) (err error) {
func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time.Time, size int64, RangeReadCloser model.RangeReadCloserIF) error {
defer RangeReadCloser.Close()
setLastModified(w, modTime)
done, rangeReq := checkPreconditions(w, r, modTime)
if done {
@@ -107,11 +113,10 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
ctx := r.Context()
switch {
case len(ranges) == 0:
reader, err := openRange(ctx, rangeReader, http_range.Range{Length: -1})
reader, err := RangeReadCloser.RangeRead(ctx, http_range.Range{Length: -1})
if err != nil {
code = http.StatusRequestedRangeNotSatisfiable
var statusCode HttpStatusCodeError
if errors.As(err, &statusCode) {
if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok {
code = int(statusCode)
}
http.Error(w, err.Error(), code)
@@ -131,11 +136,10 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
// does not request multiple parts might not support
// multipart responses."
ra := ranges[0]
sendContent, err = openRange(ctx, rangeReader, ra)
sendContent, err = RangeReadCloser.RangeRead(ctx, ra)
if err != nil {
code = http.StatusRequestedRangeNotSatisfiable
var statusCode HttpStatusCodeError
if errors.As(err, &statusCode) {
if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok {
code = int(statusCode)
}
http.Error(w, err.Error(), code)
@@ -155,6 +159,7 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
mw := multipart.NewWriter(pw)
w.Header().Set("Content-Type", "multipart/byteranges; boundary="+mw.Boundary())
sendContent = pr
defer pr.Close() // cause writing goroutine to fail and exit if CopyN doesn't finish.
go func() {
for _, ra := range ranges {
part, err := mw.CreatePart(ra.MimeHeader(contentType, size))
@@ -162,18 +167,21 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
pw.CloseWithError(err)
return
}
if err := copyRange(ctx, part, rangeReader, ra); err != nil {
reader, err := RangeReadCloser.RangeRead(ctx, ra)
if err != nil {
pw.CloseWithError(err)
return
}
if _, err := utils.CopyWithBufferN(part, reader, ra.Length); err != nil {
pw.CloseWithError(err)
return
}
}
_ = pw.CloseWithError(mw.Close())
mw.Close()
pw.Close()
}()
}
defer func() {
err = closeWithError(err, sendContent)
}()
w.Header().Set("Accept-Ranges", "bytes")
if w.Header().Get("Content-Encoding") == "" {
@@ -193,8 +201,7 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
log.Warnf("Maybe size incorrect or reader not giving correct/full data, or connection closed before finish. written bytes: %d ,sendSize:%d, ", written, sendSize)
}
code = http.StatusInternalServerError
var statusCode HttpStatusCodeError
if errors.As(err, &statusCode) {
if statusCode, ok := errs.UnwrapOrSelf(err).(HttpStatusCodeError); ok {
code = int(statusCode)
}
w.WriteHeader(code)
@@ -203,86 +210,16 @@ func ServeHTTP(w http.ResponseWriter, r *http.Request, name string, modTime time
}
return nil
}
func copyRange(ctx context.Context, dst io.Writer, rangeReader model.RangeReaderIF, requested http_range.Range) (err error) {
reader, err := openRange(ctx, rangeReader, requested)
if err != nil {
return err
}
defer func() {
err = closeWithError(err, reader)
}()
_, err = utils.CopyWithBufferN(dst, reader, requested.Length)
return err
}
func openRange(ctx context.Context, rangeReader model.RangeReaderIF, requested http_range.Range) (io.ReadCloser, error) {
reader, err := rangeReader.RangeRead(ctx, requested)
if err != nil {
if reader != nil {
err = closeWithError(err, reader)
}
return nil, err
}
if reader == nil {
return nil, errors.New("range reader returned a nil body")
}
return reader, nil
}
func closeWithError(err error, closer io.Closer) error {
closeErr := closer.Close()
if err == nil {
return closeErr
}
if closeErr == nil {
return err
}
return stderrors.Join(err, closeErr)
}
// unsafeProxyHeaders are never forwarded from the client request to the
// upstream storage, regardless of the proxy_ignore_headers setting. They either
// carry the caller's credentials, describe the hop to this server rather than
// the hop to upstream, or let the caller influence how upstream routes and
// authenticates the request.
var unsafeProxyHeaders = map[string]struct{}{
"authorization": {},
"cookie": {},
"proxy-authorization": {},
"www-authenticate": {},
"host": {},
"referer": {},
"origin": {},
"connection": {},
"keep-alive": {},
"proxy-connection": {},
"te": {},
"trailer": {},
"transfer-encoding": {},
"upgrade": {},
"forwarded": {},
"x-forwarded-for": {},
"x-forwarded-host": {},
"x-forwarded-proto": {},
"x-real-ip": {},
}
func ProcessHeader(origin, override http.Header) http.Header {
result := http.Header{}
// client header
for h, val := range origin {
lower := strings.ToLower(h)
if _, unsafe := unsafeProxyHeaders[lower]; unsafe {
continue
}
if utils.SliceContains(conf.SlicesMap[conf.ProxyIgnoreHeaders], lower) {
if utils.SliceContains(conf.SlicesMap[conf.ProxyIgnoreHeaders], strings.ToLower(h)) {
continue
}
result[h] = val
}
// needed header, produced by the storage driver rather than the client
// needed header
for h, val := range override {
result[h] = val
}
-67
View File
@@ -1,67 +0,0 @@
package net
import (
"net/http"
"testing"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
)
// The client must not be able to smuggle credential or routing headers into the
// request that this server makes to the upstream storage, even when the
// proxy_ignore_headers setting has been emptied.
func TestProcessHeaderDropsUnsafeClientHeaders(t *testing.T) {
conf.SlicesMap[conf.ProxyIgnoreHeaders] = nil
origin := http.Header{}
origin.Set("Authorization", "Bearer victim-token")
origin.Set("Cookie", "session=victim")
origin.Set("X-Forwarded-For", "127.0.0.1")
origin.Set("Host", "internal.example")
origin.Set("Range", "bytes=0-1023")
result := ProcessHeader(origin, nil)
for _, h := range []string{"Authorization", "Cookie", "X-Forwarded-For", "Host"} {
if got := result.Get(h); got != "" {
t.Errorf("header %q must not be forwarded upstream, got %q", h, got)
}
}
if got := result.Get("Range"); got != "bytes=0-1023" {
t.Errorf("Range must be preserved, got %q", got)
}
}
// Headers supplied by the storage driver still win, since they carry the
// credentials needed to reach upstream.
func TestProcessHeaderOverrideWins(t *testing.T) {
conf.SlicesMap[conf.ProxyIgnoreHeaders] = nil
origin := http.Header{}
origin.Set("Authorization", "Bearer victim-token")
override := http.Header{}
override.Set("Authorization", "Bearer driver-token")
result := ProcessHeader(origin, override)
if got := result.Get("Authorization"); got != "Bearer driver-token" {
t.Errorf("driver header must be used, got %q", got)
}
}
func TestProcessHeaderStillHonoursIgnoreSetting(t *testing.T) {
conf.SlicesMap[conf.ProxyIgnoreHeaders] = []string{"x-custom"}
t.Cleanup(func() { conf.SlicesMap[conf.ProxyIgnoreHeaders] = nil })
origin := http.Header{}
origin.Set("X-Custom", "drop-me")
origin.Set("X-Keep", "keep-me")
result := ProcessHeader(origin, nil)
if got := result.Get("X-Custom"); got != "" {
t.Errorf("configured ignore header must be dropped, got %q", got)
}
if got := result.Get("X-Keep"); got != "keep-me" {
t.Errorf("unrelated header must be preserved, got %q", got)
}
}
-236
View File
@@ -1,236 +0,0 @@
package net
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"mime"
"mime/multipart"
"net/http"
"net/http/httptest"
"reflect"
"sync"
"testing"
"time"
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
)
func TestServeHTTPClosesMultipartRangeBeforeOpeningNext(t *testing.T) {
source := newSequentialRangeSource("abc")
ctx, cancel := context.WithTimeout(t.Context(), 500*time.Millisecond)
defer cancel()
request := httptest.NewRequest(http.MethodGet, "/file", nil).WithContext(ctx)
request.Header.Set("Range", "bytes=0-0,2-2")
recorder := httptest.NewRecorder()
if err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source); err != nil {
t.Fatalf("ServeHTTP() error = %v", err)
}
response := recorder.Result()
defer response.Body.Close()
if response.StatusCode != http.StatusPartialContent {
t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusPartialContent)
}
mediaType, params, err := mime.ParseMediaType(response.Header.Get("Content-Type"))
if err != nil {
t.Fatalf("parse Content-Type: %v", err)
}
if mediaType != "multipart/byteranges" {
t.Fatalf("Content-Type = %q, want multipart/byteranges", mediaType)
}
multipartReader := multipart.NewReader(response.Body, params["boundary"])
var parts []string
for {
part, err := multipartReader.NextPart()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
t.Fatalf("read multipart part: %v", err)
}
body, err := io.ReadAll(part)
if err != nil {
t.Fatalf("read multipart body: %v", err)
}
parts = append(parts, string(body))
}
if want := []string{"a", "c"}; !reflect.DeepEqual(parts, want) {
t.Fatalf("multipart parts = %q, want %q", parts, want)
}
assertRangeLifecycle(t, source, []string{"open:0", "close:0", "open:2", "close:2"}, []int{1, 1})
}
func TestServeHTTPClosesSelectedRangeBody(t *testing.T) {
tests := []struct {
name string
method string
rangeValue string
wantStatus int
wantEvents []string
}{
{name: "full", method: http.MethodGet, wantStatus: http.StatusOK, wantEvents: []string{"open:0", "close:0"}},
{name: "single range", method: http.MethodGet, rangeValue: "bytes=1-1", wantStatus: http.StatusPartialContent, wantEvents: []string{"open:1", "close:1"}},
{name: "head", method: http.MethodHead, wantStatus: http.StatusOK, wantEvents: []string{"open:0", "close:0"}},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
source := newSequentialRangeSource("abc")
request := httptest.NewRequest(test.method, "/file", nil)
if test.rangeValue != "" {
request.Header.Set("Range", test.rangeValue)
}
recorder := httptest.NewRecorder()
if err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source); err != nil {
t.Fatalf("ServeHTTP() error = %v", err)
}
if recorder.Code != test.wantStatus {
t.Fatalf("status = %d, want %d", recorder.Code, test.wantStatus)
}
assertRangeLifecycle(t, source, test.wantEvents, []int{1})
})
}
}
func TestServeHTTPClosesRangeAfterWriteFailure(t *testing.T) {
writeErr := errors.New("write failed")
source := newSequentialRangeSource("abc")
request := httptest.NewRequest(http.MethodGet, "/file", nil)
writer := &failingResponseWriter{header: make(http.Header), err: writeErr}
err := ServeHTTP(writer, request, "file.txt", time.Time{}, 3, source)
if !errors.Is(err, writeErr) {
t.Fatalf("ServeHTTP() error = %v, want %v", err, writeErr)
}
assertRangeLifecycle(t, source, []string{"open:0", "close:0"}, []int{1})
}
func TestServeHTTPClosesBodyReturnedWithOpenError(t *testing.T) {
source := newSequentialRangeSource("abc")
source.openErr = HttpStatusCodeError(http.StatusServiceUnavailable)
source.closeErr = errors.New("close failed")
request := httptest.NewRequest(http.MethodGet, "/file", nil)
recorder := httptest.NewRecorder()
if err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source); err != nil {
t.Fatalf("ServeHTTP() error = %v", err)
}
if recorder.Code != http.StatusServiceUnavailable {
t.Fatalf("status = %d, want %d", recorder.Code, http.StatusServiceUnavailable)
}
assertRangeLifecycle(t, source, []string{"open:0", "close:0"}, []int{1})
}
func TestServeHTTPStopsMultipartAfterRangeCloseFailure(t *testing.T) {
closeErr := errors.New("close failed")
source := newSequentialRangeSource("abc")
source.closeErr = closeErr
request := httptest.NewRequest(http.MethodGet, "/file", nil)
request.Header.Set("Range", "bytes=0-0,2-2")
recorder := httptest.NewRecorder()
err := ServeHTTP(recorder, request, "file.txt", time.Time{}, 3, source)
if !errors.Is(err, closeErr) {
t.Fatalf("ServeHTTP() error = %v, want %v", err, closeErr)
}
assertRangeLifecycle(t, source, []string{"open:0", "close:0"}, []int{1})
}
type sequentialRangeSource struct {
content []byte
permit chan struct{}
mu sync.Mutex
events []string
closeCounts []int
closeErr error
openErr error
}
func newSequentialRangeSource(content string) *sequentialRangeSource {
return &sequentialRangeSource{
content: []byte(content),
permit: make(chan struct{}, 1),
}
}
func (s *sequentialRangeSource) RangeRead(ctx context.Context, requested http_range.Range) (io.ReadCloser, error) {
select {
case s.permit <- struct{}{}:
case <-ctx.Done():
return nil, ctx.Err()
}
start := int(requested.Start)
length := int(requested.Length)
if length < 0 || start+length > len(s.content) {
length = len(s.content) - start
}
end := start + length
s.mu.Lock()
index := len(s.closeCounts)
s.events = append(s.events, fmt.Sprintf("open:%d", requested.Start))
s.closeCounts = append(s.closeCounts, 0)
s.mu.Unlock()
return &testReadCloser{
Reader: bytes.NewReader(s.content[start:end]),
close: func() error {
s.mu.Lock()
s.closeCounts[index]++
closeCalls := s.closeCounts[index]
if closeCalls == 1 {
s.events = append(s.events, fmt.Sprintf("close:%d", requested.Start))
}
s.mu.Unlock()
if closeCalls != 1 {
return fmt.Errorf("body closed %d times", closeCalls)
}
<-s.permit
return s.closeErr
},
}, s.openErr
}
func (s *sequentialRangeSource) eventsSnapshot() []string {
s.mu.Lock()
defer s.mu.Unlock()
return append([]string(nil), s.events...)
}
func (s *sequentialRangeSource) closeCountsSnapshot() []int {
s.mu.Lock()
defer s.mu.Unlock()
return append([]int(nil), s.closeCounts...)
}
func assertRangeLifecycle(t *testing.T, source *sequentialRangeSource, wantEvents []string, wantCloseCounts []int) {
t.Helper()
if got := source.eventsSnapshot(); !reflect.DeepEqual(got, wantEvents) {
t.Fatalf("range lifecycle = %v, want %v", got, wantEvents)
}
if got := source.closeCountsSnapshot(); !reflect.DeepEqual(got, wantCloseCounts) {
t.Fatalf("close counts = %v, want %v", got, wantCloseCounts)
}
}
type failingResponseWriter struct {
header http.Header
err error
}
func (w *failingResponseWriter) Header() http.Header { return w.header }
func (*failingResponseWriter) WriteHeader(int) {}
func (w *failingResponseWriter) Write([]byte) (int, error) {
return 0, w.err
}
type testReadCloser struct {
io.Reader
close func() error
}
func (b *testReadCloser) Close() error { return b.close() }
+2 -1
View File
@@ -25,6 +25,7 @@ import (
"github.com/OpenListTeam/OpenList/v4/internal/op"
"github.com/OpenListTeam/OpenList/v4/internal/setting"
"github.com/OpenListTeam/OpenList/v4/internal/task"
"github.com/OpenListTeam/OpenList/v4/server/common"
"github.com/google/uuid"
"github.com/pkg/errors"
)
@@ -183,7 +184,7 @@ func AddURL(ctx context.Context, args *AddURLArgs) (task.TaskExtensionInfo, erro
t := &DownloadTask{
TaskExtension: task.TaskExtension{
Creator: taskCreator,
ApiUrl: conf.GetApiUrl(ctx),
ApiUrl: common.GetApiUrl(ctx),
},
Url: args.URL,
DstDirPath: args.DstDirPath,
+3 -2
View File
@@ -20,6 +20,7 @@ import (
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
"github.com/OpenListTeam/OpenList/v4/pkg/torrent"
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
"github.com/OpenListTeam/OpenList/v4/server/common"
"github.com/OpenListTeam/tache"
"github.com/pkg/errors"
log "github.com/sirupsen/logrus"
@@ -139,7 +140,7 @@ func transferStd(ctx context.Context, tempDir, dstDirPath string, deletePolicy D
TaskData: fs.TaskData{
TaskExtension: task.TaskExtension{
Creator: taskCreator,
ApiUrl: conf.GetApiUrl(ctx),
ApiUrl: common.GetApiUrl(ctx),
},
SrcActualPath: stdpath.Join(tempDir, entry.Name()),
DstActualPath: dstDirActualPath,
@@ -275,7 +276,7 @@ func transferObj(ctx context.Context, tempDir, dstDirPath string, deletePolicy D
TaskData: fs.TaskData{
TaskExtension: task.TaskExtension{
Creator: taskCreator,
ApiUrl: conf.GetApiUrl(ctx),
ApiUrl: common.GetApiUrl(ctx),
},
SrcActualPath: stdpath.Join(srcObjActualPath, obj.GetName()),
DstActualPath: dstDirActualPath,
+8
View File
@@ -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
View File
@@ -233,10 +233,7 @@ func Link(ctx context.Context, storage driver.Driver, path string, args model.Li
if mode == -1 {
mode = storage.(driver.LinkCacheModeResolver).ResolveLinkCacheMode(path)
}
typeKey := "proxy/" + args.Type
if args.Redirect {
typeKey = "redirect/" + args.Type
}
typeKey := args.Type
if mode&driver.LinkCacheIP != 0 {
typeKey += "/" + args.IP
}
-71
View File
@@ -1,71 +0,0 @@
package op
import (
"context"
"io"
"strings"
"testing"
"time"
"github.com/OpenListTeam/OpenList/v4/internal/driver"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/stream"
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
)
type linkModeDriver struct {
driver.Driver
storage model.Storage
calls int
}
func (d *linkModeDriver) Config() driver.Config { return driver.Config{} }
func (d *linkModeDriver) GetStorage() *model.Storage { return &d.storage }
func (d *linkModeDriver) Get(context.Context, string) (model.Obj, error) {
return &model.Object{Name: "file"}, nil
}
func (d *linkModeDriver) Link(_ context.Context, _ model.Obj, args model.LinkArgs) (*model.Link, error) {
d.calls++
expiration := time.Minute
if args.Redirect {
return &model.Link{URL: "https://example.com/file", Expiration: &expiration}, nil
}
return &model.Link{
RangeReader: stream.RangeReaderFunc(func(context.Context, http_range.Range) (io.ReadCloser, error) {
return io.NopCloser(strings.NewReader("file")), nil
}),
Expiration: &expiration,
}, nil
}
func TestLinkCacheSeparatesRedirectAndProxy(t *testing.T) {
for _, tc := range []struct {
name string
firstRedirect bool
}{
{name: "redirect then proxy", firstRedirect: true},
{name: "proxy then redirect", firstRedirect: false},
} {
t.Run(tc.name, func(t *testing.T) {
d := &linkModeDriver{storage: model.Storage{MountPath: "/" + t.Name()}}
for _, redirect := range []bool{tc.firstRedirect, !tc.firstRedirect, tc.firstRedirect, !tc.firstRedirect} {
link, _, err := Link(context.Background(), d, "/file", model.LinkArgs{Redirect: redirect})
if err != nil {
t.Fatal(err)
}
if redirect && (link.URL == "" || link.RangeReader != nil) {
t.Fatalf("redirect link has wrong shape: %+v", link)
}
if !redirect && (link.URL != "" || link.RangeReader == nil) {
t.Fatalf("proxy link has wrong shape: %+v", link)
}
}
if d.calls != 2 {
t.Fatalf("expected one driver call per mode, got %d", d.calls)
}
})
}
}
+2 -2
View File
@@ -4,11 +4,11 @@ import (
"context"
"strings"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/errs"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/op"
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
"github.com/OpenListTeam/OpenList/v4/server/common"
"github.com/pkg/errors"
)
@@ -38,7 +38,7 @@ func link(ctx context.Context, sid, path string, args *LinkArgs) (*model.Sharing
return nil, nil, nil, errors.WithMessage(err, "failed get sharing link")
}
if l.URL != "" && !strings.HasPrefix(l.URL, "http://") && !strings.HasPrefix(l.URL, "https://") {
l.URL = conf.GetApiUrl(ctx) + l.URL
l.URL = common.GetApiUrl(ctx) + l.URL
}
return sharing, l, obj, nil
}
-179
View File
@@ -1,179 +0,0 @@
package stream_test
import (
"bytes"
"context"
"io"
"math/rand"
"sync/atomic"
"testing"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/stream"
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
)
// maxReuseGap mirrors the internal continuation-reuse window (4*utils.MB).
const maxReuseGap = 4 * 1024 * 1024
// newMockSeekableStream builds a SeekableStream whose range reads are served
// from data, counting every upstream range request in gets.
func newMockSeekableStream(t *testing.T, data []byte, gets *atomic.Int64) *stream.SeekableStream {
t.Helper()
rr := stream.RangeReaderFunc(func(ctx context.Context, r http_range.Range) (io.ReadCloser, error) {
gets.Add(1)
if r.Length < 0 || r.Start+r.Length > int64(len(data)) {
r.Length = int64(len(data)) - r.Start
}
return io.NopCloser(io.NewSectionReader(bytes.NewReader(data), r.Start, r.Length)), nil
})
ss, err := stream.NewSeekableStream(&stream.FileStream{
Obj: &model.Object{Size: int64(len(data))},
Ctx: context.Background(),
}, &model.Link{
RangeReader: rr,
ContentLength: int64(len(data)),
})
if err != nil {
t.Fatalf("NewSeekableStream() error = %v", err)
}
return ss
}
// readAtFull reads len(p) bytes at off and fails the test on mismatch.
func readAtFull(t *testing.T, ra io.ReaderAt, data []byte, off int64, p []byte) {
t.Helper()
n, err := ra.ReadAt(p, off)
if err != nil {
t.Fatalf("ReadAt(off=%d) error = %v", off, err)
}
if !bytes.Equal(p, data[off:off+int64(n)]) {
t.Fatalf("ReadAt(off=%d) content mismatch", off)
}
}
func randomData(size int) []byte {
data := make([]byte, size)
x := uint64(42)
for i := range data {
x = x*6364136223846793005 + 1
data[i] = byte(x >> 33)
}
return data
}
// Sequential reads must reuse a single upstream range request.
func TestReadAtSeekerSequentialReuse(t *testing.T) {
data := randomData(16 * 1024 * 1024)
var gets atomic.Int64
ss := newMockSeekableStream(t, data, &gets)
ra, err := stream.NewReadAtSeeker(ss, 0, true)
if err != nil {
t.Fatalf("NewReadAtSeeker() error = %v", err)
}
buf := make([]byte, 128*1024)
for off := 0; off < len(data); off += len(buf) {
readAtFull(t, ra, data, int64(off), buf)
}
if n := gets.Load(); n != 1 {
t.Fatalf("sequential read issued %d range requests, want 1", n)
}
}
// A read landing up to maxReuseGap bytes past a parked reader must be served
// by advancing that reader, without a new range request.
func TestReadAtSeekerSkipsAheadWithinWindow(t *testing.T) {
data := randomData(16 * 1024 * 1024)
var gets atomic.Int64
ss := newMockSeekableStream(t, data, &gets)
ra, err := stream.NewReadAtSeeker(ss, 0, true)
if err != nil {
t.Fatalf("NewReadAtSeeker() error = %v", err)
}
// Park a continuation reader right after reading the first 2 MiB.
chunk := make([]byte, 256*1024)
for off := 0; off < 2*1024*1024; off += len(chunk) {
readAtFull(t, ra, data, int64(off), chunk)
}
skip := 512 * 1024
off := int64(2*1024*1024 + skip)
readAtFull(t, ra, data, off, chunk)
if n := gets.Load(); n != 1 {
t.Fatalf("window skip issued %d range requests, want 1", n)
}
// A second skip deeper inside the window must also be free.
off = int64(4*1024*1024) - 128*1024
readAtFull(t, ra, data, off, chunk)
if n := gets.Load(); n != 1 {
t.Fatalf("second window skip issued %d range requests, want 1", n)
}
}
// A forward jump beyond the reuse window must open a new range request but
// keep the parked reader available for later window hits.
func TestReadAtSeekerFarJumpOpensNewRequest(t *testing.T) {
data := randomData(16 * 1024 * 1024)
var gets atomic.Int64
ss := newMockSeekableStream(t, data, &gets)
ra, err := stream.NewReadAtSeeker(ss, 0, true)
if err != nil {
t.Fatalf("NewReadAtSeeker() error = %v", err)
}
buf := make([]byte, 256*1024)
for off := 0; off < 2*1024*1024; off += len(buf) {
readAtFull(t, ra, data, int64(off), buf)
}
// 2 MiB -> 10 MiB is beyond the 4 MiB reuse window.
off := int64(10 * 1024 * 1024)
readAtFull(t, ra, data, off, buf)
if n := gets.Load(); n != 2 {
t.Fatalf("far jump issued %d range requests, want 2", n)
}
// Back within the window of the 10 MiB chain: free reuse again.
readAtFull(t, ra, data, off+maxReuseGap, buf)
if n := gets.Load(); n != 2 {
t.Fatalf("jump inside new window issued %d range requests, want 2", n)
}
}
// Backward reads can never reuse a parked continuation and must open a new
// range request.
func TestReadAtSeekerBackwardJumpOpensNewRequest(t *testing.T) {
data := randomData(8 * 1024 * 1024)
var gets atomic.Int64
ss := newMockSeekableStream(t, data, &gets)
ra, err := stream.NewReadAtSeeker(ss, 0, true)
if err != nil {
t.Fatalf("NewReadAtSeeker() error = %v", err)
}
buf := make([]byte, 256*1024)
for off := 0; off < 2*1024*1024; off += len(buf) {
readAtFull(t, ra, data, int64(off), buf)
}
readAtFull(t, ra, data, int64(1024*1024), buf)
if n := gets.Load(); n != 2 {
t.Fatalf("backward jump issued %d range requests, want 2", n)
}
}
// Random reads must return correct data and keep upstream requests bounded:
// each read is either a window hit or a fresh request, never more than one.
func TestReadAtSeekerRandomReads(t *testing.T) {
data := randomData(32 * 1024 * 1024)
var gets atomic.Int64
ss := newMockSeekableStream(t, data, &gets)
ra, err := stream.NewReadAtSeeker(ss, 0, true)
if err != nil {
t.Fatalf("NewReadAtSeeker() error = %v", err)
}
const chunk = 8 * 1024
buf := make([]byte, chunk)
r := rand.New(rand.NewSource(7))
for i := 0; i < 200; i++ {
off := r.Int63n(int64(len(data)) - chunk)
readAtFull(t, ra, data, off, buf)
}
if n := gets.Load(); n > 200 {
t.Fatalf("random reads issued %d range requests, want <= 200", n)
}
}
+34 -71
View File
@@ -8,7 +8,6 @@ import (
"io"
"math"
"os"
"sort"
"sync"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
@@ -359,72 +358,10 @@ func (r *ReaderUpdatingProgress) Close() error {
type RangeReadReadAtSeeker struct {
ss *SeekableStream
masterOff int64
readers orderedReaders
readerMap sync.Map
headCache *headCache
}
type orderedReaders struct {
mu sync.Mutex
m map[int64]io.Reader
keys []int64
}
func (o *orderedReaders) store(off int64, r io.Reader) {
o.mu.Lock()
defer o.mu.Unlock()
if _, ok := o.m[off]; ok {
o.m[off] = r
return
}
if o.m == nil {
o.m = make(map[int64]io.Reader)
}
i := sort.Search(len(o.keys), func(i int) bool { return o.keys[i] >= off })
o.keys = append(o.keys, 0)
copy(o.keys[i+1:], o.keys[i:])
o.keys[i] = off
o.m[off] = r
}
func (o *orderedReaders) takeExact(off int64) (io.Reader, bool) {
o.mu.Lock()
defer o.mu.Unlock()
r, ok := o.m[off]
if ok {
delete(o.m, off)
o.removeKey(off)
}
return r, ok
}
func (o *orderedReaders) takeBest(off int64) (io.Reader, int64, bool) {
o.mu.Lock()
defer o.mu.Unlock()
if r, ok := o.m[off]; ok {
delete(o.m, off)
o.removeKey(off)
return r, off, true
}
i := sort.Search(len(o.keys), func(i int) bool { return o.keys[i] >= off })
if i == 0 {
return nil, 0, false
}
k := o.keys[i-1]
if off-k > 4*utils.MB {
return nil, 0, false
}
r := o.m[k]
delete(o.m, k)
o.removeKey(k)
return r, k, true
}
func (o *orderedReaders) removeKey(k int64) {
i := sort.Search(len(o.keys), func(i int) bool { return o.keys[i] >= k })
copy(o.keys[i:], o.keys[i+1:])
o.keys = o.keys[:len(o.keys)-1]
}
type headCache struct {
reader io.Reader
bufs [][]byte
@@ -459,7 +396,7 @@ func (r *headCache) Close() error {
func (r *RangeReadReadAtSeeker) InitHeadCache() {
if r.masterOff == 0 {
value, _ := r.readers.takeExact(0)
value, _ := r.readerMap.LoadAndDelete(int64(0))
r.headCache = &headCache{reader: value.(io.Reader)}
r.ss.Closers.Add(r.headCache)
}
@@ -485,9 +422,9 @@ func NewReadAtSeeker(ss *SeekableStream, offset int64, forceRange ...bool) (mode
if err != nil {
return nil, err
}
r.readers.store(offset, reader)
r.readerMap.Store(int64(offset), reader)
} else {
r.readers.store(0, ss)
r.readerMap.Store(int64(offset), ss)
}
return r, nil
}
@@ -505,15 +442,41 @@ func NewMultiReaderAt(ss []*SeekableStream) (readerutil.SizeReaderAt, error) {
}
func (r *RangeReadReadAtSeeker) getReaderAtOffset(off int64) (io.Reader, error) {
if rr, cur, ok := r.readers.takeBest(off); ok {
if cur == off {
for {
var cur int64 = -1
r.readerMap.Range(func(key, value any) bool {
k := key.(int64)
if off == k {
cur = k
return false
}
if off > k && off-k <= 4*utils.MB && k > cur {
cur = k
}
return true
})
if cur < 0 {
break
}
v, ok := r.readerMap.LoadAndDelete(int64(cur))
if !ok {
continue
}
rr := v.(io.Reader)
if off == int64(cur) {
// logrus.Debugf("getReaderAtOffset match_%d", off)
return rr, nil
}
n, _ := utils.CopyWithBufferN(io.Discard, rr, off-cur)
if cur+n == off {
cur += n
if cur == off {
// logrus.Debugf("getReaderAtOffset old_%d", off)
return rr, nil
}
break
}
// logrus.Debugf("getReaderAtOffset new_%d", off)
reader, err := r.ss.RangeRead(http_range.Range{Start: off, Length: -1})
if err != nil {
return nil, err
@@ -538,7 +501,7 @@ func (r *RangeReadReadAtSeeker) ReadAt(p []byte, off int64) (n int, err error) {
off += int64(n)
switch err {
case nil:
r.readers.store(off, rr)
r.readerMap.Store(int64(off), rr)
case io.ErrUnexpectedEOF:
err = io.EOF
}
-20
View File
@@ -1,20 +0,0 @@
package task_test
import (
"context"
"testing"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/task"
)
func TestTaskExtensionRestoresAPIURL(t *testing.T) {
const want = "https://openlist.example"
extension := task.TaskExtension{ApiUrl: want}
extension.SetCtx(context.Background())
if got := conf.GetApiUrl(extension.Ctx()); got != want {
t.Fatalf("restored origin = %q, want %q", got, want)
}
}
+32 -13
View File
@@ -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
View File
@@ -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
View File
@@ -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
+156
View File
@@ -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")
}
}
+160
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+2 -1
View File
@@ -31,5 +31,6 @@ func GetApiUrlFromRequest(r *http.Request) string {
}
func GetApiUrl(ctx context.Context) string {
return conf.GetApiUrl(ctx)
api, _ := ctx.Value(conf.ApiUrlKey).(string)
return api
}
+6 -2
View File
@@ -34,7 +34,9 @@ func Proxy(w http.ResponseWriter, r *http.Request, link *model.Link, file model.
if link.RangeReader == nil {
r = r.WithContext(context.WithValue(r.Context(), conf.RequestHeaderKey, r.Header))
}
return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, rrf)
return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, &model.RangeReadCloser{
RangeReader: rrf,
})
}
if link.RangeReader != nil {
@@ -43,7 +45,9 @@ func Proxy(w http.ResponseWriter, r *http.Request, link *model.Link, file model.
if size <= 0 {
size = file.GetSize()
}
return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, link.RangeReader)
return net.ServeHTTP(w, r, file.GetName(), file.ModTime(), size, &model.RangeReadCloser{
RangeReader: link.RangeReader,
})
}
//transparent proxy
-49
View File
@@ -1,49 +0,0 @@
package common
import (
"bytes"
"context"
"io"
"net/http"
"net/http/httptest"
"testing"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/stream"
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
)
func TestProxyCancelledPartitionedReaderDoesNotPanic(t *testing.T) {
oldConf := conf.Conf
conf.Conf = conf.DefaultConfig("data")
t.Cleanup(func() { conf.Conf = oldConf })
link := &model.Link{
Concurrency: 2,
PartSize: 4,
RangeReader: stream.RangeReaderFunc(func(ctx context.Context, requested http_range.Range) (io.ReadCloser, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
return io.NopCloser(bytes.NewReader([]byte("0123456789abcdef")[requested.Start : requested.Start+requested.Length])), nil
}),
}
file := &model.Object{Name: "fixture.bin", Size: 16}
for range 32 {
func() {
defer func() {
if recovered := recover(); recovered != nil {
t.Errorf("Proxy panicked on cancelled partitioned read: %v", recovered)
}
}()
r := httptest.NewRequest(http.MethodGet, "/proxy/fixture.bin", nil)
ctx, cancel := context.WithCancel(r.Context())
cancel()
w := httptest.NewRecorder()
_ = Proxy(w, r.WithContext(ctx), link, file)
if bytes.Contains(w.Body.Bytes(), []byte("0123456789abcdef")) {
t.Errorf("cancelled response contained file contents: %q", w.Body.String())
}
}()
}
}
+2 -3
View File
@@ -94,11 +94,10 @@ func (f *FileUploadProxy) Close() error {
return err
}
arr := make([]byte, 512)
n, err := f.buffer.Read(arr)
if err != nil && err != io.EOF {
if _, err := f.buffer.Read(arr); err != nil {
return err
}
contentType := http.DetectContentType(arr[:n])
contentType := http.DetectContentType(arr)
if _, err := f.buffer.Seek(0, io.SeekStart); err != nil {
return err
}
+102 -17
View File
@@ -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"`
}
-126
View File
@@ -1,126 +0,0 @@
package handles
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
_ "github.com/OpenListTeam/OpenList/v4/drivers/local"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/db"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/op"
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
func setupBackslashTraversalTest(t *testing.T, root string, permission int32) *model.User {
t.Helper()
database, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{})
if err != nil {
t.Fatal(err)
}
conf.Conf = conf.DefaultConfig(t.TempDir())
db.Init(database)
addition, err := utils.Json.MarshalToString(map[string]string{"root_folder_path": root})
if err != nil {
t.Fatal(err)
}
if _, err = op.CreateStorage(context.Background(), model.Storage{
Driver: "Local", MountPath: "/", Addition: addition,
}); err != nil {
t.Fatal(err)
}
return &model.User{
Username: "restricted-user", BasePath: "/team/a", Role: model.GENERAL,
Permission: permission,
}
}
func prepareBackslashTraversalFs(t *testing.T) (root string, secretPath string) {
t.Helper()
root = t.TempDir()
if err := os.MkdirAll(filepath.Join(root, "team", "a", "writable"), 0o700); err != nil {
t.Fatal(err)
}
if err := os.MkdirAll(filepath.Join(root, "team", "ab"), 0o700); err != nil {
t.Fatal(err)
}
secretPath = filepath.Join(root, "team", "ab", "secret.txt")
if err := os.WriteFile(secretPath, []byte("synthetic-secret"), 0o600); err != nil {
t.Fatal(err)
}
return root, secretPath
}
func invokeHandler(t *testing.T, user *model.User, payload any, handler gin.HandlerFunc) *httptest.ResponseRecorder {
t.Helper()
body, err := json.Marshal(payload)
if err != nil {
t.Fatal(err)
}
recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
req := httptest.NewRequest(http.MethodPost, "/api/fs/remove", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req = req.WithContext(context.WithValue(req.Context(), conf.UserKey, user))
ctx.Request = req
handler(ctx)
return recorder
}
func TestFsRemoveRejectsBackslashTraversal(t *testing.T) {
gin.SetMode(gin.TestMode)
root, secretPath := prepareBackslashTraversalFs(t)
user := setupBackslashTraversalTest(t, root, 1<<3|1<<7)
for _, name := range []string{"../../ab/secret.txt", `..\..\ab\secret.txt`} {
recorder := invokeHandler(t, user, map[string]any{"dir": "/writable", "names": []string{name}}, FsRemove)
if recorder.Code != http.StatusOK || !strings.Contains(recorder.Body.String(), `"code":403`) {
t.Fatalf("payload %q: got status=%d body=%s, want 403", name, recorder.Code, recorder.Body.String())
}
if _, err := os.Stat(secretPath); err != nil {
t.Fatalf("payload %q deleted sibling file: %v", name, err)
}
}
}
func TestFsMoveRejectsBackslashTraversal(t *testing.T) {
gin.SetMode(gin.TestMode)
root, secretPath := prepareBackslashTraversalFs(t)
user := setupBackslashTraversalTest(t, root, 1<<3|1<<5)
recorder := invokeHandler(t, user, map[string]any{
"src_dir": "/writable", "dst_dir": "/", "names": []string{`..\..\ab\secret.txt`},
}, FsMove)
if recorder.Code != http.StatusOK || !strings.Contains(recorder.Body.String(), `"code":403`) {
t.Fatalf("got status=%d body=%s, want 403", recorder.Code, recorder.Body.String())
}
if _, err := os.Stat(secretPath); err != nil {
t.Fatalf("backslash traversal moved sibling file: %v", err)
}
}
func TestFsCopyRejectsBackslashTraversal(t *testing.T) {
gin.SetMode(gin.TestMode)
root, secretPath := prepareBackslashTraversalFs(t)
user := setupBackslashTraversalTest(t, root, 1<<3|1<<6)
recorder := invokeHandler(t, user, map[string]any{
"src_dir": "/writable", "dst_dir": "/", "names": []string{`..\..\ab\secret.txt`},
}, FsCopy)
if recorder.Code != http.StatusOK || !strings.Contains(recorder.Body.String(), `"code":403`) {
t.Fatalf("got status=%d body=%s, want 403", recorder.Code, recorder.Body.String())
}
if _, err := os.Stat(secretPath); err != nil {
t.Fatalf("backslash traversal affected sibling file: %v", err)
}
}
+5 -7
View File
@@ -347,13 +347,11 @@ func FsGet(c *gin.Context, req *FsGetReq, user *model.User) {
}
}
}
parentPath := stdpath.Dir(reqPath)
var related []model.Obj
if !obj.IsDir() && utils.GetFileType(obj.GetName()) == conf.VIDEO {
sameLevelFiles, err := fs.List(c.Request.Context(), parentPath, &fs.ListArgs{})
if err == nil {
related = filterRelated(sameLevelFiles, obj)
}
parentPath := stdpath.Dir(reqPath)
sameLevelFiles, err := fs.List(c.Request.Context(), parentPath, &fs.ListArgs{})
if err == nil {
related = filterRelated(sameLevelFiles, obj)
}
parentMeta, _ := op.GetNearestMeta(parentPath)
thumb, _ := model.GetThumb(obj)
@@ -368,7 +366,7 @@ func FsGet(c *gin.Context, req *FsGetReq, user *model.User) {
HashInfoStr: obj.GetHash().String(),
HashInfo: obj.GetHash().Export(),
Sign: common.Sign(obj, parentPath, isEncrypt(meta, reqPath)),
Type: utils.GetObjType(obj.GetName(), obj.IsDir()),
Type: utils.GetFileType(obj.GetName()),
Thumb: thumb,
MountDetails: mountDetails,
},
+146 -2
View File
@@ -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
}
+8 -17
View File
@@ -44,7 +44,14 @@ func Search(c *gin.Context) {
return
}
nodes, total, err := search.SearchFiltered(c, req.SearchReq, func(node model.SearchNode) bool {
return isSearchNodeAccessible(user, node, req.Password, op.GetNearestMeta)
if !utils.IsSubPath(user.BasePath, node.Parent) {
return false
}
meta, err := op.GetNearestMeta(node.Parent)
if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) {
return false
}
return common.CanAccess(user, meta, path.Join(node.Parent, node.Name), req.Password)
})
if err != nil {
common.ErrorResp(c, err, 500)
@@ -56,22 +63,6 @@ func Search(c *gin.Context) {
})
}
func isSearchNodeAccessible(user *model.User, node model.SearchNode, password string, resolveMeta func(string) (*model.Meta, error)) bool {
if !utils.IsSubPath(user.BasePath, node.Parent) {
return false
}
nodePath := path.Join(node.Parent, node.Name)
metaPath := node.Parent
if node.IsDir {
metaPath = nodePath
}
meta, err := resolveMeta(metaPath)
if err != nil && !errors.Is(errors.Cause(err), errs.MetaNotFound) {
return false
}
return common.CanAccess(user, meta, nodePath, password)
}
func nodeToSearchResp(node model.SearchNode) SearchResp {
return SearchResp{
SearchNode: node,
-78
View File
@@ -1,78 +0,0 @@
package handles
import (
"path"
"testing"
"github.com/OpenListTeam/OpenList/v4/internal/errs"
"github.com/OpenListTeam/OpenList/v4/internal/model"
)
func fakeResolveMeta(metas map[string]*model.Meta) func(string) (*model.Meta, error) {
return func(p string) (*model.Meta, error) {
for {
if meta, ok := metas[p]; ok {
return meta, nil
}
if p == "/" {
return nil, errs.MetaNotFound
}
p = path.Dir(p)
}
}
}
func TestIsSearchNodeAccessible(t *testing.T) {
tests := []struct {
name string
metas map[string]*model.Meta
node model.SearchNode
want bool
wantMetaPath string
}{
{
name: "restricted directory",
metas: map[string]*model.Meta{"/private": {Path: "/private", ReadUsers: []uint{1}}},
node: model.SearchNode{Parent: "/", Name: "private", IsDir: true},
want: false,
wantMetaPath: "/private",
},
{
name: "restricted sub directory",
metas: map[string]*model.Meta{"/private": {Path: "/private", ReadUsers: []uint{1}, ReadUsersSub: true}},
node: model.SearchNode{Parent: "/private", Name: "sub", IsDir: true},
want: false,
wantMetaPath: "/private/sub",
},
{
name: "file keeps parent scope",
metas: map[string]*model.Meta{"/private": {Path: "/private", ReadUsers: []uint{1}}},
node: model.SearchNode{Parent: "/private", Name: "a.txt", IsDir: false},
want: true,
wantMetaPath: "/private",
},
{
name: "outside base path",
node: model.SearchNode{Parent: "/other", Name: "private", IsDir: true},
want: false,
wantMetaPath: "",
},
}
user := &model.User{ID: 2, BasePath: "/"}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
resolve := fakeResolveMeta(tt.metas)
var gotMetaPath string
spy := func(p string) (*model.Meta, error) {
gotMetaPath = p
return resolve(p)
}
if got := isSearchNodeAccessible(user, tt.node, "", spy); got != tt.want {
t.Fatalf("isSearchNodeAccessible() = %v, want %v", got, tt.want)
}
if tt.wantMetaPath != "" && gotMetaPath != tt.wantMetaPath {
t.Fatalf("meta resolved at %q, want %q", gotMetaPath, tt.wantMetaPath)
}
})
}
}
+1 -1
View File
@@ -54,7 +54,7 @@ func SharingGet(c *gin.Context, req *FsGetReq) {
HashInfoStr: obj.GetHash().String(),
HashInfo: obj.GetHash().Export(),
Sign: "",
Type: utils.GetObjType(obj.GetName(), obj.IsDir()),
Type: utils.GetFileType(obj.GetName()),
Thumb: thumb,
},
RawURL: url,
+36 -51
View File
@@ -122,53 +122,6 @@ func generateSSOBindingToken(c *gin.Context, purpose, ssoID string) (string, err
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(common.SecretKey)
}
// ssoTargetOrigin returns the origin that is allowed to receive the SSO result
// via postMessage. It honours the operator-configured sso_postmessage_origin so
// a frontend served from a different origin than the API can still receive the
// result; otherwise it falls back to the API origin, or "/" to restrict
// delivery to same-origin openers when that cannot be resolved.
func ssoTargetOrigin(c *gin.Context) string {
if configured := setting.GetStr(conf.SSOPostMessageOrigin); configured != "" {
if u, err := url.Parse(configured); err == nil &&
(u.Scheme == "http" || u.Scheme == "https") &&
u.Host != "" && u.User == nil &&
(u.Path == "" || u.Path == "/") &&
u.RawQuery == "" && u.Fragment == "" {
return u.Scheme + "://" + u.Host
}
}
u, err := url.Parse(common.GetApiUrl(c))
if err != nil || u.Scheme == "" || u.Host == "" {
return "/"
}
return u.Scheme + "://" + u.Host
}
// ssoPostMessage hands the SSO result back to the window that started the login.
// The target origin is pinned so that an arbitrary page cannot open the SSO
// endpoint in a popup and read the payload out of the message event.
func ssoPostMessage(c *gin.Context, payload map[string]string) {
data, err := utils.Json.MarshalToString(payload)
if err != nil {
common.ErrorResp(c, err, 500)
return
}
origin, err := utils.Json.MarshalToString(ssoTargetOrigin(c))
if err != nil {
common.ErrorResp(c, err, 500)
return
}
html := fmt.Sprintf(`<!DOCTYPE html>
<head></head>
<body>
<script>
if (window.opener) { window.opener.postMessage(%s, %s) }
window.close()
</script>
</body>`, data, origin)
c.Data(200, "text/html; charset=utf-8", []byte(html))
}
func ssoRedirectUri(c *gin.Context, useCompatibility bool, method string) string {
if useCompatibility {
return common.GetApiUrl(c) + "/api/auth/" + method
@@ -385,7 +338,15 @@ func OIDCLoginCallback(c *gin.Context) {
c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof)
return
}
ssoPostMessage(c, map[string]string{"sso_id": bindingProof})
html := fmt.Sprintf(`<!DOCTYPE html>
<head></head>
<body>
<script>
window.opener.postMessage({"sso_id": "%s"}, "*")
window.close()
</script>
</body>`, bindingProof)
c.Data(200, "text/html; charset=utf-8", []byte(html))
return
}
if method == "sso_get_token" {
@@ -406,7 +367,15 @@ func OIDCLoginCallback(c *gin.Context) {
c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token)
return
}
ssoPostMessage(c, map[string]string{"token": token})
html := fmt.Sprintf(`<!DOCTYPE html>
<head></head>
<body>
<script>
window.opener.postMessage({"token":"%s"}, "*")
window.close()
</script>
</body>`, token)
c.Data(200, "text/html; charset=utf-8", []byte(html))
return
}
}
@@ -547,7 +516,15 @@ func SSOLoginCallback(c *gin.Context) {
c.Redirect(302, common.GetApiUrl(c)+"/@manage?sso_id="+bindingProof)
return
}
ssoPostMessage(c, map[string]string{"sso_id": bindingProof})
html := fmt.Sprintf(`<!DOCTYPE html>
<head></head>
<body>
<script>
window.opener.postMessage({"sso_id": "%s"}, "*")
window.close()
</script>
</body>`, bindingProof)
c.Data(200, "text/html; charset=utf-8", []byte(html))
return
}
username := utils.Json.Get(resp.Body(), usernameField).ToString()
@@ -568,5 +545,13 @@ func SSOLoginCallback(c *gin.Context) {
c.Redirect(302, common.GetApiUrl(c)+"/@login?token="+token)
return
}
ssoPostMessage(c, map[string]string{"token": token})
html := fmt.Sprintf(`<!DOCTYPE html>
<head></head>
<body>
<script>
window.opener.postMessage({"token":"%s"}, "*")
window.close()
</script>
</body>`, token)
c.Data(200, "text/html; charset=utf-8", []byte(html))
}
-102
View File
@@ -1,102 +0,0 @@
package handles
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/op"
"github.com/gin-gonic/gin"
)
func ssoTestContext(apiUrl string) (*gin.Context, *httptest.ResponseRecorder) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, engine := gin.CreateTestContext(rec)
// Matches server.Init, which is what lets GetApiUrl reach the value the
// middleware stored on the request context.
engine.ContextWithFallback = true
req := httptest.NewRequest(http.MethodGet, "/api/auth/sso?method=sso_get_token", nil)
if apiUrl != "" {
req = req.WithContext(context.WithValue(req.Context(), conf.ApiUrlKey, apiUrl))
}
c.Request = req
// Keep setting lookups off the (uninitialised) database: ssoTargetOrigin
// reads sso_postmessage_origin through the setting cache.
op.Cache.SetSetting(conf.SSOPostMessageOrigin, &model.SettingItem{
Key: conf.SSOPostMessageOrigin,
Value: "",
})
return c, rec
}
// A page that opens the SSO endpoint in a popup must not be able to read the
// token: the postMessage target origin has to name the site, never "*".
func TestSSOPostMessagePinsTargetOrigin(t *testing.T) {
c, rec := ssoTestContext("https://openlist.example.com/base")
ssoPostMessage(c, map[string]string{"token": "secret-token"})
body := rec.Body.String()
if strings.Contains(body, `"*"`) || strings.Contains(body, `, '*'`) {
t.Fatalf("wildcard target origin present in response:\n%s", body)
}
if !strings.Contains(body, `"https://openlist.example.com"`) {
t.Errorf("expected the site origin as target, got:\n%s", body)
}
if !strings.Contains(body, "secret-token") {
t.Errorf("payload should still reach a legitimate opener, got:\n%s", body)
}
}
// A frontend served from a different origin than the API needs the operator to
// be able to point the target at the frontend origin. The configured origin
// must win over the API origin.
func TestSSOPostMessageUsesConfiguredOrigin(t *testing.T) {
c, rec := ssoTestContext("https://api.example.com/base")
op.Cache.SetSetting(conf.SSOPostMessageOrigin, &model.SettingItem{
Key: conf.SSOPostMessageOrigin,
Value: "https://frontend.example.com",
})
defer op.Cache.ClearAll()
ssoPostMessage(c, map[string]string{"token": "secret-token"})
body := rec.Body.String()
if !strings.Contains(body, `"https://frontend.example.com"`) {
t.Errorf("expected the configured origin as target, got:\n%s", body)
}
}
// If the site URL cannot be resolved the fallback must tighten delivery to
// same-origin openers, not widen it back to every origin.
func TestSSOPostMessageFallsBackToSameOrigin(t *testing.T) {
c, rec := ssoTestContext("")
ssoPostMessage(c, map[string]string{"token": "secret-token"})
body := rec.Body.String()
if strings.Contains(body, `"*"`) {
t.Fatalf("fallback must not be a wildcard origin:\n%s", body)
}
if !strings.Contains(body, `"/"`) {
t.Errorf(`expected "/" fallback origin, got:\n%s`, body)
}
}
// userID comes from the identity provider, so it must be encoded rather than
// interpolated into the JS string literal it used to land in.
func TestSSOPostMessageEscapesProviderControlledValue(t *testing.T) {
c, rec := ssoTestContext("https://openlist.example.com")
ssoPostMessage(c, map[string]string{"sso_id": `"});alert(document.domain);//`})
body := rec.Body.String()
if strings.Contains(body, `alert(document.domain)`) && !strings.Contains(body, `\"`) {
t.Fatalf("provider value was not escaped:\n%s", body)
}
if !strings.Contains(body, `\"});alert`) {
t.Errorf("expected the injected quote to be escaped, got:\n%s", body)
}
}
+1512 -12
View File
File diff suppressed because it is too large Load Diff
+259
View File
@@ -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)
}
}
+4 -6
View File
@@ -68,11 +68,9 @@ func (s *Server) callFSGet(c *gin.Context, raw json.RawMessage) (any, *rpcError)
parentPath := stdpath.Dir(reqPath)
var related []model.Obj
if !obj.IsDir() && utils.GetFileType(obj.GetName()) == conf.VIDEO {
sameLevelFiles, err := fs.List(ctx, parentPath, &fs.ListArgs{})
if err == nil {
related = filterRelatedObjs(sameLevelFiles, obj)
}
sameLevelFiles, err := fs.List(ctx, parentPath, &fs.ListArgs{})
if err == nil {
related = filterRelatedObjs(sameLevelFiles, obj)
}
parentMeta, _ := op.GetNearestMeta(parentPath)
@@ -87,7 +85,7 @@ func (s *Server) callFSGet(c *gin.Context, raw json.RawMessage) (any, *rpcError)
Created: obj.CreateTime(),
Sign: common.Sign(obj, parentPath, isEncrypt(meta, reqPath)),
Thumb: thumb,
Type: utils.GetObjType(obj.GetName(), obj.IsDir()),
Type: utils.GetFileType(obj.GetName()),
HashInfoStr: obj.GetHash().String(),
HashInfo: obj.GetHash().Export(),
MountDetails: mountDetails,
-62
View File
@@ -1,62 +0,0 @@
package middlewares
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/gin-gonic/gin"
)
func TestStoragesLoadedAdmitsRequestOrigin(t *testing.T) {
originalMode := gin.Mode()
gin.SetMode(gin.TestMode)
originalConf := conf.Conf
originalLoaded := conf.StoragesLoaded
t.Cleanup(func() {
gin.SetMode(originalMode)
conf.Conf = originalConf
conf.StoragesLoaded = originalLoaded
})
conf.StoragesLoaded = true
router := gin.New()
router.Use(StoragesLoaded)
router.GET("/", func(c *gin.Context) {
c.String(http.StatusOK, conf.GetApiUrl(c.Request.Context()))
})
assertOrigin := func(name, siteURL, target string, header http.Header, want string) {
t.Run(name, func(t *testing.T) {
conf.Conf = &conf.Config{SiteURL: siteURL}
req := httptest.NewRequest(http.MethodGet, target, nil)
req.Header = header
rec := httptest.NewRecorder()
router.ServeHTTP(rec, req)
if got := rec.Body.String(); got != want {
t.Fatalf("origin = %q, want %q", got, want)
}
})
}
assertOrigin(
"configured site URL",
"https://openlist.example/base/",
"http://ignored.example/",
nil,
"https://openlist.example/base",
)
assertOrigin(
"forwarded request",
"",
"http://internal.example/",
http.Header{
"X-Forwarded-Proto": {"https"},
"X-Forwarded-Host": {"public.example"},
},
"https://public.example",
)
}
+13
View File
@@ -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)
}
+2 -5
View File
@@ -152,8 +152,6 @@ func (b *s3Backend) HeadObject(ctx context.Context, bucketName, objectName strin
// GetObject fetchs the object from the filesystem.
func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string, rangeRequest *gofakes3.ObjectRangeRequest) (s3Obj *gofakes3.Object, err error) {
defer func() { err = mapBackendError(err) }()
bucket, err := getBucketByName(bucketName)
if err != nil {
return nil, err
@@ -195,7 +193,7 @@ func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string
return nil, fmt.Errorf("the remote storage driver need to be enhanced to support s3")
}
var rd io.ReadCloser
var rd io.Reader
if rnge != nil {
rd, err = rrf.RangeRead(ctx, http_range.Range(*rnge))
} else {
@@ -217,7 +215,6 @@ func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string
meta[k] = v
}
}
closers := utils.NewClosers(rd, link)
return &gofakes3.Object{
// Name: gofakes3.URLEncode(objectName),
@@ -226,7 +223,7 @@ func (b *s3Backend) GetObject(ctx context.Context, bucketName, objectName string
Metadata: meta,
Size: size,
Range: rnge,
Contents: utils.ReadCloser{Reader: rd, Closer: &closers},
Contents: utils.ReadCloser{Reader: rd, Closer: link},
}, nil
}
-66
View File
@@ -1,66 +0,0 @@
package s3
import (
"context"
"encoding/xml"
"errors"
"net/http"
"net/http/httptest"
"testing"
"github.com/OpenListTeam/OpenList/v4/internal/errs"
"github.com/OpenListTeam/gofakes3"
"github.com/OpenListTeam/gofakes3/s3mem"
)
func TestMapBackendErrorMapsOnlyTemporaryCapacity(t *testing.T) {
capacity := errs.NewErr(errs.TemporaryCapacity, "callback admission timed out")
if got := mapBackendError(capacity); got != gofakes3.ErrSlowDown {
t.Fatalf("capacity error mapped to %v, want %v", got, gofakes3.ErrSlowDown)
}
permanent := errors.New("permission denied")
if got := mapBackendError(permanent); got != permanent {
t.Fatalf("permanent error mapped to %v, want original error", got)
}
if got := mapBackendError(nil); got != nil {
t.Fatalf("nil error mapped to %v", got)
}
}
type capacityBackend struct {
gofakes3.Backend
}
func (b capacityBackend) GetObject(context.Context, string, string, *gofakes3.ObjectRangeRequest) (*gofakes3.Object, error) {
return nil, mapBackendError(errs.NewErr(errs.TemporaryCapacity, "callback admission timed out"))
}
func TestTemporaryCapacityProducesS3SlowDownResponse(t *testing.T) {
memory := s3mem.New()
if err := memory.CreateBucket(t.Context(), "bucket"); err != nil {
t.Fatal(err)
}
server := httptest.NewServer(gofakes3.New(capacityBackend{Backend: memory}).Server())
defer server.Close()
request, err := http.NewRequestWithContext(t.Context(), http.MethodGet, server.URL+"/bucket/object", nil)
if err != nil {
t.Fatal(err)
}
response, err := server.Client().Do(request)
if err != nil {
t.Fatal(err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusServiceUnavailable {
t.Fatalf("status = %d, want %d", response.StatusCode, http.StatusServiceUnavailable)
}
var result gofakes3.ErrorResult
if err := xml.NewDecoder(response.Body).Decode(&result); err != nil {
t.Fatal(err)
}
if result.Code != gofakes3.ErrSlowDown || result.Message != gofakes3.ErrSlowDown.Message() {
t.Fatalf("S3 error = %#v, want SlowDown with standard message", result)
}
}
-132
View File
@@ -1,132 +0,0 @@
package s3
import (
"context"
"encoding/json"
"errors"
"io"
"os"
"path/filepath"
"slices"
"strings"
"testing"
"github.com/OpenListTeam/OpenList/v4/drivers/local"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
"github.com/OpenListTeam/OpenList/v4/internal/db"
"github.com/OpenListTeam/OpenList/v4/internal/driver"
"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/op"
"github.com/OpenListTeam/OpenList/v4/internal/stream"
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
"github.com/OpenListTeam/OpenList/v4/pkg/utils"
"gorm.io/gorm"
)
const closeTrackingDriverName = "S3CloseTrackingLocal"
type closeTrackingDriver struct {
local.Local
closed *[]string
}
func (d *closeTrackingDriver) Config() driver.Config {
c := d.Local.Config()
c.Name = closeTrackingDriverName
return c
}
func (d *closeTrackingDriver) Link(context.Context, model.Obj, model.LinkArgs) (*model.Link, error) {
link := &model.Link{
ContentLength: 4,
RangeReader: stream.RangeReaderFunc(func(context.Context, http_range.Range) (io.ReadCloser, error) {
return utils.NewReadCloser(strings.NewReader("body"), func() error {
*d.closed = append(*d.closed, "body")
return nil
}), nil
}),
RequireReference: true,
}
link.SyncClosers.Add(utils.CloseFunc(func() error {
*d.closed = append(*d.closed, "link")
return nil
}))
return link, nil
}
func TestGetObjectClosesRangeBodyBeforeLink(t *testing.T) {
ctx := context.Background()
var closed []string
op.RegisterDriver(func() driver.Driver {
return &closeTrackingDriver{closed: &closed}
})
root := t.TempDir()
if err := os.WriteFile(filepath.Join(root, "fixture.txt"), []byte("body"), 0o600); err != nil {
t.Fatal(err)
}
addition, err := json.Marshal(struct {
RootFolderPath string `json:"root_folder_path"`
}{RootFolderPath: root})
if err != nil {
t.Fatal(err)
}
mount := "/" + sanitizeTestName(t.Name())
storageID, err := op.CreateStorage(ctx, model.Storage{
Driver: closeTrackingDriverName,
MountPath: mount,
Addition: string(addition),
})
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
if err := op.DeleteStorageById(ctx, storageID); err != nil {
t.Errorf("delete fixture storage: %v", err)
}
})
previousBuckets, previousBucketsErr := op.GetSettingItemByKey(conf.S3Buckets)
if previousBucketsErr != nil && !errors.Is(previousBucketsErr, gorm.ErrRecordNotFound) {
t.Fatal(previousBucketsErr)
}
if err := op.SaveSettingItem(&model.SettingItem{
Key: conf.S3Buckets,
Value: `[{"name":"close","path":"` + mount + `"}]`,
}); err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
if previousBucketsErr == nil {
if err := op.SaveSettingItem(previousBuckets); err != nil {
t.Errorf("restore S3 buckets: %v", err)
}
return
}
if err := db.DeleteSettingItemByKey(conf.S3Buckets); err != nil {
t.Errorf("delete fixture S3 buckets: %v", err)
}
op.SettingCacheUpdate()
})
object, err := newBackend().(*s3Backend).GetObject(ctx, "close", "fixture.txt", nil)
if err != nil {
t.Fatal(err)
}
contents, err := io.ReadAll(object.Contents)
if err != nil {
t.Fatal(err)
}
if string(contents) != "body" {
t.Fatalf("contents = %q, want body", contents)
}
if err := object.Contents.Close(); err != nil {
t.Fatal(err)
}
if err := object.Contents.Close(); err != nil {
t.Fatal(err)
}
if !slices.Equal(closed, []string{"body", "link"}) {
t.Fatalf("close order = %v, want [body link] exactly once", closed)
}
}
+2 -11
View File
@@ -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
}
-8
View File
@@ -5,7 +5,6 @@ package s3
import (
"context"
"encoding/json"
stderrors "errors"
"strings"
"github.com/OpenListTeam/OpenList/v4/internal/conf"
@@ -22,13 +21,6 @@ type Bucket struct {
Path string `json:"path"`
}
func mapBackendError(err error) error {
if stderrors.Is(err, errs.TemporaryCapacity) {
return gofakes3.ErrSlowDown
}
return err
}
const emptyObjectName = "ThisIsAnEmptyFolderInTheS3Bucket"
func getAndParseBuckets() ([]Bucket, error) {