1mod configure;
2
3use std::{
4 mem::take,
5 sync::{
6 Arc, Mutex,
7 atomic::{AtomicUsize, Ordering},
8 },
9 thread,
10 thread::JoinHandle,
11};
12
13use async_channel::{QueueStrategy, Receiver, RecvError, Sender};
14use futures::{TryFutureExt, channel::oneshot};
15use oneshot::Sender as ResultSender;
16use rocksdb::Direction;
17use tuwunel_core::{
18 Error, Result, Server, debug, err, error, implement,
19 result::DebugInspect,
20 smallvec::SmallVec,
21 trace,
22 utils::sys::compute::{get_affinity, set_affinity},
23};
24
25use self::configure::configure;
26use crate::{Handle, Map, keyval::KeyBuf, stream};
27
28pub(crate) struct Pool {
34 server: Arc<Server>,
35 queues: Vec<Sender<Cmd>>,
36 workers: Mutex<Vec<JoinHandle<()>>>,
37 topology: Vec<usize>,
38 busy: AtomicUsize,
39 queued_max: AtomicUsize,
40}
41
42pub(crate) enum Cmd {
47 Get(Get),
48 Iter(Seek),
49}
50
51pub(crate) struct Get {
56 pub(crate) map: Arc<Map>,
57 pub(crate) key: BatchQuery<'static>,
58 pub(crate) res: Option<ResultSender<BatchResult<'static>>>,
59}
60
61pub(crate) struct Seek {
67 pub(crate) map: Arc<Map>,
68 pub(crate) state: stream::State<'static>,
69 pub(crate) dir: Direction,
70 pub(crate) key: Option<KeyBuf>,
71 pub(crate) res: Option<ResultSender<stream::State<'static>>>,
72}
73
74pub(crate) type BatchQuery<'a> = SmallVec<[KeyBuf; BATCH_INLINE]>;
79
80pub(crate) type BatchResult<'a> = SmallVec<[ResultHandle<'a>; BATCH_INLINE]>;
85
86pub(crate) type ResultHandle<'a> = Result<Handle<'a>>;
91
92const WORKER_LIMIT: (usize, usize) = (1, 4096);
93const QUEUE_LIMIT: (usize, usize) = (1, 1024);
94const BATCH_INLINE: usize = 1;
95
96const WORKER_STACK_SIZE: usize = 1_048_576;
97const WORKER_NAME: &str = "tuwunel:db";
98
99#[implement(Pool)]
105pub(crate) fn new(server: &Arc<Server>) -> Result<Arc<Self>> {
106 const CHAN_SCHED: (QueueStrategy, QueueStrategy) = (QueueStrategy::Fifo, QueueStrategy::Lifo);
107
108 let (topology, workers, queues) = configure(server);
109
110 let (senders, receivers): (Vec<_>, Vec<_>) = queues
111 .into_iter()
112 .map(|cap| cap.max(QUEUE_LIMIT.0))
113 .map(|cap| async_channel::bounded_with_queue_strategy(cap, CHAN_SCHED))
114 .unzip();
115
116 let pool = Arc::new(Self {
117 server: server.clone(),
118 queues: senders,
119 workers: Vec::new().into(),
120 topology,
121 busy: AtomicUsize::default(),
122 queued_max: AtomicUsize::default(),
123 });
124
125 for (chan_id, &count) in workers.iter().enumerate() {
126 pool.spawn_group(&receivers, chan_id, count)
127 .inspect_err(|_| pool.close())?;
128 }
129
130 Ok(pool)
131}
132
133impl Drop for Pool {
134 fn drop(&mut self) {
135 self.close();
136
137 debug_assert!(
138 self.queues.iter().all(Sender::is_empty),
139 "channel must should not have requests queued on drop"
140 );
141 debug_assert!(
142 self.queues.iter().all(Sender::is_closed),
143 "channel should be closed on drop"
144 );
145 }
146}
147
148#[implement(Pool)]
149#[tracing::instrument(skip_all)]
150pub(crate) fn close(&self) {
151 let workers = take(&mut *self.workers.lock().expect("locked"));
152
153 let senders = self
154 .queues
155 .iter()
156 .map(Sender::sender_count)
157 .sum::<usize>();
158
159 let receivers = self
160 .queues
161 .iter()
162 .map(Sender::receiver_count)
163 .sum::<usize>();
164
165 for queue in &self.queues {
166 queue.close();
167 }
168
169 if workers.is_empty() {
170 return;
171 }
172
173 debug!(
174 senders,
175 receivers,
176 queues = self.queues.len(),
177 workers = workers.len(),
178 "Closing pool. Waiting for workers to join..."
179 );
180
181 workers
182 .into_iter()
183 .map(JoinHandle::join)
184 .map(|result| result.map_err(Error::from_panic))
185 .enumerate()
186 .for_each(|(id, result)| match result {
187 | Ok(()) => trace!(?id, "worker joined"),
188 | Err(error) => error!(?id, "worker joined with error: {error}"),
189 });
190}
191
192#[implement(Pool)]
193fn spawn_group(self: &Arc<Self>, recv: &[Receiver<Cmd>], chan_id: usize, count: usize) -> Result {
194 let mut workers = self.workers.lock().expect("locked");
195 for _ in 0..count {
196 self.clone()
197 .spawn_one(&mut workers, recv, chan_id)?;
198 }
199
200 Ok(())
201}
202
203#[implement(Pool)]
204#[tracing::instrument(
205 name = "spawn",
206 level = "trace",
207 skip_all,
208 fields(id = %workers.len())
209)]
210fn spawn_one(
211 self: Arc<Self>,
212 workers: &mut Vec<JoinHandle<()>>,
213 recv: &[Receiver<Cmd>],
214 chan_id: usize,
215) -> Result {
216 debug_assert!(!self.queues.is_empty(), "Must have at least one queue");
217 debug_assert!(!recv.is_empty(), "Must have at least one receiver");
218
219 let id = workers.len();
220 let recv = recv[chan_id].clone();
221
222 let handle = thread::Builder::new()
223 .name(WORKER_NAME.into())
224 .stack_size(WORKER_STACK_SIZE)
225 .spawn(move || self.worker(id, chan_id, &recv))?;
226
227 workers.push(handle);
228
229 Ok(())
230}
231
232#[implement(Pool)]
233#[tracing::instrument(level = "trace", name = "get", skip(self, cmd))]
234pub(crate) async fn execute_get(self: &Arc<Self>, mut cmd: Get) -> Result<BatchResult<'_>> {
235 let (send, recv) = oneshot::channel();
236 _ = cmd.res.insert(send);
237
238 let queue = self.select_queue();
239 self.execute(queue, Cmd::Get(cmd))
240 .and_then(move |()| {
241 recv.map_ok(into_recv_get)
242 .map_err(|e| err!(error!("recv failed {e:?}")))
243 })
244 .await
245}
246
247#[implement(Pool)]
248#[tracing::instrument(level = "trace", name = "iter", skip(self, cmd))]
249pub(crate) async fn execute_iter(self: &Arc<Self>, mut cmd: Seek) -> Result<stream::State<'_>> {
250 let (send, recv) = oneshot::channel();
251 _ = cmd.res.insert(send);
252
253 let queue = self.select_queue();
254 self.execute(queue, Cmd::Iter(cmd))
255 .and_then(|()| {
256 recv.map_ok(into_recv_seek)
257 .map_err(|e| err!(error!("recv failed {e:?}")))
258 })
259 .await
260}
261
262#[implement(Pool)]
271fn select_queue(&self) -> &Sender<Cmd> {
272 let core_id = get_affinity()
273 .next()
274 .expect("Affinity mask should be available.");
275
276 let chan_id = self.topology[core_id];
277
278 self.queues
279 .get(chan_id)
280 .unwrap_or_else(|| &self.queues[0])
281}
282
283#[implement(Pool)]
284#[tracing::instrument(
285 level = "trace",
286 name = "execute",
287 skip(self, cmd),
288 fields(
289 task = ?tokio::task::try_id(),
290 receivers = queue.receiver_count(),
291 queued = queue.len(),
292 queued_max = self.queued_max.load(Ordering::Relaxed),
293 ),
294)]
295async fn execute(&self, queue: &Sender<Cmd>, cmd: Cmd) -> Result {
296 if cfg!(debug_assertions) {
297 self.queued_max
298 .fetch_max(queue.len(), Ordering::Relaxed);
299 }
300
301 queue
302 .send(cmd)
303 .await
304 .map_err(|e| err!(error!("send failed {e:?}")))
305}
306
307#[implement(Pool)]
308#[tracing::instrument(
309 parent = None,
310 level = "debug",
311 skip_all,
312 fields(
313 id,
314 chan_id,
315 thread_id = ?thread::current().id(),
316 ),
317)]
318fn worker(self: Arc<Self>, id: usize, chan_id: usize, recv: &Receiver<Cmd>) {
319 self.worker_init(id, chan_id);
320 self.worker_loop(recv);
321}
322
323#[implement(Pool)]
324fn worker_init(&self, id: usize, chan_id: usize) {
325 let affinity = self
326 .topology
327 .iter()
328 .enumerate()
329 .filter(|_| self.server.config.db_pool_affinity)
330 .filter_map(|(core_id, &queue_id)| (chan_id == queue_id).then_some(core_id));
331
332 set_affinity(affinity.clone());
334
335 trace!(
336 ?id,
337 ?chan_id,
338 affinity = ?affinity.collect::<Vec<_>>(),
339 "worker ready"
340 );
341}
342
343#[implement(Pool)]
344fn worker_loop(self: &Arc<Self>, recv: &Receiver<Cmd>) {
345 self.busy.fetch_add(1, Ordering::Relaxed);
347
348 while let Ok(cmd) = self.worker_wait(recv) {
349 worker_handle(cmd);
350 }
351}
352
353#[implement(Pool)]
354#[tracing::instrument(
355 name = "wait",
356 level = "trace",
357 skip_all,
358 fields(
359 receivers = recv.receiver_count(),
360 queued = recv.len(),
361 busy = self.busy.fetch_sub(1, Ordering::AcqRel) - 1,
362 ),
363)]
364fn worker_wait(self: &Arc<Self>, recv: &Receiver<Cmd>) -> Result<Cmd, RecvError> {
365 recv.recv_blocking().debug_inspect(|_| {
366 self.busy.fetch_add(1, Ordering::Relaxed);
367 })
368}
369
370fn worker_handle(cmd: Cmd) {
371 match cmd {
372 | Cmd::Get(cmd) if cmd.key.len() == 1 => handle_get(cmd),
373 | Cmd::Get(cmd) => handle_batch(cmd),
374 | Cmd::Iter(cmd) => handle_iter(cmd),
375 }
376}
377
378#[tracing::instrument(
379 name = "iter",
380 level = "trace",
381 skip_all,
382 fields(%cmd.map),
383)]
384fn handle_iter(mut cmd: Seek) {
385 let chan = cmd.res.take().expect("missing result channel");
386
387 if chan.is_canceled() {
388 return;
389 }
390
391 let from = cmd.key.as_deref();
392
393 let result = match cmd.dir {
394 | Direction::Forward => cmd.state.init_fwd(from),
395 | Direction::Reverse => cmd.state.init_rev(from),
396 };
397
398 let chan_result = chan.send(into_send_seek(result));
399
400 let _chan_sent = chan_result.is_ok();
401}
402
403#[tracing::instrument(
404 name = "batch",
405 level = "trace",
406 skip_all,
407 fields(
408 %cmd.map,
409 keys = %cmd.key.len(),
410 ),
411)]
412fn handle_batch(mut cmd: Get) {
413 debug_assert!(cmd.key.len() > 1, "should have more than one key");
414 debug_assert!(!cmd.key.iter().any(SmallVec::is_empty), "querying for empty key");
415
416 let chan = cmd.res.take().expect("missing result channel");
417
418 if chan.is_canceled() {
419 return;
420 }
421
422 let keys = cmd.key.iter();
423
424 let result: SmallVec<_> = cmd.map.get_batch_blocking(keys).collect();
425
426 let chan_result = chan.send(into_send_get(result));
427
428 let _chan_sent = chan_result.is_ok();
429}
430
431#[tracing::instrument(
432 name = "get",
433 level = "trace",
434 skip_all,
435 fields(%cmd.map),
436)]
437fn handle_get(mut cmd: Get) {
438 debug_assert!(!cmd.key[0].is_empty(), "querying for empty key");
439
440 let chan = cmd.res.take().expect("missing result channel");
442
443 if chan.is_canceled() {
446 return;
447 }
448
449 let result = cmd.map.get_blocking(&cmd.key[0]);
453
454 let chan_result = chan.send(into_send_get([result].into()));
456
457 let _chan_sent = chan_result.is_ok();
459}
460
461fn into_send_get(result: BatchResult<'_>) -> BatchResult<'static> {
462 unsafe { std::mem::transmute(result) }
467}
468
469fn into_recv_get<'a>(result: BatchResult<'static>) -> BatchResult<'a> {
470 unsafe { std::mem::transmute(result) }
472}
473
474pub(crate) fn into_send_seek(result: stream::State<'_>) -> stream::State<'static> {
475 unsafe { std::mem::transmute(result) }
477}
478
479fn into_recv_seek<'a>(result: stream::State<'static>) -> stream::State<'a> {
480 unsafe { std::mem::transmute(result) }
482}