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: