diff --git a/geoblock.go b/geoblock.go index 3a14c34..e9800ae 100755 --- a/geoblock.go +++ b/geoblock.go @@ -114,7 +114,10 @@ func New(ctx context.Context, next http.Handler, config *Config, name string) (h return nil, err } - allowedIPAddresses, allowedIPRanges := parseAllowedIPAddresses(config.AllowedIPAddresses, infoLogger) + allowedIPAddresses, allowedIPRanges, err := parseAllowedIPAddresses(config.AllowedIPAddresses) + if err != nil { + return nil, err + } excludedPathRegexps, err := compileExcludedPathPatterns(config.ExcludedPathPatterns) if err != nil { @@ -271,6 +274,11 @@ func (a *GeoBlock) ServeHTTP(rw http.ResponseWriter, req *http.Request) { // Only keep the first IP address (should be the client, if the proxy behaves itself) // so we can check whether it is allowed or denied. if a.xForwardedForReverseProxy { + if len(requestIPAddresses) == 0 { + a.infoLogger.Printf("%s: no client IP found in request headers", a.name) + rw.WriteHeader(http.StatusForbidden) + return + } requestIPAddresses = requestIPAddresses[:1] } @@ -383,19 +391,7 @@ func (a *GeoBlock) allowDenyCachedRequestIP(requestIPAddr *net.IP, req *http.Req if !cacheHit { entry, err = a.createNewIPEntry(req, ipAddressString) if err != nil { - if a.ignoreAPIFailures { - a.infoLogger.Printf("%s: request allowed [%s] due to API failure", a.name, requestIPAddr) - return true, "" - } - - if os.IsTimeout(err) && a.ignoreAPITimeout { - a.infoLogger.Printf("%s: request allowed [%s] due to API timeout", a.name, requestIPAddr) - // TODO: this was previously an immediate response to the client - return true, "" - } - - a.infoLogger.Printf("%s: request denied [%s] due to error: %s", a.name, requestIPAddr, err) - return false, "" + return a.handleLookupError(err, requestIPAddr), "" } } else { entry = cacheEntry.(ipEntry) @@ -411,12 +407,7 @@ func (a *GeoBlock) allowDenyCachedRequestIP(requestIPAddr *net.IP, req *http.Req if a.shouldRefreshEntry(entry) { entry, err = a.createNewIPEntry(req, ipAddressString) if err != nil { - if a.ignoreAPIFailures { - a.infoLogger.Printf("%s: request allowed [%s] due to API failure", a.name, requestIPAddr) - return true, "" - } - a.infoLogger.Printf("%s: request denied [%s] due to error: %s", a.name, requestIPAddr, err) - return false, "" + return a.handleLookupError(err, requestIPAddr), "" } } @@ -460,6 +451,24 @@ func (a *GeoBlock) allowDenyCachedRequestIP(requestIPAddr *net.IP, req *http.Req return true, entry.Country } +// handleLookupError applies the configured fail-open policy to a failed +// country lookup: any API failure allows the request when ignoreApiFailures is +// set, a timeout allows it when ignoreApiTimeout is set, anything else denies. +func (a *GeoBlock) handleLookupError(err error, requestIPAddr *net.IP) bool { + if a.ignoreAPIFailures { + a.infoLogger.Printf("%s: request allowed [%s] due to API failure", a.name, requestIPAddr) + return true + } + + if os.IsTimeout(err) && a.ignoreAPITimeout { + a.infoLogger.Printf("%s: request allowed [%s] due to API timeout", a.name, requestIPAddr) + return true + } + + a.infoLogger.Printf("%s: request denied [%s] due to error: %s", a.name, requestIPAddr, err) + return false +} + func (a *GeoBlock) cachedRequestIP(requestIPAddr *net.IP, req *http.Request) (bool, string) { ipAddressString := requestIPAddr.String() cacheEntry, ok := a.database.Get(ipAddressString) @@ -729,7 +738,7 @@ func getHTTPStatusCodeDeniedRequest(code int) (int, error) { return defaultDeniedRequestHTTPStatusCode, nil } -func parseAllowedIPAddresses(entries []string, logger *log.Logger) ([]net.IP, []*net.IPNet) { +func parseAllowedIPAddresses(entries []string) ([]net.IP, []*net.IPNet, error) { var allowedIPAddresses []net.IP var allowedIPRanges []*net.IPNet @@ -746,12 +755,12 @@ func parseAllowedIPAddresses(entries []string, logger *log.Logger) ([]net.IP, [] // Attempt to parse as a single IP address ipAddress := net.ParseIP(ipAddressEntry) if ipAddress == nil { - logger.Fatal("Invalid IP address provided:", ipAddressEntry) + return nil, nil, fmt.Errorf("invalid allowed IP address or CIDR range [%s]", ipAddressEntry) } allowedIPAddresses = append(allowedIPAddresses, ipAddress) } - return allowedIPAddresses, allowedIPRanges + return allowedIPAddresses, allowedIPRanges, nil } func compileExcludedPathPatterns(patterns []string) ([]*regexp.Regexp, error) { @@ -786,6 +795,8 @@ func printConfiguration(name string, config *Config, logger *log.Logger) { logger.Printf("%s: API uri: %s", name, config.API) logger.Printf("%s: API timeout: %d", name, config.APITimeoutMs) logger.Printf("%s: ignore API timeout: %t", name, config.IgnoreAPITimeout) + logger.Printf("%s: ignore API failures: %t", name, config.IgnoreAPIFailures) + logger.Printf("%s: X-Forwarded-For reverse proxy: %t", name, config.XForwardedForReverseProxy) logger.Printf("%s: cache size: %d", name, config.CacheSize) logger.Printf("%s: cache ttl seconds: %d", name, config.CacheTTLSeconds) logger.Printf("%s: force monthly update: %t", name, config.ForceMonthlyUpdate) @@ -794,8 +805,10 @@ func printConfiguration(name string, config *Config, logger *log.Logger) { logger.Printf("%s: blacklist mode: %t", name, config.BlackListMode) logger.Printf("%s: add country header: %t", name, config.AddCountryHeader) logger.Printf("%s: countries: %v", name, config.Countries) + logger.Printf("%s: allowed IP addresses: %v", name, config.AllowedIPAddresses) logger.Printf("%s: Denied request status code: %d", name, config.HTTPStatusCodeDeniedRequest) logger.Printf("%s: Log file path: %s", name, config.LogFilePath) + logger.Printf("%s: IP database cache path: %s", name, config.IPDatabaseCachePath) if len(config.RedirectURLIfDenied) != 0 { logger.Printf("%s: Redirect URL on denied requests: %s", name, config.RedirectURLIfDenied) } diff --git a/geoblock_test.go b/geoblock_test.go index e76b3c7..6bcde5c 100755 --- a/geoblock_test.go +++ b/geoblock_test.go @@ -1981,6 +1981,90 @@ func TestCacheTTLExpiresWithoutForceMonthlyUpdate(t *testing.T) { } } +func TestInvalidAllowedIPAddressEntry(t *testing.T) { + cfg := createTesterConfig() + cfg.Countries = append(cfg.Countries, "CH") + cfg.AllowedIPAddresses = append(cfg.AllowedIPAddresses, "not-an-ip") + + ctx := context.Background() + next := http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) {}) + + _, err := geoblock.New(ctx, next, cfg, t.Name()) + + // expect error + if err == nil { + t.Fatal("invalid allowedIPAddresses entry accepted") + } +} + +func TestReverseProxyModeWithoutForwardedHeaders(t *testing.T) { + cfg := createTesterConfig() + cfg.Countries = append(cfg.Countries, "CH") + cfg.XForwardedForReverseProxy = true + + ctx := context.Background() + next := http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) {}) + + handler, err := geoblock.New(ctx, next, cfg, t.Name()) + if err != nil { + t.Fatal(err) + } + + rec := httptest.NewRecorder() + // Neither X-Forwarded-For nor X-Real-IP is set. + req := httptest.NewRequest(http.MethodGet, "http://localhost", nil) + + handler.ServeHTTP(rec, req) + + assertStatusCode(t, rec.Result(), http.StatusForbidden) +} + +func TestTimeoutOnCacheRefresh_AllowWhenIgnoreTimeoutTrue(t *testing.T) { + // Stub server that answers quickly at first (to prime the cache) and can be + // switched to respond slower than the client timeout for the refresh. + var timeoutMode atomic.Bool + apiStub := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + if timeoutMode.Load() { + time.Sleep(50 * time.Millisecond) // > APITimeoutMs below + } + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("CH")) + })) + defer apiStub.Close() + + cfg := createTesterConfig() + cfg.API = apiStub.URL + "/{ip}" + cfg.Countries = append(cfg.Countries, "CH") + cfg.APITimeoutMs = 5 // 5ms client timeout + cfg.IgnoreAPITimeout = true // timeouts should ALLOW, also on cache refresh + cfg.CacheTTLSeconds = 1 + + ctx := context.Background() + next := http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) {}) + + handler, err := geoblock.New(ctx, next, cfg, t.Name()) + if err != nil { + t.Fatal(err) + } + + // Prime the cache from the fast stub. + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "http://localhost", nil) + req.Header.Add(xForwardedFor, chExampleIP) + handler.ServeHTTP(rec, req) + assertStatusCode(t, rec.Result(), http.StatusOK) + + // Let the entry outlive its TTL, then time out on the refresh lookup. + timeoutMode.Store(true) + time.Sleep(1100 * time.Millisecond) + + rec = httptest.NewRecorder() + req = httptest.NewRequest(http.MethodGet, "http://localhost", nil) + req.Header.Add(xForwardedFor, chExampleIP) + handler.ServeHTTP(rec, req) + assertStatusCode(t, rec.Result(), http.StatusOK) +} + func createTesterConfig() *geoblock.Config { cfg := geoblock.CreateConfig() diff --git a/readme.md b/readme.md index b6c975e..4a21553 100755 --- a/readme.md +++ b/readme.md @@ -542,13 +542,13 @@ Allows customizing the HTTP status code returned if the request was denied. Allows to define a target for the logs of the middleware. The path must look like the following: `logFilePath: "/log/geoblock.log"`. Make sure the folder is writeable. -### Define a custom log file `XForwardedForReverseProxy` +### Use only the first `X-Forwarded-For` IP address `xForwardedForReverseProxy` Basically tells GeoBlock to only allow/deny a request based on the first IP address in the X-ForwardedFor HTTP header. This is useful for servers behind e.g. a Cloudflare proxy. -### Define a custom log file `redirectUrlIfDenied` +### Redirect denied requests `redirectUrlIfDenied` -Allows returning a HTTP 301 status code, which indicates that the requested resource has been moved. The URL which can be specified is used to redirect the client to. So instead of "blocking" the client, the client will be redirected to the configured URL. +Allows redirecting denied requests with a HTTP 302 (Found) status code to the configured URL. So instead of "blocking" the client, the client will be redirected to the configured URL. ### Excluded Path Patterns `excludedPathPatterns`