Compare commits

...
52 changed files with 3378 additions and 1119 deletions
@@ -354,12 +354,10 @@ fn test_config(api_endpoint: String) -> AdminConfig {
static_cdn_endpoint: "https://static.example.test".to_owned(),
admin_endpoint: "https://admin.example.test".to_owned(),
web_app_endpoint: "https://app.example.test".to_owned(),
kv_url: String::new(),
oauth_client_id: "admin-client".to_owned(),
oauth_client_secret: "admin-secret".to_owned(),
oauth_redirect_uri: "https://admin.example.test/callback".to_owned(),
build_version: "test".to_owned(),
release_channel: "test".to_owned(),
self_hosted: false,
proxy: ProxyConfig {
trust_client_ip_header: false,
+10 -11
View File
@@ -35,6 +35,7 @@ import * as FetchUtils from '@app/api/utils/FetchUtils';
import {isJsonRecord, parseJsonRecord, parseJsonWithGuard} from '@app/api/utils/JsonBoundaryUtils';
import {generateRandomUsername} from '@app/api/utils/UsernameGenerator';
import {deriveUsernameFromDisplayName} from '@app/api/utils/UsernameSuggestionUtils';
import {SSO_MOBILE_CALLBACK_URI, SSO_MOBILE_STATE_PREFIX} from '@fluxer/constants/src/SsoConstants';
import {ProfileFieldPrivacyFlags} from '@fluxer/constants/src/UserConstants';
import {ValidationErrorCodes} from '@fluxer/constants/src/ValidationErrorCodes';
import {RegistrationClosedError} from '@fluxer/errors/src/domains/auth/RegistrationClosedError';
@@ -106,7 +107,6 @@ interface JwksCacheEntry {
const CODE_VERIFIER_BYTE_LENGTH = 32;
const STATE_BYTE_LENGTH = 16;
const NONCE_BYTE_LENGTH = 16;
const MOBILE_SSO_REDIRECT_URI = 'fluxer://auth/sso/callback';
let ssoLogger: ILogger | undefined;
@@ -136,11 +136,10 @@ function buildDiscoveryCacheKey(issuer: string): string {
return `sso:oidc-discovery:${key}`;
}
function resolveSsoRedirectUri(requestedRedirectUri: string | undefined, defaultRedirectUri: string): string {
if (!requestedRedirectUri) return defaultRedirectUri;
const trimmed = requestedRedirectUri.trim();
if (!trimmed) return defaultRedirectUri;
if (trimmed === defaultRedirectUri || trimmed === MOBILE_SSO_REDIRECT_URI) return trimmed;
function isMobileSsoRedirectUri(requestedRedirectUri: string | undefined, defaultRedirectUri: string): boolean {
const trimmed = requestedRedirectUri?.trim();
if (!trimmed || trimmed === defaultRedirectUri) return false;
if (trimmed === SSO_MOBILE_CALLBACK_URI) return true;
throw InputValidationError.fromCode('redirect_uri', ValidationErrorCodes.INVALID_URL_FORMAT);
}
@@ -285,16 +284,16 @@ export class SsoService {
redirect_uri: string;
}> {
const config = await this.requireReadyConfig();
const state = randomHexToken(STATE_BYTE_LENGTH);
const isMobile = isMobileSsoRedirectUri(redirectUri, config.redirectUri);
const state = `${isMobile ? SSO_MOBILE_STATE_PREFIX : ''}${randomHexToken(STATE_BYTE_LENGTH)}`;
const codeVerifier = randomBase64UrlToken(CODE_VERIFIER_BYTE_LENGTH);
const codeChallenge = buildCodeChallenge(codeVerifier);
const nonce = randomBase64UrlToken(NONCE_BYTE_LENGTH);
const ssoRedirectUri = resolveSsoRedirectUri(redirectUri, config.redirectUri);
const statePayload: SsoStatePayload = {
codeVerifier,
nonce,
redirectTo: sanitizeSsoRedirectTo(redirectTo),
redirectUri: ssoRedirectUri,
redirectUri: config.redirectUri,
createdAt: Date.now(),
};
const {cache} = this.apiContext.services;
@@ -302,7 +301,7 @@ export class SsoService {
const searchParams = new URLSearchParams({
response_type: 'code',
client_id: config.clientId ?? '',
redirect_uri: ssoRedirectUri,
redirect_uri: config.redirectUri,
scope: config.scope,
state,
code_challenge: codeChallenge,
@@ -324,7 +323,7 @@ export class SsoService {
throw new FeatureTemporarilyDisabledError();
}
}
return {authorization_url: authorizationUrlString, state, redirect_uri: ssoRedirectUri};
return {authorization_url: authorizationUrlString, state, redirect_uri: config.redirectUri};
}
async completeLogin({code, state, request}: {code: string; state: string; request: Request}): Promise<{
@@ -131,6 +131,7 @@ describe('Auth SSO flow', () => {
.body({redirect_to: '/me'})
.execute();
expect(startData.state).toBeTruthy();
expect(startData.state.startsWith('m.')).toBe(false);
expect(startData.authorization_url).toBeTruthy();
const authUrlString = startData.authorization_url;
expect(authUrlString).toContain(`state=${startData.state}`);
@@ -179,7 +180,8 @@ describe('Auth SSO flow', () => {
expect(startData.redirect_uri).not.toContain('evil.example');
expect(startData.authorization_url).toContain(encodeURIComponent(startData.redirect_uri));
});
it('uses the requested mobile SSO redirect URI without changing the post-login redirect', async () => {
it('routes mobile SSO through the default redirect URI without changing the post-login redirect', async () => {
const status = await createBuilderWithoutAuth<{redirect_uri: string}>(harness).get('/auth/sso/status').execute();
const startData = await createBuilderWithoutAuth<SsoStartResponse>(harness)
.post('/auth/sso/start')
.body({
@@ -187,8 +189,10 @@ describe('Auth SSO flow', () => {
redirect_uri: 'fluxer://auth/sso/callback',
})
.execute();
expect(startData.redirect_uri).toBe('fluxer://auth/sso/callback');
expect(getAuthorizationUrlParam(startData.authorization_url, 'redirect_uri')).toBe('fluxer://auth/sso/callback');
expect(startData.redirect_uri).toBe(status.redirect_uri);
expect(getAuthorizationUrlParam(startData.authorization_url, 'redirect_uri')).toBe(status.redirect_uri);
expect(startData.state.startsWith('m.')).toBe(true);
expect(getAuthorizationUrlParam(startData.authorization_url, 'state')).toBe(startData.state);
const email = `sso-mobile-redirect-${Date.now()}@example.com`;
const completeData = await createBuilderWithoutAuth<SsoCompleteResponse>(harness)
.post('/auth/sso/complete')
@@ -9,6 +9,7 @@ import {
deleteChannel,
getChannel,
updateChannel,
updateGuild,
} from '@app/api/channel/tests/ChannelTestUtils';
import {type ApiTestHarness, createApiTestHarness} from '@app/api/test/ApiTestHarness';
import {HTTP_STATUS} from '@app/api/test/TestConstants';
@@ -58,6 +59,34 @@ describe('Channel Operation Permissions', () => {
.expect(HTTP_STATUS.FORBIDDEN)
.execute();
});
it('should gate channels created before a guild becomes adult-only', async () => {
const owner = await createTestAccount(harness);
const minor = await createTestAccount(harness, {dateOfBirth: '2010-01-01'});
const guild = await createGuild(harness, owner.token, 'Later Mature Guild');
const category = await createChannel(harness, owner.token, guild.id, 'category', 4);
const child = await createBuilder<{id: string; nsfw_override?: boolean | null}>(harness, owner.token)
.post(`/guilds/${guild.id}/channels`)
.body({name: 'child', type: 0, parent_id: category.id})
.execute();
const opened = await createBuilder<{id: string}>(harness, owner.token)
.post(`/guilds/${guild.id}/channels`)
.body({name: 'opened', type: 0, nsfw_override: false})
.execute();
const systemChannel = await getChannel(harness, owner.token, guild.system_channel_id!);
expect(systemChannel.nsfw_override ?? null).toBeNull();
expect(category.nsfw_override ?? null).toBeNull();
expect(child.nsfw_override ?? null).toBeNull();
const invite = await createChannelInvite(harness, owner.token, systemChannel.id);
await acceptInvite(harness, minor.token, invite.code);
await updateGuild(harness, owner.token, guild.id, {nsfw: true});
for (const channelId of [systemChannel.id, child.id]) {
await createBuilder(harness, minor.token)
.get(`/channels/${channelId}/messages`)
.expect(HTTP_STATUS.FORBIDDEN)
.execute();
}
await createBuilder(harness, minor.token).get(`/channels/${opened.id}/messages`).expect(HTTP_STATUS.OK).execute();
});
it('should reject member from updating channel without MANAGE_CHANNELS', async () => {
const owner = await createTestAccount(harness);
const member = await createTestAccount(harness);
@@ -134,7 +134,7 @@ export class ChannelOperationsService {
}
}
const requestedNsfwOverride =
params.data.nsfw_override !== undefined ? params.data.nsfw_override : (params.data.nsfw ?? null);
params.data.nsfw_override !== undefined ? params.data.nsfw_override : params.data.nsfw === true ? true : null;
const requestedContentWarningLevel =
params.data.content_warning_level === ContentWarningLevel.CONTENT_WARNING
? ContentWarningLevel.CONTENT_WARNING
@@ -828,7 +828,7 @@ export class GuildOperationsService {
position,
owner_id: null,
recipient_ids: null,
nsfw: false,
nsfw: null,
content_warning_level: null,
content_warning_text: null,
rate_limit_per_user: 0,
@@ -1048,7 +1048,7 @@ export class GuildOperationsService {
position: channel.position,
owner_id: null,
recipient_ids: null,
nsfw: channel.nsfw ?? false,
nsfw: channel.nsfw === true ? true : null,
content_warning_level: null,
content_warning_text: null,
rate_limit_per_user: channel.rate_limit_per_user ?? 0,
@@ -1100,7 +1100,7 @@ export class GuildOperationsService {
position: 0,
owner_id: null,
recipient_ids: null,
nsfw: false,
nsfw: null,
content_warning_level: null,
content_warning_text: null,
rate_limit_per_user: 0,
@@ -384,7 +384,7 @@ describe('Message Search Permissions', () => {
const guild = await createGuild(harness, owner.token, 'Age Restricted Override Guild');
const channel = await createBuilder<{id: string; nsfw_override?: boolean | null}>(harness, owner.token)
.post(`/guilds/${guild.id}/channels`)
.body({name: 'override-channel', type: ChannelTypes.GUILD_TEXT, nsfw: false})
.body({name: 'override-channel', type: ChannelTypes.GUILD_TEXT, nsfw_override: false})
.execute();
expect(channel.nsfw_override).toBe(false);
await sendChannelMessage(harness, owner.token, channel.id, 'age restricted override searchable message');
@@ -62,7 +62,9 @@ describe('queueBlocklistFeedStartupJobs', () => {
expect(await kv.exists(INITIAL_SYNC_KEY)).toBe(0);
await queueBlocklistFeedStartupJobs(kv, createWorkerService(), true);
expect(await kv.ttl(INITIAL_SYNC_KEY)).toBe(21600);
const ttl = await kv.ttl(INITIAL_SYNC_KEY);
expect(ttl).toBeGreaterThanOrEqual(21599);
expect(ttl).toBeLessThanOrEqual(21600);
});
it('a full jobs stream drops the job without failing startup', async () => {
@@ -368,6 +368,7 @@
}
.ssoRetryButton {
display: inline-block;
padding: 0.75rem 1.5rem;
border-radius: 0.625rem;
border: none;
@@ -376,6 +377,7 @@
font-weight: 600;
font-size: 0.95rem;
cursor: pointer;
text-decoration: none;
transition: background 120ms ease;
}
@@ -1,5 +1,6 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
import {PRODUCT_NAME} from '@app/features/app/config/I18nDisplayConstants';
import * as AuthenticationCommands from '@app/features/auth/commands/AuthenticationCommands';
import styles from '@app/features/auth/components/pages/LoginPage.module.css';
import {
@@ -12,6 +13,7 @@ import {safeRedirectTarget} from '@app/features/auth/utils/SafeRedirect';
import {BACK_TO_SIGN_IN_DESCRIPTOR, TRY_AGAIN_DESCRIPTOR} from '@app/features/i18n/utils/CommonMessageDescriptors';
import * as RouterUtils from '@app/features/navigation/utils/RouterUtils';
import * as FormUtils from '@app/lib/forms';
import {SSO_MOBILE_CALLBACK_URI, SSO_MOBILE_STATE_PREFIX} from '@fluxer/constants/src/SsoConstants';
import {msg} from '@lingui/core/macro';
import {Trans, useLingui} from '@lingui/react/macro';
import {observer} from 'mobx-react-lite';
@@ -29,6 +31,10 @@ const FAILED_TO_COMPLETE_SSO_SIGN_IN_DESCRIPTOR = msg({
message: 'Failed to complete SSO sign-in',
comment: 'Short label in the authentication SSO callback page. Keep the tone plain and specific.',
});
const OPEN_PRODUCT_DESCRIPTOR = msg({
message: 'Open {productName}',
comment: 'Button that hands SSO sign-in back to the mobile app. productName is the app name.',
});
const SSO_TIMEOUT_MS = 30_000;
const SsoCallbackPage = observer(function SsoCallbackPage() {
const {i18n} = useLingui();
@@ -37,6 +43,9 @@ const SsoCallbackPage = observer(function SsoCallbackPage() {
const state = params['get']('state');
const providerError = params['get']('error');
const providerErrorDescription = params['get']('error_description');
const mobileCallbackUrl = state?.startsWith(SSO_MOBILE_STATE_PREFIX)
? `${SSO_MOBILE_CALLBACK_URI}${window.location.search}`
: null;
const [error, setError] = useState<string | null>(null);
const [isProcessing, setIsProcessing] = useState(true);
const abortControllerRef = useRef<AbortController | null>(null);
@@ -54,6 +63,10 @@ const SsoCallbackPage = observer(function SsoCallbackPage() {
}
}, []);
useEffect(() => {
if (mobileCallbackUrl) {
window.location.replace(mobileCallbackUrl);
return;
}
const controller = new AbortController();
abortControllerRef.current = controller;
const timeoutId = setTimeout(() => {
@@ -97,7 +110,28 @@ const SsoCallbackPage = observer(function SsoCallbackPage() {
clearTimeout(timeoutId);
controller.abort();
};
}, [code, state, providerError, providerErrorDescription, i18n]);
}, [code, state, providerError, providerErrorDescription, mobileCallbackUrl, i18n]);
if (mobileCallbackUrl) {
return (
<div className={styles.loginContainer} data-flx="auth.sso-callback-page.login-container--mobile">
<h1 className={styles.title} data-flx="auth.sso-callback-page.title--mobile">
<Trans>Completing sign-in…</Trans>
</h1>
<p className={styles.ssoProcessingHint} data-flx="auth.sso-callback-page.sso-processing-hint--mobile">
<Trans>Jump straight to the app to continue.</Trans>
</p>
<div className={styles.ssoCallbackActions} data-flx="auth.sso-callback-page.sso-callback-actions--mobile">
<a
href={mobileCallbackUrl}
className={styles.ssoRetryButton}
data-flx="auth.sso-callback-page.sso-open-app-button"
>
{i18n._(OPEN_PRODUCT_DESCRIPTOR, {productName: PRODUCT_NAME})}
</a>
</div>
</div>
);
}
if (error) {
return (
<div className={styles.loginContainer} data-flx="auth.sso-callback-page.login-container">
@@ -73,7 +73,7 @@ Single sign-on settings for the deployment's OpenID Connect provider.
<sup>2</sup> Each entry is stored lowercased and IDNA encoded, duplicates are collapsed, and an empty entry is dropped
<sup>3</sup> The configured web application endpoint followed by `/auth/sso/callback`. No operation can set it
<sup>3</sup> The configured web application endpoint followed by `/auth/sso/callback`. It is the only redirect URI the provider needs, mobile sign-in included. No operation can set it
## Gateway rollout configuration object
@@ -423,11 +423,11 @@ Starts a single sign-on flow. Authentication is not required. Returns an [SSO st
| Field | Type | Description |
| --- | --- | --- |
| redirect_to?<sup>1</sup> | ?string | The post-authentication redirect to bind to the state |
| redirect_uri?<sup>2</sup> | ?string | The provider callback URI to use instead of the configured default |
| redirect_uri?<sup>2</sup> | ?string | The callback URI the client wants the result delivered to |
<sup>1</sup> Fluxer sanitises the value before binding it to the state and discards a value that does not survive, which the [SSO completion response](#sso-completion-response-object) reports as the empty string. Sanitisation keeps the trimmed value only when it begins with a single `/`, is at most 2,048 characters, and contains no carriage return or line feed
<sup>2</sup> The accepted values are the instance default reported as `redirect_uri` by [get SSO status](#get-sso-status) and the mobile callback `fluxer://auth/sso/callback`, and any other value returns the field code `INVALID_URL_FORMAT`. The accepted value is bound to the state and reused at the token exchange
<sup>2</sup> The accepted values are the instance default reported as `redirect_uri` by [get SSO status](#get-sso-status) and the mobile callback `fluxer://auth/sso/callback`, and any other value returns the field code `INVALID_URL_FORMAT`. The provider always receives the instance default. For the mobile callback the state starts with `m.`, and the web callback page forwards the provider's query string to `fluxer://auth/sso/callback` unchanged
### Response
@@ -349,9 +349,10 @@ safe_gen_server_call(Pid, Request, Timeout) ->
exit:_ -> error
end.
-spec safe_guild_call(integer(), pid(), term(), pos_integer()) -> {ok, term()} | error.
-spec safe_guild_call(integer(), pid(), {atom(), map()}, pos_integer()) ->
{ok, term()} | error.
safe_guild_call(GuildId, Pid, Request, Timeout) ->
try gen_server:call(Pid, Request, Timeout) of
try guild_query_handler:call(Pid, Request, Timeout) of
Reply -> {ok, Reply}
catch
exit:{timeout, _} ->
@@ -366,13 +367,13 @@ safe_guild_call(GuildId, Pid, Request, Timeout) ->
error
end.
-spec retry_after_guild_call_failure(integer(), pid(), term(), pos_integer()) ->
-spec retry_after_guild_call_failure(integer(), pid(), {atom(), map()}, pos_integer()) ->
{ok, term()} | error.
retry_after_guild_call_failure(GuildId, Pid, Request, Timeout) ->
delete_cached_guild_pid(GuildId, Pid),
retry_guild_call(GuildId, Request, Timeout).
-spec retry_guild_call(integer(), term(), pos_integer()) -> {ok, term()} | error.
-spec retry_guild_call(integer(), {atom(), map()}, pos_integer()) -> {ok, term()} | error.
retry_guild_call(GuildId, Request, Timeout) ->
case get_guild_pid(GuildId) of
{ok, NewPid} ->
@@ -381,9 +382,10 @@ retry_guild_call(GuildId, Request, Timeout) ->
error
end.
-spec retry_guild_call_pid(integer(), pid(), term(), pos_integer()) -> {ok, term()} | error.
-spec retry_guild_call_pid(integer(), pid(), {atom(), map()}, pos_integer()) ->
{ok, term()} | error.
retry_guild_call_pid(GuildId, NewPid, Request, Timeout) ->
try gen_server:call(NewPid, Request, Timeout) of
try guild_query_handler:call(NewPid, Request, Timeout) of
Reply -> {ok, Reply}
catch
throw:_Reason ->
@@ -72,7 +72,7 @@ optional_channel_id(Value) ->
-spec get_auth_context_from_guild(pid(), integer() | null, integer() | null) -> term().
get_auth_context_from_guild(Pid, UserId, ChannelId) ->
Request = {get_guild_auth_context, #{user_id => UserId, channel_id => ChannelId}},
case gen_server:call(Pid, Request, ?GUILD_CALL_TIMEOUT) of
case guild_query_handler:call(Pid, Request, ?GUILD_CALL_TIMEOUT) of
#{auth_context := null} ->
gateway_rpc_error:raise(<<"forbidden">>);
#{auth_context := AuthContext} ->
@@ -91,7 +91,8 @@ optional_user_id(Value) ->
-spec get_data_from_guild(pid(), integer() | null) -> term().
get_data_from_guild(Pid, UserId) ->
case gen_server:call(Pid, {get_guild_data, #{user_id => UserId}}, ?GUILD_CALL_TIMEOUT) of
Request = {get_guild_data, #{user_id => UserId}},
case guild_query_handler:call(Pid, Request, ?GUILD_CALL_TIMEOUT) of
#{guild_data := null, error_reason := <<"forbidden">>} ->
gateway_rpc_error:raise(<<"forbidden">>);
#{guild_data := null} ->
@@ -207,12 +207,9 @@ member_from_guild(GuildId, Pid, Msg) ->
integer(), integer()
) -> {ok, [integer()]} | error.
get_members_with_role_cached_or_rpc(GuildId, RoleId) ->
case guild_permission_cache:get_snapshot(GuildId) of
{ok, Snapshot} ->
Data = maps:get(data, Snapshot, #{}),
MemberRoleIndex = guild_data_index:member_role_index(Data),
RoleMembers = maps:get(RoleId, MemberRoleIndex, #{}),
{ok, lists:sort(maps:keys(RoleMembers))};
case guild_permission_cache:get_role_members(GuildId, RoleId) of
{ok, UserIds} ->
{ok, UserIds};
{error, not_found} ->
get_members_with_role_via_rpc(GuildId, RoleId)
end.
@@ -109,7 +109,7 @@ do_resolve_sources(GuildId, Req) ->
-spec guild_call_sources(pid(), map()) -> map().
guild_call_sources(Pid, Req) ->
Msg = {resolve_mention_sources, Req},
case gen_server:call(Pid, Msg, ?GUILD_CALL_TIMEOUT) of
case guild_query_handler:call(Pid, Msg, ?GUILD_CALL_TIMEOUT) of
#{direct_user_ids := D, role_user_ids := R, everyone_user_ids := E} ->
#{
<<"direct_user_ids">> => fmt_ids(D),
@@ -139,7 +139,7 @@ do_resolve_sources_page(GuildId, Request) ->
-spec guild_call_sources_page(pid(), map()) -> map().
guild_call_sources_page(Pid, Request) ->
Msg = {resolve_mention_sources_page, Request},
case gen_server:call(Pid, Msg, ?GUILD_CALL_TIMEOUT) of
case guild_query_handler:call(Pid, Msg, ?GUILD_CALL_TIMEOUT) of
#{mentions := Mentions, next_cursor := NextCursor} ->
#{
<<"mentions">> => fmt_mention_entries(Mentions),
@@ -161,15 +161,15 @@ build_mention_req(ChannelId, AuthorId, ME, MH, RIds, UIDs) ->
user_ids => validation:snowflake_list_or_throw(<<"user_ids">>, UIDs)
}.
-spec mention_guild_call(integer(), term(), binary()) -> term().
-spec mention_guild_call(integer(), {atom(), map()}, binary()) -> term().
mention_guild_call(GuildId, Msg, ErrorBin) ->
gateway_rpc_guild_infra:with_guild(GuildId, fun(Pid) ->
guild_call_user_ids(Pid, Msg, ErrorBin)
end).
-spec guild_call_user_ids(pid(), term(), binary()) -> map().
-spec guild_call_user_ids(pid(), {atom(), map()}, binary()) -> map().
guild_call_user_ids(Pid, Msg, ErrorBin) ->
case gen_server:call(Pid, Msg, ?GUILD_CALL_TIMEOUT) of
case guild_query_handler:call(Pid, Msg, ?GUILD_CALL_TIMEOUT) of
#{user_ids := Ids} ->
#{<<"user_ids">> => [integer_to_binary(U) || U <- Ids]};
_ ->
@@ -322,6 +322,54 @@ parse_page_params_defaults_optional_fields_test() ->
maps:get(limit, Req)
).
mention_guild_calls_carry_the_caller_deadline_test() ->
Cases = [
{
fun(Pid) -> guild_call_sources_page(Pid, #{limit => 5}) end,
#{mentions => [], next_cursor => undefined},
#{<<"mentions">> => [], <<"next_cursor">> => null}
},
{
fun(Pid) -> guild_call_sources(Pid, #{user_ids => [7]}) end,
#{direct_user_ids => [7], role_user_ids => [], everyone_user_ids => []},
#{
<<"direct_user_ids">> => [<<"7">>],
<<"role_user_ids">> => [],
<<"everyone_user_ids">> => []
}
},
{
fun(Pid) ->
guild_call_user_ids(
Pid,
{get_users_to_mention_by_user_ids, #{user_ids => [7]}},
<<"users_error">>
)
end,
#{user_ids => [7]},
#{<<"user_ids">> => [<<"7">>]}
}
],
lists:foreach(fun assert_call_carries_deadline/1, Cases).
assert_call_carries_deadline({Call, GuildReply, Expected}) ->
Self = self(),
Guild = spawn(fun() ->
receive
{'$gen_call', From, Msg} ->
Self ! {guild_request, Msg},
gen_server:reply(From, GuildReply)
end
end),
Before = os:system_time(millisecond),
?assertEqual(Expected, Call(Guild)),
receive
{guild_request, {_Tag, #{deadline := Deadline}}} ->
?assert(Deadline >= Before + ?GUILD_CALL_TIMEOUT)
after 1000 ->
?assert(false, guild_request_not_received)
end.
parse_page_params_rejects_malformed_optional_ids_test() ->
Base = #{
<<"guild_id">> => <<"1">>,
+11 -4
View File
@@ -13,6 +13,7 @@
-endif.
-define(HIBERNATE_TIMEOUT, 60000).
-define(FULLSWEEP_AFTER, 100).
-define(VOICE_MEMBERS_TABLE_WARNED, {?MODULE, voice_members_table_unavailable}).
-type guild_state() :: map().
@@ -35,7 +36,7 @@ update_counts(State) -> guild_maintenance:update_counts(State).
-spec init(map()) -> {ok, guild_state(), timeout()}.
init(GuildState) ->
process_flag(trap_exit, true),
erlang:process_flag(fullsweep_after, 10),
erlang:process_flag(fullsweep_after, ?FULLSWEEP_AFTER),
State0 = guild_init:init_base_state(GuildState),
State1 = guild_init:init_member_list(State0),
State2 = guild_init:init_counts(State1),
@@ -83,6 +84,7 @@ call_handler(Tag) -> query_call_handler(Tag).
-spec query_call_handler(atom()) -> query | voice | subscription | undefined.
query_call_handler(get_counts) -> query;
query_call_handler(get_user_counts) -> query;
query_call_handler(get_viewer_counts) -> query;
query_call_handler(get_channel_member_counts) -> query;
query_call_handler(get_large_guild_metadata) -> query;
query_call_handler(get_users_to_mention_by_roles) -> query;
@@ -214,8 +216,8 @@ handle_info(presence_reconcile, State) ->
guild_presence_reconcile:start_async(State),
_ = guild_presence_reconcile:schedule(),
{noreply, State};
handle_info({presence_reconcile_apply, PresenceById}, State) when is_map(PresenceById) ->
{noreply, guild_presence_reconcile:apply_reconcile_result(PresenceById, State)};
handle_info({presence_reconcile_apply, Mismatches}, State) when is_list(Mismatches) ->
{noreply, guild_presence_reconcile:apply_mismatches(Mismatches, State)};
handle_info({reconcile_user_presence, UserId}, State) ->
{noreply, guild_presence_reconcile:reconcile_user(UserId, State)};
handle_info({clear_stale_cached_voice_states, ConnectionIds}, State) ->
@@ -228,6 +230,11 @@ handle_info({check_auto_stop_empty, Token}, State) ->
handle_auto_stop_info(Token, State);
handle_info(check_auto_stop_empty, State) ->
{noreply, State};
handle_info({timeout, TimerRef, member_list_sync_item_cache_rotate}, State) when
is_reference(TimerRef)
->
ok = guild_member_list_subscribe:handle_sync_item_cache_timeout(TimerRef),
{noreply, State};
handle_info(timeout, State) ->
{noreply, State, hibernate};
handle_info(_, State) ->
@@ -443,7 +450,7 @@ terminate(Reason, State) ->
-spec code_change(term(), guild_state(), term()) -> {ok, guild_state()}.
code_change(_OldVsn, State, _Extra) ->
erlang:process_flag(fullsweep_after, 10),
erlang:process_flag(fullsweep_after, ?FULLSWEEP_AFTER),
erlang:garbage_collect(),
{ok, State}.
+1 -41
View File
@@ -4,7 +4,7 @@
-typing([eqwalizer]).
-behaviour(gen_server).
-export([start_link/1, cast_presence/5, cast_member_list/4, cast_event/4, ensure/1]).
-export([start_link/1, cast_member_list/4, cast_event/4, ensure/1]).
-define(MAX_MAILBOX, 500).
-define(IDLE_GC_DELAY_MS, 1000).
@@ -36,13 +36,6 @@ ensure_alive(Pid, State) ->
false -> start_and_store(State)
end.
-spec cast_presence(broadcaster_ref(), integer(), map(), snapshot(), snapshot()) -> boolean().
cast_presence(BroadcasterPid, UserId, PresenceMap, OldSnapshot, NewSnapshot) ->
maybe_cast(
BroadcasterPid,
{presence_broadcast, UserId, PresenceMap, OldSnapshot, NewSnapshot}
).
-spec cast_member_list(broadcaster_ref(), integer(), snapshot(), snapshot()) -> boolean().
cast_member_list(BroadcasterPid, UserId, OldSnapshot, NewSnapshot) ->
maybe_cast(
@@ -80,9 +73,6 @@ handle_call(_Req, _From, State) ->
{reply, {error, unknown_call}, State}.
-spec handle_cast(term(), map()) -> {noreply, map()}.
handle_cast({presence_broadcast, UserId, PresenceMap, _OldSnapshot, NewSnapshot}, State) ->
maybe_broadcast_presence(UserId, PresenceMap, NewSnapshot),
{noreply, maybe_schedule_gc(State)};
handle_cast({member_list_broadcast, _UserId, _OldSnapshot, _NewSnapshot}, State) ->
{noreply, maybe_schedule_gc(State)};
handle_cast({event_broadcast, Event, EncodedPayload, FilteredSessionPids}, State) ->
@@ -91,36 +81,6 @@ handle_cast({event_broadcast, Event, EncodedPayload, FilteredSessionPids}, State
handle_cast(_Other, State) ->
{noreply, State}.
-spec maybe_broadcast_presence(term(), term(), term()) -> ok.
maybe_broadcast_presence(UserId, PresenceMap, NewSnapshot) when
is_integer(UserId), is_map(PresenceMap), is_map(NewSnapshot)
->
try
guild_presence:broadcast_presence_update(UserId, PresenceMap, NewSnapshot)
catch
Class:Reason:Stack ->
logger:warning(
"guild_broadcaster presence_broadcast error: ~p:~p ~p",
[Class, Reason, Stack]
)
end,
maybe_sync_online_status(UserId, NewSnapshot);
maybe_broadcast_presence(_UserId, _PresenceMap, _NewSnapshot) ->
ok.
-spec maybe_sync_online_status(integer(), map()) -> ok.
maybe_sync_online_status(UserId, NewSnapshot) ->
try
guild_presence:sync_online_status(UserId, NewSnapshot)
catch
Class2:Reason2 ->
logger:warning(
"guild_broadcaster engine_online error: ~p:~p",
[Class2, Reason2]
)
end,
ok.
-spec handle_info(term(), map()) ->
{noreply, map()} | {noreply, map(), hibernate} | {stop, normal, map()}.
handle_info({'DOWN', _Ref, process, GuildPid, _Reason}, #{guild_pid := GuildPid} = State) ->
@@ -117,10 +117,8 @@ finalize_resolved(
maybe_start_session_connect_workers(State1);
finalize_resolved(GuildId, SessionId, Attempt, Result0, Computed, Request, SessionPid, State) ->
State1 = upsert_session(SessionId, SessionPid, Request, Computed, State),
UserId = maps:get(user_id, Request, undefined),
State2 = guild_sessions_connect:resection_connected_user(UserId, State, State1),
send_result(GuildId, Attempt, Result0, SessionPid),
maybe_start_session_connect_workers(State2).
maybe_start_session_connect_workers(State1).
-spec discard_pending_session(session_id() | undefined, map()) -> map().
discard_pending_session(SessionId, State) when is_binary(SessionId) ->
@@ -128,7 +126,10 @@ discard_pending_session(SessionId, State) when is_binary(SessionId) ->
case maps:find(SessionId, Sessions0) of
{ok, #{pending_connect := true} = Entry} ->
demonitor_pending_session(Entry),
State#{sessions => maps:remove(SessionId, Sessions0)};
guild_sessions_connect:remove_session_ref(
maps:get(mref, Entry, undefined),
State#{sessions => maps:remove(SessionId, Sessions0)}
);
_ ->
State
end;
@@ -306,18 +307,21 @@ upsert_pending_session(S, U, P, Request, State) ->
Sessions0 = maps:get(sessions, State, #{}),
case maps:find(S, Sessions0) of
error ->
MRef = monitor(process, P),
Entry = #{
session_id => S,
user_id => U,
pid => P,
mref => monitor(process, P),
mref => MRef,
active_guilds => maps:get(active_guilds, Request, sets:new()),
bot => maps:get(bot, Request, false),
is_staff => maps:get(is_staff, Request, false),
pending_connect => true,
viewable_channels => #{}
},
State#{sessions => Sessions0#{S => Entry}};
guild_sessions_connect:put_session_ref(S, MRef, State#{
sessions => Sessions0#{S => Entry}
});
{ok, Existing} ->
State#{sessions => Sessions0#{S => Existing#{pending_connect => true}}}
end.
@@ -485,9 +489,11 @@ upsert_session_valid(SessionId, SessionPid, UserId, Request, Computed, State) ->
store_passive_state(SessionId, GuildId, Computed),
FinalSD = maybe_mark_synced(GuildId, Computed, SessionData),
Sessions = merge_session(SessionId, FinalSD, Existing1, Sessions0),
State1 = State#{sessions => Sessions},
State1 = reindex_session_ref(SessionId, Existing, MRef, State#{sessions => Sessions}),
State2 = update_connected_tracking(UserId, Existing, State1),
update_presence_subscription(UserId, Existing, State2).
PresenceBefore = guild_member_list_connected:resolve_presence_for_user(State2, UserId),
State3 = update_presence_subscription(UserId, Existing, State2),
guild_sessions_connect:resection_connected_user(UserId, PresenceBefore, State, State3).
-spec build_session_data(session_id(), integer(), pid(), reference(), map(), map()) -> map().
build_session_data(SessionId, UserId, SessionPid, MRef, Request, Computed) ->
@@ -526,6 +532,13 @@ merge_session(SessionId, FinalSD, undefined, Sessions0) ->
merge_session(SessionId, FinalSD, Existing, Sessions0) ->
Sessions0#{SessionId => maps:merge(Existing, FinalSD)}.
-spec reindex_session_ref(session_id(), map() | undefined, reference(), map()) -> map().
reindex_session_ref(SessionId, #{mref := OldRef}, MRef, State) when OldRef =/= MRef ->
State1 = guild_sessions_connect:remove_session_ref(OldRef, State),
guild_sessions_connect:put_session_ref(SessionId, MRef, State1);
reindex_session_ref(SessionId, _Existing, MRef, State) ->
guild_sessions_connect:put_session_ref(SessionId, MRef, State).
-spec resolve_monitor(map() | undefined, pid()) -> {reference(), map() | undefined}.
resolve_monitor(undefined, SessionPid) ->
{monitor(process, SessionPid), undefined};
+8 -5
View File
@@ -303,8 +303,8 @@ get_guild_data_for_user(UserId, Data, State) ->
undefined ->
{reply, #{guild_data => null, error_reason => <<"forbidden">>}, State};
Member ->
GuildData = build_member_guild_data(UserId, Member, Data, State),
{reply, #{guild_data => GuildData}, State}
{GuildData, NewState} = build_member_guild_data(UserId, Member, Data, State),
{reply, #{guild_data => GuildData}, NewState}
end.
-spec build_complete_guild_data(map(), guild_state()) -> map().
@@ -313,14 +313,17 @@ build_complete_guild_data(Data, State) ->
Channels = map_utils:ensure_list(maps:get(<<"channels">>, Data, [])),
maps:merge(GuildProperties, build_guild_collection_data(Data, Channels, State)).
-spec build_member_guild_data(user_id(), map(), map(), guild_state()) -> map().
-spec build_member_guild_data(user_id(), map(), map(), guild_state()) -> {map(), guild_state()}.
build_member_guild_data(UserId, Member, Data, State) ->
GuildProperties = maps:get(<<"guild">>, Data, #{}),
AllChannels = guild_data_channels:channels_from_data(Data),
{ViewableChannels, _JoinedAt} = guild_data_channels:derive_member_view(
{ViewableChannels, NewState} = guild_data_channels:member_view(
UserId, Member, State, AllChannels
),
maps:merge(GuildProperties, build_guild_collection_data(Data, ViewableChannels, State)).
GuildData = maps:merge(
GuildProperties, build_guild_collection_data(Data, ViewableChannels, State)
),
{GuildData, NewState}.
-spec build_guild_collection_data(map(), [map()], guild_state()) -> map().
build_guild_collection_data(Data, Channels, State) ->
@@ -9,6 +9,7 @@
find_everyone_viewable_text_channel/2,
sort_channels_for_ordering/1,
derive_member_view/4,
member_view/4,
sanitize_voice_state/1,
voice_members_from_states/2,
merge_members/2
@@ -21,6 +22,14 @@
-type user_id() :: integer().
-type guild_id() :: integer().
-type member_view_cache() :: #{
inputs := tuple(),
exceptions := sets:set(user_id()),
views := #{term() => #{integer() => true}}
}.
-define(MEMBER_VIEW_CACHE_LIMIT, 512).
-export_type([guild_state/0, guild_data_map/0, channel_list/0, guild_member/0, user_id/0]).
-spec channels_from_state(guild_state()) -> channel_list().
@@ -66,11 +75,83 @@ derive_member_view(_UserId, undefined, _State, _Channels) ->
{[], null};
derive_member_view(UserId, Member, State, Channels) ->
Filtered = filter_viewable_channels(UserId, Member, State, Channels),
JoinedAt = maps:get(<<"joined_at">>, Member, null),
{with_missing_parents(Filtered, Channels), JoinedAt}.
-spec member_view(user_id(), guild_member(), guild_state(), channel_list()) ->
{channel_list(), guild_state()}.
member_view(UserId, Member, State, Channels) ->
#{exceptions := Exceptions, views := Views} = Cache = member_view_cache(State),
RolesKey = maps:get(<<"roles">>, Member, []),
case {sets:is_element(UserId, Exceptions), maps:find(RolesKey, Views)} of
{true, _} ->
{Viewable, _JoinedAt} = derive_member_view(UserId, Member, State, Channels),
{Viewable, State#{member_view_cache => Cache}};
{false, {ok, ViewableIds}} ->
Filtered = [C || C <- Channels, maps:is_key(channel_id(C), ViewableIds)],
{with_missing_parents(Filtered, Channels), State#{member_view_cache => Cache}};
{false, error} ->
Filtered = filter_viewable_channels(UserId, Member, State, Channels),
ViewableIds = maps:from_keys(channel_ids(Filtered), true),
NewCache = Cache#{views := put_member_view(RolesKey, ViewableIds, Views)},
{with_missing_parents(Filtered, Channels), State#{member_view_cache => NewCache}}
end.
-spec member_view_cache(guild_state()) -> member_view_cache().
member_view_cache(State) ->
Inputs = member_view_inputs(State),
case maps:get(member_view_cache, State, undefined) of
#{inputs := CachedInputs} = Cache when CachedInputs =:= Inputs ->
Cache;
_ ->
#{
inputs => Inputs,
exceptions => guild_maintenance:viewable_exceptions(State),
views => #{}
}
end.
-spec member_view_inputs(guild_state()) -> tuple().
member_view_inputs(State) ->
Data = guild_data_index:ensure_data_map(State),
{
maps:get(id, State, undefined),
maps:get(<<"id">>, State, undefined),
maps:get(virtual_channel_access, State, undefined),
maps:get(<<"guild">>, Data, undefined),
maps:get(<<"roles">>, Data, undefined),
maps:get(<<"role_index">>, Data, undefined),
maps:get(role_perms_cache, Data, undefined),
maps:get(overwrite_perms_cache, Data, undefined),
maps:get(<<"channels">>, Data, undefined),
maps:map(fun channel_view_inputs/2, guild_data_index:channel_index(Data))
}.
-spec channel_view_inputs(integer(), term()) -> term().
channel_view_inputs(_ChannelId, Channel) when is_map(Channel) ->
{
maps:get(<<"id">>, Channel, undefined),
maps:get(<<"type">>, Channel, undefined),
maps:get(<<"parent_id">>, Channel, undefined),
maps:get(<<"permission_overwrites">>, Channel, undefined)
};
channel_view_inputs(_ChannelId, Channel) ->
Channel.
-spec put_member_view(term(), #{integer() => true}, #{term() => #{integer() => true}}) ->
#{term() => #{integer() => true}}.
put_member_view(RolesKey, ViewableIds, Views) when
map_size(Views) >= ?MEMBER_VIEW_CACHE_LIMIT
->
#{RolesKey => ViewableIds};
put_member_view(RolesKey, ViewableIds, Views) ->
Views#{RolesKey => ViewableIds}.
-spec with_missing_parents(channel_list(), channel_list()) -> channel_list().
with_missing_parents(Filtered, Channels) ->
FilteredIds = sets:from_list(channel_ids(Filtered)),
MissingParentIds = find_missing_parent_ids(Filtered, FilteredIds),
ExtraCategories = collect_extra_categories(MissingParentIds, Channels),
JoinedAt = maps:get(<<"joined_at">>, Member, null),
{Filtered ++ ExtraCategories, JoinedAt}.
Filtered ++ collect_extra_categories(MissingParentIds, Channels).
-spec sanitize_voice_state(map()) -> map().
sanitize_voice_state(VS) ->
+3
View File
@@ -39,6 +39,9 @@ init_base_state(GuildState) ->
member_presence => ets:new(member_presence, [set, public]),
connected_user_ids => sets:new(),
user_session_counts => #{},
guild_session_refs => guild_sessions_connect:build_session_ref_index(
maps:get(sessions, TransferSafe, #{})
),
viewable_channels_cache => ets:new(viewable_channels_cache, [set, public])
},
guild_handoff:restore_transferred_session_state(BaseState).
@@ -13,9 +13,11 @@
maybe_prune_invalid_member_subscriptions/2,
cleanup_removed_member_subscriptions/3,
apply_everyone_perm_bit/2,
viewable_exceptions/1
viewable_exceptions/1,
new_viewable_memo/1,
memoised_member_viewable_channel_map/3
]).
-export_type([guild_state/0]).
-export_type([guild_state/0, viewable_memo/0]).
-define(COUNT_CACHE_REFRESH_INTERVAL, 30000).
@@ -8,7 +8,8 @@
unsubscribe_session/2,
send_member_list_update_to_sessions/5,
dispatch_sync_to_subscribed_list/7,
dispatch_sync_to_subscribed_sessions/6
dispatch_sync_to_subscribed_sessions/6,
handle_sync_item_cache_timeout/1
]).
-type guild_state() :: map().
@@ -18,6 +19,11 @@
-export_type([guild_state/0, list_id/0, range/0, channel_id/0]).
-define(SYNC_ITEM_CACHE_KEY, guild_member_list_sync_item_cache).
-define(SYNC_ITEM_CACHE_MAX_ENTRIES, 4096).
-define(SYNC_ITEM_CACHE_GENERATION_MS, 30000).
-define(SYNC_ITEM_CACHE_TIMEOUT_MSG, member_list_sync_item_cache_rotate).
-spec subscribe_ranges(binary(), list_id(), [range()], guild_state()) ->
{guild_state(), boolean(), [range()]}.
subscribe_ranges(SessionId, ListId, Ranges, State) ->
@@ -241,7 +247,7 @@ dispatch_sync_group(_Ranges, [], _GuildId, _SyncFun) ->
ok;
dispatch_sync_group(Ranges, Pids, GuildId, SyncFun) ->
SyncResponse = SyncFun(Ranges),
Encoded = encode_wire_payload(SyncResponse),
Encoded = encode_sync_payload(SyncResponse),
gateway_dispatch_relay:dispatch_many(Pids, guild_member_list_update, Encoded, GuildId).
-spec encode_wire_payload(map()) -> {pre_encoded, binary()}.
@@ -249,6 +255,130 @@ encode_wire_payload(Payload) ->
WirePayload = eqwalizer:dynamic_cast(guild_data_wire:payload(Payload)),
{pre_encoded, iolist_to_binary(json:encode(WirePayload))}.
-spec encode_sync_payload(map()) -> {pre_encoded, binary()}.
encode_sync_payload(#{<<"ops">> := Ops} = Payload) when is_list(Ops) ->
case sync_item_cache_enabled() of
true ->
encode_sync_payload_cached(Payload, Ops);
false ->
ok = erase_sync_item_cache(),
encode_wire_payload(Payload)
end;
encode_sync_payload(Payload) ->
encode_wire_payload(Payload).
-spec encode_sync_payload_cached(map(), list()) -> {pre_encoded, binary()}.
encode_sync_payload_cached(Payload, Ops) ->
{TimerRef, Current, Previous} = sync_item_cache(),
{FragmentOps, {Current1, Previous1}} =
lists:mapfoldl(fun fragment_op/2, {Current, Previous}, Ops),
_ = erlang:put(?SYNC_ITEM_CACHE_KEY, {TimerRef, Current1, Previous1}),
WirePayload = eqwalizer:dynamic_cast(
guild_data_wire:payload(Payload#{<<"ops">> => FragmentOps})
),
{pre_encoded, iolist_to_binary(json:encode(WirePayload, fun encode_fragment_value/2))}.
-spec fragment_op(term(), {map(), map()}) -> {term(), {map(), map()}}.
fragment_op(#{<<"items">> := Items} = Op, Generations) when is_list(Items) ->
{FragmentItems, Generations1} = lists:mapfoldl(fun fragment_item/2, Generations, Items),
{Op#{<<"items">> => FragmentItems}, Generations1};
fragment_op(Op, Generations) ->
{Op, Generations}.
-spec fragment_item(term(), {map(), map()}) -> {term(), {map(), map()}}.
fragment_item(
#{<<"member">> := #{<<"user">> := #{<<"id">> := Id}}} = Item,
{Current, Previous} = Generations
) ->
case Current of
#{Id := {Item, Fragment}} ->
{Fragment, Generations};
_ ->
Fragment = previous_or_encoded_fragment(Id, Item, Previous),
{Fragment, cache_fragment(Id, {Item, Fragment}, Generations)}
end;
fragment_item(Item, Generations) ->
{Item, Generations}.
-spec previous_or_encoded_fragment(term(), term(), map()) -> {json_fragment, binary()}.
previous_or_encoded_fragment(Id, Item, Previous) ->
case Previous of
#{Id := {Item, Fragment}} ->
Fragment;
_ ->
{json_fragment, iolist_to_binary(json:encode(guild_data_wire:payload(Item)))}
end.
-spec cache_fragment(term(), {term(), {json_fragment, binary()}}, {map(), map()}) ->
{map(), map()}.
cache_fragment(Id, Entry, {Current, _Previous}) when
map_size(Current) >= ?SYNC_ITEM_CACHE_MAX_ENTRIES
->
{#{Id => Entry}, Current};
cache_fragment(Id, Entry, {Current, Previous}) ->
{Current#{Id => Entry}, Previous}.
-spec encode_fragment_value(dynamic(), json:encoder()) -> iodata().
encode_fragment_value({json_fragment, Encoded}, _Encode) ->
Encoded;
encode_fragment_value(Value, Encode) ->
json:encode_value(Value, Encode).
-spec sync_item_cache() -> {reference(), map(), map()}.
sync_item_cache() ->
case erlang:get(?SYNC_ITEM_CACHE_KEY) of
{TimerRef, Current, Previous} = Cache when
is_reference(TimerRef), is_map(Current), is_map(Previous)
->
Cache;
_ ->
{start_sync_item_cache_timer(), #{}, #{}}
end.
-spec handle_sync_item_cache_timeout(reference()) -> ok.
handle_sync_item_cache_timeout(TimerRef) ->
case erlang:get(?SYNC_ITEM_CACHE_KEY) of
{TimerRef, Current, _Previous} when map_size(Current) =:= 0 ->
ok = erase_sync_item_cache();
{TimerRef, Current, _Previous} when is_map(Current) ->
_ = erlang:put(
?SYNC_ITEM_CACHE_KEY, {start_sync_item_cache_timer(), #{}, Current}
),
ok;
_ ->
ok
end.
-spec erase_sync_item_cache() -> ok.
erase_sync_item_cache() ->
case erlang:erase(?SYNC_ITEM_CACHE_KEY) of
{TimerRef, _, _} when is_reference(TimerRef) ->
_ = erlang:cancel_timer(TimerRef, [{async, true}, {info, false}]),
ok;
_ ->
ok
end.
-spec start_sync_item_cache_timer() -> reference().
start_sync_item_cache_timer() ->
erlang:start_timer(
sync_item_cache_generation_ms(), self(), ?SYNC_ITEM_CACHE_TIMEOUT_MSG
).
-spec sync_item_cache_generation_ms() -> pos_integer().
sync_item_cache_generation_ms() ->
case application:get_env(fluxer_gateway, member_list_sync_item_cache_generation_ms) of
{ok, Ms} when is_integer(Ms), Ms > 0 -> Ms;
_ -> ?SYNC_ITEM_CACHE_GENERATION_MS
end.
-spec sync_item_cache_enabled() -> boolean().
sync_item_cache_enabled() ->
case application:get_env(fluxer_gateway, member_list_sync_item_cache_enabled, true) of
false -> false;
_ -> true
end.
-spec session_can_view_list_members(map(), channel_id() | undefined, guild_state()) ->
boolean().
session_can_view_list_members(SessionData, undefined, State) ->
@@ -299,3 +429,320 @@ list_channel_id(ListId) when is_binary(ListId) ->
end;
list_channel_id(_) ->
undefined.
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
sync_payload_matches_uncached_encoding_test() ->
with_clean_cache(fun() ->
Members = [test_member(N) || N <- lists:seq(1, 60)],
Payloads = [
sync_payload(<<"500">>, [{0, 99}], Members),
sync_payload(<<"600">>, [{0, 99}], lists:reverse(Members)),
sync_payload(<<"600">>, [{0, 99}, {100, 199}], Members ++ Members),
sync_payload(<<"0">>, [{0, 99}], lists:sublist(Members, 7, 30))
],
[assert_matches_uncached(P) || P <- Payloads ++ Payloads],
?assertEqual(60, map_size(current_entries()))
end).
sync_payload_reencodes_changed_member_test() ->
with_clean_cache(fun() ->
Members = [test_member(N) || N <- lists:seq(1, 10)],
assert_matches_uncached(sync_payload(<<"500">>, [{0, 99}], Members)),
[First | Rest] = Members,
Renamed = put_in(First, [<<"member">>, <<"nick">>], <<"renamed \"x\"">>),
Idle = put_in(
hd(Rest), [<<"member">>, <<"presence">>, <<"status">>], <<"idle">>
),
Changed = [Renamed, Idle | tl(Rest)],
Old = encode_sync_payload(sync_payload(<<"500">>, [{0, 99}], Members)),
New = assert_matches_uncached(sync_payload(<<"500">>, [{0, 99}], Changed)),
?assertNotEqual(Old, New)
end).
sync_payload_keeps_non_member_items_inline_test() ->
with_clean_cache(fun() ->
Items = [
#{<<"group">> => #{<<"id">> => <<"online">>, <<"count">> => 3}},
#{<<"member">> => #{<<"nick">> => <<"no user">>}},
#{<<"member">> => #{<<"user">> => #{<<"username">> => <<"no id">>}}},
test_member(1)
],
assert_matches_uncached(payload_with_ops([sync_op_items({0, 99}, Items)])),
?assertEqual([<<"1427764882469228557">>], maps:keys(current_entries()))
end).
sync_payload_without_ops_matches_uncached_test() ->
with_clean_cache(fun() ->
assert_matches_uncached(#{<<"id">> => <<"500">>, <<"guild_id">> => 7}),
?assertEqual(undefined, erlang:get(?SYNC_ITEM_CACHE_KEY)),
assert_matches_uncached(payload_with_ops([#{<<"op">> => <<"INVALIDATE">>}])),
?assertEqual(#{}, current_entries())
end).
sync_item_cache_is_bounded_test() ->
with_clean_cache(fun() ->
Count = ?SYNC_ITEM_CACHE_MAX_ENTRIES + 5,
Members = [test_member(N) || N <- lists:seq(1, Count)],
Payload = sync_payload(<<"500">>, [{0, Count - 1}], Members),
assert_matches_uncached(Payload),
?assertEqual(5, map_size(current_entries())),
?assertEqual(?SYNC_ITEM_CACHE_MAX_ENTRIES, map_size(previous_entries())),
assert_matches_uncached(Payload),
?assert(
map_size(current_entries()) + map_size(previous_entries()) =<
2 * ?SYNC_ITEM_CACHE_MAX_ENTRIES
)
end).
sync_item_cache_flag_off_encodes_uncached_and_clears_test() ->
with_clean_cache(fun() ->
Payload = sync_payload(<<"500">>, [{0, 99}], [test_member(N) || N <- lists:seq(1, 5)]),
assert_matches_uncached(Payload),
?assertEqual(5, map_size(current_entries())),
with_cache_flag(false, fun() ->
assert_matches_uncached(Payload),
?assertEqual(undefined, erlang:get(?SYNC_ITEM_CACHE_KEY))
end)
end).
sync_item_cache_released_after_two_idle_generations_test() ->
with_clean_cache(fun() ->
with_generation_ms(5, fun() ->
Members = [test_member(N) || N <- lists:seq(1, 5)],
Payload = sync_payload(<<"500">>, [{0, 99}], Members),
assert_matches_uncached(Payload),
?assertNotEqual(undefined, erlang:get(?SYNC_ITEM_CACHE_KEY)),
?assertEqual(2, deliver_rotation_timers_until_released(10)),
?assertEqual(undefined, erlang:get(?SYNC_ITEM_CACHE_KEY))
end)
end).
sync_item_cache_keeps_entries_touched_each_generation_test() ->
with_clean_cache(fun() ->
Payload = sync_payload(<<"500">>, [{0, 99}], [test_member(N) || N <- lists:seq(1, 8)]),
assert_matches_uncached(Payload),
rotate_now(),
?assertEqual(#{}, current_entries()),
?assertEqual(8, map_size(previous_entries())),
assert_matches_uncached(Payload),
?assertEqual(8, map_size(current_entries())),
rotate_now(),
?assertEqual(8, map_size(previous_entries())),
rotate_now(),
?assertEqual(undefined, erlang:get(?SYNC_ITEM_CACHE_KEY)),
assert_matches_uncached(Payload),
?assertEqual(8, map_size(current_entries()))
end).
sync_item_cache_previous_generation_changed_member_reencodes_test() ->
with_clean_cache(fun() ->
Members = [test_member(N) || N <- lists:seq(1, 4)],
assert_matches_uncached(sync_payload(<<"500">>, [{0, 99}], Members)),
rotate_now(),
[First | Rest] = Members,
Renamed = put_in(First, [<<"member">>, <<"nick">>], <<"renamed">>),
assert_matches_uncached(sync_payload(<<"500">>, [{0, 99}], [Renamed | Rest])),
#{<<"member">> := #{<<"user">> := #{<<"id">> := Id}}} = Renamed,
?assertMatch(#{Id := {Renamed, _}}, current_entries())
end).
sync_item_cache_ignores_stale_timer_test() ->
with_clean_cache(fun() ->
Payload = sync_payload(<<"500">>, [{0, 99}], [test_member(N) || N <- lists:seq(1, 5)]),
assert_matches_uncached(Payload),
{StaleRef, _, _} = erlang:get(?SYNC_ITEM_CACHE_KEY),
with_cache_flag(false, fun() -> assert_matches_uncached(Payload) end),
assert_matches_uncached(Payload),
?assertMatch({noreply, #{}}, guild:handle_info(rotation_msg(StaleRef), #{})),
?assertEqual(5, map_size(current_entries())),
?assertNotMatch({StaleRef, _, _}, erlang:get(?SYNC_ITEM_CACHE_KEY))
end).
dispatch_sync_group_sends_uncached_bytes_test() ->
with_clean_cache(fun() ->
Members = [test_member(N) || N <- lists:seq(1, 20)],
SyncFun = fun(Ranges) -> sync_payload(<<"500">>, Ranges, Members) end,
Expected = uncached(SyncFun([{0, 99}])),
ok = dispatch_sync_group([{0, 99}], [self()], 7, SyncFun),
ok = dispatch_sync_group([{0, 99}], [self()], 7, SyncFun),
?assertEqual([Expected, Expected], received_dispatches())
end).
assert_matches_uncached(Payload) ->
Expected = uncached(Payload),
Actual = encode_sync_payload(Payload),
?assertEqual(Expected, Actual),
?assertEqual(
json:decode(element(2, Expected)), json:decode(element(2, Actual))
),
Actual.
uncached(Payload) ->
{pre_encoded, iolist_to_binary(json:encode(guild_data_wire:payload(Payload)))}.
received_dispatches() ->
receive
{'$gen_cast', {dispatch, guild_member_list_update, Encoded}} ->
[Encoded | received_dispatches()];
{dispatch, guild_member_list_update, Encoded} ->
[Encoded | received_dispatches()]
after 200 ->
[]
end.
sync_payload(ListId, Ranges, Items) ->
Ops = [
sync_op_items(Range, lists:sublist(Items, Start + 1, End - Start + 1))
|| {Start, End} = Range <- Ranges
],
(payload_with_ops(Ops))#{<<"id">> => ListId}.
payload_with_ops(Ops) ->
#{
<<"guild_id">> => <<"1427764882469228556">>,
<<"id">> => <<"500">>,
<<"channel_id">> => <<"500">>,
<<"member_count">> => 55278,
<<"online_count">> => 1365,
<<"groups">> => [
#{<<"id">> => <<"1427764882469228600">>, <<"count">> => 12},
#{<<"id">> => <<"online">>, <<"count">> => 1353}
],
<<"ops">> => Ops
}.
sync_op_items({Start, End}, Items) ->
#{<<"op">> => <<"SYNC">>, <<"range">> => [Start, End], <<"items">> => Items}.
test_member(N) ->
Id = integer_to_binary(1427764882469228556 + N),
User = #{
<<"id">> => Id,
<<"username">> => <<"user_", (integer_to_binary(N))/binary>>,
<<"global_name">> => pick(N, [
null, <<"Ünïcødé \\ \"q\" "/utf8, 240, 159, 152, 128>>, <<"g">>
]),
<<"avatar">> => pick(N, [null, <<"a_0123456789abcdef">>]),
<<"discriminator">> => <<"0000">>,
<<"flags">> => N * 4096,
<<"bot">> => N rem 11 =:= 0,
<<"avatar_decoration_data">> => pick(N, [
null, #{<<"sku_id">> => 99, <<"asset">> => <<"x">>}
])
},
Member = #{
<<"user">> => User,
<<"nick">> => pick(N, [null, <<"nick\n\t", 1, " ", (integer_to_binary(N))/binary>>]),
<<"roles">> => pick(N, [
[], [<<"1427764882469228600">>, <<"1427764882469228601">>], [42]
]),
<<"joined_at">> => <<"2026-01-01T00:00:00.000Z">>,
<<"communication_disabled_until">> => null,
<<"deaf">> => false,
<<"mute">> => N rem 5 =:= 0,
<<"premium_since">> => pick(N, [null, <<"2026-02-02T00:00:00Z">>]),
<<"guild_id">> => 1427764882469228556,
voice_channel_id => pick(N, [null, 1427764882469228777]),
<<"presence">> => #{
<<"user">> => #{<<"id">> => Id},
<<"status">> => pick(N, [<<"online">>, <<"idle">>, <<"dnd">>]),
<<"mobile">> => N rem 2 =:= 0,
<<"afk">> => false,
<<"custom_status">> => pick(N, [
null, #{<<"text">> => <<"hi">>, <<"emoji_id">> => 5, <<"emoji_name">> => null}
])
}
},
#{<<"member">> => Member}.
pick(N, Options) ->
lists:nth(N rem length(Options) + 1, Options).
put_in(Map, [Key], Value) ->
Map#{Key => Value};
put_in(Map, [Key | Rest], Value) ->
Map#{Key => put_in(maps:get(Key, Map), Rest, Value)}.
with_clean_cache(Fun) ->
clear_cache(),
try
Fun()
after
clear_cache()
end.
clear_cache() ->
ok = erase_sync_item_cache(),
flush_rotation_timers().
flush_rotation_timers() ->
receive
{timeout, _, ?SYNC_ITEM_CACHE_TIMEOUT_MSG} -> flush_rotation_timers()
after 0 ->
ok
end.
current_entries() ->
case erlang:get(?SYNC_ITEM_CACHE_KEY) of
{_, Current, _} -> Current;
_ -> #{}
end.
previous_entries() ->
case erlang:get(?SYNC_ITEM_CACHE_KEY) of
{_, _, Previous} -> Previous;
_ -> #{}
end.
rotation_msg(TimerRef) ->
{timeout, TimerRef, ?SYNC_ITEM_CACHE_TIMEOUT_MSG}.
rotate_now() ->
{TimerRef, _, _} = erlang:get(?SYNC_ITEM_CACHE_KEY),
_ = erlang:cancel_timer(TimerRef),
?assertMatch({noreply, #{}}, guild:handle_info(rotation_msg(TimerRef), #{})).
deliver_rotation_timers_until_released(0) ->
erlang:error(sync_item_cache_not_released);
deliver_rotation_timers_until_released(Remaining) ->
receive
{timeout, _, ?SYNC_ITEM_CACHE_TIMEOUT_MSG} = Msg ->
{noreply, _} = guild:handle_info(Msg, #{}),
case erlang:get(?SYNC_ITEM_CACHE_KEY) of
undefined -> 1;
_ -> 1 + deliver_rotation_timers_until_released(Remaining - 1)
end
after 1000 ->
erlang:error(sync_item_cache_timer_not_armed)
end.
with_generation_ms(Ms, Fun) ->
Key = member_list_sync_item_cache_generation_ms,
Previous = application:get_env(fluxer_gateway, Key),
application:set_env(fluxer_gateway, Key, Ms),
try
Fun()
after
case Previous of
{ok, Old} -> application:set_env(fluxer_gateway, Key, Old);
undefined -> application:unset_env(fluxer_gateway, Key)
end
end.
with_cache_flag(Value, Fun) ->
Previous = application:get_env(fluxer_gateway, member_list_sync_item_cache_enabled),
application:set_env(fluxer_gateway, member_list_sync_item_cache_enabled, Value),
try
Fun()
after
case Previous of
{ok, Old} ->
application:set_env(fluxer_gateway, member_list_sync_item_cache_enabled, Old);
undefined ->
application:unset_env(fluxer_gateway, member_list_sync_item_cache_enabled)
end
end.
-endif.
@@ -54,13 +54,17 @@ take_pending_list_ids(ListIds, State) ->
#{pending_list_ids := PendingListIds} = Batch when is_map(PendingListIds) ->
cancel_timer(Batch),
{
normalize_list_ids(ListIds ++ maps:keys(PendingListIds)),
normalize_list_ids(ListIds ++ unsynced_list_ids(PendingListIds)),
maps:remove(?SYNC_BATCH_STATE_KEY, State)
};
_ ->
{normalize_list_ids(ListIds), State}
end.
-spec unsynced_list_ids(map()) -> [term()].
unsynced_list_ids(PendingListIds) ->
maps:keys(maps:filter(fun(_ListId, Mark) -> Mark =/= synced end, PendingListIds)).
-spec normalize_list_ids([term()]) -> [list_id()].
normalize_list_ids(ListIds) ->
lists:usort([ListId || ListId <- ListIds, is_binary(ListId)]).
@@ -169,4 +173,23 @@ queue_list_sync_drains_preexisting_pending_batch_test() ->
},
?assertEqual(#{}, queue_list_sync(<<"500">>, State)).
flush_pending_syncs_skips_synced_lists_test() ->
Ref = erlang:send_after(60000, self(), ?FLUSH_SYNC_BATCH_MSG),
State = #{
?SYNC_BATCH_STATE_KEY => #{
timer_ref => Ref,
pending_list_ids => #{<<"500">> => synced, <<"600">> => true}
}
},
?assertEqual({[<<"600">>], #{}}, take_pending_list_ids([], State)),
?assertEqual(false, erlang:read_timer(Ref)).
queue_list_sync_keeps_explicit_list_marked_synced_test() ->
State = #{
?SYNC_BATCH_STATE_KEY => #{
pending_list_ids => #{<<"500">> => synced, <<"600">> => synced}
}
},
?assertEqual({[<<"500">>], #{}}, take_pending_list_ids([<<"500">>], State)).
-endif.
@@ -9,6 +9,8 @@
broadcast_all_member_list_updates/1,
broadcast_member_list_updates_for_channel/2,
broadcast_channel_engine_connection_change/2,
queue_synced_connection_change/2,
presence_change_resyncs_lists/2,
flush_pending_member_list_syncs/1,
resync_hoisted_member_lists/1,
rebuild_channels_for_permission_change/2,
@@ -22,6 +24,7 @@
-type channel_id() :: integer().
-type engine_ref() :: ets:table().
-type absence() :: {absent, engine_ref()} | present.
-type pending_mark() :: true | synced.
-define(MAX_MEMBER_LIST_SYNC_SKIPPED_ABSENT, 1000000000).
@@ -70,7 +73,7 @@ broadcast_member_list_updates(UserId, OldState, UpdatedState, OldPresence, NewPr
dispatch_presence_delta(UserId, OldMember, NewMember, OldPresence, NewPresence, State) ->
case presence_delta_is_inert(OldPresence, NewPresence, OldMember, NewMember) of
true ->
State;
invalidate_synced_lists(UserId, State);
false ->
SubsTab = maps:get(member_list_subscriptions, State),
dispatch_user_change_to_subscribed_lists(
@@ -87,15 +90,20 @@ dispatch_presence_delta(UserId, OldMember, NewMember, OldPresence, NewPresence,
presence_delta_is_inert(OldPresence, NewPresence, Member, Member) when
is_map(OldPresence), is_map(NewPresence), is_map(Member)
->
member_list_presence_fields(OldPresence) =:= member_list_presence_fields(NewPresence);
not presence_change_resyncs_lists(OldPresence, NewPresence);
presence_delta_is_inert(_OldPresence, _NewPresence, _OldMember, _NewMember) ->
false.
-spec member_list_presence_fields(map()) -> {binary(), term()}.
-spec presence_change_resyncs_lists(map(), map()) -> boolean().
presence_change_resyncs_lists(OldPresence, NewPresence) ->
member_list_presence_fields(OldPresence) =/= member_list_presence_fields(NewPresence).
-spec member_list_presence_fields(map()) -> {binary(), term(), term()}.
member_list_presence_fields(Presence) ->
{
maps:get(<<"status">>, Presence, <<"offline">>),
maps:get(<<"custom_status">>, Presence, null)
maps:get(<<"custom_status">>, Presence, null),
maps:get(<<"mobile">>, Presence, false)
}.
-spec find_member_in_state_data(user_id(), guild_state()) -> map() | undefined.
@@ -165,29 +173,38 @@ broadcast_channel_with_guild_id(GuildId, ChannelId, State) ->
-spec broadcast_channel_engine_connection_change(user_id(), guild_state()) -> guild_state().
broadcast_channel_engine_connection_change(UserId, State) ->
queue_connection_change(UserId, true, State).
-spec queue_synced_connection_change(user_id(), guild_state()) -> guild_state().
queue_synced_connection_change(UserId, State) ->
queue_connection_change(UserId, synced, State).
-spec queue_connection_change(user_id(), pending_mark(), guild_state()) -> guild_state().
queue_connection_change(UserId, Mark, State) ->
case maps:get(member_list_subscriptions, State, undefined) of
undefined ->
State;
SubsTab ->
broadcast_channel_engine_connection_change(UserId, State, SubsTab)
queue_connection_change(UserId, Mark, State, SubsTab)
end.
-spec broadcast_channel_engine_connection_change(user_id(), guild_state(), ets:table()) ->
-spec queue_connection_change(user_id(), pending_mark(), guild_state(), ets:table()) ->
guild_state().
broadcast_channel_engine_connection_change(UserId, State, SubsTab) ->
queue_connection_change(UserId, Mark, State, SubsTab) ->
{ok, NewState} = guild_member_list_write_context:with_guild_id(State, fun(GuildId) ->
{ok, fold_connection_change_lists(GuildId, UserId, State, SubsTab)}
{ok, fold_connection_change_lists(GuildId, UserId, Mark, State, SubsTab)}
end),
NewState.
-spec fold_connection_change_lists(integer(), user_id(), guild_state(), ets:table()) ->
guild_state().
fold_connection_change_lists(GuildId, UserId, State, SubsTab) ->
-spec fold_connection_change_lists(
integer(), user_id(), pending_mark(), guild_state(), ets:table()
) -> guild_state().
fold_connection_change_lists(GuildId, UserId, Mark, State, SubsTab) ->
Sessions = maps:get(sessions, State, #{}),
lists:foldl(
fun(ListId, AccState) ->
sync_connection_change_for_subscribed_list(
GuildId, UserId, ListId, Sessions, AccState
GuildId, UserId, ListId, Mark, Sessions, AccState
)
end,
State,
@@ -263,12 +280,12 @@ apply_channel_member_change(_UserId, _ListId, _OldMember, _NewMember, _State) ->
ok.
-spec sync_connection_change_for_subscribed_list(
integer(), user_id(), list_id(), map(), guild_state()
integer(), user_id(), list_id(), pending_mark(), map(), guild_state()
) -> guild_state().
sync_connection_change_for_subscribed_list(_GuildId, UserId, ListId, _Sessions, State) ->
sync_connection_change_for_subscribed_list(_GuildId, UserId, ListId, Mark, _Sessions, State) ->
case guild_member_list_channel_engine:is_engine_list(ListId, State) of
true ->
queue_connection_list_sync_unless_absent(UserId, ListId, State);
queue_connection_list_sync_unless_absent(UserId, ListId, Mark, State);
false ->
State
end.
@@ -301,21 +318,27 @@ rebuild_channel_store(ListId, State) ->
false -> State
end.
-spec queue_connection_list_sync_unless_absent(user_id(), list_id(), guild_state()) ->
guild_state().
queue_connection_list_sync_unless_absent(UserId, ListId, State) ->
-spec queue_connection_list_sync_unless_absent(
user_id(), list_id(), pending_mark(), guild_state()
) -> guild_state().
queue_connection_list_sync_unless_absent(UserId, ListId, Mark, State) ->
case list_member_absence(UserId, ListId, State) of
{absent, _Ref} -> record_member_list_sync_skipped_absent(State);
present -> queue_connection_list_sync(ListId, State)
present -> queue_connection_list_sync(ListId, Mark, State)
end.
-spec queue_connection_list_sync(list_id(), guild_state()) -> guild_state().
queue_connection_list_sync(ListId, State) ->
-spec queue_connection_list_sync(list_id(), pending_mark(), guild_state()) -> guild_state().
queue_connection_list_sync(ListId, Mark, State) ->
case maps:get(pending_member_list_sync_batch, State, undefined) of
#{pending_list_ids := PendingListIds} = Batch when is_map(PendingListIds) ->
State#{
pending_member_list_sync_batch => Batch#{
pending_list_ids => PendingListIds#{ListId => true}
pending_list_ids => maps:update_with(
ListId,
fun(Old) -> merge_pending_mark(Old, Mark) end,
Mark,
PendingListIds
)
}
};
_ ->
@@ -324,12 +347,40 @@ queue_connection_list_sync(ListId, State) ->
),
State#{
pending_member_list_sync_batch => #{
pending_list_ids => #{ListId => true},
pending_list_ids => #{ListId => Mark},
timer_ref => TimerRef
}
}
end.
-spec merge_pending_mark(term(), pending_mark()) -> pending_mark().
merge_pending_mark(true, _Mark) ->
true;
merge_pending_mark(_Old, Mark) ->
Mark.
-spec invalidate_synced_lists(user_id(), guild_state()) -> guild_state().
invalidate_synced_lists(UserId, State) ->
case maps:get(pending_member_list_sync_batch, State, undefined) of
#{pending_list_ids := PendingListIds} = Batch when is_map(PendingListIds) ->
Pending = maps:map(
fun(ListId, Mark) -> invalidate_synced_mark(UserId, ListId, Mark, State) end,
PendingListIds
),
State#{pending_member_list_sync_batch => Batch#{pending_list_ids => Pending}};
_ ->
State
end.
-spec invalidate_synced_mark(user_id(), list_id(), term(), guild_state()) -> term().
invalidate_synced_mark(UserId, ListId, synced, State) ->
case list_member_absence(UserId, ListId, State) of
present -> true;
{absent, _Ref} -> synced
end;
invalidate_synced_mark(_UserId, _ListId, Mark, _State) ->
Mark.
-spec connection_sync_delay_ms(guild_state()) -> pos_integer().
connection_sync_delay_ms(State) ->
case State of
@@ -530,6 +581,20 @@ list_id_fold_matches_fold_lists_test() ->
guild_member_list_subs:destroy(Tab)
end.
presence_change_resyncs_lists_on_status_custom_status_or_mobile_test() ->
Online = #{
<<"status">> => <<"online">>,
<<"mobile">> => false,
<<"afk">> => false,
<<"custom_status">> => null
},
?assert(presence_change_resyncs_lists(Online, Online#{<<"status">> => <<"idle">>})),
?assert(presence_change_resyncs_lists(Online, Online#{<<"custom_status">> => #{}})),
?assert(presence_change_resyncs_lists(Online, Online#{<<"mobile">> => true})),
?assert(presence_change_resyncs_lists(#{}, Online)),
?assertNot(presence_change_resyncs_lists(Online, Online#{<<"afk">> => true})),
?assertNot(presence_change_resyncs_lists(#{}, #{<<"status">> => <<"offline">>})).
connection_sync_delay_defaults_test() ->
?assertEqual(250, connection_sync_delay_ms(#{})),
?assertEqual(250, connection_sync_delay_ms(#{member_count => 4999})),
@@ -556,19 +621,19 @@ connection_sync_delay_rejects_invalid_env_test() ->
).
connection_list_sync_arms_one_timer_per_window_test() ->
State1 = queue_connection_list_sync(<<"500">>, #{}),
State1 = queue_connection_list_sync(<<"500">>, true, #{}),
Batch1 = maps:get(pending_member_list_sync_batch, State1),
TimerRef = maps:get(timer_ref, Batch1),
try
?assert(is_reference(TimerRef)),
?assertEqual(#{<<"500">> => true}, maps:get(pending_list_ids, Batch1)),
State2 = queue_connection_list_sync(<<"600">>, State1),
State2 = queue_connection_list_sync(<<"600">>, true, State1),
Batch2 = maps:get(pending_member_list_sync_batch, State2),
?assertEqual(TimerRef, maps:get(timer_ref, Batch2)),
?assertEqual(
#{<<"500">> => true, <<"600">> => true}, maps:get(pending_list_ids, Batch2)
),
State3 = queue_connection_list_sync(<<"500">>, State2),
State3 = queue_connection_list_sync(<<"500">>, true, State2),
?assertEqual(Batch2, maps:get(pending_member_list_sync_batch, State3))
after
_ = erlang:cancel_timer(TimerRef)
@@ -579,7 +644,7 @@ connection_change_sync_is_debounced_test() ->
try
ok = guild_member_list_engine:add_member(Ref, 7, <<"seven">>, [], true),
Next = sync_connection_change_for_subscribed_list(
1, 7, <<"500">>, #{}, engine_state(Ref)
1, 7, <<"500">>, true, #{}, engine_state(Ref)
),
Batch = maps:get(pending_member_list_sync_batch, Next),
?assertEqual(#{<<"500">> => true}, maps:get(pending_list_ids, Batch)),
@@ -592,7 +657,7 @@ connection_change_skip_absent_arms_no_timer_test() ->
Ref = guild_member_list_engine:new(),
try
Next = sync_connection_change_for_subscribed_list(
1, 7, <<"500">>, #{}, engine_state(Ref)
1, 7, <<"500">>, true, #{}, engine_state(Ref)
),
?assertNot(maps:is_key(pending_member_list_sync_batch, Next)),
?assertEqual(1, maps:get(member_list_sync_skipped_absent, Next))
@@ -606,7 +671,7 @@ connection_change_skip_absent_keeps_pending_batch_test() ->
Batch = #{pending_list_ids => #{<<"600">> => true}, timer_ref => TimerRef},
State = (engine_state(Ref))#{pending_member_list_sync_batch => Batch},
try
Next = sync_connection_change_for_subscribed_list(1, 7, <<"500">>, #{}, State),
Next = sync_connection_change_for_subscribed_list(1, 7, <<"500">>, true, #{}, State),
?assertEqual(Batch, maps:get(pending_member_list_sync_batch, Next))
after
_ = erlang:cancel_timer(TimerRef),
@@ -636,6 +701,83 @@ member_update_sync_stays_immediate_test() ->
guild_member_list_engine:destroy(Ref)
end.
synced_connection_mark_never_downgrades_pending_sync_test() ->
State1 = queue_connection_list_sync(<<"500">>, synced, #{}),
TimerRef = maps:get(timer_ref, maps:get(pending_member_list_sync_batch, State1)),
try
State2 = queue_connection_list_sync(<<"600">>, true, State1),
State3 = queue_connection_list_sync(<<"600">>, synced, State2),
State4 = queue_connection_list_sync(<<"500">>, true, State3),
Batch = maps:get(pending_member_list_sync_batch, State4),
?assertEqual(TimerRef, maps:get(timer_ref, Batch)),
?assertEqual(
#{<<"500">> => true, <<"600">> => true}, maps:get(pending_list_ids, Batch)
)
after
_ = erlang:cancel_timer(TimerRef)
end.
inert_presence_change_invalidates_synced_lists_holding_the_user_test() ->
Ref = guild_member_list_engine:new(),
Other = guild_member_list_engine:new(),
Pending = #{<<"500">> => synced, <<"600">> => synced, <<"700">> => true},
State = #{
channel_member_list_engines => #{<<"500">> => Ref, <<"600">> => Other},
member_presence => #{},
pending_member_list_sync_batch => #{pending_list_ids => Pending}
},
try
ok = guild_member_list_engine:add_member(Ref, 7, <<"seven">>, [], true),
Next = invalidate_synced_lists(7, State),
?assertEqual(
#{<<"500">> => true, <<"600">> => synced, <<"700">> => true},
maps:get(pending_list_ids, maps:get(pending_member_list_sync_batch, Next))
),
?assertEqual(#{}, invalidate_synced_lists(7, #{}))
after
guild_member_list_engine:destroy(Ref),
guild_member_list_engine:destroy(Other)
end.
inert_presence_delta_invalidates_synced_lists_test() ->
Ref = guild_member_list_engine:new(),
Member = #{<<"user">> => #{<<"id">> => <<"7">>}, <<"roles">> => []},
Old = #{
<<"status">> => <<"online">>,
<<"mobile">> => false,
<<"afk">> => false,
<<"custom_status">> => null
},
State = (engine_state(Ref))#{
pending_member_list_sync_batch => #{pending_list_ids => #{<<"500">> => synced}}
},
try
ok = guild_member_list_engine:add_member(Ref, 7, <<"seven">>, [], true),
Next = dispatch_presence_delta(
7, Member, Member, Old, Old#{<<"afk">> => true}, State
),
?assertEqual(
#{<<"500">> => true},
maps:get(pending_list_ids, maps:get(pending_member_list_sync_batch, Next))
)
after
guild_member_list_engine:destroy(Ref)
end.
synced_connection_change_marks_lists_test() ->
Ref = guild_member_list_engine:new(),
try
ok = guild_member_list_engine:add_member(Ref, 7, <<"seven">>, [], true),
Next = sync_connection_change_for_subscribed_list(
1, 7, <<"500">>, synced, #{}, engine_state(Ref)
),
Batch = maps:get(pending_member_list_sync_batch, Next),
?assertEqual(#{<<"500">> => synced}, maps:get(pending_list_ids, Batch)),
_ = erlang:cancel_timer(maps:get(timer_ref, Batch))
after
guild_member_list_engine:destroy(Ref)
end.
engine_state(Ref) ->
#{channel_member_list_engines => #{<<"500">> => Ref}, member_presence => #{}}.
+383 -38
View File
@@ -15,6 +15,7 @@
-type guild_state() :: map().
-type viewable_index() :: #{user_id() => map()}.
-type index_ctx() :: {map(), viewable_index(), guild_state()}.
-type memo() :: guild_maintenance:viewable_memo().
-spec compute_count(user_id() | term(), guild_state()) -> non_neg_integer().
compute_count(UserId, State) when is_integer(UserId), UserId > 0 ->
@@ -52,49 +53,61 @@ self_online_count(UserId, State) ->
-spec count_mutually_visible_indexed(user_id(), sets:set(), guild_state()) ->
non_neg_integer().
count_mutually_visible_indexed(UserId, ViewerSet, State) ->
{Count, _Memo} = count_and_memo(UserId, ViewerSet, State),
Count.
-spec count_and_memo(user_id(), sets:set(), guild_state()) ->
{non_neg_integer(), memo() | undefined}.
count_and_memo(UserId, ViewerSet, State) ->
Tab = maps:get(member_presence, State),
Ctx = {viewer_channel_map(ViewerSet), build_viewable_index(State), State},
ets:foldl(
fun({OtherUserId, Presence}, Acc) ->
sets:fold(
fun(OtherUserId, Acc) ->
Presence = guild_state_member:lookup_presence(Tab, OtherUserId),
count_online_member_indexed(UserId, Ctx, OtherUserId, Presence, Acc)
end,
0,
Tab
{0, undefined},
guild_member_list_connected:connected_session_user_ids(State)
).
-spec count_online_member_indexed(
user_id(), index_ctx(), term(), term(), non_neg_integer()
) -> non_neg_integer().
count_online_member_indexed(UserId, Ctx, OtherUserId, Presence, Acc) when
user_id(), index_ctx(), term(), term(), {non_neg_integer(), memo() | undefined}
) -> {non_neg_integer(), memo() | undefined}.
count_online_member_indexed(UserId, Ctx, OtherUserId, Presence, {Count, Memo} = Acc) when
is_integer(OtherUserId), is_map(Presence), OtherUserId > 0
->
case is_online(Presence) of
false -> Acc;
true when OtherUserId =:= UserId -> Acc + 1;
true -> count_if_mutually_visible_indexed(OtherUserId, Ctx, Acc)
false ->
Acc;
true when OtherUserId =:= UserId ->
{Count + 1, Memo};
true ->
case shares_viewable_channel(OtherUserId, Ctx, Memo) of
{true, Memo1} -> {Count + 1, Memo1};
{false, Memo1} -> {Count, Memo1}
end
end;
count_online_member_indexed(_UserId, _Ctx, _OtherUserId, _Presence, Acc) ->
Acc.
-spec count_if_mutually_visible_indexed(user_id(), index_ctx(), non_neg_integer()) ->
non_neg_integer().
count_if_mutually_visible_indexed(OtherUserId, Ctx, Acc) ->
case shares_viewable_channel(OtherUserId, Ctx) of
true -> Acc + 1;
false -> Acc
end.
-spec shares_viewable_channel(user_id(), index_ctx()) -> boolean().
shares_viewable_channel(OtherUserId, {ViewerMap, Index, State}) ->
-spec shares_viewable_channel(user_id(), index_ctx(), memo() | undefined) ->
{boolean(), memo() | undefined}.
shares_viewable_channel(OtherUserId, {ViewerMap, Index, State}, Memo) ->
case maps:find(OtherUserId, Index) of
{ok, OtherMap} -> maps_share_any_key(OtherMap, ViewerMap);
error -> channel_list_shares_any(OtherUserId, ViewerMap, State)
{ok, OtherMap} ->
{maps_share_any_key(OtherMap, ViewerMap), Memo};
error ->
{OtherMap, Memo1} = guild_maintenance:memoised_member_viewable_channel_map(
OtherUserId, State, ensure_memo(Memo, State)
),
{maps_share_any_key(OtherMap, ViewerMap), Memo1}
end.
-spec channel_list_shares_any(user_id(), map(), guild_state()) -> boolean().
channel_list_shares_any(OtherUserId, ViewerMap, State) ->
Channels = guild_visibility:get_user_viewable_channels(OtherUserId, State),
lists:any(fun(ChannelId) -> maps:is_key(ChannelId, ViewerMap) end, Channels).
-spec ensure_memo(memo() | undefined, guild_state()) -> memo().
ensure_memo(undefined, State) ->
guild_maintenance:new_viewable_memo(State);
ensure_memo(Memo, _State) ->
Memo.
-spec viewer_channel_map(sets:set()) -> map().
viewer_channel_map(ViewerSet) ->
@@ -167,18 +180,14 @@ key_matches_or_continue(Key, NextIterator, LargerMap) ->
-spec is_self_online(user_id(), guild_state()) -> boolean().
is_self_online(UserId, State) ->
Tab = maps:get(member_presence, State),
case ets:lookup(Tab, UserId) of
[{_, P}] -> is_online(P);
[] -> false
end.
Connected = guild_member_list_connected:connected_session_user_ids(State),
sets:is_element(UserId, Connected) andalso
is_online(guild_state_member:lookup_presence(maps:get(member_presence, State), UserId)).
-spec is_online(term()) -> boolean().
is_online(Presence) when is_map(Presence) ->
-spec is_online(map()) -> boolean().
is_online(Presence) ->
Status = maps:get(<<"status">>, Presence, <<"offline">>),
Status =/= <<"offline">> andalso Status =/= <<"invisible">>;
is_online(_) ->
false.
Status =/= <<"offline">> andalso Status =/= <<"invisible">>.
-ifdef(TEST).
@@ -242,6 +251,7 @@ mutual_visibility_state(GuildId, BotRoleId) ->
<<"channels">> => mutual_visibility_channels(BotRoleId)
},
member_presence => mutual_visibility_presence(),
connected_user_ids => sets:from_list([10, 20, 30, 40]),
sessions => #{}
}.
@@ -319,9 +329,11 @@ slow_path_returns_self_when_viewer_sees_no_channels_test() ->
<<"channels">> => []
},
member_presence => make_presence_tab(#{10 => #{<<"status">> => <<"online">>}}),
connected_user_ids => sets:from_list([10]),
sessions => #{}
},
?assertEqual(1, compute_count(10, State)).
?assertEqual(1, compute_count(10, State)),
?assertEqual(0, compute_count(10, State#{connected_user_ids => sets:new()})).
slow_path_returns_zero_when_viewer_offline_and_no_channels_test() ->
GuildId = 1,
@@ -338,6 +350,7 @@ slow_path_returns_zero_when_viewer_offline_and_no_channels_test() ->
<<"channels">> => []
},
member_presence => make_presence_tab(#{10 => #{<<"status">> => <<"offline">>}}),
connected_user_ids => sets:from_list([10]),
sessions => #{}
},
?assertEqual(0, compute_count(10, State)).
@@ -362,6 +375,67 @@ index_counts_with_cached_session_channels_test() ->
State = Base#{sessions => cached_viewable_sessions()},
?assertEqual(3, compute_count(10, State)).
disconnected_online_row_is_not_counted_test() ->
Base = mutual_visibility_state(1, 5000),
State = Base#{connected_user_ids => sets:from_list([10, 20, 40])},
?assertNot(guild_member_list_connected:user_is_online(30, State)),
?assertEqual(1, compute_count(10, State)),
ViewerSet = guild_visibility:viewable_channel_set(10, State),
?assertEqual(
reference_count_mutually_visible(10, ViewerSet, State), compute_count(10, State)
).
disconnected_viewer_is_not_counted_test() ->
Base = mutual_visibility_state(1, 5000),
State = Base#{connected_user_ids => sets:from_list([20, 30, 40])},
?assertEqual(1, compute_count(10, State)).
counts_only_members_the_member_list_shows_online_test() ->
Base = mutual_visibility_state(1, 5000),
lists:foreach(
fun(Connected) ->
State = Base#{connected_user_ids => sets:from_list(Connected)},
ViewerSet = guild_visibility:viewable_channel_set(10, State),
Expected = length([
U
|| U <- [10, 20, 30, 40],
guild_member_list_connected:user_is_online(U, State),
U =:= 10 orelse
not sets:is_empty(
sets:intersection(
ViewerSet, guild_visibility:viewable_channel_set(U, State)
)
)
]),
?assertEqual(Expected, compute_count(10, State))
end,
[[], [10], [30], [10, 30], [20, 40], [10, 20, 30, 40]]
).
fully_indexed_count_builds_no_memo_test() ->
Base = mutual_visibility_state(1, 5000),
State = Base#{sessions => fully_indexed_sessions()},
ViewerSet = guild_visibility:viewable_channel_set(10, State),
Expected = reference_count_mutually_visible(10, ViewerSet, State),
?assertEqual({Expected, undefined}, count_and_memo(10, ViewerSet, State)),
?assertEqual(Expected, compute_count(10, State)).
unindexed_online_row_builds_the_memo_test() ->
Base = mutual_visibility_state(1, 5000),
State = Base#{sessions => cached_viewable_sessions()},
ViewerSet = guild_visibility:viewable_channel_set(10, State),
Expected = reference_count_mutually_visible(10, ViewerSet, State),
{Count, Memo} = count_and_memo(10, ViewerSet, State),
?assertEqual(Expected, Count),
?assertMatch(#{exceptions := _, cache := _}, Memo).
fully_indexed_sessions() ->
#{
<<"s20">> => #{user_id => 20, viewable_channels => #{101 => true}},
<<"s30">> => #{user_id => 30, viewable_channels => #{100 => true}},
<<"s40">> => #{user_id => 40, viewable_channels => #{}}
}.
cached_viewable_sessions() ->
#{
<<"s20a">> => #{user_id => 20, viewable_channels => undefined},
@@ -387,7 +461,8 @@ reference_count_mutually_visible(UserId, ViewerSet, State) ->
reference_count_online_member(UserId, ViewerSet, State, OtherUserId, Presence, Acc) when
is_integer(OtherUserId), is_map(Presence), OtherUserId > 0
->
case is_online(Presence) of
Connected = sets:is_element(OtherUserId, maps:get(connected_user_ids, State)),
case Connected andalso is_online(Presence) of
false -> Acc;
true when OtherUserId =:= UserId -> Acc + 1;
true -> reference_count_if_mutually_visible(OtherUserId, ViewerSet, State, Acc)
@@ -405,6 +480,276 @@ reference_count_if_mutually_visible(OtherUserId, ViewerSet, State, Acc) ->
false -> Acc + 1
end.
sessionless_user_overwrite_is_not_shared_with_same_roles_test() ->
State = sessionless_role_state(
[user_overwrite(<<"30">>, 0, view_perm())], #{}, <<"999">>, #{}
),
?assertEqual(2, compute_count(10, State)),
?assertEqual(2, reference_slow_count(10, State)).
sessionless_virtual_access_is_not_shared_with_same_roles_test() ->
State = sessionless_role_state([], #{40 => sets:from_list([100])}, <<"999">>, #{}),
?assertEqual(4, compute_count(10, State)),
?assertEqual(4, reference_slow_count(10, State)).
sessionless_owner_is_not_shared_with_same_roles_test() ->
State = sessionless_role_state([], #{}, <<"40">>, #{}),
?assertEqual(4, compute_count(10, State)),
?assertEqual(4, reference_slow_count(10, State)).
sessionless_members_with_different_base_permissions_test() ->
State = sessionless_role_state(
[], #{}, <<"999">>, #{10 => [5000, 6000], 40 => [6000], 50 => [6001]}
),
?assertEqual(4, compute_count(10, State)),
?assertEqual(4, reference_slow_count(10, State)).
sessionless_members_with_reordered_and_repeated_roles_test() ->
State = sessionless_role_state(
[], #{}, <<"999">>, #{20 => [6001, 5000], 30 => [5000, 5000, 6001]}
),
?assertEqual(3, compute_count(10, State)),
?assertEqual(3, reference_slow_count(10, State)).
sessionless_role_state(ExtraOverwrites, VirtualAccess, OwnerId, RoleOverrides) ->
RoleId = 5000,
Roles = maps:merge(
#{10 => [RoleId], 20 => [RoleId], 30 => [RoleId], 40 => [], 50 => []}, RoleOverrides
),
Member = fun(Id) ->
#{
<<"user">> => #{<<"id">> => integer_to_binary(Id)},
<<"roles">> => [integer_to_binary(R) || R <- maps:get(Id, Roles)]
}
end,
#{
id => 1,
data => #{
<<"guild">> => #{<<"owner_id">> => OwnerId},
<<"roles">> =>
mutual_visibility_roles(1, RoleId) ++
[
#{
<<"id">> => <<"6000">>,
<<"permissions">> => integer_to_binary(view_perm())
},
#{<<"id">> => <<"6001">>, <<"permissions">> => <<"0">>}
],
<<"members">> => maps:from_list([{Id, Member(Id)} || Id <- [10, 20, 30, 40, 50]]),
<<"channels">> => [
#{
<<"id">> => <<"100">>,
<<"type">> => 0,
<<"permission_overwrites">> => [
overwrite(<<"1">>, 0, 0, view_perm()),
overwrite(integer_to_binary(RoleId), 0, view_perm(), 0)
| ExtraOverwrites
]
},
#{
<<"id">> => <<"101">>,
<<"type">> => 0,
<<"permission_overwrites">> => [overwrite(<<"1">>, 0, 0, view_perm())]
},
#{<<"id">> => <<"102">>, <<"type">> => 0, <<"permission_overwrites">> => []}
]
},
virtual_channel_access => VirtualAccess,
member_presence => make_presence_tab(
maps:from_keys([10, 20, 30, 40, 50], #{<<"status">> => <<"online">>})
),
connected_user_ids => sets:from_list([10, 20, 30, 40, 50]),
sessions => #{}
}.
user_overwrite(UserId, Allow, Deny) ->
overwrite(UserId, 1, Allow, Deny).
overwrite(Id, Type, Allow, Deny) ->
#{
<<"id">> => Id,
<<"type">> => Type,
<<"allow">> => integer_to_binary(Allow),
<<"deny">> => integer_to_binary(Deny)
}.
matches_reference_on_random_guilds_test_() ->
{timeout, 120, fun() ->
lists:foreach(fun assert_random_guild_matches_reference/1, lists:seq(1, 150))
end}.
assert_random_guild_matches_reference(Seed) ->
_ = rand:seed(exsss, {Seed, Seed * 7, Seed * 13}),
State = random_guild_state(),
Viewers = lists:seq(1, 60),
Results = [{Viewer, compute_count(Viewer, State)} || Viewer <- Viewers],
Expected = [{Viewer, reference_compute_count(Viewer, State)} || Viewer <- Viewers],
ets:delete(maps:get(member_presence, State)),
?assertEqual({Seed, Expected}, {Seed, Results}).
reference_compute_count(UserId, State) ->
case viewer_sees_everything(UserId, State) of
true -> guild_member_list:get_online_count(State);
false -> reference_slow_count(UserId, State)
end.
reference_slow_count(UserId, State) ->
ViewerSet = guild_visibility:viewable_channel_set(UserId, State),
case sets:is_empty(ViewerSet) of
true -> self_online_count(UserId, State);
false -> reference_count_mutually_visible(UserId, ViewerSet, State)
end.
random_guild_state() ->
GuildId = 1,
RoleIds = lists:seq(2, 9),
MemberIds = [U || U <- lists:seq(10, 55), rand:uniform(10) > 1],
Roles = [random_role(GuildId, 0) | [random_role(R, 12) || R <- RoleIds]],
Members = maps:from_list([{U, random_member(U, GuildId, RoleIds)} || U <- MemberIds]),
Channels = random_channels(GuildId, RoleIds, lists:seq(10, 55)),
Data0 = #{
<<"guild">> => #{<<"owner_id">> => integer_to_binary(pick(MemberIds))},
<<"roles">> => Roles,
<<"members">> => Members,
<<"channels">> => Channels
},
Data =
case rand:uniform(2) of
1 ->
Data0;
2 ->
guild_data_index:put_channels(
Channels, guild_data_index:put_roles(Roles, Data0)
)
end,
State0 = #{
id => GuildId,
data => Data,
virtual_channel_access => random_virtual_access(Channels),
member_presence => random_presence_tab(),
sessions => #{}
},
State1 = State0#{sessions => random_sessions(MemberIds, Channels, State0)},
State1#{
connected_user_ids => sets:from_list([U || U <- lists:seq(10, 60), rand:uniform(4) > 1])
}.
random_role(RoleId, AdminOneIn) ->
View =
case rand:uniform(4) of
1 -> 0;
_ -> view_perm()
end,
Admin =
case AdminOneIn > 0 andalso rand:uniform(AdminOneIn) =:= 1 of
true -> admin_perm();
false -> 0
end,
Other = rand:uniform(1024) bsl 20,
#{
<<"id">> => integer_to_binary(RoleId),
<<"permissions">> => integer_to_binary(View bor Admin bor Other)
}.
random_member(UserId, GuildId, RoleIds) ->
Roles = [R || R <- [GuildId | RoleIds], rand:uniform(3) =:= 1],
Shuffled = [R || {_, R} <- lists:sort([{rand:uniform(), R} || R <- Roles ++ Roles])],
#{
<<"user">> => #{<<"id">> => integer_to_binary(UserId)},
<<"roles">> => [integer_to_binary(R) || R <- lists:sublist(Shuffled, length(Roles) + 1)]
}.
random_channels(GuildId, RoleIds, UserIds) ->
Categories = [
random_channel(C, 4, null, GuildId, RoleIds, UserIds)
|| C <- [100, 101, 102]
],
Children = [
random_channel(
C, pick([0, 0, 2, 5]), pick([null, 100, 101, 102]), GuildId, RoleIds, UserIds
)
|| C <- lists:seq(200, 211)
],
Categories ++ Children.
random_channel(ChannelId, Type, ParentId, GuildId, RoleIds, UserIds) ->
Parent =
case ParentId of
null -> null;
_ -> integer_to_binary(ParentId)
end,
#{
<<"id">> => integer_to_binary(ChannelId),
<<"type">> => Type,
<<"parent_id">> => Parent,
<<"permission_overwrites">> => random_overwrites(GuildId, RoleIds, UserIds)
}.
random_overwrites(GuildId, RoleIds, UserIds) ->
Everyone = [random_overwrite(GuildId, 0) || rand:uniform(2) =:= 1],
RoleOws = [random_overwrite(R, 0) || R <- lists:sublist(RoleIds, 4), rand:uniform(5) =:= 1],
UserOws = [random_overwrite(U, 1) || U <- UserIds, rand:uniform(25) =:= 1],
Everyone ++ RoleOws ++ UserOws.
random_overwrite(Id, Type) ->
{Allow, Deny} = pick([
{view_perm(), 0}, {0, view_perm()}, {0, 0}, {view_perm(), view_perm()}
]),
overwrite(integer_to_binary(Id), Type, Allow, Deny).
random_virtual_access(Channels) ->
maps:from_list([
{U, sets:from_list([channel_int_id(pick(Channels))])}
|| U <- lists:seq(10, 60), rand:uniform(15) =:= 1
]).
random_presence_tab() ->
make_presence_tab(
maps:from_list([
{U, random_presence()}
|| U <- lists:seq(10, 60), rand:uniform(5) > 1
])
).
random_presence() ->
case rand:uniform(7) of
1 -> #{};
2 -> #{<<"status">> => <<"offline">>};
3 -> #{<<"status">> => <<"invisible">>};
4 -> #{<<"status">> => <<"idle">>};
_ -> #{<<"status">> => <<"online">>}
end.
random_sessions(MemberIds, Channels, State) ->
maps:from_list(
lists:append([random_user_sessions(U, Channels, State) || U <- MemberIds])
).
random_user_sessions(UserId, Channels, State) ->
[
{
iolist_to_binary([integer_to_list(UserId), "-", integer_to_list(N)]),
#{
user_id => UserId,
viewable_channels => random_session_channels(UserId, Channels, State)
}
}
|| N <- lists:seq(1, rand:uniform(4) - 1)
].
random_session_channels(UserId, Channels, State) ->
case rand:uniform(4) of
1 -> undefined;
2 -> maps:from_keys([channel_int_id(C) || C <- Channels, rand:uniform(3) =:= 1], true);
_ -> maps:from_keys(guild_visibility:get_user_viewable_channels(UserId, State), true)
end.
channel_int_id(Channel) ->
binary_to_integer(maps:get(<<"id">>, Channel)).
pick(List) ->
lists:nth(rand:uniform(length(List)), List).
make_presence_tab(Map) ->
Tab = ets:new(test_member_presence, [set, public]),
maps:foreach(fun(K, V) -> ets:insert(Tab, {K, V}) end, Map),
+352 -40
View File
@@ -16,6 +16,9 @@
-define(PASSIVE_SYNC_INTERVAL, 30000).
-define(LARGE_GUILD_MEMBER_COUNT, 250).
-define(HEAVY_MEMBER_DATA_KEYS, [
<<"members">>, members_normalized, <<"member_role_index">>, members_sorted_ids
]).
-type guild_state() :: map().
-type channel_id() :: binary().
@@ -35,61 +38,94 @@ handle_passive_sync(State) ->
_ = schedule_passive_sync(State),
{noreply, State}.
%% send_passive_updates/4 keeps no session unless is_large_guild/1 holds, so for any other
%% guild the spawned child returns without touching a session, a payload or the registry.
-spec maybe_spawn_passive_updates(integer(), guild_state()) -> ok.
maybe_spawn_passive_updates(GuildId, State) ->
case is_large_guild(maps:get(member_count, State, undefined)) of
false -> ok;
true -> spawn_passive_updates(GuildId, State)
case large_guild_passive_sessions(GuildId, State) of
PassiveSessions when map_size(PassiveSessions) =:= 0 ->
ok;
PassiveSessions ->
SyncState = passive_sync_state(State, PassiveSessions),
_ = spawn(fun() -> send_passive_updates(GuildId, PassiveSessions, SyncState) end),
ok
end.
-spec spawn_passive_updates(integer(), guild_state()) -> ok.
spawn_passive_updates(GuildId, State) ->
_ = spawn(fun() -> send_passive_updates_for_state(GuildId, State) end),
ok.
-spec send_passive_updates_to_sessions(guild_state()) -> guild_state().
send_passive_updates_to_sessions(State) ->
ok = send_passive_updates_for_state(maps:get(id, State), State),
GuildId = maps:get(id, State),
PassiveSessions = large_guild_passive_sessions(GuildId, State),
ok = send_passive_updates(
GuildId, PassiveSessions, passive_sync_state(State, PassiveSessions)
),
State.
-spec send_passive_updates_for_state(integer(), guild_state()) -> ok.
send_passive_updates_for_state(GuildId, State) ->
Sessions = maps:get(sessions, State, #{}),
Data = maps:get(data, State, #{}),
MemberCount = maps:get(member_count, State, undefined),
VoiceStates = maps:get(voice_states, State, #{}),
send_passive_updates(
GuildId, Sessions, passive_sync_state(State, Data, VoiceStates), MemberCount
-spec large_guild_passive_sessions(integer(), guild_state()) -> map().
large_guild_passive_sessions(GuildId, State) ->
case is_large_guild(maps:get(member_count, State, undefined)) of
false -> #{};
true -> passive_sessions(GuildId, maps:get(sessions, State, #{}))
end.
-spec passive_sessions(integer(), map()) -> map().
passive_sessions(GuildId, Sessions) ->
maps:filtermap(
fun(_SessionId, SessionData) ->
case session_passive:is_passive(GuildId, SessionData) of
true -> {true, maps:with([pid, user_id], SessionData)};
false -> false
end
end,
Sessions
).
-spec passive_sync_state(guild_state(), map(), map()) -> guild_state().
passive_sync_state(State, Data, VoiceStates) ->
State#{data => Data, voice_states => VoiceStates}.
-spec passive_sync_state(guild_state(), map()) -> guild_state().
passive_sync_state(State, PassiveSessions) ->
Base = maps:with([id, voice_server_pid, virtual_channel_access], State),
Base#{
data => passive_sync_data(maps:get(data, State, #{})),
voice_states => maps:get(voice_states, State, #{}),
sessions => first_viewable_sessions(PassiveSessions, maps:get(sessions, State, #{}))
}.
-spec passive_sync_data(map()) -> map().
passive_sync_data(#{members_ets := Tab} = Data) when is_reference(Tab) ->
maps:without(?HEAVY_MEMBER_DATA_KEYS, Data);
passive_sync_data(Data) ->
Data.
-spec first_viewable_sessions(map(), map()) -> map().
first_viewable_sessions(PassiveSessions, Sessions) ->
UserIds = maps:fold(
fun(_SessionId, SessionData, Acc) ->
Acc#{maps:get(user_id, SessionData, undefined) => true}
end,
#{},
PassiveSessions
),
first_viewable_sessions_iter(maps:next(maps:iterator(Sessions)), UserIds, #{}).
-spec first_viewable_sessions_iter(none | {term(), term(), maps:iterator()}, map(), map()) ->
map().
first_viewable_sessions_iter(none, _UserIds, Acc) ->
Acc;
first_viewable_sessions_iter(
{SessionId, #{user_id := UserId, viewable_channels := Viewable}, Next}, UserIds, Acc
) when is_map(Viewable), is_map_key(UserId, UserIds) ->
first_viewable_sessions_iter(
maps:next(Next),
maps:remove(UserId, UserIds),
Acc#{SessionId => #{user_id => UserId, viewable_channels => Viewable}}
);
first_viewable_sessions_iter({_SessionId, _SessionData, Next}, UserIds, Acc) ->
first_viewable_sessions_iter(maps:next(Next), UserIds, Acc).
-spec is_large_guild(term()) -> boolean().
is_large_guild(MemberCount) ->
is_integer(MemberCount) andalso MemberCount > ?LARGE_GUILD_MEMBER_COUNT.
-spec send_passive_updates(integer(), map(), guild_state(), non_neg_integer() | undefined) ->
ok.
send_passive_updates(GuildId, Sessions, State, MemberCount) ->
Data = maps:get(data, State, #{}),
Channels = guild_data_index:channel_list(Data),
IsLargeGuild = is_large_guild(MemberCount),
PassiveSessions = maps:filter(
fun(_SessionId, SessionData) ->
IsLargeGuild andalso session_passive:is_passive(GuildId, SessionData)
end,
Sessions
),
case map_size(PassiveSessions) of
0 ->
ok;
_ ->
send_passive_session_updates(PassiveSessions, GuildId, Channels, State)
end.
-spec send_passive_updates(integer(), map(), guild_state()) -> ok.
send_passive_updates(GuildId, PassiveSessions, SyncState) ->
Channels = guild_data_index:channel_list(maps:get(data, SyncState)),
send_passive_session_updates(PassiveSessions, GuildId, Channels, SyncState).
-spec send_passive_session_updates(map(), integer(), [map()], guild_state()) -> ok.
send_passive_session_updates(PassiveSessions, GuildId, Channels, SyncState) ->
@@ -466,4 +502,280 @@ flush_passive_dispatches() ->
_ -> flush_passive_dispatches()
end.
passive_sync_state_matches_full_state_dispatches_test() ->
with_differential_guild(fun(VoiceServer, Rounds, SessionIds, GuildId) ->
Reference = run_passive_rounds(
fun reference_send_passive_updates/1, VoiceServer, Rounds, SessionIds, GuildId
),
Projected = run_passive_rounds(
fun send_passive_updates_to_sessions/1, VoiceServer, Rounds, SessionIds, GuildId
),
[{Round1Dispatches, _}, {Round2Dispatches, _}] = Reference,
?assertEqual(4, length(Round1Dispatches)),
?assertEqual(4, length(Round2Dispatches)),
?assertEqual(Reference, Projected)
end).
passive_sync_state_drops_members_and_unrelated_sessions_test() ->
with_differential_guild(fun(_VoiceServer, [{_, State} | _], _SessionIds, GuildId) ->
PassiveSessions = large_guild_passive_sessions(GuildId, State),
SyncState = passive_sync_state(State, PassiveSessions),
?assertEqual(
[data, id, sessions, virtual_channel_access, voice_server_pid, voice_states],
lists:sort(maps:keys(SyncState))
),
SyncData = maps:get(data, SyncState),
?assertEqual([], [K || K <- ?HEAVY_MEMBER_DATA_KEYS, is_map_key(K, SyncData)]),
?assert(is_map_key(members_ets, SyncData)),
?assertEqual(
#{
<<"u2-passive">> => #{
user_id => diff_user(2), viewable_channels => #{diff_channel(1) => true}
},
<<"u3-first">> => #{
user_id => diff_user(3), viewable_channels => #{diff_channel(3) => true}
}
},
maps:get(sessions, SyncState)
),
?assertEqual(
lists:sort([
<<"u1-passive">>,
<<"u2-passive">>,
<<"u3-passive">>,
<<"u4-passive">>,
<<"u5-passive">>
]),
lists:sort(maps:keys(PassiveSessions))
)
end).
passive_sync_data_keeps_members_without_members_ets_test() ->
Data = #{<<"members">> => #{1 => #{}}, members_sorted_ids => [1], <<"channels">> => []},
?assertEqual(Data, passive_sync_data(Data)).
first_viewable_sessions_picks_the_session_the_full_state_lookup_finds_test() ->
Sessions = maps:from_list([
{integer_to_binary(N), #{user_id => N rem 3, viewable_channels => #{N => true}}}
|| N <- lists:seq(1, 64)
]),
Passive = #{<<"p">> => #{user_id => 1}, <<"q">> => #{user_id => 2}},
Projected = first_viewable_sessions(Passive, Sessions),
?assertEqual(2, map_size(Projected)),
lists:foreach(
fun(UserId) ->
?assertEqual(
guild_visibility_channels:get_cached_viewable_channel_map(
UserId, #{sessions => Sessions}
),
guild_visibility_channels:get_cached_viewable_channel_map(
UserId, #{sessions => Projected}
)
)
end,
[1, 2]
).
reference_send_passive_updates(State) ->
GuildId = maps:get(id, State),
PassiveSessions = reference_passive_sessions(
GuildId, maps:get(sessions, State), maps:get(member_count, State)
),
Data = maps:get(data, State),
Channels = guild_data_index:channel_list(Data),
SyncState = State#{data => Data, voice_states => maps:get(voice_states, State, #{})},
ok = send_passive_session_updates(PassiveSessions, GuildId, Channels, SyncState),
State.
run_passive_rounds(Send, VoiceServer, Rounds, SessionIds, GuildId) ->
flush_passive_dispatches(),
ok = passive_sync_registry:init(),
lists:foreach(
fun(SessionId) -> passive_sync_registry:delete(SessionId, GuildId) end, SessionIds
),
[
begin
ok = gen_server:call(VoiceServer, {set, VoiceStates}),
_ = Send(State),
{collect_passive_dispatches([]), [
{SessionId, passive_sync_registry:lookup(SessionId, GuildId)}
|| SessionId <- SessionIds
]}
end
|| {VoiceStates, State} <- Rounds
].
collect_passive_dispatches(Acc) ->
case receive_passive_dispatch(0) of
no_dispatch -> lists:reverse(Acc);
Payload -> collect_passive_dispatches([Payload | Acc])
end.
with_differential_guild(Fun) ->
GuildId = 1427764882469228556,
Tab = ets:new(passive_diff_members, [set, public]),
VoiceServer = spawn(fun() -> fake_voice_server(#{}) end),
try
Rounds = [
{
diff_voice_states(GuildId, Round),
differential_state(GuildId, Tab, VoiceServer, Round)
}
|| Round <- [0, 1]
],
[{_, #{sessions := Sessions}} | _] = Rounds,
Fun(VoiceServer, Rounds, lists:sort(maps:keys(Sessions)), GuildId)
after
exit(VoiceServer, kill),
ets:delete(Tab)
end.
fake_voice_server(VoiceStates) ->
receive
{'$gen_call', From, {set, NewVoiceStates}} ->
gen_server:reply(From, ok),
fake_voice_server(NewVoiceStates);
{'$gen_call', From, {get_voice_states_map}} ->
gen_server:reply(From, VoiceStates),
fake_voice_server(VoiceStates)
end.
differential_state(GuildId, Tab, VoiceServer, Round) ->
Members = [
diff_member(diff_user(0), []),
diff_member(diff_user(1), [diff_role(1)]),
diff_member(diff_user(2), [diff_role(2)]),
diff_member(diff_user(3), [diff_role(1), diff_role(2)]),
diff_member(diff_user(4), []),
diff_member(diff_user(6), [diff_role(1)]),
diff_member(diff_user(7), [])
],
true = ets:insert(Tab, [
{diff_user(N), M}
|| {N, M} <- lists:zip([0, 1, 2, 3, 4, 6, 7], Members)
]),
Data = guild_data_index:normalize_map(#{
<<"guild">> => #{<<"id">> => GuildId, <<"owner_id">> => diff_user(0)},
<<"roles">> => [
#{<<"id">> => GuildId, <<"permissions">> => <<"1024">>, <<"position">> => 0},
#{<<"id">> => diff_role(1), <<"permissions">> => <<"0">>, <<"position">> => 1},
#{<<"id">> => diff_role(2), <<"permissions">> => <<"0">>, <<"position">> => 2}
],
<<"channels">> => diff_channels(GuildId, Round),
<<"members">> => Members
}),
#{
id => GuildId,
member_count => 55278,
voice_server_pid => VoiceServer,
virtual_channel_access => #{diff_user(4) => sets:from_list([diff_channel(5)])},
voice_states => #{},
member_presence => make_ref(),
presence_subscriptions => #{diff_user(1) => true},
data => Data#{members_ets => Tab},
sessions => diff_sessions(GuildId)
}.
diff_channels(GuildId, Round) ->
Deny = fun(Id) ->
#{<<"id">> => Id, <<"type">> => 0, <<"allow">> => <<"0">>, <<"deny">> => <<"1024">>}
end,
Allow = fun(Id, Type) ->
#{<<"id">> => Id, <<"type">> => Type, <<"allow">> => <<"1024">>, <<"deny">> => <<"0">>}
end,
[
#{
<<"id">> => diff_channel(1),
<<"type">> => 0,
<<"last_message_id">> => diff_message(1, Round)
},
#{
<<"id">> => diff_channel(2),
<<"type">> => 0,
<<"last_message_id">> => diff_message(2, Round),
<<"permission_overwrites">> => [Deny(GuildId), Allow(diff_role(1), 0)]
},
#{
<<"id">> => diff_channel(3),
<<"type">> => 4,
<<"permission_overwrites">> => [Deny(GuildId)]
},
#{
<<"id">> => diff_channel(4),
<<"type">> => 0,
<<"parent_id">> => diff_channel(3),
<<"last_message_id">> => diff_message(4, 0),
<<"permission_overwrites">> => [Deny(GuildId), Allow(diff_role(2), 0)]
},
#{
<<"id">> => diff_channel(5),
<<"type">> => 2,
<<"last_message_id">> => diff_message(5, Round),
<<"permission_overwrites">> => [Deny(GuildId), Allow(diff_user(7), 1)]
},
#{<<"id">> => diff_channel(6), <<"type">> => 0, <<"last_message_id">> => null}
].
diff_sessions(GuildId) ->
Active = sets:from_list([GuildId]),
Passive = sets:new(),
Session = fun(User, ActiveGuilds, Extra) ->
maps:merge(
#{
user_id => diff_user(User),
pid => self(),
active_guilds => ActiveGuilds,
bot => false,
user_roles => [],
pending_connect => false
},
Extra
)
end,
#{
<<"u0-active">> => Session(0, Active, #{viewable_channels => #{diff_channel(2) => true}}),
<<"u1-passive">> => Session(1, Passive, #{}),
<<"u2-passive">> => Session(2, Passive, #{
viewable_channels => #{diff_channel(1) => true}
}),
<<"u3-first">> => Session(3, Active, #{viewable_channels => #{diff_channel(3) => true}}),
<<"u3-passive">> => Session(3, Passive, #{
viewable_channels => #{diff_channel(5) => true}
}),
<<"u4-passive">> => Session(4, Passive, #{viewable_channels => not_a_map}),
<<"u5-passive">> => Session(5, Passive, #{}),
<<"u6-bot">> => Session(6, Passive, #{bot => true}),
<<"u7-active">> => Session(7, Active, #{})
}.
diff_voice_states(GuildId, Round) ->
VoiceState = fun(Conn, User, Channel, Version) ->
{Conn, #{
<<"connection_id">> => Conn,
<<"guild_id">> => integer_to_binary(GuildId),
<<"channel_id">> => integer_to_binary(diff_channel(Channel)),
<<"user_id">> => integer_to_binary(diff_user(User)),
<<"version">> => Version
}}
end,
maps:from_list(
[
VoiceState(<<"c1">>, 7, 5, 1),
VoiceState(<<"c2">>, 1, 1, 1 + Round),
VoiceState(<<"c4">>, 3, 3, 1)
] ++
[VoiceState(<<"c3">>, 2, 2, 1) || Round =:= 0]
).
diff_member(UserId, Roles) ->
#{<<"user">> => #{<<"id">> => integer_to_binary(UserId)}, <<"roles">> => Roles}.
diff_user(N) -> 1130650140672000000 + N.
diff_role(N) -> 1428000118785000000 + N.
diff_channel(N) -> 1428100000000000000 + N.
diff_message(N, Round) -> 1500000000000000000 + N * 10 + Round.
-endif.
@@ -14,6 +14,7 @@
delete/1,
get_permissions/3,
get_snapshot/1,
get_role_members/2,
has_member/2,
get_member/2,
strip_data/1,
@@ -24,10 +25,11 @@
-type guild_id() :: integer().
-type user_id() :: integer().
-type channel_id() :: integer().
-type role_id() :: integer().
-type guild_state() :: map().
-type guild_data() :: map().
-export_type([guild_id/0, user_id/0, channel_id/0, guild_state/0, guild_data/0]).
-export_type([guild_id/0, user_id/0, channel_id/0, role_id/0, guild_state/0, guild_data/0]).
-define(TABLE, guild_permission_cache).
-define(STRIPPED_MEMBERS_MEMO, guild_permission_cache_stripped_members).
@@ -57,24 +59,33 @@ put_normalized_data(GuildId, NormalizedData) when is_integer(GuildId), is_map(No
ensure_table(),
StrippedData = strip_data(NormalizedData),
Snapshot = #{id => GuildId, data => StrippedData},
true = ets:insert(?TABLE, {GuildId, Snapshot}),
RoleIndex = maps:get(<<"member_role_index">>, NormalizedData, #{}),
MemberSource = maps:with([members_ets], StrippedData),
true = ets:insert(?TABLE, [
{GuildId, Snapshot},
{role_index_key(GuildId), RoleIndex, MemberSource}
]),
ok;
put_normalized_data(_, _) ->
ok.
-spec role_index_key(guild_id()) -> {member_role_index, guild_id()}.
role_index_key(GuildId) ->
{member_role_index, GuildId}.
-spec delete(guild_id()) -> ok.
delete(GuildId) when is_integer(GuildId) ->
case ets:whereis(?TABLE) of
undefined -> ok;
_ -> safe_ets_delete(GuildId)
_ -> safe_ets_delete([GuildId, role_index_key(GuildId)])
end,
ok;
delete(_) ->
ok.
-spec safe_ets_delete(guild_id()) -> ok.
safe_ets_delete(GuildId) ->
try ets:delete(?TABLE, GuildId) of
-spec safe_ets_delete([term()]) -> ok.
safe_ets_delete(Keys) ->
try lists:foreach(fun(Key) -> true = ets:delete(?TABLE, Key) end, Keys) of
_ -> ok
catch
error:badarg -> ok
@@ -133,21 +144,38 @@ get_snapshot(GuildId) when is_integer(GuildId) ->
get_snapshot(_) ->
{error, not_found}.
-spec snapshot_with_live_members(guild_state()) -> {ok, guild_state()} | {error, not_found}.
snapshot_with_live_members(Snapshot) ->
case snapshot_member_table(Snapshot) of
Tab when is_reference(Tab) -> live_member_table_snapshot(Tab, Snapshot);
undefined -> {ok, Snapshot}
-spec get_role_members(guild_id(), role_id()) -> {ok, [user_id()]} | {error, not_found}.
get_role_members(GuildId, RoleId) when is_integer(GuildId), is_integer(RoleId) ->
ensure_table(),
case ets:lookup(?TABLE, role_index_key(GuildId)) of
[{_Key, RoleIndex, MemberSource}] ->
live_role_members(maps:get(RoleId, RoleIndex, #{}), MemberSource);
[] ->
{error, not_found}
end;
get_role_members(_, _) ->
{error, not_found}.
-spec live_role_members(map(), guild_data()) -> {ok, [user_id()]} | {error, not_found}.
live_role_members(RoleMembers, MemberSource) ->
case live_members(data_member_table(MemberSource)) of
true -> {ok, lists:sort(maps:keys(RoleMembers))};
false -> {error, not_found}
end.
-spec live_member_table_snapshot(ets:tid(), guild_state()) ->
{ok, guild_state()} | {error, not_found}.
live_member_table_snapshot(Tab, Snapshot) ->
case ets:info(Tab, owner) of
undefined -> {error, not_found};
_ -> {ok, Snapshot}
-spec snapshot_with_live_members(guild_state()) -> {ok, guild_state()} | {error, not_found}.
snapshot_with_live_members(Snapshot) ->
case live_members(snapshot_member_table(Snapshot)) of
true -> {ok, Snapshot};
false -> {error, not_found}
end.
-spec live_members(ets:tid() | undefined) -> boolean().
live_members(undefined) ->
true;
live_members(Tab) ->
ets:info(Tab, owner) =/= undefined.
-spec snapshot_member_table(guild_state()) -> ets:tid() | undefined.
snapshot_member_table(#{data := Data}) when is_map(Data) ->
data_member_table(Data);
@@ -181,7 +209,6 @@ strip_data(Data) when is_map(Data) ->
Roles = strip_roles(maps:get(<<"roles">>, Data, [])),
Channels = strip_channels(maps:get(<<"channels">>, Data, [])),
ChannelIndex = strip_channel_index(maps:get(<<"channel_index">>, Data, #{})),
MemberRoleIndex = maps:get(<<"member_role_index">>, Data, #{}),
RolePermsCache = maps:get(role_perms_cache, Data, #{}),
OverwritePermsCache = maps:get(overwrite_perms_cache, Data, #{}),
with_member_source(Data, #{
@@ -189,7 +216,6 @@ strip_data(Data) when is_map(Data) ->
<<"roles">> => Roles,
<<"channels">> => Channels,
<<"channel_index">> => ChannelIndex,
<<"member_role_index">> => MemberRoleIndex,
role_perms_cache => RolePermsCache,
overwrite_perms_cache => OverwritePermsCache
});
@@ -416,21 +442,21 @@ parse_user_id(Id) ->
-spec migrate_existing_entries() -> {ok, non_neg_integer()}.
migrate_existing_entries() ->
ensure_table(),
Count = ets:foldl(
fun
({GuildId, #{data := Data} = _Snapshot}, Acc) ->
Stripped = strip_data(Data),
NewSnapshot = #{id => GuildId, data => Stripped},
true = ets:insert(?TABLE, {GuildId, NewSnapshot}),
Acc + 1;
(_, Acc) ->
Acc
end,
0,
?TABLE
),
Count = ets:foldl(fun migrate_entry/2, 0, ?TABLE),
{ok, Count}.
-spec migrate_entry(term(), non_neg_integer()) -> non_neg_integer().
migrate_entry({GuildId, #{data := #{<<"member_role_index">> := _} = Data}}, Acc) when
is_integer(GuildId)
->
ok = put_normalized_data(GuildId, Data),
Acc + 1;
migrate_entry({GuildId, #{data := Data}}, Acc) when is_integer(GuildId), is_map(Data) ->
true = ets:insert(?TABLE, {GuildId, #{id => GuildId, data => strip_data(Data)}}),
Acc + 1;
migrate_entry(_, Acc) ->
Acc.
-ifdef(TEST).
strip_member_preserves_communication_disabled_until_test() ->
+43 -36
View File
@@ -6,7 +6,10 @@
-export([
get_member_permissions/3,
compute_member_permissions/4,
member_base_permissions/3,
channel_permissions/4,
can_view_channel/4,
viewable_channel_ids/4,
can_view_channel_by_permissions/4,
can_view_channel_members/4,
can_manage_channel/3,
@@ -28,7 +31,8 @@
guild_state/0,
member/0,
maybe_member/0,
member_roles/0
member_roles/0,
base_permissions/0
]).
-define(ALL_PERMISSIONS, 16#FFFFFFFFFFFFFFFF).
@@ -42,6 +46,7 @@
-type member() :: map().
-type maybe_member() :: member() | undefined.
-type member_roles() :: [role_id()].
-type base_permissions() :: permission() | {permission(), member_roles(), role_id()}.
-spec get_member_permissions(user_id(), maybe_channel_id(), guild_state()) -> permission().
get_member_permissions(UserId, ChannelId, State) ->
@@ -49,30 +54,49 @@ get_member_permissions(UserId, ChannelId, State) ->
-spec compute_member_permissions(user_id(), maybe_channel_id(), maybe_member(), guild_state()) ->
permission().
compute_member_permissions(UserId, ChannelId, ProvidedMember, State) when is_integer(UserId) ->
compute_member_permissions(UserId, ChannelId, ProvidedMember, State) ->
channel_permissions(
member_base_permissions(UserId, ProvidedMember, State), UserId, ChannelId, State
).
-spec member_base_permissions(user_id(), maybe_member(), guild_state()) -> base_permissions().
member_base_permissions(UserId, ProvidedMember, State) when is_integer(UserId) ->
case guild_permissions_common:resolve_data_map(State) of
undefined ->
0;
Data ->
compute_permissions_for_data(UserId, ChannelId, ProvidedMember, State, Data)
base_permissions_for_data(UserId, ProvidedMember, State, Data)
end;
compute_member_permissions(_, _, _, _) ->
member_base_permissions(_, _, _) ->
0.
-spec compute_permissions_for_data(
user_id(), maybe_channel_id(), maybe_member(), guild_state(), guild_data()
) -> permission().
compute_permissions_for_data(UserId, ChannelId, ProvidedMember, State, Data) ->
-spec channel_permissions(base_permissions(), user_id(), maybe_channel_id(), guild_state()) ->
permission().
channel_permissions({Permissions, MemberRoles, GuildId}, UserId, ChannelId, State) ->
guild_permissions_overwrites:maybe_apply_channel_overwrites(
Permissions, UserId, MemberRoles, ChannelId, GuildId, State
);
channel_permissions(Permissions, _UserId, _ChannelId, _State) ->
Permissions.
-spec base_permissions_for_data(user_id(), maybe_member(), guild_state(), guild_data()) ->
base_permissions().
base_permissions_for_data(UserId, ProvidedMember, State, Data) ->
OwnerId = guild_owner_id(Data),
case UserId =:= OwnerId of
true -> ?ALL_PERMISSIONS;
false -> compute_non_owner_permissions(UserId, ChannelId, ProvidedMember, State, Data)
false -> non_owner_base_permissions(UserId, ProvidedMember, State, Data)
end.
-spec can_view_channel(user_id(), integer(), maybe_member(), guild_state()) -> boolean().
can_view_channel(UserId, ChannelId, Member, State) ->
guild_permissions_check:can_view_channel(UserId, ChannelId, Member, State).
-spec viewable_channel_ids(user_id(), base_permissions(), [map()], guild_state()) ->
#{integer() => true}.
viewable_channel_ids(UserId, Base, Channels, State) ->
guild_permissions_check:viewable_channel_ids(UserId, Base, Channels, State).
-spec can_view_channel_by_permissions(user_id(), integer(), maybe_member(), guild_state()) ->
boolean().
can_view_channel_by_permissions(UserId, ChannelId, Member, State) ->
@@ -119,21 +143,18 @@ find_channel_by_id(ChannelId, State) ->
view_inputs(ChannelId, State) ->
guild_permissions_check:view_inputs(ChannelId, State).
-spec compute_non_owner_permissions(
user_id(), maybe_channel_id(), maybe_member(), guild_state(), guild_data()
) -> permission().
compute_non_owner_permissions(UserId, ChannelId, ProvidedMember, State, Data) ->
-spec non_owner_base_permissions(user_id(), maybe_member(), guild_state(), guild_data()) ->
base_permissions().
non_owner_base_permissions(UserId, ProvidedMember, State, Data) ->
case resolve_member(UserId, ProvidedMember, State) of
undefined ->
0;
Member ->
compute_member_role_permissions(UserId, ChannelId, Member, State, Data)
member_role_base_permissions(Member, State, Data)
end.
-spec compute_member_role_permissions(
user_id(), maybe_channel_id(), member(), guild_state(), guild_data()
) -> permission().
compute_member_role_permissions(UserId, ChannelId, Member, State, Data) ->
-spec member_role_base_permissions(member(), guild_state(), guild_data()) -> base_permissions().
member_role_base_permissions(Member, State, Data) ->
case guild_id(State) of
undefined ->
0;
@@ -145,24 +166,10 @@ compute_member_role_permissions(UserId, ChannelId, Member, State, Data) ->
Permissions = aggregate_role_permissions_cached(
MemberRoles, RolePermsCache, Roles, BasePermissions
),
maybe_apply_admin_or_channel_overwrites(
Permissions, UserId, MemberRoles, ChannelId, GuildId, State
)
end.
-spec maybe_apply_admin_or_channel_overwrites(
permission(), user_id(), member_roles(), maybe_channel_id(), role_id(), guild_state()
) -> permission().
maybe_apply_admin_or_channel_overwrites(
Permissions, UserId, MemberRoles, ChannelId, GuildId, State
) ->
case permission_bits:has(Permissions, constants:administrator_permission()) of
true ->
?ALL_PERMISSIONS;
false ->
guild_permissions_overwrites:maybe_apply_channel_overwrites(
Permissions, UserId, MemberRoles, ChannelId, GuildId, State
)
case permission_bits:has(Permissions, constants:administrator_permission()) of
true -> ?ALL_PERMISSIONS;
false -> {Permissions, MemberRoles, GuildId}
end
end.
-spec resolve_member(user_id(), maybe_member(), guild_state()) -> maybe_member().
@@ -5,6 +5,7 @@
-export([
can_view_channel/4,
viewable_channel_ids/4,
can_view_channel_by_permissions/4,
can_view_channel_members/4,
can_manage_channel/3,
@@ -43,9 +44,78 @@
-spec can_view_channel(user_id(), channel_id(), maybe_member(), guild_state()) -> boolean().
can_view_channel(UserId, ChannelId, Member, State) ->
Base = guild_permissions:member_base_permissions(UserId, Member, State),
guild_virtual_channel_access:has_virtual_access(UserId, ChannelId, State) orelse
can_view_channel_by_permissions(UserId, ChannelId, Member, State) orelse
is_category_with_viewable_child(UserId, ChannelId, Member, State).
base_has_view(Base, UserId, ChannelId, State) orelse
is_category_with_viewable_child(UserId, ChannelId, Base, State).
-spec viewable_channel_ids(
user_id(), guild_permissions:base_permissions(), [map()], guild_state()
) -> #{channel_id() => true}.
viewable_channel_ids(UserId, Base, Channels, State) ->
Virtual = maps:from_keys(
guild_virtual_channel_access:get_virtual_channels_for_user(UserId, State), true
),
{Viewable, Undecided, ViewableParents} = lists:foldl(
fun(Channel, Acc) -> classify_channel(Channel, UserId, Base, Virtual, State, Acc) end,
{#{}, [], #{}},
Channels
),
lists:foldl(
fun(Id, Acc) ->
case maps:is_key(Id, ViewableParents) andalso is_category(Id, State) of
true -> Acc#{Id => true};
false -> Acc
end
end,
Viewable,
Undecided
).
-spec classify_channel(
map(),
user_id(),
guild_permissions:base_permissions(),
#{channel_id() => true},
guild_state(),
{#{channel_id() => true}, [channel_id()], map()}
) -> {#{channel_id() => true}, [channel_id()], map()}.
classify_channel(Channel, UserId, Base, Virtual, State, {Viewable, Undecided, Parents} = Acc) ->
case channel_list_id(Channel) of
undefined ->
Acc;
Id ->
case base_has_view(Base, UserId, Id, State) of
true ->
Parent = snowflake_id:parse_maybe(
maps:get(<<"parent_id">>, Channel, undefined)
),
{Viewable#{Id => true}, Undecided, Parents#{Parent => true}};
false when is_map_key(Id, Virtual) ->
{Viewable#{Id => true}, Undecided, Parents};
false ->
{Viewable, [Id | Undecided], Parents}
end
end.
-spec channel_list_id(map()) -> channel_id() | undefined.
channel_list_id(Channel) ->
snowflake_id:parse_maybe(maps:get(<<"id">>, Channel, undefined)).
-spec is_category(channel_id(), guild_state()) -> boolean().
is_category(ChannelId, State) ->
case find_channel_by_id(ChannelId, State) of
#{<<"type">> := 4} -> true;
_ -> false
end.
-spec base_has_view(
guild_permissions:base_permissions(), user_id(), channel_id(), guild_state()
) ->
boolean().
base_has_view(Base, UserId, ChannelId, State) ->
Perms = guild_permissions:channel_permissions(Base, UserId, ChannelId, State),
permission_bits:has(Perms, constants:view_channel_permission()).
-spec can_view_channel_by_permissions(user_id(), channel_id(), maybe_member(), guild_state()) ->
boolean().
@@ -246,40 +316,37 @@ find_channel_by_id(ChannelId, State) ->
undefined
end.
-spec is_category_with_viewable_child(user_id(), channel_id(), maybe_member(), guild_state()) ->
boolean().
is_category_with_viewable_child(UserId, ChannelId, Member, State) ->
case find_channel_by_id(ChannelId, State) of
#{<<"type">> := 4} -> any_child_viewable(UserId, ChannelId, Member, State);
_ -> false
end.
-spec is_category_with_viewable_child(
user_id(), channel_id(), guild_permissions:base_permissions(), guild_state()
) -> boolean().
is_category_with_viewable_child(UserId, ChannelId, Base, State) ->
is_category(ChannelId, State) andalso any_child_viewable(UserId, ChannelId, Base, State).
-spec any_child_viewable(user_id(), channel_id(), maybe_member(), guild_state()) -> boolean().
any_child_viewable(UserId, CategoryId, Member, State) ->
-spec any_child_viewable(
user_id(), channel_id(), guild_permissions:base_permissions(), guild_state()
) -> boolean().
any_child_viewable(UserId, CategoryId, Base, State) ->
case guild_permissions_common:resolve_data_map(State) of
undefined ->
false;
Data ->
any_child_viewable_in_data(UserId, CategoryId, Member, State, Data)
Channels = map_utils:ensure_list(maps:get(<<"channels">>, Data, [])),
lists:any(
fun(Channel) -> is_viewable_child(Channel, UserId, CategoryId, Base, State) end,
Channels
)
end.
-spec any_child_viewable_in_data(user_id(), channel_id(), maybe_member(), guild_state(), map()) ->
boolean().
any_child_viewable_in_data(UserId, CategoryId, Member, State, Data) ->
Channels = map_utils:ensure_list(maps:get(<<"channels">>, Data, [])),
lists:any(
fun(Channel) -> is_viewable_child(Channel, UserId, CategoryId, Member, State) end,
Channels
).
-spec is_viewable_child(map(), user_id(), channel_id(), maybe_member(), guild_state()) ->
boolean().
is_viewable_child(Channel, UserId, CategoryId, Member, State) ->
ParentId = snowflake_id:parse_maybe(maps:get(<<"parent_id">>, Channel, undefined)),
ChildId = snowflake_id:parse_maybe(maps:get(<<"id">>, Channel, undefined)),
case {ParentId, ChildId} of
{CategoryId, ResolvedChildId} when is_integer(ResolvedChildId) ->
can_view_channel_by_permissions(UserId, ResolvedChildId, Member, State);
-spec is_viewable_child(
map(), user_id(), channel_id(), guild_permissions:base_permissions(), guild_state()
) -> boolean().
is_viewable_child(Channel, UserId, CategoryId, Base, State) ->
case snowflake_id:parse_maybe(maps:get(<<"parent_id">>, Channel, undefined)) of
CategoryId ->
case channel_list_id(Channel) of
ChildId when is_integer(ChildId) -> base_has_view(Base, UserId, ChildId, State);
undefined -> false
end;
_ ->
false
end.
@@ -46,8 +46,13 @@ apply_channel_overwrites(BasePerms, UserId, MemberRoles, Channel, EveryoneRoleId
role_id()
) -> permission().
apply_cached_overwrites(BasePerms, UserId, MemberRoles, CachedOWs, EveryoneRoleId) ->
EveryonePerms = apply_cached_everyone(BasePerms, CachedOWs, EveryoneRoleId),
{RoleAllow, RoleDeny} = accumulate_cached_roles(MemberRoles, CachedOWs),
{EveryonePerms, RoleAllow, RoleDeny} = lists:foldl(
fun(CachedOW, Acc) ->
collect_cached_role_overwrite(CachedOW, MemberRoles, EveryoneRoleId, Acc)
end,
{BasePerms, 0, 0},
CachedOWs
),
RolePerms = permission_bits:apply_allow_deny(EveryonePerms, RoleAllow, RoleDeny),
apply_cached_user(RolePerms, CachedOWs, UserId).
@@ -161,46 +166,21 @@ apply_user_overwrite(Overwrite, UserId, Acc) when is_integer(UserId) ->
apply_user_overwrite(_Overwrite, _UserId, Acc) ->
Acc.
-spec apply_cached_everyone(
permission(), [{integer(), integer(), integer(), integer()}], role_id()
) -> permission().
apply_cached_everyone(BasePerms, CachedOWs, EveryoneRoleId) ->
lists:foldl(
fun
({OWId, 0, Allow, Deny}, Acc) when OWId =:= EveryoneRoleId ->
apply_allow_deny(Acc, Allow, Deny);
(_, Acc) ->
Acc
-spec collect_cached_role_overwrite(
term(), member_roles(), role_id(), {permission(), permission(), permission()}
) -> {permission(), permission(), permission()}.
collect_cached_role_overwrite({OWId, 0, Allow, Deny}, MemberRoles, EveryoneRoleId, {E, A, D}) ->
E1 =
case OWId =:= EveryoneRoleId of
true -> apply_allow_deny(E, Allow, Deny);
false -> E
end,
BasePerms,
CachedOWs
).
-spec accumulate_cached_roles(member_roles(), [{integer(), integer(), integer(), integer()}]) ->
{permission(), permission()}.
accumulate_cached_roles(MemberRoles, CachedOWs) ->
lists:foldl(
fun(RoleId, {AAcc, DAcc}) ->
accumulate_cached_role(RoleId, CachedOWs, {AAcc, DAcc})
end,
{0, 0},
MemberRoles
).
-spec accumulate_cached_role(
role_id(), [{integer(), integer(), integer(), integer()}], {permission(), permission()}
) -> {permission(), permission()}.
accumulate_cached_role(RoleId, CachedOWs, Acc) ->
lists:foldl(
fun
({OWId, 0, Allow, Deny}, {A, D}) when OWId =:= RoleId ->
{permission_bits:add(A, Allow), permission_bits:add(D, Deny)};
(_, AD) ->
AD
end,
Acc,
CachedOWs
).
case lists:member(OWId, MemberRoles) of
true -> {E1, permission_bits:add(A, Allow), permission_bits:add(D, Deny)};
false -> {E1, A, D}
end;
collect_cached_role_overwrite(_, _MemberRoles, _EveryoneRoleId, Acc) ->
Acc.
-spec apply_cached_user(
permission(), [{integer(), integer(), integer(), integer()}], user_id()
+407 -272
View File
@@ -4,9 +4,8 @@
-typing([eqwalizer]).
-export([handle_bus_presence/3, send_cached_presence_to_session/3]).
-export([broadcast_presence_update/3]).
-export([cached_presences/1, send_presence_lookup_to_session/4]).
-export([sync_online_status/2]).
-export([build_broadcast_snapshot/1]).
-export_type([guild_state/0, user_id/0]).
@@ -20,7 +19,6 @@
<<"members">>, members_normalized, <<"member_role_index">>, members_sorted_ids
]).
-define(PRESENCE_SNAPSHOT_TRIM_MEMBER_THRESHOLD_DEFAULT, 5000).
-define(SESSIONS_PROJECTION_CACHE, presence_snapshot_sessions_cache).
-ifdef(TEST).
-include_lib("eunit/include/eunit.hrl").
@@ -117,47 +115,87 @@ spawn_presence_broadcast(UserId, OldPresence, PresenceMap, OldState, NewState) -
PresenceMap
),
{Pid, NewState2} = guild_broadcaster:ensure(NewState1),
{NewSnap, NewState3} = build_broadcast_snapshot_cached(NewState2),
guild_broadcaster:cast_presence(Pid, UserId, PresenceMap, #{}, NewSnap),
NewState3.
ok = cast_presence_update(Pid, UserId, PresenceMap, NewState2),
NewState2.
-spec build_broadcast_snapshot(guild_state()) -> map().
build_broadcast_snapshot(State) ->
{Snapshot, _State} = build_broadcast_snapshot_cached(State),
Snapshot.
-spec cast_presence_update(pid() | undefined, user_id(), map(), guild_state()) -> ok.
cast_presence_update(BroadcasterPid, UserId, PresenceMap, State) when is_pid(BroadcasterPid) ->
case safe_presence_update_recipients(UserId, presence_view(State)) of
{GuildId, [_ | _] = Pids} ->
PresenceUpdate = PresenceMap#{<<"guild_id">> => integer_to_binary(GuildId)},
_ = guild_broadcaster:cast_event(
BroadcasterPid, presence_update, PresenceUpdate, Pids
),
ok;
_ ->
ok
end;
cast_presence_update(_BroadcasterPid, _UserId, _PresenceMap, _State) ->
ok.
-spec build_broadcast_snapshot_cached(guild_state()) -> {map(), guild_state()}.
build_broadcast_snapshot_cached(State) ->
maybe_trim_snapshot(build_base_snapshot(State), State).
-spec build_base_snapshot(guild_state()) -> map().
build_base_snapshot(State) ->
Keys = [
id,
data,
sessions,
member_subscriptions,
member_presence,
role_overrides,
permission_overwrites
],
lists:foldl(
fun(K, Acc) ->
put_existing_key(K, State, Acc)
end,
#{},
Keys
).
-spec maybe_trim_snapshot(map(), guild_state()) -> {map(), guild_state()}.
maybe_trim_snapshot(Snapshot, State) ->
case should_trim_snapshot(State) of
true -> trim_snapshot(Snapshot, State);
false -> {Snapshot, State}
-spec safe_presence_update_recipients(user_id(), guild_state()) -> {integer(), [pid()]} | none.
safe_presence_update_recipients(UserId, View) ->
try
presence_update_recipients(UserId, View)
catch
Class:Reason:Stack ->
logger:warning(
"guild presence_update recipients error: ~p:~p ~p",
[Class, Reason, Stack]
),
none
end.
-spec should_trim_snapshot(guild_state()) -> boolean().
should_trim_snapshot(State) ->
-spec presence_update_recipients(user_id(), guild_state()) -> {integer(), [pid()]} | none.
presence_update_recipients(UserId, View) ->
case {find_member_by_user_id(UserId, View), guild_id(View)} of
{undefined, _} ->
none;
{_Member, GuildId} when is_integer(GuildId), GuildId > 0 ->
{GuildId, subscribed_session_pids(UserId, View)};
_ ->
none
end.
-spec subscribed_session_pids(user_id(), guild_state()) -> [pid()].
subscribed_session_pids(UserId, View) ->
MemberSubs = maps:get(member_subscriptions, View, guild_subscriptions:init_state()),
case guild_subscriptions:get_subscribed_sessions(UserId, MemberSubs) of
[] ->
[];
SubscribedSessionIds ->
Sessions = maps:get(sessions, View, #{}),
TargetChannelMap = guild_presence_sync:get_user_viewable_channel_map(
UserId, Sessions, View
),
{ValidSessionIds, _InvalidSessionIds} =
guild_presence_sync:partition_subscribed_sessions(
SubscribedSessionIds, Sessions, TargetChannelMap, UserId, View
),
guild_presence_sync:session_pids(ValidSessionIds, Sessions)
end.
-spec presence_view(guild_state()) -> guild_state().
presence_view(State) ->
View = maps:with(
[
id,
data,
sessions,
member_subscriptions,
member_presence,
role_overrides,
permission_overwrites
],
State
),
case should_trim_view(State) of
true -> trim_view_data(View);
false -> View
end.
-spec should_trim_view(guild_state()) -> boolean().
should_trim_view(State) ->
presence_snapshot_trim_enabled() andalso member_count_at_or_above_threshold(State).
-spec member_count_at_or_above_threshold(guild_state()) -> boolean().
@@ -189,99 +227,37 @@ presence_snapshot_trim_member_threshold() ->
_ -> ?PRESENCE_SNAPSHOT_TRIM_MEMBER_THRESHOLD_DEFAULT
end.
-spec trim_snapshot(map(), guild_state()) -> {map(), guild_state()}.
trim_snapshot(Snapshot, State) ->
trim_snapshot_sessions(trim_snapshot_data(Snapshot), State).
-spec trim_snapshot_data(map()) -> map().
trim_snapshot_data(#{data := Data} = Snapshot) when is_map(Data) ->
Snapshot#{data => maps:without(?HEAVY_MEMBER_DATA_KEYS, Data)};
trim_snapshot_data(Snapshot) ->
Snapshot.
-spec trim_snapshot_sessions(map(), guild_state()) -> {map(), guild_state()}.
trim_snapshot_sessions(#{sessions := Sessions} = Snapshot, State) when is_map(Sessions) ->
{Projected, NewState} = cached_projected_sessions(Sessions, State),
{Snapshot#{sessions => Projected}, NewState};
trim_snapshot_sessions(Snapshot, State) ->
{Snapshot, State}.
%% Keyed by the sessions map itself, so every writer of that map invalidates the
%% projection without having to know this cache exists.
-spec cached_projected_sessions(map(), guild_state()) -> {map(), guild_state()}.
cached_projected_sessions(Sessions, State) ->
case maps:get(?SESSIONS_PROJECTION_CACHE, State, undefined) of
{Sessions, Projected} when is_map(Projected) -> {Projected, State};
_ -> store_projected_sessions(Sessions, State)
end.
-spec store_projected_sessions(map(), guild_state()) -> {map(), guild_state()}.
store_projected_sessions(Sessions, State) ->
Projected = project_sessions(Sessions),
{Projected, State#{?SESSIONS_PROJECTION_CACHE => {Sessions, Projected}}}.
-spec project_sessions(map()) -> map().
project_sessions(Sessions) ->
maps:map(
fun(_SessionId, SessionData) -> project_session(SessionData) end,
Sessions
).
-spec project_session(term()) -> term().
project_session(SessionData) when is_map(SessionData) ->
maps:with([user_id, pid, viewable_channels], SessionData);
project_session(SessionData) ->
SessionData.
-spec put_existing_key(atom(), guild_state(), map()) -> map().
put_existing_key(Key, State, Acc) ->
case maps:find(Key, State) of
{ok, Value} -> Acc#{Key => Value};
error -> Acc
end.
-spec broadcast_presence_update(user_id(), map(), guild_state()) -> ok.
broadcast_presence_update(UserId, Payload, State) ->
case find_member_by_user_id(UserId, State) of
undefined -> ok;
_Member -> broadcast_presence_update_impl(UserId, Payload, State)
end.
-spec broadcast_presence_update_impl(user_id(), map(), guild_state()) -> ok.
broadcast_presence_update_impl(UserId, Payload, State) ->
case guild_id(State) of
GuildId when is_integer(GuildId), GuildId > 0 ->
PresenceUpdate = Payload#{<<"guild_id">> => integer_to_binary(GuildId)},
Sessions = maps:get(sessions, State, #{}),
MemberSubs = maps:get(
member_subscriptions, State, guild_subscriptions:init_state()
),
SubscribedSessionIds = guild_subscriptions:get_subscribed_sessions(
UserId, MemberSubs
),
TargetChannelMap = guild_presence_sync:get_user_viewable_channel_map(
UserId, Sessions, State
),
{ValidSessionIds, InvalidSessionIds} =
guild_presence_sync:partition_subscribed_sessions(
SubscribedSessionIds, Sessions, TargetChannelMap, UserId, State
),
FinalState = guild_presence_sync:remove_invalid_subscriptions(
InvalidSessionIds, UserId, State
),
FinalSessions = maps:get(sessions, FinalState, #{}),
guild_presence_sync:dispatch_to_valid_sessions(
ValidSessionIds, FinalSessions, PresenceUpdate, GuildId
);
_ ->
ok
end.
-spec trim_view_data(guild_state()) -> guild_state().
trim_view_data(#{data := Data} = View) when is_map(Data) ->
View#{data => maps:without(?HEAVY_MEMBER_DATA_KEYS, Data)};
trim_view_data(View) ->
View.
-spec send_cached_presence_to_session(user_id(), binary(), guild_state()) -> guild_state().
send_cached_presence_to_session(UserId, SessionId, State) ->
case safe_presence_cache_get(UserId) of
{ok, Payload} -> send_presence_payload_to_session(UserId, SessionId, Payload, State);
_ -> State
send_presence_lookup_to_session(UserId, SessionId, safe_presence_cache_get(UserId), State).
-spec send_presence_lookup_to_session(
user_id(), binary(), {ok, map()} | not_found, guild_state()
) ->
guild_state().
send_presence_lookup_to_session(UserId, SessionId, {ok, Payload}, State) ->
send_presence_payload_to_session(UserId, SessionId, Payload, State);
send_presence_lookup_to_session(_UserId, _SessionId, not_found, State) ->
State.
-spec cached_presences([user_id()]) -> #{user_id() => {ok, map()} | not_found}.
cached_presences([]) ->
#{};
cached_presences(UserIds) ->
Found = safe_presence_cache_bulk_get(UserIds),
maps:from_list([{UserId, presence_lookup(UserId, Found)} || UserId <- UserIds]).
-spec presence_lookup(user_id(), #{integer() => map()}) -> {ok, map()} | not_found.
presence_lookup(UserId, Found) ->
case maps:find(UserId, Found) of
{ok, Payload} -> {ok, Payload};
error -> not_found
end.
-spec send_presence_payload_to_session(user_id(), binary(), map(), guild_state()) ->
@@ -341,6 +317,14 @@ safe_presence_cache_get(UserId) ->
_:_ -> not_found
end.
-spec safe_presence_cache_bulk_get([user_id()]) -> #{integer() => map()}.
safe_presence_cache_bulk_get(UserIds) ->
try
presence_cache:bulk_get_map(UserIds)
catch
_:_ -> #{}
end.
-spec guild_id(guild_state()) -> integer() | undefined.
guild_id(State) ->
snowflake_id:parse_optional(maps:get(id, State, undefined)).
@@ -459,10 +443,64 @@ presence_test_state() ->
member_list_subscriptions => guild_member_list_subs:new()
}.
snapshot_trim_test_state(Tab) ->
online_payload() ->
#{
<<"status">> => <<"online">>,
<<"mobile">> => true,
<<"afk">> => false,
<<"user">> => #{<<"id">> => <<"1">>, <<"username">> => <<"Alpha">>}
}.
idle_session() ->
receive
stop -> ok
after 60000 -> ok
end.
handle_bus_presence_casts_presence_update_to_broadcaster_test() ->
Subscriber = spawn(fun idle_session/0),
try
MemberSubs = guild_subscriptions:subscribe(
<<"s2">>, 1, guild_subscriptions:init_state()
),
State = (presence_test_state())#{
broadcaster_pid => self(),
member_subscriptions => MemberSubs,
sessions => #{
<<"s1">> => #{user_id => 1, pid => self(), viewable_channels => #{100 => true}},
<<"s2">> => #{
user_id => 2, pid => Subscriber, viewable_channels => #{100 => true}
}
}
},
{noreply, _NewState} = handle_bus_presence(1, online_payload(), State),
receive
{'$gen_cast', {event_broadcast, presence_update, Update, Pids}} ->
?assertEqual([Subscriber], Pids),
?assertEqual(<<"42">>, maps:get(<<"guild_id">>, Update)),
?assertEqual(<<"online">>, maps:get(<<"status">>, Update))
after 1000 ->
?assert(false)
end
after
exit(Subscriber, kill)
end.
handle_bus_presence_skips_broadcaster_without_subscribers_test() ->
State = (presence_test_state())#{broadcaster_pid => self()},
{noreply, _NewState} = handle_bus_presence(1, online_payload(), State),
receive
{'$gen_cast', {event_broadcast, presence_update, _, _}} -> ?assert(false)
after 100 ->
ok
end.
view_trim_test_state(Tab) ->
#{
id => 42,
member_count => 10,
voice_states => #{},
virtual_channel_access => #{1 => sets:from_list([100])},
data => #{
members_ets => Tab,
<<"members">> => #{1 => #{<<"user">> => #{<<"id">> => <<"1">>}}},
@@ -483,195 +521,292 @@ snapshot_trim_test_state(Tab) ->
}
}.
build_broadcast_snapshot_trims_by_default_test() ->
presence_view_trims_member_data_by_default_test() ->
application:set_env(fluxer_gateway, presence_snapshot_trim_member_threshold, 2),
Tab = ets:new(snapshot_trim_members, [set, public]),
Tab = ets:new(view_trim_members, [set, public]),
try
Snap = build_broadcast_snapshot(snapshot_trim_test_state(Tab)),
SnapData = maps:get(data, Snap),
?assertNot(maps:is_key(<<"members">>, SnapData)),
?assertNot(maps:is_key(members_normalized, SnapData)),
?assertNot(maps:is_key(<<"member_role_index">>, SnapData)),
?assertNot(maps:is_key(members_sorted_ids, SnapData)),
?assertEqual(Tab, maps:get(members_ets, SnapData)),
?assert(maps:is_key(<<"channels">>, SnapData)),
SnapSession = maps:get(<<"s1">>, maps:get(sessions, Snap)),
?assertEqual([pid, user_id, viewable_channels], lists:sort(maps:keys(SnapSession))),
?assertNot(maps:is_key(voice_states, Snap)),
?assertNot(maps:is_key(member_count, Snap))
State = view_trim_test_state(Tab),
View = presence_view(State),
ViewData = maps:get(data, View),
?assertEqual(maps:without(?HEAVY_MEMBER_DATA_KEYS, maps:get(data, State)), ViewData),
?assertEqual(Tab, maps:get(members_ets, ViewData)),
?assertEqual(maps:get(sessions, State), maps:get(sessions, View)),
?assertEqual([data, id, sessions], lists:sort(maps:keys(View)))
after
application:unset_env(fluxer_gateway, presence_snapshot_trim_member_threshold),
ets:delete(Tab)
end.
build_broadcast_snapshot_no_trim_when_disabled_test() ->
presence_view_keeps_member_data_when_trim_disabled_test() ->
application:set_env(fluxer_gateway, presence_snapshot_trim_enabled, false),
Tab = ets:new(snapshot_notrim_members, [set, public]),
Tab = ets:new(view_notrim_members, [set, public]),
try
State = (snapshot_trim_test_state(Tab))#{member_count => 100000},
Snap = build_broadcast_snapshot(State),
?assert(maps:is_key(<<"members">>, maps:get(data, Snap))),
SnapSession = maps:get(<<"s1">>, maps:get(sessions, Snap)),
?assert(maps:is_key(active_guilds, SnapSession))
State = (view_trim_test_state(Tab))#{member_count => 100000},
?assertEqual(maps:get(data, State), maps:get(data, presence_view(State)))
after
application:unset_env(fluxer_gateway, presence_snapshot_trim_enabled),
ets:delete(Tab)
end.
build_broadcast_snapshot_no_trim_below_threshold_test() ->
presence_view_keeps_member_data_below_threshold_test() ->
application:set_env(fluxer_gateway, presence_snapshot_trim_member_threshold, 50000),
Tab = ets:new(snapshot_below_members, [set, public]),
Tab = ets:new(view_below_members, [set, public]),
try
Snap = build_broadcast_snapshot(snapshot_trim_test_state(Tab)),
?assert(maps:is_key(<<"members">>, maps:get(data, Snap)))
State = view_trim_test_state(Tab),
?assertEqual(maps:get(data, State), maps:get(data, presence_view(State)))
after
application:unset_env(fluxer_gateway, presence_snapshot_trim_member_threshold),
ets:delete(Tab)
end.
reference_snapshot_keys() ->
[
id,
data,
sessions,
member_subscriptions,
member_presence,
role_overrides,
permission_overwrites
].
reference_project_session(SessionData) when is_map(SessionData) ->
maps:with([user_id, pid, viewable_channels], SessionData);
reference_project_session(SessionData) ->
SessionData.
reference_project_sessions(Sessions) ->
maps:map(
fun(_SessionId, SessionData) -> reference_project_session(SessionData) end,
Sessions
).
reference_trim_snapshot_sessions(#{sessions := Sessions} = Snapshot) when is_map(Sessions) ->
Snapshot#{sessions => reference_project_sessions(Sessions)};
reference_trim_snapshot_sessions(Snapshot) ->
Snapshot.
reference_maybe_trim_snapshot(true, Snapshot) ->
reference_trim_snapshot_sessions(trim_snapshot_data(Snapshot));
reference_maybe_trim_snapshot(false, Snapshot) ->
Snapshot.
reference_trim_snapshot(#{data := Data, sessions := Sessions} = Snapshot) ->
Snapshot#{
data => maps:without(?HEAVY_MEMBER_DATA_KEYS, Data),
sessions => maps:map(fun(_, S) -> reference_project_session(S) end, Sessions)
}.
reference_broadcast_snapshot(State) ->
reference_maybe_trim_snapshot(should_trim_snapshot(State), build_base_snapshot(State)).
Base = lists:foldl(
fun(K, Acc) ->
case maps:find(K, State) of
{ok, V} -> Acc#{K => V};
error -> Acc
end
end,
#{},
reference_snapshot_keys()
),
case should_trim_view(State) of
true -> reference_trim_snapshot(Base);
false -> Base
end.
cache_test_session(UserId, ViewableChannels) ->
reference_broadcaster_recipients(UserId, State) ->
Snapshot = reference_broadcast_snapshot(State),
case {find_member_by_user_id(UserId, Snapshot), guild_id(Snapshot)} of
{undefined, _} ->
none;
{_, GuildId} when is_integer(GuildId), GuildId > 0 ->
Sessions = maps:get(sessions, Snapshot, #{}),
MemberSubs = maps:get(
member_subscriptions, Snapshot, guild_subscriptions:init_state()
),
SubscribedSessionIds = guild_subscriptions:get_subscribed_sessions(
UserId, MemberSubs
),
TargetChannelMap = guild_presence_sync:get_user_viewable_channel_map(
UserId, Sessions, Snapshot
),
{Valid, _Invalid} = guild_presence_sync:partition_subscribed_sessions(
SubscribedSessionIds, Sessions, TargetChannelMap, UserId, Snapshot
),
{GuildId, [
P
|| Sid <- Valid, #{pid := P} <- [maps:get(Sid, Sessions, #{})], is_pid(P)
]};
_ ->
none
end.
-define(DIFF_ROLE_A, 201).
-define(DIFF_ROLE_B, 202).
-define(DIFF_CHANNELS, [500, 501, 502, 503, 504]).
diff_user_ids() ->
lists:seq(10, 29).
diff_view() ->
constants:view_channel_permission().
diff_overwrite(Id, Type) ->
#{
user_id => UserId,
pid => self(),
viewable_channels => ViewableChannels,
active_guilds => [42],
pending_connect => false
<<"id">> => integer_to_binary(Id),
<<"type">> => Type,
<<"allow">> => integer_to_binary(diff_view()),
<<"deny">> => <<"0">>
}.
cache_test_state(Sessions) ->
diff_channel(Id, Overwrites) ->
#{
<<"id">> => integer_to_binary(Id),
<<"type">> => 0,
<<"permission_overwrites">> => Overwrites
}.
diff_channels(Seed) ->
Public =
case Seed rem 2 of
0 -> [];
1 -> [diff_overwrite(42, 0)]
end,
[
diff_channel(500, [diff_overwrite(?DIFF_ROLE_A, 0)]),
diff_channel(501, [diff_overwrite(?DIFF_ROLE_B, 0)]),
diff_channel(502, [diff_overwrite(?DIFF_ROLE_A, 0), diff_overwrite(?DIFF_ROLE_B, 0)]),
diff_channel(503, [diff_overwrite(10 + Seed rem 20, 1)]),
diff_channel(504, Public)
].
diff_role(Id) ->
#{
<<"id">> => integer_to_binary(Id),
<<"name">> => integer_to_binary(Id),
<<"position">> => Id - 42,
<<"permissions">> => <<"0">>
}.
diff_member(UserId) ->
Roles = [integer_to_binary(R) || R <- [?DIFF_ROLE_A, ?DIFF_ROLE_B], rand:uniform(2) =:= 1],
#{
<<"user">> => #{
<<"id">> => integer_to_binary(UserId),
<<"username">> => integer_to_binary(UserId)
},
<<"roles">> => Roles
}.
diff_data(Seed, WithEts, Tab) ->
Data = guild_data_index:normalize_data(#{
<<"guild">> => #{<<"id">> => <<"42">>, <<"owner_id">> => <<"1">>},
<<"roles">> => [diff_role(42), diff_role(?DIFF_ROLE_A), diff_role(?DIFF_ROLE_B)],
<<"members">> => [diff_member(U) || U <- diff_user_ids()],
<<"channels">> => diff_channels(Seed)
}),
case WithEts of
true ->
true = ets:insert(Tab, maps:to_list(guild_data_index:member_map(Data))),
Data#{members_ets => Tab};
false ->
Data
end.
diff_random_user() ->
lists:nth(rand:uniform(20), diff_user_ids()).
diff_session(Pids) ->
Base = #{
user_id => diff_random_user(),
active_guilds => [42],
pending_connect => false,
user_roles => []
},
WithPid =
case rand:uniform(8) of
1 -> Base;
_ -> Base#{pid => lists:nth(rand:uniform(length(Pids)), Pids)}
end,
case rand:uniform(3) of
1 ->
WithPid;
_ ->
Viewable = maps:from_list([{C, true} || C <- ?DIFF_CHANNELS, rand:uniform(3) =:= 1]),
WithPid#{viewable_channels => Viewable}
end.
diff_member_subscriptions(SessionIds) ->
Candidates = [<<"gone">> | SessionIds],
lists:foldl(
fun(UserId, Subs) ->
lists:foldl(
fun(Sid, Acc) -> guild_subscriptions:subscribe(Sid, UserId, Acc) end,
Subs,
[Sid || Sid <- Candidates, rand:uniform(3) =:= 1]
)
end,
guild_subscriptions:init_state(),
diff_user_ids()
).
diff_state(Seed, WithEts, Tab, Pids) ->
SessionIds = [<<"s", (integer_to_binary(I))/binary>> || I <- lists:seq(1, 16)],
Sessions = maps:from_list([{Sid, diff_session(Pids)} || Sid <- SessionIds]),
#{
id => 42,
member_count => 10,
data => #{<<"channels">> => [], <<"members">> => #{}},
sessions => Sessions
member_count => 20,
data => diff_data(Seed, WithEts, Tab),
sessions => Sessions,
member_subscriptions => diff_member_subscriptions(SessionIds),
presence_subscriptions => #{10 => 3, 11 => 1},
virtual_channel_access => #{
10 => sets:from_list([500, 501]), 11 => sets:from_list([504])
},
voice_states => #{}
}.
assert_cached_snapshot_matches_reference(Sessions, State) ->
NextState = State#{sessions => Sessions},
{Snapshot, CachedState} = build_broadcast_snapshot_cached(NextState),
?assertEqual(reference_broadcast_snapshot(NextState), Snapshot),
?assertNot(maps:is_key(?SESSIONS_PROJECTION_CACHE, Snapshot)),
CachedState.
session_projection_mutation_sequence() ->
A = cache_test_session(1, #{100 => true}),
B = cache_test_session(2, #{100 => true, 200 => true}),
S1 = #{<<"a">> => A},
S2 = S1#{<<"b">> => B},
S3 = S2#{<<"a">> => cache_test_session(1, #{100 => true, 300 => true})},
S4 = maps:remove(<<"b">>, S3),
S5 = S4#{<<"a">> => (maps:get(<<"a">>, S4))#{pending_connect => true}},
S6 = S5#{<<"c">> => not_a_map},
[#{}, S1, S1, S2, S3, S4, S5, S6, S1, #{}].
cached_projection_matches_reference_across_mutations_test() ->
with_trim(on, Fun) ->
application:set_env(fluxer_gateway, presence_snapshot_trim_member_threshold, 2),
try
lists:foldl(
fun assert_cached_snapshot_matches_reference/2,
cache_test_state(#{}),
session_projection_mutation_sequence()
)
Fun()
after
application:unset_env(fluxer_gateway, presence_snapshot_trim_member_threshold)
end.
cached_projection_is_reused_when_sessions_unchanged_test() ->
application:set_env(fluxer_gateway, presence_snapshot_trim_member_threshold, 2),
try
Sessions = #{<<"a">> => cache_test_session(1, #{100 => true})},
{Snap1, State1} = build_broadcast_snapshot_cached(cache_test_state(Sessions)),
{Snap2, State2} = build_broadcast_snapshot_cached(State1),
Projected1 = maps:get(sessions, Snap1),
Projected2 = maps:get(sessions, Snap2),
?assertEqual(reference_project_sessions(Sessions), Projected1),
?assertEqual(Projected1, Projected2),
?assert(erts_debug:same(Projected1, Projected2)),
?assertEqual({Sessions, Projected1}, maps:get(?SESSIONS_PROJECTION_CACHE, State2))
after
application:unset_env(fluxer_gateway, presence_snapshot_trim_member_threshold)
end.
cached_projection_ignores_foreign_cache_entry_test() ->
application:set_env(fluxer_gateway, presence_snapshot_trim_member_threshold, 2),
try
Sessions = #{<<"a">> => cache_test_session(1, #{100 => true})},
Stale = #{<<"z">> => cache_test_session(9, #{999 => true})},
Foreign = {Stale, reference_project_sessions(Stale)},
State = (cache_test_state(Sessions))#{?SESSIONS_PROJECTION_CACHE => Foreign},
{Snapshot, _NewState} = build_broadcast_snapshot_cached(State),
?assertEqual(reference_broadcast_snapshot(State), Snapshot),
?assertEqual(reference_project_sessions(Sessions), maps:get(sessions, Snapshot))
after
application:unset_env(fluxer_gateway, presence_snapshot_trim_member_threshold)
end.
cached_projection_ignores_corrupt_cache_entry_test() ->
application:set_env(fluxer_gateway, presence_snapshot_trim_member_threshold, 2),
try
Sessions = #{<<"a">> => cache_test_session(1, #{100 => true})},
State = (cache_test_state(Sessions))#{?SESSIONS_PROJECTION_CACHE => {Sessions, junk}},
{Snapshot, _NewState} = build_broadcast_snapshot_cached(State),
?assertEqual(reference_project_sessions(Sessions), maps:get(sessions, Snapshot))
after
application:unset_env(fluxer_gateway, presence_snapshot_trim_member_threshold)
end.
cached_projection_not_stored_when_trim_disabled_test() ->
end;
with_trim(off, Fun) ->
application:set_env(fluxer_gateway, presence_snapshot_trim_enabled, false),
try
Sessions = #{<<"a">> => cache_test_session(1, #{100 => true})},
State = cache_test_state(Sessions),
{Snapshot, NewState} = build_broadcast_snapshot_cached(State),
?assertEqual(reference_broadcast_snapshot(State), Snapshot),
?assertEqual(Sessions, maps:get(sessions, Snapshot)),
?assertNot(maps:is_key(?SESSIONS_PROJECTION_CACHE, NewState))
Fun()
after
application:unset_env(fluxer_gateway, presence_snapshot_trim_enabled)
end.
handle_bus_presence_threads_projection_cache_test() ->
application:set_env(fluxer_gateway, presence_snapshot_trim_member_threshold, 2),
Sessions = #{<<"s1">> => cache_test_session(1, #{100 => true})},
State = (presence_test_state())#{member_count => 10, sessions => Sessions},
Payload = #{
<<"status">> => <<"online">>,
<<"mobile">> => true,
<<"afk">> => false,
<<"user">> => #{<<"id">> => <<"1">>, <<"username">> => <<"Alpha">>}
},
diff_compare_users(State) ->
lists:map(
fun(UserId) ->
Expected = reference_broadcaster_recipients(UserId, State),
?assertEqual(
{UserId, Expected},
{UserId, presence_update_recipients(UserId, presence_view(State))}
),
Subscribed = guild_subscriptions:get_subscribed_sessions(
UserId, maps:get(member_subscriptions, State)
),
{Expected, length(Subscribed)}
end,
[99 | diff_user_ids()]
).
diff_run(Seed, WithEts, Trim, Pids) ->
Tab = ets:new(recipients_diff_members, [set, public]),
try
{noreply, NewState} = handle_bus_presence(1, Payload, State),
?assertEqual(
{Sessions, reference_project_sessions(Sessions)},
maps:get(?SESSIONS_PROJECTION_CACHE, NewState)
)
_ = rand:seed(exsss, {Seed, 7, 11}),
State = diff_state(Seed, WithEts, Tab, Pids),
with_trim(Trim, fun() -> diff_compare_users(State) end)
after
application:unset_env(fluxer_gateway, presence_snapshot_trim_member_threshold)
ets:delete(Tab)
end.
presence_update_recipients_match_broadcaster_snapshot_path_test_() ->
{timeout, 120, fun() ->
Pids = [spawn(fun idle_session/0) || _ <- lists:seq(1, 6)],
try
Results = lists:append([
diff_run(Seed, WithEts, Trim, Pids)
|| Seed <- lists:seq(1, 60), WithEts <- [true, false], Trim <- [on, off]
]),
NonEmpty = [R || {{_, [_ | _]}, _} = R <- Results],
Filtered = [R || {{_, Ps}, Subs} = R <- Results, length(Ps) < Subs],
Unknown = [R || {none, _} = R <- Results],
?assert(length(NonEmpty) > 100),
?assert(length(Filtered) > 100),
?assert(length(Unknown) > 100)
after
[exit(P, kill) || P <- Pids]
end
end}.
-endif.
@@ -9,18 +9,20 @@
start_async/1,
reconcile_user/2,
maybe_schedule_user_repair/2,
apply_reconcile_result/2,
find_mismatches/3,
apply_mismatches/2,
connected_user_ids_list/1,
reconcile_action/3
]).
-export_type([guild_state/0, user_id/0, presence/0, presence_by_id/0]).
-export_type([guild_state/0, user_id/0, presence/0, presence_by_id/0, mismatch/0]).
-type guild_state() :: map().
-type user_id() :: integer().
-type presence() :: map().
-type presence_by_id() :: #{user_id() => presence()}.
-type display() :: {binary(), boolean(), boolean(), term()}.
-type mismatch() :: {user_id(), presence(), display()}.
-define(DEFAULT_INTERVAL_MS, 30000).
-define(MIN_INTERVAL_MS, 5000).
@@ -41,19 +43,93 @@ interval_ms() ->
start_async(State) ->
case connected_user_ids_list(State) of
[] -> ok;
UserIds -> spawn_reconcile_fetch(UserIds, self())
UserIds -> spawn_reconcile_fetch(UserIds, member_presence_tab(State), self())
end.
-spec spawn_reconcile_fetch([user_id(), ...], pid()) -> ok.
spawn_reconcile_fetch(UserIds, Parent) ->
_ = spawn(fun() -> fetch_and_reply(UserIds, Parent) end),
-spec spawn_reconcile_fetch([user_id(), ...], ets:tid() | undefined, pid()) -> ok.
spawn_reconcile_fetch(UserIds, Tab, Parent) ->
_ = spawn(fun() -> fetch_and_reply(UserIds, Tab, Parent) end),
ok.
-spec fetch_and_reply([user_id(), ...], pid()) -> ok.
fetch_and_reply(UserIds, Parent) ->
Parent ! {presence_reconcile_apply, authoritative_presence_map(UserIds)},
-spec fetch_and_reply([user_id(), ...], ets:tid() | undefined, pid()) -> ok.
fetch_and_reply(UserIds, Tab, Parent) ->
Parent ! {presence_reconcile_apply, persistent_mismatches(UserIds, Tab, ?REPAIR_DELAY_MS)},
ok.
-spec persistent_mismatches([user_id()], ets:tid() | undefined, non_neg_integer()) ->
[mismatch()].
persistent_mismatches(UserIds, Tab, ConfirmDelayMs) ->
case observe_mismatches(UserIds, Tab) of
[] ->
[];
First ->
timer:sleep(ConfirmDelayMs),
Second = observe_mismatches([UserId || {UserId, _, _} <- First], Tab),
confirmed_mismatches(First, Second)
end.
-spec observe_mismatches([user_id()], ets:tid() | undefined) -> [mismatch()].
observe_mismatches(UserIds, Tab) ->
Authoritative = authoritative_presence_map(UserIds),
try
find_mismatches(UserIds, Authoritative, Tab)
catch
error:badarg -> []
end.
-spec find_mismatches([user_id()], presence_by_id(), ets:tid() | undefined) -> [mismatch()].
find_mismatches(UserIds, Authoritative, Tab) ->
lists:filtermap(
fun(UserId) ->
Seen = tab_display(UserId, Tab),
Presence = lookup_presence_value(UserId, Authoritative),
case desired_display(Presence) =:= Seen of
true -> false;
false -> {true, {UserId, replay_payload(Presence), Seen}}
end
end,
UserIds
).
-spec confirmed_mismatches([mismatch()], [mismatch()]) -> [mismatch()].
confirmed_mismatches(First, Second) ->
Keys = sets:from_list([mismatch_key(M) || M <- First], [{version, 2}]),
[M || M <- Second, sets:is_element(mismatch_key(M), Keys)].
-spec mismatch_key(mismatch()) -> {user_id(), display(), display()}.
mismatch_key({UserId, Payload, Seen}) ->
{UserId, display_fields(Payload), Seen}.
-spec apply_mismatches([mismatch()], guild_state()) -> guild_state().
apply_mismatches(Mismatches, State) ->
Connected = guild_member_list_connected:connected_session_user_ids(State),
lists:foldl(
fun({UserId, Payload, Seen}, Acc) ->
case
sets:is_element(UserId, Connected) andalso
current_display(UserId, Acc) =:= Seen
of
true -> replay_presence(UserId, Payload, Acc);
false -> Acc
end
end,
State,
Mismatches
).
-spec member_presence_tab(guild_state()) -> ets:tid() | undefined.
member_presence_tab(State) ->
case maps:get(member_presence, State, undefined) of
Tab when is_reference(Tab) -> Tab;
_ -> undefined
end.
-spec tab_display(user_id(), ets:tid() | undefined) -> display().
tab_display(_UserId, undefined) ->
offline_display();
tab_display(UserId, Tab) ->
display_fields(guild_state_member:lookup_presence(Tab, UserId)).
-spec reconcile_user(term(), guild_state()) -> guild_state().
reconcile_user(UserId, State) when is_integer(UserId), UserId > 0 ->
case is_connected(UserId, State) of
@@ -91,18 +167,6 @@ presence_row_exists(Tab, UserId) ->
error:badarg -> false
end.
-spec apply_reconcile_result(map(), guild_state()) -> guild_state().
apply_reconcile_result(PresenceById, State) when is_map(PresenceById) ->
lists:foldl(
fun(UserId, Acc) ->
apply_user_reconcile(UserId, lookup_presence_value(UserId, PresenceById), Acc)
end,
State,
connected_user_ids_list(State)
);
apply_reconcile_result(_PresenceById, State) ->
State.
-spec lookup_presence_value(user_id(), map()) -> presence() | undefined.
lookup_presence_value(UserId, PresenceById) ->
case maps:get(UserId, PresenceById, undefined) of
@@ -142,10 +206,7 @@ desired_display(Payload) -> display_fields(Payload).
-spec current_display(user_id(), guild_state()) -> display().
current_display(UserId, State) ->
case maps:get(member_presence, State, undefined) of
undefined -> offline_display();
Tab -> display_fields(guild_state_member:lookup_presence(Tab, UserId))
end.
tab_display(UserId, member_presence_tab(State)).
-spec display_fields(presence()) -> display().
display_fields(Presence) ->
@@ -258,54 +319,142 @@ presence_user_id_test() ->
?assertEqual(7, presence_user_id(#{<<"user">> => #{<<"id">> => <<"7">>}})),
?assertEqual(undefined, presence_user_id(#{})).
apply_reconcile_result_repairs_connected_offline_member_test() ->
GuildId = 4242,
UserId = 99,
Engine = guild_member_list_engine:new(),
try
State = guild_test_state(GuildId, UserId, Engine),
apply_mismatches_repairs_connected_offline_member_test() ->
with_engine_state(fun(State, Engine) ->
ok = guild_member_list_engine:bulk_load(
Engine, [{UserId, <<"hampus">>, [], false}], []
Engine, [{99, <<"hampus">>, [], false}], []
),
?assertEqual({1, 0}, guild_member_list_engine:get_counts(Engine)),
NewState = apply_reconcile_result(#{UserId => dnd_presence(UserId)}, State),
NewState = apply_mismatches([{99, dnd_presence(99), offline_display()}], State),
?assertEqual({1, 1}, guild_member_list_engine:get_counts(Engine)),
Row = guild_state_member:lookup_presence(maps:get(member_presence, NewState), UserId),
Row = guild_state_member:lookup_presence(maps:get(member_presence, NewState), 99),
?assertEqual(<<"dnd">>, maps:get(<<"status">>, Row))
after
guild_member_list_engine:destroy(Engine)
end.
end).
apply_reconcile_result_is_noop_when_consistent_test() ->
GuildId = 4242,
UserId = 99,
Engine = guild_member_list_engine:new(),
try
State = guild_test_state(GuildId, UserId, Engine),
ets:insert(maps:get(member_presence, State), {UserId, dnd_presence(UserId)}),
apply_mismatches_marks_stale_online_offline_test() ->
with_engine_state(fun(State, Engine) ->
ets:insert(maps:get(member_presence, State), {99, dnd_presence(99)}),
ok = guild_member_list_engine:bulk_load(
Engine, [{UserId, <<"hampus">>, [], true}], []
Engine, [{99, <<"hampus">>, [], true}], []
),
?assertEqual({1, 1}, guild_member_list_engine:get_counts(Engine)),
_ = apply_reconcile_result(#{UserId => dnd_presence(UserId)}, State),
?assertEqual({1, 1}, guild_member_list_engine:get_counts(Engine))
after
guild_member_list_engine:destroy(Engine)
end.
apply_reconcile_result_marks_stale_online_offline_test() ->
GuildId = 4242,
UserId = 99,
Engine = guild_member_list_engine:new(),
try
State = guild_test_state(GuildId, UserId, Engine),
ets:insert(maps:get(member_presence, State), {UserId, dnd_presence(UserId)}),
ok = guild_member_list_engine:bulk_load(
Engine, [{UserId, <<"hampus">>, [], true}], []
),
?assertEqual({1, 1}, guild_member_list_engine:get_counts(Engine)),
_ = apply_reconcile_result(#{}, State),
Seen = display_fields(dnd_presence(99)),
_ = apply_mismatches([{99, replay_payload(undefined), Seen}], State),
?assertEqual({1, 0}, guild_member_list_engine:get_counts(Engine))
end).
apply_mismatches_skips_row_updated_after_observation_test() ->
with_engine_state(fun(State, Engine) ->
Tab = maps:get(member_presence, State),
ets:insert(Tab, {99, dnd_presence(99)}),
ok = guild_member_list_engine:bulk_load(
Engine, [{99, <<"hampus">>, [], true}], []
),
Stale = {99, idle_presence(99), offline_display()},
_ = apply_mismatches([Stale], State),
?assertEqual(
<<"dnd">>, maps:get(<<"status">>, guild_state_member:lookup_presence(Tab, 99))
),
?assertEqual({1, 1}, guild_member_list_engine:get_counts(Engine))
end).
apply_mismatches_skips_disconnected_user_test() ->
with_engine_state(fun(State, Engine) ->
ok = guild_member_list_engine:bulk_load(
Engine, [{99, <<"hampus">>, [], false}], []
),
Disconnected = State#{connected_user_ids => sets:new()},
_ = apply_mismatches([{99, dnd_presence(99), offline_display()}], Disconnected),
?assertEqual({1, 0}, guild_member_list_engine:get_counts(Engine)),
?assertEqual([], ets:lookup(maps:get(member_presence, State), 99))
end).
find_mismatches_test() ->
State = state_with_presence(#{1 => dnd_presence(1), 2 => dnd_presence(2)}),
Tab = maps:get(member_presence, State),
Authoritative = #{1 => dnd_presence(1), 3 => idle_presence(3), 4 => invisible_presence(4)},
?assertEqual(
[
{2, replay_payload(undefined), display_fields(dnd_presence(2))},
{3, idle_presence(3), offline_display()}
],
find_mismatches([1, 2, 3, 4], Authoritative, Tab)
).
find_mismatches_agrees_with_reconcile_action_test() ->
State = state_with_presence(#{1 => dnd_presence(1), 2 => dnd_presence(2)}),
Tab = maps:get(member_presence, State),
Authoritative = #{1 => dnd_presence(1), 3 => idle_presence(3), 4 => invisible_presence(4)},
Expected = [
{UserId, Payload}
|| UserId <- [1, 2, 3, 4],
{replay, Payload} <- [
reconcile_action(UserId, lookup_presence_value(UserId, Authoritative), State)
]
],
?assertEqual(
Expected,
[
{UserId, Payload}
|| {UserId, Payload, _} <- find_mismatches([1, 2, 3, 4], Authoritative, Tab)
]
).
confirmed_mismatches_keeps_only_repeated_observations_test() ->
Offline = offline_display(),
Dnd = display_fields(dnd_presence(1)),
First = [
{1, dnd_presence(1), Offline},
{2, dnd_presence(2), Offline},
{3, idle_presence(3), Offline},
{4, dnd_presence(4), Offline}
],
Second = [
{1, (dnd_presence(1))#{<<"activities">> => []}, Offline},
{2, idle_presence(2), Offline},
{3, idle_presence(3), Dnd}
],
?assertEqual([hd(Second)], confirmed_mismatches(First, Second)).
persistent_mismatches_confirms_across_two_reads_test() ->
meck:new(presence_cache, [passthrough]),
try
State = state_with_presence(#{3 => dnd_presence(3)}),
Tab = maps:get(member_presence, State),
meck:expect(
presence_cache,
bulk_get,
1,
meck:seq([
[dnd_presence(1), dnd_presence(2), dnd_presence(3)],
[dnd_presence(1), idle_presence(2)]
])
),
?assertEqual(
[{1, dnd_presence(1), offline_display()}],
persistent_mismatches([1, 2, 3], Tab, 0)
),
?assertEqual([[1, 2, 3], [1, 2]], [
Ids
|| {_, {_, bulk_get, [Ids]}, _} <- meck:history(presence_cache)
])
after
meck:unload(presence_cache)
end.
persistent_mismatches_skips_second_read_when_consistent_test() ->
meck:new(presence_cache, [passthrough]),
try
State = state_with_presence(#{1 => dnd_presence(1)}),
meck:expect(presence_cache, bulk_get, 1, [dnd_presence(1)]),
?assertEqual([], persistent_mismatches([1], maps:get(member_presence, State), 0)),
?assertEqual(1, meck:num_calls(presence_cache, bulk_get, 1))
after
meck:unload(presence_cache)
end.
with_engine_state(Fun) ->
Engine = guild_member_list_engine:new(),
try
Fun(guild_test_state(4242, 99, Engine), Engine)
after
guild_member_list_engine:destroy(Engine)
end.
@@ -355,8 +504,17 @@ state_with_presence(PresenceMap) ->
}.
dnd_presence(UserId) ->
status_presence(UserId, <<"dnd">>).
idle_presence(UserId) ->
status_presence(UserId, <<"idle">>).
invisible_presence(UserId) ->
status_presence(UserId, <<"invisible">>).
status_presence(UserId, Status) ->
#{
<<"status">> => <<"dnd">>,
<<"status">> => Status,
<<"mobile">> => false,
<<"afk">> => false,
<<"custom_status">> => null,
@@ -8,8 +8,7 @@
sync_member_data/2,
partition_subscribed_sessions/5,
get_user_viewable_channel_map/3,
remove_invalid_subscriptions/3,
dispatch_to_valid_sessions/4
session_pids/2
]).
-export_type([guild_state/0, user_id/0]).
@@ -187,71 +186,14 @@ viewable_channels_or_continue(UserId, SessionData, NextIterator) ->
find_session_viewable_channels_iter(UserId, NextIterator)
end.
-spec remove_invalid_subscriptions([binary()], user_id(), guild_state()) -> guild_state().
remove_invalid_subscriptions([], _UserId, State) ->
State;
remove_invalid_subscriptions(InvalidSessionIds, UserId, State) ->
MemberSubs = maps:get(member_subscriptions, State, guild_subscriptions:init_state()),
{NewMemberSubs, RemovedCount} = unsubscribe_sessions_for_user(
InvalidSessionIds, UserId, MemberSubs
),
State1 = State#{member_subscriptions => NewMemberSubs},
guild_sessions_presence:unsubscribe_many_from_user_presence(UserId, RemovedCount, State1).
-spec unsubscribe_sessions_for_user(
[binary()], user_id(), guild_subscriptions:subscription_state()
) ->
{guild_subscriptions:subscription_state(), non_neg_integer()}.
unsubscribe_sessions_for_user(SessionIds, UserId, MemberSubs) ->
case maps:get(UserId, MemberSubs, undefined) of
undefined ->
{MemberSubs, 0};
Subscribers ->
remove_sessions_from_subscribers(SessionIds, UserId, Subscribers, MemberSubs)
end.
-spec remove_sessions_from_subscribers(
[binary()],
user_id(),
sets:set(binary()),
guild_subscriptions:subscription_state()
) -> {guild_subscriptions:subscription_state(), non_neg_integer()}.
remove_sessions_from_subscribers(SessionIds, UserId, Subscribers, MemberSubs) ->
{NewSubscribers, RemovedCount} = lists:foldl(
fun remove_session_from_subscriber_set/2,
{Subscribers, 0},
SessionIds
),
NewMemberSubs = put_or_remove_subscribers(UserId, NewSubscribers, MemberSubs),
{NewMemberSubs, RemovedCount}.
-spec remove_session_from_subscriber_set(binary(), {sets:set(binary()), non_neg_integer()}) ->
{sets:set(binary()), non_neg_integer()}.
remove_session_from_subscriber_set(SessionId, {Subscribers, Count}) ->
case sets:is_element(SessionId, Subscribers) of
true -> {sets:del_element(SessionId, Subscribers), Count + 1};
false -> {Subscribers, Count}
end.
-spec put_or_remove_subscribers(
user_id(), sets:set(binary()), guild_subscriptions:subscription_state()
) -> guild_subscriptions:subscription_state().
put_or_remove_subscribers(UserId, Subscribers, MemberSubs) ->
case sets:size(Subscribers) of
0 -> maps:remove(UserId, MemberSubs);
_ -> MemberSubs#{UserId => Subscribers}
end.
-spec dispatch_to_valid_sessions([binary()], map(), map(), integer()) -> ok.
dispatch_to_valid_sessions(ValidSessionIds, Sessions, PresenceUpdate, GuildId) ->
Pids = lists:filtermap(
-spec session_pids([binary()], map()) -> [pid()].
session_pids(SessionIds, Sessions) ->
lists:filtermap(
fun(SessionId) ->
session_pid(SessionId, Sessions)
end,
ValidSessionIds
),
gateway_dispatch_relay:dispatch_many(Pids, presence_update, PresenceUpdate, GuildId),
ok.
SessionIds
).
-spec session_pid(binary(), map()) -> {true, pid()} | false.
session_pid(SessionId, Sessions) ->
@@ -333,21 +275,17 @@ partition_subscribed_sessions_excludes_target_user_test() ->
partition_subscribed_sessions_missing_session_test() ->
assert_partition_result(#{}, #{100 => true}, [], [<<"s1">>]).
remove_invalid_subscriptions_batches_by_user_test() ->
MemberSubs0 = guild_subscriptions:init_state(),
MemberSubs1 = guild_subscriptions:subscribe(<<"s1">>, 10, MemberSubs0),
MemberSubs2 = guild_subscriptions:subscribe(<<"s2">>, 10, MemberSubs1),
MemberSubs3 = guild_subscriptions:subscribe(<<"s3">>, 10, MemberSubs2),
State = #{
member_subscriptions => MemberSubs3,
presence_subscriptions => #{10 => 5}
session_pids_keeps_order_and_skips_unknown_and_pidless_test() ->
Pid = self(),
Sessions = #{
<<"s1">> => #{user_id => 20, pid => Pid},
<<"s2">> => #{user_id => 30},
<<"s3">> => #{user_id => 40, pid => Pid}
},
Result = remove_invalid_subscriptions([<<"s1">>, <<"s3">>, <<"missing">>], 10, State),
Remaining = guild_subscriptions:get_subscribed_sessions(
10, maps:get(member_subscriptions, Result)
),
?assertEqual([<<"s2">>], lists:sort(Remaining)),
?assertEqual(#{10 => 3}, maps:get(presence_subscriptions, Result)).
?assertEqual(
[Pid, Pid],
session_pids([<<"s3">>, <<"missing">>, <<"s2">>, <<"s1">>], Sessions)
).
sync_member_data_updates_loaded_channel_engines_test() ->
GuildId = 100,
@@ -126,12 +126,13 @@ channel_has_view_restricting_overrides(Channel, ViewBit) ->
count_online_with_access(ChannelIdSet, State) ->
TargetMap = maps:from_list([{Ch, true} || Ch <- sets:to_list(ChannelIdSet)]),
Tab = maps:get(member_presence, State),
ets:foldl(
fun({UserId, Presence}, Acc) ->
sets:fold(
fun(UserId, Acc) ->
Presence = guild_state_member:lookup_presence(Tab, UserId),
maybe_count_online_user(UserId, Presence, TargetMap, State, Acc)
end,
0,
Tab
guild_member_list_connected:connected_session_user_ids(State)
).
-spec maybe_count_online_user(term(), term(), map(), guild_state(), non_neg_integer()) ->
@@ -211,6 +212,7 @@ build_state(Channels, Members, Roles, Presences) ->
},
member_presence => Tab,
member_list_engine => build_engine(Presences),
connected_user_ids => sets:from_list([UserId || {UserId, _} <- Presences]),
sessions => #{}
}.
@@ -398,6 +400,11 @@ plain_open_channel_unchanged_by_flag_test() ->
user_deny_collapses_count_without_flag_test() ->
?assertEqual(2, compute_count(user_deny_state())).
restricted_count_skips_disconnected_online_rows_test() ->
State = user_deny_state(),
?assertEqual(1, compute_count(State#{connected_user_ids => sets:from_list([41, 42])})),
?assertEqual(0, compute_count(State#{connected_user_ids => sets:new()})).
user_deny_reports_full_count_with_flag_test() ->
State = user_deny_state(),
?assertEqual(3, with_all_perms_viewer(fun() -> compute_count(State) end)).
+149 -75
View File
@@ -3,24 +3,53 @@
-module(guild_query_handler).
-typing([eqwalizer]).
-export([handle_call/3]).
-export([handle_call/3, call/3]).
-export_type([guild_state/0]).
-ifdef(TEST).
-export([strip_member_alias/1, restore_member_alias/2]).
-endif.
-define(INLINE_MEMBER_QUERY_MAX_IDS, 100).
-type guild_state() :: map().
-type user_id() :: integer().
-type query_fun() :: fun((map(), map()) -> {reply, term(), term()}).
-spec call(pid(), {atom(), map()}, pos_integer()) -> term().
call(GuildPid, {Tag, Request}, Timeout) ->
Deadline = os:system_time(millisecond) + Timeout,
gen_server:call(GuildPid, {Tag, Request#{deadline => Deadline}}, Timeout).
-spec handle_call(term(), gen_server:from(), guild_state()) ->
{reply, term(), guild_state()}
| {noreply, guild_state()}.
handle_call({get_counts}, _From, State) ->
handle_get_counts(State);
handle_call({get_user_counts, UserId}, _From, State) when is_integer(UserId) ->
handle_get_user_counts(UserId, State);
handle_call({get_channel_member_counts, Request}, _From, State) when is_map(Request) ->
handle_get_channel_member_counts(Request, State);
handle_call({get_large_guild_metadata}, _From, State) ->
handle_get_large_guild_metadata(State);
handle_call(Msg, From, State) ->
case is_expired(Msg) of
true -> {noreply, State};
false -> handle_query(Msg, From, State)
end.
-spec is_expired(term()) -> boolean().
is_expired({_Tag, #{deadline := Deadline}}) when is_integer(Deadline) ->
os:system_time(millisecond) > Deadline;
is_expired(_Msg) ->
false.
-spec handle_query(term(), gen_server:from(), guild_state()) ->
{reply, term(), guild_state()}
| {noreply, guild_state()}.
handle_query({get_counts}, _From, State) ->
handle_get_counts(State);
handle_query({get_user_counts, UserId}, _From, State) when is_integer(UserId) ->
handle_get_user_counts(UserId, State);
handle_query({get_viewer_counts, #{user_id := UserId}}, _From, State) when is_integer(UserId) ->
handle_get_user_counts(UserId, State);
handle_query({get_channel_member_counts, Request}, _From, State) when is_map(Request) ->
handle_get_channel_member_counts(Request, State);
handle_query({get_large_guild_metadata}, _From, State) ->
handle_get_large_guild_metadata(State);
handle_query(Msg, From, State) ->
handle_call_dispatch(Msg, From, State).
-spec handle_get_counts(guild_state()) -> {reply, map(), guild_state()}.
@@ -127,97 +156,142 @@ handle_get_large_guild_metadata(State) ->
-spec handle_call_dispatch(term(), gen_server:from(), guild_state()) ->
{reply, term(), guild_state()} | {noreply, guild_state()}.
handle_call_dispatch({get_users_to_mention_by_roles, Req}, From, State) ->
async_member_query(
From, State, request_map(Req), fun guild_members:get_users_to_mention_by_roles/2
);
handle_call_dispatch({get_users_to_mention_by_user_ids, Req}, From, State) ->
handle_async_member_query({get_users_to_mention_by_user_ids, Req}, From, State);
handle_call_dispatch({check_permission, Request}, From, State) ->
handle_check_permission(request_map(Request), From, State);
handle_call_dispatch({get_user_permissions, Request}, From, State) ->
handle_get_user_permissions(request_map(Request), From, State);
handle_call_dispatch({check_permission, Request}, _From, State) ->
handle_check_permission(request_map(Request), State);
handle_call_dispatch({get_user_permissions, Request}, _From, State) ->
handle_get_user_permissions(request_map(Request), State);
handle_call_dispatch(Msg, From, State) ->
handle_async_member_query(Msg, From, State).
-spec handle_async_member_query(term(), gen_server:from(), guild_state()) ->
{reply, term(), guild_state()} | {noreply, guild_state()}.
handle_async_member_query({get_users_to_mention_by_user_ids, Req}, From, State) ->
async_member_query(
From, State, request_map(Req), fun guild_members:get_users_to_mention_by_user_ids/2
);
handle_async_member_query({get_all_users_to_mention, Req}, From, State) ->
async_member_query(
From, State, request_map(Req), fun guild_members:get_all_users_to_mention/2
);
handle_async_member_query({resolve_all_mentions, Req}, From, State) ->
async_member_query(From, State, request_map(Req), fun guild_members:resolve_all_mentions/2);
handle_async_member_query({resolve_mention_sources, Req}, From, State) ->
async_member_query(
From, State, request_map(Req), fun guild_members:resolve_mention_sources/2
);
handle_async_member_query({resolve_mention_sources_page, Req}, From, State) ->
async_member_query(
From, State, request_map(Req), fun guild_members:resolve_mention_sources_page/2
);
handle_async_member_query({resolve_channel_mentions, Req}, From, State) ->
async_member_query(
From, State, request_map(Req), fun guild_members:resolve_channel_mentions/2
);
handle_async_member_query({get_members_with_role, Req}, From, State) ->
async_member_query(
From, State, request_map(Req), fun guild_members:get_members_with_role/2
);
handle_async_member_query({get_viewable_channels, Req}, From, State) ->
async_member_query(
From, State, request_map(Req), fun guild_members:get_viewable_channels/2
);
handle_async_member_query({Tag, Req}, From, State) when is_atom(Tag) ->
case member_query_fun(Tag) of
undefined -> handle_call_sync({Tag, Req}, State);
QueryFun -> member_query(Tag, From, State, request_map(Req), QueryFun)
end;
handle_async_member_query(Msg, _From, State) ->
handle_call_sync(Msg, State).
-spec async_member_query(
gen_server:from(), guild_state(), map(), fun((map(), map()) -> {reply, term(), term()})
) ->
-spec member_query_fun(atom()) -> query_fun() | undefined.
member_query_fun(get_users_to_mention_by_roles) ->
fun guild_members:get_users_to_mention_by_roles/2;
member_query_fun(get_users_to_mention_by_user_ids) ->
fun guild_members:get_users_to_mention_by_user_ids/2;
member_query_fun(get_all_users_to_mention) ->
fun guild_members:get_all_users_to_mention/2;
member_query_fun(resolve_all_mentions) ->
fun guild_members:resolve_all_mentions/2;
member_query_fun(resolve_mention_sources) ->
fun guild_members:resolve_mention_sources/2;
member_query_fun(resolve_mention_sources_page) ->
fun guild_members:resolve_mention_sources_page/2;
member_query_fun(resolve_channel_mentions) ->
fun guild_members:resolve_channel_mentions/2;
member_query_fun(get_members_with_role) ->
fun guild_members:get_members_with_role/2;
member_query_fun(get_viewable_channels) ->
fun guild_members:get_viewable_channels/2;
member_query_fun(_Tag) ->
undefined.
-spec member_query(atom(), gen_server:from(), guild_state(), map(), query_fun()) ->
{reply, term(), guild_state()} | {noreply, guild_state()}.
member_query(Tag, From, State, Request, QueryFun) ->
case bounded_member_query(Tag, Request) of
true -> inline_member_query(State, Request, QueryFun);
false -> async_member_query(From, State, Request, QueryFun)
end.
-spec bounded_member_query(atom(), map()) -> boolean().
bounded_member_query(get_users_to_mention_by_user_ids, Request) ->
few_ids(maps:get(user_ids, Request, undefined));
bounded_member_query(resolve_all_mentions, Request) ->
direct_mentions_only(Request);
bounded_member_query(resolve_mention_sources, Request) ->
direct_mentions_only(Request);
bounded_member_query(resolve_mention_sources_page, Request) ->
direct_mentions_only(Request);
bounded_member_query(resolve_channel_mentions, Request) ->
few_ids(maps:get(channel_ids, Request, undefined));
bounded_member_query(get_viewable_channels, _Request) ->
true;
bounded_member_query(_Tag, _Request) ->
false.
-spec direct_mentions_only(map()) -> boolean().
direct_mentions_only(Request) ->
maps:get(mention_everyone, Request, undefined) =:= false andalso
maps:get(mention_here, Request, undefined) =:= false andalso
maps:get(role_ids, Request, undefined) =:= [] andalso
few_ids(maps:get(user_ids, Request, undefined)).
-spec few_ids(term()) -> boolean().
few_ids(Ids) when is_list(Ids) ->
length(Ids) =< ?INLINE_MEMBER_QUERY_MAX_IDS;
few_ids(_Ids) ->
false.
-spec inline_member_query(guild_state(), map(), query_fun()) -> {reply, term(), guild_state()}.
inline_member_query(State, Request, QueryFun) ->
QS = build_query_snapshot(State),
Reply = safe_reply(fun() ->
{reply, QueryReply, _} = QueryFun(Request, QS),
QueryReply
end),
{reply, Reply, State}.
-spec async_member_query(gen_server:from(), guild_state(), map(), query_fun()) ->
{noreply, guild_state()}.
async_member_query(From, State, Request, QueryFun) ->
QS = build_query_snapshot(State),
{Aliased, QS} = strip_member_alias(build_query_snapshot(State)),
spawn_async_reply(From, fun() ->
{reply, Reply, _} = QueryFun(Request, QS),
{reply, Reply, _} = QueryFun(Request, restore_member_alias(Aliased, QS)),
Reply
end),
{noreply, State}.
-spec handle_check_permission(map(), gen_server:from(), guild_state()) ->
{noreply, guild_state()}.
handle_check_permission(Request, From, State) ->
QS = build_query_snapshot(State),
spawn_async_reply(From, fun() ->
-spec strip_member_alias(map()) -> {boolean(), map()}.
strip_member_alias(
#{data := #{<<"members">> := Members, members_normalized := Members} = Data} = QS
) ->
{true, QS#{data := maps:remove(members_normalized, Data)}};
strip_member_alias(QS) ->
{false, QS}.
-spec restore_member_alias(boolean(), map()) -> map().
restore_member_alias(true, #{data := #{<<"members">> := Members} = Data} = QS) ->
QS#{data := Data#{members_normalized => Members}};
restore_member_alias(_Aliased, QS) ->
QS.
-spec handle_check_permission(map(), guild_state()) -> {reply, map(), guild_state()}.
handle_check_permission(Request, State) ->
Reply = safe_reply(fun() ->
#{user_id := UserId, permission := Permission, channel_id := ChannelId} = Request,
true = is_integer(Permission),
HasPermission = check_user_permission(UserId, Permission, ChannelId, QS),
HasPermission = check_user_permission(UserId, Permission, ChannelId, State),
#{has_permission => HasPermission}
end),
{noreply, State}.
{reply, Reply, State}.
-spec check_user_permission(user_id(), integer(), integer(), map()) -> boolean().
check_user_permission(UserId, Permission, ChannelId, QS) ->
case owner_id(QS) =:= UserId of
-spec check_user_permission(user_id(), integer(), integer(), guild_state()) -> boolean().
check_user_permission(UserId, Permission, ChannelId, State) ->
case owner_id(State) =:= UserId of
true ->
true;
false ->
Perms = guild_permissions:get_member_permissions(UserId, ChannelId, QS),
Perms = guild_permissions:get_member_permissions(UserId, ChannelId, State),
permission_bits:has(Perms, Permission)
end.
-spec handle_get_user_permissions(map(), gen_server:from(), guild_state()) ->
{noreply, guild_state()}.
handle_get_user_permissions(Request, From, State) ->
QS = build_query_snapshot(State),
spawn_async_reply(From, fun() ->
-spec handle_get_user_permissions(map(), guild_state()) -> {reply, map(), guild_state()}.
handle_get_user_permissions(Request, State) ->
Reply = safe_reply(fun() ->
#{user_id := UserId, channel_id := ChannelId} = Request,
#{permissions => guild_permissions:get_member_permissions(UserId, ChannelId, QS)}
#{permissions => guild_permissions:get_member_permissions(UserId, ChannelId, State)}
end),
{noreply, State}.
{reply, Reply, State}.
-spec handle_call_sync(term(), guild_state()) -> {reply, term(), guild_state()}.
handle_call_sync({can_manage_roles, Req}, State) ->
@@ -324,10 +398,10 @@ spawn_async_reply(From, ReplyFun) ->
-spec send_async_reply(gen_server:from(), fun(() -> term())) -> ok.
send_async_reply(From, ReplyFun) ->
gen_server:reply(From, safe_async_reply(ReplyFun)).
gen_server:reply(From, safe_reply(ReplyFun)).
-spec safe_async_reply(fun(() -> term())) -> term().
safe_async_reply(ReplyFun) ->
-spec safe_reply(fun(() -> term())) -> term().
safe_reply(ReplyFun) ->
try
ReplyFun()
catch
@@ -100,8 +100,9 @@ spawn_fetch_worker(Self, Tag, GuildId, GuildPid, UserId) ->
-spec worker(pid(), reference(), integer(), pid(), integer()) -> ok.
worker(Parent, Tag, GuildId, GuildPid, UserId) ->
Request = {get_viewer_counts, #{user_id => UserId}},
Result =
try gen_server:call(GuildPid, {get_user_counts, UserId}, ?GUILD_CALL_TIMEOUT_MS) of
try guild_query_handler:call(GuildPid, Request, ?GUILD_CALL_TIMEOUT_MS) of
#{member_count := MemberCount, online_count := OnlineCount} ->
{ok, MemberCount, OnlineCount};
_ ->
@@ -232,6 +233,27 @@ handle_request_echoes_nonce_test() ->
?assert(false)
end.
handle_request_fetches_viewer_counts_with_deadline_test() ->
Self = self(),
Guild = spawn(fun() ->
receive
{'$gen_call', From, {get_viewer_counts, #{user_id := 100, deadline := D}}} when
is_integer(D)
->
gen_server:reply(From, #{member_count => 50, online_count => 10})
end
end),
SessionState = #{
session_pid => Self, user_id => <<"100">>, guilds => #{7 => {Guild, make_ref()}}
},
ok = handle_request(#{<<"guild_ids">> => [<<"7">>]}, Self, SessionState),
receive
{'$gen_cast', {dispatch, guild_counts_update, Payload}} ->
?assertEqual([build_entry(7, 50, 10)], maps:get(<<"counts">>, Payload))
after 1000 ->
?assert(false)
end.
parse_nonce_test() ->
?assertEqual(<<"x">>, parse_nonce(<<"x">>)),
?assertEqual(<<"abc">>, parse_nonce(<<"abc">>)),
+251 -125
View File
@@ -30,7 +30,7 @@
build_viewable_channel_map/1
]).
-define(MAX_PERM_MEMO_ENTRIES, 8192).
-define(MAX_MEMO_ENTRIES, 8192).
-type guild_state() :: map().
-type session_id() :: binary().
@@ -41,6 +41,7 @@
-type sessions_map() :: #{session_id() => session_data()}.
-type session_pair() :: {session_id(), session_data()}.
-type perm_memo() :: #{user_id() => non_neg_integer()}.
-type view_memo() :: #{user_id() => boolean()}.
-type message_ctx() :: {channel_id(), binary(), session_id() | undefined, guild_state()}.
-export_type([
guild_state/0,
@@ -65,62 +66,23 @@ handle_session_connect(Request, Pid, State) ->
-spec handle_session_down(reference(), guild_state()) ->
{noreply, guild_state()} | {stop, normal, guild_state()}.
handle_session_down(Ref, State) ->
case pending_session_by_ref(Ref, State) of
{ok, SessionId, Session, Sessions} ->
handle_pending_ref_down(SessionId, Session, Ref, Sessions, State);
not_found ->
guild_sessions_connect:handle_session_down(Ref, State)
case guild_sessions_connect:find_session_by_ref(Ref, State) of
{SessionId, #{pending_connect := true} = Session} = Found ->
handle_pending_ref_down(SessionId, Session, Ref, Found, State);
Found ->
guild_sessions_connect:handle_session_down(Ref, Found, State)
end.
-spec handle_pending_ref_down(
session_id(), session_data(), reference(), sessions_map(), guild_state()
session_id(), session_data(), reference(), {session_id(), session_data()}, guild_state()
) -> {noreply, guild_state()} | {stop, normal, guild_state()}.
handle_pending_ref_down(SessionId, Session, Ref, Sessions, State) ->
handle_pending_ref_down(SessionId, Session, Ref, Found, State) ->
Sessions = maps:get(sessions, State, #{}),
case pending_session_owns_connected_tracking(Session, Sessions, State) of
true -> guild_sessions_connect:handle_session_down(Ref, State);
true -> guild_sessions_connect:handle_session_down(Ref, Found, State);
false -> handle_pending_session_down(SessionId, Ref, Sessions, State)
end.
-spec pending_session_by_ref(reference(), guild_state()) ->
{ok, session_id(), session_data(), sessions_map()} | not_found.
pending_session_by_ref(Ref, State) ->
Sessions = maps:get(sessions, State, #{}),
case session_id_by_ref(Ref, Sessions, State) of
SessionId when is_binary(SessionId) -> pending_session_entry(SessionId, Sessions);
_ -> not_found
end.
-spec pending_session_entry(session_id(), sessions_map()) ->
{ok, session_id(), session_data(), sessions_map()} | not_found.
pending_session_entry(SessionId, Sessions) ->
case maps:get(SessionId, Sessions, undefined) of
#{pending_connect := true} = Session -> {ok, SessionId, Session, Sessions};
_ -> not_found
end.
-spec session_id_by_ref(reference(), sessions_map(), guild_state()) -> session_id() | undefined.
session_id_by_ref(Ref, Sessions, State) ->
Refs = maps:get(guild_session_refs, State, #{}),
case maps:get(Ref, Refs, undefined) of
SessionId when is_binary(SessionId) -> SessionId;
_ -> session_id_by_ref_scan(Ref, Sessions)
end.
-spec session_id_by_ref_scan(reference(), sessions_map()) -> session_id() | undefined.
session_id_by_ref_scan(Ref, Sessions) ->
maps:fold(
fun(SessionId, Session, Found) ->
match_session_ref(Ref, SessionId, Session, Found)
end,
undefined,
Sessions
).
-spec match_session_ref(reference(), session_id(), session_data(), session_id() | undefined) ->
session_id() | undefined.
match_session_ref(Ref, SessionId, #{mref := Ref}, _Found) -> SessionId;
match_session_ref(_Ref, _SessionId, _Session, Found) -> Found.
-spec pending_session_owns_connected_tracking(session_data(), sessions_map(), guild_state()) ->
boolean().
pending_session_owns_connected_tracking(Session, Sessions, State) ->
@@ -157,11 +119,7 @@ count_active_session(_UserId, _Session, Count) ->
{noreply, guild_state()}.
handle_pending_session_down(SessionId, Ref, Sessions, State) ->
NewSessions = maps:remove(SessionId, Sessions),
SessionRefs = maps:get(guild_session_refs, State, #{}),
State1 = State#{
sessions => NewSessions,
guild_session_refs => maps:remove(Ref, SessionRefs)
},
State1 = guild_sessions_connect:remove_session_ref(Ref, State#{sessions => NewSessions}),
State2 = guild_sessions_connect_cleanup:cleanup_connect_admission_for_session(
SessionId, State1
),
@@ -247,9 +205,31 @@ handle_send_members_chunk(SessionId, ChunkData, State) ->
sessions_map(), channel_id(), session_id() | undefined, guild_state()
) -> [session_pair()].
filter_sessions_for_channel(Sessions, ChannelId, SessionIdOpt, State) ->
filter_active_sessions(Sessions, SessionIdOpt, fun(S, _Sid) ->
session_can_view_channel(S, ChannelId, State)
end).
{Acc, _Memo} = maps:fold(
fun(Sid, S, In) ->
collect_channel_session(Sid, S, ChannelId, SessionIdOpt, State, In)
end,
{[], #{}},
Sessions
),
Acc.
-spec collect_channel_session(
session_id(),
session_data(),
channel_id(),
session_id() | undefined,
guild_state(),
{[session_pair()], view_memo()}
) -> {[session_pair()], view_memo()}.
collect_channel_session(Sid, S, ChannelId, SessionIdOpt, State, {Acc, Memo}) ->
case is_pending_or_excluded(Sid, S, SessionIdOpt) of
true ->
{Acc, Memo};
false ->
{Visible, Memo1} = memo_session_can_view_channel(S, ChannelId, State, Memo),
{prepend_session(Visible, Sid, S, Acc), Memo1}
end.
-spec filter_sessions_for_message(
sessions_map(), channel_id(), binary(), session_id() | undefined, guild_state()
@@ -262,26 +242,46 @@ filter_sessions_for_message(Sessions, ChannelId, MessageId, SessionIdOpt, State)
) -> [session_pair()].
filter_message_memo(Sessions, ChannelId, MessageId, SessionIdOpt, State) ->
Ctx = {ChannelId, MessageId, SessionIdOpt, State},
{Acc, _Memo} = maps:fold(
{Acc, _ViewMemo, _PermMemo} = maps:fold(
fun(Sid, S, In) -> collect_message_session(Sid, S, Ctx, In) end,
{[], #{}},
{[], #{}, #{}},
Sessions
),
Acc.
-spec collect_message_session(
session_id(), session_data(), message_ctx(), {[session_pair()], perm_memo()}
) -> {[session_pair()], perm_memo()}.
collect_message_session(Sid, S, Ctx, {Acc, Memo}) ->
session_id(), session_data(), message_ctx(), {[session_pair()], view_memo(), perm_memo()}
) -> {[session_pair()], view_memo(), perm_memo()}.
collect_message_session(Sid, S, Ctx, {Acc, ViewMemo, PermMemo}) ->
{ChannelId, _MessageId, SessionIdOpt, State} = Ctx,
Visible =
not is_pending_or_excluded(Sid, S, SessionIdOpt) andalso
session_can_view_channel(S, ChannelId, State),
case Visible of
true -> memo_message_session(Sid, S, Ctx, Acc, Memo);
false -> {Acc, Memo}
case is_pending_or_excluded(Sid, S, SessionIdOpt) of
true ->
{Acc, ViewMemo, PermMemo};
false ->
collect_visible_message_session(
Sid,
S,
Ctx,
Acc,
memo_session_can_view_channel(S, ChannelId, State, ViewMemo),
PermMemo
)
end.
-spec collect_visible_message_session(
session_id(),
session_data(),
message_ctx(),
[session_pair()],
{boolean(), view_memo()},
perm_memo()
) -> {[session_pair()], view_memo(), perm_memo()}.
collect_visible_message_session(Sid, S, Ctx, Acc, {true, ViewMemo}, PermMemo) ->
{Acc1, PermMemo1} = memo_message_session(Sid, S, Ctx, Acc, PermMemo),
{Acc1, ViewMemo, PermMemo1};
collect_visible_message_session(_Sid, _S, _Ctx, Acc, {false, ViewMemo}, PermMemo) ->
{Acc, ViewMemo, PermMemo}.
-spec memo_message_session(
session_id(), session_data(), message_ctx(), [session_pair()], perm_memo()
) -> {[session_pair()], perm_memo()}.
@@ -303,13 +303,13 @@ memo_member_permissions(UserId, ChannelId, State, Memo) ->
{Perms, Memo};
_ ->
Computed = guild_permissions:get_member_permissions(UserId, ChannelId, State),
{Computed, store_perm_memo(UserId, Computed, Memo)}
{Computed, store_memo(UserId, Computed, Memo)}
end.
-spec store_perm_memo(user_id(), non_neg_integer(), perm_memo()) -> perm_memo().
store_perm_memo(UserId, Perms, Memo) when map_size(Memo) < ?MAX_PERM_MEMO_ENTRIES ->
Memo#{UserId => Perms};
store_perm_memo(_UserId, _Perms, Memo) ->
-spec store_memo(user_id(), V, #{user_id() => V}) -> #{user_id() => V}.
store_memo(UserId, Value, Memo) when map_size(Memo) < ?MAX_MEMO_ENTRIES ->
Memo#{UserId => Value};
store_memo(_UserId, _Value, Memo) ->
Memo.
-spec prepend_session(boolean(), session_id(), session_data(), [session_pair()]) ->
@@ -449,17 +449,31 @@ refresh_session_viewable(SessionId, SessionData, AccState) ->
AccState
end.
-spec session_can_view_channel(session_data(), channel_id(), guild_state()) -> boolean().
session_can_view_channel(SessionData, ChannelId, State) ->
-spec memo_session_can_view_channel(session_data(), channel_id(), guild_state(), view_memo()) ->
{boolean(), view_memo()}.
memo_session_can_view_channel(SessionData, ChannelId, State, Memo) ->
UserId = maps:get(user_id, SessionData, undefined),
case {UserId, maps:get(viewable_channels, SessionData, undefined)} of
{Uid, ViewableChannels} when is_integer(Uid), is_map(ViewableChannels) ->
maps:is_key(ChannelId, ViewableChannels) orelse
check_member_channel_access(Uid, ChannelId, State);
case maps:is_key(ChannelId, ViewableChannels) of
true -> {true, Memo};
false -> memo_member_channel_access(Uid, ChannelId, State, Memo)
end;
{Uid, _} when is_integer(Uid) ->
check_member_channel_access(Uid, ChannelId, State);
memo_member_channel_access(Uid, ChannelId, State, Memo);
_ ->
false
{false, Memo}
end.
-spec memo_member_channel_access(user_id(), channel_id(), guild_state(), view_memo()) ->
{boolean(), view_memo()}.
memo_member_channel_access(UserId, ChannelId, State, Memo) ->
case maps:get(UserId, Memo, undefined) of
Visible when is_boolean(Visible) ->
{Visible, Memo};
_ ->
Computed = check_member_channel_access(UserId, ChannelId, State),
{Computed, store_memo(UserId, Computed, Memo)}
end.
-spec check_member_channel_access(user_id(), channel_id(), guild_state()) -> boolean().
@@ -522,14 +536,14 @@ filter_active_sessions_with_predicate_test() ->
?assertEqual(1, length(Result)),
[{<<"b">>, _}] = Result.
store_perm_memo_test() ->
?assertEqual(#{7 => 42}, store_perm_memo(7, 42, #{})),
?assertEqual(#{7 => 42}, store_perm_memo(7, 42, #{7 => 42})).
store_memo_test() ->
?assertEqual(#{7 => 42}, store_memo(7, 42, #{})),
?assertEqual(#{7 => 42}, store_memo(7, 42, #{7 => 42})).
store_perm_memo_bound_test() ->
Full = maps:from_list([{I, 0} || I <- lists:seq(1, ?MAX_PERM_MEMO_ENTRIES)]),
?assertEqual(Full, store_perm_memo(0, 1, Full)),
?assertEqual(?MAX_PERM_MEMO_ENTRIES, map_size(store_perm_memo(0, 1, Full))).
store_memo_bound_test() ->
Full = maps:from_list([{I, 0} || I <- lists:seq(1, ?MAX_MEMO_ENTRIES)]),
?assertEqual(Full, store_memo(0, 1, Full)),
?assertEqual(?MAX_MEMO_ENTRIES, map_size(store_memo(0, 1, Full))).
memo_member_permissions_hit_test() ->
Memo = #{9 => 123},
@@ -549,15 +563,165 @@ reference_session_can_access_message(SessionData, ChannelId, MessageId, State) -
false
end.
-spec reference_session_can_view_channel(session_data(), channel_id(), guild_state()) ->
boolean().
reference_session_can_view_channel(SessionData, ChannelId, State) ->
UserId = maps:get(user_id, SessionData, undefined),
case {UserId, maps:get(viewable_channels, SessionData, undefined)} of
{Uid, ViewableChannels} when is_integer(Uid), is_map(ViewableChannels) ->
maps:is_key(ChannelId, ViewableChannels) orelse
check_member_channel_access(Uid, ChannelId, State);
{Uid, _} when is_integer(Uid) ->
check_member_channel_access(Uid, ChannelId, State);
_ ->
false
end.
-spec reference_filter_channel_direct(
sessions_map(), channel_id(), session_id() | undefined, guild_state()
) -> [session_pair()].
reference_filter_channel_direct(Sessions, ChannelId, SessionIdOpt, State) ->
filter_active_sessions(Sessions, SessionIdOpt, fun(S, _Sid) ->
reference_session_can_view_channel(S, ChannelId, State)
end).
-spec reference_filter_message_direct(
sessions_map(), channel_id(), binary(), session_id() | undefined, guild_state()
) -> [session_pair()].
reference_filter_message_direct(Sessions, ChannelId, MessageId, SessionIdOpt, State) ->
filter_active_sessions(Sessions, SessionIdOpt, fun(S, _Sid) ->
session_can_view_channel(S, ChannelId, State) andalso
reference_session_can_view_channel(S, ChannelId, State) andalso
reference_session_can_access_message(S, ChannelId, MessageId, State)
end).
memo_member_channel_access_hit_test() ->
Memo = #{1001 => true},
?assertEqual({true, Memo}, memo_member_channel_access(1001, 10, #{}, Memo)).
memo_member_channel_access_miss_test() ->
?assertEqual({false, #{1001 => false}}, memo_member_channel_access(1001, 10, #{}, #{})).
memo_session_can_view_channel_listed_skips_memo_test() ->
Session = #{user_id => 1001, viewable_channels => #{10 => true}},
?assertEqual({true, #{}}, memo_session_can_view_channel(Session, 10, #{}, #{})).
memo_session_can_view_channel_without_user_test() ->
Session = #{viewable_channels => #{10 => true}},
?assertEqual({false, #{}}, memo_session_can_view_channel(Session, 10, #{}, #{})).
channel_filters_match_per_session_reference_test() ->
State = visibility_fixture_state(),
Sessions = maps:get(sessions, State),
Excludes = [undefined | lists:sublist(lists:sort(maps:keys(Sessions)), 3)],
Results = [
assert_filters_match_reference(Sessions, ChannelId, Exclude, State)
|| ChannelId <- visibility_fixture_channels(), Exclude <- Excludes
],
Visible = lists:append([Sids || {Sids, _} <- Results]),
?assert(lists:member(<<"u1002-stale">>, Visible)),
?assert(lists:member(<<"u1005-nomap">>, Visible)),
?assert(lists:member(<<"u1007-empty">>, Visible)),
?assertNot(lists:member(<<"nouser">>, Visible)),
?assert(lists:any(fun({Sids, _}) -> Sids =/= [] end, Results)),
?assert(lists:any(fun({Sids, Msg}) -> length(Msg) < length(Sids) end, Results)).
assert_filters_match_reference(Sessions, ChannelId, Exclude, State) ->
Channel = filter_sessions_for_channel(Sessions, ChannelId, Exclude, State),
?assertEqual(reference_filter_channel_direct(Sessions, ChannelId, Exclude, State), Channel),
MessageId = <<"1430000000000000000">>,
Message = filter_sessions_for_message(Sessions, ChannelId, MessageId, Exclude, State),
?assertEqual(
reference_filter_message_direct(Sessions, ChannelId, MessageId, Exclude, State), Message
),
{[Sid || {Sid, _} <- Channel], [Sid || {Sid, _} <- Message]}.
visibility_fixture_channels() ->
[10, 20, 30, 40, 50, 60, 61, 70].
visibility_fixture_state() ->
View = constants:view_channel_permission(),
History = constants:read_message_history_permission(),
Deny = fun(Id, Type, Bits) ->
#{
<<"id">> => Id,
<<"type">> => Type,
<<"allow">> => <<"0">>,
<<"deny">> => integer_to_binary(Bits)
}
end,
Allow = fun(Id, Type, Bits) ->
#{
<<"id">> => Id,
<<"type">> => Type,
<<"allow">> => integer_to_binary(Bits),
<<"deny">> => <<"0">>
}
end,
Channel = fun(Id, Type, Parent, Overwrites) ->
#{
<<"id">> => integer_to_binary(Id),
<<"type">> => Type,
<<"parent_id">> => Parent,
<<"permission_overwrites">> => Overwrites
}
end,
Member = fun(UserId, Roles) ->
#{<<"user">> => #{<<"id">> => integer_to_binary(UserId)}, <<"roles">> => Roles}
end,
Data = #{
<<"guild">> => #{
<<"id">> => <<"1">>, <<"owner_id">> => <<"999">>, <<"features">> => []
},
<<"roles">> => [
#{<<"id">> => <<"1">>, <<"permissions">> => integer_to_binary(View bor History)},
#{<<"id">> => <<"200">>, <<"permissions">> => <<"0">>},
#{<<"id">> => <<"300">>, <<"permissions">> => <<"0">>}
],
<<"members">> => [
Member(999, []),
Member(1001, []),
Member(1002, [<<"300">>]),
Member(1003, [<<"200">>]),
Member(1004, [<<"200">>, <<"300">>]),
Member(1005, []),
Member(1006, [<<"300">>]),
Member(1007, [])
],
<<"channels">> => [
Channel(10, 0, null, []),
Channel(20, 0, null, [Deny(<<"1">>, 0, View), Allow(<<"300">>, 0, View)]),
Channel(30, 0, null, [Deny(<<"1">>, 0, View), Allow(<<"1005">>, 1, View)]),
Channel(40, 0, null, [Deny(<<"200">>, 0, View)]),
Channel(50, 0, null, [Deny(<<"1">>, 0, History), Allow(<<"300">>, 0, History)]),
Channel(60, 4, null, [Deny(<<"1">>, 0, View)]),
Channel(61, 0, <<"60">>, [Allow(<<"300">>, 0, View)])
]
},
Sessions = #{
<<"owner">> => #{user_id => 999, viewable_channels => #{}},
<<"u1001-a">> => #{user_id => 1001, viewable_channels => #{10 => true, 40 => true}},
<<"u1001-b">> => #{user_id => 1001},
<<"u1001-pending">> => #{user_id => 1001, pending_connect => true},
<<"u1002-stale">> => #{user_id => 1002, viewable_channels => #{}},
<<"u1002-full">> => #{user_id => 1002, viewable_channels => #{10 => true, 20 => true}},
<<"u1003-a">> => #{user_id => 1003, viewable_channels => #{10 => true, 40 => true}},
<<"u1003-b">> => #{user_id => 1003, viewable_channels => #{}},
<<"u1004">> => #{user_id => 1004, viewable_channels => #{20 => true}},
<<"u1005-nomap">> => #{user_id => 1005},
<<"u1005-map">> => #{user_id => 1005, viewable_channels => #{10 => true}},
<<"u1006">> => #{user_id => 1006, viewable_channels => #{61 => true}},
<<"u1007-empty">> => #{user_id => 1007, viewable_channels => #{}},
<<"u1008-nonmember">> => #{user_id => 1008, viewable_channels => #{}},
<<"u1008-listed">> => #{user_id => 1008, viewable_channels => #{50 => true}},
<<"nouser">> => #{viewable_channels => #{10 => true}}
},
#{
id => 1,
sessions => Sessions,
data => Data,
virtual_channel_access => #{1007 => sets:from_list([20])}
}.
filter_message_memo_matches_direct_test() ->
State = #{data => #{<<"guild">> => #{<<"owner_id">> => <<"1">>}}},
Owner1 = #{user_id => 1, viewable_channels => #{5 => true}},
@@ -610,44 +774,6 @@ active_session_count_test() ->
?assertEqual(1, active_session_count(2, Sessions)),
?assertEqual(0, active_session_count(3, Sessions)).
session_id_by_ref_uses_index_test() ->
Ref = make_ref(),
Sessions = #{<<"a">> => #{mref => make_ref()}},
State = #{guild_session_refs => #{Ref => <<"a">>}},
?assertEqual(<<"a">>, session_id_by_ref(Ref, Sessions, State)).
session_id_by_ref_scan_fallback_test() ->
Ref = make_ref(),
Sessions = #{<<"a">> => #{mref => make_ref()}, <<"b">> => #{mref => Ref}},
?assertEqual(<<"b">>, session_id_by_ref(Ref, Sessions, #{})),
?assertEqual(undefined, session_id_by_ref(make_ref(), Sessions, #{})).
pending_session_by_ref_found_test() ->
Ref = make_ref(),
Session = #{user_id => 1, mref => Ref, pending_connect => true},
Sessions = #{<<"a">> => Session},
State = #{sessions => Sessions, guild_session_refs => #{Ref => <<"a">>}},
?assertEqual({ok, <<"a">>, Session, Sessions}, pending_session_by_ref(Ref, State)),
?assertEqual(
{ok, <<"a">>, Session, Sessions},
pending_session_by_ref(Ref, #{sessions => Sessions})
).
pending_session_by_ref_not_pending_test() ->
Ref = make_ref(),
Session = #{user_id => 1, mref => Ref, pending_connect => false},
State = #{sessions => #{<<"a">> => Session}, guild_session_refs => #{Ref => <<"a">>}},
?assertEqual(not_found, pending_session_by_ref(Ref, State)).
pending_session_by_ref_unknown_ref_test() ->
Sessions = #{<<"a">> => #{user_id => 1, mref => make_ref(), pending_connect => true}},
?assertEqual(not_found, pending_session_by_ref(make_ref(), #{sessions => Sessions})).
pending_session_by_ref_stale_index_test() ->
Ref = make_ref(),
State = #{sessions => #{}, guild_session_refs => #{Ref => <<"gone">>}},
?assertEqual(not_found, pending_session_by_ref(Ref, State)).
pending_session_owns_connected_tracking_untracked_test() ->
Session = #{user_id => 1, pending_connect => true},
Sessions = #{<<"a">> => Session},
@@ -6,9 +6,14 @@
-export([
handle_session_connect/3,
resection_connected_user/3,
resection_connected_user/4,
build_initial_last_message_ids/1,
build_initial_channel_versions/1,
handle_session_down/2,
handle_session_down/3,
find_session_by_ref/2,
put_session_ref/3,
remove_session_ref/2,
build_session_ref_index/1,
remove_session/2,
invalidate_viewable_channels_cache/1
]).
@@ -33,7 +38,7 @@
{reply, connect_reply(), guild_state()}.
handle_session_connect(Request, Pid, State) ->
#{session_id := SessionId, user_id := UserId} = Request,
Sessions = require_sessions(maps:get(sessions, State, #{})),
Sessions = maps:get(sessions, State, #{}),
case maps:is_key(SessionId, Sessions) of
true ->
{reply, {ok, guild_data:get_guild_state(UserId, State)}, State};
@@ -46,7 +51,7 @@ handle_session_connect(Request, Pid, State) ->
) ->
{reply, connect_reply(), guild_state()}.
register_new_session(Request, Pid, UserId, SessionId, State) ->
Sessions = require_sessions(maps:get(sessions, State, #{})),
Sessions = maps:get(sessions, State, #{}),
case user_session_count(UserId, State, Sessions) >= ?MAX_SESSIONS_PER_USER_PER_GUILD of
true ->
{reply, {error, too_many_sessions}, State};
@@ -104,15 +109,16 @@ register_admitted_session(
SessionData = build_session_data(Request, Pid, UserId, SessionId, State),
store_initial_passive_state(SessionId, GuildId, GuildState),
NewSessions = Sessions#{SessionId => SessionData},
StateWithSession = put_session_ref(SessionId, SessionData, State#{
StateWithSession = put_session_ref(SessionId, maps:get(mref, SessionData), State#{
sessions => NewSessions
}),
State0a = guild_sessions_connect_cleanup:clear_auto_stop_pending(
StateWithSession
),
State1 = track_connected_user(UserId, 1, State0a),
PresenceBefore = guild_member_list_connected:resolve_presence_for_user(State1, UserId),
State2 = guild_sessions_presence:subscribe_connected_user_presence(UserId, State1),
State3 = resection_connected_user(UserId, State, State2),
State3 = resection_connected_user(UserId, PresenceBefore, State, State2),
InitialGuildId = maps:get(initial_guild_id, Request, undefined),
finalize_connect(
SessionId,
@@ -237,14 +243,15 @@ add_channel_version(Channel, Acc) when is_map(Channel) ->
add_channel_version(_, Acc) ->
Acc.
-spec handle_session_down(reference(), guild_state()) ->
-spec handle_session_down(
reference(), {session_id() | undefined, session_data() | undefined}, guild_state()
) ->
{noreply, guild_state()}.
handle_session_down(Ref, State) ->
Sessions = require_sessions(maps:get(sessions, State, #{})),
{DisconnectingSessionId, DisconnectingSession} = find_session_by_ref(Ref, Sessions, State),
handle_session_down(Ref, {DisconnectingSessionId, DisconnectingSession}, State) ->
Sessions = maps:get(sessions, State, #{}),
DisconnectUserId = disconnect_user_id(DisconnectingSession),
State1 = cleanup_disconnecting_session(DisconnectingSession, State),
NewSessions = remove_session_by_ref(DisconnectingSessionId, Ref, Sessions),
NewSessions = remove_session_id(DisconnectingSessionId, Sessions),
NewState0 = remove_session_ref(Ref, State1#{sessions => NewSessions}),
NewState = track_connected_user(DisconnectUserId, -1, NewState0),
NewState1 = maybe_resection_on_disconnect(
@@ -257,14 +264,6 @@ handle_session_down(Ref, State) ->
disconnect_user_id(#{user_id := UID}) -> UID;
disconnect_user_id(_) -> undefined.
-spec filter_sessions_by_ref(reference(), sessions_map()) ->
sessions_map().
filter_sessions_by_ref(Ref, Sessions) ->
maps:filter(
fun(_K, S) -> maps:get(mref, S) =/= Ref end,
Sessions
).
-spec finish_session_down(sessions_map(), guild_state()) ->
{noreply, guild_state()}.
finish_session_down(NewSessions, State) ->
@@ -287,7 +286,7 @@ maybe_resection_on_disconnect(_, _OldState, NewState) ->
-spec remove_session(session_id(), guild_state()) -> guild_state().
remove_session(SessionId, State) ->
Sessions = require_sessions(maps:get(sessions, State, #{})),
Sessions = maps:get(sessions, State, #{}),
case maps:get(SessionId, Sessions, undefined) of
undefined ->
State;
@@ -302,9 +301,7 @@ do_remove_session(SessionId, Session, State) ->
maybe_demonitor_session(Session),
UserId = maps:get(user_id, Session, undefined),
StateAfterCleanup = cleanup_disconnecting_session(Session, State),
SessionsAfterCleanup = require_sessions(
maps:get(sessions, StateAfterCleanup, #{})
),
SessionsAfterCleanup = maps:get(sessions, StateAfterCleanup, #{}),
NewSessions = maps:remove(SessionId, SessionsAfterCleanup),
State2 = remove_session_ref(
maps:get(mref, Session, undefined), StateAfterCleanup#{sessions => NewSessions}
@@ -322,13 +319,14 @@ maybe_demonitor_session(Session) ->
ok
end.
-spec find_session_by_ref(reference(), sessions_map(), guild_state()) ->
-spec find_session_by_ref(reference(), guild_state()) ->
{session_id() | undefined, session_data() | undefined}.
find_session_by_ref(Ref, Sessions, State) ->
Refs = session_ref_index(State, Sessions),
case maps:get(Ref, Refs, undefined) of
SessionId when is_binary(SessionId) ->
{SessionId, maps:get(SessionId, Sessions, undefined)};
find_session_by_ref(Ref, State) ->
Sessions = maps:get(sessions, State, #{}),
SessionId = maps:get(Ref, maps:get(guild_session_refs, State, #{}), undefined),
case maps:get(SessionId, Sessions, undefined) of
#{mref := Ref} = Session ->
{SessionId, Session};
_ ->
find_session_by_ref_scan(Ref, Sessions)
end.
@@ -337,66 +335,34 @@ find_session_by_ref(Ref, Sessions, State) ->
{session_id() | undefined, session_data() | undefined}.
find_session_by_ref_scan(Ref, Sessions) ->
maps:fold(
fun(SessionId, S, Acc) -> match_ref(SessionId, S, Ref, Acc) end,
fun
(SessionId, #{mref := MRef} = S, _Acc) when MRef =:= Ref -> {SessionId, S};
(_SessionId, _S, Acc) -> Acc
end,
{undefined, undefined},
Sessions
).
-spec match_ref(
session_id(), session_data(), reference(), {
session_id() | undefined, session_data() | undefined
}
) -> {session_id() | undefined, session_data() | undefined}.
match_ref(SessionId, S, Ref, Acc) ->
case maps:get(mref, S) =:= Ref of
true -> {SessionId, S};
false -> Acc
end.
-spec remove_session_by_ref(session_id() | undefined, reference(), sessions_map()) ->
sessions_map().
remove_session_by_ref(SessionId, _Ref, Sessions) when is_binary(SessionId) ->
-spec remove_session_id(session_id() | undefined, sessions_map()) -> sessions_map().
remove_session_id(SessionId, Sessions) when is_binary(SessionId) ->
maps:remove(SessionId, Sessions);
remove_session_by_ref(undefined, Ref, Sessions) ->
filter_sessions_by_ref(Ref, Sessions).
remove_session_id(undefined, Sessions) ->
Sessions.
-spec put_session_ref(session_id(), session_data(), guild_state()) -> guild_state().
put_session_ref(SessionId, Session, State) ->
case maps:get(mref, Session, undefined) of
Ref when is_reference(Ref) ->
Refs0 = session_ref_index(State, require_sessions(maps:get(sessions, State, #{}))),
State#{guild_session_refs => Refs0#{Ref => SessionId}};
_ ->
State
end.
-spec put_session_ref(session_id(), term(), guild_state()) -> guild_state().
put_session_ref(SessionId, Ref, State) when is_reference(Ref) ->
Refs = maps:get(guild_session_refs, State, #{}),
State#{guild_session_refs => Refs#{Ref => SessionId}};
put_session_ref(_SessionId, _Ref, State) ->
State.
-spec remove_session_ref(term(), guild_state()) -> guild_state().
remove_session_ref(Ref, State) when is_reference(Ref) ->
Refs0 = session_ref_index(State, require_sessions(maps:get(sessions, State, #{}))),
State#{guild_session_refs => maps:remove(Ref, Refs0)};
Refs = maps:get(guild_session_refs, State, #{}),
State#{guild_session_refs => maps:remove(Ref, Refs)};
remove_session_ref(_Ref, State) ->
State.
-spec session_ref_index(guild_state(), sessions_map()) -> #{reference() => session_id()}.
session_ref_index(State, Sessions) ->
case maps:get(guild_session_refs, State, undefined) of
Refs when is_map(Refs) -> normalize_session_ref_index(Refs);
_ -> build_session_ref_index(Sessions)
end.
-spec normalize_session_ref_index(map()) -> #{reference() => session_id()}.
normalize_session_ref_index(Refs) ->
maps:fold(
fun
(Ref, SessionId, Acc) when is_reference(Ref), is_binary(SessionId) ->
Acc#{Ref => SessionId};
(_Ref, _SessionId, Acc) ->
Acc
end,
#{},
Refs
).
-spec build_session_ref_index(sessions_map()) -> #{reference() => session_id()}.
build_session_ref_index(Sessions) ->
maps:fold(
@@ -439,31 +405,64 @@ cleanup_disconnecting_session(Session, State) ->
user_id(), guild_state(), guild_state()
) -> guild_state().
maybe_resection_disconnected_user(UserId, OldState, NewState) ->
resection_user_after_connection_change(UserId, OldState, NewState).
case became_disconnected(UserId, OldState, NewState) of
true -> resection_user_after_connection_change(UserId, OldState, NewState);
false -> NewState
end.
-spec resection_connected_user(user_id() | undefined, guild_state(), guild_state()) ->
guild_state().
resection_connected_user(UserId, OldState, NewState) ->
resection_connected_user(UserId, undefined, OldState, NewState).
-spec resection_connected_user(
user_id() | undefined, guild_state(), guild_state()
user_id() | undefined, map() | undefined, guild_state(), guild_state()
) -> guild_state().
resection_connected_user(UserId, OldState, NewState) when
resection_connected_user(UserId, PresenceBefore, OldState, NewState) when
is_integer(UserId), UserId > 0
->
case became_connected(UserId, OldState, NewState) of
true ->
ResectionedState = resection_user_after_connection_change(
UserId, OldState, NewState
ResectionedState = resection_after_connect(
UserId, PresenceBefore, OldState, NewState
),
_ = guild_presence_reconcile:maybe_schedule_user_repair(UserId, ResectionedState),
ResectionedState;
false ->
NewState
end;
resection_connected_user(_UserId, _OldState, NewState) ->
resection_connected_user(_UserId, _PresenceBefore, _OldState, NewState) ->
NewState.
-spec resection_after_connect(
user_id(), map() | undefined, guild_state(), guild_state()
) -> guild_state().
resection_after_connect(UserId, PresenceBefore, OldState, NewState) ->
case presence_resynced_lists(UserId, PresenceBefore, NewState) of
true -> resection_user_already_synced(UserId, NewState);
false -> resection_user_after_connection_change(UserId, OldState, NewState)
end.
-spec presence_resynced_lists(user_id(), map() | undefined, guild_state()) -> boolean().
presence_resynced_lists(_UserId, undefined, _State) ->
false;
presence_resynced_lists(UserId, PresenceBefore, State) ->
PresenceAfter = guild_member_list_connected:resolve_presence_for_user(State, UserId),
guild_member_list_write:presence_change_resyncs_lists(PresenceBefore, PresenceAfter).
-spec resection_user_already_synced(user_id(), guild_state()) -> guild_state().
resection_user_already_synced(UserId, State) ->
_ = guild_presence:sync_online_status(UserId, State),
guild_member_list_write:queue_synced_connection_change(UserId, State).
-spec became_connected(user_id(), guild_state(), guild_state()) -> boolean().
became_connected(UserId, OldState, NewState) ->
(not user_connected(UserId, OldState)) andalso user_connected(UserId, NewState).
-spec became_disconnected(user_id(), guild_state(), guild_state()) -> boolean().
became_disconnected(UserId, OldState, NewState) ->
user_connected(UserId, OldState) andalso not user_connected(UserId, NewState).
-spec user_connected(user_id(), guild_state()) -> boolean().
user_connected(UserId, State) ->
sets:is_element(UserId, guild_member_list_connected:connected_session_user_ids(State)).
@@ -501,7 +500,7 @@ track_connected_user(UserId, _Delta, State) when
State;
track_connected_user(UserId, Delta, State) ->
Counts = require_map(maps:get(user_session_counts, State, #{})),
Connected = require_set(maps:get(connected_user_ids, State, sets:new())),
Connected = maps:get(connected_user_ids, State, sets:new()),
OldCount = require_non_neg(maps:get(UserId, Counts, 0)),
NewCount = max(0, OldCount + Delta),
{NC, NConn} = apply_count_change(
@@ -566,28 +565,10 @@ viewable_channels_cache_table(_) ->
require_map(M) when is_map(M) -> M;
require_map(_) -> #{}.
-spec require_sessions(term()) -> sessions_map().
require_sessions(M) when is_map(M) ->
maps:fold(fun require_session_entry/3, #{}, M);
require_sessions(_) ->
#{}.
-spec require_session_entry(term(), term(), sessions_map()) -> sessions_map().
require_session_entry(K, V, Acc) when is_binary(K), is_map(V) ->
Acc#{K => V};
require_session_entry(_, _, Acc) ->
Acc.
-spec require_guild_id(term()) -> guild_id().
require_guild_id(Id) when is_integer(Id), Id > 0 -> Id;
require_guild_id(_) -> error(badarg).
-spec require_set(term()) -> sets:set(integer()).
require_set(S) when is_map(S) ->
sets:from_list([I || I <- maps:keys(S), is_integer(I)]);
require_set(_) ->
sets:new().
-spec require_non_neg(term()) -> non_neg_integer().
require_non_neg(V) when is_integer(V), V >= 0 -> V;
require_non_neg(_) -> 0.
@@ -645,16 +626,41 @@ resection_connected_user_skips_when_already_connected_test() ->
Connected = sets:from_list([42]),
Old = #{connected_user_ids => Connected},
New = #{connected_user_ids => Connected, marker => updated},
?assertEqual(New, resection_connected_user(42, Old, New)).
?assertEqual(New, resection_connected_user(42, #{}, Old, New)).
resection_connected_user_skips_when_still_disconnected_test() ->
Old = #{connected_user_ids => sets:new()},
New = #{connected_user_ids => sets:new(), marker => updated},
?assertEqual(New, resection_connected_user(42, Old, New)).
?assertEqual(New, resection_connected_user(42, #{}, Old, New)).
resection_connected_user_ignores_undefined_user_test() ->
New = #{connected_user_ids => sets:new()},
?assertEqual(New, resection_connected_user(undefined, New, New)).
?assertEqual(New, resection_connected_user(undefined, #{}, New, New)).
resection_disconnected_user_skips_when_still_connected_test() ->
Connected = sets:from_list([42]),
Old = #{connected_user_ids => Connected},
New = #{connected_user_ids => Connected, marker => updated},
?assertEqual(New, maybe_resection_disconnected_user(42, Old, New)).
resection_disconnected_user_skips_when_already_disconnected_test() ->
Old = #{connected_user_ids => sets:new()},
New = #{connected_user_ids => sets:new(), marker => updated},
?assertEqual(New, maybe_resection_disconnected_user(42, Old, New)).
became_disconnected_detects_transition_test() ->
Old = #{connected_user_ids => sets:from_list([42])},
New = #{connected_user_ids => sets:new()},
?assert(became_disconnected(42, Old, New)),
?assertNot(became_disconnected(42, Old, Old)),
?assertNot(became_disconnected(42, New, New)),
?assertNot(became_disconnected(42, New, Old)).
presence_resynced_lists_needs_a_list_visible_change_test() ->
State = #{member_presence => #{42 => #{<<"status">> => <<"online">>}}},
?assertNot(presence_resynced_lists(42, undefined, State)),
?assert(presence_resynced_lists(42, #{}, State)),
?assertNot(presence_resynced_lists(42, #{<<"status">> => <<"online">>}, State)).
became_connected_detects_transition_test() ->
Old = #{connected_user_ids => sets:new()},
@@ -19,8 +19,13 @@
-spec cleanup_connect_admission_for_session(session_id(), guild_state()) -> guild_state().
cleanup_connect_admission_for_session(SessionId, State) ->
State1 = cleanup_connect_pending(SessionId, State),
cleanup_connect_queue(SessionId, State1).
case maps:get(session_connect_pending, State, undefined) of
#{SessionId := _} = Pending ->
State1 = State#{session_connect_pending => maps:remove(SessionId, Pending)},
cleanup_connect_queue(SessionId, State1);
_ ->
State
end.
-spec normalize_connect_queue(term()) -> queue:queue() | undefined.
normalize_connect_queue({In, Out}) when is_list(In), is_list(Out) ->
@@ -51,15 +56,6 @@ clear_auto_stop_pending(State) ->
maps:remove(auto_stop_pending, State)
end.
-spec cleanup_connect_pending(session_id(), guild_state()) -> guild_state().
cleanup_connect_pending(SessionId, State) ->
case maps:get(session_connect_pending, State, undefined) of
Pending when is_map(Pending) ->
State#{session_connect_pending => maps:remove(SessionId, Pending)};
_ ->
State
end.
-spec cleanup_connect_queue(session_id(), guild_state()) -> guild_state().
cleanup_connect_queue(SessionId, State) ->
Queue0 = maps:get(session_connect_queue, State, undefined),
@@ -115,7 +115,7 @@ cleanup_removed_member_sessions(State) ->
end,
Sessions
),
State#{sessions => FilteredSessions}.
drop_removed_session_refs(Sessions, FilteredSessions, State#{sessions => FilteredSessions}).
-spec cleanup_removed_member_sessions(user_id() | undefined, guild_state()) -> guild_state().
cleanup_removed_member_sessions(UserId, State) when is_integer(UserId), UserId > 0 ->
@@ -126,10 +126,27 @@ cleanup_removed_member_sessions(UserId, State) when is_integer(UserId), UserId >
end,
Sessions
),
State#{sessions => FilteredSessions};
drop_removed_session_refs(Sessions, FilteredSessions, State#{sessions => FilteredSessions});
cleanup_removed_member_sessions(_UserId, State) ->
cleanup_removed_member_sessions(State).
-spec drop_removed_session_refs(map(), map(), guild_state()) -> guild_state().
drop_removed_session_refs(Sessions, FilteredSessions, State) ->
maps:fold(
fun(SessionId, Session, Acc) ->
case maps:is_key(SessionId, FilteredSessions) of
true ->
Acc;
false ->
guild_sessions_connect:remove_session_ref(
maps:get(mref, Session, undefined), Acc
)
end
end,
State,
Sessions
).
-spec maybe_disconnect_removed_member(user_id() | undefined, guild_state()) -> guild_state().
maybe_disconnect_removed_member(UserId, State) when is_integer(UserId), UserId > 0 ->
ok = guild_voice_lifecycle:cast_disconnect_voice_user(UserId, State),
@@ -232,6 +249,27 @@ cleanup_removed_member_sessions_for_user_removes_only_that_user_test() ->
#{} = NewSessions = maps:get(sessions, Result),
?assertEqual([<<"s1">>, <<"s3">>], lists:sort(maps:keys(NewSessions))).
cleanup_removed_member_sessions_drops_removed_session_refs_test() ->
R1 = make_ref(),
R2 = make_ref(),
R3 = make_ref(),
Sessions = #{
<<"s1">> => #{user_id => 1, mref => R1},
<<"s2">> => #{user_id => 2, mref => R2},
<<"s3">> => #{user_id => 2, mref => R3}
},
Refs = #{R1 => <<"s1">>, R2 => <<"s2">>, R3 => <<"s3">>},
Data = #{<<"members">> => #{1 => #{<<"user">> => #{<<"id">> => <<"1">>}}}},
State = #{data => Data, sessions => Sessions, guild_session_refs => Refs},
?assertEqual(
#{R1 => <<"s1">>},
maps:get(guild_session_refs, cleanup_removed_member_sessions(2, State))
),
?assertEqual(
#{R1 => <<"s1">>},
maps:get(guild_session_refs, cleanup_removed_member_sessions(State))
).
sync_member_updates_loaded_channel_engines_test() ->
GuildId = 100,
UserId = 10,
@@ -378,15 +378,32 @@ filter_member_ids_for_subscription(_GuildId, SessionUserId, MemberIds, State) ->
-spec handle_added_subscriptions([user_id()], session_id(), guild_state()) -> guild_state().
handle_added_subscriptions(Added, SessionId, State) ->
Prefetched = guild_presence:cached_presences(
[UserId || UserId <- Added, presence_watched(UserId, State)]
),
lists:foldl(
fun(UserId, Acc) ->
StateWithPresence = guild_sessions:subscribe_to_user_presence(UserId, Acc),
guild_presence:send_cached_presence_to_session(UserId, SessionId, StateWithPresence)
end,
fun(UserId, Acc) -> subscribe_added_member(UserId, SessionId, Prefetched, Acc) end,
State,
Added
).
-spec subscribe_added_member(user_id(), session_id(), map(), guild_state()) -> guild_state().
subscribe_added_member(UserId, SessionId, Prefetched, State) ->
Watched = presence_watched(UserId, State),
StateWithPresence = guild_sessions:subscribe_to_user_presence(UserId, State),
case {Watched, maps:find(UserId, Prefetched)} of
{true, {ok, Lookup}} ->
guild_presence:send_presence_lookup_to_session(
UserId, SessionId, Lookup, StateWithPresence
);
_ ->
guild_presence:send_cached_presence_to_session(UserId, SessionId, StateWithPresence)
end.
-spec presence_watched(user_id(), guild_state()) -> boolean().
presence_watched(UserId, State) ->
maps:get(UserId, maps:get(presence_subscriptions, State, #{}), 0) > 0.
-spec handle_removed_subscriptions([user_id()], guild_state()) -> guild_state().
handle_removed_subscriptions(Removed, State) ->
lists:foldl(
@@ -641,4 +658,151 @@ arm_lazy_subscribe_timer_is_idempotent_test() ->
State = #{lazy_subscribe_timer => Ref},
?assertEqual(Ref, maps:get(lazy_subscribe_timer, arm_lazy_subscribe_timer(State))).
member_subscription_test_member(UserId, RoleIds) ->
#{
<<"user">> => #{<<"id">> => integer_to_binary(UserId)},
<<"roles">> => [integer_to_binary(RoleId) || RoleId <- RoleIds]
}.
member_subscription_test_presence(UserId) ->
#{<<"status">> => <<"online">>, <<"user">> => #{<<"id">> => integer_to_binary(UserId)}}.
member_subscription_test_state() ->
GuildId = 42,
View = integer_to_binary(constants:view_channel_permission()),
Channels = [
#{
<<"id">> => <<"500">>,
<<"type">> => 0,
<<"permission_overwrites">> => [
#{
<<"id">> => <<"3000">>,
<<"type">> => 0,
<<"allow">> => <<"0">>,
<<"deny">> => View
}
]
}
],
Members = maps:from_list(
[{Id, member_subscription_test_member(Id, [])} || Id <- [10 | lists:seq(20, 27)]] ++
[{31, member_subscription_test_member(31, [3000])}]
),
Tab = ets:new(member_subscription_test_presence, [set, public]),
ets:insert(
Tab,
{23,
presence_payload:build(
maps:get(<<"user">>, maps:get(23, Members)), <<"online">>, false, false, null
)}
),
Subs = lists:foldl(
fun(UserId, Acc) -> guild_subscriptions:subscribe(<<"s1">>, UserId, Acc) end,
guild_subscriptions:init_state(),
[20, 27]
),
#{
id => GuildId,
data => #{
<<"guild">> => #{<<"owner_id">> => <<"7">>},
<<"roles">> => [
#{<<"id">> => integer_to_binary(GuildId), <<"permissions">> => View},
#{<<"id">> => <<"3000">>, <<"permissions">> => <<"0">>}
],
<<"members">> => Members,
<<"channels">> => Channels,
<<"channel_index">> => guild_data_index:build_id_index(Channels)
},
sessions => #{
<<"s1">> => #{user_id => 10, pid => self(), viewable_channels => #{500 => true}}
},
member_subscriptions => Subs,
presence_subscriptions => #{20 => 1, 21 => 1, 22 => 2, 27 => 1},
member_presence => Tab
}.
reference_update_member_subscriptions(SessionId, MemberIds, State) ->
#{user_id := Viewer, viewable_channels := SessionMap} = maps:get(
SessionId, maps:get(sessions, State)
),
Filtered = [
MemberId
|| MemberId <- MemberIds,
MemberId =/= Viewer,
lists:any(
fun(ChannelId) -> maps:is_key(ChannelId, SessionMap) end,
guild_visibility:get_user_viewable_channels(MemberId, State)
)
],
{NewSubs, Added, Removed} = guild_subscriptions:update_subscriptions_with_delta(
SessionId, Filtered, maps:get(member_subscriptions, State)
),
State1 = lists:foldl(
fun(UserId, Acc) ->
StateWithPresence = guild_sessions:subscribe_to_user_presence(UserId, Acc),
guild_presence:send_cached_presence_to_session(UserId, SessionId, StateWithPresence)
end,
State#{member_subscriptions => NewSubs},
Added
),
handle_removed_subscriptions(Removed, State1).
drain_mailbox(Acc) ->
receive
Msg -> drain_mailbox([Msg | Acc])
after 50 -> lists:reverse(Acc)
end.
ensure_started(Name, Start) ->
case whereis(Name) of
undefined ->
case Start() of
{ok, _} -> ok;
{error, {already_started, _}} -> ok
end;
_ ->
ok
end.
update_member_subscriptions_matches_reference_test() ->
ensure_started(presence_bus, fun presence_bus:start_link/0),
ensure_started(presence_cache, fun presence_cache:start_link/0),
[ok = presence_cache:put(Id, member_subscription_test_presence(Id)) || Id <- [20, 21, 23]],
_ = sys:get_state(presence_cache),
?assertMatch({ok, _}, presence_cache:get(21)),
MemberIds = [20, 21, 22, 23, 24, 25, 31, 99, 10, 21],
_ = drain_mailbox([]),
Expected = reference_update_member_subscriptions(
<<"s1">>, MemberIds, member_subscription_test_state()
),
ExpectedDispatches = drain_mailbox([]),
Actual = handle_update_member_subscriptions_local(
42, <<"s1">>, MemberIds, member_subscription_test_state()
),
ActualDispatches = drain_mailbox([]),
Strip = fun(S) -> maps:without([member_presence], S) end,
?assertEqual(Strip(Expected), Strip(Actual)),
?assertEqual(ExpectedDispatches, ActualDispatches),
?assertEqual(2, length(ActualDispatches)),
?assertEqual(
[20, 21, 22, 23, 24, 25],
lists:sort(
sets:to_list(
guild_subscriptions:get_user_ids_for_session(
<<"s1">>, maps:get(member_subscriptions, Actual)
)
)
)
).
cached_presences_reports_every_requested_user_test() ->
ensure_started(presence_cache, fun presence_cache:start_link/0),
ok = presence_cache:put(61, member_subscription_test_presence(61)),
_ = sys:get_state(presence_cache),
?assertEqual(
#{61 => presence_cache_safe:get(61), 62 => not_found},
guild_presence:cached_presences([61, 62])
),
?assertEqual(#{}, guild_presence:cached_presences([])).
-endif.
@@ -234,8 +234,7 @@ overwrite_target_id(_Overwrite, Acc) ->
-spec has_mutual_channel(user_id(), map(), guild_state()) -> boolean().
has_mutual_channel(MemberId, SessionMap, State) ->
MemberChannels = guild_visibility:get_user_viewable_channels(MemberId, State),
has_shared_channel(MemberChannels, SessionMap).
guild_visibility_channels:shares_viewable_channel(MemberId, SessionMap, State).
-spec has_shared_channel([integer()], map()) -> boolean().
has_shared_channel(MemberChannels, SessionMap) ->
@@ -298,8 +297,6 @@ test_state() ->
}
}.
%% filter_member_ids/3 as it read before the memo: every candidate materialises its
%% own complete viewable channel list.
reference_filter_member_ids(SessionUserId, MemberIds, State) ->
SessionMap = session_channel_map(SessionUserId, State),
lists:filtermap(
@@ -312,7 +309,8 @@ reference_keep(MemberId, SessionUserId, _SessionMap, _State) when
->
false;
reference_keep(MemberId, _SessionUserId, SessionMap, State) ->
case has_mutual_channel(MemberId, SessionMap, State) of
MemberChannels = guild_visibility:get_user_viewable_channels(MemberId, State),
case has_shared_channel(MemberChannels, SessionMap) of
true -> {true, MemberId};
false -> false
end.
@@ -5,6 +5,7 @@
-export([
get_user_viewable_channels/2,
shares_viewable_channel/3,
viewable_channel_set/2,
have_shared_viewable_channel/3,
viewable_channel_map/1,
@@ -36,16 +37,18 @@ get_user_viewable_channels(UserId, State) ->
undefined ->
[];
_ ->
compute_viewable_with_categories(UserId, Member, Channels, State)
Base = guild_permissions:member_base_permissions(UserId, Member, State),
compute_viewable_with_categories(UserId, Base, Channels, State)
end.
-spec compute_viewable_with_categories(
user_id(), map(), [map()], guild_state()
user_id(), guild_permissions:base_permissions(), [map()], guild_state()
) -> [channel_id()].
compute_viewable_with_categories(UserId, Member, Channels, State) ->
compute_viewable_with_categories(UserId, Base, Channels, State) ->
Viewable = guild_permissions:viewable_channel_ids(UserId, Base, Channels, State),
{ViewableIds, ViewableIdSet, NeededParentIds} = lists:foldl(
fun(Channel, {Ids, IdSet, Parents}) ->
collect_viewable_channel(Channel, UserId, Member, State, {Ids, IdSet, Parents})
collect_viewable_channel(Channel, Viewable, {Ids, IdSet, Parents})
end,
{[], #{}, #{}},
Channels
@@ -59,6 +62,38 @@ compute_viewable_with_categories(UserId, Member, Channels, State) ->
lists:reverse(ViewableIds) ++ ExtraIds
end.
-spec shares_viewable_channel(user_id(), map(), guild_state()) -> boolean().
shares_viewable_channel(UserId, ChannelMap, State) ->
Data = map_utils:ensure_map(map_utils:get_safe(State, data, #{})),
Channels = map_utils:ensure_list(maps:get(<<"channels">>, Data, [])),
case guild_permissions:find_member_by_user_id(UserId, State) of
undefined ->
false;
Member ->
lists:any(
fun(Channel) ->
shares_channel_or_parent(
Channel, UserId, Member, ChannelMap, Channels, State
)
end,
Channels
)
end.
-spec shares_channel_or_parent(map(), user_id(), map(), map(), [map()], guild_state()) ->
boolean().
shares_channel_or_parent(Channel, UserId, Member, ChannelMap, Channels, State) ->
channel_viewable(UserId, Member, Channel, State) andalso
(maps:is_key(channel_id(Channel), ChannelMap) orelse
shared_parent(channel_parent_id(Channel), ChannelMap, Channels)).
-spec shared_parent(channel_id() | undefined, map(), [map()]) -> boolean().
shared_parent(undefined, _ChannelMap, _Channels) ->
false;
shared_parent(ParentId, ChannelMap, Channels) ->
maps:is_key(ParentId, ChannelMap) andalso
lists:any(fun(C) -> channel_id(C) =:= ParentId end, Channels).
-spec collect_missing_parent_ids([map()], map()) -> [channel_id()].
collect_missing_parent_ids(Channels, MissingParents) ->
lists:filtermap(
@@ -79,10 +114,10 @@ check_missing_parent(CId, MissingParents) when is_integer(CId) ->
check_missing_parent(_, _) ->
false.
-spec collect_viewable_channel(map(), user_id(), map(), guild_state(), {list(), map(), map()}) ->
-spec collect_viewable_channel(map(), #{channel_id() => true}, {list(), map(), map()}) ->
{list(), map(), map()}.
collect_viewable_channel(Channel, UserId, Member, State, {Ids, IdSet, Parents}) ->
case channel_viewable(UserId, Member, Channel, State) of
collect_viewable_channel(Channel, Viewable, {Ids, IdSet, Parents}) ->
case maps:is_key(channel_id(Channel), Viewable) of
false ->
{Ids, IdSet, Parents};
true ->
@@ -14,6 +14,7 @@
delete/1,
get/1,
bulk_get/1,
bulk_get_map/1,
get_memory_stats/0,
pending_handoff_count/0,
get_pending_handoff_count/0,
@@ -72,6 +73,13 @@ bulk_get(UserIds) when is_list(UserIds) ->
false -> presence_cache_bulk:bulk_get_inner(UserIds)
end.
-spec bulk_get_map([integer()]) -> #{integer() => map()}.
bulk_get_map(UserIds) when is_list(UserIds) ->
case persistent_term:get(presence_noop, false) of
true -> #{};
false -> presence_cache_bulk:bulk_get_map_inner(UserIds)
end.
-spec get_memory_stats() -> {ok, map()} | {error, term()}.
get_memory_stats() ->
case presence_cache_api:safe_call_if_enabled(get_memory_stats, {error, not_available}) of
@@ -7,6 +7,7 @@
-export([
bulk_get_inner/1,
bulk_get_map_inner/1,
get_from_cluster/1,
get_local_fast/1,
local_bulk_presence_map/1,
@@ -30,11 +31,15 @@
-spec bulk_get_inner([integer()]) -> [map()].
bulk_get_inner(UserIds) ->
presence_values(bulk_get_map_inner(UserIds)).
-spec bulk_get_map_inner([integer()]) -> #{integer() => map()}.
bulk_get_map_inner(UserIds) ->
UniqueUserIds = normalize_user_ids(UserIds),
PrimaryPresenceMap = fetch_primary_presences(UniqueUserIds),
MissingUserIds = [U || U <- UniqueUserIds, not maps:is_key(U, PrimaryPresenceMap)],
FallbackPresenceMap = fetch_fallback_presences(MissingUserIds),
presence_values(maps:merge(PrimaryPresenceMap, FallbackPresenceMap)).
maps:merge(PrimaryPresenceMap, FallbackPresenceMap).
-spec get_from_cluster(integer()) -> {ok, map()} | not_found.
get_from_cluster(UserId) ->
@@ -56,6 +56,8 @@ parse_optional(Value) ->
parse(Value).
-spec parse_maybe(term()) -> t() | undefined.
parse_maybe(Value) when is_integer(Value), Value > 0, Value =< ?MAX_SNOWFLAKE ->
Value;
parse_maybe(Value) ->
try parse_optional(Value) of
Id -> Id
@@ -6,6 +6,7 @@
-include_lib("eunit/include/eunit.hrl").
-define(GENERATIONAL_FULLSWEEP_AFTER, 10).
-define(GUILD_FULLSWEEP_AFTER, 100).
session_init_keeps_generational_gc_test() ->
?assertEqual(
@@ -21,13 +22,13 @@ session_code_change_keeps_generational_gc_test() ->
guild_init_keeps_generational_gc_test() ->
?assertEqual(
?GENERATIONAL_FULLSWEEP_AFTER,
?GUILD_FULLSWEEP_AFTER,
fullsweep_after_in(fun() -> guild:init(guild_data()) end)
).
guild_code_change_keeps_generational_gc_test() ->
?assertEqual(
?GENERATIONAL_FULLSWEEP_AFTER,
?GUILD_FULLSWEEP_AFTER,
fullsweep_after_in(fun() -> guild:code_change(0, #{}, []) end)
).
@@ -327,10 +327,10 @@ run_send_member_list_update_encodes_wire_payload() ->
?assert(false)
end.
broadcast_presence_delta_outside_list_payload_skips_sync_test() ->
with_sync_dispatch_mock(fun run_broadcast_presence_delta_outside_list_payload_skips_sync/0).
broadcast_afk_only_change_skips_sync_test() ->
with_sync_dispatch_mock(fun run_broadcast_afk_only_change_skips_sync/0).
run_broadcast_presence_delta_outside_list_payload_skips_sync() ->
run_broadcast_afk_only_change_skips_sync() ->
Ref = guild_member_list_engine:new(),
try
State = presence_delta_state(Ref),
@@ -338,8 +338,8 @@ run_broadcast_presence_delta_outside_list_payload_skips_sync() ->
1,
State,
State,
presence_map(<<"online">>, false, null),
presence_map(<<"online">>, true, null)
presence_map(<<"online">>, false, false, null),
presence_map(<<"online">>, false, true, null)
),
?assertEqual(State, NewState),
assert_no_sync_dispatch()
@@ -416,10 +416,13 @@ presence_delta_state(Ref) ->
channel_list_state(Ref, make_subs_tab([{<<"500">>, <<"s1">>, [{0, 99}]}]), [Member]).
presence_map(Status, Mobile, CustomStatus) ->
presence_map(Status, Mobile, false, CustomStatus).
presence_map(Status, Mobile, Afk, CustomStatus) ->
#{
<<"status">> => Status,
<<"mobile">> => Mobile,
<<"afk">> => false,
<<"afk">> => Afk,
<<"custom_status">> => CustomStatus
}.
@@ -100,7 +100,11 @@ cleanup_connect_admission_queue_format_test() ->
#{request => #{session_id => S3}, attempt => 3}
]),
State1 = guild_sessions_connect_cleanup:cleanup_connect_admission_for_session(
S1, #{sessions => #{}, session_connect_queue => Queue}
S1, #{
sessions => #{},
session_connect_pending => #{S1 => 2},
session_connect_queue => Queue
}
),
ResultQueue = queue:to_list(maps:get(session_connect_queue, State1)),
SessionIds = [
+5
View File
@@ -0,0 +1,5 @@
// SPDX-License-Identifier: AGPL-3.0-or-later
export const SSO_MOBILE_CALLBACK_URI = 'fluxer://auth/sso/callback';
export const SSO_MOBILE_STATE_PREFIX = 'm.';