1use crate::store_utils::{DEFAULT_TIMEOUT, get_with_timeout, put_with_timeout};
11use anyhow::Result;
12use bytes::Bytes;
13use object_store::path::Path;
14use object_store::{ObjectStore, PutMode, PutOptions, UpdateVersion};
15use serde::{Deserialize, Serialize};
16use std::sync::Arc;
17use tokio::sync::Mutex;
18use uni_common::core::id::{Eid, Vid};
19
20#[derive(Serialize, Deserialize, Default, Clone)]
22struct CounterManifest {
23 next_vid_batch: u64,
25 next_eid_batch: u64,
27}
28
29struct AllocatorState {
31 manifest: CounterManifest,
32 manifest_version: Option<String>, current_vid: u64,
34 current_eid: u64,
35}
36
37pub struct IdAllocator {
44 store: Arc<dyn ObjectStore>,
45 path: Path,
46 state: Mutex<AllocatorState>,
47 batch_size: u64,
48}
49
50impl IdAllocator {
51 pub async fn new(store: Arc<dyn ObjectStore>, path: Path, batch_size: u64) -> Result<Self> {
53 let (manifest, version) = match get_with_timeout(&store, &path, DEFAULT_TIMEOUT).await {
54 Ok(get_result) => {
55 let version = get_result.meta.e_tag.clone();
56 let bytes = get_result.bytes().await?;
57 let manifest: CounterManifest = serde_json::from_slice(&bytes)?;
58 (manifest, version)
59 }
60 Err(e) if e.to_string().contains("not found") => (CounterManifest::default(), None),
61 Err(e) => return Err(e),
62 };
63
64 let current_vid = manifest.next_vid_batch;
66 let current_eid = manifest.next_eid_batch;
67
68 Ok(Self {
69 store,
70 path,
71 state: Mutex::new(AllocatorState {
72 manifest,
73 manifest_version: version,
74 current_vid,
75 current_eid,
76 }),
77 batch_size,
78 })
79 }
80
81 pub async fn allocate_vid(&self) -> Result<Vid> {
85 let mut state = self.state.lock().await;
86
87 if state.current_vid >= state.manifest.next_vid_batch {
89 let reserved = state
94 .current_vid
95 .checked_add(self.batch_size)
96 .ok_or_else(|| anyhow::anyhow!("VID space exhausted"))?;
97 let prev = state.manifest.next_vid_batch;
98 state.manifest.next_vid_batch = reserved;
99 if let Err(e) = self.persist_manifest(&mut state).await {
100 state.manifest.next_vid_batch = prev;
105 return Err(e);
106 }
107 }
108
109 let vid = Vid::new(state.current_vid);
110 state.current_vid += 1;
111 Ok(vid)
112 }
113
114 pub async fn allocate_vids(&self, count: usize) -> Result<Vec<Vid>> {
116 let mut state = self.state.lock().await;
117 let needed = count as u64;
118
119 let want = state
121 .current_vid
122 .checked_add(needed)
123 .ok_or_else(|| anyhow::anyhow!("VID space exhausted"))?;
124 if want > state.manifest.next_vid_batch {
125 let reserved = want
127 .checked_add(self.batch_size)
128 .ok_or_else(|| anyhow::anyhow!("VID space exhausted"))?;
129 let prev = state.manifest.next_vid_batch;
130 state.manifest.next_vid_batch = reserved;
131 if let Err(e) = self.persist_manifest(&mut state).await {
132 state.manifest.next_vid_batch = prev;
134 return Err(e);
135 }
136 }
137
138 let vids: Vec<Vid> = (0..count)
139 .map(|i| Vid::new(state.current_vid + i as u64))
140 .collect();
141 state.current_vid += needed;
142 Ok(vids)
143 }
144
145 pub async fn allocate_eid(&self) -> Result<Eid> {
149 let mut state = self.state.lock().await;
150
151 if state.current_eid >= state.manifest.next_eid_batch {
153 let reserved = state
155 .current_eid
156 .checked_add(self.batch_size)
157 .ok_or_else(|| anyhow::anyhow!("EID space exhausted"))?;
158 let prev = state.manifest.next_eid_batch;
159 state.manifest.next_eid_batch = reserved;
160 if let Err(e) = self.persist_manifest(&mut state).await {
161 state.manifest.next_eid_batch = prev;
163 return Err(e);
164 }
165 }
166
167 let eid = Eid::new(state.current_eid);
168 state.current_eid += 1;
169 Ok(eid)
170 }
171
172 pub async fn allocate_eids(&self, count: usize) -> Result<Vec<Eid>> {
174 let mut state = self.state.lock().await;
175 let needed = count as u64;
176
177 let want = state
179 .current_eid
180 .checked_add(needed)
181 .ok_or_else(|| anyhow::anyhow!("EID space exhausted"))?;
182 if want > state.manifest.next_eid_batch {
183 let reserved = want
185 .checked_add(self.batch_size)
186 .ok_or_else(|| anyhow::anyhow!("EID space exhausted"))?;
187 let prev = state.manifest.next_eid_batch;
188 state.manifest.next_eid_batch = reserved;
189 if let Err(e) = self.persist_manifest(&mut state).await {
190 state.manifest.next_eid_batch = prev;
192 return Err(e);
193 }
194 }
195
196 let eids: Vec<Eid> = (0..count)
197 .map(|i| Eid::new(state.current_eid + i as u64))
198 .collect();
199 state.current_eid += needed;
200 Ok(eids)
201 }
202
203 pub async fn current_vid(&self) -> u64 {
205 self.state.lock().await.current_vid
206 }
207
208 pub async fn current_eid(&self) -> u64 {
210 self.state.lock().await.current_eid
211 }
212
213 pub async fn current_hwm(&self) -> (u64, u64) {
222 let state = self.state.lock().await;
223 (state.current_vid, state.current_eid)
224 }
225
226 pub async fn checkpoint(&self) -> Result<()> {
244 let mut state = self.state.lock().await;
245 if state.manifest.next_vid_batch < state.current_vid {
248 state.manifest.next_vid_batch = state.current_vid;
249 }
250 if state.manifest.next_eid_batch < state.current_eid {
251 state.manifest.next_eid_batch = state.current_eid;
252 }
253 self.persist_manifest(&mut state).await
254 }
255
256 async fn persist_manifest(&self, state: &mut AllocatorState) -> Result<()> {
258 let json = serde_json::to_vec_pretty(&state.manifest)?;
259 let bytes = Bytes::from(json);
260
261 let put_result = if let Some(version) = &state.manifest_version {
264 let opts: PutOptions = PutMode::Update(UpdateVersion {
265 e_tag: Some(version.clone()),
266 version: None,
267 })
268 .into();
269 match tokio::time::timeout(
270 DEFAULT_TIMEOUT,
271 self.store.put_opts(&self.path, bytes.clone().into(), opts),
272 )
273 .await
274 {
275 Ok(Ok(result)) => result,
276 Ok(Err(e))
277 if e.to_string().contains("not yet implemented")
278 || e.to_string().contains("not supported") =>
279 {
280 put_with_timeout(&self.store, &self.path, bytes, DEFAULT_TIMEOUT).await?
282 }
283 Ok(Err(e)) => return Err(e.into()),
284 Err(_) => {
285 return Err(anyhow::anyhow!(
286 "Object store put_opts timed out after {:?}",
287 DEFAULT_TIMEOUT
288 ));
289 }
290 }
291 } else {
292 let opts: PutOptions = PutMode::Create.into();
294 match tokio::time::timeout(
295 DEFAULT_TIMEOUT,
296 self.store.put_opts(&self.path, bytes.clone().into(), opts),
297 )
298 .await
299 {
300 Ok(Ok(result)) => result,
301 Ok(Err(object_store::Error::AlreadyExists { .. })) => {
302 put_with_timeout(&self.store, &self.path, bytes, DEFAULT_TIMEOUT).await?
304 }
305 Ok(Err(e)) if e.to_string().contains("not yet implemented") => {
306 put_with_timeout(&self.store, &self.path, bytes, DEFAULT_TIMEOUT).await?
307 }
308 Ok(Err(e)) => return Err(e.into()),
309 Err(_) => {
310 return Err(anyhow::anyhow!(
311 "Object store put_opts timed out after {:?}",
312 DEFAULT_TIMEOUT
313 ));
314 }
315 }
316 };
317
318 state.manifest_version = put_result.e_tag;
319 Ok(())
320 }
321}
322
323#[cfg(test)]
324mod tests {
325 use super::*;
326 use object_store::memory::InMemory;
327
328 #[tokio::test]
329 async fn test_allocate_vid() {
330 let store = Arc::new(InMemory::new());
331 let path = Path::from("id_counters.json");
332 let allocator = IdAllocator::new(store, path, 100).await.unwrap();
333
334 let vid1 = allocator.allocate_vid().await.unwrap();
335 let vid2 = allocator.allocate_vid().await.unwrap();
336 let vid3 = allocator.allocate_vid().await.unwrap();
337
338 assert_eq!(vid1.as_u64(), 0);
339 assert_eq!(vid2.as_u64(), 1);
340 assert_eq!(vid3.as_u64(), 2);
341 }
342
343 #[tokio::test]
344 async fn test_allocate_eid() {
345 let store = Arc::new(InMemory::new());
346 let path = Path::from("id_counters.json");
347 let allocator = IdAllocator::new(store, path, 100).await.unwrap();
348
349 let eid1 = allocator.allocate_eid().await.unwrap();
350 let eid2 = allocator.allocate_eid().await.unwrap();
351
352 assert_eq!(eid1.as_u64(), 0);
353 assert_eq!(eid2.as_u64(), 1);
354 }
355
356 #[tokio::test]
357 async fn test_allocate_many() {
358 let store = Arc::new(InMemory::new());
359 let path = Path::from("id_counters.json");
360 let allocator = IdAllocator::new(store, path, 100).await.unwrap();
361
362 let vids = allocator.allocate_vids(5).await.unwrap();
363 assert_eq!(vids.len(), 5);
364 for (i, vid) in vids.iter().enumerate() {
365 assert_eq!(vid.as_u64(), i as u64);
366 }
367
368 let next = allocator.allocate_vid().await.unwrap();
370 assert_eq!(next.as_u64(), 5);
371 }
372
373 #[tokio::test]
374 async fn test_persistence() {
375 let store = Arc::new(InMemory::new());
376 let path = Path::from("id_counters.json");
377
378 {
380 let allocator = IdAllocator::new(store.clone(), path.clone(), 10)
381 .await
382 .unwrap();
383 for _ in 0..15 {
384 allocator.allocate_vid().await.unwrap();
385 }
386 }
387
388 {
390 let allocator = IdAllocator::new(store, path, 10).await.unwrap();
391 let vid = allocator.allocate_vid().await.unwrap();
394 assert_eq!(vid.as_u64(), 20);
395 }
396 }
397}