use std::sync::Arc;
use std::time::{Duration, Instant};
use nodedb_cluster::{
ShuffleAggregateConsumeRequest, ShuffleAggregateConsumeResponse, TypedClusterError,
};
use nodedb_physical::physical_plan::{AggregateSpec, 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_PRODUCER: u8 = 0;
pub struct RegistryShuffleAggregator {
state: Arc<SharedState>,
}
impl RegistryShuffleAggregator {
pub fn new(state: Arc<SharedState>) -> Self {
Self { state }
}
async fn aggregate(
&self,
req: ShuffleAggregateConsumeRequest,
) -> Result<Vec<u8>, TypedClusterError> {
let registry = &self.state.shuffle_registry;
let inbox = registry
.get((req.shuffle_id, req.part, SIDE_PRODUCER))
.ok_or_else(|| TypedClusterError::Internal {
code: 0,
message: format!(
"shuffle aggregate: producer inbox missing for (shuffle_id={}, part={}, \
side=0); no producer opened this part",
req.shuffle_id, req.part
),
})?;
let deadline_ms = req.deadline_remaining_ms.max(1);
if tokio::time::timeout(Duration::from_millis(deadline_ms), inbox.wait_finalized())
.await
.is_err()
{
return Err(TypedClusterError::DeadlineExceeded {
elapsed_ms: deadline_ms,
});
}
if let Some(e) = inbox.take_error() {
return Err(e);
}
let state_path = inbox.staged_path().to_string_lossy().into_owned();
let aggregates: Vec<AggregateSpec> =
zerompk::from_msgpack(&req.aggregates_bytes).map_err(|e| {
TypedClusterError::Internal {
code: 0,
message: format!("shuffle aggregate: decode aggregate specs failed: {e}"),
}
})?;
let sort_keys: Vec<(String, bool)> = req
.sort_keys
.iter()
.map(|k| (k.column.clone(), k.ascending))
.collect();
let limit = usize::try_from(req.limit).unwrap_or(usize::MAX);
let plan = PhysicalPlan::Query(QueryOp::ShuffleAggregateConsume {
state_path,
group_by: req.group_by.clone(),
aggregates,
having: req.having.clone(),
limit,
sort_keys,
});
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 aggregate 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 aggregate merge 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 aggregate result exceeded max_query_result_bytes \
({bytes} > {max_result_bytes} bytes)"
),
})
}
Ok(Err(DispatchCollectError::ChannelClosed)) => Err(TypedClusterError::Internal {
code: 0,
message: "shuffle aggregate 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::ShuffleAggregator for RegistryShuffleAggregator {
async fn on_shuffle_aggregate(
&self,
req: ShuffleAggregateConsumeRequest,
) -> ShuffleAggregateConsumeResponse {
match self.aggregate(req).await {
Ok(rows) => ShuffleAggregateConsumeResponse { rows, error: None },
Err(error) => ShuffleAggregateConsumeResponse {
rows: Vec::new(),
error: Some(error),
},
}
}
}