From c0deee08c2d1bb0ca5678f2669758fc8de2d2b81 Mon Sep 17 00:00:00 2001 From: Anish Lakhwara Date: Sun, 1 Jun 2025 18:31:51 -0700 Subject: [PATCH] fix: sync all changes while offline this implementation is so nasty, it's entirely vibe coded. But it does work! the bug we were facing was due to lack of understanding that `colVersion` also needs to be tracked alongside `dbVersion`. we also made some updates to the way the handshake happens for the protocol, introducing sync_request and having the server request changes from the client when it is ahead. The solution is entierly unoptimal right now, possibly with many duplication bugs and request (state tracking for the lastest version on the client side is particularly garbage), but it does work. I'm also certain that there's code smells surronding authentication, up to and possibly including overwriting existing keys in the authentication table. The v2 spec included in this commit has been implemented up to Phase 1: Core Protocol. We will still need to implement authorization. --- mast-react-vite/src/hooks/use-sync.ts | 20 +- mast-react-vite/src/main.tsx | 9 +- mast-react-vite/src/worker/sync-worker.ts | 541 ++++++++++++++++++++-- mast-sync-protocol-v3-spec.md | 416 +++++++++++++++++ server/auth.go | 24 + server/main.go | 429 ++++++++++++++++- 6 files changed, 1386 insertions(+), 53 deletions(-) create mode 100644 mast-sync-protocol-v3-spec.md diff --git a/mast-react-vite/src/hooks/use-sync.ts b/mast-react-vite/src/hooks/use-sync.ts index b924306..e1f0d56 100644 --- a/mast-react-vite/src/hooks/use-sync.ts +++ b/mast-react-vite/src/hooks/use-sync.ts @@ -5,19 +5,23 @@ interface SyncConfig { room: string; endpoint: string; worker: Worker; + publicKey?: string; } export function useCustomSync({ dbname, room, endpoint, - worker + worker, + publicKey }: SyncConfig) { const [status, setStatus] = useState('disconnected'); const [changesSent, setChangesSent] = useState(0); const [changesReceived, setChangesReceived] = useState(0); const [lastError, setLastError] = useState(null); const [lastSyncTime, setLastSyncTime] = useState(null); + const [roomAccess, setRoomAccess] = useState(null); + const [needsPayment, setNeedsPayment] = useState(false); // Initialize sync once on component mount useEffect(() => { @@ -31,13 +35,14 @@ export function useCustomSync({ config: { room: room, url: endpoint, + publicKey, requestUnsyncedChanges: true // Request server to send unsynced changes immediately } }); // Set up message handlers const handleMessage = (event: MessageEvent) => { - const { type, dbname: eventDbname, count, error } = event.data; + const { type, dbname: eventDbname, count, error, access, needsPayment: needsPaymentFlag } = event.data; // Only process messages for our database if (eventDbname && eventDbname !== dbname) return; @@ -53,6 +58,11 @@ export function useCustomSync({ setStatus('error'); setLastError(error); break; + case 'room_status': + setRoomAccess(access); + setNeedsPayment(needsPaymentFlag || false); + console.log(`Room access: ${access}, needs payment: ${needsPaymentFlag}`); + break; case 'CHANGES_SENT': setChangesSent(prev => prev + (count || 0)); setLastSyncTime(new Date()); @@ -75,7 +85,7 @@ export function useCustomSync({ dbname }); }; - }, [dbname, room, endpoint, worker]); + }, [dbname, room, endpoint, worker, publicKey]); // Function to trigger sync manually const syncChanges = useCallback(() => { @@ -93,7 +103,9 @@ export function useCustomSync({ changesReceived, lastError, lastSyncTime, - syncChanges + syncChanges, + roomAccess, + needsPayment }; } diff --git a/mast-react-vite/src/main.tsx b/mast-react-vite/src/main.tsx index d5c02b2..1faaf78 100644 --- a/mast-react-vite/src/main.tsx +++ b/mast-react-vite/src/main.tsx @@ -254,7 +254,7 @@ function updateUrlWithRoom(roomId: string) { } // Custom sync component using your own implementation -function CustomSyncComponent({ dbname, roomId }: { dbname: string, roomId: string }) { +function CustomSyncComponent({ dbname, roomId, publicKey }: { dbname: string, roomId: string, publicKey?: string }) { const endpoint = process.env.NODE_ENV === 'production' ? `wss://mast-server.fly.dev/sync` : `ws://localhost:8080/sync`; @@ -263,7 +263,8 @@ function CustomSyncComponent({ dbname, roomId }: { dbname: string, roomId: strin dbname, room: roomId, endpoint, - worker: syncWorker + worker: syncWorker, + publicKey }); return ; @@ -405,7 +406,7 @@ const initDb = async () => { } }); - return { db, rx, roomId, dbname }; + return { db, rx, roomId, dbname, publicKeyBase64 }; }; // Function to register a key with the server @@ -580,7 +581,7 @@ const init = async () => { {/* Use the custom sync component with the database name and room ID */} - + diff --git a/mast-react-vite/src/worker/sync-worker.ts b/mast-react-vite/src/worker/sync-worker.ts index 6bf3da4..940816a 100644 --- a/mast-react-vite/src/worker/sync-worker.ts +++ b/mast-react-vite/src/worker/sync-worker.ts @@ -13,6 +13,23 @@ function logError(message: string, ...data: any[]) { console.error(`[SyncWorker] ${message}`, ...data); } +// Version state for tracking missing changes (v2 protocol) +interface VersionColPair { + dbVersion: number; + colVersion: number; +} + +interface MissingVersionRange { + dbVersion: number; + colVersions: number[]; +} + +interface VersionState { + contiguousUpTo: VersionColPair; + missingRanges: MissingVersionRange[]; + maxVersionSeen: VersionColPair; +} + // Store active connections - only one per database const connections: Record = {}; // Handle messages from the main thread @@ -103,7 +125,7 @@ async function sendSignedMessage(dbname: string, message: any) { }; // Start syncing a database -async function startSync(dbname: string, config: { room: string, url: string }) { +async function startSync(dbname: string, config: { room: string, url: string, publicKey?: string }) { try { // Check if we already have a connection for this database if (connections[dbname]) { @@ -136,18 +158,63 @@ async function startSync(dbname: string, config: { room: string, url: string }) const siteId = siteIdResult[0].site_id; logDebug(`Site ID: ${Array.from(siteId)}`); - // Get the highest db_version from crsql_changes table + // Initialize version state for v2 protocol let lastSyncVersion = 0; + let versionState: VersionState = { + contiguousUpTo: { dbVersion: 0, colVersion: 0 }, + missingRanges: [], + maxVersionSeen: { dbVersion: 0, colVersion: 0 } + }; + try { - const versionResult = await db.execO<{max_version: number}[]>( + logDebug(`Querying version state for database ${dbname}`); + + // Get all (db_version, col_version) pairs we have (from all sites, not just our own) + const allVersionsResult = await db.execO<{db_version: number; col_version: number}[]>( + "SELECT db_version, col_version FROM crsql_changes ORDER BY db_version, col_version" + ); + + logDebug(`Found ${allVersionsResult.length} version pairs in database`); + + if (allVersionsResult.length > 0) { + const versionPairs = allVersionsResult.map(r => ({ dbVersion: r.db_version, colVersion: r.col_version })); + logDebug(`All version pairs in database: ${versionPairs.slice(0, 10).map(v => `${v.dbVersion}.${v.colVersion}`).join(',')}${versionPairs.length > 10 ? '...' : ''}`); + + // Calculate actual contiguous versions with col_version tracking + logDebug(`Calculating contiguous versions from ${versionPairs.length} pairs`); + const contiguousUpTo = calculateContiguousVersionPairs(versionPairs); + logDebug(`Calculated contiguous up to: ${contiguousUpTo.dbVersion}.${contiguousUpTo.colVersion}`); + + const maxVersionSeen = getMaxVersionPair(versionPairs); + logDebug(`Calculated max version seen: ${maxVersionSeen.dbVersion}.${maxVersionSeen.colVersion}`); + + const missingRanges = calculateMissingVersionRanges(versionPairs, maxVersionSeen); + logDebug(`Calculated ${missingRanges.length} missing ranges`); + + versionState.contiguousUpTo = contiguousUpTo; + versionState.maxVersionSeen = maxVersionSeen; + versionState.missingRanges = missingRanges; + + // For lastSyncVersion, get max from our own changes only + const ownVersionResult = await db.execO<{max_version: number}[]>( + "SELECT MAX(db_version) as max_version FROM crsql_changes WHERE site_id = crsql_site_id()" + ); + if (ownVersionResult.length > 0 && ownVersionResult[0].max_version !== null) { + lastSyncVersion = ownVersionResult[0].max_version; + } + + logDebug(`Calculated version state: contiguous_up_to=${contiguousUpTo.dbVersion}.${contiguousUpTo.colVersion}, max_version_seen=${maxVersionSeen.dbVersion}.${maxVersionSeen.colVersion}, missing_ranges=${missingRanges.length}, own_max=${lastSyncVersion}`); + } + + // Also get the global max version for reference + const globalVersionResult = await db.execO<{max_version: number}[]>( "SELECT MAX(db_version) as max_version FROM crsql_changes" ); - if (versionResult.length > 0 && versionResult[0].max_version !== null) { - lastSyncVersion = versionResult[0].max_version; - logDebug(`Retrieved lastSyncVersion ${lastSyncVersion} from database`); + if (globalVersionResult.length > 0 && globalVersionResult[0].max_version !== null) { + logDebug(`Global max version in database: ${globalVersionResult[0].max_version}`); } } catch (error) { - logError(`Error retrieving lastSyncVersion:`, error); + logError(`Error retrieving version state:`, error); } // Store or update connection info @@ -159,7 +226,10 @@ async function startSync(dbname: string, config: { room: string, url: string }) siteId, room: config.room, url: config.url, - isConnecting: true + isConnecting: true, + versionState, + publicKey: config.publicKey, + shouldReconnect: true }; } else { connections[dbname].db = db; @@ -168,15 +238,24 @@ async function startSync(dbname: string, config: { room: string, url: string }) connections[dbname].room = config.room; connections[dbname].url = config.url; connections[dbname].isConnecting = true; + connections[dbname].versionState = versionState; + connections[dbname].publicKey = config.publicKey; + connections[dbname].shouldReconnect = true; } - // Create WebSocket connection with room in query parameter + // Create WebSocket connection with room and publicKey in query parameters let wsUrl = config.url; if (!wsUrl.includes('?room=')) { const separator = wsUrl.includes('?') ? '&' : '?'; wsUrl = `${wsUrl}${separator}room=${config.room}`; } + // Add publicKey parameter for v2 protocol + if (config.publicKey) { + const separator = wsUrl.includes('?') ? '&' : '?'; + wsUrl = `${wsUrl}${separator}publicKey=${encodeURIComponent(config.publicKey)}`; + } + logDebug(`Connecting to WebSocket at ${wsUrl}`); const ws = new WebSocket(wsUrl); @@ -196,21 +275,8 @@ async function startSync(dbname: string, config: { room: string, url: string }) }, 100); }); - // Initial sync - request changes from server - const connection = connections[dbname]; - - // Encode siteId to base64 for transmission - const encodedSiteId = siteId instanceof Uint8Array - ? btoa(String.fromCharCode.apply(null, siteId)) - : btoa(String(siteId)); - - logDebug(`Requesting pull with siteId: ${Array.from(siteId)}, encoded: ${encodedSiteId}, version: ${connection.lastSyncVersion}`) - ws.send(JSON.stringify({ - type: "pull", - room: config.room, - site_id: encodedSiteId, - version: connection.lastSyncVersion - })); + // Wait for room_status before sending sync_request + // The sync_request will be sent in the room_status message handler // Notify main thread of connection self.postMessage({ type: 'SYNC_CONNECTED', dbname }); @@ -223,23 +289,98 @@ async function startSync(dbname: string, config: { room: string, url: string }) const message = JSON.parse(event.data); switch (message.type) { + case "room_status": + logDebug(`Received room status: ${message.access}`); + connections[dbname].roomAccess = message.access; + self.postMessage({ + type: 'room_status', + dbname, + access: message.access, + needsPayment: message.access === 'no_room' + }); + + // After receiving room_status, send sync_request if we have access + if (message.access === 'write' || message.access === 'read') { + const connection = connections[dbname]; + + // Encode siteId to base64 for transmission + const encodedSiteId = connection.siteId instanceof Uint8Array + ? btoa(String.fromCharCode.apply(null, connection.siteId)) + : btoa(String(connection.siteId)); + + if (config.publicKey && connection.versionState) { + // Use v2 sync_request protocol + logDebug(`Sending sync_request with v2 protocol: contiguous_up_to=${connection.versionState.contiguousUpTo.dbVersion}.${connection.versionState.contiguousUpTo.colVersion}, max_version_seen=${connection.versionState.maxVersionSeen.dbVersion}.${connection.versionState.maxVersionSeen.colVersion}`) + ws.send(JSON.stringify({ + type: "sync_request", + site_id: encodedSiteId, + contiguous_up_to: { + db_version: connection.versionState.contiguousUpTo.dbVersion, + col_version: connection.versionState.contiguousUpTo.colVersion + }, + missing_ranges: connection.versionState.missingRanges.map(range => ({ + db_version: range.dbVersion, + col_versions: range.colVersions + })), + max_version_seen: { + db_version: connection.versionState.maxVersionSeen.dbVersion, + col_version: connection.versionState.maxVersionSeen.colVersion + }, + publicKey: config.publicKey + })); + } else { + // Fall back to v1 pull protocol + logDebug(`Sending pull with v1 protocol: siteId: ${Array.from(connection.siteId)}, encoded: ${encodedSiteId}, version: ${connection.lastSyncVersion}`) + ws.send(JSON.stringify({ + type: "pull", + room: config.room, + site_id: encodedSiteId, + version: connection.lastSyncVersion + })); + } + } + break; + + case "sync_response": + logDebug(`Received sync response with ${message.changes?.length || 0} changes, server max version: ${message.current_max_version}`); + if (Array.isArray(message.changes)) { + await applyChanges(dbname, message.changes); + updateVersionState(dbname, message.changes); + } + break; + case "changes": if (Array.isArray(message.data)) { logDebug(`Received ${message.data.length} changes from server`); await applyChanges(dbname, message.data); + updateVersionState(dbname, message.data); } break; case "request_changes": - logDebug(`Server requested changes since version ${message.version}`); + logDebug(`Server requested changes since version ${JSON.stringify(message.version)}`); try { - await sendChanges(dbname, message.version); + // Extract db_version from the version pair object + const versionNumber = typeof message.version === 'object' && message.version?.db_version !== undefined + ? message.version.db_version + : (typeof message.version === 'number' ? message.version : 0); + logDebug(`Extracted version number: ${versionNumber}`); + await sendChanges(dbname, versionNumber); } catch (error) { logError(`Error responding to request_changes:`, error); } break; + case "error": + logError(`Received server error: ${message.code} - ${message.message}`); + self.postMessage({ + type: 'SYNC_ERROR', + dbname, + error: `${message.code}: ${message.message}` + }); + break; + default: logDebug(`Received unknown message type: ${message.type}`); } @@ -255,9 +396,24 @@ async function startSync(dbname: string, config: { room: string, url: string }) self.postMessage({ type: 'SYNC_ERROR', dbname, error: 'WebSocket error' }); }; - ws.onclose = () => { - logDebug(`WebSocket connection closed for ${dbname}`); + ws.onclose = (event) => { + logDebug(`WebSocket connection closed for ${dbname}`, { code: event.code, reason: event.reason }); connections[dbname].isConnecting = false; + + // Check if this was an auth failure (server termination) + // Code 1002 = protocol error, 1008 = policy violation, 4001-4999 = custom app codes + const isAuthFailure = event.code === 1002 || event.code === 1008 || (event.code >= 4000 && event.code < 5000); + + if (isAuthFailure) { + logDebug(`Connection closed due to auth failure (code: ${event.code}), not reconnecting`); + connections[dbname].shouldReconnect = false; + self.postMessage({ type: 'SYNC_AUTH_FAILED', dbname, code: event.code }); + } else if (connections[dbname].shouldReconnect !== false) { + // Normal disconnect - attempt reconnection with random delay + logDebug(`Connection lost, scheduling reconnection for ${dbname}`); + attemptReconnection(dbname); + } + // Notify main thread of disconnection self.postMessage({ type: 'SYNC_DISCONNECTED', dbname }); }; @@ -274,6 +430,44 @@ async function startSync(dbname: string, config: { room: string, url: string }) } } +// Attempt reconnection with random delay (3-9 seconds per v2 spec) +function attemptReconnection(dbname: string) { + const connection = connections[dbname]; + if (!connection || connection.shouldReconnect === false) { + return; + } + + // Clear any existing reconnection timeout + if (connection.reconnectTimeoutId) { + clearTimeout(connection.reconnectTimeoutId); + } + + // Random delay between 3 and 9 seconds to prevent stampeding herd + const delay = Math.floor(Math.random() * 6000) + 3000; // 3000-9000ms + logDebug(`Reconnecting to ${dbname} in ${delay}ms`); + + connection.reconnectTimeoutId = setTimeout(async () => { + if (connections[dbname] && connections[dbname].shouldReconnect !== false) { + logDebug(`Attempting reconnection to ${dbname}`); + try { + // Reset the WebSocket and connection state for reconnection + connections[dbname].ws = null; + connections[dbname].isConnecting = false; + + await startSync(dbname, { + room: connection.room, + url: connection.url, + publicKey: connection.publicKey + }); + } catch (error) { + logError(`Reconnection failed for ${dbname}:`, error); + // Try again after another random delay + attemptReconnection(dbname); + } + } + }, delay); +} + // Stop syncing a database function stopSync(dbname: string) { logDebug(`Stopping sync for ${dbname}`); @@ -283,6 +477,15 @@ function stopSync(dbname: string) { return; } + // Disable reconnection + connection.shouldReconnect = false; + + // Clear any pending reconnection timeout + if (connection.reconnectTimeoutId) { + clearTimeout(connection.reconnectTimeoutId); + connection.reconnectTimeoutId = undefined; + } + // Close WebSocket if (connection.ws && (connection.ws.readyState === WebSocket.OPEN || @@ -321,17 +524,34 @@ async function sendChanges(dbname: string, version?: number) { try { // Use provided version or fall back to connection's lastSyncVersion - const syncVersion = version !== undefined ? version : connection.lastSyncVersion; - logDebug(`Querying for changes since version ${syncVersion}`); + let syncVersion = version !== undefined ? version : connection.lastSyncVersion; + logDebug(`syncVersion: ${syncVersion}`) + + // Special case: if version is explicitly 0 or -1, send ALL changes + if (version === 0 || version === -1) { + syncVersion = -1; // This will get all changes since db_version > -1 means all changes + logDebug(`Sending ALL changes for ${dbname} (version was ${version})`); + } else { + logDebug(`Querying for changes since version ${syncVersion}`); + } // Query for changes since specified version const changes = await connection.db.execA( - `SELECT * FROM crsql_changes WHERE db_version > ? AND site_id = crsql_site_id()`, + `SELECT * FROM crsql_changes WHERE db_version >= ? AND site_id = crsql_site_id()`, [syncVersion] ); + // Debug: also check what ALL changes exist + const allChanges = await connection.db.execA( + `SELECT db_version, site_id FROM crsql_changes ORDER BY db_version` + ); + logDebug(`Total changes in DB: ${allChanges.length}`); + logDebug(`All versions: ${allChanges.map(c => c[0]).join(',')}`); + logDebug(`Querying for: db_version > ${syncVersion} AND site_id = crsql_site_id()`); + logDebug(`Found ${changes.length} matching changes to send`); + if (changes.length === 0) { - logDebug(`No changes to send for ${dbname}`); + logDebug(`No changes to send for ${dbname} (queried > ${syncVersion})`); return; } @@ -421,6 +641,260 @@ async function sendChanges(dbname: string, version?: number) { } } +// Update version state based on received changes +async function recalculateVersionState(dbname: string) { + const connection = connections[dbname]; + if (!connection || !connection.versionState) { + return; + } + + try { + // Re-query all version pairs from database + const allVersionsResult = await connection.db.execO<{db_version: number; col_version: number}[]>( + "SELECT db_version, col_version FROM crsql_changes ORDER BY db_version, col_version" + ); + + if (allVersionsResult.length === 0) { + connection.versionState = { + contiguousUpTo: { dbVersion: 0, colVersion: 0 }, + missingRanges: [], + maxVersionSeen: { dbVersion: 0, colVersion: 0 } + }; + return; + } + + const versionPairs = allVersionsResult.map(r => ({ dbVersion: r.db_version, colVersion: r.col_version })); + + // Recalculate all version state + const contiguousUpTo = calculateContiguousVersionPairs(versionPairs); + const maxVersionSeen = getMaxVersionPair(versionPairs); + const missingRanges = calculateMissingVersionRanges(versionPairs, maxVersionSeen); + + connection.versionState = { + contiguousUpTo, + maxVersionSeen, + missingRanges + }; + + logDebug(`Recalculated version state: contiguous_up_to=${contiguousUpTo.dbVersion}.${contiguousUpTo.colVersion}, max_version_seen=${maxVersionSeen.dbVersion}.${maxVersionSeen.colVersion}, missing_ranges=${missingRanges.length}`); + } catch (error) { + logError(`Error recalculating version state for ${dbname}:`, error); + } +} + +function updateVersionState(dbname: string, changes: any[]) { + const connection = connections[dbname]; + if (!connection || !connection.versionState) { + return; + } + + const versionState = connection.versionState; + + for (const change of changes) { + const versionPair: VersionColPair = { + dbVersion: change.DBVersion, + colVersion: change.ColVersion + }; + + // Update max seen + if (compareVersionPairs(versionPair, versionState.maxVersionSeen) > 0) { + versionState.maxVersionSeen = versionPair; + } + + // Remove this specific (db_version, col_version) from missing ranges + removeMissingVersionPair(versionState, versionPair); + } + + logDebug(`Updated version state: contiguous_up_to=${versionState.contiguousUpTo.dbVersion}.${versionState.contiguousUpTo.colVersion}, max_version_seen=${versionState.maxVersionSeen.dbVersion}.${versionState.maxVersionSeen.colVersion}, missing_ranges=${versionState.missingRanges.length}`); +} + +// Helper function to remove a specific (db_version, col_version) pair from missing ranges +function removeMissingVersionPair(versionState: VersionState, versionPair: VersionColPair) { + const newRanges: MissingVersionRange[] = []; + + for (const range of versionState.missingRanges) { + if (range.dbVersion !== versionPair.dbVersion) { + // Different db_version, keep the entire range + newRanges.push(range); + } else { + // Same db_version, remove the specific col_version + const remainingColVersions = range.colVersions.filter(cv => cv !== versionPair.colVersion); + if (remainingColVersions.length > 0) { + newRanges.push({ + dbVersion: range.dbVersion, + colVersions: remainingColVersions + }); + } + } + } + + versionState.missingRanges = newRanges; +} + +// Helper function to recalculate contiguous version state after changes +function recalculateContiguousVersionState(versionState: VersionState) { + // This function is complex and error-prone. Let's use a simpler approach: + // Re-query the database to get the current state and recalculate from scratch + logDebug("Recalculating contiguous version state - this approach is flawed, should re-query database"); + + // For now, we'll use a conservative approach: don't change contiguous state + // during incremental updates. The initial calculation should be correct. + // This is a temporary fix - ideally we should re-query the database. + + // Only update if we were at 0,0 (initial state) + if (versionState.contiguousUpTo.dbVersion === 0 && versionState.contiguousUpTo.colVersion === 0) { + // Try to set a minimal contiguous state + if (versionState.missingRanges.length === 0) { + // No missing ranges, so contiguous up to max + versionState.contiguousUpTo = versionState.maxVersionSeen; + } else { + // Find the lowest missing range and set contiguous just before it + let minMissingDb = Math.min(...versionState.missingRanges.map(r => r.dbVersion)); + if (minMissingDb > 1) { + versionState.contiguousUpTo = { dbVersion: minMissingDb - 1, colVersion: 1 }; + } else { + // Missing something in db_version 1, so contiguous up to 0 + versionState.contiguousUpTo = { dbVersion: 0, colVersion: 0 }; + } + } + } + + logDebug(`Recalculated contiguous up to: ${versionState.contiguousUpTo.dbVersion}.${versionState.contiguousUpTo.colVersion}`); +} + +// Helper function to get the maximum version pair from an array +function getMaxVersionPair(versionPairs: VersionColPair[]): VersionColPair { + if (versionPairs.length === 0) return { dbVersion: 0, colVersion: 0 }; + + let maxDbVersion = 0; + let maxColVersionForMaxDb = 0; + + for (const pair of versionPairs) { + if (pair.dbVersion > maxDbVersion || (pair.dbVersion === maxDbVersion && pair.colVersion > maxColVersionForMaxDb)) { + maxDbVersion = pair.dbVersion; + maxColVersionForMaxDb = pair.colVersion; + } + } + + return { dbVersion: maxDbVersion, colVersion: maxColVersionForMaxDb }; +} + +// Helper function to compare version pairs +function compareVersionPairs(a: VersionColPair, b: VersionColPair): number { + if (a.dbVersion !== b.dbVersion) { + return a.dbVersion - b.dbVersion; + } + return a.colVersion - b.colVersion; +} + +// Calculate the highest contiguous version pair from an array of version pairs +function calculateContiguousVersionPairs(versionPairs: VersionColPair[]): VersionColPair { + if (versionPairs.length === 0) return { dbVersion: 0, colVersion: 0 }; + + // Sort pairs by db_version, then col_version + const sorted = [...versionPairs].sort(compareVersionPairs); + + // Group by db_version + const versionGroups = new Map(); + for (const pair of sorted) { + if (!versionGroups.has(pair.dbVersion)) { + versionGroups.set(pair.dbVersion, []); + } + versionGroups.get(pair.dbVersion)!.push(pair.colVersion); + } + + // Find the highest contiguous db_version where we have all col_versions from 1 to max + let contiguousDbVersion = 0; + let contiguousColVersion = 0; + + for (let dbVersion = 1; dbVersion <= Math.max(...Array.from(versionGroups.keys())); dbVersion++) { + const colVersions = versionGroups.get(dbVersion); + if (!colVersions) { + // Missing this entire db_version + break; + } + + // Check if we have contiguous col_versions from 1 to max + const sortedColVersions = [...colVersions].sort((a, b) => a - b); + let expectedColVersion = 1; + let lastContiguousCol = 0; + + for (const colVersion of sortedColVersions) { + if (colVersion === expectedColVersion) { + lastContiguousCol = colVersion; + expectedColVersion++; + } else if (colVersion > expectedColVersion) { + break; + } + } + + if (lastContiguousCol > 0) { + contiguousDbVersion = dbVersion; + contiguousColVersion = lastContiguousCol; + } else { + break; + } + } + + return { dbVersion: contiguousDbVersion, colVersion: contiguousColVersion }; +} + +// Calculate missing version ranges given version pairs and max version +function calculateMissingVersionRanges(versionPairs: VersionColPair[], maxVersionSeen: VersionColPair): MissingVersionRange[] { + const ranges: MissingVersionRange[] = []; + + if (versionPairs.length === 0) { + if (maxVersionSeen.dbVersion > 0) { + // Missing everything from 1 to max + for (let dbVersion = 1; dbVersion <= maxVersionSeen.dbVersion; dbVersion++) { + const maxColForThisDb = dbVersion === maxVersionSeen.dbVersion ? maxVersionSeen.colVersion : Number.MAX_SAFE_INTEGER; + ranges.push({ + dbVersion, + colVersions: Array.from({ length: maxColForThisDb }, (_, i) => i + 1) + }); + } + } + return ranges; + } + + // Group by db_version + const versionGroups = new Map>(); + for (const pair of versionPairs) { + if (!versionGroups.has(pair.dbVersion)) { + versionGroups.set(pair.dbVersion, new Set()); + } + versionGroups.get(pair.dbVersion)!.add(pair.colVersion); + } + + // Find missing col_versions for each db_version from 1 to maxVersionSeen + for (let dbVersion = 1; dbVersion <= maxVersionSeen.dbVersion; dbVersion++) { + const existingColVersions = versionGroups.get(dbVersion) || new Set(); + const maxColForThisDb = dbVersion === maxVersionSeen.dbVersion ? maxVersionSeen.colVersion : Math.max(...Array.from(existingColVersions), 0); + + if (maxColForThisDb === 0) { + // No col_versions exist for this db_version, skip it for now + // This might indicate we don't have any changes for this db_version yet + continue; + } + + const missingColVersions: number[] = []; + for (let colVersion = 1; colVersion <= maxColForThisDb; colVersion++) { + if (!existingColVersions.has(colVersion)) { + missingColVersions.push(colVersion); + } + } + + if (missingColVersions.length > 0) { + ranges.push({ + dbVersion, + colVersions: missingColVersions + }); + } + } + + return ranges; +} + // Also update the applyChanges function to handle decoding correctly async function applyChanges(dbname: string, changes: any[]) { const connection = connections[dbname]; @@ -499,6 +973,9 @@ async function applyChanges(dbname: string, changes: any[]) { logDebug(`Updated lastSyncVersion to ${maxVersion}`); } + // Recalculate version state after applying changes + await recalculateVersionState(dbname); + // Notify main thread that database was updated logDebug(`🔄 DATABASE UPDATED: Notifying main thread for ${dbname} with ${changes.length} changes`); self.postMessage({ diff --git a/mast-sync-protocol-v3-spec.md b/mast-sync-protocol-v3-spec.md new file mode 100644 index 0000000..4839d9f --- /dev/null +++ b/mast-sync-protocol-v3-spec.md @@ -0,0 +1,416 @@ +# Mast Sync Protocol v2.0 Specification + +## Overview + +This is a clean v2 specification that starts fresh from v1, incorporating lessons learned to solve critical issues while maintaining simplicity. This replaces the previous over-engineered v2 implementation with a focused, production-ready protocol. + +## Core Design Principles + +1. **Secure by Default**: All operations require authentication, rooms are private by default +2. **Simple & Clean**: Remove unnecessary complexity while maintaining reliability +3. **Payment Ready**: Built-in support for paid room creation with backwards compatibility +4. **CRDT-Native**: Leverage CR-SQLite's conflict resolution, no complex coordination needed +5. **Immediate Feedback**: Clients know their room status immediately after connection + +## Authentication & Authorization + +### Key Management +- **ECDSA P-256** public key cryptography +- **Automatic registration** on first connection (when payment enforcement disabled) +- **Room-scoped permissions**: read, write, invite capabilities +- **Signature format**: `{type}:{data-json}` + +### Environment-Based Enforcement +```bash +# Development/Testing (default) +REQUIRE_AUTH=false # Auto-grant invite permissions to any connecting user + +# Production (future) +REQUIRE_AUTH=true # Only users in auth table can access rooms +``` + +## Connection Lifecycle + +### 1. WebSocket Connection +``` +URL: wss://server/sync?room={roomId}&publicKey={base64PublicKey} +``` + +### 2. Immediate Room Status +Upon connection, server immediately sends room status: + +```json +{ + "type": "room_status", + "access": "write|read|none|no_room" +} +``` + +**Access Levels**: +- `"write"`: Full read/write access to existing room +- `"read"`: Read-only access to existing room +- `"none"`: Room exists but no permissions granted (need to be invited -- later) +- `"no_room"`: Room doesn't exist (payment may be required if they wish to create one) + +### 3. Room Access Behavior + +**When `REQUIRE_AUTH=false` (Development)**: +- Any room → Auto-grant invite permissions (read + write) +- New rooms created automatically +- User receives `"access": "write"` + +**When `REQUIRE_AUTH=true` (Production)**: +- Only authenticated users can access +- User must exist in room_keys table +- User receives access level based on permissions + +## Core Sync Protocol + +### Missing Changes Problem Solution +The critical flaw in v1 was using `MAX(db_version)` which caused missing intermediate versions. + +**Problem**: Client has versions [1,2,5,10] and requests `> 10`, never getting versions 3,4,6,7,8,9. + +**Solution**: Track highest contiguous version + explicit missing ranges. + +#### Version State Tracking +```javascript +const versionState = { + contiguousUpTo: 2, // Highest version with no gaps before it + missingRanges: [ // Explicit gaps we know about + {start: 3, end: 4}, // Missing versions 3-4 + {start: 6, end: 9} // Missing versions 6-9 + ], + maxVersionSeen: 10 // Highest version ever seen +}; +``` + +### Change Synchronization + +#### Sync Request (Client → Server) +```json +{ + "type": "sync_request", + "site_id": "base64-site-id", + "contiguous_up_to": 2, + "missing_ranges": [ + {"start": 3, "end": 4}, + {"start": 6, "end": 9} + ], + "max_version_seen": 10, + "publicKey": "base64-public-key" +} +``` + +#### Sync Response (Server → Client) +```json +{ + "type": "sync_response", + "current_max_version": 15, + "changes": [ + // Missing ranges 3-4, 6-9 + // Plus new changes 11-15 + { + "TableName": "todos", + "PK": "base64-pk", + "ColumnName": "description", + "Value": "Task content", + "ColVersion": 15, + "DBVersion": 8, + "SiteID": "base64-site-id", + "CL": 1, + "Seq": 1 + } + ] +} +``` + +#### Change Push (Write Operations) +```json +{ + "type": "changes", + "publicKey": "base64-public-key", + "signature": "base64-signature", + "data": [...changes...] +} +``` + +**Signature payload**: `changes:{JSON.stringify(data)}` + +### Real-time Broadcasting +Server broadcasts changes to all authorized clients in room (excluding sender): + +```json +{ + "type": "changes", + "data": [...changes...] +} +``` + +## Server Architecture + +### Essential Components +- ✅ ECDSA signature verification +- ✅ Room-based isolation +- ✅ Permission checking +- ✅ Auto room creation (with environment flag) +- ✅ Real-time broadcasting +- ✅ Connection cleanup on auth failure + +### Database Schema +```sql +CREATE TABLE rooms ( + room_id TEXT PRIMARY KEY, + created_at INTEGER NOT NULL +); + +CREATE TABLE room_keys ( + room_id TEXT NOT NULL, + public_key TEXT NOT NULL, + can_read BOOLEAN NOT NULL DEFAULT 1, + can_write BOOLEAN NOT NULL DEFAULT 1, + can_invite BOOLEAN NOT NULL DEFAULT 0, + created_at INTEGER NOT NULL, + PRIMARY KEY (room_id, public_key) +); +``` + +### Server Version Query Logic + +#### Multi-Range SQL Query +```sql +-- Get missing ranges + new changes +SELECT * FROM crsql_changes +WHERE site_id != ? + AND ( + -- Missing range 3-4 + (db_version >= 3 AND db_version <= 4) OR + -- Missing range 6-9 + (db_version >= 6 AND db_version <= 9) OR + -- New changes 11-15 + (db_version > 10) + ) +ORDER BY db_version ASC +``` + +#### Client Version State Updates +```javascript +function updateVersionState(changes) { + for (const change of changes) { + const version = change.DBVersion; + + // Update max seen + versionState.maxVersionSeen = Math.max(versionState.maxVersionSeen, version); + + // Fill gaps and update contiguous + fillGapsAndUpdateContiguous(version); + } +} + +function fillGapsAndUpdateContiguous(version) { + // Remove version from missing ranges + removeMissingVersion(version); + + // Extend contiguous if possible + while (versionState.contiguousUpTo + 1 <= versionState.maxVersionSeen && + !isVersionMissing(versionState.contiguousUpTo + 1)) { + versionState.contiguousUpTo++; + } +} +``` + +## Auto-Registration Flow + +### Connection Logic +1. Client connects with `publicKey` parameter +2. Server checks room + key permissions +3. **When `REQUIRE_AUTH=false`**: + - Auto-grant invite permissions to any user + - Create room if it doesn't exist +4. **When `REQUIRE_AUTH=true`**: + - Only users in auth table can access + - No auto-registration + +### Permission Granting Logic +```go +// Auto-grant invite permissions (read + write) +func AutoGrantInvitePermissions(roomID, publicKey string) error { + return GrantPermissions(roomID, publicKey, true, true, true) // read, write, invite +} +``` + +## Client Implementation + +### Connection Management +```typescript +// Simple reconnection (no exponential backoff) +function attemptReconnection(room: string) { + setTimeout(() => { + connectWebSocket(room).catch(() => attemptReconnection(room)); + }, math.Rand(3, 9); // Between 3 and 9 seconds retry, to stop stampeding herd +} +``` + +### Authentication Integration +```typescript +// Include publicKey in all requests +async function sendSyncRequest(connection) { + const syncRequest = { + type: "sync_request", + site_id: connection.siteId, + last_version: connection.lastSyncVersion, + publicKey // Always include for all syncRequest + }; + + ws.send(JSON.stringify(syncRequest)); +} +``` + +### Room Status Handling +```typescript +// Handle immediate room status +case 'room_status': + self.postMessage({ + type: 'room_status', + dbname: dbname, + access: msg.access, + needsPayment: msg.access === 'no_room' + }); + break; +``` + +## Server Implementation + +### Core Logic +```go +func handleWebSocket(w http.ResponseWriter, r *http.Request) { + roomID := r.URL.Query().Get("room") + publicKey := r.URL.Query().Get("publicKey") + + // Check authentication/auto-grant permissions + access := determineAccess(roomID, publicKey) + + // Upgrade WebSocket + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + + // Send immediate room status + sendRoomStatus(conn, access) + + // Handle sync protocol + handleSyncProtocol(conn, roomID, publicKey) +} + +func determineAccess(roomID, publicKey string) string { + if !requireAuth { + // Development mode - auto-grant invite permissions + autoGrantInvitePermissions(roomID, publicKey) + return "write" + } + + // Production mode - check existing permissions + return checkUserPermissions(roomID, publicKey) +} +``` + +## Error Handling + +### Structured Error Responses +```json +{ + "type": "error", + "code": "AUTH_FAILED|ROOM_NOT_FOUND|PAYMENT_REQUIRED|PERMISSION_DENIED", + "message": "Human readable description" +} +``` + +### Authentication Failure Behavior +- **Invalid signature**: Close connection immediately +- **No read permission**: Close connection immediately +- **No write permission**: Reject change, keep connection open +- **No room**: Send `room_status` with `"access": "no_room"` + +## Migration from v1 + +### Environment Variable Control +```bash +# Start with development mode +REQUIRE_AUTH=false + +# Switch to production when ready +REQUIRE_AUTH=true +``` + +### Protocol Changes from v1 +- Replace `"pull"` message with `"sync_request"` +- Add version range tracking instead of simple `last_version` +- Add immediate `"room_status"` response +- Add `publicKey` to WebSocket URL and all requests + +### Seamless Transition +- v1 rooms continue working unchanged +- Auto-registration ensures no user disruption +- Environment variable provides clean cutoff point + +## Security Model + +### Transport Security +- **WSS required** for production (TLS 1.3 minimum) +- **Certificate validation** on client + +### Message Security +- **ECDSA P-256 signatures** for all write operations +- **Public key authentication** for all read operations +- **Room isolation** - no cross-room access + +### Key Security +- **Web Crypto API** for key generation +- **Secure storage** (recommend upgrade from localStorage for production) +- **No key transmission** (only public keys sent to server) + +### CR-SQLite Benefits +- **Offline-first**: Changes work immediately without server +- **Conflict-free**: Automatic merge resolution +- **Consistent**: Guaranteed eventual consistency +- **Efficient**: Delta-only synchronization + +## Implementation Priorities + +### Phase 1: Core Protocol +1. Implement missing changes solution with version range tracking through sync_request messages + - Version state tracking on client side + - Version requests in sync_request messages + - Response from server giving exactly the changes requested +2. Add immediate room_status messages +3. Simple reconnection (3-second retry) +4. Room status handling in UI (SyncStatus component) + +### Phase 2: Authorization +1. Environment-based auth enforcement (`REQUIRE_AUTH` flag) +2. Auto-registration for development mode +3. ECDSA signature verification (real implementation) +4. publicKey authentication for all operations +5. Connection cleanup on auth failure + +## Configuration + +### Server Configuration +```go +// Environment variables +var ( + REQUIRE_AUTH = os.Getenv("REQUIRE_AUTH") == "true" +) +``` + +### Client Configuration +```typescript +interface SyncConfig { + room: string; + endpoint: string; + autoReconnect?: boolean; // Default: true + reconnectDelay?: number; // Default: 3000ms +} +``` + diff --git a/server/auth.go b/server/auth.go index 029648b..5ae1a18 100644 --- a/server/auth.go +++ b/server/auth.go @@ -276,6 +276,30 @@ func CheckKeyPermission(roomID, publicKey string, permType string) (bool, error) return hasPerm, nil } +// AutoGrantInvitePermissions auto-grants invite permissions (read + write + invite) for development mode +func AutoGrantInvitePermissions(roomID, publicKey string) error { + logKey := publicKey + if len(logKey) > 20 { + logKey = publicKey[:20] + "..." + } + log.Printf("Auto-granting invite permissions for key %s in room %s", logKey, roomID) + + _, err := authDB.Exec( + `INSERT OR REPLACE INTO room_keys + (room_id, public_key, can_read, can_write, can_invite) + VALUES (?, ?, ?, ?, ?)`, + roomID, publicKey, true, true, true, + ) + + if err != nil { + log.Printf("Error auto-granting permissions: %v", err) + return err + } + + log.Printf("Successfully auto-granted invite permissions") + return nil +} + // VerifySignature verifies that a signature was made by the public key func VerifySignature(publicKey string, data string, signature string) (bool, error) { // This is a placeholder - the actual implementation will depend on how you're handling diff --git a/server/main.go b/server/main.go index c32fee6..7f177c1 100644 --- a/server/main.go +++ b/server/main.go @@ -4,9 +4,11 @@ import ( "database/sql" "encoding/base64" "encoding/json" + "fmt" "log" "net/http" "os" + "strings" "sync" "time" @@ -21,6 +23,9 @@ var upgrader = websocket.Upgrader{ }, } +// Environment configuration +var requireAuth = os.Getenv("REQUIRE_AUTH") == "true" + // Room management type Room struct { clients map[*websocket.Conn]bool @@ -113,13 +118,20 @@ func handleWebSocket(w http.ResponseWriter, r *http.Request) { // Set CORS headers for the WebSocket handshake w.Header().Set("Access-Control-Allow-Origin", "*") - // Extract room ID from query parameters + // Extract room ID and publicKey from query parameters roomID := r.URL.Query().Get("room") + publicKey := r.URL.Query().Get("publicKey") + if roomID == "" { http.Error(w, "Missing room parameter", http.StatusBadRequest) return } + if publicKey == "" { + http.Error(w, "Missing publicKey parameter", http.StatusBadRequest) + return + } + // Get or create the room in the auth database err := GetOrCreateRoom(roomID) if err != nil { @@ -128,13 +140,27 @@ func handleWebSocket(w http.ResponseWriter, r *http.Request) { return } + // Determine access level + access := determineAccess(roomID, publicKey) + + // Upgrade WebSocket conn, err := upgrader.Upgrade(w, r, nil) if err != nil { log.Println("Error upgrading connection:", err) return } - // Add client to room + // Send immediate room status + sendRoomStatus(conn, access) + + // Close connection if no read access + if access == "none" || access == "no_room" { + log.Printf("Closing connection for user with access: %s", access) + conn.Close() + return + } + + // Add client to room only if authorized addClientToRoom(roomID, conn) // Create database connection for this room @@ -163,9 +189,6 @@ func handleWebSocket(w http.ResponseWriter, r *http.Request) { removeClientFromRoom(roomID, conn) }() - // We don't automatically send initial changes anymore - // The client will request them with a pull message that includes requestUnsyncedChanges - // Handle incoming messages for { _, message, err := conn.ReadMessage() @@ -187,8 +210,64 @@ func handleWebSocket(w http.ResponseWriter, r *http.Request) { } switch msgType { + case "sync_request": + log.Printf("Client in room %s requested sync", roomID) + var syncMsg SyncRequestMessage + syncData, _ := json.Marshal(msg) + json.Unmarshal(syncData, &syncMsg) + + log.Printf("Sync request with site_id: %s, contiguous_up_to: %d.%d, max_version_seen: %d.%d", + syncMsg.SiteID, syncMsg.ContiguousUpTo.DBVersion, syncMsg.ContiguousUpTo.ColVersion, + syncMsg.MaxVersionSeen.DBVersion, syncMsg.MaxVersionSeen.ColVersion) + + // Verify publicKey matches + if syncMsg.PublicKey != publicKey { + log.Printf("PublicKey mismatch in sync request") + sendError(conn, "AUTH_FAILED", "PublicKey mismatch") + continue + } + + // Get current server version pair + serverMaxVersionPair, err := getLatestDBVersionCol(db) + if err != nil { + log.Printf("Error getting server's latest version pair: %v", err) + serverMaxVersionPair = VersionColPair{DBVersion: 0, ColVersion: 0} + } + log.Printf("Server max version: %d.%d, client max version seen: %d.%d", + serverMaxVersionPair.DBVersion, serverMaxVersionPair.ColVersion, + syncMsg.MaxVersionSeen.DBVersion, syncMsg.MaxVersionSeen.ColVersion) + + // Check if client has newer data than server + if compareVersionPairs(syncMsg.MaxVersionSeen, serverMaxVersionPair) > 0 { + log.Printf("Client has higher version (%d.%d) than server (%d.%d). Requesting changes.", + syncMsg.MaxVersionSeen.DBVersion, syncMsg.MaxVersionSeen.ColVersion, + serverMaxVersionPair.DBVersion, serverMaxVersionPair.ColVersion) + + // Send a request_changes message to the client + requestChangesMsg := RequestChangesMessage{ + Type: "request_changes", + RoomID: roomID, + Version: serverMaxVersionPair, + } + + requestJSON, _ := json.Marshal(requestChangesMsg) + conn.WriteMessage(websocket.TextMessage, requestJSON) + } + + // Always send sync_response with any changes the server has for the client + changes := getChangesForSyncRequest(db, syncMsg) + + response := SyncResponseMessage{ + Type: "sync_response", + CurrentMaxVersion: serverMaxVersionPair, + Changes: changes, + } + responseJSON, _ := json.Marshal(response) + conn.WriteMessage(websocket.TextMessage, responseJSON) + case "pull": - log.Printf("Client in room %s requested pull", roomID) + // Legacy support for v1 protocol + log.Printf("Client in room %s requested pull (legacy)", roomID) var pullMsg PullMessage pullData, _ := json.Marshal(msg) json.Unmarshal(pullData, &pullMsg) @@ -211,7 +290,7 @@ func handleWebSocket(w http.ResponseWriter, r *http.Request) { requestChangesMsg := RequestChangesMessage{ Type: "request_changes", RoomID: roomID, - Version: serverLatestVersion, + Version: VersionColPair{DBVersion: serverLatestVersion, ColVersion: 0}, } requestJSON, _ := json.Marshal(requestChangesMsg) @@ -229,8 +308,40 @@ func handleWebSocket(w http.ResponseWriter, r *http.Request) { case "changes": log.Printf("Received changes from client in room %s", roomID) - if publicKey, hasKey := msg["publicKey"].(string); hasKey { - log.Printf("Changes are authenticated with public key: %s...", publicKey[:20]) + + // Check write permission + if access != "write" { + log.Printf("User has no write permission for room %s", roomID) + sendError(conn, "PERMISSION_DENIED", "No write permission") + continue + } + + // For signed changes, verify the signature + if msgPublicKey, hasKey := msg["publicKey"].(string); hasKey { + log.Printf("Changes are authenticated with public key: %s...", msgPublicKey[:20]) + + // Verify publicKey matches connection + if msgPublicKey != publicKey { + log.Printf("PublicKey mismatch in changes message") + sendError(conn, "AUTH_FAILED", "PublicKey mismatch") + continue + } + + // TODO: Verify signature when ECDSA is implemented + if signature, hasSig := msg["signature"].(string); hasSig && requireAuth { + dataStr := "" + if data, ok := msg["data"]; ok { + dataBytes, _ := json.Marshal(data) + dataStr = string(dataBytes) + } + + isValid, err := VerifySignature(msgPublicKey, "changes:"+dataStr, signature) + if err != nil || !isValid { + log.Printf("Signature verification failed") + sendError(conn, "AUTH_FAILED", "Invalid signature") + continue + } + } } if data, ok := msg["data"].([]interface{}); ok { @@ -315,7 +426,7 @@ func sendInitialChanges(conn *websocket.Conn, db *sql.DB) { conn.WriteMessage(websocket.TextMessage, responseJSON) } -// getLatestDBVersion gets the latest db_version from the database +// getLatestDBVersion gets the latest db_version from the database (legacy) func getLatestDBVersion(db *sql.DB) (int, error) { // Query the maximum db_version row := db.QueryRow("SELECT MAX(db_version) FROM crsql_changes") @@ -333,6 +444,61 @@ func getLatestDBVersion(db *sql.DB) (int, error) { return int(version.Int64), nil } +// getLatestDBVersionCol gets the latest (db_version, col_version) pair from the database +func getLatestDBVersionCol(db *sql.DB) (VersionColPair, error) { + // First check if there are any changes at all + var count int + err := db.QueryRow("SELECT COUNT(*) FROM crsql_changes").Scan(&count) + if err != nil { + return VersionColPair{DBVersion: 0, ColVersion: 0}, err + } + + // If no changes exist, return 0,0 + if count == 0 { + return VersionColPair{DBVersion: 0, ColVersion: 0}, nil + } + + // Query the maximum db_version and its maximum col_version + row := db.QueryRow(` + SELECT db_version, col_version + FROM crsql_changes + ORDER BY db_version DESC, col_version DESC + LIMIT 1 + `) + + var dbVersion, colVersion sql.NullInt64 + if err := row.Scan(&dbVersion, &colVersion); err != nil { + if err == sql.ErrNoRows { + return VersionColPair{DBVersion: 0, ColVersion: 0}, nil + } + return VersionColPair{DBVersion: 0, ColVersion: 0}, err + } + + // If there are no rows or the values are null, return 0,0 + if !dbVersion.Valid || !colVersion.Valid { + return VersionColPair{DBVersion: 0, ColVersion: 0}, nil + } + + return VersionColPair{DBVersion: int(dbVersion.Int64), ColVersion: int(colVersion.Int64)}, nil +} + +// compareVersionPairs compares two version pairs, returns -1, 0, or 1 +func compareVersionPairs(a, b VersionColPair) int { + if a.DBVersion != b.DBVersion { + if a.DBVersion < b.DBVersion { + return -1 + } + return 1 + } + if a.ColVersion != b.ColVersion { + if a.ColVersion < b.ColVersion { + return -1 + } + return 1 + } + return 0 +} + func getChangesFromDB(db *sql.DB, siteID string, version int) []map[string]interface{} { // Decode the site_id from base64 decodedSiteID, err := decodeBase64(siteID) @@ -444,7 +610,7 @@ func createDirIfNotExists(path string) error { } -// PullMessage represents the structure of a 'pull' type message from clients +// PullMessage represents the structure of a 'pull' type message from clients (legacy v1) type PullMessage struct { Type string `json:"type"` SiteID string `json:"site_id"` @@ -453,9 +619,57 @@ type PullMessage struct { // RequestChangesMessage represents a request from server to client to send their changes type RequestChangesMessage struct { + Type string `json:"type"` + RoomID string `json:"room_id"` + Version VersionColPair `json:"version"` +} + +// VersionColPair represents a (db_version, col_version) pair +type VersionColPair struct { + DBVersion int `json:"db_version"` + ColVersion int `json:"col_version"` +} + +// MissingVersionRange represents missing col_versions for a specific db_version +type MissingVersionRange struct { + DBVersion int `json:"db_version"` + ColVersions []int `json:"col_versions"` +} + +// VersionRange represents a missing version range (legacy - kept for compatibility) +type VersionRange struct { + Start int `json:"start"` + End int `json:"end"` +} + +// SyncRequestMessage represents the v2 sync request with version ranges +type SyncRequestMessage struct { + Type string `json:"type"` + SiteID string `json:"site_id"` + ContiguousUpTo VersionColPair `json:"contiguous_up_to"` + MissingRanges []MissingVersionRange `json:"missing_ranges"` + MaxVersionSeen VersionColPair `json:"max_version_seen"` + PublicKey string `json:"publicKey"` +} + +// SyncResponseMessage represents the server response to sync request +type SyncResponseMessage struct { + Type string `json:"type"` + CurrentMaxVersion VersionColPair `json:"current_max_version"` + Changes []map[string]interface{} `json:"changes"` +} + +// RoomStatusMessage represents immediate room status on connection +type RoomStatusMessage struct { + Type string `json:"type"` + Access string `json:"access"` +} + +// ErrorMessage represents structured error responses +type ErrorMessage struct { Type string `json:"type"` - RoomID string `json:"room_id"` - Version int `json:"version"` + Code string `json:"code"` + Message string `json:"message"` } // AuthRequest is the structure for authentication verification requests @@ -559,7 +773,196 @@ func handleAuthVerify(w http.ResponseWriter, r *http.Request) { json.NewEncoder(w).Encode(response) } +// determineAccess determines the access level for a user connecting to a room +func determineAccess(roomID, publicKey string) string { + logKey := publicKey + if len(logKey) > 20 { + logKey = publicKey[:20] + "..." + } + log.Printf("Determining access for publicKey %s in room %s", logKey, roomID) + + // Check if room exists + roomExists, err := CheckRoomExists(roomID) + if err != nil { + log.Printf("Error checking room existence: %v", err) + return "none" + } + + if !roomExists { + log.Printf("Room %s does not exist", roomID) + if !requireAuth { + // Development mode - auto-create room and grant permissions + log.Printf("Auto-creating room %s and granting permissions", roomID) + if err := GetOrCreateRoom(roomID); err != nil { + log.Printf("Error auto-creating room: %v", err) + return "no_room" + } + if err := AutoGrantInvitePermissions(roomID, publicKey); err != nil { + log.Printf("Error auto-granting permissions: %v", err) + return "no_room" + } + return "write" + } + return "no_room" + } + + // Room exists - check permissions + if !requireAuth { + // Development mode - auto-grant permissions for existing rooms + log.Printf("Auto-granting permissions for existing room %s", roomID) + if err := AutoGrantInvitePermissions(roomID, publicKey); err != nil { + log.Printf("Error auto-granting permissions: %v", err) + } + return "write" + } + + // Production mode - check actual permissions + hasRead, err := CheckKeyPermission(roomID, publicKey, "read") + if err != nil { + log.Printf("Error checking read permission: %v", err) + return "none" + } + + if !hasRead { + return "none" + } + + hasWrite, err := CheckKeyPermission(roomID, publicKey, "write") + if err != nil { + log.Printf("Error checking write permission: %v", err) + return "read" + } + + if hasWrite { + return "write" + } + + return "read" +} + +// sendRoomStatus sends immediate room status to client +func sendRoomStatus(conn *websocket.Conn, access string) { + status := RoomStatusMessage{ + Type: "room_status", + Access: access, + } + + statusJSON, _ := json.Marshal(status) + conn.WriteMessage(websocket.TextMessage, statusJSON) + log.Printf("Sent room status: %s", access) +} + +// sendError sends structured error response to client +func sendError(conn *websocket.Conn, code, message string) { + errorMsg := ErrorMessage{ + Type: "error", + Code: code, + Message: message, + } + + errorJSON, _ := json.Marshal(errorMsg) + conn.WriteMessage(websocket.TextMessage, errorJSON) + log.Printf("Sent error: %s - %s", code, message) +} + +// getChangesForSyncRequest gets changes for missing ranges + new changes +func getChangesForSyncRequest(db *sql.DB, syncMsg SyncRequestMessage) []map[string]interface{} { + decodedSiteID, err := decodeBase64(syncMsg.SiteID) + if err != nil { + log.Printf("Error decoding site_id '%s': %v", syncMsg.SiteID, err) + decodedSiteID = []byte{} + } + + // Build query for missing (db_version, col_version) pairs + new changes + var queryParts []string + var args []interface{} + + // Add site_id filter + baseCond := "site_id != ?" + args = append(args, decodedSiteID) + + // Add missing ranges - for each db_version, get specific missing col_versions + for _, vrange := range syncMsg.MissingRanges { + if len(vrange.ColVersions) > 0 { + // Create placeholders for the col_versions + placeholders := make([]string, len(vrange.ColVersions)) + for i := range placeholders { + placeholders[i] = "?" + args = append(args, vrange.ColVersions[i]) + } + + // Add condition for this db_version with specific col_versions + queryParts = append(queryParts, + fmt.Sprintf("(db_version = ? AND col_version IN (%s))", + strings.Join(placeholders, ","))) + args = append(args, vrange.DBVersion) + } + } + + // Add new changes beyond max version seen + // This includes: db_version > max_db_version OR (db_version = max_db_version AND col_version > max_col_version) + queryParts = append(queryParts, + "(db_version > ? OR (db_version = ? AND col_version > ?))") + args = append(args, syncMsg.MaxVersionSeen.DBVersion, + syncMsg.MaxVersionSeen.DBVersion, syncMsg.MaxVersionSeen.ColVersion) + + // Combine all conditions + whereClause := baseCond + if len(queryParts) > 0 { + whereClause += " AND (" + strings.Join(queryParts, " OR ") + ")" + } + + query := "SELECT * FROM crsql_changes WHERE " + whereClause + " ORDER BY db_version ASC, col_version ASC" + + log.Printf("Executing sync query with %d missing ranges and max_version %d.%d", + len(syncMsg.MissingRanges), syncMsg.MaxVersionSeen.DBVersion, syncMsg.MaxVersionSeen.ColVersion) + + rows, err := db.Query(query, args...) + if err != nil { + log.Printf("Error querying changes for sync: %v", err) + return nil + } + defer rows.Close() + + var changes []map[string]interface{} + + for rows.Next() { + var tableName string + var pk []byte + var columnName string + var value interface{} + var colVersion, dbVersion int64 + var siteID []byte + var cl, seq int64 + + if err := rows.Scan(&tableName, &pk, &columnName, &value, &colVersion, &dbVersion, &siteID, &cl, &seq); err != nil { + log.Printf("Error scanning row: %v", err) + continue + } + + change := map[string]interface{}{ + "TableName": tableName, + "PK": encodeToBase64(pk), + "ColumnName": columnName, + "Value": value, + "ColVersion": colVersion, + "DBVersion": dbVersion, + "SiteID": encodeToBase64(siteID), + "CL": cl, + "Seq": seq, + } + + changes = append(changes, change) + } + + log.Printf("Returning %d changes for sync request", len(changes)) + return changes +} + func main() { + // Log startup configuration + log.Printf("Starting server with REQUIRE_AUTH=%v", requireAuth) + // Register SQLite with CR-SQLite extension sql.Register("sqlite3_with_extensions", &sqlite3.SQLiteDriver{ Extensions: []string{"../db/crsqlite"}, -- 2.51.2