140 lines
4.3 KiB
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")
|
|
}
|