From 3f40acfcac8671e9d05ab4d67b2d8a30f16745d2 Mon Sep 17 00:00:00 2001 From: Daniel Lavrushin Date: Wed, 29 Apr 2026 23:26:35 +0200 Subject: [PATCH] feat: enhance validation for Queue.Mark and add corresponding unit tests; refactor bypass mark application in dialers --- src/config/methods.go | 5 ++++ src/config/methods_test.go | 22 ++++++++++++++ .../src/components/dashboard/MetricsCards.tsx | 4 +-- src/http/ws/connections.go | 7 ++--- src/http/ws/discovery.go | 5 +--- src/log/connection.go | 1 - src/log/discovery.go | 1 - src/socks5/client.go | 30 +++++++++++-------- src/tproxy/listener.go | 15 +--------- 9 files changed, 50 insertions(+), 40 deletions(-) diff --git a/src/config/methods.go b/src/config/methods.go index bf436f8a..e944417a 100644 --- a/src/config/methods.go +++ b/src/config/methods.go @@ -276,6 +276,11 @@ func (c *Config) Validate() error { return fmt.Errorf("mark value 0x%x is too high for auto-derived discovery marks", c.Queue.Mark) } + const perSetReachableBits uint32 = 0x17FFF + if c.Queue.Mark != 0 && uint32(c.Queue.Mark)&^perSetReachableBits == 0 { + return fmt.Errorf("mark value 0x%x conflicts with per-set mark bits {0-14, 16}; bypass rule would catch TPROXY-redirected traffic. Use a value with at least one bit in {15, 17-31} (default 0x8000 has bit 15)", c.Queue.Mark) + } + c.System.Checker.DiscoveryFlowMark = c.DiscoveryFlowMark() c.System.Checker.DiscoveryInjectedMark = c.DiscoveryInjectedMark() diff --git a/src/config/methods_test.go b/src/config/methods_test.go index fa6d0261..6c141cde 100644 --- a/src/config/methods_test.go +++ b/src/config/methods_test.go @@ -102,6 +102,28 @@ func TestValidate(t *testing.T) { } }) + t.Run("queue mark inside per-set space fails", func(t *testing.T) { + cases := []uint{0x4000, 0x100, 0x10000, 0x12345, 0x17DFF} + for _, m := range cases { + cfg := NewConfig() + cfg.Queue.Mark = m + if err := cfg.Validate(); err == nil { + t.Errorf("expected error for Queue.Mark=%#x (collides with per-set range)", m) + } + } + }) + + t.Run("queue mark outside per-set space passes", func(t *testing.T) { + cases := []uint{0x8000, 0x18000, 0x20000, 0x80000000} + for _, m := range cases { + cfg := NewConfig() + cfg.Queue.Mark = m + if err := cfg.Validate(); err != nil { + t.Errorf("Queue.Mark=%#x should pass, got: %v", m, err) + } + } + }) + t.Run("queue num out of range", func(t *testing.T) { cfg := NewConfig() cfg.Queue.StartNum = -1 diff --git a/src/http/ui/src/components/dashboard/MetricsCards.tsx b/src/http/ui/src/components/dashboard/MetricsCards.tsx index 8c656424..5a873719 100644 --- a/src/http/ui/src/components/dashboard/MetricsCards.tsx +++ b/src/http/ui/src/components/dashboard/MetricsCards.tsx @@ -11,8 +11,8 @@ interface MetricsCardsProps { export const MetricsCards = ({ metrics }: MetricsCardsProps) => { const { t } = useTranslation(); const targetRate = - metrics.packets_processed > 0 - ? ((metrics.targeted_connections / metrics.packets_processed) * 100).toFixed(1) + metrics.total_connections > 0 + ? ((metrics.targeted_connections / metrics.total_connections) * 100).toFixed(1) : "0.0"; const isIdle = metrics.rst_dropped === 0; diff --git a/src/http/ws/connections.go b/src/http/ws/connections.go index 73e223af..f6ad0aab 100644 --- a/src/http/ws/connections.go +++ b/src/http/ws/connections.go @@ -24,8 +24,8 @@ func HandleConnectionsWebSocket(w http.ResponseWriter, r *http.Request) { ch, snapshot := hub.Subscribe() defer hub.Unsubscribe(ch) - conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) for _, msg := range snapshot { + conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) if err := conn.WriteMessage(websocket.TextMessage, []byte(msg)); err != nil { return } @@ -46,10 +46,7 @@ func HandleConnectionsWebSocket(w http.ResponseWriter, r *http.Request) { for { select { - case msg, ok := <-ch: - if !ok { - return - } + case msg := <-ch: conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) if err := conn.WriteMessage(websocket.TextMessage, []byte(msg)); err != nil { return diff --git a/src/http/ws/discovery.go b/src/http/ws/discovery.go index aab95222..40bad28c 100644 --- a/src/http/ws/discovery.go +++ b/src/http/ws/discovery.go @@ -36,10 +36,7 @@ func HandleDiscoveryWebSocket(w http.ResponseWriter, r *http.Request) { for { select { - case msg, ok := <-ch: - if !ok { - return - } + case msg := <-ch: conn.SetWriteDeadline(time.Now().Add(10 * time.Second)) if err := conn.WriteMessage(websocket.TextMessage, []byte(msg)); err != nil { return diff --git a/src/log/connection.go b/src/log/connection.go index 8a98f132..c921385f 100644 --- a/src/log/connection.go +++ b/src/log/connection.go @@ -55,7 +55,6 @@ func (h *ConnectionHub) Unsubscribe(ch chan string) { for i, l := range h.listeners { if l == ch { h.listeners = append(h.listeners[:i], h.listeners[i+1:]...) - close(ch) return } } diff --git a/src/log/discovery.go b/src/log/discovery.go index 94df9800..ab41f9e7 100644 --- a/src/log/discovery.go +++ b/src/log/discovery.go @@ -38,7 +38,6 @@ func (h *DiscoveryLogHub) Unsubscribe(ch chan string) { for i, l := range h.listeners { if l == ch { h.listeners = append(h.listeners[:i], h.listeners[i+1:]...) - close(ch) return } } diff --git a/src/socks5/client.go b/src/socks5/client.go index c989ad9b..d65668a3 100644 --- a/src/socks5/client.go +++ b/src/socks5/client.go @@ -25,6 +25,22 @@ type ClientConfig struct { BypassMark uint32 } +func ApplyBypassMark(d *net.Dialer, mark uint32) { + if mark == 0 { + return + } + d.Control = func(network, address string, c syscall.RawConn) error { + var sockErr error + err := c.Control(func(fd uintptr) { + sockErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_MARK, int(mark)) + }) + if err != nil { + return err + } + return sockErr + } +} + func DialUpstream(ctx context.Context, cfg ClientConfig, targetHost string, targetPort int) (net.Conn, error) { if cfg.Host == "" || cfg.Port < 1 || cfg.Port > 65535 { return nil, fmt.Errorf("invalid upstream config") @@ -39,19 +55,7 @@ func DialUpstream(ctx context.Context, cfg ClientConfig, targetHost string, targ } d := net.Dialer{Timeout: timeout} - if cfg.BypassMark != 0 { - mark := cfg.BypassMark - d.Control = func(network, address string, c syscall.RawConn) error { - var sockErr error - err := c.Control(func(fd uintptr) { - sockErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_MARK, int(mark)) - }) - if err != nil { - return err - } - return sockErr - } - } + ApplyBypassMark(&d, cfg.BypassMark) addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port)) conn, err := d.DialContext(ctx, "tcp", addr) if err != nil { diff --git a/src/tproxy/listener.go b/src/tproxy/listener.go index 1dd557ab..78123eda 100644 --- a/src/tproxy/listener.go +++ b/src/tproxy/listener.go @@ -18,20 +18,7 @@ import ( func markedDialer(timeout time.Duration, bypassMark uint32) net.Dialer { d := net.Dialer{Timeout: timeout} - if bypassMark == 0 { - return d - } - mark := bypassMark - d.Control = func(network, address string, c syscall.RawConn) error { - var sockErr error - err := c.Control(func(fd uintptr) { - sockErr = unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_MARK, int(mark)) - }) - if err != nil { - return err - } - return sockErr - } + socks5.ApplyBypassMark(&d, bypassMark) return d }