Skip to main content

hurray_core/
composite.rs

1//! Composite tensor: head + members, aggregated across whole `TensorDescriptor`s.
2//!
3//! Every other layout in this crate is validated and addressed within a single
4//! [`crate::descriptor::TensorDescriptor`]. Composite is different: a composite
5//! "tensor" is a **head** descriptor (`layout_tag = 0x0B`, see
6//! [`crate::layout::composite`]) plus **N** ordinary tensor descriptors — its
7//! **members** — bound by forward stream/file adjacency. This module is where that
8//! cross-descriptor aggregate lives.
9//!
10//! Distinct from [`crate::layout::composite`], which defines only the wire-level
11//! `CompositeLayout` payload carried *inside* the head descriptor's layout-specific
12//! fields (`composition_rule`, `combine_op`, `member_count`). This module aggregates
13//! a complete head plus its actual member descriptors and validates them against
14//! each other per `docs/spec/layouts/composite.md` § Validation. The two modules are
15//! kept at distinct paths (`hurray_core::layout::composite` vs
16//! `hurray_core::composite`) precisely so their similarly-named types — `CompositeLayout`
17//! / `CompositionRule` / `CombineOp` here vs [`CompositeTensor`] / [`CompositeValidator`]
18//! there — never collide in any glob or re-export.
19
20use crate::descriptor::TensorDescriptor;
21use crate::descriptor::{CompositeMemberDescriptor, MemberRole};
22use crate::layout::{CompositionRule, LayoutDescriptor};
23use crate::{Error, Result};
24
25/// Maximum composite nesting depth.
26///
27/// Mirrors the recursion-depth discipline used for nested Tiled layouts
28/// ([`crate::layout::MAX_TILED_DEPTH`]), per `docs/spec/layouts/composite.md`
29/// § Binding: "a reader MUST enforce a maximum composite nesting depth of 8
30/// levels." A member whose own `layout_tag` is the composite tag counts as one
31/// additional level.
32pub const MAX_COMPOSITE_DEPTH: u8 = 8;
33
34/// A composite tensor: a head descriptor plus its ordered member descriptors.
35///
36/// Construct via [`CompositeTensor::new`], which drives a [`CompositeValidator`]
37/// across every member and rejects a malformed head+members set — see that type's
38/// docs for the full per-member and close-time checks performed.
39///
40/// # Nesting depth — a note on this layer's limitation
41///
42/// A member here is a flat [`TensorDescriptor`], not a further-expanded
43/// [`CompositeTensor`]. If a member's own `layout_tag` happens to be the composite
44/// tag, [`CompositeTensor::new`] counts that as one additional nesting level for
45/// the depth cap, but it cannot recurse into *that* member's own members — they
46/// are not available at this layer (Layer 4 sees single descriptors, not a
47/// pre-assembled tree). Full nesting-depth validation therefore cannot be
48/// completed until a higher layer (Layer 5 streaming or Layer 6 file, which do
49/// assemble the full nested tree as they parse it) walks the tree top to bottom.
50/// [`CompositeTensor::new`] is depth `0`; a `pub(crate)` depth-aware constructor
51/// is provided for that future recursive caller.
52///
53/// # Examples
54///
55/// ```
56/// use hurray_core::{
57///     composite::CompositeTensor,
58///     descriptor::TensorDescriptor,
59///     layout::{CompositeLayout, CompositionRule, LayoutDescriptor},
60///     BufferHandle, DeviceTag, ElementType, Shape, ShardDescriptor, SyncMode,
61///     MIN_BUFFER_ALIGNMENT,
62/// };
63///
64/// // Head: float32 [8, 8], partition of 2 members.
65/// let head = TensorDescriptor::new(
66///     1, 0, ElementType::Float32, Shape::new(vec![8u64, 8]).unwrap(), 0,
67///     LayoutDescriptor::Composite(CompositeLayout::new(CompositionRule::Partition, 2).unwrap()),
68///     vec![],
69///     None, None, None, None,
70/// ).unwrap();
71///
72/// // Two [8, 4] members tiling the head's [8, 8] index space.
73/// let member = |offset: u64| {
74///     let buf = BufferHandle::new(128, MIN_BUFFER_ALIGNMENT, DeviceTag::Cpu, SyncMode::ProducerSynced).unwrap();
75///     let shard = ShardDescriptor::new(vec![8, 8], vec![0, offset]).unwrap();
76///     TensorDescriptor::new(
77///         1, 0, ElementType::Float32, Shape::new(vec![8u64, 4]).unwrap(), 0,
78///         LayoutDescriptor::RowMajor, vec![buf],
79///         None, Some(shard), None, None,
80///     ).unwrap()
81/// };
82///
83/// let composite = CompositeTensor::new(head, vec![member(0), member(4)]).unwrap();
84/// assert_eq!(composite.member_count(), 2);
85/// assert_eq!(*composite.rule(), CompositionRule::Partition);
86/// ```
87#[derive(Debug, Clone, PartialEq)]
88pub struct CompositeTensor {
89    /// The composite head descriptor (`layout_tag = 0x0B`).
90    pub head: TensorDescriptor,
91    /// The ordered member descriptors bound to the head.
92    pub members: Vec<TensorDescriptor>,
93}
94
95impl CompositeTensor {
96    /// Creates a new [`CompositeTensor`], validating `members` against `head`
97    /// via [`CompositeValidator`] (per-member checks, then close-time checks).
98    ///
99    /// Equivalent to `CompositeTensor::new_at_depth(head, members, 0)`.
100    ///
101    /// # Errors
102    ///
103    /// Propagates any [`CompositeValidator::push_member`] / [`CompositeValidator::finish`]
104    /// error, plus [`Error::InvalidLayout`] if `head.layout` is not
105    /// [`LayoutDescriptor::Composite`].
106    ///
107    /// # Examples
108    ///
109    /// See the module-level and struct-level examples.
110    pub fn new(head: TensorDescriptor, members: Vec<TensorDescriptor>) -> Result<Self> {
111        Self::new_at_depth(head, members, 0)
112    }
113
114    /// Depth-aware constructor for recursive assembly by a higher layer.
115    ///
116    /// `depth` is the nesting depth of `head` itself (the top-level call from
117    /// [`CompositeTensor::new`] uses `0`). A future Layer 5/6 caller that has
118    /// assembled a full nested composite tree can recurse with `depth + 1` for
119    /// each member that is itself a composite head — this layer cannot do that
120    /// recursion itself (see the struct-level docs), but it does track and
121    /// enforce the depth budget for the one level of nesting it *can* see: a
122    /// member whose own `layout_tag` is the composite tag.
123    pub(crate) fn new_at_depth(
124        head: TensorDescriptor,
125        members: Vec<TensorDescriptor>,
126        depth: u8,
127    ) -> Result<Self> {
128        if depth >= MAX_COMPOSITE_DEPTH {
129            return Err(Error::SubpavingNestingTooDeep);
130        }
131
132        let mut validator = CompositeValidator::new(&head)?;
133        for member in &members {
134            if matches!(member.layout, LayoutDescriptor::Composite(_))
135                && depth + 1 >= MAX_COMPOSITE_DEPTH
136            {
137                return Err(Error::SubpavingNestingTooDeep);
138            }
139            validator.push_member(member)?;
140        }
141        validator.finish()?;
142
143        Ok(Self { head, members })
144    }
145
146    /// Returns the composition rule declared by the head.
147    ///
148    /// # Examples
149    ///
150    /// See the struct-level example.
151    pub fn rule(&self) -> &CompositionRule {
152        match &self.head.layout {
153            LayoutDescriptor::Composite(c) => &c.rule,
154            // INVARIANT: CompositeTensor::new only returns Ok when head.layout is
155            // Composite (CompositeValidator::new rejects any other layout first).
156            // Struct fields are `pub` (matching TensorDescriptor's own convention
157            // elsewhere in this crate), so a caller COULD hand-build an invalid
158            // instance bypassing `new()`; this is the same trust boundary the rest
159            // of the crate already accepts for its `pub`-field descriptor types.
160            _ => unreachable!(
161                "CompositeTensor invariant: head.layout is always LayoutDescriptor::Composite"
162            ),
163        }
164    }
165
166    /// Returns the member count declared by the head.
167    ///
168    /// # Examples
169    ///
170    /// See the struct-level example.
171    pub fn member_count(&self) -> u32 {
172        match &self.head.layout {
173            LayoutDescriptor::Composite(c) => c.member_count,
174            _ => unreachable!(
175                "CompositeTensor invariant: head.layout is always LayoutDescriptor::Composite"
176            ),
177        }
178    }
179}
180
181/// Stateful, bounded validator for a composite head's members.
182///
183/// Drive it with [`CompositeValidator::new`] (extracts the composition rule from
184/// the head), then [`CompositeValidator::push_member`] once per member in stream
185/// order, then [`CompositeValidator::finish`] to run the close-time checks. This
186/// mirrors `docs/spec/layouts/composite.md` § Validation exactly: "a reader
187/// accumulates state only up to N, reaches one verdict at the Nth member, and is
188/// done" — every v1.0 composition rule uses a definite `member_count`.
189///
190/// [`CompositeTensor::new`] drives one of these internally; construct one
191/// directly only if you need to validate members incrementally as they arrive
192/// (e.g. a future streaming reader) rather than from an already-collected `Vec`.
193///
194/// # Decoded-type checking and quantized members
195///
196/// Per-member check 2 (member decoded type == head `type_tag`) is only fully
197/// checkable for **unquantized** members at this layer: [`TensorDescriptor`]
198/// stores quantization as raw `Option<Vec<u8>>` bytes with no typed schema at
199/// Layer 4 (see its `quantization` field docs), so a quantized member's decoded
200/// type cannot be resolved here. When a member carries a quantization section,
201/// this check is skipped; verifying that its dequantized type agrees with the
202/// head is deferred to a higher layer that has quantization-scheme context.
203///
204/// # Examples
205///
206/// ```
207/// use hurray_core::{
208///     composite::CompositeValidator,
209///     descriptor::TensorDescriptor,
210///     layout::{CompositeLayout, CompositionRule, LayoutDescriptor},
211///     BufferHandle, DeviceTag, ElementType, Shape, SyncMode, MIN_BUFFER_ALIGNMENT,
212/// };
213///
214/// // A group composite: members need no shard section and may differ arbitrarily.
215/// let head = TensorDescriptor::new(
216///     1, 0, ElementType::Float32, Shape::new(vec![1u64]).unwrap(), 0,
217///     LayoutDescriptor::Composite(CompositeLayout::new(CompositionRule::Group, 1).unwrap()),
218///     vec![],
219///     None, None, None, None,
220/// ).unwrap();
221///
222/// let buf = BufferHandle::new(4, MIN_BUFFER_ALIGNMENT, DeviceTag::Cpu, SyncMode::ProducerSynced).unwrap();
223/// let member = TensorDescriptor::new(
224///     1, 0, ElementType::Int8, Shape::new(vec![100u64]).unwrap(), 0,
225///     LayoutDescriptor::RowMajor, vec![buf],
226///     None, None, None, None,
227/// ).unwrap();
228///
229/// let mut validator = CompositeValidator::new(&head).unwrap();
230/// validator.push_member(&member).unwrap();
231/// assert!(validator.finish().is_ok());
232/// ```
233pub struct CompositeValidator {
234    rule: CompositionRule,
235    declared_member_count: u32,
236    head_shape_dims: Vec<u64>,
237    head_type_tag: u8,
238    members_seen: usize,
239    /// Accumulated `(shard_offset, shape)` boxes for the partition coverage check.
240    partition_boxes: Vec<(Vec<u64>, Vec<u64>)>,
241}
242
243impl CompositeValidator {
244    /// Creates a validator seeded from `head`.
245    ///
246    /// # Errors
247    ///
248    /// Returns [`Error::InvalidLayout`] if `head.layout` is not
249    /// [`LayoutDescriptor::Composite`], or if `head.shape` contains a `DYNAMIC`
250    /// dimension while the rule is partition or overlay — coverage/volume math
251    /// is undefined for a dynamic shape, so both rules require a fully static
252    /// head. This is the same check [`LayoutDescriptor::validate_against_shape`]
253    /// performs for a composite layout; re-checked here because that method is
254    /// opt-in and this validator is the layer that actually gates composite
255    /// assembly.
256    ///
257    /// # Examples
258    ///
259    /// See the type-level example.
260    pub fn new(head: &TensorDescriptor) -> Result<Self> {
261        let composite_layout = match &head.layout {
262            LayoutDescriptor::Composite(c) => c,
263            other => {
264                return Err(Error::InvalidLayout(format!(
265                    "composite head's layout_tag must be 0x0B (Composite), got 0x{:02X}",
266                    other.tag()
267                )))
268            }
269        };
270
271        let rule = composite_layout.rule.clone();
272        if !matches!(rule, CompositionRule::Group) && head.shape.has_dynamic() {
273            return Err(Error::InvalidLayout(
274                "partition/overlay composite head shape MUST NOT contain a DYNAMIC dimension \
275                 (coverage/volume math requires a fully static shape)"
276                    .to_string(),
277            ));
278        }
279
280        Ok(Self {
281            rule,
282            declared_member_count: composite_layout.member_count,
283            head_shape_dims: head.shape.dims().to_vec(),
284            head_type_tag: head.element_type.tag(),
285            members_seen: 0,
286            partition_boxes: Vec::new(),
287        })
288    }
289
290    /// Runs the per-member checks (spec § Validation › Per-member checks) for one
291    /// member, in the order it arrives.
292    ///
293    /// # Errors
294    ///
295    /// See the crate [`Error`] variants prefixed `Composite`, plus
296    /// [`Error::ShardOutOfBounds`] (reused from ADR-004 for the box-in-bounds check).
297    ///
298    /// # Examples
299    ///
300    /// See the type-level example.
301    pub fn push_member(&mut self, member: &TensorDescriptor) -> Result<()> {
302        let index = self.members_seen;
303        let rule = self.rule.clone();
304
305        match rule {
306            CompositionRule::Partition => {
307                let shard = self.shard_box(member, index)?;
308                self.check_member_type(member, index)?;
309                self.reject_composite_member_flag(member)?;
310                self.partition_boxes
311                    .push((shard.shard_offset.clone(), member.shape.dims().to_vec()));
312            }
313            CompositionRule::Overlay(_) => {
314                self.shard_box(member, index)?;
315                self.check_member_type(member, index)?;
316                self.check_overlay_member_role(member, index)?;
317            }
318            CompositionRule::Group => {
319                // Group: no shard section required, no decoded-type check (spec §
320                // Composition Semantics › Group: members MAY differ arbitrarily).
321                self.reject_composite_member_flag(member)?;
322            }
323        }
324
325        self.members_seen += 1;
326        Ok(())
327    }
328
329    /// Runs the close-time checks (spec § Validation › Close-time checks) once
330    /// every member has been pushed.
331    ///
332    /// # Errors
333    ///
334    /// - [`Error::CompositeMemberCountMismatch`] — the number of members pushed
335    ///   does not equal the head's declared `member_count`.
336    /// - [`Error::CompositeOverlayEmpty`] — an overlay composite received zero members.
337    /// - [`Error::CompositePartitionGap`] / [`Error::CompositePartitionOverlap`] —
338    ///   partition members do not exactly cover the head's index space.
339    ///
340    /// # Examples
341    ///
342    /// See the type-level example.
343    pub fn finish(self) -> Result<()> {
344        if self.members_seen as u32 != self.declared_member_count {
345            return Err(Error::CompositeMemberCountMismatch {
346                declared: self.declared_member_count,
347                actual: self.members_seen,
348            });
349        }
350
351        match self.rule {
352            CompositionRule::Partition => self.validate_partition_coverage(),
353            CompositionRule::Overlay(_) => {
354                // The base-span / base-first check already ran per-member (check 4)
355                // for whatever members did arrive; a zero-member overlay is a
356                // separate case — it never had a base to check in the first place.
357                if self.members_seen == 0 {
358                    return Err(Error::CompositeOverlayEmpty);
359                }
360                Ok(())
361            }
362            CompositionRule::Group => Ok(()),
363        }
364    }
365
366    // ── Per-member check helpers ─────────────────────────────────────────────
367
368    /// Check 1: `member` carries a shard section whose `parent_shape` equals the
369    /// head's shape and whose box is in bounds. Returns the shard on success so
370    /// callers needn't re-borrow `member.shard`.
371    fn shard_box<'a>(
372        &self,
373        member: &'a TensorDescriptor,
374        index: usize,
375    ) -> Result<&'a crate::descriptor::ShardDescriptor> {
376        let shard = member
377            .shard
378            .as_ref()
379            .ok_or(Error::CompositeMemberMissingShard { index })?;
380        if shard.parent_shape.as_slice() != self.head_shape_dims.as_slice() {
381            return Err(Error::CompositeMemberParentShapeMismatch { index });
382        }
383        // Box-in-bounds: shard_offset[k] + shape[k] <= parent_shape[k]. Reused
384        // verbatim from ADR-004 rather than a bespoke composite variant — same
385        // constraint, same failure semantics.
386        shard.validate_against_shape(&member.shape)?;
387        Ok(shard)
388    }
389
390    /// Check 2: `member`'s decoded value type equals the head's `type_tag`.
391    /// Skipped for quantized members — see the type-level docs' "Decoded-type
392    /// checking" section for why.
393    fn check_member_type(&self, member: &TensorDescriptor, index: usize) -> Result<()> {
394        if member.quantization.is_some() {
395            return Ok(());
396        }
397        let member_tag = member.element_type.tag();
398        if member_tag != self.head_type_tag {
399            return Err(Error::CompositeMemberTypeMismatch {
400                index,
401                member: member_tag,
402                head: self.head_type_tag,
403            });
404        }
405        Ok(())
406    }
407
408    /// Check 4 (overlay only): `member` carries a Composite Member section with a
409    /// valid role; the first member MUST be the base and span the index space;
410    /// every subsequent member MUST be a correction.
411    fn check_overlay_member_role(&self, member: &TensorDescriptor, index: usize) -> Result<()> {
412        let cm: &CompositeMemberDescriptor =
413            member
414                .composite_member
415                .as_ref()
416                .ok_or(Error::CompositeMemberFlagMismatch {
417                    rule: self.rule.rule_byte(),
418                    has_flag: false,
419                })?;
420
421        if index == 0 {
422            if cm.member_role != MemberRole::Base {
423                return Err(Error::CompositeOverlayBaseNotFirst);
424            }
425            // Presence of `shard` is already guaranteed by the `shard_box` call
426            // that precedes this one in `push_member`.
427            let shard = member
428                .shard
429                .as_ref()
430                .ok_or(Error::CompositeMemberMissingShard { index })?;
431            let spans_fully = shard.shard_offset.iter().all(|&o| o == 0)
432                && member.shape.dims() == self.head_shape_dims.as_slice();
433            if !spans_fully {
434                return Err(Error::CompositeOverlayBaseNotSpanning);
435            }
436        } else if cm.member_role != MemberRole::Correction {
437            // Only the first member may be the base; reusing the same error keeps
438            // the "base MUST be first" framing symmetric for both directions of
439            // the violation (missing at position 0, or appearing elsewhere).
440            return Err(Error::CompositeOverlayBaseNotFirst);
441        }
442
443        Ok(())
444    }
445
446    /// Check 5 (partition/group): `member` MUST NOT carry a Composite Member section.
447    fn reject_composite_member_flag(&self, member: &TensorDescriptor) -> Result<()> {
448        if member.composite_member.is_some() {
449            return Err(Error::CompositeMemberFlagMismatch {
450                rule: self.rule.rule_byte(),
451                has_flag: true,
452            });
453        }
454        Ok(())
455    }
456
457    // ── Close-time check helpers ──────────────────────────────────────────────
458
459    /// Partition close-time check: exact-cover + non-overlap over the accumulated
460    /// member boxes (spec § Composition Semantics › Partition).
461    ///
462    /// Mirrors the volume-sum + pairwise-overlap algorithm this crate previously
463    /// used for subpaving coverage validation (see git history:
464    /// `fix(hurray-core): enforce subpaving full-coverage validation`): pairwise
465    /// non-overlap is checked first (O(n²) over members; not a hot path), then —
466    /// only once no overlap was found — coverage is checked via volume sum. Given
467    /// every box is already confirmed in-bounds (check 1) and no two overlap, a
468    /// volume-sum shortfall can only mean a gap, never an overlap masquerading as
469    /// one; checking overlap first and gap second keeps the two error variants
470    /// unambiguous rather than defaulting everything to "gap".
471    fn validate_partition_coverage(&self) -> Result<()> {
472        let n = self.partition_boxes.len();
473        for i in 0..n {
474            for j in (i + 1)..n {
475                if boxes_overlap(&self.partition_boxes[i], &self.partition_boxes[j]) {
476                    return Err(Error::CompositePartitionOverlap { a: i, b: j });
477                }
478            }
479        }
480
481        // Full-coverage via volume sum. Skipped (accepted without further check) if
482        // any dimension is dynamic or a product overflows u64 — the total can't then
483        // be computed reliably, and CompositeValidator::new already rejects DYNAMIC
484        // head shapes for partition, so this branch is unreachable in practice; kept
485        // as a defensive fallback matching the subpaving precedent.
486        let total = self.head_shape_dims.iter().try_fold(1u64, |acc, &d| {
487            if d == crate::shape::DYNAMIC {
488                None
489            } else {
490                acc.checked_mul(d)
491            }
492        });
493        if let Some(total) = total {
494            let mut covered: Option<u64> = Some(0);
495            for (_, member_shape) in &self.partition_boxes {
496                let vol = member_shape.iter().try_fold(1u64, |acc, &d| {
497                    if d == crate::shape::DYNAMIC {
498                        None
499                    } else {
500                        acc.checked_mul(d)
501                    }
502                });
503                covered = match (covered, vol) {
504                    (Some(c), Some(v)) => c.checked_add(v),
505                    _ => None,
506                };
507            }
508            if let Some(covered) = covered {
509                if covered != total {
510                    return Err(Error::CompositePartitionGap);
511                }
512            }
513        }
514
515        Ok(())
516    }
517}
518
519/// Returns `true` if boxes `a` and `b` (each `(shard_offset, shape)`) overlap.
520///
521/// Per spec § Composition Semantics › Partition: `A` and `B` overlap if, for
522/// every dimension `k`, `A.offset[k] < B.offset[k] + B.shape[k]` AND
523/// `B.offset[k] < A.offset[k] + A.shape[k]`. Uses `saturating_add` so a
524/// pathological (already-rejected-elsewhere) huge offset/shape pair cannot wrap
525/// around and produce a false negative.
526fn boxes_overlap(a: &(Vec<u64>, Vec<u64>), b: &(Vec<u64>, Vec<u64>)) -> bool {
527    let (a_off, a_shape) = a;
528    let (b_off, b_shape) = b;
529    a_off
530        .iter()
531        .zip(a_shape)
532        .zip(b_off.iter().zip(b_shape))
533        .all(|((&ao, &asz), (&bo, &bsz))| {
534            ao < bo.saturating_add(bsz) && bo < ao.saturating_add(asz)
535        })
536}
537
538// ── Tests ─────────────────────────────────────────────────────────────────────
539
540#[cfg(test)]
541mod tests {
542    use super::*;
543    use crate::layout::{CombineOp, CompositeLayout};
544    use crate::{BufferHandle, DeviceTag, Error, Shape, ShardDescriptor, SyncMode};
545
546    // ── Helpers ───────────────────────────────────────────────────────────────
547
548    fn buf(byte_size: u64) -> BufferHandle {
549        BufferHandle::new(
550            byte_size,
551            crate::MIN_BUFFER_ALIGNMENT,
552            DeviceTag::Cpu,
553            SyncMode::ProducerSynced,
554        )
555        .unwrap()
556    }
557
558    fn composite_head(
559        rule: CompositionRule,
560        member_count: u32,
561        shape: Vec<u64>,
562        ty: crate::ElementType,
563    ) -> TensorDescriptor {
564        let layout = LayoutDescriptor::Composite(CompositeLayout::new(rule, member_count).unwrap());
565        TensorDescriptor::new(
566            1,
567            0,
568            ty,
569            Shape::new(shape).unwrap(),
570            0,
571            layout,
572            vec![],
573            None,
574            None,
575            None,
576            None,
577        )
578        .unwrap()
579    }
580
581    /// A plain dense (row-major) member with an optional shard section. Uses a
582    /// fixed 64-byte buffer — the byte size is irrelevant to composite
583    /// validation, which only inspects shape/type/shard/flags.
584    fn dense_member(
585        ty: crate::ElementType,
586        shape: Vec<u64>,
587        shard: Option<ShardDescriptor>,
588    ) -> TensorDescriptor {
589        TensorDescriptor::new(
590            1,
591            0,
592            ty,
593            Shape::new(shape).unwrap(),
594            0,
595            LayoutDescriptor::RowMajor,
596            vec![buf(64)],
597            None,
598            shard,
599            None,
600            None,
601        )
602        .unwrap()
603    }
604
605    fn shard(parent: Vec<u64>, offset: Vec<u64>) -> ShardDescriptor {
606        ShardDescriptor::new(parent, offset).unwrap()
607    }
608
609    // ── Partition ────────────────────────────────────────────────────────────
610
611    /// Spec worked example: [8, 8] head split into two [8, 4] members.
612    #[test]
613    fn partition_exact_tiling_accepted() {
614        let head = composite_head(
615            CompositionRule::Partition,
616            2,
617            vec![8, 8],
618            crate::ElementType::Float32,
619        );
620        let m0 = dense_member(
621            crate::ElementType::Float32,
622            vec![8, 4],
623            Some(shard(vec![8, 8], vec![0, 0])),
624        );
625        let m1 = dense_member(
626            crate::ElementType::Float32,
627            vec![8, 4],
628            Some(shard(vec![8, 8], vec![0, 4])),
629        );
630        let composite = CompositeTensor::new(head, vec![m0, m1]).unwrap();
631        assert_eq!(composite.member_count(), 2);
632        assert_eq!(*composite.rule(), CompositionRule::Partition);
633    }
634
635    /// A 4-way tiling of an [8, 8] head into four [4, 4] quadrants.
636    #[test]
637    fn partition_four_way_tiling_accepted() {
638        let head = composite_head(
639            CompositionRule::Partition,
640            4,
641            vec![8, 8],
642            crate::ElementType::Float32,
643        );
644        let members = [(0u64, 0u64), (0, 4), (4, 0), (4, 4)]
645            .into_iter()
646            .map(|(r, c)| {
647                dense_member(
648                    crate::ElementType::Float32,
649                    vec![4, 4],
650                    Some(shard(vec![8, 8], vec![r, c])),
651                )
652            })
653            .collect::<Vec<_>>();
654        let composite = CompositeTensor::new(head, members).unwrap();
655        assert_eq!(composite.member_count(), 4);
656    }
657
658    /// Members that leave a gap in the head's index space are rejected.
659    #[test]
660    fn partition_gap_rejected() {
661        let head = composite_head(
662            CompositionRule::Partition,
663            2,
664            vec![8, 8],
665            crate::ElementType::Float32,
666        );
667        // Non-overlapping boxes covering columns [0,3) and [3,6): columns
668        // [6, 8) are left uncovered — a genuine gap with no overlap (overlap is
669        // checked before coverage; see validate_partition_coverage's doc comment,
670        // so a gap-only fixture is needed to isolate this error variant).
671        let m0 = dense_member(
672            crate::ElementType::Float32,
673            vec![8, 3],
674            Some(shard(vec![8, 8], vec![0, 0])),
675        );
676        let m1 = dense_member(
677            crate::ElementType::Float32,
678            vec![8, 3],
679            Some(shard(vec![8, 8], vec![0, 3])),
680        );
681        let err = CompositeTensor::new(head, vec![m0, m1]).unwrap_err();
682        assert!(matches!(err, Error::CompositePartitionGap));
683    }
684
685    /// Two overlapping member boxes are rejected with the offending indices.
686    #[test]
687    fn partition_overlap_rejected() {
688        let head = composite_head(
689            CompositionRule::Partition,
690            2,
691            vec![8, 8],
692            crate::ElementType::Float32,
693        );
694        let m0 = dense_member(
695            crate::ElementType::Float32,
696            vec![8, 5],
697            Some(shard(vec![8, 8], vec![0, 0])),
698        );
699        let m1 = dense_member(
700            crate::ElementType::Float32,
701            vec![8, 5],
702            Some(shard(vec![8, 8], vec![0, 3])),
703        );
704        let err = CompositeTensor::new(head, vec![m0, m1]).unwrap_err();
705        assert!(matches!(
706            err,
707            Error::CompositePartitionOverlap { a: 0, b: 1 }
708        ));
709    }
710
711    /// A partition member without a shard section is rejected.
712    #[test]
713    fn partition_member_missing_shard_rejected() {
714        let head = composite_head(
715            CompositionRule::Partition,
716            1,
717            vec![8, 8],
718            crate::ElementType::Float32,
719        );
720        let m0 = dense_member(crate::ElementType::Float32, vec![8, 8], None);
721        let err = CompositeTensor::new(head, vec![m0]).unwrap_err();
722        assert!(matches!(
723            err,
724            Error::CompositeMemberMissingShard { index: 0 }
725        ));
726    }
727
728    /// A member's shard `parent_shape` must equal the head's shape.
729    #[test]
730    fn partition_member_parent_shape_mismatch_rejected() {
731        let head = composite_head(
732            CompositionRule::Partition,
733            1,
734            vec![8, 8],
735            crate::ElementType::Float32,
736        );
737        // parent_shape [10, 10] != head shape [8, 8].
738        let m0 = dense_member(
739            crate::ElementType::Float32,
740            vec![8, 8],
741            Some(shard(vec![10, 10], vec![0, 0])),
742        );
743        let err = CompositeTensor::new(head, vec![m0]).unwrap_err();
744        assert!(matches!(
745            err,
746            Error::CompositeMemberParentShapeMismatch { index: 0 }
747        ));
748    }
749
750    /// A member box that extends past the head's index space is rejected with
751    /// the reused ADR-004 `ShardOutOfBounds` variant (not a bespoke composite one).
752    #[test]
753    fn partition_member_out_of_bounds_rejected() {
754        let head = composite_head(
755            CompositionRule::Partition,
756            1,
757            vec![8, 8],
758            crate::ElementType::Float32,
759        );
760        // shard_offset[1]=4 + shape[1]=8 = 12 > parent_shape[1]=8.
761        let m0 = dense_member(
762            crate::ElementType::Float32,
763            vec![8, 8],
764            Some(shard(vec![8, 8], vec![0, 4])),
765        );
766        let err = CompositeTensor::new(head, vec![m0]).unwrap_err();
767        assert!(matches!(err, Error::ShardOutOfBounds { .. }));
768    }
769
770    /// An unquantized member whose decoded type disagrees with the head's
771    /// `type_tag` is rejected.
772    #[test]
773    fn partition_member_type_mismatch_rejected() {
774        let head = composite_head(
775            CompositionRule::Partition,
776            1,
777            vec![8, 8],
778            crate::ElementType::Float32,
779        );
780        let m0 = dense_member(
781            crate::ElementType::Int8,
782            vec![8, 8],
783            Some(shard(vec![8, 8], vec![0, 0])),
784        );
785        let err = CompositeTensor::new(head, vec![m0]).unwrap_err();
786        assert!(matches!(
787            err,
788            Error::CompositeMemberTypeMismatch { index: 0, .. }
789        ));
790    }
791
792    /// A *quantized* member's "mismatched" stored type is NOT rejected on that
793    /// basis alone — the decoded-type check is skipped when a quantization
794    /// section is present (Layer 4 has no typed quantization schema to resolve
795    /// the decoded type against; see the carve-out documented on
796    /// [`CompositeValidator::push_member`]).
797    #[test]
798    fn partition_quantized_member_type_mismatch_not_rejected() {
799        let head = composite_head(
800            CompositionRule::Partition,
801            1,
802            vec![8, 8],
803            crate::ElementType::Float16,
804        );
805        // Stored type int4 disagrees with the head's float16 — normally rejected,
806        // but quantization is Some(_), so the decoded-type check is skipped.
807        let m0 = TensorDescriptor::new(
808            1,
809            0,
810            crate::ElementType::Int4,
811            Shape::new(vec![8u64, 8]).unwrap(),
812            0,
813            LayoutDescriptor::RowMajor,
814            vec![buf(64)],
815            Some(vec![0x01, 0x00, 0x00, 0x00]),
816            Some(shard(vec![8, 8], vec![0, 0])),
817            None,
818            None,
819        )
820        .unwrap();
821        let composite = CompositeTensor::new(head, vec![m0]).unwrap();
822        assert_eq!(composite.member_count(), 1);
823    }
824
825    /// A partition member MUST NOT carry a Composite Member section.
826    #[test]
827    fn partition_member_illegal_composite_member_flag_rejected() {
828        let head = composite_head(
829            CompositionRule::Partition,
830            1,
831            vec![8, 8],
832            crate::ElementType::Float32,
833        );
834        let m0 = dense_member(
835            crate::ElementType::Float32,
836            vec![8, 8],
837            Some(shard(vec![8, 8], vec![0, 0])),
838        )
839        .with_composite_member(CompositeMemberDescriptor::new(MemberRole::Base));
840        let err = CompositeTensor::new(head, vec![m0]).unwrap_err();
841        assert!(matches!(
842            err,
843            Error::CompositeMemberFlagMismatch { has_flag: true, .. }
844        ));
845    }
846
847    /// Declared `member_count` not matching the number of members actually
848    /// pushed is rejected at close time.
849    #[test]
850    fn partition_member_count_mismatch_rejected() {
851        let head = composite_head(
852            CompositionRule::Partition,
853            2,
854            vec![8, 8],
855            crate::ElementType::Float32,
856        );
857        let m0 = dense_member(
858            crate::ElementType::Float32,
859            vec![8, 8],
860            Some(shard(vec![8, 8], vec![0, 0])),
861        );
862        let err = CompositeTensor::new(head, vec![m0]).unwrap_err();
863        assert!(matches!(
864            err,
865            Error::CompositeMemberCountMismatch {
866                declared: 2,
867                actual: 1,
868            }
869        ));
870    }
871
872    /// A DYNAMIC-shaped head is rejected end-to-end through the public entry
873    /// point ([`CompositeTensor::new`] / [`CompositeValidator::new`]) — this is
874    /// the layer that actually gates composite assembly; see
875    /// `LayoutDescriptor::validate_against_shape` for the opt-in Layer-3-only
876    /// check of the same constraint.
877    #[test]
878    fn partition_dynamic_head_shape_rejected() {
879        let layout = LayoutDescriptor::Composite(
880            CompositeLayout::new(CompositionRule::Partition, 0).unwrap(),
881        );
882        let head = TensorDescriptor::new(
883            1,
884            0,
885            crate::ElementType::Float32,
886            Shape::new(vec![crate::shape::DYNAMIC, 8]).unwrap(),
887            0,
888            layout,
889            vec![],
890            None,
891            None,
892            None,
893            None,
894        )
895        .unwrap();
896        let err = CompositeTensor::new(head, vec![]).unwrap_err();
897        assert!(matches!(err, Error::InvalidLayout(_)));
898    }
899
900    // ── Overlay ──────────────────────────────────────────────────────────────
901
902    fn base_member(ty: crate::ElementType, head_shape: Vec<u64>) -> TensorDescriptor {
903        let offset = vec![0u64; head_shape.len()];
904        dense_member(ty, head_shape.clone(), Some(shard(head_shape, offset)))
905            .with_composite_member(CompositeMemberDescriptor::new(MemberRole::Base))
906    }
907
908    fn correction_member(
909        ty: crate::ElementType,
910        head_shape: Vec<u64>,
911        shape: Vec<u64>,
912        offset: Vec<u64>,
913    ) -> TensorDescriptor {
914        dense_member(ty, shape, Some(shard(head_shape, offset)))
915            .with_composite_member(CompositeMemberDescriptor::new(MemberRole::Correction))
916    }
917
918    #[test]
919    fn overlay_base_and_correction_replace_accepted() {
920        let head = composite_head(
921            CompositionRule::Overlay(CombineOp::Replace),
922            2,
923            vec![4, 4],
924            crate::ElementType::Float32,
925        );
926        let base = base_member(crate::ElementType::Float32, vec![4, 4]);
927        let correction = correction_member(
928            crate::ElementType::Float32,
929            vec![4, 4],
930            vec![2, 2],
931            vec![0, 0],
932        );
933        let composite = CompositeTensor::new(head, vec![base, correction]).unwrap();
934        assert_eq!(composite.member_count(), 2);
935        assert_eq!(
936            *composite.rule(),
937            CompositionRule::Overlay(CombineOp::Replace)
938        );
939    }
940
941    #[test]
942    fn overlay_base_and_correction_add_accepted() {
943        let head = composite_head(
944            CompositionRule::Overlay(CombineOp::Add),
945            2,
946            vec![4, 4],
947            crate::ElementType::Float32,
948        );
949        let base = base_member(crate::ElementType::Float32, vec![4, 4]);
950        let correction = correction_member(
951            crate::ElementType::Float32,
952            vec![4, 4],
953            vec![2, 2],
954            vec![1, 1],
955        );
956        let composite = CompositeTensor::new(head, vec![base, correction]).unwrap();
957        assert_eq!(*composite.rule(), CompositionRule::Overlay(CombineOp::Add));
958    }
959
960    /// The first member of an overlay MUST be the base.
961    #[test]
962    fn overlay_base_not_first_rejected() {
963        let head = composite_head(
964            CompositionRule::Overlay(CombineOp::Replace),
965            1,
966            vec![4, 4],
967            crate::ElementType::Float32,
968        );
969        let correction = correction_member(
970            crate::ElementType::Float32,
971            vec![4, 4],
972            vec![4, 4],
973            vec![0, 0],
974        );
975        let err = CompositeTensor::new(head, vec![correction]).unwrap_err();
976        assert!(matches!(err, Error::CompositeOverlayBaseNotFirst));
977    }
978
979    /// A base member appearing at a position other than first is also rejected
980    /// (same error, symmetric framing).
981    #[test]
982    fn overlay_base_appears_second_rejected() {
983        let head = composite_head(
984            CompositionRule::Overlay(CombineOp::Replace),
985            2,
986            vec![4, 4],
987            crate::ElementType::Float32,
988        );
989        let correction = correction_member(
990            crate::ElementType::Float32,
991            vec![4, 4],
992            vec![4, 4],
993            vec![0, 0],
994        );
995        let base = base_member(crate::ElementType::Float32, vec![4, 4]);
996        let err = CompositeTensor::new(head, vec![correction, base]).unwrap_err();
997        assert!(matches!(err, Error::CompositeOverlayBaseNotFirst));
998    }
999
1000    /// The base member MUST span the whole index space.
1001    #[test]
1002    fn overlay_base_not_spanning_rejected() {
1003        let head = composite_head(
1004            CompositionRule::Overlay(CombineOp::Replace),
1005            1,
1006            vec![4, 4],
1007            crate::ElementType::Float32,
1008        );
1009        // Base shape [2, 2] does not equal the head's [4, 4] shape.
1010        let base = dense_member(
1011            crate::ElementType::Float32,
1012            vec![2, 2],
1013            Some(shard(vec![4, 4], vec![0, 0])),
1014        )
1015        .with_composite_member(CompositeMemberDescriptor::new(MemberRole::Base));
1016        let err = CompositeTensor::new(head, vec![base]).unwrap_err();
1017        assert!(matches!(err, Error::CompositeOverlayBaseNotSpanning));
1018    }
1019
1020    /// An overlay head with `member_count = 0` has no base member available.
1021    #[test]
1022    fn overlay_zero_members_rejected() {
1023        let head = composite_head(
1024            CompositionRule::Overlay(CombineOp::Replace),
1025            0,
1026            vec![4, 4],
1027            crate::ElementType::Float32,
1028        );
1029        let err = CompositeTensor::new(head, vec![]).unwrap_err();
1030        assert!(matches!(err, Error::CompositeOverlayEmpty));
1031    }
1032
1033    /// Corrections MAY legally overlap each other (unlike partition members).
1034    #[test]
1035    fn overlay_corrections_may_overlap() {
1036        let head = composite_head(
1037            CompositionRule::Overlay(CombineOp::Replace),
1038            3,
1039            vec![4, 4],
1040            crate::ElementType::Float32,
1041        );
1042        let base = base_member(crate::ElementType::Float32, vec![4, 4]);
1043        // Both corrections cover the same [0,0]..[2,2] box — legal for overlay.
1044        let c0 = correction_member(
1045            crate::ElementType::Float32,
1046            vec![4, 4],
1047            vec![2, 2],
1048            vec![0, 0],
1049        );
1050        let c1 = correction_member(
1051            crate::ElementType::Float32,
1052            vec![4, 4],
1053            vec![2, 2],
1054            vec![0, 0],
1055        );
1056        let composite = CompositeTensor::new(head, vec![base, c0, c1]).unwrap();
1057        assert_eq!(composite.member_count(), 3);
1058    }
1059
1060    /// A correction missing the Composite Member section is rejected.
1061    #[test]
1062    fn overlay_correction_missing_composite_member_flag_rejected() {
1063        let head = composite_head(
1064            CompositionRule::Overlay(CombineOp::Replace),
1065            2,
1066            vec![4, 4],
1067            crate::ElementType::Float32,
1068        );
1069        let base = base_member(crate::ElementType::Float32, vec![4, 4]);
1070        // Missing .with_composite_member(...) — HAS_COMPOSITE_MEMBER absent.
1071        let bad_correction = dense_member(
1072            crate::ElementType::Float32,
1073            vec![2, 2],
1074            Some(shard(vec![4, 4], vec![0, 0])),
1075        );
1076        let err = CompositeTensor::new(head, vec![base, bad_correction]).unwrap_err();
1077        assert!(matches!(
1078            err,
1079            Error::CompositeMemberFlagMismatch {
1080                has_flag: false,
1081                ..
1082            }
1083        ));
1084    }
1085
1086    // ── Group ────────────────────────────────────────────────────────────────
1087
1088    /// Group members MAY differ arbitrarily: different shapes, types, and no
1089    /// shard sections at all — no coverage or type check applies.
1090    #[test]
1091    fn group_heterogeneous_members_accepted() {
1092        let head = composite_head(
1093            CompositionRule::Group,
1094            2,
1095            vec![1],
1096            crate::ElementType::Float32,
1097        );
1098        let m0 = dense_member(crate::ElementType::Int8, vec![100], None);
1099        let m1 = dense_member(crate::ElementType::Float64, vec![3, 3, 3], None);
1100        let composite = CompositeTensor::new(head, vec![m0, m1]).unwrap();
1101        assert_eq!(composite.member_count(), 2);
1102        assert_eq!(*composite.rule(), CompositionRule::Group);
1103    }
1104
1105    /// A group member count mismatch is still rejected (the only failure mode
1106    /// group composites share with the other rules, since there is no
1107    /// coverage/type check to violate).
1108    #[test]
1109    fn group_member_count_mismatch_rejected() {
1110        let head = composite_head(
1111            CompositionRule::Group,
1112            2,
1113            vec![1],
1114            crate::ElementType::Float32,
1115        );
1116        let m0 = dense_member(crate::ElementType::Int8, vec![100], None);
1117        let err = CompositeTensor::new(head, vec![m0]).unwrap_err();
1118        assert!(matches!(
1119            err,
1120            Error::CompositeMemberCountMismatch {
1121                declared: 2,
1122                actual: 1,
1123            }
1124        ));
1125    }
1126
1127    // ── Nesting / depth ──────────────────────────────────────────────────────
1128
1129    #[test]
1130    fn max_composite_depth_constant_is_8() {
1131        assert_eq!(MAX_COMPOSITE_DEPTH, 8);
1132    }
1133
1134    /// `new_at_depth` at depth 7 (the 8th level, 0-indexed) with no nested
1135    /// composite member is accepted.
1136    #[test]
1137    fn new_at_depth_7_is_accepted() {
1138        let head = composite_head(
1139            CompositionRule::Group,
1140            1,
1141            vec![1],
1142            crate::ElementType::Float32,
1143        );
1144        let m0 = dense_member(crate::ElementType::Int8, vec![4], None);
1145        let result = CompositeTensor::new_at_depth(head, vec![m0], 7);
1146        assert!(result.is_ok(), "depth 7 should be accepted");
1147    }
1148
1149    /// `new_at_depth` at depth 8 (the 9th level) is rejected immediately,
1150    /// regardless of the members supplied.
1151    #[test]
1152    fn new_at_depth_8_is_rejected() {
1153        let head = composite_head(
1154            CompositionRule::Group,
1155            0,
1156            vec![1],
1157            crate::ElementType::Float32,
1158        );
1159        let err = CompositeTensor::new_at_depth(head, vec![], 8).unwrap_err();
1160        assert!(matches!(err, Error::SubpavingNestingTooDeep));
1161    }
1162
1163    /// A member whose own `layout_tag` is the composite tag counts as one
1164    /// additional nesting level: at depth 6, a nested composite member pushes
1165    /// the effective depth to 7 (still within the cap) and is accepted.
1166    #[test]
1167    fn new_at_depth_6_with_nested_composite_member_accepted() {
1168        let head = composite_head(
1169            CompositionRule::Group,
1170            1,
1171            vec![1],
1172            crate::ElementType::Float32,
1173        );
1174        let nested_head = composite_head(
1175            CompositionRule::Group,
1176            0,
1177            vec![1],
1178            crate::ElementType::Float32,
1179        );
1180        let result = CompositeTensor::new_at_depth(head, vec![nested_head], 6);
1181        assert!(
1182            result.is_ok(),
1183            "a nested composite member at depth 6 (effective depth 7) should be accepted"
1184        );
1185    }
1186
1187    /// The same nested-composite-member case at depth 7 pushes the effective
1188    /// depth to 8, exceeding the cap.
1189    #[test]
1190    fn new_at_depth_7_with_nested_composite_member_rejected() {
1191        let head = composite_head(
1192            CompositionRule::Group,
1193            1,
1194            vec![1],
1195            crate::ElementType::Float32,
1196        );
1197        let nested_head = composite_head(
1198            CompositionRule::Group,
1199            0,
1200            vec![1],
1201            crate::ElementType::Float32,
1202        );
1203        let err = CompositeTensor::new_at_depth(head, vec![nested_head], 7).unwrap_err();
1204        assert!(matches!(err, Error::SubpavingNestingTooDeep));
1205    }
1206
1207    // ── CompositeValidator direct use ───────────────────────────────────────
1208
1209    #[test]
1210    fn validator_rejects_non_composite_head_layout() {
1211        // CompositeValidator does not derive Debug (it holds no Debug-friendly
1212        // state worth printing), so match the Result directly rather than
1213        // unwrap_err() (which requires Ok's type: Debug).
1214        let head = dense_member(crate::ElementType::Float32, vec![4], None);
1215        let result = CompositeValidator::new(&head);
1216        assert!(matches!(result, Err(Error::InvalidLayout(_))));
1217    }
1218}