fix(api): correct user content, read state and harvest paths (#2501)

This commit is contained in:
Hampus
2026-09-06 15:30:08 +02:00
committed by GitHub
parent cc110b9f5a
commit 1f627c9cc5
14 changed files with 786 additions and 33 deletions
@@ -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<string> = [];
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);
});
});
@@ -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<void> {
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<void> {
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;
@@ -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'),
},
@@ -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<string, EntranceSound>();
failNextUpsert = false;
async listSounds(_userId: UserID): Promise<Array<EntranceSound>> {
return [...this.sounds.values()];
}
async getSound(_userId: UserID, soundId: EntranceSoundID): Promise<EntranceSound | null> {
return this.sounds.get(soundId.toString()) ?? null;
}
async upsertSound(sound: EntranceSound): Promise<EntranceSound> {
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<void> {
this.sounds.delete(soundId.toString());
}
async deleteSelectionsForSound(_userId: UserID, _soundId: EntranceSoundID): Promise<void> {}
}
function createService() {
const repository = new FakeEntranceSoundRepository();
const objects = new Set<string>();
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);
});
});
@@ -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<boolean> {
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<Array<EntranceSoundLibraryEntry>> {
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');
@@ -27,6 +27,7 @@ export interface IUserContentRepository {
deleteRecentMentions(mentions: Array<RecentMention>): Promise<void>;
deleteAllRecentMentions(userId: UserID): Promise<void>;
listSavedMessages(userId: UserID, limit?: number, before?: MessageID): Promise<Array<SavedMessage>>;
countSavedMessages(userId: UserID): Promise<number>;
createSavedMessage(userId: UserID, channelId: ChannelID, messageId: MessageID): Promise<SavedMessage>;
deleteSavedMessage(userId: UserID, messageId: MessageID): Promise<void>;
deleteAllSavedMessages(userId: UserID): Promise<void>;
@@ -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);
});
});
@@ -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<number> {
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,
@@ -186,6 +186,10 @@ export class UserContentRepository implements IUserContentRepository {
return this.savedMessageRepository.listSavedMessages(userId, limit, before);
}
async countSavedMessages(userId: UserID): Promise<number> {
return this.savedMessageRepository.countSavedMessages(userId);
}
async createSavedMessage(userId: UserID, channelId: ChannelID, messageId: MessageID): Promise<SavedMessage> {
return this.savedMessageRepository.createSavedMessage(userId, channelId, messageId);
}
@@ -605,6 +605,10 @@ export class UserRepository implements IUserRepositoryAggregate {
return this.contentRepo.listSavedMessages(userId, limit, before);
}
async countSavedMessages(userId: UserID): Promise<number> {
return this.contentRepo.countSavedMessages(userId);
}
async createSavedMessage(userId: UserID, channelId: ChannelID, messageId: MessageID): Promise<SavedMessage> {
return this.contentRepo.createSavedMessage(userId, channelId, messageId);
}
@@ -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<string, Error>(),
}: {
entries: Array<{channelId: ChannelID; messageId: MessageID}>;
readable: Array<{channelId: ChannelID; messageId: MessageID}>;
stored?: Array<{channelId: ChannelID; messageId: MessageID}>;
failures?: Map<string, Error>;
}) {
const batchCalls: Array<ChannelBatchCall> = [];
@@ -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<void>) {
const createdSavedMessageIds: Array<string> = [];
const deletedSavedMessageIds: Array<string> = [];
const deletedRecentMentionIds: Array<string> = [];
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<string> = [];
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()]);
});
});
@@ -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<void> {
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<void> {
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<Message>): Promise<Array<MessageResponse>> {
@@ -734,10 +765,18 @@ export class UserContentService {
}
async dispatchSavedMessageDelete({userId, messageId}: {userId: UserID; messageId: MessageID}): Promise<void> {
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;
});
}
}
@@ -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');
});
});
@@ -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<void> {
const harvestRepository = new UserHarvestRepository();
await harvestRepository.markAsFailed(createUserID(BigInt(userId)), BigInt(harvestId), errorMessage);
}
export async function markHarvestStarted(userId: string, harvestId: string): Promise<void> {
const harvestRepository = new UserHarvestRepository();
await harvestRepository.markAsStarted(createUserID(BigInt(userId)), BigInt(harvestId));
}
export async function findHarvest(userId: string, harvestId: string): Promise<UserHarvest | null> {
const harvestRepository = new UserHarvestRepository();
return harvestRepository.findByUserAndHarvestId(createUserID(BigInt(userId)), BigInt(harvestId));
}