11package main
22
33import (
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
1415const (
@@ -30,9 +31,6 @@ const (
3031var (
3132 activeBridge * wsConn
3233 activeBridgeMu sync.Mutex
33- bridgeUpgrader = websocket.Upgrader {
34- CheckOrigin : checkWSOrigin ,
35- }
3634)
3735
3836func 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
210218func 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}
0 commit comments