Skip to main content

s2_lite/backend/
core.rs

1use std::sync::Arc;
2
3use bytesize::ByteSize;
4use dashmap::DashMap;
5use futures::{
6    FutureExt as _,
7    future::{BoxFuture, Shared},
8};
9use s2_common::{
10    basin::BasinName,
11    config::{BasinConfig, OptionalStreamConfig},
12    encryption::{EncryptionAlgorithm, EncryptionSpec},
13    record::{NonZeroSeqNum, SeqNum, StreamPosition},
14    resources::ProvisionMode,
15    stream::StreamName,
16};
17use slatedb::config::{DurabilityLevel, ReadOptions, ScanOptions};
18use tokio::sync::{Semaphore, broadcast};
19
20use super::{
21    StreamHandle,
22    durability_notifier::DurabilityNotifier,
23    error::{
24        BasinDeletionPendingError, BasinNotFoundError, GetBasinConfigError, ProvisionStreamError,
25        StorageError, StreamDeletionPendingError, StreamNotFoundError, StreamerError,
26        StreamerMissingInActionError, TransactionConflictError,
27    },
28    kv,
29    streamer::{GuardedStreamerClient, StreamerClient, StreamerGenerationId},
30};
31use crate::{backend::bgtasks::BgtaskTrigger, stream_id::StreamId};
32
33type StreamerInitFuture = Shared<BoxFuture<'static, Result<StreamerClient, StreamerError>>>;
34
35#[derive(Clone)]
36enum StreamerClientSlot {
37    Initializing {
38        generation_id: StreamerGenerationId,
39        future: StreamerInitFuture,
40    },
41    Ready {
42        client: StreamerClient,
43    },
44}
45
46#[derive(Clone)]
47pub struct Backend {
48    pub(super) db: slatedb::Db,
49    streamer_slots: Arc<DashMap<StreamId, StreamerClientSlot>>,
50    append_inflight_bytes_sema: Arc<Semaphore>,
51    durability_notifier: DurabilityNotifier,
52    bgtask_trigger_tx: broadcast::Sender<BgtaskTrigger>,
53}
54
55impl Backend {
56    pub fn new(db: slatedb::Db, append_inflight_bytes: ByteSize) -> Self {
57        let (bgtask_trigger_tx, _) = broadcast::channel(16);
58        let append_inflight_bytes = Arc::new(Semaphore::new(
59            (append_inflight_bytes.as_u64() as usize).clamp(
60                s2_common::caps::RECORD_BATCH_MAX.bytes,
61                Semaphore::MAX_PERMITS,
62            ),
63        ));
64        let durability_notifier = DurabilityNotifier::spawn(&db);
65        Self {
66            db,
67            streamer_slots: Arc::new(DashMap::new()),
68            append_inflight_bytes_sema: append_inflight_bytes,
69            durability_notifier,
70            bgtask_trigger_tx,
71        }
72    }
73
74    pub(super) fn bgtask_trigger(&self, trigger: BgtaskTrigger) {
75        let _ = self.bgtask_trigger_tx.send(trigger);
76    }
77
78    pub(super) fn bgtask_trigger_subscribe(&self) -> broadcast::Receiver<BgtaskTrigger> {
79        self.bgtask_trigger_tx.subscribe()
80    }
81
82    /// Wait until writes up to `db_seq` are durable, i.e. visible to reads at
83    /// `DurabilityLevel::Remote`. Resolves immediately if they already are.
84    pub(super) async fn await_durable_seq(&self, db_seq: u64) -> Result<(), StorageError> {
85        let (tx, rx) = tokio::sync::oneshot::channel();
86        self.durability_notifier.subscribe(db_seq, move |res| {
87            let _ = tx.send(res);
88        });
89        let reason = match rx.await {
90            Ok(Ok(_)) => return Ok(()),
91            Ok(Err(reason)) => reason,
92            Err(_) => slatedb::CloseReason::Clean,
93        };
94        Err(slatedb::Error::closed(
95            "database closed while waiting for durability".to_owned(),
96            reason,
97        )
98        .into())
99    }
100
101    async fn start_streamer(
102        &self,
103        generation_id: StreamerGenerationId,
104        basin: BasinName,
105        stream: StreamName,
106    ) -> Result<StreamerClient, StreamerError> {
107        let stream_id = StreamId::new(&basin, &stream);
108
109        let (meta, persisted_tail, fencing_token, trim_point) = tokio::try_join!(
110            self.db_get(
111                kv::stream_meta::ser_key(&basin, &stream),
112                kv::stream_meta::deser_value,
113            ),
114            self.load_persisted_stream_tail(stream_id),
115            self.db_get(
116                kv::stream_fencing_token::ser_key(stream_id),
117                kv::stream_fencing_token::deser_value,
118            ),
119            self.db_get(
120                kv::stream_trim_point::ser_key(stream_id),
121                kv::stream_trim_point::deser_value,
122            )
123        )?;
124
125        let Some(meta) = meta else {
126            return Err(StreamNotFoundError { basin, stream }.into());
127        };
128
129        let (tail_pos, last_tail_write_timestamp) =
130            persisted_tail.unwrap_or((StreamPosition::MIN, kv::timestamp::TimestampSecs::ZERO));
131
132        self.assert_no_records_following_tail(stream_id, &basin, &stream, tail_pos)
133            .await?;
134
135        let fencing_token = fencing_token.unwrap_or_default();
136
137        if trim_point == Some(..NonZeroSeqNum::MAX) {
138            return Err(StreamDeletionPendingError.into());
139        }
140
141        let streamer_slots = self.streamer_slots.clone();
142        Ok(super::streamer::Spawner {
143            generation_id,
144            db: self.db.clone(),
145            stream_id,
146            config: meta.config,
147            cipher: meta.cipher,
148            tail_pos,
149            last_tail_write_timestamp,
150            fencing_token,
151            trim_point: ..trim_point.map_or(SeqNum::MIN, |tp| tp.end.get()),
152            append_inflight_bytes_sema: self.append_inflight_bytes_sema.clone(),
153            durability_notifier: self.durability_notifier.clone(),
154            bgtask_trigger_tx: self.bgtask_trigger_tx.clone(),
155        }
156        .spawn(move |client_id| {
157            streamer_slots.remove_if(&stream_id, |_, slot| {
158                matches!(slot, StreamerClientSlot::Ready { client } if client.generation_id() == client_id)
159            });
160        }))
161    }
162
163    async fn load_persisted_stream_tail(
164        &self,
165        stream_id: StreamId,
166    ) -> Result<Option<(StreamPosition, kv::timestamp::TimestampSecs)>, StorageError> {
167        let read_opts = ReadOptions {
168            durability_filter: DurabilityLevel::Remote,
169            ..Default::default()
170        };
171        let Some(entry) = self
172            .db
173            .get_key_value_with_options(kv::stream_tail_position::ser_key(stream_id), &read_opts)
174            .await?
175        else {
176            return Ok(None);
177        };
178        Ok(Some((
179            kv::stream_tail_position::deser_value(entry.value)?,
180            kv::timestamp::TimestampSecs::from_millis(entry.create_ts),
181        )))
182    }
183
184    async fn assert_no_records_following_tail(
185        &self,
186        stream_id: StreamId,
187        basin: &BasinName,
188        stream: &StreamName,
189        tail_pos: StreamPosition,
190    ) -> Result<(), StorageError> {
191        let prefix = kv::stream_record_data::ser_key_prefix(stream_id);
192        let start_suffix = kv::stream_record_data::ser_key_suffix(StreamPosition {
193            seq_num: tail_pos.seq_num,
194            timestamp: 0,
195        });
196        let scan_opts = ScanOptions {
197            durability_filter: DurabilityLevel::Remote,
198            ..Default::default()
199        };
200        let mut it = self
201            .db
202            .scan_prefix_with_options(prefix, start_suffix.., &scan_opts)
203            .await?;
204        let Some(kv) = it.next().await? else {
205            return Ok(());
206        };
207        let (deser_stream_id, pos) = kv::stream_record_data::deser_key(kv.key)?;
208        debug_assert_eq!(deser_stream_id, stream_id);
209        panic!(
210            "invariant violation: stream `{basin}/{stream}` tail_pos {tail_pos:?} but found record at {pos:?}"
211        )
212    }
213
214    fn streamer_client_slot(&self, basin: &BasinName, stream: &StreamName) -> StreamerClientSlot {
215        match self.streamer_slots.entry(StreamId::new(basin, stream)) {
216            dashmap::Entry::Occupied(mut oe) => {
217                if matches!(oe.get(), StreamerClientSlot::Ready { client } if client.is_dead()) {
218                    let slot = self.clone().new_initializing_slot(basin, stream);
219                    oe.insert(slot.clone());
220                    slot
221                } else {
222                    oe.get().clone()
223                }
224            }
225            dashmap::Entry::Vacant(ve) => {
226                let slot = self.clone().new_initializing_slot(basin, stream);
227                ve.insert(slot.clone());
228                slot
229            }
230        }
231    }
232
233    fn new_initializing_slot(self, basin: &BasinName, stream: &StreamName) -> StreamerClientSlot {
234        let basin = basin.clone();
235        let stream = stream.clone();
236        let generation_id = StreamerGenerationId::next();
237        let future = async move { self.start_streamer(generation_id, basin, stream).await }
238            .boxed()
239            .shared();
240        StreamerClientSlot::Initializing {
241            generation_id,
242            future,
243        }
244    }
245
246    fn streamer_finish_initialization(
247        &self,
248        stream_id: StreamId,
249        generation_id: StreamerGenerationId,
250        result: &Result<StreamerClient, StreamerError>,
251    ) {
252        if let dashmap::Entry::Occupied(mut oe) = self.streamer_slots.entry(stream_id) {
253            let is_same_init = matches!(
254                oe.get(),
255                StreamerClientSlot::Initializing {
256                    generation_id: state_generation_id,
257                    ..
258                } if *state_generation_id == generation_id
259            );
260            if is_same_init {
261                match result {
262                    Ok(client) => {
263                        debug_assert_eq!(client.generation_id(), generation_id);
264                        if client.is_dead() {
265                            oe.remove();
266                        } else {
267                            oe.insert(StreamerClientSlot::Ready {
268                                client: client.clone(),
269                            });
270                        }
271                    }
272                    Err(_) => {
273                        oe.remove();
274                    }
275                }
276            }
277        }
278    }
279
280    pub(super) async fn streamer_client(
281        &self,
282        basin: &BasinName,
283        stream: &StreamName,
284    ) -> Result<StreamerClient, StreamerError> {
285        let stream_id = StreamId::new(basin, stream);
286        match self.streamer_client_slot(basin, stream) {
287            StreamerClientSlot::Initializing {
288                generation_id,
289                future,
290            } => {
291                let result = future.await;
292                self.streamer_finish_initialization(stream_id, generation_id, &result);
293                result
294            }
295            StreamerClientSlot::Ready { client } => Ok(client),
296        }
297    }
298
299    pub(super) fn streamer_client_if_active(
300        &self,
301        basin: &BasinName,
302        stream: &StreamName,
303    ) -> Option<StreamerClient> {
304        let stream_id = StreamId::new(basin, stream);
305        let slot = self.streamer_slots.get(&stream_id)?;
306        match slot.value() {
307            StreamerClientSlot::Ready { client } if !client.is_dead() => Some(client.clone()),
308            _ => None,
309        }
310    }
311
312    pub(super) async fn streamer_client_guarded(
313        &self,
314        basin: &BasinName,
315        stream: &StreamName,
316    ) -> Result<GuardedStreamerClient, StreamerError> {
317        loop {
318            let client = self.streamer_client(basin, stream).await?;
319            match client.guard() {
320                Ok(client) => return Ok(client),
321                Err(StreamerMissingInActionError) => continue,
322            }
323        }
324    }
325
326    pub(super) async fn stream_handle_with_auto_create<E>(
327        &self,
328        basin: &BasinName,
329        stream: &StreamName,
330        should_auto_create: impl FnOnce(&BasinConfig) -> bool,
331        resolve_encryption: impl FnOnce(Option<EncryptionAlgorithm>) -> Result<EncryptionSpec, E>,
332    ) -> Result<StreamHandle, E>
333    where
334        E: From<StreamerError>
335            + From<StorageError>
336            + From<BasinNotFoundError>
337            + From<TransactionConflictError>
338            + From<BasinDeletionPendingError>
339            + From<StreamDeletionPendingError>
340            + From<StreamNotFoundError>,
341    {
342        match self.streamer_client_guarded(basin, stream).await {
343            Ok(client) => Ok(StreamHandle {
344                db: self.db.clone(),
345                encryption: resolve_encryption(client.cipher())?,
346                client,
347            }),
348            Err(StreamerError::StreamNotFound(e)) => {
349                let config = match self.get_basin_config(basin.clone()).await {
350                    Ok(config) => config,
351                    Err(GetBasinConfigError::Storage(e)) => Err(e)?,
352                    Err(GetBasinConfigError::BasinNotFound(e)) => Err(e)?,
353                };
354                if should_auto_create(&config) {
355                    if let Err(e) = self
356                        .provision_stream(
357                            basin.clone(),
358                            stream.clone(),
359                            OptionalStreamConfig::default(),
360                            ProvisionMode::CreateOnly {
361                                request_token: None,
362                            },
363                        )
364                        .await
365                    {
366                        match e {
367                            ProvisionStreamError::Storage(e) => Err(e)?,
368                            ProvisionStreamError::TransactionConflict(e) => Err(e)?,
369                            ProvisionStreamError::BasinDeletionPending(e) => Err(e)?,
370                            ProvisionStreamError::StreamDeletionPending(e) => Err(e)?,
371                            ProvisionStreamError::BasinNotFound(e) => Err(e)?,
372                            ProvisionStreamError::StreamAlreadyExists(_) => {}
373                            ProvisionStreamError::Validation(_) => {
374                                unreachable!("auto-create uses default config")
375                            }
376                        }
377                    }
378                    let client = self.streamer_client_guarded(basin, stream).await?;
379                    let encryption = resolve_encryption(client.cipher())?;
380                    Ok(StreamHandle {
381                        db: self.db.clone(),
382                        encryption,
383                        client,
384                    })
385                } else {
386                    Err(e.into())
387                }
388            }
389            Err(e) => Err(e.into()),
390        }
391    }
392}
393
394#[cfg(test)]
395mod tests {
396    use std::str::FromStr as _;
397
398    use bytes::Bytes;
399    use s2_common::{
400        config::{BasinConfig, OptionalStreamConfig, StreamConfig},
401        record::{Metered, MeteredExt as _, Record, StreamPosition},
402        resources::ProvisionMode,
403    };
404    use s2_storage::record::StoredRecord;
405    use slatedb::{WriteBatch, object_store};
406    use time::OffsetDateTime;
407
408    use super::*;
409
410    async fn new_test_backend() -> Backend {
411        let object_store: Arc<dyn object_store::ObjectStore> =
412            Arc::new(object_store::memory::InMemory::new());
413        let db = slatedb::Db::builder("test", object_store)
414            .build()
415            .await
416            .unwrap();
417        Backend::new(db, ByteSize::b(1))
418    }
419
420    #[tokio::test]
421    #[should_panic(expected = "invariant violation: stream `testbasin1/stream1` tail_pos")]
422    async fn start_streamer_fails_if_records_exist_after_tail_pos() {
423        let backend = new_test_backend().await;
424
425        let basin = BasinName::from_str("testbasin1").unwrap();
426        let stream = StreamName::from_str("stream1").unwrap();
427        let stream_id = StreamId::new(&basin, &stream);
428
429        let meta = kv::stream_meta::StreamMeta {
430            config: StreamConfig::default(),
431            cipher: None,
432            created_at: OffsetDateTime::now_utc(),
433            deleted_at: None,
434            creation_idempotency_key: None,
435        };
436
437        let tail_pos = StreamPosition {
438            seq_num: 1,
439            timestamp: 123,
440        };
441        let record_pos = StreamPosition {
442            seq_num: tail_pos.seq_num,
443            timestamp: tail_pos.timestamp,
444        };
445
446        let record = Record::try_from_parts(vec![], Bytes::from_static(b"hello")).unwrap();
447        let metered_record: Metered<StoredRecord> = StoredRecord::from(record).metered();
448
449        let mut wb = WriteBatch::new();
450        wb.put(
451            kv::stream_meta::ser_key(&basin, &stream),
452            kv::stream_meta::ser_value(&meta),
453        );
454        wb.put(
455            kv::stream_tail_position::ser_key(stream_id),
456            kv::stream_tail_position::ser_value(tail_pos),
457        );
458        wb.put(
459            kv::stream_record_data::ser_key(stream_id, record_pos),
460            kv::stream_record_data::ser_value(metered_record.as_ref()),
461        );
462        backend.db.write(wb).await.unwrap();
463
464        backend
465            .start_streamer(StreamerGenerationId::next(), basin.clone(), stream.clone())
466            .await
467            .unwrap();
468    }
469
470    #[tokio::test]
471    async fn streamer_client_slot_uses_single_initializer() {
472        let backend = new_test_backend().await;
473        let basin = BasinName::from_str("testbasin2").unwrap();
474        let stream = StreamName::from_str("stream2").unwrap();
475
476        let slot_1 = backend.streamer_client_slot(&basin, &stream);
477        let slot_2 = backend.streamer_client_slot(&basin, &stream);
478
479        let (generation_id_1, generation_id_2) = match (slot_1, slot_2) {
480            (
481                StreamerClientSlot::Initializing {
482                    generation_id: generation_id_1,
483                    ..
484                },
485                StreamerClientSlot::Initializing {
486                    generation_id: generation_id_2,
487                    ..
488                },
489            ) => (generation_id_1, generation_id_2),
490            _ => panic!("expected both slots to be Initializing"),
491        };
492        assert_eq!(generation_id_1, generation_id_2);
493        assert_eq!(backend.streamer_slots.len(), 1);
494    }
495
496    #[tokio::test]
497    async fn streamer_client_if_active_is_peek_only() {
498        let backend = new_test_backend().await;
499        let basin = BasinName::from_str("testbasin3").unwrap();
500        let stream = StreamName::from_str("stream3").unwrap();
501
502        backend
503            .provision_basin(
504                basin.clone(),
505                BasinConfig::default(),
506                ProvisionMode::CreateOnly {
507                    request_token: None,
508                },
509            )
510            .await
511            .unwrap();
512        backend
513            .provision_stream(
514                basin.clone(),
515                stream.clone(),
516                OptionalStreamConfig::default(),
517                ProvisionMode::CreateOnly {
518                    request_token: None,
519                },
520            )
521            .await
522            .unwrap();
523
524        assert!(backend.streamer_slots.is_empty());
525        assert!(backend.streamer_client_if_active(&basin, &stream).is_none());
526        assert!(backend.streamer_slots.is_empty());
527    }
528
529    #[tokio::test]
530    async fn streamer_client_failed_init_is_not_memoized() {
531        let backend = new_test_backend().await;
532        let basin = BasinName::from_str("testbasin4").unwrap();
533        let stream = StreamName::from_str("stream4").unwrap();
534        let stream_id = StreamId::new(&basin, &stream);
535
536        for _ in 0..2 {
537            let err = backend.streamer_client(&basin, &stream).await;
538            assert!(matches!(err, Err(StreamerError::StreamNotFound(_))));
539            assert!(
540                backend.streamer_slots.get(&stream_id).is_none(),
541                "failed init should not be cached"
542            );
543        }
544    }
545
546    #[tokio::test]
547    async fn streamer_finish_initialization_ignores_stale_generation_id() {
548        let backend = new_test_backend().await;
549        let basin = BasinName::from_str("testbasin5").unwrap();
550        let stream = StreamName::from_str("stream5").unwrap();
551        let stream_id = StreamId::new(&basin, &stream);
552
553        let stale_generation_id = StreamerGenerationId::next();
554        let current_generation_id = StreamerGenerationId::next();
555        let future = futures::future::pending::<Result<StreamerClient, StreamerError>>()
556            .boxed()
557            .shared();
558        backend.streamer_slots.insert(
559            stream_id,
560            StreamerClientSlot::Initializing {
561                generation_id: current_generation_id,
562                future: future.clone(),
563            },
564        );
565
566        let stale_result = Err(StreamNotFoundError { basin, stream }.into());
567        backend.streamer_finish_initialization(stream_id, stale_generation_id, &stale_result);
568
569        let Some(slot) = backend.streamer_slots.get(&stream_id) else {
570            panic!("stale init completion should not alter slot state");
571        };
572        match slot.value() {
573            StreamerClientSlot::Initializing { generation_id, .. } => {
574                assert_eq!(*generation_id, current_generation_id)
575            }
576            _ => panic!("expected initializing slot to remain unchanged"),
577        }
578    }
579
580    #[tokio::test(flavor = "multi_thread")]
581    async fn concurrent_appends_auto_create_stream_without_spurious_not_found() {
582        use s2_common::stream::{AppendInput, AppendRecord, AppendRecordParts};
583
584        use crate::backend::error::AppendError;
585
586        fn append_input(body: &str) -> AppendInput {
587            let record =
588                Record::try_from_parts(vec![], Bytes::copy_from_slice(body.as_bytes())).unwrap();
589            let record: AppendRecord = AppendRecordParts {
590                timestamp: None,
591                record: Metered::from(record),
592            }
593            .try_into()
594            .unwrap();
595            AppendInput {
596                records: vec![record].try_into().unwrap(),
597                match_seq_num: None,
598                fencing_token: None,
599            }
600        }
601
602        let backend = new_test_backend().await;
603        let basin = BasinName::from_str("autocreate").unwrap();
604        backend
605            .provision_basin(
606                basin.clone(),
607                BasinConfig {
608                    create_stream_on_append: true,
609                    ..BasinConfig::default()
610                },
611                ProvisionMode::CreateOnly {
612                    request_token: None,
613                },
614            )
615            .await
616            .unwrap();
617
618        for round in 0..10 {
619            let stream = StreamName::from_str(&format!("fresh-{round}")).unwrap();
620            let tasks: Vec<_> = (0..4)
621                .map(|i| {
622                    let backend = backend.clone();
623                    let basin = basin.clone();
624                    let stream = stream.clone();
625                    tokio::spawn(async move {
626                        let handle = backend.open_for_append(&basin, &stream, None).await?;
627                        handle.append(append_input(&format!("r{i}"))).await
628                    })
629                })
630                .collect();
631            for task in tasks {
632                match task.await.unwrap() {
633                    Ok(_) => {}
634                    // Conflict losers surface as retryable 409s to clients.
635                    Err(AppendError::TransactionConflict(_)) => {}
636                    Err(e) => panic!("concurrent auto-create append failed: {e:?}"),
637                }
638            }
639        }
640    }
641}