Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
59 changes: 36 additions & 23 deletions geoblock.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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]
}

Expand Down Expand Up @@ -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)
Expand All @@ -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), ""
}
}

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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

Expand All @@ -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) {
Expand Down Expand Up @@ -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)
Expand All @@ -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)
}
Expand Down
84 changes: 84 additions & 0 deletions geoblock_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down
6 changes: 3 additions & 3 deletions readme.md
Original file line number Diff line number Diff line change
Expand Up @@ -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`

Expand Down
Loading