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
27 changes: 27 additions & 0 deletions bench/core_bench.exs
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,33 @@ Benchee.run(
print: [configuration: false]
)

# ---------------------------------------------------------------------------
# Streaming vs batch merge benchmark
# ---------------------------------------------------------------------------

IO.puts("\n--- Streaming vs batch merge ---\n")

# SUM pipeline → streaming merge (lattice-compatible)
# The distributed path now auto-selects streaming for SUM/COUNT/MIN/MAX
Benchee.run(
%{
"streaming merge (SUM + COUNT, 2 workers)" => fn ->
Dux.from_query(medium_sql)
|> Dux.group_by(:grp)
|> Dux.summarise_with(total: "SUM(value)", n: "COUNT(*)")
|> Dux.Remote.Coordinator.execute(workers: [w1, w2])
end,
"streaming merge (MIN + MAX, 2 workers)" => fn ->
Dux.from_query(medium_sql)
|> Dux.summarise_with(lo: "MIN(value)", hi: "MAX(value)")
|> Dux.Remote.Coordinator.execute(workers: [w1, w2])
end
},
warmup: 1,
time: 5,
print: [configuration: false]
)

# ---------------------------------------------------------------------------
# Shuffle join benchmark
# ---------------------------------------------------------------------------
Expand Down
3 changes: 2 additions & 1 deletion bench/results/history.csv
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
version,sha,date,from_query_10k_ms,from_list_100_ms,from_list_10k_ms,to_rows_1k_ms,to_columns_10k_ms,join_small_ms,full_pipeline_ms,group_summarise_ms,filter_mutate_ms,distributed_2_ms,shortest_paths_ms,connected_components_ms,pagerank_ms,triangle_count_ms,out_degree_ms,communities_ms,shuffle_join_ms,broadcast_bloom_ms
version,sha,date,from_query_10k_ms,from_list_100_ms,from_list_10k_ms,to_rows_1k_ms,to_columns_10k_ms,join_small_ms,full_pipeline_ms,group_summarise_ms,filter_mutate_ms,distributed_2_ms,shortest_paths_ms,connected_components_ms,pagerank_ms,triangle_count_ms,out_degree_ms,communities_ms,shuffle_join_ms,broadcast_bloom_ms,streaming_minmax_ms,streaming_sumcount_ms
v0.1.1,2eb9e5d,2026-03-23,0.10,3.41,3675,3.69,3150,4.24,625,655,3143,,,,,,
0534abb-adbc,0534abb,2026-03-23,1.18,6.04,5.81,6.78,8.08,6.39,4.88,5.81,7.29,9.47,,,,,,
96620af-graph,96620af,2026-03-23,,,,,,,,,,,34,19,221,20,3,45
b7d9da7-shuffle-bloom,b7d9da7,2026-03-23,16.02,8.59,,,,,,2.85,,3.84,,,,,,,1500,4.68
45c613d-streaming,45c613d,2026-03-23,16.04,7.65,,,,,,2.95,,4.09,,,,,,,1490,4.80,3.29,9.41
101 changes: 95 additions & 6 deletions lib/dux/remote/coordinator.ex
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ defmodule Dux.Remote.Coordinator do
The result is a `%Dux{}` struct with the merged data.
"""

alias Dux.Remote.{Merger, Partitioner, PipelineSplitter, Shuffle, Worker}
alias Dux.Remote.{Merger, Partitioner, PipelineSplitter, Shuffle, StreamingMerger, Worker}
import Dux.SQL.Helpers, only: [qi: 1]

# Broadcast threshold: 256MB serialized Arrow IPC
Expand Down Expand Up @@ -50,8 +50,12 @@ defmodule Dux.Remote.Coordinator do
end

# Split pipeline: worker ops push down, coordinator ops apply post-merge
%{worker_ops: worker_ops, coordinator_ops: coord_ops, agg_rewrites: rewrites} =
PipelineSplitter.split(pipeline.ops)
%{
worker_ops: worker_ops,
coordinator_ops: coord_ops,
agg_rewrites: rewrites,
streaming_compatible?: streaming?
} = split = PipelineSplitter.split(pipeline.ops)

# Preprocess joins: broadcast/shuffle right sides that aren't worker-safe
case preprocess_joins(worker_ops, workers, timeout, bcast_threshold) do
Expand All @@ -60,7 +64,9 @@ defmodule Dux.Remote.Coordinator do
worker_pipeline = %{pipeline | ops: processed_ops}

try do
result = execute_fan_out(worker_pipeline, workers, strategy, timeout)
result =
execute_fan_out(worker_pipeline, workers, strategy, timeout, streaming?, split)

result = apply_avg_rewrites(result, rewrites)
finalize(result, coord_ops)
after
Expand Down Expand Up @@ -108,10 +114,26 @@ defmodule Dux.Remote.Coordinator do
# ---------------------------------------------------------------------------

# Standard distributed execution: partition → fan out → merge
defp execute_fan_out(worker_pipeline, workers, strategy, timeout) do
# When streaming_compatible? is true and a StreamingMerger can be created,
# folds results incrementally. Otherwise falls back to batch merge.
defp execute_fan_out(worker_pipeline, workers, strategy, timeout, streaming?, split) do
assignments = Partitioner.assign(worker_pipeline, workers, strategy: strategy)
n_workers = length(assignments)

# Try streaming merge for lattice-compatible pipelines
merger =
if streaming? do
StreamingMerger.new(split.worker_ops, n_workers)
end

if merger do
execute_streaming(assignments, worker_pipeline, merger, timeout)
else
execute_batch(assignments, worker_pipeline, n_workers, timeout)
end
end

defp execute_batch(assignments, worker_pipeline, n_workers, timeout) do
results =
:telemetry.span([:dux, :distributed, :fan_out], %{n_workers: n_workers}, fn ->
r = fan_out(assignments, timeout)
Expand All @@ -131,6 +153,73 @@ defmodule Dux.Remote.Coordinator do
end)
end

defp execute_streaming(assignments, _worker_pipeline, merger, timeout) do
n_workers = length(assignments)

:telemetry.span(
[:dux, :distributed, :fan_out],
%{n_workers: n_workers, streaming: true},
fn ->
final_merger = streaming_fan_out(assignments, merger, n_workers, timeout)
result = StreamingMerger.to_dux(final_merger)
{result, %{n_workers: n_workers, streaming: true}}
end
)
end

defp streaming_fan_out(assignments, merger, n_workers, timeout) do
assignments
|> Enum.with_index()
|> Task.async_stream(
fn {{worker, partition_pipeline}, idx} ->
execute_on_worker(worker, partition_pipeline, idx, n_workers, timeout)
end,
timeout: timeout,
max_concurrency: n_workers,
ordered: false
)
|> Enum.reduce(merger, &fold_streaming_result/2)
end

defp execute_on_worker(worker, pipeline, idx, n_workers, timeout) do
start_time = System.monotonic_time()
result = Worker.execute(worker, pipeline, timeout)

case result do
{:ok, ipc} ->
:telemetry.execute(
[:dux, :distributed, :worker, :stop],
%{duration: System.monotonic_time() - start_time, ipc_bytes: byte_size(ipc)},
%{worker: worker, worker_index: idx, n_workers: n_workers}
)

_ ->
:ok
end

result
end

defp fold_streaming_result({:ok, {:ok, ipc}}, merger) do
merger = StreamingMerger.fold(merger, ipc)

:telemetry.execute(
[:dux, :distributed, :streaming_merge],
%{workers_complete: merger.workers_complete, workers_total: merger.workers_total},
%{progress: StreamingMerger.progress(merger)}
)

merger
end

defp fold_streaming_result({:ok, {:error, _reason}}, merger) do
StreamingMerger.record_failure(merger)
end

defp fold_streaming_result({:exit, _reason}, merger) do
StreamingMerger.record_failure(merger)
end

# Multi-stage execution for shuffle joins:
# 1. Execute ops before the join as a distributed query
# 2. Shuffle join the result with the right side
Expand All @@ -153,7 +242,7 @@ defmodule Dux.Remote.Coordinator do
if ops_before == [] do
Dux.compute(%{pipeline | ops: []})
else
execute_fan_out(%{pipeline | ops: ops_before}, workers, strategy, timeout)
execute_fan_out(%{pipeline | ops: ops_before}, workers, strategy, timeout, false, %{})
end

# Stage 2: shuffle join
Expand Down
185 changes: 185 additions & 0 deletions lib/dux/remote/streaming_merger.ex
Original file line number Diff line number Diff line change
@@ -0,0 +1,185 @@
defmodule Dux.Remote.StreamingMerger do
@moduledoc false

# Folds worker IPC results incrementally using lattice merge operations.
# Instead of loading all IPC results as temp tables and running one big
# SQL query (batch Merger), the StreamingMerger decodes each result and
# merges aggregate values in Elixir as workers report in.
#
# Benefits:
# - Lower memory: don't hold all IPC binaries simultaneously
# - Lower latency: merge starts when first worker finishes
# - Enables progressive results (Phase C)
#
# The StreamingMerger operates on the *rewritten* worker output — after
# PipelineSplitter has transformed AVG → SUM+COUNT etc. It uses the
# re-aggregation functions (SUM→SUM, COUNT→SUM, MIN→MIN, MAX→MAX)
# matching what the batch Merger does in SQL.

alias Dux.Lattice

defstruct [
:groups,
:agg_columns,
:accumulator,
:workers_total,
:workers_complete,
:workers_failed
]

@doc """
Create a streaming merger for a pipeline with lattice-mergeable aggregates.

