mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
fix(coderd): validate webpush subscription endpoints (#24347)
Co-authored-by: Cian Johnston <cian@coder.com>
This commit is contained in:
co-authored by
Cian Johnston
parent
e317f3b239
commit
5812f84e1c
@@ -6,9 +6,12 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/SherClockHolmes/webpush-go"
|
||||
@@ -47,6 +50,7 @@ type SubscriptionCacheInvalidator interface {
|
||||
type options struct {
|
||||
clock quartz.Clock
|
||||
subscriptionCacheTTL time.Duration
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
// Option configures optional behavior for a Webpusher.
|
||||
@@ -68,6 +72,15 @@ func WithSubscriptionCacheTTL(ttl time.Duration) Option {
|
||||
}
|
||||
}
|
||||
|
||||
// WithHTTPClient overrides the default SSRF-safe HTTP client used to deliver
|
||||
// push notifications. This is intended for tests that need to deliver to
|
||||
// localhost test servers.
|
||||
func WithHTTPClient(client *http.Client) Option {
|
||||
return func(o *options) {
|
||||
o.httpClient = client
|
||||
}
|
||||
}
|
||||
|
||||
// New creates a new Dispatcher to dispatch web push notifications.
|
||||
//
|
||||
// This is *not* integrated into the enqueue system unfortunately.
|
||||
@@ -90,6 +103,9 @@ func New(ctx context.Context, log *slog.Logger, db database.Store, vapidSub stri
|
||||
if cfg.subscriptionCacheTTL <= 0 {
|
||||
cfg.subscriptionCacheTTL = defaultSubscriptionCacheTTL
|
||||
}
|
||||
if cfg.httpClient == nil {
|
||||
cfg.httpClient = newSSRFSafeHTTPClient()
|
||||
}
|
||||
|
||||
keys, err := db.GetWebpushVAPIDKeys(ctx)
|
||||
if err != nil {
|
||||
@@ -121,6 +137,7 @@ func New(ctx context.Context, log *slog.Logger, db database.Store, vapidSub stri
|
||||
subscriptionCacheTTL: cfg.subscriptionCacheTTL,
|
||||
subscriptionCache: make(map[uuid.UUID]cachedSubscriptions),
|
||||
subscriptionGenerations: make(map[uuid.UUID]uint64),
|
||||
httpClient: cfg.httpClient,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -142,6 +159,12 @@ type Webpusher struct {
|
||||
VAPIDPublicKey string
|
||||
VAPIDPrivateKey string
|
||||
|
||||
// httpClient is an SSRF-safe HTTP client that rejects connections to
|
||||
// private, loopback, and link-local IP addresses at dial time. This
|
||||
// closes the DNS rebinding TOCTOU gap where a hostname passes URL
|
||||
// validation but resolves to a private IP when the connection is made.
|
||||
httpClient *http.Client
|
||||
|
||||
clock quartz.Clock
|
||||
|
||||
cacheMu sync.RWMutex
|
||||
@@ -338,6 +361,7 @@ func (n *Webpusher) webpushSend(ctx context.Context, msg []byte, endpoint string
|
||||
Endpoint: endpoint,
|
||||
Keys: keys,
|
||||
}, &webpush.Options{
|
||||
HTTPClient: n.httpClient,
|
||||
Subscriber: n.vapidSub,
|
||||
VAPIDPublicKey: n.VAPIDPublicKey,
|
||||
VAPIDPrivateKey: n.VAPIDPrivateKey,
|
||||
@@ -407,6 +431,37 @@ func (*NoopWebpusher) PublicKey() string {
|
||||
return ""
|
||||
}
|
||||
|
||||
// newSSRFSafeHTTPClient returns an HTTP client that rejects connections to
|
||||
// private, loopback, link-local, multicast, and unspecified IP addresses.
|
||||
// This prevents DNS rebinding attacks where a hostname passes URL-level
|
||||
// validation but resolves to an internal IP at dial time.
|
||||
func newSSRFSafeHTTPClient() *http.Client {
|
||||
return &http.Client{
|
||||
Transport: &http.Transport{
|
||||
DialContext: (&net.Dialer{
|
||||
Control: func(_ string, address string, _ syscall.RawConn) error {
|
||||
host, _, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("split host/port: %w", err)
|
||||
}
|
||||
ip, err := netip.ParseAddr(host)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("parse resolved IP: %w", err)
|
||||
}
|
||||
if ip.IsPrivate() || ip.IsLoopback() || ip.IsLinkLocalUnicast() ||
|
||||
ip.IsLinkLocalMulticast() || ip.IsMulticast() ||
|
||||
ip.IsUnspecified() {
|
||||
return xerrors.Errorf(
|
||||
"webpush endpoint resolved to non-public address %s", ip.String(),
|
||||
)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}).DialContext,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// RegenerateVAPIDKeys regenerates the VAPID keys and deletes all existing
|
||||
// push subscriptions as part of the transaction, as they are no longer valid.
|
||||
func RegenerateVAPIDKeys(ctx context.Context, db database.Store) (newPrivateKey string, newPublicKey string, err error) {
|
||||
|
||||
@@ -387,6 +387,8 @@ func assertWebpushPayload(t testing.TB, r *http.Request) {
|
||||
}
|
||||
|
||||
// setupPushTest creates a common test setup for webpush notification tests.
|
||||
// The test HTTP client bypasses SSRF protection so that httptest.Server
|
||||
// (bound to 127.0.0.1) can be reached.
|
||||
func setupPushTest(ctx context.Context, t *testing.T, handlerFunc func(w http.ResponseWriter, r *http.Request)) (webpush.Dispatcher, database.Store, string) {
|
||||
t.Helper()
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
@@ -400,6 +402,9 @@ func setupPushTestWithOptions(ctx context.Context, t *testing.T, db database.Sto
|
||||
server := httptest.NewServer(http.HandlerFunc(handlerFunc))
|
||||
t.Cleanup(server.Close)
|
||||
|
||||
// Use an unrestricted HTTP client for tests. The default SSRF-safe
|
||||
// client rejects loopback addresses, which blocks httptest.Server.
|
||||
opts = append(opts, webpush.WithHTTPClient(http.DefaultClient))
|
||||
manager, err := webpush.New(ctx, &logger, db, "http://example.com", opts...)
|
||||
require.NoError(t, err, "Failed to create webpush manager")
|
||||
|
||||
@@ -423,3 +428,56 @@ func TestNoopWebpusher(t *testing.T) {
|
||||
|
||||
require.Empty(t, noop.PublicKey())
|
||||
}
|
||||
|
||||
// TestSSRFPrevention verifies that the default SSRF-safe HTTP client blocks
|
||||
// webpush delivery to loopback (and other non-public) addresses. This
|
||||
// reproduces the attack vector from the original SSRF PoC: an authenticated
|
||||
// user supplies a localhost endpoint in their webpush subscription, and the
|
||||
// server must refuse to connect.
|
||||
func TestSSRFPrevention(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
|
||||
// Start a server that records whether it received a request.
|
||||
var received atomic.Bool
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
received.Store(true)
|
||||
w.WriteHeader(http.StatusCreated)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// Create a dispatcher via New() WITHOUT WithHTTPClient so it
|
||||
// uses the default SSRF-safe client that blocks loopback.
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
logger := slogtest.Make(t, &slogtest.Options{IgnoreErrors: true}).Leveled(slog.LevelDebug)
|
||||
manager, err := webpush.New(ctx, &logger, db, "http://example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Test() calls webpushSend directly with the supplied endpoint.
|
||||
err = manager.Test(ctx, codersdk.WebpushSubscription{
|
||||
Endpoint: server.URL,
|
||||
AuthKey: validEndpointAuthKey,
|
||||
P256DHKey: validEndpointP256dhKey,
|
||||
})
|
||||
require.Error(t, err, "SSRF-safe client should reject Test() to loopback address")
|
||||
assert.False(t, received.Load(), "Test() request should not reach the localhost server")
|
||||
|
||||
// Dispatch() goes through the subscription cache → webpushSend path.
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
_, err = db.InsertWebpushSubscription(ctx, database.InsertWebpushSubscriptionParams{
|
||||
CreatedAt: dbtime.Now(),
|
||||
UserID: user.ID,
|
||||
Endpoint: server.URL,
|
||||
EndpointAuthKey: validEndpointAuthKey,
|
||||
EndpointP256dhKey: validEndpointP256dhKey,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
err = manager.Dispatch(ctx, user.ID, codersdk.WebpushMessage{
|
||||
Title: "SSRF test",
|
||||
Body: "This should not arrive.",
|
||||
})
|
||||
require.Error(t, err, "SSRF-safe client should reject Dispatch() to loopback address")
|
||||
assert.False(t, received.Load(), "Dispatch() request should not reach the localhost server")
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user