// ws.go: a minimal WebSocket (RFC 6455) implementation on the standard // library, server and client. Not using gorilla/websocket so the program // stays a single .exe with no dependencies to ship or vendor. // // Scope is deliberately narrow, just what we need: small text messages, // no fragmentation, no compression. // // TLS is handled one layer up, not in here: wsUpgrade doesn't need to // know about it (it terminates at the http.Server/listener level), and // wsDialTLS just runs the same client handshake over a *tls.Conn instead // of a plain one. See pin.go for certificate pinning. package main import ( "bufio" "crypto/rand" "crypto/sha1" "crypto/tls" "encoding/base64" "encoding/binary" "fmt" "io" "net" "net/http" "strings" "sync" "time" ) const ( wsGUID = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11" opContinuation = 0x0 opText = 0x1 opBinary = 0x2 opClose = 0x8 opPing = 0x9 opPong = 0xA maxFrameSize = 1 << 20 // 1 MiB: our messages run ~100 bytes ) type wsConn struct { conn net.Conn br *bufio.Reader isClient bool // only the client masks, per the RFC wmu sync.Mutex closed bool } func wsAcceptKey(key string) string { h := sha1.New() io.WriteString(h, key+wsGUID) return base64.StdEncoding.EncodeToString(h.Sum(nil)) } // wsUpgrade turns an incoming HTTP request into a WebSocket connection // (server side). func wsUpgrade(w http.ResponseWriter, r *http.Request) (*wsConn, error) { if !strings.Contains(strings.ToLower(r.Header.Get("Connection")), "upgrade") || !strings.EqualFold(r.Header.Get("Upgrade"), "websocket") { return nil, fmt.Errorf("not a websocket upgrade") } key := r.Header.Get("Sec-WebSocket-Key") if key == "" { return nil, fmt.Errorf("missing the Sec-WebSocket-Key header") } hj, ok := w.(http.Hijacker) if !ok { return nil, fmt.Errorf("this server doesn't support hijack") } conn, brw, err := hj.Hijack() if err != nil { return nil, err } resp := "HTTP/1.1 101 Switching Protocols\r\n" + "Upgrade: websocket\r\n" + "Connection: Upgrade\r\n" + "Sec-WebSocket-Accept: " + wsAcceptKey(key) + "\r\n\r\n" if _, err := conn.Write([]byte(resp)); err != nil { conn.Close() return nil, err } return &wsConn{conn: conn, br: brw.Reader}, nil } // wsDial opens a WebSocket connection to a hub (client side). func wsDial(addr, path string, timeout time.Duration) (*wsConn, error) { conn, err := net.DialTimeout("tcp", addr, timeout) if err != nil { return nil, err } return wsHandshake(conn, addr, path, timeout) } // wsDialTLS is wsDial over an encrypted connection: same handshake, dialed // through tlsCfg instead of a plain net.Dial. The TLS handshake itself // (including certificate verification, e.g. pinning — see pin.go) happens // inside tls.DialWithDialer before the WebSocket upgrade is attempted. func wsDialTLS(addr, path string, timeout time.Duration, tlsCfg *tls.Config) (*wsConn, error) { conn, err := tls.DialWithDialer(&net.Dialer{Timeout: timeout}, "tcp", addr, tlsCfg) if err != nil { return nil, err } return wsHandshake(conn, addr, path, timeout) } // wsHandshake does the WebSocket upgrade handshake (client side) over an // already-established connection, plain or TLS — both satisfy net.Conn. func wsHandshake(conn net.Conn, addr, path string, timeout time.Duration) (*wsConn, error) { var keyBytes [16]byte if _, err := rand.Read(keyBytes[:]); err != nil { conn.Close() return nil, err } key := base64.StdEncoding.EncodeToString(keyBytes[:]) req := "GET " + path + " HTTP/1.1\r\n" + "Host: " + addr + "\r\n" + "Upgrade: websocket\r\n" + "Connection: Upgrade\r\n" + "Sec-WebSocket-Key: " + key + "\r\n" + "Sec-WebSocket-Version: 13\r\n\r\n" conn.SetDeadline(time.Now().Add(timeout)) if _, err := conn.Write([]byte(req)); err != nil { conn.Close() return nil, err } br := bufio.NewReader(conn) resp, err := http.ReadResponse(br, nil) if err != nil { conn.Close() return nil, err } resp.Body.Close() if resp.StatusCode != http.StatusSwitchingProtocols { conn.Close() return nil, fmt.Errorf("the hub replied %s (expected 101)", resp.Status) } if !strings.EqualFold(resp.Header.Get("Sec-WebSocket-Accept"), wsAcceptKey(key)) { conn.Close() return nil, fmt.Errorf("the handshake doesn't validate (is there really a websocket on the other end?)") } conn.SetDeadline(time.Time{}) return &wsConn{conn: conn, br: br, isClient: true}, nil } func (c *wsConn) writeFrame(opcode byte, payload []byte) error { hdr := make([]byte, 0, 14) hdr = append(hdr, 0x80|opcode) // FIN + opcode maskBit := byte(0) if c.isClient { maskBit = 0x80 } n := len(payload) switch { case n <= 125: hdr = append(hdr, maskBit|byte(n)) case n <= 65535: var ext [2]byte binary.BigEndian.PutUint16(ext[:], uint16(n)) hdr = append(hdr, maskBit|126) hdr = append(hdr, ext[:]...) default: var ext [8]byte binary.BigEndian.PutUint64(ext[:], uint64(n)) hdr = append(hdr, maskBit|127) hdr = append(hdr, ext[:]...) } body := payload if c.isClient { var mask [4]byte if _, err := rand.Read(mask[:]); err != nil { return err } hdr = append(hdr, mask[:]...) body = make([]byte, n) for i := 0; i < n; i++ { body[i] = payload[i] ^ mask[i%4] } } c.wmu.Lock() defer c.wmu.Unlock() if c.closed { return io.ErrClosedPipe } if _, err := c.conn.Write(hdr); err != nil { return err } if n > 0 { if _, err := c.conn.Write(body); err != nil { return err } } return nil } func (c *wsConn) readFrame() (opcode byte, payload []byte, err error) { var h [2]byte if _, err = io.ReadFull(c.br, h[:]); err != nil { return } fin := h[0]&0x80 != 0 opcode = h[0] & 0x0F masked := h[1]&0x80 != 0 n := int(h[1] & 0x7F) switch n { case 126: var ext [2]byte if _, err = io.ReadFull(c.br, ext[:]); err != nil { return } n = int(binary.BigEndian.Uint16(ext[:])) case 127: var ext [8]byte if _, err = io.ReadFull(c.br, ext[:]); err != nil { return } v := binary.BigEndian.Uint64(ext[:]) if v > maxFrameSize { err = fmt.Errorf("frame too large (%d bytes)", v) return } n = int(v) } if n > maxFrameSize { err = fmt.Errorf("frame too large (%d bytes)", n) return } var mask [4]byte if masked { if _, err = io.ReadFull(c.br, mask[:]); err != nil { return } } payload = make([]byte, n) if n > 0 { if _, err = io.ReadFull(c.br, payload); err != nil { return } } if masked { for i := range payload { payload[i] ^= mask[i%4] } } if !fin || opcode == opContinuation { err = fmt.Errorf("fragmented frames aren't supported") } return } // ReadMessage returns the next text/binary message, answering pings // internally along the way. A close from the other side reports as io.EOF. func (c *wsConn) ReadMessage() ([]byte, error) { for { op, payload, err := c.readFrame() if err != nil { return nil, err } switch op { case opText, opBinary: return payload, nil case opPing: if err := c.writeFrame(opPong, payload); err != nil { return nil, err } case opPong: // nothing to do case opClose: c.writeFrame(opClose, nil) return nil, io.EOF default: return nil, fmt.Errorf("unknown opcode: 0x%X", op) } } } func (c *wsConn) WriteText(b []byte) error { return c.writeFrame(opText, b) } func (c *wsConn) Ping() error { return c.writeFrame(opPing, nil) } func (c *wsConn) SetReadDeadline(t time.Time) error { return c.conn.SetReadDeadline(t) } func (c *wsConn) RemoteAddr() string { return c.conn.RemoteAddr().String() } func (c *wsConn) Close() error { c.wmu.Lock() if c.closed { c.wmu.Unlock() return nil } c.closed = true c.wmu.Unlock() return c.conn.Close() }