Skip to main content

dbsp/
typed_batch.rs

1//! Strongly typed wrappers around dynamically typed batch types.
2//!
3//! These wrappers are used to implement type-safe wrappers around DBSP
4//! operators.
5
6use crate::{
7    Circuit, Error,
8    circuit::checkpointer::Checkpoint,
9    dynamic::{DataTrait, DynData, DynUnit, Erase, LeanVec, WeightTrait},
10    trace::{BatchReaderFactories, DbspSerializer, Deserializer, spine_async::WithSnapshot},
11};
12pub use crate::{
13    DBData, DBWeight, DynZWeight, Stream, Timestamp, ZWeight,
14    algebra::{
15        IndexedZSet as DynIndexedZSet, IndexedZSetReader as DynIndexedZSetReader,
16        OrdIndexedZSet as DynOrdIndexedZSet, OrdZSet as DynOrdZSet,
17        VecIndexedZSet as DynVecIndexedZSet, VecZSet as DynVecZSet, ZSet as DynZSet,
18        ZSetReader as DynZSetReader,
19    },
20    trace::{
21        Batch as DynBatch, BatchReader as DynBatchReader,
22        BatchReaderWithSnapshot as DynBatchReaderWithSnapshot,
23        FallbackIndexedWSet as DynFallbackIndexedWSet, FallbackKeyBatch as DynFallbackKeyBatch,
24        FallbackValBatch as DynFallbackValBatch, FallbackWSet as DynFallbackWSet,
25        FileIndexedWSet as DynFileIndexedWSet, FileKeyBatch as DynFileKeyBatch,
26        FileValBatch as DynFileValBatch, FileWSet as DynFileWSet,
27        OrdIndexedWSet as DynOrdIndexedWSet, OrdKeyBatch as DynOrdKeyBatch,
28        OrdValBatch as DynOrdValBatch, OrdWSet as DynOrdWSet, Spine as DynSpine,
29        SpineSnapshot as DynSpineSnapshot, Trace as DynTrace, VecIndexedWSet as DynVecIndexedWSet,
30        VecKeyBatch as DynVecKeyBatch, VecValBatch as DynVecValBatch, VecWSet as DynVecWSet,
31        merge_batches as dyn_merge_batches, merge_batches_by_reference,
32    },
33};
34use dyn_clone::clone_box;
35use rkyv::{Archive, Archived, Deserialize, Fallible, Serialize};
36use size_of::SizeOf;
37use std::{
38    fmt::Debug,
39    marker::PhantomData,
40    ops::{Deref, DerefMut, Neg},
41    sync::Arc,
42};
43
44use crate::{
45    NumEntries,
46    algebra::{AddAssignByRef, AddByRef, HasZero, NegByRef},
47    dynamic::DowncastTrait,
48    utils::Tup2,
49};
50
51/// A strongly typed wrapper around [`DynBatchReader`].
52pub trait BatchReader: 'static {
53    type Inner: DynBatchReader<Time = Self::Time, Key = Self::DynK, Val = Self::DynV, R = Self::DynR>;
54
55    /// Any batch reader can be decomposed into a list of batches.  This type represents the
56    /// type of the batches:
57    ///
58    /// * If the reader wraps an individual batch, then `IntoBatch` is the same as `Inner`.
59    /// * If the reader wraps a spine or a spine snapshot, then `IntoBatch` is the type of the batch
60    ///   in the spine.
61    type IntoBatch: DynBatch<Time = Self::Time, Key = Self::DynK, Val = Self::DynV, R = Self::DynR>;
62
63    /// Concrete key type.
64    type Key: DBData + Erase<Self::DynK>;
65
66    /// Concrete value type.
67    type Val: DBData + Erase<Self::DynV>;
68
69    /// Concrete weight typ.
70    type R: DBWeight + Erase<Self::DynR>;
71
72    /// Dynamic key type (e.g., [`DynData`]).
73    type DynK: DataTrait + ?Sized;
74
75    /// Dynamic value typ, (e.g., [`DynData`]).
76    type DynV: DataTrait + ?Sized;
77
78    /// Dynamic weight type (e.g., [`DynZWeight`]).
79    type DynR: WeightTrait + ?Sized;
80
81    type Time: Timestamp;
82
83    /// Factories for `Self::Inner`.
84    fn factories() -> <Self::Inner as DynBatchReader>::Factories {
85        BatchReaderFactories::new::<Self::Key, Self::Val, Self::R>()
86    }
87
88    fn into_batch_factories() -> <Self::IntoBatch as DynBatchReader>::Factories {
89        BatchReaderFactories::new::<Self::Key, Self::Val, Self::R>()
90    }
91
92    /// Extract the dynamically typed batch.
93    fn inner(&self) -> &Self::Inner;
94
95    fn inner_mut(&mut self) -> &mut Self::Inner;
96
97    /// Drop the statically typed wrapper and return the inner dynamic batch type.
98    fn into_inner(self) -> Self::Inner;
99
100    /// Create a statically typed wrapper around `inner`.
101    fn from_inner(inner: Self::Inner) -> Self;
102
103    /// Convert a stream of batches of statically typed batches of type `Self` into a stream
104    /// of dynamically typed batches `Self::Inner` batches.
105    ///
106    /// This operation has not runtime cost.
107    fn stream_inner<C: Clone>(stream: &Stream<C, Self>) -> Stream<C, Self::Inner>
108    where
109        Self: Sized;
110
111    /// Convert a stream of dynamically typed batches of type `Self::Inner` into a stream
112    /// of statically typed batches of type `Self`.
113    fn stream_from_inner<C: Clone>(stream: &Stream<C, Self::Inner>) -> Stream<C, Self>
114    where
115        Self: Sized;
116
117    /// Consume `self` and returns the list of dynamically typed batches comprising it.
118    fn into_dyn_batches(self) -> Vec<Arc<Self::IntoBatch>>;
119
120    /// Consume `self` and returns the list of statically typed batches comprising it.
121    fn into_batches(self) -> Vec<Arc<TypedBatch<Self::Key, Self::Val, Self::R, Self::IntoBatch>>>;
122
123    /// Returns the list of dynamically typed batches comprising `self`.
124    fn dyn_batches(&self) -> Vec<Arc<Self::IntoBatch>>;
125
126    /// Returns the list of statically typed batches comprising `self`.
127    fn batches(&self) -> Vec<Arc<TypedBatch<Self::Key, Self::Val, Self::R, Self::IntoBatch>>>;
128
129    /// Assemble batches from `self` into a spine snapshot.
130    fn dyn_snapshot(&self) -> DynSpineSnapshot<Self::IntoBatch>;
131
132    /// Convert `self` into a spine snapshot.
133    fn into_dyn_snapshot(self) -> DynSpineSnapshot<Self::IntoBatch>;
134}
135
136/// A statically typed wrapper around [`DynBatch`].
137pub trait Batch: BatchReader<Inner = Self::InnerBatch> + Clone {
138    type InnerBatch: DynBatch<Time = Self::Time, Key = Self::DynK, Val = Self::DynV, R = Self::DynR>;
139
140    fn filter<F>(&self, predicate: F) -> Self
141    where
142        F: Fn(&Self::Key, &Self::Val) -> bool,
143        Self::Time: PartialEq<()> + From<()>,
144    {
145        Self::from_inner(
146            self.inner()
147                .filter(&|k, v| unsafe { predicate(k.downcast(), v.downcast()) }),
148        )
149    }
150}
151
152impl<B> Batch for B
153where
154    B: BatchReader + Clone,
155    B::Inner: DynBatch,
156{
157    type InnerBatch = B::Inner;
158}
159
160/// A statically typed wrapper around [`DynTrace`].
161pub trait Trace: BatchReader<Inner = Self::InnerTrace> {
162    type InnerTrace: DynTrace<Time = Self::Time, Key = Self::DynK, Val = Self::DynV, R = Self::DynR>;
163}
164
165impl<T> Trace for T
166where
167    T: BatchReader,
168    T::Inner: DynTrace<Time = T::Time, Key = T::DynK, Val = T::DynV, R = T::DynR>,
169{
170    type InnerTrace = <T as BatchReader>::Inner;
171}
172
173/// A statically typed wrapper around [`DynIndexedZSetReader`].
174pub trait IndexedZSetReader: BatchReader<R = ZWeight, DynR = DynZWeight, Time = ()> {
175    fn iter(&self) -> impl Iterator<Item = (Self::Key, Self::Val, ZWeight)> + '_ {
176        self.inner().iter().map(|(boxk, boxv, w)| unsafe {
177            (
178                boxk.as_ref().downcast::<Self::Key>().clone(),
179                boxv.as_ref().downcast::<Self::Val>().clone(),
180                w,
181            )
182        })
183    }
184}
185
186impl<Z> IndexedZSetReader for Z where Z: BatchReader<R = ZWeight, DynR = DynZWeight, Time = ()> {}
187
188/// A statically typed wrapper around [`DynIndexedZSet`].
189pub trait IndexedZSet:
190    Batch<R = ZWeight, DynR = DynZWeight, Time = (), InnerBatch = Self::InnerIndexedZSet>
191{
192    type InnerIndexedZSet: DynIndexedZSet<Time = Self::Time, Key = Self::DynK, Val = Self::DynV, R = Self::DynR>;
193}
194
195impl<Z> IndexedZSet for Z
196where
197    Z: Batch<R = ZWeight, DynR = DynZWeight, Time = ()>,
198    Z::InnerBatch: DynIndexedZSet,
199{
200    type InnerIndexedZSet = Z::InnerBatch;
201}
202
203/// A statically typed wrapper around [`DynZSetReader`].
204pub trait ZSetReader: IndexedZSetReader<Val = ()> {}
205
206impl<Z> ZSetReader for Z where Z: IndexedZSetReader<Val = ()> {}
207
208/// A statically typed wrapper around [`DynZSet`].
209pub trait ZSet: IndexedZSet<Val = (), DynV = DynUnit> {
210    fn weighted_count(&self) -> ZWeight;
211}
212
213impl<Z> ZSet for Z
214where
215    Z: IndexedZSet<Val = (), DynV = DynUnit>,
216{
217    fn weighted_count(&self) -> ZWeight {
218        let mut w = 0;
219        self.inner().weighted_count(w.erase_mut());
220        w
221    }
222}
223
224/// A statically typed wrapper around a dynamically typed batch `B` with concrete key, value, and
225/// weight types `K`, `V`, and `R` respectively.
226#[derive(Clone, Eq, SizeOf)]
227// repr(transparent) guarantees that we can safely transmute this to `inner`.
228#[repr(transparent)]
229pub struct TypedBatch<K, V, R, B> {
230    inner: B,
231    phantom: PhantomData<fn(&K, &V, &R)>,
232}
233
234impl<K, V, R, B> Debug for TypedBatch<K, V, R, B>
235where
236    B: Debug,
237{
238    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
239        self.inner.fmt(f)
240    }
241}
242
243impl<K, V, R, B, B2> PartialEq<TypedBatch<K, V, R, B2>> for TypedBatch<K, V, R, B>
244where
245    B: PartialEq<B2>,
246{
247    fn eq(&self, other: &TypedBatch<K, V, R, B2>) -> bool {
248        self.inner.eq(&other.inner)
249    }
250}
251
252impl<K, V, R, B> Default for TypedBatch<K, V, R, B>
253where
254    B: DynBatch,
255    K: DBData + Erase<B::Key>,
256    V: DBData + Erase<B::Val>,
257    R: DBWeight + Erase<B::R>,
258{
259    fn default() -> Self {
260        Self::new(B::dyn_empty(&BatchReaderFactories::new::<K, V, R>()))
261    }
262}
263
264impl<K, V, R, B> Deref for TypedBatch<K, V, R, B> {
265    type Target = B;
266
267    fn deref(&self) -> &Self::Target {
268        &self.inner
269    }
270}
271
272impl<K, V, R, B> HasZero for TypedBatch<K, V, R, B>
273where
274    B: DynBatch<Time = ()>,
275    K: DBData + Erase<B::Key>,
276    V: DBData + Erase<B::Val>,
277    R: DBWeight + Erase<B::R>,
278{
279    fn zero() -> Self {
280        Self::new(B::dyn_empty(&BatchReaderFactories::new::<K, V, R>()))
281    }
282    fn is_zero(&self) -> bool {
283        self.is_empty()
284    }
285}
286
287impl<K, V, R, B> Neg for TypedBatch<K, V, R, B>
288where
289    B: DynBatchReader + Neg<Output = B>,
290    K: DBData + Erase<B::Key>,
291    V: DBData + Erase<B::Val>,
292    R: DBWeight + Erase<B::R>,
293{
294    type Output = Self;
295
296    fn neg(self) -> Self::Output {
297        Self::new(self.inner.neg())
298    }
299}
300
301impl<K, V, R, B> NegByRef for TypedBatch<K, V, R, B>
302where
303    B: DynBatchReaderWithSnapshot + NegByRef,
304    K: DBData + Erase<B::Key>,
305    V: DBData + Erase<B::Val>,
306    R: DBWeight + Erase<B::R>,
307{
308    fn neg_by_ref(&self) -> Self {
309        Self::new(self.inner().neg_by_ref())
310    }
311}
312
313impl<K, V, R, B> AddByRef for TypedBatch<K, V, R, B>
314where
315    B: DynBatchReaderWithSnapshot + AddByRef,
316    K: DBData + Erase<B::Key>,
317    V: DBData + Erase<B::Val>,
318    R: DBWeight + Erase<B::R>,
319{
320    fn add_by_ref(&self, other: &Self) -> Self {
321        Self::new(self.inner.add_by_ref(other.inner()))
322    }
323}
324
325impl<K, V, R, B> AddAssignByRef for TypedBatch<K, V, R, B>
326where
327    B: DynBatchReaderWithSnapshot + AddAssignByRef,
328    K: DBData + Erase<B::Key>,
329    V: DBData + Erase<B::Val>,
330    R: DBWeight + Erase<B::R>,
331{
332    fn add_assign_by_ref(&mut self, other: &Self) {
333        self.inner.add_assign_by_ref(other.inner())
334    }
335}
336
337impl<K, V, R, B> NumEntries for TypedBatch<K, V, R, B>
338where
339    B: DynBatchReader,
340    K: DBData + Erase<B::Key>,
341    V: DBData + Erase<B::Val>,
342    R: DBWeight + Erase<B::R>,
343{
344    const CONST_NUM_ENTRIES: Option<usize> = B::CONST_NUM_ENTRIES;
345
346    fn num_entries_shallow(&self) -> usize {
347        self.inner.num_entries_shallow()
348    }
349
350    fn num_entries_deep(&self) -> usize {
351        self.inner.num_entries_deep()
352    }
353}
354
355impl<K, V, R, B> DerefMut for TypedBatch<K, V, R, B> {
356    fn deref_mut(&mut self) -> &mut Self::Target {
357        &mut self.inner
358    }
359}
360
361impl<K, V, R, B> TypedBatch<K, V, R, B> {
362    pub fn new(inner: B) -> Self {
363        Self {
364            inner,
365            phantom: PhantomData,
366        }
367    }
368}
369
370impl<K, V, R, B> TypedBatch<K, V, R, B>
371where
372    B: DynBatch,
373    K: DBData + Erase<B::Key>,
374    V: DBData + Erase<B::Val>,
375    R: DBWeight + Erase<B::R>,
376{
377    /// Create an empty batch.
378    pub fn empty() -> Self {
379        Self::new(B::dyn_empty(&BatchReaderFactories::new::<K, V, R>()))
380    }
381
382    /// Build a batch out of `tuples`.
383    pub fn from_tuples(time: B::Time, tuples: Vec<Tup2<Tup2<K, V>, R>>) -> Self {
384        Self::new(B::dyn_from_tuples(
385            &BatchReaderFactories::new::<K, V, R>(),
386            time,
387            &mut Box::new(LeanVec::from(tuples)).erase_box(),
388        ))
389    }
390
391    /// Merge `self` with `other`.
392    pub fn merge(&self, other: &Self) -> Self {
393        Self::new(merge_batches_by_reference(
394            &self.inner.factories(),
395            [&self.inner, &other.inner],
396            &None,
397            &None,
398        ))
399    }
400
401    pub fn merge_batches<I>(batches: I) -> Self
402    where
403        I: IntoIterator<Item = Self>,
404    {
405        Self::new(dyn_merge_batches(
406            &Self::factories(),
407            batches.into_iter().map(|b| b.into_inner()),
408            &None,
409            &None,
410        ))
411    }
412}
413
414impl<K, V, R, B> TypedBatch<K, V, R, B> {
415    pub fn into_inner(self) -> B {
416        self.inner
417    }
418}
419
420impl<K, R, B> TypedBatch<K, (), R, B>
421where
422    B: DynBatch,
423    K: DBData + Erase<B::Key>,
424    (): Erase<B::Val>,
425    R: DBWeight + Erase<B::R>,
426{
427    pub fn from_keys(time: B::Time, tuples: Vec<Tup2<K, R>>) -> Self {
428        Self::from_tuples(
429            time,
430            tuples
431                .into_iter()
432                .map(|Tup2(k, r)| Tup2(Tup2(k, ()), r))
433                .collect::<Vec<_>>(),
434        )
435    }
436}
437
438impl<K, V, R, B> BatchReader for TypedBatch<K, V, R, B>
439where
440    B: DynBatchReaderWithSnapshot,
441    K: DBData + Erase<B::Key>,
442    V: DBData + Erase<B::Val>,
443    R: DBWeight + Erase<B::R>,
444{
445    type Inner = B;
446    type IntoBatch = B::Batch;
447
448    type Key = K;
449    type Val = V;
450    type R = R;
451    type Time = B::Time;
452    type DynK = B::Key;
453    type DynV = B::Val;
454    type DynR = B::R;
455
456    fn inner(&self) -> &Self::Inner {
457        &self.inner
458    }
459
460    fn inner_mut(&mut self) -> &mut Self::Inner {
461        &mut self.inner
462    }
463
464    fn into_inner(self) -> Self::Inner {
465        self.inner
466    }
467
468    fn from_inner(inner: Self::Inner) -> Self {
469        Self {
470            inner,
471            phantom: PhantomData,
472        }
473    }
474
475    fn stream_inner<C: Clone>(stream: &Stream<C, Self>) -> Stream<C, B> {
476        // Safety: repr(transparent) on TypedBatch guarantees that this is
477        // safe.
478        unsafe { stream.transmute_payload() }
479    }
480
481    fn stream_from_inner<C: Clone>(stream: &Stream<C, Self::Inner>) -> Stream<C, Self> {
482        // Safety: repr(transparent) on TypedBatch guarantees that this is
483        // safe.
484        unsafe { stream.transmute_payload() }
485    }
486
487    fn into_dyn_batches(self) -> Vec<Arc<Self::IntoBatch>> {
488        self.inner.into_ro_snapshot().into_batches()
489    }
490
491    fn into_dyn_snapshot(self) -> DynSpineSnapshot<Self::IntoBatch> {
492        self.inner.into_ro_snapshot()
493    }
494
495    fn into_batches(self) -> Vec<Arc<TypedBatch<Self::Key, Self::Val, Self::R, Self::IntoBatch>>> {
496        unsafe {
497            std::mem::transmute::<
498                Vec<Arc<Self::IntoBatch>>,
499                Vec<Arc<TypedBatch<Self::Key, Self::Val, Self::R, Self::IntoBatch>>>,
500            >(self.into_dyn_batches())
501        }
502    }
503
504    fn dyn_batches(&self) -> Vec<Arc<Self::IntoBatch>> {
505        self.inner.ro_snapshot().into_batches()
506    }
507
508    fn batches(&self) -> Vec<Arc<TypedBatch<Self::Key, Self::Val, Self::R, Self::IntoBatch>>> {
509        unsafe {
510            std::mem::transmute::<
511                Vec<Arc<Self::IntoBatch>>,
512                Vec<Arc<TypedBatch<Self::Key, Self::Val, Self::R, Self::IntoBatch>>>,
513            >(self.dyn_batches())
514        }
515    }
516
517    fn dyn_snapshot(&self) -> DynSpineSnapshot<Self::IntoBatch> {
518        self.inner.ro_snapshot()
519    }
520}
521
522impl<K, V, R, B> Checkpoint for TypedBatch<K, V, R, B>
523where
524    B: Checkpoint,
525    K: DBData,
526    V: DBData,
527    R: DBWeight,
528{
529    fn checkpoint(&self) -> Result<Vec<u8>, Error> {
530        self.inner.checkpoint()
531    }
532
533    fn restore(&mut self, data: &[u8]) -> Result<(), Error> {
534        self.inner.restore(data)
535    }
536}
537
538pub type OrdWSet<K, R, DynR> = TypedBatch<K, (), R, DynOrdWSet<DynData, DynR>>;
539pub type OrdZSet<K> = TypedBatch<K, (), ZWeight, DynOrdZSet<DynData>>;
540pub type OrdIndexedWSet<K, V, R, DynR> =
541    TypedBatch<K, V, R, DynOrdIndexedWSet<DynData, DynData, DynR>>;
542pub type OrdIndexedZSet<K, V> = TypedBatch<K, V, ZWeight, DynOrdIndexedZSet<DynData, DynData>>;
543pub type OrdKeyBatch<K, T, R, DynR> = TypedBatch<K, (), R, DynOrdKeyBatch<DynData, T, DynR>>;
544pub type OrdValBatch<K, V, T, R, DynR> =
545    TypedBatch<K, V, R, DynOrdValBatch<DynData, DynData, T, DynR>>;
546
547pub type VecWSet<K, R, DynR> = TypedBatch<K, (), R, DynVecWSet<DynData, DynR>>;
548pub type VecZSet<K> = TypedBatch<K, (), ZWeight, DynVecZSet<DynData>>;
549pub type VecIndexedWSet<K, V, R, DynR> =
550    TypedBatch<K, V, R, DynVecIndexedWSet<DynData, DynData, DynR>>;
551pub type VecIndexedZSet<K, V> = TypedBatch<K, V, ZWeight, DynVecIndexedZSet<DynData, DynData>>;
552pub type VecKeyBatch<K, T, R, DynR> = TypedBatch<K, (), R, DynVecKeyBatch<DynData, T, DynR>>;
553pub type VecValBatch<K, V, T, R, DynR> =
554    TypedBatch<K, V, R, DynVecValBatch<DynData, DynData, T, DynR>>;
555
556pub type FileWSet<K, R, DynR> = TypedBatch<K, (), R, DynFileWSet<DynData, DynR>>;
557pub type FileZSet<K> = TypedBatch<K, (), ZWeight, DynFileWSet<DynData, DynZWeight>>;
558pub type FileIndexedWSet<K, V, R, DynR> =
559    TypedBatch<K, V, R, DynFileIndexedWSet<DynData, DynData, DynR>>;
560pub type FileIndexedZSet<K, V> =
561    TypedBatch<K, V, ZWeight, DynFileIndexedWSet<DynData, DynData, DynZWeight>>;
562pub type FileKeyBatch<K, T, R, DynR> = TypedBatch<K, (), R, DynFileKeyBatch<DynData, T, DynR>>;
563pub type FileValBatch<K, V, T, R, DynR> =
564    TypedBatch<K, V, R, DynFileValBatch<DynData, DynData, T, DynR>>;
565
566pub type FallbackWSet<K, R, DynR> = TypedBatch<K, (), R, DynFallbackWSet<DynData, DynR>>;
567pub type FallbackZSet<K> = TypedBatch<K, (), ZWeight, DynFallbackWSet<DynData, DynZWeight>>;
568pub type FallbackIndexedWSet<K, V, R, DynR> =
569    TypedBatch<K, V, R, DynFallbackIndexedWSet<DynData, DynData, DynR>>;
570pub type FallbackIndexedZSet<K, V> =
571    TypedBatch<K, V, ZWeight, DynFallbackIndexedWSet<DynData, DynData, DynZWeight>>;
572pub type FallbackKeyBatch<K, T, R, DynR> =
573    TypedBatch<K, (), R, DynFallbackKeyBatch<DynData, T, DynR>>;
574pub type FallbackValBatch<K, V, T, R, DynR> =
575    TypedBatch<K, V, R, DynFallbackValBatch<DynData, DynData, T, DynR>>;
576
577pub type Spine<B> = TypedBatch<
578    <B as BatchReader>::Key,
579    <B as BatchReader>::Val,
580    <B as BatchReader>::R,
581    DynSpine<<B as BatchReader>::Inner>,
582>;
583
584pub type SpineSnapshot<B> = TypedBatch<
585    <B as BatchReader>::Key,
586    <B as BatchReader>::Val,
587    <B as BatchReader>::R,
588    DynSpineSnapshot<<B as BatchReader>::Inner>,
589>;
590
591impl<K, V, R, B> TypedBatch<K, V, R, DynSpineSnapshot<B>>
592where
593    B: DynBatch,
594    K: DBData + Erase<B::Key>,
595    V: DBData + Erase<B::Val>,
596    R: DBWeight + Erase<B::R>,
597{
598    /// Concatenate a list of snapshots into a single snapshot.
599    pub fn concat<'a, I>(snapshots: I) -> TypedBatch<K, V, R, DynSpineSnapshot<B>>
600    where
601        I: IntoIterator<Item = &'a Self>,
602    {
603        TypedBatch::new(DynSpineSnapshot::concat(
604            BatchReaderFactories::new::<K, V, R>(),
605            snapshots.into_iter().map(|snapshot| &snapshot.inner),
606        ))
607    }
608
609    /// Consolidate the batches in the snapshot.
610    pub fn consolidate(&self) -> TypedBatch<K, V, R, B> {
611        TypedBatch::new(self.inner.consolidate())
612    }
613}
614
615impl<K, V, R, B> TypedBatch<K, V, R, B>
616where
617    B: DynTrace,
618    K: DBData + Erase<B::Key>,
619    V: DBData + Erase<B::Val>,
620    R: DBWeight + Erase<B::R>,
621{
622    /// Consolidate the batches in the trace.
623    pub fn consolidate(self) -> TypedBatch<K, V, R, B::Batch> {
624        TypedBatch::new(
625            self.inner
626                .consolidate()
627                .unwrap_or_else(|| B::Batch::dyn_empty(&BatchReaderFactories::new::<K, V, R>())),
628        )
629    }
630}
631
632impl<K, V, R, B> TypedBatch<K, V, R, DynSpine<B>>
633where
634    B: DynBatch,
635    K: DBData + Erase<B::Key>,
636    V: DBData + Erase<B::Val>,
637    R: DBWeight + Erase<B::R>,
638{
639    pub fn ro_snapshot(&self) -> TypedBatch<K, V, R, DynSpineSnapshot<B>> {
640        TypedBatch::new(self.inner.ro_snapshot())
641    }
642}
643
644impl<C: Clone, B: BatchReader> Stream<C, B> {
645    pub fn inner(&self) -> Stream<C, B::Inner> {
646        BatchReader::stream_inner(self)
647    }
648}
649
650impl<C: Clone, B: DynBatchReader> Stream<C, B> {
651    pub fn typed<TB>(&self) -> Stream<C, TB>
652    where
653        TB: BatchReader<Inner = B>,
654    {
655        // Safety: repr(transparent) on TypedBatch guarantees that this is
656        // safe.
657        unsafe { self.transmute_payload() }
658    }
659}
660
661#[derive(Debug, PartialEq, Eq, PartialOrd, Ord, SizeOf)]
662#[repr(transparent)]
663pub struct TypedBox<T, D: ?Sized> {
664    inner: Box<D>,
665    phantom: PhantomData<fn(&T)>,
666}
667
668#[derive(PartialEq, Eq, PartialOrd, Ord)]
669pub struct ArchivedTypedBox<T: Archive>(<T as Archive>::Archived)
670where
671    <T as Archive>::Archived: PartialEq + Eq + PartialOrd + Ord;
672
673impl<T, D> Archive for TypedBox<T, D>
674where
675    T: DBData + Erase<D>,
676    D: DataTrait + ?Sized,
677{
678    type Archived = ArchivedTypedBox<T>;
679    type Resolver = <T as Archive>::Resolver;
680
681    unsafe fn resolve(&self, pos: usize, resolver: Self::Resolver, out: *mut Self::Archived) {
682        unsafe {
683            let val: &T = self.deref();
684            val.resolve(pos, resolver, &mut (*out).0 as *mut T::Archived);
685        }
686    }
687}
688
689impl<T, D> Serialize<DbspSerializer<'_>> for TypedBox<T, D>
690where
691    T: DBData + Erase<D>,
692    D: DataTrait + ?Sized,
693{
694    fn serialize(
695        &self,
696        serializer: &mut DbspSerializer,
697    ) -> Result<Self::Resolver, <DbspSerializer<'_> as Fallible>::Error> {
698        let val: &T = self.deref();
699        val.serialize(serializer)
700    }
701}
702
703impl<T, D> Deserialize<TypedBox<T, D>, Deserializer> for Archived<TypedBox<T, D>>
704where
705    D: DataTrait + ?Sized,
706    T: DBData + Erase<D>,
707{
708    fn deserialize(
709        &self,
710        deserializer: &mut Deserializer,
711    ) -> Result<TypedBox<T, D>, <Deserializer as Fallible>::Error> {
712        let val: T = self.0.deserialize(deserializer)?;
713        Ok(TypedBox::new(val))
714    }
715}
716
717#[cfg(test)]
718#[test]
719fn test_typedbox_rkyv() {
720    use rkyv::archived_value;
721
722    use crate::storage::file::SerializerInner;
723
724    let tbox = TypedBox::<u64, DynData>::new(12345u64);
725
726    let bytes = SerializerInner::to_fbuf_with_thread_local(|s| {
727        rkyv::ser::Serializer::serialize_value(s, &tbox)
728    })
729    .into_vec();
730
731    let archived: &<TypedBox<u64, DynData> as Archive>::Archived =
732        unsafe { archived_value::<TypedBox<u64, DynData>>(bytes.as_slice(), 0) };
733
734    let tbox2 = archived.deserialize(&mut Deserializer::default()).unwrap();
735
736    assert_eq!(tbox, tbox2);
737}
738
739impl<T, D> Deref for TypedBox<T, D>
740where
741    D: DataTrait + ?Sized,
742    T: DBData + Erase<D>,
743{
744    type Target = T;
745
746    fn deref(&self) -> &T {
747        unsafe { self.inner.downcast() }
748    }
749}
750
751impl<T, D> DerefMut for TypedBox<T, D>
752where
753    D: DataTrait + ?Sized,
754    T: DBData + Erase<D>,
755{
756    fn deref_mut(&mut self) -> &mut T {
757        unsafe { self.inner.downcast_mut() }
758    }
759}
760
761impl<T, D: ?Sized> NumEntries for TypedBox<T, D> {
762    const CONST_NUM_ENTRIES: Option<usize> = None;
763
764    fn num_entries_shallow(&self) -> usize {
765        1
766    }
767
768    fn num_entries_deep(&self) -> usize {
769        1
770    }
771}
772
773impl<T, D: DataTrait + ?Sized> TypedBox<T, D> {
774    pub fn new(v: T) -> Self
775    where
776        T: DBData + Erase<D>,
777    {
778        Self {
779            inner: Box::new(v).erase_box(),
780            phantom: PhantomData,
781        }
782    }
783
784    pub fn inner(&self) -> &D {
785        self.inner.as_ref()
786    }
787
788    pub fn into_inner(self) -> Box<D> {
789        self.inner
790    }
791}
792
793impl<T, D> Clone for TypedBox<T, D>
794where
795    D: DataTrait + ?Sized,
796{
797    fn clone(&self) -> Self {
798        Self {
799            inner: clone_box(self.inner.as_ref()),
800            phantom: PhantomData,
801        }
802    }
803}
804
805impl<C: Clone, D: DataTrait + ?Sized> Stream<C, Box<D>> {
806    /// Adds type information to `self`, wrapping each element
807    /// in a [`TypedBox`].  This function is a noop at runtime.
808    ///
809    /// # Safety
810    ///
811    /// `self` must contain concrete values of type `T`.
812    pub unsafe fn typed_data<T>(&self) -> Stream<C, TypedBox<T, D>>
813    where
814        T: DBData + Erase<D>,
815    {
816        unsafe { self.transmute_payload() }
817    }
818}
819
820impl<C: Circuit, T, D: ?Sized> Stream<C, TypedBox<T, D>> {
821    pub fn inner_data(&self) -> Stream<C, Box<D>> {
822        unsafe { self.transmute_payload() }
823    }
824}
825
826impl<C: Circuit, T: DBData, D: DataTrait + ?Sized> Stream<C, TypedBox<T, D>> {
827    pub fn inner_typed(&self) -> Stream<C, T> {
828        self.apply(|typed_box| unsafe { typed_box.inner().downcast::<T>().clone() })
829    }
830}
831
832impl<C: Circuit, T: DBData> Stream<C, T> {
833    pub fn typed_box<D>(&self) -> Stream<C, TypedBox<T, D>>
834    where
835        D: DataTrait + ?Sized,
836        T: Erase<D>,
837    {
838        self.apply(|x| TypedBox::new(x.clone()))
839    }
840}