Compare commits
18 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| ac524f0678 | |||
| bd7535b4e7 | |||
| 3ab64ea935 | |||
| 9846e74d74 | |||
| 61b7f5cd2a | |||
| 2f318cad41 | |||
| 7b479da51d | |||
| ec1b704003 | |||
| 3750bca5c2 | |||
| 0fab371ee4 | |||
| a215359dba | |||
| c3f18fd38a | |||
| c41cf7e9cf | |||
| e614d29bc5 | |||
| 8997363e72 | |||
| 01efca816b | |||
| 735182e3e2 | |||
| 6a02d6aa53 |
@@ -52,11 +52,11 @@ jobs:
|
||||
echo "BUILD_ID=${{ matrix.package }}-${sha}-${ts}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v2
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
# ..to avoid rate limits when pulling images
|
||||
- name: Login to DockerHub
|
||||
uses: docker/login-action@v2
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
username: ${{ secrets.DOCKERHUB_USERNAME }}
|
||||
password: ${{ secrets.DOCKERHUB_TOKEN }}
|
||||
@@ -67,7 +67,7 @@ jobs:
|
||||
|
||||
# ..to push image
|
||||
- name: 🐙 Login to GitHub Container Registry
|
||||
uses: docker/login-action@v2
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.repository_owner }}
|
||||
|
||||
@@ -54,10 +54,10 @@ RUN corepack enable
|
||||
ENV NODE_ENV production
|
||||
|
||||
COPY --from=cri-tools --chown=node:node /cri-tools/crictl /usr/local/bin
|
||||
COPY --from=builder --chown=node:node /app/apps/coordinator/dist/index.cjs ./index.cjs
|
||||
COPY --from=builder --chown=node:node /app/apps/coordinator/dist/index.mjs ./index.mjs
|
||||
|
||||
EXPOSE 8000
|
||||
|
||||
USER node
|
||||
|
||||
CMD [ "/usr/bin/dumb-init", "--", "/usr/local/bin/node", "./index.cjs" ]
|
||||
CMD [ "/usr/bin/dumb-init", "--", "/usr/local/bin/node", "./index.mjs" ]
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
"build": "npm run build:bundle",
|
||||
"build:bundle": "esbuild src/index.ts --bundle --outfile=dist/index.cjs --platform=node",
|
||||
"build:bundle": "esbuild src/index.ts --bundle --outfile=dist/index.mjs --platform=node --format=esm --target=esnext --banner:js=\"const require = createRequire(import.meta.url);\"",
|
||||
"build:image": "docker build -f Containerfile . -t coordinator",
|
||||
"dev": "tsx --no-warnings=ExperimentalWarning --require dotenv/config --watch src/index.ts",
|
||||
"start": "tsx src/index.ts",
|
||||
|
||||
+183
-212
@@ -1,16 +1,15 @@
|
||||
import { randomUUID } from "node:crypto";
|
||||
import { createServer } from "node:http";
|
||||
import { $ } from "execa";
|
||||
import { Namespace } from "socket.io";
|
||||
import { Server } from "socket.io";
|
||||
import { Socket, io } from "socket.io-client";
|
||||
import { DefaultEventsMap } from "socket.io/dist/typed-events";
|
||||
import {
|
||||
CoordinatorToPlatformEvents,
|
||||
CoordinatorToProdWorkerEvents,
|
||||
PlatformToCoordinatorEvents,
|
||||
CoordinatorToPlatformMessages,
|
||||
CoordinatorToProdWorkerMessages,
|
||||
PlatformToCoordinatorMessages,
|
||||
ProdWorkerSocketData,
|
||||
ProdWorkerToCoordinatorEvents,
|
||||
ProdWorkerToCoordinatorMessages,
|
||||
ZodNamespace,
|
||||
ZodSocketConnection,
|
||||
} from "@trigger.dev/core/v3";
|
||||
import { HttpReply, getTextBody, SimpleLogger } from "@trigger.dev/core-apps";
|
||||
|
||||
@@ -194,13 +193,15 @@ class TaskCoordinator {
|
||||
#httpServer: ReturnType<typeof createServer>;
|
||||
#checkpointer = new Checkpointer();
|
||||
|
||||
#prodWorkerNamespace: Namespace<
|
||||
ProdWorkerToCoordinatorEvents,
|
||||
CoordinatorToProdWorkerEvents,
|
||||
DefaultEventsMap,
|
||||
ProdWorkerSocketData
|
||||
#prodWorkerNamespace: ZodNamespace<
|
||||
typeof ProdWorkerToCoordinatorMessages,
|
||||
typeof CoordinatorToProdWorkerMessages,
|
||||
typeof ProdWorkerSocketData
|
||||
>;
|
||||
#platformSocket?: ZodSocketConnection<
|
||||
typeof CoordinatorToPlatformMessages,
|
||||
typeof PlatformToCoordinatorMessages
|
||||
>;
|
||||
#platformSocket?: Socket<PlatformToCoordinatorEvents, CoordinatorToPlatformEvents>;
|
||||
|
||||
constructor(
|
||||
private port: number,
|
||||
@@ -218,7 +219,7 @@ class TaskCoordinator {
|
||||
name: "daemon_connected_tasks_total", // don't change this without updating dashboard config
|
||||
help: "The number of tasks currently connected.",
|
||||
collect: () => {
|
||||
connectedTasksTotal.set(this.#prodWorkerNamespace.sockets.size);
|
||||
connectedTasksTotal.set(this.#prodWorkerNamespace.namespace.sockets.size);
|
||||
},
|
||||
});
|
||||
register.registerMetric(connectedTasksTotal);
|
||||
@@ -230,44 +231,28 @@ class TaskCoordinator {
|
||||
return;
|
||||
}
|
||||
|
||||
const socket: Socket<PlatformToCoordinatorEvents, CoordinatorToPlatformEvents> = io(
|
||||
`ws://${PLATFORM_HOST}:${PLATFORM_WS_PORT}/coordinator`,
|
||||
{
|
||||
transports: ["websocket"],
|
||||
auth: {
|
||||
token: PLATFORM_SECRET,
|
||||
const platformConnection = new ZodSocketConnection({
|
||||
namespace: "coordinator",
|
||||
host: PLATFORM_HOST,
|
||||
port: Number(PLATFORM_WS_PORT),
|
||||
clientMessages: CoordinatorToPlatformMessages,
|
||||
serverMessages: PlatformToCoordinatorMessages,
|
||||
authToken: PLATFORM_SECRET,
|
||||
handlers: {
|
||||
RESUME: async (message) => {
|
||||
const taskSocket = await this.#getAttemptSocket(message.attemptId);
|
||||
|
||||
if (!taskSocket) {
|
||||
logger.log("Socket for attempt not found", { attemptId: message.attemptId });
|
||||
return;
|
||||
}
|
||||
|
||||
taskSocket.emit("RESUME", message);
|
||||
},
|
||||
}
|
||||
);
|
||||
|
||||
const logger = new SimpleLogger(`[platform][${socket.id ?? "NO_ID"}]`);
|
||||
|
||||
socket.on("connect", () => {
|
||||
logger.log("connect");
|
||||
},
|
||||
});
|
||||
|
||||
socket.on("connect_error", (err) => {
|
||||
logger.error(`connect_error: ${err.message}`);
|
||||
});
|
||||
|
||||
socket.on("disconnect", () => {
|
||||
logger.log("disconnect");
|
||||
});
|
||||
|
||||
socket.on("RESUME", async (message) => {
|
||||
logger.log("[RESUME]", message);
|
||||
|
||||
const taskSocket = await this.#getAttemptSocket(message.attemptId);
|
||||
|
||||
if (!taskSocket) {
|
||||
logger.log("Socket for attempt not found", { attemptId: message.attemptId });
|
||||
return;
|
||||
}
|
||||
|
||||
taskSocket.emit("RESUME", message);
|
||||
});
|
||||
|
||||
return socket;
|
||||
return platformConnection;
|
||||
}
|
||||
|
||||
async #getAttemptSocket(attemptId: string) {
|
||||
@@ -281,182 +266,168 @@ class TaskCoordinator {
|
||||
}
|
||||
|
||||
#createProdWorkerNamespace(io: Server) {
|
||||
const namespace: Namespace<
|
||||
ProdWorkerToCoordinatorEvents,
|
||||
CoordinatorToProdWorkerEvents,
|
||||
DefaultEventsMap,
|
||||
ProdWorkerSocketData
|
||||
> = io.of("/prod-worker");
|
||||
const provider = new ZodNamespace({
|
||||
io,
|
||||
name: "prod-worker",
|
||||
clientMessages: ProdWorkerToCoordinatorMessages,
|
||||
serverMessages: CoordinatorToProdWorkerMessages,
|
||||
socketData: ProdWorkerSocketData,
|
||||
postAuth: async (socket, next, logger) => {
|
||||
function setSocketDataFromHeader(dataKey: keyof typeof socket.data, headerName: string) {
|
||||
const value = socket.handshake.headers[headerName];
|
||||
if (!value) {
|
||||
logger(`missing required header: ${headerName}`);
|
||||
throw new Error("missing header");
|
||||
}
|
||||
0;
|
||||
socket.data[dataKey] = Array.isArray(value) ? value[0] : value;
|
||||
}
|
||||
|
||||
namespace.on("connection", async (socket) => {
|
||||
const logger = new SimpleLogger(`[task][${socket.id}]`);
|
||||
try {
|
||||
setSocketDataFromHeader("podName", "x-pod-name");
|
||||
setSocketDataFromHeader("contentHash", "x-trigger-content-hash");
|
||||
setSocketDataFromHeader("cliPackageVersion", "x-trigger-cli-package-version");
|
||||
setSocketDataFromHeader("projectRef", "x-trigger-project-ref");
|
||||
setSocketDataFromHeader("attemptId", "x-trigger-attempt-id");
|
||||
setSocketDataFromHeader("envId", "x-trigger-env-id");
|
||||
} catch (error) {
|
||||
logger(error);
|
||||
socket.disconnect(true);
|
||||
return;
|
||||
}
|
||||
|
||||
this.#platformSocket?.emit("LOG", {
|
||||
version: "v1",
|
||||
metadata: {
|
||||
projectRef: socket.data.projectRef,
|
||||
attemptId: socket.data.attemptId,
|
||||
},
|
||||
text: "connected",
|
||||
});
|
||||
logger("success", socket.data);
|
||||
|
||||
logger.log("connected");
|
||||
next();
|
||||
},
|
||||
onConnection: async (socket, handler, sender) => {
|
||||
const logger = new SimpleLogger(`[task][${socket.id}]`);
|
||||
|
||||
socket.on("disconnect", (reason, description) => {
|
||||
logger.log("disconnect", { reason, description });
|
||||
this.#platformSocket?.send("LOG", {
|
||||
metadata: {
|
||||
projectRef: socket.data.projectRef,
|
||||
attemptId: socket.data.attemptId,
|
||||
},
|
||||
text: "connected",
|
||||
});
|
||||
|
||||
this.#platformSocket?.emit("LOG", {
|
||||
version: "v1",
|
||||
socket.on("LOG", (message, callback) => {
|
||||
logger.log("[LOG]", message.text);
|
||||
|
||||
callback();
|
||||
|
||||
this.#platformSocket?.send("LOG", {
|
||||
version: "v1",
|
||||
metadata: { attemptId: socket.data.attemptId },
|
||||
text: message.text,
|
||||
});
|
||||
});
|
||||
|
||||
socket.on("READY_FOR_EXECUTION", async (message) => {
|
||||
logger.log("[READY_FOR_EXECUTION]", message);
|
||||
|
||||
const executionAck = await this.#platformSocket?.sendWithAck("READY_FOR_EXECUTION", {
|
||||
version: "v1",
|
||||
attemptId: message.attemptId,
|
||||
});
|
||||
|
||||
if (!executionAck) {
|
||||
logger.error("no execution ack", { attemptId: socket.data.attemptId });
|
||||
return;
|
||||
}
|
||||
|
||||
if (!executionAck.success) {
|
||||
logger.error("execution unsuccessful", { attemptId: socket.data.attemptId });
|
||||
return;
|
||||
}
|
||||
|
||||
// FIXME: shouldn't wait for completion here
|
||||
const completionAck = await socket.emitWithAck("EXECUTE_TASK_RUN", {
|
||||
version: "v1",
|
||||
executionPayload: executionAck.payload,
|
||||
});
|
||||
|
||||
if (!completionAck.success) {
|
||||
logger.error("completion unsuccessful", { attemptId: socket.data.attemptId });
|
||||
return;
|
||||
}
|
||||
|
||||
logger.log("completed task", { completionId: completionAck.completion.id });
|
||||
|
||||
this.#platformSocket?.send("TASK_RUN_COMPLETED", {
|
||||
version: "v1",
|
||||
execution: executionAck.payload.execution,
|
||||
completion: completionAck.completion,
|
||||
});
|
||||
});
|
||||
|
||||
socket.on("WAIT_FOR_DURATION", async (message, callback) => {
|
||||
logger.log("[WAIT_FOR_DURATION]", message);
|
||||
|
||||
const checkpoint = await this.#checkpointer.checkpointAndPush(socket.data.podName);
|
||||
|
||||
if (!checkpoint) {
|
||||
logger.error("Failed to checkpoint", { podName: socket.data.podName });
|
||||
callback({ success: false });
|
||||
return;
|
||||
}
|
||||
|
||||
this.#platformSocket?.send("CHECKPOINT_CREATED", {
|
||||
version: "v1",
|
||||
attemptId: socket.data.attemptId,
|
||||
docker: checkpoint.docker,
|
||||
location: checkpoint.destination,
|
||||
reason: "WAIT_FOR_DURATION",
|
||||
});
|
||||
|
||||
callback({ success: true });
|
||||
});
|
||||
|
||||
socket.on("INDEX_TASKS", async (message, callback) => {
|
||||
logger.log("[INDEX_TASKS]", message);
|
||||
|
||||
const workerAck = await this.#platformSocket?.sendWithAck("CREATE_WORKER", {
|
||||
version: "v1",
|
||||
projectRef: socket.data.projectRef,
|
||||
envId: socket.data.envId,
|
||||
metadata: {
|
||||
cliPackageVersion: socket.data.cliPackageVersion,
|
||||
contentHash: socket.data.contentHash,
|
||||
packageVersion: message.packageVersion,
|
||||
tasks: message.tasks,
|
||||
},
|
||||
});
|
||||
|
||||
if (!workerAck) {
|
||||
logger.debug("no worker ack while indexing", message);
|
||||
}
|
||||
|
||||
callback({ success: !!workerAck?.success });
|
||||
});
|
||||
},
|
||||
onDisconnect: async (socket, handler, sender, logger) => {
|
||||
this.#platformSocket?.send("LOG", {
|
||||
metadata: {
|
||||
projectRef: socket.data.projectRef,
|
||||
attemptId: socket.data.attemptId,
|
||||
},
|
||||
text: "disconnect",
|
||||
});
|
||||
});
|
||||
|
||||
socket.on("error", (error) => {
|
||||
logger.error({ error });
|
||||
});
|
||||
|
||||
socket.on("LOG", (message, callback) => {
|
||||
logger.log("[LOG]", message.text);
|
||||
|
||||
callback();
|
||||
|
||||
this.#platformSocket?.emit("LOG", {
|
||||
version: "v1",
|
||||
metadata: { attemptId: socket.data.attemptId },
|
||||
text: message.text,
|
||||
});
|
||||
});
|
||||
|
||||
socket.on("READY_FOR_EXECUTION", async (message) => {
|
||||
logger.log("[READY_FOR_EXECUTION]", message);
|
||||
|
||||
const executionAck = await this.#platformSocket?.emitWithAck("READY_FOR_EXECUTION", {
|
||||
version: "v1",
|
||||
attemptId: message.attemptId,
|
||||
});
|
||||
|
||||
if (!executionAck) {
|
||||
logger.error("no execution ack", { attemptId: socket.data.attemptId });
|
||||
return;
|
||||
}
|
||||
|
||||
if (!executionAck.success) {
|
||||
logger.error("execution unsuccessful", { attemptId: socket.data.attemptId });
|
||||
return;
|
||||
}
|
||||
|
||||
// FIXME: shouldn't wait for completion here
|
||||
const completionAck = await socket.emitWithAck("EXECUTE_TASK_RUN", {
|
||||
version: "v1",
|
||||
payload: executionAck.payload,
|
||||
});
|
||||
|
||||
logger.log("completed task", { completionId: completionAck.completion.id });
|
||||
|
||||
this.#platformSocket?.emit("TASK_RUN_COMPLETED", {
|
||||
version: "v1",
|
||||
execution: executionAck.payload.execution,
|
||||
completion: completionAck.completion,
|
||||
});
|
||||
});
|
||||
|
||||
socket.on("TASK_HEARTBEAT", (message) => {
|
||||
logger.log("[TASK_HEARTBEAT]", message);
|
||||
|
||||
this.#platformSocket?.emit("TASK_HEARTBEAT", message);
|
||||
});
|
||||
|
||||
socket.on("WAIT_FOR_BATCH", (message) => {
|
||||
logger.log("[WAIT_FOR_BATCH]", message);
|
||||
|
||||
// this.#checkpointer.checkpointAndPush(socket.data.podName);
|
||||
});
|
||||
|
||||
socket.on("WAIT_FOR_DURATION", async (message, callback) => {
|
||||
logger.log("[WAIT_FOR_DURATION]", message);
|
||||
|
||||
const checkpoint = await this.#checkpointer.checkpointAndPush(socket.data.podName);
|
||||
|
||||
if (!checkpoint) {
|
||||
logger.error("Failed to checkpoint", { podName: socket.data.podName });
|
||||
callback({ success: false });
|
||||
return;
|
||||
}
|
||||
|
||||
this.#platformSocket?.emit("CHECKPOINT_CREATED", {
|
||||
version: "v1",
|
||||
attemptId: socket.data.attemptId,
|
||||
docker: checkpoint.docker,
|
||||
location: checkpoint.destination,
|
||||
reason: "WAIT_FOR_DURATION",
|
||||
});
|
||||
|
||||
callback({ success: true });
|
||||
});
|
||||
|
||||
socket.on("WAIT_FOR_TASK", (message) => {
|
||||
logger.log("[WAIT_FOR_TASK]", message);
|
||||
|
||||
// this.#checkpointer.checkpointAndPush(socket.data.podName);
|
||||
});
|
||||
|
||||
socket.on("INDEX_TASKS", async (message, callback) => {
|
||||
logger.log("[INDEX_TASKS]", message);
|
||||
|
||||
const workerAck = await this.#platformSocket?.emitWithAck("CREATE_WORKER", {
|
||||
version: "v1",
|
||||
projectRef: socket.data.projectRef,
|
||||
envId: socket.data.envId,
|
||||
metadata: {
|
||||
cliPackageVersion: socket.data.cliPackageVersion,
|
||||
contentHash: socket.data.contentHash,
|
||||
packageVersion: message.packageVersion,
|
||||
tasks: message.tasks,
|
||||
},
|
||||
});
|
||||
|
||||
if (!workerAck) {
|
||||
logger.debug("no worker ack while indexing", message);
|
||||
}
|
||||
|
||||
callback({ success: !!workerAck?.success });
|
||||
});
|
||||
},
|
||||
handlers: {
|
||||
TASK_HEARTBEAT: async (message) => {
|
||||
this.#platformSocket?.send("TASK_HEARTBEAT", message);
|
||||
},
|
||||
WAIT_FOR_BATCH: async (message) => {
|
||||
// this.#checkpointer.checkpointAndPush(socket.data.podName);
|
||||
},
|
||||
WAIT_FOR_TASK: async (message) => {
|
||||
// this.#checkpointer.checkpointAndPush(socket.data.podName);
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
// auth middleware
|
||||
namespace.use(async (socket, next) => {
|
||||
const logger = new SimpleLogger(`[task][${socket.id}][auth]`);
|
||||
|
||||
function setSocketDataFromHeader(dataKey: keyof typeof socket.data, headerName: string) {
|
||||
const value = socket.handshake.headers[headerName];
|
||||
if (!value) {
|
||||
logger.error(`missing required header: ${headerName}`);
|
||||
throw new Error("missing header");
|
||||
}
|
||||
socket.data[dataKey] = Array.isArray(value) ? value[0] : value;
|
||||
}
|
||||
|
||||
try {
|
||||
setSocketDataFromHeader("podName", "x-pod-name");
|
||||
setSocketDataFromHeader("contentHash", "x-trigger-content-hash");
|
||||
setSocketDataFromHeader("cliPackageVersion", "x-trigger-cli-package-version");
|
||||
setSocketDataFromHeader("projectRef", "x-trigger-project-ref");
|
||||
setSocketDataFromHeader("attemptId", "x-trigger-attempt-id");
|
||||
setSocketDataFromHeader("envId", "x-trigger-env-id");
|
||||
} catch (error) {
|
||||
return socket.disconnect(true);
|
||||
}
|
||||
|
||||
logger.log("success", socket.data);
|
||||
|
||||
next();
|
||||
});
|
||||
|
||||
return namespace;
|
||||
return provider;
|
||||
}
|
||||
|
||||
#createHttpServer() {
|
||||
|
||||
@@ -9,8 +9,8 @@ FROM base
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY --chown=node dist/index.cjs /app/
|
||||
COPY --chown=node dist/index.mjs /app/
|
||||
|
||||
EXPOSE 8000
|
||||
|
||||
ENTRYPOINT [ "/usr/bin/dumb-init", "--", "/usr/local/bin/node", "/app/index.cjs" ]
|
||||
ENTRYPOINT [ "/usr/bin/dumb-init", "--", "/usr/local/bin/node", "/app/index.mjs" ]
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
"build": "npm run build:bundle",
|
||||
"build:bundle": "esbuild src/index.ts --bundle --outfile=dist/index.cjs --platform=node",
|
||||
"build:bundle": "esbuild src/index.ts --bundle --outfile=dist/index.mjs --platform=node --format=esm --target=esnext --banner:js=\"const require = createRequire(import.meta.url);\"",
|
||||
"build:image": "docker build -f Containerfile . -t docker-provider",
|
||||
"dev": "tsx --no-warnings=ExperimentalWarning --require dotenv/config --watch src/index.ts",
|
||||
"start": "tsx src/index.ts",
|
||||
|
||||
@@ -1,37 +1,12 @@
|
||||
import { createServer } from "node:http";
|
||||
import { $ } from "execa";
|
||||
import { io, Socket } from "socket.io-client";
|
||||
import {
|
||||
clientWebsocketMessages,
|
||||
Machine,
|
||||
MessageCatalogToSocketIoEvents,
|
||||
ProviderClientToServerEvents,
|
||||
ProviderServerToClientEvents,
|
||||
serverWebsocketMessages,
|
||||
ZodMessageHandler,
|
||||
ZodMessageSender,
|
||||
} from "@trigger.dev/core/v3";
|
||||
import { HttpReply, SimpleLogger, getTextBody, getRandomPortNumber } from "@trigger.dev/core-apps";
|
||||
import { Machine } from "@trigger.dev/core/v3";
|
||||
import { SimpleLogger, TaskOperations, ProviderShell } from "@trigger.dev/core-apps";
|
||||
|
||||
const HTTP_SERVER_PORT = Number(process.env.HTTP_SERVER_PORT || getRandomPortNumber());
|
||||
const MACHINE_NAME = process.env.MACHINE_NAME || "local";
|
||||
|
||||
const PLATFORM_HOST = process.env.PLATFORM_HOST || "127.0.0.1";
|
||||
const PLATFORM_WS_PORT = process.env.PLATFORM_WS_PORT || 3030;
|
||||
const PLATFORM_SECRET = process.env.PLATFORM_SECRET || "provider-secret";
|
||||
|
||||
const COORDINATOR_PORT = process.env.COORDINATOR_PORT || 8020;
|
||||
|
||||
const logger = new SimpleLogger(`[${MACHINE_NAME}]`);
|
||||
|
||||
interface TaskOperations {
|
||||
create: (...args: any[]) => Promise<any>;
|
||||
restore: (...args: any[]) => Promise<any>;
|
||||
delete: (...args: any[]) => Promise<any>;
|
||||
get: (...args: any[]) => Promise<any>;
|
||||
index: (...args: any[]) => Promise<any>;
|
||||
}
|
||||
|
||||
class DockerTaskOperations implements TaskOperations {
|
||||
async index(opts: { contentHash: string; imageTag: string; envId: string }) {
|
||||
const containerName = this.#getIndexContainerName(opts.contentHash);
|
||||
@@ -93,242 +68,7 @@ class DockerTaskOperations implements TaskOperations {
|
||||
}
|
||||
}
|
||||
|
||||
interface Provider {
|
||||
tasks: TaskOperations;
|
||||
}
|
||||
|
||||
type DockerProviderOptions = {
|
||||
tasks: DockerTaskOperations;
|
||||
host?: string;
|
||||
port: number;
|
||||
};
|
||||
|
||||
class DockerProvider implements Provider {
|
||||
tasks: DockerTaskOperations;
|
||||
|
||||
#httpServer: ReturnType<typeof createServer>;
|
||||
#platformSocket: Socket<ProviderServerToClientEvents, ProviderClientToServerEvents>;
|
||||
|
||||
constructor(private options: DockerProviderOptions) {
|
||||
this.tasks = options.tasks;
|
||||
this.#httpServer = this.#createHttpServer();
|
||||
this.#platformSocket = this.#createPlatformSocket();
|
||||
this.#createSharedQueueSocket();
|
||||
}
|
||||
|
||||
#createSharedQueueSocket() {
|
||||
const socket: Socket<
|
||||
MessageCatalogToSocketIoEvents<typeof serverWebsocketMessages>,
|
||||
MessageCatalogToSocketIoEvents<typeof clientWebsocketMessages>
|
||||
> = io(`ws://${PLATFORM_HOST}:${PLATFORM_WS_PORT}/shared-queue`, {
|
||||
transports: ["websocket"],
|
||||
auth: {
|
||||
token: PLATFORM_SECRET,
|
||||
},
|
||||
});
|
||||
|
||||
const logger = new SimpleLogger(`[shared-queue][${socket.id ?? "NO_ID"}]`);
|
||||
|
||||
socket.on("connect_error", (err) => {
|
||||
logger.error(`connect_error: ${err.message}`);
|
||||
});
|
||||
|
||||
socket.on("connect", () => {
|
||||
logger.log("connect");
|
||||
});
|
||||
|
||||
socket.on("disconnect", () => {
|
||||
logger.log("disconnect");
|
||||
});
|
||||
|
||||
const sender = new ZodMessageSender({
|
||||
schema: clientWebsocketMessages,
|
||||
sender: async (message) => {
|
||||
return new Promise((resolve, reject) => {
|
||||
try {
|
||||
const { type, ...payload } = message;
|
||||
socket.emit(type, payload as any);
|
||||
resolve();
|
||||
} catch (err) {
|
||||
reject(err);
|
||||
}
|
||||
});
|
||||
},
|
||||
});
|
||||
|
||||
const handler = new ZodMessageHandler({
|
||||
schema: serverWebsocketMessages,
|
||||
messages: {
|
||||
SERVER_READY: async (payload) => {
|
||||
logger.log("received SERVER_READY", payload);
|
||||
|
||||
// TODO: create new schema without worker requirement
|
||||
await sender.send("READY_FOR_TASKS", {
|
||||
backgroundWorkerId: "placeholder",
|
||||
});
|
||||
},
|
||||
BACKGROUND_WORKER_MESSAGE: async (payload) => {
|
||||
logger.log("received BACKGROUND_WORKER_MESSAGE", payload);
|
||||
|
||||
if (payload.data.type === "SCHEDULE_ATTEMPT") {
|
||||
this.tasks.create({
|
||||
envId: payload.data.envId,
|
||||
attemptId: payload.data.id,
|
||||
image: payload.data.image,
|
||||
machine: {},
|
||||
});
|
||||
}
|
||||
},
|
||||
},
|
||||
});
|
||||
handler.registerHandlers(socket, logger.log.bind(logger));
|
||||
|
||||
return socket;
|
||||
}
|
||||
|
||||
#createPlatformSocket() {
|
||||
const socket: Socket<ProviderServerToClientEvents, ProviderClientToServerEvents> = io(
|
||||
`ws://${PLATFORM_HOST}:${PLATFORM_WS_PORT}/provider`,
|
||||
{
|
||||
transports: ["websocket"],
|
||||
auth: {
|
||||
token: PLATFORM_SECRET,
|
||||
},
|
||||
extraHeaders: {
|
||||
"x-trigger-provider-type": "docker",
|
||||
},
|
||||
}
|
||||
);
|
||||
|
||||
const logger = new SimpleLogger(`[platform][${socket.id ?? "NO_ID"}]`);
|
||||
|
||||
socket.on("connect_error", (err) => {
|
||||
logger.error(`connect_error: ${err.message}`);
|
||||
});
|
||||
|
||||
socket.on("connect", () => {
|
||||
logger.log("connect");
|
||||
});
|
||||
|
||||
socket.on("disconnect", () => {
|
||||
logger.log("disconnect");
|
||||
});
|
||||
|
||||
socket.on("GET", async (message) => {
|
||||
logger.log("[GET]", message);
|
||||
|
||||
this.tasks.get({ runId: message.name });
|
||||
});
|
||||
|
||||
socket.on("DELETE", async (message, callback) => {
|
||||
logger.log("[DELETE]", message);
|
||||
|
||||
callback({
|
||||
message: "delete request received",
|
||||
});
|
||||
|
||||
this.tasks.delete({ runId: message.name });
|
||||
});
|
||||
|
||||
socket.on("INDEX", async (message) => {
|
||||
logger.log("[INDEX]", message);
|
||||
try {
|
||||
await this.tasks.index({
|
||||
contentHash: message.contentHash,
|
||||
imageTag: message.imageTag,
|
||||
envId: message.envId,
|
||||
});
|
||||
} catch (error) {
|
||||
logger.error("task index failed", error);
|
||||
}
|
||||
});
|
||||
|
||||
socket.on("RESTORE", async (message) => {
|
||||
logger.log("[RESTORE]", message);
|
||||
// await this.tasks.restore({});
|
||||
});
|
||||
|
||||
socket.on("HEALTH", async (message) => {
|
||||
logger.log("[HEALTH]", message);
|
||||
});
|
||||
|
||||
return socket;
|
||||
}
|
||||
|
||||
#createHttpServer() {
|
||||
const httpServer = createServer(async (req, res) => {
|
||||
logger.log(`[${req.method}]`, req.url);
|
||||
|
||||
const reply = new HttpReply(res);
|
||||
|
||||
switch (req.url) {
|
||||
case "/health": {
|
||||
return reply.text("ok");
|
||||
}
|
||||
case "/whoami": {
|
||||
return reply.text(`${MACHINE_NAME}`);
|
||||
}
|
||||
case "/close": {
|
||||
this.#platformSocket.close();
|
||||
return reply.text("platform socket closed");
|
||||
}
|
||||
case "/delete": {
|
||||
const body = await getTextBody(req);
|
||||
|
||||
await this.tasks.delete({ runId: body });
|
||||
|
||||
return reply.text(`sent delete request: ${body}`);
|
||||
}
|
||||
case "/invoke": {
|
||||
const body = await getTextBody(req);
|
||||
|
||||
await this.tasks.create({
|
||||
attemptId: body,
|
||||
envId: "placeholder",
|
||||
image: body,
|
||||
machine: {
|
||||
cpu: "1",
|
||||
memory: "100Mi",
|
||||
},
|
||||
});
|
||||
|
||||
return reply.text(`sent restore request: ${body}`);
|
||||
}
|
||||
case "/restore": {
|
||||
const body = await getTextBody(req);
|
||||
|
||||
const items = body.split("&");
|
||||
const image = items[0];
|
||||
const baseImageTag = items[1] ?? image;
|
||||
|
||||
// await this.tasks.restore({});
|
||||
|
||||
return reply.text(`sent restore request: ${body}`);
|
||||
}
|
||||
default: {
|
||||
return reply.empty(404);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
httpServer.on("clientError", (err, socket) => {
|
||||
socket.end("HTTP/1.1 400 Bad Request\r\n\r\n");
|
||||
});
|
||||
|
||||
httpServer.on("listening", () => {
|
||||
logger.log("server listening on port", this.options.port);
|
||||
});
|
||||
|
||||
return httpServer;
|
||||
}
|
||||
|
||||
listen() {
|
||||
this.#httpServer.listen(this.options.port, this.options.host ?? "0.0.0.0");
|
||||
}
|
||||
}
|
||||
|
||||
const provider = new DockerProvider({
|
||||
port: HTTP_SERVER_PORT,
|
||||
const provider = new ProviderShell({
|
||||
tasks: new DockerTaskOperations(),
|
||||
});
|
||||
|
||||
|
||||
@@ -39,10 +39,10 @@ FROM base AS runner
|
||||
RUN corepack enable
|
||||
ENV NODE_ENV production
|
||||
|
||||
COPY --from=builder --chown=node:node /app/apps/kubernetes-provider/dist/index.cjs ./index.cjs
|
||||
COPY --from=builder --chown=node:node /app/apps/kubernetes-provider/dist/index.mjs ./index.mjs
|
||||
|
||||
EXPOSE 8000
|
||||
|
||||
USER node
|
||||
|
||||
CMD [ "/usr/bin/dumb-init", "--", "/usr/local/bin/node", "./index.cjs" ]
|
||||
CMD [ "/usr/bin/dumb-init", "--", "/usr/local/bin/node", "./index.mjs" ]
|
||||
|
||||
@@ -7,10 +7,8 @@
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
"build": "npm run build:bundle",
|
||||
"build:bundle": "esbuild src/index.ts --bundle --outfile=dist/index.cjs --platform=node",
|
||||
"build:image": "docker build -f Containerfile . -t kubernetes-provider",
|
||||
"dev": "tsx --no-warnings=ExperimentalWarning --require dotenv/config --watch src/index.ts",
|
||||
"start": "tsx src/index.ts"
|
||||
"build:bundle": "esbuild src/index.ts --bundle --outfile=dist/index.mjs --platform=node --format=esm --target=esnext --banner:js=\"const require = createRequire(import.meta.url);\"",
|
||||
"typecheck": "tsc --noEmit"
|
||||
},
|
||||
"keywords": [],
|
||||
"author": "",
|
||||
|
||||
@@ -1,23 +1,12 @@
|
||||
import { randomUUID } from "node:crypto";
|
||||
import { createServer } from "node:http";
|
||||
import k8s, { BatchV1Api, CoreV1Api, V1Job, V1Pod } from "@kubernetes/client-node";
|
||||
import { io, Socket } from "socket.io-client";
|
||||
import {
|
||||
Machine,
|
||||
ProviderClientToServerEvents,
|
||||
ProviderServerToClientEvents,
|
||||
} from "@trigger.dev/core/v3";
|
||||
import { HttpReply, SimpleLogger, getTextBody } from "@trigger.dev/core-apps";
|
||||
import { Machine } from "@trigger.dev/core/v3";
|
||||
import { ProviderShell, SimpleLogger, TaskOperations } from "@trigger.dev/core-apps";
|
||||
|
||||
const RUNTIME_ENV = process.env.KUBERNETES_PORT ? "kubernetes" : "local";
|
||||
|
||||
const HTTP_SERVER_PORT = Number(process.env.HTTP_SERVER_PORT || 8000);
|
||||
const NODE_NAME = process.env.NODE_NAME || "some-node";
|
||||
const POD_NAME = process.env.POD_NAME || "k8s-provider";
|
||||
|
||||
const PLATFORM_HOST = process.env.PLATFORM_HOST || "127.0.0.1";
|
||||
const PLATFORM_WS_PORT = process.env.PLATFORM_WS_PORT || 5080;
|
||||
const PLATFORM_SECRET = process.env.PLATFORM_SECRET || "provider-secret";
|
||||
const OTEL_EXPORTER_OTLP_ENDPOINT =
|
||||
process.env.OTEL_EXPORTER_OTLP_ENDPOINT ?? "http://0.0.0.0:4318";
|
||||
|
||||
const REGISTRY_FQDN = process.env.REGISTRY_FQDN || "localhost:5000";
|
||||
const REPO_NAME = process.env.REPO_NAME || "test";
|
||||
@@ -30,14 +19,6 @@ type Namespace = {
|
||||
};
|
||||
};
|
||||
|
||||
interface TaskOperations {
|
||||
create: (...args: any[]) => Promise<any>;
|
||||
restore: (...args: any[]) => Promise<any>;
|
||||
delete: (...args: any[]) => Promise<any>;
|
||||
get: (...args: any[]) => Promise<any>;
|
||||
index: (...args: any[]) => Promise<any>;
|
||||
}
|
||||
|
||||
class KubernetesTaskOperations implements TaskOperations {
|
||||
#namespace: Namespace;
|
||||
#k8sApi: {
|
||||
@@ -55,7 +36,7 @@ class KubernetesTaskOperations implements TaskOperations {
|
||||
this.#k8sApi = this.#createK8sApi();
|
||||
}
|
||||
|
||||
async index(opts: { contentHash: string; imageTag: string }) {
|
||||
async index(opts: { contentHash: string; imageTag: string; envId: string }) {
|
||||
await this.#createJob(
|
||||
{
|
||||
metadata: {
|
||||
@@ -101,6 +82,14 @@ class KubernetesTaskOperations implements TaskOperations {
|
||||
name: "INDEX_TASKS",
|
||||
value: "true",
|
||||
},
|
||||
{
|
||||
name: "TRIGGER_ENV_ID",
|
||||
value: opts.envId,
|
||||
},
|
||||
{
|
||||
name: "OTEL_EXPORTER_OTLP_ENDPOINT",
|
||||
value: OTEL_EXPORTER_OTLP_ENDPOINT,
|
||||
},
|
||||
{
|
||||
name: "HTTP_SERVER_PORT",
|
||||
value: "8000",
|
||||
@@ -140,19 +129,27 @@ class KubernetesTaskOperations implements TaskOperations {
|
||||
);
|
||||
}
|
||||
|
||||
async create(opts: { runId: string; image: string; machine: Machine }) {
|
||||
async create(opts: { attemptId: string; image: string; machine: Machine; envId: string }) {
|
||||
await this.#createPod(
|
||||
{
|
||||
metadata: {
|
||||
name: `${opts.runId}-${randomUUID().slice(0, 5)}`,
|
||||
name: `task-run-${opts.attemptId}-${randomUUID().slice(0, 5)}`,
|
||||
namespace: this.#namespace.metadata.name,
|
||||
labels: {
|
||||
app: "task-run",
|
||||
},
|
||||
},
|
||||
spec: {
|
||||
restartPolicy: "Never",
|
||||
imagePullSecrets: [
|
||||
{
|
||||
name: "registry-trigger",
|
||||
},
|
||||
],
|
||||
containers: [
|
||||
{
|
||||
name: opts.runId,
|
||||
image: this.#getImageFromRunId(opts.runId),
|
||||
name: opts.attemptId,
|
||||
image: opts.image,
|
||||
ports: [
|
||||
{
|
||||
containerPort: 8000,
|
||||
@@ -166,6 +163,18 @@ class KubernetesTaskOperations implements TaskOperations {
|
||||
name: "DEBUG",
|
||||
value: "true",
|
||||
},
|
||||
{
|
||||
name: "TRIGGER_ENV_ID",
|
||||
value: opts.envId,
|
||||
},
|
||||
{
|
||||
name: "TRIGGER_ATTEMPT_ID",
|
||||
value: opts.attemptId,
|
||||
},
|
||||
{
|
||||
name: "OTEL_EXPORTER_OTLP_ENDPOINT",
|
||||
value: OTEL_EXPORTER_OTLP_ENDPOINT,
|
||||
},
|
||||
{
|
||||
name: "POD_NAME",
|
||||
valueFrom: {
|
||||
@@ -200,6 +209,7 @@ class KubernetesTaskOperations implements TaskOperations {
|
||||
}
|
||||
|
||||
async restore(opts: {
|
||||
attemptId: string;
|
||||
runId: string;
|
||||
image: string;
|
||||
name: string;
|
||||
@@ -376,185 +386,7 @@ class KubernetesTaskOperations implements TaskOperations {
|
||||
}
|
||||
}
|
||||
|
||||
interface Provider {
|
||||
tasks: TaskOperations;
|
||||
}
|
||||
|
||||
type KubernetesProviderOptions = {
|
||||
tasks: KubernetesTaskOperations;
|
||||
host?: string;
|
||||
port: number;
|
||||
};
|
||||
|
||||
class KubernetesProvider implements Provider {
|
||||
tasks: KubernetesTaskOperations;
|
||||
|
||||
#httpServer: ReturnType<typeof createServer>;
|
||||
#platformSocket: Socket<ProviderServerToClientEvents, ProviderClientToServerEvents>;
|
||||
|
||||
constructor(private options: KubernetesProviderOptions) {
|
||||
this.tasks = options.tasks;
|
||||
this.#httpServer = this.#createHttpServer();
|
||||
this.#platformSocket = this.#createPlatformSocket();
|
||||
}
|
||||
|
||||
#createPlatformSocket() {
|
||||
const socket: Socket<ProviderServerToClientEvents, ProviderClientToServerEvents> = io(
|
||||
`ws://${PLATFORM_HOST}:${PLATFORM_WS_PORT}/provider`,
|
||||
{
|
||||
transports: ["websocket"],
|
||||
auth: {
|
||||
token: PLATFORM_SECRET,
|
||||
},
|
||||
extraHeaders: {
|
||||
"x-trigger-provider-type": "kubernetes",
|
||||
},
|
||||
}
|
||||
);
|
||||
|
||||
const logger = new SimpleLogger(`[platform][${socket.id ?? "NO_ID"}]`);
|
||||
|
||||
socket.on("connect_error", (err) => {
|
||||
logger.error(`connect_error: ${err.message}`);
|
||||
});
|
||||
|
||||
socket.on("connect", () => {
|
||||
logger.log("connect");
|
||||
});
|
||||
|
||||
socket.on("disconnect", () => {
|
||||
logger.log("disconnect");
|
||||
});
|
||||
|
||||
socket.on("GET", async (message) => {
|
||||
logger.log("[GET]", message);
|
||||
this.tasks.get({ runId: message.name });
|
||||
});
|
||||
|
||||
socket.on("DELETE", async (message, callback) => {
|
||||
logger.log("[DELETE]", message);
|
||||
|
||||
callback({
|
||||
message: "delete request received",
|
||||
});
|
||||
|
||||
this.tasks.delete({ runId: message.name });
|
||||
});
|
||||
|
||||
socket.on("INDEX", async (message) => {
|
||||
logger.log("[INDEX]", message);
|
||||
|
||||
await this.tasks.index({
|
||||
contentHash: message.contentHash,
|
||||
imageTag: message.imageTag,
|
||||
});
|
||||
});
|
||||
|
||||
socket.on("INVOKE", async (message) => {
|
||||
logger.log("[INVOKE]", message);
|
||||
|
||||
await this.tasks.create({
|
||||
runId: message.name,
|
||||
image: message.name,
|
||||
machine: message.machine,
|
||||
});
|
||||
});
|
||||
|
||||
socket.on("RESTORE", async (message) => {
|
||||
logger.log("[RESTORE]", message);
|
||||
|
||||
// await this.tasks.restore({});
|
||||
});
|
||||
|
||||
socket.on("HEALTH", async (message) => {
|
||||
logger.log("[HEALTH]", message);
|
||||
});
|
||||
|
||||
return socket;
|
||||
}
|
||||
|
||||
#createHttpServer() {
|
||||
const httpServer = createServer(async (req, res) => {
|
||||
logger.log(`[${req.method}]`, req.url);
|
||||
|
||||
const reply = new HttpReply(res);
|
||||
|
||||
switch (req.url) {
|
||||
case "/health": {
|
||||
return reply.text("ok");
|
||||
}
|
||||
case "/whoami": {
|
||||
return reply.text(`${POD_NAME}`);
|
||||
}
|
||||
case "/close": {
|
||||
this.#platformSocket.close();
|
||||
return reply.text("platform socket closed");
|
||||
}
|
||||
case "/delete": {
|
||||
const body = await getTextBody(req);
|
||||
|
||||
await this.tasks.delete({ runId: body });
|
||||
|
||||
return reply.text(`sent delete request: ${body}`);
|
||||
}
|
||||
case "/invoke": {
|
||||
const body = await getTextBody(req);
|
||||
|
||||
await this.tasks.create({
|
||||
runId: body,
|
||||
image: body,
|
||||
machine: {
|
||||
cpu: "1",
|
||||
memory: "100Mi",
|
||||
},
|
||||
});
|
||||
|
||||
return reply.text(`sent restore request: ${body}`);
|
||||
}
|
||||
case "/restore": {
|
||||
const body = await getTextBody(req);
|
||||
|
||||
const items = body.split("&");
|
||||
const image = items[0];
|
||||
const baseImageTag = items[1] ?? image;
|
||||
|
||||
await this.tasks.restore({
|
||||
runId: image,
|
||||
name: `${image}-restore`,
|
||||
image,
|
||||
checkpointId: baseImageTag,
|
||||
machine: {
|
||||
cpu: "1",
|
||||
memory: "100Mi",
|
||||
},
|
||||
});
|
||||
|
||||
return reply.text(`sent restore request: ${body}`);
|
||||
}
|
||||
default: {
|
||||
return reply.empty(404);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
httpServer.on("clientError", (err, socket) => {
|
||||
socket.end("HTTP/1.1 400 Bad Request\r\n\r\n");
|
||||
});
|
||||
|
||||
httpServer.on("listening", () => {
|
||||
logger.log("server listening on port", this.options.port);
|
||||
});
|
||||
|
||||
return httpServer;
|
||||
}
|
||||
|
||||
listen() {
|
||||
this.#httpServer.listen(this.options.port, this.options.host ?? "0.0.0.0");
|
||||
}
|
||||
}
|
||||
|
||||
const provider = new KubernetesProvider({
|
||||
port: HTTP_SERVER_PORT,
|
||||
const provider = new ProviderShell({
|
||||
tasks: new KubernetesTaskOperations(),
|
||||
});
|
||||
|
||||
|
||||
@@ -420,7 +420,7 @@ export class EnvironmentVariablesRepository implements Repository {
|
||||
|
||||
return [
|
||||
{
|
||||
key: "TRIGGER_API_KEY",
|
||||
key: "TRIGGER_SECRET_KEY",
|
||||
value: environment.apiKey,
|
||||
},
|
||||
{
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
import {
|
||||
ClientToSharedQueueMessages,
|
||||
CoordinatorToPlatformMessages,
|
||||
PlatformToCoordinatorMessages,
|
||||
PlatformToProviderMessages,
|
||||
ProviderToPlatformMessages,
|
||||
SharedQueueToClientMessages,
|
||||
ZodNamespace,
|
||||
clientWebsocketMessages,
|
||||
serverWebsocketMessages,
|
||||
} from "@trigger.dev/core/v3";
|
||||
import { Server } from "socket.io";
|
||||
import { env } from "~/env.server";
|
||||
@@ -46,7 +46,7 @@ function createCoordinatorNamespace(io: Server) {
|
||||
authToken: env.COORDINATOR_SECRET,
|
||||
clientMessages: CoordinatorToPlatformMessages,
|
||||
serverMessages: PlatformToCoordinatorMessages,
|
||||
messageHandler: {
|
||||
handlers: {
|
||||
READY_FOR_EXECUTION: async (message) => {
|
||||
const payload = await sharedQueueTasks.getExecutionPayloadFromAttempt(message.attemptId);
|
||||
|
||||
@@ -111,8 +111,8 @@ function createSharedQueueConsumerNamespace(io: Server) {
|
||||
io,
|
||||
name: "shared-queue",
|
||||
authToken: env.PROVIDER_SECRET,
|
||||
clientMessages: clientWebsocketMessages,
|
||||
serverMessages: serverWebsocketMessages,
|
||||
clientMessages: ClientToSharedQueueMessages,
|
||||
serverMessages: SharedQueueToClientMessages,
|
||||
onConnection: async (socket, handler, sender, logger) => {
|
||||
const sharedSocketConnection = new SharedSocketConnection(
|
||||
sharedQueue.namespace,
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { CoordinatorToPlatformEvents } from "@trigger.dev/core/v3";
|
||||
import { CoordinatorToPlatformMessages, InferSocketMessageSchema } from "@trigger.dev/core/v3";
|
||||
import type { Checkpoint } from "@trigger.dev/database";
|
||||
import { PrismaClient, prisma } from "~/db.server";
|
||||
import { logger } from "~/services/logger.server";
|
||||
@@ -13,7 +13,7 @@ export class CreateCheckpointService {
|
||||
}
|
||||
|
||||
public async call(
|
||||
params: Parameters<CoordinatorToPlatformEvents["CHECKPOINT_CREATED"]>[0]
|
||||
params: InferSocketMessageSchema<typeof CoordinatorToPlatformMessages, "CHECKPOINT_CREATED">
|
||||
): Promise<Checkpoint> {
|
||||
const attempt = await this.#prismaClient.taskRunAttempt.findUniqueOrThrow({
|
||||
where: {
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
import {
|
||||
CoordinatorToProdWorkerEvents,
|
||||
ProdWorkerToCoordinatorEvents,
|
||||
CoordinatorToProdWorkerMessages,
|
||||
ProdWorkerToCoordinatorMessages,
|
||||
TaskResource,
|
||||
ZodSocketConnection,
|
||||
} from "@trigger.dev/core/v3";
|
||||
import { HttpReply, getTextBody, SimpleLogger, getRandomPortNumber } from "@trigger.dev/core-apps";
|
||||
import { createServer } from "node:http";
|
||||
import { io, Socket } from "socket.io-client";
|
||||
import { ProdBackgroundWorker } from "./prod/backgroundWorker";
|
||||
|
||||
const HTTP_SERVER_PORT = Number(process.env.HTTP_SERVER_PORT || getRandomPortNumber());
|
||||
@@ -33,7 +33,10 @@ class ProdWorker {
|
||||
#httpPort: number;
|
||||
#backgroundWorker: ProdBackgroundWorker;
|
||||
#httpServer: ReturnType<typeof createServer>;
|
||||
#coordinatorSocket: Socket<CoordinatorToProdWorkerEvents, ProdWorkerToCoordinatorEvents>;
|
||||
#coordinatorSocket: ZodSocketConnection<
|
||||
typeof ProdWorkerToCoordinatorMessages,
|
||||
typeof CoordinatorToProdWorkerMessages
|
||||
>;
|
||||
|
||||
constructor(
|
||||
port: number,
|
||||
@@ -46,26 +49,23 @@ class ProdWorker {
|
||||
env: {
|
||||
TRIGGER_API_URL: this.apiUrl,
|
||||
TRIGGER_SECRET_KEY: this.apiKey,
|
||||
OTEL_EXPORTER_OTLP_ENDPOINT:
|
||||
process.env.OTEL_EXPORTER_OTLP_ENDPOINT ?? "http://0.0.0.0:4318",
|
||||
},
|
||||
contentHash: this.contentHash,
|
||||
});
|
||||
this.#backgroundWorker.onTaskHeartbeat.attach((attemptFriendlyId) => {
|
||||
this.#coordinatorSocket.emit("TASK_HEARTBEAT", { version: "v1", attemptFriendlyId });
|
||||
this.#coordinatorSocket.send("TASK_HEARTBEAT", { attemptFriendlyId });
|
||||
});
|
||||
this.#backgroundWorker.onWaitForBatch.attach((message) => {
|
||||
this.#coordinatorSocket.emit("WAIT_FOR_BATCH", { version: "v1", ...message });
|
||||
this.#coordinatorSocket.send("WAIT_FOR_BATCH", message);
|
||||
});
|
||||
this.#backgroundWorker.onWaitForDuration.attach((message) => {
|
||||
this.#coordinatorSocket.emit(
|
||||
"WAIT_FOR_DURATION",
|
||||
{ version: "v1", ...message },
|
||||
({ success }) => {
|
||||
logger.log("WAIT_FOR_DURATION", { success });
|
||||
}
|
||||
);
|
||||
this.#backgroundWorker.onWaitForDuration.attach(async (message) => {
|
||||
const { success } = await this.#coordinatorSocket.sendWithAck("WAIT_FOR_DURATION", message);
|
||||
logger.log("WAIT_FOR_DURATION", { success });
|
||||
});
|
||||
this.#backgroundWorker.onWaitForTask.attach((message) => {
|
||||
this.#coordinatorSocket.emit("WAIT_FOR_TASK", { version: "v1", ...message });
|
||||
this.#coordinatorSocket.send("WAIT_FOR_TASK", message);
|
||||
});
|
||||
|
||||
this.#httpPort = port;
|
||||
@@ -73,93 +73,84 @@ class ProdWorker {
|
||||
}
|
||||
|
||||
#createCoordinatorSocket() {
|
||||
const socket: Socket<CoordinatorToProdWorkerEvents, ProdWorkerToCoordinatorEvents> = io(
|
||||
`ws://${COORDINATOR_HOST}:${COORDINATOR_PORT}/prod-worker`,
|
||||
{
|
||||
transports: ["websocket"],
|
||||
extraHeaders: {
|
||||
"x-machine-name": MACHINE_NAME,
|
||||
"x-pod-name": POD_NAME,
|
||||
"x-trigger-content-hash": this.contentHash,
|
||||
"x-trigger-cli-package-version": this.cliPackageVersion,
|
||||
"x-trigger-project-ref": this.projectRef,
|
||||
"x-trigger-attempt-id": this.attemptId,
|
||||
"x-trigger-env-id": this.envId,
|
||||
const coordinatorConnection = new ZodSocketConnection({
|
||||
namespace: "prod-worker",
|
||||
host: COORDINATOR_HOST,
|
||||
port: COORDINATOR_PORT,
|
||||
clientMessages: ProdWorkerToCoordinatorMessages,
|
||||
serverMessages: CoordinatorToProdWorkerMessages,
|
||||
extraHeaders: {
|
||||
"x-machine-name": MACHINE_NAME,
|
||||
"x-pod-name": POD_NAME,
|
||||
"x-trigger-content-hash": this.contentHash,
|
||||
"x-trigger-cli-package-version": this.cliPackageVersion,
|
||||
"x-trigger-project-ref": this.projectRef,
|
||||
"x-trigger-attempt-id": this.attemptId,
|
||||
"x-trigger-env-id": this.envId,
|
||||
},
|
||||
handlers: {
|
||||
RESUME: async (message) => {
|
||||
for (let i = 0; i < message.completions.length; i++) {
|
||||
const completion = message.completions[i];
|
||||
const execution = message.executions[i];
|
||||
|
||||
if (!completion || !execution) continue;
|
||||
|
||||
this.#backgroundWorker.taskRunCompletedNotification(completion, execution);
|
||||
}
|
||||
},
|
||||
}
|
||||
);
|
||||
EXECUTE_TASK_RUN: async (message) => {
|
||||
if (this.executing || this.completed) {
|
||||
return {
|
||||
success: false,
|
||||
};
|
||||
}
|
||||
|
||||
const logger = new SimpleLogger(`[coordinator][${socket.id ?? "NO_ID"}]`);
|
||||
this.executing = true;
|
||||
const completion = await this.#backgroundWorker.executeTaskRun(message.executionPayload);
|
||||
|
||||
socket.on("connect_error", (err) => {
|
||||
logger.error(`connect_error: ${err.message}`);
|
||||
});
|
||||
logger.log("completed", completion);
|
||||
|
||||
socket.on("connect", async () => {
|
||||
logger.log("connect");
|
||||
this.completed = true;
|
||||
this.executing = false;
|
||||
|
||||
if (process.env.INDEX_TASKS === "true") {
|
||||
const taskResources = await this.#initializeWorker();
|
||||
const { success } = await socket.emitWithAck("INDEX_TASKS", {
|
||||
version: "v1",
|
||||
...taskResources,
|
||||
});
|
||||
if (success) {
|
||||
logger.log("indexing done, shutting down..");
|
||||
process.exit(0);
|
||||
setTimeout(() => {
|
||||
process.exit(0);
|
||||
}, 2000);
|
||||
|
||||
// TODO: replace ack with emit
|
||||
return {
|
||||
success: true,
|
||||
completion,
|
||||
};
|
||||
},
|
||||
},
|
||||
onConnection: async (socket, handler, sender, logger) => {
|
||||
if (process.env.INDEX_TASKS === "true") {
|
||||
const taskResources = await this.#initializeWorker();
|
||||
|
||||
const { success } = await socket.emitWithAck("INDEX_TASKS", {
|
||||
version: "v1",
|
||||
...taskResources,
|
||||
});
|
||||
|
||||
if (success) {
|
||||
logger("indexing done, shutting down..");
|
||||
process.exit(0);
|
||||
} else {
|
||||
logger("indexing failure, shutting down..");
|
||||
process.exit(1);
|
||||
}
|
||||
} else {
|
||||
logger.log("indexing failure, shutting down..");
|
||||
process.exit(1);
|
||||
socket.emit("READY_FOR_EXECUTION", {
|
||||
version: "v1",
|
||||
attemptId: this.attemptId,
|
||||
});
|
||||
}
|
||||
} else {
|
||||
socket.emit("READY_FOR_EXECUTION", {
|
||||
version: "v1",
|
||||
attemptId: process.env.TRIGGER_ATTEMPT_ID!,
|
||||
});
|
||||
}
|
||||
},
|
||||
});
|
||||
|
||||
socket.on("disconnect", () => {
|
||||
logger.log("disconnect");
|
||||
});
|
||||
|
||||
socket.on("RESUME", async (message) => {
|
||||
logger.log("[RESUME]", message);
|
||||
|
||||
for (let i = 0; i < message.completions.length; i++) {
|
||||
const completion = message.completions[i];
|
||||
const execution = message.executions[i];
|
||||
|
||||
if (!completion || !execution) continue;
|
||||
|
||||
this.#backgroundWorker.taskRunCompletedNotification(completion, execution);
|
||||
}
|
||||
});
|
||||
|
||||
socket.on("EXECUTE_TASK_RUN", async (message, callback) => {
|
||||
logger.log("[EXECUTE_TASK_RUN]", { attempt: message.payload.execution.attempt });
|
||||
|
||||
if (this.executing || this.completed) {
|
||||
return;
|
||||
}
|
||||
|
||||
this.executing = true;
|
||||
const completion = await this.#backgroundWorker.executeTaskRun(message.payload);
|
||||
|
||||
logger.log("completed", completion);
|
||||
|
||||
// TODO: replace ack with emit
|
||||
callback({ completion });
|
||||
|
||||
this.completed = true;
|
||||
this.executing = false;
|
||||
|
||||
setTimeout(() => {
|
||||
process.exit(0);
|
||||
}, 1000);
|
||||
});
|
||||
|
||||
return socket;
|
||||
return coordinatorConnection;
|
||||
}
|
||||
|
||||
#createHttpServer() {
|
||||
@@ -188,16 +179,11 @@ class ProdWorker {
|
||||
return reply.text(this.contentHash);
|
||||
|
||||
case "/wait":
|
||||
this.#coordinatorSocket.emit(
|
||||
"WAIT_FOR_DURATION",
|
||||
{
|
||||
version: "v1",
|
||||
ms: 60_000,
|
||||
},
|
||||
({ success }) => {
|
||||
logger.log("WAIT_FOR_DURATION", { success });
|
||||
}
|
||||
);
|
||||
const { success } = await this.#coordinatorSocket.sendWithAck("WAIT_FOR_DURATION", {
|
||||
version: "v1",
|
||||
ms: 60_000,
|
||||
});
|
||||
logger.log("WAIT_FOR_DURATION", { success });
|
||||
// this is required when C/Ring established connections
|
||||
this.#coordinatorSocket.close();
|
||||
return reply.text("sent WAIT");
|
||||
@@ -207,7 +193,7 @@ class ProdWorker {
|
||||
return reply.empty();
|
||||
|
||||
case "/close":
|
||||
this.#coordinatorSocket.emitWithAck("LOG", {
|
||||
this.#coordinatorSocket.sendWithAck("LOG", {
|
||||
version: "v1",
|
||||
text: "close without delay",
|
||||
});
|
||||
@@ -215,7 +201,7 @@ class ProdWorker {
|
||||
return reply.empty();
|
||||
|
||||
case "/close-delay":
|
||||
this.#coordinatorSocket.emitWithAck("LOG", {
|
||||
this.#coordinatorSocket.sendWithAck("LOG", {
|
||||
version: "v1",
|
||||
text: "close with delay",
|
||||
});
|
||||
@@ -225,7 +211,7 @@ class ProdWorker {
|
||||
return reply.empty();
|
||||
|
||||
case "/log":
|
||||
this.#coordinatorSocket.emitWithAck("LOG", {
|
||||
this.#coordinatorSocket.sendWithAck("LOG", {
|
||||
version: "v1",
|
||||
text: await getTextBody(req),
|
||||
});
|
||||
@@ -236,7 +222,7 @@ class ProdWorker {
|
||||
return reply.text("got preStop request");
|
||||
|
||||
case "/ready":
|
||||
this.#coordinatorSocket.emit("READY_FOR_EXECUTION", {
|
||||
this.#coordinatorSocket.send("READY_FOR_EXECUTION", {
|
||||
version: "v1",
|
||||
attemptId: this.attemptId,
|
||||
});
|
||||
|
||||
@@ -112,7 +112,7 @@ export class ProdBackgroundWorker {
|
||||
resolved = true;
|
||||
child.kill();
|
||||
reject(new Error("Worker timed out"));
|
||||
}, 1000);
|
||||
}, 10_000);
|
||||
|
||||
child.on("message", async (msg: any) => {
|
||||
const message = this._handler.parseMessage(msg);
|
||||
|
||||
@@ -1,2 +1,3 @@
|
||||
export * from "./http";
|
||||
export * from "./logger";
|
||||
export * from "./provider";
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
import { createServer } from "node:http";
|
||||
import {
|
||||
ClientToSharedQueueMessages,
|
||||
clientWebsocketMessages,
|
||||
PlatformToProviderMessages,
|
||||
ProviderToPlatformMessages,
|
||||
SharedQueueToClientMessages,
|
||||
ZodMessageSender,
|
||||
ZodSocketConnection,
|
||||
} from "@trigger.dev/core/v3";
|
||||
import { getRandomPortNumber, HttpReply, getTextBody } from "./http";
|
||||
import { SimpleLogger } from "./logger";
|
||||
|
||||
const HTTP_SERVER_PORT = Number(process.env.HTTP_SERVER_PORT || getRandomPortNumber());
|
||||
const MACHINE_NAME = process.env.MACHINE_NAME || "local";
|
||||
|
||||
const PLATFORM_HOST = process.env.PLATFORM_HOST || "127.0.0.1";
|
||||
const PLATFORM_WS_PORT = process.env.PLATFORM_WS_PORT || 3030;
|
||||
const PLATFORM_SECRET = process.env.PLATFORM_SECRET || "provider-secret";
|
||||
|
||||
const logger = new SimpleLogger(`[${MACHINE_NAME}]`);
|
||||
|
||||
export interface TaskOperations {
|
||||
create: (...args: any[]) => Promise<any>;
|
||||
restore: (...args: any[]) => Promise<any>;
|
||||
delete: (...args: any[]) => Promise<any>;
|
||||
get: (...args: any[]) => Promise<any>;
|
||||
index: (...args: any[]) => Promise<any>;
|
||||
}
|
||||
|
||||
type ProviderShellOptions = {
|
||||
tasks: TaskOperations;
|
||||
host?: string;
|
||||
port?: number;
|
||||
};
|
||||
|
||||
interface Provider {
|
||||
tasks: TaskOperations;
|
||||
}
|
||||
|
||||
export class ProviderShell implements Provider {
|
||||
tasks: TaskOperations;
|
||||
|
||||
#httpPort: number;
|
||||
#httpServer: ReturnType<typeof createServer>;
|
||||
#platformSocket: ZodSocketConnection<
|
||||
typeof ProviderToPlatformMessages,
|
||||
typeof PlatformToProviderMessages
|
||||
>;
|
||||
|
||||
constructor(private options: ProviderShellOptions) {
|
||||
this.tasks = options.tasks;
|
||||
this.#httpPort = options.port ?? HTTP_SERVER_PORT;
|
||||
this.#httpServer = this.#createHttpServer();
|
||||
this.#platformSocket = this.#createPlatformSocket();
|
||||
this.#createSharedQueueSocket();
|
||||
}
|
||||
|
||||
#createSharedQueueSocket() {
|
||||
const sharedQueueConnection = new ZodSocketConnection({
|
||||
namespace: "shared-queue",
|
||||
host: PLATFORM_HOST,
|
||||
port: Number(PLATFORM_WS_PORT),
|
||||
clientMessages: ClientToSharedQueueMessages,
|
||||
serverMessages: SharedQueueToClientMessages,
|
||||
authToken: PLATFORM_SECRET,
|
||||
handlers: {
|
||||
SERVER_READY: async (message) => {
|
||||
// TODO: create new schema without worker requirement
|
||||
await sender.send("READY_FOR_TASKS", {
|
||||
backgroundWorkerId: "placeholder",
|
||||
});
|
||||
},
|
||||
BACKGROUND_WORKER_MESSAGE: async (message) => {
|
||||
if (message.data.type === "SCHEDULE_ATTEMPT") {
|
||||
this.tasks.create({
|
||||
envId: message.data.envId,
|
||||
attemptId: message.data.id,
|
||||
image: message.data.image,
|
||||
machine: {},
|
||||
});
|
||||
}
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const sender = new ZodMessageSender({
|
||||
schema: clientWebsocketMessages,
|
||||
sender: async (message) => {
|
||||
return new Promise((resolve, reject) => {
|
||||
try {
|
||||
const { type, ...payload } = message;
|
||||
sharedQueueConnection.socket.emit(type, payload as any);
|
||||
resolve();
|
||||
} catch (err) {
|
||||
reject(err);
|
||||
}
|
||||
});
|
||||
},
|
||||
});
|
||||
|
||||
return sharedQueueConnection;
|
||||
}
|
||||
|
||||
#createPlatformSocket() {
|
||||
const platformConnection = new ZodSocketConnection({
|
||||
namespace: "provider",
|
||||
host: PLATFORM_HOST,
|
||||
port: Number(PLATFORM_WS_PORT),
|
||||
clientMessages: ProviderToPlatformMessages,
|
||||
serverMessages: PlatformToProviderMessages,
|
||||
authToken: PLATFORM_SECRET,
|
||||
extraHeaders: {
|
||||
"x-trigger-provider-type": "docker",
|
||||
},
|
||||
handlers: {
|
||||
DELETE: async (message) => {
|
||||
this.tasks.delete({ runId: message.name });
|
||||
|
||||
return {
|
||||
message: "delete request received",
|
||||
};
|
||||
},
|
||||
GET: async (message) => {
|
||||
this.tasks.get({ runId: message.name });
|
||||
},
|
||||
HEALTH: async (message) => {
|
||||
return {
|
||||
status: "ok",
|
||||
};
|
||||
},
|
||||
INDEX: async (message) => {
|
||||
try {
|
||||
await this.tasks.index({
|
||||
contentHash: message.contentHash,
|
||||
imageTag: message.imageTag,
|
||||
envId: message.envId,
|
||||
});
|
||||
} catch (error) {
|
||||
logger.error("task index failed", error);
|
||||
}
|
||||
},
|
||||
RESTORE: async (message) => {},
|
||||
},
|
||||
});
|
||||
|
||||
return platformConnection;
|
||||
}
|
||||
|
||||
#createHttpServer() {
|
||||
const httpServer = createServer(async (req, res) => {
|
||||
logger.log(`[${req.method}]`, req.url);
|
||||
|
||||
const reply = new HttpReply(res);
|
||||
|
||||
switch (req.url) {
|
||||
case "/health": {
|
||||
return reply.text("ok");
|
||||
}
|
||||
case "/whoami": {
|
||||
return reply.text(`${MACHINE_NAME}`);
|
||||
}
|
||||
case "/close": {
|
||||
this.#platformSocket.close();
|
||||
return reply.text("platform socket closed");
|
||||
}
|
||||
case "/delete": {
|
||||
const body = await getTextBody(req);
|
||||
|
||||
await this.tasks.delete({ runId: body });
|
||||
|
||||
return reply.text(`sent delete request: ${body}`);
|
||||
}
|
||||
case "/invoke": {
|
||||
const body = await getTextBody(req);
|
||||
|
||||
await this.tasks.create({
|
||||
attemptId: body,
|
||||
envId: "placeholder",
|
||||
image: body,
|
||||
machine: {
|
||||
cpu: "1",
|
||||
memory: "100Mi",
|
||||
},
|
||||
});
|
||||
|
||||
return reply.text(`sent restore request: ${body}`);
|
||||
}
|
||||
case "/restore": {
|
||||
const body = await getTextBody(req);
|
||||
|
||||
const items = body.split("&");
|
||||
const image = items[0];
|
||||
const baseImageTag = items[1] ?? image;
|
||||
|
||||
// await this.tasks.restore({});
|
||||
|
||||
return reply.text(`sent restore request: ${body}`);
|
||||
}
|
||||
default: {
|
||||
return reply.empty(404);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
httpServer.on("clientError", (err, socket) => {
|
||||
socket.end("HTTP/1.1 400 Bad Request\r\n\r\n");
|
||||
});
|
||||
|
||||
httpServer.on("listening", () => {
|
||||
logger.log("server listening on port", this.#httpPort);
|
||||
});
|
||||
|
||||
return httpServer;
|
||||
}
|
||||
|
||||
listen() {
|
||||
this.#httpServer.listen(this.#httpPort, this.options.host ?? "0.0.0.0");
|
||||
}
|
||||
}
|
||||
@@ -3,4 +3,7 @@ import { packageOptions, defineConfig } from "@trigger.dev/tsup";
|
||||
export default defineConfig({
|
||||
...packageOptions,
|
||||
config: "tsconfig.build.json",
|
||||
banner: {
|
||||
js: "import { createRequire } from 'module';const require = createRequire(import.meta.url);",
|
||||
},
|
||||
});
|
||||
|
||||
@@ -75,6 +75,7 @@
|
||||
"@opentelemetry/semantic-conventions": "^1.21.0",
|
||||
"humanize-duration": "^3.27.3",
|
||||
"socket.io": "^4.7.4",
|
||||
"socket.io-client": "^4.7.4",
|
||||
"ulidx": "^2.2.1",
|
||||
"zod": "3.22.3",
|
||||
"zod-error": "1.5.0"
|
||||
|
||||
@@ -4,6 +4,7 @@ export * from "./schemas";
|
||||
export * from "./apiClient";
|
||||
export * from "./zodMessageHandler";
|
||||
export * from "./zodNamespace";
|
||||
export * from "./zodSocket";
|
||||
export * from "./errors";
|
||||
export * from "./runtime-api";
|
||||
export * from "./logger-api";
|
||||
|
||||
@@ -1,7 +1,12 @@
|
||||
import { z } from "zod";
|
||||
import { RequireKeys } from "../types";
|
||||
import { TaskRunExecution, TaskRunExecutionResult } from "./common";
|
||||
import { ProdTaskRunExecution } from "./messages";
|
||||
import {
|
||||
BackgroundWorkerClientMessages,
|
||||
BackgroundWorkerServerMessages,
|
||||
ProdTaskRunExecution,
|
||||
ProdTaskRunExecutionPayload,
|
||||
} from "./messages";
|
||||
import { TaskResource } from "./resources";
|
||||
|
||||
export const Config = z.object({
|
||||
@@ -25,93 +30,288 @@ export const Machine = z.object({
|
||||
export type Machine = z.infer<typeof Machine>;
|
||||
|
||||
export const ProviderToPlatformMessages = {
|
||||
LOG: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
data: z.string(),
|
||||
}),
|
||||
LOG: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
data: z.string(),
|
||||
}),
|
||||
},
|
||||
LOG_WITH_ACK: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
data: z.string(),
|
||||
}),
|
||||
callback: z.object({
|
||||
status: z.literal("ok"),
|
||||
}),
|
||||
},
|
||||
};
|
||||
|
||||
export const PlatformToProviderMessages = {
|
||||
HEALTH: z.object({
|
||||
// TODO: callback: (ack: { status: "ok" }) => void
|
||||
version: z.literal("v1").default("v1"),
|
||||
}),
|
||||
INDEX: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
imageTag: z.string(),
|
||||
contentHash: z.string(),
|
||||
envId: z.string(),
|
||||
}),
|
||||
INVOKE: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
name: z.string(),
|
||||
machine: Machine,
|
||||
}),
|
||||
RESTORE: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
id: z.string(),
|
||||
attemptId: z.string(),
|
||||
type: z.enum(["DOCKER", "KUBERNETES"]),
|
||||
location: z.string(),
|
||||
reason: z.string().optional(),
|
||||
}),
|
||||
DELETE: z.object({
|
||||
// TODO: callback: (ack: { message: string }) => void
|
||||
version: z.literal("v1").default("v1"),
|
||||
name: z.string(),
|
||||
}),
|
||||
GET: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
name: z.string(),
|
||||
}),
|
||||
HEALTH: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
}),
|
||||
callback: z.object({
|
||||
status: z.literal("ok"),
|
||||
}),
|
||||
},
|
||||
INDEX: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
imageTag: z.string(),
|
||||
contentHash: z.string(),
|
||||
envId: z.string(),
|
||||
}),
|
||||
},
|
||||
INVOKE: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
name: z.string(),
|
||||
machine: Machine,
|
||||
}),
|
||||
},
|
||||
RESTORE: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
id: z.string(),
|
||||
attemptId: z.string(),
|
||||
type: z.enum(["DOCKER", "KUBERNETES"]),
|
||||
location: z.string(),
|
||||
reason: z.string().optional(),
|
||||
}),
|
||||
},
|
||||
DELETE: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
name: z.string(),
|
||||
}),
|
||||
callback: z.object({
|
||||
message: z.string(),
|
||||
}),
|
||||
},
|
||||
GET: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
name: z.string(),
|
||||
}),
|
||||
},
|
||||
};
|
||||
|
||||
export const CoordinatorToPlatformMessages = {
|
||||
LOG: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
metadata: z.any(),
|
||||
text: z.string(),
|
||||
}),
|
||||
CREATE_WORKER: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
projectRef: z.string(),
|
||||
envId: z.string(),
|
||||
metadata: z.object({
|
||||
cliPackageVersion: z.string(),
|
||||
contentHash: z.string(),
|
||||
packageVersion: z.string(),
|
||||
tasks: TaskResource.array(),
|
||||
LOG: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
metadata: z.any(),
|
||||
text: z.string(),
|
||||
}),
|
||||
}),
|
||||
// TODO: callback: (ack: { success: false } | { success: true; payload: ProdTaskRunExecutionPayload }) => void
|
||||
READY_FOR_EXECUTION: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptId: z.string(),
|
||||
}),
|
||||
TASK_RUN_COMPLETED: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
execution: ProdTaskRunExecution,
|
||||
completion: TaskRunExecutionResult,
|
||||
}),
|
||||
TASK_HEARTBEAT: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptFriendlyId: z.string(),
|
||||
}),
|
||||
CHECKPOINT_CREATED: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptId: z.string(),
|
||||
docker: z.boolean(),
|
||||
location: z.string(),
|
||||
reason: z.string().optional(),
|
||||
}),
|
||||
},
|
||||
CREATE_WORKER: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
projectRef: z.string(),
|
||||
envId: z.string(),
|
||||
metadata: z.object({
|
||||
cliPackageVersion: z.string(),
|
||||
contentHash: z.string(),
|
||||
packageVersion: z.string(),
|
||||
tasks: TaskResource.array(),
|
||||
}),
|
||||
}),
|
||||
callback: z.discriminatedUnion("success", [
|
||||
z.object({
|
||||
success: z.literal(false),
|
||||
}),
|
||||
z.object({
|
||||
success: z.literal(true),
|
||||
}),
|
||||
]),
|
||||
},
|
||||
READY_FOR_EXECUTION: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptId: z.string(),
|
||||
}),
|
||||
callback: z.discriminatedUnion("success", [
|
||||
z.object({
|
||||
success: z.literal(false),
|
||||
}),
|
||||
z.object({
|
||||
success: z.literal(true),
|
||||
payload: ProdTaskRunExecutionPayload,
|
||||
}),
|
||||
]),
|
||||
},
|
||||
TASK_RUN_COMPLETED: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
execution: ProdTaskRunExecution,
|
||||
completion: TaskRunExecutionResult,
|
||||
}),
|
||||
},
|
||||
TASK_HEARTBEAT: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptFriendlyId: z.string(),
|
||||
}),
|
||||
},
|
||||
CHECKPOINT_CREATED: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptId: z.string(),
|
||||
docker: z.boolean(),
|
||||
location: z.string(),
|
||||
reason: z.string().optional(),
|
||||
}),
|
||||
},
|
||||
};
|
||||
|
||||
export const PlatformToCoordinatorMessages = {
|
||||
RESUME: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptId: z.string(),
|
||||
image: z.string(),
|
||||
completions: TaskRunExecutionResult.array(),
|
||||
executions: TaskRunExecution.array(),
|
||||
}),
|
||||
RESUME: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptId: z.string(),
|
||||
image: z.string(),
|
||||
completions: TaskRunExecutionResult.array(),
|
||||
executions: TaskRunExecution.array(),
|
||||
}),
|
||||
},
|
||||
};
|
||||
|
||||
export const ClientToSharedQueueMessages = {
|
||||
READY_FOR_TASKS: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
backgroundWorkerId: z.string(),
|
||||
}),
|
||||
},
|
||||
BACKGROUND_WORKER_DEPRECATED: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
backgroundWorkerId: z.string(),
|
||||
}),
|
||||
},
|
||||
BACKGROUND_WORKER_MESSAGE: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
backgroundWorkerId: z.string(),
|
||||
data: BackgroundWorkerClientMessages,
|
||||
}),
|
||||
},
|
||||
};
|
||||
|
||||
export const SharedQueueToClientMessages = {
|
||||
SERVER_READY: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
id: z.string(),
|
||||
}),
|
||||
},
|
||||
BACKGROUND_WORKER_MESSAGE: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
backgroundWorkerId: z.string(),
|
||||
data: BackgroundWorkerServerMessages,
|
||||
}),
|
||||
},
|
||||
};
|
||||
|
||||
export const ProdWorkerToCoordinatorMessages = {
|
||||
LOG: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
text: z.string(),
|
||||
}),
|
||||
callback: z.void(),
|
||||
},
|
||||
INDEX_TASKS: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
tasks: TaskResource.array(),
|
||||
packageVersion: z.string(),
|
||||
}),
|
||||
callback: z.discriminatedUnion("success", [
|
||||
z.object({
|
||||
success: z.literal(false),
|
||||
}),
|
||||
z.object({
|
||||
success: z.literal(true),
|
||||
}),
|
||||
]),
|
||||
},
|
||||
READY_FOR_EXECUTION: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptId: z.string(),
|
||||
}),
|
||||
},
|
||||
TASK_HEARTBEAT: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptFriendlyId: z.string(),
|
||||
}),
|
||||
},
|
||||
WAIT_FOR_BATCH: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
id: z.string(),
|
||||
runs: z.string().array(),
|
||||
}),
|
||||
},
|
||||
WAIT_FOR_DURATION: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
ms: z.number(),
|
||||
}),
|
||||
callback: z.discriminatedUnion("success", [
|
||||
z.object({
|
||||
success: z.literal(false),
|
||||
}),
|
||||
z.object({
|
||||
success: z.literal(true),
|
||||
}),
|
||||
]),
|
||||
},
|
||||
WAIT_FOR_TASK: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
id: z.string(),
|
||||
}),
|
||||
},
|
||||
};
|
||||
|
||||
export const CoordinatorToProdWorkerMessages = {
|
||||
RESUME: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
attemptId: z.string(),
|
||||
image: z.string(),
|
||||
completions: TaskRunExecutionResult.array(),
|
||||
executions: TaskRunExecution.array(),
|
||||
}),
|
||||
},
|
||||
EXECUTE_TASK_RUN: {
|
||||
message: z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
executionPayload: ProdTaskRunExecutionPayload,
|
||||
}),
|
||||
callback: z.discriminatedUnion("success", [
|
||||
z.object({
|
||||
success: z.literal(false),
|
||||
}),
|
||||
z.object({
|
||||
success: z.literal(true),
|
||||
completion: TaskRunExecutionResult,
|
||||
}),
|
||||
]),
|
||||
},
|
||||
};
|
||||
|
||||
export const ProdWorkerSocketData = z.object({
|
||||
cliPackageVersion: z.string(),
|
||||
contentHash: z.string(),
|
||||
projectRef: z.string(),
|
||||
envId: z.string(),
|
||||
attemptId: z.string(),
|
||||
podName: z.string(),
|
||||
});
|
||||
|
||||
@@ -1,2 +1 @@
|
||||
export * from "./socketIo";
|
||||
export * from "./utils";
|
||||
|
||||
@@ -1,129 +0,0 @@
|
||||
import {
|
||||
TaskResource,
|
||||
ProdTaskRunExecutionPayload,
|
||||
TaskRunExecutionResult,
|
||||
TaskRunExecution,
|
||||
ProdTaskRunExecution,
|
||||
} from "../schemas";
|
||||
|
||||
export type VersionedMessage<TMessage> = { version: "v1" } & TMessage;
|
||||
|
||||
// provider <--> platform
|
||||
export interface ProviderClientToServerEvents {
|
||||
LOG: (message: VersionedMessage<{ data: string }>) => void;
|
||||
}
|
||||
|
||||
export interface ProviderServerToClientEvents {
|
||||
HEALTH: (message: VersionedMessage<{}>, callback: (ack: { status: "ok" }) => void) => void;
|
||||
INDEX: (
|
||||
message: VersionedMessage<{ imageTag: string; contentHash: string; envId: string }>
|
||||
) => void;
|
||||
RESTORE: (
|
||||
message: VersionedMessage<{
|
||||
id: string;
|
||||
attemptId: string;
|
||||
type: "DOCKER" | "KUBERNETES";
|
||||
location: string;
|
||||
reason?: string;
|
||||
}>
|
||||
) => void;
|
||||
DELETE: (
|
||||
message: VersionedMessage<{ name: string }>,
|
||||
callback: (ack: { message: string }) => void
|
||||
) => void;
|
||||
GET: (message: VersionedMessage<{ name: string }>) => void;
|
||||
}
|
||||
|
||||
// coordinator <--> prod worker
|
||||
export interface ProdWorkerToCoordinatorEvents {
|
||||
LOG: (message: VersionedMessage<{ text: string }>, callback: () => {}) => void;
|
||||
INDEX_TASKS: (
|
||||
message: VersionedMessage<{
|
||||
tasks: TaskResource[];
|
||||
packageVersion: string;
|
||||
}>,
|
||||
callback: (params: { success: boolean }) => {}
|
||||
) => void;
|
||||
READY_FOR_EXECUTION: (message: VersionedMessage<{ attemptId: string }>) => void;
|
||||
TASK_HEARTBEAT: (message: VersionedMessage<{ attemptFriendlyId: string }>) => void;
|
||||
WAIT_FOR_BATCH: (message: VersionedMessage<{ id: string; runs: string[] }>) => void;
|
||||
WAIT_FOR_DURATION: (
|
||||
message: VersionedMessage<{ ms: number }>,
|
||||
callback: (ack: { success: boolean }) => void
|
||||
) => void;
|
||||
WAIT_FOR_TASK: (message: VersionedMessage<{ id: string }>) => void;
|
||||
}
|
||||
|
||||
export interface CoordinatorToProdWorkerEvents {
|
||||
RESUME: (
|
||||
message: VersionedMessage<{
|
||||
attemptId: string;
|
||||
image: string;
|
||||
completions: TaskRunExecutionResult[];
|
||||
executions: TaskRunExecution[];
|
||||
}>
|
||||
) => void;
|
||||
EXECUTE_TASK_RUN: (
|
||||
message: VersionedMessage<{ payload: ProdTaskRunExecutionPayload }>,
|
||||
callback: (ack: { completion: TaskRunExecutionResult }) => void
|
||||
) => void;
|
||||
}
|
||||
|
||||
export interface ProdWorkerSocketData {
|
||||
cliPackageVersion: string;
|
||||
contentHash: string;
|
||||
projectRef: string;
|
||||
envId: string;
|
||||
attemptId: string;
|
||||
podName: string;
|
||||
}
|
||||
|
||||
// coordinator <--> platform
|
||||
export interface CoordinatorToPlatformEvents {
|
||||
LOG: (message: VersionedMessage<{ metadata: any; text: string }>) => void;
|
||||
CREATE_WORKER: (
|
||||
message: VersionedMessage<{
|
||||
projectRef: string;
|
||||
envId: string;
|
||||
metadata: {
|
||||
cliPackageVersion: string;
|
||||
contentHash: string;
|
||||
packageVersion: string;
|
||||
tasks: TaskResource[];
|
||||
};
|
||||
}>,
|
||||
callback: (ack: { success: boolean }) => void
|
||||
) => void;
|
||||
READY_FOR_EXECUTION: (
|
||||
message: VersionedMessage<{ attemptId: string }>,
|
||||
callback: (
|
||||
ack: { success: false } | { success: true; payload: ProdTaskRunExecutionPayload }
|
||||
) => void
|
||||
) => void;
|
||||
TASK_RUN_COMPLETED: (
|
||||
message: VersionedMessage<{
|
||||
execution: ProdTaskRunExecution;
|
||||
completion: TaskRunExecutionResult;
|
||||
}>
|
||||
) => void;
|
||||
TASK_HEARTBEAT: (message: VersionedMessage<{ attemptFriendlyId: string }>) => void;
|
||||
CHECKPOINT_CREATED: (
|
||||
message: VersionedMessage<{
|
||||
attemptId: string;
|
||||
docker: boolean;
|
||||
location: string;
|
||||
reason?: string;
|
||||
}>
|
||||
) => void;
|
||||
}
|
||||
|
||||
export interface PlatformToCoordinatorEvents {
|
||||
RESUME: (
|
||||
message: VersionedMessage<{
|
||||
attemptId: string;
|
||||
image: string;
|
||||
completions: TaskRunExecutionResult[];
|
||||
executions: TaskRunExecution[];
|
||||
}>
|
||||
) => void;
|
||||
}
|
||||
@@ -1,7 +1,11 @@
|
||||
import { z } from "zod";
|
||||
|
||||
export type ZodMessageValueSchema<TDiscriminatedUnion extends z.ZodDiscriminatedUnion<any, any>> =
|
||||
| z.ZodFirstPartySchemaTypes
|
||||
| TDiscriminatedUnion;
|
||||
|
||||
export interface ZodMessageCatalogSchema {
|
||||
[key: string]: z.ZodFirstPartySchemaTypes | z.ZodDiscriminatedUnion<any, any>;
|
||||
[key: string]: ZodMessageValueSchema<any>;
|
||||
}
|
||||
|
||||
export type ZodMessageHandlers<TCatalogSchema extends ZodMessageCatalogSchema> = Partial<{
|
||||
@@ -31,7 +35,7 @@ const messageSchema = z.object({
|
||||
payload: z.unknown(),
|
||||
});
|
||||
|
||||
interface EventEmitterLike {
|
||||
export interface EventEmitterLike {
|
||||
on(eventName: string | symbol, listener: (...args: any[]) => void): this;
|
||||
}
|
||||
|
||||
@@ -102,6 +106,7 @@ export class ZodMessageHandler<TMessageCatalog extends ZodMessageCatalogSchema>
|
||||
|
||||
let ack;
|
||||
|
||||
// FIXME: this only works if the message doesn't have genuine payload prop
|
||||
if ("payload" in message) {
|
||||
ack = await this.handleMessage({ type: eventName, ...message });
|
||||
} else {
|
||||
|
||||
@@ -1,78 +1,95 @@
|
||||
import { DisconnectReason, Namespace, Server, Socket } from "socket.io";
|
||||
import { ZodMessageSender } from "./zodMessageHandler";
|
||||
import {
|
||||
ZodMessageCatalogSchema,
|
||||
ZodMessageHandlerOptions,
|
||||
MessageCatalogToSocketIoEvents,
|
||||
ZodMessageHandler,
|
||||
ZodMessageSender,
|
||||
} from "./zodMessageHandler";
|
||||
ZodMessageCatalogToSocketIoEvents,
|
||||
ZodSocketMessageCatalogSchema,
|
||||
ZodSocketMessageHandler,
|
||||
ZodSocketMessageHandlers,
|
||||
} from "./zodSocket";
|
||||
import { DefaultEventsMap, EventsMap } from "socket.io/dist/typed-events";
|
||||
import { z } from "zod";
|
||||
|
||||
interface ExtendedError extends Error {
|
||||
data?: any;
|
||||
}
|
||||
|
||||
export type ZodSocket<
|
||||
TClientMessages extends ZodMessageCatalogSchema,
|
||||
TServerMessages extends ZodMessageCatalogSchema,
|
||||
export type ZodNamespaceSocket<
|
||||
TClientMessages extends ZodSocketMessageCatalogSchema,
|
||||
TServerMessages extends ZodSocketMessageCatalogSchema,
|
||||
TServerSideEvents extends EventsMap = DefaultEventsMap,
|
||||
TSocketData extends z.ZodObject<any, any, any> = any,
|
||||
> = Socket<
|
||||
MessageCatalogToSocketIoEvents<TClientMessages>,
|
||||
MessageCatalogToSocketIoEvents<TServerMessages>
|
||||
ZodMessageCatalogToSocketIoEvents<TClientMessages>,
|
||||
ZodMessageCatalogToSocketIoEvents<TServerMessages>,
|
||||
TServerSideEvents,
|
||||
z.infer<TSocketData>
|
||||
>;
|
||||
|
||||
interface ZodNamespaceOptions<
|
||||
TClientMessages extends ZodMessageCatalogSchema,
|
||||
TServerMessages extends ZodMessageCatalogSchema,
|
||||
TClientMessages extends ZodSocketMessageCatalogSchema,
|
||||
TServerMessages extends ZodSocketMessageCatalogSchema,
|
||||
TServerSideEvents extends EventsMap = DefaultEventsMap,
|
||||
TSocketData extends z.ZodObject<any, any, any> = any,
|
||||
> {
|
||||
io: Server;
|
||||
name: string;
|
||||
clientMessages: TClientMessages;
|
||||
serverMessages: TServerMessages;
|
||||
messageHandler?: ZodMessageHandlerOptions<TClientMessages>["messages"];
|
||||
socketData?: TSocketData;
|
||||
handlers?: ZodSocketMessageHandlers<TClientMessages>;
|
||||
authToken?: string;
|
||||
preAuth?: (
|
||||
socket: ZodSocket<TClientMessages, TServerMessages>,
|
||||
next: (err?: ExtendedError) => void
|
||||
) => void;
|
||||
socket: ZodNamespaceSocket<TClientMessages, TServerMessages, TServerSideEvents, TSocketData>,
|
||||
next: (err?: ExtendedError) => void,
|
||||
logger: (...args: any[]) => void
|
||||
) => Promise<void>;
|
||||
postAuth?: (
|
||||
socket: ZodSocket<TClientMessages, TServerMessages>,
|
||||
next: (err?: ExtendedError) => void
|
||||
) => void;
|
||||
socket: ZodNamespaceSocket<TClientMessages, TServerMessages, TServerSideEvents, TSocketData>,
|
||||
next: (err?: ExtendedError) => void,
|
||||
logger: (...args: any[]) => void
|
||||
) => Promise<void>;
|
||||
onConnection?: (
|
||||
socket: ZodSocket<TClientMessages, TServerMessages>,
|
||||
handler: ZodMessageHandler<TClientMessages>,
|
||||
socket: ZodNamespaceSocket<TClientMessages, TServerMessages, TServerSideEvents, TSocketData>,
|
||||
handler: ZodSocketMessageHandler<TClientMessages>,
|
||||
sender: ZodMessageSender<TServerMessages>,
|
||||
logger: (...args: any[]) => void
|
||||
) => Promise<void>;
|
||||
onDisconnect?: (
|
||||
socket: ZodSocket<TClientMessages, TServerMessages>,
|
||||
socket: ZodNamespaceSocket<TClientMessages, TServerMessages, TServerSideEvents, TSocketData>,
|
||||
reason: DisconnectReason,
|
||||
description: any,
|
||||
logger: (...args: any[]) => void
|
||||
) => Promise<void>;
|
||||
onError?: (
|
||||
socket: ZodSocket<TClientMessages, TServerMessages>,
|
||||
socket: ZodNamespaceSocket<TClientMessages, TServerMessages, TServerSideEvents, TSocketData>,
|
||||
err: Error,
|
||||
logger: (...args: any[]) => void
|
||||
) => Promise<void>;
|
||||
}
|
||||
|
||||
export class ZodNamespace<
|
||||
TClientMessages extends ZodMessageCatalogSchema,
|
||||
TServerMessages extends ZodMessageCatalogSchema,
|
||||
TClientMessages extends ZodSocketMessageCatalogSchema,
|
||||
TServerMessages extends ZodSocketMessageCatalogSchema,
|
||||
TSocketData extends z.ZodObject<any, any, any> = any,
|
||||
TServerSideEvents extends EventsMap = DefaultEventsMap,
|
||||
> {
|
||||
#handler: ZodMessageHandler<TClientMessages>;
|
||||
#handler: ZodSocketMessageHandler<TClientMessages>;
|
||||
sender: ZodMessageSender<TServerMessages>;
|
||||
|
||||
io: Server;
|
||||
namespace: Namespace<
|
||||
MessageCatalogToSocketIoEvents<TClientMessages>,
|
||||
MessageCatalogToSocketIoEvents<TServerMessages>
|
||||
ZodMessageCatalogToSocketIoEvents<TClientMessages>,
|
||||
ZodMessageCatalogToSocketIoEvents<TServerMessages>,
|
||||
TServerSideEvents,
|
||||
z.infer<TSocketData>
|
||||
>;
|
||||
|
||||
constructor(opts: ZodNamespaceOptions<TClientMessages, TServerMessages>) {
|
||||
this.#handler = new ZodMessageHandler({
|
||||
constructor(
|
||||
opts: ZodNamespaceOptions<TClientMessages, TServerMessages, TServerSideEvents, TSocketData>
|
||||
) {
|
||||
this.#handler = new ZodSocketMessageHandler({
|
||||
schema: opts.clientMessages,
|
||||
messages: opts.messageHandler,
|
||||
handlers: opts.handlers,
|
||||
});
|
||||
|
||||
this.io = opts.io;
|
||||
@@ -95,7 +112,13 @@ export class ZodNamespace<
|
||||
});
|
||||
|
||||
if (opts.preAuth) {
|
||||
this.namespace.use(opts.preAuth);
|
||||
this.namespace.use(async (socket, next) => {
|
||||
const logger = createLogger(`[${opts.name}][${socket.id}][preAuth]`);
|
||||
|
||||
if (typeof opts.preAuth === "function") {
|
||||
await opts.preAuth(socket, next, logger);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
if (opts.authToken) {
|
||||
@@ -121,7 +144,13 @@ export class ZodNamespace<
|
||||
}
|
||||
|
||||
if (opts.postAuth) {
|
||||
this.namespace.use(opts.postAuth);
|
||||
this.namespace.use(async (socket, next) => {
|
||||
const logger = createLogger(`[${opts.name}][${socket.id}][postAuth]`);
|
||||
|
||||
if (typeof opts.postAuth === "function") {
|
||||
await opts.postAuth(socket, next, logger);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
this.namespace.on("connection", async (socket) => {
|
||||
@@ -151,6 +180,10 @@ export class ZodNamespace<
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
fetchSockets() {
|
||||
return this.namespace.fetchSockets();
|
||||
}
|
||||
}
|
||||
|
||||
function createLogger(prefix: string) {
|
||||
|
||||
@@ -0,0 +1,359 @@
|
||||
import { io, Socket } from "socket.io-client";
|
||||
import { z } from "zod";
|
||||
import { EventEmitterLike, ZodMessageValueSchema } from "./zodMessageHandler";
|
||||
|
||||
export interface ZodSocketMessageCatalogSchema {
|
||||
[key: string]:
|
||||
| {
|
||||
message: ZodMessageValueSchema<any>;
|
||||
}
|
||||
| {
|
||||
message: ZodMessageValueSchema<any>;
|
||||
callback?: ZodMessageValueSchema<any>;
|
||||
};
|
||||
}
|
||||
|
||||
export type ZodMessageCatalogToSocketIoEvents<TCatalog extends ZodSocketMessageCatalogSchema> = {
|
||||
[K in keyof TCatalog]: SocketMessageHasCallback<TCatalog, K> extends true
|
||||
? (
|
||||
message: z.infer<GetSocketMessageSchema<TCatalog, K>>,
|
||||
callback: (ack: z.infer<GetSocketCallbackSchema<TCatalog, K>>) => void
|
||||
) => void
|
||||
: (message: z.infer<GetSocketMessageSchema<TCatalog, K>>) => void;
|
||||
};
|
||||
|
||||
type GetSocketMessageSchema<
|
||||
TRPCCatalog extends ZodSocketMessageCatalogSchema,
|
||||
TMessageType extends keyof TRPCCatalog,
|
||||
> = TRPCCatalog[TMessageType]["message"];
|
||||
|
||||
export type InferSocketMessageSchema<
|
||||
TRPCCatalog extends ZodSocketMessageCatalogSchema,
|
||||
TMessageType extends keyof TRPCCatalog,
|
||||
> = z.infer<GetSocketMessageSchema<TRPCCatalog, TMessageType>>;
|
||||
|
||||
type GetSocketCallbackSchema<
|
||||
TRPCCatalog extends ZodSocketMessageCatalogSchema,
|
||||
TMessageType extends keyof TRPCCatalog,
|
||||
> = TRPCCatalog[TMessageType] extends { callback: any }
|
||||
? TRPCCatalog[TMessageType]["callback"]
|
||||
: never;
|
||||
|
||||
export type InferSocketCallbackSchema<
|
||||
TRPCCatalog extends ZodSocketMessageCatalogSchema,
|
||||
TMessageType extends keyof TRPCCatalog,
|
||||
> = z.infer<GetSocketCallbackSchema<TRPCCatalog, TMessageType>>;
|
||||
|
||||
type SocketMessageHasCallback<
|
||||
TRPCCatalog extends ZodSocketMessageCatalogSchema,
|
||||
TMessageType extends keyof TRPCCatalog,
|
||||
> = GetSocketCallbackSchema<TRPCCatalog, TMessageType> extends never ? false : true;
|
||||
|
||||
export type ZodSocketMessageHandlers<TCatalogSchema extends ZodSocketMessageCatalogSchema> =
|
||||
Partial<{
|
||||
[K in keyof TCatalogSchema]: (
|
||||
payload: z.infer<GetSocketMessageSchema<TCatalogSchema, K>>
|
||||
) => Promise<
|
||||
SocketMessageHasCallback<TCatalogSchema, K> extends true
|
||||
? z.input<GetSocketCallbackSchema<TCatalogSchema, K>>
|
||||
: void
|
||||
>;
|
||||
}>;
|
||||
|
||||
export type ZodSocketMessageHandlerOptions<TMessageCatalog extends ZodSocketMessageCatalogSchema> =
|
||||
{
|
||||
schema: TMessageCatalog;
|
||||
handlers?: ZodSocketMessageHandlers<TMessageCatalog>;
|
||||
};
|
||||
|
||||
type MessageFromSocketSchema<
|
||||
K extends keyof TMessageCatalog,
|
||||
TMessageCatalog extends ZodSocketMessageCatalogSchema,
|
||||
> = {
|
||||
type: K;
|
||||
payload: z.input<GetSocketMessageSchema<TMessageCatalog, K>>;
|
||||
};
|
||||
|
||||
type MessagesFromSocketCatalog<TMessageCatalog extends ZodSocketMessageCatalogSchema> = {
|
||||
[K in keyof TMessageCatalog]: MessageFromSocketSchema<K, TMessageCatalog>;
|
||||
}[keyof TMessageCatalog];
|
||||
|
||||
const messageSchema = z.object({
|
||||
version: z.literal("v1").default("v1"),
|
||||
type: z.string(),
|
||||
payload: z.unknown(),
|
||||
});
|
||||
|
||||
export class ZodSocketMessageHandler<TRPCCatalog extends ZodSocketMessageCatalogSchema> {
|
||||
#schema: TRPCCatalog;
|
||||
#handlers: ZodSocketMessageHandlers<TRPCCatalog> | undefined;
|
||||
|
||||
constructor(options: ZodSocketMessageHandlerOptions<TRPCCatalog>) {
|
||||
this.#schema = options.schema;
|
||||
this.#handlers = options.handlers;
|
||||
}
|
||||
|
||||
public async handleMessage(message: unknown) {
|
||||
const parsedMessage = this.parseMessage(message);
|
||||
|
||||
if (!this.#handlers) {
|
||||
throw new Error("No handlers provided");
|
||||
}
|
||||
|
||||
const handler = this.#handlers[parsedMessage.type];
|
||||
|
||||
if (!handler) {
|
||||
console.error(`No handler for message type: ${String(parsedMessage.type)}`);
|
||||
return;
|
||||
}
|
||||
|
||||
const ack = await handler(parsedMessage.payload);
|
||||
|
||||
return ack;
|
||||
}
|
||||
|
||||
public parseMessage(message: unknown): MessagesFromSocketCatalog<TRPCCatalog> {
|
||||
const parsedMessage = messageSchema.safeParse(message);
|
||||
|
||||
if (!parsedMessage.success) {
|
||||
throw new Error(`Failed to parse message: ${JSON.stringify(parsedMessage.error)}`);
|
||||
}
|
||||
|
||||
const schema = this.#schema[parsedMessage.data.type]["message"];
|
||||
|
||||
if (!schema) {
|
||||
throw new Error(`Unknown message type: ${parsedMessage.data.type}`);
|
||||
}
|
||||
|
||||
const parsedPayload = schema.safeParse(parsedMessage.data.payload);
|
||||
|
||||
if (!parsedPayload.success) {
|
||||
throw new Error(`Failed to parse message payload: ${JSON.stringify(parsedPayload.error)}`);
|
||||
}
|
||||
|
||||
return {
|
||||
type: parsedMessage.data.type,
|
||||
payload: parsedPayload.data,
|
||||
};
|
||||
}
|
||||
|
||||
public registerHandlers(emitter: EventEmitterLike, logger?: (...args: any[]) => void) {
|
||||
const log = logger ?? console.log;
|
||||
|
||||
if (!this.#handlers) {
|
||||
log("No handlers provided");
|
||||
return;
|
||||
}
|
||||
|
||||
for (const eventName of Object.keys(this.#handlers)) {
|
||||
emitter.on(eventName, async (message: any, callback?: any): Promise<void> => {
|
||||
log(`handling ${eventName}`, message);
|
||||
|
||||
let ack;
|
||||
|
||||
// FIXME: this only works if the message doesn't have genuine payload prop
|
||||
if ("payload" in message) {
|
||||
ack = await this.handleMessage({ type: eventName, ...message });
|
||||
} else {
|
||||
// Handle messages not sent by ZodMessageSender
|
||||
const { version, ...payload } = message;
|
||||
ack = await this.handleMessage({ type: eventName, version, payload });
|
||||
}
|
||||
|
||||
if (callback && typeof callback === "function") {
|
||||
callback(ack);
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
export type ZodSocketMessageSenderOptions<TMessageCatalog extends ZodSocketMessageCatalogSchema> = {
|
||||
schema: TMessageCatalog;
|
||||
socket: ZodSocket<any, TMessageCatalog>;
|
||||
};
|
||||
|
||||
type GetSocketMessagesWithCallback<TMessageCatalog extends ZodSocketMessageCatalogSchema> = {
|
||||
[K in keyof TMessageCatalog]: SocketMessageHasCallback<TMessageCatalog, K> extends true
|
||||
? K
|
||||
: never;
|
||||
}[keyof TMessageCatalog];
|
||||
|
||||
type GetSocketMessagesWithoutCallback<TMessageCatalog extends ZodSocketMessageCatalogSchema> = {
|
||||
[K in keyof TMessageCatalog]: SocketMessageHasCallback<TMessageCatalog, K> extends true
|
||||
? never
|
||||
: K;
|
||||
}[keyof TMessageCatalog];
|
||||
|
||||
export class ZodSocketMessageSender<TMessageCatalog extends ZodSocketMessageCatalogSchema> {
|
||||
#schema: TMessageCatalog;
|
||||
#socket: ZodSocket<any, TMessageCatalog>;
|
||||
|
||||
constructor(options: ZodSocketMessageSenderOptions<TMessageCatalog>) {
|
||||
this.#schema = options.schema;
|
||||
this.#socket = options.socket;
|
||||
}
|
||||
|
||||
public send<K extends GetSocketMessagesWithoutCallback<TMessageCatalog>>(
|
||||
type: K,
|
||||
payload: z.input<GetSocketMessageSchema<TMessageCatalog, K>>
|
||||
): void {
|
||||
const schema = this.#schema[type]["message"];
|
||||
|
||||
if (!schema) {
|
||||
throw new Error(`Unknown message type: ${type as string}`);
|
||||
}
|
||||
|
||||
const parsedPayload = schema.safeParse(payload);
|
||||
|
||||
if (!parsedPayload.success) {
|
||||
throw new Error(`Failed to parse message payload: ${JSON.stringify(parsedPayload.error)}`);
|
||||
}
|
||||
|
||||
// @ts-expect-error
|
||||
this.#socket.emit(type, { payload, version: "v1" });
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
public async sendWithAck<K extends GetSocketMessagesWithCallback<TMessageCatalog>>(
|
||||
type: K,
|
||||
payload: z.input<GetSocketMessageSchema<TMessageCatalog, K>>
|
||||
): Promise<z.infer<GetSocketCallbackSchema<TMessageCatalog, K>>> {
|
||||
const schema = this.#schema[type]["message"];
|
||||
|
||||
if (!schema) {
|
||||
throw new Error(`Unknown message type: ${type as string}`);
|
||||
}
|
||||
|
||||
const parsedPayload = schema.safeParse(payload);
|
||||
|
||||
if (!parsedPayload.success) {
|
||||
throw new Error(`Failed to parse message payload: ${JSON.stringify(parsedPayload.error)}`);
|
||||
}
|
||||
|
||||
// @ts-expect-error
|
||||
const callbackResult = await this.#socket.emitWithAck(type, { payload, version: "v1" });
|
||||
|
||||
return callbackResult;
|
||||
}
|
||||
}
|
||||
|
||||
export type ZodSocket<
|
||||
TListenEvents extends ZodSocketMessageCatalogSchema,
|
||||
TEmitEvents extends ZodSocketMessageCatalogSchema,
|
||||
> = Socket<
|
||||
ZodMessageCatalogToSocketIoEvents<TListenEvents>,
|
||||
ZodMessageCatalogToSocketIoEvents<TEmitEvents>
|
||||
>;
|
||||
|
||||
interface ZodSocketConnectionOptions<
|
||||
TClientMessages extends ZodSocketMessageCatalogSchema,
|
||||
TServerMessages extends ZodSocketMessageCatalogSchema,
|
||||
> {
|
||||
host: string;
|
||||
port: number;
|
||||
namespace: string;
|
||||
clientMessages: TClientMessages;
|
||||
serverMessages: TServerMessages;
|
||||
extraHeaders?: {
|
||||
[header: string]: string;
|
||||
};
|
||||
handlers?: ZodSocketMessageHandlers<TServerMessages>;
|
||||
authToken?: string;
|
||||
onConnection?: (
|
||||
socket: ZodSocket<TServerMessages, TClientMessages>,
|
||||
handler: ZodSocketMessageHandler<TServerMessages>,
|
||||
sender: ZodSocketMessageSender<TClientMessages>,
|
||||
logger: (...args: any[]) => void
|
||||
) => Promise<void>;
|
||||
onDisconnect?: (
|
||||
socket: ZodSocket<TServerMessages, TClientMessages>,
|
||||
reason: Socket.DisconnectReason,
|
||||
description: any,
|
||||
logger: (...args: any[]) => void
|
||||
) => Promise<void>;
|
||||
onError?: (
|
||||
socket: ZodSocket<TServerMessages, TClientMessages>,
|
||||
err: Error,
|
||||
logger: (...args: any[]) => void
|
||||
) => Promise<void>;
|
||||
}
|
||||
|
||||
export class ZodSocketConnection<
|
||||
TClientMessages extends ZodSocketMessageCatalogSchema,
|
||||
TServerMessages extends ZodSocketMessageCatalogSchema,
|
||||
> {
|
||||
#sender: ZodSocketMessageSender<TClientMessages>;
|
||||
socket: ZodSocket<TServerMessages, TClientMessages>;
|
||||
|
||||
#handler: ZodSocketMessageHandler<TServerMessages>;
|
||||
#logger: (...args: any[]) => void;
|
||||
|
||||
constructor(opts: ZodSocketConnectionOptions<TClientMessages, TServerMessages>) {
|
||||
this.socket = io(`ws://${opts.host}:${opts.port}/${opts.namespace}`, {
|
||||
transports: ["websocket"],
|
||||
auth: {
|
||||
token: opts.authToken,
|
||||
},
|
||||
extraHeaders: opts.extraHeaders,
|
||||
});
|
||||
|
||||
this.#logger = createLogger(`[${opts.namespace}][${this.socket.id}]`);
|
||||
|
||||
this.#handler = new ZodSocketMessageHandler({
|
||||
schema: opts.serverMessages,
|
||||
handlers: opts.handlers,
|
||||
});
|
||||
this.#handler.registerHandlers(this.socket, this.#logger);
|
||||
|
||||
this.#sender = new ZodSocketMessageSender({
|
||||
schema: opts.clientMessages,
|
||||
socket: this.socket,
|
||||
});
|
||||
|
||||
this.socket.on("connect_error", async (error) => {
|
||||
this.#logger(`connect_error: ${error}`);
|
||||
|
||||
if (opts.onError) {
|
||||
await opts.onError(this.socket, error, this.#logger);
|
||||
}
|
||||
});
|
||||
|
||||
this.socket.on("connect", async () => {
|
||||
this.#logger("connect");
|
||||
|
||||
if (opts.onConnection) {
|
||||
await opts.onConnection(this.socket, this.#handler, this.#sender, this.#logger);
|
||||
}
|
||||
});
|
||||
|
||||
this.socket.on("disconnect", async (reason, description) => {
|
||||
this.#logger("disconnect");
|
||||
|
||||
if (opts.onDisconnect) {
|
||||
await opts.onDisconnect(this.socket, reason, description, this.#logger);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
close() {
|
||||
this.socket.close();
|
||||
}
|
||||
|
||||
connect() {
|
||||
this.socket.connect();
|
||||
}
|
||||
|
||||
get send() {
|
||||
return this.#sender.send.bind(this.#sender);
|
||||
}
|
||||
|
||||
get sendWithAck() {
|
||||
return this.#sender.sendWithAck.bind(this.#sender);
|
||||
}
|
||||
}
|
||||
|
||||
function createLogger(prefix: string) {
|
||||
return (...args: any[]) => console.log(prefix, ...args);
|
||||
}
|
||||
Generated
+3
-1
@@ -1208,6 +1208,7 @@ importers:
|
||||
jest: ^29.6.2
|
||||
rimraf: ^3.0.2
|
||||
socket.io: ^4.7.4
|
||||
socket.io-client: ^4.7.4
|
||||
ts-jest: ^29.1.1
|
||||
tsup: ^8.0.1
|
||||
typescript: ^5.3.0
|
||||
@@ -1232,6 +1233,7 @@ importers:
|
||||
'@opentelemetry/semantic-conventions': 1.21.0
|
||||
humanize-duration: 3.27.3
|
||||
socket.io: 4.7.4
|
||||
socket.io-client: 4.7.4
|
||||
ulidx: 2.2.1
|
||||
zod: 3.22.3
|
||||
zod-error: 1.5.0
|
||||
@@ -36799,7 +36801,7 @@ packages:
|
||||
dependencies:
|
||||
bs-logger: 0.2.6
|
||||
fast-json-stable-stringify: 2.1.0
|
||||
jest: 29.6.2_@types+node@18.15.13
|
||||
jest: 29.6.2_@types+node@18.17.1
|
||||
jest-util: 29.6.2
|
||||
json5: 2.2.3
|
||||
lodash.memoize: 4.1.2
|
||||
|
||||
@@ -1,2 +1,2 @@
|
||||
TRIGGER_API_KEY=
|
||||
TRIGGER_SECRET_KEY=
|
||||
TRIGGER_API_URL=
|
||||
Reference in New Issue
Block a user