scooter  ·  2026-05-28

proxy.go

  1package ysweet
  2
  3import (
  4	"context"
  5	"fmt"
  6	"io"
  7	"log/slog"
  8	"net/http"
  9	"net/url"
 10	"strings"
 11	"sync"
 12	"sync/atomic"
 13	"time"
 14
 15	"github.com/coder/websocket"
 16)
 17
 18// DisconnectReason indicates why the connection was closed
 19type DisconnectReason int
 20
 21const (
 22	DisconnectReasonClientClosed DisconnectReason = iota
 23	DisconnectReasonUpstreamClosed
 24	DisconnectReasonError
 25	DisconnectReasonContextCancelled
 26)
 27
 28// normalizeTargetURL normalizes the target URL to a websocket scheme.
 29// Accepts: http://, https://, ws://, wss://, or no scheme.
 30// Returns ws:// for http:// or no scheme (assumes internal network, no TLS).
 31// Returns wss:// for https://.
 32func normalizeTargetURL(targetURL string) string {
 33	// Trim whitespace
 34	targetURL = strings.TrimSpace(targetURL)
 35
 36	// Check for existing scheme
 37	if strings.HasPrefix(targetURL, "wss://") {
 38		return targetURL
 39	}
 40	if strings.HasPrefix(targetURL, "ws://") {
 41		return targetURL
 42	}
 43	if strings.HasPrefix(targetURL, "https://") {
 44		return "wss://" + targetURL[len("https://"):]
 45	}
 46	if strings.HasPrefix(targetURL, "http://") {
 47		return "ws://" + targetURL[len("http://"):]
 48	}
 49
 50	// No scheme provided - assume ws:// for internal network (no TLS)
 51	return "ws://" + targetURL
 52}
 53
 54// DocIDFunc extracts the document ID from the HTTP request.
 55// Users can provide a custom function to extract docID from their URL scheme.
 56type DocIDFunc func(r *http.Request) string
 57
 58// proxyConfig holds configuration for the websocket proxy
 59type proxyConfig struct {
 60	targetURL          string
 61	targetAuthToken    string
 62	onConnect          OnConnectHook
 63	onWebsocketUpgrade OnWebsocketUpgradeHook
 64	onDisconnect       OnDisconnectHook
 65	docIDFunc          DocIDFunc
 66	masterCtx          context.Context
 67	allowedOrigins     []string
 68}
 69
 70// OnConnectHook is called when a client connects, before websocket upgrade.
 71// The docID is extracted from the URL using the configured DocIDFunc (or default).
 72// Return an error to reject the connection.
 73type OnConnectHook func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error
 74
 75// OnWebsocketUpgradeHook is called after successful websocket upgrade to upstream.
 76// The docID is passed so you can take action based on which document is being synced.
 77type OnWebsocketUpgradeHook func(ctx context.Context, docID string, clientConn, upstreamConn *websocket.Conn) error
 78
 79// OnDisconnectHook is called when the connection closes.
 80// The docID identifies which document connection closed.
 81type OnDisconnectHook func(ctx context.Context, docID string, reason DisconnectReason) error
 82
 83// ProxyOption configures the websocket proxy
 84type ProxyOption func(*proxyConfig)
 85
 86// WithOnConnect sets a hook called when a client connects (before websocket upgrade).
 87// The hook receives the HTTP request and can reject the connection by returning an error.
 88func WithOnConnect(fn OnConnectHook) ProxyOption {
 89	return func(c *proxyConfig) {
 90		c.onConnect = fn
 91	}
 92}
 93
 94// WithOnWebsocketUpgrade sets a hook called after successful websocket upgrade.
 95// Both client and upstream connections are available for inspection.
 96func WithOnWebsocketUpgrade(fn OnWebsocketUpgradeHook) ProxyOption {
 97	return func(c *proxyConfig) {
 98		c.onWebsocketUpgrade = fn
 99	}
100}
101
102// WithOnDisconnect sets a hook called when a client disconnects.
103func WithOnDisconnect(fn OnDisconnectHook) ProxyOption {
104	return func(c *proxyConfig) {
105		c.onDisconnect = fn
106	}
107}
108
109// WithTargetAuthToken sets the authentication token for the upstream y-sweet server.
110func WithTargetAuthToken(token string) ProxyOption {
111	return func(c *proxyConfig) {
112		c.targetAuthToken = token
113	}
114}
115
116// WithDocIDFunc sets a custom function to extract the document ID from requests.
117// This allows you to define your own URL scheme. The docID is used to build the
118// upstream URL (translated to /d/{docID}/ws/{docID}) and passed to all hooks.
119//
120// Example:
121//
122//	mux := http.NewServeMux()
123//	mux.Handle("/ws/{docID}", ysweet.ProxyHandler("ws://upstream:8080",
124//	    ysweet.WithDocIDFunc(func(r *http.Request) string {
125//	        return r.PathValue("docID")
126//	    }),
127//	))
128func WithDocIDFunc(fn DocIDFunc) ProxyOption {
129	return func(c *proxyConfig) {
130		c.docIDFunc = fn
131	}
132}
133
134// WithContext sets a master context for the proxy handler.
135// When this context is cancelled, all active websocket connections will be closed
136// with DisconnectReasonContextCancelled. This is useful for graceful shutdown scenarios.
137// The proxy will wait 5 seconds for connections to drain to the upstream server
138// before forcefully closing them.
139//
140// Example:
141//
142//	ctx, cancel := context.WithCancel(context.Background())
143//	defer cancel()
144//
145//	handler := ysweet.ProxyHandler("ws://upstream:8080",
146//	    ysweet.WithContext(ctx),
147//	)
148//
149//	// Later, during server shutdown:
150//	cancel() // Closes all active proxy connections
151func WithContext(ctx context.Context) ProxyOption {
152	return func(c *proxyConfig) {
153		c.masterCtx = ctx
154	}
155}
156
157// WithAllowedOrigins sets the allowed origin patterns for websocket connections.
158// Patterns can include wildcards, e.g., "*" to allow all origins, or "*.example.com"
159// to allow subdomains. If not specified, the websocket library's default behavior
160// applies (same-origin only).
161//
162// Example:
163//
164//	handler := ysweet.ProxyHandler("ws://upstream:8080",
165//	    ysweet.WithAllowedOrigins("*"), // Allow all origins (development only)
166//	)
167//
168//	handler := ysweet.ProxyHandler("ws://upstream:8080",
169//	    ysweet.WithAllowedOrigins("*.example.com", "localhost:*"),
170//	)
171func WithAllowedOrigins(patterns ...string) ProxyOption {
172	return func(c *proxyConfig) {
173		c.allowedOrigins = patterns
174	}
175}
176
177// mergeContexts returns a context that cancels when either ctx1 or ctx2 is cancelled.
178// The returned context's Done channel is closed when either parent is cancelled.
179func mergeContexts(ctx1, ctx2 context.Context) (context.Context, context.CancelFunc) {
180	ctx, cancel := context.WithCancel(context.Background())
181
182	go func() {
183		select {
184		case <-ctx1.Done():
185		case <-ctx2.Done():
186		}
187		cancel()
188	}()
189
190	return ctx, cancel
191}
192
193// ProxyHandler returns an HTTP handler that proxies websocket connections to a y-sweet server.
194//
195// The targetURL can be specified in various formats:
196//   - "ws://host:port" or "wss://host:port" (websocket schemes, used as-is)
197//   - "http://host:port" → converted to "ws://host:port"
198//   - "https://host:port" → converted to "wss://host:port"
199//   - "host:port" → converted to "ws://host:port" (assumes internal network, no TLS)
200//
201// The handler:
202// 1. Extracts docID from the request (using custom DocIDFunc or default pattern)
203// 2. Calls the OnConnect hook (if configured) - return error to reject
204// 3. Accepts the websocket upgrade from the client
205// 4. Connects to the upstream y-sweet server
206// 5. Calls the OnWebsocketUpgrade hook (if configured)
207// 6. Proxies messages bidirectionally until disconnect
208// 7. Calls the OnDisconnect hook (if configured)
209//
210// Example usage:
211//
212//	handler := ysweet.ProxyHandler("target.com:8080",  // Uses ws://
213//	    ysweet.WithTargetAuthToken("my-token"),
214//	    ysweet.WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
215//	        log.Printf("Client connecting to doc %s: %s", docID, r.RemoteAddr)
216//	        return nil
217//	    }),
218//	)
219//	http.Handle("/d/", handler)
220func ProxyHandler(targetURL string, opts ...ProxyOption) http.Handler {
221	config := &proxyConfig{
222		targetURL: normalizeTargetURL(targetURL),
223	}
224	for _, opt := range opts {
225		opt(config)
226	}
227
228	return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
229		// Extract docID using configured function or default
230		docID := extractDocID(r, config)
231		if docID == "" {
232			http.Error(w, "missing doc_id", http.StatusBadRequest)
233			return
234		}
235
236		// Call OnConnect hook if configured
237		if config.onConnect != nil {
238			if err := config.onConnect(r.Context(), docID, w, r); err != nil {
239				// Hook rejected the connection
240				http.Error(w, err.Error(), http.StatusForbidden)
241				return
242			}
243		}
244
245		// Handle the websocket proxy
246		handleWebsocketProxy(w, r, config, docID)
247	})
248}
249
250// handleWebsocketProxy handles the websocket upgrade and proxy logic
251func handleWebsocketProxy(w http.ResponseWriter, r *http.Request, config *proxyConfig, docID string) {
252	ctx := r.Context()
253
254	slog.Debug("handleWebsocketProxy",
255		"doc_id", docID,
256		"method", r.Method,
257		"url", r.URL.String(),
258		"remote_addr", r.RemoteAddr,
259		"origin", r.Header.Get("Origin"),
260		"upgrade", r.Header.Get("Upgrade"),
261		"connection", r.Header.Get("Connection"),
262	)
263
264	// Merge with master context if configured
265	if config.masterCtx != nil {
266		var cancel context.CancelFunc
267		ctx, cancel = mergeContexts(ctx, config.masterCtx)
268		defer cancel()
269	}
270
271	// Accept websocket connection from client
272	// Use configured allowed origins, or nil for same-origin default
273	wsOpts := &websocket.AcceptOptions{}
274	if len(config.allowedOrigins) > 0 {
275		wsOpts.OriginPatterns = config.allowedOrigins
276	}
277	clientConn, err := websocket.Accept(w, r, wsOpts)
278	if err != nil {
279		slog.Debug("websocket accept failed", "doc_id", docID, "error", err)
280		// Connection already rejected, nothing more to do
281		return
282	}
283	slog.Debug("websocket accept succeeded", "doc_id", docID)
284	defer clientConn.Close(websocket.StatusNormalClosure, "")
285
286	// Build upstream URL (preserves query parameters including token)
287	upstreamURL := buildUpstreamURL(config.targetURL, docID, r.URL.Query())
288	slog.Debug("connecting to upstream", "doc_id", docID, "url", upstreamURL)
289
290	// Connect to upstream y-sweet server
291	upstreamOpts := &websocket.DialOptions{}
292	// Use Authorization header if configured, otherwise rely on query params
293	if config.targetAuthToken != "" {
294		upstreamOpts.HTTPHeader = http.Header{
295			"Authorization": []string{config.targetAuthToken},
296		}
297	}
298
299	upstreamConn, _, err := websocket.Dial(ctx, upstreamURL, upstreamOpts)
300	if err != nil {
301		slog.Debug("upstream dial failed", "doc_id", docID, "url", upstreamURL, "error", err)
302		// Check if this was due to context cancellation
303		if ctx.Err() != nil {
304			slog.Debug("upstream dial failed due to context cancellation", "doc_id", docID)
305			clientConn.Close(websocket.StatusNormalClosure, "")
306			if config.onDisconnect != nil {
307				config.onDisconnect(context.Background(), docID, DisconnectReasonContextCancelled)
308			}
309			return
310		}
311		slog.Debug("closing client connection due to upstream dial failure", "doc_id", docID)
312		clientConn.Close(websocket.StatusInternalError, "upstream connection failed")
313		return
314	}
315	slog.Debug("upstream dial succeeded", "doc_id", docID)
316	defer upstreamConn.Close(websocket.StatusNormalClosure, "")
317
318	// Call OnWebsocketUpgrade hook if configured
319	if config.onWebsocketUpgrade != nil {
320		if err := config.onWebsocketUpgrade(ctx, docID, clientConn, upstreamConn); err != nil {
321			clientConn.Close(websocket.StatusInternalError, "upgrade hook failed")
322			upstreamConn.Close(websocket.StatusNormalClosure, "")
323			return
324		}
325	}
326
327	// Start proxying messages
328	proxyConnections(ctx, clientConn, upstreamConn, config, docID)
329}
330
331// extractDocID extracts the document ID from the request.
332// Uses the configured DocIDFunc if provided, otherwise extracts from the
333// standard /d/{doc_id}/ws/{doc_id} pattern.
334func extractDocID(r *http.Request, config *proxyConfig) string {
335	// Use custom extractor if configured
336	if config.docIDFunc != nil {
337		return config.docIDFunc(r)
338	}
339
340	// Default: extract from /d/{doc_id}/ws/{doc_id} pattern
341	path := r.URL.Path
342	if idx := strings.Index(path, "/d/"); idx != -1 {
343		rest := path[idx+3:] // Skip "/d/"
344		if endIdx := strings.Index(rest, "/ws/"); endIdx != -1 {
345			docID := rest[:endIdx]
346			// The second doc_id comes after "/ws/"
347			secondPart := rest[endIdx+4:] // Skip "/ws/"
348			// Extract just the second doc_id (stop at next slash)
349			if slashIdx := strings.Index(secondPart, "/"); slashIdx != -1 {
350				secondPart = secondPart[:slashIdx]
351			}
352			// For Yjs compatibility, both doc IDs should match
353			if secondPart == docID {
354				return docID
355			}
356		}
357	}
358
359	return ""
360}
361
362// buildUpstreamURL constructs the upstream websocket URL.
363// Preserves all query parameters (including token) from the original request.
364func buildUpstreamURL(baseURL, docID string, query url.Values) string {
365	// Remove trailing slash from baseURL
366	baseURL = strings.TrimSuffix(baseURL, "/")
367
368	// Build the path using y-sweet's expected format: /d/{doc_id}/ws/{doc_id}
369	path := fmt.Sprintf("/d/%s/ws/%s", docID, docID)
370
371	// Reconstruct the URL
372	upstreamURL := baseURL + path
373
374	// Add all query parameters (including token) to upstream URL
375	if len(query) > 0 {
376		upstreamURL = upstreamURL + "?" + query.Encode()
377	}
378
379	return upstreamURL
380}
381
382// proxyConnections proxies messages bidirectionally between client and upstream.
383// It runs until either connection closes or the context is cancelled.
384func proxyConnections(ctx context.Context, clientConn, upstreamConn *websocket.Conn, config *proxyConfig, docID string) {
385	slog.Debug("proxy connections started", "doc_id", docID)
386	
387	// Create a cancelable context for this proxy session
388	ctx, cancel := context.WithCancel(ctx)
389	defer cancel()
390
391	// Channel to signal when either direction closes
392	done := make(chan struct{})
393	var disconnectReason DisconnectReason
394	var disconnectMu sync.Mutex
395	contextWasCancelled := int32(0)
396
397	// Proxy client -> upstream
398	go func() {
399		defer close(done)
400		err := proxyMessages(ctx, clientConn, upstreamConn, "client->upstream")
401		if err != nil {
402			slog.Debug("client to upstream proxy error", "doc_id", docID, "error", err)
403		}
404		// Check if context was cancelled when we return
405		if ctx.Err() != nil {
406			atomic.StoreInt32(&contextWasCancelled, 1)
407		}
408	}()
409
410	// Proxy upstream -> client
411	go func() {
412		err := proxyMessages(ctx, upstreamConn, clientConn, "upstream->client")
413		if err != nil {
414			slog.Debug("upstream to client proxy error", "doc_id", docID, "error", err)
415		}
416		// Check if context was cancelled when we return
417		if ctx.Err() != nil {
418			atomic.StoreInt32(&contextWasCancelled, 1)
419		}
420	}()
421
422	// Wait for either direction to close
423	<-done
424	slog.Debug("proxy connections first direction closed", "doc_id", docID)
425
426	// Check if context was cancelled (indicates shutdown)
427	if atomic.LoadInt32(&contextWasCancelled) == 1 {
428		disconnectMu.Lock()
429		disconnectReason = DisconnectReasonContextCancelled
430		disconnectMu.Unlock()
431
432		// Wait for drain timeout (5 seconds) before force-closing
433		select {
434		case <-done:
435			// Other direction already closed
436			slog.Debug("proxy connections both directions closed", "doc_id", docID)
437		case <-time.After(5 * time.Second):
438			// Drain timeout exceeded, force close
439			slog.Debug("proxy connections drain timeout exceeded", "doc_id", docID)
440		}
441	}
442
443	// Cancel context to stop both goroutines
444	cancel()
445
446	// Close both connections gracefully
447	clientConn.Close(websocket.StatusNormalClosure, "")
448	upstreamConn.Close(websocket.StatusNormalClosure, "")
449
450	slog.Debug("proxy connections ended", "doc_id", docID, "reason", disconnectReason)
451	
452	// Call OnDisconnect hook if configured
453	if config.onDisconnect != nil {
454		disconnectMu.Lock()
455		reason := disconnectReason
456		disconnectMu.Unlock()
457		config.onDisconnect(context.Background(), docID, reason)
458	}
459}
460
461// proxyMessages proxies messages from src to dst until error or context cancellation.
462func proxyMessages(ctx context.Context, src, dst *websocket.Conn, direction string) error {
463	for {
464		// Read message from source
465		msgType, reader, err := src.Reader(ctx)
466		if err != nil {
467			if ctx.Err() != nil {
468				return nil // Context cancelled, not an error
469			}
470			return fmt.Errorf("%s read error: %w", direction, err)
471		}
472
473		// Read the full message
474		data, err := io.ReadAll(reader)
475		if err != nil {
476			return fmt.Errorf("%s read body error: %w", direction, err)
477		}
478
479		// Write message to destination
480		err = dst.Write(ctx, msgType, data)
481		if err != nil {
482			if ctx.Err() != nil {
483				return nil // Context cancelled, not an error
484			}
485			return fmt.Errorf("%s write error: %w", direction, err)
486		}
487	}
488}