proxy-pool/internal/gateway/server/protection_test.go

140 lines
4.3 KiB
Go

package server
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"net/netip"
"testing"
)
func TestBasicAuthGuardUsesProxyAuthorizationAndReturnsChallenge(t *testing.T) {
t.Parallel()
guard := NewBasicAuthGuard("client", "secret")
allowed := httptest.NewRequest(http.MethodGet, "http://example.test", nil)
allowed.SetBasicAuth("ignored", "ignored")
allowed.Header.Set("Proxy-Authorization", "Basic Y2xpZW50OnNlY3JldA==")
if err := guard.Check(context.Background(), allowed); err != nil {
t.Fatalf("Check(valid) error = %v", err)
}
denied := httptest.NewRequest(http.MethodGet, "http://example.test", nil)
denied.Header.Set("Proxy-Authorization", "Basic Y2xpZW50Ondyb25n")
err := guard.Check(context.Background(), denied)
var httpError *HTTPError
if !errors.As(err, &httpError) {
t.Fatalf("Check(invalid) error = %T %v, want *HTTPError", err, err)
}
if httpError.StatusCode != http.StatusProxyAuthRequired {
t.Fatalf("status = %d, want 407", httpError.StatusCode)
}
if got := httpError.Header.Get("Proxy-Authenticate"); got != `Basic realm="proxy"` {
t.Fatalf("challenge = %q", got)
}
}
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")
}