Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 23 additions & 28 deletions lib/realtime/tenants/authorization.ex
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,8 @@ defmodule Realtime.Tenants.Authorization do
:sub => binary | nil
}

@type extension :: :broadcast | :presence

@doc """
Builds a new authorization struct which will be used to retain the information required to check Policies.

Expand Down Expand Up @@ -114,38 +116,38 @@ defmodule Realtime.Tenants.Authorization do

Automatically uses RPC if the database connection is not in the same node
"""
@spec get_write_authorizations(Policies.t(), pid(), t(), keyword()) ::
@spec get_write_authorizations(Policies.t(), pid(), t(), extension()) ::
{:ok, Policies.t()}
| {:error, :rls_policy_error, Postgrex.Error.t()}
| {:error, :query_canceled, Postgrex.Error.t()}
| {:error, :missing_partition}
| {:error, :increase_connection_pool}
| {:error, :tenant_database_unavailable}
| {:error, any()}
def get_write_authorizations(policies, db_conn, authorization_context, opts \\ [])

def get_write_authorizations(policies, db_conn, authorization_context, opts) when node() == node(db_conn) do
def get_write_authorizations(policies, db_conn, authorization_context, extension)
when extension in [:broadcast, :presence] and node() == node(db_conn) do
rate_counter = rate_counter(authorization_context.tenant_id)

if rate_counter.limit.triggered == false do
db_conn
|> get_write_policies_for_connection(authorization_context, policies, opts)
|> get_write_policies_for_connection(authorization_context, policies, extension)
|> handle_policies_result(rate_counter)
else
{:error, :increase_connection_pool}
end
end

# Remote call
def get_write_authorizations(policies, db_conn, authorization_context, opts) do
def get_write_authorizations(policies, db_conn, authorization_context, extension)
when extension in [:broadcast, :presence] do
rate_counter = rate_counter(authorization_context.tenant_id)

if rate_counter.limit.triggered == false do
case GenRpc.call(
node(db_conn),
__MODULE__,
:get_write_authorizations,
[policies, db_conn, authorization_context, opts],
[policies, db_conn, authorization_context, extension],
tenant_id: authorization_context.tenant_id,
key: authorization_context.tenant_id
) do
Expand All @@ -164,8 +166,8 @@ defmodule Realtime.Tenants.Authorization do
end
end

def get_write_authorizations(db_conn, authorization_context),
do: get_write_authorizations(%Policies{}, db_conn, authorization_context)
def get_write_authorizations(db_conn, authorization_context, extension),
do: get_write_authorizations(%Policies{}, db_conn, authorization_context, extension)

defp handle_policies_result(result, rate_counter) do
case result do
Expand Down Expand Up @@ -279,18 +281,17 @@ defmodule Realtime.Tenants.Authorization do
)
end

defp get_write_policies_for_connection(conn, authorization_context, policies, caller_opts) do
defp get_write_policies_for_connection(conn, authorization_context, policies, extension) do
tenant_id = authorization_context.tenant_id
opts = [telemetry: [:realtime, :tenants, :write_authorization_check], tenant_id: tenant_id]
metadata = [project: tenant_id, external_id: tenant_id]
extensions = extensions_to_check(caller_opts)

Database.transaction(
conn,
fn transaction_conn ->
set_conn_config(transaction_conn, authorization_context)

with {:ok, policies} <- check_write_policies(transaction_conn, authorization_context, extensions, policies) do
with {:ok, policies} <- check_write_policy(transaction_conn, authorization_context, extension, policies) do
Postgrex.query!(transaction_conn, "ROLLBACK AND CHAIN", [])
policies
else
Expand Down Expand Up @@ -328,25 +329,19 @@ defmodule Realtime.Tenants.Authorization do
end
end

defp check_write_policies(conn, authorization_context, extensions, policies) do
Enum.reduce_while(@all_extensions, {:ok, policies}, fn extension, {:ok, acc} ->
if extension in extensions do
changeset = Message.changeset(%Message{}, %{topic: authorization_context.topic, extension: extension})
defp check_write_policy(conn, authorization_context, extension, policies) do
changeset = Message.changeset(%Message{}, %{topic: authorization_context.topic, extension: extension})

case Repo.insert(conn, changeset, Message, mode: :savepoint, returning: false) do
{:ok, _} ->
{:cont, {:ok, Policies.update_policies(acc, extension, :write, true)}}
case Repo.insert(conn, changeset, Message, mode: :savepoint, returning: false) do
{:ok, _} ->
{:ok, Policies.update_policies(policies, extension, :write, true)}

{:error, %Postgrex.Error{postgres: %{code: :insufficient_privilege}}} ->
{:cont, {:ok, Policies.update_policies(acc, extension, :write, false)}}
{:error, %Postgrex.Error{postgres: %{code: :insufficient_privilege}}} ->
{:ok, Policies.update_policies(policies, extension, :write, false)}

{:error, reason} ->
{:halt, {:error, reason}}
end
else
{:cont, {:ok, Policies.update_policies(acc, extension, :write, false)}}
end
end)
{:error, reason} ->
{:error, reason}
end
end

defp rate_counter(tenant_id) do
Expand Down
2 changes: 1 addition & 1 deletion lib/realtime/tenants/batch_broadcast.ex
Original file line number Diff line number Diff line change
Expand Up @@ -159,7 +159,7 @@ defmodule Realtime.Tenants.BatchBroadcast do
|> Map.put(:topic, topic)
|> Authorization.build_authorization_params()

case Authorization.get_write_authorizations(db_conn, auth_params) do
case Authorization.get_write_authorizations(db_conn, auth_params, :broadcast) do
{:ok, policies} -> policies
{:error, :not_found} -> nil
error -> error
Expand Down
2 changes: 1 addition & 1 deletion lib/realtime/tenants/single_broadcast.ex
Original file line number Diff line number Diff line change
Expand Up @@ -204,7 +204,7 @@ defmodule Realtime.Tenants.SingleBroadcast do
defp permissions_for_message(tenant, auth_params, topic) do
with {:ok, db_conn} <- Connect.lookup_or_start_connection(tenant.external_id) do
auth_params = %{auth_params | topic: topic}
Authorization.get_write_authorizations(db_conn, auth_params)
Authorization.get_write_authorizations(db_conn, auth_params, :broadcast)
end
end

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -172,7 +172,7 @@ defmodule RealtimeWeb.RealtimeChannel.BroadcastHandler do
db_conn,
authorization_context
) do
Authorization.get_write_authorizations(policies, db_conn, authorization_context)
Authorization.get_write_authorizations(policies, db_conn, authorization_context, :broadcast)
end

