fix(gateway): stop dead sessions leaving voice states behind (#2818)

This commit is contained in:
Hampus
2026-09-17 14:55:18 +02:00
committed by GitHub
parent 4cecbf1f43
commit ed9528834d
5 changed files with 315 additions and 6 deletions
@@ -136,8 +136,19 @@ update_circuit_state(GuildPid, Result, PrevState) ->
end.
-spec is_success_result(guild_client:voice_state_update_result()) -> boolean().
is_success_result({ok, _}) -> true;
is_success_result(_) -> false.
is_success_result({ok, _}) ->
true;
is_success_result({error, Category, _}) when
Category =:= validation_error;
Category =:= not_found;
Category =:= permission_denied;
Category =:= voice_error;
Category =:= rate_limited;
Category =:= auth_failed
->
true;
is_success_result(_) ->
false.
-spec reset_failures(pid()) -> ok.
reset_failures(GuildPid) ->
@@ -329,6 +340,13 @@ record_failure_recreates_missing_table_test() ->
is_success_result_test() ->
?assertEqual(true, is_success_result({ok, #{success => true}})),
?assertEqual(false, is_success_result({error, timeout})),
?assertEqual(false, is_success_result({error, noproc})).
?assertEqual(false, is_success_result({error, noproc})),
?assertEqual(
true, is_success_result({error, validation_error, voice_missing_connection_id})
),
?assertEqual(true, is_success_result({error, not_found, voice_member_not_found})),
?assertEqual(true, is_success_result({error, permission_denied, voice_permission_denied})),
?assertEqual(false, is_success_result({error, unknown, internal_error})),
?assertEqual(false, is_success_result({error, timeout, timeout})).
-endif.
@@ -81,6 +81,16 @@ sweep_expired_pending_joins(State) ->
voice_reply().
handle_with_member(Context, UserId, State) ->
VoiceStates = voice_state_utils:voice_states(State),
case is_targeted_disconnect(Context) of
true ->
handle_disconnect(Context, VoiceStates, State);
false ->
handle_member_lookup(Context, UserId, VoiceStates, State)
end.
-spec handle_member_lookup(map(), integer(), voice_state_map(), guild_state()) ->
voice_reply().
handle_member_lookup(Context, UserId, VoiceStates, State) ->
case guild_voice_member:find_member_by_user_id(UserId, State) of
undefined ->
{reply, gateway_errors:error(voice_member_not_found), State};
@@ -88,6 +98,13 @@ handle_with_member(Context, UserId, State) ->
handle_member_voice(Context, Member, VoiceStates, State)
end.
-spec is_targeted_disconnect(map()) -> boolean().
is_targeted_disconnect(Context) ->
maps:get(channel_id, Context, undefined) =:= null andalso
(maps:get(connection_id, Context, undefined) =/= undefined orelse
voice_state_utils:normalize_session_id(maps:get(session_id, Context, undefined)) =/=
undefined).
-spec handle_member_voice(map(), map(), voice_state_map(), guild_state()) ->
voice_reply().
handle_member_voice(Context, Member, VoiceStates, State) ->
@@ -26,8 +26,14 @@
-spec handle_voice_disconnect(
binary() | undefined, term(), integer(), voice_state_map() | term(), guild_state()
) -> voice_reply().
handle_voice_disconnect(undefined, _SessionId, _UserId, _VoiceStates, State) ->
{reply, gateway_errors:error(voice_missing_connection_id), State};
handle_voice_disconnect(undefined, SessionId, UserId, VoiceStates0, State) ->
case normalize_session_id(SessionId) of
undefined ->
{reply, gateway_errors:error(voice_missing_connection_id), State};
RequestSessionId ->
VoiceStates = voice_state_utils:ensure_voice_states(VoiceStates0),
disconnect_all_user_connections(UserId, RequestSessionId, VoiceStates, State)
end;
handle_voice_disconnect(ConnectionId, _SessionId, UserId, VoiceStates0, State) ->
VoiceStates = voice_state_utils:ensure_voice_states(VoiceStates0),
case maps:get(ConnectionId, VoiceStates, undefined) of
@@ -29,6 +29,8 @@
-define(PENDING_TTL_MS, 300000).
-define(RECENT_DISCONNECT_TTL_MS, 120000).
-define(E2EE_KEY_TTL_MS, 300000).
-define(SEEDED_SESSION_CHECK_DELAY_MS, 60000).
-define(SESSION_LOOKUP_TIMEOUT_MS, 5000).
-type voice_state_map() :: #{binary() => map()}.
-type server_state() :: #{
@@ -144,6 +146,7 @@ init(#{guild_id := GuildId, guild_pid := GuildPid} = Args) ->
recently_disconnected_voice_states => #{},
e2ee_room_keys => #{}
},
ok = schedule_seeded_session_check(first, maps:keys(InitialVoiceStates)),
{ok, State}.
-spec handle_call(term(), gen_server:from(), server_state()) -> {reply, term(), server_state()}.
@@ -255,6 +258,14 @@ handle_info(sweep_pending_joins, State) ->
NewGuildState = guild_voice_connection:sweep_expired_pending_joins(GuildState),
{noreply, guild_voice_server_state:apply_guild_state(NewGuildState, State2)}
end;
handle_info({check_seeded_sessions, Round, ConnectionIds}, State) ->
ok = spawn_seeded_session_check(Round, seeded_sessions(ConnectionIds, State), State),
{noreply, State};
handle_info({seeded_sessions_gone, first, Gone}, State) ->
ok = schedule_seeded_session_check(second, maps:keys(Gone)),
{noreply, State};
handle_info({seeded_sessions_gone, second, Gone}, State) ->
{noreply, remove_gone_seeded_connections(Gone, State)};
handle_info({'EXIT', Pid, Reason}, #{guild_pid := GuildPid} = State) when Pid =:= GuildPid ->
logger:info(
"Voice server shutting down because guild process exited",
@@ -297,6 +308,104 @@ safe_lookup(Fun) ->
safe_lookup_match({ok, Pid}) when is_pid(Pid) -> {ok, Pid};
safe_lookup_match(_) -> {error, not_found}.
-spec schedule_seeded_session_check(first | second, [binary()]) -> ok.
schedule_seeded_session_check(_Round, []) ->
ok;
schedule_seeded_session_check(Round, ConnectionIds) ->
_ = erlang:send_after(
?SEEDED_SESSION_CHECK_DELAY_MS, self(), {check_seeded_sessions, Round, ConnectionIds}
),
ok.
-spec seeded_sessions([binary()], server_state()) -> #{binary() => {integer(), binary()}}.
seeded_sessions(ConnectionIds, State) ->
VoiceStates = maps:get(voice_states, State, #{}),
maps:from_list([
{ConnectionId, {UserId, SessionId}}
|| ConnectionId <- ConnectionIds,
VoiceState <- [maps:get(ConnectionId, VoiceStates, undefined)],
is_map(VoiceState),
UserId <- [voice_state_utils:voice_state_user_id(VoiceState)],
is_integer(UserId),
SessionId <- [
voice_state_utils:normalize_session_id(
maps:get(<<"session_id">>, VoiceState, undefined)
)
],
is_binary(SessionId)
]).
-spec spawn_seeded_session_check(
first | second, #{binary() => {integer(), binary()}}, server_state()
) -> ok.
spawn_seeded_session_check(_Round, Seeded, _State) when map_size(Seeded) =:= 0 ->
ok;
spawn_seeded_session_check(Round, Seeded, State) ->
Self = self(),
IsGone = session_gone_fun(State),
_ = spawn(fun() ->
SessionIds = lists:usort([SessionId || {_UserId, SessionId} <- maps:values(Seeded)]),
GoneSessions = [SessionId || SessionId <- SessionIds, IsGone(SessionId)],
Gone = maps:filter(
fun(_ConnectionId, {_UserId, SessionId}) ->
lists:member(SessionId, GoneSessions)
end,
Seeded
),
Self ! {seeded_sessions_gone, Round, Gone}
end),
ok.
-spec session_gone_fun(server_state()) -> fun((binary()) -> boolean()).
session_gone_fun(State) ->
case maps:get(test_session_gone_fun, State, undefined) of
Fun when is_function(Fun, 1) -> Fun;
_ -> fun session_gone/1
end.
-spec session_gone(binary()) -> boolean().
session_gone(SessionId) ->
try
session_manager_routing:call_owner_manager(
SessionId, {lookup, SessionId}, ?SESSION_LOOKUP_TIMEOUT_MS
)
of
{error, not_found} -> true;
_ -> false
catch
_:_ -> false
end.
-spec remove_gone_seeded_connections(#{binary() => {integer(), binary()}}, server_state()) ->
server_state().
remove_gone_seeded_connections(Gone, State) ->
maps:fold(
fun(ConnectionId, {UserId, SessionId}, AccState) ->
remove_gone_seeded_connection(ConnectionId, UserId, SessionId, AccState)
end,
State,
Gone
).
-spec remove_gone_seeded_connection(binary(), integer(), binary(), server_state()) ->
server_state().
remove_gone_seeded_connection(ConnectionId, UserId, SessionId, State) ->
case seeded_sessions([ConnectionId], State) of
#{ConnectionId := {UserId, SessionId}} ->
logger:warning(
"guild_voice_seeded_state_removed: guild_id=~p connection_id=~p"
" user_id=~p session_id=~p",
[maps:get(guild_id, State, undefined), ConnectionId, UserId, SessionId]
),
delegate_voice_cast(
fun guild_voice:disconnect_voice_user/2,
#{user_id => UserId, connection_id => ConnectionId},
State
);
_ ->
State
end.
-spec delegate_voice_cast(fun((map(), map()) -> {reply, term(), map()}), map(), server_state()) ->
server_state().
delegate_voice_cast(Fun, Request, State) ->
@@ -442,6 +551,88 @@ enforce_map_cap(Map, MaxSize) ->
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
seeded_sessions_only_tracks_states_with_a_session_test() ->
State = #{
voice_states => #{
<<"a">> => #{<<"user_id">> => <<"5">>, <<"session_id">> => <<"sess-a">>},
<<"b">> => #{<<"user_id">> => <<"6">>},
<<"c">> => #{<<"user_id">> => <<"7">>, <<"session_id">> => null}
}
},
?assertEqual(
#{<<"a">> => {5, <<"sess-a">>}},
seeded_sessions([<<"a">>, <<"b">>, <<"c">>, <<"missing">>], State)
).
seeded_session_check_reports_only_gone_sessions_test() ->
State = #{
voice_states => #{
<<"a">> => #{<<"user_id">> => <<"5">>, <<"session_id">> => <<"dead">>},
<<"b">> => #{<<"user_id">> => <<"6">>, <<"session_id">> => <<"alive">>}
},
test_session_gone_fun => fun(SessionId) -> SessionId =:= <<"dead">> end
},
ok = spawn_seeded_session_check(first, seeded_sessions([<<"a">>, <<"b">>], State), State),
receive
{seeded_sessions_gone, first, Gone} ->
?assertEqual(#{<<"a">> => {5, <<"dead">>}}, Gone)
after 1000 -> ?assert(false)
end.
gone_seeded_connection_is_removed_test() ->
TestFun = fun(_, _, _, _) -> {ok, #{success => true}} end,
GuildPid = spawn(fun() -> guild_state_reply_loop(TestFun) end),
State = seeded_test_state(GuildPid),
try
{noreply, NewState} = handle_info(
{seeded_sessions_gone, second, #{<<"ghost">> => {5, <<"dead">>}}}, State
),
?assertEqual([<<"live">>], maps:keys(maps:get(voice_states, NewState)))
after
exit(GuildPid, kill)
end.
seeded_connection_that_changed_session_is_kept_test() ->
TestFun = fun(_, _, _, _) -> {ok, #{success => true}} end,
GuildPid = spawn(fun() -> guild_state_reply_loop(TestFun) end),
State = seeded_test_state(GuildPid),
try
{noreply, NewState} = handle_info(
{seeded_sessions_gone, second, #{<<"ghost">> => {5, <<"an-older-session">>}}}, State
),
?assertEqual(
lists:sort([<<"ghost">>, <<"live">>]),
lists:sort(maps:keys(maps:get(voice_states, NewState)))
)
after
exit(GuildPid, kill)
end.
seeded_test_state(GuildPid) ->
#{
guild_id => 42,
guild_pid => GuildPid,
voice_states => #{
<<"ghost">> => #{
<<"user_id">> => <<"5">>,
<<"guild_id">> => <<"42">>,
<<"channel_id">> => <<"20">>,
<<"connection_id">> => <<"ghost">>,
<<"session_id">> => <<"dead">>
},
<<"live">> => #{
<<"user_id">> => <<"6">>,
<<"guild_id">> => <<"42">>,
<<"channel_id">> => <<"20">>,
<<"connection_id">> => <<"live">>,
<<"session_id">> => <<"alive">>
}
},
pending_voice_connections => #{},
recently_disconnected_voice_states => #{},
e2ee_room_keys => #{}
}.
disconnect_voice_user_cast_removes_the_voice_state_test() ->
Self = self(),
TestFun = fun(GId, ChId, UId, ConnId) ->
@@ -71,6 +71,53 @@ voice_state_update_guild_leave_omitted_connection_id_test() ->
),
?assertEqual(connected_voice_states(<<"10">>), maps:get(voice_states, NewState)).
voice_state_update_session_teardown_removes_that_sessions_states_test() ->
State = session_scoped_state(),
{reply, #{success := true}, NewState} =
guild_voice_connection:voice_state_update(
#{
user_id => 10,
channel_id => null,
connection_id => null,
session_id => <<"sess-a">>
},
State
),
Remaining = maps:get(voice_states, NewState),
?assertNot(maps:is_key(<<"conn-a">>, Remaining)),
?assert(maps:is_key(<<"conn-b">>, Remaining)),
?assert(maps:is_key(<<"conn-c">>, Remaining)).
voice_state_update_session_teardown_works_for_former_members_test() ->
State = maps:put(
data,
#{<<"channels">> => [base_test_channel(100)], <<"members">> => []},
session_scoped_state()
),
{reply, #{success := true}, NewState} =
guild_voice_connection:voice_state_update(
#{
user_id => 10,
channel_id => null,
connection_id => null,
session_id => <<"sess-a">>
},
State
),
?assertNot(maps:is_key(<<"conn-a">>, maps:get(voice_states, NewState))).
voice_state_update_leave_by_connection_works_for_former_members_test() ->
State = maps:put(
data,
#{<<"channels">> => [base_test_channel(100)], <<"members">> => []},
session_scoped_state()
),
{reply, #{success := true}, NewState} =
guild_voice_connection:voice_state_update(
#{user_id => 10, channel_id => null, connection_id => <<"conn-b">>}, State
),
?assertNot(maps:is_key(<<"conn-b">>, maps:get(voice_states, NewState))).
voice_state_update_guild_leave_blank_connection_id_test() ->
State = connected_state(<<"10">>),
assert_voice_state_update_error(
@@ -337,11 +384,41 @@ camera_state(SharerCount) ->
).
connected_voice_states(UserId) ->
#{<<"conn-1">> => #{<<"channel_id">> => <<"100">>, <<"user_id">> => UserId}}.
#{
<<"conn-1">> => #{
<<"guild_id">> => <<"999">>, <<"channel_id">> => <<"100">>, <<"user_id">> => UserId
}
}.
connected_state(UserId) ->
maps:put(voice_states, connected_voice_states(UserId), base_test_state()).
session_scoped_state() ->
VoiceStates = #{
<<"conn-a">> => #{
<<"guild_id">> => <<"999">>,
<<"channel_id">> => <<"100">>,
<<"user_id">> => <<"10">>,
<<"session_id">> => <<"sess-a">>
},
<<"conn-b">> => #{
<<"guild_id">> => <<"999">>,
<<"channel_id">> => <<"100">>,
<<"user_id">> => <<"10">>,
<<"session_id">> => <<"sess-b">>
},
<<"conn-c">> => #{
<<"guild_id">> => <<"999">>,
<<"channel_id">> => <<"100">>,
<<"user_id">> => <<"11">>,
<<"session_id">> => <<"sess-a">>
}
},
(base_test_state())#{
voice_states => VoiceStates,
test_force_disconnect_fun => fun(_, _, _, _) -> {ok, #{success => true}} end
}.
connected_voice_state(ConnectionId, UserId, ChannelId, ViewerKeys) ->
#{
<<"connection_id">> => ConnectionId,