diff --git a/server/private/routers/ws/ws.ts b/server/private/routers/ws/ws.ts index 4a79ff006..7327c7067 100644 --- a/server/private/routers/ws/ws.ts +++ b/server/private/routers/ws/ws.ts @@ -329,6 +329,12 @@ const initializeRedisSubscription = async (): Promise => { redisMessage.message, redisMessage.excludeClientId ); + } else if ( + redisMessage.type === "disconnect" && + redisMessage.targetClientId + ) { + // Close the client's connections if any are on this node + disconnectClientLocal(redisMessage.targetClientId); } } catch (error) { logger.error("Error processing Redis message:", error); @@ -1285,13 +1291,43 @@ if (redisManager.isRedisEnabled()) { logger.debug("WebSocket handler initialized in local mode"); } -// Disconnect a specific client and force them to reconnect +// Disconnect a specific client on every node and force them to reconnect const disconnectClient = async (clientId: string): Promise => { + const localDisconnected = disconnectClientLocal(clientId); + + if (!redisManager.isRedisEnabled()) { + return localDisconnected; + } + + try { + // Flush queued direct messages first so anything sent to this client + // just before (e.g. olm/terminate) reaches the other nodes ahead of + // the disconnect on the same channel. + await flushPendingRedisDirectMessages(); + + const redisMessage: RedisMessage = { + type: "disconnect", + targetClientId: clientId, + fromNodeId: NODE_ID + }; + await redisManager.publish(REDIS_CHANNEL, JSON.stringify(redisMessage)); + return true; + } catch (error) { + logger.error( + `Failed to publish disconnect for client ID ${clientId} via Redis:`, + error + ); + return localDisconnected; + } +}; + +// Disconnect a specific client's connections on this node only +const disconnectClientLocal = (clientId: string): boolean => { const mapKey = getClientMapKey(clientId); const clients = connectedClients.get(mapKey); if (!clients || clients.length === 0) { - logger.debug(`No connections found for client ID: ${clientId}`); + logger.debug(`No local connections found for client ID: ${clientId}`); return false; } diff --git a/server/routers/olm/getOlmToken.ts b/server/routers/olm/getOlmToken.ts index c7c0c127b..7808b558b 100644 --- a/server/routers/olm/getOlmToken.ts +++ b/server/routers/olm/getOlmToken.ts @@ -98,13 +98,13 @@ export async function getOlmToken( await validateSessionToken(userToken); if (!userSession || !user) { return next( - createHttpError(HttpCode.BAD_REQUEST, "Invalid user token") + createHttpError(HttpCode.UNAUTHORIZED, "Invalid user token") ); } if (user.userId !== existingOlm.userId) { return next( createHttpError( - HttpCode.BAD_REQUEST, + HttpCode.UNAUTHORIZED, "User token does not match olm" ) ); diff --git a/server/routers/olm/handleOlmPingMessage.ts b/server/routers/olm/handleOlmPingMessage.ts index 0e18c7f5b..aa6e8a2b9 100644 --- a/server/routers/olm/handleOlmPingMessage.ts +++ b/server/routers/olm/handleOlmPingMessage.ts @@ -1,4 +1,4 @@ -import { getClientConfigVersion } from "#dynamic/routers/ws"; +import { disconnectClient, getClientConfigVersion } from "#dynamic/routers/ws"; import { db } from "@server/db"; import { MessageHandler } from "@server/routers/ws"; import { clients, Olm } from "@server/db"; @@ -11,6 +11,29 @@ import { encodeHexLowerCase } from "@oslojs/encoding"; import { sha256 } from "@oslojs/crypto/sha2"; import { sendOlmSyncMessage } from "./sync"; import { handleFingerprintInsertion } from "./fingerprintingUtils"; +import { sendTerminateClient } from "../client/terminate"; +import { OlmErrorCodes } from "./error"; + +type OlmErrorCode = (typeof OlmErrorCodes)[keyof typeof OlmErrorCodes]; + +/** + * Tells the olm why it is being kicked, then closes its websocket so it + * does not linger connected until the offline checker notices. + */ +async function terminateOlm( + clientId: number, + olmId: string, + error: OlmErrorCode +) { + try { + await sendTerminateClient(clientId, error, olmId); + // wait a moment to ensure the message is sent + await new Promise((resolve) => setTimeout(resolve, 1000)); + await disconnectClient(olmId); + } catch (err) { + logger.error(`Error terminating olm ${olmId}`, { error: err }); + } +} /** * Handles ping messages from clients and responds with pong @@ -60,14 +83,29 @@ export const handleOlmPingMessage: MessageHandler = async (context) => { await validateSessionToken(userToken); if (!userSession || !user) { logger.warn("Invalid user session for olm ping"); - return; // by returning here we just ignore the ping and the setInterval will force it to disconnect + await terminateOlm( + client.clientId, + olm.olmId, + OlmErrorCodes.INVALID_USER_SESSION + ); + return; } if (user.userId !== olm.userId) { logger.warn("User ID mismatch for olm ping"); + await terminateOlm( + client.clientId, + olm.olmId, + OlmErrorCodes.USER_ID_MISMATCH + ); return; } if (user.userId !== client.userId) { logger.warn("Client user ID mismatch for olm ping"); + await terminateOlm( + client.clientId, + olm.olmId, + OlmErrorCodes.USER_ID_MISMATCH + ); return; } @@ -85,6 +123,18 @@ export const handleOlmPingMessage: MessageHandler = async (context) => { logger.warn( `Olm user ${olm.userId} does not pass access policies for org ${client.orgId}: ${policyCheck.error}` ); + let error: OlmErrorCode = + OlmErrorCodes.ORG_ACCESS_POLICY_DENIED; + if (policyCheck.policies?.passwordAge?.compliant === false) { + error = OlmErrorCodes.ORG_ACCESS_POLICY_PASSWORD_EXPIRED; + } else if ( + policyCheck.policies?.maxSessionLength?.compliant === false + ) { + error = OlmErrorCodes.ORG_ACCESS_POLICY_SESSION_EXPIRED; + } else if (policyCheck.policies?.requiredTwoFactor === false) { + error = OlmErrorCodes.ORG_ACCESS_POLICY_2FA_REQUIRED; + } + await terminateOlm(client.clientId, olm.olmId, error); return; } } diff --git a/server/routers/ws/types.ts b/server/routers/ws/types.ts index eeb272457..c0b61d0fd 100644 --- a/server/routers/ws/types.ts +++ b/server/routers/ws/types.ts @@ -104,4 +104,9 @@ export type RedisMessage = message: WSMessage; fromNodeId: string; options?: SendMessageOptions; + } + | { + type: "disconnect"; + targetClientId: string; + fromNodeId: string; };