Skip to main content

strat9_kernel/ipc/
semaphore.rs

1use 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    /// Returns this as u64.
10    pub fn as_u64(self) -> u64 {
11        self.0
12    }
13    /// Builds this from u64.
14    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    /// Creates a new instance.
41    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    /// Performs the wait operation.
50    pub fn wait(&self) -> Result<(), SemaphoreError> {
51        self.waitq.wait_until(|| {
52            // P2 fix: check for pending signals to avoid livelock.
53            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    /// Attempts to wait.
76    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    /// Performs the post operation.
96    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            // `>=` on i32 reduces to `==`; the explicit form avoids overflow.
103            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    /// Performs the destroy operation.
119    pub fn destroy(&self) {
120        self.destroyed.store(true, Ordering::Release);
121        self.waitq.wake_all();
122    }
123
124    /// Performs the count operation.
125    pub fn count(&self) -> i32 {
126        self.count.load(Ordering::Acquire)
127    }
128
129    /// Returns whether destroyed.
130    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
138/// Performs the ensure registry operation.
139fn ensure_registry(guard: &mut Option<BTreeMap<SemId, Arc<PosixSemaphore>>>) {
140    if guard.is_none() {
141        *guard = Some(BTreeMap::new());
142    }
143}
144
145/// Creates semaphore.
146pub 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
158/// Returns semaphore.
159pub 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
164/// Destroys semaphore.
165pub 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}