From 3ffe013886a9af9e60090a1ddbf8cb399a55b948 Mon Sep 17 00:00:00 2001 From: Elias Schneider Date: Fri, 9 Oct 2026 22:32:52 +0200 Subject: [PATCH] feat: add ability to override SSRF protection with `OUTBOUND_ALLOWED_HOSTS_*` --- .vscode/settings.json | 3 + backend/internal/bootstrap/bootstrap.go | 8 +- .../internal/bootstrap/services_bootstrap.go | 13 +- backend/internal/common/env_config.go | 32 +-- backend/internal/common/env_config_test.go | 36 --- backend/internal/oidc/cimd.go | 4 +- backend/internal/oidc/cimd_test.go | 47 ++++ backend/internal/oidc/module.go | 12 +- backend/internal/outbound/clients.go | 115 ++++++++ backend/internal/outbound/clients_test.go | 71 +++++ backend/internal/outbound/policy.go | 254 ++++++++++++++++++ backend/internal/outbound/policy_test.go | 121 +++++++++ backend/internal/outbound/transport.go | 85 ++++++ backend/internal/outbound/transport_test.go | 109 ++++++++ backend/internal/service/oidc_service.go | 33 +-- backend/internal/service/oidc_service_test.go | 32 +-- backend/internal/utils/ip_util.go | 27 -- backend/internal/utils/ip_util_test.go | 215 --------------- tests/setup/docker-compose.yml | 1 + 19 files changed, 865 insertions(+), 353 deletions(-) create mode 100644 backend/internal/outbound/clients.go create mode 100644 backend/internal/outbound/clients_test.go create mode 100644 backend/internal/outbound/policy.go create mode 100644 backend/internal/outbound/policy_test.go create mode 100644 backend/internal/outbound/transport.go create mode 100644 backend/internal/outbound/transport_test.go diff --git a/.vscode/settings.json b/.vscode/settings.json index c2225d93..6e153101 100644 --- a/.vscode/settings.json +++ b/.vscode/settings.json @@ -4,5 +4,8 @@ "oxc.fmt.disableNestedConfig": true, "[svelte]": { "editor.defaultFormatter": "oxc.oxc-vscode" + }, + "[go]": { + "editor.defaultFormatter": "golang.go" } } diff --git a/backend/internal/bootstrap/bootstrap.go b/backend/internal/bootstrap/bootstrap.go index 2d2122fe..66250de9 100644 --- a/backend/internal/bootstrap/bootstrap.go +++ b/backend/internal/bootstrap/bootstrap.go @@ -16,6 +16,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/instanceid" + "github.com/pocket-id/pocket-id/backend/internal/outbound" "github.com/pocket-id/pocket-id/backend/internal/storage" ) @@ -35,6 +36,11 @@ func Bootstrap(ctx context.Context) error { slog.InfoContext(ctx, "Pocket ID is starting") + outboundClients, err := outbound.New(&common.EnvConfig) + if err != nil { + return fmt.Errorf("failed to initialize outbound HTTP clients: %w", err) + } + // Init database db, pg, err := NewDatabase(ctx) if err != nil { @@ -95,7 +101,7 @@ func Bootstrap(ctx context.Context) error { services = append(services, actorsRun) // Create all services - svc, err := initServices(ctx, db, instanceID, actors, httpClient, imageExtensions, fileStorage) + svc, err := initServices(ctx, db, instanceID, actors, httpClient, outboundClients, imageExtensions, fileStorage) if err != nil { return fmt.Errorf("failed to initialize services: %w", err) } diff --git a/backend/internal/bootstrap/services_bootstrap.go b/backend/internal/bootstrap/services_bootstrap.go index 88ee9359..4312141a 100644 --- a/backend/internal/bootstrap/services_bootstrap.go +++ b/backend/internal/bootstrap/services_bootstrap.go @@ -24,6 +24,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/logopreset" "github.com/pocket-id/pocket-id/backend/internal/oidc" "github.com/pocket-id/pocket-id/backend/internal/onetimeaccess" + "github.com/pocket-id/pocket-id/backend/internal/outbound" "github.com/pocket-id/pocket-id/backend/internal/scimsync" "github.com/pocket-id/pocket-id/backend/internal/service" "github.com/pocket-id/pocket-id/backend/internal/storage" @@ -67,6 +68,7 @@ func initServices( instanceID string, actors francishost.Host, httpClient *http.Client, + outboundClients *outbound.Clients, imageExtensions map[string]string, fileStorage storage.FileStorage, ) (svc *services, err error) { @@ -146,7 +148,7 @@ func initServices( svc.scimSyncModule, err = scimsync.New(scimsync.Dependencies{ DB: db, Actors: actors, - HTTPClient: httpClient, + HTTPClient: outboundClients.Client(outbound.PurposeSCIM), // Disable in test environment ScheduleDisabled: common.EnvConfig.AppEnv.IsTest(), }) @@ -159,7 +161,8 @@ func initServices( svc.oidcModule, err = oidc.New(ctx, oidc.Dependencies{ DB: db, Actors: actors, - HTTPClient: httpClient, + FederatedJWKSClient: outboundClients.Client(outbound.PurposeFederatedJWKS), + CIMDTransport: outboundClients.Transport(outbound.PurposeClientMetadata), GetCIMDURLAllowlist: svc.appConfigService.GetCIMDURLAllowlist, Config: oidc.Config{ BaseURL: common.EnvConfig.AppURL, @@ -179,12 +182,12 @@ func initServices( return nil, fmt.Errorf("failed to create OIDC module: %w", err) } - backchannelLogoutService, err := backchannellogout.NewService(db, svc.jwtService, httpClient, actors) + backchannelLogoutService, err := backchannellogout.NewService(db, svc.jwtService, outboundClients.Client(outbound.PurposeBackchannelLogout), actors) if err != nil { return nil, fmt.Errorf("failed to create back-channel logout service: %w", err) } - svc.oidcService, err = service.NewOidcService(db, svc.jwtService, svc.oidcModule.Preview, svc.oidcModule, svc.scimSyncModule, backchannelLogoutService, httpClient, fileStorage) + svc.oidcService, err = service.NewOidcService(db, svc.jwtService, svc.oidcModule.Preview, svc.oidcModule, svc.scimSyncModule, backchannelLogoutService, outboundClients.Client(outbound.PurposeClientLogo), fileStorage) if err != nil { return nil, fmt.Errorf("failed to create OIDC service: %w", err) } @@ -195,7 +198,7 @@ func initServices( svc.ldapSyncModule, err = ldapsync.New(ldapsync.Dependencies{ DB: db, Actors: actors, - HTTPClient: httpClient, + HTTPClient: outboundClients.Client(outbound.PurposeLDAPPicture), FileStorage: fileStorage, Users: svc.userService, Groups: svc.userGroupService, diff --git a/backend/internal/common/env_config.go b/backend/internal/common/env_config.go index 5f6800ed..e974434d 100644 --- a/backend/internal/common/env_config.go +++ b/backend/internal/common/env_config.go @@ -8,7 +8,6 @@ import ( "net" "net/url" "os" - "path" "reflect" "strconv" "strings" @@ -92,6 +91,15 @@ type EnvConfigSchema struct { SystemdSocket bool `env:"SYSTEMD_SOCKET"` LocalIPv6Ranges string `env:"LOCAL_IPV6_RANGES"` + // The OutboundAllowedHosts* fields list the private or reserved destinations each feature's requests to admin- or third-party-controlled URLs may reach + // Each value is a comma-separated list of IP addresses, CIDR ranges, hostnames ("*.example.com" matches subdomains) and the keywords "private", "loopback" and "none" + OutboundAllowedHostsClientLogo string `env:"OUTBOUND_ALLOWED_HOSTS_CLIENT_LOGO"` + OutboundAllowedHostsClientMetadata string `env:"OUTBOUND_ALLOWED_HOSTS_CLIENT_METADATA"` + OutboundAllowedHostsSCIM string `env:"OUTBOUND_ALLOWED_HOSTS_SCIM"` + OutboundAllowedHostsBackchannelLogout string `env:"OUTBOUND_ALLOWED_HOSTS_BACKCHANNEL_LOGOUT"` + OutboundAllowedHostsFederatedJWKS string `env:"OUTBOUND_ALLOWED_HOSTS_FEDERATED_JWKS"` + OutboundAllowedHostsLDAPPicture string `env:"OUTBOUND_ALLOWED_HOSTS_LDAP_PICTURE"` + // TLS cert and key need special treatment with fsnotify, so we aren't using `options:"file"` TLSCert string `env:"TLS_CERT"` TLSKey string `env:"TLS_KEY"` @@ -178,6 +186,10 @@ func defaultConfig() EnvConfigSchema { GeoLiteDBPath: "data/GeoLite2-City.mmdb", GeoLiteDBUrl: MaxMindGeoLiteCityUrl, IconLibraryURL: DefaultIconLibraryURL, + // These targets usually run next to Pocket ID on the same Docker network, cluster or LAN + OutboundAllowedHostsSCIM: "private,loopback", + OutboundAllowedHostsBackchannelLogout: "private,loopback", + OutboundAllowedHostsFederatedJWKS: "private,loopback", } } @@ -432,24 +444,6 @@ func (c *EnvConfigSchema) IconLibraryEnabled() bool { return c.IconLibraryURL != IconLibraryDisabled } -// IsIconLibraryURL reports whether the URL points to a file inside the configured icon library -// Dot segments are rejected so a URL can't climb out of the library's path on the same host -func (c *EnvConfigSchema) IsIconLibraryURL(u *url.URL) bool { - if !c.IconLibraryEnabled() { - return false - } - - base, err := url.Parse(c.IconLibraryURL) - if err != nil { - return false - } - - return u.Scheme == base.Scheme && - strings.EqualFold(u.Host, base.Host) && - u.Path == path.Clean(u.Path) && - strings.HasPrefix(u.Path, base.Path+"/") -} - func validateFileBackend(config *EnvConfigSchema) error { switch config.FileBackend { case "s3", "database": diff --git a/backend/internal/common/env_config_test.go b/backend/internal/common/env_config_test.go index 5a0f6702..4b36ad46 100644 --- a/backend/internal/common/env_config_test.go +++ b/backend/internal/common/env_config_test.go @@ -1,7 +1,6 @@ package common import ( - "net/url" "os" "testing" @@ -699,38 +698,3 @@ func TestIconLibraryConfig(t *testing.T) { }) } } - -func TestIsIconLibraryURL(t *testing.T) { - config := EnvConfigSchema{IconLibraryURL: "http://mirror.lan:4050/icons"} - - tests := []struct { - name string - url string - want bool - }{ - {name: "file inside the library", url: "http://mirror.lan:4050/icons/svg/nextcloud.svg", want: true}, - {name: "host is matched case-insensitively", url: "http://MIRROR.lan:4050/icons/svg/nextcloud.svg", want: true}, - {name: "path outside the library", url: "http://mirror.lan:4050/admin/logo.svg", want: false}, - {name: "path sharing the library prefix", url: "http://mirror.lan:4050/icons-private/logo.svg", want: false}, - {name: "dot segments climbing out of the library", url: "http://mirror.lan:4050/icons/../admin/logo.svg", want: false}, - {name: "encoded dot segments", url: "http://mirror.lan:4050/icons/%2e%2e/admin/logo.svg", want: false}, - {name: "host that starts with the library host", url: "http://mirror.lan.evil.com:4050/icons/svg/nextcloud.svg", want: false}, - {name: "different port", url: "http://mirror.lan:8080/icons/svg/nextcloud.svg", want: false}, - {name: "different scheme", url: "https://mirror.lan:4050/icons/svg/nextcloud.svg", want: false}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - u, err := url.Parse(tt.url) - require.NoError(t, err) - assert.Equal(t, tt.want, config.IsIconLibraryURL(u)) - }) - } - - t.Run("nothing is inside a disabled library", func(t *testing.T) { - disabled := EnvConfigSchema{IconLibraryURL: IconLibraryDisabled} - u, err := url.Parse("http://disabled/svg/nextcloud.svg") - require.NoError(t, err) - assert.False(t, disabled.IsIconLibraryURL(u)) - }) -} diff --git a/backend/internal/oidc/cimd.go b/backend/internal/oidc/cimd.go index a16f6907..5a6463cc 100644 --- a/backend/internal/oidc/cimd.go +++ b/backend/internal/oidc/cimd.go @@ -34,10 +34,10 @@ var _ fosite.CIMDClientPolicy = cimdPolicy{} func newCIMDClientResolver(store *Store, config cimdResolverConfig) *cimdClientResolver { options := []fosite.CIMDFetcherOption{ fosite.WithCIMDUserAgent("pocket-id/oidc-client-metadata-fetcher"), - fosite.WithCIMDExtraPrivateRanges(utils.LocalIPv6IPNets()), } if config.transport != nil { - options = append(options, fosite.WithCIMDTransport(config.transport)) + // The provided transport is the outbound package's guarded transport so we skip Fosite's own pre-resolution checks + options = append(options, fosite.WithCIMDTransport(config.transport), fosite.WithCIMDAllowPrivateIPs(true)) } if config.transportDecorator != nil { options = append(options, fosite.WithCIMDTransportDecorator(config.transportDecorator)) diff --git a/backend/internal/oidc/cimd_test.go b/backend/internal/oidc/cimd_test.go index 4ae2af23..a4ebeb5e 100644 --- a/backend/internal/oidc/cimd_test.go +++ b/backend/internal/oidc/cimd_test.go @@ -4,6 +4,7 @@ import ( "context" "errors" "net/http" + "net/http/httptest" "strings" "testing" "time" @@ -13,11 +14,57 @@ import ( "github.com/stretchr/testify/require" "gorm.io/gorm" + "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/model" datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" + "github.com/pocket-id/pocket-id/backend/internal/outbound" testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing" ) +func TestCIMDClientLimitsResponseHeaders(t *testing.T) { + // Exercise both transport paths so hostname exceptions cannot bypass the header limit + for _, allowlist := range []string{"loopback", "localhost"} { + t.Run(allowlist, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + size := 1024 + if r.URL.Path == "/oversized" { + size = 40 * 1024 + } + w.Header().Set("X-Padding", strings.Repeat("a", size)) + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(server.Close) + + // Use the production outbound transport and CIMD client construction + clients, err := outbound.New(&common.EnvConfigSchema{OutboundAllowedHostsClientMetadata: allowlist}) + require.NoError(t, err) + resolver := newCIMDClientResolver(nil, cimdResolverConfig{transport: clients.Transport(outbound.PurposeClientMetadata)}) + provider, ok := resolver.resolver.Fetcher.(fosite.CIMDSecureHTTPClientProvider) + require.True(t, ok) + client := provider.CIMDHTTPClient() + serverURL := server.URL + if allowlist == "localhost" { + serverURL = strings.Replace(serverURL, "127.0.0.1", "localhost", 1) + } + + // Normal headers must still succeed while oversized headers fail before the body is read + for _, path := range []string{"/normal", "/oversized"} { + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, serverURL+path, nil) + require.NoError(t, err) + res, err := client.Do(req) + if res != nil { + _ = res.Body.Close() + } + if path == "/oversized" { + require.ErrorContains(t, err, "response headers exceeded") + } else { + require.NoError(t, err) + } + } + }) + } +} + func TestBuildClientFromMetadata(t *testing.T) { const id = "https://app.example.com/oauth/client" diff --git a/backend/internal/oidc/module.go b/backend/internal/oidc/module.go index eca32843..9ce6f646 100644 --- a/backend/internal/oidc/module.go +++ b/backend/internal/oidc/module.go @@ -43,10 +43,11 @@ type AuditLogger interface { } type Dependencies struct { - DB *gorm.DB - Actors francishost.Host - Config Config - HTTPClient *http.Client + DB *gorm.DB + Actors francishost.Host + Config Config + FederatedJWKSClient *http.Client + CIMDTransport http.RoundTripper GetCIMDURLAllowlist func() []string @@ -80,13 +81,14 @@ func New(ctx context.Context, deps Dependencies) (*Module, error) { store := NewStore(deps.DB, deps.APIAccess).WithIssuer(deps.Config.BaseURL) cimdResolver := newCIMDClientResolver(store, cimdResolverConfig{ getURLAllowlist: deps.GetCIMDURLAllowlist, + transport: deps.CIMDTransport, transportDecorator: func(transport http.RoundTripper) http.RoundTripper { return otelhttp.NewTransport(transport) }, }) store.clientResolver = cimdResolver - authenticator, err := newFederatedClientAuthenticator(ctx, store, deps.HTTPClient, deps.Config.BaseURL) + authenticator, err := newFederatedClientAuthenticator(ctx, store, deps.FederatedJWKSClient, deps.Config.BaseURL) if err != nil { return nil, fmt.Errorf("failed to create federated client authenticator: %w", err) } diff --git a/backend/internal/outbound/clients.go b/backend/internal/outbound/clients.go new file mode 100644 index 00000000..3bf4f066 --- /dev/null +++ b/backend/internal/outbound/clients.go @@ -0,0 +1,115 @@ +package outbound + +import ( + "fmt" + "net/http" + "net/url" + "strings" + + "github.com/ory/fosite" + "go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp" + + "github.com/pocket-id/pocket-id/backend/internal/common" +) + +const envVarPrefix = "OUTBOUND_ALLOWED_HOSTS_" + +// Purpose identifies a feature that sends requests to URLs that admins or third parties control +type Purpose string + +const ( + PurposeClientLogo Purpose = "client_logo" + PurposeClientMetadata Purpose = "client_metadata" + PurposeSCIM Purpose = "scim" + PurposeBackchannelLogout Purpose = "backchannel_logout" + PurposeFederatedJWKS Purpose = "federated_jwks" + PurposeLDAPPicture Purpose = "ldap_picture" +) + +var AllPurposes = []Purpose{ + PurposeClientLogo, + PurposeClientMetadata, + PurposeSCIM, + PurposeBackchannelLogout, + PurposeFederatedJWKS, + PurposeLDAPPicture, +} + +func (p Purpose) EnvVar() string { + return envVarPrefix + strings.ToUpper(string(p)) +} + +func (p Purpose) allowlist(config *common.EnvConfigSchema) string { + switch p { + case PurposeClientLogo: + return config.OutboundAllowedHostsClientLogo + case PurposeClientMetadata: + return config.OutboundAllowedHostsClientMetadata + case PurposeSCIM: + return config.OutboundAllowedHostsSCIM + case PurposeBackchannelLogout: + return config.OutboundAllowedHostsBackchannelLogout + case PurposeFederatedJWKS: + return config.OutboundAllowedHostsFederatedJWKS + case PurposeLDAPPicture: + return config.OutboundAllowedHostsLDAPPicture + default: + return "" + } +} + +// Clients holds one guarded transport and client per purpose +type Clients struct { + transports map[Purpose]http.RoundTripper + clients map[Purpose]*http.Client +} + +// New builds the guarded clients from the OUTBOUND_ALLOWED_HOSTS_* environment variables +func New(config *common.EnvConfigSchema) (*Clients, error) { + // LOCAL_IPV6_RANGES are blocked like private ranges and opened up by the "private" keyword + localIPv6, err := ParseLocalIPv6Ranges(config.LocalIPv6Ranges) + if err != nil { + return nil, fmt.Errorf("invalid LOCAL_IPV6_RANGES: %w", err) + } + + c := &Clients{ + transports: make(map[Purpose]http.RoundTripper, len(AllPurposes)), + clients: make(map[Purpose]*http.Client, len(AllPurposes)), + } + for _, purpose := range AllPurposes { + policy, err := ParsePolicy(purpose.allowlist(config), localIPv6) + if err != nil { + return nil, fmt.Errorf("invalid %s: %w", purpose.EnvVar(), err) + } + + // The operator configured the icon library, which lets a self-hosted mirror live on the local network + if purpose == PurposeClientLogo && config.IconLibraryEnabled() { + if u, err := url.Parse(config.IconLibraryURL); err == nil && u.Hostname() != "" { + policy.allowedHosts = append(policy.allowedHosts, normalizeHostname(u.Hostname())) + } + } + + // Preserve Fosite's header limit before wrapping the transport so it applies to trusted hosts too + var base *http.Transport + if purpose == PurposeClientMetadata { + base = http.DefaultTransport.(*http.Transport).Clone() //nolint:forcetypeassert // The default transport is always an *http.Transport + base.MaxResponseHeaderBytes = fosite.DefaultCIMDMaxSize * 8 + } + + transport := NewTransport(base, purpose, policy) + c.transports[purpose] = transport + c.clients[purpose] = &http.Client{Transport: otelhttp.NewTransport(transport)} + } + + return c, nil +} + +// Client returns the guarded HTTP client for the purpose, instrumented with OpenTelemetry +func (c *Clients) Client(purpose Purpose) *http.Client { + return c.clients[purpose] +} + +// Transport returns the guarded transport for the purpose without instrumentation, for callers that add their own +func (c *Clients) Transport(purpose Purpose) http.RoundTripper { + return c.transports[purpose] +} diff --git a/backend/internal/outbound/clients_test.go b/backend/internal/outbound/clients_test.go new file mode 100644 index 00000000..e5f3c673 --- /dev/null +++ b/backend/internal/outbound/clients_test.go @@ -0,0 +1,71 @@ +package outbound + +import ( + "errors" + "net/http" + "net/http/httptest" + "reflect" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/pocket-id/pocket-id/backend/internal/common" +) + +func TestPurposeEnvVarsMatchEnvConfig(t *testing.T) { + // Collect every env tag so a renamed field or purpose can't silently drop its allowlist + tags := map[string]bool{} + schema := reflect.TypeFor[common.EnvConfigSchema]() + for field := range schema.Fields() { + tags[strings.Split(field.Tag.Get("env"), ",")[0]] = true + } + + for _, purpose := range AllPurposes { + assert.True(t, tags[purpose.EnvVar()], "missing env field for %s", purpose.EnvVar()) + + // Every purpose must read its own field + config := &common.EnvConfigSchema{} + reflect.ValueOf(config).Elem().FieldByIndex(fieldIndexForTag(t, schema, purpose.EnvVar())).SetString("marker") + assert.Equal(t, "marker", purpose.allowlist(config)) + } +} + +func fieldIndexForTag(t *testing.T, schema reflect.Type, tag string) []int { + t.Helper() + + for field := range schema.Fields() { + if strings.Split(field.Tag.Get("env"), ",")[0] == tag { + return field.Index + } + } + t.Fatalf("no field with env tag %s", tag) + return nil +} + +func TestNewKeepsPurposeAllowlistsSeparate(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(server.Close) + + clients, err := New(&common.EnvConfigSchema{OutboundAllowedHostsSCIM: "loopback"}) + require.NoError(t, err) + + // Opening the loopback range for SCIM must not open it for any other purpose + for _, purpose := range AllPurposes { + err := get(t, clients.Client(purpose), server.URL) + _, blocked := errors.AsType[*BlockedError](err) + if purpose == PurposeSCIM { + require.NoError(t, err) + } else { + assert.True(t, blocked, "%s should be blocked, got %v", purpose, err) + } + } +} + +func TestNewNamesTheInvalidVariable(t *testing.T) { + _, err := New(&common.EnvConfigSchema{OutboundAllowedHostsBackchannelLogout: "10.0.0.0/99"}) + require.ErrorContains(t, err, "OUTBOUND_ALLOWED_HOSTS_BACKCHANNEL_LOGOUT") +} diff --git a/backend/internal/outbound/policy.go b/backend/internal/outbound/policy.go new file mode 100644 index 00000000..8040d538 --- /dev/null +++ b/backend/internal/outbound/policy.go @@ -0,0 +1,254 @@ +package outbound + +import ( + "errors" + "fmt" + "net/netip" + "net/url" + "slices" + "strings" +) + +const ( + KeywordPrivate = "private" + KeywordLoopback = "loopback" + KeywordNone = "none" +) + +// privatePrefixes are the ranges the "private" keyword expands to, in addition to LOCAL_IPV6_RANGES +var privatePrefixes = mustParsePrefixes( + "10.0.0.0/8", + "172.16.0.0/12", + "192.168.0.0/16", + "100.64.0.0/10", + "fc00::/7", +) + +// loopbackPrefixes are the ranges the "loopback" keyword expands to +var loopbackPrefixes = mustParsePrefixes( + "127.0.0.0/8", + "::1/128", +) + +// specialUsePrefixes are reserved ranges that netip's classification methods do not catch +var specialUsePrefixes = mustParsePrefixes( + "0.0.0.0/8", + "100.64.0.0/10", + "192.0.0.0/24", + "192.0.2.0/24", + "192.31.196.0/24", + "192.52.193.0/24", + "192.88.99.0/24", + "192.175.48.0/24", + "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", + "64:ff9b::/96", + "64:ff9b:1::/48", + "100::/64", + "100:0:0:1::/64", + "2001::/23", + "2001:db8::/32", + "2002::/16", + "2620:4f:8000::/48", + "3fff::/20", + "5f00::/16", + "fec0::/10", + "ff00::/8", +) + +// Policy decides which destinations an outbound request may reach +type Policy struct { + // localIPv6 holds LOCAL_IPV6_RANGES, which are public IPv6 ranges the operator uses on the local network + localIPv6 []netip.Prefix + // allowedPrefixes are the non-public ranges the allowlist opens up + allowedPrefixes []netip.Prefix + // allowedHosts are exact hostnames whose resolved addresses are trusted + allowedHosts []string + // allowedHostSuffixes come from "*.example.com" entries and include the leading dot + allowedHostSuffixes []string +} + +// ParsePolicy parses a comma-separated allowlist of keywords, IP addresses, CIDR ranges and hostnames +func ParsePolicy(value string, localIPv6 []netip.Prefix) (Policy, error) { + policy := Policy{localIPv6: localIPv6} + + for entry := range strings.SplitSeq(value, ",") { + entry = strings.ToLower(strings.TrimSpace(entry)) + if entry == "" { + continue + } + + var err error + policy, err = policy.withEntry(entry) + if err != nil { + return Policy{}, err + } + } + + return policy, nil +} + +// withEntry returns a copy of the policy that also allows the given allowlist entry +func (p Policy) withEntry(entry string) (Policy, error) { + // Keywords expand to fixed sets of ranges + switch entry { + case KeywordNone: + return p, nil + case KeywordPrivate: + p.allowedPrefixes = append(p.allowedPrefixes, privatePrefixes...) + p.allowedPrefixes = append(p.allowedPrefixes, p.localIPv6...) + return p, nil + case KeywordLoopback: + p.allowedPrefixes = append(p.allowedPrefixes, loopbackPrefixes...) + return p, nil + } + + // CIDR ranges and single addresses are compared against the IP that is actually dialed + if strings.Contains(entry, "/") { + prefix, err := netip.ParsePrefix(entry) + if err != nil { + return Policy{}, fmt.Errorf("'%s' is not a valid CIDR range", entry) + } + p.allowedPrefixes = append(p.allowedPrefixes, prefix.Masked()) + return p, nil + } + if addr, err := netip.ParseAddr(entry); err == nil { + addr = addr.Unmap() + p.allowedPrefixes = append(p.allowedPrefixes, netip.PrefixFrom(addr, addr.BitLen())) + return p, nil + } + + // Anything else must be a hostname, optionally with a leading wildcard label + if suffix, ok := strings.CutPrefix(entry, "*."); ok { + if !isValidHostname(suffix) { + return Policy{}, fmt.Errorf("'%s' is not a valid hostname pattern", entry) + } + p.allowedHostSuffixes = append(p.allowedHostSuffixes, "."+suffix) + return p, nil + } + if !isValidHostname(entry) { + return Policy{}, fmt.Errorf("'%s' is not a valid keyword, IP address, CIDR range or hostname", entry) + } + p.allowedHosts = append(p.allowedHosts, entry) + return p, nil +} + +// AllowAddr reports whether a connection to the address is permitted +func (p Policy) AllowAddr(addr netip.Addr) bool { + // Zoned addresses are never contained in a prefix, and IPv4-mapped IPv6 addresses must match IPv4 rules + addr = addr.WithZone("").Unmap() + if !addr.IsValid() { + return false + } + + if !p.isBlocked(addr) { + return true + } + + for _, prefix := range p.allowedPrefixes { + if prefix.Contains(addr) { + return true + } + } + + return false +} + +// AllowURL reports whether the URL's host is trusted without checking the addresses it resolves to +func (p Policy) AllowURL(u *url.URL) bool { + host := normalizeHostname(u.Hostname()) + if host == "" { + return false + } + + if slices.Contains(p.allowedHosts, host) { + return true + } + for _, suffix := range p.allowedHostSuffixes { + if strings.HasSuffix(host, suffix) { + return true + } + } + + return false +} + +// hasTrustedHosts reports whether some requests may skip the address check +func (p Policy) hasTrustedHosts() bool { + return len(p.allowedHosts) > 0 || len(p.allowedHostSuffixes) > 0 +} + +// isBlocked reports whether the address is anything other than a public unicast address +func (p Policy) isBlocked(addr netip.Addr) bool { + if addr.IsLoopback() || addr.IsPrivate() || addr.IsUnspecified() || + addr.IsLinkLocalUnicast() || addr.IsLinkLocalMulticast() || + addr.IsInterfaceLocalMulticast() || addr.IsMulticast() { + return true + } + + for _, prefix := range specialUsePrefixes { + if prefix.Contains(addr) { + return true + } + } + for _, prefix := range p.localIPv6 { + if prefix.Contains(addr) { + return true + } + } + + return false +} + +// normalizeHostname lowercases a hostname and drops the trailing dot of a fully qualified name +func normalizeHostname(host string) string { + return strings.TrimSuffix(strings.ToLower(host), ".") +} + +// isValidHostname accepts DNS names made of letters, digits, hyphens, underscores and dots +func isValidHostname(host string) bool { + if host == "" || len(host) > 253 || strings.HasPrefix(host, ".") || strings.Contains(host, "..") { + return false + } + + for _, r := range host { + if (r < 'a' || r > 'z') && (r < '0' || r > '9') && r != '-' && r != '.' && r != '_' { + return false + } + } + + return true +} + +// ParseLocalIPv6Ranges parses the comma-separated LOCAL_IPV6_RANGES value +func ParseLocalIPv6Ranges(value string) ([]netip.Prefix, error) { + var prefixes []netip.Prefix + for entry := range strings.SplitSeq(value, ",") { + entry = strings.TrimSpace(entry) + if entry == "" { + continue + } + + prefix, err := netip.ParsePrefix(entry) + if err != nil { + return nil, err + } + if !prefix.Addr().Is6() || prefix.Addr().Is4In6() { + return nil, errors.New("range '" + entry + "' is not a valid IPv6 range") + } + prefixes = append(prefixes, prefix.Masked()) + } + + return prefixes, nil +} + +func mustParsePrefixes(values ...string) []netip.Prefix { + prefixes := make([]netip.Prefix, len(values)) + for i, value := range values { + prefixes[i] = netip.MustParsePrefix(value) + } + return prefixes +} diff --git a/backend/internal/outbound/policy_test.go b/backend/internal/outbound/policy_test.go new file mode 100644 index 00000000..d8ffd73f --- /dev/null +++ b/backend/internal/outbound/policy_test.go @@ -0,0 +1,121 @@ +package outbound + +import ( + "net/netip" + "net/url" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestParsePolicy(t *testing.T) { + t.Run("accepts every entry type", func(t *testing.T) { + _, err := ParsePolicy(" private, LOOPBACK ,none,192.168.1.5,fd00::/8,10.0.0.0/8,scim-app,*.internal,my_host.example.com", nil) + require.NoError(t, err) + }) + + invalid := []string{ + "10.0.0.0/33", + "http://example.com", + "example.com:443", + "*.", + "exa mple.com", + ".example.com", + "example..com", + } + for _, value := range invalid { + t.Run("rejects "+value, func(t *testing.T) { + _, err := ParsePolicy(value, nil) + require.Error(t, err) + }) + } +} + +func TestPolicyAllowAddr(t *testing.T) { + localIPv6 := []netip.Prefix{netip.MustParsePrefix("2a01:abcd::/32")} + + mustPolicy := func(value string) Policy { + policy, err := ParsePolicy(value, localIPv6) + require.NoError(t, err) + return policy + } + strict := mustPolicy("") + lan := mustPolicy("private,loopback") + explicit := mustPolicy("169.254.169.254,fe80::/10,64:ff9b::/96") + + tests := []struct { + addr string + strict bool + lan bool + explicit bool + }{ + {addr: "8.8.8.8", strict: true, lan: true, explicit: true}, + {addr: "2606:4700:4700::1111", strict: true, lan: true, explicit: true}, + {addr: "127.0.0.1", strict: false, lan: true, explicit: false}, + {addr: "::1", strict: false, lan: true, explicit: false}, + {addr: "10.1.2.3", strict: false, lan: true, explicit: false}, + {addr: "172.20.0.5", strict: false, lan: true, explicit: false}, + {addr: "192.168.1.10", strict: false, lan: true, explicit: false}, + {addr: "100.100.1.1", strict: false, lan: true, explicit: false}, + {addr: "fd12::1", strict: false, lan: true, explicit: false}, + {addr: "2a01:abcd::1", strict: false, lan: true, explicit: false}, + {addr: "::ffff:10.0.0.1", strict: false, lan: true, explicit: false}, + // Cloud metadata and other link-local addresses need an explicit entry even when the LAN is allowed + {addr: "169.254.169.254", strict: false, lan: false, explicit: true}, + {addr: "fe80::1%eth0", strict: false, lan: false, explicit: true}, + // NAT64 can embed a private IPv4 address, so it is never covered by the keywords + {addr: "64:ff9b::a00:1", strict: false, lan: false, explicit: true}, + {addr: "0.0.0.0", strict: false, lan: false, explicit: false}, + {addr: "::", strict: false, lan: false, explicit: false}, + {addr: "224.0.0.1", strict: false, lan: false, explicit: false}, + {addr: "255.255.255.255", strict: false, lan: false, explicit: false}, + {addr: "198.18.0.1", strict: false, lan: false, explicit: false}, + {addr: "2002:a00:1::1", strict: false, lan: false, explicit: false}, + } + + for _, tt := range tests { + t.Run(tt.addr, func(t *testing.T) { + addr := netip.MustParseAddr(tt.addr) + assert.Equal(t, tt.strict, strict.AllowAddr(addr), "strict policy") + assert.Equal(t, tt.lan, lan.AllowAddr(addr), "private,loopback policy") + assert.Equal(t, tt.explicit, explicit.AllowAddr(addr), "explicit policy") + }) + } +} + +func TestPolicyAllowURL(t *testing.T) { + policy, err := ParsePolicy("scim-app,*.internal", nil) + require.NoError(t, err) + + tests := []struct { + url string + allowed bool + }{ + {url: "http://scim-app:8080/scim/v2", allowed: true}, + {url: "http://SCIM-APP./scim/v2", allowed: true}, + {url: "https://idp.internal/jwks", allowed: true}, + {url: "https://a.b.internal/jwks", allowed: true}, + {url: "https://internal/jwks", allowed: false}, + {url: "https://evilinternal/jwks", allowed: false}, + {url: "https://scim-app.example.com/", allowed: false}, + {url: "http://10.0.0.1/", allowed: false}, + } + + for _, tt := range tests { + t.Run(tt.url, func(t *testing.T) { + u, err := url.Parse(tt.url) + require.NoError(t, err) + assert.Equal(t, tt.allowed, policy.AllowURL(u)) + }) + } +} + +func TestParseLocalIPv6Ranges(t *testing.T) { + prefixes, err := ParseLocalIPv6Ranges("2a01:abcd::/32, 2001:db8:1::/48") + require.NoError(t, err) + assert.Len(t, prefixes, 2) + + _, err = ParseLocalIPv6Ranges("10.0.0.0/8") + require.Error(t, err) +} diff --git a/backend/internal/outbound/transport.go b/backend/internal/outbound/transport.go new file mode 100644 index 00000000..e4d0e655 --- /dev/null +++ b/backend/internal/outbound/transport.go @@ -0,0 +1,85 @@ +package outbound + +import ( + "fmt" + "net" + "net/http" + "net/netip" + "syscall" + "time" +) + +type BlockedError struct { + Purpose Purpose + Addr netip.Addr +} + +func (e *BlockedError) Error() string { + return fmt.Sprintf("connection to %s is blocked because it is a private or reserved address; allow it with %s", e.Addr, e.Purpose.EnvVar()) +} + +// guardedTransport routes each request through a dialer that refuses disallowed addresses +type guardedTransport struct { + purpose Purpose + policy Policy + guarded http.RoundTripper + trusted http.RoundTripper +} + +// NewTransport returns a transport that only connects to addresses the policy allows +func NewTransport(base *http.Transport, purpose Purpose, policy Policy) http.RoundTripper { + if base == nil { + base = http.DefaultTransport.(*http.Transport) //nolint:forcetypeassert // The default transport is always an *http.Transport + } + + // Guard the dialer itself so the check sees the IP the connection actually goes to + guarded := base.Clone() + // A proxy would hide the final destination from the dialer and bypass the check + guarded.Proxy = nil + guarded.DialTLSContext = nil + dialer := &net.Dialer{ + Timeout: 30 * time.Second, + KeepAlive: 30 * time.Second, + Control: func(_, address string, _ syscall.RawConn) error { + addrPort, err := netip.ParseAddrPort(address) + if err != nil { + return fmt.Errorf("could not parse dialed address %q: %w", address, err) + } + if !policy.AllowAddr(addrPort.Addr()) { + return &BlockedError{Purpose: purpose, Addr: addrPort.Addr().WithZone("").Unmap()} + } + return nil + }, + } + guarded.DialContext = dialer.DialContext + + t := &guardedTransport{ + purpose: purpose, + policy: policy, + guarded: guarded, + } + + // Hosts allowed by name skip the address check, so they get the base transport unchanged + if policy.hasTrustedHosts() { + t.trusted = base.Clone() + } + + return t +} + +func (t *guardedTransport) RoundTrip(req *http.Request) (*http.Response, error) { + if req.URL.Scheme != "http" && req.URL.Scheme != "https" { + if req.Body != nil { + _ = req.Body.Close() + } + return nil, fmt.Errorf("unsupported URL scheme %q", req.URL.Scheme) + } + + // Pick the transport per request, which also covers every redirect hop + transport := t.guarded + if t.trusted != nil && t.policy.AllowURL(req.URL) { + transport = t.trusted + } + + return transport.RoundTrip(req) +} diff --git a/backend/internal/outbound/transport_test.go b/backend/internal/outbound/transport_test.go new file mode 100644 index 00000000..bfb82063 --- /dev/null +++ b/backend/internal/outbound/transport_test.go @@ -0,0 +1,109 @@ +package outbound + +import ( + "bytes" + "errors" + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func newTestClient(t *testing.T, value string) *http.Client { + t.Helper() + + policy, err := ParsePolicy(value, nil) + require.NoError(t, err) + return &http.Client{Transport: NewTransport(nil, PurposeSCIM, policy)} +} + +func newTestServer(t *testing.T) *httptest.Server { + t.Helper() + + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // Redirect from the "localhost" name to the loopback IP, which is only allowed by address + if r.URL.Path == "/redirect" { + http.Redirect(w, r, server.URL+"/target", http.StatusFound) + return + } + w.WriteHeader(http.StatusOK) + })) + t.Cleanup(server.Close) + return server +} + +func get(t *testing.T, client *http.Client, rawURL string) error { + t.Helper() + + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, rawURL, nil) + require.NoError(t, err) + res, err := client.Do(req) + if err != nil { + return err + } + _ = res.Body.Close() + return nil +} + +func TestTransportBlocksDisallowedAddresses(t *testing.T) { + server := newTestServer(t) + + // Capture the log output to check that the hint names the variables to change + var logs bytes.Buffer + previous := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil))) + t.Cleanup(func() { slog.SetDefault(previous) }) + + err := get(t, newTestClient(t, ""), server.URL) + + blockedErr, ok := errors.AsType[*BlockedError](err) + require.True(t, ok, "expected a BlockedError, got %v", err) + assert.Equal(t, PurposeSCIM, blockedErr.Purpose) + assert.Equal(t, "127.0.0.1", blockedErr.Addr.String()) + assert.Contains(t, err.Error(), "OUTBOUND_ALLOWED_HOSTS_SCIM") +} + +func TestTransportAllowsConfiguredAddresses(t *testing.T) { + server := newTestServer(t) + + require.NoError(t, get(t, newTestClient(t, "loopback"), server.URL)) + require.NoError(t, get(t, newTestClient(t, "127.0.0.1"), server.URL)) + require.NoError(t, get(t, newTestClient(t, "127.0.0.0/8"), server.URL)) +} + +func TestTransportAllowsConfiguredHostnames(t *testing.T) { + server := newTestServer(t) + localhostURL := strings.Replace(server.URL, "127.0.0.1", "localhost", 1) + + require.NoError(t, get(t, newTestClient(t, "localhost"), localhostURL)) + + // Trust is tied to the name, so the loopback IP itself stays blocked + err := get(t, newTestClient(t, "localhost"), server.URL) + _, ok := errors.AsType[*BlockedError](err) + require.True(t, ok, "expected a BlockedError, got %v", err) +} + +func TestTransportChecksRedirects(t *testing.T) { + server := newTestServer(t) + localhostURL := strings.Replace(server.URL, "127.0.0.1", "localhost", 1) + + // The first hop goes to a trusted name, the redirect to an address that isn't allowed + err := get(t, newTestClient(t, "localhost"), localhostURL+"/redirect") + _, ok := errors.AsType[*BlockedError](err) + require.True(t, ok, "expected a BlockedError, got %v", err) +} + +func TestTransportRejectsUnsupportedSchemes(t *testing.T) { + transport := NewTransport(nil, PurposeSCIM, Policy{}) + + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, "ftp://8.8.8.8/file", nil) + require.NoError(t, err) + + _, err = transport.RoundTrip(req) //nolint:bodyclose // No response is returned on error + require.ErrorContains(t, err, "unsupported URL scheme") +} diff --git a/backend/internal/service/oidc_service.go b/backend/internal/service/oidc_service.go index cfe10308..d6316a49 100644 --- a/backend/internal/service/oidc_service.go +++ b/backend/internal/service/oidc_service.go @@ -21,11 +21,11 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/apperror" "github.com/pocket-id/pocket-id/backend/internal/backchannellogout" - "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/model" datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" "github.com/pocket-id/pocket-id/backend/internal/oidc" + "github.com/pocket-id/pocket-id/backend/internal/outbound" "github.com/pocket-id/pocket-id/backend/internal/storage" "github.com/pocket-id/pocket-id/backend/internal/utils" imageutil "github.com/pocket-id/pocket-id/backend/internal/utils/image" @@ -959,23 +959,6 @@ func httpClientWithCheckRedirect(source *http.Client, checkRedirect func(req *ht return client } -// checkLogoURLAllowed prevents SSRF by allowing only URLs that resolve to public IPs -// URLs inside the icon library are exempt because the operator configured it, which lets a self-hosted mirror live on the local network -func checkLogoURLAllowed(ctx context.Context, u *url.URL) error { - if common.EnvConfig.IsIconLibraryURL(u) { - return nil - } - - private, err := utils.IsURLPrivate(ctx, u) - if err != nil { - return apperror.LogoDownloadFailed(err) - } else if private { - return apperror.InvalidLogoURL(errors.New("private IP addresses are not allowed")) - } - - return nil -} - func (s *OidcService) downloadAndSaveLogoFromURL(parentCtx context.Context, clientID string, raw string, light bool) error { u, err := url.Parse(raw) if err != nil { @@ -988,18 +971,11 @@ func (s *OidcService) downloadAndSaveLogoFromURL(parentCtx context.Context, clie ctx, cancel := context.WithTimeout(parentCtx, 15*time.Second) defer cancel() - err = checkLogoURLAllowed(ctx, u) - if err != nil { - return err - } - - // We need to check this on redirects too - client := httpClientWithCheckRedirect(s.httpClient, func(r *http.Request, via []*http.Request) error { + client := httpClientWithCheckRedirect(s.httpClient, func(_ *http.Request, via []*http.Request) error { if len(via) >= 10 { return apperror.InvalidLogoURL(errors.New("stopped after 10 redirects")) } - - return checkLogoURLAllowed(r.Context(), r.URL) + return nil }) req, err := http.NewRequestWithContext(ctx, http.MethodGet, raw, nil) @@ -1014,6 +990,9 @@ func (s *OidcService) downloadAndSaveLogoFromURL(parentCtx context.Context, clie if appErr, ok := errors.AsType[*apperror.Error](err); ok { return appErr } + if blockedErr, ok := errors.AsType[*outbound.BlockedError](err); ok { + return apperror.InvalidLogoURL(blockedErr) + } return apperror.LogoDownloadFailed(err) } defer resp.Body.Close() diff --git a/backend/internal/service/oidc_service_test.go b/backend/internal/service/oidc_service_test.go index d87231fd..3c70a73e 100644 --- a/backend/internal/service/oidc_service_test.go +++ b/backend/internal/service/oidc_service_test.go @@ -3,6 +3,7 @@ package service import ( "io" "net/http" + "net/http/httptest" "strconv" "strings" "testing" @@ -18,6 +19,7 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/model" datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" "github.com/pocket-id/pocket-id/backend/internal/oidc" + "github.com/pocket-id/pocket-id/backend/internal/outbound" "github.com/pocket-id/pocket-id/backend/internal/storage" "github.com/pocket-id/pocket-id/backend/internal/utils" testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing" @@ -415,37 +417,35 @@ func TestOidcService_downloadAndSaveLogoFromURL(t *testing.T) { require.True(t, apperror.IsCode(err, apperror.CodeValidationFailed)) }) - t.Run("Allows private hosts inside the icon library only", func(t *testing.T) { - const iconLibraryURL = "http://127.0.0.1:4050/icons" + t.Run("Allows the private icon library host only", func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "image/svg+xml") + _, _ = w.Write([]byte(``)) + })) + t.Cleanup(server.Close) + + iconLibraryURL := server.URL + "/icons" originalIconLibraryURL := common.EnvConfig.IconLibraryURL common.EnvConfig.IconLibraryURL = iconLibraryURL t.Cleanup(func() { common.EnvConfig.IconLibraryURL = originalIconLibraryURL }) - //nolint:bodyclose - svgResponse := testutils.NewMockResponse(http.StatusOK, ``) - svgResponse.Header.Set("Content-Type", "image/svg+xml") - + clients, err := outbound.New(&common.EnvConfig) + require.NoError(t, err) s := &OidcService{ db: db, fileStorage: dbStorage, - httpClient: &http.Client{ - Transport: &testutils.MockRoundTripper{ - Responses: map[string]*http.Response{ - iconLibraryURL + "/svg/nextcloud.svg": svgResponse, - }, - }, - }, + httpClient: clients.Client(outbound.PurposeClientLogo), } // The operator configured the library, so its loopback address is trusted - err := s.downloadAndSaveLogoFromURL(t.Context(), client.ID, iconLibraryURL+"/svg/nextcloud.svg", true) + err = s.downloadAndSaveLogoFromURL(t.Context(), client.ID, iconLibraryURL+"/svg/nextcloud.svg", true) require.NoError(t, err) require.True(t, fileExists(t, "oidc-client-images/"+client.ID+".svg")) - // Other paths on the same private host are still blocked - err = s.downloadAndSaveLogoFromURL(t.Context(), client.ID, "http://127.0.0.1:4050/admin/logo.svg", true) + // Other private hosts are still blocked, even when they resolve to the same address + err = s.downloadAndSaveLogoFromURL(t.Context(), client.ID, strings.Replace(server.URL, "127.0.0.1", "localhost", 1)+"/icons/svg/nextcloud.svg", true) require.Error(t, err) require.True(t, apperror.IsCode(err, apperror.CodeValidationFailed)) }) diff --git a/backend/internal/utils/ip_util.go b/backend/internal/utils/ip_util.go index 3683df93..9128c8c3 100644 --- a/backend/internal/utils/ip_util.go +++ b/backend/internal/utils/ip_util.go @@ -1,11 +1,8 @@ package utils import ( - "context" - "errors" "net" "net/netip" - "net/url" "strings" "github.com/pocket-id/pocket-id/backend/internal/common" @@ -28,13 +25,6 @@ var tailscaleIPNets = []*net.IPNet{ {IP: net.IPv4(100, 64, 0, 0), Mask: net.CIDRMask(10, 32)}, // 100.64.0.0/10 } -// LocalIPv6IPNets returns the extra IPv6 ranges configured via LOCAL_IPV6_RANGES -// that are treated as local/private. It is used to extend SSRF protection in -// components that classify IPs independently (e.g. the fosite CIMD fetcher). -func LocalIPv6IPNets() []*net.IPNet { - return localIPv6Ranges -} - func IsLocalIPv6(ip net.IP) bool { if ip.To4() != nil { return false @@ -77,23 +67,6 @@ func IsPrivateIP(ip net.IP) bool { return addr.IsLoopback() || addr.IsPrivate() || addr.IsLinkLocalUnicast() || addr.IsLinkLocalMulticast() || addr.IsUnspecified() } -func IsURLPrivate(ctx context.Context, u *url.URL) (bool, error) { - var r net.Resolver - ips, err := r.LookupIPAddr(ctx, u.Hostname()) - if err != nil || len(ips) == 0 { - return false, errors.New("cannot resolve hostname") - } - - // Prevents SSRF by allowing only public IPs - for _, addr := range ips { - if IsPrivateIP(addr.IP) { - return true, nil - } - } - - return false, nil -} - func listContainsIP(ipNets []*net.IPNet, ip net.IP) bool { for _, ipNet := range ipNets { if ipNet.Contains(ip) { diff --git a/backend/internal/utils/ip_util_test.go b/backend/internal/utils/ip_util_test.go index 05d52a6f..16b8a1b9 100644 --- a/backend/internal/utils/ip_util_test.go +++ b/backend/internal/utils/ip_util_test.go @@ -1,14 +1,10 @@ package utils import ( - "context" "net" - "net/url" "testing" - "time" "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" "github.com/pocket-id/pocket-id/backend/internal/common" ) @@ -167,214 +163,3 @@ func TestInit_LocalIPv6Ranges(t *testing.T) { assert.Len(t, localIPv6Ranges, 2) } - -func TestIsURLPrivate(t *testing.T) { - ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) - defer cancel() - - tests := []struct { - name string - urlStr string - expectPriv bool - expectError bool - }{ - { - name: "localhost by name", - urlStr: "http://localhost", - expectPriv: true, - expectError: false, - }, - { - name: "localhost with port", - urlStr: "http://localhost:8080", - expectPriv: true, - expectError: false, - }, - { - name: "127.0.0.1 IP", - urlStr: "http://127.0.0.1", - expectPriv: true, - expectError: false, - }, - { - name: "127.0.0.1 with port", - urlStr: "http://127.0.0.1:3000", - expectPriv: true, - expectError: false, - }, - { - name: "IPv6 loopback", - urlStr: "http://[::1]", - expectPriv: true, - expectError: false, - }, - { - name: "IPv6 loopback with port", - urlStr: "http://[::1]:8080", - expectPriv: true, - expectError: false, - }, - { - name: "private IP 10.x.x.x", - urlStr: "http://10.0.0.1", - expectPriv: true, - expectError: false, - }, - { - name: "private IP 192.168.x.x", - urlStr: "http://192.168.1.1", - expectPriv: true, - expectError: false, - }, - { - name: "private IP 172.16.x.x", - urlStr: "http://172.16.0.1", - expectPriv: true, - expectError: false, - }, - { - name: "Tailscale IP", - urlStr: "http://100.64.0.1", - expectPriv: true, - expectError: false, - }, - { - name: "IPv4 link-local metadata IP", - urlStr: "http://169.254.169.254", - expectPriv: true, - expectError: false, - }, - { - name: "IPv4-mapped link-local metadata IP", - urlStr: "http://[::ffff:169.254.169.254]", - expectPriv: true, - expectError: false, - }, - { - name: "IPv6 link-local IP", - urlStr: "http://[fe80::1]", - expectPriv: true, - expectError: false, - }, - { - name: "IPv4 unspecified IP", - urlStr: "http://0.0.0.0", - expectPriv: true, - expectError: false, - }, - { - name: "IPv6 unspecified IP", - urlStr: "http://[::]", - expectPriv: true, - expectError: false, - }, - { - name: "public IP - Google DNS", - urlStr: "http://8.8.8.8", - expectPriv: false, - expectError: false, - }, - { - name: "public IP - Cloudflare DNS", - urlStr: "http://1.1.1.1", - expectPriv: false, - expectError: false, - }, - { - name: "invalid hostname", - urlStr: "http://this-should-not-resolve-ever-123456789.invalid", - expectPriv: false, - expectError: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - u, err := url.Parse(tt.urlStr) - require.NoError(t, err, "Failed to parse URL %s", tt.urlStr) - - isPriv, err := IsURLPrivate(ctx, u) - - if tt.expectError { - require.Error(t, err, "IsURLPrivate(%s) expected error but got none", tt.urlStr) - } else { - require.NoError(t, err, "IsURLPrivate(%s) unexpected error", tt.urlStr) - assert.Equal(t, tt.expectPriv, isPriv, "IsURLPrivate(%s)", tt.urlStr) - } - }) - } -} - -func TestIsURLPrivate_WithDomainName(t *testing.T) { - // Note: These tests rely on actual DNS resolution - // They test real public domains to ensure they are not flagged as private - ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) - defer cancel() - - tests := []struct { - name string - urlStr string - expectPriv bool - }{ - { - name: "Google public domain", - urlStr: "https://www.google.com", - expectPriv: false, - }, - { - name: "GitHub public domain", - urlStr: "https://github.com", - expectPriv: false, - }, - { - // localhost.localtest.me is a well-known domain that resolves to 127.0.0.1 - name: "localhost.localtest.me resolves to 127.0.0.1", - urlStr: "http://localhost.localtest.me", - expectPriv: true, - }, - { - // 10.0.0.1.nip.io resolves to 10.0.0.1 (private IP) - name: "nip.io domain resolving to private 10.x IP", - urlStr: "http://10.0.0.1.nip.io", - expectPriv: true, - }, - { - // 192.168.1.1.nip.io resolves to 192.168.1.1 (private IP) - name: "nip.io domain resolving to private 192.168.x IP", - urlStr: "http://192.168.1.1.nip.io", - expectPriv: true, - }, - { - // 127.0.0.1.nip.io resolves to 127.0.0.1 (localhost) - name: "nip.io domain resolving to localhost", - urlStr: "http://127.0.0.1.nip.io", - expectPriv: true, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - u, err := url.Parse(tt.urlStr) - require.NoError(t, err, "Failed to parse URL %s", tt.urlStr) - - isPriv, err := IsURLPrivate(ctx, u) - if err != nil { - t.Skipf("DNS resolution failed for %s (network issue?): %v", tt.urlStr, err) - return - } - - assert.Equal(t, tt.expectPriv, isPriv, "IsURLPrivate(%s)", tt.urlStr) - }) - } -} - -func TestIsURLPrivate_ContextCancellation(t *testing.T) { - ctx, cancel := context.WithCancel(t.Context()) - cancel() // Cancel immediately - - u, err := url.Parse("http://example.com") - require.NoError(t, err, "Failed to parse URL") - - _, err = IsURLPrivate(ctx, u) - assert.Error(t, err, "IsURLPrivate with cancelled context expected error but got none") -} diff --git a/tests/setup/docker-compose.yml b/tests/setup/docker-compose.yml index 8512455d..bdf714eb 100644 --- a/tests/setup/docker-compose.yml +++ b/tests/setup/docker-compose.yml @@ -23,6 +23,7 @@ services: APP_ENV: test ENCRYPTION_KEY: test-encryption-key FILE_BACKEND: ${FILE_BACKEND} + OUTBOUND_ALLOWED_HOSTS_BACKCHANNEL_LOGOUT: private,loopback,host.docker.internal volumes: - pocket-id-test-data:/app/data build: