add context cancelation to close all active connections
3 files changed,  +228, -14
M examples/proxy_server.go
+22, -2
 1@@ -18,6 +18,7 @@ import (
 2 	"os"
 3 	"os/signal"
 4 	"syscall"
 5+	"time"
 6 
 7 	"github.com/BTBurke/ygo/ysweet"
 8 	"github.com/coder/websocket"
 9@@ -29,8 +30,15 @@ func main() {
10 	listenAddr := getEnv("LISTEN_ADDR", ":3000")
11 	authToken := getEnv("Y_SWEET_AUTH_TOKEN", "")
12 
13+	// Create a master context for graceful shutdown
14+	// When this context is cancelled, all active websocket connections
15+	// will be closed with a 5-second drain period
16+	masterCtx, cancel := context.WithCancel(context.Background())
17+	defer cancel()
18+
19 	// Example 1: Use default URL pattern (/d/{docID}/ws/{docID})
20 	defaultHandler := ysweet.ProxyHandler(upstreamURL,
21+		ysweet.WithContext(masterCtx), // Enable graceful shutdown
22 		ysweet.WithTargetAuthToken(authToken),
23 		ysweet.WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
24 			log.Printf("[CONNECT] Client %s connecting to doc %s", r.RemoteAddr, docID)
25@@ -49,6 +57,7 @@ func main() {
26 
27 	// Example 2: Custom URL pattern using Go 1.22+ path parameters
28 	customHandler := ysweet.ProxyHandler(upstreamURL,
29+		ysweet.WithContext(masterCtx), // Enable graceful shutdown
30 		ysweet.WithDocIDFunc(func(r *http.Request) string {
31 			// Extract docID from custom URL: /ws/{docID}
32 			return r.PathValue("docID")
33@@ -86,8 +95,19 @@ func main() {
34 
35 	go func() {
36 		<-sigChan
37-		log.Println("Shutting down...")
38-		server.Shutdown(context.Background())
39+		log.Println("Shutting down - closing all websocket connections...")
40+
41+		// Cancel master context first to close all websocket connections
42+		// The proxy will wait 5 seconds for connections to drain to upstream
43+		cancel()
44+
45+		// Then shutdown the HTTP server
46+		shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 10*time.Second)
47+		defer shutdownCancel()
48+
49+		if err := server.Shutdown(shutdownCtx); err != nil {
50+			log.Printf("Server shutdown error: %v", err)
51+		}
52 	}()
53 
54 	log.Printf("Starting proxy server on %s", listenAddr)
M ysweet/proxy.go
+116, -12
  1@@ -8,6 +8,8 @@ import (
  2 	"net/url"
  3 	"strings"
  4 	"sync"
  5+	"sync/atomic"
  6+	"time"
  7 
  8 	"github.com/coder/websocket"
  9 )
 10@@ -19,8 +21,35 @@ const (
 11 	DisconnectReasonClientClosed DisconnectReason = iota
 12 	DisconnectReasonUpstreamClosed
 13 	DisconnectReasonError
 14+	DisconnectReasonContextCancelled
 15 )
 16 
 17+// normalizeTargetURL normalizes the target URL to a websocket scheme.
 18+// Accepts: http://, https://, ws://, wss://, or no scheme.
 19+// Returns ws:// for http:// or no scheme (assumes internal network, no TLS).
 20+// Returns wss:// for https://.
 21+func normalizeTargetURL(targetURL string) string {
 22+	// Trim whitespace
 23+	targetURL = strings.TrimSpace(targetURL)
 24+
 25+	// Check for existing scheme
 26+	if strings.HasPrefix(targetURL, "wss://") {
 27+		return targetURL
 28+	}
 29+	if strings.HasPrefix(targetURL, "ws://") {
 30+		return targetURL
 31+	}
 32+	if strings.HasPrefix(targetURL, "https://") {
 33+		return "wss://" + targetURL[len("https://"):]
 34+	}
 35+	if strings.HasPrefix(targetURL, "http://") {
 36+		return "ws://" + targetURL[len("http://"):]
 37+	}
 38+
 39+	// No scheme provided - assume ws:// for internal network (no TLS)
 40+	return "ws://" + targetURL
 41+}
 42+
 43 // DocIDFunc extracts the document ID from the HTTP request.
 44 // Users can provide a custom function to extract docID from their URL scheme.
 45 type DocIDFunc func(r *http.Request) string
 46@@ -33,6 +62,7 @@ type proxyConfig struct {
 47 	onWebsocketUpgrade OnWebsocketUpgradeHook
 48 	onDisconnect       OnDisconnectHook
 49 	docIDFunc          DocIDFunc
 50+	masterCtx          context.Context
 51 }
 52 
 53 // OnConnectHook is called when a client connects, before websocket upgrade.
 54@@ -99,8 +129,53 @@ func WithDocIDFunc(fn DocIDFunc) ProxyOption {
 55 	}
 56 }
 57 
 58+// WithContext sets a master context for the proxy handler.
 59+// When this context is cancelled, all active websocket connections will be closed
 60+// with DisconnectReasonContextCancelled. This is useful for graceful shutdown scenarios.
 61+// The proxy will wait 5 seconds for connections to drain to the upstream server
 62+// before forcefully closing them.
 63+//
 64+// Example:
 65+//
 66+//	ctx, cancel := context.WithCancel(context.Background())
 67+//	defer cancel()
 68+//
 69+//	handler := ysweet.ProxyHandler("ws://upstream:8080",
 70+//	    ysweet.WithContext(ctx),
 71+//	)
 72+//
 73+//	// Later, during server shutdown:
 74+//	cancel() // Closes all active proxy connections
 75+func WithContext(ctx context.Context) ProxyOption {
 76+	return func(c *proxyConfig) {
 77+		c.masterCtx = ctx
 78+	}
 79+}
 80+
 81+// mergeContexts returns a context that cancels when either ctx1 or ctx2 is cancelled.
 82+// The returned context's Done channel is closed when either parent is cancelled.
 83+func mergeContexts(ctx1, ctx2 context.Context) (context.Context, context.CancelFunc) {
 84+	ctx, cancel := context.WithCancel(context.Background())
 85+
 86+	go func() {
 87+		select {
 88+		case <-ctx1.Done():
 89+		case <-ctx2.Done():
 90+		}
 91+		cancel()
 92+	}()
 93+
 94+	return ctx, cancel
 95+}
 96+
 97 // ProxyHandler returns an HTTP handler that proxies websocket connections to a y-sweet server.
 98 //
 99+// The targetURL can be specified in various formats:
100+//   - "ws://host:port" or "wss://host:port" (websocket schemes, used as-is)
101+//   - "http://host:port" → converted to "ws://host:port"
102+//   - "https://host:port" → converted to "wss://host:port"
103+//   - "host:port" → converted to "ws://host:port" (assumes internal network, no TLS)
104+//
105 // The handler:
106 // 1. Extracts docID from the request (using custom DocIDFunc or default pattern)
107 // 2. Calls the OnConnect hook (if configured) - return error to reject
108@@ -112,7 +187,7 @@ func WithDocIDFunc(fn DocIDFunc) ProxyOption {
109 //
110 // Example usage:
111 //
112-//	handler := ysweet.ProxyHandler("ws://localhost:8080",
113+//	handler := ysweet.ProxyHandler("target.com:8080",  // Uses ws://
114 //	    ysweet.WithTargetAuthToken("my-token"),
115 //	    ysweet.WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
116 //	        log.Printf("Client connecting to doc %s: %s", docID, r.RemoteAddr)
117@@ -122,7 +197,7 @@ func WithDocIDFunc(fn DocIDFunc) ProxyOption {
118 //	http.Handle("/d/", handler)
119 func ProxyHandler(targetURL string, opts ...ProxyOption) http.Handler {
120 	config := &proxyConfig{
121-		targetURL: targetURL,
122+		targetURL: normalizeTargetURL(targetURL),
123 	}
124 	for _, opt := range opts {
125 		opt(config)
126@@ -154,6 +229,13 @@ func ProxyHandler(targetURL string, opts ...ProxyOption) http.Handler {
127 func handleWebsocketProxy(w http.ResponseWriter, r *http.Request, config *proxyConfig, docID string) {
128 	ctx := r.Context()
129 
130+	// Merge with master context if configured
131+	if config.masterCtx != nil {
132+		var cancel context.CancelFunc
133+		ctx, cancel = mergeContexts(ctx, config.masterCtx)
134+		defer cancel()
135+	}
136+
137 	// Accept websocket connection from client
138 	wsOpts := &websocket.AcceptOptions{}
139 	clientConn, err := websocket.Accept(w, r, wsOpts)
140@@ -177,6 +259,14 @@ func handleWebsocketProxy(w http.ResponseWriter, r *http.Request, config *proxyC
141 
142 	upstreamConn, _, err := websocket.Dial(ctx, upstreamURL, upstreamOpts)
143 	if err != nil {
144+		// Check if this was due to context cancellation
145+		if ctx.Err() != nil {
146+			clientConn.Close(websocket.StatusNormalClosure, "")
147+			if config.onDisconnect != nil {
148+				config.onDisconnect(context.Background(), docID, DisconnectReasonContextCancelled)
149+			}
150+			return
151+		}
152 		clientConn.Close(websocket.StatusInternalError, "upstream connection failed")
153 		return
154 	}
155@@ -257,31 +347,45 @@ func proxyConnections(ctx context.Context, clientConn, upstreamConn *websocket.C
156 	done := make(chan struct{})
157 	var disconnectReason DisconnectReason
158 	var disconnectMu sync.Mutex
159+	contextWasCancelled := int32(0)
160 
161 	// Proxy client -> upstream
162 	go func() {
163 		defer close(done)
164-		err := proxyMessages(ctx, clientConn, upstreamConn, "client->upstream")
165-		if err != nil {
166-			disconnectMu.Lock()
167-			disconnectReason = DisconnectReasonError
168-			disconnectMu.Unlock()
169+		proxyMessages(ctx, clientConn, upstreamConn, "client->upstream")
170+		// Check if context was cancelled when we return
171+		if ctx.Err() != nil {
172+			atomic.StoreInt32(&contextWasCancelled, 1)
173 		}
174 	}()
175 
176 	// Proxy upstream -> client
177 	go func() {
178-		err := proxyMessages(ctx, upstreamConn, clientConn, "upstream->client")
179-		if err != nil {
180-			disconnectMu.Lock()
181-			disconnectReason = DisconnectReasonError
182-			disconnectMu.Unlock()
183+		proxyMessages(ctx, upstreamConn, clientConn, "upstream->client")
184+		// Check if context was cancelled when we return
185+		if ctx.Err() != nil {
186+			atomic.StoreInt32(&contextWasCancelled, 1)
187 		}
188 	}()
189 
190 	// Wait for either direction to close
191 	<-done
192 
193+	// Check if context was cancelled (indicates shutdown)
194+	if atomic.LoadInt32(&contextWasCancelled) == 1 {
195+		disconnectMu.Lock()
196+		disconnectReason = DisconnectReasonContextCancelled
197+		disconnectMu.Unlock()
198+
199+		// Wait for drain timeout (5 seconds) before force-closing
200+		select {
201+		case <-done:
202+			// Other direction already closed
203+		case <-time.After(5 * time.Second):
204+			// Drain timeout exceeded, force close
205+		}
206+	}
207+
208 	// Cancel context to stop both goroutines
209 	cancel()
210 
M ysweet/proxy_test.go
+90, -0
 1@@ -344,3 +344,93 @@ func TestBuildUpstreamURL(t *testing.T) {
 2 		})
 3 	}
 4 }
 5+
 6+func TestProxyHandler_WithContext(t *testing.T) {
 7+	upstream := mockYSweetServer(t)
 8+	defer upstream.Close()
 9+	upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
10+
11+	// Create master context
12+	masterCtx, cancel := context.WithCancel(context.Background())
13+
14+	disconnectHookCalled := make(chan struct{})
15+	var disconnectDocID string
16+	var disconnectReason DisconnectReason
17+
18+	handler := ProxyHandler(upstreamURL,
19+		WithContext(masterCtx),
20+		WithOnDisconnect(func(ctx context.Context, docID string, reason DisconnectReason) error {
21+			disconnectDocID = docID
22+			disconnectReason = reason
23+			close(disconnectHookCalled)
24+			return nil
25+		}),
26+	)
27+
28+	proxyServer := httptest.NewServer(handler)
29+	defer proxyServer.Close()
30+
31+	// Connect client
32+	proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1) + "/d/test-doc/ws/test-doc"
33+	clientConn, _, err := websocket.Dial(context.Background(), proxyURL, nil)
34+	if err != nil {
35+		t.Fatalf("failed to connect: %v", err)
36+	}
37+	defer clientConn.Close(websocket.StatusNormalClosure, "")
38+
39+	// Cancel master context - should close connection
40+	cancel()
41+
42+	// Wait for disconnect hook to fire (with timeout longer than drain timeout)
43+	select {
44+	case <-disconnectHookCalled:
45+		// Hook was called
46+	case <-time.After(10 * time.Second):
47+		t.Fatal("timeout waiting for disconnect hook")
48+	}
49+
50+	if disconnectDocID != "test-doc" {
51+		t.Errorf("expected docID 'test-doc', got %s", disconnectDocID)
52+	}
53+
54+	if disconnectReason != DisconnectReasonContextCancelled {
55+		t.Errorf("expected reason DisconnectReasonContextCancelled, got %v", disconnectReason)
56+	}
57+
58+	// Try to read - should fail because connection was closed
59+	ctx, cancel2 := context.WithTimeout(context.Background(), 500*time.Millisecond)
60+	defer cancel2()
61+
62+	_, _, err = clientConn.Reader(ctx)
63+	if err == nil {
64+		t.Error("expected connection to be closed after master context cancelled")
65+	}
66+}
67+
68+func TestNormalizeTargetURL(t *testing.T) {
69+	tests := []struct {
70+		input    string
71+		expected string
72+	}{
73+		{"http://target.com", "ws://target.com"},
74+		{"http://target.com:8080", "ws://target.com:8080"},
75+		{"https://target.com", "wss://target.com"},
76+		{"https://target.com:8443", "wss://target.com:8443"},
77+		{"ws://target.com", "ws://target.com"},
78+		{"wss://target.com", "wss://target.com"},
79+		{"target.com", "ws://target.com"},
80+		{"target.com:8080", "ws://target.com:8080"},
81+		{"192.168.1.1:8080", "ws://192.168.1.1:8080"},
82+		{"  http://target.com  ", "ws://target.com"}, // whitespace trimmed
83+	}
84+
85+	for _, tt := range tests {
86+		t.Run(tt.input, func(t *testing.T) {
87+			result := normalizeTargetURL(tt.input)
88+			if result != tt.expected {
89+				t.Errorf("normalizeTargetURL(%q) = %q, want %q",
90+					tt.input, result, tt.expected)
91+			}
92+		})
93+	}
94+}