proxy_test.go
1package ysweet
2
3import (
4 "context"
5 "fmt"
6 "net/http"
7 "net/http/httptest"
8 "net/url"
9 "strings"
10 "sync/atomic"
11 "testing"
12 "time"
13
14 "github.com/coder/websocket"
15)
16
17// mockYSweetServer creates a test y-sweet server that echoes messages
18func mockYSweetServer(t *testing.T) *httptest.Server {
19 return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
20 // Accept websocket connection
21 conn, err := websocket.Accept(w, r, nil)
22 if err != nil {
23 t.Logf("accept error: %v", err)
24 return
25 }
26 defer conn.Close(websocket.StatusNormalClosure, "")
27
28 // Echo messages back
29 ctx := context.Background()
30 for {
31 msgType, reader, err := conn.Reader(ctx)
32 if err != nil {
33 return
34 }
35
36 data := make([]byte, 1024)
37 n, err := reader.Read(data)
38 if err != nil {
39 return
40 }
41
42 err = conn.Write(ctx, msgType, data[:n])
43 if err != nil {
44 return
45 }
46 }
47 }))
48}
49
50func TestProxyHandler_BasicProxy(t *testing.T) {
51 // Create mock upstream server
52 upstream := mockYSweetServer(t)
53 defer upstream.Close()
54
55 // Convert http:// to ws://
56 upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
57
58 // Create proxy handler
59 handler := ProxyHandler(upstreamURL)
60
61 // Create test server with proxy
62 proxyServer := httptest.NewServer(handler)
63 defer proxyServer.Close()
64
65 // Connect to proxy
66 proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1) + "/d/test-doc/ws/test-doc"
67 clientConn, _, err := websocket.Dial(context.Background(), proxyURL, nil)
68 if err != nil {
69 t.Fatalf("failed to connect to proxy: %v", err)
70 }
71 defer clientConn.Close(websocket.StatusNormalClosure, "")
72
73 // Send a test message
74 testMsg := []byte{0x00, 0x01, 0x02, 0x03}
75 err = clientConn.Write(context.Background(), websocket.MessageBinary, testMsg)
76 if err != nil {
77 t.Fatalf("failed to write message: %v", err)
78 }
79
80 // Read echoed message
81 ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
82 defer cancel()
83
84 msgType, reader, err := clientConn.Reader(ctx)
85 if err != nil {
86 t.Fatalf("failed to read message: %v", err)
87 }
88
89 if msgType != websocket.MessageBinary {
90 t.Errorf("expected binary message, got %v", msgType)
91 }
92
93 data := make([]byte, 1024)
94 n, err := reader.Read(data)
95 if err != nil {
96 t.Fatalf("failed to read message body: %v", err)
97 }
98
99 if string(data[:n]) != string(testMsg) {
100 t.Errorf("expected %v, got %v", testMsg, data[:n])
101 }
102}
103
104func TestProxyHandler_OnConnectHook(t *testing.T) {
105 upstream := mockYSweetServer(t)
106 defer upstream.Close()
107 upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
108
109 var hookCalled atomic.Bool
110 var hookReceivedDocID string
111 var hookReceivedPath string
112
113 handler := ProxyHandler(upstreamURL,
114 WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
115 hookCalled.Store(true)
116 hookReceivedDocID = docID
117 hookReceivedPath = r.URL.Path
118 return nil
119 }),
120 )
121
122 proxyServer := httptest.NewServer(handler)
123 defer proxyServer.Close()
124
125 proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1) + "/d/test-doc/ws/test-doc"
126 clientConn, _, err := websocket.Dial(context.Background(), proxyURL, nil)
127 if err != nil {
128 t.Fatalf("failed to connect: %v", err)
129 }
130 clientConn.Close(websocket.StatusNormalClosure, "")
131
132 if !hookCalled.Load() {
133 t.Error("OnConnect hook was not called")
134 }
135
136 if hookReceivedDocID != "test-doc" {
137 t.Errorf("expected docID 'test-doc', got %s", hookReceivedDocID)
138 }
139
140 if hookReceivedPath != "/d/test-doc/ws/test-doc" {
141 t.Errorf("expected path /d/test-doc/ws/test-doc, got %s", hookReceivedPath)
142 }
143}
144
145func TestProxyHandler_OnConnectHook_RejectsConnection(t *testing.T) {
146 upstream := mockYSweetServer(t)
147 defer upstream.Close()
148 upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
149
150 handler := ProxyHandler(upstreamURL,
151 WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
152 return fmt.Errorf("connection rejected")
153 }),
154 )
155
156 proxyServer := httptest.NewServer(handler)
157 defer proxyServer.Close()
158
159 // Try HTTP request (not websocket) to see rejection
160 resp, err := http.Get(proxyServer.URL + "/d/test-doc/ws/test-doc")
161 if err != nil {
162 t.Fatalf("http request failed: %v", err)
163 }
164 defer resp.Body.Close()
165
166 if resp.StatusCode != http.StatusForbidden {
167 t.Errorf("expected status 403, got %d", resp.StatusCode)
168 }
169}
170
171func TestProxyHandler_OnWebsocketUpgradeHook(t *testing.T) {
172 upstream := mockYSweetServer(t)
173 defer upstream.Close()
174 upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
175
176 var hookCalled atomic.Bool
177 var hookReceivedDocID string
178
179 handler := ProxyHandler(upstreamURL,
180 WithOnWebsocketUpgrade(func(ctx context.Context, docID string, clientConn, upstreamConn *websocket.Conn) error {
181 hookCalled.Store(true)
182 hookReceivedDocID = docID
183 return nil
184 }),
185 )
186
187 proxyServer := httptest.NewServer(handler)
188 defer proxyServer.Close()
189
190 proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1) + "/d/test-doc/ws/test-doc"
191 clientConn, _, err := websocket.Dial(context.Background(), proxyURL, nil)
192 if err != nil {
193 t.Fatalf("failed to connect: %v", err)
194 }
195 clientConn.Close(websocket.StatusNormalClosure, "")
196
197 if !hookCalled.Load() {
198 t.Error("OnWebsocketUpgrade hook was not called")
199 }
200
201 if hookReceivedDocID != "test-doc" {
202 t.Errorf("expected docID 'test-doc', got %s", hookReceivedDocID)
203 }
204}
205
206func TestProxyHandler_OnDisconnectHook(t *testing.T) {
207 upstream := mockYSweetServer(t)
208 defer upstream.Close()
209 upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
210
211 var hookCalled atomic.Bool
212 var hookReceivedDocID string
213 var receivedReason DisconnectReason
214
215 done := make(chan struct{})
216 handler := ProxyHandler(upstreamURL,
217 WithOnDisconnect(func(ctx context.Context, docID string, reason DisconnectReason) error {
218 hookCalled.Store(true)
219 hookReceivedDocID = docID
220 receivedReason = reason
221 close(done)
222 return nil
223 }),
224 )
225
226 proxyServer := httptest.NewServer(handler)
227 defer proxyServer.Close()
228
229 proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1) + "/d/test-doc/ws/test-doc"
230 clientConn, _, err := websocket.Dial(context.Background(), proxyURL, nil)
231 if err != nil {
232 t.Fatalf("failed to connect: %v", err)
233 }
234
235 // Close connection from client side
236 clientConn.Close(websocket.StatusNormalClosure, "")
237
238 // Wait for disconnect hook to be called
239 select {
240 case <-done:
241 // Hook was called
242 case <-time.After(2 * time.Second):
243 t.Fatal("timeout waiting for OnDisconnect hook")
244 }
245
246 if !hookCalled.Load() {
247 t.Error("OnDisconnect hook was not called")
248 }
249
250 if hookReceivedDocID != "test-doc" {
251 t.Errorf("expected docID 'test-doc', got %s", hookReceivedDocID)
252 }
253
254 // Reason should be either ClientClosed or Error depending on close timing
255 if receivedReason != DisconnectReasonClientClosed && receivedReason != DisconnectReasonError {
256 t.Errorf("unexpected disconnect reason: %v", receivedReason)
257 }
258}
259
260func TestExtractDocID(t *testing.T) {
261 tests := []struct {
262 path string
263 expected string
264 }{
265 {"/d/my-doc/ws/my-doc", "my-doc"},
266 {"/d/my-doc/ws/my-doc/extra", "my-doc"},
267 {"/d/my-doc/ws/different-doc", ""}, // mismatched doc IDs
268 {"/doc/ws/my-doc", ""}, // legacy pattern not supported
269 {"/ws/my-doc", ""}, // single-doc pattern not supported
270 {"/other/path", ""},
271 {"/", ""},
272 }
273
274 for _, tt := range tests {
275 t.Run(tt.path, func(t *testing.T) {
276 req := httptest.NewRequest("GET", tt.path, nil)
277 config := &proxyConfig{} // No custom docIDFunc, uses default
278 result := extractDocID(req, config)
279 if result != tt.expected {
280 t.Errorf("extractDocID(%q) = %q, want %q", tt.path, result, tt.expected)
281 }
282 })
283 }
284}
285
286func TestExtractDocID_WithCustomFunc(t *testing.T) {
287 // Test with custom DocIDFunc
288 config := &proxyConfig{
289 docIDFunc: func(r *http.Request) string {
290 return r.PathValue("docID")
291 },
292 }
293
294 req := httptest.NewRequest("GET", "/ws/my-custom-doc", nil)
295 // Simulate Go 1.22+ path value
296 req.SetPathValue("docID", "my-custom-doc")
297
298 result := extractDocID(req, config)
299 if result != "my-custom-doc" {
300 t.Errorf("extractDocID with custom func = %q, want %q", result, "my-custom-doc")
301 }
302}
303
304func TestBuildUpstreamURL(t *testing.T) {
305 tests := []struct {
306 baseURL string
307 docID string
308 query url.Values
309 expected string
310 }{
311 {
312 baseURL: "ws://localhost:8080",
313 docID: "test-doc",
314 query: url.Values{},
315 expected: "ws://localhost:8080/d/test-doc/ws/test-doc",
316 },
317 {
318 baseURL: "ws://localhost:8080/",
319 docID: "test-doc",
320 query: url.Values{},
321 expected: "ws://localhost:8080/d/test-doc/ws/test-doc",
322 },
323 {
324 baseURL: "ws://localhost:8080",
325 docID: "test-doc",
326 query: url.Values{"extra": []string{"value"}},
327 expected: "ws://localhost:8080/d/test-doc/ws/test-doc?extra=value",
328 },
329 {
330 baseURL: "ws://localhost:8080",
331 docID: "test-doc",
332 query: url.Values{"token": []string{"abc123"}},
333 expected: "ws://localhost:8080/d/test-doc/ws/test-doc?token=abc123",
334 },
335 }
336
337 for _, tt := range tests {
338 t.Run(tt.baseURL, func(t *testing.T) {
339 result := buildUpstreamURL(tt.baseURL, tt.docID, tt.query)
340 if result != tt.expected {
341 t.Errorf("buildUpstreamURL(%q, %q, %v) = %q, want %q",
342 tt.baseURL, tt.docID, tt.query, result, tt.expected)
343 }
344 })
345 }
346}
347
348func TestProxyHandler_WithContext(t *testing.T) {
349 upstream := mockYSweetServer(t)
350 defer upstream.Close()
351 upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
352
353 // Create master context
354 masterCtx, cancel := context.WithCancel(context.Background())
355
356 disconnectHookCalled := make(chan struct{})
357 var disconnectDocID string
358 var disconnectReason DisconnectReason
359
360 handler := ProxyHandler(upstreamURL,
361 WithContext(masterCtx),
362 WithOnDisconnect(func(ctx context.Context, docID string, reason DisconnectReason) error {
363 disconnectDocID = docID
364 disconnectReason = reason
365 close(disconnectHookCalled)
366 return nil
367 }),
368 )
369
370 proxyServer := httptest.NewServer(handler)
371 defer proxyServer.Close()
372
373 // Connect client
374 proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1) + "/d/test-doc/ws/test-doc"
375 clientConn, _, err := websocket.Dial(context.Background(), proxyURL, nil)
376 if err != nil {
377 t.Fatalf("failed to connect: %v", err)
378 }
379 defer clientConn.Close(websocket.StatusNormalClosure, "")
380
381 // Cancel master context - should close connection
382 cancel()
383
384 // Wait for disconnect hook to fire (with timeout longer than drain timeout)
385 select {
386 case <-disconnectHookCalled:
387 // Hook was called
388 case <-time.After(10 * time.Second):
389 t.Fatal("timeout waiting for disconnect hook")
390 }
391
392 if disconnectDocID != "test-doc" {
393 t.Errorf("expected docID 'test-doc', got %s", disconnectDocID)
394 }
395
396 if disconnectReason != DisconnectReasonContextCancelled {
397 t.Errorf("expected reason DisconnectReasonContextCancelled, got %v", disconnectReason)
398 }
399
400 // Try to read - should fail because connection was closed
401 ctx, cancel2 := context.WithTimeout(context.Background(), 500*time.Millisecond)
402 defer cancel2()
403
404 _, _, err = clientConn.Reader(ctx)
405 if err == nil {
406 t.Error("expected connection to be closed after master context cancelled")
407 }
408}
409
410func TestNormalizeTargetURL(t *testing.T) {
411 tests := []struct {
412 input string
413 expected string
414 }{
415 {"http://target.com", "ws://target.com"},
416 {"http://target.com:8080", "ws://target.com:8080"},
417 {"https://target.com", "wss://target.com"},
418 {"https://target.com:8443", "wss://target.com:8443"},
419 {"ws://target.com", "ws://target.com"},
420 {"wss://target.com", "wss://target.com"},
421 {"target.com", "ws://target.com"},
422 {"target.com:8080", "ws://target.com:8080"},
423 {"192.168.1.1:8080", "ws://192.168.1.1:8080"},
424 {" http://target.com ", "ws://target.com"}, // whitespace trimmed
425 }
426
427 for _, tt := range tests {
428 t.Run(tt.input, func(t *testing.T) {
429 result := normalizeTargetURL(tt.input)
430 if result != tt.expected {
431 t.Errorf("normalizeTargetURL(%q) = %q, want %q",
432 tt.input, result, tt.expected)
433 }
434 })
435 }
436}