The Pedigree Project 0.1
CallbackDelivery.h
1/*
2 * Copyright (c) 2026, Pedigree Developers
3 *
4 * Permission to use, copy, modify, and distribute this software for any
5 * purpose with or without fee is hereby granted.
6 */
7
8#ifndef USB_HCD_CALLBACK_DELIVERY_H
9#define USB_HCD_CALLBACK_DELIVERY_H
10
11#include "pedigree/kernel/Atomic.h"
12#include "pedigree/kernel/LockGuard.h"
13#include "pedigree/kernel/compiler.h"
14#include "pedigree/kernel/process/Mutex.h"
15#include "pedigree/kernel/process/TerminationDeferral.h"
16#include "pedigree/kernel/process/WaitQueue.h"
17#include "pedigree/kernel/processor/Processor.h"
18#include "pedigree/kernel/processor/ProcessorInformation.h"
19#include "pedigree/kernel/utilities/List.h"
20
21namespace UsbHcd {
30 public:
31 struct Key {
32 Key(uintptr_t transaction = 0, size_t generation = 0, size_t subscription = 0)
33 : transaction(transaction), generation(generation), subscription(subscription) {}
34
35 uintptr_t transaction;
36 size_t generation;
37 size_t subscription;
38
39 bool operator==(const Key& other) const {
40 return transaction == other.transaction && generation == other.generation &&
41 subscription == other.subscription;
42 }
43 };
44
45 using Callback = void (*)(uintptr_t, ssize_t);
46 using AfterDelivery = void (*)(void*);
47 using OnDestroy = void (*)(void*);
48
49 class Record {
50 private:
51 friend class CallbackDeliveryQueue;
52
53 enum class State {
54 Pending,
55 Running,
56 Complete,
57 };
58
59 Record(const Key& key, Callback callback, uintptr_t parameter, ssize_t result,
60 AfterDelivery afterDelivery, void* afterDeliveryContext, OnDestroy onDestroy,
61 void* destroyContext)
62 : m_Key(key),
63 m_Callback(callback),
64 m_Parameter(parameter),
65 m_Result(result),
66 m_AfterDelivery(afterDelivery),
67 m_AfterDeliveryContext(afterDeliveryContext),
68 m_OnDestroy(onDestroy),
69 m_DestroyContext(destroyContext),
70 m_References(1),
71 m_State(State::Pending),
72 m_Runner(nullptr),
73 m_CompletionWaiters(),
74 m_Completed(false) {}
75
76 ~Record() {
77 if (m_OnDestroy)
78 m_OnDestroy(m_DestroyContext);
79 }
80
81 void retain() {
82 m_References += 1;
83 }
84
85 void release() {
86 if ((m_References -= 1) == 0)
87 delete this;
88 }
89
90 void waitForCompletion() {
91 while (true) {
92 auto guard = m_CompletionWaiters.acquire();
93 if (m_Completed)
94 return;
95 const WaitQueue::WakeReason wakeReason = guard.waitForCompletion();
96 (void)wakeReason;
97 }
98 }
99
100 void complete() {
101 auto guard = m_CompletionWaiters.acquire();
102 m_Completed = true;
103 guard.wakeAll();
104 }
105
106 Key m_Key;
107 Callback m_Callback;
108 uintptr_t m_Parameter;
109 ssize_t m_Result;
110 AfterDelivery m_AfterDelivery;
111 void* m_AfterDeliveryContext;
112 OnDestroy m_OnDestroy;
113 void* m_DestroyContext;
114 Atomic<size_t> m_References;
115
117 State m_State;
118 void* m_Runner;
119
120 WaitQueue m_CompletionWaiters;
123 };
124
125 CallbackDeliveryQueue() : m_Lock(), m_Records(), m_NextGeneration(0) {}
126
128 assert(empty());
129 }
130
131 MUST_USE_RESULT Record* create(const Key& key, Callback callback, uintptr_t parameter,
132 ssize_t result, AfterDelivery afterDelivery = nullptr,
133 void* afterDeliveryContext = nullptr,
134 OnDestroy onDestroy = nullptr, void* destroyContext = nullptr) {
135 return new Record(key, callback, parameter, result, afterDelivery, afterDeliveryContext,
136 onDestroy, destroyContext);
137 }
138
139 size_t nextGeneration() {
140 size_t generation = m_NextGeneration += 1;
141 if (!generation)
142 generation = m_NextGeneration += 1;
143 return generation;
144 }
145
147 void publish(List<Record*>& records) {
148 LockGuard<Mutex> guard(m_Lock);
149 for (List<Record*>::Iterator it = records.begin(); it != records.end(); ++it) {
150 assert(findLocked((*it)->m_Key) == nullptr);
151 m_Records.pushBack(*it);
152 }
153 }
154
156 void deliver(Record* record) {
157 bool run = false;
158 bool wait = false;
159 {
160 LockGuard<Mutex> guard(m_Lock);
161 if (record->m_State == Record::State::Pending) {
162 record->m_State = Record::State::Running;
163 record->m_Runner = currentRunner();
164 run = true;
165 } else if (record->m_State == Record::State::Running) {
166 wait = record->m_Runner != currentRunner();
167 }
168 }
169
170 if (run)
171 runRecord(record);
172 else if (wait)
173 record->waitForCompletion();
174 record->release();
175 }
176
184 bool drain(const Key& key) {
185 Record* record = nullptr;
186 bool run = false;
187 bool self = false;
188 {
189 LockGuard<Mutex> guard(m_Lock);
190 record = findLocked(key);
191 if (!record)
192 return false;
193
194 record->retain();
195 if (record->m_State == Record::State::Pending) {
196 record->m_State = Record::State::Running;
197 record->m_Runner = currentRunner();
198 run = true;
199 } else if (record->m_State == Record::State::Running) {
200 self = record->m_Runner == currentRunner();
201 }
202 }
203
204 if (run)
205 runRecord(record);
206 else if (!self)
207 record->waitForCompletion();
208
209 record->release();
210 return true;
211 }
212
221 bool cancelSubscription(uintptr_t transaction, size_t subscription) {
222 TerminationDeferral cancellationLifetime;
223 bool runningTarget = false;
224 const bool callerIsCallback = inCallbackContext();
225 while (true) {
226 Record* record = nullptr;
227 bool suppressed = false;
228 bool wait = false;
229 {
230 LockGuard<Mutex> guard(m_Lock);
231 for (List<Record*>::Iterator it = m_Records.begin(); it != m_Records.end(); ++it) {
232 Record* candidate = *it;
233 if (candidate->m_Key.transaction != transaction ||
234 candidate->m_Key.subscription != subscription) {
235 continue;
236 }
237
238 if (candidate->m_State == Record::State::Running &&
239 (callerIsCallback || candidate->m_Runner == currentRunner())) {
240 runningTarget = true;
241 continue;
242 }
243
244 record = candidate;
245 record->retain();
246 if (record->m_State == Record::State::Pending) {
247 m_Records.erase(it);
248 record->m_State = Record::State::Complete;
249 record->m_Runner = nullptr;
250 suppressed = true;
251 } else if (record->m_State == Record::State::Running) {
252 wait = true;
253 }
254 break;
255 }
256 }
257
258 if (!record)
259 return !runningTarget;
260
261 if (suppressed)
262 record->complete();
263 else if (wait)
264 record->waitForCompletion();
265 record->release();
266 }
267 }
268
276 size_t drainAll() {
277 size_t drained = 0;
278 while (true) {
279 Key key = {0, 0};
280 bool found = false;
281 {
282 LockGuard<Mutex> guard(m_Lock);
283 for (List<Record*>::Iterator it = m_Records.begin(); it != m_Records.end(); ++it) {
284 Record* record = *it;
285 if (record->m_State == Record::State::Running && record->m_Runner == currentRunner()) {
286 continue;
287 }
288
289 key = record->m_Key;
290 found = true;
291 break;
292 }
293 }
294
295 if (!found)
296 return drained;
297 if (drain(key))
298 ++drained;
299 }
300 }
301
302 bool contains(const Key& key) {
303 LockGuard<Mutex> guard(m_Lock);
304 return findLocked(key) != nullptr;
305 }
306
307 size_t activeCount() {
308 LockGuard<Mutex> guard(m_Lock);
309 return m_Records.count();
310 }
311
312 bool empty() {
313 return activeCount() == 0;
314 }
315
317 static bool isInCallbackContext() {
318 return inCallbackContext();
319 }
320
321 private:
323 void* runner;
324 ActiveCallback* next;
325 };
326
328 public:
329 CallbackContext() : m_Active{currentRunner(), nullptr} {
330 LockGuard<Mutex> guard(callbackContextLock());
331 m_Active.next = callbackContexts();
332 callbackContexts() = &m_Active;
333 }
334
336 LockGuard<Mutex> guard(callbackContextLock());
337 ActiveCallback** link = &callbackContexts();
338 while (*link && *link != &m_Active)
339 link = &((*link)->next);
340 assert(*link == &m_Active);
341 if (*link)
342 *link = m_Active.next;
343 }
344
345 private:
346 ActiveCallback m_Active;
347 };
348
349 static Mutex& callbackContextLock() {
350 static Mutex lock;
351 return lock;
352 }
353
354 static ActiveCallback*& callbackContexts() {
355 static ActiveCallback* contexts = nullptr;
356 return contexts;
357 }
358
359 static bool inCallbackContext() {
360 const void* runner = currentRunner();
361 LockGuard<Mutex> guard(callbackContextLock());
362 for (ActiveCallback* active = callbackContexts(); active; active = active->next) {
363 if (active->runner == runner)
364 return true;
365 }
366 return false;
367 }
368
369 static void* currentRunner() {
370 ProcessorInformation& information = Processor::information();
371 auto* thread = information.getCurrentThread();
372 return thread ? static_cast<void*>(thread) : static_cast<void*>(&information);
373 }
374
375 Record* findLocked(const Key& key) {
376 for (List<Record*>::Iterator it = m_Records.begin(); it != m_Records.end(); ++it) {
377 if ((*it)->m_Key == key)
378 return *it;
379 }
380 return nullptr;
381 }
382
383 void finishRecord(Record* record) {
384 {
385 LockGuard<Mutex> guard(m_Lock);
386 bool removed = false;
387 for (List<Record*>::Iterator it = m_Records.begin(); it != m_Records.end(); ++it) {
388 if (*it == record) {
389 m_Records.erase(it);
390 removed = true;
391 break;
392 }
393 }
394 assert(removed);
395 record->m_State = Record::State::Complete;
396 record->m_Runner = nullptr;
397 }
398 record->complete();
399 }
400
401 void runRecord(Record* record) {
402 TerminationDeferral deliveryLifetime;
403 CallbackContext callbackContext;
404 assert(Processor::getInterrupts());
405 if (record->m_Callback)
406 record->m_Callback(record->m_Parameter, record->m_Result);
407 if (record->m_AfterDelivery)
408 record->m_AfterDelivery(record->m_AfterDeliveryContext);
409 // CallbackContext unregisters its stack record before runRecord returns.
410 finishRecord(record); // NOLINT(clang-analyzer-core.StackAddressEscape)
411 }
412
413 Mutex m_Lock;
414 List<Record*> m_Records;
415 Atomic<size_t> m_NextGeneration;
416
417 NOT_COPYABLE_OR_ASSIGNABLE(CallbackDeliveryQueue);
418};
419} // namespace UsbHcd
420
421#endif
Definition List.h:61
Iterator begin()
Definition List.h:122
::Iterator< T, node_t > Iterator
Definition List.h:67
Iterator end()
Definition List.h:132
Definition Mutex.h:56
static bool getInterrupts()
static ProcessorInformation & information()
bool cancelSubscription(uintptr_t transaction, size_t subscription)
void publish(List< Record * > &records)
MUST_USE_RESULT WakeReason waitForCompletion(const Channel &channel=Channel(), size_t debugState=0, uintptr_t debugAddress=0)
Definition WaitQueue.cc:116