1#include <sch/sched.h>
2#include <sync/condvar.h>
3#include <sync/semaphore.h>
4#include <sync/spinlock.h>
5#include <thread/thread_types.h>
6
7#define get_count(sem) atomic_load(&sem->count)
8#define set_count(sem, val) atomic_store(&sem->count, val)
9#define inc_count(sem) atomic_fetch_add(&sem->count, 1)
10#define dec_count(sem) atomic_fetch_sub(&sem->count, 1)
11#define add_count(sem, val) atomic_fetch_add(&sem->count, val)
12
13void semaphore_init(struct semaphore *s, int value, bool irq_disable) {
14 s->count = value;
15 s->irq_disable = irq_disable;
16 spinlock_init(lock: &s->lock);
17 condvar_init(cv: &s->cv, irq_disable);
18}
19
20static enum irql semaphore_lock_internal(struct semaphore *sem) {
21 if (sem->irq_disable)
22 return spin_lock_irq_disable(lock: &sem->lock);
23
24 return spin_lock(lock: &sem->lock);
25}
26
27void semaphore_wait(struct semaphore *s) {
28 enum irql irql = semaphore_lock_internal(sem: s);
29
30 while (get_count(s) == 0)
31 condvar_wait(cv: &s->cv, lock: &s->lock, irql, out: &irql);
32
33 dec_count(s);
34 spin_unlock(lock: &s->lock, old: irql);
35}
36
37bool semaphore_timedwait(struct semaphore *s, time_t timeout_ms) {
38 enum irql irql = semaphore_lock_internal(sem: s);
39
40 while (get_count(s) == 0) {
41 enum irql out;
42 if (!condvar_wait_timeout(cv: &s->cv, lock: &s->lock, timeout_ms, irql, out: &out)) {
43 spin_unlock(lock: &s->lock, old: irql);
44 return false;
45 }
46 }
47
48 dec_count(s);
49 spin_unlock(lock: &s->lock, old: irql);
50
51 return true;
52}
53
54void semaphore_post(struct semaphore *s) {
55 enum irql irql = semaphore_lock_internal(sem: s);
56
57 inc_count(s);
58
59 condvar_signal(cv: &s->cv);
60
61 spin_unlock(lock: &s->lock, old: irql);
62}
63
64void semaphore_postn(struct semaphore *s, int n) {
65 enum irql irql = semaphore_lock_internal(sem: s);
66
67 add_count(s, n);
68 for (int i = 0; i < n; i++)
69 condvar_signal(cv: &s->cv);
70
71 spin_unlock(lock: &s->lock, old: irql);
72}
73
74void semaphore_post_callback(struct semaphore *s, thread_action_callback cb) {
75 enum irql irql = semaphore_lock_internal(sem: s);
76
77 inc_count(s);
78
79 condvar_signal_callback(cv: &s->cv, cb);
80
81 spin_unlock(lock: &s->lock, old: irql);
82}
83
84void semaphore_postn_callback(struct semaphore *s, int n,
85 thread_action_callback cb) {
86 enum irql irql = semaphore_lock_internal(sem: s);
87
88 add_count(s, n);
89 for (int i = 0; i < n; i++)
90 condvar_signal_callback(cv: &s->cv, cb);
91
92 spin_unlock(lock: &s->lock, old: irql);
93}
94