fix(gateway): repair the session lifecycle, limits and dead code (#2490)

This commit is contained in:
Hampus
2026-09-06 15:23:59 +02:00
committed by GitHub
parent 9a229c1b73
commit d3170fc320
28 changed files with 727 additions and 235 deletions
@@ -4,7 +4,7 @@
-typing([eqwalizer]). -typing([eqwalizer]).
-behaviour(gen_server). -behaviour(gen_server).
-export([start_link/0, trigger/1, drain_async/0, undrain/0, diagnostic_info/0]). -export([start_link/0, drain_async/0, undrain/0, diagnostic_info/0]).
-export([ -export([
init/1, init/1,
handle_call/3, handle_call/3,
@@ -44,11 +44,6 @@
start_link() -> start_link() ->
gen_server:start_link({local, ?MODULE}, ?MODULE, [], []). gen_server:start_link({local, ?MODULE}, ?MODULE, [], []).
-spec trigger([node()]) -> ok.
trigger(Members) when is_list(Members) ->
shard_utils:safe_apply(fun() -> gen_server:cast(?MODULE, {trigger, Members}) end, ok),
ok.
-spec drain_async() -> ok. -spec drain_async() -> ok.
drain_async() -> drain_async() ->
persistent_term:put({fluxer_gateway, draining}, true), persistent_term:put({fluxer_gateway, draining}, true),
@@ -98,9 +93,6 @@ handle_call(_Request, _From, State) ->
{reply, ok, State}. {reply, ok, State}.
-spec handle_cast(term(), state()) -> {noreply, state()}. -spec handle_cast(term(), state()) -> {noreply, state()}.
handle_cast({trigger, Members}, State) ->
Normalized = gateway_cluster_handoff_transfer:normalize_members(Members),
{noreply, schedule_if_changed(Normalized, State)};
handle_cast(drain, State) -> handle_cast(drain, State) ->
cancel_timer(maps:get(timer, State, undefined)), cancel_timer(maps:get(timer, State, undefined)),
DrainMembers = gateway_cluster_membership:members(), DrainMembers = gateway_cluster_membership:members(),
@@ -111,7 +111,9 @@ websocket_info({heartbeat_check}, State) ->
gateway_handler_heartbeat:handle_legacy_heartbeat_check(State); gateway_handler_heartbeat:handle_legacy_heartbeat_check(State);
websocket_info({dispatch, Event, Data, Seq}, State) when websocket_info({dispatch, Event, Data, Seq}, State) when
is_integer(Seq), is_atom(Event), is_map(Data); is_integer(Seq), is_atom(Event), is_map(Data);
is_integer(Seq), is_binary(Event), is_map(Data) is_integer(Seq), is_binary(Event), is_map(Data);
is_integer(Seq), is_atom(Event), is_list(Data);
is_integer(Seq), is_binary(Event), is_list(Data)
-> ->
gateway_handler_dispatch:handle_dispatch(Event, Data, Seq, State); gateway_handler_dispatch:handle_dispatch(Event, Data, Seq, State);
websocket_info({dispatch, Event, null, Seq}, State) when websocket_info({dispatch, Event, null, Seq}, State) when
@@ -123,8 +125,6 @@ websocket_info({dispatch, Event, {pre_encoded, Bin} = Data, Seq}, State) when
is_integer(Seq), is_binary(Event), is_binary(Bin) is_integer(Seq), is_binary(Event), is_binary(Bin)
-> ->
gateway_handler_dispatch:handle_dispatch(Event, Data, Seq, State); gateway_handler_dispatch:handle_dispatch(Event, Data, Seq, State);
websocket_info({session_backpressure_error, _Details}, State) ->
{ok, State};
websocket_info(rollout_config_changed, State) -> websocket_info(rollout_config_changed, State) ->
gateway_handler_identify:handle_rollout_config_changed(State); gateway_handler_identify:handle_rollout_config_changed(State);
websocket_info({retry_pending_identify, Token}, State) when is_reference(Token) -> websocket_info({retry_pending_identify, Token}, State) when is_reference(Token) ->
@@ -99,7 +99,7 @@ handle_rate_limited_resume(Data, State) ->
end. end.
-spec handle_dispatch( -spec handle_dispatch(
atom() | binary(), map() | null | {pre_encoded, binary()}, integer(), state() atom() | binary(), map() | list() | null | {pre_encoded, binary()}, integer(), state()
) -> ws_result(). ) -> ws_result().
handle_dispatch(Event, Data, Seq, State) -> handle_dispatch(Event, Data, Seq, State) ->
case gateway_event_pause:is_frozen() of case gateway_event_pause:is_frozen() of
@@ -108,7 +108,7 @@ handle_dispatch(Event, Data, Seq, State) ->
end. end.
-spec do_dispatch( -spec do_dispatch(
atom() | binary(), map() | null | {pre_encoded, binary()}, integer(), state() atom() | binary(), map() | list() | null | {pre_encoded, binary()}, integer(), state()
) -> ) ->
ws_result(). ws_result().
do_dispatch(Event, {pre_encoded, EncodedData}, Seq, State) -> do_dispatch(Event, {pre_encoded, EncodedData}, Seq, State) ->
@@ -143,7 +143,9 @@ dispatch_pre_encoded(Event, EncodedData, Seq, #{compress_ctx := CompressCtx} = S
{ok, State} {ok, State}
end. end.
-spec dispatch_standard(atom() | binary(), map() | null, integer(), state()) -> ws_result(). -spec dispatch_standard(
atom() | binary(), map() | list() | null, integer(), state()
) -> ws_result().
dispatch_standard(Event, Data, Seq, State) -> dispatch_standard(Event, Data, Seq, State) ->
EventName = gateway_handler_encode:dispatch_event_name(Event), EventName = gateway_handler_encode:dispatch_event_name(Event),
Message = #{ Message = #{
@@ -350,12 +350,6 @@ session_start_error_action(_) ->
unknown. unknown.
-spec log_session_start_error(term(), state()) -> ws_result(). -spec log_session_start_error(term(), state()) -> ws_result().
log_session_start_error({retries_exhausted, Reason}, State) ->
logger:error(
"Session start failed after retries: last_error=~p peer_ip=~ts",
[Reason, maps:get(peer_ip, State, <<"unknown">>)]
),
gateway_handler_encode:close_with_reason(unknown_error, <<"Session start failed">>, State);
log_session_start_error(Reason, State) -> log_session_start_error(Reason, State) ->
logger:error( logger:error(
"Session start failed: reason=~p peer_ip=~ts", "Session start failed: reason=~p peer_ip=~ts",
@@ -134,6 +134,8 @@ handle_resume_call_result({ok, MissedEvents, CurrentSeq}, Pid, GwTimings, State)
finalize_resume(Pid, CurrentSeq, MissedEvents, GwTimings, State); finalize_resume(Pid, CurrentSeq, MissedEvents, GwTimings, State);
handle_resume_call_result(invalid_seq, _Pid, _GwTimings, State) -> handle_resume_call_result(invalid_seq, _Pid, _GwTimings, State) ->
gateway_handler_encode:close_with_reason(invalid_seq, <<"Invalid sequence">>, State); gateway_handler_encode:close_with_reason(invalid_seq, <<"Invalid sequence">>, State);
handle_resume_call_result(not_resumable, _Pid, _GwTimings, State) ->
send_invalid_session(State);
handle_resume_call_result(_ResumeResult, _Pid, _GwTimings, State) -> handle_resume_call_result(_ResumeResult, _Pid, _GwTimings, State) ->
gateway_handler_encode:close_with_reason( gateway_handler_encode:close_with_reason(
unknown_error, unknown_error,
@@ -108,7 +108,7 @@ cleanup_session_table(TableName, SessionPid) ->
-spec should_queue_voice_update(pid()) -> boolean(). -spec should_queue_voice_update(pid()) -> boolean().
should_queue_voice_update(SessionPid) -> should_queue_voice_update(SessionPid) ->
case rate_limits_disabled() of case gateway_handler_rate_limit:rate_limits_disabled() of
true -> false; true -> false;
false -> should_queue_voice_update_limited(SessionPid) false -> should_queue_voice_update_limited(SessionPid)
end. end.
@@ -141,15 +141,6 @@ check_voice_rate(SessionPid, Timestamps, Now) ->
false false
end. end.
-spec rate_limits_disabled() -> boolean().
rate_limits_disabled() ->
case os:getenv("FLUXER_DISABLE_RATE_LIMITS") of
"1" -> true;
"true" -> true;
"TRUE" -> true;
_ -> false
end.
-spec queue_voice_update(pid(), map()) -> ok. -spec queue_voice_update(pid(), map()) -> ok.
queue_voice_update(SessionPid, Data) -> queue_voice_update(SessionPid, Data) ->
ensure_voice_queue_table(), ensure_voice_queue_table(),
@@ -196,8 +196,6 @@ aggregate_node_stats(NodeStats) ->
}. }.
-spec aggregate_status([map()]) -> binary(). -spec aggregate_status([map()]) -> binary().
aggregate_status([]) ->
<<"unavailable">>;
aggregate_status(NodeStats) -> aggregate_status(NodeStats) ->
AllHealthy = lists:all( AllHealthy = lists:all(
fun(N) -> maps:get(<<"status">>, N, <<"healthy">>) =:= <<"healthy">> end, fun(N) -> maps:get(<<"status">>, N, <<"healthy">>) =:= <<"healthy">> end,
@@ -360,4 +358,10 @@ aggregate_node_stats_ignores_unknown_uptime_test() ->
safe_node_call_catches_local_errors_test() -> safe_node_call_catches_local_errors_test() ->
?assertEqual(error, safe_node_call(node(), definitely_missing_function, [], 100)). ?assertEqual(error, safe_node_call(node(), definitely_missing_function, [], 100)).
collect_and_aggregate_node_stats_with_empty_node_list_reports_local_node_test() ->
A = collect_and_aggregate_node_stats([]),
?assertEqual(1, maps:get(<<"node_count">>, A)),
?assertEqual(<<"healthy">>, maps:get(<<"status">>, A)),
?assertEqual(1, length(maps:get(<<"nodes">>, A))).
-endif. -endif.
@@ -287,7 +287,7 @@ maybe_valid_owner_node(OwnerNode) ->
false -> unavailable false -> unavailable
end. end.
-spec handle_offline_dispatch(atom(), integer(), map()) -> true. -spec handle_offline_dispatch(atom(), integer(), map() | list()) -> true.
handle_offline_dispatch(message_create, UserId, Data) -> handle_offline_dispatch(message_create, UserId, Data) ->
case offline_message_author_id(Data) of case offline_message_author_id(Data) of
{ok, AuthorId} -> {ok, AuthorId} ->
@@ -5,8 +5,6 @@
-export([ -export([
ensure_tables/0, ensure_tables/0,
is_token_banned/1,
ban_token/1,
check_user_session_limit/1, check_user_session_limit/1,
increment_user_sessions/1, increment_user_sessions/1,
decrement_user_sessions/1, decrement_user_sessions/1,
@@ -19,7 +17,6 @@
-define(IDENTIFY_TABLE, gateway_identify_rate). -define(IDENTIFY_TABLE, gateway_identify_rate).
-define(SESSION_USER_COUNTS, session_user_counts). -define(SESSION_USER_COUNTS, session_user_counts).
-define(TOKEN_BAN_TABLE, gateway_token_bans).
-define(IDENTIFY_MAX_PER_IP, 300). -define(IDENTIFY_MAX_PER_IP, 300).
-define(IDENTIFY_WINDOW_SECS, 60). -define(IDENTIFY_WINDOW_SECS, 60).
-define(IDENTIFY_CLEANUP_INTERVAL_MS, ?IDENTIFY_WINDOW_SECS * 2 * 1000). -define(IDENTIFY_CLEANUP_INTERVAL_MS, ?IDENTIFY_WINDOW_SECS * 2 * 1000).
@@ -29,48 +26,17 @@
ensure_tables() -> ensure_tables() ->
ensure_identify_table(), ensure_identify_table(),
ensure_session_user_counts_table(), ensure_session_user_counts_table(),
ensure_token_ban_table(),
ok. ok.
-spec is_token_banned(term()) -> boolean().
is_token_banned(Token) when is_binary(Token) ->
ok = ensure_token_ban_table(),
Key = utils:hash_token(Token),
try ets:member(?TOKEN_BAN_TABLE, Key) of
Member when is_boolean(Member) -> Member
catch
error:badarg -> false
end;
is_token_banned(_) ->
false.
-spec ban_token(term()) -> ok.
ban_token(Token) when is_binary(Token) ->
ok = ensure_token_ban_table(),
Key = utils:hash_token(Token),
try
_ = ets:insert(?TOKEN_BAN_TABLE, {Key, erlang:system_time(second)}),
ok
catch
error:badarg -> ok
end;
ban_token(_) ->
ok.
-spec ensure_token_ban_table() -> ok.
ensure_token_ban_table() ->
case ets:whereis(?TOKEN_BAN_TABLE) of
undefined ->
_ = create_rate_table(?TOKEN_BAN_TABLE),
ok;
_ ->
ok
end.
-spec check_user_session_limit(term()) -> ok | {error, too_many_sessions}. -spec check_user_session_limit(term()) -> ok | {error, too_many_sessions}.
check_user_session_limit(UserId) when is_integer(UserId), UserId > 0 -> check_user_session_limit(UserId) when is_integer(UserId), UserId > 0 ->
ok = ensure_session_user_counts_table(), case gateway_handler_rate_limit:rate_limits_disabled() of
do_check_user_session_limit(UserId); true ->
ok;
false ->
ok = ensure_session_user_counts_table(),
do_check_user_session_limit(UserId)
end;
check_user_session_limit(_) -> check_user_session_limit(_) ->
ok. ok.
@@ -116,8 +82,13 @@ decrement_user_sessions(_) ->
-spec check_identify_rate(term()) -> ok | {error, identify_rate_limited}. -spec check_identify_rate(term()) -> ok | {error, identify_rate_limited}.
check_identify_rate(PeerIP) when is_binary(PeerIP) -> check_identify_rate(PeerIP) when is_binary(PeerIP) ->
ok = ensure_identify_table(), case gateway_handler_rate_limit:rate_limits_disabled() of
do_check_identify_rate(PeerIP); true ->
ok;
false ->
ok = ensure_identify_table(),
do_check_identify_rate(PeerIP)
end;
check_identify_rate(_) -> check_identify_rate(_) ->
ok. ok.
@@ -218,30 +189,34 @@ check_identify_rate_allows_under_limit_test() ->
?assertEqual(ok, check_identify_rate(IP)). ?assertEqual(ok, check_identify_rate(IP)).
check_identify_rate_blocks_over_limit_test() -> check_identify_rate_blocks_over_limit_test() ->
ensure_tables(), with_rate_limits_enabled(fun() ->
IP = <<"192.0.2.200">>, ensure_tables(),
lists:foreach( IP = <<"192.0.2.200">>,
fun(_) -> check_identify_rate(IP) end, lists:foreach(
lists:seq(1, ?IDENTIFY_MAX_PER_IP) fun(_) -> check_identify_rate(IP) end,
), lists:seq(1, ?IDENTIFY_MAX_PER_IP)
?assertEqual({error, identify_rate_limited}, check_identify_rate(IP)). ),
?assertEqual({error, identify_rate_limited}, check_identify_rate(IP))
end).
check_identify_rate_disabled_by_env_test() ->
OldValue = os:getenv("FLUXER_DISABLE_RATE_LIMITS"),
os:putenv("FLUXER_DISABLE_RATE_LIMITS", "true"),
try
ensure_tables(),
IP = <<"192.0.2.201">>,
lists:foreach(
fun(_) -> ?assertEqual(ok, check_identify_rate(IP)) end,
lists:seq(1, ?IDENTIFY_MAX_PER_IP + 1)
),
?assertEqual(ok, check_identify_rate(IP))
after
restore_env("FLUXER_DISABLE_RATE_LIMITS", OldValue)
end.
check_identify_rate_non_binary_returns_ok_test() -> check_identify_rate_non_binary_returns_ok_test() ->
?assertEqual(ok, check_identify_rate(undefined)). ?assertEqual(ok, check_identify_rate(undefined)).
token_ban_round_trips_test() ->
ensure_tables(),
Token = <<"banned_token_abc">>,
?assertEqual(false, is_token_banned(Token)),
?assertEqual(ok, ban_token(Token)),
?assertEqual(true, is_token_banned(Token)),
?assertEqual(false, is_token_banned(<<"other_token">>)),
ets:delete(?TOKEN_BAN_TABLE, utils:hash_token(Token)).
token_ban_non_binary_returns_ok_test() ->
?assertEqual(false, is_token_banned(undefined)),
?assertEqual(ok, ban_token(undefined)).
user_session_limit_allows_under_limit_test() -> user_session_limit_allows_under_limit_test() ->
ensure_tables(), ensure_tables(),
UserId = 900001, UserId = 900001,
@@ -252,28 +227,49 @@ user_session_limit_allows_under_limit_test() ->
delete_user_session_count(UserId). delete_user_session_count(UserId).
user_session_limit_blocks_over_limit_test() -> user_session_limit_blocks_over_limit_test() ->
ensure_tables(), with_rate_limits_enabled(fun() ->
UserId = 900002, ensure_tables(),
delete_user_session_count(UserId), UserId = 900002,
lists:foreach( delete_user_session_count(UserId),
fun(_) -> increment_user_sessions(UserId) end, lists:foreach(
lists:seq(1, ?MAX_SESSIONS_PER_USER) fun(_) -> increment_user_sessions(UserId) end,
), lists:seq(1, ?MAX_SESSIONS_PER_USER)
?assertEqual({error, too_many_sessions}, check_user_session_limit(UserId)), ),
delete_user_session_count(UserId). ?assertEqual({error, too_many_sessions}, check_user_session_limit(UserId)),
delete_user_session_count(UserId)
end).
check_user_session_limit_disabled_by_env_test() ->
OldValue = os:getenv("FLUXER_DISABLE_RATE_LIMITS"),
os:putenv("FLUXER_DISABLE_RATE_LIMITS", "true"),
try
ensure_tables(),
UserId = 900005,
delete_user_session_count(UserId),
lists:foreach(
fun(_) -> increment_user_sessions(UserId) end,
lists:seq(1, ?MAX_SESSIONS_PER_USER + 1)
),
?assertEqual(ok, check_user_session_limit(UserId)),
delete_user_session_count(UserId)
after
restore_env("FLUXER_DISABLE_RATE_LIMITS", OldValue)
end.
user_session_decrement_works_test() -> user_session_decrement_works_test() ->
ensure_tables(), with_rate_limits_enabled(fun() ->
UserId = 900003, ensure_tables(),
delete_user_session_count(UserId), UserId = 900003,
lists:foreach( delete_user_session_count(UserId),
fun(_) -> increment_user_sessions(UserId) end, lists:foreach(
lists:seq(1, ?MAX_SESSIONS_PER_USER) fun(_) -> increment_user_sessions(UserId) end,
), lists:seq(1, ?MAX_SESSIONS_PER_USER)
?assertEqual({error, too_many_sessions}, check_user_session_limit(UserId)), ),
decrement_user_sessions(UserId), ?assertEqual({error, too_many_sessions}, check_user_session_limit(UserId)),
?assertEqual(ok, check_user_session_limit(UserId)), decrement_user_sessions(UserId),
delete_user_session_count(UserId). ?assertEqual(ok, check_user_session_limit(UserId)),
delete_user_session_count(UserId)
end).
user_session_decrement_does_not_go_negative_test() -> user_session_decrement_does_not_go_negative_test() ->
ensure_tables(), ensure_tables(),
@@ -434,4 +430,18 @@ delete_user_session_count(UserId) ->
error:badarg -> ok error:badarg -> ok
end. end.
restore_env(Key, false) ->
os:unsetenv(Key);
restore_env(Key, Value) ->
os:putenv(Key, Value).
with_rate_limits_enabled(Fun) ->
OldValue = os:getenv("FLUXER_DISABLE_RATE_LIMITS"),
os:unsetenv("FLUXER_DISABLE_RATE_LIMITS"),
try
Fun()
after
restore_env("FLUXER_DISABLE_RATE_LIMITS", OldValue)
end.
-endif. -endif.
@@ -96,7 +96,10 @@ check_full_list_bot_rate_limit(true, false, _UserId, _GuildId) ->
check_full_list_bot_rate_limit(true, true, UserId, GuildId) when check_full_list_bot_rate_limit(true, true, UserId, GuildId) when
is_integer(UserId), UserId > 0, is_integer(GuildId), GuildId > 0 is_integer(UserId), UserId > 0, is_integer(GuildId), GuildId > 0
-> ->
check_bot_rate_limit_window(UserId, GuildId); case gateway_handler_rate_limit:rate_limits_disabled() of
true -> ok;
false -> check_bot_rate_limit_window(UserId, GuildId)
end;
check_full_list_bot_rate_limit(true, true, _UserId, _GuildId) -> check_full_list_bot_rate_limit(true, true, _UserId, _GuildId) ->
ok. ok.
+49 -2
View File
@@ -8,6 +8,7 @@
new/2, new/2,
push/2, push/2,
push/3, push/3,
push_trimmed/3,
push_front/2, push_front/2,
pop/1, pop/1,
pop_front/1, pop_front/1,
@@ -51,9 +52,15 @@ push(Item, D) ->
push(Item, entry_bytes(Item), D). push(Item, entry_bytes(Item), D).
-spec push(term(), non_neg_integer(), deque()) -> deque(). -spec push(term(), non_neg_integer(), deque()) -> deque().
push(Item, ItemBytes, #{rear := Rear, count := Count, bytes := Bytes} = D) -> push(Item, ItemBytes, D) ->
{D1, _Dropped} = push_trimmed(Item, ItemBytes, D),
D1.
-spec push_trimmed(term(), non_neg_integer(), deque()) -> {deque(), [term()]}.
push_trimmed(Item, ItemBytes, #{rear := Rear, count := Count, bytes := Bytes} = D) ->
D1 = D#{rear := [Item | Rear], count := Count + 1, bytes := Bytes + ItemBytes}, D1 = D#{rear := [Item | Rear], count := Count + 1, bytes := Bytes + ItemBytes},
trim_front(D1). {D2, Dropped} = trim_front_collect(D1, []),
{D2, lists:reverse(Dropped)}.
-spec push_front(term(), deque()) -> deque(). -spec push_front(term(), deque()) -> deque().
push_front(Item, #{front := Front, count := Count, bytes := Bytes} = D) -> push_front(Item, #{front := Front, count := Count, bytes := Bytes} = D) ->
@@ -162,6 +169,25 @@ trim_front(D) ->
{_, D2} -> trim_front(D2) {_, D2} -> trim_front(D2)
end. end.
-spec trim_front_collect(deque(), [term()]) -> {deque(), [term()]}.
trim_front_collect(
#{count := Count, max_count := MaxCount, max_bytes := MaxBytes} = D, Acc
) when
Count =< MaxCount, MaxBytes =:= 0
->
{D, Acc};
trim_front_collect(
#{count := Count, max_count := MaxCount, bytes := Bytes, max_bytes := MaxBytes} = D, Acc
) when
Count =< MaxCount, Bytes =< MaxBytes
->
{D, Acc};
trim_front_collect(D, Acc) ->
case pop_front(D) of
empty -> {D, Acc};
{Item, D2} -> trim_front_collect(D2, [Item | Acc])
end.
-spec trim_rear(deque()) -> deque(). -spec trim_rear(deque()) -> deque().
trim_rear( trim_rear(
#{count := Count, max_count := MaxCount, max_bytes := MaxBytes} = D #{count := Count, max_count := MaxCount, max_bytes := MaxBytes} = D
@@ -221,6 +247,27 @@ push_trims_at_bound_test() ->
List = to_list(D1), List = to_list(D1),
?assertEqual([b, c, d], List). ?assertEqual([b, c, d], List).
push_trimmed_returns_dropped_items_test() ->
D0 = new(3, 0),
D1 = lists:foldl(fun push/2, D0, [a, b, c]),
{D2, Dropped} = push_trimmed(d, entry_bytes(d), D1),
?assertEqual([a], Dropped),
?assertEqual([b, c, d], to_list(D2)).
push_trimmed_returns_no_dropped_items_below_bound_test() ->
{D1, Dropped} = push_trimmed(a, entry_bytes(a), new(3, 0)),
?assertEqual([], Dropped),
?assertEqual([a], to_list(D1)).
push_trimmed_returns_dropped_items_oldest_first_test() ->
Big = lists:seq(1, 100),
Smalls = [[1], [2], [3]],
D0 = new(1000, entry_bytes(Big)),
D1 = lists:foldl(fun push/2, D0, Smalls),
{D2, Dropped} = push_trimmed(Big, entry_bytes(Big), D1),
?assertEqual(Smalls, Dropped),
?assertEqual([Big], to_list(D2)).
pop_front_test() -> pop_front_test() ->
assert_pop_sequence(fun pop_front/1, [1, 2, 3]). assert_pop_sequence(fun pop_front/1, [1, 2, 3]).
+10 -4
View File
@@ -97,7 +97,10 @@ handle_call(get_current_visible_presence, _From, State) ->
{reply, presence_broadcast:current_visible_presence(State), State}; {reply, presence_broadcast:current_visible_presence(State), State};
handle_call({terminate_session, SessionIdHashes}, _From, State) when is_list(SessionIdHashes) -> handle_call({terminate_session, SessionIdHashes}, _From, State) when is_list(SessionIdHashes) ->
presence_connect:handle_terminate_session_call(binary_list(SessionIdHashes), State); presence_connect:handle_terminate_session_call(binary_list(SessionIdHashes), State);
handle_call({dispatch, EventAtom, Data}, _From, State) when is_atom(EventAtom), is_map(Data) -> handle_call({dispatch, EventAtom, Data}, _From, State) when
is_atom(EventAtom), is_map(Data);
is_atom(EventAtom), is_list(Data)
->
handle_dispatch_call(EventAtom, Data, State); handle_dispatch_call(EventAtom, Data, State);
handle_call({join_guild, GuildId}, _From, State) when is_integer(GuildId) -> handle_call({join_guild, GuildId}, _From, State) when is_integer(GuildId) ->
presence_connect:handle_join_guild(GuildId, State); presence_connect:handle_join_guild(GuildId, State);
@@ -114,7 +117,10 @@ handle_call(_, _From, State) ->
{reply, ok, State}. {reply, ok, State}.
-spec handle_cast(term(), state()) -> {noreply, state()}. -spec handle_cast(term(), state()) -> {noreply, state()}.
handle_cast({dispatch, Event, Data}, State) when is_atom(Event), is_map(Data) -> handle_cast({dispatch, Event, Data}, State) when
is_atom(Event), is_map(Data);
is_atom(Event), is_list(Data)
->
handle_dispatch_cast(Event, Data, State); handle_dispatch_cast(Event, Data, State);
handle_cast(presence_rejoin, State) -> handle_cast(presence_rejoin, State) ->
handle_presence_rejoin(State); handle_presence_rejoin(State);
@@ -232,12 +238,12 @@ cast_guild_op(Fun, GuildId, State) ->
{reply, _Reply, NewState} = Fun(GuildId, State), {reply, _Reply, NewState} = Fun(GuildId, State),
NewState. NewState.
-spec handle_dispatch_call(atom(), map(), state()) -> {reply, ok, state()}. -spec handle_dispatch_call(atom(), map() | list(), state()) -> {reply, ok, state()}.
handle_dispatch_call(EventAtom, Data, State) -> handle_dispatch_call(EventAtom, Data, State) ->
presence_broadcast:dispatch_to_all_sessions(EventAtom, Data, State), presence_broadcast:dispatch_to_all_sessions(EventAtom, Data, State),
{reply, ok, process_dispatch_event(EventAtom, Data, State)}. {reply, ok, process_dispatch_event(EventAtom, Data, State)}.
-spec handle_dispatch_cast(atom(), map(), state()) -> {noreply, state()}. -spec handle_dispatch_cast(atom(), map() | list(), state()) -> {noreply, state()}.
handle_dispatch_cast(Event, Data, State) -> handle_dispatch_cast(Event, Data, State) ->
presence_broadcast:dispatch_to_all_sessions(Event, Data, State), presence_broadcast:dispatch_to_all_sessions(Event, Data, State),
{noreply, process_dispatch_event(Event, Data, State)}. {noreply, process_dispatch_event(Event, Data, State)}.
@@ -101,12 +101,12 @@ dispatch_initial_presences(Presences, State) ->
SessionPids SessionPids
). ).
-spec dispatch_to_all_sessions(atom(), map(), state()) -> ok. -spec dispatch_to_all_sessions(atom(), map() | list(), state()) -> ok.
dispatch_to_all_sessions(EventAtom, Data, State) -> dispatch_to_all_sessions(EventAtom, Data, State) ->
SessionPids = presence_connect:collect_session_pids(State), SessionPids = presence_connect:collect_session_pids(State),
dispatch_with_backpressure(SessionPids, EventAtom, Data). dispatch_with_backpressure(SessionPids, EventAtom, Data).
-spec dispatch_with_backpressure([pid()], atom(), map()) -> ok. -spec dispatch_with_backpressure([pid()], atom(), map() | list()) -> ok.
dispatch_with_backpressure(SessionPids, EventAtom, Data) -> dispatch_with_backpressure(SessionPids, EventAtom, Data) ->
Msg = {dispatch, EventAtom, Data}, Msg = {dispatch, EventAtom, Data},
lists:foreach( lists:foreach(
+4 -2
View File
@@ -132,11 +132,13 @@ handle_cast({dispatch, Event, {pre_encoded, EncodedData} = Data}, State) when
-> ->
session_dispatch:handle_dispatch(Event, Data, State); session_dispatch:handle_dispatch(Event, Data, State);
handle_cast({dispatch, Event, Data}, State) when handle_cast({dispatch, Event, Data}, State) when
is_atom(Event), is_map(Data) is_atom(Event), is_map(Data);
is_atom(Event), is_list(Data)
-> ->
session_dispatch:handle_dispatch(Event, Data, State); session_dispatch:handle_dispatch(Event, Data, State);
handle_cast({dispatch, Event, Data}, State) when handle_cast({dispatch, Event, Data}, State) when
is_binary(Event), is_map(Data) is_binary(Event), is_map(Data);
is_binary(Event), is_list(Data)
-> ->
session_dispatch:handle_dispatch(Event, Data, State); session_dispatch:handle_dispatch(Event, Data, State);
handle_cast({initial_global_presences, Presences}, State) -> handle_cast({initial_global_presences, Presences}, State) ->
+142 -29
View File
@@ -6,7 +6,8 @@
-export([ -export([
handle_dispatch/3, handle_dispatch/3,
flush_all_pending_presences/1, flush_all_pending_presences/1,
flush_reaction_buffer/1 flush_reaction_buffer/1,
replay_floor_after_eviction/2
]). ]).
-export_type([session_state/0, event/0]). -export_type([session_state/0, event/0]).
@@ -18,7 +19,7 @@
-type session_state() :: session:session_state(). -type session_state() :: session:session_state().
-type event() :: atom() | binary(). -type event() :: atom() | binary().
-spec handle_dispatch(event(), map() | {pre_encoded, binary()}, session_state()) -> -spec handle_dispatch(event(), map() | list() | {pre_encoded, binary()}, session_state()) ->
{noreply, session_state()}. {noreply, session_state()}.
handle_dispatch(Event, {pre_encoded, _} = Data, State) -> handle_dispatch(Event, {pre_encoded, _} = Data, State) ->
case case
@@ -35,7 +36,9 @@ handle_dispatch(Event, Data, State) ->
false -> route_dispatch(Event, Data, State) false -> route_dispatch(Event, Data, State)
end. end.
-spec should_skip_for_shard(event(), map() | {pre_encoded, binary()}, session_state()) -> -spec should_skip_for_shard(
event(), map() | list() | {pre_encoded, binary()}, session_state()
) ->
boolean(). boolean().
should_skip_for_shard(Event, {pre_encoded, EncodedData}, State) -> should_skip_for_shard(Event, {pre_encoded, EncodedData}, State) ->
case shard_filter_active(State) of case shard_filter_active(State) of
@@ -46,7 +49,9 @@ should_skip_for_shard(Event, Data, State) when is_map(Data) ->
case shard_filter_active(State) of case shard_filter_active(State) of
true -> not has_guild_context(Event, Data); true -> not has_guild_context(Event, Data);
false -> false false -> false
end. end;
should_skip_for_shard(_Event, Data, State) when is_list(Data) ->
shard_filter_active(State).
-spec should_skip_pre_encoded_for_shard(event(), binary()) -> boolean(). -spec should_skip_pre_encoded_for_shard(event(), binary()) -> boolean().
should_skip_pre_encoded_for_shard(Event, EncodedData) -> should_skip_pre_encoded_for_shard(Event, EncodedData) ->
@@ -74,11 +79,20 @@ decode_pre_encoded_data(EncodedData) ->
-spec has_guild_context(event(), map()) -> boolean(). -spec has_guild_context(event(), map()) -> boolean().
has_guild_context(Event, Data) -> has_guild_context(Event, Data) ->
case has_nonempty_field(<<"guild_id">>, Data) of case session_reply_event(Event) orelse has_nonempty_field(<<"guild_id">>, Data) of
true -> true; true -> true;
false -> guild_id_event(Event) andalso has_nonempty_field(<<"id">>, Data) false -> guild_id_event(Event) andalso has_nonempty_field(<<"id">>, Data)
end. end.
-spec session_reply_event(event()) -> boolean().
session_reply_event(Event) ->
case event_name(Event) of
<<"RATE_LIMITED">> -> true;
<<"GUILD_COUNTS_UPDATE">> -> true;
<<"CHANNEL_MEMBER_COUNTS_UPDATE">> -> true;
_Other -> false
end.
-spec has_nonempty_field(binary(), map()) -> boolean(). -spec has_nonempty_field(binary(), map()) -> boolean().
has_nonempty_field(Key, Data) -> has_nonempty_field(Key, Data) ->
case maps:get(Key, Data, undefined) of case maps:get(Key, Data, undefined) of
@@ -97,7 +111,7 @@ guild_id_event(Event) ->
_Other -> false _Other -> false
end. end.
-spec route_dispatch(event(), map(), session_state()) -> {noreply, session_state()}. -spec route_dispatch(event(), map() | list(), session_state()) -> {noreply, session_state()}.
route_dispatch(Event, Data, State) -> route_dispatch(Event, Data, State) ->
case session_dispatch_voice:should_buffer_reaction(Event, State) of case session_dispatch_voice:should_buffer_reaction(Event, State) of
true -> true ->
@@ -106,7 +120,8 @@ route_dispatch(Event, Data, State) ->
route_after_reaction(Event, Data, State) route_after_reaction(Event, Data, State)
end. end.
-spec route_after_reaction(event(), map(), session_state()) -> {noreply, session_state()}. -spec route_after_reaction(event(), map() | list(), session_state()) ->
{noreply, session_state()}.
route_after_reaction(Event, Data, State) -> route_after_reaction(Event, Data, State) ->
case session_dispatch_voice:maybe_cancel_buffered_reaction(Event, Data, State) of case session_dispatch_voice:maybe_cancel_buffered_reaction(Event, Data, State) of
{cancelled, NewState} -> {cancelled, NewState} ->
@@ -115,7 +130,8 @@ route_after_reaction(Event, Data, State) ->
route_after_cancel(Event, Data, State) route_after_cancel(Event, Data, State)
end. end.
-spec route_after_cancel(event(), map(), session_state()) -> {noreply, session_state()}. -spec route_after_cancel(event(), map() | list(), session_state()) ->
{noreply, session_state()}.
route_after_cancel(Event, Data, State) -> route_after_cancel(Event, Data, State) ->
case session_dispatch_presence:should_buffer_presence(Event, Data, State) of case session_dispatch_presence:should_buffer_presence(Event, Data, State) of
true -> true ->
@@ -124,7 +140,8 @@ route_after_cancel(Event, Data, State) ->
do_handle_dispatch(Event, Data, State) do_handle_dispatch(Event, Data, State)
end. end.
-spec do_handle_dispatch(event(), map(), session_state()) -> {noreply, session_state()}. -spec do_handle_dispatch(event(), map() | list(), session_state()) ->
{noreply, session_state()}.
do_handle_dispatch(Event, Data, State) -> do_handle_dispatch(Event, Data, State) ->
Seq = maps:get(seq, State), Seq = maps:get(seq, State),
NewSeq = Seq + 1, NewSeq = Seq + 1,
@@ -135,55 +152,59 @@ do_handle_dispatch(Event, Data, State) ->
dispatch_replayable_event(Event, Data, NewSeq, State) dispatch_replayable_event(Event, Data, NewSeq, State)
end. end.
-spec dispatch_replayable_event(event(), map(), non_neg_integer(), session_state()) -> -spec dispatch_replayable_event(event(), map() | list(), non_neg_integer(), session_state()) ->
{noreply, session_state()}. {noreply, session_state()}.
dispatch_replayable_event(Event, Data, NewSeq, State) -> dispatch_replayable_event(Event, Data, NewSeq, State) ->
Request = #{event => Event, data => Data, seq => NewSeq}, Request = #{event => Event, data => Data, seq => NewSeq},
RequestBytes = buffer_entry_bytes(Request), RequestBytes = buffer_entry_bytes(Request),
case is_oversized_event(RequestBytes) of case is_oversized_event(RequestBytes) of
true -> true ->
dispatch_without_replay(Event, Data, NewSeq, State); dispatch_without_replay(Event, Data, NewSeq, State#{replay_floor => NewSeq});
false -> false ->
dispatch_with_replay(Event, Data, NewSeq, Request, RequestBytes, State) dispatch_with_replay(Event, Data, NewSeq, Request, RequestBytes, State)
end. end.
-spec dispatch_with_replay( -spec dispatch_with_replay(
event(), map(), non_neg_integer(), map(), non_neg_integer(), session_state() event(), map() | list(), non_neg_integer(), map(), non_neg_integer(), session_state()
) -> ) ->
{noreply, session_state()}. {noreply, session_state()}.
dispatch_with_replay(Event, Data, NewSeq, Request, RequestBytes, State) -> dispatch_with_replay(Event, Data, NewSeq, Request, RequestBytes, State) ->
Buffer = maps:get(buffer, State), Buffer = maps:get(buffer, State),
NewBuffer = Deque =
case is_list(Buffer) of case is_list(Buffer) of
true -> true ->
D = limited_deque:from_list( limited_deque:from_list(
Buffer, ?MAX_EVENT_BUFFER_SIZE, ?MAX_TOTAL_BUFFER_BYTES Buffer, ?MAX_EVENT_BUFFER_SIZE, ?MAX_TOTAL_BUFFER_BYTES
), );
limited_deque:push(Request, RequestBytes, D);
false -> false ->
limited_deque:push(Request, RequestBytes, Buffer) Buffer
end, end,
{NewBuffer, Dropped} = limited_deque:push_trimmed(Request, RequestBytes, Deque),
send_to_socket(maps:get(socket_pid, State, undefined), Event, Data, NewSeq), send_to_socket(maps:get(socket_pid, State, undefined), Event, Data, NewSeq),
StateAfterMain = apply_state_updates(Event, Data, State, #{ StateAfterMain = apply_state_updates(Event, Data, State, #{
seq => NewSeq, buffer => NewBuffer, buffer_bytes => limited_deque:bytes(NewBuffer) seq => NewSeq,
buffer => NewBuffer,
buffer_bytes => limited_deque:bytes(NewBuffer),
replay_floor => replay_floor_after_eviction(Dropped, maps:get(replay_floor, State, 0))
}), }),
finalize_dispatch(Event, Data, StateAfterMain). finalize_dispatch(Event, Data, StateAfterMain).
-spec dispatch_without_replay(event(), map(), non_neg_integer(), session_state()) -> -spec dispatch_without_replay(event(), map() | list(), non_neg_integer(), session_state()) ->
{noreply, session_state()}. {noreply, session_state()}.
dispatch_without_replay(Event, Data, NewSeq, State) -> dispatch_without_replay(Event, Data, NewSeq, State) ->
send_to_socket(maps:get(socket_pid, State, undefined), Event, Data, NewSeq), send_to_socket(maps:get(socket_pid, State, undefined), Event, Data, NewSeq),
StateAfterMain = apply_state_updates(Event, Data, State, #{seq => NewSeq}), StateAfterMain = apply_state_updates(Event, Data, State, #{seq => NewSeq}),
finalize_dispatch(Event, Data, StateAfterMain). finalize_dispatch(Event, Data, StateAfterMain).
-spec apply_state_updates(event(), map(), session_state(), map()) -> session_state(). -spec apply_state_updates(event(), map() | list(), session_state(), map()) -> session_state().
apply_state_updates(Event, Data, State, Extra) -> apply_state_updates(Event, Data, State, Extra) ->
S1 = session_dispatch_guild:update_channels_map(Event, Data, State), S1 = session_dispatch_guild:update_channels_map(Event, Data, State),
S2 = session_dispatch_guild:update_dm_voice_states_map(Event, Data, S1), S2 = session_dispatch_guild:update_dm_voice_states_map(Event, Data, S1),
S3 = session_dispatch_guild:update_relationships_map(Event, Data, S2), S3 = session_dispatch_guild:update_relationships_map(Event, Data, S2),
maps:merge(S3, Extra). maps:merge(S3, Extra).
-spec finalize_dispatch(event(), map(), session_state()) -> {noreply, session_state()}. -spec finalize_dispatch(event(), map() | list(), session_state()) ->
{noreply, session_state()}.
finalize_dispatch(Event, Data, State) -> finalize_dispatch(Event, Data, State) ->
{S1, FlushedIds} = session_dispatch_presence:maybe_flush_pending_presences( {S1, FlushedIds} = session_dispatch_presence:maybe_flush_pending_presences(
Event, Data, State Event, Data, State
@@ -217,7 +238,7 @@ buffer_pre_encoded_event(Event, Data, NewSeq, State) ->
RequestBytes = buffer_entry_bytes(Request), RequestBytes = buffer_entry_bytes(Request),
case is_oversized_event(RequestBytes) of case is_oversized_event(RequestBytes) of
true -> true ->
State; State#{replay_floor => NewSeq};
false -> false ->
Buffer = maps:get(buffer, State), Buffer = maps:get(buffer, State),
Deque = Deque =
@@ -229,8 +250,16 @@ buffer_pre_encoded_event(Event, Data, NewSeq, State) ->
false -> false ->
Buffer Buffer
end, end,
NewBuffer = limited_deque:push(Request, RequestBytes, Deque), {NewBuffer, Dropped} = limited_deque:push_trimmed(
State#{buffer => NewBuffer, buffer_bytes => limited_deque:bytes(NewBuffer)} Request, RequestBytes, Deque
),
State#{
buffer => NewBuffer,
buffer_bytes => limited_deque:bytes(NewBuffer),
replay_floor => replay_floor_after_eviction(
Dropped, maps:get(replay_floor, State, 0)
)
}
end end
end. end.
@@ -265,16 +294,34 @@ is_oversized_event(RequestBytes) ->
-spec should_buffer_pre_encoded(event()) -> boolean(). -spec should_buffer_pre_encoded(event()) -> boolean().
should_buffer_pre_encoded(Event) -> should_buffer_pre_encoded(Event) ->
event_name(Event) =:= <<"VOICE_STATE_UPDATE">>. not should_skip_replay_buffer(Event).
-spec should_skip_replay_buffer(event()) -> boolean(). -spec should_skip_replay_buffer(event()) -> boolean().
should_skip_replay_buffer(Event) -> should_skip_replay_buffer(Event) ->
event_name(Event) =:= <<"GUILD_MEMBERS_CHUNK">>. case event_name(Event) of
<<"GUILD_MEMBERS_CHUNK">> -> true;
<<"GUILD_MEMBER_LIST_UPDATE">> -> true;
<<"GUILD_SYNC">> -> true;
_Other -> false
end.
-spec buffer_entry_bytes(term()) -> non_neg_integer(). -spec buffer_entry_bytes(term()) -> non_neg_integer().
buffer_entry_bytes(Request) -> buffer_entry_bytes(Request) ->
limited_deque:entry_bytes(Request). limited_deque:entry_bytes(Request).
-spec replay_floor_after_eviction([term()], non_neg_integer()) -> non_neg_integer().
replay_floor_after_eviction(Dropped, Floor) ->
lists:foldl(fun evicted_seq_max/2, Floor, Dropped).
-spec evicted_seq_max(term(), non_neg_integer()) -> non_neg_integer().
evicted_seq_max(Event, Floor) when is_map(Event) ->
case maps:get(seq, Event, undefined) of
Seq when is_integer(Seq), Seq > Floor -> Seq;
_Other -> Floor
end;
evicted_seq_max(_Event, Floor) ->
Floor.
-spec send_to_socket(pid() | undefined, event(), term(), non_neg_integer()) -> ok. -spec send_to_socket(pid() | undefined, event(), term(), non_neg_integer()) -> ok.
send_to_socket(undefined, _Event, _Data, _Seq) -> send_to_socket(undefined, _Event, _Data, _Seq) ->
ok; ok;
@@ -282,7 +329,7 @@ send_to_socket(Pid, Event, Data, Seq) when is_pid(Pid) ->
Pid ! {dispatch, Event, guild_data_wire:payload(Data), Seq}, Pid ! {dispatch, Event, guild_data_wire:payload(Data), Seq},
ok. ok.
-spec should_ignore_event(event(), map() | {pre_encoded, binary()}, session_state()) -> -spec should_ignore_event(event(), map() | list() | {pre_encoded, binary()}, session_state()) ->
boolean(). boolean().
should_ignore_event(Event, Data, State) -> should_ignore_event(Event, Data, State) ->
IgnoredEvents = maps:get(ignored_events, State, #{}), IgnoredEvents = maps:get(ignored_events, State, #{}),
@@ -306,7 +353,9 @@ event_name(Event) when is_atom(Event) ->
event_name(_) -> event_name(_) ->
undefined. undefined.
-spec ignored_event_must_dispatch(event(), map() | {pre_encoded, binary()}, session_state()) -> -spec ignored_event_must_dispatch(
event(), map() | list() | {pre_encoded, binary()}, session_state()
) ->
boolean(). boolean().
ignored_event_must_dispatch(message_create, {pre_encoded, EncodedData}, State) -> ignored_event_must_dispatch(message_create, {pre_encoded, EncodedData}, State) ->
case decode_pre_encoded_data(EncodedData) of case decode_pre_encoded_data(EncodedData) of
@@ -378,7 +427,8 @@ oversized_event_sent_but_not_buffered_test() ->
LargeData = make_large_data(), LargeData = make_large_data(),
{noreply, S1} = do_handle_dispatch(guild_create, LargeData, base_state(#{})), {noreply, S1} = do_handle_dispatch(guild_create, LargeData, base_state(#{})),
?assertEqual([], maps:get(buffer, S1, [])), ?assertEqual([], maps:get(buffer, S1, [])),
?assertEqual(1, maps:get(seq, S1)). ?assertEqual(1, maps:get(seq, S1)),
?assertEqual(1, maps:get(replay_floor, S1)).
guild_members_chunk_sent_but_not_buffered_test() -> guild_members_chunk_sent_but_not_buffered_test() ->
ChunkData = #{ ChunkData = #{
@@ -435,8 +485,51 @@ normal_event_buffered_test() ->
should_skip_replay_buffer_test() -> should_skip_replay_buffer_test() ->
?assertEqual(true, should_skip_replay_buffer(guild_members_chunk)), ?assertEqual(true, should_skip_replay_buffer(guild_members_chunk)),
?assertEqual(true, should_skip_replay_buffer(<<"GUILD_MEMBERS_CHUNK">>)), ?assertEqual(true, should_skip_replay_buffer(<<"GUILD_MEMBERS_CHUNK">>)),
?assertEqual(true, should_skip_replay_buffer(guild_member_list_update)),
?assertEqual(true, should_skip_replay_buffer(guild_sync)),
?assertEqual(false, should_skip_replay_buffer(message_create)). ?assertEqual(false, should_skip_replay_buffer(message_create)).
should_buffer_pre_encoded_test() ->
?assertEqual(false, should_buffer_pre_encoded(guild_members_chunk)),
?assertEqual(false, should_buffer_pre_encoded(guild_member_list_update)),
?assertEqual(false, should_buffer_pre_encoded(guild_sync)),
?assertEqual(true, should_buffer_pre_encoded(voice_state_update)),
?assertEqual(true, should_buffer_pre_encoded(message_create)),
?assertEqual(true, should_buffer_pre_encoded(channel_update)).
guild_member_list_update_map_form_not_buffered_test() ->
Data = #{<<"guild_id">> => <<"123">>, <<"ops">> => []},
{noreply, S1} = do_handle_dispatch(
guild_member_list_update, Data, base_state(#{replay_floor => 0})
),
?assertEqual([], maps:get(buffer, S1, [])),
?assertEqual(1, maps:get(seq, S1)),
?assertEqual(0, maps:get(replay_floor, S1)).
replay_floor_stays_zero_without_eviction_test() ->
{noreply, S1} = do_handle_dispatch(
message_create, #{<<"content">> => <<"hello">>}, base_state(#{})
),
?assertEqual(0, maps:get(replay_floor, S1)).
replay_floor_tracks_evicted_seq_test() ->
State0 = base_state(#{buffer => limited_deque:new(2, 0)}),
State3 = lists:foldl(
fun(_N, S) ->
{noreply, Next} = do_handle_dispatch(
message_create, #{<<"content">> => <<"hello">>}, S
),
Next
end,
State0,
lists:seq(1, 3)
),
?assertEqual(1, maps:get(replay_floor, State3)),
?assertEqual(
[2, 3],
[maps:get(seq, E) || E <- limited_deque:to_list(maps:get(buffer, State3))]
).
needs_state_update_test() -> needs_state_update_test() ->
lists:foreach( lists:foreach(
fun(E) -> ?assertEqual(true, needs_state_update(E)) end, fun(E) -> ?assertEqual(true, needs_state_update(E)) end,
@@ -494,6 +587,26 @@ guildless_dispatch_allowed_for_shard_zero_test() ->
?assert(false, dispatch_not_received) ?assert(false, dispatch_not_received)
end. end.
list_payload_dispatch_reaches_socket_test() ->
drain_mailbox(),
Data = [<<"1">>, <<"2">>],
BaseState = base_state(#{socket_pid => self()}),
{noreply, S1} = handle_dispatch(user_pinned_dms_update, Data, BaseState),
?assertEqual(1, maps:get(seq, S1)),
receive
{dispatch, user_pinned_dms_update, ReceivedData, 1} ->
?assertEqual(Data, ReceivedData)
after 100 ->
?assert(false, dispatch_not_received)
end.
list_payload_dispatch_skipped_for_nonzero_shard_test() ->
drain_mailbox(),
BaseState = base_state(#{socket_pid => self(), shard => {1, 2}}),
{noreply, S1} = handle_dispatch(user_pinned_dms_update, [<<"1">>], BaseState),
?assertEqual(0, maps:get(seq, S1)),
assert_no_dispatch().
pre_encoded_guildless_dispatch_skipped_for_nonzero_shard_test() -> pre_encoded_guildless_dispatch_skipped_for_nonzero_shard_test() ->
drain_mailbox(), drain_mailbox(),
Encoded = iolist_to_binary( Encoded = iolist_to_binary(
@@ -16,7 +16,7 @@
-type user_id() :: session:user_id(). -type user_id() :: session:user_id().
-type channel_event() :: channel_create | channel_update. -type channel_event() :: channel_create | channel_update.
-spec update_channels_map(event(), map(), session_state()) -> session_state(). -spec update_channels_map(event(), map() | list(), session_state()) -> session_state().
update_channels_map(channel_create, Data, State) when is_map(Data) -> update_channels_map(channel_create, Data, State) when is_map(Data) ->
maybe_add_dm_channel(channel_create, Data, State); maybe_add_dm_channel(channel_create, Data, State);
update_channels_map(channel_update, Data, State) when is_map(Data) -> update_channels_map(channel_update, Data, State) when is_map(Data) ->
@@ -33,7 +33,7 @@ update_channels_map(channel_recipient_remove, Data, State) when is_map(Data) ->
update_channels_map(_Event, _Data, State) -> update_channels_map(_Event, _Data, State) ->
State. State.
-spec update_dm_voice_states_map(event(), map(), session_state()) -> session_state(). -spec update_dm_voice_states_map(event(), map() | list(), session_state()) -> session_state().
update_dm_voice_states_map(voice_state_update, Data, State) when is_map(Data) -> update_dm_voice_states_map(voice_state_update, Data, State) when is_map(Data) ->
update_dm_voice_state(Data, State); update_dm_voice_state(Data, State);
update_dm_voice_states_map(call_create, Data, State) when is_map(Data) -> update_dm_voice_states_map(call_create, Data, State) when is_map(Data) ->
@@ -359,7 +359,7 @@ add_unique_user_by_id(UserMap, Id, List) ->
false -> [UserMap | List] false -> [UserMap | List]
end. end.
-spec update_relationships_map(event(), map(), session_state()) -> session_state(). -spec update_relationships_map(event(), map() | list(), session_state()) -> session_state().
update_relationships_map(relationship_add, Data, State) -> update_relationships_map(relationship_add, Data, State) ->
upsert_relationship(Data, State); upsert_relationship(Data, State);
update_relationships_map(relationship_update, Data, State) -> update_relationships_map(relationship_update, Data, State) ->
@@ -27,7 +27,7 @@
-type event() :: atom() | binary(). -type event() :: atom() | binary().
-type user_id() :: session:user_id(). -type user_id() :: session:user_id().
-spec should_buffer_presence(event(), map(), session_state()) -> boolean(). -spec should_buffer_presence(event(), map() | list(), session_state()) -> boolean().
should_buffer_presence(presence_update, Data, State) -> should_buffer_presence(presence_update, Data, State) ->
case maps:get(suppress_presence_updates, State, true) of case maps:get(suppress_presence_updates, State, true) of
true -> true ->
@@ -98,7 +98,7 @@ buffer_presence(Event, Data, State) ->
NewPending = queue:in(Entry, Trimmed), NewPending = queue:in(Entry, Trimmed),
State#{pending_presences => NewPending}. State#{pending_presences => NewPending}.
-spec maybe_flush_pending_presences(event(), map(), session_state()) -> -spec maybe_flush_pending_presences(event(), map() | list(), session_state()) ->
{session_state(), [user_id()]}. {session_state(), [user_id()]}.
maybe_flush_pending_presences(relationship_add, Data, State) -> maybe_flush_pending_presences(relationship_add, Data, State) ->
maybe_flush_relationship_pending_presences(Data, State); maybe_flush_relationship_pending_presences(Data, State);
@@ -187,12 +187,17 @@ dispatch_presence_now(P, State) ->
false -> false ->
Buffer Buffer
end, end,
NewBuffer = limited_deque:push(Request, Deque), {NewBuffer, Dropped} = limited_deque:push_trimmed(
Request, limited_deque:entry_bytes(Request), Deque
),
send_to_socket(SocketPid, Event, Data, NewSeq), send_to_socket(SocketPid, Event, Data, NewSeq),
State#{ State#{
seq => NewSeq, seq => NewSeq,
buffer => NewBuffer, buffer => NewBuffer,
buffer_bytes => limited_deque:bytes(NewBuffer) buffer_bytes => limited_deque:bytes(NewBuffer),
replay_floor => session_dispatch:replay_floor_after_eviction(
Dropped, maps:get(replay_floor, State, 0)
)
}. }.
-spec flush_all_pending_presences(session_state()) -> session_state(). -spec flush_all_pending_presences(session_state()) -> session_state().
@@ -268,7 +273,7 @@ trim_queue_from_front(Queue, MaxLen) ->
false -> Queue false -> Queue
end. end.
-spec send_to_socket(pid() | undefined, event(), map(), non_neg_integer()) -> ok. -spec send_to_socket(pid() | undefined, event(), map() | list(), non_neg_integer()) -> ok.
send_to_socket(undefined, _Event, _Data, _Seq) -> send_to_socket(undefined, _Event, _Data, _Seq) ->
ok; ok;
send_to_socket(Pid, Event, Data, Seq) when is_pid(Pid) -> send_to_socket(Pid, Event, Data, Seq) when is_pid(Pid) ->
@@ -473,6 +478,33 @@ pending_presence_buffer_drops_oldest_when_full_test() ->
?assertEqual(3, maps:get(user_id, hd(Pending))), ?assertEqual(3, maps:get(user_id, hd(Pending))),
?assertEqual(Total, maps:get(user_id, lists:last(Pending))). ?assertEqual(Total, maps:get(user_id, lists:last(Pending))).
presence_flush_eviction_raises_replay_floor_test() ->
Base = #{
user_id => 1,
seq => 0,
buffer => limited_deque:new(2, 0),
socket_pid => undefined,
pending_presences => queue:new()
},
Filled = lists:foldl(
fun(N, Acc) ->
buffer_presence(
presence_update,
presence_status_data(integer_to_binary(N), <<"online">>),
Acc
)
end,
Base,
lists:seq(2, 4)
),
Flushed = flush_all_pending_presences(Filled),
?assertEqual(3, maps:get(seq, Flushed)),
?assertEqual(1, maps:get(replay_floor, Flushed)),
?assertEqual(
[2, 3],
[maps:get(seq, E) || E <- limited_deque:to_list(maps:get(buffer, Flushed))]
).
collect_dispatched_statuses(0) -> collect_dispatched_statuses(0) ->
[]; [];
collect_dispatched_statuses(N) -> collect_dispatched_statuses(N) ->
@@ -39,7 +39,7 @@ buffer_reaction(Data, State) ->
end, end,
State#{reaction_buffer => NewBuffer, reaction_buffer_timer => NewTimer}. State#{reaction_buffer => NewBuffer, reaction_buffer_timer => NewTimer}.
-spec maybe_cancel_buffered_reaction(event(), map(), session_state()) -> -spec maybe_cancel_buffered_reaction(event(), map() | list(), session_state()) ->
{cancelled, session_state()} | not_applicable. {cancelled, session_state()} | not_applicable.
maybe_cancel_buffered_reaction(message_reaction_remove, Data, State) -> maybe_cancel_buffered_reaction(message_reaction_remove, Data, State) ->
BufferQ = ensure_queue(maps:get(reaction_buffer, State, [])), BufferQ = ensure_queue(maps:get(reaction_buffer, State, [])),
@@ -232,6 +232,7 @@ extract_core_fields(
replay_payload_bytes => replay_payload_bytes(Buffer), replay_payload_bytes => replay_payload_bytes(Buffer),
seq => Seq, seq => Seq,
ack_seq => AckSeq, ack_seq => AckSeq,
replay_floor => init_replay_floor(normalize_seq(maps:get(replay_floor, D, 0)), Seq),
properties => Properties, properties => Properties,
status => Status, status => Status,
resume_status => maps:get(resume_status, D, Status), resume_status => maps:get(resume_status, D, Status),
@@ -313,6 +314,10 @@ init_ready(false, Ready) -> Ready.
init_ack_seq(AckSeq, Seq) when AckSeq =< Seq -> AckSeq; init_ack_seq(AckSeq, Seq) when AckSeq =< Seq -> AckSeq;
init_ack_seq(_AckSeq, Seq) -> Seq. init_ack_seq(_AckSeq, Seq) -> Seq.
-spec init_replay_floor(seq(), seq()) -> seq().
init_replay_floor(Floor, Seq) when Floor =< Seq -> Floor;
init_replay_floor(_Floor, Seq) -> Seq.
-spec schedule_timers(session_state()) -> ok. -spec schedule_timers(session_state()) -> ok.
schedule_timers(#{bot := Bot, guilds := GuildsMap}) -> schedule_timers(#{bot := Bot, guilds := GuildsMap}) ->
GuildIds = maps:keys(GuildsMap), GuildIds = maps:keys(GuildsMap),
@@ -465,6 +470,17 @@ build_state_loads_relationship_ids_from_ready_test() ->
?assertEqual(#{300 => 1}, maps:get(relationships, State)), ?assertEqual(#{300 => 1}, maps:get(relationships, State)),
?assert(maps:is_key(700, maps:get(channels, State))). ?assert(maps:is_key(700, maps:get(channels, State))).
build_state_restores_replay_floor_test() ->
Data = (base_session_data(#{}))#{seq => 20, replay_floor => 7},
?assertEqual(7, maps:get(replay_floor, build_state(Data))).
build_state_clamps_replay_floor_to_seq_test() ->
Data = (base_session_data(#{}))#{seq => 3, replay_floor => 99},
?assertEqual(3, maps:get(replay_floor, build_state(Data))).
build_state_defaults_replay_floor_to_zero_test() ->
?assertEqual(0, maps:get(replay_floor, build_state(base_session_data(#{})))).
base_session_data(Ready) -> base_session_data(Ready) ->
#{ #{
id => <<"session-init-test">>, id => <<"session-init-test">>,
@@ -324,13 +324,17 @@ buffer_event_acked(_Seq, _Event) ->
false. false.
-spec handle_resume(seq(), pid(), session_state()) -> -spec handle_resume(seq(), pid(), session_state()) ->
{reply, invalid_seq | {ok, [map()], seq()}, session_state()}. {reply, invalid_seq | not_resumable | {ok, [map()], seq()}, session_state()}.
handle_resume(Seq, _SocketPid, #{seq := CurrentSeq} = State) when Seq > CurrentSeq -> handle_resume(Seq, _SocketPid, #{seq := CurrentSeq} = State) when Seq > CurrentSeq ->
{reply, invalid_seq, State}; {reply, invalid_seq, State};
handle_resume(Seq, _SocketPid, #{ack_seq := AckSeq} = State) when handle_resume(Seq, _SocketPid, #{ack_seq := AckSeq} = State) when
is_integer(AckSeq), Seq < AckSeq is_integer(AckSeq), Seq < AckSeq
-> ->
{reply, invalid_seq, State}; {reply, invalid_seq, State};
handle_resume(Seq, _SocketPid, #{replay_floor := Floor} = State) when
is_integer(Floor), Seq < Floor
->
{reply, not_resumable, State};
handle_resume(Seq, SocketPid, #{seq := CurrentSeq} = State) -> handle_resume(Seq, SocketPid, #{seq := CurrentSeq} = State) ->
#{buffer := Buffer, id := SessionId, status := Status, afk := Afk, mobile := Mobile} = #{buffer := Buffer, id := SessionId, status := Status, afk := Afk, mobile := Mobile} =
State, State,
@@ -552,6 +556,7 @@ serialize_state(State) ->
version => maps:get(version, State), version => maps:get(version, State),
seq => maps:get(seq, State), seq => maps:get(seq, State),
ack_seq => maps:get(ack_seq, State), ack_seq => maps:get(ack_seq, State),
replay_floor => maps:get(replay_floor, State, 0),
properties => maps:get(properties, State), properties => maps:get(properties, State),
status => maps:get(status, State), status => maps:get(status, State),
resume_status => maps:get(resume_status, State, maps:get(status, State)), resume_status => maps:get(resume_status, State, maps:get(status, State)),
@@ -610,6 +615,7 @@ serialize_transfer_runtime(State) ->
relationships => maps:get(relationships, State, #{}), relationships => maps:get(relationships, State, #{}),
seq => maps:get(seq, State, 0), seq => maps:get(seq, State, 0),
ack_seq => AckSeq, ack_seq => AckSeq,
replay_floor => max(maps:get(replay_floor, State, 0), AckSeq),
buffer => Buffer, buffer => Buffer,
collected_guild_states => maps:get(collected_guild_states, State, []), collected_guild_states => maps:get(collected_guild_states, State, []),
collected_sessions => maps:get(collected_sessions, State, []), collected_sessions => maps:get(collected_sessions, State, []),
+3 -7
View File
@@ -22,7 +22,6 @@
connect_permission/0, connect_permission/0,
speak_permission/0, speak_permission/0,
stream_permission/0, stream_permission/0,
use_vad_permission/0,
read_message_history_permission/0, read_message_history_permission/0,
kick_members_permission/0, kick_members_permission/0,
ban_members_permission/0, ban_members_permission/0,
@@ -79,8 +78,7 @@ close_code_to_num(rate_limited) -> 4008;
close_code_to_num(session_timeout) -> 4009; close_code_to_num(session_timeout) -> 4009;
close_code_to_num(invalid_shard) -> 4010; close_code_to_num(invalid_shard) -> 4010;
close_code_to_num(sharding_required) -> 4011; close_code_to_num(sharding_required) -> 4011;
close_code_to_num(invalid_api_version) -> 4012; close_code_to_num(invalid_api_version) -> 4012.
close_code_to_num(ack_backpressure) -> 4013.
-spec dispatch_event_atom(atom() | binary()) -> atom() | binary(). -spec dispatch_event_atom(atom() | binary()) -> atom() | binary().
dispatch_event_atom(Event) when is_atom(Event) -> dispatch_event_atom(Event) when is_atom(Event) ->
@@ -142,9 +140,6 @@ speak_permission() -> 2097152.
-spec stream_permission() -> pos_integer(). -spec stream_permission() -> pos_integer().
stream_permission() -> 512. stream_permission() -> 512.
-spec use_vad_permission() -> pos_integer().
use_vad_permission() -> 33554432.
-spec read_message_history_permission() -> pos_integer(). -spec read_message_history_permission() -> pos_integer().
read_message_history_permission() -> 65536. read_message_history_permission() -> 65536.
@@ -178,7 +173,8 @@ close_code_to_num_test() ->
?assertEqual(4000, close_code_to_num(unknown_error)), ?assertEqual(4000, close_code_to_num(unknown_error)),
?assertEqual(4004, close_code_to_num(authentication_failed)), ?assertEqual(4004, close_code_to_num(authentication_failed)),
?assertEqual(4008, close_code_to_num(rate_limited)), ?assertEqual(4008, close_code_to_num(rate_limited)),
?assertEqual(4013, close_code_to_num(ack_backpressure)). ?assertEqual(4012, close_code_to_num(invalid_api_version)),
?assertError(function_clause, close_code_to_num(ack_backpressure)).
status_type_atom_binary_to_atom_test() -> status_type_atom_binary_to_atom_test() ->
?assertEqual(online, status_type_atom(<<>>)), ?assertEqual(online, status_type_atom(<<>>)),
+1 -15
View File
@@ -16,8 +16,7 @@
get_field/2, get_field/2,
get_field/3, get_field/3,
get_required_field/3, get_required_field/3,
get_optional_field/3, get_optional_field/3
error_category_to_close_code/1
]). ]).
-spec validate_snowflake(term()) -> {ok, pos_integer()} | {error, atom(), atom()}. -spec validate_snowflake(term()) -> {ok, pos_integer()} | {error, atom(), atom()}.
@@ -144,14 +143,6 @@ get_optional_field(FieldName, Map, Validator) ->
Validator(Value) Validator(Value)
end. end.
-spec error_category_to_close_code(atom()) -> integer().
error_category_to_close_code(rate_limited) ->
constants:close_code_to_num(rate_limited);
error_category_to_close_code(auth_failed) ->
constants:close_code_to_num(authentication_failed);
error_category_to_close_code(_) ->
constants:close_code_to_num(unknown_error).
-ifdef(TEST). -ifdef(TEST).
-include_lib("eunit/include/eunit.hrl"). -include_lib("eunit/include/eunit.hrl").
@@ -221,9 +212,4 @@ get_optional_field_test() ->
?assertEqual({ok, 123}, get_optional_field(<<"id">>, Map, Validator)), ?assertEqual({ok, 123}, get_optional_field(<<"id">>, Map, Validator)),
?assertEqual({ok, undefined}, get_optional_field(<<"missing">>, Map, Validator)). ?assertEqual({ok, undefined}, get_optional_field(<<"missing">>, Map, Validator)).
error_category_to_close_code_test() ->
?assertEqual(4008, error_category_to_close_code(rate_limited)),
?assertEqual(4004, error_category_to_close_code(auth_failed)),
?assertEqual(4000, error_category_to_close_code(unknown)).
-endif. -endif.
+77 -24
View File
@@ -13,6 +13,39 @@ websocket_info_session_reconnect_sends_reconnect_then_close_test() ->
<<"Session drain requested; reconnect to continue">> <<"Session drain requested; reconnect to continue">>
). ).
websocket_info_dispatch_list_payload_emits_frame_test() ->
Decoded = decode_dispatch_frame(
gateway_handler:websocket_info(
{dispatch, user_pinned_dms_update, [<<"1">>, <<"2">>], 1}, new_json_state()
)
),
?assertEqual(<<"USER_PINNED_DMS_UPDATE">>, maps:get(<<"t">>, Decoded)),
?assertEqual([<<"1">>, <<"2">>], maps:get(<<"d">>, Decoded)),
?assertEqual(1, maps:get(<<"s">>, Decoded)).
websocket_info_sessions_replace_list_payload_emits_frame_test() ->
Sessions = [
#{<<"session_id">> => <<"abc">>, <<"status">> => <<"online">>},
#{<<"session_id">> => <<"def">>, <<"status">> => <<"idle">>}
],
Decoded = decode_dispatch_frame(
gateway_handler:websocket_info(
{dispatch, sessions_replace, Sessions, 7}, new_json_state()
)
),
?assertEqual(<<"SESSIONS_REPLACE">>, maps:get(<<"t">>, Decoded)),
?assertEqual(Sessions, maps:get(<<"d">>, Decoded)),
?assertEqual(7, maps:get(<<"s">>, Decoded)).
websocket_info_empty_list_payload_emits_json_array_test() ->
Decoded = decode_dispatch_frame(
gateway_handler:websocket_info(
{dispatch, webauthn_credentials_update, [], 3}, new_json_state()
)
),
?assertEqual(<<"WEBAUTHN_CREDENTIALS_UPDATE">>, maps:get(<<"t">>, Decoded)),
?assertEqual([], maps:get(<<"d">>, Decoded)).
start_session_with_drain_guard_holds_during_drain_test() -> start_session_with_drain_guard_holds_during_drain_test() ->
Request = #{}, Request = #{},
assert_pending_identify( assert_pending_identify(
@@ -83,6 +116,17 @@ handle_session_start_result_server_error_closes_with_unknown_error_test() ->
), ),
?assertEqual(constants:close_code_to_num(unknown_error), CloseCode). ?assertEqual(constants:close_code_to_num(unknown_error), CloseCode).
handle_session_start_result_unknown_reason_closes_failed_to_start_test() ->
State = (gateway_handler:new_state())#{
version => 1, encoding => json, compress_ctx => gateway_compress:new_context(none)
},
{[{close, CloseCode, Reason}], _NewState} =
gateway_handler_identify:handle_session_start_result(
{error, something_unmapped}, State
),
?assertEqual(constants:close_code_to_num(unknown_error), CloseCode),
?assertEqual(<<"Failed to start session">>, Reason).
websocket_init_schedules_tokened_heartbeat_timer_test() -> websocket_init_schedules_tokened_heartbeat_timer_test() ->
{[{text, _Frame}], State} = gateway_handler:websocket_init(new_json_state()), {[{text, _Frame}], State} = gateway_handler:websocket_init(new_json_state()),
?assertMatch({_TimerRef, _Token}, maps:get(heartbeat_timer, State)), ?assertMatch({_TimerRef, _Token}, maps:get(heartbeat_timer, State)),
@@ -407,30 +451,32 @@ validate_identify_data_rejects_invalid_shard_test() ->
). ).
handle_identify_runs_identify_rate_check_at_zero_rollout_test() -> handle_identify_runs_identify_rate_check_at_zero_rollout_test() ->
session_abuse_protection:ensure_tables(), with_rate_limits_enabled(fun() ->
OldConfig = gateway_rollout_config:get(), session_abuse_protection:ensure_tables(),
IP = unique_test_peer_ip(<<"zero-rollout-identify">>), OldConfig = gateway_rollout_config:get(),
try IP = unique_test_peer_ip(<<"zero-rollout-identify">>),
persistent_term:put(gateway_rollout_config, OldConfig#{ try
<<"session_rollout_percentage">> => 0 persistent_term:put(gateway_rollout_config, OldConfig#{
}), <<"session_rollout_percentage">> => 0
lists:foreach( }),
fun(_) -> lists:foreach(
?assertEqual(ok, session_abuse_protection:check_identify_rate(IP)) fun(_) ->
end, ?assertEqual(ok, session_abuse_protection:check_identify_rate(IP))
lists:seq(1, ?TEST_IDENTIFY_MAX_PER_IP - 1) end,
), lists:seq(1, ?TEST_IDENTIFY_MAX_PER_IP - 1)
{ok, HeldState} = gateway_handler_identify:handle_identify( ),
valid_identify_data(#{}), IP, new_json_state() {ok, HeldState} = gateway_handler_identify:handle_identify(
), valid_identify_data(#{}), IP, new_json_state()
cancel_pending_identify_timer(HeldState), ),
?assertEqual( cancel_pending_identify_timer(HeldState),
{error, identify_rate_limited}, ?assertEqual(
session_abuse_protection:check_identify_rate(IP) {error, identify_rate_limited},
) session_abuse_protection:check_identify_rate(IP)
after )
persistent_term:put(gateway_rollout_config, OldConfig) after
end. persistent_term:put(gateway_rollout_config, OldConfig)
end
end).
handle_session_start_result_invalid_shard_closes_4010_test() -> handle_session_start_result_invalid_shard_closes_4010_test() ->
{[{close, CloseCode, _Reason}], _NewState} = {[{close, CloseCode, _Reason}], _NewState} =
@@ -486,6 +532,13 @@ valid_identify_data(Extra) ->
Extra Extra
). ).
decode_dispatch_frame(Result) ->
{[Frame], _NewState} = Result,
{_FrameType, EncodedFrame} = Frame,
{ok, DecodedFrame} = gateway_codec:decode(EncodedFrame, json),
?assertEqual(constants:opcode_to_num(dispatch), maps:get(<<"op">>, DecodedFrame)),
DecodedFrame.
assert_reconnect_close(Result, CloseReason) -> assert_reconnect_close(Result, CloseReason) ->
{[Frame, {close, CloseCode, CloseReason}], _NewState} = Result, {[Frame, {close, CloseCode, CloseReason}], _NewState} = Result,
?assertEqual(constants:close_code_to_num(unknown_error), CloseCode), ?assertEqual(constants:close_code_to_num(unknown_error), CloseCode),
@@ -91,6 +91,49 @@ handle_resume_missing_session_sends_invalid_session_without_api_test() ->
meck:unload(session_manager) meck:unload(session_manager)
end. end.
handle_resume_below_replay_floor_sends_invalid_session_test() ->
drain_mailbox(),
SessionPid = spawn(fun() -> fake_not_resumable_session_loop(<<"resume-token">>) end),
meck:new(session_manager, [passthrough, no_link]),
meck:expect(
session_manager,
lookup_or_rehydrate,
fun(<<"evicted-session">>, <<"resume-token">>, SocketPid) when is_pid(SocketPid) ->
{ok, SessionPid}
end
),
try
{Frames, _State1} = gateway_handler_identify:handle_resume(
#{
<<"token">> => <<"resume-token">>,
<<"session_id">> => <<"evicted-session">>,
<<"seq">> => 7
},
new_json_state()
),
[{text, Payload}] = Frames,
Message = json:decode(Payload),
?assertEqual(constants:opcode_to_num(invalid_session), maps:get(<<"op">>, Message)),
?assertEqual(false, maps:get(<<"d">>, Message))
after
SessionPid ! stop,
meck:unload(session_manager)
end.
fake_not_resumable_session_loop(Token) ->
receive
{'$gen_call', From, {token_verify, Candidate}} ->
gen_server:reply(From, Candidate =:= Token),
fake_not_resumable_session_loop(Token);
{'$gen_call', From, {resume, _Seq, SocketPid}} when is_pid(SocketPid) ->
gen_server:reply(From, not_resumable),
fake_not_resumable_session_loop(Token);
stop ->
ok
after 30000 ->
ok
end.
assert_resumed_dispatch_without_timings(ExpectedSeq) -> assert_resumed_dispatch_without_timings(ExpectedSeq) ->
receive receive
{dispatch, resumed, ResumedData, ExpectedSeq} -> {dispatch, resumed, ResumedData, ExpectedSeq} ->
@@ -62,23 +62,53 @@ check_full_list_bot_rate_limit_allows_first_test() ->
clear_full_list_bot_rate_limit(UserId, GuildId). clear_full_list_bot_rate_limit(UserId, GuildId).
check_full_list_bot_rate_limit_blocks_second_within_window_test() -> check_full_list_bot_rate_limit_blocks_second_within_window_test() ->
UserId = 111111003, with_rate_limits_enabled(fun() ->
GuildId = 111111004, UserId = 111111003,
clear_full_list_bot_rate_limit(UserId, GuildId), GuildId = 111111004,
?assertEqual( clear_full_list_bot_rate_limit(UserId, GuildId),
ok, ?assertEqual(
guild_request_members_filter:check_full_list_bot_rate_limit(true, true, UserId, GuildId) ok,
), guild_request_members_filter:check_full_list_bot_rate_limit(
case true, true, UserId, GuildId
guild_request_members_filter:check_full_list_bot_rate_limit(true, true, UserId, GuildId) )
of ),
{rate_limited, RetryAfter} -> case
?assert(RetryAfter > 0), guild_request_members_filter:check_full_list_bot_rate_limit(
?assert(RetryAfter =< ?FULL_LIST_BOT_RATE_LIMIT_WINDOW_MS); true, true, UserId, GuildId
Other -> )
?assertEqual({rate_limited, expected}, Other) of
end, {rate_limited, RetryAfter} ->
clear_full_list_bot_rate_limit(UserId, GuildId). ?assert(RetryAfter > 0),
?assert(RetryAfter =< ?FULL_LIST_BOT_RATE_LIMIT_WINDOW_MS);
Other ->
?assertEqual({rate_limited, expected}, Other)
end,
clear_full_list_bot_rate_limit(UserId, GuildId)
end).
check_full_list_bot_rate_limit_disabled_by_env_test() ->
OldValue = os:getenv("FLUXER_DISABLE_RATE_LIMITS"),
os:putenv("FLUXER_DISABLE_RATE_LIMITS", "true"),
try
UserId = 111111011,
GuildId = 111111012,
clear_full_list_bot_rate_limit(UserId, GuildId),
?assertEqual(
ok,
guild_request_members_filter:check_full_list_bot_rate_limit(
true, true, UserId, GuildId
)
),
?assertEqual(
ok,
guild_request_members_filter:check_full_list_bot_rate_limit(
true, true, UserId, GuildId
)
),
clear_full_list_bot_rate_limit(UserId, GuildId)
after
restore_env("FLUXER_DISABLE_RATE_LIMITS", OldValue)
end.
check_full_list_bot_rate_limit_per_guild_isolation_test() -> check_full_list_bot_rate_limit_per_guild_isolation_test() ->
UserId = 111111005, UserId = 111111005,
@@ -180,6 +210,20 @@ check_guild_request_rate_limit_invalid_guild_test() ->
guild_request_members_filter:check_guild_request_rate_limit(invalid_guild_id()) guild_request_members_filter:check_guild_request_rate_limit(invalid_guild_id())
). ).
restore_env(Key, false) ->
os:unsetenv(Key);
restore_env(Key, Value) ->
os:putenv(Key, Value).
with_rate_limits_enabled(Fun) ->
OldValue = os:getenv("FLUXER_DISABLE_RATE_LIMITS"),
os:unsetenv("FLUXER_DISABLE_RATE_LIMITS"),
try
Fun()
after
restore_env("FLUXER_DISABLE_RATE_LIMITS", OldValue)
end.
clear_full_list_bot_rate_limit(UserId, GuildId) -> clear_full_list_bot_rate_limit(UserId, GuildId) ->
ensure_ets_table(?FULL_LIST_BOT_RATE_LIMIT_TABLE), ensure_ets_table(?FULL_LIST_BOT_RATE_LIMIT_TABLE),
ets:delete(?FULL_LIST_BOT_RATE_LIMIT_TABLE, {UserId, GuildId}). ets:delete(?FULL_LIST_BOT_RATE_LIMIT_TABLE, {UserId, GuildId}).
@@ -148,6 +148,40 @@ pre_encoded_voice_state_is_buffered_for_replay_test() ->
?assertEqual(Data, maps:get(data, Entry)), ?assertEqual(Data, maps:get(data, Entry)),
?assertEqual(1, maps:get(seq, Entry)). ?assertEqual(1, maps:get(seq, Entry)).
pre_encoded_message_create_is_buffered_for_replay_test() ->
State0 = base_state(#{}),
Data = {pre_encoded, <<"{\"id\":\"9\",\"content\":\"hi\"}">>},
{noreply, State1} = session_dispatch:handle_dispatch(message_create, Data, State0),
?assertEqual(1, limited_deque:size(maps:get(buffer, State1))),
[Entry] = limited_deque:to_list(maps:get(buffer, State1)),
?assertEqual(message_create, maps:get(event, Entry)),
?assertEqual(Data, maps:get(data, Entry)),
?assertEqual(1, maps:get(seq, Entry)).
pre_encoded_guild_sync_stays_out_of_replay_test() ->
State0 = base_state(#{}),
Data = {pre_encoded, <<"{\"id\":\"123\"}">>},
{noreply, State1} = session_dispatch:handle_dispatch(guild_sync, Data, State0),
?assertEqual(0, limited_deque:size(maps:get(buffer, State1))),
?assertEqual(1, maps:get(seq, State1)).
pre_encoded_eviction_raises_replay_floor_test() ->
State0 = base_state(#{buffer => limited_deque:new(2, 0)}),
Data = {pre_encoded, <<"{\"content\":\"hi\"}">>},
State3 = lists:foldl(
fun(_N, S) ->
{noreply, Next} = session_dispatch:handle_dispatch(message_create, Data, S),
Next
end,
State0,
lists:seq(1, 3)
),
?assertEqual(1, maps:get(replay_floor, State3)),
?assertEqual(
[2, 3],
[maps:get(seq, E) || E <- limited_deque:to_list(maps:get(buffer, State3))]
).
pre_encoded_member_list_stays_out_of_replay_test() -> pre_encoded_member_list_stays_out_of_replay_test() ->
State0 = base_state(#{}), State0 = base_state(#{}),
Data = {pre_encoded, <<"[{\"test\":true}]">>}, Data = {pre_encoded, <<"[{\"test\":true}]">>},
@@ -174,7 +208,11 @@ pre_encoded_multiple_events_seq_test() ->
message_create, {pre_encoded, <<"{\"c\":3}">>}, S2 message_create, {pre_encoded, <<"{\"c\":3}">>}, S2
), ),
?assertEqual(3, maps:get(seq, S3)), ?assertEqual(3, maps:get(seq, S3)),
?assertEqual(0, limited_deque:size(maps:get(buffer, S3))). ?assertEqual(3, limited_deque:size(maps:get(buffer, S3))),
?assertEqual(
[1, 2, 3],
[maps:get(seq, E) || E <- limited_deque:to_list(maps:get(buffer, S3))]
).
pre_encoded_sends_to_socket_test() -> pre_encoded_sends_to_socket_test() ->
State0 = base_state(#{socket_pid => self()}), State0 = base_state(#{socket_pid => self()}),
@@ -235,6 +273,59 @@ pre_encoded_roundtrip_integrity_test() ->
}, },
?assertEqual(OriginalData, json:decode(iolist_to_binary(json:encode(OriginalData)))). ?assertEqual(OriginalData, json:decode(iolist_to_binary(json:encode(OriginalData)))).
guild_counts_update_reaches_nonzero_shard_test() ->
drain_mailbox(),
State0 = base_state(#{socket_pid => self(), shard => {1, 2}}),
Data = #{counts => [], nonce => <<"n">>},
{noreply, State1} = session_dispatch:handle_dispatch(guild_counts_update, Data, State0),
?assertEqual(1, maps:get(seq, State1)),
receive
{dispatch, guild_counts_update, _Payload, 1} -> ok
after 100 -> ?assert(false, dispatch_not_received)
end.
rate_limited_reaches_nonzero_shard_test() ->
drain_mailbox(),
State0 = base_state(#{socket_pid => self(), shard => {1, 2}}),
Data = #{opcode => 8, retry_after => 12.5, meta => #{guild_id => <<"1">>}},
{noreply, State1} = session_dispatch:handle_dispatch(rate_limited, Data, State0),
?assertEqual(1, maps:get(seq, State1)),
receive
{dispatch, rate_limited, _Payload, 1} -> ok
after 100 -> ?assert(false, dispatch_not_received)
end.
channel_member_counts_update_reaches_nonzero_shard_test() ->
drain_mailbox(),
State0 = base_state(#{socket_pid => self(), shard => {1, 2}}),
Data = #{counts => [], nonce => <<"n">>},
{noreply, State1} = session_dispatch:handle_dispatch(
channel_member_counts_update, Data, State0
),
?assertEqual(1, maps:get(seq, State1)),
receive
{dispatch, channel_member_counts_update, _Payload, 1} -> ok
after 100 -> ?assert(false, dispatch_not_received)
end.
guildless_presence_still_skipped_for_nonzero_shard_test() ->
drain_mailbox(),
State0 = base_state(#{socket_pid => self(), shard => {1, 2}}),
{noreply, State1} = session_dispatch:handle_dispatch(
presence_update, #{<<"user_id">> => <<"5">>}, State0
),
?assertEqual(0, maps:get(seq, State1)),
receive
{dispatch, _Event, _Data, _Seq} -> ?assert(false, unexpected_dispatch)
after 100 -> ok
end.
drain_mailbox() ->
receive
_Message -> drain_mailbox()
after 0 -> ok
end.
dispatch_pre_encoded(Event, Json, State) -> dispatch_pre_encoded(Event, Json, State) ->
session_dispatch:handle_dispatch(Event, {pre_encoded, Json}, State). session_dispatch:handle_dispatch(Event, {pre_encoded, Json}, State).
@@ -12,12 +12,15 @@
message_create_flood_keeps_replay_buffer_bounded_test_() -> message_create_flood_keeps_replay_buffer_bounded_test_() ->
{timeout, 30, fun message_create_flood_keeps_replay_buffer_bounded/0}. {timeout, 30, fun message_create_flood_keeps_replay_buffer_bounded/0}.
pre_encoded_flood_does_not_fill_replay_buffer_test_() -> pre_encoded_flood_keeps_replay_buffer_bounded_test_() ->
{timeout, 30, fun pre_encoded_flood_does_not_fill_replay_buffer/0}. {timeout, 30, fun pre_encoded_flood_keeps_replay_buffer_bounded/0}.
guild_members_chunk_flood_does_not_fill_replay_buffer_test_() -> guild_members_chunk_flood_does_not_fill_replay_buffer_test_() ->
{timeout, 30, fun guild_members_chunk_flood_does_not_fill_replay_buffer/0}. {timeout, 30, fun guild_members_chunk_flood_does_not_fill_replay_buffer/0}.
guild_member_list_update_flood_does_not_fill_replay_buffer_test_() ->
{timeout, 30, fun guild_member_list_update_flood_does_not_fill_replay_buffer/0}.
message_create_flood_keeps_replay_buffer_bounded() -> message_create_flood_keeps_replay_buffer_bounded() ->
State0 = base_state(), State0 = base_state(),
State1 = lists:foldl( State1 = lists:foldl(
@@ -40,9 +43,10 @@ message_create_flood_keeps_replay_buffer_bounded() ->
?assertEqual(?MAX_EVENT_BUFFER_SIZE, limited_deque:size(Buffer)), ?assertEqual(?MAX_EVENT_BUFFER_SIZE, limited_deque:size(Buffer)),
?assert(limited_deque:bytes(Buffer) =< ?MAX_TOTAL_BUFFER_BYTES), ?assert(limited_deque:bytes(Buffer) =< ?MAX_TOTAL_BUFFER_BYTES),
?assertEqual(?EVENT_COUNT - ?MAX_EVENT_BUFFER_SIZE + 1, first_buffered_seq(BufferedEvents)), ?assertEqual(?EVENT_COUNT - ?MAX_EVENT_BUFFER_SIZE + 1, first_buffered_seq(BufferedEvents)),
?assertEqual(?EVENT_COUNT, last_buffered_seq(BufferedEvents)). ?assertEqual(?EVENT_COUNT, last_buffered_seq(BufferedEvents)),
?assertEqual(?EVENT_COUNT - ?MAX_EVENT_BUFFER_SIZE, maps:get(replay_floor, State1)).
pre_encoded_flood_does_not_fill_replay_buffer() -> pre_encoded_flood_keeps_replay_buffer_bounded() ->
Encoded = iolist_to_binary(json:encode(#{<<"content">> => <<"preencoded stress">>})), Encoded = iolist_to_binary(json:encode(#{<<"content">> => <<"preencoded stress">>})),
State1 = lists:foldl( State1 = lists:foldl(
fun(_Seq, State) -> fun(_Seq, State) ->
@@ -54,6 +58,27 @@ pre_encoded_flood_does_not_fill_replay_buffer() ->
base_state(), base_state(),
lists:seq(1, ?EVENT_COUNT) lists:seq(1, ?EVENT_COUNT)
), ),
Buffer = maps:get(buffer, State1),
BufferedEvents = limited_deque:to_list(Buffer),
?assertEqual(?EVENT_COUNT, maps:get(seq, State1)),
?assertEqual(?MAX_EVENT_BUFFER_SIZE, limited_deque:size(Buffer)),
?assert(limited_deque:bytes(Buffer) =< ?MAX_TOTAL_BUFFER_BYTES),
?assertEqual(?EVENT_COUNT - ?MAX_EVENT_BUFFER_SIZE + 1, first_buffered_seq(BufferedEvents)),
?assertEqual(?EVENT_COUNT, last_buffered_seq(BufferedEvents)),
?assertEqual(?EVENT_COUNT - ?MAX_EVENT_BUFFER_SIZE, maps:get(replay_floor, State1)).
guild_member_list_update_flood_does_not_fill_replay_buffer() ->
Encoded = iolist_to_binary(json:encode(#{<<"ops">> => []})),
State1 = lists:foldl(
fun(_Seq, State) ->
{noreply, NextState} = session_dispatch:handle_dispatch(
guild_member_list_update, {pre_encoded, Encoded}, State
),
NextState
end,
base_state(),
lists:seq(1, ?EVENT_COUNT)
),
?assertEqual(?EVENT_COUNT, maps:get(seq, State1)), ?assertEqual(?EVENT_COUNT, maps:get(seq, State1)),
?assertEqual(0, limited_deque:size(maps:get(buffer, State1))), ?assertEqual(0, limited_deque:size(maps:get(buffer, State1))),
?assertEqual(0, maps:get(buffer_bytes, State1)). ?assertEqual(0, maps:get(buffer_bytes, State1)).
@@ -59,6 +59,7 @@ serialize_transfer_state_includes_resume_fields_test() ->
relationships => #{}, relationships => #{},
seq => 10, seq => 10,
ack_seq => 8, ack_seq => 8,
replay_floor => 4,
buffer => [#{seq => 9}], buffer => [#{seq => 9}],
collected_guild_states => [], collected_guild_states => [],
collected_sessions => [], collected_sessions => [],
@@ -71,6 +72,7 @@ serialize_transfer_state_includes_resume_fields_test() ->
?assert(sets:is_element(123, maps:get(active_guilds, TransferState))), ?assert(sets:is_element(123, maps:get(active_guilds, TransferState))),
?assertEqual(10, maps:get(seq, TransferState)), ?assertEqual(10, maps:get(seq, TransferState)),
?assertEqual(8, maps:get(ack_seq, TransferState)), ?assertEqual(8, maps:get(ack_seq, TransferState)),
?assertEqual(8, maps:get(replay_floor, TransferState)),
?assertEqual([#{seq => 9}], maps:get(buffer, TransferState)). ?assertEqual([#{seq => 9}], maps:get(buffer, TransferState)).
serialize_transfer_state_strips_socket_pid_test() -> serialize_transfer_state_strips_socket_pid_test() ->
@@ -186,6 +188,29 @@ handle_resume_clamps_skipped_event_hole_test() ->
{reply, {ok, [], 5}, _State1} = session_lifecycle:handle_resume(2, self(), State0), {reply, {ok, [], 5}, _State1} = session_lifecycle:handle_resume(2, self(), State0),
?assertEqual([3, 5], collect_dispatched_seqs(2)). ?assertEqual([3, 5], collect_dispatched_seqs(2)).
handle_resume_below_replay_floor_returns_not_resumable_test() ->
State0 = replay_floor_resume_state(),
{reply, not_resumable, State0} = session_lifecycle:handle_resume(5, self(), State0).
handle_resume_at_replay_floor_still_replays_test() ->
ok = drain_mailbox(),
State0 = replay_floor_resume_state(),
{reply, {ok, [], 10}, State1} = session_lifecycle:handle_resume(6, self(), State0),
?assertEqual([7, 8, 9, 10], collect_dispatched_seqs(4)),
?assertEqual(self(), maps:get(socket_pid, State1)).
replay_floor_resume_state() ->
resume_test_state(#{
seq => 10,
replay_floor => 6,
buffer => [
#{seq => 7, event => message_create, data => #{}},
#{seq => 8, event => message_create, data => #{}},
#{seq => 9, event => message_create, data => #{}},
#{seq => 10, event => message_create, data => #{}}
]
}).
handle_resume_rejects_seq_ahead_of_current_test() -> handle_resume_rejects_seq_ahead_of_current_test() ->
State0 = resume_test_state(#{ State0 = resume_test_state(#{
seq => 5, seq => 5,
@@ -202,6 +227,7 @@ handle_resume_rejects_seq_below_ack_seq_test() ->
{reply, invalid_seq, State0} = session_lifecycle:handle_resume(5, self(), State0). {reply, invalid_seq, State0} = session_lifecycle:handle_resume(5, self(), State0).
handle_resume_accepts_seq_at_ack_seq_test() -> handle_resume_accepts_seq_at_ack_seq_test() ->
ok = drain_mailbox(),
State0 = resume_test_state(#{ State0 = resume_test_state(#{
seq => 10, seq => 10,
ack_seq => 8, ack_seq => 8,
@@ -215,6 +241,7 @@ handle_resume_accepts_seq_at_ack_seq_test() ->
?assertEqual([9, 10], collect_dispatched_seqs(2)). ?assertEqual([9, 10], collect_dispatched_seqs(2)).
handle_resume_accepts_contiguous_replay_buffer_test() -> handle_resume_accepts_contiguous_replay_buffer_test() ->
ok = drain_mailbox(),
State0 = resume_test_state(#{ State0 = resume_test_state(#{
seq => 5, seq => 5,
buffer => [ buffer => [
@@ -458,6 +485,13 @@ collect_dispatched_seq() ->
?assert(false) ?assert(false)
end. end.
drain_mailbox() ->
receive
_ -> drain_mailbox()
after 0 ->
ok
end.
fenced_terminate_releases_the_user_session_count_test() -> fenced_terminate_releases_the_user_session_count_test() ->
ok = session_abuse_protection:ensure_tables(), ok = session_abuse_protection:ensure_tables(),
UserId = 900201, UserId = 900201,