use std::sync::Arc;
use std::time::{Duration, Instant};
use nodedb_cluster::{ShuffleConsumeRequest, ShuffleConsumeResponse, TypedClusterError};
use nodedb_physical::physical_plan::{PhysicalPlan, QueryOp};
use crate::bridge::envelope::{Priority, Request, Status};
use crate::control::server::dispatch_utils::{DispatchCollectError, collect_bounded_response};
use crate::control::state::SharedState;
use crate::types::{DatabaseId, ReadConsistency, TenantId};
const SIDE_BUILD: u8 = 0;
const SIDE_PROBE: u8 = 1;
pub struct RegistryShuffleConsumer {
state: Arc<SharedState>,
}
impl RegistryShuffleConsumer {
pub fn new(state: Arc<SharedState>) -> Self {
Self { state }
}
async fn consume(&self, req: ShuffleConsumeRequest) -> Result<Vec<u8>, TypedClusterError> {
let registry = &self.state.shuffle_registry;
let build_inbox = registry
.get((req.shuffle_id, req.part, SIDE_BUILD))
.ok_or_else(|| TypedClusterError::Internal {
code: 0,
message: format!(
"shuffle consume: build inbox missing for (shuffle_id={}, part={}, side=0); \
no producer opened this part",
req.shuffle_id, req.part
),
})?;
let probe_inbox = registry
.get((req.shuffle_id, req.part, SIDE_PROBE))
.ok_or_else(|| TypedClusterError::Internal {
code: 0,
message: format!(
"shuffle consume: probe inbox missing for (shuffle_id={}, part={}, side=1); \
no producer opened this part",
req.shuffle_id, req.part
),
})?;
let deadline_ms = req.deadline_remaining_ms.max(1);
let wait_both = async {
build_inbox.wait_finalized().await;
probe_inbox.wait_finalized().await;
};
if tokio::time::timeout(Duration::from_millis(deadline_ms), wait_both)
.await
.is_err()
{
return Err(TypedClusterError::DeadlineExceeded {
elapsed_ms: deadline_ms,
});
}
if let Some(e) = build_inbox.take_error() {
return Err(e);
}
if let Some(e) = probe_inbox.take_error() {
return Err(e);
}
let build_path = build_inbox.staged_path().to_string_lossy().into_owned();
let probe_path = probe_inbox.staged_path().to_string_lossy().into_owned();
let on: Vec<(String, String)> = req
.on
.iter()
.map(|p| (p.left.clone(), p.right.clone()))
.collect();
let limit = usize::try_from(req.limit).unwrap_or(usize::MAX);
let plan = PhysicalPlan::Query(QueryOp::ShuffleJoinConsume {
build_path,
probe_path,
on,
join_type: req.join_type.clone(),
limit,
probe_qualifier: req.probe_qualifier.clone(),
index_qualifier: req.index_qualifier.clone(),
});
let deadline = Duration::from_millis(deadline_ms).min(Duration::from_secs(
self.state.tuning.network.default_deadline_secs,
));
let request_id = self.state.next_request_id();
let request = Request {
request_id,
tenant_id: TenantId::new(req.tenant_id),
database_id: DatabaseId::from(req.database_id),
vshard_id: crate::types::VShardId::new(0),
plan,
deadline: Instant::now() + deadline,
priority: Priority::Normal,
trace_id: nodedb_types::TraceId(req.trace_id),
consistency: ReadConsistency::Strong,
idempotency_key: None,
event_source: crate::event::EventSource::User,
user_roles: Vec::new(),
user_id: None,
statement_digest: None,
txn_id: None,
wal_lsn: None,
resolved_now_ms: None,
admission: crate::bridge::envelope::Admission::Exempt(
crate::bridge::envelope::ExemptReason::Read,
),
};
let mut rx = self.state.tracker.register(request_id);
let dispatch_result = match self.state.dispatcher.lock() {
Ok(mut d) => d.dispatch(request),
Err(poisoned) => poisoned.into_inner().dispatch(request),
};
if let Err(e) = dispatch_result {
return Err(TypedClusterError::Internal {
code: 0,
message: format!("shuffle consume dispatch failed: {e}"),
});
}
let max_result_bytes = self.state.tuning.network.max_query_result_bytes as usize;
match tokio::time::timeout(
deadline,
collect_bounded_response(&mut rx, max_result_bytes),
)
.await
{
Ok(Ok(resp)) => {
if resp.status == Status::Error {
let msg = resp
.error_code
.as_ref()
.map(|c| format!("{c:?}"))
.unwrap_or_else(|| "unknown error".into());
Err(TypedClusterError::Internal {
code: 0,
message: format!("shuffle consume join failed: {msg}"),
})
} else {
Ok(resp.payload.to_vec())
}
}
Ok(Err(DispatchCollectError::OverBudget { bytes })) => {
self.state.tracker.cancel(&request_id);
Err(TypedClusterError::Internal {
code: 0,
message: format!(
"shuffle consume join result exceeded max_query_result_bytes \
({bytes} > {max_result_bytes} bytes)"
),
})
}
Ok(Err(DispatchCollectError::ChannelClosed)) => Err(TypedClusterError::Internal {
code: 0,
message: "shuffle consume response channel closed".into(),
}),
Err(_) => {
self.state.tracker.cancel(&request_id);
Err(TypedClusterError::DeadlineExceeded {
elapsed_ms: deadline.as_millis() as u64,
})
}
}
}
}
#[async_trait::async_trait]
impl nodedb_cluster::ShuffleConsumer for RegistryShuffleConsumer {
async fn on_shuffle_consume(&self, req: ShuffleConsumeRequest) -> ShuffleConsumeResponse {
match self.consume(req).await {
Ok(rows) => ShuffleConsumeResponse { rows, error: None },
Err(error) => ShuffleConsumeResponse {
rows: Vec::new(),
error: Some(error),
},
}
}
}