diff --git a/apps/coordinator/src/index.ts b/apps/coordinator/src/index.ts index 9afdfe416..66aa69f2e 100644 --- a/apps/coordinator/src/index.ts +++ b/apps/coordinator/src/index.ts @@ -456,6 +456,45 @@ class TaskCoordinator { 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: socket.data, text: "connected", @@ -523,31 +562,17 @@ class TaskCoordinator { socket.on("TASK_RUN_COMPLETED", async ({ completion, execution }, callback) => { logger.log("completed task", { completionId: completion.id }); - type CheckpointData = { - docker: boolean; - location: string; - }; - - const confirmCompletion = ({ - didCheckpoint, - shouldExit, - checkpoint, - }: { - didCheckpoint: boolean; - shouldExit: boolean; - checkpoint?: CheckpointData; - }) => { + const completeWithoutCheckpoint = (shouldExit: boolean) => { this.#platformSocket?.send("TASK_RUN_COMPLETED", { version: "v1", execution, completion, - checkpoint, }); - callback({ didCheckpoint, shouldExit }); + callback({ willCheckpointAndRestore: false, shouldExit }); }; if (completion.ok) { - confirmCompletion({ didCheckpoint: false, shouldExit: true }); + completeWithoutCheckpoint(true); return; } @@ -555,12 +580,12 @@ 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; } @@ -569,7 +594,20 @@ class TaskCoordinator { const willCheckpointAndRestore = canCheckpoint || willSimulate; if (!willCheckpointAndRestore) { - confirmCompletion({ didCheckpoint: false, shouldExit: false }); + completeWithoutCheckpoint(false); + return; + } + + // The worker will then put itself in a checkpointable state + callback({ willCheckpointAndRestore: true, shouldExit: false }); + + const ready = await readyToCheckpoint(); + + if (!ready.success) { + logger.error("Failed to become checkpointable", { + runId: socket.data.runId, + reason: ready.reason, + }); return; } @@ -581,11 +619,16 @@ class TaskCoordinator { if (!checkpoint) { logger.error("Failed to checkpoint", { runId: socket.data.runId }); - confirmCompletion({ didCheckpoint: false, shouldExit: false }); + completeWithoutCheckpoint(false); return; } - confirmCompletion({ didCheckpoint: true, shouldExit: false, checkpoint }); + this.#platformSocket?.send("TASK_RUN_COMPLETED", { + version: "v1", + execution, + completion, + checkpoint, + }); if (!checkpoint.docker) { socket.emit("REQUEST_EXIT", { @@ -624,7 +667,7 @@ class TaskCoordinator { socket.on("WAIT_FOR_DURATION", async (message, callback) => { logger.log("[WAIT_FOR_DURATION]", message); - if (this.#checkpointableTasks.has(socket.data.runId)) { + if (checkpointInProgress()) { logger.error("Checkpoint already in progress", { runId: socket.data.runId }); callback({ willCheckpointAndRestore: false }); return; @@ -640,21 +683,14 @@ class TaskCoordinator { return; } - const isCheckpointable = new Promise((resolve, reject) => { - // We set a reasonable timeout to prevent waiting forever - setTimeout(reject, 10_000); + const ready = await readyToCheckpoint(); - this.#checkpointableTasks.set(socket.data.runId, { resolve, reject }); - }); - - try { - await isCheckpointable; - } catch (error) { - logger.error("Error while waiting for checkpointable state", { error }); - // TODO: We may want to cancel the task as it's unlikely to recover + if (!ready.success) { + logger.error("Failed to become checkpointable", { + runId: socket.data.runId, + reason: ready.reason, + }); return; - } finally { - this.#checkpointableTasks.delete(socket.data.runId); } const checkpoint = await this.#checkpointer.checkpointAndPush({ diff --git a/apps/kubernetes-provider/src/index.ts b/apps/kubernetes-provider/src/index.ts index 1b0c99cf1..5b4d290b6 100644 --- a/apps/kubernetes-provider/src/index.ts +++ b/apps/kubernetes-provider/src/index.ts @@ -49,6 +49,7 @@ class KubernetesTaskOperations implements TaskOperations { }, spec: { completions: 1, + backoffLimit: 0, ttlSecondsAfterFinished: 300, template: { metadata: { @@ -78,20 +79,6 @@ class KubernetesTaskOperations implements TaskOperations { // memory: "50Mi", // }, // }, - lifecycle: { - preStop: { - httpGet: { - path: "/preStop?cause=index", - port: 8000, - }, - }, - postStart: { - httpGet: { - path: "/postStart?cause=index", - port: 8000, - }, - }, - }, env: [ { name: "DEBUG", @@ -186,16 +173,14 @@ class KubernetesTaskOperations implements TaskOperations { // limits: opts.machine, // }, lifecycle: { - preStop: { - httpGet: { - path: "/preStop?cause=create", - port: 8000, + postStart: { + exec: { + command: this.#getLifecycleCommand("postStart", "create"), }, }, - postStart: { - httpGet: { - path: "/postStart?cause=create", - port: 8000, + preStop: { + exec: { + command: this.#getLifecycleCommand("preStop", "create"), }, }, }, @@ -328,16 +313,14 @@ class KubernetesTaskOperations implements TaskOperations { // limits: opts.machine, // }, lifecycle: { - preStop: { - httpGet: { - path: "/preStop?cause=restore", - port: 8000, + postStart: { + exec: { + command: this.#getLifecycleCommand("postStart", "restore"), }, }, - postStart: { - httpGet: { - path: "/postStart?cause=restore", - port: 8000, + preStop: { + exec: { + command: this.#getLifecycleCommand("preStop", "restore"), }, }, }, @@ -372,6 +355,10 @@ class KubernetesTaskOperations implements TaskOperations { await this.#getPod(opts.runId, this.#namespace); } + #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}`]; + } + #getIndexContainerName(suffix: string) { return `task-index-${suffix}`; } diff --git a/packages/cli-v3/src/workers/prod/entry-point.ts b/packages/cli-v3/src/workers/prod/entry-point.ts index 3b4743399..f5652137a 100644 --- a/packages/cli-v3/src/workers/prod/entry-point.ts +++ b/packages/cli-v3/src/workers/prod/entry-point.ts @@ -3,13 +3,15 @@ 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 { readFile } from "node:fs/promises"; declare const __PROJECT_CONFIG__: Config; @@ -31,13 +33,14 @@ class ProdWorker { private runId = process.env.TRIGGER_RUN_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; #httpPort: number; #backgroundWorker: ProdBackgroundWorker; @@ -93,18 +96,7 @@ class ProdWorker { } ); - logger.log("WAIT_FOR_DURATION", { willCheckpointAndRestore }); - - this.#backgroundWorker.preCheckpointNotification.post({ willCheckpointAndRestore }); - - setTimeout(async () => { - if (willCheckpointAndRestore) { - this.paused = true; - this.nextResumeAfter = "WAIT_FOR_DURATION"; - } - // Forcing a reconnect will ensure the connection handler runs to trigger automatic resume - this.#reconnect(); - }, 3_000); + this.#prepareForCheckpoint("WAIT_FOR_DURATION", willCheckpointAndRestore); }); this.#backgroundWorker.onWaitForTask.attach(async (message) => { @@ -122,18 +114,7 @@ class ProdWorker { } ); - 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.#reconnect(); - }, 3_000); + this.#prepareForCheckpoint("WAIT_FOR_TASK", willCheckpointAndRestore); }); this.#backgroundWorker.onWaitForBatch.attach(async (message) => { @@ -151,18 +132,7 @@ class ProdWorker { } ); - 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.#reconnect(); - }, 3_000); + this.#prepareForCheckpoint("WAIT_FOR_BATCH", willCheckpointAndRestore); }); this.#httpPort = port; @@ -172,6 +142,11 @@ class ProdWorker { async #reconnect() { this.#coordinatorSocket.close(); + if (!this.runningInKubernetes) { + this.#coordinatorSocket.connect(); + return; + } + try { const coordinatorHost = (await readFile("/etc/taskinfo/coordinator-host", "utf-8")).replace( "\n", @@ -180,8 +155,8 @@ class ProdWorker { logger.log("reconnecting", { coordinatorHost: { - env: COORDINATOR_HOST, - volume: coordinatorHost, + fromEnv: COORDINATOR_HOST, + fromVolume: coordinatorHost, current: this.#coordinatorSocket.socket.io.opts.hostname, }, }); @@ -193,6 +168,28 @@ class ProdWorker { } } + #prepareForCheckpoint(reason: WaitReason, willCheckpointAndRestore: boolean) { + logger.log(reason, { willCheckpointAndRestore }); + + this.#backgroundWorker.preCheckpointNotification.post({ willCheckpointAndRestore }); + + if (willCheckpointAndRestore) { + this.paused = true; + this.nextResumeAfter = reason; + } + + if (this.runningInKubernetes) { + return; + } + + // We don't have access to the postStart lifecycle hook, so we set a reconnect timer + // TODO: Implement lifecycle hook for docker + setTimeout(async () => { + // Reconnecting when paused will trigger automatic resume + await this.#reconnect(); + }, 3_000); + } + #returnValidatedExtraHeaders(headers: Record) { for (const [key, value] of Object.entries(headers)) { if (value === undefined) { @@ -264,34 +261,45 @@ class ProdWorker { logger.log("completed", completion); this.completed.add(executionPayload.execution.attempt.id); - this.executing = false; - this.attemptFriendlyId = undefined; 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 }); + // 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.#coordinatorSocket.socket.emit("READY_FOR_CHECKPOINT", { version: "v1" }); + // Give coordinator a chance to request exit + // TODO: Consider disabling automatic reconnect instead, and re-enabling it on postStart hook await new Promise((resolve) => { + // The timeout duration is below the minimum wait duration that triggers a checkpoint + // We are unlikely to restore before this, so this will have resolved on restore setTimeout(resolve, 5_000); }); - // Forcing a reconnect will ensure the connection handler runs and signals we are ready for another execution - this.#reconnect(); + // Remove executing state as late as possible to prevent further execution until we've + this.executing = false; + this.attemptFriendlyId = undefined; + + if (!this.runningInKubernetes) { + this.#reconnect(); + } }, REQUEST_ATTEMPT_CANCELLATION: async (message) => { if (!this.executing) { @@ -305,6 +313,8 @@ class ProdWorker { }, }, onConnection: async (socket, handler, sender, logger) => { + if (process.env.DEBUG === "true") return; + if (process.env.INDEX_TASKS === "true") { try { const taskResources = await this.#initializeWorker(); @@ -417,7 +427,7 @@ class ProdWorker { }, }); - this.#reconnect(); + await this.#reconnect(); }, onDisconnect: async (socket, reason, description, logger) => { // this.#reconnect(); @@ -430,71 +440,118 @@ 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 "/connect": - this.#coordinatorSocket.connect(); - return reply.empty(); + case "/close": { + await this.#coordinatorSocket.sendWithAck("LOG", { + version: "v1", + text: `[${req.method}] ${req.url}`, + }); - 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", - runId: this.runId, - totalCompletions: this.completed.size, - }); - 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(); + break; + } + default: { + logger.error("Unhandled cause", { cause: cause.data }); + break; + } + } + logger.log("postStart", { url: req.url }); + + return reply.text("postStart ok"); + } + + default: { + return reply.empty(404); + } + } + } catch (error) { + logger.error("HTTP server error", { error }); + reply.empty(500); } }); 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/src/v3/runtime/prodRuntimeManager.ts b/packages/core/src/v3/runtime/prodRuntimeManager.ts index 183cc4252..c7359802d 100644 --- a/packages/core/src/v3/runtime/prodRuntimeManager.ts +++ b/packages/core/src/v3/runtime/prodRuntimeManager.ts @@ -60,7 +60,7 @@ export class ProdRuntimeManager implements RuntimeManager { return; } - const waitForRestore = new Promise((resolve, reject) => { + const waitForRestore = new Promise((resolve, reject) => { this._waitForRestore = { resolve, reject }; }); diff --git a/packages/core/src/v3/schemas/schemas.ts b/packages/core/src/v3/schemas/schemas.ts index 0c7e6c1a8..eb1a4a0bc 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({ @@ -172,7 +176,7 @@ export const CoordinatorToPlatformMessages = { message: z.object({ version: z.literal("v1").default("v1"), attemptFriendlyId: z.string(), - type: z.enum(["WAIT_FOR_DURATION", "WAIT_FOR_TASK", "WAIT_FOR_BATCH"]), + type: WaitReason, }), }, TASK_RUN_COMPLETED: { @@ -334,7 +338,7 @@ export const ProdWorkerToCoordinatorMessages = { message: z.object({ version: z.literal("v1").default("v1"), attemptFriendlyId: z.string(), - type: z.enum(["WAIT_FOR_DURATION", "WAIT_FOR_TASK", "WAIT_FOR_BATCH"]), + type: WaitReason, }), }, READY_FOR_CHECKPOINT: { @@ -360,7 +364,7 @@ export const ProdWorkerToCoordinatorMessages = { completion: TaskRunExecutionResult, }), callback: z.object({ - didCheckpoint: z.boolean(), + willCheckpointAndRestore: z.boolean(), shouldExit: z.boolean(), }), },