|
@@ -47,6 +47,7 @@ import { i18n, type Key } from "~/i18n"
|
|
|
import { localeFromRequest } from "~/lib/language"
|
|
import { localeFromRequest } from "~/lib/language"
|
|
|
import { createModelTpmLimiter } from "./modelTpmLimiter"
|
|
import { createModelTpmLimiter } from "./modelTpmLimiter"
|
|
|
import { createModelTpsLimiter } from "./modelTpsLimiter"
|
|
import { createModelTpsLimiter } from "./modelTpsLimiter"
|
|
|
|
|
+import { createProviderBudgetTracker } from "./providerBudgetTracker"
|
|
|
import { accumulateUsage, HOT_WORKSPACES } from "./usageBatcher"
|
|
import { accumulateUsage, HOT_WORKSPACES } from "./usageBatcher"
|
|
|
|
|
|
|
|
type ZenData = Awaited<ReturnType<typeof ZenData.list>>
|
|
type ZenData = Awaited<ReturnType<typeof ZenData.list>>
|
|
@@ -132,6 +133,10 @@ export async function handler(
|
|
|
const modelTpmLimits = await modelTpmLimiter?.check()
|
|
const modelTpmLimits = await modelTpmLimiter?.check()
|
|
|
const modelTpsLimiter = createModelTpsLimiter(modelInfo.providers)
|
|
const modelTpsLimiter = createModelTpsLimiter(modelInfo.providers)
|
|
|
const modelTpsLimits = await modelTpsLimiter?.check()
|
|
const modelTpsLimits = await modelTpsLimiter?.check()
|
|
|
|
|
+ const providerBudgetTracker = createProviderBudgetTracker(
|
|
|
|
|
+ modelInfo.providers.map((provider) => ({ ...zenData.providers[provider.id], ...provider })),
|
|
|
|
|
+ )
|
|
|
|
|
+ const providerBudgetUsage = await providerBudgetTracker?.check()
|
|
|
|
|
|
|
|
const retriableRequest = async (retry: RetryOptions = { excludeProviders: [], retryCount: 0 }) => {
|
|
const retriableRequest = async (retry: RetryOptions = { excludeProviders: [], retryCount: 0 }) => {
|
|
|
const providerInfo = selectProvider(
|
|
const providerInfo = selectProvider(
|
|
@@ -145,12 +150,16 @@ export async function handler(
|
|
|
stickyProvider,
|
|
stickyProvider,
|
|
|
modelTpmLimits,
|
|
modelTpmLimits,
|
|
|
modelTpsLimits,
|
|
modelTpsLimits,
|
|
|
|
|
+ providerBudgetUsage,
|
|
|
)
|
|
)
|
|
|
validateModelSettings(billingSource, authInfo)
|
|
validateModelSettings(billingSource, authInfo)
|
|
|
updateProviderKey(authInfo, providerInfo)
|
|
updateProviderKey(authInfo, providerInfo)
|
|
|
logger.metric({
|
|
logger.metric({
|
|
|
provider: providerInfo.id,
|
|
provider: providerInfo.id,
|
|
|
"provider.model": providerInfo.model,
|
|
"provider.model": providerInfo.model,
|
|
|
|
|
+ ...(providerBudgetUsage?.[providerInfo.id]
|
|
|
|
|
+ ? { "provider.budget_usage": providerBudgetUsage?.[providerInfo.id] }
|
|
|
|
|
+ : {}),
|
|
|
})
|
|
})
|
|
|
|
|
|
|
|
const startTimestamp = Date.now()
|
|
const startTimestamp = Date.now()
|
|
@@ -257,6 +266,7 @@ export async function handler(
|
|
|
const costInfo = calculateCost(modelInfo, usageInfo)
|
|
const costInfo = calculateCost(modelInfo, usageInfo)
|
|
|
await trialLimiter?.track(usageInfo)
|
|
await trialLimiter?.track(usageInfo)
|
|
|
await modelTpmLimiter?.track(providerInfo.id, providerInfo.model, usageInfo)
|
|
await modelTpmLimiter?.track(providerInfo.id, providerInfo.model, usageInfo)
|
|
|
|
|
+ await providerBudgetTracker?.track(providerInfo.id, costInfo.totalCostInCent)
|
|
|
await trackUsage(sessionId, billingSource, authInfo, modelInfo, providerInfo, usageInfo, costInfo)
|
|
await trackUsage(sessionId, billingSource, authInfo, modelInfo, providerInfo, usageInfo, costInfo)
|
|
|
await reload(billingSource, authInfo, costInfo)
|
|
await reload(billingSource, authInfo, costInfo)
|
|
|
json.cost = calculateOccurredCost(billingSource, costInfo)
|
|
json.cost = calculateOccurredCost(billingSource, costInfo)
|
|
@@ -317,6 +327,7 @@ export async function handler(
|
|
|
timestampLastByte,
|
|
timestampLastByte,
|
|
|
usageInfo,
|
|
usageInfo,
|
|
|
)
|
|
)
|
|
|
|
|
+ await providerBudgetTracker?.track(providerInfo.id, costInfo.totalCostInCent)
|
|
|
await trackUsage(sessionId, billingSource, authInfo, modelInfo, providerInfo, usageInfo, costInfo)
|
|
await trackUsage(sessionId, billingSource, authInfo, modelInfo, providerInfo, usageInfo, costInfo)
|
|
|
await reload(billingSource, authInfo, costInfo)
|
|
await reload(billingSource, authInfo, costInfo)
|
|
|
const cost = calculateOccurredCost(billingSource, costInfo)
|
|
const cost = calculateOccurredCost(billingSource, costInfo)
|
|
@@ -483,6 +494,7 @@ export async function handler(
|
|
|
stickyProviderId: string | undefined,
|
|
stickyProviderId: string | undefined,
|
|
|
modelTpmLimits: Record<string, number> | undefined,
|
|
modelTpmLimits: Record<string, number> | undefined,
|
|
|
modelTpsLimits: Record<string, { qualify: number; unqualify: number }> | undefined,
|
|
modelTpsLimits: Record<string, { qualify: number; unqualify: number }> | undefined,
|
|
|
|
|
+ providerBudgetUsage: Record<string, number> | undefined,
|
|
|
) {
|
|
) {
|
|
|
const modelProvider = (() => {
|
|
const modelProvider = (() => {
|
|
|
// Byok is top priority b/c if user set their own API key, we should use it
|
|
// Byok is top priority b/c if user set their own API key, we should use it
|
|
@@ -505,6 +517,12 @@ export async function handler(
|
|
|
const providers = allProviders
|
|
const providers = allProviders
|
|
|
.filter((provider) => provider.weight !== 0)
|
|
.filter((provider) => provider.weight !== 0)
|
|
|
.filter((provider) => !retry.excludeProviders.includes(provider.id))
|
|
.filter((provider) => !retry.excludeProviders.includes(provider.id))
|
|
|
|
|
+ .filter((provider) => {
|
|
|
|
|
+ if (provider.budgetMode !== "fill") return true
|
|
|
|
|
+ const budget = zenData.providers[provider.id]?.budget
|
|
|
|
|
+ if (budget === undefined) return false
|
|
|
|
|
+ return (providerBudgetUsage?.[provider.id] ?? 0) < centsToMicroCents(budget * 100)
|
|
|
|
|
+ })
|
|
|
.filter((provider) => {
|
|
.filter((provider) => {
|
|
|
if (!provider.tpmLimit) return true
|
|
if (!provider.tpmLimit) return true
|
|
|
const usage = modelTpmLimits?.[`${provider.id}/${provider.model}`] ?? 0
|
|
const usage = modelTpmLimits?.[`${provider.id}/${provider.model}`] ?? 0
|