feat: add support Cloudflare location headers

This commit is contained in:
Elias Schneider
2026-10-01 20:53:38 +02:00
parent 02359840f2
commit 1a91eaa98c
15 changed files with 316 additions and 47 deletions
+2 -2
View File
@@ -103,8 +103,8 @@ func Bootstrap(ctx context.Context) error {
// Migrate the pre-actor signup tokens into their actors, once the actor host is ready
services = append(services, actorsReady.Await(svc.userSignUpModule.RunSignupTokenMigration))
// These services are only registered in non-test mode
if common.EnvConfig.AppEnv != "test" {
// Only the GeoLite provider needs a background database refresher
if !common.EnvConfig.AppEnv.IsTest() && svc.geoLiteModule != nil {
// Refresh the GeoLite database (this is cached per each replica)
services = append(services, svc.geoLiteModule.Run)
}
@@ -124,6 +124,9 @@ func shouldTraceRequest(r *http.Request) bool {
}
func registerGlobalMiddleware(r *gin.Engine) {
if common.EnvConfig.CloudflareLocationHeaders {
r.Use(middleware.CloudflareLocationMiddleware())
}
r.Use(middleware.HeadMiddleware())
r.Use(middleware.NewCacheControlMiddleware().Add())
r.Use(middleware.NewCorsMiddleware().Add())
@@ -21,10 +21,70 @@ import (
"github.com/gin-gonic/gin"
"github.com/pocket-id/pocket-id/backend/internal/apperror"
"github.com/pocket-id/pocket-id/backend/internal/common"
"github.com/pocket-id/pocket-id/backend/internal/geolite"
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
"github.com/pocket-id/pocket-id/backend/internal/middleware"
"github.com/stretchr/testify/require"
)
func TestCloudflareLocationMiddlewareOnlyWhenEnabled(t *testing.T) {
previousConfig := common.EnvConfig
t.Cleanup(func() { common.EnvConfig = previousConfig })
gin.SetMode(gin.TestMode)
for _, enabled := range []bool{false, true} {
name := "disabled"
if enabled {
name = "enabled"
}
t.Run(name, func(t *testing.T) {
common.EnvConfig.CloudflareLocationHeaders = enabled
locator, geoLiteModule, err := initIPLocationResolver(t.Context(), nil, &common.EnvConfigSchema{
GeoLiteDBPath: filepath.Join(t.TempDir(), "missing.mmdb"),
CloudflareLocationHeaders: enabled,
})
require.NoError(t, err)
if enabled {
require.IsType(t, iplocation.NewCloudflareResolver(), locator)
require.Nil(t, geoLiteModule)
} else {
require.IsType(t, &geolite.Module{}, locator)
require.Same(t, geoLiteModule, locator)
}
// Exercise the real global middleware with a Cloudflare client IP instead of the proxy's IP
router := gin.New()
router.TrustedPlatform = "CF-Connecting-IP"
registerGlobalMiddleware(router)
router.GET("/api/test", func(c *gin.Context) {
country, city, err := locator.GetLocationByIP(c.Request.Context(), c.ClientIP())
require.NoError(t, err)
c.JSON(http.StatusOK, gin.H{"ip": c.ClientIP(), "country": country, "city": city})
})
request := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/api/test", nil)
request.RemoteAddr = "192.168.1.1:1234"
request.Header.Set("cf-connecting-ip", "81.2.69.142")
request.Header.Set("cf-ipcountry", "CH")
request.Header.Set("cf-ipcity", "Zürich")
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, request)
require.Equal(t, http.StatusOK, recorder.Code)
var location map[string]string
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &location))
require.Equal(t, "81.2.69.142", location["ip"])
if enabled {
require.Equal(t, "Switzerland", location["country"])
require.Equal(t, "Zürich", location["city"])
} else {
require.Empty(t, location["country"])
require.Empty(t, location["city"])
}
})
}
}
func TestRequestLoggerUsesStructuredErrorMetadata(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -17,6 +17,7 @@ import (
"github.com/pocket-id/pocket-id/backend/internal/emailverification"
"github.com/pocket-id/pocket-id/backend/internal/environment"
"github.com/pocket-id/pocket-id/backend/internal/geolite"
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
"github.com/pocket-id/pocket-id/backend/internal/ldapsync"
"github.com/pocket-id/pocket-id/backend/internal/oidc"
"github.com/pocket-id/pocket-id/backend/internal/onetimeaccess"
@@ -33,6 +34,7 @@ type services struct {
appImagesService *service.AppImagesService
emailModule *email.Module
geoLiteModule *geolite.Module
ipLocator iplocation.Resolver
auditLogService *service.AuditLogService
jwtService *service.JwtService
userService *service.UserService
@@ -84,17 +86,13 @@ func initServices(
return nil, fmt.Errorf("failed to create email module: %w", err)
}
svc.geoLiteModule, err = geolite.New(ctx, geolite.Dependencies{
HTTPClient: httpClient,
DBPath: common.EnvConfig.GeoLiteDBPath,
DownloadURL: common.EnvConfig.GeoLiteDBUrl,
LicenseKey: common.EnvConfig.MaxMindLicenseKey,
})
// Select the location provider once so all consumers use the configured implementation
svc.ipLocator, svc.geoLiteModule, err = initIPLocationResolver(ctx, httpClient, &common.EnvConfig)
if err != nil {
return nil, fmt.Errorf("failed to create GeoLite module: %w", err)
return nil, fmt.Errorf("failed to create IP location resolver: %w", err)
}
svc.auditLogService = service.NewAuditLogService(db, svc.emailModule, svc.geoLiteModule, svc.appConfigService)
svc.auditLogService = service.NewAuditLogService(db, svc.emailModule, svc.ipLocator, svc.appConfigService)
svc.auditLogsModule, err = auditlogs.New(auditlogs.Dependencies{
DB: db,
Actors: actors,
@@ -132,7 +130,7 @@ func initServices(
Signer: svc.jwtService,
Reauth: svc.webauthnModule,
AuditLog: svc.auditLogService,
IPLocator: svc.geoLiteModule,
IPLocator: svc.ipLocator,
AppConfig: svc.appConfigService,
})
if err != nil {
@@ -262,3 +260,21 @@ func initServices(
return svc, nil
}
func initIPLocationResolver(ctx context.Context, httpClient *http.Client, config *common.EnvConfigSchema) (iplocation.Resolver, *geolite.Module, error) {
// Cloudflare locations need no database initialization or background refresh
if config.CloudflareLocationHeaders {
return iplocation.NewCloudflareResolver(), nil, nil
}
module, err := geolite.New(ctx, geolite.Dependencies{
HTTPClient: httpClient,
DBPath: config.GeoLiteDBPath,
DownloadURL: config.GeoLiteDBUrl,
LicenseKey: config.MaxMindLicenseKey,
})
if err != nil {
return nil, nil, err
}
return module, module, nil
}
+2
View File
@@ -95,6 +95,8 @@ type EnvConfigSchema struct {
MaxMindLicenseKey string `env:"MAXMIND_LICENSE_KEY" options:"file"`
GeoLiteDBPath string `env:"GEOLITE_DB_PATH"`
GeoLiteDBUrl string `env:"GEOLITE_DB_URL"`
// CloudflareLocationHeaders trusts location headers supplied by a Cloudflare proxy instead of using GeoLite
CloudflareLocationHeaders bool `env:"CLOUDFLARE_LOCATION_HEADERS"`
ActorsPort string `env:"ACTORS_PORT"`
ActorsHost string `env:"ACTORS_HOST" options:"toLower"`
@@ -121,6 +121,7 @@ func TestParseEnvConfig(t *testing.T) {
t.Setenv("PROXY_PROTOCOL", "true")
t.Setenv("ANALYTICS_DISABLED", "false")
t.Setenv("ALLOW_INSECURE_CALLBACK_URLS", "false")
t.Setenv("CLOUDFLARE_LOCATION_HEADERS", "true")
err := parseAndValidateEnvConfig(t)
require.NoError(t, err)
@@ -129,6 +130,7 @@ func TestParseEnvConfig(t *testing.T) {
assert.Equal(t, TrustProxyConfig{"0.0.0.0/0", "::/0"}, EnvConfig.ProxyProtocol)
assert.False(t, EnvConfig.AnalyticsDisabled)
assert.False(t, EnvConfig.AllowInsecureCallbackURLs)
assert.True(t, EnvConfig.CloudflareLocationHeaders)
})
t.Run("should parse trusted proxy IP addresses and CIDR ranges", func(t *testing.T) {
+2 -5
View File
@@ -11,6 +11,7 @@ import (
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
"github.com/pocket-id/pocket-id/backend/internal/model"
)
@@ -27,10 +28,6 @@ type AuditLogger interface {
DeviceStringFromUserAgent(userAgent string) string
}
type IPLocationResolver interface {
GetLocationByIP(ctx context.Context, ipAddress string) (country string, city string, err error)
}
type Dependencies struct {
DB *gorm.DB
Actors francishost.Host
@@ -39,7 +36,7 @@ type Dependencies struct {
Signer TokenService
Reauth ReauthenticationTokenConsumer
AuditLog AuditLogger
IPLocator IPLocationResolver
IPLocator iplocation.Resolver
AppConfig appconfig.AppConfigResolver
}
+3 -2
View File
@@ -13,6 +13,7 @@ import (
"github.com/pocket-id/pocket-id/backend/internal/apperror"
"github.com/pocket-id/pocket-id/backend/internal/dto"
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
"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/utils"
@@ -36,7 +37,7 @@ type Service struct {
signer TokenService
reauth ReauthenticationTokenConsumer
auditLog AuditLogger
ipLocator IPLocationResolver
ipLocator iplocation.Resolver
}
type VerificationInfo struct {
@@ -48,7 +49,7 @@ type VerificationInfo struct {
ExpiresAt datatype.DateTime
}
func NewService(actService *actor.Service, db *gorm.DB, signer TokenService, reauth ReauthenticationTokenConsumer, auditLog AuditLogger, ipLocator IPLocationResolver) *Service {
func NewService(actService *actor.Service, db *gorm.DB, signer TokenService, reauth ReauthenticationTokenConsumer, auditLog AuditLogger, ipLocator iplocation.Resolver) *Service {
return &Service{
actService: actService,
db: db,
+4 -17
View File
@@ -11,7 +11,7 @@ import (
"github.com/oschwald/maxminddb-golang/v2"
"github.com/pocket-id/pocket-id/backend/internal/utils"
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
)
// The GeoLite2 City database is kept on disk and memory-mapped (the format is optimized for random access)
@@ -19,9 +19,6 @@ import (
// The database file is considered cache, not state: it is a copy of a public artifact that any replica can rebuild on its own, so nothing is lost when a node goes away, and every replica keeps its own without needing to replicate anything
// It is also the supported way to supply a database by hand, which is what air-gapped deployments do: the file is watched, so replacing it takes effect without a restart
// internalNetworkCountry is reported for addresses that aren't routable on the public Internet
const internalNetworkCountry = "Internal Network"
// Service resolves IP addresses to locations, against a memory-mapped GeoLite2 City database
type Service struct {
log *slog.Logger
@@ -51,19 +48,9 @@ func (s *Service) GetLocationByIP(_ context.Context, ipAddress string) (country
return "", "", nil
}
// Check the IP address against known private IP ranges, which can be short-circuited
ip := net.ParseIP(ipAddress)
if ip != nil {
switch {
case utils.IsLocalIPv6(ip):
return internalNetworkCountry, "LAN", nil
case utils.IsTailscaleIP(ip):
return internalNetworkCountry, "Tailscale", nil
case utils.IsPrivateIP(ip):
return internalNetworkCountry, "LAN", nil
case utils.IsLocalhostIP(ip):
return internalNetworkCountry, "localhost", nil
}
// Keep local network labels consistent across location providers
if country, city := iplocation.LocalLocation(net.ParseIP(ipAddress)); country != "" {
return country, city, nil
}
addr, err := netip.ParseAddr(ipAddress)
+6 -6
View File
@@ -20,12 +20,12 @@ func TestServiceGetLocationByIPPrivateRanges(t *testing.T) {
city string
}{
{name: "empty address", ipAddress: ""},
{name: "private LAN IPv4", ipAddress: "192.168.1.20", country: internalNetworkCountry, city: "LAN"},
{name: "private LAN IPv4 in the 10/8 range", ipAddress: "10.4.5.6", country: internalNetworkCountry, city: "LAN"},
{name: "Tailscale IPv4", ipAddress: "100.101.102.103", country: internalNetworkCountry, city: "Tailscale"},
{name: "IPv6 unique local address", ipAddress: "fd00::1", country: internalNetworkCountry, city: "LAN"},
{name: "IPv4 loopback", ipAddress: "127.0.0.1", country: internalNetworkCountry, city: "LAN"},
{name: "IPv6 loopback", ipAddress: "::1", country: internalNetworkCountry, city: "LAN"},
{name: "private LAN IPv4", ipAddress: "192.168.1.20", country: "Internal Network", city: "LAN"},
{name: "private LAN IPv4 in the 10/8 range", ipAddress: "10.4.5.6", country: "Internal Network", city: "LAN"},
{name: "Tailscale IPv4", ipAddress: "100.101.102.103", country: "Internal Network", city: "Tailscale"},
{name: "IPv6 unique local address", ipAddress: "fd00::1", country: "Internal Network", city: "LAN"},
{name: "IPv4 loopback", ipAddress: "127.0.0.1", country: "Internal Network", city: "LAN"},
{name: "IPv6 loopback", ipAddress: "::1", country: "Internal Network", city: "LAN"},
}
for _, tt := range tests {
+66
View File
@@ -0,0 +1,66 @@
package iplocation
import (
"context"
"fmt"
"net"
"net/netip"
"strings"
"golang.org/x/text/language"
"golang.org/x/text/language/display"
)
type cloudflareLocationKey struct{}
type cloudflareLocation struct {
ipAddress string
country string
city string
}
// WithCloudflareLocation captures the client's location while the originating HTTP request is available
// Callers must only use this when the deployment explicitly trusts Cloudflare's location headers
func WithCloudflareLocation(ctx context.Context, ipAddress, countryCode, city string) context.Context {
location := cloudflareLocation{ipAddress: ipAddress, city: strings.TrimSpace(city)}
// Cloudflare supplies an ISO country code, while audit logs and notification emails display English country names
code := strings.ToUpper(strings.TrimSpace(countryCode))
if len(code) == 2 {
region, err := language.ParseRegion(code)
if err == nil && region.IsCountry() {
location.country = display.English.Regions().Name(region)
}
}
return context.WithValue(ctx, cloudflareLocationKey{}, location)
}
// CloudflareResolver resolves locations using trusted Cloudflare request headers
type CloudflareResolver struct{}
func NewCloudflareResolver() *CloudflareResolver {
return &CloudflareResolver{}
}
// GetLocationByIP returns the country and city for the client making the current request
func (r *CloudflareResolver) GetLocationByIP(ctx context.Context, ipAddress string) (country string, city string, err error) {
if ipAddress == "" {
return "", "", nil
}
// Keep local network labels consistent across location providers
if country, city := LocalLocation(net.ParseIP(ipAddress)); country != "" {
return country, city, nil
}
if _, err := netip.ParseAddr(ipAddress); err != nil {
return "", "", fmt.Errorf("failed to parse IP address: %w", err)
}
// Headers describe only the request's client, never another IP being inspected
location, ok := ctx.Value(cloudflareLocationKey{}).(cloudflareLocation)
if !ok || location.ipAddress != ipAddress {
return "", "", nil
}
return location.country, location.city, nil
}
@@ -0,0 +1,89 @@
package iplocation
import (
"testing"
"github.com/stretchr/testify/require"
)
func TestCloudflareLocation(t *testing.T) {
resolver := NewCloudflareResolver()
tests := []struct {
name string
countryCode string
city string
country string
wantCity string
}{
{name: "country and city", countryCode: "CH", city: "Zürich", country: "Switzerland", wantCity: "Zürich"},
{name: "country only", countryCode: "US", country: "United States"},
{name: "city only", city: "London", wantCity: "London"},
{name: "normalized headers", countryCode: " gb ", city: " London ", country: "United Kingdom", wantCity: "London"},
{name: "missing headers"},
{name: "unknown country", countryCode: "XX"},
{name: "Tor country", countryCode: "T1"},
{name: "unspecified country", countryCode: "ZZ"},
{name: "invalid country", countryCode: "invalid"},
{name: "continent is not a country", countryCode: "EU"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ctx := WithCloudflareLocation(t.Context(), "81.2.69.142", tt.countryCode, tt.city)
country, city, err := resolver.GetLocationByIP(ctx, "81.2.69.142")
require.NoError(t, err)
require.Equal(t, tt.country, country)
require.Equal(t, tt.wantCity, city)
})
}
}
func TestCloudflareLocationOnlyDescribesRequestClient(t *testing.T) {
resolver := NewCloudflareResolver()
ctx := WithCloudflareLocation(t.Context(), "81.2.69.142", "GB", "London")
// Looking up another device must never reuse the approving device's headers
country, city, err := resolver.GetLocationByIP(ctx, "216.160.83.56")
require.NoError(t, err)
require.Empty(t, country)
require.Empty(t, city)
// Non-request contexts have no header location to resolve
country, city, err = resolver.GetLocationByIP(t.Context(), "81.2.69.142")
require.NoError(t, err)
require.Empty(t, country)
require.Empty(t, city)
}
func TestCloudflareLocationPreservesIPHandling(t *testing.T) {
resolver := NewCloudflareResolver()
for _, ip := range []string{"192.168.1.20", "100.101.102.103", "fd00::1", "127.0.0.1"} {
t.Run(ip, func(t *testing.T) {
ctx := WithCloudflareLocation(t.Context(), ip, "CH", "Zürich")
country, city, err := resolver.GetLocationByIP(ctx, ip)
require.NoError(t, err)
require.Equal(t, "Internal Network", country)
if ip == "100.101.102.103" {
require.Equal(t, "Tailscale", city)
} else {
require.Equal(t, "LAN", city)
}
})
}
ctx := WithCloudflareLocation(t.Context(), "2001:218::1", "JP", "Tokyo")
country, city, err := resolver.GetLocationByIP(ctx, "2001:218::1")
require.NoError(t, err)
require.Equal(t, "Japan", country)
require.Equal(t, "Tokyo", city)
_, _, err = resolver.GetLocationByIP(ctx, "not-an-ip")
require.ErrorContains(t, err, "failed to parse IP address")
country, city, err = resolver.GetLocationByIP(ctx, "")
require.NoError(t, err)
require.Empty(t, country)
require.Empty(t, city)
}
+32
View File
@@ -0,0 +1,32 @@
package iplocation
import (
"context"
"net"
"github.com/pocket-id/pocket-id/backend/internal/utils"
)
type Resolver interface {
GetLocationByIP(ctx context.Context, ipAddress string) (country string, city string, err error)
}
// LocalLocation returns the labels shared by all providers for non-public addresses
func LocalLocation(ip net.IP) (country string, city string) {
if ip == nil {
return "", ""
}
switch {
case utils.IsLocalIPv6(ip):
return "Internal Network", "LAN"
case utils.IsTailscaleIP(ip):
return "Internal Network", "Tailscale"
case utils.IsPrivateIP(ip):
return "Internal Network", "LAN"
case utils.IsLocalhostIP(ip):
return "Internal Network", "localhost"
default:
return "", ""
}
}
@@ -0,0 +1,17 @@
package middleware
import (
"github.com/gin-gonic/gin"
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
)
// CloudflareLocationMiddleware captures location headers for services that only receive the request context
// Register it only when the deployment explicitly trusts Cloudflare's location headers
func CloudflareLocationMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
ctx := iplocation.WithCloudflareLocation(c.Request.Context(), c.ClientIP(), c.GetHeader("CF-IPCountry"), c.GetHeader("CF-IPCity"))
c.Request = c.Request.WithContext(ctx)
c.Next()
}
}
@@ -8,6 +8,7 @@ import (
userAgentParser "github.com/mileusna/useragent"
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
"github.com/pocket-id/pocket-id/backend/internal/model"
"github.com/pocket-id/pocket-id/backend/internal/utils"
"gorm.io/gorm"
@@ -17,18 +18,14 @@ type NewLoginEmailSender interface {
SendNewLogin(ctx context.Context, dbConfig *appconfig.AppConfigModel, userFullName, userEmail, ipAddress, country, city, device string, dateTime time.Time) error
}
type IPLocationResolver interface {
GetLocationByIP(ctx context.Context, ipAddress string) (country string, city string, err error)
}
type AuditLogService struct {
db *gorm.DB
emailSender NewLoginEmailSender
ipLocator IPLocationResolver
ipLocator iplocation.Resolver
appConfigService *appconfig.AppConfigService
}
func NewAuditLogService(db *gorm.DB, emailSender NewLoginEmailSender, ipLocator IPLocationResolver, appConfigService *appconfig.AppConfigService) *AuditLogService {
func NewAuditLogService(db *gorm.DB, emailSender NewLoginEmailSender, ipLocator iplocation.Resolver, appConfigService *appconfig.AppConfigService) *AuditLogService {
return &AuditLogService{
db: db,
emailSender: emailSender,