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}