diff --git a/fluxer_api/src/api/read_state/ReadStateService.test.ts b/fluxer_api/src/api/read_state/ReadStateService.test.ts index bd2b1a55f..3f535faf6 100644 --- a/fluxer_api/src/api/read_state/ReadStateService.test.ts +++ b/fluxer_api/src/api/read_state/ReadStateService.test.ts @@ -1,8 +1,17 @@ // SPDX-License-Identifier: AGPL-3.0-or-later +import {BadGatewayError} from '@fluxer/errors/src/domains/core/BadGatewayError'; import {describe, expect, it, vi} from 'vitest'; -import {type ChannelID, createChannelID, createMessageID, createUserID, type UserID} from '../BrandedTypes'; +import { + type ChannelID, + createChannelID, + createMessageID, + createUserID, + type MessageID, + type UserID, +} from '../BrandedTypes'; import type {IGatewayService} from '../infrastructure/IGatewayService'; +import {ReadState} from '../models/ReadState'; import type {IReadStateRepository} from './IReadStateRepository'; import {ReadStateService} from './ReadStateService'; @@ -52,3 +61,146 @@ describe('ReadStateService.bulkIncrementMentionCounts', () => { expect(invalidatePushBadgeCounts).not.toHaveBeenCalled(); }); }); + +const USER_ID = createUserID(20n); +const CHANNEL_ID = createChannelID(21n); +const MESSAGE_ID = createMessageID(22n); + +function makeReadState(channelId: ChannelID, messageId: MessageID, mentionCount = 0): ReadState { + return new ReadState({ + user_id: USER_ID, + channel_id: channelId, + message_id: messageId, + mention_count: mentionCount, + last_pin_timestamp: null, + version: 5n, + }); +} + +describe('ReadStateService gateway side effects after the write', () => { + it('returns the committed read state when the badge invalidation fails', async () => { + const stored: Array<{channelId: ChannelID; messageId: MessageID}> = []; + const repository = { + upsertReadState: vi.fn(async (_userId: UserID, channelId: ChannelID, messageId: MessageID) => { + stored.push({channelId, messageId}); + return makeReadState(channelId, messageId); + }), + } as unknown as IReadStateRepository; + const gatewayService = { + invalidatePushBadgeCount: vi.fn().mockRejectedValue(new BadGatewayError()), + clearPushChannelNotifications: vi.fn().mockResolvedValue(undefined), + dispatchPresence: vi.fn().mockResolvedValue(undefined), + } as unknown as IGatewayService; + const service = new ReadStateService(repository, gatewayService); + + const readState = await service.ackMessage({ + userId: USER_ID, + channelId: CHANNEL_ID, + messageId: MESSAGE_ID, + mentionCount: 0, + }); + + expect(readState.channelId).toBe(CHANNEL_ID); + expect(readState.lastMessageId).toBe(MESSAGE_ID); + expect(stored).toEqual([{channelId: CHANNEL_ID, messageId: MESSAGE_ID}]); + expect(gatewayService.dispatchPresence).toHaveBeenCalledTimes(1); + }); + + it('acknowledges the message when the MESSAGE_ACK dispatch fails', async () => { + const repository = { + upsertReadState: vi.fn(async (_userId: UserID, channelId: ChannelID, messageId: MessageID) => + makeReadState(channelId, messageId), + ), + } as unknown as IReadStateRepository; + const gatewayService = { + invalidatePushBadgeCount: vi.fn().mockResolvedValue(undefined), + clearPushChannelNotifications: vi.fn().mockResolvedValue(undefined), + dispatchPresence: vi.fn().mockRejectedValue(new BadGatewayError()), + } as unknown as IGatewayService; + const service = new ReadStateService(repository, gatewayService); + + const readState = await service.ackMessage({ + userId: USER_ID, + channelId: CHANNEL_ID, + messageId: MESSAGE_ID, + mentionCount: 0, + }); + + expect(readState.lastMessageId).toBe(MESSAGE_ID); + }); + + it('returns every entry of the entry-by-entry path when the dispatch fails', async () => { + const stored: Array = []; + const repository = { + upsertReadState: vi.fn(async (_userId: UserID, channelId: ChannelID, messageId: MessageID) => { + stored.push(channelId.toString()); + return makeReadState(channelId, messageId, 1); + }), + } as unknown as IReadStateRepository; + const gatewayService = { + invalidatePushBadgeCount: vi.fn().mockResolvedValue(undefined), + clearPushChannelNotifications: vi.fn().mockResolvedValue(undefined), + dispatchPresence: vi.fn().mockRejectedValue(new BadGatewayError()), + } as unknown as IGatewayService; + const service = new ReadStateService(repository, gatewayService); + + const readStates = await service.ackReadStates({ + userId: USER_ID, + readStates: [ + {channelId: CHANNEL_ID, messageId: MESSAGE_ID, manual: true}, + {channelId: createChannelID(23n), messageId: createMessageID(24n), manual: true}, + ], + }); + + expect(readStates.map((readState) => readState.channelId.toString())).toEqual(['21', '23']); + expect(stored).toEqual(['21', '23']); + }); + + it('deletes the read state when the badge invalidation fails', async () => { + const deleteReadState = vi.fn().mockResolvedValue(undefined); + const repository = {deleteReadState} as unknown as IReadStateRepository; + const gatewayService = { + invalidatePushBadgeCount: vi.fn().mockRejectedValue(new BadGatewayError()), + } as unknown as IGatewayService; + const service = new ReadStateService(repository, gatewayService); + + await expect(service.deleteReadState({userId: USER_ID, channelId: CHANNEL_ID})).resolves.toBeUndefined(); + + expect(deleteReadState).toHaveBeenCalledWith(USER_ID, CHANNEL_ID); + }); + + it('increments the mention count when the badge invalidation fails', async () => { + const incrementReadStateMentions = vi.fn().mockResolvedValue(makeReadState(CHANNEL_ID, MESSAGE_ID, 1)); + const repository = {incrementReadStateMentions} as unknown as IReadStateRepository; + const gatewayService = { + invalidatePushBadgeCount: vi.fn().mockRejectedValue(new BadGatewayError()), + } as unknown as IGatewayService; + const service = new ReadStateService(repository, gatewayService); + + await expect( + service.incrementMentionCount({userId: USER_ID, channelId: CHANNEL_ID, messageId: MESSAGE_ID}), + ).resolves.toBeUndefined(); + + expect(incrementReadStateMentions).toHaveBeenCalledTimes(1); + }); + + it('returns the bulk acknowledged states when the badge invalidation fails', async () => { + const updated = [makeReadState(CHANNEL_ID, MESSAGE_ID)]; + const repository = { + bulkAckMessages: vi.fn().mockResolvedValue(updated), + } as unknown as IReadStateRepository; + const gatewayService = { + invalidatePushBadgeCount: vi.fn().mockRejectedValue(new BadGatewayError()), + clearPushChannelNotifications: vi.fn().mockResolvedValue(undefined), + dispatchPresence: vi.fn().mockResolvedValue(undefined), + } as unknown as IGatewayService; + const service = new ReadStateService(repository, gatewayService); + + const readStates = await service.bulkAckMessages({ + userId: USER_ID, + readStates: [{channelId: CHANNEL_ID, messageId: MESSAGE_ID}], + }); + + expect(readStates).toBe(updated); + }); +}); diff --git a/fluxer_api/src/api/read_state/ReadStateService.ts b/fluxer_api/src/api/read_state/ReadStateService.ts index 600eecea6..6e899be35 100644 --- a/fluxer_api/src/api/read_state/ReadStateService.ts +++ b/fluxer_api/src/api/read_state/ReadStateService.ts @@ -34,7 +34,7 @@ export class ReadStateService { undefined, manual ?? false, ); - await this.gatewayService.invalidatePushBadgeCount({userId}); + await this.invalidatePushBadgeCount(userId); if (!silent) { await this.clearPushChannelNotifications({userId, channelId, messageId}); } @@ -115,7 +115,7 @@ export class ReadStateService { try { const updatedReadStates = await this.repository.bulkAckMessages(userId, readStates); const readStatesByChannel = new Map(updatedReadStates.map((readState) => [readState.channelId, readState])); - await this.gatewayService.invalidatePushBadgeCount({userId}); + await this.invalidatePushBadgeCount(userId); await Promise.all( readStates.map(({channelId, messageId}) => Promise.all([ @@ -145,7 +145,7 @@ export class ReadStateService { async deleteReadState({userId, channelId}: {userId: UserID; channelId: ChannelID}): Promise { await this.repository.deleteReadState(userId, channelId); - await this.gatewayService.invalidatePushBadgeCount({userId}); + await this.invalidatePushBadgeCount(userId); } async incrementMentionCount({ @@ -161,7 +161,7 @@ export class ReadStateService { if (readState == null) { return; } - await this.gatewayService.invalidatePushBadgeCount({userId}); + await this.invalidatePushBadgeCount(userId); } async bulkIncrementMentionCounts( @@ -196,6 +196,13 @@ export class ReadStateService { await this.dispatchPinsAck({userId, channelId, timestamp}); } + private async invalidatePushBadgeCount(userId: UserID): Promise { + await this.gatewayService.invalidatePushBadgeCount({userId}).catch((error) => { + Logger.error({userId: userId.toString(), error}, 'Failed to invalidate push badge count'); + return null; + }); + } + private async dispatchMessageAck(params: { userId: UserID; channelId: ChannelID; diff --git a/fluxer_api/src/api/user/UserHarvestRepository.ts b/fluxer_api/src/api/user/UserHarvestRepository.ts index e1298fd59..f4cf0c8f2 100644 --- a/fluxer_api/src/api/user/UserHarvestRepository.ts +++ b/fluxer_api/src/api/user/UserHarvestRepository.ts @@ -96,6 +96,8 @@ export class UserHarvestRepository { {user_id: userId, harvest_id: harvestId}, { started_at: Db.set(new Date()), + failed_at: Db.clear(), + error_message: Db.clear(), progress_percent: Db.set(0), progress_step: Db.set('Starting harvest'), }, diff --git a/fluxer_api/src/api/user/entrance_sound/EntranceSoundService.test.ts b/fluxer_api/src/api/user/entrance_sound/EntranceSoundService.test.ts new file mode 100644 index 000000000..076cfd821 --- /dev/null +++ b/fluxer_api/src/api/user/entrance_sound/EntranceSoundService.test.ts @@ -0,0 +1,158 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +import {describe, expect, it, vi} from 'vitest'; +import {createUserID, type EntranceSoundID, type UserID} from '../../BrandedTypes'; +import {Config} from '../../Config'; +import type {IMediaService} from '../../infrastructure/IMediaService'; +import type {IStorageService} from '../../infrastructure/IStorageService'; +import type {EntranceSound} from '../../models/EntranceSound'; +import type {EntranceSoundRepository} from './EntranceSoundRepository'; +import {EntranceSoundService} from './EntranceSoundService'; + +const USER_ID = createUserID(1234567890123456n); + +function createWavBase64(sampleCount: number): string { + const dataLength = sampleCount * 2; + const buffer = Buffer.alloc(44 + dataLength); + buffer.write('RIFF', 0, 'ascii'); + buffer.writeUInt32LE(36 + dataLength, 4); + buffer.write('WAVE', 8, 'ascii'); + buffer.write('fmt ', 12, 'ascii'); + buffer.writeUInt32LE(16, 16); + buffer.writeUInt16LE(1, 20); + buffer.writeUInt16LE(1, 22); + buffer.writeUInt32LE(8000, 24); + buffer.writeUInt32LE(16000, 28); + buffer.writeUInt16LE(2, 32); + buffer.writeUInt16LE(16, 34); + buffer.write('data', 36, 'ascii'); + buffer.writeUInt32LE(dataLength, 40); + for (let index = 0; index < sampleCount; index += 1) { + buffer.writeInt16LE(Math.round(8000 * Math.sin(index / 8)), 44 + index * 2); + } + return buffer.toString('base64'); +} + +const HALF_SECOND_WAV = createWavBase64(4000); +const OTHER_WAV = createWavBase64(6000); + +class FakeEntranceSoundRepository { + readonly sounds = new Map(); + failNextUpsert = false; + + async listSounds(_userId: UserID): Promise> { + return [...this.sounds.values()]; + } + + async getSound(_userId: UserID, soundId: EntranceSoundID): Promise { + return this.sounds.get(soundId.toString()) ?? null; + } + + async upsertSound(sound: EntranceSound): Promise { + if (this.failNextUpsert) { + this.failNextUpsert = false; + throw new Error('Failed to persist entrance sound'); + } + this.sounds.set(sound.soundId.toString(), sound); + return sound; + } + + async deleteSound(_userId: UserID, soundId: EntranceSoundID): Promise { + this.sounds.delete(soundId.toString()); + } + + async deleteSelectionsForSound(_userId: UserID, _soundId: EntranceSoundID): Promise {} +} + +function createService() { + const repository = new FakeEntranceSoundRepository(); + const objects = new Set(); + const deleteObject = vi.fn(async (_bucket: string, key: string) => { + objects.delete(key); + }); + const storageService = { + uploadObject: async (params: {key: string}) => { + objects.add(params.key); + }, + deleteObject, + } as unknown as IStorageService; + const mediaService = { + getMetadata: async () => ({ + format: 'wav', + content_type: 'audio/wav', + content_hash: 'content-hash', + size: 0, + duration: 0.5, + nsfw: false, + }), + } as unknown as IMediaService; + const service = new EntranceSoundService( + repository as unknown as EntranceSoundRepository, + storageService, + mediaService, + ); + return {service, repository, objects, deleteObject}; +} + +function keyFromUrl(url: string): string { + return url.slice(`${Config.endpoints.media}/`.length); +} + +describe('EntranceSoundService shared object references', () => { + it('keeps the shared object when one of two entries with identical audio is deleted', async () => { + const {service, objects, deleteObject} = createService(); + const first = await service.upload({userId: USER_ID, name: 'first', base64Audio: HALF_SECOND_WAV}); + const second = await service.upload({userId: USER_ID, name: 'second', base64Audio: HALF_SECOND_WAV}); + expect(keyFromUrl(second.url)).toBe(keyFromUrl(first.url)); + + await service.delete(USER_ID, first.sound.soundId); + + expect(deleteObject).not.toHaveBeenCalled(); + expect(objects.has(keyFromUrl(second.url))).toBe(true); + + await service.delete(USER_ID, second.sound.soundId); + + expect(deleteObject).toHaveBeenCalledTimes(1); + expect(deleteObject).toHaveBeenCalledWith(Config.s3.buckets.cdn, keyFromUrl(second.url)); + expect(objects.size).toBe(0); + }); + + it('keeps the shared object when a duplicate upload fails to persist', async () => { + const {service, repository, objects, deleteObject} = createService(); + const first = await service.upload({userId: USER_ID, name: 'first', base64Audio: HALF_SECOND_WAV}); + + repository.failNextUpsert = true; + await expect(service.upload({userId: USER_ID, name: 'duplicate', base64Audio: HALF_SECOND_WAV})).rejects.toThrow( + 'Failed to persist entrance sound', + ); + + expect(deleteObject).not.toHaveBeenCalled(); + expect(objects.has(keyFromUrl(first.url))).toBe(true); + await expect(service.listLibrary(USER_ID)).resolves.toHaveLength(1); + }); + + it('rolls back the object when the only entry referencing it fails to persist', async () => { + const {service, repository, objects, deleteObject} = createService(); + + repository.failNextUpsert = true; + await expect(service.upload({userId: USER_ID, name: 'only', base64Audio: HALF_SECOND_WAV})).rejects.toThrow( + 'Failed to persist entrance sound', + ); + + expect(deleteObject).toHaveBeenCalledTimes(1); + expect(objects.size).toBe(0); + }); + + it('deletes the object for a lone entry and leaves other entries alone', async () => { + const {service, objects, deleteObject} = createService(); + const lone = await service.upload({userId: USER_ID, name: 'lone', base64Audio: HALF_SECOND_WAV}); + const other = await service.upload({userId: USER_ID, name: 'other', base64Audio: OTHER_WAV}); + expect(keyFromUrl(other.url)).not.toBe(keyFromUrl(lone.url)); + + await service.delete(USER_ID, lone.sound.soundId); + + expect(deleteObject).toHaveBeenCalledTimes(1); + expect(deleteObject).toHaveBeenCalledWith(Config.s3.buckets.cdn, keyFromUrl(lone.url)); + expect(objects.has(keyFromUrl(other.url))).toBe(true); + }); +}); diff --git a/fluxer_api/src/api/user/entrance_sound/EntranceSoundService.ts b/fluxer_api/src/api/user/entrance_sound/EntranceSoundService.ts index 0909aa05c..1560d34f2 100644 --- a/fluxer_api/src/api/user/entrance_sound/EntranceSoundService.ts +++ b/fluxer_api/src/api/user/entrance_sound/EntranceSoundService.ts @@ -65,6 +65,18 @@ export class EntranceSoundService { return `${SOUND_PATH_PREFIX}/${userId}/${hash}.${extension}`; } + private async objectIsReferenced( + userId: UserID, + hash: string, + extension: string, + excludeSoundId?: EntranceSoundID, + ): Promise { + const sounds = await this.repository.listSounds(userId); + return sounds.some( + (sound) => sound.hash === hash && sound.extension === extension && sound.soundId !== excludeSoundId, + ); + } + async listLibrary(userId: UserID): Promise> { const sounds = await this.repository.listSounds(userId); return sounds.map((sound) => ({sound, url: this.cdnUrlFor(sound)})); @@ -165,7 +177,9 @@ export class EntranceSoundService { await this.repository.upsertSound(sound); } catch (error) { Logger.error({error, userId: userId.toString(), s3Key}, 'Failed to persist entrance sound; rolling back S3'); - await this.storageService.deleteObject(Config.s3.buckets.cdn, s3Key).catch(() => {}); + if (!(await this.objectIsReferenced(userId, hash, extension).catch(() => true))) { + await this.storageService.deleteObject(Config.s3.buckets.cdn, s3Key).catch(() => {}); + } throw error; } return {sound, url: this.cdnUrlFor(sound)}; @@ -191,6 +205,9 @@ export class EntranceSoundService { if (!existing) return; await this.repository.deleteSelectionsForSound(userId, soundId); await this.repository.deleteSound(userId, soundId); + if (await this.objectIsReferenced(userId, existing.hash, existing.extension, soundId).catch(() => true)) { + return; + } const s3Key = this.s3KeyFor(userId, existing.hash, existing.extension as EntranceSoundExtension); await this.storageService.deleteObject(Config.s3.buckets.cdn, s3Key).catch((error) => { Logger.error({error, userId: userId.toString(), s3Key}, 'Failed to delete entrance sound from S3'); diff --git a/fluxer_api/src/api/user/repositories/IUserContentRepository.ts b/fluxer_api/src/api/user/repositories/IUserContentRepository.ts index 84991c144..163fe495a 100644 --- a/fluxer_api/src/api/user/repositories/IUserContentRepository.ts +++ b/fluxer_api/src/api/user/repositories/IUserContentRepository.ts @@ -27,6 +27,7 @@ export interface IUserContentRepository { deleteRecentMentions(mentions: Array): Promise; deleteAllRecentMentions(userId: UserID): Promise; listSavedMessages(userId: UserID, limit?: number, before?: MessageID): Promise>; + countSavedMessages(userId: UserID): Promise; createSavedMessage(userId: UserID, channelId: ChannelID, messageId: MessageID): Promise; deleteSavedMessage(userId: UserID, messageId: MessageID): Promise; deleteAllSavedMessages(userId: UserID): Promise; diff --git a/fluxer_api/src/api/user/repositories/SavedMessageRepository.test.ts b/fluxer_api/src/api/user/repositories/SavedMessageRepository.test.ts new file mode 100644 index 000000000..62143f85d --- /dev/null +++ b/fluxer_api/src/api/user/repositories/SavedMessageRepository.test.ts @@ -0,0 +1,56 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +import {afterEach, beforeEach, describe, expect, it} from 'vitest'; +import {createChannelID, createMessageID, createUserID, type UserID} from '../../BrandedTypes'; +import {setCassandraQueryExecutorForTesting} from '../../database/CassandraQueryExecution'; +import {InMemoryCassandraQueryExecutor} from '../../test/InMemoryCassandraQueryExecutor'; +import {SavedMessageRepository} from './SavedMessageRepository'; + +const CHANNEL_ID = createChannelID(10n); + +let executor: InMemoryCassandraQueryExecutor; + +async function seedSavedMessages(repository: SavedMessageRepository, userId: UserID, count: number) { + for (let index = 0; index < count; index++) { + await repository.createSavedMessage(userId, CHANNEL_ID, createMessageID(BigInt(1000 + index))); + } +} + +describe('SavedMessageRepository.countSavedMessages', () => { + beforeEach(() => { + executor = new InMemoryCassandraQueryExecutor(); + setCassandraQueryExecutorForTesting(executor); + }); + afterEach(() => { + executor.reset(); + setCassandraQueryExecutorForTesting(null); + }); + + it('returns zero when the user has saved nothing', async () => { + const repository = new SavedMessageRepository(); + expect(await repository.countSavedMessages(createUserID(1n))).toBe(0); + }); + + it('counts past the ceiling a listing page can report', async () => { + const repository = new SavedMessageRepository(); + const userId = createUserID(1n); + await seedSavedMessages(repository, userId, 1200); + expect(await repository.listSavedMessages(userId, 1000)).toHaveLength(1000); + expect(await repository.countSavedMessages(userId)).toBe(1200); + }); + + it('follows creates and deletes and stays scoped to one user', async () => { + const repository = new SavedMessageRepository(); + const userId = createUserID(1n); + const otherUserId = createUserID(2n); + await seedSavedMessages(repository, userId, 3); + await seedSavedMessages(repository, otherUserId, 7); + expect(await repository.countSavedMessages(userId)).toBe(3); + await repository.deleteSavedMessage(userId, createMessageID(1001n)); + expect(await repository.countSavedMessages(userId)).toBe(2); + expect(await repository.countSavedMessages(otherUserId)).toBe(7); + await repository.deleteAllSavedMessages(userId); + expect(await repository.countSavedMessages(userId)).toBe(0); + expect(await repository.countSavedMessages(otherUserId)).toBe(7); + }); +}); diff --git a/fluxer_api/src/api/user/repositories/SavedMessageRepository.ts b/fluxer_api/src/api/user/repositories/SavedMessageRepository.ts index e33d98209..e3fa2a2b6 100644 --- a/fluxer_api/src/api/user/repositories/SavedMessageRepository.ts +++ b/fluxer_api/src/api/user/repositories/SavedMessageRepository.ts @@ -2,7 +2,7 @@ import {generateSnowflake} from '@fluxer/snowflake/src/Snowflake'; import {type ChannelID, createMessageID, type MessageID, type UserID} from '../../BrandedTypes'; -import {deleteOneOrMany, fetchMany, upsertOne} from '../../database/CassandraQueryExecution'; +import {deleteOneOrMany, fetchMany, fetchOne, upsertOne} from '../../database/CassandraQueryExecution'; import type {SavedMessageRow} from '../../database/types/UserTypes'; import {SavedMessage} from '../../models/SavedMessage'; import {SavedMessages} from '../../Tables'; @@ -13,7 +13,18 @@ const createFetchSavedMessagesQuery = (limit: number) => limit, }); +const COUNT_SAVED_MESSAGES_CQL = SavedMessages.selectCountCql({ + where: SavedMessages.where.eq('user_id'), +}); + export class SavedMessageRepository { + async countSavedMessages(userId: UserID): Promise { + const result = await fetchOne<{ + count: bigint; + }>(COUNT_SAVED_MESSAGES_CQL, {user_id: userId}); + return result ? Number(result.count) : 0; + } + async listSavedMessages( userId: UserID, limit: number = 25, diff --git a/fluxer_api/src/api/user/repositories/UserContentRepository.ts b/fluxer_api/src/api/user/repositories/UserContentRepository.ts index 71d07bcdd..59273b49a 100644 --- a/fluxer_api/src/api/user/repositories/UserContentRepository.ts +++ b/fluxer_api/src/api/user/repositories/UserContentRepository.ts @@ -186,6 +186,10 @@ export class UserContentRepository implements IUserContentRepository { return this.savedMessageRepository.listSavedMessages(userId, limit, before); } + async countSavedMessages(userId: UserID): Promise { + return this.savedMessageRepository.countSavedMessages(userId); + } + async createSavedMessage(userId: UserID, channelId: ChannelID, messageId: MessageID): Promise { return this.savedMessageRepository.createSavedMessage(userId, channelId, messageId); } diff --git a/fluxer_api/src/api/user/repositories/UserRepository.ts b/fluxer_api/src/api/user/repositories/UserRepository.ts index 9cb3703ee..fb5853c23 100644 --- a/fluxer_api/src/api/user/repositories/UserRepository.ts +++ b/fluxer_api/src/api/user/repositories/UserRepository.ts @@ -605,6 +605,10 @@ export class UserRepository implements IUserRepositoryAggregate { return this.contentRepo.listSavedMessages(userId, limit, before); } + async countSavedMessages(userId: UserID): Promise { + return this.contentRepo.countSavedMessages(userId); + } + async createSavedMessage(userId: UserID, channelId: ChannelID, messageId: MessageID): Promise { return this.contentRepo.createSavedMessage(userId, channelId, messageId); } diff --git a/fluxer_api/src/api/user/services/UserContentService.test.ts b/fluxer_api/src/api/user/services/UserContentService.test.ts index 35ba4d075..34b78d227 100644 --- a/fluxer_api/src/api/user/services/UserContentService.test.ts +++ b/fluxer_api/src/api/user/services/UserContentService.test.ts @@ -1,11 +1,17 @@ // SPDX-License-Identifier: AGPL-3.0-or-later +import {APIErrorCodes} from '@fluxer/constants/src/ApiErrorCodes'; import {MessageTypes} from '@fluxer/constants/src/ChannelConstants'; import {UnknownChannelError} from '@fluxer/errors/src/domains/channel/UnknownChannelError'; import {UnknownMessageError} from '@fluxer/errors/src/domains/channel/UnknownMessageError'; +import {AccessDeniedError} from '@fluxer/errors/src/domains/core/AccessDeniedError'; +import {BadGatewayError} from '@fluxer/errors/src/domains/core/BadGatewayError'; +import {MaxBookmarksError} from '@fluxer/errors/src/domains/core/MaxBookmarksError'; import {MissingPermissionsError} from '@fluxer/errors/src/domains/core/MissingPermissionsError'; import {UnknownGuildError} from '@fluxer/errors/src/domains/guild/UnknownGuildError'; -import {describe, expect, it} from 'vitest'; +import {NsfwContentRequiresAgeVerificationError} from '@fluxer/errors/src/domains/moderation/NsfwContentRequiresAgeVerificationError'; +import type {LimitConfigSnapshot} from '@fluxer/limits/src/LimitTypes'; +import {describe, expect, it, vi} from 'vitest'; import type {ApiContext} from '../../ApiContext'; import type {ChannelID, MessageID, UserID} from '../../BrandedTypes'; import {createChannelID, createMessageID, createUserID} from '../../BrandedTypes'; @@ -14,6 +20,7 @@ import type {ChannelService} from '../../channel/services/ChannelService'; import type {KVBulkMessageDeletionQueueService} from '../../infrastructure/KVBulkMessageDeletionQueueService'; import type {UserCacheService} from '../../infrastructure/UserCacheService'; import type {LimitConfigService} from '../../limits/LimitConfigService'; +import type {RequestCache} from '../../middleware/RequestCacheMiddleware'; import {Message} from '../../models/Message'; import {UserContentService, UserContentServiceTestHooks} from './UserContentService'; @@ -29,6 +36,11 @@ describe('isUnreachableEntityError', () => { expect(isUnreachableEntityError(new MissingPermissionsError())).toBe(true); }); + it('treats an age gate and an unresolved membership as unreachable', () => { + expect(isUnreachableEntityError(new AccessDeniedError())).toBe(true); + expect(isUnreachableEntityError(new NsfwContentRequiresAgeVerificationError())).toBe(true); + }); + it('leaves a deleted message to the delete path instead of marking it unavailable', () => { expect(isUnreachableEntityError(new UnknownMessageError())).toBe(false); }); @@ -78,10 +90,12 @@ function makeMessage(channelId: ChannelID, messageId: MessageID): Message { function createUserContentService({ entries, readable, + stored, failures = new Map(), }: { entries: Array<{channelId: ChannelID; messageId: MessageID}>; readable: Array<{channelId: ChannelID; messageId: MessageID}>; + stored?: Array<{channelId: ChannelID; messageId: MessageID}>; failures?: Map; }) { const batchCalls: Array = []; @@ -115,11 +129,20 @@ function createUserContentService({ }, }, }; + const storedKeys = new Set( + (stored ?? readable).map((entry) => `${entry.channelId.toString()}:${entry.messageId.toString()}`), + ); + const channelRepository = { + messages: { + getMessage: async (channelId: ChannelID, messageId: MessageID) => + storedKeys.has(`${channelId.toString()}:${messageId.toString()}`) ? makeMessage(channelId, messageId) : null, + }, + }; const service = new UserContentService( {services: {users: userRepository, gateway: {}, worker: {}, snowflake: {}}} as unknown as ApiContext, {} as unknown as UserCacheService, channelService as unknown as ChannelService, - {} as unknown as IChannelRepository, + channelRepository as unknown as IChannelRepository, {} as unknown as KVBulkMessageDeletionQueueService, {} as unknown as LimitConfigService, ); @@ -243,12 +266,38 @@ describe('getSavedMessages', () => { ]); }); - it('deletes and drops a saved message the batch read could not return', async () => { + it('keeps a saved message the batch read could not return while the message still exists', async () => { const entries = [ {channelId: CHANNEL_A, messageId: createMessageID(11n)}, {channelId: CHANNEL_A, messageId: createMessageID(12n)}, ]; - const {service, deletedSavedMessageIds} = createUserContentService({entries, readable: [entries[0]]}); + const {service, deletedSavedMessageIds} = createUserContentService({ + entries, + readable: [entries[0]], + stored: entries, + }); + + const saved = await service.getSavedMessages({userId: VIEWER_ID, limit: 50}); + + expect( + saved.map((entry) => ({id: entry.messageId.toString(), status: entry.status, hasMessage: entry.message != null})), + ).toEqual([ + {id: '12', status: 'missing_permissions', hasMessage: false}, + {id: '11', status: 'available', hasMessage: true}, + ]); + expect(deletedSavedMessageIds).toEqual([]); + }); + + it('deletes a saved message the repository no longer holds', async () => { + const entries = [ + {channelId: CHANNEL_A, messageId: createMessageID(11n)}, + {channelId: CHANNEL_A, messageId: createMessageID(12n)}, + ]; + const {service, deletedSavedMessageIds} = createUserContentService({ + entries, + readable: [entries[0]], + stored: [entries[0]], + }); const saved = await service.getSavedMessages({userId: VIEWER_ID, limit: 50}); @@ -256,3 +305,192 @@ describe('getSavedMessages', () => { expect(deletedSavedMessageIds).toEqual(['12']); }); }); + +describe('gateway dispatches after the write', () => { + function createDispatchingService(dispatchPresence: () => Promise) { + const createdSavedMessageIds: Array = []; + const deletedSavedMessageIds: Array = []; + const deletedRecentMentionIds: Array = []; + const message = makeMessage(CHANNEL_A, createMessageID(31n)); + const userRepository = { + findUnique: async () => ({ + isBot: false, + premiumType: 0, + premiumUntil: null, + premiumGiftExtensionEndsAt: null, + premiumWillCancel: false, + premiumGraceEndsAt: null, + flags: 0n, + premiumFlags: 0, + traits: null, + }), + listSavedMessages: async () => [], + countSavedMessages: async () => 0, + createSavedMessage: async (_userId: UserID, _channelId: ChannelID, messageId: MessageID) => { + createdSavedMessageIds.push(messageId.toString()); + }, + deleteSavedMessage: async (_userId: UserID, messageId: MessageID) => { + deletedSavedMessageIds.push(messageId.toString()); + }, + getRecentMention: async (_userId: UserID, messageId: MessageID) => ({messageId}), + deleteRecentMention: async (mention: {messageId: MessageID}) => { + deletedRecentMentionIds.push(mention.messageId.toString()); + }, + }; + const channelService = { + channelData: {auth: {getChannelAuthenticated: async () => ({})}}, + messages: {retrieval: {getMessage: async () => message}}, + }; + const service = new UserContentService( + { + services: {users: userRepository, gateway: {dispatchPresence}, worker: {}, snowflake: {}}, + } as unknown as ApiContext, + {} as unknown as UserCacheService, + channelService as unknown as ChannelService, + {} as unknown as IChannelRepository, + {} as unknown as KVBulkMessageDeletionQueueService, + {getConfigSnapshot: () => null} as unknown as LimitConfigService, + ); + return {service, message, createdSavedMessageIds, deletedSavedMessageIds, deletedRecentMentionIds}; + } + + function saveMessageArgs(messageId: MessageID) { + return { + userId: VIEWER_ID, + channelId: CHANNEL_A, + messageId, + userCacheService: {} as unknown as UserCacheService, + requestCache: {} as unknown as RequestCache, + }; + } + + it('keeps the saved message when SAVED_MESSAGE_CREATE fails to publish', async () => { + const {service, message, createdSavedMessageIds} = createDispatchingService(async () => { + throw new BadGatewayError(); + }); + vi.spyOn(service, 'buildMessageResponsesForUser').mockResolvedValue([]); + + await expect(service.saveMessage(saveMessageArgs(message.id))).resolves.toBeUndefined(); + + expect(createdSavedMessageIds).toEqual([message.id.toString()]); + }); + + it('still surfaces a failure of the message response build', async () => { + const {service, message} = createDispatchingService(async () => {}); + vi.spyOn(service, 'buildMessageResponsesForUser').mockRejectedValue(new Error('database is on fire')); + + await expect(service.saveMessage(saveMessageArgs(message.id))).rejects.toThrow('database is on fire'); + }); + + it('keeps the deletion when SAVED_MESSAGE_DELETE fails to publish', async () => { + const {service, deletedSavedMessageIds} = createDispatchingService(async () => { + throw new BadGatewayError(); + }); + + await expect(service.unsaveMessage({userId: VIEWER_ID, messageId: createMessageID(31n)})).resolves.toBeUndefined(); + + expect(deletedSavedMessageIds).toEqual(['31']); + }); + + it('keeps the deletion when RECENT_MENTION_DELETE fails to publish', async () => { + const {service, deletedRecentMentionIds} = createDispatchingService(async () => { + throw new BadGatewayError(); + }); + + await expect( + service.deleteRecentMention({userId: VIEWER_ID, messageId: createMessageID(31n)}), + ).resolves.toBeUndefined(); + + expect(deletedRecentMentionIds).toEqual(['31']); + }); +}); + +describe('bookmark ceiling', () => { + function createBookmarkLimitedService({ + savedMessageCount, + maxBookmarks, + }: { + savedMessageCount: number; + maxBookmarks: number; + }) { + const createdSavedMessageIds: Array = []; + const message = makeMessage(CHANNEL_A, createMessageID(41n)); + const userRepository = { + findUnique: async () => ({ + isBot: false, + premiumType: 0, + premiumUntil: null, + premiumGiftExtensionEndsAt: null, + premiumWillCancel: false, + premiumGraceEndsAt: null, + flags: 0n, + premiumFlags: 0, + traits: null, + }), + countSavedMessages: async () => savedMessageCount, + listSavedMessages: async () => { + throw new Error('the ceiling check must not page through saved messages'); + }, + createSavedMessage: async (_userId: UserID, _channelId: ChannelID, messageId: MessageID) => { + createdSavedMessageIds.push(messageId.toString()); + }, + }; + const channelService = { + channelData: {auth: {getChannelAuthenticated: async () => ({})}}, + messages: {retrieval: {getMessage: async () => message}}, + }; + const snapshot: LimitConfigSnapshot = { + traitDefinitions: [], + rules: [{id: 'default', limits: {max_bookmarks: maxBookmarks}}], + }; + const service = new UserContentService( + { + services: {users: userRepository, gateway: {dispatchPresence: async () => {}}, worker: {}, snowflake: {}}, + } as unknown as ApiContext, + {} as unknown as UserCacheService, + channelService as unknown as ChannelService, + {} as unknown as IChannelRepository, + {} as unknown as KVBulkMessageDeletionQueueService, + {getConfigSnapshot: () => snapshot} as unknown as LimitConfigService, + ); + return {service, message, createdSavedMessageIds}; + } + + function saveMessageArgs(messageId: MessageID) { + return { + userId: VIEWER_ID, + channelId: CHANNEL_A, + messageId, + userCacheService: {} as unknown as UserCacheService, + requestCache: {} as unknown as RequestCache, + }; + } + + it('enforces a configured max_bookmarks above 1000', async () => { + const {service, message, createdSavedMessageIds} = createBookmarkLimitedService({ + savedMessageCount: 1002, + maxBookmarks: 1002, + }); + vi.spyOn(service, 'buildMessageResponsesForUser').mockResolvedValue([]); + + const error = await service.saveMessage(saveMessageArgs(message.id)).catch((thrown: unknown) => thrown); + + expect(error).toBeInstanceOf(MaxBookmarksError); + expect((error as MaxBookmarksError).status).toBe(400); + expect((error as MaxBookmarksError).code).toBe(APIErrorCodes.MAX_BOOKMARKS); + expect((error as MaxBookmarksError).data?.max_bookmarks).toBe(1002); + expect(createdSavedMessageIds).toEqual([]); + }); + + it('still saves below a configured max_bookmarks above 1000', async () => { + const {service, message, createdSavedMessageIds} = createBookmarkLimitedService({ + savedMessageCount: 1001, + maxBookmarks: 1002, + }); + vi.spyOn(service, 'buildMessageResponsesForUser').mockResolvedValue([]); + + await expect(service.saveMessage(saveMessageArgs(message.id))).resolves.toBeUndefined(); + + expect(createdSavedMessageIds).toEqual([message.id.toString()]); + }); +}); diff --git a/fluxer_api/src/api/user/services/UserContentService.ts b/fluxer_api/src/api/user/services/UserContentService.ts index fbc480889..b82a23a3d 100644 --- a/fluxer_api/src/api/user/services/UserContentService.ts +++ b/fluxer_api/src/api/user/services/UserContentService.ts @@ -6,6 +6,7 @@ import {MAX_BOOKMARKS_NON_PREMIUM} from '@fluxer/constants/src/LimitConstants'; import {ValidationErrorCodes} from '@fluxer/constants/src/ValidationErrorCodes'; import {UnknownChannelError} from '@fluxer/errors/src/domains/channel/UnknownChannelError'; import {UnknownMessageError} from '@fluxer/errors/src/domains/channel/UnknownMessageError'; +import {AccessDeniedError} from '@fluxer/errors/src/domains/core/AccessDeniedError'; import {InputValidationError} from '@fluxer/errors/src/domains/core/InputValidationError'; import {MaxBookmarksError} from '@fluxer/errors/src/domains/core/MaxBookmarksError'; import {MissingPermissionsError} from '@fluxer/errors/src/domains/core/MissingPermissionsError'; @@ -13,6 +14,7 @@ import {UnknownGuildError} from '@fluxer/errors/src/domains/guild/UnknownGuildEr import {HarvestExpiredError} from '@fluxer/errors/src/domains/moderation/HarvestExpiredError'; import {HarvestFailedError} from '@fluxer/errors/src/domains/moderation/HarvestFailedError'; import {HarvestNotReadyError} from '@fluxer/errors/src/domains/moderation/HarvestNotReadyError'; +import {NsfwContentRequiresAgeVerificationError} from '@fluxer/errors/src/domains/moderation/NsfwContentRequiresAgeVerificationError'; import {UnknownHarvestError} from '@fluxer/errors/src/domains/moderation/UnknownHarvestError'; import {UnknownUserError} from '@fluxer/errors/src/domains/user/UnknownUserError'; import type {MessageResponse} from '@fluxer/schema/src/domains/message/MessageResponseSchemas'; @@ -116,7 +118,9 @@ function normalizeProviderEnvironment( const isUnreachableEntityError = (error: unknown): boolean => error instanceof MissingPermissionsError || error instanceof UnknownChannelError || - error instanceof UnknownGuildError; + error instanceof UnknownGuildError || + error instanceof AccessDeniedError || + error instanceof NsfwContentRequiresAgeVerificationError; export const UserContentServiceTestHooks = {isUnreachableEntityError}; @@ -243,7 +247,17 @@ export class UserContentService { } const message = this.pickMessage(messagesByChannel, savedMessage); if (!message) { - staleMessageIds.push(savedMessage.messageId); + const stored = await this.channelRepository.messages.getMessage(savedMessage.channelId, savedMessage.messageId); + if (!stored) { + staleMessageIds.push(savedMessage.messageId); + continue; + } + results.push({ + channelId: savedMessage.channelId, + messageId: savedMessage.messageId, + status: 'missing_permissions', + message: null, + }); continue; } results.push({ @@ -274,7 +288,7 @@ export class UserContentService { if (!user) { throw new UnknownUserError(); } - const savedMessages = await this.userRepository.listSavedMessages(userId, 1000); + const savedMessageCount = await this.userRepository.countSavedMessages(userId); const ctx = createLimitMatchContext({user}); const maxBookmarks = resolveLimitSafe( this.limitConfigService.getConfigSnapshot(), @@ -282,7 +296,7 @@ export class UserContentService { 'max_bookmarks', MAX_BOOKMARKS_NON_PREMIUM, ); - if (savedMessages.length >= maxBookmarks) { + if (savedMessageCount >= maxBookmarks) { throw new MaxBookmarksError({maxBookmarks}); } await this.channelService.channelData.auth.getChannelAuthenticated({userId, channelId}); @@ -514,12 +528,12 @@ export class UserContentService { if (!harvest) { throw new UnknownHarvestError(); } - if (!harvest.completedAt || !harvest.storageKey) { - throw new HarvestNotReadyError(); - } if (harvest.failedAt) { throw new HarvestFailedError(); } + if (!harvest.completedAt || !harvest.storageKey) { + throw new HarvestNotReadyError(); + } if (harvest.downloadUrlExpiresAt && harvest.downloadUrlExpiresAt < new Date()) { throw new HarvestExpiredError(); } @@ -696,11 +710,19 @@ export class UserContentService { } async dispatchRecentMentionDelete({userId, messageId}: {userId: UserID; messageId: MessageID}): Promise { - await this.gatewayService.dispatchPresence({ - userId, - event: 'RECENT_MENTION_DELETE', - data: {message_id: messageId.toString()}, - }); + await this.gatewayService + .dispatchPresence({ + userId, + event: 'RECENT_MENTION_DELETE', + data: {message_id: messageId.toString()}, + }) + .catch((error) => { + Logger.error( + {userId: userId.toString(), messageId: messageId.toString(), error}, + 'Failed to dispatch RECENT_MENTION_DELETE', + ); + return null; + }); } async dispatchSavedMessageCreate({ @@ -712,11 +734,20 @@ export class UserContentService { userCacheService: UserCacheService; requestCache: RequestCache; }): Promise { - await this.gatewayService.dispatchPresence({ - userId, - event: 'SAVED_MESSAGE_CREATE', - data: (await this.buildMessageResponsesForUser(userId, [message]))[0], - }); + const data = (await this.buildMessageResponsesForUser(userId, [message]))[0]; + await this.gatewayService + .dispatchPresence({ + userId, + event: 'SAVED_MESSAGE_CREATE', + data, + }) + .catch((error) => { + Logger.error( + {userId: userId.toString(), messageId: message.id.toString(), error}, + 'Failed to dispatch SAVED_MESSAGE_CREATE', + ); + return null; + }); } async buildMessageResponsesForUser(userId: UserID, messages: Array): Promise> { @@ -734,10 +765,18 @@ export class UserContentService { } async dispatchSavedMessageDelete({userId, messageId}: {userId: UserID; messageId: MessageID}): Promise { - await this.gatewayService.dispatchPresence({ - userId, - event: 'SAVED_MESSAGE_DELETE', - data: {message_id: messageId.toString()}, - }); + await this.gatewayService + .dispatchPresence({ + userId, + event: 'SAVED_MESSAGE_DELETE', + data: {message_id: messageId.toString()}, + }) + .catch((error) => { + Logger.error( + {userId: userId.toString(), messageId: messageId.toString(), error}, + 'Failed to dispatch SAVED_MESSAGE_DELETE', + ); + return null; + }); } } diff --git a/fluxer_api/src/api/user/tests/HarvestRetryClearsFailure.test.ts b/fluxer_api/src/api/user/tests/HarvestRetryClearsFailure.test.ts new file mode 100644 index 000000000..e7a0ee96a --- /dev/null +++ b/fluxer_api/src/api/user/tests/HarvestRetryClearsFailure.test.ts @@ -0,0 +1,48 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later + +import {beforeEach, describe, expect, test} from 'vitest'; +import {createTestAccount} from '../../auth/tests/AuthTestUtils'; +import {type ApiTestHarness, createApiTestHarness} from '../../test/ApiTestHarness'; +import { + expectHarvestDownloadFailsWithError, + fetchHarvestDownload, + findHarvest, + markHarvestCompleted, + markHarvestFailed, + markHarvestStarted, + requestHarvest, +} from './HarvestTestUtils'; + +describe('Harvest Retry Clears Failure', () => { + let harness: ApiTestHarness; + beforeEach(async () => { + harness = await createApiTestHarness(); + }); + test('a retry clears the previous failure and a completed retry reports completed', async () => { + const account = await createTestAccount(harness); + const {harvest_id} = await requestHarvest(harness, account.token); + await markHarvestFailed(account.userId, harvest_id, 'harvest exploded'); + const failed = await findHarvest(account.userId, harvest_id); + expect(failed?.getStatus()).toBe('failed'); + expect(failed?.errorMessage).toBe('harvest exploded'); + await markHarvestStarted(account.userId, harvest_id); + const retrying = await findHarvest(account.userId, harvest_id); + expect(retrying?.getStatus()).toBe('processing'); + expect(retrying?.failedAt).toBeNull(); + expect(retrying?.errorMessage).toBeNull(); + const validTime = new Date(Date.now() + 6 * 24 * 60 * 60 * 1000); + await markHarvestCompleted(account.userId, harvest_id, validTime); + const completed = await findHarvest(account.userId, harvest_id); + expect(completed?.getStatus()).toBe('completed'); + expect(completed?.failedAt).toBeNull(); + expect(completed?.errorMessage).toBeNull(); + const download = await fetchHarvestDownload(harness, account.token, harvest_id); + expect(download.download_url).not.toBe(''); + }); + test('download reports the failure rather than unreadiness when the latest attempt failed', async () => { + const account = await createTestAccount(harness); + const {harvest_id} = await requestHarvest(harness, account.token); + await markHarvestFailed(account.userId, harvest_id, 'harvest exploded'); + await expectHarvestDownloadFailsWithError(harness, account.token, harvest_id, 'HARVEST_FAILED'); + }); +}); diff --git a/fluxer_api/src/api/user/tests/HarvestTestUtils.ts b/fluxer_api/src/api/user/tests/HarvestTestUtils.ts index 2ab5fc4fe..4bbf140f8 100644 --- a/fluxer_api/src/api/user/tests/HarvestTestUtils.ts +++ b/fluxer_api/src/api/user/tests/HarvestTestUtils.ts @@ -5,6 +5,7 @@ import {expect} from 'vitest'; import {createUserID} from '../../BrandedTypes'; import type {ApiTestHarness} from '../../test/ApiTestHarness'; import {createBuilder} from '../../test/TestRequestBuilder'; +import type {UserHarvest} from '../UserHarvestModel'; import {UserHarvestRepository} from '../UserHarvestRepository'; interface HarvestRequestResponse { @@ -48,3 +49,18 @@ export async function markHarvestCompleted(userId: string, harvestId: string, ex const harvestIdTyped = BigInt(harvestId); await harvestRepository.markAsCompleted(userIdTyped, harvestIdTyped, `test/${harvestId}.zip`, 1024n, expiresAt); } + +export async function markHarvestFailed(userId: string, harvestId: string, errorMessage: string): Promise { + const harvestRepository = new UserHarvestRepository(); + await harvestRepository.markAsFailed(createUserID(BigInt(userId)), BigInt(harvestId), errorMessage); +} + +export async function markHarvestStarted(userId: string, harvestId: string): Promise { + const harvestRepository = new UserHarvestRepository(); + await harvestRepository.markAsStarted(createUserID(BigInt(userId)), BigInt(harvestId)); +} + +export async function findHarvest(userId: string, harvestId: string): Promise { + const harvestRepository = new UserHarvestRepository(); + return harvestRepository.findByUserAndHarvestId(createUserID(BigInt(userId)), BigInt(harvestId)); +}