mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-10-02 07:49:04 +02:00
feat: add support Cloudflare location headers
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user