mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-10-10 03:39:05 +02:00
feat: add ability to override SSRF protection with OUTBOUND_ALLOWED_HOSTS_*
This commit is contained in:
Vendored
+3
@@ -4,5 +4,8 @@
|
||||
"oxc.fmt.disableNestedConfig": true,
|
||||
"[svelte]": {
|
||||
"editor.defaultFormatter": "oxc.oxc-vscode"
|
||||
},
|
||||
"[go]": {
|
||||
"editor.defaultFormatter": "golang.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)
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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":
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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]
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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()
|
||||
|
||||
@@ -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(`<svg xmlns="http://www.w3.org/2000/svg"/>`))
|
||||
}))
|
||||
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, `<svg xmlns="http://www.w3.org/2000/svg"/>`)
|
||||
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))
|
||||
})
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user