Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 31 additions & 1 deletion agent/server/snykbroker/reflector.go
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ type RegistrationReflector struct {
config config.AgentConfig
lastTrafficTime atomic.Int64
lastStartupTime atomic.Int64
lastTunnelTime atomic.Int64
wsProxy *WebSocketProxy
}

Expand Down Expand Up @@ -137,6 +138,29 @@ func (rr *RegistrationReflector) LastStartupTime() time.Time {
return time.UnixMilli(rr.lastStartupTime.Load())
}

// RecordTunnelActivity marks that the broker server just sent something down
// the websocket tunnel. It is kept apart from lastTrafficTime: a heartbeat
// proves the tunnel is alive, not that anything was relayed.
func (rr *RegistrationReflector) RecordTunnelActivity() {
rr.lastTunnelTime.Store(time.Now().UnixMilli())
}

// LastActivityTime is the latest of the last relayed request, the last frame
// from the broker server, and the last (re)start: the idle watchdog's clock.
//
// Relayed requests alone can't tell a dead tunnel from a quiet one: the
// broker server hands a token's requests to only its newest client, so every
// other replica sharing the token looks idle. Frames cover that when the
// tunnel runs through this reflector; in "traffic" mode it doesn't, and
// relayed requests are the only signal.
func (rr *RegistrationReflector) LastActivityTime() time.Time {
return time.UnixMilli(max(
rr.lastTrafficTime.Load(),
rr.lastTunnelTime.Load(),
rr.lastStartupTime.Load(),
))
}

