Skip to content
Open
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
146 changes: 136 additions & 10 deletions src/codec/distributed_codec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -14,8 +14,8 @@ use datafusion::arrow::datatypes::SchemaRef;
use datafusion::common::Result;
use datafusion::error::DataFusionError;
use datafusion::execution::TaskContext;
use datafusion::physical_expr::EquivalenceProperties;
use datafusion::physical_expr::equivalence::{EquivalenceClass, EquivalenceGroup};
use datafusion::physical_expr::{EquivalenceProperties, PhysicalExpr};
use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType};
use datafusion::physical_plan::union::UnionExec;
use datafusion::physical_plan::{ExecutionPlan, Partitioning, PlanProperties};
Expand Down Expand Up @@ -66,11 +66,17 @@ impl PhysicalExtensionCodec for DistributedCodec {
fn parse_stage_proto(
proto: Option<StageProto>,
inputs: &[Arc<dyn ExecutionPlan>],
dynamic_filter_anchors: Vec<Arc<dyn PhysicalExpr>>,
) -> Result<Stage, DataFusionError> {
let Some(proto) = proto else {
return Err(proto_error("Empty StageProto"));
};
if let Some(input) = inputs.first().cloned() {
if !dynamic_filter_anchors.is_empty() {
return Err(proto_error(
"Dynamic filter anchors require a remote input stage",
));
}
Ok(Stage::Local(LocalStage {
query_id: deserialize_uuid(proto.query_id.as_ref())?,
num: proto.num as usize,
Expand All @@ -94,6 +100,7 @@ impl PhysicalExtensionCodec for DistributedCodec {
num: proto.num as usize,
workers: worker_urls,
runtime_stats: None,
dynamic_filter_anchors,
}))
}
}
Expand All @@ -118,6 +125,14 @@ impl PhysicalExtensionCodec for DistributedCodec {
proto_converter,
)?
.ok_or(proto_error("NetworkShuffleExec is missing partitioning"))?;
let dynamic_filter_anchors = input_stage
.as_ref()
.into_iter()
.flat_map(|stage| stage.dynamic_filter_anchors.iter())
.map(|expression| {
proto_converter.proto_to_physical_expr(expression, &schema, &decode_ctx)
})
.collect::<Result<Vec<_>>>()?;
let schema = Arc::new(schema);
let equivalence_properties = parse_equivalence_properties(
equivalence_classes,
Expand All @@ -129,7 +144,7 @@ impl PhysicalExtensionCodec for DistributedCodec {
Ok(Arc::new(new_network_hash_shuffle_exec(
partitioning,
equivalence_properties,
parse_stage_proto(input_stage, inputs)?,
parse_stage_proto(input_stage, inputs, dynamic_filter_anchors)?,
)))
}
DistributedExecNode::NetworkCoalesceTasks(NetworkCoalesceExecProto {
Expand All @@ -151,6 +166,14 @@ impl PhysicalExtensionCodec for DistributedCodec {
proto_converter,
)?
.ok_or(proto_error("NetworkCoalesceExec is missing partitioning"))?;
let dynamic_filter_anchors = input_stage
.as_ref()
.into_iter()
.flat_map(|stage| stage.dynamic_filter_anchors.iter())
.map(|expression| {
proto_converter.proto_to_physical_expr(expression, &schema, &decode_ctx)
})
.collect::<Result<Vec<_>>>()?;
let schema = Arc::new(schema);
let equivalence_properties = parse_equivalence_properties(
equivalence_classes,
Expand All @@ -162,7 +185,7 @@ impl PhysicalExtensionCodec for DistributedCodec {
Ok(Arc::new(new_network_coalesce_tasks_exec(
partitioning,
equivalence_properties,
parse_stage_proto(input_stage, inputs)?,
parse_stage_proto(input_stage, inputs, dynamic_filter_anchors)?,
)))
}
DistributedExecNode::NetworkBroadcast(NetworkBroadcastExecProto {
Expand All @@ -184,6 +207,14 @@ impl PhysicalExtensionCodec for DistributedCodec {
proto_converter,
)?
.ok_or(proto_error("NetworkBroadcastExec is missing partitioning"))?;
let dynamic_filter_anchors = input_stage
.as_ref()
.into_iter()
.flat_map(|stage| stage.dynamic_filter_anchors.iter())
.map(|expression| {
proto_converter.proto_to_physical_expr(expression, &schema, &decode_ctx)
})
.collect::<Result<Vec<_>>>()?;
let schema = Arc::new(schema);
let equivalence_properties = parse_equivalence_properties(
equivalence_classes,
Expand All @@ -195,7 +226,7 @@ impl PhysicalExtensionCodec for DistributedCodec {
Ok(Arc::new(new_network_broadcast_exec(
partitioning,
equivalence_properties,
parse_stage_proto(input_stage, inputs)?,
parse_stage_proto(input_stage, inputs, dynamic_filter_anchors)?,
)))
}
DistributedExecNode::Broadcast(BroadcastExecProto {
Expand Down Expand Up @@ -273,12 +304,22 @@ impl PhysicalExtensionCodec for DistributedCodec {
buf: &mut Vec<u8>,
proto_converter: &dyn PhysicalProtoConverterExtension,
) -> Result<()> {
fn encode_stage_proto(stage: &Stage) -> Result<StageProto, DataFusionError> {
fn encode_stage_proto(
stage: &Stage,
codec: &DistributedCodec,
proto_converter: &dyn PhysicalProtoConverterExtension,
) -> Result<StageProto, DataFusionError> {
let dynamic_filter_anchors = stage
.dynamic_filter_anchors()
.iter()
.map(|expression| proto_converter.physical_expr_to_proto(expression, codec))
.collect::<Result<Vec<_>>>()?;
Ok(match stage {
Stage::Local(local) => StageProto {
query_id: serialize_uuid(&local.query_id).into(),
num: local.num as u64,
tasks: vec![ExecutionTaskProto::default(); local.tasks],
dynamic_filter_anchors,
},
Stage::Remote(remote) => {
let mut tasks = Vec::with_capacity(remote.workers.len());
Expand All @@ -291,6 +332,7 @@ impl PhysicalExtensionCodec for DistributedCodec {
query_id: serialize_uuid(&remote.query_id).into(),
num: remote.num as u64,
tasks,
dynamic_filter_anchors,
}
}
})
Expand All @@ -304,7 +346,11 @@ impl PhysicalExtensionCodec for DistributedCodec {
self,
proto_converter,
)?),
input_stage: Some(encode_stage_proto(node.input_stage())?),
input_stage: Some(encode_stage_proto(
node.input_stage(),
self,
proto_converter,
)?),
equivalence_classes: serialize_equivalence_group(
node.properties().equivalence_properties(),
self,
Expand All @@ -325,7 +371,11 @@ impl PhysicalExtensionCodec for DistributedCodec {
self,
proto_converter,
)?),
input_stage: Some(encode_stage_proto(node.input_stage())?),
input_stage: Some(encode_stage_proto(
node.input_stage(),
self,
proto_converter,
)?),
equivalence_classes: serialize_equivalence_group(
node.properties().equivalence_properties(),
self,
Expand All @@ -346,7 +396,11 @@ impl PhysicalExtensionCodec for DistributedCodec {
self,
proto_converter,
)?),
input_stage: Some(encode_stage_proto(node.input_stage())?),
input_stage: Some(encode_stage_proto(
node.input_stage(),
self,
proto_converter,
)?),
equivalence_classes: serialize_equivalence_group(
node.properties().equivalence_properties(),
self,
Expand Down Expand Up @@ -470,6 +524,9 @@ pub struct StageProto {
/// the plan
#[prost(message, repeated, tag = "3")]
pub tasks: Vec<ExecutionTaskProto>,
/// Dynamic-filter consumers retained after a remote stage's plan has moved to its workers.
#[prost(message, repeated, tag = "4")]
pub dynamic_filter_anchors: Vec<protobuf::PhysicalExprNode>,
}

#[derive(Clone, PartialEq, ::prost::Message)]
Expand Down Expand Up @@ -649,14 +706,20 @@ fn new_network_broadcast_exec(

#[cfg(test)]
mod tests {
use super::super::physical_plan::new_proto_converter as default_proto_converter;
use super::super::physical_plan::{
new_proto_converter as default_proto_converter, roundtrip_pb,
};
use super::*;
use datafusion::arrow::datatypes::{DataType, Field};
use datafusion::physical_expr::{LexOrdering, PhysicalExpr};
use datafusion::physical_plan::empty::EmptyExec;
use datafusion::physical_plan::filter::FilterExec;
use datafusion::prelude::SessionContext;
use datafusion::{
physical_expr::{Partitioning, PhysicalSortExpr, expressions::Column, expressions::col},
physical_expr::{
Partitioning, PhysicalSortExpr,
expressions::{Column, DynamicFilterPhysicalExpr, col, lit},
},
physical_plan::{ExecutionPlan, displayable, sorts::sort::SortExec, union::UnionExec},
};

Expand All @@ -670,6 +733,7 @@ mod tests {
num: 0,
workers: vec![],
runtime_stats: None,
dynamic_filter_anchors: vec![],
})
}

Expand Down Expand Up @@ -717,6 +781,68 @@ mod tests {
Ok(())
}

#[test]
fn test_roundtrip_network_dynamic_filter_anchor() -> datafusion::common::Result<()> {
let ctx = create_context();
let schema = schema_i32("a");
let dynamic_filter = Arc::new(DynamicFilterPhysicalExpr::new(
vec![Arc::new(Column::new("a", 0))],
lit(true),
)) as Arc<dyn datafusion::physical_expr::PhysicalExpr>;
let expected_id = dynamic_filter.expression_id();
let stage = Stage::Remote(RemoteStage {
query_id: Default::default(),
num: 0,
workers: vec![],
runtime_stats: None,
dynamic_filter_anchors: vec![Arc::clone(&dynamic_filter)],
});
let network: Arc<dyn ExecutionPlan> = Arc::new(new_network_hash_shuffle_exec(
Partitioning::Hash(vec![Arc::new(Column::new("a", 0))], 4),
EquivalenceProperties::new(schema),
stage,
));

let mut buf = vec![];
DistributedCodec.try_encode(Arc::clone(&network), &mut buf, &default_proto_converter())?;
let encoded = DistributedExecProto::decode(buf.as_slice())
.map_err(|error| proto_error(format!("{error}")))?;
let Some(DistributedExecNode::NetworkHashShuffle(encoded)) = encoded.node else {
panic!("expected a network shuffle")
};
assert_eq!(
encoded
.input_stage
.expect("network shuffle should contain its input stage")
.dynamic_filter_anchors
.len(),
1,
);

let plan: Arc<dyn ExecutionPlan> = Arc::new(FilterExec::try_new(dynamic_filter, network)?);

let decoded = roundtrip_pb(plan, &ctx)?;
let filter = decoded.downcast_ref::<FilterExec>().unwrap();
let predicate = filter
.predicate()
.downcast_ref::<DynamicFilterPhysicalExpr>()
.unwrap();
let network = filter.input().downcast_ref::<NetworkShuffleExec>().unwrap();
let anchor = network.input_stage().dynamic_filter_anchors()[0]
.downcast_ref::<DynamicFilterPhysicalExpr>()
.unwrap();

assert_eq!(predicate.expression_id(), expected_id);
assert_eq!(anchor.expression_id(), expected_id);
predicate.update(lit(false))?;
assert_eq!(
anchor.current()?.to_string(),
"false",
"the filter predicate and network anchor should share state",
);
Ok(())
}

#[test]
fn test_roundtrip_union() -> datafusion::common::Result<()> {
let codec = DistributedCodec;
Expand Down
1 change: 1 addition & 0 deletions src/common/recursion.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1012,6 +1012,7 @@ mod tests {
num: 0,
workers: vec![],
runtime_stats: None,
dynamic_filter_anchors: vec![],
}))
.unwrap()
}
Expand Down
7 changes: 5 additions & 2 deletions src/coordinator/distributed.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,9 @@ use crate::coordinator::prepare_dynamic_plan::prepare_dynamic_plan;
use crate::coordinator::prepare_static_plan::prepare_static_plan;
use crate::coordinator::query_coordinator::QueryCoordinator;
use crate::coordinator::store::{Store, task_keys_for_plan};
use crate::dynamic_filtering::sever_dynamic_filter_relationships_in_plan_for_display;
use crate::dynamic_filtering::{
is_dynamic_filtering_enabled, sever_dynamic_filter_relationships_in_plan_for_display,
};
use crate::{DistributedConfig, TaskCompletedDynamicFilters, TaskKey, TaskMetrics};
use datafusion::common::internal_datafusion_err;
use datafusion::common::tree_node::TreeNodeRecursion;
Expand Down Expand Up @@ -249,7 +251,8 @@ impl ExecutionPlan for DistributedExec {
false => prepare_static_plan(&query_coordinator, &base_plan).await?,
};

prepared.plan_for_viz = match collect_dynamic_filters {
let dynamic_filtering_enabled = is_dynamic_filtering_enabled(context.session_config());
prepared.plan_for_viz = match dynamic_filtering_enabled && collect_dynamic_filters {
true => sever_dynamic_filter_relationships_in_plan_for_display(
prepared.plan_for_viz,
&context,
Expand Down
71 changes: 71 additions & 0 deletions src/coordinator/dynamic_filter_registry.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,71 @@
use crate::TaskKey;
use crate::dynamic_filtering::{
discover_dynamic_filter_consumers, discover_dynamic_filter_producers,
};
use datafusion::common::{HashMap, HashSet, Result};
use datafusion::physical_plan::ExecutionPlan;
use std::sync::{Arc, Mutex};

#[derive(Default)]
pub(super) struct PlannedDynamicFilter {
// Producer and consumer tasks for a dynamic filter.
//
// Note that it is not guaranteed that every task within a stage produces / consumes dynamic filters. For
// example, a distributed union may prevent a dynamic filter from appearing in all tasks. So, we
// store task keys rather than stage ids.
pub(super) producer_tasks: HashSet<TaskKey>,
pub(super) consumer_tasks: HashSet<TaskKey>,
}

#[derive(Default)]
pub(super) struct DynamicFilterRegistryState {
pub(super) filters: HashMap<u64, PlannedDynamicFilter>,

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

In future PRs, we add more mutex protected state in this struct.

}

/// Query-scoped hub for distributed dynamic filtering.
///
/// It stores the locations of dynamic filters and their runtime state. Informs the coordinator
/// - where dynamic filter updates are coming from
/// - how/if dynamic filter updates should be merged
/// - where dynamic filter updates should be forwarded
#[derive(Default)]
pub(crate) struct DynamicFilterRegistry {
pub(super) state: Mutex<DynamicFilterRegistryState>,
}

impl DynamicFilterRegistry {
pub(crate) fn new() -> Self {
Self::default()
}

/// Adds any dynamic filter producers and consumers found in `plan` to the registry.
pub(crate) fn register_task(
&self,
plan: &Arc<dyn ExecutionPlan>,
task_key: TaskKey,
) -> Result<()> {
let producers = discover_dynamic_filter_producers(plan)?;
// We can safely ignore anchors because they are not evaluated by network boundaries. This
// means they do not need updates forwarded to them.
let consumers = discover_dynamic_filter_consumers(plan)?.consumers;

let mut state = self.state.lock().expect("dynamic filter registry poisoned");
for producer in producers {
state
.filters
.entry(producer.id)
.or_default()
.producer_tasks
.insert(task_key);
}
for consumer in consumers {
state
.filters
.entry(consumer.id)
.or_default()
.consumer_tasks
.insert(task_key);
}
Ok(())
}
}
Loading
Loading