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: 3 additions & 1 deletion discojs/src/client/event_connection.spec.ts
Original file line number Diff line number Diff line change
Expand Up @@ -117,7 +117,9 @@ describe("sendAndWaitWithRetry", () => {
MType.NewFederatedNodeInfo,
{ retryDelayMs: RETRY_DELAY_MS, maxAttempts: 3 },
);
const rejects = expect(received).rejects.toThrow();
const rejects = expect(received).rejects.toThrow(
"no NewFederatedNodeInfo received after sending 3 ClientConnected",
);

await vi.advanceTimersByTimeAsync(3 * RETRY_DELAY_MS);

Expand Down
4 changes: 2 additions & 2 deletions discojs/src/client/event_connection.ts
Original file line number Diff line number Diff line change
Expand Up @@ -113,7 +113,7 @@ export async function sendAndWaitWithRetry<T extends MType>(
if (received !== RETRY) return received; // Return the received message if it's not a retry signal

debug(
"no %o after %dms, re-sending %o (%d/%d)",
"no %s after %dms, re-sending %s (%d/%d)",
responseType,
retryDelayMs,
request.type,
Expand All @@ -126,7 +126,7 @@ export async function sendAndWaitWithRetry<T extends MType>(
}

throw new Error(
`no ${responseType} received after ${maxAttempts} ${request.type}`,
`no ${responseType} received after sending ${maxAttempts} ${request.type}`,
);
}

Expand Down
94 changes: 49 additions & 45 deletions discojs/src/client/mtype.ts
Original file line number Diff line number Diff line change
@@ -1,77 +1,81 @@
/**
* Type of every message exchanged between clients and the server.
* The values are what goes over the wire
*
* We use a string enum rather than the default implicit integer enum
* because reordering an integer enum changes the messages' implicit value
* which can break compatibility between different builds
*/
export enum MType {
/* Both schemes */
// Sent from client to server as first point of contact to join a task.
// The server answers with an node id in a NewFederatedNodeInfo
// The server answers with a node id in a NewFederatedNodeInfo
// or NewDecentralizedNodeInfo message
ClientConnected,
ClientConnected = "ClientConnected",
// Sent by the server to notify clients that there are not enough
// participants to continue training
WaitingForMoreParticipants = "WaitingForMoreParticipants",
// Sent by the server to notify clients that there are now enough
// participants to start training collaboratively
EnoughParticipants = "EnoughParticipants",
// Sent by the server when a participant joined or left, so that
// clients don't have to wait for the end of the round to learn about it
ParticipantsUpdate = "ParticipantsUpdate",

/* Decentralized */
// When a user joins a task with a ClientConnected message, the server
// answers with its peer id and also tells the client whether we are waiting
// answers with its peer id and also tells the client whether we are waiting
// for more participants before starting training
NewDecentralizedNodeInfo,
// Message sent by peers to the server to signal they want to
// join the next round
JoinRound,
// Message sent by nodes to server signaling they are ready to
// start the next round
PeerIsReady,
NewDecentralizedNodeInfo = "NewDecentralizedNodeInfo",
// Sent by peers to the server to signal they want to join the next round
JoinRound = "JoinRound",
// Sent by nodes to the server signaling they are ready to start the next round
PeerIsReady = "PeerIsReady",
// Sent by the server to participating peers containing the list
// of peers for the round
PeersForRound,
// Message forwarded by the server from a client to another client
PeersForRound = "PeersForRound",
// Forwarded by the server from a client to another client
// to establish a peer-to-peer (WebRTC) connection
SignalForPeer,
// Message sent by nodes to server to signal all connections are established
ConnectionsReady,
SignalForPeer = "SignalForPeer",
// Sent by nodes to the server to signal all connections are established
ConnectionsReady = "ConnectionsReady",
// Sent by the server to signal nodes proceed to weight update sharing
StartWeightSharing,
StartWeightSharing = "StartWeightSharing",
// Sent by the server to signal nodes reestablish connections
RetryPeerConnections,
RetryPeerConnections = "RetryPeerConnections",
// Sent by the server to signal that the node's connection was not successful
ConnectionFail,
// The weight update
Payload,
ConnectionFail = "ConnectionFail",
// The weight update, sent from peer to peer
Payload = "Payload",
// Sent by nodes to the server to request the latest model
ModelSyncRequest,
ModelSyncRequest = "ModelSyncRequest",
// Sent by the server to nodes to share the provider node info
ModelProviderInfo,
ModelProviderInfo = "ModelProviderInfo",
// Sent by the server to the node who was selected as a model provider node
ProvideModelToPeer,
ProvideModelToPeer = "ProvideModelToPeer",
// Sent by node to node to share the latest model weights
SharedModel,
SharedModel = "SharedModel",

/* Federated */
// The server answers the ClientConnected message with the necessary information
// to start training: node id, latest model global weights, current round etc
NewFederatedNodeInfo,
// Message sent by server to notify clients that there are not enough
// participants to continue training
WaitingForMoreParticipants,
// Message sent by server to notify clients that there are now enough
// participants to start training collaboratively
EnoughParticipants,
SendPayload,
ReceiveServerPayload,

/* Both schemes */
// Message sent by the server when a participant joined or left, so that
// clients don't have to wait for the end of the round to learn about it.
// Kept last as the enum values are what goes over the wire.
ParticipantsUpdate,
CrashClient,
NewFederatedNodeInfo = "NewFederatedNodeInfo",
// Sent by clients to the server with their weight update for the round
SendPayload = "SendPayload",
// Sent by the server to clients with the aggregated weights of the round
ReceiveServerPayload = "ReceiveServerPayload",
CrashClient = "CrashClient",
}

const MTYPES: ReadonlySet<unknown> = new Set(Object.values(MType));

export function hasMessageType(
raw: unknown,
): raw is { type: MType } & Record<string, unknown> {
if (typeof raw !== "object" || raw === null) return false;

const o = raw as Record<string, unknown>;
if (!("type" in o && typeof o.type === "number" && o.type in MType)) {
return false;
}

return true;
return "type" in o && typeof o.type === "string" && MTYPES.has(o.type);
}

export interface ClientConnected {
Expand Down
6 changes: 4 additions & 2 deletions webapp/cypress.config.ts
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,12 @@ import * as msgpack from "@msgpack/msgpack";
import { loadEnv } from "vite";
import { WebSocketServer } from "ws";

import type { mtype } from "@epfml/discojs";

/**
* Messages the server answers with, depending on the type of the message it received.
*/
type ServerScript = Record<number, unknown[]>;
type ServerScript = Partial<Record<mtype.MType, unknown[]>>;

/**
* Serve the WebSockets of the server.
Expand All @@ -24,7 +26,7 @@ function serveWebSockets(port: number): (script: ServerScript) => void {
const handle = http.createServer((_, res) => res.writeHead(404).end());
new WebSocketServer({ server: handle }).on("connection", (ws) =>
ws.on("message", (data: Buffer) => {
const { type } = msgpack.decode(data) as { type: number };
const { type } = msgpack.decode(data) as { type: mtype.MType };
for (const answer of script[type] ?? []) ws.send(msgpack.encode(answer));
}),
);
Expand Down
Loading