datafusion_distributed/protocol/
channel_resolver.rs1use 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#[async_trait]
27pub trait ChannelResolver {
28 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>>);