// SetOnWSTunnelClose sets a callback invoked when a WebSocket tunnel closes.
// Used by the relay instance manager to trigger a broker restart.
func (rr *RegistrationReflector) SetOnWSTunnelClose(fn func()) {
Expand Down Expand Up @@ -396,7 +420,13 @@ func (rr *RegistrationReflector) ServeHTTP(w http.ResponseWriter, r *http.Reques
// Check if this is a WebSocket upgrade request
if rr.config.ReflectorWebSocketUpgrade && IsWebSocketUpgrade(r) {
rr.logger.Debug("Detected WebSocket upgrade request, using WebSocket proxy")
if err := rr.wsProxy.Proxy(w, r, entry.TargetURI); err != nil {
// Only the default entry is the broker's own tunnel to the server; a
// websocket to a customer origin says nothing about that tunnel.
var onActivity func()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

suuuuuuper nit but it was kind of confusing that there was this default lambda, not sure how i would implement it better though tbh

if entry.isDefault {
onActivity = rr.RecordTunnelActivity
}
if err := rr.wsProxy.Proxy(w, r, entry.TargetURI, onActivity); err != nil {
rr.logger.Error("WebSocket proxy failed", zap.Error(err))
}
return
Expand Down
12 changes: 6 additions & 6 deletions agent/server/snykbroker/relay_instance_manager.go
Original file line number Diff line number Diff line change
Expand Up @@ -308,14 +308,14 @@ func (r *relayInstanceManager) shouldRestart() (bool, string) {
if r.config.RelayIdleTimeout == 0 || r.reflector == nil {
return false, ""
}
if !r.config.HttpRelayReflectorMode.ReflectsTraffic() {
// Each reflecting mode gives the watchdog a signal: relayed requests in
// "traffic", frames on the broker's tunnel in "registration", both in
// "all". Only "disabled" leaves it blind.
mode := r.config.HttpRelayReflectorMode
if !mode.ReflectsTraffic() && !mode.ReflectsRegistration() {
return false, ""
}
lastActivity := r.reflector.LastTrafficTime()
if startup := r.reflector.LastStartupTime(); startup.After(lastActivity) {
lastActivity = startup
}
if time.Since(lastActivity) >= r.config.RelayIdleTimeout {
if time.Since(r.reflector.LastActivityTime()) >= r.config.RelayIdleTimeout {

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agreed — fixed in the latest commit. shouldRestart now runs whenever the reflector reflects traffic or registration, so registration mode (broker tunnel through the reflector, no relayed requests) is watched via tunnel frames, and only disabled is excluded since the agent sees neither signal there. TestShouldRestartCountsTunnelActivity now also covers registration (stale → restart, tunnel frame → no restart) and disabled (never restarts). Note this does enable the watchdog in registration mode where it was previously off; it only fires if the server sends nothing for RELAY_IDLE_TIMEOUT, which with a ~30s server heartbeat means the tunnel is dead.

return true, "idle_timeout"
}
return false, ""
Expand Down
46 changes: 46 additions & 0 deletions agent/server/snykbroker/relay_instance_manager_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -486,6 +486,52 @@ func TestIdleTimeoutDetectsIdleReflector(t *testing.T) {
require.NoError(t, err)
}

// Only the newest client for a token gets relayed requests, so a healthy
// replica can go a long time without traffic. Frames from the broker server
// keep it out of the idle watchdog; a tunnel that has gone quiet does not.
func TestShouldRestartCountsTunnelActivity(t *testing.T) {
cfg := config.NewAgentEnvConfig()
cfg.HttpRelayReflectorMode = config.RelayReflectorAllTraffic
cfg.RelayIdleTimeout = time.Minute

rr := NewRegistrationReflector(RegistrationReflectorParams{
Logger: zap.NewNop(),
Config: cfg,
})
r := &relayInstanceManager{config: cfg, reflector: rr}

stale := time.Now().Add(-2 * time.Minute).UnixMilli()
rr.lastStartupTime.Store(stale)
rr.lastTrafficTime.Store(stale)

restart, reason := r.shouldRestart()
require.True(t, restart, "no traffic, no startup and no tunnel frames within the window is idle")
require.Equal(t, "idle_timeout", reason)

rr.RecordTunnelActivity()
restart, _ = r.shouldRestart()
require.False(t, restart, "a frame from the broker server shows the tunnel is alive")

rr.lastTunnelTime.Store(stale)
restart, _ = r.shouldRestart()
require.True(t, restart, "a tunnel that has gone quiet is idle again")

// "registration" routes the broker's tunnel through the reflector but no
// relayed requests, so tunnel frames are its only signal.
r.config.HttpRelayReflectorMode = config.RelayReflectorRegistrationOnly
restart, _ = r.shouldRestart()
require.True(t, restart, "registration mode watches the tunnel too")
rr.RecordTunnelActivity()
restart, _ = r.shouldRestart()
require.False(t, restart, "registration mode counts tunnel frames")

// With the reflector off the agent sees neither, so it must not guess.
rr.lastTunnelTime.Store(stale)
r.config.HttpRelayReflectorMode = config.RelayReflectorDisabled
restart, _ = r.shouldRestart()
require.False(t, restart, "disabled mode has no signal and never restarts")
}

// A broker with nothing to relay looks identical, idle-wise, to one whose
// tunnel died silently: shouldRestart() can only see "no activity recorded."
// A restart must reset the idle clock (LastStartupTime) so a healthy-but-idle
Expand Down
22 changes: 14 additions & 8 deletions agent/server/snykbroker/ws_proxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,10 @@ func (wp *WebSocketProxy) ActiveConnections() int32 {

// Proxy handles a WebSocket upgrade request by establishing a tunnel to the target.
// It hijacks the client connection and proxies bidirectionally.
func (wp *WebSocketProxy) Proxy(w http.ResponseWriter, r *http.Request, targetURI string) error {
//
// onActivity, if set, is called whenever bytes arrive from the target. Bytes
// sent to the target don't count: writing into a dead tunnel can still succeed.
func (wp *WebSocketProxy) Proxy(w http.ResponseWriter, r *http.Request, targetURI string, onActivity func()) error {
targetURL, err := url.Parse(targetURI)
if err != nil {
return fmt.Errorf("invalid target URI: %w", err)
Expand Down Expand Up @@ -91,7 +94,7 @@ func (wp *WebSocketProxy) Proxy(w http.ResponseWriter, r *http.Request, targetUR
}

// Run the bidirectional tunnel
wp.runTunnel(clientConn, targetConn, targetAddr)
wp.runTunnel(clientConn, targetConn, targetAddr, onActivity)
return nil
}

Expand Down Expand Up @@ -239,7 +242,7 @@ func (wp *WebSocketProxy) hijackConnection(w http.ResponseWriter) (net.Conn, err
return conn, nil
}

func (wp *WebSocketProxy) runTunnel(clientConn, targetConn net.Conn, targetAddr string) {
func (wp *WebSocketProxy) runTunnel(clientConn, targetConn net.Conn, targetAddr string, onActivity func()) {
start := time.Now()
wp.activeConnections.Add(1)

Expand All @@ -256,21 +259,21 @@ func (wp *WebSocketProxy) runTunnel(clientConn, targetConn net.Conn, targetAddr

done := make(chan struct{}, 2)

copy := func(dst, src net.Conn, direction string) {
copy := func(dst, src net.Conn, direction string, onRead func()) {
defer func() { done <- struct{}{} }()
wp.copyWithIdleTimeout(dst, src, direction)
wp.copyWithIdleTimeout(dst, src, direction, onRead)
}

go copy(clientConn, targetConn, "target->client")
go copy(targetConn, clientConn, "client->target")
go copy(clientConn, targetConn, "target->client", onActivity)
go copy(targetConn, clientConn, "client->target", nil)

<-done
clientConn.Close()
targetConn.Close()
<-done
}

func (wp *WebSocketProxy) copyWithIdleTimeout(dst, src net.Conn, direction string) {
func (wp *WebSocketProxy) copyWithIdleTimeout(dst, src net.Conn, direction string, onRead func()) {
buf := make([]byte, 32*1024)
isFirstRead := true
for {
Expand Down Expand Up @@ -298,6 +301,9 @@ func (wp *WebSocketProxy) copyWithIdleTimeout(dst, src net.Conn, direction strin
}

if n > 0 {
if onRead != nil {
onRead()
}
dst.SetWriteDeadline(time.Now().Add(wp.HandshakeTimeout))
if _, writeErr := dst.Write(buf[:n]); writeErr != nil {
return
Expand Down
82 changes: 82 additions & 0 deletions agent/server/snykbroker/ws_proxy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -479,6 +479,88 @@ func TestWebSocketProxyCallbacks(t *testing.T) {
assert.True(t, closedDuration > 0, "duration should be positive")
}

// Frames from the broker server on the default entry (the broker's own
// tunnel) count as tunnel activity. Our own writes don't, since writing into a
// dead tunnel can still succeed, and neither does a websocket to a customer
// origin, which says nothing about the broker's tunnel.
func TestReflectorRecordsTunnelActivityForBrokerTunnelOnly(t *testing.T) {
logger := newTestLogger(t)

// A websocket server that sends one frame when told to.
newTarget := func(t *testing.T) (*httptest.Server, chan struct{}) {
send := make(chan struct{})
upgrader := websocket.Upgrader{CheckOrigin: func(r *http.Request) bool { return true }}
router := mux.NewRouter()
router.HandleFunc("/ws", func(w http.ResponseWriter, r *http.Request) {
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
return
}
defer conn.Close()
go func() {
for {
if _, _, err := conn.ReadMessage(); err != nil {
return
}
}
}()
<-send
_ = conn.WriteMessage(websocket.TextMessage, []byte("primus::ping::1"))
time.Sleep(200 * time.Millisecond)
})
server := httptest.NewServer(router)
t.Cleanup(server.Close)
return server, send
}

rr := NewRegistrationReflector(RegistrationReflectorParams{
Logger: logger,
Config: config.AgentConfig{ReflectorWebSocketUpgrade: true},
})
reflectorRouter := mux.NewRouter()
rr.RegisterRoutes(reflectorRouter)
reflectorServer := httptest.NewServer(reflectorRouter)
defer reflectorServer.Close()

// Dial through the reflector, write one frame of our own, then have the
// target send one. Returns the tunnel watermark after each step.
exchange := func(t *testing.T, proxyURI string, send chan struct{}) (afterOwnWrite, afterTargetFrame int64) {
conn, _, err := (&websocket.Dialer{}).Dial("ws"+proxyURI[4:]+"/ws", nil)
require.NoError(t, err)
defer conn.Close()

// The upgrade response is the first read from the target; only frames
// after it are under test.
time.Sleep(50 * time.Millisecond)
before := rr.lastTunnelTime.Load()

require.NoError(t, conn.WriteMessage(websocket.TextMessage, []byte("primus::pong::1")))
time.Sleep(50 * time.Millisecond)
afterOwnWrite = rr.lastTunnelTime.Load()
require.Equal(t, before, afterOwnWrite, "our own writes must not count as tunnel activity")

// Millisecond watermark: make sure a new stamp can differ from the old.
time.Sleep(5 * time.Millisecond)
close(send)
_, msg, err := conn.ReadMessage()
require.NoError(t, err)
require.Equal(t, "primus::ping::1", string(msg))
return afterOwnWrite, rr.lastTunnelTime.Load()
}

t.Run("customer origin", func(t *testing.T) {
target, send := newTarget(t)
before, after := exchange(t, rr.ProxyURI(target.URL), send)
require.Equal(t, before, after, "a websocket to a customer origin must not count as tunnel activity")
})

t.Run("broker tunnel", func(t *testing.T) {
target, send := newTarget(t)
before, after := exchange(t, rr.ProxyURI(target.URL, WithDefault(true)), send)
require.Greater(t, after, before, "a frame from the broker server must count as tunnel activity")
})
}

func TestWebSocketProxyIsConnected(t *testing.T) {
logger := newTestLogger(t)

Expand Down
Loading