Skip to main content

dbsp/operator/dynamic/
input_upsert.rs

1use crate::{
2    Circuit, DBData, NumEntries, RootCircuit, Stream, ZWeight,
3    algebra::{HasOne, HasZero, IndexedZSet, OrdZSet, ZTrace},
4    circuit::{
5        OwnershipPreference, Scope,
6        checkpointer::Checkpoint,
7        circuit_builder::{CircuitBase, RefStreamValue, register_replay_stream},
8        metadata::{BatchSizeStats, INPUT_BATCHES_STATS, OUTPUT_BATCHES_STATS, OperatorMeta},
9        operator_traits::{BinaryOperator, Operator, TernaryOperator},
10    },
11    declare_trait_object,
12    dynamic::{
13        ClonableTrait, Data, DataTrait, DynOpt, DynPairs, Erase, Factory, LeanVec, WithFactory,
14    },
15    operator::{
16        Z1,
17        dynamic::{
18            time_series::LeastUpperBoundFunc,
19            trace::{DelayedTraceId, TraceBounds, TraceId, UntimedTraceAppend, Z1Trace},
20        },
21    },
22    trace::{
23        Batch, BatchFactories, BatchReader, BatchReaderFactories, Builder, Rkyv, Spine,
24        cursor::Cursor,
25    },
26    utils::Tup2,
27};
28use feldera_macros::{IsNone, OrdRepr};
29use rkyv::{Archive, Deserialize, Serialize};
30use size_of::SizeOf;
31use std::{
32    borrow::Cow,
33    marker::PhantomData,
34    mem::take,
35    ops::{Deref, Neg},
36};
37
38use super::trace::BoundsId;
39
40#[derive(
41    Clone,
42    Debug,
43    Default,
44    SizeOf,
45    PartialEq,
46    Eq,
47    Hash,
48    PartialOrd,
49    Ord,
50    Archive,
51    Serialize,
52    Deserialize,
53    IsNone,
54    OrdRepr,
55)]
56#[archive_attr(derive(Ord, Eq, PartialEq, PartialOrd))]
57#[archive(compare(PartialEq, PartialOrd))]
58pub enum Update<V: DBData, U: DBData> {
59    Insert(V),
60    #[default]
61    Delete,
62    Update(U),
63}
64
65pub enum UpdateRef<'a, V: DataTrait + ?Sized, U: DataTrait + ?Sized> {
66    Insert(&'a V),
67    Delete,
68    Update(&'a U),
69}
70
71impl<V: DBData, U: DBData> NumEntries for Update<V, U> {
72    const CONST_NUM_ENTRIES: Option<usize> = Some(1);
73
74    fn num_entries_shallow(&self) -> usize {
75        1
76    }
77
78    fn num_entries_deep(&self) -> usize {
79        1
80    }
81}
82
83pub trait UpdateTrait<V: DataTrait + ?Sized, U: DataTrait + ?Sized>: Data {
84    fn get(&self) -> UpdateRef<'_, V, U>;
85    fn insert_ref(&mut self, val: &V);
86    fn insert_val(&mut self, val: &mut V);
87    fn delete(&mut self);
88    fn update_ref(&mut self, upd: &U);
89    fn update_val(&mut self, upd: &mut U);
90}
91
92impl<V, U, VType, UType> UpdateTrait<V, U> for Update<VType, UType>
93where
94    V: DataTrait + ?Sized,
95    U: DataTrait + ?Sized,
96    VType: DBData + Erase<V>,
97    UType: DBData + Erase<U>,
98{
99    fn get(&self) -> UpdateRef<'_, V, U> {
100        match self {
101            Update::Insert(v) => UpdateRef::Insert(v.erase()),
102            Update::Delete => UpdateRef::Delete,
103            Update::Update(u) => UpdateRef::Update(u.erase()),
104        }
105    }
106
107    fn insert_ref(&mut self, val: &V) {
108        *self = Update::Insert(unsafe { val.downcast::<VType>().clone() })
109    }
110
111    fn insert_val(&mut self, val: &mut V) {
112        *self = Update::Insert(take(unsafe { val.downcast_mut::<VType>() }))
113    }
114
115    fn delete(&mut self) {
116        *self = Update::Delete;
117    }
118
119    fn update_ref(&mut self, upd: &U) {
120        *self = Update::Update(unsafe { upd.downcast::<UType>().clone() })
121    }
122
123    fn update_val(&mut self, upd: &mut U) {
124        *self = Update::Update(take(unsafe { upd.downcast_mut::<UType>() }))
125    }
126}
127
128declare_trait_object!(DynUpdate<VTrait, UTrait> = dyn UpdateTrait<VTrait, UTrait>
129where
130    VTrait: DataTrait + ?Sized,
131    UTrait: DataTrait + ?Sized,
132);
133
134pub type PatchFunc<V, U> = Box<dyn Fn(&mut V, &U)>;
135
136pub struct InputUpsertFactories<B: IndexedZSet, U: DataTrait + ?Sized> {
137    pub batch_factories: B::Factories,
138    pub opt_key_factory: &'static dyn Factory<DynOpt<B::Key>>,
139    pub opt_val_factory: &'static dyn Factory<DynOpt<B::Val>>,
140    pub pairs_factory: &'static dyn Factory<DynPairs<B::Key, DynUpdate<B::Val, U>>>,
141}
142
143impl<B: IndexedZSet, U: DataTrait + ?Sized> Clone for InputUpsertFactories<B, U> {
144    fn clone(&self) -> Self {
145        Self {
146            batch_factories: self.batch_factories.clone(),
147            opt_key_factory: self.opt_key_factory,
148            opt_val_factory: self.opt_val_factory,
149            pairs_factory: self.pairs_factory,
150        }
151    }
152}
153
154impl<B, U> InputUpsertFactories<B, U>
155where
156    B: Batch + IndexedZSet,
157    U: DataTrait + ?Sized,
158{
159    pub fn new<KType, VType, UType>() -> Self
160    where
161        KType: DBData + Erase<B::Key>,
162        VType: DBData + Erase<B::Val>,
163        UType: DBData + Erase<U>,
164    {
165        Self {
166            batch_factories: BatchReaderFactories::new::<KType, VType, ZWeight>(),
167            opt_key_factory: WithFactory::<Option<KType>>::FACTORY,
168            opt_val_factory: WithFactory::<Option<VType>>::FACTORY,
169            pairs_factory: WithFactory::<LeanVec<Tup2<KType, Update<VType, UType>>>>::FACTORY,
170        }
171    }
172}
173
174pub struct InputUpsertWithWaterlineFactories<
175    B: IndexedZSet,
176    U: DataTrait + ?Sized,
177    E: DataTrait + ?Sized,
178> {
179    pub batch_factories: B::Factories,
180    pub opt_key_factory: &'static dyn Factory<DynOpt<B::Key>>,
181    pub opt_val_factory: &'static dyn Factory<DynOpt<B::Val>>,
182    pub val_factory: &'static dyn Factory<B::Val>,
183    pub pairs_factory: &'static dyn Factory<DynPairs<B::Key, DynUpdate<B::Val, U>>>,
184    errors_factory: <OrdZSet<E> as BatchReader>::Factories,
185}
186
187impl<B: IndexedZSet, U: DataTrait + ?Sized, E: DataTrait + ?Sized> Clone
188    for InputUpsertWithWaterlineFactories<B, U, E>
189{
190    fn clone(&self) -> Self {
191        Self {
192            batch_factories: self.batch_factories.clone(),
193            opt_key_factory: self.opt_key_factory,
194            opt_val_factory: self.opt_val_factory,
195            val_factory: self.val_factory,
196            pairs_factory: self.pairs_factory,
197            errors_factory: self.errors_factory.clone(),
198        }
199    }
200}
201
202impl<B, U, E> InputUpsertWithWaterlineFactories<B, U, E>
203where
204    B: Batch + IndexedZSet,
205    U: DataTrait + ?Sized,
206    E: DataTrait + ?Sized,
207{
208    pub fn new<KType, VType, UType, EType>() -> Self
209    where
210        KType: DBData + Erase<B::Key>,
211        VType: DBData + Erase<B::Val>,
212        UType: DBData + Erase<U>,
213        EType: DBData + Erase<E>,
214    {
215        Self {
216            batch_factories: BatchReaderFactories::new::<KType, VType, ZWeight>(),
217            opt_key_factory: WithFactory::<Option<KType>>::FACTORY,
218            opt_val_factory: WithFactory::<Option<VType>>::FACTORY,
219            val_factory: WithFactory::<VType>::FACTORY,
220            pairs_factory: WithFactory::<LeanVec<Tup2<KType, Update<VType, UType>>>>::FACTORY,
221            errors_factory: BatchReaderFactories::new::<EType, (), ZWeight>(),
222        }
223    }
224}
225
226impl<K, V, U> Stream<RootCircuit, Vec<Box<DynPairs<K, DynUpdate<V, U>>>>>
227where
228    K: DataTrait + ?Sized,
229    V: DataTrait + ?Sized,
230    U: DataTrait + ?Sized,
231{
232    /// Convert an input stream of upserts into a stream of updates to a
233    /// relation.
234    ///
235    /// The input stream carries changes to a key/value map in the form of
236    /// _upserts_.  An upsert assigns a new value to a key (or deletes the key
237    /// from the map) without explicitly removing the old value, if any.  The
238    /// operator converts upserts into batches of updates, which is the input
239    /// format of most DBSP operators.
240    ///
241    /// The operator assumes that the input vector is sorted by key; however,
242    /// unlike the [`Stream::upsert`] operator it allows the vector to
243    /// contain multiple updates per key.  Updates are applied one by one in
244    /// order, and the output of the operator reflects cumulative effect of
245    /// the updates.  Additionally, unlike the [`Stream::upsert`] operator,
246    /// which only supports inserts, which overwrite the entire value with a
247    /// new value, and deletions, this operator also supports updates that
248    /// modify the contents of a value, e.g., overwriting some of its
249    /// fields.  Type argument `U` defines the format of modifications,
250    /// and the `patch_func` function applies update of type `U` to a value of
251    /// type `V`.
252    ///
253    /// This is a stateful operator that internally maintains the trace of the
254    /// collection.
255    pub fn input_upsert<B>(
256        &self,
257        persistent_id: Option<&str>,
258        factories: &InputUpsertFactories<B, U>,
259        patch_func: PatchFunc<V, U>,
260    ) -> Stream<RootCircuit, B>
261    where
262        B: IndexedZSet<Key = K, Val = V>,
263    {
264        let circuit = self.circuit();
265
266        // We build the following circuit to implement the upsert semantics.
267        // The collection is accumulated into a trace using integrator
268        // (UntimedTraceAppend + Z1Trace = integrator).  The `InputUpsert`
269        // operator evaluates each upsert command in the input stream against
270        // the trace and computes a batch of updates to be added to the trace.
271        //
272        // ```text
273        //                               ┌────────────────────────────►
274        //                               │
275        //                               │
276        //  self        ┌───────────┐    │        ┌──────────────────┐  trace
277        // ────────────►│InputUpsert├────┴───────►│UntimedTraceAppend├────┐
278        //              └───────────┘   delta     └──────────────────┘    │
279        //                      ▲                  ▲                      │
280        //                      │                  │                      │
281        //                      │                  │   ┌───────┐          │
282        //                      └──────────────────┴───┤Z1Trace│◄─────────┘
283        //                         z1trace             └───────┘
284        // ```
285        circuit.region("input_upsert", || {
286            let bounds = <TraceBounds<K, V>>::unbounded();
287
288            let sharded = self.dyn_shard_pairs(factories.pairs_factory);
289
290            let z1 = Z1Trace::new(
291                &factories.batch_factories,
292                &factories.batch_factories,
293                false,
294                circuit.root_scope(),
295                bounds.clone(),
296            );
297
298            let (delayed_trace, z1feedback) = circuit.add_feedback_persistent(
299                persistent_id
300                    .map(|name| format!("{name}.integral"))
301                    .as_deref(),
302                z1,
303            );
304
305            delayed_trace.mark_sharded();
306
307            let delta = circuit
308                .add_binary_operator(
309                    <InputUpsert<Spine<B>, U, B>>::new(
310                        factories.batch_factories.clone(),
311                        factories.opt_val_factory,
312                        patch_func,
313                    ),
314                    &delayed_trace,
315                    &sharded,
316                )
317                .mark_distinct();
318            delta.mark_sharded();
319            let replay_stream = z1feedback.operator_mut().prepare_replay_stream(&delta);
320
321            let trace = circuit.add_binary_operator_with_preference(
322                UntimedTraceAppend::<Spine<B>>::new(),
323                (&delayed_trace, OwnershipPreference::STRONGLY_PREFER_OWNED),
324                (&delta, OwnershipPreference::PREFER_OWNED),
325            );
326            trace.mark_sharded();
327
328            z1feedback.connect_with_preference(&trace, OwnershipPreference::STRONGLY_PREFER_OWNED);
329
330            register_replay_stream(circuit, &delta, &replay_stream, &factories.batch_factories);
331
332            circuit.cache_insert(DelayedTraceId::new(trace.stream_id()), delayed_trace);
333            circuit.cache_insert(TraceId::new(delta.stream_id()), trace);
334            circuit.cache_insert(BoundsId::<B>::new(delta.stream_id()), bounds);
335            delta
336        })
337    }
338
339    // Like `input_upsert`, but additionally tracks a waterline of the input collection and
340    // rejects inputs that are below the waterline.  An input is rejected if the input record
341    // itself is below the waterline or if the existing record it replaces is below the waterline.
342    #[allow(clippy::too_many_arguments)]
343    pub fn input_upsert_with_waterline<B, W, E>(
344        &self,
345        persistent_id: Option<&str>,
346        factories: &InputUpsertWithWaterlineFactories<B, U, E>,
347        patch_func: PatchFunc<V, U>,
348        init_waterline: Box<dyn Fn() -> Box<W>>,
349        extract_ts: Box<dyn Fn(&B::Key, &B::Val, &mut W)>,
350        least_upper_bound: LeastUpperBoundFunc<W>,
351        filter_func: Box<dyn Fn(&W, &B::Key, &B::Val) -> bool>,
352        report_func: Box<dyn Fn(&W, &B::Key, &B::Val, ZWeight, &mut E)>,
353    ) -> (
354        Stream<RootCircuit, B>,
355        Stream<RootCircuit, OrdZSet<E>>,
356        Stream<RootCircuit, Box<W>>,
357    )
358    where
359        B: IndexedZSet<Key = K, Val = V>,
360        W: DataTrait + Checkpoint + ?Sized,
361        E: DataTrait + ?Sized,
362        Box<W>: Checkpoint + Clone + NumEntries + Rkyv,
363    {
364        let circuit = self.circuit();
365
366        // ```text
367        //                   ┌─────────────────────────────────────────────►
368        //                   │ waterline
369        //             ┌─────┴─────┐
370        //             │ waterline │◄─────────┬────────────────────────────►
371        //             └──────┬────┘          │
372        //                    │ waterline     │
373        //                   Z1               │
374        //  delayed_waterline │               │
375        //                    ▼               │
376        //         ┌─────────────────────┐    │        ┌──────────────────┐  trace
377        // ───────►│InputUpsertWaterline ├────┴───────►│UntimedTraceAppend├────┐
378        //         └──────┬──────────────┘   delta     └──────────────────┘    │
379        //                │          ▲                  ▲                      │
380        //                │          │                  │                      │
381        //                │          │                  │   ┌───────┐          │
382        //                │          └──────────────────┴───┤Z1Trace│◄─────────┘
383        //                │             delayed_trace       └───────┘
384        //                │
385        //                │error stream
386        //                └────────────────────────────────────────────────►
387        // ```
388
389        circuit.region("input_upsert_waterline", || {
390            let bounds = <TraceBounds<K, V>>::unbounded();
391
392            let sharded = self.dyn_shard_pairs(factories.pairs_factory);
393
394            let z1 = Z1Trace::new(
395                &factories.batch_factories,
396                &factories.batch_factories,
397                false,
398                circuit.root_scope(),
399                bounds.clone(),
400            );
401
402            let (delayed_trace, z1feedback) = circuit.add_feedback_persistent(
403                persistent_id
404                    .map(|name| format!("{name}.integral"))
405                    .as_deref(),
406                z1,
407            );
408
409            delayed_trace.mark_sharded();
410
411            let waterline_z1 = Z1::new((init_waterline)());
412
413            let (delayed_waterline, waterline_feedback) = circuit.add_feedback_persistent(
414                persistent_id
415                    .map(|name| format!("{name}.delayed_waterline"))
416                    .as_deref(),
417                waterline_z1,
418            );
419
420            let error_stream_val = RefStreamValue::empty();
421
422            let delta = circuit
423                .add_ternary_operator(
424                    <InputUpsertWithWaterline<Spine<B>, U, B, W, E>>::new(
425                        factories.clone(),
426                        patch_func,
427                        filter_func,
428                        report_func,
429                        error_stream_val.clone(),
430                    ),
431                    &delayed_trace,
432                    &sharded,
433                    &delayed_waterline,
434                )
435                .mark_distinct();
436            delta.mark_sharded();
437            let replay_stream = z1feedback.operator_mut().prepare_replay_stream(&delta);
438
439            let waterline_id = persistent_id.map(|name| format!("{name}.input_waterline"));
440
441            let waterline = delta.dyn_waterline(
442                waterline_id.as_deref(),
443                init_waterline,
444                extract_ts,
445                least_upper_bound,
446            );
447
448            waterline_feedback.connect(&waterline);
449
450            let trace = circuit.add_binary_operator_with_preference(
451                UntimedTraceAppend::<Spine<B>>::new(),
452                (&delayed_trace, OwnershipPreference::STRONGLY_PREFER_OWNED),
453                (&delta, OwnershipPreference::PREFER_OWNED),
454            );
455            trace.mark_sharded();
456
457            z1feedback.connect_with_preference(&trace, OwnershipPreference::STRONGLY_PREFER_OWNED);
458
459            register_replay_stream(circuit, &delta, &replay_stream, &factories.batch_factories);
460
461            let error_stream = Stream::with_value(
462                self.circuit().clone(),
463                delta.local_node_id(),
464                error_stream_val,
465            );
466
467            circuit.cache_insert(DelayedTraceId::new(trace.stream_id()), delayed_trace);
468            circuit.cache_insert(TraceId::new(delta.stream_id()), trace);
469            circuit.cache_insert(BoundsId::<B>::new(delta.stream_id()), bounds);
470
471            (delta, error_stream, waterline)
472        })
473    }
474}
475
476pub struct InputUpsert<T, U, B>
477where
478    T: BatchReader,
479    B: Batch,
480    U: DataTrait + ?Sized,
481{
482    batch_factories: B::Factories,
483    opt_val_factory: &'static dyn Factory<DynOpt<B::Val>>,
484    patch_func: PatchFunc<T::Val, U>,
485
486    // Input batch sizes.
487    input_batch_stats: BatchSizeStats,
488
489    // Output batch sizes.
490    output_batch_stats: BatchSizeStats,
491
492    phantom: PhantomData<B>,
493}
494
495impl<T, U, B> InputUpsert<T, U, B>
496where
497    T: BatchReader,
498    B: Batch,
499    U: DataTrait + ?Sized,
500{
501    pub fn new(
502        batch_factories: B::Factories,
503        opt_val_factory: &'static dyn Factory<DynOpt<B::Val>>,
504        patch_func: PatchFunc<T::Val, U>,
505    ) -> Self {
506        Self {
507            batch_factories,
508            opt_val_factory,
509            patch_func,
510            input_batch_stats: BatchSizeStats::new(),
511            output_batch_stats: BatchSizeStats::new(),
512            phantom: PhantomData,
513        }
514    }
515}
516
517impl<T, U, B> Operator for InputUpsert<T, U, B>
518where
519    T: BatchReader,
520    U: DataTrait + ?Sized,
521    B: Batch,
522{
523    fn name(&self) -> Cow<'static, str> {
524        Cow::from("InputUpsert")
525    }
526
527    fn metadata(&self, meta: &mut OperatorMeta) {
528        meta.extend(metadata! {
529            INPUT_BATCHES_STATS => self.input_batch_stats.metadata(),
530            OUTPUT_BATCHES_STATS => self.output_batch_stats.metadata(),
531        });
532    }
533
534    fn fixedpoint(&self, _scope: Scope) -> bool {
535        true
536    }
537}
538
539impl<T, U, B> BinaryOperator<T, Vec<Box<DynPairs<T::Key, DynUpdate<T::Val, U>>>>, B>
540    for InputUpsert<T, U, B>
541where
542    T: ZTrace<Time = ()>,
543    U: DataTrait + ?Sized,
544    B: IndexedZSet<Key = T::Key, Val = T::Val>,
545{
546    async fn eval(
547        &mut self,
548        trace: &T,
549        updates: &Vec<Box<DynPairs<T::Key, DynUpdate<T::Val, U>>>>,
550    ) -> B {
551        // Inputs must be sorted by key
552        let mut updates = updates
553            .iter()
554            .filter_map(|updates| {
555                if !updates.is_empty() {
556                    Some((&**updates, 0))
557                } else {
558                    None
559                }
560            })
561            .collect::<Vec<_>>();
562        let n_updates = updates.iter().map(|updates| updates.0.len()).sum();
563        debug_assert!(
564            updates
565                .iter()
566                .all(|updates| updates.0.is_sorted_by(&|u1, u2| u1.fst().cmp(u2.fst())))
567        );
568
569        self.input_batch_stats.add_batch(n_updates);
570
571        let mut key_updates = self.batch_factories.weighted_vals_factory().default_box();
572
573        let mut trace_cursor = trace.cursor();
574
575        let mut builder =
576            B::Builder::with_capacity(&self.batch_factories, n_updates * 2, n_updates * 2);
577
578        // Current key for which we are processing updates.
579        let mut cur_key = None;
580
581        // Current value associated with the key after applying all processed updates
582        // to it.
583        let mut cur_val: Box<DynOpt<T::Val>> = self.opt_val_factory.default_box();
584
585        while !updates.is_empty() {
586            let (index, key_upd) = updates
587                .iter()
588                .map(|(updates, index)| updates.index(*index))
589                .enumerate()
590                // Find the first update with the smallest key (compare keys, not updates, so that we apply updates in order).
591                // min_by is guaranteed to return the first among equal keys.
592                .min_by(|(_a_index, a), (_b_index, b)| a.fst().cmp(b.fst()))
593                .unwrap();
594            updates[index].1 += 1;
595            if updates[index].1 >= updates[index].0.len() {
596                updates.remove(index);
597            }
598
599            let (key, upd) = key_upd.split();
600
601            // We finished processing updates for the previous key. Push them to the
602            // builder and generate a retraction for the new key.
603            if cur_key != Some(key) {
604                // Push updates for the previous key to the builder.
605                if let Some(cur_key) = cur_key {
606                    if let Some(val) = cur_val.get_mut() {
607                        key_updates.push_with(&mut |item| {
608                            let (v, w) = item.split_mut();
609
610                            val.move_to(v);
611                            **w = HasOne::one();
612                        });
613                    }
614                    key_updates.consolidate();
615                    if !key_updates.is_empty() {
616                        for pair in key_updates.dyn_iter_mut() {
617                            let (v, d) = pair.split_mut();
618                            builder.push_val_diff_mut(v, d);
619                        }
620                        builder.push_key(cur_key);
621                    }
622                    key_updates.clear();
623                }
624
625                cur_key = Some(key);
626                cur_val.set_none();
627
628                // Generate retraction if `key` is present in the trace.
629                if trace_cursor.seek_key_exact(key, None) {
630                    // println!("{}: found key in trace_cursor", Runtime::worker_index());
631                    while trace_cursor.val_valid() {
632                        let weight = **trace_cursor.weight();
633
634                        if !weight.is_zero() {
635                            let val = trace_cursor.val();
636
637                            key_updates.push_with(&mut |item| {
638                                let (v, w) = item.split_mut();
639
640                                val.clone_to(v);
641                                **w = weight.neg()
642                            });
643                            cur_val.from_ref(val);
644                        }
645
646                        trace_cursor.step_val();
647                    }
648                }
649            }
650
651            match upd.get() {
652                UpdateRef::Delete => {
653                    // TODO: if cur_val.is_none(), report missing key.
654                    cur_val.set_none();
655                }
656                UpdateRef::Insert(val) => {
657                    cur_val.from_ref(val);
658                }
659                UpdateRef::Update(upd) => {
660                    if let Some(val) = cur_val.get_mut() {
661                        (self.patch_func)(val, upd);
662                    } else {
663                        // TODO: report missing key.
664                    }
665                }
666            }
667        }
668
669        // Push updates for the last key.
670        if let Some(cur_key) = cur_key {
671            if let Some(val) = cur_val.get_mut() {
672                key_updates.push_with(&mut |item| {
673                    let (v, w) = item.split_mut();
674
675                    val.move_to(v);
676                    **w = HasOne::one();
677                });
678            }
679
680            key_updates.consolidate();
681            if !key_updates.is_empty() {
682                for pair in key_updates.dyn_iter_mut() {
683                    let (v, d) = pair.split_mut();
684                    builder.push_val_diff_mut(v, d);
685                }
686                builder.push_key(cur_key);
687            }
688            key_updates.clear();
689        }
690
691        builder.done()
692    }
693
694    fn input_preference(&self) -> (OwnershipPreference, OwnershipPreference) {
695        (
696            OwnershipPreference::PREFER_OWNED,
697            OwnershipPreference::PREFER_OWNED,
698        )
699    }
700}
701
702pub struct InputUpsertWithWaterline<T, U, B, W, E>
703where
704    T: BatchReader,
705    B: IndexedZSet,
706    U: DataTrait + ?Sized,
707    W: DataTrait + ?Sized,
708    E: DataTrait + ?Sized,
709{
710    factories: InputUpsertWithWaterlineFactories<B, U, E>,
711    patch_func: PatchFunc<T::Val, U>,
712    filter_func: Box<dyn Fn(&W, &B::Key, &B::Val) -> bool>,
713    report_func: Box<dyn Fn(&W, &B::Key, &B::Val, ZWeight, &mut E)>,
714    error_stream_val: RefStreamValue<OrdZSet<E>>,
715
716    // Input batch sizes.
717    input_batch_stats: BatchSizeStats,
718
719    // Output batch sizes.
720    output_batch_stats: BatchSizeStats,
721
722    phantom: PhantomData<B>,
723}
724
725impl<T, U, B, W, E> InputUpsertWithWaterline<T, U, B, W, E>
726where
727    T: BatchReader,
728    B: IndexedZSet,
729    U: DataTrait + ?Sized,
730    W: DataTrait + ?Sized,
731    E: DataTrait + ?Sized,
732{
733    pub fn new(
734        factories: InputUpsertWithWaterlineFactories<B, U, E>,
735        patch_func: PatchFunc<T::Val, U>,
736        filter_func: Box<dyn Fn(&W, &B::Key, &B::Val) -> bool>,
737        report_func: Box<dyn Fn(&W, &B::Key, &B::Val, ZWeight, &mut E)>,
738        error_stream_val: RefStreamValue<OrdZSet<E>>,
739    ) -> Self {
740        Self {
741            factories,
742            patch_func,
743            filter_func,
744            report_func,
745            error_stream_val,
746            input_batch_stats: BatchSizeStats::new(),
747            output_batch_stats: BatchSizeStats::new(),
748            phantom: PhantomData,
749        }
750    }
751
752    fn passes_filter(&self, waterline: &W, key: &B::Key, val: &B::Val) -> bool {
753        (self.filter_func)(waterline, key, val)
754    }
755}
756
757impl<T, U, B, W, E> Operator for InputUpsertWithWaterline<T, U, B, W, E>
758where
759    T: BatchReader,
760    U: DataTrait + ?Sized,
761    B: IndexedZSet,
762    W: DataTrait + ?Sized,
763    E: DataTrait + ?Sized,
764{
765    fn name(&self) -> Cow<'static, str> {
766        Cow::from("InputUpsertWithWaterline")
767    }
768
769    fn metadata(&self, meta: &mut OperatorMeta) {
770        meta.extend(metadata! {
771            INPUT_BATCHES_STATS => self.input_batch_stats.metadata(),
772            OUTPUT_BATCHES_STATS => self.output_batch_stats.metadata(),
773        });
774    }
775
776    fn fixedpoint(&self, _scope: Scope) -> bool {
777        true
778    }
779}
780
781impl<T, U, B, W, E> TernaryOperator<T, Vec<Box<DynPairs<T::Key, DynUpdate<T::Val, U>>>>, Box<W>, B>
782    for InputUpsertWithWaterline<T, U, B, W, E>
783where
784    T: ZTrace<Time = ()> + Clone,
785    U: DataTrait + ?Sized,
786    B: IndexedZSet<Key = T::Key, Val = T::Val>,
787    W: DataTrait + ?Sized,
788    Box<W>: Clone,
789    E: DataTrait + ?Sized,
790{
791    async fn eval(
792        &mut self,
793        trace: Cow<'_, T>,
794        updates: Cow<'_, Vec<Box<DynPairs<T::Key, DynUpdate<T::Val, U>>>>>,
795        waterline: Cow<'_, Box<W>>,
796    ) -> B {
797        // Inputs must be sorted by key
798        let mut updates = updates
799            .iter()
800            .filter_map(|updates| {
801                if !updates.is_empty() {
802                    Some((updates, 0))
803                } else {
804                    None
805                }
806            })
807            .collect::<Vec<_>>();
808        let n_updates = updates.iter().map(|updates| updates.0.len()).sum();
809        debug_assert!(
810            updates
811                .iter()
812                .all(|updates| updates.0.is_sorted_by(&|u1, u2| u1.fst().cmp(u2.fst())))
813        );
814
815        self.input_batch_stats.add_batch(n_updates);
816
817        let mut errors = self
818            .factories
819            .errors_factory
820            .weighted_items_factory()
821            .default_box();
822
823        let waterline = waterline.deref();
824        let mut key_updates = self
825            .factories
826            .batch_factories
827            .weighted_vals_factory()
828            .default_box();
829
830        let mut trace_cursor = trace.deref().cursor();
831
832        let mut builder = B::Builder::with_capacity(
833            &self.factories.batch_factories,
834            n_updates * 2,
835            n_updates * 2,
836        );
837
838        // Current key for which we are processing updates.
839        let mut cur_key = None;
840
841        // Current value associated with the key after applying all processed updates
842        // to it.
843        let mut cur_val: Box<DynOpt<T::Val>> = self.factories.opt_val_factory.default_box();
844        let mut tmp_val: Box<T::Val> = self.factories.val_factory.default_box();
845
846        // Set to true when the value associated with the current key doesn't
847        // satisfy `val_filter`, hence refuse to remove this value and process
848        // all updates for this key.
849        let mut skip_key = false;
850
851        while !updates.is_empty() {
852            let (index, key_upd) = updates
853                .iter()
854                .map(|(updates, index)| updates.index(*index))
855                .enumerate()
856                .min_by(|(_a_index, a), (_b_index, b)| a.cmp(b))
857                .unwrap();
858            updates[index].1 += 1;
859            if updates[index].1 >= updates[index].0.len() {
860                updates.remove(index);
861            }
862
863            let (key, upd) = key_upd.split();
864
865            // We finished processing updates for the previous key. Push them to the
866            // builder and generate a retraction for the new key.
867            if cur_key != Some(key) {
868                // Push updates for the previous key to the builder.
869                if let Some(cur_key) = cur_key {
870                    if let Some(val) = cur_val.get_mut() {
871                        key_updates.push_with(&mut |item| {
872                            let (v, w) = item.split_mut();
873
874                            val.move_to(v);
875                            **w = HasOne::one();
876                        });
877                    }
878                    key_updates.consolidate();
879                    if !key_updates.is_empty() {
880                        for pair in key_updates.dyn_iter_mut() {
881                            let (v, d) = pair.split_mut();
882                            builder.push_val_diff_mut(v, d);
883                        }
884                        builder.push_key(cur_key);
885                    }
886                    key_updates.clear();
887                }
888
889                skip_key = false;
890                cur_key = Some(key);
891                cur_val.set_none();
892
893                // Generate retraction if `key` is present in the trace.
894                if trace_cursor.seek_key_exact(key, None) {
895                    // println!("{}: found key in trace_cursor", Runtime::worker_index());
896                    while trace_cursor.val_valid() {
897                        let weight = **trace_cursor.weight();
898
899                        if !weight.is_zero() {
900                            let val = trace_cursor.val();
901
902                            if self.passes_filter(waterline, key, val) {
903                                key_updates.push_with(&mut |item| {
904                                    let (v, w) = item.split_mut();
905
906                                    val.clone_to(v);
907                                    **w = weight.neg()
908                                });
909                                cur_val.from_ref(val);
910                            } else {
911                                skip_key = true;
912                                errors.push_with(&mut |item| {
913                                    let (kv, err_weight) = item.split_mut();
914                                    **err_weight = HasOne::one();
915                                    (self.report_func)(
916                                        waterline,
917                                        key,
918                                        val,
919                                        weight.neg(),
920                                        kv.fst_mut(),
921                                    );
922                                });
923                            }
924                        }
925
926                        trace_cursor.step_val();
927                    }
928                }
929            }
930
931            if !skip_key {
932                match upd.get() {
933                    UpdateRef::Delete => {
934                        // TODO: if cur_val.is_none(), report missing key.
935                        cur_val.set_none();
936                    }
937                    UpdateRef::Insert(val) => {
938                        if self.passes_filter(waterline, key, val) {
939                            cur_val.from_ref(val);
940                        } else {
941                            errors.push_with(&mut |item| {
942                                let (kv, err_weight) = item.split_mut();
943                                **err_weight = HasOne::one();
944                                (self.report_func)(waterline, key, val, 1, kv.fst_mut());
945                            });
946                        }
947                    }
948                    UpdateRef::Update(upd) => {
949                        if let Some(val) = cur_val.get_mut() {
950                            val.clone_to(&mut tmp_val);
951                            (self.patch_func)(&mut tmp_val, upd);
952                            if !self.passes_filter(waterline, key, &tmp_val) {
953                                errors.push_with(&mut |item| {
954                                    let (kv, err_weight) = item.split_mut();
955                                    **err_weight = HasOne::one();
956                                    (self.report_func)(waterline, key, &tmp_val, 1, kv.fst_mut());
957                                });
958                            } else {
959                                tmp_val.clone_to(val);
960                            }
961                        } else {
962                            // TODO: report missing key.
963                        }
964                    }
965                }
966            }
967        }
968
969        // Push updates for the last key.
970        if let Some(cur_key) = cur_key {
971            if let Some(val) = cur_val.get_mut() {
972                key_updates.push_with(&mut |item| {
973                    let (v, w) = item.split_mut();
974
975                    val.move_to(v);
976                    **w = HasOne::one();
977                });
978            }
979
980            key_updates.consolidate();
981            if !key_updates.is_empty() {
982                for pair in key_updates.dyn_iter_mut() {
983                    let (v, d) = pair.split_mut();
984                    builder.push_val_diff_mut(v, d);
985                }
986                builder.push_key(cur_key);
987            }
988            key_updates.clear();
989        }
990
991        let errors = <OrdZSet<E>>::dyn_from_tuples(&self.factories.errors_factory, (), &mut errors);
992        self.error_stream_val.put(errors);
993
994        let result = builder.done();
995        self.output_batch_stats.add_batch(result.len());
996        result
997    }
998
999    fn input_preference(
1000        &self,
1001    ) -> (
1002        OwnershipPreference,
1003        OwnershipPreference,
1004        OwnershipPreference,
1005    ) {
1006        (
1007            OwnershipPreference::PREFER_OWNED,
1008            OwnershipPreference::PREFER_OWNED,
1009            OwnershipPreference::INDIFFERENT,
1010        )
1011    }
1012}