This commit is contained in:
Nicholas Koben Kao
2024-01-24 16:26:16 -08:00
38 changed files with 7065 additions and 285 deletions
@@ -0,0 +1,175 @@
import { auth, clerkClient } from "@clerk/nextjs";
import { headers } from "next/headers";
import { redirect } from "next/navigation";
import { stripe } from "../../../../../server/stripe";
import { NextResponse } from "next/server";
import { db } from "@/db/db";
import { subscriptionStatusTable } from "@/db/schema";
import { eq } from "drizzle-orm";
const secret = process.env.STRIPE_WEBHOOK_SECRET || "";
export async function POST(req: Request) {
try {
const body = await req.text();
const signature = headers().get("stripe-signature");
if (!signature)
return new NextResponse("Signature not found.", {
status: 500,
});
const event = stripe.webhooks.constructEvent(body, signature, secret);
console.log(event);
if (event.type === "checkout.session.completed") {
if (!event.data.object.customer_details?.email) {
throw new Error(`missing user email, ${event.id}`);
}
const userId = event.data.object.metadata?.userId;
const plan = event.data.object.metadata?.plan;
if (!userId) {
throw new Error(`missing itinerary_id on metadata, ${event.id}`);
}
if (!plan) {
throw new Error(`missing plan on metadata, ${event.id}`);
}
const orgId = event.data.object.metadata?.orgId;
const customerId =
typeof event.data.object.customer == "string"
? event.data.object.customer
: event.data.object.customer?.id;
const subscriptionId =
typeof event.data.object.subscription == "string"
? event.data.object.subscription
: event.data.object.subscription?.id;
if (!customerId) {
throw new Error(`missing customerId, ${event.id}`);
}
if (!subscriptionId) {
throw new Error(`missing subscriptionId, ${event.id}`);
}
const items = await stripe.subscriptionItems.list({
subscription: subscriptionId,
limit: 5,
});
// getting the subscription item id for the api plan
const subscription_item_api_id = items.data.find(
(x) => x.price.id === process.env.STRIPE_PR_API,
)?.id;
// the plan could be either pro or enterprise
const subscription_item_plan_id =
items.data.find((x) => x.price.id === process.env.STRIPE_PR_PRO)?.id ??
items.data.find((x) => x.price.id === process.env.STRIPE_PR_ENTERPRISE)
?.id;
if (!subscription_item_api_id) {
throw new Error(
`missing plan on subscription_item_api_id, ${event.id}`,
);
}
if (!subscription_item_plan_id) {
throw new Error(
`missing plan on subscription_item_plan_id, ${event.id}`,
);
}
console.log(items);
await db
.insert(subscriptionStatusTable)
.values({
stripe_customer_id: customerId,
org_id: orgId,
user_id: userId,
subscription_id: subscriptionId,
plan: plan as "pro" | "enterprise" | "basic",
status: "active",
subscription_item_api_id: subscription_item_api_id,
subscription_item_plan_id: subscription_item_plan_id,
})
.onConflictDoNothing();
// updateDatabase(event.data.object.metadata.itinerary_id);
// sendEmail(event.data.object.customer_details.email);
} else if (event.type === "customer.subscription.paused") {
const customerId =
typeof event.data.object.customer == "string"
? event.data.object.customer
: event.data.object.customer?.id;
if (!customerId) {
throw new Error(`missing customerId, ${event.id}`);
}
await db
.update(subscriptionStatusTable)
.set({ status: "paused" })
.where(eq(subscriptionStatusTable.stripe_customer_id, customerId));
} else if (event.type === "customer.subscription.resumed") {
const customerId =
typeof event.data.object.customer == "string"
? event.data.object.customer
: event.data.object.customer?.id;
if (!customerId) {
throw new Error(`missing customerId, ${event.id}`);
}
await db
.update(subscriptionStatusTable)
.set({ status: "active" })
.where(eq(subscriptionStatusTable.stripe_customer_id, customerId));
} else if (event.type === "customer.subscription.deleted") {
const customerId =
typeof event.data.object.customer == "string"
? event.data.object.customer
: event.data.object.customer?.id;
if (!customerId) {
throw new Error(`missing customerId, ${event.id}`);
}
await db
.update(subscriptionStatusTable)
.set({ status: "deleted" })
.where(eq(subscriptionStatusTable.stripe_customer_id, customerId));
} else if (event.type === "customer.subscription.updated") {
const customerId =
typeof event.data.object.customer == "string"
? event.data.object.customer
: event.data.object.customer?.id;
if (!customerId) {
throw new Error(`missing customerId, ${event.id}`);
}
const cancel_at_period_end = event.data.object.cancel_at_period_end;
await db
.update(subscriptionStatusTable)
.set({ cancel_at_period_end: cancel_at_period_end })
.where(eq(subscriptionStatusTable.stripe_customer_id, customerId));
}
return NextResponse.json({ result: event, ok: true });
} catch (error) {
console.error(error);
return NextResponse.json(
{
message: `Something went wrong: ${error}`,
ok: false,
},
{ status: 500 },
);
}
}
@@ -0,0 +1,46 @@
import { stripe } from "@/server/stripe";
import { createCheckout } from "@/server/linkToPricing";
import { auth, clerkClient } from "@clerk/nextjs";
import { redirect } from "next/navigation";
import { getUrlServerSide } from "@/server/getUrlServerSide";
export async function GET(req: Request) {
const { userId, orgId } = auth();
const plan = new URL(req.url).searchParams.get("plan");
if (!userId) return redirect("/");
if (!plan) return redirect("/pricing");
const user = await clerkClient.users.getUser(userId);
const mapping = {
pro: process.env.STRIPE_PR_PRO,
enterprise: process.env.STRIPE_PR_ENTERPRISE,
};
const api = process.env.STRIPE_PR_API;
const session = await stripe.checkout.sessions.create({
success_url: getUrlServerSide(),
line_items: [
{
price: mapping[plan as "pro" | "enterprise"],
quantity: 1,
},
{
price: api,
},
],
metadata: {
userId: userId,
orgId: orgId ?? null,
plan: plan,
},
client_reference_id: orgId ?? userId,
customer_email: user.emailAddresses[0].emailAddress,
mode: "subscription",
});
if (session.url) redirect(session.url);
}
@@ -0,0 +1,41 @@
import { stripe } from "@/server/stripe";
import { createCheckout } from "@/server/linkToPricing";
import { auth, clerkClient } from "@clerk/nextjs";
import { redirect } from "next/navigation";
import { getUrlServerSide } from "@/server/getUrlServerSide";
import { db } from "@/db/db";
import { and, eq, isNull } from "drizzle-orm";
import { subscriptionStatusTable } from "@/db/schema";
import { getCurrentPlan } from "@/server/getCurrentPlan";
export async function GET(req: Request) {
const { userId, orgId } = auth();
if (!userId) return redirect("/");
const change = new URL(req.url).searchParams.get("change");
const sub = await getCurrentPlan({
org_id: orgId,
user_id: userId,
});
if (!sub) return redirect("/pricing");
const session = await stripe.billingPortal.sessions.create({
customer: sub.stripe_customer_id,
return_url: getUrlServerSide() + "/pricing",
// flow_data:
// change === "true" && sub.subscription_id
// ? {
// type: "subscription_update",
// subscription_update: {
// subscription: sub.subscription_id,
// },
// }
// : undefined,
});
redirect(session.url);
// if (session.url) redirect(session.url);
}
+47 -4
View File
@@ -1,6 +1,13 @@
import { parseDataSafe } from "../../../../lib/parseDataSafe";
import { db } from "@/db/db";
import { workflowRunOutputs, workflowRunsTable } from "@/db/schema";
import {
userUsageTable,
workflowRunOutputs,
workflowRunsTable,
workflowTable,
} from "@/db/schema";
import { getCurrentPlan } from "@/server/getCurrentPlan";
import { stripe } from "@/server/stripe";
import { eq } from "drizzle-orm";
import { NextResponse } from "next/server";
import { z } from "zod";
@@ -27,8 +34,7 @@ export async function POST(request: Request) {
data: output_data,
});
} else if (status) {
// console.log("status", status);
const workflow_run = await db
const [workflow_run] = await db
.update(workflowRunsTable)
.set({
status: status,
@@ -37,6 +43,43 @@ export async function POST(request: Request) {
})
.where(eq(workflowRunsTable.id, run_id))
.returning();
// Need to filter out only comfy deploy serverless
// Also multiply with the gpu selection
if (workflow_run.machine_type == "comfy-deploy-serverless") {
if (
(status === "success" || status === "failed") &&
workflow_run.user_id
) {
const sub = await getCurrentPlan({
user_id: workflow_run.user_id,
org_id: workflow_run.org_id,
});
if (sub && sub.subscription_item_api_id && workflow_run.ended_at) {
let durationInSec = Math.abs(
(workflow_run.ended_at.getTime() -
workflow_run.created_at.getTime()) /
1000,
);
durationInSec = Math.ceil(durationInSec);
switch (workflow_run.gpu) {
case "A100":
durationInSec *= 7;
break;
case "A10G":
durationInSec *= 4;
break;
}
await stripe.subscriptionItems.createUsageRecord(
sub.subscription_item_api_id,
{
quantity: durationInSec,
},
);
}
}
}
}
// const workflow_version = await db.query.workflowVersionTable.findFirst({
@@ -51,6 +94,6 @@ export async function POST(request: Request) {
},
{
status: 200,
}
},
);
}
+9
View File
@@ -0,0 +1,9 @@
"use client";
import { LoadingPageWrapper } from "@/components/LoadingWrapper";
import { usePathname } from "next/navigation";
export default function Loading() {
const pathName = usePathname();
return <LoadingPageWrapper className="h-full" tag={pathName.toLowerCase()} />;
}
+9
View File
@@ -0,0 +1,9 @@
import PricingList from "@/components/PricingPlan";
export default function Home() {
return (
<div>
<PricingList />
</div>
);
}
+8 -8
View File
@@ -1,19 +1,19 @@
// app/providers.tsx
'use client'
import posthog from 'posthog-js'
import { PostHogProvider } from 'posthog-js/react'
"use client";
import posthog from "posthog-js";
import { PostHogProvider } from "posthog-js/react";
if (typeof window !== 'undefined') {
if (typeof window !== "undefined") {
posthog.init(process.env.NEXT_PUBLIC_POSTHOG_KEY!, {
api_host: process.env.NEXT_PUBLIC_POSTHOG_HOST,
capture_pageview: false // Disable automatic pageview capture, as we capture manually
})
capture_pageview: false, // Disable automatic pageview capture, as we capture manually
});
}
export function PHProvider({
children,
}: {
children: React.ReactNode
children: React.ReactNode;
}) {
return <PostHogProvider client={posthog}>{children}</PostHogProvider>
return <PostHogProvider client={posthog}>{children}</PostHogProvider>;
}
+9
View File
@@ -0,0 +1,9 @@
"use client";
import { LoadingPageWrapper } from "@/components/LoadingWrapper";
import { usePathname } from "next/navigation";
export default function Loading() {
const pathName = usePathname();
return <LoadingPageWrapper className="h-full" tag={pathName.toLowerCase()} />;
}
+57
View File
@@ -0,0 +1,57 @@
import PricingList from "@/components/PricingPlan";
import { Badge } from "@/components/ui/badge";
import {
Card,
CardContent,
CardDescription,
CardHeader,
CardTitle,
} from "@/components/ui/card";
import { Progress } from "@/components/ui/progress";
import { getCurrentPlanWithAuth } from "@/server/getCurrentPlan";
import { stripe } from "@/server/stripe";
const freeTierSeconds = 30000;
export default async function Home() {
const sub = await getCurrentPlanWithAuth();
const data = sub?.subscription_item_api_id
? await stripe.subscriptionItems.listUsageRecordSummaries(
sub?.subscription_item_api_id,
)
: null;
return (
<div className="mt-4 flex items-center justify-center">
<Card className="p-4 w-full max-w-[600px]">
<CardHeader>
<CardTitle>Account Usage</CardTitle>
<CardDescription>View you account usage</CardDescription>
<Badge className="w-fit">{sub?.plan}</Badge>
</CardHeader>
{data && (
<CardContent className="text-sm flex flex-col gap-2">
<div className="flex justify-between gap-2">
<span>Current free gpu usage:</span>
{
<div className="flex gap-2">
<Badge>
{data.data[0].total_usage}s /{Math.floor(freeTierSeconds)}s
</Badge>
<Badge>
{Math.floor(data.data[0].total_usage / 60 / 60)}hr /
{Math.floor(freeTierSeconds / 60 / 60)}hr
</Badge>
</div>
}
</div>
<Progress
value={(data.data[0].total_usage / freeTierSeconds) * 100}
></Progress>
</CardContent>
)}
</Card>
</div>
);
}