fix(gateway): act on voice states in the voice server (#2815)

This commit is contained in:
Hampus
2026-09-17 03:59:36 +02:00
committed by GitHub
parent 34b6ecfbd2
commit b019f4a91f
6 changed files with 185 additions and 41 deletions
@@ -212,10 +212,8 @@ maybe_disconnect_voice(UserId, ProcessedUsers, State) ->
true ->
{State, ProcessedUsers};
false ->
{reply, _Result, VoiceState} = guild_voice_disconnect:disconnect_voice_user(
#{user_id => UserId, connection_id => null}, State
),
{VoiceState, sets:add_element(UserId, ProcessedUsers)}
ok = guild_voice_lifecycle:cast_disconnect_voice_user(UserId, State),
{State, sets:add_element(UserId, ProcessedUsers)}
end.
-spec schedule_availability_recheck(guild_state()) -> guild_state().
+24 -10
View File
@@ -221,24 +221,38 @@ update_data_for_event_unknown_returns_data_unchanged_test() ->
Data = #{<<"test">> => true},
?assertEqual(Data, update_data_for_event(unknown_event, #{}, Data, #{})).
guild_member_remove_disconnects_voice_test() ->
guild_member_remove_asks_the_voice_server_to_disconnect_test() ->
Self = self(),
TestFun = fun(GId, ChId, UId, ConnId) ->
Self ! {force_disconnect, GId, ChId, UId, ConnId},
{ok, #{success => true}}
end,
State = build_voice_test_state(TestFun),
VoicePid = spawn(fun() ->
receive
Message -> Self ! {voice_server_got, Message}
end
end),
State = build_voice_test_state(VoicePid),
EventData = #{<<"user">> => #{<<"id">> => <<"5">>}},
UpdatedState = update_state(guild_member_remove, EventData, State),
?assertEqual(#{}, maps:get(voice_states, UpdatedState)),
?assertEqual(#{}, maps:get(sessions, UpdatedState, #{})),
receive
{force_disconnect, 42, 20, 5, <<"conn">>} -> ok
{voice_server_got, {'$gen_cast', {disconnect_voice_user, Request}}} ->
?assertEqual(#{user_id => 5, connection_id => null}, Request)
after 200 ->
exit(VoicePid, kill),
?assert(false)
end.
build_voice_test_state(TestFun) ->
guild_member_remove_leaves_the_stale_guild_copy_alone_test() ->
VoicePid = spawn(fun() ->
receive
_ -> ok
end
end),
State = build_voice_test_state(VoicePid),
UpdatedState = update_state(
guild_member_remove, #{<<"user">> => #{<<"id">> => <<"5">>}}, State
),
?assertEqual(maps:get(voice_states, State), maps:get(voice_states, UpdatedState)).
build_voice_test_state(VoicePid) ->
#{
id => 42,
data => #{
@@ -256,7 +270,7 @@ build_voice_test_state(TestFun) ->
}
},
sessions => #{<<"s1">> => #{user_id => 5, pid => self()}},
test_force_disconnect_fun => TestFun
voice_server_pid => VoicePid
}.
make_cache_test_state(GuildId) ->
@@ -132,12 +132,8 @@ cleanup_removed_member_sessions(_UserId, State) ->
-spec maybe_disconnect_removed_member(user_id() | undefined, guild_state()) -> guild_state().
maybe_disconnect_removed_member(UserId, State) when is_integer(UserId), UserId > 0 ->
{reply, _Result, NewState} =
guild_voice_disconnect:disconnect_voice_user(
#{user_id => UserId, connection_id => null},
State
),
NewState;
ok = guild_voice_lifecycle:cast_disconnect_voice_user(UserId, State),
State;
maybe_disconnect_removed_member(_, State) ->
State.
@@ -7,7 +7,9 @@
ensure_voice_server/1,
handle_voice_server_exit/3,
reply_voice_server_pid/1,
clear_stale_cached_voice_states/2
clear_stale_cached_voice_states/2,
authoritative_voice_states/1,
cast_disconnect_voice_user/2
]).
-type guild_state() :: map().
@@ -160,32 +162,58 @@ schedule_stale_cleanup(ConnectionIds) ->
-spec read_authoritative_voice_states(guild_state()) -> {ok, map()} | {error, term()}.
read_authoritative_voice_states(State) ->
case voice_server_pid(State) of
{ok, VoiceServerPid} -> read_from_pid(VoiceServerPid);
error -> {error, no_voice_server}
end.
-spec authoritative_voice_states(guild_state()) -> map().
authoritative_voice_states(State) ->
case read_authoritative_voice_states(State) of
{ok, VoiceStates} -> VoiceStates;
{error, _Reason} -> voice_state_utils:voice_states(State)
end.
-spec cast_disconnect_voice_user(integer(), guild_state()) -> ok.
cast_disconnect_voice_user(UserId, State) when is_integer(UserId), UserId > 0 ->
case voice_server_pid(State) of
{ok, VoiceServerPid} ->
gen_server:cast(
VoiceServerPid,
{disconnect_voice_user, #{user_id => UserId, connection_id => null}}
);
error ->
ok
end;
cast_disconnect_voice_user(_UserId, _State) ->
ok.
-spec voice_server_pid(guild_state()) -> {ok, pid()} | error.
voice_server_pid(State) ->
case maps:get(voice_server_pid, State, undefined) of
VoiceServerPid when is_pid(VoiceServerPid) ->
read_from_pid_or_registry(VoiceServerPid, State);
_ ->
read_from_registry(State)
Pid when is_pid(Pid) -> live_voice_server_pid(Pid, State);
_ -> registered_voice_server_pid(State)
end.
-spec read_from_pid_or_registry(pid(), guild_state()) -> {ok, map()} | {error, term()}.
read_from_pid_or_registry(Pid, State) ->
case process_liveness:is_alive(Pid) of
true -> read_from_pid(Pid);
false -> read_from_registry(State)
-spec live_voice_server_pid(pid(), guild_state()) -> {ok, pid()} | error.
live_voice_server_pid(Pid, State) ->
case Pid =/= self() andalso process_liveness:is_alive(Pid) of
true -> {ok, Pid};
false -> registered_voice_server_pid(State)
end.
-spec read_from_registry(guild_state()) -> {ok, map()} | {error, term()}.
read_from_registry(State) ->
-spec registered_voice_server_pid(guild_state()) -> {ok, pid()} | error.
registered_voice_server_pid(State) ->
case state_guild_id(State) of
{ok, Id} -> read_registered_pid(Id);
error -> {error, no_guild_id}
{ok, Id} -> registered_pid_for_guild(Id);
error -> error
end.
-spec read_registered_pid(integer()) -> {ok, map()} | {error, term()}.
read_registered_pid(Id) ->
-spec registered_pid_for_guild(integer()) -> {ok, pid()} | error.
registered_pid_for_guild(Id) ->
case guild_voice_server:lookup_registered(Id) of
{ok, Pid} -> read_from_pid(Pid);
{error, Reason} -> {error, Reason}
{ok, Pid} when Pid =/= self() -> {ok, Pid};
_ -> error
end.
-spec state_guild_id(guild_state()) -> {ok, integer()} | error.
@@ -203,3 +231,52 @@ read_from_pid(VoiceServerPid) ->
catch
exit:Reason -> {error, Reason}
end.
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
cast_disconnect_voice_user_reaches_the_voice_server_test() ->
Self = self(),
VoicePid = spawn(fun() ->
receive
Message -> Self ! {voice_server_got, Message}
end
end),
ok = cast_disconnect_voice_user(7, #{id => 42, voice_server_pid => VoicePid}),
receive
{voice_server_got, {'$gen_cast', {disconnect_voice_user, Request}}} ->
?assertEqual(#{user_id => 7, connection_id => null}, Request)
after 200 ->
exit(VoicePid, kill),
?assert(false)
end.
cast_disconnect_voice_user_without_a_voice_server_is_a_noop_test() ->
?assertEqual(ok, cast_disconnect_voice_user(7, #{id => 987654329})),
?assertEqual(ok, cast_disconnect_voice_user(undefined, #{id => 987654329})).
authoritative_voice_states_falls_back_to_the_guild_copy_test() ->
Cached = #{<<"conn">> => #{<<"user_id">> => <<"7">>}},
State = #{id => 987654331, voice_states => Cached},
?assertEqual(Cached, authoritative_voice_states(State)).
authoritative_voice_states_prefers_the_voice_server_test() ->
Live = #{<<"live">> => #{<<"user_id">> => <<"8">>}},
VoicePid = spawn(fun() -> voice_states_reply_loop(Live) end),
State = #{id => 987654333, voice_states => #{}, voice_server_pid => VoicePid},
try
?assertEqual(Live, authoritative_voice_states(State))
after
exit(VoicePid, kill)
end.
voice_states_reply_loop(VoiceStates) ->
receive
{'$gen_call', From, {get_voice_states_map}} ->
gen_server:reply(From, VoiceStates),
voice_states_reply_loop(VoiceStates);
_ ->
voice_states_reply_loop(VoiceStates)
end.
-endif.
@@ -29,7 +29,7 @@
-spec sync_user_voice_permissions(user_id(), guild_state()) -> ok.
sync_user_voice_permissions(UserId, State) ->
VoiceStates = voice_state_utils:voice_states(State),
VoiceStates = guild_voice_lifecycle:authoritative_voice_states(State),
case state_guild_id(State) of
undefined ->
ok;
@@ -56,7 +56,7 @@ maybe_sync_user_voice_state(GuildId, UserId, VoiceState, State) ->
-spec sync_all_voice_permissions_for_channel(channel_id(), guild_state()) -> ok.
sync_all_voice_permissions_for_channel(ChannelId, State) ->
VoiceStates = voice_state_utils:voice_states(State),
VoiceStates = guild_voice_lifecycle:authoritative_voice_states(State),
case state_guild_id(State) of
undefined ->
ok;
@@ -236,7 +236,7 @@ enforce_voice_permissions_in_livekit(
-spec sync_users_with_role(integer(), guild_state()) -> ok.
sync_users_with_role(RoleId, State) ->
VoiceStates = voice_state_utils:voice_states(State),
VoiceStates = guild_voice_lifecycle:authoritative_voice_states(State),
case state_guild_id(State) of
undefined ->
ok;
@@ -232,6 +232,8 @@ handle_cast({store_pending_connection, ConnId, Meta}, State) ->
Pending = maps:get(pending_voice_connections, State, #{}),
NewPending = bounded_put(ConnId, Meta, Pending, ?MAX_PENDING_CONNECTIONS),
{noreply, State#{pending_voice_connections => NewPending}};
handle_cast({disconnect_voice_user, Request}, State) when is_map(Request) ->
{noreply, delegate_voice_cast(fun guild_voice:disconnect_voice_user/2, Request, State)};
handle_cast({cleanup_virtual_access_for_user, UserId}, State) when is_integer(UserId) ->
GS = guild_voice_server_state:build_guild_state(State),
NewGS = guild_voice_disconnect:cleanup_virtual_channel_access_for_user(UserId, GS),
@@ -295,6 +297,13 @@ safe_lookup(Fun) ->
safe_lookup_match({ok, Pid}) when is_pid(Pid) -> {ok, Pid};
safe_lookup_match(_) -> {error, not_found}.
-spec delegate_voice_cast(fun((map(), map()) -> {reply, term(), map()}), map(), server_state()) ->
server_state().
delegate_voice_cast(Fun, Request, State) ->
GS = guild_voice_server_state:build_guild_state(State),
{reply, _Reply, NewGS} = Fun(Request, GS),
guild_voice_server_state:apply_guild_state(NewGS, State).
-spec delegate_voice_call(fun((map(), map()) -> {reply, term(), map()}), map(), server_state()) ->
{reply, term(), server_state()}.
delegate_voice_call(Fun, Request, State) ->
@@ -433,6 +442,56 @@ enforce_map_cap(Map, MaxSize) ->
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
disconnect_voice_user_cast_removes_the_voice_state_test() ->
Self = self(),
TestFun = fun(GId, ChId, UId, ConnId) ->
Self ! {force_disconnect, GId, ChId, UId, ConnId},
{ok, #{success => true}}
end,
GuildPid = spawn(fun() -> guild_state_reply_loop(TestFun) end),
State = #{
guild_id => 42,
guild_pid => GuildPid,
voice_states => #{
<<"conn">> => #{
<<"user_id">> => <<"5">>,
<<"guild_id">> => <<"42">>,
<<"channel_id">> => <<"20">>,
<<"connection_id">> => <<"conn">>
}
},
pending_voice_connections => #{},
recently_disconnected_voice_states => #{},
e2ee_room_keys => #{}
},
try
{noreply, NewState} = handle_cast(
{disconnect_voice_user, #{user_id => 5, connection_id => null}}, State
),
?assertEqual(#{}, maps:get(voice_states, NewState)),
receive
{force_disconnect, 42, 20, 5, <<"conn">>} -> ok
after 200 -> ?assert(false)
end
after
exit(GuildPid, kill)
end.
guild_state_reply_loop(TestFun) ->
GuildState = #{
id => 42,
data => #{<<"guild">> => #{<<"owner_id">> => <<"999">>}},
sessions => #{},
test_force_disconnect_fun => TestFun
},
receive
{'$gen_call', From, {get_voice_guild_state}} ->
gen_server:reply(From, GuildState),
guild_state_reply_loop(TestFun);
_ ->
guild_state_reply_loop(TestFun)
end.
apply_guild_state_preserves_e2ee_test() ->
State = #{
guild_id => 1,