diff --git a/internal/handler/download_client_handler.go b/internal/handler/download_client_handler.go index c7b11b6..a2d7c7b 100644 --- a/internal/handler/download_client_handler.go +++ b/internal/handler/download_client_handler.go @@ -57,6 +57,11 @@ func (h *DownloadClientHandler) Create(c *gin.Context) { ctx := c.Request.Context() _ = 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 @@ -82,7 +87,7 @@ func (h *DownloadClientHandler) Create(c *gin.Context) { client := &model.DownloadClient{ Name: req.Name, Type: req.Type, - Host: req.Host, + Host: normalizedHost, Username: req.Username, Password: password, IsDefault: req.IsDefault, @@ -152,7 +157,16 @@ func (h *DownloadClientHandler) Update(c *gin.Context) { client.Type = req.Type } 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 != "" { client.Username = req.Username @@ -180,6 +194,12 @@ func (h *DownloadClientHandler) Update(c *gin.Context) { 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 { Error(c, http.StatusInternalServerError, ErrInternal, "更新失败") diff --git a/internal/service/aria2_adp.go b/internal/service/aria2_adp.go index 15708bf..71b6fb7 100644 --- a/internal/service/aria2_adp.go +++ b/internal/service/aria2_adp.go @@ -57,6 +57,11 @@ func NewAria2Adapter() *Aria2Adapter { func (a *Aria2Adapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error { a.mu.Lock() defer a.mu.Unlock() + endpoint, err := normalizeDownloadClientEndpoint("aria2", cfg.Host) + if err != nil { + return err + } + cfg.Host = endpoint a.cfg = cfg a.idSeq = 0 return a.getVersionLocked(ctx) @@ -71,9 +76,9 @@ func (a *Aria2Adapter) Ping(ctx context.Context) error { // getVersionLocked 内部版本检查(调用者必须持有锁)。 func (a *Aria2Adapter) getVersionLocked(ctx context.Context) error { - rpcURL := a.cfg.Host - if !strings.HasSuffix(rpcURL, "/jsonrpc") { - rpcURL = strings.TrimRight(rpcURL, "/") + "/jsonrpc" + rpcURL, err := downloadClientRPCURL("aria2", a.cfg.Host) + if err != nil { + return err } req := &aria2Request{ @@ -88,7 +93,7 @@ func (a *Aria2Adapter) getVersionLocked(ctx context.Context) error { 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 { return err } @@ -110,9 +115,9 @@ func (a *Aria2Adapter) getVersionLocked(ctx context.Context) error { // rpcLocked 发送 JSON-RPC 请求(调用者必须持有锁)。 func (a *Aria2Adapter) rpcLocked(ctx context.Context, method string, params []interface{}) (json.RawMessage, error) { - rpcURL := a.cfg.Host - if !strings.HasSuffix(rpcURL, "/jsonrpc") { - rpcURL = strings.TrimRight(rpcURL, "/") + "/jsonrpc" + rpcURL, err := downloadClientRPCURL("aria2", a.cfg.Host) + if err != nil { + return nil, err } if params == nil { @@ -145,7 +150,7 @@ func (a *Aria2Adapter) rpcLocked(ctx context.Context, method string, params []in 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 { return nil, err } diff --git a/internal/service/download_client_endpoint.go b/internal/service/download_client_endpoint.go new file mode 100644 index 0000000..58a3e0c --- /dev/null +++ b/internal/service/download_client_endpoint.go @@ -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) +} diff --git a/internal/service/download_clients.go b/internal/service/download_clients.go index 2971bb9..0d342bd 100644 --- a/internal/service/download_clients.go +++ b/internal/service/download_clients.go @@ -9,7 +9,6 @@ import ( "errors" "fmt" "net/http" - "net/url" "strings" "time" @@ -150,7 +149,14 @@ func (s *DownloadClientService) Test(ctx context.Context, id string) error { case "qbittorrent": return qbitLogin(ctx, s.client, c.Host, c.Username, c.Password) 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) if err != nil { return err @@ -175,12 +181,19 @@ func (s *DownloadClientService) Aria2GlobalStats(ctx context.Context, clientID s if c == nil || c.Type != "aria2" { return nil, errors.New("aria2 client not found") } + endpoint, err := downloadClientRPCURL("aria2", c.Host) + if err != nil { + return nil, err + } payload := fmt.Sprintf( `{"jsonrpc":"2.0","id":"x","method":"aria2.getGlobalStat","params":["token:%s"]}`, c.Password, ) - req, _ := http.NewRequestWithContext(ctx, http.MethodPost, c.Host, + req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, endpoint, strings.NewReader(payload)) + if err != nil { + return nil, err + } req.Header.Set("Content-Type", "application/json") resp, err := s.client.Do(req) if err != nil { @@ -221,14 +234,11 @@ func normalizeDownloadClientInput(in DownloadClientInput) (DownloadClientInput, if !strings.Contains(in.Host, "://") { in.Host = "http://" + in.Host } - parsed, err := url.Parse(in.Host) - if err != nil || parsed.Host == "" { - return in, errors.New("host must be a valid http(s) URL") + normalized, err := normalizeDownloadClientEndpoint(in.Type, in.Host) + if err != nil { + return in, err } - if parsed.Scheme != "http" && parsed.Scheme != "https" { - return in, errors.New("host only supports http or https") - } - in.Host = strings.TrimRight(parsed.String(), "/") + in.Host = normalized return in, nil } diff --git a/internal/service/download_clients_test.go b/internal/service/download_clients_test.go index 6c720f0..6db5d1c 100644 --- a/internal/service/download_clients_test.go +++ b/internal/service/download_clients_test.go @@ -1,6 +1,8 @@ package service import ( + "net/http" + "net/http/httptest" "testing" "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) { db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) if err != nil { diff --git a/internal/service/download_manager_svc.go b/internal/service/download_manager_svc.go index 4cb9511..8c86812 100644 --- a/internal/service/download_manager_svc.go +++ b/internal/service/download_manager_svc.go @@ -22,9 +22,9 @@ type DownloadManager struct { repo *repository.Container crypto *CryptoService - mu sync.RWMutex - clients map[string]DownloadAdapter // clientID -> adapter - configs map[string]DownloadClientConfig + mu sync.RWMutex + clients map[string]DownloadAdapter // clientID -> adapter + configs map[string]DownloadClientConfig } // NewDownloadManager 创建新的下载管理器。 @@ -230,9 +230,13 @@ func (m *DownloadManager) buildConfig(dc *model.DownloadClient) (DownloadClientC if m.crypto != nil && password != "" { password = m.crypto.Decrypt(password) } + host, err := normalizeDownloadClientEndpoint(dc.Type, dc.Host) + if err != nil { + return DownloadClientConfig{}, err + } cfg := DownloadClientConfig{ - Host: dc.Host, + Host: host, Username: dc.Username, Password: password, } diff --git a/internal/service/qbittorrent.go b/internal/service/qbittorrent.go index e38b5e4..ce2e8c3 100644 --- a/internal/service/qbittorrent.go +++ b/internal/service/qbittorrent.go @@ -78,7 +78,7 @@ var ( // an unconfigured downloader must fail closed instead of silently trying a // localhost qBittorrent instance. 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) client := NewInternalHTTPClient(20 * time.Second) client.Jar = jar @@ -93,7 +93,7 @@ func NewQBitClient(log *zap.Logger, cfg QBitConfig) *QBitClient { func (q *QBitClient) Configure(cfg QBitConfig) { q.mu.Lock() defer q.mu.Unlock() - cfg.BaseURL = strings.TrimRight(strings.TrimSpace(cfg.BaseURL), "/") + cfg.BaseURL = normalizeQBitBaseURL(cfg.BaseURL) q.cfg = cfg jar, _ := cookiejar.New(nil) q.client.Jar = jar @@ -105,6 +105,18 @@ func (q *QBitClient) IsConfigured() bool { 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. func (q *QBitClient) Login(ctx context.Context) error { if q.cfg.BaseURL == "" { @@ -193,9 +205,8 @@ func (q *QBitClient) addTorrentLocked(ctx context.Context, magnetOrURL string, t } _ = w.Close() - req, err := http.NewRequestWithContext(ctx, http.MethodPost, - strings.TrimRight(q.cfg.BaseURL, "/")+"/api/v2/torrents/add", body, - ) + req, err := newDownloadClientHTTPRequest(ctx, http.MethodPost, + strings.TrimRight(q.cfg.BaseURL, "/")+"/api/v2/torrents/add", body) if err != nil { 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) { u := strings.TrimRight(q.cfg.BaseURL, "/") + "/api/v2/torrents/info" - if filter != "" { - u += "?filter=" + url.QueryEscape(filter) - } - req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil) + req, err := newDownloadClientHTTPRequest(ctx, http.MethodGet, u, nil) if err != nil { 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) resp, err := q.client.Do(req) if err != nil { @@ -464,10 +477,9 @@ func (q *QBitClient) Delete(ctx context.Context, hash string, deleteFiles bool) } else { 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.NewReader(form.Encode()), - ) + strings.NewReader(form.Encode())) if err != nil { return err } @@ -504,10 +516,9 @@ func (q *QBitClient) SetLocation(ctx context.Context, hash, location string) err form := url.Values{} form.Set("hashes", hash) 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.NewReader(form.Encode()), - ) + strings.NewReader(form.Encode())) if err != nil { return err } diff --git a/internal/service/qbittorrent_adp.go b/internal/service/qbittorrent_adp.go index e108c33..779e2d2 100644 --- a/internal/service/qbittorrent_adp.go +++ b/internal/service/qbittorrent_adp.go @@ -42,6 +42,11 @@ func NewQBitAdapter() *QBitAdapter { func (a *QBitAdapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error { a.mu.Lock() defer a.mu.Unlock() + endpoint, err := normalizeDownloadClientEndpoint("qbittorrent", cfg.Host) + if err != nil { + return err + } + cfg.Host = endpoint a.cfg = cfg a.LoggedIn = false jar, _ := cookiejar.New(nil) @@ -73,7 +78,7 @@ func (a *QBitAdapter) AddTorrent(ctx context.Context, torrentURL, savePath strin _ = w.Close() 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) if err != nil { return "", err @@ -108,7 +113,7 @@ func (a *QBitAdapter) Pause(ctx context.Context, hash string) error { baseURL := strings.TrimRight(a.cfg.Host, "/") form := url.Values{} 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())) if err != nil { return err @@ -136,7 +141,7 @@ func (a *QBitAdapter) Resume(ctx context.Context, hash string) error { baseURL := strings.TrimRight(a.cfg.Host, "/") form := url.Values{} 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())) if err != nil { return err @@ -169,7 +174,7 @@ func (a *QBitAdapter) Remove(ctx context.Context, hash string, deleteFiles bool) } else { 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())) if err != nil { return err @@ -196,13 +201,15 @@ func (a *QBitAdapter) List(ctx context.Context, filter string) ([]TorrentInfo, e } baseURL := strings.TrimRight(a.cfg.Host, "/") u := baseURL + "/api/v2/torrents/info" - if filter != "" { - u += "?filter=" + url.QueryEscape(filter) - } - req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil) + req, err := newDownloadClientHTTPRequest(ctx, http.MethodGet, u, nil) if err != nil { return nil, err } + if filter != "" { + query := req.URL.Query() + query.Set("filter", filter) + req.URL.RawQuery = query.Encode() + } req.Header.Set("Referer", baseURL) resp, err := a.client.Do(req) if err != nil { diff --git a/internal/service/qbittorrent_login.go b/internal/service/qbittorrent_login.go index 8269dcc..f02d3fc 100644 --- a/internal/service/qbittorrent_login.go +++ b/internal/service/qbittorrent_login.go @@ -18,10 +18,15 @@ type qbitLoginVariant struct { } 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 == "" { return errors.New("qbittorrent host not configured") } + var err error + baseURL, err = normalizeDownloadClientEndpoint("qbittorrent", baseURL) + if err != nil { + return err + } var lastErr error for _, variant := range []qbitLoginVariant{ {name: "minimal"}, @@ -46,7 +51,7 @@ func qbitLoginOnce(ctx context.Context, client *http.Client, baseURL, username, form := url.Values{} form.Set("username", username) 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())) if err != nil { return err diff --git a/internal/service/transmission_adp.go b/internal/service/transmission_adp.go index 7d8434f..e6023dc 100644 --- a/internal/service/transmission_adp.go +++ b/internal/service/transmission_adp.go @@ -51,6 +51,11 @@ func NewTransmissionAdapter() *TransmissionAdapter { func (a *TransmissionAdapter) Initialize(ctx context.Context, cfg DownloadClientConfig) error { a.mu.Lock() defer a.mu.Unlock() + endpoint, err := normalizeDownloadClientEndpoint("transmission", cfg.Host) + if err != nil { + return err + } + cfg.Host = endpoint a.cfg = cfg a.sessionID = "" a.tag = 0 @@ -66,13 +71,11 @@ func (a *TransmissionAdapter) Ping(ctx context.Context) error { // pingLocked 内部 ping 实现(调用者必须持有锁)。 func (a *TransmissionAdapter) pingLocked(ctx context.Context) error { - rpcURL := a.cfg.Host - if !strings.HasSuffix(rpcURL, "/rpc") && !strings.HasSuffix(rpcURL, "/transmission/rpc") { - if !strings.Contains(rpcURL, "/rpc") { - rpcURL = strings.TrimRight(rpcURL, "/") + "/transmission/rpc" - } + rpcURL, err := downloadClientRPCURL("transmission", a.cfg.Host) + if err != nil { + return err } - req, err := http.NewRequestWithContext(ctx, http.MethodGet, rpcURL, nil) + req, err := newDownloadClientHTTPRequest(ctx, http.MethodGet, rpcURL, nil) if err != nil { return err } @@ -98,9 +101,9 @@ func (a *TransmissionAdapter) pingLocked(ctx context.Context) error { // rpcLocked 发送 RPC 请求(调用者必须持有锁)。 func (a *TransmissionAdapter) rpcLocked(ctx context.Context, method string, args map[string]interface{}) (*transmissionRPCResponse, error) { - rpcURL := a.cfg.Host - if !strings.Contains(rpcURL, "/rpc") { - rpcURL = strings.TrimRight(rpcURL, "/") + "/transmission/rpc" + rpcURL, err := downloadClientRPCURL("transmission", a.cfg.Host) + if err != nil { + return nil, err } a.tag++ @@ -114,7 +117,7 @@ func (a *TransmissionAdapter) rpcLocked(ctx context.Context, method string, args } 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 { return nil, err }