adds a y-sweet websocket proxy with hooks for connect/disconnect
4 files changed,  +1161, -0
A examples/proxy_server.go
+106, -0
  1@@ -0,0 +1,106 @@
  2+//go:build ignore
  3+// +build ignore
  4+
  5+// This example demonstrates how to use the ysweet proxy handler to create
  6+// a websocket proxy server that forwards connections to a y-sweet server.
  7+//
  8+// Usage:
  9+//
 10+//	go run proxy_server.go
 11+//
 12+// Then connect from a Yjs client to ws://localhost:3000/d/my-doc/ws/my-doc
 13+package main
 14+
 15+import (
 16+	"context"
 17+	"log"
 18+	"net/http"
 19+	"os"
 20+	"os/signal"
 21+	"syscall"
 22+
 23+	"github.com/BTBurke/ygo/ysweet"
 24+	"github.com/coder/websocket"
 25+)
 26+
 27+func main() {
 28+	// Configuration
 29+	upstreamURL := getEnv("Y_SWEET_URL", "ws://localhost:8080")
 30+	listenAddr := getEnv("LISTEN_ADDR", ":3000")
 31+	authToken := getEnv("Y_SWEET_AUTH_TOKEN", "")
 32+
 33+	// Example 1: Use default URL pattern (/d/{docID}/ws/{docID})
 34+	defaultHandler := ysweet.ProxyHandler(upstreamURL,
 35+		ysweet.WithTargetAuthToken(authToken),
 36+		ysweet.WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
 37+			log.Printf("[CONNECT] Client %s connecting to doc %s", r.RemoteAddr, docID)
 38+			// You can now take action based on docID, e.g., check permissions
 39+			return nil
 40+		}),
 41+		ysweet.WithOnWebsocketUpgrade(func(ctx context.Context, docID string, clientConn, upstreamConn *websocket.Conn) error {
 42+			log.Printf("[UPGRADE] Doc %s websocket upgrade successful", docID)
 43+			return nil
 44+		}),
 45+		ysweet.WithOnDisconnect(func(ctx context.Context, docID string, reason ysweet.DisconnectReason) error {
 46+			log.Printf("[DISCONNECT] Doc %s disconnected, reason: %v", docID, reason)
 47+			return nil
 48+		}),
 49+	)
 50+
 51+	// Example 2: Custom URL pattern using Go 1.22+ path parameters
 52+	customHandler := ysweet.ProxyHandler(upstreamURL,
 53+		ysweet.WithDocIDFunc(func(r *http.Request) string {
 54+			// Extract docID from custom URL: /ws/{docID}
 55+			return r.PathValue("docID")
 56+		}),
 57+		ysweet.WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
 58+			log.Printf("[CONNECT] Custom handler - doc %s connecting", docID)
 59+			return nil
 60+		}),
 61+	)
 62+
 63+	// Set up HTTP routes
 64+	mux := http.NewServeMux()
 65+
 66+	// Default y-sweet pattern
 67+	mux.Handle("/d/", defaultHandler)
 68+
 69+	// Custom pattern: /ws/{docID} -> proxies to upstream /d/{docID}/ws/{docID}
 70+	mux.Handle("/ws/{docID}", customHandler)
 71+
 72+	// Health check endpoint
 73+	mux.HandleFunc("/health", func(w http.ResponseWriter, r *http.Request) {
 74+		w.WriteHeader(http.StatusOK)
 75+		w.Write([]byte("ok"))
 76+	})
 77+
 78+	// Create server
 79+	server := &http.Server{
 80+		Addr:    listenAddr,
 81+		Handler: mux,
 82+	}
 83+
 84+	// Handle graceful shutdown
 85+	sigChan := make(chan os.Signal, 1)
 86+	signal.Notify(sigChan, syscall.SIGINT, syscall.SIGTERM)
 87+
 88+	go func() {
 89+		<-sigChan
 90+		log.Println("Shutting down...")
 91+		server.Shutdown(context.Background())
 92+	}()
 93+
 94+	log.Printf("Starting proxy server on %s", listenAddr)
 95+	log.Printf("Proxying to y-sweet at %s", upstreamURL)
 96+
 97+	if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed {
 98+		log.Fatalf("Server error: %v", err)
 99+	}