`worker_ops` are the ops pushed to workers (after PipelineSplitter rewrite).
`n_workers` is the total number of workers.

Returns a `%StreamingMerger{}` or `nil` if the pipeline can't be streamed
(non-lattice aggregates, no summarise, etc.).
"""
def new(worker_ops, n_workers) do
with {:ok, groups} <- find_groups(worker_ops),
{:ok, aggs} <- find_summarise(worker_ops),
{:ok, agg_columns} <- classify_rewritten_aggs(aggs) do
%__MODULE__{
groups: groups,
agg_columns: agg_columns,
accumulator: %{},
workers_total: n_workers,
workers_complete: 0,
workers_failed: 0
}
else
:not_streamable -> nil
end
end

@doc """
Fold one worker's IPC result into the accumulator.
"""
def fold(%__MODULE__{} = merger, ipc_binary) do
conn = Dux.Connection.get_conn()
ref = Dux.Backend.table_from_ipc(conn, ipc_binary)
rows = Dux.Backend.table_to_rows(conn, ref)

accumulator =
Enum.reduce(rows, merger.accumulator, fn row, acc ->
group_key = extract_group_key(row, merger.groups)
group_state = Map.get(acc, group_key, init_group(merger.agg_columns))
updated = merge_row(group_state, row, merger.agg_columns)
Map.put(acc, group_key, updated)
end)

%{merger | accumulator: accumulator, workers_complete: merger.workers_complete + 1}
end

@doc """
Record a worker failure without crashing the merge.
"""
def record_failure(%__MODULE__{} = merger) do
%{merger | workers_failed: merger.workers_failed + 1}
end

@doc """
Finalize the accumulator into a list of row maps.
"""
def finalize(%__MODULE__{} = merger) do
Enum.map(merger.accumulator, fn {group_key, agg_states} ->
finalized =
Map.new(agg_states, fn {col_name, {lattice, state}} ->
{col_name, lattice.finalize(state)}
end)

Map.merge(group_key, finalized)
end)
end

@doc """
Convert finalized rows to a `%Dux{}` struct.
"""
def to_dux(%__MODULE__{} = merger) do
rows = finalize(merger)

if rows == [] do
Dux.from_query("SELECT 1 WHERE false") |> Dux.compute()
else
Dux.from_list(rows) |> Dux.compute()
end
end

@doc """
Get progress metadata.
"""
def progress(%__MODULE__{} = merger) do
%{
workers_complete: merger.workers_complete,
workers_total: merger.workers_total,
workers_failed: merger.workers_failed,
complete?: merger.workers_complete + merger.workers_failed >= merger.workers_total
}
end

# ---------------------------------------------------------------------------
# Internals
# ---------------------------------------------------------------------------

defp find_groups(ops) do
case Enum.find(ops, &match?({:group_by, _}, &1)) do
{:group_by, cols} -> {:ok, cols}
nil -> {:ok, []}
end
end

defp find_summarise(ops) do
case Enum.find(ops, &match?({:summarise, _}, &1)) do
{:summarise, aggs} -> {:ok, aggs}
nil -> :not_streamable
end
end

# Classify the rewritten aggregate columns into lattice types.
# These are the worker output columns (after PipelineSplitter rewrite).
# SUM(x) → Sum, COUNT(x) → Count (re-aggregated as SUM), MIN → Min, MAX → Max
defp classify_rewritten_aggs(aggs) do
result =
Enum.reduce_while(aggs, [], fn {name, expr}, acc ->
case classify_rewritten(expr) do
nil -> {:halt, :not_streamable}
lattice -> {:cont, [{name, lattice} | acc]}
end
end)

case result do
:not_streamable -> :not_streamable
classified -> {:ok, Enum.reverse(classified)}
end
end

defp classify_rewritten(expr) when is_binary(expr) do
upper = String.upcase(expr)
Lattice.classify(upper)
end

defp classify_rewritten(_), do: nil

defp extract_group_key(row, groups) do
Map.take(row, groups)
end

defp init_group(agg_columns) do
Map.new(agg_columns, fn {name, lattice} ->
{name, {lattice, lattice.bottom()}}
end)
end

defp merge_row(group_state, row, agg_columns) do
Enum.reduce(agg_columns, group_state, fn {name, _lattice}, state ->
{lattice, current} = Map.fetch!(state, name)
value = Map.get(row, name, lattice.bottom())

# Coerce nil to bottom
value = if is_nil(value), do: lattice.bottom(), else: value

Map.put(state, name, {lattice, lattice.merge(current, value)})
end)
end
end
Loading
Loading