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"},