From 1a91eaa98c9f6839770c54a93f5f5d0eafb1f290 Mon Sep 17 00:00:00 2001 From: Elias Schneider Date: Thu, 1 Oct 2026 20:51:33 +0200 Subject: [PATCH] feat: add support Cloudflare location headers --- backend/internal/bootstrap/bootstrap.go | 4 +- .../internal/bootstrap/router_bootstrap.go | 3 + .../bootstrap/router_bootstrap_test.go | 60 +++++++++++++ .../internal/bootstrap/services_bootstrap.go | 34 +++++-- backend/internal/common/env_config.go | 2 + backend/internal/common/env_config_test.go | 2 + backend/internal/devicelogin/module.go | 7 +- backend/internal/devicelogin/service.go | 5 +- backend/internal/geolite/service.go | 21 +---- backend/internal/geolite/service_test.go | 12 +-- backend/internal/iplocation/cloudflare.go | 66 ++++++++++++++ .../internal/iplocation/cloudflare_test.go | 89 +++++++++++++++++++ backend/internal/iplocation/resolver.go | 32 +++++++ .../middleware/cloudflare_location.go | 17 ++++ backend/internal/service/audit_log_service.go | 9 +- 15 files changed, 316 insertions(+), 47 deletions(-) create mode 100644 backend/internal/iplocation/cloudflare.go create mode 100644 backend/internal/iplocation/cloudflare_test.go create mode 100644 backend/internal/iplocation/resolver.go create mode 100644 backend/internal/middleware/cloudflare_location.go diff --git a/backend/internal/bootstrap/bootstrap.go b/backend/internal/bootstrap/bootstrap.go index ce5508c4..2d2122fe 100644 --- a/backend/internal/bootstrap/bootstrap.go +++ b/backend/internal/bootstrap/bootstrap.go @@ -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) } diff --git a/backend/internal/bootstrap/router_bootstrap.go b/backend/internal/bootstrap/router_bootstrap.go index 5afed82a..fa222ce0 100644 --- a/backend/internal/bootstrap/router_bootstrap.go +++ b/backend/internal/bootstrap/router_bootstrap.go @@ -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()) diff --git a/backend/internal/bootstrap/router_bootstrap_test.go b/backend/internal/bootstrap/router_bootstrap_test.go index ee180ce0..39713174 100644 --- a/backend/internal/bootstrap/router_bootstrap_test.go +++ b/backend/internal/bootstrap/router_bootstrap_test.go @@ -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) diff --git a/backend/internal/bootstrap/services_bootstrap.go b/backend/internal/bootstrap/services_bootstrap.go index 5c392dec..272c0810 100644 --- a/backend/internal/bootstrap/services_bootstrap.go +++ b/backend/internal/bootstrap/services_bootstrap.go @@ -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 +} diff --git a/backend/internal/common/env_config.go b/backend/internal/common/env_config.go index beaa6b6a..a3e8eef8 100644 --- a/backend/internal/common/env_config.go +++ b/backend/internal/common/env_config.go @@ -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"` diff --git a/backend/internal/common/env_config_test.go b/backend/internal/common/env_config_test.go index 961c0a40..6845016d 100644 --- a/backend/internal/common/env_config_test.go +++ b/backend/internal/common/env_config_test.go @@ -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) { diff --git a/backend/internal/devicelogin/module.go b/backend/internal/devicelogin/module.go index bfb5f79f..bf0875ba 100644 --- a/backend/internal/devicelogin/module.go +++ b/backend/internal/devicelogin/module.go @@ -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 } diff --git a/backend/internal/devicelogin/service.go b/backend/internal/devicelogin/service.go index 61697d43..76a2b4af 100644 --- a/backend/internal/devicelogin/service.go +++ b/backend/internal/devicelogin/service.go @@ -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, diff --git a/backend/internal/geolite/service.go b/backend/internal/geolite/service.go index 61d83732..0f5da58f 100644 --- a/backend/internal/geolite/service.go +++ b/backend/internal/geolite/service.go @@ -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) diff --git a/backend/internal/geolite/service_test.go b/backend/internal/geolite/service_test.go index e238bc6e..c57dd791 100644 --- a/backend/internal/geolite/service_test.go +++ b/backend/internal/geolite/service_test.go @@ -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 { diff --git a/backend/internal/iplocation/cloudflare.go b/backend/internal/iplocation/cloudflare.go new file mode 100644 index 00000000..48e502f2 --- /dev/null +++ b/backend/internal/iplocation/cloudflare.go @@ -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 +} diff --git a/backend/internal/iplocation/cloudflare_test.go b/backend/internal/iplocation/cloudflare_test.go new file mode 100644 index 00000000..eadfd768 --- /dev/null +++ b/backend/internal/iplocation/cloudflare_test.go @@ -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) +} diff --git a/backend/internal/iplocation/resolver.go b/backend/internal/iplocation/resolver.go new file mode 100644 index 00000000..b480152f --- /dev/null +++ b/backend/internal/iplocation/resolver.go @@ -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 "", "" + } +} diff --git a/backend/internal/middleware/cloudflare_location.go b/backend/internal/middleware/cloudflare_location.go new file mode 100644 index 00000000..3fbabbec --- /dev/null +++ b/backend/internal/middleware/cloudflare_location.go @@ -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() + } +} diff --git a/backend/internal/service/audit_log_service.go b/backend/internal/service/audit_log_service.go index 2b23aaaf..f475ed9a 100644 --- a/backend/internal/service/audit_log_service.go +++ b/backend/internal/service/audit_log_service.go @@ -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,