test: verify duplicate-safe Discord reconciliation
This commit is contained in:
@@ -0,0 +1,131 @@
|
||||
import { describe, expect, test } from "bun:test";
|
||||
import {
|
||||
reconcileDiscordUser,
|
||||
type DiscordIdentity,
|
||||
type DiscordUserOperations,
|
||||
} from "@/lib/auth/discord-user-core";
|
||||
|
||||
type MemoryUser = {
|
||||
id: string;
|
||||
discordId: string;
|
||||
name: string;
|
||||
image: string | null;
|
||||
email: string | null;
|
||||
};
|
||||
|
||||
function createMemoryStore() {
|
||||
const users = new Map<string, MemoryUser>();
|
||||
const userIdByDiscord = new Map<string, string>();
|
||||
const accounts = new Map<string, { userId: string; fields: Record<string, unknown> }>();
|
||||
let nextUserId = 1;
|
||||
|
||||
const operations: DiscordUserOperations<MemoryUser> = {
|
||||
async findLinkedUserId(discordId) {
|
||||
return accounts.get(discordId)?.userId ?? null;
|
||||
},
|
||||
async upsertUser(identity) {
|
||||
let userId = userIdByDiscord.get(identity.discordId);
|
||||
if (!userId) {
|
||||
userId = `user-${nextUserId++}`;
|
||||
userIdByDiscord.set(identity.discordId, userId);
|
||||
users.set(userId, {
|
||||
id: userId,
|
||||
discordId: identity.discordId,
|
||||
name: identity.displayName,
|
||||
image: identity.avatarUrl,
|
||||
email: identity.email ?? null,
|
||||
});
|
||||
}
|
||||
return userId;
|
||||
},
|
||||
async insertAccount(userId, discordId, account) {
|
||||
if (!accounts.has(discordId)) {
|
||||
accounts.set(discordId, {
|
||||
userId,
|
||||
fields: { type: account?.type ?? "oauth" },
|
||||
});
|
||||
}
|
||||
},
|
||||
async updateUser(userId, identity) {
|
||||
const user = users.get(userId);
|
||||
if (!user) return null;
|
||||
const updated = {
|
||||
...user,
|
||||
discordId: identity.discordId,
|
||||
name: identity.displayName,
|
||||
image: identity.avatarUrl,
|
||||
...(identity.email !== undefined ? { email: identity.email } : {}),
|
||||
};
|
||||
users.set(userId, updated);
|
||||
return updated;
|
||||
},
|
||||
async updateAccount(discordId, update) {
|
||||
const account = accounts.get(discordId);
|
||||
if (account) account.fields = { ...account.fields, ...update };
|
||||
},
|
||||
};
|
||||
|
||||
return { users, accounts, operations };
|
||||
}
|
||||
|
||||
const botIdentity: DiscordIdentity = {
|
||||
discordId: "123456789012345678",
|
||||
displayName: "Bot-created user",
|
||||
avatarUrl: "https://cdn.discordapp.com/avatar.png",
|
||||
};
|
||||
|
||||
describe("Discord account reconciliation", () => {
|
||||
test("creates one user and one account for a new bot sender", async () => {
|
||||
const store = createMemoryStore();
|
||||
const user = await reconcileDiscordUser(botIdentity, store.operations);
|
||||
|
||||
expect(user.id).toBe("user-1");
|
||||
expect(store.users.size).toBe(1);
|
||||
expect(store.accounts.size).toBe(1);
|
||||
expect(store.accounts.get(botIdentity.discordId)?.userId).toBe(user.id);
|
||||
});
|
||||
|
||||
test("concurrent first uploads resolve to the same database winner", async () => {
|
||||
const store = createMemoryStore();
|
||||
const [first, second] = await Promise.all([
|
||||
reconcileDiscordUser(botIdentity, store.operations),
|
||||
reconcileDiscordUser(botIdentity, store.operations),
|
||||
]);
|
||||
|
||||
expect(first.id).toBe(second.id);
|
||||
expect(store.users.size).toBe(1);
|
||||
expect(store.accounts.size).toBe(1);
|
||||
});
|
||||
|
||||
test("OAuth reuses a bot-created user and bot refreshes preserve email", async () => {
|
||||
const store = createMemoryStore();
|
||||
const botUser = await reconcileDiscordUser(botIdentity, store.operations);
|
||||
const oauthUser = await reconcileDiscordUser(
|
||||
{
|
||||
...botIdentity,
|
||||
displayName: "OAuth name",
|
||||
email: "[email protected]",
|
||||
account: {
|
||||
type: "oauth",
|
||||
access_token: "access-token",
|
||||
refresh_token: "refresh-token",
|
||||
},
|
||||
},
|
||||
store.operations,
|
||||
);
|
||||
const refreshedByBot = await reconcileDiscordUser(
|
||||
{ ...botIdentity, displayName: "Latest Discord name" },
|
||||
store.operations,
|
||||
);
|
||||
|
||||
expect(oauthUser.id).toBe(botUser.id);
|
||||
expect(refreshedByBot.id).toBe(botUser.id);
|
||||
expect(refreshedByBot.email).toBe("[email protected]");
|
||||
expect(store.users.size).toBe(1);
|
||||
expect(store.accounts.size).toBe(1);
|
||||
expect(store.accounts.get(botIdentity.discordId)?.fields).toMatchObject({
|
||||
access_token: "access-token",
|
||||
refresh_token: "refresh-token",
|
||||
});
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,78 @@
|
||||
import type { AdapterAccount } from "next-auth/adapters";
|
||||
|
||||
export type DiscordIdentity = {
|
||||
discordId: string;
|
||||
displayName: string;
|
||||
avatarUrl: string | null;
|
||||
email?: string | null;
|
||||
account?: Partial<AdapterAccount>;
|
||||
};
|
||||
|
||||
export function discordAccountUpdate(
|
||||
account: Partial<AdapterAccount> | undefined,
|
||||
) {
|
||||
if (!account) return {};
|
||||
return {
|
||||
...(account.type !== undefined ? { type: account.type } : {}),
|
||||
...(account.refresh_token !== undefined
|
||||
? { refresh_token: account.refresh_token }
|
||||
: {}),
|
||||
...(account.access_token !== undefined
|
||||
? { access_token: account.access_token }
|
||||
: {}),
|
||||
...(account.expires_at !== undefined
|
||||
? { expires_at: account.expires_at }
|
||||
: {}),
|
||||
...(account.token_type !== undefined
|
||||
? { token_type: account.token_type }
|
||||
: {}),
|
||||
...(account.scope !== undefined ? { scope: account.scope } : {}),
|
||||
...(account.id_token !== undefined ? { id_token: account.id_token } : {}),
|
||||
...(account.session_state !== undefined
|
||||
? { session_state: account.session_state as string | null }
|
||||
: {}),
|
||||
};
|
||||
}
|
||||
|
||||
export type DiscordAccountUpdate = ReturnType<typeof discordAccountUpdate>;
|
||||
|
||||
export type DiscordUserOperations<TUser> = {
|
||||
findLinkedUserId(discordId: string): Promise<string | null>;
|
||||
upsertUser(identity: DiscordIdentity): Promise<string>;
|
||||
insertAccount(
|
||||
userId: string,
|
||||
discordId: string,
|
||||
account: Partial<AdapterAccount> | undefined,
|
||||
): Promise<void>;
|
||||
updateUser(userId: string, identity: DiscordIdentity): Promise<TUser | null>;
|
||||
updateAccount(
|
||||
discordId: string,
|
||||
update: DiscordAccountUpdate,
|
||||
): Promise<void>;
|
||||
};
|
||||
|
||||
export async function reconcileDiscordUser<TUser>(
|
||||
identity: DiscordIdentity,
|
||||
operations: DiscordUserOperations<TUser>,
|
||||
) {
|
||||
const { discordId, account } = identity;
|
||||
let userId = await operations.findLinkedUserId(discordId);
|
||||
|
||||
if (!userId) {
|
||||
userId = await operations.upsertUser(identity);
|
||||
await operations.insertAccount(userId, discordId, account);
|
||||
|
||||
const winner = await operations.findLinkedUserId(discordId);
|
||||
if (!winner) throw new Error("Failed to link Discord account");
|
||||
userId = winner;
|
||||
}
|
||||
|
||||
const user = await operations.updateUser(userId, identity);
|
||||
const nextAccount = discordAccountUpdate(account);
|
||||
if (Object.keys(nextAccount).length > 0) {
|
||||
await operations.updateAccount(discordId, nextAccount);
|
||||
}
|
||||
|
||||
if (!user) throw new Error("Failed to reconcile Discord user");
|
||||
return user;
|
||||
}
|
||||
+73
-111
@@ -1,132 +1,94 @@
|
||||
import "server-only";
|
||||
|
||||
import type { AdapterAccount } from "next-auth/adapters";
|
||||
import { and, eq } from "drizzle-orm";
|
||||
import { db } from "@/db";
|
||||
import { accounts, users } from "@/db/schema";
|
||||
import {
|
||||
discordAccountUpdate,
|
||||
reconcileDiscordUser,
|
||||
type DiscordIdentity,
|
||||
} from "@/lib/auth/discord-user-core";
|
||||
|
||||
export type DiscordIdentity = {
|
||||
discordId: string;
|
||||
displayName: string;
|
||||
avatarUrl: string | null;
|
||||
email?: string | null;
|
||||
account?: Partial<AdapterAccount>;
|
||||
};
|
||||
|
||||
function accountUpdate(account: Partial<AdapterAccount> | undefined) {
|
||||
if (!account) return {};
|
||||
return {
|
||||
...(account.type !== undefined ? { type: account.type } : {}),
|
||||
...(account.refresh_token !== undefined
|
||||
? { refresh_token: account.refresh_token }
|
||||
: {}),
|
||||
...(account.access_token !== undefined
|
||||
? { access_token: account.access_token }
|
||||
: {}),
|
||||
...(account.expires_at !== undefined
|
||||
? { expires_at: account.expires_at }
|
||||
: {}),
|
||||
...(account.token_type !== undefined
|
||||
? { token_type: account.token_type }
|
||||
: {}),
|
||||
...(account.scope !== undefined ? { scope: account.scope } : {}),
|
||||
...(account.id_token !== undefined ? { id_token: account.id_token } : {}),
|
||||
...(account.session_state !== undefined
|
||||
? { session_state: account.session_state as string | null }
|
||||
: {}),
|
||||
};
|
||||
}
|
||||
export type { DiscordIdentity } from "@/lib/auth/discord-user-core";
|
||||
|
||||
export async function ensureDiscordUser(
|
||||
identity: DiscordIdentity
|
||||
): Promise<typeof users.$inferSelect> {
|
||||
const { discordId, displayName, avatarUrl, email, account } = identity;
|
||||
|
||||
return db.transaction(async (tx) => {
|
||||
const [linked] = await tx
|
||||
.select({ userId: accounts.userId })
|
||||
.from(accounts)
|
||||
.where(
|
||||
and(
|
||||
eq(accounts.provider, "discord"),
|
||||
eq(accounts.providerAccountId, discordId)
|
||||
)
|
||||
)
|
||||
.limit(1);
|
||||
|
||||
let userId = linked?.userId;
|
||||
if (!userId) {
|
||||
const [user] = await tx
|
||||
.insert(users)
|
||||
.values({
|
||||
discordId,
|
||||
name: displayName,
|
||||
image: avatarUrl,
|
||||
...(email !== undefined ? { email } : {}),
|
||||
})
|
||||
.onConflictDoUpdate({
|
||||
target: users.discordId,
|
||||
set: {
|
||||
name: displayName,
|
||||
image: avatarUrl,
|
||||
...(email !== undefined ? { email } : {}),
|
||||
},
|
||||
})
|
||||
.returning({ id: users.id });
|
||||
userId = user.id;
|
||||
|
||||
await tx
|
||||
.insert(accounts)
|
||||
.values({
|
||||
userId,
|
||||
provider: "discord",
|
||||
providerAccountId: discordId,
|
||||
type: account?.type ?? "oauth",
|
||||
...accountUpdate(account),
|
||||
})
|
||||
.onConflictDoNothing({
|
||||
target: [accounts.provider, accounts.providerAccountId],
|
||||
});
|
||||
|
||||
const [winner] = await tx
|
||||
const findLinkedUserId = async (discordId: string) => {
|
||||
const [linked] = await tx
|
||||
.select({ userId: accounts.userId })
|
||||
.from(accounts)
|
||||
.where(
|
||||
and(
|
||||
eq(accounts.provider, "discord"),
|
||||
eq(accounts.providerAccountId, discordId)
|
||||
)
|
||||
eq(accounts.providerAccountId, discordId),
|
||||
),
|
||||
)
|
||||
.limit(1);
|
||||
if (!winner) throw new Error("Failed to link Discord account");
|
||||
userId = winner.userId;
|
||||
}
|
||||
return linked?.userId ?? null;
|
||||
};
|
||||
|
||||
const [user] = await tx
|
||||
.update(users)
|
||||
.set({
|
||||
discordId,
|
||||
name: displayName,
|
||||
image: avatarUrl,
|
||||
...(email !== undefined ? { email } : {}),
|
||||
})
|
||||
.where(eq(users.id, userId))
|
||||
.returning();
|
||||
|
||||
const nextAccount = accountUpdate(account);
|
||||
if (Object.keys(nextAccount).length > 0) {
|
||||
await tx
|
||||
.update(accounts)
|
||||
.set(nextAccount)
|
||||
.where(
|
||||
and(
|
||||
eq(accounts.provider, "discord"),
|
||||
eq(accounts.providerAccountId, discordId)
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
if (!user) throw new Error("Failed to reconcile Discord user");
|
||||
return user;
|
||||
return reconcileDiscordUser(identity, {
|
||||
findLinkedUserId,
|
||||
async upsertUser({ discordId, displayName, avatarUrl, email }) {
|
||||
const [user] = await tx
|
||||
.insert(users)
|
||||
.values({
|
||||
discordId,
|
||||
name: displayName,
|
||||
image: avatarUrl,
|
||||
...(email !== undefined ? { email } : {}),
|
||||
})
|
||||
.onConflictDoUpdate({
|
||||
target: users.discordId,
|
||||
set: {
|
||||
name: displayName,
|
||||
image: avatarUrl,
|
||||
...(email !== undefined ? { email } : {}),
|
||||
},
|
||||
})
|
||||
.returning({ id: users.id });
|
||||
return user.id;
|
||||
},
|
||||
async insertAccount(userId, discordId, account) {
|
||||
await tx
|
||||
.insert(accounts)
|
||||
.values({
|
||||
userId,
|
||||
provider: "discord",
|
||||
providerAccountId: discordId,
|
||||
type: account?.type ?? "oauth",
|
||||
...discordAccountUpdate(account),
|
||||
})
|
||||
.onConflictDoNothing({
|
||||
target: [accounts.provider, accounts.providerAccountId],
|
||||
});
|
||||
},
|
||||
async updateUser(userId, { discordId, displayName, avatarUrl, email }) {
|
||||
const [user] = await tx
|
||||
.update(users)
|
||||
.set({
|
||||
discordId,
|
||||
name: displayName,
|
||||
image: avatarUrl,
|
||||
...(email !== undefined ? { email } : {}),
|
||||
})
|
||||
.where(eq(users.id, userId))
|
||||
.returning();
|
||||
return user ?? null;
|
||||
},
|
||||
async updateAccount(discordId, update) {
|
||||
await tx
|
||||
.update(accounts)
|
||||
.set(update)
|
||||
.where(
|
||||
and(
|
||||
eq(accounts.provider, "discord"),
|
||||
eq(accounts.providerAccountId, discordId),
|
||||
),
|
||||
);
|
||||
},
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user