Skip to main content

mz_persist_client/
cache.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10//! A cache of [PersistClient]s indexed by [PersistLocation]s.
11
12use 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/// A cache of [PersistClient]s indexed by [PersistLocation]s.
50///
51/// There should be at most one of these per process. All production
52/// PersistClients should be created through this cache.
53///
54/// This is because, in production, persist is heavily limited by the number of
55/// server-side Postgres/Aurora connections. This cache allows PersistClients to
56/// share, for example, these Postgres connections.
57#[derive(Debug)]
58pub struct PersistClientCache {
59    /// The tunable knobs for persist.
60    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    /// Returns a new [PersistClientCache].
75    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    /// A test helper that returns a [PersistClientCache] disconnected from
107    /// metrics.
108    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    /// Create a [PersistClientCache] for use in turmoil tests.
118    ///
119    /// Turmoil wants to run all software under test in a single thread, so we disable the
120    /// (multi-threaded) isolated runtime.
121    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    /// Returns the [PersistConfig] being used by this cache.
150    pub fn cfg(&self) -> &PersistConfig {
151        &self.cfg
152    }
153
154    /// Returns persist `Metrics`.
155    pub fn metrics(&self) -> &Arc<Metrics> {
156        &self.metrics
157    }
158
159    /// Returns `ShardMetrics` for the given shard.
160    pub fn shard_metrics(&self, shard_id: &ShardId, name: &str) -> Arc<ShardMetrics> {
161        self.metrics.shards.shard(shard_id, name)
162    }
163
164    /// Clears the state cache, allowing for tests with disconnected states.
165    ///
166    /// Only exposed for testing.
167    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    /// Returns a new [PersistClient] for interfacing with persist shards made
176    /// durable to the given [PersistLocation].
177    ///
178    /// The same `location` may be used concurrently from multiple processes.
179    #[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    // No sense in measuring rtt latencies more often than this.
195    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                // Intentionally hold the lock, so we don't double connect under
206                // concurrency.
207                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                // Intentionally hold the lock, so we don't double connect under
242                // concurrency.
243                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                // Hedged gets need a second handle on an isolated connection
254                // pool. Built unconditionally (the wrapper reads its enable
255                // flag dynamically per call) and armed in the background, so
256                // the sibling open never runs under this lock.
257                //
258                // NOTE: HedgedBlob must stay below Tasked in this stack. Its
259                // race relies on dropping the losing future to cancel the
260                // request in flight, and on hedged gets running to
261                // completion once started (Tasked detaches). A task boundary
262                // between HedgedBlob and the backend would break the former,
263                // and an aborting layer above would slowly leak budget
264                // tokens via the latter.
265                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                // This is intentionally "outside" (wrapping) MetricsBlob so
291                // that we don't include cached responses in blob metrics.
292                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/// Starts a task to periodically measure the persist-observed latency to
301/// consensus.
302///
303/// This is a task, rather than something like looking at the latencies of prod
304/// traffic, so that we minimize any issues around Futures not being polled
305/// promptly (as can and does happen with the Timely-polled Futures).
306///
307/// The caller is responsible for shutdown via aborting the `JoinHandle`.
308///
309/// No matter whether we wrap MetricsConsensus before or after we start up the
310/// rtt latency task, there's the possibility for it being confusing at some
311/// point. Err on the side of more data (including the latency measurements) to
312/// start.
313#[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        // Use the tokio Instant for next_measurement because the reclock tests
321        // mess with the tokio sleep clock.
322        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                    // Don't spam retries if this returns an error. We're
332                    // guaranteed by the method signature that we've already got
333                    // metrics coverage of these, so we'll count the errors.
334                }
335            }
336            next_measurement = tokio::time::Instant::now() + measurement_interval;
337        }
338    })
339}
340
341/// Starts a task to periodically measure the persist-observed latency to
342/// consensus.
343///
344/// This is a task, rather than something like looking at the latencies of prod
345/// traffic, so that we minimize any issues around Futures not being polled
346/// promptly (as can and does happen with the Timely-polled Futures).
347///
348/// The caller is responsible for shutdown via aborting the `JoinHandle`.
349///
350/// No matter whether we wrap MetricsConsensus before or after we start up the
351/// rtt latency task, there's the possibility for it being confusing at some
352/// point. Err on the side of more data (including the latency measurements) to
353/// start.
354#[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        // Use the tokio Instant for next_measurement because the reclock tests
362        // mess with the tokio sleep clock.
363        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                    // Don't spam retries if this returns an error. We're
376                    // guaranteed by the method signature that we've already got
377                    // metrics coverage of these, so we'll count the errors.
378                }
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/// A cache of `TypedState`, shared between all machines for that shard.
449///
450/// This is shared between all machines that come out of the same
451/// [PersistClientCache], but in production there is one of those per process,
452/// so in practice, we have one copy of state per shard per process.
453///
454/// The mutex contention between commands is not an issue, because if two
455/// command for the same shard are executing concurrently, only one can win
456/// anyway, the other will retry. With the mutex, we even get to avoid the retry
457/// if the racing commands are on the same process.
458#[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    /// Returns a new StateCache.
474    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                            // If the Weak has lost the ability to upgrade,
522                            // we've dropped the State and it's gone. Clear the
523                            // OnceCell and init a new one.
524                            *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                        // We actually did the init work, don't bother casting back
555                        // the type erased and weak version. Additionally, inform
556                        // any listeners of this new state.
557                        return Ok(x);
558                    }
559                    let Some(state) = state.upgrade() else {
560                        // Race condition. Between when we first checked the
561                        // OnceCell and the `get_or_try_init` call, (1) the
562                        // initialization finished, (2) the other user dropped
563                        // the strong ref, and (3) the Arc noticed it was down
564                        // to only weak refs and dropped the value. Nothing we
565                        // can do except try again.
566                        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
633/// A locked decorator for TypedState that abstracts out the specific lock implementation used.
634/// Guards the private lock with public accessor fns to make locking scopes more explicit and
635/// simpler to reason about.
636pub(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    // Retained only to keep this shard's per-shard series registered for as long
643    // as the state is cached; nothing reads it through this handle anymore. Don't
644    // drop it as "unused" without moving that lifetime guarantee elsewhere.
645    shard_metrics: Arc<ShardMetrics>,
646    update_semaphore: AwaitableState<Option<tokio::time::Instant>>,
647    /// A [SchemaCacheMaps<K, V>], but stored as an Any so the `: Codec` bounds
648    /// don't propagate to basically every struct in persist.
649    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        // Bounds needed to seed the notifier with the shard's current upper.
685        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        // Bounds needed to read the shard upper. All callers are in contexts that
749        // already satisfy these (they mutate a fully-typed shard state).
750        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            // The notifier only advances the upper waiters' signal on a strict
775            // upper advance. The seqno bumps for many non-data reasons (GC,
776            // rollups, since-downgrades, other writers' CaAs) and those must not
777            // re-activate upper waiters.
778            self.notifier.notify(seqno_after, state.upper());
779        }
780        // For now, make sure to notify while under lock. It's possible to move
781        // this out of the lock window, see [StateWatchNotifier::notify].
782        drop(state);
783        ret
784    }
785
786    /// We want to _mostly_ just attempt a single CaS against the same state at once, since
787    /// only one concurrent CaS can succeed. However, we also want to guard against a
788    /// single hung update blocking all progress globally. We manage this with a shared state,
789    /// tracking whether a request is in flight and when it times out. If the timeout is never hit,
790    /// this behaves like a semaphore with limit 1... but if our requests _are_ timing out, future
791    /// requests will only wait for a bounded time before retrying, and one of those retries will
792    /// be able to claim that lease and make progress.
793    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                    // Clear the timeout if it hasn't changed since we set it.
804                    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        // Special case: if the timeout is set to zero, go ahead without taking a lease.
814        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            // Claim the lease if there isn't one, or if the current lease has expired.
823            let maybe_leased = timeout_state.maybe_modify(|state| {
824                if let Some(other_expires_at) = **state
825                    && other_expires_at > now
826                {
827                    // Still locked: sleep until the deadline and try again.
828                    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                    // Wait until either the lease has dropped or timed out, whichever is first.
841                    // If there are a lot of clients trying to update the same state, this may
842                    // cause significant lock contention... but the lock is only briefly held,
843                    // and anyways that's still cheaper than contending on the remote database.
844                    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)] // unsupported operation: returning ready events from epoll_wait is not yet implemented
876    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        // Opening a location on an empty cache saves the results.
886        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        // Opening a location with an already opened consensus reuses it, even
897        // if the blob is different.
898        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        // Ditto the other way.
909        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        // Query params and path matter, so we get new instances.
920        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        // User info and port also matter, so we get new instances.
932        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)] // unsupported operation: returning ready events from epoll_wait is not yet implemented
946    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        // The cache starts empty.
975        assert_eq!(states.states.lock().expect("lock").len(), 0);
976
977        // Panic'ing during init_fn .
978        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        // Returning an error from init_fn doesn't initialize an entry in the cache.
993        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        // Initialize one shard.
1009        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        // Trying to initialize it again does no work and returns the same state.
1029        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        // Trying to initialize with different types doesn't work.
1051        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        // We can add a shard of a different type.
1074        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        // The cache holds weak references to State so we reclaim memory if the
1096        // shards stops being used.
1097        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        // But we can re-init that shard if necessary.
1105        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)] // too slow
1121    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)] // too slow
1153    async fn update_semaphore() {
1154        // Check that the update lease mechanism is not susceptible to futurelock.
1155        // If there is an issue, this test will time out.
1156        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        // Initialize three futures, all of which will grab a lease and then poll a oneshot,
1176        // which allows us to externally trigger which ones will complete.
1177        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        // Poll all the futures, but fall through to the default case, since none are ready.
1195        tokio::select! { biased;
1196            _ = &mut one => { unreachable!() }
1197            _ = &mut two => { unreachable!() }
1198            _ = &mut three => { unreachable!() }
1199            _ = async {} => {}
1200        }
1201
1202        // Allow the third future to complete.
1203        three_tx.send(()).unwrap();
1204
1205        // Poll all the futures but the second future. This shouldn't hang, since the third future
1206        // is now ready to go and the others should eventually time out.
1207        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)] // too slow
1215    async fn update_semaphore_stress() {
1216        // Check that the update lease mechanism is not susceptible to futurelock.
1217        // If there is an issue, this test will time out.
1218        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                // Either hang forever, succeed quickly, or succeed after hitting the timeout.
1245                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        // All the futures that don't themselves hang forever should resolve.
1260        for _ in 0..(COUNT * 2) {
1261            futures.next().await.unwrap();
1262        }
1263    }
1264}