fix(message): bound batched message response requests by size (#2252)

This commit is contained in:
Hampus
2026-08-31 16:13:36 +02:00
committed by GitHub
parent 5036ac3efa
commit 8c85cce75c
2 changed files with 169 additions and 43 deletions
@@ -6,7 +6,7 @@ import type {NatsConnection} from 'nats';
import {describe, expect, it} from 'vitest';
import {createChannelID, createMessageID, createUserID} from '../../../BrandedTypes';
import {Message} from '../../../models/Message';
import {MessageResponseDataService} from './MessageResponseDataService';
import {MESSAGE_BUILD_BATCH_MAX_BYTES, MessageResponseDataService} from './MessageResponseDataService';
const encoder = new TextEncoder();
const decoder = new TextDecoder();
@@ -27,48 +27,59 @@ class FakeConnectionManager implements INatsConnectionManager {
getConnection(): NatsConnection {
return {
request: async (_subject: string, data: Uint8Array, options?: {timeout?: number}) => {
this.payloads.push(JSON.parse(decoder.decode(data)) as Record<string, unknown>);
const payload = JSON.parse(decoder.decode(data)) as Record<string, unknown>;
this.payloads.push(payload);
this.timeouts.push(options?.timeout);
if (payload.op === 'BuildResponses') {
const messages = payload.messages as Array<{message_id: string}>;
return {
data: encoder.encode(
JSON.stringify({
FoundApiMany: messages.map((message) => fakeMessageResponse(message.message_id)),
}),
),
};
}
return {
data: encoder.encode(
JSON.stringify({
FoundApi: {
id: '2',
channel_id: '1',
author: {id: '3', username: 'author', discriminator: '0001', avatar: null, flags: 0},
type: MessageTypes.DEFAULT,
flags: 0,
content: '',
timestamp: '2026-01-01T00:00:00.000Z',
edited_timestamp: null,
pinned: false,
mention_everyone: false,
tts: false,
mentions: [],
mention_roles: [],
embeds: [],
attachments: [],
stickers: [],
},
}),
),
data: encoder.encode(JSON.stringify({FoundApi: fakeMessageResponse('2')})),
};
},
} as unknown as NatsConnection;
}
}
function makeMessage(): Message {
function fakeMessageResponse(messageId: string): Record<string, unknown> {
return {
id: messageId,
channel_id: '1',
author: {id: '3', username: 'author', discriminator: '0001', avatar: null, flags: 0},
type: MessageTypes.DEFAULT,
flags: 0,
content: '',
timestamp: '2026-01-01T00:00:00.000Z',
edited_timestamp: null,
pinned: false,
mention_everyone: false,
tts: false,
mentions: [],
mention_roles: [],
embeds: [],
attachments: [],
stickers: [],
};
}
function makeMessage(messageId: bigint = 2n, content: string = ''): Message {
return new Message({
channel_id: createChannelID(1n),
bucket: 0,
message_id: createMessageID(2n),
message_id: createMessageID(messageId),
author_id: createUserID(3n),
type: MessageTypes.DEFAULT,
webhook_id: null,
webhook_name: null,
webhook_avatar_hash: null,
content: '',
content,
edited_timestamp: null,
pinned_timestamp: null,
flags: 0,
@@ -88,6 +99,26 @@ function makeMessage(): Message {
});
}
const BASE_MESSAGE_ID = 1000000000000000000n;
const VIEWER_ID = createUserID(3n);
const ACCESS = {sourceGuildId: null, messageHistoryCutoff: null, canReadMessageHistory: true};
async function measureSerializedMessageBytes(): Promise<number> {
const connectionManager = new FakeConnectionManager();
const service = new MessageResponseDataService(connectionManager);
await service.buildMessages({userId: VIEWER_ID, messages: [makeMessage(BASE_MESSAGE_ID)], access: ACCESS});
const {messages} = connectionManager.payloads[0] as {messages: Array<unknown>};
return Buffer.byteLength(JSON.stringify(messages[0]));
}
function makeSizedMessage(index: number, totalBytes: number, baseBytes: number): Message {
return makeMessage(BASE_MESSAGE_ID + BigInt(index), 'a'.repeat(totalBytes - baseBytes));
}
function batchSizes(connectionManager: FakeConnectionManager): Array<number> {
return connectionManager.payloads.map((payload) => (payload as {messages: Array<unknown>}).messages.length);
}
describe('MessageResponseDataService', () => {
it('omits reactions from broadcast message response requests', async () => {
const connectionManager = new FakeConnectionManager();
@@ -133,4 +164,73 @@ describe('MessageResponseDataService', () => {
expect(connectionManager.timeouts[0]).toBe(6000);
expect(connectionManager.timeouts[0]).toBeGreaterThan(ROUTER_SHARD_REQUEST_TIMEOUT_MS);
});
it('sends no request when there are no messages to build', async () => {
const connectionManager = new FakeConnectionManager();
const service = new MessageResponseDataService(connectionManager);
const responses = await service.buildMessages({userId: VIEWER_ID, messages: [], access: ACCESS});
expect(responses).toEqual([]);
expect(connectionManager.payloads).toEqual([]);
});
it('keeps a batch that exactly fills the byte budget in one request', async () => {
const baseBytes = await measureSerializedMessageBytes();
const connectionManager = new FakeConnectionManager();
const service = new MessageResponseDataService(connectionManager);
const messages = [
makeSizedMessage(0, MESSAGE_BUILD_BATCH_MAX_BYTES / 2, baseBytes),
makeSizedMessage(1, MESSAGE_BUILD_BATCH_MAX_BYTES / 2, baseBytes),
];
const responses = await service.buildMessages({userId: VIEWER_ID, messages, access: ACCESS});
expect(batchSizes(connectionManager)).toEqual([2]);
expect(responses.map((response) => response.id)).toEqual(messages.map((message) => message.id.toString()));
});
it('splits into ordered batches once one more byte would cross the budget', async () => {
const baseBytes = await measureSerializedMessageBytes();
const connectionManager = new FakeConnectionManager();
const service = new MessageResponseDataService(connectionManager);
const messages = [
makeSizedMessage(0, MESSAGE_BUILD_BATCH_MAX_BYTES / 2, baseBytes),
makeSizedMessage(1, MESSAGE_BUILD_BATCH_MAX_BYTES / 2 + 1, baseBytes),
makeSizedMessage(2, baseBytes, baseBytes),
];
const responses = await service.buildMessages({userId: VIEWER_ID, messages, access: ACCESS});
expect(batchSizes(connectionManager)).toEqual([1, 2]);
expect(responses.map((response) => response.id)).toEqual(messages.map((message) => message.id.toString()));
});
it('sends a single message that exceeds the budget on its own', async () => {
const baseBytes = await measureSerializedMessageBytes();
const connectionManager = new FakeConnectionManager();
const service = new MessageResponseDataService(connectionManager);
const messages = [
makeSizedMessage(0, MESSAGE_BUILD_BATCH_MAX_BYTES + 1, baseBytes),
makeSizedMessage(1, baseBytes, baseBytes),
];
const responses = await service.buildMessages({userId: VIEWER_ID, messages, access: ACCESS});
expect(batchSizes(connectionManager)).toEqual([1, 1]);
expect(responses.map((response) => response.id)).toEqual(messages.map((message) => message.id.toString()));
});
it('keeps a full page of ordinary pins in a single request', async () => {
const connectionManager = new FakeConnectionManager();
const service = new MessageResponseDataService(connectionManager);
const messages = Array.from({length: 50}, (_unused, index) =>
makeMessage(BASE_MESSAGE_ID + BigInt(index), 'a'.repeat(120)),
);
const responses = await service.buildMessages({userId: VIEWER_ID, messages, access: ACCESS});
expect(batchSizes(connectionManager)).toEqual([50]);
expect(responses.map((response) => response.id)).toEqual(messages.map((message) => message.id.toString()));
});
});
@@ -15,6 +15,7 @@ import {isJsonRecord, parseJsonRecord, parseJsonWithGuard} from '../../../utils/
const MESSAGE_RESPONSE_SERVICE_SUBJECT = 'svc.messages';
const MESSAGE_RESPONSE_SERVICE_TIMEOUT_MS = 6000;
export const MESSAGE_BUILD_BATCH_MAX_BYTES = 262_144;
let messageResponseDataService: MessageResponseDataService | undefined;
let injectedMessageResponseDataService: MessageResponseDataService | undefined;
@@ -230,23 +231,27 @@ export class MessageResponseDataService {
includeReactions?: boolean;
}): Promise<Array<MessageResponse>> {
if (params.messages.length === 0) return [];
const response = await this.request({
op: 'BuildResponses',
messages: params.messages.map(serializeMessageForService),
viewer_user_id: params.userId.toString(),
source_guild_id: params.access.sourceGuildId?.toString(),
message_history_cutoff_ms: params.access.messageHistoryCutoff
? new Date(params.access.messageHistoryCutoff).getTime()
: null,
can_read_message_history: params.access.canReadMessageHistory,
media_endpoint: Config.endpoints.media,
media_proxy_secret_key: Config.mediaProxy.secretKey,
include_reactions: params.includeReactions ?? true,
});
if (typeof response === 'object' && 'FoundApiMany' in response) {
return response.FoundApiMany;
const responses: Array<MessageResponse> = [];
for (const batch of chunkMessagesForService(params.messages)) {
const response = await this.request({
op: 'BuildResponses',
messages: batch,
viewer_user_id: params.userId.toString(),
source_guild_id: params.access.sourceGuildId?.toString(),
message_history_cutoff_ms: params.access.messageHistoryCutoff
? new Date(params.access.messageHistoryCutoff).getTime()
: null,
can_read_message_history: params.access.canReadMessageHistory,
media_endpoint: Config.endpoints.media,
media_proxy_secret_key: Config.mediaProxy.secretKey,
include_reactions: params.includeReactions ?? true,
});
if (typeof response !== 'object' || !('FoundApiMany' in response)) {
throw new Error(`[message-response-service] unexpected BuildResponses response: ${JSON.stringify(response)}`);
}
responses.push(...response.FoundApiMany);
}
throw new Error(`[message-response-service] unexpected BuildResponses response: ${JSON.stringify(response)}`);
return responses;
}
async buildMessagesForChannels(params: {
@@ -340,6 +345,27 @@ function serializeMessageForService(message: Message): Record<string, unknown> {
return row;
}
function chunkMessagesForService(messages: Array<Message>): Array<Array<Record<string, unknown>>> {
const batches: Array<Array<Record<string, unknown>>> = [];
let batch: Array<Record<string, unknown>> = [];
let batchBytes = 0;
for (const message of messages) {
const serialized = serializeMessageForService(message);
const bytes = Buffer.byteLength(JSON.stringify(serialized));
if (batch.length > 0 && batchBytes + bytes > MESSAGE_BUILD_BATCH_MAX_BYTES) {
batches.push(batch);
batch = [];
batchBytes = 0;
}
batch.push(serialized);
batchBytes += bytes;
}
if (batch.length > 0) {
batches.push(batch);
}
return batches;
}
function serializeValue(value: unknown): unknown {
if (value == null) return value;
if (typeof value === 'bigint') return value.toString();