fix(gateway): return precise rpc errors and bound snowflakes (#2491)

This commit is contained in:
Hampus
2026-09-06 15:24:54 +02:00
committed by GitHub
parent 214d19d45a
commit 31d7cb81d6
8 changed files with 183 additions and 16 deletions
@@ -59,6 +59,8 @@ execute_rpc_method(Method, PayloadBin) ->
catch
error:{gateway_rpc_error, Message} ->
handle_throw_error(Method, Message);
error:{validation, _Reason} ->
handle_throw_error(Method, <<"invalid_params">>);
throw:{error, Message} ->
handle_throw_error(Method, Message);
throw:Message ->
@@ -241,6 +243,33 @@ parse_nats_url_test() ->
?assertEqual({ok, "127.0.0.1", 4222}, parse_nats_url("nats://127.0.0.1:4222")),
?assertEqual({error, invalid_nats_url}, parse_nats_url(undefined)).
execute_rpc_method_maps_validation_failure_to_invalid_params_test() ->
Payload = iolist_to_binary(
json:encode(#{
<<"guild_id">> => <<"nope">>,
<<"user_id">> => <<"2">>,
<<"channel_id">> => <<"0">>
})
),
?assertEqual(
#{<<"ok">> => false, <<"error">> => <<"invalid_params">>},
execute_rpc_method(<<"guild.get_user_permissions">>, Payload)
).
execute_rpc_method_maps_oversized_batch_to_batch_too_large_test() ->
GuildIds = [integer_to_binary(N) || N <- lists:seq(1, 101)],
Payload = iolist_to_binary(
json:encode(#{
<<"guild_ids">> => GuildIds,
<<"user_id">> => <<"1">>,
<<"channel_id">> => <<"0">>
})
),
?assertEqual(
#{<<"ok">> => false, <<"error">> => <<"batch_too_large">>},
execute_rpc_method(<<"guild.get_user_permissions_batch">>, Payload)
).
rpc_subjects_for_role_does_not_route_rpc_to_websocket_test() ->
?assertEqual([], rpc_subjects_for_role(websocket)).
@@ -30,6 +30,7 @@ handle_get_user_permissions(
ChannelId = parse_channel_id(CIB),
case get_permissions_cached_or_rpc(GuildId, UserId, ChannelId) of
{ok, Perms} -> #{<<"permissions">> => integer_to_binary(Perms)};
{error, guild_not_found} -> gateway_rpc_error:raise(<<"guild_not_found">>);
error -> gateway_rpc_error:raise(<<"permissions_error">>)
end.
@@ -57,7 +58,7 @@ batch_permission_result(GuildId, UserId, ChannelId) when is_integer(GuildId) ->
<<"guild_id">> => integer_to_binary(GuildId),
<<"permissions">> => integer_to_binary(Permissions)
};
error ->
_ ->
undefined
end;
batch_permission_result(_, _, _) ->
@@ -79,6 +80,8 @@ handle_check_permission(
case get_permissions_cached_or_rpc(GuildId, UserId, ChannelId) of
{ok, Perms} ->
#{<<"has_permission">> => permission_bits:has(Perms, Permission)};
{error, guild_not_found} ->
gateway_rpc_error:raise(<<"guild_not_found">>);
error ->
gateway_rpc_error:raise(<<"permission_check_error">>)
end.
@@ -180,7 +183,7 @@ guild_call_max_pos(Pid, UserId) ->
-spec get_permissions_cached_or_rpc(
integer(), integer(), integer() | undefined
) -> {ok, integer()} | error.
) -> {ok, integer()} | {error, guild_not_found} | error.
get_permissions_cached_or_rpc(GuildId, UserId, ChannelId) ->
case guild_permission_cache:get_permissions(GuildId, UserId, ChannelId) of
{ok, Perms} -> {ok, Perms};
@@ -189,14 +192,14 @@ get_permissions_cached_or_rpc(GuildId, UserId, ChannelId) ->
-spec get_permissions_via_rpc(
integer(), integer(), integer() | undefined
) -> {ok, integer()} | error.
) -> {ok, integer()} | {error, guild_not_found} | error.
get_permissions_via_rpc(GuildId, UserId, ChannelId) ->
case gateway_rpc_guild_infra:ensure_guild_pid(GuildId) of
{ok, Pid} ->
Msg = {get_user_permissions, #{user_id => UserId, channel_id => ChannelId}},
get_perms_from_guild(GuildId, Pid, Msg);
error ->
error
{error, guild_not_found}
end.
-spec get_perms_from_guild(integer(), pid(), term()) -> {ok, integer()} | error.
@@ -219,7 +222,7 @@ permission_or_throw(Value) ->
try permission_bits:parse(Value) of
Permission when is_integer(Permission) -> Permission
catch
error:{invalid_bitset, _} -> gateway_rpc_error:raise(validation_invalid_params)
error:{invalid_bitset, _} -> gateway_rpc_error:raise(<<"invalid_params">>)
end.
-ifdef(TEST).
@@ -231,7 +234,30 @@ permission_or_throw_accepts_zero_test() ->
permission_or_throw_rejects_malformed_test() ->
?assertError(
{gateway_rpc_error, validation_invalid_params}, permission_or_throw(<<"bad">>)
{gateway_rpc_error, <<"invalid_params">>}, permission_or_throw(<<"bad">>)
).
check_permission_unknown_guild_raises_guild_not_found_test() ->
Payload = #{
<<"guild_id">> => <<"1">>,
<<"user_id">> => <<"2">>,
<<"permission">> => <<"1">>,
<<"channel_id">> => <<"0">>
},
?assertError(
{gateway_rpc_error, <<"guild_not_found">>},
handle(<<"guild.check_permission">>, Payload)
).
get_user_permissions_unknown_guild_raises_guild_not_found_test() ->
Payload = #{
<<"guild_id">> => <<"1">>,
<<"user_id">> => <<"2">>,
<<"channel_id">> => <<"0">>
},
?assertError(
{gateway_rpc_error, <<"guild_not_found">>},
handle(<<"guild.get_user_permissions">>, Payload)
).
get_user_permissions_batch_rejects_oversized_test() ->
@@ -242,7 +268,7 @@ get_user_permissions_batch_rejects_oversized_test() ->
<<"channel_id">> => <<"0">>
},
?assertError(
{gateway_rpc_error, _},
{gateway_rpc_error, <<"batch_too_large">>},
handle(<<"guild.get_user_permissions_batch">>, Payload)
).
-endif.
@@ -163,10 +163,7 @@ safe_guild_counts_get(GuildId) ->
-spec validate_batch_size(non_neg_integer()) -> ok.
validate_batch_size(Size) when Size > ?MAX_BATCH_SIZE ->
Max = integer_to_binary(?MAX_BATCH_SIZE),
gateway_rpc_error:raise(
<<"Batch size exceeds maximum of ", Max/binary>>
);
gateway_rpc_error:raise(<<"batch_too_large">>);
validate_batch_size(_) ->
ok.
@@ -265,7 +262,7 @@ owner_groups_for_reload_all_empty_ids_uses_active_nodes_test() ->
validate_batch_size_test() ->
?assertEqual(ok, validate_batch_size(50)),
?assertEqual(ok, validate_batch_size(100)),
?assertError({gateway_rpc_error, _}, validate_batch_size(101)).
?assertError({gateway_rpc_error, <<"batch_too_large">>}, validate_batch_size(101)).
process_batch_collects_successful_results_test() ->
Results = process_batch([1, 2, 3], fun(N) -> eqwalizer:dynamic_cast(N) * 2 end, 1000),
+21 -3
View File
@@ -35,8 +35,10 @@
-type role_id() :: t().
-type message_id() :: t().
-define(MAX_SNOWFLAKE, 16#7FFFFFFFFFFFFFFF).
-spec parse(term()) -> t().
parse(Value) when is_integer(Value), Value > 0 ->
parse(Value) when is_integer(Value), Value > 0, Value =< ?MAX_SNOWFLAKE ->
Value;
parse(Value) when is_binary(Value) ->
require_parsed(parse_binary(Value), Value);
@@ -130,7 +132,7 @@ get(_Id, _Map, Default) ->
-spec parse_binary(binary()) -> t() | undefined.
parse_binary(<<First, Rest/binary>> = Value) when First >= $1, First =< $9 ->
case all_digits(Rest) of
true -> binary_to_integer(Value);
true -> bounded(binary_to_integer(Value));
false -> undefined
end;
parse_binary(_) ->
@@ -144,12 +146,18 @@ parse_list_value(_) ->
-spec parse_digits_list([term()], pos_integer()) -> t() | undefined.
parse_digits_list([], Acc) ->
Acc;
bounded(Acc);
parse_digits_list([Digit | Rest], Acc) when is_integer(Digit), Digit >= $0, Digit =< $9 ->
parse_digits_list(Rest, Acc * 10 + Digit - $0);
parse_digits_list(_, _Acc) ->
undefined.
-spec bounded(pos_integer()) -> t() | undefined.
bounded(Id) when Id =< ?MAX_SNOWFLAKE ->
Id;
bounded(_) ->
undefined.
-spec require_parsed(t() | undefined, term()) -> t().
require_parsed(Id, _Value) when is_integer(Id) ->
Id;
@@ -180,6 +188,16 @@ parse_rejects_non_canonical_values_test() ->
?assertError({invalid_snowflake, <<"+1">>}, parse(<<"+1">>)),
?assertError({invalid_snowflake, <<"abc">>}, parse(<<"abc">>)).
parse_rejects_values_above_int64_test() ->
?assertEqual(9223372036854775807, parse(9223372036854775807)),
?assertEqual(9223372036854775807, parse(<<"9223372036854775807">>)),
?assertEqual(9223372036854775807, parse("9223372036854775807")),
?assertError({invalid_snowflake, 9223372036854775808}, parse(9223372036854775808)),
?assertError(
{invalid_snowflake, <<"9223372036854775808">>}, parse(<<"9223372036854775808">>)
),
?assertError({invalid_snowflake, "9223372036854775808"}, parse("9223372036854775808")).
member_handles_mixed_edge_values_test() ->
?assertEqual(true, member(123, [<<"123">>, 456])),
?assertError({invalid_snowflake, <<"00123">>}, member(123, [<<"00123">>, 456])).
+19 -1
View File
@@ -22,6 +22,20 @@ to_snowflake_rejects_malformed_values_test() ->
?assertError({invalid_snowflake, <<"abc">>}, snowflake_id:parse(<<"abc">>)),
?assertError({invalid_snowflake, 12.34}, snowflake_id:parse(12.34)).
to_snowflake_rejects_values_above_int64_test() ->
?assertEqual(9223372036854775807, snowflake_id:parse(9223372036854775807)),
?assertEqual(9223372036854775807, snowflake_id:parse(<<"9223372036854775807">>)),
?assertError(
{invalid_snowflake, 9223372036854775808}, snowflake_id:parse(9223372036854775808)
),
?assertError(
{invalid_snowflake, <<"9223372036854775808">>},
snowflake_id:parse(<<"9223372036854775808">>)
),
?assertError(
{invalid_snowflake, "9223372036854775808"}, snowflake_id:parse("9223372036854775808")
).
extract_id_with_atom_key_valid_test() ->
?assertEqual(123, type_conv:extract_id(#{user_id => 123}, user_id)),
?assertEqual(456, type_conv:extract_id(#{user_id => <<"456">>}, user_id)),
@@ -35,7 +49,11 @@ extract_id_with_atom_key_rejects_malformed_ids_test() ->
?assertEqual(undefined, type_conv:extract_id(#{user_id => <<"+1">>}, user_id)),
?assertEqual(undefined, type_conv:extract_id(#{user_id => "001"}, user_id)),
?assertEqual(undefined, type_conv:extract_id(#{user_id => "invalid"}, user_id)),
?assertEqual(undefined, type_conv:extract_id(#{user_id => [1, 2, 3]}, user_id)).
?assertEqual(undefined, type_conv:extract_id(#{user_id => [1, 2, 3]}, user_id)),
?assertEqual(undefined, type_conv:extract_id(#{user_id => 9223372036854775808}, user_id)),
?assertEqual(
undefined, type_conv:extract_id(#{user_id => <<"9223372036854775808">>}, user_id)
).
extract_id_with_atom_key_missing_or_invalid_test() ->
?assertEqual(undefined, type_conv:extract_id(#{other_field => 999}, user_id)),