Skip to main content

mito2/
region_write_ctx.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::mem;
16use std::sync::Arc;
17use std::sync::atomic::{AtomicU64, Ordering};
18
19use api::v1::{BulkWalEntry, Mutation, OpType, Rows, WalEntry, WriteHint};
20use futures::stream::{FuturesUnordered, StreamExt};
21use snafu::ResultExt;
22use store_api::logstore::LogStore;
23use store_api::logstore::provider::Provider;
24use store_api::storage::{RegionId, SequenceNumber};
25
26use crate::error::{Error, Result, WriteGroupSnafu};
27use crate::memtable::KeyValues;
28use crate::memtable::bulk::part::BulkPart;
29use crate::metrics;
30use crate::region::version::{VersionControlData, VersionControlRef, VersionRef};
31use crate::request::OptionOutputTx;
32use crate::wal::{EntryId, WalWriter};
33
34/// Notifier to notify write result on drop.
35struct WriteNotify {
36    /// Error to send to the waiter.
37    err: Option<Arc<Error>>,
38    /// Sender to send write result to the waiter for this mutation.
39    sender: OptionOutputTx,
40    /// Number of rows to be written.
41    num_rows: usize,
42}
43
44impl WriteNotify {
45    /// Creates a new notify from the `sender`.
46    fn new(sender: OptionOutputTx, num_rows: usize) -> WriteNotify {
47        WriteNotify {
48            err: None,
49            sender,
50            num_rows,
51        }
52    }
53
54    /// Send result to the waiter.
55    fn notify_result(&mut self) {
56        if let Some(err) = &self.err {
57            // Try to send the error to waiters.
58            self.sender
59                .send_mut(Err(err.clone()).context(WriteGroupSnafu));
60        } else {
61            // Send success result.
62            self.sender.send_mut(Ok(self.num_rows));
63        }
64    }
65}
66
67impl Drop for WriteNotify {
68    fn drop(&mut self) {
69        self.notify_result();
70    }
71}
72
73/// Context to keep region metadata and buffer write requests.
74pub(crate) struct RegionWriteCtx {
75    /// Id of region to write.
76    region_id: RegionId,
77    /// Version of the region while creating the context.
78    version: VersionRef,
79    /// VersionControl of the region.
80    version_control: VersionControlRef,
81    /// Next sequence number to write.
82    ///
83    /// The context assigns a unique sequence number for each row.
84    next_sequence: SequenceNumber,
85    /// Next entry id of WAL to write.
86    next_entry_id: EntryId,
87    /// Valid WAL entry to write.
88    ///
89    /// We keep [WalEntry] instead of mutations to avoid taking mutations
90    /// out of the context to construct the wal entry when we write to the wal.
91    wal_entry: WalEntry,
92    /// Mutations that skip WAL, paired with their write notifiers.
93    memtable_mutations: Vec<(Mutation, WriteNotify)>,
94    /// Wal options of the region being written to.
95    provider: Provider,
96    /// Notifiers to send write results to waiters.
97    ///
98    /// The i-th notify is for the i-th mutation in `wal_entry`.
99    wal_notifiers: Vec<WriteNotify>,
100    /// Notifiers for bulk requests.
101    bulk_notifiers: Vec<WriteNotify>,
102    /// Pending bulk write requests
103    pub(crate) bulk_parts: Vec<BulkPart>,
104    /// The write operation is failed and we should not write to the mutable memtable.
105    failed: bool,
106
107    // Metrics:
108    /// Rows to put.
109    pub(crate) put_num: usize,
110    /// Rows to delete.
111    pub(crate) delete_num: usize,
112    /// The total bytes written to the region.
113    pub(crate) written_bytes: Option<Arc<AtomicU64>>,
114}
115
116impl RegionWriteCtx {
117    /// Returns an empty context.
118    pub(crate) fn new(
119        region_id: RegionId,
120        version_control: &VersionControlRef,
121        provider: Provider,
122        written_bytes: Option<Arc<AtomicU64>>,
123    ) -> RegionWriteCtx {
124        let VersionControlData {
125            version,
126            committed_sequence,
127            last_entry_id,
128            ..
129        } = version_control.current();
130
131        RegionWriteCtx {
132            region_id,
133            version,
134            version_control: version_control.clone(),
135            next_sequence: committed_sequence + 1,
136            next_entry_id: last_entry_id + 1,
137            wal_entry: WalEntry::default(),
138            memtable_mutations: Vec::new(),
139            provider,
140            wal_notifiers: Vec::new(),
141            bulk_notifiers: vec![],
142            failed: false,
143            put_num: 0,
144            delete_num: 0,
145            bulk_parts: vec![],
146            written_bytes,
147        }
148    }
149
150    /// Push mutation to the context.
151    /// This method adopts the sequence number in parameters if present.
152    pub(crate) fn push_mutation(
153        &mut self,
154        op_type: i32,
155        rows: Option<Rows>,
156        write_hint: Option<WriteHint>,
157        tx: OptionOutputTx,
158        sequence: Option<SequenceNumber>,
159        skip_wal: bool,
160    ) {
161        if let Some(sequence) = sequence {
162            self.next_sequence = sequence;
163        }
164        let num_rows = rows.as_ref().map(|rows| rows.rows.len()).unwrap_or(0);
165        let mutation = Mutation {
166            op_type,
167            sequence: self.next_sequence,
168            rows,
169            write_hint,
170        };
171
172        // Assign sequences before routing so concurrent memtable writes retain
173        // their logical order regardless of the WAL policy.
174        let notify = WriteNotify::new(tx, num_rows);
175        if skip_wal {
176            self.memtable_mutations.push((mutation, notify));
177        } else {
178            self.wal_entry.mutations.push(mutation);
179            self.wal_notifiers.push(notify);
180        }
181
182        // Increase sequence number.
183        self.next_sequence += num_rows as u64;
184
185        // Update metrics.
186        match OpType::try_from(op_type) {
187            Ok(OpType::Delete) => self.delete_num += num_rows,
188            Ok(OpType::Put) => self.put_num += num_rows,
189            Err(_) => (),
190        }
191    }
192
193    /// Encode and add WAL entry to the writer.
194    pub(crate) fn add_wal_entry<S: LogStore>(
195        &mut self,
196        wal_writer: &mut WalWriter<S>,
197    ) -> Result<()> {
198        wal_writer.add_entry(
199            self.region_id,
200            self.next_entry_id,
201            &self.wal_entry,
202            &self.provider,
203        )?;
204        self.next_entry_id += 1;
205        Ok(())
206    }
207
208    pub(crate) fn version(&self) -> &VersionRef {
209        &self.version
210    }
211
212    #[cfg(test)]
213    pub(crate) fn version_control(&self) -> &VersionControlRef {
214        &self.version_control
215    }
216
217    /// Returns whether writes in this context should skip WAL.
218    pub(crate) fn skip_wal(&self) -> bool {
219        self.provider == Provider::Noop
220            || self.version.options.skip_wal
221            || (self.wal_entry.mutations.is_empty() && self.wal_entry.bulk_entries.is_empty())
222    }
223
224    /// Sets error and marks all write operations are failed.
225    pub(crate) fn set_error(&mut self, err: Arc<Error>) {
226        // Set error for all notifiers.
227        for notify in self
228            .wal_notifiers
229            .iter_mut()
230            .chain(self.memtable_mutations.iter_mut().map(|(_, notify)| notify))
231        {
232            notify.err = Some(err.clone());
233        }
234        for notify in &mut self.bulk_notifiers {
235            notify.err = Some(err.clone());
236        }
237
238        // Fail the whole write operation.
239        self.failed = true;
240    }
241
242    /// Returns whether the write operation is already marked as failed.
243    pub(crate) fn is_failed(&self) -> bool {
244        self.failed
245    }
246
247    /// Updates next entry id.
248    pub(crate) fn set_next_entry_id(&mut self, next_entry_id: EntryId) {
249        self.next_entry_id = next_entry_id
250    }
251
252    /// Returns the next entry id to write.
253    #[cfg(test)]
254    pub(crate) fn next_entry_id(&self) -> EntryId {
255        self.next_entry_id
256    }
257
258    /// Consumes mutations and writes them into mutable memtable.
259    pub(crate) async fn write_memtable(&mut self) {
260        debug_assert_eq!(self.wal_notifiers.len(), self.wal_entry.mutations.len());
261
262        if self.failed {
263            return;
264        }
265
266        let mutable_memtable = self.version.memtables.mutable.clone();
267        let prev_memory_usage = if self.written_bytes.is_some() {
268            Some(mutable_memtable.memory_usage())
269        } else {
270            None
271        };
272
273        let mut mutations = mem::take(&mut self.wal_entry.mutations)
274            .into_iter()
275            .zip(&mut self.wal_notifiers)
276            .chain(
277                self.memtable_mutations
278                    .iter_mut()
279                    // Keep notifiers in the context until all writes complete.
280                    .map(|(mutation, notify)| (mem::take(mutation), notify)),
281            )
282            .filter_map(|(mutation, notify)| {
283                let kvs = KeyValues::new(&self.version.metadata, mutation)?;
284                Some((notify, kvs))
285            })
286            .collect::<Vec<_>>();
287
288        if mutations.len() == 1 {
289            if let Err(err) = mutable_memtable.write(&mutations[0].1) {
290                mutations[0].0.err = Some(Arc::new(err));
291            }
292        } else {
293            let mut tasks = FuturesUnordered::new();
294            for (notify, kvs) in mutations {
295                let mutable = mutable_memtable.clone();
296                // use tokio runtime to schedule tasks.
297                let task = common_runtime::spawn_blocking_global(move || mutable.write(&kvs));
298                tasks.push(async move { (notify, task.await) });
299            }
300
301            while let Some((notify, result)) = tasks.next().await {
302                // First unwrap the result from `spawn` above.
303                if let Err(err) = result.unwrap() {
304                    notify.err = Some(Arc::new(err));
305                }
306            }
307        }
308
309        if let Some(written_bytes) = &self.written_bytes {
310            let new_memory_usage = mutable_memtable.memory_usage();
311            let bytes = new_memory_usage.saturating_sub(prev_memory_usage.unwrap_or_default());
312            written_bytes.fetch_add(bytes as u64, Ordering::Relaxed);
313        }
314    }
315
316    pub(crate) fn push_bulk(
317        &mut self,
318        sender: OptionOutputTx,
319        mut bulk: BulkPart,
320        sequence: Option<SequenceNumber>,
321        skip_wal: bool,
322    ) -> bool {
323        if let Some(sequence) = sequence {
324            self.next_sequence = sequence;
325        }
326        bulk.sequence = self.next_sequence;
327        bulk.min_sequence = self.next_sequence;
328        if !skip_wal {
329            let entry = match BulkWalEntry::try_from(&bulk) {
330                Ok(entry) => entry,
331                Err(e) => {
332                    sender.send(Err(e));
333                    return false;
334                }
335            };
336            self.wal_entry.bulk_entries.push(entry);
337        }
338
339        self.bulk_notifiers
340            .push(WriteNotify::new(sender, bulk.num_rows()));
341
342        self.next_sequence += bulk.num_rows() as u64;
343        self.bulk_parts.push(bulk);
344        true
345    }
346
347    pub(crate) async fn write_bulk(&mut self) {
348        if self.failed || self.bulk_parts.is_empty() {
349            return;
350        }
351        #[cfg(test)]
352        test_hooks::pause_before_bulk_install(self.region_id, &self.version_control).await;
353        let _timer = metrics::REGION_WORKER_HANDLE_WRITE_ELAPSED
354            .with_label_values(&["write_bulk"])
355            .start_timer();
356
357        let mutable_memtable = &self.version.memtables.mutable;
358        let prev_memory_usage = if self.written_bytes.is_some() {
359            Some(mutable_memtable.memory_usage())
360        } else {
361            None
362        };
363
364        if self.bulk_parts.len() == 1 {
365            let part = self.bulk_parts.swap_remove(0);
366            let num_rows = part.num_rows();
367            if let Err(e) = self.version.memtables.mutable.write_bulk(part) {
368                self.bulk_notifiers[0].err = Some(Arc::new(e));
369            } else {
370                self.put_num += num_rows;
371            }
372            return;
373        }
374
375        let mut tasks = FuturesUnordered::new();
376        for (i, part) in self.bulk_parts.drain(..).enumerate() {
377            let mutable = mutable_memtable.clone();
378            tasks.push(common_runtime::spawn_blocking_global(move || {
379                let num_rows = part.num_rows();
380                (i, mutable.write_bulk(part), num_rows)
381            }));
382        }
383        while let Some(result) = tasks.next().await {
384            // first unwrap the result from `spawn` above
385            let (i, result, num_rows) = result.unwrap();
386            if let Err(err) = result {
387                self.bulk_notifiers[i].err = Some(Arc::new(err));
388            } else {
389                self.put_num += num_rows;
390            }
391        }
392
393        if let Some(written_bytes) = &self.written_bytes {
394            let new_memory_usage = mutable_memtable.memory_usage();
395            let bytes = new_memory_usage.saturating_sub(prev_memory_usage.unwrap_or_default());
396            written_bytes.fetch_add(bytes as u64, Ordering::Relaxed);
397        }
398    }
399
400    /// Publishes the assigned sequences and entry id to the region's committed
401    /// watermark. Only call after both [`write_memtable`](Self::write_memtable)
402    /// and [`write_bulk`](Self::write_bulk) have completed; a failed context
403    /// must not publish.
404    pub(crate) fn publish_sequence_and_entry_id(&self) {
405        if self.failed {
406            return;
407        }
408        self.version_control
409            .set_sequence_and_entry_id(self.next_sequence - 1, self.next_entry_id - 1);
410    }
411}
412
413/// Test-only hooks to make write ordering races deterministic.
414#[cfg(test)]
415pub(crate) mod test_hooks {
416    use std::sync::Mutex;
417    use std::sync::atomic::{AtomicU64, Ordering};
418
419    use store_api::storage::RegionId;
420    use tokio::sync::watch;
421
422    use crate::region::version::VersionControlRef;
423
424    /// Channels of an armed bulk-install barrier; dropping the senders (by
425    /// disarming) unblocks writes paused on it.
426    struct ActiveBarrier {
427        id: u64,
428        /// Only bulk writes for this region on this version control (Arc
429        /// identity) pause at the barrier.
430        target_region_id: RegionId,
431        target_version_control: VersionControlRef,
432        reached: watch::Sender<bool>,
433        release: watch::Sender<bool>,
434    }
435
436    static ACTIVE_BARRIER: Mutex<Option<ActiveBarrier>> = Mutex::new(None);
437    static NEXT_BARRIER_ID: AtomicU64 = AtomicU64::new(1);
438
439    fn lock_active_barrier() -> std::sync::MutexGuard<'static, Option<ActiveBarrier>> {
440        // Never let a poisoned mutex (e.g. a panic in another test while
441        // holding the lock) hang or break unrelated tests.
442        ACTIVE_BARRIER
443            .lock()
444            .unwrap_or_else(|poisoned| poisoned.into_inner())
445    }
446
447    /// RAII guard: releasing (or dropping) it unblocks a paused write and
448    /// disarms the barrier.
449    pub(crate) struct BulkInstallBarrier {
450        id: u64,
451        reached_rx: watch::Receiver<bool>,
452        release_tx: watch::Sender<bool>,
453        released: bool,
454    }
455
456    impl BulkInstallBarrier {
457        pub(crate) async fn wait_until_reached(&mut self) {
458            if !*self.reached_rx.borrow() {
459                let _ = self.reached_rx.wait_for(|reached| *reached).await;
460            }
461        }
462
463        pub(crate) fn release(&mut self) {
464            if self.released {
465                return;
466            }
467            self.released = true;
468            let _ = self.release_tx.send(true);
469            disarm_barrier(self.id);
470        }
471    }
472
473    impl Drop for BulkInstallBarrier {
474        fn drop(&mut self) {
475            self.release();
476        }
477    }
478
479    /// Arms the bulk-install barrier and returns the owning guard; any
480    /// previously armed barrier is replaced.
481    pub(crate) fn arm_bulk_install_barrier(
482        target_region_id: RegionId,
483        target_version_control: VersionControlRef,
484    ) -> BulkInstallBarrier {
485        let (reached_tx, reached_rx) = watch::channel(false);
486        let (release_tx, _release_rx) = watch::channel(false);
487        let id = NEXT_BARRIER_ID.fetch_add(1, Ordering::Relaxed);
488        let mut active = lock_active_barrier();
489        *active = Some(ActiveBarrier {
490            id,
491            target_region_id,
492            target_version_control,
493            reached: reached_tx,
494            release: release_tx.clone(),
495        });
496        BulkInstallBarrier {
497            id,
498            reached_rx,
499            release_tx,
500            released: false,
501        }
502    }
503
504    fn disarm_barrier(id: u64) {
505        let mut active = lock_active_barrier();
506        if active.as_ref().is_some_and(|barrier| barrier.id == id) {
507            *active = None;
508        }
509    }
510
511    /// Pauses a bulk write before installing its parts until the barrier is
512    /// released or disarmed.
513    pub(crate) async fn pause_before_bulk_install(
514        region_id: RegionId,
515        version_control: &VersionControlRef,
516    ) {
517        let (reached_tx, release_rx) = {
518            let active = lock_active_barrier();
519            match active.as_ref() {
520                Some(barrier)
521                    if barrier.target_region_id == region_id
522                        && std::sync::Arc::ptr_eq(
523                            &barrier.target_version_control,
524                            version_control,
525                        ) =>
526                {
527                    (barrier.reached.clone(), barrier.release.subscribe())
528                }
529                _ => return,
530            }
531        };
532        let _ = reached_tx.send(true);
533        let mut release_rx = release_rx;
534        if !*release_rx.borrow() {
535            // The sender is dropped when the barrier is disarmed, which makes
536            // `wait_for` return an error instead of hanging forever.
537            let _ = release_rx.wait_for(|released| *released).await;
538        }
539    }
540}
541
542#[cfg(test)]
543mod tests {
544    use std::sync::Arc;
545
546    use common_recordbatch::DfRecordBatch;
547    use datatypes::arrow::array::{ArrayRef, TimestampMillisecondArray};
548    use datatypes::arrow::datatypes::{DataType, Field, Schema};
549    use prost::Message;
550    use store_api::logstore::provider::Provider;
551    use tokio::sync::oneshot;
552
553    use super::*;
554    use crate::error::UnexpectedSnafu;
555    use crate::memtable::bulk::part::BulkPart;
556    use crate::test_util::version_util::VersionControlBuilder;
557
558    #[test]
559    fn test_request_skip_wal_preserves_sequences_and_other_writes() {
560        // Ablate only the request flag: the workload and sequence allocation stay identical.
561        check_request_skip_wal_preserves_sequences_and_other_writes(false, false);
562        check_request_skip_wal_preserves_sequences_and_other_writes(false, true);
563        check_request_skip_wal_preserves_sequences_and_other_writes(true, false);
564        check_request_skip_wal_preserves_sequences_and_other_writes(true, true);
565    }
566
567    fn check_request_skip_wal_preserves_sequences_and_other_writes(
568        skip_wal: bool,
569        bulk_skip_wal: bool,
570    ) {
571        let builder = VersionControlBuilder::new();
572        let region_id = builder.region_id();
573        let version_control = Arc::new(builder.build());
574        let mut ctx = RegionWriteCtx::new(
575            region_id,
576            &version_control,
577            Provider::raft_engine_provider(region_id.as_u64()),
578            None,
579        );
580        for (op_type, skip, num_rows) in [
581            (OpType::Put, skip_wal, 2),
582            (OpType::Put, false, 3),
583            (OpType::Delete, false, 1),
584        ] {
585            ctx.push_mutation(
586                op_type as i32,
587                Some(Rows {
588                    schema: vec![],
589                    rows: vec![api::v1::Row::default(); num_rows],
590                }),
591                None,
592                OptionOutputTx::none(),
593                None,
594                skip,
595            );
596        }
597        // Unequal row counts make a notifier/mutation pairing mismatch visible.
598        for (mutation, notify) in ctx
599            .wal_entry
600            .mutations
601            .iter()
602            .zip(&ctx.wal_notifiers)
603            .chain(
604                ctx.memtable_mutations
605                    .iter()
606                    .map(|(mutation, notify)| (mutation, notify)),
607            )
608        {
609            assert_eq!(mutation.rows.as_ref().unwrap().rows.len(), notify.num_rows);
610        }
611        assert!(ctx.push_bulk(OptionOutputTx::none(), new_bulk_part(), None, bulk_skip_wal));
612        assert!(!ctx.skip_wal());
613        assert_eq!(ctx.next_sequence, 9);
614        assert_eq!(
615            ctx.wal_entry.bulk_entries.len(),
616            usize::from(!bulk_skip_wal)
617        );
618        assert_eq!(ctx.bulk_parts[0].sequence, 7);
619        assert_eq!(ctx.bulk_parts[0].min_sequence, 7);
620        let sequences: Vec<_> = ctx.wal_entry.mutations.iter().map(|m| m.sequence).collect();
621        assert_eq!(sequences, if skip_wal { vec![3, 6] } else { vec![1, 3, 6] });
622        // Check the actual WAL bytes after routing the mutations.
623        let encoded = crate::wal::encoder::WalEntryEncoder::new().encode_to_vec(&ctx.wal_entry);
624        let decoded = WalEntry::decode(encoded.as_slice()).unwrap();
625        assert_eq!(
626            decoded
627                .mutations
628                .iter()
629                .map(|m| m.sequence)
630                .collect::<Vec<_>>(),
631            sequences
632        );
633        assert_eq!(decoded.bulk_entries, ctx.wal_entry.bulk_entries);
634        assert_eq!(ctx.wal_entry.mutations.len(), if skip_wal { 2 } else { 3 });
635        assert_eq!(ctx.memtable_mutations.len(), usize::from(skip_wal));
636        if skip_wal {
637            assert_eq!(ctx.memtable_mutations[0].0.sequence, 1);
638        }
639        assert_eq!(
640            ctx.wal_entry.mutations.last().unwrap().op_type,
641            OpType::Delete as i32
642        );
643    }
644
645    #[test]
646    fn test_internal_delete_respects_skip_wal_flag() {
647        check_internal_delete_respects_skip_wal_flag(false);
648        check_internal_delete_respects_skip_wal_flag(true);
649    }
650
651    fn check_internal_delete_respects_skip_wal_flag(skip_wal: bool) {
652        let builder = VersionControlBuilder::new();
653        let region_id = builder.region_id();
654        let version_control = Arc::new(builder.build());
655        let mut ctx = RegionWriteCtx::new(
656            region_id,
657            &version_control,
658            Provider::raft_engine_provider(region_id.as_u64()),
659            None,
660        );
661        ctx.push_mutation(
662            OpType::Delete as i32,
663            Some(Rows {
664                schema: vec![],
665                rows: vec![api::v1::Row::default(); 2],
666            }),
667            None,
668            OptionOutputTx::none(),
669            None,
670            skip_wal,
671        );
672        assert_eq!(ctx.skip_wal(), skip_wal);
673        assert_eq!(ctx.next_sequence, 3);
674        assert_eq!(ctx.delete_num, 2);
675        assert_eq!(ctx.wal_entry.mutations.len(), usize::from(!skip_wal));
676        assert_eq!(ctx.memtable_mutations.len(), usize::from(skip_wal));
677        let encoded = crate::wal::encoder::WalEntryEncoder::new().encode_to_vec(&ctx.wal_entry);
678        let decoded = WalEntry::decode(encoded.as_slice()).unwrap();
679        if skip_wal {
680            assert!(decoded.mutations.is_empty());
681        } else {
682            assert_eq!(decoded, ctx.wal_entry);
683        }
684    }
685
686    #[test]
687    fn test_all_request_skip_wal_keeps_entry_id_and_propagates_errors() {
688        let builder = VersionControlBuilder::new();
689        let region_id = builder.region_id();
690        let version_control = Arc::new(builder.build());
691        let mut ctx = RegionWriteCtx::new(
692            region_id,
693            &version_control,
694            Provider::raft_engine_provider(region_id.as_u64()),
695            None,
696        );
697        let (tx, rx) = oneshot::channel();
698        ctx.push_mutation(
699            OpType::Put as i32,
700            Some(Rows {
701                schema: vec![],
702                rows: vec![api::v1::Row::default(); 2],
703            }),
704            None,
705            OptionOutputTx::from(tx),
706            None,
707            true,
708        );
709        assert!(ctx.skip_wal());
710        assert!(ctx.wal_entry.mutations.is_empty());
711        assert_eq!(ctx.memtable_mutations.len(), 1);
712        assert_eq!(ctx.next_entry_id(), 1);
713        assert_eq!(ctx.next_sequence, 3);
714        ctx.set_error(Arc::new(
715            UnexpectedSnafu {
716                reason: "wal failed".to_string(),
717            }
718            .build(),
719        ));
720        drop(ctx);
721        assert!(rx.blocking_recv().unwrap().is_err());
722    }
723
724    #[test]
725    fn test_bulk_skip_wal_preserves_sequences() {
726        check_bulk_skip_wal_preserves_sequences(false);
727        check_bulk_skip_wal_preserves_sequences(true);
728    }
729
730    fn check_bulk_skip_wal_preserves_sequences(skip_wal: bool) {
731        let builder = VersionControlBuilder::new();
732        let region_id = builder.region_id();
733        let version_control = Arc::new(builder.build());
734        let mut ctx = RegionWriteCtx::new(
735            region_id,
736            &version_control,
737            Provider::raft_engine_provider(region_id.as_u64()),
738            None,
739        );
740        // Alternate policies in one context, retaining all parts for the memtable.
741        for skip in [skip_wal, false, skip_wal] {
742            assert!(ctx.push_bulk(OptionOutputTx::none(), new_bulk_part(), None, skip));
743        }
744        assert_eq!(ctx.next_sequence, 7);
745        assert_eq!(ctx.bulk_notifiers.len(), 3);
746        assert_eq!(
747            ctx.bulk_parts
748                .iter()
749                .map(|p| p.sequence)
750                .collect::<Vec<_>>(),
751            vec![1, 3, 5]
752        );
753        let encoded = crate::wal::encoder::WalEntryEncoder::new().encode_to_vec(&ctx.wal_entry);
754        let decoded = WalEntry::decode(encoded.as_slice()).unwrap();
755        assert_eq!(
756            decoded
757                .bulk_entries
758                .iter()
759                .map(|e| e.sequence)
760                .collect::<Vec<_>>(),
761            if skip_wal { vec![3] } else { vec![1, 3, 5] }
762        );
763    }
764
765    #[test]
766    fn test_set_error_marks_bulk_notifiers_failed() {
767        check_set_error_marks_bulk_notifiers_failed(false);
768        check_set_error_marks_bulk_notifiers_failed(true);
769    }
770
771    fn check_set_error_marks_bulk_notifiers_failed(skip_wal: bool) {
772        let builder = VersionControlBuilder::new();
773        let region_id = builder.region_id();
774        let version_control = Arc::new(builder.build());
775        let mut ctx =
776            RegionWriteCtx::new(region_id, &version_control, Provider::noop_provider(), None);
777        let (tx, rx) = oneshot::channel();
778
779        assert!(ctx.push_bulk(OptionOutputTx::from(tx), new_bulk_part(), None, skip_wal));
780        assert_eq!(ctx.wal_entry.bulk_entries.len(), usize::from(!skip_wal));
781        assert_eq!(ctx.bulk_parts.len(), 1);
782        ctx.set_error(Arc::new(
783            UnexpectedSnafu {
784                reason: "wal failed".to_string(),
785            }
786            .build(),
787        ));
788        drop(ctx);
789
790        let result = rx.blocking_recv().unwrap();
791        assert!(result.is_err(), "bulk notifier should report WAL error");
792    }
793
794    fn new_bulk_part() -> BulkPart {
795        let schema = Arc::new(Schema::new(vec![Field::new(
796            "ts",
797            DataType::Timestamp(datatypes::arrow::datatypes::TimeUnit::Millisecond, None),
798            false,
799        )]));
800        let arrays = vec![Arc::new(TimestampMillisecondArray::from(vec![1, 2])) as ArrayRef];
801        let batch = DfRecordBatch::try_new(schema, arrays).unwrap();
802
803        BulkPart {
804            batch,
805            max_timestamp: 2,
806            min_timestamp: 1,
807            sequence: 0,
808            min_sequence: 0,
809            timestamp_index: 0,
810            raw_data: None,
811        }
812    }
813}