Convert to using ZSETs and adding env vars
This commit is contained in:
@@ -59,6 +59,16 @@ const EnvironmentSchema = z.object({
|
||||
AWS_SQS_QUEUE_URL: z.string().optional(),
|
||||
AWS_SQS_BATCH_SIZE: z.coerce.number().int().optional().default(10),
|
||||
DISABLE_SSE: z.string().optional(),
|
||||
|
||||
// Redis options
|
||||
REDIS_HOST: z.string().optional(),
|
||||
REDIS_READER_HOST: z.string().optional(),
|
||||
REDIS_READER_PORT: z.coerce.number().optional(),
|
||||
REDIS_PORT: z.coerce.number().optional(),
|
||||
REDIS_USERNAME: z.string().optional(),
|
||||
REDIS_PASSWORD: z.string().optional(),
|
||||
|
||||
DEFAULT_ORG_EXECUTION_CONCURRENCY_LIMIT: z.coerce.number().int().default(10),
|
||||
});
|
||||
|
||||
export type Environment = z.infer<typeof EnvironmentSchema>;
|
||||
|
||||
@@ -15,7 +15,15 @@ import omit from "lodash.omit";
|
||||
import { z } from "zod";
|
||||
import { PrismaClient, PrismaClientOrTransaction } from "~/db.server";
|
||||
import { workerLogger as logger, trace, workerLogger } from "~/services/logger.server";
|
||||
import { Callback, Redis, RedisOptions, Result } from "ioredis";
|
||||
import {
|
||||
Callback,
|
||||
Cluster,
|
||||
ClusterNode,
|
||||
ClusterOptions,
|
||||
Redis,
|
||||
RedisOptions,
|
||||
Result,
|
||||
} from "ioredis";
|
||||
|
||||
export interface MessageCatalogSchema {
|
||||
[key: string]: z.ZodFirstPartySchemaTypes | z.ZodDiscriminatedUnion<any, any>;
|
||||
@@ -669,7 +677,10 @@ declare module "ioredis" {
|
||||
maxSizeKey: string,
|
||||
forbiddenFlagsKey: string,
|
||||
jobId: string,
|
||||
timestamp: string,
|
||||
windowSize: string,
|
||||
forbiddenFlag: string,
|
||||
maxSize: string,
|
||||
callback?: Callback<string>
|
||||
): Result<string, Context>;
|
||||
|
||||
@@ -678,25 +689,40 @@ declare module "ioredis" {
|
||||
maxSizeKey: string,
|
||||
forbiddenFlagsKey: string,
|
||||
jobId: string,
|
||||
timestamp: string,
|
||||
windowSize: string,
|
||||
forbiddenFlag: string,
|
||||
maxSize: string,
|
||||
callback?: Callback<string>
|
||||
): Result<string, Context>;
|
||||
}
|
||||
}
|
||||
|
||||
export type RedisGraphileRateLimiterOptions = {
|
||||
redis: RedisOptions;
|
||||
redis?: RedisOptions;
|
||||
cluster?: {
|
||||
startupNodes: ClusterNode[];
|
||||
options?: ClusterOptions;
|
||||
};
|
||||
defaultConcurrency?: number;
|
||||
windowSize?: number;
|
||||
prefix?: string;
|
||||
};
|
||||
|
||||
const FORBIDDEN_FLAG_KEY = "rl:forbiddenFlags";
|
||||
|
||||
// TODO: we need to somehow seed and update the rate limit for each flag in Redis
|
||||
export class RedisGraphileRateLimiter implements GraphileRateLimiter {
|
||||
private redis: Redis;
|
||||
private prefix: string;
|
||||
private redis: Redis | Cluster;
|
||||
private defaultMaxSize: number;
|
||||
private windowSize: number;
|
||||
|
||||
constructor(options?: RedisGraphileRateLimiterOptions) {
|
||||
this.redis = new Redis(options?.redis ?? {});
|
||||
this.prefix = options?.prefix ?? "tr:gw";
|
||||
this.redis = options?.cluster
|
||||
? new Redis.Cluster(options.cluster.startupNodes, options.cluster.options)
|
||||
: new Redis(options?.redis ?? {});
|
||||
this.defaultMaxSize = options?.defaultConcurrency ?? 10;
|
||||
this.windowSize = options?.windowSize ?? 1000 * 2 * 60; // 2 minutes
|
||||
|
||||
this.redis.defineCommand("beforeTask", {
|
||||
numberOfKeys: 3,
|
||||
@@ -705,15 +731,15 @@ local setKey = KEYS[1]
|
||||
local maxSizeKey = KEYS[2]
|
||||
local forbiddenFlagsKey = KEYS[3]
|
||||
local jobId = ARGV[1]
|
||||
local forbiddenFlag = ARGV[2]
|
||||
local timestamp = ARGV[2]
|
||||
local windowSize = ARGV[3]
|
||||
local forbiddenFlag = ARGV[4]
|
||||
local defaultMaxSize = ARGV[5]
|
||||
|
||||
local maxSize = tonumber(redis.call('GET', maxSizeKey))
|
||||
if maxSize == nil then
|
||||
return false -- maxSize not set
|
||||
end
|
||||
local maxSize = tonumber(redis.call('GET', maxSizeKey) or defaultMaxSize)
|
||||
|
||||
redis.call('SADD', setKey, jobId)
|
||||
local currentSize = redis.call('SCARD', setKey)
|
||||
redis.call('ZADD', setKey, timestamp, jobId)
|
||||
local currentSize = redis.call('ZCOUNT', setKey, timestamp - windowSize, timestamp)
|
||||
|
||||
if currentSize < maxSize then
|
||||
redis.call('SREM', forbiddenFlagsKey, forbiddenFlag)
|
||||
@@ -732,50 +758,51 @@ local setKey = KEYS[1]
|
||||
local maxSizeKey = KEYS[2]
|
||||
local forbiddenFlagsKey = KEYS[3]
|
||||
local jobId = ARGV[1]
|
||||
local forbiddenFlag = ARGV[2]
|
||||
local timestamp = ARGV[2]
|
||||
local windowSize = ARGV[3]
|
||||
local forbiddenFlag = ARGV[4]
|
||||
local defaultMaxSize = ARGV[5]
|
||||
|
||||
local maxSize = tonumber(redis.call('GET', maxSizeKey))
|
||||
if maxSize == nil then
|
||||
return false -- maxSize not set
|
||||
end
|
||||
local maxSize = tonumber(redis.call('GET', maxSizeKey) or defaultMaxSize)
|
||||
|
||||
redis.call('SREM', setKey, jobId)
|
||||
local currentSize = redis.call('SCARD', setKey)
|
||||
-- Remove the job ID from the ZSET
|
||||
redis.call('ZREM', setKey, jobId)
|
||||
|
||||
-- Count the current number of jobs in the window
|
||||
local currentSize = redis.call('ZCOUNT', setKey, timestamp - windowSize, timestamp)
|
||||
|
||||
-- The cleanup of old job IDs is now an essential part of maintaining the ZSET's size
|
||||
redis.call('ZREMRANGEBYSCORE', setKey, '-inf', timestamp - windowSize)
|
||||
|
||||
-- Update the forbidden flags based on the current size
|
||||
if currentSize < maxSize then
|
||||
-- Only remove the forbidden flag if it's no longer needed
|
||||
redis.call('SREM', forbiddenFlagsKey, forbiddenFlag)
|
||||
return true
|
||||
else
|
||||
redis.call('SADD', forbiddenFlagsKey, forbiddenFlag)
|
||||
-- No need to add the forbidden flag here as it should be handled in beforeTask
|
||||
return false
|
||||
end
|
||||
|
||||
`,
|
||||
});
|
||||
|
||||
if (this.redis instanceof Redis) {
|
||||
logger.debug("⚡ RedisGraphileRateLimiter connected to Redis", {
|
||||
host: this.redis.options.host,
|
||||
port: this.redis.options.port,
|
||||
});
|
||||
} else {
|
||||
logger.debug("⚡ RedisGraphileRateLimiter connected to Redis Cluster", {
|
||||
nodes: this.redis.nodes,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
async forbiddenFlags(): Promise<string[]> {
|
||||
return this.redis.smembers(this.#prefixKey("rl:forbiddenFlags"));
|
||||
return this.redis.smembers(FORBIDDEN_FLAG_KEY);
|
||||
}
|
||||
|
||||
// wrapTask
|
||||
// Before the task is run we need to:
|
||||
// get the max concurreny for the flag
|
||||
// if there is no max concurreny for the flag, we can skip the rest of the steps
|
||||
// for each flag with the prefix "rl:"
|
||||
// add the job id to a redis set with the key "rl:flag"
|
||||
// get the length of the set
|
||||
// if the length of the set is greater or equal to the max concurrency
|
||||
// we need to add the flag to the "forbidden flags" list
|
||||
// After the task is run
|
||||
// for each flag with the prefix "rl:"
|
||||
// get the max concurreny for the flag
|
||||
// if there is no max concurreny for the flag, we can skip the rest of the steps
|
||||
// remove the job id from the redis set with the key "rl:flag"
|
||||
// get the length of the set
|
||||
// get the max concurreny for the flag
|
||||
// if the length of the set is less than the max concurrency
|
||||
// we need to remove the flag from the "forbidden flags" list
|
||||
// we need to make sure that if there are any errors thrown in the task that we still perform the "after task" steps, and then rethrow the error
|
||||
wrapTask(t: Task): Task {
|
||||
return async (payload: unknown, helpers: JobHelpers) => {
|
||||
const flags = Object.keys(helpers.job.flags ?? {}).filter((flag) => flag.startsWith("rl:"));
|
||||
@@ -784,67 +811,68 @@ end
|
||||
return t(payload, helpers);
|
||||
}
|
||||
|
||||
// Before
|
||||
// TODO: handle errors
|
||||
const beforeResults = await Promise.allSettled(
|
||||
flags.map(async (flag) => this.#callBeforeTask(flag, String(helpers.job.id)))
|
||||
);
|
||||
|
||||
// TODO: Do something with the settled promises
|
||||
|
||||
logger.debug("[rate-limiter] beforeTask results", { beforeResults, flags });
|
||||
|
||||
try {
|
||||
await t(payload, helpers);
|
||||
} finally {
|
||||
const afterResults = await Promise.allSettled(
|
||||
flags.map(async (flag) => this.#callAfterTask(flag, String(helpers.job.id)))
|
||||
);
|
||||
|
||||
// TODO: Do something with the settled promises
|
||||
|
||||
logger.debug("[rate-limiter] afterTask results", { afterResults, flags });
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
#prefixKey(key: string): string {
|
||||
return `${this.prefix}:${key}`;
|
||||
}
|
||||
|
||||
async #callBeforeTask(flag: string, jobId: string) {
|
||||
const now = performance.now();
|
||||
const results = await this.redis.beforeTask(
|
||||
this.#prefixKey(flag),
|
||||
this.#prefixKey(`${flag}:maxSize`),
|
||||
this.#prefixKey("rl:forbiddenFlags"),
|
||||
jobId,
|
||||
flag
|
||||
);
|
||||
try {
|
||||
const now = performance.now();
|
||||
const results = await this.redis.beforeTask(
|
||||
flag,
|
||||
`${flag}:maxSize`,
|
||||
FORBIDDEN_FLAG_KEY,
|
||||
jobId,
|
||||
String(Date.now()),
|
||||
String(this.windowSize),
|
||||
flag,
|
||||
String(this.defaultMaxSize)
|
||||
);
|
||||
|
||||
const durationInMs = performance.now() - now;
|
||||
const durationInMs = performance.now() - now;
|
||||
|
||||
return {
|
||||
results,
|
||||
durationInMs,
|
||||
};
|
||||
return {
|
||||
results,
|
||||
durationInMs,
|
||||
};
|
||||
} catch (error) {
|
||||
logger.error("Failed to call beforeTask", { error, flag, jobId });
|
||||
}
|
||||
}
|
||||
|
||||
async #callAfterTask(flag: string, jobId: string) {
|
||||
const now = performance.now();
|
||||
const results = await this.redis.afterTask(
|
||||
this.#prefixKey(flag),
|
||||
this.#prefixKey(`${flag}:maxSize`),
|
||||
this.#prefixKey("rl:forbiddenFlags"),
|
||||
jobId,
|
||||
flag
|
||||
);
|
||||
try {
|
||||
const now = performance.now();
|
||||
|
||||
const durationInMs = performance.now() - now;
|
||||
const results = await this.redis.afterTask(
|
||||
flag,
|
||||
`${flag}:maxSize`,
|
||||
FORBIDDEN_FLAG_KEY,
|
||||
jobId,
|
||||
String(Date.now()),
|
||||
String(this.windowSize),
|
||||
flag,
|
||||
String(this.defaultMaxSize)
|
||||
);
|
||||
|
||||
return {
|
||||
results,
|
||||
durationInMs,
|
||||
};
|
||||
const durationInMs = performance.now() - now;
|
||||
|
||||
return {
|
||||
results,
|
||||
durationInMs,
|
||||
};
|
||||
} catch (error) {
|
||||
logger.error("Failed to call afterTask", { error, flag, jobId });
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,7 +3,11 @@ import { ScheduledPayloadSchema, addMissingVersionField } from "@trigger.dev/cor
|
||||
import { z } from "zod";
|
||||
import { prisma } from "~/db.server";
|
||||
import { env } from "~/env.server";
|
||||
import { RedisGraphileRateLimiter, ZodWorker } from "~/platform/zodWorker.server";
|
||||
import {
|
||||
GraphileRateLimiter,
|
||||
RedisGraphileRateLimiter,
|
||||
ZodWorker,
|
||||
} from "~/platform/zodWorker.server";
|
||||
import { sendEmail } from "./email.server";
|
||||
import { IndexEndpointService } from "./endpoints/indexEndpoint.server";
|
||||
import { PerformEndpointIndexService } from "./endpoints/performEndpointIndexService";
|
||||
@@ -433,7 +437,7 @@ function getExecutionWorkerQueue() {
|
||||
},
|
||||
shutdownTimeoutInMs: env.GRACEFUL_SHUTDOWN_TIMEOUT,
|
||||
schema: executionWorkerCatalog,
|
||||
rateLimiter: new RedisGraphileRateLimiter(),
|
||||
rateLimiter: getRateLimiter(),
|
||||
tasks: {
|
||||
performRunExecutionV2: {
|
||||
priority: 0, // smaller number = higher priority
|
||||
@@ -508,4 +512,43 @@ function getTaskOperationWorkerQueue() {
|
||||
});
|
||||
}
|
||||
|
||||
function getRateLimiter(): GraphileRateLimiter | undefined {
|
||||
if (env.REDIS_HOST && env.REDIS_PORT) {
|
||||
if (env.REDIS_READER_HOST) {
|
||||
return new RedisGraphileRateLimiter({
|
||||
cluster: {
|
||||
startupNodes: [
|
||||
{ host: env.REDIS_HOST, port: env.REDIS_PORT },
|
||||
{ host: env.REDIS_READER_HOST, port: env.REDIS_READER_PORT ?? env.REDIS_PORT },
|
||||
],
|
||||
options: {
|
||||
keyPrefix: "tr:gw:",
|
||||
scaleReads: "slave",
|
||||
redisOptions: {
|
||||
username: env.REDIS_USERNAME,
|
||||
password: env.REDIS_PASSWORD,
|
||||
tls: {},
|
||||
enableAutoPipelining: true,
|
||||
},
|
||||
dnsLookup: (address, callback) => callback(null, address),
|
||||
},
|
||||
},
|
||||
defaultConcurrency: env.DEFAULT_ORG_EXECUTION_CONCURRENCY_LIMIT,
|
||||
});
|
||||
} else {
|
||||
return new RedisGraphileRateLimiter({
|
||||
redis: {
|
||||
keyPrefix: "tr:gw:",
|
||||
port: env.REDIS_PORT,
|
||||
host: env.REDIS_HOST,
|
||||
username: env.REDIS_USERNAME,
|
||||
password: env.REDIS_PASSWORD,
|
||||
enableAutoPipelining: true,
|
||||
},
|
||||
defaultConcurrency: env.DEFAULT_ORG_EXECUTION_CONCURRENCY_LIMIT,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
export { executionWorker, workerQueue, taskOperationWorker };
|
||||
|
||||
Reference in New Issue
Block a user