refactor(server): use user model on oauth plugin (#10031)

close CLOUD-117
This commit is contained in:
fengmk2
2025-02-10 12:01:14 +00:00
parent 23364b59a0
commit 67b6c28d67
5 changed files with 48 additions and 64 deletions

View File

@@ -187,7 +187,7 @@ test('should revoke token after change user identify', async t => {
// change password // change password
{ {
const u3Email = 'u3@affine.pro'; const u3Email = 'u3333@affine.pro';
await app.logout(); await app.logout();
const u3 = await app.signup(u3Email); const u3 = await app.signup(u3Email);

View File

@@ -1,4 +1,3 @@
import { PrismaClient } from '@prisma/client';
import ava, { TestFn } from 'ava'; import ava, { TestFn } from 'ava';
import Sinon from 'sinon'; import Sinon from 'sinon';
@@ -276,16 +275,13 @@ test('should trigger user.deleted event', async t => {
}); });
test('should paginate users', async t => { test('should paginate users', async t => {
const db = t.context.module.get(PrismaClient);
const now = Date.now(); const now = Date.now();
await Promise.all( await Promise.all(
Array.from({ length: 100 }).map((_, i) => Array.from({ length: 100 }).map((_, i) =>
db.user.create({ t.context.user.create({
data: { name: `test-paginate-${i}`,
name: `test${i}`, email: `test-paginate-${i}@affine.pro`,
email: `test${i}@affine.pro`, createdAt: new Date(now + i),
createdAt: new Date(now + i),
},
}) })
) )
); );
@@ -294,7 +290,7 @@ test('should paginate users', async t => {
t.is(users.length, 10); t.is(users.length, 10);
t.deepEqual( t.deepEqual(
users.map(user => user.email), users.map(user => user.email),
Array.from({ length: 10 }).map((_, i) => `test${i}@affine.pro`) Array.from({ length: 10 }).map((_, i) => `test-paginate-${i}@affine.pro`)
); );
}); });

View File

@@ -308,7 +308,7 @@ test('should not throw if account registered', async t => {
}); });
test('should be able to fullfil user with oauth sign in', async t => { test('should be able to fullfil user with oauth sign in', async t => {
const { app, db } = t.context; const { app, models } = t.context;
const u3 = await app.createUser('u3@affine.pro'); const u3 = await app.createUser('u3@affine.pro');
@@ -321,11 +321,11 @@ test('should be able to fullfil user with oauth sign in', async t => {
t.truthy(sessionUser); t.truthy(sessionUser);
t.is(sessionUser!.email, u3.email); t.is(sessionUser!.email, u3.email);
const account = await db.connectedAccount.findFirst({ const account = await models.user.getConnectedAccount(
where: { OAuthProviderName.Google,
userId: u3.id, '1'
}, );
});
t.truthy(account); t.truthy(account);
t.is(account!.user.id, u3.id);
}); });

View File

@@ -241,9 +241,13 @@ export class UserModel extends BaseModel {
// #region ConnectedAccount // #region ConnectedAccount
async createConnectedAccount(data: CreateConnectedAccountInput) { async createConnectedAccount(data: CreateConnectedAccountInput) {
return await this.db.connectedAccount.create({ const account = await this.db.connectedAccount.create({
data, data,
}); });
this.logger.log(
`Connected account ${account.provider}:${account.id} created`
);
return account;
} }
async getConnectedAccount(provider: string, providerAccountId: string) { async getConnectedAccount(provider: string, providerAccountId: string) {

View File

@@ -7,7 +7,7 @@ import {
Req, Req,
Res, Res,
} from '@nestjs/common'; } from '@nestjs/common';
import { ConnectedAccount, PrismaClient } from '@prisma/client'; import { ConnectedAccount } from '@prisma/client';
import type { Request, Response } from 'express'; import type { Request, Response } from 'express';
import { import {
@@ -30,8 +30,7 @@ export class OAuthController {
private readonly auth: AuthService, private readonly auth: AuthService,
private readonly oauth: OAuthService, private readonly oauth: OAuthService,
private readonly models: Models, private readonly models: Models,
private readonly providerFactory: OAuthProviderFactory, private readonly providerFactory: OAuthProviderFactory
private readonly db: PrismaClient
) {} ) {}
@Public() @Public()
@@ -120,48 +119,39 @@ export class OAuthController {
externalAccount: OAuthAccount, externalAccount: OAuthAccount,
tokens: Tokens tokens: Tokens
) { ) {
const connectedUser = await this.db.connectedAccount.findFirst({ const connectedAccount = await this.models.user.getConnectedAccount(
where: { provider,
provider, externalAccount.id
providerAccountId: externalAccount.id, );
},
include: {
user: true,
},
});
if (connectedUser) { if (connectedAccount) {
// already connected // already connected
await this.updateConnectedAccount(connectedUser, tokens); await this.updateConnectedAccount(connectedAccount, tokens);
return connectedAccount.user;
return connectedUser.user;
} }
const user = await this.models.user.fulfill(externalAccount.email, { const user = await this.models.user.fulfill(externalAccount.email, {
avatarUrl: externalAccount.avatarUrl, avatarUrl: externalAccount.avatarUrl,
}); });
await this.db.connectedAccount.create({ await this.models.user.createConnectedAccount({
data: { userId: user.id,
userId: user.id, provider,
provider, providerAccountId: externalAccount.id,
providerAccountId: externalAccount.id, ...tokens,
...tokens,
},
}); });
return user; return user;
} }
private async updateConnectedAccount( private async updateConnectedAccount(
connectedUser: ConnectedAccount, connectedAccount: ConnectedAccount,
tokens: Tokens tokens: Tokens
) { ) {
return this.db.connectedAccount.update({ return await this.models.user.updateConnectedAccount(
where: { connectedAccount.id,
id: connectedUser.id, tokens
}, );
data: tokens,
});
} }
/** /**
@@ -175,26 +165,20 @@ export class OAuthController {
externalAccount: OAuthAccount, externalAccount: OAuthAccount,
tokens: Tokens tokens: Tokens
) { ) {
const connectedUser = await this.db.connectedAccount.findFirst({ const connectedAccount = await this.models.user.getConnectedAccount(
where: { provider,
provider, externalAccount.id
providerAccountId: externalAccount.id, );
}, if (connectedAccount) {
}); if (connectedAccount.userId !== user.id) {
if (connectedUser) {
if (connectedUser.id !== user.id) {
throw new OauthAccountAlreadyConnected(); throw new OauthAccountAlreadyConnected();
} }
} else { } else {
await this.db.connectedAccount.create({ await this.models.user.createConnectedAccount({
data: { userId: user.id,
userId: user.id, provider,
provider, providerAccountId: externalAccount.id,
providerAccountId: externalAccount.id, ...tokens,
accessToken: tokens.accessToken,
refreshToken: tokens.refreshToken,
},
}); });
} }
} }