3 files changed,
+228,
-14
+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)
+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
+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+}