From 8b92a966a8fd94efe9e382944325a8c22c043f27 Mon Sep 17 00:00:00 2001 From: Rico Date: Tue, 6 Jan 2026 09:03:36 +0100 Subject: [PATCH] feat: better handling of the real ip behind a reverse proxy --- internal/web/broadcastmanager.go | 6 ++--- internal/web/client.go | 24 +++++++++++++++--- internal/web/server.go | 43 +++++++++++++++++++------------- 3 files changed, 49 insertions(+), 24 deletions(-) diff --git a/internal/web/broadcastmanager.go b/internal/web/broadcastmanager.go index 18382be..cebc82d 100644 --- a/internal/web/broadcastmanager.go +++ b/internal/web/broadcastmanager.go @@ -81,7 +81,7 @@ func (bm *BroadcastManager) GetSkippedCerts() map[string]uint64 { skippedCerts := make(map[string]uint64, len(bm.clients)) for _, c := range bm.clients { - skippedCerts[c.name] = c.skippedCerts + skippedCerts[c.Name()] = c.skippedCerts } return skippedCerts @@ -108,7 +108,7 @@ func (bm *BroadcastManager) broadcaster() { case SubTypeDomain: data = dataDomain default: - log.Printf("Unknown subscription type '%d' for client '%s'. Skipping this client!\n", c.subType, c.name) + log.Printf("Unknown subscription type '%d' for client '%s'. Skipping this client!\n", c.subType, c.Name()) continue } @@ -118,7 +118,7 @@ func (bm *BroadcastManager) broadcaster() { // Default case is executed if the client's broadcast channel is full. c.skippedCerts++ if c.skippedCerts%1000 == 1 { - log.Printf("Not providing client '%s' with cert because client's buffer is full. The client can't keep up. Skipped certs: %d\n", c.name, c.skippedCerts) + log.Printf("Not providing client '%s' with cert because client's buffer is full. The client can't keep up. Skipped certs: %d\n", c.Name(), c.skippedCerts) } } } diff --git a/internal/web/client.go b/internal/web/client.go index 464d6c0..de97c13 100644 --- a/internal/web/client.go +++ b/internal/web/client.go @@ -1,7 +1,9 @@ package web import ( + "fmt" "log" + "net" "strings" "time" @@ -18,22 +20,35 @@ type SubscriptionType int // client represents a single client's connection to the server. type client struct { - conn *websocket.Conn + conn *websocket.Conn + // when behind a proxy, this holds the real IP of the client + realIP net.IP + connectionIP net.IP + userAgent string broadcastChan chan []byte - name string subType SubscriptionType skippedCerts uint64 } -func newClient(conn *websocket.Conn, subType SubscriptionType, name string, certBufferSize int) *client { +func newClient(conn *websocket.Conn, subType SubscriptionType, realIP, connectionIP net.IP, certBufferSize int, userAgent string) *client { return &client{ conn: conn, broadcastChan: make(chan []byte, certBufferSize), - name: name, + realIP: realIP, + connectionIP: connectionIP, + userAgent: userAgent, subType: subType, } } +func (c *client) Name() string { + if c.connectionIP.Equal(c.realIP) { + return fmt.Sprintf("%s (Real IP: %s)", c.connectionIP, c.realIP) + } + + return fmt.Sprintf("%s", c.connectionIP) +} + // Each client has a broadcastHandler that runs in the background and sends out the broadcast messages to the client. func (c *client) broadcastHandler() { writeWait := 60 * time.Second @@ -128,6 +143,7 @@ func (c *client) listenWebsocket() { log.Printf("Connection to client lost: %v\n", c.conn.RemoteAddr()) } + // TODO the client's IP address is not correctly shown here log.Printf("Disconnecting client %v!\n", c.conn.RemoteAddr()) break diff --git a/internal/web/server.go b/internal/web/server.go index 50a55b6..9ce5286 100644 --- a/internal/web/server.go +++ b/internal/web/server.go @@ -3,7 +3,6 @@ package web import ( "context" "crypto/tls" - "fmt" "io" "log" "net" @@ -119,7 +118,7 @@ func initFullWebsocket(w http.ResponseWriter, r *http.Request) { return } - setupClient(connection, SubTypeFull, r.RemoteAddr) + setupClient(connection, SubTypeFull, r.RemoteAddr, r.UserAgent()) } // initLiteWebsocket is called when a client connects to the / endpoint. @@ -131,7 +130,7 @@ func initLiteWebsocket(w http.ResponseWriter, r *http.Request) { return } - setupClient(connection, SubTypeLite, r.RemoteAddr) + setupClient(connection, SubTypeLite, r.RemoteAddr, r.UserAgent()) } // initDomainWebsocket is called when a client connects to the /domains-only endpoint. @@ -143,21 +142,12 @@ func initDomainWebsocket(w http.ResponseWriter, r *http.Request) { return } - setupClient(connection, SubTypeDomain, r.RemoteAddr) + setupClient(connection, SubTypeDomain, r.RemoteAddr, r.UserAgent()) } // upgradeConnection upgrades the connection to a websocket and returns the connection. func upgradeConnection(w http.ResponseWriter, r *http.Request) (*websocket.Conn, error) { - var remoteAddr string - - xForwardedFor := r.Header.Get("X-Forwarded-For") - if xForwardedFor != "" { - remoteAddr = fmt.Sprintf("'%s' (X-Forwarded-For: '%s')", r.RemoteAddr, xForwardedFor) - } else { - remoteAddr = fmt.Sprintf("'%s'", r.RemoteAddr) - } - - log.Printf("Starting new websocket for %s - %s\n", remoteAddr, r.URL) + log.Printf("Starting new websocket for %s - %s\n", r.RemoteAddr, r.URL) connection, err := upgrader.Upgrade(w, r, nil) if err != nil { @@ -166,7 +156,7 @@ func upgradeConnection(w http.ResponseWriter, r *http.Request) (*websocket.Conn, defaultCloseHandler := connection.CloseHandler() connection.SetCloseHandler(func(code int, text string) error { - log.Printf("Stopping websocket for %s - %s\n", remoteAddr, r.URL) + log.Printf("Stopping websocket for %s - %s\n", r.RemoteAddr, r.URL) return defaultCloseHandler(code, text) }) @@ -174,8 +164,27 @@ func upgradeConnection(w http.ResponseWriter, r *http.Request) (*websocket.Conn, } // setupClient initializes a client struct and starts the broadcastHandler and websocket listener. -func setupClient(connection *websocket.Conn, subscriptionType SubscriptionType, name string) { - c := newClient(connection, subscriptionType, name, config.AppConfig.General.BufferSizes.Websocket) +func setupClient(connection *websocket.Conn, subscriptionType SubscriptionType, remoteAddr, userAgent string) { + connectionHost, _, connectionSplitErr := net.SplitHostPort(connection.RemoteAddr().String()) + if connectionSplitErr != nil { + log.Println("Error while trying to parse remote address:", connection.RemoteAddr().String()) + } + connectionIP := net.ParseIP(connectionHost) + + var realIP net.IP + if config.AppConfig.Webserver.RealIP { + // RealIP is either the real IP connecting to the server or the IP parsed from the X-Forwarded-For or X-Real-IP header if present. + realHost, _, realSplitErr := net.SplitHostPort(remoteAddr) + if realSplitErr != nil { + log.Println("Error while trying to parse remote address:", remoteAddr) + } + realIP = net.ParseIP(realHost) + } else { + // In case the RealIP option is not configured, just use the connection IP + realIP = connectionIP + } + + c := newClient(connection, subscriptionType, realIP, connectionIP, config.AppConfig.General.BufferSizes.Websocket, userAgent) go c.broadcastHandler() go c.listenWebsocket()