Skip to main content

tuwunel_database/
pool.rs

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
28/// Runs blocking database reads away from asynchronous runtime workers.
29///
30/// Operating-system threads service uncached point queries and iterator seeks
31/// submitted through bounded queues. The pool keeps those blocking calls from
32/// occupying Tokio workers.
33pub(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
42/// Represents work accepted by the database thread pool.
43///
44/// Point queries use [`Get`], while iterator initialization uses [`Seek`]. Each
45/// command carries a response slot populated before it is enqueued.
46pub(crate) enum Cmd {
47	Get(Get),
48	Iter(Seek),
49}
50
51/// Carries a batched point query to a pool worker.
52///
53/// The map and owned keys cross the thread boundary. The optional response
54/// sender is installed immediately before the command is enqueued.
55pub(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
61/// Carries an initial iterator seek to a pool worker.
62///
63/// Only the initial seek is offloaded because RocksDB prefetching is expected
64/// to keep later cursor movements nonblocking. The worker returns the
65/// positioned iterator state through the response sender.
66pub(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
74/// Stores the owned keys in a batched point query.
75///
76/// Small batches remain inline up to the crate's configured batch budget.
77/// Larger batches spill to the heap without changing query order.
78pub(crate) type BatchQuery<'a> = SmallVec<[KeyBuf; BATCH_INLINE]>;
79
80/// Stores the handles returned by a batched point query.
81///
82/// Small result batches remain inline up to the same budget as their queries.
83/// Larger batches spill to the heap.
84pub(crate) type BatchResult<'a> = SmallVec<[ResultHandle<'a>; BATCH_INLINE]>;
85
86/// Represents one point-query result handle.
87///
88/// A successful handle pins the RocksDB value until it is dropped. Its
89/// lifetime remains tied to the database that produced it.
90pub(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/// Constructs the configured database worker pool.
100///
101/// Queue topology and worker counts derive from detected hardware together
102/// with server configuration. Worker groups are spawned before the shared pool
103/// handle is returned.
104#[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/// Selects the queue assigned to the first CPU affinity entry.
263///
264/// The configured topology maps that affinity identifier to a worker group,
265/// falling back to the first queue when the mapped group is absent.
266///
267/// # Panics
268///
269/// Panics if the current thread has no available CPU affinity entry.
270#[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	// affinity is empty (no-op) if there's only one queue
333	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	// initial +1 needed prior to entering wait
346	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	// Obtain the result channel.
441	let chan = cmd.res.take().expect("missing result channel");
442
443	// It is worth checking if the future was dropped while the command was queued
444	// so we can bail without paying for any query.
445	if chan.is_canceled() {
446		return;
447	}
448
449	// Perform the actual database query. We reuse our database::Map interface but
450	// limited to the blocking calls, rather than creating another surface directly
451	// with rocksdb here.
452	let result = cmd.map.get_blocking(&cmd.key[0]);
453
454	// Send the result back to the submitter.
455	let chan_result = chan.send(into_send_get([result].into()));
456
457	// If the future was dropped during the query this will fail acceptably.
458	let _chan_sent = chan_result.is_ok();
459}
460
461fn into_send_get(result: BatchResult<'_>) -> BatchResult<'static> {
462	// SAFETY: Necessary to send the Handle (rust_rocksdb::PinnableSlice) through
463	// the channel. The lifetime on the handle is a device by rust-rocksdb to
464	// associate a database lifetime with its assets. The Handle must be dropped
465	// before the database is dropped.
466	unsafe { std::mem::transmute(result) }
467}
468
469fn into_recv_get<'a>(result: BatchResult<'static>) -> BatchResult<'a> {
470	// SAFETY: This is to receive the Handle from the channel.
471	unsafe { std::mem::transmute(result) }
472}
473
474pub(crate) fn into_send_seek(result: stream::State<'_>) -> stream::State<'static> {
475	// SAFETY: Necessary to send the State through the channel; see above.
476	unsafe { std::mem::transmute(result) }
477}
478
479fn into_recv_seek<'a>(result: stream::State<'static>) -> stream::State<'a> {
480	// SAFETY: This is to receive the State from the channel; see above.
481	unsafe { std::mem::transmute(result) }
482}