1use std::path::{Path, PathBuf};
47use std::sync::atomic::Ordering;
48use std::sync::Arc;
49use std::thread;
50use std::time::{Duration, Instant};
51
52use crate::shared_atomic::{SharedAtomicError, SharedAtomicU32, SharedAtomicU64};
53
54#[derive(Debug, Clone, Copy, PartialEq, Eq)]
55pub enum SemaphoreError {
56 Atomic(SharedAtomicError),
57 WouldBlock,
58 Timeout,
59 ReleaseOverflow,
60}
61
62impl From<SharedAtomicError> for SemaphoreError {
63 fn from(e: SharedAtomicError) -> Self { Self::Atomic(e) }
64}
65
66fn count_path(base: &Path) -> PathBuf {
67 let mut p = base.to_path_buf();
68 let stem = p.file_name().unwrap().to_string_lossy().to_string();
69 p.set_file_name(format!("{stem}.count.bin"));
70 p
71}
72fn wakeup_path(base: &Path) -> PathBuf {
73 let mut p = base.to_path_buf();
74 let stem = p.file_name().unwrap().to_string_lossy().to_string();
75 p.set_file_name(format!("{stem}.wakeup.bin"));
76 p
77}
78fn waiters_path(base: &Path) -> PathBuf {
79 let mut p = base.to_path_buf();
80 let stem = p.file_name().unwrap().to_string_lossy().to_string();
81 p.set_file_name(format!("{stem}.waiters.bin"));
82 p
83}
84
85pub struct SharedSemaphore {
86 count: Arc<SharedAtomicU32>,
87 wakeup: Arc<SharedAtomicU64>,
88 waiters: Arc<SharedAtomicU32>,
89 max_permits: u32,
90 header_sidecar: subetha_core::HandshakeHeader,
91 ring_sidecar: Box<subetha_core::ObservationRing>,
92}
93
94impl subetha_sidecar::AdaptiveInstance for SharedSemaphore {
95 fn header(&self) -> &subetha_core::HandshakeHeader { &self.header_sidecar }
96 fn ring(&self) -> &subetha_core::ObservationRing { &self.ring_sidecar }
97 fn make_policy(&self) -> Box<dyn subetha_sidecar::Policy> {
98 Box::new(subetha_sidecar::NoMigrationPolicy)
99 }
100}
101
102impl SharedSemaphore {
103 pub fn create(
107 base_path: impl AsRef<Path>,
108 initial: u32,
109 max_permits: u32,
110 ) -> Result<Self, SemaphoreError> {
111 assert!(initial <= max_permits, "initial permits must be <= max_permits");
112 let base = base_path.as_ref();
113 let count = Arc::new(SharedAtomicU32::create(count_path(base), initial)?);
114 let wakeup = Arc::new(SharedAtomicU64::create(wakeup_path(base), 0)?);
115 let waiters = Arc::new(SharedAtomicU32::create(waiters_path(base), 0)?);
116 Ok(Self {
117 count, wakeup, waiters, max_permits,
118 header_sidecar: subetha_core::HandshakeHeader::new(),
119 ring_sidecar: Box::new(subetha_core::ObservationRing::new()),
120 })
121 }
122
123 pub fn open(
128 base_path: impl AsRef<Path>,
129 max_permits: u32,
130 ) -> Result<Self, SemaphoreError> {
131 let base = base_path.as_ref();
132 let count = Arc::new(SharedAtomicU32::open(count_path(base))?);
133 let wakeup = Arc::new(SharedAtomicU64::open(wakeup_path(base))?);
134 let waiters = Arc::new(SharedAtomicU32::open(waiters_path(base))?);
135 Ok(Self {
136 count, wakeup, waiters, max_permits,
137 header_sidecar: subetha_core::HandshakeHeader::new(),
138 ring_sidecar: Box::new(subetha_core::ObservationRing::new()),
139 })
140 }
141
142 pub fn try_acquire(&self) -> Result<Permit<'_>, SemaphoreError> {
145 loop {
146 let cur = self.count.load(Ordering::Acquire);
147 if cur == 0 {
148 self.ring_sidecar
149 .push_op(crate::sidecar_ops::semaphore::OP_TRY_ACQUIRE, 1); return Err(SemaphoreError::WouldBlock);
151 }
152 match self.count.compare_exchange(
153 cur, cur - 1, Ordering::AcqRel, Ordering::Acquire,
154 ) {
155 Ok(_) => {
156 self.ring_sidecar
157 .push_op(crate::sidecar_ops::semaphore::OP_TRY_ACQUIRE, 0);
158 return Ok(Permit { sem: self });
159 }
160 Err(_) => continue, }
162 }
163 }
164
165 pub fn acquire(&self) -> Permit<'_> {
168 let mut had_contention = false;
170 loop {
171 let cur = self.count.load(Ordering::Acquire);
172 if cur > 0 {
173 if self.count.compare_exchange(
174 cur, cur - 1, Ordering::AcqRel, Ordering::Acquire,
175 ).is_ok() {
176 self.ring_sidecar.push_op(
177 crate::sidecar_ops::semaphore::OP_ACQUIRE,
178 if had_contention { 1 } else { 0 },
179 );
180 return Permit { sem: self };
181 }
182 continue;
183 }
184 had_contention = true;
185 self.waiters.fetch_add(1, Ordering::AcqRel);
187 let snapshot = self.wakeup.load(Ordering::Acquire);
188 if self.count.load(Ordering::Acquire) > 0 {
190 self.waiters.fetch_sub(1, Ordering::AcqRel);
191 continue;
192 }
193 let mut spins = 0u32;
195 loop {
196 let cur_count = self.count.load(Ordering::Acquire);
197 let cur_gen = self.wakeup.load(Ordering::Acquire);
198 if cur_count > 0 || cur_gen != snapshot {
199 self.waiters.fetch_sub(1, Ordering::AcqRel);
200 break;
201 }
202 spins += 1;
203 if spins < 32 {
204 std::hint::spin_loop();
205 } else if spins < 256 {
206 thread::yield_now();
207 } else {
208 thread::sleep(Duration::from_micros(50));
209 }
210 }
211 }
212 }
213
214 pub fn acquire_timeout(&self, timeout: Duration) -> Result<Permit<'_>, SemaphoreError> {
217 let deadline = Instant::now() + timeout;
218 let mut had_contention = false;
219 loop {
220 let cur = self.count.load(Ordering::Acquire);
221 if cur > 0 {
222 if self.count.compare_exchange(
223 cur, cur - 1, Ordering::AcqRel, Ordering::Acquire,
224 ).is_ok() {
225 self.ring_sidecar.push_op(
226 crate::sidecar_ops::semaphore::OP_ACQUIRE,
227 if had_contention { 1 } else { 0 },
228 );
229 return Ok(Permit { sem: self });
230 }
231 continue;
232 }
233 had_contention = true;
234 if Instant::now() >= deadline {
235 self.ring_sidecar
236 .push_op(crate::sidecar_ops::semaphore::OP_ACQUIRE, 1); return Err(SemaphoreError::Timeout);
238 }
239 self.waiters.fetch_add(1, Ordering::AcqRel);
240 let snapshot = self.wakeup.load(Ordering::Acquire);
241 if self.count.load(Ordering::Acquire) > 0 {
242 self.waiters.fetch_sub(1, Ordering::AcqRel);
243 continue;
244 }
245 let mut spins = 0u32;
246 loop {
247 let cur_count = self.count.load(Ordering::Acquire);
248 let cur_gen = self.wakeup.load(Ordering::Acquire);
249 if cur_count > 0 || cur_gen != snapshot {
250 self.waiters.fetch_sub(1, Ordering::AcqRel);
251 break;
252 }
253 if Instant::now() >= deadline {
254 self.waiters.fetch_sub(1, Ordering::AcqRel);
255 self.ring_sidecar
256 .push_op(crate::sidecar_ops::semaphore::OP_ACQUIRE, 1); return Err(SemaphoreError::Timeout);
258 }
259 spins += 1;
260 if spins < 32 {
261 std::hint::spin_loop();
262 } else if spins < 256 {
263 thread::yield_now();
264 } else {
265 thread::sleep(Duration::from_micros(50));
266 }
267 }
268 }
269 }
270
271 pub fn release(&self) -> Result<(), SemaphoreError> {
277 let prev = self.count.fetch_add(1, Ordering::AcqRel);
278 if prev >= self.max_permits {
279 self.count.fetch_sub(1, Ordering::AcqRel);
282 self.ring_sidecar
283 .push_op(crate::sidecar_ops::semaphore::OP_RELEASE, 1); return Err(SemaphoreError::ReleaseOverflow);
285 }
286 if self.waiters.load(Ordering::Acquire) > 0 {
287 self.wakeup.fetch_add(1, Ordering::Release);
288 }
289 self.ring_sidecar
290 .push_op(crate::sidecar_ops::semaphore::OP_RELEASE, 0);
291 Ok(())
292 }
293
294 #[inline]
296 pub fn available(&self) -> u32 {
297 self.count.load(Ordering::Acquire)
298 }
299
300 #[inline]
302 pub fn waiters(&self) -> u32 {
303 self.waiters.load(Ordering::Acquire)
304 }
305
306 #[inline]
308 pub fn max_permits(&self) -> u32 { self.max_permits }
309
310 #[inline]
317 pub fn wakeup_generation(&self) -> u64 {
318 self.wakeup.load(Ordering::Acquire)
319 }
320
321 #[inline]
332 pub fn mark_waiter_entered(&self) {
333 self.waiters.fetch_add(1, Ordering::AcqRel);
334 }
335
336 #[inline]
340 pub fn mark_waiter_left(&self) {
341 self.waiters.fetch_sub(1, Ordering::AcqRel);
342 }
343
344 pub fn flush(&self) -> Result<(), SemaphoreError> {
346 self.count.flush()?;
347 self.wakeup.flush()?;
348 self.waiters.flush()?;
349 Ok(())
350 }
351
352 pub fn flush_async(&self) -> Result<(), SemaphoreError> {
357 self.count.flush_async()?;
358 self.wakeup.flush_async()?;
359 self.waiters.flush_async()?;
360 Ok(())
361 }
362}
363
364pub struct Permit<'a> {
368 sem: &'a SharedSemaphore,
369}
370
371impl Drop for Permit<'_> {
372 fn drop(&mut self) {
373 self.sem.release().ok();
378 }
379}
380
381#[cfg(test)]
382mod tests {
383 use super::*;
384 use std::sync::atomic::{AtomicU32, Ordering as O};
385 use std::sync::Barrier;
386
387 fn tmp_base(name: &str) -> PathBuf {
388 let mut p = std::env::temp_dir();
389 let pid = std::process::id();
390 p.push(format!("subetha-semaphore-{name}-{pid}"));
391 p
392 }
393
394 fn cleanup(base: &Path) {
395 std::fs::remove_file(count_path(base)).ok();
396 std::fs::remove_file(wakeup_path(base)).ok();
397 std::fs::remove_file(waiters_path(base)).ok();
398 }
399
400 #[test]
401 fn create_initial_state_is_correct() {
402 let base = tmp_base("init");
403 let sem = SharedSemaphore::create(&base, 4, 4).unwrap();
404 assert_eq!(sem.available(), 4);
405 assert_eq!(sem.waiters(), 0);
406 assert_eq!(sem.max_permits(), 4);
407 cleanup(&base);
408 }
409
410 #[test]
411 fn try_acquire_succeeds_until_empty_then_returns_would_block() {
412 let base = tmp_base("try");
413 let sem = SharedSemaphore::create(&base, 3, 3).unwrap();
414 let _p1 = sem.try_acquire().unwrap();
415 let _p2 = sem.try_acquire().unwrap();
416 let _p3 = sem.try_acquire().unwrap();
417 assert_eq!(sem.try_acquire().err(), Some(SemaphoreError::WouldBlock));
418 cleanup(&base);
419 }
420
421 #[test]
422 fn permit_drop_releases() {
423 let base = tmp_base("drop");
424 let sem = SharedSemaphore::create(&base, 1, 1).unwrap();
425 {
426 let _p = sem.try_acquire().unwrap();
427 assert_eq!(sem.available(), 0);
428 }
429 assert_eq!(sem.available(), 1);
430 cleanup(&base);
431 }
432
433 #[test]
434 fn acquire_blocks_until_release() {
435 let base = tmp_base("block-release");
436 let sem = Arc::new(SharedSemaphore::create(&base, 1, 1).unwrap());
437 let p1 = sem.try_acquire().unwrap();
438 let sem2 = sem.clone();
440 let h = thread::spawn(move || {
441 let _p = sem2.acquire(); 42u32
443 });
444 thread::sleep(Duration::from_millis(20));
445 drop(p1);
447 let v = h.join().unwrap();
448 assert_eq!(v, 42);
449 cleanup(&base);
450 }
451
452 #[test]
453 fn acquire_timeout_returns_timeout_when_no_permit() {
454 let base = tmp_base("timeout");
455 let sem = SharedSemaphore::create(&base, 0, 1).unwrap();
456 let start = Instant::now();
457 let r = sem.acquire_timeout(Duration::from_millis(20));
458 let elapsed = start.elapsed();
459 assert_eq!(r.err(), Some(SemaphoreError::Timeout));
460 assert!(elapsed >= Duration::from_millis(20));
461 assert!(elapsed < Duration::from_millis(200), "timeout took too long: {elapsed:?}");
462 cleanup(&base);
463 }
464
465 #[test]
466 fn acquire_timeout_succeeds_when_released_before_deadline() {
467 let base = tmp_base("timeout-ok");
468 let sem = Arc::new(SharedSemaphore::create(&base, 0, 1).unwrap());
469 let sem2 = sem.clone();
470 let releaser = thread::spawn(move || {
471 thread::sleep(Duration::from_millis(20));
472 sem2.release().unwrap();
473 });
474 let p = sem.acquire_timeout(Duration::from_millis(500)).unwrap();
475 drop(p);
476 releaser.join().unwrap();
477 cleanup(&base);
478 }
479
480 #[test]
481 fn release_overflow_is_rejected_and_rolls_back() {
482 let base = tmp_base("overflow");
483 let sem = SharedSemaphore::create(&base, 1, 1).unwrap();
484 assert_eq!(sem.release().err(), Some(SemaphoreError::ReleaseOverflow));
486 assert_eq!(sem.available(), 1); cleanup(&base);
488 }
489
490 #[test]
491 fn cross_handle_acquire_release() {
492 let base = tmp_base("cross-handle");
493 let owner = SharedSemaphore::create(&base, 2, 2).unwrap();
494 let consumer = SharedSemaphore::open(&base, 2).unwrap();
495 let p = owner.try_acquire().unwrap();
496 assert_eq!(consumer.available(), 1);
498 let q = consumer.try_acquire().unwrap();
500 assert_eq!(owner.available(), 0);
501 assert_eq!(consumer.try_acquire().err(), Some(SemaphoreError::WouldBlock));
502 drop(p);
503 drop(q);
504 assert_eq!(owner.available(), 2);
505 cleanup(&base);
506 }
507
508 #[test]
509 fn contended_8_threads_bounded_to_2_permits() {
510 let base = tmp_base("contended");
511 let sem = Arc::new(SharedSemaphore::create(&base, 2, 2).unwrap());
512 let n_threads = 8;
513 let per_thread = 5;
514 let in_flight = Arc::new(AtomicU32::new(0));
515 let max_seen = Arc::new(AtomicU32::new(0));
516 let barrier = Arc::new(Barrier::new(n_threads));
517 let mut handles = vec![];
518 for _ in 0..n_threads {
519 let sem = sem.clone();
520 let in_flight = in_flight.clone();
521 let max_seen = max_seen.clone();
522 let barrier = barrier.clone();
523 handles.push(thread::spawn(move || {
524 barrier.wait();
525 for _ in 0..per_thread {
526 let _p = sem.acquire();
527 let cur = in_flight.fetch_add(1, O::AcqRel) + 1;
528 max_seen.fetch_max(cur, O::AcqRel);
529 thread::sleep(Duration::from_micros(100));
530 in_flight.fetch_sub(1, O::AcqRel);
531 }
532 }));
533 }
534 for h in handles { h.join().unwrap(); }
535 assert!(max_seen.load(O::Acquire) <= 2,
537 "saw {} concurrent holders, expected <= 2",
538 max_seen.load(O::Acquire));
539 assert_eq!(sem.available(), 2);
540 cleanup(&base);
541 }
542
543 #[test]
544 fn many_waiters_all_eventually_acquire() {
545 let base = tmp_base("many-waiters");
546 let sem = Arc::new(SharedSemaphore::create(&base, 0, 4).unwrap());
547 let n = 4;
548 let count = Arc::new(AtomicU32::new(0));
549 let mut handles = vec![];
550 for _ in 0..n {
551 let sem = sem.clone();
552 let count = count.clone();
553 handles.push(thread::spawn(move || {
554 let _p = sem.acquire();
555 count.fetch_add(1, O::AcqRel);
556 thread::sleep(Duration::from_micros(100));
557 }));
558 }
559 for _ in 0..n {
561 thread::sleep(Duration::from_millis(5));
562 sem.release().unwrap();
563 }
564 for h in handles { h.join().unwrap(); }
565 assert_eq!(count.load(O::Acquire), n);
566 cleanup(&base);
567 }
568
569 #[test]
570 fn standalone_release_works_with_forgotten_permit() {
571 let base = tmp_base("forget");
572 let sem = SharedSemaphore::create(&base, 1, 1).unwrap();
573 let p = sem.try_acquire().unwrap();
574 std::mem::forget(p); assert_eq!(sem.available(), 0);
576 sem.release().unwrap();
577 assert_eq!(sem.available(), 1);
578 cleanup(&base);
579 }
580
581 #[test]
582 fn waiters_counter_reflects_blocked_threads() {
583 let base = tmp_base("waiters");
584 let sem = Arc::new(SharedSemaphore::create(&base, 0, 4).unwrap());
585 let n = 3;
586 let mut handles = vec![];
587 for _ in 0..n {
588 let sem = sem.clone();
589 handles.push(thread::spawn(move || {
590 let _p = sem.acquire();
591 }));
592 }
593 let mut tries = 0;
595 while sem.waiters() < n as u32 && tries < 100 {
596 thread::sleep(Duration::from_millis(5));
597 tries += 1;
598 }
599 assert!(sem.waiters() >= n as u32 - 1, "expected ~{n} waiters, saw {}", sem.waiters());
601 for _ in 0..n { sem.release().unwrap(); }
603 for h in handles { h.join().unwrap(); }
604 assert_eq!(sem.waiters(), 0);
605 cleanup(&base);
606 }
607}