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}