mirror of
https://github.com/fluxerapp/fluxer
synced 2026-10-07 19:22:14 +09:00
chore: tidy request handling across services (#3168)
This commit is contained in:
@@ -56,8 +56,8 @@ async function verifySudoMode(
|
||||
const apiContext = ctx.get('apiContext');
|
||||
const credentials = await apiContext.services.users.listWebAuthnCredentials(user.id);
|
||||
const hasPasskeyCredentials = credentials.length > 0;
|
||||
const hasMfa = userHasMfa(user);
|
||||
const hasSudoCapability = userHasSudoCapability(user, hasPasskeyCredentials);
|
||||
const hasUsableMfa = userHasMfa(user) && hasSudoCapability;
|
||||
const issueSudoToken = options.issueSudoToken ?? hasSudoCapability;
|
||||
if (hasSudoCapability && ctx.get('sudoModeValid')) {
|
||||
const sudoToken = ctx.get('sudoModeToken') ?? ctx.req.header(SUDO_MODE_HEADER) ?? undefined;
|
||||
@@ -82,10 +82,10 @@ async function verifySudoMode(
|
||||
const sudoToken = issueSudoToken ? await sudoModeService.generateSudoToken(user.id) : undefined;
|
||||
return {verified: true, sudoToken, method: 'mfa'};
|
||||
}
|
||||
if (hasNoVerifiableCredential(user, hasMfa, hasPasskeyCredentials)) {
|
||||
if (hasNoVerifiableCredential(user, hasUsableMfa, hasPasskeyCredentials)) {
|
||||
return {verified: true, method: 'password'};
|
||||
}
|
||||
if (body.password && !hasMfa) {
|
||||
if (body.password && !hasUsableMfa) {
|
||||
if (!user.passwordHash) {
|
||||
throw InputValidationError.fromCode('password', ValidationErrorCodes.PASSWORD_NOT_SET);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,87 @@
|
||||
// SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
import {
|
||||
createAuthHarness,
|
||||
createTestAccount,
|
||||
createTotpSecret,
|
||||
generateTotpCode,
|
||||
type TestAccount,
|
||||
} from '@app/api/auth/tests/AuthTestUtils';
|
||||
import {setWebAuthnTwoFactor} from '@app/api/auth/tests/WebAuthnTestUtils';
|
||||
import {createUserID} from '@app/api/BrandedTypes';
|
||||
import {getUserRepository} from '@app/api/middleware/ServiceSingletons';
|
||||
import type {ApiTestHarness} from '@app/api/test/ApiTestHarness';
|
||||
import {HTTP_STATUS} from '@app/api/test/TestConstants';
|
||||
import {createBuilder} from '@app/api/test/TestRequestBuilder';
|
||||
import {UserAuthenticatorTypes} from '@fluxer/constants/src/UserConstants';
|
||||
import {afterAll, beforeAll, beforeEach, describe, expect, it} from 'vitest';
|
||||
|
||||
interface PrivateUserResponse {
|
||||
mfa_enabled: boolean;
|
||||
authenticator_types: Array<number>;
|
||||
}
|
||||
|
||||
interface SudoModeRequiredResponse {
|
||||
code: string;
|
||||
has_mfa?: boolean;
|
||||
}
|
||||
|
||||
async function setAuthenticatorTypes(account: TestAccount, types: Array<number>): Promise<void> {
|
||||
const users = getUserRepository();
|
||||
const user = (await users.findUnique(createUserID(BigInt(account.userId))))!;
|
||||
await users.patchUpsert(user.id, {authenticator_types: new Set<number>(types)}, user.toRow());
|
||||
}
|
||||
|
||||
describe('Sudo mode for accounts whose authenticator types list no usable factor', () => {
|
||||
let harness: ApiTestHarness;
|
||||
beforeAll(async () => {
|
||||
harness = await createAuthHarness();
|
||||
});
|
||||
beforeEach(async () => {
|
||||
await harness.reset();
|
||||
});
|
||||
afterAll(async () => {
|
||||
await harness?.shutdown();
|
||||
});
|
||||
|
||||
it('accepts the password and lets the account turn passkey two-factor off when no passkey remains', async () => {
|
||||
const account = await createTestAccount(harness);
|
||||
await setAuthenticatorTypes(account, [UserAuthenticatorTypes.WEBAUTHN]);
|
||||
const challenge = await createBuilder<SudoModeRequiredResponse>(harness, account.token)
|
||||
.put('/users/@me/mfa/webauthn/two-factor')
|
||||
.body({enabled: false})
|
||||
.expect(HTTP_STATUS.FORBIDDEN)
|
||||
.execute();
|
||||
expect(challenge.has_mfa).toBe(false);
|
||||
const disabled = await setWebAuthnTwoFactor(harness, account.token, false, {password: account.password});
|
||||
expect(disabled.user.authenticator_types).toEqual([]);
|
||||
const me = await createBuilder<PrivateUserResponse>(harness, account.token).get('/users/@me').execute();
|
||||
expect(me.mfa_enabled).toBe(false);
|
||||
});
|
||||
|
||||
it('accepts the password when the TOTP type is listed without a stored secret', async () => {
|
||||
const account = await createTestAccount(harness);
|
||||
await setAuthenticatorTypes(account, [UserAuthenticatorTypes.TOTP]);
|
||||
await createBuilder(harness, account.token)
|
||||
.put('/users/@me/mfa/webauthn/two-factor')
|
||||
.body({enabled: false, password: 'wrong-password'})
|
||||
.expect(HTTP_STATUS.BAD_REQUEST)
|
||||
.execute();
|
||||
await setWebAuthnTwoFactor(harness, account.token, false, {password: account.password});
|
||||
});
|
||||
|
||||
it('still refuses the password from an account with TOTP enrolled', async () => {
|
||||
const account = await createTestAccount(harness);
|
||||
const secret = createTotpSecret();
|
||||
await createBuilder(harness, account.token)
|
||||
.post('/users/@me/mfa/totp/enable')
|
||||
.body({secret, code: generateTotpCode(secret), password: account.password})
|
||||
.execute();
|
||||
const challenge = await createBuilder<SudoModeRequiredResponse>(harness, account.token)
|
||||
.put('/users/@me/mfa/webauthn/two-factor')
|
||||
.body({enabled: false, password: account.password})
|
||||
.expect(HTTP_STATUS.FORBIDDEN)
|
||||
.execute();
|
||||
expect(challenge.has_mfa).toBe(true);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,55 @@
|
||||
// SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
import {isTestSsoProvider, normalizeAndValidateSsoConfig} from '@app/api/instance/SsoConfigValidation';
|
||||
import {InputValidationError} from '@fluxer/errors/src/domains/core/InputValidationError';
|
||||
import {describe, expect, it} from 'vitest';
|
||||
|
||||
function ssoConfig(overrides: {authorizationUrl?: string | null; tokenUrl?: string | null} = {}) {
|
||||
return {
|
||||
enabled: true,
|
||||
enforced: false,
|
||||
issuer: null,
|
||||
authorizationUrl: 'test',
|
||||
tokenUrl: 'test',
|
||||
userInfoUrl: null,
|
||||
jwksUrl: null,
|
||||
clientId: 'client',
|
||||
allowedEmailDomains: [],
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
describe('isTestSsoProvider', () => {
|
||||
it('recognises the placeholder provider only when test mode is enabled', () => {
|
||||
expect(isTestSsoProvider({authorizationUrl: 'test', tokenUrl: null}, true)).toBe(true);
|
||||
expect(isTestSsoProvider({authorizationUrl: null, tokenUrl: 'test'}, true)).toBe(true);
|
||||
expect(isTestSsoProvider({authorizationUrl: 'test-provider', tokenUrl: null}, true)).toBe(true);
|
||||
});
|
||||
|
||||
it('never recognises the placeholder provider outside test mode', () => {
|
||||
expect(isTestSsoProvider({authorizationUrl: 'test', tokenUrl: null}, false)).toBe(false);
|
||||
expect(isTestSsoProvider({authorizationUrl: null, tokenUrl: 'test'}, false)).toBe(false);
|
||||
expect(isTestSsoProvider({authorizationUrl: 'test', tokenUrl: 'test'}, false)).toBe(false);
|
||||
expect(isTestSsoProvider({authorizationUrl: 'test-provider', tokenUrl: null}, false)).toBe(false);
|
||||
});
|
||||
});
|
||||
|
||||
describe('normalizeAndValidateSsoConfig placeholder endpoints', () => {
|
||||
it('accepts placeholder endpoints in test mode', async () => {
|
||||
const result = await normalizeAndValidateSsoConfig(ssoConfig(), {testModeEnabled: true});
|
||||
expect(result.ready).toBe(true);
|
||||
expect(result.authorizationUrl).toBe('test');
|
||||
});
|
||||
|
||||
it('rejects a placeholder authorization endpoint outside test mode', async () => {
|
||||
await expect(
|
||||
normalizeAndValidateSsoConfig(ssoConfig({tokenUrl: null}), {testModeEnabled: false}),
|
||||
).rejects.toBeInstanceOf(InputValidationError);
|
||||
});
|
||||
|
||||
it('rejects a placeholder token endpoint outside test mode', async () => {
|
||||
await expect(
|
||||
normalizeAndValidateSsoConfig(ssoConfig({authorizationUrl: null}), {testModeEnabled: false}),
|
||||
).rejects.toBeInstanceOf(InputValidationError);
|
||||
});
|
||||
});
|
||||
@@ -64,10 +64,13 @@ export function isTestSsoProvider(
|
||||
},
|
||||
testModeEnabled: boolean,
|
||||
): boolean {
|
||||
if (!testModeEnabled) {
|
||||
return false;
|
||||
}
|
||||
return (
|
||||
config.authorizationUrl === 'test' ||
|
||||
config.tokenUrl === 'test' ||
|
||||
(testModeEnabled && (config.authorizationUrl?.startsWith('test-') ?? false))
|
||||
(config.authorizationUrl?.startsWith('test-') ?? false)
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -95,6 +95,8 @@ const DSA_CODE_CHARSET = 'ABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789';
|
||||
const DSA_CODE_SEGMENT_LENGTH = 4;
|
||||
const DSA_CODE_SEPARATOR = '-';
|
||||
const DSA_TICKET_BYTES = 32;
|
||||
const DSA_EMAIL_SEND_RECIPIENT_MAX = 3;
|
||||
const DSA_EMAIL_SEND_RECIPIENT_WINDOW = ms('1 hour');
|
||||
|
||||
async function emitReportFiled(row: IARSubmissionRow, target: ReportTarget): Promise<void> {
|
||||
const key = row.reported_user_id ?? row.reporter_id;
|
||||
@@ -373,6 +375,20 @@ export class ReportService {
|
||||
|
||||
async sendDsaReportVerificationCode(email: string, locale: string | null = null): Promise<void> {
|
||||
const normalizedEmail = this.normalizeEmail(email);
|
||||
const recipientLimit = await this.rateLimitService.checkLimit({
|
||||
identifier: `dsa:report:email:send:recipient:${normalizedEmail}`,
|
||||
maxAttempts: DSA_EMAIL_SEND_RECIPIENT_MAX,
|
||||
windowMs: DSA_EMAIL_SEND_RECIPIENT_WINDOW,
|
||||
});
|
||||
if (!recipientLimit.allowed) {
|
||||
throw new RateLimitError({
|
||||
retryAfter: recipientLimit.retryAfter,
|
||||
retryAfterDecimal: recipientLimit.retryAfterDecimal,
|
||||
limit: recipientLimit.limit,
|
||||
resetTime: recipientLimit.resetTime,
|
||||
resetAfterDecimal: recipientLimit.resetAfterDecimal,
|
||||
});
|
||||
}
|
||||
const hasValidDns = await this.emailDnsValidationService.hasValidDnsRecords(normalizedEmail);
|
||||
if (!hasValidDns) {
|
||||
throw InputValidationError.fromCode('email', ValidationErrorCodes.EMAIL_DOMAIN_CANNOT_RECEIVE_MAIL);
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
// SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
import {clearTestEmails, createUniqueEmail, listTestEmails} from '@app/api/auth/tests/AuthTestUtils';
|
||||
import {type ApiTestHarness, createApiTestHarness} from '@app/api/test/ApiTestHarness';
|
||||
import {HTTP_STATUS} from '@app/api/test/TestConstants';
|
||||
import {createBuilderWithoutAuth} from '@app/api/test/TestRequestBuilder';
|
||||
import {afterEach, beforeEach, describe, expect, test} from 'vitest';
|
||||
|
||||
async function countVerificationEmailsTo(harness: ApiTestHarness, email: string): Promise<number> {
|
||||
const emails = await listTestEmails(harness);
|
||||
return emails.filter((sent) => sent.type === 'dsa_report_verification' && sent.to === email.toLowerCase()).length;
|
||||
}
|
||||
|
||||
describe('DSA report verification email recipient limit', () => {
|
||||
let harness: ApiTestHarness;
|
||||
beforeEach(async () => {
|
||||
harness = await createApiTestHarness();
|
||||
});
|
||||
afterEach(async () => {
|
||||
await harness?.shutdown();
|
||||
});
|
||||
|
||||
test('limits verification emails per address regardless of letter case', async () => {
|
||||
await clearTestEmails(harness);
|
||||
const email = createUniqueEmail('dsa-recipient');
|
||||
for (let attempt = 0; attempt < 3; attempt++) {
|
||||
await createBuilderWithoutAuth(harness)
|
||||
.post('/reports/dsa/email/send')
|
||||
.body({email})
|
||||
.expect(HTTP_STATUS.OK)
|
||||
.execute();
|
||||
}
|
||||
await createBuilderWithoutAuth(harness)
|
||||
.post('/reports/dsa/email/send')
|
||||
.body({email: email.toUpperCase()})
|
||||
.expect(429)
|
||||
.execute();
|
||||
expect(await countVerificationEmailsTo(harness, email)).toBe(3);
|
||||
});
|
||||
|
||||
test('keeps sending to other addresses after one address reaches its limit', async () => {
|
||||
await clearTestEmails(harness);
|
||||
const limited = createUniqueEmail('dsa-recipient');
|
||||
for (let attempt = 0; attempt < 3; attempt++) {
|
||||
await createBuilderWithoutAuth(harness).post('/reports/dsa/email/send').body({email: limited}).execute();
|
||||
}
|
||||
const other = createUniqueEmail('dsa-recipient');
|
||||
await createBuilderWithoutAuth(harness)
|
||||
.post('/reports/dsa/email/send')
|
||||
.body({email: other})
|
||||
.expect(HTTP_STATUS.OK)
|
||||
.execute();
|
||||
expect(await countVerificationEmailsTo(harness, other)).toBe(1);
|
||||
});
|
||||
});
|
||||
@@ -467,6 +467,7 @@ export function WebhookController(app: HonoApp) {
|
||||
app.post(
|
||||
'/webhooks/:webhook_id/:token/github',
|
||||
RateLimitMiddleware(RateLimitConfigs.WEBHOOK_GITHUB),
|
||||
BlockAppOriginMiddleware,
|
||||
OpenAPI({
|
||||
operationId: 'execute_github_webhook',
|
||||
summary: 'Execute GitHub webhook',
|
||||
@@ -494,6 +495,7 @@ export function WebhookController(app: HonoApp) {
|
||||
app.post(
|
||||
'/webhooks/:webhook_id/:token/slack',
|
||||
RateLimitMiddleware(RateLimitConfigs.WEBHOOK_EXECUTE),
|
||||
BlockAppOriginMiddleware,
|
||||
OpenAPI({
|
||||
operationId: 'execute_slack_webhook',
|
||||
summary: 'Execute Slack webhook',
|
||||
@@ -520,6 +522,7 @@ export function WebhookController(app: HonoApp) {
|
||||
app.post(
|
||||
'/webhooks/:webhook_id/:token/instatus',
|
||||
RateLimitMiddleware(RateLimitConfigs.WEBHOOK_INSTATUS),
|
||||
BlockAppOriginMiddleware,
|
||||
OpenAPI({
|
||||
operationId: 'execute_instatus_webhook',
|
||||
summary: 'Execute Instatus webhook',
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
// SPDX-License-Identifier: AGPL-3.0-or-later
|
||||
|
||||
import {createTestAccount} from '@app/api/auth/tests/AuthTestUtils';
|
||||
import {getConfig} from '@app/api/Config';
|
||||
import {createGuild} from '@app/api/guild/tests/GuildTestUtils';
|
||||
import {type ApiTestHarness, createApiTestHarness} from '@app/api/test/ApiTestHarness';
|
||||
import {HTTP_STATUS} from '@app/api/test/TestConstants';
|
||||
import {createBuilderWithoutAuth} from '@app/api/test/TestRequestBuilder';
|
||||
import {createWebhook} from '@app/api/webhook/tests/WebhookTestUtils';
|
||||
import {APIErrorCodes} from '@fluxer/constants/src/ApiErrorCodes';
|
||||
import {afterEach, beforeEach, describe, it} from 'vitest';
|
||||
|
||||
const TOKEN_ROUTE_SUFFIXES = ['', '/github', '/slack', '/instatus'];
|
||||
|
||||
describe('Webhook token routes and the web app origin', () => {
|
||||
let harness: ApiTestHarness;
|
||||
beforeEach(async () => {
|
||||
harness = await createApiTestHarness();
|
||||
});
|
||||
afterEach(async () => {
|
||||
await harness?.shutdown();
|
||||
});
|
||||
|
||||
for (const suffix of TOKEN_ROUTE_SUFFIXES) {
|
||||
it(`refuses POST /webhooks/:webhook_id/:token${suffix} from the web app origin`, async () => {
|
||||
const owner = await createTestAccount(harness);
|
||||
const guild = await createGuild(harness, owner.token, 'App Origin Guild');
|
||||
const webhook = await createWebhook(harness, guild.system_channel_id!, owner.token, 'App Origin Webhook');
|
||||
await createBuilderWithoutAuth(harness)
|
||||
.post(`/webhooks/${webhook.id}/${webhook.token}${suffix}`)
|
||||
.header('origin', getConfig().endpoints.webAppOrigins[0]!)
|
||||
.body({content: 'hello'})
|
||||
.expect(HTTP_STATUS.FORBIDDEN, APIErrorCodes.INVALID_API_ORIGIN)
|
||||
.execute();
|
||||
});
|
||||
}
|
||||
});
|
||||
@@ -8,6 +8,7 @@
|
||||
new_context/1,
|
||||
compress/2,
|
||||
decompress/2,
|
||||
decompress/3,
|
||||
parse_compression/1,
|
||||
parse_compression/2,
|
||||
close_context/1,
|
||||
@@ -20,6 +21,7 @@
|
||||
-define(ZSTD_STREAM_BUFFER_SIZE, 64 * 1024).
|
||||
-define(ZSTD_COMPRESSION_LEVEL, 3).
|
||||
-define(ZSTD_DECOMPRESS_WINDOW_LOG_MAX, 23).
|
||||
-define(ZSTD_STREAM_MAX_CHUNKS, 1000).
|
||||
-define(ZSTD_AVAILABLE_KEY, {?MODULE, zstd_available}).
|
||||
-define(ZSTD_STREAM_AVAILABLE_KEY, {?MODULE, zstd_stream_available}).
|
||||
|
||||
@@ -95,14 +97,19 @@ compress(Data, #{type := zstd_stream} = Ctx) ->
|
||||
zstd_frame_compress(Data, Ctx).
|
||||
|
||||
-spec decompress(binary(), compress_ctx()) -> {ok, binary(), compress_ctx()} | {error, term()}.
|
||||
decompress(Data, #{type := none} = Ctx) ->
|
||||
decompress(Data, Ctx) ->
|
||||
decompress(Data, Ctx, ?MAX_DECOMPRESSED_SIZE).
|
||||
|
||||
-spec decompress(binary(), compress_ctx(), pos_integer()) ->
|
||||
{ok, binary(), compress_ctx()} | {error, term()}.
|
||||
decompress(Data, #{type := none} = Ctx, _MaxSize) ->
|
||||
{ok, Data, Ctx};
|
||||
decompress(Data, #{type := zstd_frame} = Ctx) ->
|
||||
zstd_frame_decompress(Data, Ctx);
|
||||
decompress(Data, #{type := zstd_stream, stream_ctx := _} = Ctx) ->
|
||||
zstd_stream_decompress(Data, Ctx);
|
||||
decompress(Data, #{type := zstd_stream} = Ctx) ->
|
||||
zstd_frame_decompress(Data, Ctx).
|
||||
decompress(Data, #{type := zstd_frame} = Ctx, MaxSize) ->
|
||||
zstd_frame_decompress(Data, Ctx, MaxSize);
|
||||
decompress(Data, #{type := zstd_stream, stream_ctx := _} = Ctx, MaxSize) ->
|
||||
zstd_stream_decompress(Data, Ctx, MaxSize);
|
||||
decompress(Data, #{type := zstd_stream} = Ctx, MaxSize) ->
|
||||
zstd_frame_decompress(Data, Ctx, MaxSize).
|
||||
|
||||
-spec zstd_frame_compress(iodata(), compress_ctx()) ->
|
||||
{ok, binary(), compress_ctx()} | {error, term()}.
|
||||
@@ -199,56 +206,51 @@ handle_zstd_stream_compress_result({error, Reason}, _Ctx) ->
|
||||
handle_zstd_stream_compress_result(Other, _Ctx) ->
|
||||
{error, {compress_failed, {case_clause, Other}}}.
|
||||
|
||||
-spec zstd_frame_decompress(binary(), compress_ctx()) ->
|
||||
-spec zstd_frame_decompress(binary(), compress_ctx(), pos_integer()) ->
|
||||
{ok, binary(), compress_ctx()} | {error, term()}.
|
||||
zstd_frame_decompress(Data, Ctx) ->
|
||||
zstd_frame_decompress(Data, Ctx, MaxSize) ->
|
||||
case ezstd_available() of
|
||||
true -> zstd_frame_decompress_available(Data, Ctx);
|
||||
true -> zstd_frame_decompress_available(Data, Ctx, MaxSize);
|
||||
false -> {error, {decompress_failed, zstd_not_available}}
|
||||
end.
|
||||
|
||||
-spec zstd_frame_decompress_available(binary(), compress_ctx()) ->
|
||||
-spec zstd_frame_decompress_available(binary(), compress_ctx(), pos_integer()) ->
|
||||
{ok, binary(), compress_ctx()} | {error, term()}.
|
||||
zstd_frame_decompress_available(Data, Ctx) ->
|
||||
zstd_frame_decompress_available(Data, Ctx, MaxSize) ->
|
||||
try
|
||||
handle_zstd_decompress_result(erlang:apply(ezstd, decompress, [Data]), Ctx)
|
||||
handle_zstd_decompress_result(erlang:apply(ezstd, decompress, [Data]), Ctx, MaxSize)
|
||||
catch
|
||||
_:Exception ->
|
||||
{error, {decompress_failed, Exception}}
|
||||
end.
|
||||
|
||||
-spec handle_zstd_decompress_result(term(), compress_ctx()) ->
|
||||
-spec handle_zstd_decompress_result(term(), compress_ctx(), pos_integer()) ->
|
||||
{ok, binary(), compress_ctx()} | {error, term()}.
|
||||
handle_zstd_decompress_result(Decompressed, Ctx) when is_binary(Decompressed) ->
|
||||
check_decompressed_size(Decompressed, Ctx);
|
||||
handle_zstd_decompress_result(Decompressed, Ctx) when is_list(Decompressed) ->
|
||||
IoList = eqwalizer:dynamic_cast(Decompressed),
|
||||
case erlang:iolist_size(IoList) > ?MAX_DECOMPRESSED_SIZE of
|
||||
true -> {error, decompression_too_large};
|
||||
false -> check_decompressed_size(iolist_to_binary(IoList), Ctx)
|
||||
end;
|
||||
handle_zstd_decompress_result({error, Reason}, _Ctx) ->
|
||||
handle_zstd_decompress_result(Decompressed, Ctx, MaxSize) when is_binary(Decompressed) ->
|
||||
check_decompressed_size(Decompressed, Ctx, MaxSize);
|
||||
handle_zstd_decompress_result({error, Reason}, _Ctx, _MaxSize) ->
|
||||
{error, {decompress_failed, Reason}};
|
||||
handle_zstd_decompress_result(Other, _Ctx) ->
|
||||
handle_zstd_decompress_result(Other, _Ctx, _MaxSize) ->
|
||||
{error, {decompress_failed, {case_clause, Other}}}.
|
||||
|
||||
-spec zstd_stream_decompress(binary(), compress_ctx()) ->
|
||||
-spec zstd_stream_decompress(binary(), compress_ctx(), pos_integer()) ->
|
||||
{ok, binary(), compress_ctx()} | {error, term()}.
|
||||
zstd_stream_decompress(Data, Ctx) ->
|
||||
zstd_stream_decompress(Data, Ctx, MaxSize) ->
|
||||
case ezstd_stream_available() of
|
||||
true -> zstd_stream_decompress_available(Data, Ctx);
|
||||
true -> zstd_stream_decompress_available(Data, Ctx, MaxSize);
|
||||
false -> {error, {decompress_failed, zstd_stream_not_available}}
|
||||
end.
|
||||
|
||||
-spec zstd_stream_decompress_available(binary(), compress_ctx()) ->
|
||||
-spec zstd_stream_decompress_available(binary(), compress_ctx(), pos_integer()) ->
|
||||
{ok, binary(), compress_ctx()} | {error, term()}.
|
||||
zstd_stream_decompress_available(Data, Ctx) ->
|
||||
zstd_stream_decompress_available(Data, Ctx, MaxSize) ->
|
||||
try
|
||||
case ensure_zstd_decompress_stream_context(Ctx) of
|
||||
{ok, StreamCtx, NewCtx} ->
|
||||
handle_zstd_decompress_result(
|
||||
erlang:apply(ezstd, decompress_streaming, [StreamCtx, Data]), NewCtx
|
||||
);
|
||||
State = #{
|
||||
stream_ctx => StreamCtx, data => Data, max_size => MaxSize, ctx => NewCtx
|
||||
},
|
||||
decompress_stream_chunks(State, 0, [], 0, ?ZSTD_STREAM_MAX_CHUNKS);
|
||||
{error, Reason} ->
|
||||
{error, Reason}
|
||||
end
|
||||
@@ -257,6 +259,57 @@ zstd_stream_decompress_available(Data, Ctx) ->
|
||||
{error, {decompress_failed, Exception}}
|
||||
end.
|
||||
|
||||
-type stream_chunk_state() :: #{
|
||||
stream_ctx := reference(),
|
||||
data := binary(),
|
||||
max_size := pos_integer(),
|
||||
ctx := compress_ctx()
|
||||
}.
|
||||
|
||||
-spec decompress_stream_chunks(
|
||||
stream_chunk_state(), non_neg_integer(), iolist(), non_neg_integer(), non_neg_integer()
|
||||
) ->
|
||||
{ok, binary(), compress_ctx()} | {error, term()}.
|
||||
decompress_stream_chunks(_State, _Offset, _Acc, _Size, 0) ->
|
||||
{error, {decompress_failed, decompressor_stuck}};
|
||||
decompress_stream_chunks(
|
||||
#{stream_ctx := StreamCtx, data := Data} = State, Offset, Acc, Size, ChunksLeft
|
||||
) ->
|
||||
Result = erlang:apply(ezstd_nif, decompress_streaming_chunk, [StreamCtx, Data, Offset]),
|
||||
handle_stream_chunk(Result, State, Acc, Size, ChunksLeft).
|
||||
|
||||
-spec handle_stream_chunk(
|
||||
term(), stream_chunk_state(), iolist(), non_neg_integer(), non_neg_integer()
|
||||
) ->
|
||||
{ok, binary(), compress_ctx()} | {error, term()}.
|
||||
handle_stream_chunk(
|
||||
{ok, Chunk}, #{max_size := MaxSize, ctx := Ctx}, Acc, Size, _ChunksLeft
|
||||
) when
|
||||
is_binary(Chunk), Size + byte_size(Chunk) =< MaxSize
|
||||
->
|
||||
{ok, iolist_to_binary([Acc, Chunk]), Ctx};
|
||||
handle_stream_chunk(
|
||||
{continue, Chunk, NextOffset}, #{max_size := MaxSize} = State, Acc, Size, ChunksLeft
|
||||
) when
|
||||
is_binary(Chunk),
|
||||
is_integer(NextOffset),
|
||||
NextOffset >= 0,
|
||||
Size + byte_size(Chunk) =< MaxSize
|
||||
->
|
||||
decompress_stream_chunks(
|
||||
State, NextOffset, [Acc, Chunk], Size + byte_size(Chunk), ChunksLeft - 1
|
||||
);
|
||||
handle_stream_chunk({ok, Chunk}, _State, _Acc, _Size, _ChunksLeft) when is_binary(Chunk) ->
|
||||
{error, decompression_too_large};
|
||||
handle_stream_chunk({continue, Chunk, _NextOffset}, _State, _Acc, _Size, _ChunksLeft) when
|
||||
is_binary(Chunk)
|
||||
->
|
||||
{error, decompression_too_large};
|
||||
handle_stream_chunk({error, Reason}, _State, _Acc, _Size, _ChunksLeft) ->
|
||||
{error, {decompress_failed, Reason}};
|
||||
handle_stream_chunk(Other, _State, _Acc, _Size, _ChunksLeft) ->
|
||||
{error, {decompress_failed, {case_clause, Other}}}.
|
||||
|
||||
-spec ensure_zstd_decompress_stream_context(compress_ctx()) ->
|
||||
{ok, reference(), compress_ctx()} | {error, term()}.
|
||||
ensure_zstd_decompress_stream_context(
|
||||
@@ -287,10 +340,10 @@ set_zstd_decompress_window_log_max(StreamCtx, Ctx) ->
|
||||
Other -> {error, {decompress_failed, {case_clause, Other}}}
|
||||
end.
|
||||
|
||||
-spec check_decompressed_size(binary(), compress_ctx()) ->
|
||||
-spec check_decompressed_size(binary(), compress_ctx(), pos_integer()) ->
|
||||
{ok, binary(), compress_ctx()} | {error, term()}.
|
||||
check_decompressed_size(Decompressed, Ctx) ->
|
||||
case byte_size(Decompressed) > ?MAX_DECOMPRESSED_SIZE of
|
||||
check_decompressed_size(Decompressed, Ctx, MaxSize) ->
|
||||
case byte_size(Decompressed) > MaxSize of
|
||||
true -> {error, decompression_too_large};
|
||||
false -> {ok, Decompressed, Ctx}
|
||||
end.
|
||||
@@ -334,7 +387,17 @@ probe_ezstd_stream_available() ->
|
||||
erlang:function_exported(ezstd, set_compression_parameter, 3) andalso
|
||||
erlang:function_exported(ezstd, compress_streaming, 2) andalso
|
||||
erlang:function_exported(ezstd, create_decompression_context, 1) andalso
|
||||
erlang:function_exported(ezstd, decompress_streaming, 2);
|
||||
erlang:function_exported(ezstd, decompress_streaming, 2) andalso
|
||||
ezstd_nif_stream_chunk_available();
|
||||
_ ->
|
||||
false
|
||||
end.
|
||||
|
||||
-spec ezstd_nif_stream_chunk_available() -> boolean().
|
||||
ezstd_nif_stream_chunk_available() ->
|
||||
case code:ensure_loaded(ezstd_nif) of
|
||||
{module, ezstd_nif} ->
|
||||
erlang:function_exported(ezstd_nif, decompress_streaming_chunk, 3);
|
||||
_ ->
|
||||
false
|
||||
end.
|
||||
@@ -402,17 +465,66 @@ decompress_none_test() ->
|
||||
check_decompressed_size_allows_normal_payload_test() ->
|
||||
Ctx = new_context(zstd_frame),
|
||||
Data = <<"normal payload">>,
|
||||
?assertEqual({ok, Data, Ctx}, check_decompressed_size(Data, Ctx)).
|
||||
?assertEqual({ok, Data, Ctx}, check_decompressed_size(Data, Ctx, ?MAX_DECOMPRESSED_SIZE)).
|
||||
|
||||
check_decompressed_size_rejects_oversized_payload_test() ->
|
||||
Ctx = new_context(zstd_frame),
|
||||
Oversized = binary:copy(<<0>>, ?MAX_DECOMPRESSED_SIZE + 1),
|
||||
?assertEqual({error, decompression_too_large}, check_decompressed_size(Oversized, Ctx)).
|
||||
?assertEqual(
|
||||
{error, decompression_too_large},
|
||||
check_decompressed_size(Oversized, Ctx, ?MAX_DECOMPRESSED_SIZE)
|
||||
).
|
||||
|
||||
check_decompressed_size_allows_exact_limit_test() ->
|
||||
Ctx = new_context(zstd_frame),
|
||||
ExactLimit = binary:copy(<<0>>, ?MAX_DECOMPRESSED_SIZE),
|
||||
?assertMatch({ok, _, _}, check_decompressed_size(ExactLimit, Ctx)).
|
||||
?assertMatch({ok, _, _}, check_decompressed_size(ExactLimit, Ctx, ?MAX_DECOMPRESSED_SIZE)).
|
||||
|
||||
zstd_stream_decompress_stops_past_max_size_test() ->
|
||||
case probe_ezstd_stream_available() of
|
||||
true ->
|
||||
Compressed = stream_compress_zeros(new_context(zstd_stream), 12, <<>>),
|
||||
?assert(byte_size(Compressed) =< 4096),
|
||||
?assertEqual(
|
||||
{error, decompression_too_large},
|
||||
decompress(Compressed, new_context(zstd_stream), 4096)
|
||||
);
|
||||
false ->
|
||||
?assertEqual(skip, skip)
|
||||
end.
|
||||
|
||||
zstd_stream_decompress_allows_payload_at_max_size_test() ->
|
||||
case probe_ezstd_stream_available() of
|
||||
true ->
|
||||
Data = binary:copy(<<"a">>, 4096),
|
||||
{ok, Compressed, _} = compress(Data, new_context(zstd_stream)),
|
||||
?assertMatch({ok, Data, _}, decompress(Compressed, new_context(zstd_stream), 4096)),
|
||||
{ok, Over, _} = compress(<<Data/binary, "b">>, new_context(zstd_stream)),
|
||||
?assertEqual(
|
||||
{error, decompression_too_large},
|
||||
decompress(Over, new_context(zstd_stream), 4096)
|
||||
);
|
||||
false ->
|
||||
?assertEqual(skip, skip)
|
||||
end.
|
||||
|
||||
zstd_stream_decompress_keeps_context_across_frames_test() ->
|
||||
case probe_ezstd_stream_available() of
|
||||
true ->
|
||||
{ok, First, EncodeCtx} = compress(<<"first message">>, new_context(zstd_stream)),
|
||||
{ok, Second, _} = compress(<<"second message">>, EncodeCtx),
|
||||
{ok, <<"first message">>, DecodeCtx} =
|
||||
decompress(First, new_context(zstd_stream), 4096),
|
||||
?assertMatch({ok, <<"second message">>, _}, decompress(Second, DecodeCtx, 4096));
|
||||
false ->
|
||||
?assertEqual(skip, skip)
|
||||
end.
|
||||
|
||||
stream_compress_zeros(_Ctx, 0, Acc) ->
|
||||
Acc;
|
||||
stream_compress_zeros(Ctx, Remaining, Acc) ->
|
||||
{ok, Chunk, NextCtx} = compress(binary:copy(<<0>>, 8 * 1024 * 1024), Ctx),
|
||||
stream_compress_zeros(NextCtx, Remaining - 1, <<Acc/binary, Chunk/binary>>).
|
||||
|
||||
init_caches_availability_test() ->
|
||||
with_saved_availability(fun() ->
|
||||
|
||||
@@ -261,7 +261,7 @@ handle_incoming_data(Data, #{encoding := Encoding, compress_ctx := CompressCtx0}
|
||||
ws_result().
|
||||
handle_decompressed_incoming_data(Data, Encoding, CompressCtx, State) ->
|
||||
MaxPayloadSize = constants:max_payload_size(),
|
||||
case gateway_compress:decompress(Data, CompressCtx) of
|
||||
case gateway_compress:decompress(Data, CompressCtx, MaxPayloadSize) of
|
||||
{ok, Decompressed, NewCompressCtx} when byte_size(Decompressed) =< MaxPayloadSize ->
|
||||
Decoded = gateway_codec:decode(Decompressed, Encoding),
|
||||
handle_decode(Decoded, State#{compress_ctx => NewCompressCtx});
|
||||
@@ -269,6 +269,10 @@ handle_decompressed_incoming_data(Data, Encoding, CompressCtx, State) ->
|
||||
gateway_handler_encode:close_with_reason(
|
||||
decode_error, <<"Payload too large">>, State
|
||||
);
|
||||
{error, decompression_too_large} ->
|
||||
gateway_handler_encode:close_with_reason(
|
||||
decode_error, <<"Payload too large">>, State
|
||||
);
|
||||
{error, _Reason} ->
|
||||
gateway_handler_encode:close_with_reason(
|
||||
decode_error, <<"Decompression failed">>, State
|
||||
|
||||
@@ -506,6 +506,18 @@ voice_state_update_with_non_map_payload_closes_test() ->
|
||||
),
|
||||
?assertEqual(constants:close_code_to_num(decode_error), CloseCode).
|
||||
|
||||
websocket_handle_rejects_compressed_frame_past_max_payload_test() ->
|
||||
{ok, Compressed, _} = gateway_compress:compress(
|
||||
binary:copy(<<"a">>, constants:max_payload_size() * 64),
|
||||
gateway_compress:new_context(zstd_stream)
|
||||
),
|
||||
?assert(byte_size(Compressed) =< constants:max_payload_size()),
|
||||
State = (new_json_state())#{compress_ctx => gateway_compress:new_context(zstd_stream)},
|
||||
{[{close, CloseCode, Reason}], _} =
|
||||
gateway_handler:websocket_handle({binary, Compressed}, State),
|
||||
?assertEqual(constants:close_code_to_num(decode_error), CloseCode),
|
||||
?assertEqual(<<"Payload too large">>, Reason).
|
||||
|
||||
new_json_state() ->
|
||||
(gateway_handler:new_state())#{
|
||||
version => 1, encoding => json, compress_ctx => gateway_compress:new_context(none)
|
||||
|
||||
Reference in New Issue
Block a user