diff --git a/websocket.go b/websocket.go index 9026115..ea81ae3 100644 --- a/websocket.go +++ b/websocket.go @@ -9,6 +9,7 @@ import ( "encoding/binary" "fmt" "io" + "log/slog" "net" "net/url" "strings" @@ -52,27 +53,46 @@ var tlsDial = func(network, addr string) (net.Conn, error) { } // WSDial connects to a WebSocket endpoint and performs the HTTP upgrade handshake. +// If the initial connection attempt fails (e.g. server returns 403 over IPv6), +// it retries with IPv4 only. func WSDial(rawURL string) (*WSConn, error) { u, err := url.Parse(rawURL) if err != nil { return nil, fmt.Errorf("parse url: %w", err) } - useTLS := u.Scheme == "wss" host := u.Host if !strings.Contains(host, ":") { - if useTLS { + if u.Scheme == "wss" { host += ":443" } else { host += ":80" } } + ws, err := wsConnect(u, host, "tcp") + if err != nil { + slog.Warn("connection failed, retrying with IPv4", "error", err) + ws, err4 := wsConnect(u, host, "tcp4") + if err4 != nil { + return nil, err // return original error + } + return ws, nil + } + return ws, nil +} + +// wsConnect dials the host using the given network ("tcp", "tcp4", "tcp6") +// and performs the WebSocket upgrade handshake. +func wsConnect(u *url.URL, host, network string) (*WSConn, error) { + useTLS := u.Scheme == "wss" + var conn net.Conn + var err error if useTLS { - conn, err = tlsDial("tcp", host) + conn, err = tlsDial(network, host) } else { - conn, err = netDial("tcp", host) + conn, err = netDial(network, host) } if err != nil { return nil, fmt.Errorf("dial %s: %w", host, err) diff --git a/websocket_test.go b/websocket_test.go index ec51ea3..659cea3 100644 --- a/websocket_test.go +++ b/websocket_test.go @@ -453,17 +453,19 @@ func TestWSDialHandshakeFailed(t *testing.T) { defer func() { _ = ln.Close() }() go func() { - conn, err := ln.Accept() - if err != nil { - return + for { + conn, err := ln.Accept() + if err != nil { + return + } + go func(c net.Conn) { + defer func() { _ = c.Close() }() + buf := make([]byte, 4096) + _, _ = c.Read(buf) + resp := "HTTP/1.1 403 Forbidden\r\nContent-Length: 0\r\n\r\n" + _, _ = c.Write([]byte(resp)) + }(conn) } - defer func() { _ = conn.Close() }() - - buf := make([]byte, 4096) - _, _ = conn.Read(buf) - - resp := "HTTP/1.1 403 Forbidden\r\nContent-Length: 0\r\n\r\n" - _, _ = conn.Write([]byte(resp)) }() addr := ln.Addr().String() @@ -485,12 +487,16 @@ func TestWSDialReadHandshakeError(t *testing.T) { defer func() { _ = ln.Close() }() go func() { - conn, _ := ln.Accept() - if conn != nil { - // Read the request then close immediately - buf := make([]byte, 4096) - _, _ = conn.Read(buf) - _ = conn.Close() + for { + conn, err := ln.Accept() + if err != nil { + return + } + go func(c net.Conn) { + buf := make([]byte, 4096) + _, _ = c.Read(buf) + _ = c.Close() + }(conn) } }() @@ -597,6 +603,60 @@ func TestWSDialHandshakeTimeout(t *testing.T) { } } +func TestWSDialIPv4Fallback(t *testing.T) { + // Start a test TCP server on IPv4 that responds with valid WebSocket upgrade + ln, err := net.Listen("tcp4", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer func() { _ = ln.Close() }() + + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + go func(c net.Conn) { + defer func() { _ = c.Close() }() + buf := make([]byte, 4096) + n, _ := c.Read(buf) + key := extractHeader(string(buf[:n]), "Sec-WebSocket-Key") + accept := computeAcceptKey(key) + resp := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: " + accept + "\r\n\r\n" + _, _ = c.Write([]byte(resp)) + b := make([]byte, 1) + _, _ = c.Read(b) + }(conn) + } + }() + + addr := ln.Addr().String() + callCount := 0 + origDial := netDial + defer func() { netDial = origDial }() + + netDial = func(network, a string) (net.Conn, error) { + callCount++ + if network == "tcp" { + // Simulate IPv6 failure: TCP connects but server rejects + return nil, fmt.Errorf("dial tcp [::1]:%s: connect: connection refused", a) + } + // tcp4 fallback goes to the real server + return net.Dial(network, addr) + } + + ws, err := WSDial("ws://" + addr + "/socket") + if err != nil { + t.Fatalf("expected IPv4 fallback to succeed: %v", err) + } + _ = ws.Close() + + if callCount < 2 { + t.Errorf("expected at least 2 dial attempts (tcp + tcp4), got %d", callCount) + } +} + func TestWSDialDefaultPorts(t *testing.T) { // Test that ws:// defaults to port 80 — will fail to connect but verifies URL parsing _, err := WSDial("ws://127.0.0.1/path") @@ -659,18 +719,19 @@ func TestWSDialRejectsInvalidAcceptKey(t *testing.T) { defer func() { _ = ln.Close() }() go func() { - conn, err := ln.Accept() - if err != nil { - return + for { + conn, err := ln.Accept() + if err != nil { + return + } + go func(c net.Conn) { + defer func() { _ = c.Close() }() + buf := make([]byte, 4096) + _, _ = c.Read(buf) + resp := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: INVALID_KEY\r\n\r\n" + _, _ = c.Write([]byte(resp)) + }(conn) } - defer func() { _ = conn.Close() }() - - buf := make([]byte, 4096) - _, _ = conn.Read(buf) - - // Send 101 with a wrong accept key - resp := "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: INVALID_KEY\r\n\r\n" - _, _ = conn.Write([]byte(resp)) }() addr := ln.Addr().String()