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

121 lines
4.1 KiB
Go

package server
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"testing"
"github.com/proxy-pool/proxy-pool/internal/config"
"github.com/proxy-pool/proxy-pool/internal/gateway/policy"
)
func TestBuildProtectionFromListenerConfig(t *testing.T) {
t.Parallel()
protection, err := BuildProtection(config.Listener{
Access: config.Access{AllowCIDRs: []string{"198.51.100.0/24"}},
Auth: config.Auth{
Mode: "usernamePassword",
Username: "client",
Password: "secret",
},
Limits: config.Limits{RequestsPerMinute: 10, RequestsPerMinutePerClient: 2},
})
if err != nil {
t.Fatalf("BuildProtection() error = %v", err)
}
request := httptest.NewRequest(http.MethodGet, "http://example.test", nil)
request.RemoteAddr = "198.51.100.8:1234"
request.Header.Set("Proxy-Authorization", "Basic Y2xpZW50OnNlY3JldA==")
for name, guard := range map[string]Guard{
"auth": protection.Auth, "access": protection.Access, "admission": protection.Admission,
} {
if err := guard.Check(context.Background(), request); err != nil {
t.Fatalf("%s guard error = %v", name, err)
}
}
}
func TestTargetPolicyFromOmittedConfigDefaultsToDeny(t *testing.T) {
t.Parallel()
targets, err := TargetPolicyFromListener(config.Listener{})
if err != nil {
t.Fatalf("TargetPolicyFromListener() error = %v", err)
}
if _, err := targets.EvaluateURL(context.Background(), "http://127.0.0.1/"); !errors.Is(err, policy.ErrTargetDenied) {
t.Fatalf("EvaluateURL(loopback) error = %v, want ErrTargetDenied", err)
}
if _, err := targets.EvaluateConnectAuthority(context.Background(), "8.8.8.8:22"); !errors.Is(err, policy.ErrTargetDenied) {
t.Fatalf("EvaluateConnectAuthority(port 22) error = %v, want ErrTargetDenied", err)
}
}
func TestTargetPolicyFromListenerMapsAllowedPorts(t *testing.T) {
t.Parallel()
targets, err := TargetPolicyFromListener(config.Listener{
DestinationPolicy: config.DestinationPolicy{AllowedPorts: []uint16{8443}},
})
if err != nil {
t.Fatalf("TargetPolicyFromListener() error = %v", err)
}
if _, err := targets.EvaluateConnectAuthority(context.Background(), "8.8.8.8:8443"); err != nil {
t.Fatalf("EvaluateConnectAuthority(port 8443) error = %v", err)
}
if _, err := targets.EvaluateConnectAuthority(context.Background(), "8.8.8.8:443"); !errors.Is(err, policy.ErrTargetDenied) {
t.Fatalf("EvaluateConnectAuthority(port 443) error = %v, want ErrTargetDenied", err)
}
}
func TestTargetPolicyFromPartialConfigKeepsOmittedCategoriesDenied(t *testing.T) {
t.Parallel()
deny := true
targets, err := TargetPolicyFromListener(config.Listener{
DestinationPolicy: config.DestinationPolicy{DenyPrivateNetworks: &deny},
})
if err != nil {
t.Fatalf("TargetPolicyFromListener() error = %v", err)
}
for _, target := range []string{"http://127.0.0.1/", "http://169.254.10.20/"} {
if _, err := targets.EvaluateURL(context.Background(), target); !errors.Is(err, policy.ErrTargetDenied) {
t.Fatalf("EvaluateURL(%q) error = %v, want ErrTargetDenied", target, err)
}
}
}
func TestTargetPolicyFromListenerAllowsOnlyExplicitlyFalseCategory(t *testing.T) {
t.Parallel()
allow := false
targets, err := TargetPolicyFromListener(config.Listener{
DestinationPolicy: config.DestinationPolicy{DenyLoopback: &allow},
})
if err != nil {
t.Fatalf("TargetPolicyFromListener() error = %v", err)
}
if _, err := targets.EvaluateURL(context.Background(), "http://127.0.0.1/"); err != nil {
t.Fatalf("EvaluateURL(loopback) error = %v", err)
}
for _, target := range []string{"http://10.0.0.1/", "http://169.254.10.20/"} {
if _, err := targets.EvaluateURL(context.Background(), target); !errors.Is(err, policy.ErrTargetDenied) {
t.Fatalf("EvaluateURL(%q) error = %v, want ErrTargetDenied", target, err)
}
}
}
func TestConfigFromListenerMapsRetryAndConcurrency(t *testing.T) {
t.Parallel()
result := ConfigFromListener(config.Listener{
Limits: config.Limits{MaxConcurrentConnections: 123},
Retry: config.Retry{MaxAttempts: 2, RetryMethods: []string{"GET", "HEAD"}},
})
if result.MaxConcurrentRequests != 123 || result.MaxAttempts != 2 || len(result.RetryMethods) != 2 {
t.Fatalf("handler config = %+v", result)
}
}