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}