100+}
101+
102+func getEnv(key, defaultValue string) string {
103+	if value := os.Getenv(key); value != "" {
104+		return value
105+	}
106+	return defaultValue
107+}
A ysweet/proxy.go
+328, -0
  1@@ -0,0 +1,328 @@
  2+package ysweet
  3+
  4+import (
  5+	"context"
  6+	"fmt"
  7+	"io"
  8+	"net/http"
  9+	"net/url"
 10+	"strings"
 11+	"sync"
 12+
 13+	"github.com/coder/websocket"
 14+)
 15+
 16+// DisconnectReason indicates why the connection was closed
 17+type DisconnectReason int
 18+
 19+const (
 20+	DisconnectReasonClientClosed DisconnectReason = iota
 21+	DisconnectReasonUpstreamClosed
 22+	DisconnectReasonError
 23+)
 24+
 25+// DocIDFunc extracts the document ID from the HTTP request.
 26+// Users can provide a custom function to extract docID from their URL scheme.
 27+type DocIDFunc func(r *http.Request) string
 28+
 29+// proxyConfig holds configuration for the websocket proxy
 30+type proxyConfig struct {
 31+	targetURL          string
 32+	targetAuthToken    string
 33+	onConnect          OnConnectHook
 34+	onWebsocketUpgrade OnWebsocketUpgradeHook
 35+	onDisconnect       OnDisconnectHook
 36+	docIDFunc          DocIDFunc
 37+}
 38+
 39+// OnConnectHook is called when a client connects, before websocket upgrade.
 40+// The docID is extracted from the URL using the configured DocIDFunc (or default).
 41+// Return an error to reject the connection.
 42+type OnConnectHook func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error
 43+
 44+// OnWebsocketUpgradeHook is called after successful websocket upgrade to upstream.
 45+// The docID is passed so you can take action based on which document is being synced.
 46+type OnWebsocketUpgradeHook func(ctx context.Context, docID string, clientConn, upstreamConn *websocket.Conn) error
 47+
 48+// OnDisconnectHook is called when the connection closes.
 49+// The docID identifies which document connection closed.
 50+type OnDisconnectHook func(ctx context.Context, docID string, reason DisconnectReason) error
 51+
 52+// ProxyOption configures the websocket proxy
 53+type ProxyOption func(*proxyConfig)
 54+
 55+// WithOnConnect sets a hook called when a client connects (before websocket upgrade).
 56+// The hook receives the HTTP request and can reject the connection by returning an error.
 57+func WithOnConnect(fn OnConnectHook) ProxyOption {
 58+	return func(c *proxyConfig) {
 59+		c.onConnect = fn
 60+	}
 61+}
 62+
 63+// WithOnWebsocketUpgrade sets a hook called after successful websocket upgrade.
 64+// Both client and upstream connections are available for inspection.
 65+func WithOnWebsocketUpgrade(fn OnWebsocketUpgradeHook) ProxyOption {
 66+	return func(c *proxyConfig) {
 67+		c.onWebsocketUpgrade = fn
 68+	}
 69+}
 70+
 71+// WithOnDisconnect sets a hook called when a client disconnects.
 72+func WithOnDisconnect(fn OnDisconnectHook) ProxyOption {
 73+	return func(c *proxyConfig) {
 74+		c.onDisconnect = fn
 75+	}
 76+}
 77+
 78+// WithTargetAuthToken sets the authentication token for the upstream y-sweet server.
 79+func WithTargetAuthToken(token string) ProxyOption {
 80+	return func(c *proxyConfig) {
 81+		c.targetAuthToken = token
 82+	}
 83+}
 84+
 85+// WithDocIDFunc sets a custom function to extract the document ID from requests.
 86+// This allows you to define your own URL scheme. The docID is used to build the
 87+// upstream URL (translated to /d/{docID}/ws/{docID}) and passed to all hooks.
 88+//
 89+// Example:
 90+//
 91+//	mux := http.NewServeMux()
 92+//	mux.Handle("/ws/{docID}", ysweet.ProxyHandler("ws://upstream:8080",
 93+//	    ysweet.WithDocIDFunc(func(r *http.Request) string {
 94+//	        return r.PathValue("docID")
 95+//	    }),
 96+//	))
 97+func WithDocIDFunc(fn DocIDFunc) ProxyOption {
 98+	return func(c *proxyConfig) {
 99+		c.docIDFunc = fn
100+	}
101+}
102+
103+// ProxyHandler returns an HTTP handler that proxies websocket connections to a y-sweet server.
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+// 3. Accepts the websocket upgrade from the client
109+// 4. Connects to the upstream y-sweet server
110+// 5. Calls the OnWebsocketUpgrade hook (if configured)
111+// 6. Proxies messages bidirectionally until disconnect
112+// 7. Calls the OnDisconnect hook (if configured)
113+//
114+// Example usage:
115+//
116+//	handler := ysweet.ProxyHandler("ws://localhost:8080",
117+//	    ysweet.WithTargetAuthToken("my-token"),
118+//	    ysweet.WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
119+//	        log.Printf("Client connecting to doc %s: %s", docID, r.RemoteAddr)
120+//	        return nil
121+//	    }),
122+//	)
123+//	http.Handle("/d/", handler)
124+func ProxyHandler(targetURL string, opts ...ProxyOption) http.Handler {
125+	config := &proxyConfig{
126+		targetURL: targetURL,
127+	}
128+	for _, opt := range opts {
129+		opt(config)
130+	}
131+
132+	return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
133+		// Extract docID using configured function or default
134+		docID := extractDocID(r, config)
135+		if docID == "" {
136+			http.Error(w, "missing doc_id", http.StatusBadRequest)
137+			return
138+		}
139+
140+		// Call OnConnect hook if configured
141+		if config.onConnect != nil {
142+			if err := config.onConnect(r.Context(), docID, w, r); err != nil {
143+				// Hook rejected the connection
144+				http.Error(w, err.Error(), http.StatusForbidden)
145+				return
146+			}
147+		}
148+
149+		// Handle the websocket proxy
150+		handleWebsocketProxy(w, r, config, docID)
151+	})
152+}
153+
154+// handleWebsocketProxy handles the websocket upgrade and proxy logic
155+func handleWebsocketProxy(w http.ResponseWriter, r *http.Request, config *proxyConfig, docID string) {
156+	ctx := r.Context()
157+
158+	// Accept websocket connection from client
159+	wsOpts := &websocket.AcceptOptions{}
160+	clientConn, err := websocket.Accept(w, r, wsOpts)
161+	if err != nil {
162+		// Connection already rejected, nothing more to do
163+		return
164+	}
165+	defer clientConn.Close(websocket.StatusNormalClosure, "")
166+
167+	// Build upstream URL (preserves query parameters including token)
168+	upstreamURL := buildUpstreamURL(config.targetURL, docID, r.URL.Query())
169+
170+	// Connect to upstream y-sweet server
171+	upstreamOpts := &websocket.DialOptions{}
172+	// Use Authorization header if configured, otherwise rely on query params
173+	if config.targetAuthToken != "" {
174+		upstreamOpts.HTTPHeader = http.Header{
175+			"Authorization": []string{config.targetAuthToken},
176+		}
177+	}
178+
179+	upstreamConn, _, err := websocket.Dial(ctx, upstreamURL, upstreamOpts)
180+	if err != nil {
181+		clientConn.Close(websocket.StatusInternalError, "upstream connection failed")
182+		return
183+	}
184+	defer upstreamConn.Close(websocket.StatusNormalClosure, "")
185+
186+	// Call OnWebsocketUpgrade hook if configured
187+	if config.onWebsocketUpgrade != nil {
188+		if err := config.onWebsocketUpgrade(ctx, docID, clientConn, upstreamConn); err != nil {
189+			clientConn.Close(websocket.StatusInternalError, "upgrade hook failed")
190+			upstreamConn.Close(websocket.StatusNormalClosure, "")
191+			return
192+		}
193+	}
194+
195+	// Start proxying messages
196+	proxyConnections(ctx, clientConn, upstreamConn, config, docID)
197+}
198+
199+// extractDocID extracts the document ID from the request.
200+// Uses the configured DocIDFunc if provided, otherwise extracts from the
201+// standard /d/{doc_id}/ws/{doc_id} pattern.
202+func extractDocID(r *http.Request, config *proxyConfig) string {
203+	// Use custom extractor if configured
204+	if config.docIDFunc != nil {
205+		return config.docIDFunc(r)
206+	}
207+
208+	// Default: extract from /d/{doc_id}/ws/{doc_id} pattern
209+	path := r.URL.Path
210+	if idx := strings.Index(path, "/d/"); idx != -1 {
211+		rest := path[idx+3:] // Skip "/d/"
212+		if endIdx := strings.Index(rest, "/ws/"); endIdx != -1 {
213+			docID := rest[:endIdx]
214+			// The second doc_id comes after "/ws/"
215+			secondPart := rest[endIdx+4:] // Skip "/ws/"
216+			// Extract just the second doc_id (stop at next slash)
217+			if slashIdx := strings.Index(secondPart, "/"); slashIdx != -1 {
218+				secondPart = secondPart[:slashIdx]
219+			}
220+			// For Yjs compatibility, both doc IDs should match
221+			if secondPart == docID {
222+				return docID
223+			}
224+		}
225+	}
226+
227+	return ""
228+}
229+
230+// buildUpstreamURL constructs the upstream websocket URL.
231+// Preserves all query parameters (including token) from the original request.
232+func buildUpstreamURL(baseURL, docID string, query url.Values) string {
233+	// Remove trailing slash from baseURL
234+	baseURL = strings.TrimSuffix(baseURL, "/")
235+
236+	// Build the path using y-sweet's expected format: /d/{doc_id}/ws/{doc_id}
237+	path := fmt.Sprintf("/d/%s/ws/%s", docID, docID)
238+
239+	// Reconstruct the URL
240+	upstreamURL := baseURL + path
241+
242+	// Add all query parameters (including token) to upstream URL
243+	if len(query) > 0 {
244+		upstreamURL = upstreamURL + "?" + query.Encode()
245+	}
246+
247+	return upstreamURL
248+}
249+
250+// proxyConnections proxies messages bidirectionally between client and upstream.
251+// It runs until either connection closes or the context is cancelled.
252+func proxyConnections(ctx context.Context, clientConn, upstreamConn *websocket.Conn, config *proxyConfig, docID string) {
253+	// Create a cancelable context for this proxy session
254+	ctx, cancel := context.WithCancel(ctx)
255+	defer cancel()
256+
257+	// Channel to signal when either direction closes
258+	done := make(chan struct{})
259+	var disconnectReason DisconnectReason
260+	var disconnectMu sync.Mutex
261+
262+	// Proxy client -> upstream
263+	go func() {
264+		defer close(done)
265+		err := proxyMessages(ctx, clientConn, upstreamConn, "client->upstream")
266+		if err != nil {
267+			disconnectMu.Lock()
268+			disconnectReason = DisconnectReasonError
269+			disconnectMu.Unlock()
270+		}
271+	}()
272+
273+	// Proxy upstream -> client
274+	go func() {
275+		err := proxyMessages(ctx, upstreamConn, clientConn, "upstream->client")
276+		if err != nil {
277+			disconnectMu.Lock()
278+			disconnectReason = DisconnectReasonError
279+			disconnectMu.Unlock()
280+		}
281+	}()
282+
283+	// Wait for either direction to close
284+	<-done
285+
286+	// Cancel context to stop both goroutines
287+	cancel()
288+
289+	// Close both connections gracefully
290+	clientConn.Close(websocket.StatusNormalClosure, "")
291+	upstreamConn.Close(websocket.StatusNormalClosure, "")
292+
293+	// Call OnDisconnect hook if configured
294+	if config.onDisconnect != nil {
295+		disconnectMu.Lock()
296+		reason := disconnectReason
297+		disconnectMu.Unlock()
298+		config.onDisconnect(context.Background(), docID, reason)
299+	}
300+}
301+
302+// proxyMessages proxies messages from src to dst until error or context cancellation.
303+func proxyMessages(ctx context.Context, src, dst *websocket.Conn, direction string) error {
304+	for {
305+		// Read message from source
306+		msgType, reader, err := src.Reader(ctx)
307+		if err != nil {
308+			if ctx.Err() != nil {
309+				return nil // Context cancelled, not an error
310+			}
311+			return fmt.Errorf("%s read error: %w", direction, err)
312+		}
313+
314+		// Read the full message
315+		data, err := io.ReadAll(reader)
316+		if err != nil {
317+			return fmt.Errorf("%s read body error: %w", direction, err)
318+		}
319+
320+		// Write message to destination
321+		err = dst.Write(ctx, msgType, data)
322+		if err != nil {
323+			if ctx.Err() != nil {
324+				return nil // Context cancelled, not an error
325+			}
326+			return fmt.Errorf("%s write error: %w", direction, err)
327+		}
328+	}
329+}
A ysweet/proxy_integration_test.go
+381, -0
  1@@ -0,0 +1,381 @@
  2+package ysweet_test
  3+
  4+import (
  5+	"context"
  6+	"net/http"
  7+	"net/http/httptest"
  8+	"os"
  9+	"os/exec"
 10+	"strings"
 11+	"syscall"
 12+	"testing"
 13+	"time"
 14+
 15+	"github.com/BTBurke/ygo"
 16+	"github.com/BTBurke/ygo/ysweet"
 17+	"github.com/coder/websocket"
 18+)
 19+
 20+func skipIfNoIntegrationTests(t *testing.T) {
 21+	if os.Getenv("RUN_INTEGRATION_TESTS") == "" {
 22+		t.Skip("Skipping integration test: set RUN_INTEGRATION_TESTS=1 to run")
 23+	}
 24+}
 25+
 26+// TestProxyHandler_Integration tests the proxy handler with a live y-sweet server.
 27+// It starts a y-sweet server, creates a proxy in front of it, and verifies that:
 28+// 1. Two clients can connect through the proxy to the same document
 29+// 2. Updates from one client are received by the other
 30+// 3. Hooks fire correctly on connect and disconnect
 31+func TestProxyHandler_Integration(t *testing.T) {
 32+	skipIfNoIntegrationTests(t)
 33+
 34+	// Start y-sweet server
 35+	ctx, cancel := context.WithCancel(context.Background())
 36+	defer cancel()
 37+
 38+	cmd := exec.CommandContext(ctx, "pnpx", "y-sweet@latest", "serve")
 39+	cmd.SysProcAttr = &syscall.SysProcAttr{
 40+		Setpgid: true,
 41+	}
 42+
 43+	if err := cmd.Start(); err != nil {
 44+		t.Fatalf("failed to start y-sweet server: %v", err)
 45+	}
 46+
 47+	// Give server time to start
 48+	time.Sleep(2 * time.Second)
 49+
 50+	// Create y-sweet client to create a document
 51+	client, err := ysweet.NewClient("http://127.0.0.1:8080")
 52+	if err != nil {
 53+		t.Fatalf("failed to create y-sweet client: %v", err)
 54+	}
 55+
 56+	docID, err := client.NewDoc("")
 57+	if err != nil {
 58+		t.Fatalf("failed to create document: %v", err)
 59+	}
 60+
 61+	// Track hook invocations
 62+	connectHookCalls := make(chan string, 10)
 63+	disconnectHookCalls := make(chan string, 10)
 64+	upgradeHookCalls := make(chan string, 10)
 65+
 66+	// Create proxy handler in front of y-sweet server
 67+	proxyHandler := ysweet.ProxyHandler("ws://127.0.0.1:8080",
 68+		ysweet.WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
 69+			connectHookCalls <- docID
 70+			return nil
 71+		}),
 72+		ysweet.WithOnWebsocketUpgrade(func(ctx context.Context, docID string, clientConn, upstreamConn *websocket.Conn) error {
 73+			upgradeHookCalls <- docID
 74+			return nil
 75+		}),
 76+		ysweet.WithOnDisconnect(func(ctx context.Context, docID string, reason ysweet.DisconnectReason) error {
 77+			disconnectHookCalls <- docID
 78+			return nil
 79+		}),
 80+	)
 81+
 82+	// Start proxy server
 83+	proxyServer := httptest.NewServer(proxyHandler)
 84+	defer proxyServer.Close()
 85+
 86+	// Convert http:// to ws:// for proxy URL
 87+	proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1)
 88+
 89+	// Create first client document
 90+	doc1, err := ygo.NewDoc()
 91+	if err != nil {
 92+		t.Fatalf("failed to create doc1: %v", err)
 93+	}
 94+	defer doc1.Destroy()
 95+
 96+	// Connect first client through proxy
 97+	syncClient1, err := ygo.NewSyncClient(doc1,
 98+		ygo.WithSyncEndpoint(proxyURL+"/d/"+docID+"/ws/"+docID),
 99+	)
