mirror of
https://github.com/fosrl/pangolin.git
synced 2026-09-30 17:59:05 +02:00
Send disconnect messages if the session is bad
This commit is contained in:
@@ -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;
|
||||
}
|
||||
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
);
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -104,4 +104,9 @@ export type RedisMessage =
|
||||
message: WSMessage;
|
||||
fromNodeId: string;
|
||||
options?: SendMessageOptions;
|
||||
}
|
||||
| {
|
||||
type: "disconnect";
|
||||
targetClientId: string;
|
||||
fromNodeId: string;
|
||||
};
|
||||
|
||||
Reference in New Issue
Block a user