From 8b8e7913dc3ee0c43dfb18de339801fe4c1ab3c8 Mon Sep 17 00:00:00 2001 From: Owen Date: Mon, 10 Aug 2026 17:17:24 -0400 Subject: [PATCH] budget enforcement logic --- server/db/pg/schema/schema.ts | 96 ++++++++ server/db/sqlite/schema/schema.ts | 98 ++++++++ server/lib/aiBudgetEnforcement.ts | 341 +++++++++++++++++++++++++++ server/routers/aiGateway/pipeline.ts | 78 +++++- 4 files changed, 607 insertions(+), 6 deletions(-) create mode 100644 server/lib/aiBudgetEnforcement.ts diff --git a/server/db/pg/schema/schema.ts b/server/db/pg/schema/schema.ts index 7ded27eb7..c9a45618c 100644 --- a/server/db/pg/schema/schema.ts +++ b/server/db/pg/schema/schema.ts @@ -1758,6 +1758,100 @@ export const aiBudgets = pgTable( ] ); +export const aiUsageRecords = pgTable( + "aiUsageRecords", + { + id: serial("id").primaryKey(), + orgId: varchar("orgId") + .notNull() + .references(() => orgs.orgId, { onDelete: "cascade" }), + providerId: integer("providerId") + .notNull() + .references(() => aiProviders.providerId, { onDelete: "cascade" }), + resourceId: integer("resourceId").references( + () => resources.resourceId, + { onDelete: "cascade" } + ), + siteResourceId: integer("siteResourceId").references( + () => siteResources.siteResourceId, + { onDelete: "cascade" } + ), + userId: varchar("userId").references(() => users.userId, { + onDelete: "set null" + }), + requestedModel: varchar("requestedModel").notNull(), + promptTokens: integer("promptTokens").notNull().default(0), + cacheReadTokens: integer("cacheReadTokens").notNull().default(0), + cacheWriteTokens: integer("cacheWriteTokens").notNull().default(0), + completionTokens: integer("completionTokens").notNull().default(0), + reasoningTokens: integer("reasoningTokens").notNull().default(0), + totalTokens: integer("totalTokens").notNull().default(0), + costUsd: real("costUsd"), + estimated: boolean("estimated").notNull().default(false), + createdAt: bigint("createdAt", { mode: "number" }).notNull() + }, + (t) => [ + index("idx_ai_usage_records_org_provider_created").on( + t.orgId, + t.providerId, + t.createdAt + ), + index("idx_ai_usage_records_org_resource_created").on( + t.orgId, + t.resourceId, + t.createdAt + ), + index("idx_ai_usage_records_org_site_resource_created").on( + t.orgId, + t.siteResourceId, + t.createdAt + ), + index("idx_ai_usage_records_org_user_created").on( + t.orgId, + t.userId, + t.createdAt + ) + ] +); + +export const aiBudgetBreachEvents = pgTable( + "aiBudgetBreachEvents", + { + id: serial("id").primaryKey(), + orgId: varchar("orgId") + .notNull() + .references(() => orgs.orgId, { onDelete: "cascade" }), + budgetId: integer("budgetId") + .notNull() + .references(() => aiBudgets.budgetId, { onDelete: "cascade" }), + enforcement: varchar("enforcement").$type<"hard" | "soft">().notNull(), + unit: varchar("unit").$type<"usd" | "tokens">().notNull(), + period: varchar("period") + .$type< + | "monthly" + | "yearly" + | "lifetime" + | "daily" + | "hourly" + | "weekly" + >() + .notNull(), + amount: real("amount").notNull(), + usageAmount: real("usageAmount").notNull(), + blocked: boolean("blocked").notNull(), + requestUserId: varchar("requestUserId").references(() => users.userId, { + onDelete: "set null" + }), + createdAt: bigint("createdAt", { mode: "number" }).notNull() + }, + (t) => [ + index("idx_ai_budget_breach_events_budget_created").on( + t.budgetId, + t.createdAt + ) + ] +); + export type Org = InferSelectModel; export type User = InferSelectModel; export type Site = InferSelectModel; @@ -1845,6 +1939,8 @@ export type ResourcePolicyRule = InferSelectModel; export type AiProvider = InferSelectModel; export type AiModel = InferSelectModel; export type AiBudget = InferSelectModel; +export type AiUsageRecord = InferSelectModel; +export type AiBudgetBreachEvent = InferSelectModel; export type ResourceAiProvider = InferSelectModel; export type SiteResourceAiProvider = InferSelectModel< typeof siteResourceAiProviders diff --git a/server/db/sqlite/schema/schema.ts b/server/db/sqlite/schema/schema.ts index 34fe67a0f..8017cd3b7 100644 --- a/server/db/sqlite/schema/schema.ts +++ b/server/db/sqlite/schema/schema.ts @@ -1744,6 +1744,102 @@ export const aiBudgets = sqliteTable( ] ); +export const aiUsageRecords = sqliteTable( + "aiUsageRecords", + { + id: integer("id").primaryKey({ autoIncrement: true }), + orgId: text("orgId") + .notNull() + .references(() => orgs.orgId, { onDelete: "cascade" }), + providerId: integer("providerId") + .notNull() + .references(() => aiProviders.providerId, { onDelete: "cascade" }), + resourceId: integer("resourceId").references( + () => resources.resourceId, + { onDelete: "cascade" } + ), + siteResourceId: integer("siteResourceId").references( + () => siteResources.siteResourceId, + { onDelete: "cascade" } + ), + userId: text("userId").references(() => users.userId, { + onDelete: "set null" + }), + requestedModel: text("requestedModel").notNull(), + promptTokens: integer("promptTokens").notNull().default(0), + cacheReadTokens: integer("cacheReadTokens").notNull().default(0), + cacheWriteTokens: integer("cacheWriteTokens").notNull().default(0), + completionTokens: integer("completionTokens").notNull().default(0), + reasoningTokens: integer("reasoningTokens").notNull().default(0), + totalTokens: integer("totalTokens").notNull().default(0), + costUsd: real("costUsd"), + estimated: integer("estimated", { mode: "boolean" }) + .notNull() + .default(false), + createdAt: integer("createdAt").notNull() + }, + (t) => [ + index("idx_ai_usage_records_org_provider_created").on( + t.orgId, + t.providerId, + t.createdAt + ), + index("idx_ai_usage_records_org_resource_created").on( + t.orgId, + t.resourceId, + t.createdAt + ), + index("idx_ai_usage_records_org_site_resource_created").on( + t.orgId, + t.siteResourceId, + t.createdAt + ), + index("idx_ai_usage_records_org_user_created").on( + t.orgId, + t.userId, + t.createdAt + ) + ] +); + +export const aiBudgetBreachEvents = sqliteTable( + "aiBudgetBreachEvents", + { + id: integer("id").primaryKey({ autoIncrement: true }), + orgId: text("orgId") + .notNull() + .references(() => orgs.orgId, { onDelete: "cascade" }), + budgetId: integer("budgetId") + .notNull() + .references(() => aiBudgets.budgetId, { onDelete: "cascade" }), + enforcement: text("enforcement").$type<"hard" | "soft">().notNull(), + unit: text("unit").$type<"usd" | "tokens">().notNull(), + period: text("period") + .$type< + | "monthly" + | "yearly" + | "lifetime" + | "daily" + | "hourly" + | "weekly" + >() + .notNull(), + amount: real("amount").notNull(), + usageAmount: real("usageAmount").notNull(), + blocked: integer("blocked", { mode: "boolean" }).notNull(), + requestUserId: text("requestUserId").references(() => users.userId, { + onDelete: "set null" + }), + createdAt: integer("createdAt").notNull() + }, + (t) => [ + index("idx_ai_budget_breach_events_budget_created").on( + t.budgetId, + t.createdAt + ) + ] +); + export type Org = InferSelectModel; export type User = InferSelectModel; export type Site = InferSelectModel; @@ -1829,6 +1925,8 @@ export type UserPolicy = InferSelectModel; export type AiProvider = InferSelectModel; export type AiModel = InferSelectModel; export type AiBudget = InferSelectModel; +export type AiUsageRecord = InferSelectModel; +export type AiBudgetBreachEvent = InferSelectModel; export type ResourceAiProvider = InferSelectModel; export type SiteResourceAiProvider = InferSelectModel< typeof siteResourceAiProviders diff --git a/server/lib/aiBudgetEnforcement.ts b/server/lib/aiBudgetEnforcement.ts new file mode 100644 index 000000000..8f70340fa --- /dev/null +++ b/server/lib/aiBudgetEnforcement.ts @@ -0,0 +1,341 @@ +import { and, eq, gte, inArray, isNull, or, sql, SQL } from "drizzle-orm"; +import { + AiBudget, + aiBudgetBreachEvents, + aiBudgets, + aiModels, + aiUsageRecords, + db, + userOrgRoles +} from "@server/db"; +import { modelKeyMatches } from "@server/lib/aiModelKeyMatch"; +import type { AiUsage } from "@server/lib/aiUsageExtraction"; +import logger from "@server/logger"; + +type BudgetPeriod = AiBudget["period"]; + +const PERIOD_DURATIONS_MS: Record, number> = { + hourly: 60 * 60 * 1000, + daily: 24 * 60 * 60 * 1000, + weekly: 7 * 24 * 60 * 60 * 1000, + monthly: 30 * 24 * 60 * 60 * 1000, + yearly: 365 * 24 * 60 * 60 * 1000 +}; + +// Budget periods are trailing windows from "now", not calendar-aligned +// (e.g. "daily" = last 24h). "lifetime" has no lower bound. +function windowStart(period: BudgetPeriod, now: number): number { + if (period === "lifetime") { + return 0; + } + return now - PERIOD_DURATIONS_MS[period]; +} + +export type BudgetScopeContext = { + orgId: string; + providerId: number; + requestedModel: string; + resourceId: number | null; + siteResourceId: number | null; + roleIds: number[]; + requestUserId: string | null; +}; + +/** + * Every budget that could apply to this request: the provider itself, any + * model on that provider whose (possibly wildcarded) modelKey matches the + * requested model, the target resource/site-resource, and any role the + * requesting user holds in the org. + */ +export async function resolveApplicableBudgets( + ctx: BudgetScopeContext +): Promise { + const providerModels = await db + .select({ modelId: aiModels.modelId, modelKey: aiModels.modelKey }) + .from(aiModels) + .where( + and( + eq(aiModels.providerId, ctx.providerId), + eq(aiModels.enabled, true) + ) + ); + + const matchingModelIds = providerModels + .filter((m) => modelKeyMatches(m.modelKey, ctx.requestedModel)) + .map((m) => m.modelId); + + const scopeConditions: SQL[] = [ + and( + eq(aiBudgets.providerId, ctx.providerId), + isNull(aiBudgets.modelId) + )! + ]; + if (matchingModelIds.length > 0) { + scopeConditions.push(inArray(aiBudgets.modelId, matchingModelIds)); + } + if (ctx.resourceId != null) { + scopeConditions.push(eq(aiBudgets.resourceId, ctx.resourceId)); + } + if (ctx.siteResourceId != null) { + scopeConditions.push(eq(aiBudgets.siteResourceId, ctx.siteResourceId)); + } + if (ctx.roleIds.length > 0) { + scopeConditions.push(inArray(aiBudgets.roleId, ctx.roleIds)); + } + + return db + .select() + .from(aiBudgets) + .where( + and( + eq(aiBudgets.orgId, ctx.orgId), + eq(aiBudgets.enabled, true), + or(...scopeConditions) + ) + ); +} + +async function sumUsageAmount( + where: SQL, + unit: AiBudget["unit"] +): Promise { + const column = + unit === "usd" ? aiUsageRecords.costUsd : aiUsageRecords.totalTokens; + const [row] = await db + .select({ total: sql`coalesce(sum(${column}), 0)` }) + .from(aiUsageRecords) + .where(where); + return Number(row?.total ?? 0); +} + +/** + * Sums recorded usage for a single budget's scope + rolling window. Model + * budgets can't be pushed down to SQL because the model's key may itself be + * a glob, so those rows are fetched for the provider+window and matched in + * JS the same way access-control matching does. + */ +export async function sumUsageForBudget( + budget: AiBudget, + ctx: BudgetScopeContext, + now: number +): Promise { + const start = windowStart(budget.period, now); + + if (budget.modelId != null) { + const [model] = await db + .select({ + providerId: aiModels.providerId, + modelKey: aiModels.modelKey + }) + .from(aiModels) + .where(eq(aiModels.modelId, budget.modelId)) + .limit(1); + if (!model) { + return 0; + } + const rows = await db + .select({ + requestedModel: aiUsageRecords.requestedModel, + costUsd: aiUsageRecords.costUsd, + totalTokens: aiUsageRecords.totalTokens + }) + .from(aiUsageRecords) + .where( + and( + eq(aiUsageRecords.orgId, ctx.orgId), + eq(aiUsageRecords.providerId, model.providerId), + gte(aiUsageRecords.createdAt, start) + ) + ); + return rows + .filter((r) => modelKeyMatches(model.modelKey, r.requestedModel)) + .reduce( + (sum, r) => + sum + + (budget.unit === "usd" ? (r.costUsd ?? 0) : r.totalTokens), + 0 + ); + } + + if (budget.providerId != null) { + return sumUsageAmount( + and( + eq(aiUsageRecords.orgId, ctx.orgId), + eq(aiUsageRecords.providerId, budget.providerId), + gte(aiUsageRecords.createdAt, start) + )!, + budget.unit + ); + } + + if (budget.resourceId != null) { + return sumUsageAmount( + and( + eq(aiUsageRecords.orgId, ctx.orgId), + eq(aiUsageRecords.resourceId, budget.resourceId), + gte(aiUsageRecords.createdAt, start) + )!, + budget.unit + ); + } + + if (budget.siteResourceId != null) { + return sumUsageAmount( + and( + eq(aiUsageRecords.orgId, ctx.orgId), + eq(aiUsageRecords.siteResourceId, budget.siteResourceId), + gte(aiUsageRecords.createdAt, start) + )!, + budget.unit + ); + } + + if (budget.roleId != null) { + const members = await db + .select({ userId: userOrgRoles.userId }) + .from(userOrgRoles) + .where( + and( + eq(userOrgRoles.roleId, budget.roleId), + eq(userOrgRoles.orgId, ctx.orgId) + ) + ); + const userIds = members.map((m) => m.userId); + if (userIds.length === 0) { + return 0; + } + return sumUsageAmount( + and( + eq(aiUsageRecords.orgId, ctx.orgId), + inArray(aiUsageRecords.userId, userIds), + gte(aiUsageRecords.createdAt, start) + )!, + budget.unit + ); + } + + return 0; +} + +// Throttled to one durable event per budget per breach window, so a soft +// budget being exceeded doesn't write a row on every subsequent request +// while it stays over. +async function recordBreachEventIfNew( + budget: AiBudget, + ctx: BudgetScopeContext, + usageAmount: number, + now: number +): Promise { + try { + const start = windowStart(budget.period, now); + const [existing] = await db + .select({ id: aiBudgetBreachEvents.id }) + .from(aiBudgetBreachEvents) + .where( + and( + eq(aiBudgetBreachEvents.budgetId, budget.budgetId), + gte(aiBudgetBreachEvents.createdAt, start) + ) + ) + .limit(1); + if (existing) { + return; + } + + await db.insert(aiBudgetBreachEvents).values({ + orgId: ctx.orgId, + budgetId: budget.budgetId, + enforcement: budget.enforcement, + unit: budget.unit, + period: budget.period, + amount: budget.amount, + usageAmount, + blocked: budget.enforcement === "hard", + requestUserId: ctx.requestUserId, + createdAt: now + }); + } catch (error) { + logger.error("Failed to record AI budget breach event", { + error, + budgetId: budget.budgetId + }); + } +} + +export type BudgetCheckResult = { + blocked: boolean; + blockingBudget?: AiBudget; +}; + +export async function checkBudgets( + ctx: BudgetScopeContext +): Promise { + const budgets = await resolveApplicableBudgets(ctx); + if (budgets.length === 0) { + return { blocked: false }; + } + + const now = Date.now(); + let blockingBudget: AiBudget | undefined; + + for (const budget of budgets) { + const usage = await sumUsageForBudget(budget, ctx, now); + if (usage < budget.amount) { + continue; + } + + await recordBreachEventIfNew(budget, ctx, usage, now); + + if (budget.enforcement === "hard" && !blockingBudget) { + blockingBudget = budget; + } + } + + return blockingBudget + ? { blocked: true, blockingBudget } + : { blocked: false }; +} + +export type UsageRecordInput = { + orgId: string; + providerId: number; + resourceId: number | null; + siteResourceId: number | null; + userId: string | null; + requestedModel: string; + usage: AiUsage; + costUsd: number | null; + createdAt?: number; +}; + +export async function recordUsage(input: UsageRecordInput): Promise { + try { + const { usage } = input; + const totalTokens = + usage.promptTokens + + usage.cacheReadTokens + + usage.cacheWriteTokens + + usage.completionTokens + + usage.reasoningTokens; + + await db.insert(aiUsageRecords).values({ + orgId: input.orgId, + providerId: input.providerId, + resourceId: input.resourceId, + siteResourceId: input.siteResourceId, + userId: input.userId, + requestedModel: input.requestedModel, + promptTokens: usage.promptTokens, + cacheReadTokens: usage.cacheReadTokens, + cacheWriteTokens: usage.cacheWriteTokens, + completionTokens: usage.completionTokens, + reasoningTokens: usage.reasoningTokens, + totalTokens, + costUsd: input.costUsd, + estimated: usage.estimated, + createdAt: input.createdAt ?? Date.now() + }); + } catch (error) { + logger.error("Failed to record AI usage", { error }); + } +} diff --git a/server/routers/aiGateway/pipeline.ts b/server/routers/aiGateway/pipeline.ts index bcff1a501..802f5f453 100644 --- a/server/routers/aiGateway/pipeline.ts +++ b/server/routers/aiGateway/pipeline.ts @@ -51,6 +51,7 @@ import { } from "@server/lib/aiModelKeyMatch"; import { aiGatewayUpstreamFetch } from "@server/lib/aiGatewayUpstreamFetch"; import { getModelPricing, calculateAiCost } from "@server/lib/aiModelPricing"; +import { checkBudgets, recordUsage } from "@server/lib/aiBudgetEnforcement"; import { extractUsage, estimateUsage, @@ -141,6 +142,7 @@ export type RequestUser = { email: string | null; name: string | null; role: string | null; + roleIds: number[]; }; // Identity headers forwarded to the upstream inference endpoint when the @@ -193,7 +195,8 @@ async function buildRequestUser( username: user.username, email: user.email, name: user.name, - role: orgRoles.map((r) => r.roleName).join(", ") || null + role: orgRoles.map((r) => r.roleName).join(", ") || null, + roleIds: orgRoles.map((r) => r.roleId) }; localCache.set(cacheKey, requestUser, REQUEST_USER_TTL_SEC); @@ -526,6 +529,10 @@ function logAiUsageAndCost(args: { responseText: string; isStream: boolean; headers: Headers; + orgId: string | null; + resourceId: number | null; + siteResourceId: number | null; + requestUserId: string | null; }): void { const { capability, @@ -534,7 +541,11 @@ function logAiUsageAndCost(args: { requestBody, responseText, isStream, - headers + headers, + orgId, + resourceId, + siteResourceId, + requestUserId } = args; let usage: AiUsage | null = extractUsage( @@ -565,6 +576,19 @@ function logAiUsageAndCost(args: { pricingApproximate: pricing?.approximate ?? null, totalCostUsd: cost?.totalCost ?? null }); + + if (orgId) { + void recordUsage({ + orgId, + providerId: provider.providerId, + resourceId, + siteResourceId, + userId: requestUserId, + requestedModel: model ?? "unknown", + usage, + costUsd: cost?.totalCost ?? null + }); + } } export async function handleAiGatewayProxy( @@ -597,8 +621,13 @@ export async function handleAiGatewayProxy( }); } - const { attachments, resourceListsByProvider, resourceId, orgId } = - target; + const { + attachments, + resourceListsByProvider, + resourceId, + siteResourceId, + orgId + } = target; const capableAttachments = attachments.filter((a) => providerHasCapability(a.provider.capabilities, capability) @@ -637,6 +666,35 @@ export async function handleAiGatewayProxy( const { provider } = selection; + if (orgId) { + const budgetCheck = await checkBudgets({ + orgId, + providerId: provider.providerId, + requestedModel: requestedModel!, + resourceId, + siteResourceId, + roleIds: requestUser?.roleIds ?? [], + requestUserId: requestUser?.userId ?? null + }); + + if (budgetCheck.blocked) { + logger.warn("AI gateway request blocked by budget", { + budgetId: budgetCheck.blockingBudget?.budgetId, + orgId, + providerId: provider.providerId, + requestedModel, + resourceId, + siteResourceId, + userId: requestUser?.userId ?? null + }); + return res.status(HttpCode.TOO_MANY_REQUESTS).json({ + error: { + message: "AI usage budget exceeded for this request" + } + }); + } + } + if (provider.type === "custom" && provider.routingMode === "target") { return await proxyAiGatewayToSiteTarget( req, @@ -814,7 +872,11 @@ export async function handleAiGatewayProxy( requestBody: req.body, responseText: fullText, isStream: true, - headers: upstreamRes.headers + headers: upstreamRes.headers, + orgId, + resourceId, + siteResourceId, + requestUserId: requestUser?.userId ?? null }); } return; @@ -829,7 +891,11 @@ export async function handleAiGatewayProxy( requestBody: req.body, responseText: text, isStream: false, - headers: upstreamRes.headers + headers: upstreamRes.headers, + orgId, + resourceId, + siteResourceId, + requestUserId: requestUser?.userId ?? null }); return res.send(text); } catch (error) {