diff --git a/packages/api-gateway/src/Bootstrap/Container.ts b/packages/api-gateway/src/Bootstrap/Container.ts index 19849acbc..07660dbf9 100644 --- a/packages/api-gateway/src/Bootstrap/Container.ts +++ b/packages/api-gateway/src/Bootstrap/Container.ts @@ -23,6 +23,7 @@ import { SubscriptionTokenAuthMiddleware } from '../Controller/SubscriptionToken import { StatisticsMiddleware } from '../Controller/StatisticsMiddleware' import { CrossServiceTokenCacheInterface } from '../Service/Cache/CrossServiceTokenCacheInterface' import { RedisCrossServiceTokenCache } from '../Infra/Redis/RedisCrossServiceTokenCache' +import { WebSocketAuthMiddleware } from '../Controller/WebSocketAuthMiddleware' // eslint-disable-next-line @typescript-eslint/no-var-requires const newrelicFormatter = require('@newrelic/winston-enricher') @@ -85,6 +86,7 @@ export class ContainerConfigLoader { // Middleware container.bind(TYPES.AuthMiddleware).to(AuthMiddleware) + container.bind(TYPES.WebSocketAuthMiddleware).to(WebSocketAuthMiddleware) container .bind(TYPES.SubscriptionTokenAuthMiddleware) .to(SubscriptionTokenAuthMiddleware) diff --git a/packages/api-gateway/src/Bootstrap/Types.ts b/packages/api-gateway/src/Bootstrap/Types.ts index 448bbabf1..29ee0c638 100644 --- a/packages/api-gateway/src/Bootstrap/Types.ts +++ b/packages/api-gateway/src/Bootstrap/Types.ts @@ -18,6 +18,7 @@ const TYPES = { // Middleware StatisticsMiddleware: Symbol.for('StatisticsMiddleware'), AuthMiddleware: Symbol.for('AuthMiddleware'), + WebSocketAuthMiddleware: Symbol.for('WebSocketAuthMiddleware'), SubscriptionTokenAuthMiddleware: Symbol.for('SubscriptionTokenAuthMiddleware'), // Services HTTPService: Symbol.for('HTTPService'), diff --git a/packages/api-gateway/src/Controller/WebSocketAuthMiddleware.ts b/packages/api-gateway/src/Controller/WebSocketAuthMiddleware.ts new file mode 100644 index 000000000..92fba5164 --- /dev/null +++ b/packages/api-gateway/src/Controller/WebSocketAuthMiddleware.ts @@ -0,0 +1,95 @@ +import { CrossServiceTokenData } from '@standardnotes/security' +import { RoleName } from '@standardnotes/common' +import { NextFunction, Request, Response } from 'express' +import { inject, injectable } from 'inversify' +import { BaseMiddleware } from 'inversify-express-utils' +import { verify } from 'jsonwebtoken' +import { AxiosError, AxiosInstance } from 'axios' +import { Logger } from 'winston' + +import TYPES from '../Bootstrap/Types' + +@injectable() +export class WebSocketAuthMiddleware extends BaseMiddleware { + constructor( + @inject(TYPES.HTTPClient) private httpClient: AxiosInstance, + @inject(TYPES.AUTH_SERVER_URL) private authServerUrl: string, + @inject(TYPES.AUTH_JWT_SECRET) private jwtSecret: string, + @inject(TYPES.Logger) private logger: Logger, + ) { + super() + } + + async handler(request: Request, response: Response, next: NextFunction): Promise { + const authHeaderValue = request.headers.authorization as string + + if (!authHeaderValue) { + response.status(401).send({ + error: { + tag: 'invalid-auth', + message: 'Invalid login credentials.', + }, + }) + + return + } + + try { + const authResponse = await this.httpClient.request({ + method: 'POST', + headers: { + Authorization: authHeaderValue, + Accept: 'application/json', + }, + validateStatus: (status: number) => { + return status >= 200 && status < 500 + }, + url: `${this.authServerUrl}/sockets/tokens/validate`, + }) + + if (authResponse.status > 200) { + response.setHeader('content-type', authResponse.headers['content-type']) + response.status(authResponse.status).send(authResponse.data) + + return + } + + const crossServiceToken = authResponse.data.authToken + + response.locals.authToken = crossServiceToken + + const decodedToken = verify(crossServiceToken, this.jwtSecret, { algorithms: ['HS256'] }) + + response.locals.freeUser = + decodedToken.roles.length === 1 && + decodedToken.roles.find((role) => role.name === RoleName.CoreUser) !== undefined + response.locals.userUuid = decodedToken.user.uuid + response.locals.roles = decodedToken.roles + } catch (error) { + const errorMessage = (error as AxiosError).isAxiosError + ? JSON.stringify((error as AxiosError).response?.data) + : (error as Error).message + + this.logger.error( + `Could not pass the request to ${this.authServerUrl}/sockets/tokens/validate on underlying service: ${errorMessage}`, + ) + + this.logger.debug('Response error: %O', (error as AxiosError).response ?? error) + + if ((error as AxiosError).response?.headers['content-type']) { + response.setHeader('content-type', (error as AxiosError).response?.headers['content-type'] as string) + } + + const errorCode = + (error as AxiosError).isAxiosError && !isNaN(+((error as AxiosError).code as string)) + ? +((error as AxiosError).code as string) + : 500 + + response.status(errorCode).send(errorMessage) + + return + } + + return next() + } +} diff --git a/packages/api-gateway/src/Controller/v1/WebSocketsController.ts b/packages/api-gateway/src/Controller/v1/WebSocketsController.ts index 02791191e..ad8c844c9 100644 --- a/packages/api-gateway/src/Controller/v1/WebSocketsController.ts +++ b/packages/api-gateway/src/Controller/v1/WebSocketsController.ts @@ -20,7 +20,7 @@ export class WebSocketsController extends BaseHttpController { await this.httpService.callAuthServer(request, response, 'sockets/tokens', request.body) } - @httpPost('/', TYPES.AuthMiddleware) + @httpPost('/', TYPES.WebSocketAuthMiddleware) async createWebSocketConnection(request: Request, response: Response): Promise { if (!request.headers.connectionid) { this.logger.error('Could not create a websocket connection. Missing connection id header.') diff --git a/packages/auth/src/Bootstrap/Container.ts b/packages/auth/src/Bootstrap/Container.ts index 196423394..aea040682 100644 --- a/packages/auth/src/Bootstrap/Container.ts +++ b/packages/auth/src/Bootstrap/Container.ts @@ -212,6 +212,7 @@ import { SubscriptionInvitesController } from '../Controller/SubscriptionInvites import { CreateWebSocketConnectionToken } from '../Domain/UseCase/CreateWebSocketConnectionToken/CreateWebSocketConnectionToken' import { WebSocketsController } from '../Controller/WebSocketsController' import { WebSocketServerInterface } from '@standardnotes/api' +import { CreateCrossServiceToken } from '../Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceToken' // eslint-disable-next-line @typescript-eslint/no-var-requires const newrelicFormatter = require('@newrelic/winston-enricher') @@ -462,6 +463,7 @@ export class ContainerConfigLoader { container .bind(TYPES.CreateWebSocketConnectionToken) .to(CreateWebSocketConnectionToken) + container.bind(TYPES.CreateCrossServiceToken).to(CreateCrossServiceToken) // Handlers container.bind(TYPES.UserRegisteredEventHandler).to(UserRegisteredEventHandler) @@ -534,6 +536,11 @@ export class ContainerConfigLoader { container .bind>(TYPES.OfflineUserTokenDecoder) .toConstantValue(new TokenDecoder(container.get(TYPES.AUTH_JWT_SECRET))) + container + .bind>(TYPES.WebSocketConnectionTokenDecoder) + .toConstantValue( + new TokenDecoder(container.get(TYPES.WEB_SOCKET_CONNECTION_TOKEN_SECRET)), + ) container .bind>(TYPES.OfflineUserTokenEncoder) .toConstantValue(new TokenEncoder(container.get(TYPES.AUTH_JWT_SECRET))) diff --git a/packages/auth/src/Bootstrap/Types.ts b/packages/auth/src/Bootstrap/Types.ts index 650feb715..d8c192e21 100644 --- a/packages/auth/src/Bootstrap/Types.ts +++ b/packages/auth/src/Bootstrap/Types.ts @@ -129,6 +129,7 @@ const TYPES = { GetUserAnalyticsId: Symbol.for('GetUserAnalyticsId'), VerifyPredicate: Symbol.for('VerifyPredicate'), CreateWebSocketConnectionToken: Symbol.for('CreateWebSocketConnectionToken'), + CreateCrossServiceToken: Symbol.for('CreateCrossServiceToken'), // Handlers UserRegisteredEventHandler: Symbol.for('UserRegisteredEventHandler'), AccountDeletionRequestedEventHandler: Symbol.for('AccountDeletionRequestedEventHandler'), @@ -171,6 +172,7 @@ const TYPES = { SessionTokenEncoder: Symbol.for('SessionTokenEncoder'), ValetTokenEncoder: Symbol.for('ValetTokenEncoder'), WebSocketConnectionTokenEncoder: Symbol.for('WebSocketConnectionTokenEncoder'), + WebSocketConnectionTokenDecoder: Symbol.for('WebSocketConnectionTokenDecoder'), AuthenticationMethodResolver: Symbol.for('AuthenticationMethodResolver'), DomainEventPublisher: Symbol.for('DomainEventPublisher'), DomainEventSubscriberFactory: Symbol.for('DomainEventSubscriberFactory'), diff --git a/packages/auth/src/Controller/SessionsController.spec.ts b/packages/auth/src/Controller/SessionsController.spec.ts index f54cb747f..d15576f27 100644 --- a/packages/auth/src/Controller/SessionsController.spec.ts +++ b/packages/auth/src/Controller/SessionsController.spec.ts @@ -9,43 +9,25 @@ import { ProjectorInterface } from '../Projection/ProjectorInterface' import { GetActiveSessionsForUser } from '../Domain/UseCase/GetActiveSessionsForUser' import { AuthenticateRequest } from '../Domain/UseCase/AuthenticateRequest' import { User } from '../Domain/User/User' -import { Role } from '../Domain/Role/Role' -import { CrossServiceTokenData, TokenEncoderInterface } from '@standardnotes/security' -import { GetUserAnalyticsId } from '../Domain/UseCase/GetUserAnalyticsId/GetUserAnalyticsId' +import { CreateCrossServiceToken } from '../Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceToken' describe('SessionsController', () => { let getActiveSessionsForUser: GetActiveSessionsForUser let authenticateRequest: AuthenticateRequest - let userProjector: ProjectorInterface - let tokenEncoder: TokenEncoderInterface - const jwtTTL = 60 let sessionProjector: ProjectorInterface - let roleProjector: ProjectorInterface let session: Session let request: express.Request let response: express.Response let user: User - let role: Role - let getUserAnalyticsId: GetUserAnalyticsId + let createCrossServiceToken: CreateCrossServiceToken const createController = () => - new SessionsController( - getActiveSessionsForUser, - authenticateRequest, - userProjector, - sessionProjector, - roleProjector, - tokenEncoder, - getUserAnalyticsId, - true, - jwtTTL, - ) + new SessionsController(getActiveSessionsForUser, authenticateRequest, sessionProjector, createCrossServiceToken) beforeEach(() => { session = {} as jest.Mocked user = {} as jest.Mocked - user.roles = Promise.resolve([role]) getActiveSessionsForUser = {} as jest.Mocked getActiveSessionsForUser.execute = jest.fn().mockReturnValue({ sessions: [session] }) @@ -53,21 +35,11 @@ describe('SessionsController', () => { authenticateRequest = {} as jest.Mocked authenticateRequest.execute = jest.fn() - userProjector = {} as jest.Mocked> - userProjector.projectSimple = jest.fn().mockReturnValue({ bar: 'baz' }) - - roleProjector = {} as jest.Mocked> - roleProjector.projectSimple = jest.fn().mockReturnValue({ name: 'role1', uuid: '1-3-4' }) - sessionProjector = {} as jest.Mocked> sessionProjector.projectCustom = jest.fn().mockReturnValue({ foo: 'bar' }) - sessionProjector.projectSimple = jest.fn().mockReturnValue({ test: 'test' }) - tokenEncoder = {} as jest.Mocked> - tokenEncoder.encodeExpirableToken = jest.fn().mockReturnValue('foobar') - - getUserAnalyticsId = {} as jest.Mocked - getUserAnalyticsId.execute = jest.fn().mockReturnValue({ analyticsId: 123 }) + createCrossServiceToken = {} as jest.Mocked + createCrossServiceToken.execute = jest.fn().mockReturnValue({ token: 'foobar' }) request = { params: {}, @@ -114,75 +86,6 @@ describe('SessionsController', () => { const httpResponseContent = await result.content.readAsStringAsync() const httpResponseJSON = JSON.parse(httpResponseContent) - expect(tokenEncoder.encodeExpirableToken).toHaveBeenCalledWith( - { - analyticsId: 123, - roles: [ - { - name: 'role1', - uuid: '1-3-4', - }, - ], - session: { - test: 'test', - }, - user: { - bar: 'baz', - }, - }, - 60, - ) - - expect(httpResponseJSON.authToken).toEqual('foobar') - }) - - it('should validate a session from an incoming request - disabled analytics', async () => { - authenticateRequest.execute = jest.fn().mockReturnValue({ - success: true, - user, - session, - }) - - request.headers.authorization = 'test' - - const controller = new SessionsController( - getActiveSessionsForUser, - authenticateRequest, - userProjector, - sessionProjector, - roleProjector, - tokenEncoder, - getUserAnalyticsId, - false, - jwtTTL, - ) - - const httpResponse = await controller.validate(request) - - expect(httpResponse).toBeInstanceOf(results.JsonResult) - - const result = await httpResponse.executeAsync() - const httpResponseContent = await result.content.readAsStringAsync() - const httpResponseJSON = JSON.parse(httpResponseContent) - - expect(tokenEncoder.encodeExpirableToken).toHaveBeenCalledWith( - { - roles: [ - { - name: 'role1', - uuid: '1-3-4', - }, - ], - session: { - test: 'test', - }, - user: { - bar: 'baz', - }, - }, - 60, - ) - expect(httpResponseJSON.authToken).toEqual('foobar') }) diff --git a/packages/auth/src/Controller/SessionsController.ts b/packages/auth/src/Controller/SessionsController.ts index 71bb5039e..a7796217e 100644 --- a/packages/auth/src/Controller/SessionsController.ts +++ b/packages/auth/src/Controller/SessionsController.ts @@ -12,26 +12,18 @@ import TYPES from '../Bootstrap/Types' import { Session } from '../Domain/Session/Session' import { AuthenticateRequest } from '../Domain/UseCase/AuthenticateRequest' import { GetActiveSessionsForUser } from '../Domain/UseCase/GetActiveSessionsForUser' -import { Role } from '../Domain/Role/Role' import { User } from '../Domain/User/User' import { ProjectorInterface } from '../Projection/ProjectorInterface' import { SessionProjector } from '../Projection/SessionProjector' -import { CrossServiceTokenData, TokenEncoderInterface } from '@standardnotes/security' -import { RoleName } from '@standardnotes/common' -import { GetUserAnalyticsId } from '../Domain/UseCase/GetUserAnalyticsId/GetUserAnalyticsId' +import { CreateCrossServiceToken } from '../Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceToken' @controller('/sessions') export class SessionsController extends BaseHttpController { constructor( @inject(TYPES.GetActiveSessionsForUser) private getActiveSessionsForUser: GetActiveSessionsForUser, @inject(TYPES.AuthenticateRequest) private authenticateRequest: AuthenticateRequest, - @inject(TYPES.UserProjector) private userProjector: ProjectorInterface, @inject(TYPES.SessionProjector) private sessionProjector: ProjectorInterface, - @inject(TYPES.RoleProjector) private roleProjector: ProjectorInterface, - @inject(TYPES.CrossServiceTokenEncoder) private tokenEncoder: TokenEncoderInterface, - @inject(TYPES.GetUserAnalyticsId) private getUserAnalyticsId: GetUserAnalyticsId, - @inject(TYPES.ANALYTICS_ENABLED) private analyticsEnabled: boolean, - @inject(TYPES.AUTH_JWT_TTL) private jwtTTL: number, + @inject(TYPES.CreateCrossServiceToken) private createCrossServiceToken: CreateCrossServiceToken, ) { super() } @@ -56,25 +48,12 @@ export class SessionsController extends BaseHttpController { const user = authenticateRequestResponse.user as User - const roles = await user.roles + const result = await this.createCrossServiceToken.execute({ + user, + session: authenticateRequestResponse.session, + }) - const authTokenData: CrossServiceTokenData = { - user: this.projectUser(user), - roles: this.projectRoles(roles), - } - - if (this.analyticsEnabled) { - const { analyticsId } = await this.getUserAnalyticsId.execute({ userUuid: user.uuid }) - authTokenData.analyticsId = analyticsId - } - - if (authenticateRequestResponse.session !== undefined) { - authTokenData.session = this.projectSession(authenticateRequestResponse.session) - } - - const authToken = this.tokenEncoder.encodeExpirableToken(authTokenData, this.jwtTTL) - - return this.json({ authToken }) + return this.json({ authToken: result.token }) } @httpGet('/', TYPES.AuthMiddleware, TYPES.SessionMiddleware) @@ -93,36 +72,4 @@ export class SessionsController extends BaseHttpController { ), ) } - - private projectUser(user: User): { uuid: string; email: string } { - return <{ uuid: string; email: string }>this.userProjector.projectSimple(user) - } - - private projectSession(session: Session): { - uuid: string - api_version: string - created_at: string - updated_at: string - device_info: string - readonly_access: boolean - access_expiration: string - refresh_expiration: string - } { - return < - { - uuid: string - api_version: string - created_at: string - updated_at: string - device_info: string - readonly_access: boolean - access_expiration: string - refresh_expiration: string - } - >this.sessionProjector.projectSimple(session) - } - - private projectRoles(roles: Array): Array<{ uuid: string; name: RoleName }> { - return roles.map((role) => <{ uuid: string; name: RoleName }>this.roleProjector.projectSimple(role)) - } } diff --git a/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceToken.spec.ts b/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceToken.spec.ts new file mode 100644 index 000000000..4dcf718b0 --- /dev/null +++ b/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceToken.spec.ts @@ -0,0 +1,173 @@ +import 'reflect-metadata' + +import { TokenEncoderInterface, CrossServiceTokenData } from '@standardnotes/security' +import { ProjectorInterface } from '../../../Projection/ProjectorInterface' +import { Session } from '../../Session/Session' +import { User } from '../../User/User' +import { Role } from '../../Role/Role' +import { UserRepositoryInterface } from '../../User/UserRepositoryInterface' +import { GetUserAnalyticsId } from '../GetUserAnalyticsId/GetUserAnalyticsId' + +import { CreateCrossServiceToken } from './CreateCrossServiceToken' + +describe('CreateCrossServiceToken', () => { + let userProjector: ProjectorInterface + let sessionProjector: ProjectorInterface + let roleProjector: ProjectorInterface + let tokenEncoder: TokenEncoderInterface + let getUserAnalyticsId: GetUserAnalyticsId + let userRepository: UserRepositoryInterface + const jwtTTL = 60 + + let session: Session + let user: User + let role: Role + + const createUseCase = (analyticsEnabled = true) => + new CreateCrossServiceToken( + userProjector, + sessionProjector, + roleProjector, + tokenEncoder, + getUserAnalyticsId, + userRepository, + analyticsEnabled, + jwtTTL, + ) + + beforeEach(() => { + session = {} as jest.Mocked + + user = {} as jest.Mocked + user.roles = Promise.resolve([role]) + + userProjector = {} as jest.Mocked> + userProjector.projectSimple = jest.fn().mockReturnValue({ bar: 'baz' }) + + roleProjector = {} as jest.Mocked> + roleProjector.projectSimple = jest.fn().mockReturnValue({ name: 'role1', uuid: '1-3-4' }) + + sessionProjector = {} as jest.Mocked> + sessionProjector.projectCustom = jest.fn().mockReturnValue({ foo: 'bar' }) + sessionProjector.projectSimple = jest.fn().mockReturnValue({ test: 'test' }) + + tokenEncoder = {} as jest.Mocked> + tokenEncoder.encodeExpirableToken = jest.fn().mockReturnValue('foobar') + + getUserAnalyticsId = {} as jest.Mocked + getUserAnalyticsId.execute = jest.fn().mockReturnValue({ analyticsId: 123 }) + + userRepository = {} as jest.Mocked + userRepository.findOneByUuid = jest.fn().mockReturnValue(user) + }) + + it('should create a cross service token for user', async () => { + await createUseCase().execute({ + user, + session, + }) + + expect(tokenEncoder.encodeExpirableToken).toHaveBeenCalledWith( + { + analyticsId: 123, + roles: [ + { + name: 'role1', + uuid: '1-3-4', + }, + ], + session: { + test: 'test', + }, + user: { + bar: 'baz', + }, + }, + 60, + ) + }) + + it('should create a cross service token for user - analytics disabled', async () => { + await createUseCase(false).execute({ + user, + session, + }) + + expect(tokenEncoder.encodeExpirableToken).toHaveBeenCalledWith( + { + roles: [ + { + name: 'role1', + uuid: '1-3-4', + }, + ], + session: { + test: 'test', + }, + user: { + bar: 'baz', + }, + }, + 60, + ) + }) + + it('should create a cross service token for user without a session', async () => { + await createUseCase().execute({ + user, + }) + + expect(tokenEncoder.encodeExpirableToken).toHaveBeenCalledWith( + { + analyticsId: 123, + roles: [ + { + name: 'role1', + uuid: '1-3-4', + }, + ], + user: { + bar: 'baz', + }, + }, + 60, + ) + }) + + it('should create a cross service token for user by user uuid', async () => { + await createUseCase().execute({ + userUuid: '1-2-3', + }) + + expect(tokenEncoder.encodeExpirableToken).toHaveBeenCalledWith( + { + analyticsId: 123, + roles: [ + { + name: 'role1', + uuid: '1-3-4', + }, + ], + user: { + bar: 'baz', + }, + }, + 60, + ) + }) + + it('should throw an error if user does not exist', async () => { + userRepository.findOneByUuid = jest.fn().mockReturnValue(null) + + let caughtError = null + try { + await createUseCase().execute({ + userUuid: '1-2-3', + }) + } catch (error) { + caughtError = error + } + + expect(caughtError).not.toBeNull() + }) +}) diff --git a/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceToken.ts b/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceToken.ts new file mode 100644 index 000000000..bbda080c5 --- /dev/null +++ b/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceToken.ts @@ -0,0 +1,91 @@ +import { RoleName } from '@standardnotes/common' +import { TokenEncoderInterface, CrossServiceTokenData } from '@standardnotes/security' + +import { inject, injectable } from 'inversify' +import TYPES from '../../../Bootstrap/Types' +import { ProjectorInterface } from '../../../Projection/ProjectorInterface' +import { Role } from '../../Role/Role' +import { Session } from '../../Session/Session' +import { User } from '../../User/User' +import { UserRepositoryInterface } from '../../User/UserRepositoryInterface' +import { GetUserAnalyticsId } from '../GetUserAnalyticsId/GetUserAnalyticsId' +import { UseCaseInterface } from '../UseCaseInterface' +import { CreateCrossServiceTokenDTO } from './CreateCrossServiceTokenDTO' +import { CreateCrossServiceTokenResponse } from './CreateCrossServiceTokenResponse' + +@injectable() +export class CreateCrossServiceToken implements UseCaseInterface { + constructor( + @inject(TYPES.UserProjector) private userProjector: ProjectorInterface, + @inject(TYPES.SessionProjector) private sessionProjector: ProjectorInterface, + @inject(TYPES.RoleProjector) private roleProjector: ProjectorInterface, + @inject(TYPES.CrossServiceTokenEncoder) private tokenEncoder: TokenEncoderInterface, + @inject(TYPES.GetUserAnalyticsId) private getUserAnalyticsId: GetUserAnalyticsId, + @inject(TYPES.UserRepository) private userRepository: UserRepositoryInterface, + @inject(TYPES.ANALYTICS_ENABLED) private analyticsEnabled: boolean, + @inject(TYPES.AUTH_JWT_TTL) private jwtTTL: number, + ) {} + + async execute(dto: CreateCrossServiceTokenDTO): Promise { + let user: User | undefined | null = dto.user + if (user === undefined && dto.userUuid !== undefined) { + user = await this.userRepository.findOneByUuid(dto.userUuid) + } + + if (!user) { + throw new Error(`Could not find user with uuid ${dto.userUuid}`) + } + + const roles = await user.roles + + const authTokenData: CrossServiceTokenData = { + user: this.projectUser(user), + roles: this.projectRoles(roles), + } + + if (this.analyticsEnabled) { + const { analyticsId } = await this.getUserAnalyticsId.execute({ userUuid: user.uuid }) + authTokenData.analyticsId = analyticsId + } + + if (dto.session !== undefined) { + authTokenData.session = this.projectSession(dto.session) + } + + return { + token: this.tokenEncoder.encodeExpirableToken(authTokenData, this.jwtTTL), + } + } + + private projectUser(user: User): { uuid: string; email: string } { + return <{ uuid: string; email: string }>this.userProjector.projectSimple(user) + } + + private projectSession(session: Session): { + uuid: string + api_version: string + created_at: string + updated_at: string + device_info: string + readonly_access: boolean + access_expiration: string + refresh_expiration: string + } { + return < + { + uuid: string + api_version: string + created_at: string + updated_at: string + device_info: string + readonly_access: boolean + access_expiration: string + refresh_expiration: string + } + >this.sessionProjector.projectSimple(session) + } + + private projectRoles(roles: Array): Array<{ uuid: string; name: RoleName }> { + return roles.map((role) => <{ uuid: string; name: RoleName }>this.roleProjector.projectSimple(role)) + } +} diff --git a/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceTokenDTO.ts b/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceTokenDTO.ts new file mode 100644 index 000000000..a6451f96e --- /dev/null +++ b/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceTokenDTO.ts @@ -0,0 +1,13 @@ +import { Either, Uuid } from '@standardnotes/common' +import { Session } from '../../Session/Session' +import { User } from '../../User/User' + +export type CreateCrossServiceTokenDTO = Either< + { + user: User + session?: Session + }, + { + userUuid: Uuid + } +> diff --git a/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceTokenResponse.ts b/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceTokenResponse.ts new file mode 100644 index 000000000..ef3dcbaad --- /dev/null +++ b/packages/auth/src/Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceTokenResponse.ts @@ -0,0 +1,3 @@ +export type CreateCrossServiceTokenResponse = { + token: string +} diff --git a/packages/auth/src/Infra/InversifyExpressUtils/InversifyExpressWebSocketsController.ts b/packages/auth/src/Infra/InversifyExpressUtils/InversifyExpressWebSocketsController.ts index 43a7d431f..d8e2ff8eb 100644 --- a/packages/auth/src/Infra/InversifyExpressUtils/InversifyExpressWebSocketsController.ts +++ b/packages/auth/src/Infra/InversifyExpressUtils/InversifyExpressWebSocketsController.ts @@ -1,4 +1,6 @@ import { WebSocketServerInterface } from '@standardnotes/api' +import { ErrorTag } from '@standardnotes/common' +import { TokenDecoderInterface, WebSocketConnectionTokenData } from '@standardnotes/security' import { Request, Response } from 'express' import { inject } from 'inversify' import { @@ -11,6 +13,7 @@ import { } from 'inversify-express-utils' import TYPES from '../../Bootstrap/Types' import { AddWebSocketsConnection } from '../../Domain/UseCase/AddWebSocketsConnection/AddWebSocketsConnection' +import { CreateCrossServiceToken } from '../../Domain/UseCase/CreateCrossServiceToken/CreateCrossServiceToken' import { RemoveWebSocketsConnection } from '../../Domain/UseCase/RemoveWebSocketsConnection/RemoveWebSocketsConnection' @controller('/sockets') @@ -18,7 +21,10 @@ export class InversifyExpressWebSocketsController extends BaseHttpController { constructor( @inject(TYPES.AddWebSocketsConnection) private addWebSocketsConnection: AddWebSocketsConnection, @inject(TYPES.RemoveWebSocketsConnection) private removeWebSocketsConnection: RemoveWebSocketsConnection, + @inject(TYPES.CreateCrossServiceToken) private createCrossServiceToken: CreateCrossServiceToken, @inject(TYPES.WebSocketsController) private webSocketsController: WebSocketServerInterface, + @inject(TYPES.WebSocketConnectionTokenDecoder) + private tokenDecoder: TokenDecoderInterface, ) { super() } @@ -53,4 +59,39 @@ export class InversifyExpressWebSocketsController extends BaseHttpController { return this.json(result) } + + @httpPost('/tokens/validate') + async validateToken(request: Request): Promise { + if (!request.headers.authorization) { + return this.json( + { + error: { + tag: ErrorTag.AuthInvalid, + message: 'Invalid authorization token.', + }, + }, + 401, + ) + } + + const token: WebSocketConnectionTokenData | undefined = this.tokenDecoder.decodeToken(request.headers.authorization) + + if (token === undefined) { + return this.json( + { + error: { + tag: ErrorTag.AuthInvalid, + message: 'Invalid authorization token.', + }, + }, + 401, + ) + } + + const result = await this.createCrossServiceToken.execute({ + userUuid: token.userUuid, + }) + + return this.json({ authToken: result.token }) + } }