diff --git a/fluxer_gateway/src/guild/guild.erl b/fluxer_gateway/src/guild/guild.erl index b44133945..3bf87e5a0 100644 --- a/fluxer_gateway/src/guild/guild.erl +++ b/fluxer_gateway/src/guild/guild.erl @@ -172,10 +172,11 @@ route_cast(Tag, Msg, State) -> case cast_handler(Tag) of voice -> guild_voice_handler:handle_cast(Msg, State); subscription -> guild_subscription_handler:handle_cast(Msg, State); + dm_partners -> guild_dm_partners:handle_cast(Msg, State); undefined -> {noreply, State} end. --spec cast_handler(atom()) -> voice | subscription | undefined. +-spec cast_handler(atom()) -> voice | subscription | dm_partners | undefined. cast_handler(relay_voice_state_update) -> voice; cast_handler(relay_voice_server_update) -> voice; cast_handler(store_pending_connection) -> voice; @@ -183,6 +184,7 @@ cast_handler(add_virtual_channel_access) -> voice; cast_handler(remove_virtual_channel_access) -> voice; cast_handler(cleanup_virtual_access_for_user) -> voice; cast_handler(update_member_subscriptions) -> subscription; +cast_handler(update_dm_partners) -> dm_partners; cast_handler(_) -> undefined. -spec handle_info(term(), guild_state()) -> info_reply(). @@ -545,7 +547,7 @@ dispatch_event(Event, EventData, State) -> Event, NewState ), ok = maybe_refresh_permission_cache(Event, ParsedEventData, State, StateAfterPrune), - StateAfterPrune. + guild_dm_partners:maybe_reevaluate(Event, ParsedEventData, State, StateAfterPrune). -spec parse_event_data(term()) -> map(). parse_event_data(D) when is_binary(D) -> require_map(json:decode(D)); diff --git a/fluxer_gateway/src/guild/guild_dm_partners.erl b/fluxer_gateway/src/guild/guild_dm_partners.erl new file mode 100644 index 000000000..ef8e9b349 --- /dev/null +++ b/fluxer_gateway/src/guild/guild_dm_partners.erl @@ -0,0 +1,259 @@ +%% SPDX-License-Identifier: AGPL-3.0-or-later + +-module(guild_dm_partners). +-typing([eqwalizer]). + +-export([handle_cast/2, maybe_reevaluate/4]). + +-type guild_state() :: map(). +-type session_id() :: binary(). +-type user_id() :: integer(). +-type registration() :: #{ + user_id := user_id(), + pid := pid(), + partners := #{user_id() => true}, + eligible := #{user_id() => true} +}. +-type registrations() :: #{session_id() => registration()}. +-type scope() :: none | all | {user, user_id()}. + +-export_type([guild_state/0]). + +-define(MAX_PARTNERS, 1000). +-define(PRUNE_SLACK, 64). + +-spec handle_cast(term(), guild_state()) -> {noreply, guild_state()}. +handle_cast({update_dm_partners, SessionId, PartnerIds}, State) when + is_binary(SessionId), is_list(PartnerIds) +-> + {noreply, update(SessionId, PartnerIds, State)}; +handle_cast(_Msg, State) -> + {noreply, State}. + +-spec maybe_reevaluate(term(), term(), guild_state(), guild_state()) -> guild_state(). +maybe_reevaluate(Event, Data, PreviousState, State) -> + Registrations = registrations(State), + case map_size(Registrations) of + 0 -> + State; + _ -> + reevaluate( + effective_scope(scope(Event, Data), PreviousState, State), Registrations, State + ) + end. + +-spec effective_scope(scope(), guild_state(), guild_state()) -> scope(). +effective_scope(all, PreviousState, State) -> + case visibility_inputs(PreviousState) =:= visibility_inputs(State) of + true -> none; + false -> all + end; +effective_scope(Scope, _PreviousState, _State) -> + Scope. + +-spec visibility_inputs(guild_state()) -> term(). +visibility_inputs(State) -> + Data = map_utils:ensure_map(maps:get(data, State, #{})), + Guild = map_utils:ensure_map(maps:get(<<"guild">>, Data, #{})), + { + maps:get(<<"owner_id">>, Guild, undefined), + role_inputs(guild_data_index:role_index(Data)), + channel_inputs(guild_data_index:channel_index(Data)), + maps:get(virtual_channel_access, State, #{}) + }. + +-spec role_inputs(term()) -> term(). +role_inputs(Roles) when is_map(Roles) -> + lists:sort([ + {Id, maps:get(<<"permissions">>, Role, undefined)} + || {Id, Role} <- maps:to_list(Roles), is_map(Role) + ]); +role_inputs(Roles) -> + Roles. + +-spec channel_inputs(term()) -> term(). +channel_inputs(Channels) when is_map(Channels) -> + lists:sort([ + { + Id, + maps:get(<<"type">>, Channel, undefined), + maps:get(<<"parent_id">>, Channel, undefined), + maps:get(<<"permission_overwrites">>, Channel, []) + } + || {Id, Channel} <- maps:to_list(Channels), is_map(Channel) + ]); +channel_inputs(Channels) -> + Channels. + +-spec update(session_id(), [term()], guild_state()) -> guild_state(). +update(SessionId, PartnerIds, State) -> + Registrations = registrations(State), + Next = + case session_owner(SessionId, State) of + {ok, UserId, Pid} -> + update_owned(SessionId, UserId, Pid, PartnerIds, Registrations, State); + error -> + maps:remove(SessionId, Registrations) + end, + State#{dm_partners => maybe_prune(Next, State)}. + +-spec update_owned(session_id(), user_id(), pid(), [term()], registrations(), guild_state()) -> + registrations(). +update_owned(SessionId, UserId, Pid, PartnerIds, Registrations, State) -> + case presence_targets:dm_partner_presence_enabled(UserId) of + false -> + maps:remove(SessionId, Registrations); + true -> + Entry = #{ + user_id => UserId, + pid => Pid, + partners => partner_set(PartnerIds, UserId), + eligible => previous_eligible( + maps:get(SessionId, Registrations, undefined), Pid + ) + }, + put_all(evaluate(#{SessionId => Entry}, State), Registrations) + end. + +-spec reevaluate(scope(), registrations(), guild_state()) -> guild_state(). +reevaluate(none, _Registrations, State) -> + State; +reevaluate(all, Registrations, State) -> + State#{dm_partners => evaluate(live_registrations(Registrations, State), State)}; +reevaluate({user, UserId}, Registrations, State) -> + Affected = maps:fold( + fun(SessionId, Entry, Acc) -> + case involves_user(UserId, Entry) of + true -> Acc#{SessionId => Entry}; + false -> Acc + end + end, + #{}, + Registrations + ), + case map_size(Affected) of + 0 -> + State; + _ -> + Live = live_registrations(Affected, State), + Kept = maps:without( + maps:keys(maps:without(maps:keys(Live), Affected)), Registrations + ), + State#{dm_partners => maybe_prune(put_all(evaluate(Live, State), Kept), State)} + end. + +-spec put_all(registrations(), registrations()) -> registrations(). +put_all(Entries, Registrations) -> + maps:fold( + fun(SessionId, Entry, Acc) -> Acc#{SessionId => Entry} end, Registrations, Entries + ). + +-spec maybe_prune(registrations(), guild_state()) -> registrations(). +maybe_prune(Registrations, State) -> + Sessions = maps:get(sessions, State, #{}), + case map_size(Registrations) > map_size(Sessions) + ?PRUNE_SLACK of + true -> + maps:filter( + fun(SessionId, _Entry) -> maps:is_key(SessionId, Sessions) end, Registrations + ); + false -> + Registrations + end. + +-spec evaluate(registrations(), guild_state()) -> registrations(). +evaluate(Entries, _State) when map_size(Entries) =:= 0 -> + Entries; +evaluate(Entries, State) -> + Requests = [ + {SessionId, UserId, maps:keys(Partners)} + || {SessionId, #{user_id := UserId, partners := Partners}} <- maps:to_list(Entries) + ], + Results = guild_subscription_mutual_channels:filter_session_member_ids(Requests, State), + GuildId = maps:get(id, State), + maps:map( + fun(SessionId, Entry) -> + apply_result(GuildId, maps:get(SessionId, Results, []), Entry) + end, + Entries + ). + +-spec apply_result(integer(), [user_id()], registration()) -> registration(). +apply_result(GuildId, EligibleIds, #{pid := Pid, eligible := Previous} = Entry) -> + Eligible = presence_targets:map_from_ids(EligibleIds), + ok = notify_if_changed(Pid, GuildId, Previous, Eligible), + Entry#{eligible := Eligible}. + +-spec notify_if_changed(pid(), integer(), #{user_id() => true}, #{user_id() => true}) -> ok. +notify_if_changed(_Pid, _GuildId, Same, Same) -> + ok; +notify_if_changed(Pid, GuildId, _Previous, Eligible) -> + Pid ! {dm_partner_mutual, GuildId, maps:keys(Eligible)}, + ok. + +-spec live_registrations(registrations(), guild_state()) -> registrations(). +live_registrations(Registrations, State) -> + Sessions = maps:get(sessions, State, #{}), + maps:filter( + fun(SessionId, #{user_id := UserId}) -> + maps:is_key(SessionId, Sessions) andalso + presence_targets:dm_partner_presence_enabled(UserId) + end, + Registrations + ). + +-spec involves_user(user_id(), registration()) -> boolean(). +involves_user(UserId, #{user_id := UserId}) -> + true; +involves_user(UserId, #{partners := Partners}) -> + maps:is_key(UserId, Partners). + +-spec scope(term(), term()) -> scope(). +scope(guild_member_add, Data) -> member_scope(Data); +scope(guild_member_remove, Data) -> member_scope(Data); +scope(guild_member_update, Data) -> member_scope(Data); +scope(guild_update, _Data) -> all; +scope(guild_role_update, _Data) -> all; +scope(guild_role_update_bulk, _Data) -> all; +scope(guild_role_delete, _Data) -> all; +scope(channel_create, _Data) -> all; +scope(channel_update, _Data) -> all; +scope(channel_update_bulk, _Data) -> all; +scope(channel_delete, _Data) -> all; +scope(_Event, _Data) -> none. + +-spec member_scope(term()) -> scope(). +member_scope(#{<<"user">> := #{<<"id">> := RawUserId}}) -> + case snowflake_id:parse_maybe(RawUserId) of + UserId when is_integer(UserId) -> {user, UserId}; + _ -> all + end; +member_scope(_Data) -> + all. + +-spec session_owner(session_id(), guild_state()) -> {ok, user_id(), pid()} | error. +session_owner(SessionId, State) -> + case maps:get(SessionId, maps:get(sessions, State, #{}), undefined) of + #{user_id := UserId, pid := Pid} when is_integer(UserId), is_pid(Pid) -> + {ok, UserId, Pid}; + _ -> + error + end. + +-spec previous_eligible(registration() | undefined, pid()) -> #{user_id() => true}. +previous_eligible(#{pid := Pid, eligible := Eligible}, Pid) -> + Eligible; +previous_eligible(_Previous, _Pid) -> + #{}. + +-spec partner_set([term()], user_id()) -> #{user_id() => true}. +partner_set(PartnerIds, UserId) -> + presence_targets:map_from_ids( + lists:sublist([Id || Id <- PartnerIds, is_integer(Id), Id =/= UserId], ?MAX_PARTNERS) + ). + +-spec registrations(guild_state()) -> registrations(). +registrations(State) -> + case maps:get(dm_partners, State, #{}) of + Registrations when is_map(Registrations) -> Registrations; + _ -> #{} + end. diff --git a/fluxer_gateway/src/guild/guild_subscription_mutual_channels.erl b/fluxer_gateway/src/guild/guild_subscription_mutual_channels.erl index 404bdf93f..ef20577fb 100644 --- a/fluxer_gateway/src/guild/guild_subscription_mutual_channels.erl +++ b/fluxer_gateway/src/guild/guild_subscription_mutual_channels.erl @@ -3,7 +3,7 @@ -module(guild_subscription_mutual_channels). -typing([eqwalizer]). --export([filter_member_ids/3]). +-export([filter_member_ids/3, filter_session_member_ids/2]). -ifdef(TEST). -include_lib("eunit/include/eunit.hrl"). @@ -28,6 +28,82 @@ filter_member_ids(SessionUserId, MemberIds, State) -> ), lists:reverse(Kept). +-spec filter_session_member_ids([{term(), user_id(), [term()]}], guild_state()) -> + #{term() => [user_id()]}. +filter_session_member_ids([], _State) -> + #{}; +filter_session_member_ids(Requests, State) -> + Exceptions = exceptions(State), + Sessions = maps:get(sessions, State, #{}), + {Results, _Cache} = lists:foldl( + fun({SessionId, SessionUserId, MemberIds}, {Acc, Cache}) -> + SessionMap = request_session_map(SessionId, SessionUserId, Sessions, State), + {Kept, Cache1} = keep_session_members( + MemberIds, SessionUserId, SessionMap, Exceptions, State, Cache + ), + {Acc#{SessionId => Kept}, Cache1} + end, + {#{}, #{}}, + Requests + ), + Results. + +-spec request_session_map(term(), user_id(), map(), guild_state()) -> map(). +request_session_map(SessionId, SessionUserId, Sessions, State) -> + case maps:get(SessionId, Sessions, undefined) of + #{user_id := SessionUserId, viewable_channels := Map} = SessionData when is_map(Map) -> + case maps:get(pending_connect, SessionData, false) of + true -> session_channel_map(SessionUserId, State); + _ -> Map + end; + _ -> + session_channel_map(SessionUserId, State) + end. + +-spec keep_session_members( + [term()], user_id(), map(), sets:set(user_id()), guild_state(), #{term() => [integer()]} +) -> {[user_id()], #{term() => [integer()]}}. +keep_session_members(MemberIds, SessionUserId, SessionMap, Exceptions, State, Cache) -> + {Kept, _Seen, Cache1} = lists:foldl( + fun + (MemberId, Acc) when MemberId =:= SessionUserId; not is_integer(MemberId) -> + Acc; + (MemberId, {KeptAcc, Seen, CacheAcc}) -> + {Key, CacheAcc1} = member_channel_key(MemberId, Exceptions, State, CacheAcc), + case maps:find(Key, Seen) of + {ok, true} -> + {[MemberId | KeptAcc], Seen, CacheAcc1}; + {ok, false} -> + {KeptAcc, Seen, CacheAcc1}; + error -> + Shared = has_shared_channel(maps:get(Key, CacheAcc1), SessionMap), + {kept_if(Shared, MemberId, KeptAcc), Seen#{Key => Shared}, CacheAcc1} + end + end, + {[], #{}, Cache}, + MemberIds + ), + {lists:reverse(Kept), Cache1}. + +-spec kept_if(boolean(), user_id(), [user_id()]) -> [user_id()]. +kept_if(true, MemberId, Kept) -> [MemberId | Kept]; +kept_if(false, _MemberId, Kept) -> Kept. + +-spec member_channel_key(user_id(), sets:set(user_id()), guild_state(), #{term() => [integer()]}) -> + {term(), #{term() => [integer()]}}. +member_channel_key(MemberId, Exceptions, State, Cache) -> + Key = + case memo_key(MemberId, Exceptions, State) of + bypass -> {user, MemberId}; + {ok, RawRoles} -> {roles, RawRoles} + end, + case maps:is_key(Key, Cache) of + true -> + {Key, Cache}; + false -> + {Key, Cache#{Key => guild_visibility:get_user_viewable_channels(MemberId, State)}} + end. + -spec session_channel_map(user_id(), guild_state()) -> map(). session_channel_map(SessionUserId, State) -> case guild_visibility_channels:get_cached_viewable_channel_map(SessionUserId, State) of @@ -284,4 +360,23 @@ non_member_candidate_is_dropped_test() -> ?assertEqual([], filter_member_ids(10, [777], State)), ?assertEqual([], reference_filter_member_ids(10, [777], State)). +filter_session_member_ids_matches_filter_member_ids_test() -> + State = test_state(), + Ids = candidate_ids(), + Viewers = [10, 30, 7, 99], + Requests = [{Viewer, Viewer, Ids} || Viewer <- Viewers], + ?assertEqual( + maps:from_list([{Viewer, filter_member_ids(Viewer, Ids, State)} || Viewer <- Viewers]), + filter_session_member_ids(Requests, State) + ). + +filter_session_member_ids_reads_the_session_viewable_map_test() -> + State = (test_state())#{ + sessions => #{<<"s10">> => #{user_id => 10, viewable_channels => #{600 => true}}} + }, + ?assertEqual( + #{<<"s10">> => [30, 31, 7, 99, 4242]}, + filter_session_member_ids([{<<"s10">>, 10, [20, 30, 31, 7, 99, 4242]}], State) + ). + -endif. diff --git a/fluxer_gateway/src/presence/presence_targets.erl b/fluxer_gateway/src/presence/presence_targets.erl index aa718ec79..e2747618c 100644 --- a/fluxer_gateway/src/presence/presence_targets.erl +++ b/fluxer_gateway/src/presence/presence_targets.erl @@ -9,6 +9,8 @@ group_dm_channel_recipient_ids/2, dm_recipients_from_state/1, dm_channel_recipient_ids/2, + direct_dm_partner_ids/1, + dm_partner_presence_enabled/1, map_from_ids/1 ]). @@ -32,10 +34,11 @@ accumulate_friend_id(_UserId, _Type, Acc) -> -spec group_dm_recipients_from_state(state()) -> #{channel_id() => #{user_id() => true}}. group_dm_recipients_from_state(State) -> UserId = maps:get(user_id, State, undefined), + Eligible = eligible_dm_partner_ids(State), Channels = maps:get(channels, State, #{}), maps:fold( fun(ChannelId, Channel, Acc) -> - accumulate_dm_channel(ChannelId, Channel, UserId, Acc) + accumulate_dm_channel(ChannelId, Channel, UserId, Eligible, Acc) end, #{}, Channels @@ -45,20 +48,83 @@ group_dm_recipients_from_state(State) -> dm_recipients_from_state(State) -> group_dm_recipients_from_state(State). --spec accumulate_dm_channel(term(), term(), user_id() | undefined, map()) -> map(). -accumulate_dm_channel(ChannelId, Channel, UserId, Acc) when +-spec accumulate_dm_channel( + term(), term(), user_id() | undefined, #{user_id() => true}, map() +) -> map(). +accumulate_dm_channel(ChannelId, Channel, UserId, Eligible, Acc) when is_integer(ChannelId), is_map(Channel) -> - case is_group_dm_channel_type(maps:get(<<"type">>, Channel, 0)) of - true -> - RecipientIds = extract_recipient_ids(Channel), - Acc#{ChannelId => map_from_ids([Rid || Rid <- RecipientIds, Rid =/= UserId])}; - false -> + case maps:get(<<"type">>, Channel, 0) of + 3 -> + Acc#{ChannelId => map_from_ids(other_recipient_ids(Channel, UserId))}; + 1 -> + accumulate_direct_dm( + ChannelId, other_recipient_ids(Channel, UserId), Eligible, Acc + ); + _ -> Acc end; -accumulate_dm_channel(_ChannelId, _Channel, _UserId, Acc) -> +accumulate_dm_channel(_ChannelId, _Channel, _UserId, _Eligible, Acc) -> Acc. +-spec accumulate_direct_dm(channel_id(), [user_id()], #{user_id() => true}, map()) -> map(). +accumulate_direct_dm(ChannelId, RecipientIds, Eligible, Acc) -> + case [Rid || Rid <- RecipientIds, maps:is_key(Rid, Eligible)] of + [] -> Acc; + EligibleIds -> Acc#{ChannelId => map_from_ids(EligibleIds)} + end. + +-spec other_recipient_ids(map(), user_id() | undefined) -> [user_id()]. +other_recipient_ids(Channel, UserId) -> + [Rid || Rid <- extract_recipient_ids(Channel), Rid =/= UserId]. + +-spec eligible_dm_partner_ids(state()) -> #{user_id() => true}. +eligible_dm_partner_ids(State) -> + case dm_partner_presence_enabled(maps:get(user_id, State, undefined)) of + true -> connected_guild_dm_partners(State); + false -> #{} + end. + +-spec connected_guild_dm_partners(state()) -> #{user_id() => true}. +connected_guild_dm_partners(State) -> + Guilds = maps:get(guilds, State, #{}), + maps:fold( + fun(GuildId, PartnerIds, Acc) -> + case maps:get(GuildId, Guilds, undefined) of + {Pid, _Ref} when is_pid(Pid), is_map(PartnerIds) -> maps:merge(Acc, PartnerIds); + _ -> Acc + end + end, + #{}, + maps:get(dm_mutual_by_guild, State, #{}) + ). + +-spec direct_dm_partner_ids(state()) -> [user_id()]. +direct_dm_partner_ids(State) -> + UserId = maps:get(user_id, State, undefined), + PartnerIds = maps:fold( + fun(_ChannelId, Channel, Acc) -> accumulate_direct_partner(Channel, UserId, Acc) end, + [], + maps:get(channels, State, #{}) + ), + lists:usort(PartnerIds). + +-spec accumulate_direct_partner(term(), user_id() | undefined, [user_id()]) -> [user_id()]. +accumulate_direct_partner(#{<<"type">> := 1} = Channel, UserId, Acc) -> + other_recipient_ids(Channel, UserId) ++ Acc; +accumulate_direct_partner(_Channel, _UserId, Acc) -> + Acc. + +-spec dm_partner_presence_enabled(term()) -> boolean(). +dm_partner_presence_enabled(UserId) when is_integer(UserId) -> + case application:get_env(fluxer_gateway, dm_presence_mutual_context, true) of + true -> true; + {users, UserIds} when is_list(UserIds) -> lists:member(UserId, UserIds); + _ -> false + end; +dm_partner_presence_enabled(_UserId) -> + false. + -spec is_group_dm_channel_type(term()) -> boolean(). is_group_dm_channel_type(3) -> true; is_group_dm_channel_type(_) -> false. diff --git a/fluxer_gateway/src/session/session.erl b/fluxer_gateway/src/session/session.erl index 0f7617242..84ed0c3d1 100644 --- a/fluxer_gateway/src/session/session.erl +++ b/fluxer_gateway/src/session/session.erl @@ -238,6 +238,10 @@ handle_info({call_reconnect, ChannelId, Attempt}, State) when session_connection:handle_call_reconnect(ChannelId, Attempt, State); handle_info({gateway_timing_update, Timings}, State) -> {noreply, gateway_timings:merge_state(Timings, State)}; +handle_info({dm_partner_mutual, GuildId, PartnerIds}, State) when + is_integer(GuildId), is_list(PartnerIds) +-> + session_dm_partners:handle_mutual(GuildId, PartnerIds, State); handle_info(Msg, State) -> handle_info_lifecycle(Msg, State). diff --git a/fluxer_gateway/src/session/session_connection_guild.erl b/fluxer_gateway/src/session/session_connection_guild.erl index 56bc270b3..450eac104 100644 --- a/fluxer_gateway/src/session/session_connection_guild.erl +++ b/fluxer_gateway/src/session/session_connection_guild.erl @@ -460,7 +460,8 @@ finalize_guild_monitor(GuildId, GuildPid, Guilds0, State, ReadyFun) -> apply_ready_fun(GuildId, GuildPid, ReadyFun, State) -> case ReadyFun(State) of {noreply, ReadyState} -> - {noreply, maybe_replay_guild_subscriptions(GuildId, GuildPid, ReadyState)}; + ReplayedState = maybe_replay_guild_subscriptions(GuildId, GuildPid, ReadyState), + {noreply, session_dm_partners:register_guild(GuildId, GuildPid, ReplayedState)}; {stop, normal, ReadyState} -> {stop, normal, ReadyState} end. diff --git a/fluxer_gateway/src/session/session_dispatch_presence.erl b/fluxer_gateway/src/session/session_dispatch_presence.erl index ba28dc8c7..7bf6eadc3 100644 --- a/fluxer_gateway/src/session/session_dispatch_presence.erl +++ b/fluxer_gateway/src/session/session_dispatch_presence.erl @@ -105,14 +105,33 @@ maybe_flush_pending_presences(relationship_add, Data, State) -> maybe_flush_pending_presences(relationship_update, Data, State) -> maybe_flush_relationship_pending_presences(Data, State); maybe_flush_pending_presences(channel_create, Data, State) -> - flush_dm_channel_pending_presences(Data, State); + flush_dm_channel_pending_presences(Data, maybe_register_dm_partners(Data, State)); maybe_flush_pending_presences(channel_update, Data, State) -> - flush_dm_channel_pending_presences(Data, State); + flush_dm_channel_pending_presences(Data, maybe_register_dm_partners(Data, State)); +maybe_flush_pending_presences(channel_delete, Data, State) -> + {maybe_register_dm_partners(Data, State), []}; +maybe_flush_pending_presences(guild_delete, Data, State) -> + {forget_deleted_guild(Data, State), []}; maybe_flush_pending_presences(channel_recipient_add, Data, State) -> flush_added_recipient_pending_presences(Data, State); maybe_flush_pending_presences(_, _, State) -> {State, []}. +-spec maybe_register_dm_partners(term(), session_state()) -> session_state(). +maybe_register_dm_partners(#{<<"type">> := 1}, State) -> + session_dm_partners:register_all(State); +maybe_register_dm_partners(_Data, State) -> + State. + +-spec forget_deleted_guild(term(), session_state()) -> session_state(). +forget_deleted_guild(#{<<"id">> := RawGuildId}, State) -> + case snowflake_id:parse_maybe(RawGuildId) of + GuildId when is_integer(GuildId) -> session_dm_partners:forget_guild(GuildId, State); + _ -> State + end; +forget_deleted_guild(_Data, State) -> + State. + -spec flush_dm_channel_pending_presences(map(), session_state()) -> {session_state(), [user_id()]}. flush_dm_channel_pending_presences(Data, State) -> @@ -229,6 +248,7 @@ event_changes_presence_targets(channel_update) -> true; event_changes_presence_targets(channel_delete) -> true; event_changes_presence_targets(channel_recipient_add) -> true; event_changes_presence_targets(channel_recipient_remove) -> true; +event_changes_presence_targets(guild_delete) -> true; event_changes_presence_targets(_) -> false. -spec sync_presence_targets([user_id()], session_state()) -> session_state(). diff --git a/fluxer_gateway/src/session/session_dm_partners.erl b/fluxer_gateway/src/session/session_dm_partners.erl new file mode 100644 index 000000000..265379972 --- /dev/null +++ b/fluxer_gateway/src/session/session_dm_partners.erl @@ -0,0 +1,80 @@ +%% SPDX-License-Identifier: AGPL-3.0-or-later + +-module(session_dm_partners). +-typing([eqwalizer]). + +-export([register_all/1, register_guild/3, handle_mutual/3, forget_guild/2]). + +-type session_state() :: map(). +-type guild_id() :: integer(). +-type user_id() :: integer(). + +-define(MAX_PARTNERS, 1000). + +-spec register_all(session_state()) -> session_state(). +register_all(State) -> + case enabled(State) of + true -> + Partners = partner_ids(State), + maps:foreach( + fun(_GuildId, GuildRef) -> cast_partners(GuildRef, Partners, State) end, + maps:get(guilds, State, #{}) + ), + State; + false -> + State + end. + +-spec register_guild(guild_id(), pid(), session_state()) -> session_state(). +register_guild(_GuildId, GuildPid, State) -> + case enabled(State) of + true -> + cast_partners({GuildPid, undefined}, partner_ids(State), State), + State; + false -> + State + end. + +-spec handle_mutual(guild_id(), [user_id()], session_state()) -> {noreply, session_state()}. +handle_mutual(GuildId, PartnerIds, State) -> + case maps:get(GuildId, maps:get(guilds, State, #{}), undefined) of + {Pid, _Ref} when is_pid(Pid) -> + {noreply, put_guild_partners(GuildId, PartnerIds, State)}; + _ -> + {noreply, State} + end. + +-spec forget_guild(guild_id(), session_state()) -> session_state(). +forget_guild(GuildId, State) -> + Current = maps:get(dm_mutual_by_guild, State, #{}), + State#{dm_mutual_by_guild => maps:remove(GuildId, Current)}. + +-spec put_guild_partners(guild_id(), [user_id()], session_state()) -> session_state(). +put_guild_partners(GuildId, PartnerIds, State) -> + Current = maps:get(dm_mutual_by_guild, State, #{}), + Next = + case presence_targets:map_from_ids(PartnerIds) of + Empty when map_size(Empty) =:= 0 -> maps:remove(GuildId, Current); + Ids -> Current#{GuildId => Ids} + end, + case Next =:= Current of + true -> + State; + false -> + session_dispatch_presence:sync_presence_targets(State#{dm_mutual_by_guild => Next}) + end. + +-spec cast_partners(term(), [user_id()], session_state()) -> ok. +cast_partners({Pid, _Ref}, Partners, State) when is_pid(Pid) -> + gen_server:cast(Pid, {update_dm_partners, maps:get(id, State), Partners}); +cast_partners(_GuildRef, _Partners, _State) -> + ok. + +-spec partner_ids(session_state()) -> [user_id()]. +partner_ids(State) -> + lists:sublist(presence_targets:direct_dm_partner_ids(State), ?MAX_PARTNERS). + +-spec enabled(session_state()) -> boolean(). +enabled(State) -> + maps:get(bot, State, false) =/= true andalso + presence_targets:dm_partner_presence_enabled(maps:get(user_id, State, undefined)). diff --git a/fluxer_gateway/test/guild_dm_partners_tests.erl b/fluxer_gateway/test/guild_dm_partners_tests.erl new file mode 100644 index 000000000..92caf4a4d --- /dev/null +++ b/fluxer_gateway/test/guild_dm_partners_tests.erl @@ -0,0 +1,196 @@ +%% SPDX-License-Identifier: AGPL-3.0-or-later + +-module(guild_dm_partners_tests). +-typing([eqwalizer]). + +-include_lib("eunit/include/eunit.hrl"). + +-define(GUILD_ID, 42). +-define(VIEWER_ROLE, 1000). +-define(OTHER_ROLE, 2000). +-define(SESSION, <<"session-10">>). +-define(FLAG, dm_presence_mutual_context). + +registration_notifies_mutual_partners_test() -> + with_flag(true, fun() -> + State = register_partners(state(), [20, 30, 777]), + ?assertEqual({dm_partner_mutual, ?GUILD_ID, [20]}, receive_mutual()), + ?assert(maps:is_key(?SESSION, maps:get(dm_partners, State))) + end). + +unknown_session_is_not_registered_test() -> + with_flag(true, fun() -> + {noreply, State} = guild_dm_partners:handle_cast( + {update_dm_partners, <<"missing">>, [20]}, state() + ), + ?assertEqual(#{}, maps:get(dm_partners, State, #{})), + ?assertEqual(none, receive_mutual()) + end). + +disabled_flag_does_not_register_test() -> + with_flag(false, fun() -> + State = register_partners(state(), [20]), + ?assertEqual(#{}, maps:get(dm_partners, State, #{})), + ?assertEqual(none, receive_mutual()) + end). + +unchanged_eligibility_is_not_sent_again_test() -> + with_flag(true, fun() -> + State = register_partners(state(), [20]), + {dm_partner_mutual, _, [20]} = receive_mutual(), + _ = register_partners(State, [20, 777]), + ?assertEqual(none, receive_mutual()) + end). + +partner_leaving_the_guild_withdraws_the_partner_test() -> + with_flag(true, fun() -> + Before = register_partners(state(), [20]), + {dm_partner_mutual, _, [20]} = receive_mutual(), + After = without_member(20, Before), + _ = guild_dm_partners:maybe_reevaluate( + guild_member_remove, #{<<"user">> => #{<<"id">> => <<"20">>}}, Before, After + ), + ?assertEqual({dm_partner_mutual, ?GUILD_ID, []}, receive_mutual()) + end). + +guild_wide_event_without_visibility_change_does_nothing_test() -> + with_flag(true, fun() -> + Before = register_partners(state(), [20]), + {dm_partner_mutual, _, [20]} = receive_mutual(), + Renamed = with_channels( + [ + (channel(500, viewer_overwrites()))#{<<"name">> => <<"renamed">>}, + channel(600, other_overwrites()) + ], + Before + ), + _ = guild_dm_partners:maybe_reevaluate(channel_update, #{}, Before, Renamed), + ?assertEqual(none, receive_mutual()) + end). + +guild_wide_visibility_change_reevaluates_test() -> + with_flag(true, fun() -> + Before = register_partners(state(), [30]), + ?assertEqual(none, receive_mutual()), + Opened = with_channels( + [ + channel(500, viewer_overwrites() ++ [overwrite(?OTHER_ROLE, 0)]), + channel(600, other_overwrites()) + ], + Before + ), + _ = guild_dm_partners:maybe_reevaluate(channel_update, #{}, Before, Opened), + ?assertEqual({dm_partner_mutual, ?GUILD_ID, [30]}, receive_mutual()) + end). + +disconnected_sessions_are_dropped_on_reevaluation_test() -> + with_flag(true, fun() -> + Before = register_partners(state(), [20]), + {dm_partner_mutual, _, [20]} = receive_mutual(), + Gone = Before#{sessions => #{}}, + After = guild_dm_partners:maybe_reevaluate(guild_role_delete, #{}, Before, Gone#{ + virtual_channel_access => #{1 => sets:new()} + }), + ?assertEqual(#{}, maps:get(dm_partners, After)) + end). + +register_partners(State, PartnerIds) -> + {noreply, NewState} = guild_dm_partners:handle_cast( + {update_dm_partners, ?SESSION, PartnerIds}, State + ), + NewState. + +state() -> + Base = #{ + id => ?GUILD_ID, + virtual_channel_access => #{}, + data => data( + [channel(500, viewer_overwrites()), channel(600, other_overwrites())], members() + ) + }, + Viewable = guild_sessions:build_viewable_channel_map( + guild_visibility:get_user_viewable_channels(10, Base) + ), + Base#{ + sessions => #{ + ?SESSION => #{user_id => 10, pid => self(), viewable_channels => Viewable} + } + }. + +data(Channels, Members) -> + #{ + <<"guild">> => #{<<"owner_id">> => <<"7">>}, + <<"roles">> => [role(?GUILD_ID), role(?VIEWER_ROLE), role(?OTHER_ROLE)], + <<"members">> => Members, + <<"channels">> => Channels, + <<"channel_index">> => guild_data_index:build_id_index(Channels) + }. + +members() -> + maps:from_list( + [{7, member(7, [])}] ++ + [{Id, member(Id, [?VIEWER_ROLE])} || Id <- [10, 20]] ++ + [{30, member(30, [?OTHER_ROLE])}] + ). + +without_member(UserId, #{data := Data} = State) -> + State#{data => Data#{<<"members">> => maps:remove(UserId, maps:get(<<"members">>, Data))}}. + +with_channels(Channels, #{data := Data} = State) -> + State#{ + data => Data#{ + <<"channels">> => Channels, + <<"channel_index">> => guild_data_index:build_id_index(Channels) + } + }. + +role(RoleId) -> + #{<<"id">> => integer_to_binary(RoleId), <<"permissions">> => <<"0">>}. + +member(UserId, RoleIds) -> + #{ + <<"user">> => #{<<"id">> => integer_to_binary(UserId)}, + <<"roles">> => [integer_to_binary(RoleId) || RoleId <- RoleIds] + }. + +viewer_overwrites() -> + [overwrite(?VIEWER_ROLE, 0)]. + +other_overwrites() -> + [overwrite(?OTHER_ROLE, 0)]. + +overwrite(TargetId, Type) -> + #{ + <<"id">> => integer_to_binary(TargetId), + <<"type">> => Type, + <<"allow">> => integer_to_binary(constants:view_channel_permission()), + <<"deny">> => <<"0">> + }. + +channel(ChannelId, Overwrites) -> + #{ + <<"id">> => integer_to_binary(ChannelId), + <<"type">> => 0, + <<"permission_overwrites">> => Overwrites + }. + +receive_mutual() -> + receive + {dm_partner_mutual, _, _} = Msg -> Msg + after 100 -> + none + end. + +with_flag(Value, Fun) -> + Previous = application:get_env(fluxer_gateway, ?FLAG), + application:set_env(fluxer_gateway, ?FLAG, Value), + try + Fun() + after + restore_flag(Previous) + end. + +restore_flag(undefined) -> + application:unset_env(fluxer_gateway, ?FLAG); +restore_flag({ok, Value}) -> + application:set_env(fluxer_gateway, ?FLAG, Value). diff --git a/fluxer_gateway/test/presence_targets_dm_partner_tests.erl b/fluxer_gateway/test/presence_targets_dm_partner_tests.erl new file mode 100644 index 000000000..aa6ecefd0 --- /dev/null +++ b/fluxer_gateway/test/presence_targets_dm_partner_tests.erl @@ -0,0 +1,93 @@ +%% SPDX-License-Identifier: AGPL-3.0-or-later + +-module(presence_targets_dm_partner_tests). +-typing([eqwalizer]). + +-include_lib("eunit/include/eunit.hrl"). + +-define(SELF, 1). +-define(GUILD, 900). +-define(FLAG, dm_presence_mutual_context). + +eligible_direct_dm_is_a_presence_target_test() -> + with_flag(true, fun() -> + State = state(#{?GUILD => #{2 => true}}, connected()), + ?assertEqual( + #{100 => #{2 => true}, 300 => #{5 => true, 6 => true}}, + presence_targets:group_dm_recipients_from_state(State) + ) + end). + +direct_dm_without_mutual_context_is_not_a_target_test() -> + with_flag(true, fun() -> + State = state(#{}, connected()), + ?assertEqual( + #{300 => #{5 => true, 6 => true}}, + presence_targets:group_dm_recipients_from_state(State) + ) + end). + +mutual_context_from_a_disconnected_guild_is_ignored_test() -> + with_flag(true, fun() -> + State = state(#{?GUILD => #{2 => true}}, #{?GUILD => undefined}), + ?assertEqual( + #{300 => #{5 => true, 6 => true}}, + presence_targets:group_dm_recipients_from_state(State) + ) + end). + +disabled_flag_keeps_direct_dms_out_test() -> + with_flag(false, fun() -> + State = state(#{?GUILD => #{2 => true, 3 => true}}, connected()), + ?assertEqual( + #{300 => #{5 => true, 6 => true}}, + presence_targets:group_dm_recipients_from_state(State) + ) + end). + +user_allowlist_gates_direct_dms_test() -> + State = state(#{?GUILD => #{2 => true}}, connected()), + with_flag({users, [?SELF]}, fun() -> + ?assert(maps:is_key(100, presence_targets:group_dm_recipients_from_state(State))) + end), + with_flag({users, [42]}, fun() -> + ?assertNot(maps:is_key(100, presence_targets:group_dm_recipients_from_state(State))) + end). + +direct_dm_partner_ids_lists_one_to_one_partners_only_test() -> + ?assertEqual([2, 3], presence_targets:direct_dm_partner_ids(state(#{}, connected()))). + +state(MutualByGuild, Guilds) -> + #{ + user_id => ?SELF, + guilds => Guilds, + dm_mutual_by_guild => MutualByGuild, + channels => #{ + 100 => channel(1, [?SELF, 2]), + 200 => channel(1, [?SELF, 3]), + 300 => channel(3, [?SELF, 5, 6]) + } + }. + +channel(Type, RecipientIds) -> + #{ + <<"type">> => Type, + <<"recipients">> => [#{<<"id">> => integer_to_binary(Id)} || Id <- RecipientIds] + }. + +connected() -> + #{?GUILD => {self(), make_ref()}}. + +with_flag(Value, Fun) -> + Previous = application:get_env(fluxer_gateway, ?FLAG), + application:set_env(fluxer_gateway, ?FLAG, Value), + try + Fun() + after + restore_flag(Previous) + end. + +restore_flag(undefined) -> + application:unset_env(fluxer_gateway, ?FLAG); +restore_flag({ok, Value}) -> + application:set_env(fluxer_gateway, ?FLAG, Value). diff --git a/fluxer_gateway/test/session_dm_partners_tests.erl b/fluxer_gateway/test/session_dm_partners_tests.erl new file mode 100644 index 000000000..121963c49 --- /dev/null +++ b/fluxer_gateway/test/session_dm_partners_tests.erl @@ -0,0 +1,97 @@ +%% SPDX-License-Identifier: AGPL-3.0-or-later + +-module(session_dm_partners_tests). +-typing([eqwalizer]). + +-include_lib("eunit/include/eunit.hrl"). + +-define(SELF, 1). +-define(GUILD, 900). +-define(SESSION_ID, <<"session-1">>). +-define(FLAG, dm_presence_mutual_context). + +handle_mutual_stores_partners_for_a_connected_guild_test() -> + {noreply, State} = session_dm_partners:handle_mutual(?GUILD, [2], state(connected())), + ?assertEqual(#{?GUILD => #{2 => true}}, maps:get(dm_mutual_by_guild, State)). + +handle_mutual_ignores_a_guild_that_is_not_connected_test() -> + Initial = state(#{}), + ?assertEqual({noreply, Initial}, session_dm_partners:handle_mutual(?GUILD, [2], Initial)). + +handle_mutual_drops_the_guild_when_nothing_is_mutual_test() -> + {noreply, With} = session_dm_partners:handle_mutual(?GUILD, [2], state(connected())), + {noreply, Without} = session_dm_partners:handle_mutual(?GUILD, [], With), + ?assertEqual(#{}, maps:get(dm_mutual_by_guild, Without)). + +forget_guild_drops_its_partners_test() -> + {noreply, With} = session_dm_partners:handle_mutual(?GUILD, [2], state(connected())), + ?assertEqual( + #{}, maps:get(dm_mutual_by_guild, session_dm_partners:forget_guild(?GUILD, With)) + ). + +register_all_sends_direct_partners_to_every_connected_guild_test() -> + with_flag(true, fun() -> + _ = session_dm_partners:register_all(state(connected())), + ?assertEqual({update_dm_partners, ?SESSION_ID, [2, 3]}, receive_cast()) + end). + +register_guild_sends_direct_partners_to_that_guild_test() -> + with_flag(true, fun() -> + _ = session_dm_partners:register_guild(?GUILD, self(), state(#{})), + ?assertEqual({update_dm_partners, ?SESSION_ID, [2, 3]}, receive_cast()) + end). + +register_all_is_inert_when_disabled_test() -> + with_flag(false, fun() -> + _ = session_dm_partners:register_all(state(connected())), + ?assertEqual(none, receive_cast()) + end). + +bots_never_register_test() -> + with_flag(true, fun() -> + _ = session_dm_partners:register_all((state(connected()))#{bot => true}), + ?assertEqual(none, receive_cast()) + end). + +state(Guilds) -> + #{ + id => ?SESSION_ID, + user_id => ?SELF, + presence_pid => undefined, + guilds => Guilds, + channels => #{ + 100 => channel(1, [?SELF, 2]), + 200 => channel(1, [?SELF, 3]), + 300 => channel(3, [?SELF, 5, 6]) + } + }. + +channel(Type, RecipientIds) -> + #{ + <<"type">> => Type, + <<"recipients">> => [#{<<"id">> => integer_to_binary(Id)} || Id <- RecipientIds] + }. + +connected() -> + #{?GUILD => {self(), make_ref()}}. + +receive_cast() -> + receive + {'$gen_cast', Msg} -> Msg + after 100 -> + none + end. + +with_flag(Value, Fun) -> + Previous = application:get_env(fluxer_gateway, ?FLAG), + application:set_env(fluxer_gateway, ?FLAG, Value), + try + Fun() + after + restore_flag(Previous) + end. + +restore_flag(undefined) -> + application:unset_env(fluxer_gateway, ?FLAG); +restore_flag({ok, Value}) -> + application:set_env(fluxer_gateway, ?FLAG, Value).