From b804b0ca8d1a43da6b2c2b8becaef305a7c2244b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Filipe=20Caba=C3=A7o?= Date: Tue, 19 Sep 2023 23:48:32 +0100 Subject: [PATCH] Improve code complexity in Replication Poller --- .../postgres_cdc_rls/replication_poller.ex | 271 +++++++++--------- .../postgres_cdc_stream/cdc_stream.ex | 44 ++- lib/realtime/application.ex | 7 +- lib/realtime/helpers.ex | 25 ++ lib/realtime_web/channels/realtime_channel.ex | 79 +++-- mix.exs | 2 +- .../cluster_strategy/postgres_test.exs | 2 +- .../channels/realtime_channel_test.exs | 13 + 8 files changed, 256 insertions(+), 187 deletions(-) diff --git a/lib/extensions/postgres_cdc_rls/replication_poller.ex b/lib/extensions/postgres_cdc_rls/replication_poller.ex index b6a39f9f73..523595d845 100644 --- a/lib/extensions/postgres_cdc_rls/replication_poller.ex +++ b/lib/extensions/postgres_cdc_rls/replication_poller.ex @@ -11,15 +11,15 @@ defmodule Extensions.PostgresCdcRls.ReplicationPoller do import Realtime.Helpers, only: [cancel_timer: 1, decrypt_creds: 5, default_ssl_param: 1, maybe_enforce_ssl_config: 2] - alias Extensions.PostgresCdcRls.{Replications, MessageDispatcher} alias DBConnection.Backoff - alias Realtime.PubSub - alias Realtime.Adapters.Changes.{ - DeletedRecord, - NewRecord, - UpdatedRecord - } + alias Extensions.PostgresCdcRls.MessageDispatcher + alias Extensions.PostgresCdcRls.Replications + + alias Realtime.Adapters.Changes.DeletedRecord + alias Realtime.Adapters.Changes.NewRecord + alias Realtime.Adapters.Changes.UpdatedRecord + alias Realtime.PubSub @queue_target 5_000 @@ -45,12 +45,7 @@ defmodule Extensions.PostgresCdcRls.ReplicationPoller do tenant = args["id"] state = %{ - backoff: - Backoff.new( - backoff_min: 100, - backoff_max: 5_000, - backoff_type: :rand_exp - ), + backoff: Backoff.new(backoff_min: 100, backoff_max: 5_000, backoff_type: :rand_exp), conn: conn, db_host: args["db_host"], db_port: args["db_port"], @@ -99,76 +94,27 @@ defmodule Extensions.PostgresCdcRls.ReplicationPoller do cancel_timer(poll_ref) cancel_timer(retry_ref) - try do - {time, response} = - :timer.tc(Replications, :list_changes, [ - conn, - slot_name, - publication, - max_changes, - max_record_bytes - ]) - - Realtime.Telemetry.execute( - [:realtime, :replication, :poller, :query, :stop], - %{duration: time}, - %{tenant: tenant} + broadcast_count = + conn + |> list_changes_with_telemetry( + slot_name, + publication, + max_changes, + max_record_bytes, + tenant ) + |> handle_list_changes_result(tenant) - response - catch - {:error, reason} -> - {:error, reason} - end - |> case do - {:ok, - %Postgrex.Result{ - columns: ["wal", "is_rls_enabled", "subscription_ids", "errors"] = columns, - rows: [_ | _] = rows, - num_rows: rows_count - }} -> - Enum.reduce(rows, [], fn row, acc -> - columns - |> Enum.zip(row) - |> generate_record() - |> case do - nil -> - acc - - record_struct -> - [record_struct | acc] - end - end) - |> Enum.reverse() - |> Enum.each(fn change -> - Phoenix.PubSub.broadcast_from( - PubSub, - self(), - "realtime:postgres:" <> tenant, - change, - MessageDispatcher - ) - end) - - {:ok, rows_count} + case broadcast_count do + {:ok, 0} -> + backoff = Backoff.reset(backoff) + send(self(), :poll) - {:ok, _} -> - {:ok, 0} + {:noreply, %{state | backoff: backoff, poll_ref: nil}} - {:error, reason} -> - {:error, reason} - end - |> case do - {:ok, rows_num} -> + {:ok, _} -> backoff = Backoff.reset(backoff) - - poll_ref = - if rows_num > 0 do - send(self(), :poll) - nil - else - Process.send_after(self(), :poll, poll_interval_ms) - end + poll_ref = Process.send_after(self(), :poll, poll_interval_ms) {:noreply, %{state | backoff: backoff, poll_ref: poll_ref}} @@ -186,14 +132,9 @@ defmodule Extensions.PostgresCdcRls.ReplicationPoller do if retry_count > 3 do case Replications.terminate_backend(conn, slot_name) do - {:ok, :terminated} -> - Logger.warn("Replication slot in use - terminating") - - {:error, :slot_not_found} -> - Logger.warn("Replication slot not found") - - {:error, error} -> - Logger.warn("Error terminating backend: #{inspect(error)}") + {:ok, :terminated} -> Logger.warn("Replication slot in use - terminating") + {:error, :slot_not_found} -> Logger.warn("Replication slot not found") + {:error, error} -> Logger.warn("Error terminating backend: #{inspect(error)}") end end @@ -220,6 +161,114 @@ defmodule Extensions.PostgresCdcRls.ReplicationPoller do {:noreply, prepare_replication(state)} end + def slot_name_suffix() do + case System.get_env("SLOT_NAME_SUFFIX") do + nil -> + "" + + value -> + Logger.debug("Using slot name suffix: " <> value) + "_" <> value + end + end + + defp convert_errors([_ | _] = errors), do: errors + + defp convert_errors(_), do: nil + + defp connect_db(host, port, name, user, pass, socket_opts, ssl_enforced) do + {host, port, name, user, pass} = decrypt_creds(host, port, name, user, pass) + + [ + hostname: host, + port: port, + database: name, + password: pass, + username: user, + queue_target: @queue_target, + parameters: [application_name: "realtime_rls"], + socket_options: socket_opts + ] + |> maybe_enforce_ssl_config(ssl_enforced) + |> Postgrex.start_link() + end + + defp prepare_replication( + %{backoff: backoff, conn: conn, slot_name: slot_name, retry_count: retry_count} = state + ) do + case Replications.prepare_replication(conn, slot_name) do + {:ok, _} -> + send(self(), :poll) + state + + {:error, error} -> + Logger.error("Prepare replication error: #{inspect(error)}") + {timeout, backoff} = Backoff.backoff(backoff) + retry_ref = Process.send_after(self(), :retry, timeout) + %{state | backoff: backoff, retry_ref: retry_ref, retry_count: retry_count + 1} + end + end + + defp list_changes_with_telemetry( + conn, + slot_name, + publication, + max_changes, + max_record_bytes, + tenant + ) do + args = [ + conn, + slot_name, + publication, + max_changes, + max_record_bytes + ] + + {time, response} = :timer.tc(Replications, :list_changes, args) + + Realtime.Telemetry.execute( + [:realtime, :replication, :poller, :query, :stop], + %{duration: time}, + %{tenant: tenant} + ) + + response + catch + {:error, reason} -> {:error, reason} + end + + defp handle_list_changes_result( + {:ok, + %Postgrex.Result{ + columns: ["wal", "is_rls_enabled", "subscription_ids", "errors"] = columns, + rows: [_ | _] = rows, + num_rows: rows_count + }}, + tenant + ) do + rows + |> Enum.reduce([], fn row, acc -> + columns + |> Enum.zip(row) + |> generate_record() + |> then(fn + nil -> acc + record_struct -> [record_struct | acc] + end) + end) + |> Enum.reverse() + |> Enum.each(fn change -> + topic = "realtime:postgres:" <> tenant + Phoenix.PubSub.broadcast_from(PubSub, self(), topic, change, MessageDispatcher) + end) + + {:ok, rows_count} + end + + defp handle_list_changes_result({:ok, _}, _), do: {:ok, 0} + defp handle_list_changes_result({:error, reason}, _), do: {:error, reason} + def generate_record([ {"wal", %{ @@ -294,54 +343,4 @@ defmodule Extensions.PostgresCdcRls.ReplicationPoller do end def generate_record(_), do: nil - - def slot_name_suffix() do - case System.get_env("SLOT_NAME_SUFFIX") do - nil -> - "" - - value -> - Logger.debug("Using slot name suffix: " <> value) - "_" <> value - end - end - - defp convert_errors([_ | _] = errors), do: errors - - defp convert_errors(_), do: nil - - defp connect_db(host, port, name, user, pass, socket_opts, ssl_enforced) do - {host, port, name, user, pass} = decrypt_creds(host, port, name, user, pass) - - [ - hostname: host, - port: port, - database: name, - password: pass, - username: user, - queue_target: @queue_target, - parameters: [ - application_name: "realtime_rls" - ], - socket_options: socket_opts - ] - |> maybe_enforce_ssl_config(ssl_enforced) - |> Postgrex.start_link() - end - - defp prepare_replication( - %{backoff: backoff, conn: conn, slot_name: slot_name, retry_count: retry_count} = state - ) do - case Replications.prepare_replication(conn, slot_name) do - {:ok, _} -> - send(self(), :poll) - state - - {:error, error} -> - Logger.error("Prepare replication error: #{inspect(error)}") - {timeout, backoff} = Backoff.backoff(backoff) - retry_ref = Process.send_after(self(), :retry, timeout) - %{state | backoff: backoff, retry_ref: retry_ref, retry_count: retry_count + 1} - end - end end diff --git a/lib/extensions/postgres_cdc_stream/cdc_stream.ex b/lib/extensions/postgres_cdc_stream/cdc_stream.ex index fc0a466e7f..c63962636f 100644 --- a/lib/extensions/postgres_cdc_stream/cdc_stream.ex +++ b/lib/extensions/postgres_cdc_stream/cdc_stream.ex @@ -9,8 +9,7 @@ defmodule Extensions.PostgresCdcStream do def handle_connect(opts) do Enum.reduce_while(1..5, nil, fn retry, acc -> - get_manager_conn(opts["id"]) - |> case do + case get_manager_conn(opts["id"]) do nil -> start_distributed(opts) if retry > 1, do: Process.sleep(1_000) @@ -22,13 +21,12 @@ defmodule Extensions.PostgresCdcStream do end) end - def handle_after_connect(_, _, _) do - {:ok, nil} - end + def handle_after_connect(_, _, _), do: {:ok, nil} def handle_subscribe(pg_change_params, tenant, metadata) do Enum.each(pg_change_params, fn e -> - topic(tenant, e.params) + tenant + |> topic(e.params) |> RealtimeWeb.Endpoint.subscribe(metadata) end) end @@ -45,13 +43,9 @@ defmodule Extensions.PostgresCdcStream do @spec get_manager_conn(String.t()) :: nil | {:ok, pid(), pid()} def get_manager_conn(id) do - Phoenix.Tracker.get_by_key(Stream.Tracker, "postgres_cdc_stream", id) - |> case do - [] -> - nil - - [{_, %{manager_pid: pid, conn: conn}}] -> - {:ok, pid, conn} + case Phoenix.Tracker.get_by_key(Stream.Tracker, "postgres_cdc_stream", id) do + [] -> nil + [{_, %{manager_pid: pid, conn: conn}}] -> {:ok, pid, conn} end end @@ -81,27 +75,23 @@ defmodule Extensions.PostgresCdcStream do def start(args) do addrtype = case args["ip_version"] do - 6 -> - :inet6 - - _ -> - :inet + 6 -> :inet6 + _ -> :inet end - args = - Map.merge(args, %{ - "db_socket_opts" => [addrtype] - }) + args = Map.merge(args, %{"db_socket_opts" => [addrtype]}) Logger.debug("Starting postgres stream extension with args: #{inspect(args, pretty: true)}") + opts = %{ + id: args["id"], + start: {Stream.WorkerSupervisor, :start_link, [args]}, + restart: :transient + } + DynamicSupervisor.start_child( {:via, PartitionSupervisor, {Stream.DynamicSupervisor, self()}}, - %{ - id: args["id"], - start: {Stream.WorkerSupervisor, :start_link, [args]}, - restart: :transient - } + opts ) end diff --git a/lib/realtime/application.ex b/lib/realtime/application.ex index 8ca2d6fbf2..716ebc1724 100644 --- a/lib/realtime/application.ex +++ b/lib/realtime/application.ex @@ -62,7 +62,12 @@ defmodule Realtime.Application do RealtimeWeb.Presence, {Task.Supervisor, name: Realtime.TaskSupervisor}, Realtime.Latency, - Realtime.Telemetry.Logger + Realtime.Telemetry.Logger, + {PartitionSupervisor, + strategy: :one_for_one, + partitions: 20, + child_spec: DynamicSupervisor, + name: Realtime.DynamicSupervisors.Check.TenantDb} ] ++ extensions_supervisors() children = diff --git a/lib/realtime/helpers.ex b/lib/realtime/helpers.ex index e312dca1de..1d822968ac 100644 --- a/lib/realtime/helpers.ex +++ b/lib/realtime/helpers.ex @@ -21,6 +21,31 @@ defmodule Realtime.Helpers do |> unpad() end + @spec connect_db(%{ + host: binary(), + port: non_neg_integer(), + name: binary(), + user: binary(), + pass: binary(), + socket_opts: list(), + pool: pos_integer(), + queue_target: pos_integer(), + ssl_enforced: boolean() + }) :: {:ok, map()} | {:error, any()} + def connect_db(%{ + host: host, + port: port, + name: name, + user: user, + pass: pass, + socket_opts: socket_opts, + pool: pool, + queue_target: queue_target, + ssl_enforced: ssl_enforced + }) do + connect_db(host, port, name, user, pass, socket_opts, pool, queue_target, ssl_enforced) + end + @spec connect_db( String.t(), String.t(), diff --git a/lib/realtime_web/channels/realtime_channel.ex b/lib/realtime_web/channels/realtime_channel.ex index 87b8766538..ee8a5947a5 100644 --- a/lib/realtime_web/channels/realtime_channel.ex +++ b/lib/realtime_web/channels/realtime_channel.ex @@ -6,7 +6,7 @@ defmodule RealtimeWeb.RealtimeChannel do require Logger - import Realtime.Helpers, only: [cancel_timer: 1, decrypt!: 2] + alias Realtime.Helpers alias DBConnection.Backoff @@ -44,6 +44,8 @@ defmodule RealtimeWeb.RealtimeChannel do start_db_rate_counter(tenant) with false <- SignalHandler.shutdown_in_progress?(), + %{extensions: extensions} <- Tenants.get_tenant_by_external_id(tenant), + :ok <- check_tenant_connection(extensions), :ok <- limit_joins(socket), :ok <- limit_channels(socket), :ok <- limit_max_users(socket), @@ -70,7 +72,7 @@ defmodule RealtimeWeb.RealtimeChannel do Logger.debug("Start channel: " <> inspect(pg_change_params)) - state = %{postgres_changes: postgres_changes(pg_change_params)} + state = %{postgres_changes: add_id_to_postgres_changes(pg_change_params)} assigns = %{ ack_broadcast: !!params["config"]["broadcast"]["ack"], @@ -86,7 +88,15 @@ defmodule RealtimeWeb.RealtimeChannel do {:ok, state, assign(socket, assigns)} else - error -> handle_join_error(error) + {:error, [message: "Invalid token", claim: _claim, claim_val: _value]} = error -> + log_error_message(:warning, error) + + {:error, type} = error + when type in [:too_many_channels, :too_many_connections, :too_many_joins] -> + log_error_message(:warning, error) + + error -> + log_error_message(:error, error) end end @@ -145,7 +155,7 @@ defmodule RealtimeWeb.RealtimeChannel do } } = socket - cancel_timer(pg_sub_ref) + Helpers.cancel_timer(pg_sub_ref) args = Map.put(postgres_extension, "id", tenant) @@ -246,7 +256,7 @@ defmodule RealtimeWeb.RealtimeChannel do case confirm_token(socket) do {:ok, claims, confirm_token_ref} -> - cancel_timer(pg_sub_ref) + Helpers.cancel_timer(pg_sub_ref) pg_change_params = Enum.map(pg_change_params, &Map.put(&1, :claims, claims)) @@ -352,7 +362,7 @@ defmodule RealtimeWeb.RealtimeChannel do defp decrypt_jwt_secret(secret) do secure_key = Application.get_env(:realtime, :db_enc_key) - decrypt!(secret, secure_key) + Helpers.decrypt!(secret, secure_key) end defp postgres_subscribe(min \\ 1, max \\ 5) do @@ -488,7 +498,7 @@ defmodule RealtimeWeb.RealtimeChannel do {:ok, %{"exp" => exp} = claims} when is_integer(exp) <- ChannelsAuthorization.authorize_conn(access_token, jwt_secret_dec), exp_diff when exp_diff > 0 <- exp - Joken.current_time() do - if ref = assigns[:confirm_token_ref], do: cancel_timer(ref) + if ref = assigns[:confirm_token_ref], do: Helpers.cancel_timer(ref) interval = min(@confirm_token_ms_interval, exp_diff * 1_000) ref = Process.send_after(self(), :confirm_token, interval) @@ -580,7 +590,9 @@ defmodule RealtimeWeb.RealtimeChannel do ] end - defp postgres_cdc_subscribe(%{pg_change_params: [_ | _]} = opts) do + defp postgres_cdc_subscribe(%{pg_change_params: []}), do: [] + + defp postgres_cdc_subscribe(opts) do %{ is_new_api: is_new_api, pg_change_params: pg_change_params, @@ -608,27 +620,52 @@ defmodule RealtimeWeb.RealtimeChannel do pg_change_params end - defp postgres_cdc_subscribe(%{pg_change_params: pg_change_params}), do: pg_change_params - - defp postgres_changes(pg_change_params) do + defp add_id_to_postgres_changes(pg_change_params) do Enum.map(pg_change_params, fn %{params: params} -> id = :erlang.phash2(params) Map.put(params, :id, id) end) end - defp handle_join_error( - {:error, [message: "Invalid token", claim: _claim, claim_val: _value]} = error - ) do - log_error_message(:warning, error) - end + defp check_tenant_connection(extensions) do + extensions + |> Enum.map(fn %{settings: settings} -> + ssl_enforced = Helpers.default_ssl_param(settings) - defp handle_join_error({:error, type} = error) - when type in [:too_many_channels, :too_many_connections, :too_many_joins] do - log_error_message(:warning, error) - end + host = settings["db_host"] + port = settings["db_port"] + name = settings["db_name"] + user = settings["db_user"] + password = settings["db_password"] + socket_opts = settings["db_socket_opts"] - defp handle_join_error(error), do: log_error_message(:error, error) + opts = %{ + host: host, + port: port, + name: name, + user: user, + pass: password, + socket_opts: socket_opts, + pool: 1, + queue_target: 1000, + ssl_enforced: ssl_enforced + } + + with {:ok, conn} <- Helpers.connect_db(opts), + {:ok, _} <- Postgrex.query(conn, "SELECT 1", []) do + :ok + else + {:error, reason} -> + Logger.error("Tenant query failed with " <> inspect(reason)) + :error + end + end) + |> Enum.any?(fn res -> res == :ok end) + |> then(fn + true -> :ok + false -> {:error, :tenant_database_unavailable} + end) + end defp log_error_message(:warning, error) do error_msg = inspect(error) diff --git a/mix.exs b/mix.exs index 14a8a9cacd..261468dd1d 100644 --- a/mix.exs +++ b/mix.exs @@ -4,7 +4,7 @@ defmodule Realtime.MixProject do def project do [ app: :realtime, - version: "2.22.21", + version: "2.23.1", elixir: "~> 1.14.0", elixirc_paths: elixirc_paths(Mix.env()), start_permanent: Mix.env() == :prod, diff --git a/test/realtime/cluster_strategy/postgres_test.exs b/test/realtime/cluster_strategy/postgres_test.exs index 8891326c0c..1bc6aa681f 100644 --- a/test/realtime/cluster_strategy/postgres_test.exs +++ b/test/realtime/cluster_strategy/postgres_test.exs @@ -24,7 +24,7 @@ defmodule Realtime.Cluster.Strategy.PostgresTest do {:ok, conn_notif} = PN.start_link(state.meta.opts.()) PN.listen(conn_notif, channel_name) node = "#{node()}" - assert_receive {:notification, _, _, channel_name, ^node} + assert_receive {:notification, _, _, ^channel_name, ^node} end defp libcluster_state() do diff --git a/test/realtime_web/channels/realtime_channel_test.exs b/test/realtime_web/channels/realtime_channel_test.exs index a66b26d921..4d31e558b0 100644 --- a/test/realtime_web/channels/realtime_channel_test.exs +++ b/test/realtime_web/channels/realtime_channel_test.exs @@ -118,4 +118,17 @@ defmodule RealtimeWeb.RealtimeChannelTest do end end end + + describe "tenant not found" do + test "returns error and halts join" do + end + end + + describe "checks tenant db connectivity" do + test "successful connection proceeds with join" do + end + + test "unsuccessful connection returns error and halts join" do + end + end end