Skip to main content

nodedb_cluster/rpc_codec/
shuffle.rs

1// SPDX-License-Identifier: BUSL-1.1
2
3//! `ShufflePush` streaming RPC — cross-node streaming shuffle (E1).
4//!
5//! A producer opens one bidi stream per target partition and writes a
6//! `ShufflePushRequest` frame, then a sequence of `ShufflePushChunk` frames
7//! (each a standalone msgpack array of rows, mirroring `ExecuteStreamChunk`),
8//! terminated by exactly one `ShufflePushEnd` frame. Unlike `ExecuteStream`
9//! the direction is producer → receiver: the chunks travel on the *send*
10//! half and the receiver does not write chunks back.
11//!
12//! Discriminants 25/26/27 are permanently assigned to these variants.
13
14use super::discriminants::*;
15use super::execute::{DescriptorVersionEntry, TypedClusterError};
16use super::header::write_frame;
17use super::raft_rpc::RaftRpc;
18use crate::error::{ClusterError, Result};
19
20// ── Wire types ──────────────────────────────────────────────────────────────
21
22/// Opening frame of a shuffle push stream.
23///
24/// Carries the routing key `(shuffle_id, part, side)` plus the partition fan-out
25/// (`num_parts`) and the number of producers (`producer_count`) the receiver
26/// must see an `End` from before the per-part build barrier is complete.
27///
28/// `side` is `0` for the build side and `1` for the probe side of a hash join.
29///
30/// Cross-version safety: new optional fields should be added as `Option<T>`.
31#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
32pub struct ShufflePushRequest {
33    pub shuffle_id: u64,
34    pub part: u32,
35    /// `0` = build side, `1` = probe side.
36    pub side: u8,
37    pub num_parts: u32,
38    pub producer_count: u32,
39}
40
41/// One streamed chunk of a shuffle push stream.
42///
43/// `payload` is a standalone msgpack array of row elements — the same
44/// convention as [`ExecuteStreamChunk`](super::execute::ExecuteStreamChunk)
45/// and `RowBatch.payload`.
46#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
47pub struct ShufflePushChunk {
48    pub payload: Vec<u8>,
49}
50
51/// Terminal frame of a shuffle push stream.
52///
53/// `error: None` is a clean EOF (all chunks delivered for this producer).
54/// `error: Some(e)` is a terminal failure — any chunks already delivered are
55/// valid, but this producer's contribution is incomplete and the receiver must
56/// surface the error. Mirrors
57/// [`ExecuteStreamEnd`](super::execute::ExecuteStreamEnd).
58#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
59pub struct ShufflePushEnd {
60    pub error: Option<TypedClusterError>,
61}
62
63/// One `part -> owning node id` mapping entry of a [`ShuffleProduceRequest`]'s
64/// `part_node_map`. The coordinator computes this routing table; the producer
65/// node uses it to direct each hash-partitioned row's `(part)` to its owner
66/// (looping back for parts it owns itself).
67///
68/// Cross-version safety: new optional fields should be added as `Option<T>`.
69#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
70pub struct PartNodeEntry {
71    pub part: u32,
72    pub node_id: u64,
73}
74
75/// Cross-node shuffle PRODUCER trigger (E4a).
76///
77/// A coordinator sends this to a producer node to make it execute a LOCAL scan
78/// fragment (`plan_bytes`), hash-partition each output row on `keys`, and fan the
79/// rows out to the per-part owners (`part_node_map`) as `ShufflePush` streams —
80/// one `ShufflePushEnd` to EVERY part for this `side` so each receiver's per-part
81/// barrier reaches `producer_count`. The producer replies with exactly one
82/// [`ShuffleProduceResponse`]; it never streams the scanned rows back.
83///
84/// `side` is `0` for the build side and `1` for the probe side of a hash join.
85///
86/// The `plan_bytes` / `tenant_id` / `database_id` / `deadline_remaining_ms` /
87/// `trace_id` / `descriptor_versions` fields mirror
88/// [`ExecuteRequest`](super::execute::ExecuteRequest) so the producer can reuse
89/// the existing local streaming-execution prologue verbatim.
90///
91/// Cross-version safety: new optional fields should be added as `Option<T>`.
92#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
93pub struct ShuffleProduceRequest {
94    pub shuffle_id: u64,
95    /// `0` = build side, `1` = probe side — the local side this producer scans.
96    pub side: u8,
97    pub num_parts: u32,
98    /// How many producers will `End` each `(part, side)`; forwarded into every
99    /// emitted `ShufflePushRequest` so receivers size their build barrier.
100    pub producer_count: u32,
101    /// Local-side join field names; each output row is hashed on these.
102    pub keys: Vec<String>,
103    /// `part -> owning node id` (coordinator-computed, sorted by `part`).
104    pub part_node_map: Vec<PartNodeEntry>,
105    /// Encoded `PhysicalPlan` of the local scan fragment to execute.
106    pub plan_bytes: Vec<u8>,
107    pub tenant_id: u64,
108    pub database_id: u64,
109    pub deadline_remaining_ms: u64,
110    pub trace_id: [u8; 16],
111    pub descriptor_versions: Vec<DescriptorVersionEntry>,
112}
113
114/// Terminal reply to a [`ShuffleProduceRequest`].
115///
116/// `error: None` means the producer fanned out every row and `End`ed every part
117/// cleanly. `error: Some(e)` means the local scan failed; the producer has
118/// already `End`ed every part with the same error so each receiver fails fast.
119///
120/// Cross-version safety: new fields are appended (mirroring
121/// [`ExecuteResponse`](super::execute::ExecuteResponse)'s LSN fields).
122#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
123pub struct ShuffleProduceResponse {
124    pub error: Option<TypedClusterError>,
125    /// Max per-collection read-version LSN observed by the producer's local scan
126    /// (its scanned collection's `coll_write_lsn` at read time, a WAL LSN); 0 for
127    /// a failed produce. The sound comparand the coordinator max-folds across
128    /// producers for cross-shard OCC read validation of an in-transaction
129    /// distributed aggregate — distinct from the core-global watermark. Raw `u64`
130    /// on the wire, converted to `Lsn` at the coordinator via `Lsn::new`. Mirrors
131    /// [`ExecuteResponse::read_version_lsn`](super::execute::ExecuteResponse::read_version_lsn).
132    pub read_version_lsn: u64,
133}
134
135/// One `(left_key, right_key)` equi-join pair of a [`ShuffleConsumeRequest`]'s
136/// `on` list. The part-owner reconstructs the borrowed join spec from these.
137///
138/// Cross-version safety: new optional fields should be added as `Option<T>`.
139#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
140pub struct JoinKeyPair {
141    /// Probe-side (left) join field name.
142    pub left: String,
143    /// Build-side (right) join field name.
144    pub right: String,
145}
146
147/// Cross-node shuffle CONSUMER trigger (E4b).
148///
149/// A coordinator sends this to a part-owner node to make it complete one part of
150/// a distributed shuffle join: wait for BOTH staged sides of `(shuffle_id, part)`
151/// to finalize, run the node-local grace-hash join over them, and reply with the
152/// join-result rows. The part-owner replies with exactly one
153/// [`ShuffleConsumeResponse`].
154///
155/// The `on` / `join_type` / `limit` / `probe_qualifier` / `index_qualifier`
156/// fields are the owned join spec the consumer reconstructs into a borrowed
157/// `JoinParams` for the node-local grace join. `tenant_id` / `database_id` /
158/// `deadline_remaining_ms` / `trace_id` mirror
159/// [`ExecuteRequest`](super::execute::ExecuteRequest) so the local dispatch
160/// prologue is reused.
161///
162/// Cross-version safety: new optional fields should be added as `Option<T>`.
163#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
164pub struct ShuffleConsumeRequest {
165    pub shuffle_id: u64,
166    /// The part this node owns and must complete.
167    pub part: u32,
168    /// Equi-join key pairs `(left_key, right_key)`.
169    pub on: Vec<JoinKeyPair>,
170    /// Join type (`inner` / `left` / `right` / `full`).
171    pub join_type: String,
172    /// Output row cap. `u64::MAX` = no explicit LIMIT (budget-bounded).
173    pub limit: u64,
174    /// Column qualifier (prefix) for probe-side (left) columns.
175    pub probe_qualifier: String,
176    /// Column qualifier (prefix) for build-side (right/index) columns.
177    pub index_qualifier: String,
178    pub tenant_id: u64,
179    pub database_id: u64,
180    pub deadline_remaining_ms: u64,
181    pub trace_id: [u8; 16],
182}
183
184/// Terminal reply to a [`ShuffleConsumeRequest`].
185///
186/// `error: None` carries the join-result rows in `rows` (a msgpack array of
187/// joined rows — the same `encode_binary_rows` shape every join path emits).
188/// `error: Some(e)` means the consume failed (missing inbox, finalize timeout,
189/// producer terminal error, or join error); `rows` is empty in that case.
190#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
191pub struct ShuffleConsumeResponse {
192    /// Msgpack array of join-result rows. Empty on error.
193    pub rows: Vec<u8>,
194    pub error: Option<TypedClusterError>,
195}
196
197/// One post-aggregation sort key of a [`ShuffleAggregateConsumeRequest`]'s
198/// `sort_keys` list: a column name plus its sort direction. A named struct
199/// (rather than a `(String, bool)` tuple) so rkyv can derive its codec, mirroring
200/// how [`JoinKeyPair`] wraps an `on` tuple.
201///
202/// Cross-version safety: new optional fields should be added as `Option<T>`.
203#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
204pub struct SortKey {
205    /// Sort column name.
206    pub column: String,
207    /// `true` = ascending, `false` = descending.
208    pub ascending: bool,
209}
210
211/// Cross-node distributed GROUP BY shuffle CONSUMER trigger (E5b).
212///
213/// Single-sided sibling of [`ShuffleConsumeRequest`]: a coordinator sends this to
214/// a part-owner node to make it complete one part of a distributed GROUP BY
215/// shuffle — wait for the part's ONE staged side (side `0`) to finalize, merge
216/// the staged partial `GroupState`s, finalize / HAVING-filter / sort / LIMIT, and
217/// reply with the result rows. The part-owner replies with exactly one
218/// [`ShuffleAggregateConsumeResponse`].
219///
220/// Unlike the join consumer this waits for only the single producer side (`0`);
221/// there is no probe side. The `group_by` / `aggregates_bytes` / `having` /
222/// `limit` / `sort_keys` fields are the owned aggregate spec the consumer
223/// reconstructs into a node-local `ShuffleAggregateConsume` plan. `tenant_id` /
224/// `database_id` / `deadline_remaining_ms` / `trace_id` mirror
225/// [`ExecuteRequest`](super::execute::ExecuteRequest) so the local dispatch
226/// prologue is reused.
227///
228/// Cross-version safety: new optional fields should be added as `Option<T>`.
229#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
230pub struct ShuffleAggregateConsumeRequest {
231    pub shuffle_id: u64,
232    /// The part this node owns and must complete.
233    pub part: u32,
234    /// GROUP BY columns (the key the partial states were grouped on).
235    pub group_by: Vec<String>,
236    /// Opaque zerompk-encoded `Vec<AggregateSpec>` (nodedb-cluster does not
237    /// inspect it; the host-crate hook decodes it — mirrors how `plan_bytes`
238    /// stays opaque to the transport).
239    pub aggregates_bytes: Vec<u8>,
240    /// Msgpack HAVING predicate blob (empty = no HAVING).
241    pub having: Vec<u8>,
242    /// Output row cap after sort. `u64::MAX` = no explicit LIMIT.
243    pub limit: u64,
244    /// Post-aggregation sort keys.
245    pub sort_keys: Vec<SortKey>,
246    pub tenant_id: u64,
247    pub database_id: u64,
248    pub deadline_remaining_ms: u64,
249    pub trace_id: [u8; 16],
250}
251
252/// Terminal reply to a [`ShuffleAggregateConsumeRequest`].
253///
254/// `error: None` carries the finalized GROUP BY result rows in `rows` (a msgpack
255/// array of rows — the same shape every aggregate path emits). `error: Some(e)`
256/// means the consume failed (missing inbox, finalize timeout, producer terminal
257/// error, or merge/finalize error); `rows` is empty in that case.
258#[derive(Debug, Clone, rkyv::Archive, rkyv::Serialize, rkyv::Deserialize)]
259pub struct ShuffleAggregateConsumeResponse {
260    /// Msgpack array of finalized aggregate rows. Empty on error.
261    pub rows: Vec<u8>,
262    pub error: Option<TypedClusterError>,
263}
264
265// ── Codec ────────────────────────────────────────────────────────────────────
266
267macro_rules! to_bytes {
268    ($msg:expr) => {
269        rkyv::to_bytes::<rkyv::rancor::Error>($msg)
270            .map(|b| b.to_vec())
271            .map_err(|e| ClusterError::Codec {
272                detail: format!("rkyv serialize: {e}"),
273            })
274    };
275}
276
277macro_rules! from_bytes {
278    ($payload:expr, $T:ty, $name:expr) => {{
279        let mut aligned = rkyv::util::AlignedVec::<16>::with_capacity($payload.len());
280        aligned.extend_from_slice($payload);
281        rkyv::from_bytes::<$T, rkyv::rancor::Error>(&aligned).map_err(|e| ClusterError::Codec {
282            detail: format!("rkyv deserialize {}: {e}", $name),
283        })
284    }};
285}
286
287pub(super) fn encode_shuffle_push_req(msg: &ShufflePushRequest, out: &mut Vec<u8>) -> Result<()> {
288    write_frame(RPC_SHUFFLE_PUSH_REQ, &to_bytes!(msg)?, out)
289}
290pub(super) fn encode_shuffle_push_chunk(msg: &ShufflePushChunk, out: &mut Vec<u8>) -> Result<()> {
291    write_frame(RPC_SHUFFLE_PUSH_CHUNK, &to_bytes!(msg)?, out)
292}
293pub(super) fn encode_shuffle_push_end(msg: &ShufflePushEnd, out: &mut Vec<u8>) -> Result<()> {
294    write_frame(RPC_SHUFFLE_PUSH_END, &to_bytes!(msg)?, out)
295}
296
297pub(super) fn decode_shuffle_push_req(payload: &[u8]) -> Result<RaftRpc> {
298    Ok(RaftRpc::ShufflePushRequest(from_bytes!(
299        payload,
300        ShufflePushRequest,
301        "ShufflePushRequest"
302    )?))
303}
304pub(super) fn decode_shuffle_push_chunk(payload: &[u8]) -> Result<RaftRpc> {
305    Ok(RaftRpc::ShufflePushChunk(from_bytes!(
306        payload,
307        ShufflePushChunk,
308        "ShufflePushChunk"
309    )?))
310}
311pub(super) fn decode_shuffle_push_end(payload: &[u8]) -> Result<RaftRpc> {
312    Ok(RaftRpc::ShufflePushEnd(from_bytes!(
313        payload,
314        ShufflePushEnd,
315        "ShufflePushEnd"
316    )?))
317}
318
319pub(super) fn encode_shuffle_produce_req(
320    msg: &ShuffleProduceRequest,
321    out: &mut Vec<u8>,
322) -> Result<()> {
323    write_frame(RPC_SHUFFLE_PRODUCE_REQ, &to_bytes!(msg)?, out)
324}
325pub(super) fn encode_shuffle_produce_resp(
326    msg: &ShuffleProduceResponse,
327    out: &mut Vec<u8>,
328) -> Result<()> {
329    write_frame(RPC_SHUFFLE_PRODUCE_RESP, &to_bytes!(msg)?, out)
330}
331
332pub(super) fn decode_shuffle_produce_req(payload: &[u8]) -> Result<RaftRpc> {
333    Ok(RaftRpc::ShuffleProduceRequest(from_bytes!(
334        payload,
335        ShuffleProduceRequest,
336        "ShuffleProduceRequest"
337    )?))
338}
339pub(super) fn decode_shuffle_produce_resp(payload: &[u8]) -> Result<RaftRpc> {
340    Ok(RaftRpc::ShuffleProduceResponse(from_bytes!(
341        payload,
342        ShuffleProduceResponse,
343        "ShuffleProduceResponse"
344    )?))
345}
346
347pub(super) fn encode_shuffle_consume_req(
348    msg: &ShuffleConsumeRequest,
349    out: &mut Vec<u8>,
350) -> Result<()> {
351    write_frame(RPC_SHUFFLE_CONSUME_REQ, &to_bytes!(msg)?, out)
352}
353pub(super) fn encode_shuffle_consume_resp(
354    msg: &ShuffleConsumeResponse,
355    out: &mut Vec<u8>,
356) -> Result<()> {
357    write_frame(RPC_SHUFFLE_CONSUME_RESP, &to_bytes!(msg)?, out)
358}
359
360pub(super) fn decode_shuffle_consume_req(payload: &[u8]) -> Result<RaftRpc> {
361    Ok(RaftRpc::ShuffleConsumeRequest(from_bytes!(
362        payload,
363        ShuffleConsumeRequest,
364        "ShuffleConsumeRequest"
365    )?))
366}
367pub(super) fn decode_shuffle_consume_resp(payload: &[u8]) -> Result<RaftRpc> {
368    Ok(RaftRpc::ShuffleConsumeResponse(from_bytes!(
369        payload,
370        ShuffleConsumeResponse,
371        "ShuffleConsumeResponse"
372    )?))
373}
374
375pub(super) fn encode_shuffle_agg_consume_req(
376    msg: &ShuffleAggregateConsumeRequest,
377    out: &mut Vec<u8>,
378) -> Result<()> {
379    write_frame(RPC_SHUFFLE_AGG_CONSUME_REQ, &to_bytes!(msg)?, out)
380}
381pub(super) fn encode_shuffle_agg_consume_resp(
382    msg: &ShuffleAggregateConsumeResponse,
383    out: &mut Vec<u8>,
384) -> Result<()> {
385    write_frame(RPC_SHUFFLE_AGG_CONSUME_RESP, &to_bytes!(msg)?, out)
386}
387
388pub(super) fn decode_shuffle_agg_consume_req(payload: &[u8]) -> Result<RaftRpc> {
389    Ok(RaftRpc::ShuffleAggregateConsumeRequest(from_bytes!(
390        payload,
391        ShuffleAggregateConsumeRequest,
392        "ShuffleAggregateConsumeRequest"
393    )?))
394}
395pub(super) fn decode_shuffle_agg_consume_resp(payload: &[u8]) -> Result<RaftRpc> {
396    Ok(RaftRpc::ShuffleAggregateConsumeResponse(from_bytes!(
397        payload,
398        ShuffleAggregateConsumeResponse,
399        "ShuffleAggregateConsumeResponse"
400    )?))
401}
402
403#[cfg(test)]
404mod tests {
405    use super::*;
406
407    fn roundtrip_req(req: ShufflePushRequest) -> ShufflePushRequest {
408        let rpc = RaftRpc::ShufflePushRequest(req);
409        let encoded = super::super::encode(&rpc).unwrap();
410        match super::super::decode(&encoded).unwrap() {
411            RaftRpc::ShufflePushRequest(r) => r,
412            other => panic!("expected ShufflePushRequest, got {other:?}"),
413        }
414    }
415
416    fn roundtrip_chunk(chunk: ShufflePushChunk) -> ShufflePushChunk {
417        let rpc = RaftRpc::ShufflePushChunk(chunk);
418        let encoded = super::super::encode(&rpc).unwrap();
419        match super::super::decode(&encoded).unwrap() {
420            RaftRpc::ShufflePushChunk(c) => c,
421            other => panic!("expected ShufflePushChunk, got {other:?}"),
422        }
423    }
424
425    fn roundtrip_end(end: ShufflePushEnd) -> ShufflePushEnd {
426        let rpc = RaftRpc::ShufflePushEnd(end);
427        let encoded = super::super::encode(&rpc).unwrap();
428        match super::super::decode(&encoded).unwrap() {
429            RaftRpc::ShufflePushEnd(e) => e,
430            other => panic!("expected ShufflePushEnd, got {other:?}"),
431        }
432    }
433
434    #[test]
435    fn roundtrip_shuffle_push_request() {
436        let req = ShufflePushRequest {
437            shuffle_id: 0xDEAD_BEEF_1234_5678,
438            part: 7,
439            side: 1,
440            num_parts: 16,
441            producer_count: 3,
442        };
443        let decoded = roundtrip_req(req.clone());
444        assert_eq!(decoded.shuffle_id, req.shuffle_id);
445        assert_eq!(decoded.part, 7);
446        assert_eq!(decoded.side, 1);
447        assert_eq!(decoded.num_parts, 16);
448        assert_eq!(decoded.producer_count, 3);
449    }
450
451    #[test]
452    fn roundtrip_shuffle_push_request_build_side() {
453        let req = ShufflePushRequest {
454            shuffle_id: 1,
455            part: 0,
456            side: 0,
457            num_parts: 1,
458            producer_count: 1,
459        };
460        let decoded = roundtrip_req(req);
461        assert_eq!(decoded.side, 0);
462        assert_eq!(decoded.num_parts, 1);
463        assert_eq!(decoded.producer_count, 1);
464    }
465
466    #[test]
467    fn roundtrip_shuffle_push_chunk_payload() {
468        let chunk = ShufflePushChunk {
469            payload: vec![0x93, 0x01, 0x02, 0x03],
470        };
471        let decoded = roundtrip_chunk(chunk.clone());
472        assert_eq!(decoded.payload, chunk.payload);
473    }
474
475    #[test]
476    fn roundtrip_shuffle_push_chunk_empty_payload() {
477        let decoded = roundtrip_chunk(ShufflePushChunk { payload: vec![] });
478        assert!(decoded.payload.is_empty());
479    }
480
481    #[test]
482    fn roundtrip_shuffle_push_end_clean_eof() {
483        let decoded = roundtrip_end(ShufflePushEnd { error: None });
484        assert!(decoded.error.is_none());
485    }
486
487    #[test]
488    fn roundtrip_shuffle_push_end_terminal_error() {
489        let decoded = roundtrip_end(ShufflePushEnd {
490            error: Some(TypedClusterError::Internal {
491                code: 0xABCD,
492                message: "shuffle producer failed mid-flight".into(),
493            }),
494        });
495        match decoded.error {
496            Some(TypedClusterError::Internal { code, message }) => {
497                assert_eq!(code, 0xABCD);
498                assert!(message.contains("shuffle producer"));
499            }
500            other => panic!("expected Internal, got {other:?}"),
501        }
502    }
503
504    fn roundtrip_produce_req(req: ShuffleProduceRequest) -> ShuffleProduceRequest {
505        let rpc = RaftRpc::ShuffleProduceRequest(req);
506        let encoded = super::super::encode(&rpc).unwrap();
507        match super::super::decode(&encoded).unwrap() {
508            RaftRpc::ShuffleProduceRequest(r) => r,
509            other => panic!("expected ShuffleProduceRequest, got {other:?}"),
510        }
511    }
512
513    fn roundtrip_produce_resp(resp: ShuffleProduceResponse) -> ShuffleProduceResponse {
514        let rpc = RaftRpc::ShuffleProduceResponse(resp);
515        let encoded = super::super::encode(&rpc).unwrap();
516        match super::super::decode(&encoded).unwrap() {
517            RaftRpc::ShuffleProduceResponse(r) => r,
518            other => panic!("expected ShuffleProduceResponse, got {other:?}"),
519        }
520    }
521
522    #[test]
523    fn roundtrip_shuffle_produce_request() {
524        let req = ShuffleProduceRequest {
525            shuffle_id: 0x1234_5678_9ABC_DEF0,
526            side: 1,
527            num_parts: 4,
528            producer_count: 2,
529            keys: vec!["k".into(), "tenant.id".into()],
530            part_node_map: vec![
531                PartNodeEntry {
532                    part: 0,
533                    node_id: 7,
534                },
535                PartNodeEntry {
536                    part: 1,
537                    node_id: 9,
538                },
539                PartNodeEntry {
540                    part: 2,
541                    node_id: 7,
542                },
543                PartNodeEntry {
544                    part: 3,
545                    node_id: 9,
546                },
547            ],
548            plan_bytes: vec![0xDE, 0xAD, 0xBE, 0xEF],
549            tenant_id: 42,
550            database_id: 3,
551            deadline_remaining_ms: 7000,
552            trace_id: [5u8; 16],
553            descriptor_versions: vec![DescriptorVersionEntry {
554                collection: "orders".into(),
555                version: 11,
556            }],
557        };
558        let decoded = roundtrip_produce_req(req.clone());
559        assert_eq!(decoded.shuffle_id, req.shuffle_id);
560        assert_eq!(decoded.side, 1);
561        assert_eq!(decoded.num_parts, 4);
562        assert_eq!(decoded.producer_count, 2);
563        assert_eq!(decoded.keys, vec!["k".to_string(), "tenant.id".to_string()]);
564        assert_eq!(decoded.part_node_map.len(), 4);
565        assert_eq!(decoded.part_node_map[2].part, 2);
566        assert_eq!(decoded.part_node_map[2].node_id, 7);
567        assert_eq!(decoded.plan_bytes, vec![0xDE, 0xAD, 0xBE, 0xEF]);
568        assert_eq!(decoded.tenant_id, 42);
569        assert_eq!(decoded.database_id, 3);
570        assert_eq!(decoded.deadline_remaining_ms, 7000);
571        assert_eq!(decoded.trace_id, [5u8; 16]);
572        assert_eq!(decoded.descriptor_versions.len(), 1);
573        assert_eq!(decoded.descriptor_versions[0].collection, "orders");
574        assert_eq!(decoded.descriptor_versions[0].version, 11);
575    }
576
577    #[test]
578    fn roundtrip_shuffle_produce_request_empty_keys_and_map() {
579        let req = ShuffleProduceRequest {
580            shuffle_id: 1,
581            side: 0,
582            num_parts: 1,
583            producer_count: 1,
584            keys: vec![],
585            part_node_map: vec![],
586            plan_bytes: vec![],
587            tenant_id: 0,
588            database_id: 0,
589            deadline_remaining_ms: 1000,
590            trace_id: [0u8; 16],
591            descriptor_versions: vec![],
592        };
593        let decoded = roundtrip_produce_req(req);
594        assert!(decoded.keys.is_empty());
595        assert!(decoded.part_node_map.is_empty());
596        assert!(decoded.descriptor_versions.is_empty());
597    }
598
599    #[test]
600    fn roundtrip_shuffle_produce_response_clean() {
601        let decoded = roundtrip_produce_resp(ShuffleProduceResponse {
602            error: None,
603            read_version_lsn: 0xABCD_1234,
604        });
605        assert!(decoded.error.is_none());
606        assert_eq!(
607            decoded.read_version_lsn, 0xABCD_1234,
608            "producer read-version LSN roundtrips on the produce reply"
609        );
610    }
611
612    #[test]
613    fn roundtrip_shuffle_produce_response_error() {
614        let decoded = roundtrip_produce_resp(ShuffleProduceResponse {
615            error: Some(TypedClusterError::Internal {
616                code: 0x55,
617                message: "produce scan failed".into(),
618            }),
619            read_version_lsn: 0,
620        });
621        assert_eq!(
622            decoded.read_version_lsn, 0,
623            "a failed produce carries no read-version LSN"
624        );
625        match decoded.error {
626            Some(TypedClusterError::Internal { code, message }) => {
627                assert_eq!(code, 0x55);
628                assert!(message.contains("produce scan"));
629            }
630            other => panic!("expected Internal, got {other:?}"),
631        }
632    }
633
634    fn roundtrip_consume_req(req: ShuffleConsumeRequest) -> ShuffleConsumeRequest {
635        let rpc = RaftRpc::ShuffleConsumeRequest(req);
636        let encoded = super::super::encode(&rpc).unwrap();
637        match super::super::decode(&encoded).unwrap() {
638            RaftRpc::ShuffleConsumeRequest(r) => r,
639            other => panic!("expected ShuffleConsumeRequest, got {other:?}"),
640        }
641    }
642
643    fn roundtrip_consume_resp(resp: ShuffleConsumeResponse) -> ShuffleConsumeResponse {
644        let rpc = RaftRpc::ShuffleConsumeResponse(resp);
645        let encoded = super::super::encode(&rpc).unwrap();
646        match super::super::decode(&encoded).unwrap() {
647            RaftRpc::ShuffleConsumeResponse(r) => r,
648            other => panic!("expected ShuffleConsumeResponse, got {other:?}"),
649        }
650    }
651
652    #[test]
653    fn roundtrip_shuffle_consume_request() {
654        let req = ShuffleConsumeRequest {
655            shuffle_id: 0x0FED_CBA9_8765_4321,
656            part: 3,
657            on: vec![
658                JoinKeyPair {
659                    left: "lk".into(),
660                    right: "rk".into(),
661                },
662                JoinKeyPair {
663                    left: "tenant".into(),
664                    right: "tenant_id".into(),
665                },
666            ],
667            join_type: "inner".into(),
668            limit: 1234,
669            probe_qualifier: "l".into(),
670            index_qualifier: "r".into(),
671            tenant_id: 9,
672            database_id: 4,
673            deadline_remaining_ms: 8000,
674            trace_id: [3u8; 16],
675        };
676        let decoded = roundtrip_consume_req(req.clone());
677        assert_eq!(decoded.shuffle_id, req.shuffle_id);
678        assert_eq!(decoded.part, 3);
679        assert_eq!(decoded.on.len(), 2);
680        assert_eq!(decoded.on[0].left, "lk");
681        assert_eq!(decoded.on[0].right, "rk");
682        assert_eq!(decoded.on[1].left, "tenant");
683        assert_eq!(decoded.on[1].right, "tenant_id");
684        assert_eq!(decoded.join_type, "inner");
685        assert_eq!(decoded.limit, 1234);
686        assert_eq!(decoded.probe_qualifier, "l");
687        assert_eq!(decoded.index_qualifier, "r");
688        assert_eq!(decoded.tenant_id, 9);
689        assert_eq!(decoded.database_id, 4);
690        assert_eq!(decoded.deadline_remaining_ms, 8000);
691        assert_eq!(decoded.trace_id, [3u8; 16]);
692    }
693
694    #[test]
695    fn roundtrip_shuffle_consume_request_empty_keys() {
696        let req = ShuffleConsumeRequest {
697            shuffle_id: 1,
698            part: 0,
699            on: vec![],
700            join_type: "left".into(),
701            limit: u64::MAX,
702            probe_qualifier: String::new(),
703            index_qualifier: String::new(),
704            tenant_id: 0,
705            database_id: 0,
706            deadline_remaining_ms: 1000,
707            trace_id: [0u8; 16],
708        };
709        let decoded = roundtrip_consume_req(req);
710        assert!(decoded.on.is_empty());
711        assert_eq!(decoded.limit, u64::MAX);
712        assert_eq!(decoded.join_type, "left");
713    }
714
715    #[test]
716    fn roundtrip_shuffle_consume_response_rows() {
717        let decoded = roundtrip_consume_resp(ShuffleConsumeResponse {
718            rows: vec![0x92, 0x01, 0x02],
719            error: None,
720        });
721        assert_eq!(decoded.rows, vec![0x92, 0x01, 0x02]);
722        assert!(decoded.error.is_none());
723    }
724
725    #[test]
726    fn roundtrip_shuffle_consume_response_error() {
727        let decoded = roundtrip_consume_resp(ShuffleConsumeResponse {
728            rows: vec![],
729            error: Some(TypedClusterError::DeadlineExceeded { elapsed_ms: 8000 }),
730        });
731        assert!(decoded.rows.is_empty());
732        match decoded.error {
733            Some(TypedClusterError::DeadlineExceeded { elapsed_ms }) => {
734                assert_eq!(elapsed_ms, 8000);
735            }
736            other => panic!("expected DeadlineExceeded, got {other:?}"),
737        }
738    }
739
740    fn roundtrip_agg_consume_req(
741        req: ShuffleAggregateConsumeRequest,
742    ) -> ShuffleAggregateConsumeRequest {
743        let rpc = RaftRpc::ShuffleAggregateConsumeRequest(req);
744        let encoded = super::super::encode(&rpc).unwrap();
745        match super::super::decode(&encoded).unwrap() {
746            RaftRpc::ShuffleAggregateConsumeRequest(r) => r,
747            other => panic!("expected ShuffleAggregateConsumeRequest, got {other:?}"),
748        }
749    }
750
751    fn roundtrip_agg_consume_resp(
752        resp: ShuffleAggregateConsumeResponse,
753    ) -> ShuffleAggregateConsumeResponse {
754        let rpc = RaftRpc::ShuffleAggregateConsumeResponse(resp);
755        let encoded = super::super::encode(&rpc).unwrap();
756        match super::super::decode(&encoded).unwrap() {
757            RaftRpc::ShuffleAggregateConsumeResponse(r) => r,
758            other => panic!("expected ShuffleAggregateConsumeResponse, got {other:?}"),
759        }
760    }
761
762    #[test]
763    fn roundtrip_shuffle_agg_consume_request() {
764        let req = ShuffleAggregateConsumeRequest {
765            shuffle_id: 0x0FED_CBA9_8765_4321,
766            part: 2,
767            group_by: vec!["k".into(), "region".into()],
768            aggregates_bytes: vec![0x91, 0x01, 0x02],
769            having: vec![0xC0],
770            limit: 4321,
771            sort_keys: vec![
772                SortKey {
773                    column: "k".into(),
774                    ascending: true,
775                },
776                SortKey {
777                    column: "total".into(),
778                    ascending: false,
779                },
780            ],
781            tenant_id: 9,
782            database_id: 4,
783            deadline_remaining_ms: 8000,
784            trace_id: [3u8; 16],
785        };
786        let decoded = roundtrip_agg_consume_req(req.clone());
787        assert_eq!(decoded.shuffle_id, req.shuffle_id);
788        assert_eq!(decoded.part, 2);
789        assert_eq!(
790            decoded.group_by,
791            vec!["k".to_string(), "region".to_string()]
792        );
793        assert_eq!(decoded.aggregates_bytes, vec![0x91, 0x01, 0x02]);
794        assert_eq!(decoded.having, vec![0xC0]);
795        assert_eq!(decoded.limit, 4321);
796        assert_eq!(decoded.sort_keys.len(), 2);
797        assert_eq!(decoded.sort_keys[0].column, "k");
798        assert!(decoded.sort_keys[0].ascending);
799        assert_eq!(decoded.sort_keys[1].column, "total");
800        assert!(!decoded.sort_keys[1].ascending);
801        assert_eq!(decoded.tenant_id, 9);
802        assert_eq!(decoded.database_id, 4);
803        assert_eq!(decoded.deadline_remaining_ms, 8000);
804        assert_eq!(decoded.trace_id, [3u8; 16]);
805    }
806
807    #[test]
808    fn roundtrip_shuffle_agg_consume_request_empty() {
809        let req = ShuffleAggregateConsumeRequest {
810            shuffle_id: 1,
811            part: 0,
812            group_by: vec![],
813            aggregates_bytes: vec![],
814            having: vec![],
815            limit: u64::MAX,
816            sort_keys: vec![],
817            tenant_id: 0,
818            database_id: 0,
819            deadline_remaining_ms: 1000,
820            trace_id: [0u8; 16],
821        };
822        let decoded = roundtrip_agg_consume_req(req);
823        assert!(decoded.group_by.is_empty());
824        assert!(decoded.aggregates_bytes.is_empty());
825        assert!(decoded.sort_keys.is_empty());
826        assert_eq!(decoded.limit, u64::MAX);
827    }
828
829    #[test]
830    fn roundtrip_shuffle_agg_consume_response_rows() {
831        let decoded = roundtrip_agg_consume_resp(ShuffleAggregateConsumeResponse {
832            rows: vec![0x92, 0x01, 0x02],
833            error: None,
834        });
835        assert_eq!(decoded.rows, vec![0x92, 0x01, 0x02]);
836        assert!(decoded.error.is_none());
837    }
838
839    #[test]
840    fn roundtrip_shuffle_agg_consume_response_error() {
841        let decoded = roundtrip_agg_consume_resp(ShuffleAggregateConsumeResponse {
842            rows: vec![],
843            error: Some(TypedClusterError::DeadlineExceeded { elapsed_ms: 9000 }),
844        });
845        assert!(decoded.rows.is_empty());
846        match decoded.error {
847            Some(TypedClusterError::DeadlineExceeded { elapsed_ms }) => {
848                assert_eq!(elapsed_ms, 9000);
849            }
850            other => panic!("expected DeadlineExceeded, got {other:?}"),
851        }
852    }
853}