1use std::any::Any;
13use std::collections::BTreeMap;
14use std::collections::btree_map::Entry;
15use std::fmt::Debug;
16use std::future::Future;
17use std::sync::{Arc, RwLock, TryLockError, Weak};
18use std::time::{Duration, Instant};
19
20use differential_dataflow::difference::Monoid;
21use differential_dataflow::lattice::Lattice;
22use mz_dyncfg::{Config, ParameterScope};
23use mz_ore::instrument;
24use mz_ore::metrics::MetricsRegistry;
25use mz_ore::task::{AbortOnDropHandle, JoinHandle};
26use mz_ore::url::SensitiveUrl;
27use mz_persist::cfg::{BlobConfig, ConsensusConfig, open_hedge_sibling};
28use mz_persist::hedge::{HedgeSiblingOpener, HedgedBlob};
29use mz_persist::location::{
30 BLOB_GET_LIVENESS_KEY, Blob, CONSENSUS_HEAD_LIVENESS_KEY, Consensus, ExternalError, Tasked,
31 VersionedData,
32};
33use mz_persist_types::{Codec, Codec64};
34use timely::progress::Timestamp;
35use tokio::sync::{Mutex, OnceCell};
36use tracing::debug;
37
38use crate::async_runtime::IsolatedRuntime;
39use crate::error::{CodecConcreteType, CodecMismatch};
40use crate::internal::cache::BlobMemCache;
41use crate::internal::machine::retry_external;
42use crate::internal::metrics::{LockMetrics, Metrics, MetricsBlob, MetricsConsensus, ShardMetrics};
43use crate::internal::state::TypedState;
44use crate::internal::watch::{AwaitableState, StateWatchNotifier};
45use crate::rpc::{PubSubClientConnection, PubSubSender, ShardSubscriptionToken};
46use crate::schema::SchemaCacheMaps;
47use crate::{Diagnostics, PersistClient, PersistConfig, PersistLocation, ShardId};
48
49#[derive(Debug)]
58pub struct PersistClientCache {
59 pub cfg: PersistConfig,
61 pub(crate) metrics: Arc<Metrics>,
62 blob_by_uri: Mutex<BTreeMap<SensitiveUrl, (RttLatencyTask, Arc<dyn Blob>)>>,
63 consensus_by_uri: Mutex<BTreeMap<SensitiveUrl, (RttLatencyTask, Arc<dyn Consensus>)>>,
64 isolated_runtime: Arc<IsolatedRuntime>,
65 pub(crate) state_cache: Arc<StateCache>,
66 pubsub_sender: Arc<dyn PubSubSender>,
67 _pubsub_receiver_task: JoinHandle<()>,
68}
69
70#[derive(Debug)]
71struct RttLatencyTask(#[allow(dead_code)] AbortOnDropHandle<()>);
72
73impl PersistClientCache {
74 pub fn new<F>(cfg: PersistConfig, registry: &MetricsRegistry, pubsub: F) -> Self
76 where
77 F: FnOnce(&PersistConfig, Arc<Metrics>) -> PubSubClientConnection,
78 {
79 let metrics = Arc::new(Metrics::new(&cfg, registry));
80 let pubsub_client = pubsub(&cfg, Arc::clone(&metrics));
81
82 let state_cache = Arc::new(StateCache::new(
83 &cfg,
84 Arc::clone(&metrics),
85 Arc::clone(&pubsub_client.sender),
86 ));
87 let _pubsub_receiver_task = crate::rpc::subscribe_state_cache_to_pubsub(
88 Arc::clone(&state_cache),
89 pubsub_client.receiver,
90 );
91 let isolated_runtime =
92 IsolatedRuntime::new(registry, Some(cfg.isolated_runtime_worker_threads));
93
94 PersistClientCache {
95 cfg,
96 metrics,
97 blob_by_uri: Mutex::new(BTreeMap::new()),
98 consensus_by_uri: Mutex::new(BTreeMap::new()),
99 isolated_runtime: Arc::new(isolated_runtime),
100 state_cache,
101 pubsub_sender: pubsub_client.sender,
102 _pubsub_receiver_task,
103 }
104 }
105
106 pub fn new_no_metrics() -> Self {
109 Self::new(
110 PersistConfig::new_for_tests(),
111 &MetricsRegistry::new(),
112 |_, _| PubSubClientConnection::noop(),
113 )
114 }
115
116 #[cfg(feature = "turmoil")]
117 pub fn new_for_turmoil() -> Self {
122 use crate::rpc::NoopPubSubSender;
123
124 let cfg = PersistConfig::new_for_tests();
125 let metrics = Arc::new(Metrics::new(&cfg, &MetricsRegistry::new()));
126
127 let pubsub_sender: Arc<dyn PubSubSender> = Arc::new(NoopPubSubSender);
128 let _pubsub_receiver_task = mz_ore::task::spawn(|| "noop", async {});
129
130 let state_cache = Arc::new(StateCache::new(
131 &cfg,
132 Arc::clone(&metrics),
133 Arc::clone(&pubsub_sender),
134 ));
135 let isolated_runtime = IsolatedRuntime::new_disabled();
136
137 PersistClientCache {
138 cfg,
139 metrics,
140 blob_by_uri: Mutex::new(BTreeMap::new()),
141 consensus_by_uri: Mutex::new(BTreeMap::new()),
142 isolated_runtime: Arc::new(isolated_runtime),
143 state_cache,
144 pubsub_sender,
145 _pubsub_receiver_task,
146 }
147 }
148
149 pub fn cfg(&self) -> &PersistConfig {
151 &self.cfg
152 }
153
154 pub fn metrics(&self) -> &Arc<Metrics> {
156 &self.metrics
157 }
158
159 pub fn shard_metrics(&self, shard_id: &ShardId, name: &str) -> Arc<ShardMetrics> {
161 self.metrics.shards.shard(shard_id, name)
162 }
163
164 pub fn clear_state_cache(&mut self) {
168 self.state_cache = Arc::new(StateCache::new(
169 &self.cfg,
170 Arc::clone(&self.metrics),
171 Arc::clone(&self.pubsub_sender),
172 ))
173 }
174
175 #[instrument(level = "debug")]
180 pub async fn open(&self, location: PersistLocation) -> Result<PersistClient, ExternalError> {
181 let blob = self.open_blob(location.blob_uri).await?;
182 let consensus = self.open_consensus(location.consensus_uri).await?;
183 PersistClient::new(
184 self.cfg.clone(),
185 blob,
186 consensus,
187 Arc::clone(&self.metrics),
188 Arc::clone(&self.isolated_runtime),
189 Arc::clone(&self.state_cache),
190 Arc::clone(&self.pubsub_sender),
191 )
192 }
193
194 const PROMETHEUS_SCRAPE_INTERVAL: Duration = Duration::from_secs(60);
196
197 async fn open_consensus(
198 &self,
199 consensus_uri: SensitiveUrl,
200 ) -> Result<Arc<dyn Consensus>, ExternalError> {
201 let mut consensus_by_uri = self.consensus_by_uri.lock().await;
202 let consensus = match consensus_by_uri.entry(consensus_uri) {
203 Entry::Occupied(x) => Arc::clone(&x.get().1),
204 Entry::Vacant(x) => {
205 let consensus = ConsensusConfig::try_from(
208 x.key(),
209 Box::new(self.cfg.clone()),
210 self.metrics.postgres_consensus.clone(),
211 Arc::clone(&self.cfg().configs),
212 )?;
213 let consensus =
214 retry_external(&self.metrics.retries.external.consensus_open, || {
215 consensus.clone().open()
216 })
217 .await;
218 let consensus =
219 Arc::new(MetricsConsensus::new(consensus, Arc::clone(&self.metrics)));
220 let consensus = Arc::new(Tasked(consensus));
221 let task = consensus_rtt_latency_task(
222 Arc::clone(&consensus),
223 Arc::clone(&self.metrics),
224 Self::PROMETHEUS_SCRAPE_INTERVAL,
225 )
226 .await;
227 Arc::clone(
228 &x.insert((RttLatencyTask(task.abort_on_drop()), consensus))
229 .1,
230 )
231 }
232 };
233 Ok(consensus)
234 }
235
236 async fn open_blob(&self, blob_uri: SensitiveUrl) -> Result<Arc<dyn Blob>, ExternalError> {
237 let mut blob_by_uri = self.blob_by_uri.lock().await;
238 let blob = match blob_by_uri.entry(blob_uri) {
239 Entry::Occupied(x) => Arc::clone(&x.get().1),
240 Entry::Vacant(x) => {
241 let blob = BlobConfig::try_from(
244 x.key(),
245 Box::new(self.cfg.clone()),
246 self.metrics.s3_blob.clone(),
247 )
248 .await?;
249 let blob = retry_external(&self.metrics.retries.external.blob_open, || {
250 blob.clone().open()
251 })
252 .await;
253 let sibling_opener: HedgeSiblingOpener = {
266 let url = x.key().clone();
267 let cfg = self.cfg.clone();
268 let metrics = self.metrics.s3_blob.clone();
269 Box::new(move || {
270 let (url, cfg, metrics) = (url.clone(), cfg.clone(), metrics.clone());
271 Box::pin(
272 async move { open_hedge_sibling(&url, Box::new(cfg), metrics).await },
273 )
274 })
275 };
276 let blob = Arc::new(HedgedBlob::new_arming(
277 blob,
278 sibling_opener,
279 Arc::clone(&self.cfg.configs),
280 self.metrics.blob_hedge.clone(),
281 ));
282 let blob = Arc::new(MetricsBlob::new(blob, Arc::clone(&self.metrics)));
283 let blob = Arc::new(Tasked(blob));
284 let task = blob_rtt_latency_task(
285 Arc::clone(&blob),
286 Arc::clone(&self.metrics),
287 Self::PROMETHEUS_SCRAPE_INTERVAL,
288 )
289 .await;
290 let blob = BlobMemCache::new(&self.cfg, Arc::clone(&self.metrics), blob);
293 Arc::clone(&x.insert((RttLatencyTask(task.abort_on_drop()), blob)).1)
294 }
295 };
296 Ok(blob)
297 }
298}
299
300#[allow(clippy::unused_async)]
314async fn blob_rtt_latency_task(
315 blob: Arc<Tasked<MetricsBlob>>,
316 metrics: Arc<Metrics>,
317 measurement_interval: Duration,
318) -> JoinHandle<()> {
319 mz_ore::task::spawn(|| "persist::blob_rtt_latency", async move {
320 let mut next_measurement = tokio::time::Instant::now();
323 loop {
324 tokio::time::sleep_until(next_measurement).await;
325 let start = Instant::now();
326 match blob.get(BLOB_GET_LIVENESS_KEY).await {
327 Ok(_) => {
328 metrics.blob.rtt_latency.set(start.elapsed().as_secs_f64());
329 }
330 Err(_) => {
331 }
335 }
336 next_measurement = tokio::time::Instant::now() + measurement_interval;
337 }
338 })
339}
340
341#[allow(clippy::unused_async)]
355async fn consensus_rtt_latency_task(
356 consensus: Arc<Tasked<MetricsConsensus>>,
357 metrics: Arc<Metrics>,
358 measurement_interval: Duration,
359) -> JoinHandle<()> {
360 mz_ore::task::spawn(|| "persist::consensus_rtt_latency", async move {
361 let mut next_measurement = tokio::time::Instant::now();
364 loop {
365 tokio::time::sleep_until(next_measurement).await;
366 let start = Instant::now();
367 match consensus.head(CONSENSUS_HEAD_LIVENESS_KEY).await {
368 Ok(_) => {
369 metrics
370 .consensus
371 .rtt_latency
372 .set(start.elapsed().as_secs_f64());
373 }
374 Err(_) => {
375 }
379 }
380 next_measurement = tokio::time::Instant::now() + measurement_interval;
381 }
382 })
383}
384
385pub(crate) trait DynState: Debug + Send + Sync {
386 fn codecs(&self) -> (String, String, String, String, Option<CodecConcreteType>);
387 fn as_any(self: Arc<Self>) -> Arc<dyn Any + Send + Sync>;
388 fn push_diff(&self, diff: VersionedData);
389}
390
391impl<K, V, T, D> DynState for LockingTypedState<K, V, T, D>
392where
393 K: Codec,
394 V: Codec,
395 T: Timestamp + Lattice + Codec64 + Sync,
396 D: Codec64,
397{
398 fn codecs(&self) -> (String, String, String, String, Option<CodecConcreteType>) {
399 (
400 K::codec_name(),
401 V::codec_name(),
402 T::codec_name(),
403 D::codec_name(),
404 Some(CodecConcreteType(std::any::type_name::<(K, V, T, D)>())),
405 )
406 }
407
408 fn as_any(self: Arc<Self>) -> Arc<dyn Any + Send + Sync> {
409 self
410 }
411
412 fn push_diff(&self, diff: VersionedData) {
413 self.write_lock(&self.metrics.locks.applier_write, |state| {
414 let seqno_before = state.seqno;
415 state.apply_encoded_diffs(&self.cfg, &self.metrics, std::iter::once(&diff));
416 let seqno_after = state.seqno;
417 assert!(seqno_after >= seqno_before);
418
419 if seqno_before != seqno_after {
420 debug!(
421 "applied pushed diff {}. seqno {} -> {}.",
422 state.shard_id, seqno_before, state.seqno
423 );
424 self.metrics.pubsub_client.receiver.diff_applied.inc();
425 } else {
426 debug!(
427 "failed to apply pushed diff {}. seqno {} vs diff {}",
428 state.shard_id, seqno_before, diff.seqno
429 );
430 if diff.seqno <= seqno_before {
431 self.metrics
432 .pubsub_client
433 .receiver
434 .diff_not_applied_stale
435 .inc();
436 } else {
437 self.metrics
438 .pubsub_client
439 .receiver
440 .diff_not_applied_out_of_order
441 .inc();
442 }
443 }
444 })
445 }
446}
447
448#[derive(Debug)]
459pub struct StateCache {
460 cfg: Arc<PersistConfig>,
461 pub(crate) metrics: Arc<Metrics>,
462 states: Arc<std::sync::Mutex<BTreeMap<ShardId, Arc<OnceCell<Weak<dyn DynState>>>>>>,
463 pubsub_sender: Arc<dyn PubSubSender>,
464}
465
466#[derive(Debug)]
467enum StateCacheInit {
468 Init(Arc<dyn DynState>),
469 NeedInit(Arc<OnceCell<Weak<dyn DynState>>>),
470}
471
472impl StateCache {
473 pub fn new(
475 cfg: &PersistConfig,
476 metrics: Arc<Metrics>,
477 pubsub_sender: Arc<dyn PubSubSender>,
478 ) -> Self {
479 StateCache {
480 cfg: Arc::new(cfg.clone()),
481 metrics,
482 states: Default::default(),
483 pubsub_sender,
484 }
485 }
486
487 #[cfg(test)]
488 pub(crate) fn new_no_metrics() -> Self {
489 Self::new(
490 &PersistConfig::new_for_tests(),
491 Arc::new(Metrics::new(
492 &PersistConfig::new_for_tests(),
493 &MetricsRegistry::new(),
494 )),
495 Arc::new(crate::rpc::NoopPubSubSender),
496 )
497 }
498
499 pub(crate) async fn get<K, V, T, D, F, InitFn>(
500 &self,
501 shard_id: ShardId,
502 mut init_fn: InitFn,
503 diagnostics: &Diagnostics,
504 ) -> Result<Arc<LockingTypedState<K, V, T, D>>, Box<CodecMismatch>>
505 where
506 K: Debug + Codec,
507 V: Debug + Codec,
508 T: Timestamp + Lattice + Codec64 + Sync,
509 D: Monoid + Codec64,
510 F: Future<Output = Result<TypedState<K, V, T, D>, Box<CodecMismatch>>>,
511 InitFn: FnMut() -> F,
512 {
513 loop {
514 let init = {
515 let mut states = self.states.lock().expect("lock poisoned");
516 let state = states.entry(shard_id).or_default();
517 match state.get() {
518 Some(once_val) => match once_val.upgrade() {
519 Some(x) => StateCacheInit::Init(x),
520 None => {
521 *state = Arc::new(OnceCell::new());
525 StateCacheInit::NeedInit(Arc::clone(state))
526 }
527 },
528 None => StateCacheInit::NeedInit(Arc::clone(state)),
529 }
530 };
531
532 let state = match init {
533 StateCacheInit::Init(x) => x,
534 StateCacheInit::NeedInit(init_once) => {
535 let mut did_init: Option<Arc<LockingTypedState<K, V, T, D>>> = None;
536 let state = init_once
537 .get_or_try_init::<Box<CodecMismatch>, _, _>(|| async {
538 let init_res = init_fn().await;
539 let state = Arc::new(LockingTypedState::new(
540 shard_id,
541 init_res?,
542 Arc::clone(&self.metrics),
543 Arc::clone(&self.cfg),
544 Arc::clone(&self.pubsub_sender).subscribe(&shard_id),
545 diagnostics,
546 ));
547 let ret = Arc::downgrade(&state);
548 did_init = Some(state);
549 let ret: Weak<dyn DynState> = ret;
550 Ok(ret)
551 })
552 .await?;
553 if let Some(x) = did_init {
554 return Ok(x);
558 }
559 let Some(state) = state.upgrade() else {
560 continue;
567 };
568 state
569 }
570 };
571
572 match Arc::clone(&state)
573 .as_any()
574 .downcast::<LockingTypedState<K, V, T, D>>()
575 {
576 Ok(x) => return Ok(x),
577 Err(_) => {
578 return Err(Box::new(CodecMismatch {
579 requested: (
580 K::codec_name(),
581 V::codec_name(),
582 T::codec_name(),
583 D::codec_name(),
584 Some(CodecConcreteType(std::any::type_name::<(K, V, T, D)>())),
585 ),
586 actual: state.codecs(),
587 }));
588 }
589 }
590 }
591 }
592
593 pub(crate) fn get_state_weak(&self, shard_id: &ShardId) -> Option<Weak<dyn DynState>> {
594 self.states
595 .lock()
596 .expect("lock")
597 .get(shard_id)
598 .and_then(|x| x.get())
599 .map(Weak::clone)
600 }
601
602 #[cfg(test)]
603 fn get_cached(&self, shard_id: &ShardId) -> Option<Arc<dyn DynState>> {
604 self.states
605 .lock()
606 .expect("lock")
607 .get(shard_id)
608 .and_then(|x| x.get())
609 .and_then(|x| x.upgrade())
610 }
611
612 #[cfg(test)]
613 fn initialized_count(&self) -> usize {
614 self.states
615 .lock()
616 .expect("lock")
617 .values()
618 .filter(|x| x.initialized())
619 .count()
620 }
621
622 #[cfg(test)]
623 fn strong_count(&self) -> usize {
624 self.states
625 .lock()
626 .expect("lock")
627 .values()
628 .filter(|x| x.get().map_or(false, |x| x.upgrade().is_some()))
629 .count()
630 }
631}
632
633pub(crate) struct LockingTypedState<K, V, T, D> {
637 shard_id: ShardId,
638 state: RwLock<TypedState<K, V, T, D>>,
639 notifier: StateWatchNotifier<T>,
640 cfg: Arc<PersistConfig>,
641 metrics: Arc<Metrics>,
642 shard_metrics: Arc<ShardMetrics>,
646 update_semaphore: AwaitableState<Option<tokio::time::Instant>>,
647 schema_cache: Arc<dyn Any + Send + Sync>,
650 _subscription_token: Arc<ShardSubscriptionToken>,
651}
652
653impl<K, V, T: Debug, D> Debug for LockingTypedState<K, V, T, D> {
654 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
655 let LockingTypedState {
656 shard_id,
657 state,
658 notifier,
659 cfg: _cfg,
660 metrics: _metrics,
661 shard_metrics: _shard_metrics,
662 update_semaphore: _,
663 schema_cache: _schema_cache,
664 _subscription_token,
665 } = self;
666 f.debug_struct("LockingTypedState")
667 .field("shard_id", shard_id)
668 .field("state", state)
669 .field("notifier", notifier)
670 .finish()
671 }
672}
673
674impl<K: Codec, V: Codec, T, D> LockingTypedState<K, V, T, D> {
675 fn new(
676 shard_id: ShardId,
677 initial_state: TypedState<K, V, T, D>,
678 metrics: Arc<Metrics>,
679 cfg: Arc<PersistConfig>,
680 subscription_token: Arc<ShardSubscriptionToken>,
681 diagnostics: &Diagnostics,
682 ) -> Self
683 where
684 T: Timestamp + Lattice + Codec64,
686 D: Codec64,
687 {
688 let notifier = StateWatchNotifier::new(Arc::clone(&metrics), initial_state.upper().clone());
689 Self {
690 shard_id,
691 notifier,
692 state: RwLock::new(initial_state),
693 cfg: Arc::clone(&cfg),
694 shard_metrics: metrics.shards.shard(&shard_id, &diagnostics.shard_name),
695 update_semaphore: AwaitableState::new(None),
696 schema_cache: Arc::new(SchemaCacheMaps::<K, V>::new(&metrics.schema)),
697 metrics,
698 _subscription_token: subscription_token,
699 }
700 }
701
702 pub(crate) fn schema_cache(&self) -> Arc<SchemaCacheMaps<K, V>> {
703 Arc::clone(&self.schema_cache)
704 .downcast::<SchemaCacheMaps<K, V>>()
705 .expect("K and V match")
706 }
707}
708
709pub(crate) const STATE_UPDATE_LEASE_TIMEOUT: Config<Duration> = Config::new(
710 "persist_state_update_lease_timeout",
711 Duration::from_secs(1),
712 "The amount of time for a command to wait for a previous command to finish before executing. \
713 (If zero, commands will not wait for others to complete.) Higher values reduce database contention \
714 at the cost of higher worst-case latencies for individual requests.",
715 ParameterScope::Environment,
716);
717
718impl<K, V, T, D> LockingTypedState<K, V, T, D> {
719 pub(crate) fn shard_id(&self) -> &ShardId {
720 &self.shard_id
721 }
722
723 pub(crate) fn read_lock<R, F: FnMut(&TypedState<K, V, T, D>) -> R>(
724 &self,
725 metrics: &LockMetrics,
726 mut f: F,
727 ) -> R {
728 metrics.acquire_count.inc();
729 let state = match self.state.try_read() {
730 Ok(x) => x,
731 Err(TryLockError::WouldBlock) => {
732 metrics.blocking_acquire_count.inc();
733 let start = Instant::now();
734 let state = self.state.read().expect("lock poisoned");
735 metrics
736 .blocking_seconds
737 .inc_by(start.elapsed().as_secs_f64());
738 state
739 }
740 Err(TryLockError::Poisoned(err)) => panic!("state read lock poisoned: {}", err),
741 };
742 f(&state)
743 }
744
745 pub(crate) fn write_lock<R, F>(&self, metrics: &LockMetrics, f: F) -> R
746 where
747 F: FnOnce(&mut TypedState<K, V, T, D>) -> R,
748 K: Codec,
751 V: Codec,
752 T: Timestamp + Lattice + Codec64,
753 D: Codec64,
754 {
755 metrics.acquire_count.inc();
756 let mut state = match self.state.try_write() {
757 Ok(x) => x,
758 Err(TryLockError::WouldBlock) => {
759 metrics.blocking_acquire_count.inc();
760 let start = Instant::now();
761 let state = self.state.write().expect("lock poisoned");
762 metrics
763 .blocking_seconds
764 .inc_by(start.elapsed().as_secs_f64());
765 state
766 }
767 Err(TryLockError::Poisoned(err)) => panic!("state read lock poisoned: {}", err),
768 };
769 let seqno_before = state.seqno;
770 let ret = f(&mut state);
771 let seqno_after = state.seqno;
772 mz_ore::soft_assert_no_log!(seqno_after >= seqno_before);
773 if seqno_after > seqno_before {
774 self.notifier.notify(seqno_after, state.upper());
779 }
780 drop(state);
783 ret
784 }
785
786 pub(crate) async fn lease_for_update(&self) -> impl Drop {
794 use tokio::time::Instant;
795
796 let timeout = STATE_UPDATE_LEASE_TIMEOUT.get(&self.cfg);
797
798 struct DropLease(Option<(AwaitableState<Option<Instant>>, Instant)>);
799
800 impl Drop for DropLease {
801 fn drop(&mut self) {
802 if let Some((state, time)) = self.0.take() {
803 state.maybe_modify(|s| {
805 if s.is_some_and(|t| t == time) {
806 *s.get_mut() = None;
807 }
808 })
809 }
810 }
811 }
812
813 if timeout.is_zero() {
815 return DropLease(None);
816 }
817
818 let timeout_state = self.update_semaphore.clone();
819 loop {
820 let now = tokio::time::Instant::now();
821 let expires_at = now + timeout;
822 let maybe_leased = timeout_state.maybe_modify(|state| {
824 if let Some(other_expires_at) = **state
825 && other_expires_at > now
826 {
827 Err(other_expires_at)
829 } else {
830 *state.get_mut() = Some(expires_at);
831 Ok(())
832 }
833 });
834
835 match maybe_leased {
836 Ok(()) => {
837 break DropLease(Some((timeout_state, expires_at)));
838 }
839 Err(other_expires_at) => {
840 let _ = tokio::time::timeout_at(
845 other_expires_at,
846 timeout_state.wait_while(|s| s.is_some()),
847 )
848 .await;
849 }
850 }
851 }
852 }
853
854 pub(crate) fn notifier(&self) -> &StateWatchNotifier<T> {
855 &self.notifier
856 }
857}
858
859#[cfg(test)]
860mod tests {
861 use std::ops::Deref;
862 use std::pin::pin;
863 use std::str::FromStr;
864 use std::sync::atomic::{AtomicBool, Ordering};
865
866 use super::*;
867 use crate::rpc::NoopPubSubSender;
868 use futures::stream::{FuturesUnordered, StreamExt};
869 use mz_build_info::DUMMY_BUILD_INFO;
870 use mz_ore::task::spawn;
871 use mz_ore::{assert_err, assert_none};
872 use tokio::sync::oneshot;
873
874 #[mz_ore::test(tokio::test)]
875 #[cfg_attr(miri, ignore)] async fn client_cache() {
877 let cache = PersistClientCache::new(
878 PersistConfig::new_for_tests(),
879 &MetricsRegistry::new(),
880 |_, _| PubSubClientConnection::noop(),
881 );
882 assert_eq!(cache.blob_by_uri.lock().await.len(), 0);
883 assert_eq!(cache.consensus_by_uri.lock().await.len(), 0);
884
885 let _ = cache
887 .open(PersistLocation {
888 blob_uri: SensitiveUrl::from_str("mem://blob_zero").expect("invalid URL"),
889 consensus_uri: SensitiveUrl::from_str("mem://consensus_zero").expect("invalid URL"),
890 })
891 .await
892 .expect("failed to open location");
893 assert_eq!(cache.blob_by_uri.lock().await.len(), 1);
894 assert_eq!(cache.consensus_by_uri.lock().await.len(), 1);
895
896 let _ = cache
899 .open(PersistLocation {
900 blob_uri: SensitiveUrl::from_str("mem://blob_one").expect("invalid URL"),
901 consensus_uri: SensitiveUrl::from_str("mem://consensus_zero").expect("invalid URL"),
902 })
903 .await
904 .expect("failed to open location");
905 assert_eq!(cache.blob_by_uri.lock().await.len(), 2);
906 assert_eq!(cache.consensus_by_uri.lock().await.len(), 1);
907
908 let _ = cache
910 .open(PersistLocation {
911 blob_uri: SensitiveUrl::from_str("mem://blob_one").expect("invalid URL"),
912 consensus_uri: SensitiveUrl::from_str("mem://consensus_one").expect("invalid URL"),
913 })
914 .await
915 .expect("failed to open location");
916 assert_eq!(cache.blob_by_uri.lock().await.len(), 2);
917 assert_eq!(cache.consensus_by_uri.lock().await.len(), 2);
918
919 let _ = cache
921 .open(PersistLocation {
922 blob_uri: SensitiveUrl::from_str("mem://blob_one?foo").expect("invalid URL"),
923 consensus_uri: SensitiveUrl::from_str("mem://consensus_one/bar")
924 .expect("invalid URL"),
925 })
926 .await
927 .expect("failed to open location");
928 assert_eq!(cache.blob_by_uri.lock().await.len(), 3);
929 assert_eq!(cache.consensus_by_uri.lock().await.len(), 3);
930
931 let _ = cache
933 .open(PersistLocation {
934 blob_uri: SensitiveUrl::from_str("mem://user@blob_one").expect("invalid URL"),
935 consensus_uri: SensitiveUrl::from_str("mem://@consensus_one:123")
936 .expect("invalid URL"),
937 })
938 .await
939 .expect("failed to open location");
940 assert_eq!(cache.blob_by_uri.lock().await.len(), 4);
941 assert_eq!(cache.consensus_by_uri.lock().await.len(), 4);
942 }
943
944 #[mz_ore::test(tokio::test)]
945 #[cfg_attr(miri, ignore)] async fn state_cache() {
947 mz_ore::test::init_logging();
948 fn new_state<K, V, T, D>(shard_id: ShardId) -> TypedState<K, V, T, D>
949 where
950 K: Codec,
951 V: Codec,
952 T: Timestamp + Lattice + Codec64,
953 D: Codec64,
954 {
955 TypedState::new(
956 DUMMY_BUILD_INFO.semver_version(),
957 shard_id,
958 "host".into(),
959 0,
960 )
961 }
962 fn assert_same<K, V, T, D>(
963 state1: &LockingTypedState<K, V, T, D>,
964 state2: &LockingTypedState<K, V, T, D>,
965 ) {
966 let pointer1 = format!("{:p}", state1.state.read().expect("lock").deref());
967 let pointer2 = format!("{:p}", state2.state.read().expect("lock").deref());
968 assert_eq!(pointer1, pointer2);
969 }
970
971 let s1 = ShardId::new();
972 let states = Arc::new(StateCache::new_no_metrics());
973
974 assert_eq!(states.states.lock().expect("lock").len(), 0);
976
977 let s = Arc::clone(&states);
979 let res = spawn(|| "test", async move {
980 s.get::<(), (), u64, i64, _, _>(
981 s1,
982 || async { panic!("forced panic") },
983 &Diagnostics::for_tests(),
984 )
985 .await
986 })
987 .into_tokio_handle()
988 .await;
989 assert_err!(res);
990 assert_eq!(states.initialized_count(), 0);
991
992 let res = states
994 .get::<(), (), u64, i64, _, _>(
995 s1,
996 || async {
997 Err(Box::new(CodecMismatch {
998 requested: ("".into(), "".into(), "".into(), "".into(), None),
999 actual: ("".into(), "".into(), "".into(), "".into(), None),
1000 }))
1001 },
1002 &Diagnostics::for_tests(),
1003 )
1004 .await;
1005 assert_err!(res);
1006 assert_eq!(states.initialized_count(), 0);
1007
1008 let did_work = Arc::new(AtomicBool::new(false));
1010 let s1_state1 = states
1011 .get::<(), (), u64, i64, _, _>(
1012 s1,
1013 || {
1014 let did_work = Arc::clone(&did_work);
1015 async move {
1016 did_work.store(true, Ordering::SeqCst);
1017 Ok(new_state(s1))
1018 }
1019 },
1020 &Diagnostics::for_tests(),
1021 )
1022 .await
1023 .expect("should successfully initialize");
1024 assert_eq!(did_work.load(Ordering::SeqCst), true);
1025 assert_eq!(states.initialized_count(), 1);
1026 assert_eq!(states.strong_count(), 1);
1027
1028 let did_work = Arc::new(AtomicBool::new(false));
1030 let s1_state2 = states
1031 .get::<(), (), u64, i64, _, _>(
1032 s1,
1033 || {
1034 let did_work = Arc::clone(&did_work);
1035 async move {
1036 did_work.store(true, Ordering::SeqCst);
1037 did_work.store(true, Ordering::SeqCst);
1038 Ok(new_state(s1))
1039 }
1040 },
1041 &Diagnostics::for_tests(),
1042 )
1043 .await
1044 .expect("should successfully initialize");
1045 assert_eq!(did_work.load(Ordering::SeqCst), false);
1046 assert_eq!(states.initialized_count(), 1);
1047 assert_eq!(states.strong_count(), 1);
1048 assert_same(&s1_state1, &s1_state2);
1049
1050 let did_work = Arc::new(AtomicBool::new(false));
1052 let res = states
1053 .get::<String, (), u64, i64, _, _>(
1054 s1,
1055 || {
1056 let did_work = Arc::clone(&did_work);
1057 async move {
1058 did_work.store(true, Ordering::SeqCst);
1059 Ok(new_state(s1))
1060 }
1061 },
1062 &Diagnostics::for_tests(),
1063 )
1064 .await;
1065 assert_eq!(did_work.load(Ordering::SeqCst), false);
1066 assert_eq!(
1067 format!("{}", res.expect_err("types shouldn't match")),
1068 "requested codecs (\"String\", \"()\", \"u64\", \"i64\", Some(CodecConcreteType(\"(alloc::string::String, (), u64, i64)\"))) did not match ones in durable storage (\"()\", \"()\", \"u64\", \"i64\", Some(CodecConcreteType(\"((), (), u64, i64)\")))"
1069 );
1070 assert_eq!(states.initialized_count(), 1);
1071 assert_eq!(states.strong_count(), 1);
1072
1073 let s2 = ShardId::new();
1075 let s2_state1 = states
1076 .get::<String, (), u64, i64, _, _>(
1077 s2,
1078 || async { Ok(new_state(s2)) },
1079 &Diagnostics::for_tests(),
1080 )
1081 .await
1082 .expect("should successfully initialize");
1083 assert_eq!(states.initialized_count(), 2);
1084 assert_eq!(states.strong_count(), 2);
1085 let s2_state2 = states
1086 .get::<String, (), u64, i64, _, _>(
1087 s2,
1088 || async { Ok(new_state(s2)) },
1089 &Diagnostics::for_tests(),
1090 )
1091 .await
1092 .expect("should successfully initialize");
1093 assert_same(&s2_state1, &s2_state2);
1094
1095 drop(s1_state1);
1098 assert_eq!(states.strong_count(), 2);
1099 drop(s1_state2);
1100 assert_eq!(states.strong_count(), 1);
1101 assert_eq!(states.initialized_count(), 2);
1102 assert_none!(states.get_cached(&s1));
1103
1104 let s1_state1 = states
1106 .get::<(), (), u64, i64, _, _>(
1107 s1,
1108 || async { Ok(new_state(s1)) },
1109 &Diagnostics::for_tests(),
1110 )
1111 .await
1112 .expect("should successfully initialize");
1113 assert_eq!(states.initialized_count(), 2);
1114 assert_eq!(states.strong_count(), 2);
1115 drop(s1_state1);
1116 assert_eq!(states.strong_count(), 1);
1117 }
1118
1119 #[mz_ore::test(tokio::test(flavor = "multi_thread"))]
1120 #[cfg_attr(miri, ignore)] async fn state_cache_concurrency() {
1122 mz_ore::test::init_logging();
1123
1124 const COUNT: usize = 1000;
1125 let id = ShardId::new();
1126 let cache = StateCache::new_no_metrics();
1127 let diagnostics = Diagnostics::for_tests();
1128
1129 let mut futures = (0..COUNT)
1130 .map(|_| {
1131 cache.get::<(), (), u64, i64, _, _>(
1132 id,
1133 || async {
1134 Ok(TypedState::new(
1135 DUMMY_BUILD_INFO.semver_version(),
1136 id,
1137 "host".into(),
1138 0,
1139 ))
1140 },
1141 &diagnostics,
1142 )
1143 })
1144 .collect::<FuturesUnordered<_>>();
1145
1146 for _ in 0..COUNT {
1147 let _ = futures.next().await.unwrap();
1148 }
1149 }
1150
1151 #[mz_ore::test(tokio::test)]
1152 #[cfg_attr(miri, ignore)] async fn update_semaphore() {
1154 mz_ore::test::init_logging();
1157
1158 let shard_id = ShardId::new();
1159 let persist_config = Arc::new(PersistConfig::new_for_tests());
1160 let pubsub = Arc::new(NoopPubSubSender);
1161 let state: LockingTypedState<String, (), u64, i64> = LockingTypedState::new(
1162 shard_id,
1163 TypedState::new(
1164 DUMMY_BUILD_INFO.semver_version(),
1165 shard_id,
1166 "host".into(),
1167 0,
1168 ),
1169 Arc::new(Metrics::new(&*persist_config, &MetricsRegistry::new())),
1170 persist_config,
1171 pubsub.subscribe(&shard_id),
1172 &Diagnostics::for_tests(),
1173 );
1174
1175 let mk_future = || {
1178 let (tx, rx) = oneshot::channel();
1179 let future = async {
1180 let lease = state.lease_for_update().await;
1181 let () = rx.await.unwrap();
1182 drop(lease);
1183 };
1184 (future, tx)
1185 };
1186
1187 let (one, _one_tx) = mk_future();
1188 let (two, _two_tx) = mk_future();
1189 let (three, three_tx) = mk_future();
1190 let mut one = pin!(one);
1191 let mut two = pin!(two);
1192 let mut three = pin!(three);
1193
1194 tokio::select! { biased;
1196 _ = &mut one => { unreachable!() }
1197 _ = &mut two => { unreachable!() }
1198 _ = &mut three => { unreachable!() }
1199 _ = async {} => {}
1200 }
1201
1202 three_tx.send(()).unwrap();
1204
1205 tokio::select! { biased;
1208 _ = &mut one => { unreachable!() }
1209 _ = &mut three => { }
1210 }
1211 }
1212
1213 #[mz_ore::test(tokio::test(flavor = "multi_thread"))]
1214 #[cfg_attr(miri, ignore)] async fn update_semaphore_stress() {
1216 mz_ore::test::init_logging();
1219
1220 const TIMEOUT: Duration = Duration::from_millis(100);
1221 const COUNT: u64 = 100;
1222
1223 let shard_id = ShardId::new();
1224 let persist_config = Arc::new(PersistConfig::new_for_tests());
1225 persist_config.set_config(&STATE_UPDATE_LEASE_TIMEOUT, TIMEOUT);
1226 let pubsub = Arc::new(NoopPubSubSender);
1227 let state: LockingTypedState<String, (), u64, i64> = LockingTypedState::new(
1228 shard_id,
1229 TypedState::new(
1230 DUMMY_BUILD_INFO.semver_version(),
1231 shard_id,
1232 "host".into(),
1233 0,
1234 ),
1235 Arc::new(Metrics::new(&*persist_config, &MetricsRegistry::new())),
1236 persist_config,
1237 pubsub.subscribe(&shard_id),
1238 &Diagnostics::for_tests(),
1239 );
1240
1241 let mut futures = (0..(COUNT * 3))
1242 .map(async |i| {
1243 state.lease_for_update().await;
1244 match i % 3 {
1246 0 => {
1247 let () = std::future::pending().await;
1248 }
1249 1 => {
1250 tokio::time::sleep(Duration::from_millis(i)).await;
1251 }
1252 _ => {
1253 tokio::time::sleep(Duration::from_millis(i) + TIMEOUT).await;
1254 }
1255 }
1256 })
1257 .collect::<FuturesUnordered<_>>();
1258
1259 for _ in 0..(COUNT * 2) {
1261 futures.next().await.unwrap();
1262 }
1263 }
1264}