refactor(db): migrate app pages and APIs from Prisma facade to Drizzle (5)

Co-authored-by: Cursor <[email protected]>
This commit is contained in:
SimoandCursor committed 2026-08-01 14:15:39 +02:00
1 parent 9aa4f331bf
commit 30b54e99e7
40 files changed
+1183 -784

No files matched your search

+14 -9
View File
@@ -1,10 +1,11 @@
import { and, eq, gt, or } from "drizzle-orm";
import { headers } from "next/headers";
import { redirect } from "next/navigation";
import { auth } from "@/lib/auth";
import { activeBanWhere } from "@/lib/bans";
import { unixNow } from "@/lib/bans";
import { Ban, db } from "@/lib/db";
import { safeRedirect } from "@/lib/foundation/security";
import { logger } from "@/lib/logger";
import { prisma } from "@/lib/prisma";
import { isIpBlacklisted, recordRequest } from "@/lib/services/abuse-guard";
import { siteSettings } from "@/lib/services/site-settings";
@@ -63,13 +64,17 @@ export async function enforceSiteAccess(): Promise<void> {
try {
if (!target && session?.user?.id) {
const ban = await prisma.ban.findFirst({
where: {
userId: Number(session.user.id),
...activeBanWhere(),
},
select: { id: true },
});
const now = unixNow();
const [ban] = await db
.select({ id: Ban.id })
.from(Ban)
.where(
and(
eq(Ban.userId, Number(session.user.id)),
or(eq(Ban.banExpire, 0), gt(Ban.banExpire, now)),
),
)
.limit(1);
if (ban) target = "/banned";
}
} catch {
+30 -24
View File
@@ -1,7 +1,8 @@
import { createHash, randomBytes } from "node:crypto";
import { and, eq, gt, isNull, or } from "drizzle-orm";
import { personalTokenScope } from "@/lib/auth/personal-token-scope";
import { databaseUserId } from "@/lib/auth/session-user";
import { prisma } from "@/lib/prisma";
import { db, PersonalAccessTokens } from "@/lib/db";
/**
* Bearer-token auth for the public REST API, backed by personal_access_tokens
@@ -24,21 +25,28 @@ export async function bearerUserId(req: Request): Promise<number | null> {
if (!raw) return null;
try {
const row = await prisma.personalAccessTokens.findFirst({
where: {
token: hashToken(raw),
OR: [{ expiresAt: null }, { expiresAt: { gt: new Date() } }],
},
select: { id: true, tokenableId: true },
});
const [row] = await db
.select({
id: PersonalAccessTokens.id,
tokenableId: PersonalAccessTokens.tokenableId,
})
.from(PersonalAccessTokens)
.where(
and(
eq(PersonalAccessTokens.token, hashToken(raw)),
or(
isNull(PersonalAccessTokens.expiresAt),
gt(PersonalAccessTokens.expiresAt, new Date()),
),
),
)
.limit(1);
if (!row) return null;
// Best-effort last-used stamp (don't fail the request if it errors).
prisma.personalAccessTokens
.update({
where: { id: row.id },
data: { lastUsedAt: new Date() },
select: { id: true },
})
void db
.update(PersonalAccessTokens)
.set({ lastUsedAt: new Date() })
.where(eq(PersonalAccessTokens.id, row.id))
.catch(() => {});
return databaseUserId(row.tokenableId);
} catch {
@@ -52,17 +60,15 @@ export async function issueToken(
name = "api",
): Promise<string | null> {
const plaintext = randomBytes(32).toString("hex");
const now = new Date();
try {
await prisma.personalAccessTokens.create({
data: {
...personalTokenScope(userId),
name: name.slice(0, 100),
token: hashToken(plaintext),
abilities: '["*"]',
createdAt: new Date(),
updatedAt: new Date(),
},
select: { id: true },
await db.insert(PersonalAccessTokens).values({
...personalTokenScope(userId),
name: name.slice(0, 100),
token: hashToken(plaintext),
abilities: '["*"]',
createdAt: now,
updatedAt: now,
});
return plaintext;
} catch {
+23 -22
View File
@@ -1,4 +1,4 @@
import { sql } from "drizzle-orm";
import { eq, sql } from "drizzle-orm";
import NextAuth from "next-auth";
import Credentials from "next-auth/providers/credentials";
import { env } from "@/env";
@@ -7,9 +7,8 @@ import { LaravelEncrypter } from "@/lib/auth/laravel-encrypter";
import { checkLogin } from "@/lib/auth/password";
import { verifyTotp } from "@/lib/auth/totp";
import { cachedQuery, invalidateKey } from "@/lib/cached-db";
import { db } from "@/lib/db";
import { db, User, WebsiteLoginLogs } from "@/lib/db";
import { logger } from "@/lib/logger";
import { prisma } from "@/lib/prisma";
import { clientIp, rateLimit } from "@/lib/rate-limit";
import { siteSettings } from "@/lib/services/site-settings";
@@ -88,10 +87,14 @@ export async function invalidateLoginCache(username: string): Promise<void> {
}
async function verify2faCode(userId: number, code: string): Promise<boolean> {
const user = await prisma.user.findUnique({
where: { id: userId },
select: { twoFactorSecret: true, twoFactorRecoveryCodes: true },
});
const [user] = await db
.select({
twoFactorSecret: User.twoFactorSecret,
twoFactorRecoveryCodes: User.twoFactorRecoveryCodes,
})
.from(User)
.where(eq(User.id, userId))
.limit(1);
if (!user?.twoFactorSecret) return false;
// Try TOTP first
@@ -119,10 +122,10 @@ async function verify2faCode(userId: number, code: string): Promise<boolean> {
if (idx !== -1) {
codes.splice(idx, 1);
const remaining = codes.length > 0 ? JSON.stringify(codes) : null;
await prisma.user.update({
where: { id: userId },
data: { twoFactorRecoveryCodes: remaining },
});
await db
.update(User)
.set({ twoFactorRecoveryCodes: remaining })
.where(eq(User.id, userId));
return true;
}
}
@@ -189,10 +192,10 @@ export const { handlers, signOut, auth } = NextAuth({
}
if (res.upgradedHash) {
await prisma.user.update({
where: { id: user.id },
data: { password: res.upgradedHash },
});
await db
.update(User)
.set({ password: res.upgradedHash })
.where(eq(User.id, user.id));
invalidateLoginCache(username);
}
@@ -213,13 +216,11 @@ export const { handlers, signOut, auth } = NextAuth({
try {
const { headers } = await import("next/headers");
const ua = (await headers()).get("user-agent")?.slice(0, 512) ?? null;
await prisma.websiteLoginLogs.create({
data: {
userId: user.id,
ip,
userAgent: ua,
createdAt: new Date(),
},
await db.insert(WebsiteLoginLogs).values({
userId: user.id,
ip,
userAgent: ua,
createdAt: new Date(),
});
} catch {
logger.warn("Failed to record login log for user", {
+15 -10
View File
@@ -1,14 +1,16 @@
import { beforeEach, describe, expect, it, vi } from "vitest";
const mockFindUnique = vi.hoisted(() => vi.fn());
const mockLimit = vi.hoisted(() => vi.fn());
const mockWhere = vi.hoisted(() => vi.fn(() => ({ limit: mockLimit })));
const mockFrom = vi.hoisted(() => vi.fn(() => ({ where: mockWhere })));
const mockSelect = vi.hoisted(() => vi.fn(() => ({ from: mockFrom })));
const mockRedisGet = vi.hoisted(() => vi.fn());
const mockRedisSetex = vi.hoisted(() => vi.fn());
const mockRedisDel = vi.hoisted(() => vi.fn());
vi.mock("@/lib/prisma", () => ({
prisma: {
user: { findUnique: mockFindUnique },
},
vi.mock("@/lib/db", () => ({
db: { select: mockSelect },
User: { id: "id", websiteJwtVersion: "websiteJwtVersion" },
}));
vi.mock("@/lib/redis", () => ({
@@ -23,29 +25,32 @@ describe("jwt-version-cache", () => {
beforeEach(() => {
vi.resetModules();
vi.clearAllMocks();
mockWhere.mockReturnValue({ limit: mockLimit });
mockFrom.mockReturnValue({ where: mockWhere });
mockSelect.mockReturnValue({ from: mockFrom });
mockRedisGet.mockResolvedValue(null);
mockRedisSetex.mockResolvedValue("OK");
mockRedisDel.mockResolvedValue(1);
});
it("returns DB version and caches it", async () => {
mockFindUnique.mockResolvedValue({ websiteJwtVersion: 3 });
mockLimit.mockResolvedValue([{ websiteJwtVersion: 3 }]);
const { getCachedJwtVersion } = await import("./jwt-version-cache");
await expect(getCachedJwtVersion(42)).resolves.toBe(3);
expect(mockFindUnique).toHaveBeenCalledTimes(1);
expect(mockSelect).toHaveBeenCalledTimes(1);
await expect(getCachedJwtVersion(42)).resolves.toBe(3);
expect(mockFindUnique).toHaveBeenCalledTimes(1);
expect(mockSelect).toHaveBeenCalledTimes(1);
});
it("invalidates memory and redis entries", async () => {
mockFindUnique.mockResolvedValue({ websiteJwtVersion: 1 });
mockLimit.mockResolvedValue([{ websiteJwtVersion: 1 }]);
const { getCachedJwtVersion, invalidateJwtVersionCache } = await import(
"./jwt-version-cache"
);
await getCachedJwtVersion(7);
await invalidateJwtVersionCache(7);
expect(mockRedisDel).toHaveBeenCalled();
mockFindUnique.mockResolvedValue({ websiteJwtVersion: 2 });
mockLimit.mockResolvedValue([{ websiteJwtVersion: 2 }]);
await expect(getCachedJwtVersion(7)).resolves.toBe(2);
});
});
+7 -5
View File
@@ -1,6 +1,7 @@
import "server-only";
import { prisma } from "@/lib/prisma";
import { eq } from "drizzle-orm";
import { db, User } from "@/lib/db";
import { redis } from "@/lib/redis";
const MEMORY_TTL_MS = 60_000;
@@ -40,10 +41,11 @@ export async function getCachedJwtVersion(
}
try {
const row = await prisma.user.findUnique({
where: { id: userId },
select: { websiteJwtVersion: true },
});
const [row] = await db
.select({ websiteJwtVersion: User.websiteJwtVersion })
.from(User)
.where(eq(User.id, userId))
.limit(1);
if (!row) return null;
const version = row.websiteJwtVersion;
memory.set(userId, { version, expiresAt: now + MEMORY_TTL_MS });
+25 -8
View File
@@ -1,9 +1,18 @@
import { describe, expect, it, vi } from "vitest";
import { generateSsoTicket, issueSsoTicket } from "./sso-ticket";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { generateSsoTicket } from "./sso-ticket";
const UUID_RE =
/^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$/;
const mockWhere = vi.hoisted(() => vi.fn().mockResolvedValue(undefined));
const mockSet = vi.hoisted(() => vi.fn(() => ({ where: mockWhere })));
const mockUpdate = vi.hoisted(() => vi.fn(() => ({ set: mockSet })));
vi.mock("@/lib/db", () => ({
db: { update: mockUpdate },
User: { id: "id" },
}));
describe("generateSsoTicket", () => {
it("uses '{hotelName-without-spaces}-{uuidv4}'", () => {
const t = generateSsoTicket("Atom Hotel");
@@ -23,15 +32,23 @@ describe("generateSsoTicket", () => {
});
describe("issueSsoTicket", () => {
beforeEach(() => {
vi.clearAllMocks();
mockSet.mockReturnValue({ where: mockWhere });
mockUpdate.mockReturnValue({ set: mockSet });
mockWhere.mockResolvedValue(undefined);
});
it("writes auth_ticket AND ip_current and returns the ticket", async () => {
const update = vi.fn().mockResolvedValue(undefined);
const db = { user: { update } };
const ticket = await issueSsoTicket(db, 42, "Atom Hotel", "1.2.3.4");
const { issueSsoTicket } = await import("./sso-ticket");
const ticket = await issueSsoTicket(42, "Atom Hotel", "1.2.3.4");
expect(ticket.startsWith("AtomHotel-")).toBe(true);
expect(update).toHaveBeenCalledWith({
where: { id: 42 },
data: { authTicket: ticket, ipCurrent: "1.2.3.4" },
expect(mockUpdate).toHaveBeenCalled();
expect(mockSet).toHaveBeenCalledWith({
authTicket: ticket,
ipCurrent: "1.2.3.4",
});
expect(mockWhere).toHaveBeenCalled();
});
});
+6 -15
View File
@@ -1,4 +1,6 @@
import { randomUUID } from "node:crypto";
import { eq } from "drizzle-orm";
import { db, User } from "@/lib/db";
/**
* Build the SSO ticket exactly like AtomCMS's User::ssoTicket():
@@ -12,30 +14,19 @@ export function generateSsoTicket(hotelName: string): string {
return `${normalized}-${randomUUID()}`;
}
/** Minimal shape of the Prisma client this needs (keeps it unit-testable). */
export interface SsoUserUpdater {
user: {
update(args: {
where: { id: number };
data: { authTicket: string; ipCurrent: string };
}): Promise<unknown>;
};
}
/**
* Generate a ticket and persist it like AtomCMS: writes auth_ticket AND
* ip_current on the user, then returns the ticket for the client launcher.
*/
export async function issueSsoTicket(
db: SsoUserUpdater,
userId: number,
hotelName: string,
ip: string,
): Promise<string> {
const ticket = generateSsoTicket(hotelName);
await db.user.update({
where: { id: userId },
data: { authTicket: ticket, ipCurrent: ip },
});
await db
.update(User)
.set({ authTicket: ticket, ipCurrent: ip })
.where(eq(User.id, userId));
return ticket;
}