defp run_authorization_check(socket, _db_conn, _authorization_context) do
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -91,7 +91,12 @@ defmodule RealtimeWeb.RealtimeChannel.PresenceHandler do
# presence_diff for this socket, then authorize presence.write.
with {:ok, policies} <- maybe_authorize_presence_read(policies, db_conn, authorization_context),
{:ok, policies} <-
Authorization.get_write_authorizations(policies, db_conn, authorization_context, presence_enabled?: true) do
Authorization.get_write_authorizations(
policies,
db_conn,
authorization_context,
:presence
) do
socket = assign(socket, :policies, policies)
handle_presence_event("track", payload, db_conn, socket)
else
Expand Down
3 changes: 2 additions & 1 deletion test/realtime/monitoring/prom_ex/plugins/tenant_test.exs
Original file line number Diff line number Diff line change
Expand Up @@ -217,7 +217,8 @@ defmodule Realtime.PromEx.Plugins.TenantTest do
Authorization.get_write_authorizations(
%Policies{},
context.db_conn,
context.authorization_context
context.authorization_context,
:broadcast
)

# Wait enough time for the poll rate to be triggered at least once
Expand Down
40 changes: 31 additions & 9 deletions test/realtime/tenants/authorization_remote_test.exs
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,16 @@ defmodule Realtime.Tenants.AuthorizationRemoteTest do
Authorization.get_write_authorizations(
policies,
context.db_conn,
context.authorization_context
context.authorization_context,
:broadcast
)

{:ok, policies} =
Authorization.get_write_authorizations(
policies,
context.db_conn,
context.authorization_context,
:presence
)

assert %Policies{
Expand Down Expand Up @@ -72,7 +81,16 @@ defmodule Realtime.Tenants.AuthorizationRemoteTest do
Authorization.get_write_authorizations(
policies,
context.db_conn,
context.authorization_context
context.authorization_context,
:broadcast
)

{:ok, policies} =
Authorization.get_write_authorizations(
policies,
context.db_conn,
context.authorization_context,
:presence
)

assert %Policies{
Expand All @@ -90,7 +108,7 @@ defmodule Realtime.Tenants.AuthorizationRemoteTest do
Authorization.get_read_authorizations(%Policies{}, db_conn, context.authorization_context)

{:error, :increase_connection_pool} =
Authorization.get_write_authorizations(%Policies{}, db_conn, context.authorization_context)
Authorization.get_write_authorizations(%Policies{}, db_conn, context.authorization_context, :broadcast)
end

@tag role: "anon", policies: []
Expand Down Expand Up @@ -125,15 +143,15 @@ defmodule Realtime.Tenants.AuthorizationRemoteTest do
capture_log(fn ->
for _ <- 1..6 do
{:error, :increase_connection_pool} =
Authorization.get_write_authorizations(%Policies{}, pid, context.authorization_context)
Authorization.get_write_authorizations(%Policies{}, pid, context.authorization_context, :broadcast)
end

rate_counter = Realtime.Tenants.authorization_errors_per_second_rate(context.tenant)
RateCounterHelper.tick!(rate_counter)

for _ <- 1..10 do
{:error, :increase_connection_pool} =
Authorization.get_write_authorizations(%Policies{}, pid, context.authorization_context)
Authorization.get_write_authorizations(%Policies{}, pid, context.authorization_context, :broadcast)
end
end)

Expand Down Expand Up @@ -177,7 +195,8 @@ defmodule Realtime.Tenants.AuthorizationRemoteTest do
Authorization.get_write_authorizations(
%Policies{},
context.db_conn,
context.authorization_context
context.authorization_context,
:broadcast
)
end)

Expand Down Expand Up @@ -208,7 +227,8 @@ defmodule Realtime.Tenants.AuthorizationRemoteTest do
Authorization.get_write_authorizations(
%Policies{},
context.db_conn,
context.authorization_context
context.authorization_context,
:presence
)

assert {:error, :rls_policy_error, %Postgrex.Error{}} =
Expand All @@ -222,14 +242,16 @@ defmodule Realtime.Tenants.AuthorizationRemoteTest do
Authorization.get_write_authorizations(
%Policies{},
context.db_conn,
context.authorization_context
context.authorization_context,
:presence
)

assert {:error, :rls_policy_error, %Postgrex.Error{}} =
Authorization.get_write_authorizations(
%Policies{},
context.db_conn,
context.authorization_context
context.authorization_context,
:presence
)
end
end
Expand Down
Loading
Loading