Skip to main content

datafusion_distributed/protocol/
channel_resolver.rs

1use crate::WorkerChannel;
2use crate::distributed_planner::DistributedConfig;
3#[cfg(feature = "grpc")]
4use crate::protocol::grpc;
5use async_trait::async_trait;
6use datafusion::common::DataFusionError;
7use datafusion::execution::TaskContext;
8use datafusion::prelude::SessionConfig;
9use std::sync::Arc;
10use url::Url;
11
12/// Allows users to customize the way Worker clients are created. A common use case is to
13/// wrap the client with tower layers or schedule it in an IO-specific tokio runtime.
14///
15/// There is a default implementation of this trait that should be enough for the most common
16/// use-cases.
17///
18/// # Implementation Notes
19/// - This is called per request, so implementors of this trait should make sure that
20///   clients are reused across method calls instead of building a new Worker client every time.
21///
22/// - When implementing `get_worker_client_for_url`, it is recommended to use the
23///   [`create_worker_client`] helper function to ensure clients are configured with
24///   appropriate message size limits for internal communication. This helps avoid message
25///   size errors when transferring large datasets.
26#[async_trait]
27pub trait ChannelResolver {
28    /// For a given URL, get a Worker gRPC client for communicating to it.
29    ///
30    /// *WARNING*: This method is called for every gRPC request, so to not create
31    /// one client connection for each request, users are required to reuse generated clients.
32    /// It's recommended to rely on [DefaultChannelResolver] either by delegating method calls
33    /// to it or by copying the implementation.
34    ///
35    /// Consider using [`create_worker_client`] to create the client with appropriate
36    /// default message size limits.
37    async fn get_worker_client_for_url(
38        &self,
39        url: &Url,
40    ) -> Result<Box<dyn WorkerChannel>, DataFusionError>;
41}
42
43pub(crate) fn set_distributed_channel_resolver(
44    cfg: &mut SessionConfig,
45    channel_resolver: impl ChannelResolver + Send + Sync + 'static,
46) {
47    cfg.set_extension(Arc::new(ChannelResolverExtension(Some(Arc::new(
48        channel_resolver,
49    )))));
50    DistributedConfig::ensure_in_config(cfg);
51}
52
53pub fn get_distributed_channel_resolver(
54    task_ctx: &TaskContext,
55) -> Arc<dyn ChannelResolver + Send + Sync> {
56    let session_cfg = task_ctx.session_config();
57    if let Some(channel_resolver_ext) = session_cfg.get_extension::<ChannelResolverExtension>()
58        && let Some(cr) = &channel_resolver_ext.0
59    {
60        return Arc::clone(cr);
61    }
62
63    #[cfg(feature = "grpc")]
64    {
65        let runtime_addr = Arc::as_ptr(&task_ctx.runtime_env()) as usize;
66        grpc::DEFAULT_CHANNEL_RESOLVER_PER_RUNTIME.get_with(runtime_addr, || {
67            Arc::new(grpc::DefaultChannelResolver::default())
68        })
69    }
70
71    #[cfg(not(feature = "grpc"))]
72    {
73        panic!(
74            "gRPC feature is not enabled, and no channel resolver was provided, so no default ChannelResolver can be provided"
75        );
76    }
77}
78
79#[async_trait]
80impl ChannelResolver for Arc<dyn ChannelResolver + Send + Sync> {
81    async fn get_worker_client_for_url(
82        &self,
83        url: &Url,
84    ) -> Result<Box<dyn WorkerChannel>, DataFusionError> {
85        self.as_ref().get_worker_client_for_url(url).await
86    }
87}
88
89#[derive(Clone, Default)]
90pub(crate) struct ChannelResolverExtension(Option<Arc<dyn ChannelResolver + Send + Sync>>);