strat9_kernel/ipc/
semaphore.rs1use crate::sync::{SpinLock, WaitQueue};
2use alloc::{collections::BTreeMap, sync::Arc};
3use core::sync::atomic::{AtomicBool, AtomicI32, AtomicU64, Ordering};
4
5#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
6pub struct SemId(pub u64);
7
8impl SemId {
9 pub fn as_u64(self) -> u64 {
11 self.0
12 }
13 pub fn from_u64(raw: u64) -> Self {
15 Self(raw)
16 }
17}
18
19#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
20pub enum SemaphoreError {
21 #[error("would block")]
22 WouldBlock,
23 #[error("semaphore destroyed")]
24 Destroyed,
25 #[error("invalid initial value")]
26 InvalidValue,
27 #[error("semaphore not found")]
28 NotFound,
29 #[error("interrupted by signal")]
30 Interrupted,
31}
32
33pub struct PosixSemaphore {
34 count: AtomicI32,
35 destroyed: AtomicBool,
36 waitq: WaitQueue,
37}
38
39impl PosixSemaphore {
40 fn new(initial: u32) -> Self {
42 Self {
43 count: AtomicI32::new(initial as i32),
44 destroyed: AtomicBool::new(false),
45 waitq: WaitQueue::new(),
46 }
47 }
48
49 pub fn wait(&self) -> Result<(), SemaphoreError> {
51 self.waitq.wait_until(|| {
52 if crate::process::signal::has_pending_signals() {
54 return Some(Err(SemaphoreError::Interrupted));
55 }
56 if self.destroyed.load(Ordering::Acquire) {
57 return Some(Err(SemaphoreError::Destroyed));
58 }
59 let cur = self.count.load(Ordering::Acquire);
60 if cur <= 0 {
61 return None;
62 }
63 match self.count.compare_exchange_weak(
64 cur,
65 cur - 1,
66 Ordering::AcqRel,
67 Ordering::Acquire,
68 ) {
69 Ok(_) => Some(Ok(())),
70 Err(_) => None,
71 }
72 })
73 }
74
75 pub fn try_wait(&self) -> Result<(), SemaphoreError> {
77 if self.destroyed.load(Ordering::Acquire) {
78 return Err(SemaphoreError::Destroyed);
79 }
80 loop {
81 let cur = self.count.load(Ordering::Acquire);
82 if cur <= 0 {
83 return Err(SemaphoreError::WouldBlock);
84 }
85 if self
86 .count
87 .compare_exchange_weak(cur, cur - 1, Ordering::AcqRel, Ordering::Acquire)
88 .is_ok()
89 {
90 return Ok(());
91 }
92 }
93 }
94
95 pub fn post(&self) -> Result<(), SemaphoreError> {
97 if self.destroyed.load(Ordering::Acquire) {
98 return Err(SemaphoreError::Destroyed);
99 }
100 loop {
101 let cur = self.count.load(Ordering::Acquire);
102 if cur == i32::MAX {
104 return Err(SemaphoreError::InvalidValue);
105 }
106 if self
107 .count
108 .compare_exchange_weak(cur, cur + 1, Ordering::AcqRel, Ordering::Acquire)
109 .is_ok()
110 {
111 break;
112 }
113 }
114 self.waitq.wake_one();
115 Ok(())
116 }
117
118 pub fn destroy(&self) {
120 self.destroyed.store(true, Ordering::Release);
121 self.waitq.wake_all();
122 }
123
124 pub fn count(&self) -> i32 {
126 self.count.load(Ordering::Acquire)
127 }
128
129 pub fn is_destroyed(&self) -> bool {
131 self.destroyed.load(Ordering::Acquire)
132 }
133}
134
135static NEXT_SEM_ID: AtomicU64 = AtomicU64::new(1);
136static SEMAPHORES: SpinLock<Option<BTreeMap<SemId, Arc<PosixSemaphore>>>> = SpinLock::new(None);
137
138fn ensure_registry(guard: &mut Option<BTreeMap<SemId, Arc<PosixSemaphore>>>) {
140 if guard.is_none() {
141 *guard = Some(BTreeMap::new());
142 }
143}
144
145pub fn create_semaphore(initial: u32) -> Result<SemId, SemaphoreError> {
147 if initial > i32::MAX as u32 {
148 return Err(SemaphoreError::InvalidValue);
149 }
150 let id = SemId(NEXT_SEM_ID.fetch_add(1, Ordering::Relaxed));
151 let sem = Arc::new(PosixSemaphore::new(initial));
152 let mut reg = SEMAPHORES.lock();
153 ensure_registry(&mut *reg);
154 reg.as_mut().unwrap().insert(id, sem);
155 Ok(id)
156}
157
158pub fn get_semaphore(id: SemId) -> Option<Arc<PosixSemaphore>> {
160 let reg = SEMAPHORES.lock();
161 reg.as_ref().and_then(|m| m.get(&id).cloned())
162}
163
164pub fn destroy_semaphore(id: SemId) -> Result<(), SemaphoreError> {
166 let sem = {
167 let mut reg = SEMAPHORES.lock();
168 let map = reg.as_mut().ok_or(SemaphoreError::NotFound)?;
169 map.remove(&id).ok_or(SemaphoreError::NotFound)?
170 };
171 sem.destroy();
172 Ok(())
173}