package server import ( "context" "errors" "net/http" "net/http/httptest" "net/netip" "testing" ) func TestClientIPResolverOnlyTrustsForwardedChainFromTrustedPeer(t *testing.T) { t.Parallel() resolver, err := NewClientIPResolver([]string{"10.0.0.0/8", "192.0.2.0/24"}) if err != nil { t.Fatalf("NewClientIPResolver() error = %v", err) } request := httptest.NewRequest(http.MethodGet, "http://example.test", nil) request.RemoteAddr = "10.0.0.9:1234" request.Header.Set("X-Forwarded-For", "198.51.100.7, 192.0.2.5") address, err := resolver.Resolve(request) if err != nil { t.Fatalf("Resolve(trusted) error = %v", err) } if address != netip.MustParseAddr("198.51.100.7") { t.Fatalf("trusted forwarded address = %s", address) } request.RemoteAddr = "203.0.113.9:1234" address, err = resolver.Resolve(request) if err != nil { t.Fatalf("Resolve(untrusted) error = %v", err) } if address != netip.MustParseAddr("203.0.113.9") { t.Fatalf("untrusted forwarded address = %s", address) } } func TestClientIPResolverSupportsRFCForwardedIPv4AndIPv6(t *testing.T) { t.Parallel() resolver, err := NewClientIPResolver([]string{"10.0.0.0/8", "192.0.2.0/24"}) if err != nil { t.Fatalf("NewClientIPResolver() error = %v", err) } request := httptest.NewRequest(http.MethodGet, "http://example.test", nil) request.RemoteAddr = "10.0.0.9:1234" request.Header.Set("Forwarded", `for="[2001:db8::7]:4711";proto=https, for=192.0.2.5`) address, err := resolver.Resolve(request) if err != nil { t.Fatalf("Resolve() error = %v", err) } if address != netip.MustParseAddr("2001:db8::7") { t.Fatalf("Forwarded address = %s", address) } } func TestAccessAndAdmissionGuardsShareResolvedClientIdentity(t *testing.T) { t.Parallel() resolver, err := NewClientIPResolver(nil) if err != nil { t.Fatalf("NewClientIPResolver() error = %v", err) } access, err := NewAccessGuard(resolver, []string{"198.51.100.0/24"}) if err != nil { t.Fatalf("NewAccessGuard() error = %v", err) } admitter := &recordingAdmitter{} admission := NewAdmissionGuard(resolver, admitter) request := httptest.NewRequest(http.MethodGet, "http://example.test", nil) request.RemoteAddr = "198.51.100.8:1234" if err := access.Check(context.Background(), request); err != nil { t.Fatalf("access.Check() error = %v", err) } if err := admission.Check(context.Background(), request); err != nil { t.Fatalf("admission.Check() error = %v", err) } if admitter.key != "198.51.100.8" { t.Fatalf("admission key = %q", admitter.key) } } func TestAdmissionGuardMapsLimiterRejectionTo429(t *testing.T) { t.Parallel() resolver, _ := NewClientIPResolver(nil) guard := NewAdmissionGuard(resolver, rejectingAdmitter{}) request := httptest.NewRequest(http.MethodGet, "http://example.test", nil) request.RemoteAddr = "198.51.100.8:1234" err := guard.Check(context.Background(), request) var httpError *HTTPError if !errors.As(err, &httpError) || httpError.StatusCode != http.StatusTooManyRequests { t.Fatalf("Check() error = %T %v, want HTTP 429", err, err) } } type recordingAdmitter struct{ key string } func (admitter *recordingAdmitter) Admit(_ context.Context, key string) error { admitter.key = key return nil } type rejectingAdmitter struct{} func (rejectingAdmitter) Admit(context.Context, string) error { return errors.New("rate limited") }