From 4d7a4c809d21e5ef0d53810be89560851c75c459 Mon Sep 17 00:00:00 2001 From: Owen Date: Fri, 31 Jul 2026 11:30:32 -0400 Subject: [PATCH] sync the exit node connection --- server/lib/rebuildClientAssociations.ts | 219 +++++++++++++++++++++++- 1 file changed, 217 insertions(+), 2 deletions(-) diff --git a/server/lib/rebuildClientAssociations.ts b/server/lib/rebuildClientAssociations.ts index f8b83f5f7..ce0897aab 100644 --- a/server/lib/rebuildClientAssociations.ts +++ b/server/lib/rebuildClientAssociations.ts @@ -19,7 +19,7 @@ import { userOrgRoles, userSiteResources } from "@server/db"; -import { and, count, eq, inArray, ne } from "drizzle-orm"; +import { and, count, eq, inArray, isNotNull, ne } from "drizzle-orm"; import { deletePeersBatch as newtDeletePeersBatch } from "@server/routers/newt/peers"; import { @@ -27,6 +27,9 @@ import { deletePeersBatch as olmDeletePeersBatch } from "@server/routers/olm/peers"; import { sendToExitNode } from "#dynamic/lib/exitNodes"; +import { sendToClientsBatch } from "#dynamic/routers/ws"; +import { canCompress } from "@server/lib/clientVersionChecks"; +import config from "@server/lib/config"; import logger from "@server/logger"; import { generateAliasConfig, @@ -187,7 +190,12 @@ export async function getClientSiteResourceAccess( `rebuildClientAssociations: [getClientSiteResourceAccess] siteResourceId=${siteResource.siteResourceId} networkId=${siteResource.networkId} siteCount=${sitesList.length} siteIds=[${sitesList.map((s) => s.siteId).join(", ")}]` ); - if (sitesList.length === 0) { + if (sitesList.length === 0 && siteResource.networkId !== null) { + // A site resource with a networkId is expected to have at least one + // site attached via siteNetworks. Resources with no networkId (e.g. + // inference-mode resources, which connect clients directly to the + // exit node instead of any site) are expected to have no sites, so + // don't warn for those. logger.warn( `No sites found for siteResource ${siteResource.siteResourceId} with networkId ${siteResource.networkId}` ); @@ -687,6 +695,22 @@ async function rebuildClientAssociationsFromSiteResourceImpl( clientSiteResourcesToRemove, trx ); + + // If this resource requires clients to be connected to the exit node + // (e.g. an inference resource), re-sync the connect/disconnect state for + // every client whose access to it may have changed - both those who + // currently have access and those who just lost it. + if (siteResource.requiresExitNodeConnection) { + await syncClientExitNodeConnections( + Array.from( + new Set([ + ...mergedAllClientIds, + ...existingClientSiteResourceIds + ]) + ), + trx + ); + } } async function handleMessagesForSiteClients( @@ -1052,6 +1076,179 @@ export async function updateClientSiteDestinations( } } +// Determines, for each of the given clients, whether they currently have +// access to any enabled site resource with requiresExitNodeConnection set +// (e.g. an inference-mode resource) and tells the client's olm to connect to +// or disconnect from its assigned exit node accordingly. Site resources with +// requiresExitNodeConnection don't belong to any site/network, so this can't +// be derived from the per-site peer logic above - it has to be recomputed +// from the client's full current resource access every time that access +// changes. +async function syncClientExitNodeConnections( + clientIds: number[], + trx: Transaction | typeof db = db +): Promise { + const uniqueClientIds = Array.from(new Set(clientIds)); + if (uniqueClientIds.length === 0) { + return; + } + + // Only clients with an exit node assigned can be told to connect/disconnect. + const clientsData = await trx + .select({ + clientId: clients.clientId, + exitNodeId: clients.exitNodeId, + exitNodeSubnet: clients.exitNodeSubnet + }) + .from(clients) + .where( + and( + inArray(clients.clientId, uniqueClientIds), + isNotNull(clients.exitNodeId) + ) + ); + + if (clientsData.length === 0) { + return; + } + + const clientIdsWithExitNode = clientsData.map((c) => c.clientId); + + const requiresExitNodeRows = await trx + .select({ + clientId: clientSiteResourcesAssociationsCache.clientId + }) + .from(clientSiteResourcesAssociationsCache) + .innerJoin( + siteResources, + eq( + clientSiteResourcesAssociationsCache.siteResourceId, + siteResources.siteResourceId + ) + ) + .where( + and( + inArray( + clientSiteResourcesAssociationsCache.clientId, + clientIdsWithExitNode + ), + eq(siteResources.enabled, true), + eq(siteResources.requiresExitNodeConnection, true) + ) + ); + + const needsConnectSet = new Set(requiresExitNodeRows.map((r) => r.clientId)); + + const exitNodeIds = Array.from( + new Set( + clientsData + .map((c) => c.exitNodeId) + .filter((id): id is number => id !== null) + ) + ); + + const exitNodeRows = + exitNodeIds.length > 0 + ? await trx + .select() + .from(exitNodes) + .where(inArray(exitNodes.exitNodeId, exitNodeIds)) + : []; + const exitNodeById = new Map(exitNodeRows.map((n) => [n.exitNodeId, n])); + + const olmRows = await trx + .select({ + clientId: olms.clientId, + olmId: olms.olmId, + version: olms.version + }) + .from(olms) + .where(inArray(olms.clientId, clientIdsWithExitNode)); + const olmByClientId = new Map( + olmRows + .filter((r) => r.clientId !== null) + .map((r) => [r.clientId as number, r]) + ); + + const relayPort = config.getRawConfig().gerbil.clients_start_port; + + const connectPayloads: { + clientId: string; + message: { type: string; data: any }; + options: { compress: boolean }; + }[] = []; + const disconnectPayloads: { + clientId: string; + message: { type: string; data: any }; + options: { compress: boolean }; + }[] = []; + + for (const client of clientsData) { + const olm = olmByClientId.get(client.clientId); + if (!olm) { + // No olm registered for this client yet/anymore, nothing to send. + continue; + } + + const needsConnect = needsConnectSet.has(client.clientId); + + if (needsConnect) { + const exitNode = client.exitNodeId + ? exitNodeById.get(client.exitNodeId) + : undefined; + if (!exitNode || !client.exitNodeSubnet) { + logger.warn( + `rebuildClientAssociations: [syncClientExitNodeConnections] client ${client.clientId} needs an exit node connection but has no exit node or subnet assigned` + ); + continue; + } + + connectPayloads.push({ + clientId: olm.olmId, + message: { + type: "olm/wg/exitnode/connect", + data: { + connect: true, + endpoint: `${exitNode.endpoint}:${exitNode.listenPort}`, + relayPort, + publicKey: exitNode.publicKey, + serverIP: exitNode.address.split("/")[0], + tunnelIP: client.exitNodeSubnet.split("/")[0] + } + }, + options: { compress: canCompress(olm.version, "olm") } + }); + } else { + disconnectPayloads.push({ + clientId: olm.olmId, + message: { + type: "olm/wg/exitnode/disconnect", + data: {} + }, + options: { compress: canCompress(olm.version, "olm") } + }); + } + } + + if (connectPayloads.length > 0) { + await sendToClientsBatch(connectPayloads).catch((error) => { + logger.error( + `rebuildClientAssociations: Error sending exit node connect messages:`, + error + ); + }); + } + + if (disconnectPayloads.length > 0) { + await sendToClientsBatch(disconnectPayloads).catch((error) => { + logger.error( + `rebuildClientAssociations: Error sending exit node disconnect messages:`, + error + ); + }); + } +} + async function handleSubnetProxyTargetUpdates( siteResource: SiteResource, sitesList: Site[], @@ -1709,6 +1906,20 @@ export async function handleMessagingForUpdatedSiteResource( ); } + // If this resource requires (or required) clients to be connected to the + // exit node (e.g. an inference resource), re-sync connect/disconnect + // state for every client currently associated with it - covers toggling + // requiresExitNodeConnection on update as well as enabling/disabling it. + if ( + updatedSiteResource.requiresExitNodeConnection || + existingSiteResource?.requiresExitNodeConnection + ) { + await syncClientExitNodeConnections( + mergedAllClients.map((c) => c.clientId), + trx + ); + } + logger.debug( `handleMessagingForUpdatedSiteResource: DONE siteResourceId=${updatedSiteResource.siteResourceId}` ); @@ -1990,6 +2201,10 @@ async function rebuildClientAssociationsFromClientImpl( resourcesToRemove, trx ); + + // Re-sync exit node connect/disconnect state based on this client's + // current full set of resource access (e.g. inference resources). + await syncClientExitNodeConnections([client.clientId], trx); } async function handleMessagesForClientSites(