refactor(db): migrate app pages and APIs from Prisma facade to Drizzle (5)
Co-authored-by: Cursor <[email protected]>
This commit is contained in:
1 parent
9aa4f331bf
commit
30b54e99e7
40 files changed
+1183
-784
No files matched your search
+14
-9
@@ -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
@@ -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
@@ -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", {
|
||||
|
||||
@@ -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);
|
||||
});
|
||||
});
|
||||
@@ -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 });
|
||||
|
||||
@@ -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();
|
||||
});
|
||||
});
|
||||
@@ -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;
|
||||
}
|
||||
Reference in new issue
Block a user