diff --git a/apps/coordinator/src/index.ts b/apps/coordinator/src/index.ts index b43e45577..feb2722ba 100644 --- a/apps/coordinator/src/index.ts +++ b/apps/coordinator/src/index.ts @@ -166,7 +166,12 @@ class Checkpointer { this.#abortControllers.delete(runId); } - async #checkpointAndPush(opts: CheckpointAndPushOptions): Promise { + async #checkpointAndPush({ + runId, + leaveRunning = true, // This mirrors kubernetes behaviour more accurately + projectRef, + deploymentVersion, + }: CheckpointAndPushOptions): Promise { await this.initialize(); if (!this.#dockerMode && !this.#canCheckpoint) { @@ -174,28 +179,38 @@ class Checkpointer { return; } - if (this.#abortControllers.has(opts.runId)) { - logger.error("Checkpoint procedure already in progress", { opts }); + 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(opts.runId, controller); + this.#abortControllers.set(runId, controller); const $$ = $({ signal: controller.signal }); try { const shortCode = nanoid(8); - const imageRef = this.#getImageRef(opts.projectRef, opts.deploymentVersion, shortCode); - const exportLocation = this.#getExportLocation( - opts.projectRef, - opts.deploymentVersion, - shortCode - ); + const imageRef = this.#getImageRef(projectRef, deploymentVersion, shortCode); + const exportLocation = this.#getExportLocation(projectRef, deploymentVersion, shortCode); - this.#logger.log("Checkpointing:", { opts }); + this.#logger.log("Checkpointing:", { + options: { + runId, + leaveRunning, + projectRef, + deploymentVersion, + }, + }); - const containterName = this.#getRunContainerName(opts.runId); + const containterName = this.#getRunContainerName(runId); // Create checkpoint (docker) if (this.#dockerMode) { @@ -204,7 +219,7 @@ class Checkpointer { this.#logger.log("Simulating checkpoint"); this.#logger.debug(await $$`docker pause ${containterName}`); } else { - if (opts.leaveRunning) { + if (leaveRunning) { this.#logger.debug( await $$`docker checkpoint create --leave-running ${containterName} ${exportLocation}` ); @@ -220,7 +235,7 @@ class Checkpointer { } this.#logger.log("checkpoint created:", { - runId: opts.runId, + runId, location: exportLocation, }); @@ -279,10 +294,18 @@ class Checkpointer { docker: false, }; } catch (error) { - this.#logger.error("checkpoint failed", { options: opts, error }); + this.#logger.error("checkpoint failed", { + options: { + runId, + leaveRunning, + projectRef, + deploymentVersion, + }, + error, + }); return; } finally { - this.#abortControllers.delete(opts.runId); + this.#abortControllers.delete(runId); } } @@ -346,7 +369,7 @@ class TaskCoordinator { serverMessages: PlatformToCoordinatorMessages, authToken: PLATFORM_SECRET, handlers: { - RESUME: async (message) => { + RESUME_AFTER_DEPENDENCY: async (message) => { const taskSocket = await this.#getAttemptSocket(message.attemptFriendlyId); if (!taskSocket) { @@ -356,7 +379,10 @@ class TaskCoordinator { 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.attemptFriendlyId); @@ -647,7 +673,7 @@ class TaskCoordinator { checkpoint, }); - if (!checkpoint.docker) { + if (!checkpoint.docker || !willSimulate) { socket.emit("REQUEST_EXIT", { version: "v1", }); @@ -670,15 +696,7 @@ class TaskCoordinator { socket.on("CANCEL_CHECKPOINT", async (message) => { logger.log("[CANCEL_CHECKPOINT]", message); - const checkpointWait = this.#checkpointableTasks.get(socket.data.runId); - - if (checkpointWait) { - // Stop waiting for task to reach checkpointable state - checkpointWait.reject("Checkpoint cancelled"); - } - - // Cancel checkpointing procedure - this.#checkpointer.cancelCheckpoint(socket.data.runId); + this.#cancelCheckpoint(socket.data.runId); }); socket.on("WAIT_FOR_DURATION", async (message, callback) => { @@ -722,7 +740,7 @@ class TaskCoordinator { return; } - if (!checkpoint.docker) { + if (!checkpoint.docker || !willSimulate) { socket.emit("REQUEST_EXIT", { version: "v1", }); @@ -765,7 +783,7 @@ class TaskCoordinator { return; } - if (!checkpoint.docker) { + if (!checkpoint.docker || !willSimulate) { socket.emit("REQUEST_EXIT", { version: "v1", }); @@ -807,7 +825,7 @@ class TaskCoordinator { return; } - if (!checkpoint.docker) { + if (!checkpoint.docker || !willSimulate) { socket.emit("REQUEST_EXIT", { version: "v1", }); @@ -874,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); diff --git a/apps/docker-provider/src/index.ts b/apps/docker-provider/src/index.ts index 9ed015821..7fed78137 100644 --- a/apps/docker-provider/src/index.ts +++ b/apps/docker-provider/src/index.ts @@ -172,24 +172,7 @@ class DockerTaskOperations implements TaskOperations { throw new Error("docker unpause command failed"); } - // Emulate prod-like postStart command - // For this to work we need to first get the correct port, which is random during dev as we run with host networking and need to avoid clashes - 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"); - } - + await this.#sendPostStart(containerName); return; } @@ -200,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 }) { @@ -222,6 +207,26 @@ class DockerTaskOperations implements TaskOperations { 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, diff --git a/apps/webapp/app/v3/marqs/sharedQueueConsumer.server.ts b/apps/webapp/app/v3/marqs/sharedQueueConsumer.server.ts index 2929343fe..3b3fbd6b1 100644 --- a/apps/webapp/app/v3/marqs/sharedQueueConsumer.server.ts +++ b/apps/webapp/app/v3/marqs/sharedQueueConsumer.server.ts @@ -669,8 +669,9 @@ 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, diff --git a/apps/webapp/app/v3/services/resumeAttempt.server.ts b/apps/webapp/app/v3/services/resumeAttempt.server.ts index 001f0e520..c7ec9771c 100644 --- a/apps/webapp/app/v3/services/resumeAttempt.server.ts +++ b/apps/webapp/app/v3/services/resumeAttempt.server.ts @@ -80,7 +80,17 @@ export class ResumeAttemptService extends BaseService { 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": @@ -211,8 +221,9 @@ export class ResumeAttemptService extends BaseService { }, }); - socketIo.coordinatorNamespace.emit("RESUME", { + socketIo.coordinatorNamespace.emit("RESUME_AFTER_DEPENDENCY", { version: "v1", + runId: attempt.taskRunId, attemptId: attempt.id, attemptFriendlyId: attempt.friendlyId, completions, diff --git a/packages/cli-v3/src/workers/prod/entry-point.ts b/packages/cli-v3/src/workers/prod/entry-point.ts index e13dfc237..3e780ac8b 100644 --- a/packages/cli-v3/src/workers/prod/entry-point.ts +++ b/packages/cli-v3/src/workers/prod/entry-point.ts @@ -203,6 +203,13 @@ class ProdWorker { } } + #resumeAfterDuration() { + this.paused = false; + this.nextResumeAfter = undefined; + + this.#backgroundWorker.waitCompletedNotification(); + } + #returnValidatedExtraHeaders(headers: Record) { for (const [key, value] of Object.entries(headers)) { if (value === undefined) { @@ -243,7 +250,52 @@ class ProdWorker { 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]; @@ -254,7 +306,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) { @@ -296,6 +362,7 @@ class ProdWorker { await this.#backgroundWorker.cancelAttempt(message.attemptId); }, REQUEST_EXIT: async () => { + this.#coordinatorSocket.close(); process.exit(0); }, READY_FOR_RETRY: async (message) => { @@ -311,8 +378,6 @@ 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(); @@ -394,17 +459,17 @@ class ProdWorker { return; } + if (this.nextResumeAfter === "WAIT_FOR_DURATION") { + this.#resumeAfterDuration(); + return; + } + socket.emit("READY_FOR_RESUME", { version: "v1", attemptFriendlyId: this.attemptFriendlyId, type: this.nextResumeAfter, }); - this.#backgroundWorker.waitCompletedNotification(); - - this.paused = false; - this.nextResumeAfter = undefined; - return; } diff --git a/packages/core/src/v3/schemas/schemas.ts b/packages/core/src/v3/schemas/schemas.ts index 06db2623c..710a49329 100644 --- a/packages/core/src/v3/schemas/schemas.ts +++ b/packages/core/src/v3/schemas/schemas.ts @@ -37,9 +37,9 @@ 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 const WaitReason = z.enum(["WAIT_FOR_DURATION", "WAIT_FOR_TASK", "WAIT_FOR_BATCH"]); -export type WaitReason = z.infer +export type WaitReason = z.infer; export const ProviderToPlatformMessages = { LOG: { @@ -240,9 +240,10 @@ 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(), @@ -422,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(),