Send disconnect messages if the session is bad

This commit is contained in:
Owen
2026-09-29 10:24:59 -04:00
parent 7a655f7c6e
commit e3ce453155
4 changed files with 97 additions and 6 deletions
+38 -2
View File
@@ -329,6 +329,12 @@ const initializeRedisSubscription = async (): Promise<void> => {
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<boolean> => {
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;
}
+2 -2
View File
@@ -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"
)
);
+52 -2
View File
@@ -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;
}
}
+5
View File
@@ -104,4 +104,9 @@ export type RedisMessage =
message: WSMessage;
fromNodeId: string;
options?: SendMessageOptions;
}
| {
type: "disconnect";
targetClientId: string;
fromNodeId: string;
};