diff --git a/.github/workflows/publish-dev.yml b/.github/workflows/publish-infra.yml similarity index 64% rename from .github/workflows/publish-dev.yml rename to .github/workflows/publish-infra.yml index 34afea272..ec67436a6 100644 --- a/.github/workflows/publish-dev.yml +++ b/.github/workflows/publish-infra.yml @@ -1,9 +1,11 @@ -name: "🚢 Publish Container Images (dev)" +name: "🚢 Publish Infra Images" on: push: tags: - - "dev-*" + - "infra-dev-*" + - "infra-test-*" + - "infra-prod-*" paths: - ".github/workflows/publish.yml" - "packages/**" @@ -44,12 +46,21 @@ jobs: steps: - uses: actions/checkout@v4 - - name: Generate build ID + - name: Generate image reference id: prep + # WARNING: This step expects the workflow to have been triggered by a specific tag format of: infra-${env}-* run: | + env=$(echo ${{ github.ref_name }} | cut -d- -f2) sha=${GITHUB_SHA::7} ts=$(date +%s) - echo "BUILD_ID=${{ matrix.package }}-${sha}-${ts}" >> "$GITHUB_OUTPUT" + if [[ "${{ matrix.package }}" == *-provider ]]; then + provider_type=$(echo ${{ matrix.package }} | cut -d- -f1) + repository=provider/${provider_type} + else + repository=${{ matrix.package }} + fi + echo "IMAGE_TAG=${env}-${sha}-${ts}" >> "$GITHUB_OUTPUT" + echo "REPOSITORY=${repository}" >> "$GITHUB_OUTPUT" - name: Set up Docker Buildx uses: docker/setup-buildx-action@v3 @@ -63,7 +74,7 @@ jobs: - name: 🚢 Build Container Image run: | - docker build -t dev_image -f ./apps/${{ matrix.package }}/Containerfile . + docker build -t infra_image -f ./apps/${{ matrix.package }}/Containerfile . # ..to push image - name: 🐙 Login to GitHub Container Registry @@ -75,9 +86,9 @@ jobs: - name: 🐙 Push to GitHub Container Registry run: | - docker tag dev_image $REGISTRY/$REPOSITORY:$IMAGE_TAG + docker tag infra_image $REGISTRY/$REPOSITORY:$IMAGE_TAG docker push $REGISTRY/$REPOSITORY:$IMAGE_TAG env: REGISTRY: ghcr.io/triggerdotdev - REPOSITORY: dev - IMAGE_TAG: ${{ steps.prep.outputs.BUILD_ID }} + REPOSITORY: ${{ steps.prep.outputs.REPOSITORY }} + IMAGE_TAG: ${{ steps.prep.outputs.IMAGE_TAG }} diff --git a/apps/coordinator/Containerfile b/apps/coordinator/Containerfile index f2ca3b3e6..cd301811b 100644 --- a/apps/coordinator/Containerfile +++ b/apps/coordinator/Containerfile @@ -1,6 +1,6 @@ # syntax=docker/dockerfile:labs -FROM node:18.18.2-bullseye-slim@sha256:21479df46c3173ee0cefc6b264928e10239152c4f74df872ca9369be01a245b7 AS node-18 +FROM node:18-bullseye-slim@sha256:a4edd54dcfdcacc8a4100fee71498e8671d99556a1acf5614539214a70092426 AS node-18 WORKDIR /app @@ -35,7 +35,6 @@ COPY --from=pruner --chown=node:node /app/out/full/ . COPY --from=dev-deps --chown=node:node /app/ . COPY --chown=node:node turbo.json turbo.json -RUN pnpm run -r --filter '@trigger.dev/core*' build RUN pnpm run -r --filter coordinator build:bundle FROM alpine AS cri-tools @@ -58,6 +57,4 @@ COPY --from=builder --chown=node:node /app/apps/coordinator/dist/index.mjs ./ind EXPOSE 8000 -USER node - CMD [ "/usr/bin/dumb-init", "--", "/usr/local/bin/node", "./index.mjs" ] diff --git a/apps/coordinator/package.json b/apps/coordinator/package.json index 4cf789c16..e089175f2 100644 --- a/apps/coordinator/package.json +++ b/apps/coordinator/package.json @@ -6,7 +6,7 @@ "main": "dist/index.cjs", "scripts": { "build": "npm run build:bundle", - "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:bundle": "esbuild src/index.ts --bundle --outfile=dist/index.mjs --platform=node --format=esm --target=esnext --banner:js=\"import { createRequire } from 'module';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", @@ -19,6 +19,7 @@ "@trigger.dev/core": "workspace:*", "@trigger.dev/core-apps": "workspace:*", "execa": "^8.0.1", + "nanoid": "^5.0.6", "prom-client": "^15.1.0", "socket.io": "^4.7.4", "socket.io-client": "^4.7.4" diff --git a/apps/coordinator/src/index.ts b/apps/coordinator/src/index.ts index 9f15f5a03..feb2722ba 100644 --- a/apps/coordinator/src/index.ts +++ b/apps/coordinator/src/index.ts @@ -1,6 +1,6 @@ -import { randomUUID } from "node:crypto"; import { createServer } from "node:http"; import { $ } from "execa"; +import { nanoid } from "nanoid"; import { Server } from "socket.io"; import { CoordinatorToPlatformMessages, @@ -18,8 +18,7 @@ collectDefaultMetrics(); const HTTP_SERVER_PORT = Number(process.env.HTTP_SERVER_PORT || 8020); const NODE_NAME = process.env.NODE_NAME || "coordinator"; -const REGISTRY_FQDN = process.env.REGISTRY_FQDN || "localhost:5000"; -const REPO_NAME = process.env.REPO_NAME || "checkpoints"; +const REGISTRY_HOST = process.env.REGISTRY_HOST || "localhost:5000"; const CHECKPOINT_PATH = process.env.CHECKPOINT_PATH || "/checkpoints"; const REGISTRY_TLS_VERIFY = process.env.REGISTRY_TLS_VERIFY === "false" ? "false" : "true"; @@ -35,12 +34,25 @@ type CheckpointerInitializeReturn = { willSimulate: boolean; }; +type CheckpointAndPushOptions = { + runId: string; + leaveRunning?: boolean; + projectRef: string; + deploymentVersion: string; +}; + +type CheckpointData = { + location: string; + docker: boolean; +}; + class Checkpointer { #initialized = false; #canCheckpoint = false; #dockerMode = !process.env.KUBERNETES_PORT; #logger = new SimpleLogger("[checkptr]"); + #abortControllers = new Map(); constructor(private opts = { forceSimulate: false }) {} @@ -51,26 +63,18 @@ class Checkpointer { this.#logger.log(`${this.#dockerMode ? "Docker" : "Kubernetes"} mode`); - if (this.opts.forceSimulate) { - this.#logger.log( - "Forced simulation enabled. Will simulate regardless of checkpoint support." - ); - } - - try { - await $`criu --version`; - } catch (error) { - this.#logger.error("No checkpoint support: Missing CRIU binary"); - if (this.#dockerMode) { - this.#logger.error("Will simulate instead"); - } - this.#canCheckpoint = false; - this.#initialized = true; - - return this.#getInitializeReturn(); - } - if (this.#dockerMode) { + try { + await $`criu --version`; + } catch (error) { + this.#logger.error("No checkpoint support: Missing CRIU binary"); + this.#logger.error("Will simulate instead"); + this.#canCheckpoint = false; + this.#initialized = true; + + return this.#getInitializeReturn(); + } + try { await $`docker checkpoint`; } catch (error) { @@ -81,12 +85,24 @@ class Checkpointer { this.#canCheckpoint = false; this.#initialized = true; + return this.#getInitializeReturn(); + } + } else { + try { + await $`buildah login --get-login ${REGISTRY_HOST}`; + } catch (error) { + this.#logger.error(`No checkpoint support: Not logged in to registry ${REGISTRY_HOST}`); + this.#canCheckpoint = false; + this.#initialized = true; + return this.#getInitializeReturn(); } } this.#logger.log( - `Full checkpoint support in ${this.#dockerMode ? "docker" : "kubernetes"} mode` + `Full checkpoint support${ + this.#dockerMode && this.opts.forceSimulate ? " with forced simulation enabled." : "!" + }` ); this.#initialized = true; @@ -102,7 +118,60 @@ class Checkpointer { }; } - async checkpointAndPush(podName: string, leaveRunning = false) { + #getImageRef(projectRef: string, deploymentVersion: string, shortCode: string) { + return `${REGISTRY_HOST}/trigger/${projectRef}:${deploymentVersion}.prod-${shortCode}`; + } + + #getExportLocation(projectRef: string, deploymentVersion: string, shortCode: string) { + const basename = `${projectRef}-${deploymentVersion}-${shortCode}`; + + if (this.#dockerMode) { + return basename; + } else { + return `${CHECKPOINT_PATH}/${basename}.tar`; + } + } + + async checkpointAndPush(opts: CheckpointAndPushOptions): Promise { + const start = performance.now(); + logger.log(`checkpointAndPush() start`, { start, opts }); + + const result = await this.#checkpointAndPush(opts); + + const end = performance.now(); + logger.log(`checkpointAndPush() end`, { + start, + end, + diff: end - start, + opts, + success: !!result, + }); + + return result; + } + + isCheckpointing(runId: string) { + return this.#abortControllers.has(runId); + } + + cancelCheckpoint(runId: string) { + const controller = this.#abortControllers.get(runId); + + if (!controller) { + logger.debug("Nothing to cancel", { runId }); + return; + } + + controller.abort("cancelCheckpointing()"); + this.#abortControllers.delete(runId); + } + + async #checkpointAndPush({ + runId, + leaveRunning = true, // This mirrors kubernetes behaviour more accurately + projectRef, + deploymentVersion, + }: CheckpointAndPushOptions): Promise { await this.initialize(); if (!this.#dockerMode && !this.#canCheckpoint) { @@ -110,124 +179,138 @@ class Checkpointer { return; } - try { - const { path } = await this.#checkpointContainer(podName, leaveRunning); - const { tag } = await this.#buildImage(path, podName); - const { destination } = await this.#pushImage(tag); + if (this.#abortControllers.has(runId)) { + logger.error("Checkpoint procedure already in progress", { + options: { + runId, + leaveRunning, + projectRef, + deploymentVersion, + }, + }); + return; + } + const controller = new AbortController(); + this.#abortControllers.set(runId, controller); + + const $$ = $({ signal: controller.signal }); + + try { + const shortCode = nanoid(8); + const imageRef = this.#getImageRef(projectRef, deploymentVersion, shortCode); + const exportLocation = this.#getExportLocation(projectRef, deploymentVersion, shortCode); + + this.#logger.log("Checkpointing:", { + options: { + runId, + leaveRunning, + projectRef, + deploymentVersion, + }, + }); + + const containterName = this.#getRunContainerName(runId); + + // Create checkpoint (docker) if (this.#dockerMode) { - this.#logger.log("checkpoint created:", { podName, path }); - } else { - this.#logger.log("checkpointed and pushed image to:", destination); + try { + if (this.opts.forceSimulate || !this.#canCheckpoint) { + this.#logger.log("Simulating checkpoint"); + this.#logger.debug(await $$`docker pause ${containterName}`); + } else { + if (leaveRunning) { + this.#logger.debug( + await $$`docker checkpoint create --leave-running ${containterName} ${exportLocation}` + ); + } else { + this.#logger.debug( + await $$`docker checkpoint create ${containterName} ${exportLocation}` + ); + } + } + } catch (error: any) { + this.#logger.error(error.stderr); + return; + } + + this.#logger.log("checkpoint created:", { + runId, + location: exportLocation, + }); + + return { + location: exportLocation, + docker: true, + }; + } + + // Create checkpoint (CRI) + if (!this.#canCheckpoint) { + throw new Error("No checkpoint support in kubernetes mode."); + } + + const containerId = this.#logger.debug( + // @ts-expect-error + await $$`crictl ps` + .pipeStdout($$({ stdin: "pipe" })`grep ${containterName}`) + .pipeStdout($$({ stdin: "pipe" })`cut -f1 ${"-d "}`) + ); + + if (!containerId.stdout) { + throw new Error("could not find container id"); + } + + this.#logger.debug(await $$`crictl checkpoint --export=${exportLocation} ${containerId}`); + + // Create image from checkpoint + const container = this.#logger.debug(await $$`buildah from scratch`); + this.#logger.debug(await $$`buildah add ${container} ${exportLocation} /`); + this.#logger.debug( + await $$`buildah config --annotation=io.kubernetes.cri-o.annotations.checkpoint.name=counter ${container}` + ); + this.#logger.debug(await $$`buildah commit ${container} ${imageRef}`); + this.#logger.debug(await $$`buildah rm ${container}`); + + // Push checkpoint image + this.#logger.debug(await $$`buildah push --tls-verify=${REGISTRY_TLS_VERIFY} ${imageRef}`); + + this.#logger.log("Checkpointed and pushed image to:", { location: imageRef }); + + try { + await $$`rm ${exportLocation}`; + this.#logger.log("Deleted checkpoint archive", { exportLocation }); + + // Disabled for now as this will increase restore time by having to pull the image again + // await $`buildah rmi ${imageRef}`; + // this.#logger.log("Deleted checkpoint image", { imageRef }); + } catch (error) { + this.#logger.error("Failed during checkpoint cleanup", { exportLocation }); + this.#logger.debug(error); } return { - path, - tag, - destination: this.#dockerMode ? path : destination, - docker: this.#dockerMode, + location: imageRef, + docker: false, }; } catch (error) { - this.#logger.error("checkpoint failed", error); + this.#logger.error("checkpoint failed", { + options: { + runId, + leaveRunning, + projectRef, + deploymentVersion, + }, + error, + }); return; + } finally { + this.#abortControllers.delete(runId); } } - async #checkpointContainer(podName: string, leaveRunning = false) { - await this.initialize(); - - if (this.#dockerMode) { - this.#logger.log("Checkpointing:", podName); - - const path = randomUUID(); - - try { - if (this.opts.forceSimulate || !this.#canCheckpoint) { - this.#logger.log("Simulating checkpoint"); - this.#logger.debug(await $`docker pause ${podName}`); - } else { - if (leaveRunning) { - this.#logger.debug( - await $`docker checkpoint create --leave-running ${podName} ${path}` - ); - } else { - this.#logger.debug(await $`docker checkpoint create ${podName} ${path}`); - } - } - } catch (error: any) { - this.#logger.error(error.stderr); - } - - return { path }; - } - - if (!this.#canCheckpoint) { - throw new Error("No checkpoint support. Simulation requires docker."); - } - - const containerId = this.#logger.debug( - // @ts-expect-error - await $`crictl ps` - .pipeStdout($({ stdin: "pipe" })`grep ${podName}`) - .pipeStdout($({ stdin: "pipe" })`cut -f1 ${"-d "}`) - ); - - if (!containerId.stdout) { - throw new Error("could not find container id"); - } - - const exportPath = `${CHECKPOINT_PATH}/${podName}.tar`; - - this.#logger.debug(await $`crictl checkpoint --export=${exportPath} ${containerId}`); - - return { - path: exportPath, - }; - } - - async #buildImage(checkpointPath: string, tag: string) { - await this.initialize(); - - if (this.#dockerMode) { - // Nothing to do here - return { tag }; - } - - if (!this.#canCheckpoint) { - throw new Error("No checkpoint support. Simulation requires docker."); - } - - const container = this.#logger.debug(await $`buildah from scratch`); - this.#logger.debug(await $`buildah add ${container} ${checkpointPath} /`); - this.#logger.debug( - await $`buildah config --annotation=io.kubernetes.cri-o.annotations.checkpoint.name=counter ${container}` - ); - this.#logger.debug(await $`buildah commit ${container} ${REGISTRY_FQDN}/${REPO_NAME}:${tag}`); - this.#logger.debug(await $`buildah rm ${container}`); - - return { - tag, - }; - } - - async #pushImage(tag: string) { - await this.initialize(); - - if (this.#dockerMode) { - // Nothing to do here - return { destination: "" }; - } - - if (!this.#canCheckpoint) { - throw new Error("No checkpoint support. Simulation requires docker."); - } - - const destination = `${REGISTRY_FQDN}/${REPO_NAME}:${tag}`; - this.#logger.debug(await $`buildah push --tls-verify=${REGISTRY_TLS_VERIFY} ${destination}`); - - return { - destination, - }; + #getRunContainerName(suffix: string) { + return `task-run-${suffix}`; } } @@ -245,6 +328,11 @@ class TaskCoordinator { typeof PlatformToCoordinatorMessages >; + #checkpointableTasks = new Map< + string, + { resolve: (value: void) => void; reject: (err?: any) => void } + >(); + constructor( private port: number, private host = "0.0.0.0" @@ -281,47 +369,78 @@ class TaskCoordinator { serverMessages: PlatformToCoordinatorMessages, authToken: PLATFORM_SECRET, handlers: { - RESUME: async (message) => { - const taskSocket = await this.#getAttemptSocket(message.attemptId); + RESUME_AFTER_DEPENDENCY: async (message) => { + const taskSocket = await this.#getAttemptSocket(message.attemptFriendlyId); if (!taskSocket) { - logger.log("Socket for attempt not found", { attemptId: message.attemptId }); + logger.log("Socket for attempt not found", { + attemptFriendlyId: message.attemptFriendlyId, + }); return; } - taskSocket.emit("RESUME", message); + // In case the task resumed faster than we could checkpoint + this.#cancelCheckpoint(message.runId); + + taskSocket.emit("RESUME_AFTER_DEPENDENCY", message); }, RESUME_AFTER_DURATION: async (message) => { - const taskSocket = await this.#getAttemptSocket(message.attemptId); + const taskSocket = await this.#getAttemptSocket(message.attemptFriendlyId); if (!taskSocket) { - logger.log("Socket for attempt not found", { attemptId: message.attemptId }); + logger.log("Socket for attempt not found", { + attemptFriendlyId: message.attemptFriendlyId, + }); return; } taskSocket.emit("RESUME_AFTER_DURATION", message); }, REQUEST_ATTEMPT_CANCELLATION: async (message) => { - const taskSocket = await this.#getAttemptSocket(message.attemptId); + const taskSocket = await this.#getAttemptSocket(message.attemptFriendlyId); if (!taskSocket) { - logger.log("Socket for attempt not found", { attemptId: message.attemptId }); + logger.log("Socket for attempt not found", { + attemptFriendlyId: message.attemptFriendlyId, + }); return; } taskSocket.emit("REQUEST_ATTEMPT_CANCELLATION", message); }, + READY_FOR_RETRY: async (message) => { + const taskSocket = await this.#getRunSocket(message.runId); + + if (!taskSocket) { + logger.log("Socket for attempt not found", { + runId: message.runId, + }); + return; + } + + taskSocket.emit("READY_FOR_RETRY", message); + }, }, }); return platformConnection; } - async #getAttemptSocket(attemptId: string) { + async #getRunSocket(runId: string) { const sockets = await this.#prodWorkerNamespace.fetchSockets(); for (const socket of sockets) { - if (socket.data.attemptId === attemptId) { + if (socket.data.runId === runId) { + return socket; + } + } + } + + async #getAttemptSocket(attemptFriendlyId: string) { + const sockets = await this.#prodWorkerNamespace.fetchSockets(); + + for (const socket of sockets) { + if (socket.data.attemptFriendlyId === attemptFriendlyId) { return socket; } } @@ -335,14 +454,22 @@ class TaskCoordinator { serverMessages: CoordinatorToProdWorkerMessages, socketData: ProdWorkerSocketData, postAuth: async (socket, next, logger) => { - function setSocketDataFromHeader(dataKey: keyof typeof socket.data, headerName: string) { + function setSocketDataFromHeader( + dataKey: keyof typeof socket.data, + headerName: string, + required: boolean = true + ) { const value = socket.handshake.headers[headerName]; - if (!value) { - logger(`missing required header: ${headerName}`); + + if (value) { + socket.data[dataKey] = Array.isArray(value) ? value[0] : value; + return; + } + + if (required) { + logger.error("missing required header", { headerName }); throw new Error("missing header"); } - 0; - socket.data[dataKey] = Array.isArray(value) ? value[0] : value; } try { @@ -350,27 +477,64 @@ class TaskCoordinator { setSocketDataFromHeader("contentHash", "x-trigger-content-hash"); setSocketDataFromHeader("projectRef", "x-trigger-project-ref"); setSocketDataFromHeader("runId", "x-trigger-run-id"); - setSocketDataFromHeader("attemptId", "x-trigger-attempt-id"); + setSocketDataFromHeader("attemptFriendlyId", "x-trigger-attempt-friendly-id", false); setSocketDataFromHeader("envId", "x-trigger-env-id"); setSocketDataFromHeader("deploymentId", "x-trigger-deployment-id"); + setSocketDataFromHeader("deploymentVersion", "x-trigger-deployment-version"); } catch (error) { - logger(error); + logger.error("setSocketDataFromHeader error", { error }); socket.disconnect(true); return; } - logger("success", socket.data); + logger.debug("success", socket.data); next(); }, onConnection: async (socket, handler, sender) => { const logger = new SimpleLogger(`[prod-worker][${socket.id}]`); + const checkpointInProgress = () => { + return this.#checkpointableTasks.has(socket.data.runId); + }; + + const readyToCheckpoint = async (): Promise< + { success: true } | { success: false; reason?: string } + > => { + if (checkpointInProgress()) { + return { + success: false, + reason: "checkpoint in progress", + }; + } + + const isCheckpointable = new Promise((resolve, reject) => { + // We set a reasonable timeout to prevent waiting forever + // TODO: We may also want to cancel the task as it's unlikely to recover + setTimeout(() => reject("timeout"), 10_000); + + this.#checkpointableTasks.set(socket.data.runId, { resolve, reject }); + }); + + try { + await isCheckpointable; + this.#checkpointableTasks.delete(socket.data.runId); + + return { + success: true, + }; + } catch (error) { + logger.error("Error while waiting for checkpointable state", { error }); + + return { + success: false, + reason: typeof error === "string" ? error : "unknown", + }; + } + }; + this.#platformSocket?.send("LOG", { - metadata: { - projectRef: socket.data.projectRef, - attemptId: socket.data.attemptId, - }, + metadata: socket.data, text: "connected", }); @@ -381,7 +545,7 @@ class TaskCoordinator { this.#platformSocket?.send("LOG", { version: "v1", - metadata: { attemptId: socket.data.attemptId }, + metadata: socket.data, text: message.text, }); }); @@ -389,67 +553,64 @@ class TaskCoordinator { 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, - runId: message.runId, - }); + try { + const executionAck = await this.#platformSocket?.sendWithAck( + "READY_FOR_EXECUTION", + message + ); - if (!executionAck) { - logger.error("no execution ack", { attemptId: socket.data.attemptId }); + if (!executionAck) { + logger.error("no execution ack", { runId: socket.data.runId }); - socket.emit("REQUEST_EXIT", { + socket.emit("REQUEST_EXIT", { + version: "v1", + }); + + return; + } + + if (!executionAck.success) { + logger.error("failed to get execution payload", { runId: socket.data.runId }); + + socket.emit("REQUEST_EXIT", { + version: "v1", + }); + + return; + } + + socket.emit("EXECUTE_TASK_RUN", { version: "v1", + executionPayload: executionAck.payload, }); - return; + socket.data.attemptFriendlyId = executionAck.payload.execution.attempt.id; + } catch (error) { + logger.error("Error", { error }); } - - if (!executionAck.success) { - logger.error("failed to get execution payload", { attemptId: socket.data.attemptId }); - - socket.emit("REQUEST_EXIT", { - version: "v1", - }); - - return; - } - - socket.emit("EXECUTE_TASK_RUN", { - version: "v1", - executionPayload: executionAck.payload, - }); }); socket.on("READY_FOR_RESUME", async (message) => { logger.log("[READY_FOR_RESUME]", message); + + socket.data.attemptFriendlyId = message.attemptFriendlyId; this.#platformSocket?.send("READY_FOR_RESUME", message); }); socket.on("TASK_RUN_COMPLETED", async ({ completion, execution }, callback) => { logger.log("completed task", { completionId: completion.id }); - const sendCompletionToPlatform = () => { + const completeWithoutCheckpoint = (shouldExit: boolean) => { this.#platformSocket?.send("TASK_RUN_COMPLETED", { version: "v1", execution, completion, }); - }; - - const confirmCompletion = ({ - didCheckpoint, - shouldExit, - }: { - didCheckpoint: boolean; - shouldExit: boolean; - }) => { - sendCompletionToPlatform(); - callback({ didCheckpoint, shouldExit }); + callback({ willCheckpointAndRestore: false, shouldExit }); }; if (completion.ok) { - confirmCompletion({ didCheckpoint: false, shouldExit: true }); + completeWithoutCheckpoint(true); return; } @@ -457,12 +618,17 @@ class TaskCoordinator { completion.error.type === "INTERNAL_ERROR" && completion.error.code === "TASK_RUN_CANCELLED" ) { - confirmCompletion({ didCheckpoint: false, shouldExit: true }); + completeWithoutCheckpoint(true); return; } if (completion.retry === undefined) { - confirmCompletion({ didCheckpoint: false, shouldExit: true }); + completeWithoutCheckpoint(true); + return; + } + + if (completion.retry.delay < 10_000) { + completeWithoutCheckpoint(false); return; } @@ -471,35 +637,77 @@ class TaskCoordinator { const willCheckpointAndRestore = canCheckpoint || willSimulate; if (!willCheckpointAndRestore) { - confirmCompletion({ didCheckpoint: false, shouldExit: false }); + completeWithoutCheckpoint(false); return; } - const checkpoint = await this.#checkpointer.checkpointAndPush(socket.data.podName); + // The worker will then put itself in a checkpointable state + callback({ willCheckpointAndRestore: true, shouldExit: false }); - if (!checkpoint) { - logger.error("Failed to checkpoint", { podName: socket.data.podName }); - confirmCompletion({ didCheckpoint: false, shouldExit: false }); + const ready = await readyToCheckpoint(); + + if (!ready.success) { + logger.error("Failed to become checkpointable", { + runId: socket.data.runId, + reason: ready.reason, + }); return; } - this.#platformSocket?.send("CHECKPOINT_CREATED", { - version: "v1", - attemptId: socket.data.attemptId, - docker: checkpoint.docker, - location: checkpoint.destination, - reason: { - type: "RETRYING_AFTER_FAILURE", - attemptNumber: execution.attempt.number, - }, + const checkpoint = await this.#checkpointer.checkpointAndPush({ + runId: socket.data.runId, + projectRef: socket.data.projectRef, + deploymentVersion: socket.data.deploymentVersion, }); - confirmCompletion({ didCheckpoint: true, shouldExit: false }); + if (!checkpoint) { + logger.error("Failed to checkpoint", { runId: socket.data.runId }); + completeWithoutCheckpoint(false); + return; + } + + this.#platformSocket?.send("TASK_RUN_COMPLETED", { + version: "v1", + execution, + completion, + checkpoint, + }); + + if (!checkpoint.docker || !willSimulate) { + socket.emit("REQUEST_EXIT", { + version: "v1", + }); + } + }); + + socket.on("READY_FOR_CHECKPOINT", async (message) => { + logger.log("[READY_FOR_CHECKPOINT]", message); + + const checkpointable = this.#checkpointableTasks.get(socket.data.runId); + + if (!checkpointable) { + logger.error("No checkpoint scheduled", { runId: socket.data.runId }); + return; + } + + checkpointable.resolve(); + }); + + socket.on("CANCEL_CHECKPOINT", async (message) => { + logger.log("[CANCEL_CHECKPOINT]", message); + + this.#cancelCheckpoint(socket.data.runId); }); socket.on("WAIT_FOR_DURATION", async (message, callback) => { logger.log("[WAIT_FOR_DURATION]", message); + if (checkpointInProgress()) { + logger.error("Checkpoint already in progress", { runId: socket.data.runId }); + callback({ willCheckpointAndRestore: false }); + return; + } + const { canCheckpoint, willSimulate } = await this.#checkpointer.initialize(); const willCheckpointAndRestore = canCheckpoint || willSimulate; @@ -510,26 +718,43 @@ class TaskCoordinator { return; } - // Wait for attempt to reach checkpointable state - // TODO: The worker should let us know when to checkpoint so we don't have to guess - await new Promise((resolve) => setTimeout(resolve, 2_000)); + const ready = await readyToCheckpoint(); - const checkpoint = await this.#checkpointer.checkpointAndPush(socket.data.podName); + if (!ready.success) { + logger.error("Failed to become checkpointable", { + runId: socket.data.runId, + reason: ready.reason, + }); + return; + } + + const checkpoint = await this.#checkpointer.checkpointAndPush({ + runId: socket.data.runId, + projectRef: socket.data.projectRef, + deploymentVersion: socket.data.deploymentVersion, + }); if (!checkpoint) { - logger.error("Failed to checkpoint", { podName: socket.data.podName }); - // TODO: We have to let the worker know about failures so it can use its own timer + // The task container will keep running until the wait duration has elapsed + logger.error("Failed to checkpoint", { runId: socket.data.runId }); return; } + if (!checkpoint.docker || !willSimulate) { + socket.emit("REQUEST_EXIT", { + version: "v1", + }); + } + this.#platformSocket?.send("CHECKPOINT_CREATED", { version: "v1", - attemptId: socket.data.attemptId, + attemptFriendlyId: message.attemptFriendlyId, docker: checkpoint.docker, - location: checkpoint.destination, + location: checkpoint.location, reason: { type: "WAIT_FOR_DURATION", ms: message.ms, + now: message.now, }, }); }); @@ -547,21 +772,31 @@ class TaskCoordinator { return; } - const checkpoint = await this.#checkpointer.checkpointAndPush(socket.data.podName); + const checkpoint = await this.#checkpointer.checkpointAndPush({ + runId: socket.data.runId, + projectRef: socket.data.projectRef, + deploymentVersion: socket.data.deploymentVersion, + }); if (!checkpoint) { - logger.error("Failed to checkpoint", { podName: socket.data.podName }); + logger.error("Failed to checkpoint", { runId: socket.data.runId }); return; } + if (!checkpoint.docker || !willSimulate) { + socket.emit("REQUEST_EXIT", { + version: "v1", + }); + } + this.#platformSocket?.send("CHECKPOINT_CREATED", { version: "v1", - attemptId: socket.data.attemptId, + attemptFriendlyId: message.attemptFriendlyId, docker: checkpoint.docker, - location: checkpoint.destination, + location: checkpoint.location, reason: { type: "WAIT_FOR_TASK", - id: message.id, + friendlyId: message.friendlyId, }, }); }); @@ -579,21 +814,32 @@ class TaskCoordinator { return; } - const checkpoint = await this.#checkpointer.checkpointAndPush(socket.data.podName); + const checkpoint = await this.#checkpointer.checkpointAndPush({ + runId: socket.data.runId, + projectRef: socket.data.projectRef, + deploymentVersion: socket.data.deploymentVersion, + }); if (!checkpoint) { - logger.error("Failed to checkpoint", { podName: socket.data.podName }); + logger.error("Failed to checkpoint", { runId: socket.data.runId }); return; } + if (!checkpoint.docker || !willSimulate) { + socket.emit("REQUEST_EXIT", { + version: "v1", + }); + } + this.#platformSocket?.send("CHECKPOINT_CREATED", { version: "v1", - attemptId: socket.data.attemptId, + attemptFriendlyId: message.attemptFriendlyId, docker: checkpoint.docker, - location: checkpoint.destination, + location: checkpoint.location, reason: { type: "WAIT_FOR_BATCH", - id: message.id, + batchFriendlyId: message.batchFriendlyId, + runFriendlyIds: message.runFriendlyIds, }, }); }); @@ -632,10 +878,7 @@ class TaskCoordinator { }, onDisconnect: async (socket, handler, sender, logger) => { this.#platformSocket?.send("LOG", { - metadata: { - projectRef: socket.data.projectRef, - attemptId: socket.data.attemptId, - }, + metadata: socket.data, text: "disconnect", }); }, @@ -649,6 +892,18 @@ class TaskCoordinator { return provider; } + #cancelCheckpoint(runId: string) { + const checkpointWait = this.#checkpointableTasks.get(runId); + + if (checkpointWait) { + // Stop waiting for task to reach checkpointable state + checkpointWait.reject("Checkpoint cancelled"); + } + + // Cancel checkpointing procedure + this.#checkpointer.cancelCheckpoint(runId); + } + #createHttpServer() { const httpServer = createServer(async (req, res) => { logger.log(`[${req.method}]`, req.url); @@ -667,7 +922,7 @@ class TaskCoordinator { } case "/checkpoint": { const body = await getTextBody(req); - await this.#checkpointer.checkpointAndPush(body); + // await this.#checkpointer.checkpointAndPush(body); return reply.text(`sent restore request: ${body}`); } default: { diff --git a/apps/docker-provider/.env.example b/apps/docker-provider/.env.example index cdf3cf9cc..075b6823a 100644 --- a/apps/docker-provider/.env.example +++ b/apps/docker-provider/.env.example @@ -2,6 +2,7 @@ HTTP_SERVER_PORT=8050 PLATFORM_WS_PORT=3030 PLATFORM_SECRET=provider-secret + # Use this if you are on macOS # COORDINATOR_HOST="host.docker.internal" # OTEL_EXPORTER_OTLP_ENDPOINT="http://host.docker.internal:4318" \ No newline at end of file diff --git a/apps/docker-provider/package.json b/apps/docker-provider/package.json index 1b39e0026..59a1d522b 100644 --- a/apps/docker-provider/package.json +++ b/apps/docker-provider/package.json @@ -6,7 +6,7 @@ "main": "dist/index.cjs", "scripts": { "build": "npm run build:bundle", - "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:bundle": "esbuild src/index.ts --bundle --outfile=dist/index.mjs --platform=node --format=esm --target=esnext --banner:js=\"import { createRequire } from 'module';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", diff --git a/apps/docker-provider/src/index.ts b/apps/docker-provider/src/index.ts index 610a03cab..7fed78137 100644 --- a/apps/docker-provider/src/index.ts +++ b/apps/docker-provider/src/index.ts @@ -1,6 +1,13 @@ import { $, type ExecaChildProcess, execa } from "execa"; -import { Machine } from "@trigger.dev/core/v3"; -import { SimpleLogger, TaskOperations, ProviderShell } from "@trigger.dev/core-apps"; +import { + SimpleLogger, + TaskOperations, + ProviderShell, + TaskOperationsRestoreOptions, + TaskOperationsCreateOptions, + TaskOperationsIndexOptions, +} from "@trigger.dev/core-apps"; +import { setTimeout } from "node:timers/promises"; const MACHINE_NAME = process.env.MACHINE_NAME || "local"; const COORDINATOR_PORT = process.env.COORDINATOR_PORT || 8020; @@ -72,18 +79,12 @@ class DockerTaskOperations implements TaskOperations { }; } - async index(opts: { - contentHash: string; - imageTag: string; - envId: string; - apiKey: string; - apiUrl: string; - }) { + async index(opts: TaskOperationsIndexOptions) { await this.#initialize(); - const containerName = this.#getIndexContainerName(opts.contentHash); + const containerName = this.#getIndexContainerName(opts.shortCode); - logger.log(`Indexing task ${opts.imageTag}`, { + logger.log(`Indexing task ${opts.imageRef}`, { host: COORDINATOR_HOST, port: COORDINATOR_PORT, }); @@ -94,15 +95,16 @@ class DockerTaskOperations implements TaskOperations { "run", "--network=host", "--rm", + `--env=INDEX_TASKS=true`, `--env=TRIGGER_SECRET_KEY=${opts.apiKey}`, `--env=TRIGGER_API_URL=${opts.apiUrl}`, + `--env=TRIGGER_ENV_ID=${opts.envId}`, + `--env=OTEL_EXPORTER_OTLP_ENDPOINT=${OTEL_EXPORTER_OTLP_ENDPOINT}`, + `--env=POD_NAME=${containerName}`, `--env=COORDINATOR_HOST=${COORDINATOR_HOST}`, `--env=COORDINATOR_PORT=${COORDINATOR_PORT}`, - `--env=POD_NAME=${containerName}`, - `--env=TRIGGER_ENV_ID=${opts.envId}`, - `--env=INDEX_TASKS=true`, `--name=${containerName}`, - `${opts.imageTag}`, + `${opts.imageRef}`, ]) ); } catch (error: any) { @@ -117,21 +119,13 @@ class DockerTaskOperations implements TaskOperations { stdout: error.stdout, stderr: error.stderr, }); - - throw new Error(`Index failed with: ${error.stderr || error.stdout}`); } } - async create(opts: { - runId: string; - attemptId: string; - image: string; - machine: Machine; - envId: string; - }) { + async create(opts: TaskOperationsCreateOptions) { await this.#initialize(); - const containerName = this.#getRunContainerName(opts.attemptId); + const containerName = this.#getRunContainerName(opts.runId); try { logger.debug( @@ -139,13 +133,12 @@ class DockerTaskOperations implements TaskOperations { "run", "--network=host", "--detach", - `--env=OTEL_EXPORTER_OTLP_ENDPOINT=${OTEL_EXPORTER_OTLP_ENDPOINT}`, - `--env=COORDINATOR_HOST=${COORDINATOR_HOST}`, - `--env=COORDINATOR_PORT=${COORDINATOR_PORT}`, - `--env=POD_NAME=${containerName}`, `--env=TRIGGER_ENV_ID=${opts.envId}`, `--env=TRIGGER_RUN_ID=${opts.runId}`, - `--env=TRIGGER_ATTEMPT_ID=${opts.attemptId}`, + `--env=OTEL_EXPORTER_OTLP_ENDPOINT=${OTEL_EXPORTER_OTLP_ENDPOINT}`, + `--env=POD_NAME=${containerName}`, + `--env=COORDINATOR_HOST=${COORDINATOR_HOST}`, + `--env=COORDINATOR_PORT=${COORDINATOR_PORT}`, `--name=${containerName}`, `${opts.image}`, ]) @@ -162,30 +155,24 @@ class DockerTaskOperations implements TaskOperations { stdout: error.stdout, stderr: error.stderr, }); - - throw new Error(`Create failed with: ${error.stderr || error.stdout}`); } } - async restore(opts: { - runId: string; - attemptId: string; - checkpointRef: string; - machine: Machine; - }) { + async restore(opts: TaskOperationsRestoreOptions) { await this.#initialize(); - const containerName = this.#getRunContainerName(opts.attemptId); + const containerName = this.#getRunContainerName(opts.runId); if (!this.#canCheckpoint || this.opts.forceSimulate) { logger.log("Simulating restore"); - const { exitCode } = logger.debug(await $`docker unpause ${containerName}`); + const unpause = logger.debug(await $`docker unpause ${containerName}`); - if (exitCode !== 0) { + if (unpause.exitCode !== 0) { throw new Error("docker unpause command failed"); } + await this.#sendPostStart(containerName); return; } @@ -196,6 +183,8 @@ class DockerTaskOperations implements TaskOperations { if (exitCode !== 0) { throw new Error("docker start command failed"); } + + await this.#sendPostStart(containerName); } async delete(opts: { runId: string }) { @@ -210,12 +199,61 @@ class DockerTaskOperations implements TaskOperations { logger.log("noop: get"); } - #getIndexContainerName(contentHash: string) { - return `task-index-${contentHash}`; + #getIndexContainerName(suffix: string) { + return `task-index-${suffix}`; } - #getRunContainerName(attemptId: string) { - return `task-run-${attemptId}`; + #getRunContainerName(suffix: string) { + return `task-run-${suffix}`; + } + + async #sendPostStart(containerName: string): Promise { + // We first get the correct port, which is random during dev as we run with host networking and need to avoid clashes + // FIXME: Skip this in prod + const logs = logger.debug(await $`docker logs ${containerName}`); + const matches = logs.stdout.match(/http server listening on port (?[0-9]+)/); + + const port = Number(matches?.groups?.port); + + if (!port) { + throw new Error("failed to extract port from logs"); + } + + try { + logger.debug(await this.#runLifecycleCommand(containerName, port, "postStart", "restore")); + } catch (error) { + logger.error("postStart error", { error }); + throw new Error("postStart command failed"); + } + } + + async #runLifecycleCommand( + containerName: string, + port: number, + type: "postStart" | "preStop", + cause: "index" | "create" | "restore", + retryCount = 0 + ): Promise { + try { + return await execa("docker", [ + "exec", + containerName, + "wget", + "-q", + "-O-", + `127.0.0.1:${port}/${type}?cause=${cause}`, + ]); + } catch (error: any) { + if (retryCount < 6) { + logger.debug("retriable postStart error", { retryCount, message: error?.message }); + await setTimeout(exponentialBackoff(retryCount + 1, 2, 50, 1150, 50)); + + return this.#runLifecycleCommand(containerName, port, type, cause, retryCount + 1); + } + + logger.error("final postStart error", { message: error?.message }); + throw new Error(`postStart command failed after ${retryCount - 1} retries`); + } } } @@ -225,3 +263,20 @@ const provider = new ProviderShell({ }); provider.listen(); + +function exponentialBackoff( + retryCount: number, + exponential: number, + minDelay: number, + maxDelay: number, + jitter: number +): number { + // Calculate the delay using the exponential backoff formula + const delay = Math.min(Math.pow(exponential, retryCount) * minDelay, maxDelay); + + // Calculate the jitter + const jitterValue = Math.random() * jitter; + + // Return the calculated delay with jitter + return delay + jitterValue; +} diff --git a/apps/kubernetes-provider/.env.example b/apps/kubernetes-provider/.env.example index 352a6ca8e..10949a807 100644 --- a/apps/kubernetes-provider/.env.example +++ b/apps/kubernetes-provider/.env.example @@ -1,7 +1,8 @@ HTTP_SERVER_PORT=8060 -PLATFORM_WS_PORT=8003 +PLATFORM_WS_PORT=3030 PLATFORM_SECRET=provider-secret -REGISTRY_FQDN=docker.io -REPO_NAME=task \ No newline at end of file +# Use this if you are on macOS +# COORDINATOR_HOST="host.docker.internal" +# OTEL_EXPORTER_OTLP_ENDPOINT="http://host.docker.internal:4318" \ No newline at end of file diff --git a/apps/kubernetes-provider/Containerfile b/apps/kubernetes-provider/Containerfile index d9630c5a3..f9fad85f4 100644 --- a/apps/kubernetes-provider/Containerfile +++ b/apps/kubernetes-provider/Containerfile @@ -31,7 +31,6 @@ COPY --from=pruner --chown=node:node /app/out/full/ . COPY --from=dev-deps --chown=node:node /app/ . COPY --chown=node:node turbo.json turbo.json -RUN pnpm run -r --filter '@trigger.dev/core*' build RUN pnpm run -r --filter kubernetes-provider build:bundle FROM base AS runner diff --git a/apps/kubernetes-provider/package.json b/apps/kubernetes-provider/package.json index 92526c07c..f36a6c098 100644 --- a/apps/kubernetes-provider/package.json +++ b/apps/kubernetes-provider/package.json @@ -4,10 +4,12 @@ "version": "0.0.1", "description": "", "main": "dist/index.cjs", - "type": "module", "scripts": { "build": "npm run build:bundle", - "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:bundle": "esbuild src/index.ts --bundle --outfile=dist/index.mjs --platform=node --format=esm --target=esnext --banner:js=\"import { createRequire } from 'module';const require = createRequire(import.meta.url);\"", + "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", "typecheck": "tsc --noEmit" }, "keywords": [], diff --git a/apps/kubernetes-provider/src/index.ts b/apps/kubernetes-provider/src/index.ts index c790497b4..5b4d290b6 100644 --- a/apps/kubernetes-provider/src/index.ts +++ b/apps/kubernetes-provider/src/index.ts @@ -1,17 +1,21 @@ -import { randomUUID } from "node:crypto"; -import k8s, { BatchV1Api, CoreV1Api, V1Job, V1Pod } from "@kubernetes/client-node"; -import { Machine } from "@trigger.dev/core/v3"; -import { ProviderShell, SimpleLogger, TaskOperations } from "@trigger.dev/core-apps"; +import * as k8s from "@kubernetes/client-node"; +import { + ProviderShell, + SimpleLogger, + TaskOperations, + TaskOperationsCreateOptions, + TaskOperationsIndexOptions, + TaskOperationsRestoreOptions, +} from "@trigger.dev/core-apps"; +import { randomUUID } from "crypto"; const RUNTIME_ENV = process.env.KUBERNETES_PORT ? "kubernetes" : "local"; -const NODE_NAME = process.env.NODE_NAME || "some-node"; +const NODE_NAME = process.env.NODE_NAME || "local"; 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"; - const logger = new SimpleLogger(`[${NODE_NAME}]`); +logger.log(`running in ${RUNTIME_ENV} mode`); type Namespace = { metadata: { @@ -22,8 +26,8 @@ type Namespace = { class KubernetesTaskOperations implements TaskOperations { #namespace: Namespace; #k8sApi: { - core: CoreV1Api; - batch: BatchV1Api; + core: k8s.CoreV1Api; + batch: k8s.BatchV1Api; }; constructor(namespace = "default") { @@ -36,15 +40,16 @@ class KubernetesTaskOperations implements TaskOperations { this.#k8sApi = this.#createK8sApi(); } - async index(opts: { contentHash: string; imageTag: string; envId: string }) { + async index(opts: TaskOperationsIndexOptions) { await this.#createJob( { metadata: { - name: `task-index-${opts.contentHash}`, + name: this.#getIndexContainerName(opts.shortCode), namespace: this.#namespace.metadata.name, }, spec: { completions: 1, + backoffLimit: 0, ttlSecondsAfterFinished: 300, template: { metadata: { @@ -61,19 +66,19 @@ class KubernetesTaskOperations implements TaskOperations { ], containers: [ { - name: opts.contentHash, - image: opts.imageTag, + name: this.#getIndexContainerName(opts.shortCode), + image: opts.imageRef, ports: [ { containerPort: 8000, }, ], - resources: { - limits: { - cpu: "100m", - memory: "50Mi", - }, - }, + // resources: { + // limits: { + // cpu: "100m", + // memory: "50Mi", + // }, + // }, env: [ { name: "DEBUG", @@ -83,6 +88,14 @@ class KubernetesTaskOperations implements TaskOperations { name: "INDEX_TASKS", value: "true", }, + { + name: "TRIGGER_SECRET_KEY", + value: opts.apiKey, + }, + { + name: "TRIGGER_API_URL", + value: opts.apiUrl, + }, { name: "TRIGGER_ENV_ID", value: opts.envId, @@ -130,11 +143,11 @@ class KubernetesTaskOperations implements TaskOperations { ); } - async create(opts: { attemptId: string; image: string; machine: Machine; envId: string }) { + async create(opts: TaskOperationsCreateOptions) { await this.#createPod( { metadata: { - name: `task-run-${opts.attemptId}-${randomUUID().slice(0, 5)}`, + name: this.#getRunContainerName(opts.runId), namespace: this.#namespace.metadata.name, labels: { app: "task-run", @@ -149,7 +162,7 @@ class KubernetesTaskOperations implements TaskOperations { ], containers: [ { - name: opts.attemptId, + name: this.#getRunContainerName(opts.runId), image: opts.image, ports: [ { @@ -159,18 +172,38 @@ class KubernetesTaskOperations implements TaskOperations { // resources: { // limits: opts.machine, // }, + lifecycle: { + postStart: { + exec: { + command: this.#getLifecycleCommand("postStart", "create"), + }, + }, + preStop: { + exec: { + command: this.#getLifecycleCommand("preStop", "create"), + }, + }, + }, env: [ { name: "DEBUG", value: "true", }, + { + name: "HTTP_SERVER_PORT", + value: "8000", + }, { name: "TRIGGER_ENV_ID", value: opts.envId, }, { - name: "TRIGGER_ATTEMPT_ID", - value: opts.attemptId, + name: "TRIGGER_RUN_ID", + value: opts.runId, + }, + { + name: "TRIGGER_WORKER_VERSION", + value: opts.version, }, { name: "OTEL_EXPORTER_OTLP_ENDPOINT", @@ -201,6 +234,18 @@ class KubernetesTaskOperations implements TaskOperations { }, }, ], + volumeMounts: [ + { + name: "taskinfo", + mountPath: "/etc/taskinfo", + }, + ], + }, + ], + volumes: [ + { + name: "taskinfo", + emptyDir: {}, }, ], }, @@ -209,37 +254,56 @@ class KubernetesTaskOperations implements TaskOperations { ); } - async restore(opts: { - attemptId: string; - runId: string; - image: string; - name: string; - checkpointId: string; - machine: Machine; - }) { + async restore(opts: TaskOperationsRestoreOptions) { await this.#createPod( { metadata: { - name: opts.name, + name: `${this.#getRunContainerName(opts.runId)}-${randomUUID().slice(0, 8)}`, namespace: this.#namespace.metadata.name, + labels: { + app: "task-run", + }, }, spec: { + restartPolicy: "Never", imagePullSecrets: [ { - name: "regcred", + name: "registry-trigger", }, ], initContainers: [ { name: "pull-base-image", - image: this.#getRestoreImage(opts.runId, opts.checkpointId), + image: opts.imageRef, command: ["sleep", "0"], }, + { + name: "populate-taskinfo", + image: "busybox", + command: ["/bin/sh", "-c"], + args: ["printenv COORDINATOR_HOST | tee /etc/taskinfo/coordinator-host"], + env: [ + { + name: "COORDINATOR_HOST", + valueFrom: { + fieldRef: { + fieldPath: "status.hostIP", + }, + }, + }, + ], + volumeMounts: [ + { + name: "taskinfo", + mountPath: "/etc/taskinfo", + }, + ], + }, ], containers: [ { - name: opts.runId, - image: this.#getImageFromRunId(opts.runId), + name: this.#getRunContainerName(opts.runId), + image: opts.checkpointRef, ports: [ { containerPort: 8000, @@ -250,44 +314,30 @@ class KubernetesTaskOperations implements TaskOperations { // }, lifecycle: { postStart: { - httpGet: { - path: "/connect", - port: 8000, + exec: { + command: this.#getLifecycleCommand("postStart", "restore"), + }, + }, + preStop: { + exec: { + command: this.#getLifecycleCommand("preStop", "restore"), }, }, }, - env: [ + volumeMounts: [ { - name: "DEBUG", - value: "true", - }, - { - name: "POD_NAME", - valueFrom: { - fieldRef: { - fieldPath: "metadata.name", - }, - }, - }, - { - name: "COORDINATOR_HOST", - valueFrom: { - fieldRef: { - fieldPath: "status.hostIP", - }, - }, - }, - { - name: "NODE_NAME", - valueFrom: { - fieldRef: { - fieldPath: "spec.nodeName", - }, - }, + name: "taskinfo", + mountPath: "/etc/taskinfo", }, ], }, ], + volumes: [ + { + name: "taskinfo", + emptyDir: {}, + }, + ], }, }, this.#namespace @@ -296,7 +346,7 @@ class KubernetesTaskOperations implements TaskOperations { async delete(opts: { runId: string }) { await this.#deletePod({ - podName: opts.runId, + runId: opts.runId, namespace: this.#namespace, }); } @@ -305,12 +355,16 @@ class KubernetesTaskOperations implements TaskOperations { await this.#getPod(opts.runId, this.#namespace); } - #getImageFromRunId(runId: string) { - return `${REGISTRY_FQDN}/${REPO_NAME}:${runId}`; + #getLifecycleCommand(type: "postStart" | "preStop", cause: "index" | "create" | "restore") { + return ["/bin/sh", "-c", `sleep 1; wget -q -O- 127.0.0.1:8000/${type}?cause=${cause}`]; } - #getRestoreImage(runId: string, checkpointId: string) { - return `${REGISTRY_FQDN}/${REPO_NAME}:${checkpointId}`; + #getIndexContainerName(suffix: string) { + return `task-index-${suffix}`; + } + + #getRunContainerName(suffix: string) { + return `task-run-${suffix}`; } #createK8sApi() { @@ -330,59 +384,67 @@ class KubernetesTaskOperations implements TaskOperations { }; } - async #createPod(pod: V1Pod, namespace: Namespace) { + async #createPod(pod: k8s.V1Pod, namespace: Namespace) { try { const res = await this.#k8sApi.core.createNamespacedPod(namespace.metadata.name, pod); logger.debug(res.body); - } catch (err: any) { - if ("body" in err) { - logger.error(err.body); - } else { - logger.error(err); - } + } catch (err: unknown) { + this.#handleK8sError(err); } } - async #deletePod(opts: { podName: string; namespace: Namespace }) { + async #deletePod(opts: { runId: string; namespace: Namespace }) { try { const res = await this.#k8sApi.core.deleteNamespacedPod( - opts.podName, + opts.runId, opts.namespace.metadata.name ); logger.debug(res.body); - } catch (err: any) { - if ("body" in err) { - logger.error(err.body); - } else { - logger.error(err); - } + } catch (err: unknown) { + this.#handleK8sError(err); } } - async #getPod(podName: string, namespace: Namespace) { + async #getPod(runId: string, namespace: Namespace) { try { - const res = await this.#k8sApi.core.readNamespacedPod(podName, namespace.metadata.name); + const res = await this.#k8sApi.core.readNamespacedPod(runId, namespace.metadata.name); logger.debug(res.body); return res.body; - } catch (err: any) { - if ("body" in err) { - logger.error(err.body); - } else { - logger.error(err); - } + } catch (err: unknown) { + this.#handleK8sError(err); } } - async #createJob(job: V1Job, namespace: Namespace) { + async #createJob(job: k8s.V1Job, namespace: Namespace) { try { const res = await this.#k8sApi.batch.createNamespacedJob(namespace.metadata.name, job); logger.debug(res.body); - } catch (err: any) { - if ("body" in err) { - logger.error(err.body); + } catch (err: unknown) { + this.#handleK8sError(err); + } + } + + #throwUnlessRecord(candidate: unknown): asserts candidate is Record { + if (typeof candidate !== "object" || candidate === null) { + throw candidate; + } + } + + #handleK8sError(err: unknown) { + this.#throwUnlessRecord(err); + + if ("body" in err && err.body) { + logger.error(err.body); + this.#throwUnlessRecord(err.body); + + if (typeof err.body.message === "string") { + throw new Error(err.body?.message); } else { - logger.error(err); + throw err.body; } + } else { + logger.error(err); + throw err; } } } diff --git a/apps/kubernetes-provider/tsconfig.json b/apps/kubernetes-provider/tsconfig.json index a491b740b..345326da8 100644 --- a/apps/kubernetes-provider/tsconfig.json +++ b/apps/kubernetes-provider/tsconfig.json @@ -7,10 +7,10 @@ "strict": true, "skipLibCheck": true, "paths": { - "@trigger.dev/core/v3": ["../core/src/v3"], - "@trigger.dev/core/v3/*": ["../core/src/v3/*"], - "@trigger.dev/core-apps": ["../core-apps/src"], - "@trigger.dev/core-apps/*": ["../core-apps/src/*"] + "@trigger.dev/core/v3": ["../../packages/core/src/v3"], + "@trigger.dev/core/v3/*": ["../../packages/core/src/v3/*"], + "@trigger.dev/core-apps": ["../../packages/core-apps/src"], + "@trigger.dev/core-apps/*": ["../../packages/core-apps/src/*"] } } } diff --git a/apps/webapp/app/services/logger.server.ts b/apps/webapp/app/services/logger.server.ts index c2a413e2d..e38365072 100644 --- a/apps/webapp/app/services/logger.server.ts +++ b/apps/webapp/app/services/logger.server.ts @@ -30,3 +30,14 @@ export const workerLogger = new Logger( return fields ? { ...fields } : {}; } ); + +export const socketLogger = new Logger( + "socket", + (process.env.APP_LOG_LEVEL ?? "debug") as LogLevel, + [], + sensitiveDataReplacer, + () => { + const fields = currentFieldsStore.getStore(); + return fields ? { ...fields } : {}; + } +); diff --git a/apps/webapp/app/v3/handleSocketIo.server.ts b/apps/webapp/app/v3/handleSocketIo.server.ts index 6a2189660..7bdb7ee6e 100644 --- a/apps/webapp/app/v3/handleSocketIo.server.ts +++ b/apps/webapp/app/v3/handleSocketIo.server.ts @@ -52,7 +52,8 @@ function createCoordinatorNamespace(io: Server) { READY_FOR_EXECUTION: async (message) => { const payload = await sharedQueueTasks.getLatestExecutionPayloadFromRun( message.runId, - true + true, + !!message.totalCompletions ); if (!payload) { @@ -67,7 +68,11 @@ function createCoordinatorNamespace(io: Server) { }, TASK_RUN_COMPLETED: async (message) => { const completeAttempt = new CompleteAttemptService(); - await completeAttempt.call(message.completion, message.execution); + await completeAttempt.call({ + completion: message.completion, + execution: message.execution, + checkpoint: message.checkpoint, + }); }, TASK_HEARTBEAT: async (message) => { await sharedQueueTasks.taskHeartbeat(message.attemptFriendlyId); @@ -132,14 +137,14 @@ function createSharedQueueConsumerNamespace(io: Server) { clientMessages: ClientToSharedQueueMessages, serverMessages: SharedQueueToClientMessages, onConnection: async (socket, handler, sender, logger) => { - const sharedSocketConnection = new SharedSocketConnection( - sharedQueue.namespace, + const sharedSocketConnection = new SharedSocketConnection({ + namespace: sharedQueue.namespace, socket, - logger - ); + logger, + }); sharedSocketConnection.onClose.attach((closeEvent) => { - logger("Socket closed", { closeEvent }); + logger.info("Socket closed", { closeEvent }); }); await sharedSocketConnection.initialize(); diff --git a/apps/webapp/app/v3/marqs/devQueueConsumer.server.ts b/apps/webapp/app/v3/marqs/devQueueConsumer.server.ts index 637ae2c5c..46bad40f7 100644 --- a/apps/webapp/app/v3/marqs/devQueueConsumer.server.ts +++ b/apps/webapp/app/v3/marqs/devQueueConsumer.server.ts @@ -125,7 +125,7 @@ export class DevQueueConsumer { logger.debug("Task run completed", { taskRunCompletion: completion, execution }); const service = new CompleteAttemptService(); - const result = await service.call(completion, execution, this.env); + const result = await service.call({ completion, execution, env: this.env }); if (result === "COMPLETED") { this._inProgressRuns.delete(execution.run.id); diff --git a/apps/webapp/app/v3/marqs/sharedQueueConsumer.server.ts b/apps/webapp/app/v3/marqs/sharedQueueConsumer.server.ts index 48c68727d..3b3fbd6b1 100644 --- a/apps/webapp/app/v3/marqs/sharedQueueConsumer.server.ts +++ b/apps/webapp/app/v3/marqs/sharedQueueConsumer.server.ts @@ -10,7 +10,12 @@ import { ZodMessageSender, serverWebsocketMessages, } from "@trigger.dev/core/v3"; -import { BackgroundWorker, BackgroundWorkerTask } from "@trigger.dev/database"; +import { + BackgroundWorker, + BackgroundWorkerTask, + TaskRunAttemptStatus, + TaskRunStatus, +} from "@trigger.dev/database"; import { z } from "zod"; import { prisma } from "~/db.server"; import { logger } from "~/services/logger.server"; @@ -29,15 +34,18 @@ const MessageBody = z.discriminatedUnion("type", [ z.object({ type: z.literal("EXECUTE"), taskIdentifier: z.string(), + checkpointEventId: z.string().optional(), }), z.object({ type: z.literal("RESUME"), completedAttemptIds: z.string().array(), resumableAttemptId: z.string(), + checkpointEventId: z.string().optional(), }), z.object({ type: z.literal("RESUME_AFTER_DURATION"), resumableAttemptId: z.string(), + checkpointEventId: z.string(), }), ]); @@ -130,11 +138,6 @@ export class SharedQueueConsumer { logger.debug("Stopping shared queue consumer"); this._enabled = false; - - // TODO: think about automatic prod cancellation - - // We need to cancel all the in progress task run attempts and ack the messages so they will stop processing - // await this.#cancelInProgressAttempts(reason); } async #cancelInProgressAttempts(reason: string) { @@ -180,7 +183,7 @@ export class SharedQueueConsumer { this._taskFailures = 0; this._taskSuccesses = 0; - this.#doWork().finally(() => { }); + this.#doWork().finally(() => {}); } async #doWork() { @@ -241,7 +244,7 @@ export class SharedQueueConsumer { const message = await marqs?.dequeueMessageInSharedQueue(); if (!message) { - setTimeout(() => this.#doWork(), this._options.nextTickInterval); + this.#doMoreWork(this._options.nextTickInterval); return; } @@ -264,8 +267,8 @@ export class SharedQueueConsumer { queueMessage: message.data, envId, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + this.#ackAndDoMoreWork(message.messageId); return; } @@ -278,9 +281,7 @@ export class SharedQueueConsumer { env: environment, }); - await marqs?.acknowledgeMessage(message.messageId); - - setTimeout(() => this.#doWork(), this._options.interval); + this.#ackAndDoMoreWork(message.messageId); return; } @@ -297,23 +298,36 @@ export class SharedQueueConsumer { queueMessage: message.data, messageId: message.messageId, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + this.#ackAndDoMoreWork(message.messageId); return; } + const retryingFromCheckpoint = !!messageBody.data.checkpointEventId; + + const EXECUTABLE_RUN_STATUSES: { + fromCheckpoint: TaskRunStatus[]; + withoutCheckpoint: TaskRunStatus[]; + } = { + fromCheckpoint: ["WAITING_TO_RESUME"], + withoutCheckpoint: ["PENDING", "RETRYING_AFTER_FAILURE"], + }; + if ( - existingTaskRun.status !== "PENDING" && - existingTaskRun.status !== "RETRYING_AFTER_FAILURE" + (retryingFromCheckpoint && + !EXECUTABLE_RUN_STATUSES.fromCheckpoint.includes(existingTaskRun.status)) || + (!retryingFromCheckpoint && + !EXECUTABLE_RUN_STATUSES.withoutCheckpoint.includes(existingTaskRun.status)) ) { - logger.debug("Task run is not pending, aborting", { + logger.debug("Task run has invalid status for execution", { queueMessage: message.data, messageId: message.messageId, taskRun: existingTaskRun.id, status: existingTaskRun.status, + retryingFromCheckpoint, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + this.#ackAndDoMoreWork(message.messageId); return; } @@ -324,8 +338,8 @@ export class SharedQueueConsumer { queueMessage: message.data, messageId: message.messageId, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + this.#ackAndDoMoreWork(message.messageId); return; } @@ -335,8 +349,8 @@ export class SharedQueueConsumer { messageId: message.messageId, deployment: deployment.id, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + this.#ackAndDoMoreWork(message.messageId); return; } @@ -353,9 +367,7 @@ export class SharedQueueConsumer { taskSlugs: deployment.worker.tasks.map((task) => task.slug), }); - await marqs?.acknowledgeMessage(message.messageId); - - setTimeout(() => this.#doWork(), this._options.interval); + this.#ackAndDoMoreWork(message.messageId); return; } @@ -391,9 +403,7 @@ export class SharedQueueConsumer { messageId: message.messageId, }); - await marqs?.acknowledgeMessage(message.messageId); - - setTimeout(() => this.#doWork(), this._options.interval); + this.#ackAndDoMoreWork(message.messageId); return; } @@ -407,8 +417,7 @@ export class SharedQueueConsumer { }); if (!queue) { - await marqs?.nackMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.nextTickInterval); + await this.#nackAndDoMoreWork(message.messageId, this._options.nextTickInterval); return; } @@ -431,24 +440,31 @@ export class SharedQueueConsumer { }, }); - try { - const latestCheckpoint = lockedTaskRun.checkpoints[0]; + const isRetry = taskRunAttempt.number > 1; - if (lockedTaskRun.status === "RETRYING_AFTER_FAILURE" && latestCheckpoint) { - if (latestCheckpoint.reason !== "RETRYING_AFTER_FAILURE") { - logger.error("Latest checkpoint is invalid", { + try { + if (messageBody.data.checkpointEventId) { + const restoreService = new RestoreCheckpointService(); + + const checkpoint = await restoreService.call({ + eventId: messageBody.data.checkpointEventId, + isRetry, + }); + + if (!checkpoint) { + logger.error("Failed to restore checkpoint", { queueMessage: message.data, messageId: message.messageId, - resumableAttemptId: taskRunAttempt.id, - latestCheckpointId: latestCheckpoint.id, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + await this.#ackAndDoMoreWork(message.messageId); return; } - - const restoreService = new RestoreCheckpointService(); - await restoreService.call({ checkpointId: latestCheckpoint.id }); + } else if (isRetry) { + socketIo.coordinatorNamespace.emit("READY_FOR_RETRY", { + version: "v1", + runId: taskRunAttempt.taskRunId, + }); } else { await this._sender.send("BACKGROUND_WORKER_MESSAGE", { backgroundWorkerId: deployment.worker.friendlyId, @@ -458,6 +474,7 @@ export class SharedQueueConsumer { image: deployment.imageReference, envId: environment.id, runId: taskRunAttempt.taskRunId, + version: deployment.version, }, }); } @@ -490,22 +507,55 @@ export class SharedQueueConsumer { }), ]); - // Finally we need to nack the message so it can be retried - await marqs?.nackMessage(message.messageId); - } finally { - setTimeout(() => this.#doWork(), this._options.interval); + await this.#nackAndDoMoreWork(message.messageId); + return; } + break; } // Resume after dependency completed with no remaining retries case "RESUME": { + if (messageBody.data.checkpointEventId) { + try { + const restoreService = new RestoreCheckpointService(); + + const checkpoint = await restoreService.call({ + eventId: messageBody.data.checkpointEventId, + }); + + if (!checkpoint) { + logger.error("Failed to restore checkpoint", { + queueMessage: message.data, + messageId: message.messageId, + }); + + await this.#ackAndDoMoreWork(message.messageId); + return; + } + } catch (e) { + if (e instanceof Error) { + this._currentSpan?.recordException(e); + } else { + this._currentSpan?.recordException(new Error(String(e))); + } + + this._endSpanInNextIteration = true; + + await this.#nackAndDoMoreWork(message.messageId); + return; + } + + this.#doMoreWork(); + return; + } + if (messageBody.data.completedAttemptIds.length < 1) { logger.error("No attempt IDs provided", { queueMessage: message.data, messageId: message.messageId, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + await this.#ackAndDoMoreWork(message.messageId); return; } @@ -520,8 +570,8 @@ export class SharedQueueConsumer { queueMessage: message.data, messageId: message.messageId, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + await this.#ackAndDoMoreWork(message.messageId); return; } @@ -544,8 +594,8 @@ export class SharedQueueConsumer { queueMessage: message.data, messageId: message.messageId, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + await this.#ackAndDoMoreWork(message.messageId); return; } @@ -559,8 +609,7 @@ export class SharedQueueConsumer { }); if (!queue) { - await marqs?.nackMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.nextTickInterval); + await this.#nackAndDoMoreWork(message.messageId, this._options.nextTickInterval); return; } @@ -569,28 +618,6 @@ export class SharedQueueConsumer { return; } - if (resumableAttempt.status === "PAUSED") { - // We need to restore the attempt from the latest checkpoint before we can resume - const latestCheckpoint = resumableAttempt.checkpoints[0]; - - if (!latestCheckpoint) { - logger.error("No checkpoint found", { - queueMessage: message.data, - messageId: message.messageId, - resumableAttemptId: resumableAttempt.id, - }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); - return; - } - - const restoreService = new RestoreCheckpointService(); - await restoreService.call({ checkpointId: latestCheckpoint.id }); - - setTimeout(() => this.#doWork(), this._options.interval); - return; - } - const completions: TaskRunExecutionResult[] = []; const executions: TaskRunExecution[] = []; @@ -614,16 +641,15 @@ export class SharedQueueConsumer { queueMessage: message.data, messageId: message.messageId, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + await this.#ackAndDoMoreWork(message.messageId); return; } const completion = await this._tasks.getCompletionPayloadFromAttempt(completedAttempt.id); if (!completion) { - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + await this.#ackAndDoMoreWork(message.messageId); return; } @@ -634,8 +660,7 @@ export class SharedQueueConsumer { ); if (!executionPayload) { - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + await this.#ackAndDoMoreWork(message.messageId); return; } @@ -644,9 +669,11 @@ export class SharedQueueConsumer { try { // The attempt should still be running so we can broadcast to all coordinators to resume immediately - socketIo.coordinatorNamespace.emit("RESUME", { + socketIo.coordinatorNamespace.emit("RESUME_AFTER_DEPENDENCY", { version: "v1", + runId: resumableAttempt.taskRunId, attemptId: resumableAttempt.id, + attemptFriendlyId: resumableAttempt.friendlyId, completions, executions, }); @@ -659,68 +686,30 @@ export class SharedQueueConsumer { this._endSpanInNextIteration = true; - // Finally we need to nack the message so it can be retried - await marqs?.nackMessage(message.messageId); - } finally { - setTimeout(() => this.#doWork(), this._options.interval); + await this.#nackAndDoMoreWork(message.messageId); + return; } + break; } // Resume after duration-based wait case "RESUME_AFTER_DURATION": { - const resumableAttempt = await prisma.taskRunAttempt.findUnique({ - where: { - id: messageBody.data.resumableAttemptId, - }, - include: { - checkpoints: { - take: 1, - orderBy: { - createdAt: "desc", - }, - }, - taskRun: true, - }, - }); - - if (!resumableAttempt) { - logger.error("Resumable attempt not found", { - queueMessage: message.data, - messageId: message.messageId, - }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); - return; - } - - if (resumableAttempt.status !== "PAUSED") { - logger.error("Attempt not paused", { - queueMessage: message.data, - messageId: message.messageId, - }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); - return; - } - try { - // We need to restore the attempt from the latest checkpoint before we can resume - const latestCheckpoint = resumableAttempt.checkpoints[0]; + const restoreService = new RestoreCheckpointService(); - if (!latestCheckpoint) { - logger.error("No checkpoint found", { + const checkpoint = await restoreService.call({ + eventId: messageBody.data.checkpointEventId, + }); + + if (!checkpoint) { + logger.error("Failed to restore checkpoint", { queueMessage: message.data, messageId: message.messageId, - resumableAttemptId: resumableAttempt.id, }); - await marqs?.acknowledgeMessage(message.messageId); - setTimeout(() => this.#doWork(), this._options.interval); + + await this.#ackAndDoMoreWork(message.messageId); return; } - - // The attempt will resume automatically after restore - const restoreService = new RestoreCheckpointService(); - await restoreService.call({ checkpointId: latestCheckpoint.id }); } catch (e) { if (e instanceof Error) { this._currentSpan?.recordException(e); @@ -730,19 +719,35 @@ export class SharedQueueConsumer { this._endSpanInNextIteration = true; - // Finally we need to nack the message so it can be retried - await marqs?.nackMessage(message.messageId); - } finally { - setTimeout(() => this.#doWork(), this._options.interval); + await this.#nackAndDoMoreWork(message.messageId); + return; } + break; } } + + this.#doMoreWork(); + return; } #envIdFromQueue(queueName: string) { return queueName.split(":")[1]; } + + #doMoreWork(intervalInMs = this._options.interval) { + setTimeout(() => this.#doWork(), intervalInMs); + } + + async #ackAndDoMoreWork(messageId: string, intervalInMs?: number) { + await marqs?.acknowledgeMessage(messageId); + this.#doMoreWork(intervalInMs); + } + + async #nackAndDoMoreWork(messageId: string, intervalInMs?: number) { + await marqs?.nackMessage(messageId); + this.#doMoreWork(intervalInMs); + } } class SharedQueueTasks { @@ -799,7 +804,8 @@ class SharedQueueTasks { async getExecutionPayloadFromAttempt( id: string, - setToExecuting?: boolean + setToExecuting?: boolean, + isRetrying?: boolean ): Promise { const attempt = await prisma.taskRunAttempt.findUnique({ where: { @@ -858,26 +864,42 @@ class SharedQueueTasks { } if (setToExecuting) { + const FINAL_RUN_STATUSES: TaskRunStatus[] = [ + "CANCELED", + "COMPLETED_SUCCESSFULLY", + "COMPLETED_WITH_ERRORS", + "INTERRUPTED", + "SYSTEM_FAILURE", + ]; + const FINAL_ATTEMPT_STATUSES: TaskRunAttemptStatus[] = ["CANCELED", "COMPLETED", "FAILED"]; + + if ( + FINAL_ATTEMPT_STATUSES.includes(attempt.status) || + FINAL_RUN_STATUSES.includes(attempt.taskRun.status) + ) { + logger.error("Status already in final state", { + attempt: { + id: attempt.id, + status: attempt.status, + }, + run: { + id: attempt.taskRunId, + status: attempt.taskRun.status, + }, + }); + return; + } + await prisma.taskRunAttempt.update({ where: { id, - taskRun: { - status: { - notIn: [ - "CANCELED", - "COMPLETED_SUCCESSFULLY", - "COMPLETED_WITH_ERRORS", - "SYSTEM_FAILURE", - ], - }, - }, }, data: { status: "EXECUTING", taskRun: { update: { data: { - status: "EXECUTING", + status: isRetrying ? "RETRYING_AFTER_FAILURE" : "EXECUTING", }, }, }, @@ -960,7 +982,8 @@ class SharedQueueTasks { async getLatestExecutionPayloadFromRun( id: string, - setToExecuting?: boolean + setToExecuting?: boolean, + isRetrying?: boolean ): Promise { const run = await prisma.taskRun.findUnique({ where: { @@ -983,7 +1006,7 @@ class SharedQueueTasks { return; } - return this.getExecutionPayloadFromAttempt(latestAttempt.id, setToExecuting); + return this.getExecutionPayloadFromAttempt(latestAttempt.id, setToExecuting, isRetrying); } async taskHeartbeat(attemptFriendlyId: string, seconds: number = 60) { diff --git a/apps/webapp/app/v3/services/cancelTaskRun.server.ts b/apps/webapp/app/v3/services/cancelTaskRun.server.ts index 3d2c29e99..fc9534e31 100644 --- a/apps/webapp/app/v3/services/cancelTaskRun.server.ts +++ b/apps/webapp/app/v3/services/cancelTaskRun.server.ts @@ -106,6 +106,7 @@ export class CancelTaskRunService extends BaseService { socketIo.coordinatorNamespace.emit("REQUEST_ATTEMPT_CANCELLATION", { version: "v1", attemptId: attempt.id, + attemptFriendlyId: attempt.friendlyId, }); break; diff --git a/apps/webapp/app/v3/services/completeAttempt.server.ts b/apps/webapp/app/v3/services/completeAttempt.server.ts index 9c76a3436..8a5e729d8 100644 --- a/apps/webapp/app/v3/services/completeAttempt.server.ts +++ b/apps/webapp/app/v3/services/completeAttempt.server.ts @@ -17,15 +17,28 @@ import { BaseService } from "./baseService.server"; import { CancelAttemptService } from "./cancelAttempt.server"; import { ResumeTaskRunDependenciesService } from "./resumeTaskRunDependencies.server"; import { MAX_TASK_RUN_ATTEMPTS } from "~/consts"; +import { CreateCheckpointService } from "./createCheckpoint.server"; +import { TaskRun } from "@trigger.dev/database"; type FoundAttempt = Awaited>; +type CheckpointData = { + docker: boolean; + location: string; +}; + export class CompleteAttemptService extends BaseService { - public async call( - completion: TaskRunExecutionResult, - execution: TaskRunExecution, - env?: AuthenticatedEnvironment - ): Promise<"COMPLETED" | "RETRIED"> { + public async call({ + completion, + execution, + env, + checkpoint, + }: { + completion: TaskRunExecutionResult; + execution: TaskRunExecution; + env?: AuthenticatedEnvironment; + checkpoint?: CheckpointData; + }): Promise<"COMPLETED" | "RETRIED"> { const taskRunAttempt = await findAttempt(this._prisma, completion.id); if (!taskRunAttempt) { @@ -47,7 +60,13 @@ export class CompleteAttemptService extends BaseService { if (completion.ok) { return await this.#completeAttemptSuccessfully(completion, taskRunAttempt, env); } else { - return await this.#completeAttemptFailed(completion, execution, taskRunAttempt, env); + return await this.#completeAttemptFailed( + completion, + execution, + taskRunAttempt, + env, + checkpoint + ); } } @@ -55,7 +74,7 @@ export class CompleteAttemptService extends BaseService { completion: TaskRunSuccessfulExecutionResult, taskRunAttempt: NonNullable, env?: AuthenticatedEnvironment - ): Promise<"COMPLETED" | "RETRIED"> { + ): Promise<"COMPLETED"> { await this._prisma.taskRunAttempt.update({ where: { friendlyId: completion.id }, data: { @@ -97,8 +116,9 @@ export class CompleteAttemptService extends BaseService { completion: TaskRunFailedExecutionResult, execution: TaskRunExecution, taskRunAttempt: NonNullable, - env?: AuthenticatedEnvironment - ) { + env?: AuthenticatedEnvironment, + checkpoint?: CheckpointData + ): Promise<"COMPLETED" | "RETRIED"> { if ( completion.error.type === "INTERNAL_ERROR" && completion.error.code === "TASK_RUN_CANCELLED" @@ -166,18 +186,47 @@ export class CompleteAttemptService extends BaseService { if (environment.type === "DEVELOPMENT") { // This is already an EXECUTE message so we can just NACK await marqs?.nackMessage(taskRunAttempt.taskRunId, completion.retry.timestamp); - } else { - // We have to replace a potential RESUME with EXECUTE to correctly retry the attempt - await marqs?.replaceMessage( - taskRunAttempt.taskRunId, - { - type: "EXECUTE", - taskIdentifier: taskRunAttempt.taskRun.taskIdentifier, - }, - completion.retry.timestamp - ); + return "RETRIED"; } + if (!checkpoint) { + await this.#enqueueRetry(taskRunAttempt.taskRun, completion.retry.timestamp); + return "RETRIED"; + } + + const createCheckpoint = new CreateCheckpointService(this._prisma); + const checkpointCreateResult = await createCheckpoint.call({ + attemptFriendlyId: execution.attempt.id, + docker: checkpoint.docker, + location: checkpoint.location, + reason: { + type: "RETRYING_AFTER_FAILURE", + attemptNumber: execution.attempt.number, + }, + }); + + if (!checkpointCreateResult) { + logger.error("Failed to create checkpoint", { checkpoint, execution: execution.run.id }); + + // Update the task run to be failed + await this._prisma.taskRun.update({ + where: { + friendlyId: execution.run.id, + }, + data: { + status: "SYSTEM_FAILURE", + }, + }); + + return "COMPLETED"; + } + + await this.#enqueueRetry( + taskRunAttempt.taskRun, + completion.retry.timestamp, + checkpointCreateResult.event.id + ); + return "RETRIED"; } else { // No more retries, we need to fail the task run @@ -210,6 +259,19 @@ export class CompleteAttemptService extends BaseService { } } + async #enqueueRetry(run: TaskRun, retryTimestamp: number, checkpointEventId?: string) { + // We have to replace a potential RESUME with EXECUTE to correctly retry the attempt + return await marqs?.replaceMessage( + run.id, + { + type: "EXECUTE", + taskIdentifier: run.taskIdentifier, + checkpointEventId: checkpointEventId, + }, + retryTimestamp + ); + } + #generateMetadataAttributesForNextAttempt(execution: TaskRunExecution) { const context = TaskRunContext.parse(execution); diff --git a/apps/webapp/app/v3/services/createCheckpoint.server.ts b/apps/webapp/app/v3/services/createCheckpoint.server.ts index f59092fdf..17500dba2 100644 --- a/apps/webapp/app/v3/services/createCheckpoint.server.ts +++ b/apps/webapp/app/v3/services/createCheckpoint.server.ts @@ -1,33 +1,79 @@ import { CoordinatorToPlatformMessages, InferSocketMessageSchema } from "@trigger.dev/core/v3"; -import type { Checkpoint } from "@trigger.dev/database"; -import { PrismaClient, prisma } from "~/db.server"; +import type { + CheckpointRestoreEvent, + TaskRunAttemptStatus, + TaskRunStatus, +} from "@trigger.dev/database"; import { logger } from "~/services/logger.server"; import { generateFriendlyId } from "../friendlyIdentifiers"; import { marqs } from "../marqs.server"; import { CreateCheckpointRestoreEventService } from "./createCheckpointRestoreEvent.server"; +import { BaseService } from "./baseService.server"; -export class CreateCheckpointService { - #prismaClient: PrismaClient; - - constructor(prismaClient: PrismaClient = prisma) { - this.#prismaClient = prismaClient; - } +const FREEZABLE_RUN_STATUSES: TaskRunStatus[] = ["EXECUTING", "RETRYING_AFTER_FAILURE"]; +const FREEZABLE_ATTEMPT_STATUSES: TaskRunAttemptStatus[] = ["EXECUTING", "FAILED"]; +export class CreateCheckpointService extends BaseService { public async call( - params: InferSocketMessageSchema - ): Promise { + params: Omit< + InferSocketMessageSchema, + "version" + > + ) { logger.debug(`Creating checkpoint`, params); - const attempt = await this.#prismaClient.taskRunAttempt.findUniqueOrThrow({ + const attempt = await this._prisma.taskRunAttempt.findUnique({ where: { - id: params.attemptId, + friendlyId: params.attemptFriendlyId, }, include: { taskRun: true, + backgroundWorker: { + select: { + id: true, + deployment: { + select: { + imageReference: true, + }, + }, + }, + }, }, }); - const checkpoint = await this.#prismaClient.checkpoint.create({ + if (!attempt) { + logger.error("Attempt not found", { attemptFriendlyId: params.attemptFriendlyId }); + return; + } + + if ( + !FREEZABLE_ATTEMPT_STATUSES.includes(attempt.status) || + !FREEZABLE_RUN_STATUSES.includes(attempt.taskRun.status) + ) { + logger.error("Unfreezable state", { + attempt: { + id: attempt.id, + status: attempt.status, + }, + run: { + id: attempt.taskRunId, + status: attempt.taskRun.status, + }, + }); + return; + } + + const imageRef = attempt.backgroundWorker.deployment?.imageReference; + + if (!imageRef) { + logger.error("Missing deployment or image ref", { + attemptId: attempt.id, + workerId: attempt.backgroundWorker.id, + }); + return; + } + + const checkpoint = await this._prisma.checkpoint.create({ data: { friendlyId: generateFriendlyId("checkpoint"), runtimeEnvironmentId: attempt.taskRun.runtimeEnvironmentId, @@ -38,44 +84,60 @@ export class CreateCheckpointService { type: params.docker ? "DOCKER" : "KUBERNETES", reason: params.reason.type, metadata: JSON.stringify(params.reason), + imageRef, }, }); - const eventService = new CreateCheckpointRestoreEventService(this.#prismaClient); - await eventService.call({ checkpointId: checkpoint.id, type: "CHECKPOINT" }); + const eventService = new CreateCheckpointRestoreEventService(this._prisma); - await this.#prismaClient.taskRunAttempt.update({ + await this._prisma.taskRunAttempt.update({ where: { - id: params.attemptId, + id: attempt.id, }, data: { - status: "PAUSED", + status: params.reason.type === "RETRYING_AFTER_FAILURE" ? undefined : "PAUSED", taskRun: { update: { - status: - params.reason.type === "RETRYING_AFTER_FAILURE" - ? "RETRYING_AFTER_FAILURE" - : "WAITING_TO_RESUME", + status: "WAITING_TO_RESUME", }, }, }, }); - switch (params.reason.type) { + const { reason } = params; + let checkpointEvent: CheckpointRestoreEvent | undefined; + + switch (reason.type) { case "WAIT_FOR_DURATION": { - await marqs?.replaceMessage( - attempt.taskRunId, - { type: "RESUME_AFTER_DURATION", resumableAttemptId: attempt.id }, - Date.now() + params.reason.ms - ); + checkpointEvent = await eventService.checkpoint({ + checkpointId: checkpoint.id, + }); + + break; + } + case "WAIT_FOR_TASK": { + checkpointEvent = await eventService.checkpoint({ + checkpointId: checkpoint.id, + dependencyFriendlyRunId: reason.friendlyId, + }); + + await marqs?.acknowledgeMessage(attempt.taskRunId); break; } - case "WAIT_FOR_TASK": case "WAIT_FOR_BATCH": { + checkpointEvent = await eventService.checkpoint({ + checkpointId: checkpoint.id, + batchDependencyFriendlyId: reason.batchFriendlyId, + }); + await marqs?.acknowledgeMessage(attempt.taskRunId); break; } case "RETRYING_AFTER_FAILURE": { + checkpointEvent = await eventService.checkpoint({ + checkpointId: checkpoint.id, + }); + // ACK is already handled by attempt completion break; } @@ -84,6 +146,30 @@ export class CreateCheckpointService { } } - return checkpoint; + if (!checkpointEvent) { + logger.error("No checkpoint event", { + attemptId: attempt.id, + checkpointId: checkpoint.id, + }); + await marqs?.acknowledgeMessage(attempt.taskRunId); + return; + } + + if (reason.type === "WAIT_FOR_DURATION") { + await marqs?.replaceMessage( + attempt.taskRunId, + { + type: "RESUME_AFTER_DURATION", + resumableAttemptId: attempt.id, + checkpointEventId: checkpointEvent.id, + }, + reason.now + reason.ms + ); + } + + return { + checkpoint, + event: checkpointEvent, + }; } } diff --git a/apps/webapp/app/v3/services/createCheckpointRestoreEvent.server.ts b/apps/webapp/app/v3/services/createCheckpointRestoreEvent.server.ts index f4acf5adc..cb6c75328 100644 --- a/apps/webapp/app/v3/services/createCheckpointRestoreEvent.server.ts +++ b/apps/webapp/app/v3/services/createCheckpointRestoreEvent.server.ts @@ -2,19 +2,69 @@ import type { CheckpointRestoreEvent, CheckpointRestoreEventType } from "@trigge import { logger } from "~/services/logger.server"; import { BaseService } from "./baseService.server"; -export class CreateCheckpointRestoreEventService extends BaseService { +interface CheckpointRestoreEventCallParams { + checkpointId: string; + type: CheckpointRestoreEventType; + dependencyFriendlyRunId?: string; + batchDependencyFriendlyId?: string; +} - public async call(params: { - checkpointId: string; - type: CheckpointRestoreEventType; - }): Promise { - const checkpoint = await this._prisma.checkpoint.findUniqueOrThrow({ +type CheckpointRestoreEventParams = Omit; + +export class CreateCheckpointRestoreEventService extends BaseService { + async checkpoint(params: CheckpointRestoreEventParams) { + return this.#call({ ...params, type: "CHECKPOINT" }); + } + + async restore(params: CheckpointRestoreEventParams) { + return this.#call({ ...params, type: "RESTORE" }); + } + + async #call( + params: CheckpointRestoreEventCallParams + ): Promise { + if (params.dependencyFriendlyRunId && params.batchDependencyFriendlyId) { + logger.error("Only one dependency can be set", { params }); + return; + } + + const checkpoint = await this._prisma.checkpoint.findUnique({ where: { id: params.checkpointId, }, }); - logger.debug(`Creating checkpoint/restore event`, params); + if (!checkpoint) { + logger.error("Checkpoint not found", { id: params.checkpointId }); + return; + } + + logger.debug(`Creating checkpoint/restore event`, { params }); + + let taskRunDependencyId: string | undefined; + + if (params.dependencyFriendlyRunId) { + const run = await this._prisma.taskRun.findUnique({ + where: { + friendlyId: params.dependencyFriendlyRunId, + }, + select: { + id: true, + dependency: { + select: { + id: true, + }, + }, + }, + }); + + taskRunDependencyId = run?.dependency?.id; + + if (!taskRunDependencyId) { + logger.error("Dependency or run not found", { runId: params.dependencyFriendlyRunId }); + return; + } + } const checkpointEvent = await this._prisma.checkpointRestoreEvent.create({ data: { @@ -26,6 +76,24 @@ export class CreateCheckpointRestoreEventService extends BaseService { type: params.type, reason: checkpoint.reason, metadata: checkpoint.metadata, + ...(taskRunDependencyId + ? { + taskRunDependency: { + connect: { + id: taskRunDependencyId, + }, + }, + } + : undefined), + ...(params.batchDependencyFriendlyId + ? { + batchTaskRunDependency: { + connect: { + friendlyId: params.batchDependencyFriendlyId, + }, + }, + } + : undefined), }, }); diff --git a/apps/webapp/app/v3/services/indexDeployment.server.ts b/apps/webapp/app/v3/services/indexDeployment.server.ts index 6ae1eaca7..854dcb034 100644 --- a/apps/webapp/app/v3/services/indexDeployment.server.ts +++ b/apps/webapp/app/v3/services/indexDeployment.server.ts @@ -47,7 +47,7 @@ export class IndexDeploymentService extends BaseService { try { const responses = await socketIo.providerNamespace.timeout(30_000).emitWithAck("INDEX", { version: "v1", - contentHash: deployment.contentHash, + shortCode: deployment.shortCode, imageTag: deployment.imageReference, envId: deployment.environmentId, apiKey: deployment.environment.apiKey, diff --git a/apps/webapp/app/v3/services/restoreCheckpoint.server.ts b/apps/webapp/app/v3/services/restoreCheckpoint.server.ts index 3704cb4e5..a61f95603 100644 --- a/apps/webapp/app/v3/services/restoreCheckpoint.server.ts +++ b/apps/webapp/app/v3/services/restoreCheckpoint.server.ts @@ -1,36 +1,79 @@ -import { type Checkpoint } from "@trigger.dev/database"; -import { PrismaClient, prisma } from "~/db.server"; +import { TaskRunStatus, type Checkpoint, TaskRunAttemptStatus } from "@trigger.dev/database"; import { logger } from "~/services/logger.server"; import { socketIo } from "../handleSocketIo.server"; import { CreateCheckpointRestoreEventService } from "./createCheckpointRestoreEvent.server"; +import { BaseService } from "./baseService.server"; -export class RestoreCheckpointService { - #prismaClient: PrismaClient; +const RESTORABLE_RUN_STATUSES: TaskRunStatus[] = ["WAITING_TO_RESUME"]; +const RESTORABLE_ATTEMPT_STATUSES: TaskRunAttemptStatus[] = ["PAUSED"]; - constructor(prismaClient: PrismaClient = prisma) { - this.#prismaClient = prismaClient; - } - - public async call(params: { checkpointId: string }): Promise { +export class RestoreCheckpointService extends BaseService { + public async call(params: { + eventId: string; + isRetry?: boolean; + }): Promise { logger.debug(`Restoring checkpoint`, params); - const checkpoint = await this.#prismaClient.checkpoint.findUniqueOrThrow({ + const checkpointEvent = await this._prisma.checkpointRestoreEvent.findUnique({ where: { - id: params.checkpointId, + id: params.eventId, + type: "CHECKPOINT", + }, + include: { + checkpoint: { + include: { + run: { + select: { + status: true, + }, + }, + attempt: { + select: { + status: true, + }, + }, + }, + }, }, }); - const eventService = new CreateCheckpointRestoreEventService(this.#prismaClient); - await eventService.call({ checkpointId: checkpoint.id, type: "RESTORE" }); + if (!checkpointEvent) { + logger.error("Checkpoint event not found", params); + return; + } + + const checkpoint = checkpointEvent.checkpoint; + + const runIsRestorable = RESTORABLE_RUN_STATUSES.includes(checkpoint.run.status); + const attemptIsRestorable = RESTORABLE_ATTEMPT_STATUSES.includes(checkpoint.attempt.status); + + if (!runIsRestorable) { + logger.error("Run is unrestorable", { + id: checkpoint.runId, + status: checkpoint.run.status, + }); + return; + } + + if (!attemptIsRestorable && !params.isRetry) { + logger.error("Attempt is unrestorable", { + id: checkpoint.attemptId, + status: checkpoint.attempt.status, + }); + return; + } + + const eventService = new CreateCheckpointRestoreEventService(this._prisma); + await eventService.restore({ checkpointId: checkpoint.id }); socketIo.providerNamespace.emit("RESTORE", { version: "v1", checkpointId: checkpoint.id, runId: checkpoint.runId, - attemptId: checkpoint.attemptId, type: checkpoint.type, location: checkpoint.location, reason: checkpoint.reason ?? undefined, + imageRef: checkpoint.imageRef, }); return checkpoint; diff --git a/apps/webapp/app/v3/services/resumeAttempt.server.ts b/apps/webapp/app/v3/services/resumeAttempt.server.ts index 0e4c20e5f..c7ec9771c 100644 --- a/apps/webapp/app/v3/services/resumeAttempt.server.ts +++ b/apps/webapp/app/v3/services/resumeAttempt.server.ts @@ -4,28 +4,23 @@ import { TaskRunExecution, TaskRunExecutionResult, } from "@trigger.dev/core/v3"; -import { $transaction, PrismaClient, prisma } from "~/db.server"; +import { $transaction } from "~/db.server"; import { logger } from "~/services/logger.server"; import { marqs } from "../marqs.server"; import { socketIo } from "../handleSocketIo.server"; import { sharedQueueTasks } from "../marqs/sharedQueueConsumer.server"; +import { BaseService } from "./baseService.server"; -export class ResumeAttemptService { - #prismaClient: PrismaClient; - - constructor(prismaClient: PrismaClient = prisma) { - this.#prismaClient = prismaClient; - } - +export class ResumeAttemptService extends BaseService { public async call( params: InferSocketMessageSchema ): Promise { logger.debug(`ResumeAttemptService.call()`, params); - await $transaction(this.#prismaClient, async (tx) => { + await $transaction(this._prisma, async (tx) => { const attempt = await tx.taskRunAttempt.findUnique({ where: { - id: params.attemptId, + friendlyId: params.attemptFriendlyId, }, include: { taskRun: true, @@ -71,13 +66,13 @@ export class ResumeAttemptService { }); if (!attempt) { - logger.error("Could not find attempt", { attemptId: params.attemptId }); + logger.error("Could not find attempt", { attemptFriendlyId: params.attemptFriendlyId }); return; } if (attempt.taskRun.status !== "WAITING_TO_RESUME") { logger.error("Run is not resumable", { - attemptId: params.attemptId, + attemptId: attempt.id, runId: attempt.taskRunId, }); return; @@ -85,7 +80,17 @@ export class ResumeAttemptService { switch (params.type) { case "WAIT_FOR_DURATION": { - // Nothing to do, but thanks for checking in! + logger.error( + "Attempt requested resume after duration wait, this is unexpected and likely a bug", + { attemptId: attempt.id } + ); + + // Attempts should not request resume for duration waits, this is just here as a backup + socketIo.coordinatorNamespace.emit("RESUME_AFTER_DURATION", { + version: "v1", + attemptId: attempt.id, + attemptFriendlyId: attempt.friendlyId, + }); break; } case "WAIT_FOR_TASK": @@ -96,7 +101,7 @@ export class ResumeAttemptService { const dependentAttempt = attempt.taskRunDependency.taskRun.attempts[0]; if (!dependentAttempt) { - logger.error("No dependent attempt", { attemptId: params.attemptId }); + logger.error("No dependent attempt", { attemptId: attempt.id }); return; } @@ -104,7 +109,7 @@ export class ResumeAttemptService { await tx.taskRunAttempt.update({ where: { - id: params.attemptId, + id: attempt.id, }, data: { taskRunDependency: { @@ -116,7 +121,7 @@ export class ResumeAttemptService { const dependentBatchItems = attempt.batchTaskRunDependency.items; if (!dependentBatchItems) { - logger.error("No dependent batch items", { attemptId: params.attemptId }); + logger.error("No dependent batch items", { attemptId: attempt.id }); return; } @@ -124,7 +129,7 @@ export class ResumeAttemptService { await tx.taskRunAttempt.update({ where: { - id: params.attemptId, + id: attempt.id, }, data: { batchTaskRunDependency: { @@ -133,12 +138,12 @@ export class ResumeAttemptService { }, }); } else { - logger.error("No dependencies", { attemptId: params.attemptId }); + logger.error("No dependencies", { attemptId: attempt.id }); return; } if (completedAttemptIds.length === 0) { - logger.error("No completed attempt IDs", { attemptId: params.attemptId }); + logger.error("No completed attempt IDs", { attemptId: attempt.id }); return; } @@ -146,7 +151,7 @@ export class ResumeAttemptService { const executions: TaskRunExecution[] = []; for (const completedAttemptId of completedAttemptIds) { - const completedAttempt = await prisma.taskRunAttempt.findUnique({ + const completedAttempt = await tx.taskRunAttempt.findUnique({ where: { id: completedAttemptId, taskRun: { @@ -162,7 +167,7 @@ export class ResumeAttemptService { if (!completedAttempt) { logger.error("Completed attempt not found", { - attemptId: params.attemptId, + attemptId: attempt.id, completedAttemptId, }); await marqs?.acknowledgeMessage(attempt.taskRunId); @@ -175,7 +180,7 @@ export class ResumeAttemptService { if (!completion) { logger.error("Failed to get completion payload", { - attemptId: params.attemptId, + attemptId: attempt.id, completedAttemptId, }); await marqs?.acknowledgeMessage(attempt.taskRunId); @@ -190,7 +195,7 @@ export class ResumeAttemptService { if (!executionPayload) { logger.error("Failed to get execution payload", { - attemptId: params.attemptId, + attemptId: attempt.id, completedAttemptId, }); await marqs?.acknowledgeMessage(attempt.taskRunId); @@ -200,25 +205,27 @@ export class ResumeAttemptService { executions.push(executionPayload.execution); } - await prisma.taskRunAttempt.update({ + const updated = await tx.taskRunAttempt.update({ where: { - id: params.attemptId, + id: attempt.id, }, data: { status: "EXECUTING", taskRun: { update: { data: { - status: "EXECUTING", + status: attempt.number > 1 ? "RETRYING_AFTER_FAILURE" : "EXECUTING", }, }, }, }, }); - socketIo.coordinatorNamespace.emit("RESUME", { + socketIo.coordinatorNamespace.emit("RESUME_AFTER_DEPENDENCY", { version: "v1", - attemptId: params.attemptId, + runId: attempt.taskRunId, + attemptId: attempt.id, + attemptFriendlyId: attempt.friendlyId, completions, executions, }); diff --git a/apps/webapp/app/v3/services/resumeBatchRun.server.ts b/apps/webapp/app/v3/services/resumeBatchRun.server.ts index deb19657c..5866b0b0b 100644 --- a/apps/webapp/app/v3/services/resumeBatchRun.server.ts +++ b/apps/webapp/app/v3/services/resumeBatchRun.server.ts @@ -2,6 +2,7 @@ import { PrismaClientOrTransaction } from "~/db.server"; import { workerQueue } from "~/services/worker.server"; import { marqs } from "../marqs.server"; import { BaseService } from "./baseService.server"; +import { logger } from "~/services/logger.server"; export class ResumeBatchRunService extends BaseService { public async call(batchRunId: string, sourceTaskAttemptId: string) { @@ -40,6 +41,7 @@ export class ResumeBatchRunService extends BaseService { return; } + // We need to update the batchRun status so we don't resume it again await this._prisma.batchTaskRun.update({ where: { id: batchRun.id, @@ -49,8 +51,6 @@ export class ResumeBatchRunService extends BaseService { }, }); - // We need to update the batchRun status so we don't resume it again - // This batch has a dependent attempt and just finalized, we should resume that attempt const environment = batchRun.dependentTaskAttempt.runtimeEnvironment; @@ -62,6 +62,15 @@ export class ResumeBatchRunService extends BaseService { const dependentRun = batchRun.dependentTaskAttempt.taskRun; if (batchRun.dependentTaskAttempt.status === "PAUSED") { + if (!batchRun.checkpointEventId) { + logger.error("Can't resume paused attempt without checkpoint event", { + batchRunId: batchRun.id, + }); + + await marqs?.acknowledgeMessage(dependentRun.id); + return; + } + await marqs?.enqueueMessage( environment, dependentRun.queue, @@ -70,6 +79,7 @@ export class ResumeBatchRunService extends BaseService { type: "RESUME", completedAttemptIds: [sourceTaskAttemptId], resumableAttemptId: batchRun.dependentTaskAttempt.id, + checkpointEventId: batchRun.checkpointEventId, }, dependentRun.concurrencyKey ?? undefined ); diff --git a/apps/webapp/app/v3/services/resumeTaskDependency.server.ts b/apps/webapp/app/v3/services/resumeTaskDependency.server.ts index cbfdfd3da..f4ebe1bdd 100644 --- a/apps/webapp/app/v3/services/resumeTaskDependency.server.ts +++ b/apps/webapp/app/v3/services/resumeTaskDependency.server.ts @@ -2,6 +2,7 @@ import { PrismaClientOrTransaction } from "~/db.server"; import { workerQueue } from "~/services/worker.server"; import { marqs } from "../marqs.server"; import { BaseService } from "./baseService.server"; +import { logger } from "~/services/logger.server"; export class ResumeTaskDependencyService extends BaseService { public async call(dependencyId: string, sourceTaskAttemptId: string) { @@ -34,9 +35,19 @@ export class ResumeTaskDependencyService extends BaseService { if (dependency.taskRun.runtimeEnvironment.type === "DEVELOPMENT") { return; } + const dependentRun = dependency.dependentAttempt.taskRun; if (dependency.dependentAttempt.status === "PAUSED") { + if (!dependency.checkpointEventId) { + logger.error("Can't resume paused attempt without checkpoint event", { + attemptId: dependency.id, + }); + + await marqs?.acknowledgeMessage(dependentRun.id); + return; + } + await marqs?.enqueueMessage( dependency.taskRun.runtimeEnvironment, dependentRun.queue, @@ -45,6 +56,7 @@ export class ResumeTaskDependencyService extends BaseService { type: "RESUME", completedAttemptIds: [sourceTaskAttemptId], resumableAttemptId: dependency.dependentAttempt.id, + checkpointEventId: dependency.checkpointEventId, }, dependentRun.concurrencyKey ?? undefined ); diff --git a/apps/webapp/app/v3/sharedSocketConnection.ts b/apps/webapp/app/v3/sharedSocketConnection.ts index 1e7992dfc..ac6d0fb76 100644 --- a/apps/webapp/app/v3/sharedSocketConnection.ts +++ b/apps/webapp/app/v3/sharedSocketConnection.ts @@ -1,5 +1,6 @@ import { MessageCatalogToSocketIoEvents, + StructuredLogger, ZodMessageHandler, ZodMessageSender, clientWebsocketMessages, @@ -11,6 +12,18 @@ import { logger } from "~/services/logger.server"; import { SharedQueueConsumer } from "./marqs/sharedQueueConsumer.server"; import { DisconnectReason, Namespace, Socket } from "socket.io"; +interface SharedSocketConnectionOptions { + namespace: Namespace< + MessageCatalogToSocketIoEvents, + MessageCatalogToSocketIoEvents + >; + socket: Socket< + MessageCatalogToSocketIoEvents, + MessageCatalogToSocketIoEvents + >; + logger?: StructuredLogger; +} + export class SharedSocketConnection { public id: string; public onClose: Evt = new Evt(); @@ -19,17 +32,7 @@ export class SharedSocketConnection { private _sharedConsumer: SharedQueueConsumer; private _messageHandler: ZodMessageHandler; - constructor( - namespace: Namespace< - MessageCatalogToSocketIoEvents, - MessageCatalogToSocketIoEvents - >, - private socket: Socket< - MessageCatalogToSocketIoEvents, - MessageCatalogToSocketIoEvents - >, - logger?: (...args: any[]) => void - ) { + constructor(opts: SharedSocketConnectionOptions) { this.id = randomUUID(); this._sender = new ZodMessageSender({ @@ -38,7 +41,7 @@ export class SharedSocketConnection { return new Promise((resolve, reject) => { try { const { type, ...payload } = message; - namespace.emit(type, payload as any); + opts.namespace.emit(type, payload as any); resolve(); } catch (err) { reject(err); @@ -52,8 +55,8 @@ export class SharedSocketConnection { nextTickInterval: 1000, }); - socket.on("disconnect", this.#handleClose.bind(this)); - socket.on("error", this.#handleError.bind(this)); + opts.socket.on("disconnect", this.#handleClose.bind(this)); + opts.socket.on("error", this.#handleError.bind(this)); this._messageHandler = new ZodMessageHandler({ schema: clientWebsocketMessages, @@ -78,7 +81,7 @@ export class SharedSocketConnection { }, }, }); - this._messageHandler.registerHandlers(this.socket, logger); + this._messageHandler.registerHandlers(opts.socket, opts.logger ?? logger); } async initialize() { diff --git a/packages/cli-v3/src/Containerfile.prod b/packages/cli-v3/src/Containerfile.prod index bd3672a67..ec3a0e410 100644 --- a/packages/cli-v3/src/Containerfile.prod +++ b/packages/cli-v3/src/Containerfile.prod @@ -1,4 +1,4 @@ -FROM node:18-alpine@sha256:ca9f6cb0466f9638e59e0c249d335a07c867cd50c429b5c7830dda1bed584649 AS base +FROM node:20-alpine@sha256:bf77dc26e48ea95fca9d1aceb5acfa69d2e546b765ec2abfb502975f1a2d4def AS base RUN apk add --no-cache dumb-init diff --git a/packages/cli-v3/src/commands/deploy.ts b/packages/cli-v3/src/commands/deploy.ts index 547b4c09a..d3bc2d894 100644 --- a/packages/cli-v3/src/commands/deploy.ts +++ b/packages/cli-v3/src/commands/deploy.ts @@ -38,6 +38,7 @@ import { detectPackageNameFromImportPath, parsePackageName } from "../utilities/ import { logger } from "../utilities/logger.js"; import { createTaskFileImports, gatherTaskFiles } from "../utilities/taskFiles"; import { login } from "./login"; +import { SetOptional } from "type-fest"; const DeployCommandOptions = CommonCommandOptions.extend({ skipTypecheck: z.boolean().default(false), @@ -47,7 +48,7 @@ const DeployCommandOptions = CommonCommandOptions.extend({ buildPlatform: z.enum(["linux/amd64", "linux/arm64"]).default("linux/amd64"), selfHosted: z.boolean().default(false), registry: z.string().optional(), - pushImage: z.boolean().default(false), + push: z.boolean().default(false), config: z.string().optional(), projectRef: z.string().optional(), outputMetafile: z.string().optional(), @@ -82,14 +83,14 @@ export function configureDeployCommand(program: Command) { ) .addOption( new CommandOption( - "--push-image", - "(Coming soon) When using the --self-hosted flag, push the image to the default registry. (defaults to false when not using --registry)" + "--push", + "When using the --self-hosted flag, push the image to the default registry. (defaults to false when not using --registry)" ).hideHelp() ) .addOption( new CommandOption( "--registry ", - "(Coming soon) The registry to push the image to when using --self-hosted" + "The registry to push the image to when using --self-hosted" ).hideHelp() ) .addOption( @@ -233,12 +234,13 @@ async function _deployCommand(dir: string, options: DeployCommandOptions) { const deploymentSpinner = spinner(); deploymentSpinner.start(`Deploying version ${version}`); - const registryHost = - deploymentResponse.data.registryHost ?? options.registry ?? "registry.trigger.dev"; + const selfHostedRegistryHost = deploymentResponse.data.registryHost ?? options.registry; + const registryHost = selfHostedRegistryHost ?? "registry.trigger.dev"; const buildImage = async () => { if (options.selfHosted) { return buildAndPushSelfHostedImage({ + registryHost: selfHostedRegistryHost, imageTag: deploymentResponse.data.imageTag, cwd: compilation.path, projectId: resolvedConfig.config.project, @@ -247,6 +249,8 @@ async function _deployCommand(dir: string, options: DeployCommandOptions) { contentHash: deploymentResponse.data.contentHash, projectRef: resolvedConfig.config.project, buildPlatform: options.buildPlatform, + pushImage: options.push, + selfHostedRegistry: !!options.registry, }); } @@ -283,7 +287,9 @@ async function _deployCommand(dir: string, options: DeployCommandOptions) { } const imageReference = options.selfHosted - ? `${image.image}${image.digest ? `@${image.digest}` : ""}` + ? `${selfHostedRegistryHost ? `${selfHostedRegistryHost}/` : ""}${image.image}${ + image.digest ? `@${image.digest}` : "" + }` : `${registryHost}/${image.image}${image.digest ? `@${image.digest}` : ""}`; span?.setAttributes({ @@ -637,10 +643,16 @@ async function buildAndPushImage( }); } -type BuildAndPushSelfHostedImageOptions = Omit< - BuildAndPushImageOptions, - "registryHost" | "buildId" | "buildToken" | "buildProjectId" | "auth" | "loadImage" ->; +type BuildAndPushSelfHostedImageOptions = SetOptional< + Omit< + BuildAndPushImageOptions, + "buildId" | "buildToken" | "buildProjectId" | "auth" | "loadImage" + >, + "registryHost" +> & { + pushImage: boolean; + selfHostedRegistry: boolean; +}; async function buildAndPushSelfHostedImage( options: BuildAndPushSelfHostedImageOptions @@ -656,7 +668,9 @@ async function buildAndPushSelfHostedImage( "options.projectRef": options.projectRef, }); - const args = [ + const imageRef = `${options.registryHost ? `${options.registryHost}/` : ""}${options.imageTag}`; + + const buildArgs = [ "build", "-f", "Containerfile", @@ -673,48 +687,41 @@ async function buildAndPushSelfHostedImage( "--build-arg", `TRIGGER_PROJECT_REF=${options.projectRef}`, "-t", - `${options.imageTag}`, + imageRef, ".", // The build context ].filter(Boolean) as string[]; - logger.debug(`docker ${args.join(" ")}`); + logger.debug(`docker ${buildArgs.join(" ")}`); - span.setAttribute("docker.command", `docker ${args.join(" ")}`); + span.setAttribute("docker.command.build", `docker ${buildArgs.join(" ")}`); - // Step 4: Build and push the image - const childProcess = execa("docker", args, { + // Build the image + const buildProcess = execa("docker", buildArgs, { cwd: options.cwd, }); const errors: string[] = []; + let digest: string | undefined; try { await new Promise((res, rej) => { // For some reason everything is output on stderr, not stdout - childProcess.stderr?.on("data", (data: Buffer) => { + buildProcess.stderr?.on("data", (data: Buffer) => { const text = data.toString(); errors.push(text); logger.debug(text); }); - childProcess.on("error", (e) => rej(e)); - childProcess.on("close", () => res()); + buildProcess.on("error", (e) => rej(e)); + buildProcess.on("close", () => res()); }); - const digest = extractImageDigest(errors); + digest = extractImageDigest(errors); span.setAttributes({ "image.digest": digest, }); - - span.end(); - - return { - ok: true as const, - image: options.imageTag, - digest, - }; } catch (e) { recordSpanException(span, e); @@ -725,6 +732,57 @@ async function buildAndPushSelfHostedImage( error: e instanceof Error ? e.message : JSON.stringify(e), }; } + + const pushArgs = ["push", imageRef].filter(Boolean) as string[]; + + logger.debug(`docker ${pushArgs.join(" ")}`); + + span.setAttribute("docker.command.push", `docker ${pushArgs.join(" ")}`); + + if (options.selfHostedRegistry || options.pushImage) { + // Push the image + const pushProcess = execa("docker", pushArgs, { + cwd: options.cwd, + }); + + try { + await new Promise((res, rej) => { + pushProcess.stdout?.on("data", (data: Buffer) => { + const text = data.toString(); + + logger.debug(text); + }); + + pushProcess.stderr?.on("data", (data: Buffer) => { + const text = data.toString(); + + logger.debug(text); + }); + + pushProcess.on("error", (e) => rej(e)); + pushProcess.on("close", () => res()); + }); + + span.end(); + } catch (e) { + recordSpanException(span, e); + + span.end(); + + return { + ok: false as const, + error: e instanceof Error ? e.message : JSON.stringify(e), + }; + } + } + + span.end(); + + return { + ok: true as const, + image: options.imageTag, + digest, + }; }); } diff --git a/packages/cli-v3/src/workers/prod/backgroundWorker.ts b/packages/cli-v3/src/workers/prod/backgroundWorker.ts index 4fc2bfc25..26ed569f0 100644 --- a/packages/cli-v3/src/workers/prod/backgroundWorker.ts +++ b/packages/cli-v3/src/workers/prod/backgroundWorker.ts @@ -2,6 +2,7 @@ import { BackgroundWorkerProperties, Config, CreateBackgroundWorkerResponse, + InferSocketMessageSchema, ProdChildToWorkerMessages, ProdTaskRunExecutionPayload, ProdWorkerToChildMessages, @@ -17,7 +18,6 @@ import { } from "@trigger.dev/core/v3"; import { Evt } from "evt"; import { ChildProcess, fork } from "node:child_process"; -import { safeDeleteFileSync } from "../../utilities/fileSystem"; import { UncaughtExceptionError } from "../common/errors"; class UnexpectedExitError extends Error { @@ -56,11 +56,19 @@ export class ProdBackgroundWorker { public onTaskHeartbeat: Evt = new Evt(); - public onWaitForDuration: Evt<{ version?: "v1"; ms: number }> = new Evt(); - public onWaitForTask: Evt<{ version?: "v1"; id: string }> = new Evt(); - public onWaitForBatch: Evt<{ version?: "v1"; id: string; runs: string[] }> = new Evt(); + public onWaitForBatch: Evt< + InferSocketMessageSchema + > = new Evt(); + public onWaitForDuration: Evt< + InferSocketMessageSchema + > = new Evt(); + public onWaitForTask: Evt< + InferSocketMessageSchema + > = new Evt(); public preCheckpointNotification = Evt.create<{ willCheckpointAndRestore: boolean }>(); + public onReadyForCheckpoint = Evt.create<{ version?: "v1" }>(); + public onCancelCheckpoint = Evt.create<{ version?: "v1" }>(); private _onClose: Evt = new Evt(); @@ -221,6 +229,15 @@ export class ProdBackgroundWorker { this.onWaitForTask.post(message); }); + taskRunProcess.onReadyForCheckpoint.attach((message) => { + this.onReadyForCheckpoint.post(message); + }); + + taskRunProcess.onCancelCheckpoint.attach((message) => { + this.onCancelCheckpoint.post(message); + }); + + // Notify down the chain this.preCheckpointNotification.attach((message) => { taskRunProcess.preCheckpointNotification.post(message); }); @@ -342,11 +359,19 @@ class TaskRunProcess { public onTaskHeartbeat: Evt = new Evt(); public onExit: Evt = new Evt(); - public onWaitForBatch: Evt<{ version?: "v1"; id: string; runs: string[] }> = new Evt(); - public onWaitForDuration: Evt<{ version?: "v1"; ms: number }> = new Evt(); - public onWaitForTask: Evt<{ version?: "v1"; id: string }> = new Evt(); + public onWaitForBatch: Evt< + InferSocketMessageSchema + > = new Evt(); + public onWaitForDuration: Evt< + InferSocketMessageSchema + > = new Evt(); + public onWaitForTask: Evt< + InferSocketMessageSchema + > = new Evt(); public preCheckpointNotification = Evt.create<{ willCheckpointAndRestore: boolean }>(); + public onReadyForCheckpoint = Evt.create<{ version?: "v1" }>(); + public onCancelCheckpoint = Evt.create<{ version?: "v1" }>(); constructor( private path: string, @@ -415,6 +440,12 @@ class TaskRunProcess { WAIT_FOR_TASK: async (message) => { this.onWaitForTask.post(message); }, + READY_FOR_CHECKPOINT: async (message) => { + this.onReadyForCheckpoint.post(message); + }, + CANCEL_CHECKPOINT: async (message) => { + this.onCancelCheckpoint.post(message); + }, }, }); diff --git a/packages/cli-v3/src/workers/prod/entry-point.ts b/packages/cli-v3/src/workers/prod/entry-point.ts index 6d5de2ac7..39a853798 100644 --- a/packages/cli-v3/src/workers/prod/entry-point.ts +++ b/packages/cli-v3/src/workers/prod/entry-point.ts @@ -3,12 +3,16 @@ import { CoordinatorToProdWorkerMessages, ProdWorkerToCoordinatorMessages, TaskResource, + WaitReason, ZodSocketConnection, } from "@trigger.dev/core/v3"; -import { HttpReply, getTextBody, SimpleLogger, getRandomPortNumber } from "@trigger.dev/core-apps"; +import { HttpReply, SimpleLogger, getRandomPortNumber } from "@trigger.dev/core-apps"; +import { readFile } from "node:fs/promises"; import { createServer } from "node:http"; +import { z } from "zod"; import { ProdBackgroundWorker } from "./backgroundWorker"; import { UncaughtExceptionError } from "../common/errors"; +import { setTimeout } from "node:timers/promises"; declare const __PROJECT_CONFIG__: Config; @@ -28,14 +32,17 @@ class ProdWorker { private projectRef = process.env.TRIGGER_PROJECT_REF!; private envId = process.env.TRIGGER_ENV_ID!; private runId = process.env.TRIGGER_RUN_ID || "index-only"; - private attemptId = process.env.TRIGGER_ATTEMPT_ID || "index-only"; private deploymentId = process.env.TRIGGER_DEPLOYMENT_ID!; + private deploymentVersion = process.env.TRIGGER_DEPLOYMENT_VERSION!; + private runningInKubernetes = !!process.env.KUBERNETES_PORT; private executing = false; private completed = new Set(); private paused = false; + private attemptFriendlyId?: string; - private nextResumeAfter: "WAIT_FOR_DURATION" | "WAIT_FOR_TASK" | "WAIT_FOR_BATCH" | undefined; + private nextResumeAfter?: WaitReason; + private waitForPostStart = false; #httpPort: number; #backgroundWorker: ProdBackgroundWorker; @@ -49,7 +56,7 @@ class ProdWorker { port: number, private host = "0.0.0.0" ) { - this.#coordinatorSocket = this.#createCoordinatorSocket(); + this.#coordinatorSocket = this.#createCoordinatorSocket(COORDINATOR_HOST); this.#backgroundWorker = new ProdBackgroundWorker("worker.js", { projectConfig: __PROJECT_CONFIG__, @@ -68,76 +75,154 @@ class ProdWorker { this.#coordinatorSocket.socket.emit("TASK_HEARTBEAT", { version: "v1", attemptFriendlyId }); }); + this.#backgroundWorker.onReadyForCheckpoint.attach(async (message) => { + this.#coordinatorSocket.socket.emit("READY_FOR_CHECKPOINT", { version: "v1" }); + }); + + this.#backgroundWorker.onCancelCheckpoint.attach(async (message) => { + logger.log("onCancelCheckpoint() clearing paused state, don't wait for post start hook", { + paused: this.paused, + nextResumeAfter: this.nextResumeAfter, + waitForPostStart: this.waitForPostStart, + }); + + this.paused = false; + this.nextResumeAfter = undefined; + this.waitForPostStart = false; + + this.#coordinatorSocket.socket.emit("CANCEL_CHECKPOINT", { version: "v1" }); + }); + this.#backgroundWorker.onWaitForDuration.attach(async (message) => { - // TODO: Switch to .send() once coordinator uses zod handler for all messages + if (!this.attemptFriendlyId) { + logger.error("Failed to send wait message, attempt friendly ID not set", { message }); + return; + } + const { willCheckpointAndRestore } = await this.#coordinatorSocket.socket.emitWithAck( "WAIT_FOR_DURATION", - { version: "v1", ...message } + { + ...message, + attemptFriendlyId: this.attemptFriendlyId, + } ); - logger.log("WAIT_FOR_DURATION", { willCheckpointAndRestore }); - - this.#backgroundWorker.preCheckpointNotification.post({ willCheckpointAndRestore }); - - setTimeout(() => { - if (willCheckpointAndRestore) { - this.paused = true; - this.nextResumeAfter = "WAIT_FOR_DURATION"; - } - // Forcing a reconnect will ensure the connection handler runs to trigger automatic resume - this.#coordinatorSocket.close(); - this.#coordinatorSocket.connect(); - }, 3_000); + this.#prepareForWait("WAIT_FOR_DURATION", willCheckpointAndRestore); }); this.#backgroundWorker.onWaitForTask.attach(async (message) => { - // TODO: Switch to .send() once coordinator uses zod handler for all messages + if (!this.attemptFriendlyId) { + logger.error("Failed to send wait message, attempt friendly ID not set", { message }); + return; + } + const { willCheckpointAndRestore } = await this.#coordinatorSocket.socket.emitWithAck( "WAIT_FOR_TASK", - { version: "v1", ...message } + { + ...message, + attemptFriendlyId: this.attemptFriendlyId, + } ); - logger.log("WAIT_FOR_TASK", { willCheckpointAndRestore }); - - this.#backgroundWorker.preCheckpointNotification.post({ willCheckpointAndRestore }); - - setTimeout(() => { - if (willCheckpointAndRestore) { - this.paused = true; - this.nextResumeAfter = "WAIT_FOR_TASK"; - } - // Forcing a reconnect will ensure the connection handler runs to trigger automatic resume - this.#coordinatorSocket.close(); - this.#coordinatorSocket.connect(); - }, 3_000); + this.#prepareForWait("WAIT_FOR_TASK", willCheckpointAndRestore); }); this.#backgroundWorker.onWaitForBatch.attach(async (message) => { - // TODO: Switch to .send() once coordinator uses zod handler for all messages + if (!this.attemptFriendlyId) { + logger.error("Failed to send wait message, attempt friendly ID not set", { message }); + return; + } + const { willCheckpointAndRestore } = await this.#coordinatorSocket.socket.emitWithAck( "WAIT_FOR_BATCH", - { version: "v1", ...message } + { + ...message, + attemptFriendlyId: this.attemptFriendlyId, + } ); - logger.log("WAIT_FOR_BATCH", { willCheckpointAndRestore }); - - this.#backgroundWorker.preCheckpointNotification.post({ willCheckpointAndRestore }); - - setTimeout(() => { - if (willCheckpointAndRestore) { - this.paused = true; - this.nextResumeAfter = "WAIT_FOR_BATCH"; - } - // Forcing a reconnect will ensure the connection handler runs to trigger automatic resume - this.#coordinatorSocket.close(); - this.#coordinatorSocket.connect(); - }, 3_000); + this.#prepareForWait("WAIT_FOR_BATCH", willCheckpointAndRestore); }); this.#httpPort = port; this.#httpServer = this.#createHttpServer(); } + async #reconnect(isPostStart = false) { + if (isPostStart) { + this.waitForPostStart = false; + } + + this.#coordinatorSocket.close(); + + if (!this.runningInKubernetes) { + this.#coordinatorSocket.connect(); + return; + } + + try { + const coordinatorHost = (await readFile("/etc/taskinfo/coordinator-host", "utf-8")).replace( + "\n", + "" + ); + + logger.log("reconnecting", { + coordinatorHost: { + fromEnv: COORDINATOR_HOST, + fromVolume: coordinatorHost, + current: this.#coordinatorSocket.socket.io.opts.hostname, + }, + }); + + this.#coordinatorSocket = this.#createCoordinatorSocket(coordinatorHost); + } catch (error) { + logger.error("taskinfo read error during reconnect", { error }); + this.#coordinatorSocket.connect(); + } + } + + #prepareForWait(reason: WaitReason, willCheckpointAndRestore: boolean) { + logger.log(`prepare for ${reason}`, { willCheckpointAndRestore }); + + this.#backgroundWorker.preCheckpointNotification.post({ willCheckpointAndRestore }); + + if (willCheckpointAndRestore) { + this.paused = true; + this.nextResumeAfter = reason; + this.waitForPostStart = true; + } + } + + async #prepareForRetry(willCheckpointAndRestore: boolean, shouldExit: boolean) { + logger.log("prepare for retry", { willCheckpointAndRestore, shouldExit }); + + // Graceful shutdown on final attempt + if (shouldExit) { + if (willCheckpointAndRestore) { + logger.log("WARNING: Will checkpoint but also requested exit. This won't end well."); + } + + await this.#backgroundWorker.close(); + process.exit(0); + } + + this.executing = false; + this.attemptFriendlyId = undefined; + + if (willCheckpointAndRestore) { + this.waitForPostStart = true; + this.#coordinatorSocket.socket.emit("READY_FOR_CHECKPOINT", { version: "v1" }); + return; + } + } + + #resumeAfterDuration() { + this.paused = false; + this.nextResumeAfter = undefined; + + this.#backgroundWorker.waitCompletedNotification(); + } + #returnValidatedExtraHeaders(headers: Record) { for (const [key, value] of Object.entries(headers)) { if (value === undefined) { @@ -148,33 +233,82 @@ class ProdWorker { return headers; } - #createCoordinatorSocket() { + #createCoordinatorSocket(host: string) { const extraHeaders = this.#returnValidatedExtraHeaders({ "x-machine-name": MACHINE_NAME, "x-pod-name": POD_NAME, "x-trigger-content-hash": this.contentHash, "x-trigger-project-ref": this.projectRef, - "x-trigger-attempt-id": this.attemptId, "x-trigger-env-id": this.envId, "x-trigger-deployment-id": this.deploymentId, "x-trigger-run-id": this.runId, + "x-trigger-deployment-version": this.deploymentVersion, }); + if (this.attemptFriendlyId) { + extraHeaders["x-trigger-attempt-friendly-id"] = this.attemptFriendlyId; + } + logger.log("connecting to coordinator", { - host: COORDINATOR_HOST, + host, port: COORDINATOR_PORT, extraHeaders, }); const coordinatorConnection = new ZodSocketConnection({ namespace: "prod-worker", - host: COORDINATOR_HOST, + host, port: COORDINATOR_PORT, clientMessages: ProdWorkerToCoordinatorMessages, serverMessages: CoordinatorToProdWorkerMessages, extraHeaders, handlers: { - RESUME: async (message) => { + RESUME_AFTER_DEPENDENCY: async (message) => { + if (!this.paused) { + logger.error("worker not paused", { + completions: message.completions, + executions: message.executions, + }); + return; + } + + if (message.completions.length !== message.executions.length) { + logger.error("did not receive the same number of completions and executions", { + completions: message.completions, + executions: message.executions, + }); + return; + } + + if (message.completions.length === 0 || message.executions.length === 0) { + logger.error("no completions or executions", { + completions: message.completions, + executions: message.executions, + }); + return; + } + + if ( + this.nextResumeAfter !== "WAIT_FOR_TASK" && + this.nextResumeAfter !== "WAIT_FOR_BATCH" + ) { + logger.error("not waiting to resume after dependency", { + nextResumeAfter: this.nextResumeAfter, + }); + return; + } + + if (this.nextResumeAfter === "WAIT_FOR_TASK" && message.completions.length > 1) { + logger.error("waiting for single task but got multiple completions", { + completions: message.completions, + executions: message.executions, + }); + return; + } + + this.paused = false; + this.nextResumeAfter = undefined; + for (let i = 0; i < message.completions.length; i++) { const completion = message.completions[i]; const execution = message.executions[i]; @@ -185,7 +319,21 @@ class ProdWorker { } }, RESUME_AFTER_DURATION: async (message) => { - this.#backgroundWorker.waitCompletedNotification(); + if (!this.paused) { + logger.error("worker not paused", { + attemptId: message.attemptId, + }); + return; + } + + if (this.nextResumeAfter !== "WAIT_FOR_DURATION") { + logger.error("not waiting to resume after duration", { + nextResumeAfter: this.nextResumeAfter, + }); + return; + } + + this.#resumeAfterDuration(); }, EXECUTE_TASK_RUN: async ({ executionPayload }) => { if (this.executing) { @@ -199,34 +347,25 @@ class ProdWorker { } this.executing = true; + this.attemptFriendlyId = executionPayload.execution.attempt.id; const completion = await this.#backgroundWorker.executeTaskRun(executionPayload); logger.log("completed", completion); this.completed.add(executionPayload.execution.attempt.id); - this.executing = false; await this.#backgroundWorker.flushTelemetry(); - const { didCheckpoint, shouldExit } = await this.#coordinatorSocket.socket.emitWithAck( - "TASK_RUN_COMPLETED", - { + const { willCheckpointAndRestore, shouldExit } = + await this.#coordinatorSocket.socket.emitWithAck("TASK_RUN_COMPLETED", { version: "v1", execution: executionPayload.execution, completion, - } - ); + }); - logger.log("completion acknowledged", { didCheckpoint, shouldExit }); + logger.log("completion acknowledged", { willCheckpointAndRestore, shouldExit }); - if (shouldExit) { - await this.#backgroundWorker.close(); - process.exit(0); - } - - // Forcing a reconnect will ensure the connection handler runs and signals we are ready for another execution - this.#coordinatorSocket.close(); - this.#coordinatorSocket.connect(); + this.#prepareForRetry(willCheckpointAndRestore, shouldExit); }, REQUEST_ATTEMPT_CANCELLATION: async (message) => { if (!this.executing) { @@ -236,10 +375,27 @@ class ProdWorker { await this.#backgroundWorker.cancelAttempt(message.attemptId); }, REQUEST_EXIT: async () => { + this.#coordinatorSocket.close(); process.exit(0); }, + READY_FOR_RETRY: async (message) => { + if (this.completed.size < 1) { + return; + } + + this.#coordinatorSocket.socket.emit("READY_FOR_EXECUTION", { + version: "v1", + runId: this.runId, + totalCompletions: this.completed.size, + }); + }, }, onConnection: async (socket, handler, sender, logger) => { + if (this.waitForPostStart) { + logger.log("skip connection handler, waiting for post start hook"); + return; + } + if (process.env.INDEX_TASKS === "true") { try { const taskResources = await this.#initializeWorker(); @@ -251,15 +407,15 @@ class ProdWorker { }); if (success) { - logger("indexing done, shutting down.."); + logger.info("indexing done, shutting down.."); process.exit(0); } else { - logger("indexing failure, shutting down.."); + logger.info("indexing failure, shutting down.."); process.exit(1); } } catch (e) { if (e instanceof UncaughtExceptionError) { - logger("uncaught exception", e.originalError.message); + logger.error("uncaught exception", { message: e.originalError.message }); socket.emit("INDEXING_FAILED", { version: "v1", @@ -271,7 +427,7 @@ class ProdWorker { }, }); } else if (e instanceof Error) { - logger("error", e.message); + logger.error("error", { message: e.message }); socket.emit("INDEXING_FAILED", { version: "v1", @@ -283,7 +439,7 @@ class ProdWorker { }, }); } else if (typeof e === "string") { - logger("string error", e); + logger.error("string error", { message: e }); socket.emit("INDEXING_FAILED", { version: "v1", @@ -294,7 +450,7 @@ class ProdWorker { }, }); } else { - logger("unknown error", e); + logger.error("unknown error", { error: e }); socket.emit("INDEXING_FAILED", { version: "v1", @@ -306,9 +462,8 @@ class ProdWorker { }); } - setTimeout(() => { - process.exit(1); - }, 200); + await setTimeout(200); + process.exit(1); } } @@ -317,15 +472,22 @@ class ProdWorker { return; } + if (!this.attemptFriendlyId) { + logger.error("Missing friendly ID"); + return; + } + + if (this.nextResumeAfter === "WAIT_FOR_DURATION") { + this.#resumeAfterDuration(); + return; + } + socket.emit("READY_FOR_RESUME", { version: "v1", - attemptId: this.attemptId, + attemptFriendlyId: this.attemptFriendlyId, type: this.nextResumeAfter, }); - this.#backgroundWorker.waitCompletedNotification(); - this.paused = false; - return; } @@ -335,10 +497,23 @@ class ProdWorker { socket.emit("READY_FOR_EXECUTION", { version: "v1", - attemptId: this.attemptId, runId: this.runId, + totalCompletions: this.completed.size, }); }, + onError: async (socket, err, logger) => { + logger.error("onError", { + error: { + name: err.name, + message: err.message, + }, + }); + + await this.#reconnect(); + }, + onDisconnect: async (socket, reason, description, logger) => { + // this.#reconnect(); + }, }); return coordinatorConnection; @@ -347,84 +522,117 @@ class ProdWorker { #createHttpServer() { const httpServer = createServer(async (req, res) => { logger.log(`[${req.method}]`, req.url); - const reply = new HttpReply(res); - switch (req.url) { - case "/complete": - setTimeout(() => process.exit(0), 1000); - return reply.text("ok"); + try { + const url = new URL(req.url ?? "", `http://${req.headers.host}`); - case "/date": - const date = new Date(); - return reply.text(date.toString()); + switch (url.pathname) { + case "/health": { + return reply.text("ok"); + } - case "/fail": - setTimeout(() => process.exit(1), 1000); - return reply.text("ok"); + case "/status": { + return reply.json({ + executing: this.executing, + pause: this.paused, + nextResumeAfter: this.nextResumeAfter, + }); + } - case "/health": - return reply.text("ok"); + case "/connect": { + this.#coordinatorSocket.connect(); - case "/whoami": - return reply.text(this.contentHash); + return reply.text("Connected to coordinator"); + } - case "/wait": - const { willCheckpointAndRestore } = await this.#coordinatorSocket.sendWithAck( - "WAIT_FOR_DURATION", - { + case "/close": { + await this.#coordinatorSocket.sendWithAck("LOG", { version: "v1", - ms: 60_000, - } - ); - logger.log("WAIT_FOR_DURATION", { willCheckpointAndRestore }); - // this is required when C/Ring established connections - this.#coordinatorSocket.close(); - return reply.text("sent WAIT"); + text: `[${req.method}] ${req.url}`, + }); - case "/connect": - this.#coordinatorSocket.connect(); - return reply.empty(); - - case "/close": - this.#coordinatorSocket.sendWithAck("LOG", { - version: "v1", - text: "close without delay", - }); - this.#coordinatorSocket.close(); - return reply.empty(); - - case "/close-delay": - this.#coordinatorSocket.sendWithAck("LOG", { - version: "v1", - text: "close with delay", - }); - setTimeout(() => { this.#coordinatorSocket.close(); - }, 200); - return reply.empty(); - case "/log": - this.#coordinatorSocket.sendWithAck("LOG", { - version: "v1", - text: await getTextBody(req), - }); - return reply.empty(); + return reply.text("Disconnected from coordinator"); + } - case "/preStop": - logger.log("should do preStop stuff, e.g. checkpoint and graceful shutdown"); - return reply.text("got preStop request"); + case "/test": { + await this.#coordinatorSocket.sendWithAck("LOG", { + version: "v1", + text: `[${req.method}] ${req.url}`, + }); - case "/ready": - this.#coordinatorSocket.send("READY_FOR_EXECUTION", { - version: "v1", - attemptId: this.attemptId, - runId: this.runId, - }); - return reply.empty(); + return reply.text("Received ACK from coordinator"); + } - default: - return reply.empty(404); + case "/preStop": { + const schema = z.enum(["index", "create", "restore"]); + + const cause = schema.safeParse(url.searchParams.get("cause")); + + if (!cause.success) { + logger.error("Failed to parse cause", { cause }); + return; + } + + switch (cause.data) { + case "index": { + break; + } + case "create": { + break; + } + case "restore": { + break; + } + default: { + logger.error("Unhandled cause", { cause: cause.data }); + break; + } + } + logger.log("preStop", { url: req.url }); + + return reply.text("preStop ok"); + } + + case "/postStart": { + const schema = z.enum(["index", "create", "restore"]); + + const cause = schema.safeParse(url.searchParams.get("cause")); + + if (!cause.success) { + logger.error("Failed to parse cause", { cause }); + return; + } + + switch (cause.data) { + case "index": { + break; + } + case "create": { + break; + } + case "restore": { + await this.#reconnect(true); + break; + } + default: { + logger.error("Unhandled cause", { cause: cause.data }); + break; + } + } + + return reply.text("postStart ok"); + } + + default: { + return reply.empty(404); + } + } + } catch (error) { + logger.error("HTTP server error", { error }); + reply.empty(500); } }); @@ -436,7 +644,7 @@ class ProdWorker { logger.log("http server listening on port", this.#httpPort); }); - httpServer.on("error", (error) => { + httpServer.on("error", async (error) => { // @ts-expect-error if (error.code != "EADDRINUSE") { return; @@ -446,9 +654,8 @@ class ProdWorker { this.#httpPort = getRandomPortNumber(); - setTimeout(() => { - this.start(); - }, 100); + await setTimeout(100); + this.start(); }); return httpServer; diff --git a/packages/core-apps/package.json b/packages/core-apps/package.json index bebf33384..5eeb914b6 100644 --- a/packages/core-apps/package.json +++ b/packages/core-apps/package.json @@ -21,17 +21,12 @@ "./package.json": "./package.json" }, "scripts": { - "clean": "rimraf dist", - "build": "npm run clean && npm run build:tsup", - "build:tsup": "tsup --dts-resolve", "typecheck": "tsc --noEmit" }, "devDependencies": { + "@trigger.dev/core": "workspace:*", "@trigger.dev/tsconfig": "workspace:*", - "@trigger.dev/tsup": "workspace:*", "@types/node": "18", - "rimraf": "^3.0.2", - "tsup": "^8.0.1", "typescript": "^5.3.0" }, "engines": { diff --git a/packages/core-apps/src/http.ts b/packages/core-apps/src/http.ts index 60eeb67a0..0a7a8560d 100644 --- a/packages/core-apps/src/http.ts +++ b/packages/core-apps/src/http.ts @@ -26,6 +26,14 @@ export class HttpReply { .writeHead(status ?? 200, { "Content-Type": contentType || "text/plain" }) .end(text.endsWith("\n") ? text : `${text}\n`); } + + json(value: any, pretty?: boolean) { + return this.text( + JSON.stringify(value, undefined, pretty ? 2 : undefined), + 200, + "application/json" + ); + } } function getRandomInteger(min: number, max: number) { diff --git a/packages/core-apps/src/provider.ts b/packages/core-apps/src/provider.ts index 1b120664b..d304b4fa5 100644 --- a/packages/core-apps/src/provider.ts +++ b/packages/core-apps/src/provider.ts @@ -2,6 +2,7 @@ import { createServer } from "node:http"; import { ClientToSharedQueueMessages, clientWebsocketMessages, + Machine, PlatformToProviderMessages, ProviderToPlatformMessages, SharedQueueToClientMessages, @@ -20,12 +21,36 @@ const PLATFORM_SECRET = process.env.PLATFORM_SECRET || "provider-secret"; const logger = new SimpleLogger(`[${MACHINE_NAME}]`); +export interface TaskOperationsIndexOptions { + shortCode: string; + imageRef: string; + envId: string; + apiKey: string; + apiUrl: string; +} + +export interface TaskOperationsCreateOptions { + runId: string; + image: string; + machine: Machine; + envId: string; + version: string; +} + +export interface TaskOperationsRestoreOptions { + runId: string; + imageRef: string; + checkpointRef: string; + machine: Machine; +} + export interface TaskOperations { - create: (...args: any[]) => Promise; - restore: (...args: any[]) => Promise; + index: (opts: TaskOperationsIndexOptions) => Promise; + create: (opts: TaskOperationsCreateOptions) => Promise; + restore: (opts: TaskOperationsRestoreOptions) => Promise; + delete: (...args: any[]) => Promise; get: (...args: any[]) => Promise; - index: (...args: any[]) => Promise; } type ProviderShellOptions = { @@ -74,13 +99,17 @@ export class ProviderShell implements Provider { }, BACKGROUND_WORKER_MESSAGE: async (message) => { if (message.data.type === "SCHEDULE_ATTEMPT") { - this.tasks.create({ - envId: message.data.envId, - runId: message.data.runId, - attemptId: message.data.id, - image: message.data.image, - machine: {}, - }); + try { + this.tasks.create({ + envId: message.data.envId, + runId: message.data.runId, + image: message.data.image, + machine: {}, + version: message.version, + }); + } catch (error) { + logger.error("create failed", error); + } } }, }, @@ -134,8 +163,8 @@ export class ProviderShell implements Provider { INDEX: async (message) => { try { await this.tasks.index({ - contentHash: message.contentHash, - imageTag: message.imageTag, + shortCode: message.shortCode, + imageRef: message.imageTag, envId: message.envId, apiKey: message.apiKey, apiUrl: message.apiUrl, @@ -178,10 +207,12 @@ export class ProviderShell implements Provider { try { await this.tasks.restore({ runId: message.runId, - attemptId: message.attemptId, checkpointRef: message.location, - // TODO - // machine: message.machine, + machine: { + cpu: "1", + memory: "100Mi", + }, + imageRef: message.imageRef, }); } catch (error) { logger.error("restore failed", error); @@ -221,13 +252,14 @@ export class ProviderShell implements Provider { const body = await getTextBody(req); await this.tasks.create({ - attemptId: body, envId: "placeholder", image: body, machine: { cpu: "1", memory: "100Mi", }, + runId: "", + version: "", }); return reply.text(`sent restore request: ${body}`); diff --git a/packages/core-apps/tsconfig.build.json b/packages/core-apps/tsconfig.build.json deleted file mode 100644 index f444efc19..000000000 --- a/packages/core-apps/tsconfig.build.json +++ /dev/null @@ -1,11 +0,0 @@ -{ - "extends": "@trigger.dev/tsconfig/node18.json", - "include": ["src/globals.d.ts", "./src/**/*.ts", "tsup.config.ts"], - "compilerOptions": { - "experimentalDecorators": true, - "emitDecoratorMetadata": true, - "declaration": false, - "declarationMap": false - }, - "exclude": ["node_modules"] -} diff --git a/packages/core-apps/tsup.config.ts b/packages/core-apps/tsup.config.ts deleted file mode 100644 index 5a0b9c86a..000000000 --- a/packages/core-apps/tsup.config.ts +++ /dev/null @@ -1,9 +0,0 @@ -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);", - }, -}); diff --git a/packages/core/src/v3/runtime/prodRuntimeManager.ts b/packages/core/src/v3/runtime/prodRuntimeManager.ts index 0f9889bb6..c7359802d 100644 --- a/packages/core/src/v3/runtime/prodRuntimeManager.ts +++ b/packages/core/src/v3/runtime/prodRuntimeManager.ts @@ -49,6 +49,8 @@ export class ProdRuntimeManager implements RuntimeManager { async waitForDuration(ms: number): Promise { let timeout: NodeJS.Timeout | undefined; + const now = Date.now(); + const resolveAfterDuration = new Promise((resolve) => { timeout = setTimeout(resolve, ms); }); @@ -58,21 +60,28 @@ export class ProdRuntimeManager implements RuntimeManager { return; } - const waitForRestore = new Promise((resolve, reject) => { + const waitForRestore = new Promise((resolve, reject) => { this._waitForRestore = { resolve, reject }; }); - // There is a slight delay before actually checkpointing, so this has a chance to return - const { willCheckpointAndRestore } = await this.ipc.sendWithAck("WAIT_FOR_DURATION", { ms }); + const { willCheckpointAndRestore } = await this.ipc.sendWithAck("WAIT_FOR_DURATION", { + ms, + now, + }); if (!willCheckpointAndRestore) { await resolveAfterDuration; return; } - // Checkpointing should happen after this line + this.ipc.send("READY_FOR_CHECKPOINT", {}); + + // Don't wait for checkpoint beyond the requested wait duration + await Promise.race([waitForRestore, resolveAfterDuration]); + + // The coordinator can then cancel any in-progress checkpoints + this.ipc.send("CANCEL_CHECKPOINT", {}); - await waitForRestore; clearTimeout(timeout); } @@ -95,7 +104,7 @@ export class ProdRuntimeManager implements RuntimeManager { }); await this.ipc.send("WAIT_FOR_TASK", { - id: params.id, + friendlyId: params.id, }); return await promise; @@ -119,8 +128,8 @@ export class ProdRuntimeManager implements RuntimeManager { ); await this.ipc.send("WAIT_FOR_BATCH", { - id: params.id, - runs: params.runs, + batchFriendlyId: params.id, + runFriendlyIds: params.runs, }); const results = await promise; diff --git a/packages/core/src/v3/schemas/messages.ts b/packages/core/src/v3/schemas/messages.ts index 2cb724d43..f37ad1c27 100644 --- a/packages/core/src/v3/schemas/messages.ts +++ b/packages/core/src/v3/schemas/messages.ts @@ -43,6 +43,7 @@ export const BackgroundWorkerServerMessages = z.discriminatedUnion("type", [ image: z.string(), envId: z.string(), runId: z.string(), + version: z.string(), }), ]); @@ -302,10 +303,21 @@ export const ProdChildToWorkerMessages = { READY_TO_DISPOSE: { message: z.undefined(), }, + READY_FOR_CHECKPOINT: { + message: z.object({ + version: z.literal("v1").default("v1"), + }), + }, + CANCEL_CHECKPOINT: { + message: z.object({ + version: z.literal("v1").default("v1"), + }), + }, WAIT_FOR_DURATION: { message: z.object({ version: z.literal("v1").default("v1"), ms: z.number(), + now: z.number(), }), callback: z.object({ willCheckpointAndRestore: z.boolean(), @@ -314,14 +326,14 @@ export const ProdChildToWorkerMessages = { WAIT_FOR_TASK: { message: z.object({ version: z.literal("v1").default("v1"), - id: z.string(), + friendlyId: z.string(), }), }, WAIT_FOR_BATCH: { message: z.object({ version: z.literal("v1").default("v1"), - id: z.string(), - runs: z.string().array(), + batchFriendlyId: z.string(), + runFriendlyIds: z.string().array(), }), }, UNCAUGHT_EXCEPTION: { diff --git a/packages/core/src/v3/schemas/schemas.ts b/packages/core/src/v3/schemas/schemas.ts index 06640caec..710a49329 100644 --- a/packages/core/src/v3/schemas/schemas.ts +++ b/packages/core/src/v3/schemas/schemas.ts @@ -37,6 +37,10 @@ export const Machine = z.object({ export type Machine = z.infer; +export const WaitReason = z.enum(["WAIT_FOR_DURATION", "WAIT_FOR_TASK", "WAIT_FOR_BATCH"]); + +export type WaitReason = z.infer; + export const ProviderToPlatformMessages = { LOG: { message: z.object({ @@ -68,7 +72,7 @@ export const PlatformToProviderMessages = { message: z.object({ version: z.literal("v1").default("v1"), imageTag: z.string(), - contentHash: z.string(), + shortCode: z.string(), envId: z.string(), apiKey: z.string(), apiUrl: z.string(), @@ -99,10 +103,10 @@ export const PlatformToProviderMessages = { version: z.literal("v1").default("v1"), checkpointId: z.string(), runId: z.string(), - attemptId: z.string(), type: z.enum(["DOCKER", "KUBERNETES"]), location: z.string(), reason: z.string().optional(), + imageRef: z.string(), }), }, DELETE: { @@ -155,8 +159,8 @@ export const CoordinatorToPlatformMessages = { READY_FOR_EXECUTION: { message: z.object({ version: z.literal("v1").default("v1"), - attemptId: z.string(), runId: z.string(), + totalCompletions: z.number(), }), callback: z.discriminatedUnion("success", [ z.object({ @@ -171,8 +175,8 @@ export const CoordinatorToPlatformMessages = { READY_FOR_RESUME: { message: z.object({ version: z.literal("v1").default("v1"), - attemptId: z.string(), - type: z.enum(["WAIT_FOR_DURATION", "WAIT_FOR_TASK", "WAIT_FOR_BATCH"]), + attemptFriendlyId: z.string(), + type: WaitReason, }), }, TASK_RUN_COMPLETED: { @@ -180,6 +184,12 @@ export const CoordinatorToPlatformMessages = { version: z.literal("v1").default("v1"), execution: ProdTaskRunExecution, completion: TaskRunExecutionResult, + checkpoint: z + .object({ + docker: z.boolean(), + location: z.string(), + }) + .optional(), }), }, TASK_HEARTBEAT: { @@ -191,21 +201,23 @@ export const CoordinatorToPlatformMessages = { CHECKPOINT_CREATED: { message: z.object({ version: z.literal("v1").default("v1"), - attemptId: z.string(), + attemptFriendlyId: z.string(), docker: z.boolean(), location: z.string(), reason: z.discriminatedUnion("type", [ z.object({ type: z.literal("WAIT_FOR_DURATION"), ms: z.number(), + now: z.number(), }), z.object({ type: z.literal("WAIT_FOR_BATCH"), - id: z.string(), + batchFriendlyId: z.string(), + runFriendlyIds: z.string().array(), }), z.object({ type: z.literal("WAIT_FOR_TASK"), - id: z.string(), + friendlyId: z.string(), }), z.object({ type: z.literal("RETRYING_AFTER_FAILURE"), @@ -228,10 +240,12 @@ export const CoordinatorToPlatformMessages = { }; export const PlatformToCoordinatorMessages = { - RESUME: { + RESUME_AFTER_DEPENDENCY: { message: z.object({ version: z.literal("v1").default("v1"), + runId: z.string(), attemptId: z.string(), + attemptFriendlyId: z.string(), completions: TaskRunExecutionResult.array(), executions: TaskRunExecution.array(), }), @@ -240,12 +254,20 @@ export const PlatformToCoordinatorMessages = { message: z.object({ version: z.literal("v1").default("v1"), attemptId: z.string(), + attemptFriendlyId: z.string(), }), }, REQUEST_ATTEMPT_CANCELLATION: { message: z.object({ version: z.literal("v1").default("v1"), attemptId: z.string(), + attemptFriendlyId: z.string(), + }), + }, + READY_FOR_RETRY: { + message: z.object({ + version: z.literal("v1").default("v1"), + runId: z.string(), }), }, }; @@ -315,15 +337,25 @@ export const ProdWorkerToCoordinatorMessages = { READY_FOR_EXECUTION: { message: z.object({ version: z.literal("v1").default("v1"), - attemptId: z.string(), runId: z.string(), + totalCompletions: z.number(), }), }, READY_FOR_RESUME: { message: z.object({ version: z.literal("v1").default("v1"), - attemptId: z.string(), - type: z.enum(["WAIT_FOR_DURATION", "WAIT_FOR_TASK", "WAIT_FOR_BATCH"]), + attemptFriendlyId: z.string(), + type: WaitReason, + }), + }, + READY_FOR_CHECKPOINT: { + message: z.object({ + version: z.literal("v1").default("v1"), + }), + }, + CANCEL_CHECKPOINT: { + message: z.object({ + version: z.literal("v1").default("v1"), }), }, TASK_HEARTBEAT: { @@ -339,7 +371,7 @@ export const ProdWorkerToCoordinatorMessages = { completion: TaskRunExecutionResult, }), callback: z.object({ - didCheckpoint: z.boolean(), + willCheckpointAndRestore: z.boolean(), shouldExit: z.boolean(), }), }, @@ -347,6 +379,8 @@ export const ProdWorkerToCoordinatorMessages = { message: z.object({ version: z.literal("v1").default("v1"), ms: z.number(), + now: z.number(), + attemptFriendlyId: z.string(), }), callback: z.object({ willCheckpointAndRestore: z.boolean(), @@ -355,7 +389,9 @@ export const ProdWorkerToCoordinatorMessages = { WAIT_FOR_TASK: { message: z.object({ version: z.literal("v1").default("v1"), - id: z.string(), + friendlyId: z.string(), + // This is the attempt that is waiting + attemptFriendlyId: z.string(), }), callback: z.object({ willCheckpointAndRestore: z.boolean(), @@ -364,8 +400,10 @@ export const ProdWorkerToCoordinatorMessages = { WAIT_FOR_BATCH: { message: z.object({ version: z.literal("v1").default("v1"), - id: z.string(), - runs: z.string().array(), + batchFriendlyId: z.string(), + runFriendlyIds: z.string().array(), + // This is the attempt that is waiting + attemptFriendlyId: z.string(), }), callback: z.object({ willCheckpointAndRestore: z.boolean(), @@ -385,7 +423,7 @@ export const ProdWorkerToCoordinatorMessages = { }; export const CoordinatorToProdWorkerMessages = { - RESUME: { + RESUME_AFTER_DEPENDENCY: { message: z.object({ version: z.literal("v1").default("v1"), attemptId: z.string(), @@ -416,6 +454,12 @@ export const CoordinatorToProdWorkerMessages = { version: z.literal("v1").default("v1"), }), }, + READY_FOR_RETRY: { + message: z.object({ + version: z.literal("v1").default("v1"), + runId: z.string(), + }), + }, }; export const ProdWorkerSocketData = z.object({ @@ -423,7 +467,8 @@ export const ProdWorkerSocketData = z.object({ projectRef: z.string(), envId: z.string(), runId: z.string(), - attemptId: z.string(), + attemptFriendlyId: z.string().optional(), podName: z.string(), deploymentId: z.string(), + deploymentVersion: z.string(), }); diff --git a/packages/core/src/v3/zodMessageHandler.ts b/packages/core/src/v3/zodMessageHandler.ts index b4a64c1b6..4627cb29d 100644 --- a/packages/core/src/v3/zodMessageHandler.ts +++ b/packages/core/src/v3/zodMessageHandler.ts @@ -1,4 +1,5 @@ import { z } from "zod"; +import { StructuredLogger } from "./zodNamespace"; export type ZodMessageValueSchema> = | z.ZodFirstPartySchemaTypes @@ -92,17 +93,20 @@ export class ZodMessageHandler }; } - public registerHandlers(emitter: EventEmitterLike, logger?: (...args: any[]) => void) { - const log = logger ?? console.log; + public registerHandlers(emitter: EventEmitterLike, logger?: StructuredLogger) { + const log = logger ?? console; if (!this.#handlers) { - log("No handlers provided"); + log.info("No handlers provided"); return; } for (const eventName of Object.keys(this.#schema)) { emitter.on(eventName, async (message: any, callback?: any): Promise => { - log(`handling ${eventName}`, message); + log.info(`handling ${eventName}`, { + payload: message, + hasCallback: !!callback, + }); let ack; diff --git a/packages/core/src/v3/zodNamespace.ts b/packages/core/src/v3/zodNamespace.ts index d6dd5976c..fac0a109f 100644 --- a/packages/core/src/v3/zodNamespace.ts +++ b/packages/core/src/v3/zodNamespace.ts @@ -25,6 +25,87 @@ export type ZodNamespaceSocket< z.infer >; +type StructuredArgs = (Record | undefined)[]; + +export interface StructuredLogger { + log: (message: string, ...args: StructuredArgs) => any; + error: (message: string, ...args: StructuredArgs) => any; + warn: (message: string, ...args: StructuredArgs) => any; + info: (message: string, ...args: StructuredArgs) => any; + debug: (message: string, ...args: StructuredArgs) => any; + child: (fields: Record) => StructuredLogger; +} + +export enum LogLevel { + "log", + "error", + "warn", + "info", + "debug", +} + +export class SimpleStructuredLogger implements StructuredLogger { + constructor( + private name: string, + private level: LogLevel = ["1", "true"].includes(process.env.DEBUG ?? "") + ? LogLevel.debug + : LogLevel.info, + private fields?: Record + ) {} + + child(fields: Record, level?: LogLevel) { + return new SimpleStructuredLogger(this.name, level, { ...this.fields, ...fields }); + } + + log(message: string, ...args: StructuredArgs) { + if (this.level < LogLevel.log) return; + + this.#structuredLog(console.log, message, "log", ...args); + } + + error(message: string, ...args: StructuredArgs) { + if (this.level < LogLevel.error) return; + + this.#structuredLog(console.error, message, "error", ...args); + } + + warn(message: string, ...args: StructuredArgs) { + if (this.level < LogLevel.warn) return; + + this.#structuredLog(console.warn, message, "warn", ...args); + } + + info(message: string, ...args: StructuredArgs) { + if (this.level < LogLevel.info) return; + + this.#structuredLog(console.info, message, "info", ...args); + } + + debug(message: string, ...args: StructuredArgs) { + if (this.level < LogLevel.debug) return; + + this.#structuredLog(console.debug, message, "debug", ...args); + } + + #structuredLog( + loggerFunction: (message: string, ...args: any[]) => void, + message: string, + level: string, + ...args: StructuredArgs + ) { + const structuredLog = { + ...(args.length === 1 ? args[0] : args), + ...this.fields, + timestamp: new Date(), + name: this.name, + message, + level, + }; + + loggerFunction(JSON.stringify(structuredLog)); + } +} + interface ZodNamespaceOptions< TClientMessages extends ZodSocketMessageCatalogSchema, TServerMessages extends ZodSocketMessageCatalogSchema, @@ -38,32 +119,33 @@ interface ZodNamespaceOptions< socketData?: TSocketData; handlers?: ZodSocketMessageHandlers; authToken?: string; + logger?: StructuredLogger; preAuth?: ( socket: ZodNamespaceSocket, next: (err?: ExtendedError) => void, - logger: (...args: any[]) => void + logger: StructuredLogger ) => Promise; postAuth?: ( socket: ZodNamespaceSocket, next: (err?: ExtendedError) => void, - logger: (...args: any[]) => void + logger: StructuredLogger ) => Promise; onConnection?: ( socket: ZodNamespaceSocket, handler: ZodSocketMessageHandler, sender: ZodMessageSender, - logger: (...args: any[]) => void + logger: StructuredLogger ) => Promise; onDisconnect?: ( socket: ZodNamespaceSocket, reason: DisconnectReason, description: any, - logger: (...args: any[]) => void + logger: StructuredLogger ) => Promise; onError?: ( socket: ZodNamespaceSocket, err: Error, - logger: (...args: any[]) => void + logger: StructuredLogger ) => Promise; } @@ -73,6 +155,7 @@ export class ZodNamespace< TSocketData extends z.ZodObject = any, TServerSideEvents extends EventsMap = DefaultEventsMap, > { + #logger: StructuredLogger; #handler: ZodSocketMessageHandler; sender: ZodMessageSender; @@ -87,6 +170,8 @@ export class ZodNamespace< constructor( opts: ZodNamespaceOptions ) { + this.#logger = opts.logger ?? new SimpleStructuredLogger(opts.name); + this.#handler = new ZodSocketMessageHandler({ schema: opts.clientMessages, handlers: opts.handlers, @@ -114,7 +199,7 @@ export class ZodNamespace< if (opts.preAuth) { this.namespace.use(async (socket, next) => { - const logger = createLogger(`[${opts.name}][${socket.id}][preAuth]`); + const logger = this.#logger.child({ socketId: socket.id, socketStage: "preAuth" }); if (typeof opts.preAuth === "function") { await opts.preAuth(socket, next, logger); @@ -124,21 +209,21 @@ export class ZodNamespace< if (opts.authToken) { this.namespace.use((socket, next) => { - const logger = createLogger(`[${opts.name}][${socket.id}][auth]`); + const logger = this.#logger.child({ socketId: socket.id, socketStage: "auth" }); const { auth } = socket.handshake; if (!("token" in auth)) { - logger("no token"); + logger.error("no token"); return socket.disconnect(true); } if (auth.token !== opts.authToken) { - logger("invalid token"); + logger.error("invalid token"); return socket.disconnect(true); } - logger("success"); + logger.info("success"); next(); }); @@ -146,7 +231,7 @@ export class ZodNamespace< if (opts.postAuth) { this.namespace.use(async (socket, next) => { - const logger = createLogger(`[${opts.name}][${socket.id}][postAuth]`); + const logger = this.#logger.child({ socketId: socket.id, socketStage: "auth" }); if (typeof opts.postAuth === "function") { await opts.postAuth(socket, next, logger); @@ -155,13 +240,13 @@ export class ZodNamespace< } this.namespace.on("connection", async (socket) => { - const logger = createLogger(`[${opts.name}][${socket.id}]`); - logger("connection"); + const logger = this.#logger.child({ socketId: socket.id, socketStage: "connection" }); + logger.info("connected"); this.#handler.registerHandlers(socket, logger); socket.on("disconnect", async (reason, description) => { - logger("disconnect", { reason, description }); + logger.info("disconnect", { reason, description }); if (opts.onDisconnect) { await opts.onDisconnect(socket, reason, description, logger); @@ -169,7 +254,7 @@ export class ZodNamespace< }); socket.on("error", async (error) => { - logger("error", error); + logger.error("error", { error }); if (opts.onError) { await opts.onError(socket, error, logger); @@ -186,7 +271,3 @@ export class ZodNamespace< return this.namespace.fetchSockets(); } } - -function createLogger(prefix: string) { - return (...args: any[]) => console.log(prefix, ...args); -} diff --git a/packages/core/src/v3/zodSocket.ts b/packages/core/src/v3/zodSocket.ts index 68d489796..294cdc89f 100644 --- a/packages/core/src/v3/zodSocket.ts +++ b/packages/core/src/v3/zodSocket.ts @@ -1,6 +1,7 @@ import { io, Socket } from "socket.io-client"; import { z } from "zod"; import { EventEmitterLike, ZodMessageValueSchema } from "./zodMessageHandler"; +import { LogLevel, SimpleStructuredLogger, StructuredLogger } from "./zodNamespace"; export interface ZodSocketMessageCatalogSchema { [key: string]: @@ -137,27 +138,35 @@ export class ZodSocketMessageHandler void) { - const log = logger ?? console.log; + public registerHandlers(emitter: EventEmitterLike, logger?: StructuredLogger) { + const log = logger ?? console; if (!this.#handlers) { - log("No handlers provided"); + log.info("No handlers provided"); return; } for (const eventName of Object.keys(this.#handlers)) { emitter.on(eventName, async (message: any, callback?: any): Promise => { - log(`handling ${eventName}`, message); + log.info(`handling ${eventName}`, { + payload: message, + hasCallback: !!callback, + }); 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 }); + try { + // 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 }); + } + } catch (error) { + log.error("Error while handling message", { error }); + return; } if (callback && typeof callback === "function") { @@ -267,18 +276,18 @@ interface ZodSocketConnectionOptions< socket: ZodSocket, handler: ZodSocketMessageHandler, sender: ZodSocketMessageSender, - logger: (...args: any[]) => void + logger: StructuredLogger ) => Promise; onDisconnect?: ( socket: ZodSocket, reason: Socket.DisconnectReason, description: any, - logger: (...args: any[]) => void + logger: StructuredLogger ) => Promise; onError?: ( socket: ZodSocket, err: Error, - logger: (...args: any[]) => void + logger: StructuredLogger ) => Promise; } @@ -290,7 +299,7 @@ export class ZodSocketConnection< socket: ZodSocket; #handler: ZodSocketMessageHandler; - #logger: (...args: any[]) => void; + #logger: StructuredLogger; constructor(opts: ZodSocketConnectionOptions) { this.socket = io(`ws://${opts.host}:${opts.port}/${opts.namespace}`, { @@ -301,7 +310,9 @@ export class ZodSocketConnection< extraHeaders: opts.extraHeaders, }); - this.#logger = createLogger(`[${opts.namespace}][${this.socket.id}]`); + this.#logger = new SimpleStructuredLogger(opts.namespace, LogLevel.info, { + socketId: this.socket.id, + }); this.#handler = new ZodSocketMessageHandler({ schema: opts.serverMessages, @@ -315,7 +326,7 @@ export class ZodSocketConnection< }); this.socket.on("connect_error", async (error) => { - this.#logger(`connect_error: ${error}`); + this.#logger.error(`connect_error: ${error}`); if (opts.onError) { await opts.onError(this.socket, error, this.#logger); @@ -323,7 +334,7 @@ export class ZodSocketConnection< }); this.socket.on("connect", async () => { - this.#logger("connect"); + this.#logger.info("connect"); if (opts.onConnection) { await opts.onConnection(this.socket, this.#handler, this.#sender, this.#logger); @@ -331,7 +342,7 @@ export class ZodSocketConnection< }); this.socket.on("disconnect", async (reason, description) => { - this.#logger("disconnect"); + this.#logger.info("disconnect", { reason, description }); if (opts.onDisconnect) { await opts.onDisconnect(this.socket, reason, description, this.#logger); diff --git a/packages/database/prisma/migrations/20240318170823_add_image_ref_to_checkpoint/migration.sql b/packages/database/prisma/migrations/20240318170823_add_image_ref_to_checkpoint/migration.sql new file mode 100644 index 000000000..1f12a8b59 --- /dev/null +++ b/packages/database/prisma/migrations/20240318170823_add_image_ref_to_checkpoint/migration.sql @@ -0,0 +1,8 @@ +/* + Warnings: + + - Added the required column `imageRef` to the `Checkpoint` table without a default value. This is not possible if the table is not empty. + +*/ +-- AlterTable +ALTER TABLE "Checkpoint" ADD COLUMN "imageRef" TEXT NOT NULL; diff --git a/packages/database/prisma/migrations/20240322172035_add_checkpoint_event_to_dependencies/migration.sql b/packages/database/prisma/migrations/20240322172035_add_checkpoint_event_to_dependencies/migration.sql new file mode 100644 index 000000000..91fb3276a --- /dev/null +++ b/packages/database/prisma/migrations/20240322172035_add_checkpoint_event_to_dependencies/migration.sql @@ -0,0 +1,24 @@ +/* + Warnings: + + - A unique constraint covering the columns `[checkpointEventId]` on the table `BatchTaskRun` will be added. If there are existing duplicate values, this will fail. + - A unique constraint covering the columns `[checkpointEventId]` on the table `TaskRunDependency` will be added. If there are existing duplicate values, this will fail. + +*/ +-- AlterTable +ALTER TABLE "BatchTaskRun" ADD COLUMN "checkpointEventId" TEXT; + +-- AlterTable +ALTER TABLE "TaskRunDependency" ADD COLUMN "checkpointEventId" TEXT; + +-- CreateIndex +CREATE UNIQUE INDEX "BatchTaskRun_checkpointEventId_key" ON "BatchTaskRun"("checkpointEventId"); + +-- CreateIndex +CREATE UNIQUE INDEX "TaskRunDependency_checkpointEventId_key" ON "TaskRunDependency"("checkpointEventId"); + +-- AddForeignKey +ALTER TABLE "TaskRunDependency" ADD CONSTRAINT "TaskRunDependency_checkpointEventId_fkey" FOREIGN KEY ("checkpointEventId") REFERENCES "CheckpointRestoreEvent"("id") ON DELETE CASCADE ON UPDATE CASCADE; + +-- AddForeignKey +ALTER TABLE "BatchTaskRun" ADD CONSTRAINT "BatchTaskRun_checkpointEventId_fkey" FOREIGN KEY ("checkpointEventId") REFERENCES "CheckpointRestoreEvent"("id") ON DELETE CASCADE ON UPDATE CASCADE; diff --git a/packages/database/prisma/schema.prisma b/packages/database/prisma/schema.prisma index 750a98697..dd075354f 100644 --- a/packages/database/prisma/schema.prisma +++ b/packages/database/prisma/schema.prisma @@ -1653,6 +1653,9 @@ model TaskRunDependency { taskRun TaskRun @relation(fields: [taskRunId], references: [id], onDelete: Cascade, onUpdate: Cascade) taskRunId String @unique + checkpointEvent CheckpointRestoreEvent? @relation(fields: [checkpointEventId], references: [id], onDelete: Cascade, onUpdate: Cascade) + checkpointEventId String? @unique + /// An attempt that is dependent on this task run. dependentAttempt TaskRunAttempt? @relation("dependentAttempt", fields: [dependentAttemptId], references: [id]) dependentAttemptId String? @unique @@ -1878,6 +1881,9 @@ model BatchTaskRun { idempotencyKey String taskIdentifier String + checkpointEvent CheckpointRestoreEvent? @relation(fields: [checkpointEventId], references: [id], onDelete: Cascade, onUpdate: Cascade) + checkpointEventId String? @unique + runtimeEnvironment RuntimeEnvironment @relation(fields: [runtimeEnvironmentId], references: [id], onDelete: Cascade, onUpdate: Cascade) runtimeEnvironmentId String @@ -1959,6 +1965,7 @@ model Checkpoint { type CheckpointType location String + imageRef String reason String? metadata String? @@ -2007,6 +2014,9 @@ model CheckpointRestoreEvent { runtimeEnvironment RuntimeEnvironment @relation(fields: [runtimeEnvironmentId], references: [id], onDelete: Cascade, onUpdate: Cascade) runtimeEnvironmentId String + taskRunDependency TaskRunDependency? + batchTaskRunDependency BatchTaskRun? + createdAt DateTime @default(now()) updatedAt DateTime @updatedAt } diff --git a/packages/trigger-sdk/src/v3/wait.ts b/packages/trigger-sdk/src/v3/wait.ts index 9600d7a44..c3130af58 100644 --- a/packages/trigger-sdk/src/v3/wait.ts +++ b/packages/trigger-sdk/src/v3/wait.ts @@ -29,9 +29,12 @@ export const wait = { return tracer.startActiveSpan( `wait.for()`, async (span) => { + const start = Date.now(); const durationInMs = calculateDurationInMs(options); await runtime.waitForDuration(durationInMs); + + span.end(start + durationInMs); }, { attributes: { @@ -53,13 +56,17 @@ export const wait = { return tracer.startActiveSpan( `wait.until()`, async (span) => { + const start = Date.now(); + if (options.throwIfInThePast && options.date < new Date()) { throw new Error("Date is in the past"); } - const durationInMs = options.date.getTime() - new Date().getTime(); + const durationInMs = options.date.getTime() - start; await runtime.waitForDuration(durationInMs); + + span.end(start + durationInMs); }, { attributes: { diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index d85c3e932..91e4c0c8b 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -55,6 +55,7 @@ importers: dotenv: ^16.4.2 esbuild: ^0.19.11 execa: ^8.0.1 + nanoid: ^5.0.6 prom-client: ^15.1.0 socket.io: ^4.7.4 socket.io-client: ^4.7.4 @@ -64,6 +65,7 @@ importers: '@trigger.dev/core': link:../../packages/core '@trigger.dev/core-apps': link:../../packages/core-apps execa: 8.0.1 + nanoid: 5.0.6 prom-client: 15.1.0 socket.io: 4.7.4 socket.io-client: 4.7.4 @@ -1267,18 +1269,14 @@ importers: packages/core-apps: specifiers: + '@trigger.dev/core': workspace:* '@trigger.dev/tsconfig': workspace:* - '@trigger.dev/tsup': workspace:* '@types/node': '18' - rimraf: ^3.0.2 - tsup: ^8.0.1 typescript: ^5.3.0 devDependencies: + '@trigger.dev/core': link:../core '@trigger.dev/tsconfig': link:../../config-packages/tsconfig - '@trigger.dev/tsup': link:../../config-packages/tsup '@types/node': 18.17.1 - rimraf: 3.0.2 - tsup: 8.0.2_typescript@5.3.3 typescript: 5.3.3 packages/core-backend: @@ -30803,6 +30801,12 @@ packages: hasBin: true dev: false + /nanoid/5.0.6: + resolution: {integrity: sha512-rRq0eMHoGZxlvaFOUdK1Ev83Bd1IgzzR+WJ3IbDJ7QOSdAxYjlurSPqFs9s4lJg29RT6nPwizFtJhQS6V5xgiA==} + engines: {node: ^18 || >=20} + hasBin: true + dev: false + /nanomatch/1.2.13: resolution: {integrity: sha512-fpoe2T0RbHwBTBUOftAfBPaDEi06ufaUai0mE6Yn1kacc3SnTErfb/h+X94VXzI64rKFHYImXSvdwGGCmwOqCA==} engines: {node: '>=0.10.0'} @@ -37618,45 +37622,6 @@ packages: - ts-node dev: true - /tsup/8.0.2_typescript@5.3.3: - resolution: {integrity: sha512-NY8xtQXdH7hDUAZwcQdY/Vzlw9johQsaqf7iwZ6g1DOUlFYQ5/AtVAjTvihhEyeRlGo4dLRVHtrRaL35M1daqQ==} - engines: {node: '>=18'} - hasBin: true - peerDependencies: - '@microsoft/api-extractor': ^7.36.0 - '@swc/core': ^1 - postcss: ^8.4.12 - typescript: '>=4.5.0' - peerDependenciesMeta: - '@microsoft/api-extractor': - optional: true - '@swc/core': - optional: true - postcss: - optional: true - typescript: - optional: true - dependencies: - bundle-require: 4.0.1_esbuild@0.19.11 - cac: 6.7.14 - chokidar: 3.5.3 - debug: 4.3.4 - esbuild: 0.19.11 - execa: 5.1.1 - globby: 11.1.0 - joycon: 3.1.1 - postcss-load-config: 4.0.1 - resolve-from: 5.0.0 - rollup: 4.6.1 - source-map: 0.8.0-beta.0 - sucrase: 3.32.0 - tree-kill: 1.2.2 - typescript: 5.3.3 - transitivePeerDependencies: - - supports-color - - ts-node - dev: true - /tsutils/3.21.0: resolution: {integrity: sha512-mHKK3iUXL+3UF6xL5k0PEhKRUBKPBCv/+RkEOpjRWxxx27KKRBmmA60A9pgOUvMi8GKhRMPEmjBRPzs2W7O1OA==} engines: {node: '>= 6'} diff --git a/references/v3-catalog/src/trigger/retries.ts b/references/v3-catalog/src/trigger/retries.ts index dab7222e1..d067dcea7 100644 --- a/references/v3-catalog/src/trigger/retries.ts +++ b/references/v3-catalog/src/trigger/retries.ts @@ -1,4 +1,4 @@ -import { logger, retry, task } from "@trigger.dev/sdk/v3"; +import { logger, retry, task, wait } from "@trigger.dev/sdk/v3"; import { cache } from "./utils/cache"; import { interceptor } from "./utils/interceptor"; diff --git a/references/v3-catalog/src/trigger/subtasks.ts b/references/v3-catalog/src/trigger/subtasks.ts index 6f0f9692d..a1e6a742d 100644 --- a/references/v3-catalog/src/trigger/subtasks.ts +++ b/references/v3-catalog/src/trigger/subtasks.ts @@ -1,4 +1,5 @@ -import { Context, logger, task } from "@trigger.dev/sdk/v3"; +import { logger, task } from "@trigger.dev/sdk/v3"; +import { taskWithRetries } from "./retries"; export const simpleParentTask = task({ id: "simple-parent-task", @@ -47,3 +48,69 @@ export const simpleChildTask = task({ logger.log("Simple child task payload", { payload, ctx }); }, }); + +export const subtasksWithRetries = task({ + id: "subtasks-with-retries", + run: async (payload: { message: string }) => { + await taskWithRetries.triggerAndWait({ + payload: { + message: `${payload.message} - 2.b`, + }, + }); + + await taskWithRetries.batchTrigger({ + items: [ + { + payload: { + message: `${payload.message} - 2.c`, + }, + }, + { + payload: { + message: `${payload.message} - 2.cc`, + }, + }, + ], + }); + + await taskWithRetries.batchTriggerAndWait({ + items: [ + { + payload: { + message: `${payload.message} - 2.d`, + }, + }, + { + payload: { + message: `${payload.message} - 2.dd`, + }, + }, + ], + }); + + await taskWithRetries.triggerAndWait({ + payload: { + message: `${payload.message} - 2.e`, + }, + }); + + await taskWithRetries.batchTriggerAndWait({ + items: [ + { + payload: { + message: `${payload.message} - 2.f`, + }, + }, + { + payload: { + message: `${payload.message} - 2.ff`, + }, + }, + ], + }); + + return { + hello: "world", + }; + }, +});