Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions discojs/src/client/event_connection.ts
Original file line number Diff line number Diff line change
Expand Up @@ -262,6 +262,10 @@ export class WebSocketServer
}

disconnect(): Promise<void> {
// the server may have closed the connection already, such as when it
// crashes the client
if (this.socket.readyState === WebSocket.CLOSED) return Promise.resolve();

return new Promise((resolve, reject) => {
this.socket.onclose = () => resolve();
this.socket.onerror = (event) => reject(new Error(describeError(event)));
Expand Down
165 changes: 165 additions & 0 deletions server/src/controllers/federated_controller.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import {
weightsDecode,
weightsEncode,
} from "@epfml/discojs";
import * as msgpack from "@msgpack/msgpack";
import { assert, describe, expect, it, vi } from "vitest";
import { FederatedController } from "../../src/controllers/federated_controller.js";
import type { FederatedFakeWebSocket } from "../../tests/fake_websocket.js";
Expand Down Expand Up @@ -893,3 +894,167 @@ describe("The departure mechanism of a participant", () => {
}
});
});

describe("Crashing a participant", () => {
/** Checks that the server told a participant to crash, then disconnected it */
async function expectCrashed(ws: FederatedFakeWebSocket): Promise<void> {
expectLastMessageOfType(ws, MessageTypes.CrashClient);
await vi.waitFor(() => {
expect(ws.nbOfTerminations).toBe(1);
});
}

/** Checks that a participant was told it is now alone */
async function expectToWaitAlone(ws: FederatedFakeWebSocket): Promise<void> {
await vi.waitFor(() => {
expect(
expectLastMessageOfType(ws, MessageTypes.WaitingForMoreParticipants)
.nbOfParticipants,
).toBe(1);
});
expect(ws.nbOfTerminations).toBe(0);
}

it("disconnects a participant contributing without having connected", async () => {
const controller = await makeController(2);
const ws = makeFederatedFakeWebSocket();

controller.handle(ws); // never sends ClientConnected
await contribute(ws, DUMMY_WEIGHTS_1, 0);

await expectCrashed(ws);
});

it("disconnects a participant resending ClientConnected too many times", async () => {
const controller = await makeController(2);
const ws1 = makeFederatedFakeWebSocket();
const ws2 = makeFederatedFakeWebSocket();

connect(controller, ws1);
connect(controller, ws2);
for (
let i = 0;
i < FederatedController.MAX_CLIENT_CONNECTED_PER_SOCKET;
i++
)
ws1.emitMessage({ type: MessageTypes.ClientConnected });

await expectCrashed(ws1);
await expectToWaitAlone(ws2);
});

it("disconnects a participant contributing for a future round", async () => {
const controller = await makeController(2);
const ws1 = makeFederatedFakeWebSocket();
const ws2 = makeFederatedFakeWebSocket();

connect(controller, ws1);
connect(controller, ws2);
await contribute(ws1, DUMMY_WEIGHTS_1, 10);

await expectCrashed(ws1);
await expectToWaitAlone(ws2);
});

it("disconnects a participant sending an invalid message", async () => {
const controller = await makeController(2);
const ws1 = makeFederatedFakeWebSocket();
const ws2 = makeFederatedFakeWebSocket();

connect(controller, ws1);
connect(controller, ws2);
ws1.emit("message", msgpack.encode({ type: "not a message type" }));

await expectCrashed(ws1);
await expectToWaitAlone(ws2);
});

it("disconnects a participant sending an undecodable message", async () => {
const controller = await makeController(2);
const ws1 = makeFederatedFakeWebSocket();
const ws2 = makeFederatedFakeWebSocket();

connect(controller, ws1);
connect(controller, ws2);
ws1.emit("message", Buffer.from([0xc1])); // never used msgpack byte

await expectCrashed(ws1);
await expectToWaitAlone(ws2);
});

it("disconnects a participant sending undecodable weights", async () => {
const controller = await makeController(2);
const ws1 = makeFederatedFakeWebSocket();
const ws2 = makeFederatedFakeWebSocket();

connect(controller, ws1);
connect(controller, ws2);
ws1.emitMessage({
type: MessageTypes.SendPayload,
payload: new Uint8Array([1, 2, 3]),
round: 0,
});

await expectCrashed(ws1);
await expectToWaitAlone(ws2);
});

it("disconnects a participant whose websocket errored", async () => {
const controller = await makeController(2);
const ws1 = makeFederatedFakeWebSocket();
const ws2 = makeFederatedFakeWebSocket();

connect(controller, ws1);
connect(controller, ws2);
ws1.emit("error", new Error("socket failure"));

await vi.waitFor(() => {
expect(ws1.nbOfTerminations).toBe(1);
});
await expectToWaitAlone(ws2);
});

it("disconnects a participant it fails to send messages to", async () => {
const controller = await makeController(2);
const ws1 = makeFederatedFakeWebSocket();
const ws2 = makeFederatedFakeWebSocket();
ws2.send = vi.fn((_: unknown, ...args: unknown[]) => {
const cb = args.find((arg) => typeof arg === "function");
(cb as ((err: Error) => void) | undefined)?.(new Error("send failure"));
}) as typeof ws2.send;

connect(controller, ws1);
connect(controller, ws2);

await vi.waitFor(() => {
expect(ws2.nbOfTerminations).toBe(1);
});
await expectToWaitAlone(ws1);
});

it("aggregates without waiting for a crashed participant", async () => {
const controller = await makeController(2);
const ws1 = makeFederatedFakeWebSocket();
const ws2 = makeFederatedFakeWebSocket();
const ws3 = makeFederatedFakeWebSocket();

connect(controller, ws1);
connect(controller, ws2);
connect(controller, ws3);

// ws3 crashes, then the two others contribute
await contribute(ws3, DUMMY_WEIGHTS_3, 10);
await expectCrashed(ws3);
await contribute(ws1, DUMMY_WEIGHTS_1, 0);
await contribute(ws2, DUMMY_WEIGHTS_2, 0);

for (const ws of [ws1, ws2])
await vi.waitFor(() => {
checkReceiveServerPayload(ws, {
weights: MEAN_WEIGHTS_12,
round: 1,
nbOfParticipants: 2,
});
});
});
});
Loading
Loading