tuwunel_service/globals/
data.rs1use 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
16pub struct Data {
21 global: Arc<Map>,
22 retires: Sender<u64>,
23 counter: Arc<Counter>,
24 pub(super) db: Arc<Database>,
26}
27
28pub(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 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 #[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 #[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 #[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 #[inline]
109 pub(super) fn current_count(&self) -> u64 { self.counter.current() }
110
111 #[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 pub fn bump_database_version(&self, new_version: u64) {
145 self.global.raw_put(b"version", new_version);
146 }
147
148 pub async fn database_version(&self) -> u64 {
152 self.global
153 .get(b"version")
154 .await
155 .deserialized()
156 .unwrap_or(0)
157 }
158}