4 files changed,
+1161,
-0
+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+}
+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+}
+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+}
+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+}