Skip to content

Commit abea078

Browse files
committed
migrate gorilla/websocket to coder/websocket
1 parent cd6c6ac commit abea078

8 files changed

Lines changed: 225 additions & 107 deletions

File tree

bridge.go

Lines changed: 33 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,15 @@
11
package main
22

33
import (
4+
"context"
45
"encoding/json"
56
"log/slog"
67
"net/http"
78
"strings"
89
"sync"
910
"time"
1011

11-
"github.com/gorilla/websocket"
12+
"github.com/coder/websocket"
1213
)
1314

1415
const (
@@ -30,9 +31,6 @@ const (
3031
var (
3132
activeBridge *wsConn
3233
activeBridgeMu sync.Mutex
33-
bridgeUpgrader = websocket.Upgrader{
34-
CheckOrigin: checkWSOrigin,
35-
}
3634
)
3735

3836
func newBridge() http.Handler {
@@ -42,11 +40,18 @@ func newBridge() http.Handler {
4240
return
4341
}
4442

45-
raw, err := bridgeUpgrader.Upgrade(w, r, nil)
43+
if !checkWSOrigin(r) {
44+
http.Error(w, "Origin not allowed", http.StatusForbidden)
45+
return
46+
}
47+
48+
c, err := websocket.Accept(w, r, &websocket.AcceptOptions{
49+
InsecureSkipVerify: true,
50+
})
4651
if err != nil {
4752
return
4853
}
49-
client := &wsConn{Conn: raw}
54+
client := &wsConn{Conn: c, remoteAddr: r.RemoteAddr}
5055
client.SetReadLimit(maxMessageSize)
5156

5257
claims := getClaims(r)
@@ -58,7 +63,7 @@ func newBridge() http.Handler {
5863
if claims != nil {
5964
emitEvent("bridge.superseded", clientIP(r), claims.UID, r.UserAgent(), 0, map[string]any{
6065
"session_id": claims.SID,
61-
"previous_ip": prev.RemoteAddr().String(),
66+
"previous_ip": prev.remoteAddr,
6267
})
6368
}
6469
}
@@ -71,7 +76,7 @@ func newBridge() http.Handler {
7176
activeBridge = nil
7277
}
7378
activeBridgeMu.Unlock()
74-
client.Close()
79+
client.CloseNow()
7580
}()
7681

7782
// Dial backend (claw gateway)
@@ -80,14 +85,17 @@ func newBridge() http.Handler {
8085
if o := r.Header.Get("Origin"); o != "" {
8186
hdr.Set("Origin", o)
8287
}
83-
dialer := websocket.Dialer{HandshakeTimeout: dialTimeout}
84-
backend, _, err := dialer.Dial(wsURL, hdr)
88+
dialCtx, dialCancel := context.WithTimeout(r.Context(), dialTimeout)
89+
defer dialCancel()
90+
backend, _, err := websocket.Dial(dialCtx, wsURL, &websocket.DialOptions{
91+
HTTPHeader: hdr,
92+
})
8593
if err != nil {
8694
logError("bridge.dial", err)
8795
closeBridgeWS(client, codeBackendDown, "Backend unavailable")
8896
return
8997
}
90-
defer backend.Close()
98+
defer backend.CloseNow()
9199
backend.SetReadLimit(maxMessageSize)
92100

93101
done := make(chan struct{}, 3)
@@ -120,16 +128,13 @@ func newBridge() http.Handler {
120128
go func() {
121129
defer func() { done <- struct{}{} }()
122130
for {
123-
backend.SetReadDeadline(time.Now().Add(idleTimeout))
124-
mt, msg, err := backend.ReadMessage()
131+
readCtx, readCancel := context.WithTimeout(r.Context(), idleTimeout)
132+
mt, msg, err := backend.Read(readCtx)
133+
readCancel()
125134
if err != nil {
126135
return
127136
}
128-
client.wmu.Lock()
129-
client.Conn.SetWriteDeadline(time.Now().Add(writeWait))
130-
werr := client.Conn.WriteMessage(mt, msg)
131-
client.wmu.Unlock()
132-
if werr != nil {
137+
if err := client.safeWrite(mt, msg); err != nil {
133138
return
134139
}
135140
}
@@ -141,12 +146,13 @@ func newBridge() http.Handler {
141146
var count int
142147
win := time.Now()
143148
for {
144-
client.SetReadDeadline(time.Now().Add(idleTimeout))
145-
mt, msg, err := client.ReadMessage()
149+
readCtx, readCancel := context.WithTimeout(r.Context(), idleTimeout)
150+
mt, msg, err := client.Read(readCtx)
151+
readCancel()
146152
if err != nil {
147153
return
148154
}
149-
if mt == websocket.BinaryMessage {
155+
if mt == websocket.MessageBinary {
150156
closeBridgeWS(client, 1003, "Binary not supported")
151157
return
152158
}
@@ -160,11 +166,13 @@ func newBridge() http.Handler {
160166
closeBridgeWS(client, codeRateLimited, "Rate limited")
161167
return
162168
}
163-
if mt == websocket.TextMessage {
169+
if mt == websocket.MessageText {
164170
msg = injectToken(msg, cfg.GatewayToken)
165171
}
166-
backend.SetWriteDeadline(time.Now().Add(writeWait))
167-
if err := backend.WriteMessage(mt, msg); err != nil {
172+
wctx, wcancel := context.WithTimeout(r.Context(), writeWait)
173+
werr := backend.Write(wctx, mt, msg)
174+
wcancel()
175+
if werr != nil {
168176
return
169177
}
170178
}
@@ -208,7 +216,5 @@ func toWS(u string) string {
208216
}
209217

210218
func closeBridgeWS(c *wsConn, code int, reason string) {
211-
msg := websocket.FormatCloseMessage(code, reason)
212-
c.safeWriteControl(websocket.CloseMessage, msg, time.Now().Add(time.Second))
213-
c.Close()
219+
c.Close(websocket.StatusCode(code), reason)
214220
}

go.mod

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -3,9 +3,11 @@ module birdcage
33
go 1.26.1
44

55
require (
6+
github.com/charmbracelet/bubbles v1.0.0
7+
github.com/charmbracelet/bubbletea v1.3.10
8+
github.com/charmbracelet/lipgloss v1.1.0
69
github.com/coder/websocket v1.8.14
710
github.com/golang-jwt/jwt/v5 v5.3.1
8-
github.com/gorilla/websocket v1.5.3
911
github.com/joho/godotenv v1.5.1
1012
github.com/matoous/go-nanoid/v2 v2.1.0
1113
golang.org/x/crypto v0.48.0
@@ -15,10 +17,7 @@ require (
1517

1618
require (
1719
github.com/aymanbagabas/go-osc52/v2 v2.0.1 // indirect
18-
github.com/charmbracelet/bubbles v1.0.0 // indirect
19-
github.com/charmbracelet/bubbletea v1.3.10 // indirect
2020
github.com/charmbracelet/colorprofile v0.4.1 // indirect
21-
github.com/charmbracelet/lipgloss v1.1.0 // indirect
2221
github.com/charmbracelet/x/ansi v0.11.6 // indirect
2322
github.com/charmbracelet/x/cellbuf v0.0.15 // indirect
2423
github.com/charmbracelet/x/term v0.2.2 // indirect

go.sum

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
github.com/aymanbagabas/go-osc52/v2 v2.0.1 h1:HwpRHbFMcZLEVr42D4p7XBqjyuxQH5SMiErDT4WkJ2k=
22
github.com/aymanbagabas/go-osc52/v2 v2.0.1/go.mod h1:uYgXzlJ7ZpABp8OJ+exZzJJhRNQ2ASbcXHWsFqH8hp8=
3+
github.com/aymanbagabas/go-udiff v0.3.1 h1:LV+qyBQ2pqe0u42ZsUEtPiCaUoqgA9gYRDs3vj1nolY=
4+
github.com/aymanbagabas/go-udiff v0.3.1/go.mod h1:G0fsKmG+P6ylD0r6N/KgQD/nWzgfnl8ZBcNLgcbrw8E=
35
github.com/charmbracelet/bubbles v1.0.0 h1:12J8/ak/uCZEMQ6KU7pcfwceyjLlWsDLAxB5fXonfvc=
46
github.com/charmbracelet/bubbles v1.0.0/go.mod h1:9d/Zd5GdnauMI5ivUIVisuEm3ave1XwXtD1ckyV6r3E=
57
github.com/charmbracelet/bubbletea v1.3.10 h1:otUDHWMMzQSB0Pkc87rm691KZ3SWa4KUlvF9nRvCICw=
@@ -12,6 +14,8 @@ github.com/charmbracelet/x/ansi v0.11.6 h1:GhV21SiDz/45W9AnV2R61xZMRri5NlLnl6CVF
1214
github.com/charmbracelet/x/ansi v0.11.6/go.mod h1:2JNYLgQUsyqaiLovhU2Rv/pb8r6ydXKS3NIttu3VGZQ=
1315
github.com/charmbracelet/x/cellbuf v0.0.15 h1:ur3pZy0o6z/R7EylET877CBxaiE1Sp1GMxoFPAIztPI=
1416
github.com/charmbracelet/x/cellbuf v0.0.15/go.mod h1:J1YVbR7MUuEGIFPCaaZ96KDl5NoS0DAWkskup+mOY+Q=
17+
github.com/charmbracelet/x/exp/golden v0.0.0-20241011142426-46044092ad91 h1:payRxjMjKgx2PaCWLZ4p3ro9y97+TVLZNaRZgJwSVDQ=
18+
github.com/charmbracelet/x/exp/golden v0.0.0-20241011142426-46044092ad91/go.mod h1:wDlXFlCrmJ8J+swcL/MnGUuYnqgQdW9rhSD61oNMb6U=
1519
github.com/charmbracelet/x/term v0.2.2 h1:xVRT/S2ZcKdhhOuSP4t5cLi5o+JxklsoEObBSgfgZRk=
1620
github.com/charmbracelet/x/term v0.2.2/go.mod h1:kF8CY5RddLWrsgVwpw4kAa6TESp6EB5y3uxGLeCqzAI=
1721
github.com/clipperhouse/displaywidth v0.9.0 h1:Qb4KOhYwRiN3viMv1v/3cTBlz3AcAZX3+y9OLhMtAtA=
@@ -34,8 +38,6 @@ github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17k
3438
github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA=
3539
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
3640
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
37-
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
38-
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
3941
github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
4042
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
4143
github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0=

node.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@ import (
1212
"sync"
1313
"time"
1414

15-
"github.com/gorilla/websocket"
15+
"github.com/coder/websocket"
1616
"golang.org/x/crypto/curve25519"
1717
)
1818

@@ -267,7 +267,7 @@ func notifyNodeSync() {
267267
}
268268
data, _ := json.Marshal(msg)
269269
if target.conn != nil {
270-
target.conn.safeWrite(websocket.TextMessage, data)
270+
target.conn.safeWrite(websocket.MessageText, data)
271271
}
272272
}
273273
}

relay.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@ import (
66
"sync"
77
"time"
88

9-
"github.com/gorilla/websocket"
9+
"github.com/coder/websocket"
1010
)
1111

1212
const (
@@ -141,5 +141,5 @@ func handleRelayPacket(sourceNodeID int, raw []byte) {
141141
binary.BigEndian.PutUint32(header, uint32(sourceNodeID)) // #nosec G115 — node IDs are small positive integers from SQLite autoincrement
142142
packet := append(header, raw[4:]...)
143143

144-
destNode.conn.safeWrite(websocket.BinaryMessage, packet)
144+
destNode.conn.safeWrite(websocket.MessageBinary, packet)
145145
}

0 commit comments

Comments
 (0)