chore: tidy request handling across services (#3168)

This commit is contained in:
Hampus
2026-10-03 15:36:33 +02:00
committed by GitHub
parent c6941d5905
commit da9e9ff0be
11 changed files with 428 additions and 44 deletions
+151 -39
View File
@@ -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)