diff --git a/model/notification.go b/model/notification.go index a2304cf0..59ebb109 100644 --- a/model/notification.go +++ b/model/notification.go @@ -1,14 +1,10 @@ package model import ( - "context" - "crypto/tls" "errors" "fmt" "io" - "net" "net/http" - "net/netip" "net/url" "strings" "time" @@ -29,35 +25,6 @@ const ( NotificationRequestMethodPOST ) -var errNotificationURLNotAllowed = errors.New("notification URL target is not allowed") - -var notificationBlockedCIDRs = mustParseNotificationCIDRs([]string{ - "0.0.0.0/8", - "10.0.0.0/8", - "100.64.0.0/10", - "127.0.0.0/8", - "169.254.0.0/16", - "172.16.0.0/12", - "192.0.0.0/24", - "192.0.2.0/24", - "192.168.0.0/16", - "198.18.0.0/15", - "198.51.100.0/24", - "203.0.113.0/24", - "224.0.0.0/4", - "240.0.0.0/4", - "::/128", - "::1/128", - "::ffff:0:0/96", - "64:ff9b::/96", - "100::/64", - "2001::/23", - "2001:db8::/32", - "fc00::/7", - "fe80::/10", - "ff00::/8", -}) - type NotificationServerBundle struct { Notification *Notification Server *Server @@ -191,105 +158,13 @@ func (ns *NotificationServerBundle) Send(message string) error { return nil } -func newNotificationHTTPClient(rawURL string, verifyTLS bool) (*http.Client, error) { - parsedURL, ip, err := resolveNotificationTarget(rawURL) - if err != nil { - return nil, err - } - - port := parsedURL.Port() - if port == "" { - if parsedURL.Scheme == "https" { - port = "443" - } else { - port = "80" - } - } - targetAddress := net.JoinHostPort(ip.String(), port) - dialer := &net.Dialer{} - - return &http.Client{ - Transport: &http.Transport{ - DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { - return dialer.DialContext(ctx, network, targetAddress) - }, - TLSClientConfig: &tls.Config{InsecureSkipVerify: !verifyTLS, ServerName: parsedURL.Hostname()}, - }, - CheckRedirect: func(req *http.Request, via []*http.Request) error { - return http.ErrUseLastResponse - }, - Timeout: time.Minute * 10, - }, nil -} - func notificationResponseError(resp *http.Response) error { _, _ = io.CopyN(io.Discard, resp.Body, 4096) return fmt.Errorf("%d@%s", resp.StatusCode, resp.Status) } -func resolveNotificationTarget(rawURL string) (*url.URL, net.IP, error) { - parsedURL, err := url.Parse(rawURL) - if err != nil { - return nil, nil, err - } - if parsedURL.Scheme != "http" && parsedURL.Scheme != "https" { - return nil, nil, errNotificationURLNotAllowed - } - - host := parsedURL.Hostname() - if host == "" { - return nil, nil, errNotificationURLNotAllowed - } - if ip := net.ParseIP(host); ip != nil { - if !notificationIPAllowed(ip) { - return nil, nil, errNotificationURLNotAllowed - } - return parsedURL, ip, nil - } - - ips, err := net.LookupIP(host) - if err != nil { - return nil, nil, err - } - if len(ips) == 0 { - return nil, nil, errNotificationURLNotAllowed - } - for _, ip := range ips { - if !notificationIPAllowed(ip) { - return nil, nil, errNotificationURLNotAllowed - } - } - - return parsedURL, ips[0], nil -} - -func notificationIPAllowed(ip net.IP) bool { - parsedIP, ok := netipFromIP(ip) - if !ok { - return false - } - for _, cidr := range notificationBlockedCIDRs { - if cidr.Contains(parsedIP) { - return false - } - } - return parsedIP.IsGlobalUnicast() -} - -func netipFromIP(ip net.IP) (netip.Addr, bool) { - parsedIP, ok := netip.AddrFromSlice(ip) - if !ok { - return netip.Addr{}, false - } - return parsedIP.Unmap(), true -} - -func mustParseNotificationCIDRs(cidrs []string) []netip.Prefix { - prefixes := make([]netip.Prefix, 0, len(cidrs)) - for _, cidr := range cidrs { - prefixes = append(prefixes, netip.MustParsePrefix(cidr)) - } - return prefixes +func newNotificationHTTPClient(rawURL string, verifyTLS bool) (*http.Client, error) { + return utils.NewRestrictedHTTPClient(rawURL, !verifyTLS) } // replaceParamInString 替换字符串中的占位符 diff --git a/model/notification_test.go b/model/notification_test.go index 28dd4873..25e6599a 100644 --- a/model/notification_test.go +++ b/model/notification_test.go @@ -6,6 +6,8 @@ import ( "strings" "testing" "time" + + "github.com/nezhahq/nezha/pkg/utils" ) var ( @@ -307,7 +309,7 @@ func TestNotificationTargetRejectsBlockedRanges(t *testing.T) { for _, rawURL := range cases { t.Run(rawURL, func(t *testing.T) { - if _, _, err := resolveNotificationTarget(rawURL); err == nil { + if _, _, err := utils.ResolveAllowedHTTPURL(rawURL); err == nil { t.Fatalf("expected %s to be rejected", rawURL) } }) @@ -323,7 +325,7 @@ func TestNotificationTargetAllowsPublicAddresses(t *testing.T) { for _, rawURL := range cases { t.Run(rawURL, func(t *testing.T) { - parsedURL, _, err := resolveNotificationTarget(rawURL) + parsedURL, _, err := utils.ResolveAllowedHTTPURL(rawURL) if err != nil { t.Fatalf("expected %s to be allowed, got %v", rawURL, err) } @@ -334,31 +336,36 @@ func TestNotificationTargetAllowsPublicAddresses(t *testing.T) { } } -func TestNotificationHTTPClientPreservesTLSServerName(t *testing.T) { - client, err := newNotificationHTTPClient("https://1.1.1.1/webhook", true) - if err != nil { - t.Fatalf("expected public HTTPS URL to create client: %v", err) +func TestNotificationHTTPClientInvertsVerifyTLSFlag(t *testing.T) { + // newNotificationHTTPClient takes verifyTLS, utils.NewRestrictedHTTPClient + // takes skipVerifyTLS. The wrapper must invert the boolean; if a future + // refactor drops the negation, TLS verification silently turns off. + // SNI / redirect / IP-pinning are covered by pkg/utils/http_test.go. + cases := []struct { + name string + verifyTLS bool + wantSkipVerifyOn bool + }{ + {"verifyTLS_true_means_skipVerify_false", true, false}, + {"verifyTLS_false_means_skipVerify_true", false, true}, } - transport, ok := client.Transport.(*http.Transport) - if !ok { - t.Fatalf("expected http.Transport, got %T", client.Transport) - } - if transport.TLSClientConfig == nil || transport.TLSClientConfig.ServerName != "1.1.1.1" { - t.Fatalf("expected TLS ServerName 1.1.1.1, got %#v", transport.TLSClientConfig) - } -} - -func TestNotificationHTTPClientRejectsRedirects(t *testing.T) { - client, err := newNotificationHTTPClient("https://1.1.1.1/webhook", true) - if err != nil { - t.Fatalf("expected client construction: %v", err) - } - req, err := http.NewRequest(http.MethodGet, "https://1.1.1.1/start", nil) - if err != nil { - t.Fatalf("new request: %v", err) - } - via := []*http.Request{req} - if err := client.CheckRedirect(req, via); err != http.ErrUseLastResponse { - t.Fatalf("expected ErrUseLastResponse, got %v", err) + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + client, err := newNotificationHTTPClient("https://1.1.1.1/webhook", tc.verifyTLS) + if err != nil { + t.Fatalf("expected client construction: %v", err) + } + transport, ok := client.Transport.(*http.Transport) + if !ok { + t.Fatalf("expected *http.Transport, got %T", client.Transport) + } + if transport.TLSClientConfig == nil { + t.Fatalf("expected TLSClientConfig to be set") + } + if got := transport.TLSClientConfig.InsecureSkipVerify; got != tc.wantSkipVerifyOn { + t.Fatalf("verifyTLS=%v: expected InsecureSkipVerify=%v, got %v", + tc.verifyTLS, tc.wantSkipVerifyOn, got) + } + }) } } diff --git a/pkg/ddns/webhook/webhook.go b/pkg/ddns/webhook/webhook.go index 480df4c3..61e0fbe2 100644 --- a/pkg/ddns/webhook/webhook.go +++ b/pkg/ddns/webhook/webhook.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "io" "net/http" "net/url" "strings" @@ -57,13 +58,19 @@ func (provider *Provider) SetRecords(ctx context.Context, zone string, provider.ipAddr = rr.Data provider.domain = fmt.Sprintf("%s.%s", rr.Name, strings.TrimSuffix(zone, ".")) - req, err := provider.prepareRequest(ctx) + // WebhookURL is attacker-controlled (GHSA-6x26-5727-rrm9); the request and + // the client are paired so URL validation and DialContext pinning are driven + // by a single DNS resolution. Do not swap the client for utils.HttpClient. + req, client, err := provider.prepareRequest(ctx) if err != nil { return nil, fmt.Errorf("failed to update a domain: %s. Cause by: %v", provider.domain, err) } - if _, err := utils.HttpClient.Do(req); err != nil { + resp, err := client.Do(req) + if err != nil { return nil, fmt.Errorf("failed to update a domain: %s. Cause by: %v", provider.domain, err) } + _, _ = io.Copy(io.Discard, resp.Body) + resp.Body.Close() default: return nil, fmt.Errorf("unsupported record type: %T", rec) } @@ -72,26 +79,32 @@ func (provider *Provider) SetRecords(ctx context.Context, zone string, return recs, nil } -func (provider *Provider) prepareRequest(ctx context.Context) (*http.Request, error) { +func (provider *Provider) prepareRequest(ctx context.Context) (*http.Request, *http.Client, error) { u, err := provider.reqUrl() if err != nil { - return nil, err + return nil, nil, err + } + // Single SSRF check + dial pin; the returned client must be used by callers + // so the dialer's pinned IP and the validated URL stay in sync. + client, err := utils.NewRestrictedHTTPClient(u.String(), false) + if err != nil { + return nil, nil, err } body, err := provider.reqBody() if err != nil { - return nil, err + return nil, nil, err } headers, err := utils.GjsonIter( provider.formatWebhookString(provider.DDNSProfile.WebhookHeaders)) if err != nil { - return nil, err + return nil, nil, err } req, err := http.NewRequestWithContext(ctx, requestTypes[provider.DDNSProfile.WebhookMethod], u.String(), strings.NewReader(body)) if err != nil { - return nil, err + return nil, nil, err } provider.setContentType(req) @@ -100,7 +113,7 @@ func (provider *Provider) prepareRequest(ctx context.Context) (*http.Request, er req.Header.Set(k, v) } - return req, nil + return req, client, nil } func (provider *Provider) setContentType(req *http.Request) { diff --git a/pkg/ddns/webhook/webhook_test.go b/pkg/ddns/webhook/webhook_test.go index 028a245c..8bbe9b3e 100644 --- a/pkg/ddns/webhook/webhook_test.go +++ b/pkg/ddns/webhook/webhook_test.go @@ -2,6 +2,7 @@ package webhook import ( "context" + "strings" "testing" "github.com/nezhahq/nezha/model" @@ -44,7 +45,7 @@ func execCase(t *testing.T, item testSt) { t.Fatalf("Expected %s, but got %s", item.expectBody, reqBody) } - req, err := pw.prepareRequest(context.Background()) + req, _, err := pw.prepareRequest(context.Background()) if err != nil { t.Fatalf("Error: %s", err) } @@ -69,11 +70,11 @@ func TestWebhookRequest(t *testing.T) { Domains: []string{"www.example.com"}, MaxRetries: 1, EnableIPv4: &ipv4, - WebhookURL: "http://ddns.example.com/?ip=#ip#", + WebhookURL: "http://1.1.1.1/?ip=#ip#", WebhookMethod: methodGET, WebhookHeaders: `{"ip":"#ip#","record":"#record#"}`, }, - expectURL: "http://ddns.example.com/?ip=1.1.1.1", + expectURL: "http://1.1.1.1/?ip=1.1.1.1", expectContentType: "", expectHeader: map[string]string{ "ip": "1.1.1.1", @@ -85,12 +86,12 @@ func TestWebhookRequest(t *testing.T) { Domains: []string{"www.example.com"}, MaxRetries: 1, EnableIPv4: &ipv4, - WebhookURL: "http://ddns.example.com/api", + WebhookURL: "http://1.1.1.1/api", WebhookMethod: methodPOST, WebhookRequestType: requestTypeJSON, WebhookRequestBody: `{"ip":"#ip#","record":"#record#"}`, }, - expectURL: "http://ddns.example.com/api", + expectURL: "http://1.1.1.1/api", expectContentType: reqTypeJSON, expectBody: `{"ip":"1.1.1.1","record":"A"}`, }, @@ -99,12 +100,12 @@ func TestWebhookRequest(t *testing.T) { Domains: []string{"www.example.com"}, MaxRetries: 1, EnableIPv4: &ipv4, - WebhookURL: "http://ddns.example.com/api", + WebhookURL: "http://1.1.1.1/api", WebhookMethod: methodPOST, WebhookRequestType: requestTypeForm, WebhookRequestBody: `{"ip":"#ip#","record":"#record#"}`, }, - expectURL: "http://ddns.example.com/api", + expectURL: "http://1.1.1.1/api", expectContentType: reqTypeForm, expectBody: "ip=1.1.1.1&record=A", }, @@ -114,3 +115,58 @@ func TestWebhookRequest(t *testing.T) { execCase(t, c) } } + +func TestWebhookTargetRejectsBlockedRanges(t *testing.T) { + cases := []string{ + "http://0.0.0.0/", + "http://10.1.2.3/", + "http://100.64.0.1/", + "http://127.0.0.1/", + "http://127.255.255.254/", + "http://169.254.169.254/", + "http://172.16.0.1/", + "http://192.0.0.1/", + "http://192.0.2.1/", + "http://192.168.1.1/", + "http://198.18.0.1/", + "http://198.51.100.1/", + "http://203.0.113.1/", + "http://224.0.0.1/", + "http://240.0.0.1/", + "http://[::]/", + "http://[::1]/", + "http://[::ffff:127.0.0.1]/", + "http://[64:ff9b::1]/", + "http://[100::1]/", + "http://[2001:db8::1]/", + "http://[fc00::1]/", + "http://[fe80::1]/", + "http://[ff00::1]/", + "ftp://example.com/", + "file:///etc/passwd", + "http:///path", + } + + for _, rawURL := range cases { + t.Run(rawURL, func(t *testing.T) { + provider := Provider{DDNSProfile: &model.DDNSProfile{ + Domains: []string{"www.example.com"}, + WebhookURL: rawURL, + WebhookMethod: methodGET, + WebhookHeaders: `{}`, + }} + provider.ipAddr = "1.1.1.1" + provider.domain = provider.DDNSProfile.Domains[0] + provider.ipType = "ipv4" + provider.recordType = "A" + + _, _, err := provider.prepareRequest(context.Background()) + if err == nil { + t.Fatalf("expected %s to be rejected", rawURL) + } + if !strings.Contains(err.Error(), "not allowed") { + t.Fatalf("expected not allowed error, got %q", err.Error()) + } + }) + } +} diff --git a/pkg/utils/http.go b/pkg/utils/http.go index 19227c47..72aab42b 100644 --- a/pkg/utils/http.go +++ b/pkg/utils/http.go @@ -1,16 +1,53 @@ package utils import ( + "context" "crypto/tls" + "errors" + "net" "net/http" + "net/netip" + "net/url" "time" ) +// HttpClient / HttpClientSkipTlsVerify must not be used to dispatch +// requests to user-controlled URLs (SSRF risk, GHSA-6x26-5727-rrm9). +// For any attacker-controlled URL use NewRestrictedHTTPClient instead. var ( HttpClientSkipTlsVerify *http.Client HttpClient *http.Client ) +var ErrHTTPURLTargetNotAllowed = errors.New("HTTP URL target is not allowed") + +var blockedHTTPClientCIDRs = mustParseHTTPClientCIDRs([]string{ + "0.0.0.0/8", + "10.0.0.0/8", + "100.64.0.0/10", + "127.0.0.0/8", + "169.254.0.0/16", + "172.16.0.0/12", + "192.0.0.0/24", + "192.0.2.0/24", + "192.168.0.0/16", + "198.18.0.0/15", + "198.51.100.0/24", + "203.0.113.0/24", + "224.0.0.0/4", + "240.0.0.0/4", + "::/128", + "::1/128", + "::ffff:0:0/96", + "64:ff9b::/96", + "100::/64", + "2001::/23", + "2001:db8::/32", + "fc00::/7", + "fe80::/10", + "ff00::/8", +}) + func init() { HttpClientSkipTlsVerify = httpClient(_httpClient{ Transport: httpTransport(_httpTransport{ @@ -47,3 +84,106 @@ func httpClient(conf _httpClient) *http.Client { Timeout: time.Minute * 10, } } + +func NewRestrictedHTTPClient(rawURL string, skipVerifyTLS bool) (*http.Client, error) { + parsedURL, ip, err := ResolveAllowedHTTPURL(rawURL) + if err != nil { + return nil, err + } + return buildRestrictedHTTPClient(parsedURL, ip, skipVerifyTLS), nil +} + +// buildRestrictedHTTPClient assembles a client whose DialContext is pinned to +// the already-vetted IP. Separated from NewRestrictedHTTPClient so tests can +// exercise the SNI / redirect behavior without relying on live DNS. +func buildRestrictedHTTPClient(parsedURL *url.URL, ip net.IP, skipVerifyTLS bool) *http.Client { + port := parsedURL.Port() + if port == "" { + if parsedURL.Scheme == "https" { + port = "443" + } else { + port = "80" + } + } + // Pin outbound webhooks to the vetted IP so DNS changes cannot retarget private hosts. + targetAddress := net.JoinHostPort(ip.String(), port) + dialer := &net.Dialer{} + + return &http.Client{ + Transport: &http.Transport{ + DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { + return dialer.DialContext(ctx, network, targetAddress) + }, + TLSClientConfig: &tls.Config{InsecureSkipVerify: skipVerifyTLS, ServerName: parsedURL.Hostname()}, + }, + CheckRedirect: func(req *http.Request, via []*http.Request) error { + return http.ErrUseLastResponse + }, + Timeout: time.Minute * 10, + } +} + +func ResolveAllowedHTTPURL(rawURL string) (*url.URL, net.IP, error) { + parsedURL, err := url.Parse(rawURL) + if err != nil { + return nil, nil, err + } + if parsedURL.Scheme != "http" && parsedURL.Scheme != "https" { + return nil, nil, ErrHTTPURLTargetNotAllowed + } + + host := parsedURL.Hostname() + if host == "" { + return nil, nil, ErrHTTPURLTargetNotAllowed + } + if ip := net.ParseIP(host); ip != nil { + if !HTTPURLTargetIPAllowed(ip) { + return nil, nil, ErrHTTPURLTargetNotAllowed + } + return parsedURL, ip, nil + } + + ips, err := net.LookupIP(host) + if err != nil { + return nil, nil, err + } + if len(ips) == 0 { + return nil, nil, ErrHTTPURLTargetNotAllowed + } + for _, ip := range ips { + if !HTTPURLTargetIPAllowed(ip) { + return nil, nil, ErrHTTPURLTargetNotAllowed + } + } + + return parsedURL, ips[0], nil +} + +func HTTPURLTargetIPAllowed(ip net.IP) bool { + parsedIP, ok := netipFromIP(ip) + if !ok { + return false + } + for _, cidr := range blockedHTTPClientCIDRs { + if cidr.Contains(parsedIP) { + return false + } + } + return parsedIP.IsGlobalUnicast() +} + +func netipFromIP(ip net.IP) (netip.Addr, bool) { + parsedIP, ok := netip.AddrFromSlice(ip) + if !ok { + return netip.Addr{}, false + } + return parsedIP.Unmap(), true +} + +func mustParseHTTPClientCIDRs(cidrs []string) []netip.Prefix { + prefixes := make([]netip.Prefix, 0, len(cidrs)) + for _, cidr := range cidrs { + prefixes = append(prefixes, netip.MustParsePrefix(cidr)) + } + return prefixes +} diff --git a/pkg/utils/http_test.go b/pkg/utils/http_test.go new file mode 100644 index 00000000..b0cd9924 --- /dev/null +++ b/pkg/utils/http_test.go @@ -0,0 +1,110 @@ +package utils + +import ( + "net" + "net/http" + "net/url" + "testing" + "time" +) + +func TestBuildRestrictedHTTPClientPreservesHostnameAsTLSServerName(t *testing.T) { + // Construct a hostname URL paired with an arbitrary public IP so we exercise + // the SNI preservation path without depending on live DNS in unit tests. + parsed, err := url.Parse("https://example.com/webhook") + if err != nil { + t.Fatalf("parse url: %v", err) + } + pinnedIP := net.ParseIP("1.1.1.1") + if pinnedIP == nil { + t.Fatalf("expected valid pinned IP") + } + + client := buildRestrictedHTTPClient(parsed, pinnedIP, false) + transport, ok := client.Transport.(*http.Transport) + if !ok { + t.Fatalf("expected *http.Transport, got %T", client.Transport) + } + if transport.TLSClientConfig == nil { + t.Fatalf("expected TLSClientConfig to be set") + } + // SNI must come from the original URL hostname so the certificate validates + // the intended hostname, not the pinned dial IP. + if got := transport.TLSClientConfig.ServerName; got != "example.com" { + t.Fatalf("expected ServerName example.com, got %q", got) + } + if transport.TLSClientConfig.ServerName == pinnedIP.String() { + t.Fatalf("ServerName must not be the pinned IP, got %q", transport.TLSClientConfig.ServerName) + } + if transport.TLSClientConfig.InsecureSkipVerify { + t.Fatalf("expected verifyTLS path (InsecureSkipVerify=false)") + } +} + +func TestBuildRestrictedHTTPClientHonorsSkipVerifyTLS(t *testing.T) { + parsed, _ := url.Parse("https://example.com/webhook") + client := buildRestrictedHTTPClient(parsed, net.ParseIP("1.1.1.1"), true) + transport := client.Transport.(*http.Transport) + if !transport.TLSClientConfig.InsecureSkipVerify { + t.Fatalf("expected InsecureSkipVerify=true when skipVerifyTLS=true") + } +} + +func TestBuildRestrictedHTTPClientRejectsRedirects(t *testing.T) { + parsed, _ := url.Parse("https://example.com/start") + client := buildRestrictedHTTPClient(parsed, net.ParseIP("1.1.1.1"), false) + req, err := http.NewRequest(http.MethodGet, "https://example.com/start", nil) + if err != nil { + t.Fatalf("new request: %v", err) + } + if err := client.CheckRedirect(req, []*http.Request{req}); err != http.ErrUseLastResponse { + t.Fatalf("expected ErrUseLastResponse, got %v", err) + } +} + +// TestBuildRestrictedHTTPClientPinsDialToVettedIP confirms DialContext routes +// to the pinned IP even when the request URL uses a different hostname, +// preventing DNS rebinding from retargeting traffic. +func TestBuildRestrictedHTTPClientPinsDialToVettedIP(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + defer listener.Close() + _, port, err := net.SplitHostPort(listener.Addr().String()) + if err != nil { + t.Fatalf("split host port: %v", err) + } + + accepted := make(chan string, 1) + go func() { + conn, err := listener.Accept() + if err != nil { + accepted <- "" + return + } + accepted <- conn.LocalAddr().String() + conn.Close() + }() + + requestURL := "http://example.com:" + port + "/" + parsed, _ := url.Parse(requestURL) + pinned := net.ParseIP("127.0.0.1") + client := buildRestrictedHTTPClient(parsed, pinned, false) + client.Timeout = 2 * time.Second + + req, _ := http.NewRequest(http.MethodGet, requestURL, nil) + resp, _ := client.Do(req) + if resp != nil { + resp.Body.Close() + } + + select { + case addr := <-accepted: + if addr == "" { + t.Fatalf("listener accept failed") + } + case <-time.After(2 * time.Second): + t.Fatalf("expected dial to reach pinned IP 127.0.0.1:%s, listener did not accept", port) + } +}