diff --git a/internal-packages/run-engine/src/engine/index.ts b/internal-packages/run-engine/src/engine/index.ts index e8dc83237..45d998c3a 100644 --- a/internal-packages/run-engine/src/engine/index.ts +++ b/internal-packages/run-engine/src/engine/index.ts @@ -1540,6 +1540,7 @@ export class RunEngine { createdAt: true, completedAt: true, taskEventStore: true, + parentTaskRunId: true, runtimeEnvironment: { select: { organizationId: true, @@ -1559,7 +1560,9 @@ export class RunEngine { }); //remove it from the queue and release concurrency - await this.runQueue.acknowledgeMessage(run.runtimeEnvironment.organizationId, runId); + await this.runQueue.acknowledgeMessage(run.runtimeEnvironment.organizationId, runId, { + messageId: run.parentTaskRunId ?? undefined, + }); //if executing, we need to message the worker to cancel the run and put it into `PENDING_CANCEL` status if (isExecuting(latestSnapshot.executionStatus)) { @@ -2790,10 +2793,13 @@ export class RunEngine { createdAt: true, completedAt: true, taskEventStore: true, + parentTaskRunId: true, }, }); - await this.runQueue.acknowledgeMessage(updatedRun.runtimeEnvironment.organizationId, runId); + await this.runQueue.acknowledgeMessage(updatedRun.runtimeEnvironment.organizationId, runId, { + messageId: updatedRun.parentTaskRunId ?? undefined, + }); if (!updatedRun.associatedWaitpoint) { throw new ServiceValidationError("No associated waitpoint found", 400); @@ -2931,10 +2937,14 @@ export class RunEngine { createdAt: true, completedAt: true, taskEventStore: true, + parentTaskRunId: true, }, }); const newSnapshot = await getLatestExecutionSnapshot(prisma, runId); - await this.runQueue.acknowledgeMessage(run.project.organizationId, runId); + + await this.runQueue.acknowledgeMessage(run.project.organizationId, runId, { + messageId: run.parentTaskRunId ?? undefined, + }); // We need to manually emit this as we created the final snapshot as part of the task run update this.eventBus.emit("executionSnapshotCreated", { @@ -3248,6 +3258,7 @@ export class RunEngine { attemptNumber: true, spanId: true, batchId: true, + parentTaskRunId: true, associatedWaitpoint: { select: { id: true, @@ -3282,13 +3293,15 @@ export class RunEngine { throw new ServiceValidationError("No associated waitpoint found", 400); } + await this.runQueue.acknowledgeMessage(run.runtimeEnvironment.organizationId, runId, { + messageId: run.parentTaskRunId ?? undefined, + }); + await this.completeWaitpoint({ id: run.associatedWaitpoint.id, output: { value: JSON.stringify(error), isError: true }, }); - await this.runQueue.acknowledgeMessage(run.runtimeEnvironment.organizationId, runId); - this.eventBus.emit("runFailed", { time: failedAt, run: { diff --git a/internal-packages/run-engine/src/run-queue/index.ts b/internal-packages/run-engine/src/run-queue/index.ts index e30a00d2f..a4b9e98d7 100644 --- a/internal-packages/run-engine/src/run-queue/index.ts +++ b/internal-packages/run-engine/src/run-queue/index.ts @@ -481,18 +481,20 @@ export class RunQueue { * This is done when the run is in a final state. * @param messageId */ - public async acknowledgeMessage(orgId: string, messageId: string) { + public async acknowledgeMessage( + orgId: string, + messageId: string, + reserveConcurrency?: { + messageId?: string; + } + ) { return this.#trace( "acknowledgeMessage", async (span) => { const message = await this.readMessage(orgId, messageId); if (!message) { - this.logger.log(`[${this.name}].acknowledgeMessage() message not found`, { - messageId, - service: this.name, - }); - return; + throw new MessageNotFoundError(messageId); } span.setAttributes({ @@ -504,6 +506,7 @@ export class RunQueue { await this.#callAcknowledgeMessage({ message, + reserveConcurrency, }); }, { @@ -1006,7 +1009,15 @@ export class RunQueue { }; } - async #callAcknowledgeMessage({ message }: { message: OutputPayload }) { + async #callAcknowledgeMessage({ + message, + reserveConcurrency, + }: { + message: OutputPayload; + reserveConcurrency?: { + messageId?: string; + }; + }) { const messageId = message.runId; const messageKey = this.keys.messageKey(message.orgId, messageId); const messageQueue = message.queue; @@ -1026,22 +1037,37 @@ export class RunQueue { service: this.name, }); - const queueReserveConcurrencyKey = this.keys.reserveConcurrencyKeyFromQueue(messageQueue); - const envReserveConcurrencyKey = this.keys.envReserveConcurrencyKeyFromQueue(messageQueue); + if (!reserveConcurrency?.messageId) { + return this.redis.acknowledgeMessage( + messageKey, + messageQueue, + queueCurrentConcurrencyKey, + envCurrentConcurrencyKey, + envQueueKey, + messageId, + messageQueue, + JSON.stringify(masterQueues), + this.options.redis.keyPrefix ?? "" + ); + } else { + const queueReserveConcurrencyKey = this.keys.reserveConcurrencyKeyFromQueue(messageQueue); + const envReserveConcurrencyKey = this.keys.envReserveConcurrencyKeyFromQueue(messageQueue); - return this.redis.acknowledgeMessage( - messageKey, - messageQueue, - queueCurrentConcurrencyKey, - envCurrentConcurrencyKey, - envQueueKey, - queueReserveConcurrencyKey, - envReserveConcurrencyKey, - messageId, - messageQueue, - JSON.stringify(masterQueues), - this.options.redis.keyPrefix ?? "" - ); + return this.redis.acknowledgeMessageWithReserveConcurrency( + messageKey, + messageQueue, + queueCurrentConcurrencyKey, + envCurrentConcurrencyKey, + envQueueKey, + envReserveConcurrencyKey, + queueReserveConcurrencyKey, + messageId, + messageQueue, + JSON.stringify(masterQueues), + this.options.redis.keyPrefix ?? "", + reserveConcurrency.messageId + ); + } } async #callNackMessage({ message, retryAt }: { message: OutputPayload; retryAt?: number }) { @@ -1393,16 +1419,14 @@ return {messageId, messageScore, messagePayload} -- Return message details }); this.redis.defineCommand("acknowledgeMessage", { - numberOfKeys: 7, + numberOfKeys: 5, lua: ` -- Keys: local messageKey = KEYS[1] -local messageQueue = KEYS[2] -local concurrencyKey = KEYS[3] +local messageQueueKey = KEYS[2] +local queueCurrentConcurrencyKey = KEYS[3] local envCurrentConcurrencyKey = KEYS[4] local envQueueKey = KEYS[5] -local queueReserveConcurrencyKey = KEYS[6] -local envReserveConcurrencyKey = KEYS[7] -- Args: local messageId = ARGV[1] @@ -1414,11 +1438,11 @@ local keyPrefix = ARGV[4] redis.call('DEL', messageKey) -- Remove the message from the queue -redis.call('ZREM', messageQueue, messageId) +redis.call('ZREM', messageQueueKey, messageId) redis.call('ZREM', envQueueKey, messageId) -- Rebalance the parent queues -local earliestMessage = redis.call('ZRANGE', messageQueue, 0, 0, 'WITHSCORES') +local earliestMessage = redis.call('ZRANGE', messageQueueKey, 0, 0, 'WITHSCORES') for _, parentQueue in ipairs(parentQueues) do local prefixedParentQueue = keyPrefix .. parentQueue if #earliestMessage == 0 then @@ -1429,12 +1453,55 @@ for _, parentQueue in ipairs(parentQueues) do end -- Update the concurrency keys -redis.call('SREM', concurrencyKey, messageId) +redis.call('SREM', queueCurrentConcurrencyKey, messageId) +redis.call('SREM', envCurrentConcurrencyKey, messageId) +`, + }); + + this.redis.defineCommand("acknowledgeMessageWithReserveConcurrency", { + numberOfKeys: 7, + lua: ` +-- Keys: +local messageKey = KEYS[1] +local messageQueueKey = KEYS[2] +local queueCurrentConcurrencyKey = KEYS[3] +local envCurrentConcurrencyKey = KEYS[4] +local envQueueKey = KEYS[5] +local queueReserveConcurrencyKey = KEYS[6] +local envReserveConcurrencyKey = KEYS[7] + +-- Args: +local messageId = ARGV[1] +local messageQueueName = ARGV[2] +local parentQueues = cjson.decode(ARGV[3]) +local keyPrefix = ARGV[4] +local reserveMessageId = ARGV[5] + +-- Remove the message from the message key +redis.call('DEL', messageKey) + +-- Remove the message from the queue +redis.call('ZREM', messageQueueKey, messageId) +redis.call('ZREM', envQueueKey, messageId) + +-- Rebalance the parent queues +local earliestMessage = redis.call('ZRANGE', messageQueueKey, 0, 0, 'WITHSCORES') +for _, parentQueue in ipairs(parentQueues) do + local prefixedParentQueue = keyPrefix .. parentQueue + if #earliestMessage == 0 then + redis.call('ZREM', prefixedParentQueue, messageQueueName) + else + redis.call('ZADD', prefixedParentQueue, earliestMessage[2], messageQueueName) + end +end + +-- Update the concurrency keys +redis.call('SREM', queueCurrentConcurrencyKey, messageId) redis.call('SREM', envCurrentConcurrencyKey, messageId) -- Clear reserve concurrency -redis.call('SREM', queueReserveConcurrencyKey, messageId) -redis.call('SREM', envReserveConcurrencyKey, messageId) +redis.call('SREM', queueReserveConcurrencyKey, reserveMessageId) +redis.call('SREM', envReserveConcurrencyKey, reserveMessageId) `, }); @@ -1585,12 +1652,6 @@ end redis.call('SADD', queueCurrentConcurrencyKey, messageId) redis.call('SADD', envCurrentConcurrencyKey, messageId) --- Remove the message from the queue reserve concurrency set -redis.call('SREM', queueReserveConcurrencyKey, messageId) - --- Remove the message from the env reserve concurrency set -redis.call('SREM', envReserveConcurrencyKey, messageId) - return true `, }); @@ -1701,8 +1762,6 @@ declare module "@internal/redis" { concurrencyKey: string, envConcurrencyKey: string, envQueueKey: string, - queueReserveConcurrencyKey: string, - envReserveConcurrencyKey: string, messageId: string, messageQueueName: string, masterQueues: string, @@ -1710,6 +1769,22 @@ declare module "@internal/redis" { callback?: Callback ): Result; + acknowledgeMessageWithReserveConcurrency( + messageKey: string, + messageQueue: string, + concurrencyKey: string, + envConcurrencyKey: string, + envQueueKey: string, + envReserveConcurrencyKey: string, + queueReserveConcurrencyKey: string, + messageId: string, + messageQueueName: string, + masterQueues: string, + keyPrefix: string, + reserveMessageId: string, + callback?: Callback + ): Result; + nackMessage( messageKey: string, messageQueue: string, diff --git a/internal-packages/run-engine/src/run-queue/tests/ack.test.ts b/internal-packages/run-engine/src/run-queue/tests/ack.test.ts index 3887dad53..6805984d9 100644 --- a/internal-packages/run-engine/src/run-queue/tests/ack.test.ts +++ b/internal-packages/run-engine/src/run-queue/tests/ack.test.ts @@ -274,7 +274,7 @@ describe("RunQueue.acknowledgeMessage", () => { }); redisTest( - "acknowledging a message clears reserve concurrency sets even when not dequeued", + "acknowledging a message clears env reserve concurrency when recursive queue is false", async ({ redisContainer }) => { const queue = new RunQueue({ ...testOptions, @@ -302,8 +302,8 @@ describe("RunQueue.acknowledgeMessage", () => { message: messageDev, masterQueues: ["main", envMasterQueue], reserveConcurrency: { - messageId: messageDev.runId, - recursiveQueue: true, + messageId: "r1235", + recursiveQueue: false, }, }); @@ -317,7 +317,7 @@ describe("RunQueue.acknowledgeMessage", () => { authenticatedEnvDev, messageDev.queue ); - expect(queueReserveConcurrency).toBe(1); + expect(queueReserveConcurrency).toBe(0); // Verify message is in queue before acknowledging const queueLengthBefore = await queue.lengthOfQueue(authenticatedEnvDev, messageDev.queue); @@ -327,7 +327,9 @@ describe("RunQueue.acknowledgeMessage", () => { expect(envQueueLengthBefore).toBe(1); // Acknowledge the message before dequeuing - await queue.acknowledgeMessage(messageDev.orgId, messageDev.runId); + await queue.acknowledgeMessage(messageDev.orgId, messageDev.runId, { + messageId: "r1235", + }); // Verify reserve concurrency is cleared const envReserveConcurrencyAfter = await queue.reserveConcurrencyOfEnvironment(