mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-04 12:39:06 +02:00
merge main
This commit is contained in:
+1
-1
@@ -10,7 +10,7 @@ FROM gcr.io/distroless/base:debug
|
||||
COPY netbird-proxy /go/bin/netbird-proxy
|
||||
COPY --from=builder /tmp/passwd /etc/passwd
|
||||
COPY --from=builder /tmp/group /etc/group
|
||||
COPY --from=builder /tmp/var/lib/netbird /var/lib/netbird
|
||||
COPY --from=builder --chown=1000:1000 /tmp/var/lib/netbird /var/lib/netbird
|
||||
COPY --from=builder --chown=1000:1000 --chmod=755 /tmp/certs /certs
|
||||
USER netbird:netbird
|
||||
ENV HOME=/var/lib/netbird
|
||||
|
||||
@@ -28,7 +28,7 @@ FROM gcr.io/distroless/base:debug
|
||||
COPY --from=builder /app/netbird-proxy /usr/bin/netbird-proxy
|
||||
COPY --from=builder /tmp/passwd /etc/passwd
|
||||
COPY --from=builder /tmp/group /etc/group
|
||||
COPY --from=builder /tmp/var/lib/netbird /var/lib/netbird
|
||||
COPY --from=builder --chown=1000:1000 /tmp/var/lib/netbird /var/lib/netbird
|
||||
COPY --from=builder --chown=1000:1000 --chmod=755 /tmp/certs /certs
|
||||
USER netbird:netbird
|
||||
ENV HOME=/var/lib/netbird
|
||||
|
||||
+2
-1
@@ -13,10 +13,11 @@ import (
|
||||
|
||||
type Method string
|
||||
|
||||
var (
|
||||
const (
|
||||
MethodPassword Method = "password"
|
||||
MethodPIN Method = "pin"
|
||||
MethodOIDC Method = "oidc"
|
||||
MethodHeader Method = "header"
|
||||
)
|
||||
|
||||
func (m Method) String() string {
|
||||
|
||||
+64
-31
@@ -7,6 +7,7 @@ import (
|
||||
"os/signal"
|
||||
"strconv"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/spf13/cobra"
|
||||
@@ -34,28 +35,35 @@ var (
|
||||
)
|
||||
|
||||
var (
|
||||
debugLogs bool
|
||||
mgmtAddr string
|
||||
addr string
|
||||
proxyDomain string
|
||||
certDir string
|
||||
acmeCerts bool
|
||||
acmeAddr string
|
||||
acmeDir string
|
||||
acmeEABKID string
|
||||
acmeEABHMACKey string
|
||||
acmeChallengeType string
|
||||
debugEndpoint bool
|
||||
debugEndpointAddr string
|
||||
healthAddr string
|
||||
forwardedProto string
|
||||
trustedProxies string
|
||||
certFile string
|
||||
certKeyFile string
|
||||
certLockMethod string
|
||||
wgPort int
|
||||
proxyProtocol bool
|
||||
preSharedKey string
|
||||
logLevel string
|
||||
debugLogs bool
|
||||
mgmtAddr string
|
||||
addr string
|
||||
proxyDomain string
|
||||
maxDialTimeout time.Duration
|
||||
maxSessionIdleTimeout time.Duration
|
||||
certDir string
|
||||
acmeCerts bool
|
||||
acmeAddr string
|
||||
acmeDir string
|
||||
acmeEABKID string
|
||||
acmeEABHMACKey string
|
||||
acmeChallengeType string
|
||||
debugEndpoint bool
|
||||
debugEndpointAddr string
|
||||
healthAddr string
|
||||
forwardedProto string
|
||||
trustedProxies string
|
||||
certFile string
|
||||
certKeyFile string
|
||||
certLockMethod string
|
||||
wildcardCertDir string
|
||||
wgPort uint16
|
||||
proxyProtocol bool
|
||||
preSharedKey string
|
||||
supportsCustomPorts bool
|
||||
requireSubdomain bool
|
||||
geoDataDir string
|
||||
)
|
||||
|
||||
var rootCmd = &cobra.Command{
|
||||
@@ -68,7 +76,9 @@ var rootCmd = &cobra.Command{
|
||||
}
|
||||
|
||||
func init() {
|
||||
rootCmd.PersistentFlags().StringVar(&logLevel, "log-level", envStringOrDefault("NB_PROXY_LOG_LEVEL", "info"), "Log level: panic, fatal, error, warn, info, debug, trace")
|
||||
rootCmd.PersistentFlags().BoolVar(&debugLogs, "debug", envBoolOrDefault("NB_PROXY_DEBUG_LOGS", false), "Enable debug logs")
|
||||
_ = rootCmd.PersistentFlags().MarkDeprecated("debug", "use --log-level instead")
|
||||
rootCmd.Flags().StringVar(&mgmtAddr, "mgmt", envStringOrDefault("NB_PROXY_MANAGEMENT_ADDRESS", DefaultManagementURL), "Management address to connect to")
|
||||
rootCmd.Flags().StringVar(&addr, "addr", envStringOrDefault("NB_PROXY_ADDRESS", ":443"), "Reverse proxy address to listen on")
|
||||
rootCmd.Flags().StringVar(&proxyDomain, "domain", envStringOrDefault("NB_PROXY_DOMAIN", ""), "The Domain at which this proxy will be reached. e.g., netbird.example.com")
|
||||
@@ -87,9 +97,15 @@ func init() {
|
||||
rootCmd.Flags().StringVar(&certFile, "cert-file", envStringOrDefault("NB_PROXY_CERTIFICATE_FILE", "tls.crt"), "TLS certificate filename within the certificate directory")
|
||||
rootCmd.Flags().StringVar(&certKeyFile, "cert-key-file", envStringOrDefault("NB_PROXY_CERTIFICATE_KEY_FILE", "tls.key"), "TLS certificate key filename within the certificate directory")
|
||||
rootCmd.Flags().StringVar(&certLockMethod, "cert-lock-method", envStringOrDefault("NB_PROXY_CERT_LOCK_METHOD", "auto"), "Certificate lock method for cross-replica coordination: auto, flock, or k8s-lease")
|
||||
rootCmd.Flags().IntVar(&wgPort, "wg-port", envIntOrDefault("NB_PROXY_WG_PORT", 0), "WireGuard listen port (0 = random). Fixed port only works with single-account deployments")
|
||||
rootCmd.Flags().StringVar(&wildcardCertDir, "wildcard-cert-dir", envStringOrDefault("NB_PROXY_WILDCARD_CERT_DIR", ""), "Directory containing wildcard certificate pairs (<name>.crt/<name>.key). Wildcard patterns are extracted from SANs automatically")
|
||||
rootCmd.Flags().Uint16Var(&wgPort, "wg-port", envUint16OrDefault("NB_PROXY_WG_PORT", 0), "WireGuard listen port (0 = random). Fixed port only works with single-account deployments")
|
||||
rootCmd.Flags().BoolVar(&proxyProtocol, "proxy-protocol", envBoolOrDefault("NB_PROXY_PROXY_PROTOCOL", false), "Enable PROXY protocol on TCP listeners to preserve client IPs behind L4 proxies")
|
||||
rootCmd.Flags().StringVar(&preSharedKey, "preshared-key", envStringOrDefault("NB_PROXY_PRESHARED_KEY", ""), "Define a pre-shared key for the tunnel between proxy and peers")
|
||||
rootCmd.Flags().BoolVar(&supportsCustomPorts, "supports-custom-ports", envBoolOrDefault("NB_PROXY_SUPPORTS_CUSTOM_PORTS", true), "Whether the proxy can bind arbitrary ports for UDP/TCP passthrough")
|
||||
rootCmd.Flags().BoolVar(&requireSubdomain, "require-subdomain", envBoolOrDefault("NB_PROXY_REQUIRE_SUBDOMAIN", false), "Require a subdomain label in front of the cluster domain")
|
||||
rootCmd.Flags().DurationVar(&maxDialTimeout, "max-dial-timeout", envDurationOrDefault("NB_PROXY_MAX_DIAL_TIMEOUT", 0), "Cap per-service backend dial timeout (0 = no cap)")
|
||||
rootCmd.Flags().DurationVar(&maxSessionIdleTimeout, "max-session-idle-timeout", envDurationOrDefault("NB_PROXY_MAX_SESSION_IDLE_TIMEOUT", 0), "Cap per-service session idle timeout (0 = no cap)")
|
||||
rootCmd.Flags().StringVar(&geoDataDir, "geo-data-dir", envStringOrDefault("NB_PROXY_GEO_DATA_DIR", "/var/lib/netbird/geolocation"), "Directory for the GeoLite2 MMDB file (auto-downloaded if missing)")
|
||||
}
|
||||
|
||||
// Execute runs the root command.
|
||||
@@ -115,7 +131,7 @@ func runServer(cmd *cobra.Command, args []string) error {
|
||||
return fmt.Errorf("proxy token is required: set %s environment variable", envProxyToken)
|
||||
}
|
||||
|
||||
level := "error"
|
||||
level := logLevel
|
||||
if debugLogs {
|
||||
level = "debug"
|
||||
}
|
||||
@@ -162,19 +178,21 @@ func runServer(cmd *cobra.Command, args []string) error {
|
||||
ForwardedProto: forwardedProto,
|
||||
TrustedProxies: parsedTrustedProxies,
|
||||
CertLockMethod: nbacme.CertLockMethod(certLockMethod),
|
||||
WildcardCertDir: wildcardCertDir,
|
||||
WireguardPort: wgPort,
|
||||
ProxyProtocol: proxyProtocol,
|
||||
PreSharedKey: preSharedKey,
|
||||
SupportsCustomPorts: supportsCustomPorts,
|
||||
RequireSubdomain: requireSubdomain,
|
||||
MaxDialTimeout: maxDialTimeout,
|
||||
MaxSessionIdleTimeout: maxSessionIdleTimeout,
|
||||
GeoDataDir: geoDataDir,
|
||||
}
|
||||
|
||||
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGTERM, syscall.SIGINT)
|
||||
defer stop()
|
||||
|
||||
if err := srv.ListenAndServe(ctx, addr); err != nil {
|
||||
logger.Error(err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
return srv.ListenAndServe(ctx, addr)
|
||||
}
|
||||
|
||||
func envBoolOrDefault(key string, def bool) bool {
|
||||
@@ -184,6 +202,7 @@ func envBoolOrDefault(key string, def bool) bool {
|
||||
}
|
||||
parsed, err := strconv.ParseBool(v)
|
||||
if err != nil {
|
||||
log.Warnf("parse %s=%q: %v, using default %v", key, v, err, def)
|
||||
return def
|
||||
}
|
||||
return parsed
|
||||
@@ -197,13 +216,27 @@ func envStringOrDefault(key string, def string) string {
|
||||
return v
|
||||
}
|
||||
|
||||
func envIntOrDefault(key string, def int) int {
|
||||
func envUint16OrDefault(key string, def uint16) uint16 {
|
||||
v, exists := os.LookupEnv(key)
|
||||
if !exists {
|
||||
return def
|
||||
}
|
||||
parsed, err := strconv.Atoi(v)
|
||||
parsed, err := strconv.ParseUint(v, 10, 16)
|
||||
if err != nil {
|
||||
log.Warnf("parse %s=%q: %v, using default %d", key, v, err, def)
|
||||
return def
|
||||
}
|
||||
return uint16(parsed)
|
||||
}
|
||||
|
||||
func envDurationOrDefault(key string, def time.Duration) time.Duration {
|
||||
v, exists := os.LookupEnv(key)
|
||||
if !exists {
|
||||
return def
|
||||
}
|
||||
parsed, err := time.ParseDuration(v)
|
||||
if err != nil {
|
||||
log.Warnf("parse %s=%q: %v, using default %s", key, v, err, def)
|
||||
return def
|
||||
}
|
||||
return parsed
|
||||
|
||||
@@ -38,11 +38,18 @@ func (m *mockMappingStream) Context() context.Context { return context.Backgroun
|
||||
func (m *mockMappingStream) SendMsg(any) error { return nil }
|
||||
func (m *mockMappingStream) RecvMsg(any) error { return nil }
|
||||
|
||||
func closedChan() chan struct{} {
|
||||
ch := make(chan struct{})
|
||||
close(ch)
|
||||
return ch
|
||||
}
|
||||
|
||||
func TestHandleMappingStream_SyncCompleteFlag(t *testing.T) {
|
||||
checker := health.NewChecker(nil, nil)
|
||||
s := &Server{
|
||||
Logger: log.StandardLogger(),
|
||||
healthChecker: checker,
|
||||
routerReady: closedChan(),
|
||||
}
|
||||
|
||||
stream := &mockMappingStream{
|
||||
@@ -62,6 +69,7 @@ func TestHandleMappingStream_NoSyncFlagDoesNotMarkDone(t *testing.T) {
|
||||
s := &Server{
|
||||
Logger: log.StandardLogger(),
|
||||
healthChecker: checker,
|
||||
routerReady: closedChan(),
|
||||
}
|
||||
|
||||
stream := &mockMappingStream{
|
||||
@@ -78,7 +86,8 @@ func TestHandleMappingStream_NoSyncFlagDoesNotMarkDone(t *testing.T) {
|
||||
|
||||
func TestHandleMappingStream_NilHealthChecker(t *testing.T) {
|
||||
s := &Server{
|
||||
Logger: log.StandardLogger(),
|
||||
Logger: log.StandardLogger(),
|
||||
routerReady: closedChan(),
|
||||
}
|
||||
|
||||
stream := &mockMappingStream{
|
||||
|
||||
@@ -4,13 +4,16 @@ import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/rs/xid"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
@@ -19,6 +22,17 @@ const (
|
||||
bytesThreshold = 1024 * 1024 * 1024 // Log every 1GB
|
||||
usageCleanupPeriod = 1 * time.Hour // Clean up stale counters every hour
|
||||
usageInactiveWindow = 24 * time.Hour // Consider domain inactive if no traffic for 24 hours
|
||||
logSendTimeout = 10 * time.Second
|
||||
|
||||
// denyCooldown is the min interval between deny log entries per service+reason
|
||||
// to prevent flooding from denied connections (e.g. UDP packets from blocked IPs).
|
||||
denyCooldown = 10 * time.Second
|
||||
|
||||
// maxDenyBuckets caps tracked deny rate-limit entries to bound memory under DDoS.
|
||||
maxDenyBuckets = 10000
|
||||
|
||||
// maxLogWorkers caps concurrent gRPC send goroutines.
|
||||
maxLogWorkers = 4096
|
||||
)
|
||||
|
||||
type domainUsage struct {
|
||||
@@ -35,6 +49,18 @@ type gRPCClient interface {
|
||||
SendAccessLog(ctx context.Context, in *proto.SendAccessLogRequest, opts ...grpc.CallOption) (*proto.SendAccessLogResponse, error)
|
||||
}
|
||||
|
||||
// denyBucketKey identifies a rate-limited deny log stream.
|
||||
type denyBucketKey struct {
|
||||
ServiceID types.ServiceID
|
||||
Reason string
|
||||
}
|
||||
|
||||
// denyBucket tracks rate-limited deny log entries.
|
||||
type denyBucket struct {
|
||||
lastLogged time.Time
|
||||
suppressed int64
|
||||
}
|
||||
|
||||
// Logger sends access log entries to the management server via gRPC.
|
||||
type Logger struct {
|
||||
client gRPCClient
|
||||
@@ -44,7 +70,12 @@ type Logger struct {
|
||||
usageMux sync.Mutex
|
||||
domainUsage map[string]*domainUsage
|
||||
|
||||
denyMu sync.Mutex
|
||||
denyBuckets map[denyBucketKey]*denyBucket
|
||||
|
||||
logSem chan struct{}
|
||||
cleanupCancel context.CancelFunc
|
||||
dropped atomic.Int64
|
||||
}
|
||||
|
||||
// NewLogger creates a new access log Logger. The trustedProxies parameter
|
||||
@@ -61,6 +92,8 @@ func NewLogger(client gRPCClient, logger *log.Logger, trustedProxies []netip.Pre
|
||||
logger: logger,
|
||||
trustedProxies: trustedProxies,
|
||||
domainUsage: make(map[string]*domainUsage),
|
||||
denyBuckets: make(map[denyBucketKey]*denyBucket),
|
||||
logSem: make(chan struct{}, maxLogWorkers),
|
||||
cleanupCancel: cancel,
|
||||
}
|
||||
|
||||
@@ -79,22 +112,104 @@ func (l *Logger) Close() {
|
||||
|
||||
type logEntry struct {
|
||||
ID string
|
||||
AccountID string
|
||||
ServiceId string
|
||||
AccountID types.AccountID
|
||||
ServiceID types.ServiceID
|
||||
Host string
|
||||
Path string
|
||||
DurationMs int64
|
||||
Method string
|
||||
ResponseCode int32
|
||||
SourceIp string
|
||||
SourceIP netip.Addr
|
||||
AuthMechanism string
|
||||
UserId string
|
||||
UserID string
|
||||
AuthSuccess bool
|
||||
BytesUpload int64
|
||||
BytesDownload int64
|
||||
Protocol Protocol
|
||||
}
|
||||
|
||||
func (l *Logger) log(ctx context.Context, entry logEntry) {
|
||||
// Protocol identifies the transport protocol of an access log entry.
|
||||
type Protocol string
|
||||
|
||||
const (
|
||||
ProtocolHTTP Protocol = "http"
|
||||
ProtocolTCP Protocol = "tcp"
|
||||
ProtocolUDP Protocol = "udp"
|
||||
ProtocolTLS Protocol = "tls"
|
||||
)
|
||||
|
||||
// L4Entry holds the data for a layer-4 (TCP/UDP) access log entry.
|
||||
type L4Entry struct {
|
||||
AccountID types.AccountID
|
||||
ServiceID types.ServiceID
|
||||
Protocol Protocol
|
||||
Host string // SNI hostname or listen address
|
||||
SourceIP netip.Addr
|
||||
DurationMs int64
|
||||
BytesUpload int64
|
||||
BytesDownload int64
|
||||
// DenyReason, when non-empty, indicates the connection was denied.
|
||||
// Values match the HTTP auth mechanism strings: "ip_restricted",
|
||||
// "country_restricted", "geo_unavailable".
|
||||
DenyReason string
|
||||
}
|
||||
|
||||
// LogL4 sends an access log entry for a layer-4 connection (TCP or UDP).
|
||||
// The call is non-blocking: the gRPC send happens in a background goroutine.
|
||||
func (l *Logger) LogL4(entry L4Entry) {
|
||||
le := logEntry{
|
||||
ID: xid.New().String(),
|
||||
AccountID: entry.AccountID,
|
||||
ServiceID: entry.ServiceID,
|
||||
Protocol: entry.Protocol,
|
||||
Host: entry.Host,
|
||||
SourceIP: entry.SourceIP,
|
||||
DurationMs: entry.DurationMs,
|
||||
BytesUpload: entry.BytesUpload,
|
||||
BytesDownload: entry.BytesDownload,
|
||||
}
|
||||
if entry.DenyReason != "" {
|
||||
if !l.allowDenyLog(entry.ServiceID, entry.DenyReason) {
|
||||
return
|
||||
}
|
||||
le.AuthMechanism = entry.DenyReason
|
||||
le.AuthSuccess = false
|
||||
}
|
||||
l.log(le)
|
||||
l.trackUsage(entry.Host, entry.BytesUpload+entry.BytesDownload)
|
||||
}
|
||||
|
||||
// allowDenyLog rate-limits deny log entries per service+reason combination.
|
||||
func (l *Logger) allowDenyLog(serviceID types.ServiceID, reason string) bool {
|
||||
key := denyBucketKey{ServiceID: serviceID, Reason: reason}
|
||||
now := time.Now()
|
||||
|
||||
l.denyMu.Lock()
|
||||
defer l.denyMu.Unlock()
|
||||
|
||||
b, ok := l.denyBuckets[key]
|
||||
if !ok {
|
||||
if len(l.denyBuckets) >= maxDenyBuckets {
|
||||
return false
|
||||
}
|
||||
l.denyBuckets[key] = &denyBucket{lastLogged: now}
|
||||
return true
|
||||
}
|
||||
|
||||
if now.Sub(b.lastLogged) >= denyCooldown {
|
||||
if b.suppressed > 0 {
|
||||
l.logger.Debugf("access restriction: suppressed %d deny log entries for %s (%s)", b.suppressed, serviceID, reason)
|
||||
}
|
||||
b.lastLogged = now
|
||||
b.suppressed = 0
|
||||
return true
|
||||
}
|
||||
|
||||
b.suppressed++
|
||||
return false
|
||||
}
|
||||
|
||||
func (l *Logger) log(entry logEntry) {
|
||||
// Fire off the log request in a separate routine.
|
||||
// This increases the possibility of losing a log message
|
||||
// (although it should still get logged in the event of an error),
|
||||
@@ -103,43 +218,58 @@ func (l *Logger) log(ctx context.Context, entry logEntry) {
|
||||
// There is also a chance that log messages will arrive at
|
||||
// the server out of order; however, the timestamp should
|
||||
// allow for resolving that on the server.
|
||||
now := timestamppb.Now() // Grab the timestamp before launching the goroutine to try to prevent weird timing issues. This is probably unnecessary.
|
||||
now := timestamppb.Now()
|
||||
select {
|
||||
case l.logSem <- struct{}{}:
|
||||
default:
|
||||
total := l.dropped.Add(1)
|
||||
l.logger.Debugf("access log send dropped: worker limit reached (total dropped: %d)", total)
|
||||
return
|
||||
}
|
||||
go func() {
|
||||
logCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer func() { <-l.logSem }()
|
||||
logCtx, cancel := context.WithTimeout(context.Background(), logSendTimeout)
|
||||
defer cancel()
|
||||
// Only OIDC sessions have a meaningful user identity.
|
||||
if entry.AuthMechanism != auth.MethodOIDC.String() {
|
||||
entry.UserId = ""
|
||||
entry.UserID = ""
|
||||
}
|
||||
|
||||
var sourceIP string
|
||||
if entry.SourceIP.IsValid() {
|
||||
sourceIP = entry.SourceIP.String()
|
||||
}
|
||||
|
||||
if _, err := l.client.SendAccessLog(logCtx, &proto.SendAccessLogRequest{
|
||||
Log: &proto.AccessLog{
|
||||
LogId: entry.ID,
|
||||
AccountId: entry.AccountID,
|
||||
AccountId: string(entry.AccountID),
|
||||
Timestamp: now,
|
||||
ServiceId: entry.ServiceId,
|
||||
ServiceId: string(entry.ServiceID),
|
||||
Host: entry.Host,
|
||||
Path: entry.Path,
|
||||
DurationMs: entry.DurationMs,
|
||||
Method: entry.Method,
|
||||
ResponseCode: entry.ResponseCode,
|
||||
SourceIp: entry.SourceIp,
|
||||
SourceIp: sourceIP,
|
||||
AuthMechanism: entry.AuthMechanism,
|
||||
UserId: entry.UserId,
|
||||
UserId: entry.UserID,
|
||||
AuthSuccess: entry.AuthSuccess,
|
||||
BytesUpload: entry.BytesUpload,
|
||||
BytesDownload: entry.BytesDownload,
|
||||
Protocol: string(entry.Protocol),
|
||||
},
|
||||
}); err != nil {
|
||||
// If it fails to send on the gRPC connection, then at least log it to the error log.
|
||||
l.logger.WithFields(log.Fields{
|
||||
"service_id": entry.ServiceId,
|
||||
"service_id": entry.ServiceID,
|
||||
"host": entry.Host,
|
||||
"path": entry.Path,
|
||||
"duration": entry.DurationMs,
|
||||
"method": entry.Method,
|
||||
"response_code": entry.ResponseCode,
|
||||
"source_ip": entry.SourceIp,
|
||||
"source_ip": sourceIP,
|
||||
"auth_mechanism": entry.AuthMechanism,
|
||||
"user_id": entry.UserId,
|
||||
"user_id": entry.UserID,
|
||||
"auth_success": entry.AuthSuccess,
|
||||
"error": err,
|
||||
}).Error("Error sending access log on gRPC connection")
|
||||
@@ -198,7 +328,7 @@ func (l *Logger) trackUsage(domain string, bytesTransferred int64) {
|
||||
}
|
||||
}
|
||||
|
||||
// cleanupStaleUsage removes usage entries for domains that have been inactive.
|
||||
// cleanupStaleUsage removes usage and deny-rate-limit entries that have been inactive.
|
||||
func (l *Logger) cleanupStaleUsage(ctx context.Context) {
|
||||
ticker := time.NewTicker(usageCleanupPeriod)
|
||||
defer ticker.Stop()
|
||||
@@ -208,20 +338,41 @@ func (l *Logger) cleanupStaleUsage(ctx context.Context) {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
l.usageMux.Lock()
|
||||
now := time.Now()
|
||||
removed := 0
|
||||
for domain, usage := range l.domainUsage {
|
||||
if now.Sub(usage.lastActivity) > usageInactiveWindow {
|
||||
delete(l.domainUsage, domain)
|
||||
removed++
|
||||
}
|
||||
}
|
||||
l.usageMux.Unlock()
|
||||
|
||||
if removed > 0 {
|
||||
l.logger.Debugf("cleaned up %d stale domain usage entries", removed)
|
||||
}
|
||||
l.cleanupDomainUsage(now)
|
||||
l.cleanupDenyBuckets(now)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (l *Logger) cleanupDomainUsage(now time.Time) {
|
||||
l.usageMux.Lock()
|
||||
defer l.usageMux.Unlock()
|
||||
|
||||
removed := 0
|
||||
for domain, usage := range l.domainUsage {
|
||||
if now.Sub(usage.lastActivity) > usageInactiveWindow {
|
||||
delete(l.domainUsage, domain)
|
||||
removed++
|
||||
}
|
||||
}
|
||||
if removed > 0 {
|
||||
l.logger.Debugf("cleaned up %d stale domain usage entries", removed)
|
||||
}
|
||||
}
|
||||
|
||||
func (l *Logger) cleanupDenyBuckets(now time.Time) {
|
||||
l.denyMu.Lock()
|
||||
defer l.denyMu.Unlock()
|
||||
|
||||
removed := 0
|
||||
for key, bucket := range l.denyBuckets {
|
||||
if now.Sub(bucket.lastLogged) > usageInactiveWindow {
|
||||
delete(l.denyBuckets, key)
|
||||
removed++
|
||||
}
|
||||
}
|
||||
if removed > 0 {
|
||||
l.logger.Debugf("cleaned up %d stale deny rate-limit entries", removed)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
"github.com/netbirdio/netbird/proxy/web"
|
||||
)
|
||||
|
||||
// Middleware wraps an HTTP handler to log access entries and resolve client IPs.
|
||||
func (l *Logger) Middleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// Skip logging for internal proxy assets (CSS, JS, etc.)
|
||||
@@ -47,8 +48,9 @@ func (l *Logger) Middleware(next http.Handler) http.Handler {
|
||||
// Create a mutable struct to capture data from downstream handlers.
|
||||
// We pass a pointer in the context - the pointer itself flows down immutably,
|
||||
// but the struct it points to can be mutated by inner handlers.
|
||||
capturedData := &proxy.CapturedData{RequestID: requestID}
|
||||
capturedData := proxy.NewCapturedData(requestID)
|
||||
capturedData.SetClientIP(sourceIp)
|
||||
|
||||
ctx := proxy.WithCapturedData(r.Context(), capturedData)
|
||||
|
||||
start := time.Now()
|
||||
@@ -66,24 +68,25 @@ func (l *Logger) Middleware(next http.Handler) http.Handler {
|
||||
|
||||
entry := logEntry{
|
||||
ID: requestID,
|
||||
ServiceId: capturedData.GetServiceId(),
|
||||
AccountID: string(capturedData.GetAccountId()),
|
||||
ServiceID: capturedData.GetServiceID(),
|
||||
AccountID: capturedData.GetAccountID(),
|
||||
Host: host,
|
||||
Path: r.URL.Path,
|
||||
DurationMs: duration.Milliseconds(),
|
||||
Method: r.Method,
|
||||
ResponseCode: int32(sw.status),
|
||||
SourceIp: sourceIp,
|
||||
SourceIP: sourceIp,
|
||||
AuthMechanism: capturedData.GetAuthMethod(),
|
||||
UserId: capturedData.GetUserID(),
|
||||
UserID: capturedData.GetUserID(),
|
||||
AuthSuccess: sw.status != http.StatusUnauthorized && sw.status != http.StatusForbidden,
|
||||
BytesUpload: bytesUpload,
|
||||
BytesDownload: bytesDownload,
|
||||
Protocol: ProtocolHTTP,
|
||||
}
|
||||
l.logger.Debugf("response: request_id=%s method=%s host=%s path=%s status=%d duration=%dms source=%s origin=%s service=%s account=%s",
|
||||
requestID, r.Method, host, r.URL.Path, sw.status, duration.Milliseconds(), sourceIp, capturedData.GetOrigin(), capturedData.GetServiceId(), capturedData.GetAccountId())
|
||||
requestID, r.Method, host, r.URL.Path, sw.status, duration.Milliseconds(), sourceIp, capturedData.GetOrigin(), capturedData.GetServiceID(), capturedData.GetAccountID())
|
||||
|
||||
l.log(r.Context(), entry)
|
||||
l.log(entry)
|
||||
|
||||
// Track usage for cost monitoring (upload + download) by domain
|
||||
l.trackUsage(host, bytesUpload+bytesDownload)
|
||||
|
||||
@@ -11,6 +11,6 @@ import (
|
||||
// proxy configuration. When trustedProxies is non-empty and the direct
|
||||
// connection is from a trusted source, it walks X-Forwarded-For right-to-left
|
||||
// skipping trusted IPs. Otherwise it returns RemoteAddr directly.
|
||||
func extractSourceIP(r *http.Request, trustedProxies []netip.Prefix) string {
|
||||
func extractSourceIP(r *http.Request, trustedProxies []netip.Prefix) netip.Addr {
|
||||
return proxy.ResolveClientIP(r.RemoteAddr, r.Header.Get("X-Forwarded-For"), trustedProxies)
|
||||
}
|
||||
|
||||
+213
-21
@@ -11,6 +11,8 @@ import (
|
||||
"fmt"
|
||||
"math/rand/v2"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
@@ -20,6 +22,8 @@ import (
|
||||
"golang.org/x/crypto/acme"
|
||||
"golang.org/x/crypto/acme/autocert"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/certwatch"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
)
|
||||
|
||||
@@ -27,7 +31,7 @@ import (
|
||||
var oidSCTList = asn1.ObjectIdentifier{1, 3, 6, 1, 4, 1, 11129, 2, 4, 2}
|
||||
|
||||
type certificateNotifier interface {
|
||||
NotifyCertificateIssued(ctx context.Context, accountID, serviceID, domain string) error
|
||||
NotifyCertificateIssued(ctx context.Context, accountID types.AccountID, serviceID types.ServiceID, domain string) error
|
||||
}
|
||||
|
||||
type domainState int
|
||||
@@ -39,8 +43,8 @@ const (
|
||||
)
|
||||
|
||||
type domainInfo struct {
|
||||
accountID string
|
||||
serviceID string
|
||||
accountID types.AccountID
|
||||
serviceID types.ServiceID
|
||||
state domainState
|
||||
err string
|
||||
}
|
||||
@@ -49,6 +53,34 @@ type metricsRecorder interface {
|
||||
RecordCertificateIssuance(duration time.Duration)
|
||||
}
|
||||
|
||||
// wildcardEntry maps a domain suffix (e.g. ".example.com") to a certwatch
|
||||
// watcher that hot-reloads the corresponding wildcard certificate from disk.
|
||||
type wildcardEntry struct {
|
||||
suffix string // e.g. ".example.com"
|
||||
pattern string // e.g. "*.example.com"
|
||||
watcher *certwatch.Watcher
|
||||
}
|
||||
|
||||
// ManagerConfig holds the configuration values for the ACME certificate manager.
|
||||
type ManagerConfig struct {
|
||||
// CertDir is the directory used for caching ACME certificates.
|
||||
CertDir string
|
||||
// ACMEURL is the ACME directory URL (e.g. Let's Encrypt).
|
||||
ACMEURL string
|
||||
// EABKID and EABHMACKey are optional External Account Binding credentials
|
||||
// required by some CAs (e.g. ZeroSSL). EABHMACKey is the base64
|
||||
// URL-encoded string provided by the CA.
|
||||
EABKID string
|
||||
EABHMACKey string
|
||||
// LockMethod controls the cross-replica coordination strategy.
|
||||
LockMethod CertLockMethod
|
||||
// WildcardDir is an optional path to a directory containing wildcard
|
||||
// certificate pairs (<name>.crt / <name>.key). Wildcard patterns are
|
||||
// extracted from the certificates' SAN lists. Domains matching a
|
||||
// wildcard are served from disk; all others go through ACME.
|
||||
WildcardDir string
|
||||
}
|
||||
|
||||
// Manager wraps autocert.Manager with domain tracking and cross-replica
|
||||
// coordination via a pluggable locking strategy. The locker prevents
|
||||
// duplicate ACME requests when multiple replicas share a certificate cache.
|
||||
@@ -60,54 +92,182 @@ type Manager struct {
|
||||
mu sync.RWMutex
|
||||
domains map[domain.Domain]*domainInfo
|
||||
|
||||
// wildcards holds all loaded wildcard certificates, keyed by suffix.
|
||||
wildcards []wildcardEntry
|
||||
|
||||
certNotifier certificateNotifier
|
||||
logger *log.Logger
|
||||
metrics metricsRecorder
|
||||
}
|
||||
|
||||
// NewManager creates a new ACME certificate manager. The certDir is used
|
||||
// for caching certificates. The lockMethod controls cross-replica
|
||||
// coordination strategy (see CertLockMethod constants).
|
||||
// eabKID and eabHMACKey are optional External Account Binding credentials
|
||||
// required for some CAs like ZeroSSL. The eabHMACKey should be the base64
|
||||
// URL-encoded string provided by the CA.
|
||||
func NewManager(certDir, acmeURL, eabKID, eabHMACKey string, notifier certificateNotifier, logger *log.Logger, lockMethod CertLockMethod, metrics metricsRecorder) *Manager {
|
||||
// NewManager creates a new ACME certificate manager.
|
||||
func NewManager(cfg ManagerConfig, notifier certificateNotifier, logger *log.Logger, metrics metricsRecorder) (*Manager, error) {
|
||||
if logger == nil {
|
||||
logger = log.StandardLogger()
|
||||
}
|
||||
mgr := &Manager{
|
||||
certDir: certDir,
|
||||
locker: newCertLocker(lockMethod, certDir, logger),
|
||||
certDir: cfg.CertDir,
|
||||
locker: newCertLocker(cfg.LockMethod, cfg.CertDir, logger),
|
||||
domains: make(map[domain.Domain]*domainInfo),
|
||||
certNotifier: notifier,
|
||||
logger: logger,
|
||||
metrics: metrics,
|
||||
}
|
||||
|
||||
if cfg.WildcardDir != "" {
|
||||
entries, err := loadWildcardDir(cfg.WildcardDir, logger)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load wildcard certificates from %q: %w", cfg.WildcardDir, err)
|
||||
}
|
||||
mgr.wildcards = entries
|
||||
}
|
||||
|
||||
var eab *acme.ExternalAccountBinding
|
||||
if eabKID != "" && eabHMACKey != "" {
|
||||
decodedKey, err := base64.RawURLEncoding.DecodeString(eabHMACKey)
|
||||
if cfg.EABKID != "" && cfg.EABHMACKey != "" {
|
||||
decodedKey, err := base64.RawURLEncoding.DecodeString(cfg.EABHMACKey)
|
||||
if err != nil {
|
||||
logger.Errorf("failed to decode EAB HMAC key: %v", err)
|
||||
} else {
|
||||
eab = &acme.ExternalAccountBinding{
|
||||
KID: eabKID,
|
||||
KID: cfg.EABKID,
|
||||
Key: decodedKey,
|
||||
}
|
||||
logger.Infof("configured External Account Binding with KID: %s", eabKID)
|
||||
logger.Infof("configured External Account Binding with KID: %s", cfg.EABKID)
|
||||
}
|
||||
}
|
||||
|
||||
mgr.Manager = &autocert.Manager{
|
||||
Prompt: autocert.AcceptTOS,
|
||||
HostPolicy: mgr.hostPolicy,
|
||||
Cache: autocert.DirCache(certDir),
|
||||
Cache: autocert.DirCache(cfg.CertDir),
|
||||
ExternalAccountBinding: eab,
|
||||
Client: &acme.Client{
|
||||
DirectoryURL: acmeURL,
|
||||
DirectoryURL: cfg.ACMEURL,
|
||||
},
|
||||
}
|
||||
return mgr
|
||||
return mgr, nil
|
||||
}
|
||||
|
||||
// WatchWildcards starts watching all wildcard certificate files for changes.
|
||||
// It blocks until ctx is cancelled. It is a no-op if no wildcards are loaded.
|
||||
func (mgr *Manager) WatchWildcards(ctx context.Context) {
|
||||
if len(mgr.wildcards) == 0 {
|
||||
return
|
||||
}
|
||||
seen := make(map[*certwatch.Watcher]struct{})
|
||||
var wg sync.WaitGroup
|
||||
for i := range mgr.wildcards {
|
||||
w := mgr.wildcards[i].watcher
|
||||
if _, ok := seen[w]; ok {
|
||||
continue
|
||||
}
|
||||
seen[w] = struct{}{}
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
w.Watch(ctx)
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
// loadWildcardDir scans dir for .crt files, pairs each with a matching .key
|
||||
// file, loads them, and extracts wildcard SANs (*.example.com) to build
|
||||
// the suffix lookup entries.
|
||||
func loadWildcardDir(dir string, logger *log.Logger) ([]wildcardEntry, error) {
|
||||
crtFiles, err := filepath.Glob(filepath.Join(dir, "*.crt"))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("glob certificate files: %w", err)
|
||||
}
|
||||
|
||||
if len(crtFiles) == 0 {
|
||||
return nil, fmt.Errorf("no .crt files found in %s", dir)
|
||||
}
|
||||
|
||||
var entries []wildcardEntry
|
||||
|
||||
for _, crtPath := range crtFiles {
|
||||
base := strings.TrimSuffix(filepath.Base(crtPath), ".crt")
|
||||
keyPath := filepath.Join(dir, base+".key")
|
||||
if _, err := os.Stat(keyPath); err != nil {
|
||||
logger.Warnf("skipping %s: no matching key file %s", crtPath, keyPath)
|
||||
continue
|
||||
}
|
||||
|
||||
watcher, err := certwatch.NewWatcher(crtPath, keyPath, logger)
|
||||
if err != nil {
|
||||
logger.Warnf("skipping %s: %v", crtPath, err)
|
||||
continue
|
||||
}
|
||||
|
||||
leaf := watcher.Leaf()
|
||||
if leaf == nil {
|
||||
logger.Warnf("skipping %s: no parsed leaf certificate", crtPath)
|
||||
continue
|
||||
}
|
||||
|
||||
for _, san := range leaf.DNSNames {
|
||||
suffix, ok := parseWildcard(san)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
entries = append(entries, wildcardEntry{
|
||||
suffix: suffix,
|
||||
pattern: san,
|
||||
watcher: watcher,
|
||||
})
|
||||
logger.Infof("wildcard certificate loaded: %s (from %s)", san, filepath.Base(crtPath))
|
||||
}
|
||||
}
|
||||
|
||||
if len(entries) == 0 {
|
||||
return nil, fmt.Errorf("no wildcard SANs (*.example.com) found in certificates in %s", dir)
|
||||
}
|
||||
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
// parseWildcard validates a wildcard domain pattern like "*.example.com"
|
||||
// and returns the suffix ".example.com" for matching.
|
||||
func parseWildcard(pattern string) (suffix string, ok bool) {
|
||||
if !strings.HasPrefix(pattern, "*.") {
|
||||
return "", false
|
||||
}
|
||||
parent := pattern[1:] // ".example.com"
|
||||
if strings.Count(parent, ".") < 1 {
|
||||
return "", false
|
||||
}
|
||||
return strings.ToLower(parent), true
|
||||
}
|
||||
|
||||
// findWildcardEntry returns the wildcard entry that covers host, or nil.
|
||||
func (mgr *Manager) findWildcardEntry(host string) *wildcardEntry {
|
||||
if len(mgr.wildcards) == 0 {
|
||||
return nil
|
||||
}
|
||||
host = strings.ToLower(host)
|
||||
for i := range mgr.wildcards {
|
||||
e := &mgr.wildcards[i]
|
||||
if !strings.HasSuffix(host, e.suffix) {
|
||||
continue
|
||||
}
|
||||
// Single-level match: prefix before suffix must have no dots.
|
||||
prefix := strings.TrimSuffix(host, e.suffix)
|
||||
if len(prefix) > 0 && !strings.Contains(prefix, ".") {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// WildcardPatterns returns the wildcard patterns that are currently loaded.
|
||||
func (mgr *Manager) WildcardPatterns() []string {
|
||||
patterns := make([]string, len(mgr.wildcards))
|
||||
for i, e := range mgr.wildcards {
|
||||
patterns[i] = e.pattern
|
||||
}
|
||||
slices.Sort(patterns)
|
||||
return patterns
|
||||
}
|
||||
|
||||
func (mgr *Manager) hostPolicy(_ context.Context, host string) error {
|
||||
@@ -123,8 +283,39 @@ func (mgr *Manager) hostPolicy(_ context.Context, host string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// AddDomain registers a domain for ACME certificate prefetching.
|
||||
func (mgr *Manager) AddDomain(d domain.Domain, accountID, serviceID string) {
|
||||
// GetCertificate returns the TLS certificate for the given ClientHello.
|
||||
// If the requested domain matches a loaded wildcard, the static wildcard
|
||||
// certificate is returned. Otherwise, the ACME autocert manager handles
|
||||
// the request.
|
||||
func (mgr *Manager) GetCertificate(hello *tls.ClientHelloInfo) (*tls.Certificate, error) {
|
||||
if e := mgr.findWildcardEntry(hello.ServerName); e != nil {
|
||||
return e.watcher.GetCertificate(hello)
|
||||
}
|
||||
return mgr.Manager.GetCertificate(hello)
|
||||
}
|
||||
|
||||
// AddDomain registers a domain for certificate management. Domains that
|
||||
// match a loaded wildcard are marked ready immediately (they use the
|
||||
// static wildcard certificate) and the method returns true. All other
|
||||
// domains go through ACME prefetch and the method returns false.
|
||||
//
|
||||
// When AddDomain returns true the caller is responsible for sending any
|
||||
// certificate-ready notifications after the surrounding operation (e.g.
|
||||
// mapping update) has committed successfully.
|
||||
func (mgr *Manager) AddDomain(d domain.Domain, accountID types.AccountID, serviceID types.ServiceID) (wildcardHit bool) {
|
||||
name := d.PunycodeString()
|
||||
if e := mgr.findWildcardEntry(name); e != nil {
|
||||
mgr.mu.Lock()
|
||||
mgr.domains[d] = &domainInfo{
|
||||
accountID: accountID,
|
||||
serviceID: serviceID,
|
||||
state: domainReady,
|
||||
}
|
||||
mgr.mu.Unlock()
|
||||
mgr.logger.Debugf("domain %q matches wildcard %q, using static certificate", name, e.pattern)
|
||||
return true
|
||||
}
|
||||
|
||||
mgr.mu.Lock()
|
||||
mgr.domains[d] = &domainInfo{
|
||||
accountID: accountID,
|
||||
@@ -134,6 +325,7 @@ func (mgr *Manager) AddDomain(d domain.Domain, accountID, serviceID string) {
|
||||
mgr.mu.Unlock()
|
||||
|
||||
go mgr.prefetchCertificate(d)
|
||||
return false
|
||||
}
|
||||
|
||||
// prefetchCertificate proactively triggers certificate generation for a domain.
|
||||
|
||||
@@ -2,16 +2,29 @@ package acme
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/ecdsa"
|
||||
"crypto/elliptic"
|
||||
"crypto/rand"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/pem"
|
||||
"math/big"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
)
|
||||
|
||||
func TestHostPolicy(t *testing.T) {
|
||||
mgr := NewManager(t.TempDir(), "https://acme.example.com/directory", "", "", nil, nil, "", nil)
|
||||
mgr.AddDomain("example.com", "acc1", "rp1")
|
||||
mgr, err := NewManager(ManagerConfig{CertDir: t.TempDir(), ACMEURL: "https://acme.example.com/directory"}, nil, nil, nil)
|
||||
require.NoError(t, err)
|
||||
mgr.AddDomain("example.com", types.AccountID("acc1"), types.ServiceID("rp1"))
|
||||
|
||||
// Wait for the background prefetch goroutine to finish so the temp dir
|
||||
// can be cleaned up without a race.
|
||||
@@ -70,7 +83,8 @@ func TestHostPolicy(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDomainStates(t *testing.T) {
|
||||
mgr := NewManager(t.TempDir(), "https://acme.example.com/directory", "", "", nil, nil, "", nil)
|
||||
mgr, err := NewManager(ManagerConfig{CertDir: t.TempDir(), ACMEURL: "https://acme.example.com/directory"}, nil, nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, 0, mgr.PendingCerts(), "initially zero")
|
||||
assert.Equal(t, 0, mgr.TotalDomains(), "initially zero domains")
|
||||
@@ -80,8 +94,8 @@ func TestDomainStates(t *testing.T) {
|
||||
|
||||
// AddDomain starts as pending, then the prefetch goroutine will fail
|
||||
// (no real ACME server) and transition to failed.
|
||||
mgr.AddDomain("a.example.com", "acc1", "rp1")
|
||||
mgr.AddDomain("b.example.com", "acc1", "rp1")
|
||||
mgr.AddDomain("a.example.com", types.AccountID("acc1"), types.ServiceID("rp1"))
|
||||
mgr.AddDomain("b.example.com", types.AccountID("acc1"), types.ServiceID("rp1"))
|
||||
|
||||
assert.Equal(t, 2, mgr.TotalDomains(), "two domains registered")
|
||||
|
||||
@@ -100,3 +114,193 @@ func TestDomainStates(t *testing.T) {
|
||||
assert.Contains(t, failed, "b.example.com")
|
||||
assert.Empty(t, mgr.ReadyDomains())
|
||||
}
|
||||
|
||||
func TestParseWildcard(t *testing.T) {
|
||||
tests := []struct {
|
||||
pattern string
|
||||
wantSuffix string
|
||||
wantOK bool
|
||||
}{
|
||||
{"*.example.com", ".example.com", true},
|
||||
{"*.foo.example.com", ".foo.example.com", true},
|
||||
{"*.COM", ".com", true}, // single-label TLD
|
||||
{"example.com", "", false}, // no wildcard prefix
|
||||
{"*example.com", "", false}, // missing dot
|
||||
{"**.example.com", "", false}, // double star
|
||||
{"", "", false},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.pattern, func(t *testing.T) {
|
||||
suffix, ok := parseWildcard(tc.pattern)
|
||||
assert.Equal(t, tc.wantOK, ok)
|
||||
if ok {
|
||||
assert.Equal(t, tc.wantSuffix, suffix)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMatchesWildcard(t *testing.T) {
|
||||
wcDir := t.TempDir()
|
||||
generateSelfSignedCert(t, wcDir, "example", "*.example.com")
|
||||
|
||||
acmeDir := t.TempDir()
|
||||
mgr, err := NewManager(ManagerConfig{CertDir: acmeDir, ACMEURL: "https://acme.example.com/directory", WildcardDir: wcDir}, nil, nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
tests := []struct {
|
||||
host string
|
||||
match bool
|
||||
}{
|
||||
{"foo.example.com", true},
|
||||
{"bar.example.com", true},
|
||||
{"FOO.Example.COM", true}, // case insensitive
|
||||
{"example.com", false}, // bare parent
|
||||
{"sub.foo.example.com", false}, // multi-level
|
||||
{"notexample.com", false},
|
||||
{"", false},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.host, func(t *testing.T) {
|
||||
assert.Equal(t, tc.match, mgr.findWildcardEntry(tc.host) != nil)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// generateSelfSignedCert creates a temporary self-signed certificate and key
|
||||
// for testing purposes. The baseName controls the output filenames:
|
||||
// <baseName>.crt and <baseName>.key.
|
||||
func generateSelfSignedCert(t *testing.T, dir, baseName string, dnsNames ...string) {
|
||||
t.Helper()
|
||||
|
||||
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||
require.NoError(t, err)
|
||||
|
||||
template := &x509.Certificate{
|
||||
SerialNumber: big.NewInt(1),
|
||||
Subject: pkix.Name{CommonName: dnsNames[0]},
|
||||
DNSNames: dnsNames,
|
||||
NotBefore: time.Now().Add(-time.Hour),
|
||||
NotAfter: time.Now().Add(24 * time.Hour),
|
||||
}
|
||||
|
||||
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
|
||||
require.NoError(t, err)
|
||||
|
||||
certFile, err := os.Create(filepath.Join(dir, baseName+".crt"))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, pem.Encode(certFile, &pem.Block{Type: "CERTIFICATE", Bytes: certDER}))
|
||||
require.NoError(t, certFile.Close())
|
||||
|
||||
keyDER, err := x509.MarshalECPrivateKey(key)
|
||||
require.NoError(t, err)
|
||||
keyFile, err := os.Create(filepath.Join(dir, baseName+".key"))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, pem.Encode(keyFile, &pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}))
|
||||
require.NoError(t, keyFile.Close())
|
||||
}
|
||||
|
||||
func TestWildcardAddDomainSkipsACME(t *testing.T) {
|
||||
wcDir := t.TempDir()
|
||||
generateSelfSignedCert(t, wcDir, "example", "*.example.com")
|
||||
|
||||
acmeDir := t.TempDir()
|
||||
mgr, err := NewManager(ManagerConfig{CertDir: acmeDir, ACMEURL: "https://acme.example.com/directory", WildcardDir: wcDir}, nil, nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Add a wildcard-matching domain — should be immediately ready.
|
||||
mgr.AddDomain("foo.example.com", types.AccountID("acc1"), types.ServiceID("svc1"))
|
||||
assert.Equal(t, 0, mgr.PendingCerts(), "wildcard domain should not be pending")
|
||||
assert.Equal(t, []string{"foo.example.com"}, mgr.ReadyDomains())
|
||||
|
||||
// Add a non-wildcard domain — should go through ACME (pending then failed).
|
||||
mgr.AddDomain("other.net", types.AccountID("acc2"), types.ServiceID("svc2"))
|
||||
assert.Equal(t, 2, mgr.TotalDomains())
|
||||
|
||||
// Wait for the ACME prefetch to fail.
|
||||
assert.Eventually(t, func() bool {
|
||||
return mgr.PendingCerts() == 0
|
||||
}, 30*time.Second, 100*time.Millisecond)
|
||||
|
||||
assert.Equal(t, []string{"foo.example.com"}, mgr.ReadyDomains())
|
||||
assert.Contains(t, mgr.FailedDomains(), "other.net")
|
||||
}
|
||||
|
||||
func TestWildcardGetCertificate(t *testing.T) {
|
||||
wcDir := t.TempDir()
|
||||
generateSelfSignedCert(t, wcDir, "example", "*.example.com")
|
||||
|
||||
acmeDir := t.TempDir()
|
||||
mgr, err := NewManager(ManagerConfig{CertDir: acmeDir, ACMEURL: "https://acme.example.com/directory", WildcardDir: wcDir}, nil, nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
mgr.AddDomain("foo.example.com", types.AccountID("acc1"), types.ServiceID("svc1"))
|
||||
|
||||
// GetCertificate for a wildcard-matching domain should return the static cert.
|
||||
cert, err := mgr.GetCertificate(&tls.ClientHelloInfo{ServerName: "foo.example.com"})
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, cert)
|
||||
assert.Contains(t, cert.Leaf.DNSNames, "*.example.com")
|
||||
}
|
||||
|
||||
func TestMultipleWildcards(t *testing.T) {
|
||||
wcDir := t.TempDir()
|
||||
generateSelfSignedCert(t, wcDir, "example", "*.example.com")
|
||||
generateSelfSignedCert(t, wcDir, "other", "*.other.org")
|
||||
|
||||
acmeDir := t.TempDir()
|
||||
mgr, err := NewManager(ManagerConfig{CertDir: acmeDir, ACMEURL: "https://acme.example.com/directory", WildcardDir: wcDir}, nil, nil, nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.ElementsMatch(t, []string{"*.example.com", "*.other.org"}, mgr.WildcardPatterns())
|
||||
|
||||
// Both wildcards should resolve.
|
||||
mgr.AddDomain("foo.example.com", types.AccountID("acc1"), types.ServiceID("svc1"))
|
||||
mgr.AddDomain("bar.other.org", types.AccountID("acc2"), types.ServiceID("svc2"))
|
||||
|
||||
assert.Equal(t, 0, mgr.PendingCerts())
|
||||
assert.ElementsMatch(t, []string{"foo.example.com", "bar.other.org"}, mgr.ReadyDomains())
|
||||
|
||||
// GetCertificate routes to the correct cert.
|
||||
cert1, err := mgr.GetCertificate(&tls.ClientHelloInfo{ServerName: "foo.example.com"})
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, cert1.Leaf.DNSNames, "*.example.com")
|
||||
|
||||
cert2, err := mgr.GetCertificate(&tls.ClientHelloInfo{ServerName: "bar.other.org"})
|
||||
require.NoError(t, err)
|
||||
assert.Contains(t, cert2.Leaf.DNSNames, "*.other.org")
|
||||
|
||||
// Non-matching domain falls through to ACME.
|
||||
mgr.AddDomain("custom.net", types.AccountID("acc3"), types.ServiceID("svc3"))
|
||||
assert.Eventually(t, func() bool {
|
||||
return mgr.PendingCerts() == 0
|
||||
}, 30*time.Second, 100*time.Millisecond)
|
||||
assert.Contains(t, mgr.FailedDomains(), "custom.net")
|
||||
}
|
||||
|
||||
func TestWildcardDirEmpty(t *testing.T) {
|
||||
wcDir := t.TempDir()
|
||||
// Empty directory — no .crt files.
|
||||
_, err := NewManager(ManagerConfig{CertDir: t.TempDir(), ACMEURL: "https://acme.example.com/directory", WildcardDir: wcDir}, nil, nil, nil)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "no .crt files found")
|
||||
}
|
||||
|
||||
func TestWildcardDirNonWildcardCert(t *testing.T) {
|
||||
wcDir := t.TempDir()
|
||||
// Certificate without a wildcard SAN.
|
||||
generateSelfSignedCert(t, wcDir, "plain", "plain.example.com")
|
||||
|
||||
_, err := NewManager(ManagerConfig{CertDir: t.TempDir(), ACMEURL: "https://acme.example.com/directory", WildcardDir: wcDir}, nil, nil, nil)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "no wildcard SANs")
|
||||
}
|
||||
|
||||
func TestNoWildcardDir(t *testing.T) {
|
||||
// Empty string means no wildcard dir — pure ACME mode.
|
||||
mgr, err := NewManager(ManagerConfig{CertDir: t.TempDir(), ACMEURL: "https://acme.example.com/directory"}, nil, nil, nil)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, mgr.WildcardPatterns())
|
||||
}
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
// ErrHeaderAuthFailed indicates that the header was present but the
|
||||
// credential did not validate. Callers should return 401 instead of
|
||||
// falling through to other auth schemes.
|
||||
var ErrHeaderAuthFailed = errors.New("header authentication failed")
|
||||
|
||||
// Header implements header-based authentication. The proxy checks for the
|
||||
// configured header in each request and validates its value via gRPC.
|
||||
type Header struct {
|
||||
id types.ServiceID
|
||||
accountId types.AccountID
|
||||
headerName string
|
||||
client authenticator
|
||||
}
|
||||
|
||||
// NewHeader creates a Header authentication scheme for the given header name.
|
||||
func NewHeader(client authenticator, id types.ServiceID, accountId types.AccountID, headerName string) Header {
|
||||
return Header{
|
||||
id: id,
|
||||
accountId: accountId,
|
||||
headerName: headerName,
|
||||
client: client,
|
||||
}
|
||||
}
|
||||
|
||||
// Type returns auth.MethodHeader.
|
||||
func (Header) Type() auth.Method {
|
||||
return auth.MethodHeader
|
||||
}
|
||||
|
||||
// Authenticate checks for the configured header in the request. If absent,
|
||||
// returns empty (unauthenticated). If present, validates via gRPC.
|
||||
func (h Header) Authenticate(r *http.Request) (string, string, error) {
|
||||
value := r.Header.Get(h.headerName)
|
||||
if value == "" {
|
||||
return "", "", nil
|
||||
}
|
||||
|
||||
res, err := h.client.Authenticate(r.Context(), &proto.AuthenticateRequest{
|
||||
Id: string(h.id),
|
||||
AccountId: string(h.accountId),
|
||||
Request: &proto.AuthenticateRequest_HeaderAuth{
|
||||
HeaderAuth: &proto.HeaderAuthRequest{
|
||||
HeaderValue: value,
|
||||
HeaderName: h.headerName,
|
||||
},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("authenticate header: %w", err)
|
||||
}
|
||||
|
||||
if res.GetSuccess() {
|
||||
return res.GetSessionToken(), "", nil
|
||||
}
|
||||
|
||||
return "", "", ErrHeaderAuthFailed
|
||||
}
|
||||
@@ -4,9 +4,12 @@ import (
|
||||
"context"
|
||||
"crypto/ed25519"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"html"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -16,11 +19,16 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/proxy"
|
||||
"github.com/netbirdio/netbird/proxy/internal/restrict"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/proxy/web"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
// errValidationUnavailable indicates that session validation failed due to
|
||||
// an infrastructure error (e.g. gRPC unavailable), not an invalid token.
|
||||
var errValidationUnavailable = errors.New("session validation unavailable")
|
||||
|
||||
type authenticator interface {
|
||||
Authenticate(ctx context.Context, in *proto.AuthenticateRequest, opts ...grpc.CallOption) (*proto.AuthenticateResponse, error)
|
||||
}
|
||||
@@ -40,12 +48,14 @@ type Scheme interface {
|
||||
Authenticate(*http.Request) (token string, promptData string, err error)
|
||||
}
|
||||
|
||||
// DomainConfig holds the authentication and restriction settings for a protected domain.
|
||||
type DomainConfig struct {
|
||||
Schemes []Scheme
|
||||
SessionPublicKey ed25519.PublicKey
|
||||
SessionExpiration time.Duration
|
||||
AccountID string
|
||||
ServiceID string
|
||||
AccountID types.AccountID
|
||||
ServiceID types.ServiceID
|
||||
IPRestrictions *restrict.Filter
|
||||
}
|
||||
|
||||
type validationResult struct {
|
||||
@@ -54,17 +64,18 @@ type validationResult struct {
|
||||
DeniedReason string
|
||||
}
|
||||
|
||||
// Middleware applies per-domain authentication and IP restriction checks.
|
||||
type Middleware struct {
|
||||
domainsMux sync.RWMutex
|
||||
domains map[string]DomainConfig
|
||||
logger *log.Logger
|
||||
sessionValidator SessionValidator
|
||||
geo restrict.GeoResolver
|
||||
}
|
||||
|
||||
// NewMiddleware creates a new authentication middleware.
|
||||
// The sessionValidator is optional; if nil, OIDC session tokens will be validated
|
||||
// locally without group access checks.
|
||||
func NewMiddleware(logger *log.Logger, sessionValidator SessionValidator) *Middleware {
|
||||
// NewMiddleware creates a new authentication middleware. The sessionValidator is
|
||||
// optional; if nil, OIDC session tokens are validated locally without group access checks.
|
||||
func NewMiddleware(logger *log.Logger, sessionValidator SessionValidator, geo restrict.GeoResolver) *Middleware {
|
||||
if logger == nil {
|
||||
logger = log.StandardLogger()
|
||||
}
|
||||
@@ -72,18 +83,12 @@ func NewMiddleware(logger *log.Logger, sessionValidator SessionValidator) *Middl
|
||||
domains: make(map[string]DomainConfig),
|
||||
logger: logger,
|
||||
sessionValidator: sessionValidator,
|
||||
geo: geo,
|
||||
}
|
||||
}
|
||||
|
||||
// Protect applies authentication middleware to the passed handler.
|
||||
// For each incoming request it will be checked against the middleware's
|
||||
// internal list of protected domains.
|
||||
// If the Host domain in the inbound request is not present, then it will
|
||||
// simply be passed through.
|
||||
// However, if the Host domain is present, then the specified authentication
|
||||
// schemes for that domain will be applied to the request.
|
||||
// In the event that no authentication schemes are defined for the domain,
|
||||
// then the request will also be simply passed through.
|
||||
// Protect wraps next with per-domain authentication and IP restriction checks.
|
||||
// Requests whose Host is not registered pass through unchanged.
|
||||
func (mw *Middleware) Protect(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
host, _, err := net.SplitHostPort(r.Host)
|
||||
@@ -94,8 +99,7 @@ func (mw *Middleware) Protect(next http.Handler) http.Handler {
|
||||
config, exists := mw.getDomainConfig(host)
|
||||
mw.logger.Debugf("checking authentication for host: %s, exists: %t", host, exists)
|
||||
|
||||
// Domains that are not configured here or have no authentication schemes applied should simply pass through.
|
||||
if !exists || len(config.Schemes) == 0 {
|
||||
if !exists {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
@@ -103,6 +107,16 @@ func (mw *Middleware) Protect(next http.Handler) http.Handler {
|
||||
// Set account and service IDs in captured data for access logging.
|
||||
setCapturedIDs(r, config)
|
||||
|
||||
if !mw.checkIPRestrictions(w, r, config) {
|
||||
return
|
||||
}
|
||||
|
||||
// Domains with no authentication schemes pass through after IP checks.
|
||||
if len(config.Schemes) == 0 {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
if mw.handleOAuthCallbackError(w, r) {
|
||||
return
|
||||
}
|
||||
@@ -111,6 +125,10 @@ func (mw *Middleware) Protect(next http.Handler) http.Handler {
|
||||
return
|
||||
}
|
||||
|
||||
if mw.forwardWithHeaderAuth(w, r, host, config, next) {
|
||||
return
|
||||
}
|
||||
|
||||
mw.authenticateWithSchemes(w, r, host, config)
|
||||
})
|
||||
}
|
||||
@@ -124,11 +142,65 @@ func (mw *Middleware) getDomainConfig(host string) (DomainConfig, bool) {
|
||||
|
||||
func setCapturedIDs(r *http.Request, config DomainConfig) {
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
cd.SetAccountId(types.AccountID(config.AccountID))
|
||||
cd.SetServiceId(config.ServiceID)
|
||||
cd.SetAccountID(config.AccountID)
|
||||
cd.SetServiceID(config.ServiceID)
|
||||
}
|
||||
}
|
||||
|
||||
// checkIPRestrictions validates the client IP against the domain's IP restrictions.
|
||||
// Uses the resolved client IP from CapturedData (which accounts for trusted proxies)
|
||||
// rather than r.RemoteAddr directly.
|
||||
func (mw *Middleware) checkIPRestrictions(w http.ResponseWriter, r *http.Request, config DomainConfig) bool {
|
||||
if config.IPRestrictions == nil {
|
||||
return true
|
||||
}
|
||||
|
||||
clientIP := mw.resolveClientIP(r)
|
||||
if !clientIP.IsValid() {
|
||||
mw.logger.Debugf("IP restriction: cannot resolve client address for %q, denying", r.RemoteAddr)
|
||||
http.Error(w, "Forbidden", http.StatusForbidden)
|
||||
return false
|
||||
}
|
||||
|
||||
verdict := config.IPRestrictions.Check(clientIP, mw.geo)
|
||||
if verdict == restrict.Allow {
|
||||
return true
|
||||
}
|
||||
|
||||
reason := verdict.String()
|
||||
mw.blockIPRestriction(r, reason)
|
||||
http.Error(w, "Forbidden", http.StatusForbidden)
|
||||
return false
|
||||
}
|
||||
|
||||
// resolveClientIP extracts the real client IP from CapturedData, falling back to r.RemoteAddr.
|
||||
func (mw *Middleware) resolveClientIP(r *http.Request) netip.Addr {
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
if ip := cd.GetClientIP(); ip.IsValid() {
|
||||
return ip
|
||||
}
|
||||
}
|
||||
|
||||
clientIPStr, _, _ := net.SplitHostPort(r.RemoteAddr)
|
||||
if clientIPStr == "" {
|
||||
clientIPStr = r.RemoteAddr
|
||||
}
|
||||
addr, err := netip.ParseAddr(clientIPStr)
|
||||
if err != nil {
|
||||
return netip.Addr{}
|
||||
}
|
||||
return addr.Unmap()
|
||||
}
|
||||
|
||||
// blockIPRestriction sets captured data fields for an IP-restriction block event.
|
||||
func (mw *Middleware) blockIPRestriction(r *http.Request, reason string) {
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
cd.SetOrigin(proxy.OriginAuth)
|
||||
cd.SetAuthMethod(reason)
|
||||
}
|
||||
mw.logger.Debugf("IP restriction: %s for %s", reason, r.RemoteAddr)
|
||||
}
|
||||
|
||||
// handleOAuthCallbackError checks for error query parameters from an OAuth
|
||||
// callback and renders the access denied page if present.
|
||||
func (mw *Middleware) handleOAuthCallbackError(w http.ResponseWriter, r *http.Request) bool {
|
||||
@@ -146,6 +218,8 @@ func (mw *Middleware) handleOAuthCallbackError(w http.ResponseWriter, r *http.Re
|
||||
errDesc := r.URL.Query().Get("error_description")
|
||||
if errDesc == "" {
|
||||
errDesc = "An error occurred during authentication"
|
||||
} else {
|
||||
errDesc = html.EscapeString(errDesc)
|
||||
}
|
||||
web.ServeAccessDeniedPage(w, r, http.StatusForbidden, "Access Denied", errDesc, requestID)
|
||||
return true
|
||||
@@ -170,6 +244,85 @@ func (mw *Middleware) forwardWithSessionCookie(w http.ResponseWriter, r *http.Re
|
||||
return true
|
||||
}
|
||||
|
||||
// forwardWithHeaderAuth checks for a Header auth scheme. If the header validates,
|
||||
// the request is forwarded directly (no redirect), which is important for API clients.
|
||||
func (mw *Middleware) forwardWithHeaderAuth(w http.ResponseWriter, r *http.Request, host string, config DomainConfig, next http.Handler) bool {
|
||||
for _, scheme := range config.Schemes {
|
||||
hdr, ok := scheme.(Header)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
handled := mw.tryHeaderScheme(w, r, host, config, hdr, next)
|
||||
if handled {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (mw *Middleware) tryHeaderScheme(w http.ResponseWriter, r *http.Request, host string, config DomainConfig, hdr Header, next http.Handler) bool {
|
||||
token, _, err := hdr.Authenticate(r)
|
||||
if err != nil {
|
||||
return mw.handleHeaderAuthError(w, r, err)
|
||||
}
|
||||
if token == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
result, err := mw.validateSessionToken(r.Context(), host, token, config.SessionPublicKey, auth.MethodHeader)
|
||||
if err != nil {
|
||||
setHeaderCapturedData(r.Context(), "")
|
||||
status := http.StatusBadRequest
|
||||
msg := "invalid session token"
|
||||
if errors.Is(err, errValidationUnavailable) {
|
||||
status = http.StatusBadGateway
|
||||
msg = "authentication service unavailable"
|
||||
}
|
||||
http.Error(w, msg, status)
|
||||
return true
|
||||
}
|
||||
|
||||
if !result.Valid {
|
||||
setHeaderCapturedData(r.Context(), result.UserID)
|
||||
http.Error(w, "Unauthorized", http.StatusUnauthorized)
|
||||
return true
|
||||
}
|
||||
|
||||
setSessionCookie(w, token, config.SessionExpiration)
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
cd.SetUserID(result.UserID)
|
||||
cd.SetAuthMethod(auth.MethodHeader.String())
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
return true
|
||||
}
|
||||
|
||||
func (mw *Middleware) handleHeaderAuthError(w http.ResponseWriter, r *http.Request, err error) bool {
|
||||
if errors.Is(err, ErrHeaderAuthFailed) {
|
||||
setHeaderCapturedData(r.Context(), "")
|
||||
http.Error(w, "Unauthorized", http.StatusUnauthorized)
|
||||
return true
|
||||
}
|
||||
mw.logger.WithField("scheme", "header").Warnf("header auth infrastructure error: %v", err)
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
cd.SetOrigin(proxy.OriginAuth)
|
||||
}
|
||||
http.Error(w, "authentication service unavailable", http.StatusBadGateway)
|
||||
return true
|
||||
}
|
||||
|
||||
func setHeaderCapturedData(ctx context.Context, userID string) {
|
||||
cd := proxy.CapturedDataFromContext(ctx)
|
||||
if cd == nil {
|
||||
return
|
||||
}
|
||||
cd.SetOrigin(proxy.OriginAuth)
|
||||
cd.SetAuthMethod(auth.MethodHeader.String())
|
||||
cd.SetUserID(userID)
|
||||
}
|
||||
|
||||
// authenticateWithSchemes tries each configured auth scheme in order.
|
||||
// On success it sets a session cookie and redirects; on failure it renders the login page.
|
||||
func (mw *Middleware) authenticateWithSchemes(w http.ResponseWriter, r *http.Request, host string, config DomainConfig) {
|
||||
@@ -217,7 +370,13 @@ func (mw *Middleware) handleAuthenticatedToken(w http.ResponseWriter, r *http.Re
|
||||
cd.SetOrigin(proxy.OriginAuth)
|
||||
cd.SetAuthMethod(scheme.Type().String())
|
||||
}
|
||||
http.Error(w, err.Error(), http.StatusBadRequest)
|
||||
status := http.StatusBadRequest
|
||||
msg := "invalid session token"
|
||||
if errors.Is(err, errValidationUnavailable) {
|
||||
status = http.StatusBadGateway
|
||||
msg = "authentication service unavailable"
|
||||
}
|
||||
http.Error(w, msg, status)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -233,7 +392,21 @@ func (mw *Middleware) handleAuthenticatedToken(w http.ResponseWriter, r *http.Re
|
||||
return
|
||||
}
|
||||
|
||||
expiration := config.SessionExpiration
|
||||
setSessionCookie(w, token, config.SessionExpiration)
|
||||
|
||||
// Redirect instead of forwarding the auth POST to the backend.
|
||||
// The browser will follow with a GET carrying the new session cookie.
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
cd.SetOrigin(proxy.OriginAuth)
|
||||
cd.SetUserID(result.UserID)
|
||||
cd.SetAuthMethod(scheme.Type().String())
|
||||
}
|
||||
redirectURL := stripSessionTokenParam(r.URL)
|
||||
http.Redirect(w, r, redirectURL, http.StatusSeeOther)
|
||||
}
|
||||
|
||||
// setSessionCookie writes a session cookie with secure defaults.
|
||||
func setSessionCookie(w http.ResponseWriter, token string, expiration time.Duration) {
|
||||
if expiration == 0 {
|
||||
expiration = auth.DefaultSessionExpiry
|
||||
}
|
||||
@@ -245,16 +418,6 @@ func (mw *Middleware) handleAuthenticatedToken(w http.ResponseWriter, r *http.Re
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
MaxAge: int(expiration.Seconds()),
|
||||
})
|
||||
|
||||
// Redirect instead of forwarding the auth POST to the backend.
|
||||
// The browser will follow with a GET carrying the new session cookie.
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
cd.SetOrigin(proxy.OriginAuth)
|
||||
cd.SetUserID(result.UserID)
|
||||
cd.SetAuthMethod(scheme.Type().String())
|
||||
}
|
||||
redirectURL := stripSessionTokenParam(r.URL)
|
||||
http.Redirect(w, r, redirectURL, http.StatusSeeOther)
|
||||
}
|
||||
|
||||
// wasCredentialSubmitted checks if credentials were submitted for the given auth method.
|
||||
@@ -275,13 +438,14 @@ func wasCredentialSubmitted(r *http.Request, method auth.Method) bool {
|
||||
// session JWTs. Returns an error if the key is missing or invalid.
|
||||
// Callers must not serve the domain if this returns an error, to avoid
|
||||
// exposing an unauthenticated service.
|
||||
func (mw *Middleware) AddDomain(domain string, schemes []Scheme, publicKeyB64 string, expiration time.Duration, accountID, serviceID string) error {
|
||||
func (mw *Middleware) AddDomain(domain string, schemes []Scheme, publicKeyB64 string, expiration time.Duration, accountID types.AccountID, serviceID types.ServiceID, ipRestrictions *restrict.Filter) error {
|
||||
if len(schemes) == 0 {
|
||||
mw.domainsMux.Lock()
|
||||
defer mw.domainsMux.Unlock()
|
||||
mw.domains[domain] = DomainConfig{
|
||||
AccountID: accountID,
|
||||
ServiceID: serviceID,
|
||||
AccountID: accountID,
|
||||
ServiceID: serviceID,
|
||||
IPRestrictions: ipRestrictions,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -302,30 +466,28 @@ func (mw *Middleware) AddDomain(domain string, schemes []Scheme, publicKeyB64 st
|
||||
SessionExpiration: expiration,
|
||||
AccountID: accountID,
|
||||
ServiceID: serviceID,
|
||||
IPRestrictions: ipRestrictions,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveDomain unregisters authentication for the given domain.
|
||||
func (mw *Middleware) RemoveDomain(domain string) {
|
||||
mw.domainsMux.Lock()
|
||||
defer mw.domainsMux.Unlock()
|
||||
delete(mw.domains, domain)
|
||||
}
|
||||
|
||||
// validateSessionToken validates a session token, optionally checking group access via gRPC.
|
||||
// For OIDC tokens with a configured validator, it calls ValidateSession to check group access.
|
||||
// For other auth methods (PIN, password), it validates the JWT locally.
|
||||
// Returns a validationResult with user ID and validity status, or error for invalid tokens.
|
||||
// validateSessionToken validates a session token. OIDC tokens with a configured
|
||||
// validator go through gRPC for group access checks; other methods validate locally.
|
||||
func (mw *Middleware) validateSessionToken(ctx context.Context, host, token string, publicKey ed25519.PublicKey, method auth.Method) (*validationResult, error) {
|
||||
// For OIDC with a session validator, call the gRPC service to check group access
|
||||
if method == auth.MethodOIDC && mw.sessionValidator != nil {
|
||||
resp, err := mw.sessionValidator.ValidateSession(ctx, &proto.ValidateSessionRequest{
|
||||
Domain: host,
|
||||
SessionToken: token,
|
||||
})
|
||||
if err != nil {
|
||||
mw.logger.WithError(err).Error("ValidateSession gRPC call failed")
|
||||
return nil, fmt.Errorf("session validation failed")
|
||||
return nil, fmt.Errorf("%w: %w", errValidationUnavailable, err)
|
||||
}
|
||||
if !resp.Valid {
|
||||
mw.logger.WithFields(log.Fields{
|
||||
@@ -342,7 +504,6 @@ func (mw *Middleware) validateSessionToken(ctx context.Context, host, token stri
|
||||
return &validationResult{UserID: resp.UserId, Valid: true}, nil
|
||||
}
|
||||
|
||||
// For non-OIDC methods or when no validator is configured, validate JWT locally
|
||||
userID, _, err := auth.ValidateSessionJWT(token, host, publicKey)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -14,10 +17,13 @@ import (
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/proxy"
|
||||
"github.com/netbirdio/netbird/proxy/internal/restrict"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
func generateTestKeyPair(t *testing.T) *sessionkey.KeyPair {
|
||||
@@ -52,11 +58,11 @@ func newPassthroughHandler() http.Handler {
|
||||
}
|
||||
|
||||
func TestAddDomain_ValidKey(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "")
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil)
|
||||
require.NoError(t, err)
|
||||
|
||||
mw.domainsMux.RLock()
|
||||
@@ -70,10 +76,10 @@ func TestAddDomain_ValidKey(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAddDomain_EmptyKey(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, "", time.Hour, "", "")
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, "", time.Hour, "", "", nil)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid session public key size")
|
||||
|
||||
@@ -84,10 +90,10 @@ func TestAddDomain_EmptyKey(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAddDomain_InvalidBase64(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, "not-valid-base64!!!", time.Hour, "", "")
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, "not-valid-base64!!!", time.Hour, "", "", nil)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "decode session public key")
|
||||
|
||||
@@ -98,11 +104,11 @@ func TestAddDomain_InvalidBase64(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAddDomain_WrongKeySize(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
shortKey := base64.StdEncoding.EncodeToString([]byte("tooshort"))
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, shortKey, time.Hour, "", "")
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, shortKey, time.Hour, "", "", nil)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid session public key size")
|
||||
|
||||
@@ -113,9 +119,9 @@ func TestAddDomain_WrongKeySize(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAddDomain_NoSchemes_NoKeyRequired(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
err := mw.AddDomain("example.com", nil, "", time.Hour, "", "")
|
||||
err := mw.AddDomain("example.com", nil, "", time.Hour, "", "", nil)
|
||||
require.NoError(t, err, "domains with no auth schemes should not require a key")
|
||||
|
||||
mw.domainsMux.RLock()
|
||||
@@ -125,14 +131,14 @@ func TestAddDomain_NoSchemes_NoKeyRequired(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAddDomain_OverwritesPreviousConfig(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp1 := generateTestKeyPair(t)
|
||||
kp2 := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp1.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp2.PublicKey, 2*time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp1.PublicKey, time.Hour, "", "", nil))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp2.PublicKey, 2*time.Hour, "", "", nil))
|
||||
|
||||
mw.domainsMux.RLock()
|
||||
config := mw.domains["example.com"]
|
||||
@@ -144,11 +150,11 @@ func TestAddDomain_OverwritesPreviousConfig(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRemoveDomain(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
mw.RemoveDomain("example.com")
|
||||
|
||||
@@ -159,7 +165,7 @@ func TestRemoveDomain(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_UnknownDomainPassesThrough(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://unknown.com/", nil)
|
||||
@@ -171,8 +177,8 @@ func TestProtect_UnknownDomainPassesThrough(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_DomainWithNoSchemesPassesThrough(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
require.NoError(t, mw.AddDomain("example.com", nil, "", time.Hour, "", ""))
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
require.NoError(t, mw.AddDomain("example.com", nil, "", time.Hour, "", "", nil))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -185,11 +191,11 @@ func TestProtect_DomainWithNoSchemesPassesThrough(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_UnauthenticatedRequestIsBlocked(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
var backendCalled bool
|
||||
backend := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -206,11 +212,11 @@ func TestProtect_UnauthenticatedRequestIsBlocked(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_HostWithPortIsMatched(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
var backendCalled bool
|
||||
backend := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -227,16 +233,16 @@ func TestProtect_HostWithPortIsMatched(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_ValidSessionCookiePassesThrough(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "test-user", "example.com", auth.MethodPIN, time.Hour)
|
||||
require.NoError(t, err)
|
||||
|
||||
capturedData := &proxy.CapturedData{}
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
cd := proxy.CapturedDataFromContext(r.Context())
|
||||
require.NotNil(t, cd)
|
||||
@@ -257,11 +263,11 @@ func TestProtect_ValidSessionCookiePassesThrough(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_ExpiredSessionCookieIsRejected(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
// Sign a token that expired 1 second ago.
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "test-user", "example.com", auth.MethodPIN, -time.Second)
|
||||
@@ -283,11 +289,11 @@ func TestProtect_ExpiredSessionCookieIsRejected(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_WrongDomainCookieIsRejected(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
// Token signed for a different domain audience.
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "test-user", "other.com", auth.MethodPIN, time.Hour)
|
||||
@@ -309,12 +315,12 @@ func TestProtect_WrongDomainCookieIsRejected(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_WrongKeyCookieIsRejected(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp1 := generateTestKeyPair(t)
|
||||
kp2 := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp1.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp1.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
// Token signed with a different private key.
|
||||
token, err := sessionkey.SignToken(kp2.PrivateKey, "test-user", "example.com", auth.MethodPIN, time.Hour)
|
||||
@@ -336,7 +342,7 @@ func TestProtect_WrongKeyCookieIsRejected(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_SchemeAuthRedirectsWithCookie(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "pin-user", "example.com", auth.MethodPIN, time.Hour)
|
||||
@@ -351,7 +357,7 @@ func TestProtect_SchemeAuthRedirectsWithCookie(t *testing.T) {
|
||||
return "", "pin", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
var backendCalled bool
|
||||
backend := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -386,7 +392,7 @@ func TestProtect_SchemeAuthRedirectsWithCookie(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_FailedAuthDoesNotSetCookie(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{
|
||||
@@ -395,7 +401,7 @@ func TestProtect_FailedAuthDoesNotSetCookie(t *testing.T) {
|
||||
return "", "pin", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -409,7 +415,7 @@ func TestProtect_FailedAuthDoesNotSetCookie(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_MultipleSchemes(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "password-user", "example.com", auth.MethodPassword, time.Hour)
|
||||
@@ -431,7 +437,7 @@ func TestProtect_MultipleSchemes(t *testing.T) {
|
||||
return "", "password", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{pinScheme, passwordScheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{pinScheme, passwordScheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
var backendCalled bool
|
||||
backend := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -451,7 +457,7 @@ func TestProtect_MultipleSchemes(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_InvalidTokenFromSchemeReturns400(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
// Return a garbage token that won't validate.
|
||||
@@ -461,7 +467,7 @@ func TestProtect_InvalidTokenFromSchemeReturns400(t *testing.T) {
|
||||
return "invalid-jwt-token", "", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -473,7 +479,7 @@ func TestProtect_InvalidTokenFromSchemeReturns400(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestAddDomain_RandomBytes32NotEd25519(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
// 32 random bytes that happen to be valid base64 and correct size
|
||||
// but are actually a valid ed25519 public key length-wise.
|
||||
@@ -485,19 +491,19 @@ func TestAddDomain_RandomBytes32NotEd25519(t *testing.T) {
|
||||
key := base64.StdEncoding.EncodeToString(randomBytes)
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
|
||||
err = mw.AddDomain("example.com", []Scheme{scheme}, key, time.Hour, "", "")
|
||||
err = mw.AddDomain("example.com", []Scheme{scheme}, key, time.Hour, "", "", nil)
|
||||
require.NoError(t, err, "any 32-byte key should be accepted at registration time")
|
||||
}
|
||||
|
||||
func TestAddDomain_InvalidKeyDoesNotCorruptExistingConfig(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
// Attempt to overwrite with an invalid key.
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, "bad", time.Hour, "", "")
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, "bad", time.Hour, "", "", nil)
|
||||
require.Error(t, err)
|
||||
|
||||
// The original valid config should still be intact.
|
||||
@@ -511,7 +517,7 @@ func TestAddDomain_InvalidKeyDoesNotCorruptExistingConfig(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_FailedPinAuthCapturesAuthMethod(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
// Scheme that always fails authentication (returns empty token)
|
||||
@@ -521,9 +527,9 @@ func TestProtect_FailedPinAuthCapturesAuthMethod(t *testing.T) {
|
||||
return "", "pin", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
capturedData := &proxy.CapturedData{}
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
// Submit wrong PIN - should capture auth method
|
||||
@@ -539,7 +545,7 @@ func TestProtect_FailedPinAuthCapturesAuthMethod(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_FailedPasswordAuthCapturesAuthMethod(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{
|
||||
@@ -548,9 +554,9 @@ func TestProtect_FailedPasswordAuthCapturesAuthMethod(t *testing.T) {
|
||||
return "", "password", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
capturedData := &proxy.CapturedData{}
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
// Submit wrong password - should capture auth method
|
||||
@@ -566,7 +572,7 @@ func TestProtect_FailedPasswordAuthCapturesAuthMethod(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestProtect_NoCredentialsDoesNotCaptureAuthMethod(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil)
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{
|
||||
@@ -575,9 +581,9 @@ func TestProtect_NoCredentialsDoesNotCaptureAuthMethod(t *testing.T) {
|
||||
return "", "pin", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", ""))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil))
|
||||
|
||||
capturedData := &proxy.CapturedData{}
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
// No credentials submitted - should not capture auth method
|
||||
@@ -658,3 +664,271 @@ func TestWasCredentialSubmitted(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckIPRestrictions_UnparseableAddress(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
err := mw.AddDomain("example.com", nil, "", 0, "acc1", "svc1",
|
||||
restrict.ParseFilter([]string{"10.0.0.0/8"}, nil, nil, nil))
|
||||
require.NoError(t, err)
|
||||
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
remoteAddr string
|
||||
wantCode int
|
||||
}{
|
||||
{"unparsable address denies", "not-an-ip:1234", http.StatusForbidden},
|
||||
{"empty address denies", "", http.StatusForbidden},
|
||||
{"allowed address passes", "10.1.2.3:5678", http.StatusOK},
|
||||
{"denied address blocked", "192.168.1.1:5678", http.StatusForbidden},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil)
|
||||
req.RemoteAddr = tt.remoteAddr
|
||||
req.Host = "example.com"
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
assert.Equal(t, tt.wantCode, rr.Code)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckIPRestrictions_UsesCapturedDataClientIP(t *testing.T) {
|
||||
// When CapturedData is set (by the access log middleware, which resolves
|
||||
// trusted proxies), checkIPRestrictions should use that IP, not RemoteAddr.
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
err := mw.AddDomain("example.com", nil, "", 0, "acc1", "svc1",
|
||||
restrict.ParseFilter([]string{"203.0.113.0/24"}, nil, nil, nil))
|
||||
require.NoError(t, err)
|
||||
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
|
||||
// RemoteAddr is a trusted proxy, but CapturedData has the real client IP.
|
||||
req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil)
|
||||
req.RemoteAddr = "10.0.0.1:5000"
|
||||
req.Host = "example.com"
|
||||
|
||||
cd := proxy.NewCapturedData("")
|
||||
cd.SetClientIP(netip.MustParseAddr("203.0.113.50"))
|
||||
ctx := proxy.WithCapturedData(req.Context(), cd)
|
||||
req = req.WithContext(ctx)
|
||||
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
assert.Equal(t, http.StatusOK, rr.Code, "should use CapturedData IP (203.0.113.50), not RemoteAddr (10.0.0.1)")
|
||||
|
||||
// Same request but CapturedData has a blocked IP.
|
||||
req2 := httptest.NewRequest(http.MethodGet, "http://example.com/", nil)
|
||||
req2.RemoteAddr = "203.0.113.50:5000"
|
||||
req2.Host = "example.com"
|
||||
|
||||
cd2 := proxy.NewCapturedData("")
|
||||
cd2.SetClientIP(netip.MustParseAddr("10.0.0.1"))
|
||||
ctx2 := proxy.WithCapturedData(req2.Context(), cd2)
|
||||
req2 = req2.WithContext(ctx2)
|
||||
|
||||
rr2 := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr2, req2)
|
||||
assert.Equal(t, http.StatusForbidden, rr2.Code, "should use CapturedData IP (10.0.0.1), not RemoteAddr (203.0.113.50)")
|
||||
}
|
||||
|
||||
func TestCheckIPRestrictions_NilGeoWithCountryRules(t *testing.T) {
|
||||
// Geo is nil, country restrictions are configured: must deny (fail-close).
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
err := mw.AddDomain("example.com", nil, "", 0, "acc1", "svc1",
|
||||
restrict.ParseFilter(nil, nil, []string{"US"}, nil))
|
||||
require.NoError(t, err)
|
||||
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil)
|
||||
req.RemoteAddr = "1.2.3.4:5678"
|
||||
req.Host = "example.com"
|
||||
rr := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rr, req)
|
||||
assert.Equal(t, http.StatusForbidden, rr.Code, "country restrictions with nil geo must deny")
|
||||
}
|
||||
|
||||
// mockAuthenticator is a minimal mock for the authenticator gRPC interface
|
||||
// used by the Header scheme.
|
||||
type mockAuthenticator struct {
|
||||
fn func(ctx context.Context, req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error)
|
||||
}
|
||||
|
||||
func (m *mockAuthenticator) Authenticate(ctx context.Context, in *proto.AuthenticateRequest, _ ...grpc.CallOption) (*proto.AuthenticateResponse, error) {
|
||||
return m.fn(ctx, in)
|
||||
}
|
||||
|
||||
// newHeaderSchemeWithToken creates a Header scheme backed by a mock that
|
||||
// returns a signed session token when the expected header value is provided.
|
||||
func newHeaderSchemeWithToken(t *testing.T, kp *sessionkey.KeyPair, headerName, expectedValue string) Header {
|
||||
t.Helper()
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "header-user", "example.com", auth.MethodHeader, time.Hour)
|
||||
require.NoError(t, err)
|
||||
|
||||
mock := &mockAuthenticator{fn: func(_ context.Context, req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
|
||||
ha := req.GetHeaderAuth()
|
||||
if ha != nil && ha.GetHeaderValue() == expectedValue {
|
||||
return &proto.AuthenticateResponse{Success: true, SessionToken: token}, nil
|
||||
}
|
||||
return &proto.AuthenticateResponse{Success: false}, nil
|
||||
}}
|
||||
return NewHeader(mock, "svc1", "acc1", headerName)
|
||||
}
|
||||
|
||||
func TestProtect_HeaderAuth_ForwardsOnSuccess(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
hdr := newHeaderSchemeWithToken(t, kp, "X-API-Key", "secret-key")
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil))
|
||||
|
||||
var backendCalled bool
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
backendCalled = true
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write([]byte("ok"))
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://example.com/path", nil)
|
||||
req.Header.Set("X-API-Key", "secret-key")
|
||||
req = req.WithContext(proxy.WithCapturedData(req.Context(), capturedData))
|
||||
rec := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec, req)
|
||||
|
||||
assert.True(t, backendCalled, "backend should be called directly for header auth (no redirect)")
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
assert.Equal(t, "ok", rec.Body.String())
|
||||
|
||||
// Session cookie should be set.
|
||||
var sessionCookie *http.Cookie
|
||||
for _, c := range rec.Result().Cookies() {
|
||||
if c.Name == auth.SessionCookieName {
|
||||
sessionCookie = c
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotNil(t, sessionCookie, "session cookie should be set after successful header auth")
|
||||
assert.True(t, sessionCookie.HttpOnly)
|
||||
assert.True(t, sessionCookie.Secure)
|
||||
|
||||
assert.Equal(t, "header-user", capturedData.GetUserID())
|
||||
assert.Equal(t, "header", capturedData.GetAuthMethod())
|
||||
}
|
||||
|
||||
func TestProtect_HeaderAuth_MissingHeaderFallsThrough(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
hdr := newHeaderSchemeWithToken(t, kp, "X-API-Key", "secret-key")
|
||||
// Also add a PIN scheme so we can verify fallthrough behavior.
|
||||
pinScheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr, pinScheme}, kp.PublicKey, time.Hour, "acc1", "svc1", nil))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
// No X-API-Key header: should fall through to PIN login page (401).
|
||||
req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil)
|
||||
rec := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, rec.Code, "missing header should fall through to login page")
|
||||
}
|
||||
|
||||
func TestProtect_HeaderAuth_WrongValueReturns401(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
mock := &mockAuthenticator{fn: func(_ context.Context, _ *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
|
||||
return &proto.AuthenticateResponse{Success: false}, nil
|
||||
}}
|
||||
hdr := NewHeader(mock, "svc1", "acc1", "X-API-Key")
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil))
|
||||
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil)
|
||||
req.Header.Set("X-API-Key", "wrong-key")
|
||||
req = req.WithContext(proxy.WithCapturedData(req.Context(), capturedData))
|
||||
rec := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, rec.Code)
|
||||
assert.Equal(t, "header", capturedData.GetAuthMethod())
|
||||
}
|
||||
|
||||
func TestProtect_HeaderAuth_InfraErrorReturns502(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
mock := &mockAuthenticator{fn: func(_ context.Context, _ *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
|
||||
return nil, errors.New("gRPC unavailable")
|
||||
}}
|
||||
hdr := NewHeader(mock, "svc1", "acc1", "X-API-Key")
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://example.com/", nil)
|
||||
req.Header.Set("X-API-Key", "some-key")
|
||||
rec := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadGateway, rec.Code)
|
||||
}
|
||||
|
||||
func TestProtect_HeaderAuth_SubsequentRequestUsesSessionCookie(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
hdr := newHeaderSchemeWithToken(t, kp, "X-API-Key", "secret-key")
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil))
|
||||
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
|
||||
// First request with header auth.
|
||||
req1 := httptest.NewRequest(http.MethodGet, "http://example.com/", nil)
|
||||
req1.Header.Set("X-API-Key", "secret-key")
|
||||
req1 = req1.WithContext(proxy.WithCapturedData(req1.Context(), proxy.NewCapturedData("")))
|
||||
rec1 := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec1, req1)
|
||||
require.Equal(t, http.StatusOK, rec1.Code)
|
||||
|
||||
// Extract session cookie.
|
||||
var sessionCookie *http.Cookie
|
||||
for _, c := range rec1.Result().Cookies() {
|
||||
if c.Name == auth.SessionCookieName {
|
||||
sessionCookie = c
|
||||
break
|
||||
}
|
||||
}
|
||||
require.NotNil(t, sessionCookie)
|
||||
|
||||
// Second request with only the session cookie (no header).
|
||||
capturedData2 := proxy.NewCapturedData("")
|
||||
req2 := httptest.NewRequest(http.MethodGet, "http://example.com/other", nil)
|
||||
req2.AddCookie(sessionCookie)
|
||||
req2 = req2.WithContext(proxy.WithCapturedData(req2.Context(), capturedData2))
|
||||
rec2 := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec2, req2)
|
||||
|
||||
assert.Equal(t, http.StatusOK, rec2.Code)
|
||||
assert.Equal(t, "header-user", capturedData2.GetUserID())
|
||||
assert.Equal(t, "header", capturedData2.GetAuthMethod())
|
||||
}
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"google.golang.org/grpc"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
@@ -17,14 +18,14 @@ type urlGenerator interface {
|
||||
}
|
||||
|
||||
type OIDC struct {
|
||||
id string
|
||||
accountId string
|
||||
id types.ServiceID
|
||||
accountId types.AccountID
|
||||
forwardedProto string
|
||||
client urlGenerator
|
||||
}
|
||||
|
||||
// NewOIDC creates a new OIDC authentication scheme
|
||||
func NewOIDC(client urlGenerator, id, accountId, forwardedProto string) OIDC {
|
||||
func NewOIDC(client urlGenerator, id types.ServiceID, accountId types.AccountID, forwardedProto string) OIDC {
|
||||
return OIDC{
|
||||
id: id,
|
||||
accountId: accountId,
|
||||
@@ -53,8 +54,8 @@ func (o OIDC) Authenticate(r *http.Request) (string, string, error) {
|
||||
}
|
||||
|
||||
res, err := o.client.GetOIDCURL(r.Context(), &proto.GetOIDCURLRequest{
|
||||
Id: o.id,
|
||||
AccountId: o.accountId,
|
||||
Id: string(o.id),
|
||||
AccountId: string(o.accountId),
|
||||
RedirectUrl: redirectURL.String(),
|
||||
})
|
||||
if err != nil {
|
||||
|
||||
@@ -5,17 +5,19 @@ import (
|
||||
"net/http"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
const passwordFormId = "password"
|
||||
|
||||
type Password struct {
|
||||
id, accountId string
|
||||
client authenticator
|
||||
id types.ServiceID
|
||||
accountId types.AccountID
|
||||
client authenticator
|
||||
}
|
||||
|
||||
func NewPassword(client authenticator, id, accountId string) Password {
|
||||
func NewPassword(client authenticator, id types.ServiceID, accountId types.AccountID) Password {
|
||||
return Password{
|
||||
id: id,
|
||||
accountId: accountId,
|
||||
@@ -41,8 +43,8 @@ func (p Password) Authenticate(r *http.Request) (string, string, error) {
|
||||
}
|
||||
|
||||
res, err := p.client.Authenticate(r.Context(), &proto.AuthenticateRequest{
|
||||
Id: p.id,
|
||||
AccountId: p.accountId,
|
||||
Id: string(p.id),
|
||||
AccountId: string(p.accountId),
|
||||
Request: &proto.AuthenticateRequest_Password{
|
||||
Password: &proto.PasswordRequest{
|
||||
Password: password,
|
||||
|
||||
@@ -5,17 +5,19 @@ import (
|
||||
"net/http"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
const pinFormId = "pin"
|
||||
|
||||
type Pin struct {
|
||||
id, accountId string
|
||||
client authenticator
|
||||
id types.ServiceID
|
||||
accountId types.AccountID
|
||||
client authenticator
|
||||
}
|
||||
|
||||
func NewPin(client authenticator, id, accountId string) Pin {
|
||||
func NewPin(client authenticator, id types.ServiceID, accountId types.AccountID) Pin {
|
||||
return Pin{
|
||||
id: id,
|
||||
accountId: accountId,
|
||||
@@ -41,8 +43,8 @@ func (p Pin) Authenticate(r *http.Request) (string, string, error) {
|
||||
}
|
||||
|
||||
res, err := p.client.Authenticate(r.Context(), &proto.AuthenticateRequest{
|
||||
Id: p.id,
|
||||
AccountId: p.accountId,
|
||||
Id: string(p.id),
|
||||
AccountId: string(p.accountId),
|
||||
Request: &proto.AuthenticateRequest_Pin{
|
||||
Pin: &proto.PinRequest{
|
||||
Pin: pin,
|
||||
|
||||
@@ -67,6 +67,13 @@ func (w *Watcher) GetCertificate(_ *tls.ClientHelloInfo) (*tls.Certificate, erro
|
||||
return w.cert, nil
|
||||
}
|
||||
|
||||
// Leaf returns the parsed leaf certificate, or nil if not yet loaded.
|
||||
func (w *Watcher) Leaf() *x509.Certificate {
|
||||
w.mu.RLock()
|
||||
defer w.mu.RUnlock()
|
||||
return w.leaf
|
||||
}
|
||||
|
||||
// Watch starts watching for certificate file changes. It blocks until
|
||||
// ctx is cancelled. It uses fsnotify for immediate detection and falls
|
||||
// back to polling if fsnotify is unavailable (e.g. on NFS).
|
||||
|
||||
@@ -10,10 +10,11 @@ import (
|
||||
type trackedConn struct {
|
||||
net.Conn
|
||||
tracker *HijackTracker
|
||||
host string
|
||||
}
|
||||
|
||||
func (c *trackedConn) Close() error {
|
||||
c.tracker.conns.Delete(c)
|
||||
c.tracker.remove(c)
|
||||
return c.Conn.Close()
|
||||
}
|
||||
|
||||
@@ -22,6 +23,7 @@ func (c *trackedConn) Close() error {
|
||||
type trackingWriter struct {
|
||||
http.ResponseWriter
|
||||
tracker *HijackTracker
|
||||
host string
|
||||
}
|
||||
|
||||
func (w *trackingWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
|
||||
@@ -33,8 +35,8 @@ func (w *trackingWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
tc := &trackedConn{Conn: conn, tracker: w.tracker}
|
||||
w.tracker.conns.Store(tc, struct{}{})
|
||||
tc := &trackedConn{Conn: conn, tracker: w.tracker, host: w.host}
|
||||
w.tracker.add(tc)
|
||||
return tc, buf, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package conntrack
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"sync"
|
||||
)
|
||||
@@ -10,10 +9,14 @@ import (
|
||||
// upgrades). http.Server.Shutdown does not close hijacked connections, so
|
||||
// they must be tracked and closed explicitly during graceful shutdown.
|
||||
//
|
||||
// Connections are indexed by the request Host so they can be closed
|
||||
// per-domain when a service mapping is removed.
|
||||
//
|
||||
// Use Middleware as the outermost HTTP middleware to ensure hijacked
|
||||
// connections are tracked and automatically deregistered when closed.
|
||||
type HijackTracker struct {
|
||||
conns sync.Map // net.Conn → struct{}
|
||||
mu sync.Mutex
|
||||
conns map[*trackedConn]struct{}
|
||||
}
|
||||
|
||||
// Middleware returns an HTTP middleware that wraps the ResponseWriter so that
|
||||
@@ -21,21 +24,73 @@ type HijackTracker struct {
|
||||
// tracker when closed. This should be the outermost middleware in the chain.
|
||||
func (t *HijackTracker) Middleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
next.ServeHTTP(&trackingWriter{ResponseWriter: w, tracker: t}, r)
|
||||
next.ServeHTTP(&trackingWriter{
|
||||
ResponseWriter: w,
|
||||
tracker: t,
|
||||
host: hostOnly(r.Host),
|
||||
}, r)
|
||||
})
|
||||
}
|
||||
|
||||
// CloseAll closes all tracked hijacked connections and returns the number
|
||||
// of connections that were closed.
|
||||
// CloseAll closes all tracked hijacked connections and returns the count.
|
||||
func (t *HijackTracker) CloseAll() int {
|
||||
var count int
|
||||
t.conns.Range(func(key, _ any) bool {
|
||||
if conn, ok := key.(net.Conn); ok {
|
||||
_ = conn.Close()
|
||||
count++
|
||||
t.mu.Lock()
|
||||
conns := t.conns
|
||||
t.conns = nil
|
||||
t.mu.Unlock()
|
||||
|
||||
for tc := range conns {
|
||||
_ = tc.Conn.Close()
|
||||
}
|
||||
return len(conns)
|
||||
}
|
||||
|
||||
// CloseByHost closes all tracked hijacked connections for the given host
|
||||
// and returns the number of connections closed.
|
||||
func (t *HijackTracker) CloseByHost(host string) int {
|
||||
host = hostOnly(host)
|
||||
t.mu.Lock()
|
||||
var toClose []*trackedConn
|
||||
for tc := range t.conns {
|
||||
if tc.host == host {
|
||||
toClose = append(toClose, tc)
|
||||
}
|
||||
t.conns.Delete(key)
|
||||
return true
|
||||
})
|
||||
return count
|
||||
}
|
||||
for _, tc := range toClose {
|
||||
delete(t.conns, tc)
|
||||
}
|
||||
t.mu.Unlock()
|
||||
|
||||
for _, tc := range toClose {
|
||||
_ = tc.Conn.Close()
|
||||
}
|
||||
return len(toClose)
|
||||
}
|
||||
|
||||
func (t *HijackTracker) add(tc *trackedConn) {
|
||||
t.mu.Lock()
|
||||
if t.conns == nil {
|
||||
t.conns = make(map[*trackedConn]struct{})
|
||||
}
|
||||
t.conns[tc] = struct{}{}
|
||||
t.mu.Unlock()
|
||||
}
|
||||
|
||||
func (t *HijackTracker) remove(tc *trackedConn) {
|
||||
t.mu.Lock()
|
||||
delete(t.conns, tc)
|
||||
t.mu.Unlock()
|
||||
}
|
||||
|
||||
// hostOnly strips the port from a host:port string.
|
||||
func hostOnly(hostport string) string {
|
||||
for i := len(hostport) - 1; i >= 0; i-- {
|
||||
if hostport[i] == ':' {
|
||||
return hostport[:i]
|
||||
}
|
||||
if hostport[i] < '0' || hostport[i] > '9' {
|
||||
return hostport
|
||||
}
|
||||
}
|
||||
return hostport
|
||||
}
|
||||
|
||||
@@ -0,0 +1,142 @@
|
||||
package conntrack
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// fakeHijackWriter implements http.ResponseWriter and http.Hijacker for testing.
|
||||
type fakeHijackWriter struct {
|
||||
http.ResponseWriter
|
||||
conn net.Conn
|
||||
}
|
||||
|
||||
func (f *fakeHijackWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
|
||||
rw := bufio.NewReadWriter(bufio.NewReader(f.conn), bufio.NewWriter(f.conn))
|
||||
return f.conn, rw, nil
|
||||
}
|
||||
|
||||
func TestCloseByHost(t *testing.T) {
|
||||
var tracker HijackTracker
|
||||
|
||||
// Simulate hijacking two connections for different hosts.
|
||||
connA1, connA2 := net.Pipe()
|
||||
defer connA2.Close()
|
||||
connB1, connB2 := net.Pipe()
|
||||
defer connB2.Close()
|
||||
|
||||
twA := &trackingWriter{
|
||||
ResponseWriter: httptest.NewRecorder(),
|
||||
tracker: &tracker,
|
||||
host: "a.example.com",
|
||||
}
|
||||
twB := &trackingWriter{
|
||||
ResponseWriter: httptest.NewRecorder(),
|
||||
tracker: &tracker,
|
||||
host: "b.example.com",
|
||||
}
|
||||
|
||||
// Use fakeHijackWriter to provide the Hijack method.
|
||||
twA.ResponseWriter = &fakeHijackWriter{ResponseWriter: twA.ResponseWriter, conn: connA1}
|
||||
twB.ResponseWriter = &fakeHijackWriter{ResponseWriter: twB.ResponseWriter, conn: connB1}
|
||||
|
||||
_, _, err := twA.Hijack()
|
||||
require.NoError(t, err)
|
||||
_, _, err = twB.Hijack()
|
||||
require.NoError(t, err)
|
||||
|
||||
tracker.mu.Lock()
|
||||
assert.Equal(t, 2, len(tracker.conns), "should track 2 connections")
|
||||
tracker.mu.Unlock()
|
||||
|
||||
// Close only host A.
|
||||
n := tracker.CloseByHost("a.example.com")
|
||||
assert.Equal(t, 1, n, "should close 1 connection for host A")
|
||||
|
||||
tracker.mu.Lock()
|
||||
assert.Equal(t, 1, len(tracker.conns), "should have 1 remaining connection")
|
||||
tracker.mu.Unlock()
|
||||
|
||||
// Verify host A's conn is actually closed.
|
||||
buf := make([]byte, 1)
|
||||
_, err = connA2.Read(buf)
|
||||
assert.Error(t, err, "host A pipe should be closed")
|
||||
|
||||
// Host B should still be alive.
|
||||
go func() { _, _ = connB1.Write([]byte("x")) }()
|
||||
|
||||
// Close all remaining.
|
||||
n = tracker.CloseAll()
|
||||
assert.Equal(t, 1, n, "should close remaining 1 connection")
|
||||
|
||||
tracker.mu.Lock()
|
||||
assert.Equal(t, 0, len(tracker.conns), "should have 0 connections after CloseAll")
|
||||
tracker.mu.Unlock()
|
||||
}
|
||||
|
||||
func TestCloseAll(t *testing.T) {
|
||||
var tracker HijackTracker
|
||||
|
||||
for range 5 {
|
||||
c1, c2 := net.Pipe()
|
||||
defer c2.Close()
|
||||
tc := &trackedConn{Conn: c1, tracker: &tracker, host: "test.com"}
|
||||
tracker.add(tc)
|
||||
}
|
||||
|
||||
tracker.mu.Lock()
|
||||
assert.Equal(t, 5, len(tracker.conns))
|
||||
tracker.mu.Unlock()
|
||||
|
||||
n := tracker.CloseAll()
|
||||
assert.Equal(t, 5, n)
|
||||
|
||||
// Double CloseAll is safe.
|
||||
n = tracker.CloseAll()
|
||||
assert.Equal(t, 0, n)
|
||||
}
|
||||
|
||||
func TestTrackedConn_AutoDeregister(t *testing.T) {
|
||||
var tracker HijackTracker
|
||||
|
||||
c1, c2 := net.Pipe()
|
||||
defer c2.Close()
|
||||
|
||||
tc := &trackedConn{Conn: c1, tracker: &tracker, host: "auto.com"}
|
||||
tracker.add(tc)
|
||||
|
||||
tracker.mu.Lock()
|
||||
assert.Equal(t, 1, len(tracker.conns))
|
||||
tracker.mu.Unlock()
|
||||
|
||||
// Close the tracked conn: should auto-deregister.
|
||||
require.NoError(t, tc.Close())
|
||||
|
||||
tracker.mu.Lock()
|
||||
assert.Equal(t, 0, len(tracker.conns), "should auto-deregister on close")
|
||||
tracker.mu.Unlock()
|
||||
}
|
||||
|
||||
func TestHostOnly(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{"example.com:443", "example.com"},
|
||||
{"example.com", "example.com"},
|
||||
{"127.0.0.1:8080", "127.0.0.1"},
|
||||
{"[::1]:443", "[::1]"},
|
||||
{"", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.input, func(t *testing.T) {
|
||||
assert.Equal(t, tt.want, hostOnly(tt.input))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -152,7 +152,7 @@ func (c *Client) printClients(data map[string]any) {
|
||||
return
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintf(c.out, "%-38s %-12s %-40s %s\n", "ACCOUNT ID", "AGE", "DOMAINS", "HAS CLIENT")
|
||||
_, _ = fmt.Fprintf(c.out, "%-38s %-12s %-40s %s\n", "ACCOUNT ID", "AGE", "SERVICES", "HAS CLIENT")
|
||||
_, _ = fmt.Fprintln(c.out, strings.Repeat("-", 110))
|
||||
|
||||
for _, item := range clients {
|
||||
@@ -166,7 +166,7 @@ func (c *Client) printClientRow(item any) {
|
||||
return
|
||||
}
|
||||
|
||||
domains := c.extractDomains(client)
|
||||
services := c.extractServiceKeys(client)
|
||||
hasClient := "no"
|
||||
if hc, ok := client["has_client"].(bool); ok && hc {
|
||||
hasClient = "yes"
|
||||
@@ -175,20 +175,20 @@ func (c *Client) printClientRow(item any) {
|
||||
_, _ = fmt.Fprintf(c.out, "%-38s %-12v %s %s\n",
|
||||
client["account_id"],
|
||||
client["age"],
|
||||
domains,
|
||||
services,
|
||||
hasClient,
|
||||
)
|
||||
}
|
||||
|
||||
func (c *Client) extractDomains(client map[string]any) string {
|
||||
d, ok := client["domains"].([]any)
|
||||
func (c *Client) extractServiceKeys(client map[string]any) string {
|
||||
d, ok := client["service_keys"].([]any)
|
||||
if !ok || len(d) == 0 {
|
||||
return "-"
|
||||
}
|
||||
|
||||
parts := make([]string, len(d))
|
||||
for i, domain := range d {
|
||||
parts[i] = fmt.Sprint(domain)
|
||||
for i, key := range d {
|
||||
parts[i] = fmt.Sprint(key)
|
||||
}
|
||||
return strings.Join(parts, ", ")
|
||||
}
|
||||
|
||||
@@ -189,7 +189,7 @@ type indexData struct {
|
||||
Version string
|
||||
Uptime string
|
||||
ClientCount int
|
||||
TotalDomains int
|
||||
TotalServices int
|
||||
CertsTotal int
|
||||
CertsReady int
|
||||
CertsPending int
|
||||
@@ -202,7 +202,7 @@ type indexData struct {
|
||||
|
||||
type clientData struct {
|
||||
AccountID string
|
||||
Domains string
|
||||
Services string
|
||||
Age string
|
||||
Status string
|
||||
}
|
||||
@@ -211,9 +211,9 @@ func (h *Handler) handleIndex(w http.ResponseWriter, _ *http.Request, wantJSON b
|
||||
clients := h.provider.ListClientsForDebug()
|
||||
sortedIDs := sortedAccountIDs(clients)
|
||||
|
||||
totalDomains := 0
|
||||
totalServices := 0
|
||||
for _, info := range clients {
|
||||
totalDomains += info.DomainCount
|
||||
totalServices += info.ServiceCount
|
||||
}
|
||||
|
||||
var certsTotal, certsReady, certsPending, certsFailed int
|
||||
@@ -234,24 +234,24 @@ func (h *Handler) handleIndex(w http.ResponseWriter, _ *http.Request, wantJSON b
|
||||
for _, id := range sortedIDs {
|
||||
info := clients[id]
|
||||
clientsJSON = append(clientsJSON, map[string]interface{}{
|
||||
"account_id": info.AccountID,
|
||||
"domain_count": info.DomainCount,
|
||||
"domains": info.Domains,
|
||||
"has_client": info.HasClient,
|
||||
"created_at": info.CreatedAt,
|
||||
"age": time.Since(info.CreatedAt).Round(time.Second).String(),
|
||||
"account_id": info.AccountID,
|
||||
"service_count": info.ServiceCount,
|
||||
"service_keys": info.ServiceKeys,
|
||||
"has_client": info.HasClient,
|
||||
"created_at": info.CreatedAt,
|
||||
"age": time.Since(info.CreatedAt).Round(time.Second).String(),
|
||||
})
|
||||
}
|
||||
resp := map[string]interface{}{
|
||||
"version": version.NetbirdVersion(),
|
||||
"uptime": time.Since(h.startTime).Round(time.Second).String(),
|
||||
"client_count": len(clients),
|
||||
"total_domains": totalDomains,
|
||||
"certs_total": certsTotal,
|
||||
"certs_ready": certsReady,
|
||||
"certs_pending": certsPending,
|
||||
"certs_failed": certsFailed,
|
||||
"clients": clientsJSON,
|
||||
"version": version.NetbirdVersion(),
|
||||
"uptime": time.Since(h.startTime).Round(time.Second).String(),
|
||||
"client_count": len(clients),
|
||||
"total_services": totalServices,
|
||||
"certs_total": certsTotal,
|
||||
"certs_ready": certsReady,
|
||||
"certs_pending": certsPending,
|
||||
"certs_failed": certsFailed,
|
||||
"clients": clientsJSON,
|
||||
}
|
||||
if len(certsPendingDomains) > 0 {
|
||||
resp["certs_pending_domains"] = certsPendingDomains
|
||||
@@ -278,7 +278,7 @@ func (h *Handler) handleIndex(w http.ResponseWriter, _ *http.Request, wantJSON b
|
||||
Version: version.NetbirdVersion(),
|
||||
Uptime: time.Since(h.startTime).Round(time.Second).String(),
|
||||
ClientCount: len(clients),
|
||||
TotalDomains: totalDomains,
|
||||
TotalServices: totalServices,
|
||||
CertsTotal: certsTotal,
|
||||
CertsReady: certsReady,
|
||||
CertsPending: certsPending,
|
||||
@@ -291,9 +291,9 @@ func (h *Handler) handleIndex(w http.ResponseWriter, _ *http.Request, wantJSON b
|
||||
|
||||
for _, id := range sortedIDs {
|
||||
info := clients[id]
|
||||
domains := info.Domains.SafeString()
|
||||
if domains == "" {
|
||||
domains = "-"
|
||||
services := strings.Join(info.ServiceKeys, ", ")
|
||||
if services == "" {
|
||||
services = "-"
|
||||
}
|
||||
status := "No client"
|
||||
if info.HasClient {
|
||||
@@ -301,7 +301,7 @@ func (h *Handler) handleIndex(w http.ResponseWriter, _ *http.Request, wantJSON b
|
||||
}
|
||||
data.Clients = append(data.Clients, clientData{
|
||||
AccountID: string(info.AccountID),
|
||||
Domains: domains,
|
||||
Services: services,
|
||||
Age: time.Since(info.CreatedAt).Round(time.Second).String(),
|
||||
Status: status,
|
||||
})
|
||||
@@ -324,12 +324,12 @@ func (h *Handler) handleListClients(w http.ResponseWriter, _ *http.Request, want
|
||||
for _, id := range sortedIDs {
|
||||
info := clients[id]
|
||||
clientsJSON = append(clientsJSON, map[string]interface{}{
|
||||
"account_id": info.AccountID,
|
||||
"domain_count": info.DomainCount,
|
||||
"domains": info.Domains,
|
||||
"has_client": info.HasClient,
|
||||
"created_at": info.CreatedAt,
|
||||
"age": time.Since(info.CreatedAt).Round(time.Second).String(),
|
||||
"account_id": info.AccountID,
|
||||
"service_count": info.ServiceCount,
|
||||
"service_keys": info.ServiceKeys,
|
||||
"has_client": info.HasClient,
|
||||
"created_at": info.CreatedAt,
|
||||
"age": time.Since(info.CreatedAt).Round(time.Second).String(),
|
||||
})
|
||||
}
|
||||
h.writeJSON(w, map[string]interface{}{
|
||||
@@ -347,9 +347,9 @@ func (h *Handler) handleListClients(w http.ResponseWriter, _ *http.Request, want
|
||||
|
||||
for _, id := range sortedIDs {
|
||||
info := clients[id]
|
||||
domains := info.Domains.SafeString()
|
||||
if domains == "" {
|
||||
domains = "-"
|
||||
services := strings.Join(info.ServiceKeys, ", ")
|
||||
if services == "" {
|
||||
services = "-"
|
||||
}
|
||||
status := "No client"
|
||||
if info.HasClient {
|
||||
@@ -357,7 +357,7 @@ func (h *Handler) handleListClients(w http.ResponseWriter, _ *http.Request, want
|
||||
}
|
||||
data.Clients = append(data.Clients, clientData{
|
||||
AccountID: string(info.AccountID),
|
||||
Domains: domains,
|
||||
Services: services,
|
||||
Age: time.Since(info.CreatedAt).Round(time.Second).String(),
|
||||
Status: status,
|
||||
})
|
||||
@@ -409,17 +409,13 @@ func (h *Handler) handleClientStatus(w http.ResponseWriter, r *http.Request, acc
|
||||
}
|
||||
|
||||
pbStatus := nbstatus.ToProtoFullStatus(fullStatus)
|
||||
overview := nbstatus.ConvertToStatusOutputOverview(
|
||||
pbStatus,
|
||||
false,
|
||||
version.NetbirdVersion(),
|
||||
statusFilter,
|
||||
prefixNamesFilter,
|
||||
prefixNamesFilterMap,
|
||||
ipsFilterMap,
|
||||
connectionTypeFilter,
|
||||
"",
|
||||
)
|
||||
overview := nbstatus.ConvertToStatusOutputOverview(pbStatus, nbstatus.ConvertOptions{
|
||||
StatusFilter: statusFilter,
|
||||
PrefixNamesFilter: prefixNamesFilter,
|
||||
PrefixNamesFilterMap: prefixNamesFilterMap,
|
||||
IPsFilter: ipsFilterMap,
|
||||
ConnectionTypeFilter: connectionTypeFilter,
|
||||
})
|
||||
|
||||
if wantJSON {
|
||||
h.writeJSON(w, map[string]interface{}{
|
||||
|
||||
@@ -12,14 +12,14 @@
|
||||
<table>
|
||||
<tr>
|
||||
<th>Account ID</th>
|
||||
<th>Domains</th>
|
||||
<th>Services</th>
|
||||
<th>Age</th>
|
||||
<th>Status</th>
|
||||
</tr>
|
||||
{{range .Clients}}
|
||||
<tr>
|
||||
<td><a href="/debug/clients/{{.AccountID}}/tools">{{.AccountID}}</a></td>
|
||||
<td>{{.Domains}}</td>
|
||||
<td>{{.Services}}</td>
|
||||
<td>{{.Age}}</td>
|
||||
<td>{{.Status}}</td>
|
||||
</tr>
|
||||
|
||||
@@ -27,19 +27,19 @@
|
||||
<ul>{{range .CertsFailedDomains}}<li>{{.Domain}}: {{.Error}}</li>{{end}}</ul>
|
||||
</details>
|
||||
{{end}}
|
||||
<h2>Clients ({{.ClientCount}}) | Domains ({{.TotalDomains}})</h2>
|
||||
<h2>Clients ({{.ClientCount}}) | Services ({{.TotalServices}})</h2>
|
||||
{{if .Clients}}
|
||||
<table>
|
||||
<tr>
|
||||
<th>Account ID</th>
|
||||
<th>Domains</th>
|
||||
<th>Services</th>
|
||||
<th>Age</th>
|
||||
<th>Status</th>
|
||||
</tr>
|
||||
{{range .Clients}}
|
||||
<tr>
|
||||
<td><a href="/debug/clients/{{.AccountID}}/tools">{{.AccountID}}</a></td>
|
||||
<td>{{.Domains}}</td>
|
||||
<td>{{.Services}}</td>
|
||||
<td>{{.Age}}</td>
|
||||
<td>{{.Status}}</td>
|
||||
</tr>
|
||||
|
||||
@@ -0,0 +1,264 @@
|
||||
package geolocation
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"bufio"
|
||||
"compress/gzip"
|
||||
"crypto/sha256"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
const (
|
||||
mmdbTarGZURL = "https://pkgs.netbird.io/geolocation-dbs/GeoLite2-City/download?suffix=tar.gz"
|
||||
mmdbSha256URL = "https://pkgs.netbird.io/geolocation-dbs/GeoLite2-City/download?suffix=tar.gz.sha256"
|
||||
mmdbInnerName = "GeoLite2-City.mmdb"
|
||||
|
||||
downloadTimeout = 2 * time.Minute
|
||||
maxMMDBSize = 256 << 20 // 256 MB
|
||||
)
|
||||
|
||||
// ensureMMDB checks for an existing MMDB file in dataDir. If none is found,
|
||||
// it downloads from pkgs.netbird.io with SHA256 verification.
|
||||
func ensureMMDB(logger *log.Logger, dataDir string) (string, error) {
|
||||
if err := os.MkdirAll(dataDir, 0o755); err != nil {
|
||||
return "", fmt.Errorf("create geo data directory %s: %w", dataDir, err)
|
||||
}
|
||||
|
||||
pattern := filepath.Join(dataDir, mmdbGlob)
|
||||
if files, _ := filepath.Glob(pattern); len(files) > 0 {
|
||||
mmdbPath := files[len(files)-1]
|
||||
logger.Debugf("using existing geolocation database: %s", mmdbPath)
|
||||
return mmdbPath, nil
|
||||
}
|
||||
|
||||
logger.Info("geolocation database not found, downloading from pkgs.netbird.io")
|
||||
return downloadMMDB(logger, dataDir)
|
||||
}
|
||||
|
||||
func downloadMMDB(logger *log.Logger, dataDir string) (string, error) {
|
||||
client := &http.Client{Timeout: downloadTimeout}
|
||||
|
||||
datedName, err := fetchRemoteFilename(client, mmdbTarGZURL)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("get remote filename: %w", err)
|
||||
}
|
||||
|
||||
mmdbFilename := deriveMMDBFilename(datedName)
|
||||
mmdbPath := filepath.Join(dataDir, mmdbFilename)
|
||||
|
||||
tmp, err := os.MkdirTemp("", "geolite-proxy-*")
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("create temp directory: %w", err)
|
||||
}
|
||||
defer os.RemoveAll(tmp)
|
||||
|
||||
checksumFile := filepath.Join(tmp, "checksum.sha256")
|
||||
if err := downloadToFile(client, mmdbSha256URL, checksumFile); err != nil {
|
||||
return "", fmt.Errorf("download checksum: %w", err)
|
||||
}
|
||||
|
||||
expectedHash, err := readChecksumFile(checksumFile)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("read checksum: %w", err)
|
||||
}
|
||||
|
||||
tarFile := filepath.Join(tmp, datedName)
|
||||
logger.Debugf("downloading geolocation database (%s)", datedName)
|
||||
if err := downloadToFile(client, mmdbTarGZURL, tarFile); err != nil {
|
||||
return "", fmt.Errorf("download database: %w", err)
|
||||
}
|
||||
|
||||
if err := verifySHA256(tarFile, expectedHash); err != nil {
|
||||
return "", fmt.Errorf("verify database checksum: %w", err)
|
||||
}
|
||||
|
||||
if err := extractMMDBFromTarGZ(tarFile, mmdbPath); err != nil {
|
||||
return "", fmt.Errorf("extract database: %w", err)
|
||||
}
|
||||
|
||||
logger.Infof("geolocation database downloaded: %s", mmdbPath)
|
||||
return mmdbPath, nil
|
||||
}
|
||||
|
||||
// deriveMMDBFilename converts a tar.gz filename to an MMDB filename.
|
||||
// Example: GeoLite2-City_20240101.tar.gz -> GeoLite2-City_20240101.mmdb
|
||||
func deriveMMDBFilename(tarName string) string {
|
||||
base, _, _ := strings.Cut(tarName, ".")
|
||||
if !strings.Contains(base, "_") {
|
||||
return "GeoLite2-City.mmdb"
|
||||
}
|
||||
return base + ".mmdb"
|
||||
}
|
||||
|
||||
func fetchRemoteFilename(client *http.Client, url string) (string, error) {
|
||||
resp, err := client.Head(url)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("HEAD request: HTTP %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
cd := resp.Header.Get("Content-Disposition")
|
||||
if cd == "" {
|
||||
return "", errors.New("no Content-Disposition header")
|
||||
}
|
||||
|
||||
_, params, err := mime.ParseMediaType(cd)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("parse Content-Disposition: %w", err)
|
||||
}
|
||||
|
||||
name := filepath.Base(params["filename"])
|
||||
if name == "" || name == "." {
|
||||
return "", errors.New("no filename in Content-Disposition")
|
||||
}
|
||||
return name, nil
|
||||
}
|
||||
|
||||
func downloadToFile(client *http.Client, url, dest string) error {
|
||||
resp, err := client.Get(url) //nolint:gosec
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 1024))
|
||||
return fmt.Errorf("HTTP %d: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
f, err := os.Create(dest) //nolint:gosec
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
// Cap download at 256 MB to prevent unbounded reads from a compromised server.
|
||||
if _, err := io.Copy(f, io.LimitReader(resp.Body, maxMMDBSize)); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func readChecksumFile(path string) (string, error) {
|
||||
f, err := os.Open(path) //nolint:gosec
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
scanner := bufio.NewScanner(f)
|
||||
if scanner.Scan() {
|
||||
parts := strings.Fields(scanner.Text())
|
||||
if len(parts) > 0 {
|
||||
return parts[0], nil
|
||||
}
|
||||
}
|
||||
if err := scanner.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "", errors.New("empty checksum file")
|
||||
}
|
||||
|
||||
func verifySHA256(path, expected string) error {
|
||||
f, err := os.Open(path) //nolint:gosec
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
h := sha256.New()
|
||||
if _, err := io.Copy(h, f); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
actual := fmt.Sprintf("%x", h.Sum(nil))
|
||||
if actual != expected {
|
||||
return fmt.Errorf("SHA256 mismatch: expected %s, got %s", expected, actual)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func extractMMDBFromTarGZ(tarGZPath, destPath string) error {
|
||||
f, err := os.Open(tarGZPath) //nolint:gosec
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
gz, err := gzip.NewReader(f)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer gz.Close()
|
||||
|
||||
tr := tar.NewReader(gz)
|
||||
for {
|
||||
hdr, err := tr.Next()
|
||||
if err != nil {
|
||||
if errors.Is(err, io.EOF) {
|
||||
break
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
if hdr.Typeflag == tar.TypeReg && filepath.Base(hdr.Name) == mmdbInnerName {
|
||||
if hdr.Size < 0 || hdr.Size > maxMMDBSize {
|
||||
return fmt.Errorf("mmdb entry size %d exceeds limit %d", hdr.Size, maxMMDBSize)
|
||||
}
|
||||
if err := extractToFileAtomic(io.LimitReader(tr, hdr.Size), destPath); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
return fmt.Errorf("%s not found in archive", mmdbInnerName)
|
||||
}
|
||||
|
||||
// extractToFileAtomic writes r to a temporary file in the same directory as
|
||||
// destPath, then renames it into place so a crash never leaves a truncated file.
|
||||
func extractToFileAtomic(r io.Reader, destPath string) error {
|
||||
dir := filepath.Dir(destPath)
|
||||
tmp, err := os.CreateTemp(dir, ".mmdb-*.tmp")
|
||||
if err != nil {
|
||||
return fmt.Errorf("create temp file: %w", err)
|
||||
}
|
||||
tmpPath := tmp.Name()
|
||||
|
||||
if _, err := io.Copy(tmp, r); err != nil { //nolint:gosec // G110: caller bounds with LimitReader
|
||||
if closeErr := tmp.Close(); closeErr != nil {
|
||||
log.Debugf("failed to close temp file %s: %v", tmpPath, closeErr)
|
||||
}
|
||||
if removeErr := os.Remove(tmpPath); removeErr != nil {
|
||||
log.Debugf("failed to remove temp file %s: %v", tmpPath, removeErr)
|
||||
}
|
||||
return fmt.Errorf("write mmdb: %w", err)
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
if removeErr := os.Remove(tmpPath); removeErr != nil {
|
||||
log.Debugf("failed to remove temp file %s: %v", tmpPath, removeErr)
|
||||
}
|
||||
return fmt.Errorf("close temp file: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmpPath, destPath); err != nil {
|
||||
if removeErr := os.Remove(tmpPath); removeErr != nil {
|
||||
log.Debugf("failed to remove temp file %s: %v", tmpPath, removeErr)
|
||||
}
|
||||
return fmt.Errorf("rename to %s: %w", destPath, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,152 @@
|
||||
// Package geolocation provides IP-to-country lookups using MaxMind GeoLite2 databases.
|
||||
package geolocation
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"os"
|
||||
"strconv"
|
||||
"sync"
|
||||
|
||||
"github.com/oschwald/maxminddb-golang"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
const (
|
||||
// EnvDisable disables geolocation lookups entirely when set to a truthy value.
|
||||
EnvDisable = "NB_PROXY_DISABLE_GEOLOCATION"
|
||||
|
||||
mmdbGlob = "GeoLite2-City_*.mmdb"
|
||||
)
|
||||
|
||||
type record struct {
|
||||
Country struct {
|
||||
ISOCode string `maxminddb:"iso_code"`
|
||||
} `maxminddb:"country"`
|
||||
City struct {
|
||||
Names struct {
|
||||
En string `maxminddb:"en"`
|
||||
} `maxminddb:"names"`
|
||||
} `maxminddb:"city"`
|
||||
Subdivisions []struct {
|
||||
ISOCode string `maxminddb:"iso_code"`
|
||||
Names struct {
|
||||
En string `maxminddb:"en"`
|
||||
} `maxminddb:"names"`
|
||||
} `maxminddb:"subdivisions"`
|
||||
}
|
||||
|
||||
// Result holds the outcome of a geo lookup.
|
||||
type Result struct {
|
||||
CountryCode string
|
||||
CityName string
|
||||
SubdivisionCode string
|
||||
SubdivisionName string
|
||||
}
|
||||
|
||||
// Lookup provides IP geolocation lookups.
|
||||
type Lookup struct {
|
||||
mu sync.RWMutex
|
||||
db *maxminddb.Reader
|
||||
logger *log.Logger
|
||||
}
|
||||
|
||||
// NewLookup opens or downloads the GeoLite2-City MMDB in dataDir.
|
||||
// Returns nil without error if geolocation is disabled via environment
|
||||
// variable, no data directory is configured, or the download fails
|
||||
// (graceful degradation: country restrictions will deny all requests).
|
||||
func NewLookup(logger *log.Logger, dataDir string) (*Lookup, error) {
|
||||
if isDisabledByEnv(logger) {
|
||||
logger.Info("geolocation disabled via environment variable")
|
||||
return nil, nil //nolint:nilnil
|
||||
}
|
||||
|
||||
if dataDir == "" {
|
||||
return nil, nil //nolint:nilnil
|
||||
}
|
||||
|
||||
mmdbPath, err := ensureMMDB(logger, dataDir)
|
||||
if err != nil {
|
||||
logger.Warnf("geolocation database unavailable: %v", err)
|
||||
logger.Warn("country-based access restrictions will deny all requests until a database is available")
|
||||
return nil, nil //nolint:nilnil
|
||||
}
|
||||
|
||||
db, err := maxminddb.Open(mmdbPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open GeoLite2 database %s: %w", mmdbPath, err)
|
||||
}
|
||||
|
||||
logger.Infof("geolocation database loaded from %s", mmdbPath)
|
||||
return &Lookup{db: db, logger: logger}, nil
|
||||
}
|
||||
|
||||
// LookupAddr returns the country ISO code and city name for the given IP.
|
||||
// Returns an empty Result if the database is nil or the lookup fails.
|
||||
func (l *Lookup) LookupAddr(addr netip.Addr) Result {
|
||||
if l == nil {
|
||||
return Result{}
|
||||
}
|
||||
|
||||
l.mu.RLock()
|
||||
defer l.mu.RUnlock()
|
||||
|
||||
if l.db == nil {
|
||||
return Result{}
|
||||
}
|
||||
|
||||
addr = addr.Unmap()
|
||||
|
||||
var rec record
|
||||
if err := l.db.Lookup(addr.AsSlice(), &rec); err != nil {
|
||||
l.logger.Debugf("geolocation lookup %s: %v", addr, err)
|
||||
return Result{}
|
||||
}
|
||||
r := Result{
|
||||
CountryCode: rec.Country.ISOCode,
|
||||
CityName: rec.City.Names.En,
|
||||
}
|
||||
if len(rec.Subdivisions) > 0 {
|
||||
r.SubdivisionCode = rec.Subdivisions[0].ISOCode
|
||||
r.SubdivisionName = rec.Subdivisions[0].Names.En
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// Available reports whether the lookup has a loaded database.
|
||||
func (l *Lookup) Available() bool {
|
||||
if l == nil {
|
||||
return false
|
||||
}
|
||||
l.mu.RLock()
|
||||
defer l.mu.RUnlock()
|
||||
return l.db != nil
|
||||
}
|
||||
|
||||
// Close releases the database resources.
|
||||
func (l *Lookup) Close() error {
|
||||
if l == nil {
|
||||
return nil
|
||||
}
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
if l.db != nil {
|
||||
err := l.db.Close()
|
||||
l.db = nil
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func isDisabledByEnv(logger *log.Logger) bool {
|
||||
val := os.Getenv(EnvDisable)
|
||||
if val == "" {
|
||||
return false
|
||||
}
|
||||
disabled, err := strconv.ParseBool(val)
|
||||
if err != nil {
|
||||
logger.Warnf("parse %s=%q: %v", EnvDisable, val, err)
|
||||
return false
|
||||
}
|
||||
return disabled
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package metrics_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"reflect"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
promexporter "go.opentelemetry.io/otel/exporters/prometheus"
|
||||
sdkmetric "go.opentelemetry.io/otel/sdk/metric"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/metrics"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
)
|
||||
|
||||
func newTestMetrics(t *testing.T) *metrics.Metrics {
|
||||
t.Helper()
|
||||
|
||||
exporter, err := promexporter.New()
|
||||
if err != nil {
|
||||
t.Fatalf("create prometheus exporter: %v", err)
|
||||
}
|
||||
|
||||
provider := sdkmetric.NewMeterProvider(sdkmetric.WithReader(exporter))
|
||||
pkg := reflect.TypeOf(metrics.Metrics{}).PkgPath()
|
||||
meter := provider.Meter(pkg)
|
||||
|
||||
m, err := metrics.New(context.Background(), meter)
|
||||
if err != nil {
|
||||
t.Fatalf("create metrics: %v", err)
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
func TestL4ServiceGauge(t *testing.T) {
|
||||
m := newTestMetrics(t)
|
||||
|
||||
m.L4ServiceAdded(types.ServiceModeTCP)
|
||||
m.L4ServiceAdded(types.ServiceModeTCP)
|
||||
m.L4ServiceAdded(types.ServiceModeUDP)
|
||||
m.L4ServiceRemoved(types.ServiceModeTCP)
|
||||
}
|
||||
|
||||
func TestTCPRelayMetrics(t *testing.T) {
|
||||
m := newTestMetrics(t)
|
||||
|
||||
acct := types.AccountID("acct-1")
|
||||
|
||||
m.TCPRelayStarted(acct)
|
||||
m.TCPRelayStarted(acct)
|
||||
m.TCPRelayEnded(acct, 10*time.Second, 1000, 500)
|
||||
m.TCPRelayDialError(acct)
|
||||
m.TCPRelayRejected(acct)
|
||||
}
|
||||
|
||||
func TestUDPSessionMetrics(t *testing.T) {
|
||||
m := newTestMetrics(t)
|
||||
|
||||
acct := types.AccountID("acct-2")
|
||||
|
||||
m.UDPSessionStarted(acct)
|
||||
m.UDPSessionStarted(acct)
|
||||
m.UDPSessionEnded(acct)
|
||||
m.UDPSessionDialError(acct)
|
||||
m.UDPSessionRejected(acct)
|
||||
m.UDPPacketRelayed(types.RelayDirectionClientToBackend, 100)
|
||||
m.UDPPacketRelayed(types.RelayDirectionClientToBackend, 200)
|
||||
m.UDPPacketRelayed(types.RelayDirectionBackendToClient, 150)
|
||||
}
|
||||
@@ -6,12 +6,15 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"go.opentelemetry.io/otel/attribute"
|
||||
"go.opentelemetry.io/otel/metric"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/proxy"
|
||||
"github.com/netbirdio/netbird/proxy/internal/responsewriter"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
)
|
||||
|
||||
// Metrics collects OpenTelemetry metrics for the proxy.
|
||||
type Metrics struct {
|
||||
ctx context.Context
|
||||
requestsTotal metric.Int64Counter
|
||||
@@ -22,85 +25,188 @@ type Metrics struct {
|
||||
backendDuration metric.Int64Histogram
|
||||
certificateIssueDuration metric.Int64Histogram
|
||||
|
||||
// L4 service-level metrics.
|
||||
l4Services metric.Int64UpDownCounter
|
||||
|
||||
// L4 TCP connection-level metrics.
|
||||
tcpActiveConns metric.Int64UpDownCounter
|
||||
tcpConnsTotal metric.Int64Counter
|
||||
tcpConnDuration metric.Int64Histogram
|
||||
tcpBytesTotal metric.Int64Counter
|
||||
|
||||
// L4 UDP session-level metrics.
|
||||
udpActiveSess metric.Int64UpDownCounter
|
||||
udpSessionsTotal metric.Int64Counter
|
||||
udpPacketsTotal metric.Int64Counter
|
||||
udpBytesTotal metric.Int64Counter
|
||||
|
||||
mappingsMux sync.Mutex
|
||||
mappingPaths map[string]int
|
||||
}
|
||||
|
||||
// New creates a Metrics instance using the given OpenTelemetry meter.
|
||||
func New(ctx context.Context, meter metric.Meter) (*Metrics, error) {
|
||||
requestsTotal, err := meter.Int64Counter(
|
||||
m := &Metrics{
|
||||
ctx: ctx,
|
||||
mappingPaths: make(map[string]int),
|
||||
}
|
||||
|
||||
if err := m.initHTTPMetrics(meter); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := m.initL4Metrics(meter); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return m, nil
|
||||
}
|
||||
|
||||
func (m *Metrics) initHTTPMetrics(meter metric.Meter) error {
|
||||
var err error
|
||||
|
||||
m.requestsTotal, err = meter.Int64Counter(
|
||||
"proxy.http.request.counter",
|
||||
metric.WithUnit("1"),
|
||||
metric.WithDescription("Total number of requests made to the netbird proxy"),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
activeRequests, err := meter.Int64UpDownCounter(
|
||||
m.activeRequests, err = meter.Int64UpDownCounter(
|
||||
"proxy.http.active_requests",
|
||||
metric.WithUnit("1"),
|
||||
metric.WithDescription("Current in-flight requests handled by the netbird proxy"),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
configuredDomains, err := meter.Int64UpDownCounter(
|
||||
m.configuredDomains, err = meter.Int64UpDownCounter(
|
||||
"proxy.domains.count",
|
||||
metric.WithUnit("1"),
|
||||
metric.WithDescription("Current number of domains configured on the netbird proxy"),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
totalPaths, err := meter.Int64UpDownCounter(
|
||||
m.totalPaths, err = meter.Int64UpDownCounter(
|
||||
"proxy.paths.count",
|
||||
metric.WithUnit("1"),
|
||||
metric.WithDescription("Total number of paths configured on the netbird proxy"),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
requestDuration, err := meter.Int64Histogram(
|
||||
m.requestDuration, err = meter.Int64Histogram(
|
||||
"proxy.http.request.duration.ms",
|
||||
metric.WithUnit("milliseconds"),
|
||||
metric.WithDescription("Duration of requests made to the netbird proxy"),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
backendDuration, err := meter.Int64Histogram(
|
||||
m.backendDuration, err = meter.Int64Histogram(
|
||||
"proxy.backend.duration.ms",
|
||||
metric.WithUnit("milliseconds"),
|
||||
metric.WithDescription("Duration of peer round trip time from the netbird proxy"),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
certificateIssueDuration, err := meter.Int64Histogram(
|
||||
m.certificateIssueDuration, err = meter.Int64Histogram(
|
||||
"proxy.certificate.issue.duration.ms",
|
||||
metric.WithUnit("milliseconds"),
|
||||
metric.WithDescription("Duration of ACME certificate issuance"),
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (m *Metrics) initL4Metrics(meter metric.Meter) error {
|
||||
var err error
|
||||
|
||||
m.l4Services, err = meter.Int64UpDownCounter(
|
||||
"proxy.l4.services.count",
|
||||
metric.WithUnit("1"),
|
||||
metric.WithDescription("Current number of configured L4 services (TCP/TLS/UDP) by mode"),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return err
|
||||
}
|
||||
|
||||
return &Metrics{
|
||||
ctx: ctx,
|
||||
requestsTotal: requestsTotal,
|
||||
activeRequests: activeRequests,
|
||||
configuredDomains: configuredDomains,
|
||||
totalPaths: totalPaths,
|
||||
requestDuration: requestDuration,
|
||||
backendDuration: backendDuration,
|
||||
certificateIssueDuration: certificateIssueDuration,
|
||||
mappingPaths: make(map[string]int),
|
||||
}, nil
|
||||
m.tcpActiveConns, err = meter.Int64UpDownCounter(
|
||||
"proxy.tcp.active_connections",
|
||||
metric.WithUnit("1"),
|
||||
metric.WithDescription("Current number of active TCP/TLS relay connections"),
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
m.tcpConnsTotal, err = meter.Int64Counter(
|
||||
"proxy.tcp.connections.total",
|
||||
metric.WithUnit("1"),
|
||||
metric.WithDescription("Total TCP/TLS relay connections by result and account"),
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
m.tcpConnDuration, err = meter.Int64Histogram(
|
||||
"proxy.tcp.connection.duration.ms",
|
||||
metric.WithUnit("milliseconds"),
|
||||
metric.WithDescription("Duration of TCP/TLS relay connections"),
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
m.tcpBytesTotal, err = meter.Int64Counter(
|
||||
"proxy.tcp.bytes.total",
|
||||
metric.WithUnit("bytes"),
|
||||
metric.WithDescription("Total bytes transferred through TCP/TLS relay by direction"),
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
m.udpActiveSess, err = meter.Int64UpDownCounter(
|
||||
"proxy.udp.active_sessions",
|
||||
metric.WithUnit("1"),
|
||||
metric.WithDescription("Current number of active UDP relay sessions"),
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
m.udpSessionsTotal, err = meter.Int64Counter(
|
||||
"proxy.udp.sessions.total",
|
||||
metric.WithUnit("1"),
|
||||
metric.WithDescription("Total UDP relay sessions by result and account"),
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
m.udpPacketsTotal, err = meter.Int64Counter(
|
||||
"proxy.udp.packets.total",
|
||||
metric.WithUnit("1"),
|
||||
metric.WithDescription("Total UDP packets relayed by direction"),
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
m.udpBytesTotal, err = meter.Int64Counter(
|
||||
"proxy.udp.bytes.total",
|
||||
metric.WithUnit("bytes"),
|
||||
metric.WithDescription("Total bytes transferred through UDP relay by direction"),
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
type responseInterceptor struct {
|
||||
@@ -120,6 +226,13 @@ func (w *responseInterceptor) Write(b []byte) (int, error) {
|
||||
return size, err
|
||||
}
|
||||
|
||||
// Unwrap returns the underlying ResponseWriter so http.ResponseController
|
||||
// can reach through to the original writer for Hijack/Flush operations.
|
||||
func (w *responseInterceptor) Unwrap() http.ResponseWriter {
|
||||
return w.PassthroughWriter
|
||||
}
|
||||
|
||||
// Middleware wraps an HTTP handler with request metrics.
|
||||
func (m *Metrics) Middleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
m.requestsTotal.Add(m.ctx, 1)
|
||||
@@ -144,6 +257,7 @@ func (f roundTripperFunc) RoundTrip(r *http.Request) (*http.Response, error) {
|
||||
return f(r)
|
||||
}
|
||||
|
||||
// RoundTripper wraps an http.RoundTripper with backend duration metrics.
|
||||
func (m *Metrics) RoundTripper(next http.RoundTripper) http.RoundTripper {
|
||||
return roundTripperFunc(func(req *http.Request) (*http.Response, error) {
|
||||
start := time.Now()
|
||||
@@ -156,6 +270,7 @@ func (m *Metrics) RoundTripper(next http.RoundTripper) http.RoundTripper {
|
||||
})
|
||||
}
|
||||
|
||||
// AddMapping records that a domain mapping was added.
|
||||
func (m *Metrics) AddMapping(mapping proxy.Mapping) {
|
||||
m.mappingsMux.Lock()
|
||||
defer m.mappingsMux.Unlock()
|
||||
@@ -175,13 +290,13 @@ func (m *Metrics) AddMapping(mapping proxy.Mapping) {
|
||||
m.mappingPaths[mapping.Host] = newPathCount
|
||||
}
|
||||
|
||||
// RemoveMapping records that a domain mapping was removed.
|
||||
func (m *Metrics) RemoveMapping(mapping proxy.Mapping) {
|
||||
m.mappingsMux.Lock()
|
||||
defer m.mappingsMux.Unlock()
|
||||
|
||||
oldPathCount, exists := m.mappingPaths[mapping.Host]
|
||||
if !exists {
|
||||
// Nothing to remove
|
||||
return
|
||||
}
|
||||
|
||||
@@ -195,3 +310,80 @@ func (m *Metrics) RemoveMapping(mapping proxy.Mapping) {
|
||||
func (m *Metrics) RecordCertificateIssuance(duration time.Duration) {
|
||||
m.certificateIssueDuration.Record(m.ctx, duration.Milliseconds())
|
||||
}
|
||||
|
||||
// L4ServiceAdded increments the L4 service gauge for the given mode.
|
||||
func (m *Metrics) L4ServiceAdded(mode types.ServiceMode) {
|
||||
m.l4Services.Add(m.ctx, 1, metric.WithAttributes(attribute.String("mode", string(mode))))
|
||||
}
|
||||
|
||||
// L4ServiceRemoved decrements the L4 service gauge for the given mode.
|
||||
func (m *Metrics) L4ServiceRemoved(mode types.ServiceMode) {
|
||||
m.l4Services.Add(m.ctx, -1, metric.WithAttributes(attribute.String("mode", string(mode))))
|
||||
}
|
||||
|
||||
// TCPRelayStarted records a new TCP relay connection starting.
|
||||
func (m *Metrics) TCPRelayStarted(accountID types.AccountID) {
|
||||
acct := attribute.String("account_id", string(accountID))
|
||||
m.tcpActiveConns.Add(m.ctx, 1, metric.WithAttributes(acct))
|
||||
m.tcpConnsTotal.Add(m.ctx, 1, metric.WithAttributes(acct, attribute.String("result", "success")))
|
||||
}
|
||||
|
||||
// TCPRelayEnded records a TCP relay connection ending and accumulates bytes and duration.
|
||||
func (m *Metrics) TCPRelayEnded(accountID types.AccountID, duration time.Duration, srcToDst, dstToSrc int64) {
|
||||
acct := attribute.String("account_id", string(accountID))
|
||||
m.tcpActiveConns.Add(m.ctx, -1, metric.WithAttributes(acct))
|
||||
m.tcpConnDuration.Record(m.ctx, duration.Milliseconds(), metric.WithAttributes(acct))
|
||||
m.tcpBytesTotal.Add(m.ctx, srcToDst, metric.WithAttributes(attribute.String("direction", "client_to_backend")))
|
||||
m.tcpBytesTotal.Add(m.ctx, dstToSrc, metric.WithAttributes(attribute.String("direction", "backend_to_client")))
|
||||
}
|
||||
|
||||
// TCPRelayDialError records a dial failure for a TCP relay.
|
||||
func (m *Metrics) TCPRelayDialError(accountID types.AccountID) {
|
||||
m.tcpConnsTotal.Add(m.ctx, 1, metric.WithAttributes(
|
||||
attribute.String("account_id", string(accountID)),
|
||||
attribute.String("result", "dial_error"),
|
||||
))
|
||||
}
|
||||
|
||||
// TCPRelayRejected records a rejected TCP relay (semaphore full).
|
||||
func (m *Metrics) TCPRelayRejected(accountID types.AccountID) {
|
||||
m.tcpConnsTotal.Add(m.ctx, 1, metric.WithAttributes(
|
||||
attribute.String("account_id", string(accountID)),
|
||||
attribute.String("result", "rejected"),
|
||||
))
|
||||
}
|
||||
|
||||
// UDPSessionStarted records a new UDP session starting.
|
||||
func (m *Metrics) UDPSessionStarted(accountID types.AccountID) {
|
||||
acct := attribute.String("account_id", string(accountID))
|
||||
m.udpActiveSess.Add(m.ctx, 1, metric.WithAttributes(acct))
|
||||
m.udpSessionsTotal.Add(m.ctx, 1, metric.WithAttributes(acct, attribute.String("result", "success")))
|
||||
}
|
||||
|
||||
// UDPSessionEnded records a UDP session ending.
|
||||
func (m *Metrics) UDPSessionEnded(accountID types.AccountID) {
|
||||
m.udpActiveSess.Add(m.ctx, -1, metric.WithAttributes(attribute.String("account_id", string(accountID))))
|
||||
}
|
||||
|
||||
// UDPSessionDialError records a dial failure for a UDP session.
|
||||
func (m *Metrics) UDPSessionDialError(accountID types.AccountID) {
|
||||
m.udpSessionsTotal.Add(m.ctx, 1, metric.WithAttributes(
|
||||
attribute.String("account_id", string(accountID)),
|
||||
attribute.String("result", "dial_error"),
|
||||
))
|
||||
}
|
||||
|
||||
// UDPSessionRejected records a rejected UDP session (limit or rate limited).
|
||||
func (m *Metrics) UDPSessionRejected(accountID types.AccountID) {
|
||||
m.udpSessionsTotal.Add(m.ctx, 1, metric.WithAttributes(
|
||||
attribute.String("account_id", string(accountID)),
|
||||
attribute.String("result", "rejected"),
|
||||
))
|
||||
}
|
||||
|
||||
// UDPPacketRelayed records a packet relayed in the given direction with its size in bytes.
|
||||
func (m *Metrics) UDPPacketRelayed(direction types.RelayDirection, bytes int) {
|
||||
dir := attribute.String("direction", string(direction))
|
||||
m.udpPacketsTotal.Add(m.ctx, 1, metric.WithAttributes(dir))
|
||||
m.udpBytesTotal.Add(m.ctx, int64(bytes), metric.WithAttributes(dir))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
package netutil
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math"
|
||||
"net"
|
||||
"syscall"
|
||||
)
|
||||
|
||||
// ValidatePort converts an int32 proto port to uint16, returning an error
|
||||
// if the value is out of the valid 1–65535 range.
|
||||
func ValidatePort(port int32) (uint16, error) {
|
||||
if port <= 0 || port > math.MaxUint16 {
|
||||
return 0, fmt.Errorf("invalid port %d: must be 1–65535", port)
|
||||
}
|
||||
return uint16(port), nil
|
||||
}
|
||||
|
||||
// IsExpectedError returns true for errors that are normal during
|
||||
// connection teardown and should not be logged as warnings.
|
||||
func IsExpectedError(err error) bool {
|
||||
return errors.Is(err, net.ErrClosed) ||
|
||||
errors.Is(err, context.Canceled) ||
|
||||
errors.Is(err, io.EOF) ||
|
||||
errors.Is(err, syscall.ECONNRESET) ||
|
||||
errors.Is(err, syscall.EPIPE) ||
|
||||
errors.Is(err, syscall.ECONNABORTED)
|
||||
}
|
||||
|
||||
// IsTimeout checks whether the error is a network timeout.
|
||||
func IsTimeout(err error) bool {
|
||||
var netErr net.Error
|
||||
if errors.As(err, &netErr) {
|
||||
return netErr.Timeout()
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
package netutil
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"syscall"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestValidatePort(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
port int32
|
||||
want uint16
|
||||
wantErr bool
|
||||
}{
|
||||
{"valid min", 1, 1, false},
|
||||
{"valid mid", 8080, 8080, false},
|
||||
{"valid max", 65535, 65535, false},
|
||||
{"zero", 0, 0, true},
|
||||
{"negative", -1, 0, true},
|
||||
{"too large", 65536, 0, true},
|
||||
{"way too large", 100000, 0, true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := ValidatePort(tt.port)
|
||||
if tt.wantErr {
|
||||
assert.Error(t, err)
|
||||
assert.Zero(t, got)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, tt.want, got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsExpectedError(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
want bool
|
||||
}{
|
||||
{"net.ErrClosed", net.ErrClosed, true},
|
||||
{"context.Canceled", context.Canceled, true},
|
||||
{"io.EOF", io.EOF, true},
|
||||
{"ECONNRESET", syscall.ECONNRESET, true},
|
||||
{"EPIPE", syscall.EPIPE, true},
|
||||
{"ECONNABORTED", syscall.ECONNABORTED, true},
|
||||
{"wrapped expected", fmt.Errorf("wrap: %w", net.ErrClosed), true},
|
||||
{"unexpected EOF", io.ErrUnexpectedEOF, false},
|
||||
{"generic error", errors.New("something"), false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
assert.Equal(t, tt.want, IsExpectedError(tt.err))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type timeoutErr struct{ timeout bool }
|
||||
|
||||
func (e *timeoutErr) Error() string { return "timeout" }
|
||||
func (e *timeoutErr) Timeout() bool { return e.timeout }
|
||||
func (e *timeoutErr) Temporary() bool { return false }
|
||||
|
||||
func TestIsTimeout(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
want bool
|
||||
}{
|
||||
{"net timeout", &timeoutErr{timeout: true}, true},
|
||||
{"net non-timeout", &timeoutErr{timeout: false}, false},
|
||||
{"wrapped timeout", fmt.Errorf("wrap: %w", &timeoutErr{timeout: true}), true},
|
||||
{"generic error", errors.New("not a timeout"), false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
assert.Equal(t, tt.want, IsTimeout(tt.err))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -2,6 +2,7 @@ package proxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"sync"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
@@ -10,8 +11,6 @@ import (
|
||||
type requestContextKey string
|
||||
|
||||
const (
|
||||
serviceIdKey requestContextKey = "serviceId"
|
||||
accountIdKey requestContextKey = "accountId"
|
||||
capturedDataKey requestContextKey = "capturedData"
|
||||
)
|
||||
|
||||
@@ -46,112 +45,117 @@ func (o ResponseOrigin) String() string {
|
||||
// to pass data back up the middleware chain.
|
||||
type CapturedData struct {
|
||||
mu sync.RWMutex
|
||||
RequestID string
|
||||
ServiceId string
|
||||
AccountId types.AccountID
|
||||
Origin ResponseOrigin
|
||||
ClientIP string
|
||||
UserID string
|
||||
AuthMethod string
|
||||
requestID string
|
||||
serviceID types.ServiceID
|
||||
accountID types.AccountID
|
||||
origin ResponseOrigin
|
||||
clientIP netip.Addr
|
||||
userID string
|
||||
authMethod string
|
||||
}
|
||||
|
||||
// GetRequestID safely gets the request ID
|
||||
// NewCapturedData creates a CapturedData with the given request ID.
|
||||
func NewCapturedData(requestID string) *CapturedData {
|
||||
return &CapturedData{requestID: requestID}
|
||||
}
|
||||
|
||||
// GetRequestID returns the request ID.
|
||||
func (c *CapturedData) GetRequestID() string {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.RequestID
|
||||
return c.requestID
|
||||
}
|
||||
|
||||
// SetServiceId safely sets the service ID
|
||||
func (c *CapturedData) SetServiceId(serviceId string) {
|
||||
// SetServiceID sets the service ID.
|
||||
func (c *CapturedData) SetServiceID(serviceID types.ServiceID) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.ServiceId = serviceId
|
||||
c.serviceID = serviceID
|
||||
}
|
||||
|
||||
// GetServiceId safely gets the service ID
|
||||
func (c *CapturedData) GetServiceId() string {
|
||||
// GetServiceID returns the service ID.
|
||||
func (c *CapturedData) GetServiceID() types.ServiceID {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.ServiceId
|
||||
return c.serviceID
|
||||
}
|
||||
|
||||
// SetAccountId safely sets the account ID
|
||||
func (c *CapturedData) SetAccountId(accountId types.AccountID) {
|
||||
// SetAccountID sets the account ID.
|
||||
func (c *CapturedData) SetAccountID(accountID types.AccountID) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.AccountId = accountId
|
||||
c.accountID = accountID
|
||||
}
|
||||
|
||||
// GetAccountId safely gets the account ID
|
||||
func (c *CapturedData) GetAccountId() types.AccountID {
|
||||
// GetAccountID returns the account ID.
|
||||
func (c *CapturedData) GetAccountID() types.AccountID {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.AccountId
|
||||
return c.accountID
|
||||
}
|
||||
|
||||
// SetOrigin safely sets the response origin
|
||||
// SetOrigin sets the response origin.
|
||||
func (c *CapturedData) SetOrigin(origin ResponseOrigin) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.Origin = origin
|
||||
c.origin = origin
|
||||
}
|
||||
|
||||
// GetOrigin safely gets the response origin
|
||||
// GetOrigin returns the response origin.
|
||||
func (c *CapturedData) GetOrigin() ResponseOrigin {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.Origin
|
||||
return c.origin
|
||||
}
|
||||
|
||||
// SetClientIP safely sets the resolved client IP.
|
||||
func (c *CapturedData) SetClientIP(ip string) {
|
||||
// SetClientIP sets the resolved client IP.
|
||||
func (c *CapturedData) SetClientIP(ip netip.Addr) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.ClientIP = ip
|
||||
c.clientIP = ip
|
||||
}
|
||||
|
||||
// GetClientIP safely gets the resolved client IP.
|
||||
func (c *CapturedData) GetClientIP() string {
|
||||
// GetClientIP returns the resolved client IP.
|
||||
func (c *CapturedData) GetClientIP() netip.Addr {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.ClientIP
|
||||
return c.clientIP
|
||||
}
|
||||
|
||||
// SetUserID safely sets the authenticated user ID.
|
||||
// SetUserID sets the authenticated user ID.
|
||||
func (c *CapturedData) SetUserID(userID string) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.UserID = userID
|
||||
c.userID = userID
|
||||
}
|
||||
|
||||
// GetUserID safely gets the authenticated user ID.
|
||||
// GetUserID returns the authenticated user ID.
|
||||
func (c *CapturedData) GetUserID() string {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.UserID
|
||||
return c.userID
|
||||
}
|
||||
|
||||
// SetAuthMethod safely sets the authentication method used.
|
||||
// SetAuthMethod sets the authentication method used.
|
||||
func (c *CapturedData) SetAuthMethod(method string) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.AuthMethod = method
|
||||
c.authMethod = method
|
||||
}
|
||||
|
||||
// GetAuthMethod safely gets the authentication method used.
|
||||
// GetAuthMethod returns the authentication method used.
|
||||
func (c *CapturedData) GetAuthMethod() string {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.AuthMethod
|
||||
return c.authMethod
|
||||
}
|
||||
|
||||
// WithCapturedData adds a CapturedData struct to the context
|
||||
// WithCapturedData adds a CapturedData struct to the context.
|
||||
func WithCapturedData(ctx context.Context, data *CapturedData) context.Context {
|
||||
return context.WithValue(ctx, capturedDataKey, data)
|
||||
}
|
||||
|
||||
// CapturedDataFromContext retrieves the CapturedData from context
|
||||
// CapturedDataFromContext retrieves the CapturedData from context.
|
||||
func CapturedDataFromContext(ctx context.Context) *CapturedData {
|
||||
v := ctx.Value(capturedDataKey)
|
||||
data, ok := v.(*CapturedData)
|
||||
@@ -160,28 +164,3 @@ func CapturedDataFromContext(ctx context.Context) *CapturedData {
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
func withServiceId(ctx context.Context, serviceId string) context.Context {
|
||||
return context.WithValue(ctx, serviceIdKey, serviceId)
|
||||
}
|
||||
|
||||
func ServiceIdFromContext(ctx context.Context) string {
|
||||
v := ctx.Value(serviceIdKey)
|
||||
serviceId, ok := v.(string)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
return serviceId
|
||||
}
|
||||
func withAccountId(ctx context.Context, accountId types.AccountID) context.Context {
|
||||
return context.WithValue(ctx, accountIdKey, accountId)
|
||||
}
|
||||
|
||||
func AccountIdFromContext(ctx context.Context) types.AccountID {
|
||||
v := ctx.Value(accountIdKey)
|
||||
accountId, ok := v.(types.AccountID)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
return accountId
|
||||
}
|
||||
|
||||
@@ -25,7 +25,7 @@ func (nopTransport) RoundTrip(*http.Request) (*http.Response, error) {
|
||||
func BenchmarkServeHTTP(b *testing.B) {
|
||||
rp := proxy.NewReverseProxy(nopTransport{}, "http", nil, nil)
|
||||
rp.AddMapping(proxy.Mapping{
|
||||
ID: rand.Text(),
|
||||
ID: types.ServiceID(rand.Text()),
|
||||
AccountID: types.AccountID(rand.Text()),
|
||||
Host: "app.example.com",
|
||||
Paths: map[string]*proxy.PathTarget{
|
||||
@@ -66,7 +66,7 @@ func BenchmarkServeHTTPHostCount(b *testing.B) {
|
||||
target = id
|
||||
}
|
||||
rp.AddMapping(proxy.Mapping{
|
||||
ID: id,
|
||||
ID: types.ServiceID(id),
|
||||
AccountID: types.AccountID(rand.Text()),
|
||||
Host: host,
|
||||
Paths: map[string]*proxy.PathTarget{
|
||||
@@ -118,7 +118,7 @@ func BenchmarkServeHTTPPathCount(b *testing.B) {
|
||||
}
|
||||
}
|
||||
rp.AddMapping(proxy.Mapping{
|
||||
ID: rand.Text(),
|
||||
ID: types.ServiceID(rand.Text()),
|
||||
AccountID: types.AccountID(rand.Text()),
|
||||
Host: "app.example.com",
|
||||
Paths: paths,
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/roundtrip"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/proxy/web"
|
||||
)
|
||||
|
||||
@@ -65,19 +66,16 @@ func (p *ReverseProxy) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
// Set the serviceId in the context for later retrieval.
|
||||
ctx := withServiceId(r.Context(), result.serviceID)
|
||||
// Set the accountId in the context for later retrieval (for middleware).
|
||||
ctx = withAccountId(ctx, result.accountID)
|
||||
// Set the accountId in the context for the roundtripper to use.
|
||||
ctx := r.Context()
|
||||
// Set the account ID in the context for the roundtripper to use.
|
||||
ctx = roundtrip.WithAccountID(ctx, result.accountID)
|
||||
|
||||
// Also populate captured data if it exists (allows middleware to read after handler completes).
|
||||
// Populate captured data if it exists (allows middleware to read after handler completes).
|
||||
// This solves the problem of passing data UP the middleware chain: we put a mutable struct
|
||||
// pointer in the context, and mutate the struct here so outer middleware can read it.
|
||||
if capturedData := CapturedDataFromContext(ctx); capturedData != nil {
|
||||
capturedData.SetServiceId(result.serviceID)
|
||||
capturedData.SetAccountId(result.accountID)
|
||||
capturedData.SetServiceID(result.serviceID)
|
||||
capturedData.SetAccountID(result.accountID)
|
||||
}
|
||||
|
||||
pt := result.target
|
||||
@@ -86,9 +84,7 @@ func (p *ReverseProxy) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
ctx = roundtrip.WithSkipTLSVerify(ctx)
|
||||
}
|
||||
if pt.RequestTimeout > 0 {
|
||||
var cancel context.CancelFunc
|
||||
ctx, cancel = context.WithTimeout(ctx, pt.RequestTimeout)
|
||||
defer cancel()
|
||||
ctx = types.WithDialTimeout(ctx, pt.RequestTimeout)
|
||||
}
|
||||
|
||||
rewriteMatchedPath := result.matchedPath
|
||||
@@ -97,10 +93,10 @@ func (p *ReverseProxy) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
rp := &httputil.ReverseProxy{
|
||||
Rewrite: p.rewriteFunc(pt.URL, rewriteMatchedPath, result.passHostHeader, pt.PathRewrite, pt.CustomHeaders),
|
||||
Rewrite: p.rewriteFunc(pt.URL, rewriteMatchedPath, result.passHostHeader, pt.PathRewrite, pt.CustomHeaders, result.stripAuthHeaders),
|
||||
Transport: p.transport,
|
||||
FlushInterval: -1,
|
||||
ErrorHandler: proxyErrorHandler,
|
||||
ErrorHandler: p.proxyErrorHandler,
|
||||
}
|
||||
if result.rewriteRedirects {
|
||||
rp.ModifyResponse = p.rewriteLocationFunc(pt.URL, rewriteMatchedPath, r) //nolint:bodyclose
|
||||
@@ -114,7 +110,7 @@ func (p *ReverseProxy) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
// When passHostHeader is true, the original client Host header is preserved
|
||||
// instead of being rewritten to the backend's address.
|
||||
// The pathRewrite parameter controls how the request path is transformed.
|
||||
func (p *ReverseProxy) rewriteFunc(target *url.URL, matchedPath string, passHostHeader bool, pathRewrite PathRewriteMode, customHeaders map[string]string) func(r *httputil.ProxyRequest) {
|
||||
func (p *ReverseProxy) rewriteFunc(target *url.URL, matchedPath string, passHostHeader bool, pathRewrite PathRewriteMode, customHeaders map[string]string, stripAuthHeaders []string) func(r *httputil.ProxyRequest) {
|
||||
return func(r *httputil.ProxyRequest) {
|
||||
switch pathRewrite {
|
||||
case PathRewritePreserve:
|
||||
@@ -138,13 +134,17 @@ func (p *ReverseProxy) rewriteFunc(target *url.URL, matchedPath string, passHost
|
||||
r.Out.Host = target.Host
|
||||
}
|
||||
|
||||
for _, h := range stripAuthHeaders {
|
||||
r.Out.Header.Del(h)
|
||||
}
|
||||
|
||||
for k, v := range customHeaders {
|
||||
r.Out.Header.Set(k, v)
|
||||
}
|
||||
|
||||
clientIP := extractClientIP(r.In.RemoteAddr)
|
||||
clientIP := extractHostIP(r.In.RemoteAddr)
|
||||
|
||||
if IsTrustedProxy(clientIP, p.trustedProxies) {
|
||||
if isTrustedAddr(clientIP, p.trustedProxies) {
|
||||
p.setTrustedForwardingHeaders(r, clientIP)
|
||||
} else {
|
||||
p.setUntrustedForwardingHeaders(r, clientIP)
|
||||
@@ -214,12 +214,14 @@ func normalizeHost(u *url.URL) string {
|
||||
// setTrustedForwardingHeaders appends to the existing forwarding header chain
|
||||
// and preserves upstream-provided headers when the direct connection is from
|
||||
// a trusted proxy.
|
||||
func (p *ReverseProxy) setTrustedForwardingHeaders(r *httputil.ProxyRequest, clientIP string) {
|
||||
func (p *ReverseProxy) setTrustedForwardingHeaders(r *httputil.ProxyRequest, clientIP netip.Addr) {
|
||||
ipStr := clientIP.String()
|
||||
|
||||
// Append the direct connection IP to the existing X-Forwarded-For chain.
|
||||
if existing := r.In.Header.Get("X-Forwarded-For"); existing != "" {
|
||||
r.Out.Header.Set("X-Forwarded-For", existing+", "+clientIP)
|
||||
r.Out.Header.Set("X-Forwarded-For", existing+", "+ipStr)
|
||||
} else {
|
||||
r.Out.Header.Set("X-Forwarded-For", clientIP)
|
||||
r.Out.Header.Set("X-Forwarded-For", ipStr)
|
||||
}
|
||||
|
||||
// Preserve upstream X-Real-IP if present; otherwise resolve through the chain.
|
||||
@@ -227,7 +229,7 @@ func (p *ReverseProxy) setTrustedForwardingHeaders(r *httputil.ProxyRequest, cli
|
||||
r.Out.Header.Set("X-Real-IP", realIP)
|
||||
} else {
|
||||
resolved := ResolveClientIP(r.In.RemoteAddr, r.In.Header.Get("X-Forwarded-For"), p.trustedProxies)
|
||||
r.Out.Header.Set("X-Real-IP", resolved)
|
||||
r.Out.Header.Set("X-Real-IP", resolved.String())
|
||||
}
|
||||
|
||||
// Preserve upstream X-Forwarded-Host if present.
|
||||
@@ -257,10 +259,11 @@ func (p *ReverseProxy) setTrustedForwardingHeaders(r *httputil.ProxyRequest, cli
|
||||
// sets them fresh based on the direct connection. This is the default
|
||||
// behavior when no trusted proxies are configured or the direct connection
|
||||
// is from an untrusted source.
|
||||
func (p *ReverseProxy) setUntrustedForwardingHeaders(r *httputil.ProxyRequest, clientIP string) {
|
||||
func (p *ReverseProxy) setUntrustedForwardingHeaders(r *httputil.ProxyRequest, clientIP netip.Addr) {
|
||||
ipStr := clientIP.String()
|
||||
proto := auth.ResolveProto(p.forwardedProto, r.In.TLS)
|
||||
r.Out.Header.Set("X-Forwarded-For", clientIP)
|
||||
r.Out.Header.Set("X-Real-IP", clientIP)
|
||||
r.Out.Header.Set("X-Forwarded-For", ipStr)
|
||||
r.Out.Header.Set("X-Real-IP", ipStr)
|
||||
r.Out.Header.Set("X-Forwarded-Host", r.In.Host)
|
||||
r.Out.Header.Set("X-Forwarded-Proto", proto)
|
||||
r.Out.Header.Set("X-Forwarded-Port", extractForwardedPort(r.In.Host, proto))
|
||||
@@ -288,16 +291,6 @@ func stripSessionTokenQuery(r *httputil.ProxyRequest) {
|
||||
}
|
||||
}
|
||||
|
||||
// extractClientIP extracts the IP address from an http.Request.RemoteAddr
|
||||
// which is always in host:port format.
|
||||
func extractClientIP(remoteAddr string) string {
|
||||
ip, _, err := net.SplitHostPort(remoteAddr)
|
||||
if err != nil {
|
||||
return remoteAddr
|
||||
}
|
||||
return ip
|
||||
}
|
||||
|
||||
// extractForwardedPort returns the port from the Host header if present,
|
||||
// otherwise defaults to the standard port for the resolved protocol.
|
||||
func extractForwardedPort(host, resolvedProto string) string {
|
||||
@@ -313,7 +306,7 @@ func extractForwardedPort(host, resolvedProto string) string {
|
||||
|
||||
// proxyErrorHandler handles errors from the reverse proxy and serves
|
||||
// user-friendly error pages instead of raw error responses.
|
||||
func proxyErrorHandler(w http.ResponseWriter, r *http.Request, err error) {
|
||||
func (p *ReverseProxy) proxyErrorHandler(w http.ResponseWriter, r *http.Request, err error) {
|
||||
if cd := CapturedDataFromContext(r.Context()); cd != nil {
|
||||
cd.SetOrigin(OriginProxyError)
|
||||
}
|
||||
@@ -321,16 +314,18 @@ func proxyErrorHandler(w http.ResponseWriter, r *http.Request, err error) {
|
||||
clientIP := getClientIP(r)
|
||||
title, message, code, status := classifyProxyError(err)
|
||||
|
||||
log.Warnf("proxy error: request_id=%s client_ip=%s method=%s host=%s path=%s status=%d title=%q err=%v",
|
||||
p.logger.Warnf("proxy error: request_id=%s client_ip=%s method=%s host=%s path=%s status=%d title=%q err=%v",
|
||||
requestID, clientIP, r.Method, r.Host, r.URL.Path, code, title, err)
|
||||
|
||||
web.ServeErrorPage(w, r, code, title, message, requestID, status)
|
||||
}
|
||||
|
||||
// getClientIP retrieves the resolved client IP from context.
|
||||
// getClientIP retrieves the resolved client IP string from context.
|
||||
func getClientIP(r *http.Request) string {
|
||||
if capturedData := CapturedDataFromContext(r.Context()); capturedData != nil {
|
||||
return capturedData.GetClientIP()
|
||||
if ip := capturedData.GetClientIP(); ip.IsValid() {
|
||||
return ip.String()
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
@@ -28,7 +28,7 @@ func TestRewriteFunc_HostRewriting(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
|
||||
t.Run("rewrites host to backend by default", func(t *testing.T) {
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "https://public.example.com/path", "203.0.113.1:12345")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -37,7 +37,7 @@ func TestRewriteFunc_HostRewriting(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("preserves original host when passHostHeader is true", func(t *testing.T) {
|
||||
rewrite := p.rewriteFunc(target, "", true, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", true, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "https://public.example.com/path", "203.0.113.1:12345")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -52,7 +52,7 @@ func TestRewriteFunc_HostRewriting(t *testing.T) {
|
||||
func TestRewriteFunc_XForwardedForStripping(t *testing.T) {
|
||||
target, _ := url.Parse("http://backend.internal:8080")
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
t.Run("sets X-Forwarded-For from direct connection IP", func(t *testing.T) {
|
||||
pr := newProxyRequest(t, "http://example.com/", "203.0.113.50:9999")
|
||||
@@ -89,7 +89,7 @@ func TestRewriteFunc_ForwardedHostAndProto(t *testing.T) {
|
||||
|
||||
t.Run("sets X-Forwarded-Host to original host", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://myapp.example.com:8443/path", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -99,7 +99,7 @@ func TestRewriteFunc_ForwardedHostAndProto(t *testing.T) {
|
||||
|
||||
t.Run("sets X-Forwarded-Port from explicit host port", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://example.com:8443/path", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -109,7 +109,7 @@ func TestRewriteFunc_ForwardedHostAndProto(t *testing.T) {
|
||||
|
||||
t.Run("defaults X-Forwarded-Port to 443 for https", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "https://example.com/", "1.2.3.4:5000")
|
||||
pr.In.TLS = &tls.ConnectionState{}
|
||||
|
||||
@@ -120,7 +120,7 @@ func TestRewriteFunc_ForwardedHostAndProto(t *testing.T) {
|
||||
|
||||
t.Run("defaults X-Forwarded-Port to 80 for http", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -130,7 +130,7 @@ func TestRewriteFunc_ForwardedHostAndProto(t *testing.T) {
|
||||
|
||||
t.Run("auto detects https from TLS", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "https://example.com/", "1.2.3.4:5000")
|
||||
pr.In.TLS = &tls.ConnectionState{}
|
||||
|
||||
@@ -141,7 +141,7 @@ func TestRewriteFunc_ForwardedHostAndProto(t *testing.T) {
|
||||
|
||||
t.Run("auto detects http without TLS", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -151,7 +151,7 @@ func TestRewriteFunc_ForwardedHostAndProto(t *testing.T) {
|
||||
|
||||
t.Run("forced proto overrides TLS detection", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "https"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/", "1.2.3.4:5000")
|
||||
// No TLS, but forced to https
|
||||
|
||||
@@ -162,7 +162,7 @@ func TestRewriteFunc_ForwardedHostAndProto(t *testing.T) {
|
||||
|
||||
t.Run("forced http proto", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "http"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "https://example.com/", "1.2.3.4:5000")
|
||||
pr.In.TLS = &tls.ConnectionState{}
|
||||
|
||||
@@ -175,7 +175,7 @@ func TestRewriteFunc_ForwardedHostAndProto(t *testing.T) {
|
||||
func TestRewriteFunc_SessionCookieStripping(t *testing.T) {
|
||||
target, _ := url.Parse("http://backend.internal:8080")
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
t.Run("strips nb_session cookie", func(t *testing.T) {
|
||||
pr := newProxyRequest(t, "http://example.com/", "1.2.3.4:5000")
|
||||
@@ -220,7 +220,7 @@ func TestRewriteFunc_SessionCookieStripping(t *testing.T) {
|
||||
func TestRewriteFunc_SessionTokenQueryStripping(t *testing.T) {
|
||||
target, _ := url.Parse("http://backend.internal:8080")
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
t.Run("strips session_token query parameter", func(t *testing.T) {
|
||||
pr := newProxyRequest(t, "http://example.com/callback?session_token=secret123&other=keep", "1.2.3.4:5000")
|
||||
@@ -248,7 +248,7 @@ func TestRewriteFunc_URLRewriting(t *testing.T) {
|
||||
|
||||
t.Run("rewrites URL to target with path prefix", func(t *testing.T) {
|
||||
target, _ := url.Parse("http://backend.internal:8080/app")
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/somepath", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -261,7 +261,7 @@ func TestRewriteFunc_URLRewriting(t *testing.T) {
|
||||
|
||||
t.Run("strips matched path prefix to avoid duplication", func(t *testing.T) {
|
||||
target, _ := url.Parse("https://backend.example.org:443/app")
|
||||
rewrite := p.rewriteFunc(target, "/app", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "/app", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/app", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -274,7 +274,7 @@ func TestRewriteFunc_URLRewriting(t *testing.T) {
|
||||
|
||||
t.Run("strips matched prefix and preserves subpath", func(t *testing.T) {
|
||||
target, _ := url.Parse("https://backend.example.org:443/app")
|
||||
rewrite := p.rewriteFunc(target, "/app", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "/app", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/app/article/123", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -284,23 +284,23 @@ func TestRewriteFunc_URLRewriting(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestExtractClientIP(t *testing.T) {
|
||||
func TestExtractHostIP(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
remoteAddr string
|
||||
expected string
|
||||
expected netip.Addr
|
||||
}{
|
||||
{"IPv4 with port", "192.168.1.1:12345", "192.168.1.1"},
|
||||
{"IPv6 with port", "[::1]:12345", "::1"},
|
||||
{"IPv6 full with port", "[2001:db8::1]:443", "2001:db8::1"},
|
||||
{"IPv4 without port fallback", "192.168.1.1", "192.168.1.1"},
|
||||
{"IPv6 without brackets fallback", "::1", "::1"},
|
||||
{"empty string fallback", "", ""},
|
||||
{"public IP", "203.0.113.50:9999", "203.0.113.50"},
|
||||
{"IPv4 with port", "192.168.1.1:12345", netip.MustParseAddr("192.168.1.1")},
|
||||
{"IPv6 with port", "[::1]:12345", netip.MustParseAddr("::1")},
|
||||
{"IPv6 full with port", "[2001:db8::1]:443", netip.MustParseAddr("2001:db8::1")},
|
||||
{"IPv4 without port fallback", "192.168.1.1", netip.MustParseAddr("192.168.1.1")},
|
||||
{"IPv6 without brackets fallback", "::1", netip.MustParseAddr("::1")},
|
||||
{"empty string fallback", "", netip.Addr{}},
|
||||
{"public IP", "203.0.113.50:9999", netip.MustParseAddr("203.0.113.50")},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
assert.Equal(t, tt.expected, extractClientIP(tt.remoteAddr))
|
||||
assert.Equal(t, tt.expected, extractHostIP(tt.remoteAddr))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -332,7 +332,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("appends to X-Forwarded-For", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto", trustedProxies: trusted}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://example.com/", "10.0.0.1:5000")
|
||||
pr.In.Header.Set("X-Forwarded-For", "203.0.113.50")
|
||||
@@ -344,7 +344,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("preserves upstream X-Real-IP", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto", trustedProxies: trusted}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://example.com/", "10.0.0.1:5000")
|
||||
pr.In.Header.Set("X-Forwarded-For", "203.0.113.50")
|
||||
@@ -357,7 +357,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("resolves X-Real-IP from XFF when not set by upstream", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto", trustedProxies: trusted}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://example.com/", "10.0.0.1:5000")
|
||||
pr.In.Header.Set("X-Forwarded-For", "203.0.113.50, 10.0.0.2")
|
||||
@@ -370,7 +370,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("preserves upstream X-Forwarded-Host", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto", trustedProxies: trusted}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://proxy.internal/", "10.0.0.1:5000")
|
||||
pr.In.Header.Set("X-Forwarded-Host", "original.example.com")
|
||||
@@ -382,7 +382,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("preserves upstream X-Forwarded-Proto", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto", trustedProxies: trusted}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://example.com/", "10.0.0.1:5000")
|
||||
pr.In.Header.Set("X-Forwarded-Proto", "https")
|
||||
@@ -394,7 +394,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("preserves upstream X-Forwarded-Port", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto", trustedProxies: trusted}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://example.com/", "10.0.0.1:5000")
|
||||
pr.In.Header.Set("X-Forwarded-Port", "8443")
|
||||
@@ -406,7 +406,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("falls back to local proto when upstream does not set it", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "https", trustedProxies: trusted}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://example.com/", "10.0.0.1:5000")
|
||||
|
||||
@@ -418,7 +418,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("sets X-Forwarded-Host from request when upstream does not set it", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto", trustedProxies: trusted}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://example.com/", "10.0.0.1:5000")
|
||||
|
||||
@@ -429,7 +429,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("untrusted RemoteAddr strips headers even with trusted list", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto", trustedProxies: trusted}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://example.com/", "203.0.113.50:9999")
|
||||
pr.In.Header.Set("X-Forwarded-For", "10.0.0.1, 172.16.0.1")
|
||||
@@ -454,7 +454,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("empty trusted list behaves as untrusted", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto", trustedProxies: nil}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://example.com/", "10.0.0.1:5000")
|
||||
pr.In.Header.Set("X-Forwarded-For", "203.0.113.50")
|
||||
@@ -467,7 +467,7 @@ func TestRewriteFunc_TrustedProxy(t *testing.T) {
|
||||
|
||||
t.Run("XFF starts fresh when trusted proxy has no upstream XFF", func(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto", trustedProxies: trusted}
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "", false, PathRewriteDefault, nil, nil)
|
||||
|
||||
pr := newProxyRequest(t, "http://example.com/", "10.0.0.1:5000")
|
||||
|
||||
@@ -490,7 +490,7 @@ func TestRewriteFunc_PathForwarding(t *testing.T) {
|
||||
t.Run("path prefix baked into target URL is a no-op", func(t *testing.T) {
|
||||
// Management builds: path="/heise", target="https://heise.de:443/heise"
|
||||
target, _ := url.Parse("https://heise.de:443/heise")
|
||||
rewrite := p.rewriteFunc(target, "/heise", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "/heise", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://external.test/heise", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -501,7 +501,7 @@ func TestRewriteFunc_PathForwarding(t *testing.T) {
|
||||
|
||||
t.Run("subpath under prefix also preserved", func(t *testing.T) {
|
||||
target, _ := url.Parse("https://heise.de:443/heise")
|
||||
rewrite := p.rewriteFunc(target, "/heise", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "/heise", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://external.test/heise/article/123", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -513,7 +513,7 @@ func TestRewriteFunc_PathForwarding(t *testing.T) {
|
||||
// What the behavior WOULD be if target URL had no path (true stripping)
|
||||
t.Run("target without path prefix gives true stripping", func(t *testing.T) {
|
||||
target, _ := url.Parse("https://heise.de:443")
|
||||
rewrite := p.rewriteFunc(target, "/heise", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "/heise", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://external.test/heise", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -524,7 +524,7 @@ func TestRewriteFunc_PathForwarding(t *testing.T) {
|
||||
|
||||
t.Run("target without path prefix strips and preserves subpath", func(t *testing.T) {
|
||||
target, _ := url.Parse("https://heise.de:443")
|
||||
rewrite := p.rewriteFunc(target, "/heise", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "/heise", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://external.test/heise/article/123", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -536,7 +536,7 @@ func TestRewriteFunc_PathForwarding(t *testing.T) {
|
||||
// Root path "/" — no stripping expected
|
||||
t.Run("root path forwards full request path unchanged", func(t *testing.T) {
|
||||
target, _ := url.Parse("https://backend.example.com:443/")
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://external.test/heise", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -551,7 +551,7 @@ func TestRewriteFunc_PreservePath(t *testing.T) {
|
||||
target, _ := url.Parse("http://backend.internal:8080")
|
||||
|
||||
t.Run("preserve keeps full request path", func(t *testing.T) {
|
||||
rewrite := p.rewriteFunc(target, "/api", false, PathRewritePreserve, nil)
|
||||
rewrite := p.rewriteFunc(target, "/api", false, PathRewritePreserve, nil, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/api/users/123", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -561,7 +561,7 @@ func TestRewriteFunc_PreservePath(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("preserve with root matchedPath", func(t *testing.T) {
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewritePreserve, nil)
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewritePreserve, nil, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/anything", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -579,7 +579,7 @@ func TestRewriteFunc_CustomHeaders(t *testing.T) {
|
||||
"X-Custom-Auth": "token-abc",
|
||||
"X-Env": "production",
|
||||
}
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewriteDefault, headers)
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewriteDefault, headers, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -589,7 +589,7 @@ func TestRewriteFunc_CustomHeaders(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("nil customHeaders is fine", func(t *testing.T) {
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewriteDefault, nil)
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewriteDefault, nil, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
@@ -599,7 +599,7 @@ func TestRewriteFunc_CustomHeaders(t *testing.T) {
|
||||
|
||||
t.Run("custom headers override existing request headers", func(t *testing.T) {
|
||||
headers := map[string]string{"X-Override": "new-value"}
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewriteDefault, headers)
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewriteDefault, headers, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/", "1.2.3.4:5000")
|
||||
pr.In.Header.Set("X-Override", "old-value")
|
||||
|
||||
@@ -609,11 +609,38 @@ func TestRewriteFunc_CustomHeaders(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestRewriteFunc_StripsAuthorizationHeader(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
target, _ := url.Parse("http://backend.internal:8080")
|
||||
|
||||
t.Run("strips incoming Authorization when no custom Authorization set", func(t *testing.T) {
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewriteDefault, nil, []string{"Authorization"})
|
||||
pr := newProxyRequest(t, "http://example.com/", "1.2.3.4:5000")
|
||||
pr.In.Header.Set("Authorization", "Bearer proxy-token")
|
||||
|
||||
rewrite(pr)
|
||||
|
||||
assert.Empty(t, pr.Out.Header.Get("Authorization"), "Authorization should be stripped")
|
||||
})
|
||||
|
||||
t.Run("custom Authorization replaces incoming", func(t *testing.T) {
|
||||
headers := map[string]string{"Authorization": "Basic YmFja2VuZDpzZWNyZXQ="}
|
||||
rewrite := p.rewriteFunc(target, "/", false, PathRewriteDefault, headers, []string{"Authorization"})
|
||||
pr := newProxyRequest(t, "http://example.com/", "1.2.3.4:5000")
|
||||
pr.In.Header.Set("Authorization", "Bearer proxy-token")
|
||||
|
||||
rewrite(pr)
|
||||
|
||||
assert.Equal(t, "Basic YmFja2VuZDpzZWNyZXQ=", pr.Out.Header.Get("Authorization"),
|
||||
"backend Authorization from custom headers should be set")
|
||||
})
|
||||
}
|
||||
|
||||
func TestRewriteFunc_PreservePathWithCustomHeaders(t *testing.T) {
|
||||
p := &ReverseProxy{forwardedProto: "auto"}
|
||||
target, _ := url.Parse("http://backend.internal:8080")
|
||||
|
||||
rewrite := p.rewriteFunc(target, "/api", false, PathRewritePreserve, map[string]string{"X-Via": "proxy"})
|
||||
rewrite := p.rewriteFunc(target, "/api", false, PathRewritePreserve, map[string]string{"X-Via": "proxy"}, nil)
|
||||
pr := newProxyRequest(t, "http://example.com/api/deep/path", "1.2.3.4:5000")
|
||||
|
||||
rewrite(pr)
|
||||
|
||||
@@ -30,22 +30,29 @@ type PathTarget struct {
|
||||
CustomHeaders map[string]string
|
||||
}
|
||||
|
||||
// Mapping describes how a domain is routed by the HTTP reverse proxy.
|
||||
type Mapping struct {
|
||||
ID string
|
||||
ID types.ServiceID
|
||||
AccountID types.AccountID
|
||||
Host string
|
||||
Paths map[string]*PathTarget
|
||||
PassHostHeader bool
|
||||
RewriteRedirects bool
|
||||
// StripAuthHeaders are header names used for header-based auth.
|
||||
// These headers are stripped from requests before forwarding.
|
||||
StripAuthHeaders []string
|
||||
// sortedPaths caches the paths sorted by length (longest first).
|
||||
sortedPaths []string
|
||||
}
|
||||
|
||||
type targetResult struct {
|
||||
target *PathTarget
|
||||
matchedPath string
|
||||
serviceID string
|
||||
serviceID types.ServiceID
|
||||
accountID types.AccountID
|
||||
passHostHeader bool
|
||||
rewriteRedirects bool
|
||||
stripAuthHeaders []string
|
||||
}
|
||||
|
||||
func (p *ReverseProxy) findTargetForRequest(req *http.Request) (targetResult, bool) {
|
||||
@@ -64,16 +71,7 @@ func (p *ReverseProxy) findTargetForRequest(req *http.Request) (targetResult, bo
|
||||
return targetResult{}, false
|
||||
}
|
||||
|
||||
// Sort paths by length (longest first) in a naive attempt to match the most specific route first.
|
||||
paths := make([]string, 0, len(m.Paths))
|
||||
for path := range m.Paths {
|
||||
paths = append(paths, path)
|
||||
}
|
||||
sort.Slice(paths, func(i, j int) bool {
|
||||
return len(paths[i]) > len(paths[j])
|
||||
})
|
||||
|
||||
for _, path := range paths {
|
||||
for _, path := range m.sortedPaths {
|
||||
if strings.HasPrefix(req.URL.Path, path) {
|
||||
pt := m.Paths[path]
|
||||
if pt == nil || pt.URL == nil {
|
||||
@@ -88,6 +86,7 @@ func (p *ReverseProxy) findTargetForRequest(req *http.Request) (targetResult, bo
|
||||
accountID: m.AccountID,
|
||||
passHostHeader: m.PassHostHeader,
|
||||
rewriteRedirects: m.RewriteRedirects,
|
||||
stripAuthHeaders: m.StripAuthHeaders,
|
||||
}, true
|
||||
}
|
||||
}
|
||||
@@ -95,14 +94,30 @@ func (p *ReverseProxy) findTargetForRequest(req *http.Request) (targetResult, bo
|
||||
return targetResult{}, false
|
||||
}
|
||||
|
||||
// AddMapping registers a host-to-backend mapping for the reverse proxy.
|
||||
func (p *ReverseProxy) AddMapping(m Mapping) {
|
||||
// Sort paths longest-first to match the most specific route first.
|
||||
paths := make([]string, 0, len(m.Paths))
|
||||
for path := range m.Paths {
|
||||
paths = append(paths, path)
|
||||
}
|
||||
sort.Slice(paths, func(i, j int) bool {
|
||||
return len(paths[i]) > len(paths[j])
|
||||
})
|
||||
m.sortedPaths = paths
|
||||
|
||||
p.mappingsMux.Lock()
|
||||
defer p.mappingsMux.Unlock()
|
||||
p.mappings[m.Host] = m
|
||||
}
|
||||
|
||||
func (p *ReverseProxy) RemoveMapping(m Mapping) {
|
||||
// RemoveMapping removes the mapping for the given host and reports whether it existed.
|
||||
func (p *ReverseProxy) RemoveMapping(m Mapping) bool {
|
||||
p.mappingsMux.Lock()
|
||||
defer p.mappingsMux.Unlock()
|
||||
if _, ok := p.mappings[m.Host]; !ok {
|
||||
return false
|
||||
}
|
||||
delete(p.mappings, m.Host)
|
||||
return true
|
||||
}
|
||||
|
||||
@@ -7,21 +7,11 @@ import (
|
||||
|
||||
// IsTrustedProxy checks if the given IP string falls within any of the trusted prefixes.
|
||||
func IsTrustedProxy(ipStr string, trusted []netip.Prefix) bool {
|
||||
if len(trusted) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
addr, err := netip.ParseAddr(ipStr)
|
||||
if err != nil {
|
||||
if err != nil || len(trusted) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
for _, prefix := range trusted {
|
||||
if prefix.Contains(addr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
return isTrustedAddr(addr.Unmap(), trusted)
|
||||
}
|
||||
|
||||
// ResolveClientIP extracts the real client IP from X-Forwarded-For using the trusted proxy list.
|
||||
@@ -30,10 +20,10 @@ func IsTrustedProxy(ipStr string, trusted []netip.Prefix) bool {
|
||||
//
|
||||
// If the trusted list is empty or remoteAddr is not trusted, it returns the
|
||||
// remoteAddr IP directly (ignoring any forwarding headers).
|
||||
func ResolveClientIP(remoteAddr, xff string, trusted []netip.Prefix) string {
|
||||
remoteIP := extractClientIP(remoteAddr)
|
||||
func ResolveClientIP(remoteAddr, xff string, trusted []netip.Prefix) netip.Addr {
|
||||
remoteIP := extractHostIP(remoteAddr)
|
||||
|
||||
if len(trusted) == 0 || !IsTrustedProxy(remoteIP, trusted) {
|
||||
if len(trusted) == 0 || !isTrustedAddr(remoteIP, trusted) {
|
||||
return remoteIP
|
||||
}
|
||||
|
||||
@@ -47,14 +37,45 @@ func ResolveClientIP(remoteAddr, xff string, trusted []netip.Prefix) string {
|
||||
if ip == "" {
|
||||
continue
|
||||
}
|
||||
if !IsTrustedProxy(ip, trusted) {
|
||||
return ip
|
||||
addr, err := netip.ParseAddr(ip)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
addr = addr.Unmap()
|
||||
if !isTrustedAddr(addr, trusted) {
|
||||
return addr
|
||||
}
|
||||
}
|
||||
|
||||
// All IPs in XFF are trusted; return the leftmost as best guess.
|
||||
if first := strings.TrimSpace(parts[0]); first != "" {
|
||||
return first
|
||||
if addr, err := netip.ParseAddr(first); err == nil {
|
||||
return addr.Unmap()
|
||||
}
|
||||
}
|
||||
return remoteIP
|
||||
}
|
||||
|
||||
// extractHostIP parses the IP from a host:port string and returns it unmapped.
|
||||
func extractHostIP(hostPort string) netip.Addr {
|
||||
if ap, err := netip.ParseAddrPort(hostPort); err == nil {
|
||||
return ap.Addr().Unmap()
|
||||
}
|
||||
if addr, err := netip.ParseAddr(hostPort); err == nil {
|
||||
return addr.Unmap()
|
||||
}
|
||||
return netip.Addr{}
|
||||
}
|
||||
|
||||
// isTrustedAddr checks if the given address falls within any of the trusted prefixes.
|
||||
func isTrustedAddr(addr netip.Addr, trusted []netip.Prefix) bool {
|
||||
if !addr.IsValid() {
|
||||
return false
|
||||
}
|
||||
for _, prefix := range trusted {
|
||||
if prefix.Contains(addr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -48,77 +48,77 @@ func TestResolveClientIP(t *testing.T) {
|
||||
remoteAddr string
|
||||
xff string
|
||||
trusted []netip.Prefix
|
||||
want string
|
||||
want netip.Addr
|
||||
}{
|
||||
{
|
||||
name: "empty trusted list returns RemoteAddr",
|
||||
remoteAddr: "203.0.113.50:9999",
|
||||
xff: "1.2.3.4",
|
||||
trusted: nil,
|
||||
want: "203.0.113.50",
|
||||
want: netip.MustParseAddr("203.0.113.50"),
|
||||
},
|
||||
{
|
||||
name: "untrusted RemoteAddr ignores XFF",
|
||||
remoteAddr: "203.0.113.50:9999",
|
||||
xff: "1.2.3.4, 10.0.0.1",
|
||||
trusted: trusted,
|
||||
want: "203.0.113.50",
|
||||
want: netip.MustParseAddr("203.0.113.50"),
|
||||
},
|
||||
{
|
||||
name: "trusted RemoteAddr with single client in XFF",
|
||||
remoteAddr: "10.0.0.1:5000",
|
||||
xff: "203.0.113.50",
|
||||
trusted: trusted,
|
||||
want: "203.0.113.50",
|
||||
want: netip.MustParseAddr("203.0.113.50"),
|
||||
},
|
||||
{
|
||||
name: "trusted RemoteAddr walks past trusted entries in XFF",
|
||||
remoteAddr: "10.0.0.1:5000",
|
||||
xff: "203.0.113.50, 10.0.0.2, 172.16.0.5",
|
||||
trusted: trusted,
|
||||
want: "203.0.113.50",
|
||||
want: netip.MustParseAddr("203.0.113.50"),
|
||||
},
|
||||
{
|
||||
name: "trusted RemoteAddr with empty XFF falls back to RemoteAddr",
|
||||
remoteAddr: "10.0.0.1:5000",
|
||||
xff: "",
|
||||
trusted: trusted,
|
||||
want: "10.0.0.1",
|
||||
want: netip.MustParseAddr("10.0.0.1"),
|
||||
},
|
||||
{
|
||||
name: "all XFF IPs trusted returns leftmost",
|
||||
remoteAddr: "10.0.0.1:5000",
|
||||
xff: "10.0.0.2, 172.16.0.1, 10.0.0.3",
|
||||
trusted: trusted,
|
||||
want: "10.0.0.2",
|
||||
want: netip.MustParseAddr("10.0.0.2"),
|
||||
},
|
||||
{
|
||||
name: "XFF with whitespace",
|
||||
remoteAddr: "10.0.0.1:5000",
|
||||
xff: " 203.0.113.50 , 10.0.0.2 ",
|
||||
trusted: trusted,
|
||||
want: "203.0.113.50",
|
||||
want: netip.MustParseAddr("203.0.113.50"),
|
||||
},
|
||||
{
|
||||
name: "XFF with empty segments",
|
||||
remoteAddr: "10.0.0.1:5000",
|
||||
xff: "203.0.113.50,,10.0.0.2",
|
||||
trusted: trusted,
|
||||
want: "203.0.113.50",
|
||||
want: netip.MustParseAddr("203.0.113.50"),
|
||||
},
|
||||
{
|
||||
name: "multi-hop with mixed trust",
|
||||
remoteAddr: "10.0.0.1:5000",
|
||||
xff: "8.8.8.8, 203.0.113.50, 172.16.0.1",
|
||||
trusted: trusted,
|
||||
want: "203.0.113.50",
|
||||
want: netip.MustParseAddr("203.0.113.50"),
|
||||
},
|
||||
{
|
||||
name: "RemoteAddr without port",
|
||||
remoteAddr: "10.0.0.1",
|
||||
xff: "203.0.113.50",
|
||||
trusted: trusted,
|
||||
want: "203.0.113.50",
|
||||
want: netip.MustParseAddr("203.0.113.50"),
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
|
||||
@@ -0,0 +1,183 @@
|
||||
// Package restrict provides connection-level access control based on
|
||||
// IP CIDR ranges and geolocation (country codes).
|
||||
package restrict
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strings"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/geolocation"
|
||||
)
|
||||
|
||||
// GeoResolver resolves an IP address to geographic information.
|
||||
type GeoResolver interface {
|
||||
LookupAddr(addr netip.Addr) geolocation.Result
|
||||
Available() bool
|
||||
}
|
||||
|
||||
// Filter evaluates IP restrictions. CIDR checks are performed first
|
||||
// (cheap), followed by country lookups (more expensive) only when needed.
|
||||
type Filter struct {
|
||||
AllowedCIDRs []netip.Prefix
|
||||
BlockedCIDRs []netip.Prefix
|
||||
AllowedCountries []string
|
||||
BlockedCountries []string
|
||||
}
|
||||
|
||||
// ParseFilter builds a Filter from the raw string slices. Returns nil
|
||||
// if all slices are empty.
|
||||
func ParseFilter(allowedCIDRs, blockedCIDRs, allowedCountries, blockedCountries []string) *Filter {
|
||||
if len(allowedCIDRs) == 0 && len(blockedCIDRs) == 0 &&
|
||||
len(allowedCountries) == 0 && len(blockedCountries) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
f := &Filter{
|
||||
AllowedCountries: normalizeCountryCodes(allowedCountries),
|
||||
BlockedCountries: normalizeCountryCodes(blockedCountries),
|
||||
}
|
||||
for _, cidr := range allowedCIDRs {
|
||||
prefix, err := netip.ParsePrefix(cidr)
|
||||
if err != nil {
|
||||
log.Warnf("skip invalid allowed CIDR %q: %v", cidr, err)
|
||||
continue
|
||||
}
|
||||
f.AllowedCIDRs = append(f.AllowedCIDRs, prefix.Masked())
|
||||
}
|
||||
for _, cidr := range blockedCIDRs {
|
||||
prefix, err := netip.ParsePrefix(cidr)
|
||||
if err != nil {
|
||||
log.Warnf("skip invalid blocked CIDR %q: %v", cidr, err)
|
||||
continue
|
||||
}
|
||||
f.BlockedCIDRs = append(f.BlockedCIDRs, prefix.Masked())
|
||||
}
|
||||
return f
|
||||
}
|
||||
|
||||
func normalizeCountryCodes(codes []string) []string {
|
||||
if len(codes) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make([]string, len(codes))
|
||||
for i, c := range codes {
|
||||
out[i] = strings.ToUpper(c)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// Verdict is the result of an access check.
|
||||
type Verdict int
|
||||
|
||||
const (
|
||||
// Allow indicates the address passed all checks.
|
||||
Allow Verdict = iota
|
||||
// DenyCIDR indicates the address was blocked by a CIDR rule.
|
||||
DenyCIDR
|
||||
// DenyCountry indicates the address was blocked by a country rule.
|
||||
DenyCountry
|
||||
// DenyGeoUnavailable indicates that country restrictions are configured
|
||||
// but the geo lookup is unavailable.
|
||||
DenyGeoUnavailable
|
||||
)
|
||||
|
||||
// String returns the deny reason string matching the HTTP auth mechanism names.
|
||||
func (v Verdict) String() string {
|
||||
switch v {
|
||||
case Allow:
|
||||
return "allow"
|
||||
case DenyCIDR:
|
||||
return "ip_restricted"
|
||||
case DenyCountry:
|
||||
return "country_restricted"
|
||||
case DenyGeoUnavailable:
|
||||
return "geo_unavailable"
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
}
|
||||
|
||||
// Check evaluates whether addr is permitted. CIDR rules are evaluated
|
||||
// first because they are O(n) prefix comparisons. Country rules run
|
||||
// only when CIDR checks pass and require a geo lookup.
|
||||
func (f *Filter) Check(addr netip.Addr, geo GeoResolver) Verdict {
|
||||
if f == nil {
|
||||
return Allow
|
||||
}
|
||||
|
||||
// Normalize v4-mapped-v6 (e.g. ::ffff:10.1.2.3) to plain v4 so that
|
||||
// IPv4 CIDR rules match regardless of how the address was received.
|
||||
addr = addr.Unmap()
|
||||
|
||||
if v := f.checkCIDR(addr); v != Allow {
|
||||
return v
|
||||
}
|
||||
return f.checkCountry(addr, geo)
|
||||
}
|
||||
|
||||
func (f *Filter) checkCIDR(addr netip.Addr) Verdict {
|
||||
if len(f.AllowedCIDRs) > 0 {
|
||||
allowed := false
|
||||
for _, prefix := range f.AllowedCIDRs {
|
||||
if prefix.Contains(addr) {
|
||||
allowed = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !allowed {
|
||||
return DenyCIDR
|
||||
}
|
||||
}
|
||||
|
||||
for _, prefix := range f.BlockedCIDRs {
|
||||
if prefix.Contains(addr) {
|
||||
return DenyCIDR
|
||||
}
|
||||
}
|
||||
return Allow
|
||||
}
|
||||
|
||||
func (f *Filter) checkCountry(addr netip.Addr, geo GeoResolver) Verdict {
|
||||
if len(f.AllowedCountries) == 0 && len(f.BlockedCountries) == 0 {
|
||||
return Allow
|
||||
}
|
||||
|
||||
if geo == nil || !geo.Available() {
|
||||
return DenyGeoUnavailable
|
||||
}
|
||||
|
||||
result := geo.LookupAddr(addr)
|
||||
if result.CountryCode == "" {
|
||||
// Unknown country: deny if an allowlist is active, allow otherwise.
|
||||
// Blocklists are best-effort: unknown countries pass through since
|
||||
// the default policy is allow.
|
||||
if len(f.AllowedCountries) > 0 {
|
||||
return DenyCountry
|
||||
}
|
||||
return Allow
|
||||
}
|
||||
|
||||
if len(f.AllowedCountries) > 0 {
|
||||
if !slices.Contains(f.AllowedCountries, result.CountryCode) {
|
||||
return DenyCountry
|
||||
}
|
||||
}
|
||||
|
||||
if slices.Contains(f.BlockedCountries, result.CountryCode) {
|
||||
return DenyCountry
|
||||
}
|
||||
|
||||
return Allow
|
||||
}
|
||||
|
||||
// HasRestrictions returns true if any restriction rules are configured.
|
||||
func (f *Filter) HasRestrictions() bool {
|
||||
if f == nil {
|
||||
return false
|
||||
}
|
||||
return len(f.AllowedCIDRs) > 0 || len(f.BlockedCIDRs) > 0 ||
|
||||
len(f.AllowedCountries) > 0 || len(f.BlockedCountries) > 0
|
||||
}
|
||||
@@ -0,0 +1,278 @@
|
||||
package restrict
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/geolocation"
|
||||
)
|
||||
|
||||
type mockGeo struct {
|
||||
countries map[string]string
|
||||
}
|
||||
|
||||
func (m *mockGeo) LookupAddr(addr netip.Addr) geolocation.Result {
|
||||
return geolocation.Result{CountryCode: m.countries[addr.String()]}
|
||||
}
|
||||
|
||||
func (m *mockGeo) Available() bool { return true }
|
||||
|
||||
func newMockGeo(entries map[string]string) *mockGeo {
|
||||
return &mockGeo{countries: entries}
|
||||
}
|
||||
|
||||
func TestFilter_Check_NilFilter(t *testing.T) {
|
||||
var f *Filter
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("1.2.3.4"), nil))
|
||||
}
|
||||
|
||||
func TestFilter_Check_AllowedCIDR(t *testing.T) {
|
||||
f := ParseFilter([]string{"10.0.0.0/8"}, nil, nil, nil)
|
||||
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("10.1.2.3"), nil))
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("192.168.1.1"), nil))
|
||||
}
|
||||
|
||||
func TestFilter_Check_BlockedCIDR(t *testing.T) {
|
||||
f := ParseFilter(nil, []string{"10.0.0.0/8"}, nil, nil)
|
||||
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("10.1.2.3"), nil))
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("192.168.1.1"), nil))
|
||||
}
|
||||
|
||||
func TestFilter_Check_AllowedAndBlockedCIDR(t *testing.T) {
|
||||
f := ParseFilter([]string{"10.0.0.0/8"}, []string{"10.1.0.0/16"}, nil, nil)
|
||||
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("10.2.3.4"), nil), "allowed by allowlist, not in blocklist")
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("10.1.2.3"), nil), "allowed by allowlist but in blocklist")
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("192.168.1.1"), nil), "not in allowlist")
|
||||
}
|
||||
|
||||
func TestFilter_Check_AllowedCountry(t *testing.T) {
|
||||
geo := newMockGeo(map[string]string{
|
||||
"1.1.1.1": "US",
|
||||
"2.2.2.2": "DE",
|
||||
"3.3.3.3": "CN",
|
||||
})
|
||||
f := ParseFilter(nil, nil, []string{"US", "DE"}, nil)
|
||||
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("1.1.1.1"), geo), "US in allowlist")
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("2.2.2.2"), geo), "DE in allowlist")
|
||||
assert.Equal(t, DenyCountry, f.Check(netip.MustParseAddr("3.3.3.3"), geo), "CN not in allowlist")
|
||||
}
|
||||
|
||||
func TestFilter_Check_BlockedCountry(t *testing.T) {
|
||||
geo := newMockGeo(map[string]string{
|
||||
"1.1.1.1": "CN",
|
||||
"2.2.2.2": "RU",
|
||||
"3.3.3.3": "US",
|
||||
})
|
||||
f := ParseFilter(nil, nil, nil, []string{"CN", "RU"})
|
||||
|
||||
assert.Equal(t, DenyCountry, f.Check(netip.MustParseAddr("1.1.1.1"), geo), "CN in blocklist")
|
||||
assert.Equal(t, DenyCountry, f.Check(netip.MustParseAddr("2.2.2.2"), geo), "RU in blocklist")
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("3.3.3.3"), geo), "US not in blocklist")
|
||||
}
|
||||
|
||||
func TestFilter_Check_AllowedAndBlockedCountry(t *testing.T) {
|
||||
geo := newMockGeo(map[string]string{
|
||||
"1.1.1.1": "US",
|
||||
"2.2.2.2": "DE",
|
||||
"3.3.3.3": "CN",
|
||||
})
|
||||
// Allow US and DE, but block DE explicitly.
|
||||
f := ParseFilter(nil, nil, []string{"US", "DE"}, []string{"DE"})
|
||||
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("1.1.1.1"), geo), "US allowed and not blocked")
|
||||
assert.Equal(t, DenyCountry, f.Check(netip.MustParseAddr("2.2.2.2"), geo), "DE allowed but also blocked, block wins")
|
||||
assert.Equal(t, DenyCountry, f.Check(netip.MustParseAddr("3.3.3.3"), geo), "CN not in allowlist")
|
||||
}
|
||||
|
||||
func TestFilter_Check_UnknownCountryWithAllowlist(t *testing.T) {
|
||||
geo := newMockGeo(map[string]string{
|
||||
"1.1.1.1": "US",
|
||||
})
|
||||
f := ParseFilter(nil, nil, []string{"US"}, nil)
|
||||
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("1.1.1.1"), geo), "known US in allowlist")
|
||||
assert.Equal(t, DenyCountry, f.Check(netip.MustParseAddr("9.9.9.9"), geo), "unknown country denied when allowlist is active")
|
||||
}
|
||||
|
||||
func TestFilter_Check_UnknownCountryWithBlocklistOnly(t *testing.T) {
|
||||
geo := newMockGeo(map[string]string{
|
||||
"1.1.1.1": "CN",
|
||||
})
|
||||
f := ParseFilter(nil, nil, nil, []string{"CN"})
|
||||
|
||||
assert.Equal(t, DenyCountry, f.Check(netip.MustParseAddr("1.1.1.1"), geo), "known CN in blocklist")
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("9.9.9.9"), geo), "unknown country allowed when only blocklist is active")
|
||||
}
|
||||
|
||||
func TestFilter_Check_CountryWithoutGeo(t *testing.T) {
|
||||
f := ParseFilter(nil, nil, []string{"US"}, nil)
|
||||
assert.Equal(t, DenyGeoUnavailable, f.Check(netip.MustParseAddr("1.2.3.4"), nil), "nil geo with country allowlist")
|
||||
}
|
||||
|
||||
func TestFilter_Check_CountryBlocklistWithoutGeo(t *testing.T) {
|
||||
f := ParseFilter(nil, nil, nil, []string{"CN"})
|
||||
assert.Equal(t, DenyGeoUnavailable, f.Check(netip.MustParseAddr("1.2.3.4"), nil), "nil geo with country blocklist")
|
||||
}
|
||||
|
||||
func TestFilter_Check_GeoUnavailable(t *testing.T) {
|
||||
geo := &unavailableGeo{}
|
||||
|
||||
f := ParseFilter(nil, nil, []string{"US"}, nil)
|
||||
assert.Equal(t, DenyGeoUnavailable, f.Check(netip.MustParseAddr("1.2.3.4"), geo), "unavailable geo with country allowlist")
|
||||
|
||||
f2 := ParseFilter(nil, nil, nil, []string{"CN"})
|
||||
assert.Equal(t, DenyGeoUnavailable, f2.Check(netip.MustParseAddr("1.2.3.4"), geo), "unavailable geo with country blocklist")
|
||||
}
|
||||
|
||||
func TestFilter_Check_CIDROnlySkipsGeo(t *testing.T) {
|
||||
f := ParseFilter([]string{"10.0.0.0/8"}, nil, nil, nil)
|
||||
|
||||
// CIDR-only filter should never touch geo, so nil geo is fine.
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("10.1.2.3"), nil))
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("192.168.1.1"), nil))
|
||||
}
|
||||
|
||||
func TestFilter_Check_CIDRAllowThenCountryBlock(t *testing.T) {
|
||||
geo := newMockGeo(map[string]string{
|
||||
"10.1.2.3": "CN",
|
||||
"10.2.3.4": "US",
|
||||
})
|
||||
f := ParseFilter([]string{"10.0.0.0/8"}, nil, nil, []string{"CN"})
|
||||
|
||||
assert.Equal(t, DenyCountry, f.Check(netip.MustParseAddr("10.1.2.3"), geo), "CIDR allowed but country blocked")
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("10.2.3.4"), geo), "CIDR allowed and country not blocked")
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("192.168.1.1"), geo), "CIDR denied before country check")
|
||||
}
|
||||
|
||||
func TestParseFilter_Empty(t *testing.T) {
|
||||
f := ParseFilter(nil, nil, nil, nil)
|
||||
assert.Nil(t, f)
|
||||
}
|
||||
|
||||
func TestParseFilter_InvalidCIDR(t *testing.T) {
|
||||
f := ParseFilter([]string{"invalid", "10.0.0.0/8"}, nil, nil, nil)
|
||||
|
||||
assert.NotNil(t, f)
|
||||
assert.Len(t, f.AllowedCIDRs, 1, "invalid CIDR should be skipped")
|
||||
assert.Equal(t, netip.MustParsePrefix("10.0.0.0/8"), f.AllowedCIDRs[0])
|
||||
}
|
||||
|
||||
func TestFilter_HasRestrictions(t *testing.T) {
|
||||
assert.False(t, (*Filter)(nil).HasRestrictions())
|
||||
assert.False(t, (&Filter{}).HasRestrictions())
|
||||
assert.True(t, ParseFilter([]string{"10.0.0.0/8"}, nil, nil, nil).HasRestrictions())
|
||||
assert.True(t, ParseFilter(nil, nil, []string{"US"}, nil).HasRestrictions())
|
||||
}
|
||||
|
||||
func TestFilter_Check_IPv6CIDR(t *testing.T) {
|
||||
f := ParseFilter([]string{"2001:db8::/32"}, nil, nil, nil)
|
||||
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("2001:db8::1"), nil), "v6 addr in v6 allowlist")
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("2001:db9::1"), nil), "v6 addr not in v6 allowlist")
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("10.1.2.3"), nil), "v4 addr not in v6 allowlist")
|
||||
}
|
||||
|
||||
func TestFilter_Check_IPv4MappedIPv6(t *testing.T) {
|
||||
f := ParseFilter([]string{"10.0.0.0/8"}, nil, nil, nil)
|
||||
|
||||
// A v4-mapped-v6 address like ::ffff:10.1.2.3 must match a v4 CIDR.
|
||||
v4mapped := netip.MustParseAddr("::ffff:10.1.2.3")
|
||||
assert.True(t, v4mapped.Is4In6(), "precondition: address is v4-in-v6")
|
||||
assert.Equal(t, Allow, f.Check(v4mapped, nil), "v4-mapped-v6 must match v4 CIDR after Unmap")
|
||||
|
||||
v4mappedOutside := netip.MustParseAddr("::ffff:192.168.1.1")
|
||||
assert.Equal(t, DenyCIDR, f.Check(v4mappedOutside, nil), "v4-mapped-v6 outside v4 CIDR")
|
||||
}
|
||||
|
||||
func TestFilter_Check_MixedV4V6CIDRs(t *testing.T) {
|
||||
f := ParseFilter([]string{"10.0.0.0/8", "2001:db8::/32"}, nil, nil, nil)
|
||||
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("10.1.2.3"), nil), "v4 in v4 CIDR")
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("2001:db8::1"), nil), "v6 in v6 CIDR")
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("::ffff:10.1.2.3"), nil), "v4-mapped matches v4 CIDR")
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("192.168.1.1"), nil), "v4 not in either CIDR")
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("fe80::1"), nil), "v6 not in either CIDR")
|
||||
}
|
||||
|
||||
func TestParseFilter_CanonicalizesNonMaskedCIDR(t *testing.T) {
|
||||
// 1.1.1.1/24 has host bits set; ParseFilter should canonicalize to 1.1.1.0/24.
|
||||
f := ParseFilter([]string{"1.1.1.1/24"}, nil, nil, nil)
|
||||
assert.Equal(t, netip.MustParsePrefix("1.1.1.0/24"), f.AllowedCIDRs[0])
|
||||
|
||||
// Verify it still matches correctly.
|
||||
assert.Equal(t, Allow, f.Check(netip.MustParseAddr("1.1.1.100"), nil))
|
||||
assert.Equal(t, DenyCIDR, f.Check(netip.MustParseAddr("1.1.2.1"), nil))
|
||||
}
|
||||
|
||||
func TestFilter_Check_CountryCodeCaseInsensitive(t *testing.T) {
|
||||
geo := newMockGeo(map[string]string{
|
||||
"1.1.1.1": "US",
|
||||
"2.2.2.2": "DE",
|
||||
"3.3.3.3": "CN",
|
||||
})
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
allowedCountries []string
|
||||
blockedCountries []string
|
||||
addr string
|
||||
want Verdict
|
||||
}{
|
||||
{
|
||||
name: "lowercase allowlist matches uppercase MaxMind code",
|
||||
allowedCountries: []string{"us", "de"},
|
||||
addr: "1.1.1.1",
|
||||
want: Allow,
|
||||
},
|
||||
{
|
||||
name: "mixed-case allowlist matches",
|
||||
allowedCountries: []string{"Us", "dE"},
|
||||
addr: "2.2.2.2",
|
||||
want: Allow,
|
||||
},
|
||||
{
|
||||
name: "lowercase allowlist rejects non-matching country",
|
||||
allowedCountries: []string{"us", "de"},
|
||||
addr: "3.3.3.3",
|
||||
want: DenyCountry,
|
||||
},
|
||||
{
|
||||
name: "lowercase blocklist blocks matching country",
|
||||
blockedCountries: []string{"cn"},
|
||||
addr: "3.3.3.3",
|
||||
want: DenyCountry,
|
||||
},
|
||||
{
|
||||
name: "mixed-case blocklist blocks matching country",
|
||||
blockedCountries: []string{"Cn"},
|
||||
addr: "3.3.3.3",
|
||||
want: DenyCountry,
|
||||
},
|
||||
{
|
||||
name: "lowercase blocklist does not block non-matching country",
|
||||
blockedCountries: []string{"cn"},
|
||||
addr: "1.1.1.1",
|
||||
want: Allow,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
f := ParseFilter(nil, nil, tc.allowedCountries, tc.blockedCountries)
|
||||
got := f.Check(netip.MustParseAddr(tc.addr), geo)
|
||||
assert.Equal(t, tc.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// unavailableGeo simulates a GeoResolver whose database is not loaded.
|
||||
type unavailableGeo struct{}
|
||||
|
||||
func (u *unavailableGeo) LookupAddr(_ netip.Addr) geolocation.Result { return geolocation.Result{} }
|
||||
func (u *unavailableGeo) Available() bool { return false }
|
||||
+140
-116
@@ -5,6 +5,7 @@ import (
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -14,11 +15,12 @@ import (
|
||||
"golang.org/x/exp/maps"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/codes"
|
||||
grpcstatus "google.golang.org/grpc/status"
|
||||
|
||||
"github.com/netbirdio/netbird/client/embed"
|
||||
nberrors "github.com/netbirdio/netbird/client/errors"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/util"
|
||||
)
|
||||
@@ -26,7 +28,22 @@ import (
|
||||
const deviceNamePrefix = "ingress-proxy-"
|
||||
|
||||
// backendKey identifies a backend by its host:port from the target URL.
|
||||
type backendKey = string
|
||||
type backendKey string
|
||||
|
||||
// ServiceKey uniquely identifies a service (HTTP reverse proxy or L4 service)
|
||||
// that holds a reference to an embedded NetBird client. Callers should use the
|
||||
// DomainServiceKey and L4ServiceKey constructors to avoid namespace collisions.
|
||||
type ServiceKey string
|
||||
|
||||
// DomainServiceKey returns a ServiceKey for an HTTP/TLS domain-based service.
|
||||
func DomainServiceKey(domain string) ServiceKey {
|
||||
return ServiceKey("domain:" + domain)
|
||||
}
|
||||
|
||||
// L4ServiceKey returns a ServiceKey for an L4 service (TCP/UDP).
|
||||
func L4ServiceKey(id types.ServiceID) ServiceKey {
|
||||
return ServiceKey("l4:" + id)
|
||||
}
|
||||
|
||||
var (
|
||||
// ErrNoAccountID is returned when a request context is missing the account ID.
|
||||
@@ -39,24 +56,24 @@ var (
|
||||
ErrTooManyInflight = errors.New("too many in-flight requests")
|
||||
)
|
||||
|
||||
// domainInfo holds metadata about a registered domain.
|
||||
type domainInfo struct {
|
||||
serviceID string
|
||||
// serviceInfo holds metadata about a registered service.
|
||||
type serviceInfo struct {
|
||||
serviceID types.ServiceID
|
||||
}
|
||||
|
||||
type domainNotification struct {
|
||||
domain domain.Domain
|
||||
serviceID string
|
||||
type serviceNotification struct {
|
||||
key ServiceKey
|
||||
serviceID types.ServiceID
|
||||
}
|
||||
|
||||
// clientEntry holds an embedded NetBird client and tracks which domains use it.
|
||||
// clientEntry holds an embedded NetBird client and tracks which services use it.
|
||||
type clientEntry struct {
|
||||
client *embed.Client
|
||||
transport *http.Transport
|
||||
// insecureTransport is a clone of transport with TLS verification disabled,
|
||||
// used when per-target skip_tls_verify is set.
|
||||
insecureTransport *http.Transport
|
||||
domains map[domain.Domain]domainInfo
|
||||
services map[ServiceKey]serviceInfo
|
||||
createdAt time.Time
|
||||
started bool
|
||||
// Per-backend in-flight limiting keyed by target host:port.
|
||||
@@ -93,12 +110,12 @@ func (e *clientEntry) acquireInflight(backend backendKey) (release func(), ok bo
|
||||
// ClientConfig holds configuration for the embedded NetBird client.
|
||||
type ClientConfig struct {
|
||||
MgmtAddr string
|
||||
WGPort int
|
||||
WGPort uint16
|
||||
PreSharedKey string
|
||||
}
|
||||
|
||||
type statusNotifier interface {
|
||||
NotifyStatus(ctx context.Context, accountID, serviceID, domain string, connected bool) error
|
||||
NotifyStatus(ctx context.Context, accountID types.AccountID, serviceID types.ServiceID, connected bool) error
|
||||
}
|
||||
|
||||
type managementClient interface {
|
||||
@@ -107,7 +124,7 @@ type managementClient interface {
|
||||
|
||||
// NetBird provides an http.RoundTripper implementation
|
||||
// backed by underlying NetBird connections.
|
||||
// Clients are keyed by AccountID, allowing multiple domains to share the same connection.
|
||||
// Clients are keyed by AccountID, allowing multiple services to share the same connection.
|
||||
type NetBird struct {
|
||||
proxyID string
|
||||
proxyAddr string
|
||||
@@ -124,11 +141,11 @@ type NetBird struct {
|
||||
|
||||
// ClientDebugInfo contains debug information about a client.
|
||||
type ClientDebugInfo struct {
|
||||
AccountID types.AccountID
|
||||
DomainCount int
|
||||
Domains domain.List
|
||||
HasClient bool
|
||||
CreatedAt time.Time
|
||||
AccountID types.AccountID
|
||||
ServiceCount int
|
||||
ServiceKeys []string
|
||||
HasClient bool
|
||||
CreatedAt time.Time
|
||||
}
|
||||
|
||||
// accountIDContextKey is the context key for storing the account ID.
|
||||
@@ -137,37 +154,37 @@ type accountIDContextKey struct{}
|
||||
// skipTLSVerifyContextKey is the context key for requesting insecure TLS.
|
||||
type skipTLSVerifyContextKey struct{}
|
||||
|
||||
// AddPeer registers a domain for an account. If the account doesn't have a client yet,
|
||||
// AddPeer registers a service for an account. If the account doesn't have a client yet,
|
||||
// one is created by authenticating with the management server using the provided token.
|
||||
// Multiple domains can share the same client.
|
||||
func (n *NetBird) AddPeer(ctx context.Context, accountID types.AccountID, d domain.Domain, authToken, serviceID string) error {
|
||||
// Multiple services can share the same client.
|
||||
func (n *NetBird) AddPeer(ctx context.Context, accountID types.AccountID, key ServiceKey, authToken string, serviceID types.ServiceID) error {
|
||||
si := serviceInfo{serviceID: serviceID}
|
||||
|
||||
n.clientsMux.Lock()
|
||||
|
||||
entry, exists := n.clients[accountID]
|
||||
if exists {
|
||||
// Client already exists for this account, just register the domain
|
||||
entry.domains[d] = domainInfo{serviceID: serviceID}
|
||||
entry.services[key] = si
|
||||
started := entry.started
|
||||
n.clientsMux.Unlock()
|
||||
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"domain": d,
|
||||
}).Debug("registered domain with existing client")
|
||||
"account_id": accountID,
|
||||
"service_key": key,
|
||||
}).Debug("registered service with existing client")
|
||||
|
||||
// If client is already started, notify this domain as connected immediately
|
||||
if started && n.statusNotifier != nil {
|
||||
if err := n.statusNotifier.NotifyStatus(ctx, string(accountID), serviceID, string(d), true); err != nil {
|
||||
if err := n.statusNotifier.NotifyStatus(ctx, accountID, serviceID, true); err != nil {
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"domain": d,
|
||||
"account_id": accountID,
|
||||
"service_key": key,
|
||||
}).WithError(err).Warn("failed to notify status for existing client")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
entry, err := n.createClientEntry(ctx, accountID, d, authToken, serviceID)
|
||||
entry, err := n.createClientEntry(ctx, accountID, key, authToken, si)
|
||||
if err != nil {
|
||||
n.clientsMux.Unlock()
|
||||
return err
|
||||
@@ -177,8 +194,8 @@ func (n *NetBird) AddPeer(ctx context.Context, accountID types.AccountID, d doma
|
||||
n.clientsMux.Unlock()
|
||||
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"domain": d,
|
||||
"account_id": accountID,
|
||||
"service_key": key,
|
||||
}).Info("created new client for account")
|
||||
|
||||
// Attempt to start the client in the background; if this fails we will
|
||||
@@ -190,7 +207,8 @@ func (n *NetBird) AddPeer(ctx context.Context, accountID types.AccountID, d doma
|
||||
|
||||
// createClientEntry generates a WireGuard keypair, authenticates with management,
|
||||
// and creates an embedded NetBird client. Must be called with clientsMux held.
|
||||
func (n *NetBird) createClientEntry(ctx context.Context, accountID types.AccountID, d domain.Domain, authToken, serviceID string) (*clientEntry, error) {
|
||||
func (n *NetBird) createClientEntry(ctx context.Context, accountID types.AccountID, key ServiceKey, authToken string, si serviceInfo) (*clientEntry, error) {
|
||||
serviceID := si.serviceID
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"service_id": serviceID,
|
||||
@@ -209,7 +227,7 @@ func (n *NetBird) createClientEntry(ctx context.Context, accountID types.Account
|
||||
}).Debug("authenticating new proxy peer with management")
|
||||
|
||||
resp, err := n.mgmtClient.CreateProxyPeer(ctx, &proto.CreateProxyPeerRequest{
|
||||
ServiceId: serviceID,
|
||||
ServiceId: string(serviceID),
|
||||
AccountId: string(accountID),
|
||||
Token: authToken,
|
||||
WireguardPublicKey: publicKey.String(),
|
||||
@@ -240,13 +258,14 @@ func (n *NetBird) createClientEntry(ctx context.Context, accountID types.Account
|
||||
|
||||
// Create embedded NetBird client with the generated private key.
|
||||
// The peer has already been created via CreateProxyPeer RPC with the public key.
|
||||
wgPort := int(n.clientCfg.WGPort)
|
||||
client, err := embed.New(embed.Options{
|
||||
DeviceName: deviceNamePrefix + n.proxyID,
|
||||
ManagementURL: n.clientCfg.MgmtAddr,
|
||||
PrivateKey: privateKey.String(),
|
||||
LogLevel: log.WarnLevel.String(),
|
||||
BlockInbound: true,
|
||||
WireguardPort: &n.clientCfg.WGPort,
|
||||
WireguardPort: &wgPort,
|
||||
PreSharedKey: n.clientCfg.PreSharedKey,
|
||||
})
|
||||
if err != nil {
|
||||
@@ -257,7 +276,7 @@ func (n *NetBird) createClientEntry(ctx context.Context, accountID types.Account
|
||||
// the client's HTTPClient to avoid issues with request validation that do
|
||||
// not work with reverse proxied requests.
|
||||
transport := &http.Transport{
|
||||
DialContext: client.DialContext,
|
||||
DialContext: dialWithTimeout(client.DialContext),
|
||||
ForceAttemptHTTP2: true,
|
||||
MaxIdleConns: n.transportCfg.maxIdleConns,
|
||||
MaxIdleConnsPerHost: n.transportCfg.maxIdleConnsPerHost,
|
||||
@@ -276,7 +295,7 @@ func (n *NetBird) createClientEntry(ctx context.Context, accountID types.Account
|
||||
|
||||
return &clientEntry{
|
||||
client: client,
|
||||
domains: map[domain.Domain]domainInfo{d: {serviceID: serviceID}},
|
||||
services: map[ServiceKey]serviceInfo{key: si},
|
||||
transport: transport,
|
||||
insecureTransport: insecureTransport,
|
||||
createdAt: time.Now(),
|
||||
@@ -286,7 +305,7 @@ func (n *NetBird) createClientEntry(ctx context.Context, accountID types.Account
|
||||
}, nil
|
||||
}
|
||||
|
||||
// runClientStartup starts the client and notifies registered domains on success.
|
||||
// runClientStartup starts the client and notifies registered services on success.
|
||||
func (n *NetBird) runClientStartup(ctx context.Context, accountID types.AccountID, client *embed.Client) {
|
||||
startCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
@@ -300,16 +319,16 @@ func (n *NetBird) runClientStartup(ctx context.Context, accountID types.AccountI
|
||||
return
|
||||
}
|
||||
|
||||
// Mark client as started and collect domains to notify outside the lock.
|
||||
// Mark client as started and collect services to notify outside the lock.
|
||||
n.clientsMux.Lock()
|
||||
entry, exists := n.clients[accountID]
|
||||
if exists {
|
||||
entry.started = true
|
||||
}
|
||||
var domainsToNotify []domainNotification
|
||||
var toNotify []serviceNotification
|
||||
if exists {
|
||||
for dom, info := range entry.domains {
|
||||
domainsToNotify = append(domainsToNotify, domainNotification{domain: dom, serviceID: info.serviceID})
|
||||
for key, info := range entry.services {
|
||||
toNotify = append(toNotify, serviceNotification{key: key, serviceID: info.serviceID})
|
||||
}
|
||||
}
|
||||
n.clientsMux.Unlock()
|
||||
@@ -317,24 +336,24 @@ func (n *NetBird) runClientStartup(ctx context.Context, accountID types.AccountI
|
||||
if n.statusNotifier == nil {
|
||||
return
|
||||
}
|
||||
for _, dn := range domainsToNotify {
|
||||
if err := n.statusNotifier.NotifyStatus(ctx, string(accountID), dn.serviceID, string(dn.domain), true); err != nil {
|
||||
for _, sn := range toNotify {
|
||||
if err := n.statusNotifier.NotifyStatus(ctx, accountID, sn.serviceID, true); err != nil {
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"domain": dn.domain,
|
||||
"account_id": accountID,
|
||||
"service_key": sn.key,
|
||||
}).WithError(err).Warn("failed to notify tunnel connection status")
|
||||
} else {
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"domain": dn.domain,
|
||||
"account_id": accountID,
|
||||
"service_key": sn.key,
|
||||
}).Info("notified management about tunnel connection")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RemovePeer unregisters a domain from an account. The client is only stopped
|
||||
// when no domains are using it anymore.
|
||||
func (n *NetBird) RemovePeer(ctx context.Context, accountID types.AccountID, d domain.Domain) error {
|
||||
// RemovePeer unregisters a service from an account. The client is only stopped
|
||||
// when no services are using it anymore.
|
||||
func (n *NetBird) RemovePeer(ctx context.Context, accountID types.AccountID, key ServiceKey) error {
|
||||
n.clientsMux.Lock()
|
||||
|
||||
entry, exists := n.clients[accountID]
|
||||
@@ -344,74 +363,65 @@ func (n *NetBird) RemovePeer(ctx context.Context, accountID types.AccountID, d d
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get domain info before deleting
|
||||
domInfo, domainExists := entry.domains[d]
|
||||
if !domainExists {
|
||||
si, svcExists := entry.services[key]
|
||||
if !svcExists {
|
||||
n.clientsMux.Unlock()
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"domain": d,
|
||||
}).Debug("remove peer: domain not registered")
|
||||
"account_id": accountID,
|
||||
"service_key": key,
|
||||
}).Debug("remove peer: service not registered")
|
||||
return nil
|
||||
}
|
||||
|
||||
delete(entry.domains, d)
|
||||
|
||||
// If there are still domains using this client, keep it running
|
||||
if len(entry.domains) > 0 {
|
||||
n.clientsMux.Unlock()
|
||||
delete(entry.services, key)
|
||||
|
||||
stopClient := len(entry.services) == 0
|
||||
var client *embed.Client
|
||||
var transport, insecureTransport *http.Transport
|
||||
if stopClient {
|
||||
n.logger.WithField("account_id", accountID).Info("stopping client, no more services")
|
||||
client = entry.client
|
||||
transport = entry.transport
|
||||
insecureTransport = entry.insecureTransport
|
||||
delete(n.clients, accountID)
|
||||
} else {
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"domain": d,
|
||||
"remaining_domains": len(entry.domains),
|
||||
}).Debug("unregistered domain, client still in use")
|
||||
|
||||
// Notify this domain as disconnected
|
||||
if n.statusNotifier != nil {
|
||||
if err := n.statusNotifier.NotifyStatus(ctx, string(accountID), domInfo.serviceID, string(d), false); err != nil {
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"domain": d,
|
||||
}).WithError(err).Warn("failed to notify tunnel disconnection status")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
"account_id": accountID,
|
||||
"service_key": key,
|
||||
"remaining_services": len(entry.services),
|
||||
}).Debug("unregistered service, client still in use")
|
||||
}
|
||||
|
||||
// No more domains using this client, stop it
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
}).Info("stopping client, no more domains")
|
||||
|
||||
client := entry.client
|
||||
transport := entry.transport
|
||||
insecureTransport := entry.insecureTransport
|
||||
delete(n.clients, accountID)
|
||||
n.clientsMux.Unlock()
|
||||
|
||||
// Notify disconnection before stopping
|
||||
if n.statusNotifier != nil {
|
||||
if err := n.statusNotifier.NotifyStatus(ctx, string(accountID), domInfo.serviceID, string(d), false); err != nil {
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"domain": d,
|
||||
}).WithError(err).Warn("failed to notify tunnel disconnection status")
|
||||
n.notifyDisconnect(ctx, accountID, key, si.serviceID)
|
||||
|
||||
if stopClient {
|
||||
transport.CloseIdleConnections()
|
||||
insecureTransport.CloseIdleConnections()
|
||||
if err := client.Stop(ctx); err != nil {
|
||||
n.logger.WithField("account_id", accountID).WithError(err).Warn("failed to stop netbird client")
|
||||
}
|
||||
}
|
||||
|
||||
transport.CloseIdleConnections()
|
||||
insecureTransport.CloseIdleConnections()
|
||||
|
||||
if err := client.Stop(ctx); err != nil {
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
}).WithError(err).Warn("failed to stop netbird client")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (n *NetBird) notifyDisconnect(ctx context.Context, accountID types.AccountID, key ServiceKey, serviceID types.ServiceID) {
|
||||
if n.statusNotifier == nil {
|
||||
return
|
||||
}
|
||||
if err := n.statusNotifier.NotifyStatus(ctx, accountID, serviceID, false); err != nil {
|
||||
if s, ok := grpcstatus.FromError(err); ok && s.Code() == codes.NotFound {
|
||||
n.logger.WithField("service_key", key).Debug("service already removed, skipping disconnect notification")
|
||||
} else {
|
||||
n.logger.WithFields(log.Fields{
|
||||
"account_id": accountID,
|
||||
"service_key": key,
|
||||
}).WithError(err).Warn("failed to notify tunnel disconnection status")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RoundTrip implements http.RoundTripper. It looks up the client for the account
|
||||
// specified in the request context and uses it to dial the backend.
|
||||
func (n *NetBird) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
@@ -435,7 +445,7 @@ func (n *NetBird) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
}
|
||||
n.clientsMux.RUnlock()
|
||||
|
||||
release, ok := entry.acquireInflight(req.URL.Host)
|
||||
release, ok := entry.acquireInflight(backendKey(req.URL.Host))
|
||||
defer release()
|
||||
if !ok {
|
||||
return nil, ErrTooManyInflight
|
||||
@@ -496,16 +506,16 @@ func (n *NetBird) HasClient(accountID types.AccountID) bool {
|
||||
return exists
|
||||
}
|
||||
|
||||
// DomainCount returns the number of domains registered for the given account.
|
||||
// ServiceCount returns the number of services registered for the given account.
|
||||
// Returns 0 if the account has no client.
|
||||
func (n *NetBird) DomainCount(accountID types.AccountID) int {
|
||||
func (n *NetBird) ServiceCount(accountID types.AccountID) int {
|
||||
n.clientsMux.RLock()
|
||||
defer n.clientsMux.RUnlock()
|
||||
entry, exists := n.clients[accountID]
|
||||
if !exists {
|
||||
return 0
|
||||
}
|
||||
return len(entry.domains)
|
||||
return len(entry.services)
|
||||
}
|
||||
|
||||
// ClientCount returns the total number of active clients.
|
||||
@@ -533,16 +543,16 @@ func (n *NetBird) ListClientsForDebug() map[types.AccountID]ClientDebugInfo {
|
||||
|
||||
result := make(map[types.AccountID]ClientDebugInfo)
|
||||
for accountID, entry := range n.clients {
|
||||
domains := make(domain.List, 0, len(entry.domains))
|
||||
for d := range entry.domains {
|
||||
domains = append(domains, d)
|
||||
keys := make([]string, 0, len(entry.services))
|
||||
for k := range entry.services {
|
||||
keys = append(keys, string(k))
|
||||
}
|
||||
result[accountID] = ClientDebugInfo{
|
||||
AccountID: accountID,
|
||||
DomainCount: len(entry.domains),
|
||||
Domains: domains,
|
||||
HasClient: entry.client != nil,
|
||||
CreatedAt: entry.createdAt,
|
||||
AccountID: accountID,
|
||||
ServiceCount: len(entry.services),
|
||||
ServiceKeys: keys,
|
||||
HasClient: entry.client != nil,
|
||||
CreatedAt: entry.createdAt,
|
||||
}
|
||||
}
|
||||
return result
|
||||
@@ -581,6 +591,20 @@ func NewNetBird(proxyID, proxyAddr string, clientCfg ClientConfig, logger *log.L
|
||||
}
|
||||
}
|
||||
|
||||
// dialWithTimeout wraps a DialContext function so that any dial timeout
|
||||
// stored in the context (via types.WithDialTimeout) is applied only to
|
||||
// the connection establishment phase, not the full request lifetime.
|
||||
func dialWithTimeout(dial func(ctx context.Context, network, addr string) (net.Conn, error)) func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
return func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
if d, ok := types.DialTimeoutFromContext(ctx); ok {
|
||||
var cancel context.CancelFunc
|
||||
ctx, cancel = context.WithTimeout(ctx, d)
|
||||
defer cancel()
|
||||
}
|
||||
return dial(ctx, network, addr)
|
||||
}
|
||||
}
|
||||
|
||||
// WithAccountID adds the account ID to the context.
|
||||
func WithAccountID(ctx context.Context, accountID types.AccountID) context.Context {
|
||||
return context.WithValue(ctx, accountIDContextKey{}, accountID)
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package roundtrip
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"math/big"
|
||||
"sync"
|
||||
@@ -8,7 +9,6 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
)
|
||||
|
||||
// Simple benchmark for comparison with AddPeer contention.
|
||||
@@ -29,9 +29,9 @@ func BenchmarkHasClient(b *testing.B) {
|
||||
target = id
|
||||
}
|
||||
nb.clients[id] = &clientEntry{
|
||||
domains: map[domain.Domain]domainInfo{
|
||||
domain.Domain(rand.Text()): {
|
||||
serviceID: rand.Text(),
|
||||
services: map[ServiceKey]serviceInfo{
|
||||
ServiceKey(rand.Text()): {
|
||||
serviceID: types.ServiceID(rand.Text()),
|
||||
},
|
||||
},
|
||||
createdAt: time.Now(),
|
||||
@@ -70,9 +70,9 @@ func BenchmarkHasClientDuringAddPeer(b *testing.B) {
|
||||
target = id
|
||||
}
|
||||
nb.clients[id] = &clientEntry{
|
||||
domains: map[domain.Domain]domainInfo{
|
||||
domain.Domain(rand.Text()): {
|
||||
serviceID: rand.Text(),
|
||||
services: map[ServiceKey]serviceInfo{
|
||||
ServiceKey(rand.Text()): {
|
||||
serviceID: types.ServiceID(rand.Text()),
|
||||
},
|
||||
},
|
||||
createdAt: time.Now(),
|
||||
@@ -81,19 +81,22 @@ func BenchmarkHasClientDuringAddPeer(b *testing.B) {
|
||||
}
|
||||
|
||||
// Launch workers that continuously call AddPeer with new random accountIDs.
|
||||
ctx, cancel := context.WithCancel(b.Context())
|
||||
var wg sync.WaitGroup
|
||||
for range addPeerWorkers {
|
||||
wg.Go(func() {
|
||||
for {
|
||||
if err := nb.AddPeer(b.Context(),
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for ctx.Err() == nil {
|
||||
if err := nb.AddPeer(ctx,
|
||||
types.AccountID(rand.Text()),
|
||||
domain.Domain(rand.Text()),
|
||||
ServiceKey(rand.Text()),
|
||||
rand.Text(),
|
||||
rand.Text()); err != nil {
|
||||
b.Log(err)
|
||||
types.ServiceID(rand.Text())); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
})
|
||||
}()
|
||||
}
|
||||
|
||||
// Benchmark calling HasClient during AddPeer contention.
|
||||
@@ -104,4 +107,6 @@ func BenchmarkHasClientDuringAddPeer(b *testing.B) {
|
||||
}
|
||||
})
|
||||
b.StopTimer()
|
||||
cancel()
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
@@ -11,7 +11,6 @@ import (
|
||||
"google.golang.org/grpc"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
@@ -27,16 +26,15 @@ type mockStatusNotifier struct {
|
||||
}
|
||||
|
||||
type statusCall struct {
|
||||
accountID string
|
||||
serviceID string
|
||||
domain string
|
||||
accountID types.AccountID
|
||||
serviceID types.ServiceID
|
||||
connected bool
|
||||
}
|
||||
|
||||
func (m *mockStatusNotifier) NotifyStatus(_ context.Context, accountID, serviceID, domain string, connected bool) error {
|
||||
func (m *mockStatusNotifier) NotifyStatus(_ context.Context, accountID types.AccountID, serviceID types.ServiceID, connected bool) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.statuses = append(m.statuses, statusCall{accountID, serviceID, domain, connected})
|
||||
m.statuses = append(m.statuses, statusCall{accountID, serviceID, connected})
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -62,36 +60,34 @@ func TestNetBird_AddPeer_CreatesClientForNewAccount(t *testing.T) {
|
||||
|
||||
// Initially no client exists.
|
||||
assert.False(t, nb.HasClient(accountID), "should not have client before AddPeer")
|
||||
assert.Equal(t, 0, nb.DomainCount(accountID), "domain count should be 0")
|
||||
assert.Equal(t, 0, nb.ServiceCount(accountID), "service count should be 0")
|
||||
|
||||
// Add first domain - this should create a new client.
|
||||
// Note: This will fail to actually connect since we use an invalid URL,
|
||||
// but the client entry should still be created.
|
||||
err := nb.AddPeer(context.Background(), accountID, domain.Domain("domain1.test"), "setup-key-1", "proxy-1")
|
||||
// Add first service - this should create a new client.
|
||||
err := nb.AddPeer(context.Background(), accountID, "domain1.test", "setup-key-1", types.ServiceID("proxy-1"))
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.True(t, nb.HasClient(accountID), "should have client after AddPeer")
|
||||
assert.Equal(t, 1, nb.DomainCount(accountID), "domain count should be 1")
|
||||
assert.Equal(t, 1, nb.ServiceCount(accountID), "service count should be 1")
|
||||
}
|
||||
|
||||
func TestNetBird_AddPeer_ReuseClientForSameAccount(t *testing.T) {
|
||||
nb := mockNetBird()
|
||||
accountID := types.AccountID("account-1")
|
||||
|
||||
// Add first domain.
|
||||
err := nb.AddPeer(context.Background(), accountID, domain.Domain("domain1.test"), "setup-key-1", "proxy-1")
|
||||
// Add first service.
|
||||
err := nb.AddPeer(context.Background(), accountID, "domain1.test", "setup-key-1", types.ServiceID("proxy-1"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 1, nb.DomainCount(accountID))
|
||||
assert.Equal(t, 1, nb.ServiceCount(accountID))
|
||||
|
||||
// Add second domain for the same account - should reuse existing client.
|
||||
err = nb.AddPeer(context.Background(), accountID, domain.Domain("domain2.test"), "setup-key-1", "proxy-2")
|
||||
// Add second service for the same account - should reuse existing client.
|
||||
err = nb.AddPeer(context.Background(), accountID, "domain2.test", "setup-key-1", types.ServiceID("proxy-2"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 2, nb.DomainCount(accountID), "domain count should be 2 after adding second domain")
|
||||
assert.Equal(t, 2, nb.ServiceCount(accountID), "service count should be 2 after adding second service")
|
||||
|
||||
// Add third domain.
|
||||
err = nb.AddPeer(context.Background(), accountID, domain.Domain("domain3.test"), "setup-key-1", "proxy-3")
|
||||
// Add third service.
|
||||
err = nb.AddPeer(context.Background(), accountID, "domain3.test", "setup-key-1", types.ServiceID("proxy-3"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 3, nb.DomainCount(accountID), "domain count should be 3 after adding third domain")
|
||||
assert.Equal(t, 3, nb.ServiceCount(accountID), "service count should be 3 after adding third service")
|
||||
|
||||
// Still only one client.
|
||||
assert.True(t, nb.HasClient(accountID))
|
||||
@@ -102,64 +98,62 @@ func TestNetBird_AddPeer_SeparateClientsForDifferentAccounts(t *testing.T) {
|
||||
account1 := types.AccountID("account-1")
|
||||
account2 := types.AccountID("account-2")
|
||||
|
||||
// Add domain for account 1.
|
||||
err := nb.AddPeer(context.Background(), account1, domain.Domain("domain1.test"), "setup-key-1", "proxy-1")
|
||||
// Add service for account 1.
|
||||
err := nb.AddPeer(context.Background(), account1, "domain1.test", "setup-key-1", types.ServiceID("proxy-1"))
|
||||
require.NoError(t, err)
|
||||
|
||||
// Add domain for account 2.
|
||||
err = nb.AddPeer(context.Background(), account2, domain.Domain("domain2.test"), "setup-key-2", "proxy-2")
|
||||
// Add service for account 2.
|
||||
err = nb.AddPeer(context.Background(), account2, "domain2.test", "setup-key-2", types.ServiceID("proxy-2"))
|
||||
require.NoError(t, err)
|
||||
|
||||
// Both accounts should have their own clients.
|
||||
assert.True(t, nb.HasClient(account1), "account1 should have client")
|
||||
assert.True(t, nb.HasClient(account2), "account2 should have client")
|
||||
assert.Equal(t, 1, nb.DomainCount(account1), "account1 domain count should be 1")
|
||||
assert.Equal(t, 1, nb.DomainCount(account2), "account2 domain count should be 1")
|
||||
assert.Equal(t, 1, nb.ServiceCount(account1), "account1 service count should be 1")
|
||||
assert.Equal(t, 1, nb.ServiceCount(account2), "account2 service count should be 1")
|
||||
}
|
||||
|
||||
func TestNetBird_RemovePeer_KeepsClientWhenDomainsRemain(t *testing.T) {
|
||||
func TestNetBird_RemovePeer_KeepsClientWhenServicesRemain(t *testing.T) {
|
||||
nb := mockNetBird()
|
||||
accountID := types.AccountID("account-1")
|
||||
|
||||
// Add multiple domains.
|
||||
err := nb.AddPeer(context.Background(), accountID, domain.Domain("domain1.test"), "setup-key-1", "proxy-1")
|
||||
// Add multiple services.
|
||||
err := nb.AddPeer(context.Background(), accountID, "domain1.test", "setup-key-1", types.ServiceID("proxy-1"))
|
||||
require.NoError(t, err)
|
||||
err = nb.AddPeer(context.Background(), accountID, domain.Domain("domain2.test"), "setup-key-1", "proxy-2")
|
||||
err = nb.AddPeer(context.Background(), accountID, "domain2.test", "setup-key-1", types.ServiceID("proxy-2"))
|
||||
require.NoError(t, err)
|
||||
err = nb.AddPeer(context.Background(), accountID, domain.Domain("domain3.test"), "setup-key-1", "proxy-3")
|
||||
err = nb.AddPeer(context.Background(), accountID, "domain3.test", "setup-key-1", types.ServiceID("proxy-3"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 3, nb.DomainCount(accountID))
|
||||
assert.Equal(t, 3, nb.ServiceCount(accountID))
|
||||
|
||||
// Remove one domain - client should remain.
|
||||
// Remove one service - client should remain.
|
||||
err = nb.RemovePeer(context.Background(), accountID, "domain1.test")
|
||||
require.NoError(t, err)
|
||||
assert.True(t, nb.HasClient(accountID), "client should remain after removing one domain")
|
||||
assert.Equal(t, 2, nb.DomainCount(accountID), "domain count should be 2")
|
||||
assert.True(t, nb.HasClient(accountID), "client should remain after removing one service")
|
||||
assert.Equal(t, 2, nb.ServiceCount(accountID), "service count should be 2")
|
||||
|
||||
// Remove another domain - client should still remain.
|
||||
// Remove another service - client should still remain.
|
||||
err = nb.RemovePeer(context.Background(), accountID, "domain2.test")
|
||||
require.NoError(t, err)
|
||||
assert.True(t, nb.HasClient(accountID), "client should remain after removing second domain")
|
||||
assert.Equal(t, 1, nb.DomainCount(accountID), "domain count should be 1")
|
||||
assert.True(t, nb.HasClient(accountID), "client should remain after removing second service")
|
||||
assert.Equal(t, 1, nb.ServiceCount(accountID), "service count should be 1")
|
||||
}
|
||||
|
||||
func TestNetBird_RemovePeer_RemovesClientWhenLastDomainRemoved(t *testing.T) {
|
||||
func TestNetBird_RemovePeer_RemovesClientWhenLastServiceRemoved(t *testing.T) {
|
||||
nb := mockNetBird()
|
||||
accountID := types.AccountID("account-1")
|
||||
|
||||
// Add single domain.
|
||||
err := nb.AddPeer(context.Background(), accountID, domain.Domain("domain1.test"), "setup-key-1", "proxy-1")
|
||||
// Add single service.
|
||||
err := nb.AddPeer(context.Background(), accountID, "domain1.test", "setup-key-1", types.ServiceID("proxy-1"))
|
||||
require.NoError(t, err)
|
||||
assert.True(t, nb.HasClient(accountID))
|
||||
|
||||
// Remove the only domain - client should be removed.
|
||||
// Note: Stop() may fail since the client never actually connected,
|
||||
// but the entry should still be removed from the map.
|
||||
// Remove the only service - client should be removed.
|
||||
_ = nb.RemovePeer(context.Background(), accountID, "domain1.test")
|
||||
|
||||
// After removing all domains, client should be gone.
|
||||
assert.False(t, nb.HasClient(accountID), "client should be removed after removing last domain")
|
||||
assert.Equal(t, 0, nb.DomainCount(accountID), "domain count should be 0")
|
||||
// After removing all services, client should be gone.
|
||||
assert.False(t, nb.HasClient(accountID), "client should be removed after removing last service")
|
||||
assert.Equal(t, 0, nb.ServiceCount(accountID), "service count should be 0")
|
||||
}
|
||||
|
||||
func TestNetBird_RemovePeer_NonExistentAccountIsNoop(t *testing.T) {
|
||||
@@ -171,21 +165,21 @@ func TestNetBird_RemovePeer_NonExistentAccountIsNoop(t *testing.T) {
|
||||
assert.NoError(t, err, "removing from non-existent account should not error")
|
||||
}
|
||||
|
||||
func TestNetBird_RemovePeer_NonExistentDomainIsNoop(t *testing.T) {
|
||||
func TestNetBird_RemovePeer_NonExistentServiceIsNoop(t *testing.T) {
|
||||
nb := mockNetBird()
|
||||
accountID := types.AccountID("account-1")
|
||||
|
||||
// Add one domain.
|
||||
err := nb.AddPeer(context.Background(), accountID, domain.Domain("domain1.test"), "setup-key-1", "proxy-1")
|
||||
// Add one service.
|
||||
err := nb.AddPeer(context.Background(), accountID, "domain1.test", "setup-key-1", types.ServiceID("proxy-1"))
|
||||
require.NoError(t, err)
|
||||
|
||||
// Remove non-existent domain - should not affect existing domain.
|
||||
err = nb.RemovePeer(context.Background(), accountID, domain.Domain("nonexistent.test"))
|
||||
// Remove non-existent service - should not affect existing service.
|
||||
err = nb.RemovePeer(context.Background(), accountID, "nonexistent.test")
|
||||
require.NoError(t, err)
|
||||
|
||||
// Original domain should still be registered.
|
||||
// Original service should still be registered.
|
||||
assert.True(t, nb.HasClient(accountID))
|
||||
assert.Equal(t, 1, nb.DomainCount(accountID), "original domain should remain")
|
||||
assert.Equal(t, 1, nb.ServiceCount(accountID), "original service should remain")
|
||||
}
|
||||
|
||||
func TestWithAccountID_AndAccountIDFromContext(t *testing.T) {
|
||||
@@ -216,19 +210,17 @@ func TestNetBird_StopAll_StopsAllClients(t *testing.T) {
|
||||
account2 := types.AccountID("account-2")
|
||||
account3 := types.AccountID("account-3")
|
||||
|
||||
// Add domains for multiple accounts.
|
||||
err := nb.AddPeer(context.Background(), account1, domain.Domain("domain1.test"), "key-1", "proxy-1")
|
||||
// Add services for multiple accounts.
|
||||
err := nb.AddPeer(context.Background(), account1, "domain1.test", "key-1", types.ServiceID("proxy-1"))
|
||||
require.NoError(t, err)
|
||||
err = nb.AddPeer(context.Background(), account2, domain.Domain("domain2.test"), "key-2", "proxy-2")
|
||||
err = nb.AddPeer(context.Background(), account2, "domain2.test", "key-2", types.ServiceID("proxy-2"))
|
||||
require.NoError(t, err)
|
||||
err = nb.AddPeer(context.Background(), account3, domain.Domain("domain3.test"), "key-3", "proxy-3")
|
||||
err = nb.AddPeer(context.Background(), account3, "domain3.test", "key-3", types.ServiceID("proxy-3"))
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, 3, nb.ClientCount(), "should have 3 clients")
|
||||
|
||||
// Stop all clients.
|
||||
// Note: StopAll may return errors since clients never actually connected,
|
||||
// but the clients should still be removed from the map.
|
||||
_ = nb.StopAll(context.Background())
|
||||
|
||||
assert.Equal(t, 0, nb.ClientCount(), "should have 0 clients after StopAll")
|
||||
@@ -243,18 +235,18 @@ func TestNetBird_ClientCount(t *testing.T) {
|
||||
assert.Equal(t, 0, nb.ClientCount(), "should start with 0 clients")
|
||||
|
||||
// Add clients for different accounts.
|
||||
err := nb.AddPeer(context.Background(), types.AccountID("account-1"), domain.Domain("domain1.test"), "key-1", "proxy-1")
|
||||
err := nb.AddPeer(context.Background(), types.AccountID("account-1"), "domain1.test", "key-1", types.ServiceID("proxy-1"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 1, nb.ClientCount())
|
||||
|
||||
err = nb.AddPeer(context.Background(), types.AccountID("account-2"), domain.Domain("domain2.test"), "key-2", "proxy-2")
|
||||
err = nb.AddPeer(context.Background(), types.AccountID("account-2"), "domain2.test", "key-2", types.ServiceID("proxy-2"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 2, nb.ClientCount())
|
||||
|
||||
// Adding domain to existing account should not increase count.
|
||||
err = nb.AddPeer(context.Background(), types.AccountID("account-1"), domain.Domain("domain1b.test"), "key-1", "proxy-1b")
|
||||
// Adding service to existing account should not increase count.
|
||||
err = nb.AddPeer(context.Background(), types.AccountID("account-1"), "domain1b.test", "key-1", types.ServiceID("proxy-1b"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 2, nb.ClientCount(), "adding domain to existing account should not increase client count")
|
||||
assert.Equal(t, 2, nb.ClientCount(), "adding service to existing account should not increase client count")
|
||||
}
|
||||
|
||||
func TestNetBird_RoundTrip_RequiresAccountIDInContext(t *testing.T) {
|
||||
@@ -293,8 +285,8 @@ func TestNetBird_AddPeer_ExistingStartedClient_NotifiesStatus(t *testing.T) {
|
||||
}, nil, notifier, &mockMgmtClient{})
|
||||
accountID := types.AccountID("account-1")
|
||||
|
||||
// Add first domain — creates a new client entry.
|
||||
err := nb.AddPeer(context.Background(), accountID, domain.Domain("domain1.test"), "key-1", "svc-1")
|
||||
// Add first service — creates a new client entry.
|
||||
err := nb.AddPeer(context.Background(), accountID, "domain1.test", "key-1", types.ServiceID("svc-1"))
|
||||
require.NoError(t, err)
|
||||
|
||||
// Manually mark client as started to simulate background startup completing.
|
||||
@@ -302,15 +294,14 @@ func TestNetBird_AddPeer_ExistingStartedClient_NotifiesStatus(t *testing.T) {
|
||||
nb.clients[accountID].started = true
|
||||
nb.clientsMux.Unlock()
|
||||
|
||||
// Add second domain — should notify immediately since client is already started.
|
||||
err = nb.AddPeer(context.Background(), accountID, domain.Domain("domain2.test"), "key-1", "svc-2")
|
||||
// Add second service — should notify immediately since client is already started.
|
||||
err = nb.AddPeer(context.Background(), accountID, "domain2.test", "key-1", types.ServiceID("svc-2"))
|
||||
require.NoError(t, err)
|
||||
|
||||
calls := notifier.calls()
|
||||
require.Len(t, calls, 1)
|
||||
assert.Equal(t, string(accountID), calls[0].accountID)
|
||||
assert.Equal(t, "svc-2", calls[0].serviceID)
|
||||
assert.Equal(t, "domain2.test", calls[0].domain)
|
||||
assert.Equal(t, accountID, calls[0].accountID)
|
||||
assert.Equal(t, types.ServiceID("svc-2"), calls[0].serviceID)
|
||||
assert.True(t, calls[0].connected)
|
||||
}
|
||||
|
||||
@@ -323,18 +314,18 @@ func TestNetBird_RemovePeer_NotifiesDisconnection(t *testing.T) {
|
||||
}, nil, notifier, &mockMgmtClient{})
|
||||
accountID := types.AccountID("account-1")
|
||||
|
||||
err := nb.AddPeer(context.Background(), accountID, domain.Domain("domain1.test"), "key-1", "svc-1")
|
||||
err := nb.AddPeer(context.Background(), accountID, "domain1.test", "key-1", types.ServiceID("svc-1"))
|
||||
require.NoError(t, err)
|
||||
err = nb.AddPeer(context.Background(), accountID, domain.Domain("domain2.test"), "key-1", "svc-2")
|
||||
err = nb.AddPeer(context.Background(), accountID, "domain2.test", "key-1", types.ServiceID("svc-2"))
|
||||
require.NoError(t, err)
|
||||
|
||||
// Remove one domain — client stays, but disconnection notification fires.
|
||||
// Remove one service — client stays, but disconnection notification fires.
|
||||
err = nb.RemovePeer(context.Background(), accountID, "domain1.test")
|
||||
require.NoError(t, err)
|
||||
assert.True(t, nb.HasClient(accountID))
|
||||
|
||||
calls := notifier.calls()
|
||||
require.Len(t, calls, 1)
|
||||
assert.Equal(t, "domain1.test", calls[0].domain)
|
||||
assert.Equal(t, types.ServiceID("svc-1"), calls[0].serviceID)
|
||||
assert.False(t, calls[0].connected)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
package tcp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/tls"
|
||||
"io"
|
||||
"net"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// BenchmarkPeekClientHello_TLS measures the overhead of peeking at a real
|
||||
// TLS ClientHello and extracting the SNI. This is the per-connection cost
|
||||
// added to every TLS connection on the main listener.
|
||||
func BenchmarkPeekClientHello_TLS(b *testing.B) {
|
||||
// Pre-generate a ClientHello by capturing what crypto/tls sends.
|
||||
clientConn, serverConn := net.Pipe()
|
||||
go func() {
|
||||
tlsConn := tls.Client(clientConn, &tls.Config{
|
||||
ServerName: "app.example.com",
|
||||
InsecureSkipVerify: true, //nolint:gosec
|
||||
})
|
||||
_ = tlsConn.Handshake()
|
||||
}()
|
||||
|
||||
var hello []byte
|
||||
buf := make([]byte, 16384)
|
||||
n, _ := serverConn.Read(buf)
|
||||
hello = make([]byte, n)
|
||||
copy(hello, buf[:n])
|
||||
clientConn.Close()
|
||||
serverConn.Close()
|
||||
|
||||
b.ResetTimer()
|
||||
b.ReportAllocs()
|
||||
|
||||
for b.Loop() {
|
||||
r := bytes.NewReader(hello)
|
||||
conn := &readerConn{Reader: r}
|
||||
sni, wrapped, err := PeekClientHello(conn)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if sni != "app.example.com" {
|
||||
b.Fatalf("unexpected SNI: %q", sni)
|
||||
}
|
||||
// Simulate draining the peeked bytes (what the HTTP server would do).
|
||||
_, _ = io.Copy(io.Discard, wrapped)
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkPeekClientHello_NonTLS measures peek overhead for non-TLS
|
||||
// connections that hit the fast non-handshake exit path.
|
||||
func BenchmarkPeekClientHello_NonTLS(b *testing.B) {
|
||||
httpReq := []byte("GET / HTTP/1.1\r\nHost: example.com\r\n\r\n")
|
||||
|
||||
b.ResetTimer()
|
||||
b.ReportAllocs()
|
||||
|
||||
for b.Loop() {
|
||||
r := bytes.NewReader(httpReq)
|
||||
conn := &readerConn{Reader: r}
|
||||
_, wrapped, err := PeekClientHello(conn)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
_, _ = io.Copy(io.Discard, wrapped)
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkPeekedConn_Read measures the read overhead of the peekedConn
|
||||
// wrapper compared to a plain connection read. The peeked bytes use
|
||||
// io.MultiReader which adds one indirection per Read call.
|
||||
func BenchmarkPeekedConn_Read(b *testing.B) {
|
||||
data := make([]byte, 4096)
|
||||
peeked := make([]byte, 512)
|
||||
buf := make([]byte, 1024)
|
||||
|
||||
b.ResetTimer()
|
||||
b.ReportAllocs()
|
||||
|
||||
for b.Loop() {
|
||||
r := bytes.NewReader(data)
|
||||
conn := &readerConn{Reader: r}
|
||||
pc := newPeekedConn(conn, peeked)
|
||||
for {
|
||||
_, err := pc.Read(buf)
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BenchmarkExtractSNI measures just the in-memory SNI parsing cost,
|
||||
// excluding I/O.
|
||||
func BenchmarkExtractSNI(b *testing.B) {
|
||||
clientConn, serverConn := net.Pipe()
|
||||
go func() {
|
||||
tlsConn := tls.Client(clientConn, &tls.Config{
|
||||
ServerName: "app.example.com",
|
||||
InsecureSkipVerify: true, //nolint:gosec
|
||||
})
|
||||
_ = tlsConn.Handshake()
|
||||
}()
|
||||
|
||||
buf := make([]byte, 16384)
|
||||
n, _ := serverConn.Read(buf)
|
||||
payload := make([]byte, n-tlsRecordHeaderLen)
|
||||
copy(payload, buf[tlsRecordHeaderLen:n])
|
||||
clientConn.Close()
|
||||
serverConn.Close()
|
||||
|
||||
b.ResetTimer()
|
||||
b.ReportAllocs()
|
||||
|
||||
for b.Loop() {
|
||||
sni := extractSNI(payload)
|
||||
if sni != "app.example.com" {
|
||||
b.Fatalf("unexpected SNI: %q", sni)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// readerConn wraps an io.Reader as a net.Conn for benchmarking.
|
||||
// Only Read is functional; all other methods are no-ops.
|
||||
type readerConn struct {
|
||||
io.Reader
|
||||
net.Conn
|
||||
}
|
||||
|
||||
func (c *readerConn) Read(b []byte) (int, error) {
|
||||
return c.Reader.Read(b)
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package tcp
|
||||
|
||||
import (
|
||||
"net"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// chanListener implements net.Listener by reading connections from a channel.
|
||||
// It allows the SNI router to feed HTTP connections to http.Server.ServeTLS.
|
||||
type chanListener struct {
|
||||
ch chan net.Conn
|
||||
addr net.Addr
|
||||
once sync.Once
|
||||
closed chan struct{}
|
||||
}
|
||||
|
||||
func newChanListener(ch chan net.Conn, addr net.Addr) *chanListener {
|
||||
return &chanListener{
|
||||
ch: ch,
|
||||
addr: addr,
|
||||
closed: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Accept waits for and returns the next connection from the channel.
|
||||
func (l *chanListener) Accept() (net.Conn, error) {
|
||||
for {
|
||||
select {
|
||||
case conn, ok := <-l.ch:
|
||||
if !ok {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
return conn, nil
|
||||
case <-l.closed:
|
||||
// Drain buffered connections before returning.
|
||||
for {
|
||||
select {
|
||||
case conn, ok := <-l.ch:
|
||||
if !ok {
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
_ = conn.Close()
|
||||
default:
|
||||
return nil, net.ErrClosed
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Close signals the listener to stop accepting connections and drains
|
||||
// any buffered connections that have not yet been accepted.
|
||||
func (l *chanListener) Close() error {
|
||||
l.once.Do(func() {
|
||||
close(l.closed)
|
||||
for {
|
||||
select {
|
||||
case conn, ok := <-l.ch:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
_ = conn.Close()
|
||||
default:
|
||||
return
|
||||
}
|
||||
}
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
// Addr returns the listener's network address.
|
||||
func (l *chanListener) Addr() net.Addr {
|
||||
return l.addr
|
||||
}
|
||||
|
||||
var _ net.Listener = (*chanListener)(nil)
|
||||
@@ -0,0 +1,39 @@
|
||||
package tcp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"net"
|
||||
)
|
||||
|
||||
// peekedConn wraps a net.Conn and prepends previously peeked bytes
|
||||
// so that readers see the full original stream transparently.
|
||||
type peekedConn struct {
|
||||
net.Conn
|
||||
reader io.Reader
|
||||
}
|
||||
|
||||
func newPeekedConn(conn net.Conn, peeked []byte) *peekedConn {
|
||||
return &peekedConn{
|
||||
Conn: conn,
|
||||
reader: io.MultiReader(bytes.NewReader(peeked), conn),
|
||||
}
|
||||
}
|
||||
|
||||
// Read replays the peeked bytes first, then reads from the underlying conn.
|
||||
func (c *peekedConn) Read(b []byte) (int, error) {
|
||||
return c.reader.Read(b)
|
||||
}
|
||||
|
||||
// CloseWrite delegates to the underlying connection if it supports
|
||||
// half-close (e.g. *net.TCPConn). Without this, embedding net.Conn
|
||||
// as an interface hides the concrete type's CloseWrite method, making
|
||||
// half-close a silent no-op for all SNI-routed connections.
|
||||
func (c *peekedConn) CloseWrite() error {
|
||||
if hc, ok := c.Conn.(halfCloser); ok {
|
||||
return hc.CloseWrite()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
var _ halfCloser = (*peekedConn)(nil)
|
||||
@@ -0,0 +1,29 @@
|
||||
package tcp
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
|
||||
"github.com/pires/go-proxyproto"
|
||||
)
|
||||
|
||||
// writeProxyProtoV2 sends a PROXY protocol v2 header to the backend connection,
|
||||
// conveying the real client address.
|
||||
func writeProxyProtoV2(client, backend net.Conn) error {
|
||||
tp := proxyproto.TCPv4
|
||||
if addr, ok := client.RemoteAddr().(*net.TCPAddr); ok && addr.IP.To4() == nil {
|
||||
tp = proxyproto.TCPv6
|
||||
}
|
||||
|
||||
header := &proxyproto.Header{
|
||||
Version: 2,
|
||||
Command: proxyproto.PROXY,
|
||||
TransportProtocol: tp,
|
||||
SourceAddr: client.RemoteAddr(),
|
||||
DestinationAddr: client.LocalAddr(),
|
||||
}
|
||||
if _, err := header.WriteTo(backend); err != nil {
|
||||
return fmt.Errorf("write PROXY protocol v2 header: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,128 @@
|
||||
package tcp
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/pires/go-proxyproto"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestWriteProxyProtoV2_IPv4(t *testing.T) {
|
||||
// Set up a real TCP listener and dial to get connections with real addresses.
|
||||
ln, err := net.Listen("tcp4", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer ln.Close()
|
||||
|
||||
var serverConn net.Conn
|
||||
accepted := make(chan struct{})
|
||||
go func() {
|
||||
var err error
|
||||
serverConn, err = ln.Accept()
|
||||
if err != nil {
|
||||
t.Error("accept failed:", err)
|
||||
}
|
||||
close(accepted)
|
||||
}()
|
||||
|
||||
clientConn, err := net.Dial("tcp4", ln.Addr().String())
|
||||
require.NoError(t, err)
|
||||
defer clientConn.Close()
|
||||
|
||||
<-accepted
|
||||
defer serverConn.Close()
|
||||
|
||||
// Use a pipe as the backend: write the header to one end, read from the other.
|
||||
backendRead, backendWrite := net.Pipe()
|
||||
defer backendRead.Close()
|
||||
defer backendWrite.Close()
|
||||
|
||||
// serverConn is the "client" arg: RemoteAddr is the source, LocalAddr is the destination.
|
||||
writeDone := make(chan error, 1)
|
||||
go func() {
|
||||
writeDone <- writeProxyProtoV2(serverConn, backendWrite)
|
||||
}()
|
||||
|
||||
// Read the PROXY protocol header from the backend read side.
|
||||
header, err := proxyproto.Read(bufio.NewReader(backendRead))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, header, "should have received a proxy protocol header")
|
||||
|
||||
writeErr := <-writeDone
|
||||
require.NoError(t, writeErr)
|
||||
|
||||
assert.Equal(t, byte(2), header.Version, "version should be 2")
|
||||
assert.Equal(t, proxyproto.PROXY, header.Command, "command should be PROXY")
|
||||
assert.Equal(t, proxyproto.TCPv4, header.TransportProtocol, "transport should be TCPv4")
|
||||
|
||||
// serverConn.RemoteAddr() is the client's address (source in the header).
|
||||
expectedSrc := serverConn.RemoteAddr().(*net.TCPAddr)
|
||||
actualSrc := header.SourceAddr.(*net.TCPAddr)
|
||||
assert.Equal(t, expectedSrc.IP.String(), actualSrc.IP.String(), "source IP should match client remote addr")
|
||||
assert.Equal(t, expectedSrc.Port, actualSrc.Port, "source port should match client remote addr")
|
||||
|
||||
// serverConn.LocalAddr() is the server's address (destination in the header).
|
||||
expectedDst := serverConn.LocalAddr().(*net.TCPAddr)
|
||||
actualDst := header.DestinationAddr.(*net.TCPAddr)
|
||||
assert.Equal(t, expectedDst.IP.String(), actualDst.IP.String(), "destination IP should match server local addr")
|
||||
assert.Equal(t, expectedDst.Port, actualDst.Port, "destination port should match server local addr")
|
||||
}
|
||||
|
||||
func TestWriteProxyProtoV2_IPv6(t *testing.T) {
|
||||
// Set up a real TCP6 listener on loopback.
|
||||
ln, err := net.Listen("tcp6", "[::1]:0")
|
||||
if err != nil {
|
||||
t.Skip("IPv6 not available:", err)
|
||||
}
|
||||
defer ln.Close()
|
||||
|
||||
var serverConn net.Conn
|
||||
accepted := make(chan struct{})
|
||||
go func() {
|
||||
var err error
|
||||
serverConn, err = ln.Accept()
|
||||
if err != nil {
|
||||
t.Error("accept failed:", err)
|
||||
}
|
||||
close(accepted)
|
||||
}()
|
||||
|
||||
clientConn, err := net.Dial("tcp6", ln.Addr().String())
|
||||
require.NoError(t, err)
|
||||
defer clientConn.Close()
|
||||
|
||||
<-accepted
|
||||
defer serverConn.Close()
|
||||
|
||||
backendRead, backendWrite := net.Pipe()
|
||||
defer backendRead.Close()
|
||||
defer backendWrite.Close()
|
||||
|
||||
writeDone := make(chan error, 1)
|
||||
go func() {
|
||||
writeDone <- writeProxyProtoV2(serverConn, backendWrite)
|
||||
}()
|
||||
|
||||
header, err := proxyproto.Read(bufio.NewReader(backendRead))
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, header, "should have received a proxy protocol header")
|
||||
|
||||
writeErr := <-writeDone
|
||||
require.NoError(t, writeErr)
|
||||
|
||||
assert.Equal(t, byte(2), header.Version, "version should be 2")
|
||||
assert.Equal(t, proxyproto.PROXY, header.Command, "command should be PROXY")
|
||||
assert.Equal(t, proxyproto.TCPv6, header.TransportProtocol, "transport should be TCPv6")
|
||||
|
||||
expectedSrc := serverConn.RemoteAddr().(*net.TCPAddr)
|
||||
actualSrc := header.SourceAddr.(*net.TCPAddr)
|
||||
assert.Equal(t, expectedSrc.IP.String(), actualSrc.IP.String(), "source IP should match client remote addr")
|
||||
assert.Equal(t, expectedSrc.Port, actualSrc.Port, "source port should match client remote addr")
|
||||
|
||||
expectedDst := serverConn.LocalAddr().(*net.TCPAddr)
|
||||
actualDst := header.DestinationAddr.(*net.TCPAddr)
|
||||
assert.Equal(t, expectedDst.IP.String(), actualDst.IP.String(), "destination IP should match server local addr")
|
||||
assert.Equal(t, expectedDst.Port, actualDst.Port, "destination port should match server local addr")
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
package tcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/netutil"
|
||||
)
|
||||
|
||||
// errIdleTimeout is returned when a relay connection is closed due to inactivity.
|
||||
var errIdleTimeout = errors.New("idle timeout")
|
||||
|
||||
// DefaultIdleTimeout is the default idle timeout for TCP relay connections.
|
||||
// A zero value disables idle timeout checking.
|
||||
const DefaultIdleTimeout = 5 * time.Minute
|
||||
|
||||
// halfCloser is implemented by connections that support half-close
|
||||
// (e.g. *net.TCPConn). When one copy direction finishes, we signal
|
||||
// EOF to the remote by closing the write side while keeping the read
|
||||
// side open so the other direction can drain.
|
||||
type halfCloser interface {
|
||||
CloseWrite() error
|
||||
}
|
||||
|
||||
// copyBufPool avoids allocating a new 32KB buffer per io.Copy call.
|
||||
var copyBufPool = sync.Pool{
|
||||
New: func() any {
|
||||
buf := make([]byte, 32*1024)
|
||||
return &buf
|
||||
},
|
||||
}
|
||||
|
||||
// Relay copies data bidirectionally between src and dst until both
|
||||
// sides are done or the context is canceled. When idleTimeout is
|
||||
// non-zero, each direction's read is deadline-guarded; if no data
|
||||
// flows within the timeout the connection is torn down. When one
|
||||
// direction finishes, it half-closes the write side of the
|
||||
// destination (if supported) to signal EOF, allowing the other
|
||||
// direction to drain gracefully before the full connection teardown.
|
||||
func Relay(ctx context.Context, logger *log.Entry, src, dst net.Conn, idleTimeout time.Duration) (srcToDst, dstToSrc int64) {
|
||||
ctx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
|
||||
go func() {
|
||||
<-ctx.Done()
|
||||
_ = src.Close()
|
||||
_ = dst.Close()
|
||||
}()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(2)
|
||||
|
||||
var errSrcToDst, errDstToSrc error
|
||||
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
srcToDst, errSrcToDst = copyWithIdleTimeout(dst, src, idleTimeout)
|
||||
halfClose(dst)
|
||||
cancel()
|
||||
}()
|
||||
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
dstToSrc, errDstToSrc = copyWithIdleTimeout(src, dst, idleTimeout)
|
||||
halfClose(src)
|
||||
cancel()
|
||||
}()
|
||||
|
||||
wg.Wait()
|
||||
|
||||
if errors.Is(errSrcToDst, errIdleTimeout) || errors.Is(errDstToSrc, errIdleTimeout) {
|
||||
logger.Debug("relay closed due to idle timeout")
|
||||
}
|
||||
if errSrcToDst != nil && !isExpectedCopyError(errSrcToDst) {
|
||||
logger.Debugf("relay copy error (src→dst): %v", errSrcToDst)
|
||||
}
|
||||
if errDstToSrc != nil && !isExpectedCopyError(errDstToSrc) {
|
||||
logger.Debugf("relay copy error (dst→src): %v", errDstToSrc)
|
||||
}
|
||||
|
||||
return srcToDst, dstToSrc
|
||||
}
|
||||
|
||||
// copyWithIdleTimeout copies from src to dst using a pooled buffer.
|
||||
// When idleTimeout > 0 it sets a read deadline on src before each
|
||||
// read and treats a timeout as an idle-triggered close.
|
||||
func copyWithIdleTimeout(dst io.Writer, src io.Reader, idleTimeout time.Duration) (int64, error) {
|
||||
bufp := copyBufPool.Get().(*[]byte)
|
||||
defer copyBufPool.Put(bufp)
|
||||
|
||||
if idleTimeout <= 0 {
|
||||
return io.CopyBuffer(dst, src, *bufp)
|
||||
}
|
||||
|
||||
conn, ok := src.(net.Conn)
|
||||
if !ok {
|
||||
return io.CopyBuffer(dst, src, *bufp)
|
||||
}
|
||||
|
||||
buf := *bufp
|
||||
var total int64
|
||||
for {
|
||||
if err := conn.SetReadDeadline(time.Now().Add(idleTimeout)); err != nil {
|
||||
return total, err
|
||||
}
|
||||
nr, readErr := src.Read(buf)
|
||||
if nr > 0 {
|
||||
n, err := checkedWrite(dst, buf[:nr])
|
||||
total += n
|
||||
if err != nil {
|
||||
return total, err
|
||||
}
|
||||
}
|
||||
if readErr != nil {
|
||||
if netutil.IsTimeout(readErr) {
|
||||
return total, errIdleTimeout
|
||||
}
|
||||
return total, readErr
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// checkedWrite writes buf to dst and returns the number of bytes written.
|
||||
// It guards against short writes and negative counts per io.Copy convention.
|
||||
func checkedWrite(dst io.Writer, buf []byte) (int64, error) {
|
||||
nw, err := dst.Write(buf)
|
||||
if nw < 0 || nw > len(buf) {
|
||||
nw = 0
|
||||
}
|
||||
if err != nil {
|
||||
return int64(nw), err
|
||||
}
|
||||
if nw != len(buf) {
|
||||
return int64(nw), io.ErrShortWrite
|
||||
}
|
||||
return int64(nw), nil
|
||||
}
|
||||
|
||||
func isExpectedCopyError(err error) bool {
|
||||
return errors.Is(err, errIdleTimeout) || netutil.IsExpectedError(err)
|
||||
}
|
||||
|
||||
// halfClose attempts to half-close the write side of the connection.
|
||||
// If the connection does not support half-close, this is a no-op.
|
||||
func halfClose(conn net.Conn) {
|
||||
if hc, ok := conn.(halfCloser); ok {
|
||||
// Best-effort; the full close will follow shortly.
|
||||
_ = hc.CloseWrite()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,210 @@
|
||||
package tcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/netutil"
|
||||
)
|
||||
|
||||
func TestRelay_BidirectionalCopy(t *testing.T) {
|
||||
srcClient, srcServer := net.Pipe()
|
||||
dstClient, dstServer := net.Pipe()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
ctx := context.Background()
|
||||
|
||||
srcData := []byte("hello from src")
|
||||
dstData := []byte("hello from dst")
|
||||
|
||||
// dst side: write response first, then read + close.
|
||||
go func() {
|
||||
_, _ = dstClient.Write(dstData)
|
||||
buf := make([]byte, 256)
|
||||
_, _ = dstClient.Read(buf)
|
||||
dstClient.Close()
|
||||
}()
|
||||
|
||||
// src side: read the response, then send data + close.
|
||||
go func() {
|
||||
buf := make([]byte, 256)
|
||||
_, _ = srcClient.Read(buf)
|
||||
_, _ = srcClient.Write(srcData)
|
||||
srcClient.Close()
|
||||
}()
|
||||
|
||||
s2d, d2s := Relay(ctx, logger, srcServer, dstServer, 0)
|
||||
|
||||
assert.Equal(t, int64(len(srcData)), s2d, "bytes src→dst")
|
||||
assert.Equal(t, int64(len(dstData)), d2s, "bytes dst→src")
|
||||
}
|
||||
|
||||
func TestRelay_ContextCancellation(t *testing.T) {
|
||||
srcClient, srcServer := net.Pipe()
|
||||
dstClient, dstServer := net.Pipe()
|
||||
defer srcClient.Close()
|
||||
defer dstClient.Close()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
Relay(ctx, logger, srcServer, dstServer, 0)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Cancel should cause Relay to return.
|
||||
cancel()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("Relay did not return after context cancellation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRelay_OneSideClosed(t *testing.T) {
|
||||
srcClient, srcServer := net.Pipe()
|
||||
dstClient, dstServer := net.Pipe()
|
||||
defer dstClient.Close()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
ctx := context.Background()
|
||||
|
||||
// Close src immediately. Relay should complete without hanging.
|
||||
srcClient.Close()
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
Relay(ctx, logger, srcServer, dstServer, 0)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("Relay did not return after one side closed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRelay_LargeTransfer(t *testing.T) {
|
||||
srcClient, srcServer := net.Pipe()
|
||||
dstClient, dstServer := net.Pipe()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
ctx := context.Background()
|
||||
|
||||
// 1MB of data.
|
||||
data := make([]byte, 1<<20)
|
||||
for i := range data {
|
||||
data[i] = byte(i % 256)
|
||||
}
|
||||
|
||||
go func() {
|
||||
_, _ = srcClient.Write(data)
|
||||
srcClient.Close()
|
||||
}()
|
||||
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
received, err := io.ReadAll(dstClient)
|
||||
if err != nil {
|
||||
errCh <- err
|
||||
return
|
||||
}
|
||||
if len(received) != len(data) {
|
||||
errCh <- fmt.Errorf("expected %d bytes, got %d", len(data), len(received))
|
||||
return
|
||||
}
|
||||
errCh <- nil
|
||||
dstClient.Close()
|
||||
}()
|
||||
|
||||
s2d, _ := Relay(ctx, logger, srcServer, dstServer, 0)
|
||||
assert.Equal(t, int64(len(data)), s2d, "should transfer all bytes")
|
||||
require.NoError(t, <-errCh)
|
||||
}
|
||||
|
||||
func TestRelay_IdleTimeout(t *testing.T) {
|
||||
// Use real TCP connections so SetReadDeadline works (net.Pipe
|
||||
// does not support deadlines).
|
||||
srcLn, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer srcLn.Close()
|
||||
|
||||
dstLn, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer dstLn.Close()
|
||||
|
||||
srcClient, err := net.Dial("tcp", srcLn.Addr().String())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer srcClient.Close()
|
||||
|
||||
srcServer, err := srcLn.Accept()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
dstClient, err := net.Dial("tcp", dstLn.Addr().String())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer dstClient.Close()
|
||||
|
||||
dstServer, err := dstLn.Accept()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
ctx := context.Background()
|
||||
|
||||
// Send initial data to prove the relay works.
|
||||
go func() {
|
||||
_, _ = srcClient.Write([]byte("ping"))
|
||||
}()
|
||||
|
||||
done := make(chan struct{})
|
||||
var s2d, d2s int64
|
||||
go func() {
|
||||
s2d, d2s = Relay(ctx, logger, srcServer, dstServer, 200*time.Millisecond)
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// Read the forwarded data on the dst side.
|
||||
buf := make([]byte, 64)
|
||||
n, err := dstClient.Read(buf)
|
||||
assert.NoError(t, err)
|
||||
assert.Equal(t, "ping", string(buf[:n]))
|
||||
|
||||
// Now stop sending. The relay should close after the idle timeout.
|
||||
select {
|
||||
case <-done:
|
||||
assert.Greater(t, s2d, int64(0), "should have transferred initial data")
|
||||
_ = d2s
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("Relay did not exit after idle timeout")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsExpectedError(t *testing.T) {
|
||||
assert.True(t, netutil.IsExpectedError(net.ErrClosed))
|
||||
assert.True(t, netutil.IsExpectedError(context.Canceled))
|
||||
assert.True(t, netutil.IsExpectedError(io.EOF))
|
||||
assert.False(t, netutil.IsExpectedError(io.ErrUnexpectedEOF))
|
||||
}
|
||||
@@ -0,0 +1,658 @@
|
||||
package tcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/accesslog"
|
||||
"github.com/netbirdio/netbird/proxy/internal/restrict"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
)
|
||||
|
||||
// defaultDialTimeout is the fallback dial timeout when no per-route
|
||||
// timeout is configured.
|
||||
const defaultDialTimeout = 30 * time.Second
|
||||
|
||||
// errAccessRestricted is returned by relayTCP for access restriction
|
||||
// denials so callers can skip warn-level logging (already logged at debug).
|
||||
var errAccessRestricted = errors.New("rejected by access restrictions")
|
||||
|
||||
// SNIHost is a typed key for SNI hostname lookups.
|
||||
type SNIHost string
|
||||
|
||||
// RouteType specifies how a connection should be handled.
|
||||
type RouteType int
|
||||
|
||||
const (
|
||||
// RouteHTTP routes the connection through the HTTP reverse proxy.
|
||||
RouteHTTP RouteType = iota
|
||||
// RouteTCP relays the connection directly to the backend (TLS passthrough).
|
||||
RouteTCP
|
||||
)
|
||||
|
||||
const (
|
||||
// sniPeekTimeout is the deadline for reading the TLS ClientHello.
|
||||
sniPeekTimeout = 5 * time.Second
|
||||
// DefaultDrainTimeout is the default grace period for in-flight relay
|
||||
// connections to finish during shutdown.
|
||||
DefaultDrainTimeout = 30 * time.Second
|
||||
// DefaultMaxRelayConns is the default cap on concurrent TCP relay connections per router.
|
||||
DefaultMaxRelayConns = 4096
|
||||
// httpChannelBuffer is the capacity of the channel feeding HTTP connections.
|
||||
httpChannelBuffer = 4096
|
||||
)
|
||||
|
||||
// DialResolver returns a DialContextFunc for the given account.
|
||||
type DialResolver func(accountID types.AccountID) (types.DialContextFunc, error)
|
||||
|
||||
// Route describes where a connection for a given SNI should be sent.
|
||||
type Route struct {
|
||||
Type RouteType
|
||||
AccountID types.AccountID
|
||||
ServiceID types.ServiceID
|
||||
// Domain is the service's configured domain, used for access log entries.
|
||||
Domain string
|
||||
// Protocol is the frontend protocol (tcp, tls), used for access log entries.
|
||||
Protocol accesslog.Protocol
|
||||
// Target is the backend address for TCP relay (e.g. "10.0.0.5:5432").
|
||||
Target string
|
||||
// ProxyProtocol enables sending a PROXY protocol v2 header to the backend.
|
||||
ProxyProtocol bool
|
||||
// DialTimeout overrides the default dial timeout for this route.
|
||||
// Zero uses defaultDialTimeout.
|
||||
DialTimeout time.Duration
|
||||
// SessionIdleTimeout overrides the default idle timeout for relay connections.
|
||||
// Zero uses DefaultIdleTimeout.
|
||||
SessionIdleTimeout time.Duration
|
||||
// Filter holds connection-level IP/geo restrictions. Nil means no restrictions.
|
||||
Filter *restrict.Filter
|
||||
}
|
||||
|
||||
// l4Logger sends layer-4 access log entries to the management server.
|
||||
type l4Logger interface {
|
||||
LogL4(entry accesslog.L4Entry)
|
||||
}
|
||||
|
||||
// RelayObserver receives callbacks for TCP relay lifecycle events.
|
||||
// All methods must be safe for concurrent use.
|
||||
type RelayObserver interface {
|
||||
TCPRelayStarted(accountID types.AccountID)
|
||||
TCPRelayEnded(accountID types.AccountID, duration time.Duration, srcToDst, dstToSrc int64)
|
||||
TCPRelayDialError(accountID types.AccountID)
|
||||
TCPRelayRejected(accountID types.AccountID)
|
||||
}
|
||||
|
||||
// Router accepts raw TCP connections on a shared listener, peeks at
|
||||
// the TLS ClientHello to extract the SNI, and routes the connection
|
||||
// to either the HTTP reverse proxy or a direct TCP relay.
|
||||
type Router struct {
|
||||
logger *log.Logger
|
||||
// httpCh is immutable after construction: set only in NewRouter, nil in NewPortRouter.
|
||||
httpCh chan net.Conn
|
||||
httpListener *chanListener
|
||||
mu sync.RWMutex
|
||||
routes map[SNIHost][]Route
|
||||
fallback *Route
|
||||
draining bool
|
||||
dialResolve DialResolver
|
||||
activeConns sync.WaitGroup
|
||||
activeRelays sync.WaitGroup
|
||||
relaySem chan struct{}
|
||||
drainDone chan struct{}
|
||||
observer RelayObserver
|
||||
accessLog l4Logger
|
||||
geo restrict.GeoResolver
|
||||
// svcCtxs tracks a context per service ID. All relay goroutines for a
|
||||
// service derive from its context; canceling it kills them immediately.
|
||||
svcCtxs map[types.ServiceID]context.Context
|
||||
svcCancels map[types.ServiceID]context.CancelFunc
|
||||
}
|
||||
|
||||
// NewRouter creates a new SNI-based connection router.
|
||||
func NewRouter(logger *log.Logger, dialResolve DialResolver, addr net.Addr) *Router {
|
||||
httpCh := make(chan net.Conn, httpChannelBuffer)
|
||||
return &Router{
|
||||
logger: logger,
|
||||
httpCh: httpCh,
|
||||
httpListener: newChanListener(httpCh, addr),
|
||||
routes: make(map[SNIHost][]Route),
|
||||
dialResolve: dialResolve,
|
||||
relaySem: make(chan struct{}, DefaultMaxRelayConns),
|
||||
svcCtxs: make(map[types.ServiceID]context.Context),
|
||||
svcCancels: make(map[types.ServiceID]context.CancelFunc),
|
||||
}
|
||||
}
|
||||
|
||||
// NewPortRouter creates a Router for a dedicated port without an HTTP
|
||||
// channel. Connections that don't match any SNI route fall through to
|
||||
// the fallback relay (if set) or are closed.
|
||||
func NewPortRouter(logger *log.Logger, dialResolve DialResolver) *Router {
|
||||
return &Router{
|
||||
logger: logger,
|
||||
routes: make(map[SNIHost][]Route),
|
||||
dialResolve: dialResolve,
|
||||
relaySem: make(chan struct{}, DefaultMaxRelayConns),
|
||||
svcCtxs: make(map[types.ServiceID]context.Context),
|
||||
svcCancels: make(map[types.ServiceID]context.CancelFunc),
|
||||
}
|
||||
}
|
||||
|
||||
// HTTPListener returns a net.Listener that yields connections routed
|
||||
// to the HTTP handler. Use this with http.Server.ServeTLS.
|
||||
func (r *Router) HTTPListener() net.Listener {
|
||||
return r.httpListener
|
||||
}
|
||||
|
||||
// AddRoute registers an SNI route. Multiple routes for the same host are
|
||||
// stored and resolved by priority at lookup time (HTTP > TCP).
|
||||
// Empty host is ignored to prevent conflicts with ECH/ESNI fallback.
|
||||
func (r *Router) AddRoute(host SNIHost, route Route) {
|
||||
host = SNIHost(strings.ToLower(string(host)))
|
||||
if host == "" {
|
||||
return
|
||||
}
|
||||
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
routes := r.routes[host]
|
||||
for i, existing := range routes {
|
||||
if existing.ServiceID == route.ServiceID {
|
||||
r.cancelServiceLocked(route.ServiceID)
|
||||
routes[i] = route
|
||||
return
|
||||
}
|
||||
}
|
||||
r.routes[host] = append(routes, route)
|
||||
}
|
||||
|
||||
// RemoveRoute removes the route for the given host and service ID.
|
||||
// Active relay connections for the service are closed immediately.
|
||||
// If other routes remain for the host, they are preserved.
|
||||
func (r *Router) RemoveRoute(host SNIHost, svcID types.ServiceID) {
|
||||
host = SNIHost(strings.ToLower(string(host)))
|
||||
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
r.routes[host] = slices.DeleteFunc(r.routes[host], func(route Route) bool {
|
||||
return route.ServiceID == svcID
|
||||
})
|
||||
if len(r.routes[host]) == 0 {
|
||||
delete(r.routes, host)
|
||||
}
|
||||
r.cancelServiceLocked(svcID)
|
||||
}
|
||||
|
||||
// SetFallback registers a catch-all route for connections that don't
|
||||
// match any SNI route. On a port router this handles plain TCP relay;
|
||||
// on the main router it takes priority over the HTTP channel.
|
||||
func (r *Router) SetFallback(route Route) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.fallback = &route
|
||||
}
|
||||
|
||||
// RemoveFallback clears the catch-all fallback route and closes any
|
||||
// active relay connections for the given service.
|
||||
func (r *Router) RemoveFallback(svcID types.ServiceID) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.fallback = nil
|
||||
r.cancelServiceLocked(svcID)
|
||||
}
|
||||
|
||||
// SetObserver sets the relay lifecycle observer. Must be called before Serve.
|
||||
func (r *Router) SetObserver(obs RelayObserver) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.observer = obs
|
||||
}
|
||||
|
||||
// SetAccessLogger sets the L4 access logger. Must be called before Serve.
|
||||
func (r *Router) SetAccessLogger(l l4Logger) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.accessLog = l
|
||||
}
|
||||
|
||||
// getObserver returns the current relay observer under the read lock.
|
||||
func (r *Router) getObserver() RelayObserver {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return r.observer
|
||||
}
|
||||
|
||||
// IsEmpty returns true when the router has no SNI routes and no fallback.
|
||||
func (r *Router) IsEmpty() bool {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return len(r.routes) == 0 && r.fallback == nil
|
||||
}
|
||||
|
||||
// Serve accepts connections from ln and routes them based on SNI.
|
||||
// It blocks until ctx is canceled or ln is closed, then drains
|
||||
// active relay connections up to DefaultDrainTimeout.
|
||||
func (r *Router) Serve(ctx context.Context, ln net.Listener) error {
|
||||
done := make(chan struct{})
|
||||
defer close(done)
|
||||
|
||||
go func() {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
_ = ln.Close()
|
||||
if r.httpListener != nil {
|
||||
r.httpListener.Close()
|
||||
}
|
||||
case <-done:
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
if ctx.Err() != nil || errors.Is(err, net.ErrClosed) {
|
||||
if ok := r.Drain(DefaultDrainTimeout); !ok {
|
||||
r.logger.Warn("timed out waiting for connections to drain")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
r.logger.Debugf("SNI router accept: %v", err)
|
||||
continue
|
||||
}
|
||||
r.activeConns.Add(1)
|
||||
go func() {
|
||||
defer r.activeConns.Done()
|
||||
r.handleConn(ctx, conn)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
// handleConn peeks at the TLS ClientHello and routes the connection.
|
||||
func (r *Router) handleConn(ctx context.Context, conn net.Conn) {
|
||||
// Fast path: when no SNI routes and no HTTP channel exist (pure TCP
|
||||
// fallback port), skip the TLS peek entirely to avoid read errors on
|
||||
// non-TLS connections and reduce latency.
|
||||
if r.isFallbackOnly() {
|
||||
r.handleUnmatched(ctx, conn)
|
||||
return
|
||||
}
|
||||
|
||||
if err := conn.SetReadDeadline(time.Now().Add(sniPeekTimeout)); err != nil {
|
||||
r.logger.Debugf("set SNI peek deadline: %v", err)
|
||||
_ = conn.Close()
|
||||
return
|
||||
}
|
||||
|
||||
sni, wrapped, err := PeekClientHello(conn)
|
||||
if err != nil {
|
||||
r.logger.Debugf("SNI peek: %v", err)
|
||||
if wrapped != nil {
|
||||
r.handleUnmatched(ctx, wrapped)
|
||||
} else {
|
||||
_ = conn.Close()
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
if err := wrapped.SetReadDeadline(time.Time{}); err != nil {
|
||||
r.logger.Debugf("clear SNI peek deadline: %v", err)
|
||||
_ = wrapped.Close()
|
||||
return
|
||||
}
|
||||
|
||||
host := SNIHost(strings.ToLower(sni))
|
||||
route, ok := r.lookupRoute(host)
|
||||
if !ok {
|
||||
r.handleUnmatched(ctx, wrapped)
|
||||
return
|
||||
}
|
||||
|
||||
if route.Type == RouteHTTP {
|
||||
r.sendToHTTP(wrapped)
|
||||
return
|
||||
}
|
||||
|
||||
if err := r.relayTCP(ctx, wrapped, host, route); err != nil {
|
||||
if !errors.Is(err, errAccessRestricted) {
|
||||
r.logger.WithFields(log.Fields{
|
||||
"sni": host,
|
||||
"service_id": route.ServiceID,
|
||||
"target": route.Target,
|
||||
}).Warnf("TCP relay: %v", err)
|
||||
}
|
||||
_ = wrapped.Close()
|
||||
}
|
||||
}
|
||||
|
||||
// isFallbackOnly returns true when the router has no SNI routes and no HTTP
|
||||
// channel, meaning all connections should go directly to the fallback relay.
|
||||
func (r *Router) isFallbackOnly() bool {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return len(r.routes) == 0 && r.httpCh == nil
|
||||
}
|
||||
|
||||
// handleUnmatched routes a connection that didn't match any SNI route.
|
||||
// This includes ECH/ESNI connections where the cleartext SNI is empty.
|
||||
// It tries the fallback relay first, then the HTTP channel, and closes
|
||||
// the connection if neither is available.
|
||||
func (r *Router) handleUnmatched(ctx context.Context, conn net.Conn) {
|
||||
r.mu.RLock()
|
||||
fb := r.fallback
|
||||
r.mu.RUnlock()
|
||||
|
||||
if fb != nil {
|
||||
if err := r.relayTCP(ctx, conn, SNIHost("fallback"), *fb); err != nil {
|
||||
if !errors.Is(err, errAccessRestricted) {
|
||||
r.logger.WithFields(log.Fields{
|
||||
"service_id": fb.ServiceID,
|
||||
"target": fb.Target,
|
||||
}).Warnf("TCP relay (fallback): %v", err)
|
||||
}
|
||||
_ = conn.Close()
|
||||
}
|
||||
return
|
||||
}
|
||||
r.sendToHTTP(conn)
|
||||
}
|
||||
|
||||
// lookupRoute returns the highest-priority route for the given SNI host.
|
||||
// HTTP routes take precedence over TCP routes.
|
||||
func (r *Router) lookupRoute(host SNIHost) (Route, bool) {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
routes, ok := r.routes[host]
|
||||
if !ok || len(routes) == 0 {
|
||||
return Route{}, false
|
||||
}
|
||||
best := routes[0]
|
||||
for _, route := range routes[1:] {
|
||||
if route.Type < best.Type {
|
||||
best = route
|
||||
}
|
||||
}
|
||||
return best, true
|
||||
}
|
||||
|
||||
// sendToHTTP feeds the connection to the HTTP handler via the channel.
|
||||
// If no HTTP channel is configured (port router), the router is
|
||||
// draining, or the channel is full, the connection is closed.
|
||||
func (r *Router) sendToHTTP(conn net.Conn) {
|
||||
if r.httpCh == nil {
|
||||
_ = conn.Close()
|
||||
return
|
||||
}
|
||||
|
||||
r.mu.RLock()
|
||||
draining := r.draining
|
||||
r.mu.RUnlock()
|
||||
|
||||
if draining {
|
||||
_ = conn.Close()
|
||||
return
|
||||
}
|
||||
|
||||
select {
|
||||
case r.httpCh <- conn:
|
||||
default:
|
||||
r.logger.Warnf("HTTP channel full, dropping connection from %s", conn.RemoteAddr())
|
||||
_ = conn.Close()
|
||||
}
|
||||
}
|
||||
|
||||
// Drain prevents new relay connections from starting and waits for all
|
||||
// in-flight connection handlers and active relays to finish, up to the
|
||||
// given timeout. Returns true if all completed, false on timeout.
|
||||
func (r *Router) Drain(timeout time.Duration) bool {
|
||||
r.mu.Lock()
|
||||
r.draining = true
|
||||
if r.drainDone == nil {
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
r.activeConns.Wait()
|
||||
r.activeRelays.Wait()
|
||||
close(done)
|
||||
}()
|
||||
r.drainDone = done
|
||||
}
|
||||
done := r.drainDone
|
||||
r.mu.Unlock()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
return true
|
||||
case <-time.After(timeout):
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// cancelServiceLocked cancels and removes the context for the given service,
|
||||
// closing all its active relay connections. Must be called with mu held.
|
||||
func (r *Router) cancelServiceLocked(svcID types.ServiceID) {
|
||||
if cancel, ok := r.svcCancels[svcID]; ok {
|
||||
cancel()
|
||||
delete(r.svcCtxs, svcID)
|
||||
delete(r.svcCancels, svcID)
|
||||
}
|
||||
}
|
||||
|
||||
// SetGeo sets the geolocation lookup used for country-based restrictions.
|
||||
func (r *Router) SetGeo(geo restrict.GeoResolver) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.geo = geo
|
||||
}
|
||||
|
||||
// checkRestrictions evaluates the route's access filter against the
|
||||
// connection's remote address. Returns Allow if the connection is
|
||||
// permitted, or a deny verdict indicating the reason.
|
||||
func (r *Router) checkRestrictions(conn net.Conn, route Route) restrict.Verdict {
|
||||
if route.Filter == nil {
|
||||
return restrict.Allow
|
||||
}
|
||||
|
||||
addr, err := addrFromConn(conn)
|
||||
if err != nil {
|
||||
r.logger.Debugf("cannot parse client address %s for restriction check, denying", conn.RemoteAddr())
|
||||
return restrict.DenyCIDR
|
||||
}
|
||||
|
||||
r.mu.RLock()
|
||||
geo := r.geo
|
||||
r.mu.RUnlock()
|
||||
|
||||
return route.Filter.Check(addr, geo)
|
||||
}
|
||||
|
||||
// relayTCP sets up and runs a bidirectional TCP relay.
|
||||
// The caller owns conn and must close it if this method returns an error.
|
||||
// On success (nil error), both conn and backend are closed by the relay.
|
||||
func (r *Router) relayTCP(ctx context.Context, conn net.Conn, sni SNIHost, route Route) error {
|
||||
if verdict := r.checkRestrictions(conn, route); verdict != restrict.Allow {
|
||||
r.logger.Debugf("connection from %s rejected by access restrictions: %s", conn.RemoteAddr(), verdict)
|
||||
r.logL4Deny(route, conn, verdict)
|
||||
return errAccessRestricted
|
||||
}
|
||||
|
||||
svcCtx, err := r.acquireRelay(ctx, route)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() {
|
||||
<-r.relaySem
|
||||
r.activeRelays.Done()
|
||||
}()
|
||||
|
||||
backend, err := r.dialBackend(svcCtx, route)
|
||||
if err != nil {
|
||||
obs := r.getObserver()
|
||||
if obs != nil {
|
||||
obs.TCPRelayDialError(route.AccountID)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
if route.ProxyProtocol {
|
||||
if err := writeProxyProtoV2(conn, backend); err != nil {
|
||||
_ = backend.Close()
|
||||
return fmt.Errorf("write PROXY protocol header: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
obs := r.getObserver()
|
||||
if obs != nil {
|
||||
obs.TCPRelayStarted(route.AccountID)
|
||||
}
|
||||
|
||||
entry := r.logger.WithFields(log.Fields{
|
||||
"sni": sni,
|
||||
"service_id": route.ServiceID,
|
||||
"target": route.Target,
|
||||
})
|
||||
entry.Debug("TCP relay started")
|
||||
|
||||
idleTimeout := route.SessionIdleTimeout
|
||||
if idleTimeout <= 0 {
|
||||
idleTimeout = DefaultIdleTimeout
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
s2d, d2s := Relay(svcCtx, entry, conn, backend, idleTimeout)
|
||||
elapsed := time.Since(start)
|
||||
|
||||
if obs != nil {
|
||||
obs.TCPRelayEnded(route.AccountID, elapsed, s2d, d2s)
|
||||
}
|
||||
entry.Debugf("TCP relay ended (client→backend: %d bytes, backend→client: %d bytes)", s2d, d2s)
|
||||
|
||||
r.logL4Entry(route, conn, elapsed, s2d, d2s)
|
||||
return nil
|
||||
}
|
||||
|
||||
// acquireRelay checks draining state, increments activeRelays, and acquires
|
||||
// a semaphore slot. Returns the per-service context on success.
|
||||
// The caller must release the semaphore and call activeRelays.Done() when done.
|
||||
func (r *Router) acquireRelay(ctx context.Context, route Route) (context.Context, error) {
|
||||
r.mu.Lock()
|
||||
if r.draining {
|
||||
r.mu.Unlock()
|
||||
return nil, errors.New("router is draining")
|
||||
}
|
||||
r.activeRelays.Add(1)
|
||||
svcCtx := r.getOrCreateServiceCtxLocked(ctx, route.ServiceID)
|
||||
r.mu.Unlock()
|
||||
|
||||
select {
|
||||
case r.relaySem <- struct{}{}:
|
||||
return svcCtx, nil
|
||||
default:
|
||||
r.activeRelays.Done()
|
||||
obs := r.getObserver()
|
||||
if obs != nil {
|
||||
obs.TCPRelayRejected(route.AccountID)
|
||||
}
|
||||
return nil, errors.New("TCP relay connection limit reached")
|
||||
}
|
||||
}
|
||||
|
||||
// dialBackend resolves the dialer for the route's account and dials the backend.
|
||||
func (r *Router) dialBackend(svcCtx context.Context, route Route) (net.Conn, error) {
|
||||
dialFn, err := r.dialResolve(route.AccountID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("resolve dialer: %w", err)
|
||||
}
|
||||
|
||||
dialTimeout := route.DialTimeout
|
||||
if dialTimeout <= 0 {
|
||||
dialTimeout = defaultDialTimeout
|
||||
}
|
||||
dialCtx, dialCancel := context.WithTimeout(svcCtx, dialTimeout)
|
||||
backend, err := dialFn(dialCtx, "tcp", route.Target)
|
||||
dialCancel()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("dial backend %s: %w", route.Target, err)
|
||||
}
|
||||
return backend, nil
|
||||
}
|
||||
|
||||
// logL4Entry sends a TCP relay access log entry if an access logger is configured.
|
||||
func (r *Router) logL4Entry(route Route, conn net.Conn, duration time.Duration, bytesUp, bytesDown int64) {
|
||||
r.mu.RLock()
|
||||
al := r.accessLog
|
||||
r.mu.RUnlock()
|
||||
|
||||
if al == nil {
|
||||
return
|
||||
}
|
||||
|
||||
sourceIP, _ := addrFromConn(conn)
|
||||
|
||||
al.LogL4(accesslog.L4Entry{
|
||||
AccountID: route.AccountID,
|
||||
ServiceID: route.ServiceID,
|
||||
Protocol: route.Protocol,
|
||||
Host: route.Domain,
|
||||
SourceIP: sourceIP,
|
||||
DurationMs: duration.Milliseconds(),
|
||||
BytesUpload: bytesUp,
|
||||
BytesDownload: bytesDown,
|
||||
})
|
||||
}
|
||||
|
||||
// logL4Deny sends an access log entry for a denied connection.
|
||||
func (r *Router) logL4Deny(route Route, conn net.Conn, verdict restrict.Verdict) {
|
||||
r.mu.RLock()
|
||||
al := r.accessLog
|
||||
r.mu.RUnlock()
|
||||
|
||||
if al == nil {
|
||||
return
|
||||
}
|
||||
|
||||
sourceIP, _ := addrFromConn(conn)
|
||||
|
||||
al.LogL4(accesslog.L4Entry{
|
||||
AccountID: route.AccountID,
|
||||
ServiceID: route.ServiceID,
|
||||
Protocol: route.Protocol,
|
||||
Host: route.Domain,
|
||||
SourceIP: sourceIP,
|
||||
DenyReason: verdict.String(),
|
||||
})
|
||||
}
|
||||
|
||||
// getOrCreateServiceCtxLocked returns the context for a service, creating one
|
||||
// if it doesn't exist yet. The context is a child of the server context.
|
||||
// Must be called with mu held.
|
||||
func (r *Router) getOrCreateServiceCtxLocked(parent context.Context, svcID types.ServiceID) context.Context {
|
||||
if ctx, ok := r.svcCtxs[svcID]; ok {
|
||||
return ctx
|
||||
}
|
||||
ctx, cancel := context.WithCancel(parent)
|
||||
r.svcCtxs[svcID] = ctx
|
||||
r.svcCancels[svcID] = cancel
|
||||
return ctx
|
||||
}
|
||||
|
||||
// addrFromConn extracts a netip.Addr from a connection's remote address.
|
||||
func addrFromConn(conn net.Conn) (netip.Addr, error) {
|
||||
remote := conn.RemoteAddr()
|
||||
if remote == nil {
|
||||
return netip.Addr{}, errors.New("no remote address")
|
||||
}
|
||||
ap, err := netip.ParseAddrPort(remote.String())
|
||||
if err != nil {
|
||||
return netip.Addr{}, err
|
||||
}
|
||||
return ap.Addr().Unmap(), nil
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,191 @@
|
||||
package tcp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
)
|
||||
|
||||
const (
|
||||
// TLS record header is 5 bytes: ContentType(1) + Version(2) + Length(2).
|
||||
tlsRecordHeaderLen = 5
|
||||
// TLS handshake type for ClientHello.
|
||||
handshakeTypeClientHello = 1
|
||||
// TLS ContentType for handshake messages.
|
||||
contentTypeHandshake = 22
|
||||
// SNI extension type (RFC 6066).
|
||||
extensionServerName = 0
|
||||
// SNI host name type.
|
||||
sniHostNameType = 0
|
||||
// maxClientHelloLen caps the ClientHello size we're willing to buffer.
|
||||
maxClientHelloLen = 16384
|
||||
// maxSNILen is the maximum valid DNS hostname length per RFC 1035.
|
||||
maxSNILen = 253
|
||||
)
|
||||
|
||||
// PeekClientHello reads the TLS ClientHello from conn, extracts the SNI
|
||||
// server name, and returns a wrapped connection that replays the peeked
|
||||
// bytes transparently. If the data is not a valid TLS ClientHello or
|
||||
// contains no SNI extension, sni is empty and err is nil.
|
||||
//
|
||||
// ECH/ESNI: When the client uses Encrypted Client Hello (TLS 1.3), the
|
||||
// real server name is encrypted inside the encrypted_client_hello
|
||||
// extension. This parser only reads the cleartext server_name extension
|
||||
// (type 0x0000), so ECH connections return sni="" and are routed through
|
||||
// the fallback path (or HTTP channel), which is the correct behavior
|
||||
// for a transparent proxy that does not terminate TLS.
|
||||
func PeekClientHello(conn net.Conn) (sni string, wrapped net.Conn, err error) {
|
||||
// Read the 5-byte TLS record header into a small stack-friendly buffer.
|
||||
var header [tlsRecordHeaderLen]byte
|
||||
if _, err := io.ReadFull(conn, header[:]); err != nil {
|
||||
return "", nil, fmt.Errorf("read TLS record header: %w", err)
|
||||
}
|
||||
|
||||
if header[0] != contentTypeHandshake {
|
||||
return "", newPeekedConn(conn, header[:]), nil
|
||||
}
|
||||
|
||||
recordLen := int(binary.BigEndian.Uint16(header[3:5]))
|
||||
if recordLen == 0 || recordLen > maxClientHelloLen {
|
||||
return "", newPeekedConn(conn, header[:]), nil
|
||||
}
|
||||
|
||||
// Single allocation for header + payload. The peekedConn takes
|
||||
// ownership of this buffer, so no further copies are needed.
|
||||
buf := make([]byte, tlsRecordHeaderLen+recordLen)
|
||||
copy(buf, header[:])
|
||||
|
||||
n, err := io.ReadFull(conn, buf[tlsRecordHeaderLen:])
|
||||
if err != nil {
|
||||
return "", newPeekedConn(conn, buf[:tlsRecordHeaderLen+n]), fmt.Errorf("read TLS handshake payload: %w", err)
|
||||
}
|
||||
|
||||
sni = extractSNI(buf[tlsRecordHeaderLen:])
|
||||
return sni, newPeekedConn(conn, buf), nil
|
||||
}
|
||||
|
||||
// extractSNI parses a TLS handshake payload to find the SNI extension.
|
||||
// Returns empty string if the payload is not a ClientHello or has no SNI.
|
||||
func extractSNI(payload []byte) string {
|
||||
if len(payload) < 4 {
|
||||
return ""
|
||||
}
|
||||
|
||||
if payload[0] != handshakeTypeClientHello {
|
||||
return ""
|
||||
}
|
||||
|
||||
// Handshake length (3 bytes, big-endian).
|
||||
handshakeLen := int(payload[1])<<16 | int(payload[2])<<8 | int(payload[3])
|
||||
if handshakeLen > len(payload)-4 {
|
||||
return ""
|
||||
}
|
||||
|
||||
return parseSNIFromClientHello(payload[4 : 4+handshakeLen])
|
||||
}
|
||||
|
||||
// parseSNIFromClientHello walks the ClientHello message fields to reach
|
||||
// the extensions block and extract the server_name extension value.
|
||||
func parseSNIFromClientHello(msg []byte) string {
|
||||
// ClientHello layout:
|
||||
// ProtocolVersion(2) + Random(32) = 34 bytes minimum before session_id
|
||||
if len(msg) < 34 {
|
||||
return ""
|
||||
}
|
||||
|
||||
pos := 34
|
||||
|
||||
// Session ID (variable, 1 byte length prefix).
|
||||
if pos >= len(msg) {
|
||||
return ""
|
||||
}
|
||||
sessionIDLen := int(msg[pos])
|
||||
pos++
|
||||
pos += sessionIDLen
|
||||
|
||||
// Cipher suites (variable, 2 byte length prefix).
|
||||
if pos+2 > len(msg) {
|
||||
return ""
|
||||
}
|
||||
cipherSuitesLen := int(binary.BigEndian.Uint16(msg[pos : pos+2]))
|
||||
pos += 2 + cipherSuitesLen
|
||||
|
||||
// Compression methods (variable, 1 byte length prefix).
|
||||
if pos >= len(msg) {
|
||||
return ""
|
||||
}
|
||||
compMethodsLen := int(msg[pos])
|
||||
pos++
|
||||
pos += compMethodsLen
|
||||
|
||||
// Extensions (variable, 2 byte length prefix).
|
||||
if pos+2 > len(msg) {
|
||||
return ""
|
||||
}
|
||||
extensionsLen := int(binary.BigEndian.Uint16(msg[pos : pos+2]))
|
||||
pos += 2
|
||||
|
||||
extensionsEnd := pos + extensionsLen
|
||||
if extensionsEnd > len(msg) {
|
||||
return ""
|
||||
}
|
||||
|
||||
return findSNIExtension(msg[pos:extensionsEnd])
|
||||
}
|
||||
|
||||
// findSNIExtension iterates over TLS extensions and returns the host
|
||||
// name from the server_name extension, if present.
|
||||
func findSNIExtension(extensions []byte) string {
|
||||
pos := 0
|
||||
for pos+4 <= len(extensions) {
|
||||
extType := binary.BigEndian.Uint16(extensions[pos : pos+2])
|
||||
extLen := int(binary.BigEndian.Uint16(extensions[pos+2 : pos+4]))
|
||||
pos += 4
|
||||
|
||||
if pos+extLen > len(extensions) {
|
||||
return ""
|
||||
}
|
||||
|
||||
if extType == extensionServerName {
|
||||
return parseSNIExtensionData(extensions[pos : pos+extLen])
|
||||
}
|
||||
pos += extLen
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// parseSNIExtensionData parses the ServerNameList structure inside an
|
||||
// SNI extension to extract the host name.
|
||||
func parseSNIExtensionData(data []byte) string {
|
||||
if len(data) < 2 {
|
||||
return ""
|
||||
}
|
||||
listLen := int(binary.BigEndian.Uint16(data[0:2]))
|
||||
if listLen > len(data)-2 {
|
||||
return ""
|
||||
}
|
||||
|
||||
list := data[2 : 2+listLen]
|
||||
pos := 0
|
||||
for pos+3 <= len(list) {
|
||||
nameType := list[pos]
|
||||
nameLen := int(binary.BigEndian.Uint16(list[pos+1 : pos+3]))
|
||||
pos += 3
|
||||
|
||||
if pos+nameLen > len(list) {
|
||||
return ""
|
||||
}
|
||||
|
||||
if nameType == sniHostNameType {
|
||||
name := list[pos : pos+nameLen]
|
||||
if nameLen > maxSNILen || bytes.ContainsRune(name, 0) {
|
||||
return ""
|
||||
}
|
||||
return string(name)
|
||||
}
|
||||
pos += nameLen
|
||||
}
|
||||
return ""
|
||||
}
|
||||
@@ -0,0 +1,251 @@
|
||||
package tcp
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"io"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestPeekClientHello_ValidSNI(t *testing.T) {
|
||||
clientConn, serverConn := net.Pipe()
|
||||
defer clientConn.Close()
|
||||
defer serverConn.Close()
|
||||
|
||||
const expectedSNI = "example.com"
|
||||
trailingData := []byte("trailing data after handshake")
|
||||
|
||||
go func() {
|
||||
tlsConn := tls.Client(clientConn, &tls.Config{
|
||||
ServerName: expectedSNI,
|
||||
InsecureSkipVerify: true, //nolint:gosec
|
||||
})
|
||||
// The Handshake will send the ClientHello. It will fail because
|
||||
// our server side isn't doing a real TLS handshake, but that's
|
||||
// fine: we only need the ClientHello to be sent.
|
||||
_ = tlsConn.Handshake()
|
||||
}()
|
||||
|
||||
sni, wrapped, err := PeekClientHello(serverConn)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, expectedSNI, sni, "should extract SNI from ClientHello")
|
||||
assert.NotNil(t, wrapped, "wrapped connection should not be nil")
|
||||
|
||||
// Verify the wrapped connection replays the peeked bytes.
|
||||
// Read the first 5 bytes (TLS record header) to confirm replay.
|
||||
buf := make([]byte, 5)
|
||||
n, err := wrapped.Read(buf)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 5, n)
|
||||
assert.Equal(t, byte(contentTypeHandshake), buf[0], "first byte should be TLS handshake content type")
|
||||
|
||||
// Write trailing data from the client side and verify it arrives
|
||||
// through the wrapped connection after the peeked bytes.
|
||||
go func() {
|
||||
_, _ = clientConn.Write(trailingData)
|
||||
}()
|
||||
|
||||
// Drain the rest of the peeked ClientHello first.
|
||||
peekedRest := make([]byte, 16384)
|
||||
_, _ = wrapped.Read(peekedRest)
|
||||
|
||||
got := make([]byte, len(trailingData))
|
||||
n, err = io.ReadFull(wrapped, got)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, trailingData, got[:n])
|
||||
}
|
||||
|
||||
func TestPeekClientHello_MultipleSNIs(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
serverName string
|
||||
expectedSNI string
|
||||
}{
|
||||
{"simple domain", "example.com", "example.com"},
|
||||
{"subdomain", "sub.example.com", "sub.example.com"},
|
||||
{"deep subdomain", "a.b.c.example.com", "a.b.c.example.com"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
clientConn, serverConn := net.Pipe()
|
||||
defer clientConn.Close()
|
||||
defer serverConn.Close()
|
||||
|
||||
go func() {
|
||||
tlsConn := tls.Client(clientConn, &tls.Config{
|
||||
ServerName: tt.serverName,
|
||||
InsecureSkipVerify: true, //nolint:gosec
|
||||
})
|
||||
_ = tlsConn.Handshake()
|
||||
}()
|
||||
|
||||
sni, wrapped, err := PeekClientHello(serverConn)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tt.expectedSNI, sni)
|
||||
assert.NotNil(t, wrapped)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPeekClientHello_NonTLSData(t *testing.T) {
|
||||
clientConn, serverConn := net.Pipe()
|
||||
defer clientConn.Close()
|
||||
defer serverConn.Close()
|
||||
|
||||
// Send plain HTTP data (not TLS).
|
||||
httpData := []byte("GET / HTTP/1.1\r\nHost: example.com\r\n\r\n")
|
||||
go func() {
|
||||
_, _ = clientConn.Write(httpData)
|
||||
}()
|
||||
|
||||
sni, wrapped, err := PeekClientHello(serverConn)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, sni, "should return empty SNI for non-TLS data")
|
||||
assert.NotNil(t, wrapped)
|
||||
|
||||
// Verify the wrapped connection still provides the original data.
|
||||
buf := make([]byte, len(httpData))
|
||||
n, err := io.ReadFull(wrapped, buf)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, httpData, buf[:n], "wrapped connection should replay original data")
|
||||
}
|
||||
|
||||
func TestPeekClientHello_TruncatedHeader(t *testing.T) {
|
||||
clientConn, serverConn := net.Pipe()
|
||||
defer serverConn.Close()
|
||||
|
||||
// Write only 3 bytes then close, fewer than the 5-byte TLS header.
|
||||
go func() {
|
||||
_, _ = clientConn.Write([]byte{0x16, 0x03, 0x01})
|
||||
clientConn.Close()
|
||||
}()
|
||||
|
||||
_, _, err := PeekClientHello(serverConn)
|
||||
assert.Error(t, err, "should error on truncated header")
|
||||
}
|
||||
|
||||
func TestPeekClientHello_TruncatedPayload(t *testing.T) {
|
||||
clientConn, serverConn := net.Pipe()
|
||||
defer serverConn.Close()
|
||||
|
||||
// Write a valid TLS header claiming 100 bytes, but only send 10.
|
||||
go func() {
|
||||
header := []byte{0x16, 0x03, 0x01, 0x00, 0x64} // 100 bytes claimed
|
||||
_, _ = clientConn.Write(header)
|
||||
_, _ = clientConn.Write(make([]byte, 10))
|
||||
clientConn.Close()
|
||||
}()
|
||||
|
||||
_, _, err := PeekClientHello(serverConn)
|
||||
assert.Error(t, err, "should error on truncated payload")
|
||||
}
|
||||
|
||||
func TestPeekClientHello_ZeroLengthRecord(t *testing.T) {
|
||||
clientConn, serverConn := net.Pipe()
|
||||
defer clientConn.Close()
|
||||
defer serverConn.Close()
|
||||
|
||||
// TLS handshake header with zero-length payload.
|
||||
go func() {
|
||||
_, _ = clientConn.Write([]byte{0x16, 0x03, 0x01, 0x00, 0x00})
|
||||
}()
|
||||
|
||||
sni, wrapped, err := PeekClientHello(serverConn)
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, sni)
|
||||
assert.NotNil(t, wrapped)
|
||||
}
|
||||
|
||||
func TestExtractSNI_InvalidPayload(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
payload []byte
|
||||
}{
|
||||
{"nil", nil},
|
||||
{"empty", []byte{}},
|
||||
{"too short", []byte{0x01, 0x00}},
|
||||
{"wrong handshake type", []byte{0x02, 0x00, 0x00, 0x05, 0x03, 0x03, 0x00, 0x00, 0x00}},
|
||||
{"truncated client hello", []byte{0x01, 0x00, 0x00, 0x20}}, // claims 32 bytes but has none
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
assert.Empty(t, extractSNI(tt.payload))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPeekedConn_CloseWrite(t *testing.T) {
|
||||
t.Run("delegates to underlying TCPConn", func(t *testing.T) {
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer ln.Close()
|
||||
|
||||
accepted := make(chan net.Conn, 1)
|
||||
go func() {
|
||||
c, err := ln.Accept()
|
||||
if err == nil {
|
||||
accepted <- c
|
||||
}
|
||||
}()
|
||||
|
||||
client, err := net.Dial("tcp", ln.Addr().String())
|
||||
require.NoError(t, err)
|
||||
defer client.Close()
|
||||
|
||||
server := <-accepted
|
||||
defer server.Close()
|
||||
|
||||
wrapped := newPeekedConn(server, []byte("peeked"))
|
||||
|
||||
// CloseWrite should succeed on a real TCP connection.
|
||||
err = wrapped.CloseWrite()
|
||||
assert.NoError(t, err)
|
||||
|
||||
// The client should see EOF on reads after CloseWrite.
|
||||
buf := make([]byte, 1)
|
||||
_, err = client.Read(buf)
|
||||
assert.Equal(t, io.EOF, err, "client should see EOF after half-close")
|
||||
})
|
||||
|
||||
t.Run("no-op on non-halfcloser", func(t *testing.T) {
|
||||
// net.Pipe does not implement CloseWrite.
|
||||
_, server := net.Pipe()
|
||||
defer server.Close()
|
||||
|
||||
wrapped := newPeekedConn(server, []byte("peeked"))
|
||||
err := wrapped.CloseWrite()
|
||||
assert.NoError(t, err, "should be no-op on non-halfcloser")
|
||||
})
|
||||
}
|
||||
|
||||
func TestPeekedConn_ReplayAndPassthrough(t *testing.T) {
|
||||
clientConn, serverConn := net.Pipe()
|
||||
defer clientConn.Close()
|
||||
defer serverConn.Close()
|
||||
|
||||
peeked := []byte("peeked-data")
|
||||
subsequent := []byte("subsequent-data")
|
||||
|
||||
wrapped := newPeekedConn(serverConn, peeked)
|
||||
|
||||
go func() {
|
||||
_, _ = clientConn.Write(subsequent)
|
||||
}()
|
||||
|
||||
// Read should return peeked data first.
|
||||
buf := make([]byte, len(peeked))
|
||||
n, err := io.ReadFull(wrapped, buf)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, peeked, buf[:n])
|
||||
|
||||
// Then subsequent data from the real connection.
|
||||
buf = make([]byte, len(subsequent))
|
||||
n, err = io.ReadFull(wrapped, buf)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, subsequent, buf[:n])
|
||||
}
|
||||
@@ -1,5 +1,56 @@
|
||||
// Package types defines common types used across the proxy package.
|
||||
package types
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"time"
|
||||
)
|
||||
|
||||
// AccountID represents a unique identifier for a NetBird account.
|
||||
type AccountID string
|
||||
|
||||
// ServiceID represents a unique identifier for a proxy service.
|
||||
type ServiceID string
|
||||
|
||||
// ServiceMode describes how a reverse proxy service is exposed.
|
||||
type ServiceMode string
|
||||
|
||||
const (
|
||||
ServiceModeHTTP ServiceMode = "http"
|
||||
ServiceModeTCP ServiceMode = "tcp"
|
||||
ServiceModeUDP ServiceMode = "udp"
|
||||
ServiceModeTLS ServiceMode = "tls"
|
||||
)
|
||||
|
||||
// IsL4 returns true for TCP, UDP, and TLS modes.
|
||||
func (m ServiceMode) IsL4() bool {
|
||||
return m == ServiceModeTCP || m == ServiceModeUDP || m == ServiceModeTLS
|
||||
}
|
||||
|
||||
// RelayDirection indicates the direction of a relayed packet.
|
||||
type RelayDirection string
|
||||
|
||||
const (
|
||||
RelayDirectionClientToBackend RelayDirection = "client_to_backend"
|
||||
RelayDirectionBackendToClient RelayDirection = "backend_to_client"
|
||||
)
|
||||
|
||||
// DialContextFunc dials a backend through the WireGuard tunnel.
|
||||
type DialContextFunc func(ctx context.Context, network, address string) (net.Conn, error)
|
||||
|
||||
// dialTimeoutKey is the context key for a per-request dial timeout.
|
||||
type dialTimeoutKey struct{}
|
||||
|
||||
// WithDialTimeout returns a context carrying a dial timeout that
|
||||
// DialContext wrappers can use to scope the timeout to just the
|
||||
// connection establishment phase.
|
||||
func WithDialTimeout(ctx context.Context, d time.Duration) context.Context {
|
||||
return context.WithValue(ctx, dialTimeoutKey{}, d)
|
||||
}
|
||||
|
||||
// DialTimeoutFromContext returns the dial timeout from the context, if set.
|
||||
func DialTimeoutFromContext(ctx context.Context) (time.Duration, bool) {
|
||||
d, ok := ctx.Value(dialTimeoutKey{}).(time.Duration)
|
||||
return d, ok && d > 0
|
||||
}
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
func TestServiceMode_IsL4(t *testing.T) {
|
||||
tests := []struct {
|
||||
mode ServiceMode
|
||||
want bool
|
||||
}{
|
||||
{ServiceModeHTTP, false},
|
||||
{ServiceModeTCP, true},
|
||||
{ServiceModeUDP, true},
|
||||
{ServiceModeTLS, true},
|
||||
{ServiceMode("unknown"), false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(string(tt.mode), func(t *testing.T) {
|
||||
assert.Equal(t, tt.want, tt.mode.IsL4())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDialTimeoutContext(t *testing.T) {
|
||||
t.Run("round trip", func(t *testing.T) {
|
||||
ctx := WithDialTimeout(context.Background(), 5*time.Second)
|
||||
d, ok := DialTimeoutFromContext(ctx)
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, 5*time.Second, d)
|
||||
})
|
||||
|
||||
t.Run("missing", func(t *testing.T) {
|
||||
_, ok := DialTimeoutFromContext(context.Background())
|
||||
assert.False(t, ok)
|
||||
})
|
||||
|
||||
t.Run("zero returns false", func(t *testing.T) {
|
||||
ctx := WithDialTimeout(context.Background(), 0)
|
||||
_, ok := DialTimeoutFromContext(ctx)
|
||||
assert.False(t, ok, "zero duration should return ok=false")
|
||||
})
|
||||
|
||||
t.Run("negative returns false", func(t *testing.T) {
|
||||
ctx := WithDialTimeout(context.Background(), -1*time.Second)
|
||||
_, ok := DialTimeoutFromContext(ctx)
|
||||
assert.False(t, ok, "negative duration should return ok=false")
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,560 @@
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/netip"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/time/rate"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/accesslog"
|
||||
"github.com/netbirdio/netbird/proxy/internal/netutil"
|
||||
"github.com/netbirdio/netbird/proxy/internal/restrict"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
)
|
||||
|
||||
const (
|
||||
// DefaultSessionTTL is the default idle timeout for UDP sessions before cleanup.
|
||||
DefaultSessionTTL = 30 * time.Second
|
||||
// cleanupInterval is how often the cleaner goroutine runs.
|
||||
cleanupInterval = time.Minute
|
||||
// maxPacketSize is the maximum UDP packet size we'll handle.
|
||||
maxPacketSize = 65535
|
||||
// DefaultMaxSessions is the default cap on concurrent UDP sessions per relay.
|
||||
DefaultMaxSessions = 1024
|
||||
// sessionCreateRate limits new session creation per second.
|
||||
sessionCreateRate = 50
|
||||
// sessionCreateBurst is the burst allowance for session creation.
|
||||
sessionCreateBurst = 100
|
||||
// defaultDialTimeout is the fallback dial timeout for backend connections.
|
||||
defaultDialTimeout = 30 * time.Second
|
||||
)
|
||||
|
||||
// l4Logger sends layer-4 access log entries to the management server.
|
||||
type l4Logger interface {
|
||||
LogL4(entry accesslog.L4Entry)
|
||||
}
|
||||
|
||||
// SessionObserver receives callbacks for UDP session lifecycle events.
|
||||
// All methods must be safe for concurrent use.
|
||||
type SessionObserver interface {
|
||||
UDPSessionStarted(accountID types.AccountID)
|
||||
UDPSessionEnded(accountID types.AccountID)
|
||||
UDPSessionDialError(accountID types.AccountID)
|
||||
UDPSessionRejected(accountID types.AccountID)
|
||||
UDPPacketRelayed(direction types.RelayDirection, bytes int)
|
||||
}
|
||||
|
||||
// clientAddr is a typed key for UDP session lookups.
|
||||
type clientAddr string
|
||||
|
||||
// Relay listens for incoming UDP packets on a dedicated port and
|
||||
// maintains per-client sessions that relay packets to a backend
|
||||
// through the WireGuard tunnel.
|
||||
type Relay struct {
|
||||
logger *log.Entry
|
||||
listener net.PacketConn
|
||||
target string
|
||||
domain string
|
||||
accountID types.AccountID
|
||||
serviceID types.ServiceID
|
||||
dialFunc types.DialContextFunc
|
||||
dialTimeout time.Duration
|
||||
sessionTTL time.Duration
|
||||
maxSessions int
|
||||
filter *restrict.Filter
|
||||
geo restrict.GeoResolver
|
||||
|
||||
mu sync.RWMutex
|
||||
sessions map[clientAddr]*session
|
||||
|
||||
bufPool sync.Pool
|
||||
sessLimiter *rate.Limiter
|
||||
sessWg sync.WaitGroup
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
observer SessionObserver
|
||||
accessLog l4Logger
|
||||
}
|
||||
|
||||
type session struct {
|
||||
backend net.Conn
|
||||
addr net.Addr
|
||||
createdAt time.Time
|
||||
// lastSeen stores the last activity timestamp as unix nanoseconds.
|
||||
lastSeen atomic.Int64
|
||||
cancel context.CancelFunc
|
||||
// bytesIn tracks total bytes received from the client.
|
||||
bytesIn atomic.Int64
|
||||
// bytesOut tracks total bytes sent back to the client.
|
||||
bytesOut atomic.Int64
|
||||
}
|
||||
|
||||
func (s *session) updateLastSeen() {
|
||||
s.lastSeen.Store(time.Now().UnixNano())
|
||||
}
|
||||
|
||||
func (s *session) idleDuration() time.Duration {
|
||||
return time.Since(time.Unix(0, s.lastSeen.Load()))
|
||||
}
|
||||
|
||||
// RelayConfig holds the configuration for a UDP relay.
|
||||
type RelayConfig struct {
|
||||
Logger *log.Entry
|
||||
Listener net.PacketConn
|
||||
Target string
|
||||
Domain string
|
||||
AccountID types.AccountID
|
||||
ServiceID types.ServiceID
|
||||
DialFunc types.DialContextFunc
|
||||
DialTimeout time.Duration
|
||||
SessionTTL time.Duration
|
||||
MaxSessions int
|
||||
AccessLog l4Logger
|
||||
// Filter holds connection-level IP/geo restrictions. Nil means no restrictions.
|
||||
Filter *restrict.Filter
|
||||
// Geo is the geolocation lookup used for country-based restrictions.
|
||||
Geo restrict.GeoResolver
|
||||
}
|
||||
|
||||
// New creates a UDP relay for the given listener and backend target.
|
||||
// MaxSessions caps the number of concurrent sessions; use 0 for DefaultMaxSessions.
|
||||
// DialTimeout controls how long to wait for backend connections; use 0 for default.
|
||||
// SessionTTL is the idle timeout before a session is reaped; use 0 for DefaultSessionTTL.
|
||||
func New(parentCtx context.Context, cfg RelayConfig) *Relay {
|
||||
maxSessions := cfg.MaxSessions
|
||||
dialTimeout := cfg.DialTimeout
|
||||
sessionTTL := cfg.SessionTTL
|
||||
if maxSessions <= 0 {
|
||||
maxSessions = DefaultMaxSessions
|
||||
}
|
||||
if dialTimeout <= 0 {
|
||||
dialTimeout = defaultDialTimeout
|
||||
}
|
||||
if sessionTTL <= 0 {
|
||||
sessionTTL = DefaultSessionTTL
|
||||
}
|
||||
ctx, cancel := context.WithCancel(parentCtx)
|
||||
return &Relay{
|
||||
logger: cfg.Logger,
|
||||
listener: cfg.Listener,
|
||||
target: cfg.Target,
|
||||
domain: cfg.Domain,
|
||||
accountID: cfg.AccountID,
|
||||
serviceID: cfg.ServiceID,
|
||||
accessLog: cfg.AccessLog,
|
||||
dialFunc: cfg.DialFunc,
|
||||
dialTimeout: dialTimeout,
|
||||
sessionTTL: sessionTTL,
|
||||
maxSessions: maxSessions,
|
||||
filter: cfg.Filter,
|
||||
geo: cfg.Geo,
|
||||
sessions: make(map[clientAddr]*session),
|
||||
bufPool: sync.Pool{
|
||||
New: func() any {
|
||||
buf := make([]byte, maxPacketSize)
|
||||
return &buf
|
||||
},
|
||||
},
|
||||
sessLimiter: rate.NewLimiter(sessionCreateRate, sessionCreateBurst),
|
||||
ctx: ctx,
|
||||
cancel: cancel,
|
||||
}
|
||||
}
|
||||
|
||||
// ServiceID returns the service ID associated with this relay.
|
||||
func (r *Relay) ServiceID() types.ServiceID {
|
||||
return r.serviceID
|
||||
}
|
||||
|
||||
// SetObserver sets the session lifecycle observer. Must be called before Serve.
|
||||
func (r *Relay) SetObserver(obs SessionObserver) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.observer = obs
|
||||
}
|
||||
|
||||
// getObserver returns the current session lifecycle observer.
|
||||
func (r *Relay) getObserver() SessionObserver {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
return r.observer
|
||||
}
|
||||
|
||||
// Serve starts the relay loop. It blocks until the context is canceled
|
||||
// or the listener is closed.
|
||||
func (r *Relay) Serve() {
|
||||
go r.cleanupLoop()
|
||||
|
||||
for {
|
||||
bufp := r.bufPool.Get().(*[]byte)
|
||||
buf := *bufp
|
||||
|
||||
n, addr, err := r.listener.ReadFrom(buf)
|
||||
if err != nil {
|
||||
r.bufPool.Put(bufp)
|
||||
if r.ctx.Err() != nil || errors.Is(err, net.ErrClosed) {
|
||||
return
|
||||
}
|
||||
r.logger.Debugf("UDP read: %v", err)
|
||||
continue
|
||||
}
|
||||
|
||||
data := buf[:n]
|
||||
sess, err := r.getOrCreateSession(addr)
|
||||
if err != nil {
|
||||
r.bufPool.Put(bufp)
|
||||
r.logger.Debugf("create UDP session for %s: %v", addr, err)
|
||||
continue
|
||||
}
|
||||
|
||||
sess.updateLastSeen()
|
||||
|
||||
nw, err := sess.backend.Write(data)
|
||||
if err != nil {
|
||||
r.bufPool.Put(bufp)
|
||||
if !netutil.IsExpectedError(err) {
|
||||
r.logger.Debugf("UDP write to backend for %s: %v", addr, err)
|
||||
}
|
||||
r.removeSession(sess)
|
||||
continue
|
||||
}
|
||||
sess.bytesIn.Add(int64(nw))
|
||||
|
||||
if obs := r.getObserver(); obs != nil {
|
||||
obs.UDPPacketRelayed(types.RelayDirectionClientToBackend, nw)
|
||||
}
|
||||
r.bufPool.Put(bufp)
|
||||
}
|
||||
}
|
||||
|
||||
// getOrCreateSession returns an existing session or creates a new one.
|
||||
func (r *Relay) getOrCreateSession(addr net.Addr) (*session, error) {
|
||||
key := clientAddr(addr.String())
|
||||
|
||||
r.mu.RLock()
|
||||
sess, ok := r.sessions[key]
|
||||
r.mu.RUnlock()
|
||||
if ok && sess != nil {
|
||||
return sess, nil
|
||||
}
|
||||
|
||||
// Check before taking the write lock: if the relay is shutting down,
|
||||
// don't create new sessions. This prevents orphaned goroutines when
|
||||
// Serve() processes a packet that was already read before Close().
|
||||
if r.ctx.Err() != nil {
|
||||
return nil, r.ctx.Err()
|
||||
}
|
||||
|
||||
if err := r.checkAccessRestrictions(addr); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
r.mu.Lock()
|
||||
|
||||
if sess, ok = r.sessions[key]; ok && sess != nil {
|
||||
r.mu.Unlock()
|
||||
return sess, nil
|
||||
}
|
||||
if ok {
|
||||
// Another goroutine is dialing for this key, skip.
|
||||
r.mu.Unlock()
|
||||
return nil, fmt.Errorf("session dial in progress for %s", key)
|
||||
}
|
||||
|
||||
if len(r.sessions) >= r.maxSessions {
|
||||
r.mu.Unlock()
|
||||
if obs := r.getObserver(); obs != nil {
|
||||
obs.UDPSessionRejected(r.accountID)
|
||||
}
|
||||
return nil, fmt.Errorf("session limit reached (%d)", r.maxSessions)
|
||||
}
|
||||
|
||||
if !r.sessLimiter.Allow() {
|
||||
r.mu.Unlock()
|
||||
if obs := r.getObserver(); obs != nil {
|
||||
obs.UDPSessionRejected(r.accountID)
|
||||
}
|
||||
return nil, fmt.Errorf("session creation rate limited")
|
||||
}
|
||||
|
||||
// Reserve the slot with a nil session so concurrent callers for the same
|
||||
// key see it exists and wait. Release the lock before dialing.
|
||||
r.sessions[key] = nil
|
||||
r.mu.Unlock()
|
||||
|
||||
dialCtx, dialCancel := context.WithTimeout(r.ctx, r.dialTimeout)
|
||||
backend, err := r.dialFunc(dialCtx, "udp", r.target)
|
||||
dialCancel()
|
||||
if err != nil {
|
||||
r.mu.Lock()
|
||||
delete(r.sessions, key)
|
||||
r.mu.Unlock()
|
||||
if obs := r.getObserver(); obs != nil {
|
||||
obs.UDPSessionDialError(r.accountID)
|
||||
}
|
||||
return nil, fmt.Errorf("dial backend %s: %w", r.target, err)
|
||||
}
|
||||
|
||||
sessCtx, sessCancel := context.WithCancel(r.ctx)
|
||||
sess = &session{
|
||||
backend: backend,
|
||||
addr: addr,
|
||||
createdAt: time.Now(),
|
||||
cancel: sessCancel,
|
||||
}
|
||||
sess.updateLastSeen()
|
||||
|
||||
r.mu.Lock()
|
||||
r.sessions[key] = sess
|
||||
r.mu.Unlock()
|
||||
|
||||
if obs := r.getObserver(); obs != nil {
|
||||
obs.UDPSessionStarted(r.accountID)
|
||||
}
|
||||
|
||||
r.sessWg.Go(func() {
|
||||
r.relayBackendToClient(sessCtx, sess)
|
||||
})
|
||||
|
||||
r.logger.Debugf("UDP session created for %s", addr)
|
||||
return sess, nil
|
||||
}
|
||||
|
||||
func (r *Relay) checkAccessRestrictions(addr net.Addr) error {
|
||||
if r.filter == nil {
|
||||
return nil
|
||||
}
|
||||
clientIP, err := addrFromUDPAddr(addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse client address %s for restriction check: %w", addr, err)
|
||||
}
|
||||
if v := r.filter.Check(clientIP, r.geo); v != restrict.Allow {
|
||||
r.logDeny(clientIP, v)
|
||||
return fmt.Errorf("access restricted for %s", addr)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// relayBackendToClient reads packets from the backend and writes them
|
||||
// back to the client through the public-facing listener.
|
||||
func (r *Relay) relayBackendToClient(ctx context.Context, sess *session) {
|
||||
bufp := r.bufPool.Get().(*[]byte)
|
||||
defer r.bufPool.Put(bufp)
|
||||
defer r.removeSession(sess)
|
||||
|
||||
for ctx.Err() == nil {
|
||||
data, ok := r.readBackendPacket(sess, *bufp)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if data == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
sess.updateLastSeen()
|
||||
|
||||
nw, err := r.listener.WriteTo(data, sess.addr)
|
||||
if err != nil {
|
||||
if !netutil.IsExpectedError(err) {
|
||||
r.logger.Debugf("UDP write to client %s: %v", sess.addr, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
sess.bytesOut.Add(int64(nw))
|
||||
|
||||
if obs := r.getObserver(); obs != nil {
|
||||
obs.UDPPacketRelayed(types.RelayDirectionBackendToClient, nw)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// readBackendPacket reads one packet from the backend with an idle deadline.
|
||||
// Returns (data, true) on success, (nil, true) on idle timeout that should
|
||||
// retry, or (nil, false) when the session should be torn down.
|
||||
func (r *Relay) readBackendPacket(sess *session, buf []byte) ([]byte, bool) {
|
||||
if err := sess.backend.SetReadDeadline(time.Now().Add(r.sessionTTL)); err != nil {
|
||||
r.logger.Debugf("set backend read deadline for %s: %v", sess.addr, err)
|
||||
return nil, false
|
||||
}
|
||||
|
||||
n, err := sess.backend.Read(buf)
|
||||
if err != nil {
|
||||
if netutil.IsTimeout(err) {
|
||||
if sess.idleDuration() > r.sessionTTL {
|
||||
return nil, false
|
||||
}
|
||||
return nil, true
|
||||
}
|
||||
if !netutil.IsExpectedError(err) {
|
||||
r.logger.Debugf("UDP read from backend for %s: %v", sess.addr, err)
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
return buf[:n], true
|
||||
}
|
||||
|
||||
// cleanupLoop periodically removes idle sessions.
|
||||
func (r *Relay) cleanupLoop() {
|
||||
ticker := time.NewTicker(cleanupInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-r.ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
r.cleanupIdleSessions()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// cleanupIdleSessions closes sessions that have been idle for too long.
|
||||
func (r *Relay) cleanupIdleSessions() {
|
||||
var expired []*session
|
||||
|
||||
r.mu.Lock()
|
||||
for key, sess := range r.sessions {
|
||||
if sess == nil {
|
||||
continue
|
||||
}
|
||||
idle := sess.idleDuration()
|
||||
if idle > r.sessionTTL {
|
||||
r.logger.Debugf("UDP session %s idle for %s, closing (client→backend: %d bytes, backend→client: %d bytes)",
|
||||
sess.addr, idle, sess.bytesIn.Load(), sess.bytesOut.Load())
|
||||
delete(r.sessions, key)
|
||||
sess.cancel()
|
||||
if err := sess.backend.Close(); err != nil {
|
||||
r.logger.Debugf("close idle session %s backend: %v", sess.addr, err)
|
||||
}
|
||||
expired = append(expired, sess)
|
||||
}
|
||||
}
|
||||
r.mu.Unlock()
|
||||
|
||||
obs := r.getObserver()
|
||||
for _, sess := range expired {
|
||||
if obs != nil {
|
||||
obs.UDPSessionEnded(r.accountID)
|
||||
}
|
||||
r.logSessionEnd(sess)
|
||||
}
|
||||
}
|
||||
|
||||
// removeSession removes a session from the map if it still matches the
|
||||
// given pointer. This is safe to call concurrently with cleanupIdleSessions
|
||||
// because the identity check prevents double-close when both paths race.
|
||||
func (r *Relay) removeSession(sess *session) {
|
||||
r.mu.Lock()
|
||||
key := clientAddr(sess.addr.String())
|
||||
removed := r.sessions[key] == sess
|
||||
if removed {
|
||||
delete(r.sessions, key)
|
||||
sess.cancel()
|
||||
if err := sess.backend.Close(); err != nil {
|
||||
r.logger.Debugf("close session %s backend: %v", sess.addr, err)
|
||||
}
|
||||
}
|
||||
r.mu.Unlock()
|
||||
|
||||
if removed {
|
||||
r.logger.Debugf("UDP session %s ended (client→backend: %d bytes, backend→client: %d bytes)",
|
||||
sess.addr, sess.bytesIn.Load(), sess.bytesOut.Load())
|
||||
if obs := r.getObserver(); obs != nil {
|
||||
obs.UDPSessionEnded(r.accountID)
|
||||
}
|
||||
r.logSessionEnd(sess)
|
||||
}
|
||||
}
|
||||
|
||||
// logSessionEnd sends an access log entry for a completed UDP session.
|
||||
func (r *Relay) logSessionEnd(sess *session) {
|
||||
if r.accessLog == nil {
|
||||
return
|
||||
}
|
||||
|
||||
var sourceIP netip.Addr
|
||||
if ap, err := netip.ParseAddrPort(sess.addr.String()); err == nil {
|
||||
sourceIP = ap.Addr().Unmap()
|
||||
}
|
||||
|
||||
r.accessLog.LogL4(accesslog.L4Entry{
|
||||
AccountID: r.accountID,
|
||||
ServiceID: r.serviceID,
|
||||
Protocol: accesslog.ProtocolUDP,
|
||||
Host: r.domain,
|
||||
SourceIP: sourceIP,
|
||||
DurationMs: time.Unix(0, sess.lastSeen.Load()).Sub(sess.createdAt).Milliseconds(),
|
||||
BytesUpload: sess.bytesIn.Load(),
|
||||
BytesDownload: sess.bytesOut.Load(),
|
||||
})
|
||||
}
|
||||
|
||||
// logDeny sends an access log entry for a denied UDP packet.
|
||||
func (r *Relay) logDeny(clientIP netip.Addr, verdict restrict.Verdict) {
|
||||
if r.accessLog == nil {
|
||||
return
|
||||
}
|
||||
|
||||
r.accessLog.LogL4(accesslog.L4Entry{
|
||||
AccountID: r.accountID,
|
||||
ServiceID: r.serviceID,
|
||||
Protocol: accesslog.ProtocolUDP,
|
||||
Host: r.domain,
|
||||
SourceIP: clientIP,
|
||||
DenyReason: verdict.String(),
|
||||
})
|
||||
}
|
||||
|
||||
// Close stops the relay, waits for all session goroutines to exit,
|
||||
// and cleans up remaining sessions.
|
||||
func (r *Relay) Close() {
|
||||
r.cancel()
|
||||
if err := r.listener.Close(); err != nil {
|
||||
r.logger.Debugf("close UDP listener: %v", err)
|
||||
}
|
||||
|
||||
var closedSessions []*session
|
||||
r.mu.Lock()
|
||||
for key, sess := range r.sessions {
|
||||
if sess == nil {
|
||||
delete(r.sessions, key)
|
||||
continue
|
||||
}
|
||||
r.logger.Debugf("UDP session %s closed (client→backend: %d bytes, backend→client: %d bytes)",
|
||||
sess.addr, sess.bytesIn.Load(), sess.bytesOut.Load())
|
||||
sess.cancel()
|
||||
if err := sess.backend.Close(); err != nil {
|
||||
r.logger.Debugf("close session %s backend: %v", sess.addr, err)
|
||||
}
|
||||
delete(r.sessions, key)
|
||||
closedSessions = append(closedSessions, sess)
|
||||
}
|
||||
r.mu.Unlock()
|
||||
|
||||
obs := r.getObserver()
|
||||
for _, sess := range closedSessions {
|
||||
if obs != nil {
|
||||
obs.UDPSessionEnded(r.accountID)
|
||||
}
|
||||
r.logSessionEnd(sess)
|
||||
}
|
||||
|
||||
r.sessWg.Wait()
|
||||
}
|
||||
|
||||
// addrFromUDPAddr extracts a netip.Addr from a net.Addr.
|
||||
func addrFromUDPAddr(addr net.Addr) (netip.Addr, error) {
|
||||
ap, err := netip.ParseAddrPort(addr.String())
|
||||
if err != nil {
|
||||
return netip.Addr{}, err
|
||||
}
|
||||
return ap.Addr().Unmap(), nil
|
||||
}
|
||||
@@ -0,0 +1,493 @@
|
||||
package udp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
)
|
||||
|
||||
func TestRelay_BasicPacketExchange(t *testing.T) {
|
||||
// Set up a UDP backend that echoes packets.
|
||||
backend, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer backend.Close()
|
||||
|
||||
go func() {
|
||||
buf := make([]byte, 65535)
|
||||
for {
|
||||
n, addr, err := backend.ReadFrom(buf)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_, _ = backend.WriteTo(buf[:n], addr)
|
||||
}
|
||||
}()
|
||||
|
||||
// Set up the relay's public-facing listener.
|
||||
listener, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer listener.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
backendAddr := backend.LocalAddr().String()
|
||||
|
||||
dialFunc := func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
return net.Dial(network, address)
|
||||
}
|
||||
|
||||
relay := New(ctx, RelayConfig{Logger: logger, Listener: listener, Target: backendAddr, DialFunc: dialFunc})
|
||||
go relay.Serve()
|
||||
defer relay.Close()
|
||||
|
||||
// Create a client and send a packet to the relay.
|
||||
client, err := net.Dial("udp", listener.LocalAddr().String())
|
||||
require.NoError(t, err)
|
||||
defer client.Close()
|
||||
|
||||
testData := []byte("hello UDP relay")
|
||||
_, err = client.Write(testData)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Read the echoed response.
|
||||
if err := client.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
buf := make([]byte, 1024)
|
||||
n, err := client.Read(buf)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testData, buf[:n], "should receive echoed packet")
|
||||
}
|
||||
|
||||
func TestRelay_MultipleClients(t *testing.T) {
|
||||
backend, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer backend.Close()
|
||||
|
||||
go func() {
|
||||
buf := make([]byte, 65535)
|
||||
for {
|
||||
n, addr, err := backend.ReadFrom(buf)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_, _ = backend.WriteTo(buf[:n], addr)
|
||||
}
|
||||
}()
|
||||
|
||||
listener, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer listener.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
dialFunc := func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
return net.Dial(network, address)
|
||||
}
|
||||
|
||||
relay := New(ctx, RelayConfig{Logger: logger, Listener: listener, Target: backend.LocalAddr().String(), DialFunc: dialFunc})
|
||||
go relay.Serve()
|
||||
defer relay.Close()
|
||||
|
||||
// Two clients, each should get their own session.
|
||||
for i, msg := range []string{"client-1", "client-2"} {
|
||||
client, err := net.Dial("udp", listener.LocalAddr().String())
|
||||
require.NoError(t, err, "client %d", i)
|
||||
defer client.Close()
|
||||
|
||||
_, err = client.Write([]byte(msg))
|
||||
require.NoError(t, err)
|
||||
|
||||
if err := client.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
buf := make([]byte, 1024)
|
||||
n, err := client.Read(buf)
|
||||
require.NoError(t, err, "client %d read", i)
|
||||
assert.Equal(t, msg, string(buf[:n]), "client %d should get own echo", i)
|
||||
}
|
||||
|
||||
// Verify two sessions were created.
|
||||
relay.mu.RLock()
|
||||
sessionCount := len(relay.sessions)
|
||||
relay.mu.RUnlock()
|
||||
assert.Equal(t, 2, sessionCount, "should have two sessions")
|
||||
}
|
||||
|
||||
func TestRelay_Close(t *testing.T) {
|
||||
listener, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
dialFunc := func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
return net.Dial(network, address)
|
||||
}
|
||||
|
||||
relay := New(ctx, RelayConfig{Logger: logger, Listener: listener, Target: "127.0.0.1:9999", DialFunc: dialFunc})
|
||||
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
relay.Serve()
|
||||
close(done)
|
||||
}()
|
||||
|
||||
relay.Close()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("Serve did not return after Close")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRelay_SessionCleanup(t *testing.T) {
|
||||
backend, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer backend.Close()
|
||||
|
||||
go func() {
|
||||
buf := make([]byte, 65535)
|
||||
for {
|
||||
n, addr, err := backend.ReadFrom(buf)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_, _ = backend.WriteTo(buf[:n], addr)
|
||||
}
|
||||
}()
|
||||
|
||||
listener, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer listener.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
dialFunc := func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
return net.Dial(network, address)
|
||||
}
|
||||
|
||||
relay := New(ctx, RelayConfig{Logger: logger, Listener: listener, Target: backend.LocalAddr().String(), DialFunc: dialFunc})
|
||||
go relay.Serve()
|
||||
defer relay.Close()
|
||||
|
||||
// Create a session.
|
||||
client, err := net.Dial("udp", listener.LocalAddr().String())
|
||||
require.NoError(t, err)
|
||||
_, err = client.Write([]byte("hello"))
|
||||
require.NoError(t, err)
|
||||
|
||||
if err := client.SetReadDeadline(time.Now().Add(2 * time.Second)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
buf := make([]byte, 1024)
|
||||
_, err = client.Read(buf)
|
||||
require.NoError(t, err)
|
||||
client.Close()
|
||||
|
||||
// Verify session exists.
|
||||
relay.mu.RLock()
|
||||
assert.Equal(t, 1, len(relay.sessions))
|
||||
relay.mu.RUnlock()
|
||||
|
||||
// Make session appear idle by setting lastSeen to the past.
|
||||
relay.mu.Lock()
|
||||
for _, sess := range relay.sessions {
|
||||
sess.lastSeen.Store(time.Now().Add(-2 * DefaultSessionTTL).UnixNano())
|
||||
}
|
||||
relay.mu.Unlock()
|
||||
|
||||
// Trigger cleanup manually.
|
||||
relay.cleanupIdleSessions()
|
||||
|
||||
relay.mu.RLock()
|
||||
assert.Equal(t, 0, len(relay.sessions), "idle sessions should be cleaned up")
|
||||
relay.mu.RUnlock()
|
||||
}
|
||||
|
||||
// TestRelay_CloseAndRecreate verifies that closing a relay and creating a new
|
||||
// one on the same port works cleanly (simulates port mapping modify cycle).
|
||||
func TestRelay_CloseAndRecreate(t *testing.T) {
|
||||
backend, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer backend.Close()
|
||||
|
||||
go func() {
|
||||
buf := make([]byte, 65535)
|
||||
for {
|
||||
n, addr, err := backend.ReadFrom(buf)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_, _ = backend.WriteTo(buf[:n], addr)
|
||||
}
|
||||
}()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
dialFunc := func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
return net.Dial(network, address)
|
||||
}
|
||||
|
||||
// First relay.
|
||||
ln1, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
relay1 := New(ctx, RelayConfig{Logger: logger, Listener: ln1, Target: backend.LocalAddr().String(), DialFunc: dialFunc})
|
||||
go relay1.Serve()
|
||||
|
||||
client1, err := net.Dial("udp", ln1.LocalAddr().String())
|
||||
require.NoError(t, err)
|
||||
_, err = client1.Write([]byte("relay1"))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, client1.SetReadDeadline(time.Now().Add(2*time.Second)))
|
||||
buf := make([]byte, 1024)
|
||||
n, err := client1.Read(buf)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "relay1", string(buf[:n]))
|
||||
client1.Close()
|
||||
|
||||
// Close first relay.
|
||||
relay1.Close()
|
||||
|
||||
// Second relay on same port.
|
||||
port := ln1.LocalAddr().(*net.UDPAddr).Port
|
||||
ln2, err := net.ListenPacket("udp", fmt.Sprintf("127.0.0.1:%d", port))
|
||||
require.NoError(t, err)
|
||||
|
||||
relay2 := New(ctx, RelayConfig{Logger: logger, Listener: ln2, Target: backend.LocalAddr().String(), DialFunc: dialFunc})
|
||||
go relay2.Serve()
|
||||
defer relay2.Close()
|
||||
|
||||
client2, err := net.Dial("udp", ln2.LocalAddr().String())
|
||||
require.NoError(t, err)
|
||||
defer client2.Close()
|
||||
_, err = client2.Write([]byte("relay2"))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, client2.SetReadDeadline(time.Now().Add(2*time.Second)))
|
||||
n, err = client2.Read(buf)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "relay2", string(buf[:n]), "second relay should work on same port")
|
||||
}
|
||||
|
||||
func TestRelay_SessionLimit(t *testing.T) {
|
||||
backend, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer backend.Close()
|
||||
|
||||
go func() {
|
||||
buf := make([]byte, 65535)
|
||||
for {
|
||||
n, addr, err := backend.ReadFrom(buf)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_, _ = backend.WriteTo(buf[:n], addr)
|
||||
}
|
||||
}()
|
||||
|
||||
listener, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer listener.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
dialFunc := func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
return net.Dial(network, address)
|
||||
}
|
||||
|
||||
// Create a relay with a max of 2 sessions.
|
||||
relay := New(ctx, RelayConfig{Logger: logger, Listener: listener, Target: backend.LocalAddr().String(), DialFunc: dialFunc, MaxSessions: 2})
|
||||
go relay.Serve()
|
||||
defer relay.Close()
|
||||
|
||||
// Create 2 clients to fill up the session limit.
|
||||
for i := range 2 {
|
||||
client, err := net.Dial("udp", listener.LocalAddr().String())
|
||||
require.NoError(t, err, "client %d", i)
|
||||
defer client.Close()
|
||||
|
||||
_, err = client.Write([]byte("hello"))
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, client.SetReadDeadline(time.Now().Add(2*time.Second)))
|
||||
buf := make([]byte, 1024)
|
||||
_, err = client.Read(buf)
|
||||
require.NoError(t, err, "client %d should get response", i)
|
||||
}
|
||||
|
||||
relay.mu.RLock()
|
||||
assert.Equal(t, 2, len(relay.sessions), "should have exactly 2 sessions")
|
||||
relay.mu.RUnlock()
|
||||
|
||||
// Third client should get its packet dropped (session creation fails).
|
||||
client3, err := net.Dial("udp", listener.LocalAddr().String())
|
||||
require.NoError(t, err)
|
||||
defer client3.Close()
|
||||
|
||||
_, err = client3.Write([]byte("should be dropped"))
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, client3.SetReadDeadline(time.Now().Add(500*time.Millisecond)))
|
||||
buf := make([]byte, 1024)
|
||||
_, err = client3.Read(buf)
|
||||
assert.Error(t, err, "third client should time out because session was rejected")
|
||||
|
||||
relay.mu.RLock()
|
||||
assert.Equal(t, 2, len(relay.sessions), "session count should not exceed limit")
|
||||
relay.mu.RUnlock()
|
||||
}
|
||||
|
||||
// testObserver records UDP session lifecycle events for test assertions.
|
||||
type testObserver struct {
|
||||
mu sync.Mutex
|
||||
started int
|
||||
ended int
|
||||
rejected int
|
||||
dialErr int
|
||||
packets int
|
||||
bytes int
|
||||
}
|
||||
|
||||
func (o *testObserver) UDPSessionStarted(types.AccountID) { o.mu.Lock(); o.started++; o.mu.Unlock() }
|
||||
func (o *testObserver) UDPSessionEnded(types.AccountID) { o.mu.Lock(); o.ended++; o.mu.Unlock() }
|
||||
func (o *testObserver) UDPSessionDialError(types.AccountID) { o.mu.Lock(); o.dialErr++; o.mu.Unlock() }
|
||||
func (o *testObserver) UDPSessionRejected(types.AccountID) { o.mu.Lock(); o.rejected++; o.mu.Unlock() }
|
||||
func (o *testObserver) UDPPacketRelayed(_ types.RelayDirection, b int) {
|
||||
o.mu.Lock()
|
||||
o.packets++
|
||||
o.bytes += b
|
||||
o.mu.Unlock()
|
||||
}
|
||||
|
||||
func TestRelay_CloseFiresObserverEnded(t *testing.T) {
|
||||
backend, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer backend.Close()
|
||||
|
||||
go func() {
|
||||
buf := make([]byte, 65535)
|
||||
for {
|
||||
n, addr, err := backend.ReadFrom(buf)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_, _ = backend.WriteTo(buf[:n], addr)
|
||||
}
|
||||
}()
|
||||
|
||||
listener, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer listener.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
dialFunc := func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
return net.Dial(network, address)
|
||||
}
|
||||
|
||||
obs := &testObserver{}
|
||||
relay := New(ctx, RelayConfig{Logger: logger, Listener: listener, Target: backend.LocalAddr().String(), AccountID: "test-acct", DialFunc: dialFunc})
|
||||
relay.SetObserver(obs)
|
||||
go relay.Serve()
|
||||
|
||||
// Create two sessions.
|
||||
for i := range 2 {
|
||||
client, err := net.Dial("udp", listener.LocalAddr().String())
|
||||
require.NoError(t, err, "client %d", i)
|
||||
|
||||
_, err = client.Write([]byte("hello"))
|
||||
require.NoError(t, err)
|
||||
|
||||
require.NoError(t, client.SetReadDeadline(time.Now().Add(2*time.Second)))
|
||||
buf := make([]byte, 1024)
|
||||
_, err = client.Read(buf)
|
||||
require.NoError(t, err)
|
||||
client.Close()
|
||||
}
|
||||
|
||||
obs.mu.Lock()
|
||||
assert.Equal(t, 2, obs.started, "should have 2 started events")
|
||||
obs.mu.Unlock()
|
||||
|
||||
// Close should fire UDPSessionEnded for all remaining sessions.
|
||||
relay.Close()
|
||||
|
||||
obs.mu.Lock()
|
||||
assert.Equal(t, 2, obs.ended, "Close should fire UDPSessionEnded for each session")
|
||||
obs.mu.Unlock()
|
||||
}
|
||||
|
||||
func TestRelay_SessionRateLimit(t *testing.T) {
|
||||
backend, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer backend.Close()
|
||||
|
||||
go func() {
|
||||
buf := make([]byte, 65535)
|
||||
for {
|
||||
n, addr, err := backend.ReadFrom(buf)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
_, _ = backend.WriteTo(buf[:n], addr)
|
||||
}
|
||||
}()
|
||||
|
||||
listener, err := net.ListenPacket("udp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer listener.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
logger := log.NewEntry(log.StandardLogger())
|
||||
dialFunc := func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
return net.Dial(network, address)
|
||||
}
|
||||
|
||||
obs := &testObserver{}
|
||||
// High max sessions (1000) but the relay uses a rate limiter internally
|
||||
// (default: 50/s burst 100). We exhaust the burst by creating sessions
|
||||
// rapidly, then verify that subsequent creates are rejected.
|
||||
relay := New(ctx, RelayConfig{Logger: logger, Listener: listener, Target: backend.LocalAddr().String(), AccountID: "test-acct", DialFunc: dialFunc, MaxSessions: 1000})
|
||||
relay.SetObserver(obs)
|
||||
go relay.Serve()
|
||||
defer relay.Close()
|
||||
|
||||
// Exhaust the burst by calling getOrCreateSession directly with
|
||||
// synthetic addresses. This is faster than real UDP round-trips.
|
||||
for i := range sessionCreateBurst + 20 {
|
||||
addr := &net.UDPAddr{IP: net.IPv4(10, 0, byte(i/256), byte(i%256)), Port: 10000 + i}
|
||||
_, _ = relay.getOrCreateSession(addr)
|
||||
}
|
||||
|
||||
obs.mu.Lock()
|
||||
rejected := obs.rejected
|
||||
obs.mu.Unlock()
|
||||
|
||||
assert.Greater(t, rejected, 0, "some sessions should be rate-limited")
|
||||
}
|
||||
@@ -209,7 +209,7 @@ func (m *testProxyManager) Disconnect(_ context.Context, _ string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *testProxyManager) Heartbeat(_ context.Context, _ string) error {
|
||||
func (m *testProxyManager) Heartbeat(_ context.Context, _, _, _ string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -221,6 +221,10 @@ func (m *testProxyManager) GetActiveClusterAddressesForAccount(_ context.Context
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (m *testProxyManager) GetActiveClusters(_ context.Context) ([]nbproxy.Cluster, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (m *testProxyManager) CleanupStale(_ context.Context, _ time.Duration) error {
|
||||
return nil
|
||||
}
|
||||
@@ -264,6 +268,14 @@ func (c *testProxyController) GetProxiesForCluster(_ string) []string {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *testProxyController) ClusterSupportsCustomPorts(_ string) *bool {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *testProxyController) ClusterRequireSubdomain(_ string) *bool {
|
||||
return nil
|
||||
}
|
||||
|
||||
// storeBackedServiceManager reads directly from the real store.
|
||||
type storeBackedServiceManager struct {
|
||||
store store.Store
|
||||
@@ -344,6 +356,10 @@ func (m *storeBackedServiceManager) GetServiceByDomain(ctx context.Context, doma
|
||||
return m.store.GetServiceByDomain(ctx, domain)
|
||||
}
|
||||
|
||||
func (m *storeBackedServiceManager) GetActiveClusters(_ context.Context, _, _ string) ([]nbproxy.Cluster, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func strPtr(s string) *string {
|
||||
return &s
|
||||
}
|
||||
@@ -511,7 +527,7 @@ func TestIntegration_ProxyConnection_ReconnectDoesNotDuplicateState(t *testing.T
|
||||
logger := log.New()
|
||||
logger.SetLevel(log.WarnLevel)
|
||||
|
||||
authMw := auth.NewMiddleware(logger, nil)
|
||||
authMw := auth.NewMiddleware(logger, nil, nil)
|
||||
proxyHandler := proxy.NewReverseProxy(nil, "auto", nil, logger)
|
||||
|
||||
clusterAddress := "test.proxy.io"
|
||||
@@ -530,15 +546,16 @@ func TestIntegration_ProxyConnection_ReconnectDoesNotDuplicateState(t *testing.T
|
||||
nil,
|
||||
"",
|
||||
0,
|
||||
mapping.GetAccountId(),
|
||||
mapping.GetId(),
|
||||
proxytypes.AccountID(mapping.GetAccountId()),
|
||||
proxytypes.ServiceID(mapping.GetId()),
|
||||
nil,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Apply to real proxy (idempotent)
|
||||
proxyHandler.AddMapping(proxy.Mapping{
|
||||
Host: mapping.GetDomain(),
|
||||
ID: mapping.GetId(),
|
||||
ID: proxytypes.ServiceID(mapping.GetId()),
|
||||
AccountID: proxytypes.AccountID(mapping.GetAccountId()),
|
||||
})
|
||||
}
|
||||
|
||||
+879
-62
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user