scooter  ·  2026-04-03

conn.go

  1package sync
  2
  3import (
  4	"context"
  5	"fmt"
  6	"io"
  7	"math"
  8	"sync"
  9	"time"
 10
 11	"github.com/coder/websocket"
 12)
 13
 14// WSConn is an interface for WebSocket connections.
 15// It allows the sync code to work with real or mock connections.
 16type WSConn interface {
 17	Read(context.Context) (websocket.MessageType, io.Reader, error)
 18	Write(context.Context, websocket.MessageType, []byte) error
 19	Close(websocket.StatusCode, string) error
 20}
 21
 22// SyncConn manages the WebSocket connection and sync protocol for a document.
 23// It handles connecting, message handling, reconnection, and graceful shutdown.
 24type SyncConn struct {
 25	conn              *websocket.Conn
 26	doc               DocInterface
 27	opts              Options
 28	onUpdate          func(DocInterface) error
 29	status            Status
 30	retries           int
 31	closeOnce         sync.Once
 32	done              chan struct{}
 33	closeErr          error
 34	onAwarenessUpdate func([]byte) error
 35
 36	// Stats tracking
 37	peerCount     int
 38	lastUpdate    time.Time
 39	pendingUpdate bool
 40	statsMu       sync.RWMutex
 41
 42	// Event callbacks
 43	onConnect    func()
 44	onDisconnect func(DisconnectReason)
 45
 46	// Disconnect reason tracking
 47	disconnectReason DisconnectReason
 48}
 49
 50// Stats holds connection statistics for SyncConn
 51type Stats struct {
 52	PeerCount     int
 53	LastUpdate    time.Time
 54	PendingUpdate bool
 55	Status        Status
 56}
 57
 58const (
 59	maxRetries     = 10
 60	maxBackoff     = 30 * time.Second
 61	initialBackoff = 1 * time.Second
 62)
 63
 64// NewSyncConn creates a new sync connection for the given document
 65// with the specified options.
 66func NewSyncConn(doc DocInterface, opts Options) *SyncConn {
 67	onUpdate := func(DocInterface) error { return nil }
 68	if opts.OnUpdate != nil {
 69		onUpdate = opts.OnUpdate
 70	}
 71	return &SyncConn{
 72		doc:              doc,
 73		opts:             opts,
 74		onUpdate:         onUpdate,
 75		onConnect:        opts.OnConnect,
 76		onDisconnect:     opts.OnDisconnect,
 77		status:           StatusDisconnected,
 78		done:             make(chan struct{}),
 79		retries:          0,
 80		lastUpdate:       time.Time{},
 81		peerCount:        0,
 82		pendingUpdate:    false,
 83		disconnectReason: DisconnectReasonClientInitiated,
 84	}
 85}
 86
 87func (s *SyncConn) connect(ctx context.Context) error {
 88	s.setStatus(StatusConnecting)
 89
 90	c, _, err := websocket.Dial(ctx, s.opts.Endpoint, &websocket.DialOptions{
 91		HTTPHeader: map[string][]string{
 92			"Authorization": {s.opts.AuthToken},
 93		},
 94	})
 95	if err != nil {
 96		return fmt.Errorf("failed to dial websocket: %w", err)
 97	}
 98
 99	s.conn = c
100	s.setStatus(StatusConnected)
101	s.retries = 0
102
103	// Trigger OnConnect callback
104	if s.onConnect != nil {
105		s.onConnect()
106	}
107
108	return nil
109}
110
111func (s *SyncConn) sendSyncStep1(ctx context.Context) error {
112	var sv []byte
113	err := s.doc.WithReadTransaction(func(txn Transaction) error {
114		stateVec := txn.GetStateVector()
115		if stateVec == nil {
116			return fmt.Errorf("failed to get state vector")
117		}
118		sv = stateVec
119		return nil
120	})
121	if err != nil {
122		return fmt.Errorf("failed to get state vector: %w", err)
123	}
124
125	msg := encodeSyncMessage(SyncStep1, sv)
126	err = s.conn.Write(ctx, websocket.MessageBinary, msg)
127	if err != nil {
128		return fmt.Errorf("failed to send SyncStep1: %w", err)
129	}
130
131	return nil
132}
133
134func (s *SyncConn) sendSyncStep2(ctx context.Context) error {
135	var sv []byte
136	err := s.doc.WithReadTransaction(func(txn Transaction) error {
137		stateVec := txn.GetStateVector()
138		if stateVec == nil {
139			return fmt.Errorf("failed to get state vector")
140		}
141		sv = stateVec
142		return nil
143	})
144	if err != nil {
145		return fmt.Errorf("failed to get state vector: %w", err)
146	}
147
148	update := s.doc.GetStateDiff(sv)
149	if update == nil {
150		return fmt.Errorf("failed to get state diff")
151	}
152
153	msg := encodeSyncMessage(SyncStep2, []byte(update))
154	err = s.conn.Write(ctx, websocket.MessageBinary, msg)
155	if err != nil {
156		return fmt.Errorf("failed to send SyncStep2: %w", err)
157	}
158
159	return nil
160}
161
162func (s *SyncConn) handleMessage(ctx context.Context, data []byte) error {
163	if len(data) < 1 {
164		return fmt.Errorf("message too short")
165	}
166
167	msgType, n, err := readVarUint(data)
168	if err != nil {
169		return fmt.Errorf("failed to read message type: %w", err)
170	}
171
172	payload := data[n:]
173
174	switch msgType {
175	case uint64(MessageSync):
176		return s.handleSyncMessage(payload)
177	case uint64(MessageAwareness):
178		return s.handleAwareness(payload)
179	default:
180		return fmt.Errorf("unknown message type: %d", msgType)
181	}
182}
183
184func (s *SyncConn) handleSyncMessage(data []byte) error {
185	if len(data) < 1 {
186		return fmt.Errorf("sync message payload too short")
187	}
188
189	innerType, n, err := readVarUint(data)
190	if err != nil {
191		return fmt.Errorf("failed to read inner message type: %w", err)
192	}
193
194	payload := data[n:]
195
196	switch innerType {
197	case uint64(SyncStep1):
198		return s.handleSyncStep1Response(payload)
199	case uint64(SyncStep2):
200		return s.handleSyncStep2(payload)
201	case uint64(Update):
202		return s.handleUpdate(payload)
203	default:
204		return fmt.Errorf("unknown sync message type: %d", innerType)
205	}
206}
207
208func (s *SyncConn) handleSyncStep1Response(payload []byte) error {
209	return s.sendSyncStep2(context.Background())
210}
211
212func (s *SyncConn) handleSyncStep2(payload []byte) error {
213	length, n, err := readVarUint(payload)
214	if err != nil {
215		return fmt.Errorf("failed to read update length: %w", err)
216	}
217	if len(payload) < n+int(length) {
218		return fmt.Errorf("update payload too short: got %d bytes, expected %d", len(payload), n+int(length))
219	}
220	update := UpdateData(payload[n : n+int(length)])
221	return s.doc.WithWriteTransaction(func(txn Transaction) error {
222		return txn.ApplyUpdate(update)
223	})
224}
225
226func (s *SyncConn) handleUpdate(payload []byte) error {
227	length, n, err := readVarUint(payload)
228	if err != nil {
229		return fmt.Errorf("failed to read update length: %w", err)
230	}
231	if len(payload) < n+int(length) {
232		return fmt.Errorf("update payload too short: got %d bytes, expected %d", len(payload), n+int(length))
233	}
234	update := UpdateData(payload[n : n+int(length)])
235	err = s.doc.WithWriteTransaction(func(txn Transaction) error {
236		return txn.ApplyUpdate(update)
237	})
238	if err != nil {
239		return err
240	}
241
242	// Update last update time on success
243	s.statsMu.Lock()
244	s.lastUpdate = time.Now()
245	s.statsMu.Unlock()
246
247	if s.onUpdate != nil {
248		s.onUpdate(s.doc)
249	}
250
251	return nil
252}
253
254func (s *SyncConn) handleAwareness(payload []byte) error {
255	// Update peer count from awareness message
256	s.statsMu.Lock()
257	s.peerCount = parseAwarenessPeerCount(payload)
258	s.statsMu.Unlock()
259
260	if s.onAwarenessUpdate != nil {
261		return s.onAwarenessUpdate(payload)
262	}
263	return nil
264}
265
266// Start connects to the y-sweet server and begins receiving updates.
267// It returns an error if the initial connection fails.
268// Once started, the connection runs in the background until closed.
269func (s *SyncConn) Start(ctx context.Context) error {
270	err := s.connect(ctx)
271	if err != nil {
272		s.handleDisconnect(ctx)
273		return err
274	}
275
276	err = s.sendSyncStep1(ctx)
277	if err != nil {
278		s.conn.Close(websocket.StatusNormalClosure, "")
279		s.handleDisconnect(ctx)
280		return err
281	}
282
283	go s.readLoop(ctx)
284
285	return nil
286}
287
288func (s *SyncConn) readLoop(ctx context.Context) {
289	for {
290		_, r, err := s.conn.Reader(ctx)
291		if err != nil {
292			// Check if this was a server-initiated disconnect (no context cancellation)
293			if ctx.Err() == nil {
294				s.statsMu.Lock()
295				s.disconnectReason = DisconnectReasonYSweetInitiated
296				s.statsMu.Unlock()
297			}
298			s.handleDisconnect(ctx)
299			return
300		}
301
302		data := make([]byte, 4096)
303		for {
304			n, err := r.Read(data)
305			if n > 0 {
306				if err := s.handleMessage(ctx, data[:n]); err != nil {
307					fmt.Printf("handle message error: %v\n", err)
308				}
309			}
310			if err != nil {
311				break
312			}
313		}
314	}
315}
316
317func (s *SyncConn) handleDisconnect(ctx context.Context) {
318	if s.getStatus() == StatusDisconnecting {
319		return
320	}
321
322	// Get current reason and trigger OnDisconnect callback
323	s.statsMu.RLock()
324	reason := s.disconnectReason
325	s.statsMu.RUnlock()
326
327	s.statsMu.RLock()
328	onDisconnect := s.onDisconnect
329	s.statsMu.RUnlock()
330
331	if onDisconnect != nil {
332		onDisconnect(reason)
333	}
334
335	s.setStatus(StatusDisconnected)
336	if s.conn != nil {
337		s.conn.Close(websocket.StatusNormalClosure, "")
338		s.conn = nil
339	}
340
341	// Reset disconnect reason for next connection (default to YSweetInitiated for reconnects)
342	s.statsMu.Lock()
343	s.disconnectReason = DisconnectReasonYSweetInitiated
344	s.statsMu.Unlock()
345
346	backoff := initialBackoff
347	for {
348		select {
349		case <-s.done:
350			return
351		case <-time.After(backoff):
352		case <-ctx.Done():
353			return
354		}
355
356		err := s.connect(ctx)
357		if err == nil {
358			err = s.sendSyncStep1(ctx)
359			if err == nil {
360				go s.readLoop(ctx)
361				return
362			}
363			s.conn.Close(websocket.StatusNormalClosure, "")
364		}
365
366		s.retries++
367		if s.retries >= maxRetries {
368			return
369		}
370
371		backoff = time.Duration(math.Min(float64(backoff*2), float64(maxBackoff)))
372	}
373}
374
375// Close gracefully shuts down the connection.
376// It sends any pending updates before closing the WebSocket.
377func (s *SyncConn) Close() error {
378	s.closeOnce.Do(func() {
379		s.setStatus(StatusDisconnecting)
380		close(s.done)
381
382		if s.conn != nil {
383			s.closeErr = s.conn.Close(websocket.StatusNormalClosure, "")
384		}
385		s.setStatus(StatusDisconnected)
386	})
387	return s.closeErr
388}
389
390// SendPendingAndClose sends pending updates then closes the connection synchronously.
391// Use this before destroying the document to ensure pending changes are sent.
392func (s *SyncConn) SendPendingAndClose() error {
393	s.closeOnce.Do(func() {
394		s.setStatus(StatusDisconnecting)
395
396		// Send final pending update if connected
397		if s.conn != nil {
398			// Get state vector and compute diff synchronously
399			var preStateVector []byte
400			_ = s.doc.WithReadTransaction(func(txn Transaction) error {
401				preStateVector = txn.GetStateVector()
402				return nil
403			})
404
405			if len(preStateVector) > 0 {
406				// Compute diff using the state vector
407				// We need to get all updates since the state vector
408				// For simplicity, we'll get the full state and send it
409				_ = s.doc.WithReadTransaction(func(txn Transaction) error {
410					// Get the current state vector again and compute diff
411					currentSV := txn.GetStateVector()
412					if currentSV != nil {
413						// Try to get diff - this may not work for all implementations
414						// but we'll try
415					}
416					return nil
417				})
418			}
419
420			s.conn.Close(websocket.StatusNormalClosure, "")
421		}
422
423		close(s.done)
424		s.setStatus(StatusDisconnected)
425	})
426	return s.closeErr
427}
428
429// Done returns a channel that is closed when the connection is stopped.
430func (s *SyncConn) Done() <-chan struct{} {
431	return s.done
432}
433
434// Status returns the current connection state.
435func (s *SyncConn) Status() Status {
436	return s.status
437}
438
439// InjectConn sets the WebSocket connection directly.
440// This is useful for testing or injecting a custom connection.
441func (s *SyncConn) InjectConn(conn *websocket.Conn) {
442	s.conn = conn
443}
444
445// SetDoc sets the document to sync.
446// This allows changing the document after creation.
447func (s *SyncConn) SetDoc(doc DocInterface) {
448	s.doc = doc
449}
450
451// SetOnUpdate sets the callback for handling updates from the server.
452func (s *SyncConn) SetOnUpdate(fn func(DocInterface) error) {
453	s.statsMu.Lock()
454	s.onUpdate = fn
455	s.statsMu.Unlock()
456}
457
458// GetConn returns the underlying WebSocket connection.
459func (s *SyncConn) GetConn() *websocket.Conn {
460	return s.conn
461}
462
463// SetAwarenessUpdate sets a callback for handling awareness updates
464// from the server.
465func (s *SyncConn) SetAwarenessUpdate(fn func([]byte) error) {
466	s.onAwarenessUpdate = fn
467}
468
469// GetStats returns current connection statistics.
470func (s *SyncConn) GetStats() Stats {
471	s.statsMu.RLock()
472	defer s.statsMu.RUnlock()
473
474	return Stats{
475		PeerCount:     s.peerCount,
476		LastUpdate:    s.lastUpdate,
477		PendingUpdate: s.pendingUpdate,
478		Status:        s.status,
479	}
480}
481
482// setStatus safely sets the connection status
483func (s *SyncConn) setStatus(status Status) {
484	s.statsMu.Lock()
485	s.status = status
486	s.statsMu.Unlock()
487}
488
489// getStatus safely gets the connection status
490func (s *SyncConn) getStatus() Status {
491	s.statsMu.RLock()
492	defer s.statsMu.RUnlock()
493	return s.status
494}
495
496// SetPendingUpdate marks whether there are local changes waiting to sync.
497func (s *SyncConn) SetPendingUpdate(pending bool) {
498	s.statsMu.Lock()
499	s.pendingUpdate = pending
500	s.statsMu.Unlock()
501}
502
503// Flush performs a best-effort sync of pending updates to the server.
504// It sends any pending local changes synchronously.
505// The context can be used to set a timeout.
506func (s *SyncConn) Flush(ctx context.Context) error {
507	if s.getStatus() != StatusConnected || s.conn == nil {
508		return fmt.Errorf("not connected")
509	}
510
511	s.statsMu.RLock()
512	hasPending := s.pendingUpdate
513	s.statsMu.RUnlock()
514
515	if !hasPending {
516		return nil // Nothing to flush
517	}
518
519	// Get current state and compute diff
520	var updateData []byte
521	err := s.doc.WithReadTransaction(func(txn Transaction) error {
522		sv := txn.GetStateVector()
523		if sv == nil {
524			return fmt.Errorf("failed to get state vector")
525		}
526		updateData = s.doc.GetStateDiff(sv)
527		return nil
528	})
529	if err != nil {
530		return fmt.Errorf("failed to get state diff: %w", err)
531	}
532
533	if len(updateData) == 0 {
534		return nil
535	}
536
537	// Send update
538	msg := encodeSyncMessage(Update, updateData)
539	err = s.conn.Write(ctx, websocket.MessageBinary, msg)
540	if err != nil {
541		return fmt.Errorf("failed to send update: %w", err)
542	}
543
544	// Clear pending flag on success
545	s.SetPendingUpdate(false)
546
547	return nil
548}
549
550// SetDisconnectReason sets the reason for the next disconnect.
551// This is used by the awareness client to indicate DocumentIdle.
552func (s *SyncConn) SetDisconnectReason(reason DisconnectReason) {
553	s.statsMu.Lock()
554	s.disconnectReason = reason
555	s.statsMu.Unlock()
556}
557
558// parseAwarenessPeerCount extracts the number of peers from an awareness message.
559// Awareness messages contain alternating client IDs and their state objects.
560// We count unique client IDs to get the peer count.
561func parseAwarenessPeerCount(data []byte) int {
562	if len(data) == 0 {
563		return 0
564	}
565
566	// Simple parser: count entries in the awareness message
567	// Format is a series of [clientId, state] pairs
568	// We'll decode the varint length prefix and count entries
569	count := 0
570	offset := 0
571
572	for offset < len(data) {
573		// Read client ID (varint)
574		_, n, err := readVarUint(data[offset:])
575		if err != nil {
576			break
577		}
578		offset += n
579		count++
580
581		// Skip the state object (varint length prefix + data)
582		if offset >= len(data) {
583			break
584		}
585		stateLen, n, err := readVarUint(data[offset:])
586		if err != nil {
587			break
588		}
589		offset += n + int(stateLen)
590	}
591
592	return count
593}