client_test.mx raw

   1  package ws
   2  
   3  import (
   4  	"errors"
   5  	"fmt"
   6  	"syscall"
   7  	"testing"
   8  
   9  	"git.smesh.lol/nostr/pkg/filter"
  10  )
  11  
  12  // The two tests below run against a loopback WebSocket server instead of
  13  // wss://relay.damus.io. The package under test is the client, so the server
  14  // only has to accept the HTTP upgrade, optionally consume one client frame,
  15  // send one text frame, and close. It runs in a spawn domain (fork) so the
  16  // client's blocking Dial has a peer; the listening fd is inherited at fork.
  17  
  18  func TestDialAndSubscribe(t *testing.T) {
  19  	fd, port, err := listenLoopback()
  20  	if err != nil {
  21  		t.Fatalf("listen: %v", err)
  22  		return
  23  	}
  24  	done := spawn(serveSubscribe, fd, []byte("[\"EOSE\",\"s1\"]"))
  25  
  26  	c, cerr := Connect("ws://127.0.0.1:" | fmt.Sprintf("%d", port) | "/")
  27  	if cerr != nil {
  28  		t.Fatalf("connect: %v", cerr)
  29  		return
  30  	}
  31  	sub, serr := c.Subscribe(&filter.F{})
  32  	if serr != nil {
  33  		t.Fatalf("subscribe: %v", serr)
  34  		c.Close()
  35  		return
  36  	}
  37  	op, payload, rerr := c.ws.ReadMessage()
  38  	if rerr != nil {
  39  		t.Fatalf("read EOSE: %v", rerr)
  40  		c.Close()
  41  		return
  42  	}
  43  	if op != OpText {
  44  		t.Errorf("op = %d, want %d", int32(op), int32(OpText))
  45  	}
  46  	c.dispatch(payload)
  47  	select {
  48  	case <-sub.EOSE:
  49  	default:
  50  		t.Error("EOSE frame was not dispatched to the subscription")
  51  	}
  52  	c.Close()
  53  	<-done
  54  }
  55  
  56  func TestWebSocketFraming(t *testing.T) {
  57  	fd, port, err := listenLoopback()
  58  	if err != nil {
  59  		t.Fatalf("listen: %v", err)
  60  		return
  61  	}
  62  	done := spawn(serveOnce, fd, []byte("hello"))
  63  
  64  	c, cerr := Dial("ws://127.0.0.1:" | fmt.Sprintf("%d", port) | "/")
  65  	if cerr != nil {
  66  		t.Fatalf("dial: %v", cerr)
  67  		return
  68  	}
  69  	op, payload, rerr := c.ReadMessage()
  70  	if rerr != nil {
  71  		t.Fatalf("read: %v", rerr)
  72  		c.Close()
  73  		return
  74  	}
  75  	if op != OpText {
  76  		t.Errorf("op = %d, want %d", int32(op), int32(OpText))
  77  	}
  78  	if string(payload) != "hello" {
  79  		t.Errorf("payload = %s, want hello", string(payload))
  80  	}
  81  	c.Close()
  82  	<-done
  83  }
  84  
  85  // listenLoopback binds 127.0.0.1:0 and returns the listening fd and port.
  86  // The receive timeout on the listener makes a blocked accept() in the spawned
  87  // server return instead of parking a domain forever when a test fails early.
  88  func listenLoopback() (fd int32, port int32, err error) {
  89  	fd, err = syscall.Socket(syscall.AF_INET, syscall.SOCK_STREAM, 0)
  90  	if err != nil {
  91  		return 0, 0, err
  92  	}
  93  	syscall.SetsockoptInt(fd, syscall.SOL_SOCKET, syscall.SO_REUSEADDR, 1)
  94  	tv := syscall.Timeval{Sec: 5}
  95  	syscall.SetsockoptTimeval(fd, syscall.SOL_SOCKET, syscall.SO_RCVTIMEO, &tv)
  96  	sa := &syscall.SockaddrInet4{Port: 0, Addr: [4]byte{127, 0, 0, 1}}
  97  	if err = syscall.Bind(fd, sa); err != nil {
  98  		syscall.Close(fd)
  99  		return 0, 0, err
 100  	}
 101  	if err = syscall.Listen(fd, 8); err != nil {
 102  		syscall.Close(fd)
 103  		return 0, 0, err
 104  	}
 105  	got, gerr := syscall.Getsockname(fd)
 106  	if gerr != nil {
 107  		syscall.Close(fd)
 108  		return 0, 0, gerr
 109  	}
 110  	sa4, ok := got.(*syscall.SockaddrInet4)
 111  	if !ok {
 112  		syscall.Close(fd)
 113  		return 0, 0, errors.New("ws: loopback listener is not inet4")
 114  	}
 115  	return fd, sa4.Port, nil
 116  }
 117  
 118  // serveOnce completes the upgrade handshake and sends one text frame.
 119  func serveOnce(fd int32, payload []byte) {
 120  	nfd, _, err := syscall.Accept(fd)
 121  	if err != nil {
 122  		syscall.Close(fd)
 123  		return
 124  	}
 125  	req := readUpgradeRequest(nfd)
 126  	out := []byte(upgradeHead(req)) | serverFrame(OpText, payload)
 127  	syscall.Write(nfd, out)
 128  	syscall.Close(nfd)
 129  	syscall.Close(fd)
 130  }
 131  
 132  // serveSubscribe completes the handshake, waits for the client's REQ frame,
 133  // then sends one text frame. Reading the REQ before closing keeps the client's
 134  // Subscribe write from racing a peer that has already closed.
 135  func serveSubscribe(fd int32, payload []byte) {
 136  	nfd, _, err := syscall.Accept(fd)
 137  	if err != nil {
 138  		syscall.Close(fd)
 139  		return
 140  	}
 141  	req := readUpgradeRequest(nfd)
 142  	syscall.Write(nfd, []byte(upgradeHead(req)))
 143  	buf := []byte{:512}
 144  	syscall.Read(nfd, buf)
 145  	syscall.Write(nfd, serverFrame(OpText, payload))
 146  	syscall.Close(nfd)
 147  	syscall.Close(fd)
 148  }
 149  
 150  // upgradeHead is the 101 response for the client's Sec-WebSocket-Key.
 151  func upgradeHead(req []byte) (head string) {
 152  	accept := ComputeAccept(websocketKey(req))
 153  	return "HTTP/1.1 101 Switching Protocols\r\n" |
 154  		"Upgrade: websocket\r\n" |
 155  		"Connection: Upgrade\r\n" |
 156  		"Sec-WebSocket-Accept: " | accept | "\r\n\r\n"
 157  }
 158  
 159  // websocketKey extracts the Sec-WebSocket-Key header value. A manual scan, not
 160  // bytes.Index: that helper is known to miss a present needle in this tree.
 161  func websocketKey(req []byte) (key string) {
 162  	needle := "Sec-WebSocket-Key: "
 163  	n := int32(len(needle))
 164  	for i := int32(0); i+n <= int32(len(req)); i++ {
 165  		if string(req[i:i+n]) == needle {
 166  			j := i + n
 167  			for j < int32(len(req)) && req[j] != '\r' && req[j] != '\n' {
 168  				j++
 169  			}
 170  			return string(req[i+n : j])
 171  		}
 172  	}
 173  	return ""
 174  }
 175  
 176  // readUpgradeRequest reads until the blank line that ends the HTTP headers.
 177  func readUpgradeRequest(fd int32) (req []byte) {
 178  	buf := []byte{:1024}
 179  	for {
 180  		n, err := syscall.Read(fd, buf)
 181  		if n <= 0 || err != nil {
 182  			return req
 183  		}
 184  		req = req | buf[:n]
 185  		if hasHeaderEnd(req) {
 186  			return req
 187  		}
 188  	}
 189  }
 190  
 191  func hasHeaderEnd(b []byte) (ok bool) {
 192  	for i := int32(0); i+3 < int32(len(b)); i++ {
 193  		if b[i] == '\r' && b[i+1] == '\n' && b[i+2] == '\r' && b[i+3] == '\n' {
 194  			return true
 195  		}
 196  	}
 197  	return false
 198  }
 199  
 200  // serverFrame builds an unmasked frame (RFC 6455: servers MUST NOT mask).
 201  func serverFrame(op byte, payload []byte) (buf []byte) {
 202  	plen := int32(len(payload))
 203  	if plen < 126 {
 204  		buf = []byte{:2}
 205  		buf[1] = byte(plen)
 206  	} else {
 207  		buf = []byte{:4}
 208  		buf[1] = 126
 209  		buf[2] = byte(plen >> 8)
 210  		buf[3] = byte(plen)
 211  	}
 212  	buf[0] = 0x80 | op
 213  	return buf | payload
 214  }
 215