doc_manager.go

  1package rtw
  2
  3import (
  4	"context"
  5	"log/slog"
  6	"strings"
  7	"sync"
  8	"sync/atomic"
  9	"time"
 10
 11	"git.kilimanjaro.io/rtw/document"
 12	"git.kilimanjaro.io/rtw/routing"
 13	"git.kilimanjaro.io/rtw/waypoint"
 14	"git.kilimanjaro.io/ygo"
 15	"git.kilimanjaro.io/ygo/ysweet"
 16)
 17
 18// toHTTPURL converts a WebSocket URL to an HTTP URL for REST API calls.
 19// ws://  -> http://
 20// wss:// -> https://
 21func toHTTPURL(wsURL string) string {
 22	if strings.HasPrefix(wsURL, "wss://") {
 23		return "https://" + wsURL[len("wss://"):]
 24	}
 25	if strings.HasPrefix(wsURL, "ws://") {
 26		return "http://" + wsURL[len("ws://"):]
 27	}
 28	return wsURL
 29}
 30
 31// DocManager manages local YDoc copies for all active documents.
 32// It integrates with the y-sweet proxy to track client connections
 33// and runs debounced operations on documents.
 34type DocManager struct {
 35	targetURL  string
 36	sessions   map[string]*docSession
 37	mu         sync.RWMutex
 38	docService *document.Service
 39}
 40
 41// docSession represents a single document's lifecycle.
 42// Each session runs in its own goroutine.
 43type docSession struct {
 44	docID            string
 45	targetURL        string
 46	docService       *document.Service
 47	refCount         int
 48	clientCountCh    chan int      // +1 for connect, -1 for disconnect
 49	shutdownCh       chan struct{} // signal to shutdown immediately
 50	doc              *ygo.Doc
 51	syncClient       *ygo.SyncClient
 52	done             chan struct{} // closed when goroutine exits
 53	processingUpdate int32         // atomic flag for reentry guard in WithOnUpdate
 54}
 55
 56// NewDocManager creates a new DocManager.
 57func NewDocManager(targetURL string, docService *document.Service) *DocManager {
 58	return &DocManager{
 59		targetURL:  targetURL,
 60		sessions:   make(map[string]*docSession),
 61		docService: docService,
 62	}
 63}
 64
 65// run is the main goroutine for a document session.
 66// It manages the YDoc lifecycle and runs debounced operations.
 67func (s *docSession) run() {
 68	defer close(s.done)
 69
 70	// Create YDoc
 71	doc, err := ygo.NewDoc()
 72	if err != nil {
 73		slog.Error("failed to create doc", "doc_id", s.docID, "error", err)
 74		return
 75	}
 76	s.doc = doc
 77
 78	// Get auth token for y-sweet (convert WebSocket URL to HTTP for REST API)
 79	httpURL := toHTTPURL(s.targetURL)
 80	client, err := ysweet.NewClient(httpURL)
 81	if err != nil {
 82		slog.Error("failed to create ysweet client", "doc_id", s.docID, "error", err)
 83		return
 84	}
 85
 86	auth, err := client.AuthDoc(s.docID)
 87	if err != nil {
 88		slog.Error("failed to auth doc", "doc_id", s.docID, "error", err)
 89		return
 90	}
 91
 92	// Create and start sync client with 1-minute awareness timeout
 93	ctx, cancel := context.WithCancel(context.Background())
 94	defer cancel()
 95
 96	syncClient, err := ygo.NewSyncClient(doc,
 97		ygo.WithSyncEndpoint(auth.WebsocketURL()),
 98		ygo.WithDisconnectOnNoClientsAfter(1*time.Minute),
 99	)
100	if err != nil {
101		slog.Error("failed to create sync client", "doc_id", s.docID, "error", err)
102		return
103	}
104	s.syncClient = syncClient
105
106	// Set up OnUpdate callback using method (not as option - options don't work for callbacks)
107	s.syncClient.OnUpdate(func(d *ygo.Doc, stats ygo.SyncStats) error {
108		// Reentry guard: skip if already processing
109		if !atomic.CompareAndSwapInt32(&s.processingUpdate, 0, 1) {
110			slog.Debug("skipping recursive update", "doc_id", s.docID)
111			return nil
112		}
113		defer atomic.StoreInt32(&s.processingUpdate, 0)
114
115		// Run operations on remote update
116		return s.runOperations()
117	})
118
119	// Set up OnDisconnect callback
120	s.syncClient.OnDisconnect(func(d *ygo.Doc, stats ygo.SyncStats, reason ygo.DisconnectReason) error {
121		slog.Info("disconnected", "doc_id", s.docID, "reason", reason)
122		if reason == ygo.DisconnectReasonDocumentIdle {
123			// Trigger shutdown when alone for timeout period
124			select {
125			case <-s.shutdownCh:
126				// Already shutting down
127			default:
128				close(s.shutdownCh)
129			}
130		}
131		return nil
132	})
133
134	if err := s.syncClient.Connect(ctx); err != nil {
135		slog.Error("failed to connect sync", "doc_id", s.docID, "error", err)
136		return
137	}
138
139	slog.Info("doc session started", "doc_id", s.docID, "ref_count", s.refCount)
140
141	for {
142		select {
143		case delta := <-s.clientCountCh:
144			s.refCount += delta
145			slog.Info("refCount changed", "doc_id", s.docID, "ref_count", s.refCount)
146
147		case <-s.shutdownCh:
148			slog.Info("shutting down doc session", "doc_id", s.docID)
149			s.cleanup()
150			return
151		}
152	}
153}
154
155// runOperations executes the waypoint operations.
156// Returns error if operations fail.
157func (s *docSession) runOperations() error {
158	if s.doc == nil {
159		return nil
160	}
161
162	// (1) Update route sources
163	if err := waypoint.FindRouteSource(s.doc); err != nil {
164		slog.Error("FindRouteSource error", "doc_id", s.docID, "error", err)
165		return err
166	}
167
168	// (2) Extract waypoints for routing calculation
169	waypoints, err := waypoint.ExtractWaypoints(s.doc)
170	if err != nil {
171		slog.Error("ExtractWaypoints error", "doc_id", s.docID, "error", err)
172		return err
173	}
174
175	// Log waypoint count and names for debugging
176	if len(waypoints) > 0 {
177		names := make([]string, len(waypoints))
178		for i, wp := range waypoints {
179			names[i] = wp.Label
180		}
181		slog.Debug("waypoints extracted", "doc_id", s.docID, "count", len(waypoints), "names", names)
182	}
183
184	// (3) Calculate missing routes
185	storage := routing.NewRouteStorage("./data/route")
186	calculatedRoutes, err := waypoint.CalculateMissingRoutes(waypoints, storage)
187	if err != nil {
188		slog.Error("CalculateMissingRoutes error", "doc_id", s.docID, "error", err)
189		// Don't return error - we can continue even if routing fails
190		// The routes will be calculated on next update
191	}
192
193	// (4) Update waypoints with calculated routes
194	if len(calculatedRoutes) > 0 {
195		// Convert to updates with point snapping
196		updates := waypoint.ToRouteUpdates(calculatedRoutes)
197		sourcePointUpdates := waypoint.ToSourcePointUpdates(calculatedRoutes)
198
199		// Apply route and destination point updates
200		if len(updates) > 0 {
201			if err := waypoint.UpdateWaypointRoutes(s.doc, updates); err != nil {
202				slog.Error("UpdateWaypointRoutes error", "doc_id", s.docID, "error", err)
203				return err
204			}
205			slog.Info("updated routes", "doc_id", s.docID, "count", len(updates))
206		}
207
208		// Apply source point updates (snapping to route network)
209		if len(sourcePointUpdates) > 0 {
210			if err := waypoint.UpdateSourceWaypointPoints(s.doc, sourcePointUpdates); err != nil {
211				slog.Error("UpdateSourceWaypointPoints error", "doc_id", s.docID, "error", err)
212				return err
213			}
214			slog.Info("updated source waypoint locations", "doc_id", s.docID, "count", len(sourcePointUpdates))
215		}
216	}
217
218	// (5) Persist document content
219	start := time.Now()
220	jsonBytes, err := s.doc.MarshalJSON(ygo.WithXmlFragment("main"))
221	if err != nil {
222		slog.Error("failed to marshal document content", "doc_id", s.docID, "error", err)
223		return err
224	}
225
226	if s.docService != nil {
227		if _, err := s.docService.SetContent(s.docID, jsonBytes); err != nil {
228			slog.Error("failed to save document content", "doc_id", s.docID, "error", err)
229			// Don't return error — sync should not fail due to DB issue
230		} else {
231			slog.Debug("persisted document content", "doc_id", s.docID, "duration_ms", time.Since(start).Milliseconds())
232		}
233	}
234
235	return nil
236}
237
238// cleanup flushes changes and destroys resources.
239func (s *docSession) cleanup() {
240	if s.syncClient != nil {
241		// Flush any pending updates
242		ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
243		if err := s.syncClient.Flush(ctx); err != nil {
244			slog.Error("flush error", "doc_id", s.docID, "error", err)
245		}
246		cancel()
247
248		// Close sync client (this sends pending updates internally)
249		if err := s.syncClient.Close(); err != nil {
250			slog.Error("error closing sync", "doc_id", s.docID, "error", err)
251		}
252	}
253
254	if s.doc != nil {
255		s.doc.Destroy()
256	}
257
258	slog.Info("cleanup complete", "doc_id", s.docID)
259}
260
261// ClientConnect is called when a client connects to the document.
262// It creates a new session if one doesn't exist, or increments the ref count.
263func (dm *DocManager) ClientConnect(docID string) {
264	dm.mu.Lock()
265	defer dm.mu.Unlock()
266
267	session, exists := dm.sessions[docID]
268	if !exists {
269		// Create new session
270		session = &docSession{
271			docID:         docID,
272			targetURL:     dm.targetURL,
273			docService:    dm.docService,
274			refCount:      1, // Start at 1 for the connecting client
275			clientCountCh: make(chan int, 10),
276			shutdownCh:    make(chan struct{}),
277			done:          make(chan struct{}),
278		}
279		dm.sessions[docID] = session
280
281		// Start the goroutine
282		go session.run()
283		slog.Info("created new session", "doc_id", docID)
284	} else {
285		// Increment ref count
286		select {
287		case session.clientCountCh <- 1:
288		default:
289			slog.Warn("clientCountCh full", "doc_id", docID)
290		}
291	}
292}
293
294// ClientDisconnect is called when a client disconnects from the document.
295// It decrements the ref count and may trigger the grace period.
296func (dm *DocManager) ClientDisconnect(docID string) {
297	dm.mu.RLock()
298	session, exists := dm.sessions[docID]
299	dm.mu.RUnlock()
300
301	if !exists {
302		slog.Warn("disconnect for unknown doc", "doc_id", docID)
303		return
304	}
305
306	select {
307	case session.clientCountCh <- -1:
308	default:
309		slog.Warn("clientCountCh full", "doc_id", docID)
310	}
311}
312
313// Shutdown gracefully stops all document sessions.
314func (dm *DocManager) Shutdown() {
315	dm.mu.Lock()
316	sessions := make([]*docSession, 0, len(dm.sessions))
317	for _, s := range dm.sessions {
318		sessions = append(sessions, s)
319	}
320	dm.mu.Unlock()
321
322	// Signal all sessions to shutdown
323	for _, s := range sessions {
324		select {
325		case <-s.shutdownCh:
326			// Already shutting down
327		default:
328			close(s.shutdownCh)
329		}
330	}
331
332	// Wait for all sessions to complete
333	for _, s := range sessions {
334		<-s.done
335	}
336
337	slog.Info("all sessions shutdown")
338}