package netutil import ( "testing" "github.com/gorilla/websocket" ) func TestSSRFWebsocketDialerPreservesHandshakeTimeout(t *testing.T) { dialer := SSRFWebsocketDialer(false) if dialer.HandshakeTimeout != websocket.DefaultDialer.HandshakeTimeout { t.Fatalf("HandshakeTimeout = %v, want %v", dialer.HandshakeTimeout, websocket.DefaultDialer.HandshakeTimeout) } if dialer.NetDialContext == nil { t.Fatal("NetDialContext is nil; public-only dialing is not enforced") } } func TestEnforceWSSURLAllowsProductionLoopback(t *testing.T) { for _, rawURL := range []string{ "ws://127.0.0.1:6555/mill", "ws://[::1]:6555/mill", } { t.Run(rawURL, func(t *testing.T) { if _, err := EnforceWSSURL(rawURL, false); err != nil { t.Fatal(err) } }) } } func TestEnforceWSSURLRejectsProductionCleartextOffHost(t *testing.T) { for _, rawURL := range []string{ "ws://192.0.2.1:6555/mill", "ws://mill.example.com/mill", } { t.Run(rawURL, func(t *testing.T) { if _, err := EnforceWSSURL(rawURL, false); err == nil { t.Fatalf("EnforceWSSURL(%q) accepted production cleartext", rawURL) } }) } }