import "server-only"; import { and, eq, sql } from "drizzle-orm"; import { drizzle } from "drizzle-orm/postgres-js"; import postgres from "postgres"; import { subscription, type Subscription, usage, type Usage, } from "@/lib/db/schema"; import { type LimitKey, LIMITS, type Tier, todayPeriodKey, } from "./config"; const client = postgres(process.env.DATABASE_URL ?? ""); const db = drizzle(client); export class LimitExceededError extends Error { kind: LimitKey; tier: Tier; used: number; limit: number; constructor(args: { kind: LimitKey; tier: Tier; used: number; limit: number; }) { super( `Limit exceeded for ${args.kind}: ${args.used}/${args.limit} on ${args.tier} tier` ); this.kind = args.kind; this.tier = args.tier; this.used = args.used; this.limit = args.limit; } } export async function ensureSubscription(userId: string): Promise { const existing = await db .select() .from(subscription) .where(eq(subscription.userId, userId)) .limit(1); if (existing[0]) { return existing[0]; } const [row] = await db .insert(subscription) .values({ userId, tier: "free", status: "active" }) .returning(); return row; } export async function getSubscription( userId: string ): Promise { const [row] = await db .select() .from(subscription) .where(eq(subscription.userId, userId)) .limit(1); return row ?? null; } export async function getCurrentUsage(userId: string): Promise { const periodKey = todayPeriodKey(); const [existing] = await db .select() .from(usage) .where(and(eq(usage.userId, userId), eq(usage.periodKey, periodKey))) .limit(1); if (existing) { return existing; } const [created] = await db .insert(usage) .values({ userId, periodKey, messageCount: 0, toolCallCount: 0 }) .onConflictDoNothing() .returning(); if (created) { return created; } const [fallback] = await db .select() .from(usage) .where(and(eq(usage.userId, userId), eq(usage.periodKey, periodKey))) .limit(1); return fallback; } export async function checkAndConsume( userId: string, kind: LimitKey, units = 1 ): Promise<{ remaining: number; limit: number; tier: Tier }> { const sub = await ensureSubscription(userId); const tier = sub.tier as Tier; const limit = LIMITS[tier][kind]; const column = kind === "chat_message" ? usage.messageCount : usage.toolCallCount; const column_name = kind === "chat_message" ? "messageCount" : "toolCallCount"; const periodKey = todayPeriodKey(); await ensureUsageRow(userId, periodKey); const result = await db .update(usage) .set({ [column_name]: sql`${column} + ${units}`, updatedAt: new Date(), }) .where( and( eq(usage.userId, userId), eq(usage.periodKey, periodKey), sql`${column} + ${units} <= ${limit}` ) ) .returning(); if (result.length === 0) { const current = await getCurrentUsage(userId); const used = kind === "chat_message" ? current.messageCount : current.toolCallCount; throw new LimitExceededError({ kind, tier, used, limit }); } const updated = result[0]; const used = kind === "chat_message" ? updated.messageCount : updated.toolCallCount; return { remaining: Math.max(0, limit - used), limit, tier }; } async function ensureUsageRow(userId: string, periodKey: string) { await db .insert(usage) .values({ userId, periodKey, messageCount: 0, toolCallCount: 0 }) .onConflictDoNothing(); } export async function upgradeToPro( userId: string, args: { stripeCustomerId?: string; stripeSubscriptionId?: string; currentPeriodEnd?: Date; } = {} ): Promise { await ensureSubscription(userId); const [row] = await db .update(subscription) .set({ tier: "pro", status: "active", stripeCustomerId: args.stripeCustomerId, stripeSubscriptionId: args.stripeSubscriptionId, currentPeriodEnd: args.currentPeriodEnd, updatedAt: new Date(), }) .where(eq(subscription.userId, userId)) .returning(); return row; } export async function cancelSubscription( userId: string ): Promise { const [row] = await db .update(subscription) .set({ status: "canceled", updatedAt: new Date() }) .where(eq(subscription.userId, userId)) .returning(); return row; } export async function getQuotaSummary(userId: string) { const sub = await ensureSubscription(userId); const tier = sub.tier as Tier; const usage = await getCurrentUsage(userId); const tierLimits = LIMITS[tier]; return { tier, status: sub.status, periodKey: usage.periodKey, messages: { used: usage.messageCount, limit: tierLimits.chat_message, remaining: Math.max(0, tierLimits.chat_message - usage.messageCount), }, toolCalls: { used: usage.toolCallCount, limit: tierLimits.tool_call, remaining: Math.max(0, tierLimits.tool_call - usage.toolCallCount), }, }; } export type QuotaSummary = Awaited>;