1#![cfg(feature = "runtime")]
16
17use core::future::Future;
18use core::ops::{Deref, DerefMut};
19use core::pin::Pin;
20use core::task::{Context, Poll, Waker};
21
22use alloc::boxed::Box;
23use alloc::collections::VecDeque;
24use alloc::sync::Arc;
25use alloc::vec::Vec;
26
27use spin::Mutex;
28
29use crate::error::G2gError;
30
31#[derive(Debug)]
32pub struct BufferPool<T> {
33 inner: Arc<PoolInner<T>>,
34}
35
36#[derive(Debug)]
37struct PoolInner<T> {
38 state: Mutex<PoolState<T>>,
39 capacity: usize,
40}
41
42#[derive(Debug)]
43struct PoolState<T> {
44 free: Vec<T>,
45 waiters: VecDeque<Waker>,
46}
47
48impl<T> BufferPool<T> {
49 pub fn from_buffers(buffers: Vec<T>) -> Self {
52 let capacity = buffers.len();
53 Self {
54 inner: Arc::new(PoolInner {
55 state: Mutex::new(PoolState {
56 free: buffers,
57 waiters: VecDeque::new(),
58 }),
59 capacity,
60 }),
61 }
62 }
63
64 pub fn capacity(&self) -> usize {
65 self.inner.capacity
66 }
67
68 pub fn available(&self) -> usize {
70 self.inner.state.lock().free.len()
71 }
72
73 pub fn outstanding(&self) -> usize {
75 self.inner.capacity - self.available()
76 }
77
78 pub fn try_acquire(&self) -> Option<PooledBuffer<T>> {
80 let value = self.inner.state.lock().free.pop()?;
81 Some(PooledBuffer {
82 value: Some(value),
83 pool: self.inner.clone(),
84 })
85 }
86
87 pub fn try_acquire_or_err(&self) -> Result<PooledBuffer<T>, G2gError> {
91 self.try_acquire().ok_or(G2gError::PoolExhausted)
92 }
93
94 pub fn acquire(&self) -> AcquireFuture<'_, T> {
96 AcquireFuture { pool: self }
97 }
98}
99
100#[allow(missing_debug_implementations)]
101pub struct AcquireFuture<'a, T> {
102 pool: &'a BufferPool<T>,
103}
104
105impl<'a, T> Future for AcquireFuture<'a, T> {
106 type Output = PooledBuffer<T>;
107
108 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
109 let inner = &self.pool.inner;
110 let mut state = inner.state.lock();
111 if let Some(v) = state.free.pop() {
112 return Poll::Ready(PooledBuffer {
113 value: Some(v),
114 pool: inner.clone(),
115 });
116 }
117 if !state.waiters.iter().any(|w| w.will_wake(cx.waker())) {
121 state.waiters.push_back(cx.waker().clone());
122 }
123 Poll::Pending
124 }
125}
126
127impl<T> Clone for BufferPool<T> {
128 fn clone(&self) -> Self {
129 Self {
130 inner: self.inner.clone(),
131 }
132 }
133}
134
135impl BufferPool<Box<[u8]>> {
136 pub fn new_byte_pool(count: usize, bytes: usize) -> Self {
138 let mut buffers: Vec<Box<[u8]>> = Vec::with_capacity(count);
139 for _ in 0..count {
140 buffers.push(alloc::vec![0u8; bytes].into_boxed_slice());
141 }
142 Self::from_buffers(buffers)
143 }
144}
145
146#[derive(Debug)]
147pub struct PooledBuffer<T> {
148 value: Option<T>,
149 pool: Arc<PoolInner<T>>,
150}
151
152impl<T> Deref for PooledBuffer<T> {
153 type Target = T;
154 fn deref(&self) -> &T {
155 self.value
156 .as_ref()
157 .expect("PooledBuffer accessed after drop")
158 }
159}
160
161impl<T> DerefMut for PooledBuffer<T> {
162 fn deref_mut(&mut self) -> &mut T {
163 self.value
164 .as_mut()
165 .expect("PooledBuffer accessed after drop")
166 }
167}
168
169impl<T: AsRef<[u8]>> AsRef<[u8]> for PooledBuffer<T> {
170 fn as_ref(&self) -> &[u8] {
171 self.deref().as_ref()
172 }
173}
174
175impl<T: AsMut<[u8]>> AsMut<[u8]> for PooledBuffer<T> {
176 fn as_mut(&mut self) -> &mut [u8] {
177 self.deref_mut().as_mut()
178 }
179}
180
181impl<T> Drop for PooledBuffer<T> {
182 fn drop(&mut self) {
183 if let Some(v) = self.value.take() {
184 let waiters = {
190 let mut state = self.pool.state.lock();
191 state.free.push(v);
192 core::mem::take(&mut state.waiters)
193 };
194 for w in waiters {
195 w.wake();
196 }
197 }
198 }
199}
200
201#[cfg(test)]
202mod tests {
203 use super::*;
204
205 #[test]
206 fn pool_capacity_and_available_match_on_construction() {
207 let pool = BufferPool::new_byte_pool(4, 16);
208 assert_eq!(pool.capacity(), 4);
209 assert_eq!(pool.available(), 4);
210 assert_eq!(pool.outstanding(), 0);
211 }
212
213 #[test]
214 fn acquire_decrements_available_drop_returns_buffer() {
215 let pool = BufferPool::new_byte_pool(4, 16);
216 {
217 let _a = pool.try_acquire_or_err().expect("a");
218 let _b = pool.try_acquire_or_err().expect("b");
219 assert_eq!(pool.available(), 2);
220 assert_eq!(pool.outstanding(), 2);
221 }
222 assert_eq!(pool.available(), 4);
223 assert_eq!(pool.outstanding(), 0);
224 }
225
226 #[test]
227 fn exhausted_pool_returns_pool_exhausted() {
228 let pool = BufferPool::new_byte_pool(2, 8);
229 let _a = pool.try_acquire_or_err().unwrap();
230 let _b = pool.try_acquire_or_err().unwrap();
231 assert!(matches!(
232 pool.try_acquire_or_err(),
233 Err(G2gError::PoolExhausted)
234 ));
235 assert!(pool.try_acquire().is_none());
236 }
237
238 #[test]
239 fn clones_share_the_same_free_list() {
240 let pool = BufferPool::new_byte_pool(2, 8);
241 let pool2 = pool.clone();
242 let _a = pool.try_acquire_or_err().unwrap();
243 assert_eq!(pool2.available(), 1);
244 }
245
246 #[test]
247 fn buffer_size_matches_byte_pool_argument() {
248 let pool = BufferPool::new_byte_pool(1, 32);
249 let buf = pool.try_acquire_or_err().unwrap();
250 assert_eq!(buf.as_ref().len(), 32);
251 }
252
253 struct CountWaker(core::sync::atomic::AtomicUsize);
256 impl alloc::task::Wake for CountWaker {
257 fn wake(self: Arc<Self>) {
258 self.0.fetch_add(1, core::sync::atomic::Ordering::SeqCst);
259 }
260 fn wake_by_ref(self: &Arc<Self>) {
261 self.0.fetch_add(1, core::sync::atomic::Ordering::SeqCst);
262 }
263 }
264
265 #[test]
266 fn dropped_pending_acquirer_does_not_starve_a_live_waiter() {
267 use core::task::{Context, Poll};
268
269 let pool = BufferPool::new_byte_pool(1, 8);
270 let held = pool.try_acquire_or_err().unwrap(); let live = Arc::new(CountWaker(core::sync::atomic::AtomicUsize::new(0)));
273 let live_waker = live.clone().into();
274 let mut live_cx = Context::from_waker(&live_waker);
275
276 let mut live_fut = pool.acquire();
279 assert!(matches!(
280 core::pin::Pin::new(&mut live_fut).poll(&mut live_cx),
281 Poll::Pending
282 ));
283 {
284 let cancelled = Arc::new(CountWaker(core::sync::atomic::AtomicUsize::new(0)));
285 let cancelled_waker = cancelled.clone().into();
286 let mut cancelled_cx = Context::from_waker(&cancelled_waker);
287 let mut cancelled_fut = pool.acquire();
288 assert!(matches!(
289 core::pin::Pin::new(&mut cancelled_fut).poll(&mut cancelled_cx),
290 Poll::Pending
291 ));
292 }
294
295 drop(held);
298 assert!(
299 live.0.load(core::sync::atomic::Ordering::SeqCst) >= 1,
300 "the live acquirer must be woken when the buffer returns"
301 );
302 assert!(matches!(
303 core::pin::Pin::new(&mut live_fut).poll(&mut live_cx),
304 Poll::Ready(_)
305 ));
306 }
307
308 #[test]
309 fn repolling_a_pending_acquirer_does_not_accumulate_wakers() {
310 use core::task::{Context, Poll};
311
312 let pool = BufferPool::new_byte_pool(1, 8);
313 let _held = pool.try_acquire_or_err().unwrap();
314
315 let w = Arc::new(CountWaker(core::sync::atomic::AtomicUsize::new(0)));
316 let waker = w.clone().into();
317 let mut cx = Context::from_waker(&waker);
318
319 let mut fut = pool.acquire();
320 for _ in 0..5 {
321 assert!(matches!(
322 core::pin::Pin::new(&mut fut).poll(&mut cx),
323 Poll::Pending
324 ));
325 }
326 assert_eq!(
327 pool.inner.state.lock().waiters.len(),
328 1,
329 "re-polling the same future must not push duplicate wakers"
330 );
331 }
332}