100+	if err != nil {
101+		t.Fatalf("failed to create sync client 1: %v", err)
102+	}
103+
104+	connectCtx1, cancel1 := context.WithTimeout(context.Background(), 5*time.Second)
105+	defer cancel1()
106+
107+	if err := syncClient1.Connect(connectCtx1); err != nil {
108+		t.Fatalf("failed to connect client 1: %v", err)
109+	}
110+	defer syncClient1.Close()
111+
112+	// Wait for first client's hooks to fire
113+	select {
114+	case receivedDocID := <-connectHookCalls:
115+		if receivedDocID != docID {
116+			t.Errorf("client 1 connect hook: expected docID %s, got %s", docID, receivedDocID)
117+		}
118+		t.Logf("✓ Client 1 connect hook fired for doc %s", receivedDocID)
119+	case <-time.After(2 * time.Second):
120+		t.Fatal("timeout waiting for client 1 connect hook")
121+	}
122+
123+	select {
124+	case receivedDocID := <-upgradeHookCalls:
125+		if receivedDocID != docID {
126+			t.Errorf("client 1 upgrade hook: expected docID %s, got %s", docID, receivedDocID)
127+		}
128+		t.Logf("✓ Client 1 upgrade hook fired for doc %s", receivedDocID)
129+	case <-time.After(2 * time.Second):
130+		t.Fatal("timeout waiting for client 1 upgrade hook")
131+	}
132+
133+	// Create second client document
134+	doc2, err := ygo.NewDoc()
135+	if err != nil {
136+		t.Fatalf("failed to create doc2: %v", err)
137+	}
138+	defer doc2.Destroy()
139+
140+	// Connect second client through proxy (same document)
141+	syncClient2, err := ygo.NewSyncClient(doc2,
142+		ygo.WithSyncEndpoint(proxyURL+"/d/"+docID+"/ws/"+docID),
143+	)
144+	if err != nil {
145+		t.Fatalf("failed to create sync client 2: %v", err)
146+	}
147+
148+	connectCtx2, cancel2 := context.WithTimeout(context.Background(), 5*time.Second)
149+	defer cancel2()
150+
151+	if err := syncClient2.Connect(connectCtx2); err != nil {
152+		t.Fatalf("failed to connect client 2: %v", err)
153+	}
154+	defer syncClient2.Close()
155+
156+	// Wait for second client's hooks to fire
157+	select {
158+	case receivedDocID := <-connectHookCalls:
159+		if receivedDocID != docID {
160+			t.Errorf("client 2 connect hook: expected docID %s, got %s", docID, receivedDocID)
161+		}
162+		t.Logf("✓ Client 2 connect hook fired for doc %s", receivedDocID)
163+	case <-time.After(2 * time.Second):
164+		t.Fatal("timeout waiting for client 2 connect hook")
165+	}
166+
167+	select {
168+	case receivedDocID := <-upgradeHookCalls:
169+		if receivedDocID != docID {
170+			t.Errorf("client 2 upgrade hook: expected docID %s, got %s", docID, receivedDocID)
171+		}
172+		t.Logf("✓ Client 2 upgrade hook fired for doc %s", receivedDocID)
173+	case <-time.After(2 * time.Second):
174+		t.Fatal("timeout waiting for client 2 upgrade hook")
175+	}
176+
177+	// Give clients time to sync
178+	time.Sleep(500 * time.Millisecond)
179+
180+	// Make an update on client 1
181+	txt1, err := doc1.GetText("content")
182+	if err != nil {
183+		t.Fatalf("failed to get text from doc1: %v", err)
184+	}
185+
186+	err = doc1.WithWriteTransaction(func(txn *ygo.Transaction) error {
187+		return txt1.Insert(txn, 0, "Hello from client 1!")
188+	})
189+	if err != nil {
190+		t.Fatalf("failed to insert text: %v", err)
191+	}
192+
193+	// Wait for sync
194+	time.Sleep(1 * time.Second)
195+
196+	// Verify client 2 received the update
197+	txt2, err := doc2.GetText("content")
198+	if err != nil {
199+		t.Fatalf("failed to get text from doc2: %v", err)
200+	}
201+
202+	var content2 string
203+	err = doc2.WithReadTransaction(func(txn *ygo.Transaction) error {
204+		var err error
205+		content2, err = txt2.String(txn)
206+		return err
207+	})
208+	if err != nil {
209+		t.Fatalf("failed to get string from text: %v", err)
210+	}
211+
212+	if content2 != "Hello from client 1!" {
213+		t.Errorf("expected 'Hello from client 1!', got '%s'", content2)
214+	}
215+	t.Logf("✓ Client 2 received update from client 1: '%s'", content2)
216+
217+	// Make an update on client 2
218+	err = doc2.WithWriteTransaction(func(txn *ygo.Transaction) error {
219+		return txt2.Insert(txn, uint32(len(content2)), " And hello from client 2!")
220+	})
221+	if err != nil {
222+		t.Fatalf("failed to insert text: %v", err)
223+	}
224+
225+	// Wait for sync
226+	time.Sleep(1 * time.Second)
227+
228+	// Verify client 1 received the update
229+	var content1 string
230+	err = doc1.WithReadTransaction(func(txn *ygo.Transaction) error {
231+		var err error
232+		content1, err = txt1.String(txn)
233+		return err
234+	})
235+	if err != nil {
236+		t.Fatalf("failed to get string from text: %v", err)
237+	}
238+
239+	expectedContent := "Hello from client 1! And hello from client 2!"
240+	if content1 != expectedContent {
241+		t.Errorf("expected '%s', got '%s'", expectedContent, content1)
242+	}
243+	t.Logf("✓ Client 1 received update from client 2: '%s'", content1)
244+
245+	// Close first client and verify disconnect hook
246+	syncClient1.Close()
247+
248+	select {
249+	case receivedDocID := <-disconnectHookCalls:
250+		if receivedDocID != docID {
251+			t.Errorf("client 1 disconnect hook: expected docID %s, got %s", docID, receivedDocID)
252+		}
253+		t.Logf("✓ Client 1 disconnect hook fired for doc %s", receivedDocID)
254+	case <-time.After(2 * time.Second):
255+		t.Fatal("timeout waiting for client 1 disconnect hook")
256+	}
257+
258+	// Close second client and verify disconnect hook
259+	syncClient2.Close()
260+
261+	select {
262+	case receivedDocID := <-disconnectHookCalls:
263+		if receivedDocID != docID {
264+			t.Errorf("client 2 disconnect hook: expected docID %s, got %s", docID, receivedDocID)
265+		}
266+		t.Logf("✓ Client 2 disconnect hook fired for doc %s", receivedDocID)
267+	case <-time.After(2 * time.Second):
268+		t.Fatal("timeout waiting for client 2 disconnect hook")
269+	}
270+
271+	t.Logf("✓ Integration test completed successfully")
272+
273+	// Cleanup server
274+	if err := syscall.Kill(-cmd.Process.Pid, syscall.SIGINT); err != nil {
275+		t.Logf("warning: failed to kill server: %v", err)
276+	}
277+}
278+
279+// TestProxyHandler_IntegrationWithCustomURLPattern tests the proxy with a custom URL pattern.
280+func TestProxyHandler_IntegrationWithCustomURLPattern(t *testing.T) {
281+	skipIfNoIntegrationTests(t)
282+
283+	// Start y-sweet server
284+	ctx, cancel := context.WithCancel(context.Background())
285+	defer cancel()
286+
287+	cmd := exec.CommandContext(ctx, "pnpx", "y-sweet@latest", "serve")
288+	cmd.SysProcAttr = &syscall.SysProcAttr{
289+		Setpgid: true,
290+	}
291+
292+	if err := cmd.Start(); err != nil {
293+		t.Fatalf("failed to start y-sweet server: %v", err)
294+	}
295+
296+	// Give server time to start
297+	time.Sleep(2 * time.Second)
298+
299+	// Create y-sweet client to create a document
300+	client, err := ysweet.NewClient("http://127.0.0.1:8080")
301+	if err != nil {
302+		t.Fatalf("failed to create y-sweet client: %v", err)
303+	}
304+
305+	docID, err := client.NewDoc("")
306+	if err != nil {
307+		t.Fatalf("failed to create document: %v", err)
308+	}
309+
310+	// Track hook invocations
311+	connectHookCalls := make(chan string, 10)
312+
313+	// Create proxy handler with custom URL pattern: /api/doc/{docID}
314+	proxyHandler := ysweet.ProxyHandler("ws://127.0.0.1:8080",
315+		ysweet.WithDocIDFunc(func(r *http.Request) string {
316+			// Extract docID from /api/doc/{docID}
317+			path := r.URL.Path
318+			prefix := "/api/doc/"
319+			if strings.HasPrefix(path, prefix) {
320+				docID := path[len(prefix):]
321+				// Remove any trailing segments
322+				if idx := strings.Index(docID, "/"); idx != -1 {
323+					docID = docID[:idx]
324+				}
325+				return docID
326+			}
327+			return ""
328+		}),
329+		ysweet.WithOnConnect(func(ctx context.Context, receivedDocID string, w http.ResponseWriter, r *http.Request) error {
330+			connectHookCalls <- receivedDocID
331+			return nil
332+		}),
333+	)
334+
335+	// Start proxy server
336+	proxyServer := httptest.NewServer(proxyHandler)
337+	defer proxyServer.Close()
338+
339+	// Convert http:// to ws:// for proxy URL
340+	proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1)
341+
342+	// Create client document
343+	doc, err := ygo.NewDoc()
344+	if err != nil {
345+		t.Fatalf("failed to create doc: %v", err)
346+	}
347+	defer doc.Destroy()
348+
349+	// Connect through proxy using custom URL pattern
350+	syncClient, err := ygo.NewSyncClient(doc,
351+		ygo.WithSyncEndpoint(proxyURL+"/api/doc/"+docID),
352+	)
353+	if err != nil {
354+		t.Fatalf("failed to create sync client: %v", err)
355+	}
356+
357+	connectCtx, cancelConnect := context.WithTimeout(context.Background(), 5*time.Second)
358+	defer cancelConnect()
359+
360+	if err := syncClient.Connect(connectCtx); err != nil {
361+		t.Fatalf("failed to connect: %v", err)
362+	}
363+	defer syncClient.Close()
364+
365+	// Wait for connect hook to fire with correct docID
366+	select {
367+	case receivedDocID := <-connectHookCalls:
368+		if receivedDocID != docID {
369+			t.Errorf("connect hook: expected docID %s, got %s", docID, receivedDocID)
370+		}
371+		t.Logf("✓ Connect hook fired for doc %s (via custom URL pattern)", receivedDocID)
372+	case <-time.After(2 * time.Second):
373+		t.Fatal("timeout waiting for connect hook")
374+	}
375+
376+	t.Logf("✓ Custom URL pattern integration test completed successfully")
377+
378+	// Cleanup server
379+	if err := syscall.Kill(-cmd.Process.Pid, syscall.SIGINT); err != nil {
380+		t.Logf("warning: failed to kill server: %v", err)
381+	}
382+}
A ysweet/proxy_test.go
+346, -0
  1@@ -0,0 +1,346 @@
  2+package ysweet
  3+
  4+import (
  5+	"context"
  6+	"fmt"
  7+	"net/http"
  8+	"net/http/httptest"
  9+	"net/url"
 10+	"strings"
 11+	"sync/atomic"
 12+	"testing"
 13+	"time"
 14+
 15+	"github.com/coder/websocket"
 16+)
 17+
 18+// mockYSweetServer creates a test y-sweet server that echoes messages
 19+func mockYSweetServer(t *testing.T) *httptest.Server {
 20+	return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
 21+		// Accept websocket connection
 22+		conn, err := websocket.Accept(w, r, nil)
 23+		if err != nil {
 24+			t.Logf("accept error: %v", err)
 25+			return
 26+		}
 27+		defer conn.Close(websocket.StatusNormalClosure, "")
 28+
 29+		// Echo messages back
 30+		ctx := context.Background()
 31+		for {
 32+			msgType, reader, err := conn.Reader(ctx)
 33+			if err != nil {
 34+				return
 35+			}
 36+
 37+			data := make([]byte, 1024)
 38+			n, err := reader.Read(data)
 39+			if err != nil {
 40+				return
 41+			}
 42+
 43+			err = conn.Write(ctx, msgType, data[:n])
 44+			if err != nil {
 45+				return
 46+			}
 47+		}
 48+	}))
 49+}
 50+
 51+func TestProxyHandler_BasicProxy(t *testing.T) {
 52+	// Create mock upstream server
 53+	upstream := mockYSweetServer(t)
 54+	defer upstream.Close()
 55+
 56+	// Convert http:// to ws://
 57+	upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
 58+
 59+	// Create proxy handler
 60+	handler := ProxyHandler(upstreamURL)
 61+
 62+	// Create test server with proxy
 63+	proxyServer := httptest.NewServer(handler)
 64+	defer proxyServer.Close()
 65+
 66+	// Connect to proxy
 67+	proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1) + "/d/test-doc/ws/test-doc"
 68+	clientConn, _, err := websocket.Dial(context.Background(), proxyURL, nil)
 69+	if err != nil {
 70+		t.Fatalf("failed to connect to proxy: %v", err)
 71+	}
 72+	defer clientConn.Close(websocket.StatusNormalClosure, "")
 73+
 74+	// Send a test message
 75+	testMsg := []byte{0x00, 0x01, 0x02, 0x03}
 76+	err = clientConn.Write(context.Background(), websocket.MessageBinary, testMsg)
 77+	if err != nil {
 78+		t.Fatalf("failed to write message: %v", err)
 79+	}
 80+
 81+	// Read echoed message
 82+	ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
 83+	defer cancel()
 84+
 85+	msgType, reader, err := clientConn.Reader(ctx)
 86+	if err != nil {
 87+		t.Fatalf("failed to read message: %v", err)
 88+	}
 89+
 90+	if msgType != websocket.MessageBinary {
 91+		t.Errorf("expected binary message, got %v", msgType)
 92+	}
 93+
 94+	data := make([]byte, 1024)
 95+	n, err := reader.Read(data)
 96+	if err != nil {
 97+		t.Fatalf("failed to read message body: %v", err)
 98+	}
 99+
