Compare commits

...

18 Commits

Author SHA1 Message Date
nicktrn ac524f0678 pass otlp env var to prod worker
🚢 Publish Container Images (dev) / build (coordinator) (push) Has been cancelled
🚢 Publish Container Images (dev) / build (kubernetes-provider) (push) Has been cancelled
2024-03-05 18:33:55 +00:00
nicktrn bd7535b4e7 set otlp endpoint for on runs 2024-03-05 14:07:44 +00:00
nicktrn 3ab64ea935 increase prod worker timeout 2024-03-05 13:59:49 +00:00
nicktrn 9846e74d74 ensure attempt id is always set 2024-03-05 13:59:35 +00:00
nicktrn 61b7f5cd2a set task run label 2024-03-05 13:49:25 +00:00
nicktrn 2f318cad41 fix coordinator build 2024-03-05 12:58:47 +00:00
nicktrn 7b479da51d update env example to new v3 key var 2024-03-05 12:55:26 +00:00
nicktrn ec1b704003 complete socket.io types to schemas migration 2024-03-05 12:54:53 +00:00
nicktrn 3750bca5c2 update injected secret key env var name 2024-03-05 12:52:21 +00:00
nicktrn 0fab371ee4 remove unused types 2024-03-05 09:17:32 +00:00
nicktrn a215359dba update docker actions and fix builds 2024-03-05 09:11:04 +00:00
nicktrn c3f18fd38a set otlp endpoint 2024-03-04 19:24:01 +00:00
nicktrn c41cf7e9cf Merge branch 'main' into v3/provider-updates 2024-03-04 19:10:18 +00:00
nicktrn e614d29bc5 update k8s provider task ops 2024-03-04 19:05:56 +00:00
nicktrn 8997363e72 use shared provider shell 2024-03-04 18:38:26 +00:00
nicktrn 01efca816b use zod socket for shared queue 2024-03-04 18:10:37 +00:00
nicktrn 735182e3e2 start using zod socket 2024-03-04 18:01:51 +00:00
nicktrn 6a02d6aa53 add zod socket 2024-03-04 17:24:35 +00:00
28 changed files with 1283 additions and 1061 deletions
+3 -3
View File
@@ -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 }}
+2 -2
View File
@@ -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" ]
+1 -1
View File
@@ -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
View File
@@ -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() {
+2 -2
View File
@@ -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" ]
+1 -1
View File
@@ -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",
+3 -263
View File
@@ -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(),
});
+2 -2
View File
@@ -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" ]
+2 -4
View File
@@ -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": "",
+39 -207
View File
@@ -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,
},
{
+5 -5
View File
@@ -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: {
+93 -107
View File
@@ -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,
});
+1 -1
View File
@@ -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
View File
@@ -1,2 +1,3 @@
export * from "./http";
export * from "./logger";
export * from "./provider";
+220
View File
@@ -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
View File
@@ -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);",
},
});
+1
View File
@@ -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"
+1
View File
@@ -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";
+280 -80
View File
@@ -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
View File
@@ -1,2 +1 @@
export * from "./socketIo";
export * from "./utils";
-129
View File
@@ -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;
}
+7 -2
View File
@@ -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 {
+67 -34
View File
@@ -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) {
+359
View File
@@ -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);
}
+3 -1
View File
@@ -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 -1
View File
@@ -1,2 +1,2 @@
TRIGGER_API_KEY=
TRIGGER_SECRET_KEY=
TRIGGER_API_URL=