Skip to main content

tuwunel_service/globals/
data.rs

1//! Persistent global counter and database version state.
2//!
3//! The counter records sequence numbers before handing them to writers and tracks their ordered
4//! retirement through permits. The same map stores the database schema version used by migrations.
5
6use std::{ops::Range, sync::Arc};
7
8use futures::TryFutureExt;
9use tokio::sync::watch::Sender;
10use tuwunel_core::{
11	Result, err, utils,
12	utils::two_phase_counter::{Counter as TwoPhaseCounter, Permit as TwoPhasePermit},
13};
14use tuwunel_database::{Database, Deserialized, Map};
15
16/// Owns persistent global state and the two-phase sequence counter.
17///
18/// Dispatched counter values are recorded before use. A watch channel publishes the retirement
19/// frontier as outstanding permits are dropped.
20pub struct Data {
21	global: Arc<Map>,
22	retires: Sender<u64>,
23	counter: Arc<Counter>,
24	/// Database handle used to report process-wide read-only state.
25	pub(super) db: Arc<Database>,
26}
27
28/// Permit guarding one dispatched global sequence number.
29///
30/// The permit exposes the allocated number and retires it on drop. Retirement advances in dispatch
31/// order even when later permits finish first.
32pub(super) type Permit = TwoPhasePermit<Callback>;
33type Counter = TwoPhaseCounter<Callback>;
34type Callback = Box<dyn Fn(u64) -> Result + Send + Sync>;
35
36const COUNTER: &[u8] = b"c";
37
38impl Data {
39	/// Restores the global counter and initializes retirement notifications.
40	///
41	/// A fresh database starts at zero, making one the first dispatched sequence number.
42	///
43	/// # Panics
44	///
45	/// Panics when a successfully read stored counter cannot be decoded.
46	pub(super) fn new(args: &crate::Args<'_>) -> Self {
47		let db = args.db.clone();
48		let count = Self::stored_count(&args.db["global"]).expect("initialize global counter");
49		let retires = Sender::new(count);
50		Self {
51			db: args.db.clone(),
52			global: args.db["global"].clone(),
53			retires: retires.clone(),
54			counter: Counter::new(
55				count,
56				Box::new(move |count| Self::store_count(&db, &db["global"], count)),
57				Box::new(move |count| Self::handle_retire(&retires, count)),
58			),
59		}
60	}
61
62	/// Waits for all sequence numbers dispatched at call time to retire.
63	///
64	/// The dispatched frontier is sampled before subscribing. The returned frontier is guaranteed
65	/// to be at least that sample.
66	#[inline]
67	pub(super) async fn wait_pending(&self) -> Result<u64> {
68		let count = self.counter.dispatched();
69		self.wait_count(&count).await.inspect(|retired| {
70			debug_assert!(
71				*retired >= count,
72				"Expecting retired sequence number >= snapshotted dispatch number"
73			);
74		})
75	}
76
77	/// Waits until the retirement frontier reaches `count`.
78	///
79	/// The returned frontier may exceed `count` when additional permits retire before the waiter is
80	/// notified.
81	#[inline]
82	pub(super) async fn wait_count(&self, count: &u64) -> Result<u64> {
83		self.retires
84			.subscribe()
85			.wait_for(|retired| retired.ge(count))
86			.map_ok(|retired| *retired)
87			.map_err(|e| err!(debug_error!("counter channel error {e:?}")))
88			.await
89	}
90
91	/// Dispatches the next sequence number and returns its retirement permit.
92	///
93	/// Dispatch records the new number before exposing it to the caller.
94	///
95	/// # Panics
96	///
97	/// Panics when the counter is exhausted or the dispatched value cannot be recorded.
98	#[inline]
99	pub(super) fn next_count(&self) -> Permit {
100		self.counter
101			.next()
102			.expect("failed to obtain next sequence number")
103	}
104
105	/// Returns the highest fully retired sequence number.
106	///
107	/// All writes through this frontier are safe for readers to observe.
108	#[inline]
109	pub(super) fn current_count(&self) -> u64 { self.counter.current() }
110
111	/// Returns the retired-to-dispatched counter range.
112	///
113	/// The start is the reader-visible frontier and the end is the latest dispatched value.
114	#[inline]
115	pub(super) fn pending_count(&self) -> Range<u64> { self.counter.range() }
116
117	#[tracing::instrument(name = "retire", level = "debug", skip(sender))]
118	fn handle_retire(sender: &Sender<u64>, count: u64) -> Result {
119		let _prev = sender.send_replace(count);
120
121		Ok(())
122	}
123
124	#[tracing::instrument(name = "dispatch", level = "debug", skip(db, global))]
125	fn store_count(db: &Arc<Database>, global: &Arc<Map>, count: u64) -> Result {
126		let _cork = db.cork();
127		global.insert(COUNTER, count.to_be_bytes());
128
129		Ok(())
130	}
131
132	fn stored_count(global: &Arc<Map>) -> Result<u64> {
133		global
134			.get_blocking(COUNTER)
135			.as_deref()
136			.map_or(Ok(0_u64), utils::u64_from_bytes)
137	}
138}
139
140impl Data {
141	/// Stores a new database schema version.
142	///
143	/// The value replaces the existing version in the global metadata map.
144	pub fn bump_database_version(&self, new_version: u64) {
145		self.global.raw_put(b"version", new_version);
146	}
147
148	/// Loads the current database schema version.
149	///
150	/// Missing or undecodable values are treated as version zero.
151	pub async fn database_version(&self) -> u64 {
152		self.global
153			.get(b"version")
154			.await
155			.deserialized()
156			.unwrap_or(0)
157	}
158}