100+	if string(data[:n]) != string(testMsg) {
101+		t.Errorf("expected %v, got %v", testMsg, data[:n])
102+	}
103+}
104+
105+func TestProxyHandler_OnConnectHook(t *testing.T) {
106+	upstream := mockYSweetServer(t)
107+	defer upstream.Close()
108+	upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
109+
110+	var hookCalled atomic.Bool
111+	var hookReceivedDocID string
112+	var hookReceivedPath string
113+
114+	handler := ProxyHandler(upstreamURL,
115+		WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
116+			hookCalled.Store(true)
117+			hookReceivedDocID = docID
118+			hookReceivedPath = r.URL.Path
119+			return nil
120+		}),
121+	)
122+
123+	proxyServer := httptest.NewServer(handler)
124+	defer proxyServer.Close()
125+
126+	proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1) + "/d/test-doc/ws/test-doc"
127+	clientConn, _, err := websocket.Dial(context.Background(), proxyURL, nil)
128+	if err != nil {
129+		t.Fatalf("failed to connect: %v", err)
130+	}
131+	clientConn.Close(websocket.StatusNormalClosure, "")
132+
133+	if !hookCalled.Load() {
134+		t.Error("OnConnect hook was not called")
135+	}
136+
137+	if hookReceivedDocID != "test-doc" {
138+		t.Errorf("expected docID 'test-doc', got %s", hookReceivedDocID)
139+	}
140+
141+	if hookReceivedPath != "/d/test-doc/ws/test-doc" {
142+		t.Errorf("expected path /d/test-doc/ws/test-doc, got %s", hookReceivedPath)
143+	}
144+}
145+
146+func TestProxyHandler_OnConnectHook_RejectsConnection(t *testing.T) {
147+	upstream := mockYSweetServer(t)
148+	defer upstream.Close()
149+	upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
150+
151+	handler := ProxyHandler(upstreamURL,
152+		WithOnConnect(func(ctx context.Context, docID string, w http.ResponseWriter, r *http.Request) error {
153+			return fmt.Errorf("connection rejected")
154+		}),
155+	)
156+
157+	proxyServer := httptest.NewServer(handler)
158+	defer proxyServer.Close()
159+
160+	// Try HTTP request (not websocket) to see rejection
161+	resp, err := http.Get(proxyServer.URL + "/d/test-doc/ws/test-doc")
162+	if err != nil {
163+		t.Fatalf("http request failed: %v", err)
164+	}
165+	defer resp.Body.Close()
166+
167+	if resp.StatusCode != http.StatusForbidden {
168+		t.Errorf("expected status 403, got %d", resp.StatusCode)
169+	}
170+}
171+
172+func TestProxyHandler_OnWebsocketUpgradeHook(t *testing.T) {
173+	upstream := mockYSweetServer(t)
174+	defer upstream.Close()
175+	upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
176+
177+	var hookCalled atomic.Bool
178+	var hookReceivedDocID string
179+
180+	handler := ProxyHandler(upstreamURL,
181+		WithOnWebsocketUpgrade(func(ctx context.Context, docID string, clientConn, upstreamConn *websocket.Conn) error {
182+			hookCalled.Store(true)
183+			hookReceivedDocID = docID
184+			return nil
185+		}),
186+	)
187+
188+	proxyServer := httptest.NewServer(handler)
189+	defer proxyServer.Close()
190+
191+	proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1) + "/d/test-doc/ws/test-doc"
192+	clientConn, _, err := websocket.Dial(context.Background(), proxyURL, nil)
193+	if err != nil {
194+		t.Fatalf("failed to connect: %v", err)
195+	}
196+	clientConn.Close(websocket.StatusNormalClosure, "")
197+
198+	if !hookCalled.Load() {
199+		t.Error("OnWebsocketUpgrade hook was not called")
200+	}
201+
202+	if hookReceivedDocID != "test-doc" {
203+		t.Errorf("expected docID 'test-doc', got %s", hookReceivedDocID)
204+	}
205+}
206+
207+func TestProxyHandler_OnDisconnectHook(t *testing.T) {
208+	upstream := mockYSweetServer(t)
209+	defer upstream.Close()
210+	upstreamURL := strings.Replace(upstream.URL, "http://", "ws://", 1)
211+
212+	var hookCalled atomic.Bool
213+	var hookReceivedDocID string
214+	var receivedReason DisconnectReason
215+
216+	done := make(chan struct{})
217+	handler := ProxyHandler(upstreamURL,
218+		WithOnDisconnect(func(ctx context.Context, docID string, reason DisconnectReason) error {
219+			hookCalled.Store(true)
220+			hookReceivedDocID = docID
221+			receivedReason = reason
222+			close(done)
223+			return nil
224+		}),
225+	)
226+
227+	proxyServer := httptest.NewServer(handler)
228+	defer proxyServer.Close()
229+
230+	proxyURL := strings.Replace(proxyServer.URL, "http://", "ws://", 1) + "/d/test-doc/ws/test-doc"
231+	clientConn, _, err := websocket.Dial(context.Background(), proxyURL, nil)
232+	if err != nil {
233+		t.Fatalf("failed to connect: %v", err)
234+	}
235+
236+	// Close connection from client side
237+	clientConn.Close(websocket.StatusNormalClosure, "")
238+
239+	// Wait for disconnect hook to be called
240+	select {
241+	case <-done:
242+		// Hook was called
243+	case <-time.After(2 * time.Second):
244+		t.Fatal("timeout waiting for OnDisconnect hook")
245+	}
246+
247+	if !hookCalled.Load() {
248+		t.Error("OnDisconnect hook was not called")
249+	}
250+
251+	if hookReceivedDocID != "test-doc" {
252+		t.Errorf("expected docID 'test-doc', got %s", hookReceivedDocID)
253+	}
254+
255+	// Reason should be either ClientClosed or Error depending on close timing
256+	if receivedReason != DisconnectReasonClientClosed && receivedReason != DisconnectReasonError {
257+		t.Errorf("unexpected disconnect reason: %v", receivedReason)
258+	}
259+}
260+
261+func TestExtractDocID(t *testing.T) {
262+	tests := []struct {
263+		path     string
264+		expected string
265+	}{
266+		{"/d/my-doc/ws/my-doc", "my-doc"},
267+		{"/d/my-doc/ws/my-doc/extra", "my-doc"},
268+		{"/d/my-doc/ws/different-doc", ""}, // mismatched doc IDs
269+		{"/doc/ws/my-doc", ""},             // legacy pattern not supported
270+		{"/ws/my-doc", ""},                 // single-doc pattern not supported
271+		{"/other/path", ""},
272+		{"/", ""},
273+	}
274+
275+	for _, tt := range tests {
276+		t.Run(tt.path, func(t *testing.T) {
277+			req := httptest.NewRequest("GET", tt.path, nil)
278+			config := &proxyConfig{} // No custom docIDFunc, uses default
279+			result := extractDocID(req, config)
280+			if result != tt.expected {
281+				t.Errorf("extractDocID(%q) = %q, want %q", tt.path, result, tt.expected)
282+			}
283+		})
284+	}
285+}
286+
287+func TestExtractDocID_WithCustomFunc(t *testing.T) {
288+	// Test with custom DocIDFunc
289+	config := &proxyConfig{
290+		docIDFunc: func(r *http.Request) string {
291+			return r.PathValue("docID")
292+		},
293+	}
294+
295+	req := httptest.NewRequest("GET", "/ws/my-custom-doc", nil)
296+	// Simulate Go 1.22+ path value
297+	req.SetPathValue("docID", "my-custom-doc")
298+
299+	result := extractDocID(req, config)
300+	if result != "my-custom-doc" {
301+		t.Errorf("extractDocID with custom func = %q, want %q", result, "my-custom-doc")
302+	}
303+}
304+
305+func TestBuildUpstreamURL(t *testing.T) {
306+	tests := []struct {
307+		baseURL  string
308+		docID    string
309+		query    url.Values
310+		expected string
311+	}{
312+		{
313+			baseURL:  "ws://localhost:8080",
314+			docID:    "test-doc",
315+			query:    url.Values{},
316+			expected: "ws://localhost:8080/d/test-doc/ws/test-doc",
317+		},
318+		{
319+			baseURL:  "ws://localhost:8080/",
320+			docID:    "test-doc",
321+			query:    url.Values{},
322+			expected: "ws://localhost:8080/d/test-doc/ws/test-doc",
323+		},
324+		{
325+			baseURL:  "ws://localhost:8080",
326+			docID:    "test-doc",
327+			query:    url.Values{"extra": []string{"value"}},
328+			expected: "ws://localhost:8080/d/test-doc/ws/test-doc?extra=value",
329+		},
330+		{
331+			baseURL:  "ws://localhost:8080",
332+			docID:    "test-doc",
333+			query:    url.Values{"token": []string{"abc123"}},
334+			expected: "ws://localhost:8080/d/test-doc/ws/test-doc?token=abc123",
335+		},
336+	}
337+
338+	for _, tt := range tests {
339+		t.Run(tt.baseURL, func(t *testing.T) {
340+			result := buildUpstreamURL(tt.baseURL, tt.docID, tt.query)
341+			if result != tt.expected {
342+				t.Errorf("buildUpstreamURL(%q, %q, %v) = %q, want %q",
343+					tt.baseURL, tt.docID, tt.query, result, tt.expected)
344+			}
345+		})
346+	}
347+}