mirror of
https://github.com/truewhile/MeBox.git
synced 2026-10-06 21:36:37 +08:00
fix: validate downloader RPC endpoints
This commit is contained in:
@@ -57,6 +57,11 @@ func (h *DownloadClientHandler) Create(c *gin.Context) {
|
|||||||
|
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
_ = h.svc.Repo.Setting.Set(ctx, "download_clients.managed", "true")
|
_ = h.svc.Repo.Setting.Set(ctx, "download_clients.managed", "true")
|
||||||
|
normalizedHost, err := service.NormalizeDownloadClientHost(req.Type, req.Host)
|
||||||
|
if err != nil {
|
||||||
|
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
// 加密密码
|
// 加密密码
|
||||||
password := req.Password
|
password := req.Password
|
||||||
@@ -82,7 +87,7 @@ func (h *DownloadClientHandler) Create(c *gin.Context) {
|
|||||||
client := &model.DownloadClient{
|
client := &model.DownloadClient{
|
||||||
Name: req.Name,
|
Name: req.Name,
|
||||||
Type: req.Type,
|
Type: req.Type,
|
||||||
Host: req.Host,
|
Host: normalizedHost,
|
||||||
Username: req.Username,
|
Username: req.Username,
|
||||||
Password: password,
|
Password: password,
|
||||||
IsDefault: req.IsDefault,
|
IsDefault: req.IsDefault,
|
||||||
@@ -152,7 +157,16 @@ func (h *DownloadClientHandler) Update(c *gin.Context) {
|
|||||||
client.Type = req.Type
|
client.Type = req.Type
|
||||||
}
|
}
|
||||||
if req.Host != "" {
|
if req.Host != "" {
|
||||||
client.Host = req.Host
|
clientType := client.Type
|
||||||
|
if req.Type != "" {
|
||||||
|
clientType = req.Type
|
||||||
|
}
|
||||||
|
normalizedHost, err := service.NormalizeDownloadClientHost(clientType, req.Host)
|
||||||
|
if err != nil {
|
||||||
|
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
client.Host = normalizedHost
|
||||||
}
|
}
|
||||||
if req.Username != "" {
|
if req.Username != "" {
|
||||||
client.Username = req.Username
|
client.Username = req.Username
|
||||||
@@ -180,6 +194,12 @@ func (h *DownloadClientHandler) Update(c *gin.Context) {
|
|||||||
client.Extra = extraStr
|
client.Extra = extraStr
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
normalizedHost, err := service.NormalizeDownloadClientHost(client.Type, client.Host)
|
||||||
|
if err != nil {
|
||||||
|
Error(c, http.StatusBadRequest, ErrInvalidParams, err.Error())
|
||||||
|
return
|
||||||
|
}
|
||||||
|
client.Host = normalizedHost
|
||||||
|
|
||||||
if err := h.svc.Repo.DownloadClient.Update(ctx, client); err != nil {
|
if err := h.svc.Repo.DownloadClient.Update(ctx, client); err != nil {
|
||||||
Error(c, http.StatusInternalServerError, ErrInternal, "更新失败")
|
Error(c, http.StatusInternalServerError, ErrInternal, "更新失败")
|
||||||
|
|||||||
@@ -57,6 +57,11 @@ func NewAria2Adapter() *Aria2Adapter {
|
|||||||
func (a *Aria2Adapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error {
|
func (a *Aria2Adapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error {
|
||||||
a.mu.Lock()
|
a.mu.Lock()
|
||||||
defer a.mu.Unlock()
|
defer a.mu.Unlock()
|
||||||
|
endpoint, err := normalizeDownloadClientEndpoint("aria2", cfg.Host)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
cfg.Host = endpoint
|
||||||
a.cfg = cfg
|
a.cfg = cfg
|
||||||
a.idSeq = 0
|
a.idSeq = 0
|
||||||
return a.getVersionLocked(ctx)
|
return a.getVersionLocked(ctx)
|
||||||
@@ -71,9 +76,9 @@ func (a *Aria2Adapter) Ping(ctx context.Context) error {
|
|||||||
|
|
||||||
// getVersionLocked 内部版本检查(调用者必须持有锁)。
|
// getVersionLocked 内部版本检查(调用者必须持有锁)。
|
||||||
func (a *Aria2Adapter) getVersionLocked(ctx context.Context) error {
|
func (a *Aria2Adapter) getVersionLocked(ctx context.Context) error {
|
||||||
rpcURL := a.cfg.Host
|
rpcURL, err := downloadClientRPCURL("aria2", a.cfg.Host)
|
||||||
if !strings.HasSuffix(rpcURL, "/jsonrpc") {
|
if err != nil {
|
||||||
rpcURL = strings.TrimRight(rpcURL, "/") + "/jsonrpc"
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
req := &aria2Request{
|
req := &aria2Request{
|
||||||
@@ -88,7 +93,7 @@ func (a *Aria2Adapter) getVersionLocked(ctx context.Context) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
|
httpReq, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -110,9 +115,9 @@ func (a *Aria2Adapter) getVersionLocked(ctx context.Context) error {
|
|||||||
|
|
||||||
// rpcLocked 发送 JSON-RPC 请求(调用者必须持有锁)。
|
// rpcLocked 发送 JSON-RPC 请求(调用者必须持有锁)。
|
||||||
func (a *Aria2Adapter) rpcLocked(ctx context.Context, method string, params []interface{}) (json.RawMessage, error) {
|
func (a *Aria2Adapter) rpcLocked(ctx context.Context, method string, params []interface{}) (json.RawMessage, error) {
|
||||||
rpcURL := a.cfg.Host
|
rpcURL, err := downloadClientRPCURL("aria2", a.cfg.Host)
|
||||||
if !strings.HasSuffix(rpcURL, "/jsonrpc") {
|
if err != nil {
|
||||||
rpcURL = strings.TrimRight(rpcURL, "/") + "/jsonrpc"
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
if params == nil {
|
if params == nil {
|
||||||
@@ -145,7 +150,7 @@ func (a *Aria2Adapter) rpcLocked(ctx context.Context, method string, params []in
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
|
httpReq, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,115 @@
|
|||||||
|
package service
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"regexp"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
var downloadClientEndpointPattern = regexp.MustCompile(`^https?://(?:[A-Za-z0-9.-]+|\[[0-9A-Fa-f:.]+\])(?::[0-9]{1,5})?(?:/[A-Za-z0-9._~%!$&'()*+,;=:@/-]*)?$`)
|
||||||
|
|
||||||
|
func NormalizeDownloadClientHost(clientType, raw string) (string, error) {
|
||||||
|
return normalizeDownloadClientEndpoint(clientType, raw)
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeDownloadClientEndpoint(clientType, raw string) (string, error) {
|
||||||
|
raw = strings.TrimSpace(raw)
|
||||||
|
if raw == "" {
|
||||||
|
return "", errors.New("host required")
|
||||||
|
}
|
||||||
|
if strings.ContainsAny(raw, "\r\n\t") {
|
||||||
|
return "", errors.New("host contains invalid control characters")
|
||||||
|
}
|
||||||
|
if !strings.Contains(raw, "://") {
|
||||||
|
raw = "http://" + raw
|
||||||
|
}
|
||||||
|
if !downloadClientEndpointPattern.MatchString(raw) {
|
||||||
|
return "", errors.New("host must be a valid http(s) URL without username, query, or fragment")
|
||||||
|
}
|
||||||
|
parsed, err := url.Parse(raw)
|
||||||
|
if err != nil || parsed.Scheme == "" || parsed.Host == "" {
|
||||||
|
return "", errors.New("host must be a valid http(s) URL")
|
||||||
|
}
|
||||||
|
scheme := strings.ToLower(parsed.Scheme)
|
||||||
|
if scheme != "http" && scheme != "https" {
|
||||||
|
return "", errors.New("host only supports http or https")
|
||||||
|
}
|
||||||
|
if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" {
|
||||||
|
return "", errors.New("host must not include username, query, or fragment")
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(parsed.Hostname()) == "" {
|
||||||
|
return "", errors.New("host must include a hostname")
|
||||||
|
}
|
||||||
|
if port := parsed.Port(); port != "" {
|
||||||
|
n, err := strconv.Atoi(port)
|
||||||
|
if err != nil || n < 1 || n > 65535 {
|
||||||
|
return "", errors.New("host port must be between 1 and 65535")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := validateDownloadClientPath(clientType, parsed.Path); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
parsed.Scheme = scheme
|
||||||
|
parsed.Path = strings.TrimRight(parsed.Path, "/")
|
||||||
|
parsed.RawPath = ""
|
||||||
|
parsed.RawQuery = ""
|
||||||
|
parsed.Fragment = ""
|
||||||
|
return strings.TrimRight(parsed.String(), "/"), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func validateDownloadClientPath(clientType, rawPath string) error {
|
||||||
|
rawPath = strings.TrimSpace(rawPath)
|
||||||
|
if rawPath == "" || rawPath == "/" {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
for _, segment := range strings.Split(rawPath, "/") {
|
||||||
|
if segment == "." || segment == ".." {
|
||||||
|
return errors.New("host path must not contain traversal segments")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
switch clientType {
|
||||||
|
case "qbittorrent", "aria2", "transmission":
|
||||||
|
return nil
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("unsupported client type %q", clientType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func downloadClientRPCURL(clientType, host string) (string, error) {
|
||||||
|
base, err := normalizeDownloadClientEndpoint(clientType, host)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
u, err := url.Parse(base)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
switch clientType {
|
||||||
|
case "aria2":
|
||||||
|
if !strings.HasSuffix(strings.ToLower(u.Path), "/jsonrpc") {
|
||||||
|
u.Path = strings.TrimRight(u.Path, "/") + "/jsonrpc"
|
||||||
|
}
|
||||||
|
case "transmission":
|
||||||
|
if !strings.Contains(strings.ToLower(u.Path), "/rpc") {
|
||||||
|
u.Path = strings.TrimRight(u.Path, "/") + "/transmission/rpc"
|
||||||
|
}
|
||||||
|
case "qbittorrent":
|
||||||
|
default:
|
||||||
|
return "", fmt.Errorf("unsupported client type %q", clientType)
|
||||||
|
}
|
||||||
|
return u.String(), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func newDownloadClientHTTPRequest(ctx context.Context, method, endpoint string, body io.Reader) (*http.Request, error) {
|
||||||
|
endpoint = strings.TrimSpace(endpoint)
|
||||||
|
if !downloadClientEndpointPattern.MatchString(endpoint) {
|
||||||
|
return nil, errors.New("download client endpoint failed safety validation")
|
||||||
|
}
|
||||||
|
return http.NewRequestWithContext(ctx, method, endpoint, body)
|
||||||
|
}
|
||||||
@@ -9,7 +9,6 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -150,7 +149,14 @@ func (s *DownloadClientService) Test(ctx context.Context, id string) error {
|
|||||||
case "qbittorrent":
|
case "qbittorrent":
|
||||||
return qbitLogin(ctx, s.client, c.Host, c.Username, c.Password)
|
return qbitLogin(ctx, s.client, c.Host, c.Username, c.Password)
|
||||||
case "aria2", "transmission":
|
case "aria2", "transmission":
|
||||||
req, _ := http.NewRequestWithContext(ctx, http.MethodGet, c.Host, nil)
|
endpoint, err := downloadClientRPCURL(c.Type, c.Host)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodGet, endpoint, nil)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
resp, err := s.client.Do(req)
|
resp, err := s.client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -175,12 +181,19 @@ func (s *DownloadClientService) Aria2GlobalStats(ctx context.Context, clientID s
|
|||||||
if c == nil || c.Type != "aria2" {
|
if c == nil || c.Type != "aria2" {
|
||||||
return nil, errors.New("aria2 client not found")
|
return nil, errors.New("aria2 client not found")
|
||||||
}
|
}
|
||||||
|
endpoint, err := downloadClientRPCURL("aria2", c.Host)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
payload := fmt.Sprintf(
|
payload := fmt.Sprintf(
|
||||||
`{"jsonrpc":"2.0","id":"x","method":"aria2.getGlobalStat","params":["token:%s"]}`,
|
`{"jsonrpc":"2.0","id":"x","method":"aria2.getGlobalStat","params":["token:%s"]}`,
|
||||||
c.Password,
|
c.Password,
|
||||||
)
|
)
|
||||||
req, _ := http.NewRequestWithContext(ctx, http.MethodPost, c.Host,
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, endpoint,
|
||||||
strings.NewReader(payload))
|
strings.NewReader(payload))
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
req.Header.Set("Content-Type", "application/json")
|
req.Header.Set("Content-Type", "application/json")
|
||||||
resp, err := s.client.Do(req)
|
resp, err := s.client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -221,14 +234,11 @@ func normalizeDownloadClientInput(in DownloadClientInput) (DownloadClientInput,
|
|||||||
if !strings.Contains(in.Host, "://") {
|
if !strings.Contains(in.Host, "://") {
|
||||||
in.Host = "http://" + in.Host
|
in.Host = "http://" + in.Host
|
||||||
}
|
}
|
||||||
parsed, err := url.Parse(in.Host)
|
normalized, err := normalizeDownloadClientEndpoint(in.Type, in.Host)
|
||||||
if err != nil || parsed.Host == "" {
|
if err != nil {
|
||||||
return in, errors.New("host must be a valid http(s) URL")
|
return in, err
|
||||||
}
|
}
|
||||||
if parsed.Scheme != "http" && parsed.Scheme != "https" {
|
in.Host = normalized
|
||||||
return in, errors.New("host only supports http or https")
|
|
||||||
}
|
|
||||||
in.Host = strings.TrimRight(parsed.String(), "/")
|
|
||||||
return in, nil
|
return in, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
package service
|
package service
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/glebarez/sqlite"
|
"github.com/glebarez/sqlite"
|
||||||
@@ -81,6 +83,95 @@ func TestDownloadClientRejectsUnsupportedHostScheme(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDownloadClientRejectsUnsafeEndpointParts(t *testing.T) {
|
||||||
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := db.AutoMigrate(&model.DownloadClient{}, &model.Setting{}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
svc := NewDownloadClientService(zap.NewNop(), repository.New(db))
|
||||||
|
|
||||||
|
for _, host := range []string{
|
||||||
|
"http://user:pass@127.0.0.1:6800",
|
||||||
|
"http://127.0.0.1:6800/jsonrpc?target=http://169.254.169.254",
|
||||||
|
"http://127.0.0.1:6800/jsonrpc#fragment",
|
||||||
|
"http://127.0.0.1:70000",
|
||||||
|
"file:///etc/passwd",
|
||||||
|
} {
|
||||||
|
if _, err := svc.Create(t.Context(), DownloadClientInput{
|
||||||
|
Name: "bad",
|
||||||
|
Type: "aria2",
|
||||||
|
Host: host,
|
||||||
|
Enabled: true,
|
||||||
|
}); err == nil {
|
||||||
|
t.Fatalf("Create allowed unsafe host %q", host)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDownloadClientRPCURLAppendsExpectedPath(t *testing.T) {
|
||||||
|
cases := []struct {
|
||||||
|
clientType string
|
||||||
|
host string
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"aria2", "127.0.0.1:6800", "http://127.0.0.1:6800/jsonrpc"},
|
||||||
|
{"aria2", "http://nas.local:6800/rpc", "http://nas.local:6800/rpc/jsonrpc"},
|
||||||
|
{"transmission", "http://nas.local:9091", "http://nas.local:9091/transmission/rpc"},
|
||||||
|
{"transmission", "http://nas.local:9091/transmission/rpc", "http://nas.local:9091/transmission/rpc"},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
got, err := downloadClientRPCURL(tc.clientType, tc.host)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("downloadClientRPCURL(%q, %q) error: %v", tc.clientType, tc.host, err)
|
||||||
|
}
|
||||||
|
if got != tc.want {
|
||||||
|
t.Fatalf("downloadClientRPCURL(%q, %q) = %q, want %q", tc.clientType, tc.host, got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAria2AdapterRejectsUnsafeHostBeforeHTTPRequest(t *testing.T) {
|
||||||
|
adapter := NewAria2Adapter()
|
||||||
|
called := false
|
||||||
|
adapter.client = &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||||
|
called = true
|
||||||
|
return &http.Response{StatusCode: http.StatusOK, Body: http.NoBody}, nil
|
||||||
|
})}
|
||||||
|
|
||||||
|
if err := adapter.Initialize(t.Context(), DownloadClientConfig{
|
||||||
|
Host: "http://user:pass@127.0.0.1:6800",
|
||||||
|
Password: "secret",
|
||||||
|
}); err == nil {
|
||||||
|
t.Fatal("expected unsafe host error")
|
||||||
|
}
|
||||||
|
if called {
|
||||||
|
t.Fatal("unsafe aria2 host should be rejected before any HTTP request")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAria2AdapterUsesNormalizedRPCURL(t *testing.T) {
|
||||||
|
var gotPath string
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
gotPath = r.URL.Path
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
adapter := NewAria2Adapter()
|
||||||
|
if err := adapter.Initialize(t.Context(), DownloadClientConfig{
|
||||||
|
Host: server.URL,
|
||||||
|
Password: "secret",
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if gotPath != "/jsonrpc" {
|
||||||
|
t.Fatalf("aria2 request path = %q, want /jsonrpc", gotPath)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestDownloadClientDeleteClearsLegacyQBitConnectionWhenNoDefault(t *testing.T) {
|
func TestDownloadClientDeleteClearsLegacyQBitConnectionWhenNoDefault(t *testing.T) {
|
||||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -22,9 +22,9 @@ type DownloadManager struct {
|
|||||||
repo *repository.Container
|
repo *repository.Container
|
||||||
crypto *CryptoService
|
crypto *CryptoService
|
||||||
|
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
clients map[string]DownloadAdapter // clientID -> adapter
|
clients map[string]DownloadAdapter // clientID -> adapter
|
||||||
configs map[string]DownloadClientConfig
|
configs map[string]DownloadClientConfig
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewDownloadManager 创建新的下载管理器。
|
// NewDownloadManager 创建新的下载管理器。
|
||||||
@@ -230,9 +230,13 @@ func (m *DownloadManager) buildConfig(dc *model.DownloadClient) (DownloadClientC
|
|||||||
if m.crypto != nil && password != "" {
|
if m.crypto != nil && password != "" {
|
||||||
password = m.crypto.Decrypt(password)
|
password = m.crypto.Decrypt(password)
|
||||||
}
|
}
|
||||||
|
host, err := normalizeDownloadClientEndpoint(dc.Type, dc.Host)
|
||||||
|
if err != nil {
|
||||||
|
return DownloadClientConfig{}, err
|
||||||
|
}
|
||||||
|
|
||||||
cfg := DownloadClientConfig{
|
cfg := DownloadClientConfig{
|
||||||
Host: dc.Host,
|
Host: host,
|
||||||
Username: dc.Username,
|
Username: dc.Username,
|
||||||
Password: password,
|
Password: password,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -78,7 +78,7 @@ var (
|
|||||||
// an unconfigured downloader must fail closed instead of silently trying a
|
// an unconfigured downloader must fail closed instead of silently trying a
|
||||||
// localhost qBittorrent instance.
|
// localhost qBittorrent instance.
|
||||||
func NewQBitClient(log *zap.Logger, cfg QBitConfig) *QBitClient {
|
func NewQBitClient(log *zap.Logger, cfg QBitConfig) *QBitClient {
|
||||||
cfg.BaseURL = strings.TrimRight(strings.TrimSpace(cfg.BaseURL), "/")
|
cfg.BaseURL = normalizeQBitBaseURL(cfg.BaseURL)
|
||||||
jar, _ := cookiejar.New(nil)
|
jar, _ := cookiejar.New(nil)
|
||||||
client := NewInternalHTTPClient(20 * time.Second)
|
client := NewInternalHTTPClient(20 * time.Second)
|
||||||
client.Jar = jar
|
client.Jar = jar
|
||||||
@@ -93,7 +93,7 @@ func NewQBitClient(log *zap.Logger, cfg QBitConfig) *QBitClient {
|
|||||||
func (q *QBitClient) Configure(cfg QBitConfig) {
|
func (q *QBitClient) Configure(cfg QBitConfig) {
|
||||||
q.mu.Lock()
|
q.mu.Lock()
|
||||||
defer q.mu.Unlock()
|
defer q.mu.Unlock()
|
||||||
cfg.BaseURL = strings.TrimRight(strings.TrimSpace(cfg.BaseURL), "/")
|
cfg.BaseURL = normalizeQBitBaseURL(cfg.BaseURL)
|
||||||
q.cfg = cfg
|
q.cfg = cfg
|
||||||
jar, _ := cookiejar.New(nil)
|
jar, _ := cookiejar.New(nil)
|
||||||
q.client.Jar = jar
|
q.client.Jar = jar
|
||||||
@@ -105,6 +105,18 @@ func (q *QBitClient) IsConfigured() bool {
|
|||||||
return strings.TrimSpace(q.cfg.BaseURL) != ""
|
return strings.TrimSpace(q.cfg.BaseURL) != ""
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func normalizeQBitBaseURL(raw string) string {
|
||||||
|
raw = strings.TrimSpace(raw)
|
||||||
|
if raw == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
normalized, err := normalizeDownloadClientEndpoint("qbittorrent", raw)
|
||||||
|
if err != nil {
|
||||||
|
return strings.TrimRight(raw, "/")
|
||||||
|
}
|
||||||
|
return normalized
|
||||||
|
}
|
||||||
|
|
||||||
// Login performs POST /api/v2/auth/login.
|
// Login performs POST /api/v2/auth/login.
|
||||||
func (q *QBitClient) Login(ctx context.Context) error {
|
func (q *QBitClient) Login(ctx context.Context) error {
|
||||||
if q.cfg.BaseURL == "" {
|
if q.cfg.BaseURL == "" {
|
||||||
@@ -193,9 +205,8 @@ func (q *QBitClient) addTorrentLocked(ctx context.Context, magnetOrURL string, t
|
|||||||
}
|
}
|
||||||
_ = w.Close()
|
_ = w.Close()
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost,
|
||||||
strings.TrimRight(q.cfg.BaseURL, "/")+"/api/v2/torrents/add", body,
|
strings.TrimRight(q.cfg.BaseURL, "/")+"/api/v2/torrents/add", body)
|
||||||
)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -427,13 +438,15 @@ func (q *QBitClient) List(ctx context.Context, filter string) ([]QBitTorrent, er
|
|||||||
|
|
||||||
func (q *QBitClient) listLocked(ctx context.Context, filter string) ([]QBitTorrent, error) {
|
func (q *QBitClient) listLocked(ctx context.Context, filter string) ([]QBitTorrent, error) {
|
||||||
u := strings.TrimRight(q.cfg.BaseURL, "/") + "/api/v2/torrents/info"
|
u := strings.TrimRight(q.cfg.BaseURL, "/") + "/api/v2/torrents/info"
|
||||||
if filter != "" {
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodGet, u, nil)
|
||||||
u += "?filter=" + url.QueryEscape(filter)
|
|
||||||
}
|
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
if filter != "" {
|
||||||
|
query := req.URL.Query()
|
||||||
|
query.Set("filter", filter)
|
||||||
|
req.URL.RawQuery = query.Encode()
|
||||||
|
}
|
||||||
req.Header.Set("Referer", q.cfg.BaseURL)
|
req.Header.Set("Referer", q.cfg.BaseURL)
|
||||||
resp, err := q.client.Do(req)
|
resp, err := q.client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -464,10 +477,9 @@ func (q *QBitClient) Delete(ctx context.Context, hash string, deleteFiles bool)
|
|||||||
} else {
|
} else {
|
||||||
form.Set("deleteFiles", "false")
|
form.Set("deleteFiles", "false")
|
||||||
}
|
}
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost,
|
||||||
strings.TrimRight(q.cfg.BaseURL, "/")+"/api/v2/torrents/delete",
|
strings.TrimRight(q.cfg.BaseURL, "/")+"/api/v2/torrents/delete",
|
||||||
strings.NewReader(form.Encode()),
|
strings.NewReader(form.Encode()))
|
||||||
)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -504,10 +516,9 @@ func (q *QBitClient) SetLocation(ctx context.Context, hash, location string) err
|
|||||||
form := url.Values{}
|
form := url.Values{}
|
||||||
form.Set("hashes", hash)
|
form.Set("hashes", hash)
|
||||||
form.Set("location", location)
|
form.Set("location", location)
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost,
|
||||||
strings.TrimRight(q.cfg.BaseURL, "/")+"/api/v2/torrents/setLocation",
|
strings.TrimRight(q.cfg.BaseURL, "/")+"/api/v2/torrents/setLocation",
|
||||||
strings.NewReader(form.Encode()),
|
strings.NewReader(form.Encode()))
|
||||||
)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -42,6 +42,11 @@ func NewQBitAdapter() *QBitAdapter {
|
|||||||
func (a *QBitAdapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error {
|
func (a *QBitAdapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error {
|
||||||
a.mu.Lock()
|
a.mu.Lock()
|
||||||
defer a.mu.Unlock()
|
defer a.mu.Unlock()
|
||||||
|
endpoint, err := normalizeDownloadClientEndpoint("qbittorrent", cfg.Host)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
cfg.Host = endpoint
|
||||||
a.cfg = cfg
|
a.cfg = cfg
|
||||||
a.LoggedIn = false
|
a.LoggedIn = false
|
||||||
jar, _ := cookiejar.New(nil)
|
jar, _ := cookiejar.New(nil)
|
||||||
@@ -73,7 +78,7 @@ func (a *QBitAdapter) AddTorrent(ctx context.Context, torrentURL, savePath strin
|
|||||||
_ = w.Close()
|
_ = w.Close()
|
||||||
|
|
||||||
baseURL := strings.TrimRight(a.cfg.Host, "/")
|
baseURL := strings.TrimRight(a.cfg.Host, "/")
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost,
|
||||||
baseURL+"/api/v2/torrents/add", body)
|
baseURL+"/api/v2/torrents/add", body)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
@@ -108,7 +113,7 @@ func (a *QBitAdapter) Pause(ctx context.Context, hash string) error {
|
|||||||
baseURL := strings.TrimRight(a.cfg.Host, "/")
|
baseURL := strings.TrimRight(a.cfg.Host, "/")
|
||||||
form := url.Values{}
|
form := url.Values{}
|
||||||
form.Set("hashes", hash)
|
form.Set("hashes", hash)
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost,
|
||||||
baseURL+"/api/v2/torrents/pause", strings.NewReader(form.Encode()))
|
baseURL+"/api/v2/torrents/pause", strings.NewReader(form.Encode()))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -136,7 +141,7 @@ func (a *QBitAdapter) Resume(ctx context.Context, hash string) error {
|
|||||||
baseURL := strings.TrimRight(a.cfg.Host, "/")
|
baseURL := strings.TrimRight(a.cfg.Host, "/")
|
||||||
form := url.Values{}
|
form := url.Values{}
|
||||||
form.Set("hashes", hash)
|
form.Set("hashes", hash)
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost,
|
||||||
baseURL+"/api/v2/torrents/resume", strings.NewReader(form.Encode()))
|
baseURL+"/api/v2/torrents/resume", strings.NewReader(form.Encode()))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -169,7 +174,7 @@ func (a *QBitAdapter) Remove(ctx context.Context, hash string, deleteFiles bool)
|
|||||||
} else {
|
} else {
|
||||||
form.Set("deleteFiles", "false")
|
form.Set("deleteFiles", "false")
|
||||||
}
|
}
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost,
|
||||||
baseURL+"/api/v2/torrents/delete", strings.NewReader(form.Encode()))
|
baseURL+"/api/v2/torrents/delete", strings.NewReader(form.Encode()))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -196,13 +201,15 @@ func (a *QBitAdapter) List(ctx context.Context, filter string) ([]TorrentInfo, e
|
|||||||
}
|
}
|
||||||
baseURL := strings.TrimRight(a.cfg.Host, "/")
|
baseURL := strings.TrimRight(a.cfg.Host, "/")
|
||||||
u := baseURL + "/api/v2/torrents/info"
|
u := baseURL + "/api/v2/torrents/info"
|
||||||
if filter != "" {
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodGet, u, nil)
|
||||||
u += "?filter=" + url.QueryEscape(filter)
|
|
||||||
}
|
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
if filter != "" {
|
||||||
|
query := req.URL.Query()
|
||||||
|
query.Set("filter", filter)
|
||||||
|
req.URL.RawQuery = query.Encode()
|
||||||
|
}
|
||||||
req.Header.Set("Referer", baseURL)
|
req.Header.Set("Referer", baseURL)
|
||||||
resp, err := a.client.Do(req)
|
resp, err := a.client.Do(req)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -18,10 +18,15 @@ type qbitLoginVariant struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func qbitLogin(ctx context.Context, client *http.Client, baseURL, username, password string) error {
|
func qbitLogin(ctx context.Context, client *http.Client, baseURL, username, password string) error {
|
||||||
baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
|
baseURL = strings.TrimSpace(baseURL)
|
||||||
if baseURL == "" {
|
if baseURL == "" {
|
||||||
return errors.New("qbittorrent host not configured")
|
return errors.New("qbittorrent host not configured")
|
||||||
}
|
}
|
||||||
|
var err error
|
||||||
|
baseURL, err = normalizeDownloadClientEndpoint("qbittorrent", baseURL)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
var lastErr error
|
var lastErr error
|
||||||
for _, variant := range []qbitLoginVariant{
|
for _, variant := range []qbitLoginVariant{
|
||||||
{name: "minimal"},
|
{name: "minimal"},
|
||||||
@@ -46,7 +51,7 @@ func qbitLoginOnce(ctx context.Context, client *http.Client, baseURL, username,
|
|||||||
form := url.Values{}
|
form := url.Values{}
|
||||||
form.Set("username", username)
|
form.Set("username", username)
|
||||||
form.Set("password", password)
|
form.Set("password", password)
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost,
|
||||||
baseURL+"/api/v2/auth/login", strings.NewReader(form.Encode()))
|
baseURL+"/api/v2/auth/login", strings.NewReader(form.Encode()))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|||||||
@@ -51,6 +51,11 @@ func NewTransmissionAdapter() *TransmissionAdapter {
|
|||||||
func (a *TransmissionAdapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error {
|
func (a *TransmissionAdapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error {
|
||||||
a.mu.Lock()
|
a.mu.Lock()
|
||||||
defer a.mu.Unlock()
|
defer a.mu.Unlock()
|
||||||
|
endpoint, err := normalizeDownloadClientEndpoint("transmission", cfg.Host)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
cfg.Host = endpoint
|
||||||
a.cfg = cfg
|
a.cfg = cfg
|
||||||
a.sessionID = ""
|
a.sessionID = ""
|
||||||
a.tag = 0
|
a.tag = 0
|
||||||
@@ -66,13 +71,11 @@ func (a *TransmissionAdapter) Ping(ctx context.Context) error {
|
|||||||
|
|
||||||
// pingLocked 内部 ping 实现(调用者必须持有锁)。
|
// pingLocked 内部 ping 实现(调用者必须持有锁)。
|
||||||
func (a *TransmissionAdapter) pingLocked(ctx context.Context) error {
|
func (a *TransmissionAdapter) pingLocked(ctx context.Context) error {
|
||||||
rpcURL := a.cfg.Host
|
rpcURL, err := downloadClientRPCURL("transmission", a.cfg.Host)
|
||||||
if !strings.HasSuffix(rpcURL, "/rpc") && !strings.HasSuffix(rpcURL, "/transmission/rpc") {
|
if err != nil {
|
||||||
if !strings.Contains(rpcURL, "/rpc") {
|
return err
|
||||||
rpcURL = strings.TrimRight(rpcURL, "/") + "/transmission/rpc"
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, rpcURL, nil)
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodGet, rpcURL, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -98,9 +101,9 @@ func (a *TransmissionAdapter) pingLocked(ctx context.Context) error {
|
|||||||
|
|
||||||
// rpcLocked 发送 RPC 请求(调用者必须持有锁)。
|
// rpcLocked 发送 RPC 请求(调用者必须持有锁)。
|
||||||
func (a *TransmissionAdapter) rpcLocked(ctx context.Context, method string, args map[string]interface{}) (*transmissionRPCResponse, error) {
|
func (a *TransmissionAdapter) rpcLocked(ctx context.Context, method string, args map[string]interface{}) (*transmissionRPCResponse, error) {
|
||||||
rpcURL := a.cfg.Host
|
rpcURL, err := downloadClientRPCURL("transmission", a.cfg.Host)
|
||||||
if !strings.Contains(rpcURL, "/rpc") {
|
if err != nil {
|
||||||
rpcURL = strings.TrimRight(rpcURL, "/") + "/transmission/rpc"
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
a.tag++
|
a.tag++
|
||||||
@@ -114,7 +117,7 @@ func (a *TransmissionAdapter) rpcLocked(ctx context.Context, method string, args
|
|||||||
}
|
}
|
||||||
|
|
||||||
for attempt := 0; attempt < 2; attempt++ {
|
for attempt := 0; attempt < 2; attempt++ {
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
|
req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, rpcURL, bytes.NewReader(body))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user