Skip to main content

mito2/worker/
handle_write.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
15//! Handling write requests.
16
17use std::collections::{HashMap, HashSet, hash_map};
18use std::sync::Arc;
19
20use api::v1::OpType;
21use common_telemetry::{debug, error};
22use snafu::ensure;
23use store_api::codec::PrimaryKeyEncoding;
24use store_api::logstore::LogStore;
25use store_api::storage::RegionId;
26
27use crate::error::{
28    InvalidRequestSnafu, PartitionExprVersionMismatchSnafu, RegionNotFoundSnafu, RegionStateSnafu,
29    RejectWriteSnafu, Result,
30};
31use crate::metrics;
32use crate::metrics::{
33    WRITE_REJECT_TOTAL, WRITE_ROWS_TOTAL, WRITE_STAGE_ELAPSED, WRITE_STALL_TOTAL,
34};
35use crate::region::{RegionLeaderState, RegionRoleState};
36use crate::region_write_ctx::RegionWriteCtx;
37use crate::request::{SenderBulkRequest, SenderWriteRequest, WriteRequest};
38use crate::wal::Wal;
39use crate::worker::RegionWorkerLoop;
40
41impl<S: LogStore> RegionWorkerLoop<S> {
42    /// Takes and handles all write requests.
43    pub(crate) async fn handle_write_requests(
44        &mut self,
45        write_requests: &mut Vec<SenderWriteRequest>,
46        bulk_requests: &mut Vec<SenderBulkRequest>,
47        allow_stall: bool,
48    ) {
49        if write_requests.is_empty() && bulk_requests.is_empty() {
50            return;
51        }
52
53        let write_region_ids = write_region_ids(write_requests, bulk_requests);
54
55        // Check region pressure before writes to match the global write buffer behavior.
56        self.maybe_flush_worker();
57        let pressure = self.maybe_flush_write_regions(write_region_ids);
58
59        if self.should_reject_write() {
60            // The memory pressure is still too high, reject write requests.
61            reject_write_requests(write_requests, bulk_requests);
62            // Also reject all stalled requests.
63            self.reject_stalled_requests();
64            return;
65        }
66
67        if !pressure.rejected_region_ids.is_empty() {
68            reject_region_write_requests(
69                &pressure.rejected_region_ids,
70                write_requests,
71                bulk_requests,
72            );
73            for region_id in &pressure.rejected_region_ids {
74                self.reject_region_stalled_requests(region_id);
75            }
76            if write_requests.is_empty() && bulk_requests.is_empty() {
77                return;
78            }
79        }
80
81        if self.write_buffer_manager.should_stall() && allow_stall {
82            let stalled_count = (write_requests.len() + bulk_requests.len()) as i64;
83            self.stalling_count.add(stalled_count);
84            WRITE_STALL_TOTAL.inc_by(stalled_count as u64);
85            self.stalled_requests.append(write_requests, bulk_requests);
86            self.listener.on_write_stall();
87            return;
88        }
89
90        if allow_stall {
91            self.stall_region_write_requests(
92                &pressure.stalled_region_ids,
93                write_requests,
94                bulk_requests,
95            );
96            if write_requests.is_empty() && bulk_requests.is_empty() {
97                return;
98            }
99        }
100
101        // Prepare write context.
102        let mut region_ctxs = {
103            let _timer = WRITE_STAGE_ELAPSED
104                .with_label_values(&["prepare_ctx"])
105                .start_timer();
106            self.prepare_region_write_ctx(write_requests, bulk_requests)
107        };
108
109        // Write to WAL.
110        {
111            let _timer = WRITE_STAGE_ELAPSED
112                .with_label_values(&["write_wal"])
113                .start_timer();
114            if !write_wal(&self.wal, &mut region_ctxs).await {
115                // Failed to write to the WAL, all waiters are notified with the error.
116                return;
117            }
118        }
119
120        let (mut put_rows, mut delete_rows) = (0, 0);
121        // Write to memtables.
122        {
123            let _timer = WRITE_STAGE_ELAPSED
124                .with_label_values(&["write_memtable"])
125                .start_timer();
126            if region_ctxs.len() == 1 {
127                // fast path for single region.
128                let mut region_ctx = region_ctxs.into_values().next().unwrap();
129                region_ctx.write_memtable().await;
130                region_ctx.write_bulk().await;
131                put_rows += region_ctx.put_num;
132                delete_rows += region_ctx.delete_num;
133            } else {
134                let region_write_task = region_ctxs
135                    .into_values()
136                    .map(|mut region_ctx| {
137                        // use tokio runtime to schedule tasks.
138                        common_runtime::spawn_global(async move {
139                            region_ctx.write_memtable().await;
140                            region_ctx.write_bulk().await;
141                            (region_ctx.put_num, region_ctx.delete_num)
142                        })
143                    })
144                    .collect::<Vec<_>>();
145
146                for result in futures::future::join_all(region_write_task).await {
147                    match result {
148                        Ok((put, delete)) => {
149                            put_rows += put;
150                            delete_rows += delete;
151                        }
152                        Err(e) => {
153                            error!(e; "unexpected error when joining region write tasks");
154                        }
155                    }
156                }
157            }
158        }
159        WRITE_ROWS_TOTAL
160            .with_label_values(&["put"])
161            .inc_by(put_rows as u64);
162        WRITE_ROWS_TOTAL
163            .with_label_values(&["delete"])
164            .inc_by(delete_rows as u64);
165    }
166
167    /// Handles stalled write requests whose regions no longer need to stall.
168    pub(crate) async fn handle_stalled_requests(&mut self) {
169        let region_ids = self
170            .stalled_requests
171            .requests
172            .keys()
173            .copied()
174            .collect::<HashSet<_>>();
175        let pressure = self.maybe_flush_write_regions(region_ids);
176        for region_id in &pressure.rejected_region_ids {
177            self.reject_region_stalled_requests(region_id);
178        }
179        let ready_region_ids = self
180            .stalled_requests
181            .requests
182            .keys()
183            .filter(|region_id| !pressure.stalled_region_ids.contains(region_id))
184            .copied()
185            .collect::<Vec<_>>();
186
187        // These requests have already been stalled. Retry ready regions without stalling the
188        // same requests again. Regions that still exceed their limit remain in the queue until
189        // their own flush releases the pressure.
190        for region_id in ready_region_ids {
191            self.handle_region_stalled_requests(&region_id, false).await;
192        }
193    }
194
195    /// Rejects all stalled requests.
196    pub(crate) fn reject_stalled_requests(&mut self) {
197        let stalled = std::mem::take(&mut self.stalled_requests);
198        self.stalling_count.sub(stalled.stalled_count() as i64);
199        for (_, (_, mut requests, mut bulk)) in stalled.requests {
200            reject_write_requests(&mut requests, &mut bulk);
201        }
202    }
203
204    /// Rejects a specific region's stalled requests.
205    pub(crate) fn reject_region_stalled_requests(&mut self, region_id: &RegionId) {
206        debug!("Rejects stalled requests for region {}", region_id);
207        let (mut requests, mut bulk) = self.stalled_requests.remove(region_id);
208        self.stalling_count
209            .sub((requests.len() + bulk.len()) as i64);
210        reject_write_requests(&mut requests, &mut bulk);
211    }
212
213    /// Fails a specific region's stalled requests if the region no longer exists.
214    pub(crate) fn fail_region_stalled_requests_as_not_found(&mut self, region_id: &RegionId) {
215        debug!(
216            "Fails stalled requests for region {} as region not found",
217            region_id
218        );
219        let (requests, bulk) = self.stalled_requests.remove(region_id);
220        self.stalling_count
221            .sub((requests.len() + bulk.len()) as i64);
222
223        for req in requests {
224            req.sender.send(
225                RegionNotFoundSnafu {
226                    region_id: req.request.region_id,
227                }
228                .fail(),
229            );
230        }
231        for req in bulk {
232            req.sender.send(
233                RegionNotFoundSnafu {
234                    region_id: req.region_id,
235                }
236                .fail(),
237            );
238        }
239    }
240
241    /// Handles a specific region's stalled requests.
242    ///
243    /// `allow_stall` should be false for backpressure retry paths to avoid stalling the same
244    /// requests again. It should remain true for non-backpressure retries, such as requests stalled
245    /// by alter, staging, and region editing. Global reject backpressure still applies before the
246    /// stall check.
247    pub(crate) async fn handle_region_stalled_requests(
248        &mut self,
249        region_id: &RegionId,
250        allow_stall: bool,
251    ) {
252        debug!("Handles stalled requests for region {}", region_id);
253        let (mut requests, mut bulk) = self.stalled_requests.remove(region_id);
254        self.stalling_count
255            .sub((requests.len() + bulk.len()) as i64);
256        self.handle_write_requests(&mut requests, &mut bulk, allow_stall)
257            .await;
258    }
259
260    /// Processes same-batch writes for a region before handling its edit-completion notification.
261    ///
262    /// The worker dispatch loop handles background notifications before the current batch's write
263    /// buffer. Without this step, writes that arrived during edit N could be classified only after
264    /// edit N+1 is started, placing them behind that next edit.
265    pub(crate) async fn handle_buffered_region_write_requests(
266        &mut self,
267        region_id: &RegionId,
268        write_requests: &mut Vec<SenderWriteRequest>,
269        bulk_requests: &mut Vec<SenderBulkRequest>,
270    ) {
271        let mut current_region_write_requests = write_requests
272            .extract_if(.., |r| r.request.region_id == *region_id)
273            .collect::<Vec<_>>();
274
275        let mut current_region_bulk_requests = bulk_requests
276            .extract_if(.., |r| r.region_id == *region_id)
277            .collect::<Vec<_>>();
278
279        self.handle_write_requests(
280            &mut current_region_write_requests,
281            &mut current_region_bulk_requests,
282            true,
283        )
284        .await;
285    }
286}
287
288impl<S> RegionWorkerLoop<S> {
289    /// Validates and groups requests by region.
290    fn prepare_region_write_ctx(
291        &mut self,
292        write_requests: &mut Vec<SenderWriteRequest>,
293        bulk_requests: &mut Vec<SenderBulkRequest>,
294    ) -> HashMap<RegionId, RegionWriteCtx> {
295        // Initialize region write context map.
296        let mut region_ctxs = HashMap::new();
297        self.process_write_requests(&mut region_ctxs, write_requests);
298        self.process_bulk_requests(&mut region_ctxs, bulk_requests);
299        region_ctxs
300    }
301
302    fn process_write_requests(
303        &mut self,
304        region_ctxs: &mut HashMap<RegionId, RegionWriteCtx>,
305        write_requests: &mut Vec<SenderWriteRequest>,
306    ) {
307        for mut sender_req in write_requests.drain(..) {
308            let region_id = sender_req.request.region_id;
309
310            // If region is waiting for alteration, add requests to pending writes.
311            if self.flush_scheduler.has_pending_ddls(region_id) {
312                // TODO(yingwen): consider adding some metrics for this.
313                // Safety: The region has pending ddls.
314                self.flush_scheduler
315                    .add_write_request_to_pending(sender_req);
316                continue;
317            }
318
319            // Checks whether the region exists and is it stalling.
320            if let hash_map::Entry::Vacant(e) = region_ctxs.entry(region_id) {
321                let Some(region) = self
322                    .regions
323                    .get_region_or(region_id, &mut sender_req.sender)
324                else {
325                    // No such region.
326                    continue;
327                };
328                #[cfg(test)]
329                debug!(
330                    "Handling write request for region {}, state: {:?}",
331                    region_id,
332                    region.state()
333                );
334                match region.state() {
335                    RegionRoleState::Leader(RegionLeaderState::Writable)
336                    | RegionRoleState::Leader(RegionLeaderState::Staging) => {
337                        if region.reject_all_writes_in_staging() {
338                            sender_req
339                                .sender
340                                .send(RejectWriteSnafu { region_id }.fail());
341                            continue;
342                        }
343
344                        let region_ctx = RegionWriteCtx::new(
345                            region.region_id,
346                            &region.version_control,
347                            region.provider.clone(),
348                            Some(region.region_stats.written_bytes.clone()),
349                        );
350
351                        e.insert(region_ctx);
352                    }
353                    RegionRoleState::Leader(RegionLeaderState::Altering)
354                    | RegionRoleState::Leader(RegionLeaderState::Editing) => {
355                        // Editing is transient: queue the write so edit completion can drain it
356                        // before starting the next queued edit.
357                        debug!(
358                            "Region {} is {:?}, add request to pending writes",
359                            region.region_id,
360                            region.state()
361                        );
362                        self.stalling_count.add(1);
363                        WRITE_STALL_TOTAL.inc();
364                        self.stalled_requests.push(sender_req);
365                        continue;
366                    }
367                    RegionRoleState::Leader(RegionLeaderState::EnteringStaging) => {
368                        debug!(
369                            "Region {} is entering staging, add request to pending writes",
370                            region.region_id
371                        );
372                        self.stalling_count.add(1);
373                        WRITE_STALL_TOTAL.inc();
374                        self.stalled_requests.push(sender_req);
375                        continue;
376                    }
377                    state => {
378                        // The region is not writable.
379                        sender_req.sender.send(
380                            RegionStateSnafu {
381                                region_id,
382                                state,
383                                expect: RegionRoleState::Leader(RegionLeaderState::Writable),
384                            }
385                            .fail(),
386                        );
387                        continue;
388                    }
389                }
390            }
391
392            // Safety: Now we ensure the region exists.
393            let region_ctx = region_ctxs.get_mut(&region_id).unwrap();
394            let Some(region) = self
395                .regions
396                .get_region_or(region_id, &mut sender_req.sender)
397            else {
398                continue;
399            };
400            if region.reject_all_writes_in_staging() {
401                sender_req
402                    .sender
403                    .send(RejectWriteSnafu { region_id }.fail());
404                continue;
405            }
406            let expected_version = region.expected_partition_expr_version();
407            if let Err(e) = check_partition_expr_version(
408                region_id,
409                expected_version,
410                sender_req.request.partition_expr_version,
411            ) {
412                sender_req.sender.send(Err(e));
413                continue;
414            }
415
416            if let Err(e) = check_op_type(
417                region_ctx.version().options.append_mode,
418                &sender_req.request,
419            ) {
420                // Do not allow non-put op under append mode.
421                sender_req.sender.send(Err(e));
422
423                continue;
424            }
425
426            // Double check the request schema
427            let need_fill_missing_columns =
428                if let Some(ref region_metadata) = sender_req.request.region_metadata {
429                    region_ctx.version().metadata.schema_version != region_metadata.schema_version
430                } else {
431                    true
432                };
433            // Only fill missing columns if primary key is dense encoded.
434            if need_fill_missing_columns
435                && sender_req.request.primary_key_encoding() == PrimaryKeyEncoding::Dense
436                && let Err(e) = sender_req
437                    .request
438                    .maybe_fill_missing_columns(&region_ctx.version().metadata)
439            {
440                sender_req.sender.send(Err(e));
441
442                continue;
443            }
444
445            // Collect requests by region.
446            region_ctx.push_mutation(
447                sender_req.request.op_type as i32,
448                Some(sender_req.request.rows),
449                sender_req.request.hint,
450                sender_req.sender,
451                None,
452            );
453        }
454    }
455
456    /// Processes bulk insert requests.
457    fn process_bulk_requests(
458        &mut self,
459        region_ctxs: &mut HashMap<RegionId, RegionWriteCtx>,
460        requests: &mut Vec<SenderBulkRequest>,
461    ) {
462        let _timer = metrics::REGION_WORKER_HANDLE_WRITE_ELAPSED
463            .with_label_values(&["prepare_bulk_request"])
464            .start_timer();
465        for mut bulk_req in requests.drain(..) {
466            let region_id = bulk_req.region_id;
467            // If region is waiting for alteration, add requests to pending writes.
468            if self.flush_scheduler.has_pending_ddls(region_id) {
469                // Safety: The region has pending ddls.
470                self.flush_scheduler.add_bulk_request_to_pending(bulk_req);
471                continue;
472            }
473
474            // Checks whether the region exists and is it stalling.
475            if let hash_map::Entry::Vacant(e) = region_ctxs.entry(region_id) {
476                let Some(region) = self.regions.get_region_or(region_id, &mut bulk_req.sender)
477                else {
478                    continue;
479                };
480                match region.state() {
481                    RegionRoleState::Leader(RegionLeaderState::Writable)
482                    | RegionRoleState::Leader(RegionLeaderState::Staging) => {
483                        if region.reject_all_writes_in_staging() {
484                            bulk_req.sender.send(RejectWriteSnafu { region_id }.fail());
485                            continue;
486                        }
487                        let region_ctx = RegionWriteCtx::new(
488                            region.region_id,
489                            &region.version_control,
490                            region.provider.clone(),
491                            Some(region.region_stats.written_bytes.clone()),
492                        );
493
494                        e.insert(region_ctx);
495                    }
496                    RegionRoleState::Leader(RegionLeaderState::Altering)
497                    | RegionRoleState::Leader(RegionLeaderState::Editing) => {
498                        // Editing is transient: queue the bulk write so edit completion can drain
499                        // it before starting the next queued edit.
500                        debug!(
501                            "Region {} is {:?}, add request to pending writes",
502                            region.region_id,
503                            region.state()
504                        );
505                        self.stalling_count.add(1);
506                        WRITE_STALL_TOTAL.inc();
507                        self.stalled_requests.push_bulk(bulk_req);
508                        continue;
509                    }
510                    state => {
511                        // The region is not writable.
512                        bulk_req.sender.send(
513                            RegionStateSnafu {
514                                region_id,
515                                state,
516                                expect: RegionRoleState::Leader(RegionLeaderState::Writable),
517                            }
518                            .fail(),
519                        );
520                        continue;
521                    }
522                }
523            }
524
525            // Safety: Now we ensure the region exists.
526            let region_ctx = region_ctxs.get_mut(&region_id).unwrap();
527            let Some(region) = self.regions.get_region_or(region_id, &mut bulk_req.sender) else {
528                continue;
529            };
530            if region.reject_all_writes_in_staging() {
531                bulk_req.sender.send(RejectWriteSnafu { region_id }.fail());
532                continue;
533            }
534            let expected_version = region.expected_partition_expr_version();
535            if let Err(e) = check_partition_expr_version(
536                region_id,
537                expected_version,
538                bulk_req.partition_expr_version,
539            ) {
540                bulk_req.sender.send(Err(e));
541                continue;
542            }
543
544            // Double-check the request schema
545            let need_fill_missing_columns =
546                !bulk_req.region_metadata.is_some_and(|aligned_schema| {
547                    aligned_schema.schema_version == region_ctx.version().metadata.schema_version
548                });
549
550            // Fill missing columns if needed
551            if need_fill_missing_columns
552                && let Err(e) = bulk_req
553                    .request
554                    .fill_missing_columns(&region_ctx.version().metadata)
555            {
556                bulk_req.sender.send(Err(e));
557                continue;
558            }
559
560            // Collect requests by region.
561            if !region_ctx.push_bulk(bulk_req.sender, bulk_req.request, None) {
562                return;
563            }
564        }
565    }
566
567    /// Returns true if the engine needs to reject some write requests.
568    pub(crate) fn should_reject_write(&self) -> bool {
569        // If memory usage reaches high threshold (we should also consider stalled requests) returns true.
570        self.write_buffer_manager.memory_usage() + self.stalled_requests.estimated_size
571            >= self.config.global_write_buffer_reject_size.as_bytes() as usize
572    }
573
574    fn stall_region_write_requests(
575        &mut self,
576        stalled_region_ids: &HashSet<RegionId>,
577        write_requests: &mut Vec<SenderWriteRequest>,
578        bulk_requests: &mut Vec<SenderBulkRequest>,
579    ) {
580        let mut stalled_count = 0;
581        let mut stalled_write_requests = write_requests
582            .extract_if(.., |req| {
583                stalled_region_ids.contains(&req.request.region_id)
584            })
585            .collect::<Vec<_>>();
586        let mut stalled_bulk_requests = bulk_requests
587            .extract_if(.., |req| stalled_region_ids.contains(&req.region_id))
588            .collect::<Vec<_>>();
589
590        stalled_count += stalled_write_requests.len() + stalled_bulk_requests.len();
591        self.stalled_requests
592            .append(&mut stalled_write_requests, &mut stalled_bulk_requests);
593
594        if stalled_count > 0 {
595            let stalled_count = stalled_count as i64;
596            self.stalling_count.add(stalled_count);
597            WRITE_STALL_TOTAL.inc_by(stalled_count as u64);
598            self.listener.on_write_stall();
599        }
600    }
601}
602
603/// Writes WAL entries of all region contexts to the WAL in one batch and updates
604/// the next entry id of each region on success.
605///
606/// Returns `false` if the batch fails to be written to the WAL. In this case all
607/// contexts are consumed and their waiters are notified with the error, so the
608/// caller should skip the memtable phase.
609async fn write_wal<S: LogStore>(
610    wal: &Wal<S>,
611    region_ctxs: &mut HashMap<RegionId, RegionWriteCtx>,
612) -> bool {
613    let mut wal_writer = wal.writer();
614    for region_ctx in region_ctxs.values_mut() {
615        if region_ctx.skip_wal() {
616            continue;
617        }
618        if let Err(e) = region_ctx.add_wal_entry(&mut wal_writer).map_err(Arc::new) {
619            region_ctx.set_error(e);
620        }
621    }
622    match wal_writer.write_to_wal().await.map_err(Arc::new) {
623        Ok(response) => {
624            for (region_id, region_ctx) in region_ctxs.iter_mut() {
625                if region_ctx.skip_wal() {
626                    continue;
627                }
628                // The entry of a failed region (e.g. failed to build its WAL entry) is
629                // not in the batch so the response has no last entry id for it. Its
630                // waiters are already notified with the error.
631                if region_ctx.is_failed() {
632                    continue;
633                }
634
635                // Safety: the log store implementation ensures that either the `write_to_wal` fails and no
636                // response is returned or the last entry ids for each region in the batch do exist.
637                let last_entry_id = response.last_entry_ids.get(region_id).unwrap();
638                region_ctx.set_next_entry_id(last_entry_id + 1);
639            }
640            true
641        }
642        Err(e) => {
643            // Failed to write wal.
644            for (_, mut region_ctx) in region_ctxs.drain() {
645                region_ctx.set_error(e.clone());
646            }
647            false
648        }
649    }
650}
651
652/// Send rejected error to all `write_requests`.
653fn reject_write_requests(
654    write_requests: &mut Vec<SenderWriteRequest>,
655    bulk_requests: &mut Vec<SenderBulkRequest>,
656) {
657    WRITE_REJECT_TOTAL.inc_by(write_requests.len() as u64);
658
659    for req in write_requests.drain(..) {
660        req.sender.send(
661            RejectWriteSnafu {
662                region_id: req.request.region_id,
663            }
664            .fail(),
665        );
666    }
667    for req in bulk_requests.drain(..) {
668        let region_id = req.region_id;
669        req.sender.send(RejectWriteSnafu { region_id }.fail());
670    }
671}
672
673fn reject_region_write_requests(
674    rejected_region_ids: &HashSet<RegionId>,
675    write_requests: &mut Vec<SenderWriteRequest>,
676    bulk_requests: &mut Vec<SenderBulkRequest>,
677) {
678    let mut rejected_write_requests = write_requests
679        .extract_if(.., |req| {
680            rejected_region_ids.contains(&req.request.region_id)
681        })
682        .collect::<Vec<_>>();
683    let mut rejected_bulk_requests = bulk_requests
684        .extract_if(.., |req| rejected_region_ids.contains(&req.region_id))
685        .collect::<Vec<_>>();
686    reject_write_requests(&mut rejected_write_requests, &mut rejected_bulk_requests);
687}
688
689fn write_region_ids(
690    write_requests: &[SenderWriteRequest],
691    bulk_requests: &[SenderBulkRequest],
692) -> HashSet<RegionId> {
693    write_requests
694        .iter()
695        .map(|req| req.request.region_id)
696        .chain(bulk_requests.iter().map(|req| req.region_id))
697        .collect()
698}
699
700/// Rejects delete request under append mode.
701fn check_op_type(append_mode: bool, request: &WriteRequest) -> Result<()> {
702    if append_mode {
703        ensure!(
704            request.op_type == OpType::Put,
705            InvalidRequestSnafu {
706                region_id: request.region_id,
707                reason: "DELETE is not allowed under append mode",
708            }
709        );
710    }
711
712    Ok(())
713}
714
715fn check_partition_expr_version(
716    region_id: RegionId,
717    expected_version: u64,
718    request_version: Option<u64>,
719) -> Result<()> {
720    let request_version = match request_version {
721        None => return Ok(()),
722        Some(value) => value,
723    };
724    if request_version != expected_version {
725        return PartitionExprVersionMismatchSnafu {
726            region_id,
727            request_version,
728            expected_version,
729        }
730        .fail();
731    }
732    Ok(())
733}
734
735#[cfg(test)]
736mod tests {
737    use api::v1::{Row, Rows};
738    use futures::stream;
739    use log_store::error::{
740        Error as LogStoreError, IllegalStateSnafu, InvalidProviderSnafu, Result as LogStoreResult,
741    };
742    use store_api::logstore::entry::{Entry, NaiveEntry};
743    use store_api::logstore::provider::Provider;
744    use store_api::logstore::{AppendBatchResponse, EntryId, SendableEntryStream, WalIndex};
745    use store_api::region_request::AffectedRows;
746    use tokio::sync::oneshot;
747
748    use super::*;
749    use crate::request::OptionOutputTx;
750    use crate::test_util::version_util::VersionControlBuilder;
751
752    /// A log store that fails to build entries for `failing_region` and fails the
753    /// whole batch when `fail_append` is true.
754    #[derive(Debug, Default)]
755    struct MockLogStore {
756        failing_region: Option<RegionId>,
757        fail_append: bool,
758    }
759
760    #[async_trait::async_trait]
761    impl LogStore for MockLogStore {
762        type Error = LogStoreError;
763
764        async fn stop(&self) -> LogStoreResult<()> {
765            Ok(())
766        }
767
768        async fn append_batch(&self, entries: Vec<Entry>) -> LogStoreResult<AppendBatchResponse> {
769            if self.fail_append {
770                return IllegalStateSnafu {}.fail();
771            }
772            let mut last_entry_ids = HashMap::new();
773            for entry in &entries {
774                let last_entry_id = last_entry_ids.entry(entry.region_id()).or_insert(0);
775                *last_entry_id = entry.entry_id().max(*last_entry_id);
776            }
777            Ok(AppendBatchResponse { last_entry_ids })
778        }
779
780        async fn read(
781            &self,
782            _provider: &Provider,
783            _entry_id: EntryId,
784            _index: Option<WalIndex>,
785        ) -> LogStoreResult<SendableEntryStream<'static, Entry, Self::Error>> {
786            Ok(Box::pin(stream::empty()))
787        }
788
789        async fn create_namespace(&self, _ns: &Provider) -> LogStoreResult<()> {
790            Ok(())
791        }
792
793        async fn delete_namespace(&self, _ns: &Provider) -> LogStoreResult<()> {
794            Ok(())
795        }
796
797        async fn list_namespaces(&self) -> LogStoreResult<Vec<Provider>> {
798            Ok(vec![])
799        }
800
801        async fn obsolete(
802            &self,
803            _provider: &Provider,
804            _region_id: RegionId,
805            _entry_id: EntryId,
806        ) -> LogStoreResult<()> {
807            Ok(())
808        }
809
810        async fn obsolete_all(
811            &self,
812            _provider: &Provider,
813            _region_id: RegionId,
814        ) -> LogStoreResult<()> {
815            Ok(())
816        }
817
818        fn entry(
819            &self,
820            data: Vec<u8>,
821            entry_id: EntryId,
822            region_id: RegionId,
823            provider: &Provider,
824        ) -> LogStoreResult<Entry> {
825            if self.failing_region == Some(region_id) {
826                return InvalidProviderSnafu {
827                    expected: "raft_engine",
828                    actual: "mock",
829                }
830                .fail();
831            }
832            Ok(Entry::Naive(NaiveEntry {
833                provider: provider.clone(),
834                region_id,
835                entry_id,
836                data,
837            }))
838        }
839
840        fn latest_entry_id(&self, _provider: &Provider) -> LogStoreResult<EntryId> {
841            Ok(0)
842        }
843    }
844
845    /// Creates a write context for `region_id` with one pending mutation of one row.
846    fn new_region_ctx(
847        region_id: RegionId,
848    ) -> (RegionWriteCtx, oneshot::Receiver<Result<AffectedRows>>) {
849        let version_control = Arc::new(VersionControlBuilder::new().build());
850        let mut ctx = RegionWriteCtx::new(
851            region_id,
852            &version_control,
853            Provider::raft_engine_provider(region_id.as_u64()),
854            None,
855        );
856        let (tx, rx) = oneshot::channel();
857        ctx.push_mutation(
858            OpType::Put as i32,
859            Some(Rows {
860                schema: vec![],
861                rows: vec![Row { values: vec![] }],
862            }),
863            None,
864            OptionOutputTx::from(tx),
865            None,
866        );
867        (ctx, rx)
868    }
869
870    #[tokio::test]
871    async fn test_write_wal_skips_region_failed_to_build_entry() {
872        let failing_region = RegionId::new(1, 1);
873        let ok_region = RegionId::new(1, 2);
874        let wal = Wal::new(Arc::new(MockLogStore {
875            failing_region: Some(failing_region),
876            ..Default::default()
877        }));
878
879        let mut region_ctxs = HashMap::new();
880        let (ctx, failing_rx) = new_region_ctx(failing_region);
881        region_ctxs.insert(failing_region, ctx);
882        let (ctx, ok_rx) = new_region_ctx(ok_region);
883        region_ctxs.insert(ok_region, ctx);
884        let entry_id = region_ctxs[&ok_region].next_entry_id();
885
886        // The failed region must not fail the batch or panic the worker.
887        assert!(write_wal(&wal, &mut region_ctxs).await);
888
889        assert!(region_ctxs[&failing_region].is_failed());
890        assert!(!region_ctxs[&ok_region].is_failed());
891        assert_eq!(entry_id + 1, region_ctxs[&ok_region].next_entry_id());
892
893        // Waiters of the failed region get the error while others get the result.
894        drop(region_ctxs);
895        assert!(failing_rx.await.unwrap().is_err());
896        assert_eq!(1, ok_rx.await.unwrap().unwrap());
897    }
898
899    #[tokio::test]
900    async fn test_write_wal_all_regions_failed_to_build_entries() {
901        let failing_region = RegionId::new(1, 1);
902        let wal = Wal::new(Arc::new(MockLogStore {
903            failing_region: Some(failing_region),
904            ..Default::default()
905        }));
906
907        let mut region_ctxs = HashMap::new();
908        let (ctx, rx) = new_region_ctx(failing_region);
909        region_ctxs.insert(failing_region, ctx);
910
911        // Writing an empty batch to the WAL succeeds, the failed region must not panic
912        // the worker.
913        assert!(write_wal(&wal, &mut region_ctxs).await);
914
915        assert!(region_ctxs[&failing_region].is_failed());
916        drop(region_ctxs);
917        assert!(rx.await.unwrap().is_err());
918    }
919
920    #[tokio::test]
921    async fn test_write_wal_append_batch_failure() {
922        let region_id = RegionId::new(1, 1);
923        let wal = Wal::new(Arc::new(MockLogStore {
924            fail_append: true,
925            ..Default::default()
926        }));
927
928        let mut region_ctxs = HashMap::new();
929        let (ctx, rx) = new_region_ctx(region_id);
930        region_ctxs.insert(region_id, ctx);
931
932        assert!(!write_wal(&wal, &mut region_ctxs).await);
933
934        // All contexts are consumed and waiters are notified with the error.
935        assert!(region_ctxs.is_empty());
936        assert!(rx.await.unwrap().is_err());
937    }
938}