优化
This commit is contained in:
truewhile
2026-08-24 17:00:18 +08:00
parent afae22ffe2
commit 0aff87c69d
39 changed files with 846 additions and 219 deletions
+1
View File
@@ -72,6 +72,7 @@ require (
golang.org/x/exp v0.0.0-20230905200255-921286631fa9 // indirect golang.org/x/exp v0.0.0-20230905200255-921286631fa9 // indirect
golang.org/x/net v0.21.0 // indirect golang.org/x/net v0.21.0 // indirect
golang.org/x/text v0.20.0 // indirect golang.org/x/text v0.20.0 // indirect
golang.org/x/time v0.15.0 // indirect
google.golang.org/protobuf v1.31.0 // indirect google.golang.org/protobuf v1.31.0 // indirect
gopkg.in/ini.v1 v1.67.0 // indirect gopkg.in/ini.v1 v1.67.0 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect
+2
View File
@@ -177,6 +177,8 @@ golang.org/x/sys v0.20.0 h1:Od9JTbYCk261bKm4M/mw7AklTlFYIa0bIp9BgSm1S8Y=
golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA= golang.org/x/sys v0.20.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/text v0.20.0 h1:gK/Kv2otX8gz+wn7Rmb3vT96ZwuoxnQlY+HlJVj7Qug= golang.org/x/text v0.20.0 h1:gK/Kv2otX8gz+wn7Rmb3vT96ZwuoxnQlY+HlJVj7Qug=
golang.org/x/text v0.20.0/go.mod h1:D4IsuqiFMhST5bX19pQ9ikHC2GsaKyk/oF+pn3ducp4= golang.org/x/text v0.20.0/go.mod h1:D4IsuqiFMhST5bX19pQ9ikHC2GsaKyk/oF+pn3ducp4=
golang.org/x/time v0.15.0 h1:bbrp8t3bGUeFOx08pvsMYRTCVSMk89u4tKbNOZbp88U=
golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0= golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp09yW+WbY/TyQbw= google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp09yW+WbY/TyQbw=
google.golang.org/protobuf v1.31.0 h1:g0LDEJHgrBl9N9r17Ru3sqWhkIx2NB67okBHPwC7hs8= google.golang.org/protobuf v1.31.0 h1:g0LDEJHgrBl9N9r17Ru3sqWhkIx2NB67okBHPwC7hs8=
+73 -42
View File
@@ -1,18 +1,16 @@
// 115 开放平台 HTTP 客户端(移植自 QMediaSync 的 v115open,去掉 resty 依赖, // 115 开放平台 HTTP 客户端(集成全局三级令牌桶限流、熔断保护与防 405 重定向策略)。
// 使用 net/http + 简单限流重试;只保留只读能力:授权/列目录/详情/下载直链)。
package cloud115 package cloud115
import ( import (
"bytes" "bytes"
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
"net/url" "net/url"
"strings" "strings"
"sync"
"sync/atomic"
"time" "time"
) )
@@ -22,21 +20,35 @@ type OpenClient struct {
HTTP *http.Client HTTP *http.Client
AccessToken string AccessToken string
RefreshTokenStr string RefreshTokenStr string
executor *QueueExecutor
// 全局 QPS 限流(115 开放平台免费额度较低)
lastSecond int64
reqInSecond int64
} }
var openClientMu sync.Mutex // default115HTTPClient 创建带有防 405 重定向保护的 http.Client。
func default115HTTPClient() *http.Client {
return &http.Client{
Timeout: 60 * time.Second,
// 防止 Go 标准库在遇到 301/302/307 重定向时将 POST 降级为 GET 导致 115 报 405 Method Not Allowed
CheckRedirect: func(req *http.Request, via []*http.Request) error {
if len(via) >= 5 {
return errors.New("stopped after 5 redirects")
}
// 如果原请求是 POST,重定向时不允许静默转为 GET(直接终止自动重定向,由业务层处理响应)
if len(via) > 0 && via[len(via)-1].Method == http.MethodPost {
return http.ErrUseLastResponse
}
return nil
},
}
}
// NewOpenClient 构造客户端。 // NewOpenClient 构造客户端。
func NewOpenClient(appID, accessToken, refreshToken string) *OpenClient { func NewOpenClient(appID, accessToken, refreshToken string) *OpenClient {
return &OpenClient{ return &OpenClient{
AppID: appID, AppID: appID,
HTTP: &http.Client{Timeout: 60 * time.Second}, HTTP: default115HTTPClient(),
AccessToken: accessToken, AccessToken: accessToken,
RefreshTokenStr: refreshToken, RefreshTokenStr: refreshToken,
executor: GetGlobalExecutor(),
} }
} }
@@ -46,30 +58,6 @@ func (c *OpenClient) SetAuthToken(accessToken, refreshToken string) {
c.RefreshTokenStr = refreshToken c.RefreshTokenStr = refreshToken
} }
// throttle 简单的每秒限流(默认 4 QPS,115 免费应用限额约 5 QPS)。
func (c *OpenClient) throttle(n int) {
for i := 0; i < n; i++ {
now := time.Now().Unix()
last := atomic.LoadInt64(&c.lastSecond)
if last != now {
if atomic.CompareAndSwapInt64(&c.lastSecond, last, now) {
atomic.StoreInt64(&c.reqInSecond, 0)
}
}
count := atomic.LoadInt64(&c.reqInSecond)
if count >= 4 {
time.Sleep(300 * time.Millisecond)
i--
continue
}
if atomic.CompareAndSwapInt64(&c.reqInSecond, count, count+1) {
return
}
time.Sleep(50 * time.Millisecond)
i--
}
}
// RespState 兼容 115 不同端点返回的 state 类型(proapi 返回布尔、passport 返回数字)。 // RespState 兼容 115 不同端点返回的 state 类型(proapi 返回布尔、passport 返回数字)。
type RespState bool type RespState bool
@@ -101,38 +89,65 @@ type RespBase struct {
Raw json.RawMessage `json:"-"` // 原始响应体(外层附加字段用) Raw json.RawMessage `json:"-"` // 原始响应体(外层附加字段用)
} }
// doJSON 执行 GET 请求并解析为统一响应;带 AccessToken(access=true 时)。 // doJSON 执行 HTTP 请求并解析为统一响应;带 AccessToken(access=true 时)。
func (c *OpenClient) doJSON(ctx context.Context, method, rawURL string, form map[string]string, access bool, retries int) (*RespBase, error) { func (c *OpenClient) doJSON(ctx context.Context, method, rawURL string, form map[string]string, access bool, retries int) (*RespBase, error) {
executor := c.executor
if executor == nil {
executor = GetGlobalExecutor()
}
var lastErr error var lastErr error
for attempt := 0; attempt <= retries; attempt++ { for attempt := 0; attempt <= retries; attempt++ {
c.throttle(1) // 1. 获取全局三级令牌桶(QPS/QPM/QPH)令牌,若在熔断状态则阻塞等待冷却
if err := executor.Acquire(ctx); err != nil {
return nil, err
}
req, err := c.buildRequest(ctx, method, rawURL, form, access) req, err := c.buildRequest(ctx, method, rawURL, form, access)
if err != nil { if err != nil {
return nil, err return nil, err
} }
resp, err := c.HTTP.Do(req) resp, err := c.HTTP.Do(req)
if err != nil { if err != nil {
lastErr = err lastErr = err
if attempt < retries { if attempt < retries {
time.Sleep(time.Duration(attempt+1) * 500 * time.Millisecond) time.Sleep(time.Duration(attempt+1) * 1 * time.Second)
continue continue
} }
return nil, err return nil, err
} }
body, readErr := io.ReadAll(io.LimitReader(resp.Body, 16<<20)) body, readErr := io.ReadAll(io.LimitReader(resp.Body, 16<<20))
_ = resp.Body.Close() _ = resp.Body.Close()
if readErr != nil { if readErr != nil {
lastErr = readErr lastErr = readErr
continue continue
} }
// 检查 HTTP 状态码:特别处理 405/406/429 等限流和异常阻断
if resp.StatusCode < 200 || resp.StatusCode >= 300 { if resp.StatusCode < 200 || resp.StatusCode >= 300 {
lastErr = fmt.Errorf("115 接口返回 HTTP %d:%s", resp.StatusCode, strings.TrimSpace(string(body))) if resp.StatusCode == http.StatusMethodNotAllowed || resp.StatusCode == http.StatusNotAcceptable || resp.StatusCode == http.StatusTooManyRequests {
// 触发全局熔断冷却,防止短时间连续重试导致 IP 被完全拉黑
executor.MarkThrottled()
lastErr = fmt.Errorf("115 接口触发频控/安全拦截(HTTP %d):%s", resp.StatusCode, strings.TrimSpace(string(body)))
if attempt < retries { if attempt < retries {
time.Sleep(time.Duration(attempt+1) * 500 * time.Millisecond) // 阻塞等待 115 限流冷却解封后自动重试当前请求
if waitErr := executor.throttleManager.WaitThrottleRecovery(ctx); waitErr != nil {
return nil, waitErr
}
continue continue
} }
return nil, lastErr return nil, lastErr
} }
lastErr = fmt.Errorf("115 接口返回 HTTP %d:%s", resp.StatusCode, strings.TrimSpace(string(body)))
if attempt < retries {
time.Sleep(time.Duration(attempt+1) * 1 * time.Second)
continue
}
return nil, lastErr
}
var base RespBase var base RespBase
if err := json.Unmarshal(body, &base); err != nil { if err := json.Unmarshal(body, &base); err != nil {
return nil, fmt.Errorf("115 接口响应解析失败:%w", err) return nil, fmt.Errorf("115 接口响应解析失败:%w", err)
@@ -141,13 +156,29 @@ func (c *OpenClient) doJSON(ctx context.Context, method, rawURL string, form map
if base.State { if base.State {
return &base, nil return &base, nil
} }
// 业务失败:限流/Token 错误不重试,其余按配置重试
if IsThrottleCode(base.Code) || isTokenCode(base.Code) { // 业务返回限流错误码(406 / 770004):标记全局熔断,等待冷却后重试
if IsThrottleCode(base.Code) {
executor.MarkThrottled()
lastErr = NewOpenAPIResponseError(base.Code, base.Errno, base.Message, base.Error, "115 访问频率达到上限,已进入冷却")
if attempt < retries {
// 阻塞等待限流冷却解封后自动重试当前请求
if waitErr := executor.throttleManager.WaitThrottleRecovery(ctx); waitErr != nil {
return nil, waitErr
}
continue
}
return &base, lastErr
}
// Token 失效不重试
if isTokenCode(base.Code) {
return &base, nil return &base, nil
} }
lastErr = NewOpenAPIResponseError(base.Code, base.Errno, base.Message, base.Error, "115 接口调用失败") lastErr = NewOpenAPIResponseError(base.Code, base.Errno, base.Message, base.Error, "115 接口调用失败")
if attempt < retries { if attempt < retries {
time.Sleep(time.Duration(attempt+1) * 500 * time.Millisecond) time.Sleep(time.Duration(attempt+1) * 1 * time.Second)
continue continue
} }
return &base, lastErr return &base, lastErr
@@ -7,6 +7,7 @@ import (
"net/url" "net/url"
"strings" "strings"
"testing" "testing"
"time"
) )
// mockTransport 用 httptest server 替换 API 基址。 // mockTransport 用 httptest server 替换 API 基址。
@@ -321,3 +322,56 @@ func TestSourceCatalog(t *testing.T) {
t.Fatal("relay/thrid-party sources broken") t.Fatal("relay/thrid-party sources broken")
} }
} }
func TestHTTP405AndThrottleRecovery(t *testing.T) {
tm := NewThrottleManager(100 * time.Millisecond)
qe := NewQueueExecutor(10, 200, 12000)
qe.throttleManager = tm
mockAPI(t, func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusMethodNotAllowed)
w.Write([]byte("Method Not Allowed"))
})
c := NewOpenClient("100195125", "at1", "rt1")
c.executor = qe
_, err := c.GetDownloadURL(context.Background(), "pickTest")
if err == nil {
t.Fatal("expected 405 error")
}
if !tm.IsThrottled() {
t.Fatal("405 should trigger throttle status")
}
// 验证熔断冷却后能正常恢复
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Second)
defer cancel()
if err := tm.WaitThrottleRecovery(ctx); err != nil {
t.Fatalf("wait throttle recovery failed: %v", err)
}
if tm.IsThrottled() {
t.Fatal("throttle status should be cleared after duration")
}
}
func TestThrottleCodeHandling(t *testing.T) {
tm := NewThrottleManager(100 * time.Millisecond)
qe := NewQueueExecutor(10, 200, 12000)
qe.throttleManager = tm
mockAPI(t, func(w http.ResponseWriter, r *http.Request) {
w.Write([]byte(`{"state":false,"code":770004,"message":"访问频率过高"}`))
})
c := NewOpenClient("100195125", "at1", "rt1")
c.executor = qe
_, _, err := c.GetFsList(context.Background(), "0", 0, 100)
if err == nil {
t.Fatal("expected throttle error")
}
if !tm.IsThrottled() {
t.Fatal("code 770004 should trigger throttle status")
}
}
+2 -2
View File
@@ -37,8 +37,8 @@ const (
RefreshTokenCheckFailed = 40140120 RefreshTokenCheckFailed = 40140120
) )
// DefaultUA 请求 UA。 // DefaultUA 请求 UA(标准 API 客户端格式,避免伪装浏览器触发 115 WAF 拦截与 405 错误)。
const DefaultUA = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/124.0 Safari/537.36 MMTL-115-OpenAPI/1.0" const DefaultUA = "MMTL-115-GoClient/1.0"
// OpenAPIError 保留 115 开放平台返回的原始错误信息。 // OpenAPIError 保留 115 开放平台返回的原始错误信息。
type OpenAPIError struct { type OpenAPIError struct {
+106
View File
@@ -0,0 +1,106 @@
package cloud115
import (
"context"
"sync"
"time"
"golang.org/x/time/rate"
)
// QueueExecutor 全局 115 请求调度执行器,通过三级令牌桶(QPS / QPM / QPH)平滑所有外发请求。
type QueueExecutor struct {
sync.RWMutex
qpsLimiter *rate.Limiter
qpmLimiter *rate.Limiter
qphLimiter *rate.Limiter
throttleManager *ThrottleManager
qpsConfig int
qpmConfig int
qphConfig int
}
var (
globalExecutor *QueueExecutor
executorOnce sync.Once
)
// GetGlobalExecutor 获取全局队列执行器单例(默认 QPS=2, QPM=120, QPH=6000,保障 115 API 调用安全不超频)。
func GetGlobalExecutor() *QueueExecutor {
executorOnce.Do(func() {
globalExecutor = NewQueueExecutor(2, 120, 6000)
})
return globalExecutor
}
// NewQueueExecutor 创建新的队列调度执行器。
func NewQueueExecutor(qps, qpm, qph int) *QueueExecutor {
if qps <= 0 {
qps = 3
}
if qpm <= 0 {
qpm = 200
}
if qph <= 0 {
qph = 12000
}
return &QueueExecutor{
qpsLimiter: rate.NewLimiter(rate.Limit(qps), qps),
qpmLimiter: rate.NewLimiter(rate.Every(time.Minute/time.Duration(qpm)), qpm),
qphLimiter: rate.NewLimiter(rate.Every(time.Hour/time.Duration(qph)), qph),
throttleManager: GetGlobalThrottleManager(),
qpsConfig: qps,
qpmConfig: qpm,
qphConfig: qph,
}
}
// Acquire 统一在发送请求前获取令牌,并等待熔断恢复。
func (qe *QueueExecutor) Acquire(ctx context.Context) error {
// 1. 如果处于限流状态,先阻塞等待 60s 静默冷却结束
if err := qe.throttleManager.WaitThrottleRecovery(ctx); err != nil {
return err
}
// 2. 依次获取小时级、分钟级、秒级令牌
if err := qe.qphLimiter.Wait(ctx); err != nil {
return err
}
if err := qe.qpmLimiter.Wait(ctx); err != nil {
return err
}
if err := qe.qpsLimiter.Wait(ctx); err != nil {
return err
}
return nil
}
// SetRateLimitConfig 动态调整速率限制。
func (qe *QueueExecutor) SetRateLimitConfig(qps, qpm, qph int) {
qe.Lock()
defer qe.Unlock()
if qps > 0 {
qe.qpsConfig = qps
qe.qpsLimiter.SetLimit(rate.Limit(qps))
qe.qpsLimiter.SetBurst(qps)
}
if qpm > 0 {
qe.qpmConfig = qpm
qe.qpmLimiter.SetLimit(rate.Every(time.Minute / time.Duration(qpm)))
qe.qpmLimiter.SetBurst(qpm)
}
if qph > 0 {
qe.qphConfig = qph
qe.qphLimiter.SetLimit(rate.Every(time.Hour / time.Duration(qph)))
qe.qphLimiter.SetBurst(qph)
}
}
// MarkThrottled 触发限流熔断。
func (qe *QueueExecutor) MarkThrottled() {
qe.throttleManager.MarkThrottled()
}
@@ -0,0 +1,119 @@
package cloud115
import (
"context"
"sync"
"time"
)
// ThrottleManager 全局限流/熔断管理器,用于控制 115 API 访问频率与被风控时的全局静默冷却。
type ThrottleManager struct {
sync.RWMutex
isThrottled bool
throttleStartTime time.Time
throttleNotify chan struct{}
throttleDuration time.Duration
}
var (
globalThrottleManager *ThrottleManager
throttleOnce sync.Once
)
// GetGlobalThrottleManager 获取全局限流管理器单例。
func GetGlobalThrottleManager() *ThrottleManager {
throttleOnce.Do(func() {
globalThrottleManager = NewThrottleManager(60 * time.Second)
})
return globalThrottleManager
}
// NewThrottleManager 创建限流管理器。
func NewThrottleManager(duration time.Duration) *ThrottleManager {
if duration <= 0 {
duration = 60 * time.Second
}
return &ThrottleManager{
isThrottled: false,
throttleNotify: make(chan struct{}),
throttleDuration: duration,
}
}
// IsThrottled 检查是否处于限流冷却状态。
func (tm *ThrottleManager) IsThrottled() bool {
tm.RLock()
defer tm.RUnlock()
if !tm.isThrottled {
return false
}
if time.Since(tm.throttleStartTime) >= tm.throttleDuration {
return false
}
return true
}
// MarkThrottled 标记为限流状态,并启动恢复计时器。
func (tm *ThrottleManager) MarkThrottled() {
tm.Lock()
defer tm.Unlock()
if tm.isThrottled && time.Since(tm.throttleStartTime) < tm.throttleDuration {
return
}
tm.isThrottled = true
tm.throttleStartTime = time.Now()
go tm.startRecoveryTimer()
}
func (tm *ThrottleManager) startRecoveryTimer() {
time.Sleep(tm.throttleDuration)
tm.Lock()
defer tm.Unlock()
tm.isThrottled = false
select {
case tm.throttleNotify <- struct{}{}:
default:
}
}
// WaitThrottleRecovery 如果处于限流状态,挂起等待直到冷却恢复。
func (tm *ThrottleManager) WaitThrottleRecovery(ctx context.Context) error {
for {
if !tm.IsThrottled() {
return nil
}
ticker := time.NewTicker(200 * time.Millisecond)
select {
case <-ctx.Done():
ticker.Stop()
return ctx.Err()
case <-ticker.C:
ticker.Stop()
if !tm.IsThrottled() {
return nil
}
}
}
}
// Reset 重置限流状态(主要用于单元测试)。
func (tm *ThrottleManager) Reset() {
tm.Lock()
defer tm.Unlock()
tm.isThrottled = false
}
// SetDuration 调整限流冷却时间(用于测试或动态配置)。
func (tm *ThrottleManager) SetDuration(d time.Duration) {
tm.Lock()
defer tm.Unlock()
tm.throttleDuration = d
}
@@ -148,6 +148,7 @@ func joinRemotePath(base, rel string) string {
} }
return path.Clean(path.Join(parts...)) return path.Clean(path.Join(parts...))
} }
// 通用库路径辅助(原云盘库/STRM 生成逻辑使用;网盘后端移除后保留为纯工具函数, // 通用库路径辅助(原云盘库/STRM 生成逻辑使用;网盘后端移除后保留为纯工具函数,
// 供库路径构建与既有测试作为稳定夹具使用)。 // 供库路径构建与既有测试作为稳定夹具使用)。
+11 -2
View File
@@ -5,13 +5,22 @@ import (
"strings" "strings"
) )
// sanitizeFilename removes characters not safe for filesystem names. // sanitizeFilename removes characters not safe for filesystem names across Windows, Linux, and macOS.
func sanitizeFilename(s string) string { func sanitizeFilename(s string) string {
r := strings.NewReplacer( r := strings.NewReplacer(
"/", " ", "\\", " ", ":", " ", "*", "", "?", "", "/", " ", "\\", " ", ":", " ", "*", "", "?", "",
"\"", "", "<", "", ">", "", "|", "", "\"", "", "<", "", ">", "", "|", "",
) )
return strings.TrimSpace(r.Replace(s)) replaced := r.Replace(s)
var sb strings.Builder
for _, ch := range replaced {
// Strip ASCII control characters (0-31, 127)
if ch < 32 || ch == 127 {
continue
}
sb.WriteRune(ch)
}
return strings.Join(strings.Fields(sb.String()), " ")
} }
func (o *OrganizerService) organizeRoot(libraryPath, mediaType, category string) string { func (o *OrganizerService) organizeRoot(libraryPath, mediaType, category string) string {
@@ -93,6 +93,7 @@ func (s *ScraperService) writeArtworkDataToPath(dir, name, ctype string, data []
if len(data) == 0 { if len(data) == 0 {
return "" return ""
} }
dir = sanitizeLocalPath(dir)
if err := os.MkdirAll(dir, 0o755); err != nil { if err := os.MkdirAll(dir, 0o755); err != nil {
s.log.Warn("scrape artwork mkdir failed", zap.String("dir", dir), zap.Error(err)) s.log.Warn("scrape artwork mkdir failed", zap.String("dir", dir), zap.Error(err))
return "" return ""
+6
View File
@@ -59,6 +59,11 @@ func (s *StrmService) downloadWorker(ctx context.Context) {
} }
func (s *StrmService) processDownloadTask(ctx context.Context, task *model.StrmDownloadTask) { func (s *StrmService) processDownloadTask(ctx context.Context, task *model.StrmDownloadTask) {
cleanPath := sanitizeLocalPath(task.LocalPath)
if cleanPath != "" && cleanPath != task.LocalPath {
task.LocalPath = cleanPath
_ = s.repo.StrmDownload.Update(context.Background(), task)
}
finish := func(status, message string) { finish := func(status, message string) {
now := time.Now() now := time.Now()
task.Status = status task.Status = status
@@ -210,6 +215,7 @@ func retryTask(retryCount *int, status *string, errMsg *string, nextTryAt **time
// downloadToFile 把直链内容下载到目标文件(临时文件 + 原子改名)。 // downloadToFile 把直链内容下载到目标文件(临时文件 + 原子改名)。
func downloadToFile(ctx context.Context, link *cloud.DirectLink, target string, client *http.Client) error { func downloadToFile(ctx context.Context, link *cloud.DirectLink, target string, client *http.Client) error {
target = sanitizeLocalPath(target)
if link == nil || link.URL == "" { if link == nil || link.URL == "" {
return errors.New("空下载地址") return errors.New("空下载地址")
} }
+150 -3
View File
@@ -572,14 +572,161 @@ const defaultStrmUA = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537
// ensureLocalDir 创建本地输出目录。 // ensureLocalDir 创建本地输出目录。
func ensureLocalDir(dir string) error { func ensureLocalDir(dir string) error {
return os.MkdirAll(dir, 0o755) return os.MkdirAll(sanitizeLocalPath(dir), 0o755)
}
// windowsReservedNames 包含 Windows 系统底层保留的设备名称(大小写不敏感)。
var windowsReservedNames = map[string]bool{
"CON": true, "PRN": true, "AUX": true, "NUL": true,
"COM1": true, "COM2": true, "COM3": true, "COM4": true, "COM5": true,
"COM6": true, "COM7": true, "COM8": true, "COM9": true,
"LPT1": true, "LPT2": true, "LPT3": true, "LPT4": true, "LPT5": true,
"LPT6": true, "LPT7": true, "LPT8": true, "LPT9": true,
}
func isWindowsReservedName(name string) bool {
return windowsReservedNames[strings.ToUpper(name)]
}
// truncateStringRuneSafe 安全截断 UTF-8 字符串至指定字节长度,不切断中文字符或多字节 rune。
func truncateStringRuneSafe(s string, maxBytes int) string {
if len(s) <= maxBytes {
return s
}
b := []byte(s)
if len(b) <= maxBytes {
return s
}
for maxBytes > 0 && (b[maxBytes]&0xC0 == 0x80) {
maxBytes--
}
return strings.TrimRight(string(b[:maxBytes]), ". ")
}
// cleanEntryName 清理单个目录名或文件名中的非法字符、控制字符、尾部点空格及 Windows 保留字,
// 确保在 Windows (NTFS/FAT)、Linux (ext4/btrfs/xfs) 及 NAS/SMB 挂载环境下均安全可用。
func cleanEntryName(name string, isDir bool) string {
name = strings.TrimSpace(name)
if name == "" {
return "unnamed"
}
if isDir {
clean := sanitizeFilename(name)
clean = strings.Trim(clean, ". ")
if clean == "" {
return "unnamed"
}
if isWindowsReservedName(clean) {
clean = "_" + clean
}
return truncateStringRuneSafe(clean, 255)
}
ext := filepath.Ext(name)
base := strings.TrimSuffix(name, ext)
// 清洗扩展名中的非法字符
cleanExt := sanitizeFilename(ext)
cleanExt = strings.TrimSpace(cleanExt)
if cleanExt != "" && !strings.HasPrefix(cleanExt, ".") {
cleanExt = "." + cleanExt
}
cleanBase := sanitizeFilename(base)
cleanBase = strings.Trim(cleanBase, ". ")
if cleanBase == "" {
cleanBase = "unnamed"
}
if isWindowsReservedName(cleanBase) {
cleanBase = "_" + cleanBase
}
maxBaseBytes := 255 - len(cleanExt)
if maxBaseBytes < 10 {
maxBaseBytes = 255
cleanExt = ""
}
cleanBase = truncateStringRuneSafe(cleanBase, maxBaseBytes)
return cleanBase + cleanExt
}
// sanitizeRelativePath 清理相对路径中的非法字符与首尾点空格,确保跨操作系统路径合法。
func sanitizeRelativePath(rel string) string {
rel = strings.TrimSpace(rel)
if rel == "" {
return ""
}
rel = strings.ReplaceAll(rel, "\\", "/")
parts := strings.Split(rel, "/")
out := make([]string, 0, len(parts))
for i, part := range parts {
part = strings.TrimSpace(part)
if part == "" || part == "." || part == ".." {
continue
}
isDir := i < len(parts)-1
clean := cleanEntryName(part, isDir)
if clean == "" || clean == "." || clean == ".." {
continue
}
out = append(out, clean)
}
if len(out) == 0 {
return ""
}
return filepath.Join(out...)
}
// sanitizeLocalPath 对本地全路径中除根目录/盘符/UNC之外的各级目录及文件名进行跨平台非法字符清洗。
func sanitizeLocalPath(p string) string {
p = strings.TrimSpace(p)
if p == "" || p == "." {
return p
}
vol := filepath.VolumeName(p)
rest := p[len(vol):]
hasRootSlash := len(rest) > 0 && (rest[0] == '/' || rest[0] == '\\')
parts := strings.FieldsFunc(rest, func(r rune) bool {
return r == '/' || r == '\\'
})
cleanedParts := make([]string, 0, len(parts))
for i, part := range parts {
part = strings.TrimSpace(part)
if part == "" || part == "." || part == ".." {
continue
}
isDir := i < len(parts)-1
clean := cleanEntryName(part, isDir)
if clean == "" || clean == "." || clean == ".." {
continue
}
cleanedParts = append(cleanedParts, clean)
}
prefix := vol
if hasRootSlash {
prefix += string(filepath.Separator)
}
if len(cleanedParts) == 0 {
return filepath.Clean(prefix)
}
if prefix == "" {
return filepath.Clean(filepath.Join(cleanedParts...))
}
return filepath.Clean(filepath.Join(prefix, filepath.Join(cleanedParts...)))
} }
// joinLocalRel 拼接本地目标路径并校验不越出根目录。 // joinLocalRel 拼接本地目标路径并校验不越出根目录。
func joinLocalRel(root, rel string) (string, error) { func joinLocalRel(root, rel string) (string, error) {
root = filepath.Clean(root) root = filepath.Clean(root)
target := filepath.Clean(filepath.Join(root, filepath.FromSlash(rel))) cleanRel := sanitizeRelativePath(rel)
if target != root && !strings.HasPrefix(target, root+string(filepath.Separator)) { if cleanRel == "" {
return "", errors.New("无效相对路径")
}
target := filepath.Clean(filepath.Join(root, cleanRel))
relToRoot, err := filepath.Rel(root, target)
if err != nil || strings.HasPrefix(relToRoot, "..") || (relToRoot == "." && target != root) {
return "", errors.New("本地路径越界") return "", errors.New("本地路径越界")
} }
return target, nil return target, nil
+20 -8
View File
@@ -253,9 +253,10 @@ func (st *strmSyncState) walkRemote() error {
return fmt.Errorf("列出远端目录 %s 失败:%w", task.id, err) return fmt.Errorf("列出远端目录 %s 失败:%w", task.id, err)
} }
for _, entry := range entries { for _, entry := range entries {
rel := entry.Name cleanName := cleanEntryName(entry.Name, entry.IsDir)
rel := cleanName
if task.rel != "" { if task.rel != "" {
rel = task.rel + "/" + entry.Name rel = task.rel + "/" + cleanName
} }
if entry.IsDir { if entry.IsDir {
queue = append(queue, dirTask{id: entry.ID, rel: rel}) queue = append(queue, dirTask{id: entry.ID, rel: rel})
@@ -277,8 +278,13 @@ func (st *strmSyncState) processRemoteFile(entry cloud.FileEntry, rel string) {
switch { switch {
case st.isVideoExt(ext, entry.Size): case st.isVideoExt(ext, entry.Size):
st.handleVideo(entry, rel, ext) st.handleVideo(entry, rel, ext)
case st.cfg.DownloadMeta && st.isMetaExt(ext): case st.isMetaExt(ext):
st.recordRemoteMeta(entry, rel)
if st.cfg.DownloadMeta {
st.handleMeta(entry, rel, ext) st.handleMeta(entry, rel, ext)
} else {
st.touchProgress()
}
default: default:
st.touchProgress() st.touchProgress()
} }
@@ -414,12 +420,17 @@ func (st *strmSyncState) strmPathParam(rel string) string {
} }
} }
// handleMeta 元数据入下载队列(本地已存在且大小一致则跳过)。 // recordRemoteMeta 记录远端存在的元数据索引及文件大小。
func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) { func (st *strmSyncState) recordRemoteMeta(entry cloud.FileEntry, rel string) {
st.mu.Lock() st.mu.Lock()
st.seenMeta["m:"+rel] = true st.seenMeta["m:"+rel] = true
st.remoteMeta["m:"+rel] = entry.Size st.remoteMeta["m:"+rel] = entry.Size
st.mu.Unlock() st.mu.Unlock()
}
// handleMeta 元数据入下载队列(本地已存在且大小一致则跳过)。
func (st *strmSyncState) handleMeta(entry cloud.FileEntry, rel, ext string) {
st.recordRemoteMeta(entry, rel)
target, err := joinLocalRel(st.p.LocalPath, rel) target, err := joinLocalRel(st.p.LocalPath, rel)
if err != nil { if err != nil {
@@ -569,12 +580,13 @@ func (st *strmSyncState) scanLocalMetaForUpload() error {
return nil return nil
} }
st.mu.Lock() st.mu.Lock()
remoteSize, exists := st.remoteMeta["m:"+rel] _, exists := st.remoteMeta["m:"+rel]
st.mu.Unlock() st.mu.Unlock()
if exists && remoteSize == info.Size() { if exists {
// 网盘端已存在该元数据文件,跳过上传
return nil return nil
} }
if st.taskExists("upload", st.p.ID, rel) { if st.taskExists("upload", st.p.ID, path) {
return nil return nil
} }
task := &model.StrmUploadTask{ task := &model.StrmUploadTask{
+138
View File
@@ -202,3 +202,141 @@ func TestStrmDefaultBaseURL(t *testing.T) {
t.Fatal("effective config base url should fall back to default") t.Fatal("effective config base url should fall back to default")
} }
} }
func TestSanitizePathWithSpecialChars(t *testing.T) {
cases := []struct {
input string
isDir bool
wantName string
}{
{"数码宝贝:拯救者", true, "数码宝贝 拯救者"},
{"数码宝贝:拯救者.mp4", false, "数码宝贝 拯救者.mp4"},
{"Season 01.", true, "Season 01"},
{"file:name?test*<foo>|bar\".mkv", false, "file nametestfoobar.mkv"},
{"poster.jpg", false, "poster.jpg"},
{" trailing space ", true, "trailing space"},
{"CON", true, "_CON"},
{"aux.mp4", false, "_aux.mp4"},
{"NUL.nfo", false, "_NUL.nfo"},
{"COM1", false, "_COM1"},
{"test\x00\x1f\x7fcontrol.mkv", false, "testcontrol.mkv"},
}
for _, tc := range cases {
got := cleanEntryName(tc.input, tc.isDir)
if got != tc.wantName {
t.Errorf("cleanEntryName(%q, %v) = %q, want %q", tc.input, tc.isDir, got, tc.wantName)
}
}
// Test sanitizeRelativePath
rel := "动漫/数码宝贝:拯救者/数码宝贝:拯救者.mp4"
wantRel := filepath.Join("动漫", "数码宝贝 拯救者", "数码宝贝 拯救者.mp4")
if got := sanitizeRelativePath(rel); got != wantRel {
t.Errorf("sanitizeRelativePath(%q) = %q, want %q", rel, got, wantRel)
}
// Test sanitizeLocalPath - Windows Drive
localWin := `D:\test\动漫\数码宝贝:拯救者\poster.jpg`
wantLocalWin := filepath.Join(`D:\`, "test", "动漫", "数码宝贝 拯救者", "poster.jpg")
if got := sanitizeLocalPath(localWin); got != wantLocalWin {
t.Errorf("sanitizeLocalPath(%q) = %q, want %q", localWin, got, wantLocalWin)
}
// Test sanitizeLocalPath - Linux root
localLinux := `/media/data/动漫/数码宝贝:拯救者/ep01:拯救者.mkv`
cleanedLinux := sanitizeLocalPath(localLinux)
if strings.Contains(cleanedLinux, ":") && !strings.HasPrefix(cleanedLinux, "D:") && !strings.HasPrefix(cleanedLinux, "C:") {
t.Errorf("sanitizeLocalPath should not contain colons in path components: %q", cleanedLinux)
}
// Test joinLocalRel
root := t.TempDir()
joined, err := joinLocalRel(root, "动漫/数码宝贝:拯救者/poster.jpg")
if err != nil {
t.Fatalf("joinLocalRel failed: %v", err)
}
wantJoined := filepath.Join(root, "动漫", "数码宝贝 拯救者", "poster.jpg")
if joined != wantJoined {
t.Errorf("joinLocalRel = %q, want %q", joined, wantJoined)
}
// Test joinLocalRel prevents path traversal
if _, err := joinLocalRel(root, "../../"); err == nil {
t.Error("joinLocalRel should error on empty/traversal-only path")
}
traversalCleaned, err := joinLocalRel(root, "../../etc/passwd")
if err != nil {
t.Fatalf("joinLocalRel failed on traversal path: %v", err)
}
if !strings.HasPrefix(traversalCleaned, root) {
t.Errorf("joinLocalRel allowed path escape: %s", traversalCleaned)
}
// Verify mkdir succeeds on this joined path
if err := os.MkdirAll(filepath.Dir(joined), 0o755); err != nil {
t.Fatalf("MkdirAll on joined path failed: %v", err)
}
if info, err := os.Stat(filepath.Dir(joined)); err != nil || !info.IsDir() {
t.Fatalf("Dir not created: %v", err)
}
}
// TestScanLocalMetaForUpload 验证元数据上传比对逻辑:网盘已存在跳过,网盘不存在才入队上传。
func TestScanLocalMetaForUpload(t *testing.T) {
svc := testStrmService(t)
localDir := t.TempDir()
// 本地有 2 个元数据:poster.jpg 和 fanart.jpg
writeFile(t, filepath.Join(localDir, "动漫", "poster.jpg"), "poster-data")
writeFile(t, filepath.Join(localDir, "动漫", "fanart.jpg"), "fanart-data")
p := &model.StrmSyncPath{
Base: model.Base{ID: "test-path-upload"},
AccountID: "acct-1",
Provider: model.StrmProviderCloudDrive,
RemotePath: "/Media",
LocalPath: localDir,
UploadMeta: true,
}
st := &strmSyncState{
s: svc,
ctx: context.Background(),
p: p,
cfg: &strmPathConfig{UploadMeta: true, MetaExt: []string{"jpg", "nfo"}},
rec: &model.StrmSyncRecord{},
seenMeta: map[string]bool{},
remoteMeta: map[string]int64{},
}
// 模拟远端已存在 poster.jpg
st.remoteMeta["m:动漫/poster.jpg"] = 1000
if err := st.scanLocalMetaForUpload(); err != nil {
t.Fatalf("scanLocalMetaForUpload failed: %v", err)
}
// 此时应该只有 fanart.jpg 入队上传,poster.jpg 被跳过
tasks, _, err := svc.repo.StrmUpload.List(context.Background(), "", 1, 10)
if err != nil {
t.Fatal(err)
}
if len(tasks) != 1 {
t.Fatalf("expected 1 upload task (fanart.jpg), got %d", len(tasks))
}
if tasks[0].FileName != "fanart.jpg" {
t.Errorf("expected upload task for fanart.jpg, got %s", tasks[0].FileName)
}
// 再次扫描:fanart.jpg 已经在队列中,应自动去重,不重复入队
if err := st.scanLocalMetaForUpload(); err != nil {
t.Fatalf("second scanLocalMetaForUpload failed: %v", err)
}
tasks, _, err = svc.repo.StrmUpload.List(context.Background(), "", 1, 10)
if err != nil {
t.Fatal(err)
}
if len(tasks) != 1 {
t.Fatalf("expected still 1 upload task after dedup, got %d", len(tasks))
}
}