1 // Package ws provides a minimal WebSocket client for Nostr relays.
2 // Implements RFC 6455 over raw TCP/TLS - no net/http dependency.
3 package ws
4
5 import (
6 "bufio"
7 "bytes"
8 "crypto/rand"
9 "crypto/sha1"
10 "crypto/tls"
11 "encoding/base64"
12 "encoding/binary"
13 "fmt"
14 "io"
15 "net"
16 "net/url"
17 "time"
18 )
19
20 const (
21 OpText byte = 0x1
22 OpBinary byte = 0x2
23 OpClose byte = 0x8
24 OpPing byte = 0x9
25 OpPong byte = 0xA
26
27 wsMagic = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"
28 maxPayload = 33 << 20 // 33 MB
29 )
30
31 // Conn is a WebSocket connection (client or server mode).
32 type Conn struct {
33 raw net.Conn
34 br *bufio.Reader
35 server bool // true = server mode (unmasked writes)
36 }
37
38 // NewServerConn wraps an already-upgraded connection in server mode.
39 func NewServerConn(conn net.Conn, br *bufio.Reader) (c *Conn) {
40 return &Conn{raw: conn, br: br, server: true}
41 }
42
43 // ComputeAccept calculates Sec-WebSocket-Accept from a client key.
44 func ComputeAccept(key string) (s string) { return computeAccept(key) }
45
46 // Dial opens a WebSocket connection to the given URL (ws:// or wss://).
47 func Dial(rawURL string) (c *Conn, derr error) {
48 u, err := url.Parse(rawURL)
49 if err != nil {
50 return nil, err
51 }
52 host := u.Hostname()
53 port := u.Port()
54 useTLS := u.Scheme == "wss"
55 if port == "" {
56 if useTLS {
57 port = "443"
58 } else {
59 port = "80"
60 }
61 }
62 // Resolve hostname via DNS cache (24h TTL).
63 ip := host
64 if net.ParseIP(host) == nil {
65 ip, err = resolveHost(host)
66 if err != nil {
67 return nil, fmt.Errorf("ws: resolve %s: %w", host, err)
68 }
69 }
70 addr := net.JoinHostPort(ip, port)
71
72 var conn net.Conn
73 conn, err = net.Dial("tcp", addr)
74 if err != nil {
75 return nil, fmt.Errorf("ws: dial %s: %w", addr, err)
76 }
77 if useTLS {
78 tlsConn := tls.Client(conn, &tls.Config{ServerName: []byte(host)})
79 if err = tlsConn.Handshake(); err != nil {
80 conn.Close()
81 return nil, fmt.Errorf("ws: tls %s: %w", host, err)
82 }
83 conn = tlsConn
84 }
85
86 path := u.RequestURI()
87 if path == "" {
88 path = "/"
89 }
90
91 // Generate Sec-WebSocket-Key.
92 var keyRaw [16]byte
93 io.ReadFull(rand.Reader(), keyRaw[:])
94 wsKey := base64.StdEncoding().EncodeToString(keyRaw[:])
95
96 // Send HTTP upgrade.
97 req := "GET " | path | " HTTP/1.1\r\n" |
98 "Host: " | host | "\r\n" |
99 "Upgrade: websocket\r\n" |
100 "Connection: Upgrade\r\n" |
101 "Sec-WebSocket-Key: " | wsKey | "\r\n" |
102 "Sec-WebSocket-Version: 13\r\n" |
103 "\r\n"
104 if _, err = conn.Write([]byte(req)); err != nil {
105 conn.Close()
106 return nil, fmt.Errorf("ws: write upgrade: %w", err)
107 }
108
109 br := bufio.NewReaderSize(conn, 32768)
110
111 // Read status line.
112 status, err := br.ReadString('\n')
113 if err != nil {
114 conn.Close()
115 return nil, fmt.Errorf("ws: read status: %w", err)
116 }
117 if !bytes.Contains(status, "101") {
118 conn.Close()
119 return nil, fmt.Errorf("ws: upgrade rejected: %s", bytes.TrimSpace(status))
120 }
121
122 // Consume headers, validate accept.
123 expectedAccept := computeAccept(wsKey)
124 var accepted bool
125 for {
126 line, lerr := br.ReadString('\n')
127 if lerr != nil {
128 conn.Close()
129 return nil, fmt.Errorf("ws: read header: %w", lerr)
130 }
131 trimmed := bytes.TrimSpace(line)
132 if trimmed == "" {
133 break
134 }
135 lower := bytes.ToLower(trimmed)
136 if bytes.HasPrefix(lower, "sec-websocket-accept:") {
137 val := bytes.TrimSpace(trimmed[len("sec-websocket-accept:"):])
138 if val == expectedAccept {
139 accepted = true
140 }
141 }
142 }
143 if !accepted {
144 conn.Close()
145 return nil, fmt.Errorf("ws: bad accept key")
146 }
147 return &Conn{raw: conn, br: br}, nil
148 }
149
150 func computeAccept(key string) (s string) {
151 h := sha1.New()
152 h.Write([]byte(key))
153 h.Write([]byte(wsMagic))
154 return base64.StdEncoding().EncodeToString(h.Sum(nil))
155 }
156
157 // WriteText sends a text frame.
158 func (c *Conn) WriteText(msg []byte) (err error) { return c.writeFrame(OpText, msg) }
159
160 // WritePong sends a pong control frame.
161 func (c *Conn) WritePong(data []byte) (err error) { return c.writeFrame(OpPong, data) }
162
163 func (c *Conn) writeFrame(op byte, payload []byte) (err error) {
164 if c.server {
165 return c.writeServerFrame(op, payload)
166 }
167 return c.writeClientFrame(op, payload)
168 }
169
170 // writeServerFrame sends an unmasked frame (RFC 6455: servers MUST NOT mask).
171 func (c *Conn) writeServerFrame(op byte, payload []byte) (err error) {
172 plen := len(payload)
173 var hdr []byte
174 if plen < 126 {
175 hdr = []byte{:2}
176 hdr[1] = byte(plen)
177 } else if plen < 65536 {
178 hdr = []byte{:4}
179 hdr[1] = 126
180 binary.BigEndian().PutUint16(hdr[2:], uint16(plen))
181 } else {
182 hdr = []byte{:10}
183 hdr[1] = 127
184 binary.BigEndian().PutUint64(hdr[2:], uint64(plen))
185 }
186 hdr[0] = 0x80 | op
187 if _, err = c.raw.Write(hdr); err != nil {
188 return err
189 }
190 _, err := c.raw.Write(payload)
191 return err
192 }
193
194 // writeClientFrame sends a masked frame (RFC 6455: clients MUST mask).
195 func (c *Conn) writeClientFrame(op byte, payload []byte) (err error) {
196 plen := len(payload)
197 var hdr []byte
198 if plen < 126 {
199 hdr = []byte{:6} // 2 + 4 mask
200 hdr[1] = 0x80 | byte(plen)
201 } else if plen < 65536 {
202 hdr = []byte{:8} // 4 + 4 mask
203 hdr[1] = 0x80 | 126
204 binary.BigEndian().PutUint16(hdr[2:], uint16(plen))
205 } else {
206 hdr = []byte{:14} // 10 + 4 mask
207 hdr[1] = 0x80 | 127
208 binary.BigEndian().PutUint64(hdr[2:], uint64(plen))
209 }
210 hdr[0] = 0x80 | op
211
212 maskOff := len(hdr) - 4
213 io.ReadFull(rand.Reader(), hdr[maskOff:])
214 mask := [4]byte{hdr[maskOff], hdr[maskOff+1], hdr[maskOff+2], hdr[maskOff+3]}
215
216 masked := []byte{:plen}
217 // An index loop, not `for i, b := range payload`: stage4 typed the range
218 // value as a pointer here and emitted `xor ptr`, which clang refused. The
219 // elements are the same either way.
220 for i := int32(0); i < int32(len(payload)); i++ {
221 masked[i] = payload[i] ^ mask[i%4]
222 }
223 if _, err = c.raw.Write(hdr); err != nil {
224 return err
225 }
226 _, err := c.raw.Write(masked)
227 return err
228 }
229
230 // ReadMessage reads the next data frame, automatically handling ping/pong.
231 // Returns OpText, OpBinary, or OpClose. Reserved opcodes and bare
232 // continuation frames (without a preceding fragment, which this library
233 // does not support) are reported as protocol errors so callers don't spin
234 // in `if op != OpText { continue }` loops on malformed upstream traffic.
235 func (c *Conn) ReadMessage() (op byte, payload []byte, err error) {
236 for {
237 op, payload, err = c.readFrame()
238 if err != nil {
239 return
240 }
241 switch op {
242 case OpText, OpBinary:
243 return
244 case OpPing:
245 c.WritePong(payload)
246 case OpPong:
247 // ignore
248 case OpClose:
249 c.writeFrame(OpClose, payload)
250 return
251 default:
252 err = fmt.Errorf("ws: unexpected opcode 0x%x", op)
253 return
254 }
255 }
256 }
257
258 func (c *Conn) readFrame() (op byte, payload []byte, err error) {
259 var hdr [2]byte
260 if _, err = io.ReadFull(c.br, hdr[:]); err != nil {
261 return
262 }
263 op = hdr[0] & 0x0F
264 masked := hdr[1]&0x80 != 0
265 plen := uint64(hdr[1] & 0x7F)
266
267 if plen == 126 {
268 var ext [2]byte
269 if _, err = io.ReadFull(c.br, ext[:]); err != nil {
270 return
271 }
272 plen = uint64(binary.BigEndian().Uint16(ext[:]))
273 } else if plen == 127 {
274 var ext [8]byte
275 if _, err = io.ReadFull(c.br, ext[:]); err != nil {
276 return
277 }
278 plen = binary.BigEndian().Uint64(ext[:])
279 }
280 if plen > uint64(maxPayload) {
281 err = fmt.Errorf("ws: payload %d exceeds limit %d", plen, maxPayload)
282 return
283 }
284
285 var mask [4]byte
286 if masked {
287 if _, err = io.ReadFull(c.br, mask[:]); err != nil {
288 return
289 }
290 }
291
292 // Read into a local, unmask it there, and hand it to the named result at
293 // the end. Writing through the result variable (`payload[i] ^= ...`) is the
294 // one shape stage4 could not type: it emitted `xor ptr` for the element and
295 // clang refused the module. A named result is an immutable borrow in
296 // Moxie; the element write belongs on a local slice.
297 buf := []byte{:plen}
298 if _, err = io.ReadFull(c.br, buf); err != nil {
299 return
300 }
301 if masked {
302 for i := range buf {
303 buf[i] = buf[i] ^ mask[i%4]
304 }
305 }
306 payload = buf
307 return
308 }
309
310 // SetReadDeadline sets the read deadline on the underlying connection.
311 func (c *Conn) SetReadDeadline(t time.Time) (err error) {
312 return c.raw.SetReadDeadline(t)
313 }
314
315 // Close sends a close frame and closes the TCP connection.
316 func (c *Conn) Close() (err error) {
317 data := []byte{0x03, 0xE8} // code 1000 normal closure
318 c.writeFrame(OpClose, data)
319 return c.raw.Close()
320 }
321