diff --git a/apps/kubernetes-provider/src/index.ts b/apps/kubernetes-provider/src/index.ts index 60c27d1d7..9ff7dc71f 100644 --- a/apps/kubernetes-provider/src/index.ts +++ b/apps/kubernetes-provider/src/index.ts @@ -4,6 +4,7 @@ import { TaskOperations, TaskOperationsCreateOptions, TaskOperationsIndexOptions, + TaskOperationsPrePullImageOptions, TaskOperationsRestoreOptions, } from "@trigger.dev/core-apps/provider"; import { SimpleLogger } from "@trigger.dev/core-apps/logger"; @@ -49,6 +50,7 @@ class KubernetesTaskOperations implements TaskOperations { #k8sApi: { core: k8s.CoreV1Api; batch: k8s.BatchV1Api; + apps: k8s.AppsV1Api; }; constructor(namespace = "default") { @@ -313,6 +315,72 @@ class KubernetesTaskOperations implements TaskOperations { await this.#getPod(opts.runId, this.#namespace); } + async prePullImage(opts: TaskOperationsPrePullImageOptions) { + const metaName = this.#getPrePullContainerName(opts.shortCode); + + const metaLabels = { + ...this.#getSharedLabels(opts), + app: "task-prepull", + "app.kubernetes.io/part-of": "trigger-worker", + "app.kubernetes.io/component": "prepull", + deployment: opts.deploymentId, + name: metaName, + } satisfies k8s.V1ObjectMeta["labels"]; + + await this.#createDaemonSet( + { + metadata: { + name: metaName, + namespace: this.#namespace.metadata.name, + labels: metaLabels, + }, + spec: { + selector: { + matchLabels: { + name: metaName, + }, + }, + template: { + metadata: { + labels: metaLabels, + }, + spec: { + ...this.#defaultPodSpec, + restartPolicy: "Always", + initContainers: [ + { + name: "prepull", + image: opts.imageRef, + command: ["/usr/bin/true"], + resources: { + limits: { + cpu: "0.25", + memory: "100Mi", + "ephemeral-storage": "1Gi", + }, + }, + }, + ], + containers: [ + { + name: "pause", + image: "registry.k8s.io/pause:3.9", + resources: { + limits: { + cpu: "1m", + memory: "12Mi", + }, + }, + }, + ], + }, + }, + }, + }, + this.#namespace + ); + } + #envTypeToLabelValue(type: EnvironmentType) { switch (type) { case "PRODUCTION": @@ -402,7 +470,11 @@ class KubernetesTaskOperations implements TaskOperations { } #getSharedLabels( - opts: TaskOperationsIndexOptions | TaskOperationsCreateOptions | TaskOperationsRestoreOptions + opts: + | TaskOperationsIndexOptions + | TaskOperationsCreateOptions + | TaskOperationsRestoreOptions + | TaskOperationsPrePullImageOptions ): Record { return { env: opts.envId, @@ -446,6 +518,10 @@ class KubernetesTaskOperations implements TaskOperations { return `task-run-${suffix}`; } + #getPrePullContainerName(suffix: string) { + return `task-prepull-${suffix}`; + } + #createK8sApi() { const kubeConfig = new k8s.KubeConfig(); @@ -460,6 +536,7 @@ class KubernetesTaskOperations implements TaskOperations { return { core: kubeConfig.makeApiClient(k8s.CoreV1Api), batch: kubeConfig.makeApiClient(k8s.BatchV1Api), + apps: kubeConfig.makeApiClient(k8s.AppsV1Api), }; } @@ -503,6 +580,18 @@ class KubernetesTaskOperations implements TaskOperations { } } + async #createDaemonSet(daemonSet: k8s.V1DaemonSet, namespace: Namespace) { + try { + const res = await this.#k8sApi.apps.createNamespacedDaemonSet( + namespace.metadata.name, + daemonSet + ); + logger.debug(res.body); + } catch (err: unknown) { + this.#handleK8sError(err); + } + } + #throwUnlessRecord(candidate: unknown): asserts candidate is Record { if (typeof candidate !== "object" || candidate === null) { throw candidate; diff --git a/apps/webapp/app/v3/services/createDeployedBackgroundWorker.server.ts b/apps/webapp/app/v3/services/createDeployedBackgroundWorker.server.ts index 6406786de..82ba42f51 100644 --- a/apps/webapp/app/v3/services/createDeployedBackgroundWorker.server.ts +++ b/apps/webapp/app/v3/services/createDeployedBackgroundWorker.server.ts @@ -11,6 +11,7 @@ import { logger } from "~/services/logger.server"; import { ExecuteTasksWaitingForDeployService } from "./executeTasksWaitingForDeploy"; import { PerformDeploymentAlertsService } from "./alerts/performDeploymentAlerts.server"; import { TimeoutDeploymentService } from "./timeoutDeployment.server"; +import { socketIo } from "../handleSocketIo.server"; export class CreateDeployedBackgroundWorkerService extends BaseService { public async call( @@ -132,6 +133,20 @@ export class CreateDeployedBackgroundWorkerService extends BaseService { logger.error("Failed to publish WORKER_CREATED event", { err }); } + if (deployment.imageReference) { + socketIo.providerNamespace.emit("PRE_PULL_IMAGE", { + version: "v1", + imageRef: deployment.imageReference, + shortCode: deployment.shortCode, + // identifiers + deploymentId: deployment.id, + envId: environment.id, + envType: environment.type, + orgId: environment.organizationId, + projectId: deployment.projectId, + }); + } + await ExecuteTasksWaitingForDeployService.enqueue(backgroundWorker.id, this._prisma); await PerformDeploymentAlertsService.enqueue(deployment.id, this._prisma); await TimeoutDeploymentService.dequeue(deployment.id, this._prisma); diff --git a/packages/core-apps/src/provider.ts b/packages/core-apps/src/provider.ts index 8ce381ba2..da53e6996 100644 --- a/packages/core-apps/src/provider.ts +++ b/packages/core-apps/src/provider.ts @@ -64,6 +64,17 @@ export interface TaskOperationsRestoreOptions { checkpointId: string; } +export interface TaskOperationsPrePullImageOptions { + shortCode: string; + imageRef: string; + // identifiers + envId: string; + envType: EnvironmentType; + orgId: string; + projectId: string; + deploymentId: string; +} + export interface TaskOperations { init: () => Promise; @@ -73,8 +84,10 @@ export interface TaskOperations { restore: (opts: TaskOperationsRestoreOptions) => Promise; // unimplemented - delete: (...args: any[]) => Promise; - get: (...args: any[]) => Promise; + delete?: (...args: any[]) => Promise; + get?: (...args: any[]) => Promise; + + prePullImage?: (opts: TaskOperationsPrePullImageOptions) => Promise; } type ProviderShellOptions = { @@ -277,6 +290,27 @@ export class ProviderShell implements Provider { logger.error("restore failed", error); } }, + PRE_PULL_IMAGE: async (message) => { + if (!this.tasks.prePullImage) { + logger.debug("prePullImage not implemented", message); + return; + } + + try { + await this.tasks.prePullImage({ + shortCode: message.shortCode, + imageRef: message.imageRef, + // identifiers + envId: message.envId, + envType: message.envType, + orgId: message.orgId, + projectId: message.projectId, + deploymentId: message.deploymentId, + }); + } catch (error) { + logger.error("prePullImage failed", error); + } + }, }, }); @@ -306,9 +340,12 @@ export class ProviderShell implements Provider { case "/delete": { const body = await getTextBody(req); - await this.tasks.delete({ runId: body }); - - return reply.text(`sent delete request: ${body}`); + if (this.tasks.delete) { + await this.tasks.delete({ runId: body }); + return reply.text(`sent delete request: ${body}`); + } else { + return reply.text("delete not implemented", 501); + } } default: { return reply.empty(404); diff --git a/packages/core/src/v3/schemas/messages.ts b/packages/core/src/v3/schemas/messages.ts index b57eb9f4b..bd86eb4ca 100644 --- a/packages/core/src/v3/schemas/messages.ts +++ b/packages/core/src/v3/schemas/messages.ts @@ -367,6 +367,19 @@ export const PlatformToProviderMessages = { runId: z.string(), }), }, + PRE_PULL_IMAGE: { + message: z.object({ + version: z.literal("v1").default("v1"), + imageRef: z.string(), + shortCode: z.string(), + // identifiers + envId: z.string(), + envType: EnvironmentType, + orgId: z.string(), + projectId: z.string(), + deploymentId: z.string(), + }), + }, }; const CreateWorkerMessage = z.object({