defmodule SuperWorker.Supervisor.Chain do @moduledoc """ Documentation for `SuperWorker.Supervisor.Chain`. """ @chain_params [:id, :restart_strategy, :finished_callback, :queue_length, :send_type] @chain_restart_strategies [:one_for_one, :one_for_all, :rest_for_one, :before_for_one] @send_types [:broadcast, :random, :partition, :round_robin] alias :ets, as: Ets alias __MODULE__ defstruct [ :id, # chain id, unique in supervior. :first_worker_id, # first worker id in the chain. where the data is sent. restart_strategy: :one_for_one, supervisor: nil, partition: nil, finished_callback: nil, queue_length: 50, send_type: :random, # :broadcast, :random, :partition, :round_robin data_table: nil, # ets table of supervisor. ] @type t :: %__MODULE__{ id: any, first_worker_id: any, restart_strategy: atom, supervisor: atom, partition: atom, finished_callback: nil | {:fun, fun} | {module, atom, [any]}, queue_length: non_neg_integer, send_type: :broadcast | :random | :partition | :round_robin, data_table: atom } import SuperWorker.Supervisor.Utils alias SuperWorker.Supervisor, as: Sup alias SuperWorker.Supervisor.{Worker, Message, MapQueue} require Logger ## Public functions @spec check_options([atom() | keyword()]) :: {:error, atom | {atom, any}} | {:ok, Chain.t} def check_options(opts) do with {:ok, opts} <- normalize_opts(opts, @chain_params), {:ok, chain} <- map_to_struct(opts), {:ok, chain} <- validate_opts(chain) do {:ok, chain} end end @spec get_worker(Chain.t, any()) :: {:error, :not_found} | {:ok, Worker.t} def get_worker(chain, worker_id) do case Ets.lookup(chain.data_table, {:worker, {:chain, chain.id}, worker_id}) do [{_, worker}] -> {:ok, worker} [] -> {:error, :not_found} end end @spec worker_exists?(Chain.t, any()) :: boolean() def worker_exists?(chain, worker_id) do case get_worker(chain, worker_id) do {:ok, _} -> true {:error, _} -> false end end @spec get_all_workers(Chain.t) :: {:ok, list(Worker.t)} def get_all_workers(chain) do result = Ets.match(chain.data_table, {{:worker, {:chain, chain.id}, :_}, :"$1"}) |> List.flatten() {:ok, result} end @spec add_worker(Chain.t, Worker.t) :: {:error, :already_exists} | {:ok, Chain.t} def add_worker(chain, %Worker{} = worker) do if worker_exists?(chain, worker.id) do {:error, :already_exists} else worker = worker |> Map.put(:order, get_and_update_chain_order(chain)) |> Map.put(:parent, chain.id) if worker.num_workers == 1 do # has 1 worker per chain node. do_add_worker(chain, worker) else chain = Enum.reduce(1..worker.num_workers, chain, fn index, acc -> worker = Map.put(worker, :id, {:multi_workers, worker.id, index}) {:ok, chain} = do_add_worker(acc, worker) chain end) Logger.debug("SuperWorker, added multi workers (#{inspect worker.id}) to the chain #{inspect chain.id}") {:ok, chain} end end end @spec do_add_worker(Chain.t, Worker.t) :: {:error, :already_exists} | {:ok, Chain.t} defp do_add_worker(chain, %Worker{} = worker) do Logger.debug("SuperWorker, adding worker #{inspect worker.id} to the chain #{inspect chain.id}") if worker_exists?(chain, worker.id) do {:error, :already_exists} else Ets.insert(chain.data_table, {{:worker, {:chain, chain.id}, worker.id}, worker}) chain |> update_chain_first(worker) |> spawn_worker(worker.id) end end @spec restart_worker(Chain.t, any()) :: {:error, any} | {:ok, Chain.t} def restart_worker(chain, worker_id) do if worker_exists?(chain, worker_id) do kill_worker(chain, worker_id) spawn_worker(chain, worker_id) else {:error, "Worker not found"} end end @spec restart_all_workers(Chain.t) :: {:ok, Chain.t} # TO-DO: support restart workers depend on host partition. def restart_all_workers(chain) do workers = Ets.match_object(chain.data_table, {:worker, {:chain, chain.id}, :_}) |> Enum.map(fn {_, worker} -> worker end) Enum.map(workers, fn worker -> Logger.info("SuperWorker, restarting worker #{worker.id}, pid: #{worker.pid}") Process.exit(worker.pid, :kill) worker = do_spawn_worker(worker) worker.id end) {:ok, chain} end @spec remove_worker(Chain.t, any()) :: {:error, any} | {:ok, Chain.t} def remove_worker(chain, worker_id) do case get_worker(chain, worker_id) do {:ok, _} -> Ets.delete(chain.data_table, {:worker, {:chain, chain.id}, worker_id}) # TO-DO: remove other info of worker. {:ok, chain} {:error, reason} = error -> Logger.error("SuperWorker, failed to remove worker #{inspect(worker_id)} in chain #{inspect chain.id}, error: #{inspect reason}") error end end @spec kill_worker(Chain.t, any()) :: {:error, any} | {:ok, Chain.t} def kill_worker(chain, worker_id) do case get_worker(chain, worker_id) do {:ok, worker} -> Process.exit(worker.pid, :kill) {:ok, chain} {:error, reason} = error -> Logger.error("SuperWorker, failed to kill worker #{inspect(worker_id)} in chain #{inspect chain.id}, error: #{inspect reason}") error end end @spec kill_all_workers(Chain.t) :: {:ok, Chain.t} # TO-DO: refactor this function, remove ref & pid from worker def kill_all_workers(chain) do workers = chain.workers Enum.each(workers, fn worker_id -> {:ok, worker} = get_worker(chain, worker_id) Process.exit(worker.pid, :kill) end) {:ok, chain} end @spec new_data(Chain.t, Message.t) :: any def new_data(chain = %Chain{}, msg = %Message{}) do send_next(chain, 1, msg) end ## Private functions @spec send_next(Chain.t, non_neg_integer, Message.t) :: any defp send_next(chain = %Chain{}, order, msg = %Message{}) do case Registry.lookup(chain.supervisor, {:chain_order, chain.id, order}) do [] -> Logger.debug("SuperWorker, not found next worker for order #{order}, chain: #{chain.id}, go to finished callback.") # TO-DO: catch throw, error from outside. case chain.finished_callback do nil -> Logger.debug("SuperWorker, not found callback for chain #{chain.id}") {:error, :no_worker_or_callback} {:fun, fun} -> fun.(msg.data) {:ok, :call_back} {m, f, a} -> apply(m, f, [msg.data|a]) {:ok, :call_back} end [{pid, worker_id}] -> # just one worker doesn't check type. Logger.debug("SuperWorker, chain #{inspect chain.id}, order: #{order}, found a next worker: #{inspect worker_id}, send msg #{inspect msg.id}") send(pid, {:new_data, msg}) {:ok, :send_one} [_|_] = entries -> Logger.debug("SuperWorker, chain #{inspect chain.id}, order: #{order}, found next workers: #{inspect entries}") case chain.send_type do :broadcast -> Enum.each(entries, fn {pid, worker_id} -> Logger.debug("SuperWorker, sending data to the next worker #{inspect worker_id} (broadcast)") send(pid, {:new_data, msg}) end) {:ok, :send_all} :random -> {pid, _} = Enum.random(entries) Logger.debug("SuperWorker, sending data to the next worker #{inspect pid} (random)") send(pid, {:new_data, msg}) {:ok, :send_random} :partition -> order = get_hash_order(msg.data, length(entries)) {pid, {:multi_workers, worker_id, index}} = Enum.at(entries, order) Logger.debug("SuperWorker, sending data to the next worker #{inspect worker_id} (partition, #{index})") send(pid, {:new_data, msg}) {:ok, :send_partition} :round_robin -> [{_, {:multi_workers, worker_id, _}} | _] = entries order = get_next_round_robin_order(chain, worker_id, length(entries) - 1) {pid, {_, _, index}} = Enum.at(entries, order) Logger.debug("SuperWorker, sending data to the next worker #{inspect worker_id} (round_robin, #{index})") send(pid, {:new_data, msg}) {:ok, :send_round_robin} end end end defp get_next_round_robin_order(chain, worker_id, max_order) do Logger.debug("SuperWorker, getting next round robin order for worker #{inspect worker_id}, max_order: #{max_order}") Ets.update_counter(chain.data_table, {:round_robin, {:chain, chain.id}, worker_id}, {2, 1, max_order, 0}, {{:round_robin, {:chain, chain.id}, worker_id}, 0}) end defp update_chain_first(chain, worker) do if chain.first_worker_id do chain else Map.put(chain, :first_worker_id, worker.id) end end defp spawn_worker(chain, worker_id) do {:ok, worker} = get_worker(chain, worker_id) worker = worker |> Map.put(:supervisor, chain.supervisor) |> Map.put(:first_worker_id, chain.first_worker_id) |> do_spawn_worker() Ets.insert(chain.data_table, {{:worker, {:chain, chain.id}, worker.id}, worker}) Ets.insert(chain.data_table, {{:worker, :ref, worker.ref}, worker.id, worker.pid, {:chain, chain.id}}) {:ok, chain} end defp do_spawn_worker(%Worker{} = worker) do {pid, ref} = spawn_monitor(fn -> # Store for user can directly access to the worker. Process.put({:supervisor, :sup_id}, worker.supervisor) Process.put({:supervisor,:chain}, worker.parent) Process.put({:supervisor, :worker_id}, worker.id) Registry.register(worker.supervisor, {:chain, worker.parent}, :worker) Registry.register(worker.supervisor, {:chain_order, worker.parent, worker.order}, worker.id) case worker.id do {:multi_workers, root_id, index} -> # subsribe to the root worker id for get data. Registry.register(worker.supervisor, {:worker, {:chain, worker.parent}, root_id}, index) # Registry.register(worker.supervisor, {:worker, worker.id}, 1) _ -> Registry.register(worker.supervisor, {:worker, {:chain, worker.parent}, worker.id}, 0) end loop_chain(%MapQueue{}, worker) end) # Link to child for case supervisor is down. # TO-DO: Improve case worker crash immediately. Process.link(pid) worker |> Map.put(:pid, pid) |> Map.put(:ref, ref) end # Support receive data from the previous process in the chain and pass it to the next process. defp loop_chain(queue, %Worker{id: id, parent: chain_id} = worker) do receive do {:processed, msg_id, worker_id} -> Logger.debug("SuperWorker, worker #{inspect worker_id} processed the data, msg_id: #{msg_id}") {:ok, queue} = MapQueue.remove(queue, msg_id) loop_chain(queue, worker) {:new_data, msg = %Message{}} -> result = # TO-DO: catch throw, error from outside. case worker.fun do {:fun, f} -> f.(msg.data) {m, f, a} -> apply(m, f, [msg.data | a]) end if worker.first_worker_id != id do send(msg.from, {:processed, msg.id, id}) end case result do {:next, new_data} -> if MapQueue.is_full?(queue) do Logger.debug("SuperWorker, worker #{inspect(id)}, queue is full, go to loop waiting for consume last data.") loop_send(queue, worker) end Logger.debug("SuperWorker, worker #{inspect(id)}, passing data to the next process, chain: #{inspect(chain_id)}") {:ok, queue, msg_id} = MapQueue.add(queue, new_data) {:ok,chain} = Sup.get_chain(get_my_supervisor(), chain_id) msg = Message.new(self(), nil, new_data, msg_id) send_next(chain, worker.order + 1, msg) loop_chain(queue, worker) {:error, reason} -> Logger.error("SuperWorker, worker #{inspect(id)}, error in chain process, chain: #{inspect(chain_id)}: #{inspect(reason)}") # TO-DO: decide to ignore or stop the chain. {:drop, reason} -> Logger.info("SuperWorker, worker #{inspect(id)}, dropping chain process, chain: #{inspect(chain_id)}: #{inspect(reason)}") loop_chain(queue, worker) {:stop, reason} -> Logger.info("SuperWorker, worker #{inspect(id)}, stopping chain process, chain: #{inspect(chain_id)}") exit(reason) data -> Logger.debug("SuperWorker, worker #{inspect(id)}, passing data (default) to the next process, chain: #{inspect(chain_id)}") if MapQueue.is_full?(queue) do Logger.debug("SuperWorker, worker #{inspect(id)}, queue is full, go to loop waiting for consume last data.") loop_send(queue, worker) end {:ok, queue, msg_id} = MapQueue.add(queue, data) chain = Sup.get_chain(get_my_supervisor(), chain_id) msg = Message.new(self(), nil, data, msg_id) send_next(chain, worker.order + 1, msg) loop_chain(queue, worker) end {:kill, reason} -> Logger.debug("SuperWorker, worker #{inspect(id)}, killing chain, chain: #{inspect(chain_id)}") exit(reason) {:stop, ^chain_id} -> Logger.debug("SuperWorker, worker #{inspect(id)}, stopping chain, chain: #{inspect(chain_id)}") end end defp loop_send(queue, %Worker{id: id, parent: chain_id} = _worker) do receive do {:processed, msg_id, worker_id} -> Logger.debug("SuperWorker, worker #{worker_id} processed the data, msg_id: #{msg_id}") {:ok, MapQueue.remove(queue, msg_id)} {:kill, reason} -> Logger.debug("SuperWorker, worker #{id}, killing chain, chain: #{inspect(chain_id)}") exit(reason) {:stop, ^chain_id} -> Logger.debug("SuperWorker, worker #{id}, stopping chain process, chain: #{inspect(chain_id)}") :stop end end defp validate_restart_strategy(opts) do if opts.restart_strategy in @chain_restart_strategies do {:ok, opts} else {:error, "Invalid group restart strategy, #{inspect opts.restart_strategy}"} end end defp validate_send_type(opts) do if opts.send_type in @send_types do {:ok, opts} else {:error, "Invalid send type, #{inspect opts.send_type}"} end end defp validate_callback(opts) do case opts.finished_callback do nil -> {:ok, opts} {:fun, fun} when is_function(fun) -> {:ok, opts} {m, f, a} when is_atom(m) and is_atom(f) and is_list(a) -> {:ok, opts} _ -> {:error, "Invalid callback"} end end defp validate_queue_length(opts) do case opts.queue_length do n when is_integer(n) and n > 0 -> {:ok, opts} _ -> {:error, "Invalid queue length"} end end defp validate_opts(chain) do with {:ok, chain} <- validate_restart_strategy(chain), {:ok, chain} <- validate_send_type(chain), {:ok, chain} <- validate_callback(chain), {:ok, chain} <- validate_queue_length(chain) do {:ok, chain} end end defp map_to_struct(opts) when is_map(opts) do {:ok, struct(__MODULE__, opts)} end defp get_my_chain_id() do Process.get({:supervisor, :chain_id}) end defp get_my_supervisor() do Process.get({:supervisor, :sup_id}) end defp get_and_update_chain_order(chain) do Ets.update_counter(chain.data_table, {:last_chain_order, chain.id}, {2, 1}, {{:last_chain_order, chain.id}, 0}) end end